mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
feat(usage): 修复缓存命中率计算并新增用户端 API 格式统计
- 新增 input_context_expr() 按 api_format 区分 input_tokens 语义 (OpenAI/Gemini input_tokens 已含 cache_read,Claude 需额外加上) - 缓存命中率统一改为基于归一化后的 total_input_context 计算 - 用户 /me/usage 接口新增 summary_by_api_format 后端聚合字段 - 前端 API 格式统计改用后端聚合数据,移除前端逐条记录手动统计 - 提取 formatHitRate 到 utils/format.ts 消除三处重复定义 - 移除 PoolManager 中未使用的 select_key 方法 Co-Authored-By: AAEE86 <ppk0227@hotmail.com>
This commit is contained in:
@@ -29,6 +29,7 @@ from src.models.database import (
|
||||
)
|
||||
from src.services.system.stats_aggregator import AggregatedStats, StatsFilter, query_stats_hybrid
|
||||
from src.services.system.time_range import TimeRangeParams
|
||||
from src.services.usage.query import input_context_expr
|
||||
from src.services.usage.service import UsageService
|
||||
from src.utils.cache_decorator import cache_result
|
||||
|
||||
@@ -78,6 +79,25 @@ def _build_time_range_params(
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
|
||||
def _calculate_token_cache_hit_rate(
|
||||
total_input_context: int | None,
|
||||
cache_read_tokens: int | None,
|
||||
) -> float:
|
||||
"""计算缓存命中率。
|
||||
|
||||
Args:
|
||||
total_input_context: 已归一化的总输入上下文 token 数。
|
||||
Claude 格式: input_tokens + cache_read_input_tokens
|
||||
OpenAI/Gemini 格式: input_tokens(已包含 cache_read)
|
||||
cache_read_tokens: 缓存读取 token 数
|
||||
"""
|
||||
context = int(total_input_context or 0)
|
||||
cached = int(cache_read_tokens or 0)
|
||||
if context <= 0:
|
||||
return 0.0
|
||||
return round(cached / context * 100, 2)
|
||||
|
||||
|
||||
# ==================== RESTful Routes ====================
|
||||
|
||||
|
||||
@@ -107,10 +127,10 @@ async def get_usage_aggregation(
|
||||
- `limit`: 返回数量限制,默认 20,最大 100
|
||||
|
||||
**返回字段**:
|
||||
- 按模型聚合时:model, request_count, total_tokens, total_cost, actual_cost
|
||||
- 按模型聚合时:model, request_count, total_tokens, total_cost, actual_cost, cache_read_tokens, cache_hit_rate
|
||||
- 按用户聚合时:user_id, email, username, request_count, total_tokens, total_cost
|
||||
- 按提供商聚合时:provider_id, provider, request_count, total_tokens, total_cost, actual_cost, avg_response_time_ms, success_rate, error_count
|
||||
- 按 API 格式聚合时:api_format, request_count, total_tokens, total_cost, actual_cost, avg_response_time_ms
|
||||
- 按提供商聚合时:provider_id, provider, request_count, total_tokens, total_cost, actual_cost, avg_response_time_ms, success_rate, error_count, cache_read_tokens, cache_hit_rate
|
||||
- 按 API 格式聚合时:api_format, request_count, total_tokens, total_cost, actual_cost, avg_response_time_ms, cache_read_tokens, cache_hit_rate
|
||||
"""
|
||||
time_range = _apply_admin_default_range(
|
||||
_build_time_range_params(start_date, end_date, preset, timezone_name, tz_offset_minutes)
|
||||
@@ -524,6 +544,8 @@ class AdminUsageByModelAdapter(AdminApiAdapter):
|
||||
func.sum(Usage.total_tokens).label("total_tokens"),
|
||||
func.sum(Usage.total_cost_usd).label("total_cost"),
|
||||
func.sum(Usage.actual_total_cost_usd).label("actual_cost"),
|
||||
func.sum(Usage.cache_read_input_tokens).label("cache_read_tokens"),
|
||||
func.sum(input_context_expr()).label("total_input_context"),
|
||||
)
|
||||
# 过滤掉 pending/streaming 状态的请求(尚未完成的请求不应计入统计)
|
||||
query = query.filter(Usage.status.notin_(["pending", "streaming"]))
|
||||
@@ -553,8 +575,13 @@ class AdminUsageByModelAdapter(AdminApiAdapter):
|
||||
"total_tokens": int(tokens or 0),
|
||||
"total_cost": float(cost or 0),
|
||||
"actual_cost": float(actual_cost or 0),
|
||||
"cache_read_tokens": int(cache_read_tokens or 0),
|
||||
"cache_hit_rate": _calculate_token_cache_hit_rate(
|
||||
total_input_context=total_input_context,
|
||||
cache_read_tokens=cache_read_tokens,
|
||||
),
|
||||
}
|
||||
for model, count, tokens, cost, actual_cost in stats
|
||||
for model, count, tokens, cost, actual_cost, cache_read_tokens, total_input_context in stats
|
||||
]
|
||||
|
||||
|
||||
@@ -678,6 +705,8 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
func.sum(Usage.total_cost_usd).label("total_cost"),
|
||||
func.sum(Usage.actual_total_cost_usd).label("actual_cost"),
|
||||
func.avg(Usage.response_time_ms).label("avg_response_time_ms"),
|
||||
func.sum(Usage.cache_read_input_tokens).label("cache_read_tokens"),
|
||||
func.sum(input_context_expr()).label("total_input_context"),
|
||||
).filter(
|
||||
Usage.provider_id.isnot(None),
|
||||
# 过滤掉 pending/streaming 状态的请求
|
||||
@@ -742,6 +771,13 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
"avg_response_time_ms": float(stat.avg_latency_ms or 0),
|
||||
"success_rate": round(success_rate, 2),
|
||||
"error_count": failed_count,
|
||||
"cache_read_tokens": (
|
||||
int(usage_stat.cache_read_tokens or 0) if usage_stat else 0
|
||||
),
|
||||
"cache_hit_rate": _calculate_token_cache_hit_rate(
|
||||
total_input_context=(usage_stat.total_input_context if usage_stat else 0),
|
||||
cache_read_tokens=(usage_stat.cache_read_tokens if usage_stat else 0),
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -773,6 +809,8 @@ class AdminUsageByApiFormatAdapter(AdminApiAdapter):
|
||||
func.sum(Usage.total_cost_usd).label("total_cost"),
|
||||
func.sum(Usage.actual_total_cost_usd).label("actual_cost"),
|
||||
func.avg(Usage.response_time_ms).label("avg_response_time_ms"),
|
||||
func.sum(Usage.cache_read_input_tokens).label("cache_read_tokens"),
|
||||
func.sum(input_context_expr()).label("total_input_context"),
|
||||
)
|
||||
# 过滤掉 pending/streaming 状态的请求
|
||||
query = query.filter(Usage.status.notin_(["pending", "streaming"]))
|
||||
@@ -808,8 +846,22 @@ class AdminUsageByApiFormatAdapter(AdminApiAdapter):
|
||||
"total_cost": float(cost or 0),
|
||||
"actual_cost": float(actual_cost or 0),
|
||||
"avg_response_time_ms": float(avg_response_time or 0),
|
||||
"cache_read_tokens": int(cache_read_tokens or 0),
|
||||
"cache_hit_rate": _calculate_token_cache_hit_rate(
|
||||
total_input_context=total_input_context,
|
||||
cache_read_tokens=cache_read_tokens,
|
||||
),
|
||||
}
|
||||
for api_format, count, tokens, cost, actual_cost, avg_response_time in stats
|
||||
for (
|
||||
api_format,
|
||||
count,
|
||||
tokens,
|
||||
cost,
|
||||
actual_cost,
|
||||
avg_response_time,
|
||||
cache_read_tokens,
|
||||
total_input_context,
|
||||
) in stats
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -48,6 +48,7 @@ from src.models.database import (
|
||||
from src.services.cache.user_cache import UserCacheService
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.services.system.time_range import TimeRangeParams
|
||||
from src.services.usage.query import input_context_expr
|
||||
from src.services.usage.service import UsageService
|
||||
from src.services.user.apikey import ApiKeyService
|
||||
from src.services.user.bulk_cleanup import pre_clean_api_key
|
||||
@@ -59,6 +60,20 @@ router = APIRouter(prefix="/api/users/me", tags=["User Profile"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
def _calculate_token_cache_hit_rate(total_input_context: int, cache_read_tokens: int) -> float:
|
||||
"""计算缓存命中率。
|
||||
|
||||
Args:
|
||||
total_input_context: 已归一化的总输入上下文 token 数(由 query.py 按 API 格式精确计算)。
|
||||
cache_read_tokens: 缓存读取 token 数。
|
||||
"""
|
||||
context = max(0, int(total_input_context))
|
||||
cached = max(0, int(cache_read_tokens))
|
||||
if context == 0:
|
||||
return 0.0
|
||||
return round(cached / context * 100, 2)
|
||||
|
||||
|
||||
def _update_profile_sync(
|
||||
user_id: str,
|
||||
request: UpdateProfileRequest,
|
||||
@@ -508,8 +523,8 @@ async def get_my_usage(
|
||||
- `total_requests`: 总请求数
|
||||
- `total_tokens`: 总 Token 数
|
||||
- `total_cost`: 总成本(USD)
|
||||
- `summary_by_model`: 按模型分组统计
|
||||
- `summary_by_provider`: 按提供商分组统计
|
||||
- `summary_by_model`: 按模型分组统计(含 `cache_read_tokens`、`cache_hit_rate`)
|
||||
- `summary_by_provider`: 按提供商分组统计(含 `cache_read_tokens`、`cache_hit_rate`)
|
||||
- `records`: 详细使用记录列表
|
||||
- `pagination`: 分页信息
|
||||
"""
|
||||
@@ -1045,6 +1060,9 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
"total_tokens": 0,
|
||||
"cache_read_tokens": 0,
|
||||
"total_input_context": 0,
|
||||
"cache_hit_rate": 0.0,
|
||||
"total_cost_usd": 0.0,
|
||||
}
|
||||
# 管理员可以看到真实成本
|
||||
@@ -1056,6 +1074,8 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
stats["input_tokens"] += item["input_tokens"]
|
||||
stats["output_tokens"] += item["output_tokens"]
|
||||
stats["total_tokens"] += item["total_tokens"]
|
||||
stats["cache_read_tokens"] += int(item.get("cache_read_tokens", 0) or 0)
|
||||
stats["total_input_context"] += int(item.get("total_input_context", 0) or 0)
|
||||
stats["total_cost_usd"] += item["total_cost_usd"]
|
||||
# 管理员可以看到真实成本
|
||||
if user.role == UserRole.ADMIN:
|
||||
@@ -1066,6 +1086,8 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
"provider": provider_name,
|
||||
"requests": 0,
|
||||
"total_tokens": 0,
|
||||
"cache_read_tokens": 0,
|
||||
"total_input_context": 0,
|
||||
"total_cost_usd": 0.0,
|
||||
"success_count": 0,
|
||||
"total_response_time_ms": 0.0,
|
||||
@@ -1074,6 +1096,8 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
provider_stats = provider_summary.setdefault(provider_name, provider_base_stats)
|
||||
provider_stats["requests"] += item["requests"]
|
||||
provider_stats["total_tokens"] += item["total_tokens"]
|
||||
provider_stats["cache_read_tokens"] += int(item.get("cache_read_tokens", 0) or 0)
|
||||
provider_stats["total_input_context"] += int(item.get("total_input_context", 0) or 0)
|
||||
provider_stats["total_cost_usd"] += item["total_cost_usd"]
|
||||
provider_stats["success_count"] += int(item.get("success_count", 0) or 0)
|
||||
success_response_time_count = int(item.get("success_response_time_count", 0) or 0)
|
||||
@@ -1083,6 +1107,13 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
)
|
||||
provider_stats["response_time_count"] += success_response_time_count
|
||||
|
||||
for model_stats in model_summary.values():
|
||||
model_stats["cache_hit_rate"] = _calculate_token_cache_hit_rate(
|
||||
total_input_context=int(model_stats.get("total_input_context", 0) or 0),
|
||||
cache_read_tokens=int(model_stats.get("cache_read_tokens", 0) or 0),
|
||||
)
|
||||
model_stats.pop("total_input_context", None)
|
||||
|
||||
summary_by_model = sorted(model_summary.values(), key=lambda x: x["requests"], reverse=True)
|
||||
summary_by_provider = []
|
||||
for provider_stats in provider_summary.values():
|
||||
@@ -1101,6 +1132,11 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
"provider": provider_stats["provider"],
|
||||
"requests": provider_stats["requests"],
|
||||
"total_tokens": provider_stats["total_tokens"],
|
||||
"cache_read_tokens": provider_stats["cache_read_tokens"],
|
||||
"cache_hit_rate": _calculate_token_cache_hit_rate(
|
||||
total_input_context=int(provider_stats.get("total_input_context", 0) or 0),
|
||||
cache_read_tokens=int(provider_stats.get("cache_read_tokens", 0) or 0),
|
||||
),
|
||||
"total_cost_usd": provider_stats["total_cost_usd"],
|
||||
"success_rate": round(success_rate, 2),
|
||||
"avg_response_time_ms": round(avg_response_time_ms, 2),
|
||||
@@ -1108,6 +1144,52 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
)
|
||||
summary_by_provider = sorted(summary_by_provider, key=lambda x: x["requests"], reverse=True)
|
||||
|
||||
# 按 api_format 聚合统计(独立查询,因为 get_usage_summary 按 provider+model 分组无此维度)
|
||||
api_format_query = db.query(
|
||||
Usage.api_format,
|
||||
func.count(Usage.id).label("request_count"),
|
||||
func.sum(Usage.total_tokens).label("total_tokens"),
|
||||
func.sum(Usage.cache_read_input_tokens).label("cache_read_tokens"),
|
||||
func.sum(input_context_expr()).label("total_input_context"),
|
||||
func.sum(Usage.total_cost_usd).label("total_cost_usd"),
|
||||
func.avg(Usage.response_time_ms).label("avg_response_time_ms"),
|
||||
).filter(
|
||||
Usage.user_id == user.id,
|
||||
Usage.status.notin_(["pending", "streaming"]),
|
||||
Usage.provider_name.notin_(["unknown", "pending"]),
|
||||
Usage.api_format.isnot(None),
|
||||
)
|
||||
if start_utc and end_utc:
|
||||
api_format_query = api_format_query.filter(
|
||||
Usage.created_at >= start_utc, Usage.created_at < end_utc
|
||||
)
|
||||
api_format_stats = (
|
||||
api_format_query.group_by(Usage.api_format).order_by(func.count(Usage.id).desc()).all()
|
||||
)
|
||||
summary_by_api_format = [
|
||||
{
|
||||
"api_format": api_format or "unknown",
|
||||
"request_count": count,
|
||||
"total_tokens": int(total_tokens or 0),
|
||||
"cache_read_tokens": int(cache_read_tokens or 0),
|
||||
"cache_hit_rate": _calculate_token_cache_hit_rate(
|
||||
total_input_context=total_input_context,
|
||||
cache_read_tokens=cache_read_tokens,
|
||||
),
|
||||
"total_cost_usd": float(total_cost_usd or 0),
|
||||
"avg_response_time_ms": float(avg_response_time_ms or 0),
|
||||
}
|
||||
for (
|
||||
api_format,
|
||||
count,
|
||||
total_tokens,
|
||||
cache_read_tokens,
|
||||
total_input_context,
|
||||
total_cost_usd,
|
||||
avg_response_time_ms,
|
||||
) in api_format_stats
|
||||
]
|
||||
|
||||
query = (
|
||||
db.query(Usage, ApiKey, ProviderEndpoint)
|
||||
.outerjoin(ApiKey, Usage.api_key_id == ApiKey.id)
|
||||
@@ -1199,6 +1281,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
"billing": WalletService.serialize_wallet_summary(wallet),
|
||||
"summary_by_model": summary_by_model,
|
||||
"summary_by_provider": summary_by_provider,
|
||||
"summary_by_api_format": summary_by_api_format,
|
||||
# 分页信息
|
||||
"pagination": {
|
||||
"total": total_records,
|
||||
|
||||
@@ -471,186 +471,6 @@ class PoolManager:
|
||||
|
||||
return ordered_keys, trace
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Single-key selection (used by CandidateBuilder for pooled providers)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def select_key(
|
||||
self,
|
||||
session_uuid: str | None,
|
||||
keys: list[ProviderAPIKey],
|
||||
) -> ProviderAPIKey | None:
|
||||
"""Select the best key from *keys* according to pool rules.
|
||||
|
||||
Same logic as :meth:`reorder_candidates` but operates directly on
|
||||
:class:`ProviderAPIKey` objects instead of candidates:
|
||||
|
||||
1. Sticky session hit (if bound and still healthy).
|
||||
2. Filter out keys in account-blocked / cooldown / cost-exhausted.
|
||||
3. LRU sort among remaining keys.
|
||||
4. Random tiebreak for identical LRU scores.
|
||||
5. Return the first available key, or ``None``.
|
||||
"""
|
||||
if not keys:
|
||||
return None
|
||||
|
||||
pid = self.provider_id
|
||||
|
||||
# --- 1. Sticky session ------------------------------------------------
|
||||
sticky_key_id: str | None = None
|
||||
if session_uuid and self.config.sticky_session_ttl_seconds > 0:
|
||||
sticky_key_id = await redis_ops.get_sticky_binding(
|
||||
pid, session_uuid, self.config.sticky_session_ttl_seconds
|
||||
)
|
||||
|
||||
# --- 2. Batch fetch pool state (parallel) -----------------------------
|
||||
key_ids = [str(k.id) for k in keys]
|
||||
|
||||
_cooldown_coro = redis_ops.batch_get_cooldowns(pid, key_ids)
|
||||
_cost_coro = (
|
||||
redis_ops.batch_get_cost_totals(pid, key_ids, self.config.cost_window_seconds)
|
||||
if (
|
||||
self.config.cost_limit_per_key_tokens is not None
|
||||
or self.config.scheduling_mode == "multi_score"
|
||||
)
|
||||
else None
|
||||
)
|
||||
_need_lru_sk = self.config.lru_enabled or self.config.scheduling_mode == "multi_score"
|
||||
_lru_coro = redis_ops.get_lru_scores(pid, key_ids) if _need_lru_sk else None
|
||||
_latency_coro = (
|
||||
redis_ops.batch_get_latency_avgs(pid, key_ids, self.config.latency_window_seconds)
|
||||
if self.config.scheduling_mode == "multi_score"
|
||||
else None
|
||||
)
|
||||
|
||||
coros_sk: list[Any] = [_cooldown_coro]
|
||||
_cost_idx_sk = -1
|
||||
_lru_idx_sk = -1
|
||||
_latency_idx_sk = -1
|
||||
if _cost_coro is not None:
|
||||
_cost_idx_sk = len(coros_sk)
|
||||
coros_sk.append(_cost_coro)
|
||||
if _lru_coro is not None:
|
||||
_lru_idx_sk = len(coros_sk)
|
||||
coros_sk.append(_lru_coro)
|
||||
if _latency_coro is not None:
|
||||
_latency_idx_sk = len(coros_sk)
|
||||
coros_sk.append(_latency_coro)
|
||||
|
||||
gathered_sk = await asyncio.gather(*coros_sk)
|
||||
|
||||
cooldowns = gathered_sk[0]
|
||||
|
||||
cost_exhausted: set[str] = set()
|
||||
cost_totals: dict[str, int] = {}
|
||||
if _cost_idx_sk >= 0:
|
||||
cost_totals = gathered_sk[_cost_idx_sk]
|
||||
if self.config.cost_limit_per_key_tokens is not None:
|
||||
for kid, total in cost_totals.items():
|
||||
if total >= self.config.cost_limit_per_key_tokens:
|
||||
cost_exhausted.add(kid)
|
||||
|
||||
lru_scores: dict[str, float] = {}
|
||||
if _lru_idx_sk >= 0:
|
||||
lru_scores = gathered_sk[_lru_idx_sk]
|
||||
|
||||
latency_avgs: dict[str, float] = {}
|
||||
if _latency_idx_sk >= 0:
|
||||
latency_avgs = gathered_sk[_latency_idx_sk]
|
||||
|
||||
health_scores: dict[str, float] = {}
|
||||
if self.config.scheduling_mode == "multi_score":
|
||||
health_scores = get_health_scores(pid, keys)
|
||||
|
||||
provider_type = self.provider_type
|
||||
if provider_type is None and keys:
|
||||
first_provider = getattr(keys[0], "provider", None)
|
||||
provider_type = str(getattr(first_provider, "provider_type", "") or "").strip().lower()
|
||||
if not provider_type:
|
||||
provider_type = None
|
||||
|
||||
# --- Strategy: compute_score ------------------------------------------
|
||||
strategies = _get_active_strategies(self.config)
|
||||
strategy_context: dict[str, Any] = {
|
||||
"session_uuid": session_uuid,
|
||||
"provider_type": provider_type,
|
||||
"all_key_ids": key_ids,
|
||||
"lru_scores": lru_scores,
|
||||
"cost_totals": cost_totals if _cost_idx_sk >= 0 else {},
|
||||
"cost_limit_per_key_tokens": self.config.cost_limit_per_key_tokens,
|
||||
"latency_avgs": latency_avgs,
|
||||
"health_scores": health_scores,
|
||||
"keys_by_id": {str(k.id): k for k in keys},
|
||||
}
|
||||
for strategy in strategies:
|
||||
if hasattr(strategy, "compute_score"):
|
||||
for kid in key_ids:
|
||||
try:
|
||||
custom = strategy.compute_score(
|
||||
key_id=kid,
|
||||
config=self.config,
|
||||
context=strategy_context,
|
||||
)
|
||||
if custom is not None:
|
||||
lru_scores[kid] = float(custom)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# --- 3. Classify keys -------------------------------------------------
|
||||
# Use precomputed account states when available.
|
||||
account_states_sk: dict[str, Any] = {}
|
||||
for k in keys:
|
||||
kid = str(k.id)
|
||||
if kid not in account_states_sk:
|
||||
precomputed = getattr(k, "_pool_account_state", None)
|
||||
if precomputed is not None:
|
||||
account_states_sk[kid] = precomputed
|
||||
else:
|
||||
account_states_sk[kid] = resolve_pool_account_state(
|
||||
provider_type=self.provider_type,
|
||||
upstream_metadata=getattr(k, "upstream_metadata", None),
|
||||
oauth_invalid_reason=getattr(k, "oauth_invalid_reason", None),
|
||||
)
|
||||
|
||||
sticky_key: ProviderAPIKey | None = None
|
||||
available: list[ProviderAPIKey] = []
|
||||
|
||||
for k in keys:
|
||||
kid = str(k.id)
|
||||
|
||||
if account_states_sk[kid].blocked:
|
||||
continue
|
||||
|
||||
if cooldowns.get(kid) is not None:
|
||||
continue
|
||||
if kid in cost_exhausted:
|
||||
continue
|
||||
|
||||
if sticky_key_id and kid == sticky_key_id:
|
||||
sticky_key = k
|
||||
continue
|
||||
|
||||
available.append(k)
|
||||
|
||||
# --- 4. Sort by LRU ---------------------------------------------------
|
||||
if lru_scores and available:
|
||||
available.sort(key=lambda k: lru_scores.get(str(k.id), 0.0))
|
||||
|
||||
# Random tiebreak within same-score groups
|
||||
if len(available) > 1 and lru_scores:
|
||||
_shuffle_same_score_keys(available, lru_scores)
|
||||
|
||||
# --- 5. Pick the winner -----------------------------------------------
|
||||
if sticky_key is not None:
|
||||
logger.debug(
|
||||
"Pool[{}]: sticky select key={}",
|
||||
pid[:8],
|
||||
sticky_key_id and sticky_key_id[:8],
|
||||
)
|
||||
return sticky_key
|
||||
|
||||
return available[0] if available else None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Post-request hooks
|
||||
# ------------------------------------------------------------------
|
||||
@@ -799,10 +619,3 @@ def _shuffle_same_score_groups(
|
||||
lru_scores: dict[str, float],
|
||||
) -> None:
|
||||
_shuffle_same_score(candidates, lru_scores, lambda c: str(c.key.id))
|
||||
|
||||
|
||||
def _shuffle_same_score_keys(
|
||||
keys: list[ProviderAPIKey],
|
||||
lru_scores: dict[str, float],
|
||||
) -> None:
|
||||
_shuffle_same_score(keys, lru_scores, lambda k: str(k.id))
|
||||
|
||||
@@ -4,13 +4,28 @@ from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import case, func
|
||||
from sqlalchemy import Case, case, func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, Usage, User
|
||||
|
||||
|
||||
def input_context_expr() -> Case:
|
||||
"""构造 SQL CASE 表达式,根据 api_format 精确计算每条记录的总输入上下文 token 数。
|
||||
|
||||
- OpenAI/Gemini: input_tokens 已包含 cache_read_input_tokens,直接使用
|
||||
- Claude/未知: input_tokens 不含 cache_read,需要加上
|
||||
"""
|
||||
return case(
|
||||
(
|
||||
Usage.api_format.like("openai:%") | Usage.api_format.like("gemini:%"),
|
||||
Usage.input_tokens,
|
||||
),
|
||||
else_=Usage.input_tokens + Usage.cache_read_input_tokens,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class RequestBalanceCheckResult:
|
||||
allowed: bool
|
||||
@@ -245,6 +260,8 @@ class UsageQueryMixin:
|
||||
func.sum(Usage.input_tokens).label("input_tokens"),
|
||||
func.sum(Usage.output_tokens).label("output_tokens"),
|
||||
func.sum(Usage.total_tokens).label("total_tokens"),
|
||||
func.sum(Usage.cache_read_input_tokens).label("cache_read_tokens"),
|
||||
func.sum(input_context_expr()).label("total_input_context"),
|
||||
func.sum(Usage.total_cost_usd).label("total_cost_usd"),
|
||||
func.sum(Usage.actual_total_cost_usd).label("actual_total_cost_usd"),
|
||||
func.sum(case((Usage.status_code == 200, 1), else_=0)).label("success_count"),
|
||||
@@ -289,6 +306,8 @@ class UsageQueryMixin:
|
||||
"input_tokens": row.input_tokens,
|
||||
"output_tokens": row.output_tokens,
|
||||
"total_tokens": row.total_tokens,
|
||||
"cache_read_tokens": int(row.cache_read_tokens or 0),
|
||||
"total_input_context": int(row.total_input_context or 0),
|
||||
"total_cost_usd": float(row.total_cost_usd or 0.0),
|
||||
"actual_total_cost_usd": float(row.actual_total_cost_usd or 0.0),
|
||||
"success_count": int(row.success_count or 0),
|
||||
|
||||
Reference in New Issue
Block a user