packages feed

tls-2.4.9: test/HandshakeSpec.hs

{-# LANGUAGE OverloadedStrings #-}

module HandshakeSpec where

import Control.Concurrent (threadDelay)
import Control.Concurrent.Async (concurrently_)
import qualified Control.Exception as E
import Control.Monad
import Crypto.Cipher.AES (AES128)
import Crypto.Cipher.Types (AEADMode (..), AuthTag (..), aeadInit, aeadSimpleEncrypt, cipherInit)
import Crypto.Error (throwCryptoError)
import Data.Bits (shiftR, xor)
import qualified Data.ByteArray as BA
import qualified Data.ByteString as B
import qualified Data.ByteString.Lazy as L
import Data.IORef
import Data.List
import Data.Maybe
import Data.Word (Word16, Word64, Word8)
import Data.X509 (ExtKeyUsageFlag (..), ExtKeyUsagePurpose (..))
import Network.TLS
import Network.TLS.Extra.Cipher
import Network.TLS.Extra.CipherCBC
import Network.TLS.Internal
import Network.TLS.QUIC (hkdfExpandLabel)
import Test.Hspec
import Test.Hspec.QuickCheck
import System.Timeout (timeout)
import Test.QuickCheck

import API
import Arbitrary
import Certificate (arbitraryRSACredentialWithPurpose)
import PipeChan
import Run
import Session

spec :: Spec
spec = do
    describe "pipe" $ do
        it "can setup a channel" pipe_work
    describe "channel binding" $ do
        prop "is unavailable before the handshake" binding_before_handshake
    describe "handshake" $ do
        prop "can run TLS 1.2" handshake_simple
        prop "can run TLS 1.3" handshake13_simple
        prop "can update key for TLS 1.3" handshake_update_key
        it
            "rejects more than 32 consecutive TLS 1.3 KeyUpdates"
            handshake_key_update_flood
        it
            "can disable the consecutive TLS 1.3 KeyUpdate limit"
            handshake_key_update_unlimited
        it
            "does not disable TLS 1.3 KeyUpdates with non-positive limits"
            handshake_key_update_non_positive
        prop "can prevent downgrade attack" handshake13_downgrade
        it "negotiates TLS 1.2 with a ClientHello naming a higher version" $
            handshake_high_legacy_version TLS12 (Version 0x0309)
        it "ignores legacy_version when supported_versions is present" $
            handshake_high_legacy_version TLS13 TLS13
        it "rejects ec_point_formats without uncompressed" $
            handshake12_ec_point_formats
                (B.pack [1, 1])
                rejectedAsIllegalParameter
        it "rejects an empty ec_point_formats" $
            handshake12_ec_point_formats (B.pack [0]) rejectedAsDecodeError
        prop "can negotiate hash and signature" handshake_hashsignatures
        prop "can negotiate cipher suite" handshake_ciphersuites
        it "rejects a cipher outside the server callback candidates" $
            handshake_rejects_server_cipher_callback_escape
        it "rejects a TLS 1.2-only cipher selected for TLS 1.3" $
            handshake_rejects_legacy_cipher_in_tls13
        prop "can negotiate group" handshake_groups
        prop "can negotiate elliptic curve" handshake_ec
        prop "can fallback for certificate with cipher" handshake_cert_fallback_cipher
        prop
            "can fallback for certificate with hash and signature"
            handshake_cert_fallback_hs
        prop "can handle server key usage" handshake_server_key_usage
        it "accepts a TLS 1.2 server certificate permitting server auth" $
            handshake_server_key_purpose TLS12 KeyUsagePurpose_ServerAuth True
        it "accepts a TLS 1.3 server certificate permitting server auth" $
            handshake_server_key_purpose TLS13 KeyUsagePurpose_ServerAuth True
        it "rejects a TLS 1.2 server certificate restricted to client auth" $
            handshake_server_key_purpose TLS12 KeyUsagePurpose_ClientAuth False
        it "rejects a TLS 1.3 server certificate restricted to client auth" $
            handshake_server_key_purpose TLS13 KeyUsagePurpose_ClientAuth False
        prop "can handle client key usage" handshake_client_key_usage
        prop "can authenticate client" handshake_client_auth
        it "rejects a TLS 1.3 CertificateVerify algorithm unfit for the key" $
            handshake13_client_cert_verify_unfit_sigalg
        it "rejects a TLS 1.2 CertificateVerify algorithm for another key type" $
            handshake12_client_cert_verify_sigalg
                (HashSHA256, SignatureECDSA)
                rejectedAsIllegalParameter
        it "rejects a TLS 1.2 CertificateVerify that does not fit the RSA key" $
            handshake12_client_cert_verify_sigalg
                (HashIntrinsic, SignatureRSApsspssSHA256)
                rejectedAsDecryptError
        prop "can receive client authentication failure" handshake_client_auth_fail
        it "accepts an empty TLS 1.2 client certificate when the hook does" $
            handshake_client_auth_empty TLS12
        it "accepts an empty TLS 1.3 client certificate when the hook does" $
            handshake_client_auth_empty TLS13
        it "requests only defined certificate types in TLS 1.2" $
            handshake12_cert_request_types
        prop "can handle extended main secret" handshake_ems
        prop "can resume with extended main secret" handshake_resumption_ems
        prop "can handle ALPN" handshake_alpn
        it "rejects an unoffered ALPN selection from the server hook" $
            handshake_alpn_rejects_unoffered_server_selection
        it "rejects an unoffered ALPN selection received by the client" $
            handshake_alpn_rejects_unoffered_client_selection
        prop "can handle SNI" handshake_sni
        it "sends no SNI for an empty server name" handshake_sni_empty
        it "rejects multiple host_names in SNI" $
            handshake_sni_illegal ["example.com", "example.org"]
        it "rejects a host_name with a control character in SNI" $
            handshake_sni_illegal ["example\0.com"]
        it "rejects a non-ASCII host_name in SNI" $
            handshake_sni_illegal ["ex\xc4\x85mple.com"]
        prop "can handshake with TLS 1.2 CBC" handshake_cbc
        prop "can re-negotiate with TLS 1.2" handshake12_renegotiation
        it "rejects SCSV in a secure renegotiation" $
            handshake12_renegotiation_tampered $ \ch ->
                ch{chCiphers = chCiphers ch ++ [CipherId 0xff]}
        it "rejects a secure renegotiation without renegotiation_info" $
            handshake12_renegotiation_tampered $ \ch ->
                ch
                    { chExtensions =
                        filter
                            (\(ExtensionRaw eid _) -> eid /= EID_SecureRenegotiation)
                            (chExtensions ch)
                    }
        prop "can resume session with TLS 1.2" handshake12_session_resumption
        prop
            "rejects resuming a TLS 1.2 session without its cipher"
            handshake12_session_resumption_cipher_missing
        prop "can resume session ticket with TLS 1.2" handshake12_session_ticket
        it "sends no session ticket to a TLS 1.2 client that did not ask" $
            handshake12_session_ticket_unoffered
        prop "can handshake with TLS 1.3 Full" handshake13_full
        prop "can handshake with TLS 1.3 HRR" handshake13_hrr
        prop "can handshake with TLS 1.3 PSK" handshake13_psk
        it "does not resume a TLS 1.2 session with a TLS 1.3 PSK" $
            handshake13_psk_tls12_session
        prop "can handshake with TLS 1.3 PSK ticket" handshake13_psk_ticket
        prop "can handshake with TLS 1.3 PSK -> HRR" handshake13_psk_fallback
        prop "can handshake with TLS 1.3 0RTT" handshake13_0rtt
        it "rejects TLS 1.3 early data when ALPN changes" $
            handshake13_0rtt_alpn
        prop "can handshake with TLS 1.3 0RTT -> PSK" handshake13_0rtt_fallback
        prop "can handshake with TLS 1.3 EE" handshake13_ee_groups
        prop "can handshake with TLS 1.3 EC groups" handshake13_ec
        prop "can handshake with TLS 1.3 FFDHE groups" handshake13_ffdhe
        it "rejects an X25519MLKEM768 key share with a zero X25519 part" $
            handshake13_x25519mlkem768_zero_x25519
        prop "can handshake with TLS 1.3 Post-handshake auth" post_handshake_auth
        it
            "keeps record alignment when a slow record follows client auth"
            handshake13_client_auth_slow_record
        it "rejects a too short TLS 1.2 AEAD record with bad_record_mac" $
            short_record_bad_record_mac TLS12 cipher_ECDHE_RSA_WITH_AES_128_GCM_SHA256
        it "rejects a too short TLS 1.2 CBC record with bad_record_mac" $
            short_record_bad_record_mac TLS12 cipher_ECDHE_RSA_AES128CBC_SHA256
        it "rejects a too short TLS 1.3 record with bad_record_mac" $
            short_record_bad_record_mac TLS13 cipher13_AES_128_GCM_SHA256
        it "rejects a ServerHello as the first client message" $
            server_first_message_unexpected 2
        it "rejects a Finished as the first client message" $
            server_first_message_unexpected 20
        it "rejects an unknown handshake type as the first client message" $
            server_first_message_unexpected 254
        it "rejects a ChangeCipherSpec before the first ClientHello" $
            server_ccs_interleaved 0
        it "rejects a ChangeCipherSpec inside a fragmented ClientHello" $
            server_ccs_interleaved 2
        it "rejects an SSLv2-style record header at once" $
            server_first_record_type_unexpected 0x80 0x3fff
        it "rejects an unknown record type" $
            server_first_record_type_unexpected 24 0
        it "rejects application data inside a TLS 1.3 handshake message" $
            handshake13_interleaved_app_data
        it "rejects a two-byte TLS 1.2 ChangeCipherSpec" $
            malformed_ccs_unexpected
                TLS12
                cipher_ECDHE_RSA_WITH_AES_128_GCM_SHA256
                [P256]
                [P256]
        it "rejects a two-byte TLS 1.3 ChangeCipherSpec" $
            malformed_ccs_unexpected TLS13 cipher13_AES_128_GCM_SHA256 [X25519] [X25519]
        it "rejects a two-byte TLS 1.3 ChangeCipherSpec after HelloRetryRequest" $
            malformed_ccs_unexpected
                TLS13
                cipher13_AES_128_GCM_SHA256
                [P256, X25519]
                [X25519]

--------------------------------------------------------------

-- | Both channel bindings already answer with 'Maybe', and a caller may
-- reasonably ask for one on a context whose handshake has not run -- or has
-- failed.  The answer is that there is no binding, not a crash.
binding_before_handshake :: (ClientParams, ServerParams) -> IO ()
binding_before_handshake params = withPairContext params $ \(cCtx, sCtx) ->
    forM_ [cCtx, sCtx] $ \ctx -> do
        getTLSUnique ctx `shouldReturn` Nothing
        getTLSExporter ctx `shouldReturn` Nothing

pipe_work :: IO ()
pipe_work = do
    pipe <- newPipe
    _ <- runPipe pipe

    let bSize = 16
    n <- generate (choose (1, 32))

    let d1 = B.replicate (bSize * n) 40
    let d2 = B.replicate (bSize * n) 45

    d1' <- writePipeC pipe d1 >> readPipeS pipe (B.length d1)
    d1' `shouldBe` d1

    d2' <- writePipeS pipe d2 >> readPipeC pipe (B.length d2)
    d2' `shouldBe` d2

--------------------------------------------------------------

handshake_simple :: (ClientParams, ServerParams) -> IO ()
handshake_simple = runTLSSimple

--------------------------------------------------------------

newtype CSP13 = CSP13 (ClientParams, ServerParams) deriving (Show)

instance Arbitrary CSP13 where
    arbitrary = CSP13 <$> arbitraryPairParams13

handshake13_simple :: CSP13 -> IO ()
handshake13_simple (CSP13 params) = runTLSSimple13 params hs
  where
    cgrps = supportedGroups $ clientSupported $ fst params
    sgrps = supportedGroups $ serverSupported $ snd params
    hs = if unsafeHead cgrps `elem` sgrps then FullHandshake else HelloRetryRequest

handshake_rejects_server_cipher_callback_escape :: IO ()
handshake_rejects_server_cipher_callback_escape = do
    (clientParam, serverParam) <- generate arbitraryPairParams13
    let params = cipherSelectionParams clientParam serverParam selectLegacy
    withPairContextWith (id, id) params $ \(cctx, sctx) ->
        concurrently_
            (handshake sctx `shouldThrow` serverRejectedCipherEscape)
            (handshake cctx `shouldThrow` anyTLSException)
  where
    selectLegacy _ _ = cipher_ECDHE_RSA_AES128CBC_SHA256

handshake_rejects_legacy_cipher_in_tls13 :: IO ()
handshake_rejects_legacy_cipher_in_tls13 = do
    (clientParam, serverParam) <- generate arbitraryPairParams13
    let params = cipherSelectionParams clientParam serverParam defaultSelection
    withPairContextWith (id, id) params $ \(cctx, sctx) -> do
        contextHookSetHandshakeRecv cctx tamperCipher
        concurrently_
            (handshake sctx `shouldThrow` anyTLSException)
            (handshake cctx `shouldThrow` clientRejectedLegacyCipher)
  where
    defaultSelection _ = unsafeHead
    tamperCipher (ServerHello sh) =
        pure $
            ServerHello
                sh
                    { shCipher =
                        CipherId $ cipherID cipher_ECDHE_RSA_AES128CBC_SHA256
                    }
    tamperCipher hs = pure hs

cipherSelectionParams
    :: ClientParams
    -> ServerParams
    -> (Version -> [Cipher] -> Cipher)
    -> (ClientParams, ServerParams)
cipherSelectionParams clientParam serverParam select =
    ( clientParam{clientSupported = supported}
    , serverParam
        { serverSupported = supported
        , serverHooks =
            (serverHooks serverParam)
                { onCipherChoosing = select
                }
        }
    )
  where
    supported =
        defaultSupported
            { supportedVersions = [TLS13]
            , supportedCiphers =
                [ cipher13_AES_128_GCM_SHA256
                , cipher_ECDHE_RSA_AES128CBC_SHA256
                ]
            }

serverRejectedCipherEscape :: TLSException -> Bool
serverRejectedCipherEscape (HandshakeFailed (Error_Protocol msg alert)) =
    msg == "onCipherChoosing selected a cipher outside the candidate list"
        && alert == InternalError
serverRejectedCipherEscape _ = False

clientRejectedLegacyCipher :: TLSException -> Bool
clientRejectedLegacyCipher (HandshakeFailed (Error_Protocol msg alert)) =
    msg == "server selected a cipher invalid for the negotiated version"
        && alert == IllegalParameter
clientRejectedLegacyCipher _ = False

anyTLSException :: TLSException -> Bool
anyTLSException = const True

handshake_key_update_flood :: IO ()
handshake_key_update_flood = do
    params <- generate arbitraryPairParams13
    withPairContextWith (id, id) params $ \(cctx, sctx) ->
        concurrently_
            ( do
                handshake sctx
                recvData sctx `shouldReturn` "after 32 key updates"
                recvData sctx `shouldThrow` excessiveKeyUpdate
            )
            ( do
                handshake cctx
                replicateM_ 32 $ void $ updateKey cctx OneWay
                sendData cctx "after 32 key updates"
                replicateM_ 33 $ void $ updateKey cctx OneWay
                sendData cctx "after 33 key updates"
            )
  where
    excessiveKeyUpdate
        (Terminated _ _ (Error_Misc "too many consecutive KeyUpdate messages")) = True
    excessiveKeyUpdate _ = False

handshake_key_update_unlimited :: IO ()
handshake_key_update_unlimited = do
    (cparams, sparams0) <- generate arbitraryPairParams13
    let shared0 = serverShared sparams0
        limits = (sharedLimit shared0){limitKeyUpdate = Nothing}
        sparams = sparams0{serverShared = shared0{sharedLimit = limits}}
    withPairContextWith (id, id) (cparams, sparams) $ \(cctx, sctx) ->
        concurrently_
            ( do
                handshake sctx
                recvData sctx `shouldReturn` "after 33 key updates"
            )
            ( do
                handshake cctx
                replicateM_ 33 $ void $ updateKey cctx OneWay
                sendData cctx "after 33 key updates"
            )

handshake_key_update_non_positive :: IO ()
handshake_key_update_non_positive =
    forM_ [0, -1] $ \limit -> do
        (cparams, sparams0) <- generate arbitraryPairParams13
        let shared0 = serverShared sparams0
            limits = (sharedLimit shared0){limitKeyUpdate = Just limit}
            sparams = sparams0{serverShared = shared0{sharedLimit = limits}}
        withPairContextWith (id, id) (cparams, sparams) $ \(cctx, sctx) ->
            concurrently_
                ( do
                    handshake sctx
                    recvData sctx `shouldReturn` "after key update"
                )
                ( do
                    handshake cctx
                    void $ updateKey cctx OneWay
                    sendData cctx "after key update"
                )

--------------------------------------------------------------

handshake_cbc :: IO ()
handshake_cbc = do
    clientCiphers <- generate $ cipherGen >>= shuffle
    serverCiphers <- generate $ cipherGen >>= shuffle
    clientGroups <- generate $ groupGen >>= shuffle
    serverGroups <- generate $ groupGen >>= shuffle
    (clientParam, serverParam) <- generate $
        arbitraryPairParamsWithVersionsAndCiphers
            ([TLS12], [TLS12])
            (clientCiphers, serverCiphers)
    let clientParam' = clientParam {
            clientSupported = (clientSupported clientParam)
                { supportedGroups = clientGroups } }
        serverParam' = serverParam {
            serverSupported = (serverSupported serverParam)
                { supportedGroups = serverGroups } }
    let ciphers = clientCiphers `intersect` serverCiphers
        groups = clientGroups `intersect` serverGroups
     in if compat ciphers groups
        then runTLSSimple (clientParam', serverParam')
        else runTLSFailure (clientParam', serverParam') handshake handshake
  where
    groupGen :: Gen [Group]
    groupGen = sublistOf grps `suchThat` (not . null)
      where
        grps = [X25519, P256, P384, FFDHE2048, FFDHE3072, FFDHE4096]

    cipherGen :: Gen [Cipher]
    cipherGen = sublistOf ciphersuite_pfs_sha2_cbc `suchThat` (not . null)

    compat :: [Cipher] -> [Group] -> Bool
    compat [] _ = False
    compat _ [] = False
    compat ciphers groups =
        let mustdh = all (== CipherKeyExchange_DHE_RSA) $ map cipherKeyExchange ciphers
            mustec = all (/= CipherKeyExchange_DHE_RSA) $ map cipherKeyExchange ciphers
            havedh = any (`elem` [FFDHE2048, FFDHE3072, FFDHE4096]) groups
            haveec = any (`elem` [X25519, P256, P384]) groups
         in ((not mustdh || havedh) && (not mustec || haveec))

