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

185 lines
7.1 KiB
Python

from __future__ import annotations
import json
import mimetypes
import os
import time
from typing import Any
import httpx
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask
from app.models.model_config import ModelConfig
from app.models.token_usage import TokenUsage
from app.services.generation_log_service import log_provider_call
from app.services.provider_limit import provider_limit
from app.utils.id_gen import generate_id
def _absolute_url(url: str) -> str:
if url.startswith("http://") or url.startswith("https://") or url.startswith("data:"):
return url
base = settings.BASE_URL.rstrip("/")
return f"{base}/{url.lstrip('/')}"
def _load_refs(record: ChatGenerationTask) -> list[dict]:
if not record.media_references:
return []
try:
data = json.loads(record.media_references)
return data if isinstance(data, list) else []
except Exception:
return []
def _build_user_content(record: ChatGenerationTask) -> list[dict[str, Any]]:
if record.gen_type == "image":
params = f"图片参数:分辨率档位={record.image_size or '2K'},比例={record.image_proportion or '1:1'},像素={record.image_px or '2048x2048'}"
else:
params = f"视频参数:时长={record.duration or 4}秒,比例={record.aspect_ratio or '16:9'},分辨率={record.resolution or '480p'}"
text = (
f"生成类型:{record.gen_type}\n"
f"{params}\n"
f"用户描述:{record.original_prompt}\n\n"
"请只输出最终可直接用于图片/视频生成模型的 prompt,不要说你已经生成了图片或视频。"
)
parts: list[dict[str, Any]] = [{"type": "text", "text": text}]
for ref in _load_refs(record):
ref_type = ref.get("type")
ref_url = ref.get("url") or ""
if not ref_url:
continue
url = _absolute_url(ref_url)
if ref_type == "image":
parts.append({"type": "image_url", "image_url": {"url": url}})
elif ref_type == "video":
parts.append({"type": "video_url", "video_url": {"url": url, "fps": settings.CHATAPI_VIDEO_FPS}})
return parts
async def _get_model_config(db: AsyncSession) -> ModelConfig:
result = await db.execute(
select(ModelConfig)
.where(ModelConfig.is_active == True)
.order_by(ModelConfig.priority.desc())
.limit(1)
)
config = result.scalar_one_or_none()
if not config:
raise ValueError("没有可用的ChatAPI模型配置")
if config.provider == "mock":
return config
if not config.api_base or not config.api_key or not config.model_name:
raise ValueError("ChatAPI模型配置不完整")
return config
async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask) -> tuple[str, dict]:
"""Call ChatAPI once with current request params and attachments. No history context."""
config = await _get_model_config(db)
if config.provider == "mock":
return record.original_prompt, {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
system_prompt = (
"你是图片/视频生成提示词整理助手。你的职责是根据用户文字、上传图片/视频和生成参数,"
"整理最终可直接用于生成模型的 prompt。不要声称你已经生成图片或视频,不要调用工具。"
"输出中文为主,内容具体、可执行,保留用户关键要求。"
)
request_data = {
"model": config.model_name,
"messages": [
{"role": "system", "content": system_prompt},
{"role": "user", "content": _build_user_content(record)},
],
"max_tokens": config.max_tokens,
"temperature": config.temperature,
}
started = time.perf_counter()
async with provider_limit("ark_chat_prompt", settings.ARK_CHAT_PROMPT_MAX_CONCURRENCY):
async with httpx.AsyncClient(timeout=settings.CHATAPI_REQUEST_TIMEOUT_SECONDS) as client:
try:
response = await client.post(
f"{config.api_base.rstrip('/')}/chat/completions",
headers={
"Authorization": f"Bearer {config.api_key}",
"Content-Type": "application/json",
},
json=request_data,
)
latency_ms = int((time.perf_counter() - started) * 1000)
if response.status_code >= 400:
await log_provider_call(
record,
provider=config.provider,
api_type="chat_prompt",
model=config.model_name,
engine_id=record.engine_id,
status="failed",
latency_ms=latency_ms,
http_status=response.status_code,
request_data=request_data,
response_data=response.text,
error_message=response.text[:1000],
)
raise RuntimeError(f"ChatAPI HTTP {response.status_code}: {response.text}")
data = response.json()
except Exception as exc:
latency_ms = int((time.perf_counter() - started) * 1000)
await log_provider_call(
record,
provider=config.provider,
api_type="chat_prompt",
model=config.model_name,
engine_id=record.engine_id,
status="failed",
latency_ms=latency_ms,
request_data=request_data,
response_data=None,
error_message=str(exc),
)
raise
usage = data.get("usage", {}) or {}
input_tokens = int(usage.get("prompt_tokens", 0) or 0)
output_tokens = int(usage.get("completion_tokens", 0) or 0)
total_tokens = int(usage.get("total_tokens", input_tokens + output_tokens) or 0)
content = data.get("choices", [{}])[0].get("message", {}).get("content", "").strip()
if not content:
raise RuntimeError("ChatAPI未返回有效prompt")
db.add(TokenUsage(
id=generate_id(),
model_config_id=config.id,
user_id=record.user_id,
input_tokens=input_tokens,
output_tokens=output_tokens,
total_tokens=total_tokens,
))
await db.flush()
await log_provider_call(
record,
provider=config.provider,
api_type="chat_prompt",
model=config.model_name,
engine_id=record.engine_id,
status="success",
latency_ms=int((time.perf_counter() - started) * 1000),
http_status=200,
request_data=request_data,
response_data=data,
prompt_tokens=input_tokens,
completion_tokens=output_tokens,
total_tokens=total_tokens,
)
return content, {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"total_tokens": total_tokens,
}