488 lines
19 KiB
Python
488 lines
19 KiB
Python
import json
|
|
import logging
|
|
import os
|
|
import time
|
|
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.operation_log_service import build_exception_detail, log_ai_model_event
|
|
from app.utils.id_gen import generate_id
|
|
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 _provider_log_context(engine, record, *, call_id: str, step_code: str) -> dict:
|
|
generation_mode = str(getattr(record, "generation_mode", "") or "generation_record")
|
|
owner_type = "chat_generation_task" if generation_mode != "generation_record" else "generation_record"
|
|
return {
|
|
"module": generation_mode,
|
|
"step_code": step_code,
|
|
"call_id": call_id,
|
|
"source": "app.services.image_gen",
|
|
"user_id": str(getattr(record, "user_id", "") or "") or None,
|
|
"project_id": str(getattr(record, "project_id", "") or "") or None,
|
|
"task_id": str(getattr(record, "id", "") or "") or None,
|
|
"owner_type": owner_type,
|
|
"owner_id": str(getattr(record, "id", "") or "") or None,
|
|
"generation_attempt_no": int(getattr(record, "generation_attempt_no", 1) or 1),
|
|
"model_config_id": str(getattr(engine, "id", "") or "") or None,
|
|
"model_config_name": str(getattr(engine, "name", "") or "") or None,
|
|
"model_name": str(getattr(engine, "model_name", "") or "") or None,
|
|
"provider": str(getattr(engine, "provider", "") or "") or None,
|
|
"api_base": str(getattr(engine, "api_base", "") or "") or None,
|
|
}
|
|
|
|
|
|
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
|
|
|
|
call_id = generate_id()
|
|
started = time.perf_counter()
|
|
api_step = "image_sync_batch_create" if count > 1 else "image_sync_create"
|
|
log_context = _provider_log_context(engine, record, call_id=call_id, step_code=api_step)
|
|
log_ai_model_event(
|
|
event_type="REQUEST",
|
|
event_phase="REQUEST",
|
|
event_status="started",
|
|
remote_action=api_step,
|
|
request=request_log_payload,
|
|
**log_context,
|
|
)
|
|
|
|
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_ai_model_event(
|
|
event_type="RESPONSE",
|
|
event_phase="RESPONSE",
|
|
event_status="success",
|
|
remote_action=api_step,
|
|
latency_ms=int((time.perf_counter() - started) * 1000),
|
|
response=response_data,
|
|
token_usage=response_data.get("usage"),
|
|
**log_context,
|
|
)
|
|
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_ai_model_event(
|
|
event_type="ERROR",
|
|
event_phase="ERROR",
|
|
event_status="failed",
|
|
remote_action=api_step,
|
|
http_status=provider_error.http_status,
|
|
remote_request_id=provider_error.provider_request_id,
|
|
latency_ms=int((time.perf_counter() - started) * 1000),
|
|
detail=build_exception_detail(exc, provider_error.as_dict()),
|
|
error=provider_error.safe_message,
|
|
**log_context,
|
|
)
|
|
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
|