--------------------------------------------------------------

handshake13_downgrade :: (ClientParams, ServerParams) -> IO ()
handshake13_downgrade (cparam, sparam) = do
    versionForced <-
        generate $ elements (supportedVersions $ clientSupported cparam)
    let debug' = (serverDebug sparam){debugVersionForced = Just versionForced}
        sparam' = sparam{serverDebug = debug'}
        params = (cparam, sparam')
        downgraded =
            (isVersionEnabled TLS13 params && versionForced < TLS13)
                || (isVersionEnabled TLS12 params && versionForced < TLS12)
    if downgraded
        then runTLSFailure params handshake handshake
        else runTLSSimple params

handshake_update_key :: (ClientParams, ServerParams) -> IO ()
handshake_update_key = runTLSSimpleKeyUpdate

--------------------------------------------------------------

handshake_hashsignatures
    :: ([HashAndSignatureAlgorithm], [HashAndSignatureAlgorithm]) -> IO ()
handshake_hashsignatures (clientHashSigs, serverHashSigs) = do
    tls13 <- generate arbitrary
    let version = if tls13 then TLS13 else TLS12
        ciphers =
            [ cipher_ECDHE_RSA_WITH_AES_256_GCM_SHA384
            , cipher_ECDHE_ECDSA_WITH_AES_256_GCM_SHA384
            , cipher13_AES_128_GCM_SHA256
            ]
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                ([version], [version])
                (ciphers, ciphers)
    let clientParam' =
            clientParam
                { clientSupported =
                    (clientSupported clientParam)
                        { supportedHashSignatures = clientHashSigs
                        }
                }
        serverParam' =
            serverParam
                { serverSupported =
                    (serverSupported serverParam)
                        { supportedHashSignatures = serverHashSigs
                        }
                }
        commonHashSigs = clientHashSigs `intersect` serverHashSigs
        shouldFail
            | tls13 = all incompatibleWithDefaultCurve commonHashSigs
            | otherwise = null commonHashSigs
    if shouldFail
        then runTLSFailure (clientParam', serverParam') handshake handshake
        else runTLSSimple (clientParam', serverParam')
  where
    incompatibleWithDefaultCurve (h, SignatureECDSA) = h /= HashSHA256
    incompatibleWithDefaultCurve _ = False

handshake_ciphersuites :: ([Cipher], [Cipher]) -> IO ()
handshake_ciphersuites (clientCiphers, serverCiphers) = do
    tls13 <- generate arbitrary
    let version = if tls13 then TLS13 else TLS12
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                ([version], [version])
                (clientCiphers, serverCiphers)
    let adequate = cipherAllowedForVersion version
        shouldSucceed = any adequate (clientCiphers `intersect` serverCiphers)
    if shouldSucceed
        then runTLSSimple (clientParam, serverParam)
        else runTLSFailure (clientParam, serverParam) handshake handshake

--------------------------------------------------------------

handshake_groups :: GGP -> IO ()
handshake_groups (GGP clientGroups serverGroups) = do
    tls13 <- generate arbitrary
    let versions = if tls13 then [TLS13] else [TLS12]
        ciphers = ciphersuite_strong
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                (versions, versions)
                (ciphers, ciphers)
    denyCustom <- generate arbitrary
    let groupUsage =
            if denyCustom
                then GroupUsageUnsupported "custom group denied"
                else GroupUsageValid
        clientParam' =
            clientParam
                { clientSupported =
                    (clientSupported clientParam)
                        { supportedGroups = clientGroups
                        }
                , clientHooks =
                    (clientHooks clientParam)
                        { onCustomFFDHEGroup = \_ _ -> return groupUsage
                        }
                }
        serverParam' =
            serverParam
                { serverSupported =
                    (serverSupported serverParam)
                        { supportedGroups = serverGroups
                        , supportedGroupsTLS13 = [serverGroups]
                        }
                }
        commonGroups = clientGroups `intersect` serverGroups
        shouldFail = null commonGroups
        p minfo = isNothing (minfo >>= infoSupportedGroup) == null commonGroups
    if shouldFail
        then runTLSFailure (clientParam', serverParam') handshake handshake
        else runTLSPredicate (clientParam', serverParam') p

--------------------------------------------------------------

newtype SG = SG [Group] deriving (Show)

instance Arbitrary SG where
    arbitrary = SG <$> shuffle sigGroups
      where
        sigGroups = [P256, P521]

handshake_ec :: SG -> IO ()
handshake_ec (SG sigGroups) = do
    let versions = [TLS12]
        ciphers =
            [ cipher_ECDHE_ECDSA_WITH_AES_256_GCM_SHA384
            ]
        hashSignatures =
            [ (HashSHA256, SignatureECDSA)
            ]
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                (versions, versions)
                (ciphers, ciphers)
    clientGroups <- generate $ shuffle sigGroups
    clientHashSignatures <- generate $ sublistOf hashSignatures
    serverHashSignatures <- generate $ sublistOf hashSignatures
    credentials <- generate arbitraryCredentialsOfEachCurve
    let clientParam' =
            clientParam
                { clientSupported =
                    (clientSupported clientParam)
                        { supportedGroups = clientGroups
                        , supportedHashSignatures = clientHashSignatures
                        }
                }
        serverParam' =
            serverParam
                { serverSupported =
                    (serverSupported serverParam)
                        { supportedGroups = sigGroups
                        , supportedGroupsTLS13 = [sigGroups]
                        , supportedHashSignatures = serverHashSignatures
                        }
                , serverShared =
                    (serverShared serverParam)
                        { sharedCredentials = Credentials credentials
                        }
                }
        sigAlgs = map snd (clientHashSignatures `intersect` serverHashSignatures)
        ecdsaDenied = SignatureECDSA `notElem` sigAlgs
    if ecdsaDenied
        then runTLSFailure (clientParam', serverParam') handshake handshake
        else runTLSSimple (clientParam', serverParam')

-- Tests ability to use or ignore client "signature_algorithms" extension when
-- choosing a server certificate.  Here peers allow DHE_RSA_AES128_SHA1 but
-- the server RSA certificate has a SHA-1 signature that the client does not
-- support.  Server may choose the DSA certificate only when cipher
-- DHE_DSA_AES128_SHA1 is allowed.  Otherwise it must fallback to the RSA
-- certificate.

data OC = OC [Cipher] [Cipher] deriving (Show)

instance Arbitrary OC where
    arbitrary = OC <$> sublistOf otherCiphers <*> sublistOf otherCiphers
      where
        otherCiphers =
            [ cipher_ECDHE_RSA_WITH_AES_256_GCM_SHA384
            , cipher_ECDHE_RSA_WITH_AES_128_GCM_SHA256
            ]

handshake_cert_fallback_cipher :: OC -> IO ()
handshake_cert_fallback_cipher (OC clientCiphers serverCiphers) = do
    let clientVersions = [TLS12]
        serverVersions = [TLS12]
        commonCiphers = [cipher_ECDHE_RSA_WITH_AES_128_GCM_SHA256]
        hashSignatures = [(HashSHA256, SignatureRSA), (HashSHA1, SignatureDSA)]
    chainRef <- newIORef Nothing
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                (clientVersions, serverVersions)
                (clientCiphers ++ commonCiphers, serverCiphers ++ commonCiphers)
    let clientParam' =
            clientParam
                { clientSupported =
                    (clientSupported clientParam)
                        { supportedHashSignatures = hashSignatures
                        }
                , clientHooks =
                    (clientHooks clientParam)
                        { onServerCertificate = \_ _ _ chain ->
                            writeIORef chainRef (Just chain) >> return []
                        }
                }
    runTLSSimple (clientParam', serverParam)
    serverChain <- readIORef chainRef
    isLeafRSA serverChain `shouldBe` True

-- Same as above but testing with supportedHashSignatures directly instead of
-- ciphers, and thus allowing TLS13.  Peers accept RSA with SHA-256 but the
-- server RSA certificate has a SHA-1 signature.  When Ed25519 is allowed by
-- both client and server, the Ed25519 certificate is selected.  Otherwise the
-- server fallbacks to RSA.
--
-- Note: SHA-1 is supposed to be disallowed in X.509 signatures with TLS13
-- unless client advertises explicit support.  Currently this is not enforced by
-- the library, which is useful to test this scenario.  SHA-1 could be replaced
-- by another algorithm.

data OHS = OHS [HashAndSignatureAlgorithm] [HashAndSignatureAlgorithm]
    deriving (Show)

instance Arbitrary OHS where
    arbitrary = OHS <$> sublistOf otherHS <*> sublistOf otherHS
      where
        otherHS = [(HashIntrinsic, SignatureEd25519)]

handshake_cert_fallback_hs :: OHS -> IO ()
handshake_cert_fallback_hs (OHS clientHS serverHS) = do
    tls13 <- generate arbitrary
    let versions = if tls13 then [TLS13] else [TLS12]
        ciphers =
            [ cipher_ECDHE_RSA_WITH_AES_128_GCM_SHA256
            , cipher_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256
            , cipher13_AES_128_GCM_SHA256
            ]
        commonHS =
            [ (HashSHA256, SignatureRSA)
            , (HashIntrinsic, SignatureRSApssRSAeSHA256)
            ]
    chainRef <- newIORef Nothing
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                (versions, versions)
                (ciphers, ciphers)
    let clientParam' =
            clientParam
                { clientSupported =
                    (clientSupported clientParam)
                        { supportedHashSignatures = commonHS ++ clientHS
                        }
                , clientHooks =
                    (clientHooks clientParam)
                        { onServerCertificate = \_ _ _ chain ->
                            writeIORef chainRef (Just chain) >> return []
                        }
                }
        serverParam' =
            serverParam
                { serverSupported =
                    (serverSupported serverParam)
                        { supportedHashSignatures = commonHS ++ serverHS
                        }
                }
        eddsaDisallowed =
            (HashIntrinsic, SignatureEd25519) `notElem` clientHS
                || (HashIntrinsic, SignatureEd25519) `notElem` serverHS
    runTLSSimple (clientParam', serverParam')
    serverChain <- readIORef chainRef
    isLeafRSA serverChain `shouldBe` eddsaDisallowed

--------------------------------------------------------------

handshake_server_key_usage :: [ExtKeyUsageFlag] -> IO ()
handshake_server_key_usage usageFlags = do
    tls13 <- generate arbitrary
    let versions = if tls13 then [TLS13] else [TLS12]
        ciphers = ciphersuite_all
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                (versions, versions)
                (ciphers, ciphers)
    cred <- generate $ arbitraryRSACredentialWithUsage usageFlags
    let serverParam' =
            serverParam
                { serverShared =
                    (serverShared serverParam)
                        { sharedCredentials = Credentials [cred]
                        }
                }
        shouldSucceed = KeyUsage_digitalSignature `elem` usageFlags
    if shouldSucceed
        then runTLSSimple (clientParam, serverParam')
        else runTLSFailure (clientParam, serverParam') handshake handshake

handshake_server_key_purpose :: Version -> ExtKeyUsagePurpose -> Bool -> IO ()
handshake_server_key_purpose version purpose shouldSucceed = do
    let cipher
            | version == TLS13 = cipher13_AES_128_GCM_SHA256
            | otherwise = cipher_ECDHE_RSA_WITH_AES_128_GCM_SHA256
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                ([version], [version])
                ([cipher], [cipher])
    cred <- generate $ arbitraryRSACredentialWithPurpose purpose
    let clientParam' =
            clientParam
                { clientHooks =
                    (clientHooks clientParam)
                        { onServerCertificate = \_ _ _ _ -> return []
                        }
                }
        serverParam' =
            serverParam
                { serverShared =
                    (serverShared serverParam)
                        { sharedCredentials = Credentials [cred]
                        }
                }
    if shouldSucceed
        then runTLSSimple (clientParam', serverParam')
        else runTLSFailure (clientParam', serverParam') handshake handshake

handshake_client_key_usage :: [ExtKeyUsageFlag] -> IO ()
handshake_client_key_usage usageFlags = do
    (clientParam, serverParam) <- generate arbitrary
    cred <- generate $ arbitraryRSACredentialWithUsage usageFlags
    let clientParam' =
            clientParam
                { clientHooks =
                    (clientHooks clientParam)
                        { onCertificateRequest = \_ -> return $ Just cred
                        }
                }
        serverParam' =
            serverParam
                { serverWantClientCert = True
                , serverHooks =
                    (serverHooks serverParam)
                        { onClientCertificate = \_ -> return CertificateUsageAccept
                        }
                }
        shouldSucceed = KeyUsage_digitalSignature `elem` usageFlags
    if shouldSucceed
        then runTLSSimple (clientParam', serverParam')
        else runTLSFailure (clientParam', serverParam') handshake handshake

--------------------------------------------------------------

handshake_client_auth :: (ClientParams, ServerParams) -> IO ()
handshake_client_auth (clientParam, serverParam) = do
    let clientVersions = supportedVersions $ clientSupported clientParam
        serverVersions = supportedVersions $ serverSupported serverParam
        version = maximum (clientVersions `intersect` serverVersions)
    cred <- generate (arbitraryClientCredential version)
    let clientParam' =
            clientParam
                { clientHooks =
                    (clientHooks clientParam)
                        { onCertificateRequest = \_ -> return $ Just cred
                        }
                }
        serverParam' =
            serverParam
                { serverWantClientCert = True
                , serverHooks =
                    (serverHooks serverParam)
                        { onClientCertificate = validateChain cred
                        }
                }
    runTLSSimple (clientParam', serverParam')
  where
    validateChain cred chain
        | chain == fst cred = return CertificateUsageAccept
        | otherwise = return (CertificateUsageReject CertificateRejectUnknownCA)

-- A client without a certificate answers CertificateRequest with an empty
-- Certificate and, in TLS 1.2, sends no CertificateVerify.  A server whose
-- hook accepts that must go on to ChangeCipherSpec rather than wait for a
-- CertificateVerify.
handshake_client_auth_empty :: Version -> IO ()
handshake_client_auth_empty version = do
    let cipher
            | version == TLS13 = cipher13_AES_128_GCM_SHA256
            | otherwise = cipher_ECDHE_RSA_WITH_AES_128_GCM_SHA256
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                ([version], [version])
                ([cipher], [cipher])
    let clientParam' =
            clientParam
                { clientHooks =
                    (clientHooks clientParam)
                        { onCertificateRequest = \_ -> return Nothing
                        }
                }
        serverParam' =
            serverParam
                { serverWantClientCert = True
                , serverHooks =
                    (serverHooks serverParam)
                        { onClientCertificate = acceptEmpty
                        }
                }
    runTLSSimple (clientParam', serverParam')
  where
    acceptEmpty chain
        | isNullCertificateChain chain = return CertificateUsageAccept
        | otherwise = return (CertificateUsageReject CertificateRejectUnknownCA)

-- In TLS 1.2 a client certificate with an Ed25519 or Ed448 key is asked
-- for with ecdsa_sign (RFC 8422 Section 3.1).  The Ed25519 and Ed448
-- certificate types of this library are synthetic and have no code point,
-- so they must not reach the wire.
handshake12_cert_request_types :: IO ()
handshake12_cert_request_types = do
    let cipher = cipher_ECDHE_RSA_WITH_AES_128_GCM_SHA256
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                ([TLS12], [TLS12])
                ([cipher], [cipher])
    ref <- newIORef Nothing
    let clientParam' =
            clientParam
                { clientHooks =
                    (clientHooks clientParam)
                        { onCertificateRequest = \_ -> return Nothing
                        }
                }
        serverParam' =
            serverParam
                { serverWantClientCert = True
                , serverSupported =
                    (serverSupported serverParam)
                        { supportedHashSignatures =
                            supportedHashSignatures defaultSupported
                        }
                , serverHooks =
                    (serverHooks serverParam)
                        { onClientCertificate = \_ -> return CertificateUsageAccept
                        }
                }
        record hs@(CertRequest certTypes _ _) = writeIORef ref (Just certTypes) >> return hs
        record hs = return hs
    withPairContextWith (id, id) (clientParam', serverParam') $ \(cctx, sctx) -> do
        contextHookSetHandshakeRecv cctx record
        concurrently_ (handshake sctx) (handshake cctx)
    mtypes <- readIORef ref
    case mtypes of
        Nothing -> expectationFailure "no CertificateRequest received"
        Just certTypes -> do
            certTypes `shouldSatisfy` all (`elem` defined)
            certTypes `shouldContain` [CertificateType_ECDSA_Sign]
  where
    defined =
        [ CertificateType_RSA_Sign
        , CertificateType_DSA_Sign
        , CertificateType_ECDSA_Sign
        ]

-- RFC 5246 Section 7.4.8: a TLS 1.2 CertificateVerify algorithm for
-- another type of key than the certificate's has a field that is
-- incorrect, an illegal_parameter, even when onUnverifiedClientCert would
-- accept a signature that does not verify.  One for an RSA key that still
-- does not fit it, RSASSA-PSS for an rsaEncryption key, is a signature
-- that does not verify, a decrypt_error.  The client signs with an
-- rsaEncryption key and the server reads the given algorithm.
handshake12_client_cert_verify_sigalg
    :: HashAndSignatureAlgorithm -> Selector TLSException -> IO ()
handshake12_client_cert_verify_sigalg alg rejected = do
    let cipher = cipher_ECDHE_RSA_WITH_AES_128_GCM_SHA256
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                ([TLS12], [TLS12])
                ([cipher], [cipher])
    cred <- generate $ arbitraryRSACredentialWithPurpose KeyUsagePurpose_ClientAuth
    let clientParam' =
            clientParam
                { clientHooks =
                    (clientHooks clientParam)
                        { onCertificateRequest = \_ -> return $ Just cred
                        }
                }
        acceptUnverified = alg == (HashSHA256, SignatureECDSA)
        serverParam' =
            serverParam
                { serverWantClientCert = True
                , serverHooks =
                    (serverHooks serverParam)
                        { onClientCertificate = \_ -> return CertificateUsageAccept
                        , onUnverifiedClientCert = return acceptUnverified
                        }
                }
        unfit (CertVerify (DigitallySigned _ sig)) =
            pure $ CertVerify (DigitallySigned alg sig)
        unfit hs = pure hs
    r <- timeout 10000000 $
        withPairContextWith (id, id) (clientParam', serverParam') $ \(cctx, sctx) -> do
            contextHookSetHandshakeRecv sctx unfit
            concurrently_
                (handshake sctx `shouldThrow` rejected)
                (void (E.try (handshake cctx) :: IO (Either TLSException ())))
    r `shouldSatisfy` isJust

rejectedAsDecryptError :: TLSException -> Bool
rejectedAsDecryptError (HandshakeFailed (Error_Protocol _ DecryptError)) = True
rejectedAsDecryptError _ = False

-- RFC 8446 Section 6.2: a CertificateVerify whose algorithm may not be used
-- with the certificate's key has a field that is incorrect, which is an
-- illegal_parameter; decrypt_error is for a signature that does not verify.
-- The client signs with RSA-PSS for an rsaEncryption key and the server
-- reads the algorithm as rsa_pss_pss_sha256, which needs an RSASSA-PSS key.
handshake13_client_cert_verify_unfit_sigalg :: IO ()
handshake13_client_cert_verify_unfit_sigalg = do
    let cipher = cipher13_AES_128_GCM_SHA256
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                ([TLS13], [TLS13])
                ([cipher], [cipher])
    cred <- generate $ arbitraryRSACredentialWithPurpose KeyUsagePurpose_ClientAuth
    let clientParam' =
            clientParam
                { clientHooks =
                    (clientHooks clientParam)
                        { onCertificateRequest = \_ -> return $ Just cred
                        }
                }
        serverParam' =
            serverParam
                { serverWantClientCert = True
                , serverHooks =
                    (serverHooks serverParam)
                        { onClientCertificate = \_ -> return CertificateUsageAccept
                        }
                }
        unfit (CertVerify13 (DigitallySigned _ sig)) =
            pure $
                CertVerify13 (DigitallySigned (HashIntrinsic, SignatureRSApsspssSHA256) sig)
        unfit hs = pure hs
    r <- timeout 10000000 $
        withPairContextWith (id, id) (clientParam', serverParam') $ \(cctx, sctx) -> do
            contextHookSetHandshake13Recv sctx unfit
            concurrently_
                ((handshake sctx >> recvData sctx) `shouldThrow` rejectedAsIllegalParameter)
                ( void
                    (E.try (handshake cctx >> recvData cctx) :: IO (Either TLSException B.ByteString))
                )
    r `shouldSatisfy` isJust

rejectedAsIllegalParameter :: TLSException -> Bool
rejectedAsIllegalParameter (HandshakeFailed (Error_Protocol _ IllegalParameter)) = True
rejectedAsIllegalParameter (Terminated _ _ (Error_Protocol _ IllegalParameter)) = True
rejectedAsIllegalParameter _ = False

-- A server that receives a ClientHello whose legacy_version is higher than
-- its own negotiates the highest version it supports (RFC 5246 Appendix
-- E.1); with supported_versions present it does not use legacy_version at
-- all (RFC 8446 Section 4.2.1).  The ClientHello's legacy_version is raised
-- on its way to the server, and the handshake must still reach the version
-- both sides support.
handshake_high_legacy_version :: Version -> Version -> IO ()
handshake_high_legacy_version version legacy = do
    let cipher
            | version == TLS13 = cipher13_AES_128_GCM_SHA256
            | otherwise = cipher_ECDHE_RSA_WITH_AES_128_GCM_SHA256
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                ([version], [version])
                ([cipher], [cipher])
    let raise (ClientHello ch) = pure $ ClientHello ch{chVersion = legacy}
        raise hs = pure hs
    withPairContextWith (id, id) (clientParam, serverParam) $ \(cctx, sctx) -> do
        contextHookSetHandshakeRecv sctx raise
        concurrently_ (handshake sctx) (handshake cctx)
        info <- contextGetInformation sctx
        (infoVersion <$> info) `shouldBe` Just version

-- RFC 8422 Section 5.1.2: ec_point_format_list is <1..2^8-1>, and a
-- client naming one of its curves in supported_groups must offer the
-- uncompressed format.  The ClientHello's ec_point_formats is replaced
-- on its way to the server, which must refuse it.
handshake12_ec_point_formats
    :: B.ByteString -> Selector TLSException -> IO ()
handshake12_ec_point_formats formats rejected = do
    CSP12 (cparams, sparams) <- generate arbitrary
    let cparams' =
            cparams
                { clientSupported =
                    (clientSupported cparams){supportedGroups = [X25519]}
                }
        sparams' =
            sparams
                { serverSupported =
                    (serverSupported sparams){supportedGroups = [X25519]}
                }
        replace (ClientHello ch) =
            pure $
                ClientHello
                    ch
                        { chExtensions =
                            filter
                                (\(ExtensionRaw eid _) -> eid /= EID_EcPointFormats)
                                (chExtensions ch)
                                ++ [ExtensionRaw EID_EcPointFormats formats]
                        }
        replace hs = pure hs
    r <- timeout 10000000 $
        withPairContextWith (id, id) (cparams', sparams') $ \(cctx, sctx) -> do
            contextHookSetHandshakeRecv sctx replace
            concurrently_
                (handshake sctx `shouldThrow` rejected)
                (void (E.try (handshake cctx) :: IO (Either TLSException ())))
    r `shouldSatisfy` isJust

rejectedAsDecodeError :: TLSException -> Bool
rejectedAsDecodeError (HandshakeFailed (Error_Protocol _ DecodeError)) = True
rejectedAsDecodeError _ = False

handshake_client_auth_fail :: (ClientParams, ServerParams) -> IO ()
handshake_client_auth_fail (clientParam, serverParam) = do
    let clientVersions = supportedVersions $ clientSupported clientParam
        serverVersions = supportedVersions $ serverSupported serverParam
        version = maximum (clientVersions `intersect` serverVersions)
    cred <- generate (arbitraryClientCredential version)
    let clientParam' =
            clientParam
                { clientHooks =
                    (clientHooks clientParam)
                        { onCertificateRequest = \_ -> return $ Just cred
                        }
                }
        serverParam' =
            serverParam
                { serverWantClientCert = True
                , serverHooks =
                    (serverHooks serverParam)
                        { onClientCertificate = validateChain cred
                        }
                }
    runTLSFailure (clientParam', serverParam') handshake handshake
  where
    validateChain _ _ = return (CertificateUsageReject CertificateRejectUnknownCA)

--------------------------------------------------------------

handshake_ems :: (EMSMode, EMSMode) -> IO ()
handshake_ems (cems, sems) = do
    params <- generate arbitrary
    let params' = setEMSMode (cems, sems) params
        version = getConnectVersion params'
        emsVersion = version >= TLS10 && version <= TLS12
        use = cems /= NoEMS && sems /= NoEMS
        require = cems == RequireEMS || sems == RequireEMS
        p info = infoExtendedMainSecret info == (emsVersion && use)
    if emsVersion && require && not use
        then runTLSFailure params' handshake handshake
        else runTLSPredicate params' (maybe False p)

newtype CompatEMS = CompatEMS (EMSMode, EMSMode) deriving (Show)

instance Arbitrary CompatEMS where
    arbitrary = CompatEMS <$> (arbitrary `suchThat` compatible)
      where
        compatible (NoEMS, RequireEMS) = False
        compatible (RequireEMS, NoEMS) = False
        compatible _ = True

handshake_resumption_ems :: (CompatEMS, CompatEMS) -> IO ()
handshake_resumption_ems (CompatEMS ems, CompatEMS ems2) = do
    sessionRefs <- twoSessionRefs
    let sessionManagers = twoSessionManagers sessionRefs

    plainParams <- generate arbitrary
    let params =
            setEMSMode ems $
                setPairParamsSessionManagers sessionManagers plainParams

    runTLSSimple params

    -- and resume
    sessionParams <- readClientSessionRef sessionRefs
    expectJust "session param should be Just" sessionParams
    let params2 =
            setEMSMode ems2 $
                setPairParamsSessionResuming (fromJust sessionParams) params

    let version = getConnectVersion params2
        emsVersion = version >= TLS10 && version <= TLS12

    if emsVersion && use ems && not (use ems2)
        then runTLSFailure params2 handshake handshake
        else do
            runTLSSimple params2
            mSessionParams2 <- readClientSessionRef sessionRefs
            let sameSession = sessionParams == mSessionParams2
                sameUse = use ems == use ems2
            when emsVersion (sameSession `shouldBe` sameUse)
  where
    use (NoEMS, _) = False
    use (_, NoEMS) = False
    use _ = True

--------------------------------------------------------------

handshake_alpn :: (ClientParams, ServerParams) -> IO ()
handshake_alpn (clientParam, serverParam) = do
    let clientParam' =
            clientParam
                { clientHooks =
                    (clientHooks clientParam)
                        { onSuggestALPN = return $ Just ["h2", "http/1.1"]
                        }
                }
        serverParam' =
            serverParam
                { serverHooks =
                    (serverHooks serverParam)
                        { onALPNClientSuggest = Just alpn
                        }
                }
        params' = (clientParam', serverParam')
    runTLSSuccess params' hsClient hsServer
  where
    hsClient ctx = do
        handshake ctx
        proto <- getNegotiatedProtocol ctx
        proto `shouldBe` Just "h2"
    hsServer ctx = do
        handshake ctx
        proto <- getNegotiatedProtocol ctx
        proto `shouldBe` Just "h2"
    alpn xs
        | "h2" `elem` xs = return "h2"
        | otherwise = return "http/1.1"

handshake_alpn_rejects_unoffered_server_selection :: IO ()
handshake_alpn_rejects_unoffered_server_selection = do
    (clientParam, serverParam) <- generate arbitraryPairParams13
    let params = alpnParams clientParam serverParam (const $ pure "h2")
    withPairContextWith (id, id) params $ \(cctx, sctx) ->
        concurrently_
            (handshake sctx `shouldThrow` serverRejectedUnofferedALPN)
            (handshake cctx `shouldThrow` anyTLSException)

handshake_alpn_rejects_unoffered_client_selection :: IO ()
handshake_alpn_rejects_unoffered_client_selection = do
    (clientParam, serverParam) <- generate arbitraryPairParams13
    let params = alpnParams clientParam serverParam (pure . unsafeHead)
    withPairContextWith (id, id) params $ \(cctx, sctx) -> do
        contextHookSetHandshake13Recv cctx tamperALPN
        concurrently_
            (handshake sctx `shouldThrow` anyTLSException)
            (handshake cctx `shouldThrow` clientRejectedUnofferedALPN)
  where
    tamperALPN (EncryptedExtensions13 exts) =
        pure $ EncryptedExtensions13 $ map replaceALPN exts
    tamperALPN hs = pure hs
    replaceALPN ext@(ExtensionRaw eid _)
        | eid == EID_ApplicationLayerProtocolNegotiation =
            toExtensionRaw $ ApplicationLayerProtocolNegotiation ["h2"]
        | otherwise = ext

alpnParams
    :: ClientParams
    -> ServerParams
    -> ([B.ByteString] -> IO B.ByteString)
    -> (ClientParams, ServerParams)
alpnParams clientParam serverParam select =
    ( clientParam
        { clientHooks =
            (clientHooks clientParam)
                { onSuggestALPN = pure $ Just ["http/1.1"]
                }
        }
    , serverParam
        { serverHooks =
            (serverHooks serverParam)
                { onALPNClientSuggest = Just select
                }
        }
    )

serverRejectedUnofferedALPN :: TLSException -> Bool
serverRejectedUnofferedALPN (HandshakeFailed (Error_Protocol msg alert)) =
    msg == "ALPN callback selected a protocol not offered by the client"
        && alert == NoApplicationProtocol
serverRejectedUnofferedALPN _ = False

clientRejectedUnofferedALPN :: TLSException -> Bool
clientRejectedUnofferedALPN (HandshakeFailed (Error_Protocol msg alert)) =
    msg == "server selected an ALPN protocol not offered by the client"
        && alert == IllegalParameter
clientRejectedUnofferedALPN _ = False

-- RFC 6066 Section 3: HostName is <1..2^16-1>, so a client with an
-- empty server name sends no server_name, and the server sees none.
handshake_sni_empty :: IO ()
handshake_sni_empty = do
    CSP12 (clientParam, serverParam) <- generate arbitrary
    let clientParam' = clientParam{clientServerIdentification = ("", "")}
    runTLSSuccess (clientParam', serverParam) hs hs
  where
    hs ctx = do
        handshake ctx
        msni <- getClientSNI ctx
        msni `shouldBe` Nothing

-- RFC 6066 Section 3: the server_name_list MUST NOT contain more than one
-- name of the same name_type, and HostName is an ASCII DNS host name.  The
-- ClientHello's server_name is replaced on its way to the server, which
-- must refuse it with illegal_parameter.
handshake_sni_illegal :: [B.ByteString] -> IO ()
handshake_sni_illegal names = do
    CSP12 (clientParam, serverParam) <- generate arbitrary
    let entry name =
            B.concat [B.pack [0, len `shiftR` 8, len], name]
          where
            len = fromIntegral $ B.length name
        list = B.concat $ map entry names
        listLen = fromIntegral $ B.length list
        sni = B.concat [B.pack [listLen `shiftR` 8, listLen], list]
        replace (ClientHello ch) =
            pure $
                ClientHello
                    ch
                        { chExtensions =
                            ExtensionRaw EID_ServerName sni
                                : filter
                                    (\(ExtensionRaw eid _) -> eid /= EID_ServerName)
                                    (chExtensions ch)
                        }
        replace hs = pure hs
    r <- timeout 10000000 $
        withPairContextWith (id, id) (clientParam, serverParam) $ \(cctx, sctx) -> do
            contextHookSetHandshakeRecv sctx replace
            concurrently_
                (handshake sctx `shouldThrow` rejectedAsIllegalParameter)
                (void (E.try (handshake cctx) :: IO (Either TLSException ())))
    r `shouldSatisfy` isJust

handshake_sni :: (ClientParams, ServerParams) -> IO ()
handshake_sni (clientParam, serverParam) = do
    ref <- newIORef Nothing
    let clientParam' =
            clientParam
                { clientServerIdentification = (serverName, "")
                }
        serverParam' =
            serverParam
                { serverHooks =
                    (serverHooks serverParam)
                        { onServerNameIndication = onSNI ref
                        }
                }
        params' = (clientParam', serverParam')
    runTLSSuccess params' hsClient hsServer
    receivedName <- readIORef ref
    receivedName `shouldBe` Just (Just serverName)
  where
    hsClient ctx = do
        handshake ctx
        msni <- getClientSNI ctx
        expectMaybe "C: SNI should be Just" serverName msni
    hsServer ctx = do
        handshake ctx
        msni <- getClientSNI ctx
        expectMaybe "S: SNI should be Just" serverName msni
    onSNI ref name = do
        mx <- readIORef ref
        mx `shouldBe` Nothing
        writeIORef ref (Just name)
        return (Credentials [])
    serverName = "haskell.org"

--------------------------------------------------------------

newtype CSP12 = CSP12 (ClientParams, ServerParams) deriving (Show)

instance Arbitrary CSP12 where
    arbitrary = CSP12 <$> arbitraryPairParams12

handshake12_renegotiation :: CSP12 -> IO ()
handshake12_renegotiation (CSP12 (cparams, sparams)) = do
    renegDisabled <- generate arbitrary
    let sparams' =
            sparams
                { serverSupported =
                    (serverSupported sparams)
                        { supportedClientInitiatedRenegotiation = not renegDisabled
                        }
                }
    if renegDisabled
        then runTLSFailure (cparams, sparams') hsClient hsServer
        else runTLSSimple (cparams, sparams')
  where
    hsClient ctx = handshake ctx >> handshake ctx
    -- recvData receives the alert from the second handshake
    hsServer ctx = handshake ctx >> void (recvData ctx)

-- RFC 5746 Section 3.7: when a connection with secure renegotiation is
-- renegotiated, ClientHello must not contain the SCSV and must contain
-- the renegotiation_info extension.  The second ClientHello is tampered
-- with on its way to the server, which must abort with handshake_failure.
handshake12_renegotiation_tampered :: (ClientHello -> ClientHello) -> IO ()
handshake12_renegotiation_tampered tamper = do
    CSP12 (cparams, sparams) <- generate arbitrary
    let cparams' =
            cparams
                { clientSupported =
                    (clientSupported cparams)
                        { supportedSecureRenegotiation = True
                        }
                }
        sparams' =
            sparams
                { serverSupported =
                    (serverSupported sparams)
                        { supportedSecureRenegotiation = True
                        , supportedClientInitiatedRenegotiation = True
                        }
                }
    count <- newIORef (0 :: Int)
    let tamperSecond (ClientHello ch) = do
            n <- atomicModifyIORef' count $ \i -> (i + 1, i)
            pure $ ClientHello $ if n == 0 then ch else tamper ch
        tamperSecond hs = pure hs
    r <- timeout 10000000 $
        withPairContextWith (id, id) (cparams', sparams') $ \(cctx, sctx) -> do
            contextHookSetHandshakeRecv sctx tamperSecond
            concurrently_ (handshake sctx) (handshake cctx)
            concurrently_
                (recvData sctx `shouldThrow` rejectedAsHandshakeFailure)
                ( void
                    (E.try (handshake cctx >> recvData cctx) :: IO (Either TLSException B.ByteString))
                )
    r `shouldSatisfy` isJust

rejectedAsHandshakeFailure :: TLSException -> Bool
rejectedAsHandshakeFailure (HandshakeFailed (Error_Protocol _ HandshakeFailure)) = True
rejectedAsHandshakeFailure (Terminated _ _ (Error_Protocol _ HandshakeFailure)) = True
rejectedAsHandshakeFailure _ = False

handshake12_session_resumption :: CSP12 -> IO ()
handshake12_session_resumption (CSP12 plainParams) = do
    sessionRefs <- twoSessionRefs
    let sessionManagers = twoSessionManagers sessionRefs

    let params = setPairParamsSessionManagers sessionManagers plainParams

    runTLSSimple params

    -- and resume
    sessionParams <- readClientSessionRef sessionRefs
    expectJust "session param should be Just" sessionParams
    let params2 = setPairParamsSessionResuming (fromJust sessionParams) params

    runTLSPredicate params2 (maybe False infoTLS12Resumption)

-- RFC 5246 Section 7.4.1.2: a client resuming a session MUST offer its
-- cipher suite.  When it asks to resume one and offers none the server can
-- use, that is what is reported -- illegal_parameter, as when a cipher could
-- be chosen -- rather than the missing common cipher.
handshake12_session_resumption_cipher_missing :: CSP12 -> IO ()
handshake12_session_resumption_cipher_missing (CSP12 plainParams) = do
    sessionRefs <- twoSessionRefs
    let sessionManagers = twoSessionManagers sessionRefs
        params = setPairParamsSessionManagers sessionManagers plainParams
    runTLSSimple params
    sessionParams <- readClientSessionRef sessionRefs
    expectJust "session param should be Just" sessionParams
    let params2 = setPairParamsSessionResuming (fromJust sessionParams) params
        -- TLS_RSA_WITH_NULL_MD5, which the server does not support
        nullCiphers (ClientHello ch) = pure $ ClientHello ch{chCiphers = [CipherId 0x0001]}
        nullCiphers hs = pure hs
    withPairContextWith (id, id) params2 $ \(cctx, sctx) -> do
        contextHookSetHandshakeRecv sctx nullCiphers
        concurrently_
            (handshake sctx `shouldThrow` serverRejectedMissingCipher)
            (handshake cctx `shouldThrow` anyTLSException)

serverRejectedMissingCipher :: TLSException -> Bool
serverRejectedMissingCipher (HandshakeFailed (Error_Protocol _ IllegalParameter)) = True
serverRejectedMissingCipher _ = False

-- RFC 5077 Sections 3.2 and 3.3: a server sends the session_ticket
-- extension and NewSessionTicket only to a client that sent the
-- extension.  The client's session_ticket is removed on its way to a
-- server whose session manager uses tickets, and the client must then
-- receive neither.
handshake12_session_ticket_unoffered :: IO ()
handshake12_session_ticket_unoffered = do
    CSP12 (cparams, sparams) <- generate arbitrary
    let sparams' =
            sparams
                { serverShared =
                    (serverShared sparams){sharedSessionManager = oneSessionTicket}
                }
        unoffer (ClientHello ch) =
            pure $
                ClientHello
                    ch
                        { chExtensions =
                            filter
                                (\(ExtensionRaw eid _) -> eid /= EID_SessionTicket)
                                (chExtensions ch)
                        }
        unoffer hs = pure hs
    received <- newIORef []
    let record hs = modifyIORef received (hs :) >> pure hs
    withPairContextWith (id, id) (cparams, sparams') $ \(cctx, sctx) -> do
        contextHookSetHandshakeRecv sctx unoffer
        contextHookSetHandshakeRecv cctx record
        concurrently_ (handshake sctx) (handshake cctx)
    hss <- readIORef received
    let ticketExt (ServerHello sh) =
            any (\(ExtensionRaw eid _) -> eid == EID_SessionTicket) (shExtensions sh)
        ticketExt _ = False
        newTicket NewSessionTicket{} = True
        newTicket _ = False
    filter (\h -> ticketExt h || newTicket h) hss `shouldBe` []

handshake12_session_ticket :: CSP12 -> IO ()
handshake12_session_ticket (CSP12 plainParams) = do
    sessionRefs <- twoSessionRefs
    let sessionManagers0 = twoSessionManagers sessionRefs
        sessionManagers = (fst sessionManagers0, oneSessionTicket)

    let params = setPairParamsSessionManagers sessionManagers plainParams

    runTLSSimple params

    -- and resume
    sessionParams <- readClientSessionRef sessionRefs
    expectJust "session param should be Just" sessionParams
    let params2 = setPairParamsSessionResuming (fromJust sessionParams) params

    runTLSPredicate params2 (maybe False infoTLS12Resumption)

--------------------------------------------------------------

handshake13_full :: CSP13 -> IO ()
handshake13_full (CSP13 (cli, srv)) = do
    let cliSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [X25519]
                }
        svrSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [X25519]
                , supportedGroupsTLS13 = [[X25519]]
                }
        params =
            ( cli{clientSupported = cliSupported}
            , srv{serverSupported = svrSupported}
            )
    runTLSSimple13 params FullHandshake

handshake13_hrr :: CSP13 -> IO ()
handshake13_hrr (CSP13 (cli, srv)) = do
    let cliSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [P256, X25519]
                }
        svrSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [X25519]
                , supportedGroupsTLS13 = [[X25519]]
                }
        params =
            ( cli{clientSupported = cliSupported}
            , srv{serverSupported = svrSupported}
            )
    runTLSSimple13 params HelloRetryRequest

handshake13_psk :: CSP13 -> IO ()
handshake13_psk (CSP13 (cli, srv)) = do
    let cliSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [P256, X25519]
                }
        svrSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [X25519]
                , supportedGroupsTLS13 = [[X25519]]
                }
        params0 =
            ( cli{clientSupported = cliSupported}
            , srv{serverSupported = svrSupported}
            )

    sessionRefs <- twoSessionRefs
    let sessionManagers = twoSessionManagers sessionRefs

    let params = setPairParamsSessionManagers sessionManagers params0

    runTLSSimple13 params HelloRetryRequest

    -- and resume
    sessionParams <- readClientSessionRef sessionRefs
    expectJust "session param should be Just" sessionParams
    let params2 = setPairParamsSessionResuming (fromJust sessionParams) params

    runTLSSimple13 params2 PreSharedKey

-- RFC 8446 Section 4.6.1: a TLS 1.3 PSK resumes only a TLS 1.3 session.
-- The server's session manager answers the client's PSK identity with a
-- TLS 1.2 session, which has no ticket information; the server must
-- fall back to a full handshake rather than fail.
handshake13_psk_tls12_session :: IO ()
handshake13_psk_tls12_session = do
    CSP12 params12 <- generate arbitrary
    refs12 <- twoSessionRefs
    runTLSSimple $ setPairParamsSessionManagers (twoSessionManagers refs12) params12
    Just (_, sdata12) <- readIORef (snd refs12)
    sessionVersion sdata12 `shouldBe` TLS12

    CSP13 (cli, srv) <- generate arbitrary
    let cliSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [P256, X25519]
                }
        svrSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [X25519]
                , supportedGroupsTLS13 = [[X25519]]
                }
        params0 =
            ( cli{clientSupported = cliSupported}
            , srv{serverSupported = svrSupported}
            )
    refs13 <- twoSessionRefs
    let params = setPairParamsSessionManagers (twoSessionManagers refs13) params0
    runTLSSimple13 params HelloRetryRequest
    Just sessionParams <- readClientSessionRef refs13

    let tls12Manager =
            noSessionManager
                { sessionResume = \_ -> return $ Just sdata12
                , sessionResumeOnlyOnce = \_ -> return $ Just sdata12
                }
        params2 =
            setPairParamsSessionResuming sessionParams $
                setPairParamsSessionManagers (fst (twoSessionManagers refs13), tls12Manager) params0
    -- The key share is for the group of the earlier session, so no
    -- HelloRetryRequest, and the PSK is not used.
    runTLSSimple13 params2 FullHandshake

