250 lines
8.7 KiB
Python
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 |