from __future__ import annotations import json from typing import Any, Mapping from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings from app.enums.credit_record import ( CreditRecordAction, CreditRecordChargeKind, CreditRecordOwnerType, CreditRecordSourceModule, ) 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.utils.id_gen import generate_id def _safe_int(value: Any, default: int = 0) -> int: try: if value is None or value == "": return default return int(value) except Exception: return default def _safe_json_dict(value: Any) -> dict[str, Any]: if not value: return {} if isinstance(value, dict): return value try: parsed = json.loads(value) return parsed if isinstance(parsed, dict) else {} except Exception: return {} def _extract_usage(provider_response: Any) -> dict[str, Any]: data = _safe_json_dict(provider_response) usage = data.get("usage") return usage if isinstance(usage, dict) else {} def _normalize_media_tokens( *, gen_type: str | None, provider_response: Any = None, fallback_total: int | None = None, ) -> tuple[int, int, int]: usage = _extract_usage(provider_response) input_tokens = _safe_int(usage.get("input_tokens"), 0) output_tokens = _safe_int( usage.get("output_tokens"), _safe_int(usage.get("generated_tokens"), 0), ) total_tokens = _safe_int(usage.get("total_tokens"), 0) if total_tokens <= 0: total_tokens = _safe_int(fallback_total, 0) if output_tokens <= 0: output_tokens = max(0, total_tokens - input_tokens) if total_tokens <= 0: total_tokens = input_tokens + output_tokens # 图片生成多数供应商只返回 output/total,没有 input;保持 input=0。视频同理兼容缺字段。 return input_tokens, output_tokens, total_tokens def _engine_model_from_provider_response(provider_response: Any) -> str | None: data = _safe_json_dict(provider_response) model = data.get("model") return str(model) if model else None async def _find_latest_media_charge( db: AsyncSession, *, user_id: str, owner_type: str, owner_id: str, 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) query = query.order_by(CreditRecord.attempt_no.desc().nullslast(), CreditRecord.created_at.desc()).limit(1) result = await db.execute(query) return result.scalar_one_or_none() async def _get_or_create_token_usage( db: AsyncSession, *, charge: CreditRecord, input_tokens: int, output_tokens: int, total_tokens: int, model_config_id: str | None = None, ) -> TokenUsage: token_usage: TokenUsage | None = None if charge.token_usage_id: result = await db.execute(select(TokenUsage).where(TokenUsage.id == charge.token_usage_id).limit(1)) token_usage = result.scalar_one_or_none() if token_usage is None and charge.biz_key: result = await db.execute(select(TokenUsage).where(TokenUsage.biz_key == charge.biz_key).limit(1)) token_usage = result.scalar_one_or_none() if token_usage is None: token_usage = TokenUsage( id=generate_id(), user_id=charge.user_id, model_config_id=model_config_id, 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.user_id = token_usage.user_id or charge.user_id token_usage.model_config_id = token_usage.model_config_id or model_config_id token_usage.owner_type = token_usage.owner_type or charge.owner_type token_usage.owner_id = token_usage.owner_id or charge.owner_id token_usage.biz_key = token_usage.biz_key or charge.biz_key token_usage.source_module = token_usage.source_module or charge.source_module token_usage.source_step_code = token_usage.source_step_code or charge.source_step_code 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, gen_type: str | None, provider_response: Any = None, fallback_total: int | None = None, ) -> CreditRecord | None: if not charge: return None input_tokens, output_tokens, total_tokens = _normalize_media_tokens( gen_type=gen_type, provider_response=provider_response, fallback_total=fallback_total, ) if total_tokens <= 0: return charge token_usage = await _get_or_create_token_usage( db, charge=charge, input_tokens=input_tokens, output_tokens=output_tokens, total_tokens=total_tokens, model_config_id=None, ) charge.token_usage_id = token_usage.id charge.input_tokens = input_tokens charge.output_tokens = output_tokens charge.total_tokens = total_tokens # 兼容旧流水扣费时未冷备 engine_model_name 的场景,能从 provider response 推出来就补充。 provider_model = _engine_model_from_provider_response(provider_response) if provider_model and not charge.engine_model_name: charge.engine_model_name = provider_model return charge async def sync_chat_generation_task_media_token_snapshot( db: AsyncSession, task: ChatGenerationTask, *, provider_response: Any = None, ) -> CreditRecord | None: """把 ChatGenerationTask 图片/视频媒体生成 token 后置快照回填到积分流水。 媒体扣费发生在创建任务前,供应商 usage 只能在创建/轮询成功后拿到, 所以这里按 owner_type + owner_id + media_type 找到对应 media charge 流水并回填。 """ if not bool(getattr(settings, "MEDIA_TOKEN_SNAPSHOT_ENABLED", True)): return None if not task: return None gen_type = (getattr(task, "gen_type", None) or "").lower().strip() fallback_total = task.image_tokens_used if gen_type == "image" else task.video_tokens_used response = provider_response if provider_response is not None else getattr(task, "provider_response_json", None) charge = await _find_latest_media_charge( db, user_id=task.user_id, owner_type=CreditRecordOwnerType.CHAT_GENERATION_TASK.value, owner_id=task.id, media_type=gen_type or None, ) return await _sync_charge_snapshot( db, charge=charge, gen_type=gen_type, provider_response=response, fallback_total=fallback_total, ) async def sync_generation_record_media_token_snapshot( db: AsyncSession, record: GenerationRecord, *, provider_response: Any = None, ) -> CreditRecord | None: """把旧 GenerationRecord 图片/视频媒体生成 token 后置快照回填到积分流水。""" if not bool(getattr(settings, "MEDIA_TOKEN_SNAPSHOT_ENABLED", True)): return None if not record: return None gen_type = (getattr(record, "gen_type", None) or "").lower().strip() fallback_total = record.image_tokens_used if gen_type == "image" else record.video_tokens_used charge = await _find_latest_media_charge( db, user_id=record.user_id, owner_type=CreditRecordOwnerType.GENERATION_RECORD.value, owner_id=record.id, media_type=gen_type or None, ) return await _sync_charge_snapshot( db, charge=charge, gen_type=gen_type, provider_response=provider_response, fallback_total=fallback_total, )