diff --git a/examples/Cargo.lock b/examples/Cargo.lock index 98bb3e1..6cf6bb0 100644 --- a/examples/Cargo.lock +++ b/examples/Cargo.lock @@ -2377,6 +2377,17 @@ dependencies = [ "winapi-util", ] +[[package]] +name = "save_to_disk" +version = "0.1.0" +dependencies = [ + "bytes", + "futures", + "livekit", + "tokio", + "tokio-util", +] + [[package]] name = "schannel" version = "0.1.21" diff --git a/examples/save_to_disk/Cargo.toml b/examples/save_to_disk/Cargo.toml new file mode 100644 index 0000000..9821761 --- /dev/null +++ b/examples/save_to_disk/Cargo.toml @@ -0,0 +1,13 @@ +[package] +name = "save_to_disk" +version = "0.1.0" +edition = "2021" + +# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html + +[dependencies] +tokio = { version = "1", features = ["full"] } +livekit = { path = "../../livekit", version = "0.1.1" } +bytes = "1.4.0" +tokio-util = "0.7.8" +futures = "0.3.28" diff --git a/examples/save_to_disk/src/main.rs b/examples/save_to_disk/src/main.rs new file mode 100644 index 0000000..61b36d3 --- /dev/null +++ b/examples/save_to_disk/src/main.rs @@ -0,0 +1,143 @@ +use bytes::{BufMut, BytesMut}; +use futures::StreamExt; +use livekit::prelude::*; +use livekit::webrtc::audio_stream::native::NativeAudioStream; +use std::env; +use tokio::fs::File; +use tokio::io::{AsyncWriteExt, BufWriter}; + +const WAV_HEADER_SIZE: usize = 44; +const FILE_PATH: &str = "record.wav"; + +#[derive(Debug, Clone, Copy)] +pub struct WavHeader { + pub sample_rate: u32, + pub bit_depth: u16, + pub num_channels: u16, +} + +pub struct WavWriter { + header: WavHeader, + data: BytesMut, + writer: BufWriter, +} + +impl WavWriter { + pub async fn create>( + path: P, + header: WavHeader, + ) -> Result { + let file = File::create(path).await?; + let writer = BufWriter::new(file); + + let mut wav_writer = WavWriter { + header, + data: BytesMut::new(), + writer, + }; + + wav_writer.write_header()?; + Ok(wav_writer) + } + + fn write_header(&mut self) -> Result<(), std::io::Error> { + let byte_rate = (self.header.sample_rate + * self.header.bit_depth as u32 + * self.header.num_channels as u32); + + let block_align = byte_rate / self.header.sample_rate as u32; + + self.data.put_slice(b"RIFF"); + self.data.put_u32_le(0); // Placeholder for file size + self.data.put_slice(b"WAVE"); + self.data.put_slice(b"fmt "); + self.data.put_u32_le(16); // Subchunk1Size (16 for PCM) + self.data.put_u16_le(1); // AudioFormat (1 for PCM) + self.data.put_u16_le(self.header.num_channels); + self.data.put_u32_le(self.header.sample_rate); + self.data.put_u32_le(byte_rate); + self.data.put_u16_le(32); + self.data.put_u16_le(self.header.bit_depth); + self.data.put_slice(b"data"); + self.data.put_u32_le(0); // Placeholder for data size + + assert_eq!(self.data.len(), WAV_HEADER_SIZE); + + Ok(()) + } + + pub async fn write_sample(&mut self, sample: i16) -> Result<(), std::io::Error> { + self.data.put_i16_le(sample); + Ok(()) + } + + pub async fn finalize(mut self) -> Result<(), std::io::Error> { + let data_size = self.data.len() as u32 - WAV_HEADER_SIZE as u32; + let file_size = data_size + WAV_HEADER_SIZE as u32 - 8; + self.data.as_mut()[4..8].copy_from_slice(&file_size.to_le_bytes()); + self.data.as_mut()[40..44].copy_from_slice(&data_size.to_le_bytes()); + + self.writer.write_all(&self.data).await?; + self.writer.flush().await?; + Ok(()) + } +} + +#[tokio::main] +async fn main() { + let url = env::var("LIVEKIT_URL").expect("LIVEKIT_URL is not set"); + let token = env::var("LIVEKIT_TOKEN").expect("LIVEKIT_TOKEN is not set"); + + let (room, mut rx) = Room::connect(&url, &token).await.unwrap(); + let session = room.session(); + println!("Connected to room: {} - {}", session.name(), session.sid()); + + while let Some(msg) = rx.recv().await { + match msg { + RoomEvent::TrackSubscribed { + track, + publication: _, + participant: _, + } => { + if let RemoteTrack::Audio(audio_track) = track { + record_track(audio_track).await.unwrap(); + break; + } + } + _ => {} + } + } + + println!("Done"); +} + +async fn record_track(audio_track: RemoteAudioTrack) -> Result<(), std::io::Error> { + println!("Recording track {:?}", audio_track.sid()); + let rtc_track = audio_track.rtc_track(); + + // TODO(theomonnom): Remove hardcoded values + let header = WavHeader { + sample_rate: 48000, + bit_depth: 16, + num_channels: 1, + }; + + let mut wav_writer = WavWriter::create(FILE_PATH, header).await?; + let mut audio_stream = NativeAudioStream::new(rtc_track); + + let max_record = 5 * header.sample_rate * header.num_channels as u32; + let mut sample_count = 0; + 'recv_loop: while let Some(frame) = audio_stream.next().await { + for sample in frame.data { + wav_writer.write_sample(sample).await.unwrap(); + sample_count += 1; + + if sample_count >= max_record { + break 'recv_loop; + } + } + } + + wav_writer.finalize().await?; + Ok(()) +} diff --git a/livekit-ffi/src/server/audio_frame.rs b/livekit-ffi/src/server/audio_frame.rs index 0070820..eb61cb4 100644 --- a/livekit-ffi/src/server/audio_frame.rs +++ b/livekit-ffi/src/server/audio_frame.rs @@ -62,7 +62,7 @@ impl FfiAudioSream { close_tx, track_sid, }; - tokio::spawn(Self::native_audio_stream_task( + server.async_runtime.spawn(Self::native_audio_stream_task( server, audio_stream.handle_id, NativeAudioStream::new(track), diff --git a/livekit-ffi/src/server/mod.rs b/livekit-ffi/src/server/mod.rs index a201977..7f0e2b6 100644 --- a/livekit-ffi/src/server/mod.rs +++ b/livekit-ffi/src/server/mod.rs @@ -200,7 +200,7 @@ impl FfiServer { publish: proto::PublishTrackRequest, ) -> FfiResult { let async_id = self.next_id() as FfiAsyncId; - tokio::spawn(async move { + self.async_runtime.spawn(async move { let res = async { let room_handle = publish .room_handle diff --git a/livekit-ffi/src/server/room.rs b/livekit-ffi/src/server/room.rs index 9909f92..a80b085 100644 --- a/livekit-ffi/src/server/room.rs +++ b/livekit-ffi/src/server/room.rs @@ -21,7 +21,7 @@ impl FfiRoom { let session = room.session(); let next_id = server.next_id() as FfiHandleId; - let handle = tokio::spawn(room_task( + let handle = server.async_runtime.spawn(room_task( server, session.clone(), next_id, @@ -60,9 +60,11 @@ async fn room_task( mut events: mpsc::UnboundedReceiver, mut close_rx: oneshot::Receiver<()>, ) { - tokio::spawn(participant_task(Participant::Local( - session.local_participant(), - ))); + server + .async_runtime + .spawn(participant_task(Participant::Local( + session.local_participant(), + ))); loop { tokio::select! { @@ -73,7 +75,7 @@ async fn room_task( match event { RoomEvent::ParticipantConnected(p) => { - tokio::spawn(participant_task(Participant::Remote(p))); + server.async_runtime.spawn(participant_task(Participant::Remote(p))); } _ => {} } diff --git a/livekit-ffi/src/server/video_frame.rs b/livekit-ffi/src/server/video_frame.rs index 73c1691..157c047 100644 --- a/livekit-ffi/src/server/video_frame.rs +++ b/livekit-ffi/src/server/video_frame.rs @@ -62,7 +62,7 @@ impl FfiVideoStream { stream_type, track_sid, }; - tokio::spawn(Self::native_video_stream_task( + server.async_runtime.spawn(Self::native_video_stream_task( server, video_stream.handle_id, NativeVideoStream::new(track), diff --git a/webrtc-sys/src/audio_device.cpp b/webrtc-sys/src/audio_device.cpp index 7bc35c5..392f73a 100644 --- a/webrtc-sys/src/audio_device.cpp +++ b/webrtc-sys/src/audio_device.cpp @@ -16,7 +16,7 @@ #include "livekit/audio_device.h" -const int kBitsPerSample = 16; +const int kBytesPerSample = 2; const int kSampleRate = 48000; const int kChannels = 2; const int kSamplesPer10Ms = kSampleRate / 100; @@ -54,13 +54,14 @@ int32_t AudioDevice::Init() { if (playing_) { int64_t elapsed_time_ms = -1; int64_t ntp_time_ms = -1; + size_t n_samples_out = 0; void* data = data_.data(); // Request the AudioData, otherwise WebRTC will ignore the packets. // 10ms of audio data. - audio_transport_->PullRenderData(kBitsPerSample, kSampleRate, - kChannels, kSamplesPer10Ms, data, - &elapsed_time_ms, &ntp_time_ms); + audio_transport_->NeedMorePlayData( + kSamplesPer10Ms, kBytesPerSample, kChannels, kSampleRate, data, + n_samples_out, &elapsed_time_ms, &ntp_time_ms); } return webrtc::TimeDelta::Millis(10);