handshake13_psk_ticket :: CSP13 -> IO ()
handshake13_psk_ticket (CSP13 (cli, srv)) = do
    let cliSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [P256, X25519]
                }
        svrSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [X25519]
                , supportedGroupsTLS13 = [[X25519]]
                }
        params0 =
            ( cli{clientSupported = cliSupported}
            , srv{serverSupported = svrSupported}
            )

    sessionRefs <- twoSessionRefs
    let sessionManagers0 = twoSessionManagers sessionRefs
        sessionManagers = (fst sessionManagers0, oneSessionTicket)

    let params = setPairParamsSessionManagers sessionManagers params0

    runTLSSimple13 params HelloRetryRequest

    -- and resume
    sessionParams <- readClientSessionRef sessionRefs
    expectJust "session param should be Just" sessionParams
    let params2 = setPairParamsSessionResuming (fromJust sessionParams) params

    runTLSSimple13 params2 PreSharedKey

handshake13_psk_fallback :: CSP13 -> IO ()
handshake13_psk_fallback (CSP13 (cli, srv)) = do
    let cliSupported =
            defaultSupported
                { supportedCiphers =
                    [ cipher13_AES_128_GCM_SHA256
                    , cipher13_AES_128_CCM_SHA256
                    ]
                , supportedGroups = [P256, X25519]
                }
        svrSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [X25519]
                , supportedGroupsTLS13 = [[X25519]]
                }
        params0 =
            ( cli{clientSupported = cliSupported}
            , srv{serverSupported = svrSupported}
            )

    sessionRefs <- twoSessionRefs
    let sessionManagers = twoSessionManagers sessionRefs

    let params = setPairParamsSessionManagers sessionManagers params0

    runTLSSimple13 params HelloRetryRequest

    -- resumption fails because GCM cipher is not supported anymore, full
    -- handshake is not possible because X25519 has been removed, so we are
    -- back with P256 after hello retry
    sessionParams <- readClientSessionRef sessionRefs
    expectJust "session param should be Just" sessionParams
    let (cli2, srv2) = setPairParamsSessionResuming (fromJust sessionParams) params
        srv2' =
            srv2{serverSupported = svrSupported'}
        svrSupported' =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_CCM_SHA256]
                , supportedGroups = [P256]
                , supportedGroupsTLS13 = [[P256]]
                }

    runTLSSimple13 (cli2, srv2') HelloRetryRequest

handshake13_0rtt :: CSP13 -> IO ()
handshake13_0rtt (CSP13 (cli, srv)) = do
    let cliSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [P256, X25519]
                }
        svrSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [X25519]
                , supportedGroupsTLS13 = [[X25519]]
                }
        cliHooks =
            defaultClientHooks
                { onSuggestALPN = return $ Just ["h2"]
                }
        svrHooks =
            defaultServerHooks
                { onALPNClientSuggest = Just (return . unsafeHead)
                }
        params0 =
            ( cli
                { clientSupported = cliSupported
                , clientHooks = cliHooks
                }
            , srv
                { serverSupported = svrSupported
                , serverHooks = svrHooks
                , serverEarlyDataSize = 2048
                }
            )

    sessionRefs <- twoSessionRefs
    let sessionManagers = twoSessionManagers sessionRefs

    let params = setPairParamsSessionManagers sessionManagers params0

    runTLSSimple13 params HelloRetryRequest
    runTLS0rtt params sessionRefs
    runTLS0rtt params sessionRefs
  where
    runTLS0rtt params sessionRefs = do
        -- and resume
        sessionParams <- readClientSessionRef sessionRefs
        expectJust "session param should be Just" sessionParams
        clearClientSessionRef sessionRefs
        earlyData <- B.pack <$> generate (someWords8 256)
        let (pc, ps) = setPairParamsSessionResuming (fromJust sessionParams) params
            params2 = (pc{clientUseEarlyData = True}, ps)

        runTLS0RTT params2 RTT0 earlyData

