{-# LANGUAGE FlexibleContexts #-}

-- | Instantiate 'ObjectPoolReader' and 'ObjectPoolWriter' using Peras
-- certificates from the 'PerasCertDB' (or the 'ChainDB' which is wrapping the
-- 'PerasCertDB').
module Ouroboros.Consensus.MiniProtocol.ObjectDiffusion.ObjectPool.PerasCert
  ( makePerasCertPoolReaderFromChainDB
  , makePerasCertPoolWriterFromChainDB

    -- * For testing purposes
  , makeTestPerasCertPoolReaderFromCertDB
  , makeTestPerasCertPoolWriterFromCertDB
  ) where

import Data.Foldable (traverse_)
import Data.Map (Map)
import qualified Data.Map as Map
import qualified Data.Set as Set
import Ouroboros.Consensus.Block.SupportsPeras
  ( BlockSupportsPeras (..)
  , IsPerasCert (..)
  , PerasRoundNo
  , ValidatedPerasCert (..)
  )
import Ouroboros.Consensus.BlockchainTime.WallClock.Types
  ( SystemTime (..)
  , WithArrivalTime (..)
  )
import Ouroboros.Consensus.MiniProtocol.ObjectDiffusion.ObjectPool.API
  ( ObjectPoolReader (..)
  , ObjectPoolWriter (..)
  )
import Ouroboros.Consensus.Peras.Context
  ( PerasEpochContextResolverHandle
  , verifyPerasCertWithHandle
  )
import Ouroboros.Consensus.Storage.ChainDB.API (ChainDB)
import qualified Ouroboros.Consensus.Storage.ChainDB.API as ChainDB
import Ouroboros.Consensus.Storage.PerasCertDB.API
  ( PerasCertDB
  , PerasCertTicketNo
  )
import qualified Ouroboros.Consensus.Storage.PerasCertDB.API as PerasCertDB
import Ouroboros.Consensus.Util.IOLike (IOLike, MonadSTM (..))

-- | TODO: replace by `Data.Map.take` as soon as we move to GHC 9.8
takeAscMap :: Int -> Map k v -> Map k v
takeAscMap :: forall k v. Int -> Map k v -> Map k v
takeAscMap Int
n = [(k, v)] -> Map k v
forall k a. [(k, a)] -> Map k a
Map.fromDistinctAscList ([(k, v)] -> Map k v)
-> (Map k v -> [(k, v)]) -> Map k v -> Map k v
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Int -> [(k, v)] -> [(k, v)]
forall a. Int -> [a] -> [a]
take Int
n ([(k, v)] -> [(k, v)])
-> (Map k v -> [(k, v)]) -> Map k v -> [(k, v)]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Map k v -> [(k, v)]
forall k a. Map k a -> [(k, a)]
Map.toAscList

-------------------------------------------------------------------------------
-- Readers
-------------------------------------------------------------------------------

-- | Internal helper: create a pool reader from a @getCertsAfter@ function.
makePerasCertPoolReader ::
  ( IOLike m
  , IsPerasCert (PerasCert blk) blk
  ) =>
  ( PerasCertTicketNo ->
    STM m (Map PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk))))
  ) ->
  ObjectPoolReader PerasRoundNo (PerasCert blk) PerasCertTicketNo m
makePerasCertPoolReader :: forall (m :: * -> *) blk.
(IOLike m, IsPerasCert (PerasCert blk) blk) =>
(PerasCertTicketNo
 -> STM
      m
      (Map
         PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk)))))
-> ObjectPoolReader
     PerasRoundNo (PerasCert blk) PerasCertTicketNo m
makePerasCertPoolReader PerasCertTicketNo
-> STM
     m
     (Map
        PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk))))
