use thiserror::Error;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use super::{Frame, FramedRecv, FramedSend};
use crate::{zakura::ZakuraPeerId, BoxError};
#[derive(Debug)]
pub struct PeerStreamSession {
peer_id: ZakuraPeerId,
stream_kind: u16,
recv: FramedRecv,
send: FramedSend,
cancel_token: CancellationToken,
}
impl PeerStreamSession {
pub fn new(
peer_id: ZakuraPeerId,
stream_kind: u16,
recv: FramedRecv,
send: FramedSend,
cancel_token: CancellationToken,
) -> Self {
Self {
peer_id,
stream_kind,
recv,
send,
cancel_token,
}
}
pub fn peer_id(&self) -> &ZakuraPeerId {
&self.peer_id
}
pub fn stream_kind(&self) -> u16 {
self.stream_kind
}
pub fn cancel_token(&self) -> CancellationToken {
self.cancel_token.clone()
}
pub fn sender(&self) -> FramedSend {
self.send.clone()
}
pub fn try_send_frame(&self, frame: Frame) -> Result<(), OrderedSendError> {
try_send_frame(&self.send, frame)
}
pub fn try_send_encoded(
&self,
encode: impl FnOnce() -> Result<Frame, BoxError>,
) -> Result<(), OrderedSendError> {
let frame = encode().map_err(OrderedSendError::Encode)?;
self.try_send_frame(frame)
}
pub fn into_parts(self) -> (ZakuraPeerId, u16, FramedRecv, FramedSend, CancellationToken) {
(
self.peer_id,
self.stream_kind,
self.recv,
self.send,
self.cancel_token,
)
}
}
#[derive(Debug, Error)]
pub enum OrderedSendError {
#[error("ordered stream send queue is full")]
Full,
#[error("ordered stream send queue is closed")]
Closed,
#[error("failed to encode ordered stream frame: {0}")]
Encode(#[source] BoxError),
}
fn try_send_frame(send: &FramedSend, frame: Frame) -> Result<(), OrderedSendError> {
match send.try_send(frame) {
Ok(()) => Ok(()),
Err(mpsc::error::TrySendError::Full(_frame)) => Err(OrderedSendError::Full),
Err(mpsc::error::TrySendError::Closed(_frame)) => Err(OrderedSendError::Closed),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::zakura::framed_channel;
fn frame(message_type: u16) -> Frame {
Frame {
message_type,
flags: 0,
payload: Vec::new(),
}
}
#[test]
fn try_send_succeeds_with_capacity() {
let (send, _recv) = framed_channel(1);
assert!(try_send_frame(&send, frame(1)).is_ok());
}
#[test]
fn try_send_returns_full_without_waiting() {
let (send, _recv) = framed_channel(1);
try_send_frame(&send, frame(1)).expect("first send has capacity");
assert!(matches!(
try_send_frame(&send, frame(2)),
Err(OrderedSendError::Full)
));
}
#[test]
fn try_send_returns_closed_when_worker_is_gone() {
let (send, recv) = framed_channel(1);
drop(recv);
assert!(matches!(
try_send_frame(&send, frame(1)),
Err(OrderedSendError::Closed)
));
}
}