use alloc::{
collections::{BTreeMap, btree_map},
string::String,
sync::Arc,
vec::Vec,
};
use core::{cmp::Ordering, iter::Peekable, marker::PhantomData, mem, ops::Bound};
use buggy::{Bug, bug};
use yoke::{Yoke, Yokeable};
use crate::{
Address, Bytes, Checkpoint, ClientError, ClientState, CmdId, Command, Fact, FactPerspective,
GraphId, Keys, MaxCut, NullSink, Perspective, Policy, PolicyId, PolicyStore, Prior, Priority,
Query, QueryMut, Revertable, Segment as _, Sink, Storage, StorageError, StorageProvider,
policy::{ActionPlacement, CommandPlacement},
};
pub struct Session<SP: StorageProvider, PS> {
graph_id: GraphId,
policy_id: PolicyId,
base_facts: <SP::Storage as Storage>::FactIndex,
fact_log: Vec<(String, Keys, Option<Bytes>)>,
current_facts: Arc<BTreeMap<String, BTreeMap<Keys, Option<Bytes>>>>,
_policy_store: PhantomData<PS>,
}
struct SessionPerspective<'a, SP: StorageProvider, PS, MS> {
session: &'a mut Session<SP, PS>,
message_sink: &'a mut MS,
}
impl<SP: StorageProvider, PS> Session<SP, PS> {
pub(super) fn new(provider: &mut SP, graph_id: GraphId) -> Result<Self, ClientError> {
let storage = provider.get_storage(graph_id)?;
let head_loc = storage.get_head()?;
let seg = storage.get_segment(head_loc)?;
let base_facts = seg.facts()?;
let result = Self {
graph_id,
policy_id: seg.policy(),
base_facts,
fact_log: Vec::new(),
current_facts: Arc::default(),
_policy_store: PhantomData,
};
Ok(result)
}
}
impl<SP: StorageProvider, PS: PolicyStore> Session<SP, PS> {
pub fn action<ES, MS>(
&mut self,
client: &ClientState<PS, SP>,
effect_sink: &mut ES,
message_sink: &mut MS,
action: <PS::Policy as Policy>::Action<'_>,
) -> Result<(), ClientError>
where
ES: Sink<PS::Effect>,
MS: for<'b> Sink<&'b [u8]>,
{
let policy = client.policy_store.get_policy(self.policy_id)?;
let mut perspective = SessionPerspective {
session: self,
message_sink,
};
let checkpoint = perspective.checkpoint();
effect_sink.begin();
match policy.call_action(
action,
&mut perspective,
effect_sink,
ActionPlacement::OffGraph,
) {
Ok(()) => {
effect_sink.commit();
Ok(())
}
Err(e) => {
perspective.revert(checkpoint)?;
perspective.message_sink.rollback();
effect_sink.rollback();
Err(e.into())
}
}
}
pub fn receive(
&mut self,
client: &ClientState<PS, SP>,
sink: &mut impl Sink<PS::Effect>,
command_bytes: &[u8],
) -> Result<(), ClientError> {
let command = SessionCommand::deserialize(self.graph_id, command_bytes)
.ok_or(ClientError::SessionDeserialize)?;
let policy = client.policy_store.get_policy(self.policy_id)?;
let mut perspective = SessionPerspective {
session: self,
message_sink: &mut NullSink,
};
sink.begin();
let checkpoint = perspective.checkpoint();
if let Err(e) =
policy.call_rule(&command, &mut perspective, sink, CommandPlacement::OffGraph)
{
perspective.revert(checkpoint)?;
sink.rollback();
return Err(e.into());
}
sink.commit();
Ok(())
}
}
fn session_parent(graph_id: GraphId) -> Prior<Address> {
Prior::Single(Address {
id: CmdId::transmute(graph_id),
max_cut: MaxCut::new(0),
})
}
struct SessionCommand<'a> {
graph_id: GraphId,
id: CmdId,
data: &'a [u8],
}
impl Command for SessionCommand<'_> {
fn priority(&self) -> Priority {
Priority::Basic(0)
}
fn id(&self) -> CmdId {
self.id
}
fn parent(&self) -> Prior<Address> {
session_parent(self.graph_id)
}
fn policy(&self) -> Option<&[u8]> {
None
}
fn bytes(&self) -> &[u8] {
self.data
}
}
impl<'sc> SessionCommand<'sc> {
fn from_cmd(graph_id: GraphId, command: &'sc impl Command) -> Result<Self, Bug> {
if command.policy().is_some() {
bug!("session command should have no policy");
}
if !matches!(command.priority(), Priority::Basic(_)) {
bug!("session command has bad priority");
}
if command.parent() != session_parent(graph_id) {
bug!("session command has bad parent");
}
Ok(SessionCommand {
graph_id,
id: command.id(),
data: command.bytes(),
})
}
fn serialize(&self) -> Vec<u8> {
[self.id.as_bytes(), self.data].concat()
}
fn deserialize(graph_id: GraphId, bytes: &'sc [u8]) -> Option<Self> {
let (id, data) = bytes.split_first_chunk()?;
Some(Self {
graph_id,
id: CmdId::from_bytes(*id),
data,
})
}
}
struct QueryIterator<I1: Iterator, I2: Iterator> {
prior: Peekable<I1>,
current: Peekable<I2>,
}
impl<I1, I2> QueryIterator<I1, I2>
where
I1: Iterator<Item = Result<Fact, StorageError>>,
I2: Iterator<Item = (Keys, Option<Bytes>)>,
{
fn new(prior: I1, current: I2) -> Self {
Self {
prior: prior.peekable(),
current: current.peekable(),
}
}
}
impl<I1, I2> Iterator for QueryIterator<I1, I2>
where
I1: Iterator<Item = Result<Fact, StorageError>>,
I2: Iterator<Item = (Keys, Option<Bytes>)>,
{
type Item = Result<Fact, StorageError>;
fn next(&mut self) -> Option<Self::Item> {
loop {
let Some(new) = self.current.peek() else {
return self.prior.next();
};
if let Some(old) = self.prior.peek() {
let Ok(old) = old else {
return self.prior.next();
};
match new.0.cmp(&old.key) {
Ordering::Equal => {
let _ = self.prior.next();
}
Ordering::Greater => {
return self.prior.next();
}
Ordering::Less => {
}
}
}
let Some(slot) = self.current.next() else {
bug!("expected Some after peek")
};
if let (k, Some(v)) = slot {
return Some(Ok(Fact {
key: k.iter().cloned().collect(),
value: v,
}));
}
}
}
}
impl<SP, PS, MS> FactPerspective for SessionPerspective<'_, SP, PS, MS> where SP: StorageProvider {}
impl<SP, PS, MS> Query for SessionPerspective<'_, SP, PS, MS>
where
SP: StorageProvider,
{
fn query(&self, name: &str, keys: &[Bytes]) -> Result<Option<Bytes>, StorageError> {
if let Some(slot) = self
.session
.current_facts
.get(name)
.and_then(|m| m.get(keys))
{
return Ok(slot.clone());
}
self.session.base_facts.query(name, keys)
}
type QueryIterator = QueryIterator<
<<SP::Storage as Storage>::FactIndex as Query>::QueryIterator,
YokeIter<PrefixIter<'static>, Arc<BTreeMap<String, BTreeMap<Keys, Option<Bytes>>>>>,
>;
fn query_prefix(
&self,
name: &str,
prefix: &[Bytes],
) -> Result<Self::QueryIterator, StorageError> {
let prior = self.session.base_facts.query_prefix(name, prefix)?;
let current = Yoke::<PrefixIter<'static>, _>::attach_to_cart(
Arc::clone(&self.session.current_facts),
|map| match map.get(name) {
Some(facts) => PrefixIter::new(facts, prefix.iter().cloned().collect()),
None => PrefixIter::default(),
},
);
Ok(QueryIterator::new(prior, YokeIter::new(current)))
}
}
#[derive(Default, Yokeable)]
struct PrefixIter<'map> {
range: btree_map::Range<'map, Keys, Option<Bytes>>,
prefix: Keys,
}
impl<'map> PrefixIter<'map> {
fn new(map: &'map BTreeMap<Keys, Option<Bytes>>, prefix: Keys) -> Self {
let range = map.range::<[Bytes], _>((Bound::Included(prefix.as_ref()), Bound::Unbounded));
Self { range, prefix }
}
}
impl Iterator for PrefixIter<'_> {
type Item = (Keys, Option<Bytes>);
fn next(&mut self) -> Option<Self::Item> {
self.range
.next()
.filter(|(k, _)| k.starts_with(&self.prefix))
.map(|(k, v)| (k.clone(), v.clone()))
}
}
struct YokeIter<I: for<'a> Yokeable<'a>, C>(Option<Yoke<I, C>>);
impl<I: for<'a> Yokeable<'a>, C> YokeIter<I, C> {
fn new(yoke: Yoke<I, C>) -> Self {
Self(Some(yoke))
}
}
impl<I, C> Iterator for YokeIter<I, C>
where
I: Iterator + for<'a> Yokeable<'a>,
for<'a> <I as Yokeable<'a>>::Output: Iterator<Item = I::Item>,
{
type Item = I::Item;
fn next(&mut self) -> Option<Self::Item> {
let mut item = None;
self.0 = Some(self.0.take()?.map_project::<I, _>(|mut it, _| {
item = it.next();
it
}));
item
}
}
impl<SP: StorageProvider, PS, MS> QueryMut for SessionPerspective<'_, SP, PS, MS> {
fn insert(&mut self, name: String, keys: Keys, value: Bytes) -> Result<(), StorageError> {
self.session
.fact_log
.push((name.clone(), keys.clone(), Some(value.clone())));
Arc::make_mut(&mut self.session.current_facts)
.entry(name)
.or_default()
.insert(keys, Some(value));
Ok(())
}
fn delete(&mut self, name: String, keys: Keys) -> Result<(), StorageError> {
self.session
.fact_log
.push((name.clone(), keys.clone(), None));
Arc::make_mut(&mut self.session.current_facts)
.entry(name)
.or_default()
.insert(keys, None);
Ok(())
}
}
impl<SP, PS, MS> Perspective for SessionPerspective<'_, SP, PS, MS>
where
SP: StorageProvider,
MS: for<'b> Sink<&'b [u8]>,
{
fn policy(&self) -> PolicyId {
self.session.policy_id
}
fn add_command(&mut self, command: &impl Command) -> Result<usize, StorageError> {
let command = SessionCommand::from_cmd(self.session.graph_id, command)?;
self.message_sink.consume(&command.serialize());
Ok(0)
}
fn includes(&self, _id: CmdId) -> bool {
debug_assert!(false, "only used in transactions");
false
}
fn head_address(&self) -> Result<Prior<Address>, Bug> {
Ok(session_parent(self.session.graph_id))
}
}
impl<SP, PS, MS> Revertable for SessionPerspective<'_, SP, PS, MS>
where
SP: StorageProvider,
{
fn checkpoint(&self) -> Checkpoint {
Checkpoint {
index: self.session.fact_log.len(),
}
}
fn revert(&mut self, checkpoint: Checkpoint) -> Result<(), StorageError> {
if checkpoint.index == self.session.fact_log.len() {
return Ok(());
}
if checkpoint.index > self.session.fact_log.len() {
bug!(
"A checkpoint's index should always be less than or equal to the length of a session's fact log!"
);
}
self.session.fact_log.truncate(checkpoint.index);
let mut facts =
Arc::get_mut(&mut self.session.current_facts).map_or_else(BTreeMap::new, mem::take);
facts.clear();
for (n, k, v) in self.session.fact_log.iter().cloned() {
facts.entry(n).or_default().insert(k, v);
}
self.session.current_facts = Arc::new(facts);
Ok(())
}
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn test_query_iterator() {
#![allow(clippy::type_complexity)]
let prior: Vec<Result<(&[&[u8]], &[u8]), _>> = vec![
Ok((&[b"a"], b"a0")),
Ok((&[b"c"], b"c0")),
Ok((&[b"d"], b"d0")),
Ok((&[b"f"], b"f0")),
Err(StorageError::IoError),
];
let current: Vec<([Bytes; 1], Option<&[u8]>)> = vec![
([Bytes::from(*b"a")], None),
([Bytes::from(*b"b")], Some(b"b1")),
([Bytes::from(*b"e")], None),
([Bytes::from(*b"j")], None),
];
let merged: Vec<Result<(&[&[u8]], &[u8]), _>> = vec![
Ok((&[b"b"], b"b1")),
Ok((&[b"c"], b"c0")),
Ok((&[b"d"], b"d0")),
Ok((&[b"f"], b"f0")),
Err(StorageError::IoError),
];
let got: Vec<_> = QueryIterator::new(
prior.into_iter().map(|r| {
r.map(|(k, v)| Fact {
key: k.into(),
value: v.into(),
})
}),
current
.into_iter()
.map(|(k, v)| (k.into_iter().collect(), v.map(Bytes::from))),
)
.collect();
let want: Vec<_> = merged
.into_iter()
.map(|r| {
r.map(|(k, v)| Fact {
key: k.into(),
value: v.into(),
})
})
.collect();
assert_eq!(got, want);
}
}