From 24194182fc9ca43e4701b0bdea428a420a82f78c Mon Sep 17 00:00:00 2001 From: jhqxxx <18280426169@163.com> Date: Thu, 8 Jan 2026 19:22:21 +0800 Subject: [PATCH] update get audio data --- Cargo.lock | 4 +- src/utils/audio_utils.rs | 101 ++++++++++++++++++++++++++++++++++++--- tests/messy_test.rs | 9 ++++ 3 files changed, 105 insertions(+), 9 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 9e14dc2..943e39b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -431,9 +431,9 @@ checksum = "8f1fe948ff07f4bd06c30984e69f5b4899c516a3ef74f34df92a2df2ab535495" [[package]] name = "bytes" -version = "1.10.1" +version = "1.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d71b6127be86fdcfddb610f7182ac57211d4b18a3e9c82eb2d17662f2227ad6a" +checksum = "b35204fbdc0b3f4446b89fc1ac2cf84a8a68971995d0bf2e925ec7cd960f9cb3" dependencies = [ "serde", ] diff --git a/src/utils/audio_utils.rs b/src/utils/audio_utils.rs index 6a43275..ed7fb46 100644 --- a/src/utils/audio_utils.rs +++ b/src/utils/audio_utils.rs @@ -269,6 +269,21 @@ pub fn load_audio_from_url(url: &str) -> Result { }) } +pub fn load_audio_bytes_from_url(url: &str) -> Result> { + tokio::task::block_in_place(|| { + let client = reqwest::blocking::Client::new(); + let response = client.get(url).send()?; + if !response.status().is_success() { + return Err(anyhow::anyhow!( + "Failed to download file: {}", + response.status() + )); + } + let bytes = response.bytes()?.to_vec(); + Ok(bytes) + }) +} + pub fn get_audio_path(path_str: &str) -> Result { if path_str.starts_with("http://") || path_str.starts_with("https://") { // Download file from network @@ -308,6 +323,34 @@ pub fn get_audio_path(path_str: &str) -> Result { } } +pub fn get_audio_bytes_vec(path_str: &str) -> Result> { + if path_str.starts_with("http://") || path_str.starts_with("https://") { + // Download file from network + load_audio_bytes_from_url(path_str) + } else if path_str.starts_with("file://") { + // Convert file:// URL to local path + let path = url::Url::parse(path_str)?; + let path = path.to_file_path(); + let path = match path { + Ok(path) => path, + Err(_) => { + let mut path = path_str.to_owned(); + path = path.split_off(7); + PathBuf::from(path) + } + }; + let bytes = std::fs::read(path)?; + Ok(bytes) + } else if path_str.starts_with("data:audio") && path_str.contains("base64,") { + let data: Vec<&str> = path_str.split("base64,").collect(); + let data = data[1]; + let data = BASE64_STANDARD.decode(data)?; + Ok(data) + } else { + Err(anyhow::anyhow!("get audio path error {}", path_str)) + } +} + pub fn load_audio_mono_vec(path: &str) -> Result<(Vec, usize)> { let audio_path = get_audio_path(path)?; let mut reader = WavReader::open(audio_path)?; @@ -410,13 +453,55 @@ pub fn load_audio_use_hound(audio_path: PathBuf, device: &Device) -> Result<(Ten Ok((audio_tensor, sample_rate as usize)) } -pub fn load_audio_use_symphonia(path: PathBuf, device: &Device) -> Result<(Tensor, usize)> { - let file = File::open(path.clone())?; - let mss = MediaSourceStream::new(Box::new(file), Default::default()); +pub fn get_audio_format_from_bytes(bytes: &[u8]) -> Result { + if bytes.len() < 12 { + return Err(anyhow::anyhow!("bytes too short: {}", bytes.len())); + } + // Check for different audio formats based on their magic bytes + if bytes.starts_with(&[0x52, 0x49, 0x46, 0x46]) && bytes.len() >= 12 { + // RIFF header - typically WAV files + if bytes.len() >= 8 && bytes[8..12] == [0x57, 0x41, 0x56, 0x45] { + Ok("wav".to_string()) + } else { + Ok("riff".to_string()) + } + } else if bytes.starts_with(&[0xFF, 0xFB]) + || bytes.starts_with(&[0xFF, 0xF3]) + || bytes.starts_with(&[0xFF, 0xF2]) + { + // MP3 header with different bitrates and options + Ok("mp3".to_string()) + } else if bytes.len() >= 3 && bytes[0..3] == [0x49, 0x44, 0x33] { + // ID3 tag - typically MP3 files + Ok("mp3".to_string()) + } else if bytes.len() >= 4 && bytes[0..4] == [0x46, 0x4F, 0x52, 0x4D] { + // FORM header - AIFF files + Ok("aiff".to_string()) + } else if bytes.len() >= 8 && bytes[0..4] == [0x4F, 0x67, 0x67, 0x53] { + // OggS header - OGG files + Ok("ogg".to_string()) + } else if bytes.len() >= 4 && bytes[0..4] == [0x66, 0x4C, 0x61, 0x43] { + // fLaC header - FLAC files + Ok("flac".to_string()) + } else if bytes.len() >= 8 && bytes[4..8] == [0x6D, 0x70, 0x34, 0x20] { + // M4A header + Ok("m4a".to_string()) + } else if bytes.len() >= 8 && bytes[4..8] == [0x6D, 0x70, 0x34, 0x61] { + // MP4A header + Ok("mp4".to_string()) + } else { + Err(anyhow::anyhow!("Unknown format ")) + } +} + +pub fn load_audio_use_symphonia(audio_vec: Vec, device: &Device) -> Result<(Tensor, usize)> { + let extension = get_audio_format_from_bytes(&audio_vec)?; + let content = Cursor::new(audio_vec); + let mss = MediaSourceStream::new(Box::new(content), Default::default()); let mut hint = Hint::new(); - let extension = path.extension().and_then(|ext| ext.to_str()).unwrap_or(""); - hint.with_extension(extension); + + hint.with_extension(&extension); let probed = symphonia::default::get_probe().format( &hint, @@ -509,10 +594,12 @@ pub fn load_audio_with_resample( device: &Device, target_sample_rate: Option, ) -> Result { - let audio_path = get_audio_path(path)?; // hound 只支持wav文件 + // let audio_path = get_audio_path(path)?; // let (mut audio, sr) = load_audio_use_hound(audio_path, device)?; - let (mut audio, sr) = load_audio_use_symphonia(audio_path, device)?; + + let audio_vec = get_audio_bytes_vec(path)?; + let (mut audio, sr) = load_audio_use_symphonia(audio_vec, device)?; if let Some(target_sample_rate) = target_sample_rate && target_sample_rate != sr { diff --git a/tests/messy_test.rs b/tests/messy_test.rs index e0d083f..9a9e150 100644 --- a/tests/messy_test.rs +++ b/tests/messy_test.rs @@ -1,11 +1,20 @@ +// use std::io::Cursor; + use aha::utils::audio_utils::create_hann_window; use anyhow::Result; use candle_core::DType; +// use symphonia::core::io::MediaSourceStream; #[test] fn messy_test() -> Result<()> { // RUST_BACKTRACE=1 cargo test -F cuda,ffmpeg messy_test -r -- --nocapture let device = &candle_core::Device::Cpu; + // let url = "https://sis-sample-audio.obs.cn-north-1.myhuaweicloud.com/16k16bit.mp3"; + // let client = reqwest::blocking::Client::new(); + // let response = client.get(url).send()?; + // let vec_u8 = response.bytes()?.to_vec(); + // let mut content = Cursor::new(vec_u8); + // let mss = MediaSourceStream::new(Box::new(content), Default::default()); let window = create_hann_window(400, DType::F32, device)?; println!("window: {}", window); // let audio_path = "file:///home/jhq/Videos/voice_01.wav";