{-# LANGUAGE FlexibleContexts #-}

-- | Instantiate 'ObjectPoolReader' and 'ObjectPoolWriter' using Peras
-- votes from the 'PerasVoteDB' (or the 'ChainDB' which is wrapping the
-- 'PerasVoteDB').
module Ouroboros.Consensus.MiniProtocol.ObjectDiffusion.ObjectPool.PerasVote
  ( makePerasVotePoolReaderFromChainDB
  , makePerasVotePoolWriterFromChainDB

    -- * For testing purposes
  , makeTestPerasVotePoolReaderFromVoteDB
  , makeTestPerasVotePoolWriterFromVoteDB
  ) 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 (..)
  , IsPerasVote
  , PerasVoteId
  , ValidatedPerasVote (..)
  , getPerasVoteId
  )
import Ouroboros.Consensus.BlockchainTime.WallClock.Types
  ( SystemTime (..)
  , WithArrivalTime (..)
  )
import Ouroboros.Consensus.MiniProtocol.ObjectDiffusion.ObjectPool.API
  ( ObjectPoolReader (..)
  , ObjectPoolWriter (..)
  )
import Ouroboros.Consensus.Peras.Context
  ( PerasEpochContextResolverHandle
  , verifyPerasVoteWithHandle
  )
import Ouroboros.Consensus.Storage.ChainDB.API (ChainDB)
import qualified Ouroboros.Consensus.Storage.ChainDB.API as ChainDB
import Ouroboros.Consensus.Storage.PerasVoteDB.API
  ( PerasVoteDB
  , PerasVoteTicketNo
  )
import qualified Ouroboros.Consensus.Storage.PerasVoteDB.API as PerasVoteDB
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 @getVotesAfter@ function.
makePerasVotePoolReader ::
  ( IOLike m
  , IsPerasVote (PerasVote blk) blk
  ) =>
  ( PerasVoteTicketNo ->
    STM m (Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk)))
  ) ->
  ObjectPoolReader PerasVoteId (PerasVote blk) PerasVoteTicketNo m
makePerasVotePoolReader :: forall (m :: * -> *) blk.
(IOLike m, IsPerasVote (PerasVote blk) blk) =>
(PerasVoteTicketNo
 -> STM
      m
      (Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk))))
-> ObjectPoolReader PerasVoteId (PerasVote blk) PerasVoteTicketNo m
makePerasVotePoolReader PerasVoteTicketNo
-> STM
     m
     (Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk)))
getVotesAfterSTM =
  ObjectPoolReader
    { oprObjectId :: PerasVote blk -> PerasVoteId
oprObjectId = PerasVote blk -> PerasVoteId
forall vote blk. IsPerasVote vote blk => vote -> PerasVoteId
getPerasVoteId
    , oprZeroTicketNo :: PerasVoteTicketNo
oprZeroTicketNo = PerasVoteTicketNo
PerasVoteDB.zeroPerasVoteTicketNo
    , oprObjectsAfter :: PerasVoteTicketNo
-> Word64
-> STM m (Maybe (m (Map PerasVoteTicketNo (PerasVote blk))))
oprObjectsAfter = \PerasVoteTicketNo
lastKnown Word64
limit -> do
        votesAfterLastKnownNoLimit <- PerasVoteTicketNo
-> STM
     m
     (Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk)))
getVotesAfterSTM PerasVoteTicketNo
lastKnown
        if Map.null votesAfterLastKnownNoLimit
          then pure Nothing
          else pure . Just $ do
            let votesAfterLastKnown = Int
-> Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk))
-> Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote 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 PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk))
votesAfterLastKnownNoLimit
            pure $ Map.map (vpvVote . forgetArrivalTime) votesAfterLastKnown
    }

makeTestPerasVotePoolReaderFromVoteDB ::
  ( IOLike m
  , IsPerasVote (PerasVote blk) blk
  ) =>
  PerasVoteDB m blk ->
  ObjectPoolReader PerasVoteId (PerasVote blk) PerasVoteTicketNo m
makeTestPerasVotePoolReaderFromVoteDB :: forall (m :: * -> *) blk.
(IOLike m, IsPerasVote (PerasVote blk) blk) =>
PerasVoteDB m blk
-> ObjectPoolReader PerasVoteId (PerasVote blk) PerasVoteTicketNo m
makeTestPerasVotePoolReaderFromVoteDB PerasVoteDB m blk
perasVoteDB =
  (PerasVoteTicketNo
 -> STM
      m
      (Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk))))
