feat(pool): 记录并展示 Provider Key 累计 Token 与费用

Closes #219

Co-authored-by: AAEE86 <ppk0227@hotmail.com>
This commit is contained in:
fawney19
2026-03-11 15:56:58 +08:00
parent 02e2f4f500
commit 380d69e096
10 changed files with 452 additions and 16 deletions

View File

@@ -30,6 +30,7 @@ from src.core.exceptions import NotFoundException
from src.core.logger import logger
from src.database import get_db
from src.models.database import Provider, ProviderAPIKey
from src.services.billing.precision import to_money_decimal
from src.services.provider.fingerprint import generate_fingerprint
from src.services.provider.pool import redis_ops as pool_redis
from src.services.provider.pool.account_state import resolve_pool_account_state
@@ -216,6 +217,10 @@ def _to_float(value: Any) -> float | None:
return None
def _serialize_money(value: Any) -> str:
return format(to_money_decimal(value), "f")
def _is_known_banned_key(key: ProviderAPIKey, provider_type: str) -> bool:
from src.services.provider.pool.account_state import resolve_pool_account_state
@@ -677,6 +682,8 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
ProviderAPIKey.health_by_format,
ProviderAPIKey.circuit_breaker_by_format,
ProviderAPIKey.request_count,
ProviderAPIKey.total_tokens,
ProviderAPIKey.total_cost_usd,
ProviderAPIKey.last_used_at,
ProviderAPIKey.created_at,
ProviderAPIKey.upstream_metadata,
@@ -864,6 +871,8 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
else []
)
key_request_count = int(getattr(k, "request_count", 0) or 0)
key_total_tokens = int(getattr(k, "total_tokens", 0) or 0)
key_total_cost_usd = _serialize_money(getattr(k, "total_cost_usd", 0.0))
key_last_used_at = getattr(k, "last_used_at", None)
oauth_auth_config = _extract_oauth_auth_config(k)
@@ -923,6 +932,8 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
cost_window_usage=cost_usage,
cost_limit=cost_limit,
request_count=key_request_count,
total_tokens=key_total_tokens,
total_cost_usd=key_total_cost_usd,
sticky_sessions=sticky_counts.get(kid, 0),
lru_score=lru_scores.get(kid),
created_at=(

View File

@@ -103,6 +103,8 @@ class PoolKeyDetail(BaseModel):
cost_window_usage: int = 0
cost_limit: int | None = None
request_count: int = 0
total_tokens: int = 0
total_cost_usd: str = "0.00000000"
sticky_sessions: int = 0
lru_score: float | None = None
created_at: str | None = None

View File

@@ -2869,6 +2869,8 @@ def _purge_stats_and_reset_counters(db: Session) -> None:
db.query(ProviderAPIKey).update(
{
ProviderAPIKey.request_count: 0,
ProviderAPIKey.total_tokens: 0,
ProviderAPIKey.total_cost_usd: 0.0,
ProviderAPIKey.success_count: 0,
ProviderAPIKey.error_count: 0,
ProviderAPIKey.total_response_time_ms: 0,

View File

@@ -1808,6 +1808,8 @@ class ProviderAPIKey(ExportMixin, Base):
"health_by_format",
"circuit_breaker_by_format",
"request_count",
"total_tokens",
"total_cost_usd",
"success_count",
"error_count",
"total_response_time_ms",
@@ -1913,6 +1915,8 @@ class ProviderAPIKey(ExportMixin, Base):
# 使用统计
request_count = Column(Integer, default=0) # 请求次数
total_tokens = Column(BigInteger, default=0, nullable=False) # 累计 Token 数
total_cost_usd = Column(Numeric(20, 8), default=0.0, nullable=False) # 累计成本
success_count = Column(Integer, default=0) # 成功次数
error_count = Column(Integer, default=0) # 错误次数
total_response_time_ms = Column(Integer, default=0) # 总响应时间(用于计算平均值)

View File

@@ -9,7 +9,15 @@ from typing import Any
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey, Provider, ProxyNode, Usage, User, UserModelUsageCount
from src.models.database import (
ApiKey,
Provider,
ProviderAPIKey,
ProxyNode,
Usage,
User,
UserModelUsageCount,
)
from src.services.billing.precision import to_money_decimal
from src.services.provider_keys.codex_quota_sync_dispatcher import (
dispatch_codex_quota_sync_from_response_headers,
@@ -82,6 +90,40 @@ def _increment_proxy_node_requests(
)
def _increment_provider_api_key_totals(
db: Session,
provider_api_key_id: str | None,
*,
total_tokens: int = 0,
total_cost: float = 0.0,
) -> None:
"""原子递增 ProviderAPIKey 的累计 Token/成本。"""
if not provider_api_key_id:
return
token_increment = int(total_tokens or 0)
cost_increment = to_money_decimal(total_cost)
if token_increment <= 0 and cost_increment <= 0:
return
from sqlalchemy import update
values: dict[str, Any] = {}
if token_increment > 0:
values["total_tokens"] = ProviderAPIKey.total_tokens + token_increment
if cost_increment > 0:
values["total_cost_usd"] = ProviderAPIKey.total_cost_usd + Decimal(str(cost_increment))
db.execute(
update(ProviderAPIKey).where(ProviderAPIKey.id == provider_api_key_id).values(**values)
)
def _get_actual_total_cost_usd(usage_params: dict[str, Any]) -> float:
return float(to_money_decimal(usage_params.get("actual_total_cost_usd") or 0.0))
class UsageRecordingMixin(UsageBillingIntegrationMixin):
"""记录用量相关方法"""
@@ -313,7 +355,7 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
)
cls._finalize_usage_billing(
accounted, _charge_applied = cls._finalize_usage_billing(
db,
usage=usage,
total_cost=total_cost,
@@ -321,6 +363,14 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
finalized_at=finalized_at,
)
if accounted:
_increment_provider_api_key_totals(
db,
provider_api_key_id,
total_tokens=int(usage_params.get("total_tokens") or 0),
total_cost=_get_actual_total_cost_usd(usage_params),
)
dispatch_codex_quota_sync_from_response_headers(
provider_api_key_id=provider_api_key_id,
response_headers=response_headers,
@@ -493,6 +543,13 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
sa_update(ApiKeyModel).where(ApiKeyModel.id == api_key.id).values(**values)
)
_increment_provider_api_key_totals(
db,
provider_api_key_id,
total_tokens=int(usage_params.get("total_tokens") or 0),
total_cost=_get_actual_total_cost_usd(usage_params),
)
# 更新 GlobalModel 使用计数
db.execute(
sa_update(GlobalModel)
@@ -719,6 +776,13 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
)
db.execute(update(ApiKeyModel).where(ApiKeyModel.id == api_key.id).values(**values))
_increment_provider_api_key_totals(
db,
provider_api_key_id,
total_tokens=int(usage_params.get("total_tokens") or 0),
total_cost=_get_actual_total_cost_usd(usage_params),
)
# 更新 GlobalModel 使用计数
db.execute(
update(GlobalModel)
@@ -855,6 +919,9 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
apikey_stats: dict[str, dict[str, Any]] = defaultdict(
lambda: {"requests": 0, "cost": 0.0, "is_standalone": False}
)
provider_key_stats: dict[str, dict[str, Any]] = defaultdict(
lambda: {"tokens": 0, "actual_cost": 0.0}
)
model_counts: dict[str, int] = defaultdict(int) # model -> count
user_model_counts: dict[tuple[str, str], int] = defaultdict(
int
@@ -1021,6 +1088,15 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
apikey_stats[key_id]["cost"] += total_cost
apikey_stats[key_id]["is_standalone"] = api_key.is_standalone
provider_api_key_id = record.get("provider_api_key_id")
if isinstance(provider_api_key_id, str) and provider_api_key_id:
provider_key_stats[provider_api_key_id]["tokens"] += int(
usage_params.get("total_tokens") or 0
)
provider_key_stats[provider_api_key_id][
"actual_cost"
] += _get_actual_total_cost_usd(usage_params)
manual_nid = _extract_manual_proxy_node_id(record.get("metadata"))
if manual_nid:
proxy_node_counts[manual_nid] += 1
@@ -1085,6 +1161,15 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
apikey_stats[key_id]["cost"] += total_cost
apikey_stats[key_id]["is_standalone"] = api_key.is_standalone
provider_api_key_id = record.get("provider_api_key_id")
if isinstance(provider_api_key_id, str) and provider_api_key_id:
provider_key_stats[provider_api_key_id]["tokens"] += int(
usage_params.get("total_tokens") or 0
)
provider_key_stats[provider_api_key_id][
"actual_cost"
] += _get_actual_total_cost_usd(usage_params)
manual_nid = _extract_manual_proxy_node_id(record.get("metadata"))
if manual_nid:
proxy_node_counts[manual_nid] += 1
@@ -1173,6 +1258,14 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
)
)
for provider_key_id, stats in provider_key_stats.items():
_increment_provider_api_key_totals(
db,
provider_key_id,
total_tokens=int(stats["tokens"]),
total_cost=float(stats["actual_cost"]),
)
# 批量更新手动代理节点请求计数
_increment_proxy_node_requests(db, proxy_node_counts, proxy_node_failed)