648 lines
25 KiB
Python
648 lines
25 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
from dataclasses import dataclass
|
|
from typing import Any
|
|
|
|
from fastapi import HTTPException
|
|
from sqlalchemy import select
|
|
from sqlalchemy.exc import IntegrityError
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.enums.common import LogEventStatusEnum
|
|
from app.enums.credit_record import (
|
|
CreditRecordBillingScene,
|
|
CreditRecordChargeKind,
|
|
CreditRecordSourceModule,
|
|
)
|
|
from app.enums.generation_record import (
|
|
GenerationRecordConfigSourceEnum,
|
|
GenerationRecordEventTypeEnum,
|
|
)
|
|
from app.enums.generation_status import (
|
|
ASPECT_RATIOS,
|
|
DURATIONS,
|
|
IMAGE_SIZES,
|
|
RESOLUTIONS,
|
|
GenerationStatus,
|
|
GenerationType,
|
|
)
|
|
from app.enums.llm_billing import LlmBillingConfigKey
|
|
from app.models.generation_record import GenerationRecord
|
|
from app.models.project import Project
|
|
from app.schemas.generation import OptimizeParams
|
|
from app.services.error_codes import extract_error_message
|
|
from app.services.generation.ai.engine_service import (
|
|
get_image_engine,
|
|
get_video_engine,
|
|
image_supported_sizes,
|
|
parse_json_list,
|
|
)
|
|
from app.services.generation.billing_service import OWNER_GENERATION_RECORD
|
|
from app.services.generation.media_reference_service import (
|
|
calculate_media_reference_usage,
|
|
validate_media_reference_usage_for_engine,
|
|
)
|
|
from app.services.generation.pipeline.generation_record_config_service import (
|
|
freeze_generation_record_config_with_log,
|
|
is_generation_record_config_complete,
|
|
)
|
|
from app.services.llm import optimize_prompt
|
|
from app.services.llm_billing import (
|
|
LlmBillingContext,
|
|
log_provider_failure,
|
|
log_provider_start,
|
|
log_provider_success,
|
|
release_on_failure,
|
|
settle_success,
|
|
start_hold,
|
|
)
|
|
from app.services.operation_log_service import log_operation_error, log_operation_event
|
|
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
|
|
from app.utils.id_gen import generate_id
|
|
|
|
logger = logging.getLogger("videogen")
|
|
|
|
_PROMPT_ATTEMPT_NO = 1
|
|
_LOG_DOMAIN = "generation_record"
|
|
_LOG_MODULE = "generation_record"
|
|
_LOG_SOURCE = GenerationRecordConfigSourceEnum.PROMPT_OPTIMIZE.value
|
|
|
|
|
|
@dataclass(slots=True, frozen=True)
|
|
class PromptOptimizeServiceResult:
|
|
record_id: str
|
|
idempotent: bool = False
|
|
|
|
|
|
def _log_event(
|
|
event: GenerationRecordEventTypeEnum,
|
|
*,
|
|
status: LogEventStatusEnum = LogEventStatusEnum.SUCCESS,
|
|
user_id: str,
|
|
project_id: str | None,
|
|
record_id: str | None,
|
|
detail: dict[str, Any] | None = None,
|
|
error: str | None = None,
|
|
) -> None:
|
|
log_operation_event(
|
|
domain=_LOG_DOMAIN,
|
|
module=_LOG_MODULE,
|
|
event_type=event.value,
|
|
event_status=status.value,
|
|
source=_LOG_SOURCE,
|
|
user_id=user_id,
|
|
project_id=project_id,
|
|
task_id=record_id,
|
|
detail=detail,
|
|
error=error,
|
|
)
|
|
|
|
|
|
def _billing_context(*, user_id: str, record_id: str, request_id: str | None) -> LlmBillingContext:
|
|
return LlmBillingContext(
|
|
user_id=user_id,
|
|
owner_type=OWNER_GENERATION_RECORD,
|
|
owner_id=record_id,
|
|
attempt_no=_PROMPT_ATTEMPT_NO,
|
|
charge_kind=CreditRecordChargeKind.TEXT_PROMPT.value,
|
|
billing_scene=CreditRecordBillingScene.GENERATION_RECORD_TEXT_PROMPT_OPTIMIZE.value,
|
|
source_module=CreditRecordSourceModule.GENERATION_RECORD.value,
|
|
related_id=record_id,
|
|
hold_config_key=LlmBillingConfigKey.HOLD_GENERATION_RECORD_PROMPT.value,
|
|
description_prefix="AI创作提示词优化",
|
|
trace_id=f"generation-optimize:{record_id}",
|
|
request_id=request_id,
|
|
)
|
|
|
|
|
|
def _canonical_references(value: object) -> str:
|
|
return json.dumps(value or [], ensure_ascii=False, sort_keys=True, separators=(",", ":"))
|
|
|
|
|
|
def _idempotency_config_matches(record: GenerationRecord, req: OptimizeParams) -> bool:
|
|
try:
|
|
existing_references = json.loads(record.media_references) if record.media_references else []
|
|
except (TypeError, json.JSONDecodeError):
|
|
return False
|
|
if (
|
|
str(record.project_id) != str(req.project_id)
|
|
or record.original_prompt != req.prompt
|
|
or record.gen_type != req.gen_type.value
|
|
or str(record.engine_id or "") != str(req.engine_id)
|
|
or bool(record.include_media_references) != bool(req.include_media_references)
|
|
or _canonical_references(existing_references) != _canonical_references(req.references)
|
|
):
|
|
return False
|
|
if req.gen_type == GenerationType.video:
|
|
return (
|
|
record.duration == req.duration
|
|
and record.aspect_ratio == req.aspect_ratio
|
|
and record.resolution == req.resolution
|
|
)
|
|
return (
|
|
record.image_size == req.image_size
|
|
and record.image_proportion == req.image_proportion
|
|
and record.image_px == req.image_px
|
|
)
|
|
|
|
|
|
def _validate_video_engine_selection(engine: Any, *, aspect_ratio: str, resolution: str, duration: int) -> None:
|
|
ratios = [str(item) for item in parse_json_list(engine.supported_ratios, [])]
|
|
resolutions = [str(item) for item in parse_json_list(engine.supported_resolutions, [])]
|
|
durations = [int(item) for item in parse_json_list(engine.supported_durations, []) if str(item).isdigit()]
|
|
if ratios and aspect_ratio not in ratios:
|
|
raise HTTPException(status_code=400, detail="当前视频引擎不支持所选画面比例")
|
|
if resolutions and resolution not in resolutions:
|
|
raise HTTPException(status_code=400, detail="当前视频引擎不支持所选分辨率")
|
|
if durations and duration not in durations:
|
|
raise HTTPException(status_code=400, detail="当前视频引擎不支持所选时长")
|
|
if int(engine.max_duration or 0) > 0 and duration > int(engine.max_duration):
|
|
raise HTTPException(status_code=400, detail="生成时长超过当前视频引擎上限")
|
|
|
|
|
|
def _validate_image_engine_selection(engine: Any, *, image_size: str) -> None:
|
|
sizes = image_supported_sizes(engine)
|
|
if sizes and image_size not in sizes:
|
|
raise HTTPException(status_code=400, detail="当前图片引擎不支持所选画面分辨率")
|
|
|
|
|
|
async def _find_idempotency_record(
|
|
db: AsyncSession,
|
|
*,
|
|
user_id: str,
|
|
idempotency_key: str | None,
|
|
) -> GenerationRecord | None:
|
|
if not idempotency_key:
|
|
return None
|
|
stmt = (
|
|
select(GenerationRecord)
|
|
.where(
|
|
GenerationRecord.user_id == user_id,
|
|
GenerationRecord.idempotency_key == idempotency_key,
|
|
GenerationRecord.deleted_at.is_(None),
|
|
)
|
|
.order_by(GenerationRecord.created_at.desc(), GenerationRecord.id.desc())
|
|
.limit(1)
|
|
)
|
|
result = await db.execute(stmt)
|
|
return result.scalar_one_or_none()
|
|
|
|
|
|
async def _settle_staged_result(
|
|
db: AsyncSession,
|
|
*,
|
|
record_id: str,
|
|
user_id: str,
|
|
project_name: str,
|
|
request_id: str | None,
|
|
) -> PromptOptimizeServiceResult:
|
|
result = await db.execute(
|
|
select(GenerationRecord)
|
|
.where(
|
|
GenerationRecord.id == record_id,
|
|
GenerationRecord.user_id == user_id,
|
|
GenerationRecord.deleted_at.is_(None),
|
|
)
|
|
.with_for_update()
|
|
.limit(1)
|
|
)
|
|
record = result.scalar_one_or_none()
|
|
if record is None:
|
|
raise HTTPException(status_code=404, detail="生成记录不存在")
|
|
if record.status in {
|
|
GenerationStatus.prompt_optimized.value,
|
|
GenerationStatus.generating.value,
|
|
GenerationStatus.completed.value,
|
|
}:
|
|
await db.rollback()
|
|
return PromptOptimizeServiceResult(record_id=record_id, idempotent=True)
|
|
if record.status != GenerationStatus.settlement_pending.value:
|
|
await db.rollback()
|
|
raise HTTPException(status_code=409, detail=f"当前提词状态不可结算:{record.status}")
|
|
if not record.optimized_prompt or not record.prompt_usage_snapshot_json:
|
|
await db.rollback()
|
|
raise HTTPException(status_code=409, detail="提词结果或计费快照缺失,需人工排查")
|
|
|
|
try:
|
|
usage = json.loads(record.prompt_usage_snapshot_json)
|
|
except (TypeError, json.JSONDecodeError) as exc:
|
|
await db.rollback()
|
|
raise HTTPException(status_code=409, detail="提词计费快照损坏,需人工排查") from exc
|
|
if not isinstance(usage, dict):
|
|
await db.rollback()
|
|
raise HTTPException(status_code=409, detail="提词计费快照格式错误,需人工排查")
|
|
|
|
project_id_snapshot = str(record.project_id)
|
|
status_snapshot = str(record.status)
|
|
ctx = _billing_context(user_id=user_id, record_id=record_id, request_id=request_id)
|
|
_log_event(
|
|
GenerationRecordEventTypeEnum.PROMPT_OPTIMIZE_SETTLEMENT_PENDING,
|
|
status=LogEventStatusEnum.STARTED,
|
|
user_id=user_id,
|
|
project_id=project_id_snapshot,
|
|
record_id=record_id,
|
|
detail={"status": status_snapshot, "attempt_no": _PROMPT_ATTEMPT_NO},
|
|
)
|
|
try:
|
|
billing = await settle_success(
|
|
db,
|
|
ctx,
|
|
usage=usage,
|
|
description=f"提示词优化 - {project_name}",
|
|
)
|
|
charge_item = next(
|
|
(item for item in billing.items if item.biz_key == ctx.charge_biz_key),
|
|
None,
|
|
)
|
|
record.text_credits_cost = round(float(charge_item.amount if charge_item else 0.0), 2)
|
|
record.text_tokens_used = int(usage.get("total_tokens", 0) or 0)
|
|
record.status = GenerationStatus.prompt_optimized.value
|
|
record.pipeline_stage = None
|
|
record.error_message = None
|
|
credits_snapshot = float(record.text_credits_cost or 0.0)
|
|
tokens_snapshot = int(record.text_tokens_used or 0)
|
|
await db.commit()
|
|
except Exception as exc:
|
|
await db.rollback()
|
|
_log_event(
|
|
GenerationRecordEventTypeEnum.PROMPT_OPTIMIZE_SETTLEMENT_PENDING,
|
|
status=LogEventStatusEnum.FAILED,
|
|
user_id=user_id,
|
|
project_id=project_id_snapshot,
|
|
record_id=record_id,
|
|
detail={"attempt_no": _PROMPT_ATTEMPT_NO, "error_type": type(exc).__name__},
|
|
error=str(exc),
|
|
)
|
|
raise HTTPException(status_code=503, detail="提词已生成,积分结算暂未完成,请使用相同幂等键重试") from exc
|
|
|
|
_log_event(
|
|
GenerationRecordEventTypeEnum.PROMPT_OPTIMIZE_SETTLEMENT_SUCCESS,
|
|
user_id=user_id,
|
|
project_id=project_id_snapshot,
|
|
record_id=record_id,
|
|
detail={
|
|
"attempt_no": _PROMPT_ATTEMPT_NO,
|
|
"text_credits_cost": credits_snapshot,
|
|
"text_tokens_used": tokens_snapshot,
|
|
},
|
|
)
|
|
return PromptOptimizeServiceResult(record_id=record_id, idempotent=False)
|
|
|
|
|
|
async def _handle_existing_record(
|
|
db: AsyncSession,
|
|
*,
|
|
record: GenerationRecord,
|
|
req: OptimizeParams,
|
|
user_id: str,
|
|
project_name: str,
|
|
) -> PromptOptimizeServiceResult:
|
|
record_id = str(record.id)
|
|
project_id = str(record.project_id)
|
|
status = str(record.status)
|
|
if not _idempotency_config_matches(record, req):
|
|
await db.rollback()
|
|
raise HTTPException(status_code=409, detail="幂等键已绑定其他生成配置,请重新提交")
|
|
|
|
_log_event(
|
|
GenerationRecordEventTypeEnum.PROMPT_OPTIMIZE_IDEMPOTENCY_HIT,
|
|
status=LogEventStatusEnum.SKIPPED,
|
|
user_id=user_id,
|
|
project_id=project_id,
|
|
record_id=record_id,
|
|
detail={"record_status": status, "idempotency_key": req.idempotency_key},
|
|
)
|
|
if status == GenerationStatus.settlement_pending.value:
|
|
await db.rollback()
|
|
return await _settle_staged_result(
|
|
db,
|
|
record_id=record_id,
|
|
user_id=user_id,
|
|
project_name=project_name,
|
|
request_id=req.idempotency_key,
|
|
)
|
|
if status == GenerationStatus.optimizing.value:
|
|
await db.rollback()
|
|
raise HTTPException(status_code=409, detail="相同幂等请求正在处理,请勿重复提交")
|
|
if status == GenerationStatus.failed.value:
|
|
await db.rollback()
|
|
raise HTTPException(status_code=409, detail="该幂等请求已失败,请使用新的幂等键重新提交")
|
|
if record.optimized_prompt and is_generation_record_config_complete(record):
|
|
await db.rollback()
|
|
return PromptOptimizeServiceResult(record_id=record_id, idempotent=True)
|
|
await db.rollback()
|
|
raise HTTPException(status_code=409, detail="幂等记录配置不完整,需人工排查或使用新的幂等键")
|
|
|
|
|
|
async def optimize_generation_prompt(
|
|
db: AsyncSession,
|
|
*,
|
|
req: OptimizeParams,
|
|
user_id: str,
|
|
) -> PromptOptimizeServiceResult:
|
|
project_result = await db.execute(
|
|
select(Project).where(
|
|
Project.id == req.project_id,
|
|
Project.user_id == user_id,
|
|
Project.deleted_at.is_(None),
|
|
).limit(1)
|
|
)
|
|
project = project_result.scalar_one_or_none()
|
|
if project is None:
|
|
raise HTTPException(status_code=404, detail="项目不存在")
|
|
project_name = str(project.name)
|
|
project_industry = str(project.industry or "")
|
|
project_id = str(project.id)
|
|
|
|
existing = await _find_idempotency_record(
|
|
db,
|
|
user_id=user_id,
|
|
idempotency_key=req.idempotency_key,
|
|
)
|
|
if existing is not None:
|
|
return await _handle_existing_record(
|
|
db,
|
|
record=existing,
|
|
req=req,
|
|
user_id=user_id,
|
|
project_name=project_name,
|
|
)
|
|
|
|
if req.gen_type == GenerationType.video:
|
|
if req.duration not in DURATIONS:
|
|
raise HTTPException(status_code=400, detail=f"视频时长必须为{DURATIONS}秒之一")
|
|
if req.aspect_ratio not in ASPECT_RATIOS:
|
|
raise HTTPException(status_code=400, detail="不支持的画面比例")
|
|
if req.resolution not in RESOLUTIONS:
|
|
raise HTTPException(status_code=400, detail="不支持的分辨率")
|
|
engine = await get_video_engine(db, req.engine_id)
|
|
_validate_video_engine_selection(
|
|
engine,
|
|
aspect_ratio=str(req.aspect_ratio),
|
|
resolution=str(req.resolution),
|
|
duration=int(req.duration),
|
|
)
|
|
else:
|
|
if req.image_size not in IMAGE_SIZES:
|
|
raise HTTPException(status_code=400, detail=f"图片分辨率必须为{IMAGE_SIZES}之一")
|
|
if not req.image_proportion or not req.image_px:
|
|
raise HTTPException(status_code=400, detail="图片生成需要指定比例和像素尺寸")
|
|
engine = await get_image_engine(db, req.engine_id)
|
|
_validate_image_engine_selection(engine, image_size=str(req.image_size))
|
|
|
|
reference_usage = calculate_media_reference_usage(
|
|
json.dumps(req.references, ensure_ascii=False) if req.references else None,
|
|
include=bool(req.include_media_references),
|
|
)
|
|
validate_media_reference_usage_for_engine(
|
|
reference_usage,
|
|
gen_type=req.gen_type.value,
|
|
engine=engine,
|
|
)
|
|
|
|
record_id = generate_id()
|
|
record = GenerationRecord(
|
|
id=record_id,
|
|
user_id=user_id,
|
|
project_id=project_id,
|
|
original_prompt=req.prompt,
|
|
optimized_prompt=None,
|
|
prompt_usage_snapshot_json=None,
|
|
gen_type=req.gen_type.value,
|
|
duration=req.duration if req.gen_type == GenerationType.video else None,
|
|
aspect_ratio=req.aspect_ratio if req.gen_type == GenerationType.video else None,
|
|
resolution=req.resolution if req.gen_type == GenerationType.video else None,
|
|
image_size=req.image_size if req.gen_type == GenerationType.image else None,
|
|
image_proportion=req.image_proportion if req.gen_type == GenerationType.image else None,
|
|
image_px=req.image_px if req.gen_type == GenerationType.image else None,
|
|
status=GenerationStatus.optimizing.value,
|
|
pipeline_stage=None,
|
|
credits_cost=0,
|
|
text_credits_cost=0,
|
|
text_tokens_used=0,
|
|
media_references=json.dumps(req.references, ensure_ascii=False) if req.references else None,
|
|
include_media_references=bool(req.include_media_references),
|
|
idempotency_key=req.idempotency_key,
|
|
engine_id=req.engine_id,
|
|
)
|
|
if req.gen_type == GenerationType.video:
|
|
provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot(
|
|
db,
|
|
target_resolution=str(req.resolution),
|
|
aspect_ratio=str(req.aspect_ratio),
|
|
supported_provider_resolutions=parse_json_list(engine.supported_resolutions, []),
|
|
)
|
|
record.provider_generation_resolution = provider_resolution
|
|
record.video_upscale_enabled_snapshot = upscale_enabled
|
|
record.video_upscale_snapshot_json = upscale_snapshot_json
|
|
else:
|
|
record.provider_generation_resolution = None
|
|
record.video_upscale_enabled_snapshot = False
|
|
record.video_upscale_snapshot_json = None
|
|
freeze_generation_record_config_with_log(
|
|
record,
|
|
engine=engine,
|
|
source=GenerationRecordConfigSourceEnum.PROMPT_OPTIMIZE,
|
|
)
|
|
db.add(record)
|
|
try:
|
|
await db.flush()
|
|
except IntegrityError:
|
|
await db.rollback()
|
|
conflicting = await _find_idempotency_record(
|
|
db,
|
|
user_id=user_id,
|
|
idempotency_key=req.idempotency_key,
|
|
)
|
|
if conflicting is None:
|
|
raise
|
|
return await _handle_existing_record(
|
|
db,
|
|
record=conflicting,
|
|
req=req,
|
|
user_id=user_id,
|
|
project_name=project_name,
|
|
)
|
|
|
|
ctx = _billing_context(user_id=user_id, record_id=record_id, request_id=req.idempotency_key)
|
|
try:
|
|
await start_hold(db, ctx)
|
|
await db.commit()
|
|
except Exception:
|
|
await db.rollback()
|
|
raise
|
|
|
|
_log_event(
|
|
GenerationRecordEventTypeEnum.PROMPT_OPTIMIZE_PLACEHOLDER_CREATED,
|
|
user_id=user_id,
|
|
project_id=project_id,
|
|
record_id=record_id,
|
|
detail={
|
|
"idempotency_key": req.idempotency_key,
|
|
"gen_type": req.gen_type.value,
|
|
"engine_id": req.engine_id,
|
|
"attempt_no": _PROMPT_ATTEMPT_NO,
|
|
"reference_count": len(req.references or []),
|
|
},
|
|
)
|
|
|
|
log_provider_start(ctx, detail={"gen_type": req.gen_type.value})
|
|
try:
|
|
optimized_prompt, usage = await optimize_prompt(
|
|
db,
|
|
req.prompt,
|
|
user_id=user_id,
|
|
industry_key=project_industry,
|
|
duration=req.duration if req.gen_type == GenerationType.video else None,
|
|
image_size=req.image_size if req.gen_type == GenerationType.image else None,
|
|
image_proportion=req.image_proportion if req.gen_type == GenerationType.image else None,
|
|
image_px=req.image_px if req.gen_type == GenerationType.image else None,
|
|
references=req.references,
|
|
gen_type=req.gen_type.value,
|
|
log_module="generation_record",
|
|
log_step="prompt_optimize",
|
|
log_project_id=project_id,
|
|
log_owner_type=OWNER_GENERATION_RECORD,
|
|
log_owner_id=record_id,
|
|
generation_attempt_no=_PROMPT_ATTEMPT_NO,
|
|
)
|
|
log_provider_success(ctx, usage=usage)
|
|
except Exception as exc:
|
|
await db.rollback()
|
|
log_provider_failure(ctx, error=str(exc))
|
|
compensated = False
|
|
try:
|
|
failed_result = await db.execute(
|
|
select(GenerationRecord)
|
|
.where(GenerationRecord.id == record_id, GenerationRecord.deleted_at.is_(None))
|
|
.with_for_update()
|
|
.limit(1)
|
|
)
|
|
failed_record = failed_result.scalar_one_or_none()
|
|
if failed_record is not None and failed_record.status == GenerationStatus.optimizing.value:
|
|
failed_record.status = GenerationStatus.failed.value
|
|
failed_record.error_message = extract_error_message(exc, "提示词")
|
|
await release_on_failure(db, ctx, error=str(exc))
|
|
compensated = True
|
|
await db.commit()
|
|
except Exception:
|
|
await db.rollback()
|
|
logger.exception("prompt optimize failure compensation failed: record_id=%s", record_id)
|
|
raise
|
|
_log_event(
|
|
GenerationRecordEventTypeEnum.PROMPT_OPTIMIZE_FAILED_RELEASED,
|
|
status=(LogEventStatusEnum.FAILED if compensated else LogEventStatusEnum.SKIPPED),
|
|
user_id=user_id,
|
|
project_id=project_id,
|
|
record_id=record_id,
|
|
detail={
|
|
"attempt_no": _PROMPT_ATTEMPT_NO,
|
|
"error_type": type(exc).__name__,
|
|
"compensated": compensated,
|
|
},
|
|
error=str(exc),
|
|
)
|
|
raise HTTPException(
|
|
status_code=502,
|
|
detail=f"AI模型调用失败: {extract_error_message(exc, '提示词')}",
|
|
) from exc
|
|
|
|
usage_snapshot = dict(usage or {})
|
|
try:
|
|
input_tokens = int(usage_snapshot.get("input_tokens", 0) or 0)
|
|
output_tokens = int(usage_snapshot.get("output_tokens", 0) or 0)
|
|
except (TypeError, ValueError):
|
|
# 保留原始 usage 交给统一账务校验拒绝;这里只避免展示字段写入异常。
|
|
input_tokens = 0
|
|
output_tokens = 0
|
|
if input_tokens >= 0 and output_tokens >= 0:
|
|
reported_total = usage_snapshot.get("total_tokens")
|
|
normalized_total = input_tokens + output_tokens
|
|
if reported_total not in (None, ""):
|
|
try:
|
|
if int(reported_total) != normalized_total:
|
|
usage_snapshot["reported_total_tokens"] = int(reported_total)
|
|
except (TypeError, ValueError):
|
|
usage_snapshot["reported_total_tokens"] = reported_total
|
|
usage_snapshot["total_tokens"] = normalized_total
|
|
usage_snapshot.setdefault("source_module", "generation_record")
|
|
usage_snapshot.setdefault("source_step_code", "prompt_optimize")
|
|
staged = False
|
|
last_stage_error: Exception | None = None
|
|
for _ in range(2):
|
|
try:
|
|
await db.rollback()
|
|
stage_result = await db.execute(
|
|
select(GenerationRecord)
|
|
.where(
|
|
GenerationRecord.id == record_id,
|
|
GenerationRecord.user_id == user_id,
|
|
GenerationRecord.deleted_at.is_(None),
|
|
)
|
|
.with_for_update()
|
|
.limit(1)
|
|
)
|
|
staged_record = stage_result.scalar_one_or_none()
|
|
if staged_record is None:
|
|
raise RuntimeError("prompt optimize owner record missing")
|
|
if staged_record.status in {
|
|
GenerationStatus.prompt_optimized.value,
|
|
GenerationStatus.generating.value,
|
|
GenerationStatus.completed.value,
|
|
}:
|
|
await db.rollback()
|
|
return PromptOptimizeServiceResult(record_id=record_id, idempotent=True)
|
|
staged_record.optimized_prompt = optimized_prompt
|
|
staged_record.prompt_usage_snapshot_json = json.dumps(
|
|
usage_snapshot,
|
|
ensure_ascii=False,
|
|
sort_keys=True,
|
|
default=str,
|
|
)
|
|
staged_record.text_tokens_used = int(usage_snapshot.get("total_tokens", 0) or 0)
|
|
staged_record.status = GenerationStatus.settlement_pending.value
|
|
staged_record.pipeline_stage = None
|
|
staged_record.error_message = None
|
|
await db.commit()
|
|
staged = True
|
|
break
|
|
except Exception as exc:
|
|
last_stage_error = exc
|
|
await db.rollback()
|
|
logger.exception("prompt optimize provider result staging failed: record_id=%s", record_id)
|
|
if not staged:
|
|
log_operation_error(
|
|
domain=_LOG_DOMAIN,
|
|
event_type=GenerationRecordEventTypeEnum.PROMPT_OPTIMIZE_PROVIDER_RESULT_STAGED.value,
|
|
module=_LOG_MODULE,
|
|
source=_LOG_SOURCE,
|
|
user_id=user_id,
|
|
project_id=project_id,
|
|
task_id=record_id,
|
|
detail={"attempt_no": _PROMPT_ATTEMPT_NO, "stage": "provider_result_persistence"},
|
|
exc=last_stage_error or RuntimeError("unknown staging failure"),
|
|
)
|
|
raise HTTPException(status_code=503, detail="提词已生成但本地暂存失败,请联系管理员根据模型日志处理")
|
|
|
|
_log_event(
|
|
GenerationRecordEventTypeEnum.PROMPT_OPTIMIZE_PROVIDER_RESULT_STAGED,
|
|
user_id=user_id,
|
|
project_id=project_id,
|
|
record_id=record_id,
|
|
detail={
|
|
"attempt_no": _PROMPT_ATTEMPT_NO,
|
|
"input_tokens": usage_snapshot.get("input_tokens"),
|
|
"output_tokens": usage_snapshot.get("output_tokens"),
|
|
"total_tokens": usage_snapshot.get("total_tokens"),
|
|
},
|
|
)
|
|
return await _settle_staged_result(
|
|
db,
|
|
record_id=record_id,
|
|
user_id=user_id,
|
|
project_name=project_name,
|
|
request_id=req.idempotency_key,
|
|
)
|