-- Message.hs: conduit-backed OpenPGP message helpers
-- Copyright © 2026  Clint Adams
-- This software is released under the terms of the Expat license.
-- (See the LICENSE file).

module Data.Conduit.OpenPGP.Message
  ( VerificationPolicy(..)
  , VerificationOptions(..)
  , defaultVerificationOptions
  , verifyMessagePackets
  , verifyMessage
  , VerificationMode(..)
  ) where

import qualified Data.ByteString.Lazy as BL
import Data.Conduit ((.|), runConduitPure)
import qualified Data.Conduit.List as CL
import Data.Time.Clock (UTCTime)

import Codec.Encryption.OpenPGP.Compression (decompressPkt)
import Codec.Encryption.OpenPGP.Policy (defaultVerificationDefaults, verificationDefaultStreaming, verificationDefaultStrict)
import Codec.Encryption.OpenPGP.Serialize (parsePkts)
import Codec.Encryption.OpenPGP.Signatures (VerificationError)
import Codec.Encryption.OpenPGP.Types
import Data.Conduit.OpenPGP.Verify
  ( VerificationMode(..)
  , VerificationModeW(..)
  , verifyPacketsWithModeTyped
  , verifyPacketsBatch
  )

data VerificationPolicy
  = VerifyInformational
  | VerifyStrict
  deriving (VerificationPolicy -> VerificationPolicy -> Bool
(VerificationPolicy -> VerificationPolicy -> Bool)
-> (VerificationPolicy -> VerificationPolicy -> Bool)
-> Eq VerificationPolicy
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: VerificationPolicy -> VerificationPolicy -> Bool
== :: VerificationPolicy -> VerificationPolicy -> Bool
$c/= :: VerificationPolicy -> VerificationPolicy -> Bool
/= :: VerificationPolicy -> VerificationPolicy -> Bool
Eq, Int -> VerificationPolicy -> ShowS
[VerificationPolicy] -> ShowS
VerificationPolicy -> String
(Int -> VerificationPolicy -> ShowS)
-> (VerificationPolicy -> String)
-> ([VerificationPolicy] -> ShowS)
-> Show VerificationPolicy
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> VerificationPolicy -> ShowS
showsPrec :: Int -> VerificationPolicy -> ShowS
$cshow :: VerificationPolicy -> String
show :: VerificationPolicy -> String
$cshowList :: [VerificationPolicy] -> ShowS
showList :: [VerificationPolicy] -> ShowS
Show)

data VerificationOptions = VerificationOptions
  { VerificationOptions -> VerificationPolicy
verificationPolicy :: VerificationPolicy
  , VerificationOptions -> VerificationMode
verificationMode :: VerificationMode
  , VerificationOptions -> Maybe UTCTime
verificationTime :: Maybe UTCTime
  }
  deriving (VerificationOptions -> VerificationOptions -> Bool
(VerificationOptions -> VerificationOptions -> Bool)
-> (VerificationOptions -> VerificationOptions -> Bool)
-> Eq VerificationOptions
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: VerificationOptions -> VerificationOptions -> Bool
== :: VerificationOptions -> VerificationOptions -> Bool
$c/= :: VerificationOptions -> VerificationOptions -> Bool
/= :: VerificationOptions -> VerificationOptions -> Bool
Eq, Int -> VerificationOptions -> ShowS
[VerificationOptions] -> ShowS
VerificationOptions -> String
(Int -> VerificationOptions -> ShowS)
-> (VerificationOptions -> String)
-> ([VerificationOptions] -> ShowS)
-> Show VerificationOptions
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> VerificationOptions -> ShowS
showsPrec :: Int -> VerificationOptions -> ShowS
$cshow :: VerificationOptions -> String
show :: VerificationOptions -> String
$cshowList :: [VerificationOptions] -> ShowS
showList :: [VerificationOptions] -> ShowS
Show)

defaultVerificationOptions :: VerificationOptions
defaultVerificationOptions :: VerificationOptions
defaultVerificationOptions =
  VerificationOptions
    { verificationPolicy :: VerificationPolicy
verificationPolicy =
        if VerificationDefaults -> Bool
verificationDefaultStrict VerificationDefaults
defaultVerificationDefaults
          then VerificationPolicy
VerifyStrict
          else VerificationPolicy
VerifyInformational
    , verificationMode :: VerificationMode
verificationMode =
        if VerificationDefaults -> Bool
verificationDefaultStreaming VerificationDefaults
defaultVerificationDefaults
          then VerificationMode
VerificationStreaming
          else VerificationMode
VerificationBatch
    , verificationTime :: Maybe UTCTime
verificationTime = Maybe UTCTime
forall a. Maybe a
Nothing
    }

verifyMessagePackets ::
     VerificationOptions
  -> PublicKeyring
  -> [Pkt]
  -> [Either VerificationError Verification]
verifyMessagePackets :: VerificationOptions
-> PublicKeyring
-> [Pkt]
-> [Either VerificationError Verification]
verifyMessagePackets VerificationOptions
options PublicKeyring
keyring [Pkt]
packets =
  VerificationPolicy
-> [Either VerificationError Verification]
-> [Either VerificationError Verification]
applyVerificationPolicy (VerificationOptions -> VerificationPolicy
verificationPolicy VerificationOptions
options) [Either VerificationError Verification]
rawResults
  where
    rawResults :: [Either VerificationError Verification]
rawResults =
      case VerificationOptions -> VerificationMode
verificationMode VerificationOptions
options of
        VerificationMode
