会员积分改版V1
This commit is contained in:
@@ -9,6 +9,7 @@ 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.services.llm_billing.context import LlmProviderPostprocessError
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
|
||||
@@ -74,66 +75,52 @@ async def optimize_prompt(
|
||||
log_owner_type: str | None = None,
|
||||
log_owner_id: str | None = None,
|
||||
generation_attempt_no: int | None = None,
|
||||
fixed_model_config_id: str | None = None,
|
||||
fixed_model_snapshot: dict | 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())
|
||||
stmt = select(ModelConfig).where(
|
||||
ModelConfig.is_active.is_(True),
|
||||
ModelConfig.deleted_at.is_(None),
|
||||
ModelConfig.provider != "mock",
|
||||
)
|
||||
if fixed_model_config_id:
|
||||
stmt = stmt.where(ModelConfig.id == fixed_model_config_id)
|
||||
stmt = stmt.order_by(ModelConfig.priority.desc(), ModelConfig.id.asc()).limit(1)
|
||||
result = await db.execute(stmt)
|
||||
item = result.scalar_one_or_none()
|
||||
if item is None:
|
||||
raise LLMProviderCallError("没有可用的固定LLM模型,本版本禁止Mock、自动切换和降级")
|
||||
snapshot = dict(fixed_model_snapshot or {})
|
||||
selected = SimpleNamespace(
|
||||
id=item.id,
|
||||
name=snapshot.get("name") or item.name,
|
||||
provider=snapshot.get("provider") or item.provider,
|
||||
api_base=snapshot.get("api_base") or item.api_base,
|
||||
api_key=item.api_key,
|
||||
model_name=snapshot.get("model_name") or item.model_name,
|
||||
max_tokens=snapshot.get("max_tokens") if snapshot.get("max_tokens") is not None else item.max_tokens,
|
||||
temperature=snapshot.get("temperature") if snapshot.get("temperature") is not None else item.temperature,
|
||||
)
|
||||
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)
|
||||
if selected.provider not in ("openai_compatible", "sdk"):
|
||||
raise LLMProviderCallError(f"固定模型供应商不受支持:{selected.provider}")
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
def _mock_optimize(prompt: str, gen_type: str = "video") -> tuple[str, dict]:
|
||||
@@ -435,19 +422,28 @@ async def _call_openai_compatible(
|
||||
)
|
||||
raise LLMProviderCallError(f"{type(exc).__name__}: {exc}") from exc
|
||||
|
||||
raw_usage = data.get("usage") if isinstance(data, dict) else None
|
||||
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)
|
||||
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,
|
||||
}
|
||||
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("模型未返回有效提示词")
|
||||
@@ -461,18 +457,5 @@ async def _call_openai_compatible(
|
||||
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,
|
||||
}
|
||||
raise LlmProviderPostprocessError(f"模型响应解析失败: {exc}", usage=token_usage) from exc
|
||||
return content, token_usage
|
||||
|
||||
Reference in New Issue
Block a user