use std::time::Duration;
#[cfg(feature = "image")]
pub mod image;
#[cfg(feature = "tts")]
pub mod tts;
#[cfg(feature = "video")]
pub mod video;
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum MediaError {
#[error("参数无效: {0}")]
InvalidInput(String),
#[error("{0}")]
Failed(String),
}
impl MediaError {
pub fn message(&self) -> &str {
match self {
MediaError::InvalidInput(m) | MediaError::Failed(m) => m,
}
}
}
#[derive(Clone, Default)]
pub struct MediaHttp {
base: Option<std::sync::Arc<dyn Fn() -> reqwest::ClientBuilder + Send + Sync>>,
}
impl MediaHttp {
pub fn from_fn(f: impl Fn() -> reqwest::ClientBuilder + Send + Sync + 'static) -> Self {
Self {
base: Some(std::sync::Arc::new(f)),
}
}
pub(crate) fn builder(&self) -> reqwest::ClientBuilder {
match &self.base {
Some(f) => f(),
None => reqwest::Client::builder(),
}
}
}
impl std::fmt::Debug for MediaHttp {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MediaHttp")
.field("custom_base", &self.base.is_some())
.finish()
}
}
pub(crate) fn truncate(s: &str, max: usize) -> String {
if s.chars().count() <= max {
s.to_string()
} else {
let head: String = s.chars().take(max).collect();
format!("{head}…")
}
}
pub(crate) mod http {
use super::Duration;
const CONNECT_TIMEOUT_SECS: u64 = 15;
#[cfg(any(feature = "video", feature = "tts"))]
const DEFAULT_TIMEOUT_SECS: u64 = 240;
#[cfg(feature = "image")]
const IMAGE_READ_TIMEOUT_SECS: u64 = 360;
#[cfg(any(feature = "video", feature = "tts"))]
pub(crate) fn default_client(base: &super::MediaHttp) -> MediaClient {
MediaClient::build(
base.builder()
.connect_timeout(Duration::from_secs(CONNECT_TIMEOUT_SECS))
.timeout(Duration::from_secs(DEFAULT_TIMEOUT_SECS)),
)
}
#[cfg(feature = "image")]
pub(crate) fn image_client(base: &super::MediaHttp) -> MediaClient {
MediaClient::build(
base.builder()
.connect_timeout(Duration::from_secs(CONNECT_TIMEOUT_SECS))
.read_timeout(Duration::from_secs(IMAGE_READ_TIMEOUT_SECS)),
)
}
#[derive(Clone)]
pub(crate) struct MediaClient(pub(super) Result<reqwest::Client, String>);
impl MediaClient {
fn build(b: reqwest::ClientBuilder) -> Self {
Self(b.build().map_err(|e| describe_reqwest_error(&e)))
}
pub(crate) fn get(&self) -> Result<&reqwest::Client, super::MediaError> {
self.0.as_ref().map_err(|e| {
super::MediaError::Failed(format!("HTTP 客户端初始化失败(请检查代理设置): {e}"))
})
}
}
fn reqwest_cause_chain(e: &reqwest::Error) -> String {
use std::error::Error as _;
let mut causes: Vec<String> = Vec::new();
let mut cur = e.source();
let mut depth = 0;
while let Some(src) = cur {
if depth >= 8 {
break;
}
let s = src.to_string();
if causes.last() != Some(&s) {
causes.push(s);
}
cur = src.source();
depth += 1;
}
causes.join(" → ")
}
pub(crate) fn describe_reqwest_error(e: &reqwest::Error) -> String {
let kind = if e.is_timeout() {
"超时"
} else if e.is_connect() {
"连接失败(端点不可达/DNS/TLS)"
} else if e.is_request() {
"请求发送失败"
} else if e.is_body() {
"请求/响应体错误"
} else if e.is_decode() {
"响应解码失败"
} else if e.is_redirect() {
"重定向过多"
} else {
"网络错误"
};
let chain = reqwest_cause_chain(e);
if chain.is_empty() {
format!("{kind}: {e}")
} else {
format!("{kind}: {e} —— 底层原因: {chain}")
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn display_matches_storyloom_app_error() {
assert_eq!(
MediaError::InvalidInput("配音文本为空".into()).to_string(),
"参数无效: 配音文本为空"
);
assert_eq!(
MediaError::Failed("上游 500".into()).to_string(),
"上游 500"
);
assert_eq!(MediaError::InvalidInput("x".into()).message(), "x");
}
#[test]
fn broken_client_reports_instead_of_falling_back() {
let broken = http::MediaClient(Err("代理地址无效".into()));
let err = broken.get().unwrap_err().to_string();
assert!(
err.contains("HTTP 客户端初始化失败") && err.contains("代理地址无效"),
"{err}"
);
}
#[test]
#[cfg(any(feature = "image", feature = "video"))]
fn media_http_factory_is_used_by_every_entry() {
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
let calls = Arc::new(AtomicUsize::new(0));
let c = calls.clone();
let http = MediaHttp::from_fn(move || {
c.fetch_add(1, Ordering::SeqCst);
reqwest::Client::builder()
});
let mut expected = 0;
#[cfg(feature = "image")]
{
let _ = image::AnyImageProvider::from_config_with(
image::ImageGenConfig {
endpoint: "https://a/v1".into(),
model: "m".into(),
api_key: "k".into(),
},
&http,
);
expected += 1;
}
#[cfg(feature = "video")]
{
let _ = video::AnyVideoProvider::from_config_with(
video::VideoGenConfig {
endpoint: "https://a/v1".into(),
model: "m".into(),
api_key: "k".into(),
},
"",
&http,
);
expected += 1;
}
assert_eq!(calls.load(Ordering::SeqCst), expected);
assert!(format!("{http:?}").contains("custom_base: true"));
assert!(format!("{:?}", MediaHttp::default()).contains("custom_base: false"));
}
#[test]
fn truncate_is_char_safe() {
assert_eq!(truncate("你好世界", 2), "你好…");
assert_eq!(truncate("abc", 3), "abc");
}
}