handshake13_0rtt_alpn :: IO ()
handshake13_0rtt_alpn = do
    (cli, srv) <- generate arbitraryPairParams13
    let cliSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [X25519]
                }
        svrSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [X25519]
                , supportedGroupsTLS13 = [[X25519]]
                }
        cliHooks =
            defaultClientHooks
                { onSuggestALPN = return $ Just ["h2"]
                }
        svrHooks =
            defaultServerHooks
                { onALPNClientSuggest = Just (return . unsafeHead)
                }
        params0 =
            ( cli
                { clientSupported = cliSupported
                , clientHooks = cliHooks
                }
            , srv
                { serverSupported = svrSupported
                , serverHooks = svrHooks
                , serverEarlyDataSize = 2048
                }
            )
    sessionRefs <- twoSessionRefs
    let params =
            setPairParamsSessionManagers
                (twoSessionManagers sessionRefs)
                params0
    runTLSSimple13 params FullHandshake

    sessionParams <- readClientSessionRef sessionRefs
    expectJust "session param should be Just" sessionParams
    sessionALPN (snd $ fromJust sessionParams) `shouldBe` Just "h2"
    let (pc, ps) = setPairParamsSessionResuming (fromJust sessionParams) params
        pc' =
            pc
                { clientUseEarlyData = True
                , clientHooks =
                    (clientHooks pc)
                        { onSuggestALPN = return $ Just ["http/1.1"]
                        }
                }
        ps' =
            ps
                { serverHooks =
                    (serverHooks ps)
                        { onALPNClientSuggest = Just (return . unsafeHead)
                        }
                }
    runTLS0RTT (pc', ps') PreSharedKey "GET /admin HTTP/1.1\r\n\r\n"

