use std::collections::HashMap;
use rtsp_types::{Message, Method, Request, StatusCode, Version, headers};
use crate::auth::{Authenticator, Credentials};
use crate::error::{Error, Result};
use crate::interleaved::{self, MAGIC};
use crate::state::{SessionState, client_next_state};
use crate::transport::Transport;
type Body = Vec<u8>;
#[derive(Debug, Clone)]
struct Pending {
method: Method,
uri: String,
request: Request<Body>,
auth_retried: bool,
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ClientEvent {
Response {
cseq: u32,
method: Method,
status: StatusCode,
body: Vec<u8>,
},
AuthRetry {
method: Method,
cseq: u32,
request: Vec<u8>,
},
MediaData {
channel: u8,
data: Vec<u8>,
},
}
#[derive(Debug)]
pub struct ClientSession {
state: SessionState,
next_cseq: u32,
session_id: Option<String>,
session_timeout: Option<u64>,
credentials: Option<Credentials>,
authenticator: Option<Authenticator>,
negotiated_transport: Option<Transport>,
pending: HashMap<u32, Pending>,
inbound: Vec<u8>,
user_agent: String,
}
impl Default for ClientSession {
fn default() -> Self {
Self::new()
}
}
impl ClientSession {
pub fn new() -> Self {
ClientSession {
state: SessionState::Init,
next_cseq: 1,
session_id: None,
session_timeout: None,
credentials: None,
authenticator: None,
negotiated_transport: None,
pending: HashMap::new(),
inbound: Vec::new(),
user_agent: "rtsp-runtime".to_string(),
}
}
pub fn with_credentials(mut self, credentials: Credentials) -> Self {
self.credentials = Some(credentials);
self
}
pub fn with_user_agent(mut self, ua: impl Into<String>) -> Self {
self.user_agent = ua.into();
self
}
pub fn state(&self) -> SessionState {
self.state
}
pub fn session_id(&self) -> Option<&str> {
self.session_id.as_deref()
}
pub fn session_timeout(&self) -> Option<u64> {
self.session_timeout
}
pub fn negotiated_transport(&self) -> Option<&Transport> {
self.negotiated_transport.as_ref()
}
pub fn options(&mut self, uri: &str) -> Result<Vec<u8>> {
self.build_request(Method::Options, uri, None, &[])
}
pub fn describe(&mut self, uri: &str) -> Result<Vec<u8>> {
self.build_request(
Method::Describe,
uri,
None,
&[(headers::ACCEPT, "application/sdp".to_string())],
)
}
pub fn setup(&mut self, uri: &str, transport: &Transport) -> Result<Vec<u8>> {
self.build_request(
Method::Setup,
uri,
None,
&[(headers::TRANSPORT, transport.to_header_value())],
)
}
pub fn play(&mut self, uri: &str) -> Result<Vec<u8>> {
self.build_request(Method::Play, uri, None, &[])
}
pub fn pause(&mut self, uri: &str) -> Result<Vec<u8>> {
self.build_request(Method::Pause, uri, None, &[])
}
pub fn teardown(&mut self, uri: &str) -> Result<Vec<u8>> {
self.build_request(Method::Teardown, uri, None, &[])
}
pub fn get_parameter(&mut self, uri: &str, body: &[u8]) -> Result<Vec<u8>> {
self.build_request_with_body(Method::GetParameter, uri, body, &[])
}
fn build_request(
&mut self,
method: Method,
uri: &str,
_range: Option<&str>,
extra: &[(headers::HeaderName, String)],
) -> Result<Vec<u8>> {
self.build_request_with_body(method, uri, &[], extra)
}
fn build_request_with_body(
&mut self,
method: Method,
uri: &str,
body: &[u8],
extra: &[(headers::HeaderName, String)],
) -> Result<Vec<u8>> {
client_next_state(self.state, &method)?;
let cseq = self.next_cseq;
let request = self.assemble(method.clone(), uri, cseq, body, extra)?;
let bytes = serialize(&Message::from(request.clone()))?;
self.next_cseq += 1;
self.pending.insert(
cseq,
Pending {
method,
uri: uri.to_string(),
request,
auth_retried: false,
},
);
Ok(bytes)
}
fn assemble(
&mut self,
method: Method,
uri: &str,
cseq: u32,
body: &[u8],
extra: &[(headers::HeaderName, String)],
) -> Result<Request<Body>> {
let url = rtsp_types::Url::parse(uri)
.map_err(|e| Error::TransportParse(format!("invalid request URI {uri:?}: {e}")))?;
let mut builder = Request::builder(method.clone(), Version::V1_0)
.request_uri(url)
.header(headers::CSEQ, cseq.to_string())
.header(headers::USER_AGENT, self.user_agent.clone());
if let Some(sid) = &self.session_id {
builder = builder.header(headers::SESSION, sid.clone());
}
for (name, value) in extra {
builder = builder.header(name.clone(), value.clone());
}
if let Some(auth) = &mut self.authenticator {
let value = auth.authorization(<&str>::from(&method), uri)?;
builder = builder.header(headers::AUTHORIZATION, value);
}
let request = if body.is_empty() {
builder.build(Vec::new())
} else {
builder.build(body.to_vec())
};
Ok(request)
}
pub fn handle_data(&mut self, data: &[u8]) -> Result<Vec<ClientEvent>> {
self.inbound.extend_from_slice(data);
let mut events = Vec::new();
loop {
if self.inbound.is_empty() {
break;
}
if self.inbound[0] == MAGIC {
match interleaved::InterleavedFrame::parse(&self.inbound)? {
Some((frame, consumed)) => {
events.push(ClientEvent::MediaData {
channel: frame.channel,
data: frame.payload,
});
self.inbound.drain(..consumed);
}
None => break, }
continue;
}
match Message::<Body>::parse(&self.inbound) {
Ok((message, consumed)) => {
self.inbound.drain(..consumed);
self.process_message(message, &mut events)?;
}
Err(rtsp_types::ParseError::Incomplete(_)) => break,
Err(rtsp_types::ParseError::Error) => {
return Err(Error::MessageParse("malformed RTSP message".into()));
}
}
}
Ok(events)
}
fn process_message(
&mut self,
message: Message<Body>,
events: &mut Vec<ClientEvent>,
) -> Result<()> {
match message {
Message::Response(response) => {
let cseq = header_value(response.header(&headers::CSEQ))
.and_then(|s| s.trim().parse::<u32>().ok())
.ok_or(Error::MissingCSeq)?;
let status = response.status();
if status == StatusCode::Unauthorized {
if let Some(retry) = self.try_auth_retry(cseq, &response)? {
events.push(retry);
return Ok(());
}
}
let pending = self.pending.remove(&cseq).ok_or(Error::UnknownCSeq(cseq))?;
if let Some(session_hdr) = header_value(response.header(&headers::SESSION)) {
let (id, timeout) = parse_session(session_hdr);
self.session_id = Some(id);
if timeout.is_some() {
self.session_timeout = timeout;
}
}
if pending.method == Method::Setup {
if let Some(t) = header_value(response.header(&headers::TRANSPORT)) {
self.negotiated_transport = Some(Transport::parse(t)?);
}
}
if status.is_success() {
self.state = client_next_state(self.state, &pending.method)?;
if pending.method == Method::Teardown {
self.session_id = None;
self.session_timeout = None;
self.authenticator = None;
}
} else if status.is_redirection() {
self.state = SessionState::Init;
}
events.push(ClientEvent::Response {
cseq,
method: pending.method,
status,
body: response.into_body(),
});
Ok(())
}
Message::Data(data) => {
events.push(ClientEvent::MediaData {
channel: data.channel_id(),
data: data.into_body(),
});
Ok(())
}
Message::Request(_) => {
Ok(())
}
}
}
fn try_auth_retry(
&mut self,
cseq: u32,
response: &rtsp_types::Response<Body>,
) -> Result<Option<ClientEvent>> {
let creds = match &self.credentials {
Some(c) => c.clone(),
None => return Ok(None),
};
let (method, uri, already) = match self.pending.get(&cseq) {
Some(p) => (p.method.clone(), p.uri.clone(), p.auth_retried),
None => return Ok(None),
};
let challenge = header_value(response.header(&headers::WWW_AUTHENTICATE))
.ok_or_else(|| Error::Auth("401 without WWW-Authenticate".into()))?;
let stale = challenge.to_ascii_lowercase().contains("stale=true");
if self.authenticator.is_none() || already || stale {
self.authenticator = Some(Authenticator::from_challenge(challenge, creds)?);
}
if already && !stale {
return Ok(None);
}
let extra = self.replay_extra(&method, cseq);
self.pending.remove(&cseq);
let new_cseq = self.next_cseq;
let request = self.assemble(method.clone(), &uri, new_cseq, &[], &extra)?;
let bytes = serialize(&Message::from(request.clone()))?;
self.next_cseq += 1;
self.pending.insert(
new_cseq,
Pending {
method: method.clone(),
uri,
request,
auth_retried: true,
},
);
Ok(Some(ClientEvent::AuthRetry {
method,
cseq: new_cseq,
request: bytes,
}))
}
fn replay_extra(&self, _method: &Method, old_cseq: u32) -> Vec<(headers::HeaderName, String)> {
let mut extra = Vec::new();
if let Some(p) = self.pending.get(&old_cseq) {
for name in [headers::ACCEPT, headers::TRANSPORT, headers::RANGE] {
if let Some(v) = header_value(p.request.header(&name)) {
extra.push((name, v.to_string()));
}
}
}
extra
}
}
fn header_value(h: Option<&headers::HeaderValue>) -> Option<&str> {
h.map(|v| v.as_str())
}
fn parse_session(value: &str) -> (String, Option<u64>) {
let mut parts = value.split(';').map(str::trim);
let id = parts.next().unwrap_or("").to_string();
let timeout = value
.split(';')
.filter_map(|s| s.trim().strip_prefix("timeout="))
.find_map(|s| s.trim().parse::<u64>().ok());
(id, timeout)
}
fn serialize(message: &Message<Body>) -> Result<Vec<u8>> {
let mut out = Vec::new();
message
.write(&mut out)
.map_err(|e| Error::MessageWrite(e.to_string()))?;
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn play_in_init_bites() {
let mut c = ClientSession::new();
assert!(c.play("rtsp://h/s").is_err());
}
#[test]
fn setup_allowed_in_init() {
let mut c = ClientSession::new();
let t = Transport::single(crate::transport::TransportSpec::rtp_avp_tcp_interleaved(
0, 1,
));
assert!(c.setup("rtsp://h/s", &t).is_ok());
}
#[test]
fn cseq_increments() {
let mut c = ClientSession::new();
let a = c.options("rtsp://h/s").unwrap();
let b = c.describe("rtsp://h/s").unwrap();
assert!(String::from_utf8_lossy(&a).contains("CSeq: 1"));
assert!(String::from_utf8_lossy(&b).contains("CSeq: 2"));
}
}