warp-s2n-tls-0.1.1.0: test/ProtocolTests.hs
{-# LANGUAGE NumericUnderscores #-}
module ProtocolTests (protocolSpec) where
import Control.Applicative
import Control.Exception (SomeException, bracket, try)
import Data.ByteString (ByteString)
import Data.Default.Class (def)
import Network.HTTP.Types (status200)
import Network.Socket qualified as Socket
import Network.Socket.ByteString qualified as SocketBS
import Network.TLS qualified as TLS
import Network.TLS.Extra.Cipher qualified as TLS
import Network.Wai (responseLBS)
import Network.Wai.Handler.Warp (
Settings,
defaultSettings,
setBeforeMainLoop,
setHTTP2Disabled,
)
import Network.Wai.Handler.WarpS2N (
TLSSettings (..),
runTLSSocketLib,
tlsSettings,
)
import S2nTls (S2nTls (..))
import Test.Hspec (SpecWith, describe, expectationFailure, it, shouldBe)
import TestUtils
import UnliftIO.Async
import UnliftIO.Concurrent
import UnliftIO.STM
protocolSpec :: SpecWith S2nTls
protocolSpec = do
describe "TLS version negotiation" versionNegotiationSpec
describe "Cipher preferences" cipherPreferencesSpec
describe "ALPN" alpnSpec
alpnSpec :: SpecWith S2nTls
alpnSpec = do
it "negotiates h2 when the client offers it" $ \tls -> do
let serverSettings = tlsSettings testCertPath testKeyPath
proto <- withTestServerGetALPN tls serverSettings defaultSettings ["h2", "http/1.1"]
proto `shouldBe` Just "h2"
it "selects http/1.1 when the client does not offer h2" $ \tls -> do
let serverSettings = tlsSettings testCertPath testKeyPath
proto <- withTestServerGetALPN tls serverSettings defaultSettings ["http/1.1"]
proto `shouldBe` Just "http/1.1"
-- The advertised list is derived from Warp's own settingsHTTP2Enabled, so
-- setHTTP2Disabled has to withdraw h2 from the offer entirely. Leaving it
-- advertised would let a client select a protocol Warp then refuses to
-- speak, which is worse than never offering it.
it "withdraws h2 when the Warp settings disable HTTP/2" $ \tls -> do
let serverSettings = tlsSettings testCertPath testKeyPath
proto <-
withTestServerGetALPN
tls
serverSettings
(setHTTP2Disabled defaultSettings)
["h2", "http/1.1"]
proto `shouldBe` Just "http/1.1"
versionNegotiationSpec :: SpecWith S2nTls
versionNegotiationSpec = do
it "negotiates TLS 1.3 with default settings" $ \tls -> do
let serverSettings = tlsSettings testCertPath testKeyPath
negotiatedVersion <- withTestServerGetVersion tls serverSettings [TLS.TLS13, TLS.TLS12]
case negotiatedVersion of
Just TLS.TLS13 -> pure ()
Just v -> expectationFailure $ "Expected TLS 1.3, got " ++ show v
Nothing -> expectationFailure "Failed to get negotiated version"
it "negotiates TLS 1.2 when client only supports TLS 1.2" $ \tls -> do
let serverSettings =
(tlsSettings testCertPath testKeyPath)
{ tlsCipherPreferences = "default"
}
negotiatedVersion <- withTestServerGetVersion tls serverSettings [TLS.TLS12]
case negotiatedVersion of
Just TLS.TLS12 -> pure ()
Just v -> expectationFailure $ "Expected TLS 1.2, got " ++ show v
Nothing -> expectationFailure "Failed to get negotiated version"
it "fails when no common TLS version" $ \tls -> do
let serverSettings =
(tlsSettings testCertPath testKeyPath)
{ tlsCipherPreferences = "default_tls13"
}
result <-
try @SomeException $
withTestServerGetVersion tls serverSettings [TLS.TLS10]
case result of
Left _ -> pure ()
Right (Just v) -> expectationFailure $ "Should have failed but got " ++ show v
Right Nothing -> pure ()
cipherPreferencesSpec :: SpecWith S2nTls
cipherPreferencesSpec = do
it "accepts connection with default_tls13 policy" $ \tls -> do
let serverSettings =
(tlsSettings testCertPath testKeyPath)
{ tlsCipherPreferences = "default_tls13"
}
result <- withTestServerConnect tls serverSettings
result `shouldBe` True
it "accepts connection with default policy" $ \tls -> do
let serverSettings =
(tlsSettings testCertPath testKeyPath)
{ tlsCipherPreferences = "default"
}
result <- withTestServerConnect tls serverSettings
result `shouldBe` True
-- | Helper: Start server and get negotiated TLS version from client perspective
withTestServerGetVersion :: S2nTls -> TLSSettings -> [TLS.Version] -> IO (Maybe TLS.Version)
withTestServerGetVersion tls tlsSet clientVersions =
bracket bindFreePort (Socket.close . fst) $ \(sock, port) -> do
serverReady <- newEmptyTMVarIO
let app _ respond = respond $ responseLBS status200 [] "blarg!"
warpSet = setBeforeMainLoop (atomically $ putTMVar serverReady ()) defaultSettings
withAsync (runTLSSocketLib tls tlsSet warpSet sock app) $ \as -> do
atomically $ waitSTM as <|> takeTMVar serverReady
threadDelay 10_000
backend <- makeClientSocket port
params <- makeClientParams clientVersions
ctx <- TLS.contextNew backend params
TLS.handshake ctx
info <- TLS.contextGetInformation ctx
let version = TLS.infoVersion <$> info
TLS.bye ctx
pure version
{- | Helper: Start server and report the ALPN protocol the client settles on.
Takes the Warp 'Settings' explicitly, because the advertised protocol list is
derived from them.
-}
withTestServerGetALPN :: S2nTls -> TLSSettings -> Settings -> [ByteString] -> IO (Maybe ByteString)
withTestServerGetALPN tls tlsSet baseWarpSet clientProtos =
bracket bindFreePort (Socket.close . fst) $ \(sock, port) -> do
serverReady <- newEmptyTMVarIO
let app _ respond = respond $ responseLBS status200 [] "blarg!"
warpSet = setBeforeMainLoop (atomically $ putTMVar serverReady ()) baseWarpSet
withAsync (runTLSSocketLib tls tlsSet warpSet sock app) $ \as -> do
atomically $ waitSTM as <|> takeTMVar serverReady
threadDelay 10_000
backend <- makeClientSocket port
params <- makeClientParams [TLS.TLS13, TLS.TLS12]
let alpnParams =
params
{ TLS.clientHooks =
(TLS.clientHooks params)
{ TLS.onSuggestALPN = pure (Just clientProtos)
}
}
ctx <- TLS.contextNew backend alpnParams
TLS.handshake ctx
proto <- TLS.getNegotiatedProtocol ctx
TLS.bye ctx
pure proto
-- | Helper: Start server and test if connection succeeds
withTestServerConnect :: S2nTls -> TLSSettings -> IO Bool
withTestServerConnect tls tlsSet =
bracket bindFreePort (Socket.close . fst) $ \(sock, port) -> do
serverReady <- newEmptyTMVarIO
let app _ respond = respond $ responseLBS status200 [] "blarg!"
warpSet = setBeforeMainLoop (atomically $ putTMVar serverReady ()) defaultSettings
withAsync (runTLSSocketLib tls tlsSet warpSet sock app) $ \as -> do
atomically $ waitSTM as <|> takeTMVar serverReady
threadDelay 10_000
result <- try @SomeException $ do
backend <- makeClientSocket port
params <- makeClientParams [TLS.TLS13, TLS.TLS12]
ctx <- TLS.contextNew backend params
TLS.handshake ctx
TLS.sendData ctx "GET / HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n"
_ <- TLS.recvData ctx
TLS.bye ctx
pure $ case result of
Left _ -> False
Right _ -> True
-- | Create client TLS parameters
makeClientParams :: [TLS.Version] -> IO TLS.ClientParams
makeClientParams versions = do
pure $
(TLS.defaultParamsClient "localhost" "")
{ TLS.clientSupported =
def
{ TLS.supportedCiphers = TLS.ciphersuite_strong
, TLS.supportedVersions = versions
}
, TLS.clientShared = def
, TLS.clientHooks =
def
{ TLS.onServerCertificate = \_ _ _ _ -> pure []
}
}
-- | Create a client socket connected to localhost:port
makeClientSocket :: Int -> IO TLS.Backend
makeClientSocket port = do
sock <- Socket.socket Socket.AF_INET Socket.Stream Socket.defaultProtocol
Socket.connect sock (Socket.SockAddrInet (fromIntegral port) (Socket.tupleToHostAddress (127, 0, 0, 1)))
pure $
TLS.Backend
{ TLS.backendFlush = pure ()
, TLS.backendClose = Socket.close sock
, TLS.backendSend = SocketBS.sendAll sock
, TLS.backendRecv = SocketBS.recv sock
}
-- | Bind to a free port
bindFreePort :: IO (Socket.Socket, Int)
bindFreePort = do
sock <- Socket.socket Socket.AF_INET Socket.Stream Socket.defaultProtocol
Socket.setSocketOption sock Socket.ReuseAddr 1
Socket.bind sock (Socket.SockAddrInet 0 (Socket.tupleToHostAddress (127, 0, 0, 1)))
Socket.listen sock 5
port <- Socket.socketPort sock
pure (sock, fromIntegral port)