diff -ruN -x 'dist*' tls-1.9.0-orig/Network/TLS/Crypto.hs tls-fork/Network/TLS/Crypto.hs --- tls-1.9.0-orig/Network/TLS/Crypto.hs 2026-09-30 18:44:40.083085887 +0000 +++ tls-fork/Network/TLS/Crypto.hs 2026-09-30 18:44:40.106519923 +0000 @@ -8,6 +8,7 @@ , hashUpdate , hashUpdateSSL , hashFinal + , hashCopy , module Network.TLS.Crypto.DH , module Network.TLS.Crypto.IES @@ -44,6 +45,8 @@ import qualified Crypto.Hash as H import qualified Data.ByteString as B import qualified Data.ByteArray as B (convert) +import qualified Data.ByteArray as BA +import Unsafe.Coerce (unsafeCoerce) import Crypto.Error import Crypto.Number.Basic (numBits) import Crypto.Random @@ -135,6 +138,17 @@ hashUpdate (HashContextSSL sha1Ctx md5Ctx) b = HashContextSSL (H.hashUpdate sha1Ctx b) (H.hashUpdate md5Ctx b) +-- | Copies the context into a new allocation. +hashCopy :: HashContext -> IO HashContext +hashCopy (HashContext (ContextSimple h)) = HashContext . ContextSimple <$> copyContext h +hashCopy (HashContextSSL sha1Ctx md5Ctx) = HashContextSSL <$> copyContext sha1Ctx <*> copyContext md5Ctx + +-- H.Context is a newtype over Bytes, with the constructor in a hidden crypton module +copyContext :: H.Context alg -> IO (H.Context alg) +copyContext h = do + b <- BA.copy h (\_ -> return ()) :: IO BA.Bytes + return $! unsafeCoerce b + hashUpdateSSL :: HashCtx -> (B.ByteString,B.ByteString) -- ^ (for the md5 context, for the sha1 context) -> HashCtx diff -ruN -x 'dist*' tls-1.9.0-orig/Network/TLS/Handshake/Compact.hs tls-fork/Network/TLS/Handshake/Compact.hs --- tls-1.9.0-orig/Network/TLS/Handshake/Compact.hs 1970-01-01 00:00:00.000000000 +0000 +++ tls-fork/Network/TLS/Handshake/Compact.hs 2026-09-30 18:45:10.137984909 +0000 @@ -0,0 +1,111 @@ +{-# LANGUAGE OverloadedStrings #-} +{-# OPTIONS_HADDOCK hide #-} +module Network.TLS.Handshake.Compact + ( compactContext + ) where + +import Control.Concurrent.MVar +import Control.Exception (evaluate) +import Crypto.Random (drgNew) +import qualified Data.ByteString as B +import Data.IORef + +import Network.TLS.Cipher +import Network.TLS.Context.Internal +import Network.TLS.Crypto +import Network.TLS.Extension (Cookie (..)) +import Network.TLS.Handshake.State +import Network.TLS.Imports +import Network.TLS.KeySchedule (hkdfExpandLabel) +import Network.TLS.Record.State +import Network.TLS.RNG +import Network.TLS.State +import Network.TLS.Struct +import Network.TLS.Types + +-- | Re-allocates the state a context keeps after the handshake, next to each other. +-- Each small pinned value created during the handshake shares a 4K block with short-lived +-- crypto temporaries, and keeps the whole block alive, with every dead ScrubbedBytes in it. +compactContext :: Context -> IO () +compactContext ctx = do + tx <- readMVar $ ctxTxState ctx + rx <- readMVar $ ctxRxState ctx + st <- readMVar $ ctxState ctx + hst <- readMVar $ ctxHandshake ctx + finished <- readIORef $ ctxFinished ctx + peerFinished <- readIORef $ ctxPeerFinished ctx + -- key derivation and reseeding allocate temporaries, so they precede the copies + txKey <- trafficKey tx + rxKey <- trafficKey rx + rng <- evaluate $ StateRNG $ fst $ withTLSRNG (stRandomGen st) drgNew + tx' <- copyRecordState BulkEncrypt txKey tx + rx' <- copyRecordState BulkDecrypt rxKey rx + st' <- copyTLSState rng st + hst' <- traverse copyHandshakeState hst + finished' <- traverse copyBytes finished + peerFinished' <- traverse copyBytes peerFinished + modifyMVar_ (ctxTxState ctx) $ \_ -> return tx' + modifyMVar_ (ctxRxState ctx) $ \_ -> return rx' + modifyMVar_ (ctxState ctx) $ \_ -> return st' + modifyMVar_ (ctxHandshake ctx) $ \_ -> return hst' + writeIORef (ctxFinished ctx) finished' + writeIORef (ctxPeerFinished ctx) peerFinished' + +copyBytes :: ByteString -> IO ByteString +copyBytes = evaluate . B.copy + +-- TLS 1.3 record keys are derived from the traffic secret kept in the record state; +-- earlier versions keep the bulk key only inside the cipher state. +trafficKey :: RecordState -> IO (Maybe ByteString) +trafficKey RecordState{stCipher = Just cipher, stCryptLevel = CryptApplicationSecret, stCryptState = cst} = + Just <$> evaluate (hkdfExpandLabel (cipherHash cipher) (cstMacSecret cst) "key" "" (bulkKeySize $ cipherBulk cipher)) +trafficKey _ = return Nothing + +copyRecordState :: BulkDirection -> Maybe ByteString -> RecordState -> IO RecordState +copyRecordState dir key_ rs@RecordState{stCryptState = cst} = do + iv <- copyBytes $ cstIV cst + secret <- copyBytes $ cstMacSecret cst + bulkState <- case (key_, stCipher rs) of + (Just key, Just cipher) -> bulkInit (cipherBulk cipher) dir <$> copyBytes key + _ -> return $ cstKey cst + evaluate rs{stCryptState = CryptState{cstKey = bulkState, cstIV = iv, cstMacSecret = secret}} + +copyTLSState :: StateRNG -> TLSState -> IO TLSState +copyTLSState rng st = do + session <- case stSession st of + Session sid -> Session <$> traverse copyBytes sid + clientVerified <- copyBytes $ stClientVerifiedData st + serverVerified <- copyBytes $ stServerVerifiedData st + proto <- traverse copyBytes $ stNegotiatedProtocol st + alpn <- traverse (traverse copyBytes) $ stClientALPNSuggest st + cookie <- traverse (\(Cookie c) -> Cookie <$> copyBytes c) $ stTLS13Cookie st + exporter <- traverse copyBytes $ stExporterMasterSecret st + evaluate + st + { stSession = session + , stClientVerifiedData = clientVerified + , stServerVerifiedData = serverVerified + , stNegotiatedProtocol = proto + , stClientALPNSuggest = alpn + , stTLS13Cookie = cookie + , stExporterMasterSecret = exporter + , stRandomGen = rng + } + +copyHandshakeState :: HandshakeState -> IO HandshakeState +copyHandshakeState hst = do + clientRandom <- ClientRandom <$> copyBytes (unClientRandom $ hstClientRandom hst) + serverRandom <- traverse (fmap ServerRandom . copyBytes . unServerRandom) $ hstServerRandom hst + mainSecret <- traverse copyBytes $ hstMasterSecret hst + digest <- case hstHandshakeDigest hst of + HandshakeMessages msgs -> HandshakeMessages <$> traverse copyBytes msgs + HandshakeDigestContext h -> HandshakeDigestContext <$> hashCopy h + resumption <- traverse (\(BaseSecret s) -> BaseSecret <$> copyBytes s) $ hstTLS13ResumptionSecret hst + evaluate + hst + { hstClientRandom = clientRandom + , hstServerRandom = serverRandom + , hstMasterSecret = mainSecret + , hstHandshakeDigest = digest + , hstTLS13ResumptionSecret = resumption + } diff -ruN -x 'dist*' tls-1.9.0-orig/Network/TLS/Handshake.hs tls-fork/Network/TLS/Handshake.hs --- tls-1.9.0-orig/Network/TLS/Handshake.hs 2026-09-30 18:44:40.083878116 +0000 +++ tls-fork/Network/TLS/Handshake.hs 2026-09-30 18:45:14.390780619 +0000 @@ -18,6 +18,7 @@ import Network.TLS.Struct import Network.TLS.Handshake.Common +import Network.TLS.Handshake.Compact import Network.TLS.Handshake.Client import Network.TLS.Handshake.Server @@ -27,7 +28,7 @@ -- This is to be called at the beginning of a connection, and during renegotiation handshake :: MonadIO m => Context -> m () handshake ctx = - liftIO $ withRWLock ctx $ handleException ctx (ctxDoHandshake ctx ctx) + liftIO $ withRWLock ctx $ handleException ctx (ctxDoHandshake ctx ctx) >> compactContext ctx -- Handshake when requested by the remote end -- This is called automatically by 'recvData', in a context where the read lock diff -ruN -x 'dist*' tls-1.9.0-orig/tls.cabal tls-fork/tls.cabal --- tls-1.9.0-orig/tls.cabal 2026-09-30 18:44:40.090442445 +0000 +++ tls-fork/tls.cabal 2026-09-30 19:12:16.099341583 +0000 @@ -85,6 +85,7 @@ Network.TLS.Handshake.Signature Network.TLS.Handshake.State Network.TLS.Handshake.State13 + Network.TLS.Handshake.Compact Network.TLS.Hooks Network.TLS.IO Network.TLS.Imports