getCertsAfterSTM =
  ObjectPoolReader
    { oprObjectId :: PerasCert blk -> PerasRoundNo
oprObjectId = PerasCert blk -> PerasRoundNo
forall cert blk. IsPerasCert cert blk => cert -> PerasRoundNo
getPerasCertRound
    , oprZeroTicketNo :: PerasCertTicketNo
oprZeroTicketNo = PerasCertTicketNo
PerasCertDB.zeroPerasCertTicketNo
    , oprObjectsAfter :: PerasCertTicketNo
-> Word64
-> STM m (Maybe (m (Map PerasCertTicketNo (PerasCert blk))))
oprObjectsAfter = \PerasCertTicketNo
lastKnown Word64
limit -> do
        certsAfterLastKnownNoLimit <- PerasCertTicketNo
-> STM
     m
     (Map
        PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk))))
getCertsAfterSTM PerasCertTicketNo
lastKnown
        if Map.null certsAfterLastKnownNoLimit
          then pure Nothing
          else pure . Just $ do
            let certsAfterLastKnown = Int
-> Map
     PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk)))
-> Map
     PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk)))
forall k v. Int -> Map k v -> Map k v
takeAscMap (Word64 -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral Word64
limit) Map
  PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk)))
certsAfterLastKnownNoLimit
            traverse
              (\m (WithArrivalTime (ValidatedPerasCert blk))
loadCertAction -> (ValidatedPerasCert blk -> PerasCert blk
forall blk. ValidatedPerasCert blk -> PerasCert blk
vpcCert (ValidatedPerasCert blk -> PerasCert blk)
-> (WithArrivalTime (ValidatedPerasCert blk)
    -> ValidatedPerasCert blk)
-> WithArrivalTime (ValidatedPerasCert blk)
-> PerasCert blk
forall b c a. (b -> c) -> (a -> b) -> a -> c
. WithArrivalTime (ValidatedPerasCert blk) -> ValidatedPerasCert blk
forall a. WithArrivalTime a -> a
forgetArrivalTime) (WithArrivalTime (ValidatedPerasCert blk) -> PerasCert blk)
-> m (WithArrivalTime (ValidatedPerasCert blk))
-> m (PerasCert blk)
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> m (WithArrivalTime (ValidatedPerasCert blk))
loadCertAction)
              certsAfterLastKnown
    }

makeTestPerasCertPoolReaderFromCertDB ::
  ( IOLike m
  , IsPerasCert (PerasCert blk) blk
  ) =>
  PerasCertDB m blk ->
  ObjectPoolReader PerasRoundNo (PerasCert blk) PerasCertTicketNo m
makeTestPerasCertPoolReaderFromCertDB :: forall (m :: * -> *) blk.
(IOLike m, IsPerasCert (PerasCert blk) blk) =>
PerasCertDB m blk
-> ObjectPoolReader
     PerasRoundNo (PerasCert blk) PerasCertTicketNo m
makeTestPerasCertPoolReaderFromCertDB PerasCertDB m blk
perasCertDB =
  (PerasCertTicketNo
 -> STM
      m
      (Map
         PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk)))))
-> ObjectPoolReader
     PerasRoundNo (PerasCert blk) PerasCertTicketNo m
forall (m :: * -> *) blk.
(IOLike m, IsPerasCert (PerasCert blk) blk) =>
(PerasCertTicketNo
 -> STM
      m
      (Map
         PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk)))))
-> ObjectPoolReader
     PerasRoundNo (PerasCert blk) PerasCertTicketNo m
makePerasCertPoolReader
    (PerasCertDB m blk
-> PerasCertTicketNo
-> STM
     m
     (Map
        PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk))))
forall (m :: * -> *) blk.
PerasCertDB m blk
-> PerasCertTicketNo
-> STM
     m
     (Map
        PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk))))
PerasCertDB.getCertsAfter PerasCertDB m blk
perasCertDB)

makePerasCertPoolReaderFromChainDB ::
  ( IOLike m
  , IsPerasCert (PerasCert blk) blk
  ) =>
  ChainDB m blk ->
  ObjectPoolReader PerasRoundNo (PerasCert blk) PerasCertTicketNo m
makePerasCertPoolReaderFromChainDB :: forall (m :: * -> *) blk.
(IOLike m, IsPerasCert (PerasCert blk) blk) =>
ChainDB m blk
-> ObjectPoolReader
     PerasRoundNo (PerasCert blk) PerasCertTicketNo m
