From 51a9750891db2ce133b172b652a0514f378de192 Mon Sep 17 00:00:00 2001 From: Evgeny Poberezkin <2769109+epoberezkin@users.noreply.github.com> Date: Sat, 25 Dec 2021 17:13:53 +0000 Subject: [PATCH] double ratchet algorithm implementation (#236) * started double ratchet implementation * initialize ratchets * started ratchet encryption * ratchet encryption * simplify / narrow down Ratchet type * double ratchet decryption "framework" * advance receive ratched on skipped messages * more ratchet decryption * double ratchet encrypt/decrypt (TODO tests) * double ratchet tests * double ratchet tests * use ratchet AD in header encryption, use header and ratchet AD as AD in message encryption * change ratchet message error, remove Show instances * Update tests/AgentTests/DoubleRatchetTests.hs Co-authored-by: Efim Poberezkin <8711996+efim-poberezkin@users.noreply.github.com> * Update tests/AgentTests/DoubleRatchetTests.hs Co-authored-by: Efim Poberezkin <8711996+efim-poberezkin@users.noreply.github.com> * Update tests/AgentTests/DoubleRatchetTests.hs Co-authored-by: Efim Poberezkin <8711996+efim-poberezkin@users.noreply.github.com> * Update tests/AgentTests/DoubleRatchetTests.hs Co-authored-by: Efim Poberezkin <8711996+efim-poberezkin@users.noreply.github.com> * Update tests/AgentTests/DoubleRatchetTests.hs Co-authored-by: Efim Poberezkin <8711996+efim-poberezkin@users.noreply.github.com> * Update src/Simplex/Messaging/Crypto/Ratchet.hs Co-authored-by: Efim Poberezkin <8711996+efim-poberezkin@users.noreply.github.com> * test in the same ratchet step * merge tests * Update src/Simplex/Messaging/Crypto/Ratchet.hs * Update src/Simplex/Messaging/Crypto/Ratchet.hs * remove HMAC comment Co-authored-by: Efim Poberezkin <8711996+efim-poberezkin@users.noreply.github.com> --- simplexmq.cabal | 2 + src/Simplex/Messaging/Crypto.hs | 57 ++- src/Simplex/Messaging/Crypto/Ratchet.hs | 397 +++++++++++++++++++++ src/Simplex/Messaging/Parsers.hs | 7 + src/Simplex/Messaging/Util.hs | 15 + tests/AgentTests.hs | 2 + tests/AgentTests/ConnectionRequestTests.hs | 2 +- tests/AgentTests/DoubleRatchetTests.hs | 189 ++++++++++ 8 files changed, 658 insertions(+), 13 deletions(-) create mode 100644 src/Simplex/Messaging/Crypto/Ratchet.hs create mode 100644 tests/AgentTests/DoubleRatchetTests.hs diff --git a/simplexmq.cabal b/simplexmq.cabal index e3f7f7acd..af4204243 100644 --- a/simplexmq.cabal +++ b/simplexmq.cabal @@ -45,6 +45,7 @@ library Simplex.Messaging.Agent.Store.SQLite.Migrations Simplex.Messaging.Client Simplex.Messaging.Crypto + Simplex.Messaging.Crypto.Ratchet Simplex.Messaging.Parsers Simplex.Messaging.Protocol Simplex.Messaging.Server @@ -219,6 +220,7 @@ test-suite smp-server-test other-modules: AgentTests AgentTests.ConnectionRequestTests + AgentTests.DoubleRatchetTests AgentTests.FunctionalAPITests AgentTests.SQLiteTests ProtocolErrorTests diff --git a/src/Simplex/Messaging/Crypto.hs b/src/Simplex/Messaging/Crypto.hs index 57e5a0bb1..163ae6069 100644 --- a/src/Simplex/Messaging/Crypto.hs +++ b/src/Simplex/Messaging/Crypto.hs @@ -32,6 +32,7 @@ module Simplex.Messaging.Crypto SAlgorithm (..), Alg (..), SignAlg (..), + DhAlgorithm, PrivateKey (..), PublicKey (..), APrivateKey (..), @@ -54,6 +55,8 @@ module Simplex.Messaging.Crypto privateToX509, -- * E2E hybrid encryption scheme + E2EEncryptionVersion, + currentE2EVersion, encrypt, encrypt', decrypt, @@ -85,6 +88,8 @@ module Simplex.Messaging.Crypto IV (..), encryptAES, decryptAES, + encryptAEAD, + decryptAEAD, authTagSize, authTagToBS, bsToAuthTag, @@ -92,6 +97,7 @@ module Simplex.Messaging.Crypto randomIV, aesKeyP, ivP, + ivSize, -- * NaCl crypto_box cbEncrypt, @@ -143,14 +149,20 @@ import Data.Kind (Constraint, Type) import Data.String import Data.Type.Equality import Data.Typeable (Typeable) +import Data.Word (Word16) import Data.X509 import Database.SQLite.Simple.FromField (FromField (..)) import Database.SQLite.Simple.ToField (ToField (..)) import GHC.TypeLits (ErrorMessage (..), TypeError) import Network.Transport.Internal (decodeWord32, encodeWord32) -import Simplex.Messaging.Parsers (base64P, base64UriP, blobFieldParser, parseAll, parseString) +import Simplex.Messaging.Parsers (base64P, base64UriP, blobFieldParser, parseAll, parseE, parseString) import Simplex.Messaging.Util (liftEitherError, (<$?>)) +type E2EEncryptionVersion = Word16 + +currentE2EVersion :: E2EEncryptionVersion +currentE2EVersion = 1 + -- | Cryptographic algorithms. data Algorithm = RSA | Ed25519 | Ed448 | X25519 | X448 @@ -757,8 +769,16 @@ data CryptoError CBDecryptError | -- | message does not fit in SMP block CryptoLargeMsgError - | -- | failure parsing RSA-encrypted message header + | -- | failure parsing message header CryptoHeaderError String + | -- | no sending chain key in ratchet state + CERatchetState + | -- | header decryption error (could indicate that another key should be tried) + CERatchetHeader + | -- | too many skipped messages + CERatchetTooManySkipped + | -- | duplicate message number (or, possibly, skipped message that failed to decrypt?) + CERatchetDuplicateMessage deriving (Eq, Show, Exception) pubExpRange :: Integer @@ -805,6 +825,7 @@ data Header = Header -- | AES key newtype. newtype Key = Key {unKey :: ByteString} + deriving (Eq, Ord) -- | IV bytes newtype. newtype IV = IV {unIV :: ByteString} @@ -845,8 +866,8 @@ aesKeyP = Key <$> A.take aesKeySize ivP :: Parser IV ivP = IV <$> A.take (ivSize @AES256) -parseHeader :: ByteString -> Either CryptoError Header -parseHeader = first CryptoHeaderError . parseAll headerP +parseHeader :: ByteString -> ExceptT CryptoError IO Header +parseHeader = parseE CryptoHeaderError headerP -- * E2E hybrid encryption scheme @@ -870,7 +891,7 @@ decrypt' :: PrivateKey a -> ByteString -> ExceptT CryptoError IO ByteString decrypt' pk@(PrivateKeyRSA _) msg'' = do let (encHeader, msg') = B.splitAt (keySize pk) msg'' header <- decryptOAEP pk encHeader - Header {aesKey, ivBytes, authTag, msgSize} <- except $ parseHeader header + Header {aesKey, ivBytes, authTag, msgSize} <- parseHeader header msg <- decryptAES aesKey ivBytes msg' authTag return $ B.take msgSize msg decrypt' _ _ = throwE UnsupportedAlgorithm @@ -881,27 +902,39 @@ encrypt (APublicEncryptKey _ k) = encrypt' k decrypt :: APrivateDecryptKey -> ByteString -> ExceptT CryptoError IO ByteString decrypt (APrivateDecryptKey _ pk) = decrypt' pk --- | AEAD-GCM encryption. +-- | AEAD-GCM encryption with empty associated data. -- -- Used as part of hybrid E2E encryption scheme and for SMP transport blocks encryption. encryptAES :: Key -> IV -> Int -> ByteString -> ExceptT CryptoError IO (AES.AuthTag, ByteString) -encryptAES aesKey ivBytes paddedSize msg = do +encryptAES key iv paddedLen = encryptAEAD key iv paddedLen "" + +-- | AEAD-GCM encryption. +-- +-- Used as part of hybrid E2E encryption scheme and for SMP transport blocks encryption. +encryptAEAD :: Key -> IV -> Int -> ByteString -> ByteString -> ExceptT CryptoError IO (AES.AuthTag, ByteString) +encryptAEAD aesKey ivBytes paddedSize ad msg = do aead <- initAEAD @AES256 aesKey ivBytes msg' <- paddedMsg - return $ AES.aeadSimpleEncrypt aead B.empty msg' authTagSize + return $ AES.aeadSimpleEncrypt aead ad msg' authTagSize where len = B.length msg paddedMsg - | len >= paddedSize = throwE CryptoLargeMsgError + | len > paddedSize = throwE CryptoLargeMsgError | otherwise = return (msg <> B.replicate (paddedSize - len) '#') +-- | AEAD-GCM decryption with empty associated data. +-- +-- Used as part of hybrid E2E encryption scheme and for SMP transport blocks decryption. +decryptAES :: Key -> IV -> ByteString -> AES.AuthTag -> ExceptT CryptoError IO ByteString +decryptAES key iv = decryptAEAD key iv "" + -- | AEAD-GCM decryption. -- -- Used as part of hybrid E2E encryption scheme and for SMP transport blocks decryption. -decryptAES :: Key -> IV -> ByteString -> AES.AuthTag -> ExceptT CryptoError IO ByteString -decryptAES aesKey ivBytes msg authTag = do +decryptAEAD :: Key -> IV -> ByteString -> ByteString -> AES.AuthTag -> ExceptT CryptoError IO ByteString +decryptAEAD aesKey ivBytes ad msg authTag = do aead <- initAEAD @AES256 aesKey ivBytes - maybeError AESDecryptError $ AES.aeadSimpleDecrypt aead B.empty msg authTag + maybeError AESDecryptError $ AES.aeadSimpleDecrypt aead ad msg authTag initAEAD :: forall c. AES.BlockCipher c => Key -> IV -> ExceptT CryptoError IO (AES.AEAD c) initAEAD (Key aesKey) (IV ivBytes) = do diff --git a/src/Simplex/Messaging/Crypto/Ratchet.hs b/src/Simplex/Messaging/Crypto/Ratchet.hs new file mode 100644 index 000000000..26bfe70ce --- /dev/null +++ b/src/Simplex/Messaging/Crypto/Ratchet.hs @@ -0,0 +1,397 @@ +{-# LANGUAGE DuplicateRecordFields #-} +{-# LANGUAGE GADTs #-} +{-# LANGUAGE LambdaCase #-} +{-# LANGUAGE NamedFieldPuns #-} +{-# LANGUAGE OverloadedStrings #-} +{-# LANGUAGE ScopedTypeVariables #-} +{-# LANGUAGE TupleSections #-} +{-# LANGUAGE TypeApplications #-} + +module Simplex.Messaging.Crypto.Ratchet where + +import Control.Monad.Except +import Control.Monad.Trans.Except +import Crypto.Cipher.AES (AES256) +import qualified Crypto.Cipher.Types as AES +import Crypto.Hash (SHA512) +import qualified Crypto.KDF.HKDF as H +import Data.Attoparsec.ByteString.Char8 (Parser) +import qualified Data.Attoparsec.ByteString.Char8 as A +import Data.ByteString.Char8 (ByteString) +import qualified Data.ByteString.Char8 as B +import Data.Map.Strict (Map) +import qualified Data.Map.Strict as M +import Data.Maybe (fromMaybe) +import Data.Word (Word16, Word32) +import Network.Transport.Internal (decodeWord16, decodeWord32, encodeWord16, encodeWord32) +import Simplex.Messaging.Crypto +import Simplex.Messaging.Parsers (parseAll, parseE, parseE') +import Simplex.Messaging.Util (tryE, (<$?>)) + +data Ratchet a = Ratchet + { -- current ratchet version + rcVersion :: E2EEncryptionVersion, + -- associated data - must be the same in both parties ratchets + rcAD :: ByteString, + rcDHRs :: KeyPair a, + rcRK :: RatchetKey, + rcSnd :: Maybe (SndRatchet a), + rcRcv :: Maybe RcvRatchet, + rcMKSkipped :: Map HeaderKey SkippedMsgKeys, + rcNs :: Word32, + rcNr :: Word32, + rcPN :: Word32, + rcNHKs :: HeaderKey, + rcNHKr :: HeaderKey + } + +data SndRatchet a = SndRatchet + { rcDHRr :: PublicKey a, + rcCKs :: RatchetKey, + rcHKs :: HeaderKey + } + +data RcvRatchet = RcvRatchet + { rcCKr :: RatchetKey, + rcHKr :: HeaderKey + } + +type SkippedMsgKeys = Map Word32 MessageKey + +type HeaderKey = Key + +data MessageKey = MessageKey Key IV + +data ARatchet + = forall a. + (AlgorithmI a, DhAlgorithm a) => + ARatchet (SAlgorithm a) (Ratchet a) + +-- | Input key material for double ratchet HKDF functions +newtype RatchetKey = RatchetKey ByteString + +-- | Sending ratchet initialization, equivalent to RatchetInitAliceHE in double ratchet spec +-- +-- Please note that sPKey is not stored, and its public part together with random salt +-- is sent to the recipient. +initSndRatchet' :: + forall a. (AlgorithmI a, DhAlgorithm a) => PublicKey a -> PrivateKey a -> ByteString -> ByteString -> IO (Ratchet a) +initSndRatchet' rcDHRr sPKey salt rcAD = do + rcDHRs@(_, pk) <- generateKeyPair' @a 0 + let (sk, rcHKs, rcNHKr) = initKdf salt rcDHRr sPKey + -- state.RK, state.CKs, state.NHKs = KDF_RK_HE(SK, DH(state.DHRs, state.DHRr)) + (rcRK, rcCKs, rcNHKs) = rootKdf sk rcDHRr pk + pure + Ratchet + { rcVersion = currentE2EVersion, + rcAD, + rcDHRs, + rcRK, + rcSnd = Just SndRatchet {rcDHRr, rcCKs, rcHKs}, + rcRcv = Nothing, + rcMKSkipped = M.empty, + rcPN = 0, + rcNs = 0, + rcNr = 0, + rcNHKs, + rcNHKr + } + +-- | Receiving ratchet initialization, equivalent to RatchetInitBobHE in double ratchet spec +-- +-- Please note that the public part of rcDHRs was sent to the sender +-- as part of the connection request and random salt was received from the sender. +initRcvRatchet' :: + forall a. (AlgorithmI a, DhAlgorithm a) => PublicKey a -> KeyPair a -> ByteString -> ByteString -> IO (Ratchet a) +initRcvRatchet' sKey rcDHRs@(_, pk) salt rcAD = do + let (sk, rcNHKr, rcNHKs) = initKdf salt sKey pk + pure + Ratchet + { rcVersion = currentE2EVersion, + rcAD, + rcDHRs, + rcRK = sk, + rcSnd = Nothing, + rcRcv = Nothing, + rcMKSkipped = M.empty, + rcPN = 0, + rcNs = 0, + rcNr = 0, + rcNHKs, + rcNHKr + } + +data MsgHeader a = MsgHeader + { -- | current E2E version + msgVersion :: E2EEncryptionVersion, + -- | latest E2E version supported by sending clients (to simplify version upgrade) + msgLatestVersion :: E2EEncryptionVersion, + msgDHRs :: PublicKey a, + msgPN :: Word32, + msgNs :: Word32, + msgLen :: Word16 + } + deriving (Eq, Show) + +data AMsgHeader + = forall a. + (AlgorithmI a, DhAlgorithm a) => + AMsgHeader (SAlgorithm a) (MsgHeader a) + +paddedHeaderLen :: Int +paddedHeaderLen = 128 + +fullHeaderLen :: Int +fullHeaderLen = paddedHeaderLen + authTagSize + ivSize @AES256 + +serializeMsgHeader' :: AlgorithmI a => MsgHeader a -> ByteString +serializeMsgHeader' MsgHeader {msgVersion, msgLatestVersion, msgDHRs, msgPN, msgNs, msgLen} = + encodeWord16 msgVersion + <> encodeWord16 msgLatestVersion + <> encodeWord16 (fromIntegral $ B.length key) + <> key + <> encodeWord32 msgPN + <> encodeWord32 msgNs + <> encodeWord16 msgLen + where + key = encodeKey msgDHRs + +msgHeaderP' :: AlgorithmI a => Parser (MsgHeader a) +msgHeaderP' = do + msgVersion <- word16 + msgLatestVersion <- word16 + keyLen <- fromIntegral <$> word16 + msgDHRs <- parseAll binaryKeyP <$?> A.take keyLen + msgPN <- word32 + msgNs <- word32 + msgLen <- word16 + pure MsgHeader {msgVersion, msgLatestVersion, msgDHRs, msgPN, msgNs, msgLen} + where + word16 = decodeWord16 <$> A.take 2 + word32 = decodeWord32 <$> A.take 4 + +data EncHeader = EncHeader + { ehBody :: ByteString, + ehAuthTag :: AES.AuthTag, + ehIV :: IV + } + +serializeEncHeader :: EncHeader -> ByteString +serializeEncHeader EncHeader {ehBody, ehAuthTag, ehIV} = + ehBody <> authTagToBS ehAuthTag <> unIV ehIV + +encHeaderP :: Parser EncHeader +encHeaderP = do + ehBody <- A.take paddedHeaderLen + ehAuthTag <- bsToAuthTag <$> A.take authTagSize + ehIV <- ivP + pure EncHeader {ehBody, ehAuthTag, ehIV} + +data EncMessage = EncMessage + { emHeader :: ByteString, + emBody :: ByteString, + emAuthTag :: AES.AuthTag + } + +serializeEncMessage :: EncMessage -> ByteString +serializeEncMessage EncMessage {emHeader, emBody, emAuthTag} = + emHeader <> emBody <> authTagToBS emAuthTag + +encMessageP :: Parser EncMessage +encMessageP = do + emHeader <- A.take fullHeaderLen + s <- A.takeByteString + when (B.length s <= authTagSize) $ fail "message too short" + let (emBody, aTag) = B.splitAt (B.length s - authTagSize) s + emAuthTag = bsToAuthTag aTag + pure EncMessage {emHeader, emBody, emAuthTag} + +rcEncrypt' :: AlgorithmI a => Ratchet a -> Int -> ByteString -> ExceptT CryptoError IO (ByteString, Ratchet a) +rcEncrypt' Ratchet {rcSnd = Nothing} _ _ = throwE CERatchetState +rcEncrypt' rc@Ratchet {rcSnd = Just sr@SndRatchet {rcCKs, rcHKs}, rcNs, rcAD} paddedMsgLen msg = do + -- state.CKs, mk = KDF_CK(state.CKs) + let (ck', mk, iv, ehIV) = chainKdf rcCKs + -- enc_header = HENCRYPT(state.HKs, header) + (ehAuthTag, ehBody) <- encryptAEAD rcHKs ehIV paddedHeaderLen rcAD msgHeader + -- return enc_header, ENCRYPT(mk, plaintext, CONCAT(AD, enc_header)) + let emHeader = serializeEncHeader EncHeader {ehBody, ehAuthTag, ehIV} + (emAuthTag, emBody) <- encryptAEAD mk iv paddedMsgLen (rcAD <> emHeader) msg + let msg' = serializeEncMessage EncMessage {emHeader, emBody, emAuthTag} + -- state.Ns += 1 + rc' = rc {rcSnd = Just sr {rcCKs = ck'}, rcNs = rcNs + 1} + pure (msg', rc') + where + -- header = HEADER(state.DHRs, state.PN, state.Ns) + msgHeader = + serializeMsgHeader' + MsgHeader + { msgVersion = rcVersion rc, + msgLatestVersion = currentE2EVersion, + msgDHRs = fst $ rcDHRs rc, + msgPN = rcPN rc, + msgNs = rcNs, + msgLen = fromIntegral $ B.length msg + } + +data SkippedMessage a + = SMMessage (Either CryptoError ByteString) (Ratchet a) + | SMHeader (Maybe RatchetStep) (MsgHeader a) + | SMNone + +data RatchetStep = AdvanceRatchet | SameRatchet + deriving (Eq) + +type DecryptResult a = (Either CryptoError ByteString, Ratchet a) + +maxSkip :: Word32 +maxSkip = 512 + +rcDecrypt' :: + forall a. + (AlgorithmI a, DhAlgorithm a) => + Ratchet a -> + ByteString -> + ExceptT CryptoError IO (DecryptResult a) +rcDecrypt' rc@Ratchet {rcRcv, rcMKSkipped, rcAD} msg' = do + encMsg@EncMessage {emHeader} <- parseE CryptoHeaderError encMessageP msg' + encHdr <- parseE CryptoHeaderError encHeaderP emHeader + -- plaintext = TrySkippedMessageKeysHE(state, enc_header, ciphertext, AD) + decryptSkipped encHdr encMsg >>= \case + SMNone -> do + (rcStep, hdr) <- decryptRcHeader rcRcv encHdr + decryptRcMessage rcStep hdr encMsg + SMHeader rcStep_ hdr -> + case rcStep_ of + Just rcStep -> decryptRcMessage rcStep hdr encMsg + Nothing -> throwE CERatchetHeader + SMMessage msg rc' -> pure (msg, rc') + where + decryptRcMessage :: RatchetStep -> MsgHeader a -> EncMessage -> ExceptT CryptoError IO (DecryptResult a) + decryptRcMessage rcStep hdr@MsgHeader {msgDHRs, msgPN, msgNs} encMsg = do + -- if dh_ratchet: + rc' <- ratchetStep rcStep + case skipMessageKeys msgNs rc' of + Left e -> pure (Left e, rc') + Right rc''@Ratchet {rcRcv = Just rr@RcvRatchet {rcCKr}, rcNr} -> do + -- state.CKr, mk = KDF_CK(state.CKr) + let (rcCKr', mk, iv, _) = chainKdf rcCKr + -- return DECRYPT (mk, ciphertext, CONCAT (AD, enc_header)) + msg <- decryptMessage (MessageKey mk iv) hdr encMsg + -- state . Nr += 1 + pure (msg, rc'' {rcRcv = Just rr {rcCKr = rcCKr'}, rcNr = rcNr + 1}) + Right rc'' -> pure (Left CERatchetState, rc'') + where + ratchetStep :: RatchetStep -> ExceptT CryptoError IO (Ratchet a) + ratchetStep SameRatchet = pure rc + ratchetStep AdvanceRatchet = + -- SkipMessageKeysHE(state, header.pn) + case skipMessageKeys msgPN rc of + Left e -> throwE e + Right rc'@Ratchet {rcDHRs, rcRK, rcNHKs, rcNHKr} -> do + -- DHRatchetHE(state, header) + rcDHRs' <- liftIO $ generateKeyPair' @a 0 + -- state.RK, state.CKr, state.NHKr = KDF_RK_HE(state.RK, DH(state.DHRs, state.DHRr)) + let (rcRK', rcCKr', rcNHKr') = rootKdf rcRK msgDHRs (snd rcDHRs) + -- state.RK, state.CKs, state.NHKs = KDF_RK_HE(state.RK, DH(state.DHRs, state.DHRr)) + (rcRK'', rcCKs', rcNHKs') = rootKdf rcRK' msgDHRs (snd rcDHRs') + pure + rc' + { rcDHRs = rcDHRs', + rcRK = rcRK'', + rcSnd = Just SndRatchet {rcDHRr = msgDHRs, rcCKs = rcCKs', rcHKs = rcNHKs}, + rcRcv = Just RcvRatchet {rcCKr = rcCKr', rcHKr = rcNHKr}, + rcPN = rcNs rc, + rcNs = 0, + rcNr = 0, + rcNHKs = rcNHKs', + rcNHKr = rcNHKr' + } + skipMessageKeys :: Word32 -> Ratchet a -> Either CryptoError (Ratchet a) + skipMessageKeys _ r@Ratchet {rcRcv = Nothing} = Right r + skipMessageKeys untilN r@Ratchet {rcRcv = Just rr@RcvRatchet {rcCKr, rcHKr}, rcNr, rcMKSkipped = mkSkipped} + | rcNr > untilN = Left CERatchetDuplicateMessage + | rcNr + maxSkip < untilN = Left CERatchetTooManySkipped + | rcNr == untilN = Right r + | otherwise = + let mks = fromMaybe M.empty $ M.lookup rcHKr mkSkipped + (rcCKr', rcNr', mks') = advanceRcvRatchet (untilN - rcNr) rcCKr rcNr mks + in Right + r + { rcRcv = Just rr {rcCKr = rcCKr'}, + rcNr = rcNr', + rcMKSkipped = M.insert rcHKr mks' mkSkipped + } + advanceRcvRatchet :: Word32 -> RatchetKey -> Word32 -> SkippedMsgKeys -> (RatchetKey, Word32, SkippedMsgKeys) + advanceRcvRatchet 0 ck msgNs mks = (ck, msgNs, mks) + advanceRcvRatchet n ck msgNs mks = + let (ck', mk, iv, _) = chainKdf ck + mks' = M.insert msgNs (MessageKey mk iv) mks + in advanceRcvRatchet (n - 1) ck' (msgNs + 1) mks' + decryptSkipped :: EncHeader -> EncMessage -> ExceptT CryptoError IO (SkippedMessage a) + decryptSkipped encHdr encMsg = tryDecryptSkipped SMNone $ M.assocs rcMKSkipped + where + tryDecryptSkipped :: SkippedMessage a -> [(HeaderKey, SkippedMsgKeys)] -> ExceptT CryptoError IO (SkippedMessage a) + tryDecryptSkipped SMNone ((hk, mks) : hks) = do + tryE (decryptHeader hk encHdr) >>= \case + Left CERatchetHeader -> tryDecryptSkipped SMNone hks + Left e -> throwE e + Right hdr@MsgHeader {msgNs} -> + case M.lookup msgNs mks of + Nothing -> + let nextRc + | maybe False ((== hk) . rcHKr) rcRcv = Just SameRatchet + | hk == rcNHKr rc = Just AdvanceRatchet + | otherwise = Nothing + in pure $ SMHeader nextRc hdr + Just mk -> do + let mks' = M.delete msgNs mks + mksSkipped + | M.null mks' = M.delete hk rcMKSkipped + | otherwise = M.insert hk mks' rcMKSkipped + rc' = rc {rcMKSkipped = mksSkipped} + msg <- decryptMessage mk hdr encMsg + pure $ SMMessage msg rc' + tryDecryptSkipped r _ = pure r + decryptRcHeader :: Maybe RcvRatchet -> EncHeader -> ExceptT CryptoError IO (RatchetStep, MsgHeader a) + decryptRcHeader Nothing hdr = decryptNextHeader hdr + decryptRcHeader (Just RcvRatchet {rcHKr}) hdr = + -- header = HDECRYPT(state.HKr, enc_header) + ((SameRatchet,) <$> decryptHeader rcHKr hdr) `catchE` \case + CERatchetHeader -> decryptNextHeader hdr + e -> throwE e + -- header = HDECRYPT(state.NHKr, enc_header) + decryptNextHeader hdr = (AdvanceRatchet,) <$> decryptHeader (rcNHKr rc) hdr + decryptHeader k EncHeader {ehBody, ehAuthTag, ehIV} = do + header <- decryptAEAD k ehIV rcAD ehBody ehAuthTag `catchE` \_ -> throwE CERatchetHeader + parseE' CryptoHeaderError msgHeaderP' header + decryptMessage :: MessageKey -> MsgHeader a -> EncMessage -> ExceptT CryptoError IO (Either CryptoError ByteString) + decryptMessage (MessageKey mk iv) MsgHeader {msgLen} EncMessage {emHeader, emBody, emAuthTag} = + -- DECRYPT(mk, ciphertext, CONCAT(AD, enc_header)) + -- TODO add associated data + tryE (B.take (fromIntegral msgLen) <$> decryptAEAD mk iv (rcAD <> emHeader) emBody emAuthTag) + +initKdf :: (AlgorithmI a, DhAlgorithm a) => ByteString -> PublicKey a -> PrivateKey a -> (RatchetKey, Key, Key) +initKdf salt k pk = + let dhOut = dhSecretBytes $ dh' k pk + (sk, hk, nhk) = hkdf3 salt dhOut "SimpleXInitRatchet" + in (RatchetKey sk, Key hk, Key nhk) + +rootKdf :: (AlgorithmI a, DhAlgorithm a) => RatchetKey -> PublicKey a -> PrivateKey a -> (RatchetKey, RatchetKey, Key) +rootKdf (RatchetKey rk) k pk = + let dhOut = dhSecretBytes $ dh' k pk + (rk', ck, nhk) = hkdf3 rk dhOut "SimpleXRootRatchet" + in (RatchetKey rk', RatchetKey ck, Key nhk) + +chainKdf :: RatchetKey -> (RatchetKey, Key, IV, IV) +chainKdf (RatchetKey ck) = + let (ck', mk, ivs) = hkdf3 "" ck "SimpleXChainRatchet" + (iv1, iv2) = B.splitAt 16 ivs + in (RatchetKey ck', Key mk, IV iv1, IV iv2) + +hkdf3 :: ByteString -> ByteString -> ByteString -> (ByteString, ByteString, ByteString) +hkdf3 salt ikm info = (s1, s2, s3) + where + prk = H.extract salt ikm :: H.PRK SHA512 + out = H.expand prk info 96 + (s1, rest) = B.splitAt 32 out + (s2, s3) = B.splitAt 32 rest diff --git a/src/Simplex/Messaging/Parsers.hs b/src/Simplex/Messaging/Parsers.hs index 8e82852b8..d14419bf0 100644 --- a/src/Simplex/Messaging/Parsers.hs +++ b/src/Simplex/Messaging/Parsers.hs @@ -3,6 +3,7 @@ module Simplex.Messaging.Parsers where +import Control.Monad.Trans.Except import Data.Attoparsec.ByteString.Char8 (Parser) import qualified Data.Attoparsec.ByteString.Char8 as A import Data.Bifunctor (first) @@ -53,6 +54,12 @@ parse parser err = first (const err) . parseAll parser parseAll :: Parser a -> (ByteString -> Either String a) parseAll parser = A.parseOnly (parser <* A.endOfInput) +parseE :: (String -> e) -> Parser a -> (ByteString -> ExceptT e IO a) +parseE err parser = except . first err . parseAll parser + +parseE' :: (String -> e) -> Parser a -> (ByteString -> ExceptT e IO a) +parseE' err parser = except . first err . A.parseOnly parser + parseRead :: Read a => Parser ByteString -> Parser a parseRead = (>>= maybe (fail "cannot read") pure . readMaybe . B.unpack) diff --git a/src/Simplex/Messaging/Util.hs b/src/Simplex/Messaging/Util.hs index d558a636a..616101f19 100644 --- a/src/Simplex/Messaging/Util.hs +++ b/src/Simplex/Messaging/Util.hs @@ -7,6 +7,7 @@ module Simplex.Messaging.Util where import Control.Monad.Except import Control.Monad.IO.Unlift +import Control.Monad.Trans.Except import Data.Bifunctor (first) import Data.ByteString.Char8 (ByteString) import qualified Data.ByteString.Char8 as B @@ -36,30 +37,44 @@ infixl 4 <$$>, <$?> (<$$>) :: (Functor f, Functor g) => (a -> b) -> f (g a) -> f (g b) (<$$>) = fmap . fmap +{-# INLINE (<$$>) #-} (<$?>) :: MonadFail m => (a -> Either String b) -> m a -> m b f <$?> m = m >>= either fail pure . f +{-# INLINE (<$?>) #-} bshow :: Show a => a -> ByteString bshow = B.pack . show +{-# INLINE bshow #-} maybeWord :: (a -> ByteString) -> Maybe a -> ByteString maybeWord f = maybe "" $ B.cons ' ' . f +{-# INLINE maybeWord #-} liftIOEither :: (MonadIO m, MonadError e m) => IO (Either e a) -> m a liftIOEither a = liftIO a >>= liftEither +{-# INLINE liftIOEither #-} liftError :: (MonadIO m, MonadError e' m) => (e -> e') -> ExceptT e IO a -> m a liftError f = liftEitherError f . runExceptT +{-# INLINE liftError #-} liftEitherError :: (MonadIO m, MonadError e' m) => (e -> e') -> IO (Either e a) -> m a liftEitherError f a = liftIOEither (first f <$> a) +{-# INLINE liftEitherError #-} tryError :: MonadError e m => m a -> m (Either e a) tryError action = (Right <$> action) `catchError` (pure . Left) +{-# INLINE tryError #-} + +tryE :: Monad m => ExceptT e m a -> ExceptT e m (Either e a) +tryE m = (Right <$> m) `catchE` (pure . Left) +{-# INLINE tryE #-} ifM :: Monad m => m Bool -> m a -> m a -> m a ifM ba t f = ba >>= \b -> if b then t else f +{-# INLINE ifM #-} unlessM :: Monad m => m Bool -> m () -> m () unlessM b = ifM b $ pure () +{-# INLINE unlessM #-} diff --git a/tests/AgentTests.hs b/tests/AgentTests.hs index 266d899de..23855700a 100644 --- a/tests/AgentTests.hs +++ b/tests/AgentTests.hs @@ -10,6 +10,7 @@ module AgentTests (agentTests) where import AgentTests.ConnectionRequestTests +import AgentTests.DoubleRatchetTests (doubleRatchetTests) import AgentTests.FunctionalAPITests (functionalAPITests) import AgentTests.SQLiteTests (storeTests) import Control.Concurrent @@ -31,6 +32,7 @@ import Test.Hspec agentTests :: ATransport -> Spec agentTests (ATransport t) = do describe "Connection request" connectionRequestTests + describe "Double ratchet tests" doubleRatchetTests describe "Functional API" $ functionalAPITests (ATransport t) describe "SQLite store" storeTests describe "SMP agent protocol syntax" $ syntaxTests t diff --git a/tests/AgentTests/ConnectionRequestTests.hs b/tests/AgentTests/ConnectionRequestTests.hs index 49c6c9661..0bfb09b2e 100644 --- a/tests/AgentTests/ConnectionRequestTests.hs +++ b/tests/AgentTests/ConnectionRequestTests.hs @@ -40,7 +40,7 @@ connectionRequest = } connectionRequestTests :: Spec -connectionRequestTests = do +connectionRequestTests = describe "connection request parsing / serializing" $ do it "should serialize SMP queue URIs" $ do serializeSMPQueueUri queue {smpServer = srv {port = Nothing}} diff --git a/tests/AgentTests/DoubleRatchetTests.hs b/tests/AgentTests/DoubleRatchetTests.hs new file mode 100644 index 000000000..16d62dd0f --- /dev/null +++ b/tests/AgentTests/DoubleRatchetTests.hs @@ -0,0 +1,189 @@ +{-# LANGUAGE DataKinds #-} +{-# LANGUAGE LambdaCase #-} +{-# LANGUAGE OverloadedStrings #-} +{-# LANGUAGE PatternSynonyms #-} +{-# LANGUAGE RankNTypes #-} +{-# LANGUAGE ScopedTypeVariables #-} +{-# LANGUAGE TypeApplications #-} +{-# OPTIONS_GHC -fno-warn-unticked-promoted-constructors #-} + +module AgentTests.DoubleRatchetTests where + +import Control.Concurrent.STM +import Control.Monad.Except +import Crypto.Random (getRandomBytes) +import Data.ByteString.Char8 (ByteString) +import qualified Data.ByteString.Char8 as B +import Simplex.Messaging.Crypto (Algorithm (..), AlgorithmI, CryptoError, DhAlgorithm) +import qualified Simplex.Messaging.Crypto as C +import Simplex.Messaging.Crypto.Ratchet +import Simplex.Messaging.Parsers (parseAll) +import Test.Hspec + +doubleRatchetTests :: Spec +doubleRatchetTests = do + describe "double-ratchet encryption/decryption" $ do + it "should serialize and parse message header" testMessageHeader + it "should encrypt and decrypt messages" $ do + withRatchets @X25519 testEncryptDecrypt + withRatchets @X448 testEncryptDecrypt + it "should encrypt and decrypt skipped messages" $ do + withRatchets @X25519 testSkippedMessages + withRatchets @X448 testSkippedMessages + it "should encrypt and decrypt many messages" $ do + withRatchets @X25519 testManyMessages + it "should allow skipped after ratchet advance" $ do + withRatchets @X25519 testSkippedAfterRatchetAdvance + +paddedMsgLen :: Int +paddedMsgLen = 100 + +fullMsgLen :: Int +fullMsgLen = fullHeaderLen + paddedMsgLen + C.authTagSize + +testMessageHeader :: Expectation +testMessageHeader = do + (k, _) <- C.generateKeyPair' @X25519 0 + let hdr = MsgHeader {msgVersion = 1, msgLatestVersion = 1, msgDHRs = k, msgPN = 0, msgNs = 0, msgLen = 11} + parseAll (msgHeaderP' @X25519) (serializeMsgHeader' hdr) `shouldBe` Right hdr + +pattern Decrypted :: ByteString -> Either CryptoError (Either CryptoError ByteString) +pattern Decrypted msg <- Right (Right msg) + +type TestRatchets a = (AlgorithmI a, DhAlgorithm a) => TVar (Ratchet a) -> TVar (Ratchet a) -> IO () + +testEncryptDecrypt :: TestRatchets a +testEncryptDecrypt alice bob = do + (bob, "hello alice") #> alice + (alice, "hello bob") #> bob + Right b1 <- encrypt bob "how are you, alice?" + Right b2 <- encrypt bob "are you there?" + Right b3 <- encrypt bob "hey?" + Right a1 <- encrypt alice "how are you, bob?" + Right a2 <- encrypt alice "are you there?" + Right a3 <- encrypt alice "hey?" + Decrypted "how are you, alice?" <- decrypt alice b1 + Decrypted "are you there?" <- decrypt alice b2 + Decrypted "hey?" <- decrypt alice b3 + Decrypted "how are you, bob?" <- decrypt bob a1 + Decrypted "are you there?" <- decrypt bob a2 + Decrypted "hey?" <- decrypt bob a3 + (bob, "I'm here, all good") #> alice + (alice, "I'm here too, same") #> bob + pure () + +testSkippedMessages :: TestRatchets a +testSkippedMessages alice bob = do + Right msg1 <- encrypt bob "hello alice" + Right msg2 <- encrypt bob "hello there again" + Right msg3 <- encrypt bob "are you there?" + Decrypted "are you there?" <- decrypt alice msg3 + Right (Left C.CERatchetDuplicateMessage) <- decrypt alice msg3 + Decrypted "hello there again" <- decrypt alice msg2 + Decrypted "hello alice" <- decrypt alice msg1 + pure () + +testManyMessages :: TestRatchets a +testManyMessages alice bob = do + (bob, "b1") #> alice + (bob, "b2") #> alice + (bob, "b3") #> alice + (bob, "b4") #> alice + (alice, "a5") #> bob + (alice, "a6") #> bob + (alice, "a7") #> bob + (bob, "b8") #> alice + (alice, "a9") #> bob + (alice, "a10") #> bob + (bob, "b11") #> alice + (bob, "b12") #> alice + (alice, "a14") #> bob + (bob, "b15") #> alice + (bob, "b16") #> alice + +testSkippedAfterRatchetAdvance :: TestRatchets a +testSkippedAfterRatchetAdvance alice bob = do + (bob, "b1") #> alice + Right b2 <- encrypt bob "b2" + Right b3 <- encrypt bob "b3" + Right b4 <- encrypt bob "b4" + (alice, "a5") #> bob + Right b5 <- encrypt bob "b5" + Right b6 <- encrypt bob "b6" + (bob, "b7") #> alice + Right b8 <- encrypt bob "b8" + Right b9 <- encrypt bob "b9" + (alice, "a10") #> bob + Right b11 <- encrypt bob "b11" + Right b12 <- encrypt bob "b12" + (alice, "a14") #> bob + Decrypted "b12" <- decrypt alice b12 + Decrypted "b2" <- decrypt alice b2 + -- fails on duplicate message + Left C.CERatchetHeader <- decrypt alice b2 + (alice, "a15") #> bob + Right a16 <- encrypt bob "a16" + Right a17 <- encrypt bob "a17" + Decrypted "b8" <- decrypt alice b8 + Decrypted "b3" <- decrypt alice b3 + Decrypted "b4" <- decrypt alice b4 + Decrypted "b5" <- decrypt alice b5 + Decrypted "b6" <- decrypt alice b6 + (alice, "a18") #> bob + Decrypted "a16" <- decrypt alice a16 + Decrypted "a17" <- decrypt alice a17 + Decrypted "b9" <- decrypt alice b9 + Decrypted "b11" <- decrypt alice b11 + pure () + +(#>) :: (AlgorithmI a, DhAlgorithm a) => (TVar (Ratchet a), ByteString) -> TVar (Ratchet a) -> Expectation +(alice, msg) #> bob = do + Right msg' <- encrypt alice msg + Decrypted msg'' <- decrypt bob msg' + msg'' `shouldBe` msg + +withRatchets :: forall a. (AlgorithmI a, DhAlgorithm a) => (TVar (Ratchet a) -> TVar (Ratchet a) -> IO ()) -> Expectation +withRatchets test = do + (a, b) <- initRatchets @a + alice <- newTVarIO a + bob <- newTVarIO b + test alice bob `shouldReturn` () + +initRatchets :: (AlgorithmI a, DhAlgorithm a) => IO (Ratchet a, Ratchet a) +initRatchets = do + salt <- getRandomBytes 16 + (ak, apk) <- C.generateKeyPair' 0 + (bk, bpk) <- C.generateKeyPair' 0 + bob <- initSndRatchet' ak bpk salt "bob -> alice" + alice <- initRcvRatchet' bk (ak, apk) salt "bob -> alice" + pure (alice, bob) + +encrypt_ :: AlgorithmI a => Ratchet a -> ByteString -> IO (Either CryptoError (ByteString, Ratchet a)) +encrypt_ rc msg = + runExceptT (rcEncrypt' rc paddedMsgLen msg) + >>= either (pure . Left) checkLength + where + checkLength r@(msg', _) = do + B.length msg' `shouldBe` fullMsgLen + pure $ Right r + +decrypt_ :: (AlgorithmI a, DhAlgorithm a) => Ratchet a -> ByteString -> IO (Either CryptoError (Either CryptoError ByteString, Ratchet a)) +decrypt_ rc msg = runExceptT $ rcDecrypt' rc msg + +encrypt :: AlgorithmI a => TVar (Ratchet a) -> ByteString -> IO (Either CryptoError ByteString) +encrypt = withTVar encrypt_ + +decrypt :: (AlgorithmI a, DhAlgorithm a) => TVar (Ratchet a) -> ByteString -> IO (Either CryptoError (Either CryptoError ByteString)) +decrypt = withTVar decrypt_ + +withTVar :: + (Ratchet a -> ByteString -> IO (Either e (r, Ratchet a))) -> + TVar (Ratchet a) -> + ByteString -> + IO (Either e r) +withTVar op rcVar msg = + readTVarIO rcVar + >>= (`op` msg) + >>= \case + Right (res, rc') -> atomically (writeTVar rcVar rc') >> pure (Right res) + Left e -> pure $ Left e