增强生成音频下载重试
为生成音频下载添加 identity 请求头和响应体读取失败重试。
补充音频下载响应体失败后重试成功的单元测试。
(cherry picked from commit 54b18bf0e2)
This commit is contained in:
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user