use crate::{
connection::Connection,
event::TransportEvent,
protocol::ProtocolRegistry,
transport::{
config::TransportConfig, connection_state::ConnectionStateManager,
context::TransportContext, memory_pool::OptimizedMemoryPool,
},
Packet, SessionId, TransportError,
};
use bytes::Bytes;
use std::sync::{
atomic::{AtomicU32, Ordering},
Arc,
};
use tokio::sync::{broadcast, oneshot, Mutex};
pub struct Transport {
config: TransportConfig,
protocol_registry: Arc<ProtocolRegistry>,
memory_pool: Arc<OptimizedMemoryPool>,
connection: Arc<Mutex<Option<Box<dyn Connection>>>>,
session_id: Arc<Mutex<Option<SessionId>>>,
state_manager: ConnectionStateManager,
event_sender: broadcast::Sender<TransportEvent>,
request_tracker: Arc<RequestTracker>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct RequestTrackerKey {
pub session_id: Option<SessionId>,
pub message_id: u32,
}
impl RequestTrackerKey {
pub fn new(session_id: Option<SessionId>, message_id: u32) -> Self {
Self {
session_id,
message_id,
}
}
}
const REQUEST_WAITER_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
pub struct RequestTracker {
registry: Arc<crate::transport::request_registry::RequestRegistry>,
next_id: AtomicU32,
}
impl RequestTracker {
pub fn new() -> Self {
Self {
registry: Arc::new(crate::transport::request_registry::RequestRegistry::new()),
next_id: AtomicU32::new(1),
}
}
pub fn new_with_start_id(start_id: u32) -> Self {
Self {
registry: Arc::new(crate::transport::request_registry::RequestRegistry::new()),
next_id: AtomicU32::new(start_id),
}
}
fn register_waiter(&self, session_id: Option<SessionId>, id: u32) -> oneshot::Receiver<Packet> {
match self
.registry
.try_register_waiter(id, session_id, 0, REQUEST_WAITER_TIMEOUT)
{
Ok(rx) => rx,
Err(_) => {
tracing::warn!(
"[REQUEST] Duplicate pending request refused: session_id={:?}, message_id={}",
session_id,
id
);
let (_tx, rx) = oneshot::channel();
rx
}
}
}
pub fn register(&self) -> (u32, oneshot::Receiver<Packet>) {
let id = self.next_id.fetch_add(1, Ordering::Relaxed);
(id, self.register_waiter(None, id))
}
pub fn register_with_id(&self, id: u32) -> (u32, oneshot::Receiver<Packet>) {
self.register_with_session_id(None, id)
}
pub fn register_with_session(
&self,
session_id: SessionId,
id: u32,
) -> (u32, oneshot::Receiver<Packet>) {
self.register_with_session_id(Some(session_id), id)
}
pub fn register_with_session_id(
&self,
session_id: Option<SessionId>,
id: u32,
) -> (u32, oneshot::Receiver<Packet>) {
(id, self.register_waiter(session_id, id))
}
pub fn complete(&self, id: u32, packet: Packet) -> bool {
self.complete_with_session_id(None, id, packet)
}
pub fn complete_with_session(&self, session_id: SessionId, id: u32, packet: Packet) -> bool {
self.complete_with_session_id(Some(session_id), id, packet)
}
pub fn complete_with_session_id(
&self,
session_id: Option<SessionId>,
id: u32,
packet: Packet,
) -> bool {
self.registry.complete_waiter(session_id, id, packet)
}
pub fn remove(&self, id: u32) -> bool {
self.remove_with_session_id(None, id)
}
pub fn remove_with_session(&self, session_id: SessionId, id: u32) -> bool {
self.remove_with_session_id(Some(session_id), id)
}
pub fn remove_with_session_id(&self, session_id: Option<SessionId>, id: u32) -> bool {
self.registry.abort_waiter(session_id, id)
}
pub fn fail_session(&self, session_id: Option<SessionId>) -> usize {
match session_id {
Some(sid) => self.registry.close_session_pending(sid),
None => 0,
}
}
pub fn fail_all(&self) -> usize {
self.registry.abort_all()
}
pub fn clear(&self) {
self.fail_all();
}
pub fn next_message_id(&self) -> u32 {
self.next_id.fetch_add(1, Ordering::Relaxed)
}
}
impl Transport {
pub fn with_context(config: TransportConfig, ctx: &TransportContext) -> Self {
let (event_sender, _) = broadcast::channel(8192);
Self {
config,
protocol_registry: ctx.protocol_registry.clone(),
memory_pool: ctx.memory_pool.clone(),
connection: Arc::new(Mutex::new(None)),
session_id: Arc::new(Mutex::new(None)),
state_manager: ConnectionStateManager::new(),
event_sender,
request_tracker: Arc::new(RequestTracker::new()),
}
}
pub async fn connect_with_config<T>(
self: &Arc<Self>,
config: T,
) -> Result<SessionId, TransportError>
where
T: crate::protocol::client_config::ConnectableConfig,
{
config.connect(Arc::clone(self)).await
}
pub async fn send(&self, packet: Packet) -> Result<(), TransportError> {
let mut guard = self.connection.lock().await;
match guard.as_mut() {
Some(conn) => conn.send(packet).await,
None => Err(TransportError::connection_error("Not connected", false)),
}
}
pub(crate) async fn set_frame_policy(&self, policy: crate::packet::FramePolicy) {
if let Some(conn) = self.connection.lock().await.as_ref() {
conn.set_frame_policy(policy);
}
}
pub async fn disconnect(&self) -> Result<(), TransportError> {
if let Some(session_id) = self.current_session_id().await {
self.close_session(session_id).await
} else {
Err(TransportError::connection_error("Not connected", false))
}
}
pub async fn close_session(&self, session_id: SessionId) -> Result<(), TransportError> {
if !self.state_manager.try_start_closing(session_id).await {
tracing::debug!(
"Session {} already closing or closed, skipping close logic",
session_id
);
return Ok(());
}
tracing::info!("[CONN] Starting graceful session shutdown: {}", session_id);
let failed_pending = self.request_tracker.fail_all();
if failed_pending > 0 {
tracing::debug!(
"[REQUEST] Failed {} pending requests during session {} shutdown",
failed_pending,
session_id
);
}
self.do_close_session(session_id).await?;
self.state_manager.mark_closed(session_id).await;
if self.session_id.lock().await.as_ref() == Some(&session_id) {
*self.session_id.lock().await = None;
*self.connection.lock().await = None;
}
tracing::info!("[SUCCESS] Session {} shutdown complete", session_id);
Ok(())
}
pub async fn force_close_session(&self, session_id: SessionId) -> Result<(), TransportError> {
if !self.state_manager.try_start_closing(session_id).await {
tracing::debug!(
"Session {} already closing or closed, skipping force close",
session_id
);
return Ok(());
}
tracing::info!("[CONN] Force closing session: {}", session_id);
let failed_pending = self.request_tracker.fail_all();
if failed_pending > 0 {
tracing::debug!(
"[REQUEST] Failed {} pending requests during session {} force close",
failed_pending,
session_id
);
}
if let Some(conn) = self.connection.lock().await.as_mut() {
let _ = conn.close().await;
}
self.state_manager.mark_closed(session_id).await;
if self.session_id.lock().await.as_ref() == Some(&session_id) {
*self.session_id.lock().await = None;
*self.connection.lock().await = None;
}
tracing::info!("[SUCCESS] Session {} force close complete", session_id);
Ok(())
}
async fn do_close_session(&self, session_id: SessionId) -> Result<(), TransportError> {
let mut guard = self.connection.lock().await;
if let Some(conn) = guard.as_mut() {
match tokio::time::timeout(
self.config.graceful_timeout,
self.try_graceful_close(&mut **conn),
)
.await
{
Ok(Ok(_)) => {
tracing::debug!("[SUCCESS] Session {} graceful close successful", session_id);
}
Ok(Err(e)) => {
tracing::warn!(
"[WARN] Session {} graceful close failed, executing force close: {:?}",
session_id,
e
);
let _ = conn.close().await;
}
Err(_) => {
tracing::warn!(
"[WARN] Session {} graceful close timeout, executing force close",
session_id
);
let _ = conn.close().await;
}
}
}
Ok(())
}
async fn try_graceful_close(&self, conn: &mut dyn Connection) -> Result<(), TransportError> {
tracing::debug!("[CONN] Using underlying protocol graceful close mechanism");
conn.close().await
}
pub async fn should_ignore_messages(&self, session_id: SessionId) -> bool {
self.state_manager.should_ignore_messages(session_id).await
}
pub async fn is_connected(&self) -> bool {
self.session_id.lock().await.is_some()
}
pub async fn current_session_id(&self) -> Option<SessionId> {
self.session_id.lock().await.as_ref().cloned()
}
pub async fn set_connection(
self: &Arc<Self>,
mut connection: Box<dyn Connection>,
session_id: SessionId,
) {
connection.set_session_id(session_id);
let event_receiver_opt = connection.event_stream();
*self.connection.lock().await = Some(connection);
*self.session_id.lock().await = Some(session_id);
self.state_manager.add_connection(session_id);
tracing::debug!("[SUCCESS] Transport connection set: {}", session_id);
if let Some(mut event_receiver) = event_receiver_opt {
let this = Arc::clone(self);
tokio::spawn(async move {
tracing::debug!(
"[LISTEN] Transport event consumer started (session: {})",
session_id
);
while let Ok(event) = event_receiver.recv().await {
this.on_event(event).await;
}
let failed_pending = this.request_tracker.fail_all();
if failed_pending > 0 {
tracing::debug!(
"[REQUEST] Failed {} pending requests after event stream ended (session: {})",
failed_pending,
session_id
);
}
tracing::debug!(
"[LISTEN] Transport event consumer ended (session: {})",
session_id
);
});
}
}
pub async fn set_connection_no_consumer(
&self,
mut connection: Box<dyn Connection>,
session_id: SessionId,
) {
connection.set_session_id(session_id);
*self.connection.lock().await = Some(connection);
*self.session_id.lock().await = Some(session_id);
self.state_manager.add_connection(session_id);
tracing::debug!(
"[SUCCESS] Transport connection set (no consumer): {}",
session_id
);
}
pub fn protocol_registry(&self) -> &ProtocolRegistry {
&self.protocol_registry
}
pub fn config(&self) -> &TransportConfig {
&self.config
}
pub fn memory_pool_stats(&self) -> crate::transport::memory_pool::OptimizedMemoryStatsSnapshot {
self.memory_pool.get_stats()
}
pub async fn get_event_stream(
&self,
) -> Option<tokio::sync::broadcast::Receiver<crate::event::TransportEvent>> {
if self.connection.lock().await.is_some() {
Some(self.event_sender.subscribe())
} else {
None
}
}
pub async fn request(&self, packet: Packet) -> Result<Packet, TransportError> {
if packet.header.packet_type != crate::packet::PacketType::Request {
return Err(TransportError::connection_error(
"Not a Request packet",
false,
));
}
let client_message_id = packet.header.message_id;
let session_id = self.current_session_id().await;
let (_, rx) = self
.request_tracker
.register_with_session_id(session_id, client_message_id);
if let Err(e) = self.send(packet).await {
self.request_tracker
.remove_with_session_id(session_id, client_message_id);
return Err(e);
}
let timeout_duration = std::time::Duration::from_secs(10);
match tokio::time::timeout(timeout_duration, rx).await {
Ok(Ok(resp)) => Ok(resp),
Ok(Err(_)) => Err(TransportError::connection_error("Connection closed", true)),
Err(_) => {
self.request_tracker
.remove_with_session_id(session_id, client_message_id);
Err(TransportError::timeout_error("request", timeout_duration))
}
}
}
fn decode_payload(&self, packet: &Packet) -> Result<Vec<u8>, TransportError> {
if packet.header.compression != crate::packet::CompressionType::None {
let mut packet_copy = packet.clone();
match packet_copy.decompress_payload() {
Ok(_) => Ok(packet_copy.payload),
Err(e) => {
tracing::warn!("[WARN] Failed to decompress packet: {}", e);
Err(TransportError::protocol_error(
"packet",
format!("Failed to decompress packet: {}", e),
))
}
}
} else {
Ok(packet.payload.clone())
}
}
pub async fn on_event(&self, event: crate::event::TransportEvent) {
match event {
crate::event::TransportEvent::MessageReceived(packet) => {
tracing::debug!(
"[TARGET] Transport::on_event processing message packet: ID={}, type={:?}",
packet.header.message_id,
packet.header.packet_type
);
match packet.header.packet_type {
crate::packet::PacketType::Response => {
let id = packet.header.message_id;
tracing::info!(
"[RECV] Processing response packet: ID={}, type={:?}, biz_type={}",
id,
packet.header.packet_type,
packet.header.biz_type
);
let session_id = self.current_session_id().await;
let completed = self.request_tracker.complete_with_session_id(
session_id,
id,
packet.clone(),
);
tracing::info!(
"[PROC] Response packet processing result: ID={}, completed={}",
id,
completed
);
if !completed {
tracing::warn!("[WARN] Response packet ID={} not found in request tracker, may be timeout or duplicate", id);
let _ = self
.event_sender
.send(crate::event::TransportEvent::MessageReceived(packet));
}
}
crate::packet::PacketType::Request => {
let id = packet.header.message_id;
tracing::debug!("[PROC] Received request packet, creating unified TransportContext: ID={}, type={:?}", id, packet.header.packet_type);
tracing::debug!(
"[SEND] Sending unified MessageReceived event (Request): ID={}",
id
);
let _ = self
.event_sender
.send(crate::event::TransportEvent::MessageReceived(packet));
}
crate::packet::PacketType::OneWay => {
tracing::debug!(
"[RECV] Processing one-way message packet: ID={}, type={:?}",
packet.header.message_id,
packet.header.packet_type
);
match self.decode_payload(&packet) {
Ok(data) => {
let session_id = self.session_id.lock().await.as_ref().cloned();
let _message = crate::event::Message {
peer: session_id,
data,
message_id: packet.header.message_id,
};
let _ = self
.event_sender
.send(crate::event::TransportEvent::MessageReceived(packet));
}
Err(e) => {
tracing::error!("[ERROR] Failed to unpack message data: {}", e);
let _ = self.event_sender.send(
crate::event::TransportEvent::TransportError { error: e },
);
}
}
}
}
}
crate::event::TransportEvent::ConnectionClosed { reason } => {
let failed_pending = self.request_tracker.fail_all();
if failed_pending > 0 {
tracing::debug!(
"[REQUEST] Failed {} pending requests after connection closed: {:?}",
failed_pending,
reason
);
}
let _ = self
.event_sender
.send(crate::event::TransportEvent::ConnectionClosed { reason });
}
_ => {
tracing::trace!("[SEND] Forwarding other event: {:?}", event);
let _ = self.event_sender.send(event);
}
}
}
pub fn subscribe_events(&self) -> broadcast::Receiver<TransportEvent> {
self.event_sender.subscribe()
}
pub async fn request_with_options(
&self,
data: Bytes,
options: super::TransportOptions,
) -> Result<Bytes, TransportError> {
let message_id = options
.message_id
.unwrap_or_else(|| self.request_tracker.next_message_id());
let mut packet = crate::packet::Packet {
header: crate::packet::FixedHeader {
version: 1,
compression: options
.compression
.unwrap_or(crate::packet::CompressionType::None),
packet_type: crate::packet::PacketType::Request,
biz_type: options.biz_type.unwrap_or(0),
message_id,
ext_header_len: options.ext_header.as_ref().map_or(0, |h| h.len() as u16),
payload_len: data.len() as u32,
reserved: crate::packet::ReservedFlags::new(),
},
ext_header: options.ext_header.unwrap_or_default().to_vec(),
payload: data.to_vec(),
};
if options.compression.is_some()
&& options.compression != Some(crate::packet::CompressionType::None)
{
if let Err(e) = packet.compress_payload() {
tracing::warn!("[WARN] Failed to compress packet: {}, using raw data", e);
}
}
let session_id = self.current_session_id().await;
let (_id, rx) = self
.request_tracker
.register_with_session_id(session_id, message_id);
tracing::info!(
"[SEND] Sending request: message_id={}, biz_type={}, timeout={:?}",
message_id,
packet.header.biz_type,
options.timeout
);
if let Err(e) = self.send(packet).await {
self.request_tracker
.remove_with_session_id(session_id, message_id);
return Err(e);
}
tracing::info!(
"[WAIT] Waiting for response: message_id={}, timeout={:?}",
message_id,
options.timeout
);
let timeout_duration = options
.timeout
.unwrap_or(std::time::Duration::from_secs(10));
match tokio::time::timeout(timeout_duration, rx).await {
Ok(Ok(resp)) => {
tracing::info!(
"[SUCCESS] Received response: message_id={}, biz_type={}, payload_len={}",
message_id,
resp.header.biz_type,
resp.payload.len()
);
self.decode_payload(&resp).map(Bytes::from)
}
Ok(Err(_)) => {
tracing::warn!("[WARN] Response channel closed: message_id={}", message_id);
Err(TransportError::connection_error("Connection closed", true))
}
Err(_) => {
self.request_tracker
.remove_with_session_id(session_id, message_id);
tracing::warn!(
"[WARN] Request timeout: message_id={}, timeout={:?}",
message_id,
timeout_duration
);
Err(TransportError::timeout_error("request", timeout_duration))
}
}
}
pub(crate) fn next_message_id(&self) -> u32 {
self.request_tracker.next_message_id()
}
pub async fn send_with_options(
&self,
data: Bytes,
options: super::TransportOptions,
) -> Result<(), TransportError> {
let message_id = options.message_id.unwrap_or_else(|| {
self.request_tracker
.next_id
.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
});
let mut packet = crate::packet::Packet {
header: crate::packet::FixedHeader {
version: 1,
compression: options
.compression
.unwrap_or(crate::packet::CompressionType::None),
packet_type: crate::packet::PacketType::OneWay,
biz_type: options.biz_type.unwrap_or(0),
message_id,
ext_header_len: options.ext_header.as_ref().map_or(0, |h| h.len() as u16),
payload_len: data.len() as u32,
reserved: crate::packet::ReservedFlags::new(),
},
ext_header: options.ext_header.unwrap_or_default().to_vec(),
payload: data.to_vec(),
};
if options.compression.is_some()
&& options.compression != Some(crate::packet::CompressionType::None)
{
if let Err(e) = packet.compress_payload() {
tracing::warn!("[WARN] Failed to compress packet: {}, using raw data", e);
}
}
self.send(packet).await?;
Ok(())
}
}
impl Clone for Transport {
fn clone(&self) -> Self {
Self {
config: self.config.clone(),
protocol_registry: self.protocol_registry.clone(),
memory_pool: self.memory_pool.clone(),
connection: self.connection.clone(),
session_id: self.session_id.clone(),
state_manager: self.state_manager.clone(),
event_sender: self.event_sender.clone(),
request_tracker: self.request_tracker.clone(),
}
}
}
impl std::fmt::Debug for Transport {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Transport")
.field("connected", &"<async>")
.field("session_id", &"<async>")
.finish()
}
}
#[cfg(test)]
mod request_tracker_tests {
use super::*;
#[test]
fn cross_session_response_cannot_complete_another_sessions_request() {
let tracker = RequestTracker::new();
let victim = SessionId(1);
let attacker = SessionId(2);
let (_, _rx) = tracker.register_with_session(victim, 100);
assert!(
!tracker.complete_with_session(attacker, 100, Packet::response(100, Vec::new())),
"attacker session must not complete victim's request"
);
assert!(
tracker.complete_with_session(victim, 100, Packet::response(100, Vec::new())),
"victim session must complete its own request"
);
}
#[test]
fn none_keyed_and_session_keyed_requests_are_distinct() {
let tracker = RequestTracker::new();
let (_, _rx) = tracker.register_with_id(42);
assert!(!tracker.complete_with_session(SessionId(9), 42, Packet::response(42, Vec::new())));
assert!(tracker.complete(42, Packet::response(42, Vec::new())));
}
#[test]
fn fail_session_clears_only_matching_session() {
let tracker = RequestTracker::new();
let (_, _a1) = tracker.register_with_session(SessionId(1), 10);
let (_, _a2) = tracker.register_with_session(SessionId(1), 11);
let (_, _b1) = tracker.register_with_session(SessionId(2), 10);
assert_eq!(tracker.fail_session(Some(SessionId(1))), 2);
assert!(tracker.complete_with_session(SessionId(2), 10, Packet::response(10, Vec::new())));
assert!(!tracker.complete_with_session(SessionId(1), 10, Packet::response(10, Vec::new())));
}
#[test]
fn fail_all_clears_every_pending_request() {
let tracker = RequestTracker::new();
let (_, _a) = tracker.register_with_session(SessionId(1), 1);
let (_, _b) = tracker.register_with_session(SessionId(2), 2);
let (_, _c) = tracker.register_with_id(3);
assert_eq!(tracker.fail_all(), 3);
assert!(!tracker.complete_with_session(SessionId(1), 1, Packet::response(1, Vec::new())));
}
}