diff --git a/apps/smp-agent/Main.hs b/apps/smp-agent/Main.hs index b2bb22e9d..1087530f9 100644 --- a/apps/smp-agent/Main.hs +++ b/apps/smp-agent/Main.hs @@ -6,8 +6,8 @@ module Main where import Control.Logger.Simple import qualified Data.List.NonEmpty as L -import Simplex.Messaging.Agent (runSMPAgent) import Simplex.Messaging.Agent.Env.SQLite +import Simplex.Messaging.Agent.Server (runSMPAgent) import Simplex.Messaging.Transport (TLS, Transport (..)) cfg :: AgentConfig diff --git a/apps/smp-server/Main.hs b/apps/smp-server/Main.hs index 97d386d27..fbc1356c1 100644 --- a/apps/smp-server/Main.hs +++ b/apps/smp-server/Main.hs @@ -22,12 +22,13 @@ import Simplex.Messaging.Encoding.String import Simplex.Messaging.Server (runSMPServer) import Simplex.Messaging.Server.Env.STM import Simplex.Messaging.Server.StoreLog (StoreLog, openReadStoreLog, storeLogFilePath) -import Simplex.Messaging.Transport (ATransport (..), TLS, Transport (..), loadFingerprint, simplexMQVersion) +import Simplex.Messaging.Transport (ATransport (..), TLS, Transport (..), simplexMQVersion) +import Simplex.Messaging.Transport.Server (loadFingerprint) import Simplex.Messaging.Transport.WebSockets (WS) import System.Directory (createDirectoryIfMissing, doesDirectoryExist, doesFileExist, removeDirectoryRecursive) import System.Exit (exitFailure) import System.FilePath (combine) -import System.IO (BufferMode (..), IOMode (..), hGetLine, withFile, hSetBuffering, stderr, stdout) +import System.IO (BufferMode (..), IOMode (..), hGetLine, hSetBuffering, stderr, stdout, withFile) import System.Process (readCreateProcess, shell) import Text.Read (readMaybe) diff --git a/simplexmq.cabal b/simplexmq.cabal index 3fb4de0ba..f99491e46 100644 --- a/simplexmq.cabal +++ b/simplexmq.cabal @@ -35,6 +35,7 @@ library Simplex.Messaging.Agent.Protocol Simplex.Messaging.Agent.QueryString Simplex.Messaging.Agent.RetryInterval + Simplex.Messaging.Agent.Server Simplex.Messaging.Agent.Store Simplex.Messaging.Agent.Store.SQLite Simplex.Messaging.Agent.Store.SQLite.Migrations @@ -54,6 +55,8 @@ library Simplex.Messaging.Server.QueueStore.STM Simplex.Messaging.Server.StoreLog Simplex.Messaging.Transport + Simplex.Messaging.Transport.Client + Simplex.Messaging.Transport.Server Simplex.Messaging.Transport.WebSockets Simplex.Messaging.Util Simplex.Messaging.Version diff --git a/src/Simplex/Messaging/Agent.hs b/src/Simplex/Messaging/Agent.hs index f79f1f68e..41d9a073f 100644 --- a/src/Simplex/Messaging/Agent.hs +++ b/src/Simplex/Messaging/Agent.hs @@ -26,11 +26,7 @@ -- -- See https://github.com/simplex-chat/simplexmq/blob/master/protocol/agent-protocol.md module Simplex.Messaging.Agent - ( -- * SMP agent over TCP - runSMPAgent, - runSMPAgentBlocking, - - -- * queue-based SMP agent + ( -- * queue-based SMP agent getAgentClient, runAgentClient, @@ -51,6 +47,7 @@ module Simplex.Messaging.Agent ackMessage, suspendConnection, deleteConnection, + logConnection, ) where @@ -62,7 +59,6 @@ import Control.Monad.Reader import Crypto.Random (MonadRandom) import Data.Bifunctor (first, second) import Data.ByteString.Char8 (ByteString) -import qualified Data.ByteString.Char8 as B import Data.Composition ((.:), (.:.)) import Data.Functor (($>)) import Data.List.NonEmpty (NonEmpty (..)) @@ -70,7 +66,6 @@ import qualified Data.List.NonEmpty as L import qualified Data.Map.Strict as M import Data.Maybe (isJust) import qualified Data.Text as T -import Data.Text.Encoding (decodeUtf8) import Data.Time.Clock import Data.Time.Clock.System (systemToUTCTime) import Database.SQLite.Simple (SQLError) @@ -87,7 +82,6 @@ import Simplex.Messaging.Encoding import Simplex.Messaging.Parsers (parse) import Simplex.Messaging.Protocol (MsgBody) import qualified Simplex.Messaging.Protocol as SMP -import Simplex.Messaging.Transport (ATransport (..), TProxy, Transport (..), loadTLSServerParams, runTransportServer, simplexMQVersion) import Simplex.Messaging.Util (bshow, liftError, tryError, unlessM) import Simplex.Messaging.Version import System.Random (randomR) @@ -95,33 +89,6 @@ import UnliftIO.Async (async, race_) import qualified UnliftIO.Exception as E import UnliftIO.STM --- | Runs an SMP agent as a TCP service using passed configuration. --- --- See a full agent executable here: https://github.com/simplex-chat/simplexmq/blob/master/apps/smp-agent/Main.hs -runSMPAgent :: (MonadRandom m, MonadUnliftIO m) => ATransport -> AgentConfig -> m () -runSMPAgent t cfg = do - started <- newEmptyTMVarIO - runSMPAgentBlocking t started cfg - --- | Runs an SMP agent as a TCP service using passed configuration with signalling. --- --- This function uses passed TMVar to signal when the server is ready to accept TCP requests (True) --- and when it is disconnected from the TCP socket once the server thread is killed (False). -runSMPAgentBlocking :: (MonadRandom m, MonadUnliftIO m) => ATransport -> TMVar Bool -> AgentConfig -> m () -runSMPAgentBlocking (ATransport t) started cfg@AgentConfig {tcpPort, caCertificateFile, certificateFile, privateKeyFile} = do - runReaderT (smpAgent t) =<< newSMPAgentEnv cfg - where - smpAgent :: forall c m'. (Transport c, MonadUnliftIO m', MonadReader Env m') => TProxy c -> m' () - smpAgent _ = do - -- tlsServerParams is not in Env to avoid breaking functional API w/t key and certificate generation - tlsServerParams <- liftIO $ loadTLSServerParams caCertificateFile certificateFile privateKeyFile - runTransportServer started tcpPort tlsServerParams $ \(h :: c) -> do - liftIO . putLn h $ "Welcome to SMP agent v" <> B.pack simplexMQVersion - c <- getAgentClient - logConnection c True - race_ (connectClient h c) (runAgentClient c) - `E.finally` disconnectAgentClient c - -- | Creates an SMP agent client instance getSMPAgentClient :: (MonadRandom m, MonadUnliftIO m) => AgentConfig -> m AgentClient getSMPAgentClient cfg = newSMPAgentEnv cfg >>= runReaderT runAgent @@ -186,9 +153,6 @@ withAgentEnv c = (`runReaderT` agentEnv c) getAgentClient :: (MonadUnliftIO m, MonadReader Env m) => m AgentClient getAgentClient = ask >>= atomically . newAgentClient -connectClient :: Transport c => MonadUnliftIO m => c -> AgentClient -> m () -connectClient h c = race_ (send h c) (receive h c) - logConnection :: MonadUnliftIO m => AgentClient -> Bool -> m () logConnection c connected = let event = if connected then "connected to" else "disconnected from" @@ -198,28 +162,6 @@ logConnection c connected = runAgentClient :: (MonadUnliftIO m, MonadReader Env m) => AgentClient -> m () runAgentClient c = race_ (subscriber c) (client c) -receive :: forall c m. (Transport c, MonadUnliftIO m) => c -> AgentClient -> m () -receive h c@AgentClient {rcvQ, subQ} = forever $ do - (corrId, connId, cmdOrErr) <- tGet SClient h - case cmdOrErr of - Right cmd -> write rcvQ (corrId, connId, cmd) - Left e -> write subQ (corrId, connId, ERR e) - where - write :: TBQueue (ATransmission p) -> ATransmission p -> m () - write q t = do - logClient c "-->" t - atomically $ writeTBQueue q t - -send :: (Transport c, MonadUnliftIO m) => c -> AgentClient -> m () -send h c@AgentClient {subQ} = forever $ do - t <- atomically $ readTBQueue subQ - tPut h t - logClient c "<--" t - -logClient :: MonadUnliftIO m => AgentClient -> ByteString -> ATransmission a -> m () -logClient AgentClient {clientId} dir (corrId, connId, cmd) = do - logInfo . decodeUtf8 $ B.unwords [bshow clientId, dir, "A :", corrId, connId, B.takeWhile (/= ' ') $ serializeCommand cmd] - client :: forall m. (MonadUnliftIO m, MonadReader Env m) => AgentClient -> m () client c@AgentClient {rcvQ, subQ} = forever $ do (corrId, connId, cmd) <- atomically $ readTBQueue rcvQ diff --git a/src/Simplex/Messaging/Agent/Server.hs b/src/Simplex/Messaging/Agent/Server.hs new file mode 100644 index 000000000..b3c70c9ac --- /dev/null +++ b/src/Simplex/Messaging/Agent/Server.hs @@ -0,0 +1,81 @@ +{-# LANGUAGE FlexibleContexts #-} +{-# LANGUAGE NamedFieldPuns #-} +{-# LANGUAGE OverloadedStrings #-} +{-# LANGUAGE ScopedTypeVariables #-} + +module Simplex.Messaging.Agent.Server + ( -- * SMP agent over TCP + runSMPAgent, + runSMPAgentBlocking, + ) +where + +import Control.Logger.Simple (logInfo) +import Control.Monad.Except +import Control.Monad.IO.Unlift (MonadUnliftIO) +import Control.Monad.Reader +import Crypto.Random (MonadRandom) +import Data.ByteString.Char8 (ByteString) +import qualified Data.ByteString.Char8 as B +import Data.Text.Encoding (decodeUtf8) +import Simplex.Messaging.Agent +import Simplex.Messaging.Agent.Env.SQLite +import Simplex.Messaging.Agent.Protocol +import Simplex.Messaging.Transport (ATransport (..), TProxy, Transport (..), simplexMQVersion) +import Simplex.Messaging.Transport.Server (loadTLSServerParams, runTransportServer) +import Simplex.Messaging.Util (bshow) +import UnliftIO.Async (race_) +import qualified UnliftIO.Exception as E +import UnliftIO.STM + +-- | Runs an SMP agent as a TCP service using passed configuration. +-- +-- See a full agent executable here: https://github.com/simplex-chat/simplexmq/blob/master/apps/smp-agent/Main.hs +runSMPAgent :: (MonadRandom m, MonadUnliftIO m) => ATransport -> AgentConfig -> m () +runSMPAgent t cfg = do + started <- newEmptyTMVarIO + runSMPAgentBlocking t started cfg + +-- | Runs an SMP agent as a TCP service using passed configuration with signalling. +-- +-- This function uses passed TMVar to signal when the server is ready to accept TCP requests (True) +-- and when it is disconnected from the TCP socket once the server thread is killed (False). +runSMPAgentBlocking :: (MonadRandom m, MonadUnliftIO m) => ATransport -> TMVar Bool -> AgentConfig -> m () +runSMPAgentBlocking (ATransport t) started cfg@AgentConfig {tcpPort, caCertificateFile, certificateFile, privateKeyFile} = do + runReaderT (smpAgent t) =<< newSMPAgentEnv cfg + where + smpAgent :: forall c m'. (Transport c, MonadUnliftIO m', MonadReader Env m') => TProxy c -> m' () + smpAgent _ = do + -- tlsServerParams is not in Env to avoid breaking functional API w/t key and certificate generation + tlsServerParams <- liftIO $ loadTLSServerParams caCertificateFile certificateFile privateKeyFile + runTransportServer started tcpPort tlsServerParams $ \(h :: c) -> do + liftIO . putLn h $ "Welcome to SMP agent v" <> B.pack simplexMQVersion + c <- getAgentClient + logConnection c True + race_ (connectClient h c) (runAgentClient c) + `E.finally` disconnectAgentClient c + +connectClient :: Transport c => MonadUnliftIO m => c -> AgentClient -> m () +connectClient h c = race_ (send h c) (receive h c) + +receive :: forall c m. (Transport c, MonadUnliftIO m) => c -> AgentClient -> m () +receive h c@AgentClient {rcvQ, subQ} = forever $ do + (corrId, connId, cmdOrErr) <- tGet SClient h + case cmdOrErr of + Right cmd -> write rcvQ (corrId, connId, cmd) + Left e -> write subQ (corrId, connId, ERR e) + where + write :: TBQueue (ATransmission p) -> ATransmission p -> m () + write q t = do + logClient c "-->" t + atomically $ writeTBQueue q t + +send :: (Transport c, MonadUnliftIO m) => c -> AgentClient -> m () +send h c@AgentClient {subQ} = forever $ do + t <- atomically $ readTBQueue subQ + tPut h t + logClient c "<--" t + +logClient :: MonadUnliftIO m => AgentClient -> ByteString -> ATransmission a -> m () +logClient AgentClient {clientId} dir (corrId, connId, cmd) = do + logInfo . decodeUtf8 $ B.unwords [bshow clientId, dir, "A :", corrId, connId, B.takeWhile (/= ' ') $ serializeCommand cmd] diff --git a/src/Simplex/Messaging/Client.hs b/src/Simplex/Messaging/Client.hs index 92eaf120a..962efa323 100644 --- a/src/Simplex/Messaging/Client.hs +++ b/src/Simplex/Messaging/Client.hs @@ -63,7 +63,8 @@ import Network.Socket (ServiceName) import Numeric.Natural import qualified Simplex.Messaging.Crypto as C import Simplex.Messaging.Protocol -import Simplex.Messaging.Transport (ATransport (..), THandle (..), TLS, TProxy, Transport (..), TransportError, clientHandshake, runTransportClient) +import Simplex.Messaging.Transport (ATransport (..), THandle (..), TLS, TProxy, Transport (..), TransportError, clientHandshake) +import Simplex.Messaging.Transport.Client (runTransportClient) import Simplex.Messaging.Transport.WebSockets (WS) import Simplex.Messaging.Util (bshow, liftError, raceAny_) import System.Timeout (timeout) diff --git a/src/Simplex/Messaging/Server.hs b/src/Simplex/Messaging/Server.hs index 3ae73aa7f..73b5780ef 100644 --- a/src/Simplex/Messaging/Server.hs +++ b/src/Simplex/Messaging/Server.hs @@ -48,6 +48,7 @@ import Simplex.Messaging.Server.QueueStore import Simplex.Messaging.Server.QueueStore.STM (QueueStore) import Simplex.Messaging.Server.StoreLog import Simplex.Messaging.Transport +import Simplex.Messaging.Transport.Server import Simplex.Messaging.Util import UnliftIO.Concurrent import UnliftIO.Exception diff --git a/src/Simplex/Messaging/Server/Env/STM.hs b/src/Simplex/Messaging/Server/Env/STM.hs index 689096cac..d1b4c51fc 100644 --- a/src/Simplex/Messaging/Server/Env/STM.hs +++ b/src/Simplex/Messaging/Server/Env/STM.hs @@ -21,7 +21,8 @@ import Simplex.Messaging.Server.MsgStore.STM import Simplex.Messaging.Server.QueueStore (QueueRec (..)) import Simplex.Messaging.Server.QueueStore.STM import Simplex.Messaging.Server.StoreLog -import Simplex.Messaging.Transport (ATransport, loadFingerprint, loadTLSServerParams) +import Simplex.Messaging.Transport (ATransport) +import Simplex.Messaging.Transport.Server (loadFingerprint, loadTLSServerParams) import System.IO (IOMode (..)) import UnliftIO.STM diff --git a/src/Simplex/Messaging/Transport.hs b/src/Simplex/Messaging/Transport.hs index 8a0c66d2e..8cc1b7e57 100644 --- a/src/Simplex/Messaging/Transport.hs +++ b/src/Simplex/Messaging/Transport.hs @@ -36,15 +36,11 @@ module Simplex.Messaging.Transport ATransport (..), TransportPeer (..), - -- * Transport over TLS - runTransportServer, - runTransportClient, - loadTLSServerParams, - loadFingerprint, - -- * TLS Transport TLS (..), + connectTLS, closeTLS, + supportedParameters, withTlsUnique, -- * SMP transport @@ -64,9 +60,7 @@ where import Control.Applicative ((<|>)) import Control.Monad.Except -import Control.Monad.IO.Unlift import Control.Monad.Trans.Except (throwE) -import qualified Crypto.Store.X509 as SX import Data.Attoparsec.ByteString.Char8 (Parser) import Data.Bifunctor (first) import Data.Bitraversable (bimapM) @@ -75,14 +69,7 @@ import qualified Data.ByteString.Char8 as B import qualified Data.ByteString.Lazy as BL import Data.Default (def) import Data.Functor (($>)) -import Data.Set (Set) -import qualified Data.Set as S -import qualified Data.X509 as X -import qualified Data.X509.CertificateStore as XS -import Data.X509.Validation (Fingerprint (..)) -import qualified Data.X509.Validation as XV import GHC.Generics (Generic) -import GHC.IO.Exception (IOErrorType (..)) import GHC.IO.Handle.Internals (ioe_EOF) import Generic.Random (genericArbitraryU) import Network.Socket @@ -93,11 +80,8 @@ import Simplex.Messaging.Encoding import Simplex.Messaging.Parsers (parse, parseRead1) import Simplex.Messaging.Util (bshow) import Simplex.Messaging.Version -import System.Exit (exitFailure) -import System.IO.Error import Test.QuickCheck (Arbitrary (..)) -import UnliftIO.Concurrent -import UnliftIO.Exception (Exception, IOException) +import UnliftIO.Exception (Exception) import qualified UnliftIO.Exception as E import UnliftIO.STM @@ -154,104 +138,6 @@ data TProxy c = TProxy data ATransport = forall c. Transport c => ATransport (TProxy c) --- * Transport over TLS - --- | Run transport server (plain TCP or WebSockets) on passed TCP port and signal when server started and stopped via passed TMVar. --- --- All accepted connections are passed to the passed function. -runTransportServer :: forall c m. (Transport c, MonadUnliftIO m) => TMVar Bool -> ServiceName -> T.ServerParams -> (c -> m ()) -> m () -runTransportServer started port serverParams server = do - u <- askUnliftIO - liftIO $ do - clients <- newTVarIO S.empty - E.bracket - (startTCPServer started port) - (closeServer clients) - $ \sock -> forever $ do - (connSock, _) <- accept sock - tid <- forkIO $ connectClient u connSock `E.catch` \(_ :: E.SomeException) -> pure () - atomically . modifyTVar clients $ S.insert tid - where - connectClient :: UnliftIO m -> Socket -> IO () - connectClient u connSock = - E.bracket - (connectTLS serverParams connSock >>= getServerConnection) - closeConnection - (unliftIO u . server) - closeServer :: TVar (Set ThreadId) -> Socket -> IO () - closeServer clients sock = do - readTVarIO clients >>= mapM_ killThread - close sock - void . atomically $ tryPutTMVar started False - -startTCPServer :: TMVar Bool -> ServiceName -> IO Socket -startTCPServer started port = withSocketsDo $ resolve >>= open >>= setStarted - where - resolve = - let hints = defaultHints {addrFlags = [AI_PASSIVE], addrSocketType = Stream} - in head <$> getAddrInfo (Just hints) Nothing (Just port) - open addr = do - sock <- socket (addrFamily addr) (addrSocketType addr) (addrProtocol addr) - setSocketOption sock ReuseAddr 1 - withFdSocket sock setCloseOnExecIfNeeded - bind sock $ addrAddress addr - listen sock 1024 - return sock - setStarted sock = atomically (tryPutTMVar started True) >> pure sock - --- | Connect to passed TCP host:port and pass handle to the client. -runTransportClient :: Transport c => MonadUnliftIO m => HostName -> ServiceName -> C.KeyHash -> (c -> m a) -> m a -runTransportClient host port keyHash client = do - let clientParams = mkTLSClientParams host port keyHash - c <- liftIO $ startTCPClient host port clientParams - client c `E.finally` liftIO (closeConnection c) - -startTCPClient :: forall c. Transport c => HostName -> ServiceName -> T.ClientParams -> IO c -startTCPClient host port clientParams = withSocketsDo $ resolve >>= tryOpen err - where - err :: IOException - err = mkIOError NoSuchThing "no address" Nothing Nothing - - resolve :: IO [AddrInfo] - resolve = - let hints = defaultHints {addrSocketType = Stream} - in getAddrInfo (Just hints) (Just host) (Just port) - - tryOpen :: IOException -> [AddrInfo] -> IO c - tryOpen e [] = E.throwIO e - tryOpen _ (addr : as) = - E.try (open addr) >>= either (`tryOpen` as) pure - - open :: AddrInfo -> IO c - open addr = do - sock <- socket (addrFamily addr) (addrSocketType addr) (addrProtocol addr) - connect sock $ addrAddress addr - ctx <- connectTLS clientParams sock - getClientConnection ctx - -loadTLSServerParams :: FilePath -> FilePath -> FilePath -> IO T.ServerParams -loadTLSServerParams caCertificateFile certificateFile privateKeyFile = - fromCredential <$> loadServerCredential - where - loadServerCredential :: IO T.Credential - loadServerCredential = - T.credentialLoadX509Chain certificateFile [caCertificateFile] privateKeyFile >>= \case - Right credential -> pure credential - Left _ -> putStrLn "invalid credential" >> exitFailure - fromCredential :: T.Credential -> T.ServerParams - fromCredential credential = - def - { T.serverWantClientCert = False, - T.serverShared = def {T.sharedCredentials = T.Credentials [credential]}, - T.serverHooks = def, - T.serverSupported = supportedParameters - } - -loadFingerprint :: FilePath -> IO Fingerprint -loadFingerprint certificateFile = do - (cert : _) <- SX.readSignedObject certificateFile - pure $ XV.getFingerprint (cert :: X.SignedExact X.Certificate) X.HashSHA256 - -- * TLS Transport data TLS = TLS @@ -290,33 +176,6 @@ closeTLS ctx = (T.bye ctx >> T.contextClose ctx) -- sometimes socket was closed before 'TLS.bye' `E.catch` (\(_ :: E.SomeException) -> pure ()) -- so we catch the 'Broken pipe' error here -mkTLSClientParams :: HostName -> ServiceName -> C.KeyHash -> T.ClientParams -mkTLSClientParams host port keyHash = do - let p = B.pack port - (T.defaultParamsClient host p) - { T.clientShared = def, - T.clientHooks = def {T.onServerCertificate = \_ _ _ -> validateCertificateChain keyHash host p}, - T.clientSupported = supportedParameters - } - -validateCertificateChain :: C.KeyHash -> HostName -> ByteString -> X.CertificateChain -> IO [XV.FailedReason] -validateCertificateChain _ _ _ (X.CertificateChain []) = pure [XV.EmptyChain] -validateCertificateChain _ _ _ (X.CertificateChain [_]) = pure [XV.EmptyChain] -validateCertificateChain (C.KeyHash kh) host port cc@(X.CertificateChain sc@[_, caCert]) = - if Fingerprint kh == XV.getFingerprint caCert X.HashSHA256 - then x509validate - else pure [XV.UnknownCA] - where - x509validate :: IO [XV.FailedReason] - x509validate = XV.validate X.HashSHA256 hooks checks certStore cache serviceID cc - where - hooks = XV.defaultHooks - checks = XV.defaultChecks - certStore = XS.makeCertificateStore sc - cache = XV.exceptionValidationCache [] -- we manually check fingerprint only of the identity certificate (ca.crt) - serviceID = (host, port) -validateCertificateChain _ _ _ _ = pure [XV.AuthorityTooDeep] - supportedParameters :: T.Supported supportedParameters = def diff --git a/src/Simplex/Messaging/Transport/Client.hs b/src/Simplex/Messaging/Transport/Client.hs new file mode 100644 index 000000000..29bff966c --- /dev/null +++ b/src/Simplex/Messaging/Transport/Client.hs @@ -0,0 +1,82 @@ +{-# LANGUAGE ScopedTypeVariables #-} + +module Simplex.Messaging.Transport.Client + ( runTransportClient, + clientHandshake, + ) +where + +import Control.Monad.Except +import Control.Monad.IO.Unlift +import Data.ByteString.Char8 (ByteString) +import qualified Data.ByteString.Char8 as B +import Data.Default (def) +import qualified Data.X509 as X +import qualified Data.X509.CertificateStore as XS +import Data.X509.Validation (Fingerprint (..)) +import qualified Data.X509.Validation as XV +import GHC.IO.Exception (IOErrorType (..)) +import Network.Socket +import qualified Network.TLS as T +import qualified Simplex.Messaging.Crypto as C +import Simplex.Messaging.Transport +import System.IO.Error +import UnliftIO.Exception (IOException) +import qualified UnliftIO.Exception as E + +-- | Connect to passed TCP host:port and pass handle to the client. +runTransportClient :: Transport c => MonadUnliftIO m => HostName -> ServiceName -> C.KeyHash -> (c -> m a) -> m a +runTransportClient host port keyHash client = do + let clientParams = mkTLSClientParams host port keyHash + c <- liftIO $ startTCPClient host port clientParams + client c `E.finally` liftIO (closeConnection c) + +startTCPClient :: forall c. Transport c => HostName -> ServiceName -> T.ClientParams -> IO c +startTCPClient host port clientParams = withSocketsDo $ resolve >>= tryOpen err + where + err :: IOException + err = mkIOError NoSuchThing "no address" Nothing Nothing + + resolve :: IO [AddrInfo] + resolve = + let hints = defaultHints {addrSocketType = Stream} + in getAddrInfo (Just hints) (Just host) (Just port) + + tryOpen :: IOException -> [AddrInfo] -> IO c + tryOpen e [] = E.throwIO e + tryOpen _ (addr : as) = + E.try (open addr) >>= either (`tryOpen` as) pure + + open :: AddrInfo -> IO c + open addr = do + sock <- socket (addrFamily addr) (addrSocketType addr) (addrProtocol addr) + connect sock $ addrAddress addr + ctx <- connectTLS clientParams sock + getClientConnection ctx + +mkTLSClientParams :: HostName -> ServiceName -> C.KeyHash -> T.ClientParams +mkTLSClientParams host port keyHash = do + let p = B.pack port + (T.defaultParamsClient host p) + { T.clientShared = def, + T.clientHooks = def {T.onServerCertificate = \_ _ _ -> validateCertificateChain keyHash host p}, + T.clientSupported = supportedParameters + } + +validateCertificateChain :: C.KeyHash -> HostName -> ByteString -> X.CertificateChain -> IO [XV.FailedReason] +validateCertificateChain _ _ _ (X.CertificateChain []) = pure [XV.EmptyChain] +validateCertificateChain _ _ _ (X.CertificateChain [_]) = pure [XV.EmptyChain] +validateCertificateChain (C.KeyHash kh) host port cc@(X.CertificateChain sc@[_, caCert]) = + if Fingerprint kh == XV.getFingerprint caCert X.HashSHA256 + then x509validate + else pure [XV.UnknownCA] + where + x509validate :: IO [XV.FailedReason] + x509validate = XV.validate X.HashSHA256 hooks checks certStore cache serviceID cc + where + hooks = XV.defaultHooks + checks = XV.defaultChecks + certStore = XS.makeCertificateStore sc + cache = XV.exceptionValidationCache [] -- we manually check fingerprint only of the identity certificate (ca.crt) + serviceID = (host, port) +validateCertificateChain _ _ _ _ = pure [XV.AuthorityTooDeep] diff --git a/src/Simplex/Messaging/Transport/Server.hs b/src/Simplex/Messaging/Transport/Server.hs new file mode 100644 index 000000000..8057b4967 --- /dev/null +++ b/src/Simplex/Messaging/Transport/Server.hs @@ -0,0 +1,95 @@ +{-# LANGUAGE DuplicateRecordFields #-} +{-# LANGUAGE LambdaCase #-} +{-# LANGUAGE NamedFieldPuns #-} +{-# LANGUAGE ScopedTypeVariables #-} + +module Simplex.Messaging.Transport.Server + ( runTransportServer, + loadTLSServerParams, + loadFingerprint, + serverHandshake, + ) +where + +import Control.Monad.Except +import Control.Monad.IO.Unlift +import qualified Crypto.Store.X509 as SX +import Data.Default (def) +import Data.Set (Set) +import qualified Data.Set as S +import qualified Data.X509 as X +import Data.X509.Validation (Fingerprint (..)) +import qualified Data.X509.Validation as XV +import Network.Socket +import qualified Network.TLS as T +import Simplex.Messaging.Transport +import System.Exit (exitFailure) +import UnliftIO.Concurrent +import qualified UnliftIO.Exception as E +import UnliftIO.STM + +-- | Run transport server (plain TCP or WebSockets) on passed TCP port and signal when server started and stopped via passed TMVar. +-- +-- All accepted connections are passed to the passed function. +runTransportServer :: forall c m. (Transport c, MonadUnliftIO m) => TMVar Bool -> ServiceName -> T.ServerParams -> (c -> m ()) -> m () +runTransportServer started port serverParams server = do + u <- askUnliftIO + liftIO $ do + clients <- newTVarIO S.empty + E.bracket + (startTCPServer started port) + (closeServer clients) + $ \sock -> forever $ do + (connSock, _) <- accept sock + tid <- forkIO $ connectClient u connSock `E.catch` \(_ :: E.SomeException) -> pure () + atomically . modifyTVar clients $ S.insert tid + where + connectClient :: UnliftIO m -> Socket -> IO () + connectClient u connSock = + E.bracket + (connectTLS serverParams connSock >>= getServerConnection) + closeConnection + (unliftIO u . server) + closeServer :: TVar (Set ThreadId) -> Socket -> IO () + closeServer clients sock = do + readTVarIO clients >>= mapM_ killThread + close sock + void . atomically $ tryPutTMVar started False + +startTCPServer :: TMVar Bool -> ServiceName -> IO Socket +startTCPServer started port = withSocketsDo $ resolve >>= open >>= setStarted + where + resolve = + let hints = defaultHints {addrFlags = [AI_PASSIVE], addrSocketType = Stream} + in head <$> getAddrInfo (Just hints) Nothing (Just port) + open addr = do + sock <- socket (addrFamily addr) (addrSocketType addr) (addrProtocol addr) + setSocketOption sock ReuseAddr 1 + withFdSocket sock setCloseOnExecIfNeeded + bind sock $ addrAddress addr + listen sock 1024 + return sock + setStarted sock = atomically (tryPutTMVar started True) >> pure sock + +loadTLSServerParams :: FilePath -> FilePath -> FilePath -> IO T.ServerParams +loadTLSServerParams caCertificateFile certificateFile privateKeyFile = + fromCredential <$> loadServerCredential + where + loadServerCredential :: IO T.Credential + loadServerCredential = + T.credentialLoadX509Chain certificateFile [caCertificateFile] privateKeyFile >>= \case + Right credential -> pure credential + Left _ -> putStrLn "invalid credential" >> exitFailure + fromCredential :: T.Credential -> T.ServerParams + fromCredential credential = + def + { T.serverWantClientCert = False, + T.serverShared = def {T.sharedCredentials = T.Credentials [credential]}, + T.serverHooks = def, + T.serverSupported = supportedParameters + } + +loadFingerprint :: FilePath -> IO Fingerprint +loadFingerprint certificateFile = do + (cert : _) <- SX.readSignedObject certificateFile + pure $ XV.getFingerprint (cert :: X.SignedExact X.Certificate) X.HashSHA256 diff --git a/tests/SMPAgentClient.hs b/tests/SMPAgentClient.hs index 1922a33fa..4cf4f308a 100644 --- a/tests/SMPAgentClient.hs +++ b/tests/SMPAgentClient.hs @@ -20,12 +20,13 @@ import SMPClient withSmpServerOn, withSmpServerThreadOn, ) -import Simplex.Messaging.Agent (runSMPAgentBlocking) import Simplex.Messaging.Agent.Env.SQLite import Simplex.Messaging.Agent.Protocol import Simplex.Messaging.Agent.RetryInterval +import Simplex.Messaging.Agent.Server (runSMPAgentBlocking) import Simplex.Messaging.Client (SMPClientConfig (..), smpDefaultConfig) import Simplex.Messaging.Transport +import Simplex.Messaging.Transport.Client import Test.Hspec import UnliftIO.Concurrent import UnliftIO.Directory diff --git a/tests/SMPClient.hs b/tests/SMPClient.hs index 1dddd47cc..7d4352c1a 100644 --- a/tests/SMPClient.hs +++ b/tests/SMPClient.hs @@ -21,6 +21,7 @@ import Simplex.Messaging.Server (runSMPServerBlocking) import Simplex.Messaging.Server.Env.STM import Simplex.Messaging.Server.StoreLog (openReadStoreLog) import Simplex.Messaging.Transport +import Simplex.Messaging.Transport.Client import Test.Hspec import UnliftIO.Concurrent import qualified UnliftIO.Exception as E