use crate::adapters::outbound::SEND_QUEUE_WAIT;
use crate::transport::request_registry::{MarkResult, RequestRegistry};
use crate::{event::TransportEvent, packet::Packet, transport::transport::Transport, SessionId};
use async_trait::async_trait;
use flume::{bounded, Receiver, Sender};
use std::sync::Arc;
#[derive(Debug)]
pub enum ActorMessage {
InboundEvent(TransportEvent),
Send(Packet),
SendWithReply {
packet: Packet,
reply: tokio::sync::oneshot::Sender<Result<(), crate::TransportError>>,
},
Close,
}
#[async_trait]
pub trait SessionHandler: Send + Sync + 'static {
async fn on_message(&self, session_id: SessionId, packet: Packet, sender: SessionSender);
async fn on_connected(&self, session_id: SessionId) {
let _ = session_id; }
async fn on_disconnected(&self, session_id: SessionId, reason: crate::error::CloseReason) {
let _ = (session_id, reason); }
async fn on_error(&self, session_id: SessionId, error: crate::TransportError) {
let _ = (session_id, error); }
}
#[derive(Clone)]
pub struct SessionSender {
session_id: SessionId,
transport: Arc<Transport>,
inbound_registry: Option<Arc<RequestRegistry>>,
}
impl SessionSender {
pub(crate) fn new(
session_id: SessionId,
transport: Arc<Transport>,
inbound_registry: Option<Arc<RequestRegistry>>,
) -> Self {
Self {
session_id,
transport,
inbound_registry,
}
}
pub async fn send(&self, packet: Packet) -> Result<(), crate::TransportError> {
self.transport.send(packet).await
}
pub async fn send_data(&self, data: Vec<u8>) -> Result<(), crate::TransportError> {
let packet = Packet::one_way(0, data);
self.transport.send(packet).await
}
pub async fn respond(
&self,
message_id: u32,
biz_type: u8,
data: Vec<u8>,
) -> Result<(), crate::TransportError> {
if let Some(registry) = &self.inbound_registry {
if registry.mark_responded(Some(self.session_id), message_id) != MarkResult::Updated {
tracing::debug!(
"[ACTOR] Skip duplicate/late/unknown response: session={}, id={}",
self.session_id,
message_id
);
return Ok(());
}
}
let response_packet = Packet {
header: crate::packet::FixedHeader {
version: 1,
compression: crate::packet::CompressionType::None,
packet_type: crate::packet::PacketType::Response,
biz_type,
message_id,
ext_header_len: 0,
payload_len: data.len() as u32,
reserved: crate::packet::ReservedFlags::new(),
},
ext_header: Vec::new(),
payload: data,
};
let result = self.transport.send(response_packet).await;
if result.is_err() {
if let Some(registry) = &self.inbound_registry {
registry.record_response_send_failed();
}
}
result
}
pub fn session_id(&self) -> SessionId {
self.session_id
}
}
pub struct Responder {
session_id: SessionId,
message_id: u32,
biz_type: u8,
transport: Arc<Transport>,
responded: std::sync::atomic::AtomicBool,
}
impl Responder {
pub(crate) fn new(
session_id: SessionId,
message_id: u32,
biz_type: u8,
transport: Arc<Transport>,
) -> Self {
Self {
session_id,
message_id,
biz_type,
transport,
responded: std::sync::atomic::AtomicBool::new(false),
}
}
pub async fn respond(self, data: Vec<u8>) -> Result<(), crate::TransportError> {
if self
.responded
.swap(true, std::sync::atomic::Ordering::SeqCst)
{
return Err(crate::TransportError::protocol_error(
"session",
"Already responded to this request",
));
}
let response_packet = Packet {
header: crate::packet::FixedHeader {
version: 1,
compression: crate::packet::CompressionType::None,
packet_type: crate::packet::PacketType::Response,
biz_type: self.biz_type,
message_id: self.message_id,
ext_header_len: 0,
payload_len: data.len() as u32,
reserved: crate::packet::ReservedFlags::new(),
},
ext_header: Vec::new(),
payload: data,
};
self.transport.send(response_packet).await
}
pub fn session_id(&self) -> SessionId {
self.session_id
}
pub fn message_id(&self) -> u32 {
self.message_id
}
}
#[derive(Clone)]
pub struct SessionHandle {
pub(crate) tx: Sender<ActorMessage>,
pub(crate) transport: Arc<Transport>,
}
impl SessionHandle {
pub async fn send_event(
&self,
event: TransportEvent,
) -> Result<(), tokio::sync::mpsc::error::SendError<TransportEvent>> {
self.tx
.send_async(ActorMessage::InboundEvent(event))
.await
.map_err(|e| {
match e.0 {
ActorMessage::InboundEvent(evt) => tokio::sync::mpsc::error::SendError(evt),
_ => unreachable!(),
}
})
}
async fn enqueue(&self, message: ActorMessage) -> Result<(), crate::TransportError> {
match tokio::time::timeout(SEND_QUEUE_WAIT, self.tx.send_async(message)).await {
Ok(Ok(())) => Ok(()),
Ok(Err(_)) => Err(crate::TransportError::connection_error(
"Actor channel closed",
false,
)),
Err(_elapsed) => Err(crate::TransportError::resource_error(
"session_actor_outbound_queue",
self.tx.len(),
self.tx.capacity().unwrap_or(DEFAULT_ACTOR_BUFFER_SIZE),
)),
}
}
pub async fn send_packet(&self, packet: Packet) -> Result<(), crate::TransportError> {
self.enqueue(ActorMessage::Send(packet)).await
}
pub async fn send_packet_with_reply(
&self,
packet: Packet,
) -> Result<(), crate::TransportError> {
let (reply_tx, reply_rx) = tokio::sync::oneshot::channel();
self.enqueue(ActorMessage::SendWithReply {
packet,
reply: reply_tx,
})
.await?;
reply_rx.await.map_err(|_| {
crate::TransportError::connection_error("Actor dropped before reply", false)
})?
}
pub fn transport(&self) -> &Arc<Transport> {
&self.transport
}
}
const BATCH_SIZE: usize = 64;
pub struct SessionActor {
session_id: SessionId,
transport: Arc<Transport>,
rx: Receiver<ActorMessage>,
handler: Arc<dyn SessionHandler>,
inbound_registry: Option<Arc<RequestRegistry>>,
}
impl SessionActor {
pub fn new(
session_id: SessionId,
transport: Arc<Transport>,
rx: Receiver<ActorMessage>,
handler: Arc<dyn SessionHandler>,
) -> Self {
Self {
session_id,
transport,
rx,
handler,
inbound_registry: None,
}
}
pub(crate) fn with_inbound_registry(
mut self,
inbound_registry: Option<Arc<RequestRegistry>>,
) -> Self {
self.inbound_registry = inbound_registry;
self
}
pub async fn run(self) {
tracing::debug!(
"[ACTOR] SessionActor started for session {}",
self.session_id
);
self.handler.on_connected(self.session_id).await;
let mut batch: Vec<ActorMessage> = Vec::with_capacity(BATCH_SIZE);
let sender = SessionSender::new(
self.session_id,
self.transport.clone(),
self.inbound_registry.clone(),
);
loop {
batch.clear();
match self.rx.recv_async().await {
Ok(first) => batch.push(first),
Err(_) => break,
}
while batch.len() < BATCH_SIZE {
match self.rx.try_recv() {
Ok(next) => batch.push(next),
Err(flume::TryRecvError::Empty) => break,
Err(flume::TryRecvError::Disconnected) => break,
}
}
let mut should_break = false;
for msg in batch.drain(..) {
match msg {
ActorMessage::InboundEvent(event) => {
match event {
TransportEvent::MessageReceived(packet) => {
self.handler
.on_message(self.session_id, packet, sender.clone())
.await;
}
TransportEvent::MessageSent { packet_id } => {
tracing::trace!(
"[ACTOR] Message {} sent for session {}",
packet_id,
self.session_id
);
}
TransportEvent::ConnectionClosed { reason } => {
tracing::debug!(
"[ACTOR] Session {} closed: {:?}",
self.session_id,
reason
);
self.handler.on_disconnected(self.session_id, reason).await;
should_break = true;
}
TransportEvent::TransportError { error } => {
tracing::warn!(
"[ACTOR] Session {} error: {:?}",
self.session_id,
error
);
self.handler.on_error(self.session_id, error).await;
}
_ => {
tracing::trace!(
"[ACTOR] Session {} received unhandled event",
self.session_id
);
}
}
}
ActorMessage::Send(packet) => {
if let Err(e) = self.transport.send(packet).await {
tracing::warn!(
"[ACTOR] Session {} send failed: {:?}",
self.session_id,
e
);
}
}
ActorMessage::SendWithReply { packet, reply } => {
let result = self.transport.send(packet).await;
let _ = reply.send(result);
}
ActorMessage::Close => {
tracing::debug!(
"[ACTOR] Session {} received close command",
self.session_id
);
should_break = true;
}
}
}
if should_break {
break;
}
}
tracing::debug!(
"[ACTOR] SessionActor stopped for session {}",
self.session_id
);
}
}
pub const DEFAULT_ACTOR_BUFFER_SIZE: usize = 512;
pub fn create_session_actor(
session_id: SessionId,
transport: Arc<Transport>,
handler: Arc<dyn SessionHandler>,
buffer_size: usize,
) -> (SessionHandle, SessionActor) {
let (tx, rx) = bounded(buffer_size);
let handle = SessionHandle {
tx,
transport: transport.clone(),
};
let actor = SessionActor::new(session_id, transport, rx, handler);
(handle, actor)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::transport::{config::TransportConfig, context::TransportContext};
async fn test_handle(tx: Sender<ActorMessage>) -> SessionHandle {
let context = TransportContext::new().await.unwrap();
let transport = Arc::new(Transport::with_context(
TransportConfig::default(),
&context,
));
SessionHandle { tx, transport }
}
#[tokio::test]
async fn outbound_send_fails_when_actor_mailbox_stays_full() {
let (tx, _rx) = bounded(1);
tx.try_send(ActorMessage::Send(Packet::one_way(1, vec![1])))
.unwrap();
let handle = test_handle(tx).await;
let error = handle
.send_packet_with_reply(Packet::one_way(2, vec![2]))
.await
.unwrap_err();
assert!(matches!(
error,
crate::TransportError::Resource { ref resource, current: 1, limit: 1 }
if resource == "session_actor_outbound_queue"
));
}
#[tokio::test]
async fn transient_full_actor_mailbox_is_absorbed() {
let (tx, rx) = bounded(1);
tx.try_send(ActorMessage::Send(Packet::one_way(1, vec![1])))
.unwrap();
let handle = test_handle(tx).await;
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
let _ = rx.recv_async().await;
tokio::time::sleep(std::time::Duration::from_secs(1)).await;
});
assert!(handle
.send_packet(Packet::one_way(2, vec![2]))
.await
.is_ok());
}
}