259 lines
8.8 KiB
Python
259 lines
8.8 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,
|
|
attempt_no: int | None = 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 == int(attempt_no))
|
|
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,
|
|
attempt_no=int(getattr(task, "generation_attempt_no", 1) or 1),
|
|
)
|
|
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,
|
|
attempt_no=int(getattr(record, "generation_attempt_no", 1) or 1),
|
|
)
|
|
return await _sync_charge_snapshot(
|
|
db,
|
|
charge=charge,
|
|
gen_type=gen_type,
|
|
provider_response=provider_response,
|
|
fallback_total=fallback_total,
|
|
)
|