use anyhow::{Context, Result};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use super::codec::OpusEncoder;
use super::codec::{CHANNELS, SAMPLES_PER_FRAME};
const AUDIO_STREAM_ID: u8 = 0x11;
const MAGIC: &[u8; 4] = b"PCAU";
pub struct AudioSender {
encoder: OpusEncoder,
origin: Instant,
video_pts_us: Arc<AtomicU64>,
sent: AtomicU64,
dropped: AtomicU64,
}
impl std::fmt::Debug for AudioSender {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AudioSender")
.field("sent", &self.sent.load(Ordering::Relaxed))
.field("dropped", &self.dropped.load(Ordering::Relaxed))
.finish()
}
}
impl AudioSender {
pub fn new(_unused: (), video_pts_us: Arc<AtomicU64>) -> Result<Self> {
Ok(Self {
encoder: OpusEncoder::new()?,
origin: Instant::now(),
video_pts_us,
sent: AtomicU64::new(0),
dropped: AtomicU64::new(0),
})
}
pub fn note_video_capture(&self, pts_us: u64) {
self.video_pts_us.store(pts_us, Ordering::Relaxed);
}
pub fn encode_frame(&mut self, pcm: &[f32]) -> Result<Vec<u8>> {
self.encoder.encode(pcm)
}
pub async fn send_encoded<W>(&self, stream: &mut W, frame: &EncodedFrame) -> Result<()>
where
W: AsyncWrite + Unpin,
{
let mut out = Vec::with_capacity(16 + frame.opus.len());
out.extend_from_slice(&frame.pts_us.to_le_bytes());
out.extend_from_slice(&frame.pcm_len.to_le_bytes());
out.extend_from_slice(&(frame.opus.len() as u32).to_le_bytes());
out.extend_from_slice(&frame.opus);
stream.write_all(&out).await?;
self.sent.fetch_add(1, Ordering::Relaxed);
Ok(())
}
pub async fn send_frame<W>(
&mut self,
stream: &mut W,
pcm: &[f32],
capture: Instant,
) -> Result<()>
where
W: AsyncWrite + Unpin,
{
let encoded = self.encoder.encode(pcm)?;
let pts_us = {
let v = self.video_pts_us.load(Ordering::Relaxed);
if v > 0 {
v
} else {
capture.duration_since(self.origin).as_micros() as u64
}
};
let mut frame = Vec::with_capacity(MAGIC.len() + 20 + encoded.len());
frame.extend_from_slice(MAGIC);
frame.extend_from_slice(&pts_us.to_le_bytes());
frame.extend_from_slice(&(pcm.len() as u32).to_le_bytes());
frame.extend_from_slice(&(encoded.len() as u32).to_le_bytes());
frame.extend_from_slice(&encoded);
stream.write_all(&frame).await?;
self.sent.fetch_add(1, Ordering::Relaxed);
Ok(())
}
pub fn sent(&self) -> u64 {
self.sent.load(Ordering::Relaxed)
}
pub fn dropped(&self) -> u64 {
self.dropped.load(Ordering::Relaxed)
}
pub fn note_drop(&self) {
self.dropped.fetch_add(1, Ordering::Relaxed);
}
}
pub async fn open_stream<W>(stream: &mut W) -> Result<()>
where
W: AsyncWrite + Unpin,
{
let mut header = Vec::with_capacity(MAGIC.len() + 8);
header.extend_from_slice(MAGIC);
header.extend_from_slice(&(SAMPLES_PER_FRAME as u32).to_le_bytes());
header.extend_from_slice(&[2u8]); stream.write_all(&header).await?;
Ok(())
}
pub struct AudioReceiver {
decoder: super::codec::OpusDecoder,
seen: u64,
gaps: u64,
last_pts_us: u64,
}
impl std::fmt::Debug for AudioReceiver {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AudioReceiver")
.field("seen", &self.seen)
.field("gaps", &self.gaps)
.finish()
}
}
impl AudioReceiver {
pub fn new() -> Result<Self> {
Ok(Self {
decoder: super::codec::OpusDecoder::new()?,
seen: 0,
gaps: 0,
last_pts_us: 0,
})
}
pub async fn read_header<R>(reader: &mut R) -> Result<bool>
where
R: AsyncRead + Unpin,
{
let mut magic = [0u8; 4];
reader.read_exact(&mut magic).await?;
if &magic != MAGIC {
return Ok(false);
}
let mut rest = [0u8; 5];
reader.read_exact(&mut rest).await?;
Ok(true)
}
pub async fn read_frame<R>(&mut self, reader: &mut R) -> Result<Option<DecodedAudio>>
where
R: AsyncRead + Unpin,
{
let mut head = [0u8; 16];
match reader.read_exact(&mut head).await {
Ok(_) => {}
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(None),
Err(e) => return Err(e.into()),
}
let pts_us = u64::from_le_bytes(head[0..8].try_into().unwrap());
let pcm_len = u32::from_le_bytes(head[8..12].try_into().unwrap()) as usize;
let enc_len = u32::from_le_bytes(head[12..16].try_into().unwrap()) as usize;
if pcm_len > SAMPLES_PER_FRAME * 2 * 4 {
anyhow::bail!(
"audio frame claims {pcm_len} samples, more than one {}ms frame can hold",
super::FRAME_MS
);
}
let mut encoded = vec![0u8; enc_len];
reader.read_exact(&mut encoded).await?;
if self.last_pts_us > pts_us {
self.gaps += 1;
}
self.last_pts_us = pts_us;
self.seen += 1;
let mut pcm = vec![0f32; pcm_len];
let per_channel = self.decoder.decode(&encoded, &mut pcm)?;
let expected = pcm_len / CHANNELS;
if per_channel != expected {
anyhow::bail!("Opus decoded {per_channel} samples per channel, expected {expected}");
}
Ok(Some(DecodedAudio { pts_us, pcm }))
}
pub fn frames_seen(&self) -> u64 {
self.seen
}
pub fn gaps(&self) -> u64 {
self.gaps
}
}
#[derive(Debug, Clone)]
pub struct EncodedFrame {
pub pts_us: u64,
pub pcm_len: u32,
pub opus: Arc<Vec<u8>>,
}
#[derive(Debug, Clone)]
pub struct DecodedAudio {
pub pts_us: u64,
pub pcm: Vec<f32>,
}
impl DecodedAudio {
pub fn pts(&self) -> Duration {
Duration::from_micros(self.pts_us)
}
}
pub const AUDIO_UNI_STREAM_TYPE: u8 = AUDIO_STREAM_ID;
pub fn encode_frame_for_test(pcm: &[f32]) -> Result<Vec<u8>> {
let mut encoder = OpusEncoder::new().context("test encoder")?;
encoder.encode(pcm)
}
#[cfg(test)]
mod tests {
use super::*;
fn tone(frames: usize) -> Vec<f32> {
(0..frames * SAMPLES_PER_FRAME)
.map(|i| ((i as f32) * 0.01).sin() * 0.2)
.collect()
}
#[test]
fn an_audio_frame_round_trips_through_opus() {
let pcm = tone(1);
let encoded = encode_frame_for_test(&pcm).unwrap();
let mut decoder = super::super::codec::OpusDecoder::new().unwrap();
let mut out = vec![0f32; SAMPLES_PER_FRAME];
let n = decoder.decode(&encoded, &mut out).unwrap();
assert!(n > 0, "a round trip must produce samples");
}
#[tokio::test]
async fn a_frame_carries_the_shared_video_timestamp() {
let video_pts = Arc::new(AtomicU64::new(0));
let mut sender = AudioSender::new((), video_pts.clone()).unwrap();
let mut stream: Vec<u8> = Vec::new();
sender
.send_frame(&mut stream, &tone(1), Instant::now())
.await
.unwrap();
assert_eq!(&stream[0..4], MAGIC);
assert_eq!(sender.sent(), 1);
video_pts.store(123_456, Ordering::Relaxed);
let mut second: Vec<u8> = Vec::new();
sender
.send_frame(&mut second, &tone(1), Instant::now())
.await
.unwrap();
let pts = u64::from_le_bytes(second[4..12].try_into().unwrap());
assert_eq!(pts, 123_456, "audio must ride the video clock");
}
fn frame(pts_us: u64, pcm: &[f32], encoder: &mut OpusEncoder) -> Vec<u8> {
let encoded = encoder.encode(pcm).unwrap();
let mut out = Vec::new();
out.extend_from_slice(&pts_us.to_le_bytes());
out.extend_from_slice(&(pcm.len() as u32).to_le_bytes());
out.extend_from_slice(&(encoded.len() as u32).to_le_bytes());
out.extend_from_slice(&encoded);
out
}
#[tokio::test]
async fn a_stream_header_is_distinguishable_from_anything_else() {
let mut wire: Vec<u8> = Vec::new();
open_stream(&mut wire).await.unwrap();
let mut cursor = wire.as_slice();
assert!(AudioReceiver::read_header(&mut cursor).await.unwrap());
let mut other = b"NOTAU".to_vec();
other.extend_from_slice(&[0u8; 8]);
let mut c2 = other.as_slice();
assert!(!AudioReceiver::read_header(&mut c2).await.unwrap());
}
#[tokio::test]
async fn the_receiver_reports_gaps_instead_of_stalling() {
let mut encoder = OpusEncoder::new().unwrap();
let pcm = tone(1);
let mut wire: Vec<u8> = Vec::new();
open_stream(&mut wire).await.unwrap();
wire.extend(frame(1_000, &pcm, &mut encoder));
wire.extend(frame(0, &pcm, &mut encoder));
wire.extend(frame(2_000, &pcm, &mut encoder));
let mut receiver = AudioReceiver::new().unwrap();
let mut cursor = wire.as_slice();
assert!(AudioReceiver::read_header(&mut cursor).await.unwrap());
let mut seen = 0;
while let Some(frame) = receiver.read_frame(&mut cursor).await.unwrap() {
assert!(!frame.pcm.is_empty(), "a frame must decode to samples");
seen += 1;
}
assert_eq!(seen, 3, "a gap must not stop the stream");
assert_eq!(receiver.gaps(), 1, "the reorder must be counted");
}
#[tokio::test]
async fn a_lying_frame_length_is_refused_before_allocating() {
let mut wire: Vec<u8> = Vec::new();
open_stream(&mut wire).await.unwrap();
wire.extend_from_slice(&1u64.to_le_bytes()); wire.extend_from_slice(&u32::MAX.to_le_bytes()); wire.extend_from_slice(&0u32.to_le_bytes()); let mut receiver = AudioReceiver::new().unwrap();
let mut cursor = wire.as_slice();
assert!(AudioReceiver::read_header(&mut cursor).await.unwrap());
let err = receiver
.read_frame(&mut cursor)
.await
.unwrap_err()
.to_string();
assert!(err.contains("more than one"), "unhelpful: {err}");
}
}