1use error::{Error, SapResult};
19use lazy_static::lazy_static;
20use murmur3::murmur3_32;
21use sdp::SessionDescription;
22use socket2::{Domain, Protocol, SockAddr, Socket, Type};
23use std::{
24 collections::HashMap,
25 io::Cursor,
26 net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr},
27 time::{Duration, SystemTime, UNIX_EPOCH},
28};
29use tokio::{
30 net::UdpSocket,
31 select,
32 sync::{mpsc, oneshot},
33 time::interval,
34};
35use tosub::SubsystemHandle;
36use tracing::{debug, error, info, trace};
37
38pub mod error;
39
40const DEFAULT_PAYLOAD_TYPE: &str = "application/sdp";
41const DEFAULT_SAP_PORT: u16 = 9875;
42const DEFAULT_MULTICAST_ADDRESS: [u8; 4] = [239, 255, 255, 255];
43
44lazy_static! {
45 static ref HASH_SEED: u32 = SystemTime::now()
46 .duration_since(UNIX_EPOCH)
47 .expect("something is wrong with the system clock")
48 .as_secs() as u32;
49}
50
51#[derive(Debug, Clone)]
52pub struct SessionAnnouncement {
53 pub deletion: bool,
54 pub encrypted: bool,
55 pub compressed: bool,
56 pub msg_id_hash: u16,
57 pub auth_data: Option<String>,
58 pub originating_source: IpAddr,
59 pub payload_type: Option<String>,
60 pub sdp: SessionDescription,
61}
62
63impl SessionAnnouncement {
64 pub fn new(sdp: SessionDescription) -> SapResult<Self> {
65 Ok(Self {
66 deletion: false,
67 encrypted: false,
68 compressed: false,
69 msg_id_hash: sdp_hash(&sdp),
70 auth_data: None,
71 originating_source: sdp.origin.unicast_address.parse()?,
72 payload_type: Some(DEFAULT_PAYLOAD_TYPE.to_owned()),
73 sdp,
74 })
75 }
76
77 pub fn deletion(sdp: SessionDescription) -> SapResult<Self> {
78 Ok(Self {
79 deletion: true,
80 encrypted: false,
81 compressed: false,
82 msg_id_hash: sdp_hash(&sdp),
83 auth_data: None,
84 originating_source: sdp.origin.unicast_address.parse()?,
85 payload_type: Some(DEFAULT_PAYLOAD_TYPE.to_owned()),
86 sdp,
87 })
88 }
89}
90
91pub struct SapActor {
92 subsys: SubsystemHandle,
93 rx: mpsc::Receiver<Vec<u8>>,
94 deletion_announcements: HashMap<u64, SubsystemHandle>,
95 event_tx: mpsc::Sender<Event>,
96 msg_rx: mpsc::Receiver<Message>,
97 announcement_sender: mpsc::Sender<SessionAnnouncement>,
98}
99
100pub enum Event {
101 SessionFound(SessionAnnouncement),
102 SessionLost(SessionAnnouncement),
103}
104
105enum Message {
106 AnnounceSession(Box<SessionAnnouncement>, oneshot::Sender<SapResult<()>>),
107 DeleteSession(u64, oneshot::Sender<SapResult<()>>),
108 DeleteAllSessions(oneshot::Sender<SapResult<()>>),
109}
110
111impl SapActor {
112 async fn run(mut self) -> SapResult<()> {
113 loop {
114 select! {
115 recv = self.msg_rx.recv() => if let Some(msg) = recv {
116 self.process_api_msg(msg).await?;
117 } else {
118 info!("Message channel closed, shutting down SAP actor.");
119 break;
120 },
121 recv = self.rx.recv() => if let Some(data) = recv {
122 self.forward_announcement(&data).await;
123 } else {
124 info!("Socket channel closed, shutting down SAP actor.");
125 break;
126 },
127 _ = self.subsys.shutdown_requested() => {
128 info!("Shutdown requested, shutting down SAP actor.");
129 break;
130 },
131 }
132 }
133
134 info!("SAP actor stopped.");
135
136 Ok(())
137 }
138
139 async fn process_api_msg(&mut self, msg: Message) -> SapResult<()> {
140 match msg {
141 Message::AnnounceSession(sa, tx) => {
142 tx.send(self.announce_session(*sa).await).ok();
143 }
144 Message::DeleteSession(id, tx) => {
145 tx.send(self.delete_session(id).await).ok();
146 }
147 Message::DeleteAllSessions(tx) => {
148 tx.send(self.delete_all_sessions().await).ok();
149 }
150 }
151
152 Ok(())
153 }
154
155 async fn forward_announcement(&self, buf: &[u8]) {
156 trace!("forwarding SAP message");
157 match decode_sap(buf) {
158 Ok(sap) => {
159 let event = if sap.deletion {
160 Event::SessionLost(sap)
161 } else {
162 Event::SessionFound(sap)
163 };
164 if let Err(e) = self.event_tx.send(event).await {
165 error!("Error forwarding SAP message error: {e}");
166 } else {
167 trace!("SAP message forwarded");
168 }
169 }
170 Err(e) => {
171 error!("error decoding SAP message: {e}");
172 }
173 }
174 }
175
176 async fn announce_session(&mut self, announcement: SessionAnnouncement) -> SapResult<()> {
177 let session_id = announcement.sdp.origin.session_id;
178
179 info!(
180 "Announcing new session with hash {}.",
181 announcement.msg_id_hash
182 );
183
184 self.delete_session(announcement.sdp.origin.session_id)
185 .await?;
186
187 let mut deletion_announcement = announcement.clone();
188 deletion_announcement.deletion = true;
189
190 let tx = self.announcement_sender.clone();
191
192 let announcement = self.subsys.spawn(
193 format!("announcement/{}", announcement.msg_id_hash),
194 |s| async move {
195 let mut interval = interval(Duration::from_secs(5));
196
197 loop {
198 select! {
202 _ = interval.tick() => tx.send(announcement.clone()).await?,
203 _ = s.shutdown_requested() => break,
204 }
205 }
206
207 tx.send(deletion_announcement).await.ok();
208
209 Ok::<(), error::Error>(())
210 },
211 );
212
213 self.deletion_announcements.insert(session_id, announcement);
214
215 Ok(())
216 }
217
218 async fn delete_session(&mut self, session_id: u64) -> SapResult<()> {
219 if let Some(subsys) = self.deletion_announcements.remove(&session_id) {
220 info!("Deleting active session {session_id}.");
221 subsys.request_local_shutdown();
222 } else {
223 debug!("No session active, nothing to delete.");
224 }
225
226 Ok(())
227 }
228
229 async fn delete_all_sessions(&mut self) -> SapResult<()> {
230 let sessions = self.deletion_announcements.drain().collect::<Vec<_>>();
231
232 for (session_id, subsys) in sessions {
233 info!("Deleting active session {session_id}.");
234 subsys.request_local_shutdown();
235 }
236
237 Ok(())
238 }
239}
240
241async fn send_announcement(
242 socket: &UdpSocket,
243 multicast_addr: &SocketAddr,
244 announcement: &SessionAnnouncement,
245) -> SapResult<()> {
246 debug!(
247 "Broadcasting session description to {}:\n{}\n",
248 multicast_addr,
249 announcement.sdp.marshal()
250 );
251 let msg = encode_sap(announcement);
252 socket.send_to(&msg, multicast_addr).await?;
253 Ok(())
254}
255
256#[derive(Clone)]
257pub struct Sap {
258 msg_tx: mpsc::Sender<Message>,
259}
260
261impl Sap {
262 pub async fn new(
263 subsys: &SubsystemHandle,
264 iface_name: String,
265 ) -> SapResult<(Self, mpsc::Receiver<Event>)> {
266 let socket = create_socket(iface_name).await?;
267
268 let deletion_announcements = HashMap::new();
269
270 let (event_tx, event_rx) = mpsc::channel(1);
271 let (msg_tx, msg_rx) = mpsc::channel(100);
272 let (socket_tx, socket_rx) = mpsc::channel(100);
273
274 subsys.spawn("sap", move |s| {
275 let multicast_addr = SocketAddr::new(
276 IpAddr::V4(Ipv4Addr::from(DEFAULT_MULTICAST_ADDRESS)),
277 DEFAULT_SAP_PORT,
278 );
279
280 let (announce_tx, announce_rx) = mpsc::channel(1);
281
282 s.spawn("socket", move |s| {
283 IoLoop {
284 s,
285 socket,
286 multicast_addr,
287 socket_tx,
288 announce_rx,
289 }
290 .io_loop()
291 });
292
293 SapActor {
294 subsys: s,
295 deletion_announcements,
296 event_tx,
297 msg_rx,
298 announcement_sender: announce_tx,
299 rx: socket_rx,
300 }
301 .run()
302 });
303
304 Ok((Sap { msg_tx }, event_rx))
305 }
306
307 pub async fn announce_session(&self, sd: SessionDescription) -> SapResult<()> {
308 let sa = SessionAnnouncement::new(sd)?;
309 let (tx, rx) = oneshot::channel();
310 self.msg_tx
311 .send(Message::AnnounceSession(Box::new(sa), tx))
312 .await?;
313 rx.await?
314 }
315
316 pub async fn delete_session(&self, session_id: u64) -> SapResult<()> {
317 let (tx, rx) = oneshot::channel();
318 self.msg_tx
319 .send(Message::DeleteSession(session_id, tx))
320 .await?;
321 rx.await?
322 }
323
324 pub async fn delete_all_sessions(&self) -> SapResult<()> {
325 let (tx, rx) = oneshot::channel();
326 self.msg_tx.send(Message::DeleteAllSessions(tx)).await?;
327 rx.await?
328 }
329}
330
331struct IoLoop {
332 s: SubsystemHandle,
333 socket: UdpSocket,
334 multicast_addr: SocketAddr,
335 socket_tx: mpsc::Sender<Vec<u8>>,
336 announce_rx: mpsc::Receiver<SessionAnnouncement>,
337}
338impl IoLoop {
339 async fn io_loop(mut self) -> SapResult<()> {
340 let mut buf = [0; 1024];
341
342 loop {
343 select! {
344 len = self.socket.recv(&mut buf) => self.socket_tx.send(buf[..len?].to_vec()).await?,
345 recv = self.announce_rx.recv() => if let Some(announcement) = recv {
346 send_announcement(&self.socket, &self.multicast_addr, &announcement).await?
347 } else {
348 break;
349 },
350 }
351 }
352
353 self.s.request_local_shutdown();
354
355 info!("SAP socket closed.");
356
357 Ok(())
358 }
359}
360
361pub fn decode_sap(msg: &[u8]) -> SapResult<SessionAnnouncement> {
362 let mut min_length = 4;
363
364 if msg.len() < min_length {
365 return Err(Error::MalformedPacket(msg.to_owned()));
366 }
367
368 let header = msg[0];
369 let auth_len = msg[1];
370 let msg_id_hash = u16::from_be_bytes([msg[2], msg[3]]);
371
372 let ipv6 = (header & 0b00001000) >> 3 == 1;
373 let deletion = (header & 0b00000100) >> 2 == 1;
374 let encrypted = (header & 0b00000010) >> 1 == 1;
375 let compressed = header & 0b00000001 == 1;
376
377 if encrypted {
379 return Err(Error::NotImplemented("encryption"));
380 }
381 if compressed {
383 return Err(Error::NotImplemented("encryption"));
384 }
385
386 if ipv6 {
387 min_length += 16;
388 } else {
389 min_length += 4;
390 }
391
392 if msg.len() < min_length {
393 return Err(Error::MalformedPacket(msg.to_owned()));
394 }
395
396 let originating_source = if ipv6 {
397 let bits = u128::from_be_bytes([
398 msg[4], msg[5], msg[6], msg[7], msg[8], msg[9], msg[10], msg[11], msg[12], msg[13],
399 msg[14], msg[15], msg[16], msg[17], msg[18], msg[19],
400 ]);
401 IpAddr::V6(Ipv6Addr::from_bits(bits))
402 } else {
403 let bits = u32::from_be_bytes([msg[4], msg[5], msg[6], msg[7]]);
404 IpAddr::V4(Ipv4Addr::from_bits(bits))
405 };
406
407 let auth_data_start = min_length;
408
409 min_length += auth_len as usize;
410
411 if msg.len() <= min_length {
412 return Err(Error::MalformedPacket(msg.to_owned()));
413 }
414
415 let auth_data = if auth_len > 0 {
416 Some(String::from_utf8_lossy(&msg[auth_data_start..min_length]).to_string())
417 } else {
418 None
419 };
420
421 let payload = String::from_utf8_lossy(&msg[min_length..]).to_string();
422 let split: Vec<&str> = payload.split('\0').collect();
423
424 let payload_type = if split.len() >= 2 {
425 Some(split[0].to_owned())
426 } else {
427 None
428 };
429
430 let payload = if split.len() == 1 {
431 split[0]
432 } else {
433 &split[1..].join("\0")
434 };
435
436 let sdp = SessionDescription::unmarshal(&mut Cursor::new(payload))?;
437
438 Ok(SessionAnnouncement {
439 deletion,
440 encrypted,
441 compressed,
442 msg_id_hash,
443 auth_data,
444 originating_source,
445 payload_type,
446 sdp,
447 })
448}
449
450pub fn encode_sap(msg: &SessionAnnouncement) -> Vec<u8> {
451 let v = 1u8;
452 let (a, originating_source): (u8, &[u8]) = match msg.originating_source {
453 IpAddr::V4(addr) => (0u8, &addr.octets()),
454 IpAddr::V6(addr) => (1u8, &addr.octets()),
455 };
456 let r = 0u8;
457 let t = if msg.deletion { 1u8 } else { 0u8 };
458 let e = if msg.encrypted { 1u8 } else { 0u8 };
459 let c = if msg.compressed { 1u8 } else { 0u8 };
460 let header = v << 5 | a << 4 | r << 3 | t << 2 | e << 1 | c;
461 let auth_len = msg.auth_data.as_ref().map(|d| d.len()).unwrap_or(0) as u8;
462 let msg_id_hash = msg.msg_id_hash.to_be_bytes();
463
464 let mut data = Vec::new();
465 data.push(header);
466 data.push(auth_len);
467 data.extend_from_slice(&msg_id_hash);
468 data.extend_from_slice(originating_source);
469 if let Some(auth_data) = &msg.auth_data {
470 data.extend_from_slice(auth_data.as_bytes());
471 }
472 if let Some(payload_type) = &msg.payload_type {
473 data.extend_from_slice(payload_type.as_bytes());
474 data.push(b'\0');
475 }
476 data.extend_from_slice(msg.sdp.marshal().as_bytes());
477
478 data
479}
480
481fn sdp_hash(sdp: &SessionDescription) -> u16 {
482 info!("computing message hash ...");
483 let res = murmur3_32(&mut Cursor::new(sdp.marshal()), *HASH_SEED).unwrap_or(0) as u16;
484 info!("computing message hash done");
485 res
486}
487
488fn get_iface_ipv4(iface_name: &str) -> Option<Ipv4Addr> {
489 for iface in if_addrs::get_if_addrs().ok()? {
490 if iface.name == iface_name
491 && let IpAddr::V4(addr) = iface.addr.ip()
492 && !addr.is_loopback()
493 && !addr.is_link_local()
494 && !addr.is_broadcast()
495 {
496 return Some(addr);
497 }
498 }
499 None
500}
501
502async fn create_socket(iface_name: String) -> SapResult<UdpSocket> {
503 let multicast_addr = Ipv4Addr::from(DEFAULT_MULTICAST_ADDRESS);
504 let iface_ip =
505 get_iface_ipv4(&iface_name).ok_or_else(|| Error::NoIpAddress(iface_name.clone()))?;
506 let local_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), DEFAULT_SAP_PORT);
507
508 let socket = Socket::new(Domain::IPV4, Type::DGRAM, Some(Protocol::UDP))?;
509 socket.set_reuse_address(true)?;
510 socket.set_nonblocking(true)?;
511 socket.bind(&SockAddr::from(local_addr))?;
512 socket.join_multicast_v4(&multicast_addr, &iface_ip)?;
513 socket.set_multicast_if_v4(&iface_ip)?;
514
515 info!(
516 "SAP socket joined multicast group {} on interface {} with IP {}.",
517 multicast_addr, iface_name, iface_ip
518 );
519
520 let socket = UdpSocket::from_std(socket.into())?;
521
522 Ok(socket)
523}
524
525#[cfg(test)]
526mod tests {
527
528 use super::*;
529
530 #[test]
531 fn sdp_gets_hashed_correctly() {
532 let sdp = SessionDescription::unmarshal(&mut Cursor::new(
533 "v=0
534o=- 123456 123458 IN IP4 10.0.1.2
535s=My sample flow
536i=4 channels: c1, c2, c3, c4
537t=0 0
538a=recvonly
539m=audio 5004 RTP/AVP 98
540c=IN IP4 239.69.11.44/32
541a=rtpmap:98 L24/48000/4
542a=ptime:1
543a=ts-refclk:ptp=IEEE1588-2008:00-11-22-FF-FE-33-44-55:0
544a=mediaclk:direct=0",
545 ))
546 .unwrap();
547 assert!(sdp_hash(&sdp) != 0);
548 }
549
550 #[test]
551 fn encode_decode_roundtrip_is_successful() {
552 let sdp = "v=0
553o=- 123456 123458 IN IP4 10.0.1.2
554s=My sample flow
555i=4 channels: c1, c2, c3, c4
556t=0 0
557a=recvonly
558m=audio 5004 RTP/AVP 98
559c=IN IP4 239.69.11.44/32
560a=rtpmap:98 L24/48000/4
561a=ptime:1
562a=ts-refclk:ptp=IEEE1588-2008:00-11-22-FF-FE-33-44-55:0
563a=mediaclk:direct=0
564";
565
566 let sa = SessionAnnouncement {
567 auth_data: None,
568 payload_type: None,
569 compressed: false,
570 deletion: true,
571 encrypted: false,
572 msg_id_hash: 1234,
573 originating_source: "127.0.0.1".parse().unwrap(),
574 sdp: SessionDescription::unmarshal(&mut Cursor::new(sdp)).unwrap(),
575 };
576
577 let sa_msg = encode_sap(&sa);
578
579 let decoded = decode_sap(&sa_msg).unwrap();
580
581 assert_eq!(sa.auth_data, decoded.auth_data);
582 assert_eq!(sa.compressed, decoded.compressed);
583 assert_eq!(sa.deletion, decoded.deletion);
584 assert_eq!(sa.encrypted, decoded.encrypted);
585 assert_eq!(sa.msg_id_hash, decoded.msg_id_hash);
586 assert_eq!(sa.originating_source, decoded.originating_source);
587 assert_eq!(sa.payload_type, decoded.payload_type);
588 assert_eq!(sa.sdp.marshal().replace('\r', ""), sdp);
589 }
590}