handshake13_0rtt_fallback :: CSP13 -> IO ()
handshake13_0rtt_fallback (CSP13 (cli, srv)) = do
    group0 <- generate $ elements [P256, X25519]
    let cliSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [P256, X25519]
                }
        svrSupported =
            defaultSupported
                { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                , supportedGroups = [group0]
                , supportedGroupsTLS13 = [[group0]]
                }
        params =
            ( cli{clientSupported = cliSupported}
            , srv
                { serverSupported = svrSupported
                , serverEarlyDataSize = 1024
                }
            )

    sessionRefs <- twoSessionRefs
    let sessionManagers = twoSessionManagers sessionRefs

    let params0 = setPairParamsSessionManagers sessionManagers params

    let mode = if group0 == P256 then FullHandshake else HelloRetryRequest
    runTLSSimple13 params0 mode

    -- and resume
    mSessionParams <- readClientSessionRef sessionRefs
    case mSessionParams of
        Nothing -> expectationFailure "session params: Just is expected"
        Just sessionParams -> do
            earlyData <- B.pack <$> generate (someWords8 256)
            group1 <- generate $ elements [P256, X25519]
            let (pc, ps) = setPairParamsSessionResuming sessionParams params0
                svrSupported1 =
                    defaultSupported
                        { supportedCiphers = [cipher13_AES_128_GCM_SHA256]
                        , supportedGroups = [group1]
                        , supportedGroupsTLS13 = [[group1]]
                        }
                params1 =
                    ( pc{clientUseEarlyData = True}
                    , ps
                        { serverEarlyDataSize = 0
                        , serverSupported = svrSupported1
                        }
                    )
            -- C: [P256, X25519]
            -- S: [group0]
            -- C: [P256, X25519]
            -- S: [group1]
            if group0 == group1
                -- 0-RTT is not allowed, so fallback to PreSharedKey
                then runTLS0RTT params1 PreSharedKey earlyData
                -- HRR but not allowed for 0-RTT
                else runTLSFailure params1 (tlsClient earlyData) tlsServer
  where
    tlsClient earlyData ctx = do
        handshake ctx
        sendData ctx $ L.fromStrict earlyData
        _ <- recvData ctx
        bye ctx
    tlsServer ctx = do
        handshake ctx
        _ <- recvData ctx
        bye ctx

