-- SEIPDv2.hs: OpenPGP (RFC9580) SEIPDv2 and SKESK v6 crypto helpers
-- Copyright © 2012-2026  Clint Adams
-- This software is released under the terms of the Expat license.
-- (See the LICENSE file).
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE PackageImports #-}
{-# LANGUAGE TypeApplications #-}

module Codec.Encryption.OpenPGP.SEIPDv2
    ( aeadModeAndNonceSizeForSEIPDv2
    , supportedSEIPDv2AEADAlgorithms
    , supportedSEIPDv2SymmetricAlgorithms
    , seipdv2SymmetricKeySize
    , deriveSKESK6KEK
    , encryptSKESK6SessionKey
    , decryptSKESK6SessionKey
    ) where

import Control.Error.Util (note)
import qualified Crypto.Hash.Algorithms as CHA
import Crypto.KDF.HKDF (expand, extract)
import Data.Bifunctor (first)
import qualified Data.ByteArray as BA
import qualified Data.ByteString as B
import Data.Either (isRight)
import qualified Data.Set as Set
import qualified "crypton" Crypto.Cipher.Types as CCT

import Codec.Encryption.OpenPGP.BlockCipher
    ( withAEADCipher
    )
import Codec.Encryption.OpenPGP.Internal.HOBlockCipher
    ( HOBlockCipher (..)
    )
import Codec.Encryption.OpenPGP.Internal.RFC7253OCB
    ( decryptWithOCBRFC7253With
    , encryptWithOCBRFC7253
    )
import Codec.Encryption.OpenPGP.Types

aeadModeAndNonceSizeForSEIPDv2
    :: AEADAlgorithm -> Either SEIPDv2Failure (CCT.AEADMode, Int)
aeadModeAndNonceSizeForSEIPDv2 :: AEADAlgorithm -> Either SEIPDv2Failure (AEADMode, Int)
aeadModeAndNonceSizeForSEIPDv2 AEADAlgorithm
EAX =
    SEIPDv2Failure -> Either SEIPDv2Failure (AEADMode, Int)
forall a b. a -> Either a b
Left (SEIPDv2Failure -> Either SEIPDv2Failure (AEADMode, Int))
-> SEIPDv2Failure -> Either SEIPDv2Failure (AEADMode, Int)
forall a b. (a -> b) -> a -> b
$ AEADAlgorithm -> SEIPDv2Failure
SEIPDv2UnsupportedAEADAlgorithm AEADAlgorithm
EAX
aeadModeAndNonceSizeForSEIPDv2 AEADAlgorithm
OCB = (AEADMode, Int) -> Either SEIPDv2Failure (AEADMode, Int)
forall a b. b -> Either a b
Right (AEADMode
CCT.AEAD_OCB, Int
15)
aeadModeAndNonceSizeForSEIPDv2 AEADAlgorithm
GCM = (AEADMode, Int) -> Either SEIPDv2Failure (AEADMode, Int)
forall a b. b -> Either a b
Right (AEADMode
CCT.AEAD_GCM, Int
12)
aeadModeAndNonceSizeForSEIPDv2 (OtherAEADAlgo Word8
_) =
    SEIPDv2Failure -> Either SEIPDv2Failure (AEADMode, Int)
forall a b. a -> Either a b
Left (SEIPDv2Failure -> Either SEIPDv2Failure (AEADMode, Int))
-> (AEADAlgorithm -> SEIPDv2Failure)
-> AEADAlgorithm
-> Either SEIPDv2Failure (AEADMode, Int)
forall b c a. (b -> c) -> (a -> b) -> a -> c
. AEADAlgorithm -> SEIPDv2Failure
SEIPDv2UnsupportedAEADAlgorithm (AEADAlgorithm -> Either SEIPDv2Failure (AEADMode, Int))
-> AEADAlgorithm -> Either SEIPDv2Failure (AEADMode, Int)
forall a b. (a -> b) -> a -> b
$ Word8 -> AEADAlgorithm
OtherAEADAlgo Word8
0

{- | AEAD algorithms that the SEIPDv2 encryption backend can actually use,
in descending preference order. Derived directly from
'aeadModeAndNonceSizeForSEIPDv2' so that enabling backend support for an
algorithm (e.g. flipping the EAX case to 'Right') automatically makes it
available to capability negotiation without touching the negotiation code.
-}
supportedSEIPDv2AEADAlgorithms :: Set.Set AEADAlgorithm
supportedSEIPDv2AEADAlgorithms :: Set AEADAlgorithm
supportedSEIPDv2AEADAlgorithms =
    [AEADAlgorithm] -> Set AEADAlgorithm
forall a. Ord a => [a] -> Set a
Set.fromList
        [ AEADAlgorithm
a
        | AEADAlgorithm
a <- [AEADAlgorithm
OCB, AEADAlgorithm
EAX, AEADAlgorithm
GCM]
        , Either SEIPDv2Failure (AEADMode, Int) -> Bool
forall a b. Either a b -> Bool
isRight (AEADAlgorithm -> Either SEIPDv2Failure (AEADMode, Int)
aeadModeAndNonceSizeForSEIPDv2 AEADAlgorithm
a)
        ]

{- | Symmetric algorithms that the SEIPDv2 encryption backend can actually use.
Derived directly from 'seipdv2SymmetricKeySize' so that enabling backend support
for an algorithm automatically makes it available to capability negotiation.
-}
supportedSEIPDv2SymmetricAlgorithms :: Set.Set SymmetricAlgorithm
supportedSEIPDv2SymmetricAlgorithms :: Set SymmetricAlgorithm
supportedSEIPDv2SymmetricAlgorithms =
    [SymmetricAlgorithm] -> Set SymmetricAlgorithm
forall a. Ord a => [a] -> Set a
Set.fromList
        [ SymmetricAlgorithm
a
        | SymmetricAlgorithm
a <- [SymmetricAlgorithm
AES128, SymmetricAlgorithm
AES192, SymmetricAlgorithm
AES256]
        , Either SEIPDv2Failure Int -> Bool
forall a b. Either a b -> Bool
isRight (SymmetricAlgorithm -> Either SEIPDv2Failure Int
seipdv2SymmetricKeySize SymmetricAlgorithm
a)
        ]

seipdv2SymmetricKeySize
    :: SymmetricAlgorithm -> Either SEIPDv2Failure Int
seipdv2SymmetricKeySize :: SymmetricAlgorithm -> Either SEIPDv2Failure Int
seipdv2SymmetricKeySize SymmetricAlgorithm
symalgo =
    case SymmetricAlgorithm
symalgo of
        SymmetricAlgorithm
AES128 -> Int -> Either SEIPDv2Failure Int
forall a b. b -> Either a b
Right Int
16
        SymmetricAlgorithm
AES192 -> Int -> Either SEIPDv2Failure Int
forall a b. b -> Either a b
Right Int
24
        SymmetricAlgorithm
AES256 -> Int -> Either SEIPDv2Failure Int
forall a b. b -> Either a b
Right Int
32
        SymmetricAlgorithm
_ -> SEIPDv2Failure -> Either SEIPDv2Failure Int
forall a b. a -> Either a b
Left (SEIPDv2Failure -> Either SEIPDv2Failure Int)
-> SEIPDv2Failure -> Either SEIPDv2Failure Int
forall a b. (a -> b) -> a -> b
$ SymmetricAlgorithm -> SEIPDv2Failure
SEIPDv2UnsupportedSymmetricAlgorithm SymmetricAlgorithm
symalgo

skeskV6Info
    :: SymmetricAlgorithm -> AEADAlgorithm -> B.ByteString
skeskV6Info :: SymmetricAlgorithm -> AEADAlgorithm -> ByteString
skeskV6Info SymmetricAlgorithm
symalgo AEADAlgorithm
aead = [Word8] -> ByteString
B.pack [Word8
0xc3, Word8
6, SymmetricAlgorithm -> Word8
forall a. FutureVal a => a -> Word8
fromFVal SymmetricAlgorithm
symalgo, AEADAlgorithm -> Word8
forall a. FutureVal a => a -> Word8
fromFVal AEADAlgorithm
aead]

deriveSKESK6KEK
    :: SymmetricAlgorithm
    -> AEADAlgorithm
    -> B.ByteString
    -> Either SEIPDv2Failure B.ByteString
deriveSKESK6KEK :: SymmetricAlgorithm
-> AEADAlgorithm -> ByteString -> Either SEIPDv2Failure ByteString
deriveSKESK6KEK SymmetricAlgorithm
symalgo AEADAlgorithm
aead ByteString
ikm = do
    keyLen <- SymmetricAlgorithm -> Either SEIPDv2Failure Int
seipdv2SymmetricKeySize SymmetricAlgorithm
symalgo
    let prk = forall a salt ikm.
(HashAlgorithm a, ByteArrayAccess salt, ByteArrayAccess ikm) =>
salt -> ikm -> PRK a
extract @CHA.SHA256 ByteString
B.empty ByteString
ikm
    pure (expand @CHA.SHA256 prk (skeskV6Info symalgo aead) keyLen)

encryptSKESK6SessionKey
    :: SymmetricAlgorithm
    -> AEADAlgorithm
    -> B.ByteString
    -> B.ByteString
    -> B.ByteString
    -> Either SEIPDv2Failure (B.ByteString, B.ByteString)
encryptSKESK6SessionKey :: SymmetricAlgorithm
-> AEADAlgorithm
-> ByteString
-> ByteString
-> ByteString
-> Either SEIPDv2Failure (ByteString, ByteString)
encryptSKESK6SessionKey SymmetricAlgorithm
symalgo AEADAlgorithm
aead ByteString
kek ByteString
iv ByteString
sessionKey = do
    (mode, nonceSize) <- AEADAlgorithm -> Either SEIPDv2Failure (AEADMode, Int)
aeadModeAndNonceSizeForSEIPDv2 AEADAlgorithm
aead
    if B.length iv /= nonceSize
        then Left SEIPDv2InvalidIVLength
        else
            first
                SEIPDv2CipherFailed
                ( withAEADCipher
                    symalgo
                    kek
                    ( \cipher
cipher ->
                        if AEADMode
mode AEADMode -> AEADMode -> Bool
forall a. Eq a => a -> a -> Bool
== AEADMode
CCT.AEAD_OCB
                            then do
                                (tag, ct) <-
                                    cipher
-> ByteString
-> ByteString
-> ByteString
-> Either CipherError (AuthTag, ByteString)
forall c e.
HOBlockCipher c =>
c
-> ByteString
-> ByteString
-> ByteString
-> Either e (AuthTag, ByteString)
encryptWithOCBRFC7253
                                        cipher
cipher
                                        ByteString
iv
                                        (SymmetricAlgorithm -> AEADAlgorithm -> ByteString
skeskV6Info SymmetricAlgorithm
symalgo AEADAlgorithm
aead)
                                        ByteString
sessionKey
                                pure (ct, authTagToBS tag)
                            else do
                                aeadCtx <-
                                    AEADMode
-> cipher -> ByteString -> Either CipherError (AEAD cipher)
forall cipher.
HOBlockCipher cipher =>
AEADMode
-> cipher -> ByteString -> Either CipherError (AEAD cipher)
aeadInit AEADMode
mode cipher
cipher ByteString
iv
                                let (tag, ciphertext) =
                                        aeadSimpleEncrypt
                                            aeadCtx
                                            (skeskV6Info symalgo aead)
                                            sessionKey
                                            16
                                pure (ciphertext, authTagToBS tag)
                    )
                )

decryptSKESK6SessionKey
    :: SymmetricAlgorithm
    -> AEADAlgorithm
    -> B.ByteString
    -> B.ByteString
    -> B.ByteString
    -> B.ByteString
    -> Either SEIPDv2Failure B.ByteString
decryptSKESK6SessionKey :: SymmetricAlgorithm
-> AEADAlgorithm
-> ByteString
-> ByteString
-> ByteString
-> ByteString
-> Either SEIPDv2Failure ByteString
decryptSKESK6SessionKey SymmetricAlgorithm
symalgo AEADAlgorithm
aead ByteString
kek ByteString
iv ByteString
ciphertext ByteString
tag = do
    (mode, nonceSize) <- AEADAlgorithm -> Either SEIPDv2Failure (AEADMode, Int)
aeadModeAndNonceSizeForSEIPDv2 AEADAlgorithm
aead
    if B.length iv /= nonceSize
        then Left SEIPDv2InvalidIVLength
        else
            first
                SEIPDv2CipherFailed
                ( withAEADCipher
                    symalgo
                    kek
                    ( \cipher
cipher ->
                        if AEADMode
mode AEADMode -> AEADMode -> Bool
forall a. Eq a => a -> a -> Bool
== AEADMode
CCT.AEAD_OCB
                            then
                                (ByteString
 -> ByteString
 -> ByteString
 -> ByteString
 -> ByteString
 -> ByteString
 -> CipherError)
-> cipher
-> ByteString
-> ByteString
-> ByteString
-> AuthTag
-> Either CipherError ByteString
forall c e.
HOBlockCipher c =>
(ByteString
 -> ByteString
 -> ByteString
 -> ByteString
 -> ByteString
 -> ByteString
 -> e)
-> c
-> ByteString
-> ByteString
-> ByteString
-> AuthTag
-> Either e ByteString
decryptWithOCBRFC7253With
                                    (\ByteString
_ ByteString
_ ByteString
_ ByteString
_ ByteString
_ ByteString
_ -> CipherError
CipherAEADAuthFailed)
                                    cipher
cipher
                                    ByteString
iv
                                    (SymmetricAlgorithm -> AEADAlgorithm -> ByteString
skeskV6Info SymmetricAlgorithm
symalgo AEADAlgorithm
aead)
                                    ByteString
ciphertext
                                    (ByteString -> AuthTag
mkAuthTag ByteString
tag)
                            else do
                                aeadCtx <-
                                    AEADMode
-> cipher -> ByteString -> Either CipherError (AEAD cipher)
forall cipher.
HOBlockCipher cipher =>
AEADMode
-> cipher -> ByteString -> Either CipherError (AEAD cipher)
aeadInit AEADMode
mode cipher
cipher ByteString
iv
                                let decrypted =
                                        AEAD cipher
-> ByteString -> ByteString -> AuthTag -> Maybe ByteString
forall cipher ct aad.
(HOBlockCipher cipher, ByteArray ct, ByteArrayAccess aad) =>
AEAD cipher -> aad -> ct -> AuthTag -> Maybe ct
forall ct aad.
(ByteArray ct, ByteArrayAccess aad) =>
AEAD cipher -> aad -> ct -> AuthTag -> Maybe ct
aeadSimpleDecrypt
                                            AEAD cipher
aeadCtx
                                            (SymmetricAlgorithm -> AEADAlgorithm -> ByteString
skeskV6Info SymmetricAlgorithm
symalgo AEADAlgorithm
aead)
                                            ByteString
ciphertext
                                            (ByteString -> AuthTag
mkAuthTag ByteString
tag)
                                note
                                    CipherAEADDecryptFailed
                                    decrypted
                    )
                )

authTagToBS :: CCT.AuthTag -> B.ByteString
authTagToBS :: AuthTag -> ByteString
authTagToBS = Bytes -> ByteString
forall bin bout.
(ByteArrayAccess bin, ByteArray bout) =>
bin -> bout
BA.convert (Bytes -> ByteString)
-> (AuthTag -> Bytes) -> AuthTag -> ByteString
forall b c a. (b -> c) -> (a -> b) -> a -> c
. AuthTag -> Bytes
CCT.unAuthTag

mkAuthTag :: B.ByteString -> CCT.AuthTag
mkAuthTag :: ByteString -> AuthTag
mkAuthTag = Bytes -> AuthTag
CCT.AuthTag (Bytes -> AuthTag)
-> (ByteString -> Bytes) -> ByteString -> AuthTag
forall b c a. (b -> c) -> (a -> b) -> a -> c
. ByteString -> Bytes
forall bin bout.
(ByteArrayAccess bin, ByteArray bout) =>
bin -> bout
BA.convert