use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio::net::TcpStream;
use rtsp_types::Message;
use crate::client::{ClientEvent, ClientSession};
use crate::error::{Error, Result};
use crate::interleaved::MAGIC;
use crate::server::{ServerEvent, ServerSession};
use crate::transport::Transport;
type Body = Vec<u8>;
pub const RTSP_DEFAULT_PORT: u16 = 554;
pub const RTSPS_DEFAULT_PORT: u16 = 322;
const READ_CHUNK: usize = 8192;
fn io_err(context: &str, e: std::io::Error) -> Error {
Error::Io(format!("{context}: {e}"))
}
#[derive(Debug)]
pub struct AsyncRtspClient<S> {
stream: S,
session: ClientSession,
read_buf: Vec<u8>,
pending_media: std::collections::VecDeque<ClientEvent>,
}
impl AsyncRtspClient<TcpStream> {
pub async fn connect<A: tokio::net::ToSocketAddrs>(addr: A) -> Result<Self> {
let stream = TcpStream::connect(addr)
.await
.map_err(|e| io_err("connect", e))?;
Ok(Self::with_stream(stream, ClientSession::new()))
}
pub async fn connect_with<A: tokio::net::ToSocketAddrs>(
addr: A,
session: ClientSession,
) -> Result<Self> {
let stream = TcpStream::connect(addr)
.await
.map_err(|e| io_err("connect", e))?;
Ok(Self::with_stream(stream, session))
}
}
impl<S> AsyncRtspClient<S>
where
S: AsyncRead + AsyncWrite + Unpin,
{
pub fn with_stream(stream: S, session: ClientSession) -> Self {
AsyncRtspClient {
stream,
session,
read_buf: Vec::new(),
pending_media: std::collections::VecDeque::new(),
}
}
pub fn state(&self) -> crate::SessionState {
self.session.state()
}
pub fn session_id(&self) -> Option<&str> {
self.session.session_id()
}
pub fn session(&self) -> &ClientSession {
&self.session
}
pub async fn options(&mut self, uri: &str) -> Result<ClientEvent> {
let bytes = self.session.options(uri)?;
self.exchange(bytes).await
}
pub async fn describe(&mut self, uri: &str) -> Result<ClientEvent> {
let bytes = self.session.describe(uri)?;
self.exchange(bytes).await
}
pub async fn setup(&mut self, uri: &str, transport: &Transport) -> Result<ClientEvent> {
let bytes = self.session.setup(uri, transport)?;
self.exchange(bytes).await
}
pub async fn play(&mut self, uri: &str) -> Result<ClientEvent> {
let bytes = self.session.play(uri)?;
self.exchange(bytes).await
}
pub async fn pause(&mut self, uri: &str) -> Result<ClientEvent> {
let bytes = self.session.pause(uri)?;
self.exchange(bytes).await
}
pub async fn teardown(&mut self, uri: &str) -> Result<ClientEvent> {
let bytes = self.session.teardown(uri)?;
self.exchange(bytes).await
}
pub async fn get_parameter(&mut self, uri: &str, body: &[u8]) -> Result<ClientEvent> {
let bytes = self.session.get_parameter(uri, body)?;
self.exchange(bytes).await
}
async fn exchange(&mut self, request: Vec<u8>) -> Result<ClientEvent> {
self.stream
.write_all(&request)
.await
.map_err(|e| io_err("write request", e))?;
self.stream.flush().await.map_err(|e| io_err("flush", e))?;
loop {
let events = self
.session
.handle_data(&std::mem::take(&mut self.read_buf))?;
let mut response = None;
for event in events {
match event {
ClientEvent::Response { .. } => response = Some(event),
ClientEvent::AuthRetry { ref request, .. } => {
let retry = request.clone();
self.stream
.write_all(&retry)
.await
.map_err(|e| io_err("write auth retry", e))?;
self.stream.flush().await.map_err(|e| io_err("flush", e))?;
}
ClientEvent::MediaData { .. } => self.pending_media.push_back(event),
}
}
if let Some(response) = response {
return Ok(response);
}
self.fill_from_socket().await?;
}
}
pub async fn recv_interleaved(&mut self) -> Result<Option<ClientEvent>> {
loop {
if let Some(event) = self.pending_media.pop_front() {
return Ok(Some(event));
}
let events = self
.session
.handle_data(&std::mem::take(&mut self.read_buf))?;
for event in events {
if matches!(event, ClientEvent::MediaData { .. }) {
self.pending_media.push_back(event);
}
}
if let Some(event) = self.pending_media.pop_front() {
return Ok(Some(event));
}
let n = self.read_once().await?;
if n == 0 {
return Ok(None);
}
}
}
async fn fill_from_socket(&mut self) -> Result<()> {
let n = self.read_once().await?;
if n == 0 {
return Err(Error::Io("peer closed connection before response".into()));
}
Ok(())
}
async fn read_once(&mut self) -> Result<usize> {
let mut chunk = [0u8; READ_CHUNK];
let n = self
.stream
.read(&mut chunk)
.await
.map_err(|e| io_err("read", e))?;
self.read_buf.extend_from_slice(&chunk[..n]);
Ok(n)
}
}
#[cfg(feature = "tls")]
impl AsyncRtspClient<tokio_rustls::client::TlsStream<TcpStream>> {
pub async fn connect_tls<A: tokio::net::ToSocketAddrs>(
addr: A,
server_name: &str,
config: rustls::ClientConfig,
) -> Result<Self> {
use std::sync::Arc;
use tokio_rustls::TlsConnector;
let tcp = TcpStream::connect(addr)
.await
.map_err(|e| io_err("connect", e))?;
let connector = TlsConnector::from(Arc::new(config));
let dns = rustls::pki_types::ServerName::try_from(server_name.to_string())
.map_err(|e| Error::Tls(format!("invalid server name {server_name:?}: {e}")))?;
let stream = connector
.connect(dns, tcp)
.await
.map_err(|e| io_err("TLS handshake", e))?;
Ok(Self::with_stream(stream, ClientSession::new()))
}
}
#[cfg(feature = "tls")]
pub fn default_tls_client_config() -> rustls::ClientConfig {
let mut roots = rustls::RootCertStore::empty();
roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
rustls::ClientConfig::builder()
.with_root_certificates(roots)
.with_no_client_auth()
}
#[derive(Debug)]
pub struct AsyncRtspServer<S> {
stream: S,
session: ServerSession,
read_buf: Vec<u8>,
}
impl AsyncRtspServer<TcpStream> {
pub fn accept(stream: TcpStream) -> Self {
Self::with_stream(stream, ServerSession::new())
}
pub fn accept_with(stream: TcpStream, session: ServerSession) -> Self {
Self::with_stream(stream, session)
}
}
#[cfg(feature = "tls")]
impl AsyncRtspServer<tokio_rustls::server::TlsStream<TcpStream>> {
pub async fn accept_tls(stream: TcpStream, config: rustls::ServerConfig) -> Result<Self> {
use std::sync::Arc;
use tokio_rustls::TlsAcceptor;
let acceptor = TlsAcceptor::from(Arc::new(config));
let tls = acceptor
.accept(stream)
.await
.map_err(|e| io_err("TLS handshake", e))?;
Ok(Self::with_stream(tls, ServerSession::new()))
}
}
impl<S> AsyncRtspServer<S>
where
S: AsyncRead + AsyncWrite + Unpin,
{
pub fn with_stream(stream: S, session: ServerSession) -> Self {
AsyncRtspServer {
stream,
session,
read_buf: Vec::new(),
}
}
pub fn state(&self) -> crate::SessionState {
self.session.state()
}
pub fn session_id(&self) -> Option<&str> {
self.session.session_id()
}
pub fn stream_mut(&mut self) -> &mut S {
&mut self.stream
}
pub async fn next_request(&mut self) -> Result<Option<Vec<ServerEvent>>> {
loop {
if let Some(consumed) = complete_request_len(&self.read_buf)? {
let request: Vec<u8> = self.read_buf.drain(..consumed).collect();
let (response, events) = self.session.handle_request(&request)?;
self.stream
.write_all(&response)
.await
.map_err(|e| io_err("write response", e))?;
self.stream.flush().await.map_err(|e| io_err("flush", e))?;
return Ok(Some(events));
}
let mut chunk = [0u8; READ_CHUNK];
let n = self
.stream
.read(&mut chunk)
.await
.map_err(|e| io_err("read", e))?;
if n == 0 {
if self.read_buf.is_empty() {
return Ok(None);
}
return Err(Error::Io("peer closed connection mid-request".into()));
}
self.read_buf.extend_from_slice(&chunk[..n]);
}
}
pub async fn send_interleaved(&mut self, channel: u8, payload: &[u8]) -> Result<()> {
let frame = crate::interleaved::InterleavedFrame::new(channel, payload.to_vec());
let bytes = frame.to_bytes()?;
self.stream
.write_all(&bytes)
.await
.map_err(|e| io_err("write interleaved frame", e))?;
self.stream.flush().await.map_err(|e| io_err("flush", e))?;
Ok(())
}
}
fn complete_request_len(buf: &[u8]) -> Result<Option<usize>> {
if buf.is_empty() {
return Ok(None);
}
if buf[0] == MAGIC {
return Err(Error::MessageParse(
"interleaved '$' frame received where a request was expected".into(),
));
}
match Message::<Body>::parse(buf) {
Ok((_, consumed)) => Ok(Some(consumed)),
Err(rtsp_types::ParseError::Incomplete(_)) => Ok(None),
Err(rtsp_types::ParseError::Error) => {
Err(Error::MessageParse("malformed RTSP request".into()))
}
}
}