diff --git a/src/Simplex/Messaging/Agent.hs b/src/Simplex/Messaging/Agent.hs index 67edbd458..f5747e51c 100644 --- a/src/Simplex/Messaging/Agent.hs +++ b/src/Simplex/Messaging/Agent.hs @@ -99,7 +99,7 @@ import qualified Data.ByteString.Char8 as B import Data.Composition ((.:), (.:.), (.::)) import Data.Foldable (foldl') import Data.Functor (($>)) -import Data.List (deleteFirstsBy, find) +import Data.List (deleteFirstsBy, find, (\\)) import Data.List.NonEmpty (NonEmpty (..), (<|)) import qualified Data.List.NonEmpty as L import Data.Map.Strict (Map) @@ -126,7 +126,7 @@ import Simplex.Messaging.Notifications.Protocol (DeviceToken, NtfRegCode (NtfReg import Simplex.Messaging.Notifications.Server.Push.APNS (PNMessageData (..)) import Simplex.Messaging.Notifications.Types import Simplex.Messaging.Parsers (parse) -import Simplex.Messaging.Protocol (BrokerMsg, ErrorType (AUTH), MsgBody, MsgFlags, NtfServer, SMPMsgMeta, SndPublicVerifyKey, sameSrvAddr, sameSrvAddr') +import Simplex.Messaging.Protocol (BrokerMsg, ErrorType (AUTH), MsgBody, MsgFlags, NtfServer, SMPMsgMeta, SndPublicVerifyKey, protoServer, sameSrvAddr, sameSrvAddr') import qualified Simplex.Messaging.Protocol as SMP import qualified Simplex.Messaging.TMap as TM import Simplex.Messaging.Util @@ -196,7 +196,7 @@ deleteConnectionAsync c = withAgentEnv c .: deleteConnectionAsync' c -- | Create SMP agent connection (NEW command) createConnection :: AgentErrorMonad m => AgentClient -> UserId -> Bool -> SConnectionMode c -> Maybe CRClientData -> m (ConnId, ConnectionRequestUri c) -createConnection c userId enableNtfs cMode clientData = withAgentEnv c $ newConn c userId "" False enableNtfs cMode clientData +createConnection c userId enableNtfs cMode clientData = withAgentEnv c $ newConn c userId "" enableNtfs cMode clientData -- | Join SMP agent connection (JOIN command) joinConnection :: AgentErrorMonad m => AgentClient -> UserId -> Bool -> ConnectionRequestUri c -> ConnInfo -> m ConnId @@ -357,7 +357,7 @@ client c@AgentClient {rcvQ, subQ} = forever $ do -- | execute any SMP agent command processCommand :: forall m. AgentMonad m => AgentClient -> (ConnId, ACommand 'Client) -> m (ConnId, ACommand 'Agent) processCommand c (connId, cmd) = case cmd of - NEW enableNtfs (ACM cMode) -> second (INV . ACR cMode) <$> newConn c userId connId False enableNtfs cMode Nothing + NEW enableNtfs (ACM cMode) -> second (INV . ACR cMode) <$> newConn c userId connId enableNtfs cMode Nothing JOIN enableNtfs (ACR _ cReq) connInfo -> (,OK) <$> joinConn c userId connId False enableNtfs cReq connInfo LET confId ownCInfo -> allowConnection' c connId confId ownCInfo $> (connId, OK) ACPT invId ownCInfo -> (,OK) <$> acceptContact' c connId True invId ownCInfo @@ -470,27 +470,31 @@ switchConnectionAsync' c corrId connId = SomeConn _ DuplexConnection {} -> enqueueCommand c corrId connId Nothing $ AClientCommand SWCH _ -> throwError $ CMD PROHIBITED -newConn :: AgentMonad m => AgentClient -> UserId -> ConnId -> Bool -> Bool -> SConnectionMode c -> Maybe CRClientData -> m (ConnId, ConnectionRequestUri c) -newConn c userId connId asyncMode enableNtfs cMode clientData = - getSMPServer c userId >>= newConnSrv c userId connId asyncMode enableNtfs cMode clientData +newConn :: AgentMonad m => AgentClient -> UserId -> ConnId -> Bool -> SConnectionMode c -> Maybe CRClientData -> m (ConnId, ConnectionRequestUri c) +newConn c userId connId enableNtfs cMode clientData = + getSMPServer c userId >>= newConnSrv c userId connId enableNtfs cMode clientData -newConnSrv :: AgentMonad m => AgentClient -> UserId -> ConnId -> Bool -> Bool -> SConnectionMode c -> Maybe CRClientData -> SMPServerWithAuth -> m (ConnId, ConnectionRequestUri c) -newConnSrv c userId connId asyncMode enableNtfs cMode clientData srv = do +newConnSrv :: AgentMonad m => AgentClient -> UserId -> ConnId -> Bool -> SConnectionMode c -> Maybe CRClientData -> SMPServerWithAuth -> m (ConnId, ConnectionRequestUri c) +newConnSrv c userId connId enableNtfs cMode clientData srv = do + connId' <- newConnNoQueues c userId connId enableNtfs cMode + newRcvConnSrv c userId connId' enableNtfs cMode clientData srv + +newRcvConnSrv :: AgentMonad m => AgentClient -> UserId -> ConnId -> Bool -> SConnectionMode c -> Maybe CRClientData -> SMPServerWithAuth -> m (ConnId, ConnectionRequestUri c) +newRcvConnSrv c userId connId enableNtfs cMode clientData srv = do AgentConfig {smpClientVRange, smpAgentVRange, e2eEncryptVRange} <- asks config - connId' <- if asyncMode then pure connId else newConnNoQueues c userId connId enableNtfs cMode - (rq, qUri) <- newRcvQueue c userId connId' srv smpClientVRange `catchError` \e -> liftIO (print e) >> throwError e - void . withStore c $ \db -> updateNewConnRcv db connId' rq + (rq, qUri) <- newRcvQueue c userId connId srv smpClientVRange `catchError` \e -> liftIO (print e) >> throwError e + void . withStore c $ \db -> updateNewConnRcv db connId rq addSubscription c rq when enableNtfs $ do ns <- asks ntfSupervisor - atomically $ sendNtfSubCommand ns (connId', NSCCreate) + atomically $ sendNtfSubCommand ns (connId, NSCCreate) let crData = ConnReqUriData simplexChat smpAgentVRange [qUri] clientData case cMode of - SCMContact -> pure (connId', CRContactUri crData) + SCMContact -> pure (connId, CRContactUri crData) SCMInvitation -> do (pk1, pk2, e2eRcvParams) <- liftIO . CR.generateE2EParams $ maxVersion e2eEncryptVRange - withStore' c $ \db -> createRatchetX3dhKeys db connId' pk1 pk2 - pure (connId', CRInvitationUri crData $ toVersionRangeT e2eRcvParams e2eEncryptVRange) + withStore' c $ \db -> createRatchetX3dhKeys db connId pk1 pk2 + pure (connId, CRInvitationUri crData $ toVersionRangeT e2eRcvParams e2eEncryptVRange) joinConn :: AgentMonad m => AgentClient -> UserId -> ConnId -> Bool -> Bool -> ConnectionRequestUri c -> ConnInfo -> m ConnId joinConn c userId connId asyncMode enableNtfs cReq cInfo = do @@ -545,7 +549,7 @@ joinConnSrv c userId connId False enableNtfs (CRContactUri ConnReqUriData {crAge crAgentVRange `compatibleVersion` aVRange ) of (Just qInfo, Just vrsn) -> do - (connId', cReq) <- newConnSrv c userId connId False enableNtfs SCMInvitation Nothing srv + (connId', cReq) <- newConnSrv c userId connId enableNtfs SCMInvitation Nothing srv sendInvitation c userId qInfo vrsn cReq cInfo pure connId' _ -> throwError $ AGENT A_VERSION @@ -804,22 +808,20 @@ runCommandProcessing c@AgentClient {subQ} server_ = do atomically $ beginAgentOperation c AOSndNetwork E.try (withStore c $ \db -> getPendingCommand db cmdId) >>= \case Left (e :: E.SomeException) -> atomically $ writeTBQueue subQ ("", "", ERR . INTERNAL $ show e) - Right (corrId, connId, cmd) -> processCmd (riFast ri) corrId connId cmdId cmd + Right cmd -> processCmd (riFast ri) cmdId cmd where - processCmd :: RetryInterval -> ACorrId -> ConnId -> AsyncCmdId -> AgentCommand -> m () - processCmd ri corrId connId cmdId command = case command of + processCmd :: RetryInterval -> AsyncCmdId -> PendingCommand -> m () + processCmd ri cmdId PendingCommand {corrId, userId, connId, command} = case command of AClientCommand cmd -> case cmd of NEW enableNtfs (ACM cMode) -> noServer $ do usedSrvs <- newTVarIO ([] :: [SMPServer]) tryCommand . withNextSrv usedSrvs [] $ \srv -> do - let userId = 1 -- TODO - (_, cReq) <- newConnSrv c userId connId True enableNtfs cMode Nothing srv + (_, cReq) <- newRcvConnSrv c userId connId enableNtfs cMode Nothing srv notify $ INV (ACR cMode cReq) JOIN enableNtfs (ACR _ cReq@(CRInvitationUri ConnReqUriData {crSmpQueues = q :| _} _)) connInfo -> noServer $ do let initUsed = [qServer q] usedSrvs <- newTVarIO initUsed tryCommand . withNextSrv usedSrvs initUsed $ \srv -> do - let userId = 1 -- TODO void $ joinConnSrv c userId connId True enableNtfs cReq connInfo srv notify OK LET confId ownCInfo -> withServer' . tryCommand $ allowConnection' c connId confId ownCInfo >> notify OK @@ -916,12 +918,11 @@ runCommandProcessing c@AgentClient {subQ} server_ = do withNextSrv :: TVar [SMPServer] -> [SMPServer] -> (SMPServerWithAuth -> m ()) -> m () withNextSrv usedSrvs initUsed action = do used <- readTVarIO usedSrvs - let userId = 1 -- TODO srvAuth@(ProtoServerWithAuth srv _) <- getNextSMPServer c userId used atomically $ do srvs_ <- TM.lookup userId $ smpServers c - -- TODO this condition does not account for servers change, it has to be changed to see if there are any remaining unused servers configured for the user - let used' = if length used + 1 >= maybe 0 L.length srvs_ then initUsed else srv : used + let unused = maybe [] ((\\ used) . map protoServer . L.toList) srvs_ + used' = if null unused then initUsed else srv : used writeTVar usedSrvs $! used' action srvAuth -- ^ ^ ^ async command processing / @@ -1505,7 +1506,7 @@ subscriber c@AgentClient {msgQ} = forever $ do Right _ -> return () processSMPTransmission :: forall m. AgentMonad m => AgentClient -> ServerTransmission BrokerMsg -> m () -processSMPTransmission c@AgentClient {smpClients, subQ} (srv, v, sessId, rId, cmd) = do +processSMPTransmission c@AgentClient {smpClients, subQ} (tSess@(_, srv, _), v, sessId, rId, cmd) = do (rq, SomeConn _ conn) <- withStore c (\db -> getRcvConn db srv rId) processSMP rq conn $ connData conn where @@ -1608,9 +1609,7 @@ processSMPTransmission c@AgentClient {smpClients, subQ} (srv, v, sessId, rId, cm ackDel = enqueueCmd . ICAckDel rId srvMsgId handleNotifyAck :: m () -> m () handleNotifyAck m = m `catchError` \e -> notify (ERR e) >> ack - SMP.END -> do - -- TODO is race condition possible here on session mode change? - tSess <- mkSMPTransportSession c rq + SMP.END -> atomically (TM.lookup tSess smpClients $>>= tryReadTMVar >>= processEND) >>= logServer "<--" c srv rId where diff --git a/src/Simplex/Messaging/Agent/Client.hs b/src/Simplex/Messaging/Agent/Client.hs index 6efbb8b0d..641b7e86e 100644 --- a/src/Simplex/Messaging/Agent/Client.hs +++ b/src/Simplex/Messaging/Agent/Client.hs @@ -48,9 +48,6 @@ module Simplex.Messaging.Agent.Client agentNtfCreateSubscription, agentNtfCheckSubscription, agentNtfDeleteSubscription, - -- TODO possibly, this should not be exported? - -- this is used in END processing - mkSMPTransportSession, agentCbEncrypt, agentCbDecrypt, cryptoError, @@ -170,9 +167,6 @@ type SMPClientVar = TMVar (Either AgentErrorType SMPClient) type NtfClientVar = TMVar (Either AgentErrorType NtfClient) --- | Transport session key - includes entity ID if `sessionMode = TSMEntity`. -type TransportSession msg = (UserId, ProtoServer msg, Maybe EntityId) - type SMPTransportSession = TransportSession SMP.BrokerMsg type NtfTransportSession = TransportSession NtfResponse @@ -299,7 +293,7 @@ instance ProtocolServerClient NtfResponse where clientProtocolError = NTF getSMPServerClient :: forall m. AgentMonad m => AgentClient -> SMPTransportSession -> m SMPClient -getSMPServerClient c@AgentClient {active, smpClients, msgQ} tSess@(userId, srv, entityId_) = do +getSMPServerClient c@AgentClient {active, smpClients, msgQ} tSess@(userId, srv, _) = do unlessM (readTVarIO active) . throwError $ INTERNAL "agent is stopped" atomically (getClientVar tSess smpClients) >>= either @@ -310,8 +304,7 @@ getSMPServerClient c@AgentClient {active, smpClients, msgQ} tSess@(userId, srv, connectClient = do cfg <- getClientConfig c smpCfg u <- askUnliftIO - let proxyUsername = C.sha256Hash $ bshow userId <> maybe "" (":" <>) entityId_ - liftEitherError (protocolClientError SMP $ B.unpack $ strEncode srv) (getProtocolClient srv cfg (Just proxyUsername) (Just msgQ) $ clientDisconnected u) + liftEitherError (protocolClientError SMP $ B.unpack $ strEncode srv) (getProtocolClient tSess cfg (Just msgQ) $ clientDisconnected u) clientDisconnected :: UnliftIO m -> SMPClient -> IO () clientDisconnected u client = do @@ -376,7 +369,7 @@ getNtfServerClient c@AgentClient {active, ntfClients} tSess@(userId, srv, _) = d connectClient :: m NtfClient connectClient = do cfg <- getClientConfig c ntfCfg - liftEitherError (protocolClientError NTF $ B.unpack $ strEncode srv) (getProtocolClient srv cfg Nothing Nothing clientDisconnected) + liftEitherError (protocolClientError NTF $ B.unpack $ strEncode srv) (getProtocolClient tSess cfg Nothing clientDisconnected) clientDisconnected :: NtfClient -> IO () clientDisconnected client = do @@ -583,8 +576,8 @@ runSMPServerTest c userId (ProtoServerWithAuth srv auth) = do cfg <- getClientConfig c smpCfg C.SignAlg a <- asks $ cmdSignAlg . config liftIO $ do - let proxyUsername = C.sha256Hash $ bshow userId - getProtocolClient srv cfg (Just proxyUsername) Nothing (\_ -> pure ()) >>= \case + let tSess = (userId, srv, Nothing) + getProtocolClient tSess cfg Nothing (\_ -> pure ()) >>= \case Right smp -> do (rKey, rpKey) <- C.generateSignatureKeyPair a (sKey, _) <- C.generateSignatureKeyPair a diff --git a/src/Simplex/Messaging/Agent/Store.hs b/src/Simplex/Messaging/Agent/Store.hs index cff4ca125..7f9f8bb90 100644 --- a/src/Simplex/Messaging/Agent/Store.hs +++ b/src/Simplex/Messaging/Agent/Store.hs @@ -254,6 +254,13 @@ data ConnData = ConnData } deriving (Eq, Show) +data PendingCommand = PendingCommand + { corrId :: ACorrId, + userId :: UserId, + connId :: ConnId, + command :: AgentCommand + } + data AgentCmdType = ACClient | ACInternal type UserId = Int64 diff --git a/src/Simplex/Messaging/Agent/Store/SQLite.hs b/src/Simplex/Messaging/Agent/Store/SQLite.hs index dfdc73103..59e6b41e2 100644 --- a/src/Simplex/Messaging/Agent/Store/SQLite.hs +++ b/src/Simplex/Messaging/Agent/Store/SQLite.hs @@ -836,13 +836,20 @@ getPendingCommands db connId = do where srvCmdId (host, port, keyHash, cmdId) = (SMPServer <$> host <*> port <*> keyHash, cmdId) -getPendingCommand :: DB.Connection -> AsyncCmdId -> IO (Either StoreError (ACorrId, ConnId, AgentCommand)) +getPendingCommand :: DB.Connection -> AsyncCmdId -> IO (Either StoreError PendingCommand) getPendingCommand db msgId = do - firstRow id SECmdNotFound $ + firstRow pendingCommand SECmdNotFound $ DB.query db - "SELECT corr_id, conn_id, command FROM commands WHERE command_id = ?" + [sql| + SELECT c.corr_id, cs.user_id, c.conn_id, c.command + FROM commands c + JOIN connections cs USING (conn_id) + WHERE c.command_id = ? + |] (Only msgId) + where + pendingCommand (corrId, userId, connId, command) = PendingCommand {corrId, userId, connId, command} deleteCommand :: DB.Connection -> AsyncCmdId -> IO () deleteCommand db cmdId = diff --git a/src/Simplex/Messaging/Agent/Store/SQLite/Migrations/M20230110_users.hs b/src/Simplex/Messaging/Agent/Store/SQLite/Migrations/M20230110_users.hs index a56930ad0..73758cdeb 100644 --- a/src/Simplex/Messaging/Agent/Store/SQLite/Migrations/M20230110_users.hs +++ b/src/Simplex/Messaging/Agent/Store/SQLite/Migrations/M20230110_users.hs @@ -21,6 +21,8 @@ REFERENCES users ON DELETE CASCADE; CREATE INDEX idx_connections_user ON connections(user_id); +CREATE INDEX idx_commands_conn_id ON commands(conn_id); + UPDATE connections SET user_id = 1; PRAGMA ignore_check_constraints=OFF; diff --git a/src/Simplex/Messaging/Agent/Store/SQLite/Migrations/agent_schema.sql b/src/Simplex/Messaging/Agent/Store/SQLite/Migrations/agent_schema.sql index 548aaa63c..3121dd736 100644 --- a/src/Simplex/Messaging/Agent/Store/SQLite/Migrations/agent_schema.sql +++ b/src/Simplex/Messaging/Agent/Store/SQLite/Migrations/agent_schema.sql @@ -232,3 +232,4 @@ CREATE INDEX idx_snd_message_deliveries ON snd_message_deliveries( ); CREATE TABLE users(user_id INTEGER PRIMARY KEY AUTOINCREMENT); CREATE INDEX idx_connections_user ON connections(user_id); +CREATE INDEX idx_commands_conn_id ON commands(conn_id); diff --git a/src/Simplex/Messaging/Client.hs b/src/Simplex/Messaging/Client.hs index dcf1e3e10..5f67bcc95 100644 --- a/src/Simplex/Messaging/Client.hs +++ b/src/Simplex/Messaging/Client.hs @@ -26,6 +26,7 @@ -- See https://github.com/simplex-chat/simplexmq/blob/master/protocol/simplex-messaging.md module Simplex.Messaging.Client ( -- * Connect (disconnect) client to (from) SMP server + TransportSession, ProtocolClient (thVersion, sessionId, sessionTs), SMPClient, getProtocolClient, @@ -76,6 +77,7 @@ import Data.ByteString.Char8 (ByteString) import qualified Data.ByteString.Char8 as B import Data.Either (rights) import Data.Functor (($>)) +import Data.Int (Int64) import Data.List (find) import Data.List.NonEmpty (NonEmpty) import qualified Data.List.NonEmpty as L @@ -111,7 +113,7 @@ data ProtocolClient msg = ProtocolClient data PClient msg = PClient { connected :: TVar Bool, - protocolServer :: ProtoServer msg, + transportSession :: TransportSession msg, transportHost :: TransportHost, tcpTimeout :: Int, clientCorrId :: TVar Natural, @@ -127,7 +129,7 @@ type SMPClient = ProtocolClient SMP.BrokerMsg type ClientCommand msg = (Maybe C.APrivateSignKey, QueueId, ProtoCommand msg) -- | Type synonym for transmission from some SPM server queue. -type ServerTransmission msg = (ProtoServer msg, Version, SessionId, QueueId, msg) +type ServerTransmission msg = (TransportSession msg, Version, SessionId, QueueId, msg) data HostMode = -- | prefer (or require) onion hosts when connecting via SOCKS proxy @@ -243,19 +245,26 @@ chooseTransportHost NetworkConfig {socksProxy, hostMode, requiredHostMode} hosts publicHost = find (not . isOnionHost) hosts clientServer :: ProtocolTypeI (ProtoType msg) => ProtocolClient msg -> String -clientServer = B.unpack . strEncode . protocolServer . client_ +clientServer = B.unpack . strEncode . snd3 . transportSession . client_ + where + snd3 (_, s, _) = s transportHost' :: ProtocolClient msg -> TransportHost transportHost' = transportHost . client_ +type UserId = Int64 + +-- | Transport session key - includes entity ID if `sessionMode = TSMEntity`. +type TransportSession msg = (UserId, ProtoServer msg, Maybe EntityId) + -- | Connects to 'ProtocolServer' using passed client configuration -- and queue for messages and notifications. -- -- A single queue can be used for multiple 'SMPClient' instances, -- as 'SMPServerTransmission' includes server information. -getProtocolClient :: forall msg. Protocol msg => ProtoServer msg -> ProtocolClientConfig -> Maybe ByteString -> Maybe (TBQueue (ServerTransmission msg)) -> (ProtocolClient msg -> IO ()) -> IO (Either ProtocolClientError (ProtocolClient msg)) -getProtocolClient protocolServer cfg@ProtocolClientConfig {qSize, networkConfig, smpServerVRange} proxyUsername msgQ disconnected = do - case chooseTransportHost networkConfig (host protocolServer) of +getProtocolClient :: forall msg. Protocol msg => TransportSession msg -> ProtocolClientConfig -> Maybe (TBQueue (ServerTransmission msg)) -> (ProtocolClient msg -> IO ()) -> IO (Either ProtocolClientError (ProtocolClient msg)) +getProtocolClient transportSession@(_, srv, _) cfg@ProtocolClientConfig {qSize, networkConfig, smpServerVRange} msgQ disconnected = do + case chooseTransportHost networkConfig (host srv) of Right useHost -> (atomically (mkProtocolClient useHost) >>= runClient useTransport useHost) `catch` \(e :: IOException) -> pure . Left $ PCEIOError e @@ -272,7 +281,7 @@ getProtocolClient protocolServer cfg@ProtocolClientConfig {qSize, networkConfig, return PClient { connected, - protocolServer, + transportSession, transportHost, tcpTimeout, clientCorrId, @@ -286,9 +295,10 @@ getProtocolClient protocolServer cfg@ProtocolClientConfig {qSize, networkConfig, runClient (port', ATransport t) useHost c = do cVar <- newEmptyTMVarIO let tcConfig = transportClientConfig networkConfig + username = proxyUsername transportSession action <- async $ - runTransportClient tcConfig proxyUsername useHost port' (Just $ keyHash protocolServer) (client t c cVar) + runTransportClient tcConfig (Just username) useHost port' (Just $ keyHash srv) (client t c cVar) `finally` atomically (putTMVar cVar $ Left PCENetworkError) c_ <- tcpConnectTimeout `timeout` atomically (takeTMVar cVar) pure $ case c_ of @@ -296,15 +306,18 @@ getProtocolClient protocolServer cfg@ProtocolClientConfig {qSize, networkConfig, Just (Left e) -> Left e Nothing -> Left PCENetworkError + proxyUsername :: TransportSession msg -> ByteString + proxyUsername (userId, _, entityId_) = C.sha256Hash $ bshow userId <> maybe "" (":" <>) entityId_ + useTransport :: (ServiceName, ATransport) - useTransport = case port protocolServer of + useTransport = case port srv of "" -> defaultTransport cfg "80" -> ("80", transport @WS) p -> (p, transport @TLS) client :: forall c. Transport c => TProxy c -> PClient msg -> TMVar (Either ProtocolClientError (ProtocolClient msg)) -> c -> IO () client _ c cVar h = - runExceptT (protocolClientHandshake @msg h (keyHash protocolServer) smpServerVRange) >>= \case + runExceptT (protocolClientHandshake @msg h (keyHash srv) smpServerVRange) >>= \case Left e -> atomically . putTMVar cVar . Left $ PCETransportError e Right th@THandle {sessionId, thVersion} -> do sessionTs <- getCurrentTime @@ -426,8 +439,8 @@ writeSMPMessage :: SMPClient -> RecipientId -> BrokerMsg -> IO () writeSMPMessage c rId msg = atomically $ mapM_ (`writeTBQueue` serverTransmission c rId msg) (msgQ $ client_ c) serverTransmission :: ProtocolClient msg -> RecipientId -> msg -> ServerTransmission msg -serverTransmission ProtocolClient {thVersion, sessionId, client_ = PClient {protocolServer}} entityId message = - (protocolServer, thVersion, sessionId, entityId, message) +serverTransmission ProtocolClient {thVersion, sessionId, client_ = PClient {transportSession}} entityId message = + (transportSession, thVersion, sessionId, entityId, message) -- | Get message from SMP queue. The server returns ERR PROHIBITED if a client uses SUB and GET via the same transport connection for the same queue -- diff --git a/src/Simplex/Messaging/Client/Agent.hs b/src/Simplex/Messaging/Client/Agent.hs index f2fe97666..feefacecd 100644 --- a/src/Simplex/Messaging/Client/Agent.hs +++ b/src/Simplex/Messaging/Client/Agent.hs @@ -160,7 +160,7 @@ getSMPServerClient' ca@SMPClientAgent {agentCfg, smpClients, msgQ} srv = void $ tryConnectClient (const reconnectClient) loop connectClient :: ExceptT ProtocolClientError IO SMPClient - connectClient = ExceptT $ getProtocolClient srv (smpCfg agentCfg) Nothing (Just msgQ) clientDisconnected + connectClient = ExceptT $ getProtocolClient (1, srv, Nothing) (smpCfg agentCfg) (Just msgQ) clientDisconnected clientDisconnected :: SMPClient -> IO () clientDisconnected _ = do diff --git a/src/Simplex/Messaging/Notifications/Server.hs b/src/Simplex/Messaging/Notifications/Server.hs index 80069a534..02f07da00 100644 --- a/src/Simplex/Messaging/Notifications/Server.hs +++ b/src/Simplex/Messaging/Notifications/Server.hs @@ -190,7 +190,7 @@ ntfSubscriber NtfSubscriber {smpSubscribers, newSubQ, smpAgent = ca@SMPClientAge receiveSMP :: M () receiveSMP = forever $ do - (srv, _, _, ntfId, msg) <- atomically $ readTBQueue msgQ + ((_, srv, _), _, _, ntfId, msg) <- atomically $ readTBQueue msgQ let smpQueue = SMPQueueNtf srv ntfId case msg of SMP.NMSG nmsgNonce encNMsgMeta -> do diff --git a/src/Simplex/Messaging/Protocol.hs b/src/Simplex/Messaging/Protocol.hs index a6897074d..59643f6b4 100644 --- a/src/Simplex/Messaging/Protocol.hs +++ b/src/Simplex/Messaging/Protocol.hs @@ -746,7 +746,7 @@ basicAuth s where valid c = isPrint c && not (isSpace c) && c /= '@' && c /= ':' && c /= '/' -data ProtoServerWithAuth p = ProtoServerWithAuth (ProtocolServer p) (Maybe BasicAuth) +data ProtoServerWithAuth p = ProtoServerWithAuth {protoServer :: ProtocolServer p, serverBasicAuth :: Maybe BasicAuth} deriving (Show) instance ProtocolTypeI p => IsString (ProtoServerWithAuth p) where diff --git a/tests/AgentTests.hs b/tests/AgentTests.hs index 3c6bfa696..705bf4550 100644 --- a/tests/AgentTests.hs +++ b/tests/AgentTests.hs @@ -338,8 +338,8 @@ testServerConnectionAfterError t _ = do withServer $ do alice <#= \case ("", "bob", SENT 4) -> True; ("", "", UP s ["bob"]) -> s == server; _ -> False alice <#= \case ("", "bob", SENT 4) -> True; ("", "", UP s ["bob"]) -> s == server; _ -> False - bob <# ("", "", UP server ["alice"]) - bob <#= \case ("", "alice", Msg "hello") -> True; _ -> False + bob <#= \case ("", "alice", Msg "hello") -> True; ("", "", UP s ["alice"]) -> s == server; _ -> False + bob <#= \case ("", "alice", Msg "hello") -> True; ("", "", UP s ["alice"]) -> s == server; _ -> False bob #: ("2", "alice", "ACK 4") #> ("2", "alice", OK) alice #: ("1", "bob", "SEND F 11\nhello again") #> ("1", "bob", MID 5) alice <# ("", "bob", SENT 5) diff --git a/tests/AgentTests/FunctionalAPITests.hs b/tests/AgentTests/FunctionalAPITests.hs index 989d65230..6109b85f4 100644 --- a/tests/AgentTests/FunctionalAPITests.hs +++ b/tests/AgentTests/FunctionalAPITests.hs @@ -156,7 +156,7 @@ functionalAPITests t = do it "JOIN fail, no auth " $ testBasicAuth t True (Just "abcd", 5) (Just "abcd", 5) (Nothing, 5) `shouldReturn` 1 it "JOIN fail, bad auth " $ testBasicAuth t True (Just "abcd", 5) (Just "abcd", 5) (Just "wrong", 5) `shouldReturn` 1 it "JOIN fail, version " $ testBasicAuth t True (Just "abcd", 5) (Just "abcd", 5) (Just "abcd", 4) `shouldReturn` 1 - describe "no server auth" $ do + fdescribe "no server auth" $ do it "success " $ testBasicAuth t True (Nothing, 5) (Nothing, 5) (Nothing, 5) `shouldReturn` 2 it "srv disabled" $ testBasicAuth t False (Nothing, 5) (Nothing, 5) (Nothing, 5) `shouldReturn` 0 it "version srv " $ testBasicAuth t True (Nothing, 4) (Nothing, 5) (Nothing, 5) `shouldReturn` 2