use bytes::Bytes;
use serde::{Deserialize, Serialize};
use std::collections::{BTreeMap, VecDeque};
use std::sync::Mutex;
use std::time::{Duration, Instant};
use tracing::{debug, trace, warn};
use crate::binary_protocol::{flags, ClientMessage, MessageType, PayloadType};
use crate::errors::{Error, Result};
const INITIAL_RTT: Duration = Duration::from_millis(100);
const INITIAL_RTO: Duration = Duration::from_millis(200);
const MIN_RTO: Duration = Duration::from_millis(50);
const MAX_RTO: Duration = Duration::from_secs(30);
const RTO_VARIANCE_FACTOR: u32 = 4;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AcknowledgeContent {
#[serde(rename = "AcknowledgedMessageType")]
pub message_type: String,
#[serde(rename = "AcknowledgedMessageId")]
pub message_id: String,
#[serde(rename = "AcknowledgedMessageSequenceNumber")]
pub sequence_number: i64,
#[serde(rename = "IsSequentialMessage")]
pub is_sequential: bool,
}
pub fn build_ack(received: &ClientMessage, is_sequential: bool) -> Result<ClientMessage> {
let content = AcknowledgeContent {
message_type: received.message_type.as_str().to_owned(),
message_id: received.message_id.to_string(),
sequence_number: received.sequence_number,
is_sequential,
};
let payload = serde_json::to_vec(&content)?;
let mut ack = ClientMessage::new(
MessageType::Acknowledge,
0,
PayloadType::Undefined,
Bytes::from(payload),
);
ack.flags = flags::SYN | flags::FIN;
Ok(ack)
}
pub fn parse_ack(message: &ClientMessage) -> Result<AcknowledgeContent> {
serde_json::from_slice(&message.payload)
.map_err(|e| Error::protocol(format!("malformed acknowledge payload: {e}")))
}
#[derive(Debug, Clone, Copy)]
pub struct RttEstimate {
pub smoothed: Duration,
pub variance: Duration,
pub rto: Duration,
pub samples: u64,
}
impl Default for RttEstimate {
fn default() -> Self {
Self {
smoothed: INITIAL_RTT,
variance: INITIAL_RTT / 2,
rto: INITIAL_RTO,
samples: 0,
}
}
}
impl RttEstimate {
pub fn record(&mut self, sample: Duration) {
if self.samples == 0 {
self.smoothed = sample;
self.variance = sample / 2;
} else {
let deviation = self.smoothed.abs_diff(sample);
self.variance = (self.variance * 3 + deviation) / 4;
self.smoothed = (self.smoothed * 7 + sample) / 8;
}
self.samples += 1;
self.rto = (self.smoothed + self.variance * RTO_VARIANCE_FACTOR).clamp(MIN_RTO, MAX_RTO);
}
}
#[derive(Debug)]
struct Pending {
wire: Bytes,
sequence: i64,
last_sent: Instant,
attempts: u32,
}
#[derive(Debug)]
pub enum RetransmitAction {
Idle,
Resend {
sequence: i64,
wire: Bytes,
},
GaveUp {
sequence: i64,
attempts: u32,
},
}
#[derive(Debug)]
pub struct OutgoingBuffer {
state: Mutex<OutgoingState>,
capacity: usize,
max_attempts: u32,
}
#[derive(Debug)]
struct OutgoingState {
pending: VecDeque<Pending>,
rtt: RttEstimate,
}
impl OutgoingBuffer {
pub fn new(capacity: usize, max_attempts: u32) -> Self {
Self {
state: Mutex::new(OutgoingState {
pending: VecDeque::new(),
rtt: RttEstimate::default(),
}),
capacity,
max_attempts,
}
}
#[must_use]
pub fn track(&self, wire: Bytes, sequence: i64) -> bool {
let mut state = self.lock();
if state.pending.len() >= self.capacity {
warn!(
sequence,
capacity = self.capacity,
"outgoing buffer is full; refusing to track another unacknowledged message"
);
return false;
}
state.pending.push_back(Pending {
wire,
sequence,
last_sent: Instant::now(),
attempts: 1,
});
true
}
pub fn acknowledge(&self, sequence: i64) -> bool {
let mut state = self.lock();
let Some(index) = state.pending.iter().position(|p| p.sequence == sequence) else {
trace!(
sequence,
"acknowledge for an unknown or already-retired message"
);
return false;
};
let entry = state.pending.remove(index).expect("index from position()");
if entry.attempts == 1 {
let sample = entry.last_sent.elapsed();
state.rtt.record(sample);
trace!(
sequence,
rtt_ms = sample.as_millis(),
rto_ms = state.rtt.rto.as_millis(),
"acknowledged, RTT updated"
);
}
true
}
pub fn poll_retransmit(&self) -> RetransmitAction {
let mut state = self.lock();
let rto = state.rtt.rto;
let Some(head) = state.pending.front_mut() else {
return RetransmitAction::Idle;
};
if head.last_sent.elapsed() <= rto {
return RetransmitAction::Idle;
}
if head.attempts >= self.max_attempts {
return RetransmitAction::GaveUp {
sequence: head.sequence,
attempts: head.attempts,
};
}
head.attempts += 1;
head.last_sent = Instant::now();
debug!(
sequence = head.sequence,
attempt = head.attempts,
rto_ms = rto.as_millis(),
"retransmitting unacknowledged message"
);
RetransmitAction::Resend {
sequence: head.sequence,
wire: head.wire.clone(),
}
}
pub fn len(&self) -> usize {
self.lock().pending.len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn rtt(&self) -> RttEstimate {
self.lock().rtt
}
pub fn clear(&self) {
self.lock().pending.clear();
}
fn lock(&self) -> std::sync::MutexGuard<'_, OutgoingState> {
self.state.lock().unwrap_or_else(|e| e.into_inner())
}
}
#[derive(Debug)]
pub struct IncomingBuffer {
messages: Mutex<BTreeMap<i64, ClientMessage>>,
capacity: usize,
}
impl IncomingBuffer {
pub fn new(capacity: usize) -> Self {
Self {
messages: Mutex::new(BTreeMap::new()),
capacity,
}
}
#[must_use]
pub fn insert(&self, message: ClientMessage) -> bool {
let mut messages = self.lock();
if messages.len() >= self.capacity && !messages.contains_key(&message.sequence_number) {
return false;
}
messages.insert(message.sequence_number, message);
true
}
pub fn take(&self, sequence: i64) -> Option<ClientMessage> {
self.lock().remove(&sequence)
}
pub fn len(&self) -> usize {
self.lock().len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
fn lock(&self) -> std::sync::MutexGuard<'_, BTreeMap<i64, ClientMessage>> {
self.messages.lock().unwrap_or_else(|e| e.into_inner())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn message(sequence: i64) -> ClientMessage {
ClientMessage::new(
MessageType::InputStreamData,
sequence,
PayloadType::Output,
Bytes::from_static(b"payload"),
)
}
#[test]
fn ack_uses_sequence_zero_and_syn_fin() {
let ack = build_ack(&message(42), true).unwrap();
assert_eq!(ack.sequence_number, 0);
assert_eq!(ack.flags, flags::SYN | flags::FIN);
assert_eq!(ack.message_type, MessageType::Acknowledge);
let content = parse_ack(&ack).unwrap();
assert_eq!(content.sequence_number, 42);
assert!(content.is_sequential);
assert_eq!(content.message_type, "input_stream_data");
}
#[test]
fn ack_json_uses_the_pascal_case_names_the_agent_expects() {
let ack = build_ack(&message(1), false).unwrap();
let json = std::str::from_utf8(&ack.payload).unwrap();
for key in [
"AcknowledgedMessageType",
"AcknowledgedMessageId",
"AcknowledgedMessageSequenceNumber",
"IsSequentialMessage",
] {
assert!(json.contains(key), "missing {key} in {json}");
}
}
#[test]
fn rtt_first_sample_seeds_the_estimate() {
let mut rtt = RttEstimate::default();
rtt.record(Duration::from_millis(50));
assert_eq!(rtt.smoothed, Duration::from_millis(50));
assert_eq!(rtt.samples, 1);
assert!(rtt.rto >= MIN_RTO && rtt.rto <= MAX_RTO);
}
#[test]
fn rtt_converges_towards_the_mean() {
let mut rtt = RttEstimate::default();
for _ in 0..40 {
rtt.record(Duration::from_millis(60));
}
assert!(
(55..=65).contains(&rtt.smoothed.as_millis()),
"srtt drifted to {:?}",
rtt.smoothed
);
assert!(rtt.rto >= MIN_RTO);
}
#[test]
fn rto_is_clamped_at_both_ends() {
let mut fast = RttEstimate::default();
fast.record(Duration::from_micros(10));
assert!(fast.rto >= MIN_RTO);
let mut slow = RttEstimate::default();
slow.record(Duration::from_secs(120));
assert!(slow.rto <= MAX_RTO);
}
#[test]
fn outgoing_buffer_retires_acknowledged_messages() {
let buffer = OutgoingBuffer::new(8, 3);
assert!(buffer.track(Bytes::from_static(b"a"), 0));
assert!(buffer.track(Bytes::from_static(b"b"), 1));
assert_eq!(buffer.len(), 2);
assert!(buffer.acknowledge(0));
assert_eq!(buffer.len(), 1);
assert!(!buffer.acknowledge(0));
assert_eq!(buffer.len(), 1);
}
#[test]
fn outgoing_buffer_refuses_to_overflow() {
let buffer = OutgoingBuffer::new(2, 3);
assert!(buffer.track(Bytes::from_static(b"a"), 0));
assert!(buffer.track(Bytes::from_static(b"b"), 1));
assert!(!buffer.track(Bytes::from_static(b"c"), 2));
assert_eq!(buffer.len(), 2);
}
#[test]
fn poll_retransmit_is_idle_until_the_rto_elapses() {
let buffer = OutgoingBuffer::new(8, 3);
assert!(matches!(buffer.poll_retransmit(), RetransmitAction::Idle));
assert!(buffer.track(Bytes::from_static(b"a"), 0));
assert!(matches!(buffer.poll_retransmit(), RetransmitAction::Idle));
}
#[test]
fn poll_retransmit_gives_up_after_the_attempt_budget() {
let buffer = OutgoingBuffer::new(8, 2);
assert!(buffer.track(Bytes::from_static(b"a"), 7));
{
let mut state = buffer.lock();
state.pending[0].last_sent = Instant::now() - Duration::from_secs(60);
}
match buffer.poll_retransmit() {
RetransmitAction::Resend { sequence, .. } => assert_eq!(sequence, 7),
other => panic!("expected a resend, got {other:?}"),
}
{
let mut state = buffer.lock();
state.pending[0].last_sent = Instant::now() - Duration::from_secs(60);
}
match buffer.poll_retransmit() {
RetransmitAction::GaveUp { sequence, attempts } => {
assert_eq!(sequence, 7);
assert_eq!(attempts, 2);
}
other => panic!("expected to give up, got {other:?}"),
}
}
#[test]
fn retransmitted_messages_do_not_pollute_the_rtt_estimate() {
let buffer = OutgoingBuffer::new(8, 5);
assert!(buffer.track(Bytes::from_static(b"a"), 0));
{
let mut state = buffer.lock();
state.pending[0].last_sent = Instant::now() - Duration::from_secs(60);
}
let _ = buffer.poll_retransmit();
let before = buffer.rtt();
assert!(buffer.acknowledge(0));
let after = buffer.rtt();
assert_eq!(
before.samples, after.samples,
"sample count must not change"
);
assert_eq!(before.smoothed, after.smoothed);
}
#[test]
fn incoming_buffer_holds_and_releases_by_sequence() {
let buffer = IncomingBuffer::new(4);
assert!(buffer.insert(message(5)));
assert!(buffer.insert(message(6)));
assert_eq!(buffer.len(), 2);
assert!(buffer.take(7).is_none());
assert_eq!(buffer.take(5).unwrap().sequence_number, 5);
assert_eq!(buffer.len(), 1);
}
#[test]
fn incoming_buffer_rejects_new_sequences_when_full() {
let buffer = IncomingBuffer::new(2);
assert!(buffer.insert(message(1)));
assert!(buffer.insert(message(2)));
assert!(!buffer.insert(message(3)), "must refuse when full");
assert!(buffer.insert(message(2)));
}
}