use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, Instant};
use tokio::sync::watch;
use weida_core::Error;
use weida_protocol::header::MAX_CURSOR_RECORD_LEN;
use weida_protocol::header::limits::MAX_REPORT_LEVELS;
use weida_protocol::{
CursorHeader, CursorLevel, FrameKind, ReportMode, encode_cursor_record, encode_preamble,
};
use crate::conn::ConnHandle;
use crate::transport::SendHalf;
pub(crate) struct ReportTable {
next: AtomicU64,
live: std::sync::Mutex<HashMap<u64, Arc<watch::Sender<CursorSet>>>>,
}
impl ReportTable {
pub(crate) fn new() -> ReportTable {
ReportTable {
next: AtomicU64::new(1),
live: std::sync::Mutex::new(HashMap::new()),
}
}
pub(crate) fn claim(&self, id: u64) -> Option<Arc<watch::Sender<CursorSet>>> {
self.live
.lock()
.unwrap_or_else(|e| e.into_inner())
.get(&id)
.map(Arc::clone)
}
pub(crate) fn release(&self, id: u64) {
self.live
.lock()
.unwrap_or_else(|e| e.into_inner())
.remove(&id);
}
}
pub(crate) fn order_report(conn: &ConnHandle) -> (u64, Cursors) {
let id = conn.reports.next.fetch_add(1, Ordering::Relaxed);
let (tx, rx) = watch::channel(CursorSet::default());
conn.reports
.live
.lock()
.unwrap_or_else(|e| e.into_inner())
.insert(id, Arc::new(tx));
(
id,
Cursors {
rx,
guard: ReportGuard {
id,
conn: ConnHandle::clone(conn),
},
},
)
}
struct ReportGuard {
id: u64,
conn: ConnHandle,
}
impl Drop for ReportGuard {
fn drop(&mut self) {
self.conn.reports.release(self.id);
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct CursorSet {
entries: [(u64, u64); MAX_REPORT_LEVELS],
len: u8,
}
impl CursorSet {
pub fn offset(&self, level: CursorLevel) -> Option<u64> {
let wire = level.to_wire();
self.used()
.iter()
.find(|(w, _)| *w == wire)
.map(|(_, offset)| *offset)
}
pub fn iter(&self) -> impl Iterator<Item = (CursorLevel, u64)> + '_ {
self.used().iter().map(|(wire, offset)| {
(
CursorLevel::from_wire(*wire).expect("only decoded levels are stored"),
*offset,
)
})
}
pub fn len(&self) -> usize {
usize::from(self.len)
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
fn used(&self) -> &[(u64, u64)] {
&self.entries[..usize::from(self.len)]
}
pub(crate) fn advance(&mut self, level: CursorLevel, offset: u64) -> bool {
let wire = level.to_wire();
let used = usize::from(self.len);
match self.entries[..used].binary_search_by_key(&wire, |(w, _)| *w) {
Ok(at) => {
if offset <= self.entries[at].1 {
return false;
}
self.entries[at].1 = offset;
true
}
Err(at) => {
if used == MAX_REPORT_LEVELS {
return false;
}
self.entries[at..=used].rotate_right(1);
self.entries[at] = (wire, offset);
self.len += 1;
true
}
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Reported {
Changed,
Waiting,
Ended,
}
#[derive(Debug)]
pub struct Cursors {
rx: watch::Receiver<CursorSet>,
guard: ReportGuard,
}
impl Cursors {
pub fn snapshot(&self) -> CursorSet {
*self.rx.borrow()
}
pub fn offset(&self, level: CursorLevel) -> Option<u64> {
self.rx.borrow().offset(level)
}
pub async fn changed(&mut self) -> Option<CursorSet> {
tokio::select! {
changed = self.rx.changed() => {
changed.ok()?;
Some(*self.rx.borrow_and_update())
}
_ = self.guard.conn.conn.closed() => None,
}
}
pub async fn changed_within(&mut self, deadline: Duration) -> Reported {
let exec = self.guard.conn.exec.clone();
match exec.within(deadline, self.changed()).await {
None => Reported::Waiting,
Some(None) => Reported::Ended,
Some(Some(_)) => Reported::Changed,
}
}
}
impl std::fmt::Debug for ReportGuard {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ReportGuard").field("id", &self.id).finish()
}
}
pub struct Reporter {
conn: ConnHandle,
report_id: u64,
levels: Vec<CursorLevel>,
mode: ReportMode,
stream: Option<SendHalf>,
broken: bool,
bytes: u64,
interval: Duration,
pending: [Option<Pending>; MAX_REPORT_LEVELS],
}
#[derive(Clone, Copy, Debug)]
struct Pending {
level: CursorLevel,
latest: u64,
emitted: Option<u64>,
at: Option<Instant>,
}
impl Reporter {
pub const DEFAULT_BYTES: u64 = 1024 * 1024;
pub const DEFAULT_INTERVAL: Duration = Duration::from_millis(100);
pub(crate) fn new(
conn: ConnHandle,
report_id: u64,
levels: Vec<CursorLevel>,
mode: ReportMode,
) -> Reporter {
Reporter {
conn,
report_id,
levels,
mode,
stream: None,
broken: false,
bytes: Reporter::DEFAULT_BYTES,
interval: Reporter::DEFAULT_INTERVAL,
pending: [None; MAX_REPORT_LEVELS],
}
}
pub fn with_granularity(mut self, bytes: u64, interval: Duration) -> Reporter {
self.bytes = bytes;
self.interval = interval;
self
}
pub fn levels(&self) -> &[CursorLevel] {
&self.levels
}
pub fn mode(&self) -> ReportMode {
self.mode
}
pub async fn report(&mut self, level: CursorLevel, offset: u64) -> Result<(), Error> {
if !self.levels.contains(&level) {
return Ok(());
}
let Some(slot) = self.slot(level) else {
return Ok(());
};
let entry = self.pending[slot].get_or_insert(Pending {
level,
latest: offset,
emitted: None,
at: None,
});
entry.latest = entry.latest.max(offset);
let latest = entry.latest;
if self.mode == ReportMode::FinalOnly || !self.worth_writing(slot) {
return Ok(());
}
self.emit(slot, latest).await;
Ok(())
}
pub async fn finish(mut self) -> Result<(), Error> {
for slot in 0..MAX_REPORT_LEVELS {
let Some(entry) = self.pending[slot] else {
continue;
};
if entry.emitted == Some(entry.latest) {
continue;
}
self.emit(slot, entry.latest).await;
}
if let Some(mut stream) = self.stream.take()
&& let Err(e) = stream.finish()
{
tracing::debug!(error = %e, "failed to finish a cursor stream");
}
Ok(())
}
fn slot(&self, level: CursorLevel) -> Option<usize> {
self.levels.iter().position(|l| *l == level)
}
fn worth_writing(&self, slot: usize) -> bool {
let Some(entry) = self.pending[slot] else {
return false;
};
match (entry.emitted, entry.at) {
(None, _) | (_, None) => true,
(Some(emitted), Some(at)) => {
entry.latest.saturating_sub(emitted) >= self.bytes || at.elapsed() >= self.interval
}
}
}
async fn emit(&mut self, slot: usize, offset: u64) {
if self.broken {
return;
}
let Some(entry) = self.pending[slot] else {
return;
};
let mut record = Vec::with_capacity(MAX_CURSOR_RECORD_LEN);
if let Err(e) = encode_cursor_record(entry.level, offset, &mut record) {
tracing::debug!(error = %e, "a cursor level has no wire representation");
return;
}
if self.stream.is_none() {
match self.open().await {
Ok(stream) => self.stream = Some(stream),
Err(e) => {
tracing::debug!(error = %e, "failed to open a cursor stream");
self.broken = true;
return;
}
}
}
let stream = self.stream.as_mut().expect("opened just above");
if let Err(e) = stream.write_all(&record).await {
tracing::debug!(error = %e, "failed to write a cursor record");
self.broken = true;
return;
}
if let Some(entry) = self.pending[slot].as_mut() {
entry.emitted = Some(offset);
entry.at = Some(Instant::now());
}
}
async fn open(&self) -> Result<SendHalf, Error> {
let mut stream = self.conn.open_uni().await?;
let head = CursorHeader {
report_id: self.report_id,
}
.encode();
let mut frame = Vec::with_capacity(head.len() + 10);
encode_preamble(FrameKind::Cursor, head.len() as u64, &mut frame);
frame.extend_from_slice(&head);
stream.write_all(&frame).await?;
Ok(stream)
}
}
impl std::fmt::Debug for Reporter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Reporter")
.field("report_id", &self.report_id)
.field("levels", &self.levels)
.field("mode", &self.mode)
.finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use super::*;
use weida_protocol::header::Acknowledgement;
#[test]
fn a_set_keeps_the_maximum_offset_per_level() {
let mut set = CursorSet::default();
let stored = CursorLevel::Known(Acknowledgement::Stored);
assert!(set.advance(stored, 100));
assert!(!set.advance(stored, 40));
assert!(!set.advance(stored, 100));
assert_eq!(set.offset(stored), Some(100));
assert_eq!(set.len(), 1);
}
#[test]
fn levels_are_ordered_by_their_wire_value() {
let mut set = CursorSet::default();
let app = CursorLevel::Application(17);
let accepted = CursorLevel::Known(Acknowledgement::Accepted);
assert!(set.advance(app, 1));
assert!(set.advance(accepted, 2));
assert_eq!(
set.iter().collect::<Vec<_>>(),
vec![(accepted, 2), (app, 1)]
);
}
#[test]
fn a_set_holds_exactly_the_cap() {
let mut set = CursorSet::default();
for i in 0..MAX_REPORT_LEVELS as u64 {
assert!(set.advance(
CursorLevel::Application(CursorLevel::APPLICATION_FLOOR + i),
i
));
}
assert_eq!(set.len(), MAX_REPORT_LEVELS);
assert!(!set.advance(CursorLevel::Known(Acknowledgement::Stored), 9));
assert_eq!(
set.offset(CursorLevel::Known(Acknowledgement::Stored)),
None
);
}
}