handshake13_ee_groups :: CSP13 -> IO ()
handshake13_ee_groups (CSP13 (cli, srv)) = do
    let -- The client prefers P256
        cliSupported = (clientSupported cli){supportedGroups = [P256, X25519]}
        -- The server prefers X25519
        svrSupported =
            (serverSupported srv)
                { supportedGroups = [X25519, P256]
                , supportedGroupsTLS13 = [[X25519, P256]]
                }
        params =
            ( cli{clientSupported = cliSupported}
            , srv{serverSupported = svrSupported}
            )
    (_, serverMessages) <- runTLSCapture13 params
    -- The server should tell X25519 in supported_groups in EE to client
    let isSupportedGroups (ExtensionRaw eid _) = eid == EID_SupportedGroups
        eeMessagesHaveExt =
            [ any isSupportedGroups exts
            | EncryptedExtensions13 exts <- serverMessages
            ]
    eeMessagesHaveExt `shouldBe` [True]

handshake13_ec :: CSP13 -> IO ()
handshake13_ec (CSP13 (cli, srv)) = do
    EC cgrps <- generate arbitrary
    EC sgrps <- generate arbitrary
    let cliSupported = (clientSupported cli){supportedGroups = cgrps}
        svrSupported =
            (serverSupported srv)
                { supportedGroups = sgrps
                , supportedGroupsTLS13 = [sgrps]
                }
        params =
            ( cli{clientSupported = cliSupported}
            , srv{serverSupported = svrSupported}
            )
    runTLSSimple13 params FullHandshake

