366 lines
14 KiB
Python
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,
|
|
)
|