use crate::media::processor::ProcessorChain;
use crate::media::recorder::RecorderOption;
use crate::media::track::TrackConfig;
use crate::{
event::EventSender,
media::AudioFrame,
media::Samples,
media::TrackId,
media::{
stream::MediaStreamBuilder,
track::{Track, TrackPacketSender},
},
};
use anyhow::Result;
use async_trait::async_trait;
use std::sync::Arc;
use tempfile::tempdir;
use tokio::sync::Mutex;
use tokio::time::Duration;
use tracing::warn;
pub struct TestTrack {
id: TrackId,
config: TrackConfig,
sender: Option<TrackPacketSender>,
processor_chain: ProcessorChain,
received_packets: Arc<Mutex<Vec<AudioFrame>>>,
}
impl TestTrack {
pub fn new(id: TrackId) -> Self {
Self {
id,
config: TrackConfig::default(),
sender: None,
processor_chain: ProcessorChain::new(16000),
received_packets: Arc::new(Mutex::new(Vec::new())),
}
}
}
#[async_trait]
impl Track for TestTrack {
fn ssrc(&self) -> u32 {
0 }
fn id(&self) -> &TrackId {
&self.id
}
fn config(&self) -> &TrackConfig {
&self.config
}
fn processor_chain(&mut self) -> &mut ProcessorChain {
&mut self.processor_chain
}
async fn handshake(&mut self, _offer: String, _timeout: Option<Duration>) -> Result<String> {
Ok("".to_string())
}
async fn update_remote_description(&mut self, _answer: &String) -> Result<()> {
Ok(())
}
async fn start(
&mut self,
_event_sender: EventSender,
packet_sender: TrackPacketSender,
) -> Result<()> {
self.sender = Some(packet_sender);
Ok(())
}
async fn stop(&self) -> Result<()> {
Ok(())
}
async fn send_packet(&mut self, packet: &AudioFrame) -> Result<()> {
{
let mut received = self.received_packets.lock().await;
received.push(packet.clone());
}
let mut packet_clone = packet.clone();
if let Err(e) = self.processor_chain.process_frame(&mut packet_clone) {
warn!("Error processing packet: {}", e);
}
if let Some(sender) = &self.sender {
match sender.send(packet_clone) {
Ok(_) => {}
Err(e) => {
warn!("Failed to send packet: {}", e);
}
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_stream_add_track() {
let event_sender = crate::event::create_event_sender();
let stream = MediaStreamBuilder::new(event_sender).build();
let track = Box::new(TestTrack::new("test1".to_string()));
stream.update_track(track, None).await;
}
#[tokio::test]
async fn test_stream_remove_track() {
let event_sender = crate::event::create_event_sender();
let stream = MediaStreamBuilder::new(event_sender.clone())
.with_id("ms:test".to_string())
.build();
let track_id = "test1".to_string();
stream
.update_track(Box::new(TestTrack::new(track_id.clone())), None)
.await;
stream.remove_track(&track_id, false).await;
}
}
#[tokio::test]
async fn test_media_stream_basic() -> Result<()> {
let event_sender = crate::event::create_event_sender();
let stream = MediaStreamBuilder::new(event_sender).build();
let track = Box::new(TestTrack::new("test1".to_string()));
stream.update_track(track, None).await;
let handle = tokio::spawn(async move {
stream.serve().await.unwrap();
});
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
handle.abort();
Ok(())
}
#[tokio::test]
async fn test_media_stream_events() -> Result<()> {
let event_sender = crate::event::create_event_sender();
let stream = MediaStreamBuilder::new(event_sender.clone()).build();
let _events = event_sender.subscribe();
let track = Box::new(TestTrack::new("test1".to_string()));
stream.update_track(track, None).await;
let handle = tokio::spawn(async move {
stream.serve().await.unwrap();
});
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
handle.abort();
Ok(())
}
#[tokio::test]
async fn test_stream_forward_packets() -> Result<()> {
let event_sender = crate::event::create_event_sender();
let stream = MediaStreamBuilder::new(event_sender).build();
let track1 = TestTrack::new("test1".to_string());
let track2 = TestTrack::new("test2".to_string());
let track2_id = track2.id().clone();
stream.update_track(Box::new(track1), None).await;
stream.update_track(Box::new(track2), None).await;
let packet_sender = stream.packet_sender.clone();
let handle = tokio::spawn(async move {
stream.serve().await.unwrap();
});
tokio::time::sleep(Duration::from_millis(50)).await;
let samples = vec![16000, 8000, 12000, 4000];
let packet = AudioFrame {
track_id: track2_id.clone(),
timestamp: 1000,
samples: Samples::PCM { samples: samples },
sample_rate: 16000,
channels: 1,
..Default::default()
};
let _ = packet_sender.send(packet);
tokio::time::sleep(Duration::from_millis(50)).await;
handle.abort();
Ok(())
}
#[tokio::test]
async fn test_stream_recorder() -> Result<()> {
let event_sender = crate::event::create_event_sender();
let temp_dir = tempdir()?;
let file_path = temp_dir.path().join("test_recording.wav");
let stream = Arc::new(
MediaStreamBuilder::new(event_sender)
.with_recorder_config(RecorderOption {
recorder_file: file_path.to_string_lossy().to_string(),
..Default::default()
})
.build(),
);
let track1 = Box::new(TestTrack::new("test1".to_string()));
let track2 = Box::new(TestTrack::new("test2".to_string()));
let track2_id = track2.id().clone();
stream.update_track(track1, None).await;
stream.update_track(track2, None).await;
let stream_clone = stream.clone();
let handle = tokio::spawn(async move {
stream_clone.serve().await.unwrap();
});
tokio::time::sleep(Duration::from_millis(50)).await;
let packet_sender = stream.packet_sender.clone();
let samples1 = vec![3000, 6000, 9000, 12000];
let samples2 = vec![15000, 18000, 21000, 24000];
let packet1 = AudioFrame {
track_id: track2_id.clone(),
timestamp: 1000,
samples: Samples::PCM { samples: samples1 },
sample_rate: 16000,
channels: 1,
..Default::default()
};
let packet2 = AudioFrame {
track_id: track2_id,
timestamp: 1020,
samples: Samples::PCM { samples: samples2 },
sample_rate: 16000,
channels: 1,
..Default::default()
};
packet_sender.send(packet1).unwrap();
tokio::time::sleep(Duration::from_millis(20)).await;
packet_sender.send(packet2).unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
handle.abort();
Ok(())
}
#[tokio::test]
async fn test_stream_forward_payload_conversion() -> Result<()> {
let event_sender = crate::event::create_event_sender();
let stream = Arc::new(MediaStreamBuilder::new(event_sender).build());
let track1 = TestTrack::new("track1".to_string()); let track2 = TestTrack::new("track2".to_string());
stream.update_track(Box::new(track1), None).await;
stream.update_track(Box::new(track2), None).await;
let stream_clone = stream.clone();
let handle = tokio::spawn(async move {
stream_clone.serve().await.unwrap();
});
tokio::time::sleep(Duration::from_millis(50)).await;
let packet_sender = stream.packet_sender.clone();
let rtp_packet = AudioFrame {
track_id: "track2".to_string(),
timestamp: 1000,
samples: Samples::RTP {
payload_type: 0,
payload: vec![1, 2, 3, 4],
sequence_number: 1,
},
sample_rate: 16000,
channels: 1,
..Default::default()
};
let _ = packet_sender.send(rtp_packet);
let pcm_packet = AudioFrame {
track_id: "track1".to_string(),
timestamp: 2000,
samples: Samples::PCM {
samples: vec![3000, 6000, 9000, 12000],
},
sample_rate: 16000,
channels: 1,
..Default::default()
};
let _ = packet_sender.send(pcm_packet);
tokio::time::sleep(Duration::from_millis(50)).await;
handle.abort();
Ok(())
}
#[tokio::test]
async fn test_remove_processor() -> Result<()> {
use crate::media::processor::Processor;
struct TestProcessor {
#[allow(unused)]
name: String,
}
impl Processor for TestProcessor {
fn process_frame(&mut self, _frame: &mut AudioFrame) -> Result<()> {
Ok(())
}
}
let event_sender = crate::event::create_event_sender();
let stream = MediaStreamBuilder::new(event_sender).build();
let track_id = "test-track".to_string();
let mut track = TestTrack::new(track_id.clone());
track
.processor_chain
.append_processor(Box::new(TestProcessor {
name: "processor1".to_string(),
}));
track
.processor_chain
.append_processor(Box::new(TestProcessor {
name: "processor2".to_string(),
}));
stream.update_track(Box::new(track), None).await;
let result = stream.remove_processor::<TestProcessor>(&track_id).await;
assert!(result.is_ok());
Ok(())
}
#[tokio::test]
async fn test_append_processor() -> Result<()> {
use crate::media::processor::Processor;
struct AppendTestProcessor {
_value: u32,
}
impl Processor for AppendTestProcessor {
fn process_frame(&mut self, _frame: &mut AudioFrame) -> Result<()> {
Ok(())
}
}
let event_sender = crate::event::create_event_sender();
let stream = MediaStreamBuilder::new(event_sender).build();
let track_id = "test-track".to_string();
let track = TestTrack::new(track_id.clone());
stream.update_track(Box::new(track), None).await;
let processor = Box::new(AppendTestProcessor { _value: 42 });
let result = stream.append_processor(&track_id, processor).await;
assert!(result.is_ok());
Ok(())
}
#[tokio::test]
async fn test_remove_processor_from_nonexistent_track() -> Result<()> {
use crate::media::processor::Processor;
struct NonexistentProcessor;
impl Processor for NonexistentProcessor {
fn process_frame(&mut self, _frame: &mut AudioFrame) -> Result<()> {
Ok(())
}
}
let event_sender = crate::event::create_event_sender();
let stream = MediaStreamBuilder::new(event_sender).build();
let result = stream
.remove_processor::<NonexistentProcessor>(&"nonexistent-track".to_string())
.await;
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("not found"));
Ok(())
}
pub struct StoppableTestTrack {
id: TrackId,
config: TrackConfig,
processor_chain: ProcessorChain,
stopped: Arc<std::sync::atomic::AtomicBool>,
}
impl StoppableTestTrack {
pub fn new(id: TrackId, stopped: Arc<std::sync::atomic::AtomicBool>) -> Self {
Self {
id,
config: TrackConfig::default(),
processor_chain: ProcessorChain::new(16000),
stopped,
}
}
}
#[async_trait]
impl Track for StoppableTestTrack {
fn ssrc(&self) -> u32 {
0
}
fn id(&self) -> &TrackId {
&self.id
}
fn config(&self) -> &TrackConfig {
&self.config
}
fn processor_chain(&mut self) -> &mut ProcessorChain {
&mut self.processor_chain
}
async fn handshake(&mut self, _offer: String, _timeout: Option<Duration>) -> Result<String> {
Ok("".to_string())
}
async fn update_remote_description(&mut self, _answer: &String) -> Result<()> {
Ok(())
}
async fn start(
&mut self,
_event_sender: EventSender,
_packet_sender: TrackPacketSender,
) -> Result<()> {
Ok(())
}
async fn stop(&self) -> Result<()> {
self.stopped
.store(true, std::sync::atomic::Ordering::SeqCst);
Ok(())
}
async fn send_packet(&mut self, _packet: &AudioFrame) -> Result<()> {
Ok(())
}
}
#[tokio::test]
async fn test_cleanup_drains_all_tracks() -> Result<()> {
let event_sender = crate::event::create_event_sender();
let stream = MediaStreamBuilder::new(event_sender)
.with_id("test-cleanup".to_string())
.build();
let stopped1 = Arc::new(std::sync::atomic::AtomicBool::new(false));
let stopped2 = Arc::new(std::sync::atomic::AtomicBool::new(false));
let stopped3 = Arc::new(std::sync::atomic::AtomicBool::new(false));
stream
.update_track(
Box::new(StoppableTestTrack::new(
"track1".to_string(),
stopped1.clone(),
)),
None,
)
.await;
stream
.update_track(
Box::new(StoppableTestTrack::new(
"track2".to_string(),
stopped2.clone(),
)),
None,
)
.await;
stream
.update_track(
Box::new(StoppableTestTrack::new(
"track3".to_string(),
stopped3.clone(),
)),
None,
)
.await;
assert_eq!(stream.track_count().await, 3);
stream.cleanup().await.unwrap();
assert!(
stopped1.load(std::sync::atomic::Ordering::SeqCst),
"track1 should have been stopped"
);
assert!(
stopped2.load(std::sync::atomic::Ordering::SeqCst),
"track2 should have been stopped"
);
assert!(
stopped3.load(std::sync::atomic::Ordering::SeqCst),
"track3 should have been stopped"
);
assert_eq!(
stream.track_count().await,
0,
"all tracks should be drained after cleanup"
);
Ok(())
}
#[tokio::test]
async fn test_cleanup_is_idempotent() -> Result<()> {
let event_sender = crate::event::create_event_sender();
let stream = MediaStreamBuilder::new(event_sender)
.with_id("test-idempotent".to_string())
.build();
let stopped = Arc::new(std::sync::atomic::AtomicBool::new(false));
stream
.update_track(
Box::new(StoppableTestTrack::new(
"track1".to_string(),
stopped.clone(),
)),
None,
)
.await;
stream.cleanup().await.unwrap();
assert!(stopped.load(std::sync::atomic::Ordering::SeqCst));
assert_eq!(stream.track_count().await, 0);
stream.cleanup().await.unwrap();
assert_eq!(stream.track_count().await, 0);
Ok(())
}