mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
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:
@@ -644,22 +644,16 @@ class AdminListSchedulingPresetsAdapter(AdminApiAdapter):
|
||||
class AdminPoolOverviewAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
providers = (
|
||||
db.query(Provider)
|
||||
.filter(Provider.is_active.is_(True))
|
||||
.order_by(Provider.provider_priority.asc())
|
||||
.all()
|
||||
)
|
||||
providers = db.query(Provider).order_by(Provider.provider_priority.asc()).all()
|
||||
|
||||
# 先批量计算池化 Provider 的 Key 总数/启用数,避免每个 Provider 单独查询(N+1)。
|
||||
# 仅保留号池调度已开启的 Provider。
|
||||
enabled_providers: list[Provider] = []
|
||||
pool_provider_ids: list[str] = []
|
||||
pool_enabled_map: dict[str, bool] = {}
|
||||
for p in providers:
|
||||
pid = str(p.id)
|
||||
enabled = parse_pool_config(getattr(p, "config", None)) is not None
|
||||
pool_enabled_map[pid] = enabled
|
||||
if enabled:
|
||||
pool_provider_ids.append(pid)
|
||||
if parse_pool_config(getattr(p, "config", None)) is None:
|
||||
continue
|
||||
enabled_providers.append(p)
|
||||
pool_provider_ids.append(str(p.id))
|
||||
|
||||
key_ids_by_provider: dict[str, list[str]] = {pid: [] for pid in pool_provider_ids}
|
||||
key_stats_by_provider: dict[str, dict[str, int]] = {
|
||||
@@ -710,19 +704,8 @@ class AdminPoolOverviewAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
items: list[PoolOverviewItem] = []
|
||||
for p in providers:
|
||||
for p in enabled_providers:
|
||||
pid = str(p.id)
|
||||
if not pool_enabled_map.get(pid, False):
|
||||
items.append(
|
||||
PoolOverviewItem(
|
||||
provider_id=pid,
|
||||
provider_name=p.name,
|
||||
provider_type=str(getattr(p, "provider_type", "custom") or "custom"),
|
||||
pool_enabled=False,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
key_stats = key_stats_by_provider.get(pid, {"total": 0, "active": 0})
|
||||
|
||||
items.append(
|
||||
|
||||
@@ -40,6 +40,7 @@ from src.core.provider_templates.types import ProviderType
|
||||
from src.database import create_session
|
||||
from src.database.database import get_db
|
||||
from src.models.database import Provider, ProviderAPIKey, User
|
||||
from src.services.provider.pool.config import parse_pool_config
|
||||
from src.utils.auth_utils import require_admin
|
||||
|
||||
router = APIRouter(prefix="/api/admin/provider-oauth", tags=["Provider OAuth"])
|
||||
@@ -1879,6 +1880,7 @@ async def _batch_import_standard_oauth_internal(
|
||||
proxy_config: dict[str, Any] | None = None,
|
||||
key_proxy: dict[str, Any] | None = None,
|
||||
progress_hook: BatchImportProgressHook | None = None,
|
||||
concurrency: int = 1,
|
||||
) -> BatchImportResponse:
|
||||
"""标准 OAuth Provider 批量导入(不含 Kiro)。"""
|
||||
template = _require_oauth_template(provider_type)
|
||||
@@ -1893,238 +1895,253 @@ async def _batch_import_standard_oauth_internal(
|
||||
is_json = "anthropic.com" in token_url
|
||||
scope_str = " ".join(template.oauth.scopes) if template.oauth.scopes else ""
|
||||
|
||||
results: list[BatchImportResultItem] = []
|
||||
total = len(import_entries)
|
||||
results: list[BatchImportResultItem] = [None] * total # type: ignore[list-item]
|
||||
success_count = 0
|
||||
failed_count = 0
|
||||
processed_count = 0
|
||||
db_lock = asyncio.Lock()
|
||||
sem = asyncio.Semaphore(max(concurrency, 1))
|
||||
|
||||
for idx, import_entry in enumerate(import_entries):
|
||||
refresh_token = import_entry.get("refresh_token", "")
|
||||
async def _process_entry(idx: int, import_entry: dict[str, Any]) -> None:
|
||||
nonlocal success_count, failed_count, processed_count
|
||||
result_item: BatchImportResultItem
|
||||
try:
|
||||
if not refresh_token or len(refresh_token) < 10:
|
||||
|
||||
async with sem:
|
||||
try:
|
||||
refresh_token = import_entry.get("refresh_token", "")
|
||||
if not refresh_token or len(refresh_token) < 10:
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error="Token 无效或过短",
|
||||
)
|
||||
failed_count += 1
|
||||
else:
|
||||
if is_json:
|
||||
body: dict[str, Any] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": refresh_token,
|
||||
}
|
||||
if scope_str:
|
||||
body["scope"] = scope_str
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = None
|
||||
json_body = body
|
||||
else:
|
||||
form: dict[str, str] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": refresh_token,
|
||||
}
|
||||
if scope_str:
|
||||
form["scope"] = scope_str
|
||||
if template.oauth.client_secret:
|
||||
form["client_secret"] = template.oauth.client_secret
|
||||
headers = {
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = form
|
||||
json_body = None
|
||||
|
||||
try:
|
||||
resp = await post_oauth_token(
|
||||
provider_type=provider_type,
|
||||
token_url=token_url,
|
||||
headers=headers,
|
||||
data=data,
|
||||
json_body=json_body,
|
||||
proxy_config=proxy_config,
|
||||
timeout_seconds=timeout_seconds,
|
||||
)
|
||||
except Exception as exc:
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=f"Token 刷新请求失败: {exc}",
|
||||
)
|
||||
failed_count += 1
|
||||
processed_count += 1
|
||||
results[idx] = result_item
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
total, processed_count, success_count, failed_count, result_item
|
||||
)
|
||||
return
|
||||
|
||||
if resp.status_code < 200 or resp.status_code >= 300:
|
||||
error_reason = f"HTTP {resp.status_code}"
|
||||
try:
|
||||
error_body = resp.json()
|
||||
if "error" in error_body:
|
||||
error_reason = str(
|
||||
error_body.get("error_description") or error_body.get("error")
|
||||
)
|
||||
except Exception:
|
||||
error_reason = (
|
||||
resp.text[:100] if resp.text else f"HTTP {resp.status_code}"
|
||||
)
|
||||
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=f"Token 验证失败: {error_reason}",
|
||||
)
|
||||
failed_count += 1
|
||||
processed_count += 1
|
||||
results[idx] = result_item
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
total, processed_count, success_count, failed_count, result_item
|
||||
)
|
||||
return
|
||||
|
||||
token_data = resp.json()
|
||||
access_token = str(token_data.get("access_token") or "")
|
||||
new_refresh_token = str(token_data.get("refresh_token") or "") or refresh_token
|
||||
|
||||
if not access_token:
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error="Token 刷新返回缺少 access_token",
|
||||
)
|
||||
failed_count += 1
|
||||
processed_count += 1
|
||||
results[idx] = result_item
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
total, processed_count, success_count, failed_count, result_item
|
||||
)
|
||||
return
|
||||
|
||||
expires_in = token_data.get("expires_in")
|
||||
expires_at: int | None = None
|
||||
try:
|
||||
if expires_in is not None:
|
||||
expires_at = int(time.time()) + int(expires_in)
|
||||
except Exception:
|
||||
expires_at = None
|
||||
|
||||
auth_config: dict[str, Any] = {
|
||||
"provider_type": provider_type,
|
||||
"token_type": token_data.get("token_type"),
|
||||
"refresh_token": new_refresh_token or None,
|
||||
"expires_at": expires_at,
|
||||
"scope": token_data.get("scope"),
|
||||
"updated_at": int(time.time()),
|
||||
}
|
||||
|
||||
try:
|
||||
auth_config = await enrich_auth_config(
|
||||
provider_type=provider_type,
|
||||
auth_config=auth_config,
|
||||
token_response=token_data,
|
||||
access_token=access_token,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("批量导入: enrich_auth_config 失败 (index={}): {}", idx, exc)
|
||||
|
||||
if provider_type == ProviderType.CODEX.value:
|
||||
_apply_codex_import_hints(auth_config, import_entry)
|
||||
|
||||
async with db_lock:
|
||||
try:
|
||||
existing_key = _check_duplicate_oauth_account(
|
||||
db, provider_id, auth_config
|
||||
)
|
||||
except InvalidRequestException as exc:
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=str(exc),
|
||||
)
|
||||
failed_count += 1
|
||||
processed_count += 1
|
||||
results[idx] = result_item
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
total,
|
||||
processed_count,
|
||||
success_count,
|
||||
failed_count,
|
||||
result_item,
|
||||
)
|
||||
return
|
||||
|
||||
replaced = False
|
||||
if existing_key:
|
||||
new_key = _update_existing_oauth_key(
|
||||
db,
|
||||
existing_key,
|
||||
access_token,
|
||||
auth_config,
|
||||
flush_only=True,
|
||||
proxy=key_proxy,
|
||||
)
|
||||
name = existing_key.name
|
||||
replaced = True
|
||||
else:
|
||||
email = auth_config.get("email")
|
||||
if email:
|
||||
name = f"{provider_type}_{email}"
|
||||
else:
|
||||
name = f"{provider_type}_{int(time.time())}_{idx}"
|
||||
if len(name) > 100:
|
||||
name = name[:100]
|
||||
|
||||
new_key = _create_oauth_key(
|
||||
db,
|
||||
provider_id=provider_id,
|
||||
name=name,
|
||||
access_token=access_token,
|
||||
auth_config=auth_config,
|
||||
api_formats=api_formats,
|
||||
flush_only=True,
|
||||
proxy=key_proxy,
|
||||
)
|
||||
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="success",
|
||||
key_id=str(new_key.id),
|
||||
key_name=name,
|
||||
replaced=replaced,
|
||||
)
|
||||
success_count += 1
|
||||
|
||||
except Exception as exc:
|
||||
logger.error("批量导入 OAuth 凭据失败 (index={}): {}", idx, exc)
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error="Token 无效或过短",
|
||||
error=f"导入失败: {exc}",
|
||||
)
|
||||
failed_count += 1
|
||||
else:
|
||||
if is_json:
|
||||
body: dict[str, Any] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": refresh_token,
|
||||
}
|
||||
if scope_str:
|
||||
body["scope"] = scope_str
|
||||
headers = {"Content-Type": "application/json", "Accept": "application/json"}
|
||||
data = None
|
||||
json_body = body
|
||||
else:
|
||||
form: dict[str, str] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": refresh_token,
|
||||
}
|
||||
if scope_str:
|
||||
form["scope"] = scope_str
|
||||
if template.oauth.client_secret:
|
||||
form["client_secret"] = template.oauth.client_secret
|
||||
headers = {
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = form
|
||||
json_body = None
|
||||
|
||||
try:
|
||||
resp = await post_oauth_token(
|
||||
provider_type=provider_type,
|
||||
token_url=token_url,
|
||||
headers=headers,
|
||||
data=data,
|
||||
json_body=json_body,
|
||||
proxy_config=proxy_config,
|
||||
timeout_seconds=timeout_seconds,
|
||||
)
|
||||
except Exception as exc:
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=f"Token 刷新请求失败: {exc}",
|
||||
)
|
||||
failed_count += 1
|
||||
results.append(result_item)
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
len(import_entries),
|
||||
idx + 1,
|
||||
success_count,
|
||||
failed_count,
|
||||
result_item,
|
||||
)
|
||||
continue
|
||||
|
||||
if resp.status_code < 200 or resp.status_code >= 300:
|
||||
error_reason = f"HTTP {resp.status_code}"
|
||||
try:
|
||||
error_body = resp.json()
|
||||
if "error" in error_body:
|
||||
error_reason = str(
|
||||
error_body.get("error_description") or error_body.get("error")
|
||||
)
|
||||
except Exception:
|
||||
error_reason = resp.text[:100] if resp.text else f"HTTP {resp.status_code}"
|
||||
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=f"Token 验证失败: {error_reason}",
|
||||
)
|
||||
failed_count += 1
|
||||
results.append(result_item)
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
len(import_entries),
|
||||
idx + 1,
|
||||
success_count,
|
||||
failed_count,
|
||||
result_item,
|
||||
)
|
||||
continue
|
||||
|
||||
token_data = resp.json()
|
||||
access_token = str(token_data.get("access_token") or "")
|
||||
new_refresh_token = str(token_data.get("refresh_token") or "") or refresh_token
|
||||
|
||||
if not access_token:
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error="Token 刷新返回缺少 access_token",
|
||||
)
|
||||
failed_count += 1
|
||||
results.append(result_item)
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
len(import_entries),
|
||||
idx + 1,
|
||||
success_count,
|
||||
failed_count,
|
||||
result_item,
|
||||
)
|
||||
continue
|
||||
|
||||
expires_in = token_data.get("expires_in")
|
||||
expires_at: int | None = None
|
||||
try:
|
||||
if expires_in is not None:
|
||||
expires_at = int(time.time()) + int(expires_in)
|
||||
except Exception:
|
||||
expires_at = None
|
||||
|
||||
auth_config: dict[str, Any] = {
|
||||
"provider_type": provider_type,
|
||||
"token_type": token_data.get("token_type"),
|
||||
"refresh_token": new_refresh_token or None,
|
||||
"expires_at": expires_at,
|
||||
"scope": token_data.get("scope"),
|
||||
"updated_at": int(time.time()),
|
||||
}
|
||||
|
||||
try:
|
||||
auth_config = await enrich_auth_config(
|
||||
provider_type=provider_type,
|
||||
auth_config=auth_config,
|
||||
token_response=token_data,
|
||||
access_token=access_token,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("批量导入: enrich_auth_config 失败 (index={}): {}", idx, exc)
|
||||
|
||||
if provider_type == ProviderType.CODEX.value:
|
||||
_apply_codex_import_hints(auth_config, import_entry)
|
||||
|
||||
try:
|
||||
existing_key = _check_duplicate_oauth_account(db, provider_id, auth_config)
|
||||
except InvalidRequestException as exc:
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=str(exc),
|
||||
)
|
||||
failed_count += 1
|
||||
results.append(result_item)
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
len(import_entries),
|
||||
idx + 1,
|
||||
success_count,
|
||||
failed_count,
|
||||
result_item,
|
||||
)
|
||||
continue
|
||||
|
||||
replaced = False
|
||||
if existing_key:
|
||||
new_key = _update_existing_oauth_key(
|
||||
db,
|
||||
existing_key,
|
||||
access_token,
|
||||
auth_config,
|
||||
flush_only=True,
|
||||
proxy=key_proxy,
|
||||
)
|
||||
name = existing_key.name
|
||||
replaced = True
|
||||
else:
|
||||
email = auth_config.get("email")
|
||||
if email:
|
||||
name = f"{provider_type}_{email}"
|
||||
else:
|
||||
name = f"{provider_type}_{int(time.time())}_{idx}"
|
||||
if len(name) > 100:
|
||||
name = name[:100]
|
||||
|
||||
new_key = _create_oauth_key(
|
||||
db,
|
||||
provider_id=provider_id,
|
||||
name=name,
|
||||
access_token=access_token,
|
||||
auth_config=auth_config,
|
||||
api_formats=api_formats,
|
||||
flush_only=True,
|
||||
proxy=key_proxy,
|
||||
)
|
||||
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="success",
|
||||
key_id=str(new_key.id),
|
||||
key_name=name,
|
||||
replaced=replaced,
|
||||
)
|
||||
success_count += 1
|
||||
|
||||
except Exception as exc:
|
||||
logger.error("批量导入 OAuth 凭据失败 (index={}): {}", idx, exc)
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=f"导入失败: {exc}",
|
||||
)
|
||||
failed_count += 1
|
||||
|
||||
results.append(result_item)
|
||||
processed_count += 1
|
||||
results[idx] = result_item
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
len(import_entries), idx + 1, success_count, failed_count, result_item
|
||||
)
|
||||
await progress_hook(total, processed_count, success_count, failed_count, result_item)
|
||||
|
||||
await asyncio.gather(
|
||||
*[_process_entry(i, e) for i, e in enumerate(import_entries)],
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
if success_count > 0:
|
||||
db.commit()
|
||||
|
||||
final_results = [r for r in results if r is not None]
|
||||
if len(final_results) != total:
|
||||
logger.warning("[BATCH_IMPORT] 结果不完整: expected={}, got={}", total, len(final_results))
|
||||
|
||||
logger.info(
|
||||
"[BATCH_IMPORT] Provider {} ({}): 成功 {}/{}, 失败 {}",
|
||||
provider_id,
|
||||
@@ -2138,7 +2155,7 @@ async def _batch_import_standard_oauth_internal(
|
||||
total=len(import_entries),
|
||||
success=success_count,
|
||||
failed=failed_count,
|
||||
results=results,
|
||||
results=final_results,
|
||||
)
|
||||
|
||||
|
||||
@@ -2174,6 +2191,10 @@ async def batch_import_oauth(
|
||||
getattr(provider, "proxy", None), payload.proxy_node_id
|
||||
)
|
||||
|
||||
# 从 pool_advanced 读取批量并发数
|
||||
_pool_cfg = parse_pool_config(getattr(provider, "config", None))
|
||||
batch_concurrency = (_pool_cfg.batch_concurrency or 8) if _pool_cfg else 8
|
||||
|
||||
if provider_type == ProviderType.KIRO.value:
|
||||
return await _batch_import_kiro_internal(
|
||||
provider_id=provider_id,
|
||||
@@ -2182,6 +2203,7 @@ async def batch_import_oauth(
|
||||
db=db,
|
||||
proxy_config=proxy_config,
|
||||
key_proxy=key_proxy,
|
||||
concurrency=batch_concurrency,
|
||||
)
|
||||
|
||||
return await _batch_import_standard_oauth_internal(
|
||||
@@ -2192,6 +2214,7 @@ async def batch_import_oauth(
|
||||
db=db,
|
||||
proxy_config=proxy_config,
|
||||
key_proxy=key_proxy,
|
||||
concurrency=batch_concurrency,
|
||||
)
|
||||
|
||||
|
||||
@@ -2224,6 +2247,10 @@ async def _run_batch_import_task(
|
||||
getattr(provider, "proxy", None), payload.proxy_node_id
|
||||
)
|
||||
|
||||
# 从 pool_advanced 读取批量并发数
|
||||
_pool_cfg = parse_pool_config(getattr(provider, "config", None))
|
||||
batch_concurrency = (_pool_cfg.batch_concurrency or 8) if _pool_cfg else 8
|
||||
|
||||
async def progress_hook(
|
||||
total: int,
|
||||
processed: int,
|
||||
@@ -2261,6 +2288,7 @@ async def _run_batch_import_task(
|
||||
proxy_config=proxy_config,
|
||||
key_proxy=key_proxy,
|
||||
progress_hook=progress_hook,
|
||||
concurrency=batch_concurrency,
|
||||
)
|
||||
else:
|
||||
result = await _batch_import_standard_oauth_internal(
|
||||
@@ -2272,6 +2300,7 @@ async def _run_batch_import_task(
|
||||
proxy_config=proxy_config,
|
||||
key_proxy=key_proxy,
|
||||
progress_hook=progress_hook,
|
||||
concurrency=batch_concurrency,
|
||||
)
|
||||
|
||||
state["status"] = "completed"
|
||||
@@ -2391,12 +2420,14 @@ async def _batch_import_kiro_internal(
|
||||
proxy_config: dict[str, Any] | None = None,
|
||||
key_proxy: dict[str, Any] | None = None,
|
||||
progress_hook: BatchImportProgressHook | None = None,
|
||||
concurrency: int = 1,
|
||||
) -> BatchImportResponse:
|
||||
"""Kiro 批量导入内部实现(供通用端点调用)。
|
||||
|
||||
Args:
|
||||
proxy_config: 本次操作使用的代理配置(已由调用方解析)
|
||||
key_proxy: 需要保存到 Key 上的代理配置
|
||||
concurrency: 并发数(从 pool_advanced.batch_concurrency 读取)
|
||||
"""
|
||||
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
|
||||
from src.services.provider.adapters.kiro.token_manager import refresh_access_token
|
||||
@@ -2409,141 +2440,154 @@ async def _batch_import_kiro_internal(
|
||||
|
||||
api_formats = _get_provider_api_formats(provider)
|
||||
|
||||
results: list[BatchImportResultItem] = []
|
||||
total = len(credentials)
|
||||
results: list[BatchImportResultItem] = [None] * total # type: ignore[list-item]
|
||||
success_count = 0
|
||||
failed_count = 0
|
||||
processed_count = 0
|
||||
db_lock = asyncio.Lock()
|
||||
sem = asyncio.Semaphore(max(concurrency, 1))
|
||||
|
||||
for idx, cred in enumerate(credentials):
|
||||
async def _process_entry(idx: int, cred: dict[str, Any]) -> None:
|
||||
nonlocal success_count, failed_count, processed_count
|
||||
result_item: BatchImportResultItem
|
||||
try:
|
||||
is_valid, error_msg = KiroAuthConfig.validate_required_fields(cred)
|
||||
if not is_valid:
|
||||
|
||||
async with sem:
|
||||
try:
|
||||
is_valid, error_msg = KiroAuthConfig.validate_required_fields(cred)
|
||||
if not is_valid:
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=error_msg,
|
||||
)
|
||||
failed_count += 1
|
||||
processed_count += 1
|
||||
results[idx] = result_item
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
total, processed_count, success_count, failed_count, result_item
|
||||
)
|
||||
return
|
||||
|
||||
cfg = KiroAuthConfig.from_dict(cred)
|
||||
cfg.provider_type = ProviderType.KIRO.value
|
||||
|
||||
try:
|
||||
access_token, new_cfg = await refresh_access_token(
|
||||
cfg,
|
||||
proxy_config=proxy_config,
|
||||
timeout_seconds=timeout_seconds,
|
||||
)
|
||||
except Exception as exc:
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=f"Token 验证失败: {exc}",
|
||||
)
|
||||
failed_count += 1
|
||||
processed_count += 1
|
||||
results[idx] = result_item
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
total, processed_count, success_count, failed_count, result_item
|
||||
)
|
||||
return
|
||||
|
||||
email = await _fetch_kiro_email(new_cfg.to_dict(), proxy_config=proxy_config)
|
||||
if email and not new_cfg.email:
|
||||
new_cfg.email = email
|
||||
|
||||
async with db_lock:
|
||||
try:
|
||||
existing_key = _check_duplicate_oauth_account(
|
||||
db, provider_id, new_cfg.to_dict()
|
||||
)
|
||||
except InvalidRequestException as exc:
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=str(exc),
|
||||
)
|
||||
failed_count += 1
|
||||
processed_count += 1
|
||||
results[idx] = result_item
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
total,
|
||||
processed_count,
|
||||
success_count,
|
||||
failed_count,
|
||||
result_item,
|
||||
)
|
||||
return
|
||||
|
||||
replaced = False
|
||||
if existing_key:
|
||||
new_key = _update_existing_oauth_key(
|
||||
db,
|
||||
existing_key,
|
||||
access_token,
|
||||
new_cfg.to_dict(),
|
||||
flush_only=True,
|
||||
proxy=key_proxy,
|
||||
)
|
||||
name = existing_key.name
|
||||
replaced = True
|
||||
else:
|
||||
name = _build_kiro_key_name(
|
||||
email, new_cfg.auth_method, new_cfg.refresh_token
|
||||
)
|
||||
new_key = _create_oauth_key(
|
||||
db,
|
||||
provider_id=provider_id,
|
||||
name=name,
|
||||
access_token=access_token,
|
||||
auth_config=new_cfg.to_dict(),
|
||||
api_formats=api_formats,
|
||||
flush_only=True,
|
||||
proxy=key_proxy,
|
||||
)
|
||||
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=error_msg,
|
||||
status="success",
|
||||
key_id=str(new_key.id),
|
||||
key_name=name,
|
||||
auth_method=new_cfg.auth_method or "social",
|
||||
replaced=replaced,
|
||||
)
|
||||
failed_count += 1
|
||||
results.append(result_item)
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
len(credentials),
|
||||
idx + 1,
|
||||
success_count,
|
||||
failed_count,
|
||||
result_item,
|
||||
)
|
||||
continue
|
||||
success_count += 1
|
||||
|
||||
cfg = KiroAuthConfig.from_dict(cred)
|
||||
cfg.provider_type = ProviderType.KIRO.value
|
||||
|
||||
try:
|
||||
access_token, new_cfg = await refresh_access_token(
|
||||
cfg,
|
||||
proxy_config=proxy_config,
|
||||
timeout_seconds=timeout_seconds,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error("批量导入 Kiro 凭据失败 (index={}): {}", idx, exc)
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=f"Token 验证失败: {exc}",
|
||||
error=f"导入失败: {exc}",
|
||||
)
|
||||
failed_count += 1
|
||||
results.append(result_item)
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
len(credentials),
|
||||
idx + 1,
|
||||
success_count,
|
||||
failed_count,
|
||||
result_item,
|
||||
)
|
||||
continue
|
||||
|
||||
email = await _fetch_kiro_email(new_cfg.to_dict(), proxy_config=proxy_config)
|
||||
if email and not new_cfg.email:
|
||||
new_cfg.email = email
|
||||
|
||||
try:
|
||||
existing_key = _check_duplicate_oauth_account(db, provider_id, new_cfg.to_dict())
|
||||
except InvalidRequestException as exc:
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=str(exc),
|
||||
)
|
||||
failed_count += 1
|
||||
results.append(result_item)
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
len(credentials),
|
||||
idx + 1,
|
||||
success_count,
|
||||
failed_count,
|
||||
result_item,
|
||||
)
|
||||
continue
|
||||
|
||||
replaced = False
|
||||
if existing_key:
|
||||
new_key = _update_existing_oauth_key(
|
||||
db,
|
||||
existing_key,
|
||||
access_token,
|
||||
new_cfg.to_dict(),
|
||||
flush_only=True,
|
||||
proxy=key_proxy,
|
||||
)
|
||||
name = existing_key.name
|
||||
replaced = True
|
||||
else:
|
||||
name = _build_kiro_key_name(email, new_cfg.auth_method, new_cfg.refresh_token)
|
||||
new_key = _create_oauth_key(
|
||||
db,
|
||||
provider_id=provider_id,
|
||||
name=name,
|
||||
access_token=access_token,
|
||||
auth_config=new_cfg.to_dict(),
|
||||
api_formats=api_formats,
|
||||
flush_only=True,
|
||||
proxy=key_proxy,
|
||||
)
|
||||
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="success",
|
||||
key_id=str(new_key.id),
|
||||
key_name=name,
|
||||
auth_method=new_cfg.auth_method or "social",
|
||||
replaced=replaced,
|
||||
)
|
||||
success_count += 1
|
||||
|
||||
except Exception as exc:
|
||||
logger.error("批量导入 Kiro 凭据失败 (index={}): {}", idx, exc)
|
||||
result_item = BatchImportResultItem(
|
||||
index=idx,
|
||||
status="error",
|
||||
error=f"导入失败: {exc}",
|
||||
)
|
||||
failed_count += 1
|
||||
|
||||
results.append(result_item)
|
||||
processed_count += 1
|
||||
results[idx] = result_item
|
||||
if progress_hook is not None:
|
||||
await progress_hook(
|
||||
len(credentials),
|
||||
idx + 1,
|
||||
success_count,
|
||||
failed_count,
|
||||
result_item,
|
||||
)
|
||||
await progress_hook(total, processed_count, success_count, failed_count, result_item)
|
||||
|
||||
await asyncio.gather(
|
||||
*[_process_entry(i, c) for i, c in enumerate(credentials)],
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
# 提交所有成功的记录
|
||||
if success_count > 0:
|
||||
db.commit()
|
||||
|
||||
final_results = [r for r in results if r is not None]
|
||||
if len(final_results) != total:
|
||||
logger.warning(
|
||||
"[KIRO_BATCH_IMPORT] 结果不完整: expected={}, got={}", total, len(final_results)
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"[KIRO_BATCH_IMPORT] Provider {}: 成功 {}/{}, 失败 {}",
|
||||
provider_id,
|
||||
@@ -2556,7 +2600,7 @@ async def _batch_import_kiro_internal(
|
||||
total=len(credentials),
|
||||
success=success_count,
|
||||
failed=failed_count,
|
||||
results=results,
|
||||
results=final_results,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -77,8 +77,6 @@ from src.models.database import (
|
||||
User,
|
||||
)
|
||||
from src.services.provider.behavior import get_provider_behavior
|
||||
from src.services.provider.fingerprint import ensure_key_fingerprint
|
||||
from src.services.provider.request_context import set_current_fingerprint
|
||||
from src.services.provider.stream_policy import (
|
||||
enforce_stream_mode_for_upstream,
|
||||
get_upstream_stream_policy,
|
||||
@@ -721,8 +719,6 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
else:
|
||||
request_body = dict(original_request_body)
|
||||
|
||||
set_current_fingerprint(ensure_key_fingerprint(key, persist_if_missing=True))
|
||||
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=provider_type,
|
||||
@@ -751,6 +747,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
key_id=str(getattr(key, "id", "") or ""),
|
||||
is_stream=upstream_is_stream,
|
||||
provider_id=str(getattr(provider, "id", "") or ""),
|
||||
key=key,
|
||||
)
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
|
||||
@@ -40,8 +40,6 @@ from src.core.exceptions import (
|
||||
)
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.behavior import get_provider_behavior
|
||||
from src.services.provider.fingerprint import ensure_key_fingerprint
|
||||
from src.services.provider.request_context import set_current_fingerprint
|
||||
from src.services.provider.stream_policy import (
|
||||
enforce_stream_mode_for_upstream,
|
||||
get_upstream_stream_policy,
|
||||
@@ -327,8 +325,6 @@ class CliStreamMixin:
|
||||
)
|
||||
ctx.needs_conversion = needs_conversion
|
||||
|
||||
set_current_fingerprint(ensure_key_fingerprint(key, persist_if_missing=True))
|
||||
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=provider_type,
|
||||
@@ -357,6 +353,7 @@ class CliStreamMixin:
|
||||
key_id=str(getattr(key, "id", "") or ""),
|
||||
is_stream=upstream_is_stream,
|
||||
provider_id=str(getattr(provider, "id", "") or ""),
|
||||
key=key,
|
||||
)
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
|
||||
@@ -33,8 +33,6 @@ from src.core.exceptions import (
|
||||
)
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.behavior import get_provider_behavior
|
||||
from src.services.provider.fingerprint import ensure_key_fingerprint
|
||||
from src.services.provider.request_context import set_current_fingerprint
|
||||
from src.services.provider.stream_policy import (
|
||||
enforce_stream_mode_for_upstream,
|
||||
get_upstream_stream_policy,
|
||||
@@ -143,8 +141,6 @@ class CliSyncMixin:
|
||||
)
|
||||
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
|
||||
|
||||
set_current_fingerprint(ensure_key_fingerprint(key, persist_if_missing=True))
|
||||
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
behavior = get_provider_behavior(
|
||||
provider_type=provider_type,
|
||||
@@ -173,6 +169,7 @@ class CliSyncMixin:
|
||||
key_id=str(getattr(key, "id", "") or ""),
|
||||
is_stream=upstream_is_stream,
|
||||
provider_id=str(getattr(provider, "id", "") or ""),
|
||||
key=key,
|
||||
)
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
|
||||
18
src/main.py
18
src/main.py
@@ -235,6 +235,9 @@ async def lifespan(app: FastAPI) -> Any:
|
||||
# 启动月卡额度重置调度器(仅一个 worker 执行)
|
||||
logger.info("启动月卡额度重置调度器...")
|
||||
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
|
||||
from src.services.provider_keys.pool_quota_probe_scheduler import (
|
||||
get_pool_quota_probe_scheduler,
|
||||
)
|
||||
from src.services.system.maintenance_scheduler import get_maintenance_scheduler
|
||||
from src.services.task.task_poller import get_task_poller
|
||||
from src.services.usage.quota_scheduler import get_quota_scheduler
|
||||
@@ -243,6 +246,7 @@ async def lifespan(app: FastAPI) -> Any:
|
||||
quota_scheduler = get_quota_scheduler()
|
||||
maintenance_scheduler = get_maintenance_scheduler()
|
||||
model_fetch_scheduler = get_model_fetch_scheduler()
|
||||
pool_quota_probe_scheduler = get_pool_quota_probe_scheduler()
|
||||
task_poller = get_task_poller()
|
||||
task_coordinator = StartupTaskCoordinator(redis_client)
|
||||
|
||||
@@ -272,6 +276,15 @@ async def lifespan(app: FastAPI) -> Any:
|
||||
logger.info("检测到其他 worker 已运行模型获取调度器,本实例跳过")
|
||||
model_fetch_scheduler = None # type: ignore[assignment]
|
||||
|
||||
# 启动号池额度主动探测调度器
|
||||
pool_quota_probe_scheduler_active = await task_coordinator.acquire("pool_quota_probe_scheduler")
|
||||
if pool_quota_probe_scheduler_active:
|
||||
logger.info("启动号池额度主动探测调度器...")
|
||||
await pool_quota_probe_scheduler.start()
|
||||
else:
|
||||
logger.info("检测到其他 worker 已运行号池额度主动探测调度器,本实例跳过")
|
||||
pool_quota_probe_scheduler = None # type: ignore[assignment]
|
||||
|
||||
# 启动异步任务轮询服务(当前仅视频)
|
||||
task_poller_active = await task_coordinator.acquire("task_poller:video")
|
||||
if task_poller_active:
|
||||
@@ -339,6 +352,11 @@ async def lifespan(app: FastAPI) -> Any:
|
||||
await model_fetch_scheduler.stop()
|
||||
await task_coordinator.release("model_fetch_scheduler")
|
||||
|
||||
if pool_quota_probe_scheduler:
|
||||
logger.info("停止号池额度主动探测调度器...")
|
||||
await pool_quota_probe_scheduler.stop()
|
||||
await task_coordinator.release("pool_quota_probe_scheduler")
|
||||
|
||||
if task_poller:
|
||||
logger.info("停止 TaskPoller(video)...")
|
||||
await task_poller.stop()
|
||||
|
||||
@@ -283,6 +283,20 @@ class PoolAdvancedConfig(BaseModel):
|
||||
None,
|
||||
description="关键词临时不可调度规则: [{'keyword': '...', 'duration_minutes': 5}]",
|
||||
)
|
||||
batch_concurrency: int | None = Field(
|
||||
None,
|
||||
ge=1,
|
||||
le=32,
|
||||
description="批量操作并发数(前端批量刷新 OAuth/额度等)。默认 8",
|
||||
)
|
||||
probing_enabled: bool = Field(False, description="启用主动探测(定期检查 Key 可用性)")
|
||||
probing_interval_minutes: int | None = Field(
|
||||
None,
|
||||
ge=1,
|
||||
le=1440,
|
||||
description="主动探测间隔(分钟)。默认 10",
|
||||
)
|
||||
auto_remove_banned_keys: bool = Field(False, description="检测到封号时自动清除账号")
|
||||
|
||||
|
||||
class ClaudeCodeAdvancedConfig(BaseModel):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
|
||||
367
src/services/provider_keys/pool_quota_probe_scheduler.py
Normal file
367
src/services/provider_keys/pool_quota_probe_scheduler.py
Normal 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",
|
||||
]
|
||||
Reference in New Issue
Block a user