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.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, record_provider_exception, log_provider_start, log_provider_success, finalize_llm_business_failure, mark_business_success, charge_llm_credits, ) 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, 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 mark_business_success( db, ctx, usage=usage, description=f"提示词优化 - {project_name}", ) charge_item = next( (item for item in billing.items if item.biz_key == ctx.billing_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 charge_llm_credits(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 []), }, ) await log_provider_start(db, 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, fixed_model_config_id=ctx.model_config_id, fixed_model_snapshot=ctx.model_parameters_snapshot, ) await log_provider_success(db, ctx, usage=usage) except Exception as exc: await db.rollback() provider_succeeded, recovered_usage = await record_provider_exception(db, ctx, exc) if recovered_usage: usage = recovered_usage 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() business_already_succeeded = bool( failed_record is not None and failed_record.status in { GenerationStatus.prompt_optimized.value, GenerationStatus.generating.value, GenerationStatus.completed.value, } ) 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 db.commit() if not business_already_succeeded: await finalize_llm_business_failure(ctx, error=str(exc)) compensated = True 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: failure = last_stage_error or RuntimeError("unknown staging failure") 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=failure, ) # 供应商调用和 Token 已由独立调用审计事务保存。结果连续暂存失败后, # 当前 API 已没有可继续恢复的本地业务结果,必须先结束业务状态,再按 # 原积分来源退回场景积分消费,不能让账务永久停留在 processing。 business_already_succeeded = False try: await db.rollback() failed_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) ) failed_record = failed_result.scalar_one_or_none() business_already_succeeded = bool( failed_record is not None and failed_record.status in { GenerationStatus.prompt_optimized.value, GenerationStatus.generating.value, GenerationStatus.completed.value, } ) if failed_record is not None and not business_already_succeeded: failed_record.status = GenerationStatus.failed.value failed_record.pipeline_stage = None failed_record.error_message = "提词已生成但本地结果暂存失败" await db.commit() except Exception: await db.rollback() logger.exception("prompt optimize staging final failure persistence failed: record_id=%s", record_id) if business_already_succeeded: return PromptOptimizeServiceResult(record_id=record_id, idempotent=True) await finalize_llm_business_failure(ctx, error=str(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, )