handshake13_ffdhe :: CSP13 -> IO ()
handshake13_ffdhe (CSP13 (cli, srv)) = do
    FFDHE cgrps <- generate arbitrary
    FFDHE sgrps <- generate arbitrary
    let cliSupported = (clientSupported cli){supportedGroups = cgrps}
        svrSupported =
            (serverSupported srv)
                { supportedGroups = sgrps
                , supportedGroupsTLS13 = [sgrps]
                }
        params =
            ( cli{clientSupported = cliSupported}
            , srv{serverSupported = svrSupported}
            )
    runTLSSimple13 params FullHandshake

-- An all-zero X25519 public key decodes, but the shared secret computed
-- from it is all zero and is rejected.  The server must answer with
-- illegal_parameter, as it does for X25519 alone, rather than crash.
handshake13_x25519mlkem768_zero_x25519 :: IO ()
handshake13_x25519mlkem768_zero_x25519 = do
    CSP13 (cli, srv) <- generate arbitrary
    let cliSupported =
            (clientSupported cli){supportedGroups = [X25519MLKEM768]}
        svrSupported =
            (serverSupported srv)
                { supportedGroups = [X25519MLKEM768]
                , supportedGroupsTLS13 = [[X25519MLKEM768]]
                }
        params =
            ( cli{clientSupported = cliSupported}
            , srv{serverSupported = svrSupported}
            )
    withPairContextWith (id, id) params $ \(cctx, sctx) -> do
        contextHookSetHandshakeRecv sctx zeroX25519
        concurrently_
            (handshake sctx `shouldThrow` serverRejectedZeroX25519)
            (handshake cctx `shouldThrow` anyTLSException)
  where
    zeroX25519 (ClientHello ch) =
        pure $ ClientHello ch{chExtensions = map zeroKeyShare $ chExtensions ch}
    zeroX25519 hs = pure hs
    zeroKeyShare ext@(ExtensionRaw eid bs)
        | eid == EID_KeyShare
        , Just (KeyShareClientHello kses) <- extensionDecode MsgTClientHello bs =
            toExtensionRaw $ KeyShareClientHello $ map zeroEntry kses
        | otherwise = ext
    -- The ML-KEM-768 encapsulation key (1184 bytes) is followed by the
    -- X25519 public key (32 bytes).
    zeroEntry (KeyShareEntry grp key)
        | grp == X25519MLKEM768 =
            KeyShareEntry grp $ B.take 1184 key <> B.replicate 32 0
    zeroEntry kse = kse

serverRejectedZeroX25519 :: TLSException -> Bool
serverRejectedZeroX25519 (HandshakeFailed (Error_Protocol msg alert)) =
    msg == "invalid client X25519MLKEM768 public key"
        && alert == IllegalParameter
serverRejectedZeroX25519 _ = False

post_handshake_auth :: CSP13 -> IO ()
post_handshake_auth (CSP13 (clientParam, serverParam)) = do
    cred <- generate (arbitraryClientCredential TLS13)
    let clientParam' =
            clientParam
                { clientHooks =
                    (clientHooks clientParam)
                        { onCertificateRequest = \_ -> return $ Just cred
                        }
                }
        serverParam' =
            serverParam
                { serverHooks =
                    (serverHooks serverParam)
                        { onClientCertificate = validateChain cred
                        }
                }
    if isCredentialDSA cred
        then runTLSFailure (clientParam', serverParam') hsClient hsServer
        else runTLSSuccess (clientParam', serverParam') hsClient hsServer
  where
    validateChain cred chain
        | chain == fst cred = return CertificateUsageAccept
        | otherwise = return (CertificateUsageReject CertificateRejectUnknownCA)
    hsClient ctx = do
        handshake ctx
        sendData ctx "request 1"
        recvDataAssert ctx "response 1"
        sendData ctx "request 2"
        recvDataAssert ctx "response 2"
    hsServer ctx = do
        handshake ctx
        recvDataAssert ctx "request 1"
        _ <- requestCertificate ctx -- single request
        sendData ctx "response 1"
        recvDataAssert ctx "request 2"
        _ <- requestCertificate ctx
        _ <- requestCertificate ctx -- two simultaneously
        sendData ctx "response 2"

-- | After sending a Certificate message, a TLS 1.3 client peeks for a
-- client-authentication alert with a deadline of a few RTTs.  That peek must
-- not abandon a record it has already started reading: the record layer has no
-- receive buffer, so the bytes consumed for the record header would be lost and
-- the caller's next 'recvData' would decode part of a record body as a header.
--
-- Here the server's first write is one full-size record whose body is made to
-- arrive late, so the peek does hit its deadline with the header already
-- consumed.  Before the fix this failed with
-- @Error_Protocol "record exceeding maximum size" RecordOverflow@.
handshake13_client_auth_slow_record :: IO ()
handshake13_client_auth_slow_record = do
    (clientParam, serverParam) <- generate arbitraryPairParams13
    cred <- generate (arbitraryClientCredential TLS13)
    let clientParam' =
            clientParam
                { clientHooks =
                    (clientHooks clientParam)
                        { onCertificateRequest = \_ -> return $ Just cred
                        }
                }
        serverParam' =
            serverParam
                { serverWantClientCert = True
                , serverHooks =
                    (serverHooks serverParam)
                        { onClientCertificate = \_ -> return CertificateUsageAccept
                        }
                }
        payload = B.replicate 16384 65
    withPairContextWith (delayBigReads, id) (clientParam', serverParam') $
        \(cCtx, sCtx) ->
            concurrently_
                ( do
                    handshake sCtx
                    sendData sCtx $ L.fromStrict payload
                )
                ( do
                    handshake cCtx
                    recvData cCtx `shouldReturn` payload
                )
  where
    -- Only the body of the big record is held back; every handshake record is
    -- far smaller than this threshold and so arrives immediately.
    delayBigReads be =
        be
            { backendRecv = \n -> do
                when (n > 4096) $ threadDelay 300000
                backendRecv be n
            }

