use std::{borrow::Cow, collections::HashMap, ops::Bound};
use reifydb_codec::{
key::encoded::{EncodedKey, EncodedKeyRange},
row::{pod::EncodedPodRow, shape::RowShape},
};
use reifydb_core::{
common::CommitVersion,
interface::catalog::{
config::{ConfigKey, GetConfig},
flow::OperatorId,
},
key::{
any::TaggedKey,
operator::{
keyspace::join::JoinRowMappingKey,
state::{GroupId, GroupStateKey, OperatorStateKey, group_inner_range_split, node_prefix},
},
},
state::timer::{GroupSweep, StateStore, TimerKind, TimerStore},
};
use reifydb_transaction::multi::RangeScope;
use reifydb_value::{
Result,
byte_size::ByteSize,
value::{
Value,
datetime::DateTime,
dictionary::{DictionaryEntryId, DictionaryId},
row_number::RowNumber,
value_type::ValueType,
},
};
use crate::{
operator::state::{iter::StateIterator, reaper::IdentityReclaim, reclaim::ReclaimOutcome},
timer::{Timer, extension::TimerExtension},
transaction::{
FlowTransaction,
dictionary::DictionaryExtension,
join_expiry::{DueStart, JoinDueEntry, JoinDuePage, JoinRowExpiryExtension},
reclaim::ReclaimExtension,
row_number::RowNumberExtension,
state::{StateExtension, StateRange},
},
};
pub trait HostContext: StateStore + TimerStore + IdentityReclaim {
fn version(&self) -> CommitVersion;
fn disarm_timer_by_key(&mut self, kind: TimerKind, key: &EncodedKey) -> Result<()>;
fn join_expiry_arm(&mut self, group: GroupId, side: u8, row_number: RowNumber, at: DateTime) -> Result<()>;
fn join_expiry_clear(&mut self, group: GroupId, side: u8, row_number: RowNumber) -> Result<Option<DateTime>>;
fn join_expiry_free(&mut self, entry: &JoinDueEntry) -> Result<()>;
fn join_expiry_min(&mut self) -> Result<Option<DateTime>>;
fn join_due_page(&mut self, at: DateTime, budget: usize, start: &DueStart) -> Result<JoinDuePage>;
fn config_uint8(&self, key: ConfigKey) -> u64;
fn state_get_many(&mut self, keys: &[GroupStateKey]) -> Result<Vec<(GroupStateKey, EncodedPodRow)>>;
fn row_shape_cache(&mut self) -> &mut HashMap<EncodedKey, RowShape>;
fn state_range(&mut self, range: EncodedKeyRange) -> Result<Vec<(GroupStateKey, EncodedPodRow)>>;
fn state_range_limited(
&mut self,
range: EncodedKeyRange,
limit: Option<usize>,
) -> Result<Vec<(GroupStateKey, EncodedPodRow)>>;
fn state_range_limited_visit(
&mut self,
range: EncodedKeyRange,
limit: Option<usize>,
visit: &mut dyn FnMut(GroupStateKey, EncodedPodRow) -> Result<()>,
) -> Result<()> {
for (key, row) in self.state_range_limited(range, limit)? {
visit(key, row)?;
}
Ok(())
}
fn state_range_iter(&mut self, range: EncodedKeyRange) -> StateIterator<'_>;
fn state_clear(&mut self) -> Result<()>;
fn reclaim_group_identity(&mut self, group: GroupId, limit: usize) -> Result<ReclaimOutcome>;
fn reclaim_group_identity_keys(&mut self, group: GroupId, keys: &[GroupStateKey]) -> Result<ReclaimOutcome>;
fn get_row_numbers(&mut self, group: GroupId, keys: &[EncodedKey]) -> Result<Vec<Option<RowNumber>>>;
fn get_row_numbers_for_groups(&mut self, groups: &[GroupId]) -> Result<Vec<Option<RowNumber>>>;
fn get_join_row_numbers(&mut self, keys: &[JoinRowMappingKey]) -> Result<Vec<Option<RowNumber>>>;
fn get_or_create_join_row_numbers(&mut self, keys: &[JoinRowMappingKey]) -> Result<Vec<(RowNumber, bool)>>;
fn remove_join_row_numbers(&mut self, keys: &[JoinRowMappingKey]) -> Result<()>;
fn remove_join_row_numbers_for_left(&mut self, tag: u8, left: u64) -> Result<()>;
fn dictionary_id_by_name(&mut self, name: &str) -> Result<Option<DictionaryId>>;
fn dictionary_value_type(&mut self, dictionary: DictionaryId) -> Option<ValueType>;
fn dictionary_id_type(&mut self, dictionary: DictionaryId) -> Option<ValueType>;
fn dictionary_find(&mut self, dictionary: DictionaryId, value: &Value) -> Result<Option<DictionaryEntryId>>;
fn dictionary_get(&mut self, dictionary: DictionaryId, id: DictionaryEntryId) -> Result<Option<Value>>;
}
pub struct TxnHostContext<'a, T: FlowTransaction> {
txn: &'a mut T,
operator: OperatorId,
now: DateTime,
}
impl<'a, T: FlowTransaction> TxnHostContext<'a, T> {
pub fn new(txn: &'a mut T, operator: OperatorId) -> Self {
let now = txn.written_at();
Self {
txn,
operator,
now,
}
}
}
impl<T: FlowTransaction> TimerStore for TxnHostContext<'_, T> {
fn arm_timer(&mut self, due: DateTime, kind: TimerKind, key: &EncodedKey) -> Result<()> {
self.txn.arm_timer(
self.operator,
&Timer {
due,
kind,
key: key.clone(),
},
)
}
fn disarm_timer(&mut self, due: DateTime, kind: TimerKind, key: &EncodedKey) -> Result<()> {
self.txn.disarm_timer(
self.operator,
&Timer {
due,
kind,
key: key.clone(),
},
)
}
fn flow_watermark(&mut self) -> Result<Option<DateTime>> {
Ok(self.txn.flow_watermark())
}
}
impl<T: FlowTransaction> StateStore for TxnHostContext<'_, T> {
fn state_get(&mut self, key: &GroupStateKey) -> Result<Option<EncodedPodRow>> {
self.txn.state_get(self.operator, key)
}
fn state_get_many_visit(
&mut self,
keys: &[GroupStateKey],
visit: &mut dyn FnMut(GroupStateKey, EncodedPodRow) -> Result<()>,
) -> Result<()> {
let batch = self.txn.state_get_many(self.operator, keys)?;
for r in batch.items {
let TaggedKey::OperatorState(decoded) = &r.key else {
continue;
};
let Some(inner) = GroupStateKey::from_framed(decoded.inner()) else {
continue;
};
visit(inner, EncodedPodRow::from(r.bytes))?;
}
Ok(())
}
fn state_classify(&mut self, key: &GroupStateKey, pre: Option<ByteSize>) {
self.txn.state_classify(self.operator, key, pre);
}
fn state_set(&mut self, key: &GroupStateKey, payload: EncodedPodRow) -> Result<()> {
self.txn.state_set(self.operator, key, payload)
}
fn state_remove(&mut self, key: &GroupStateKey) -> Result<()> {
self.txn.state_remove(self.operator, key)
}
fn state_page_inner(
&mut self,
range: EncodedKeyRange,
limit: Option<usize>,
) -> Result<Vec<(GroupStateKey, EncodedPodRow)>> {
let batch = self.txn.state_range(
self.operator,
StateRange {
range,
limit,
site: "operator::host_page",
},
)?;
let mut out = Vec::with_capacity(batch.items.len());
for r in batch.items {
if let TaggedKey::OperatorState(decoded) = &r.key
&& let Some(inner) = GroupStateKey::from_framed(decoded.inner())
{
out.push((inner, EncodedPodRow::from(r.bytes)));
}
}
Ok(out)
}
fn group_sweep_many(&mut self, groups: &[GroupId], limit: usize) -> Result<GroupSweep> {
let batch = self.txn.state_group_range(self.operator, groups, limit)?;
let mut rows = Vec::with_capacity(batch.items.len());
for r in batch.items {
if let TaggedKey::OperatorState(decoded) = &r.key
&& let Some(inner) = GroupStateKey::from_framed(decoded.inner())
{
rows.push((inner, EncodedPodRow::from(r.bytes)));
}
}
Ok(GroupSweep {
rows,
complete: !batch.has_more,
})
}
fn state_last(&mut self, range: EncodedKeyRange) -> Result<Option<(GroupStateKey, EncodedPodRow)>> {
let Some(r) = self.txn.state_last(self.operator, range)? else {
return Ok(None);
};
if let TaggedKey::OperatorState(decoded) = &r.key
&& let Some(inner) = GroupStateKey::from_framed(decoded.inner())
{
return Ok(Some((inner, EncodedPodRow::from(r.bytes))));
}
Ok(None)
}
fn get_or_create_row_numbers(&mut self, group: GroupId, keys: &[EncodedKey]) -> Result<Vec<(RowNumber, bool)>> {
self.txn.get_or_create_row_numbers(self.operator, group, keys)
}
fn get_or_create_row_numbers_for_groups(&mut self, groups: &[GroupId]) -> Result<Vec<(RowNumber, bool)>> {
self.txn.get_or_create_row_numbers_for_groups(self.operator, groups)
}
fn remove_row_number(&mut self, group: GroupId, key: &EncodedKey) -> Result<()> {
self.txn.remove_row_number(self.operator, group, key)
}
fn remove_row_number_for_group(&mut self, group: GroupId) -> Result<()> {
self.txn.remove_row_number_for_group(self.operator, group)
}
fn remove_row_numbers(&mut self, group: GroupId, keys: &[EncodedKey]) -> Result<()> {
for key in keys {
self.txn.remove_row_number(self.operator, group, key)?;
}
Ok(())
}
fn written_at(&self) -> DateTime {
self.now
}
}
impl<T: FlowTransaction> IdentityReclaim for TxnHostContext<'_, T> {
fn reclaim_identity(&mut self, group: GroupId, limit: usize) -> Result<ReclaimOutcome> {
self.txn.reclaim_group_identity(self.operator, group, limit)
}
fn reclaim_identity_keys(&mut self, group: GroupId, keys: &[GroupStateKey]) -> Result<ReclaimOutcome> {
self.txn.reclaim_group_identity_keys(self.operator, group, keys)
}
}
impl<T: FlowTransaction> HostContext for TxnHostContext<'_, T> {
fn version(&self) -> CommitVersion {
self.txn.version()
}
fn row_shape_cache(&mut self) -> &mut HashMap<EncodedKey, RowShape> {
self.txn.row_shape_cache(self.operator)
}
fn disarm_timer_by_key(&mut self, kind: TimerKind, key: &EncodedKey) -> Result<()> {
self.txn.disarm_timer_by_key(self.operator, kind, key)
}
fn join_expiry_arm(&mut self, group: GroupId, side: u8, row_number: RowNumber, at: DateTime) -> Result<()> {
self.txn.join_expiry_arm(self.operator, group, side, row_number, at)
}
fn join_expiry_clear(&mut self, group: GroupId, side: u8, row_number: RowNumber) -> Result<Option<DateTime>> {
self.txn.join_expiry_clear(self.operator, group, side, row_number)
}
fn join_expiry_free(&mut self, entry: &JoinDueEntry) -> Result<()> {
self.txn.join_expiry_free(self.operator, entry)
}
fn join_expiry_min(&mut self) -> Result<Option<DateTime>> {
self.txn.join_expiry_min(self.operator)
}
fn join_due_page(&mut self, at: DateTime, budget: usize, start: &DueStart) -> Result<JoinDuePage> {
self.txn.join_due_page(self.operator, at, budget, start)
}
fn config_uint8(&self, key: ConfigKey) -> u64 {
self.txn.catalog().get_config_uint8(key)
}
fn state_get_many(&mut self, keys: &[GroupStateKey]) -> Result<Vec<(GroupStateKey, EncodedPodRow)>> {
let batch = self.txn.state_get_many(self.operator, keys)?;
let mut out = Vec::with_capacity(batch.items.len());
for r in batch.items {
let Some(key) = unscope(&r.key) else {
continue;
};
out.push((key, EncodedPodRow::from(r.bytes)));
}
Ok(out)
}
fn state_range(&mut self, range: EncodedKeyRange) -> Result<Vec<(GroupStateKey, EncodedPodRow)>> {
self.state_range_limited(range, None)
}
fn state_range_limited(
&mut self,
range: EncodedKeyRange,
limit: Option<usize>,
) -> Result<Vec<(GroupStateKey, EncodedPodRow)>> {
let mut out = Vec::new();
self.state_range_limited_visit(range, limit, &mut |key, row| {
out.push((key, row));
Ok(())
})?;
Ok(out)
}
fn state_range_limited_visit(
&mut self,
range: EncodedKeyRange,
limit: Option<usize>,
visit: &mut dyn FnMut(GroupStateKey, EncodedPodRow) -> Result<()>,
) -> Result<()> {
let site = range_site(&range);
let mut query = StateRange::forward(range, site);
query.limit = limit;
let batch = self.txn.state_range(self.operator, query)?;
for r in batch.items {
let Some(key) = unscope(&r.key) else {
continue;
};
visit(key, EncodedPodRow::from(r.bytes))?;
}
Ok(())
}
fn state_range_iter(&mut self, range: EncodedKeyRange) -> StateIterator<'_> {
let prefixed = range.with_prefix(EncodedKey::new(node_prefix(self.operator)));
StateIterator::new(self.txn.range(prefixed, RangeScope::All, 1024))
}
fn state_clear(&mut self) -> Result<()> {
self.txn.state_clear(self.operator)
}
fn reclaim_group_identity(&mut self, group: GroupId, limit: usize) -> Result<ReclaimOutcome> {
self.txn.reclaim_group_identity(self.operator, group, limit)
}
fn reclaim_group_identity_keys(&mut self, group: GroupId, keys: &[GroupStateKey]) -> Result<ReclaimOutcome> {
self.txn.reclaim_group_identity_keys(self.operator, group, keys)
}
fn get_row_numbers(&mut self, group: GroupId, keys: &[EncodedKey]) -> Result<Vec<Option<RowNumber>>> {
self.txn.get_row_numbers(self.operator, group, keys)
}
fn get_row_numbers_for_groups(&mut self, groups: &[GroupId]) -> Result<Vec<Option<RowNumber>>> {
self.txn.get_row_numbers_for_groups(self.operator, groups)
}
fn get_join_row_numbers(&mut self, keys: &[JoinRowMappingKey]) -> Result<Vec<Option<RowNumber>>> {
self.txn.get_join_row_numbers(self.operator, keys)
}
fn get_or_create_join_row_numbers(&mut self, keys: &[JoinRowMappingKey]) -> Result<Vec<(RowNumber, bool)>> {
self.txn.get_or_create_join_row_numbers(self.operator, keys)
}
fn remove_join_row_numbers(&mut self, keys: &[JoinRowMappingKey]) -> Result<()> {
self.txn.remove_join_row_numbers(self.operator, keys)
}
fn remove_join_row_numbers_for_left(&mut self, tag: u8, left: u64) -> Result<()> {
self.txn.remove_join_row_numbers_for_left(self.operator, tag, left)
}
fn dictionary_id_by_name(&mut self, name: &str) -> Result<Option<DictionaryId>> {
Ok(self.txn.find_dictionary_by_name(name).map(|d| d.id))
}
fn dictionary_value_type(&mut self, dictionary: DictionaryId) -> Option<ValueType> {
self.txn.find_dictionary(dictionary).map(|d| d.value_type)
}
fn dictionary_id_type(&mut self, dictionary: DictionaryId) -> Option<ValueType> {
self.txn.find_dictionary(dictionary).map(|d| d.id_type)
}
fn dictionary_find(&mut self, dictionary: DictionaryId, value: &Value) -> Result<Option<DictionaryEntryId>> {
match self.txn.find_dictionary(dictionary) {
Some(dict) => self.txn.find_in_dictionary(&dict, value),
None => Ok(None),
}
}
fn dictionary_get(&mut self, dictionary: DictionaryId, id: DictionaryEntryId) -> Result<Option<Value>> {
match self.txn.find_dictionary(dictionary) {
Some(dict) => self.txn.get_from_dictionary(&dict, id),
None => Ok(None),
}
}
}
fn unscope(key: &TaggedKey) -> Option<GroupStateKey> {
let TaggedKey::OperatorState(key) = key else {
return None;
};
GroupStateKey::from_framed(key.inner())
}
fn range_site(range: &EncodedKeyRange) -> &'static str {
if group_inner_range_split(range).is_some() {
return "operator::host_sweep";
}
let key = match &range.start {
Bound::Included(key) | Bound::Excluded(key) => key,
Bound::Unbounded => return "operator::host_range",
};
match OperatorStateKey::decode_inner(key.as_slice()).map(|(_, keyspace, _)| keyspace.name()) {
Some(Cow::Borrowed(name)) => name,
_ => "operator::host_range",
}
}