makePerasCertPoolReaderFromChainDB ChainDB m blk
chainDB =
  (PerasCertTicketNo
 -> STM
      m
      (Map
         PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk)))))
-> ObjectPoolReader
     PerasRoundNo (PerasCert blk) PerasCertTicketNo m
forall (m :: * -> *) blk.
(IOLike m, IsPerasCert (PerasCert blk) blk) =>
(PerasCertTicketNo
 -> STM
      m
      (Map
         PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk)))))
-> ObjectPoolReader
     PerasRoundNo (PerasCert blk) PerasCertTicketNo m
makePerasCertPoolReader
    (ChainDB m blk
-> PerasCertTicketNo
-> STM
     m
     (Map
        PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk))))
forall (m :: * -> *) blk.
ChainDB m blk
-> PerasCertTicketNo
-> STM
     m
     (Map
        PerasCertTicketNo (m (WithArrivalTime (ValidatedPerasCert blk))))
ChainDB.getPerasCertsAfter ChainDB m blk
chainDB)

-------------------------------------------------------------------------------
-- Writers
-------------------------------------------------------------------------------

-- | Create a pool writer directly from a 'PerasCertDB'. This is mostly meant
-- for tests against the 'PerasCertDB' in isolation; for actual production use,
-- see 'makePerasCertPoolWriterFromChainDB' which creates a pool writer from the
-- 'ChainDB' with proper handling of chain selection side-effects.
makeTestPerasCertPoolWriterFromCertDB ::
  ( IOLike m
  , BlockSupportsPeras blk
  ) =>
  SystemTime m ->
  PerasCertDB m blk ->
  PerasEpochContextResolverHandle m blk ->
  ObjectPoolWriter PerasRoundNo (PerasCert blk) m
makeTestPerasCertPoolWriterFromCertDB :: forall (m :: * -> *) blk.
(IOLike m, BlockSupportsPeras blk) =>
SystemTime m
-> PerasCertDB m blk
-> PerasEpochContextResolverHandle m blk
-> ObjectPoolWriter PerasRoundNo (PerasCert blk) m
makeTestPerasCertPoolWriterFromCertDB SystemTime m
systemTime PerasCertDB m blk
perasCertDB PerasEpochContextResolverHandle m blk
resolverHandle =
  ObjectPoolWriter
    { opwObjectId :: PerasCert blk -> PerasRoundNo
opwObjectId = PerasCert blk -> PerasRoundNo
forall cert blk. IsPerasCert cert blk => cert -> PerasRoundNo
getPerasCertRound
    , opwAddObjects :: [PerasCert blk] -> m ()
opwAddObjects = \[PerasCert blk]
certs -> do
        now <- SystemTime m -> m RelativeTime
forall (m :: * -> *). SystemTime m -> m RelativeTime
systemTimeCurrent SystemTime m
systemTime
        atomically $ do
          alreadyInDb <- PerasCertDB.getCertIds perasCertDB
          let certsNotAlreadyInDb = (PerasCert blk -> Bool) -> [PerasCert blk] -> [PerasCert blk]
forall a. (a -> Bool) -> [a] -> [a]
filter ((PerasRoundNo -> Set PerasRoundNo -> Bool
forall a. Ord a => a -> Set a -> Bool
`Set.notMember` Set PerasRoundNo
alreadyInDb) (PerasRoundNo -> Bool)
-> (PerasCert blk -> PerasRoundNo) -> PerasCert blk -> Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. PerasCert blk -> PerasRoundNo
forall cert blk. IsPerasCert cert blk => cert -> PerasRoundNo
getPerasCertRound) [PerasCert blk]
certs
          validatedCerts <- traverse (verifyPerasCertWithHandle resolverHandle) certsNotAlreadyInDb
          -- Some certs are invalid => reject the whole batch
          --
          -- NOTE: we could combine the two 'traverse' operations into one in
          -- which case any validated cert would be immediately added no matter
          -- what is the validity of the other certs in the batch.
          traverse_ (PerasCertDB.addCert perasCertDB . WithArrivalTime now) validatedCerts
    , opwHasObject :: STM m (PerasRoundNo -> Bool)
opwHasObject = do
        certIds <- PerasCertDB m blk -> STM m (Set PerasRoundNo)
forall (m :: * -> *) blk.
PerasCertDB m blk -> STM m (Set PerasRoundNo)
PerasCertDB.getCertIds PerasCertDB m blk
perasCertDB
        pure $ \PerasRoundNo
roundNo -> PerasRoundNo -> Set PerasRoundNo -> Bool
forall a. Ord a => a -> Set a -> Bool
Set.member PerasRoundNo
roundNo Set PerasRoundNo
certIds
    }

