feat: 强化用量计费状态机,新增钱包每日消费汇总分类账

- 将 usage.billing_status 默认值从 settled 改为 pending,完善
  pending -> settled/void 的状态转换逻辑,确保终态不可逆
- 新增 WalletDailyUsageLedger 模型和聚合服务,按账单日汇总
  每个钱包的消费金额、请求数和 token 用量
- 前端钱包中心页面集成每日消费流水展示,支持与充值记录混合
  排序和分页
- 新增两个数据库迁移:修复历史数据状态一致性、创建每日汇总表
- 补充计费状态机单元测试

Closes #218

Co-authored-by: LewisPen <LewisPen@nyadoo.com>
This commit is contained in:
fawney19
2026-03-11 15:05:11 +08:00
parent 6235c772ac
commit 04ab4bd9f2
25 changed files with 1212 additions and 121 deletions

View File

@@ -37,6 +37,7 @@ QUIET_POLLING_PATHS: set[str] = {
"/api/admin/usage/stats",
"/api/admin/usage/aggregation/stats",
"/api/admin/health/status",
"/api/wallet/today-cost",
}

View File

@@ -492,6 +492,9 @@ class StreamTelemetryRecorder:
usage = db.query(Usage).filter(Usage.request_id == request_id).first()
if usage:
if getattr(usage, "billing_status", None) in {"settled", "void"}:
logger.debug("[{}] Usage 已终态,跳过快速状态更新: {}", self.request_id, status)
return
setattr(usage, "status", status)
setattr(usage, "status_code", status_code)
setattr(usage, "response_time_ms", response_time_ms)

View File

