Files
video-gen/video-gen-api/app/services/media_token_usage_snapshot_service.py
T

366 lines
14 KiB
Python

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