feat(pool): 号池额度主动探测、封禁自动清除、调度硬优先级与前端重构

- 新增 PoolQuotaProbeScheduler,按 probing_interval_minutes 主动探测静默 Key 额度
- pool_advanced 增加 probing_enabled / auto_remove_banned_keys 配置项
- error_handler 和 quota_service 支持封禁 Key 自动删除及缓存清理
- multi_score 策略从加权混合重构为硬优先级排序,引入 mutex_group 互斥组
- 指纹注入从 handler 层下移至 ClaudeCode envelope 层
- OAuth 批量导入支持 concurrency 并发参数
- 前端号池管理拆分高级设置/账号批量/代理设置为独立组件
- 号池总览接口精简,仅返回已启用调度的 Provider
This commit is contained in:
fawney19
2026-03-05 15:15:26 +08:00
parent fdb50a065b
commit b1be413dc0
29 changed files with 3098 additions and 852 deletions

View File

@@ -7,6 +7,7 @@
from __future__ import annotations
import asyncio
import json
import re
from typing import Any
@@ -25,6 +26,7 @@ from src.core.logger import logger
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.health.monitor import health_monitor
from src.services.provider.format import normalize_endpoint_signature
from src.services.provider.pool.config import parse_pool_config
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.services.rate_limit.detector import RateLimitType, detect_rate_limit_type
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
@@ -168,7 +170,7 @@ class ErrorHandlerService:
and str(getattr(key, "auth_type", "") or "").lower() == "oauth"
and self._is_account_validation_required(error_response_text)
):
self._mark_oauth_key_blocked(key, request_id)
self._mark_oauth_key_blocked(key, request_id, provider=provider)
# 403 suspended -> 标记 OAuth key 为账号被暂停
elif (
status_code == 403
@@ -176,7 +178,12 @@ class ErrorHandlerService:
and str(getattr(key, "auth_type", "") or "").lower() == "oauth"
and self._is_account_suspended(error_response_text)
):
self._mark_oauth_key_blocked(key, request_id, reason="AWS 账号被暂停")
self._mark_oauth_key_blocked(
key,
request_id,
reason="AWS 账号被暂停",
provider=provider,
)
return
# 限流错误
@@ -345,6 +352,8 @@ class ErrorHandlerService:
key: ProviderAPIKey,
request_id: str | None,
reason: str = "Google 要求验证账号",
*,
provider: Provider,
) -> None:
"""标记 OAuth key 为账号级别封禁"""
try:
@@ -355,6 +364,26 @@ class ErrorHandlerService:
key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = f"{OAUTH_ACCOUNT_BLOCK_PREFIX}{reason}"
key.is_active = False
pool_cfg = parse_pool_config(getattr(provider, "config", None))
auto_remove_enabled = bool(pool_cfg and pool_cfg.auto_remove_banned_keys)
if auto_remove_enabled:
key_id = str(getattr(key, "id", "") or "")
provider_id = str(getattr(key, "provider_id", "") or "")
display = self._format_key_display(key)
self.db.delete(key)
self.db.commit()
self._schedule_auto_cleanup_after_delete(provider_id=provider_id, key_id=key_id)
logger.warning(
" [{}] {}{} 已标记为账号异常并自动清除",
request_id,
display,
reason,
)
return
self.db.commit()
logger.warning(
" [{}] {}{} 已标记为账号异常并自动停用",
@@ -364,3 +393,31 @@ class ErrorHandlerService:
)
except Exception as mark_exc:
logger.debug(" [{}] 标记 oauth_invalid 失败: {}", request_id, mark_exc)
@staticmethod
def _schedule_auto_cleanup_after_delete(*, provider_id: str, key_id: str) -> None:
if not provider_id or not key_id:
return
async def _cleanup() -> None:
from src.api.base.models_service import invalidate_models_list_cache
from src.services.cache.provider_cache import ProviderCacheService
from src.services.provider.pool import redis_ops as pool_redis
await ProviderCacheService.invalidate_provider_api_key_cache(key_id)
await invalidate_models_list_cache()
await asyncio.gather(
pool_redis.clear_cooldown(provider_id, key_id),
pool_redis.clear_cost(provider_id, key_id),
return_exceptions=True,
)
task = asyncio.get_running_loop().create_task(_cleanup())
def _log_async_error(done_task: asyncio.Task[Any]) -> None:
try:
done_task.result()
except Exception as exc:
logger.debug("auto cleanup side effect failed for key {}: {}", key_id[:8], exc)
task.add_done_callback(_log_async_error)

