1use std::collections::HashMap;
26
27use rtsp_types::{Message, Method, Request, StatusCode, Version, headers};
28
29use crate::auth::{Authenticator, Credentials, RequestContext};
30use crate::error::{Error, Result};
31use crate::interleaved::{self, MAGIC};
32use crate::state::{SessionState, client_next_state};
33use crate::transport::Transport;
34
35type Body = Vec<u8>;
37
38#[derive(Debug, Clone)]
40struct Pending {
41 method: Method,
42 uri: String,
43 request: Request<Body>,
45 auth_retried: bool,
47}
48
49#[non_exhaustive]
51#[derive(Debug, Clone, PartialEq, Eq)]
52pub enum ClientEvent {
53 Response {
55 cseq: u32,
57 method: Method,
59 status: StatusCode,
61 body: Vec<u8>,
63 },
64 AuthRetry {
67 method: Method,
69 cseq: u32,
71 request: Vec<u8>,
73 },
74 MediaData {
76 channel: u8,
78 data: Vec<u8>,
80 },
81}
82
83#[derive(Debug)]
85pub struct ClientSession {
86 state: SessionState,
87 next_cseq: u32,
88 session_id: Option<String>,
89 session_timeout: Option<u64>,
90 credentials: Option<Credentials>,
91 authenticator: Option<Authenticator>,
92 negotiated_transport: Option<Transport>,
93 pending: HashMap<u32, Pending>,
94 inbound: Vec<u8>,
97 user_agent: String,
98}
99
100impl Default for ClientSession {
101 fn default() -> Self {
102 Self::new()
103 }
104}
105
106impl ClientSession {
107 pub fn new() -> Self {
110 ClientSession {
111 state: SessionState::Init,
112 next_cseq: 1,
113 session_id: None,
114 session_timeout: None,
115 credentials: None,
116 authenticator: None,
117 negotiated_transport: None,
118 pending: HashMap::new(),
119 inbound: Vec::new(),
120 user_agent: "rtsp-runtime".to_string(),
121 }
122 }
123
124 pub fn with_credentials(mut self, credentials: Credentials) -> Self {
126 self.credentials = Some(credentials);
127 self
128 }
129
130 pub fn with_user_agent(mut self, ua: impl Into<String>) -> Self {
132 self.user_agent = ua.into();
133 self
134 }
135
136 pub fn state(&self) -> SessionState {
138 self.state
139 }
140
141 pub fn session_id(&self) -> Option<&str> {
143 self.session_id.as_deref()
144 }
145
146 pub fn session_timeout(&self) -> Option<u64> {
148 self.session_timeout
149 }
150
151 pub fn negotiated_transport(&self) -> Option<&Transport> {
153 self.negotiated_transport.as_ref()
154 }
155
156 pub fn options(&mut self, uri: &str) -> Result<Vec<u8>> {
160 self.build_request(Method::Options, uri, None, &[])
161 }
162
163 pub fn describe(&mut self, uri: &str) -> Result<Vec<u8>> {
165 self.build_request(
166 Method::Describe,
167 uri,
168 None,
169 &[(headers::ACCEPT, "application/sdp".to_string())],
170 )
171 }
172
173 pub fn setup(&mut self, uri: &str, transport: &Transport) -> Result<Vec<u8>> {
175 self.build_request(
176 Method::Setup,
177 uri,
178 None,
179 &[(headers::TRANSPORT, transport.to_header_value())],
180 )
181 }
182
183 pub fn play(&mut self, uri: &str) -> Result<Vec<u8>> {
185 self.build_request(Method::Play, uri, None, &[])
186 }
187
188 pub fn pause(&mut self, uri: &str) -> Result<Vec<u8>> {
190 self.build_request(Method::Pause, uri, None, &[])
191 }
192
193 pub fn teardown(&mut self, uri: &str) -> Result<Vec<u8>> {
195 self.build_request(Method::Teardown, uri, None, &[])
196 }
197
198 pub fn get_parameter(&mut self, uri: &str, body: &[u8]) -> Result<Vec<u8>> {
201 self.build_request_with_body(Method::GetParameter, uri, body, &[])
202 }
203
204 fn build_request(
205 &mut self,
206 method: Method,
207 uri: &str,
208 _range: Option<&str>,
209 extra: &[(headers::HeaderName, String)],
210 ) -> Result<Vec<u8>> {
211 self.build_request_with_body(method, uri, &[], extra)
212 }
213
214 fn build_request_with_body(
215 &mut self,
216 method: Method,
217 uri: &str,
218 body: &[u8],
219 extra: &[(headers::HeaderName, String)],
220 ) -> Result<Vec<u8>> {
221 client_next_state(self.state, &method)?;
223
224 let cseq = self.next_cseq;
225 let request = self.assemble(method.clone(), uri, cseq, body, extra)?;
226 let bytes = serialize(&Message::from(request.clone()))?;
227 self.next_cseq += 1;
228 self.pending.insert(
229 cseq,
230 Pending {
231 method,
232 uri: uri.to_string(),
233 request,
234 auth_retried: false,
235 },
236 );
237 Ok(bytes)
238 }
239
240 fn assemble(
243 &mut self,
244 method: Method,
245 uri: &str,
246 cseq: u32,
247 body: &[u8],
248 extra: &[(headers::HeaderName, String)],
249 ) -> Result<Request<Body>> {
250 let url = rtsp_types::Url::parse(uri)
251 .map_err(|e| Error::TransportParse(format!("invalid request URI {uri:?}: {e}")))?;
252 let mut builder = Request::builder(method.clone(), Version::V1_0)
253 .request_uri(url)
254 .header(headers::CSEQ, cseq.to_string())
255 .header(headers::USER_AGENT, self.user_agent.clone());
256 if let Some(sid) = &self.session_id {
257 builder = builder.header(headers::SESSION, sid.clone());
258 }
259 for (name, value) in extra {
260 builder = builder.header(name.clone(), value.clone());
261 }
262 if let Some(auth) = &mut self.authenticator {
263 let ctx = RequestContext::new(<&str>::from(&method), uri);
264 let value = auth.authorization(&ctx)?;
265 builder = builder.header(headers::AUTHORIZATION, value);
266 }
267 let request = if body.is_empty() {
268 builder.build(Vec::new())
269 } else {
270 builder.build(body.to_vec())
271 };
272 Ok(request)
273 }
274
275 pub fn handle_data(&mut self, data: &[u8]) -> Result<Vec<ClientEvent>> {
280 self.inbound.extend_from_slice(data);
281 let mut events = Vec::new();
282
283 loop {
284 if self.inbound.is_empty() {
285 break;
286 }
287 if self.inbound[0] == MAGIC {
288 match interleaved::InterleavedFrame::parse(&self.inbound)? {
290 Some((frame, consumed)) => {
291 events.push(ClientEvent::MediaData {
292 channel: frame.channel,
293 data: frame.payload,
294 });
295 self.inbound.drain(..consumed);
296 }
297 None => break, }
299 continue;
300 }
301
302 match Message::<Body>::parse(&self.inbound) {
304 Ok((message, consumed)) => {
305 self.inbound.drain(..consumed);
306 self.process_message(message, &mut events)?;
307 }
308 Err(rtsp_types::ParseError::Incomplete(_)) => break,
309 Err(rtsp_types::ParseError::Error) => {
310 return Err(Error::MessageParse("malformed RTSP message".into()));
311 }
312 }
313 }
314 Ok(events)
315 }
316
317 fn process_message(
318 &mut self,
319 message: Message<Body>,
320 events: &mut Vec<ClientEvent>,
321 ) -> Result<()> {
322 match message {
323 Message::Response(response) => {
324 let cseq = header_value(response.header(&headers::CSEQ))
325 .and_then(|s| s.trim().parse::<u32>().ok())
326 .ok_or(Error::MissingCSeq)?;
327 let status = response.status();
328
329 if status == StatusCode::Unauthorized {
331 if let Some(retry) = self.try_auth_retry(cseq, &response)? {
332 events.push(retry);
333 return Ok(());
334 }
335 }
336
337 let pending = self.pending.remove(&cseq).ok_or(Error::UnknownCSeq(cseq))?;
338
339 if let Some(session_hdr) = header_value(response.header(&headers::SESSION)) {
341 let (id, timeout) = parse_session(session_hdr);
342 self.session_id = Some(id);
343 if timeout.is_some() {
344 self.session_timeout = timeout;
345 }
346 }
347 if pending.method == Method::Setup {
349 if let Some(t) = header_value(response.header(&headers::TRANSPORT)) {
350 self.negotiated_transport = Some(Transport::parse(t)?);
351 }
352 }
353
354 if status.is_success() {
356 self.state = client_next_state(self.state, &pending.method)?;
357 if pending.method == Method::Teardown {
359 self.session_id = None;
360 self.session_timeout = None;
361 self.authenticator = None;
362 }
363 } else if status.is_redirection() {
364 self.state = SessionState::Init;
365 }
366 events.push(ClientEvent::Response {
369 cseq,
370 method: pending.method,
371 status,
372 body: response.into_body(),
373 });
374 Ok(())
375 }
376 Message::Data(data) => {
377 events.push(ClientEvent::MediaData {
378 channel: data.channel_id(),
379 data: data.into_body(),
380 });
381 Ok(())
382 }
383 Message::Request(_) => {
384 Ok(())
387 }
388 }
389 }
390
391 fn try_auth_retry(
395 &mut self,
396 cseq: u32,
397 response: &rtsp_types::Response<Body>,
398 ) -> Result<Option<ClientEvent>> {
399 let creds = match &self.credentials {
400 Some(c) => c.clone(),
401 None => return Ok(None),
402 };
403 let (method, uri, already) = match self.pending.get(&cseq) {
405 Some(p) => (p.method.clone(), p.uri.clone(), p.auth_retried),
406 None => return Ok(None),
407 };
408
409 let challenge = header_value(response.header(&headers::WWW_AUTHENTICATE))
410 .ok_or_else(|| Error::Auth("401 without WWW-Authenticate".into()))?;
411 let stale = challenge.to_ascii_lowercase().contains("stale=true");
412
413 if self.authenticator.is_none() || already || stale {
416 self.authenticator = Some(Authenticator::from_challenge(challenge, creds)?);
417 }
418 if already && !stale {
421 return Ok(None);
422 }
423
424 let extra = self.replay_extra(&method, cseq);
427 self.pending.remove(&cseq);
428 let new_cseq = self.next_cseq;
429 let request = self.assemble(method.clone(), &uri, new_cseq, &[], &extra)?;
430 let bytes = serialize(&Message::from(request.clone()))?;
431 self.next_cseq += 1;
432 self.pending.insert(
433 new_cseq,
434 Pending {
435 method: method.clone(),
436 uri,
437 request,
438 auth_retried: true,
439 },
440 );
441 Ok(Some(ClientEvent::AuthRetry {
442 method,
443 cseq: new_cseq,
444 request: bytes,
445 }))
446 }
447
448 fn replay_extra(&self, _method: &Method, old_cseq: u32) -> Vec<(headers::HeaderName, String)> {
451 let mut extra = Vec::new();
452 if let Some(p) = self.pending.get(&old_cseq) {
453 for name in [headers::ACCEPT, headers::TRANSPORT, headers::RANGE] {
454 if let Some(v) = header_value(p.request.header(&name)) {
455 extra.push((name, v.to_string()));
456 }
457 }
458 }
459 extra
460 }
461}
462
463fn header_value(h: Option<&headers::HeaderValue>) -> Option<&str> {
465 h.map(|v| v.as_str())
466}
467
468fn parse_session(value: &str) -> (String, Option<u64>) {
470 let mut parts = value.split(';').map(str::trim);
471 let id = parts.next().unwrap_or("").to_string();
472 let timeout = value
473 .split(';')
474 .filter_map(|s| s.trim().strip_prefix("timeout="))
475 .find_map(|s| s.trim().parse::<u64>().ok());
476 (id, timeout)
477}
478
479fn serialize(message: &Message<Body>) -> Result<Vec<u8>> {
481 let mut out = Vec::new();
482 message
483 .write(&mut out)
484 .map_err(|e| Error::MessageWrite(e.to_string()))?;
485 Ok(out)
486}
487
488#[cfg(test)]
489mod tests {
490 use super::*;
491
492 #[test]
493 fn play_in_init_bites() {
494 let mut c = ClientSession::new();
495 assert!(c.play("rtsp://h/s").is_err());
496 }
497
498 #[test]
499 fn setup_allowed_in_init() {
500 let mut c = ClientSession::new();
501 let t = Transport::single(crate::transport::TransportSpec::rtp_avp_tcp_interleaved(
502 0, 1,
503 ));
504 assert!(c.setup("rtsp://h/s", &t).is_ok());
505 }
506
507 #[test]
508 fn cseq_increments() {
509 let mut c = ClientSession::new();
510 let a = c.options("rtsp://h/s").unwrap();
511 let b = c.describe("rtsp://h/s").unwrap();
512 assert!(String::from_utf8_lossy(&a).contains("CSeq: 1"));
513 assert!(String::from_utf8_lossy(&b).contains("CSeq: 2"));
514 }
515
516 #[test]
520 fn client_session_debug_does_not_leak_embedded_credentials_secret() {
521 let c = ClientSession::new()
522 .with_credentials(Credentials::new("admin", "extremely-secret-password"));
523 let debug = format!("{c:?}");
524 assert!(
525 !debug.contains("extremely-secret-password"),
526 "leaked via ClientSession Debug: {debug}"
527 );
528 assert!(debug.contains("***"), "expected redaction marker: {debug}");
529 }
530}