-- | Create a pool writer from the 'ChainDB'. This properly handles any needed
-- chain selection side-effects.
makePerasCertPoolWriterFromChainDB ::
  ( IOLike m
  , BlockSupportsPeras blk
  ) =>
  SystemTime m ->
  ChainDB m blk ->
  ObjectPoolWriter PerasRoundNo (PerasCert blk) m
makePerasCertPoolWriterFromChainDB :: forall (m :: * -> *) blk.
(IOLike m, BlockSupportsPeras blk) =>
SystemTime m
-> ChainDB m blk -> ObjectPoolWriter PerasRoundNo (PerasCert blk) m
makePerasCertPoolWriterFromChainDB SystemTime m
systemTime ChainDB m blk
chainDB =
  let resolverHandle :: PerasEpochContextResolverHandle m blk
resolverHandle = ChainDB m blk -> PerasEpochContextResolverHandle m blk
forall (m :: * -> *) blk.
ChainDB m blk -> PerasEpochContextResolverHandle m blk
ChainDB.getPerasEpochContextResolverHandle ChainDB m blk
chainDB
   in ObjectPoolWriter
        { opwObjectId :: PerasCert blk -> PerasRoundNo
opwObjectId = PerasCert blk -> PerasRoundNo
forall cert blk. IsPerasCert cert blk => cert -> PerasRoundNo
getPerasCertRound
        , opwAddObjects :: [PerasCert blk] -> m ()
opwAddObjects = \[PerasCert blk]
certs -> do
            now <- SystemTime m -> m RelativeTime
forall (m :: * -> *). SystemTime m -> m RelativeTime
systemTimeCurrent SystemTime m
systemTime
            validatedCerts <- atomically $ do
              alreadyInDb <- ChainDB.getPerasCertIds chainDB
              let certsNotAlreadyInDb = (PerasCert blk -> Bool) -> [PerasCert blk] -> [PerasCert blk]
forall a. (a -> Bool) -> [a] -> [a]
filter ((PerasRoundNo -> Set PerasRoundNo -> Bool
forall a. Ord a => a -> Set a -> Bool
`Set.notMember` Set PerasRoundNo
alreadyInDb) (PerasRoundNo -> Bool)
-> (PerasCert blk -> PerasRoundNo) -> PerasCert blk -> Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. PerasCert blk -> PerasRoundNo
forall cert blk. IsPerasCert cert blk => cert -> PerasRoundNo
getPerasCertRound) [PerasCert blk]
certs
              traverse (verifyPerasCertWithHandle resolverHandle) certsNotAlreadyInDb
            -- Some certs are invalid => reject the whole batch
            --
            -- NOTE: we could combine the two 'traverse' operations into one in
            -- which case any validated cert would be immediately added no
            -- matter what is the validity of the other certs in the batch.
            traverse_ (ChainDB.addPerasCertAsync chainDB . WithArrivalTime now) validatedCerts
        , opwHasObject :: STM m (PerasRoundNo -> Bool)
opwHasObject = do
            certIds <- ChainDB m blk -> STM m (Set PerasRoundNo)
forall (m :: * -> *) blk. ChainDB m blk -> STM m (Set PerasRoundNo)
ChainDB.getPerasCertIds ChainDB m blk
chainDB
            pure $ \PerasRoundNo
roundNo -> PerasRoundNo -> Set PerasRoundNo -> Bool
forall a. Ord a => a -> Set a -> Bool
Set.member PerasRoundNo
roundNo Set PerasRoundNo
certIds
        }