use crate::command::ConnectionInfo;
use crate::error::TransportError;
use crate::packet::Packet;
use crate::transport::request_registry::{MarkResult, RequestRegistry};
use crate::{CloseReason, PacketId, SessionId};
use std::net::SocketAddr;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Instant;
#[derive(Debug, Clone)]
pub enum TransportEvent {
ConnectionEstablished {
info: ConnectionInfo,
},
ConnectionClosed {
reason: CloseReason,
},
MessageReceived(Packet),
MessageSent {
packet_id: PacketId,
},
TransportError {
error: TransportError,
},
ServerStarted {
address: SocketAddr,
},
ServerStopped,
ClientConnected {
address: SocketAddr,
},
ClientDisconnected,
RequestReceived(RequestContext),
}
pub trait ProtocolEvent: Clone + Send + std::fmt::Debug + 'static {
fn into_transport_event(self) -> TransportEvent;
fn session_id(&self) -> Option<SessionId>;
fn is_data_event(&self) -> bool;
fn is_error_event(&self) -> bool;
}
#[derive(Debug, Clone)]
pub enum ConnectionEvent {
Established {
session_id: SessionId,
info: ConnectionInfo,
},
Closed {
session_id: SessionId,
reason: CloseReason,
},
}
impl TransportEvent {
pub fn session_id(&self) -> Option<SessionId> {
match self {
TransportEvent::ConnectionEstablished { .. } => None,
TransportEvent::ConnectionClosed { .. } => None,
TransportEvent::MessageReceived(..) => None,
TransportEvent::MessageSent { .. } => None,
TransportEvent::TransportError { .. } => None,
TransportEvent::ServerStarted { .. } => None,
TransportEvent::ServerStopped => None,
TransportEvent::ClientConnected { .. } => None,
TransportEvent::ClientDisconnected => None,
TransportEvent::RequestReceived(context) => context.peer,
}
}
pub fn is_connection_event(&self) -> bool {
matches!(
self,
TransportEvent::ConnectionEstablished { .. } | TransportEvent::ConnectionClosed { .. }
)
}
pub fn is_data_event(&self) -> bool {
matches!(
self,
TransportEvent::MessageReceived(..) | TransportEvent::RequestReceived(..)
)
}
pub fn is_error_event(&self) -> bool {
matches!(self, TransportEvent::TransportError { .. })
}
pub fn is_server_event(&self) -> bool {
matches!(
self,
TransportEvent::ServerStarted { .. } | TransportEvent::ServerStopped
)
}
pub fn is_client_event(&self) -> bool {
matches!(
self,
TransportEvent::ClientConnected { .. } | TransportEvent::ClientDisconnected
)
}
}
#[derive(Debug, Clone)]
pub enum TcpEvent {
ListenerBound { addr: SocketAddr },
AcceptError { error: String },
ConnectionTimeout { session_id: SessionId },
}
impl ProtocolEvent for TcpEvent {
fn into_transport_event(self) -> TransportEvent {
match self {
TcpEvent::ListenerBound { addr } => TransportEvent::ServerStarted { address: addr },
TcpEvent::AcceptError { error } => TransportEvent::TransportError {
error: TransportError::connection_error(
format!(
"IO error: {:?}",
std::io::Error::new(std::io::ErrorKind::Other, error)
),
true,
),
},
TcpEvent::ConnectionTimeout { session_id } => TransportEvent::ConnectionClosed {
reason: CloseReason::Timeout,
},
}
}
fn session_id(&self) -> Option<SessionId> {
match self {
TcpEvent::ConnectionTimeout { session_id } => Some(*session_id),
_ => None,
}
}
fn is_data_event(&self) -> bool {
false
}
fn is_error_event(&self) -> bool {
matches!(self, TcpEvent::AcceptError { .. })
}
}
#[derive(Debug, Clone)]
pub enum WebSocketEvent {
HandshakeCompleted {
session_id: SessionId,
},
PingReceived {
session_id: SessionId,
},
PongReceived {
session_id: SessionId,
},
InvalidFrame {
session_id: SessionId,
error: String,
},
}
impl ProtocolEvent for WebSocketEvent {
fn into_transport_event(self) -> TransportEvent {
match self {
WebSocketEvent::HandshakeCompleted { session_id } => {
TransportEvent::ConnectionEstablished {
info: ConnectionInfo::default(), }
}
WebSocketEvent::InvalidFrame { session_id, error } => TransportEvent::TransportError {
error: TransportError::protocol_error("generic", error),
},
_ => {
TransportEvent::TransportError {
error: TransportError::protocol_error(
"generic",
"Unhandled WebSocket event".to_string(),
),
}
}
}
}
fn session_id(&self) -> Option<SessionId> {
match self {
WebSocketEvent::HandshakeCompleted { session_id } => Some(*session_id),
WebSocketEvent::PingReceived { session_id } => Some(*session_id),
WebSocketEvent::PongReceived { session_id } => Some(*session_id),
WebSocketEvent::InvalidFrame { session_id, .. } => Some(*session_id),
}
}
fn is_data_event(&self) -> bool {
matches!(
self,
WebSocketEvent::PingReceived { .. } | WebSocketEvent::PongReceived { .. }
)
}
fn is_error_event(&self) -> bool {
matches!(self, WebSocketEvent::InvalidFrame { .. })
}
}
#[derive(Debug, Clone)]
pub enum QuicEvent {
StreamOpened {
session_id: SessionId,
stream_id: u64,
},
StreamClosed {
session_id: SessionId,
stream_id: u64,
},
CertificateVerified {
session_id: SessionId,
},
ConnectionIdRetired {
session_id: SessionId,
connection_id: u64,
},
}
impl ProtocolEvent for QuicEvent {
fn into_transport_event(self) -> TransportEvent {
match self {
QuicEvent::StreamOpened { session_id, .. } => {
TransportEvent::ConnectionEstablished {
info: ConnectionInfo::default(),
}
}
QuicEvent::StreamClosed { session_id, .. } => TransportEvent::ConnectionClosed {
reason: CloseReason::Normal,
},
_ => {
TransportEvent::TransportError {
error: TransportError::protocol_error(
"generic",
"Unhandled QUIC event".to_string(),
),
}
}
}
}
fn session_id(&self) -> Option<SessionId> {
match self {
QuicEvent::StreamOpened { session_id, .. } => Some(*session_id),
QuicEvent::StreamClosed { session_id, .. } => Some(*session_id),
QuicEvent::CertificateVerified { session_id } => Some(*session_id),
QuicEvent::ConnectionIdRetired { session_id, .. } => Some(*session_id),
}
}
fn is_data_event(&self) -> bool {
false
}
fn is_error_event(&self) -> bool {
false
}
}
#[derive(Debug, Clone)]
pub struct Message {
pub peer: Option<SessionId>,
pub data: Vec<u8>,
pub message_id: u32,
}
impl Message {
pub fn as_text(&self) -> Result<String, std::string::FromUtf8Error> {
String::from_utf8(self.data.clone())
}
pub fn as_text_lossy(&self) -> String {
String::from_utf8_lossy(&self.data).to_string()
}
pub fn as_bytes(&self) -> &[u8] {
&self.data
}
}
pub struct RequestContext {
pub peer: Option<SessionId>,
pub data: Vec<u8>,
pub request_id: u32,
pub biz_type: u8,
responder: Arc<dyn Fn(Vec<u8>) + Send + Sync + 'static>,
responded: Arc<std::sync::atomic::AtomicBool>,
is_primary: bool,
}
impl RequestContext {
pub fn new(
peer: Option<SessionId>,
data: Vec<u8>,
request_id: u32,
biz_type: u8,
responder: Arc<dyn Fn(Vec<u8>) + Send + Sync + 'static>,
) -> Self {
Self {
peer,
data,
request_id,
biz_type,
responder,
responded: Arc::new(std::sync::atomic::AtomicBool::new(false)),
is_primary: false, }
}
pub fn as_text(&self) -> Result<String, std::string::FromUtf8Error> {
String::from_utf8(self.data.clone())
}
pub fn as_text_lossy(&self) -> String {
String::from_utf8_lossy(&self.data).to_string()
}
pub fn as_bytes(&self) -> &[u8] {
&self.data
}
pub fn respond_text(&mut self, response: &str) {
self.respond_bytes(response.as_bytes());
}
pub fn respond_bytes(&mut self, response: &[u8]) {
if self
.responded
.compare_exchange(
false,
true,
std::sync::atomic::Ordering::SeqCst,
std::sync::atomic::Ordering::SeqCst,
)
.is_ok()
{
(self.responder)(response.to_vec());
} else {
tracing::warn!(
"[WARN] RequestContext already responded (ID: {})",
self.request_id
);
}
}
pub(crate) fn set_primary(&mut self) {
self.is_primary = true;
}
}
impl Clone for RequestContext {
fn clone(&self) -> Self {
Self {
peer: self.peer,
data: self.data.clone(),
request_id: self.request_id,
biz_type: self.biz_type, responder: self.responder.clone(), responded: self.responded.clone(), is_primary: false, }
}
}
impl std::fmt::Debug for RequestContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RequestContext")
.field("peer", &self.peer)
.field("data", &format!("{} bytes", self.data.len()))
.field("request_id", &self.request_id)
.field("biz_type", &self.biz_type)
.field(
"responded",
&self.responded.load(std::sync::atomic::Ordering::SeqCst),
)
.finish()
}
}
impl Drop for RequestContext {
fn drop(&mut self) {
if self.is_primary && !self.responded.load(std::sync::atomic::Ordering::SeqCst) {
tracing::warn!(
"[WARN] RequestContext dropped without response (ID: {})",
self.request_id
);
}
}
}
#[derive(Debug, Clone)]
pub enum ClientEvent {
Connected { info: ConnectionInfo },
Disconnected { reason: CloseReason },
MessageReceived(TransportContext),
MessageSent { message_id: u32 },
Error { error: TransportError },
}
#[derive(Debug, Clone)]
pub enum ServerEvent {
ConnectionEstablished {
session_id: SessionId,
info: ConnectionInfo,
},
ConnectionClosed {
session_id: SessionId,
reason: CloseReason,
},
MessageReceived {
session_id: SessionId,
context: TransportContext,
},
MessageSent {
session_id: SessionId,
message_id: u32,
},
TransportError {
session_id: Option<SessionId>,
error: TransportError,
},
ServerStarted {
address: std::net::SocketAddr,
},
ServerStopped,
}
impl ServerEvent {
pub fn from_transport_event(event: TransportEvent) -> Option<Self> {
match event {
TransportEvent::ConnectionEstablished { info } => None,
TransportEvent::ConnectionClosed { reason } => None,
TransportEvent::MessageReceived(packet) => None,
TransportEvent::MessageSent { packet_id } => None,
TransportEvent::TransportError { error } => None,
TransportEvent::ServerStarted { address } => {
Some(ServerEvent::ServerStarted { address })
}
TransportEvent::ServerStopped => Some(ServerEvent::ServerStopped),
TransportEvent::RequestReceived(ctx) => None,
_ => None,
}
}
pub fn from_transport_event_with_session(
event: TransportEvent,
session_id: SessionId,
) -> Option<Self> {
match event {
TransportEvent::ConnectionEstablished { info } => {
Some(ServerEvent::ConnectionEstablished { session_id, info })
}
TransportEvent::ConnectionClosed { reason } => {
Some(ServerEvent::ConnectionClosed { session_id, reason })
}
TransportEvent::MessageReceived(packet) => {
let context = TransportContext::new_oneway(
Some(session_id),
packet.header.message_id,
packet.header.biz_type,
if packet.ext_header.is_empty() {
None
} else {
Some(packet.ext_header.clone())
},
packet.payload.clone(),
);
Some(ServerEvent::MessageReceived {
session_id,
context,
})
}
TransportEvent::MessageSent { packet_id } => Some(ServerEvent::MessageSent {
session_id,
message_id: packet_id,
}),
TransportEvent::TransportError { error } => Some(ServerEvent::TransportError {
session_id: Some(session_id),
error,
}),
TransportEvent::RequestReceived(ctx) => {
let transport_ctx = TransportContext::new_request(
Some(session_id),
ctx.request_id,
ctx.biz_type, None, ctx.data.clone(),
ctx.responder.clone(),
);
Some(ServerEvent::MessageReceived {
session_id,
context: transport_ctx,
})
}
TransportEvent::ServerStarted { address } => {
Some(ServerEvent::ServerStarted { address })
}
TransportEvent::ServerStopped => Some(ServerEvent::ServerStopped),
_ => None,
}
}
}
impl ClientEvent {
pub fn from_transport_event(event: TransportEvent) -> Option<Self> {
match event {
TransportEvent::ConnectionEstablished { info } => Some(ClientEvent::Connected { info }),
TransportEvent::ConnectionClosed { reason } => {
Some(ClientEvent::Disconnected { reason })
}
TransportEvent::MessageReceived(packet) => {
match packet.header.packet_type {
crate::packet::PacketType::Request => {
None
}
_ => {
let context = TransportContext::new_oneway(
None,
packet.header.message_id,
packet.header.biz_type,
if packet.ext_header.is_empty() {
None
} else {
Some(packet.ext_header.clone())
},
packet.payload.clone(),
);
Some(ClientEvent::MessageReceived(context))
}
}
}
TransportEvent::MessageSent { packet_id } => Some(ClientEvent::MessageSent {
message_id: packet_id,
}),
TransportEvent::TransportError { error } => Some(ClientEvent::Error { error }),
_ => None,
}
}
pub fn is_connection_event(&self) -> bool {
matches!(
self,
ClientEvent::Connected { .. } | ClientEvent::Disconnected { .. }
)
}
pub fn is_data_event(&self) -> bool {
matches!(
self,
ClientEvent::MessageReceived(..) | ClientEvent::MessageSent { .. }
)
}
pub fn is_error_event(&self) -> bool {
matches!(self, ClientEvent::Error { .. })
}
}
pub struct TransportContext {
pub peer: Option<SessionId>,
pub message_id: u32,
pub biz_type: u8,
pub ext_header: Option<Vec<u8>>,
pub data: Vec<u8>,
pub timestamp: Instant,
kind: TransportContextKind,
}
type ResponderFn = Arc<
dyn Fn(Vec<u8>) -> futures::future::BoxFuture<'static, Result<(), crate::error::TransportError>>
+ Send
+ Sync,
>;
enum TransportContextKind {
OneWay,
Request {
responder: ResponderFn,
responded: Arc<AtomicBool>,
is_primary: bool, request_registry: Option<Arc<RequestRegistry>>,
},
}
impl TransportContext {
pub fn new_oneway(
peer: Option<SessionId>,
message_id: u32,
biz_type: u8,
ext_header: Option<Vec<u8>>,
data: Vec<u8>,
) -> Self {
Self {
peer,
message_id,
biz_type,
ext_header,
data,
timestamp: Instant::now(),
kind: TransportContextKind::OneWay,
}
}
pub fn new_request(
peer: Option<SessionId>,
message_id: u32,
biz_type: u8,
ext_header: Option<Vec<u8>>,
data: Vec<u8>,
responder: Arc<dyn Fn(Vec<u8>) + Send + Sync + 'static>,
) -> Self {
let responder: ResponderFn = Arc::new(
move |data: Vec<u8>| -> futures::future::BoxFuture<
'static,
Result<(), crate::error::TransportError>,
> {
responder(data);
Box::pin(async { Ok(()) })
},
);
Self::new_request_with_registry(
peer, message_id, biz_type, ext_header, data, responder, None,
)
}
pub(crate) fn new_request_with_registry(
peer: Option<SessionId>,
message_id: u32,
biz_type: u8,
ext_header: Option<Vec<u8>>,
data: Vec<u8>,
responder: ResponderFn,
request_registry: Option<Arc<RequestRegistry>>,
) -> Self {
Self {
peer,
message_id,
biz_type,
ext_header,
data,
timestamp: Instant::now(),
kind: TransportContextKind::Request {
responder,
responded: Arc::new(AtomicBool::new(false)),
is_primary: false, request_registry,
},
}
}
pub(crate) fn set_primary(&mut self) {
if let TransportContextKind::Request { is_primary, .. } = &mut self.kind {
*is_primary = true;
}
}
pub fn is_request(&self) -> bool {
matches!(self.kind, TransportContextKind::Request { .. })
}
pub fn as_text_lossy(&self) -> String {
String::from_utf8_lossy(&self.data).to_string()
}
pub fn respond(mut self, response: Vec<u8>) {
match &mut self.kind {
TransportContextKind::Request {
responder,
responded,
request_registry,
..
} => {
if let Some(registry) = request_registry {
match registry.mark_responded(self.peer, self.message_id) {
MarkResult::Updated => {}
MarkResult::Already(state) => {
tracing::debug!(
"[RESPOND] Skip duplicate/late response: request_id={}, state={:?}",
self.message_id,
state
);
return;
}
MarkResult::NotFound => {
tracing::debug!(
"[RESPOND] Skip response for inactive request: request_id={}, session_id={:?}",
self.message_id,
self.peer
);
return;
}
}
}
if responded
.compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst)
.is_ok()
{
let fut = responder(response);
tokio::spawn(async move {
if let Err(e) = fut.await {
tracing::debug!(
"[RESPOND] fire-and-forget response send failed: {:?}",
e
);
}
});
} else {
tracing::debug!(
"[RESPOND] TransportContext already responded locally (ID: {})",
self.message_id
);
}
}
TransportContextKind::OneWay => {
tracing::warn!(
"[WARN] Cannot respond to one-way message (ID: {})",
self.message_id
);
}
}
}
pub fn respond_bytes(self, response: &[u8]) {
self.respond(response.to_vec());
}
pub async fn respond_checked(
mut self,
response: Vec<u8>,
) -> Result<(), crate::error::TransportError> {
let fut = match &mut self.kind {
TransportContextKind::Request {
responder,
responded,
request_registry,
..
} => {
if let Some(registry) = request_registry {
match registry.mark_responded(self.peer, self.message_id) {
MarkResult::Updated => {}
_ => return Ok(()),
}
}
if responded
.compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst)
.is_ok()
{
Some(responder(response))
} else {
None
}
}
TransportContextKind::OneWay => {
return Err(crate::error::TransportError::protocol_error(
"generic",
"Cannot respond to one-way message",
));
}
};
match fut {
Some(f) => f.await,
None => Ok(()),
}
}
}
impl Clone for TransportContext {
fn clone(&self) -> Self {
let kind = match &self.kind {
TransportContextKind::OneWay => TransportContextKind::OneWay,
TransportContextKind::Request {
responder,
responded,
request_registry,
..
} => TransportContextKind::Request {
responder: responder.clone(),
responded: responded.clone(),
is_primary: false,
request_registry: request_registry.clone(),
},
};
Self {
peer: self.peer,
message_id: self.message_id,
biz_type: self.biz_type,
ext_header: self.ext_header.clone(),
data: self.data.clone(),
timestamp: self.timestamp,
kind,
}
}
}
impl std::fmt::Debug for TransportContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TransportContext")
.field("peer", &self.peer)
.field("message_id", &self.message_id)
.field("data", &format!("{} bytes", self.data.len()))
.field("timestamp", &self.timestamp)
.field("is_request", &self.is_request())
.finish()
}
}
impl Drop for TransportContext {
fn drop(&mut self) {
if let TransportContextKind::Request {
responded,
is_primary,
..
} = &self.kind
{
if *is_primary && !responded.load(Ordering::SeqCst) {
tracing::debug!(
"[DROP] TransportContext dropped before response (request_id={}, session_id={:?}, biz_type={})",
self.message_id,
self.peer,
self.biz_type
);
}
}
}
}
#[derive(Debug, Clone)]
pub struct TransportResult {
pub peer: Option<SessionId>,
pub message_id: u32,
pub timestamp: Instant,
pub data: Option<Vec<u8>>,
pub status: TransportStatus,
}
#[derive(Debug, Clone, PartialEq)]
pub enum TransportStatus {
Sent,
Timeout,
ConnectionError,
Completed,
}
impl TransportResult {
pub fn new_sent(peer: Option<SessionId>, message_id: u32) -> Self {
Self {
peer,
message_id,
timestamp: Instant::now(),
data: None,
status: TransportStatus::Sent,
}
}
pub fn new_completed(peer: Option<SessionId>, message_id: u32, data: Vec<u8>) -> Self {
Self {
peer,
message_id,
timestamp: Instant::now(),
data: Some(data),
status: TransportStatus::Completed,
}
}
pub fn new_timeout(peer: Option<SessionId>, message_id: u32) -> Self {
Self {
peer,
message_id,
timestamp: Instant::now(),
data: None,
status: TransportStatus::Timeout,
}
}
pub fn new_connection_error(peer: Option<SessionId>, message_id: u32) -> Self {
Self {
peer,
message_id,
timestamp: Instant::now(),
data: None,
status: TransportStatus::ConnectionError,
}
}
pub fn is_sent(&self) -> bool {
matches!(
self.status,
TransportStatus::Sent | TransportStatus::Completed
)
}
pub fn has_response(&self) -> bool {
self.data.is_some()
}
}
#[cfg(test)]
mod respond_checked_tests {
use super::*;
#[tokio::test]
async fn respond_checked_returns_ok_from_responder() {
let ok: ResponderFn = Arc::new(|_data| Box::pin(async { Ok(()) }));
let ctx = TransportContext::new_request_with_registry(
Some(SessionId(1)),
42,
0,
None,
b"req".to_vec(),
ok,
None,
);
assert!(ctx.respond_checked(b"resp".to_vec()).await.is_ok());
}
#[tokio::test]
async fn respond_checked_propagates_send_error() {
let err: ResponderFn = Arc::new(|_data| {
Box::pin(async {
Err(crate::error::TransportError::connection_error(
"boom", false,
))
})
});
let ctx = TransportContext::new_request_with_registry(
Some(SessionId(2)),
43,
0,
None,
b"req".to_vec(),
err,
None,
);
assert!(ctx.respond_checked(b"resp".to_vec()).await.is_err());
}
#[tokio::test]
async fn respond_checked_on_one_way_is_error() {
let ctx = TransportContext::new_oneway(Some(SessionId(3)), 44, 0, None, b"data".to_vec());
assert!(ctx.respond_checked(b"resp".to_vec()).await.is_err());
}
}