-> ObjectPoolReader PerasVoteId (PerasVote blk) PerasVoteTicketNo m
forall (m :: * -> *) blk.
(IOLike m, IsPerasVote (PerasVote blk) blk) =>
(PerasVoteTicketNo
 -> STM
      m
      (Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk))))
-> ObjectPoolReader PerasVoteId (PerasVote blk) PerasVoteTicketNo m
makePerasVotePoolReader
    (PerasVoteDB m blk
-> PerasVoteTicketNo
-> STM
     m
     (Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk)))
forall (m :: * -> *) blk.
PerasVoteDB m blk
-> PerasVoteTicketNo
-> STM
     m
     (Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk)))
PerasVoteDB.getVotesAfter PerasVoteDB m blk
perasVoteDB)

makePerasVotePoolReaderFromChainDB ::
  ( IOLike m
  , IsPerasVote (PerasVote blk) blk
  ) =>
  ChainDB m blk ->
  ObjectPoolReader PerasVoteId (PerasVote blk) PerasVoteTicketNo m
makePerasVotePoolReaderFromChainDB :: forall (m :: * -> *) blk.
(IOLike m, IsPerasVote (PerasVote blk) blk) =>
ChainDB m blk
-> ObjectPoolReader PerasVoteId (PerasVote blk) PerasVoteTicketNo m
makePerasVotePoolReaderFromChainDB ChainDB m blk
chainDB =
  (PerasVoteTicketNo
 -> STM
      m
      (Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk))))
-> ObjectPoolReader PerasVoteId (PerasVote blk) PerasVoteTicketNo m
forall (m :: * -> *) blk.
(IOLike m, IsPerasVote (PerasVote blk) blk) =>
(PerasVoteTicketNo
 -> STM
      m
      (Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk))))
-> ObjectPoolReader PerasVoteId (PerasVote blk) PerasVoteTicketNo m
makePerasVotePoolReader
    (ChainDB m blk
-> PerasVoteTicketNo
-> STM
     m
     (Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk)))
forall (m :: * -> *) blk.
ChainDB m blk
-> PerasVoteTicketNo
-> STM
     m
     (Map PerasVoteTicketNo (WithArrivalTime (ValidatedPerasVote blk)))
ChainDB.getPerasVotesAfter ChainDB m blk
chainDB)

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

-- | Create a pool writer directly from a 'PerasVoteDB'.
-- In particular, the result of 'addVote' is ignored, so any produced cert will
-- have to be handled manually by another mean. This function is mostly meant
-- for tests against the 'PerasVoteDB' in isolation; for actual production use,
-- see 'makePerasVotePoolWriterFromChainDB' which creates a pool writer from the
-- 'ChainDB' and thus properly handles the produced certs.
makeTestPerasVotePoolWriterFromVoteDB ::
  ( IOLike m
  , BlockSupportsPeras blk
  ) =>
  SystemTime m ->
  PerasVoteDB m blk ->
  PerasEpochContextResolverHandle m blk ->
  ObjectPoolWriter PerasVoteId (PerasVote blk) m
makeTestPerasVotePoolWriterFromVoteDB :: forall (m :: * -> *) blk.
(IOLike m, BlockSupportsPeras blk) =>
SystemTime m
-> PerasVoteDB m blk
-> PerasEpochContextResolverHandle m blk
-> ObjectPoolWriter PerasVoteId (PerasVote blk) m
makeTestPerasVotePoolWriterFromVoteDB SystemTime m
systemTime PerasVoteDB m blk
perasVoteDB PerasEpochContextResolverHandle m blk
resolverHandle =
  ObjectPoolWriter
    { opwObjectId :: PerasVote blk -> PerasVoteId
opwObjectId = PerasVote blk -> PerasVoteId
forall vote blk. IsPerasVote vote blk => vote -> PerasVoteId
getPerasVoteId
    , opwAddObjects :: [PerasVote blk] -> m ()
opwAddObjects = \[PerasVote blk]
votes -> do
        now <- SystemTime m -> m RelativeTime
forall (m :: * -> *). SystemTime m -> m RelativeTime
systemTimeCurrent SystemTime m
systemTime
        atomically $ do
          alreadyInDb <- PerasVoteDB.getVoteIds perasVoteDB
          let votesNotAlreadyInDb = (PerasVote blk -> Bool) -> [PerasVote blk] -> [PerasVote blk]
forall a. (a -> Bool) -> [a] -> [a]
filter ((PerasVoteId -> Set PerasVoteId -> Bool
forall a. Ord a => a -> Set a -> Bool
`Set.notMember` Set PerasVoteId
alreadyInDb) (PerasVoteId -> Bool)
-> (PerasVote blk -> PerasVoteId) -> PerasVote blk -> Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. PerasVote blk -> PerasVoteId
forall vote blk. IsPerasVote vote blk => vote -> PerasVoteId
getPerasVoteId) [PerasVote blk]
votes
          validatedVotes <- traverse (verifyPerasVoteWithHandle resolverHandle) votesNotAlreadyInDb
          -- Some votes are invalid => reject the whole batch
          --
          -- NOTE: we could combine the two 'traverse' operations into one in
          -- which case any validated vote would be immediately added no matter
          -- what is the validity of the other votes in the batch.
          traverse_ (PerasVoteDB.addVote perasVoteDB . WithArrivalTime now) validatedVotes
    , opwHasObject :: STM m (PerasVoteId -> Bool)
opwHasObject = do
        voteIds <- PerasVoteDB m blk -> STM m (Set PerasVoteId)
forall (m :: * -> *) blk.
PerasVoteDB m blk -> STM m (Set PerasVoteId)
PerasVoteDB.getVoteIds PerasVoteDB m blk
perasVoteDB
        pure $ \PerasVoteId
voteId -> PerasVoteId -> Set PerasVoteId -> Bool
forall a. Ord a => a -> Set a -> Bool
Set.member PerasVoteId
voteId Set PerasVoteId
voteIds
    }

