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 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 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 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 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 chain.remove_processor::<HoldProcessor>();
562 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 chain.remove_processor::<HoldProcessor>();
572 }
573}