增强生成音频下载重试

为生成音频下载添加 identity 请求头和响应体读取失败重试。

补充音频下载响应体失败后重试成功的单元测试。

(cherry picked from commit 54b18bf0e2)
This commit is contained in:
2026-06-30 13:46:08 +08:00
parent 53c5639620
commit c556d42d22
+141 -23
View File
@@ -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<DownloadedAudio, AudioError> {
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<DownloadedAudio, AudioError> {
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");
}
}