from __future__ import annotations from copy import deepcopy from datetime import datetime, timezone from typing import Any from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings from app.enums.credit_record import CreditRecordAction, CreditRecordChargeKind, CreditRecordOwnerType from app.enums.model_pricing import PricingSnapshotStage, ProviderCostStatus from app.models.chat_generation_task import ChatGenerationTask from app.models.credit_record import CreditRecord from app.models.generation_record import GenerationRecord from app.models.token_usage import TokenUsage from app.services.model_pricing.attachment_snapshot_service import build_attachment_snapshot, build_generation_snapshot from app.services.model_pricing.usage_normalizer import ( normalize_provider_media_usage, safe_int, safe_json_dict, ) from app.services.model_pricing.snapshot_service import finalize_credit_record_pricing from app.services.operation_log_service import log_model_pricing_event from app.utils.id_gen import generate_id def _provider_model_from_response(provider_response: Any) -> str | None: data = safe_json_dict(provider_response) for candidate in ( data.get("model"), (data.get("data") or {}).get("model") if isinstance(data.get("data"), dict) else None, (data.get("result") or {}).get("model") if isinstance(data.get("result"), dict) else None, ): if candidate: return str(candidate) return None def _provider_task_id_from_response(provider_response: Any) -> str | None: data = safe_json_dict(provider_response) for candidate in (data.get("task_id"), data.get("id"), data.get("provider_task_id")): if candidate: return str(candidate) return None def _engine_snapshot_from_owner(owner: Any) -> dict[str, Any]: snapshot = safe_json_dict(getattr(owner, "engine_snapshot_json", None)) return { "provider": snapshot.get("provider") or snapshot.get("engine_provider"), "model_name": snapshot.get("model_name") or snapshot.get("engine_model_name"), "engine_name": snapshot.get("name") or snapshot.get("engine_name"), "engine_id": snapshot.get("id") or snapshot.get("engine_id") or getattr(owner, "engine_id", None), } async def _find_media_charge( db: AsyncSession, *, user_id: str, owner_type: str, owner_id: str, attempt_no: int | None, media_type: str | None, ) -> CreditRecord | None: query = ( select(CreditRecord) .where(CreditRecord.user_id == user_id) .where(CreditRecord.type == "consume") .where(CreditRecord.owner_type == owner_type) .where(CreditRecord.owner_id == owner_id) .where(CreditRecord.charge_kind == CreditRecordChargeKind.MEDIA.value) .where(CreditRecord.charge_action == CreditRecordAction.CHARGE.value) ) if media_type: query = query.where(CreditRecord.media_type == media_type) if attempt_no is not None: query = query.where(CreditRecord.attempt_no == attempt_no) query = query.order_by(CreditRecord.created_at.desc()).limit(1).with_for_update() return (await db.execute(query)).scalar_one_or_none() # 兼容迁移上线时仍在执行、尚未写 current_billing_attempt_no 的旧任务: # 只有候选消费流水唯一时才允许绑定;多次重试产生多条流水时宁可跳过,也不能猜“最新一条”。 candidates = ( await db.execute(query.order_by(CreditRecord.created_at.desc()).limit(2).with_for_update()) ).scalars().all() return candidates[0] if len(candidates) == 1 else None async def _get_or_create_token_usage( db: AsyncSession, *, charge: CreditRecord, input_tokens: int, output_tokens: int, total_tokens: int, ) -> TokenUsage: token_usage: TokenUsage | None = None if charge.token_usage_id: token_usage = ( await db.execute(select(TokenUsage).where(TokenUsage.id == charge.token_usage_id).limit(1)) ).scalar_one_or_none() if token_usage is None and charge.biz_key: token_usage = ( await db.execute(select(TokenUsage).where(TokenUsage.biz_key == charge.biz_key).limit(1)) ).scalar_one_or_none() if token_usage is None: token_usage = TokenUsage( id=generate_id(), user_id=charge.user_id, model_config_id=None, owner_type=charge.owner_type, owner_id=charge.owner_id, biz_key=charge.biz_key, source_module=charge.source_module, source_step_code=charge.source_step_code, input_tokens=input_tokens, output_tokens=output_tokens, total_tokens=total_tokens, ) db.add(token_usage) await db.flush() else: token_usage.input_tokens = input_tokens token_usage.output_tokens = output_tokens token_usage.total_tokens = total_tokens return token_usage async def _sync_charge_snapshot( db: AsyncSession, *, charge: CreditRecord | None, owner: Any, gen_type: str, stage: str, provider_response: Any = None, fallback_total: int = 0, ) -> CreditRecord | None: if not charge: log_model_pricing_event( event_type="pricing_snapshot_skip", event_status="warning", owner_type=owner.__class__.__name__, owner_id=getattr(owner, "id", None), message="未找到与 current_billing_attempt_no 匹配的媒体消费流水", detail={"attempt_no": getattr(owner, "current_billing_attempt_no", None), "stage": stage}, ) return None response = provider_response if provider_response is not None else getattr(owner, "provider_response_json", None) attachment_snapshot, attachment_counts = build_attachment_snapshot(getattr(owner, "media_references", None)) generation_snapshot, generation_counts, generation_usage = build_generation_snapshot( owner, provider_response=response, stage=stage, ) existing_usage = deepcopy(dict(charge.usage_snapshot_json or {})) provider_usage = normalize_provider_media_usage( response, gen_type=gen_type, fallback_total_tokens=fallback_total, request_image_px=getattr(owner, "image_px", None), requested_output_count=max(1, safe_int(generation_counts.get("requested_output_count"), 1)), provider_input_image_count=safe_int(attachment_counts.get("provider_input_image_count")), ) usage = {**existing_usage, **generation_usage, **provider_usage} usage.update( { "has_input_video": safe_int(attachment_counts.get("provider_input_video_count")) > 0, "provider_input_image_count": safe_int( provider_usage.get("provider_input_image_count"), safe_int(attachment_counts.get("provider_input_image_count")), ), "input_image_count": safe_int( provider_usage.get("provider_input_image_count"), safe_int(attachment_counts.get("provider_input_image_count")), ), "input_video_duration_seconds": float(attachment_counts.get("attachment_video_duration_seconds") or 0), "input_audio_duration_seconds": float(attachment_counts.get("attachment_audio_duration_seconds") or 0), "usage_stage": stage, } ) input_tokens = max(0, safe_int(usage.get("input_tokens"))) output_tokens = max(0, safe_int(usage.get("output_tokens"))) total_tokens = max(0, safe_int(usage.get("total_tokens"), input_tokens + output_tokens)) if total_tokens > 0: token_usage = await _get_or_create_token_usage( db, charge=charge, input_tokens=input_tokens, output_tokens=output_tokens, total_tokens=total_tokens, ) charge.token_usage_id = token_usage.id charge.input_tokens = input_tokens charge.output_tokens = output_tokens charge.total_tokens = total_tokens provider_model = _provider_model_from_response(response) engine_snapshot = _engine_snapshot_from_owner(owner) charge.engine_provider = charge.engine_provider or engine_snapshot.get("provider") charge.engine_model_name = charge.engine_model_name or provider_model or engine_snapshot.get("model_name") charge.engine_name = charge.engine_name or engine_snapshot.get("engine_name") charge.engine_id = charge.engine_id or engine_snapshot.get("engine_id") await finalize_credit_record_pricing( db, charge=charge, usage=usage, stage=stage, attachment_snapshot=attachment_snapshot, attachment_counts=attachment_counts, generation_snapshot=generation_snapshot, generation_counts=generation_counts, allow_upgrade_estimated=True, ) return charge async def mark_media_provider_cost_status( db: AsyncSession, *, owner: ChatGenerationTask | GenerationRecord, status: str, reason: str, usage_stage: str, ) -> CreditRecord | None: """Finalize a media charge when a synchronous provider call failed or became uncertain. This helper never commits and always resolves the charge by owner + billing attempt. """ if isinstance(owner, ChatGenerationTask): owner_type = CreditRecordOwnerType.CHAT_GENERATION_TASK.value else: owner_type = CreditRecordOwnerType.GENERATION_RECORD.value gen_type = str(getattr(owner, "gen_type", None) or "").lower().strip() charge = await _find_media_charge( db, user_id=owner.user_id, owner_type=owner_type, owner_id=owner.id, attempt_no=getattr(owner, "current_billing_attempt_no", None), media_type=gen_type or None, ) if not charge: log_model_pricing_event( event_type="pricing_snapshot_skip", event_status="warning", owner_type=owner_type, owner_id=owner.id, message="同步图片异常时未找到唯一媒体消费流水", detail={ "attempt_no": getattr(owner, "current_billing_attempt_no", None), "target_status": status, "reason": reason, }, ) return None if charge.provider_cost_status in {ProviderCostStatus.CALCULATED.value, ProviderCostStatus.ESTIMATED.value}: return charge usage_snapshot = deepcopy(dict(charge.usage_snapshot_json or {})) usage_snapshot.update( { "usage_stage": usage_stage, "provider_error_reason": reason, "provider_result_uncertain": status == ProviderCostStatus.PROVIDER_RESULT_UNCERTAIN.value, } ) charge.usage_snapshot_json = usage_snapshot charge.provider_cost_status = status charge.provider_cost_amount = None if status == ProviderCostStatus.PROVIDER_RESULT_UNCERTAIN.value else 0 charge.provider_cost_is_estimated = False charge.provider_cost_calculated_at = datetime.now(timezone.utc) charge.provider_cost_finalized_at = datetime.now(timezone.utc) return charge async def sync_chat_generation_task_media_token_snapshot( db: AsyncSession, task: ChatGenerationTask, *, provider_response: Any = None, stage: str | None = None, ) -> CreditRecord | None: if not bool(getattr(settings, "MEDIA_TOKEN_SNAPSHOT_ENABLED", True)) or not task: return None gen_type = (task.gen_type or "").lower().strip() stage = stage or ( PricingSnapshotStage.PROVIDER_SYNC_COMPLETED.value if gen_type == "image" else PricingSnapshotStage.PROVIDER_ASYNC_COMPLETED.value ) response = provider_response if provider_response is not None else task.provider_response_json callback_task_id = _provider_task_id_from_response(response) current_task_id = task.seedance_task_id or task.provider_task_id if gen_type == "video" and callback_task_id and current_task_id and callback_task_id != current_task_id: log_model_pricing_event( event_type="pricing_stale_callback_skip", event_status="warning", owner_type=CreditRecordOwnerType.CHAT_GENERATION_TASK.value, owner_id=task.id, message="旧 Provider 回调与当前任务 ID 不一致,已跳过", detail={"callback_task_id": callback_task_id, "current_task_id": current_task_id}, ) return None charge = await _find_media_charge( db, user_id=task.user_id, owner_type=CreditRecordOwnerType.CHAT_GENERATION_TASK.value, owner_id=task.id, attempt_no=task.current_billing_attempt_no, media_type=gen_type or None, ) fallback_total = task.image_tokens_used if gen_type == "image" else task.video_tokens_used return await _sync_charge_snapshot( db, charge=charge, owner=task, gen_type=gen_type, stage=stage, provider_response=response, fallback_total=fallback_total or 0, ) async def sync_generation_record_media_token_snapshot( db: AsyncSession, record: GenerationRecord, *, provider_response: Any = None, stage: str | None = None, ) -> CreditRecord | None: if not bool(getattr(settings, "MEDIA_TOKEN_SNAPSHOT_ENABLED", True)) or not record: return None gen_type = (record.gen_type or "").lower().strip() stage = stage or ( PricingSnapshotStage.PROVIDER_SYNC_COMPLETED.value if gen_type == "image" else PricingSnapshotStage.PROVIDER_ASYNC_COMPLETED.value ) response = provider_response if provider_response is not None else record.provider_response_json charge = await _find_media_charge( db, user_id=record.user_id, owner_type=CreditRecordOwnerType.GENERATION_RECORD.value, owner_id=record.id, attempt_no=record.current_billing_attempt_no, media_type=gen_type or None, ) fallback_total = record.image_tokens_used if gen_type == "image" else record.video_tokens_used return await _sync_charge_snapshot( db, charge=charge, owner=record, gen_type=gen_type, stage=stage, provider_response=response, fallback_total=fallback_total or 0, )