use std::collections::BTreeMap;
use std::sync::{Arc, Mutex, RwLock};
use std::time::{Duration, Instant};
use crossbeam::channel::{
Receiver, RecvTimeoutError, SendError, SendTimeoutError, Sender, TryRecvError, bounded,
};
use deimos_shared::peripherals::PeripheralId;
const CHANNEL_CAPACITY: usize = 10;
const SEND_POLL_INTERVAL: Duration = Duration::from_millis(2);
type EndpointInner = (Sender<Vec<u8>>, Receiver<Vec<u8>>);
#[derive(Clone, Debug, Default)]
pub struct SocketChannels {
channels: Arc<RwLock<BTreeMap<PeripheralId, Arc<SocketChannel>>>>,
}
impl SocketChannels {
pub(crate) fn claim_controller(&self, id: PeripheralId) -> Result<SocketEndpoint, String> {
self.channel(id)?.claim(id, Side::Controller)
}
pub(crate) fn claim_peripheral(&self, id: PeripheralId) -> Result<SocketEndpoint, String> {
self.channel(id)?.claim(id, Side::Peripheral)
}
fn channel(&self, id: PeripheralId) -> Result<Arc<SocketChannel>, String> {
let mut channels = self
.channels
.write()
.map_err(|_| "ThreadChannelSocket channel registry is poisoned".to_owned())?;
Ok(channels
.entry(id)
.or_insert_with(|| Arc::new(SocketChannel::default()))
.clone())
}
}
#[derive(Clone, Copy, Debug)]
enum Side {
Controller,
Peripheral,
}
#[derive(Debug)]
struct ChannelEnds {
controller: Option<EndpointInner>,
peripheral: Option<EndpointInner>,
}
impl ChannelEnds {
fn new() -> Self {
let (controller_tx, peripheral_rx) = bounded(CHANNEL_CAPACITY);
let (peripheral_tx, controller_rx) = bounded(CHANNEL_CAPACITY);
Self {
controller: Some((controller_tx, controller_rx)),
peripheral: Some((peripheral_tx, peripheral_rx)),
}
}
fn slot(&mut self, side: Side) -> &mut Option<EndpointInner> {
match side {
Side::Controller => &mut self.controller,
Side::Peripheral => &mut self.peripheral,
}
}
fn is_active(&self, side: Side) -> bool {
match side {
Side::Controller => self.controller.is_none(),
Side::Peripheral => self.peripheral.is_none(),
}
}
}
impl Side {
fn opposite(self) -> Self {
match self {
Self::Controller => Self::Peripheral,
Self::Peripheral => Self::Controller,
}
}
}
#[derive(Debug)]
struct SocketChannel {
ends: Mutex<ChannelEnds>,
}
impl Default for SocketChannel {
fn default() -> Self {
Self {
ends: Mutex::new(ChannelEnds::new()),
}
}
}
impl SocketChannel {
fn claim(self: &Arc<Self>, id: PeripheralId, side: Side) -> Result<SocketEndpoint, String> {
let mut ends = self
.ends
.lock()
.map_err(|_| format!("ThreadChannelSocket channel {id:?} is poisoned"))?;
let inner = ends
.slot(side)
.take()
.ok_or_else(|| format!("ThreadChannelSocket {side:?} endpoint for {id:?} is active"))?;
Ok(SocketEndpoint {
side,
channel: self.clone(),
inner: Some(inner),
})
}
}
#[derive(Debug)]
pub struct SocketEndpoint {
side: Side,
channel: Arc<SocketChannel>,
inner: Option<EndpointInner>,
}
impl SocketEndpoint {
pub fn send(&self, packet: Vec<u8>) -> Result<(), SendError<Vec<u8>>> {
let mut packet = packet;
loop {
if !self.peer_is_active() {
return Err(SendError(packet));
}
match self
.inner
.as_ref()
.unwrap()
.0
.send_timeout(packet, SEND_POLL_INTERVAL)
{
Ok(()) => return Ok(()),
Err(SendTimeoutError::Timeout(unsent)) => packet = unsent,
Err(SendTimeoutError::Disconnected(unsent)) => return Err(SendError(unsent)),
}
}
}
pub fn send_timeout(
&self,
packet: Vec<u8>,
timeout: Duration,
) -> Result<(), SendTimeoutError<Vec<u8>>> {
let deadline = Instant::now().checked_add(timeout);
let mut packet = packet;
loop {
if !self.peer_is_active() {
return Err(SendTimeoutError::Disconnected(packet));
}
let wait = deadline.map_or(SEND_POLL_INTERVAL, |deadline| {
deadline
.saturating_duration_since(Instant::now())
.min(SEND_POLL_INTERVAL)
});
match self.inner.as_ref().unwrap().0.send_timeout(packet, wait) {
Ok(()) => return Ok(()),
Err(SendTimeoutError::Timeout(unsent)) => {
packet = unsent;
if deadline.is_some_and(|deadline| Instant::now() >= deadline) {
return Err(SendTimeoutError::Timeout(packet));
}
}
Err(SendTimeoutError::Disconnected(unsent)) => {
return Err(SendTimeoutError::Disconnected(unsent));
}
}
}
}
pub fn try_recv(&self) -> Result<Vec<u8>, TryRecvError> {
self.inner.as_ref().unwrap().1.try_recv()
}
pub fn recv_timeout(&self, timeout: Duration) -> Result<Vec<u8>, RecvTimeoutError> {
self.inner.as_ref().unwrap().1.recv_timeout(timeout)
}
fn peer_is_active(&self) -> bool {
self.channel
.ends
.lock()
.is_ok_and(|ends| ends.is_active(self.side.opposite()))
}
}
impl Drop for SocketEndpoint {
fn drop(&mut self) {
let Some(inner) = self.inner.take() else {
return;
};
let Ok(mut ends) = self.channel.ends.lock() else {
return;
};
let slot = ends.slot(self.side);
debug_assert!(slot.is_none());
if slot.is_none() {
*slot = Some(inner);
}
if ends.controller.is_some() && ends.peripheral.is_some() {
*ends = ChannelEnds::new();
}
}
}
pub(crate) fn socket_channels_default() -> SocketChannels {
SocketChannels::default()
}
#[cfg(test)]
mod tests {
use super::*;
use std::thread;
const ID: PeripheralId = PeripheralId {
model_number: 1,
serial_number: 2,
};
#[test]
fn endpoint_drop_releases_the_role_and_clears_completed_connection_packets() {
let channels = SocketChannels::default();
let controller = channels.claim_controller(ID).unwrap();
let peripheral = channels.claim_peripheral(ID).unwrap();
assert!(channels.claim_peripheral(ID).is_err());
controller.send(vec![1]).unwrap();
assert_eq!(peripheral.recv_timeout(Duration::ZERO).unwrap(), vec![1]);
peripheral.send(vec![2]).unwrap();
drop(controller);
drop(peripheral);
let controller = channels.claim_controller(ID).unwrap();
let peripheral = channels.claim_peripheral(ID).unwrap();
assert!(controller.try_recv().is_err());
controller.send(vec![3]).unwrap();
assert_eq!(peripheral.recv_timeout(Duration::ZERO).unwrap(), vec![3]);
}
#[test]
fn send_fails_when_the_peer_is_inactive() {
let channels = SocketChannels::default();
let controller = channels.claim_controller(ID).unwrap();
assert!(controller.send(vec![1]).is_err());
let peripheral = channels.claim_peripheral(ID).unwrap();
drop(peripheral);
assert!(matches!(
controller.send_timeout(vec![2], Duration::from_secs(1)),
Err(SendTimeoutError::Disconnected(_))
));
}
#[test]
fn dropping_a_peer_unblocks_a_full_channel_send() {
let channels = SocketChannels::default();
let controller = channels.claim_controller(ID).unwrap();
let peripheral = channels.claim_peripheral(ID).unwrap();
for packet in 0..CHANNEL_CAPACITY {
controller.send(vec![packet as u8]).unwrap();
}
let blocked = thread::spawn(move || controller.send(vec![255]));
thread::sleep(Duration::from_millis(10));
drop(peripheral);
assert!(blocked.join().unwrap().is_err());
}
}