use std::collections::{HashMap, HashSet, VecDeque};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Weak};
use tokio::sync::{Mutex, Notify};
use crate::as4::As4TopologyCoordination;
use crate::as4::coordination::{
ConversationGuardHandle, ConversationOrderGate, ConversationTurnHandle,
};
use crate::core::{AsxError, ErrorCode, ErrorContext, Result};
pub struct As4ConversationOrderGate {
slots: Mutex<HashMap<Arc<str>, Weak<ConversationTurnSlot>>>,
completed_messages: Mutex<HashMap<Arc<str>, CompletedMessageTracker>>,
capacity: usize,
}
impl std::fmt::Debug for As4ConversationOrderGate {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("As4ConversationOrderGate")
.field("capacity", &self.capacity)
.finish_non_exhaustive()
}
}
pub struct ConversationGuard {
_turn: ConversationTurnGuard,
}
pub struct ConversationTurnReservation {
slot: Arc<ConversationTurnSlot>,
ticket: u64,
advance_on_drop: bool,
}
pub struct ConversationTurnGuard {
slot: Arc<ConversationTurnSlot>,
}
struct ConversationTurnState {
waiters: HashMap<u64, Arc<Notify>>,
abandoned: HashSet<u64>,
}
#[derive(Default)]
struct CompletedMessageTracker {
ids: HashSet<Arc<str>>,
order: VecDeque<Arc<str>>,
}
struct ConversationTurnSlot {
next_ticket: AtomicU64,
serving_ticket: AtomicU64,
state: std::sync::Mutex<ConversationTurnState>,
}
impl Default for ConversationTurnSlot {
fn default() -> Self {
Self {
next_ticket: AtomicU64::new(0),
serving_ticket: AtomicU64::new(0),
state: std::sync::Mutex::new(ConversationTurnState {
waiters: HashMap::new(),
abandoned: HashSet::new(),
}),
}
}
}
const MAX_TRACKED_COMPLETED_MESSAGE_IDS_PER_CONVERSATION: usize = 4096;
impl std::fmt::Debug for ConversationGuard {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ConversationGuard").finish_non_exhaustive()
}
}
impl std::fmt::Debug for ConversationTurnReservation {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ConversationTurnReservation")
.field("ticket", &self.ticket)
.finish_non_exhaustive()
}
}
impl std::fmt::Debug for ConversationTurnGuard {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ConversationTurnGuard")
.finish_non_exhaustive()
}
}
impl As4ConversationOrderGate {
pub fn new(capacity: usize) -> Self {
assert!(
capacity > 0,
"As4ConversationOrderGate capacity must be > 0"
);
Self {
slots: Mutex::new(HashMap::with_capacity(capacity.min(512))),
completed_messages: Mutex::new(HashMap::with_capacity(capacity.min(512))),
capacity,
}
}
#[inline]
pub fn cluster_safe(&self) -> bool {
false
}
pub(crate) async fn enforce_reply_predecessor_and_record(
&self,
conversation_id: &str,
message_id: &str,
ref_to_message_id: Option<&str>,
) -> Result<()> {
let mut completed = self.completed_messages.lock().await;
let tracker = completed.entry(Arc::from(conversation_id)).or_default();
if let Some(ref_to) = ref_to_message_id
&& !tracker.ids.contains(ref_to)
{
return Err(AsxError::new(
ErrorCode::ReliabilityFailure,
format!(
"ordered AS4 Two-Way response references unknown predecessor MessageId '{}'; expected predecessor completion on this replica before response processing",
ref_to
),
ErrorContext::new("as4_receive_push_ordered"),
));
}
let message_id: Arc<str> = Arc::from(message_id);
if tracker.ids.insert(Arc::clone(&message_id)) {
tracker.order.push_back(message_id);
while tracker.order.len() > MAX_TRACKED_COMPLETED_MESSAGE_IDS_PER_CONVERSATION {
if let Some(expired) = tracker.order.pop_front() {
tracker.ids.remove(&expired);
}
}
}
Ok(())
}
pub async fn acquire(&self, conversation_id: &str) -> Result<ConversationGuard> {
let turn = self.reserve_turn(conversation_id).await?;
let turn_guard = turn.wait_turn().await?;
Ok(ConversationGuard { _turn: turn_guard })
}
pub async fn reserve_turn(&self, conversation_id: &str) -> Result<ConversationTurnReservation> {
let slot = self.lookup_or_create_turn_slot(conversation_id).await?;
let ticket = slot.next_ticket.fetch_add(1, Ordering::AcqRel);
if ticket == u64::MAX {
return Err(AsxError::new(
ErrorCode::PolicyViolation,
"as4 conversation order gate ticket counter overflow",
ErrorContext::new("as4_conversation_order_gate"),
));
}
Ok(ConversationTurnReservation {
slot,
ticket,
advance_on_drop: true,
})
}
async fn lookup_or_create_turn_slot(
&self,
conversation_id: &str,
) -> Result<Arc<ConversationTurnSlot>> {
let mut map = self.slots.lock().await;
if let Some(slot_ref) = map.get_mut(conversation_id) {
if let Some(slot) = slot_ref.upgrade() {
return Ok(slot);
}
let new_slot = Arc::new(ConversationTurnSlot::default());
*slot_ref = Arc::downgrade(&new_slot);
return Ok(new_slot);
}
if map.len() >= self.capacity {
map.retain(|_, w| w.strong_count() > 0);
let active_keys: HashSet<Arc<str>> = map.keys().cloned().collect();
let mut completed = self.completed_messages.lock().await;
completed.retain(|k, _| active_keys.contains(k));
if map.len() >= self.capacity {
return Err(AsxError::new(
ErrorCode::CapacityExhausted,
format!(
"as4 conversation order gate capacity exhausted \
({} active conversations)",
self.capacity
),
ErrorContext::new("as4_conversation_order_gate"),
));
}
}
let new_slot = Arc::new(ConversationTurnSlot::default());
map.insert(Arc::from(conversation_id), Arc::downgrade(&new_slot));
Ok(new_slot)
}
#[inline]
pub fn capacity(&self) -> usize {
self.capacity
}
}
impl As4TopologyCoordination for As4ConversationOrderGate {
fn cluster_safe(&self) -> bool {
self.cluster_safe()
}
fn topology_component(&self) -> &'static str {
"conversation-order-gate"
}
}
struct InProcessConversationGuardHandle(Option<ConversationTurnGuard>);
impl ConversationGuardHandle for InProcessConversationGuardHandle {
fn release(mut self: Box<Self>) {
drop(self.0.take());
}
}
#[derive(Debug)]
struct InProcessTurnHandle(ConversationTurnReservation);
impl ConversationTurnHandle for InProcessTurnHandle {
fn wait_for_turn(
self: Box<Self>,
) -> std::pin::Pin<
Box<
dyn std::future::Future<Output = crate::core::Result<Box<dyn ConversationGuardHandle>>>
+ Send,
>,
> {
Box::pin(async move {
let guard = self.0.wait_turn().await?;
Ok(Box::new(InProcessConversationGuardHandle(Some(guard)))
as Box<dyn ConversationGuardHandle>)
})
}
}
impl Drop for InProcessConversationGuardHandle {
fn drop(&mut self) {
let _ = self.0.take();
}
}
#[allow(clippy::manual_async_fn)]
impl ConversationOrderGate for As4ConversationOrderGate {
fn reserve_ordered_turn<'a>(
&'a self,
conversation_id: &'a str,
_session: &'a crate::core::SessionContext,
) -> std::pin::Pin<
Box<
dyn std::future::Future<Output = crate::core::Result<Box<dyn ConversationTurnHandle>>>
+ Send
+ 'a,
>,
> {
Box::pin(async move {
let reservation = self.reserve_turn(conversation_id).await?;
Ok(Box::new(InProcessTurnHandle(reservation)) as Box<dyn ConversationTurnHandle>)
})
}
fn record_message_ordering<'a>(
&'a self,
conversation_id: &'a str,
message_id: &'a str,
ref_to_message_id: Option<&'a str>,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = crate::core::Result<()>> + Send + 'a>>
{
Box::pin(async move {
self.enforce_reply_predecessor_and_record(
conversation_id,
message_id,
ref_to_message_id,
)
.await
})
}
}
impl ConversationTurnReservation {
pub async fn wait_turn(mut self) -> Result<ConversationTurnGuard> {
loop {
if self.ticket == self.slot.serving_ticket.load(Ordering::Acquire) {
self.advance_on_drop = false;
return Ok(ConversationTurnGuard {
slot: self.slot.clone(),
});
}
let maybe_waiter = {
let mut state = self
.slot
.state
.lock()
.expect("conversation gate state lock");
if self.ticket == self.slot.serving_ticket.load(Ordering::Acquire) {
continue;
}
state
.waiters
.entry(self.ticket)
.or_insert_with(|| Arc::new(Notify::new()))
.clone()
};
maybe_waiter.notified().await;
}
}
}
impl Drop for ConversationTurnReservation {
fn drop(&mut self) {
if !self.advance_on_drop {
return;
}
{
let mut state = self
.slot
.state
.lock()
.expect("conversation gate state lock");
state.waiters.remove(&self.ticket);
state.abandoned.insert(self.ticket);
}
self.slot.drain_abandoned_head();
}
}
impl Drop for ConversationTurnGuard {
fn drop(&mut self) {
let next_ticket = self.slot.serving_ticket.fetch_add(1, Ordering::AcqRel) + 1;
let maybe_waiter = {
let mut state = self
.slot
.state
.lock()
.expect("conversation gate state lock");
state.waiters.remove(&next_ticket)
};
if let Some(waiter) = maybe_waiter {
waiter.notify_one();
}
self.slot.drain_abandoned_head();
}
}
impl ConversationTurnSlot {
fn drain_abandoned_head(&self) {
const MAX_DRAIN_ITERS: u32 = 64;
for _ in 0..MAX_DRAIN_ITERS {
let serving = self.serving_ticket.load(Ordering::Acquire);
let is_abandoned = {
let state = self.state.lock().expect("conversation gate state lock");
state.abandoned.contains(&serving)
};
if !is_abandoned {
break;
}
if self
.serving_ticket
.compare_exchange(serving, serving + 1, Ordering::AcqRel, Ordering::Acquire)
.is_err()
{
continue;
}
let maybe_waiter = {
let mut state = self.state.lock().expect("conversation gate state lock");
state.abandoned.remove(&serving);
state.waiters.remove(&(serving + 1))
};
if let Some(waiter) = maybe_waiter {
waiter.notify_one();
}
}
}
}
#[cfg(test)]
mod tests {
use super::As4ConversationOrderGate;
use crate::core::ErrorCode;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::oneshot;
use tokio::sync::{Barrier, Notify};
use tokio::time::{Duration, timeout};
#[tokio::test]
async fn same_conversation_is_serialized() {
let gate = Arc::new(As4ConversationOrderGate::new(8));
let guard = gate.acquire("conv-1").await.expect("first acquire");
let (acquired_tx, acquired_rx) = oneshot::channel::<()>();
let gate_ref = Arc::clone(&gate);
let waiter = tokio::spawn(async move {
let _g = gate_ref.acquire("conv-1").await.expect("second acquire");
let _ = acquired_tx.send(());
});
assert!(
timeout(Duration::from_millis(30), acquired_rx)
.await
.is_err()
);
drop(guard);
waiter.await.expect("waiter task join");
}
#[tokio::test]
async fn different_conversations_proceed_in_parallel() {
let gate = As4ConversationOrderGate::new(8);
let _guard_a = gate.acquire("conv-a").await.expect("acquire conv-a");
let guard_b = timeout(Duration::from_millis(50), gate.acquire("conv-b"))
.await
.expect("conv-b acquire timeout")
.expect("acquire conv-b");
drop(guard_b);
}
#[tokio::test]
async fn capacity_is_never_overshot_under_contention() {
const CAPACITY: usize = 8;
const CONTENDERS: usize = 64;
let gate = Arc::new(As4ConversationOrderGate::new(CAPACITY));
let start = Arc::new(Barrier::new(CONTENDERS + 1));
let release = Arc::new(Notify::new());
let success = Arc::new(AtomicUsize::new(0));
let capacity_exhausted = Arc::new(AtomicUsize::new(0));
let mut tasks = Vec::with_capacity(CONTENDERS);
for idx in 0..CONTENDERS {
let gate_ref = Arc::clone(&gate);
let start_ref = Arc::clone(&start);
let release_ref = Arc::clone(&release);
let success_ref = Arc::clone(&success);
let exhausted_ref = Arc::clone(&capacity_exhausted);
tasks.push(tokio::spawn(async move {
start_ref.wait().await;
match gate_ref.acquire(&format!("conv-{idx}")).await {
Ok(_guard) => {
success_ref.fetch_add(1, Ordering::SeqCst);
release_ref.notified().await;
}
Err(err) if err.code == ErrorCode::CapacityExhausted => {
exhausted_ref.fetch_add(1, Ordering::SeqCst);
}
Err(err) => panic!("unexpected error: {err}"),
}
}));
}
start.wait().await;
tokio::time::sleep(Duration::from_millis(40)).await;
assert_eq!(success.load(Ordering::SeqCst), CAPACITY);
assert_eq!(
capacity_exhausted.load(Ordering::SeqCst),
CONTENDERS - CAPACITY
);
release.notify_waiters();
for task in tasks {
task.await.expect("join contender");
}
}
#[tokio::test]
async fn dropping_reserved_turn_skips_ticket_for_following_waiters() {
let gate = As4ConversationOrderGate::new(8);
let first = gate.reserve_turn("conv-drop").await.expect("reserve first");
let dropped = gate
.reserve_turn("conv-drop")
.await
.expect("reserve dropped");
let third = gate.reserve_turn("conv-drop").await.expect("reserve third");
let first_guard = first.wait_turn().await.expect("first guard");
drop(dropped);
drop(first_guard);
let third_guard = timeout(Duration::from_millis(100), third.wait_turn())
.await
.expect("third wait timeout")
.expect("third guard");
drop(third_guard);
}
#[tokio::test]
async fn cancelled_wait_turn_does_not_deadlock_later_ticket() {
let gate = Arc::new(As4ConversationOrderGate::new(8));
let first = gate
.reserve_turn("conv-cancel")
.await
.expect("reserve first");
let cancelled = gate
.reserve_turn("conv-cancel")
.await
.expect("reserve cancelled");
let third = gate
.reserve_turn("conv-cancel")
.await
.expect("reserve third");
let first_guard = first.wait_turn().await.expect("first guard");
let cancelled_task = tokio::spawn(async move {
let _ = cancelled.wait_turn().await;
});
tokio::time::sleep(Duration::from_millis(10)).await;
cancelled_task.abort();
let _ = cancelled_task.await;
drop(first_guard);
let third_guard = timeout(Duration::from_millis(100), third.wait_turn())
.await
.expect("third wait timeout")
.expect("third guard");
drop(third_guard);
}
}