AI创作批量生成任务 main V1 init
This commit is contained in:
@@ -1,9 +1,8 @@
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
import mimetypes
|
||||
import os
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy import select
|
||||
@@ -11,10 +10,16 @@ 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 is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data
|
||||
from app.services.generation_provider_types import (
|
||||
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,
|
||||
)
|
||||
@@ -22,8 +27,39 @@ from app.services.generation_provider_types import (
|
||||
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):
|
||||
"""Log image generation request to log/AiModel/YYYY-MM-DD.log"""
|
||||
if not is_enabled():
|
||||
return
|
||||
try:
|
||||
@@ -41,14 +77,13 @@ def _log_image_request(engine: ProviderImageEngineLike, record_id: str, request_
|
||||
"request": request_encrypted,
|
||||
"request_length": len(request_str),
|
||||
}
|
||||
with open(log_file, "a", encoding="utf-8") as f:
|
||||
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
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):
|
||||
"""Log image generation response to log/AiModel/YYYY-MM-DD.log"""
|
||||
if not is_enabled():
|
||||
return
|
||||
try:
|
||||
@@ -63,17 +98,13 @@ def _log_image_response(record_id: str, response_data: dict, error: str | None =
|
||||
"response": response_encrypted,
|
||||
"error": error,
|
||||
}
|
||||
with open(log_file, "a", encoding="utf-8") as f:
|
||||
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
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:
|
||||
"""Get the active image engine with highest priority."""
|
||||
result = await db.execute(
|
||||
select(ImageEngine)
|
||||
.where(ImageEngine.is_active == True)
|
||||
@@ -103,105 +134,291 @@ def _resolve_url(url: str) -> str:
|
||||
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,
|
||||
) -> dict:
|
||||
"""Submit an image generation task via Ark SDK. Returns image_url."""
|
||||
from volcenginesdkarkruntime import Ark
|
||||
|
||||
client = Ark(
|
||||
base_url=engine.api_base,
|
||||
api_key=engine.api_key,
|
||||
timeout=300,
|
||||
)
|
||||
generation_count: int = 1,
|
||||
) -> ImageProviderBatchResult:
|
||||
"""通过 Ark 同步图片接口生成单图或单次组图。
|
||||
|
||||
prompt = record.optimized_prompt or record.original_prompt
|
||||
image_urls = []
|
||||
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:
|
||||
ref_type = ref.get("type")
|
||||
ref_url = ref.get("url", "")
|
||||
if ref_type == "image" and ref_url:
|
||||
resolved = _resolve_url(ref_url)
|
||||
image_urls.append(resolved)
|
||||
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):
|
||||
pass
|
||||
image_urls = []
|
||||
|
||||
request_payload = {
|
||||
request_log_payload: dict[str, Any] = {
|
||||
"model": engine.model_name,
|
||||
"prompt": prompt,
|
||||
"prompt": provider_prompt,
|
||||
"size": record.image_size or engine.default_size,
|
||||
"sequential_image_generation": "disabled",
|
||||
"output_format": "png",
|
||||
"response_format": "url",
|
||||
"watermark": False,
|
||||
"include_media_references": include_media_references,
|
||||
}
|
||||
|
||||
request_sdk_payload: dict[str, Any] = dict(request_log_payload)
|
||||
if image_urls:
|
||||
request_payload["image"] = 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
|
||||
|
||||
_log_image_request(engine, record.id, request_payload)
|
||||
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(
|
||||
model=engine.model_name,
|
||||
prompt=prompt,
|
||||
size=record.image_size or engine.default_size,
|
||||
output_format="png",
|
||||
response_format="url",
|
||||
watermark=False,
|
||||
image=image_urls if image_urls else None,
|
||||
)
|
||||
image_url = result.data[0].url
|
||||
|
||||
response_data = {
|
||||
"model": result.model,
|
||||
"created": result.created,
|
||||
"data": [{"url": item.url, "size": item.size} for item in result.data] if result.data else [],
|
||||
"usage": {
|
||||
"generated_images": result.usage.generated_images if hasattr(result.usage, 'generated_images') else 0,
|
||||
"output_tokens": result.usage.output_tokens if hasattr(result.usage, 'output_tokens') else 0,
|
||||
"total_tokens": result.usage.total_tokens if hasattr(result.usage, 'total_tokens') else 0,
|
||||
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,
|
||||
},
|
||||
}
|
||||
except httpx.TimeoutException:
|
||||
error_msg = "图片生成超时,请稍后重试"
|
||||
logger.error(f"Image generation timeout for record {record.id}")
|
||||
_log_image_response(record.id, {}, error_msg)
|
||||
raise TimeoutError(error_msg)
|
||||
except Exception as e:
|
||||
error_msg = str(e)
|
||||
logger.error(f"Image generation failed for record {record.id}: {error_msg}")
|
||||
_log_image_response(record.id, {}, error_msg)
|
||||
raise
|
||||
_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:
|
||||
client.close()
|
||||
|
||||
return {
|
||||
"image_url": image_url,
|
||||
"image_tokens": getattr(result.usage, "total_tokens", 0),
|
||||
"response_data": json.dumps(response_data, ensure_ascii=False, default=str),
|
||||
"error": str(result.error) if result.error else "",
|
||||
}
|
||||
try:
|
||||
client.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def poll_image_task_status(engine: ImageEngine, task_id: str) -> dict:
|
||||
"""Query image task status via Ark SDK. Returns {status, image_url, response_data}."""
|
||||
client = AsyncArk(
|
||||
base_url=engine.api_base,
|
||||
api_key=engine.api_key,
|
||||
)
|
||||
|
||||
result = await client.image_generation.tasks.get(task_id=task_id)
|
||||
await client.close()
|
||||
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,
|
||||
@@ -239,13 +456,11 @@ async def poll_image_task_status(engine: ImageEngine, task_id: str) -> dict:
|
||||
|
||||
|
||||
async def download_image(image_url: str, dest_path: str) -> str:
|
||||
"""Download image to local storage."""
|
||||
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 f:
|
||||
with open(dest_path, "wb") as file:
|
||||
async for chunk in response.aiter_bytes(chunk_size=8192):
|
||||
f.write(chunk)
|
||||
return dest_path
|
||||
file.write(chunk)
|
||||
return dest_path
|
||||
|
||||
Reference in New Issue
Block a user