use std::{collections::HashMap, sync::Arc};
use tokio::sync::{oneshot, Mutex, OwnedSemaphorePermit, Semaphore};
use crate::DigMessage;
#[derive(Debug)]
pub(crate) struct Request {
sender: oneshot::Sender<DigMessage>,
expected: Vec<u8>,
undeclared: Option<u8>,
_permit: OwnedSemaphorePermit,
}
impl Request {
pub(crate) fn send(self, message: DigMessage) {
self.sender.send(message).ok();
}
pub(crate) fn into_diagnosis(self) -> Option<(Vec<u8>, u8)> {
self.undeclared.map(|found| (self.expected, found))
}
}
pub(crate) enum Correlation {
Unknown,
Answer(Request),
Undeclared,
}
#[derive(Debug)]
pub(crate) struct RequestMap {
items: Mutex<HashMap<u16, Request>>,
capacity: Arc<Semaphore>,
next_id: Mutex<u16>,
}
impl RequestMap {
pub(crate) fn new() -> Self {
Self {
items: Mutex::new(HashMap::new()),
capacity: Arc::new(Semaphore::new(u16::MAX as usize)),
next_id: Mutex::new(0),
}
}
pub(crate) async fn insert(
&self,
sender: oneshot::Sender<DigMessage>,
expected: Vec<u8>,
) -> u16 {
let permit = self
.capacity
.clone()
.acquire_owned()
.await
.expect("request capacity semaphore is never closed");
let mut items = self.items.lock().await;
items.retain(|_, request| !request.sender.is_closed());
let mut next_id = self.next_id.lock().await;
let mut id = *next_id;
loop {
if !items.contains_key(&id) {
break;
}
id = id.wrapping_add(1);
}
*next_id = id.wrapping_add(1);
drop(next_id);
items.insert(
id,
Request {
sender,
expected,
undeclared: None,
_permit: permit,
},
);
id
}
pub(crate) async fn take(&self, id: u16, msg_type: u8) -> Correlation {
let mut items = self.items.lock().await;
let Some(request) = items.get_mut(&id) else {
return Correlation::Unknown;
};
if !request.expected.contains(&msg_type) {
request.undeclared = Some(msg_type);
return Correlation::Undeclared;
}
Correlation::Answer(
items
.remove(&id)
.expect("the entry was observed under the same lock"),
)
}
pub(crate) async fn cancel(&self, id: u16) -> Option<Request> {
self.items.lock().await.remove(&id)
}
}