Skip to main content

sap_rs/
lib.rs

1/*
2 *  Copyright (C) 2024 Michael Bachmann
3 *
4 *  This program is free software: you can redistribute it and/or modify
5 *  it under the terms of the GNU Affero General Public License as published by
6 *  the Free Software Foundation, either version 3 of the License, or
7 *  (at your option) any later version.
8 *
9 *  This program is distributed in the hope that it will be useful,
10 *  but WITHOUT ANY WARRANTY; without even the implied warranty of
11 *  MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE.  See the
12 *  GNU Affero General Public License for more details.
13 *
14 *  You should have received a copy of the GNU Affero General Public License
15 *  along with this program.  If not, see <https://www.gnu.org/licenses/>.
16 */
17
18use 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                    // TODO receive other announcements and update delay
199                    // TODO send announcement in according intervals
200                    //
201                    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    // TODO implement decryption
378    if encrypted {
379        return Err(Error::NotImplemented("encryption"));
380    }
381    // TODO implement decompression
382    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}