Files

479 lines
20 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import json
import time
from types import SimpleNamespace
import httpx
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.models.model_config import ModelConfig
from app.services.operation_log_service import build_exception_detail, log_ai_model_event
from app.utils.id_gen import generate_id
class LLMProviderCallError(RuntimeError):
"""Remote model call or response validation failed and may use fallback."""
MOCK_OPTIMIZED_PROMPTS = {
"直播": "专业直播间场景,45度斜角机位,暖色柔光打光,主播居中构图,背景虚化处理,产品特写切换流畅,镜头推进节奏感强,画面色彩饱和度高,适合电商直播推广视频。",
"产品": "高端产品展示视频,360度旋转环绕拍摄,纯白/深色渐变背景,柔光箱打光消除阴影,微距镜头捕捉产品细节,金属/玻璃材质高光反射,品牌Logo水印角落显示。",
"课程": "在线教育课程片头,明亮书房环境,讲师半身出镜,板书/屏幕录制无缝切换,字幕条底部滚动,知识要点动画弹出,背景轻音乐辅助,整体色调清新专业。",
"美食": "美食制作过程记录,俯拍+侧拍双机位切换,暖色灯光突出食材质感,慢镜头捕捉烹饪瞬间(蒸汽、油花、酱汁淋洒),成品摆盘精致特写,色调偏暖黄增进食欲。",
"品牌": "品牌形象宣传片,电影级调色(青橙对比色),航拍+地面多角度取景,城市/自然场景交替,人物情感特写穿插,品牌故事旁白叠加,片尾Logo定版动画。",
"游戏": "游戏宣传CG风格视频,高速运镜+粒子特效,角色动态捕捉流畅,技能释放光效炸裂,UI界面模拟叠加,BGM史诗感配乐,画面帧率60fps丝滑体验。",
"product": "Premium product showcase video with 360-degree rotation, clean gradient background, professional studio lighting, macro lens capturing fine details, metallic and glass material highlights, brand watermark in corner.",
"live": "Professional livestream scene with 45-degree angle camera, warm soft lighting, host centered with bokeh background, smooth product close-up transitions, vibrant saturated colors.",
"course": "Online education intro with bright study environment, instructor half-body shot, seamless screen recording transitions, animated key points overlay, clean professional tone.",
"food": "Food preparation recording with overhead and side camera switching, warm lighting highlighting textures, slow-motion cooking moments, elegant plating close-up.",
"brand": "Cinematic brand film with teal-orange color grading, aerial and ground multi-angle shots, urban and nature scenes alternating, emotional character close-ups, brand story narration.",
"game": "Game promotional CG-style video with dynamic camera movements, particle effects, character motion capture, explosive skill light effects, epic BGM, smooth 60fps visuals.",
}
def _get_default_prompt(prompt: str, gen_type: str = "video") -> tuple[str, dict]:
if gen_type == "image":
optimized = (
f"Professional high-quality image with expert composition, precise color grading, "
f"sharp focus and rich details. Theme: {prompt}. Photographic style with "
f"professional lighting and strong visual impact, suitable for commercial use."
)
else:
optimized = (
f"Professionally crafted video with expert composition, precise color grading, "
f"smooth camera movements. Theme: {prompt}. Cinematic shooting techniques with "
f"rich lighting layers and strong visual impact, suitable for commercial distribution."
)
return optimized, {
"input_tokens": 0,
"output_tokens": 0,
"total_tokens": 0,
"usage_reported": True,
"billing_free": True,
}
async def optimize_prompt(
db: AsyncSession,
original_prompt: str,
user_id: str | None = None,
industry_key: str | None = None,
duration: int | None = None,
image_size: str | None = None,
image_proportion: str | None = None,
image_px: str | None = None,
references: list[dict] | None = None,
gen_type: str = "video",
*,
log_module: str = "generation_ai",
log_step: str = "prompt_optimize",
log_project_id: str | None = None,
log_task_id: str | None = None,
log_owner_type: str | None = None,
log_owner_id: str | None = None,
generation_attempt_no: int | None = None,
) -> tuple[str, dict]:
"""Optimize user prompt using LLM. Returns (optimized_text, token_usage_dict)."""
result = await db.execute(
select(ModelConfig)
.where(ModelConfig.is_active == True, ModelConfig.deleted_at.is_(None))
.order_by(ModelConfig.priority.desc())
)
configs = [
SimpleNamespace(
id=item.id,
name=item.name,
provider=item.provider,
api_base=item.api_base,
api_key=item.api_key,
model_name=item.model_name,
max_tokens=item.max_tokens,
temperature=item.temperature,
)
for item in result.scalars().all()
]
# Release the read transaction before the external LLM request. Only
# plain scalar snapshots are used afterwards, so expire_on_commit does
# not trigger an ORM refresh while the provider request is in flight.
await db.commit()
if configs:
# 按 priority 从大到小依次尝试,跳过 mock,失败则用下一个
for selected in configs:
if selected.provider == "mock":
continue
if selected.provider in ("openai_compatible", "sdk"):
try:
return await _call_openai_compatible(
selected, original_prompt, db, user_id, industry_key, duration,
references=references,
gen_type=gen_type,
image_size=image_size,
image_proportion=image_proportion,
image_px=image_px,
log_module=log_module,
log_step=log_step,
log_project_id=log_project_id,
log_task_id=log_task_id,
log_owner_type=log_owner_type,
log_owner_id=log_owner_id,
generation_attempt_no=generation_attempt_no,
)
except LLMProviderCallError:
continue
# 所有真实模型都失败,降级到 mock
mock_cfg = next((c for c in configs if c.provider == "mock"), None)
if mock_cfg:
return _mock_optimize(original_prompt, gen_type)
if settings.LLM_MOCK:
return _mock_optimize(original_prompt, gen_type)
return _get_default_prompt(original_prompt, gen_type)
def _mock_optimize(prompt: str, gen_type: str = "video") -> tuple[str, dict]:
"""Return a keyword-matched mock optimized prompt."""
for keyword, optimized in MOCK_OPTIMIZED_PROMPTS.items():
if keyword in prompt:
return optimized, {
"input_tokens": 0,
"output_tokens": 0,
"total_tokens": 0,
"usage_reported": True,
"billing_free": True,
}
return _get_default_prompt(prompt, gen_type)
async def _call_openai_compatible(
config: ModelConfig,
original_prompt: str,
db: AsyncSession | None = None,
user_id: str | None = None,
industry_key: str | None = None,
duration: int | None = None,
references: list[dict] | None = None,
gen_type: str = "video",
image_size: str | None = None,
image_proportion: str | None = None,
image_px: str | None = None,
*,
log_module: str = "generation_ai",
log_step: str = "prompt_optimize",
log_project_id: str | None = None,
log_task_id: str | None = None,
log_owner_type: str | None = None,
log_owner_id: str | None = None,
generation_attempt_no: int | None = None,
) -> tuple[str, dict]:
"""Call an OpenAI-compatible API to optimize the prompt. Returns (content, token_usage)."""
system_prompt = None
if industry_key and db is not None:
from app.models.industry_config import IndustryConfig
result = await db.execute(
select(IndustryConfig).where(IndustryConfig.key == industry_key).limit(1)
)
ind = result.scalar_one_or_none()
if ind and ind.skills:
try:
skills = json.loads(ind.skills)
if gen_type == "image":
skill_keys = ("文图理解生成图片提示词", "文图理解生成图片", "文图理解")
else:
skill_keys = ("文图理解生成视频提示词", "文图理解")
for s in skills:
if s.get("key") in skill_keys and s.get("label"):
system_prompt = s["label"]
break
except (json.JSONDecodeError, TypeError):
pass
await db.commit()
if not system_prompt:
if gen_type == "image":
system_prompt = (
"你是一位专业的摄影师和图像创作专家。请根据用户提供的主题,"
"生成一段详细的、专业的图片生成提示词。要求:\n"
"1. 包含具体的画面构图\n"
"2. 描述光影效果和色调\n"
"3. 指定拍摄角度和镜头类型\n"
"4. 画面风格和质感\n"
"5. 整体不超过200字"
)
else:
system_prompt = (
"你是一位专业的视频导演和文案专家。请根据用户提供的视频主题,"
"生成一段详细的、专业的视频生成提示词。要求:\n"
"1. 包含具体的镜头语言(机位、运镜方式)\n"
"2. 描述光影效果和色调\n"
"3. 画面构图和视觉层次\n"
"4. 适合的节奏感和转场\n"
"5. 整体不超过200字"
)
if gen_type == "video" and duration:
system_prompt += f"\n\n请根据视频总时长{duration}秒,合理分配镜头节奏,生成适合{duration}秒视频的提示词。"
user_content = f"{original_prompt}\n\n请生成一段{duration}秒的视频提示词。"
elif gen_type == "image" and image_size:
system_prompt += f"\n\n请根据画面分辨率:{image_size},宽高比:{image_proportion},宽高像素值:{image_px},生成适合该分辨率的图片提示词。"
user_content = f"{original_prompt}\n\n请生成适合分辨率:{image_size},宽高比:{image_proportion},宽高像素值:{image_px}的图片提示词。"
else:
user_content = original_prompt
# Build multimodal user message content when images are present
image_urls = []
video_urls = []
if references:
for ref in references:
if not isinstance(ref, dict):
continue
ref_type = ref.get("type")
ref_url = ref.get("url")
if not ref_url:
continue
if ref_type == "image":
image_urls.append(ref_url)
if ref_type == "video":
video_urls.append(ref_url)
async def _build_multimodal_content(
user_content: str,
image_urls: list[str],
video_urls: list[str],
*,
as_base64: bool,
) -> tuple[dict, dict | None]:
"""构建多模态 user_message。返回 (actual_message, log_message)。"""
from app.utils.media import media_to_base64
content_parts = [{"type": "text", "text": user_content}]
for img in image_urls:
if as_base64:
url = await media_to_base64(img, "image/png")
else:
url = _file_url_or_data_uri(img, "image/png")
content_parts.append({
"type": "image_url",
"image_url": {"url": url},
})
for video in video_urls:
if as_base64:
url = await media_to_base64(video, "video/mp4")
else:
url = _file_url_or_data_uri(video, "video/mp4")
content_parts.append({
"type": "video_url",
"video_url": {"url": url},
})
user_message = {
"role": "user",
"content": content_parts,
}
# Log-friendly version: keep original paths instead of base64
log_content_parts = [{"type": "text", "text": user_content}]
for img in image_urls:
log_content_parts.append({
"type": "image_url",
"image_url": {"url": img},
})
for video in video_urls:
log_content_parts.append({
"type": "video_url",
"video_url": {"url": video},
})
log_user_message = {
"role": "user",
"content": log_content_parts,
}
return user_message, log_user_message
def _file_url_or_data_uri(file_url: str, fallback_mime: str) -> str:
"""
远程 URL 保持原样;本地上传路径拼接 BASE_URL 转为完整链接。
"""
if file_url.startswith(("http://", "https://", "data:")):
return file_url
# 本地上传路径拼接 BASE_URL,例如 /uploads/images/xxx.png → https://domain.com/uploads/images/xxx.png
base = settings.BASE_URL.rstrip("/")
path = file_url if file_url.startswith("/") else f"/{file_url}"
return f"{base}{path}"
if image_urls or video_urls:
from app.utils.media import get_llm_media_as_base64
as_base64 = await get_llm_media_as_base64(db)
# 配置读取后立即结束事务;后续 URL 下载/Base64 转换属于外部 I/O
# 不能继续占用数据库连接。
if db is not None:
await db.commit()
user_message, log_user_message = await _build_multimodal_content(
user_content,
image_urls,
video_urls,
as_base64=as_base64,
)
else:
user_message = {
"role": "user",
"content": user_content,
}
log_user_message = None
# 防御性结束可能由配置读取开启的只读事务;HTTP 请求期间不占用数据库连接。
if db is not None:
await db.commit()
async with httpx.AsyncClient(timeout=120) as client:
request_data = {
"model": config.model_name,
"messages": [
{"role": "system", "content": system_prompt},
user_message,
],
"max_tokens": config.max_tokens,
"temperature": config.temperature,
}
call_id = generate_id()
started = time.perf_counter()
common_log = {
"module": log_module,
"step_code": log_step,
"call_id": call_id,
"source": "app.services.llm",
"user_id": user_id,
"project_id": log_project_id,
"task_id": log_task_id,
"owner_type": log_owner_type,
"owner_id": log_owner_id,
"generation_attempt_no": generation_attempt_no,
"model_config_id": config.id,
"model_config_name": config.name,
"model_name": config.model_name,
"provider": config.provider,
"api_base": config.api_base,
"remote_action": "chat_completions",
}
log_ai_model_event(
event_type="REQUEST",
event_phase="REQUEST",
event_status="started",
request=request_data,
**common_log,
)
try:
response = await client.post(
f"{config.api_base}/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:
error_body = response.text
log_ai_model_event(
event_type="RESPONSE",
event_phase="RESPONSE",
event_status="failed",
http_status=response.status_code,
latency_ms=latency_ms,
response={"body": error_body},
error=f"HTTP {response.status_code}",
**common_log,
)
error = LLMProviderCallError(f"HTTP {response.status_code}: {error_body}")
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
http_status=response.status_code,
latency_ms=latency_ms,
detail=build_exception_detail(error),
error=str(error),
**common_log,
)
raise error
data = response.json()
log_ai_model_event(
event_type="RESPONSE",
event_phase="RESPONSE",
event_status="success",
http_status=response.status_code,
latency_ms=latency_ms,
response=data,
token_usage=data.get("usage") if isinstance(data, dict) else None,
**common_log,
)
except LLMProviderCallError:
raise
except Exception as exc:
latency_ms = int((time.perf_counter() - started) * 1000)
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
latency_ms=latency_ms,
detail=build_exception_detail(exc),
error=str(exc),
**common_log,
)
raise LLMProviderCallError(f"{type(exc).__name__}: {exc}") from exc
try:
raw_usage = data.get("usage")
usage_reported = bool(
isinstance(raw_usage, dict)
and any(
key in raw_usage
for key in ("prompt_tokens", "completion_tokens", "total_tokens")
)
)
usage = raw_usage if isinstance(raw_usage, dict) else {}
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["choices"][0]["message"]["content"].strip()
if not content:
raise ValueError("模型未返回有效提示词")
except Exception as exc:
log_ai_model_event(
event_type="ERROR",
event_phase="ERROR",
event_status="failed",
latency_ms=int((time.perf_counter() - started) * 1000),
detail=build_exception_detail(exc, {"stage": "response_validation"}),
error=str(exc),
**common_log,
)
raise LLMProviderCallError(f"模型响应解析失败: {exc}") from exc
token_usage = {
"model_config_id": config.id,
"model_config_name": config.name,
"model_provider": config.provider,
"model_name": config.model_name,
"source_module": log_module,
"source_step_code": log_step,
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"total_tokens": total_tokens,
"usage_reported": usage_reported,
}
return content, token_usage