Files
video-gen/video-gen-api/app/services/media_token_usage_snapshot_service.py
T
2026-07-11 13:00:06 +08:00

254 lines
8.5 KiB
Python

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,
)