use serde::{Deserialize, Serialize};
use super::truncate;
use super::MediaError;
#[derive(Debug, Clone)]
pub struct ImageGenConfig {
pub endpoint: String,
pub model: String,
pub api_key: String,
}
#[derive(Debug, Clone)]
pub struct ImageGenParams {
pub prompt: String,
pub negative_prompt: String,
pub seed: Option<i64>,
pub size: String,
pub images: Vec<String>,
}
impl Default for ImageGenParams {
fn default() -> Self {
Self {
prompt: String::new(),
negative_prompt: String::new(),
seed: None,
size: "720x1280".to_string(), images: Vec::new(),
}
}
}
#[derive(Debug, Clone)]
pub struct ImageResult {
pub bytes: Vec<u8>,
pub ext: String,
pub seed: Option<i64>,
}
#[allow(async_fn_in_trait)]
pub trait ImageProvider {
async fn generate(&self, params: &ImageGenParams) -> Result<ImageResult, MediaError>;
}
pub struct OpenAiImageProvider {
client: super::http::MediaClient,
config: ImageGenConfig,
}
impl OpenAiImageProvider {
pub fn new(config: ImageGenConfig) -> Self {
Self::with_http(config, &super::MediaHttp::default())
}
pub fn with_http(config: ImageGenConfig, http: &super::MediaHttp) -> Self {
Self {
client: super::http::image_client(http),
config,
}
}
fn images_url(&self) -> String {
let e = self.config.endpoint.trim_end_matches('/');
if e.ends_with("/images/generations") {
e.to_string()
} else {
format!("{e}/images/generations")
}
}
}
#[derive(Serialize)]
struct ImageRequest<'a> {
model: &'a str,
prompt: &'a str,
#[serde(skip_serializing_if = "str::is_empty")]
negative_prompt: &'a str,
image_size: &'a str,
batch_size: u8,
#[serde(skip_serializing_if = "Option::is_none")]
seed: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
image: Option<serde_json::Value>,
}
#[derive(Deserialize)]
struct ImageResponse {
#[serde(default)]
images: Vec<ImageItem>,
#[serde(default)]
seed: Option<i64>,
#[serde(default)]
data: Vec<ImageItem>,
}
#[derive(Deserialize, Default)]
struct ImageItem {
#[serde(default)]
url: Option<String>,
#[serde(default)]
b64_json: Option<String>,
}
impl ImageProvider for OpenAiImageProvider {
async fn generate(&self, params: &ImageGenParams) -> Result<ImageResult, MediaError> {
let image = match params.images.as_slice() {
[] => None,
[one] => Some(serde_json::Value::String(one.clone())),
many => Some(serde_json::Value::Array(
many.iter()
.map(|s| serde_json::Value::String(s.clone()))
.collect(),
)),
};
let body = ImageRequest {
model: &self.config.model,
prompt: ¶ms.prompt,
negative_prompt: ¶ms.negative_prompt,
image_size: ¶ms.size,
batch_size: 1,
seed: params.seed,
image,
};
let resp = self
.client
.get()?
.post(self.images_url())
.bearer_auth(&self.config.api_key)
.json(&body)
.send()
.await
.map_err(|e| {
MediaError::Failed(format!(
"出图请求失败: {}",
super::http::describe_reqwest_error(&e)
))
})?;
let status = resp.status();
if !status.is_success() {
let text = resp.text().await.unwrap_or_default();
let model_not_found = status.as_u16() == 404
|| text.contains("InvalidEndpointOrModel")
|| text.contains("does not exist")
|| (text.contains("model") && text.contains("Not Found"));
if model_not_found {
return Err(MediaError::Failed(format!(
"出图失败:模型「{}」不存在或未开通。请到 设置 → 生图 把模型改成你在服务商控制台已开通的模型 ID(火山方舟 Seedream 带日期后缀,需与控制台「开通模型」一致)。原始响应 HTTP {}: {}",
self.config.model,
status.as_u16(),
truncate(&text, 200)
)));
}
return Err(MediaError::Failed(explain_image_http_error(
status.as_u16(),
&text,
)));
}
let parsed: ImageResponse = resp
.json()
.await
.map_err(|e| MediaError::Failed(format!("解析出图响应失败: {e}")))?;
let item = parsed
.images
.into_iter()
.next()
.or_else(|| parsed.data.into_iter().next())
.ok_or_else(|| {
MediaError::Failed(
"出图失败:供应商返回成功但未给出图片(通常是内容安全审核拦截,或额度异常)。\
建议:① 调整提示词 / 参考图,避开露肤、暴力、政治等敏感元素重试;\
② 或改用其他生图供应商(如火山方舟)"
.into(),
)
})?;
if let Some(b64) = item.b64_json.filter(|s| !s.is_empty()) {
let bytes = base64_decode(&b64)?;
validate_image_bytes(&bytes, "供应商内联返回(b64_json)")?;
return Ok(ImageResult {
bytes,
ext: "png".into(),
seed: parsed.seed,
});
}
let url = item
.url
.filter(|s| !s.is_empty())
.ok_or_else(|| MediaError::Failed("出图响应缺少图片 URL".into()))?;
let (bytes, ext) = download_image_checked(self.client.get()?, &url, "下载出图结果").await?;
Ok(ImageResult {
bytes,
ext,
seed: parsed.seed,
})
}
}
pub struct DashScopeImageProvider {
client: super::http::MediaClient,
config: ImageGenConfig,
}
impl DashScopeImageProvider {
pub fn new(config: ImageGenConfig) -> Self {
Self::with_http(config, &super::MediaHttp::default())
}
pub fn with_http(config: ImageGenConfig, http: &super::MediaHttp) -> Self {
Self {
client: super::http::image_client(http),
config,
}
}
fn base(&self) -> String {
self.config.endpoint.trim_end_matches('/').to_string()
}
async fn submit(&self, params: &ImageGenParams) -> Result<String, MediaError> {
let size = params.size.replace(['x', 'X'], "*");
let mut input = serde_json::json!({ "prompt": params.prompt });
if !params.negative_prompt.is_empty() {
input["negative_prompt"] = serde_json::Value::String(params.negative_prompt.clone());
}
let mut parameters = serde_json::json!({ "size": size, "n": 1 });
if let Some(seed) = params.seed {
parameters["seed"] = serde_json::json!(seed);
}
let body = serde_json::json!({
"model": self.config.model,
"input": input,
"parameters": parameters,
});
let resp = self
.client
.get()?
.post(format!(
"{}/services/aigc/text2image/image-synthesis",
self.base()
))
.bearer_auth(&self.config.api_key)
.header("X-DashScope-Async", "enable")
.json(&body)
.send()
.await
.map_err(|e| {
MediaError::Failed(format!(
"提交文生图任务失败: {}",
super::http::describe_reqwest_error(&e)
))
})?;
let status = resp.status();
if !status.is_success() {
let text = resp.text().await.unwrap_or_default();
return Err(MediaError::Failed(format!(
"文生图提交 HTTP {}: {}",
status.as_u16(),
truncate(&text, 300)
)));
}
let parsed: DashImageResponse = resp
.json()
.await
.map_err(|e| MediaError::Failed(format!("解析文生图任务创建响应失败: {e}")))?;
parsed
.output
.and_then(|o| o.task_id)
.filter(|s| !s.is_empty())
.ok_or_else(|| {
MediaError::Failed(format!(
"文生图任务创建未返回 task_id{}",
parsed.message.map(|m| format!(": {m}")).unwrap_or_default()
))
})
}
async fn poll_until_url(&self, task_id: &str) -> Result<String, MediaError> {
for _ in 0..100 {
let resp = self
.client
.get()?
.get(format!("{}/tasks/{}", self.base(), task_id))
.bearer_auth(&self.config.api_key)
.send()
.await
.map_err(|e| {
MediaError::Failed(format!(
"查询文生图任务失败: {}",
super::http::describe_reqwest_error(&e)
))
})?;
let status = resp.status();
if !status.is_success() {
let text = resp.text().await.unwrap_or_default();
return Err(MediaError::Failed(format!(
"文生图查询 HTTP {}: {}",
status.as_u16(),
truncate(&text, 300)
)));
}
let parsed: DashImageResponse = resp
.json()
.await
.map_err(|e| MediaError::Failed(format!("解析文生图任务状态失败: {e}")))?;
let out = parsed.output.unwrap_or_default();
match out.task_status.as_deref() {
Some("SUCCEEDED") => {
return out
.results
.into_iter()
.find_map(|r| r.url.filter(|s| !s.is_empty()))
.ok_or_else(|| MediaError::Failed("文生图成功但未返回图片 URL".into()));
}
Some("FAILED") | Some("CANCELED") | Some("UNKNOWN") => {
return Err(MediaError::Failed(format!(
"文生图任务失败: {}",
out.message.unwrap_or_else(|| "未知原因".into())
)));
}
_ => tokio::time::sleep(std::time::Duration::from_secs(3)).await,
}
}
Err(MediaError::Failed("文生图任务超时(>300s 未完成)".into()))
}
}
#[derive(Deserialize, Default)]
struct DashImageOutput {
#[serde(default)]
task_id: Option<String>,
#[serde(default)]
task_status: Option<String>,
#[serde(default)]
results: Vec<DashImageItem>,
#[serde(default)]
message: Option<String>,
}
#[derive(Deserialize, Default)]
struct DashImageItem {
#[serde(default)]
url: Option<String>,
}
#[derive(Deserialize)]
struct DashImageResponse {
#[serde(default)]
output: Option<DashImageOutput>,
#[serde(default)]
message: Option<String>,
}
impl ImageProvider for DashScopeImageProvider {
async fn generate(&self, params: &ImageGenParams) -> Result<ImageResult, MediaError> {
let task_id = self.submit(params).await?;
let url = self.poll_until_url(&task_id).await?;
let (bytes, ext) =
download_image_checked(self.client.get()?, &url, "下载文生图结果").await?;
Ok(ImageResult {
bytes,
ext,
seed: params.seed,
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ImageProtocol {
OpenAi,
DashScope,
}
impl ImageProtocol {
pub fn detect(endpoint: &str) -> Self {
if endpoint.to_ascii_lowercase().contains("dashscope") {
Self::DashScope
} else {
Self::OpenAi
}
}
}
pub enum AnyImageProvider {
OpenAi(OpenAiImageProvider),
DashScope(DashScopeImageProvider),
}
impl AnyImageProvider {
pub fn from_config(config: ImageGenConfig) -> Self {
Self::from_config_with(config, &super::MediaHttp::default())
}
pub fn from_config_with(config: ImageGenConfig, http: &super::MediaHttp) -> Self {
match ImageProtocol::detect(&config.endpoint) {
ImageProtocol::DashScope => {
Self::DashScope(DashScopeImageProvider::with_http(config, http))
}
ImageProtocol::OpenAi => Self::OpenAi(OpenAiImageProvider::with_http(config, http)),
}
}
pub async fn generate(&self, params: &ImageGenParams) -> Result<ImageResult, MediaError> {
match self {
Self::OpenAi(p) => p.generate(params).await,
Self::DashScope(p) => p.generate(params).await,
}
}
}
fn explain_image_http_error(status: u16, body: &str) -> String {
let low = body.to_ascii_lowercase();
let is_balance = low.contains("insufficient") || low.contains("quota") || body.contains("余额");
let is_moderation = low.contains("sensitive")
|| low.contains("moderation")
|| low.contains("risk")
|| low.contains("content_policy")
|| body.contains("审核")
|| body.contains("违规");
let hint = match status {
400 if is_balance => "供应商余额不足,请充值后重试",
400 if is_moderation || low.contains("fail_to_fetch_task") || low.contains("upstream returned status") => {
"该出图请求被供应商拒绝(通常是内容安全审核:提示词/参考图含露肤、暴力、政治等敏感元素)。\
建议:① 调整提示词与参考图重试;② 或改用其他生图供应商(如火山方舟)"
}
400 => "出图被供应商拒绝(参数或内容不被接受)。建议调整提示词/尺寸或换供应商重试",
401 | 403 if is_balance => "供应商余额不足或无权调用该模型,请充值 / 确认该 Key 已开通此生图模型",
403 => "无权调用该生图模型(Key 未开通或越权),请在供应商后台确认权限",
413 => "请求体过大(参考图过大),请压缩参考图后重试",
429 => "供应商限流,请稍后重试",
_ => "",
};
if hint.is_empty() {
format!("出图 HTTP {status}: {}", truncate(body, 300))
} else {
format!(
"出图失败(HTTP {status}):{hint}。原始响应:{}",
truncate(body, 200)
)
}
}
pub const MIN_IMAGE_BYTES: usize = 1024;
fn validate_image_bytes(bytes: &[u8], source: &str) -> Result<(), MediaError> {
if bytes.is_empty() {
return Err(MediaError::Failed(format!(
"出图失败:{source}返回了空内容(0 字节)。通常是供应商内容审核拦截、额度异常或中转站回源失败。\
建议:① 调整提示词/参考图重试;② 或改用其他生图供应商"
)));
}
if bytes.len() < MIN_IMAGE_BYTES {
return Err(MediaError::Failed(format!(
"出图失败:{source}返回的内容只有 {} 字节,不是有效图片(通常是错误页或占位内容)。请重试或更换生图供应商",
bytes.len()
)));
}
let is_png = bytes.starts_with(&[0x89, 0x50, 0x4E, 0x47]);
let is_jpeg = bytes.starts_with(&[0xFF, 0xD8, 0xFF]);
let is_webp = bytes.starts_with(b"RIFF") && bytes.len() > 12 && &bytes[8..12] == b"WEBP";
let is_gif = bytes.starts_with(b"GIF8");
let is_bmp = bytes.starts_with(b"BM");
if !(is_png || is_jpeg || is_webp || is_gif || is_bmp) {
let head = String::from_utf8_lossy(&bytes[..bytes.len().min(120)]).to_string();
return Err(MediaError::Failed(format!(
"出图失败:{source}返回的内容不是有效图片格式(非 PNG/JPEG/WEBP)。内容开头:{}",
truncate(&head, 120)
)));
}
Ok(())
}
async fn download_image_checked(
client: &reqwest::Client,
url: &str,
label: &str,
) -> Result<(Vec<u8>, String), MediaError> {
let mut last_err: Option<MediaError> = None;
for attempt in 0..2 {
if attempt > 0 {
tokio::time::sleep(std::time::Duration::from_secs(1)).await;
log::warn!("{label}下载失败,重试第 {attempt} 次(图已生成,重试不重复计费)");
}
match download_image_once(client, url, label).await {
Ok(v) => return Ok(v),
Err(e) => last_err = Some(e),
}
}
Err(last_err.unwrap_or_else(|| MediaError::Failed(format!("{label}下载失败"))))
}
async fn download_image_once(
client: &reqwest::Client,
url: &str,
label: &str,
) -> Result<(Vec<u8>, String), MediaError> {
let resp = client.get(url).send().await.map_err(|e| {
MediaError::Failed(format!(
"{label}失败: {}",
super::http::describe_reqwest_error(&e)
))
})?;
if !resp.status().is_success() {
return Err(MediaError::Failed(format!(
"{label} HTTP {}",
resp.status().as_u16()
)));
}
let ext = ext_from_url_or_ct(url, resp.headers());
let bytes = resp
.bytes()
.await
.map_err(|e| MediaError::Failed(format!("读取出图字节失败: {e}")))?
.to_vec();
validate_image_bytes(&bytes, "结果图下载")?;
Ok((bytes, ext))
}
fn ext_from_url_or_ct(url: &str, headers: &reqwest::header::HeaderMap) -> String {
let lower = url.split('?').next().unwrap_or(url).to_ascii_lowercase();
for e in ["png", "jpeg", "jpg", "webp"] {
if lower.ends_with(&format!(".{e}")) {
return if e == "jpeg" { "jpg".into() } else { e.into() };
}
}
if let Some(ct) = headers
.get(reqwest::header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
{
if ct.contains("png") {
return "png".into();
} else if ct.contains("webp") {
return "webp".into();
} else if ct.contains("jpeg") || ct.contains("jpg") {
return "jpg".into();
}
}
"png".into()
}
fn base64_decode(s: &str) -> Result<Vec<u8>, MediaError> {
use base64::Engine;
base64::engine::general_purpose::STANDARD
.decode(s.trim())
.map_err(|e| MediaError::Failed(format!("图片 base64 解码失败: {e}")))
}
#[cfg(test)]
mod tests {
use super::*;
fn padded(magic: &[u8]) -> Vec<u8> {
let mut v = magic.to_vec();
v.resize(MIN_IMAGE_BYTES + 64, 0x00);
v
}
#[test]
fn rejects_empty_and_undersized_payloads() {
assert!(
validate_image_bytes(&[], "测试").is_err(),
"0 字节必须判失败"
);
assert!(
validate_image_bytes(&[0x89, 0x50, 0x4E, 0x47], "测试").is_err(),
"带正确魔数但长度不足 1KB 仍应判失败"
);
}
#[test]
fn rejects_non_image_payload() {
let mut json_err = br#"{"error":"upstream returned status 500"}"#.to_vec();
json_err.resize(MIN_IMAGE_BYTES + 64, b' ');
assert!(
validate_image_bytes(&json_err, "测试").is_err(),
"长度达标但非图片格式必须判失败"
);
}
#[test]
fn accepts_common_image_formats() {
assert!(
validate_image_bytes(&padded(&[0x89, 0x50, 0x4E, 0x47]), "测试").is_ok(),
"PNG"
);
assert!(
validate_image_bytes(&padded(&[0xFF, 0xD8, 0xFF]), "测试").is_ok(),
"JPEG"
);
assert!(
validate_image_bytes(&padded(b"GIF8"), "测试").is_ok(),
"GIF"
);
let mut webp = b"RIFF".to_vec();
webp.extend_from_slice(&[0, 0, 0, 0]);
webp.extend_from_slice(b"WEBP");
webp.resize(MIN_IMAGE_BYTES + 64, 0x00);
assert!(validate_image_bytes(&webp, "测试").is_ok(), "WEBP");
}
}
#[cfg(test)]
mod protocol_tests {
use super::*;
#[test]
fn every_image_preset_detects_expected_protocol() {
for p in crate::preset::presets_for(crate::Kind::Image) {
let Some(url) = p.base_url else { continue };
let want = if p.key == "wan_image" {
ImageProtocol::DashScope
} else {
ImageProtocol::OpenAi
};
assert_eq!(ImageProtocol::detect(url), want, "{}", p.key);
}
}
}