use std::io;
use std::net::SocketAddr;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use rsip::SipMessage;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf};
use tokio::io::{ReadHalf, WriteHalf};
use tokio::net::TcpStream;
use tokio::sync::{mpsc, Mutex};
use tokio::task::JoinHandle;
use tracing::{debug, trace, warn};
use super::framing::Framer;
use super::transaction::Reliability;
use super::transport::serialize;
pub(crate) enum SipStream {
Tcp(TcpStream),
#[cfg(feature = "tls")]
Tls(Box<tokio_rustls::client::TlsStream<TcpStream>>),
}
impl AsyncRead for SipStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
match self.get_mut() {
Self::Tcp(s) => Pin::new(s).poll_read(cx, buf),
#[cfg(feature = "tls")]
Self::Tls(s) => Pin::new(s.as_mut()).poll_read(cx, buf),
}
}
}
impl AsyncWrite for SipStream {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
match self.get_mut() {
Self::Tcp(s) => Pin::new(s).poll_write(cx, buf),
#[cfg(feature = "tls")]
Self::Tls(s) => Pin::new(s.as_mut()).poll_write(cx, buf),
}
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.get_mut() {
Self::Tcp(s) => Pin::new(s).poll_flush(cx),
#[cfg(feature = "tls")]
Self::Tls(s) => Pin::new(s.as_mut()).poll_flush(cx),
}
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.get_mut() {
Self::Tcp(s) => Pin::new(s).poll_shutdown(cx),
#[cfg(feature = "tls")]
Self::Tls(s) => Pin::new(s.as_mut()).poll_shutdown(cx),
}
}
}
const READ_CHUNK: usize = 8 * 1024;
const INBOUND_DEPTH: usize = 64;
pub(crate) struct StreamTransport {
write: Arc<Mutex<WriteHalf<SipStream>>>,
inbound: Mutex<mpsc::Receiver<(SipMessage, SocketAddr)>>,
reader: JoinHandle<()>,
local: SocketAddr,
peer: SocketAddr,
}
impl Drop for StreamTransport {
fn drop(&mut self) {
self.reader.abort();
}
}
impl StreamTransport {
pub(crate) async fn connect(
peer: SocketAddr,
setup: &super::transport::TransportSetup,
) -> io::Result<Self> {
let sock = TcpStream::connect(peer).await?;
sock.set_nodelay(true)?;
let local = sock.local_addr()?;
let stream = match setup.transport {
#[cfg(feature = "tls")]
crate::account::Transport::Tls => {
let tls = setup.tls.as_ref().ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidInput,
"TLS transport selected without a TLS setup",
)
})?;
SipStream::Tls(Box::new(
super::tls::connect(sock, &tls.server_name, &tls.policy).await?,
))
}
#[cfg(not(feature = "tls"))]
crate::account::Transport::Tls => {
return Err(io::Error::new(
io::ErrorKind::Unsupported,
"TLS transport requires the `tls` feature",
))
}
_ => SipStream::Tcp(sock),
};
let (read, write) = tokio::io::split(stream);
let (tx, rx) = mpsc::channel(INBOUND_DEPTH);
let reader = tokio::spawn(read_task(read, peer, tx));
Ok(Self {
write: Arc::new(Mutex::new(write)),
inbound: Mutex::new(rx),
reader,
local,
peer,
})
}
pub(crate) fn local_addr(&self) -> io::Result<SocketAddr> {
Ok(self.local)
}
pub(crate) fn reliability(&self) -> Reliability {
Reliability::Reliable
}
pub(crate) async fn send_to(&self, msg: &SipMessage, dst: SocketAddr) -> io::Result<()> {
if dst != self.peer {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("stream transport is connected to {}, not {dst}", self.peer),
));
}
let bytes = serialize(msg);
debug!(
dst = %self.peer,
bytes = bytes.len(),
"\n>>> SEND to {} >>>\n{}",
self.peer,
String::from_utf8_lossy(&bytes).trim_end(),
);
write_message(&mut *self.write.lock().await, &bytes).await
}
pub(crate) async fn recv(&self) -> io::Result<(SipMessage, SocketAddr)> {
self.inbound.lock().await.recv().await.ok_or_else(|| {
io::Error::new(io::ErrorKind::ConnectionReset, "stream connection closed")
})
}
}
async fn write_message<W: AsyncWrite + Unpin + ?Sized>(w: &mut W, bytes: &[u8]) -> io::Result<()> {
w.write_all(bytes).await?;
w.flush().await
}
async fn read_task(
mut read: ReadHalf<SipStream>,
peer: SocketAddr,
tx: mpsc::Sender<(SipMessage, SocketAddr)>,
) {
let mut framer = Framer::new();
let mut chunk = vec![0u8; READ_CHUNK];
loop {
let n = match read.read(&mut chunk).await {
Ok(0) => {
debug!(%peer, "stream closed by peer");
return;
}
Ok(n) => n,
Err(e) => {
warn!(%peer, error = %e, "stream read failed");
return;
}
};
trace!(%peer, bytes = n, "stream read");
framer.push(&chunk[..n]);
loop {
match framer.next_message() {
Ok(Some(msg)) => {
debug!(
src = %peer,
"\n<<< RECV from {peer} <<<\n{}",
String::from_utf8_lossy(&serialize(&msg)).trim_end(),
);
if tx.send((msg, peer)).await.is_err() {
return;
}
}
Ok(None) => break,
Err(e) => {
warn!(%peer, error = %e, "stream framing failed; dropping connection");
return;
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::account::Transport;
use tokio::net::TcpListener;
fn options_to(dst: SocketAddr) -> SipMessage {
let raw = format!(
"OPTIONS sip:bob@example.com SIP/2.0\r\n\
Via: SIP/2.0/TCP {dst};branch=z9hG4bK-opt\r\n\
From: <sip:alice@example.com>;tag=alice\r\n\
To: <sip:bob@example.com>\r\n\
Call-ID: call-opt\r\n\
CSeq: 4 OPTIONS\r\n\
Content-Length: 0\r\n\r\n"
);
rsip::Request::try_from(raw.as_bytes())
.expect("valid request")
.into()
}
fn ok_response() -> Vec<u8> {
b"SIP/2.0 200 OK\r\n\
Via: SIP/2.0/TCP 127.0.0.1:5060;branch=z9hG4bK-opt\r\n\
From: <sip:alice@example.com>;tag=alice\r\n\
To: <sip:bob@example.com>;tag=bob\r\n\
Call-ID: call-opt\r\n\
CSeq: 4 OPTIONS\r\n\
Content-Length: 0\r\n\r\n"
.to_vec()
}
#[tokio::test]
async fn connects_and_reports_a_local_addr() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
tokio::spawn(async move {
let _ = listener.accept().await;
});
let t = StreamTransport::connect(addr, &Transport::Tcp.into())
.await
.expect("connects");
let local = t.local_addr().expect("local addr");
assert_eq!(local.ip().to_string(), "127.0.0.1");
assert_ne!(local.port(), 0, "the OS assigned a real port");
}
#[tokio::test]
async fn is_always_reliable() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
tokio::spawn(async move {
let _ = listener.accept().await;
});
let t = StreamTransport::connect(addr, &Transport::Tcp.into())
.await
.expect("connects");
assert!(t.reliability().is_reliable());
}
#[tokio::test]
async fn sends_a_message_the_peer_can_read() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
let server = tokio::spawn(async move {
let (mut sock, _) = listener.accept().await.expect("accept");
let mut buf = vec![0u8; 4096];
let n = sock.read(&mut buf).await.expect("read");
String::from_utf8_lossy(&buf[..n]).to_string()
});
let t = StreamTransport::connect(addr, &Transport::Tcp.into())
.await
.expect("connects");
t.send_to(&options_to(addr), addr).await.expect("sends");
let got = server.await.expect("server task");
assert!(
got.starts_with("OPTIONS sip:bob@example.com SIP/2.0\r\n"),
"{got}"
);
}
#[tokio::test]
async fn rejects_a_send_to_an_unconnected_peer() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
tokio::spawn(async move {
let _ = listener.accept().await;
});
let t = StreamTransport::connect(addr, &Transport::Tcp.into())
.await
.expect("connects");
let elsewhere: SocketAddr = "127.0.0.1:1".parse().expect("addr");
assert!(t.send_to(&options_to(addr), elsewhere).await.is_err());
}
#[tokio::test]
async fn receives_a_response_from_the_same_connection() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
tokio::spawn(async move {
let (mut sock, _) = listener.accept().await.expect("accept");
let mut buf = vec![0u8; 4096];
let _ = sock.read(&mut buf).await.expect("read");
sock.write_all(&ok_response()).await.expect("write");
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
});
let t = StreamTransport::connect(addr, &Transport::Tcp.into())
.await
.expect("connects");
t.send_to(&options_to(addr), addr).await.expect("sends");
let (msg, src) = t.recv().await.expect("receives");
assert!(matches!(msg, SipMessage::Response(_)));
assert_eq!(src, addr, "source is the connected peer");
}
#[tokio::test]
async fn receives_two_messages_written_together() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
tokio::spawn(async move {
let (mut sock, _) = listener.accept().await.expect("accept");
let mut both = ok_response();
both.extend_from_slice(&ok_response());
sock.write_all(&both).await.expect("write");
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
});
let t = StreamTransport::connect(addr, &Transport::Tcp.into())
.await
.expect("connects");
assert!(t.recv().await.is_ok(), "first message");
assert!(t.recv().await.is_ok(), "second message");
}
#[derive(Default)]
struct BufferedUntilFlush {
pending: Vec<u8>,
wire: Vec<u8>,
}
impl AsyncWrite for BufferedUntilFlush {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
self.get_mut().pending.extend_from_slice(buf);
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let this = self.get_mut();
this.wire.append(&mut this.pending);
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
self.poll_flush(cx)
}
}
#[tokio::test]
async fn a_sent_message_is_flushed_out_of_a_buffering_stream() {
let mut w = BufferedUntilFlush::default();
write_message(&mut w, b"OPTIONS sip:bob@example.com SIP/2.0\r\n\r\n")
.await
.expect("writes");
assert_eq!(
w.wire, b"OPTIONS sip:bob@example.com SIP/2.0\r\n\r\n",
"the message must reach the wire, not sit in the session buffer"
);
}
#[tokio::test]
async fn dropping_the_transport_closes_an_idle_connection() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
let server = tokio::spawn(async move {
let (mut sock, _) = listener.accept().await.expect("accept");
let mut buf = [0u8; 64];
tokio::time::timeout(std::time::Duration::from_secs(2), sock.read(&mut buf)).await
});
let t = StreamTransport::connect(addr, &Transport::Tcp.into())
.await
.expect("connects");
drop(t);
let read = server
.await
.expect("server task")
.expect("the peer must see the connection close, not wait forever");
assert_eq!(read.expect("read"), 0, "EOF");
}
#[tokio::test]
async fn connect_to_a_closed_port_fails() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
drop(listener);
assert!(StreamTransport::connect(addr, &Transport::Tcp.into())
.await
.is_err());
}
#[cfg(feature = "tls")]
#[tokio::test]
async fn tls_transport_without_a_tls_setup_is_rejected() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
tokio::spawn(async move {
let _ = listener.accept().await;
});
let result = StreamTransport::connect(addr, &Transport::Tls.into()).await;
match result {
Ok(_) => panic!("must not connect without a TLS setup"),
Err(err) => assert_eq!(err.kind(), io::ErrorKind::InvalidInput),
}
}
#[cfg(not(feature = "tls"))]
#[tokio::test]
async fn tls_transport_without_the_feature_is_unsupported() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
tokio::spawn(async move {
let _ = listener.accept().await;
});
let result = StreamTransport::connect(addr, &Transport::Tls.into()).await;
match result {
Ok(_) => panic!("must not connect without the tls feature"),
Err(err) => assert_eq!(err.kind(), io::ErrorKind::Unsupported),
}
}
#[tokio::test]
async fn sip_stream_delegates_reads_and_writes() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
let server = tokio::spawn(async move {
let (mut sock, _) = listener.accept().await.expect("accept");
sock.write_all(b"hello through the enum")
.await
.expect("write");
let mut buf = vec![0u8; 64];
let n = sock.read(&mut buf).await.expect("read");
buf.truncate(n);
buf
});
let client = TcpStream::connect(addr).await.expect("connect");
let mut stream = SipStream::Tcp(client);
let mut buf = vec![0u8; 64];
let n = stream.read(&mut buf).await.expect("read through SipStream");
assert_eq!(&buf[..n], b"hello through the enum");
stream
.write_all(b"back through the enum")
.await
.expect("write through SipStream");
stream.flush().await.expect("flush through SipStream");
let got = server.await.expect("server task");
assert_eq!(got, b"back through the enum");
}
}