mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat(pool): 记录并展示 Provider Key 累计 Token 与费用
Closes #219 Co-authored-by: AAEE86 <ppk0227@hotmail.com>
This commit is contained in:
@@ -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=(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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) # 总响应时间(用于计算平均值)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user