diff --git a/server-rs/crates/platform-audio/src/download.rs b/server-rs/crates/platform-audio/src/download.rs index 4aceda4a7..0c252f756 100644 --- a/server-rs/crates/platform-audio/src/download.rs +++ b/server-rs/crates/platform-audio/src/download.rs @@ -1,9 +1,12 @@ -use std::error::Error; +use std::{error::Error, time::Duration}; use reqwest::header; use crate::{AudioError, DownloadedAudio, MAX_GENERATED_AUDIO_BYTES}; +const GENERATED_AUDIO_DOWNLOAD_MAX_ATTEMPTS: u32 = 2; +const GENERATED_AUDIO_DOWNLOAD_RETRY_DELAY_MS: u64 = 300; + pub fn normalize_audio_mime_type(content_type: &str, audio_url: &str) -> String { let mime_type = content_type .split(';') @@ -38,21 +41,70 @@ pub fn audio_mime_to_extension(mime_type: &str) -> &'static str { pub async fn download_generated_audio( http_client: &reqwest::Client, audio_url: &str, - _provider: &str, + provider: &str, ) -> Result { - let response = http_client.get(audio_url).send().await.map_err(|error| { - AudioError::request( - format!("下载生成音频失败:{error}"), - Some(audio_url.to_string()), - error.is_timeout(), - error.is_connect(), - error.is_request(), - error.is_body(), - error.status().map(|status| status.as_u16()), - Error::source(&error).map(ToString::to_string), - ) - })?; + let mut latest_error = None; + for attempt in 1..=GENERATED_AUDIO_DOWNLOAD_MAX_ATTEMPTS { + match download_generated_audio_once(http_client, audio_url).await { + Ok(audio) => return Ok(audio), + Err(error) if should_retry_generated_audio_download(&error, attempt) => { + tracing::warn!( + provider, + endpoint = audio_url, + attempt, + max_attempts = GENERATED_AUDIO_DOWNLOAD_MAX_ATTEMPTS, + error = %error, + "读取生成音频内容失败,准备重试" + ); + latest_error = Some(error); + tokio::time::sleep(Duration::from_millis( + GENERATED_AUDIO_DOWNLOAD_RETRY_DELAY_MS, + )) + .await; + } + Err(error) => return Err(error), + } + } + + Err(latest_error.unwrap_or_else(|| AudioError::missing_audio("生成音频下载失败"))) +} + +async fn download_generated_audio_once( + http_client: &reqwest::Client, + audio_url: &str, +) -> Result { + let response = http_client + .get(audio_url) + .header(header::ACCEPT, "audio/*,*/*") + .header(header::ACCEPT_ENCODING, "identity") + .send() + .await + .map_err(|error| { + AudioError::request( + format!("下载生成音频失败:{error}"), + Some(audio_url.to_string()), + error.is_timeout(), + error.is_connect(), + error.is_request(), + error.is_body(), + error.status().map(|status| status.as_u16()), + Error::source(&error).map(ToString::to_string), + ) + })?; let status = response.status(); + if !status.is_success() { + let raw_excerpt = response + .text() + .await + .map(|body| truncate_raw(body.as_str())) + .unwrap_or_default(); + return Err(AudioError::upstream( + format!("下载生成音频失败:HTTP {}", status.as_u16()), + status.as_u16(), + raw_excerpt, + )); + } + let content_type = response .headers() .get(header::CONTENT_TYPE) @@ -67,17 +119,10 @@ pub async fn download_generated_audio( false, false, true, - None, - None, + Some(status.as_u16()), + Error::source(&error).map(ToString::to_string), ) })?; - if !status.is_success() { - return Err(AudioError::upstream( - format!("下载生成音频失败:HTTP {}", status.as_u16()), - status.as_u16(), - truncate_raw(""), - )); - } if body.is_empty() || body.len() > MAX_GENERATED_AUDIO_BYTES { return Err(AudioError::missing_audio("生成音频内容为空或超过大小上限")); } @@ -90,6 +135,11 @@ pub async fn download_generated_audio( }) } +fn should_retry_generated_audio_download(error: &AudioError, attempt: u32) -> bool { + attempt < GENERATED_AUDIO_DOWNLOAD_MAX_ATTEMPTS + && matches!(error, AudioError::Request { body: true, .. }) +} + fn mime_type_from_audio_url(audio_url: &str) -> String { let path = audio_url .split('?') @@ -116,3 +166,71 @@ fn mime_type_from_audio_url(audio_url: &str) -> String { fn truncate_raw(raw_text: &str) -> String { raw_text.chars().take(800).collect() } + +#[cfg(test)] +mod tests { + use std::{ + io::{Read, Write}, + net::TcpListener, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + thread, + }; + + use super::*; + + #[test] + fn download_generated_audio_retries_body_decode_failure_once() { + let listener = TcpListener::bind("127.0.0.1:0").expect("mock server should bind"); + let address = listener + .local_addr() + .expect("mock server address should be readable"); + let request_count = Arc::new(AtomicUsize::new(0)); + let request_count_for_server = Arc::clone(&request_count); + let server = thread::spawn(move || { + for _ in 0..2 { + let (mut stream, _) = listener.accept().expect("request should connect"); + let mut request = [0_u8; 1024]; + let _ = stream.read(&mut request); + let index = request_count_for_server.fetch_add(1, Ordering::SeqCst); + if index == 0 { + let _ = stream.write_all( + b"HTTP/1.1 200 OK\r\nContent-Type: audio/mpeg\r\nTransfer-Encoding: chunked\r\n\r\n10\r\nabc\r\n", + ); + continue; + } + + let body = b"RIFF"; + let response = format!( + "HTTP/1.1 200 OK\r\nContent-Type: audio/wav\r\nContent-Length: {}\r\n\r\n", + body.len() + ); + let _ = stream.write_all(response.as_bytes()); + let _ = stream.write_all(body); + } + }); + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .expect("tokio runtime should build"); + let http_client = reqwest::Client::builder() + .timeout(Duration::from_secs(2)) + .build() + .expect("http client should build"); + + let audio = runtime + .block_on(download_generated_audio( + &http_client, + format!("http://{address}/audio.wav").as_str(), + "test-provider", + )) + .expect("second attempt should download audio"); + + assert_eq!(audio.bytes, b"RIFF"); + assert_eq!(audio.mime_type, "audio/wav"); + assert_eq!(request_count.load(Ordering::SeqCst), 2); + server.join().expect("mock server should finish"); + } +}