VerificationBatch ->
          PublicKeyring
-> Maybe UTCTime
-> [Pkt]
-> [Either VerificationError Verification]
verifyPacketsBatch PublicKeyring
keyring (VerificationOptions -> Maybe UTCTime
verificationTime VerificationOptions
options) [Pkt]
packets
        VerificationMode
VerificationStreaming ->
          ConduitT () Void Identity [Either VerificationError Verification]
-> [Either VerificationError Verification]
forall r. ConduitT () Void Identity r -> r
runConduitPure (ConduitT () Void Identity [Either VerificationError Verification]
 -> [Either VerificationError Verification])
-> ConduitT
     () Void Identity [Either VerificationError Verification]
-> [Either VerificationError Verification]
forall a b. (a -> b) -> a -> b
$
          [Pkt] -> ConduitT () Pkt Identity ()
forall (m :: * -> *) a i. Monad m => [a] -> ConduitT i a m ()
CL.sourceList [Pkt]
packets ConduitT () Pkt Identity ()
-> ConduitT
     Pkt Void Identity [Either VerificationError Verification]
-> ConduitT
     () Void Identity [Either VerificationError Verification]
forall (m :: * -> *) a b c r.
Monad m =>
ConduitT a b m () -> ConduitT b c m r -> ConduitT a c m r
.|
          VerificationModeW 'VerificationStreaming
-> PublicKeyring
-> Maybe UTCTime
-> ConduitT Pkt (Either VerificationError Verification) Identity ()
forall (m :: * -> *) (mode :: VerificationMode).
Monad m =>
VerificationModeW mode
-> PublicKeyring
-> Maybe UTCTime
-> ConduitT Pkt (Either VerificationError Verification) m ()
verifyPacketsWithModeTyped VerificationModeW 'VerificationStreaming
VerificationStreamingW
            PublicKeyring
keyring
            (VerificationOptions -> Maybe UTCTime
verificationTime VerificationOptions
options) ConduitT Pkt (Either VerificationError Verification) Identity ()
-> ConduitT
     (Either VerificationError Verification)
     Void
     Identity
     [Either VerificationError Verification]
-> ConduitT
     Pkt Void Identity [Either VerificationError Verification]
forall (m :: * -> *) a b c r.
Monad m =>
ConduitT a b m () -> ConduitT b c m r -> ConduitT a c m r
.|
          ConduitT
  (Either VerificationError Verification)
  Void
  Identity
  [Either VerificationError Verification]
forall (m :: * -> *) a o. Monad m => ConduitT a o m [a]
CL.consume

verifyMessage ::
     VerificationOptions
  -> PublicKeyring
  -> BL.ByteString
  -> [Either VerificationError Verification]
verifyMessage :: VerificationOptions
-> PublicKeyring
-> ByteString
-> [Either VerificationError Verification]
verifyMessage VerificationOptions
options PublicKeyring
keyring ByteString
signedMessage =
  VerificationOptions
-> PublicKeyring
-> [Pkt]
-> [Either VerificationError Verification]
verifyMessagePackets
    VerificationOptions
options
    PublicKeyring
keyring
    ((Pkt -> [Pkt]) -> [Pkt] -> [Pkt]
forall (t :: * -> *) a b. Foldable t => (a -> [b]) -> t a -> [b]
concatMap ((CompressionError -> [Pkt])
-> ([Pkt] -> [Pkt]) -> Either CompressionError [Pkt] -> [Pkt]
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either ([Pkt] -> CompressionError -> [Pkt]
forall a b. a -> b -> a
const []) [Pkt] -> [Pkt]
forall a. a -> a
id (Either CompressionError [Pkt] -> [Pkt])
-> (Pkt -> Either CompressionError [Pkt]) -> Pkt -> [Pkt]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Pkt -> Either CompressionError [Pkt]
decompressPkt) (ByteString -> [Pkt]
parsePkts ByteString
signedMessage))

applyVerificationPolicy ::
     VerificationPolicy
  -> [Either VerificationError Verification]
  -> [Either VerificationError Verification]
applyVerificationPolicy :: VerificationPolicy
-> [Either VerificationError Verification]
-> [Either VerificationError Verification]
applyVerificationPolicy VerificationPolicy
VerifyInformational [Either VerificationError Verification]
results = [Either VerificationError Verification]
results
applyVerificationPolicy VerificationPolicy
VerifyStrict [Either VerificationError Verification]
results =
  (VerificationError -> [Either VerificationError Verification])
-> ([Verification] -> [Either VerificationError Verification])
-> Either VerificationError [Verification]
-> [Either VerificationError Verification]
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (Either VerificationError Verification
-> [Either VerificationError Verification]
forall a. a -> [a]
forall (f :: * -> *) a. Applicative f => a -> f a
pure (Either VerificationError Verification
 -> [Either VerificationError Verification])
-> (VerificationError -> Either VerificationError Verification)
-> VerificationError
-> [Either VerificationError Verification]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. VerificationError -> Either VerificationError Verification
forall a b. a -> Either a b
Left) (Verification -> Either VerificationError Verification
forall a b. b -> Either a b
Right (Verification -> Either VerificationError Verification)
-> [Verification] -> [Either VerificationError Verification]
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$>) ([Either VerificationError Verification]
-> Either VerificationError [Verification]
forall (t :: * -> *) (m :: * -> *) a.
(Traversable t, Monad m) =>
t (m a) -> m (t a)
forall (m :: * -> *) a. Monad m => [m a] -> m [a]
sequence [Either VerificationError Verification]
results)