use crate::{
command::{ConnectionInfo, ConnectionState},
connection::Connection,
error::TransportError,
event::TransportEvent,
packet::{Packet, PacketError},
protocol::{AdapterStats, TcpClientConfig, TcpServerConfig},
transport::memory_pool::{shared_memory_pool, BufferSize, OptimizedMemoryPool},
SessionId,
};
use async_trait::async_trait;
use bytes::BytesMut;
use std::io;
use std::sync::Arc;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream},
sync::{broadcast, mpsc},
};
fn apply_tcp_keepalive(stream: &TcpStream, duration: std::time::Duration) {
let sock_ref = socket2::SockRef::from(&stream);
let keepalive = socket2::TcpKeepalive::new().with_time(duration);
if let Err(e) = sock_ref.set_tcp_keepalive(&keepalive) {
tracing::warn!("Failed to set TCP keepalive: {}", e);
}
}
#[derive(Debug, thiserror::Error)]
pub enum TcpError {
#[error("IO error: {0}")]
Io(#[from] io::Error),
#[error("Connection timeout")]
Timeout,
#[error("Connection closed")]
ConnectionClosed,
#[error("Packet error: {0}")]
Packet(#[from] PacketError),
#[error("Buffer overflow")]
BufferOverflow,
#[error("Configuration error: {0}")]
Config(String),
}
impl From<TcpError> for TransportError {
fn from(error: TcpError) -> Self {
match error {
TcpError::Io(io_err) => {
TransportError::connection_error(format!("TCP IO error: {:?}", io_err), true)
}
TcpError::Timeout => TransportError::connection_error("TCP connection timeout", true),
TcpError::ConnectionClosed => {
TransportError::connection_error("TCP connection closed", true)
}
TcpError::Packet(packet_err) => TransportError::protocol_error(
"packet",
format!("TCP packet error: {}", packet_err),
),
TcpError::BufferOverflow => {
TransportError::protocol_error("generic", "TCP buffer overflow".to_string())
}
TcpError::Config(msg) => TransportError::config_error("tcp", msg),
}
}
}
const MAX_PAYLOAD_SIZE: usize = 1024 * 1024;
const MAX_EXT_HEADER_SIZE: usize = 64 * 1024;
const MAX_RESYNC_SCAN_DISTANCE: usize = 4096;
const FIXED_HEADER_SIZE: usize = 16;
use crate::adapters::outbound::SEND_QUEUE_CAPACITY;
struct OptimizedReadBuffer {
buffer: BytesMut,
target_capacity: usize,
stats: ReadBufferStats,
pool: Arc<OptimizedMemoryPool>,
buffer_tier: BufferSize,
}
#[derive(Debug, Default)]
struct ReadBufferStats {
reads: u64,
packets_parsed: u64,
reallocations: u64,
bytes_read: u64,
resync_attempts: u64,
bytes_discarded: u64,
}
impl Drop for OptimizedReadBuffer {
fn drop(&mut self) {
if self.buffer.capacity() > 0 {
let mut buf = std::mem::replace(&mut self.buffer, BytesMut::new());
buf.clear();
self.pool.return_buffer(buf, self.buffer_tier);
}
}
}
impl OptimizedReadBuffer {
fn new_with_pool(initial_capacity: usize, pool: Arc<OptimizedMemoryPool>) -> Self {
let buffer_tier = if initial_capacity <= 1024 {
BufferSize::Small
} else if initial_capacity <= 8192 {
BufferSize::Medium
} else {
BufferSize::Large
};
let buffer = pool.get_buffer(buffer_tier);
Self {
buffer,
target_capacity: initial_capacity,
stats: ReadBufferStats::default(),
pool,
buffer_tier,
}
}
fn is_valid_header_at(&self, offset: usize) -> bool {
if self.buffer.len() < offset + FIXED_HEADER_SIZE {
return false;
}
let header = &self.buffer[offset..offset + FIXED_HEADER_SIZE];
let version = header[0];
if version != 1 {
return false;
}
let compression = header[1];
if compression > 2 {
return false;
}
let packet_type = header[2];
if packet_type > 2 {
return false;
}
let ext_header_len = u16::from_be_bytes([header[8], header[9]]) as usize;
if ext_header_len > MAX_EXT_HEADER_SIZE {
return false;
}
let payload_len =
u32::from_be_bytes([header[10], header[11], header[12], header[13]]) as usize;
if payload_len > MAX_PAYLOAD_SIZE {
return false;
}
true
}
fn try_resync_frame(&mut self) -> bool {
self.stats.resync_attempts += 1;
let scan_limit = self.buffer.len().min(MAX_RESYNC_SCAN_DISTANCE);
for offset in 1..scan_limit {
if self.is_valid_header_at(offset) {
tracing::warn!(
"[RESYNC] Frame resync successful, discarded {} bytes",
offset
);
self.stats.bytes_discarded += offset as u64;
let _ = self.buffer.split_to(offset);
return true;
}
}
if self.buffer.len() > MAX_RESYNC_SCAN_DISTANCE {
tracing::warn!(
"[RESYNC] No valid frame found in {} bytes, discarding",
MAX_RESYNC_SCAN_DISTANCE
);
self.stats.bytes_discarded += MAX_RESYNC_SCAN_DISTANCE as u64;
let _ = self.buffer.split_to(MAX_RESYNC_SCAN_DISTANCE);
return true;
}
false
}
fn try_parse_next_packet(&mut self) -> Result<Option<Packet>, TcpError> {
loop {
if self.buffer.len() < FIXED_HEADER_SIZE {
return Ok(None);
}
if !self.is_valid_header_at(0) {
if self.stats.packets_parsed == 0 {
tracing::warn!(
"[PARSE] Invalid protocol header on first packet, closing connection"
);
return Err(TcpError::Config(
"Invalid protocol header on first packet".to_string(),
));
}
tracing::debug!("[PARSE] Invalid header detected, attempting resync");
if !self.try_resync_frame() {
return Err(TcpError::BufferOverflow);
}
continue;
}
let header_bytes = &self.buffer[0..FIXED_HEADER_SIZE];
let ext_header_len = u16::from_be_bytes([header_bytes[8], header_bytes[9]]) as usize;
let payload_len = u32::from_be_bytes([
header_bytes[10],
header_bytes[11],
header_bytes[12],
header_bytes[13],
]) as usize;
let total_packet_len = FIXED_HEADER_SIZE + ext_header_len + payload_len;
if self.buffer.len() < total_packet_len {
return Ok(None);
}
let packet_bytes = self.buffer.split_to(total_packet_len);
match Packet::from_bytes(&packet_bytes) {
Ok(packet) => {
self.stats.packets_parsed += 1;
return Ok(Some(packet));
}
Err(e) => {
tracing::warn!("[PARSE] Packet parse error: {:?}, attempting resync", e);
if !self.try_resync_frame() {
return Err(TcpError::Packet(e));
}
continue;
}
}
}
}
async fn fill_from_stream(
&mut self,
read_half: &mut tokio::net::tcp::OwnedReadHalf,
) -> Result<usize, TcpError> {
if self.buffer.capacity() - self.buffer.len() < 4096 {
self.buffer.reserve(self.target_capacity);
self.stats.reallocations += 1;
}
let bytes_read = read_half
.read_buf(&mut self.buffer)
.await
.map_err(TcpError::Io)?;
self.stats.reads += 1;
self.stats.bytes_read += bytes_read as u64;
Ok(bytes_read)
}
fn stats(&self) -> &ReadBufferStats {
&self.stats
}
fn clear(&mut self) {
self.buffer.clear();
}
}
pub struct TcpAdapter<C> {
state: crate::adapters::core::ConnState,
config: C,
stats: AdapterStats,
connection_info: ConnectionInfo,
send_queue: mpsc::Sender<Packet>,
event_sender: broadcast::Sender<TransportEvent>,
shutdown_sender: mpsc::UnboundedSender<()>,
event_loop_handle: Option<tokio::task::JoinHandle<()>>,
}
impl<C> TcpAdapter<C> {
pub async fn new(
stream: TcpStream,
config: C,
event_sender: broadcast::Sender<TransportEvent>,
) -> Result<Self, TcpError> {
stream.set_nodelay(true)?;
let local_addr = stream.local_addr()?;
let peer_addr = stream.peer_addr()?;
let mut connection_info = ConnectionInfo::default();
connection_info.local_addr = local_addr;
connection_info.peer_addr = peer_addr;
connection_info.protocol = "tcp".to_string();
connection_info.state = ConnectionState::Connected;
connection_info.established_at = std::time::SystemTime::now();
let state =
crate::adapters::core::ConnState::new(crate::adapters::core::ConnStatus::Connected);
let (send_queue_tx, send_queue_rx) = mpsc::channel(SEND_QUEUE_CAPACITY);
let (shutdown_tx, shutdown_rx) = mpsc::unbounded_channel();
let memory_pool = shared_memory_pool();
let event_loop_handle = Self::start_event_loop(
stream,
state.clone(),
send_queue_rx,
shutdown_rx,
event_sender.clone(),
memory_pool,
)
.await;
Ok(Self {
state,
config,
stats: AdapterStats::new(),
connection_info,
send_queue: send_queue_tx,
event_sender,
shutdown_sender: shutdown_tx,
event_loop_handle: Some(event_loop_handle),
})
}
pub fn subscribe_events(&self) -> broadcast::Receiver<TransportEvent> {
self.event_sender.subscribe()
}
async fn start_event_loop(
stream: TcpStream,
state: crate::adapters::core::ConnState,
mut send_queue: mpsc::Receiver<Packet>,
mut shutdown_signal: mpsc::UnboundedReceiver<()>,
event_sender: broadcast::Sender<TransportEvent>,
memory_pool: Arc<OptimizedMemoryPool>,
) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
let current_session_id = state.session_id();
tracing::debug!(
"[START] TCP event loop started (session: {})",
current_session_id
);
let (mut read_half, mut write_half) = stream.into_split();
let mut read_buffer = OptimizedReadBuffer::new_with_pool(8192, memory_pool);
'event_loop: loop {
let current_session_id = state.session_id();
tokio::select! {
read_result = read_buffer.fill_from_stream(&mut read_half) => {
match read_result {
Ok(0) => {
tracing::debug!("[RECV] Peer actively closed TCP connection (session: {})", current_session_id);
let close_event = TransportEvent::ConnectionClosed { reason: crate::error::CloseReason::Normal };
if let Err(e) = event_sender.send(close_event) {
tracing::debug!("[CONNECT] Failed to notify upper layer of connection close: session {} - {:?}", current_session_id, e);
} else {
tracing::debug!("[NOTIFY] Notified upper layer of connection close: session {}", current_session_id);
}
break;
}
Ok(_) => {
loop {
match read_buffer.try_parse_next_packet() {
Ok(Some(packet)) => {
tracing::debug!("[RECV] TCP received packet: {} bytes (session: {})", packet.payload.len(), current_session_id);
tracing::debug!("[DETAIL] Packet details: ID={}, type={:?}, payload_len={}", packet.header.message_id, packet.header.packet_type, packet.payload.len());
let event = TransportEvent::MessageReceived(packet);
if let Err(e) = event_sender.send(event) {
tracing::warn!("[RECV] Failed to send receive event: {:?}", e);
}
}
Ok(None) => break,
Err(e) => {
tracing::error!("[RECV] TCP parse error: {:?} (session: {})", e, current_session_id);
let close_event = TransportEvent::ConnectionClosed {
reason: crate::error::CloseReason::Error(format!("{:?}", e)),
};
let _ = event_sender.send(close_event);
break 'event_loop;
}
}
}
}
Err(e) => {
tracing::error!("[RECV] TCP connection error: {:?} (session: {})", e, current_session_id);
let close_event = TransportEvent::ConnectionClosed { reason: crate::error::CloseReason::Error(format!("{:?}", e)) };
if let Err(e) = event_sender.send(close_event) {
tracing::debug!("[CONNECT] Failed to notify upper layer of connection error: session {} - {:?}", current_session_id, e);
} else {
tracing::debug!("[NOTIFY] Notified upper layer of connection error: session {}", current_session_id);
}
break;
}
}
}
packet = send_queue.recv() => {
if let Some(packet) = packet {
match Self::write_packet_to_stream(&mut write_half, &packet).await {
Ok(_) => {
tracing::debug!("[SEND] TCP send successful: {} bytes (session: {})", packet.payload.len(), current_session_id);
let event = TransportEvent::MessageSent { packet_id: packet.header.message_id };
if let Err(e) = event_sender.send(event) {
tracing::warn!("[SEND] Failed to send send event: {:?}", e);
}
}
Err(e) => {
tracing::error!("[SEND] TCP send error: {:?} (session: {})", e, current_session_id);
let close_event = TransportEvent::ConnectionClosed { reason: crate::error::CloseReason::Error(format!("{:?}", e)) };
if let Err(e) = event_sender.send(close_event) {
tracing::debug!("[CONNECT] Failed to notify upper layer of send error: session {} - {:?}", current_session_id, e);
} else {
tracing::debug!("[NOTIFY] Notified upper layer of send error: session {}", current_session_id);
}
break;
}
}
}
}
_ = shutdown_signal.recv() => {
tracing::info!("[STOP] Received shutdown signal, stopping TCP event loop (session: {})", current_session_id);
tracing::debug!("[CLOSE] Active close, not sending close event");
break;
}
}
}
state.set_status(crate::adapters::core::ConnStatus::Closed);
tracing::debug!(
"[SUCCESS] TCP event loop ended (session: {})",
current_session_id
);
})
}
async fn write_packet_to_stream(
write_half: &mut tokio::net::tcp::OwnedWriteHalf,
packet: &Packet,
) -> Result<(), TcpError> {
let packet_bytes = packet.to_bytes();
write_half
.write_all(&packet_bytes)
.await
.map_err(TcpError::Io)?;
Ok(())
}
}
impl TcpAdapter<TcpClientConfig> {
pub async fn connect(
addr: std::net::SocketAddr,
config: TcpClientConfig,
) -> Result<Self, TcpError> {
tracing::debug!("[CONNECT] TCP client connecting to: {}", addr);
let stream = if config.connect_timeout != std::time::Duration::from_secs(0) {
tokio::time::timeout(config.connect_timeout, TcpStream::connect(addr))
.await
.map_err(|_| TcpError::Timeout)?
.map_err(TcpError::Io)?
} else {
TcpStream::connect(addr).await.map_err(TcpError::Io)?
};
tracing::debug!("[SUCCESS] TCP connection established successfully");
if let Some(keepalive) = config.keepalive {
apply_tcp_keepalive(&stream, keepalive);
}
Self::new(stream, config, broadcast::channel(8192).0).await
}
}
#[async_trait]
impl<C: Send + Sync + 'static> Connection for TcpAdapter<C> {
async fn send(&mut self, packet: Packet) -> Result<(), TransportError> {
crate::adapters::outbound::send_bounded(
&self.send_queue,
packet,
"tcp_outbound_queue",
"TCP connection closed",
)
.await
}
async fn close(&mut self) -> Result<(), TransportError> {
let _ = self.shutdown_sender.send(());
if let Some(handle) = self.event_loop_handle.take() {
let _ = handle.await;
}
self.state
.set_status(crate::adapters::core::ConnStatus::Closed);
self.connection_info.state = ConnectionState::Closed;
self.connection_info.closed_at = Some(std::time::SystemTime::now());
Ok(())
}
fn session_id(&self) -> SessionId {
self.state.session_id()
}
fn set_session_id(&mut self, session_id: SessionId) {
self.state.set_session_id(session_id);
self.connection_info.session_id = session_id;
}
fn connection_info(&self) -> ConnectionInfo {
self.connection_info.clone()
}
fn is_connected(&self) -> bool {
self.state.is_connected()
}
async fn flush(&mut self) -> Result<(), TransportError> {
Ok(())
}
fn event_stream(
&self,
) -> Option<tokio::sync::broadcast::Receiver<crate::event::TransportEvent>> {
Some(self.event_sender.subscribe())
}
}
pub(crate) struct TcpServerBuilder {
config: TcpServerConfig,
bind_address: Option<std::net::SocketAddr>,
}
impl TcpServerBuilder {
pub(crate) fn new() -> Self {
Self {
config: TcpServerConfig::default(),
bind_address: None,
}
}
pub(crate) fn bind_address(mut self, addr: std::net::SocketAddr) -> Self {
self.bind_address = Some(addr);
self
}
pub(crate) fn config(mut self, config: TcpServerConfig) -> Self {
self.config = config;
self
}
pub(crate) async fn build(self) -> Result<TcpServer, TcpError> {
let bind_addr = self.bind_address.unwrap_or(self.config.bind_address);
tracing::debug!("[START] TCP server starting on: {}", bind_addr);
let listener = TcpListener::bind(bind_addr).await?;
tracing::info!(
"[SUCCESS] TCP server successfully started on: {}",
listener.local_addr()?
);
Ok(TcpServer {
listener: Some(listener),
config: self.config,
})
}
}
impl Default for TcpServerBuilder {
fn default() -> Self {
Self::new()
}
}
pub(crate) struct TcpServer {
listener: Option<TcpListener>,
config: TcpServerConfig,
}
impl TcpServer {
pub(crate) fn builder() -> TcpServerBuilder {
TcpServerBuilder::new()
}
pub(crate) async fn accept(&mut self) -> Result<TcpAdapter<TcpServerConfig>, TcpError> {
let listener = self
.listener
.as_mut()
.ok_or_else(|| TcpError::Config("TCP server is shut down".to_string()))?;
let (stream, peer_addr) = listener.accept().await?;
tracing::debug!("[CONNECT] TCP new connection from: {}", peer_addr);
if let Some(keepalive) = self.config.keepalive {
apply_tcp_keepalive(&stream, keepalive);
}
TcpAdapter::new(stream, self.config.clone(), broadcast::channel(8192).0).await
}
pub(crate) fn local_addr(&self) -> Result<std::net::SocketAddr, TcpError> {
let listener = self
.listener
.as_ref()
.ok_or_else(|| TcpError::Config("TCP server is shut down".to_string()))?;
Ok(listener.local_addr()?)
}
pub(crate) async fn shutdown(&mut self) -> Result<(), TcpError> {
self.listener.take();
Ok(())
}
}
pub(crate) struct TcpClientBuilder {
config: TcpClientConfig,
target_address: Option<std::net::SocketAddr>,
}
impl TcpClientBuilder {
pub(crate) fn new() -> Self {
Self {
config: TcpClientConfig::default(),
target_address: None,
}
}
pub(crate) fn target_address(mut self, addr: std::net::SocketAddr) -> Self {
self.target_address = Some(addr);
self
}
pub(crate) fn config(mut self, config: TcpClientConfig) -> Self {
self.config = config;
self
}
pub(crate) async fn connect(self) -> Result<TcpAdapter<TcpClientConfig>, TcpError> {
let target_addr = self.target_address.unwrap_or(self.config.target_address);
TcpAdapter::connect(target_addr, self.config).await
}
}
impl Default for TcpClientBuilder {
fn default() -> Self {
Self::new()
}
}