@@ -5,6 +5,7 @@ from .wallet_payment import (
serialize_admin_wallet_transaction,
serialize_payment_callback,
serialize_payment_order,
serialize_wallet_daily_usage,
serialize_wallet_payload,
serialize_wallet_refund,
serialize_wallet_transaction,
@@ -17,6 +18,7 @@ __all__ = [
"serialize_admin_wallet_transaction",
"serialize_payment_callback",
"serialize_payment_order",
"serialize_wallet_daily_usage",
"serialize_wallet_payload",
"serialize_wallet_refund",
"serialize_wallet_transaction",

View File

@@ -7,9 +7,10 @@ from src.models.database import (
PaymentOrder,
RefundRequest,
Wallet,
WalletDailyUsageLedger,
WalletTransaction,
)
from src.services.wallet import WalletService
from src.services.wallet import WalletDailyUsageSnapshot, WalletService
def safe_gateway_response(raw: dict[str, Any] | None) -> dict[str, Any]:
@@ -124,6 +125,27 @@ def serialize_wallet_transaction(tx: WalletTransaction) -> dict[str, Any]:
}
def serialize_wallet_daily_usage(
ledger: WalletDailyUsageLedger | WalletDailyUsageSnapshot,
) -> dict[str, Any]:
billing_date = getattr(ledger, "billing_date", None)
return {
"id": getattr(ledger, "id", None),
"date": billing_date.isoformat() if billing_date is not None else None,
"timezone": getattr(ledger, "billing_timezone", None),
"total_cost": float(getattr(ledger, "total_cost_usd", 0) or 0),
"total_requests": int(getattr(ledger, "total_requests", 0) or 0),
"input_tokens": int(getattr(ledger, "input_tokens", 0) or 0),
"output_tokens": int(getattr(ledger, "output_tokens", 0) or 0),
"cache_creation_tokens": int(getattr(ledger, "cache_creation_tokens", 0) or 0),
"cache_read_tokens": int(getattr(ledger, "cache_read_tokens", 0) or 0),
"first_finalized_at": getattr(ledger, "first_finalized_at", None),
"last_finalized_at": getattr(ledger, "last_finalized_at", None),
"aggregated_at": getattr(ledger, "aggregated_at", None),
"is_today": bool(getattr(ledger, "is_today", False)),
}
def serialize_wallet_refund(refund: RefundRequest) -> dict[str, Any]:
return {
"id": refund.id,

View File

@@ -3,9 +3,10 @@
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime, timezone
from datetime import date, datetime, timezone
from typing import Any
from uuid import uuid4
from zoneinfo import ZoneInfo
from fastapi import APIRouter, Depends, Query, Request
from pydantic import BaseModel, Field, ValidationError
@@ -18,15 +19,22 @@ from src.api.base.pipeline import ApiRequestPipeline
from src.api.serializers import (
safe_gateway_response,
serialize_payment_order,
serialize_wallet_daily_usage,
serialize_wallet_payload,
serialize_wallet_refund,
serialize_wallet_transaction,
)
from src.core.exceptions import InvalidRequestException, NotFoundException, translate_pydantic_error
from src.database import get_db
from src.models.database import PaymentOrder, RefundRequest, Wallet, WalletTransaction
from src.models.database import (
PaymentOrder,
RefundRequest,
Wallet,
WalletDailyUsageLedger,
WalletTransaction,
)
from src.services.payment import PaymentService
from src.services.wallet import WalletService
from src.services.wallet import WalletDailyUsageLedgerService, WalletService
router = APIRouter(prefix="/api/wallet", tags=["Wallet"])
pipeline = ApiRequestPipeline()
@@ -61,6 +69,43 @@ def _build_refund_no() -> str:
return f"rf_{ts}_{uuid4().hex[:8]}"
def _resolve_user_wallet(db: Session, user: Any) -> Wallet | None:
existing_wallet = WalletService.get_wallet(db, user_id=user.id)
wallet = existing_wallet or WalletService.get_or_create_wallet(db, user=user)
if wallet is not None and existing_wallet is None:
db.commit()
db.refresh(wallet)
return wallet
def _flow_sort_key(item: dict[str, Any], billing_tz: ZoneInfo) -> tuple[date, int, datetime]:
item_type = item.get("type")
data = item.get("data") or {}
if item_type == "daily_usage":
raw_date = data.get("date")
billing_date = (
date.fromisoformat(raw_date) if isinstance(raw_date, str) and raw_date else date.min
)
sort_dt = (
data.get("last_finalized_at")
or data.get("aggregated_at")
or datetime.min.replace(tzinfo=timezone.utc)
)
if isinstance(sort_dt, str):
sort_dt = datetime.fromisoformat(sort_dt.replace("Z", "+00:00"))
return billing_date, 1, sort_dt
created_at = data.get("created_at")
if isinstance(created_at, str):
created_at = datetime.fromisoformat(created_at.replace("Z", "+00:00"))
if not isinstance(created_at, datetime):
created_at = datetime.min.replace(tzinfo=timezone.utc)
elif created_at.tzinfo is None:
created_at = created_at.replace(tzinfo=timezone.utc)
local_date = created_at.astimezone(billing_tz).date()
return local_date, 0, created_at
@router.get("/balance")
async def get_wallet_balance(request: Request, db: Session = Depends(get_db)) -> Any:
adapter = WalletBalanceAdapter()
@@ -78,6 +123,23 @@ async def list_wallet_transactions(
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/flow")
async def get_wallet_flow(
request: Request,
limit: int = Query(50, ge=1, le=200),
offset: int = Query(0, ge=0, le=5000),
db: Session = Depends(get_db),
) -> Any:
adapter = WalletFlowAdapter(limit=limit, offset=offset)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/today-cost")
async def get_wallet_today_cost(request: Request, db: Session = Depends(get_db)) -> Any:
adapter = WalletTodayCostAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.post("/recharge")
async def create_recharge_order(request: Request, db: Session = Depends(get_db)) -> Any:
adapter = WalletRechargeCreateAdapter()
@@ -135,8 +197,7 @@ class WalletTransactionsAdapter(AuthenticatedApiAdapter):
if user is None:
raise InvalidRequestException("未登录")
existing_wallet = WalletService.get_wallet(db, user_id=user.id)
wallet = existing_wallet or WalletService.get_or_create_wallet(db, user=user)
wallet = _resolve_user_wallet(db, user)
if wallet is None:
return {
"items": [],
@@ -146,10 +207,6 @@ class WalletTransactionsAdapter(AuthenticatedApiAdapter):
**serialize_wallet_payload(None),
}
if existing_wallet is None:
db.commit()
db.refresh(wallet)
base_query = db.query(WalletTransaction).filter(WalletTransaction.wallet_id == wallet.id)
total = base_query.count()
items = (
@@ -168,6 +225,80 @@ class WalletTransactionsAdapter(AuthenticatedApiAdapter):
}
@dataclass
class WalletFlowAdapter(AuthenticatedApiAdapter):
limit: int
offset: int
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
db = context.db
user = context.user
if user is None:
raise InvalidRequestException("未登录")
wallet = _resolve_user_wallet(db, user)
if wallet is None:
return {
"today_entry": None,
"items": [],
"total": 0,
"limit": self.limit,
"offset": self.offset,
**serialize_wallet_payload(None),
}
today_entry = WalletDailyUsageLedgerService.get_today_snapshot(db, wallet.id)
billing_tz = WalletDailyUsageLedgerService.get_timezone(today_entry.billing_timezone)
fetch_size = min(self.offset + self.limit, 5200)
tx_query = db.query(WalletTransaction).filter(WalletTransaction.wallet_id == wallet.id)
tx_total = tx_query.count()
tx_items = tx_query.order_by(WalletTransaction.created_at.desc()).limit(fetch_size).all()
daily_query = db.query(WalletDailyUsageLedger).filter(
WalletDailyUsageLedger.wallet_id == wallet.id,
WalletDailyUsageLedger.billing_timezone == today_entry.billing_timezone,
WalletDailyUsageLedger.billing_date < today_entry.billing_date,
)
daily_total = daily_query.count()
daily_items = (
daily_query.order_by(WalletDailyUsageLedger.billing_date.desc()).limit(fetch_size).all()
)
merged = [
{"type": "transaction", "data": serialize_wallet_transaction(item)} for item in tx_items
] + [
{"type": "daily_usage", "data": serialize_wallet_daily_usage(item)}
for item in daily_items
]
merged.sort(key=lambda item: _flow_sort_key(item, billing_tz), reverse=True)
paged = merged[self.offset : self.offset + self.limit]
return {
"today_entry": serialize_wallet_daily_usage(today_entry),
"items": paged,
"total": tx_total + daily_total,
"limit": self.limit,
"offset": self.offset,
**serialize_wallet_payload(wallet),
}
class WalletTodayCostAdapter(AuthenticatedApiAdapter):
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
db = context.db
user = context.user
if user is None:
raise InvalidRequestException("未登录")
wallet = _resolve_user_wallet(db, user)
snapshot = WalletDailyUsageLedgerService.get_today_snapshot(
db,
wallet.id if wallet is not None else None,
)
return serialize_wallet_daily_usage(snapshot)
class WalletBalanceAdapter(AuthenticatedApiAdapter):
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
db = context.db
@@ -175,15 +306,10 @@ class WalletBalanceAdapter(AuthenticatedApiAdapter):
if user is None:
raise InvalidRequestException("未登录")
existing_wallet = WalletService.get_wallet(db, user_id=user.id)
wallet = existing_wallet or WalletService.get_or_create_wallet(db, user=user)
wallet = _resolve_user_wallet(db, user)
if wallet is None:
return serialize_wallet_payload(None)
if existing_wallet is None:
db.commit()
db.refresh(wallet)
pending_refunds = (
db.query(RefundRequest)
.filter(

View File

@@ -7,7 +7,7 @@ from __future__ import annotations
import hashlib
import secrets
import uuid
from datetime import datetime, timezone
from datetime import date, datetime, timezone
from enum import Enum as PyEnum
from typing import Any, ClassVar
@@ -18,6 +18,7 @@ from sqlalchemy import (
Boolean,
CheckConstraint,
Column,
Date,
DateTime,
Enum,
Float,
@@ -413,10 +414,10 @@ class Usage(Base):
status = Column(String(20), default="completed", nullable=False, index=True)
# 结算状态(与 status 解耦)
# - pending: 等待结算(任务未完成 / 流式未结束
# - settled: 已结算cost 已写入,可能 > 0 或 = 0
# - void: 作废(不收费,如任务未开始就取消
billing_status = Column(String(20), default="settled", nullable=False, index=True)
# - pending: 等待结算(请求已创建,但账务尚未进入最终状态
# - settled: 已结算cost 已写入,且钱包侧结算动作已完成
# - void: 作废(明确不收费)
billing_status = Column(String(20), default="pending", nullable=False, index=True)
finalized_at = Column(DateTime(timezone=True), nullable=True) # 结算完成时间(可选)
wallet_balance_before = Column(Numeric(20, 8), nullable=True) # 结算前可用总余额快照
wallet_balance_after = Column(Numeric(20, 8), nullable=True) # 结算后可用总余额快照
@@ -556,6 +557,9 @@ class Wallet(Base):
transactions = relationship(
"WalletTransaction", back_populates="wallet", cascade="all, delete-orphan"
)
daily_usage_ledgers = relationship(
"WalletDailyUsageLedger", back_populates="wallet", cascade="all, delete-orphan"
)
payment_orders = relationship("PaymentOrder", back_populates="wallet")
refund_requests = relationship("RefundRequest", back_populates="wallet")
@@ -609,6 +613,50 @@ class WalletTransaction(Base):
operator = relationship("User")
class WalletDailyUsageLedger(Base):
"""钱包按天汇总的消费流水投影。"""
__tablename__ = "wallet_daily_usage_ledgers"
__table_args__ = (
UniqueConstraint(
"wallet_id",
"billing_date",
"billing_timezone",
name="uq_wallet_daily_usage_ledgers_wallet_date_tz",
),
Index("idx_wallet_daily_usage_wallet_date", "wallet_id", "billing_date"),
Index("idx_wallet_daily_usage_date", "billing_date"),
)
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
wallet_id = Column(String(36), ForeignKey("wallets.id", ondelete="CASCADE"), nullable=False)
billing_date = Column(Date, nullable=False)
billing_timezone = Column(String(64), nullable=False, default="Asia/Shanghai")
total_cost_usd = Column(Numeric(20, 8), nullable=False, default=0)
total_requests = Column(Integer, nullable=False, default=0)
input_tokens = Column(BigInteger, nullable=False, default=0)
output_tokens = Column(BigInteger, nullable=False, default=0)
cache_creation_tokens = Column(BigInteger, nullable=False, default=0)
cache_read_tokens = Column(BigInteger, nullable=False, default=0)
first_finalized_at = Column(DateTime(timezone=True), nullable=True)
last_finalized_at = Column(DateTime(timezone=True), nullable=True)
aggregated_at = Column(DateTime(timezone=True), nullable=False)
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
wallet = relationship("Wallet", back_populates="daily_usage_ledgers")
class PaymentOrder(Base):
"""充值订单"""

View File

@@ -32,6 +32,7 @@ from src.services.system.config import SystemConfigService
from src.services.system.scheduler import get_scheduler
from src.services.system.stats_aggregator import StatsAggregatorService
from src.services.user.apikey import ApiKeyService
from src.services.wallet import WalletDailyUsageLedgerService
from src.utils.compression import compress_json
@@ -45,6 +46,7 @@ class MaintenanceScheduler:
self.running = False
self._interval_tasks = []
self._stats_aggregation_lock = asyncio.Lock()
self._wallet_daily_usage_lock = asyncio.Lock()
def _get_checkin_time(self) -> tuple[int, int]:
"""获取签到任务的执行时间
@@ -144,6 +146,13 @@ class MaintenanceScheduler:
name="统计小时数据聚合",
timezone="UTC",
)
scheduler.add_cron_job(
self._scheduled_wallet_daily_usage_aggregation,
hour=0,
minute=10,
job_id="wallet_daily_usage_aggregation",
name="钱包每日消费汇总",
)
# 清理任务 - 凌晨 3 点执行
scheduler.add_cron_job(
self._scheduled_cleanup,
@@ -268,6 +277,10 @@ class MaintenanceScheduler:
"""统计聚合任务(定时调用)"""
await self._perform_stats_aggregation(backfill=backfill)
async def _scheduled_wallet_daily_usage_aggregation(self) -> None:
"""钱包每日消费汇总任务(定时调用)"""
await self._perform_wallet_daily_usage_aggregation()
async def _scheduled_hourly_stats_aggregation(self) -> None:
"""小时统计聚合任务(定时调用)"""
await self._perform_hourly_stats_aggregation()
@@ -499,6 +512,32 @@ class MaintenanceScheduler:
finally:
db.close()
async def _perform_wallet_daily_usage_aggregation(self) -> None:
if self._wallet_daily_usage_lock.locked():
logger.info("钱包每日消费汇总任务正在运行,跳过本次触发")
return
async with self._wallet_daily_usage_lock:
db = create_session()
try:
logger.info("开始执行钱包每日消费汇总...")
billing_today = WalletDailyUsageLedgerService.get_today_billing_date()
billing_yesterday = billing_today - timedelta(days=1)
affected = WalletDailyUsageLedgerService.aggregate_day(db, billing_yesterday)
logger.info(
"钱包每日消费汇总完成: date={}, wallets={}",
billing_yesterday.isoformat(),
affected,
)
except Exception as e:
logger.exception("钱包每日消费汇总任务执行失败: {}", e)
try:
db.rollback()
except Exception:
pass
finally:
db.close()
async def _perform_hourly_stats_aggregation(self) -> None:
"""执行小时统计聚合任务"""
db = create_session()

View File

@@ -92,6 +92,7 @@ class VideoTaskBillingService:
provider_api_key_id=getattr(task, "key_id", None),
status="completed" if task.status == "completed" else "failed",
target_model=None,
finalized_at=getattr(task, "completed_at", None),
)
return True
except Exception as exc:
@@ -122,6 +123,8 @@ class VideoTaskBillingService:
if not request_id:
return False
# Advisory check无锁仅用于快速跳过实际状态转换由 update_settled_billing 的
# with_for_update() 保证原子性。
existing = self.db.query(Usage).filter(Usage.request_id == request_id).first()
if not existing:
logger.warning(
@@ -131,12 +134,12 @@ class VideoTaskBillingService:
)
return await self._create_fallback_usage_for_video_task(task, request_id)
metadata = existing.request_metadata or {}
if metadata.get("billing_updated_at"):
if getattr(existing, "billing_status", None) != "pending":
logger.debug(
"Video task billing already updated: task_id={} request_id={}",
"Skip video task billing finalize because Usage is already terminal: task_id={} request_id={} billing_status={}",
getattr(task, "id", None),
request_id,
getattr(existing, "billing_status", None),
)
return False
@@ -300,6 +303,7 @@ class VideoTaskBillingService:
"field": "video_tasks.request_metadata.poll_raw_response",
},
},
finalized_at=getattr(task, "completed_at", None),
)
if updated:

View File

@@ -44,6 +44,19 @@ class VideoTaskCancelService:
from src.services.provider.auth import get_provider_auth
from src.services.provider.transport import build_provider_url
current_status = str(getattr(task, "status", "") or "")
non_cancellable_statuses = {
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
VideoStatus.EXPIRED.value,
}
if current_status in non_cancellable_statuses:
raise HTTPException(
status_code=409,
detail=f"Task cannot be cancelled in status: {current_status}",
)
external_task_id = getattr(task, "external_task_id", None)
if not external_task_id:
raise HTTPException(status_code=500, detail="Task missing external_task_id")
@@ -122,8 +135,10 @@ class VideoTaskCancelService:
detail=f"Cancel not supported for provider format: {provider_format}",
)
now = datetime.now(timezone.utc)
task.status = VideoStatus.CANCELLED.value
task.updated_at = datetime.now(timezone.utc)
task.completed_at = getattr(task, "completed_at", None) or now
task.updated_at = now
# Void Usage (no charge)
try:
@@ -131,12 +146,13 @@ class VideoTaskCancelService:
self.db,
request_id=task.request_id,
reason="cancelled_by_user",
finalized_at=task.completed_at,
)
if not voided:
UsageService.void_settled(
self.db,
request_id=task.request_id,
reason="cancelled_by_user",
logger.warning(
"Skip voiding video usage because billing is already terminal: task_id={} request_id={}",
getattr(task, "id", task_id),
getattr(task, "request_id", None),
)
except Exception as exc:
logger.warning(

View File

@@ -260,6 +260,19 @@ class VideoTaskPollerAdapter:
logger.warning("Task {} disappeared during poll update", task_id)
return
if task.status in {
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
VideoStatus.EXPIRED.value,
}:
logger.debug(
"Skip poll update for terminal task {} with status {}",
task_id,
task.status,
)
return
if error_exception is not None and ctx is not None:
# HTTP 请求失败(需要 ctx 来计算 backoff
self._handle_poll_error(task, error_exception, ctx)
@@ -382,6 +395,17 @@ class VideoTaskPollerAdapter:
"""
兼容入口:复用三阶段轮询流程,避免维护重复逻辑。
"""
if task.status in {
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
VideoStatus.EXPIRED.value,
}:
logger.debug(
"Skip legacy poll for terminal task {} with status {}", task.id, task.status
)
return
ctx_or_result = await self.prepare_poll_context(db, task)
if isinstance(ctx_or_result, InternalVideoPollResult):

View File

@@ -130,7 +130,7 @@ class UsageActiveRequestsMixin:
while True:
stale_requests = (
db.query(Usage.id, Usage.request_id, Usage.status)
db.query(Usage.id, Usage.request_id, Usage.status, Usage.billing_status)
.filter(
Usage.status.in_(["pending", "streaming"]),
Usage.created_at < cutoff_time,
@@ -143,13 +143,13 @@ class UsageActiveRequestsMixin:
if not stale_requests:
break
stale_request_ids = [request_id for _, request_id, _ in stale_requests if request_id]
stale_request_ids = [request_id for _, request_id, _, _ in stale_requests if request_id]
completed_request_ids = cls._find_completed_request_ids(db, stale_request_ids)
usage_updates = []
failed_request_ids: list[str] = []
for usage_id, request_id, old_status in stale_requests:
for usage_id, request_id, old_status, billing_status in stale_requests:
if request_id and request_id in completed_request_ids:
usage_updates.append(
{
@@ -161,16 +161,22 @@ class UsageActiveRequestsMixin:
)
recovered_count += 1
else:
usage_updates.append(
{
"id": usage_id,
"status": "failed",
"status_code": 504,
"error_message": (
f"请求超时: 状态 '{old_status}' 超过 {timeout_minutes} 分钟未完成"
),
}
)
entry: dict[str, Any] = {
"id": usage_id,
"status": "failed",
"status_code": 504,
"error_message": (
f"请求超时: 状态 '{old_status}' 超过 {timeout_minutes} 分钟未完成"
),
}
if billing_status == "pending":
entry["billing_status"] = "void"
entry["finalized_at"] = now
entry["total_cost_usd"] = 0.0
entry["request_cost_usd"] = 0.0
entry["actual_total_cost_usd"] = 0.0
entry["actual_request_cost_usd"] = 0.0
usage_updates.append(entry)
failed_count += 1
if request_id:
failed_request_ids.append(request_id)

View File

@@ -11,6 +11,7 @@ import json
import os
import socket
import time
from datetime import datetime, timezone
from typing import Any
from redis.exceptions import ConnectionError as RedisConnectionError
@@ -58,6 +59,10 @@ def _event_to_record(event: UsageEvent) -> dict[str, Any]:
elif event.event_type == UsageEventType.CANCELLED:
status = "cancelled"
finalized_at = None
if event.timestamp_ms > 0:
finalized_at = datetime.fromtimestamp(event.timestamp_ms / 1000, tz=timezone.utc)
return {
"request_id": event.request_id,
"user_id": data.get("user_id"),
@@ -95,6 +100,7 @@ def _event_to_record(event: UsageEvent) -> dict[str, Any]:
"provider_api_key_id": data.get("provider_api_key_id"),
"status": status,
"target_model": data.get("target_model"),
"finalized_at": finalized_at,
}
@@ -491,6 +497,11 @@ class UsageQueueConsumer:
provider_api_key_id=data.get("provider_api_key_id"),
status=status,
target_model=data.get("target_model"),
finalized_at=(
datetime.fromtimestamp(event.timestamp_ms / 1000, tz=timezone.utc)
if event.timestamp_ms > 0
else None
),
)
finally:
if own_session:

View File

@@ -18,6 +18,12 @@ from src.services.wallet import WalletService
class UsageLifecycleMixin:
"""使用记录生命周期管理方法"""
@staticmethod
def _is_billing_terminal(usage: Usage | None) -> bool:
return bool(
usage is not None and getattr(usage, "billing_status", None) in {"settled", "void"}
)
@classmethod
def begin_pending_usage(
cls,
@@ -152,6 +158,7 @@ class UsageLifecycleMixin:
response_time_ms: int | None = None,
billing_snapshot: dict[str, Any] | None = None,
extra_metadata: dict[str, Any] | None = None,
finalized_at: datetime | None = None,
) -> bool:
"""
并发安全的幂等 finalizesettled
@@ -160,7 +167,7 @@ class UsageLifecycleMixin:
- 仅当 billing_status='pending' 时才会生效rowcount==1
- 不在本方法内 commit由调用方决定事务提交时机
"""
now = datetime.now(timezone.utc)
now = finalized_at or datetime.now(timezone.utc)
cost = to_money_decimal(total_cost_usd)
request_cost = to_money_decimal(request_cost_usd) if request_cost_usd is not None else cost
@@ -197,6 +204,7 @@ class UsageLifecycleMixin:
*,
reason: str | None = None,
status_code: int = 499,
finalized_at: datetime | None = None,
) -> bool:
"""
并发安全的幂等 finalizevoid不收费
@@ -205,7 +213,7 @@ class UsageLifecycleMixin:
- 仅当 billing_status='pending' 时才会生效rowcount==1
- 不在本方法内 commit由调用方决定事务提交时机
"""
now = datetime.now(timezone.utc)
now = finalized_at or datetime.now(timezone.utc)
usage = db.query(Usage).filter(Usage.request_id == request_id).with_for_update().first()
if not usage or usage.billing_status != "pending":
return False
@@ -214,6 +222,8 @@ class UsageLifecycleMixin:
usage.finalized_at = now
usage.total_cost_usd = to_money_decimal(0)
usage.request_cost_usd = to_money_decimal(0)
usage.actual_total_cost_usd = to_money_decimal(0)
usage.actual_request_cost_usd = to_money_decimal(0)
usage.status = "cancelled"
usage.status_code = status_code
usage.error_message = reason
@@ -313,28 +323,24 @@ class UsageLifecycleMixin:
response_time_ms: int | None = None,
billing_snapshot: dict[str, Any] | None = None,
extra_metadata: dict[str, Any] | None = None,
finalized_at: datetime | None = None,
) -> bool:
"""
写入异步任务最终账单(轮询完成后调用)。
语义:
- 正常路径:pending -> settled / void首次最终结算
- 补写路径:已写入 0 成本但尚未扣钱包的记录,可补写一次最终值
- 已 void 的记录不可再结算
- 已扣钱包wallet_balance_after 已存在)的记录不可重复扣费
- 仅允许 pending -> settled / void首次最终结算
- settled / void 一旦进入即不可再修改
约定:
- 不在本方法内 commit由调用方决定事务提交时机
"""
now = datetime.now(timezone.utc)
now = finalized_at or datetime.now(timezone.utc)
cost = to_money_decimal(total_cost_usd)
request_cost = to_money_decimal(request_cost_usd) if request_cost_usd is not None else cost
usage = db.query(Usage).filter(Usage.request_id == request_id).with_for_update().first()
if not usage or usage.billing_status == "void":
return False
if usage.billing_status == "settled" and usage.wallet_balance_after is not None:
if not usage or usage.billing_status != "pending":
return False
usage.total_cost_usd = cost
@@ -345,7 +351,7 @@ class UsageLifecycleMixin:
usage.error_message = error_message
if response_time_ms is not None:
usage.response_time_ms = response_time_ms
usage.finalized_at = usage.finalized_at or now
usage.finalized_at = now
if cost > 0:
usage.billing_status = "settled"
WalletService.apply_usage_charge(db, usage=usage, amount_usd=cost)
@@ -373,32 +379,15 @@ class UsageLifecycleMixin:
status_code: int = 499,
) -> bool:
"""
将已结算的记录作废(用于异步任务取消)
与 finalize_void 不同:
- finalize_void: pending -> void未结算时作废
- void_settled: settled -> void已结算后取消费用归零
约定:
- 仅当 billing_status='settled' 时才会生效
- 不在本方法内 commit由调用方决定事务提交时机
已废弃settled 为账务终态,不允许再回滚为 void
"""
now = datetime.now(timezone.utc)
usage = db.query(Usage).filter(Usage.request_id == request_id).with_for_update().first()
if not usage or usage.billing_status != "settled":
return False
if usage.wallet_balance_after is not None and to_money_decimal(usage.total_cost_usd) > 0:
# 已实际扣费的记录当前不做自动回滚,避免 silent inconsistency。
return False
usage.billing_status = "void"
usage.finalized_at = now
usage.total_cost_usd = to_money_decimal(0)
usage.request_cost_usd = to_money_decimal(0)
usage.status = "cancelled"
usage.status_code = status_code
usage.error_message = reason
return True
logger.warning(
"void_settled is deprecated and ignored: request_id={}, reason={}, status_code={}",
request_id,
reason,
status_code,
)
return False
@classmethod
def update_usage_status(
@@ -453,6 +442,15 @@ class UsageLifecycleMixin:
logger.warning("未找到 request_id={} 的使用记录,无法更新状态", request_id)
return None
if cls._is_billing_terminal(usage):
logger.debug(
"跳过已终态 Usage 状态更新: request_id={}, status={}, billing_status={}",
request_id,
status,
getattr(usage, "billing_status", None),
)
return usage
# 避免状态回退streaming 只能从 pending/streaming 进入
if status == "streaming" and usage.status not in ("pending", "streaming"):
logger.debug(
@@ -529,6 +527,10 @@ class UsageLifecycleMixin:
usage.billing_status = "void"
if getattr(usage, "finalized_at", None) is None:
usage.finalized_at = datetime.now(timezone.utc)
usage.total_cost_usd = to_money_decimal(0)
usage.request_cost_usd = to_money_decimal(0)
usage.actual_total_cost_usd = to_money_decimal(0)
usage.actual_request_cost_usd = to_money_decimal(0)
db.commit()

View File

@@ -26,6 +26,18 @@ from src.services.usage._types import UsageCostInfo, UsageRecordParams
from src.services.wallet import WalletService
def _coerce_finalized_at(value: Any, fallback: datetime | None = None) -> datetime | None:
if isinstance(value, datetime):
return value if value.tzinfo is not None else value.replace(tzinfo=timezone.utc)
if isinstance(value, str):
try:
parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
return parsed if parsed.tzinfo is not None else parsed.replace(tzinfo=timezone.utc)
except ValueError:
return fallback
return fallback
def _extract_manual_proxy_node_id(metadata: dict[str, Any] | None) -> str | None:
"""从 request_metadata 中提取手动代理节点 ID仅 is_manual 节点)。
@@ -131,10 +143,8 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
@staticmethod
def _is_usage_finalized(usage: Usage) -> bool:
return (
getattr(usage, "billing_status", None) in {"settled", "void"}
and getattr(usage, "finalized_at", None) is not None
)
# 与 UsageLifecycleMixin._is_billing_terminal 语义一致
return getattr(usage, "billing_status", None) in {"settled", "void"}
@classmethod
def _finalize_usage_billing(
@@ -157,10 +167,7 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
usage.billing_status = "pending"
return False, False
if (
getattr(usage, "billing_status", None) in {"settled", "void"}
and getattr(usage, "finalized_at", None) is not None
):
if getattr(usage, "billing_status", None) in {"settled", "void"}:
return False, False
now = finalized_at or datetime.now(timezone.utc)
@@ -221,6 +228,7 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
cache_ttl_minutes: int | None = None,
use_tiered_pricing: bool = True,
target_model: str | None = None,
finalized_at: datetime | None = None,
) -> Usage:
"""异步记录使用量(简化版,仅插入新记录)
@@ -310,6 +318,7 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
usage=usage,
total_cost=total_cost,
status=status,
finalized_at=finalized_at,
)
dispatch_codex_quota_sync_from_response_headers(
@@ -363,6 +372,7 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
cache_ttl_minutes: int | None = None,
use_tiered_pricing: bool = True,
target_model: str | None = None,
finalized_at: datetime | None = None,
) -> Usage:
"""记录使用量(完整版,支持更新已存在记录和用户统计)
@@ -464,6 +474,7 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
usage=usage,
total_cost=total_cost,
status=status,
finalized_at=finalized_at,
)
if accounted:
@@ -568,6 +579,7 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
provider_api_key_id: str | None = None,
status: str = "completed",
target_model: str | None = None,
finalized_at: datetime | None = None,
) -> Usage:
"""
记录"已计算好的"成本(用于 Video/Image/Audio 等异步任务的 FormulaEngine 计费结果)。
@@ -690,6 +702,7 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
usage=usage,
total_cost=total_cost,
status=status,
finalized_at=finalized_at,
)
if accounted:
@@ -961,7 +974,7 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
update_results = prepared_results[: len(update_params_list)]
insert_results = prepared_results[len(update_params_list) :]
finalized_at = datetime.now(timezone.utc)
batch_finalized_at = datetime.now(timezone.utc)
# 1. 处理需要更新的记录
for i, (record, request_id, params) in enumerate(update_params_list):
@@ -982,7 +995,9 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
usage=existing_usage,
total_cost=total_cost,
status=usage_params.get("status"),
finalized_at=finalized_at,
finalized_at=_coerce_finalized_at(
record.get("finalized_at"), batch_finalized_at
),
)
usages.append(existing_usage)
updated_count += 1
@@ -1043,7 +1058,9 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
usage=usage,
total_cost=total_cost,
status=usage_params.get("status"),
finalized_at=finalized_at,
finalized_at=_coerce_finalized_at(
record.get("finalized_at"), batch_finalized_at
),
)
usages.append(usage)
inserted_count += 1

View File

@@ -1,3 +1,12 @@
from src.services.wallet.daily_usage_ledger import (
WalletDailyUsageLedgerService,
WalletDailyUsageSnapshot,
)
from src.services.wallet.service import WalletAccessResult, WalletService
__all__ = ["WalletAccessResult", "WalletService"]
__all__ = [
"WalletAccessResult",
"WalletDailyUsageLedgerService",
"WalletDailyUsageSnapshot",
"WalletService",
]

View File

@@ -0,0 +1,204 @@
from __future__ import annotations
from dataclasses import dataclass
from datetime import date, datetime, time, timedelta, timezone
from decimal import Decimal
from zoneinfo import ZoneInfo
from sqlalchemy import func
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import Usage, WalletDailyUsageLedger
from src.services.billing.precision import to_money_decimal
from src.services.system.scheduler import APP_TIMEZONE
@dataclass(slots=True)
class WalletDailyUsageSnapshot:
wallet_id: str | None
billing_date: date
billing_timezone: str
total_cost_usd: Decimal
total_requests: int
input_tokens: int
output_tokens: int
cache_creation_tokens: int
cache_read_tokens: int
first_finalized_at: datetime | None
last_finalized_at: datetime | None
aggregated_at: datetime
is_today: bool = False
id: str | None = None
class WalletDailyUsageLedgerService:
"""钱包每日消费汇总服务。"""
@staticmethod
def get_timezone(timezone_name: str | None = None) -> ZoneInfo:
tz_name = timezone_name or APP_TIMEZONE
try:
return ZoneInfo(tz_name)
except Exception:
logger.warning("Invalid billing timezone {}, fallback to UTC", tz_name)
return ZoneInfo("UTC")
@classmethod
def get_today_billing_date(cls, timezone_name: str | None = None) -> date:
tz = cls.get_timezone(timezone_name)
return datetime.now(tz).date()
@classmethod
def get_day_window_utc(
cls,
billing_date: date,
timezone_name: str | None = None,
) -> tuple[datetime, datetime]:
tz = cls.get_timezone(timezone_name)
local_start = datetime.combine(billing_date, time.min, tzinfo=tz)
local_end = local_start + timedelta(days=1)
return local_start.astimezone(timezone.utc), local_end.astimezone(timezone.utc)
@classmethod
def aggregate_day(
cls,
db: Session,
billing_date: date,
*,
timezone_name: str | None = None,
commit: bool = True,
) -> int:
tz_name = timezone_name or APP_TIMEZONE
start_utc, end_utc = cls.get_day_window_utc(billing_date, tz_name)
now_utc = datetime.now(timezone.utc)
rows = (
db.query(
Usage.wallet_id.label("wallet_id"),
func.count(Usage.id).label("total_requests"),
func.coalesce(func.sum(Usage.total_cost_usd), 0).label("total_cost_usd"),
func.coalesce(func.sum(Usage.input_tokens), 0).label("input_tokens"),
func.coalesce(func.sum(Usage.output_tokens), 0).label("output_tokens"),
func.coalesce(func.sum(Usage.cache_creation_input_tokens), 0).label(
"cache_creation_tokens"
),
func.coalesce(func.sum(Usage.cache_read_input_tokens), 0).label(
"cache_read_tokens"
),
func.min(Usage.finalized_at).label("first_finalized_at"),
func.max(Usage.finalized_at).label("last_finalized_at"),
)
.filter(
Usage.wallet_id.isnot(None),
Usage.billing_status == "settled",
Usage.total_cost_usd > 0,
Usage.finalized_at >= start_utc,
Usage.finalized_at < end_utc,
)
.group_by(Usage.wallet_id)
.all()
)
existing_ledgers = (
db.query(WalletDailyUsageLedger)
.filter(
WalletDailyUsageLedger.billing_date == billing_date,
WalletDailyUsageLedger.billing_timezone == tz_name,
)
.all()
)
existing_map = {ledger.wallet_id: ledger for ledger in existing_ledgers}
seen_wallet_ids: set[str] = set()
for row in rows:
wallet_id = getattr(row, "wallet_id", None)
if not wallet_id:
continue
seen_wallet_ids.add(str(wallet_id))
ledger = existing_map.get(str(wallet_id))
if ledger is None:
ledger = WalletDailyUsageLedger(
wallet_id=str(wallet_id),
billing_date=billing_date,
billing_timezone=tz_name,
aggregated_at=now_utc,
)
db.add(ledger)
ledger.total_cost_usd = to_money_decimal(getattr(row, "total_cost_usd", 0) or 0)
ledger.total_requests = int(getattr(row, "total_requests", 0) or 0)
ledger.input_tokens = int(getattr(row, "input_tokens", 0) or 0)
ledger.output_tokens = int(getattr(row, "output_tokens", 0) or 0)
ledger.cache_creation_tokens = int(getattr(row, "cache_creation_tokens", 0) or 0)
ledger.cache_read_tokens = int(getattr(row, "cache_read_tokens", 0) or 0)
ledger.first_finalized_at = getattr(row, "first_finalized_at", None)
ledger.last_finalized_at = getattr(row, "last_finalized_at", None)
ledger.aggregated_at = now_utc
stale_ledgers = [
ledger for ledger in existing_ledgers if ledger.wallet_id not in seen_wallet_ids
]
for ledger in stale_ledgers:
db.delete(ledger)
if commit:
db.commit()
return len(seen_wallet_ids)
@classmethod
def get_today_snapshot(
cls,
db: Session,
wallet_id: str | None,
*,
timezone_name: str | None = None,
) -> WalletDailyUsageSnapshot:
tz_name = timezone_name or APP_TIMEZONE
billing_date = cls.get_today_billing_date(tz_name)
start_utc, end_utc = cls.get_day_window_utc(billing_date, tz_name)
now_utc = datetime.now(timezone.utc)
if wallet_id:
row = (
db.query(
func.count(Usage.id).label("total_requests"),
func.coalesce(func.sum(Usage.total_cost_usd), 0).label("total_cost_usd"),
func.coalesce(func.sum(Usage.input_tokens), 0).label("input_tokens"),
func.coalesce(func.sum(Usage.output_tokens), 0).label("output_tokens"),
func.coalesce(func.sum(Usage.cache_creation_input_tokens), 0).label(
"cache_creation_tokens"
),
func.coalesce(func.sum(Usage.cache_read_input_tokens), 0).label(
"cache_read_tokens"
),
func.min(Usage.finalized_at).label("first_finalized_at"),
func.max(Usage.finalized_at).label("last_finalized_at"),
)
.filter(
Usage.wallet_id == wallet_id,
Usage.billing_status == "settled",
Usage.total_cost_usd > 0,
Usage.finalized_at >= start_utc,
Usage.finalized_at < end_utc,
)
.first()
)
else:
row = None
return WalletDailyUsageSnapshot(
wallet_id=wallet_id,
billing_date=billing_date,
billing_timezone=tz_name,
total_cost_usd=to_money_decimal(getattr(row, "total_cost_usd", 0) or 0),
total_requests=int(getattr(row, "total_requests", 0) or 0),
input_tokens=int(getattr(row, "input_tokens", 0) or 0),
output_tokens=int(getattr(row, "output_tokens", 0) or 0),
cache_creation_tokens=int(getattr(row, "cache_creation_tokens", 0) or 0),
cache_read_tokens=int(getattr(row, "cache_read_tokens", 0) or 0),
first_finalized_at=getattr(row, "first_finalized_at", None),
last_finalized_at=getattr(row, "last_finalized_at", None),
aggregated_at=now_utc,
is_today=True,
)