use std::{collections::HashMap, sync::Mutex};
use bytes::Bytes;
use tokio::{sync::mpsc, task::AbortHandle};
pub type CallId = u64;
struct Call {
body: Option<mpsc::Sender<Bytes>>,
pump: Option<AbortHandle>,
cancelled: bool,
}
pub(crate) struct Registration<'a> {
registry: &'a CallRegistry,
id: CallId,
}
impl Registration<'_> {
pub(crate) fn attach(self, pump: AbortHandle) {
let mut calls = self.registry.lock();
let Some(call) = calls.get_mut(&self.id) else {
return;
};
if call.cancelled {
calls.remove(&self.id);
drop(calls);
pump.abort();
} else {
call.pump = Some(pump);
}
}
}
#[derive(Default)]
pub(crate) struct CallRegistry {
calls: Mutex<HashMap<CallId, Call>>,
}
impl CallRegistry {
pub(crate) fn begin(&self, id: CallId, body: Option<mpsc::Sender<Bytes>>) -> Registration<'_> {
self.lock().insert(
id,
Call {
body,
pump: None,
cancelled: false,
},
);
Registration { registry: self, id }
}
pub(crate) fn body_sender(&self, id: CallId) -> Option<mpsc::Sender<Bytes>> {
self.lock().get(&id).and_then(|c| c.body.clone())
}
pub(crate) fn close_request_body(&self, id: CallId) {
if let Some(call) = self.lock().get_mut(&id) {
call.body = None;
}
}
pub(crate) fn remove(&self, id: CallId) {
let mut calls = self.lock();
let Some(call) = calls.get_mut(&id) else {
return;
};
match call.pump.take() {
Some(pump) => {
calls.remove(&id);
drop(calls);
pump.abort();
}
None => {
call.cancelled = true;
call.body = None;
}
}
}
pub(crate) fn forget(&self, id: CallId) {
self.lock().remove(&id);
}
#[cfg(feature = "testing")]
pub(crate) fn len(&self) -> usize {
self.lock().len()
}
fn lock(&self) -> std::sync::MutexGuard<'_, HashMap<CallId, Call>> {
self.calls.lock().unwrap_or_else(|e| e.into_inner())
}
}