Merge pull request #220 from AAEE86/master

fix(pool): recent_refresh 按 provider_type 解析 reset_seconds 并锁定 Codex 周重置语义
This commit is contained in:
fawney19
2026-03-11 16:02:14 +08:00
committed by GitHub
7 changed files with 148 additions and 3 deletions

View File

@@ -5,7 +5,7 @@ from __future__ import annotations
import math
from typing import Any
from src.core.provider_types import ProviderType
from src.core.provider_types import ProviderType, normalize_provider_type
from src.services.provider_keys.quota_reader import get_quota_reader
@@ -108,8 +108,54 @@ def extract_plan_type(key_obj: Any) -> str | None:
return None
def extract_reset_seconds(key_obj: Any) -> float | None:
def _resolve_key_provider_type(key_obj: Any, provider_type: str | None = None) -> str | None:
explicit = normalize_provider_type(provider_type)
if explicit:
return explicit
direct = normalize_provider_type(getattr(key_obj, "provider_type", None))
if direct:
return direct
provider = getattr(key_obj, "provider", None)
related = normalize_provider_type(getattr(provider, "provider_type", None))
if related:
return related
metadata = safe_metadata(key_obj)
candidates = [
provider.value
for provider in (ProviderType.CODEX, ProviderType.KIRO, ProviderType.ANTIGRAVITY)
if isinstance(metadata.get(provider.value), dict)
]
if len(candidates) == 1:
return candidates[0]
return None
def _extract_codex_weekly_reset_seconds(metadata: dict[str, Any]) -> float | None:
codex = metadata.get(ProviderType.CODEX.value)
if not isinstance(codex, dict):
return None
parsed = safe_float(codex.get("primary_reset_seconds"))
if parsed is None or parsed < 0:
return None
return parsed
def extract_reset_seconds(key_obj: Any, provider_type: str | None = None) -> float | None:
metadata = safe_metadata(key_obj)
resolved_provider_type = _resolve_key_provider_type(key_obj, provider_type)
if resolved_provider_type == ProviderType.CODEX:
# Codex metadata 已统一约定primary_* 表示周限额secondary_* 表示 5H 限额。
return _extract_codex_weekly_reset_seconds(metadata)
if resolved_provider_type in (ProviderType.KIRO, ProviderType.ANTIGRAVITY):
return get_quota_reader(resolved_provider_type, metadata).reset_seconds()
candidates: list[float] = []
for provider_type in (ProviderType.CODEX, ProviderType.KIRO, ProviderType.ANTIGRAVITY):

View File

@@ -39,9 +39,10 @@ class RecentRefreshDimension(PresetDimensionBase):
context: dict[str, Any],
mode: str | None,
) -> float:
provider_type = context.get("provider_type")
reset_scores: dict[str, float] = {}
for kid in all_key_ids:
reset_seconds = extract_reset_seconds(keys_by_id.get(kid))
reset_seconds = extract_reset_seconds(keys_by_id.get(kid), provider_type=provider_type)
if reset_seconds is not None:
reset_scores[kid] = reset_seconds
if not reset_scores:

View File

@@ -193,6 +193,7 @@ class PoolManager:
strategy_context.update(
{
"provider_type": provider_type,
"all_key_ids": all_key_ids,
"lru_scores": lru_scores,
"cost_totals": cost_totals,
@@ -561,10 +562,18 @@ class PoolManager:
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 {},