use crate::observability::VeloMetrics;
use bytes::{Buf, BufMut, Bytes, BytesMut};
use dashmap::DashSet;
use futures::future::BoxFuture;
use futures::task::AtomicWaker;
use parking_lot::Mutex;
use std::collections::{HashMap, VecDeque};
use std::fmt;
use std::future::Future;
use std::mem::size_of;
use std::pin::Pin;
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use std::task::{Context, Poll};
use thiserror::Error;
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
use tracing::{debug, trace, warn};
use uuid::Uuid;
use super::events::Outcome;
type WorkerId = u64;
type ArenaAllocation<T, E> = (usize, u64, Arc<Slot<T, E>>);
const RESPONSE_SLOT_CAPACITY: usize = u16::MAX as usize;
const MAX_GENERATION: u64 = (1u64 << 48) - 1;
#[derive(Debug, Error)]
pub(crate) enum DecodeError {
#[error("Response header too short: expected at least 18 bytes, got {0}")]
HeaderTooShort(usize),
#[error("Invalid headers length")]
InvalidHeadersLength,
#[error("Failed to deserialize headers: {0}")]
HeaderDeserializationError(#[from] rmp_serde::decode::Error),
}
#[allow(clippy::type_complexity)]
pub(crate) fn decode_response_header(
header: Bytes,
) -> Result<(ResponseId, Outcome, Option<HashMap<String, String>>), DecodeError> {
let mut header = header;
if header.len() < 19 {
return Err(DecodeError::HeaderTooShort(header.len()));
}
let response_id_value = header.get_u128_le();
let response_id = ResponseId::from_u128(response_id_value);
let outcome_byte = header.get_u8();
let outcome = if outcome_byte == 0 {
Outcome::Ok
} else {
Outcome::Error
};
let headers_len = header.get_u16_le() as usize;
let headers = if headers_len > 0 {
if header.remaining() < headers_len {
return Err(DecodeError::InvalidHeadersLength);
}
let headers_bytes = header.copy_to_bytes(headers_len);
let headers_map: HashMap<String, String> = rmp_serde::from_slice(&headers_bytes)?;
Some(headers_map)
} else {
None
};
Ok((response_id, outcome, headers))
}
#[derive(Debug, Error)]
pub(crate) enum EncodeError {
#[error("Response headers too large: {0} bytes exceeds u16 maximum of 65535")]
HeadersTooLarge(usize),
#[error("Failed to serialize headers: {0}")]
HeaderSerializationError(#[from] rmp_serde::encode::Error),
}
#[inline]
pub(crate) fn encode_response_header(
response_id: ResponseId,
outcome: Outcome,
headers: Option<HashMap<String, String>>,
) -> Result<Bytes, EncodeError> {
let headers_bytes = if let Some(ref h) = headers {
let msgpack_bytes = rmp_serde::to_vec(h)?;
Some(msgpack_bytes)
} else {
None
};
let headers_len = headers_bytes.as_ref().map(|b| b.len()).unwrap_or(0);
if headers_len > u16::MAX as usize {
return Err(EncodeError::HeadersTooLarge(headers_len));
}
let capacity = size_of::<u128>() + 1 + 2 + headers_len;
let mut bytes = BytesMut::with_capacity(capacity);
bytes.extend_from_slice(&response_id.as_u128().to_le_bytes());
let outcome_byte: u8 = match outcome {
Outcome::Ok => 0,
Outcome::Error => 1,
};
bytes.put_u8(outcome_byte);
bytes.extend_from_slice(&(headers_len as u16).to_le_bytes());
if let Some(hbytes) = headers_bytes {
bytes.extend_from_slice(&hbytes);
}
Ok(bytes.freeze())
}
#[inline]
fn encode_response_key(worker_id: WorkerId, slot_index: usize, generation: u64) -> u128 {
let worker_bits = worker_id as u128;
let slot_bits = ((slot_index as u16) as u128) << 64;
let gen_bits = (generation as u128) << 80;
worker_bits | slot_bits | gen_bits
}
#[inline]
fn decode_response_key(raw: u128) -> (WorkerId, usize, u64) {
let worker_id = (raw & 0xFFFF_FFFF_FFFF_FFFF) as u64;
let slot_index = ((raw >> 64) & 0xFFFF) as u16;
let generation = ((raw >> 80) & 0xFFFF_FFFF_FFFF) as u64;
(worker_id, slot_index as usize, generation)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct ResponseId(Uuid);
impl ResponseId {
pub(crate) fn from_u128(val: u128) -> Self {
Self(Uuid::from_u128(val))
}
pub(crate) fn as_u128(&self) -> u128 {
self.0.as_u128()
}
pub(crate) fn worker_id(&self) -> WorkerId {
let (worker_id, _, _) = decode_response_key(self.as_u128());
worker_id
}
pub(crate) fn slot_index(&self) -> usize {
let (_, slot_index, _) = decode_response_key(self.as_u128());
slot_index
}
pub(crate) fn generation(&self) -> u64 {
let (_, _, generation) = decode_response_key(self.as_u128());
generation
}
}
impl fmt::Display for ResponseId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.0)
}
}
pub struct ResponseAwaiter {
response_id: ResponseId,
manager: Arc<ResponseManagerInner>,
slot: Arc<Slot<Option<Bytes>, String>>,
index: usize,
consumed: bool,
permit: Option<OwnedSemaphorePermit>,
}
impl ResponseAwaiter {
fn new(
manager: Arc<ResponseManagerInner>,
slot: Arc<Slot<Option<Bytes>, String>>,
index: usize,
generation: u64,
permit: OwnedSemaphorePermit,
) -> Self {
let response_id = manager.encode_key(index, generation);
Self {
response_id,
manager,
slot,
index,
consumed: false,
permit: Some(permit),
}
}
pub fn response_id(&self) -> ResponseId {
self.response_id
}
pub async fn recv(&mut self) -> Result<Option<Bytes>, String> {
if self.consumed {
return Err("response awaiter already consumed".to_string());
}
let result = self.slot.wait_and_take().await;
self.consumed = true;
self.recycle();
match result {
Some(outcome) => outcome,
None => Err("response awaiter dropped before completion".to_string()),
}
}
pub fn poll_recv(
&mut self,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<Option<Bytes>, String>> {
use std::task::Poll;
if self.consumed {
return Poll::Ready(Err("response awaiter already consumed".to_string()));
}
match self.slot.poll_wait(cx) {
Poll::Ready(result) => {
self.consumed = true;
self.recycle();
Poll::Ready(match result {
Some(outcome) => outcome,
None => Err("response awaiter dropped before completion".to_string()),
})
}
Poll::Pending => Poll::Pending,
}
}
fn recycle(&mut self) {
if let Some(permit) = self.permit.take() {
self.manager.recycle_slot(self.index, permit);
}
}
}
impl Drop for ResponseAwaiter {
fn drop(&mut self) {
if !self.consumed {
self.recycle();
}
}
}
impl fmt::Debug for ResponseAwaiter {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ResponseAwaiter")
.field("response_id", &self.response_id)
.field("consumed", &self.consumed)
.finish()
}
}
#[derive(Debug, thiserror::Error)]
pub enum ResponseRegistrationError {
#[error("response slot capacity ({capacity}) exhausted; {pending} in flight")]
Exhausted { capacity: usize, pending: usize },
}
#[must_use = "RegisterOutcome::Backpressured must be awaited to acquire a slot"]
pub enum RegisterOutcome {
Allocated(ResponseAwaiter),
Backpressured(SlotBackpressure),
}
#[must_use = "SlotBackpressure must be awaited to acquire a response slot"]
pub struct SlotBackpressure {
fut: BoxFuture<'static, ResponseAwaiter>,
}
impl Future for SlotBackpressure {
type Output = ResponseAwaiter;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<ResponseAwaiter> {
self.fut.as_mut().poll(cx)
}
}
impl fmt::Debug for SlotBackpressure {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SlotBackpressure").finish_non_exhaustive()
}
}
pub(crate) struct ResponseManager {
inner: Arc<ResponseManagerInner>,
}
impl ResponseManager {
#[allow(dead_code)]
pub fn new(worker_id: WorkerId) -> Self {
Self::with_observability(worker_id, None)
}
pub fn with_observability(
worker_id: WorkerId,
observability: Option<Arc<VeloMetrics>>,
) -> Self {
Self {
inner: Arc::new(ResponseManagerInner::new(
worker_id,
observability,
RESPONSE_SLOT_CAPACITY,
)),
}
}
#[cfg(test)]
pub fn with_capacity(
worker_id: WorkerId,
capacity: usize,
observability: Option<Arc<VeloMetrics>>,
) -> Self {
assert!(
(1..=RESPONSE_SLOT_CAPACITY).contains(&capacity),
"response slot capacity {capacity} out of range (1..={RESPONSE_SLOT_CAPACITY})"
);
Self {
inner: Arc::new(ResponseManagerInner::new(
worker_id,
observability,
capacity,
)),
}
}
pub fn register_outcome(&self) -> Result<ResponseAwaiter, ResponseRegistrationError> {
self.inner.try_acquire().ok_or_else(|| {
self.inner.record_exhaustion();
ResponseRegistrationError::Exhausted {
capacity: self.inner.capacity,
pending: self.inner.pending_outcome_count(),
}
})
}
pub fn try_register_outcome(&self) -> RegisterOutcome {
if let Some(awaiter) = self.inner.try_acquire() {
RegisterOutcome::Allocated(awaiter)
} else {
self.inner.record_exhaustion();
RegisterOutcome::Backpressured(SlotBackpressure {
fut: Box::pin(ResponseManagerInner::acquire_owned(Arc::clone(&self.inner))),
})
}
}
pub fn complete_outcome(
&self,
response_id: ResponseId,
outcome: Result<Option<Bytes>, String>,
) -> bool {
self.inner.complete_outcome(response_id, outcome)
}
pub fn pending_outcome_count(&self) -> usize {
self.inner.pending_outcome_count()
}
}
struct ResponseManagerInner {
worker_id: WorkerId,
arena: Arc<SlotArena<Option<Bytes>, String>>,
slot_sem: Arc<Semaphore>,
pending: AtomicUsize,
capacity: usize,
observability: Option<Arc<VeloMetrics>>,
}
impl ResponseManagerInner {
fn new(worker_id: WorkerId, observability: Option<Arc<VeloMetrics>>, capacity: usize) -> Self {
let arena = SlotArena::with_capacity(capacity);
Self {
worker_id,
arena,
slot_sem: Arc::new(Semaphore::new(capacity)),
pending: AtomicUsize::new(0),
capacity,
observability,
}
}
fn try_acquire(self: &Arc<Self>) -> Option<ResponseAwaiter> {
let permit = Arc::clone(&self.slot_sem).try_acquire_owned().ok()?;
let (index, generation, slot) = self
.arena
.allocate()
.expect("arena and semaphore permits must stay in sync");
self.mark_pending();
Some(ResponseAwaiter::new(
Arc::clone(self),
slot,
index,
generation,
permit,
))
}
async fn acquire_owned(inner: Arc<Self>) -> ResponseAwaiter {
let permit = Arc::clone(&inner.slot_sem)
.acquire_owned()
.await
.expect("response slot semaphore must not be closed");
let (index, generation, slot) = inner
.arena
.allocate()
.expect("arena and semaphore permits must stay in sync");
inner.mark_pending();
ResponseAwaiter::new(Arc::clone(&inner), slot, index, generation, permit)
}
fn recycle_slot(&self, index: usize, permit: OwnedSemaphorePermit) {
let retired = self.arena.recycle(index);
if retired {
permit.forget();
}
let pending = self.pending.fetch_sub(1, Ordering::Release) - 1;
if let Some(metrics) = self.observability.as_ref() {
metrics.set_pending_responses(pending);
}
}
fn mark_pending(&self) {
let pending = self.pending.fetch_add(1, Ordering::AcqRel) + 1;
if let Some(metrics) = self.observability.as_ref() {
metrics.set_pending_responses(pending);
}
}
fn record_exhaustion(&self) {
if let Some(metrics) = self.observability.as_ref() {
metrics.inc_response_slot_exhausted();
}
}
fn encode_key(&self, slot_index: usize, generation: u64) -> ResponseId {
ResponseId::from_u128(encode_response_key(self.worker_id, slot_index, generation))
}
fn decode_key(&self, response_id: ResponseId) -> Option<(u64, usize, u64)> {
Some(decode_response_key(response_id.as_u128()))
}
fn complete_outcome(
&self,
response_id: ResponseId,
outcome: Result<Option<Bytes>, String>,
) -> bool {
trace!(
response_id = %response_id,
"ResponseManager.complete_outcome() called - decoding response_id"
);
let (worker_id, slot_index, expected_generation) = match self.decode_key(response_id) {
Some(parts) => parts,
None => {
warn!(response_id = %response_id, "invalid response identifier");
return false;
}
};
trace!(
response_id = %response_id,
worker_id,
slot_index,
expected_generation,
"ResponseManager decoded response_id successfully"
);
if worker_id != self.worker_id {
warn!(
response_id = %response_id,
expected_worker = self.worker_id,
received_worker = worker_id,
"response targeted wrong worker"
);
return false;
}
if slot_index >= self.capacity {
warn!(
response_id = %response_id,
slot_index,
capacity = self.capacity,
"response slot index out of bounds"
);
return false;
}
if !self.arena.is_allocated(slot_index) {
warn!(
response_id = %response_id,
slot_index,
"response slot has been recycled - discarding stale response"
);
return false;
}
let slot = match self.arena.slot(slot_index) {
Some(slot) => {
trace!(
response_id = %response_id,
slot_index,
"ResponseManager found slot in arena"
);
slot
}
None => {
warn!(
response_id = %response_id,
slot_index,
"response slot not found (likely freed)"
);
return false;
}
};
trace!(
response_id = %response_id,
slot_index,
expected_generation,
"ResponseManager completing slot outcome"
);
let completed = match outcome {
Ok(payload) => {
trace!(
response_id = %response_id,
slot_index,
payload_present = payload.is_some(),
expected_generation,
"ResponseManager calling slot.complete_ok()"
);
slot.complete_ok(payload, expected_generation)
}
Err(err) => {
trace!(
response_id = %response_id,
slot_index,
error = %err,
expected_generation,
"ResponseManager calling slot.complete_err()"
);
slot.complete_err(err, expected_generation)
}
};
if completed {
debug!(response_id = %response_id, slot_index, "ResponseManager: slot.complete_ok/err RETURNED TRUE - awaiter should wake");
} else {
warn!(
response_id = %response_id,
slot_index,
expected_generation,
"ResponseManager: slot.complete_ok/err RETURNED FALSE - response outcome already completed, cancelled, or generation mismatch"
);
}
completed
}
fn pending_outcome_count(&self) -> usize {
self.pending.load(Ordering::Acquire)
}
}
impl Clone for ResponseManager {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
}
}
}
impl fmt::Debug for ResponseManager {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ResponseManager")
.field("worker_id", &self.inner.worker_id)
.field("pending", &self.pending_outcome_count())
.field("capacity", &self.inner.capacity)
.finish()
}
}
struct SlotState<T, E> {
value: Option<Result<T, E>>,
generation: u64,
did_finish: bool,
}
impl<T, E> SlotState<T, E> {
fn new(generation: u64) -> Self {
Self {
value: None,
generation,
did_finish: false,
}
}
fn can_complete(&self, expected_gen: u64) -> bool {
self.generation == expected_gen && self.value.is_none()
}
fn complete(&mut self, value: Result<T, E>, expected_gen: u64) -> bool {
if !self.can_complete(expected_gen) {
return false;
}
self.value = Some(value);
self.generation = self.generation.wrapping_add(1);
self.did_finish = true;
true
}
fn take_value(&mut self) -> Option<Result<T, E>> {
self.value.take()
}
fn recycle(&mut self) -> u64 {
if !self.did_finish {
self.generation = self.generation.wrapping_add(1);
}
self.value = None;
self.did_finish = false;
self.generation
}
fn generation(&self) -> u64 {
self.generation
}
}
struct Slot<T, E> {
waker: AtomicWaker,
state: Mutex<SlotState<T, E>>,
}
impl<T, E> Slot<T, E> {
pub fn new() -> Self {
Self {
waker: AtomicWaker::new(),
state: Mutex::new(SlotState::new(0)),
}
}
pub fn complete_ok(&self, val: T, expected_generation: u64) -> bool {
self.finish(Ok(val), expected_generation)
}
pub fn complete_err(&self, err: E, expected_generation: u64) -> bool {
self.finish(Err(err), expected_generation)
}
fn finish(&self, res: Result<T, E>, expected_generation: u64) -> bool {
use tracing::{debug, trace};
trace!("Slot.finish() called - locking state");
let mut guard = self.state.lock();
let success = guard.complete(res, expected_generation);
if success {
trace!("Slot.finish() - value set, dropping lock");
drop(guard);
trace!("Slot.finish() - waking waiter");
self.waker.wake();
debug!("Slot.finish() - waiter woken, returning true");
} else {
debug!("Slot.finish() - generation mismatch or already completed, returning false");
}
success
}
pub async fn wait_and_take(&self) -> Option<Result<T, E>> {
std::future::poll_fn(|cx| self.poll_wait(cx)).await
}
pub fn poll_wait(
&self,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Result<T, E>>> {
use std::task::Poll;
self.waker.register(cx.waker());
if let Some(val) = self.state.lock().take_value() {
return Poll::Ready(Some(val));
}
Poll::Pending
}
#[expect(dead_code)]
#[doc(hidden)]
pub fn try_take(&self) -> Option<Result<T, E>> {
self.state.lock().take_value()
}
pub fn recycle(&self) -> u64 {
self.state.lock().recycle()
}
pub fn current_generation(&self) -> u64 {
self.state.lock().generation()
}
}
struct SlotArena<T, E> {
slots: Vec<Arc<Slot<T, E>>>,
free: parking_lot::Mutex<VecDeque<usize>>,
allocated: DashSet<usize>,
}
impl<T, E> SlotArena<T, E> {
pub fn with_capacity(cap: usize) -> Arc<Self> {
let slots = (0..cap).map(|_| Arc::new(Slot::new())).collect();
Arc::new(Self {
slots,
free: parking_lot::Mutex::new((0..cap).collect()),
allocated: DashSet::new(),
})
}
pub fn allocate(&self) -> Option<ArenaAllocation<T, E>> {
let mut free = self.free.lock();
free.pop_front().map(|i| {
self.allocated.insert(i);
let generation = self.slots[i].current_generation();
(i, generation, self.slots[i].clone())
})
}
pub fn slot(&self, index: usize) -> Option<Arc<Slot<T, E>>> {
self.slots.get(index).cloned()
}
#[expect(dead_code)]
pub fn complete(&self, index: usize, val: Result<T, E>, expected_generation: u64) -> bool {
let slot = &self.slots[index];
match val {
Ok(v) => slot.complete_ok(v, expected_generation),
Err(e) => slot.complete_err(e, expected_generation),
}
}
pub fn recycle(&self, index: usize) -> bool {
let new_generation = self.slots[index].recycle();
self.allocated.remove(&index);
if new_generation <= MAX_GENERATION {
self.free.lock().push_back(index);
false
} else {
true
}
}
pub fn is_allocated(&self, index: usize) -> bool {
self.allocated.contains(&index)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn outcome_registration_and_completion() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let mut awaiter = manager.register_outcome().expect("allocate slot");
let response_id = awaiter.response_id();
assert!(manager.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"ping")))));
let bytes = awaiter.recv().await.unwrap().unwrap();
let hash = xxhash_rust::xxh3::xxh3_64(&bytes);
assert_eq!(hash, xxhash_rust::xxh3::xxh3_64(b"ping"));
}
#[tokio::test]
async fn deferred_send_failure_completes_awaiter_via_header_decode() {
use super::super::messages::{
ActiveMessage, MessageMetadata, decode_response_id_from_request_header,
};
let worker_id = 7;
let manager = ResponseManager::new(worker_id);
let mut awaiter = manager.register_outcome().expect("allocate slot");
let response_id = awaiter.response_id();
let (header, _payload, _mt) = ActiveMessage {
metadata: MessageMetadata::new_unary(response_id, "any_handler".to_string(), None),
payload: Bytes::from_static(b""),
}
.encode()
.expect("encode");
let decoded = decode_response_id_from_request_header(&header).expect("decode id");
assert_eq!(decoded.as_u128(), response_id.as_u128());
assert!(manager.complete_outcome(decoded, Err("peer disconnected".to_string())));
let err = awaiter.recv().await.expect_err("should be Err");
assert_eq!(err, "peer disconnected");
}
#[tokio::test]
async fn malformed_header_does_not_spuriously_complete_awaiter() {
use super::super::messages::decode_response_id_from_request_header;
let worker_id = 7;
let manager = ResponseManager::new(worker_id);
let _awaiter = manager.register_outcome().expect("allocate slot");
let short = Bytes::from_static(&[1u8, 2, 3]);
assert!(decode_response_id_from_request_header(&short).is_none());
let bogus_id = ResponseId::from_u128(0u128);
assert!(!manager.complete_outcome(bogus_id, Err("x".to_string())));
}
#[tokio::test]
async fn drop_recycles_slot() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let awaiter = manager.register_outcome().expect("allocate slot");
let response_id = awaiter.response_id();
drop(awaiter);
assert!(!manager.complete_outcome(response_id, Ok(None)));
assert_eq!(manager.pending_outcome_count(), 0);
}
const TEST_CAPACITY: usize = 16;
#[tokio::test]
async fn allocation_exhaustion() {
let worker_id = 42;
let manager = ResponseManager::with_capacity(worker_id, TEST_CAPACITY, None);
let mut awaiters = Vec::with_capacity(TEST_CAPACITY);
for _ in 0..TEST_CAPACITY {
let awaiter = manager.register_outcome().expect("allocate slot");
awaiters.push(awaiter);
}
match manager.register_outcome() {
Err(ResponseRegistrationError::Exhausted { capacity, pending }) => {
assert_eq!(capacity, TEST_CAPACITY);
assert_eq!(pending, TEST_CAPACITY);
}
other => panic!("expected Exhausted, got {:?}", other),
}
let awaiter = awaiters.pop().expect("awaiter");
drop(awaiter);
let awaiter = manager.register_outcome().expect("allocate after recycle");
drop(awaiter);
}
#[tokio::test]
async fn try_register_outcome_reports_backpressure_when_full() {
let manager = ResponseManager::with_capacity(0, TEST_CAPACITY, None);
let mut awaiters = Vec::with_capacity(TEST_CAPACITY);
for _ in 0..TEST_CAPACITY {
awaiters.push(manager.register_outcome().expect("allocate"));
}
match manager.try_register_outcome() {
RegisterOutcome::Allocated(_) => panic!("expected backpressure at capacity"),
RegisterOutcome::Backpressured(_) => {}
}
}
#[tokio::test]
async fn slot_backpressure_resolves_after_recycle() {
use tokio::time::{Duration, timeout};
let manager = ResponseManager::with_capacity(0, TEST_CAPACITY, None);
let mut awaiters = Vec::with_capacity(TEST_CAPACITY);
for _ in 0..TEST_CAPACITY {
awaiters.push(manager.register_outcome().expect("allocate"));
}
let bp = match manager.try_register_outcome() {
RegisterOutcome::Backpressured(bp) => bp,
RegisterOutcome::Allocated(_) => panic!("expected backpressure"),
};
let mut handle = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(25)).await;
drop(awaiters.pop());
awaiters
});
let awaiter = timeout(Duration::from_secs(1), bp)
.await
.expect("backpressure resolved");
drop(awaiter);
let _remaining = (&mut handle).await.expect("drop task");
}
#[tokio::test]
async fn cancelling_slot_backpressure_is_safe() {
let manager = ResponseManager::with_capacity(0, TEST_CAPACITY, None);
let mut awaiters = Vec::with_capacity(TEST_CAPACITY);
for _ in 0..TEST_CAPACITY {
awaiters.push(manager.register_outcome().expect("allocate"));
}
if let RegisterOutcome::Backpressured(bp) = manager.try_register_outcome() {
drop(bp);
} else {
panic!("expected backpressure");
}
drop(awaiters.pop());
let awaiter = manager
.register_outcome()
.expect("allocate after bp cancel");
drop(awaiter);
}
#[tokio::test]
async fn slot_backpressure_debug_format() {
let manager = ResponseManager::with_capacity(0, TEST_CAPACITY, None);
let mut awaiters = Vec::with_capacity(TEST_CAPACITY);
for _ in 0..TEST_CAPACITY {
awaiters.push(manager.register_outcome().expect("allocate"));
}
let bp = match manager.try_register_outcome() {
RegisterOutcome::Backpressured(bp) => bp,
RegisterOutcome::Allocated(_) => panic!("expected backpressure"),
};
assert!(
format!("{:?}", bp).contains("SlotBackpressure"),
"Debug fmt should name the type"
);
}
#[tokio::test]
async fn try_register_outcome_fires_exhaustion_metric() {
use crate::observability::VeloMetrics;
use crate::observability::test_helpers::MetricSnapshot;
let registry = prometheus::Registry::new();
let metrics = Arc::new(VeloMetrics::register(®istry).expect("metrics"));
let manager = ResponseManager::with_capacity(0, TEST_CAPACITY, Some(metrics));
let mut awaiters = Vec::with_capacity(TEST_CAPACITY);
for _ in 0..TEST_CAPACITY {
awaiters.push(manager.register_outcome().expect("allocate"));
}
for _ in 0..2 {
match manager.try_register_outcome() {
RegisterOutcome::Backpressured(_) => {}
RegisterOutcome::Allocated(_) => panic!("expected backpressure"),
}
}
let snapshot = MetricSnapshot::from_registry(®istry);
let value = snapshot.counter("velo_messenger_response_slot_exhausted_total", &[]);
assert!(
value >= 2.0,
"try_register_outcome should also fire the exhaustion counter"
);
}
#[tokio::test]
async fn exhaustion_records_metric_and_details() {
use crate::observability::VeloMetrics;
use crate::observability::test_helpers::MetricSnapshot;
let registry = prometheus::Registry::new();
let metrics = Arc::new(VeloMetrics::register(®istry).expect("metrics"));
let manager = ResponseManager::with_capacity(0, TEST_CAPACITY, Some(metrics));
let mut awaiters = Vec::with_capacity(TEST_CAPACITY);
for _ in 0..TEST_CAPACITY {
awaiters.push(manager.register_outcome().expect("allocate"));
}
for _ in 0..3 {
match manager.register_outcome() {
Err(ResponseRegistrationError::Exhausted { capacity, pending }) => {
assert_eq!(capacity, TEST_CAPACITY);
assert_eq!(pending, TEST_CAPACITY);
}
Ok(_) => panic!("expected Exhausted"),
}
}
let snapshot = MetricSnapshot::from_registry(®istry);
let value = snapshot.counter("velo_messenger_response_slot_exhausted_total", &[]);
assert!(value >= 3.0, "exhaustion counter should fire per attempt");
}
#[tokio::test]
async fn recv_works_with_tokio_select() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let mut awaiter = manager.register_outcome().expect("allocate slot");
let response_id = awaiter.response_id();
let manager_clone = manager.clone();
tokio::spawn(async move {
tokio::time::sleep(tokio::time::Duration::from_millis(50)).await;
manager_clone.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"delayed"))));
});
let result = tokio::select! {
res = awaiter.recv() => res,
_ = tokio::time::sleep(tokio::time::Duration::from_secs(1)) => {
Err("timeout".to_string())
}
};
assert!(result.is_ok());
assert_eq!(result.unwrap().unwrap(), Bytes::from_static(b"delayed"));
}
#[tokio::test]
async fn recv_prevents_double_consumption() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let mut awaiter = manager.register_outcome().expect("allocate slot");
let response_id = awaiter.response_id();
manager.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"data"))));
let first = awaiter.recv().await;
assert!(first.is_ok());
let second = awaiter.recv().await;
assert!(second.is_err());
assert_eq!(second.unwrap_err(), "response awaiter already consumed");
}
#[tokio::test]
async fn complete_with_wrong_worker_id() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let mut awaiter = manager.register_outcome().expect("allocate slot");
let response_id = awaiter.response_id();
let other_manager = ResponseManager::new(999);
assert!(
!other_manager.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"data"))))
);
assert!(manager.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"correct")))));
let result = awaiter.recv().await.unwrap().unwrap();
assert_eq!(result, Bytes::from_static(b"correct"));
}
#[tokio::test]
async fn complete_with_out_of_bounds_slot() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let fake_slot_index = (RESPONSE_SLOT_CAPACITY + 1000) as u16;
let worker_bits = worker_id as u128;
let slot_bits = (fake_slot_index as u128) << 64;
let fake_id = ResponseId::from_u128(worker_bits | slot_bits);
assert!(!manager.complete_outcome(fake_id, Ok(None)));
}
#[tokio::test]
async fn complete_after_recycle_is_rejected() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let awaiter = manager.register_outcome().expect("allocate slot");
let response_id = awaiter.response_id();
drop(awaiter);
let _new_awaiter = manager.register_outcome().expect("allocate new slot");
assert!(!manager.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"stale")))));
}
#[tokio::test]
async fn double_completion_fails() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let mut awaiter = manager.register_outcome().expect("allocate slot");
let response_id = awaiter.response_id();
assert!(manager.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"first")))));
assert!(!manager.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"second")))));
let result = awaiter.recv().await.unwrap().unwrap();
assert_eq!(result, Bytes::from_static(b"first"));
}
#[tokio::test]
async fn complete_with_error_outcome() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let mut awaiter = manager.register_outcome().expect("allocate slot");
let response_id = awaiter.response_id();
assert!(manager.complete_outcome(response_id, Err("operation failed".to_string())));
let result = awaiter.recv().await;
assert!(result.is_err());
assert_eq!(result.unwrap_err(), "operation failed");
}
#[tokio::test]
async fn pending_count_tracking() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
assert_eq!(manager.pending_outcome_count(), 0);
let mut awaiter1 = manager.register_outcome().expect("allocate 1");
assert_eq!(manager.pending_outcome_count(), 1);
let awaiter2 = manager.register_outcome().expect("allocate 2");
assert_eq!(manager.pending_outcome_count(), 2);
let mut awaiter3 = manager.register_outcome().expect("allocate 3");
assert_eq!(manager.pending_outcome_count(), 3);
manager.complete_outcome(awaiter1.response_id(), Ok(None));
awaiter1.recv().await.unwrap();
assert_eq!(manager.pending_outcome_count(), 2);
drop(awaiter2);
assert_eq!(manager.pending_outcome_count(), 1);
manager.complete_outcome(awaiter3.response_id(), Ok(None));
awaiter3.recv().await.unwrap();
assert_eq!(manager.pending_outcome_count(), 0);
}
#[tokio::test]
async fn slot_reuse_after_recycling() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let mut first_awaiter = manager.register_outcome().expect("allocate first");
let first_id = first_awaiter.response_id();
manager.complete_outcome(first_id, Ok(Some(Bytes::from_static(b"first"))));
first_awaiter.recv().await.unwrap();
let mut second_awaiter = manager.register_outcome().expect("allocate second");
let second_id = second_awaiter.response_id();
assert_ne!(first_id, second_id);
assert!(!manager.complete_outcome(first_id, Ok(Some(Bytes::from_static(b"stale")))));
assert!(manager.complete_outcome(second_id, Ok(Some(Bytes::from_static(b"second")))));
let result = second_awaiter.recv().await.unwrap().unwrap();
assert_eq!(result, Bytes::from_static(b"second"));
}
#[tokio::test]
async fn allocated_set_accuracy() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let mut awaiters = vec![];
for _ in 0..10 {
awaiters.push(manager.register_outcome().expect("allocate"));
}
assert_eq!(manager.pending_outcome_count(), 10);
awaiters.truncate(5);
assert_eq!(manager.pending_outcome_count(), 5);
for _ in 0..5 {
awaiters.push(manager.register_outcome().expect("allocate"));
}
assert_eq!(manager.pending_outcome_count(), 10);
}
#[tokio::test]
async fn none_payload_handling() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let mut awaiter = manager.register_outcome().expect("allocate slot");
let response_id = awaiter.response_id();
assert!(manager.complete_outcome(response_id, Ok(None)));
let result = awaiter.recv().await.unwrap();
assert!(result.is_none());
}
#[tokio::test]
async fn concurrent_allocation() {
let worker_id = 42;
let manager = Arc::new(ResponseManager::new(worker_id));
let mut handles = vec![];
let allocation_count = 100;
for _ in 0..allocation_count {
let mgr = Arc::clone(&manager);
let handle = tokio::spawn(async move { mgr.register_outcome().expect("allocate") });
handles.push(handle);
}
let mut awaiters = vec![];
for handle in handles {
awaiters.push(handle.await.unwrap());
}
assert_eq!(awaiters.len(), allocation_count);
let mut ids = std::collections::HashSet::new();
for awaiter in &awaiters {
assert!(ids.insert(awaiter.response_id()));
}
assert_eq!(ids.len(), allocation_count);
assert_eq!(manager.pending_outcome_count(), allocation_count);
}
#[tokio::test]
async fn concurrent_completion() {
let worker_id = 42;
let manager = Arc::new(ResponseManager::new(worker_id));
let mut awaiters = vec![];
for _ in 0..50 {
awaiters.push(manager.register_outcome().expect("allocate"));
}
let mut handles = vec![];
for (i, awaiter) in awaiters.iter().enumerate() {
let mgr = Arc::clone(&manager);
let response_id = awaiter.response_id();
let handle = tokio::spawn(async move {
tokio::time::sleep(tokio::time::Duration::from_micros(i as u64 * 10)).await;
mgr.complete_outcome(response_id, Ok(Some(Bytes::from(format!("data-{}", i)))))
});
handles.push(handle);
}
for handle in handles {
assert!(handle.await.unwrap());
}
for (i, mut awaiter) in awaiters.into_iter().enumerate() {
let result = awaiter.recv().await.unwrap().unwrap();
assert_eq!(result, Bytes::from(format!("data-{}", i)));
}
assert_eq!(manager.pending_outcome_count(), 0);
}
#[tokio::test]
async fn race_drop_and_complete() {
let worker_id = 42;
let manager = Arc::new(ResponseManager::new(worker_id));
for iteration in 0..100 {
let awaiter = manager.register_outcome().expect("allocate");
let response_id = awaiter.response_id();
let mgr = Arc::clone(&manager);
let complete_handle = tokio::spawn(async move {
tokio::time::sleep(tokio::time::Duration::from_micros(iteration % 3)).await;
mgr.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"data"))))
});
let drop_handle = tokio::spawn(async move {
tokio::time::sleep(tokio::time::Duration::from_micros((iteration + 1) % 3)).await;
drop(awaiter);
});
let complete_result = complete_handle.await.unwrap();
drop_handle.await.unwrap();
let _ = complete_result;
}
assert_eq!(manager.pending_outcome_count(), 0);
}
#[tokio::test]
async fn encode_decode_boundary_values() {
let max_worker_id = u64::MAX;
let manager = ResponseManager::new(max_worker_id);
let awaiter = manager.register_outcome().expect("allocate");
let response_id = awaiter.response_id();
let raw = response_id.as_u128();
let decoded_worker = (raw & 0xFFFF_FFFF_FFFF_FFFF) as u64;
let decoded_slot = ((raw >> 64) & 0xFFFF) as u16;
assert_eq!(decoded_worker, max_worker_id);
assert_eq!(decoded_slot, 0); }
#[tokio::test]
async fn uuid_round_trip_correctness() {
let worker_id = 0x1234_5678_9ABC_DEF0u64;
let manager = ResponseManager::new(worker_id);
for expected_slot in 0..10 {
let awaiter = manager.register_outcome().expect("allocate");
let response_id = awaiter.response_id();
let raw = response_id.as_u128();
let decoded_worker = (raw & 0xFFFF_FFFF_FFFF_FFFF) as u64;
let decoded_slot = ((raw >> 64) & 0xFFFF) as u16;
assert_eq!(decoded_worker, worker_id);
assert_eq!(decoded_slot as usize, expected_slot);
assert!(manager.complete_outcome(response_id, Ok(None)));
}
}
#[tokio::test]
async fn manager_clone_shares_state() {
let worker_id = 42;
let manager1 = ResponseManager::new(worker_id);
let manager2 = manager1.clone();
let mut awaiter1 = manager1.register_outcome().expect("allocate with manager1");
let response_id1 = awaiter1.response_id();
assert!(manager2.complete_outcome(response_id1, Ok(Some(Bytes::from_static(b"shared")))));
let result = awaiter1.recv().await.unwrap().unwrap();
assert_eq!(result, Bytes::from_static(b"shared"));
assert_eq!(manager1.pending_outcome_count(), 0);
assert_eq!(manager2.pending_outcome_count(), 0);
let mut awaiter2 = manager2.register_outcome().expect("allocate with manager2");
let response_id2 = awaiter2.response_id();
assert!(manager1.complete_outcome(response_id2, Ok(Some(Bytes::from_static(b"reverse")))));
let result = awaiter2.recv().await.unwrap().unwrap();
assert_eq!(result, Bytes::from_static(b"reverse"));
}
#[tokio::test]
async fn generation_mismatch_rejection() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let awaiter = manager.register_outcome().expect("allocate slot");
let response_id = awaiter.response_id();
drop(awaiter);
assert!(!manager.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"stale")))));
}
#[tokio::test]
async fn generation_validation_on_complete() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let mut awaiter = manager.register_outcome().expect("allocate slot");
let response_id = awaiter.response_id();
assert!(manager.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"data")))));
assert!(!manager.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"second")))));
let result = awaiter.recv().await.unwrap().unwrap();
assert_eq!(result, Bytes::from_static(b"data"));
}
#[tokio::test]
async fn did_finish_prevents_double_increment() {
let worker_id = 42;
let manager = ResponseManager::new(worker_id);
let mut awaiter = manager.register_outcome().expect("allocate slot");
let response_id = awaiter.response_id();
assert!(manager.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"data")))));
awaiter.recv().await.unwrap();
let awaiter2 = manager.register_outcome().expect("allocate again");
let response_id2 = awaiter2.response_id();
assert!(!manager.complete_outcome(response_id, Ok(Some(Bytes::from_static(b"stale")))));
assert!(manager.complete_outcome(response_id2, Ok(Some(Bytes::from_static(b"new")))));
}
#[tokio::test]
async fn concurrent_complete_with_mismatched_generations() {
let worker_id = 42;
let manager = Arc::new(ResponseManager::new(worker_id));
let awaiter = manager.register_outcome().expect("allocate");
let response_id = awaiter.response_id();
drop(awaiter);
let mut handles = vec![];
for _ in 0..10 {
let mgr = Arc::clone(&manager);
let rid = response_id;
let handle = tokio::spawn(async move {
mgr.complete_outcome(rid, Ok(Some(Bytes::from_static(b"stale"))))
});
handles.push(handle);
}
for handle in handles {
assert!(!handle.await.unwrap());
}
}
#[test]
fn test_decode_response_header_too_short() {
let short_header = Bytes::from_static(&[1, 2, 3, 4, 5]);
let result = decode_response_header(short_header);
assert!(result.is_err(), "Should error on short header");
match result {
Err(DecodeError::HeaderTooShort(len)) => {
assert_eq!(len, 5);
}
_ => panic!("Expected HeaderTooShort error"),
}
}
#[test]
fn test_decode_response_header_empty() {
let empty_header = Bytes::new();
let result = decode_response_header(empty_header);
assert!(result.is_err(), "Should error on empty header");
match result {
Err(DecodeError::HeaderTooShort(len)) => {
assert_eq!(len, 0);
}
_ => panic!("Expected HeaderTooShort error"),
}
}
#[test]
fn test_decode_response_header_exactly_19_bytes() {
let valid_header = Bytes::from(vec![0u8; 19]);
let result = decode_response_header(valid_header);
assert!(result.is_ok(), "Should succeed with exactly 19 bytes");
let (_response_id, outcome, headers) = result.unwrap();
assert!(matches!(outcome, Outcome::Ok));
assert!(headers.is_none(), "Should have no headers");
}
#[test]
fn test_decode_response_header_more_than_19_bytes() {
let mut data = vec![0u8; 19];
data.extend_from_slice(&[1, 2, 3, 4]); let long_header = Bytes::from(data);
let result = decode_response_header(long_header);
assert!(result.is_ok(), "Should handle extra bytes");
}
#[test]
fn test_decode_response_header_18_bytes() {
let short_header = Bytes::from(vec![0u8; 18]);
let result = decode_response_header(short_header);
assert!(result.is_err(), "Should error with 18 bytes");
match result {
Err(DecodeError::HeaderTooShort(len)) => {
assert_eq!(len, 18);
}
_ => panic!("Expected HeaderTooShort error"),
}
}
#[test]
fn test_decode_response_header_round_trip() {
let response_id = ResponseId::from_u128(0x1234_5678_9ABC_DEF0_1234_5678_9ABC_DEF0);
let encoded = encode_response_header(response_id, Outcome::Ok, None).unwrap();
assert_eq!(
encoded.len(),
19,
"Encoded header should be 19 bytes (16 response_id + 1 outcome + 2 headers_len)"
);
let (decoded_id, decoded_outcome, decoded_headers) =
decode_response_header(encoded).unwrap();
assert_eq!(decoded_id.as_u128(), response_id.as_u128());
assert!(matches!(decoded_outcome, Outcome::Ok));
assert!(decoded_headers.is_none(), "Headers should be None");
}
#[test]
fn test_decode_response_header_error_outcome_round_trip() {
let response_id = ResponseId::from_u128(0xABCD_EF01_2345_6789_ABCD_EF01_2345_6789);
let encoded = encode_response_header(response_id, Outcome::Error, None).unwrap();
let (decoded_id, decoded_outcome, decoded_headers) =
decode_response_header(encoded).unwrap();
assert_eq!(decoded_id.as_u128(), response_id.as_u128());
assert!(matches!(decoded_outcome, Outcome::Error));
assert!(decoded_headers.is_none());
}
#[test]
fn test_response_headers_encode_decode_round_trip() {
let response_id = ResponseId::from_u128(0x1234_5678_9ABC_DEF0_1234_5678_9ABC_DEF0);
let mut headers = HashMap::new();
headers.insert("trace-id".to_string(), "abc123".to_string());
headers.insert("span-id".to_string(), "def456".to_string());
let encoded =
encode_response_header(response_id, Outcome::Ok, Some(headers.clone())).unwrap();
assert!(encoded.len() > 19, "Should be larger with headers");
let (decoded_id, decoded_outcome, decoded_headers) =
decode_response_header(encoded).unwrap();
assert_eq!(decoded_id.as_u128(), response_id.as_u128());
assert!(matches!(decoded_outcome, Outcome::Ok));
assert!(decoded_headers.is_some());
let decoded_headers = decoded_headers.unwrap();
assert_eq!(decoded_headers.len(), 2);
assert_eq!(decoded_headers.get("trace-id").unwrap(), "abc123");
assert_eq!(decoded_headers.get("span-id").unwrap(), "def456");
}
#[test]
fn test_response_headers_empty_map() {
let response_id = ResponseId::from_u128(12345);
let headers = HashMap::new();
let encoded = encode_response_header(response_id, Outcome::Ok, Some(headers)).unwrap();
let (decoded_id, decoded_outcome, decoded_headers) =
decode_response_header(encoded).unwrap();
assert_eq!(decoded_id.as_u128(), response_id.as_u128());
assert!(matches!(decoded_outcome, Outcome::Ok));
assert!(decoded_headers.is_some());
assert_eq!(decoded_headers.unwrap().len(), 0);
}
#[test]
fn test_response_headers_with_unicode() {
let response_id = ResponseId::from_u128(12345);
let mut headers = HashMap::new();
headers.insert("emoji".to_string(), "🚀".to_string());
headers.insert("chinese".to_string(), "ä½ å¥½".to_string());
let encoded =
encode_response_header(response_id, Outcome::Ok, Some(headers.clone())).unwrap();
let (decoded_id, decoded_outcome, decoded_headers) =
decode_response_header(encoded).unwrap();
assert_eq!(decoded_id.as_u128(), response_id.as_u128());
assert!(matches!(decoded_outcome, Outcome::Ok));
let decoded_headers = decoded_headers.unwrap();
assert_eq!(decoded_headers.get("emoji").unwrap(), "🚀");
assert_eq!(decoded_headers.get("chinese").unwrap(), "ä½ å¥½");
}
#[test]
fn test_response_headers_many_entries() {
let response_id = ResponseId::from_u128(12345);
let mut headers = HashMap::new();
for i in 0..50 {
headers.insert(format!("key-{}", i), format!("value-{}", i));
}
let encoded =
encode_response_header(response_id, Outcome::Ok, Some(headers.clone())).unwrap();
let (decoded_id, decoded_outcome, decoded_headers) =
decode_response_header(encoded).unwrap();
assert_eq!(decoded_id.as_u128(), response_id.as_u128());
assert!(matches!(decoded_outcome, Outcome::Ok));
let decoded_headers = decoded_headers.unwrap();
assert_eq!(decoded_headers.len(), 50);
assert_eq!(decoded_headers.get("key-25").unwrap(), "value-25");
}
#[test]
fn test_response_headers_none_vs_empty() {
let response_id = ResponseId::from_u128(12345);
let encoded_none = encode_response_header(response_id, Outcome::Ok, None).unwrap();
let none_len = encoded_none.len();
let (_, outcome_none, headers_none) = decode_response_header(encoded_none).unwrap();
assert!(matches!(outcome_none, Outcome::Ok));
assert!(headers_none.is_none());
let empty_map = HashMap::new();
let encoded_empty =
encode_response_header(response_id, Outcome::Ok, Some(empty_map)).unwrap();
let empty_len = encoded_empty.len();
let (_, outcome_empty, headers_empty) = decode_response_header(encoded_empty).unwrap();
assert!(matches!(outcome_empty, Outcome::Ok));
assert!(headers_empty.is_some());
assert_eq!(headers_empty.unwrap().len(), 0);
assert!(empty_len > none_len, "Empty map should be larger than None");
}
}