refactor: 优化音频读取与编码传输

This commit is contained in:
Del Wang
2026-01-06 12:08:06 +08:00
parent c0bc4b8ae6
commit 0125a1478e
5 changed files with 178 additions and 246 deletions
+1 -14
View File
@@ -512,29 +512,16 @@ impl Server {
let session = self.sessions.get(&addr).context("Session not found")?;
let reader = WavReader::open(file_path)?;
let opus_rate = if reader.sample_rate > 24000 {
48000
} else {
16000
};
let config = AudioConfig {
sample_rate: opus_rate,
channels: reader.channels,
frame_size: (opus_rate / 50) as usize, // 20ms
..AudioConfig::music_48k()
};
// 通知客户端准备接收音频
session
.send(&ControlPacket::StartPlayback {
config: config.clone(),
config: reader.config.clone(),
})
.await?;
// 创建播放流
let handle = FilePlaybackStream::spawn(
config,
reader,
self.audio_bus.socket(),
session.audio_addr,
+109 -153
View File
@@ -28,6 +28,7 @@ use crate::net::protocol::AudioPacket;
use crate::net::sync::now_us;
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::broadcast;
use tokio_util::sync::CancellationToken;
@@ -61,6 +62,96 @@ impl Drop for StreamHandle {
}
}
/// 音频流发送控制器
pub struct StreamSender {
socket: Arc<AudioSocket>,
target: SocketAddr,
codec: OpusCodec,
// 缓存区:避免重复分配内存
opus_buffer: Vec<u8>,
// 状态变量
seq: u32,
stream_start_ts: Option<u128>,
// 配置参数
input_len: usize,
frame_duration_us: u128,
max_lead_us: u128, // 允许的最大超前时间,例如 1_000_000 (1s)
}
impl StreamSender {
pub fn new(
config: AudioConfig,
socket: Arc<AudioSocket>,
target: SocketAddr,
) -> anyhow::Result<Self> {
let codec = OpusCodec::new(&config)?;
let input_len: usize = config.frame_size * (config.channels as usize);
// 预分配缓冲区
let opus_buffer = vec![0u8; 4096];
let frame_duration_us =
(config.frame_size as f64 / config.sample_rate as f64 * 1_000_000.0) as u128;
Ok(Self {
socket,
target,
codec,
opus_buffer,
seq: 0,
stream_start_ts: None,
input_len,
frame_duration_us,
max_lead_us: 1_000_000,
})
}
/// 发送音频帧
pub async fn send(&mut self, pcm_input: &[i16]) -> anyhow::Result<()> {
// 0. 处理输入(不足时填充静音)
let pcm_buffer = if pcm_input.len() < self.input_len {
let mut pcm_buffer = vec![0i16; self.input_len];
pcm_buffer[..pcm_input.len()].copy_from_slice(pcm_input);
pcm_buffer
} else {
pcm_input.to_vec()
};
// 1. 编码
let encoded_len = self.codec.encode(&pcm_buffer, &mut self.opus_buffer)?;
// 2. 时间戳处理
let now = now_us();
let start_ts = *self.stream_start_ts.get_or_insert(now);
let target_ts = start_ts + (self.seq as u128 * self.frame_duration_us);
// 3. 构建并发送
let packet = AudioPacket {
seq: self.seq,
timestamp: target_ts,
data: self.opus_buffer[..encoded_len].to_vec(),
};
self.socket.send(&packet, self.target).await?;
self.seq += 1;
// 4. 平滑流控
// 如果发送进度超过当前时间 + 允许的缓冲量,则进行睡眠
if target_ts > now + self.max_lead_us {
let drift = target_ts - (now + self.max_lead_us);
let sleep_ms = (drift / 1000).min(100) as u64;
if sleep_ms > 0 {
tokio::time::sleep(Duration::from_millis(sleep_ms)).await;
}
}
Ok(())
}
}
/// 文件播放流 - 从 WAV 文件读取并发送到指定客户端
pub struct FilePlaybackStream;
@@ -74,7 +165,6 @@ impl FilePlaybackStream {
/// * `target` - 目标客户端地址
/// * `parent_cancel` - 父级取消令牌(用于 session 级别取消)
pub fn spawn(
config: AudioConfig,
reader: WavReader,
socket: Arc<AudioSocket>,
target: SocketAddr,
@@ -84,7 +174,7 @@ impl FilePlaybackStream {
let token = cancel.clone();
tokio::spawn(async move {
if let Err(e) = Self::run(config, reader, socket, target, token).await {
if let Err(e) = Self::run(reader, socket, target, token).await {
eprintln!("[FilePlayback] Error: {}", e);
}
});
@@ -93,86 +183,32 @@ impl FilePlaybackStream {
}
async fn run(
config: AudioConfig,
mut reader: WavReader,
socket: Arc<AudioSocket>,
target: SocketAddr,
cancel: CancellationToken,
) -> anyhow::Result<()> {
#[cfg(not(target_os = "linux"))]
{
use crate::audio::reader::AudioReader;
let mut codec = OpusCodec::new(&config)?;
let mut pcm = vec![0i16; config.frame_size * config.channels as usize];
let mut opus_buf = vec![0u8; 4096];
println!("[FilePlayback] Started -> {}", target);
println!("[FilePlayback] Started -> {}", target);
let mut sender = StreamSender::new(reader.config.clone(), socket, target)?;
let mut seq = 0u32;
let delay_us = 0_000; // 100ms 基础延迟
let mut stream_start_ts = 0;
let frame_duration_us =
(config.frame_size as f64 / config.sample_rate as f64 * 1_000_000.0) as u128;
let mut reader = AudioReader::new("temp/test.wav")?;
loop {
tokio::select! {
_ = cancel.cancelled() => break,
result = async{
if let Some((left_pcm, right_pcm)) = reader.read_chunk(config.frame_size)?{
let actual_len = left_pcm.len();
// 1. 先将 pcm 缓冲区清零(处理尾帧时的静音填充)
pcm.fill(0);
// 2. 只循环实际读取到的长度
if config.channels == 2 {
for i in 0..actual_len {
pcm[i * 2] = left_pcm[i];
pcm[i * 2 + 1] = right_pcm[i];
}
} else {
pcm[..actual_len].copy_from_slice(&left_pcm);
}
// 3. 编码时依然使用固定的 frame_size
let input_len = config.frame_size * config.channels as usize;
if let Ok(len) = codec.encode(&pcm[..input_len], &mut opus_buf) {
let now = now_us();
if stream_start_ts == 0 {
// 初始化流开始时间戳
stream_start_ts = now;
}
let target_ts = stream_start_ts + ((seq) as u128 * frame_duration_us) + delay_us;
let packet = AudioPacket {
seq,
timestamp: target_ts,
data: opus_buf[..len].to_vec(),
};
let _ = socket.send(&packet, target).await;
seq += 1;
// 音频发送时长超过音频播放 1s 时进行等待(控制数据超前缓冲 1s)
if target_ts > now + 1_000_000 {
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
}
return Ok(())
}
}
Err(anyhow::anyhow!("Unexpected EOF"))
} => {
if let Err(e) = result {
if e.to_string() == "EOF" {
break;
}
loop {
tokio::select! {
_ = cancel.cancelled() => break,
result = async {
// 读取一帧 PCM 数据
if let Some(pcm) = reader.read_one_frame()? {
// 编码并发送
sender.send(pcm).await
} else {
Err(anyhow::anyhow!("EOF"))
}
} => {
if let Err(e) = result {
if e.to_string() != "EOF" {
eprintln!("[FilePlayback] Error: {}", e);
break;
}
break;
}
}
}
@@ -265,83 +301,3 @@ impl RecorderStream {
Ok(())
}
}
/// 音频转发流 - 从总线订阅并转发到指定客户端
/// 用于实现"监听"功能或者服务端音频源推送
pub struct ForwardStream;
impl ForwardStream {
/// 启动转发流
///
/// # Arguments
/// * `socket` - UDP socket
/// * `target` - 目标地址
/// * `bus_rx` - 音频总线接收器
/// * `source_filter` - 源过滤(可选)
/// * `parent_cancel` - 父级取消令牌
pub fn spawn(
socket: Arc<AudioSocket>,
target: SocketAddr,
bus_rx: broadcast::Receiver<AudioFrame>,
source_filter: Option<SocketAddr>,
parent_cancel: CancellationToken,
) -> StreamHandle {
let cancel = parent_cancel.child_token();
let token = cancel.clone();
tokio::spawn(async move {
if let Err(e) = Self::run(socket, target, bus_rx, source_filter, token).await {
eprintln!("[Forward] Error: {}", e);
}
});
StreamHandle::new(cancel)
}
async fn run(
socket: Arc<AudioSocket>,
target: SocketAddr,
mut bus_rx: broadcast::Receiver<AudioFrame>,
source_filter: Option<SocketAddr>,
cancel: CancellationToken,
) -> anyhow::Result<()> {
println!(
"[Forward] Started -> {} (filter: {:?})",
target, source_filter
);
loop {
tokio::select! {
_ = cancel.cancelled() => break,
result = bus_rx.recv() => {
match result {
Ok(frame) => {
// 应用源过滤
if let Some(filter_addr) = &source_filter {
if frame.source.as_ref() != Some(filter_addr) {
continue;
}
}
// 不转发给自己
if frame.source.as_ref() == Some(&target) {
continue;
}
let _ = socket.send(&frame.packet, target).await;
}
Err(broadcast::error::RecvError::Lagged(n)) => {
eprintln!("[Forward] Lagged {} frames", n);
}
Err(broadcast::error::RecvError::Closed) => {
break;
}
}
}
}
}
println!("[Forward] Stopped");
Ok(())
}
}
+4 -14
View File
@@ -9,24 +9,14 @@ pub struct OpusCodec {
impl OpusCodec {
pub fn new(config: &AudioConfig) -> Result<Self> {
// Opus 仅支持这些采样率,强制进行转换以防万一
let opus_rate = match config.sample_rate {
8000 => 8000,
12000 => 12000,
16000 => 16000,
24000 => 24000,
48000 => 48000,
_ => {
let fallback = if config.sample_rate < 24000 {
16000
} else {
48000
};
println!(
"Warning: Opus does not support {}Hz, falling back to {}Hz",
config.sample_rate, fallback
);
fallback
return Err(anyhow::anyhow!(
"Unsupported sample rate for Opus: {}",
config.sample_rate
));
}
};
+10 -7
View File
@@ -1,4 +1,5 @@
#![cfg(not(target_os = "linux"))]
use anyhow::{Context, Result};
use std::fs::File;
use std::path::Path;
@@ -14,22 +15,22 @@ pub struct AudioReader {
decoder: Box<dyn Decoder>,
track_id: u32,
sample_buf: Option<SampleBuffer<i16>>,
channels: usize,
left_buffer: Vec<i16>,
right_buffer: Vec<i16>,
pub channels: usize,
pub sample_rate: u32,
}
impl AudioReader {
pub fn new(path: &str) -> Result<Self> {
pub fn new(path: impl AsRef<Path>) -> Result<Self> {
let path_ref = path.as_ref();
let src =
File::open(Path::new(path)).context(format!("Failed to open audio file: {}", path))?;
File::open(path_ref).context(format!("Failed to open audio file: {:?}", path_ref))?;
let mss = MediaSourceStream::new(Box::new(src), Default::default());
let mut hint = Hint::new();
if path.ends_with(".wav") {
hint.with_extension("wav");
} else if path.ends_with(".mp3") {
hint.with_extension("mp3");
if let Some(ext) = path_ref.extension().and_then(|s| s.to_str()) {
hint.with_extension(ext);
}
let probed = symphonia::default::get_probe()
@@ -54,6 +55,7 @@ impl AudioReader {
.context("Failed to create decoder")?;
let channels = track.codec_params.channels.map(|c| c.count()).unwrap_or(1);
let sample_rate = track.codec_params.sample_rate.unwrap_or(44100);
Ok(Self {
format,
@@ -61,6 +63,7 @@ impl AudioReader {
track_id,
sample_buf: None,
channels,
sample_rate,
left_buffer: Vec::new(),
right_buffer: Vec::new(),
})
+54 -58
View File
@@ -1,6 +1,7 @@
use crate::audio::config::AudioConfig;
use anyhow::Result;
use std::fs::File;
use std::io::{BufReader, BufWriter, Read, Seek, SeekFrom, Write};
use std::io::{BufWriter, Seek, SeekFrom, Write};
pub struct WavWriter {
writer: BufWriter<File>,
@@ -64,72 +65,67 @@ impl WavWriter {
}
pub struct WavReader {
reader: BufReader<File>,
pub sample_rate: u32,
pub channels: u16,
pub data_size: u32,
#[cfg(not(target_os = "linux"))]
reader: crate::audio::reader::AudioReader,
pcm_buffer: Vec<i16>,
pub config: AudioConfig,
}
impl WavReader {
pub fn open(path: &str) -> Result<Self> {
let file = File::open(path)?;
let mut reader = BufReader::new(file);
let mut header = [0u8; 44];
reader.read_exact(&mut header)?;
#[cfg(not(target_os = "linux"))]
{
let reader = crate::audio::reader::AudioReader::new(path)?;
let channels = reader.channels as u16;
if &header[0..4] != b"RIFF" || &header[8..12] != b"WAVE" {
return Err(anyhow::anyhow!("Not a WAV file"));
let config = AudioConfig {
channels,
sample_rate: reader.sample_rate,
frame_size: reader.sample_rate as usize / 50, // 20ms
..AudioConfig::music_48k()
};
let input_len = config.frame_size * (channels as usize);
Ok(Self {
reader,
config,
pcm_buffer: vec![0i16; input_len],
})
}
#[cfg(target_os = "linux")]
{
Err(anyhow::anyhow!(
"WavReader is only supported on non-Linux platforms"
))
}
let channels = u16::from_le_bytes([header[22], header[23]]);
let sample_rate = u32::from_le_bytes([header[24], header[25], header[26], header[27]]);
let data_size = u32::from_le_bytes([header[40], header[41], header[42], header[43]]);
Ok(Self {
reader,
sample_rate,
channels,
data_size,
})
}
/// 读取指定的采样数。如果是立体声,则返回左右声道独立的 Vec。
/// chunk_size: 每个声道需要读取的采样点数。
pub fn read_chunk(&mut self, chunk_size: usize) -> Result<Option<(Vec<i16>, Vec<i16>)>> {
let channels = self.channels as usize;
let total_samples_to_read = chunk_size * channels;
let mut bytes = vec![0u8; total_samples_to_read * 2]; // 预分配完整大小的字节数组
let bytes_read = self.reader.read(&mut bytes)?;
// 如果连 1 字节都没读到,说明到文件结尾了
if bytes_read == 0 {
return Ok(None);
}
let mut left = vec![0i16; chunk_size]; // 预填 0
let mut right = vec![0i16; chunk_size]; // 预填 0
// 实际读到了多少个采样点(总数)
let total_samples_read = bytes_read / 2;
// 计算每个声道实际读到了多少个点
let samples_per_channel = total_samples_read / channels;
for i in 0..samples_per_channel {
let base = i * channels * 2;
// 左声道 (或单声道)
left[i] = i16::from_le_bytes([bytes[base], bytes[base + 1]]);
if channels == 2 {
// 右声道
right[i] = i16::from_le_bytes([bytes[base + 2], bytes[base + 3]]);
} else {
// 单声道填充双声道时,复制左声道
right[i] = left[i];
pub fn read_one_frame(&mut self) -> Result<Option<&[i16]>> {
#[cfg(not(target_os = "linux"))]
{
match self.reader.read_chunk(self.config.frame_size)? {
Some((left_pcm, right_pcm)) => {
let actual_len = left_pcm.len();
if self.config.channels == 2 {
for i in 0..actual_len {
self.pcm_buffer[i * 2] = left_pcm[i];
self.pcm_buffer[i * 2 + 1] = right_pcm[i];
}
} else {
self.pcm_buffer[..actual_len].copy_from_slice(&left_pcm);
}
let total_len = actual_len * (self.config.channels as usize);
Ok(Some(&self.pcm_buffer[..total_len]))
}
None => Ok(None),
}
}
Ok(Some((left, right)))
#[cfg(target_os = "linux")]
{
Err(anyhow::anyhow!(
"WavReader is only supported on non-Linux platforms"
))
}
}
}