use crate::core::Result;
use bytes::{Bytes, BytesMut};
use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Instant;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::{mpsc, Mutex, RwLock};
use tracing::{debug, error, trace, warn};
use crate::transport::keepalive::{
strip_leading_keepalives, KeepAliveConfig, KEEPALIVE_PING, KEEPALIVE_PONG,
};
use crate::transport::traits::{IncomingMessage, OutgoingMessage, TransportProtocol};
#[cfg(test)]
use std::sync::{
atomic::{AtomicU64, Ordering},
LazyLock, Mutex as StdMutex,
};
pub const MAX_TCP_SIZE: usize = 65536;
const INITIAL_BUF_SIZE: usize = 4096;
#[cfg(test)]
static FORCE_ACCEPT_ERROR_V1: AtomicU64 = AtomicU64::new(0);
#[cfg(test)]
static FORCE_ACCEPT_ERROR_V2: AtomicU64 = AtomicU64::new(0);
#[cfg(test)]
static FORCE_ACCEPT_ERROR_V3: AtomicU64 = AtomicU64::new(0);
#[cfg(test)]
static FORCE_READ_ERROR: AtomicU64 = AtomicU64::new(0);
#[cfg(test)]
static FORCE_WRITE_ERROR: AtomicU64 = AtomicU64::new(0);
#[cfg(test)]
static FORCE_BIND_ERROR: AtomicU64 = AtomicU64::new(0);
#[cfg(test)]
static FORCE_LOCAL_ADDR_ERROR: AtomicU64 = AtomicU64::new(0);
#[cfg(test)]
static FORCE_SKIP_CONNECT_INSERT: LazyLock<StdMutex<Option<SocketAddr>>> =
LazyLock::new(|| StdMutex::new(None));
#[cfg(test)]
fn force_accept_error_once() {
FORCE_ACCEPT_ERROR_V1.store(current_thread_id(), Ordering::SeqCst);
}
#[cfg(test)]
fn force_accept_error_other_message_once() {
FORCE_ACCEPT_ERROR_V2.store(current_thread_id(), Ordering::SeqCst);
}
#[cfg(test)]
fn force_accept_error_other_kind_once() {
FORCE_ACCEPT_ERROR_V3.store(current_thread_id(), Ordering::SeqCst);
}
#[cfg(test)]
fn force_read_error_once() {
FORCE_READ_ERROR.store(current_thread_id(), Ordering::SeqCst);
}
#[cfg(test)]
fn force_write_error_once() {
FORCE_WRITE_ERROR.store(current_thread_id(), Ordering::SeqCst);
}
#[cfg(test)]
fn force_bind_error_once() {
FORCE_BIND_ERROR.store(current_thread_id(), Ordering::SeqCst);
}
#[cfg(test)]
fn force_local_addr_error_once() {
FORCE_LOCAL_ADDR_ERROR.store(current_thread_id(), Ordering::SeqCst);
}
#[cfg(test)]
fn force_skip_connect_insert_once(dest: SocketAddr) {
let mut guard = FORCE_SKIP_CONNECT_INSERT.lock().unwrap();
*guard = Some(dest);
}
#[cfg(test)]
fn take_skip_connect_insert_for(dest: SocketAddr) -> bool {
let mut guard = FORCE_SKIP_CONNECT_INSERT.lock().unwrap();
if guard.as_ref() == Some(&dest) {
*guard = None;
true
} else {
false
}
}
#[cfg(test)]
fn try_take(flag: &AtomicU64) -> bool {
let current = current_thread_id();
flag.load(Ordering::SeqCst) == current
&& flag
.compare_exchange(current, 0, Ordering::SeqCst, Ordering::SeqCst)
.is_ok()
}
#[cfg(test)]
fn take_forced_accept_error() -> Option<std::io::Error> {
if try_take(&FORCE_ACCEPT_ERROR_V1) {
Some(std::io::Error::other("forced accept error"))
} else if try_take(&FORCE_ACCEPT_ERROR_V2) {
Some(std::io::Error::other("forced accept error other"))
} else if try_take(&FORCE_ACCEPT_ERROR_V3) {
Some(std::io::Error::new(
std::io::ErrorKind::ConnectionAborted,
"forced accept error",
))
} else {
None
}
}
#[cfg(test)]
fn take_forced_error(flag: &AtomicU64, message: &str) -> Option<std::io::Error> {
if try_take(flag) {
Some(std::io::Error::other(message))
} else {
None
}
}
#[cfg(test)]
fn current_thread_id() -> u64 {
use std::hash::{Hash, Hasher};
let mut hasher = std::collections::hash_map::DefaultHasher::new();
std::thread::current().id().hash(&mut hasher);
normalize_thread_id(hasher.finish())
}
#[cfg(test)]
fn normalize_thread_id(id: u64) -> u64 {
if id == 0 {
1
} else {
id
}
}
struct TcpConnection {
stream: TcpStream,
#[allow(dead_code)]
remote_addr: SocketAddr,
read_buf: BytesMut,
last_ping_sent: Option<Instant>,
}
impl TcpConnection {
fn new(stream: TcpStream, remote_addr: SocketAddr) -> Self {
Self {
stream,
remote_addr,
read_buf: BytesMut::with_capacity(INITIAL_BUF_SIZE),
last_ping_sent: None,
}
}
async fn read_message(&mut self, keepalive: &KeepAliveConfig) -> Result<Option<Bytes>> {
loop {
let pings = strip_leading_keepalives(&mut self.read_buf);
for _ in 0..pings {
stream_write_all(&mut self.stream, KEEPALIVE_PONG).await?;
trace!("Sent CRLF pong to {}", self.remote_addr);
}
if let Some(msg) = self.try_parse_message() {
return Ok(Some(msg));
}
let mut temp_buf = [0u8; 4096];
let read_result = if keepalive.send_pings {
let last = self.last_ping_sent.unwrap_or_else(Instant::now);
let elapsed = last.elapsed();
let remaining = keepalive
.ping_interval
.checked_sub(elapsed)
.unwrap_or_default();
tokio::time::timeout(remaining, stream_read(&mut self.stream, &mut temp_buf))
.await
.ok()
} else {
Some(stream_read(&mut self.stream, &mut temp_buf).await)
};
let n = match read_result {
Some(Ok(n)) => n,
Some(Err(e)) => return Err(e.into()),
None => {
stream_write_all(&mut self.stream, KEEPALIVE_PING).await?;
self.last_ping_sent = Some(Instant::now());
trace!("Sent CRLF ping to {}", self.remote_addr);
continue;
}
};
if n == 0 {
if self.read_buf.is_empty() {
return Ok(None);
}
return Ok(None);
}
self.read_buf.extend_from_slice(&temp_buf[..n]);
if self.read_buf.len() > MAX_TCP_SIZE {
return Err(crate::core::TransportError::MessageTooLarge {
size: self.read_buf.len(),
max: MAX_TCP_SIZE,
}
.into());
}
}
}
fn try_parse_message(&mut self) -> Option<Bytes> {
let data = &self.read_buf[..];
let header_end = find_header_end(data)?;
let headers = &data[..header_end];
let content_length = parse_content_length(headers);
let total_length = header_end + content_length;
if data.len() < total_length {
return None;
}
let msg = self.read_buf.split_to(total_length).freeze();
Some(msg)
}
async fn write_message(&mut self, data: &[u8]) -> Result<()> {
stream_write_all(&mut self.stream, data).await?;
Ok(())
}
}
async fn stream_read(stream: &mut TcpStream, buf: &mut [u8]) -> std::io::Result<usize> {
#[cfg(test)]
if let Some(err) = take_forced_error(&FORCE_READ_ERROR, "forced read error") {
return Err(err);
}
stream.read(buf).await
}
async fn stream_write_all(stream: &mut TcpStream, data: &[u8]) -> std::io::Result<()> {
#[cfg(test)]
if let Some(err) = take_forced_error(&FORCE_WRITE_ERROR, "forced write error") {
return Err(err);
}
stream.write_all(data).await
}
async fn bind_listener(addr: SocketAddr) -> std::io::Result<TcpListener> {
#[cfg(test)]
if let Some(err) = take_forced_error(&FORCE_BIND_ERROR, "forced bind error") {
return Err(err);
}
TcpListener::bind(addr).await
}
fn listener_local_addr(listener: &TcpListener) -> std::io::Result<SocketAddr> {
#[cfg(test)]
if let Some(err) = take_forced_error(&FORCE_LOCAL_ADDR_ERROR, "forced local_addr error") {
return Err(err);
}
listener.local_addr()
}
fn find_header_end(data: &[u8]) -> Option<usize> {
for i in 0..data.len().saturating_sub(3) {
if &data[i..i + 4] == b"\r\n\r\n" {
return Some(i + 4);
}
}
None
}
fn parse_content_length(headers: &[u8]) -> usize {
let headers_str = match std::str::from_utf8(headers) {
Ok(s) => s,
Err(_) => return 0,
};
for line in headers_str.lines() {
let line_lower = line.to_lowercase();
if line_lower.starts_with("content-length:") || line_lower.starts_with("l:") {
let value = line.split_once(':').map(|(_, value)| value).unwrap_or("");
if let Ok(len) = value.trim().parse() {
return len;
}
}
}
0
}
pub struct TcpTransport {
local_addr: SocketAddr,
listener: Option<TcpListener>,
connections: Arc<RwLock<HashMap<SocketAddr, Arc<Mutex<TcpConnection>>>>>,
incoming_tx: Option<mpsc::Sender<IncomingMessage>>,
keepalive: KeepAliveConfig,
}
impl TcpTransport {
pub async fn bind(addr: SocketAddr) -> Result<Self> {
let listener = bind_listener(addr).await?;
let local_addr = listener_local_addr(&listener)?;
debug!("TCP transport bound to {}", local_addr);
Ok(Self {
local_addr,
listener: Some(listener),
connections: Arc::new(RwLock::new(HashMap::new())),
incoming_tx: None,
keepalive: KeepAliveConfig::default(),
})
}
pub fn new_client(local_addr: SocketAddr) -> Self {
Self {
local_addr,
listener: None,
connections: Arc::new(RwLock::new(HashMap::new())),
incoming_tx: None,
keepalive: KeepAliveConfig::default(),
}
}
pub fn with_keepalive(mut self, keepalive: KeepAliveConfig) -> Self {
self.keepalive = keepalive;
self
}
pub fn local_addr(&self) -> SocketAddr {
self.local_addr
}
pub async fn connect(&self, addr: SocketAddr) -> Result<()> {
{
let connections = self.connections.read().await;
if connections.contains_key(&addr) {
return Ok(());
}
}
debug!("Connecting to {}", addr);
let stream = TcpStream::connect(addr).await?;
let conn = TcpConnection::new(stream, addr);
let mut connections = self.connections.write().await;
#[cfg(test)]
if take_skip_connect_insert_for(addr) {
return Ok(());
}
connections.insert(addr, Arc::new(Mutex::new(conn)));
Ok(())
}
pub async fn send(&self, msg: OutgoingMessage) -> Result<()> {
let dest = msg.destination;
self.connect(dest).await?;
let conn_arc = {
let connections = self.connections.read().await;
connections.get(&dest).cloned()
}
.ok_or(crate::core::TransportError::ConnectionClosed)?;
let mut conn = conn_arc.lock().await;
trace!("Sending {} bytes to {} over TCP", msg.data.len(), dest);
conn.write_message(&msg.data).await?;
Ok(())
}
pub async fn send_to(&self, data: &[u8], dest: SocketAddr) -> Result<()> {
self.send(OutgoingMessage::new(Bytes::copy_from_slice(data), dest))
.await
}
pub fn start(mut self) -> (mpsc::Receiver<IncomingMessage>, TcpSender) {
let (tx, rx) = mpsc::channel(256);
self.incoming_tx = Some(tx.clone());
let connections = self.connections.clone();
let listener = self.listener.take();
let keepalive = self.keepalive.clone();
if let Some(listener) = listener {
let tx_clone = tx.clone();
let connections_clone = connections.clone();
let keepalive_clone = keepalive.clone();
tokio::spawn(async move {
loop {
let accept_result = async {
#[cfg(test)]
{
if let Some(err) = take_forced_accept_error() {
return Err(err);
}
}
listener.accept().await
}
.await;
match accept_result {
Ok((stream, remote_addr)) => {
debug!("Accepted connection from {}", remote_addr);
let conn = TcpConnection::new(stream, remote_addr);
let conn_arc = Arc::new(Mutex::new(conn));
{
let mut conns = connections_clone.write().await;
conns.insert(remote_addr, conn_arc.clone());
}
let tx = tx_clone.clone();
let conns = connections_clone.clone();
let ka = keepalive_clone.clone();
tokio::spawn(async move {
Self::read_loop(conn_arc, remote_addr, tx, conns, ka).await;
});
}
Err(e) => {
error!("TCP accept error: {}", e);
#[cfg(test)]
if e.kind() == std::io::ErrorKind::Other
&& e.to_string() == "forced accept error"
{
break;
}
}
}
}
});
}
let tx_clone = tx;
let connections_clone = connections.clone();
let keepalive_existing = keepalive;
tokio::spawn(async move {
let conns = connections_clone.read().await;
for (addr, conn_arc) in conns.iter() {
let tx = tx_clone.clone();
let addr = *addr;
let conn_arc = conn_arc.clone();
let conns = connections_clone.clone();
let ka = keepalive_existing.clone();
tokio::spawn(async move {
Self::read_loop(conn_arc, addr, tx, conns, ka).await;
});
}
});
let sender = TcpSender {
connections: self.connections.clone(),
};
(rx, sender)
}
async fn read_loop(
conn_arc: Arc<Mutex<TcpConnection>>,
remote_addr: SocketAddr,
tx: mpsc::Sender<IncomingMessage>,
connections: Arc<RwLock<HashMap<SocketAddr, Arc<Mutex<TcpConnection>>>>>,
keepalive: KeepAliveConfig,
) {
loop {
let result = {
let mut conn = conn_arc.lock().await;
conn.read_message(&keepalive).await
};
match result {
Ok(Some(data)) => {
trace!(
"Received {} bytes from {} over TCP",
data.len(),
remote_addr
);
let msg = IncomingMessage {
data,
source: remote_addr,
transport: TransportProtocol::Tcp,
};
if tx.send(msg).await.is_err() {
debug!("Receiver dropped, stopping TCP read loop");
break;
}
}
Ok(None) => {
debug!("Connection closed by {}", remote_addr);
break;
}
Err(e) => {
warn!("TCP read error from {}: {}", remote_addr, e);
break;
}
}
}
let mut conns = connections.write().await;
conns.remove(&remote_addr);
debug!("Removed connection to {}", remote_addr);
}
pub fn sender(&self) -> TcpSender {
TcpSender {
connections: self.connections.clone(),
}
}
}
#[derive(Clone)]
pub struct TcpSender {
connections: Arc<RwLock<HashMap<SocketAddr, Arc<Mutex<TcpConnection>>>>>,
}
impl TcpSender {
pub async fn send(&self, msg: OutgoingMessage) -> Result<()> {
let dest = msg.destination;
let conn_arc = {
let connections = self.connections.read().await;
connections.get(&dest).cloned()
};
if let Some(conn_arc) = conn_arc {
let mut conn = conn_arc.lock().await;
trace!("Sending {} bytes to {} over TCP", msg.data.len(), dest);
conn.write_message(&msg.data).await?;
Ok(())
} else {
Err(crate::core::TransportError::ConnectionClosed.into())
}
}
pub async fn send_to(&self, data: &[u8], dest: SocketAddr) -> Result<()> {
self.send(OutgoingMessage::new(Bytes::copy_from_slice(data), dest))
.await
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::{IpAddr, Ipv4Addr};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Once};
fn init_tracing() {
static INIT: Once = Once::new();
INIT.call_once(|| {
let _ = tracing_subscriber::fmt()
.with_max_level(tracing::Level::TRACE)
.with_test_writer()
.try_init();
});
}
async fn wait_for_forced_accept_error_clear_inner(flag: &AtomicU64) {
for _ in 0..20 {
if flag.load(Ordering::SeqCst) == 0 {
return;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
panic!("forced accept error not consumed");
}
async fn wait_for_forced_accept_error_v1_clear() {
wait_for_forced_accept_error_clear_inner(&FORCE_ACCEPT_ERROR_V1).await;
}
async fn wait_for_forced_accept_error_v2_clear() {
wait_for_forced_accept_error_clear_inner(&FORCE_ACCEPT_ERROR_V2).await;
}
async fn wait_for_forced_accept_error_v3_clear() {
wait_for_forced_accept_error_clear_inner(&FORCE_ACCEPT_ERROR_V3).await;
}
#[test]
fn test_max_tcp_size() {
assert_eq!(MAX_TCP_SIZE, 65536);
}
#[test]
fn test_initial_buf_size() {
assert_eq!(INITIAL_BUF_SIZE, 4096);
}
#[test]
fn test_find_header_end() {
let data = b"INVITE sip:test SIP/2.0\r\nContent-Length: 0\r\n\r\n";
assert_eq!(find_header_end(data), Some(46));
let data = b"INVITE sip:test SIP/2.0\r\nContent-Length: 0\r\n";
assert_eq!(find_header_end(data), None);
}
#[test]
fn test_find_header_end_empty() {
let data = b"";
assert_eq!(find_header_end(data), None);
}
#[test]
fn test_find_header_end_only_crlf() {
let data = b"\r\n\r\n";
assert_eq!(find_header_end(data), Some(4));
}
#[test]
fn test_find_header_end_partial_crlf() {
let data = b"\r\n\r";
assert_eq!(find_header_end(data), None);
}
#[test]
fn test_find_header_end_multiple_crlf() {
let data = b"\r\n\r\nmore data\r\n\r\n";
assert_eq!(find_header_end(data), Some(4));
}
#[test]
fn test_find_header_end_short_data() {
let data = b"abc";
assert_eq!(find_header_end(data), None);
}
#[test]
fn test_parse_content_length() {
let headers = b"INVITE sip:test SIP/2.0\r\nContent-Length: 123\r\n\r\n";
assert_eq!(parse_content_length(headers), 123);
let headers = b"INVITE sip:test SIP/2.0\r\nl: 456\r\n\r\n";
assert_eq!(parse_content_length(headers), 456);
let headers = b"INVITE sip:test SIP/2.0\r\n\r\n";
assert_eq!(parse_content_length(headers), 0);
}
#[test]
fn test_parse_content_length_with_spaces() {
let headers = b"INVITE sip:test SIP/2.0\r\nContent-Length: 789 \r\n\r\n";
assert_eq!(parse_content_length(headers), 789);
}
#[test]
fn test_parse_content_length_uppercase() {
let headers = b"INVITE sip:test SIP/2.0\r\nCONTENT-LENGTH: 100\r\n\r\n";
assert_eq!(parse_content_length(headers), 100);
}
#[test]
fn test_parse_content_length_mixed_case() {
let headers = b"INVITE sip:test SIP/2.0\r\nContent-length: 200\r\n\r\n";
assert_eq!(parse_content_length(headers), 200);
}
#[test]
fn test_parse_content_length_short_form_uppercase() {
let headers = b"INVITE sip:test SIP/2.0\r\nL: 300\r\n\r\n";
assert_eq!(parse_content_length(headers), 300);
}
#[test]
fn test_parse_content_length_invalid_value() {
let headers = b"INVITE sip:test SIP/2.0\r\nContent-Length: invalid\r\n\r\n";
assert_eq!(parse_content_length(headers), 0);
}
#[test]
fn test_parse_content_length_empty() {
let headers = b"";
assert_eq!(parse_content_length(headers), 0);
}
#[test]
fn test_parse_content_length_invalid_utf8() {
let headers = &[0xFF, 0xFE, 0x00, 0x01];
assert_eq!(parse_content_length(headers), 0);
}
#[test]
fn test_parse_content_length_multiple_headers() {
let headers =
b"INVITE sip:test SIP/2.0\r\nContent-Length: 100\r\nContent-Length: 200\r\n\r\n";
assert_eq!(parse_content_length(headers), 100);
}
#[test]
fn test_parse_content_length_no_colon_value() {
let headers = b"INVITE sip:test SIP/2.0\r\nContent-Length:\r\n\r\n";
assert_eq!(parse_content_length(headers), 0);
}
#[tokio::test]
async fn test_tcp_bind() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let transport = TcpTransport::bind(addr).await.unwrap();
assert_ne!(transport.local_addr().port(), 0);
}
#[tokio::test]
async fn test_tcp_bind_ipv6() {
let addr = SocketAddr::new(IpAddr::V6("::1".parse().unwrap()), 0);
let transport = TcpTransport::bind(addr).await.unwrap();
assert!(transport.local_addr().is_ipv6());
}
#[tokio::test]
async fn test_tcp_bind_forced_error() {
force_bind_error_once();
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let result = TcpTransport::bind(addr).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_tcp_bind_forced_local_addr_error() {
force_local_addr_error_once();
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let result = TcpTransport::bind(addr).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_tcp_connection_read_and_write() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let client_task = tokio::spawn(async move {
let mut stream = TcpStream::connect(server_addr).await.unwrap();
let msg = b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
stream.write_all(msg).await.unwrap();
let mut buf = [0u8; 64];
let n = stream.read(&mut buf).await.unwrap();
String::from_utf8_lossy(&buf[..n]).to_string()
});
let (stream, remote) = listener.accept().await.unwrap();
let mut conn = TcpConnection::new(stream, remote);
let msg = conn
.read_message(&KeepAliveConfig::default())
.await
.unwrap()
.unwrap();
assert!(msg.starts_with(b"INVITE"));
conn.write_message(b"PONG").await.unwrap();
let received = client_task.await.unwrap();
assert!(received.contains("PONG"));
}
#[tokio::test]
async fn test_tcp_connection_forced_read_error() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let client_task =
tokio::spawn(async move { TcpStream::connect(server_addr).await.unwrap() });
let (stream, remote) = listener.accept().await.unwrap();
let mut conn = TcpConnection::new(stream, remote);
let _client = client_task.await.unwrap();
force_read_error_once();
let result = conn.read_message(&KeepAliveConfig::default()).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_tcp_connection_forced_write_error() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let client_task = tokio::spawn(async move {
let _stream = TcpStream::connect(server_addr).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
});
let (stream, remote) = listener.accept().await.unwrap();
let mut conn = TcpConnection::new(stream, remote);
force_write_error_once();
let result = conn.write_message(b"PING").await;
assert!(result.is_err());
client_task.await.unwrap();
}
#[tokio::test]
async fn test_tcp_sender_send_existing_connection() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let client_task = tokio::spawn(async move {
let mut client = TcpStream::connect(server_addr).await.unwrap();
let mut buf = [0u8; 16];
let n =
tokio::time::timeout(std::time::Duration::from_millis(500), client.read(&mut buf))
.await
.unwrap()
.unwrap();
buf[..n].to_vec()
});
let (stream, remote_addr) = listener.accept().await.unwrap();
let mut map = HashMap::new();
map.insert(
remote_addr,
Arc::new(Mutex::new(TcpConnection::new(stream, remote_addr))),
);
let sender = TcpSender {
connections: Arc::new(RwLock::new(map)),
};
let msg = OutgoingMessage::new(Bytes::from_static(b"PING"), remote_addr);
sender.send(msg).await.unwrap();
let data = client_task.await.unwrap();
assert_eq!(data, b"PING");
}
#[tokio::test]
async fn test_tcp_accept_error_logged() {
init_tracing();
force_accept_error_once();
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let (_rx, _sender) = server.start();
wait_for_forced_accept_error_v1_clear().await;
}
#[tokio::test]
async fn test_tcp_accept_error_other_message() {
init_tracing();
force_accept_error_other_message_once();
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let (_rx, _sender) = server.start();
wait_for_forced_accept_error_v2_clear().await;
}
#[tokio::test]
async fn test_tcp_accept_error_other_kind() {
init_tracing();
force_accept_error_other_kind_once();
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let (_rx, _sender) = server.start();
wait_for_forced_accept_error_v3_clear().await;
}
#[tokio::test]
async fn test_wait_for_forced_accept_error_clear_panics() {
let flag = Arc::new(AtomicU64::new(1));
let flag_handle = flag.clone();
let handle = tokio::spawn(async move {
wait_for_forced_accept_error_clear_inner(&flag_handle).await;
});
let result = handle.await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_tcp_start_spawns_read_loop_for_existing_connections() {
init_tracing();
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server_task = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let msg = b"OPTIONS sip:test@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
stream.write_all(msg).await.unwrap();
});
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let transport = TcpTransport::new_client(client_addr);
transport.connect(server_addr).await.unwrap();
let (mut rx, _sender) = transport.start();
let received = tokio::time::timeout(std::time::Duration::from_millis(200), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(received.transport, TransportProtocol::Tcp);
let _ = server_task.await;
}
#[test]
fn test_tcp_new_client() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 5060);
let transport = TcpTransport::new_client(addr);
assert_eq!(transport.local_addr(), addr);
assert!(transport.listener.is_none());
}
#[tokio::test]
async fn test_tcp_local_addr() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let transport = TcpTransport::bind(addr).await.unwrap();
let local = transport.local_addr();
assert!(local.port() > 0);
assert_eq!(local.ip(), IpAddr::V4(Ipv4Addr::LOCALHOST));
}
#[tokio::test]
async fn test_tcp_sender() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let transport = TcpTransport::bind(addr).await.unwrap();
let _sender = transport.sender();
}
#[tokio::test]
async fn test_tcp_sender_clone() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let transport = TcpTransport::bind(addr).await.unwrap();
let sender1 = transport.sender();
let _sender2 = sender1.clone();
}
#[tokio::test]
async fn test_tcp_sender_no_connection() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let transport = TcpTransport::bind(addr).await.unwrap();
let sender = transport.sender();
let dest = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 9999);
let msg = OutgoingMessage::new(Bytes::from_static(b"test"), dest);
let result = sender.send(msg).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_tcp_send_missing_connection() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server_handle = tokio::spawn(async move {
let _ = listener.accept().await;
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
});
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
force_skip_connect_insert_once(server_addr);
let msg = OutgoingMessage::new(Bytes::from_static(b"test"), server_addr);
let result = client.send(msg).await;
assert!(result.is_err());
let _ = server_handle.await;
}
#[tokio::test]
async fn test_tcp_send_forced_write_error() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (_rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
client.connect(server_addr).await.unwrap();
force_write_error_once();
let msg = OutgoingMessage::new(Bytes::from_static(b"test"), server_addr);
let result = client.send(msg).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_tcp_sender_forced_write_error() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (_rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
client.connect(server_addr).await.unwrap();
let sender = client.sender();
force_write_error_once();
let msg = OutgoingMessage::new(Bytes::from_static(b"test"), server_addr);
let result = sender.send(msg).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_tcp_connect_and_send() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
let msg = b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
client.send_to(msg, server_addr).await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(&received.data[..], msg);
assert_eq!(received.transport, TransportProtocol::Tcp);
}
#[tokio::test]
async fn test_tcp_connect_already_connected() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (_rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
client.connect(server_addr).await.unwrap();
client.connect(server_addr).await.unwrap();
}
#[tokio::test]
async fn test_tcp_send_with_body() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
let msg = b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 11\r\n\r\nHello World";
client.send_to(msg, server_addr).await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(&received.data[..], msg);
}
#[tokio::test]
async fn test_tcp_send_multiple_messages() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
for i in 0..3 {
let msg = format!("MESSAGE sip:test{} SIP/2.0\r\nContent-Length: 0\r\n\r\n", i);
client.send_to(msg.as_bytes(), server_addr).await.unwrap();
}
for _ in 0..3 {
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert!(received.data.starts_with(b"MESSAGE"));
}
}
#[tokio::test]
async fn test_tcp_connect_fail() {
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
let bad_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 1);
let result = client.connect(bad_addr).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_tcp_start_client_only() {
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 5060);
let transport = TcpTransport::new_client(addr);
let (_rx, _sender) = transport.start();
}
#[tokio::test]
async fn test_tcp_bidirectional() {
init_tracing();
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut server_rx, _server_sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
let (_client_rx, _client_sender) = client.start();
let client_transport =
TcpTransport::bind(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0))
.await
.unwrap();
client_transport
.send_to(b"PING\r\n\r\n", server_addr)
.await
.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(1), server_rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(&received.data[..], b"PING\r\n\r\n");
}
#[tokio::test]
async fn test_tcp_connection_new() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let connect_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
let (stream, remote_addr) = listener.accept().await.unwrap();
let conn = TcpConnection::new(stream, remote_addr);
assert_eq!(conn.remote_addr, remote_addr);
assert!(conn.read_buf.is_empty());
connect_handle.await.unwrap();
}
#[test]
fn test_outgoing_message_new() {
let dest = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 5060);
let msg = OutgoingMessage::new(Bytes::from_static(b"test"), dest);
assert_eq!(msg.destination, dest);
assert_eq!(&msg.data[..], b"test");
}
#[tokio::test]
async fn test_tcp_message_too_large() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let mut stream = TcpStream::connect(server_addr).await.unwrap();
let headers = format!(
"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: {}\r\n\r\n",
MAX_TCP_SIZE + 1
);
stream.write_all(headers.as_bytes()).await.unwrap();
let chunk = vec![b'X'; 4096];
for _ in 0..(MAX_TCP_SIZE / 4096 + 2) {
stream.write_all(&chunk).await.unwrap();
}
let result = tokio::time::timeout(std::time::Duration::from_millis(500), rx.recv()).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_tcp_partial_message_on_close() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let mut stream = TcpStream::connect(server_addr).await.unwrap();
stream
.write_all(b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 10\r\n")
.await
.unwrap();
drop(stream);
let result = tokio::time::timeout(std::time::Duration::from_millis(500), rx.recv()).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_tcp_multiple_messages_in_buffer() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let mut stream = TcpStream::connect(server_addr).await.unwrap();
let msg1 = b"INVITE sip:test1@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
let msg2 = b"INVITE sip:test2@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
let msg3 = b"INVITE sip:test3@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
let mut combined = Vec::new();
combined.extend_from_slice(msg1);
combined.extend_from_slice(msg2);
combined.extend_from_slice(msg3);
stream.write_all(&combined).await.unwrap();
stream.flush().await.unwrap();
for i in 1..=3 {
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
let expected = format!("test{}@example.com", i);
assert!(String::from_utf8_lossy(&received.data).contains(&expected));
}
}
#[tokio::test]
async fn test_tcp_message_with_large_body() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
let body = vec![b'X'; 10000];
let headers = format!(
"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: {}\r\n\r\n",
body.len()
);
let mut msg = headers.into_bytes();
msg.extend_from_slice(&body);
client.send_to(&msg, server_addr).await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(2), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(received.data.len(), msg.len());
assert_eq!(&received.data[..], &msg[..]);
}
#[tokio::test]
async fn test_tcp_incremental_message_assembly() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let mut stream = TcpStream::connect(server_addr).await.unwrap();
let msg = b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 5\r\n\r\nHELLO";
for chunk in msg.chunks(5) {
stream.write_all(chunk).await.unwrap();
stream.flush().await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(&received.data[..], msg);
}
#[tokio::test]
async fn test_tcp_message_with_partial_body() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let mut stream = TcpStream::connect(server_addr).await.unwrap();
stream
.write_all(b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 100\r\n\r\n")
.await
.unwrap();
stream.flush().await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
stream.write_all(&[b'A'; 50]).await.unwrap();
stream.flush().await.unwrap();
let result = tokio::time::timeout(std::time::Duration::from_millis(200), rx.recv()).await;
assert!(result.is_err());
stream.write_all(&[b'B'; 50]).await.unwrap();
stream.flush().await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(received.data.len(), 60 + 100); }
#[tokio::test]
async fn test_tcp_sender_send_to() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
client.connect(server_addr).await.unwrap();
let sender = client.sender();
let msg = b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
sender.send_to(msg, server_addr).await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(&received.data[..], msg);
}
#[tokio::test]
async fn test_tcp_connection_close_cleanup() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (_rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
client.connect(server_addr).await.unwrap();
{
let connections = client.connections.read().await;
assert_eq!(connections.len(), 1);
}
drop(client);
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
}
#[tokio::test]
async fn test_tcp_receiver_dropped() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
client.connect(server_addr).await.unwrap();
drop(rx);
let msg = b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
let result = client.send_to(msg, server_addr).await;
assert!(result.is_ok());
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
}
#[tokio::test]
async fn test_tcp_connection_read_write_errors() {
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
let bad_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 1);
let msg = b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
let result = client.send_to(msg, bad_addr).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_tcp_empty_message_after_headers() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
let msg = b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
client.send_to(msg, server_addr).await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(&received.data[..], msg);
}
#[tokio::test]
async fn test_tcp_connection_graceful_close() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let stream = TcpStream::connect(server_addr).await.unwrap();
drop(stream);
let result = tokio::time::timeout(std::time::Duration::from_millis(500), rx.recv()).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_tcp_message_spanning_multiple_reads() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let mut stream = TcpStream::connect(server_addr).await.unwrap();
let body = vec![b'X'; 8192]; let headers = format!(
"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: {}\r\n\r\n",
body.len()
);
let mut msg = headers.into_bytes();
msg.extend_from_slice(&body);
for chunk in msg.chunks(1024) {
stream.write_all(chunk).await.unwrap();
stream.flush().await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(5)).await;
}
let received = tokio::time::timeout(std::time::Duration::from_secs(2), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(received.data.len(), msg.len());
}
#[tokio::test]
async fn test_tcp_compact_header_form() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
let msg = b"INVITE sip:test@example.com SIP/2.0\r\nl: 5\r\n\r\nHELLO";
client.send_to(msg, server_addr).await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(&received.data[..], msg);
}
#[tokio::test]
async fn test_tcp_no_content_length_header() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
let msg = b"INVITE sip:test@example.com SIP/2.0\r\nVia: SIP/2.0/TCP test\r\n\r\n";
client.send_to(msg, server_addr).await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(&received.data[..], msg);
}
#[test]
fn test_parse_content_length_negative() {
let headers = b"INVITE sip:test SIP/2.0\r\nContent-Length: -100\r\n\r\n";
assert_eq!(parse_content_length(headers), 0);
}
#[test]
fn test_parse_content_length_overflow() {
let headers =
b"INVITE sip:test SIP/2.0\r\nContent-Length: 999999999999999999999999\r\n\r\n";
assert_eq!(parse_content_length(headers), 0);
}
#[test]
fn test_find_header_end_exact_boundary() {
let data = b"abc";
assert_eq!(find_header_end(data), None);
let data = b"ab";
assert_eq!(find_header_end(data), None);
let data = b"a";
assert_eq!(find_header_end(data), None);
}
#[tokio::test]
async fn test_tcp_send_after_connection_established() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
let data =
Bytes::from_static(b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n");
let msg = OutgoingMessage::new(data.clone(), server_addr);
client.send(msg).await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(&received.data[..], &data[..]);
}
#[tokio::test]
async fn test_tcp_headers_with_multiple_colons() {
let headers = b"INVITE sip:test@example.com:5060 SIP/2.0\r\nContent-Length: 10\r\n\r\n";
assert_eq!(parse_content_length(headers), 10);
}
#[tokio::test]
async fn test_tcp_connection_buffer_initialization() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let connect_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
let (stream, remote_addr) = listener.accept().await.unwrap();
let conn = TcpConnection::new(stream, remote_addr);
assert_eq!(conn.read_buf.capacity(), INITIAL_BUF_SIZE);
assert_eq!(conn.read_buf.len(), 0);
connect_handle.await.unwrap();
}
#[tokio::test]
async fn test_tcp_start_with_existing_connections() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut server_rx, _server_sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
client.connect(server_addr).await.unwrap();
{
let connections = client.connections.read().await;
assert_eq!(connections.len(), 1);
}
let (_client_rx, _client_sender) = client.start();
let new_client = TcpTransport::bind(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0))
.await
.unwrap();
let msg = b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
new_client.send_to(msg, server_addr).await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(1), server_rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(&received.data[..], msg);
}
#[tokio::test]
async fn test_tcp_write_error_on_closed_connection() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
client.connect(server_addr).await.unwrap();
drop(rx);
drop(_sender);
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
let msg = b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
let _ = client.send_to(msg, server_addr).await;
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
}
#[tokio::test]
async fn test_tcp_multiple_clients_to_server() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let client1_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client1 = TcpTransport::bind(client1_addr).await.unwrap();
let client2_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client2 = TcpTransport::bind(client2_addr).await.unwrap();
let msg1 = b"INVITE sip:client1@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
let msg2 = b"INVITE sip:client2@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
client1.send_to(msg1, server_addr).await.unwrap();
client2.send_to(msg2, server_addr).await.unwrap();
let mut received_count = 0;
for _ in 0..2 {
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
let is_client1 = received.data.windows(7).any(|w| w == b"client1");
let is_client2 = received.data.windows(7).any(|w| w == b"client2");
received_count += usize::from(is_client1) + usize::from(is_client2);
}
assert_eq!(received_count, 2);
}
#[tokio::test]
async fn test_tcp_message_exactly_at_buffer_boundary() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
let body_size = 4000; let body = vec![b'X'; body_size];
let headers = format!(
"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: {}\r\n\r\n",
body_size
);
let mut msg = headers.into_bytes();
msg.extend_from_slice(&body);
client.send_to(&msg, server_addr).await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(2), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(received.data.len(), msg.len());
}
#[tokio::test]
async fn test_tcp_connection_state_after_multiple_messages() {
let server_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = TcpTransport::bind(server_addr).await.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let client_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let client = TcpTransport::bind(client_addr).await.unwrap();
let messages = vec![
b"INVITE sip:test1@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n".to_vec(),
b"INVITE sip:test2@example.com SIP/2.0\r\nContent-Length: 5\r\n\r\nHELLO".to_vec(),
b"INVITE sip:test3@example.com SIP/2.0\r\nContent-Length: 10\r\n\r\nHELLOWORLD"
.to_vec(),
];
for msg in &messages {
client.send_to(msg, server_addr).await.unwrap();
}
for _ in 0..messages.len() {
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert!(!received.data.is_empty());
}
}
#[tokio::test]
async fn test_keepalive_ping_replied_with_pong() {
let server = TcpTransport::bind(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0))
.await
.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let mut client = TcpStream::connect(server_addr).await.unwrap();
client.write_all(b"\r\n\r\n").await.unwrap();
let mut buf = [0u8; 8];
let n = tokio::time::timeout(std::time::Duration::from_secs(1), client.read(&mut buf))
.await
.unwrap()
.unwrap();
assert_eq!(&buf[..n], b"\r\n", "expected CRLF pong");
let msg = b"INVITE sip:test@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
client.write_all(msg).await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(&received.data[..], msg);
}
#[tokio::test]
async fn test_keepalive_outbound_ping_emitted() {
let server = TcpTransport::bind(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0))
.await
.unwrap()
.with_keepalive(KeepAliveConfig::enabled_with_interval(
std::time::Duration::from_millis(50),
));
let server_addr = server.local_addr();
let (_rx, _sender) = server.start();
let mut client = TcpStream::connect(server_addr).await.unwrap();
let mut buf = [0u8; 8];
let n = tokio::time::timeout(std::time::Duration::from_secs(1), client.read(&mut buf))
.await
.unwrap()
.unwrap();
assert_eq!(&buf[..n], b"\r\n\r\n", "expected CRLF-CRLF ping");
}
#[tokio::test]
async fn test_keepalive_leading_crlfs_before_message_are_consumed() {
let server = TcpTransport::bind(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0))
.await
.unwrap();
let server_addr = server.local_addr();
let (mut rx, _sender) = server.start();
let mut client = TcpStream::connect(server_addr).await.unwrap();
let payload =
b"\r\n\r\n\r\n\r\nINVITE sip:test@example.com SIP/2.0\r\nContent-Length: 0\r\n\r\n";
client.write_all(payload).await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert!(received.data.starts_with(b"INVITE"));
let mut buf = [0u8; 8];
let n = tokio::time::timeout(std::time::Duration::from_secs(1), client.read(&mut buf))
.await
.unwrap()
.unwrap();
assert!(n >= 2, "expected at least one CRLF pong, got {} bytes", n);
assert!(&buf[..n].starts_with(b"\r\n"));
}
}