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) .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) -> 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: async for chunk in response.aiter_bytes(chunk_size=8192): file.write(chunk) return dest_path