Skip to main content

active_call/media/
stream.rs

1use crate::event::{EventSender, SessionEvent};
2use crate::media::dtmf::DtmfDetector;
3use crate::media::volume_control::HoldProcessor;
4use crate::media::{AudioFrame, Samples, TrackId};
5use crate::media::{
6    processor::Processor,
7    recorder::{Recorder, RecorderOption},
8    track::{Track, TrackPacketReceiver, TrackPacketSender},
9};
10use anyhow::Result;
11use std::collections::{HashMap, HashSet};
12use std::path::Path;
13use std::time::Duration;
14use tokio::task::JoinHandle;
15use tokio::{
16    select,
17    sync::{Mutex, mpsc},
18};
19use tokio_util::sync::CancellationToken;
20use tracing::{debug, info, warn};
21use uuid;
22
23pub struct MediaStream {
24    id: String,
25    pub cancel_token: CancellationToken,
26    recorder_option: Mutex<Option<RecorderOption>>,
27    tracks: Mutex<HashMap<TrackId, (Box<dyn Track>, DtmfDetector)>>,
28    suppressed_sources: Mutex<HashSet<TrackId>>,
29    event_sender: EventSender,
30    pub packet_sender: TrackPacketSender,
31    packet_receiver: Mutex<Option<TrackPacketReceiver>>,
32    recorder_sender: mpsc::UnboundedSender<AudioFrame>,
33    recorder_receiver: Mutex<Option<mpsc::UnboundedReceiver<AudioFrame>>>,
34    recorder_handle: Mutex<Option<JoinHandle<()>>>,
35}
36
37const CALLEE_TRACK_ID: &str = "callee-track";
38const QUEUE_HOLD_TRACK_ID: &str = "queue-hold-track";
39pub const SERVER_SIDE_TRACK_ID: &str = "server-side-track";
40
41pub struct MediaStreamBuilder {
42    cancel_token: Option<CancellationToken>,
43    id: Option<String>,
44    event_sender: EventSender,
45    recorder_config: Option<RecorderOption>,
46}
47
48impl MediaStreamBuilder {
49    pub fn new(event_sender: EventSender) -> Self {
50        Self {
51            id: Some(format!("ms:{}", uuid::Uuid::new_v4())),
52            cancel_token: None,
53            event_sender,
54            recorder_config: None,
55        }
56    }
57    pub fn with_id(mut self, id: String) -> Self {
58        self.id = Some(id);
59        self
60    }
61
62    pub fn with_cancel_token(mut self, cancel_token: CancellationToken) -> Self {
63        self.cancel_token = Some(cancel_token);
64        self
65    }
66
67    pub fn with_recorder_config(mut self, recorder_config: RecorderOption) -> Self {
68        self.recorder_config = Some(recorder_config);
69        self
70    }
71
72    pub fn build(self) -> MediaStream {
73        let cancel_token = self
74            .cancel_token
75            .unwrap_or_else(|| CancellationToken::new());
76        let tracks = Mutex::new(HashMap::new());
77        let (track_packet_sender, track_packet_receiver) = mpsc::unbounded_channel();
78        let (recorder_sender, recorder_receiver) = mpsc::unbounded_channel();
79        MediaStream {
80            id: self.id.unwrap_or_default(),
81            cancel_token,
82            recorder_option: Mutex::new(self.recorder_config),
83            tracks,
84            suppressed_sources: Mutex::new(HashSet::new()),
85            event_sender: self.event_sender,
86            packet_sender: track_packet_sender,
87            packet_receiver: Mutex::new(Some(track_packet_receiver)),
88            recorder_sender,
89            recorder_receiver: Mutex::new(Some(recorder_receiver)),
90            recorder_handle: Mutex::new(None),
91        }
92    }
93}
94
95impl MediaStream {
96    pub async fn serve(&self) -> Result<()> {
97        let packet_receiver = match self.packet_receiver.lock().await.take() {
98            Some(receiver) => receiver,
99            None => {
100                warn!(
101                    session_id = self.id,
102                    "MediaStream::serve() called multiple times, stream already serving"
103                );
104                return Ok(());
105            }
106        };
107        self.start_recorder().await.ok();
108        info!(session_id = self.id, "mediastream serving");
109        select! {
110            _ = self.cancel_token.cancelled() => {}
111            r = self.handle_forward_track(packet_receiver) => {
112                info!(session_id = self.id, "track packet receiver stopped {:?}", r);
113            }
114        }
115        Ok(())
116    }
117
118    pub fn stop(&self, _reason: Option<String>, _initiator: Option<String>) {
119        self.cancel_token.cancel()
120    }
121
122    pub async fn cleanup(&self) -> Result<()> {
123        self.cancel_token.cancel();
124        {
125            let mut tracks = self.tracks.lock().await;
126            for (id, (track, _)) in tracks.drain() {
127                if let Err(e) = track.stop().await {
128                    warn!(session_id = self.id, track_id = %id, "failed to stop track during cleanup: {}", e);
129                }
130            }
131        }
132        self.suppressed_sources.lock().await.clear();
133
134        if let Some(recorder_handle) = self.recorder_handle.lock().await.take() {
135            if let Ok(Ok(_)) = tokio::time::timeout(Duration::from_secs(30), recorder_handle).await
136            {
137                info!(session_id = self.id, "recorder stopped");
138            } else {
139                warn!(session_id = self.id, "recorder timeout");
140            }
141        }
142        Ok(())
143    }
144    pub async fn track_count(&self) -> usize {
145        self.tracks.lock().await.len()
146    }
147
148    pub async fn update_recorder_option(&self, recorder_config: RecorderOption) {
149        *self.recorder_option.lock().await = Some(recorder_config);
150        self.start_recorder().await.ok();
151    }
152
153    pub async fn remove_track(&self, id: &TrackId, graceful: bool) {
154        let track_entry = { self.tracks.lock().await.remove(id) };
155        if let Some((track, _)) = track_entry {
156            self.suppressed_sources.lock().await.remove(id);
157            let res = if !graceful {
158                track.stop().await
159            } else {
160                track.stop_graceful().await
161            };
162            match res {
163                Ok(_) => {}
164                Err(e) => {
165                    warn!(session_id = self.id, "failed to stop track: {}", e);
166                }
167            }
168        }
169    }
170    pub async fn update_remote_description(
171        &self,
172        track_id: &TrackId,
173        answer: &String,
174    ) -> Result<()> {
175        let track_entry = { self.tracks.lock().await.remove(track_id) };
176        if let Some((mut track, dtmf)) = track_entry {
177            let res = track.update_remote_description(answer).await;
178            self.tracks
179                .lock()
180                .await
181                .insert(track_id.clone(), (track, dtmf));
182            res?;
183        }
184        Ok(())
185    }
186
187    pub async fn update_remote_description_force(
188        &self,
189        track_id: &TrackId,
190        answer: &String,
191    ) -> Result<()> {
192        let track_entry = { self.tracks.lock().await.remove(track_id) };
193        if let Some((mut track, dtmf)) = track_entry {
194            let res = track.update_remote_description_force(answer).await;
195            self.tracks
196                .lock()
197                .await
198                .insert(track_id.clone(), (track, dtmf));
199            res?;
200        }
201        Ok(())
202    }
203
204    pub async fn handshake(
205        &self,
206        track_id: &TrackId,
207        offer: String,
208        timeout: Option<Duration>,
209    ) -> Result<String> {
210        let track_entry = { self.tracks.lock().await.remove(track_id) };
211        if let Some((mut track, dtmf)) = track_entry {
212            let res = track.handshake(offer, timeout).await;
213            self.tracks
214                .lock()
215                .await
216                .insert(track_id.clone(), (track, dtmf));
217            res
218        } else {
219            anyhow::bail!("track not found: {}", track_id)
220        }
221    }
222
223    pub async fn update_track(&self, mut track: Box<dyn Track>, play_id: Option<String>) {
224        self.remove_track(track.id(), false).await;
225        if self.recorder_option.lock().await.is_some() {
226            track.append_processor(Box::new(RecorderProcessor::new(
227                self.recorder_sender.clone(),
228            )));
229        }
230        match track
231            .start(self.event_sender.clone(), self.packet_sender.clone())
232            .await
233        {
234            Ok(_) => {
235                info!(session_id = self.id, track_id = track.id(), "track started");
236                let track_id = track.id().clone();
237                self.tracks
238                    .lock()
239                    .await
240                    .insert(track_id.clone(), (track, DtmfDetector::new()));
241                self.event_sender
242                    .send(SessionEvent::TrackStart {
243                        track_id,
244                        timestamp: crate::media::get_timestamp(),
245                        play_id,
246                    })
247                    .ok();
248            }
249            Err(e) => {
250                warn!(
251                    session_id = self.id,
252                    track_id = track.id(),
253                    play_id = play_id.as_deref(),
254                    "Failed to start track: {}",
255                    e
256                );
257            }
258        }
259    }
260
261    pub async fn mute_track(&self, id: Option<TrackId>) {
262        if let Some(id) = id {
263            if let Some((track, _)) = self.tracks.lock().await.get_mut(&id) {
264                MuteProcessor::mute_track(track.as_mut());
265            }
266        } else {
267            for (track, _) in self.tracks.lock().await.values_mut() {
268                MuteProcessor::mute_track(track.as_mut());
269            }
270        }
271    }
272
273    pub async fn unmute_track(&self, id: Option<TrackId>) {
274        if let Some(id) = id {
275            if let Some((track, _)) = self.tracks.lock().await.get_mut(&id) {
276                MuteProcessor::unmute_track(track.as_mut());
277            }
278        } else {
279            for (track, _) in self.tracks.lock().await.values_mut() {
280                MuteProcessor::unmute_track(track.as_mut());
281            }
282        }
283    }
284
285    pub async fn pause_playback(&self, id: TrackId) -> Result<()> {
286        self.set_playback_paused(id, true).await
287    }
288
289    pub async fn resume_playback(&self, id: TrackId) -> Result<()> {
290        self.set_playback_paused(id, false).await
291    }
292
293    async fn set_playback_paused(&self, id: TrackId, paused: bool) -> Result<()> {
294        if let Some((track, _)) = self.tracks.lock().await.get_mut(&id) {
295            if track.set_paused(paused) {
296                Ok(())
297            } else {
298                warn!(
299                    session_id = self.id,
300                    track_id = %id,
301                    paused,
302                    "pause state requested for track that does not support pausing"
303                );
304                Err(anyhow::anyhow!("track does not support pausing: {}", id))
305            }
306        } else {
307            warn!(
308                session_id = self.id,
309                track_id = %id,
310                paused,
311                "pause state requested for unknown track"
312            );
313            Err(anyhow::anyhow!("track not found: {}", id))
314        }
315    }
316
317    pub async fn hold_track(&self, id: Option<TrackId>) {
318        if let Some(id) = id {
319            if let Some((track, _)) = self.tracks.lock().await.get_mut(&id) {
320                HoldTrack::hold_track(track.as_mut());
321            }
322        } else {
323            for (track, _) in self.tracks.lock().await.values_mut() {
324                HoldTrack::hold_track(track.as_mut());
325            }
326        }
327    }
328
329    pub async fn resume_track(&self, id: Option<TrackId>) {
330        if let Some(id) = id {
331            if let Some((track, _)) = self.tracks.lock().await.get_mut(&id) {
332                HoldTrack::resume_track(track.as_mut());
333            }
334        } else {
335            for (track, _) in self.tracks.lock().await.values_mut() {
336                HoldTrack::resume_track(track.as_mut());
337            }
338        }
339    }
340
341    pub async fn suppress_forwarding(&self, track_id: &TrackId) {
342        self.suppressed_sources
343            .lock()
344            .await
345            .insert(track_id.clone());
346    }
347
348    pub async fn resume_forwarding(&self, track_id: &TrackId) {
349        self.suppressed_sources.lock().await.remove(track_id);
350    }
351
352    pub async fn remove_processor<T: 'static>(&self, track_id: &TrackId) -> Result<()> {
353        if let Some((track, _)) = self.tracks.lock().await.get_mut(track_id) {
354            track.as_mut().processor_chain().remove_processor::<T>();
355            Ok(())
356        } else {
357            Err(anyhow::anyhow!("Track {} not found", track_id))
358        }
359    }
360
361    pub async fn append_processor(
362        &self,
363        track_id: &TrackId,
364        processor: Box<dyn crate::media::processor::Processor>,
365    ) -> Result<()> {
366        if let Some((track, _)) = self.tracks.lock().await.get_mut(track_id) {
367            track.as_mut().processor_chain().append_processor(processor);
368            Ok(())
369        } else {
370            Err(anyhow::anyhow!("Track {} not found", track_id))
371        }
372    }
373
374}
375
376#[derive(Clone)]
377pub struct RecorderProcessor {
378    sender: mpsc::UnboundedSender<AudioFrame>,
379}
380
381impl RecorderProcessor {
382    pub fn new(sender: mpsc::UnboundedSender<AudioFrame>) -> Self {
383        Self { sender }
384    }
385}
386
387impl Processor for RecorderProcessor {
388    fn process_frame(&mut self, frame: &mut AudioFrame) -> Result<()> {
389        let frame_clone = frame.clone();
390        let _ = self.sender.send(frame_clone);
391        Ok(())
392    }
393}
394
395impl MediaStream {
396    pub async fn start_recorder(&self) -> Result<()> {
397        let recorder_option = self.recorder_option.lock().await.clone();
398        if let Some(recorder_option) = recorder_option {
399            if recorder_option.recorder_file.is_empty() {
400                warn!(
401                    session_id = self.id,
402                    "recorder file is empty, skipping recorder start"
403                );
404                return Ok(());
405            }
406            let recorder_receiver = match self.recorder_receiver.lock().await.take() {
407                Some(receiver) => receiver,
408                None => {
409                    return Ok(());
410                }
411            };
412            let cancel_token = self.cancel_token.child_token();
413            let session_id_clone = self.id.clone();
414
415            info!(
416                session_id = session_id_clone,
417                sample_rate = recorder_option.samplerate,
418                ptime = recorder_option.ptime,
419                "start recorder",
420            );
421
422            let recorder_handle = crate::spawn(async move {
423                let recorder_file = recorder_option.recorder_file.clone();
424                let recorder =
425                    Recorder::new(cancel_token, session_id_clone.clone(), recorder_option);
426                match recorder
427                    .process_recording(Path::new(&recorder_file), recorder_receiver)
428                    .await
429                {
430                    Ok(_) => {}
431                    Err(e) => {
432                        warn!(
433                            session_id = session_id_clone,
434                            "Failed to process recorder: {}", e
435                        );
436                    }
437                }
438            });
439            *self.recorder_handle.lock().await = Some(recorder_handle);
440
441            // Inject RecorderProcessor into tracks that were added before the recorder started
442            for (track, _) in self.tracks.lock().await.values_mut() {
443                track.insert_processor(Box::new(RecorderProcessor::new(
444                    self.recorder_sender.clone(),
445                )));
446            }
447        }
448        Ok(())
449    }
450
451    pub async fn set_track_refer(&self, track_id: &TrackId, refer: Option<bool>) {
452        if let Some((_, dtmf)) = self.tracks.lock().await.get_mut(track_id) {
453            dtmf.refer = refer;
454        }
455    }
456
457    pub async fn set_track_dtmf_forward(&self, track_id: &TrackId, forward: bool) {
458        if let Some((_, dtmf)) = self.tracks.lock().await.get_mut(track_id) {
459            dtmf.suppress_dtmf_forward = !forward;
460        }
461    }
462
463    async fn handle_forward_track(&self, mut packet_receiver: TrackPacketReceiver) {
464        let event_sender = self.event_sender.clone();
465        while let Some(packet) = packet_receiver.recv().await {
466            let suppressed = {
467                self.suppressed_sources
468                    .lock()
469                    .await
470                    .contains(&packet.track_id)
471            };
472
473            let is_dtmf = matches!(&packet.samples,
474                Samples::RTP { payload_type, .. } if *payload_type >= 96 && *payload_type <= 127);
475
476            let mut tracks = self.tracks.lock().await;
477
478            // Check once whether the source track suppresses DTMF forwarding.
479            let source_suppresses_dtmf = is_dtmf
480                && tracks
481                    .get(&packet.track_id)
482                    .map(|(_, d)| d.suppress_dtmf_forward)
483                    .unwrap_or(false);
484
485            for (track, dtmf_detector) in tracks.values_mut() {
486                if track.id() == &packet.track_id {
487                    if let Samples::RTP { payload_type, payload, .. } = &packet.samples {
488                        if let Some(digit) = dtmf_detector.detect_rtp(*payload_type, payload) {
489                            debug!(track_id = track.id(), digit, "DTMF detected");
490                            event_sender
491                                .send(SessionEvent::Dtmf {
492                                    track_id: packet.track_id.to_string(),
493                                    timestamp: packet.timestamp,
494                                    digit,
495                                    refer: dtmf_detector.refer,
496                                })
497                                .ok();
498                        }
499                    }
500                    continue;
501                }
502                if suppressed {
503                    continue;
504                }
505                // Skip DTMF forwarding if source or destination has it suppressed.
506                if source_suppresses_dtmf || (is_dtmf && dtmf_detector.suppress_dtmf_forward) {
507                    continue;
508                }
509                if packet.track_id == QUEUE_HOLD_TRACK_ID && track.id() == CALLEE_TRACK_ID {
510                    continue;
511                }
512                if let Err(e) = track.send_packet(&packet).await {
513                    warn!(
514                        id = track.id(),
515                        "media_stream: Failed to send packet to track: {}", e
516                    );
517                }
518            }
519        }
520    }
521}
522
523pub struct MuteProcessor;
524
525impl MuteProcessor {
526    pub fn mute_track(track: &mut dyn Track) {
527        let chain = track.processor_chain();
528        if !chain.has_processor::<MuteProcessor>() {
529            chain.insert_processor(Box::new(MuteProcessor));
530        }
531    }
532
533    pub fn unmute_track(track: &mut dyn Track) {
534        let chain = track.processor_chain();
535        chain.remove_processor::<MuteProcessor>();
536    }
537}
538
539impl Processor for MuteProcessor {
540    fn process_frame(&mut self, frame: &mut AudioFrame) -> Result<()> {
541        match &mut frame.samples {
542            Samples::PCM { samples } => {
543                samples.fill(0);
544            }
545            // discard DTMF frames
546            Samples::RTP { payload_type, .. } if *payload_type >= 96 && *payload_type <= 127 => {
547                frame.samples = Samples::Empty;
548            }
549            _ => {}
550        }
551        Ok(())
552    }
553}
554
555pub struct HoldTrack;
556
557impl HoldTrack {
558    pub fn hold_track(track: &mut dyn Track) {
559        let chain = track.processor_chain();
560        // Remove existing processor if present
561        chain.remove_processor::<HoldProcessor>();
562        // Add a new processor with hold state set to true
563        let processor = HoldProcessor::new();
564        processor.set_hold(true);
565        chain.insert_processor(Box::new(processor));
566    }
567
568    pub fn resume_track(track: &mut dyn Track) {
569        let chain = track.processor_chain();
570        // Simply remove the hold processor to resume normal operation
571        chain.remove_processor::<HoldProcessor>();
572    }
573}