Files
video-gen/video-gen-api/app/services/image_gen.py
T

250 lines
8.7 KiB
Python

import base64
import json
import logging
import mimetypes
import os
from datetime import datetime
import httpx
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from volcenginesdkarkruntime import AsyncArk
from app.config import settings
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 (
ProviderGenerationRecordLike,
ProviderImageEngineLike,
)
logger = logging.getLogger("videogen")
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:
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)
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 f:
f.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:
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) 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 f:
f.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)
.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:")):
return url
return f"{settings.BASE_URL.rstrip('/')}/{url.lstrip('/')}"
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,
)
prompt = record.optimized_prompt or record.original_prompt
image_urls = []
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)
except (json.JSONDecodeError, TypeError):
pass
request_payload = {
"model": engine.model_name,
"prompt": 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,
}
if image_urls:
request_payload["image"] = image_urls
_log_image_request(engine, record.id, request_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,
}
}
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
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 "",
}
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()
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:
"""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:
async for chunk in response.aiter_bytes(chunk_size=8192):
f.write(chunk)
return dest_path