-- | Create a pool writer from the 'ChainDB'.
-- This properly handles the produced certs by letting the ChainDB take care
-- of them (see 'ChainDB.addPerasVoteWithAsyncCertHandling').
makePerasVotePoolWriterFromChainDB ::
  ( IOLike m
  , BlockSupportsPeras blk
  ) =>
  SystemTime m ->
  ChainDB m blk ->
  ObjectPoolWriter PerasVoteId (PerasVote blk) m
makePerasVotePoolWriterFromChainDB :: forall (m :: * -> *) blk.
(IOLike m, BlockSupportsPeras blk) =>
SystemTime m
-> ChainDB m blk -> ObjectPoolWriter PerasVoteId (PerasVote blk) m
makePerasVotePoolWriterFromChainDB 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 :: PerasVote blk -> PerasVoteId
opwObjectId = PerasVote blk -> PerasVoteId
forall vote blk. IsPerasVote vote blk => vote -> PerasVoteId
getPerasVoteId
        , opwAddObjects :: [PerasVote blk] -> m ()
opwAddObjects = \[PerasVote blk]
votes -> do
            now <- SystemTime m -> m RelativeTime
forall (m :: * -> *). SystemTime m -> m RelativeTime
systemTimeCurrent SystemTime m
systemTime
            validatedVotes <- atomically $ do
              alreadyInDb <- ChainDB.getPerasVoteIds chainDB
              let votesNotAlreadyInDb = (PerasVote blk -> Bool) -> [PerasVote blk] -> [PerasVote blk]
forall a. (a -> Bool) -> [a] -> [a]
filter ((PerasVoteId -> Set PerasVoteId -> Bool
forall a. Ord a => a -> Set a -> Bool
`Set.notMember` Set PerasVoteId
alreadyInDb) (PerasVoteId -> Bool)
-> (PerasVote blk -> PerasVoteId) -> PerasVote blk -> Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. PerasVote blk -> PerasVoteId
forall vote blk. IsPerasVote vote blk => vote -> PerasVoteId
getPerasVoteId) [PerasVote blk]
votes
              traverse (verifyPerasVoteWithHandle resolverHandle) votesNotAlreadyInDb
            -- Some votes are invalid => reject the whole batch
            --
            -- NOTE: we could combine the two 'traverse' operations into one in
            -- which case any validated vote would be immediately added no
            -- matter what is the validity of the other votes in the batch.
            traverse_ (ChainDB.addPerasVoteWithAsyncCertHandling chainDB . WithArrivalTime now) validatedVotes
        , opwHasObject :: STM m (PerasVoteId -> Bool)
opwHasObject = do
            voteIds <- ChainDB m blk -> STM m (Set PerasVoteId)
forall (m :: * -> *) blk. ChainDB m blk -> STM m (Set PerasVoteId)
ChainDB.getPerasVoteIds ChainDB m blk
chainDB
            pure $ \PerasVoteId
voteId -> PerasVoteId -> Set PerasVoteId -> Bool
forall a. Ord a => a -> Set a -> Bool
Set.member PerasVoteId
voteId Set PerasVoteId
voteIds
        }