use std::collections::BTreeMap;
use std::collections::HashMap;
use std::collections::HashSet;
use std::collections::VecDeque;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::sync::Arc;
use bytes::Bytes;
use uuid::Uuid;
use super::Chunk;
use super::ReassemblyLimits;
use crate::consts::MAX_TTL_MS;
use crate::consts::TS_OFFSET_TOLERANCE_MS;
use crate::fair_admission::try_reserve_atomic;
use crate::utils::get_epoch_ms;
pub(super) struct Pending {
total: usize,
pub(super) slots: BTreeMap<usize, Bytes>,
pub(super) data_bytes: usize,
ts_ms: u128,
ttl_ms: u64,
failure_charged: bool,
local_capacity_rejected: bool,
peer_attributable: bool,
}
impl Pending {
fn new(total: usize, ts_ms: u128, ttl_ms: u64, peer_attributable: bool) -> Self {
Self {
total,
slots: BTreeMap::new(),
data_bytes: 0,
ts_ms,
ttl_ms,
failure_charged: false,
local_capacity_rejected: false,
peer_attributable,
}
}
fn is_complete(&self) -> bool {
self.slots.len() == self.total
}
pub(super) fn cost(&self, slot_overhead: usize) -> usize {
self.slots
.len()
.saturating_mul(slot_overhead)
.saturating_add(self.data_bytes)
}
fn assemble(self) -> Bytes {
self.slots.into_values().flatten().collect()
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
struct LogicalTransmission {
id: Uuid,
ts_ms: u128,
ttl_ms: u64,
}
impl LogicalTransmission {
const fn new(id: Uuid, ts_ms: u128, ttl_ms: u64) -> Self {
Self { id, ts_ms, ttl_ms }
}
}
pub struct MessageReassembler {
pub(super) pending: HashMap<Uuid, Pending>,
pub(super) buffered_cost: usize,
completed: VecDeque<(Uuid, u128)>,
pub(super) completed_ids: HashSet<Uuid>,
failed: VecDeque<(LogicalTransmission, u128)>,
failed_ids: HashSet<LogicalTransmission>,
failure_tracking_saturated_until: u128,
capacity_rejected: VecDeque<(Uuid, u128)>,
capacity_rejected_ids: HashSet<Uuid>,
capacity_tracking_saturated_until: u128,
limits: ReassemblyLimits,
budget: Arc<ReassemblyBudget>,
}
pub(crate) struct ReassemblyBudget {
pub(super) buffered_cost: AtomicUsize,
limit: usize,
}
impl ReassemblyBudget {
pub(crate) fn new(limits: ReassemblyLimits) -> Self {
Self {
buffered_cost: AtomicUsize::new(0),
limit: limits.normalized().max_total_buffered_cost,
}
}
fn try_reserve(&self, cost: usize) -> bool {
try_reserve_atomic(&self.buffered_cost, cost, self.limit)
}
fn release(&self, cost: usize) {
if self
.buffered_cost
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |current| {
current.checked_sub(cost)
})
.is_err()
{
tracing::error!(cost, "reassembly budget release exceeded retained cost");
}
}
#[cfg(all(test, feature = "dummy", not(target_family = "wasm")))]
pub(crate) fn buffered_cost_for_test(&self) -> usize {
self.buffered_cost.load(Ordering::Acquire)
}
}
pub(crate) struct RetainedReassembly {
bytes: Bytes,
budget: Arc<ReassemblyBudget>,
cost: usize,
}
pub(crate) enum ReassemblyOutcome {
Incomplete,
Complete(RetainedReassembly),
Rejected(ReassemblyRejection),
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum ReassemblyRejection {
Invalid,
Capacity,
Replay,
}
impl RetainedReassembly {
fn into_bytes(mut self) -> Bytes {
self.budget.release(self.cost);
self.cost = 0;
std::mem::take(&mut self.bytes)
}
}
impl AsRef<[u8]> for RetainedReassembly {
fn as_ref(&self) -> &[u8] {
&self.bytes
}
}
impl Drop for RetainedReassembly {
fn drop(&mut self) {
self.budget.release(self.cost);
}
}
impl Default for MessageReassembler {
fn default() -> Self {
Self::with_limits(ReassemblyLimits::production())
}
}
impl MessageReassembler {
pub fn new() -> Self {
Self::default()
}
pub fn with_limits(limits: ReassemblyLimits) -> Self {
let budget = Arc::new(ReassemblyBudget::new(limits));
Self::with_limits_and_budget(limits, budget)
}
pub(crate) fn with_limits_and_budget(
limits: ReassemblyLimits,
budget: Arc<ReassemblyBudget>,
) -> Self {
Self {
pending: HashMap::new(),
buffered_cost: 0,
completed: VecDeque::new(),
completed_ids: HashSet::new(),
failed: VecDeque::new(),
failed_ids: HashSet::new(),
failure_tracking_saturated_until: 0,
capacity_rejected: VecDeque::new(),
capacity_rejected_ids: HashSet::new(),
capacity_tracking_saturated_until: 0,
limits: limits.normalized(),
budget,
}
}
fn mark_completed(&mut self, id: Uuid, expiry: u128) {
if self.completed_ids.insert(id) {
self.completed.push_back((id, expiry));
}
while self.completed.len() > self.limits.max_completed_ids {
if let Some((old, _)) = self.completed.pop_front() {
self.completed_ids.remove(&old);
}
}
}
pub fn pending_count(&self) -> usize {
self.pending.len()
}
pub fn remove_expired(&mut self) {
let _ = self.remove_expired_at(get_epoch_ms());
}
pub(crate) fn has_pending(&self) -> bool {
!self.pending.is_empty()
}
pub(crate) fn prepare_for_close(&mut self) -> bool {
let buffered_cost = &mut self.buffered_cost;
let budget = &self.budget;
let slot_overhead = self.limits.slot_overhead;
self.pending.retain(|_, pending| {
let retained = pending.peer_attributable
&& !(pending.failure_charged || pending.local_capacity_rejected);
if !retained {
let cost = pending.cost(slot_overhead);
*buffered_cost = buffered_cost.saturating_sub(cost);
budget.release(cost);
}
retained
});
self.clear_terminal_history();
!self.pending.is_empty()
}
pub(crate) fn discard_after_close_timer_failure(&mut self) {
for pending in self.pending.values() {
self.budget.release(pending.cost(self.limits.slot_overhead));
}
self.pending.clear();
self.buffered_cost = 0;
self.clear_terminal_history();
}
fn clear_terminal_history(&mut self) {
self.completed.clear();
self.completed_ids.clear();
self.failed.clear();
self.failed_ids.clear();
self.failure_tracking_saturated_until = 0;
self.capacity_rejected.clear();
self.capacity_rejected_ids.clear();
self.capacity_tracking_saturated_until = 0;
}
pub(crate) fn remove_expired_at(&mut self, now: u128) -> usize {
self.evict_expired_terminal_history(now);
let mut expired_count = 0_usize;
let mut expired_transmissions = Vec::new();
let buffered_cost = &mut self.buffered_cost;
let budget = &self.budget;
let slot_overhead = self.limits.slot_overhead;
self.pending.retain(|id, p| {
let alive = p.ts_ms.saturating_add(p.ttl_ms as u128) > now;
if !alive {
let cost = p.cost(slot_overhead);
*buffered_cost = buffered_cost.saturating_sub(cost);
budget.release(cost);
expired_transmissions.push(LogicalTransmission::new(*id, p.ts_ms, p.ttl_ms));
if !(p.failure_charged || p.local_capacity_rejected) && p.peer_attributable {
expired_count = expired_count.saturating_add(1);
}
}
alive
});
for transmission in expired_transmissions {
if !self.mark_failed_transmission(transmission, now) {
self.extend_failure_tracking_saturation(now);
}
}
expired_count
}
fn evict_expired_terminal_history(&mut self, now: u128) {
let completed_ids = &mut self.completed_ids;
self.completed.retain(|&(id, expiry)| {
let alive = expiry > now;
if !alive {
completed_ids.remove(&id);
}
alive
});
let failed_ids = &mut self.failed_ids;
self.failed.retain(|&(transmission, expiry)| {
let alive = expiry > now;
if !alive {
failed_ids.remove(&transmission);
}
alive
});
let capacity_rejected_ids = &mut self.capacity_rejected_ids;
self.capacity_rejected.retain(|&(id, expiry)| {
let alive = expiry > now;
if !alive {
capacity_rejected_ids.remove(&id);
}
alive
});
}
pub fn remove(&mut self, id: Uuid) {
if let Some(p) = self.pending.remove(&id) {
let cost = p.cost(self.limits.slot_overhead);
self.buffered_cost = self.buffered_cost.saturating_sub(cost);
self.budget.release(cost);
}
}
pub fn handle(&mut self, chunk: Chunk) -> Option<Bytes> {
self.handle_at(chunk, get_epoch_ms())
}
#[cfg(test)]
pub(crate) fn handle_retained(&mut self, chunk: Chunk) -> Option<RetainedReassembly> {
match self.handle_retained_outcome(chunk) {
ReassemblyOutcome::Complete(bytes) => Some(bytes),
ReassemblyOutcome::Incomplete | ReassemblyOutcome::Rejected(_) => None,
}
}
#[cfg(test)]
pub(crate) fn handle_retained_outcome(&mut self, chunk: Chunk) -> ReassemblyOutcome {
self.handle_retained_at(chunk, get_epoch_ms()).0
}
#[cfg(test)]
pub(crate) fn handle_retained_outcome_at(
&mut self,
chunk: Chunk,
now: u128,
) -> ReassemblyOutcome {
self.handle_retained_at(chunk, now).0
}
pub(crate) fn handle_retained_outcome_with_expiry(
&mut self,
chunk: Chunk,
peer_attributable: bool,
) -> (ReassemblyOutcome, usize) {
self.handle_retained_at_with_attribution(chunk, get_epoch_ms(), peer_attributable)
}
pub(super) fn handle_at(&mut self, chunk: Chunk, now: u128) -> Option<Bytes> {
match self.handle_retained_at(chunk, now).0 {
ReassemblyOutcome::Complete(bytes) => Some(bytes.into_bytes()),
ReassemblyOutcome::Incomplete | ReassemblyOutcome::Rejected(_) => None,
}
}
fn handle_retained_at(&mut self, chunk: Chunk, now: u128) -> (ReassemblyOutcome, usize) {
self.handle_retained_at_with_attribution(chunk, now, true)
}
fn handle_retained_at_with_attribution(
&mut self,
chunk: Chunk,
now: u128,
peer_attributable: bool,
) -> (ReassemblyOutcome, usize) {
let expired = self.remove_expired_at(now);
let outcome = match self.classify(&chunk, now) {
Ok(cost) => self.admit(chunk, cost, peer_attributable),
Err(reason) => {
tracing::debug!(?reason, id = ?chunk.meta.id, "reassembler dropped chunk");
let rejection = match reason.rejection() {
ReassemblyRejection::Invalid => {
if self.mark_logical_failure(&chunk, now) {
ReassemblyRejection::Invalid
} else {
ReassemblyRejection::Replay
}
}
ReassemblyRejection::Capacity => {
self.mark_pending_capacity_rejection(&chunk);
ReassemblyRejection::Capacity
}
ReassemblyRejection::Replay => ReassemblyRejection::Replay,
};
ReassemblyOutcome::Rejected(rejection)
}
};
(outcome, expired)
}
fn classify(&self, chunk: &Chunk, now: u128) -> std::result::Result<usize, Rejected> {
let meta = &chunk.meta;
let transmission = LogicalTransmission::new(meta.id, meta.ts_ms, meta.ttl_ms);
if self
.pending
.get(&meta.id)
.is_some_and(|pending| pending.failure_charged)
|| self.failed_ids.contains(&transmission)
{
return Err(Rejected::AlreadyFailed);
}
if self.capacity_rejected_ids.contains(&meta.id) {
return Err(Rejected::CapacityRejectedId);
}
if meta.ttl_ms > MAX_TTL_MS {
return Err(Rejected::TtlTooLarge);
}
if meta.ts_ms.saturating_sub(TS_OFFSET_TOLERANCE_MS) > now {
return Err(Rejected::FutureTimestamp);
}
if meta.ts_ms.saturating_add(meta.ttl_ms as u128) <= now {
return Err(Rejected::Expired);
}
let [position, total] = chunk.chunk;
if total == 0 || position >= total {
return Err(Rejected::Malformed);
}
if total > self.limits.max_chunks_per_message {
return Err(Rejected::TooManyChunks);
}
if chunk.data.len() > self.limits.max_chunk_data_len {
return Err(Rejected::ChunkTooLarge);
}
if self.completed_ids.contains(&meta.id) {
return Err(Rejected::AlreadyCompleted);
}
let buffered_for_id = match self.pending.get(&meta.id) {
None if self.capacity_tracking_saturated_until > now => {
return Err(Rejected::CapacityTrackingFull);
}
None if self.pending.len() >= self.limits.max_pending_messages => {
return Err(Rejected::PendingFull);
}
None => 0,
Some(p) => {
if p.total != total {
return Err(Rejected::TotalMismatch);
}
if p.ts_ms != meta.ts_ms || p.ttl_ms != meta.ttl_ms {
return Err(Rejected::MetadataMismatch);
}
if let Some(existing) = p.slots.get(&position) {
return if existing == &chunk.data {
Err(Rejected::DuplicatePosition)
} else {
Err(Rejected::ConflictingPosition)
};
}
p.data_bytes
}
};
if buffered_for_id.saturating_add(chunk.data.len()) > self.limits.max_message_bytes {
return Err(Rejected::PerMessageBytes);
}
let cost = chunk.data.len().saturating_add(self.limits.slot_overhead);
if self.buffered_cost.saturating_add(cost) > self.limits.max_peer_buffered_cost() {
return Err(Rejected::PeerBudget);
}
Ok(cost)
}
fn mark_logical_failure(&mut self, chunk: &Chunk, now: u128) -> bool {
let meta = &chunk.meta;
if let Some(pending) = self.pending.get_mut(&meta.id) {
if pending.failure_charged {
return false;
}
pending.failure_charged = true;
return true;
}
self.mark_failed_transmission(
LogicalTransmission::new(meta.id, meta.ts_ms, meta.ttl_ms),
now,
)
}
fn mark_failed_transmission(&mut self, transmission: LogicalTransmission, now: u128) -> bool {
if self.failed_ids.contains(&transmission) || self.failure_tracking_saturated_until > now {
return false;
}
if self.failed_ids.len() >= self.limits.max_completed_ids {
self.extend_failure_tracking_saturation(now);
return false;
}
let expiry = now.saturating_add(MAX_TTL_MS as u128);
self.failed_ids.insert(transmission);
self.failed.push_back((transmission, expiry));
true
}
fn extend_failure_tracking_saturation(&mut self, now: u128) {
self.failure_tracking_saturated_until = self
.failure_tracking_saturated_until
.max(now.saturating_add(MAX_TTL_MS as u128));
}
fn mark_pending_capacity_rejection(&mut self, chunk: &Chunk) {
let meta = &chunk.meta;
if let Some(pending) = self.pending.get_mut(&meta.id) {
if pending.ts_ms == meta.ts_ms && pending.ttl_ms == meta.ttl_ms {
pending.local_capacity_rejected = true;
return;
}
}
self.mark_capacity_rejected_id(meta.id, meta.ts_ms, meta.ttl_ms);
}
fn mark_capacity_rejected_id(&mut self, id: Uuid, ts_ms: u128, ttl_ms: u64) {
let expiry = ts_ms.saturating_add(ttl_ms.min(MAX_TTL_MS) as u128);
if self.capacity_rejected_ids.contains(&id) {
return;
}
if self.capacity_rejected_ids.len() >= self.limits.max_completed_ids {
self.capacity_tracking_saturated_until =
self.capacity_tracking_saturated_until.max(expiry);
return;
}
self.capacity_rejected_ids.insert(id);
self.capacity_rejected.push_back((id, expiry));
}
fn admit(&mut self, chunk: Chunk, cost: usize, peer_attributable: bool) -> ReassemblyOutcome {
if !self.budget.try_reserve(cost) {
self.mark_pending_capacity_rejection(&chunk);
tracing::debug!(
reason = ?Rejected::GlobalBudget,
id = ?chunk.meta.id,
"reassembler dropped chunk"
);
return ReassemblyOutcome::Rejected(ReassemblyRejection::Capacity);
}
let id = chunk.meta.id;
let [position, total] = chunk.chunk;
let mut pending = self.pending.remove(&id).unwrap_or_else(|| {
Pending::new(
total,
chunk.meta.ts_ms,
chunk.meta.ttl_ms,
peer_attributable,
)
});
pending.peer_attributable &= peer_attributable;
pending.data_bytes = pending.data_bytes.saturating_add(chunk.data.len());
pending.slots.insert(position, chunk.data);
self.buffered_cost = self.buffered_cost.saturating_add(cost);
if !pending.is_complete() {
self.pending.insert(id, pending);
#[cfg(all(test, feature = "dummy", not(target_family = "wasm")))]
crate::simulation::observe_reassembly_capacity(
self.budget.buffered_cost.load(Ordering::Acquire),
self.budget.limit,
self.buffered_cost,
self.limits.max_peer_buffered_cost(),
self.pending.len(),
self.limits.max_pending_messages,
);
return ReassemblyOutcome::Incomplete;
}
let output_cost = pending.data_bytes;
if !self.budget.try_reserve(output_cost) {
self.mark_capacity_rejected_id(id, pending.ts_ms, pending.ttl_ms);
let dropped_cost = pending.cost(self.limits.slot_overhead);
self.buffered_cost = self.buffered_cost.saturating_sub(dropped_cost);
self.budget.release(dropped_cost);
tracing::debug!(
reason = ?Rejected::GlobalBudget,
?id,
output_cost,
"reassembler dropped completed message before output allocation"
);
return ReassemblyOutcome::Rejected(ReassemblyRejection::Capacity);
}
#[cfg(all(test, feature = "dummy", not(target_family = "wasm")))]
crate::simulation::observe_reassembly_capacity(
self.budget.buffered_cost.load(Ordering::Acquire),
self.budget.limit,
self.buffered_cost,
self.limits.max_peer_buffered_cost(),
self.pending.len().saturating_add(1),
self.limits.max_pending_messages,
);
let done = pending;
let done_cost = done.cost(self.limits.slot_overhead);
let expiry = done.ts_ms.saturating_add(done.ttl_ms as u128);
self.buffered_cost = self.buffered_cost.saturating_sub(done_cost);
let bytes = done.assemble();
self.budget.release(done_cost);
self.mark_completed(id, expiry);
ReassemblyOutcome::Complete(RetainedReassembly {
bytes,
budget: self.budget.clone(),
cost: output_cost,
})
}
}
impl Drop for MessageReassembler {
fn drop(&mut self) {
self.budget.release(self.buffered_cost);
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Rejected {
TtlTooLarge,
FutureTimestamp,
Expired,
Malformed,
TooManyChunks,
ChunkTooLarge,
AlreadyCompleted,
AlreadyFailed,
CapacityRejectedId,
CapacityTrackingFull,
PendingFull,
TotalMismatch,
MetadataMismatch,
DuplicatePosition,
ConflictingPosition,
PerMessageBytes,
PeerBudget,
GlobalBudget,
}
impl Rejected {
const fn rejection(self) -> ReassemblyRejection {
match self {
Self::AlreadyCompleted
| Self::AlreadyFailed
| Self::DuplicatePosition => ReassemblyRejection::Replay,
Self::CapacityRejectedId
| Self::CapacityTrackingFull
| Self::PendingFull
| Self::PeerBudget
| Self::GlobalBudget => ReassemblyRejection::Capacity,
Self::TtlTooLarge
| Self::FutureTimestamp
| Self::Expired
| Self::Malformed
| Self::TooManyChunks
| Self::ChunkTooLarge
| Self::TotalMismatch
| Self::MetadataMismatch
| Self::ConflictingPosition
| Self::PerMessageBytes => ReassemblyRejection::Invalid,
}
}
}