use rtsp_types::{headers, Message, Method, Request, Response, StatusCode, Version};
use crate::error::{Error, Result};
use crate::state::{server_next_state, SessionState};
use crate::transport::Transport;
type Body = Vec<u8>;
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ServerEvent {
RequestAccepted {
method: Method,
cseq: u32,
state: SessionState,
},
MethodNotValid {
method: Method,
state: SessionState,
},
SessionSetup {
session_id: String,
transport: Transport,
},
}
const DEFAULT_SESSION_SEED: u64 = 0x1234_5678;
#[derive(Debug)]
pub struct ServerSession {
state: SessionState,
session_id: Option<String>,
session_timeout: Option<u64>,
next_session_seed: u64,
negotiated_transport: Option<Transport>,
server_header: String,
}
impl Default for ServerSession {
fn default() -> Self {
Self::new()
}
}
impl ServerSession {
pub fn new() -> Self {
ServerSession {
state: SessionState::Init,
session_id: None,
session_timeout: None,
next_session_seed: DEFAULT_SESSION_SEED,
negotiated_transport: None,
server_header: "rtsp-runtime".to_string(),
}
}
pub fn with_session_timeout(mut self, seconds: u64) -> Self {
self.session_timeout = Some(seconds);
self
}
pub fn with_session_seed(mut self, seed: u64) -> Self {
self.next_session_seed = seed;
self
}
pub fn state(&self) -> SessionState {
self.state
}
pub fn session_id(&self) -> Option<&str> {
self.session_id.as_deref()
}
pub fn negotiated_transport(&self) -> Option<&Transport> {
self.negotiated_transport.as_ref()
}
pub fn handle_request(&mut self, data: &[u8]) -> Result<(Vec<u8>, Vec<ServerEvent>)> {
let (message, _consumed) =
Message::<Body>::parse(data).map_err(|e| Error::MessageParse(format!("{e:?}")))?;
let request = match message {
Message::Request(r) => r,
_ => return Err(Error::MessageParse("expected an RTSP request".into())),
};
self.handle_parsed(request)
}
fn handle_parsed(&mut self, request: Request<Body>) -> Result<(Vec<u8>, Vec<ServerEvent>)> {
let method = request.method().clone();
let cseq = header_value(request.header(&headers::CSEQ))
.and_then(|s| s.trim().parse::<u32>().ok())
.ok_or(Error::MissingCSeq)?;
let next_state = match server_next_state(self.state, &method) {
Ok(s) => s,
Err(_) => {
let resp = self.build_response(StatusCode::MethodNotValidInThisState, cseq, |b| b);
let bytes = serialize(&Message::from(resp))?;
return Ok((
bytes,
vec![ServerEvent::MethodNotValid {
method,
state: self.state,
}],
));
}
};
let mut events = Vec::new();
if method == Method::Setup {
let transport = match header_value(request.header(&headers::TRANSPORT)) {
Some(t) => Transport::parse(t)?,
None => {
let resp = self.build_response(StatusCode::UnsupportedTransport, cseq, |b| b);
return Ok((serialize(&Message::from(resp))?, events));
}
};
if transport.first().is_none() {
let resp = self.build_response(StatusCode::UnsupportedTransport, cseq, |b| b);
return Ok((serialize(&Message::from(resp))?, events));
}
let session_id = self
.session_id
.clone()
.unwrap_or_else(|| self.allocate_session());
self.session_id = Some(session_id.clone());
self.negotiated_transport = Some(transport.clone());
let sid = session_id.clone();
let session_hdr = match self.session_timeout {
Some(t) => format!("{sid};timeout={t}"),
None => sid.clone(),
};
let transport_hdr = transport.to_header_value();
let resp = self.build_response(StatusCode::Ok, cseq, |b| {
b.header(headers::SESSION, session_hdr)
.header(headers::TRANSPORT, transport_hdr)
});
self.state = next_state;
events.push(ServerEvent::SessionSetup {
session_id,
transport,
});
events.push(ServerEvent::RequestAccepted {
method,
cseq,
state: self.state,
});
return Ok((serialize(&Message::from(resp))?, events));
}
let session_hdr = self.session_id.clone();
let resp = self.build_response(StatusCode::Ok, cseq, |mut b| {
if let Some(sid) = &session_hdr {
b = b.header(headers::SESSION, sid.clone());
}
b
});
self.state = next_state;
if method == Method::Teardown {
self.session_id = None;
self.negotiated_transport = None;
}
events.push(ServerEvent::RequestAccepted {
method,
cseq,
state: self.state,
});
Ok((serialize(&Message::from(resp))?, events))
}
fn allocate_session(&mut self) -> String {
let id = self.next_session_seed;
self.next_session_seed = self.next_session_seed.wrapping_add(1);
format!("{id:08}")
}
fn build_response<F>(&self, status: StatusCode, cseq: u32, f: F) -> Response<Body>
where
F: FnOnce(rtsp_types::ResponseBuilder) -> rtsp_types::ResponseBuilder,
{
let builder = Response::builder(Version::V1_0, status)
.header(headers::CSEQ, cseq.to_string())
.header(headers::SERVER, self.server_header.clone());
f(builder).build(Vec::new())
}
}
fn header_value(h: Option<&headers::HeaderValue>) -> Option<&str> {
h.map(|v| v.as_str())
}
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::*;
fn req(bytes: &str) -> Vec<u8> {
bytes.replace('\n', "\r\n").into_bytes()
}
#[test]
fn setup_transitions_init_to_ready() {
let mut s = ServerSession::new();
let (resp, events) = s
.handle_request(&req(
"SETUP rtsp://h/s RTSP/1.0\nCSeq: 1\nTransport: RTP/AVP/TCP;interleaved=0-1\n\n",
))
.unwrap();
assert_eq!(s.state(), SessionState::Ready);
let text = String::from_utf8_lossy(&resp);
assert!(text.contains("200"));
assert!(text.contains("Session:"));
assert!(text.contains("Transport:"));
assert!(events
.iter()
.any(|e| matches!(e, ServerEvent::SessionSetup { .. })));
}
#[test]
fn play_in_init_returns_455() {
let mut s = ServerSession::new();
let (resp, events) = s
.handle_request(&req("PLAY rtsp://h/s RTSP/1.0\nCSeq: 1\n\n"))
.unwrap();
assert_eq!(s.state(), SessionState::Init);
assert!(String::from_utf8_lossy(&resp).contains("455"));
assert!(events
.iter()
.any(|e| matches!(e, ServerEvent::MethodNotValid { .. })));
}
}