-- A protected record too short for the cipher cannot be deprotected, which
-- RFC 5246 Section 7.2.2 and RFC 8446 Section 5.2 answer with
-- bad_record_mac.  The client's first application data record is cut down
-- to a single byte of ciphertext on its way to the server.
short_record_bad_record_mac :: Version -> Cipher -> IO ()
short_record_bad_record_mac version cipher = do
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                ([version], [version])
                ([cipher], [cipher])
    armed <- newIORef False
    let truncateRecord be =
            be
                { backendSend = \bs -> do
                    cut <- readIORef armed
                    if cut && B.length bs > 6 && B.head bs == 23
                        then do
                            writeIORef armed False
                            backendSend be $ B.take 3 bs <> B.pack [0, 1] <> B.take 1 (B.drop 5 bs)
                        else backendSend be bs
                }
    withPairContextWith (truncateRecord, id) (clientParam, serverParam) $
        \(cctx, sctx) ->
            concurrently_
                ( do
                    handshake sctx
                    recvData sctx `shouldThrow` serverRejectedShortRecord
                )
                ( do
                    handshake cctx
                    writeIORef armed True
                    sendData cctx "hello"
                    void (E.try (recvData cctx) :: IO (Either E.SomeException B.ByteString))
                )

serverRejectedShortRecord :: TLSException -> Bool
serverRejectedShortRecord (Terminated _ _ (Error_Protocol _ BadRecordMac)) = True
serverRejectedShortRecord _ = False

-- The first handshake message a server receives must be a ClientHello; any
-- other is out of order and answered with unexpected_message (RFC 8446
-- Section 4), whatever its body would decode to.  The handshake type of the
-- client's first record is replaced on its way to the server.
server_first_message_unexpected :: Word8 -> IO ()
server_first_message_unexpected ty = do
    (clientParam, serverParam) <- generate arbitrary
    armed <- newIORef True
    let retype be =
            be
                { backendSend = \bs -> do
                    first <- atomicModifyIORef' armed (\a -> (False, a))
                    if first && B.length bs > 5 && B.head bs == 22
                        then backendSend be $ B.take 5 bs <> B.singleton ty <> B.drop 6 bs
                        else backendSend be bs
                }
    withPairContextWith (retype, id) (clientParam, serverParam) $ \(cctx, sctx) ->
        concurrently_
            (handshake sctx `shouldThrow` serverRejectedUnexpectedFirst)
            (handshake cctx `shouldThrow` anyTLSException)

-- ChangeCipherSpec comes between complete handshake messages (RFC 5246
-- Section 7.1), and a server has no handshake state before its first
-- ClientHello (CVE-2004-0079).  The client's first record is split after
-- the given number of bytes of its body, and a ChangeCipherSpec record is
-- sent in between; with 0 it is sent before the whole ClientHello.
server_ccs_interleaved :: Int -> IO ()
server_ccs_interleaved n = do
    (clientParam, serverParam) <- generate arbitrary
    armed <- newIORef True
    let interleave be =
            be
                { backendSend = \bs -> do
                    first <- atomicModifyIORef' armed (\a -> (False, a))
                    if first && B.length bs > 5 + n && B.head bs == 22
                        then do
                            let (hdr, body) = B.splitAt 5 bs
                                (body1, body2) = B.splitAt n body
                                record b = B.take 3 hdr <> encodeWord16 (fromIntegral $ B.length b) <> b
                                ccs = B.pack [20, 3, 3, 0, 1, 1]
                            if n == 0
                                then backendSend be $ ccs <> bs
                                else backendSend be $ record body1 <> ccs <> record body2
                        else backendSend be bs
                }
    withPairContextWith (interleave, id) (clientParam, serverParam) $ \(cctx, sctx) ->
        concurrently_
            (handshake sctx `shouldThrow` serverRejectedUnexpectedFirst)
            (handshake cctx `shouldThrow` anyTLSException)

-- An unknown record type is answered with unexpected_message (RFC 8446
-- Section 5).  The type of the client's first record is replaced; with a
-- length larger than what follows, as an SSLv2 ClientHello reads when taken
-- for a TLS record header, the server must answer from the header alone
-- rather than wait for a body that never comes.  A length of 0 keeps the
-- original one.
server_first_record_type_unexpected :: Word8 -> Word16 -> IO ()
server_first_record_type_unexpected ty len = do
    (clientParam, serverParam) <- generate arbitrary
    armed <- newIORef True
    let retype be =
            be
                { backendSend = \bs -> do
                    first <- atomicModifyIORef' armed (\a -> (False, a))
                    if first && B.length bs > 5 && B.head bs == 22
                        then do
                            let lenBytes
                                    | len == 0 = B.take 2 (B.drop 3 bs)
                                    | otherwise =
                                        B.pack [fromIntegral (len `div` 256), fromIntegral (len `mod` 256)]
                            backendSend be $
                                B.singleton ty <> B.take 2 (B.drop 1 bs) <> lenBytes <> B.drop 5 bs
                        else backendSend be bs
                }
    r <- timeout 5000000 $
        withPairContextWith (retype, id) (clientParam, serverParam) $ \(cctx, sctx) ->
            concurrently_
                (handshake sctx `shouldThrow` serverRejectedUnexpectedFirst)
                (handshake cctx `shouldThrow` anyTLSException)
    r `shouldSatisfy` isJust

-- RFC 8446 Section 5.1: handshake messages MUST NOT be interleaved with
-- other record types.  After the handshake the client sends, protected under
-- its application traffic secret, a record holding only the first two bytes
-- of a KeyUpdate and then a record of application data.  The server must not
-- deliver the data and must answer with unexpected_message.
handshake13_interleaved_app_data :: IO ()
handshake13_interleaved_app_data = do
    let cipher = cipher13_AES_128_GCM_SHA256
    (clientParam, serverParam) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                ([TLS13], [TLS13])
                ([cipher], [cipher])
    secretRef <- newIORef Nothing
    sendRef <- newIORef (\_ -> return ())
    let logKey line = case words line of
            ["CLIENT_TRAFFIC_SECRET_0", _, h] -> writeIORef secretRef (Just h)
            _ -> return ()
        clientParam' =
            clientParam
                { clientDebug = (clientDebug clientParam){debugKeyLogger = logKey}
                }
        -- remember the raw sender so that records can be written by hand
        capture be =
            be
                { backendSend = \bs -> do
                    writeIORef sendRef (backendSend be)
                    backendSend be bs
                }
    withPairContextWith (capture, id) (clientParam', serverParam) $ \(cctx, sctx) ->
        concurrently_
            ( do
                handshake sctx
                r <- E.try (recvData sctx) :: IO (Either TLSException B.ByteString)
                case r of
                    Right d -> expectationFailure $ "server delivered " ++ show d
                    Left _ -> return ()
            )
            ( do
                handshake cctx
                Just h <- readIORef secretRef
                send <- readIORef sendRef
                let secret = BA.convert (unhex h) :: BA.ScrubbedBytes
                    key = hkdfExpandLabel SHA256 secret "key" "" 16 :: B.ByteString
                    iv = hkdfExpandLabel SHA256 secret "iv" "" 12 :: B.ByteString
                    -- the first two bytes of KeyUpdate(update_not_requested)
                    partialKeyUpdate = B.pack [24, 0]
                send $ protect13 key iv 0 22 partialKeyUpdate
                send $ protect13 key iv 1 23 "hello"
                recvData cctx `shouldThrow` peerSentUnexpected
            )
  where
    unhex [] = B.empty
    unhex (a : b : rest) = B.cons (read ['0', 'x', a, b]) (unhex rest)
    unhex _ = error "unhex"

-- A ChangeCipherSpec is the single byte 1; any other value is answered with
-- unexpected_message (RFC 8446 Section 5), and so is one carrying two of
-- them.  The client's ChangeCipherSpec record is made two bytes long.  The
-- groups are fixed so that a TLS 1.3 handshake goes through a
-- HelloRetryRequest -- after which the client sends its ChangeCipherSpec
-- before the second ClientHello -- or not, as the test asks.
malformed_ccs_unexpected :: Version -> Cipher -> [Group] -> [Group] -> IO ()
malformed_ccs_unexpected version cipher cgroups sgroups = do
    (clientParam0, serverParam0) <-
        generate $
            arbitraryPairParamsWithVersionsAndCiphers
                ([version], [version])
                ([cipher], [cipher])
    let clientParam =
            clientParam0
                { clientSupported = (clientSupported clientParam0){supportedGroups = cgroups}
                }
        serverParam =
            serverParam0
                { serverSupported =
                    (serverSupported serverParam0)
                        { supportedGroups = sgroups
                        , supportedGroupsTLS13 = [sgroups]
                        }
                }
    seen <- newIORef False
    let ccs = B.pack [20, 3, 3, 0, 1, 1]
        doubled = B.pack [20, 3, 3, 0, 2, 1, 1]
        double be =
            be
                { backendSend = \bs ->
                    if ccs `B.isPrefixOf` bs
                        then do
                            writeIORef seen True
                            backendSend be $ doubled <> B.drop 6 bs
                        else backendSend be bs
                }
    -- a TLS 1.3 server takes the client's ChangeCipherSpec and Finished in
    -- its first recvData, a TLS 1.2 one in handshake
    r <- timeout 10000000 $
        withPairContextWith (double, id) (clientParam, serverParam) $ \(cctx, sctx) ->
            concurrently_
                ((handshake sctx >> recvData sctx) `shouldThrow` rejectedAsUnexpected)
                ( void
                    (E.try (handshake cctx >> recvData cctx) :: IO (Either TLSException B.ByteString))
                )
    r `shouldSatisfy` isJust
    readIORef seen `shouldReturn` True

rejectedAsUnexpected :: TLSException -> Bool
rejectedAsUnexpected (HandshakeFailed (Error_Packet_unexpected _ _)) = True
rejectedAsUnexpected (Terminated _ _ (Error_Packet_unexpected _ _)) = True
rejectedAsUnexpected _ = False

-- An AES-128-GCM TLS 1.3 record: content and inner type, protected with the
-- record header as additional data and the sequence number in the nonce.
protect13 :: B.ByteString -> B.ByteString -> Word64 -> Word8 -> B.ByteString -> B.ByteString
protect13 key iv sqn innerType content = hdr <> ct <> BA.convert tag
  where
    len = B.length content + 1 + 16
    hdr = B.pack [23, 3, 3, fromIntegral (len `div` 256), fromIntegral (len `mod` 256)]
    sqnBytes = B.pack [fromIntegral (sqn `shiftR` (8 * i)) | i <- [7, 6 .. 0]]
    nonce = B.pack $ B.zipWith xor iv (B.replicate 4 0 <> sqnBytes)
    aes = throwCryptoError (cipherInit key) :: AES128
    aead = throwCryptoError (aeadInit AEAD_GCM aes nonce)
    (AuthTag tag, ct) = aeadSimpleEncrypt aead hdr (content <> B.singleton innerType) 16

peerSentUnexpected :: TLSException -> Bool
peerSentUnexpected (Terminated True _ (Error_Protocol _ UnexpectedMessage)) = True
peerSentUnexpected _ = False

serverRejectedUnexpectedFirst :: TLSException -> Bool
serverRejectedUnexpectedFirst (HandshakeFailed (Error_Packet_unexpected _ _)) = True
serverRejectedUnexpectedFirst _ = False

expectJust :: String -> Maybe a -> Expectation
expectJust tag mx = case mx of
    Nothing -> expectationFailure tag
    Just _ -> return ()