use std::collections::{BTreeSet, HashMap};
use uuid::Uuid;
use crate::limits::{
MAX_PENDING_REPORTS_PER_CONNECTION, MAX_REPORT_PAGES, REPORT_IDLE_TIMEOUT,
REPORT_TOTAL_TIMEOUT, WireValidationError,
};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PageOutcome {
Pending,
Final {
accumulated_discovered_count: u32,
},
}
#[derive(Debug)]
struct PendingReport {
pages_received: BTreeSet<u32>,
total_pages: u32,
started_at: tokio::time::Instant,
last_page_at: tokio::time::Instant,
discovered_count: u32,
}
#[derive(Debug)]
pub struct ReportTracker {
pending: HashMap<Uuid, PendingReport>,
}
impl ReportTracker {
pub fn new() -> Self {
Self {
pending: HashMap::new(),
}
}
pub fn register_page(
&mut self,
report_id: Uuid,
page: u32,
total_pages: u32,
) -> Result<PageOutcome, WireValidationError> {
self.evict_expired();
if let Some(existing) = self.pending.get_mut(&report_id) {
if existing.total_pages != total_pages {
return Err(WireValidationError {
field: "pagination.total_pages",
message: format!(
"total_pages changed from {} to {total_pages} within report {report_id}",
existing.total_pages
),
});
}
if existing.pages_received.contains(&page) {
return Err(WireValidationError {
field: "pagination.page",
message: format!("duplicate page {page} for report {report_id}"),
});
}
existing.pages_received.insert(page);
existing.last_page_at = tokio::time::Instant::now();
if existing.pages_received.len() == existing.total_pages as usize {
#[expect(
clippy::expect_used,
reason = "infallible: we are inside a branch that already accessed this key via get_mut; the entry is guaranteed to exist"
)]
let report = self
.pending
.remove(&report_id)
.expect("report was just accessed");
Ok(PageOutcome::Final {
accumulated_discovered_count: report.discovered_count,
})
} else {
Ok(PageOutcome::Pending)
}
} else {
if self.pending.len() >= MAX_PENDING_REPORTS_PER_CONNECTION {
return Err(WireValidationError {
field: "pagination.report_id",
message: format!(
"too many concurrent paginated reports (max {MAX_PENDING_REPORTS_PER_CONNECTION})"
),
});
}
if total_pages > MAX_REPORT_PAGES {
return Err(WireValidationError {
field: "pagination.total_pages",
message: format!("total_pages is {total_pages}, max {MAX_REPORT_PAGES}"),
});
}
let now = tokio::time::Instant::now();
let mut pages_received = BTreeSet::new();
pages_received.insert(page);
if total_pages == 1 {
return Ok(PageOutcome::Final {
accumulated_discovered_count: 0,
});
}
self.pending.insert(
report_id,
PendingReport {
pages_received,
total_pages,
started_at: now,
last_page_at: now,
discovered_count: 0,
},
);
Ok(PageOutcome::Pending)
}
}
pub fn add_discovered_count(&mut self, report_id: Uuid, count: u32) {
if let Some(report) = self.pending.get_mut(&report_id) {
report.discovered_count = report.discovered_count.saturating_add(count);
}
}
pub fn evict_expired(&mut self) {
let now = tokio::time::Instant::now();
self.pending.retain(|id, report| {
let total_elapsed = now.duration_since(report.started_at);
let idle_elapsed = now.duration_since(report.last_page_at);
if total_elapsed >= REPORT_TOTAL_TIMEOUT {
tracing::warn!(
report_id = %id,
pages_received = report.pages_received.len(),
total_pages = report.total_pages,
"paginated report timed out (total timeout {}s exceeded)",
REPORT_TOTAL_TIMEOUT.as_secs()
);
return false;
}
if idle_elapsed >= REPORT_IDLE_TIMEOUT {
tracing::warn!(
report_id = %id,
pages_received = report.pages_received.len(),
total_pages = report.total_pages,
"paginated report timed out (idle timeout {}s exceeded)",
REPORT_IDLE_TIMEOUT.as_secs()
);
return false;
}
true
});
}
pub fn pending_count(&self) -> usize {
self.pending.len()
}
}
impl Default for ReportTracker {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test(start_paused = true)]
async fn single_page_report_returns_final() {
let mut tracker = ReportTracker::new();
let id = Uuid::new_v4();
let outcome = tracker.register_page(id, 1, 1).unwrap();
assert_eq!(
outcome,
PageOutcome::Final {
accumulated_discovered_count: 0
}
);
assert_eq!(tracker.pending_count(), 0);
}
#[tokio::test(start_paused = true)]
async fn multi_page_report_lifecycle() {
let mut tracker = ReportTracker::new();
let id = Uuid::new_v4();
let outcome = tracker.register_page(id, 1, 3).unwrap();
assert_eq!(outcome, PageOutcome::Pending);
assert_eq!(tracker.pending_count(), 1);
tracker.add_discovered_count(id, 100);
let outcome = tracker.register_page(id, 2, 3).unwrap();
assert_eq!(outcome, PageOutcome::Pending);
tracker.add_discovered_count(id, 50);
let outcome = tracker.register_page(id, 3, 3).unwrap();
assert_eq!(
outcome,
PageOutcome::Final {
accumulated_discovered_count: 150
}
);
assert_eq!(tracker.pending_count(), 0);
}
#[tokio::test(start_paused = true)]
async fn pages_out_of_order() {
let mut tracker = ReportTracker::new();
let id = Uuid::new_v4();
assert_eq!(
tracker.register_page(id, 3, 3).unwrap(),
PageOutcome::Pending
);
assert_eq!(
tracker.register_page(id, 1, 3).unwrap(),
PageOutcome::Pending
);
let outcome = tracker.register_page(id, 2, 3).unwrap();
assert_eq!(
outcome,
PageOutcome::Final {
accumulated_discovered_count: 0
}
);
}
#[tokio::test(start_paused = true)]
async fn duplicate_page_rejected() {
let mut tracker = ReportTracker::new();
let id = Uuid::new_v4();
tracker.register_page(id, 1, 3).unwrap();
let err = tracker.register_page(id, 1, 3).unwrap_err();
assert!(err.message.contains("duplicate page"));
}
#[tokio::test(start_paused = true)]
async fn total_pages_mismatch_rejected() {
let mut tracker = ReportTracker::new();
let id = Uuid::new_v4();
tracker.register_page(id, 1, 3).unwrap();
let err = tracker.register_page(id, 2, 5).unwrap_err();
assert!(err.message.contains("total_pages changed"));
}
#[tokio::test(start_paused = true)]
async fn max_concurrent_reports_enforced() {
let mut tracker = ReportTracker::new();
for i in 0..MAX_PENDING_REPORTS_PER_CONNECTION {
let id = Uuid::from_u128(i as u128);
tracker.register_page(id, 1, 2).unwrap();
}
let extra_id = Uuid::from_u128(999);
let err = tracker.register_page(extra_id, 1, 2).unwrap_err();
assert!(err.message.contains("too many concurrent"));
}
#[tokio::test(start_paused = true)]
async fn idle_timeout_evicts() {
let mut tracker = ReportTracker::new();
let id = Uuid::new_v4();
tracker.register_page(id, 1, 3).unwrap();
assert_eq!(tracker.pending_count(), 1);
tokio::time::advance(REPORT_IDLE_TIMEOUT + std::time::Duration::from_secs(1)).await;
tracker.evict_expired();
assert_eq!(tracker.pending_count(), 0);
}
#[tokio::test(start_paused = true)]
async fn total_timeout_evicts() {
let mut tracker = ReportTracker::new();
let id = Uuid::new_v4();
tracker.register_page(id, 1, 3).unwrap();
for step in 0..25 {
tokio::time::advance(std::time::Duration::from_secs(13)).await;
if step == 0 {
}
}
tokio::time::advance(REPORT_TOTAL_TIMEOUT + std::time::Duration::from_secs(1)).await;
tracker.evict_expired();
assert_eq!(tracker.pending_count(), 0);
}
#[tokio::test(start_paused = true)]
async fn add_discovered_count_no_op_for_unknown() {
let mut tracker = ReportTracker::new();
tracker.add_discovered_count(Uuid::new_v4(), 42);
assert_eq!(tracker.pending_count(), 0);
}
}