use alloc::collections::{BTreeMap, BTreeSet};
use alloc::vec::Vec;
use crate::wire::{
ConversationId, DeliverySeq, InvalidObserverEpoch, InvalidObserverEpochList, ObserverEpoch,
ObserverProgressStatus, ObserverRecoveryAccepted, ObserverRecoveryHandshake,
ObserverRecoveryResponse,
};
#[derive(Debug, PartialEq, Eq)]
pub struct ObserverProgressProjection {
conversation_id: ConversationId,
new_observer_progress: DeliverySeq,
}
impl ObserverProgressProjection {
pub(in crate::lifecycle) const fn new(
conversation_id: ConversationId,
new_observer_progress: DeliverySeq,
) -> Self {
Self {
conversation_id,
new_observer_progress,
}
}
#[must_use]
pub const fn conversation_id(&self) -> ConversationId {
self.conversation_id
}
#[must_use]
pub const fn new_observer_progress(&self) -> DeliverySeq {
self.new_observer_progress
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ObserverRecoveryAggregateRestoreError {
DuplicateProgress {
conversation_id: ConversationId,
},
DuplicateArm {
conversation_id: ConversationId,
},
ArmWithoutProgress {
conversation_id: ConversationId,
},
ArmEpochMismatch {
conversation_id: ConversationId,
armed_epoch: ObserverEpoch,
current_observer_progress: DeliverySeq,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ObserverProgressTrackError {
AlreadyTracked {
conversation_id: ConversationId,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ObserverProgressAdvanceError {
ConversationUnknown {
conversation_id: ConversationId,
},
NotAdvancing {
conversation_id: ConversationId,
current_observer_progress: DeliverySeq,
presented_progress: DeliverySeq,
},
}
#[derive(Debug, Default, PartialEq, Eq)]
pub struct ObserverRecoveryAggregate {
progress: BTreeMap<ConversationId, DeliverySeq>,
armed: BTreeMap<ConversationId, ObserverEpoch>,
}
impl ObserverRecoveryAggregate {
#[must_use]
pub const fn new() -> Self {
Self {
progress: BTreeMap::new(),
armed: BTreeMap::new(),
}
}
pub fn restore(
progress_rows: &[(ConversationId, DeliverySeq)],
armed_rows: &[(ConversationId, ObserverEpoch)],
) -> Result<Self, ObserverRecoveryAggregateRestoreError> {
let mut progress = BTreeMap::new();
for (conversation_id, observer_progress) in progress_rows {
if progress
.insert(*conversation_id, *observer_progress)
.is_some()
{
return Err(ObserverRecoveryAggregateRestoreError::DuplicateProgress {
conversation_id: *conversation_id,
});
}
}
let mut armed = BTreeMap::new();
for (conversation_id, armed_epoch) in armed_rows {
let Some(current_observer_progress) = progress.get(conversation_id).copied() else {
return Err(ObserverRecoveryAggregateRestoreError::ArmWithoutProgress {
conversation_id: *conversation_id,
});
};
if current_observer_progress != *armed_epoch {
return Err(ObserverRecoveryAggregateRestoreError::ArmEpochMismatch {
conversation_id: *conversation_id,
armed_epoch: *armed_epoch,
current_observer_progress,
});
}
if armed.insert(*conversation_id, *armed_epoch).is_some() {
return Err(ObserverRecoveryAggregateRestoreError::DuplicateArm {
conversation_id: *conversation_id,
});
}
}
Ok(Self { progress, armed })
}
#[must_use]
pub fn observer_progress(&self, conversation_id: ConversationId) -> Option<DeliverySeq> {
self.progress.get(&conversation_id).copied()
}
#[must_use]
pub fn armed_epoch(&self, conversation_id: ConversationId) -> Option<ObserverEpoch> {
self.armed.get(&conversation_id).copied()
}
#[must_use]
pub fn progress_rows(&self) -> Vec<(ConversationId, DeliverySeq)> {
self.progress
.iter()
.map(|(conversation_id, observer_progress)| (*conversation_id, *observer_progress))
.collect()
}
#[must_use]
pub fn armed_rows(&self) -> Vec<(ConversationId, ObserverEpoch)> {
self.armed
.iter()
.map(|(conversation_id, armed_epoch)| (*conversation_id, *armed_epoch))
.collect()
}
#[must_use]
pub fn decide_track(
self,
conversation_id: ConversationId,
observer_progress: DeliverySeq,
) -> ObserverProgressTrackDecision {
if self.progress.contains_key(&conversation_id) {
return ObserverProgressTrackDecision::Refuse {
aggregate: self,
error: ObserverProgressTrackError::AlreadyTracked { conversation_id },
};
}
ObserverProgressTrackDecision::Commit(ObserverProgressTrackTransaction {
aggregate: self,
conversation_id,
observer_progress,
})
}
#[must_use]
pub fn decide_progress_advance(
self,
conversation_id: ConversationId,
presented_progress: DeliverySeq,
) -> ObserverProgressAdvanceDecision {
let Some(current) = self.progress.get(&conversation_id).copied() else {
return ObserverProgressAdvanceDecision::Refuse {
aggregate: self,
error: ObserverProgressAdvanceError::ConversationUnknown { conversation_id },
};
};
if presented_progress <= current {
return ObserverProgressAdvanceDecision::Refuse {
aggregate: self,
error: ObserverProgressAdvanceError::NotAdvancing {
conversation_id,
current_observer_progress: current,
presented_progress,
},
};
}
let fired = self
.armed
.get(&conversation_id)
.copied()
.map(|refused_epoch| ObserverRecoveryArm {
conversation_id,
refused_epoch,
});
ObserverProgressAdvanceDecision::Commit(ObserverProgressAdvanceTransaction {
aggregate: self,
conversation_id,
presented_progress,
fired,
})
}
#[must_use]
pub fn decide_recovery(
self,
request: &ObserverRecoveryHandshake,
max_entries: u64,
connection_conversation_limit: u64,
tracked_conversations: &[ConversationId],
) -> ObserverRecoveryTransactionDecision {
let decision = apply_observer_recovery(
request,
max_entries,
connection_conversation_limit,
tracked_conversations,
|conversation_id| self.progress.get(&conversation_id).copied(),
);
match decision {
ObserverRecoveryDecision::Respond(response) => {
ObserverRecoveryTransactionDecision::Respond {
aggregate: self,
response,
}
}
ObserverRecoveryDecision::Commit(commit) => {
ObserverRecoveryTransactionDecision::Commit(ObserverRecoveryTransaction {
aggregate: self,
commit,
})
}
}
}
}
#[derive(Debug, PartialEq, Eq)]
pub enum ObserverRecoveryTransactionDecision {
Respond {
aggregate: ObserverRecoveryAggregate,
response: ObserverRecoveryResponse,
},
Commit(ObserverRecoveryTransaction),
}
#[derive(Debug, PartialEq, Eq)]
pub struct ObserverRecoveryTransaction {
aggregate: ObserverRecoveryAggregate,
commit: ObserverRecoveryCommit,
}
impl ObserverRecoveryTransaction {
#[must_use]
pub fn arms(&self) -> &[ObserverRecoveryArm] {
self.commit.arms()
}
#[must_use]
pub const fn outcome(&self) -> &ObserverRecoveryAccepted {
self.commit.outcome()
}
#[must_use]
pub fn commit(mut self) -> (ObserverRecoveryAggregate, ObserverRecoveryAccepted) {
let (arms, outcome) = self.commit.into_parts();
for arm in arms {
self.aggregate
.armed
.insert(arm.conversation_id(), arm.refused_epoch());
}
(self.aggregate, outcome)
}
#[must_use]
pub fn abort(self) -> ObserverRecoveryAggregate {
self.aggregate
}
}
#[derive(Debug, PartialEq, Eq)]
pub enum ObserverProgressAdvanceDecision {
Refuse {
aggregate: ObserverRecoveryAggregate,
error: ObserverProgressAdvanceError,
},
Commit(ObserverProgressAdvanceTransaction),
}
#[derive(Debug, PartialEq, Eq)]
pub struct ObserverProgressAdvanceTransaction {
aggregate: ObserverRecoveryAggregate,
conversation_id: ConversationId,
presented_progress: DeliverySeq,
fired: Option<ObserverRecoveryArm>,
}
impl ObserverProgressAdvanceTransaction {
#[must_use]
pub const fn conversation_id(&self) -> ConversationId {
self.conversation_id
}
#[must_use]
pub const fn presented_progress(&self) -> DeliverySeq {
self.presented_progress
}
#[must_use]
pub const fn fired_arm(&self) -> Option<ObserverRecoveryArm> {
self.fired
}
#[must_use]
pub fn commit(mut self) -> (ObserverRecoveryAggregate, Option<ObserverRecoveryArm>) {
self.aggregate
.progress
.insert(self.conversation_id, self.presented_progress);
if self.fired.is_some() {
self.aggregate.armed.remove(&self.conversation_id);
}
(self.aggregate, self.fired)
}
#[must_use]
pub fn abort(self) -> ObserverRecoveryAggregate {
self.aggregate
}
}
#[derive(Debug, PartialEq, Eq)]
pub enum ObserverProgressTrackDecision {
Refuse {
aggregate: ObserverRecoveryAggregate,
error: ObserverProgressTrackError,
},
Commit(ObserverProgressTrackTransaction),
}
#[derive(Debug, PartialEq, Eq)]
pub struct ObserverProgressTrackTransaction {
aggregate: ObserverRecoveryAggregate,
conversation_id: ConversationId,
observer_progress: DeliverySeq,
}
impl ObserverProgressTrackTransaction {
#[must_use]
pub const fn conversation_id(&self) -> ConversationId {
self.conversation_id
}
#[must_use]
pub const fn observer_progress(&self) -> DeliverySeq {
self.observer_progress
}
#[must_use]
pub fn commit(mut self) -> ObserverRecoveryAggregate {
self.aggregate
.progress
.insert(self.conversation_id, self.observer_progress);
self.aggregate
}
#[must_use]
pub fn abort(self) -> ObserverRecoveryAggregate {
self.aggregate
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ObserverRecoveryArm {
conversation_id: ConversationId,
refused_epoch: ObserverEpoch,
}
impl ObserverRecoveryArm {
#[must_use]
pub const fn conversation_id(self) -> ConversationId {
self.conversation_id
}
#[must_use]
pub const fn refused_epoch(self) -> ObserverEpoch {
self.refused_epoch
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ObserverRecoveryCommit {
arms: Vec<ObserverRecoveryArm>,
outcome: ObserverRecoveryAccepted,
}
impl ObserverRecoveryCommit {
#[must_use]
pub fn arms(&self) -> &[ObserverRecoveryArm] {
&self.arms
}
#[must_use]
pub const fn outcome(&self) -> &ObserverRecoveryAccepted {
&self.outcome
}
#[must_use]
pub fn into_parts(self) -> (Vec<ObserverRecoveryArm>, ObserverRecoveryAccepted) {
(self.arms, self.outcome)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ObserverRecoveryDecision {
Respond(ObserverRecoveryResponse),
Commit(ObserverRecoveryCommit),
}
fn wire_count(value: usize) -> u64 {
u64::try_from(value).map_or(u64::MAX, core::convert::identity)
}
#[must_use]
pub(super) fn apply_observer_recovery<F>(
request: &ObserverRecoveryHandshake,
max_entries: u64,
connection_conversation_limit: u64,
tracked_conversations: &[ConversationId],
mut observer_progress: F,
) -> ObserverRecoveryDecision
where
F: FnMut(ConversationId) -> Option<DeliverySeq>,
{
let presented_entries = wire_count(request.observer_refusals.len());
if presented_entries > max_entries {
return ObserverRecoveryDecision::Respond(
ObserverRecoveryResponse::invalid_observer_epoch_list(
InvalidObserverEpochList::TooManyEntries {
presented_entries,
max_entries,
},
),
);
}
let mut first_indices = BTreeMap::new();
for (index, refusal) in request.observer_refusals.iter().enumerate() {
let request_index = wire_count(index);
if let Some(first_index) = first_indices.insert(refusal.conversation_id, request_index) {
return ObserverRecoveryDecision::Respond(
ObserverRecoveryResponse::invalid_observer_epoch_list(
InvalidObserverEpochList::DuplicateConversation {
conversation_id: refusal.conversation_id,
first_index,
duplicate_index: request_index,
},
),
);
}
}
let mut tracked: BTreeSet<_> = tracked_conversations.iter().copied().collect();
for refusal in &request.observer_refusals {
if tracked.contains(&refusal.conversation_id) {
continue;
}
let occupied = wire_count(tracked.len());
if occupied >= connection_conversation_limit {
return ObserverRecoveryDecision::Respond(
ObserverRecoveryResponse::connection_capacity_exceeded(
refusal.conversation_id,
connection_conversation_limit,
),
);
}
tracked.insert(refusal.conversation_id);
}
let mut validated = Vec::with_capacity(request.observer_refusals.len());
for refusal in &request.observer_refusals {
let Some(current_observer_progress) = observer_progress(refusal.conversation_id) else {
return ObserverRecoveryDecision::Respond(
ObserverRecoveryResponse::invalid_observer_epoch(
InvalidObserverEpoch::ConversationUnknown {
conversation_id: refusal.conversation_id,
presented_epoch: refusal.refused_epoch,
},
),
);
};
if refusal.refused_epoch > current_observer_progress {
return ObserverRecoveryDecision::Respond(
ObserverRecoveryResponse::invalid_observer_epoch(
InvalidObserverEpoch::EpochAhead {
conversation_id: refusal.conversation_id,
presented_epoch: refusal.refused_epoch,
current_observer_progress,
},
),
);
}
validated.push((refusal, current_observer_progress));
}
let mut arms = Vec::new();
let mut statuses = Vec::with_capacity(validated.len());
for (refusal, current_observer_progress) in validated {
let armed = refusal.refused_epoch == current_observer_progress;
if armed {
arms.push(ObserverRecoveryArm {
conversation_id: refusal.conversation_id,
refused_epoch: refusal.refused_epoch,
});
}
statuses.push(ObserverProgressStatus {
conversation_id: refusal.conversation_id,
refused_epoch: refusal.refused_epoch,
current_observer_progress,
armed,
progressed: !armed,
});
}
ObserverRecoveryDecision::Commit(ObserverRecoveryCommit {
arms,
outcome: ObserverRecoveryAccepted { statuses },
})
}