use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, OnceLock};
use parking_lot::Mutex;
use tokio::sync::mpsc;
use tracing::{debug, warn};
use crate::callback::CallbackId;
use crate::client::direct::unified::UnifiedWriter;
use crate::packet::puback::PubAckPacket;
use crate::packet::publish::PublishPacket;
use crate::packet::pubrec::PubRecPacket;
use crate::packet::Packet;
use crate::protocol::v5::properties::Properties;
use crate::protocol::v5::reason_codes::ReasonCode;
use crate::session::state::AckResolution;
use crate::session::SessionState;
use crate::transport::PacketWriter;
use crate::validation::strip_shared_subscription_prefix;
use crate::QoS;
type WriterHandle = Arc<tokio::sync::Mutex<UnifiedWriter>>;
type WriterSlot = Arc<tokio::sync::Mutex<Option<WriterHandle>>>;
const DROP_REASON: ReasonCode = ReasonCode::UnspecifiedError;
pub(crate) enum AckKind {
Ack,
Reject(ReasonCode),
DropAuto,
}
pub(crate) struct AckRequest {
packet_id: u16,
qos: QoS,
kind: AckKind,
}
pub struct AckToken {
packet_id: u16,
qos: QoS,
armed: bool,
sender: mpsc::UnboundedSender<AckRequest>,
}
impl AckToken {
#[must_use]
pub fn packet_id(&self) -> u16 {
self.packet_id
}
#[must_use]
pub fn qos(&self) -> QoS {
self.qos
}
pub fn ack(mut self) {
self.emit(AckKind::Ack);
}
pub fn reject(mut self, reason: ReasonCode) {
let reason = if reason.is_error() {
reason
} else {
ReasonCode::UnspecifiedError
};
self.emit(AckKind::Reject(reason));
}
fn emit(&mut self, kind: AckKind) {
if !self.armed {
return;
}
self.armed = false;
let _ = self.sender.send(AckRequest {
packet_id: self.packet_id,
qos: self.qos,
kind,
});
}
}
impl Drop for AckToken {
fn drop(&mut self) {
if self.armed {
warn!(
packet_id = self.packet_id,
qos = ?self.qos,
"AckToken dropped without ack/reject; auto-acknowledging with a non-success reason"
);
self.emit(AckKind::DropAuto);
}
}
}
pub(crate) struct AckDispatcher {
tx: mpsc::UnboundedSender<AckRequest>,
writer_slot: WriterSlot,
session: Arc<tokio::sync::RwLock<SessionState>>,
pending_rx: tokio::sync::Mutex<Option<mpsc::UnboundedReceiver<AckRequest>>>,
}
impl AckDispatcher {
pub(crate) fn new(session: Arc<tokio::sync::RwLock<SessionState>>) -> Self {
let (tx, rx) = mpsc::unbounded_channel::<AckRequest>();
Self {
tx,
writer_slot: Arc::new(tokio::sync::Mutex::new(None)),
session,
pending_rx: tokio::sync::Mutex::new(Some(rx)),
}
}
async fn ensure_started(&self) {
let Some(mut rx) = self.pending_rx.lock().await.take() else {
return;
};
let slot = Arc::clone(&self.writer_slot);
let session = Arc::clone(&self.session);
tokio::spawn(async move {
while let Some(request) = rx.recv().await {
Self::handle(&request, &slot, &session).await;
}
});
}
pub(crate) fn token(&self, packet_id: u16, qos: QoS) -> AckToken {
AckToken {
packet_id,
qos,
armed: true,
sender: self.tx.clone(),
}
}
pub(crate) async fn set_writer(&self, writer: WriterHandle) {
self.ensure_started().await;
*self.writer_slot.lock().await = Some(writer);
}
pub(crate) async fn clear_writer(&self) {
*self.writer_slot.lock().await = None;
}
pub(crate) fn enqueue(&self, packet_id: u16, qos: QoS, kind: AckKind) {
let _ = self.tx.send(AckRequest {
packet_id,
qos,
kind,
});
}
async fn handle(
request: &AckRequest,
slot: &WriterSlot,
session: &Arc<tokio::sync::RwLock<SessionState>>,
) {
let reason = match request.kind {
AckKind::Ack => ReasonCode::Success,
AckKind::Reject(r) => r,
AckKind::DropAuto => DROP_REASON,
};
let packet = match request.qos {
QoS::AtMostOnce => return,
QoS::AtLeastOnce => Packet::PubAck(PubAckPacket {
packet_id: request.packet_id,
reason_code: reason,
properties: Properties::default(),
}),
QoS::ExactlyOnce => Packet::PubRec(PubRecPacket {
packet_id: request.packet_id,
reason_code: reason,
properties: Properties::default(),
}),
};
let is_success = reason == ReasonCode::Success;
{
let session = session.read().await;
match request.qos {
QoS::AtMostOnce => {}
QoS::ExactlyOnce if is_success => {
session.mark_pubrec_sent(request.packet_id).await;
session
.set_resolution(request.packet_id, AckResolution::Acked)
.await;
}
QoS::AtLeastOnce | QoS::ExactlyOnce => {
session.acknowledge_inbound(request.packet_id).await;
session.clear_inbound_state(request.packet_id).await;
}
}
}
let writer = slot.lock().await.clone();
let written = match &writer {
Some(handle) => handle.lock().await.write_packet(packet).await.is_ok(),
None => false,
};
if !written {
debug!(
packet_id = request.packet_id,
"Deferred ack not written (disconnected); resolution recorded for replay"
);
}
}
}
pub(crate) type AckPublishCallback = Arc<dyn Fn(PublishPacket, AckToken) + Send + Sync>;
struct AckCallbackEntry {
callback: AckPublishCallback,
topic_filter: String,
}
struct AckDispatchItem {
callback: AckPublishCallback,
message: PublishPacket,
token: AckToken,
}
pub(crate) struct AckCallbackManager {
exact: Mutex<HashMap<String, AckCallbackEntry>>,
wildcard: Mutex<Vec<AckCallbackEntry>>,
next_id: AtomicU64,
dispatch_tx: OnceLock<mpsc::UnboundedSender<AckDispatchItem>>,
}
impl AckCallbackManager {
pub(crate) fn new() -> Self {
Self {
exact: Mutex::new(HashMap::new()),
wildcard: Mutex::new(Vec::new()),
next_id: AtomicU64::new(1),
dispatch_tx: OnceLock::new(),
}
}
fn dispatch_sender(&self) -> &mpsc::UnboundedSender<AckDispatchItem> {
self.dispatch_tx.get_or_init(|| {
let (tx, mut rx) = mpsc::unbounded_channel::<AckDispatchItem>();
tokio::spawn(async move {
while let Some(item) = rx.recv().await {
(item.callback)(item.message, item.token);
}
});
tx
})
}
pub(crate) fn register(&self, topic_filter: &str, callback: AckPublishCallback) -> CallbackId {
let id = self.next_id.fetch_add(1, Ordering::SeqCst);
let entry = AckCallbackEntry {
callback,
topic_filter: topic_filter.to_string(),
};
let actual = strip_shared_subscription_prefix(topic_filter).to_string();
if actual.contains('+') || actual.contains('#') {
self.wildcard.lock().push(entry);
} else {
self.exact.lock().insert(actual, entry);
}
id
}
pub(crate) fn unregister(&self, topic_filter: &str) -> bool {
let actual = strip_shared_subscription_prefix(topic_filter);
let removed_exact = self.exact.lock().remove(actual).is_some();
let mut wildcard = self.wildcard.lock();
let before = wildcard.len();
wildcard.retain(|e| e.topic_filter != topic_filter);
removed_exact || wildcard.len() < before
}
pub(crate) fn find_one(&self, topic: &str) -> Option<AckPublishCallback> {
if let Some(entry) = self.exact.lock().get(topic) {
return Some(Arc::clone(&entry.callback));
}
let wildcard = self.wildcard.lock();
for entry in wildcard.iter() {
let filter = strip_shared_subscription_prefix(&entry.topic_filter);
if crate::topic_matching::matches(topic, filter) {
return Some(Arc::clone(&entry.callback));
}
}
None
}
pub(crate) fn dispatch(
&self,
callback: AckPublishCallback,
message: PublishPacket,
token: AckToken,
) {
let _ = self.dispatch_sender().send(AckDispatchItem {
callback,
message,
token,
});
}
}
#[cfg(test)]
mod tests {
use super::{AckCallbackManager, AckDispatcher, AckPublishCallback};
use crate::packet::publish::PublishPacket;
use crate::session::state::{SessionConfig, SessionState};
use crate::QoS;
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::Arc;
#[tokio::test]
async fn duplicate_wildcard_subscription_dispatches_to_a_single_callback() {
let mgr = AckCallbackManager::new();
let hits_first = Arc::new(AtomicU32::new(0));
let hits_second = Arc::new(AtomicU32::new(0));
let f = Arc::clone(&hits_first);
let s = Arc::clone(&hits_second);
let cb_first: AckPublishCallback = Arc::new(move |_p, _t| {
f.fetch_add(1, Ordering::SeqCst);
});
let cb_second: AckPublishCallback = Arc::new(move |_p, _t| {
s.fetch_add(1, Ordering::SeqCst);
});
mgr.register("jobs/#", cb_first);
mgr.register("jobs/#", cb_second);
let dispatcher = AckDispatcher::new(Arc::new(tokio::sync::RwLock::new(SessionState::new(
"t".to_string(),
SessionConfig::default(),
true,
))));
let callback = mgr
.find_one("jobs/build")
.expect("a wildcard entry matches jobs/build");
let token = dispatcher.token(1, QoS::ExactlyOnce);
callback(
PublishPacket::new("jobs/build", b"x".to_vec(), QoS::ExactlyOnce),
token,
);
assert_eq!(
hits_first.load(Ordering::SeqCst) + hits_second.load(Ordering::SeqCst),
1,
"a matching publish invokes exactly one ack callback, never both"
);
assert_eq!(
hits_second.load(Ordering::SeqCst),
0,
"the duplicate registration is shadowed and never fires"
);
}
}