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

@@ -0,0 +1,59 @@
"""add provider_api_keys usage total columns
Revision ID: 9b7c6d5e4f3a
Revises: d4e5f6a7b8c9
Create Date: 2026-03-11 22:00:00.000000+00:00
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "9b7c6d5e4f3a"
down_revision: str | None = "d4e5f6a7b8c9"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [c["name"] for c in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
if not column_exists("provider_api_keys", "total_tokens"):
op.add_column(
"provider_api_keys",
sa.Column("total_tokens", sa.BigInteger(), nullable=False, server_default="0"),
)
if column_exists("provider_api_keys", "total_tokens"):
op.alter_column("provider_api_keys", "total_tokens", server_default=None)
if not column_exists("provider_api_keys", "total_cost_usd"):
op.add_column(
"provider_api_keys",
sa.Column(
"total_cost_usd",
sa.Numeric(20, 8),
nullable=False,
server_default="0.0",
),
)
if column_exists("provider_api_keys", "total_cost_usd"):
op.alter_column("provider_api_keys", "total_cost_usd", server_default=None)
def downgrade() -> None:
if column_exists("provider_api_keys", "total_cost_usd"):
op.drop_column("provider_api_keys", "total_cost_usd")
if column_exists("provider_api_keys", "total_tokens"):
op.drop_column("provider_api_keys", "total_tokens")

View File

@@ -29,7 +29,7 @@ export interface PoolStatusResponse {
* 获取 Provider 的号池状态 * 获取 Provider 的号池状态
*/ */
export async function getPoolStatus(providerId: string): Promise<PoolStatusResponse> { export async function getPoolStatus(providerId: string): Promise<PoolStatusResponse> {
const response = await client.get(`/api/admin/providers/${providerId}/pool-status`) const response = await client.get<PoolStatusResponse>(`/api/admin/providers/${providerId}/pool-status`)
return response.data return response.data
} }
@@ -40,7 +40,7 @@ export async function clearPoolCooldown(
providerId: string, providerId: string,
keyId: string, keyId: string,
): Promise<{ message: string }> { ): Promise<{ message: string }> {
const response = await client.post( const response = await client.post<{ message: string }>(
`/api/admin/providers/${providerId}/pool/clear-cooldown/${keyId}`, `/api/admin/providers/${providerId}/pool/clear-cooldown/${keyId}`,
) )
return response.data return response.data
@@ -53,7 +53,7 @@ export async function resetPoolCost(
providerId: string, providerId: string,
keyId: string, keyId: string,
): Promise<{ message: string }> { ): Promise<{ message: string }> {
const response = await client.post( const response = await client.post<{ message: string }>(
`/api/admin/providers/${providerId}/pool/reset-cost/${keyId}`, `/api/admin/providers/${providerId}/pool/reset-cost/${keyId}`,
) )
return response.data return response.data
@@ -125,6 +125,8 @@ export interface PoolKeyDetail {
cost_window_usage: number cost_window_usage: number
cost_limit: number | null cost_limit: number | null
request_count: number request_count: number
total_tokens: number
total_cost_usd: string
sticky_sessions: number sticky_sessions: number
lru_score: number | null lru_score: number | null
created_at: string | null created_at: string | null

View File

@@ -543,13 +543,25 @@
>-</span> >-</span>
</TableCell> </TableCell>
<TableCell class="py-3 px-2 align-middle"> <TableCell class="py-3 px-2 align-middle">
<div class="w-[136px] mx-auto text-[10px] leading-4"> <div class="grid grid-rows-3 gap-0.5 w-[136px] mx-auto text-[10px] leading-4">
<div class="flex items-center justify-between gap-2"> <div class="flex items-center justify-between gap-2">
<span class="text-muted-foreground">请求</span> <span class="text-muted-foreground">请求</span>
<span class="tabular-nums text-foreground/90"> <span class="tabular-nums text-foreground/90">
{{ formatStatInteger(key.request_count) }} {{ formatStatInteger(key.request_count) }}
</span> </span>
</div> </div>
<div class="flex items-center justify-between gap-2">
<span class="text-muted-foreground">Token</span>
<span class="tabular-nums text-foreground/90">
{{ formatTokenCount(key.total_tokens) }}
</span>
</div>
<div class="flex items-center justify-between gap-2">
<span class="text-muted-foreground">费用</span>
<span class="tabular-nums text-foreground/90">
{{ formatStatUsd(key.total_cost_usd) }}
</span>
</div>
</div> </div>
</TableCell> </TableCell>
<TableCell class="py-3 text-center"> <TableCell class="py-3 text-center">
@@ -912,11 +924,19 @@
<div class="text-muted-foreground mb-0.5"> <div class="text-muted-foreground mb-0.5">
统计 统计
</div> </div>
<div class="text-[10px]"> <div class="space-y-0.5 text-[10px]">
<div class="flex items-center justify-between gap-2"> <div class="flex items-center justify-between gap-2">
<span class="text-muted-foreground">请求</span> <span class="text-muted-foreground">请求</span>
<span class="tabular-nums">{{ formatStatInteger(key.request_count) }}</span> <span class="tabular-nums">{{ formatStatInteger(key.request_count) }}</span>
</div> </div>
<div class="flex items-center justify-between gap-2">
<span class="text-muted-foreground">Token</span>
<span class="tabular-nums">{{ formatTokenCount(key.total_tokens) }}</span>
</div>
<div class="flex items-center justify-between gap-2">
<span class="text-muted-foreground">费用</span>
<span class="tabular-nums">{{ formatStatUsd(key.total_cost_usd) }}</span>
</div>
</div> </div>
</div> </div>
<div <div
@@ -1195,6 +1215,7 @@ async function loadOverview() {
selectedProviderData.value = null selectedProviderData.value = null
showAccountBatchDialog.value = false showAccountBatchDialog.value = false
closeProviderProxyPopovers() closeProviderProxyPopovers()
resetKeyPage()
} }
} }
} catch (err) { } catch (err) {
@@ -1382,6 +1403,7 @@ const desktopColumnWidths = computed(() => {
async function selectProvider(id: string) { async function selectProvider(id: string) {
const requestId = ++selectProviderRequestId const requestId = ++selectProviderRequestId
selectedProviderId.value = id selectedProviderId.value = id
selectedProviderData.value = null
editingKeyDetail.value = null editingKeyDetail.value = null
showAccountBatchDialog.value = false showAccountBatchDialog.value = false
keyPermissionsDialogOpen.value = false keyPermissionsDialogOpen.value = false
@@ -1399,6 +1421,7 @@ async function selectProvider(id: string) {
clearTimeout(keysSearchDebounceTimer) clearTimeout(keysSearchDebounceTimer)
keysSearchDebounceTimer = null keysSearchDebounceTimer = null
} }
resetKeyPage(1, pageSize.value)
const keysTask = loadKeys() const keysTask = loadKeys()
// Provider summary is non-blocking for key list rendering. // Provider summary is non-blocking for key list rendering.
void loadProviderData(id) void loadProviderData(id)
@@ -1423,7 +1446,11 @@ async function refresh() {
} }
// --- Keys --- // --- Keys ---
const keyPage = ref<PoolKeysPageResponse>({ total: 0, page: 1, page_size: 50, keys: [] }) function createEmptyKeyPage(page = 1, pageSizeValue = 50): PoolKeysPageResponse {
return { total: 0, page, page_size: pageSizeValue, keys: [] }
}
const keyPage = ref<PoolKeysPageResponse>(createEmptyKeyPage())
const keysLoading = ref(false) const keysLoading = ref(false)
const refreshingCurrentPageQuota = ref(false) const refreshingCurrentPageQuota = ref(false)
const searchQuery = ref('') const searchQuery = ref('')
@@ -1432,7 +1459,6 @@ const currentPage = ref(1)
const pageSize = ref(50) const pageSize = ref(50)
const MANUAL_QUOTA_REFRESH_COOLDOWN_SECONDS = 5 * 60 const MANUAL_QUOTA_REFRESH_COOLDOWN_SECONDS = 5 * 60
const refreshingOAuthKeyId = ref<string | null>(null) const refreshingOAuthKeyId = ref<string | null>(null)
const revealedKeys = ref<Map<string, string>>(new Map())
const recoveringHealthKeyId = ref<string | null>(null) const recoveringHealthKeyId = ref<string | null>(null)
const savingProxyKeyId = ref<string | null>(null) const savingProxyKeyId = ref<string | null>(null)
const proxyDesktopPopoverOpenKeyId = ref<string | null>(null) const proxyDesktopPopoverOpenKeyId = ref<string | null>(null)
@@ -1473,6 +1499,14 @@ const refreshCurrentPageLoading = computed(() => {
return keysLoading.value || refreshingCurrentPageQuota.value return keysLoading.value || refreshingCurrentPageQuota.value
}) })
function resetKeyPage(page = currentPage.value, pageSizeValue = pageSize.value): void {
keyPage.value = createEmptyKeyPage(page, pageSizeValue)
}
function refreshOverviewInBackground(): void {
void loadOverview()
}
function normalizeQuotaUpdatedAt(raw: number | null | undefined): number | null { function normalizeQuotaUpdatedAt(raw: number | null | undefined): number | null {
const value = Number(raw ?? 0) const value = Number(raw ?? 0)
if (!Number.isFinite(value) || value <= 0) return null if (!Number.isFinite(value) || value <= 0) return null
@@ -1608,6 +1642,7 @@ async function loadKeys() {
keyPage.value = nextPage keyPage.value = nextPage
} catch (err) { } catch (err) {
if (requestId !== keysRequestId || selectedProviderId.value !== providerId) return if (requestId !== keysRequestId || selectedProviderId.value !== providerId) return
resetKeyPage(page, pageSizeValue)
showError(parseApiError(err)) showError(parseApiError(err))
} finally { } finally {
if (requestId === keysRequestId) { if (requestId === keysRequestId) {
@@ -1879,6 +1914,7 @@ async function handleDeleteKey(key: PoolKeyDetail) {
if (keyPage.value.keys.length === 0 && currentPage.value > 1) { if (keyPage.value.keys.length === 0 && currentPage.value > 1) {
currentPage.value-- currentPage.value--
} }
refreshOverviewInBackground()
} catch (err) { } catch (err) {
showError(parseApiError(err, '删除账号失败')) showError(parseApiError(err, '删除账号失败'))
} finally { } finally {
@@ -1887,12 +1923,6 @@ async function handleDeleteKey(key: PoolKeyDetail) {
} }
async function copyFullKey(key: PoolKeyDetail) { async function copyFullKey(key: PoolKeyDetail) {
const cached = revealedKeys.value.get(key.key_id)
if (cached) {
await copyToClipboard(cached)
return
}
try { try {
const result = await revealEndpointKey(key.key_id) const result = await revealEndpointKey(key.key_id)
let textToCopy = '' let textToCopy = ''
@@ -1912,7 +1942,6 @@ async function copyFullKey(key: PoolKeyDetail) {
return return
} }
revealedKeys.value.set(key.key_id, textToCopy)
await copyToClipboard(textToCopy) await copyToClipboard(textToCopy)
} catch (err) { } catch (err) {
showError(parseApiError(err, '获取密钥失败')) showError(parseApiError(err, '获取密钥失败'))
@@ -1969,6 +1998,7 @@ async function clearCooldown(keyId: string) {
const res = await clearPoolCooldown(selectedProviderId.value, keyId) const res = await clearPoolCooldown(selectedProviderId.value, keyId)
success(res.message) success(res.message)
await loadKeys() await loadKeys()
refreshOverviewInBackground()
} catch (err) { } catch (err) {
showError(parseApiError(err)) showError(parseApiError(err))
} }
@@ -1994,6 +2024,7 @@ async function toggleKeyActive(key: PoolKeyDetail) {
} }
success(nextStatus ? '账号已启用' : '账号已停用') success(nextStatus ? '账号已启用' : '账号已停用')
await loadKeys() await loadKeys()
refreshOverviewInBackground()
} catch (err) { } catch (err) {
showError(parseApiError(err)) showError(parseApiError(err))
} finally { } finally {
@@ -2496,6 +2527,23 @@ function formatStatInteger(value: number | null | undefined): string {
return Math.round(n).toLocaleString('en-US') return Math.round(n).toLocaleString('en-US')
} }
function formatTokenCount(value: number | null | undefined): string {
const n = Number(value ?? 0)
if (!Number.isFinite(n) || n <= 0) return '0'
if (n >= 1_000_000) return `${(n / 1_000_000).toFixed(1)}M`
if (n >= 1_000) return `${(n / 1_000).toFixed(1)}K`
return String(Math.round(n))
}
function formatStatUsd(value: number | string | null | undefined): string {
const n = Number(value ?? 0)
if (!Number.isFinite(n) || n <= 0) return '$0.00'
if (n < 0.01) return `$${n.toFixed(4)}`
if (n < 1) return `$${n.toFixed(3)}`
if (n < 1000) return `$${n.toFixed(2)}`
return `$${n.toLocaleString('en-US', { minimumFractionDigits: 2, maximumFractionDigits: 2 })}`
}
function formatRelativeTime(isoStr: string): string { function formatRelativeTime(isoStr: string): string {
const date = new Date(isoStr) const date = new Date(isoStr)
const pad = (n: number) => String(n).padStart(2, '0') const pad = (n: number) => String(n).padStart(2, '0')

View File

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

View File

@@ -103,6 +103,8 @@ class PoolKeyDetail(BaseModel):
cost_window_usage: int = 0 cost_window_usage: int = 0
cost_limit: int | None = None cost_limit: int | None = None
request_count: int = 0 request_count: int = 0
total_tokens: int = 0
total_cost_usd: str = "0.00000000"
sticky_sessions: int = 0 sticky_sessions: int = 0
lru_score: float | None = None lru_score: float | None = None
created_at: str | 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( db.query(ProviderAPIKey).update(
{ {
ProviderAPIKey.request_count: 0, ProviderAPIKey.request_count: 0,
ProviderAPIKey.total_tokens: 0,
ProviderAPIKey.total_cost_usd: 0.0,
ProviderAPIKey.success_count: 0, ProviderAPIKey.success_count: 0,
ProviderAPIKey.error_count: 0, ProviderAPIKey.error_count: 0,
ProviderAPIKey.total_response_time_ms: 0, ProviderAPIKey.total_response_time_ms: 0,

View File

@@ -1808,6 +1808,8 @@ class ProviderAPIKey(ExportMixin, Base):
"health_by_format", "health_by_format",
"circuit_breaker_by_format", "circuit_breaker_by_format",
"request_count", "request_count",
"total_tokens",
"total_cost_usd",
"success_count", "success_count",
"error_count", "error_count",
"total_response_time_ms", "total_response_time_ms",
@@ -1913,6 +1915,8 @@ class ProviderAPIKey(ExportMixin, Base):
# 使用统计 # 使用统计
request_count = Column(Integer, default=0) # 请求次数 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) # 成功次数 success_count = Column(Integer, default=0) # 成功次数
error_count = Column(Integer, default=0) # 错误次数 error_count = Column(Integer, default=0) # 错误次数
total_response_time_ms = 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 sqlalchemy.orm import Session
from src.core.logger import logger 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.billing.precision import to_money_decimal
from src.services.provider_keys.codex_quota_sync_dispatcher import ( from src.services.provider_keys.codex_quota_sync_dispatcher import (
dispatch_codex_quota_sync_from_response_headers, 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): class UsageRecordingMixin(UsageBillingIntegrationMixin):
"""记录用量相关方法""" """记录用量相关方法"""
@@ -313,7 +355,7 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost) .values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
) )
cls._finalize_usage_billing( accounted, _charge_applied = cls._finalize_usage_billing(
db, db,
usage=usage, usage=usage,
total_cost=total_cost, total_cost=total_cost,
@@ -321,6 +363,14 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
finalized_at=finalized_at, 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( dispatch_codex_quota_sync_from_response_headers(
provider_api_key_id=provider_api_key_id, provider_api_key_id=provider_api_key_id,
response_headers=response_headers, response_headers=response_headers,
@@ -493,6 +543,13 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
sa_update(ApiKeyModel).where(ApiKeyModel.id == api_key.id).values(**values) 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 使用计数 # 更新 GlobalModel 使用计数
db.execute( db.execute(
sa_update(GlobalModel) sa_update(GlobalModel)
@@ -719,6 +776,13 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
) )
db.execute(update(ApiKeyModel).where(ApiKeyModel.id == api_key.id).values(**values)) 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 使用计数 # 更新 GlobalModel 使用计数
db.execute( db.execute(
update(GlobalModel) update(GlobalModel)
@@ -855,6 +919,9 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
apikey_stats: dict[str, dict[str, Any]] = defaultdict( apikey_stats: dict[str, dict[str, Any]] = defaultdict(
lambda: {"requests": 0, "cost": 0.0, "is_standalone": False} 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 model_counts: dict[str, int] = defaultdict(int) # model -> count
user_model_counts: dict[tuple[str, str], int] = defaultdict( user_model_counts: dict[tuple[str, str], int] = defaultdict(
int int
@@ -1021,6 +1088,15 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
apikey_stats[key_id]["cost"] += total_cost apikey_stats[key_id]["cost"] += total_cost
apikey_stats[key_id]["is_standalone"] = api_key.is_standalone 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")) manual_nid = _extract_manual_proxy_node_id(record.get("metadata"))
if manual_nid: if manual_nid:
proxy_node_counts[manual_nid] += 1 proxy_node_counts[manual_nid] += 1
@@ -1085,6 +1161,15 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
apikey_stats[key_id]["cost"] += total_cost apikey_stats[key_id]["cost"] += total_cost
apikey_stats[key_id]["is_standalone"] = api_key.is_standalone 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")) manual_nid = _extract_manual_proxy_node_id(record.get("metadata"))
if manual_nid: if manual_nid:
proxy_node_counts[manual_nid] += 1 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) _increment_proxy_node_requests(db, proxy_node_counts, proxy_node_failed)

View File

@@ -0,0 +1,204 @@
from __future__ import annotations
from typing import Any
from unittest.mock import AsyncMock, MagicMock
import pytest
import src.services.usage.recording as recording_module
from src.services.usage.service import UsageService
class DummyQuery:
def __init__(self, result: list[Any]) -> None:
self._result = result
def options(self, *args: Any, **kwargs: Any) -> "DummyQuery":
return self
def filter(self, *args: Any, **kwargs: Any) -> "DummyQuery":
return self
def with_for_update(self) -> "DummyQuery":
return self
def all(self) -> list[Any]:
return self._result
def first(self) -> Any | None:
return self._result[0] if self._result else None
@pytest.mark.asyncio
async def test_record_usage_updates_provider_key_totals(monkeypatch: Any) -> None:
db = MagicMock()
db.query.side_effect = lambda _model: DummyQuery([])
usage_params = {
"request_id": "req-provider-key-1",
"provider_name": "openai",
"model": "gpt-4o",
"status": "completed",
"total_tokens": 123,
"actual_total_cost_usd": 1.75,
}
monkeypatch.setattr(
UsageService,
"_prepare_usage_record",
AsyncMock(return_value=(usage_params, 1.25)),
)
monkeypatch.setattr(
UsageService,
"_finalize_usage_billing",
MagicMock(return_value=(True, True)),
)
helper = MagicMock()
monkeypatch.setattr(recording_module, "_increment_provider_api_key_totals", helper)
monkeypatch.setattr(
recording_module,
"dispatch_codex_quota_sync_from_response_headers",
MagicMock(),
)
await UsageService.record_usage(
db=db,
user=None,
api_key=None,
provider="openai",
model="gpt-4o",
input_tokens=100,
output_tokens=23,
provider_api_key_id="provider-key-1",
request_id="req-provider-key-1",
status="completed",
)
helper.assert_called_once()
_, provider_key_id = helper.call_args.args
assert provider_key_id == "provider-key-1"
assert helper.call_args.kwargs["total_tokens"] == 123
assert float(helper.call_args.kwargs["total_cost"]) == 1.75
@pytest.mark.asyncio
async def test_record_usage_async_updates_provider_key_totals(monkeypatch: Any) -> None:
db = MagicMock()
usage_params = {
"request_id": "req-provider-key-async",
"provider_name": "openai",
"model": "gpt-4o-mini",
"status": "completed",
"total_tokens": 77,
"actual_total_cost_usd": 0.75,
}
monkeypatch.setattr(
UsageService,
"_prepare_usage_record",
AsyncMock(return_value=(usage_params, 0.5)),
)
monkeypatch.setattr(
UsageService,
"_finalize_usage_billing",
MagicMock(return_value=(True, False)),
)
helper = MagicMock()
monkeypatch.setattr(recording_module, "_increment_provider_api_key_totals", helper)
monkeypatch.setattr(
recording_module,
"dispatch_codex_quota_sync_from_response_headers",
MagicMock(),
)
await UsageService.record_usage_async(
db=db,
user=None,
api_key=None,
provider="openai",
model="gpt-4o-mini",
input_tokens=50,
output_tokens=27,
provider_api_key_id="provider-key-async",
request_id="req-provider-key-async",
status="completed",
)
helper.assert_called_once()
_, provider_key_id = helper.call_args.args
assert provider_key_id == "provider-key-async"
assert helper.call_args.kwargs["total_tokens"] == 77
assert helper.call_args.kwargs["total_cost"] == 0.75
@pytest.mark.asyncio
async def test_record_usage_batch_aggregates_provider_key_totals(monkeypatch: Any) -> None:
db = MagicMock()
db.query.side_effect = lambda _model: DummyQuery([])
usage_params_1 = {
"request_id": "req-provider-key-batch-1",
"provider_name": "anthropic",
"model": "claude-sonnet",
"status": "completed",
"total_tokens": 321,
"actual_total_cost_usd": 2.5,
}
usage_params_2 = {
"request_id": "req-provider-key-batch-2",
"provider_name": "anthropic",
"model": "claude-sonnet",
"status": "completed",
"total_tokens": 79,
"actual_total_cost_usd": 0.75,
}
monkeypatch.setattr(
UsageService,
"_prepare_usage_records_batch",
AsyncMock(
return_value=[
(usage_params_1, 2.0, None),
(usage_params_2, 0.5, None),
]
),
)
monkeypatch.setattr(
UsageService,
"_finalize_usage_billing",
MagicMock(return_value=(True, True)),
)
helper = MagicMock()
monkeypatch.setattr(recording_module, "_increment_provider_api_key_totals", helper)
monkeypatch.setattr(
recording_module,
"dispatch_codex_quota_sync_from_response_headers",
MagicMock(),
)
await UsageService.record_usage_batch(
db,
[
{
"request_id": "req-provider-key-batch-1",
"provider": "anthropic",
"model": "claude-sonnet",
"status": "completed",
"provider_api_key_id": "provider-key-batch",
},
{
"request_id": "req-provider-key-batch-2",
"provider": "anthropic",
"model": "claude-sonnet",
"status": "completed",
"provider_api_key_id": "provider-key-batch",
},
],
)
helper.assert_called_once()
_, provider_key_id = helper.call_args.args
assert provider_key_id == "provider-key-batch"
assert helper.call_args.kwargs["total_tokens"] == 400
assert helper.call_args.kwargs["total_cost"] == 3.25

View File

@@ -0,0 +1,11 @@
from __future__ import annotations
from decimal import Decimal
from src.api.admin.pool.routes import _serialize_money
def test_serialize_money_preserves_storage_precision() -> None:
assert _serialize_money(Decimal("12.34")) == "12.34000000"
assert _serialize_money("0.00000001") == "0.00000001"
assert _serialize_money(None) == "0.00000000"