from __future__ import annotations from dataclasses import dataclass from typing import Any, Iterable, Mapping from sqlalchemy import select from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession from app.enums.credit_record import ( CreditRecordAction, CreditRecordBillingScene, CreditRecordChargeKind, CreditRecordOwnerType, CreditRecordSourceModule, CreditRecordSubject, ) from app.enums.llm_billing import LlmBillingDomain, LlmBillingEvent, LlmBillingLedgerState from app.models.credit_record import CreditRecord from app.models.module_generation_step import ModuleGenerationStep from app.models.token_usage import TokenUsage from app.services.credit_record_meta_service import ( CreditRecordMeta, build_generation_record_prompt_meta, build_module_step_prompt_meta, build_shot_video_analysis_meta, ) from app.services.credits import add_credits_result, calc_text_credits, deduct_credits_result from app.services.generation.billing_service import BillingItem, BillingSummary from app.services.llm_billing.config import get_llm_billing_policy from app.services.llm_billing.context import ( LlmBillingConfigurationError, LlmBillingContext, LlmBillingStateError, LlmHoldResult, LlmHoldValidation, ) from app.services.operation_log_service import log_operation_event from app.utils.exceptions import InsufficientCreditsError from app.utils.id_gen import generate_id _LEDGER_QUERY_BATCH_SIZE = 1000 @dataclass(slots=True) class _LedgerRecords: state: LlmBillingLedgerState hold: CreditRecord | None = None release: CreditRecord | None = None charge: CreditRecord | None = None reason: str | None = None @property def hold_amount(self) -> float: return _round2(abs(float(self.hold.amount or 0))) if self.hold else 0.0 def _round2(value: Any) -> float: try: return round(float(value or 0), 2) except (TypeError, ValueError): return 0.0 def _safe_int(value: Any, default: int = 0) -> int: try: if value is None or value == "": return default return int(value) except (TypeError, ValueError): return default def _context_detail(ctx: LlmBillingContext, **extra: Any) -> dict[str, Any]: detail = { "user_id": ctx.user_id, "owner_type": ctx.owner_type, "owner_id": ctx.owner_id, "attempt_no": ctx.attempt_no, "charge_kind": ctx.charge_kind, "billing_scene": ctx.billing_scene, "source_module": ctx.source_module, "source_project_id": ctx.source_project_id, "source_step_id": ctx.source_step_id, "source_step_code": ctx.source_step_code, "related_id": ctx.related_id, "hold_biz_key": ctx.hold_biz_key, "hold_release_biz_key": ctx.hold_release_biz_key, "charge_biz_key": ctx.charge_biz_key, "celery_task_id": ctx.celery_task_id, "provider": ctx.provider, "model_name": ctx.model_name, "token_usage_id": ctx.token_usage_id, } detail.update({key: value for key, value in extra.items() if value is not None}) return {key: value for key, value in detail.items() if value is not None} def _log( ctx: LlmBillingContext, event: LlmBillingEvent, *, status: str = "success", message: str | None = None, detail: dict[str, Any] | None = None, error: str | None = None, ) -> None: log_operation_event( domain=LlmBillingDomain.LLM_BILLING.value, module=ctx.source_module or LlmBillingDomain.LLM_BILLING.value, event_type=event.value, event_status=status, source="app.services.llm_billing.service", trace_id=ctx.trace_id, request_id=ctx.request_id, user_id=ctx.user_id, project_id=ctx.source_project_id, task_id=ctx.owner_id, step_id=ctx.source_step_id, message=message, detail=detail or _context_detail(ctx), error=error, ) def log_provider_start(ctx: LlmBillingContext, *, detail: Mapping[str, Any] | None = None) -> None: _log( ctx, LlmBillingEvent.PROVIDER_START, status="started", detail=_context_detail(ctx, **dict(detail or {})), ) def log_provider_success(ctx: LlmBillingContext, *, usage: Mapping[str, Any] | None = None) -> None: usage_snapshot = dict(usage or {}) ctx.provider = str(usage_snapshot.get("provider") or usage_snapshot.get("model_provider") or "") or ctx.provider ctx.model_name = str(usage_snapshot.get("model_name") or usage_snapshot.get("model") or "") or ctx.model_name ctx.token_usage_id = str(usage_snapshot.get("token_usage_id") or "") or ctx.token_usage_id _log( ctx, LlmBillingEvent.PROVIDER_SUCCESS, detail=_context_detail( ctx, input_tokens=_safe_int(usage_snapshot.get("input_tokens")), output_tokens=_safe_int(usage_snapshot.get("output_tokens")), total_tokens=_safe_int(usage_snapshot.get("total_tokens")), ), ) def log_provider_failure(ctx: LlmBillingContext, *, error: str) -> None: _log( ctx, LlmBillingEvent.PROVIDER_FAILURE, status="failed", detail=_context_detail(ctx, error_type="provider_call_failed"), error=error, ) def log_celery_dispatch_start(ctx: LlmBillingContext) -> None: _log(ctx, LlmBillingEvent.CELERY_DISPATCH_START, status="started") def log_celery_dispatch_success(ctx: LlmBillingContext) -> None: _log(ctx, LlmBillingEvent.CELERY_DISPATCH_SUCCESS) def log_celery_dispatch_failure(ctx: LlmBillingContext, *, error: str) -> None: _log(ctx, LlmBillingEvent.CELERY_DISPATCH_FAILURE, status="failed", error=error) def log_celery_dispatch_compensated(ctx: LlmBillingContext, *, error: str) -> None: _log(ctx, LlmBillingEvent.CELERY_DISPATCH_COMPENSATED, detail=_context_detail(ctx, compensation_error=error)) def _hold_meta(ctx: LlmBillingContext, *, action: str) -> CreditRecordMeta: subject = ( CreditRecordSubject.ANALYSIS.value if ctx.charge_kind == CreditRecordChargeKind.VIDEO_ANALYSIS.value else CreditRecordSubject.TEXT.value ) return CreditRecordMeta( owner_type=ctx.owner_type, owner_id=ctx.owner_id, attempt_no=ctx.attempt_no, charge_kind=ctx.charge_kind, charge_action=action, credit_subject=subject, media_type="video" if ctx.charge_kind == CreditRecordChargeKind.VIDEO_ANALYSIS.value else None, billing_scene=ctx.billing_scene, source_module=ctx.source_module, source_project_id=ctx.source_project_id, source_step_id=ctx.source_step_id, source_step_code=ctx.source_step_code, ) def _action_valid(record: CreditRecord | None, expected: CreditRecordAction) -> bool: if record is None: return True # 兼容旧数据:正式 biz_key 已明确动作、charge_action 为空时仍可识别。 return record.charge_action in (None, expected.value) def _record_meta_matches_context(record: CreditRecord | None, ctx: LlmBillingContext) -> bool: if record is None: return True checks = ( (record.owner_type, ctx.owner_type), (record.owner_id, ctx.owner_id), (record.attempt_no, ctx.attempt_no), (record.charge_kind, ctx.charge_kind), ) return all(actual is None or str(actual) == str(expected) for actual, expected in checks) def _classify_ledger( *, ctx: LlmBillingContext, hold: CreditRecord | None, release: CreditRecord | None, charge: CreditRecord | None, ) -> _LedgerRecords: if not _action_valid(hold, CreditRecordAction.HOLD): return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "hold_action_mismatch") if not _action_valid(release, CreditRecordAction.HOLD_RELEASE): return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "release_action_mismatch") if not _action_valid(charge, CreditRecordAction.CHARGE): return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "charge_action_mismatch") if not _record_meta_matches_context(hold, ctx): return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "hold_context_mismatch") if not _record_meta_matches_context(release, ctx): return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "release_context_mismatch") if not _record_meta_matches_context(charge, ctx): return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "charge_context_mismatch") if hold is None: if release is not None or charge is not None: return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "hold_missing_with_followup") return _LedgerRecords(LlmBillingLedgerState.MISSING) if release is None and charge is None: if hold.type != "consume" or _round2(float(hold.amount or 0)) >= 0: return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "hold_type_or_sign_invalid") if _round2(abs(float(hold.amount or 0))) <= 0: return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "hold_amount_not_positive") return _LedgerRecords(LlmBillingLedgerState.ACTIVE, hold) if release is not None and charge is None: if release.type != "refund" or _round2(float(release.amount or 0)) <= 0: return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "release_type_or_sign_invalid") if str(release.refund_for_biz_key or "") != str(ctx.hold_biz_key): return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "release_target_mismatch") if _round2(release.amount) != _round2(abs(float(hold.amount or 0))): return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "release_amount_mismatch") return _LedgerRecords(LlmBillingLedgerState.RELEASED, hold, release) if release is not None and charge is not None: if release.type != "refund" or _round2(float(release.amount or 0)) <= 0: return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "release_type_or_sign_invalid") if str(release.refund_for_biz_key or "") != str(ctx.hold_biz_key): return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "release_target_mismatch") if _round2(release.amount) != _round2(abs(float(hold.amount or 0))): return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "release_amount_mismatch") if charge.type != "consume" or _round2(float(charge.amount or 0)) > 0: return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "charge_type_or_sign_invalid") return _LedgerRecords(LlmBillingLedgerState.CHARGED, hold, release, charge) return _LedgerRecords(LlmBillingLedgerState.INVALID, hold, release, charge, "charge_without_release") async def _load_ledgers( db: AsyncSession, contexts: Iterable[LlmBillingContext], ) -> dict[str, _LedgerRecords]: context_list = list(contexts) if not context_list: return {} record_map: dict[tuple[str, str], CreditRecord] = {} all_keys = list(dict.fromkeys(key for ctx in context_list for key in ctx.ledger_biz_keys)) user_ids = list(dict.fromkeys(ctx.user_id for ctx in context_list)) for offset in range(0, len(all_keys), _LEDGER_QUERY_BATCH_SIZE): chunk = all_keys[offset : offset + _LEDGER_QUERY_BATCH_SIZE] result = await db.execute( select(CreditRecord).where( CreditRecord.user_id.in_(user_ids), CreditRecord.biz_key.in_(chunk), ) ) for record in result.scalars().all(): if record.biz_key: record_map[(str(record.user_id), str(record.biz_key))] = record output: dict[str, _LedgerRecords] = {} for ctx in context_list: hold = record_map.get((ctx.user_id, ctx.hold_biz_key)) release = record_map.get((ctx.user_id, ctx.hold_release_biz_key)) charge = record_map.get((ctx.user_id, ctx.charge_biz_key)) output[ctx.hold_biz_key] = _classify_ledger(ctx=ctx, hold=hold, release=release, charge=charge) return output async def _load_ledger(db: AsyncSession, ctx: LlmBillingContext) -> _LedgerRecords: return (await _load_ledgers(db, [ctx]))[ctx.hold_biz_key] async def get_llm_ledger_states( db: AsyncSession, contexts: Iterable[LlmBillingContext], ) -> dict[str, LlmHoldValidation]: """批量读取 attempt 的三类流水;供恢复任务收集 ID 后统一过滤。""" context_list = list(contexts) ledgers = await _load_ledgers(db, context_list) return { ctx.hold_biz_key: LlmHoldValidation( can_execute=ledgers[ctx.hold_biz_key].state == LlmBillingLedgerState.ACTIVE, amount=ledgers[ctx.hold_biz_key].hold_amount, state=ledgers[ctx.hold_biz_key].state, reason=ledgers[ctx.hold_biz_key].reason, hold_record_id=ledgers[ctx.hold_biz_key].hold.id if ledgers[ctx.hold_biz_key].hold else None, ) for ctx in context_list } async def validate_retryable_previous_attempt( db: AsyncSession, ctx: LlmBillingContext, ) -> LlmHoldValidation: """校验失败 attempt 的账务是否已关闭,供业务创建下一 attempt 前调用。 只有已释放 HOLD,或当前计费明确关闭且旧 attempt 没有流水时,才允许创建 新 attempt。这里不主动退款,避免重试入口承担失败补偿职责。 """ _log( ctx, LlmBillingEvent.RETRY_PREVIOUS_ATTEMPT_VALIDATE_START, status="started", ) ledger = await _load_ledger(db, ctx) if ledger.state == LlmBillingLedgerState.RELEASED: result = LlmHoldValidation( True, ledger.hold_amount, ledger.state, hold_record_id=ledger.hold.id if ledger.hold else None, ) _log( ctx, LlmBillingEvent.RETRY_PREVIOUS_ATTEMPT_VALIDATE_SUCCESS, detail=_context_detail( ctx, ledger_state=result.state.value, hold_credits=result.amount, hold_record_id=result.hold_record_id, ), ) return result if ledger.state == LlmBillingLedgerState.MISSING: policy = await get_llm_billing_policy( db, config_key=ctx.hold_config_key, explicit_hold_credits=None, ) if policy.bypassed: result = LlmHoldValidation( True, 0.0, LlmBillingLedgerState.BILLING_BYPASSED, "billing_disabled", ) _log( ctx, LlmBillingEvent.RETRY_PREVIOUS_ATTEMPT_VALIDATE_SUCCESS, status="skipped", detail=_context_detail( ctx, ledger_state=result.state.value, billing_bypassed=True, ), ) return result if not policy.valid: result = LlmHoldValidation( False, 0.0, LlmBillingLedgerState.INVALID, policy.error or "billing_config_invalid", ) else: result = LlmHoldValidation( False, 0.0, ledger.state, "previous_attempt_hold_missing", ) else: result = LlmHoldValidation( False, ledger.hold_amount, ledger.state, ledger.reason or f"previous_attempt_ledger_{ledger.state.value}", ledger.hold.id if ledger.hold else None, ) _log( ctx, LlmBillingEvent.RETRY_PREVIOUS_ATTEMPT_BLOCKED, status="failed", detail=_context_detail( ctx, ledger_state=result.state.value, hold_credits=result.amount, hold_record_id=result.hold_record_id, skip_reason=result.reason, ), error="旧attempt账务尚未关闭,拒绝创建新的分析attempt", ) return result async def start_hold(db: AsyncSession, ctx: LlmBillingContext) -> LlmHoldResult: # 幂等/异常 attempt 优先由已落库流水判定;只有全新 attempt 才读取配置。 ledger = await _load_ledger(db, ctx) # 配置可能在任务执行期间被关闭或修改:已经存在的 active HOLD 必须继续沿用, # 否则会留下永久冻结流水。只有“没有任何历史流水”的新 attempt 才允许按关闭配置绕过。 if ledger.state == LlmBillingLedgerState.ACTIVE and ledger.hold: amount = ledger.hold_amount ctx.hold_credits = amount _log( ctx, LlmBillingEvent.HOLD_SUCCESS, detail=_context_detail( ctx, hold_credits=amount, hold_record_id=ledger.hold.id, idempotent=True, ledger_state=ledger.state.value, ), ) return LlmHoldResult(amount, ledger.state, created=False, record_id=ledger.hold.id) if ledger.state != LlmBillingLedgerState.MISSING: error = f"当前attempt账务状态为{ledger.state.value},不能复用旧预扣" _log( ctx, LlmBillingEvent.EXECUTION_BLOCKED, status="failed", detail=_context_detail(ctx, ledger_state=ledger.state.value, reason=ledger.reason), error=error, ) raise LlmBillingStateError(error) policy = await get_llm_billing_policy( db, config_key=ctx.hold_config_key, explicit_hold_credits=ctx.hold_credits, ) if policy.bypassed: _log( ctx, LlmBillingEvent.HOLD_BYPASSED, status="skipped", detail=_context_detail(ctx, ledger_state=LlmBillingLedgerState.BILLING_BYPASSED.value), ) return LlmHoldResult( 0.0, LlmBillingLedgerState.BILLING_BYPASSED, reason="billing_disabled", ) if not policy.valid: _log( ctx, LlmBillingEvent.HOLD_CONFIG_INVALID, status="failed", detail=_context_detail( ctx, hold_credits=policy.hold_credits, config_key=policy.config_key, config_source=policy.source_key, ), error=policy.error, ) raise LlmBillingConfigurationError(policy.error or "LLM计费配置无效") amount = policy.hold_credits ctx.hold_credits = amount _log( ctx, LlmBillingEvent.HOLD_START, status="started", detail=_context_detail( ctx, hold_credits=amount, config_key=policy.config_key, config_source=policy.source_key, ), ) try: mutation = await deduct_credits_result( db, user_id=ctx.user_id, amount=amount, description=f"{ctx.description_prefix}预扣积分", related_id=ctx.related_id or ctx.owner_id, biz_key=ctx.hold_biz_key, record_meta=_hold_meta(ctx, action=CreditRecordAction.HOLD.value), allow_negative=False, ) except InsufficientCreditsError: _log( ctx, LlmBillingEvent.HOLD_INSUFFICIENT, status="failed", detail=_context_detail(ctx, hold_credits=amount), error="积分不足,无法预扣", ) raise if not mutation.created: # 并发幂等命中后重新读取三类流水,避免复用已被另一事务释放的 HOLD。 ledger = await _load_ledger(db, ctx) if ledger.state != LlmBillingLedgerState.ACTIVE or ledger.hold is None: error = f"并发预扣后账务状态为{ledger.state.value},拒绝继续执行" _log( ctx, LlmBillingEvent.EXECUTION_BLOCKED, status="failed", detail=_context_detail(ctx, ledger_state=ledger.state.value), error=error, ) raise LlmBillingStateError(error) mutation_record = ledger.hold amount = ledger.hold_amount else: mutation_record = mutation.record _log( ctx, LlmBillingEvent.HOLD_SUCCESS, detail=_context_detail( ctx, hold_credits=amount, hold_record_id=mutation_record.id if mutation_record else None, idempotent=not mutation.created, balance_before=mutation.balance_before, balance_after=mutation.balance_after, ledger_state=LlmBillingLedgerState.ACTIVE.value, ), ) return LlmHoldResult( amount, LlmBillingLedgerState.ACTIVE, created=mutation.created, record_id=mutation_record.id if mutation_record else None, ) async def ensure_hold_exists(db: AsyncSession, ctx: LlmBillingContext) -> LlmHoldValidation: """worker 调用模型前确认计费绕过或 active HOLD;不在 worker 首次预扣。""" _log(ctx, LlmBillingEvent.EXECUTION_VALIDATE_START, status="started") ledger = await _load_ledger(db, ctx) # 先尊重已经落库的 attempt 账务状态,再处理当前配置。这样关闭计费不会 # 把运行中的 active HOLD 遗留为永久冻结;而新建且没有 HOLD 的任务才会绕过。 if ledger.state == LlmBillingLedgerState.ACTIVE and ledger.hold: amount = ledger.hold_amount ctx.hold_credits = amount result = LlmHoldValidation(True, amount, ledger.state, hold_record_id=ledger.hold.id) _log( ctx, LlmBillingEvent.EXECUTION_VALIDATE_SUCCESS, detail=_context_detail( ctx, hold_credits=amount, hold_record_id=ledger.hold.id, ledger_state=ledger.state.value, ), ) return result if ledger.state == LlmBillingLedgerState.MISSING: policy = await get_llm_billing_policy( db, config_key=ctx.hold_config_key, explicit_hold_credits=None, ) else: policy = None if policy is not None and policy.bypassed: result = LlmHoldValidation( True, 0.0, LlmBillingLedgerState.BILLING_BYPASSED, "billing_disabled", ) _log( ctx, LlmBillingEvent.EXECUTION_VALIDATE_SUCCESS, detail=_context_detail( ctx, ledger_state=result.state.value, billing_bypassed=True, ), ) return result if policy is not None and not policy.valid: result = LlmHoldValidation(False, 0.0, LlmBillingLedgerState.INVALID, "billing_config_invalid") _log( ctx, LlmBillingEvent.EXECUTION_BLOCKED, status="failed", detail=_context_detail(ctx, ledger_state=result.state.value), error=policy.error, ) return result result = LlmHoldValidation( False, ledger.hold_amount, ledger.state, ledger.reason or f"ledger_{ledger.state.value}", ledger.hold.id if ledger.hold else None, ) event = ( LlmBillingEvent.HOLD_MISSING if ledger.state == LlmBillingLedgerState.MISSING else LlmBillingEvent.EXECUTION_BLOCKED ) _log( ctx, event, status="failed", detail=_context_detail( ctx, hold_credits=result.amount, hold_record_id=result.hold_record_id, ledger_state=result.state.value, skip_reason=result.reason, ), error="LLM预扣不是有效冻结状态,拒绝调用模型", ) return result async def _release_active_hold( db: AsyncSession, ctx: LlmBillingContext, *, hold_record: CreditRecord, reason: str, ) -> BillingItem: amount = _round2(abs(float(hold_record.amount or 0))) ctx.hold_credits = amount if amount <= 0: _log(ctx, LlmBillingEvent.HOLD_RELEASE_SKIPPED, status="skipped", detail=_context_detail(ctx, hold_record_id=hold_record.id, reason=reason, skip_reason="hold_amount_not_positive")) return BillingItem( charge_key=CreditRecordAction.HOLD_RELEASE.value, amount=0.0, charged=False, skipped_reason="hold_amount_not_positive", biz_key=ctx.hold_release_biz_key, attempt_no=ctx.attempt_no, ) _log(ctx, LlmBillingEvent.HOLD_RELEASE_START, status="started", detail=_context_detail(ctx, hold_credits=amount, hold_record_id=hold_record.id, reason=reason)) mutation = await add_credits_result( db, user_id=ctx.user_id, amount=amount, description=f"{ctx.description_prefix}预扣积分释放", related_id=ctx.related_id or ctx.owner_id, record_type="refund", biz_key=ctx.hold_release_biz_key, refund_for_biz_key=ctx.hold_biz_key, record_meta=_hold_meta(ctx, action=CreditRecordAction.HOLD_RELEASE.value), ) _log( ctx, LlmBillingEvent.HOLD_RELEASE_SUCCESS, detail=_context_detail( ctx, hold_credits=amount, hold_record_id=hold_record.id, hold_release_record_id=mutation.record.id if mutation.record else None, reason=reason, idempotent=not mutation.created, balance_before=mutation.balance_before, balance_after=mutation.balance_after, ), ) return BillingItem( charge_key=CreditRecordAction.HOLD_RELEASE.value, amount=amount, charged=False, skipped_reason=None if mutation.created else "already_released", biz_key=ctx.hold_release_biz_key, attempt_no=ctx.attempt_no, ) async def release_hold(db: AsyncSession, ctx: LlmBillingContext, *, reason: str = "failure") -> BillingItem: ledger = await _load_ledger(db, ctx) # 即使管理员已经关闭计费,历史 active HOLD 也必须按真实冻结流水释放。 if ledger.state == LlmBillingLedgerState.ACTIVE and ledger.hold: return await _release_active_hold(db, ctx, hold_record=ledger.hold, reason=reason) if ledger.state in (LlmBillingLedgerState.RELEASED, LlmBillingLedgerState.CHARGED): amount = _round2(abs(float(ledger.release.amount or 0))) if ledger.release else ledger.hold_amount _log( ctx, LlmBillingEvent.HOLD_RELEASE_SKIPPED, status="skipped", detail=_context_detail( ctx, reason=reason, hold_credits=amount, ledger_state=ledger.state.value, skip_reason="already_released", idempotent=True, ), ) return BillingItem( charge_key=CreditRecordAction.HOLD_RELEASE.value, amount=amount, charged=False, skipped_reason="already_released", biz_key=ctx.hold_release_biz_key, attempt_no=ctx.attempt_no, ) if ledger.state == LlmBillingLedgerState.MISSING: policy = await get_llm_billing_policy(db, config_key=ctx.hold_config_key) else: policy = None if policy is not None and policy.bypassed: _log( ctx, LlmBillingEvent.HOLD_RELEASE_SKIPPED, status="skipped", detail=_context_detail( ctx, reason=reason, ledger_state=LlmBillingLedgerState.BILLING_BYPASSED.value, skip_reason="billing_disabled", ), ) return BillingItem( charge_key=CreditRecordAction.HOLD_RELEASE.value, amount=0.0, charged=False, skipped_reason="billing_disabled", biz_key=ctx.hold_release_biz_key, attempt_no=ctx.attempt_no, ) _log( ctx, LlmBillingEvent.HOLD_RELEASE_SKIPPED, status="skipped", detail=_context_detail( ctx, reason=reason, ledger_state=ledger.state.value, skip_reason=ledger.reason or ledger.state.value, ), error="没有可释放的有效LLM预扣流水", ) return BillingItem( charge_key=CreditRecordAction.HOLD_RELEASE.value, amount=0.0, charged=False, skipped_reason=ledger.reason or ledger.state.value, biz_key=ctx.hold_release_biz_key, attempt_no=ctx.attempt_no, ) async def release_on_failure(db: AsyncSession, ctx: LlmBillingContext, *, error: str | None = None) -> BillingSummary: _log(ctx, LlmBillingEvent.FAILURE_RELEASE_START, status="started", detail=_context_detail(ctx, error=error), error=error) item = await release_hold(db, ctx, reason="failure") if item.amount > 0 and not item.skipped_reason: _log(ctx, LlmBillingEvent.FAILURE_RELEASE_SUCCESS, detail=_context_detail(ctx, hold_credits=item.amount, error=error), error=error) else: _log(ctx, LlmBillingEvent.FAILURE_RELEASE_SKIPPED, status="skipped", detail=_context_detail(ctx, hold_credits=item.amount, error=error, skip_reason=item.skipped_reason or "hold_not_active"), error=error) return BillingSummary(record_id=ctx.owner_id, user_id=ctx.user_id, items=[item]) def _normalize_settlement_usage( ctx: LlmBillingContext, usage: Mapping[str, Any], ) -> dict[str, Any]: snapshot = dict(usage or {}) try: if snapshot.get("usage_reported") is False and not bool(snapshot.get("billing_free")): raise ValueError("provider_usage_missing") if "input_tokens" not in snapshot or "output_tokens" not in snapshot: raise ValueError("input_or_output_tokens_missing") input_tokens = int(snapshot.get("input_tokens")) output_tokens = int(snapshot.get("output_tokens")) if input_tokens < 0 or output_tokens < 0: raise ValueError("token_count_negative") normalized_total = input_tokens + output_tokens raw_total = snapshot.get("total_tokens") if raw_total not in (None, ""): total_tokens = int(raw_total) if total_tokens < 0: raise ValueError("total_tokens_negative") else: total_tokens = normalized_total if total_tokens != normalized_total: snapshot["reported_total_tokens"] = total_tokens total_tokens = normalized_total snapshot["input_tokens"] = input_tokens snapshot["output_tokens"] = output_tokens snapshot["total_tokens"] = total_tokens return snapshot except (TypeError, ValueError) as exc: _log( ctx, LlmBillingEvent.USAGE_INVALID, status="failed", detail=_context_detail(ctx, usage=snapshot, reason=str(exc)), error="LLM usage 无效,拒绝释放 HOLD 和创建真实扣费", ) raise LlmBillingStateError(f"LLM usage 无效:{exc}") from exc async def _ensure_token_usage_once( db: AsyncSession, ctx: LlmBillingContext, usage: dict[str, Any], ) -> TokenUsage: supplied_id = str(usage.get("token_usage_id") or "").strip() or None if supplied_id: supplied_result = await db.execute( select(TokenUsage).where(TokenUsage.id == supplied_id).limit(1) ) supplied = supplied_result.scalar_one_or_none() if supplied is not None: if supplied.user_id not in (None, ctx.user_id): raise LlmBillingStateError("TokenUsage 用户归属与当前账务上下文不一致") if supplied.biz_key not in (None, ctx.charge_biz_key): raise LlmBillingStateError("TokenUsage biz_key 与当前 charge 不一致") supplied.user_id = supplied.user_id or ctx.user_id supplied.owner_type = supplied.owner_type or ctx.owner_type supplied.owner_id = supplied.owner_id or ctx.owner_id supplied.biz_key = supplied.biz_key or ctx.charge_biz_key supplied.source_module = supplied.source_module or ctx.source_module supplied.source_step_code = supplied.source_step_code or ctx.source_step_code usage["token_usage_id"] = supplied.id ctx.token_usage_id = supplied.id _log( ctx, LlmBillingEvent.TOKEN_USAGE_REUSED, detail=_context_detail(ctx, token_usage_id=supplied.id, source="supplied_id"), ) return supplied result = await db.execute( select(TokenUsage) .where( TokenUsage.user_id == ctx.user_id, TokenUsage.biz_key == ctx.charge_biz_key, ) .limit(1) ) existing = result.scalar_one_or_none() if existing is not None: usage["token_usage_id"] = existing.id ctx.token_usage_id = existing.id _log( ctx, LlmBillingEvent.TOKEN_USAGE_REUSED, detail=_context_detail(ctx, token_usage_id=existing.id, source="biz_key"), ) return existing token_usage = TokenUsage( id=generate_id(), model_config_id=usage.get("model_config_id"), user_id=ctx.user_id, input_tokens=int(usage["input_tokens"]), output_tokens=int(usage["output_tokens"]), total_tokens=int(usage["total_tokens"]), owner_type=ctx.owner_type, owner_id=ctx.owner_id, biz_key=ctx.charge_biz_key, source_module=ctx.source_module, source_step_code=ctx.source_step_code, ) try: async with db.begin_nested(): db.add(token_usage) await db.flush() except IntegrityError: if token_usage in db.sync_session: db.sync_session.expunge(token_usage) result = await db.execute( select(TokenUsage) .where( TokenUsage.user_id == ctx.user_id, TokenUsage.biz_key == ctx.charge_biz_key, ) .limit(1) ) token_usage = result.scalar_one_or_none() if token_usage is None: raise event = LlmBillingEvent.TOKEN_USAGE_REUSED else: event = LlmBillingEvent.TOKEN_USAGE_CREATED usage["token_usage_id"] = token_usage.id ctx.token_usage_id = token_usage.id _log( ctx, event, detail=_context_detail( ctx, token_usage_id=token_usage.id, input_tokens=token_usage.input_tokens, output_tokens=token_usage.output_tokens, total_tokens=token_usage.total_tokens, ), ) return token_usage async def _build_charge_meta(db: AsyncSession, ctx: LlmBillingContext, usage: Mapping[str, Any]) -> CreditRecordMeta: usage_snapshot = dict(usage or {}) ctx.provider = str(usage_snapshot.get("provider") or usage_snapshot.get("model_provider") or "") or ctx.provider ctx.model_name = str(usage_snapshot.get("model_name") or usage_snapshot.get("model") or "") or ctx.model_name ctx.token_usage_id = str(usage_snapshot.get("token_usage_id") or "") or ctx.token_usage_id if ctx.owner_type == CreditRecordOwnerType.GENERATION_RECORD.value: return await build_generation_record_prompt_meta( db, record_id=ctx.owner_id, attempt_no=ctx.attempt_no, charge_kind=ctx.charge_kind, usage=usage_snapshot, ) if ctx.owner_type == CreditRecordOwnerType.MODULE_GENERATION_STEP.value: return await build_module_step_prompt_meta( db, step_id=ctx.owner_id, attempt_no=ctx.attempt_no, usage=usage_snapshot, ) if ctx.charge_kind == CreditRecordChargeKind.VIDEO_ANALYSIS.value: return await build_shot_video_analysis_meta( db, owner_type=ctx.owner_type, owner_id=ctx.owner_id, attempt_no=ctx.attempt_no, usage=usage_snapshot, billing_scene=ctx.billing_scene or CreditRecordBillingScene.SHOT_VIDEO_ANALYSIS.value, source_project_id=ctx.source_project_id, source_step_id=ctx.source_step_id, ) return CreditRecordMeta( owner_type=ctx.owner_type, owner_id=ctx.owner_id, attempt_no=ctx.attempt_no, charge_kind=ctx.charge_kind, charge_action=CreditRecordAction.CHARGE.value, credit_subject=CreditRecordSubject.TEXT.value, billing_scene=ctx.billing_scene, source_module=ctx.source_module, source_project_id=ctx.source_project_id, source_step_id=ctx.source_step_id, source_step_code=ctx.source_step_code, token_usage_id=usage_snapshot.get("token_usage_id"), input_tokens=_safe_int(usage_snapshot.get("input_tokens")), output_tokens=_safe_int(usage_snapshot.get("output_tokens")), total_tokens=_safe_int(usage_snapshot.get("total_tokens")), ) async def _settle_success_impl( db: AsyncSession, ctx: LlmBillingContext, *, usage: Mapping[str, Any], description: str | None = None, ) -> BillingSummary: ledger = await _load_ledger(db, ctx) if ledger.state == LlmBillingLedgerState.CHARGED and ledger.hold and ledger.release and ledger.charge: release_item = BillingItem( charge_key=CreditRecordAction.HOLD_RELEASE.value, amount=abs(_round2(ledger.release.amount)), charged=False, skipped_reason="already_released", biz_key=ctx.hold_release_biz_key, attempt_no=ctx.attempt_no, ) charge_item = BillingItem( charge_key=ctx.charge_kind, amount=abs(_round2(ledger.charge.amount)), charged=False, skipped_reason="already_charged", biz_key=ctx.charge_biz_key, attempt_no=ctx.attempt_no, ) _log( ctx, LlmBillingEvent.SETTLE_SUCCESS, detail=_context_detail( ctx, actual_credits=charge_item.amount, ledger_state=ledger.state.value, idempotent=True, ), ) return BillingSummary(record_id=ctx.owner_id, user_id=ctx.user_id, items=[release_item, charge_item]) # 新任务在关闭计费时没有 HOLD,成功后直接绕过;历史 active HOLD 则必须继续结算。 if ledger.state == LlmBillingLedgerState.MISSING: policy = await get_llm_billing_policy(db, config_key=ctx.hold_config_key) else: policy = None if policy is not None and policy.bypassed: _log( ctx, LlmBillingEvent.SETTLE_SUCCESS, status="skipped", detail=_context_detail( ctx, ledger_state=LlmBillingLedgerState.BILLING_BYPASSED.value, billing_bypassed=True, ), ) return BillingSummary(record_id=ctx.owner_id, user_id=ctx.user_id, items=[]) if ledger.state != LlmBillingLedgerState.ACTIVE or ledger.hold is None: # 统一由 settle_success 外层记录一次 SETTLE_FAILED,避免同一异常产生重复日志。 raise LlmBillingStateError( f"当前attempt账务状态为{ledger.state.value},不能执行成功结算" ) normalized_usage = _normalize_settlement_usage(ctx, usage) await _ensure_token_usage_once(db, ctx, normalized_usage) _log( ctx, LlmBillingEvent.SETTLE_START, status="started", detail=_context_detail( ctx, ledger_state=ledger.state.value, input_tokens=normalized_usage["input_tokens"], output_tokens=normalized_usage["output_tokens"], total_tokens=normalized_usage["total_tokens"], provider=normalized_usage.get("provider") or normalized_usage.get("model_provider"), model_name=normalized_usage.get("model_name") or normalized_usage.get("model"), ), ) input_tokens = int(normalized_usage["input_tokens"]) output_tokens = int(normalized_usage["output_tokens"]) amount = await calc_text_credits(db, input_tokens, output_tokens) meta = await _build_charge_meta(db, ctx, normalized_usage) release_item = await _release_active_hold(db, ctx, hold_record=ledger.hold, reason="success") if meta.charge_action is None: meta.charge_action = CreditRecordAction.CHARGE.value meta.billing_scene = meta.billing_scene or ctx.billing_scene meta.source_module = meta.source_module or ctx.source_module meta.source_project_id = meta.source_project_id or ctx.source_project_id meta.source_step_id = meta.source_step_id or ctx.source_step_id meta.source_step_code = meta.source_step_code or ctx.source_step_code mutation = await deduct_credits_result( db, user_id=ctx.user_id, amount=amount, description=description or f"{ctx.description_prefix}真实扣费", related_id=ctx.related_id or ctx.owner_id, biz_key=ctx.charge_biz_key, record_meta=meta, allow_negative=True, create_zero_record=True, ) charged_amount = mutation.amount charge_item = BillingItem( charge_key=ctx.charge_kind, amount=charged_amount, charged=mutation.created, skipped_reason=None if mutation.created else "already_charged", biz_key=ctx.charge_biz_key, attempt_no=ctx.attempt_no, ) if ctx.owner_type == CreditRecordOwnerType.MODULE_GENERATION_STEP.value: result = await db.execute(select(ModuleGenerationStep).where(ModuleGenerationStep.id == ctx.owner_id).limit(1)) step = result.scalar_one_or_none() if step: step.token_usage_id = meta.token_usage_id step.model_config_id = normalized_usage.get("model_config_id") step.input_tokens = meta.input_tokens step.output_tokens = meta.output_tokens step.total_tokens = meta.total_tokens step.text_credits_cost = charged_amount after_balance = mutation.balance_after _log( ctx, LlmBillingEvent.CHARGE_SUCCESS, detail=_context_detail( ctx, actual_credits=charged_amount, charge_record_id=mutation.record.id if mutation.record else None, idempotent=not mutation.created, balance_before=mutation.balance_before, balance_after=after_balance, allow_negative=True, ), ) if after_balance < 0: _log(ctx, LlmBillingEvent.CHARGE_NEGATIVE_BALANCE, status="warning", detail=_context_detail(ctx, actual_credits=charged_amount, balance_after=after_balance, allow_negative=True)) _log(ctx, LlmBillingEvent.SETTLE_SUCCESS, detail=_context_detail(ctx, actual_credits=charged_amount, balance_after=after_balance, ledger_state=LlmBillingLedgerState.CHARGED.value, idempotent=not mutation.created)) return BillingSummary(record_id=ctx.owner_id, user_id=ctx.user_id, items=[release_item, charge_item]) async def settle_success( db: AsyncSession, ctx: LlmBillingContext, *, usage: Mapping[str, Any], description: str | None = None, ) -> BillingSummary: """成功结算统一入口;任何异常都留下可检索的 SETTLE_FAILED 日志。""" try: return await _settle_success_impl( db, ctx, usage=usage, description=description, ) except Exception as exc: _log( ctx, LlmBillingEvent.SETTLE_FAILED, status="failed", detail=_context_detail(ctx, error_type=type(exc).__name__), error=str(exc), ) raise