View File

@@ -585,11 +585,19 @@ class ClaudeCodeEnvelope:
key_id: str,
is_stream: bool,
provider_id: str | None = None,
key: Any = None,
) -> str | None:
from src.services.provider.adapters.claude_code.context import (
build_and_set_claude_code_request_context,
)
# 在 envelope 层设置指纹 context var仅 Claude Code 需要指纹注入)
if key is not None:
from src.services.provider.fingerprint import ensure_key_fingerprint
from src.services.provider.request_context import set_current_fingerprint
set_current_fingerprint(ensure_key_fingerprint(key, persist_if_missing=True))
_ctx, tls_profile = build_and_set_claude_code_request_context(
provider_config=provider_config,
key_id=key_id,

View File

@@ -61,6 +61,7 @@ class ProviderEnvelope(Protocol):
key_id: str,
is_stream: bool,
provider_id: str | None = None,
key: Any = None,
) -> str | None:
"""Pre-wrap hook: build provider-specific request context.

View File

@@ -82,6 +82,11 @@ class PoolConfig:
# -- Temporary Unschedulable Rules ----------------------------------------
unschedulable_rules: list[UnschedulableRule] = field(default_factory=list)
# -- Quota Probing --------------------------------------------------------
probing_enabled: bool = False
probing_interval_minutes: int = 10
auto_remove_banned_keys: bool = False
# -- Stream Timeout Auto-Pause --------------------------------------------
stream_timeout_threshold: int = 3 # N timeouts within window trigger cooldown
stream_timeout_window_seconds: int = 1800 # 30 min counting window
@@ -188,6 +193,9 @@ def parse_pool_config(provider_config: Any) -> PoolConfig | None:
proactive_refresh_seconds=_int_or("proactive_refresh_seconds", 180),
health_policy_enabled=_bool_or("health_policy_enabled", True),
unschedulable_rules=rules,
probing_enabled=_bool_or("probing_enabled", False),
probing_interval_minutes=max(1, min(_int_or("probing_interval_minutes", 10), 1440)),
auto_remove_banned_keys=_bool_or("auto_remove_banned_keys", False),
stream_timeout_threshold=_int_or("stream_timeout_threshold", 3),
stream_timeout_window_seconds=_int_or("stream_timeout_window_seconds", 1800),
stream_timeout_cooldown_seconds=_int_or("stream_timeout_cooldown_seconds", 300),

View File

@@ -9,18 +9,31 @@ from src.services.provider.pool.dimensions import get_preset_dimension, get_pres
from src.services.provider.pool.dimensions._helpers import rank_ascending, safe_float
from src.services.provider.pool.strategy import register_pool_strategy
# When LRU is enabled alongside presets, this fraction of the final score
# comes from the LRU rank (tiebreaker to avoid same-score collisions).
_LRU_BLEND_FACTOR = 0.04
# Positional weight decay factor: weight = 1 / (1 + DECAY * index).
_POSITIONAL_DECAY = 0.6
def _normalize_mutex_group(value: Any) -> str | None:
if not isinstance(value, str):
return None
normalized = value.strip().lower()
return normalized or None
def _get_preset_mutex_group(preset_name: str) -> str | None:
# LRU is a built-in preset (not in registry) but shares the distribution mutex group.
if preset_name == "lru":
return "distribution_mode"
dim = get_preset_dimension(preset_name)
if dim is None:
return None
return _normalize_mutex_group(getattr(dim, "mutex_group", None))
def _normalize_presets_from_config(config: Any) -> tuple[tuple[str, str | None], ...]:
"""Extract enabled (preset_name, mode) tuples from config.scheduling_presets.
Supports both new SchedulingPreset objects and legacy string lists.
Excludes ``lru`` since LRU is handled separately as a blend factor.
Excludes ``lru`` from output (LRU is a final tie-breaker only).
For mutex groups, enabled members inherit the group's first appearance index
so the selected member keeps the group's visible priority slot.
"""
raw = getattr(config, "scheduling_presets", ())
@@ -28,9 +41,9 @@ def _normalize_presets_from_config(config: Any) -> tuple[tuple[str, str | None],
return ()
allowed = get_preset_names() | {"lru"}
ordered: list[tuple[str, str | None]] = []
entries: list[tuple[int, str, bool, str | None]] = []
seen: set[str] = set()
for item in raw:
for idx, item in enumerate(raw):
preset_name: str | None = None
enabled = True
mode: str | None = None
@@ -48,13 +61,36 @@ def _normalize_presets_from_config(config: Any) -> tuple[tuple[str, str | None],
if not preset_name or preset_name not in allowed or preset_name in seen:
continue
if not enabled:
continue
if preset_name == "lru":
continue
seen.add(preset_name)
ordered.append((preset_name, mode))
return tuple(ordered)
entries.append((idx, preset_name, enabled, mode))
if not entries:
return ()
group_anchor_index: dict[str, int] = {}
for idx, preset_name, _enabled, _mode in entries:
mutex_group = _get_preset_mutex_group(preset_name)
if mutex_group and mutex_group not in group_anchor_index:
group_anchor_index[mutex_group] = idx
ordered_enabled: list[tuple[int, int, str, str | None]] = []
group_enabled: dict[str, tuple[int, int, str, str | None]] = {}
for idx, preset_name, enabled, mode in entries:
if not enabled or preset_name == "lru":
continue
mutex_group = _get_preset_mutex_group(preset_name)
if not mutex_group:
ordered_enabled.append((idx, idx, preset_name, mode))
continue
anchor = group_anchor_index.get(mutex_group, idx)
existing = group_enabled.get(mutex_group)
if existing is None or idx < existing[1]:
group_enabled[mutex_group] = (anchor, idx, preset_name, mode)
ordered_enabled.extend(group_enabled.values())
ordered_enabled.sort(key=lambda item: (item[0], item[1]))
return tuple((preset_name, mode) for _anchor, _idx, preset_name, mode in ordered_enabled)
class MultiScoreStrategy:
@@ -143,34 +179,64 @@ class MultiScoreStrategy:
keys_by_id: dict[str, Any],
context: dict[str, Any],
) -> float:
lru_rank_asc = rank_ascending(key_id, lru_scores, all_key_ids)
cache_signature = (tuple(all_key_ids), presets, bool(lru_enabled))
cache = context.get("_preset_hard_order_cache")
if (
isinstance(cache, dict)
and cache.get("signature") == cache_signature
and isinstance(cache.get("ranks"), dict)
):
cached_rank = safe_float(cache["ranks"].get(key_id))
if cached_rank is not None:
return max(0.0, min(cached_rank, 1.0))
weighted_sum = 0.0
weight_sum = 0.0
for idx, (preset_name, mode) in enumerate(presets):
metric = 0.5
dim = get_preset_dimension(preset_name)
if dim is not None:
metric = dim.compute_metric(
key_id=key_id,
all_key_ids=all_key_ids,
keys_by_id=keys_by_id,
lru_scores=lru_scores,
context=context,
mode=mode,
# Hard-priority semantics:
# 1) Compare by preset[0] metric first;
# 2) only if tied, compare preset[1], preset[2], ...
# 3) if all preset metrics tie and LRU is enabled, use LRU as final tiebreak.
metric_vectors: dict[str, tuple[float, ...]] = {}
for kid in all_key_ids:
vector_parts: list[float] = []
for preset_name, mode in presets:
metric = 0.5
dim = get_preset_dimension(preset_name)
if dim is not None:
metric = dim.compute_metric(
key_id=kid,
all_key_ids=all_key_ids,
keys_by_id=keys_by_id,
lru_scores=lru_scores,
context=context,
mode=mode,
)
metric_value = safe_float(metric)
vector_parts.append(
max(0.0, min(metric_value, 1.0)) if metric_value is not None else 0.5
)
weight = 1.0 / (1.0 + _POSITIONAL_DECAY * idx)
weighted_sum += metric * weight
weight_sum += weight
if weight_sum <= 0:
return lru_rank_asc
if lru_enabled:
vector_parts.append(rank_ascending(kid, lru_scores, all_key_ids))
lru_blend = _LRU_BLEND_FACTOR if lru_enabled else 0.0
preset_blend = 1.0 - lru_blend
blended = (weighted_sum / weight_sum) * preset_blend + lru_rank_asc * lru_blend
return max(0.0, min(blended, 1.0))
metric_vectors[kid] = tuple(vector_parts)
decorated = [
(metric_vectors.get(kid, (0.5,)), idx, kid) for idx, kid in enumerate(all_key_ids)
]
decorated.sort(key=lambda item: (item[0], item[1]))
total = len(decorated)
ranks: dict[str, float] = {}
for rank_idx, (_vec, _idx, kid) in enumerate(decorated):
ranks[kid] = 0.0 if total <= 1 else rank_idx / float(total - 1)
context["_preset_hard_order_cache"] = {
"signature": cache_signature,
"ranks": ranks,
}
rank = safe_float(ranks.get(key_id))
if rank is None:
return 0.5
return max(0.0, min(rank, 1.0))
register_pool_strategy("multi_score", MultiScoreStrategy())

View File

@@ -13,6 +13,10 @@ from src.core.logger import logger
from src.core.provider_types import ProviderType, normalize_provider_type
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.model.upstream_fetcher import merge_upstream_metadata
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.config import parse_pool_config
from src.services.provider_keys.key_side_effects import run_delete_key_side_effects
from src.services.provider_keys.quota_refresh import (
refresh_antigravity_key_quota,
refresh_codex_key_quota,
@@ -74,6 +78,8 @@ async def refresh_provider_quota_for_provider(
provider_type = normalize_provider_type(getattr(provider, "provider_type", ""))
if provider_type not in {ProviderType.CODEX, ProviderType.ANTIGRAVITY, ProviderType.KIRO}:
raise InvalidRequestException("仅支持 Codex / Antigravity / Kiro 类型的 Provider 刷新限额")
pool_cfg = parse_pool_config(getattr(provider, "config", None))
auto_remove_banned_keys = bool(pool_cfg and pool_cfg.auto_remove_banned_keys)
selected_key_ids: list[str] | None = None
if key_ids is not None:
@@ -100,6 +106,7 @@ async def refresh_provider_quota_for_provider(
"total": 0,
"results": [],
"message": "未提供可刷新的 Key",
"auto_removed": 0,
}
keys_query = keys_query.filter(ProviderAPIKey.id.in_(selected_key_ids))
@@ -111,6 +118,7 @@ async def refresh_provider_quota_for_provider(
"total": 0,
"results": [],
"message": "没有可刷新的 Key",
"auto_removed": 0,
}
endpoint = _select_refresh_endpoint(provider, provider_type)
@@ -159,7 +167,14 @@ async def refresh_provider_quota_for_provider(
failed_count += 1
# 统一更新数据库(避免在并发任务中操作 session
if metadata_updates or state_updates:
auto_removed_contexts: list[tuple[str, str | None, list[str] | None]] = []
result_index_by_key_id: dict[str, dict[str, Any]] = {}
for result in results:
rid = str(result.get("key_id", "")).strip()
if rid:
result_index_by_key_id[rid] = result
if metadata_updates or state_updates or auto_remove_banned_keys:
for key in keys:
key_dirty = False
if key.id in metadata_updates:
@@ -173,11 +188,58 @@ async def refresh_provider_quota_for_provider(
for field_name, field_value in updates.items():
setattr(key, field_name, field_value)
key_dirty = True
if auto_remove_banned_keys:
account_state = resolve_pool_account_state(
provider_type=provider_type,
upstream_metadata=getattr(key, "upstream_metadata", None),
oauth_invalid_reason=getattr(key, "oauth_invalid_reason", None),
)
if account_state.blocked:
key_id = str(getattr(key, "id", "") or "")
auto_removed_contexts.append(
(
key_id,
(
str(getattr(key, "provider_id", "") or "")
if getattr(key, "provider_id", None)
else None
),
getattr(key, "allowed_models", None),
)
)
if key_id and key_id in result_index_by_key_id:
result_index_by_key_id[key_id]["auto_removed"] = True
db.delete(key)
continue
if key_dirty:
db.add(key)
db.commit()
if auto_removed_contexts:
cleanup_coros = []
for key_id, pid, _allowed_models in auto_removed_contexts:
if not key_id or not pid:
continue
cleanup_coros.append(pool_redis.clear_cooldown(pid, key_id))
cleanup_coros.append(pool_redis.clear_cost(pid, key_id))
if cleanup_coros:
await asyncio.gather(*cleanup_coros, return_exceptions=True)
for _key_id, pid, allowed_models in auto_removed_contexts:
await run_delete_key_side_effects(
db=db,
provider_id=pid,
deleted_key_allowed_models=allowed_models,
)
logger.warning(
"[QUOTA_REFRESH] Provider {}: auto removed {} banned key(s): {}",
provider_id,
len(auto_removed_contexts),
[ctx[0][:8] for ctx in auto_removed_contexts if ctx[0]],
)
failed_details = [
f"{r.get('key_name', r.get('key_id', '?'))}: {r.get('message', 'unknown')}"
for r in results
@@ -205,4 +267,5 @@ async def refresh_provider_quota_for_provider(
"failed": failed_count,
"total": len(keys),
"results": results,
"auto_removed": len(auto_removed_contexts),
}

View File

@@ -0,0 +1,367 @@
"""
号池额度主动探测调度器。
行为:
- 当 provider.pool_advanced.probing_enabled=true 时启用
- Key 在静默超过 probing_interval_minutes 后,主动触发额度刷新
- Key 一旦被实际请求使用last_used_at 变新),探测冷却自动重置
"""
from __future__ import annotations
import os
import time
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any
from src.clients.redis_client import get_redis_client
from src.core.logger import logger
from src.core.provider_types import ProviderType, normalize_provider_type
from src.database import create_session
from src.models.database import Provider, ProviderAPIKey
from src.services.provider.pool.config import parse_pool_config
from src.services.provider_keys.key_quota_service import refresh_provider_quota_for_provider
from src.services.system.scheduler import get_scheduler
# 与 admin 刷新额度 API 保持一致
_CODEX_WHAM_USAGE_URL = "https://chatgpt.com/backend-api/wham/usage"
_REDIS_PREFIX = "ap:quota_probe:last"
_DEFAULT_INTERVAL_MINUTES = 10
_DEFAULT_SCAN_INTERVAL_SECONDS = 60
_DEFAULT_MAX_KEYS_PER_PROVIDER = 50
_MAX_INTERVAL_MINUTES = 1440
_SUPPORTED_PROVIDER_TYPES = {
ProviderType.CODEX.value,
ProviderType.KIRO.value,
ProviderType.ANTIGRAVITY.value,
}
def _probe_stamp_key(provider_id: str, key_id: str) -> str:
return f"{_REDIS_PREFIX}:{provider_id}:{key_id}"
def _to_unix_seconds(value: datetime | None) -> int | None:
if not isinstance(value, datetime):
return None
dt = value if value.tzinfo else value.replace(tzinfo=timezone.utc)
try:
return int(dt.timestamp())
except Exception:
return None
def _to_float(value: Any) -> float | None:
if isinstance(value, bool):
return None
if isinstance(value, (int, float)):
return float(value)
if isinstance(value, str):
text = value.strip()
if not text:
return None
try:
return float(text)
except ValueError:
return None
return None
def _extract_quota_updated_at(provider_type: str, upstream_metadata: Any) -> int | None:
if not isinstance(upstream_metadata, dict):
return None
normalized = normalize_provider_type(provider_type)
if normalized == ProviderType.CODEX.value:
bucket = upstream_metadata.get("codex")
elif normalized == ProviderType.KIRO.value:
bucket = upstream_metadata.get("kiro")
elif normalized == ProviderType.ANTIGRAVITY.value:
bucket = upstream_metadata.get("antigravity")
else:
return None
if not isinstance(bucket, dict):
return None
updated_at = _to_float(bucket.get("updated_at"))
if updated_at is None or updated_at <= 0:
return None
# 兼容毫秒时间戳
if updated_at > 1_000_000_000_000:
updated_at /= 1000
return int(updated_at)
def _parse_probe_stamp(raw_value: Any) -> int | None:
parsed = _to_float(raw_value)
if parsed is None or parsed <= 0:
return None
return int(parsed)
def _normalize_probe_interval_minutes(raw_value: Any) -> int:
parsed = _to_float(raw_value)
if parsed is None:
return _DEFAULT_INTERVAL_MINUTES
return max(1, min(int(parsed), _MAX_INTERVAL_MINUTES))
def _select_probe_key_ids(
*,
keys: list[ProviderAPIKey],
provider_type: str,
now_ts: int,
interval_seconds: int,
last_probe_timestamps: dict[str, int],
limit: int,
) -> list[str]:
stale: list[tuple[int, str]] = []
for key in keys:
key_id = str(getattr(key, "id", "") or "")
if not key_id:
continue
last_used_ts = _to_unix_seconds(getattr(key, "last_used_at", None))
quota_updated_ts = _extract_quota_updated_at(
provider_type,
getattr(key, "upstream_metadata", None),
)
last_probe_ts = last_probe_timestamps.get(key_id)
anchor_ts = max(last_used_ts or 0, quota_updated_ts or 0, last_probe_ts or 0)
if anchor_ts <= 0 or (now_ts - anchor_ts) >= interval_seconds:
stale.append((anchor_ts, key_id))
# anchor 越小说明越久未被探测/使用,优先探测
stale.sort(key=lambda item: item[0])
if limit > 0:
stale = stale[:limit]
return [key_id for _, key_id in stale]
@dataclass(frozen=True, slots=True)
class _ProviderProbeTask:
provider_id: str
provider_type: str
probe_key_ids: list[str]
interval_seconds: int
class PoolQuotaProbeScheduler:
"""按号池高级配置执行额度主动探测。"""
def __init__(self) -> None:
scan_interval_raw = os.getenv(
"POOL_QUOTA_PROBE_SCAN_INTERVAL_SECONDS",
str(_DEFAULT_SCAN_INTERVAL_SECONDS),
)
max_keys_raw = os.getenv(
"POOL_QUOTA_PROBE_MAX_KEYS_PER_PROVIDER",
str(_DEFAULT_MAX_KEYS_PER_PROVIDER),
)
self.scan_interval_seconds = max(
15, int(_to_float(scan_interval_raw) or _DEFAULT_SCAN_INTERVAL_SECONDS)
)
self.max_keys_per_provider = max(
0, int(_to_float(max_keys_raw) or _DEFAULT_MAX_KEYS_PER_PROVIDER)
)
self.running = False
async def start(self) -> Any:
if self.running:
logger.warning("PoolQuotaProbeScheduler already running")
return
self.running = True
logger.info(
"PoolQuotaProbeScheduler started: scan={}s, max_keys_per_provider={}",
self.scan_interval_seconds,
self.max_keys_per_provider,
)
scheduler = get_scheduler()
scheduler.add_interval_job(
self._scheduled_probe_check,
seconds=self.scan_interval_seconds,
job_id="pool_quota_probe_check",
name="号池额度主动探测检查",
)
# 启动时立即执行一次,避免首次等待一个轮询周期
await self._run_probe_cycle()
async def stop(self) -> Any:
if not self.running:
return
self.running = False
logger.info("PoolQuotaProbeScheduler stopped")
async def _scheduled_probe_check(self) -> None:
if not self.running:
return
await self._run_probe_cycle()
async def _load_probe_timestamps(
self,
*,
redis_client: Any,
provider_id: str,
key_ids: list[str],
) -> dict[str, int]:
if redis_client is None or not key_ids:
return {}
redis_keys = [_probe_stamp_key(provider_id, key_id) for key_id in key_ids]
try:
values = await redis_client.mget(redis_keys)
except Exception as exc:
logger.debug("PoolQuotaProbeScheduler mget probe stamps failed: {}", exc)
return {}
mapping: dict[str, int] = {}
for key_id, raw in zip(key_ids, values, strict=False):
parsed = _parse_probe_stamp(raw)
if parsed is not None:
mapping[key_id] = parsed
return mapping
async def _mark_probe_timestamps(
self,
*,
redis_client: Any,
provider_id: str,
key_ids: list[str],
now_ts: int,
interval_seconds: int,
) -> None:
if redis_client is None or not key_ids:
return
ttl_seconds = max(interval_seconds * 2, 120)
try:
pipe = redis_client.pipeline(transaction=False)
value = str(now_ts)
for key_id in key_ids:
pipe.set(_probe_stamp_key(provider_id, key_id), value, ex=ttl_seconds)
await pipe.execute()
except Exception as exc:
logger.debug("PoolQuotaProbeScheduler set probe stamps failed: {}", exc)
async def _run_probe_cycle(self) -> None:
now_ts = int(time.time())
redis_client = await get_redis_client(require_redis=False)
# 第一阶段:用一个短生命周期 session 查出需要探测的 provider / key 信息
probe_tasks: list[_ProviderProbeTask] = []
db = create_session()
try:
providers = db.query(Provider).filter(Provider.is_active == True).all() # noqa: E712
for provider in providers:
provider_id = str(getattr(provider, "id", "") or "")
provider_type = normalize_provider_type(getattr(provider, "provider_type", ""))
if not provider_id or provider_type not in _SUPPORTED_PROVIDER_TYPES:
continue
pool_cfg = parse_pool_config(getattr(provider, "config", None))
if pool_cfg is None or not pool_cfg.probing_enabled:
continue
interval_minutes = _normalize_probe_interval_minutes(
pool_cfg.probing_interval_minutes
)
interval_seconds = interval_minutes * 60
keys = (
db.query(ProviderAPIKey)
.filter(
ProviderAPIKey.provider_id == provider_id,
ProviderAPIKey.is_active == True, # noqa: E712
)
.all()
)
if not keys:
continue
key_ids = [str(key.id) for key in keys if getattr(key, "id", None)]
probe_stamps = await self._load_probe_timestamps(
redis_client=redis_client,
provider_id=provider_id,
key_ids=key_ids,
)
probe_key_ids = _select_probe_key_ids(
keys=keys,
provider_type=provider_type,
now_ts=now_ts,
interval_seconds=interval_seconds,
last_probe_timestamps=probe_stamps,
limit=self.max_keys_per_provider,
)
if not probe_key_ids:
continue
probe_tasks.append(
_ProviderProbeTask(
provider_id=provider_id,
provider_type=provider_type,
probe_key_ids=probe_key_ids,
interval_seconds=interval_seconds,
)
)
finally:
db.close()
# 第二阶段:每个 provider 使用独立 session 执行探测
for task in probe_tasks:
# 先写探测节流时间戳,避免异常时高频重入
await self._mark_probe_timestamps(
redis_client=redis_client,
provider_id=task.provider_id,
key_ids=task.probe_key_ids,
now_ts=now_ts,
interval_seconds=task.interval_seconds,
)
probe_db = create_session()
try:
result = await refresh_provider_quota_for_provider(
db=probe_db,
provider_id=task.provider_id,
codex_wham_usage_url=_CODEX_WHAM_USAGE_URL,
key_ids=task.probe_key_ids,
)
logger.info(
"[POOL_PROBE] Provider {} ({}) 静默探测完成: selected={}, success={}, failed={}",
task.provider_id[:8],
task.provider_type,
len(task.probe_key_ids),
int(result.get("success") or 0),
int(result.get("failed") or 0),
)
except Exception as exc:
try:
probe_db.rollback()
except Exception:
pass
logger.warning(
"[POOL_PROBE] Provider {} ({}) 静默探测失败: {}",
task.provider_id[:8],
task.provider_type,
exc,
)
finally:
probe_db.close()
_pool_quota_probe_scheduler: PoolQuotaProbeScheduler | None = None
def get_pool_quota_probe_scheduler() -> PoolQuotaProbeScheduler:
global _pool_quota_probe_scheduler
if _pool_quota_probe_scheduler is None:
_pool_quota_probe_scheduler = PoolQuotaProbeScheduler()
return _pool_quota_probe_scheduler
__all__ = [
"PoolQuotaProbeScheduler",
"get_pool_quota_probe_scheduler",
"_select_probe_key_ids",
]