Files
video-gen/video-gen-api/app/services/image_gen.py
T
2026-07-20 14:01:22 +08:00

478 lines
18 KiB
Python

import json
import logging
import os
from datetime import datetime
from typing import Any
import httpx
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from volcenginesdkarkruntime import AsyncArk
from app.config import settings
from app.enums.generation_provider import (
MULTI_IMAGE_PROMPT_TEMPLATE,
ImageProviderErrorType,
)
from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX
from app.models.image_engine import ImageEngine
from app.services.log_config import LOG_DATE_FORMAT, LOG_DIR, encrypt_data, is_enabled
from app.types.generation.provider import (
ImageProviderBatchResult,
ImageProviderItem,
ProviderGenerationRecordLike,
ProviderImageEngineLike,
)
logger = logging.getLogger("videogen")
class ImageProviderError(RuntimeError):
"""可被生成任务状态机安全收敛的图片供应商异常。"""
def __init__(
self,
message: str,
*,
error_type: ImageProviderErrorType = ImageProviderErrorType.UNKNOWN,
error_code: str | None = None,
retryable: bool = False,
http_status: int | None = None,
provider_request_id: str | None = None,
):
super().__init__(message)
self.safe_message = message
self.error_type = error_type
self.error_code = error_code
self.retryable = retryable
self.http_status = http_status
self.provider_request_id = provider_request_id
def as_dict(self) -> dict[str, Any]:
return {
"error_type": self.error_type.value,
"error_code": self.error_code,
"message": self.safe_message,
"retryable": self.retryable,
"http_status": self.http_status,
"provider_request_id": self.provider_request_id,
}
def _log_image_request(engine: ProviderImageEngineLike, record_id: str, request_data: dict):
if not is_enabled():
return
try:
os.makedirs(LOG_DIR, exist_ok=True)
today = datetime.now().strftime(LOG_DATE_FORMAT)
log_file = os.path.join(LOG_DIR, f"{today}.log")
request_str = json.dumps(request_data, ensure_ascii=False)
request_encrypted = encrypt_data(request_data, True)
entry = {
"timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"type": "image_gen_request",
"engine": engine.name,
"model": engine.model_name,
"record_id": record_id,
"request": request_encrypted,
"request_length": len(request_str),
}
with open(log_file, "a", encoding="utf-8") as file:
file.write(json.dumps(entry, ensure_ascii=False) + "\n")
except Exception:
pass
def _log_image_response(record_id: str, response_data: dict, error: str | None = None):
if not is_enabled():
return
try:
os.makedirs(LOG_DIR, exist_ok=True)
today = datetime.now().strftime(LOG_DATE_FORMAT)
log_file = os.path.join(LOG_DIR, f"{today}.log")
response_encrypted = encrypt_data(response_data, True) if response_data else ""
entry = {
"timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"type": "image_gen_response",
"record_id": record_id,
"response": response_encrypted,
"error": error,
}
with open(log_file, "a", encoding="utf-8") as file:
file.write(json.dumps(entry, ensure_ascii=False) + "\n")
except Exception:
pass
async def get_active_image_engine(db: AsyncSession) -> ImageEngine:
result = await db.execute(
select(ImageEngine)
.where(ImageEngine.is_active == True, ImageEngine.deleted_at.is_(None))
.order_by(ImageEngine.priority.desc())
.limit(1)
)
engine = result.scalar_one_or_none()
if not engine:
raise ValueError("没有可用的图片引擎,请联系管理员配置")
return engine
def _resolve_url(url: str) -> str:
"""Convert local path to base64 data URI, pass through remote URLs."""
# if url.startswith("http"):
# return url
# file_path = os.path.join(settings.UPLOAD_LOCAL_PATH, url.replace("/uploads/", ""))
# if not os.path.exists(file_path):
# settings.BASE_URL + mime = mimetypes.guess_type(file_path)[0] or "application/octet-stream"
# with open(file_path, "rb") as f:
# b64 = base64.b64encode(f.read()).decode()
# return f"data:{mime};base64,{b64}"
if url.startswith(("http://", "https://", "data:", PRIVATE_PORTRAIT_ASSET_URI_PREFIX)):
return url
return f"{settings.BASE_URL.rstrip('/')}/{url.lstrip('/')}"
def _value(obj: Any, name: str, default: Any = None) -> Any:
if obj is None:
return default
if isinstance(obj, dict):
return obj.get(name, default)
return getattr(obj, name, default)
def _jsonable(value: Any) -> Any:
if value is None or isinstance(value, (str, int, float, bool)):
return value
if isinstance(value, dict):
return {str(key): _jsonable(item) for key, item in value.items()}
if isinstance(value, (list, tuple)):
return [_jsonable(item) for item in value]
if hasattr(value, "model_dump"):
try:
return _jsonable(value.model_dump())
except Exception:
pass
if hasattr(value, "to_dict"):
try:
return _jsonable(value.to_dict())
except Exception:
pass
result: dict[str, Any] = {}
for key in ("url", "b64_json", "size", "output_format", "error", "code", "message"):
item = getattr(value, key, None)
if item is not None:
result[key] = _jsonable(item)
return result or str(value)
def _safe_text(value: Any, *, limit: int = 1000) -> str:
text = str(value or "").strip()
return text[:limit]
def _classify_provider_exception(exc: Exception) -> ImageProviderError:
if isinstance(exc, ImageProviderError):
return exc
if isinstance(exc, (httpx.TimeoutException, TimeoutError)):
return ImageProviderError(
"图片生成请求超时,请稍后重试",
error_type=ImageProviderErrorType.TIMEOUT,
retryable=True,
)
status_code = getattr(exc, "status_code", None)
request_id = getattr(exc, "request_id", None) or getattr(exc, "x_request_id", None)
code = getattr(exc, "code", None)
raw_message = _safe_text(getattr(exc, "message", None) or exc)
lowered = raw_message.lower()
if status_code == 429 or "rate limit" in lowered or "限流" in raw_message:
error_type = ImageProviderErrorType.RATE_LIMIT
retryable = True
message = "图片生成请求过于频繁,请稍后重试"
elif status_code in {401, 403} or "api key" in lowered or "unauthorized" in lowered:
error_type = ImageProviderErrorType.AUTH
retryable = False
message = "图片引擎鉴权失败,请联系管理员检查配置"
elif status_code and int(status_code) >= 500:
error_type = ImageProviderErrorType.PROVIDER_INTERNAL
retryable = True
message = "图片供应商服务异常,请稍后重试"
elif "sequential_image_generation" in lowered or "not support" in lowered or "unsupported" in lowered:
error_type = ImageProviderErrorType.CAPABILITY_MISMATCH
retryable = False
message = "图片引擎组图能力配置与供应商实际能力不匹配,请联系管理员"
elif "content" in lowered and ("risk" in lowered or "moderation" in lowered or "policy" in lowered):
error_type = ImageProviderErrorType.CONTENT_REJECTED
retryable = False
message = "图片内容未通过供应商审核,请调整提示词后重试"
elif status_code and 400 <= int(status_code) < 500:
error_type = ImageProviderErrorType.INVALID_REQUEST
retryable = False
message = "图片生成参数不被供应商支持,请联系管理员检查引擎配置"
elif isinstance(exc, httpx.HTTPError):
error_type = ImageProviderErrorType.NETWORK
retryable = True
message = "图片供应商网络连接异常,请稍后重试"
else:
error_type = ImageProviderErrorType.UNKNOWN
retryable = False
message = raw_message or "图片生成失败"
return ImageProviderError(
message,
error_type=error_type,
error_code=_safe_text(code, limit=128) or None,
retryable=retryable,
http_status=int(status_code) if status_code is not None else None,
provider_request_id=_safe_text(request_id, limit=128) or None,
)
def build_multi_image_provider_prompt(prompt: str, generation_count: int) -> str:
base_prompt = (prompt or "").strip()
if generation_count <= 1:
return base_prompt
suffix = MULTI_IMAGE_PROMPT_TEMPLATE.format(count=generation_count)
return f"{base_prompt}\n\n{suffix}" if base_prompt else suffix
def submit_image_task(
db,
engine: ProviderImageEngineLike,
record: ProviderGenerationRecordLike,
*,
include_media_references: bool,
generation_count: int = 1,
) -> ImageProviderBatchResult:
"""通过 Ark 同步图片接口生成单图或单次组图。
generation_count > 1 时只执行一次 sequential_auto 请求;任何失败都直接抛出,
绝不退化为多次单图请求。
"""
from volcenginesdkarkruntime import Ark
count = max(1, int(generation_count or 1))
multi_generation_enabled = bool(getattr(engine, "multi_generation_enabled", False))
max_generation_count = max(1, min(5, int(getattr(engine, "max_generation_count", 1) or 1)))
if count > 1 and not multi_generation_enabled:
raise ImageProviderError(
"当前图片引擎未开启多份生成",
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
)
if count > max_generation_count:
raise ImageProviderError(
f"当前图片引擎最多允许生成 {max_generation_count} 份",
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
)
client = Ark(base_url=engine.api_base, api_key=engine.api_key, timeout=300)
original_prompt = record.optimized_prompt or record.original_prompt
provider_prompt = build_multi_image_provider_prompt(original_prompt, count)
image_urls: list[str] = []
if include_media_references and record.media_references:
try:
refs = json.loads(record.media_references)
for ref in refs if isinstance(refs, list) else []:
if (ref.get("type") or "").lower() == "image" and ref.get("url"):
image_urls.append(_resolve_url(ref["url"]))
except (json.JSONDecodeError, TypeError):
image_urls = []
request_log_payload: dict[str, Any] = {
"model": engine.model_name,
"prompt": provider_prompt,
"size": record.image_size or engine.default_size,
"response_format": "url",
"watermark": False,
}
request_sdk_payload: dict[str, Any] = dict(request_log_payload)
if image_urls:
request_log_payload["image"] = image_urls
request_sdk_payload["image"] = image_urls
output_format = (getattr(engine, "output_format", "") or "").lower().strip()
if output_format:
request_log_payload["output_format"] = output_format
request_sdk_payload["output_format"] = output_format
if count > 1:
try:
from volcenginesdkarkruntime.types.images import SequentialImageGenerationOptions
except Exception:
try:
from volcenginesdkarkruntime.types.images.image_generate_params import (
SequentialImageGenerationOptions,
)
except Exception as import_exc:
raise ImageProviderError(
"当前图片引擎运行依赖缺少组图参数对象,请升级火山 Ark SDK 后重试",
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
) from import_exc
request_log_payload["sequential_image_generation"] = "auto"
request_log_payload["sequential_image_generation_options"] = {"max_images": count}
request_log_payload["stream"] = False
request_sdk_payload["sequential_image_generation"] = "auto"
request_sdk_payload["sequential_image_generation_options"] = SequentialImageGenerationOptions(
max_images=count,
)
request_sdk_payload["stream"] = False
_log_image_request(engine, record.id, request_log_payload)
try:
result = client.images.generate(**request_sdk_payload)
top_error = _value(result, "error")
if top_error:
error_code = _value(top_error, "code")
error_message = _value(top_error, "message") or str(top_error)
raise ImageProviderError(
_safe_text(error_message) or "图片供应商返回失败",
error_type=ImageProviderErrorType.INVALID_REQUEST,
error_code=_safe_text(error_code, limit=128) or None,
)
raw_data = _value(result, "data", []) or []
if not isinstance(raw_data, (list, tuple)):
raise ImageProviderError(
"图片供应商返回 data 结构异常",
error_type=ImageProviderErrorType.INVALID_RESPONSE,
)
items: list[ImageProviderItem] = []
response_items: list[dict[str, Any]] = []
for index, raw_item in enumerate(raw_data, start=1):
item_error = _value(raw_item, "error")
if item_error:
error_code = _safe_text(_value(item_error, "code"), limit=128)
error_message = _safe_text(_value(item_error, "message") or item_error)
items.append({
"generation_index": index,
"error_code": error_code,
"error_message": error_message or "单张图片生成失败",
"response_data": _jsonable(raw_item),
})
response_items.append(_jsonable(raw_item))
continue
url = _safe_text(_value(raw_item, "url"), limit=4000)
b64_json = _safe_text(_value(raw_item, "b64_json"), limit=100) if not url else ""
item: ImageProviderItem = {
"generation_index": index,
"remote_result_url": url,
"size": _safe_text(_value(raw_item, "size"), limit=64),
"output_format": _safe_text(_value(raw_item, "output_format"), limit=32),
"response_data": _jsonable(raw_item),
}
if b64_json:
item["b64_json"] = b64_json
items.append(item)
response_items.append(_jsonable(raw_item))
usage = _value(result, "usage")
generated_images = int(_value(usage, "generated_images", 0) or 0)
total_tokens = int(_value(usage, "total_tokens", 0) or 0)
response_data = {
"model": _value(result, "model", engine.model_name),
"created": _value(result, "created"),
"data": response_items,
"usage": {
"generated_images": generated_images,
"input_images": int(_value(usage, "input_images", 0) or 0),
"output_tokens": int(_value(usage, "output_tokens", 0) or 0),
"total_tokens": total_tokens,
},
}
_log_image_response(record.id, response_data)
return {
"items": items,
"model": str(response_data["model"] or ""),
"created": int(response_data["created"] or 0),
"generated_images": generated_images,
"image_tokens": total_tokens,
"response_data": response_data,
}
except Exception as exc:
provider_error = _classify_provider_exception(exc)
logger.error(
"Image generation failed for record %s: type=%s code=%s message=%s",
record.id,
provider_error.error_type.value,
provider_error.error_code,
provider_error.safe_message,
)
_log_image_response(record.id, provider_error.as_dict(), provider_error.safe_message)
raise provider_error from exc
finally:
try:
client.close()
except Exception:
pass
async def poll_image_task_status(engine: ImageEngine, task_id: str) -> dict:
client = AsyncArk(base_url=engine.api_base, api_key=engine.api_key)
try:
result = await client.image_generation.tasks.get(task_id=task_id)
finally:
await client.close()
response_dict = {
"id": result.id,
"model": result.model,
"status": result.status,
"created_at": result.created_at,
"updated_at": result.updated_at,
}
image_url = None
image_tokens = 0
if result.status == "succeeded" and result.content:
image_url = getattr(result.content, "image_url", None)
response_dict["image_url"] = image_url
response_dict["ratio"] = getattr(result, "ratio", None)
response_dict["size"] = getattr(result, "size", None)
usage = getattr(result, "usage", None)
if usage:
response_dict["usage"] = {
"input_tokens": getattr(usage, "input_tokens", 0),
"output_tokens": getattr(usage, "output_tokens", 0),
"total_tokens": getattr(usage, "total_tokens", 0),
}
image_tokens = getattr(usage, "total_tokens", 0)
elif result.status == "failed":
response_dict["error"] = str(getattr(result, "error", "图片生成失败"))
return {
"status": result.status,
"image_url": image_url,
"image_tokens": image_tokens,
"response_data": json.dumps(response_dict, ensure_ascii=False, default=str),
"error": response_dict.get("error"),
}
async def download_image(
image_url: str,
dest_path: str,
*,
execution_guard=None,
) -> str:
os.makedirs(os.path.dirname(dest_path), exist_ok=True)
async with httpx.AsyncClient(timeout=300) as client:
async with client.stream("GET", image_url) as response:
response.raise_for_status()
with open(dest_path, "wb") as file:
chunk_no = 0
async for chunk in response.aiter_bytes(chunk_size=8192):
file.write(chunk)
chunk_no += 1
if execution_guard is not None and chunk_no % 32 == 0:
await execution_guard()
if execution_guard is not None:
await execution_guard()
return dest_path