diff --git a/e2e/conftest.py b/e2e/conftest.py index c9eaa51ece..1aed39c52d 100644 --- a/e2e/conftest.py +++ b/e2e/conftest.py @@ -56,9 +56,6 @@ def __init__( ProviderConfig( api_key="browser-test", base_url="https://provider.invalid/v1", - rate_limit=1_000, - rate_window=1, - max_concurrency=100, http_read_timeout=1.0, http_write_timeout=1.0, http_connect_timeout=1.0, @@ -252,15 +249,20 @@ def admin_base_url( ), } - async def fixture_provider(provider_id: str, _settings: Settings) -> BaseProvider: + async def fixture_provider( + provider_id: str, _settings: Settings, _admission_registry + ) -> BaseProvider: if provider_id not in providers: raise AssertionError(f"Missing browser fixture provider: {provider_id}") return providers[provider_id] manager = ProviderRuntimeManager( get_settings(), - runtime_factory=lambda snapshot: ProviderRuntime( - snapshot, dict(providers), provider_constructor=fixture_provider + runtime_factory=lambda snapshot, admission_registry: ProviderRuntime( + snapshot, + admission_registry, + dict(providers), + provider_constructor=fixture_provider, ), ) runtime = ApplicationRuntime( diff --git a/scripts/install.ps1 b/scripts/install.ps1 index f580b08c70..b82c65de2b 100644 --- a/scripts/install.ps1 +++ b/scripts/install.ps1 @@ -1006,7 +1006,7 @@ function Install-Hermes { Invoke-DownloadedPowerShellInstaller ` -Url $HermesInstallUrl ` -Name "Hermes Agent" ` - -ScriptArguments @("-NonInteractive", "-SkipSetup") + -ScriptArguments @("-NonInteractive") Add-KnownBinDirectories } diff --git a/smoke/installers/test_installers.py b/smoke/installers/test_installers.py index 042daa5eb9..9c2801717c 100644 --- a/smoke/installers/test_installers.py +++ b/smoke/installers/test_installers.py @@ -2327,15 +2327,15 @@ def powershell_harness( encoding="utf-8", ) (fixtures / "hermes-installer.ps1").write_text( - r"""param( - [switch] $NonInteractive, - [switch] $SkipSetup + r"""[CmdletBinding(PositionalBinding=$false)] +param( + [switch] $NonInteractive ) if ($env:FAIL_STEP -eq "hermes-install") { exit 65 } $bin = Join-Path $env:LOCALAPPDATA "hermes\hermes-agent\bin" New-Item -ItemType Directory -Force -Path $bin | Out-Null Copy-Item (Join-Path $env:FAKE_FIXTURES "hermes-command.cmd") (Join-Path $bin "hermes.cmd") -Force -Add-Content -LiteralPath $env:CALL_LOG -Value "hermes-install:${NonInteractive}:${SkipSetup}" +Add-Content -LiteralPath $env:CALL_LOG -Value "hermes-install:${NonInteractive}" """, encoding="utf-8", ) @@ -3173,7 +3173,7 @@ def test_install_ps1_fresh_install_is_verified( ) assert calls.index("npm:install -g cline") < calls.index("cline:--version") assert any("hermes-agent.nousresearch.com/install.ps1" in call for call in calls) - assert "hermes-install:True:True" in calls + assert "hermes-install:True" in calls assert calls.index("npm:install -g @deepseek-ai/dsh@0.1.0-rc.8") < calls.index( "dsh:--version" ) @@ -3251,7 +3251,7 @@ def test_install_ps1_discovers_grok_in_custom_bin_directory( ("client", "install_call"), [ ("cline", "npm:install -g cline"), - ("hermes", "hermes-install:True:True"), + ("hermes", "hermes-install:True"), ("grok", "grok-install"), ( "aider", @@ -3538,6 +3538,16 @@ def test_install_ps1_stops_when_selected_dsh_install_fails( _assert_uv_ready_without_fcc_install(powershell_harness.calls()) +def test_install_ps1_installs_hermes_noninteractively( + powershell_harness: PowerShellHarness, +) -> None: + result = powershell_harness.run_functions("Ensure-Hermes") + + assert result.returncode == 0, result.stderr + calls = powershell_harness.calls() + assert calls.index("hermes-install:True") < calls.index("hermes:--version") + + def test_install_ps1_rejects_unsupported_hermes_architecture_before_download( powershell_harness: PowerShellHarness, ) -> None: diff --git a/src/free_claude_code/config/settings.py b/src/free_claude_code/config/settings.py index 6ce7cbfe85..00a00a5c30 100644 --- a/src/free_claude_code/config/settings.py +++ b/src/free_claude_code/config/settings.py @@ -584,12 +584,14 @@ def validate_provider_references(self) -> Settings: default=None, validation_alias="OLLAMA_CLOUD_PROXY" ) # ==================== Provider Rate Limiting ==================== - provider_rate_limit: int = Field(default=1, validation_alias="PROVIDER_RATE_LIMIT") + provider_rate_limit: int = Field( + default=1, gt=0, validation_alias="PROVIDER_RATE_LIMIT" + ) provider_rate_window: int = Field( - default=2, validation_alias="PROVIDER_RATE_WINDOW" + default=2, gt=0, validation_alias="PROVIDER_RATE_WINDOW" ) provider_max_concurrency: int = Field( - default=2, validation_alias="PROVIDER_MAX_CONCURRENCY" + default=2, gt=0, validation_alias="PROVIDER_MAX_CONCURRENCY" ) provider_progress_timeout: float = Field( default=600.0, diff --git a/src/free_claude_code/providers/admission.py b/src/free_claude_code/providers/admission.py index 3bd9490a7f..51f3373f58 100644 --- a/src/free_claude_code/providers/admission.py +++ b/src/free_claude_code/providers/admission.py @@ -5,6 +5,7 @@ import random import time import uuid +from collections import deque from collections.abc import Awaitable, Callable from dataclasses import dataclass, field from datetime import UTC, datetime @@ -14,8 +15,8 @@ from loguru import logger -from free_claude_code.core.rate_limit import StrictSlidingWindowLimiter from free_claude_code.core.trace import trace_event +from free_claude_code.providers.admission_policy import ProviderAdmissionLimits from free_claude_code.providers.failure_policy import ( ProviderFailureOverride, ProviderRecoveryExhausted, @@ -416,12 +417,7 @@ def __init__( max_delay: float = DEFAULT_UPSTREAM_MAX_DELAY, jitter: float = DEFAULT_UPSTREAM_JITTER, ) -> None: - if rate_limit <= 0: - raise ValueError("rate_limit must be > 0") - if rate_window <= 0: - raise ValueError("rate_window must be > 0") - if max_concurrency <= 0: - raise ValueError("max_concurrency must be > 0") + limits = ProviderAdmissionLimits(rate_limit, rate_window, max_concurrency) if max_attempts <= 0: raise ValueError("max_attempts must be > 0") if base_delay < 0: @@ -436,10 +432,11 @@ def __init__( self._base_delay = base_delay self._max_delay = max_delay self._jitter = jitter - self._proactive_limiter = StrictSlidingWindowLimiter( - rate_limit, float(rate_window) - ) - self._concurrency_sem = asyncio.Semaphore(max_concurrency) + self._limits = limits + self._admission_times: deque[float] = deque() + self._last_discarded_at: float | None = None + self._active_attempts = 0 + self._capacity_changed = asyncio.Event() self._condition = asyncio.Condition() self._episode: _RecoveryEpisode | None = None self._next_generation = 1 @@ -485,19 +482,11 @@ async def _open_attempt( slot_acquired = False claim: _AttemptClaim | None = None try: - admitted = await self._proactive_limiter.acquire_if( - lambda permit=permit: self._permit_is_current(execution, permit) - ) + admitted = await self._acquire_capacity(execution, permit) if not admitted: await self._abandon_probe_permit(execution, permit) continue - await self._concurrency_sem.acquire() slot_acquired = True - if not self._permit_is_current(execution, permit): - self._concurrency_sem.release() - slot_acquired = False - await self._abandon_probe_permit(execution, permit) - continue claim = execution._claim_attempt(operation_kind) attempt = ProviderAttempt(self, execution, permit, claim) trace_event( @@ -518,7 +507,7 @@ async def _open_attempt( if claim is not None: execution._close_attempt(claim) if slot_acquired: - self._concurrency_sem.release() + self._release_concurrency() await self._abandon_probe_permit(execution, permit) raise @@ -639,7 +628,7 @@ async def _attempt_accepted( if episode is None: return self._episode = None - self._condition.notify_all() + self._notify_recovery_waiters() trace_event( stage="provider", event="provider.recovery.closed", @@ -664,7 +653,7 @@ async def _attempt_corrected( return episode.probe_active = False episode.ready_at = time.monotonic() - self._condition.notify_all() + self._notify_recovery_waiters() async def _attempt_rejected( self, @@ -679,7 +668,7 @@ async def _attempt_rejected( if episode is None: return self._episode = None - self._condition.notify_all() + self._notify_recovery_waiters() trace_event( stage="provider", event="provider.recovery.closed", @@ -705,7 +694,7 @@ async def _attempt_abandoned( episode.leader = None episode.probe_active = False episode.ready_at = time.monotonic() - self._condition.notify_all() + self._notify_recovery_waiters() async def _attempt_failed( self, @@ -740,7 +729,7 @@ async def _attempt_failed( execution._fail_recovery(episode.last_error) else: episode.waiters.add(execution) - self._condition.notify_all() + self._notify_recovery_waiters() elif episode is None or matching_probe is not None: terminal_delay = self._retry_delay( error, @@ -761,7 +750,7 @@ async def _attempt_failed( waiter._fail_recovery(error) episode.waiters.clear() exhausted_episode = True - self._condition.notify_all() + self._notify_recovery_waiters() label = self._failure_label(status, error) if became_leader: @@ -871,7 +860,7 @@ async def _abandon_probe_permit( episode.leader = None episode.probe_active = False episode.ready_at = time.monotonic() - self._condition.notify_all() + self._notify_recovery_waiters() async def _abandon_waiting_leader( self, @@ -888,10 +877,64 @@ async def _abandon_waiting_leader( ): return episode.leader = None - self._condition.notify_all() + self._notify_recovery_waiters() + + def reconfigure(self, limits: ProviderAdmissionLimits) -> None: + """Update capacity without resetting occupancy, history, or recovery.""" + if limits != self._limits: + self._limits = limits + self._wake_capacity() + + def _wake_capacity(self) -> None: + changed = self._capacity_changed + self._capacity_changed = asyncio.Event() + changed.set() + + def _notify_recovery_waiters(self) -> None: + self._condition.notify_all() + self._wake_capacity() + + async def _acquire_capacity( + self, execution: ProviderExecution, permit: _GatePermit + ) -> bool: + while self._permit_is_current(execution, permit): + changed = self._capacity_changed + now = time.monotonic() + limits = self._limits + cutoff = now - limits.rate_window + while self._admission_times and self._admission_times[0] <= cutoff: + self._last_discarded_at = self._admission_times.popleft() + + ready_at = now + if self._last_discarded_at is not None: + ready_at = max(ready_at, self._last_discarded_at + limits.rate_window) + if len(self._admission_times) >= limits.rate_limit: + ready_at = max( + ready_at, + self._admission_times[-limits.rate_limit] + limits.rate_window, + ) + if self._active_attempts < limits.max_concurrency and ready_at <= now: + self._admission_times.append(now) + self._active_attempts += 1 + return True + + # Occupancy changes explicitly wake us. Avoid an expired rate timer + # spinning while all concurrency slots remain occupied. + delay = ( + max(0.0, ready_at - now) + if self._active_attempts < limits.max_concurrency + else None + ) + try: + async with asyncio.timeout(delay): + await changed.wait() + except TimeoutError: + pass + return False def _release_concurrency(self) -> None: - self._concurrency_sem.release() + self._active_attempts -= 1 + self._wake_capacity() def _retry_delay(self, error: Exception, attempt: int) -> float: exponent = max(0, attempt - 1) diff --git a/src/free_claude_code/providers/admission_policy.py b/src/free_claude_code/providers/admission_policy.py new file mode 100644 index 0000000000..3b423d8f97 --- /dev/null +++ b/src/free_claude_code/providers/admission_policy.py @@ -0,0 +1,31 @@ +"""Validated capacity policy, independent of clients and controller lifetime.""" + +import math +from dataclasses import dataclass +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from free_claude_code.config.settings import Settings + + +@dataclass(frozen=True, slots=True) +class ProviderAdmissionLimits: + rate_limit: int + rate_window: float + max_concurrency: int + + def __post_init__(self) -> None: + for name in ("rate_limit", "max_concurrency"): + value = getattr(self, name) + if not isinstance(value, int) or isinstance(value, bool) or value <= 0: + raise ValueError(f"{name} must be a positive integer") + if not math.isfinite(self.rate_window) or self.rate_window <= 0: + raise ValueError("rate_window must be finite and > 0") + + @classmethod + def from_settings(cls, settings: Settings) -> ProviderAdmissionLimits: + return cls( + settings.provider_rate_limit, + settings.provider_rate_window, + settings.provider_max_concurrency, + ) diff --git a/src/free_claude_code/providers/admission_registry.py b/src/free_claude_code/providers/admission_registry.py new file mode 100644 index 0000000000..772c11bf1d --- /dev/null +++ b/src/free_claude_code/providers/admission_registry.py @@ -0,0 +1,50 @@ +"""Loop-owned admission policy and controller lifetimes, independent of clients.""" + +from collections.abc import Iterable +from typing import TYPE_CHECKING + +from free_claude_code.providers.admission_policy import ProviderAdmissionLimits + +if TYPE_CHECKING: + from free_claude_code.providers.admission import ProviderAdmissionController + + +class ProviderAdmissionRegistry: + """One protection budget per configured provider for one runtime manager.""" + + def __init__(self, limits: ProviderAdmissionLimits) -> None: + self._limits = limits + self._controllers: dict[str, ProviderAdmissionController] = {} + self._closed = False + + def get(self, provider_id: str) -> ProviderAdmissionController: + if self._closed: + raise RuntimeError("Provider admission registry is closed") + if provider_id not in self._controllers: + # Provider preparation loads SDKs in a worker before reaching here. + from free_claude_code.providers.admission import ProviderAdmissionController + + self._controllers[provider_id] = ProviderAdmissionController( + provider_name=provider_id, + rate_limit=self._limits.rate_limit, + rate_window=self._limits.rate_window, + max_concurrency=self._limits.max_concurrency, + ) + return self._controllers[provider_id] + + def reconfigure(self, limits: ProviderAdmissionLimits) -> None: + """Publish validated limits without yielding to waiting requests.""" + self._limits = limits + for controller in self._controllers.values(): + controller.reconfigure(limits) + + def retain_custom(self, provider_ids: Iterable[str]) -> None: + retained = set(provider_ids) + for provider_id in tuple(self._controllers): + if provider_id.startswith("custom_") and provider_id not in retained: + del self._controllers[provider_id] + + def close(self) -> None: + """Forget state only after client generations have drained and closed.""" + self._closed = True + self._controllers.clear() diff --git a/src/free_claude_code/providers/base.py b/src/free_claude_code/providers/base.py index ecf11ab817..a8ab440528 100644 --- a/src/free_claude_code/providers/base.py +++ b/src/free_claude_code/providers/base.py @@ -20,9 +20,6 @@ class ProviderConfig: api_key: str | None base_url: str - rate_limit: int - rate_window: int - max_concurrency: int http_read_timeout: float http_write_timeout: float http_connect_timeout: float diff --git a/src/free_claude_code/providers/runtime/config.py b/src/free_claude_code/providers/runtime/config.py index a36e9b921c..0c6a3dd45c 100644 --- a/src/free_claude_code/providers/runtime/config.py +++ b/src/free_claude_code/providers/runtime/config.py @@ -80,9 +80,6 @@ def build_provider_config( return ProviderConfig( api_key=credential, base_url=resolved_base_url, - rate_limit=settings.provider_rate_limit, - rate_window=settings.provider_rate_window, - max_concurrency=settings.provider_max_concurrency, http_read_timeout=settings.http_read_timeout, http_write_timeout=settings.http_write_timeout, http_connect_timeout=settings.http_connect_timeout, @@ -99,9 +96,6 @@ def build_custom_provider_config( api_key=definition.api_key.get_secret_value() if definition.api_key else None, base_url=definition.base_url, proxy=None, - rate_limit=settings.provider_rate_limit, - rate_window=settings.provider_rate_window, - max_concurrency=settings.provider_max_concurrency, http_read_timeout=settings.http_read_timeout, http_write_timeout=settings.http_write_timeout, http_connect_timeout=settings.http_connect_timeout, diff --git a/src/free_claude_code/providers/runtime/factory.py b/src/free_claude_code/providers/runtime/factory.py index cd08d7b9f5..9533451201 100644 --- a/src/free_claude_code/providers/runtime/factory.py +++ b/src/free_claude_code/providers/runtime/factory.py @@ -11,6 +11,7 @@ from free_claude_code.config.provider_catalog import PROVIDER_CATALOG from free_claude_code.config.settings import Settings from free_claude_code.providers.admission import ProviderAdmissionController +from free_claude_code.providers.admission_registry import ProviderAdmissionRegistry from free_claude_code.providers.base import BaseProvider, ProviderConfig from free_claude_code.providers.openai_chat import ( OPENAI_CHAT_PROFILES, @@ -251,7 +252,7 @@ def prepare_provider( provider_id: str, provider_loaders: Mapping[str, Callable[[], ProviderFactory]], custom_definition: CustomProviderDefinition | None = None, -) -> Callable[[Settings], BaseProvider]: +) -> Callable[[Settings, ProviderAdmissionRegistry], BaseProvider]: """Load implementation modules in a worker; return a loop-owned constructor.""" # The SDK lazily imports these on first client resource access. Keep that @@ -262,14 +263,11 @@ def prepare_provider( from .config import build_custom_provider_config - def construct_custom(settings: Settings) -> BaseProvider: + def construct_custom( + settings: Settings, admission_registry: ProviderAdmissionRegistry + ) -> BaseProvider: config = build_custom_provider_config(custom_definition, settings) - admission = ProviderAdmissionController( - provider_name=provider_id, - rate_limit=config.rate_limit, - rate_window=config.rate_window, - max_concurrency=config.max_concurrency, - ) + admission = admission_registry.get(provider_id) return CustomProvider( config, definition=custom_definition, admission=admission ) @@ -287,14 +285,11 @@ def construct_custom(settings: Settings) -> BaseProvider: ) factory = loader() if loader is not None else None - def construct(settings: Settings) -> BaseProvider: + def construct( + settings: Settings, admission_registry: ProviderAdmissionRegistry + ) -> BaseProvider: config = build_provider_config(descriptor, settings) - admission = ProviderAdmissionController( - provider_name=provider_id, - rate_limit=config.rate_limit, - rate_window=config.rate_window, - max_concurrency=config.max_concurrency, - ) + admission = admission_registry.get(provider_id) if factory is not None: return factory(config, settings, admission) return create_openai_chat_provider(provider_id, config, admission) diff --git a/src/free_claude_code/providers/runtime/runtime.py b/src/free_claude_code/providers/runtime/runtime.py index 0a2ebd09e7..a615af40df 100644 --- a/src/free_claude_code/providers/runtime/runtime.py +++ b/src/free_claude_code/providers/runtime/runtime.py @@ -8,18 +8,21 @@ from free_claude_code.application.errors import ApplicationUnavailableError from free_claude_code.config.settings import Settings from free_claude_code.core.async_tasks import run_sync_owned +from free_claude_code.providers.admission_registry import ProviderAdmissionRegistry from free_claude_code.providers.base import BaseProvider if TYPE_CHECKING: from .factory import ProviderFactory -type ProviderConstructor = Callable[[str, Settings], Awaitable[BaseProvider]] +type ProviderConstructor = Callable[ + [str, Settings, ProviderAdmissionRegistry], Awaitable[BaseProvider] +] type ProviderLoader = Callable[[], ProviderFactory] def _load_constructor( provider_id: str, provider_loaders: Mapping[str, ProviderLoader], settings: Settings -) -> Callable[[Settings], BaseProvider]: +) -> Callable[[Settings, ProviderAdmissionRegistry], BaseProvider]: from .factory import prepare_provider return prepare_provider( @@ -34,13 +37,14 @@ def _load_constructor( async def create_provider( provider_id: str, settings: Settings, + admission_registry: ProviderAdmissionRegistry, *, provider_loaders: Mapping[str, ProviderLoader] | None = None, ) -> BaseProvider: constructor = await run_sync_owned( partial(_load_constructor, provider_id, provider_loaders or {}, settings) ) - return constructor(settings) + return constructor(settings, admission_registry) class ProviderRuntime: @@ -49,11 +53,13 @@ class ProviderRuntime: def __init__( self, settings: Settings, + admission_registry: ProviderAdmissionRegistry, providers: MutableMapping[str, BaseProvider] | None = None, *, provider_constructor: ProviderConstructor = create_provider, ) -> None: self.settings = settings + self._admission_registry = admission_registry self._providers = providers if providers is not None else {} self._provider_constructor = provider_constructor self._creations: dict[str, asyncio.Task[BaseProvider]] = {} @@ -85,7 +91,9 @@ def _creation_done( task.exception() # Observe failures even when every caller stopped waiting. async def _construct(self, provider_id: str) -> BaseProvider: - provider = await self._provider_constructor(provider_id, self.settings) + provider = await self._provider_constructor( + provider_id, self.settings, self._admission_registry + ) self._providers[provider_id] = provider return provider diff --git a/src/free_claude_code/runtime/provider_manager.py b/src/free_claude_code/runtime/provider_manager.py index 7b124d6eda..e04bc9e8cf 100644 --- a/src/free_claude_code/runtime/provider_manager.py +++ b/src/free_claude_code/runtime/provider_manager.py @@ -20,6 +20,8 @@ from free_claude_code.core.json_types import JsonObject from free_claude_code.core.token_estimation import initialize_token_estimation from free_claude_code.core.trace import trace_event +from free_claude_code.providers.admission_policy import ProviderAdmissionLimits +from free_claude_code.providers.admission_registry import ProviderAdmissionRegistry from free_claude_code.providers.base import BaseProvider from free_claude_code.providers.model_listing import model_infos_from_ids from free_claude_code.providers.runtime.discovery import ( @@ -31,7 +33,9 @@ from free_claude_code.providers.runtime.model_cache import ProviderModelCache from free_claude_code.providers.runtime.runtime import ProviderRuntime -ProviderRuntimeFactory = Callable[[Settings], ProviderRuntime] +ProviderRuntimeFactory = Callable[ + [Settings, ProviderAdmissionRegistry], ProviderRuntime +] ConnectedProviderIds = Callable[[], tuple[str, ...]] CommitConfig = Callable[[], Awaitable[None]] @@ -157,6 +161,9 @@ def __init__( model_catalog_publisher: ModelCatalogPublisher | None = None, ) -> None: self._runtime_factory = runtime_factory + self._admission_registry = ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(settings) + ) self._connected_provider_ids = connected_provider_ids self._model_catalog_publisher = model_catalog_publisher self._catalog_changed: Callable[[], None] | None = None @@ -174,7 +181,7 @@ def __init__( self._current = _ProviderGeneration( generation_id=1, settings=settings, - runtime=runtime_factory(settings), + runtime=runtime_factory(settings, self._admission_registry), cache=ProviderModelCache( model_cache_provider_ids_for_settings( settings, connected_provider_ids() @@ -527,10 +534,13 @@ async def replace( async with self._replace_lock: self._ensure_open() await self._retry_unpublished_cleanup() + limits = ProviderAdmissionLimits.from_settings(settings) candidate_id = self._next_generation_id candidate_runtime: ProviderRuntime | None = None try: - candidate_runtime = self._runtime_factory(settings) + candidate_runtime = self._runtime_factory( + settings, self._admission_registry + ) await commit() except BaseException as exc: trace_event( @@ -565,6 +575,7 @@ async def replace( runtime=candidate_runtime, cache=cache, ) + self._admission_registry.reconfigure(limits) self._current = candidate self._catalog_revision += 1 previous.retired = True @@ -621,8 +632,21 @@ async def close(self) -> None: if not all(generation_results) or not unpublished_closed: raise RuntimeError("One or more provider runtimes failed to close.") self._current.cache.clear() + self._admission_registry.close() self._closed = True + def _prune_admission(self) -> None: + runtimes = ( + self._current.runtime, + *(generation.runtime for generation in self._retired.values()), + *self._unpublished, + ) + self._admission_registry.retain_custom( + definition.provider_id + for runtime in runtimes + for definition in runtime.settings.custom_providers + ) + async def _release(self, generation: _ProviderGeneration) -> None: if generation.active_leases <= 0: return @@ -645,6 +669,7 @@ async def _cleanup_unpublished(self, runtime: ProviderRuntime) -> bool: ) return False self._unpublished.discard(runtime) + self._prune_admission() return True async def _retry_unpublished_cleanup(self) -> bool: @@ -696,6 +721,7 @@ async def _run_generation_cleanup( generation.closed = True self._retired.pop(generation.generation_id, None) + self._prune_admission() trace_event( stage="runtime", event="provider_generation.closed", diff --git a/tests/api/support.py b/tests/api/support.py index ff7494daff..c5543727b7 100644 --- a/tests/api/support.py +++ b/tests/api/support.py @@ -67,8 +67,9 @@ def connected_provider_ids() -> tuple[str, ...]: else: manager = ApiTestRuntime( settings, - runtime_factory=lambda snapshot: ProviderRuntime( + runtime_factory=lambda snapshot, admission_registry: ProviderRuntime( snapshot, + admission_registry, dict(providers), ), connected_provider_ids=connected_provider_ids, diff --git a/tests/api/test_openai_codex_compatibility.py b/tests/api/test_openai_codex_compatibility.py index d1e152409c..2d48f0ad1c 100644 --- a/tests/api/test_openai_codex_compatibility.py +++ b/tests/api/test_openai_codex_compatibility.py @@ -71,9 +71,6 @@ def handler(request: httpx2.Request) -> httpx2.Response: make_provider_config( api_key="", base_url="https://chatgpt.com/backend-api/codex", - rate_limit=100, - rate_window=1, - max_concurrency=2, ), auth=_FakeOpenAIAuth(), admission=ProviderAdmissionController( diff --git a/tests/cli/test_config_restart_lifecycle.py b/tests/cli/test_config_restart_lifecycle.py index 404a58a391..d01ed8ab76 100644 --- a/tests/cli/test_config_restart_lifecycle.py +++ b/tests/cli/test_config_restart_lifecycle.py @@ -35,7 +35,10 @@ def test_supervised_http_apply_finishes_and_reconnects(monkeypatch, stop_during_ def build(settings, restart_callback): manager = ProviderRuntimeManager( - settings, runtime_factory=lambda snapshot: ProviderRuntime(snapshot, {}) + settings, + runtime_factory=lambda snapshot, admission_registry: ProviderRuntime( + snapshot, admission_registry, {} + ), ) monkeypatch.setattr(manager, "start_model_list_refresh", lambda: None) monkeypatch.setattr(manager, "_start_pass", lambda *args, **kwargs: None) diff --git a/tests/conftest.py b/tests/conftest.py index 3da3ad974b..e2799bf4c0 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -85,9 +85,6 @@ def provider_config(): return make_provider_config( api_key="test_key", base_url="https://test.api.nvidia.com/v1", - rate_limit=10, - rate_window=60, - max_concurrency=5, http_read_timeout=300.0, http_write_timeout=10.0, http_connect_timeout=10.0, @@ -123,9 +120,6 @@ def lmstudio_provider(provider_config): lmstudio_config = make_provider_config( api_key="lm-studio", base_url="http://localhost:1234/v1", - rate_limit=provider_config.rate_limit, - rate_window=provider_config.rate_window, - max_concurrency=provider_config.max_concurrency, http_read_timeout=provider_config.http_read_timeout, http_write_timeout=provider_config.http_write_timeout, http_connect_timeout=provider_config.http_connect_timeout, @@ -143,9 +137,6 @@ def llamacpp_provider(provider_config): llamacpp_config = make_provider_config( api_key="llamacpp", base_url="http://localhost:8080/v1", - rate_limit=10, - rate_window=60, - max_concurrency=5, http_read_timeout=300.0, http_write_timeout=10.0, http_connect_timeout=10.0, diff --git a/tests/contracts/test_startup_import_boundaries.py b/tests/contracts/test_startup_import_boundaries.py index 0bb62026d2..f76fd192f7 100644 --- a/tests/contracts/test_startup_import_boundaries.py +++ b/tests/contracts/test_startup_import_boundaries.py @@ -41,6 +41,8 @@ def test_sdk_preparation_runs_in_worker_and_provider_construction_runs_on_owner_ from free_claude_code.config.settings import Settings from free_claude_code.providers.base import BaseProvider from free_claude_code.providers.runtime.runtime import create_provider +from free_claude_code.providers.admission_policy import ProviderAdmissionLimits +from free_claude_code.providers.admission_registry import ProviderAdmissionRegistry async def main(): loop = asyncio.get_running_loop() @@ -54,7 +56,7 @@ def construct(*args): assert threading.get_ident() == thread return provider return construct - assert await create_provider("groq", Settings(groq_api_key="test"), provider_loaders={"groq": load}) is provider + assert await create_provider("groq", Settings(groq_api_key="test"), ProviderAdmissionRegistry(ProviderAdmissionLimits(1, 2, 2)), provider_loaders={"groq": load}) is provider await provider.cleanup() asyncio.run(main()) diff --git a/tests/providers/support.py b/tests/providers/support.py index 63cf6fde7d..4da6541453 100644 --- a/tests/providers/support.py +++ b/tests/providers/support.py @@ -52,9 +52,6 @@ async def close(self) -> None: def make_provider_config( api_key: str | None, base_url: str, - rate_limit: int = 1_000_000, - rate_window: int = 1, - max_concurrency: int = 1_000, http_read_timeout: float = 120.0, http_write_timeout: float = 10.0, http_connect_timeout: float = 10.0, @@ -67,9 +64,6 @@ def make_provider_config( return ProviderConfig( api_key=api_key, base_url=base_url, - rate_limit=rate_limit, - rate_window=rate_window, - max_concurrency=max_concurrency, http_read_timeout=http_read_timeout, http_write_timeout=http_write_timeout, http_connect_timeout=http_connect_timeout, diff --git a/tests/providers/test_admission_reconfiguration.py b/tests/providers/test_admission_reconfiguration.py new file mode 100644 index 0000000000..1d0b025eda --- /dev/null +++ b/tests/providers/test_admission_reconfiguration.py @@ -0,0 +1,146 @@ +"""Live policy changes preserve capacity and wake the correct waiting calls.""" + +import asyncio +from types import SimpleNamespace + +import pytest +from pydantic import ValidationError + +from free_claude_code.config.settings import Settings +from free_claude_code.providers import admission as admission_module +from free_claude_code.providers.admission import ( + ProviderAdmissionController, + ProviderOperationKind, +) +from free_claude_code.providers.admission_policy import ProviderAdmissionLimits + + +async def admit(controller): + return await controller.start_execution().open_attempt( + ProviderOperationKind.GENERATION + ) + + +async def completed(controller): + attempt = await admit(controller) + await attempt.accept() + await attempt.aclose() + + +@pytest.mark.asyncio +async def test_limit_changes_preserve_occupied_slots_and_wake_waiters(): + controller = ProviderAdmissionController( + provider_name="test", rate_limit=1000, max_concurrency=1 + ) + first = await admit(controller) + pending = asyncio.create_task(admit(controller)) + await asyncio.sleep(0) + assert not pending.done() + controller.reconfigure(ProviderAdmissionLimits(1000, 60, 2)) + second = await asyncio.wait_for(pending, 1) + controller.reconfigure(ProviderAdmissionLimits(1000, 60, 1)) + third = asyncio.create_task(admit(controller)) + await first.aclose() + await asyncio.sleep(0) + assert not third.done() + await second.aclose() + await (await asyncio.wait_for(third, 1)).aclose() + + +@pytest.mark.asyncio +async def test_waiting_for_concurrency_does_not_spend_rate_capacity(monkeypatch): + clock = SimpleNamespace(now=0.0) + monkeypatch.setattr( + admission_module, "time", SimpleNamespace(monotonic=lambda: clock.now) + ) + controller = ProviderAdmissionController( + provider_name="test", rate_limit=1, rate_window=10, max_concurrency=1 + ) + first = await admit(controller) + clock.now = 11 + pending = asyncio.create_task(admit(controller)) + await asyncio.sleep(0) + pending.cancel() + await asyncio.gather(pending, return_exceptions=True) + await first.aclose() + await asyncio.wait_for(completed(controller), 1) + + +@pytest.mark.asyncio +async def test_rate_count_increase_wakes_waiter_without_clearing_history(): + controller = ProviderAdmissionController( + provider_name="test", rate_limit=1, rate_window=60 + ) + await completed(controller) + pending = asyncio.create_task(admit(controller)) + await asyncio.sleep(0) + assert not pending.done() + controller.reconfigure(ProviderAdmissionLimits(2, 60, 5)) + await (await asyncio.wait_for(pending, 1)).aclose() + blocked = asyncio.create_task(admit(controller)) + await asyncio.sleep(0) + assert not blocked.done() + blocked.cancel() + await asyncio.gather(blocked, return_exceptions=True) + + +@pytest.mark.asyncio +async def test_longer_window_waits_only_until_discarded_history_expires(monkeypatch): + clock = SimpleNamespace(now=0.0) + monkeypatch.setattr( + admission_module, "time", SimpleNamespace(monotonic=lambda: clock.now) + ) + controller = ProviderAdmissionController( + provider_name="test", rate_limit=1, rate_window=2 + ) + await completed(controller) + clock.now = 3 + await completed(controller) # Drops the timestamp at zero. + controller.reconfigure(ProviderAdmissionLimits(100, 10, 5)) + pending = asyncio.create_task(admit(controller)) + await asyncio.sleep(0) + assert not pending.done() # Capacity is ample, but older history is incomplete. + clock.now = 9 + controller.reconfigure(ProviderAdmissionLimits(101, 10, 5)) + await asyncio.sleep(0) + assert not pending.done() + clock.now = 10 + controller.reconfigure(ProviderAdmissionLimits(102, 10, 5)) + await (await asyncio.wait_for(pending, 1)).aclose() + + +@pytest.mark.asyncio +async def test_shorter_window_wakes_waiter_and_longer_window_uses_retained_history( + monkeypatch, +): + clock = SimpleNamespace(now=0.0) + monkeypatch.setattr( + admission_module, "time", SimpleNamespace(monotonic=lambda: clock.now) + ) + controller = ProviderAdmissionController( + provider_name="test", rate_limit=1, rate_window=60 + ) + await completed(controller) + controller.reconfigure(ProviderAdmissionLimits(2, 120, 5)) + await asyncio.wait_for(completed(controller), 1) # No artificial expansion pause. + pending = asyncio.create_task(admit(controller)) + await asyncio.sleep(0) + assert not pending.done() + clock.now = 2 + controller.reconfigure(ProviderAdmissionLimits(1, 2, 5)) + await (await asyncio.wait_for(pending, 1)).aclose() + + +@pytest.mark.parametrize( + "field", ["provider_rate_limit", "provider_rate_window", "provider_max_concurrency"] +) +@pytest.mark.parametrize("value", [0, -1]) +def test_settings_reject_invalid_admission_limits(field, value): + with pytest.raises(ValidationError): + Settings.model_validate({field: value}) + + +@pytest.mark.parametrize("window", [float("inf"), float("nan"), 0, -1]) +def test_internal_policy_rejects_invalid_window(window): + with pytest.raises(ValueError, match="rate_window"): + ProviderAdmissionLimits(1, window, 1) diff --git a/tests/providers/test_agnes.py b/tests/providers/test_agnes.py index 77dc384b57..f4640bd408 100644 --- a/tests/providers/test_agnes.py +++ b/tests/providers/test_agnes.py @@ -31,8 +31,6 @@ def agnes_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-agnes-key", base_url=AGNES_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="agnes"), ) @@ -211,8 +209,6 @@ def build_client(*args: Any, **kwargs: Any) -> AsyncOpenAI: make_provider_config( api_key="wire-agnes-key", base_url=AGNES_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="agnes"), ) diff --git a/tests/providers/test_cerebras.py b/tests/providers/test_cerebras.py index 501633e289..e8f3231196 100644 --- a/tests/providers/test_cerebras.py +++ b/tests/providers/test_cerebras.py @@ -57,8 +57,6 @@ def cerebras_config(): return make_provider_config( api_key="test_cerebras_key", base_url=CEREBRAS_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ) @@ -127,8 +125,6 @@ def test_replay_is_independent_of_current_turn_reasoning_control(): make_provider_config( api_key="test_cerebras_key", base_url=CEREBRAS_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_cheaperinference.py b/tests/providers/test_cheaperinference.py index bf71d9d47e..f2ecc4fa83 100644 --- a/tests/providers/test_cheaperinference.py +++ b/tests/providers/test_cheaperinference.py @@ -32,8 +32,6 @@ def cheaperinference_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-cheaperinference-key", base_url=CHEAPERINFERENCE_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="cheaperinference", max_attempts=1), ) diff --git a/tests/providers/test_chutes.py b/tests/providers/test_chutes.py index eb8692e361..2c5d5f48bb 100644 --- a/tests/providers/test_chutes.py +++ b/tests/providers/test_chutes.py @@ -32,8 +32,6 @@ def chutes_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-chutes-key", base_url=CHUTES_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="chutes"), ) diff --git a/tests/providers/test_cline_pass.py b/tests/providers/test_cline_pass.py index 2c98230748..4cad67e1d5 100644 --- a/tests/providers/test_cline_pass.py +++ b/tests/providers/test_cline_pass.py @@ -40,8 +40,6 @@ def cline_pass_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-cline-key", base_url=CLINE_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="cline_pass"), ) @@ -67,8 +65,6 @@ def _provider_with_transport( make_provider_config( api_key=api_key, base_url=CLINE_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="cline_pass"), ) diff --git a/tests/providers/test_cloudflare.py b/tests/providers/test_cloudflare.py index 7bcff676d3..4e7c90866a 100644 --- a/tests/providers/test_cloudflare.py +++ b/tests/providers/test_cloudflare.py @@ -34,8 +34,6 @@ def cloudflare_config() -> ProviderConfig: return make_provider_config( api_key="test-cloudflare-token", base_url=CLOUDFLARE_AI_REST_ROOT, - rate_limit=10, - rate_window=60, ) diff --git a/tests/providers/test_codestral.py b/tests/providers/test_codestral.py index 18bd740bb8..bb1085b9f1 100644 --- a/tests/providers/test_codestral.py +++ b/tests/providers/test_codestral.py @@ -25,8 +25,6 @@ def codestral_config(): return make_provider_config( api_key="test_codestral_key", base_url=CODESTRAL_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ) @@ -72,8 +70,6 @@ def test_build_request_body_global_disable_blocks_reasoning_mapping(): make_provider_config( api_key="test_codestral_key", base_url=CODESTRAL_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_cohere.py b/tests/providers/test_cohere.py index 2e5acb906d..27b6a56ff1 100644 --- a/tests/providers/test_cohere.py +++ b/tests/providers/test_cohere.py @@ -26,8 +26,6 @@ def cohere_config(): return make_provider_config( api_key="test_cohere_key", base_url=COHERE_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ) @@ -147,8 +145,6 @@ def test_build_request_body_maps_reasoning_off_to_none(): make_provider_config( api_key="test_cohere_key", base_url=COHERE_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_deepinfra.py b/tests/providers/test_deepinfra.py index b2734849a9..c8da28532f 100644 --- a/tests/providers/test_deepinfra.py +++ b/tests/providers/test_deepinfra.py @@ -71,8 +71,6 @@ def deepinfra_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-deepinfra-key", base_url=DEEPINFRA_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="deepinfra"), ) @@ -290,8 +288,6 @@ def build_client(*args: Any, **kwargs: Any) -> AsyncOpenAI: make_provider_config( api_key="wire-deepinfra-key", base_url=DEEPINFRA_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="deepinfra"), ) diff --git a/tests/providers/test_deepseek.py b/tests/providers/test_deepseek.py index ae03d119ac..78730315d0 100644 --- a/tests/providers/test_deepseek.py +++ b/tests/providers/test_deepseek.py @@ -40,8 +40,6 @@ def deepseek_config(): return make_provider_config( api_key="test_deepseek_key", base_url=DEEPSEEK_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ) @@ -368,8 +366,6 @@ def test_build_request_body_encodes_reasoning_off(): make_provider_config( api_key="k", base_url=DEEPSEEK_DEFAULT_BASE, - rate_limit=1, - rate_window=1, ), admission=immediate_admission(), ) @@ -689,8 +685,6 @@ def test_thinking_off_preserves_historical_reasoning(): make_provider_config( api_key="k", base_url=DEEPSEEK_DEFAULT_BASE, - rate_limit=1, - rate_window=1, ), admission=immediate_admission(), ) @@ -718,8 +712,6 @@ def test_thinking_off_still_replays_required_tool_reasoning(): make_provider_config( api_key="k", base_url=DEEPSEEK_DEFAULT_BASE, - rate_limit=1, - rate_window=1, ), admission=immediate_admission(), ) @@ -866,8 +858,6 @@ def test_vision_model_strips_user_document(): make_provider_config( api_key="k", base_url=DEEPSEEK_DEFAULT_BASE, - rate_limit=1, - rate_window=1, ), admission=immediate_admission(), ) @@ -903,8 +893,6 @@ def test_startup_rejects_mcp_servers(): make_provider_config( api_key="k", base_url=DEEPSEEK_DEFAULT_BASE, - rate_limit=1, - rate_window=1, ), admission=immediate_admission(), ) @@ -922,8 +910,6 @@ def test_startup_rejects_listed_server_tools_in_tools_list(): make_provider_config( api_key="k", base_url=DEEPSEEK_DEFAULT_BASE, - rate_limit=1, - rate_window=1, ), admission=immediate_admission(), ) @@ -959,8 +945,6 @@ def test_startup_preserves_completed_server_tool_history(): make_provider_config( api_key="k", base_url=DEEPSEEK_DEFAULT_BASE, - rate_limit=1, - rate_window=1, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_execution_failure_boundary.py b/tests/providers/test_execution_failure_boundary.py index 20d27a79f6..bdbdee2646 100644 --- a/tests/providers/test_execution_failure_boundary.py +++ b/tests/providers/test_execution_failure_boundary.py @@ -62,8 +62,6 @@ def _provider() -> NvidiaNimProvider: make_provider_config( api_key="test_key", base_url="https://test.api.nvidia.com/v1", - rate_limit=10, - rate_window=60, ), nim_settings=NimSettings(), admission=immediate_admission(), diff --git a/tests/providers/test_experiential.py b/tests/providers/test_experiential.py index f192276ac9..36f042bc88 100644 --- a/tests/providers/test_experiential.py +++ b/tests/providers/test_experiential.py @@ -32,8 +32,6 @@ def experiential_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-experiential-key", base_url=EXPERIENTIAL_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="experiential", max_attempts=1), ) diff --git a/tests/providers/test_featherless.py b/tests/providers/test_featherless.py index 4dd33ea024..eaf6b36315 100644 --- a/tests/providers/test_featherless.py +++ b/tests/providers/test_featherless.py @@ -35,8 +35,6 @@ def featherless_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-featherless-key", base_url=FEATHERLESS_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission( provider_name="featherless", diff --git a/tests/providers/test_fireworks.py b/tests/providers/test_fireworks.py index 36f32ed9c8..7f33675c8c 100644 --- a/tests/providers/test_fireworks.py +++ b/tests/providers/test_fireworks.py @@ -25,8 +25,6 @@ def fireworks_provider(): make_provider_config( api_key="test_fireworks_key", base_url=FIREWORKS_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) @@ -81,8 +79,6 @@ def test_replay_is_independent_of_current_turn_reasoning_control(): make_provider_config( api_key="k", base_url=FIREWORKS_DEFAULT_BASE, - rate_limit=1, - rate_window=1, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_gemini.py b/tests/providers/test_gemini.py index cfa342d93a..60b2c2d238 100644 --- a/tests/providers/test_gemini.py +++ b/tests/providers/test_gemini.py @@ -49,8 +49,6 @@ def gemini_config(): return make_provider_config( api_key="test_gemini_key", base_url=GEMINI_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ) @@ -123,8 +121,6 @@ def test_build_request_body_reasoning_off_sets_reasoning_none(): make_provider_config( api_key="test_gemini_key", base_url=GEMINI_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_groq.py b/tests/providers/test_groq.py index d329f56e17..bc997e2d24 100644 --- a/tests/providers/test_groq.py +++ b/tests/providers/test_groq.py @@ -24,8 +24,6 @@ def groq_config(): return make_provider_config( api_key="test_groq_key", base_url=GROQ_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ) @@ -144,8 +142,6 @@ def test_build_request_body_global_disable_blocks_reasoning_mapping(): make_provider_config( api_key="test_groq_key", base_url=GROQ_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_groq_reasoning.py b/tests/providers/test_groq_reasoning.py index 88909a6c02..aaa22785c7 100644 --- a/tests/providers/test_groq_reasoning.py +++ b/tests/providers/test_groq_reasoning.py @@ -64,8 +64,6 @@ def _provider(*, max_attempts: int = 5) -> GroqProvider: make_provider_config( api_key="test_groq_key", base_url=GROQ_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission( provider_name="GROQ", diff --git a/tests/providers/test_huggingface.py b/tests/providers/test_huggingface.py index f460d97ad9..e84f482275 100644 --- a/tests/providers/test_huggingface.py +++ b/tests/providers/test_huggingface.py @@ -29,8 +29,6 @@ def huggingface_config(): return make_provider_config( api_key="test_hf_key", base_url=HUGGINGFACE_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ) diff --git a/tests/providers/test_kilo.py b/tests/providers/test_kilo.py index 37fd73b070..b023ac4826 100644 --- a/tests/providers/test_kilo.py +++ b/tests/providers/test_kilo.py @@ -34,8 +34,6 @@ def kilo_config(): return make_provider_config( api_key="test_kilo_key", base_url=KILO_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ) diff --git a/tests/providers/test_kimi.py b/tests/providers/test_kimi.py index 66cb8dde50..aa88a43acc 100644 --- a/tests/providers/test_kimi.py +++ b/tests/providers/test_kimi.py @@ -26,8 +26,6 @@ def kimi_provider(): make_provider_config( api_key="test_kimi_key", base_url=KIMI_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_kimi_code.py b/tests/providers/test_kimi_code.py index 80e7aa2286..445b0d57d8 100644 --- a/tests/providers/test_kimi_code.py +++ b/tests/providers/test_kimi_code.py @@ -34,8 +34,6 @@ def kimi_code_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-subscription-key", base_url=KIMI_CODE_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_lightning.py b/tests/providers/test_lightning.py index 19451cc9e1..3f65391a6f 100644 --- a/tests/providers/test_lightning.py +++ b/tests/providers/test_lightning.py @@ -35,8 +35,6 @@ def lightning_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-lightning-key", base_url=LIGHTNING_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="lightning", max_attempts=1), ) diff --git a/tests/providers/test_llm7.py b/tests/providers/test_llm7.py index ba871968fc..f7fbaed862 100644 --- a/tests/providers/test_llm7.py +++ b/tests/providers/test_llm7.py @@ -36,8 +36,6 @@ def llm7_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-llm7-key", base_url=LLM7_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="llm7", max_attempts=1), ) diff --git a/tests/providers/test_lmstudio.py b/tests/providers/test_lmstudio.py index 4b762b6ebd..6d9de49f1b 100644 --- a/tests/providers/test_lmstudio.py +++ b/tests/providers/test_lmstudio.py @@ -49,8 +49,6 @@ def lmstudio_config(): return make_provider_config( api_key="lm-studio", base_url=LMSTUDIO_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ) diff --git a/tests/providers/test_minimax.py b/tests/providers/test_minimax.py index 2157079562..a8a789ce23 100644 --- a/tests/providers/test_minimax.py +++ b/tests/providers/test_minimax.py @@ -46,8 +46,6 @@ def minimax_provider(): make_provider_config( api_key="test-minimax-key", base_url=MINIMAX_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_mistral.py b/tests/providers/test_mistral.py index e3b054437b..1fce2eca2f 100644 --- a/tests/providers/test_mistral.py +++ b/tests/providers/test_mistral.py @@ -31,8 +31,6 @@ def mistral_config(): return make_provider_config( api_key="test_mistral_key", base_url=MISTRAL_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ) @@ -209,8 +207,6 @@ def test_build_request_body_reasoning_off_uses_native_none(): make_provider_config( api_key="test_mistral_key", base_url=MISTRAL_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) @@ -226,8 +222,6 @@ def test_reasoning_off_keeps_replay_separate_from_new_turn_compute(): make_provider_config( api_key="test_mistral_key", base_url=MISTRAL_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_model_discovery.py b/tests/providers/test_model_discovery.py index 3a6565386c..087d47c095 100644 --- a/tests/providers/test_model_discovery.py +++ b/tests/providers/test_model_discovery.py @@ -73,7 +73,9 @@ def _manager( providers = providers or {} return ProviderRuntimeManager( settings, - runtime_factory=lambda snapshot: ProviderRuntime(snapshot, dict(providers)), + runtime_factory=lambda snapshot, admission_registry: ProviderRuntime( + snapshot, admission_registry, dict(providers) + ), ) diff --git a/tests/providers/test_nararoute.py b/tests/providers/test_nararoute.py index b7d1b5b99a..783c5a5582 100644 --- a/tests/providers/test_nararoute.py +++ b/tests/providers/test_nararoute.py @@ -22,8 +22,6 @@ async def test_model_catalog_extracts_strict_optional_reasoning_boolean() -> Non make_provider_config( api_key="test-nararoute-key", base_url=NARAROUTE_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="nararoute"), ) diff --git a/tests/providers/test_nebius.py b/tests/providers/test_nebius.py index d50e01961d..40f1d0d546 100644 --- a/tests/providers/test_nebius.py +++ b/tests/providers/test_nebius.py @@ -33,8 +33,6 @@ def nebius_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-nebius-key", base_url=NEBIUS_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="nebius"), ) diff --git a/tests/providers/test_novita.py b/tests/providers/test_novita.py index 8b20436669..8751d1446d 100644 --- a/tests/providers/test_novita.py +++ b/tests/providers/test_novita.py @@ -27,8 +27,6 @@ def novita_provider(): make_provider_config( api_key="test_novita_key", base_url=NOVITA_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_nvidia_nim_degraded_retry.py b/tests/providers/test_nvidia_nim_degraded_retry.py index fef3449455..2d2a68fbf0 100644 --- a/tests/providers/test_nvidia_nim_degraded_retry.py +++ b/tests/providers/test_nvidia_nim_degraded_retry.py @@ -31,9 +31,6 @@ def _config(base_url: str) -> ProviderConfig: return make_provider_config( api_key="test_key", base_url=base_url, - rate_limit=1_000_000, - rate_window=1, - max_concurrency=1_000, http_read_timeout=30.0, http_write_timeout=15.0, http_connect_timeout=5.0, diff --git a/tests/providers/test_open_router.py b/tests/providers/test_open_router.py index e6ad157367..fcbfde7f13 100644 --- a/tests/providers/test_open_router.py +++ b/tests/providers/test_open_router.py @@ -156,8 +156,6 @@ def open_router_provider(): make_provider_config( api_key="test_openrouter_key", base_url="https://openrouter.ai/api/v1", - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_openai_chat_output_cap.py b/tests/providers/test_openai_chat_output_cap.py index c70e98bcf1..fbe28c15bb 100644 --- a/tests/providers/test_openai_chat_output_cap.py +++ b/tests/providers/test_openai_chat_output_cap.py @@ -233,8 +233,6 @@ def groq_provider(): make_provider_config( api_key="test_groq_key", base_url=GROQ_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_openai_chat_usage.py b/tests/providers/test_openai_chat_usage.py index e0660e4928..8ab19d2d6f 100644 --- a/tests/providers/test_openai_chat_usage.py +++ b/tests/providers/test_openai_chat_usage.py @@ -56,8 +56,6 @@ def __init__(self): make_provider_config( api_key="test_key", base_url="https://provider.example/v1", - rate_limit=100, - rate_window=60, ), behavior=_UsageTestBehavior( OpenAIChatProfile( diff --git a/tests/providers/test_openai_codex_provider.py b/tests/providers/test_openai_codex_provider.py index 87f1801e54..9dfdddb79d 100644 --- a/tests/providers/test_openai_codex_provider.py +++ b/tests/providers/test_openai_codex_provider.py @@ -92,9 +92,6 @@ def _config() -> ProviderConfig: return make_provider_config( api_key="", base_url="https://chatgpt.com/backend-api/codex", - rate_limit=100, - rate_window=1, - max_concurrency=2, ) diff --git a/tests/providers/test_openai_compat_5xx_retry.py b/tests/providers/test_openai_compat_5xx_retry.py index b848335bdb..bee395b42b 100644 --- a/tests/providers/test_openai_compat_5xx_retry.py +++ b/tests/providers/test_openai_compat_5xx_retry.py @@ -40,8 +40,6 @@ async def test_nim_stream_retries_on_openai_5xx_then_streams(status_code): config = make_provider_config( api_key="test_key", base_url="https://test.api.nvidia.com/v1", - rate_limit=100, - rate_window=60, http_read_timeout=600.0, http_write_timeout=15.0, http_connect_timeout=5.0, @@ -87,8 +85,6 @@ async def test_nim_stream_retries_on_pre_stream_connection_error_then_streams(): config = make_provider_config( api_key="test_key", base_url="https://test.api.nvidia.com/v1", - rate_limit=100, - rate_window=60, http_read_timeout=600.0, http_write_timeout=15.0, http_connect_timeout=5.0, @@ -131,8 +127,6 @@ async def test_nim_stream_connection_error_exhausted_emits_cause_chain(): config = make_provider_config( api_key="test_key", base_url="https://test.api.nvidia.com/v1", - rate_limit=100, - rate_window=60, http_read_timeout=600.0, http_write_timeout=15.0, http_connect_timeout=5.0, @@ -186,8 +180,6 @@ async def test_nim_stream_openai_5xx_exhausted_emits_user_message( config = make_provider_config( api_key="test_key", base_url="https://test.api.nvidia.com/v1", - rate_limit=100, - rate_window=60, http_read_timeout=600.0, http_write_timeout=15.0, http_connect_timeout=5.0, diff --git a/tests/providers/test_opencode.py b/tests/providers/test_opencode.py index 13040dcf99..aea6338c9e 100644 --- a/tests/providers/test_opencode.py +++ b/tests/providers/test_opencode.py @@ -58,8 +58,6 @@ def _config(): return make_provider_config( api_key="test_opencode_key", base_url="https://opencode.ai/zen/v1", - rate_limit=100, - rate_window=1, ) diff --git a/tests/providers/test_orcarouter.py b/tests/providers/test_orcarouter.py index 1ee38f95cb..b67539e034 100644 --- a/tests/providers/test_orcarouter.py +++ b/tests/providers/test_orcarouter.py @@ -35,8 +35,6 @@ def orcarouter_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-orcarouter-key", base_url=ORCAROUTER_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="orcarouter", max_attempts=1), ) diff --git a/tests/providers/test_poolside.py b/tests/providers/test_poolside.py index 49eadc81ad..efa7422a85 100644 --- a/tests/providers/test_poolside.py +++ b/tests/providers/test_poolside.py @@ -34,8 +34,6 @@ def poolside_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-poolside-key", base_url=POOLSIDE_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="poolside", max_attempts=1), ) diff --git a/tests/providers/test_provider_client_tool_discovery.py b/tests/providers/test_provider_client_tool_discovery.py index f4bc940f81..0910c6356d 100644 --- a/tests/providers/test_provider_client_tool_discovery.py +++ b/tests/providers/test_provider_client_tool_discovery.py @@ -9,6 +9,8 @@ from free_claude_code.config.settings import Settings from free_claude_code.core.anthropic.stream_contracts import parse_sse_text from free_claude_code.core.openai_responses import OpenAIResponsesRequest +from free_claude_code.providers.admission_policy import ProviderAdmissionLimits +from free_claude_code.providers.admission_registry import ProviderAdmissionRegistry from free_claude_code.providers.openai_chat import OpenAIChatProvider from free_claude_code.providers.runtime.runtime import create_provider from tests.core.openai_responses.test_client_tool_discovery import AGENTS, SEARCH @@ -257,6 +259,16 @@ def upstream(request: httpx2.Request) -> httpx2.Response: groq_api_key="test", mistral_api_key="test", ), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings( + Settings( + open_router_api_key="test", + nvidia_nim_api_key="test", + groq_api_key="test", + mistral_api_key="test", + ) + ) + ), ) ), ) diff --git a/tests/providers/test_provider_runtime.py b/tests/providers/test_provider_runtime.py index 02ecce18ad..da9a216f24 100644 --- a/tests/providers/test_provider_runtime.py +++ b/tests/providers/test_provider_runtime.py @@ -44,6 +44,8 @@ ZENMUX_DEFAULT_BASE, ) from free_claude_code.providers.admission import ProviderAdmissionController +from free_claude_code.providers.admission_policy import ProviderAdmissionLimits +from free_claude_code.providers.admission_registry import ProviderAdmissionRegistry from free_claude_code.providers.cloudflare import CloudflareProvider from free_claude_code.providers.deepseek import DeepSeekProvider from free_claude_code.providers.gemini import GeminiProvider @@ -261,7 +263,11 @@ async def test_poolside_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("poolside", settings) + provider = await create_provider( + "poolside", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "Poolside AI" assert descriptor.credential_env == "POOLSIDE_API_KEY" @@ -286,7 +292,11 @@ async def test_llm7_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("llm7", settings) + provider = await create_provider( + "llm7", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "LLM7.io" assert descriptor.credential_env == "LLM7_API_KEY" @@ -311,7 +321,11 @@ async def test_experiential_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("experiential", settings) + provider = await create_provider( + "experiential", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "Experiential Labs" assert descriptor.credential_env == "EXPLABS_API_KEY" @@ -339,7 +353,11 @@ async def test_cheaperinference_provider_config_uses_key_base_and_proxy() -> Non config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("cheaperinference", settings) + provider = await create_provider( + "cheaperinference", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "Cheaper Inference" assert descriptor.credential_env == "CHEAPER_INFERENCE_API_KEY" @@ -365,7 +383,11 @@ async def test_orcarouter_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("orcarouter", settings) + provider = await create_provider( + "orcarouter", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "OrcaRouter" assert descriptor.credential_env == "ORCAROUTER_API_KEY" @@ -390,7 +412,11 @@ async def test_xai_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("xai", settings) + provider = await create_provider( + "xai", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "xAI (Grok)" assert descriptor.credential_env == "XAI_API_KEY" @@ -410,7 +436,11 @@ async def test_qwencloud_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("qwencloud", settings) + provider = await create_provider( + "qwencloud", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "QwenCloud Token Plan" assert descriptor.credential_env == "QWENCLOUD_API_KEY" @@ -430,7 +460,11 @@ async def test_cline_pass_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("cline_pass", settings) + provider = await create_provider( + "cline_pass", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "ClinePass" assert descriptor.credential_env == "CLINE_API_KEY" @@ -454,7 +488,11 @@ async def test_qwencloud_coding_provider_config_uses_key_base_and_proxy() -> Non config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("qwencloud_coding", settings) + provider = await create_provider( + "qwencloud_coding", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "QwenCloud Coding Plan" assert descriptor.credential_env == "QWENCLOUD_CODING_API_KEY" @@ -474,7 +512,11 @@ async def test_together_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("together", settings) + provider = await create_provider( + "together", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "Together AI" assert descriptor.credential_env == "TOGETHER_API_KEY" @@ -494,7 +536,11 @@ async def test_deepinfra_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("deepinfra", settings) + provider = await create_provider( + "deepinfra", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "DeepInfra" assert descriptor.credential_env == "DEEPINFRA_API_KEY" @@ -514,7 +560,11 @@ async def test_siliconflow_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("siliconflow", settings) + provider = await create_provider( + "siliconflow", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "SiliconFlow" assert descriptor.credential_env == "SILICONFLOW_API_KEY" @@ -535,7 +585,11 @@ async def test_nebius_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("nebius", settings) + provider = await create_provider( + "nebius", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "Nebius Token Factory" assert descriptor.credential_env == "NEBIUS_API_KEY" @@ -559,7 +613,11 @@ async def test_scaleway_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("scaleway", settings) + provider = await create_provider( + "scaleway", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "Scaleway" assert descriptor.credential_env == "SCW_SECRET_KEY" @@ -581,7 +639,11 @@ async def test_chutes_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("chutes", settings) + provider = await create_provider( + "chutes", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "Chutes" assert descriptor.credential_env == "CHUTES_API_KEY" @@ -605,7 +667,11 @@ async def test_featherless_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("featherless", settings) + provider = await create_provider( + "featherless", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "Featherless AI" assert descriptor.credential_env == "FEATHERLESS_API_KEY" @@ -627,7 +693,11 @@ async def test_agnes_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("agnes", settings) + provider = await create_provider( + "agnes", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "Agnes AI" assert descriptor.credential_env == "AGNES_API_KEY" @@ -648,7 +718,11 @@ async def test_zenmux_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("zenmux", settings) + provider = await create_provider( + "zenmux", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "ZenMux" assert descriptor.credential_env == "ZENMUX_API_KEY" @@ -669,7 +743,11 @@ async def test_wandb_provider_config_uses_key_base_and_proxy() -> None: config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("wandb", settings) + provider = await create_provider( + "wandb", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "W&B Inference" assert descriptor.credential_env == "WANDB_API_KEY" @@ -752,7 +830,11 @@ async def test_local_provider_factory_resolves_catalog_static_credential( config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider(provider_id, settings) + provider = await create_provider( + provider_id, + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert config.api_key == expected_api_key assert isinstance(provider, OpenAIChatProvider) @@ -788,7 +870,11 @@ async def test_zai_api_provider_config_uses_shared_key_general_base_and_own_prox config = build_provider_config(descriptor, settings) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("zai_api", settings) + provider = await create_provider( + "zai_api", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert descriptor.display_name == "Z.ai API" assert descriptor.credential_env == "ZAI_API_KEY" @@ -840,7 +926,11 @@ async def test_create_cloudflare_provider_uses_account_scoped_base_url(): ) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("cloudflare", settings) + provider = await create_provider( + "cloudflare", + settings, + ProviderAdmissionRegistry(ProviderAdmissionLimits.from_settings(settings)), + ) assert isinstance(provider, CloudflareProvider) assert provider._base_url == ( @@ -851,7 +941,13 @@ async def test_create_cloudflare_provider_uses_account_scoped_base_url(): @pytest.mark.asyncio async def test_opencode_zen_provider_config_uses_explicit_id_and_name(): with patch("httpx.AsyncClient"): - provider = await create_provider("opencode_zen", _make_settings()) + provider = await create_provider( + "opencode_zen", + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + ) assert isinstance(provider, OpenCodeProvider) assert str(provider._client.base_url).rstrip("/") == "https://opencode.ai/zen/v1" @@ -863,7 +959,13 @@ async def test_opencode_zen_provider_config_uses_explicit_id_and_name(): @pytest.mark.asyncio async def test_opencode_go_provider_config_uses_correct_base_url_and_name(): with patch("httpx.AsyncClient"): - provider = await create_provider("opencode_go", _make_settings()) + provider = await create_provider( + "opencode_go", + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + ) assert isinstance(provider, OpenCodeProvider) assert str(provider._client.base_url).rstrip("/") == "https://opencode.ai/zen/go/v1" @@ -954,7 +1056,13 @@ def test_build_provider_config_cohere_uses_api_key_and_proxy() -> None: @pytest.mark.asyncio async def test_create_provider_uses_openai_chat_openrouter_by_default(): with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): - provider = await create_provider("open_router", _make_settings()) + provider = await create_provider( + "open_router", + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + ) assert isinstance(provider, OpenRouterProvider) @@ -1062,7 +1170,7 @@ async def test_create_provider_instantiates_each_builtin(): patch("free_claude_code.providers.openai_api.provider.AsyncOpenAI"), patch("httpx.AsyncClient"), patch( - "free_claude_code.providers.runtime.factory.ProviderAdmissionController", + "free_claude_code.providers.admission.ProviderAdmissionController", return_value=sentinel_admission, ) as admission_factory, ): @@ -1070,6 +1178,9 @@ async def test_create_provider_instantiates_each_builtin(): provider = await create_provider( provider_id, settings, + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(settings) + ), provider_loaders={ key: lambda factory=factory: factory for key, factory in injected_factories.items() @@ -1095,7 +1206,12 @@ async def test_create_provider_instantiates_each_builtin(): @pytest.mark.asyncio async def test_provider_runtime_caches_by_provider_id(): - runtime = ProviderRuntime(_make_settings()) + runtime = ProviderRuntime( + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + ) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): first = await runtime.resolve_provider("nvidia_nim") @@ -1111,7 +1227,7 @@ async def test_provider_creation_retries_after_shared_failure(cancelled): provider = MagicMock(cleanup=AsyncMock()) attempts = 0 - async def construct(_id, _settings): + async def construct(_id, _settings, _admission_registry): nonlocal attempts attempts += 1 attempt = attempts @@ -1123,7 +1239,13 @@ async def construct(_id, _settings): raise RuntimeError("construction failed") return provider - runtime = ProviderRuntime(_make_settings(), provider_constructor=construct) + runtime = ProviderRuntime( + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + provider_constructor=construct, + ) waiting = [] try: for attempt in (1, 2): @@ -1161,7 +1283,7 @@ async def test_finished_creation_callback_cannot_forget_a_new_retry(): attempts = 0 retry = None - async def construct(_id, _settings): + async def construct(_id, _settings, _admission_registry): nonlocal attempts, retry attempts += 1 if attempts == 1: @@ -1172,7 +1294,13 @@ async def construct(_id, _settings): await release.wait() return provider - runtime = ProviderRuntime(_make_settings(), provider_constructor=construct) + runtime = ProviderRuntime( + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + provider_constructor=construct, + ) later = None try: with pytest.raises(RuntimeError, match="construction failed"): @@ -1200,14 +1328,20 @@ async def test_cancelled_acquisition_keeps_construction_available_to_next_caller provider = MagicMock(cleanup=AsyncMock()) attempts = 0 - async def construct(_id, _settings): + async def construct(_id, _settings, _admission_registry): nonlocal attempts attempts += 1 entered.set() await release.wait() return provider - runtime = ProviderRuntime(_make_settings(), provider_constructor=construct) + runtime = ProviderRuntime( + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + provider_constructor=construct, + ) first = asyncio.create_task(runtime.resolve_provider("nvidia_nim")) second = None try: @@ -1236,13 +1370,19 @@ async def test_abandoned_creation_failure_has_no_unretrieved_exception(): previous_handler = loop.get_exception_handler() errors = [] - async def construct(_id, _settings): + async def construct(_id, _settings, _admission_registry): entered.set() await release.wait() failed.set() raise RuntimeError("abandoned construction failed") - runtime = ProviderRuntime(_make_settings(), provider_constructor=construct) + runtime = ProviderRuntime( + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + provider_constructor=construct, + ) waiting = asyncio.create_task(runtime.resolve_provider("nvidia_nim")) loop.set_exception_handler(lambda _loop, context: errors.append(context)) try: @@ -1272,14 +1412,20 @@ async def test_shutdown_drains_construction_that_finishes_during_cancellation(): entered = asyncio.Event() provider = MagicMock(cleanup=AsyncMock()) - async def construct(_id, _settings): + async def construct(_id, _settings, _admission_registry): entered.set() try: await asyncio.Event().wait() except asyncio.CancelledError: return provider - runtime = ProviderRuntime(_make_settings(), provider_constructor=construct) + runtime = ProviderRuntime( + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + provider_constructor=construct, + ) waiting = asyncio.create_task(runtime.resolve_provider("nvidia_nim")) try: await entered.wait() @@ -1297,7 +1443,12 @@ async def construct(_id, _settings): @pytest.mark.asyncio async def test_provider_runtime_provider_owns_one_admission_controller() -> None: - runtime = ProviderRuntime(_make_settings()) + runtime = ProviderRuntime( + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + ) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): first = await runtime.resolve_provider("nvidia_nim") @@ -1310,8 +1461,18 @@ async def test_provider_runtime_provider_owns_one_admission_controller() -> None @pytest.mark.asyncio async def test_separate_provider_runtimes_never_share_admission_controllers() -> None: - first_runtime = ProviderRuntime(_make_settings()) - second_runtime = ProviderRuntime(_make_settings()) + first_runtime = ProviderRuntime( + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + ) + second_runtime = ProviderRuntime( + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + ) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): first = await first_runtime.resolve_provider("nvidia_nim") @@ -1325,7 +1486,12 @@ async def test_separate_provider_runtimes_never_share_admission_controllers() -> @pytest.mark.asyncio async def test_different_providers_have_independent_admission_controllers() -> None: - runtime = ProviderRuntime(_make_settings()) + runtime = ProviderRuntime( + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + ) with patch("free_claude_code.providers.openai_chat.client.AsyncOpenAI"): nim = await runtime.resolve_provider("nvidia_nim") @@ -1339,7 +1505,15 @@ async def test_different_providers_have_independent_admission_controllers() -> N @pytest.mark.asyncio async def test_unknown_provider_raises_unknown_provider_type_error(): with pytest.raises(UnknownProviderError, match="Unknown provider_type"): - (await create_provider("unknown", _make_settings())) + ( + await create_provider( + "unknown", + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + ) + ) @pytest.mark.asyncio @@ -1349,7 +1523,13 @@ async def test_provider_runtime_cleanup_runs_all_even_if_one_fails() -> None: p1.cleanup = AsyncMock(side_effect=RuntimeError("first")) p2 = MagicMock() p2.cleanup = AsyncMock() - runtime = ProviderRuntime(_make_settings(), {"a": p1, "b": p2}) + runtime = ProviderRuntime( + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + {"a": p1, "b": p2}, + ) with pytest.raises(RuntimeError, match="first"): await runtime.cleanup() @@ -1386,6 +1566,9 @@ async def cleanup_second() -> None: third.cleanup = AsyncMock() runtime = ProviderRuntime( _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), {"first": first, "second": second, "third": third}, ) cleanup_task = asyncio.create_task(runtime.cleanup()) @@ -1417,7 +1600,13 @@ async def test_provider_runtime_cleanup_exceptiongroup_on_multiple_failures() -> p1.cleanup = AsyncMock(side_effect=RuntimeError("a")) p2 = MagicMock() p2.cleanup = AsyncMock(side_effect=RuntimeError("b")) - runtime = ProviderRuntime(_make_settings(), {"x": p1, "y": p2}) + runtime = ProviderRuntime( + _make_settings(), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(_make_settings()) + ), + {"x": p1, "y": p2}, + ) with pytest.raises(ExceptionGroup) as exc_info: await runtime.cleanup() diff --git a/tests/providers/test_qwencloud.py b/tests/providers/test_qwencloud.py index 03738553b3..8e6a1e5e6b 100644 --- a/tests/providers/test_qwencloud.py +++ b/tests/providers/test_qwencloud.py @@ -30,8 +30,6 @@ def qwencloud_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-qwencloud-key", base_url=QWENCLOUD_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="qwencloud"), ) diff --git a/tests/providers/test_qwencloud_coding.py b/tests/providers/test_qwencloud_coding.py index 7a3be80953..fca1b05a61 100644 --- a/tests/providers/test_qwencloud_coding.py +++ b/tests/providers/test_qwencloud_coding.py @@ -28,8 +28,6 @@ def qwencloud_coding_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-qwencloud-coding-key", base_url=QWENCLOUD_CODING_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="qwencloud_coding"), ) diff --git a/tests/providers/test_sambanova.py b/tests/providers/test_sambanova.py index dc82b2195d..cf5d56780b 100644 --- a/tests/providers/test_sambanova.py +++ b/tests/providers/test_sambanova.py @@ -25,8 +25,6 @@ def sambanova_config(): return make_provider_config( api_key="test_sambanova_key", base_url=SAMBANOVA_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ) diff --git a/tests/providers/test_scaleway.py b/tests/providers/test_scaleway.py index 96969b166d..5d38e64020 100644 --- a/tests/providers/test_scaleway.py +++ b/tests/providers/test_scaleway.py @@ -33,8 +33,6 @@ def scaleway_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-scaleway-key", base_url=SCALEWAY_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="scaleway"), ) diff --git a/tests/providers/test_siliconflow.py b/tests/providers/test_siliconflow.py index ee1ec0d9a8..7b46ed4d1b 100644 --- a/tests/providers/test_siliconflow.py +++ b/tests/providers/test_siliconflow.py @@ -45,8 +45,6 @@ def siliconflow_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-siliconflow-key", base_url=SILICONFLOW_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="siliconflow"), ) @@ -235,8 +233,6 @@ def build_client(*args: Any, **kwargs: Any) -> AsyncOpenAI: make_provider_config( api_key="wire-siliconflow-key", base_url=SILICONFLOW_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="siliconflow"), ) diff --git a/tests/providers/test_streaming_errors.py b/tests/providers/test_streaming_errors.py index 9c19e95fe1..6b05a5d8dc 100644 --- a/tests/providers/test_streaming_errors.py +++ b/tests/providers/test_streaming_errors.py @@ -127,8 +127,6 @@ def _make_provider(): config = make_provider_config( api_key="test_key", base_url="https://test.api.nvidia.com/v1", - rate_limit=10, - rate_window=60, ) return NvidiaNimProvider( config, diff --git a/tests/providers/test_together.py b/tests/providers/test_together.py index 7bd1bc4776..7dcceb79cc 100644 --- a/tests/providers/test_together.py +++ b/tests/providers/test_together.py @@ -31,8 +31,6 @@ def together_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-together-key", base_url=TOGETHER_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="together"), ) diff --git a/tests/providers/test_vercel.py b/tests/providers/test_vercel.py index daaa3883a6..8ba68f4946 100644 --- a/tests/providers/test_vercel.py +++ b/tests/providers/test_vercel.py @@ -27,8 +27,6 @@ def vercel_config(): return make_provider_config( api_key="test_vercel_key", base_url=VERCEL_AI_GATEWAY_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ) diff --git a/tests/providers/test_wafer.py b/tests/providers/test_wafer.py index 1172c7261f..d46aeb02ae 100644 --- a/tests/providers/test_wafer.py +++ b/tests/providers/test_wafer.py @@ -25,8 +25,6 @@ def wafer_config(): return make_provider_config( api_key="test-wafer-key", base_url=WAFER_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ) diff --git a/tests/providers/test_wandb.py b/tests/providers/test_wandb.py index 0184f68e67..9cf0d665fb 100644 --- a/tests/providers/test_wandb.py +++ b/tests/providers/test_wandb.py @@ -33,8 +33,6 @@ def wandb_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-wandb-key", base_url=WANDB_INFERENCE_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="wandb"), ) @@ -209,8 +207,6 @@ def build_client(*args: Any, **kwargs: Any) -> AsyncOpenAI: make_provider_config( api_key="wire-wandb-key", base_url=WANDB_INFERENCE_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="wandb"), ) diff --git a/tests/providers/test_xai.py b/tests/providers/test_xai.py index f023d49c9d..5fe601ed3f 100644 --- a/tests/providers/test_xai.py +++ b/tests/providers/test_xai.py @@ -32,8 +32,6 @@ def xai_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-xai-key", base_url=XAI_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="xai"), ) diff --git a/tests/providers/test_zai.py b/tests/providers/test_zai.py index f9f3f7e3b2..25dfeb8b13 100644 --- a/tests/providers/test_zai.py +++ b/tests/providers/test_zai.py @@ -33,8 +33,6 @@ def zai_provider(request: pytest.FixtureRequest): make_provider_config( api_key="test_zai_key", base_url=_ZAI_BASES[provider_id], - rate_limit=10, - rate_window=60, ), admission=immediate_admission(), ) diff --git a/tests/providers/test_zenmux.py b/tests/providers/test_zenmux.py index 91cd6031f6..fcf7b2eae5 100644 --- a/tests/providers/test_zenmux.py +++ b/tests/providers/test_zenmux.py @@ -43,8 +43,6 @@ def zenmux_provider() -> OpenAIChatProvider: make_provider_config( api_key="test-zenmux-key", base_url=ZENMUX_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="zenmux"), ) @@ -456,8 +454,6 @@ def build_client(*args: Any, **kwargs: Any) -> AsyncOpenAI: make_provider_config( api_key="wire-zenmux-key", base_url=ZENMUX_DEFAULT_BASE, - rate_limit=10, - rate_window=60, ), admission=immediate_admission(provider_name="zenmux"), ) diff --git a/tests/runtime/test_application_runtime.py b/tests/runtime/test_application_runtime.py index 1feb66d1af..7ee11375f5 100644 --- a/tests/runtime/test_application_runtime.py +++ b/tests/runtime/test_application_runtime.py @@ -37,8 +37,8 @@ class TrackingRuntime(ProviderRuntime): - def __init__(self, settings: Settings) -> None: - super().__init__(settings) + def __init__(self, settings: Settings, admission_registry) -> None: + super().__init__(settings, admission_registry) self.cleanup_calls = 0 async def cleanup(self) -> None: @@ -105,11 +105,11 @@ def __init__(self) -> None: self.fail = False self.events: list[str] = [] - def __call__(self, settings: Settings) -> ProviderRuntime: + def __call__(self, settings: Settings, admission_registry) -> ProviderRuntime: self.events.append(f"construct:{settings.model}") if self.fail: raise RuntimeError("candidate failed") - runtime = TrackingRuntime(settings) + runtime = TrackingRuntime(settings, admission_registry) self.runtimes.append(runtime) return runtime @@ -233,8 +233,9 @@ def _runtime_with_admin_provider( ) manager = ProviderRuntimeManager( settings, - runtime_factory=lambda snapshot: ProviderRuntime( + runtime_factory=lambda snapshot, admission_registry: ProviderRuntime( snapshot, + admission_registry, {"nvidia_nim": provider}, ), ) diff --git a/tests/runtime/test_custom_provider_lifecycle.py b/tests/runtime/test_custom_provider_lifecycle.py index 96b7835200..647a56a8d6 100644 --- a/tests/runtime/test_custom_provider_lifecycle.py +++ b/tests/runtime/test_custom_provider_lifecycle.py @@ -6,6 +6,8 @@ from free_claude_code.application.model_metadata import ProviderModelInfo from free_claude_code.config.custom_providers import CustomProviderDefinition from free_claude_code.config.settings import Settings +from free_claude_code.providers.admission_policy import ProviderAdmissionLimits +from free_claude_code.providers.admission_registry import ProviderAdmissionRegistry from free_claude_code.providers.custom import CustomProvider from free_claude_code.providers.runtime import ProviderRuntime from free_claude_code.runtime.provider_manager import ProviderRuntimeManager @@ -32,7 +34,12 @@ async def commit(): async def test_runtime_constructs_one_custom_owner_lazily(): - runtime = ProviderRuntime(settings(model_ids=["manual"])) + runtime = ProviderRuntime( + settings(model_ids=["manual"]), + ProviderAdmissionRegistry( + ProviderAdmissionLimits.from_settings(settings(model_ids=["manual"])) + ), + ) try: assert not runtime.is_cached(ID) one, two = await asyncio.gather( @@ -70,8 +77,8 @@ async def wait(): await release.wait() return frozenset() - def blocked(settings): - runtime = factory(settings) + def blocked(settings, admission_registry): + runtime = factory(settings, admission_registry) assert isinstance(runtime, FakeRuntime) runtime.provider.list_model_infos.side_effect = wait return runtime diff --git a/tests/runtime/test_integration_startup.py b/tests/runtime/test_integration_startup.py index d0ce318f3e..9efe7cbea6 100644 --- a/tests/runtime/test_integration_startup.py +++ b/tests/runtime/test_integration_startup.py @@ -50,13 +50,13 @@ def runtime(): return_value=frozenset({ProviderModelInfo("one")}) ) - async def construct(_id, _settings): + async def construct(_id, _settings, _admission_registry): return provider manager = ProviderRuntimeManager( settings, - runtime_factory=lambda snapshot: ProviderRuntime( - snapshot, provider_constructor=construct + runtime_factory=lambda snapshot, admission_registry: ProviderRuntime( + snapshot, admission_registry, provider_constructor=construct ), ) return ApplicationRuntime( diff --git a/tests/runtime/test_progressive_startup.py b/tests/runtime/test_progressive_startup.py index c166ca0c90..ede123bbdb 100644 --- a/tests/runtime/test_progressive_startup.py +++ b/tests/runtime/test_progressive_startup.py @@ -47,13 +47,15 @@ async def list_models(): provider.list_model_infos = AsyncMock(side_effect=list_models) - async def construct(_provider_id, _settings): + async def construct(_provider_id, _settings, _admission_registry): return provider manager = ProviderRuntimeManager( _settings(), - runtime_factory=lambda settings: ProviderRuntime( - settings, provider_constructor=construct + runtime_factory=lambda settings, admission_registry: ProviderRuntime( + settings, + admission_registry, + provider_constructor=construct, ), ) runtime = ApplicationRuntime( @@ -94,13 +96,15 @@ async def slow_models(): return_value=frozenset({ProviderModelInfo("one")}) ) - async def construct(provider_id, _settings): + async def construct(provider_id, _settings, _admission_registry): return slow if provider_id == "groq" else fast manager = ProviderRuntimeManager( _settings(), - runtime_factory=lambda settings: ProviderRuntime( - settings, provider_constructor=construct + runtime_factory=lambda settings, admission_registry: ProviderRuntime( + settings, + admission_registry, + provider_constructor=construct, ), ) manager.start_model_list_refresh() @@ -153,7 +157,7 @@ async def test_failed_discovery_construction_recovers_in_same_generation(refresh ) attempts = 0 - async def construct(_id, _settings): + async def construct(_id, _settings, _admission_registry): nonlocal attempts attempts += 1 if attempts == 1: @@ -162,8 +166,10 @@ async def construct(_id, _settings): manager = ProviderRuntimeManager( _settings(), - runtime_factory=lambda settings: ProviderRuntime( - settings, provider_constructor=construct + runtime_factory=lambda settings, admission_registry: ProviderRuntime( + settings, + admission_registry, + provider_constructor=construct, ), ) generation = manager.current_generation_id @@ -193,7 +199,7 @@ async def test_recovery_wait_uses_remaining_budget_without_orphaning_waiter( attempts = 0 construction = None - async def construct(_id, _settings): + async def construct(_id, _settings, _admission_registry): nonlocal attempts, construction attempts += 1 if attempts == 1: @@ -205,8 +211,10 @@ async def construct(_id, _settings): manager = ProviderRuntimeManager( _settings(), - runtime_factory=lambda settings: ProviderRuntime( - settings, provider_constructor=construct + runtime_factory=lambda settings, admission_registry: ProviderRuntime( + settings, + admission_registry, + provider_constructor=construct, ), ) wait = InitializationWait(1 if cancel else 0.01) @@ -269,7 +277,7 @@ async def test_late_retired_discovery_cannot_replace_current_metadata(): entered, release = asyncio.Event(), asyncio.Event() providers = [] - def factory(settings): + def factory(settings, admission_registry): provider = MagicMock(spec=BaseProvider) providers.append(provider) ordinal = len(providers) @@ -282,10 +290,14 @@ async def models(): provider.list_model_infos = AsyncMock(side_effect=models) - async def construct(_provider_id, _settings): + async def construct(_provider_id, _settings, _admission_registry): return provider - return ProviderRuntime(settings, provider_constructor=construct) + return ProviderRuntime( + settings, + admission_registry, + provider_constructor=construct, + ) manager = ProviderRuntimeManager(_settings(), runtime_factory=factory) old = await manager.acquire() @@ -326,13 +338,15 @@ async def models(): provider.list_model_infos = AsyncMock(side_effect=models) - async def construct(_provider_id, _settings): + async def construct(_provider_id, _settings, _admission_registry): return provider manager = ProviderRuntimeManager( _settings(), - runtime_factory=lambda settings: ProviderRuntime( - settings, provider_constructor=construct + runtime_factory=lambda settings, admission_registry: ProviderRuntime( + settings, + admission_registry, + provider_constructor=construct, ), model_catalog_publisher=CodexModelCatalogPublisher(path), ) @@ -360,7 +374,7 @@ async def test_catalog_wait_follows_replacement( entered, release = asyncio.Event(), asyncio.Event() count = 0 - def runtime(settings): + def runtime(settings, admission_registry): nonlocal count count += 1 old = count == 1 @@ -374,10 +388,14 @@ async def models(): provider.list_model_infos = AsyncMock(side_effect=models) - async def construct(_id, _settings): + async def construct(_id, _settings, _admission_registry): return provider - return ProviderRuntime(settings, provider_constructor=construct) + return ProviderRuntime( + settings, + admission_registry, + provider_constructor=construct, + ) path = tmp_path / "catalog.json" manager = ProviderRuntimeManager( @@ -442,13 +460,15 @@ async def models(): provider.list_model_infos = AsyncMock(side_effect=models) - async def construct(_id, _settings): + async def construct(_id, _settings, _admission_registry): return provider manager = ProviderRuntimeManager( _settings(), - runtime_factory=lambda settings: ProviderRuntime( - settings, provider_constructor=construct + runtime_factory=lambda settings, admission_registry: ProviderRuntime( + settings, + admission_registry, + provider_constructor=construct, ), ) wait = InitializationWait(1 if cancel else 0.01) diff --git a/tests/runtime/test_provider_manager.py b/tests/runtime/test_provider_manager.py index 616b977d7c..5270b67aab 100644 --- a/tests/runtime/test_provider_manager.py +++ b/tests/runtime/test_provider_manager.py @@ -45,7 +45,7 @@ def __init__(self) -> None: self.runtimes: list[FakeRuntime] = [] self.error: Exception | None = None - def __call__(self, settings: Settings) -> ProviderRuntime: + def __call__(self, settings: Settings, admission_registry) -> ProviderRuntime: if self.error is not None: raise self.error runtime = FakeRuntime(settings) @@ -302,7 +302,7 @@ async def test_replacement_keeps_leased_generation_open_until_final_release() -> @pytest.mark.asyncio -async def test_hot_replacement_owns_admission_per_provider_generation() -> None: +async def test_hot_replacement_shares_admission_and_retires_clients() -> None: first_settings = _settings("nvidia_nim/one") second_settings = _settings("nvidia_nim/two") clients: list[MagicMock] = [] @@ -331,7 +331,7 @@ def create_client(*_args: object, **_kwargs: object) -> MagicMock: assert isinstance(old_provider, NvidiaNimProvider) assert isinstance(new_provider, NvidiaNimProvider) assert new_provider is not old_provider - assert new_provider._admission is not old_provider._admission + assert new_provider._admission is old_provider._admission assert (await old_lease.resolve_provider("nvidia_nim")) is old_provider clients[0].close.assert_not_awaited() diff --git a/tests/runtime/test_retired_chat.py b/tests/runtime/test_retired_chat.py index 80a9dfa260..4f1151c4d3 100644 --- a/tests/runtime/test_retired_chat.py +++ b/tests/runtime/test_retired_chat.py @@ -40,7 +40,9 @@ def _runtime(): store.initialize() manager = ProviderRuntimeManager( store.read().settings, - runtime_factory=lambda snapshot: ProviderRuntime(snapshot, {}), + runtime_factory=lambda snapshot, admission_registry: ProviderRuntime( + snapshot, admission_registry, {} + ), ) return ApplicationRuntime( manager, configuration=ConfigurationService(store), transcriber=None diff --git a/tests/runtime/test_shared_provider_admission.py b/tests/runtime/test_shared_provider_admission.py new file mode 100644 index 0000000000..8c623159cf --- /dev/null +++ b/tests/runtime/test_shared_provider_admission.py @@ -0,0 +1,266 @@ +"""Admission protections survive replacement of their provider clients.""" + +import asyncio +from typing import cast +from unittest.mock import AsyncMock, patch + +import httpx +import pytest + +from free_claude_code.config.settings import Settings +from free_claude_code.providers.admission import ProviderOperationKind +from free_claude_code.providers.admission_policy import ProviderAdmissionLimits +from free_claude_code.providers.admission_registry import ProviderAdmissionRegistry +from free_claude_code.providers.nvidia_nim import NvidiaNimProvider +from free_claude_code.runtime.provider_manager import ProviderRuntimeManager +from tests.runtime.test_provider_manager import RuntimeFactory + + +@pytest.mark.asyncio +@pytest.mark.parametrize("protection", ["concurrency", "rate", "recovery"]) +async def test_replacement_cannot_bypass_existing_protection(protection: str) -> None: + settings = Settings( + nvidia_nim_api_key="test-key", + provider_rate_limit=1 if protection == "rate" else 1000, + provider_rate_window=60, + provider_max_concurrency=1, + ) + manager = ProviderRuntimeManager(settings) + old_lease = await manager.acquire() + pending = None + old_attempt = None + try: + old = cast( + NvidiaNimProvider, + await manager._current.runtime.resolve_provider("nvidia_nim"), + ) + old_attempt = await old._admission.start_execution().open_attempt( + ProviderOperationKind.GENERATION + ) + if protection == "recovery": + response = httpx.Response( + 503, + headers={"retry-after": "60"}, + request=httpx.Request("POST", "https://provider.test"), + ) + await old_attempt.fail( + httpx.HTTPStatusError( + "unavailable", request=response.request, response=response + ) + ) + else: + await old_attempt.accept() + if protection != "concurrency": + await old_attempt.aclose() + + with patch.object( + NvidiaNimProvider, "list_model_infos", AsyncMock(return_value=frozenset()) + ): + await manager.replace( + settings.model_copy(update={"log_raw_sse_events": True}), + commit=AsyncMock(), + ) + new = cast( + NvidiaNimProvider, + await manager._current.runtime.resolve_provider("nvidia_nim"), + ) + pending = asyncio.create_task( + new._admission.start_execution().open_attempt( + ProviderOperationKind.GENERATION + ) + ) + await asyncio.sleep(0) + assert not pending.done(), f"replacement bypassed {protection}" + finally: + if pending is not None: + pending.cancel() + outcomes = await asyncio.gather(pending, return_exceptions=True) + if not isinstance(outcomes[0], BaseException): + await outcomes[0].aclose() + if old_attempt is not None: + await old_attempt.aclose() + await old_lease.release() + await manager.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("outcome", ["success", "failure", "cancel"]) +async def test_limits_publish_only_with_successful_settings(outcome): + settings = Settings(provider_rate_limit=1000, provider_max_concurrency=1) + manager = ProviderRuntimeManager(settings, runtime_factory=RuntimeFactory()) + controller = manager._admission_registry.get("nvidia_nim") + first = await controller.start_execution().open_attempt( + ProviderOperationKind.GENERATION + ) + pending = asyncio.create_task( + controller.start_execution().open_attempt(ProviderOperationKind.GENERATION) + ) + entered, release = asyncio.Event(), asyncio.Event() + + async def commit(): + entered.set() + await release.wait() + if outcome == "failure": + raise OSError("cannot save") + + replacement = asyncio.create_task( + manager.replace( + settings.model_copy(update={"provider_max_concurrency": 2}), commit=commit + ) + ) + await entered.wait() + assert not pending.done() + if outcome == "cancel": + replacement.cancel() + else: + release.set() + results = await asyncio.gather(replacement, return_exceptions=True) + if outcome == "success": + assert results == [2] + await (await asyncio.wait_for(pending, 1)).aclose() + else: + assert isinstance(results[0], BaseException) + await asyncio.sleep(0) + assert not pending.done() + pending.cancel() + await asyncio.gather(pending, return_exceptions=True) + await first.aclose() + await manager.close() + + +@pytest.mark.asyncio +async def test_blocked_publication_keeps_old_limits_until_generation_switch(): + settings = Settings(provider_rate_limit=1000, provider_max_concurrency=1) + manager = ProviderRuntimeManager(settings, runtime_factory=RuntimeFactory()) + controller = manager._admission_registry.get("nvidia_nim") + first = await controller.start_execution().open_attempt( + ProviderOperationKind.GENERATION + ) + pending = asyncio.create_task( + controller.start_execution().open_attempt(ProviderOperationKind.GENERATION) + ) + committed = asyncio.Event() + + async def commit(): + committed.set() + + await manager._publication_lock.acquire() + replacement = asyncio.create_task( + manager.replace( + settings.model_copy(update={"provider_max_concurrency": 2}), commit=commit + ) + ) + await committed.wait() + assert not pending.done() + assert manager.current_generation_id == 1 + manager._publication_lock.release() + await replacement + await (await asyncio.wait_for(pending, 1)).aclose() + await first.aclose() + await manager.close() + + +@pytest.mark.asyncio +async def test_old_lazy_construction_uses_current_limits_and_survives_retirement(): + settings = Settings( + nvidia_nim_api_key="test", provider_rate_limit=1000, provider_max_concurrency=1 + ) + manager = ProviderRuntimeManager(settings) + old_lease = await manager.acquire() + old_runtime = manager._current.runtime + with patch.object( + NvidiaNimProvider, "list_model_infos", AsyncMock(return_value=frozenset()) + ): + await manager.replace( + settings.model_copy(update={"provider_max_concurrency": 2}), + commit=AsyncMock(), + ) + old = cast(NvidiaNimProvider, await old_runtime.resolve_provider("nvidia_nim")) + new = cast( + NvidiaNimProvider, + await manager._current.runtime.resolve_provider("nvidia_nim"), + ) + first = await old._admission.start_execution().open_attempt( + ProviderOperationKind.GENERATION + ) + second = await asyncio.wait_for( + new._admission.start_execution().open_attempt( + ProviderOperationKind.GENERATION + ), + 1, + ) + blocked = asyncio.create_task( + new._admission.start_execution().open_attempt( + ProviderOperationKind.GENERATION + ) + ) + await asyncio.sleep(0) + assert not blocked.done() + blocked.cancel() + await asyncio.gather(blocked, return_exceptions=True) + await first.aclose() + await second.aclose() + await old_lease.release() + await asyncio.wait_for( + new._admission.start_execution().run_call( + AsyncMock(return_value="ok"), + operation_kind=ProviderOperationKind.GENERATION, + ), + 1, + ) + await manager.close() + + +@pytest.mark.asyncio +async def test_separate_provider_ids_and_managers_have_independent_capacity(): + limits = ProviderAdmissionLimits(1, 60, 1) + registry = ProviderAdmissionRegistry(limits) + other = ProviderAdmissionRegistry(limits) + attempts = [] + for owner, provider in [(registry, "one"), (registry, "two"), (other, "one")]: + attempts.append( + await asyncio.wait_for( + owner.get(provider) + .start_execution() + .open_attempt(ProviderOperationKind.GENERATION), + 1, + ) + ) + for attempt in attempts: + await attempt.aclose() + registry.close() + other.close() + + +@pytest.mark.asyncio +async def test_custom_registry_entry_lives_until_retired_generation_cleanup_succeeds(): + from free_claude_code.config.custom_providers import CustomProviderDefinition + + custom_id = "custom_12345678123412341234123456789abc" + settings = Settings( + custom_providers=( + CustomProviderDefinition( + provider_id=custom_id, + display_name="Example", + base_url="https://example.test/v1", + api_format="openai_chat", + ), + ) + ) + factory = RuntimeFactory() + manager = ProviderRuntimeManager(settings, runtime_factory=factory) + lease = await manager.acquire() + controller = manager._admission_registry.get(custom_id) + await manager.replace( + settings.model_copy(update={"custom_providers": ()}), commit=AsyncMock() + ) + assert manager._admission_registry.get(custom_id) is controller + factory.runtimes[0].cleanup_error = OSError("not closed yet") + await lease.release() + assert manager._admission_registry.get(custom_id) is controller + factory.runtimes[0].cleanup_error = None + await manager._close_generation(manager._retired[1], forced=False) + assert custom_id not in manager._admission_registry._controllers + await manager.close() + with pytest.raises(RuntimeError, match="closed"): + manager._admission_registry.get("nvidia_nim") diff --git a/tests/runtime/test_vscode_chat_sync.py b/tests/runtime/test_vscode_chat_sync.py index 89e95f7447..a6847f4bca 100644 --- a/tests/runtime/test_vscode_chat_sync.py +++ b/tests/runtime/test_vscode_chat_sync.py @@ -35,7 +35,11 @@ async def construct(*_): manager = ProviderRuntimeManager( settings, - runtime_factory=lambda s: ProviderRuntime(s, provider_constructor=construct), + runtime_factory=lambda s, admission_registry: ProviderRuntime( + s, + admission_registry, + provider_constructor=construct, + ), ) app = ApplicationRuntime( manager,