diff --git a/ARCHITECTURE.md b/ARCHITECTURE.md index b2b827a0925b2bf2a4317027cb08750056e9e139..159eb2620f3411f34957c2c5c05077b95c9d7490 100644 --- a/ARCHITECTURE.md +++ b/ARCHITECTURE.md @@ -153,7 +153,8 @@ On final shutdown it best-effort kills registered child processes. [runtime/bootstrap.py](src/free_claude_code/runtime/bootstrap.py) is the single production composition function. The CLI supervisor supplies one settings snapshot and its restart callback; bootstrap -configures logging, constructs the runtime owners, supplies +configures logging, constructs the runtime owners and the configured voice +transcriber, supplies [api/ports.py](src/free_claude_code/api/ports.py) to the pure API factory, and returns the ASGI application. [api/app.py](src/free_claude_code/api/app.py) registers routers, HTTP correlation middleware, and exception handlers around @@ -162,8 +163,17 @@ runtime resources. `app.state.services` is the only runtime state published to FastAPI. [runtime/application.py](src/free_claude_code/runtime/application.py) owns process startup and shutdown, optional messaging, -the managed CLI session manager, Admin pending state, and the injected restart -callback. [runtime/asgi.py](src/free_claude_code/runtime/asgi.py) drives that owner from ASGI lifespan messages and preserves +the selected transcriber, the managed CLI session manager, Admin pending state, +and the injected restart callback. Shutdown is serialized and ordered: quiesce +messaging ingress, cancel and drain workflow/CLI work, flush persistence, close +delivery, close transcription, then close providers. An owner reference is +released only after its cleanup succeeds; cancellation or failure leaves the +incomplete graph retryable. Teardown stops at a failed dependency gate rather +than closing resources that still-live upstream work may need, and the ASGI +adapter reports that incomplete graph as lifespan shutdown failure. Cleanup is +completion-driven: generic timeouts do not cancel half-closed external resources; +the process supervisor owns any force-termination deadline. +[runtime/asgi.py](src/free_claude_code/runtime/asgi.py) drives that owner from ASGI lifespan messages and preserves the concise startup-failure contract. [runtime/provider_manager.py](src/free_claude_code/runtime/provider_manager.py) is the only owner that constructs, publishes, @@ -173,7 +183,11 @@ streaming responses release it from the response iterator's `finally` path on completion, failure, cancellation, or disconnect. A provider-only Admin Apply prepares a candidate and commits configuration before publication. New requests then use the candidate while old streams finish on the retired generation; its -last lease closes it exactly once. +last lease closes it exactly once. Final shutdown rejects new acquisition and +replacement, waits every lease, and awaits the same manager-owned cleanup task +even if the initiating request or lease release is cancelled. Failed generation +or unpublished-candidate cleanup remains owned and retryable; the manager does +not become terminal or clear its model catalog until every owned runtime closes. The manager also owns one application-lifetime provider model catalog and its single best-effort discovery task. The catalog survives provider replacement. @@ -366,7 +380,14 @@ credential env var, default base URL, settings attribute names, and proxy suppor [providers/runtime/](src/free_claude_code/providers/runtime/) owns construction details for one closable provider generation: factory wiring, provider configuration, lazy -provider instances, and transport cleanup. Application-level generation +provider instances, provider-owned rate limiters, and transport cleanup. Each +lazy provider receives a fresh `ProviderRateLimiter`; there is no process +singleton or second limiter registry. The provider cache already guarantees one +provider and limiter per provider ID within a generation. Retired generations +retain their own synchronization state until request leases drain, while new +generations and separate server instances never reuse it. Hot replacement +therefore begins with fresh quota state; an old and new generation enforce +independent budgets while old request leases drain. Application-level generation publication, request leases, model metadata, discovery orchestration, and configured-model validation belong to `ProviderRuntimeManager` in the runtime package. This separates a single generation's resources from process-lifetime @@ -632,13 +653,14 @@ If `MESSAGING_PLATFORM` is `none`, or if the selected platform token is missing, the messaging bridge is skipped. `ApplicationRuntime` privately owns the selected platform runtime, the -`MessagingWorkflow`, and its managed CLI session manager. The workflow owns +`MessagingWorkflow`, configured `Transcriber`, and managed CLI session manager. +The workflow owns conversation snapshot restoration and final persistence flush. The API sees only the `SessionControlPort` used to preserve `/stop` behavior. The platform factory returns a `MessagingPlatformComponents` bundle from [messaging/platforms/ports.py](src/free_claude_code/messaging/platforms/ports.py): a -`MessagingRuntime` for lifecycle and inbound callbacks, an `OutboundMessenger` +`MessagingRuntime` with separate `quiesce()` and `close()` phases, an `OutboundMessenger` for queued sends/edits/deletes, and an optional `VoiceCancellation` port for reply-scoped `/clear` during voice transcription. Workflow code depends on these ports, not on Telegram or Discord SDK objects. @@ -646,7 +668,16 @@ these ports, not on Telegram or Discord SDK objects. Runtime adapters in [messaging/platforms/telegram.py](src/free_claude_code/messaging/platforms/telegram.py) and [messaging/platforms/discord.py](src/free_claude_code/messaging/platforms/discord.py) own SDK client -lifecycle, event subscription, inbound handoff, and voice-note handoff. Inbound +lifecycle, event subscription, inbound handoff, voice-note handoff, and one +injected `MessagingRateLimiter`. The platform factory creates a fresh limiter +for the selected runtime. `quiesce()` stops new SDK ingress and drains active +handlers while delivery remains available; after workflow tasks settle, +`close()` drains the outbox and limiter. Discord additionally retains, observes, +and drains its long-lived client task and inbound-handler tasks, so an SDK exit +after initial readiness immediately withdraws the runtime's connected state. +Telegram retries initialization and polling as separate repeatable steps; it +never restarts an already-running SDK application after polling bootstrap fails. +Separate application runtimes cannot share or stop each other's queue. Inbound normalization lives in [messaging/platforms/telegram_inbound.py](src/free_claude_code/messaging/platforms/telegram_inbound.py) and [messaging/platforms/discord_inbound.py](src/free_claude_code/messaging/platforms/discord_inbound.py). @@ -654,14 +685,27 @@ Outbound SDK calls live in [messaging/platforms/telegram_io.py](src/free_claude_code/messaging/platforms/telegram_io.py) and [messaging/platforms/discord_io.py](src/free_claude_code/messaging/platforms/discord_io.py). Shared delivery policy lives in [messaging/platforms/outbox.py](src/free_claude_code/messaging/platforms/outbox.py), -which owns queued send/edit/list-based delete, dedup keys, limiter delegation, -and fire-and-forget behavior. Workflow and command code request deletion of +which requires that limiter directly and owns queued send/edit/list-based delete, +dedup keys, and retained fire-and-forget tasks. Shutdown cancels and awaits both +queued limiter work and arbitrary outbox work; there is no optional unthrottled +fallback, and both owners reject admission once close begins. Workflow and command code request deletion of message ID lists; platform IO decides whether to use native batch deletion (Telegram) or internal per-message deletion (Discord). Shared voice-note orchestration lives in [messaging/platforms/voice_flow.py](src/free_claude_code/messaging/platforms/voice_flow.py), which owns -pending voice registration, temp-file cleanup, transcription, cancellation, error -replies, and the handoff to `IncomingMessage`. +pending voice registration, file-size validation, temp-file cleanup, +transcription, cancellation, error replies, and the handoff to +`IncomingMessage`. It depends only on the consumer-owned `Transcriber` protocol +from [messaging/voice.py](src/free_claude_code/messaging/voice.py). Bootstrap selects either the +instance-owned local Whisper `TranscriptionService` or the provider-owned +`NvidiaNimTranscriber`. Messaging no longer imports a provider adapter, and the +local service retains only one lazy pipeline for its immutable runtime settings; +caller cancellation waits for thread-backed transcription to actually exit +before temporary files, pipelines, or credentials are released. The NIM adapter +closes its per-call authenticated gRPC channel before that worker exits. Changing the +credential used by an active voice backend through Admin is therefore +restart-required, while the same provider credential remains hot-replaceable +when voice does not use it. [messaging/workflow.py](src/free_claude_code/messaging/workflow.py) contains `MessagingWorkflow`, the platform-agnostic coordinator. It owns dependencies, callback wiring, stop/clear @@ -755,6 +799,10 @@ Logging defaults are conservative: message logging are opt-in. - Messaging text, transcription previews, CLI diagnostics, and detailed messaging exception strings are controlled by separate diagnostic flags. +- Process logging, server/managed-CLI authentication, and messaging diagnostics + are captured by their lifecycle owners at construction. Admin marks those + settings restart-required so an Apply cannot report success while an existing + runtime continues using stale security or privacy policy. - Values under keys that look like API keys, authorization, tokens, or secrets are redacted by trace helpers where structured traces are emitted. diff --git a/pyproject.toml b/pyproject.toml index 194e28b46ad95705e110b47017c81d84866e925d..98268242d2073dc223ba512c45622dbc4b6404db 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "hatchling.build" [project] name = "free-claude-code" -version = "3.4.17" +version = "3.4.18" description = "Middleware between Claude Code CLI (Anthropic API) and NVIDIA NIM" readme = "README.md" requires-python = ">=3.14.0" diff --git a/smoke/capabilities.py b/smoke/capabilities.py index 2f8a46e2f9c5e9f1044aacd617cee7bbacd2615f..71ebb2a281ab99747c003ec0b2daf0da471d23db 100644 --- a/smoke/capabilities.py +++ b/smoke/capabilities.py @@ -150,9 +150,13 @@ CAPABILITY_CONTRACTS: tuple[CapabilityContract, ...] = ( "provider_proxy_timeout_config", "free_claude_code.providers.runtime.ProviderRuntime", "provider proxy, timeout, and rate-limit settings", - "provider client and scoped limiter config", + "provider client and instance-owned limiter config", "provider construction failure", - ("tests/api/test_dependencies.py", "tests/providers/test_provider_runtime.py"), + ( + "tests/api/test_dependencies.py", + "tests/providers/test_provider_runtime.py", + "tests/providers/test_provider_rate_limit.py", + ), ), CapabilityContract( "provider_routing", @@ -271,7 +275,7 @@ CAPABILITY_CONTRACTS: tuple[CapabilityContract, ...] = ( "provider_runtime", "rate_limit_and_disconnect", "smart_rate_limiting", - "free_claude_code.providers.rate_limit.GlobalRateLimiter", + "free_claude_code.providers.rate_limit.ProviderRateLimiter", "concurrent provider requests and 429/disconnect failures", "proactive throttle, retry, cleanup", "mapped provider error or smoke skip for upstream disconnect", @@ -357,6 +361,9 @@ CAPABILITY_CONTRACTS: tuple[CapabilityContract, ...] = ( ( "tests/messaging/test_discord_platform.py", "tests/messaging/test_telegram.py", + "tests/messaging/test_limiter.py", + "tests/messaging/test_platform_outbox.py", + "tests/runtime/test_application_runtime.py", ), ("test_telegram_bot_api_permissions", "test_discord_bot_api_permissions"), ), @@ -410,13 +417,15 @@ CAPABILITY_CONTRACTS: tuple[CapabilityContract, ...] = ( "voice", "voice_transcription", "voice_notes", - "free_claude_code.messaging.voice.VoiceTranscriptionService", + "free_claude_code.messaging.voice.Transcriber", "Discord/Telegram audio file and voice backend settings", "transcribed prompt routed to handler", "missing optional extra or backend error shown to user", ( "tests/messaging/test_voice_handlers.py", "tests/messaging/test_transcription.py", + "tests/messaging/test_transcription_nim.py", + "tests/messaging/test_platform_voice_flow.py", ), ("test_voice_transcription_backend_when_explicitly_enabled",), ), diff --git a/smoke/features.py b/smoke/features.py index 1e7823a7ab63fdefc5978a222da7b68521620436..f5efeb96f37051e36c829ca4b42ccf1132144e3d 100644 --- a/smoke/features.py +++ b/smoke/features.py @@ -214,6 +214,9 @@ FEATURE_INVENTORY: tuple[FeatureCoverage, ...] = ( ( "tests/messaging/test_discord_platform.py", "tests/messaging/test_telegram.py", + "tests/messaging/test_limiter.py", + "tests/messaging/test_platform_outbox.py", + "tests/runtime/test_application_runtime.py", ), ( "test_telegram_bot_api_permissions", diff --git a/smoke/lib/e2e.py b/smoke/lib/e2e.py index bf46d45798798aec25552c1b16f6546931e5b52f..d90f4415a1f02705d9b1b087d92ce5a18bee49b4 100644 --- a/smoke/lib/e2e.py +++ b/smoke/lib/e2e.py @@ -323,7 +323,10 @@ class FakePlatform: async def start(self) -> None: return None - async def stop(self) -> None: + async def quiesce(self) -> None: + return None + + async def close(self) -> None: for task in self._tasks: if not task.done(): task.cancel() diff --git a/smoke/prereq/test_voice_prereq_live.py b/smoke/prereq/test_voice_prereq_live.py index b0e6a9a5087d3d6c6c3a5c7849659ab724843ba3..78091eb1ecd5979e9567253166017cc1a3e78817 100644 --- a/smoke/prereq/test_voice_prereq_live.py +++ b/smoke/prereq/test_voice_prereq_live.py @@ -5,13 +5,16 @@ from pathlib import Path import pytest -from free_claude_code.messaging.transcription import transcribe_audio +from free_claude_code.messaging.transcription import TranscriptionService +from free_claude_code.messaging.voice import Transcriber +from free_claude_code.providers.nvidia_nim.voice import NvidiaNimTranscriber from smoke.lib.config import SmokeConfig pytestmark = [pytest.mark.live, pytest.mark.smoke_target("voice")] -def test_voice_transcription_backend_when_explicitly_enabled( +@pytest.mark.asyncio +async def test_voice_transcription_backend_when_explicitly_enabled( smoke_config: SmokeConfig, tmp_path: Path ) -> None: if not smoke_config.settings.voice_note_enabled: @@ -21,16 +24,24 @@ def test_voice_transcription_backend_when_explicitly_enabled( wav_path = tmp_path / "smoke-tone.wav" _write_tone_wav(wav_path) + transcriber: Transcriber + if smoke_config.settings.whisper_device == "nvidia_nim": + transcriber = NvidiaNimTranscriber( + model=smoke_config.settings.whisper_model, + api_key=smoke_config.settings.nvidia_nim_api_key, + ) + else: + transcriber = TranscriptionService( + model=smoke_config.settings.whisper_model, + device=smoke_config.settings.whisper_device, + huggingface_api_key=smoke_config.settings.huggingface_api_key, + ) try: - t_kw: dict[str, str] = { - "whisper_model": smoke_config.settings.whisper_model, - "whisper_device": smoke_config.settings.whisper_device, - } - if smoke_config.settings.whisper_device == "nvidia_nim": - t_kw["nvidia_nim_api_key"] = smoke_config.settings.nvidia_nim_api_key - text = transcribe_audio(wav_path, "audio/wav", **t_kw) + text = await transcriber.transcribe(wav_path) except ImportError as exc: pytest.skip(str(exc)) + finally: + await transcriber.close() assert isinstance(text, str) assert text.strip() diff --git a/smoke/product/test_voice_product_live.py b/smoke/product/test_voice_product_live.py index 1c691f4dd7209bf47854d3dd36489bec1e5d112b..23947d27b1e9723c43b8797fcd114606294f12ae 100644 --- a/smoke/product/test_voice_product_live.py +++ b/smoke/product/test_voice_product_live.py @@ -3,7 +3,8 @@ from pathlib import Path import pytest -from free_claude_code.messaging.transcription import transcribe_audio +from free_claude_code.messaging.transcription import TranscriptionService +from free_claude_code.providers.nvidia_nim.voice import NvidiaNimTranscriber from smoke.lib.config import SmokeConfig from smoke.lib.e2e import VoiceFixtureDriver @@ -11,7 +12,10 @@ pytestmark = [pytest.mark.live] @pytest.mark.smoke_target("voice") -def test_voice_local_backend_e2e(smoke_config: SmokeConfig, tmp_path: Path) -> None: +@pytest.mark.asyncio +async def test_voice_local_backend_e2e( + smoke_config: SmokeConfig, tmp_path: Path +) -> None: if not smoke_config.settings.voice_note_enabled: pytest.skip("missing_env: VOICE_NOTE_ENABLED is false") if os.getenv("FCC_SMOKE_RUN_VOICE") != "1": @@ -21,22 +25,25 @@ def test_voice_local_backend_e2e(smoke_config: SmokeConfig, tmp_path: Path) -> N wav_path = tmp_path / "voice-local-product.wav" VoiceFixtureDriver.write_tone_wav(wav_path) + transcriber = TranscriptionService( + model=smoke_config.settings.whisper_model, + device=smoke_config.settings.whisper_device, + huggingface_api_key=smoke_config.settings.huggingface_api_key, + ) try: - text = transcribe_audio( - wav_path, - "audio/wav", - whisper_model=smoke_config.settings.whisper_model, - whisper_device=smoke_config.settings.whisper_device, - ) + text = await transcriber.transcribe(wav_path) except ImportError as exc: pytest.skip(f"missing_env: {exc}") + finally: + await transcriber.close() assert isinstance(text, str) assert text.strip() @pytest.mark.smoke_target("voice") -def test_voice_nim_backend_e2e(smoke_config: SmokeConfig, tmp_path: Path) -> None: +@pytest.mark.asyncio +async def test_voice_nim_backend_e2e(smoke_config: SmokeConfig, tmp_path: Path) -> None: if not smoke_config.settings.voice_note_enabled: pytest.skip("missing_env: VOICE_NOTE_ENABLED is false") if os.getenv("FCC_SMOKE_RUN_VOICE") != "1": @@ -48,13 +55,14 @@ def test_voice_nim_backend_e2e(smoke_config: SmokeConfig, tmp_path: Path) -> Non wav_path = tmp_path / "voice-nim-product.wav" VoiceFixtureDriver.write_tone_wav(wav_path) - text = transcribe_audio( - wav_path, - "audio/wav", - whisper_model=smoke_config.settings.whisper_model, - whisper_device="nvidia_nim", - nvidia_nim_api_key=smoke_config.settings.nvidia_nim_api_key, + transcriber = NvidiaNimTranscriber( + model=smoke_config.settings.whisper_model, + api_key=smoke_config.settings.nvidia_nim_api_key, ) + try: + text = await transcriber.transcribe(wav_path) + finally: + await transcriber.close() assert isinstance(text, str) assert text.strip() diff --git a/src/free_claude_code/config/admin/manifest.py b/src/free_claude_code/config/admin/manifest.py index 0565967df07862ba1bf1688aa56b7214b05bbfd0..3db30b914258a3fa33fea66c1ca49fdcae22cf70 100644 --- a/src/free_claude_code/config/admin/manifest.py +++ b/src/free_claude_code/config/admin/manifest.py @@ -168,6 +168,7 @@ _NON_PROVIDER_FIELDS: tuple[ConfigFieldSpec, ...] = ( settings_attr="anthropic_auth_token", default="freecc", secret=True, + restart_required=True, description="Protects Claude/API access. It is not admin-page login.", ), ConfigFieldSpec( @@ -424,6 +425,7 @@ _NON_PROVIDER_FIELDS: tuple[ConfigFieldSpec, ...] = ( settings_attr="debug_platform_edits", default="false", advanced=True, + restart_required=True, ), ConfigFieldSpec( "DEBUG_SUBAGENT_STACK", @@ -433,6 +435,7 @@ _NON_PROVIDER_FIELDS: tuple[ConfigFieldSpec, ...] = ( settings_attr="debug_subagent_stack", default="false", advanced=True, + restart_required=True, ), ConfigFieldSpec( "LOG_RAW_API_PAYLOADS", @@ -442,6 +445,7 @@ _NON_PROVIDER_FIELDS: tuple[ConfigFieldSpec, ...] = ( settings_attr="log_raw_api_payloads", default="false", advanced=True, + restart_required=True, ), ConfigFieldSpec( "LOG_RAW_SSE_EVENTS", @@ -460,6 +464,7 @@ _NON_PROVIDER_FIELDS: tuple[ConfigFieldSpec, ...] = ( settings_attr="log_api_error_tracebacks", default="false", advanced=True, + restart_required=True, ), ConfigFieldSpec( "LOG_RAW_MESSAGING_CONTENT", @@ -469,6 +474,7 @@ _NON_PROVIDER_FIELDS: tuple[ConfigFieldSpec, ...] = ( settings_attr="log_raw_messaging_content", default="false", advanced=True, + restart_required=True, ), ConfigFieldSpec( "LOG_RAW_CLI_DIAGNOSTICS", @@ -478,6 +484,7 @@ _NON_PROVIDER_FIELDS: tuple[ConfigFieldSpec, ...] = ( settings_attr="log_raw_cli_diagnostics", default="false", advanced=True, + restart_required=True, ), ConfigFieldSpec( "LOG_MESSAGING_ERROR_DETAILS", @@ -487,6 +494,7 @@ _NON_PROVIDER_FIELDS: tuple[ConfigFieldSpec, ...] = ( settings_attr="log_messaging_error_details", default="false", advanced=True, + restart_required=True, ), ConfigFieldSpec( "FCC_SMOKE_MODEL_NVIDIA_NIM", diff --git a/src/free_claude_code/config/admin/persistence.py b/src/free_claude_code/config/admin/persistence.py index 5dbd82d88b6fa47123a0a2d1520a3ce3cfa19a69..0e1320d3f1d8417c8926f73ef4c3f8eef9b34bf5 100644 --- a/src/free_claude_code/config/admin/persistence.py +++ b/src/free_claude_code/config/admin/persistence.py @@ -106,14 +106,25 @@ def validate_updates(updates: Mapping[str, Any]) -> dict[str, Any]: return prepare_admin_update(updates).validation_response() -def changed_pending_fields(updates: Mapping[str, Any]) -> list[str]: +def changed_pending_fields( + updates: Mapping[str, Any], + *, + settings: Settings, +) -> list[str]: """Return changed fields that require manual runtime action.""" state = load_value_state() pending: list[str] = [] for key, value in updates.items(): field = FIELD_BY_KEY.get(key) - if field is None or not (field.restart_required or field.session_sensitive): + if field is None or is_locked_source(state[key]["source"]): + continue + if field.secret and value == MASKED_SECRET: + continue + requires_restart = field.restart_required or field.session_sensitive + if not requires_restart: + requires_restart = _active_voice_credential(settings) == key + if not requires_restart: continue if normalize_for_env(value) == str(state[key]["value"]): continue @@ -121,6 +132,14 @@ def changed_pending_fields(updates: Mapping[str, Any]) -> list[str]: return pending +def _active_voice_credential(settings: Settings) -> str | None: + if not settings.voice_note_enabled: + return None + if settings.whisper_device == "nvidia_nim": + return "NVIDIA_NIM_API_KEY" + return "HUGGINGFACE_API_KEY" + + def prepare_admin_update(updates: Mapping[str, Any]) -> PreparedAdminUpdate: """Validate an update and construct its prospective Settings snapshot.""" @@ -128,7 +147,9 @@ def prepare_admin_update(updates: Mapping[str, Any]) -> PreparedAdminUpdate: effective_values = effective_values_for_validation(target_values) settings, errors = settings_from_values(effective_values) pending_fields = ( - tuple(changed_pending_fields(updates)) if settings is not None else () + tuple(changed_pending_fields(updates, settings=settings)) + if settings is not None + else () ) return PreparedAdminUpdate( target_values=target_values, diff --git a/src/free_claude_code/messaging/limiter.py b/src/free_claude_code/messaging/limiter.py index 8a75b78aedbc90fb65b98dd506efc8427fe8eb5a..584289ed3ebf4d1605d9f9e12fd9a6d10d19069c 100644 --- a/src/free_claude_code/messaging/limiter.py +++ b/src/free_claude_code/messaging/limiter.py @@ -1,9 +1,4 @@ -""" -Global Rate Limiter for Messaging Platforms. - -Centralizes outgoing message requests and ensures compliance with rate limits -using a strict sliding window algorithm and a task queue. -""" +"""Runtime-owned queued delivery for one messaging platform.""" import asyncio from collections import deque @@ -12,7 +7,6 @@ from typing import Any from loguru import logger -from free_claude_code.config.settings import get_settings from free_claude_code.core.rate_limit import ( StrictSlidingWindowLimiter as SlidingWindowLimiter, ) @@ -22,43 +16,21 @@ from .safe_diagnostics import format_exception_for_log class MessagingRateLimiter: """ - A thread-safe global rate limiter for messaging. + Rate limiter and compacting work queue for one messaging runtime. Uses a custom queue with task compaction (deduplication) to ensure only the latest version of a message update is processed. """ - _instance: MessagingRateLimiter | None = None - _lock = asyncio.Lock() - - def __new__(cls, *args, **kwargs): - return super().__new__(cls) - - @classmethod - async def get_instance( - cls, + def __init__( + self, *, - rate_limit: int = 1, - rate_window: float = 1.0, - ) -> MessagingRateLimiter: - """Get the singleton instance of the limiter. - - ``rate_limit`` and ``rate_window`` apply only when the singleton is first - created. Call :meth:`shutdown_instance` before changing parameters. - """ - async with cls._lock: - if cls._instance is None: - cls._instance = cls(rate_limit=rate_limit, rate_window=rate_window) - # Start the background worker (tracked for graceful shutdown). - cls._instance._start_worker() - return cls._instance - - def __init__(self, *, rate_limit: int, rate_window: float) -> None: - # Prevent double initialization in singleton - if hasattr(self, "_initialized"): - return - + rate_limit: int, + rate_window: float, + log_error_details: bool = False, + ) -> None: self.limiter = SlidingWindowLimiter(rate_limit, rate_window) + self._log_error_details = log_error_details # Custom queue state - using deque for O(1) popleft self._queue_list: deque[str] = deque() # Deque of dedup_keys in order self._queue_map: dict[ @@ -66,25 +38,27 @@ class MessagingRateLimiter: ] = {} self._condition = asyncio.Condition() self._shutdown = asyncio.Event() - self._worker_task: asyncio.Task | None = None - - self._initialized = True + self._worker_task: asyncio.Task[None] | None = None + self._background_tasks: set[asyncio.Task[None]] = set() + self._active_futures: list[asyncio.Future[Any]] = [] + self._closed = False self._paused_until = 0 logger.info( f"MessagingRateLimiter initialized ({rate_limit} req / {rate_window}s with Task Compaction)" ) - def _start_worker(self) -> None: - """Ensure the worker task exists.""" + def start(self) -> None: + """Start the owned worker on the current event loop.""" + if self._closed: + raise RuntimeError("Messaging rate limiter is closed.") if self._worker_task and not self._worker_task.done(): return - # Named task helps debugging shutdown hangs. self._worker_task = asyncio.create_task( self._worker(), name="msg-limiter-worker" ) - async def _worker(self): + async def _worker(self) -> None: """Background worker that processes queued messaging tasks.""" logger.info("MessagingRateLimiter worker started") while not self._shutdown.is_set(): @@ -99,6 +73,7 @@ class MessagingRateLimiter: dedup_key = self._queue_list.popleft() func, futures = self._queue_map.pop(dedup_key) + self._active_futures = futures # Check for manual pause (FloodWait) now = asyncio.get_event_loop().time() @@ -116,6 +91,19 @@ class MessagingRateLimiter: for f in futures: if not f.done(): f.set_result(result) + except asyncio.CancelledError: + for f in futures: + if not f.done(): + f.cancel() + worker = asyncio.current_task() + if self._shutdown.is_set() or ( + worker is not None and worker.cancelling() + ): + raise + logger.debug( + "Messaging operation cancelled for key {}; worker remains active", + dedup_key, + ) except Exception as e: # Report error to all futures and log it for f in futures: @@ -148,17 +136,26 @@ class MessagingRateLimiter: asyncio.get_event_loop().time() + wait_secs ) else: - d = get_settings().log_messaging_error_details logger.error( "Error in limiter worker for key {}: {}", dedup_key, - format_exception_for_log(e, log_full_message=d), + format_exception_for_log( + e, + log_full_message=self._log_error_details, + ), ) + finally: + self._active_futures = [] except asyncio.CancelledError: - break + for future in self._active_futures: + if not future.done(): + future.cancel() + self._active_futures = [] + if self._shutdown.is_set(): + break + raise except Exception as e: - d = get_settings().log_messaging_error_details - if d: + if self._log_error_details: logger.error( "MessagingRateLimiter worker critical error: {}", e, @@ -171,53 +168,84 @@ class MessagingRateLimiter: ) await asyncio.sleep(1) - async def shutdown(self, timeout: float = 2.0) -> None: - """Stop the background worker so process shutdown doesn't hang.""" + async def shutdown(self, timeout: float | None = None) -> None: + """Cancel queued work and stop every task owned by this limiter.""" + self._closed = True self._shutdown.set() - try: - async with self._condition: - self._condition.notify_all() - except Exception: - # Best-effort: condition may be bound to a closing loop. - pass - + async with self._condition: + queued_futures = [ + future + for _func, futures in self._queue_map.values() + for future in futures + ] + self._queue_list.clear() + self._queue_map.clear() + for future in queued_futures: + if not future.done(): + future.cancel() + for future in self._active_futures: + if not future.done(): + future.cancel() + self._condition.notify_all() + + cancellation: asyncio.CancelledError | None = None + timeout_error: TimeoutError | None = None task = self._worker_task - if not task or task.done(): - self._worker_task = None - return - - task.cancel() - try: - await asyncio.wait_for(task, timeout=timeout) - except TimeoutError: - logger.warning("MessagingRateLimiter worker did not stop before timeout") - except asyncio.CancelledError: - pass - except Exception as e: - d = get_settings().log_messaging_error_details - logger.debug( - "MessagingRateLimiter worker shutdown error: {}", - format_exception_for_log(e, log_full_message=d), - ) - finally: + if task and not task.done(): + task.cancel() + try: + drain = asyncio.gather(task, return_exceptions=True) + if timeout is None: + await drain + else: + await asyncio.wait_for(drain, timeout=timeout) + except TimeoutError as exc: + timeout_error = exc + except asyncio.CancelledError as exc: + cancellation = exc + if task is None or task.done(): self._worker_task = None - @classmethod - async def shutdown_instance(cls, timeout: float = 2.0) -> None: - """Shutdown and clear the singleton instance (safe to call multiple times).""" - inst = cls._instance - if not inst: - return - try: - await inst.shutdown(timeout=timeout) - finally: - cls._instance = None + background_tasks = tuple(self._background_tasks) + for background_task in background_tasks: + background_task.cancel() + if background_tasks: + try: + await asyncio.gather(*background_tasks, return_exceptions=True) + except asyncio.CancelledError as exc: + cancellation = exc + self._background_tasks.difference_update( + task for task in background_tasks if task.done() + ) - async def _enqueue_internal(self, func, future, dedup_key, front=False): + if cancellation is not None: + raise cancellation + if timeout_error is not None: + raise TimeoutError( + "MessagingRateLimiter worker did not stop before timeout" + ) from timeout_error + + async def _enqueue_internal( + self, + func: Callable[[], Awaitable[Any]], + future: asyncio.Future[Any], + dedup_key: str, + front: bool = False, + ) -> None: await self._enqueue_internal_multi(func, [future], dedup_key, front) - async def _enqueue_internal_multi(self, func, futures, dedup_key, front=False): + async def _enqueue_internal_multi( + self, + func: Callable[[], Awaitable[Any]], + futures: list[asyncio.Future[Any]], + dedup_key: str, + front: bool = False, + ) -> None: async with self._condition: + if self._closed: + raise RuntimeError("Messaging rate limiter is closed.") + if self._worker_task is None or self._worker_task.done(): + raise RuntimeError("Messaging rate limiter has not been started.") if dedup_key in self._queue_map: # Compaction: Update existing task with new func, append new futures _old_func, old_futures = self._queue_map[dedup_key] @@ -241,28 +269,33 @@ class MessagingRateLimiter: Enqueue a messaging task and return its future result. If dedup_key is provided, subsequent tasks with the same key will replace this one. """ + self._require_running() if dedup_key is None: # Unique key to avoid deduplication - dedup_key = f"task_{id(func)}_{asyncio.get_event_loop().time()}" + dedup_key = f"task_{id(func)}_{asyncio.get_running_loop().time()}" - future = asyncio.get_event_loop().create_future() - await self._enqueue_internal(func, future, dedup_key) + future = asyncio.get_running_loop().create_future() + try: + await self._enqueue_internal(func, future, dedup_key) + except BaseException: + future.cancel() + raise return await future def fire_and_forget( self, func: Callable[[], Awaitable[Any]], dedup_key: str | None = None - ): + ) -> None: """Enqueue a task without waiting for the result.""" + self._require_running() if dedup_key is None: - dedup_key = f"task_{id(func)}_{asyncio.get_event_loop().time()}" + dedup_key = f"task_{id(func)}_{asyncio.get_running_loop().time()}" - future = asyncio.get_event_loop().create_future() - - async def _wrapped(): + async def _wrapped() -> None: max_retries = 2 for attempt in range(max_retries + 1): try: - return await self.enqueue(func, dedup_key) + await self.enqueue(func, dedup_key) + return except Exception as e: error_msg = str(e).lower() # Only retry transient connectivity issues that might have slipped through @@ -271,8 +304,7 @@ class MessagingRateLimiter: x in error_msg for x in ["connect", "timeout", "broken"] ): wait = 2**attempt - d = get_settings().log_messaging_error_details - if d: + if self._log_error_details: logger.warning( "Limiter fire_and_forget transient error (attempt {}): {}. Retrying in {}s...", attempt + 1, @@ -289,14 +321,22 @@ class MessagingRateLimiter: await asyncio.sleep(wait) continue - d = get_settings().log_messaging_error_details logger.error( "Final error in fire_and_forget for key {}: {}", dedup_key, - format_exception_for_log(e, log_full_message=d), + format_exception_for_log( + e, + log_full_message=self._log_error_details, + ), ) - if not future.done(): - future.set_exception(e) break - _ = asyncio.create_task(_wrapped()) + task = asyncio.create_task(_wrapped(), name=f"msg-limiter:{dedup_key}") + self._background_tasks.add(task) + task.add_done_callback(self._background_tasks.discard) + + def _require_running(self) -> None: + if self._closed: + raise RuntimeError("Messaging rate limiter is closed.") + if self._worker_task is None or self._worker_task.done(): + raise RuntimeError("Messaging rate limiter has not been started.") diff --git a/src/free_claude_code/messaging/platforms/discord.py b/src/free_claude_code/messaging/platforms/discord.py index 892d6f8dd55b6373484d96a7538ede527153f0d7..655c65fedf1af639af6cd04f9daedd026e7b3739 100644 --- a/src/free_claude_code/messaging/platforms/discord.py +++ b/src/free_claude_code/messaging/platforms/discord.py @@ -9,8 +9,10 @@ from loguru import logger from free_claude_code.core.anthropic import format_user_error_preview +from ..limiter import MessagingRateLimiter from ..models import IncomingMessage from ..rendering.discord_markdown import format_status_discord +from ..voice import Transcriber from .discord_inbound import ( discord_text_message_from_event, discord_voice_request_from_event, @@ -55,8 +57,7 @@ if DISCORD_AVAILABLE and _discord_module is not None: self._runtime = runtime async def on_ready(self) -> None: - self._runtime._connected = True - logger.info("Discord platform connected") + self._runtime._mark_connected() async def on_message(self, message: Any) -> None: await self._runtime._handle_client_message(message) @@ -74,13 +75,8 @@ class DiscordRuntime: bot_token: str | None = None, allowed_channel_ids: str | None = None, *, - voice_note_enabled: bool = True, - whisper_model: str = "base", - whisper_device: str = "cpu", - huggingface_api_key: str = "", - nvidia_nim_api_key: str = "", - messaging_rate_limit: int = 1, - messaging_rate_window: float = 1.0, + limiter: MessagingRateLimiter, + transcriber: Transcriber | None, log_raw_messaging_content: bool = False, log_api_error_tracebacks: bool = False, ) -> None: @@ -102,30 +98,45 @@ class DiscordRuntime: self._client = _DiscordClient(self, intents) self._message_handler: InboundMessageHandler | None = None self._connected = False - self._limiter: Any | None = None - self._start_task: asyncio.Task | None = None + self._accepting_messages = False + self._ready = asyncio.Event() + self._inbound_tasks: set[asyncio.Task[Any]] = set() + self._limiter = limiter + self._start_task: asyncio.Task[None] | None = None self.outbound = DiscordMessenger( get_client=lambda: self._client, get_discord=_get_discord, - get_limiter=lambda: self._limiter, + limiter=limiter, ) self._voice_flow = VoiceNoteFlow( - voice_note_enabled=voice_note_enabled, - whisper_model=whisper_model, - whisper_device=whisper_device, - huggingface_api_key=huggingface_api_key, - nvidia_nim_api_key=nvidia_nim_api_key, + transcriber=transcriber, log_raw_messaging_content=log_raw_messaging_content, log_api_error_tracebacks=log_api_error_tracebacks, ) - self._messaging_rate_limit = messaging_rate_limit - self._messaging_rate_window = messaging_rate_window self._log_raw_messaging_content = log_raw_messaging_content self._log_api_error_tracebacks = log_api_error_tracebacks async def _handle_client_message(self, message: Any) -> None: """Adapter entry point used by the internal Discord client.""" - await self._on_discord_message(message) + if not self._accepting_messages: + return + task = asyncio.current_task() + if task is not None: + self._inbound_tasks.add(task) + try: + if self._accepting_messages: + await self._on_discord_message(message) + finally: + if task is not None: + self._inbound_tasks.discard(task) + + def _mark_connected(self) -> None: + """Publish Discord readiness while this runtime accepts ingress.""" + if not self._accepting_messages: + return + self._connected = True + self._ready.set() + logger.info("Discord platform connected") async def cancel_pending_voice( self, chat_id: str, reply_id: str @@ -185,46 +196,105 @@ class DiscordRuntime: if not self.bot_token: raise ValueError("DISCORD_BOT_TOKEN is required") - from ..limiter import MessagingRateLimiter - - self._limiter = await MessagingRateLimiter.get_instance( - rate_limit=self._messaging_rate_limit, - rate_window=self._messaging_rate_window, - ) + self._limiter.start() + self._accepting_messages = True + self._ready.clear() self._start_task = asyncio.create_task( self._client.start(self.bot_token), name="discord-client-start", ) + self._start_task.add_done_callback(self._observe_client_exit) + ready_task = asyncio.create_task( + self._ready.wait(), + name="discord-client-ready", + ) + try: + done, _pending = await asyncio.wait( + (self._start_task, ready_task), + timeout=30.0, + return_when=asyncio.FIRST_COMPLETED, + ) + if not done: + raise RuntimeError("Discord client failed to connect within timeout") + if self._start_task in done: + await self._start_task + raise RuntimeError("Discord client stopped before becoming ready") + if self._start_task.done(): + await self._start_task + raise RuntimeError("Discord client stopped unexpectedly") + finally: + ready_task.cancel() + await asyncio.gather(ready_task, return_exceptions=True) - max_wait = 30 - waited = 0.0 - while not self._connected and waited < max_wait: - await asyncio.sleep(0.5) - waited += 0.5 + logger.info("Discord platform started") - if not self._connected: - raise RuntimeError("Discord client failed to connect within timeout") + def _observe_client_exit(self, task: asyncio.Task[None]) -> None: + """Observe the long-lived Discord client task and publish lost readiness.""" + if task.cancelled(): + exception: BaseException | None = None + else: + exception = task.exception() - logger.info("Discord platform started") + was_connected = self._connected + if not self._accepting_messages: + return - async def stop(self) -> None: - """Stop Discord SDK resources.""" - if self._client.is_closed(): - self._connected = False + self._connected = False + self._ready.clear() + if not was_connected: return - await self._client.close() - if self._start_task and not self._start_task.done(): - try: - await asyncio.wait_for(self._start_task, timeout=5.0) - except TimeoutError, asyncio.CancelledError: - self._start_task.cancel() - with contextlib.suppress(asyncio.CancelledError): - await self._start_task + if exception is None: + logger.error("Discord client stopped unexpectedly") + elif self._log_api_error_tracebacks: + logger.error("Discord client stopped unexpectedly: {}", exception) + else: + logger.error( + "Discord client stopped unexpectedly: exc_type={}", + type(exception).__name__, + ) - self._connected = False - logger.info("Discord platform stopped") + async def quiesce(self) -> None: + """Stop Discord ingress and drain active SDK handlers.""" + self._accepting_messages = False + try: + if not self._client.is_closed(): + await self._client.close() + finally: + try: + await self._drain_start_task() + finally: + try: + await self._drain_inbound_tasks() + finally: + self._connected = False + self._ready.clear() + + async def close(self) -> None: + """Close Discord delivery resources after ingress is quiescent.""" + try: + await self.outbound.close() + finally: + await self._limiter.shutdown() + logger.info("Discord platform closed") + + async def _drain_start_task(self) -> None: + task = self._start_task + if task is None: + return + if not task.done(): + task.cancel() + try: + await asyncio.gather(task, return_exceptions=True) + finally: + if task.done() and self._start_task is task: + self._start_task = None + + async def _drain_inbound_tasks(self) -> None: + tasks = tuple(self._inbound_tasks) + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) def on_message(self, handler: Callable[[IncomingMessage], Awaitable[None]]) -> None: """Register the workflow callback for inbound messages.""" diff --git a/src/free_claude_code/messaging/platforms/discord_io.py b/src/free_claude_code/messaging/platforms/discord_io.py index 854986394babd89bdb66b79949fd0d740eb476fa..ae8dfd406a8a86d7fd01a42cdb78e144f43a8033 100644 --- a/src/free_claude_code/messaging/platforms/discord_io.py +++ b/src/free_claude_code/messaging/platforms/discord_io.py @@ -3,13 +3,13 @@ from collections.abc import Awaitable, Callable from typing import Any, cast +from ..limiter import MessagingRateLimiter from .outbox import PlatformOutbox DISCORD_MESSAGE_LIMIT = 2000 ClientGetter = Callable[[], Any] DiscordGetter = Callable[[], Any] -LimiterGetter = Callable[[], Any | None] def truncate_discord_message(text: str, limit: int = DISCORD_MESSAGE_LIMIT) -> str: @@ -27,12 +27,12 @@ class DiscordMessenger: *, get_client: ClientGetter, get_discord: DiscordGetter, - get_limiter: LimiterGetter, + limiter: MessagingRateLimiter, ) -> None: self._get_client = get_client self._get_discord = get_discord self._outbox = PlatformOutbox( - get_limiter=get_limiter, + limiter=limiter, send=self.send_message, edit=self.edit_message, delete_many=self.delete_messages, @@ -161,3 +161,7 @@ class DiscordMessenger: def fire_and_forget(self, task: Awaitable[Any]) -> None: """Execute a coroutine without awaiting it.""" self._outbox.fire_and_forget(task) + + async def close(self) -> None: + """Cancel outstanding outbound work.""" + await self._outbox.close() diff --git a/src/free_claude_code/messaging/platforms/factory.py b/src/free_claude_code/messaging/platforms/factory.py index 5d917914ff9b90c2a0df62c4b592682df44bcabf..1e3d5ba4f2fe21ea71258e248bc4fd54c810445c 100644 --- a/src/free_claude_code/messaging/platforms/factory.py +++ b/src/free_claude_code/messaging/platforms/factory.py @@ -4,6 +4,8 @@ from dataclasses import dataclass from loguru import logger +from ..limiter import MessagingRateLimiter +from ..voice import Transcriber from .ports import MessagingPlatformComponents @@ -16,14 +18,11 @@ class MessagingPlatformOptions: telegram_proxy_url: str = "" discord_bot_token: str | None = None allowed_discord_channels: str | None = None - voice_note_enabled: bool = True - whisper_model: str = "base" - whisper_device: str = "cpu" - huggingface_api_key: str = "" - nvidia_nim_api_key: str = "" + transcriber: Transcriber | None = None messaging_rate_limit: int = 1 messaging_rate_window: float = 1.0 log_raw_messaging_content: bool = False + log_messaging_error_details: bool = False log_api_error_tracebacks: bool = False @@ -45,17 +44,17 @@ def create_messaging_components( from .telegram import TelegramRuntime + limiter = MessagingRateLimiter( + rate_limit=opts.messaging_rate_limit, + rate_window=opts.messaging_rate_window, + log_error_details=opts.log_messaging_error_details, + ) runtime = TelegramRuntime( bot_token=bot_token, allowed_user_id=opts.allowed_telegram_user_id, telegram_proxy_url=opts.telegram_proxy_url, - voice_note_enabled=opts.voice_note_enabled, - whisper_model=opts.whisper_model, - whisper_device=opts.whisper_device, - huggingface_api_key=opts.huggingface_api_key, - nvidia_nim_api_key=opts.nvidia_nim_api_key, - messaging_rate_limit=opts.messaging_rate_limit, - messaging_rate_window=opts.messaging_rate_window, + limiter=limiter, + transcriber=opts.transcriber, log_raw_messaging_content=opts.log_raw_messaging_content, log_api_error_tracebacks=opts.log_api_error_tracebacks, ) @@ -74,16 +73,16 @@ def create_messaging_components( from .discord import DiscordRuntime + limiter = MessagingRateLimiter( + rate_limit=opts.messaging_rate_limit, + rate_window=opts.messaging_rate_window, + log_error_details=opts.log_messaging_error_details, + ) runtime = DiscordRuntime( bot_token=bot_token, allowed_channel_ids=opts.allowed_discord_channels, - voice_note_enabled=opts.voice_note_enabled, - whisper_model=opts.whisper_model, - whisper_device=opts.whisper_device, - huggingface_api_key=opts.huggingface_api_key, - nvidia_nim_api_key=opts.nvidia_nim_api_key, - messaging_rate_limit=opts.messaging_rate_limit, - messaging_rate_window=opts.messaging_rate_window, + limiter=limiter, + transcriber=opts.transcriber, log_raw_messaging_content=opts.log_raw_messaging_content, log_api_error_tracebacks=opts.log_api_error_tracebacks, ) diff --git a/src/free_claude_code/messaging/platforms/outbox.py b/src/free_claude_code/messaging/platforms/outbox.py index 6afeef58c4369e1e2e7e5cb0200b9d1c3e4d66e1..0b9c51d3b6b3540b012f7b1789226cab214e9661 100644 --- a/src/free_claude_code/messaging/platforms/outbox.py +++ b/src/free_claude_code/messaging/platforms/outbox.py @@ -5,13 +5,16 @@ import hashlib from collections.abc import Awaitable, Callable from typing import Any, cast +from loguru import logger + +from ..limiter import MessagingRateLimiter + SendOperation = Callable[ [str, str, str | None, str | None, str | None], Awaitable[str], ] EditOperation = Callable[[str, str, str, str | None], Awaitable[None]] DeleteManyOperation = Callable[[str, list[str]], Awaitable[None]] -LimiterGetter = Callable[[], Any | None] class PlatformOutbox: @@ -20,15 +23,17 @@ class PlatformOutbox: def __init__( self, *, - get_limiter: LimiterGetter, + limiter: MessagingRateLimiter, send: SendOperation, edit: EditOperation, delete_many: DeleteManyOperation, ) -> None: - self._get_limiter = get_limiter + self._limiter = limiter self._send = send self._edit = edit self._delete_many = delete_many + self._background_tasks: set[asyncio.Future[Any]] = set() + self._closed = False async def queue_send_message( self, @@ -40,15 +45,7 @@ class PlatformOutbox: message_thread_id: str | None = None, ) -> str | None: """Queue or immediately send a platform message.""" - limiter = self._get_limiter() - if limiter is None: - return await self._send( - chat_id, - text, - reply_to, - parse_mode, - message_thread_id, - ) + self._require_open() async def _send() -> str: return await self._send( @@ -60,9 +57,9 @@ class PlatformOutbox: ) if fire_and_forget: - limiter.fire_and_forget(_send) + self._limiter.fire_and_forget(_send) return None - return cast(str | None, await limiter.enqueue(_send)) + return cast(str | None, await self._limiter.enqueue(_send)) async def queue_edit_message( self, @@ -73,19 +70,16 @@ class PlatformOutbox: fire_and_forget: bool = True, ) -> None: """Queue or immediately edit a platform message.""" - limiter = self._get_limiter() - if limiter is None: - await self._edit(chat_id, message_id, text, parse_mode) - return + self._require_open() async def _edit() -> None: await self._edit(chat_id, message_id, text, parse_mode) dedup_key = f"edit:{chat_id}:{message_id}" if fire_and_forget: - limiter.fire_and_forget(_edit, dedup_key=dedup_key) + self._limiter.fire_and_forget(_edit, dedup_key=dedup_key) else: - await limiter.enqueue(_edit, dedup_key=dedup_key) + await self._limiter.enqueue(_edit, dedup_key=dedup_key) async def queue_delete_messages( self, @@ -94,28 +88,56 @@ class PlatformOutbox: fire_and_forget: bool = True, ) -> None: """Queue or immediately bulk-delete platform messages.""" + self._require_open() ids_snapshot = tuple(str(message_id) for message_id in message_ids) if not ids_snapshot: return - limiter = self._get_limiter() - if limiter is None: - await self._delete_many(chat_id, list(ids_snapshot)) - return - async def _delete_many() -> None: await self._delete_many(chat_id, list(ids_snapshot)) digest = hashlib.sha256("\x1f".join(ids_snapshot).encode()).hexdigest()[:16] dedup_key = f"del_bulk:{chat_id}:{digest}" if fire_and_forget: - limiter.fire_and_forget(_delete_many, dedup_key=dedup_key) + self._limiter.fire_and_forget(_delete_many, dedup_key=dedup_key) else: - await limiter.enqueue(_delete_many, dedup_key=dedup_key) + await self._limiter.enqueue(_delete_many, dedup_key=dedup_key) def fire_and_forget(self, task: Awaitable[Any]) -> None: - """Execute a coroutine or future without awaiting it.""" - if asyncio.iscoroutine(task): - _ = asyncio.create_task(task) - else: - _ = asyncio.ensure_future(task) + """Run and retain arbitrary outbound work until completion or shutdown.""" + future = asyncio.ensure_future(task) + if self._closed: + future.cancel() + raise RuntimeError("Platform outbox is closed.") + self._background_tasks.add(future) + future.add_done_callback(self._complete_background_task) + + async def close(self) -> None: + """Cancel and await arbitrary outbound work owned by this outbox.""" + if not self._closed: + self._closed = True + tasks = tuple(self._background_tasks) + for task in tasks: + task.cancel() + try: + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) + finally: + self._background_tasks.difference_update( + task for task in tasks if task.done() + ) + + def _complete_background_task(self, task: asyncio.Future[Any]) -> None: + self._background_tasks.discard(task) + if task.cancelled(): + return + error = task.exception() + if error is not None: + logger.error( + "Outbound background task failed: exc_type={}", + type(error).__name__, + ) + + def _require_open(self) -> None: + if self._closed: + raise RuntimeError("Platform outbox is closed.") diff --git a/src/free_claude_code/messaging/platforms/ports.py b/src/free_claude_code/messaging/platforms/ports.py index 0a1007586428e20b0f4ef7c79b5aff0e3e0c1661..ba847b7002a1c630142d003776469e049f4ec497 100644 --- a/src/free_claude_code/messaging/platforms/ports.py +++ b/src/free_claude_code/messaging/platforms/ports.py @@ -11,14 +11,16 @@ InboundMessageHandler = Callable[[IncomingMessage], Awaitable[None]] @runtime_checkable class MessagingRuntime(Protocol): - """Owns inbound SDK lifecycle for one messaging platform.""" + """Owns ingress and delivery lifecycle for one messaging platform.""" @property def name(self) -> str: ... async def start(self) -> None: ... - async def stop(self) -> None: ... + async def quiesce(self) -> None: ... + + async def close(self) -> None: ... def on_message(self, handler: InboundMessageHandler) -> None: ... diff --git a/src/free_claude_code/messaging/platforms/telegram.py b/src/free_claude_code/messaging/platforms/telegram.py index 52d3611d420dab13b684456ba94cebed3ebbbc16..9aecc67915ea59183da62b79cc7fe42f40396188 100644 --- a/src/free_claude_code/messaging/platforms/telegram.py +++ b/src/free_claude_code/messaging/platforms/telegram.py @@ -4,7 +4,6 @@ import asyncio import contextlib import os from collections.abc import Awaitable, Callable -from typing import Any # Opt-in to future behavior for python-telegram-bot (retry_after as timedelta). os.environ["PTB_TIMEDELTA"] = "1" @@ -13,8 +12,10 @@ from loguru import logger from free_claude_code.core.anthropic import format_user_error_preview +from ..limiter import MessagingRateLimiter from ..models import IncomingMessage from ..rendering.telegram_markdown import escape_md_v2 +from ..voice import Transcriber from .ports import InboundMessageHandler from .telegram_inbound import ( telegram_text_message_from_update, @@ -50,13 +51,8 @@ class TelegramRuntime: allowed_user_id: str | None = None, *, telegram_proxy_url: str = "", - voice_note_enabled: bool = True, - whisper_model: str = "base", - whisper_device: str = "cpu", - huggingface_api_key: str = "", - nvidia_nim_api_key: str = "", - messaging_rate_limit: int = 1, - messaging_rate_window: float = 1.0, + limiter: MessagingRateLimiter, + transcriber: Transcriber | None, log_raw_messaging_content: bool = False, log_api_error_tracebacks: bool = False, ) -> None: @@ -74,22 +70,16 @@ class TelegramRuntime: self._application: Application | None = None self._message_handler: InboundMessageHandler | None = None self._connected = False - self._limiter: Any | None = None + self._limiter = limiter self.outbound = TelegramMessenger( get_application=lambda: self._application, - get_limiter=lambda: self._limiter, + limiter=limiter, ) self._voice_flow = VoiceNoteFlow( - voice_note_enabled=voice_note_enabled, - whisper_model=whisper_model, - whisper_device=whisper_device, - huggingface_api_key=huggingface_api_key, - nvidia_nim_api_key=nvidia_nim_api_key, + transcriber=transcriber, log_raw_messaging_content=log_raw_messaging_content, log_api_error_tracebacks=log_api_error_tracebacks, ) - self._messaging_rate_limit = messaging_rate_limit - self._messaging_rate_window = messaging_rate_window self._log_raw_messaging_content = log_raw_messaging_content self._log_api_error_tracebacks = log_api_error_tracebacks @@ -128,51 +118,31 @@ class TelegramRuntime: connection_pool_size=8, connect_timeout=30.0, read_timeout=30.0 ) builder = Application.builder().token(self.bot_token).request(request) - self._application = builder.build() + application = builder.build() + self._application = application - self._application.add_handler( + application.add_handler( MessageHandler(filters.TEXT & (~filters.COMMAND), self._on_telegram_message) ) - self._application.add_handler(CommandHandler("start", self._on_start_command)) - self._application.add_handler( + application.add_handler(CommandHandler("start", self._on_start_command)) + application.add_handler( MessageHandler(filters.COMMAND, self._on_telegram_message) ) - self._application.add_handler( - MessageHandler(filters.VOICE, self._on_telegram_voice) - ) + application.add_handler(MessageHandler(filters.VOICE, self._on_telegram_voice)) - max_retries = 3 - for attempt in range(max_retries): - try: - await self._application.initialize() - await self._application.start() - if self._application.updater: - await self._application.updater.start_polling( - drop_pending_updates=False - ) - self._connected = True - break - except Exception as e: - if attempt < max_retries - 1: - wait_time = 2 * (attempt + 1) - logger.warning( - "Connection failed (attempt {}/{}): {}. Retrying in {}s...", - attempt + 1, - max_retries, - e, - wait_time, - ) - await asyncio.sleep(wait_time) - else: - logger.error("Failed to connect after {} attempts", max_retries) - raise - - from ..limiter import MessagingRateLimiter - - self._limiter = await MessagingRateLimiter.get_instance( - rate_limit=self._messaging_rate_limit, - rate_window=self._messaging_rate_window, + await self._retry_connection_step( + application.initialize, + step="initialization", ) + await application.start() + self._limiter.start() + updater = application.updater + if updater is not None: + await self._retry_connection_step( + lambda: updater.start_polling(drop_pending_updates=False), + step="polling", + ) + self._connected = True try: target = self.allowed_user_id @@ -193,15 +163,75 @@ class TelegramRuntime: logger.info("Telegram platform started (Bot API)") - async def stop(self) -> None: - """Stop Telegram polling and SDK resources.""" - if self._application and self._application.updater: - await self._application.updater.stop() - await self._application.stop() - await self._application.shutdown() + async def _retry_connection_step( + self, + operation: Callable[[], Awaitable[object]], + *, + step: str, + ) -> None: + """Retry one independently repeatable Telegram connection step.""" + max_attempts = 3 + for attempt in range(1, max_attempts + 1): + try: + await operation() + return + except Exception as exc: + if attempt == max_attempts: + logger.error( + "Telegram {} failed after {} attempts", + step, + max_attempts, + ) + raise + wait_time = 2 * attempt + if self._log_api_error_tracebacks: + logger.warning( + "Telegram {} failed (attempt {}/{}): {}. Retrying in {}s...", + step, + attempt, + max_attempts, + exc, + wait_time, + ) + else: + logger.warning( + "Telegram {} failed (attempt {}/{}): exc_type={}. Retrying in {}s...", + step, + attempt, + max_attempts, + type(exc).__name__, + wait_time, + ) + await asyncio.sleep(wait_time) - self._connected = False - logger.info("Telegram platform stopped") + async def quiesce(self) -> None: + """Stop Telegram ingress after draining active SDK handlers.""" + application = self._application + updater = application.updater if application is not None else None + try: + if updater is not None and updater.running: + await updater.stop() + finally: + try: + if application is not None and application.running: + await application.stop() + finally: + self._connected = False + + async def close(self) -> None: + """Close Telegram delivery and initialized SDK resources.""" + application = self._application + try: + await self.outbound.close() + finally: + try: + await self._limiter.shutdown() + finally: + try: + if application is not None: + await application.shutdown() + finally: + logger.info("Telegram platform closed") def on_message(self, handler: Callable[[IncomingMessage], Awaitable[None]]) -> None: """Register the workflow callback for inbound messages.""" diff --git a/src/free_claude_code/messaging/platforms/telegram_io.py b/src/free_claude_code/messaging/platforms/telegram_io.py index 8bccef70df8d447fd8646fbb63550be52ee13ac4..dff1fbbb16f6b97f6168183b37f77841f3038dd2 100644 --- a/src/free_claude_code/messaging/platforms/telegram_io.py +++ b/src/free_claude_code/messaging/platforms/telegram_io.py @@ -7,6 +7,7 @@ from typing import Any from loguru import logger +from ..limiter import MessagingRateLimiter from .outbox import PlatformOutbox TELEGRAM_DELETE_MESSAGES_BATCH_SIZE = 100 @@ -34,7 +35,6 @@ except ImportError: TelegramBaseError = Exception ApplicationGetter = Callable[[], Any | None] -LimiterGetter = Callable[[], Any | None] class TelegramMessenger: @@ -44,11 +44,11 @@ class TelegramMessenger: self, *, get_application: ApplicationGetter, - get_limiter: LimiterGetter, + limiter: MessagingRateLimiter, ) -> None: self._get_application = get_application self._outbox = PlatformOutbox( - get_limiter=get_limiter, + limiter=limiter, send=self.send_message, edit=self.edit_message, delete_many=self.delete_messages, @@ -281,3 +281,7 @@ class TelegramMessenger: def fire_and_forget(self, task: Awaitable[Any]) -> None: """Execute a coroutine without awaiting it.""" self._outbox.fire_and_forget(task) + + async def close(self) -> None: + """Cancel outstanding outbound work.""" + await self._outbox.close() diff --git a/src/free_claude_code/messaging/platforms/voice_flow.py b/src/free_claude_code/messaging/platforms/voice_flow.py index 8f077a2882e512ae734076de56546800a941b87d..31926130a3dad8b96eaf381d19d8c7de8227a367 100644 --- a/src/free_claude_code/messaging/platforms/voice_flow.py +++ b/src/free_claude_code/messaging/platforms/voice_flow.py @@ -1,5 +1,6 @@ """Shared voice-note flow for messaging platform adapters.""" +import asyncio import contextlib import tempfile from collections.abc import Awaitable, Callable @@ -12,9 +13,10 @@ from loguru import logger from free_claude_code.core.anthropic import format_user_error_preview from ..models import IncomingMessage -from ..voice import PendingVoiceRegistry, VoiceTranscriptionService +from ..voice import PendingVoiceRegistry, Transcriber AUDIO_EXTENSIONS = (".ogg", ".mp4", ".mp3", ".wav", ".m4a") +MAX_AUDIO_SIZE_BYTES = 25 * 1024 * 1024 VOICE_DISABLED_MESSAGE = "Voice notes are disabled." VOICE_TRANSCRIPTION_ERROR_MESSAGE = ( "Could not transcribe voice note. Please try again or send text." @@ -87,35 +89,25 @@ class VoiceNoteFlow: def __init__( self, *, - voice_note_enabled: bool, - whisper_model: str, - whisper_device: str, - huggingface_api_key: str, - nvidia_nim_api_key: str, + transcriber: Transcriber | None, log_raw_messaging_content: bool, log_api_error_tracebacks: bool, ) -> None: - self._voice_note_enabled = voice_note_enabled - self._whisper_model = whisper_model - self._whisper_device = whisper_device + self._transcriber = transcriber self._log_raw_messaging_content = log_raw_messaging_content self._log_api_error_tracebacks = log_api_error_tracebacks self._pending_voice = PendingVoiceRegistry() - self._voice_transcription = VoiceTranscriptionService( - huggingface_api_key=huggingface_api_key, - nvidia_nim_api_key=nvidia_nim_api_key, - ) @property def is_enabled(self) -> bool: """Return whether voice-note handling is enabled.""" - return self._voice_note_enabled + return self._transcriber is not None async def reply_if_disabled( self, reply_text: Callable[[str], Awaitable[None]] ) -> bool: """Reply with the disabled message when voice-note handling is disabled.""" - if self._voice_note_enabled: + if self.is_enabled: return False await reply_text(VOICE_DISABLED_MESSAGE) return True @@ -180,13 +172,12 @@ class VoiceNoteFlow: try: await request.download_to(tmp_path) + _validate_audio_file(tmp_path) - transcribed = await self._voice_transcription.transcribe( - tmp_path, - request.content_type, - whisper_model=self._whisper_model, - whisper_device=self._whisper_device, - ) + transcriber = self._transcriber + if transcriber is None: + raise RuntimeError("Voice transcription is not configured.") + transcribed = await transcriber.transcribe(tmp_path) if not await self.is_voice_still_pending( request.chat_id, @@ -218,6 +209,14 @@ class VoiceNoteFlow: self._log_transcription(request, transcribed) await message_handler(incoming) return True + except asyncio.CancelledError: + await self._clear_failed_pending_voice( + request, + status_msg_id_text, + queue_delete_messages, + handed_off=handed_off, + ) + raise except ValueError as e: await self._clear_failed_pending_voice( request, @@ -291,3 +290,13 @@ class VoiceNoteFlow: request.message_id, len(transcribed), ) + + +def _validate_audio_file(file_path: Path) -> None: + if not file_path.exists(): + raise FileNotFoundError(f"Audio file not found: {file_path}") + size = file_path.stat().st_size + if size > MAX_AUDIO_SIZE_BYTES: + raise ValueError( + f"Audio file too large ({size} bytes). Max {MAX_AUDIO_SIZE_BYTES} bytes." + ) diff --git a/src/free_claude_code/messaging/transcription.py b/src/free_claude_code/messaging/transcription.py index 48223e9eaed5be309eeff75128b3f7d901f7b82b..82eb42617ca50d6ab387753e0b54344172fab7e1 100644 --- a/src/free_claude_code/messaging/transcription.py +++ b/src/free_claude_code/messaging/transcription.py @@ -1,23 +1,11 @@ -"""Voice note transcription for messaging platforms. - -Supports: -- Local Whisper (cpu/cuda): Hugging Face transformers pipeline -- NVIDIA NIM: NVIDIA NIM Whisper/Parakeet -""" +"""Instance-owned local Whisper transcription.""" +import asyncio from pathlib import Path from typing import Any from loguru import logger -from free_claude_code.providers.nvidia_nim.voice import ( - transcribe_audio_file as transcribe_nvidia_nim_audio, -) - -# Max file size in bytes (25 MB) -MAX_AUDIO_SIZE_BYTES = 25 * 1024 * 1024 - -# Short model names -> full Hugging Face model IDs (for local Whisper) _MODEL_MAP: dict[str, str] = { "tiny": "openai/whisper-tiny", "base": "openai/whisper-base", @@ -27,143 +15,118 @@ _MODEL_MAP: dict[str, str] = { "large-v3": "openai/whisper-large-v3", "large-v3-turbo": "openai/whisper-large-v3-turbo", } +_WHISPER_SAMPLE_RATE = 16000 -# Lazy-loaded pipelines: (model_id, device, Hugging Face API key fingerprint) -> pipeline -_pipeline_cache: dict[tuple[str, str, str], Any] = {} - - -def _resolve_model_id(whisper_model: str) -> str: - """Resolve short name to full Hugging Face model ID.""" - return _MODEL_MAP.get(whisper_model, whisper_model) +class TranscriptionService: + """Own one lazily loaded local Whisper pipeline.""" -def _get_pipeline(model_id: str, device: str, huggingface_api_key: str = "") -> Any: - """Lazy-load transformers Whisper pipeline. Raises ImportError if not installed.""" - global _pipeline_cache - if device not in ("cpu", "cuda"): - raise ValueError(f"whisper_device must be 'cpu' or 'cuda', got {device!r}") - resolved_token = huggingface_api_key or "" - cache_key = (model_id, device, resolved_token) - if cache_key not in _pipeline_cache: + def __init__( + self, + *, + model: str, + device: str, + huggingface_api_key: str = "", + ) -> None: + if device not in {"cpu", "cuda"}: + raise ValueError( + f"Local Whisper device must be 'cpu' or 'cuda', got {device!r}" + ) + self._model_id = _MODEL_MAP.get(model, model) + self._device = device + self._huggingface_api_key = huggingface_api_key + self._pipeline: Any | None = None + self._lock = asyncio.Lock() + self._closed = False + + async def transcribe(self, file_path: Path) -> str: + """Transcribe one audio file without blocking the event loop.""" + async with self._lock: + if self._closed: + raise RuntimeError("Transcription service is closed.") + worker = asyncio.create_task( + asyncio.to_thread(self._transcribe_sync, file_path) + ) + try: + return await asyncio.shield(worker) + except asyncio.CancelledError: + await _wait_for_thread_exit(worker) + raise + + async def close(self) -> None: + """Prevent new work and release the owned model pipeline.""" + self._closed = True + async with self._lock: + self._pipeline = None + self._huggingface_api_key = "" + + def _transcribe_sync(self, file_path: Path) -> str: + pipe = self._get_pipeline() + audio = _load_audio(file_path) + result = pipe(audio, generate_kwargs={"language": "en", "task": "transcribe"}) + text = result.get("text", "") or "" + if isinstance(text, list): + text = " ".join(text) if text else "" + result_text = text.strip() + logger.debug("Local transcription: {} chars", len(result_text)) + return result_text or "(no speech detected)" + + def _get_pipeline(self) -> Any: + if self._pipeline is not None: + return self._pipeline try: import torch from transformers import AutoModelForSpeechSeq2Seq, AutoProcessor, pipeline - - hf_auth_token = resolved_token or None - - use_cuda = device == "cuda" and torch.cuda.is_available() - pipe_device = "cuda:0" if use_cuda else "cpu" - model_dtype = torch.float16 if use_cuda else torch.float32 - - model = AutoModelForSpeechSeq2Seq.from_pretrained( - model_id, - dtype=model_dtype, - low_cpu_mem_usage=True, - attn_implementation="sdpa", - token=hf_auth_token, - ) - model = model.to(pipe_device) - processor = AutoProcessor.from_pretrained(model_id, token=hf_auth_token) - - pipe = pipeline( - "automatic-speech-recognition", - model=model, - tokenizer=processor.tokenizer, - feature_extractor=processor.feature_extractor, - device=pipe_device, - ) - _pipeline_cache[cache_key] = pipe - logger.debug( - f"Loaded Whisper pipeline: model={model_id} device={pipe_device}" - ) - except ImportError as e: + except ImportError as exc: raise ImportError( - "Local Whisper requires the voice_local extra. Install with: uv sync --extra voice_local" - ) from e - return _pipeline_cache[cache_key] - - -def transcribe_audio( - file_path: Path, - mime_type: str, - *, - whisper_model: str = "base", - whisper_device: str = "cpu", - huggingface_api_key: str = "", - nvidia_nim_api_key: str = "", -) -> str: - """ - Transcribe audio file to text. - - Supports: - - whisper_device="cpu"/"cuda": local Whisper (requires voice_local extra) - - whisper_device="nvidia_nim": NVIDIA NIM Whisper API (requires voice extra) - - Args: - file_path: Path to audio file (OGG, MP3, MP4, WAV, M4A supported) - mime_type: MIME type of the audio (e.g. "audio/ogg") - whisper_model: Model ID or short name (local) or NVIDIA NIM model - whisper_device: "cpu" | "cuda" | "nvidia_nim" - - Returns: - Transcribed text - - Raises: - FileNotFoundError: If file does not exist - ValueError: If file too large - ImportError: If voice_local extra not installed (for local Whisper) - """ - - if not file_path.exists(): - raise FileNotFoundError(f"Audio file not found: {file_path}") - - size = file_path.stat().st_size - if size > MAX_AUDIO_SIZE_BYTES: - raise ValueError( - f"Audio file too large ({size} bytes). Max {MAX_AUDIO_SIZE_BYTES} bytes." + "Local Whisper requires the voice_local extra. " + "Install with: uv sync --extra voice_local" + ) from exc + + token = self._huggingface_api_key or None + use_cuda = self._device == "cuda" and torch.cuda.is_available() + pipeline_device = "cuda:0" if use_cuda else "cpu" + model_dtype = torch.float16 if use_cuda else torch.float32 + model = AutoModelForSpeechSeq2Seq.from_pretrained( + self._model_id, + dtype=model_dtype, + low_cpu_mem_usage=True, + attn_implementation="sdpa", + token=token, ) - - if whisper_device == "nvidia_nim": - return transcribe_nvidia_nim_audio( - file_path, whisper_model, api_key=nvidia_nim_api_key + model = model.to(pipeline_device) + processor = AutoProcessor.from_pretrained(self._model_id, token=token) + self._pipeline = pipeline( + "automatic-speech-recognition", + model=model, + tokenizer=processor.tokenizer, + feature_extractor=processor.feature_extractor, + device=pipeline_device, ) - return _transcribe_local( - file_path, - whisper_model, - whisper_device, - huggingface_api_key=huggingface_api_key, - ) + logger.debug( + "Loaded Whisper pipeline: model={} device={}", + self._model_id, + pipeline_device, + ) + return self._pipeline -# Whisper expects 16 kHz sample rate -_WHISPER_SAMPLE_RATE = 16000 +async def _wait_for_thread_exit(worker: asyncio.Task[str]) -> None: + """Wait through repeated caller cancellation without cancelling thread work.""" + while not worker.done(): + try: + await asyncio.shield(asyncio.wait((worker,))) + except asyncio.CancelledError: + continue + if not worker.cancelled(): + worker.exception() def _load_audio(file_path: Path) -> dict[str, Any]: - """Load audio file to waveform dict. No ffmpeg required.""" + """Load an audio file into the waveform shape expected by Whisper.""" import librosa - waveform, sr = librosa.load(str(file_path), sr=_WHISPER_SAMPLE_RATE, mono=True) - return {"array": waveform, "sampling_rate": sr} - - -def _transcribe_local( - file_path: Path, - whisper_model: str, - whisper_device: str, - *, - huggingface_api_key: str = "", -) -> str: - """Transcribe using transformers Whisper pipeline.""" - model_id = _resolve_model_id(whisper_model) - pipe = _get_pipeline( - model_id, whisper_device, huggingface_api_key=huggingface_api_key + waveform, sample_rate = librosa.load( + str(file_path), sr=_WHISPER_SAMPLE_RATE, mono=True ) - audio = _load_audio(file_path) - result = pipe(audio, generate_kwargs={"language": "en", "task": "transcribe"}) - text = result.get("text", "") or "" - if isinstance(text, list): - text = " ".join(text) if text else "" - result_text = text.strip() - logger.debug(f"Local transcription: {len(result_text)} chars") - return result_text or "(no speech detected)" + return {"array": waveform, "sampling_rate": sample_rate} diff --git a/src/free_claude_code/messaging/voice.py b/src/free_claude_code/messaging/voice.py index a1971f16581f5d512ede28529a28f766c06676b4..908c15887b07d0f8de6f7d2aa299b0e0e7773bc6 100644 --- a/src/free_claude_code/messaging/voice.py +++ b/src/free_claude_code/messaging/voice.py @@ -2,6 +2,15 @@ import asyncio from pathlib import Path +from typing import Protocol + + +class Transcriber(Protocol): + """Consumer-owned voice transcription boundary.""" + + async def transcribe(self, file_path: Path) -> str: ... + + async def close(self) -> None: ... class PendingVoiceRegistry: @@ -39,36 +48,3 @@ class PendingVoiceRegistry: async with self._lock: self._pending.pop((chat_id, voice_msg_id), None) self._pending.pop((chat_id, status_msg_id), None) - - -class VoiceTranscriptionService: - """Run configured transcription backends off the event loop.""" - - def __init__( - self, - *, - huggingface_api_key: str = "", - nvidia_nim_api_key: str = "", - ) -> None: - self._huggingface_api_key = huggingface_api_key - self._nvidia_nim_api_key = nvidia_nim_api_key - - async def transcribe( - self, - file_path: Path, - mime_type: str, - *, - whisper_model: str, - whisper_device: str, - ) -> str: - from .transcription import transcribe_audio - - return await asyncio.to_thread( - transcribe_audio, - file_path, - mime_type, - whisper_model=whisper_model, - whisper_device=whisper_device, - huggingface_api_key=self._huggingface_api_key, - nvidia_nim_api_key=self._nvidia_nim_api_key, - ) diff --git a/src/free_claude_code/providers/cerebras/client.py b/src/free_claude_code/providers/cerebras/client.py index 0433f001a4e28ad02d268e0c48f920d8d4316f9c..9fbbfdf2c34bfa1ec23fb4c908865b9bd147248e 100644 --- a/src/free_claude_code/providers/cerebras/client.py +++ b/src/free_claude_code/providers/cerebras/client.py @@ -4,6 +4,7 @@ from typing import Any from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import CEREBRAS_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -20,12 +21,13 @@ _REQUEST_POLICY = OpenAIChatRequestPolicy( class CerebrasProvider(OpenAIChatTransport): """Cerebras API at ``https://api.cerebras.ai/v1/chat/completions``.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="CEREBRAS", base_url=config.base_url or CEREBRAS_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/cloudflare/client.py b/src/free_claude_code/providers/cloudflare/client.py index f6d4dea2603bb3c61734940eee309c7ab31a2fd9..493880d75ee34547470257ef9bada5ec6fc0d8b8 100644 --- a/src/free_claude_code/providers/cloudflare/client.py +++ b/src/free_claude_code/providers/cloudflare/client.py @@ -17,6 +17,7 @@ from free_claude_code.providers.model_listing import ( extract_openai_model_ids, model_infos_from_ids, ) +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.http import maybe_await_aclose from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, @@ -59,7 +60,13 @@ def _cloudflare_account_api_url(api_root: str | None, account_id: str) -> str: class CloudflareProvider(OpenAIChatTransport): """Cloudflare Workers AI OpenAI-compatible chat provider.""" - def __init__(self, config: ProviderConfig, *, account_id: str): + def __init__( + self, + config: ProviderConfig, + *, + account_id: str, + rate_limiter: ProviderRateLimiter, + ): base_url = cloudflare_ai_base_url(config.base_url, account_id) self._model_search_url = _cloudflare_model_search_url( config.base_url, account_id @@ -78,6 +85,7 @@ class CloudflareProvider(OpenAIChatTransport): provider_name="CLOUDFLARE", base_url=base_url, api_key=config.api_key, + rate_limiter=rate_limiter, ) async def cleanup(self) -> None: diff --git a/src/free_claude_code/providers/codestral/client.py b/src/free_claude_code/providers/codestral/client.py index b802fcef74ac73bb6a0070a8320bcfc095740dde..3bbb1bc943b23756eaddf2c93762a090468cb0cc 100644 --- a/src/free_claude_code/providers/codestral/client.py +++ b/src/free_claude_code/providers/codestral/client.py @@ -4,6 +4,7 @@ from typing import Any from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import CODESTRAL_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -20,12 +21,13 @@ class CodestralProvider(OpenAIChatTransport): Request shaping matches Mistral La Plateforme. """ - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="CODESTRAL", base_url=config.base_url or CODESTRAL_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/cohere/client.py b/src/free_claude_code/providers/cohere/client.py index cc018b5d322afd4fcdd890c6744e59ca68cb0947..4ee70798f9576645bba9756bacc4c80d143e2882 100644 --- a/src/free_claude_code/providers/cohere/client.py +++ b/src/free_claude_code/providers/cohere/client.py @@ -7,6 +7,7 @@ from typing import Any from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import COHERE_DEFAULT_BASE from free_claude_code.providers.exceptions import InvalidRequestError +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -44,12 +45,13 @@ _REQUEST_POLICY = OpenAIChatRequestPolicy( class CohereProvider(OpenAIChatTransport): """Cohere Compatibility API at ``https://api.cohere.ai/compatibility/v1``.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="COHERE", base_url=config.base_url or COHERE_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/deepseek/client.py b/src/free_claude_code/providers/deepseek/client.py index 14a0e47ec0b46482fdd0ebfdbdbd0c3b0fd9d66c..ec70b8dc98a87c49767508bac95165d82f21184e 100644 --- a/src/free_claude_code/providers/deepseek/client.py +++ b/src/free_claude_code/providers/deepseek/client.py @@ -4,6 +4,7 @@ from typing import Any from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import DEEPSEEK_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import OpenAIChatTransport from free_claude_code.providers.transports.openai_chat.usage import usage_int @@ -13,12 +14,13 @@ from .compat import build_deepseek_request_body class DeepSeekProvider(OpenAIChatTransport): """DeepSeek using ``https://api.deepseek.com`` Chat Completions.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="DEEPSEEK", base_url=config.base_url or DEEPSEEK_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/error_mapping.py b/src/free_claude_code/providers/error_mapping.py index 49ce80dd04a1098b171bb2951c3a78bcbf621126..7cd33e91f3e21644e2f6a082051c784973bbfc50 100644 --- a/src/free_claude_code/providers/error_mapping.py +++ b/src/free_claude_code/providers/error_mapping.py @@ -24,7 +24,7 @@ from free_claude_code.providers.exceptions import ( ProviderError, RateLimitError, ) -from free_claude_code.providers.rate_limit import GlobalRateLimiter +from free_claude_code.providers.rate_limit import ProviderRateLimiter _BODY_ATTR = "_fcc_provider_error_body" _BODY_TRUNCATED_ATTR = "_fcc_provider_error_body_truncated" @@ -311,7 +311,7 @@ def map_stream_start_error( provider_name: str, read_timeout_s: float | None, request_id: str | None, - rate_limiter: GlobalRateLimiter | None = None, + rate_limiter: ProviderRateLimiter, ) -> ProviderError: """Map a final pre-start stream failure into an HTTP-serializable provider error. @@ -339,17 +339,14 @@ def map_stream_start_error( return APIError(message, status_code=502, raw_error=str(error)) -def map_error( - e: Exception, *, rate_limiter: GlobalRateLimiter | None = None -) -> Exception: +def map_error(e: Exception, *, rate_limiter: ProviderRateLimiter) -> Exception: """Map OpenAI or HTTPX exception to specific ProviderError. - Streaming transports should pass their scoped limiter (``self._global_rate_limiter``) - so reactive 429 handling applies to the correct provider. Tests may omit - ``rate_limiter`` to use the process-wide singleton. + Streaming transports pass their owned limiter so reactive 429 handling + applies only to that provider instance. """ message = get_user_facing_error_message(e) - limiter = rate_limiter or GlobalRateLimiter.get_instance() + limiter = rate_limiter if isinstance(e, openai.AuthenticationError): return AuthenticationError(message, raw_error=str(e)) diff --git a/src/free_claude_code/providers/fireworks/client.py b/src/free_claude_code/providers/fireworks/client.py index bac1f4d1f243515852937ebf4ce2022bd4b651e8..b5371096a5b0048f542a3aa43fc2b36c916b1632 100644 --- a/src/free_claude_code/providers/fireworks/client.py +++ b/src/free_claude_code/providers/fireworks/client.py @@ -5,6 +5,7 @@ from typing import Any from free_claude_code.config.constants import ANTHROPIC_DEFAULT_MAX_OUTPUT_TOKENS from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import FIREWORKS_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -27,12 +28,13 @@ _REQUEST_POLICY = OpenAIChatRequestPolicy( class FireworksProvider(OpenAIChatTransport): """Fireworks AI using ``https://api.fireworks.ai/inference/v1/chat/completions``.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="FIREWORKS", base_url=config.base_url or FIREWORKS_BASE_URL, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/gemini/client.py b/src/free_claude_code/providers/gemini/client.py index 136f81474b897091684d0dd8ed20a36f6d350b84..8f199ddcb8dc6e81867802b79447dbb80cda2f10 100644 --- a/src/free_claude_code/providers/gemini/client.py +++ b/src/free_claude_code/providers/gemini/client.py @@ -5,6 +5,7 @@ from typing import Any from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import GEMINI_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -20,12 +21,13 @@ _REQUEST_POLICY = OpenAIChatRequestPolicy(provider_name="GEMINI") class GeminiProvider(OpenAIChatTransport): """Gemini API using ``https://generativelanguage.googleapis.com/v1beta/openai/``.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="GEMINI", base_url=config.base_url or GEMINI_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) self._tool_call_extra_content_by_id: dict[str, dict[str, Any]] = {} diff --git a/src/free_claude_code/providers/github_models/client.py b/src/free_claude_code/providers/github_models/client.py index 9df1a858a5530421ec27364c5fe0f560ab34a4aa..d4bdeeff426b0fd3ffee716753ec671104eaa031 100644 --- a/src/free_claude_code/providers/github_models/client.py +++ b/src/free_claude_code/providers/github_models/client.py @@ -12,6 +12,7 @@ from free_claude_code.providers.model_listing import ( ProviderModelInfo, model_infos_from_ids, ) +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.http import maybe_await_aclose from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, @@ -31,7 +32,7 @@ _REQUIRED_MODEL_CAPABILITIES = frozenset({"streaming", "tool-calling"}) class GitHubModelsProvider(OpenAIChatTransport): """GitHub Models OpenAI-compatible inference provider.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): self._catalog_url = GITHUB_MODELS_CATALOG_URL self._model_list_client = httpx.AsyncClient( proxy=config.proxy or None, @@ -47,6 +48,7 @@ class GitHubModelsProvider(OpenAIChatTransport): provider_name="GITHUB_MODELS", base_url=config.base_url or GITHUB_MODELS_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, default_headers=_github_models_default_headers(), ) diff --git a/src/free_claude_code/providers/groq/client.py b/src/free_claude_code/providers/groq/client.py index 1129f4d5fc5e5b348403fbcafe9145da4bd82370..630f6f93038a4a99152f751270cc4f61b6d8b6f2 100644 --- a/src/free_claude_code/providers/groq/client.py +++ b/src/free_claude_code/providers/groq/client.py @@ -4,6 +4,7 @@ from typing import Any from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import GROQ_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -23,12 +24,13 @@ _REQUEST_POLICY = OpenAIChatRequestPolicy( class GroqProvider(OpenAIChatTransport): """Groq API using ``https://api.groq.com/openai/v1/chat/completions``.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="GROQ", base_url=config.base_url or GROQ_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/huggingface/client.py b/src/free_claude_code/providers/huggingface/client.py index 701a46a0559ba089bcc0cd092a17f5e573e6c647..7fff464056a2ed16e01f0d54578d63c059d2049e 100644 --- a/src/free_claude_code/providers/huggingface/client.py +++ b/src/free_claude_code/providers/huggingface/client.py @@ -8,18 +8,20 @@ from free_claude_code.core.anthropic.conversion import OpenAIConversionError from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import HUGGINGFACE_DEFAULT_BASE from free_claude_code.providers.exceptions import InvalidRequestError +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import OpenAIChatTransport class HuggingFaceProvider(OpenAIChatTransport): """Hugging Face Inference Providers router at ``https://router.huggingface.co/v1``.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="HUGGINGFACE", base_url=config.base_url or HUGGINGFACE_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/kimi/client.py b/src/free_claude_code/providers/kimi/client.py index 2b486df01956ce6083d24e994073a291659a9d7e..530ab4254232a68f1173be006a297d209e923f25 100644 --- a/src/free_claude_code/providers/kimi/client.py +++ b/src/free_claude_code/providers/kimi/client.py @@ -5,6 +5,7 @@ from typing import Any from free_claude_code.config.constants import ANTHROPIC_DEFAULT_MAX_OUTPUT_TOKENS from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import KIMI_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -23,12 +24,13 @@ _REQUEST_POLICY = OpenAIChatRequestPolicy( class KimiProvider(OpenAIChatTransport): """Kimi provider using ``https://api.moonshot.ai/v1/chat/completions``.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="KIMI", base_url=config.base_url or KIMI_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/llamacpp/client.py b/src/free_claude_code/providers/llamacpp/client.py index 3cd0032b735a355196a1b8788760372ae2e6d703..cc8202af8c2875844d58a4eb280cff5d47c12520 100644 --- a/src/free_claude_code/providers/llamacpp/client.py +++ b/src/free_claude_code/providers/llamacpp/client.py @@ -2,6 +2,7 @@ from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import LLAMACPP_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.anthropic_messages import ( AnthropicMessagesTransport, ) @@ -10,9 +11,10 @@ from free_claude_code.providers.transports.anthropic_messages import ( class LlamaCppProvider(AnthropicMessagesTransport): """Llama.cpp provider using native Anthropic Messages endpoint.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="LLAMACPP", default_base_url=LLAMACPP_DEFAULT_BASE, + rate_limiter=rate_limiter, ) diff --git a/src/free_claude_code/providers/lmstudio/client.py b/src/free_claude_code/providers/lmstudio/client.py index 5b045a460603405189e65695f479f7b712101eed..6240cf30a8bb9f73e396dff61995b1d8bcd9de73 100644 --- a/src/free_claude_code/providers/lmstudio/client.py +++ b/src/free_claude_code/providers/lmstudio/client.py @@ -26,6 +26,7 @@ from free_claude_code.core.anthropic.conversion import OpenAIConversionError from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import LMSTUDIO_DEFAULT_BASE from free_claude_code.providers.exceptions import InvalidRequestError +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import OpenAIChatTransport @@ -39,12 +40,13 @@ class LMStudioProvider(OpenAIChatTransport): # dying mid-stream. _CONTEXT_CACHE_TTL_S = 30.0 - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="LMSTUDIO", base_url=config.base_url or LMSTUDIO_DEFAULT_BASE, api_key=config.api_key or "lm-studio", + rate_limiter=rate_limiter, ) self._loaded_context_cache: tuple[float, int | None] = (0.0, None) diff --git a/src/free_claude_code/providers/minimax/client.py b/src/free_claude_code/providers/minimax/client.py index 357d38124646f5bffd51388ed679a39d40b1fe70..d59153bff981bc785703494cbd23e222513c7f9b 100644 --- a/src/free_claude_code/providers/minimax/client.py +++ b/src/free_claude_code/providers/minimax/client.py @@ -5,6 +5,7 @@ from typing import Any from free_claude_code.config.constants import ANTHROPIC_DEFAULT_MAX_OUTPUT_TOKENS from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import MINIMAX_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -21,12 +22,13 @@ _REQUEST_POLICY = OpenAIChatRequestPolicy( class MiniMaxProvider(OpenAIChatTransport): """MiniMax using ``https://api.minimax.io/v1/chat/completions``.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="MINIMAX", base_url=config.base_url or MINIMAX_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/mistral/client.py b/src/free_claude_code/providers/mistral/client.py index 2654386223379da8ca6bdf728934a6b94aeb32ee..157313297691ed8093e47be51eaf9a0f1415a2c9 100644 --- a/src/free_claude_code/providers/mistral/client.py +++ b/src/free_claude_code/providers/mistral/client.py @@ -6,6 +6,7 @@ from loguru import logger from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import MISTRAL_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -25,12 +26,13 @@ _REQUEST_POLICY = OpenAIChatRequestPolicy(provider_name="MISTRAL") class MistralProvider(OpenAIChatTransport): """Mistral API using ``https://api.mistral.ai/v1/chat/completions``.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="MISTRAL", base_url=config.base_url or MISTRAL_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/nvidia_nim/client.py b/src/free_claude_code/providers/nvidia_nim/client.py index 6bb970724d53d3887f2d7c6c6a3c8b524746c232..61cd6db4b0512127102aa0c4b006137cc9e67858 100644 --- a/src/free_claude_code/providers/nvidia_nim/client.py +++ b/src/free_claude_code/providers/nvidia_nim/client.py @@ -9,6 +9,7 @@ from loguru import logger from free_claude_code.config.nim import NimSettings from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import NVIDIA_NIM_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import OpenAIChatTransport from .request_options import build_nim_request_body @@ -26,12 +27,19 @@ from .tool_schema import ( class NvidiaNimProvider(OpenAIChatTransport): """NVIDIA NIM provider using official OpenAI client.""" - def __init__(self, config: ProviderConfig, *, nim_settings: NimSettings): + def __init__( + self, + config: ProviderConfig, + *, + nim_settings: NimSettings, + rate_limiter: ProviderRateLimiter, + ): super().__init__( config, provider_name="NIM", base_url=config.base_url or NVIDIA_NIM_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) self._nim_settings = nim_settings diff --git a/src/free_claude_code/providers/nvidia_nim/voice.py b/src/free_claude_code/providers/nvidia_nim/voice.py index bc46da3c0d22ef3528776155e2a2e73ed75ba57d..89496f7561a20e1e680b0700b770a7910eadebad 100644 --- a/src/free_claude_code/providers/nvidia_nim/voice.py +++ b/src/free_claude_code/providers/nvidia_nim/voice.py @@ -1,5 +1,6 @@ """NVIDIA NIM / Riva offline ASR for voice notes (provider-owned transport).""" +import asyncio from pathlib import Path from loguru import logger @@ -23,71 +24,90 @@ _NIM_ASR_MODEL_MAP: dict[str, tuple[str, str]] = { _RIVA_SERVER = "grpc.nvcf.nvidia.com:443" -def transcribe_audio_file( - file_path: Path, - model: str, - *, - api_key: str, -) -> str: - """Transcribe audio using NVIDIA NIM / Riva gRPC (offline recognition). - - Args: - file_path: Path to encoded audio bytes readable by Riva. - model: Hugging Face-style NIM model id (see ``_NIM_ASR_MODEL_MAP``). - api_key: NVIDIA API key (Bearer token); must be non-empty. - - Returns: - Transcript text, or ``(no speech detected)`` when empty. - """ - key = (api_key or "").strip() - if not key: - raise ValueError( - "NVIDIA NIM transcription requires a non-empty nvidia_nim_api_key " - "(configure NVIDIA_NIM_API_KEY or pass api_key explicitly)." - ) - - try: - import riva.client - except ImportError as e: - raise ImportError( - "NVIDIA NIM transcription requires the voice extra. " - "Install with: uv sync --extra voice" - ) from e - - model_config = _NIM_ASR_MODEL_MAP.get(model) - if not model_config: - raise ValueError( - f"No NVIDIA NIM config found for model: {model}. " - f"Supported models: {', '.join(_NIM_ASR_MODEL_MAP.keys())}" +class NvidiaNimTranscriber: + """Own configured NVIDIA NIM / Riva transcription.""" + + def __init__(self, *, model: str, api_key: str) -> None: + self._model = model + self._key = api_key.strip() + self._lock = asyncio.Lock() + self._closed = False + + async def transcribe(self, file_path: Path) -> str: + """Transcribe one audio file without blocking the event loop.""" + async with self._lock: + if self._closed: + raise RuntimeError("NVIDIA NIM transcriber is closed.") + worker = asyncio.create_task( + asyncio.to_thread(self._transcribe_sync, file_path) + ) + try: + return await asyncio.shield(worker) + except asyncio.CancelledError: + await _wait_for_thread_exit(worker) + raise + + async def close(self) -> None: + """Close this stateless adapter to future work.""" + self._closed = True + async with self._lock: + self._key = "" + + def _transcribe_sync(self, file_path: Path) -> str: + if not self._key: + raise ValueError( + "NVIDIA NIM transcription requires a non-empty " + "nvidia_nim_api_key (configure NVIDIA_NIM_API_KEY)." + ) + model_config = _NIM_ASR_MODEL_MAP.get(self._model) + if model_config is None: + raise ValueError( + f"No NVIDIA NIM config found for model: {self._model}. " + f"Supported models: {', '.join(_NIM_ASR_MODEL_MAP)}" + ) + function_id, language_code = model_config + try: + import riva.client + except ImportError as exc: + raise ImportError( + "NVIDIA NIM transcription requires the voice extra. " + "Install with: uv sync --extra voice" + ) from exc + + auth = riva.client.Auth( + use_ssl=True, + uri=_RIVA_SERVER, + metadata_args=[ + ["function-id", function_id], + ["authorization", f"Bearer {self._key}"], + ], ) - function_id, language_code = model_config - - auth = riva.client.Auth( - use_ssl=True, - uri=_RIVA_SERVER, - metadata_args=[ - ["function-id", function_id], - ["authorization", f"Bearer {key}"], - ], - ) - - asr_service = riva.client.ASRService(auth) - - config = riva.client.RecognitionConfig( - language_code=language_code, - max_alternatives=1, - verbatim_transcripts=True, - ) - - with open(file_path, "rb") as f: - data = f.read() - - response = asr_service.offline_recognize(data, config) - - transcript = "" - results = getattr(response, "results", None) - if results and results[0].alternatives: - transcript = results[0].alternatives[0].transcript - - logger.debug(f"NIM transcription: {len(transcript)} chars") - return transcript or "(no speech detected)" + try: + asr_service = riva.client.ASRService(auth) + config = riva.client.RecognitionConfig( + language_code=language_code, + max_alternatives=1, + verbatim_transcripts=True, + ) + data = file_path.read_bytes() + response = asr_service.offline_recognize(data, config) + + transcript = "" + results = getattr(response, "results", None) + if results and results[0].alternatives: + transcript = results[0].alternatives[0].transcript + logger.debug("NIM transcription: {} chars", len(transcript)) + return transcript or "(no speech detected)" + finally: + auth.channel.close() + + +async def _wait_for_thread_exit(worker: asyncio.Task[str]) -> None: + """Wait through repeated caller cancellation without cancelling thread work.""" + while not worker.done(): + try: + await asyncio.shield(asyncio.wait((worker,))) + except asyncio.CancelledError: + continue + if not worker.cancelled(): + worker.exception() diff --git a/src/free_claude_code/providers/ollama/client.py b/src/free_claude_code/providers/ollama/client.py index 6266565a2040fffdb8c312f33b1e6b5987bf0be9..8f9dc7bd6a5f089139aa5abb7ca04fd455bad9bc 100644 --- a/src/free_claude_code/providers/ollama/client.py +++ b/src/free_claude_code/providers/ollama/client.py @@ -5,6 +5,7 @@ import httpx from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import OLLAMA_DEFAULT_BASE from free_claude_code.providers.model_listing import extract_ollama_model_ids +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.anthropic_messages import ( AnthropicMessagesTransport, ) @@ -13,11 +14,12 @@ from free_claude_code.providers.transports.anthropic_messages import ( class OllamaProvider(AnthropicMessagesTransport): """Ollama provider using native Anthropic Messages API.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="OLLAMA", default_base_url=OLLAMA_DEFAULT_BASE, + rate_limiter=rate_limiter, ) self._api_key = config.api_key or "ollama" diff --git a/src/free_claude_code/providers/open_router/client.py b/src/free_claude_code/providers/open_router/client.py index f03e8b6d040d486aead1d689eef97410cc5852e0..b2bd5ddc31d53c40697a8b5bc8408d7848c14b25 100644 --- a/src/free_claude_code/providers/open_router/client.py +++ b/src/free_claude_code/providers/open_router/client.py @@ -13,6 +13,7 @@ from free_claude_code.providers.model_listing import ( extract_openrouter_tool_model_ids, extract_openrouter_tool_model_infos, ) +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -33,12 +34,13 @@ _REQUEST_POLICY = OpenAIChatRequestPolicy( class OpenRouterProvider(OpenAIChatTransport): """OpenRouter provider using the OpenAI-compatible Chat Completions API.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="OPENROUTER", base_url=config.base_url or OPENROUTER_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/opencode/client.py b/src/free_claude_code/providers/opencode/client.py index 2f6172feb479fc018c90b6ffbfb57d9428a37436..aa7c6d1ddc6bb9c3e0505e40976b506f744a21d3 100644 --- a/src/free_claude_code/providers/opencode/client.py +++ b/src/free_claude_code/providers/opencode/client.py @@ -4,6 +4,7 @@ from typing import Any from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import OPENCODE_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -14,12 +15,19 @@ from free_claude_code.providers.transports.openai_chat import ( class OpenCodeProvider(OpenAIChatTransport): """OpenCode Zen provider using ``https://opencode.ai/zen/v1/chat/completions``.""" - def __init__(self, config: ProviderConfig, provider_name: str = "OPENCODE"): + def __init__( + self, + config: ProviderConfig, + provider_name: str = "OPENCODE", + *, + rate_limiter: ProviderRateLimiter, + ): super().__init__( config, provider_name=provider_name, base_url=config.base_url or OPENCODE_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) self._request_policy = OpenAIChatRequestPolicy(provider_name=provider_name) diff --git a/src/free_claude_code/providers/rate_limit.py b/src/free_claude_code/providers/rate_limit.py index c99d583efd5d453419639a6b08db053c07762506..51e0e55393eb53edb7efbe5da1cbfd8bee033baf 100644 --- a/src/free_claude_code/providers/rate_limit.py +++ b/src/free_claude_code/providers/rate_limit.py @@ -1,11 +1,11 @@ -"""Global rate limiter for API requests.""" +"""Provider-owned upstream rate limiting and retry policy.""" import asyncio import random import time from collections.abc import AsyncIterator, Callable from contextlib import asynccontextmanager -from typing import Any, ClassVar, TypeVar +from typing import Any, TypeVar import httpx import openai @@ -58,11 +58,13 @@ def retryable_upstream_transport_error(exc: BaseException) -> bool: ) -class GlobalRateLimiter: +class ProviderRateLimiter: """ - Global singleton rate limiter that blocks all requests - when a rate limit error is encountered (reactive) and - throttles requests (proactive) using a strict rolling window. + Rate limiter owned by one provider instance. + + Blocks that provider's requests when a rate-limit error is encountered + (reactive) and throttles its requests with a strict rolling window + (proactive). Optionally enforces a max_concurrency cap: at most N provider streams may be open simultaneously, independent of the sliding window. @@ -72,19 +74,12 @@ class GlobalRateLimiter: Concurrency limit - caps simultaneously open streams. """ - _instance: ClassVar[GlobalRateLimiter | None] = None - _scoped_instances: ClassVar[dict[str, GlobalRateLimiter]] = {} - def __init__( self, rate_limit: int = 40, rate_window: float = 60.0, max_concurrency: int = 5, ): - # Prevent re-initialization on singleton reuse - if hasattr(self, "_initialized"): - return - if rate_limit <= 0: raise ValueError("rate_limit must be > 0") if rate_window <= 0: @@ -100,69 +95,10 @@ class GlobalRateLimiter: ) self._blocked_until: float = 0 self._concurrency_sem = asyncio.Semaphore(max_concurrency) - self._initialized = True - logger.info( - f"GlobalRateLimiter (Provider) initialized ({rate_limit} req / {rate_window}s, max_concurrency={max_concurrency})" - ) - - @classmethod - def get_instance( - cls, - rate_limit: int | None = None, - rate_window: float | None = None, - max_concurrency: int = 5, - ) -> GlobalRateLimiter: - """Get or create the singleton instance. - - Args: - rate_limit: Requests per window (only used on first creation) - rate_window: Window in seconds (only used on first creation) - max_concurrency: Max simultaneous open streams (only used on first creation) - """ - if cls._instance is None: - cls._instance = cls( - rate_limit=rate_limit or 40, - rate_window=rate_window or 60.0, - max_concurrency=max_concurrency, - ) - return cls._instance - - @classmethod - def get_scoped_instance( - cls, - scope: str, - *, - rate_limit: int | None = None, - rate_window: float | None = None, - max_concurrency: int = 5, - ) -> GlobalRateLimiter: - """Get or create a provider-scoped limiter instance.""" - if not scope: - raise ValueError("scope must be non-empty") - desired_rate_limit = rate_limit or 40 - desired_rate_window = float(rate_window or 60.0) - existing = cls._scoped_instances.get(scope) - if existing and existing.matches_config( - desired_rate_limit, desired_rate_window, max_concurrency - ): - return existing - if existing: - logger.info( - "Rebuilding provider rate limiter for updated scope '{}'", scope - ) - cls._scoped_instances[scope] = cls( - rate_limit=desired_rate_limit, - rate_window=desired_rate_window, - max_concurrency=max_concurrency, + "ProviderRateLimiter initialized " + f"({rate_limit} req / {rate_window}s, max_concurrency={max_concurrency})" ) - return cls._scoped_instances[scope] - - @classmethod - def reset_instance(cls) -> None: - """Reset singleton (for testing).""" - cls._instance = None - cls._scoped_instances = {} async def wait_if_blocked(self) -> bool: """ @@ -177,7 +113,7 @@ class GlobalRateLimiter: if now < self._blocked_until: wait_time = self._blocked_until - now logger.warning( - f"Global provider rate limit active (reactive), waiting {wait_time:.1f}s..." + f"Provider rate limit active (reactive), waiting {wait_time:.1f}s..." ) await asyncio.sleep(wait_time) waited_reactively = True @@ -197,28 +133,18 @@ class GlobalRateLimiter: def set_blocked(self, seconds: float = 60) -> None: """ - Set global block for specified seconds (reactive). + Set this provider's block for the specified seconds (reactive). Args: seconds: How long to block (default 60s) """ self._blocked_until = time.monotonic() + seconds - logger.warning(f"Global provider rate limit set for {seconds:.1f}s (reactive)") + logger.warning(f"Provider rate limit set for {seconds:.1f}s (reactive)") def is_blocked(self) -> bool: """Check if currently reactively blocked.""" return time.monotonic() < self._blocked_until - def matches_config( - self, rate_limit: int, rate_window: float, max_concurrency: int - ) -> bool: - """Return whether this limiter matches the requested runtime config.""" - return ( - self._rate_limit == rate_limit - and self._rate_window == float(rate_window) - and self._max_concurrency == max_concurrency - ) - def remaining_wait(self) -> float: """Get remaining reactive wait time in seconds.""" return max(0.0, self._blocked_until - time.monotonic()) diff --git a/src/free_claude_code/providers/runtime/cache.py b/src/free_claude_code/providers/runtime/cache.py index 7492711c223a827906005c1c875c67f1780e4c85..413f06e0effbaf7fe55e5c94566a1aaf099d01fc 100644 --- a/src/free_claude_code/providers/runtime/cache.py +++ b/src/free_claude_code/providers/runtime/cache.py @@ -1,5 +1,6 @@ """Provider instance cache and cleanup.""" +import asyncio from collections.abc import Callable, MutableMapping from free_claude_code.config.settings import Settings @@ -35,17 +36,18 @@ class ProviderCache: return self._providers[provider_id] async def cleanup(self) -> None: - """Clean up every cached provider, then clear the cache.""" + """Clean every cached provider, retaining unfinished entries for retry.""" items = list(self._providers.items()) errors: list[Exception] = [] - try: - for _provider_id, provider in items: - try: - await provider.cleanup() - except Exception as exc: - errors.append(exc) - finally: - self._providers.clear() + for provider_id, provider in items: + try: + await provider.cleanup() + except asyncio.CancelledError: + raise + except Exception as exc: + errors.append(exc) + else: + self._providers.pop(provider_id, None) if len(errors) == 1: raise errors[0] if len(errors) > 1: diff --git a/src/free_claude_code/providers/runtime/factory.py b/src/free_claude_code/providers/runtime/factory.py index e91f8b9c2772ad0259250da004ecf83aaadca449..50f829122a2105fc60699c300fcd242736e90d9d 100644 --- a/src/free_claude_code/providers/runtime/factory.py +++ b/src/free_claude_code/providers/runtime/factory.py @@ -9,156 +9,265 @@ from free_claude_code.config.provider_catalog import ( from free_claude_code.config.settings import Settings from free_claude_code.providers.base import BaseProvider, ProviderConfig from free_claude_code.providers.exceptions import UnknownProviderTypeError +from free_claude_code.providers.rate_limit import ProviderRateLimiter from .config import build_provider_config -ProviderFactory = Callable[[ProviderConfig, Settings], BaseProvider] +ProviderFactory = Callable[ + [ProviderConfig, Settings, ProviderRateLimiter], BaseProvider +] -def _create_nvidia_nim(config: ProviderConfig, settings: Settings) -> BaseProvider: +def _create_nvidia_nim( + config: ProviderConfig, + settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.nvidia_nim import NvidiaNimProvider - return NvidiaNimProvider(config, nim_settings=settings.nim) + return NvidiaNimProvider( + config, + nim_settings=settings.nim, + rate_limiter=rate_limiter, + ) -def _create_open_router(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_open_router( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.open_router import OpenRouterProvider - return OpenRouterProvider(config) + return OpenRouterProvider(config, rate_limiter=rate_limiter) -def _create_mistral(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_mistral( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.mistral import MistralProvider - return MistralProvider(config) + return MistralProvider(config, rate_limiter=rate_limiter) def _create_mistral_codestral( - config: ProviderConfig, _settings: Settings + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, ) -> BaseProvider: from free_claude_code.providers.codestral import CodestralProvider - return CodestralProvider(config) + return CodestralProvider(config, rate_limiter=rate_limiter) -def _create_deepseek(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_deepseek( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.deepseek import DeepSeekProvider - return DeepSeekProvider(config) + return DeepSeekProvider(config, rate_limiter=rate_limiter) -def _create_lmstudio(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_lmstudio( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.lmstudio import LMStudioProvider - return LMStudioProvider(config) + return LMStudioProvider(config, rate_limiter=rate_limiter) -def _create_llamacpp(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_llamacpp( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.llamacpp import LlamaCppProvider - return LlamaCppProvider(config) + return LlamaCppProvider(config, rate_limiter=rate_limiter) -def _create_ollama(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_ollama( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.ollama import OllamaProvider - return OllamaProvider(config) + return OllamaProvider(config, rate_limiter=rate_limiter) -def _create_kimi(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_kimi( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.kimi import KimiProvider - return KimiProvider(config) + return KimiProvider(config, rate_limiter=rate_limiter) -def _create_wafer(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_wafer( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.wafer import WaferProvider - return WaferProvider(config) + return WaferProvider(config, rate_limiter=rate_limiter) -def _create_minimax(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_minimax( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.minimax import MiniMaxProvider - return MiniMaxProvider(config) + return MiniMaxProvider(config, rate_limiter=rate_limiter) -def _create_opencode(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_opencode( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.opencode import OpenCodeProvider - return OpenCodeProvider(config) + return OpenCodeProvider(config, rate_limiter=rate_limiter) -def _create_opencode_go(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_opencode_go( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.opencode import OpenCodeProvider - return OpenCodeProvider(config, provider_name="OPENCODE_GO") + return OpenCodeProvider( + config, + provider_name="OPENCODE_GO", + rate_limiter=rate_limiter, + ) -def _create_vercel(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_vercel( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.vercel import VercelProvider - return VercelProvider(config) + return VercelProvider(config, rate_limiter=rate_limiter) -def _create_huggingface(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_huggingface( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.huggingface import HuggingFaceProvider - return HuggingFaceProvider(config) + return HuggingFaceProvider(config, rate_limiter=rate_limiter) -def _create_cohere(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_cohere( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.cohere import CohereProvider - return CohereProvider(config) + return CohereProvider(config, rate_limiter=rate_limiter) -def _create_github_models(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_github_models( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.github_models import GitHubModelsProvider - return GitHubModelsProvider(config) + return GitHubModelsProvider(config, rate_limiter=rate_limiter) -def _create_zai(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_zai( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.zai import ZaiProvider - return ZaiProvider(config) + return ZaiProvider(config, rate_limiter=rate_limiter) -def _create_fireworks(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_fireworks( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.fireworks import FireworksProvider - return FireworksProvider(config) + return FireworksProvider(config, rate_limiter=rate_limiter) -def _create_cloudflare(config: ProviderConfig, settings: Settings) -> BaseProvider: +def _create_cloudflare( + config: ProviderConfig, + settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.cloudflare import CloudflareProvider - return CloudflareProvider(config, account_id=settings.cloudflare_account_id) + return CloudflareProvider( + config, + account_id=settings.cloudflare_account_id, + rate_limiter=rate_limiter, + ) -def _create_gemini(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_gemini( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.gemini import GeminiProvider - return GeminiProvider(config) + return GeminiProvider(config, rate_limiter=rate_limiter) -def _create_groq(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_groq( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.groq import GroqProvider - return GroqProvider(config) + return GroqProvider(config, rate_limiter=rate_limiter) -def _create_sambanova(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_sambanova( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.sambanova import SambaNovaProvider - return SambaNovaProvider(config) + return SambaNovaProvider(config, rate_limiter=rate_limiter) -def _create_cerebras(config: ProviderConfig, _settings: Settings) -> BaseProvider: +def _create_cerebras( + config: ProviderConfig, + _settings: Settings, + rate_limiter: ProviderRateLimiter, +) -> BaseProvider: from free_claude_code.providers.cerebras import CerebrasProvider - return CerebrasProvider(config) + return CerebrasProvider(config, rate_limiter=rate_limiter) PROVIDER_FACTORIES: dict[str, ProviderFactory] = { @@ -210,4 +319,10 @@ def create_provider(provider_id: str, settings: Settings) -> BaseProvider: factory = PROVIDER_FACTORIES.get(provider_id) if factory is None: raise AssertionError(f"Unhandled provider descriptor: {provider_id}") - return factory(build_provider_config(descriptor, settings), settings) + config = build_provider_config(descriptor, settings) + rate_limiter = ProviderRateLimiter( + rate_limit=config.rate_limit or 40, + rate_window=config.rate_window or 60.0, + max_concurrency=config.max_concurrency, + ) + return factory(config, settings, rate_limiter) diff --git a/src/free_claude_code/providers/sambanova/client.py b/src/free_claude_code/providers/sambanova/client.py index 9605d4eeb41aa8cb71dee4ada8fe3a8eb419b546..094b0c6248404aa884294b08e8b57f93a01e4a74 100644 --- a/src/free_claude_code/providers/sambanova/client.py +++ b/src/free_claude_code/providers/sambanova/client.py @@ -4,6 +4,7 @@ from typing import Any from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import SAMBANOVA_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -19,12 +20,13 @@ _REQUEST_POLICY = OpenAIChatRequestPolicy( class SambaNovaProvider(OpenAIChatTransport): """SambaNova Cloud API at ``https://api.sambanova.ai/v1``.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="SAMBANOVA", base_url=config.base_url or SAMBANOVA_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/transports/anthropic_messages/recovery.py b/src/free_claude_code/providers/transports/anthropic_messages/recovery.py index fb632260f5d405f74c6af60a0bceaebad90e0c68..3e974d7038bf35d8753a65f30465fc3660687b0f 100644 --- a/src/free_claude_code/providers/transports/anthropic_messages/recovery.py +++ b/src/free_claude_code/providers/transports/anthropic_messages/recovery.py @@ -48,10 +48,8 @@ class AnthropicMessagesRecovery: for attempt in range(MIDSTREAM_RECOVERY_ATTEMPTS): response: httpx.Response | None = None try: - response = ( - await self._transport._global_rate_limiter.execute_with_retry( - self._transport._validated_stream_send, body, req_tag=req_tag - ) + response = await self._transport._rate_limiter.execute_with_retry( + self._transport._validated_stream_send, body, req_tag=req_tag ) state = self._transport._new_stream_state( None, thinking_enabled=thinking_enabled diff --git a/src/free_claude_code/providers/transports/anthropic_messages/stream.py b/src/free_claude_code/providers/transports/anthropic_messages/stream.py index 71a2c8bab42346527fb49303be42e5e0dbcfed15..2d9fa89d4a987deb85a4a653ff260292da68bd62 100644 --- a/src/free_claude_code/providers/transports/anthropic_messages/stream.py +++ b/src/free_claude_code/providers/transports/anthropic_messages/stream.py @@ -91,16 +91,14 @@ class AnthropicMessagesStreamAdapter: ledger = self._new_ledger() recovery = RecoveryController(provider_name=tag, request_id=self._request_id) - async with self._transport._global_rate_limiter.concurrency_slot(): + async with self._transport._rate_limiter.concurrency_slot(): while True: stream_opened = False try: - response = ( - await self._transport._global_rate_limiter.execute_with_retry( - self._transport._validated_stream_send, - body, - req_tag=req_tag, - ) + response = await self._transport._rate_limiter.execute_with_retry( + self._transport._validated_stream_send, + body, + req_tag=req_tag, ) stream_opened = True chunk_count = 0 @@ -261,7 +259,7 @@ class AnthropicMessagesStreamAdapter: provider_name=tag, read_timeout_s=self._transport._config.http_read_timeout, request_id=self._request_id, - rate_limiter=self._transport._global_rate_limiter, + rate_limiter=self._transport._rate_limiter, ) from error return finally: diff --git a/src/free_claude_code/providers/transports/anthropic_messages/transport.py b/src/free_claude_code/providers/transports/anthropic_messages/transport.py index 2a9466eb1bcd4973a6e3d48a5457f4a735b44433..6c74338f28cf2d29f0425c126380328d40a8e41e 100644 --- a/src/free_claude_code/providers/transports/anthropic_messages/transport.py +++ b/src/free_claude_code/providers/transports/anthropic_messages/transport.py @@ -20,7 +20,7 @@ from free_claude_code.providers.model_listing import ( extract_openai_model_ids, model_infos_from_ids, ) -from free_claude_code.providers.rate_limit import GlobalRateLimiter +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.http import maybe_await_aclose from .http import model_list_json, raise_for_status_with_body @@ -44,18 +44,14 @@ class AnthropicMessagesTransport(BaseProvider): *, provider_name: str, default_base_url: str, + rate_limiter: ProviderRateLimiter, ): super().__init__(config) self._provider_name = provider_name self._api_key = config.api_key self._base_url = (config.base_url or default_base_url).rstrip("/") self._request_policy = NativeMessagesRequestPolicy(provider_name=provider_name) - self._global_rate_limiter = GlobalRateLimiter.get_scoped_instance( - provider_name.lower(), - rate_limit=config.rate_limit, - rate_window=config.rate_window, - max_concurrency=config.max_concurrency, - ) + self._rate_limiter = rate_limiter self._client = httpx.AsyncClient( base_url=self._base_url, proxy=config.proxy or None, @@ -182,7 +178,7 @@ class AnthropicMessagesTransport(BaseProvider): self, error: Exception, request_id: str | None ) -> tuple[Exception, str]: """Map an exception into a user-facing provider error message.""" - mapped_error = map_error(error, rate_limiter=self._global_rate_limiter) + mapped_error = map_error(error, rate_limiter=self._rate_limiter) return ( mapped_error, user_visible_message_for_mapped_provider_error( diff --git a/src/free_claude_code/providers/transports/openai_chat/stream.py b/src/free_claude_code/providers/transports/openai_chat/stream.py index 0571516c17ab622cf49202dc106ededadaa14c91..263e3b6cff547b4f5ea46342cbdd5e0677eee756 100644 --- a/src/free_claude_code/providers/transports/openai_chat/stream.py +++ b/src/free_claude_code/providers/transports/openai_chat/stream.py @@ -104,7 +104,7 @@ class OpenAIChatStreamAdapter: tool_argument_aliases: dict[str, dict[str, str]] = {} tool_argument_alias_buffers: dict[int, str] = {} - async with self._transport._global_rate_limiter.concurrency_slot(): + async with self._transport._rate_limiter.concurrency_slot(): while True: if not ledger.message_started: for event in hold_event(ledger.message_start()): @@ -306,7 +306,7 @@ class OpenAIChatStreamAdapter: provider_name=tag, read_timeout_s=self._transport._config.http_read_timeout, request_id=self._request_id, - rate_limiter=self._transport._global_rate_limiter, + rate_limiter=self._transport._rate_limiter, ) from error for event in ledger.terminal_error_tail( error_message, diff --git a/src/free_claude_code/providers/transports/openai_chat/transport.py b/src/free_claude_code/providers/transports/openai_chat/transport.py index 8f0891e2b29e5dd0704892041b752014c7e4a6d4..ad8183064458821124d16bf3e20393732094fb18 100644 --- a/src/free_claude_code/providers/transports/openai_chat/transport.py +++ b/src/free_claude_code/providers/transports/openai_chat/transport.py @@ -16,7 +16,7 @@ from free_claude_code.providers.error_mapping import ( user_visible_message_for_mapped_provider_error, ) from free_claude_code.providers.model_listing import extract_openai_model_ids -from free_claude_code.providers.rate_limit import GlobalRateLimiter +from free_claude_code.providers.rate_limit import ProviderRateLimiter from .output_cap import clamp_output_tokens, parse_output_token_cap from .stream import OpenAIChatStreamAdapter @@ -33,6 +33,7 @@ class OpenAIChatTransport(BaseProvider): provider_name: str, base_url: str, api_key: str, + rate_limiter: ProviderRateLimiter, default_headers: Mapping[str, str] | None = None, ): super().__init__(config) @@ -42,12 +43,7 @@ class OpenAIChatTransport(BaseProvider): # Learned per-model output-token caps from upstream 400 rejections, so # later requests clamp proactively instead of paying the 400 each time. self._model_output_caps: dict[str, int] = {} - self._global_rate_limiter = GlobalRateLimiter.get_scoped_instance( - provider_name.lower(), - rate_limit=config.rate_limit, - rate_window=config.rate_window, - max_concurrency=config.max_concurrency, - ) + self._rate_limiter = rate_limiter http_client = None if config.proxy: http_client = httpx.AsyncClient( @@ -125,7 +121,7 @@ class OpenAIChatTransport(BaseProvider): while True: try: create_body = self._prepare_create_body(body) - stream = await self._global_rate_limiter.execute_with_retry( + stream = await self._rate_limiter.execute_with_retry( self._client.chat.completions.create, **create_body, stream=True ) return stream, body @@ -198,7 +194,7 @@ class OpenAIChatTransport(BaseProvider): def _map_error_details( self, error: Exception, request_id: str | None ) -> tuple[Exception, str]: - mapped_error = map_error(error, rate_limiter=self._global_rate_limiter) + mapped_error = map_error(error, rate_limiter=self._rate_limiter) return ( mapped_error, user_visible_message_for_mapped_provider_error( diff --git a/src/free_claude_code/providers/vercel/client.py b/src/free_claude_code/providers/vercel/client.py index 6bbdc79d507f7dfa706ec3d516e99d3dbec5ebb8..11235481622899ce6bddfc35714916d290317e0a 100644 --- a/src/free_claude_code/providers/vercel/client.py +++ b/src/free_claude_code/providers/vercel/client.py @@ -4,6 +4,7 @@ from typing import Any from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import VERCEL_AI_GATEWAY_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -19,12 +20,13 @@ _REQUEST_POLICY = OpenAIChatRequestPolicy( class VercelProvider(OpenAIChatTransport): """Vercel AI Gateway at ``https://ai-gateway.vercel.sh/v1``.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="VERCEL", base_url=config.base_url or VERCEL_AI_GATEWAY_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/wafer/client.py b/src/free_claude_code/providers/wafer/client.py index 01e448c159ecc5f6406dda229a232467386a8080..54022b6e1ac2821f980f1ce80227a7dd6231011f 100644 --- a/src/free_claude_code/providers/wafer/client.py +++ b/src/free_claude_code/providers/wafer/client.py @@ -5,6 +5,7 @@ from typing import Any from free_claude_code.config.constants import ANTHROPIC_DEFAULT_MAX_OUTPUT_TOKENS from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import WAFER_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -20,12 +21,13 @@ _REQUEST_POLICY = OpenAIChatRequestPolicy( class WaferProvider(OpenAIChatTransport): """Wafer using ``https://pass.wafer.ai/v1/chat/completions``.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="WAFER", base_url=config.base_url or WAFER_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/providers/zai/client.py b/src/free_claude_code/providers/zai/client.py index 0b9dd7c7894a1dd1e460ad77894557dfea4f24ba..e5d85b57582972e27d0116342bbb824104e6bfb9 100644 --- a/src/free_claude_code/providers/zai/client.py +++ b/src/free_claude_code/providers/zai/client.py @@ -5,6 +5,7 @@ from typing import Any from free_claude_code.config.constants import ANTHROPIC_DEFAULT_MAX_OUTPUT_TOKENS from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.defaults import ZAI_DEFAULT_BASE +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import ( OpenAIChatRequestPolicy, OpenAIChatTransport, @@ -23,12 +24,13 @@ _REQUEST_POLICY = OpenAIChatRequestPolicy( class ZaiProvider(OpenAIChatTransport): """Z.ai Coding Plan via ``https://api.z.ai/api/coding/paas/v4``.""" - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="ZAI", base_url=config.base_url or ZAI_DEFAULT_BASE, api_key=config.api_key, + rate_limiter=rate_limiter, ) def _build_request_body( diff --git a/src/free_claude_code/runtime/application.py b/src/free_claude_code/runtime/application.py index caf6428a23366c7d27a4ef236a636ab4646afaf2..8c49980b7f74bf2dfaf11a28b2c082fd2db93ee2 100644 --- a/src/free_claude_code/runtime/application.py +++ b/src/free_claude_code/runtime/application.py @@ -11,7 +11,6 @@ from typing import Any from loguru import logger import free_claude_code.cli.managed as cli_managed -import free_claude_code.messaging.limiter as messaging_limiter import free_claude_code.messaging.session as messaging_session import free_claude_code.messaging.workflow as messaging_workflow_module from free_claude_code.api.ports import StopResult @@ -36,26 +35,29 @@ from free_claude_code.messaging.platforms.ports import ( MessagingPlatformComponents, MessagingRuntime, ) +from free_claude_code.messaging.voice import Transcriber from free_claude_code.providers.exceptions import ServiceUnavailableError from .provider_manager import ProviderRuntimeManager -_SHUTDOWN_TIMEOUT_S = 5.0 RestartCallback = Callable[[], Awaitable[None] | None] async def best_effort( name: str, awaitable: Awaitable[Any], - timeout_s: float = _SHUTDOWN_TIMEOUT_S, *, log_verbose_errors: bool = False, -) -> None: - """Run one bounded cleanup step without masking later cleanup.""" +) -> bool: + """Run one cleanup step and report whether it completed. + + The lifecycle owner intentionally applies no generic timeout here. Cancelling + an arbitrary cleanup at a deadline can abandon a half-closed SDK, thread, or + provider resource; resource-specific cleanup or the process supervisor owns + any force-termination deadline. + """ try: - await asyncio.wait_for(awaitable, timeout=timeout_s) - except TimeoutError: - logger.warning("Shutdown step timed out: {} ({}s)", name, timeout_s) + await awaitable except Exception as exc: if log_verbose_errors: logger.warning( @@ -70,6 +72,8 @@ async def best_effort( name, type(exc).__name__, ) + return False + return True def warn_if_process_auth_token(settings: Settings) -> None: @@ -100,9 +104,11 @@ class ApplicationRuntime: self, provider_manager: ProviderRuntimeManager, *, + transcriber: Transcriber | None, restart_callback: RestartCallback | None = None, ) -> None: self.provider_manager = provider_manager + self._transcriber = transcriber self._restart_callback = restart_callback self._config_lock = asyncio.Lock() self._pending_fields: list[str] = [] @@ -113,6 +119,8 @@ class ApplicationRuntime: self._cli_manager: cli_managed.ManagedClaudeSessionManager | None = None self._started = False self._closed = False + self._provider_manager_closed = False + self._close_lock = asyncio.Lock() @property def settings(self) -> Settings: @@ -132,32 +140,31 @@ class ApplicationRuntime: local_admin_url(self.settings), ) self._started = True + except asyncio.CancelledError: + await self.close() + raise except Exception as exc: logger.error( "Startup failed:\n{}", startup_failure_message(self.settings, exc), ) - await self._cleanup_messaging() - await best_effort( - "provider_manager.close", - self.provider_manager.close(), - log_verbose_errors=self.settings.log_api_error_tracebacks, - ) + await self.close() raise - async def close(self) -> None: - if self._closed: - return - self._closed = True - logger.info("Shutdown requested, cleaning up...") - await self._cleanup_messaging() - await best_effort( - "provider_manager.close", - self.provider_manager.close(), - log_verbose_errors=self.settings.log_api_error_tracebacks, - ) - await self._shutdown_limiter() - logger.info("Server shut down cleanly") + async def close(self) -> bool: + async with self._close_lock: + if self._closed: + return True + logger.info("Shutdown requested, cleaning up...") + self._closed = await self._close_owned_resources() + if self._closed: + self._started = False + logger.info("Server shut down cleanly") + else: + logger.warning( + "Server shutdown incomplete; owned resources remain for retry" + ) + return self._closed async def apply_admin_config( self, @@ -326,14 +333,11 @@ class ApplicationRuntime: telegram_proxy_url=settings.telegram_proxy_url, discord_bot_token=settings.discord_bot_token, allowed_discord_channels=settings.allowed_discord_channels, - voice_note_enabled=settings.voice_note_enabled, - whisper_model=settings.whisper_model, - whisper_device=settings.whisper_device, - huggingface_api_key=settings.huggingface_api_key, - nvidia_nim_api_key=settings.nvidia_nim_api_key, + transcriber=self._transcriber, messaging_rate_limit=settings.messaging_rate_limit, messaging_rate_window=settings.messaging_rate_window, log_raw_messaging_content=settings.log_raw_messaging_content, + log_messaging_error_details=settings.log_messaging_error_details, log_api_error_tracebacks=settings.log_api_error_tracebacks, ) @@ -342,6 +346,7 @@ class ApplicationRuntime: components: MessagingPlatformComponents, ) -> None: settings = self.settings + self._messaging_runtime = components.runtime workspace = ( os.path.abspath(settings.allowed_dir) if settings.allowed_dir @@ -366,7 +371,6 @@ class ApplicationRuntime: storage_path=os.path.join(data_path, "sessions.json"), message_log_cap=settings.max_message_log_entries_per_chat, ) - self._messaging_runtime = components.runtime self._messaging_workflow = messaging_workflow_module.MessagingWorkflow( platform_name=components.name, outbound=components.outbound, @@ -384,16 +388,48 @@ class ApplicationRuntime: await components.runtime.start() logger.info("{} platform started with messaging workflow", components.name) - async def _cleanup_messaging(self) -> None: + async def _close_owned_resources(self) -> bool: + if not await self._cleanup_messaging(): + return False + if not await self._cleanup_transcriber(): + return False + if self._provider_manager_closed: + return True + verbose = self.settings.log_api_error_tracebacks + self._provider_manager_closed = await best_effort( + "provider_manager.close", + self.provider_manager.close(), + log_verbose_errors=verbose, + ) + return self._provider_manager_closed + + async def _cleanup_messaging(self) -> bool: verbose = self.settings.log_api_error_tracebacks workflow = self._messaging_workflow runtime = self._messaging_runtime cli_manager = self._cli_manager - self._messaging_workflow = None - self._messaging_runtime = None - self._cli_manager = None + + if runtime is not None: + quiesced = await best_effort( + "messaging_runtime.quiesce", + runtime.quiesce(), + log_verbose_errors=verbose, + ) + if not quiesced: + # Delivery must remain available until ingress is known stopped. + # Retaining the graph lets the next close retry this exact gate. + return False if workflow is not None: + drained = await best_effort( + "messaging_workflow.stop_all_tasks", + workflow.stop_all_tasks(), + log_verbose_errors=verbose, + ) + if not drained: + # Active workflow tasks may still need delivery, transcription, + # CLI sessions, and providers while a later close retries drain. + return False try: workflow.close() except Exception as exc: @@ -404,37 +440,43 @@ class ApplicationRuntime: "Session store flush on shutdown: exc_type={}", type(exc).__name__, ) - if runtime is not None: - await best_effort( - "messaging_runtime.stop", - runtime.stop(), - log_verbose_errors=verbose, - ) - if cli_manager is not None: - await best_effort( + return False + if self._messaging_workflow is workflow: + self._messaging_workflow = None + if self._cli_manager is cli_manager: + self._cli_manager = None + elif cli_manager is not None: + drained = await best_effort( "cli_manager.stop_all", cli_manager.stop_all(), log_verbose_errors=verbose, ) + if not drained: + return False + if self._cli_manager is cli_manager: + self._cli_manager = None - async def _shutdown_limiter(self) -> None: - verbose = self.settings.log_api_error_tracebacks - try: - await best_effort( - "MessagingRateLimiter.shutdown_instance", - messaging_limiter.MessagingRateLimiter.shutdown_instance(), - timeout_s=2.0, + if runtime is not None: + closed = await best_effort( + "messaging_runtime.close", + runtime.close(), log_verbose_errors=verbose, ) - except Exception as exc: - if verbose: - logger.debug( - "Rate limiter shutdown skipped: {}: {}", - type(exc).__name__, - exc, - ) - else: - logger.debug( - "Rate limiter shutdown skipped: exc_type={}", - type(exc).__name__, - ) + if not closed: + return False + if self._messaging_runtime is runtime: + self._messaging_runtime = None + return True + + async def _cleanup_transcriber(self) -> bool: + transcriber = self._transcriber + if transcriber is None: + return True + closed = await best_effort( + "transcriber.close", + transcriber.close(), + log_verbose_errors=self.settings.log_api_error_tracebacks, + ) + if closed and self._transcriber is transcriber: + self._transcriber = None + return closed diff --git a/src/free_claude_code/runtime/asgi.py b/src/free_claude_code/runtime/asgi.py index 37a0ad47baa4a2e49b07cca0fd0a74389b5cd3f8..654d53b5e2b0e6b37caaa424b777f0615449d70d 100644 --- a/src/free_claude_code/runtime/asgi.py +++ b/src/free_claude_code/runtime/asgi.py @@ -49,7 +49,7 @@ class RuntimeASGIApp: if message["type"] == "lifespan.shutdown": if started: try: - await self.runtime.close() + closed = await self.runtime.close() except Exception as exc: logger.error( "Shutdown failed: exc_type={}", @@ -57,5 +57,8 @@ class RuntimeASGIApp: ) await send({"type": "lifespan.shutdown.failed", "message": ""}) return + if not closed: + await send({"type": "lifespan.shutdown.failed", "message": ""}) + return await send({"type": "lifespan.shutdown.complete"}) return diff --git a/src/free_claude_code/runtime/bootstrap.py b/src/free_claude_code/runtime/bootstrap.py index 7206862d54e60c624e697724952da311e35d9d34..83c672aa3ffc9b90bc53a3dc782481a2f255d24e 100644 --- a/src/free_claude_code/runtime/bootstrap.py +++ b/src/free_claude_code/runtime/bootstrap.py @@ -8,6 +8,9 @@ from free_claude_code.api.ports import ApiServices from free_claude_code.config.logging_config import configure_logging from free_claude_code.config.paths import server_log_path from free_claude_code.config.settings import Settings +from free_claude_code.messaging.transcription import TranscriptionService +from free_claude_code.messaging.voice import Transcriber +from free_claude_code.providers.nvidia_nim.voice import NvidiaNimTranscriber from .application import ApplicationRuntime, RestartCallback from .asgi import RuntimeASGIApp @@ -24,6 +27,7 @@ def build_asgi_app( provider_manager = ProviderRuntimeManager(settings) runtime = ApplicationRuntime( provider_manager, + transcriber=_create_transcriber(settings), restart_callback=restart_callback, ) services = ApiServices( @@ -32,3 +36,18 @@ def build_asgi_app( sessions=runtime, ) return RuntimeASGIApp(create_app(services), runtime) + + +def _create_transcriber(settings: Settings) -> Transcriber | None: + if not settings.voice_note_enabled: + return None + if settings.whisper_device == "nvidia_nim": + return NvidiaNimTranscriber( + model=settings.whisper_model, + api_key=settings.nvidia_nim_api_key, + ) + return TranscriptionService( + model=settings.whisper_model, + device=settings.whisper_device, + huggingface_api_key=settings.huggingface_api_key, + ) diff --git a/src/free_claude_code/runtime/provider_manager.py b/src/free_claude_code/runtime/provider_manager.py index 8f5ab01422dd460c0aa0bb39de30c63f5e66a97e..041aa9a322324b1acd0b938b4cf84e067fa023ec 100644 --- a/src/free_claude_code/runtime/provider_manager.py +++ b/src/free_claude_code/runtime/provider_manager.py @@ -2,7 +2,6 @@ import asyncio from collections.abc import Callable, Iterable -from contextlib import suppress from dataclasses import dataclass, field from loguru import logger @@ -30,6 +29,7 @@ class _ProviderGeneration: retired: bool = False closed: bool = False drained: asyncio.Event = field(default_factory=asyncio.Event) + cleanup_task: asyncio.Task[bool] | None = None def __post_init__(self) -> None: self.drained.set() @@ -85,10 +85,13 @@ class ProviderRuntimeManager: ) -> None: self._runtime_factory = runtime_factory self._replace_lock = asyncio.Lock() + self._close_lock = asyncio.Lock() self._model_cache = ProviderModelCache() self._refresh_task: asyncio.Task[None] | None = None self._next_generation_id = 2 self._retired: dict[int, _ProviderGeneration] = {} + self._unpublished: set[ProviderRuntime] = set() + self._closing = False self._closed = False self._current = _ProviderGeneration( generation_id=1, @@ -102,7 +105,7 @@ class ProviderRuntimeManager: return self._current.generation_id async def acquire(self) -> ProviderGenerationLease: - if self._closed: + if self._closing or self._closed: raise ServiceUnavailableError("Provider runtime is shutting down.") generation = self._current generation.active_leases += 1 @@ -144,6 +147,8 @@ class ProviderRuntimeManager: def start_model_list_refresh(self) -> None: """Start one non-blocking refresh for the current generation.""" + if self._closing or self._closed: + return if self._refresh_task is not None and not self._refresh_task.done(): return generation = self._current @@ -154,6 +159,8 @@ class ProviderRuntimeManager: async def refresh_model_list_cache(self) -> None: """Run an explicit full refresh without racing replacement.""" async with self._replace_lock: + if self._closing or self._closed: + raise ServiceUnavailableError("Provider runtime is shutting down.") await self._cancel_refresh() await self._refresh_generation(self._current, only_missing=False) @@ -166,7 +173,10 @@ class ProviderRuntimeManager: ) -> int: """Prepare, commit, and atomically publish one replacement generation.""" async with self._replace_lock: + if self._closing or self._closed: + raise ServiceUnavailableError("Provider runtime is shutting down.") await self._cancel_refresh() + await self._retry_unpublished_cleanup() candidate_id = self._next_generation_id candidate_runtime: ProviderRuntime | None = None try: @@ -200,31 +210,42 @@ class ProviderRuntimeManager: self._trace_published(candidate, previous=previous, reason=reason) self._trace_retired(previous, reason=reason) - if previous.active_leases == 0: - await self._close_generation(previous, forced=False) self._refresh_task = asyncio.create_task( self._refresh_generation(candidate, only_missing=False) ) + if previous.active_leases == 0: + await self._close_generation(previous, forced=False) return candidate.generation_id async def close(self) -> None: """Reject new leases, drain existing work, and close every generation.""" - async with self._replace_lock: + async with self._close_lock: if self._closed: return + async with self._replace_lock: + self._closing = True + await self._cancel_refresh() + current = self._current + if not current.retired: + current.retired = True + self._retired[current.generation_id] = current + self._trace_retired(current, reason="shutdown") + generations = tuple(self._retired.values()) + + await asyncio.gather( + *(generation.drained.wait() for generation in generations) + ) + generation_results = await asyncio.gather( + *( + self._close_generation(generation, forced=False) + for generation in generations + ) + ) + unpublished_closed = await self._retry_unpublished_cleanup() + if not all(generation_results) or not unpublished_closed: + raise RuntimeError("One or more provider runtimes failed to close.") + self._model_cache.clear() self._closed = True - await self._cancel_refresh() - current = self._current - if not current.retired: - current.retired = True - self._retired[current.generation_id] = current - self._trace_retired(current, reason="shutdown") - generations = tuple(self._retired.values()) - - await asyncio.gather(*(generation.drained.wait() for generation in generations)) - for generation in generations: - await self._close_generation(generation, forced=False) - self._model_cache.clear() async def _release(self, generation: _ProviderGeneration) -> None: if generation.active_leases <= 0: @@ -233,7 +254,7 @@ class ProviderRuntimeManager: if generation.active_leases != 0: return generation.drained.set() - if generation.retired: + if generation.retired and not self._closing: await self._close_generation(generation, forced=False) async def _refresh_generation( @@ -269,38 +290,70 @@ class ProviderRuntimeManager: if task is None or task.done(): return task.cancel() - with suppress(asyncio.CancelledError): - await task + await asyncio.gather(task, return_exceptions=True) - async def _cleanup_unpublished(self, runtime: ProviderRuntime) -> None: + async def _cleanup_unpublished(self, runtime: ProviderRuntime) -> bool: + self._unpublished.add(runtime) try: await runtime.cleanup() + except asyncio.CancelledError: + raise except Exception as exc: logger.warning( "Unpublished provider generation cleanup failed: exc_type={}", type(exc).__name__, ) + return False + self._unpublished.discard(runtime) + return True + + async def _retry_unpublished_cleanup(self) -> bool: + all_closed = True + for runtime in tuple(self._unpublished): + if not await self._cleanup_unpublished(runtime): + all_closed = False + return all_closed async def _close_generation( self, generation: _ProviderGeneration, *, forced: bool, - ) -> None: - if generation.closed or generation.active_leases != 0: - return - generation.closed = True - outcome = "ok" - try: - await generation.runtime.cleanup() - except Exception as exc: - outcome = "error" - logger.warning( - "Provider generation cleanup failed: generation_id={} exc_type={}", - generation.generation_id, - type(exc).__name__, + ) -> bool: + if generation.closed: + return True + if generation.active_leases != 0: + return False + task = generation.cleanup_task + if task is None: + task = asyncio.create_task( + self._run_generation_cleanup(generation, forced=forced), + name=f"provider-generation-cleanup-{generation.generation_id}", ) - finally: + generation.cleanup_task = task + return await asyncio.shield(task) + + async def _run_generation_cleanup( + self, + generation: _ProviderGeneration, + *, + forced: bool, + ) -> bool: + task = asyncio.current_task() + try: + try: + await generation.runtime.cleanup() + except asyncio.CancelledError: + raise + except Exception as exc: + logger.warning( + "Provider generation cleanup failed: generation_id={} exc_type={}", + generation.generation_id, + type(exc).__name__, + ) + return False + + generation.closed = True self._retired.pop(generation.generation_id, None) trace_event( stage="runtime", @@ -309,8 +362,12 @@ class ProviderRuntimeManager: generation_id=generation.generation_id, active_leases=generation.active_leases, forced=forced, - outcome=outcome, + outcome="ok", ) + return True + finally: + if not generation.closed and generation.cleanup_task is task: + generation.cleanup_task = None @staticmethod def _trace_published( diff --git a/tests/api/support.py b/tests/api/support.py index 4f21eb20e85b00b3f2daebdb3cebca7e669a26cb..64024b6840a2097939efcbc861da24f6461254c9 100644 --- a/tests/api/support.py +++ b/tests/api/support.py @@ -31,7 +31,11 @@ def create_test_app( dict(providers), ), ) - runtime = ApplicationRuntime(manager, restart_callback=restart_callback) + runtime = ApplicationRuntime( + manager, + transcriber=None, + restart_callback=restart_callback, + ) return create_app( ApiServices( requests=manager, diff --git a/tests/api/test_admin.py b/tests/api/test_admin.py index fb7a465dde828875b9271fb5c6c648ee6bd54649..add2d082556636485c90f1caf78e8e8d4fa3839a 100644 --- a/tests/api/test_admin.py +++ b/tests/api/test_admin.py @@ -2,6 +2,7 @@ from pathlib import Path from unittest.mock import patch import httpx +import pytest from fastapi.testclient import TestClient from free_claude_code.config.admin.values import MASKED_SECRET @@ -24,6 +25,7 @@ def _clear_process_config(monkeypatch) -> None: for key in ( "MODEL", "NVIDIA_NIM_API_KEY", + "HUGGINGFACE_API_KEY", "OPENROUTER_API_KEY", "ANTHROPIC_AUTH_TOKEN", "TELEGRAM_PROXY_URL", @@ -34,6 +36,8 @@ def _clear_process_config(monkeypatch) -> None: "SAMBANOVA_API_KEY", "HOST", "PORT", + "VOICE_NOTE_ENABLED", + "WHISPER_DEVICE", "LOG_FILE", "ZAI_BASE_URL", "CLAUDE_WORKSPACE", @@ -132,6 +136,19 @@ def test_admin_config_masks_secrets_and_exposes_manifest(monkeypatch, tmp_path): field for field in body["fields"] if field["key"] == "TELEGRAM_PROXY_URL" ) assert telegram_proxy_field["secret"] is True + restart_required = { + field["key"] for field in body["fields"] if field["restart_required"] is True + } + assert { + "ANTHROPIC_AUTH_TOKEN", + "DEBUG_PLATFORM_EDITS", + "DEBUG_SUBAGENT_STACK", + "LOG_RAW_API_PAYLOADS", + "LOG_API_ERROR_TRACEBACKS", + "LOG_RAW_MESSAGING_CONTENT", + "LOG_RAW_CLI_DIAGNOSTICS", + "LOG_MESSAGING_ERROR_DETAILS", + } <= restart_required def test_admin_config_preserves_managed_env_source_contract(monkeypatch, tmp_path): @@ -393,6 +410,7 @@ def test_admin_apply_writes_huggingface_key_and_masks_preview(monkeypatch, tmp_p assert response.status_code == 200 body = response.json() assert body["applied"] is True + assert body["pending_fields"] == [] assert "HUGGINGFACE_API_KEY=********" in body["env_preview"] env_file = tmp_path / ".fcc" / ".env" text = env_file.read_text(encoding="utf-8") @@ -400,6 +418,97 @@ def test_admin_apply_writes_huggingface_key_and_masks_preview(monkeypatch, tmp_p assert "HUGGINGFACE_API_KEY=hf-secret" in text +@pytest.mark.parametrize( + ("device", "credential_key"), + [ + ("nvidia_nim", "NVIDIA_NIM_API_KEY"), + ("cpu", "HUGGINGFACE_API_KEY"), + ], +) +def test_admin_key_change_requires_restart_for_active_voice_backend( + monkeypatch, + tmp_path, + device, + credential_key, +): + _set_home(monkeypatch, tmp_path) + _clear_process_config(monkeypatch) + env_file = tmp_path / ".fcc" / ".env" + env_file.parent.mkdir(parents=True) + env_file.write_text( + "\n".join( + [ + "VOICE_NOTE_ENABLED=true", + f"WHISPER_DEVICE={device}", + f"{credential_key}=old-key", + "", + ] + ), + encoding="utf-8", + ) + app = create_test_app() + + response = _local_client(app).post( + "/admin/api/config/apply", + json={"values": {credential_key: "new-key"}}, + ) + + assert response.status_code == 200 + body = response.json() + assert body["applied"] is True + assert body["pending_fields"] == [credential_key] + assert body["restart"] == { + "required": True, + "automatic": False, + "admin_url": None, + "fields": [credential_key], + } + + +@pytest.mark.parametrize( + ("key", "initial", "updated"), + [ + ("ANTHROPIC_AUTH_TOKEN", "old-token", "new-token"), + ("DEBUG_PLATFORM_EDITS", "true", "false"), + ("DEBUG_SUBAGENT_STACK", "true", "false"), + ("LOG_RAW_API_PAYLOADS", "true", "false"), + ("LOG_API_ERROR_TRACEBACKS", "true", "false"), + ("LOG_RAW_MESSAGING_CONTENT", "true", "false"), + ("LOG_RAW_CLI_DIAGNOSTICS", "true", "false"), + ("LOG_MESSAGING_ERROR_DETAILS", "true", "false"), + ], +) +def test_admin_constructor_captured_setting_requires_restart( + monkeypatch, + tmp_path, + key, + initial, + updated, +): + _set_home(monkeypatch, tmp_path) + _clear_process_config(monkeypatch) + env_file = tmp_path / ".fcc" / ".env" + env_file.parent.mkdir(parents=True) + env_file.write_text(f"{key}={initial}\n", encoding="utf-8") + app = create_test_app() + + response = _local_client(app).post( + "/admin/api/config/apply", + json={"values": {key: updated}}, + ) + + assert response.status_code == 200 + body = response.json() + assert body["applied"] is True + assert body["pending_fields"] == [key] + assert body["restart"] == { + "required": True, + "automatic": False, + "admin_url": None, + "fields": [key], + } + + def test_admin_apply_writes_cohere_key_and_masks_preview(monkeypatch, tmp_path): _set_home(monkeypatch, tmp_path) _clear_process_config(monkeypatch) diff --git a/tests/api/test_app_lifespan_and_errors.py b/tests/api/test_app_lifespan_and_errors.py index 97999806c1551d5a72f779d1e682b422b0f90bea..e81b25e2e482567c2e929e548b6493487a08df3c 100644 --- a/tests/api/test_app_lifespan_and_errors.py +++ b/tests/api/test_app_lifespan_and_errors.py @@ -8,17 +8,20 @@ from fastapi import FastAPI from fastapi.testclient import TestClient from free_claude_code.config.settings import Settings +from free_claude_code.messaging.transcription import TranscriptionService from free_claude_code.providers.exceptions import ( AuthenticationError, ServiceUnavailableError, ) +from free_claude_code.providers.nvidia_nim.client import NvidiaNimProvider +from free_claude_code.providers.nvidia_nim.voice import NvidiaNimTranscriber from free_claude_code.runtime.application import ( ApplicationRuntime, startup_failure_message, warn_if_process_auth_token, ) from free_claude_code.runtime.asgi import RuntimeASGIApp -from free_claude_code.runtime.bootstrap import build_asgi_app +from free_claude_code.runtime.bootstrap import _create_transcriber, build_asgi_app from free_claude_code.runtime.provider_manager import ProviderRuntimeManager from tests.api.support import create_test_app @@ -65,7 +68,7 @@ async def test_runtime_startup_logs_admin_url_without_printed_server_banner(): port=9099, ) manager = ProviderRuntimeManager(settings) - runtime = ApplicationRuntime(manager) + runtime = ApplicationRuntime(manager, transcriber=None) uvicorn_logger = MagicMock() with ( @@ -166,7 +169,7 @@ def test_general_exception_default_log_excludes_exception_message(): async def test_model_validation_failure_does_not_block_runtime_startup(): settings = _settings(messaging_platform="none") manager = ProviderRuntimeManager(settings) - runtime = ApplicationRuntime(manager) + runtime = ApplicationRuntime(manager, transcriber=None) validation = AsyncMock(side_effect=ServiceUnavailableError("bad model")) with ( @@ -209,7 +212,7 @@ async def test_runtime_asgi_app_starts_and_closes_owner_once(): runtime = MagicMock(spec=ApplicationRuntime) runtime.settings = _settings() runtime.start = AsyncMock() - runtime.close = AsyncMock() + runtime.close = AsyncMock(return_value=True) app = RuntimeASGIApp(AsyncMock(), runtime) received = iter( [ @@ -235,6 +238,35 @@ async def test_runtime_asgi_app_starts_and_closes_owner_once(): ] +@pytest.mark.asyncio +async def test_runtime_asgi_app_reports_incomplete_owned_shutdown() -> None: + runtime = MagicMock(spec=ApplicationRuntime) + runtime.settings = _settings() + runtime.start = AsyncMock() + runtime.close = AsyncMock(return_value=False) + app = RuntimeASGIApp(AsyncMock(), runtime) + received = iter( + [ + {"type": "lifespan.startup"}, + {"type": "lifespan.shutdown"}, + ] + ) + sent: list[dict[str, str]] = [] + + async def receive(): + return next(received) + + async def send(message): + sent.append(message) + + await app({"type": "lifespan"}, receive, send) + + assert sent == [ + {"type": "lifespan.startup.complete"}, + {"type": "lifespan.shutdown.failed", "message": ""}, + ] + + @pytest.mark.asyncio async def test_runtime_asgi_app_reports_concise_startup_failure(): runtime = MagicMock(spec=ApplicationRuntime) @@ -290,3 +322,58 @@ def test_bootstrap_honors_process_log_file_override(monkeypatch, tmp_path): build_asgi_app(_settings()) assert configure.call_args.args[0] == log_path + + +def test_bootstrap_constructs_fresh_runtime_owned_transcribers() -> None: + settings = _settings(voice_note_enabled=True, whisper_device="cpu") + + first = _create_transcriber(settings) + second = _create_transcriber(settings) + + assert isinstance(first, TranscriptionService) + assert isinstance(second, TranscriptionService) + assert first is not second + + +@pytest.mark.asyncio +async def test_bootstrap_constructs_isolated_runtime_resource_graphs() -> None: + settings = _settings( + model="nvidia_nim/test-model", + voice_note_enabled=True, + whisper_device="cpu", + ) + + with patch("free_claude_code.runtime.bootstrap.configure_logging"): + first = build_asgi_app(settings) + second = build_asgi_app(settings) + + first_lease = await first.runtime.provider_manager.acquire() + second_lease = await second.runtime.provider_manager.acquire() + try: + first_provider = first_lease.resolve_provider("nvidia_nim") + second_provider = second_lease.resolve_provider("nvidia_nim") + + assert isinstance(first_provider, NvidiaNimProvider) + assert isinstance(second_provider, NvidiaNimProvider) + assert first_provider._rate_limiter is not second_provider._rate_limiter + assert first.runtime._transcriber is not second.runtime._transcriber + finally: + await first_lease.release() + await second_lease.release() + await first.runtime.close() + await second.runtime.close() + + +def test_bootstrap_selects_nvidia_transcriber_without_loading_riva() -> None: + settings = _settings( + voice_note_enabled=True, + whisper_device="nvidia_nim", + whisper_model="openai/whisper-large-v3", + nvidia_nim_api_key="nvapi-test", + ) + + assert isinstance(_create_transcriber(settings), NvidiaNimTranscriber) + + +def test_bootstrap_disables_transcription_as_one_owned_resource() -> None: + assert _create_transcriber(_settings(voice_note_enabled=False)) is None diff --git a/tests/api/test_runtime_safe_logging.py b/tests/api/test_runtime_safe_logging.py index e6b32ad3a5fb1ebfd84762f6857b253536e54a8a..faa6ec892a03c8d7981184af2238542d15e10bc8 100644 --- a/tests/api/test_runtime_safe_logging.py +++ b/tests/api/test_runtime_safe_logging.py @@ -20,7 +20,10 @@ async def test_messaging_start_failure_default_logs_exclude_traceback(caplog): "log_api_error_tracebacks": False, } ) - runtime = ApplicationRuntime(ProviderRuntimeManager(settings)) + runtime = ApplicationRuntime( + ProviderRuntimeManager(settings), + transcriber=None, + ) with ( patch( diff --git a/tests/conftest.py b/tests/conftest.py index 26ba301c5b91abcc03f7777dcc423a881570c457..a07368702ff4597b9f777acf1a7b27b3b12a4ac2 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -8,6 +8,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest from free_claude_code.config.settings import Settings +from tests.providers.support import passthrough_rate_limiter # Set mock environment BEFORE any imports that use Settings os.environ.setdefault("NVIDIA_NIM_API_KEY", "test_key") @@ -45,14 +46,18 @@ def nim_provider(provider_config): from free_claude_code.config.nim import NimSettings from free_claude_code.providers.nvidia_nim import NvidiaNimProvider - return NvidiaNimProvider(provider_config, nim_settings=NimSettings()) + return NvidiaNimProvider( + provider_config, + nim_settings=NimSettings(), + rate_limiter=passthrough_rate_limiter(), + ) @pytest.fixture def open_router_provider(provider_config): from free_claude_code.providers.open_router import OpenRouterProvider - return OpenRouterProvider(provider_config) + return OpenRouterProvider(provider_config, rate_limiter=passthrough_rate_limiter()) @pytest.fixture @@ -66,7 +71,7 @@ def lmstudio_provider(provider_config): rate_limit=provider_config.rate_limit, rate_window=provider_config.rate_window, ) - return LMStudioProvider(lmstudio_config) + return LMStudioProvider(lmstudio_config, rate_limiter=passthrough_rate_limiter()) @pytest.fixture @@ -80,7 +85,7 @@ def llamacpp_provider(provider_config): rate_limit=10, rate_window=60, ) - return LlamaCppProvider(llamacpp_config) + return LlamaCppProvider(llamacpp_config, rate_limiter=passthrough_rate_limiter()) @pytest.fixture diff --git a/tests/contracts/test_import_boundaries.py b/tests/contracts/test_import_boundaries.py index 88d8773a897b10f0f07248ad5fe7829b890c2f3d..aa6faaef1f30ef576ca81423d2c7739027bb441a 100644 --- a/tests/contracts/test_import_boundaries.py +++ b/tests/contracts/test_import_boundaries.py @@ -225,13 +225,8 @@ def test_settings_stays_schema_only() -> None: assert removed_api not in settings_text -_MESSAGING_ALLOWED_PROVIDER_MODULES = frozenset( - {"free_claude_code.providers.nvidia_nim.voice"} -) - - def test_messaging_does_not_import_disallowed_modules() -> None: - """Runtime composition is external; only NIM voice ASR crosses provider bounds.""" + """Runtime composition keeps messaging independent of concrete products.""" repo_root = Path(__file__).resolve().parents[2] offenders: list[str] = [] for path in (repo_root / "src" / "free_claude_code" / "messaging").rglob("*.py"): @@ -245,12 +240,7 @@ def test_messaging_does_not_import_disallowed_modules() -> None: or imported.startswith("free_claude_code.cli.") or imported == "smoke" or imported.startswith("smoke.") - ): - rel = path.relative_to(repo_root) - offenders.append(f"{rel}: {imported}") - elif imported.startswith("free_claude_code.providers."): - if imported in _MESSAGING_ALLOWED_PROVIDER_MODULES: - continue + ) or imported.startswith("free_claude_code.providers."): rel = path.relative_to(repo_root) offenders.append(f"{rel}: {imported}") @@ -801,7 +791,7 @@ def test_messaging_platforms_use_shared_outbox_and_voice_flow() -> None: text = runtime.read_text(encoding="utf-8") assert "PlatformOutbox" not in text assert "VoiceNoteFlow" in text - assert "from ..voice" not in text + assert "PendingVoiceRegistry" not in text assert "NamedTemporaryFile" not in text for messenger in { diff --git a/tests/messaging/test_discord_platform.py b/tests/messaging/test_discord_platform.py index 5e924072c438d7fcd71932c4f440fd90aa538255..460ee9c2009880ab088f832349efaf54480d728c 100644 --- a/tests/messaging/test_discord_platform.py +++ b/tests/messaging/test_discord_platform.py @@ -17,6 +17,22 @@ from free_claude_code.messaging.platforms.discord_inbound import ( from free_claude_code.messaging.platforms.discord_io import truncate_discord_message +def _limiter_mock() -> MagicMock: + limiter = MagicMock() + limiter.start = MagicMock() + limiter.shutdown = AsyncMock() + return limiter + + +def _discord_runtime(*args, limiter=None, transcriber=None, **kwargs) -> DiscordRuntime: + return DiscordRuntime( + *args, + limiter=limiter or _limiter_mock(), + transcriber=transcriber, + **kwargs, + ) + + class TestGetDiscord: """Tests for _get_discord helper.""" @@ -84,7 +100,7 @@ class TestDiscordRuntime: """Tests for Discord runtime and messenger behavior.""" def test_init_with_token(self): - platform = DiscordRuntime( + platform = _discord_runtime( bot_token="test_token", allowed_channel_ids="123,456", ) @@ -93,13 +109,13 @@ class TestDiscordRuntime: def test_init_without_allowed_channels(self): with patch.dict("os.environ", {"ALLOWED_DISCORD_CHANNELS": ""}, clear=False): - platform = DiscordRuntime(bot_token="token", allowed_channel_ids="") + platform = _discord_runtime(bot_token="token", allowed_channel_ids="") assert platform.allowed_channel_ids == set() def test_empty_allowed_channels_rejects_all_messages(self): """When allowed_channel_ids is empty, no channels are allowed (secure default).""" with patch.dict("os.environ", {"ALLOWED_DISCORD_CHANNELS": ""}, clear=False): - platform = DiscordRuntime(bot_token="token", allowed_channel_ids="") + platform = _discord_runtime(bot_token="token", allowed_channel_ids="") assert platform.allowed_channel_ids == set() # Empty set means: not self.allowed_channel_ids is True -> reject @@ -128,7 +144,7 @@ class TestDiscordRuntime: @pytest.mark.asyncio async def test_send_message_returns_message_id(self): - platform = DiscordRuntime(bot_token="token") + platform = _discord_runtime(bot_token="token") mock_msg = MagicMock() mock_msg.id = 999 mock_channel = AsyncMock() @@ -142,7 +158,7 @@ class TestDiscordRuntime: @pytest.mark.asyncio async def test_edit_message(self): - platform = DiscordRuntime(bot_token="token") + platform = _discord_runtime(bot_token="token") mock_msg = AsyncMock() mock_channel = AsyncMock() mock_channel.fetch_message = AsyncMock(return_value=mock_msg) @@ -155,7 +171,7 @@ class TestDiscordRuntime: @pytest.mark.asyncio async def test_send_message_channel_not_found_raises(self): - platform = DiscordRuntime(bot_token="token") + platform = _discord_runtime(bot_token="token") platform._connected = True with ( patch.object(platform._client, "get_channel", MagicMock(return_value=None)), @@ -165,7 +181,7 @@ class TestDiscordRuntime: @pytest.mark.asyncio async def test_send_message_channel_no_send_raises(self): - platform = DiscordRuntime(bot_token="token") + platform = _discord_runtime(bot_token="token") platform._connected = True mock_channel = MagicMock(spec=[]) # No send attr with ( @@ -177,9 +193,13 @@ class TestDiscordRuntime: await platform.outbound.send_message("123", "Hello") @pytest.mark.asyncio - async def test_queue_send_message_without_limiter_calls_send_message(self): - platform = DiscordRuntime(bot_token="token") - platform._limiter = None + async def test_queue_send_message_uses_required_limiter(self): + platform = _discord_runtime(bot_token="token") + + async def enqueue(operation, dedup_key=None): + return await operation() + + platform._limiter.enqueue = AsyncMock(side_effect=enqueue) platform._connected = True mock_channel = AsyncMock() mock_msg = MagicMock() @@ -188,14 +208,21 @@ class TestDiscordRuntime: with patch.object( platform._client, "get_channel", MagicMock(return_value=mock_channel) ): - result = await platform.outbound.queue_send_message("123", "hi") + result = await platform.outbound.queue_send_message( + "123", "hi", fire_and_forget=False + ) assert result == "42" + platform._limiter.enqueue.assert_awaited_once() mock_channel.send.assert_awaited_once() @pytest.mark.asyncio - async def test_queue_edit_message_without_limiter_calls_edit_message(self): - platform = DiscordRuntime(bot_token="token") - platform._limiter = None + async def test_queue_edit_message_uses_required_limiter(self): + platform = _discord_runtime(bot_token="token") + + async def enqueue(operation, dedup_key=None): + return await operation() + + platform._limiter.enqueue = AsyncMock(side_effect=enqueue) platform._connected = True mock_msg = AsyncMock() mock_channel = AsyncMock() @@ -203,12 +230,15 @@ class TestDiscordRuntime: with patch.object( platform._client, "get_channel", MagicMock(return_value=mock_channel) ): - await platform.outbound.queue_edit_message("123", "456", "Updated") + await platform.outbound.queue_edit_message( + "123", "456", "Updated", fire_and_forget=False + ) + platform._limiter.enqueue.assert_awaited_once() mock_msg.edit.assert_called_once_with(content="Updated") @pytest.mark.asyncio async def test_on_discord_message_bot_ignored(self): - platform = DiscordRuntime(bot_token="token", allowed_channel_ids="123") + platform = _discord_runtime(bot_token="token", allowed_channel_ids="123") handler = AsyncMock() platform.on_message(handler) msg = MagicMock() @@ -220,7 +250,7 @@ class TestDiscordRuntime: @pytest.mark.asyncio async def test_on_discord_message_empty_content_ignored(self): - platform = DiscordRuntime(bot_token="token", allowed_channel_ids="123") + platform = _discord_runtime(bot_token="token", allowed_channel_ids="123") handler = AsyncMock() platform.on_message(handler) msg = MagicMock() @@ -232,7 +262,7 @@ class TestDiscordRuntime: @pytest.mark.asyncio async def test_on_discord_message_channel_not_allowed_ignored(self): - platform = DiscordRuntime(bot_token="token", allowed_channel_ids="123") + platform = _discord_runtime(bot_token="token", allowed_channel_ids="123") handler = AsyncMock() platform.on_message(handler) msg = MagicMock() @@ -244,7 +274,7 @@ class TestDiscordRuntime: @pytest.mark.asyncio async def test_on_discord_message_valid_calls_handler(self): - platform = DiscordRuntime(bot_token="token", allowed_channel_ids="123") + platform = _discord_runtime(bot_token="token", allowed_channel_ids="123") handler = AsyncMock() platform.on_message(handler) msg = MagicMock() @@ -266,7 +296,7 @@ class TestDiscordRuntime: @pytest.mark.asyncio async def test_send_message_with_reply_to(self): - platform = DiscordRuntime(bot_token="token") + platform = _discord_runtime(bot_token="token") mock_msg = MagicMock() mock_msg.id = 999 mock_channel = AsyncMock() @@ -295,7 +325,7 @@ class TestDiscordRuntime: async def test_edit_message_not_found_returns_gracefully(self): import discord as discord_pkg - platform = DiscordRuntime(bot_token="token") + platform = _discord_runtime(bot_token="token") mock_channel = AsyncMock() mock_resp = MagicMock() mock_resp.status = 404 @@ -311,7 +341,7 @@ class TestDiscordRuntime: @pytest.mark.asyncio async def test_delete_message(self): - platform = DiscordRuntime(bot_token="token") + platform = _discord_runtime(bot_token="token") mock_msg = AsyncMock() mock_channel = AsyncMock() mock_channel.fetch_message = AsyncMock(return_value=mock_msg) @@ -331,24 +361,19 @@ class TestDiscordRuntime: @pytest.mark.asyncio async def test_fire_and_forget_with_coroutine(self): - platform = DiscordRuntime(bot_token="token") + platform = _discord_runtime(bot_token="token") + completed = asyncio.Event() async def _task(): - pass + completed.set() - coro = _task() - with patch("asyncio.create_task") as mock_create: - - def _run(c): - return asyncio.ensure_future(c) - - mock_create.side_effect = _run - platform.outbound.fire_and_forget(coro) - mock_create.assert_called_once() - await asyncio.sleep(0) + platform.outbound.fire_and_forget(_task()) + await completed.wait() + await platform.outbound.close() + assert platform.outbound._outbox._background_tasks == set() def test_on_message_registers_handler(self): - platform = DiscordRuntime(bot_token="token") + platform = _discord_runtime(bot_token="token") handler = AsyncMock() platform.on_message(handler) assert platform._message_handler is handler @@ -356,45 +381,85 @@ class TestDiscordRuntime: @pytest.mark.asyncio async def test_start_requires_token(self): with patch.dict("os.environ", {"DISCORD_BOT_TOKEN": ""}, clear=False): - platform = DiscordRuntime(bot_token="") + platform = _discord_runtime(bot_token="") with pytest.raises(ValueError, match="DISCORD_BOT_TOKEN"): await platform.start() @pytest.mark.asyncio async def test_start_connects(self): - platform = DiscordRuntime(bot_token="token") + limiter = _limiter_mock() + platform = _discord_runtime(bot_token="token", limiter=limiter) + keep_running = asyncio.Event() async def _fake_start(_token): - platform._connected = True + platform._mark_connected() + await keep_running.wait() - with ( - patch.object( - platform._client, - "start", - new_callable=AsyncMock, - side_effect=_fake_start, - ), - patch( - "free_claude_code.messaging.limiter.MessagingRateLimiter.get_instance", - new_callable=AsyncMock, - ), + with patch.object( + platform._client, + "start", + new_callable=AsyncMock, + side_effect=_fake_start, ): await platform.start() assert platform.is_connected is True + limiter.start.assert_called_once_with() + await platform.quiesce() + await platform.close() @pytest.mark.asyncio - async def test_stop_when_already_closed(self): - platform = DiscordRuntime(bot_token="token") + async def test_client_failure_after_readiness_marks_runtime_disconnected(self): + limiter = _limiter_mock() + platform = _discord_runtime(bot_token="token", limiter=limiter) + fail_client = asyncio.Event() + + async def _fake_start(_token): + platform._mark_connected() + await fail_client.wait() + raise RuntimeError("connection lost") + + with patch.object( + platform._client, + "start", + new_callable=AsyncMock, + side_effect=_fake_start, + ): + await platform.start() + assert platform.is_connected is True + + fail_client.set() + assert platform._start_task is not None + result = await asyncio.wait_for( + asyncio.gather(platform._start_task, return_exceptions=True), + timeout=1.0, + ) + assert isinstance(result[0], RuntimeError) + await asyncio.sleep(0) + + assert platform.is_connected is False + await platform.quiesce() + await platform.close() + + @pytest.mark.asyncio + async def test_quiesce_when_client_already_closed_then_close_delivery(self): + limiter = _limiter_mock() + platform = _discord_runtime(bot_token="token", limiter=limiter) platform._connected = True with patch.object( platform._client, "is_closed", new_callable=MagicMock, return_value=True ): - await platform.stop() + await platform.quiesce() assert platform.is_connected is False + limiter.shutdown.assert_not_awaited() + + await platform.close() + + limiter.shutdown.assert_awaited_once_with() @pytest.mark.asyncio - async def test_stop_closes_client(self): - platform = DiscordRuntime(bot_token="token") + async def test_quiesce_closes_client_without_closing_delivery(self): + limiter = _limiter_mock() + platform = _discord_runtime(bot_token="token", limiter=limiter) platform._connected = True mock_close = AsyncMock() with ( @@ -407,6 +472,131 @@ class TestDiscordRuntime: patch.object(platform._client, "close", mock_close), ): platform._start_task = None - await platform.stop() + await platform.quiesce() mock_close.assert_awaited_once() assert platform.is_connected is False + limiter.shutdown.assert_not_awaited() + + await platform.close() + + limiter.shutdown.assert_awaited_once_with() + + @pytest.mark.asyncio + async def test_quiesce_drains_start_task_when_client_close_fails(self): + limiter = _limiter_mock() + platform = _discord_runtime(bot_token="token", limiter=limiter) + platform._connected = True + started = asyncio.Event() + + async def pending_start() -> None: + started.set() + await asyncio.Event().wait() + + platform._start_task = asyncio.create_task(pending_start()) + await started.wait() + with ( + patch.object( + platform._client, + "is_closed", + new_callable=MagicMock, + return_value=False, + ), + patch.object( + platform._client, + "close", + new_callable=AsyncMock, + side_effect=RuntimeError("close failed"), + ), + pytest.raises(RuntimeError, match="close failed"), + ): + await platform.quiesce() + + assert platform._start_task is None + limiter.shutdown.assert_not_awaited() + assert platform.is_connected is False + + await platform.close() + + limiter.shutdown.assert_awaited_once_with() + + @pytest.mark.asyncio + async def test_start_propagates_client_failure_immediately(self): + limiter = _limiter_mock() + platform = _discord_runtime(bot_token="token", limiter=limiter) + + with ( + patch.object( + platform._client, + "start", + new_callable=AsyncMock, + side_effect=RuntimeError("invalid token"), + ), + pytest.raises(RuntimeError, match="invalid token"), + ): + await platform.start() + + assert platform._start_task is not None + assert platform._start_task.done() + await platform.quiesce() + await platform.close() + + @pytest.mark.asyncio + async def test_quiesce_waits_for_tracked_inbound_handler(self): + platform = _discord_runtime(bot_token="token") + platform._accepting_messages = True + entered = asyncio.Event() + release = asyncio.Event() + + async def handle(_message) -> None: + entered.set() + await release.wait() + + with ( + patch.object(platform, "_on_discord_message", side_effect=handle), + patch.object( + platform._client, + "is_closed", + new_callable=MagicMock, + return_value=True, + ), + ): + handler_task = asyncio.create_task( + platform._handle_client_message(MagicMock()) + ) + await entered.wait() + quiesce_task = asyncio.create_task(platform.quiesce()) + await asyncio.sleep(0) + + assert not quiesce_task.done() + release.set() + await handler_task + await quiesce_task + + assert platform._inbound_tasks == set() + await platform.close() + + @pytest.mark.asyncio + async def test_quiesce_preserves_caller_cancellation(self): + platform = _discord_runtime(bot_token="token") + entered = asyncio.Event() + + async def close_client() -> None: + entered.set() + await asyncio.Event().wait() + + with ( + patch.object( + platform._client, + "is_closed", + new_callable=MagicMock, + return_value=False, + ), + patch.object(platform._client, "close", side_effect=close_client), + ): + quiesce_task = asyncio.create_task(platform.quiesce()) + await entered.wait() + quiesce_task.cancel() + with pytest.raises(asyncio.CancelledError): + await quiesce_task + + await platform.close() diff --git a/tests/messaging/test_limiter.py b/tests/messaging/test_limiter.py index 1ddcaabe31d7e679ec818a47a8b55e1f176bb66f..3ed42f69c995638a2630e16e934cc046cb48d1e1 100644 --- a/tests/messaging/test_limiter.py +++ b/tests/messaging/test_limiter.py @@ -1,6 +1,7 @@ import asyncio import contextlib import time +from collections.abc import Callable import pytest import pytest_asyncio @@ -12,27 +13,63 @@ class TestMessagingRateLimiter: """Tests for MessagingRateLimiter.""" @pytest_asyncio.fixture(autouse=True) - async def reset_limiter(self): - """Reset singleton before each test.""" - await MessagingRateLimiter.shutdown_instance(timeout=0.1) + async def limiter_factory(self): + """Build started limiters and stop every instance after the test.""" + instances: list[MessagingRateLimiter] = [] + + def create( + *, rate_limit: int = 1, rate_window: float = 1.0 + ) -> MessagingRateLimiter: + limiter = MessagingRateLimiter( + rate_limit=rate_limit, + rate_window=rate_window, + ) + limiter.start() + instances.append(limiter) + return limiter + self.create_limiter: Callable[..., MessagingRateLimiter] = create yield - - await MessagingRateLimiter.shutdown_instance(timeout=0.1) + for limiter in reversed(instances): + await limiter.shutdown(timeout=0.1) @pytest.mark.asyncio - async def test_singleton_pattern(self): - """Test that get_instance returns the same object.""" - limiter1 = await MessagingRateLimiter.get_instance( - rate_limit=1, rate_window=0.5 - ) - limiter2 = await MessagingRateLimiter.get_instance( - rate_limit=99, rate_window=99.0 - ) - assert limiter1 is limiter2 - # First-construction wins for rate parameters + async def test_instances_are_independent(self): + """Each messaging runtime receives independent limiter state.""" + limiter1 = self.create_limiter(rate_limit=1, rate_window=0.5) + limiter2 = self.create_limiter(rate_limit=99, rate_window=99.0) + + assert limiter1 is not limiter2 assert limiter1.limiter._rate_limit == 1 assert limiter1.limiter._rate_window == 0.5 + assert limiter2.limiter._rate_limit == 99 + assert limiter2.limiter._rate_window == 99.0 + + await limiter1.shutdown(timeout=0.1) + + async def succeed() -> str: + return "still running" + + assert await limiter2.enqueue(succeed) == "still running" + + @pytest.mark.asyncio + async def test_start_is_required_and_shutdown_is_idempotent(self): + limiter = MessagingRateLimiter(rate_limit=1, rate_window=1.0) + + async def succeed() -> str: + return "ok" + + with pytest.raises(RuntimeError, match="has not been started"): + await limiter.enqueue(succeed) + + limiter.start() + limiter.start() + assert await limiter.enqueue(succeed) == "ok" + await limiter.shutdown(timeout=0.1) + await limiter.shutdown(timeout=0.1) + + with pytest.raises(RuntimeError, match="is closed"): + await limiter.enqueue(succeed) @pytest.mark.asyncio async def test_compaction(self): @@ -40,8 +77,7 @@ class TestMessagingRateLimiter: Verify multiple rapid requests with same dedup_key are compacted. Logic ported from verify_limiter.py """ - await MessagingRateLimiter.shutdown_instance(timeout=0.1) - limiter = await MessagingRateLimiter.get_instance(rate_limit=1, rate_window=1.0) + limiter = self.create_limiter(rate_limit=1, rate_window=1.0) call_counts = {} @@ -71,8 +107,7 @@ class TestMessagingRateLimiter: Verify that even when compacted, all futures resolve to the result of the LAST execution. Logic ported from verify_limiter_v2.py """ - await MessagingRateLimiter.shutdown_instance(timeout=0.1) - limiter = await MessagingRateLimiter.get_instance(rate_limit=1, rate_window=0.5) + limiter = self.create_limiter(rate_limit=1, rate_window=0.5) call_counts = {} msg_id = "test_msg_hang" @@ -107,8 +142,7 @@ class TestMessagingRateLimiter: @pytest.mark.asyncio async def test_flood_wait_handling(self): """Test that FloodWait exceptions pause the worker.""" - await MessagingRateLimiter.shutdown_instance(timeout=0.1) - limiter = await MessagingRateLimiter.get_instance(rate_limit=1, rate_window=1.0) + limiter = self.create_limiter(rate_limit=1, rate_window=1.0) # Mock exception with .seconds attribute class FloodWait(Exception): @@ -148,8 +182,7 @@ class TestMessagingRateLimiter: @pytest.mark.asyncio async def test_flood_wait_retry_after_parsing(self): """Error message with 'retry after N' parses the wait seconds.""" - await MessagingRateLimiter.shutdown_instance(timeout=0.1) - limiter = await MessagingRateLimiter.get_instance(rate_limit=1, rate_window=1.0) + limiter = self.create_limiter(rate_limit=1, rate_window=1.0) async def mock_flood(): raise Exception("Flood wait: retry after 2 seconds") @@ -163,8 +196,7 @@ class TestMessagingRateLimiter: @pytest.mark.asyncio async def test_non_flood_exception_no_pause(self): """Non-flood exception doesn't trigger pause.""" - await MessagingRateLimiter.shutdown_instance(timeout=0.1) - limiter = await MessagingRateLimiter.get_instance(rate_limit=1, rate_window=1.0) + limiter = self.create_limiter(rate_limit=1, rate_window=1.0) async def mock_error(): raise ValueError("some regular error") @@ -178,8 +210,7 @@ class TestMessagingRateLimiter: @pytest.mark.asyncio async def test_flood_with_seconds_attribute(self): """Exception with .seconds attribute uses that value for pause.""" - await MessagingRateLimiter.shutdown_instance(timeout=0.1) - limiter = await MessagingRateLimiter.get_instance(rate_limit=1, rate_window=1.0) + limiter = self.create_limiter(rate_limit=1, rate_window=1.0) class FloodWaitCustom(Exception): def __init__(self): @@ -200,8 +231,7 @@ class TestMessagingRateLimiter: Proactive limiter should enforce a strict sliding window: for any i, t[i+rate_limit] - t[i] >= rate_window (within tolerance). """ - await MessagingRateLimiter.shutdown_instance(timeout=0.1) - limiter = await MessagingRateLimiter.get_instance(rate_limit=2, rate_window=0.5) + limiter = self.create_limiter(rate_limit=2, rate_window=0.5) async def acquire(i: int) -> float: async def _do() -> float: @@ -224,8 +254,7 @@ class TestMessagingRateLimiter: @pytest.mark.asyncio async def test_compaction_last_task_fails_all_futures_get_exception(self): """When compacted task's last func fails, all futures get the exception.""" - await MessagingRateLimiter.shutdown_instance(timeout=0.1) - limiter = await MessagingRateLimiter.get_instance(rate_limit=1, rate_window=1.0) + limiter = self.create_limiter(rate_limit=1, rate_window=1.0) async def ok_task(): return "ok" @@ -244,8 +273,7 @@ class TestMessagingRateLimiter: @pytest.mark.asyncio async def test_fire_and_forget_failure_logged(self, caplog): """fire_and_forget with failing task logs error and does not re-raise.""" - await MessagingRateLimiter.shutdown_instance(timeout=0.1) - limiter = await MessagingRateLimiter.get_instance(rate_limit=1, rate_window=1.0) + limiter = self.create_limiter(rate_limit=1, rate_window=1.0) async def fail_task(): raise ValueError("fire_and_forget failed") @@ -256,3 +284,148 @@ class TestMessagingRateLimiter: joined = " ".join(str(r.message) for r in caplog.records) assert "ValueError" in joined assert "fire_and_forget failed" not in joined + + @pytest.mark.asyncio + async def test_shutdown_settles_active_queued_and_background_work(self): + limiter = self.create_limiter(rate_limit=1, rate_window=60.0) + active_started = asyncio.Event() + never_finish = asyncio.Event() + + async def active_operation() -> None: + active_started.set() + await never_finish.wait() + + active = asyncio.create_task( + limiter.enqueue(active_operation, dedup_key="active") + ) + await active_started.wait() + + async def queued_operation() -> None: + await never_finish.wait() + + queued = asyncio.create_task( + limiter.enqueue(queued_operation, dedup_key="queued") + ) + limiter.fire_and_forget(queued_operation, dedup_key="background") + await asyncio.sleep(0) + + await limiter.shutdown(timeout=0.1) + results = await asyncio.gather(active, queued, return_exceptions=True) + + assert all(isinstance(result, asyncio.CancelledError) for result in results) + assert limiter._background_tasks == set() + assert limiter._queue_map == {} + assert not limiter._queue_list + + @pytest.mark.asyncio + async def test_enqueue_cannot_enter_after_shutdown_begins(self): + limiter = self.create_limiter(rate_limit=1, rate_window=1.0) + await limiter._condition.acquire() + + async def succeed() -> str: + return "unexpected" + + enqueue_task = asyncio.create_task( + limiter.enqueue(succeed, dedup_key="shutdown-race") + ) + await asyncio.sleep(0) + shutdown_task = asyncio.create_task(limiter.shutdown(timeout=0.1)) + await asyncio.sleep(0) + assert limiter._closed is True + + limiter._condition.release() + await shutdown_task + + with pytest.raises(RuntimeError, match="is closed"): + await enqueue_task + + @pytest.mark.asyncio + async def test_shutdown_preserves_external_cancellation(self): + limiter = self.create_limiter(rate_limit=1, rate_window=1.0) + release = asyncio.Event() + + async def cancellation_resistant_worker() -> None: + try: + await release.wait() + except asyncio.CancelledError: + await release.wait() + + worker_task = limiter._worker_task + assert worker_task is not None + await asyncio.sleep(0) + worker_task.cancel() + await asyncio.gather(worker_task, return_exceptions=True) + limiter._worker_task = asyncio.create_task(cancellation_resistant_worker()) + shutdown_task = asyncio.create_task(limiter.shutdown(timeout=1.0)) + await asyncio.sleep(0) + + shutdown_task.cancel() + release.set() + + with pytest.raises(asyncio.CancelledError): + await shutdown_task + + @pytest.mark.asyncio + async def test_cancelled_shutdown_retries_queued_future_settlement(self): + limiter = self.create_limiter(rate_limit=1, rate_window=60.0) + active_started = asyncio.Event() + never_finish = asyncio.Event() + + async def active_operation() -> None: + active_started.set() + await never_finish.wait() + + active = asyncio.create_task( + limiter.enqueue(active_operation, dedup_key="active") + ) + await active_started.wait() + + async def queued_operation() -> None: + await never_finish.wait() + + queued = asyncio.create_task( + limiter.enqueue(queued_operation, dedup_key="queued") + ) + while "queued" not in limiter._queue_map: + await asyncio.sleep(0) + + await limiter._condition.acquire() + shutdown_task = asyncio.create_task(limiter.shutdown()) + await asyncio.sleep(0) + assert limiter._closed is True + + shutdown_task.cancel() + with pytest.raises(asyncio.CancelledError): + await shutdown_task + limiter._condition.release() + + await limiter.shutdown(timeout=0.1) + results = await asyncio.gather(active, queued, return_exceptions=True) + + assert all(isinstance(result, asyncio.CancelledError) for result in results) + assert limiter._queue_map == {} + assert not limiter._queue_list + assert limiter._worker_task is None + + @pytest.mark.asyncio + async def test_cancelled_operation_does_not_stop_owned_worker(self): + limiter = self.create_limiter(rate_limit=2, rate_window=1.0) + + async def cancelled_operation() -> None: + raise asyncio.CancelledError + + async def successful_operation() -> str: + return "delivered" + + first = asyncio.create_task( + limiter.enqueue(cancelled_operation, dedup_key="cancelled") + ) + second = asyncio.create_task( + limiter.enqueue(successful_operation, dedup_key="next") + ) + results = await asyncio.gather(first, second, return_exceptions=True) + + assert isinstance(results[0], asyncio.CancelledError) + assert results[1] == "delivered" + assert limiter._worker_task is not None + assert not limiter._worker_task.done() diff --git a/tests/messaging/test_messaging.py b/tests/messaging/test_messaging.py index 2d94bbd2725dbbc97daa2a824196ca8a3956d6eb..1d5513650ee625173abb2856557b5402876b3527 100644 --- a/tests/messaging/test_messaging.py +++ b/tests/messaging/test_messaging.py @@ -55,7 +55,8 @@ class TestMessagingPorts: runtime = MagicMock() runtime.name = "telegram" runtime.start = AsyncMock() - runtime.stop = AsyncMock() + runtime.quiesce = AsyncMock() + runtime.close = AsyncMock() runtime.on_message = MagicMock() outbound = MagicMock() outbound.queue_send_message = AsyncMock() diff --git a/tests/messaging/test_messaging_factory.py b/tests/messaging/test_messaging_factory.py index b183b5d929c4acddd310600ecfeb7ca72f716817..e23600bd3fd46daf5932dd6e3e07e78d91d6c930 100644 --- a/tests/messaging/test_messaging_factory.py +++ b/tests/messaging/test_messaging_factory.py @@ -16,7 +16,13 @@ class TestCreateMessagingComponents: mock_runtime = MagicMock() mock_runtime.name = "telegram" mock_runtime.outbound = MagicMock() + limiter = MagicMock() + transcriber = MagicMock() with ( + patch( + "free_claude_code.messaging.platforms.factory.MessagingRateLimiter", + return_value=limiter, + ) as limiter_cls, patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ), @@ -31,9 +37,9 @@ class TestCreateMessagingComponents: telegram_bot_token="test_token", allowed_telegram_user_id="12345", telegram_proxy_url="socks5://127.0.0.1:1080", - voice_note_enabled=False, - whisper_model="large-v3", - whisper_device="cuda", + transcriber=transcriber, + messaging_rate_limit=7, + messaging_rate_window=2.5, ), ) @@ -41,17 +47,17 @@ class TestCreateMessagingComponents: assert result.runtime is mock_runtime assert result.outbound is mock_runtime.outbound assert result.voice_cancellation is mock_runtime + limiter_cls.assert_called_once_with( + rate_limit=7, + rate_window=2.5, + log_error_details=False, + ) runtime_cls.assert_called_once_with( bot_token="test_token", allowed_user_id="12345", telegram_proxy_url="socks5://127.0.0.1:1080", - voice_note_enabled=False, - whisper_model="large-v3", - whisper_device="cuda", - huggingface_api_key="", - nvidia_nim_api_key="", - messaging_rate_limit=1, - messaging_rate_window=1.0, + limiter=limiter, + transcriber=transcriber, log_raw_messaging_content=False, log_api_error_tracebacks=False, ) @@ -73,7 +79,13 @@ class TestCreateMessagingComponents: mock_runtime = MagicMock() mock_runtime.name = "discord" mock_runtime.outbound = MagicMock() + limiter = MagicMock() + transcriber = MagicMock() with ( + patch( + "free_claude_code.messaging.platforms.factory.MessagingRateLimiter", + return_value=limiter, + ) as limiter_cls, patch( "free_claude_code.messaging.platforms.discord.DISCORD_AVAILABLE", True ), @@ -87,9 +99,9 @@ class TestCreateMessagingComponents: MessagingPlatformOptions( discord_bot_token="test_token", allowed_discord_channels="123,456", - voice_note_enabled=False, - whisper_model="small", - whisper_device="nvidia_nim", + transcriber=transcriber, + messaging_rate_limit=3, + messaging_rate_window=4.5, ), ) @@ -97,16 +109,16 @@ class TestCreateMessagingComponents: assert result.runtime is mock_runtime assert result.outbound is mock_runtime.outbound assert result.voice_cancellation is mock_runtime + limiter_cls.assert_called_once_with( + rate_limit=3, + rate_window=4.5, + log_error_details=False, + ) runtime_cls.assert_called_once_with( bot_token="test_token", allowed_channel_ids="123,456", - voice_note_enabled=False, - whisper_model="small", - whisper_device="nvidia_nim", - huggingface_api_key="", - nvidia_nim_api_key="", - messaging_rate_limit=1, - messaging_rate_window=1.0, + limiter=limiter, + transcriber=transcriber, log_raw_messaging_content=False, log_api_error_tracebacks=False, ) @@ -138,3 +150,34 @@ class TestCreateMessagingComponents: "slack", MessagingPlatformOptions(telegram_bot_token="token") ) assert result is None + + def test_separate_factory_calls_construct_distinct_limiters(self): + """Each selected platform runtime owns a new limiter instance.""" + runtime = MagicMock(name="runtime") + runtime.name = "telegram" + runtime.outbound = MagicMock() + with ( + patch( + "free_claude_code.messaging.platforms.telegram.TelegramRuntime", + return_value=runtime, + ) as runtime_cls, + patch( + "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True + ), + ): + first = create_messaging_components( + "telegram", + MessagingPlatformOptions(telegram_bot_token="one"), + ) + second = create_messaging_components( + "telegram", + MessagingPlatformOptions(telegram_bot_token="two"), + ) + + assert first is not None + assert second is not None + first_limiter = runtime_cls.call_args_list[0].kwargs["limiter"] + second_limiter = runtime_cls.call_args_list[1].kwargs["limiter"] + assert first_limiter is not second_limiter + assert runtime_cls.call_args_list[0].kwargs["transcriber"] is None + assert runtime_cls.call_args_list[1].kwargs["transcriber"] is None diff --git a/tests/messaging/test_platform_outbox.py b/tests/messaging/test_platform_outbox.py index efc03a778649dc094a0a8a75f25ba518464e5e98..c582cdb29008c333adc93408df6b3367d34af826 100644 --- a/tests/messaging/test_platform_outbox.py +++ b/tests/messaging/test_platform_outbox.py @@ -1,4 +1,5 @@ -from unittest.mock import AsyncMock, MagicMock +import asyncio +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -27,7 +28,7 @@ def _noop_outbox(*, limiter=None, delete_many=None) -> PlatformOutbox: return None return PlatformOutbox( - get_limiter=lambda: limiter, + limiter=limiter or MagicMock(), send=send, edit=edit, delete_many=delete_many or default_delete_many, @@ -35,8 +36,14 @@ def _noop_outbox(*, limiter=None, delete_many=None) -> PlatformOutbox: @pytest.mark.asyncio -async def test_queue_send_without_limiter_calls_raw_send() -> None: - outbox = _noop_outbox() +async def test_queue_send_awaits_required_limiter() -> None: + limiter = MagicMock() + + async def enqueue(operation, dedup_key=None): + return await operation() + + limiter.enqueue = AsyncMock(side_effect=enqueue) + outbox = _noop_outbox(limiter=limiter) result = await outbox.queue_send_message( "chat", @@ -48,6 +55,7 @@ async def test_queue_send_without_limiter_calls_raw_send() -> None: ) assert result == "chat:hello:reply:MarkdownV2:thread" + limiter.enqueue.assert_awaited_once() @pytest.mark.asyncio @@ -111,3 +119,68 @@ async def test_queue_delete_many_snapshots_ids_before_queueing() -> None: await operation() assert deleted == [["1", "2"]] + + +@pytest.mark.asyncio +async def test_close_cancels_and_settles_owned_background_work() -> None: + outbox = _noop_outbox() + started = asyncio.Event() + cancelled = asyncio.Event() + + async def pending() -> None: + started.set() + try: + await asyncio.Event().wait() + finally: + cancelled.set() + + outbox.fire_and_forget(pending()) + await started.wait() + + await outbox.close() + + assert cancelled.is_set() + assert outbox._background_tasks == set() + + +@pytest.mark.asyncio +async def test_completed_background_failure_is_observed_and_released() -> None: + outbox = _noop_outbox() + + async def fail() -> None: + raise RuntimeError("background failed") + + with patch("free_claude_code.messaging.platforms.outbox.logger.error") as error_log: + outbox.fire_and_forget(fail()) + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert outbox._background_tasks == set() + error_log.assert_called_once_with( + "Outbound background task failed: exc_type={}", + "RuntimeError", + ) + + +@pytest.mark.asyncio +async def test_close_rejects_all_later_work() -> None: + limiter = MagicMock() + outbox = _noop_outbox(limiter=limiter) + await outbox.close() + + with pytest.raises(RuntimeError, match="outbox is closed"): + await outbox.queue_send_message("chat", "message") + + ran = False + + async def late_task() -> None: + nonlocal ran + ran = True + + with pytest.raises(RuntimeError, match="outbox is closed"): + outbox.fire_and_forget(late_task()) + await asyncio.sleep(0) + + assert ran is False + assert outbox._background_tasks == set() + limiter.fire_and_forget.assert_not_called() diff --git a/tests/messaging/test_platform_voice_flow.py b/tests/messaging/test_platform_voice_flow.py index af5bdbd0c2d4f7348426fc18910728d17a05acb4..dc5070fb96d593c0f735b164771f81b27b0eae93 100644 --- a/tests/messaging/test_platform_voice_flow.py +++ b/tests/messaging/test_platform_voice_flow.py @@ -1,3 +1,4 @@ +import asyncio from pathlib import Path from unittest.mock import AsyncMock @@ -11,17 +12,33 @@ from free_claude_code.messaging.platforms.voice_flow import ( audio_suffix_from_metadata, is_audio_metadata, ) +from free_claude_code.messaging.voice import Transcriber -def _flow(*, enabled: bool = True) -> VoiceNoteFlow: - return VoiceNoteFlow( - voice_note_enabled=enabled, - whisper_model="base", - whisper_device="cpu", - huggingface_api_key="", - nvidia_nim_api_key="", - log_raw_messaging_content=False, - log_api_error_tracebacks=False, +class MockTranscriber: + def __init__(self, result: str = "hello from voice") -> None: + self.run = AsyncMock(return_value=result) + self.close_run = AsyncMock() + self.paths: list[Path] = [] + + async def transcribe(self, file_path: Path) -> str: + self.paths.append(file_path) + return await self.run(file_path) + + async def close(self) -> None: + await self.close_run() + + +def _flow(*, enabled: bool = True) -> tuple[VoiceNoteFlow, MockTranscriber]: + transcriber = MockTranscriber() + configured: Transcriber | None = transcriber if enabled else None + return ( + VoiceNoteFlow( + transcriber=configured, + log_raw_messaging_content=False, + log_api_error_tracebacks=False, + ), + transcriber, ) @@ -52,10 +69,8 @@ def _request( @pytest.mark.asyncio -async def test_voice_flow_success_builds_incoming_message(monkeypatch) -> None: - flow = _flow() - transcribe = AsyncMock(return_value="hello from voice") - monkeypatch.setattr(flow._voice_transcription, "transcribe", transcribe) +async def test_voice_flow_success_builds_incoming_message() -> None: + flow, transcriber = _flow() handler = AsyncMock() queue_send = AsyncMock(return_value="status") queue_delete = AsyncMock() @@ -91,14 +106,14 @@ async def test_voice_flow_success_builds_incoming_message(monkeypatch) -> None: assert incoming.reply_to_message_id == "reply" assert incoming.message_thread_id == "thread" assert incoming.status_message_id == "status" + transcriber.run.assert_awaited_once() + assert transcriber.paths == downloaded_paths assert downloaded_paths and not downloaded_paths[0].exists() @pytest.mark.asyncio -async def test_voice_flow_disabled_replies_without_transcribing(monkeypatch) -> None: - flow = _flow(enabled=False) - transcribe = AsyncMock(return_value="should not run") - monkeypatch.setattr(flow._voice_transcription, "transcribe", transcribe) +async def test_voice_flow_disabled_replies_without_transcribing() -> None: + flow, transcriber = _flow(enabled=False) reply_text = AsyncMock() handled = await flow.handle( @@ -110,22 +125,18 @@ async def test_voice_flow_disabled_replies_without_transcribing(monkeypatch) -> assert handled is True reply_text.assert_awaited_once_with(VOICE_DISABLED_MESSAGE) - transcribe.assert_not_awaited() + transcriber.run.assert_not_awaited() @pytest.mark.asyncio -async def test_voice_flow_cancelled_transcription_deletes_status(monkeypatch) -> None: - flow = _flow() +async def test_voice_flow_cancelled_transcription_deletes_status() -> None: + flow, transcriber = _flow() - async def canceling_transcribe(*args, **kwargs) -> str: + async def canceling_transcribe(_path: Path) -> str: await flow.cancel_pending_voice("chat", "voice") return "ignored" - monkeypatch.setattr( - flow._voice_transcription, - "transcribe", - AsyncMock(side_effect=canceling_transcribe), - ) + transcriber.run.side_effect = canceling_transcribe handler = AsyncMock() queue_send = AsyncMock(return_value="status") queue_delete = AsyncMock() @@ -143,10 +154,57 @@ async def test_voice_flow_cancelled_transcription_deletes_status(monkeypatch) -> @pytest.mark.asyncio -async def test_voice_flow_download_failure_cleans_pending_state(monkeypatch) -> None: - flow = _flow() - transcribe = AsyncMock(return_value="should not run") - monkeypatch.setattr(flow._voice_transcription, "transcribe", transcribe) +async def test_voice_flow_task_cancellation_waits_then_cleans_pending_state() -> None: + flow, transcriber = _flow() + started = asyncio.Event() + cancellation_received = asyncio.Event() + release = asyncio.Event() + stopped = asyncio.Event() + + async def cancellation_safe_transcribe(_path: Path) -> str: + started.set() + try: + await asyncio.Event().wait() + return "unreachable" + except asyncio.CancelledError: + cancellation_received.set() + await release.wait() + stopped.set() + raise + + transcriber.run.side_effect = cancellation_safe_transcribe + handler = AsyncMock() + queue_delete = AsyncMock() + handle_task = asyncio.create_task( + flow.handle( + _request(), + message_handler=handler, + queue_send_message=AsyncMock(return_value="status"), + queue_delete_messages=queue_delete, + ) + ) + + await started.wait() + handle_task.cancel() + await cancellation_received.wait() + + assert not handle_task.done() + assert await flow.is_voice_still_pending("chat", "voice") is True + queue_delete.assert_not_awaited() + + release.set() + with pytest.raises(asyncio.CancelledError): + await handle_task + + assert stopped.is_set() + handler.assert_not_awaited() + queue_delete.assert_awaited_once_with("chat", ["status"]) + assert await flow.cancel_pending_voice("chat", "voice") is None + + +@pytest.mark.asyncio +async def test_voice_flow_download_failure_cleans_pending_state() -> None: + flow, transcriber = _flow() reply_text = AsyncMock() queue_delete = AsyncMock() @@ -161,22 +219,16 @@ async def test_voice_flow_download_failure_cleans_pending_state(monkeypatch) -> ) assert handled is True - transcribe.assert_not_awaited() + transcriber.run.assert_not_awaited() queue_delete.assert_awaited_once_with("chat", ["status"]) reply_text.assert_awaited_once_with(VOICE_TRANSCRIPTION_ERROR_MESSAGE) assert await flow.cancel_pending_voice("chat", "voice") is None @pytest.mark.asyncio -async def test_voice_flow_transcription_failure_cleans_pending_state( - monkeypatch, -) -> None: - flow = _flow() - monkeypatch.setattr( - flow._voice_transcription, - "transcribe", - AsyncMock(side_effect=RuntimeError("transcription failed")), - ) +async def test_voice_flow_transcription_failure_cleans_pending_state() -> None: + flow, transcriber = _flow() + transcriber.run.side_effect = RuntimeError("transcription failed") reply_text = AsyncMock() queue_delete = AsyncMock() @@ -194,15 +246,10 @@ async def test_voice_flow_transcription_failure_cleans_pending_state( @pytest.mark.asyncio -async def test_voice_flow_handler_failure_cleans_pending_without_deleting_status( - monkeypatch, -) -> None: - flow = _flow() - monkeypatch.setattr( - flow._voice_transcription, - "transcribe", - AsyncMock(return_value="hello from voice"), - ) +async def test_voice_flow_handler_failure_cleans_pending_without_deleting_status() -> ( + None +): + flow, _transcriber = _flow() reply_text = AsyncMock() queue_delete = AsyncMock() @@ -222,6 +269,35 @@ async def test_voice_flow_handler_failure_cleans_pending_without_deleting_status assert await flow.cancel_pending_voice("chat", "voice") is None +@pytest.mark.asyncio +async def test_voice_flow_rejects_oversized_audio_before_transcription( + monkeypatch, +) -> None: + monkeypatch.setattr( + "free_claude_code.messaging.platforms.voice_flow.MAX_AUDIO_SIZE_BYTES", + 3, + ) + flow, transcriber = _flow() + reply_text = AsyncMock() + queue_delete = AsyncMock() + + async def download(path: Path) -> None: + path.write_bytes(b"four") + + handled = await flow.handle( + _request(download_to=download, reply_text=reply_text), + message_handler=AsyncMock(), + queue_send_message=AsyncMock(return_value="status"), + queue_delete_messages=queue_delete, + ) + + assert handled is True + transcriber.run.assert_not_awaited() + queue_delete.assert_awaited_once_with("chat", ["status"]) + assert reply_text.await_args is not None + assert "too large" in reply_text.await_args.args[0] + + def test_audio_metadata_helpers() -> None: assert is_audio_metadata("voice.ogg", "application/octet-stream") is True assert is_audio_metadata("file.txt", "audio/ogg") is True diff --git a/tests/messaging/test_reliability.py b/tests/messaging/test_reliability.py index eb076a0bc152771f8246cffa2e55b83a8708b815..8cbc544b3367a38ecc57534f98f9c76659e79ca5 100644 --- a/tests/messaging/test_reliability.py +++ b/tests/messaging/test_reliability.py @@ -3,6 +3,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from telegram.error import NetworkError, RetryAfter, TelegramError +from free_claude_code.messaging.limiter import MessagingRateLimiter from free_claude_code.messaging.platforms.telegram import TelegramRuntime @@ -11,7 +12,12 @@ def telegram_platform(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - platform = TelegramRuntime(bot_token="test_token", allowed_user_id="12345") + platform = TelegramRuntime( + bot_token="test_token", + allowed_user_id="12345", + limiter=MessagingRateLimiter(rate_limit=1, rate_window=1.0), + transcriber=None, + ) return platform diff --git a/tests/messaging/test_telegram.py b/tests/messaging/test_telegram.py index 0ba272f584e0073eb897546d40832893cabfac97..f5b4c73ea68a9a9d01c9d6bf307f1b0d13599ef4 100644 --- a/tests/messaging/test_telegram.py +++ b/tests/messaging/test_telegram.py @@ -6,18 +6,36 @@ from telegram.error import TelegramError from free_claude_code.messaging.platforms.telegram import TelegramRuntime +def _limiter_mock() -> MagicMock: + limiter = MagicMock() + limiter.start = MagicMock() + limiter.shutdown = AsyncMock() + return limiter + + +def _telegram_runtime( + *args, limiter=None, transcriber=None, **kwargs +) -> TelegramRuntime: + return TelegramRuntime( + *args, + limiter=limiter or _limiter_mock(), + transcriber=transcriber, + **kwargs, + ) + + @pytest.fixture def telegram_platform(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - platform = TelegramRuntime(bot_token="test_token", allowed_user_id="12345") + platform = _telegram_runtime(bot_token="test_token", allowed_user_id="12345") return platform def test_telegram_platform_init_no_token(): with patch.dict("os.environ", {}, clear=True): - platform = TelegramRuntime(bot_token=None) + platform = _telegram_runtime(bot_token=None) assert platform.bot_token is None @@ -31,27 +49,25 @@ async def test_telegram_platform_start_success(telegram_platform): mock_builder.return_value.token.return_value.request.return_value.build.return_value = mock_app - # Mock MessagingRateLimiter - with patch( - "free_claude_code.messaging.limiter.MessagingRateLimiter.get_instance", - AsyncMock(), - ): - await telegram_platform.start() + await telegram_platform.start() - assert telegram_platform._connected is True - mock_app.initialize.assert_called_once() - mock_app.start.assert_called_once() + assert telegram_platform._connected is True + mock_app.initialize.assert_called_once() + mock_app.start.assert_called_once() + telegram_platform._limiter.start.assert_called_once_with() @pytest.mark.asyncio async def test_telegram_platform_start_with_proxy(): + limiter = _limiter_mock() with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - platform = TelegramRuntime( + platform = _telegram_runtime( bot_token="test_token", allowed_user_id="12345", telegram_proxy_url="socks5://127.0.0.1:1080", + limiter=limiter, ) with ( @@ -74,11 +90,7 @@ async def test_telegram_platform_start_with_proxy(): update_request = MagicMock() request_cls.side_effect = [request, update_request] - with patch( - "free_claude_code.messaging.limiter.MessagingRateLimiter.get_instance", - AsyncMock(), - ): - await platform.start() + await platform.start() assert request_cls.call_count == 2 request_cls.assert_any_call( @@ -90,6 +102,7 @@ async def test_telegram_platform_start_with_proxy(): builder.request.assert_called_once_with(request) builder.get_updates_request.assert_called_once_with(update_request) assert platform._connected is True + limiter.start.assert_called_once_with() @pytest.mark.asyncio @@ -254,8 +267,8 @@ async def test_telegram_platform_single_delete_still_swallows_known_error( @pytest.mark.asyncio async def test_telegram_platform_queue_send_message(telegram_platform): - mock_limiter = AsyncMock() - telegram_platform._limiter = mock_limiter + mock_limiter = telegram_platform._limiter + mock_limiter.enqueue = AsyncMock() await telegram_platform.outbound.queue_send_message( "chat_1", "hello", fire_and_forget=False diff --git a/tests/messaging/test_telegram_edge_cases.py b/tests/messaging/test_telegram_edge_cases.py index 0dad7b9b9570d58ee46602b20c650a5f0ff944da..2da9fe8b4deee69f96402de8a7e21b1a4412de75 100644 --- a/tests/messaging/test_telegram_edge_cases.py +++ b/tests/messaging/test_telegram_edge_cases.py @@ -6,14 +6,32 @@ import pytest from telegram.error import NetworkError, RetryAfter, TelegramError +def _limiter_mock() -> MagicMock: + limiter = MagicMock() + limiter.start = MagicMock() + limiter.shutdown = AsyncMock() + return limiter + + +def _telegram_runtime(*args, limiter=None, transcriber=None, **kwargs): + from free_claude_code.messaging.platforms.telegram import TelegramRuntime + + return TelegramRuntime( + *args, + limiter=limiter or _limiter_mock(), + transcriber=transcriber, + **kwargs, + ) + + def test_telegram_platform_init_raises_when_dependency_missing(): - with patch( - "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", False + with ( + patch( + "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", False + ), + pytest.raises(ImportError), ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - with pytest.raises(ImportError): - TelegramRuntime(bot_token="x") + _telegram_runtime(bot_token="x") @pytest.mark.asyncio @@ -22,35 +40,179 @@ async def test_telegram_platform_start_requires_token(): patch.dict("os.environ", {}, clear=True), patch("free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True), ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token=None) + platform = _telegram_runtime(bot_token=None) with pytest.raises(ValueError): await platform.start() @pytest.mark.asyncio -async def test_telegram_platform_stop_no_application_is_noop(): +async def test_telegram_platform_quiesce_and_close_without_application(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t") + platform = _telegram_runtime(bot_token="t") platform._application = None platform._connected = True - await platform.stop() + await platform.quiesce() + await platform.close() assert platform.is_connected is False + platform._limiter.shutdown.assert_awaited_once_with() @pytest.mark.asyncio -async def test_with_retry_returns_none_when_message_not_modified_network_error(): +async def test_telegram_close_cleans_up_partially_initialized_application(): + with patch( + "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True + ): + platform = _telegram_runtime(bot_token="t") + platform._application = MagicMock() + platform._application.running = False + platform._application.updater.running = False + platform._application.updater.stop = AsyncMock() + platform._application.stop = AsyncMock() + platform._application.shutdown = AsyncMock() + platform.outbound.close = AsyncMock() + + await platform.quiesce() + await platform.close() + + platform._application.updater.stop.assert_not_awaited() + platform._application.stop.assert_not_awaited() + platform.outbound.close.assert_awaited_once_with() + platform._limiter.shutdown.assert_awaited_once_with() + platform._application.shutdown.assert_awaited_once_with() + + +@pytest.mark.asyncio +async def test_telegram_two_phase_lifecycle_drains_before_delivery_close(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime + platform = _telegram_runtime(bot_token="t") + order: list[str] = [] + platform._application = MagicMock() + platform._application.running = True + platform._application.updater.running = True + platform._application.updater.stop = AsyncMock( + side_effect=lambda: order.append("updater.stop") + ) + platform._application.stop = AsyncMock( + side_effect=lambda: order.append("application.stop") + ) + platform.outbound.close = AsyncMock( + side_effect=lambda: order.append("outbound.close") + ) + platform._limiter.shutdown = AsyncMock( + side_effect=lambda: order.append("limiter.shutdown") + ) + platform._application.shutdown = AsyncMock( + side_effect=lambda: order.append("application.shutdown") + ) - platform = TelegramRuntime(bot_token="t") + await platform.quiesce() + assert order == ["updater.stop", "application.stop"] + + await platform.close() + + assert order == [ + "updater.stop", + "application.stop", + "outbound.close", + "limiter.shutdown", + "application.shutdown", + ] + assert platform.is_connected is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "failing_step", + ["updater.stop", "application.stop"], +) +async def test_telegram_quiesce_attempts_all_steps_after_failure(failing_step): + with patch( + "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True + ): + platform = _telegram_runtime(bot_token="t") + order: list[str] = [] + + async def record(step: str) -> None: + order.append(step) + if step == failing_step: + raise RuntimeError(step) + + def action(step: str): + async def run() -> None: + await record(step) + + return run + + platform._application = MagicMock() + platform._application.running = True + platform._application.updater.running = True + platform._application.updater.stop = AsyncMock( + side_effect=action("updater.stop") + ) + platform._application.stop = AsyncMock(side_effect=action("application.stop")) + platform.outbound.close = AsyncMock(side_effect=action("outbound.close")) + platform._limiter.shutdown = AsyncMock(side_effect=action("limiter.shutdown")) + platform._application.shutdown = AsyncMock( + side_effect=action("application.shutdown") + ) + + with pytest.raises(RuntimeError, match=failing_step): + await platform.quiesce() + + assert order == ["updater.stop", "application.stop"] + assert platform.is_connected is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "failing_step", + ["outbound.close", "limiter.shutdown", "application.shutdown"], +) +async def test_telegram_close_attempts_all_steps_after_failure(failing_step): + with patch( + "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True + ): + platform = _telegram_runtime(bot_token="t") + order: list[str] = [] + + async def record(step: str) -> None: + order.append(step) + if step == failing_step: + raise RuntimeError(step) + + def action(step: str): + async def run() -> None: + await record(step) + + return run + + platform._application = MagicMock() + platform._application.shutdown = AsyncMock( + side_effect=action("application.shutdown") + ) + platform.outbound.close = AsyncMock(side_effect=action("outbound.close")) + platform._limiter.shutdown = AsyncMock(side_effect=action("limiter.shutdown")) + + with pytest.raises(RuntimeError, match=failing_step): + await platform.close() + + assert order == [ + "outbound.close", + "limiter.shutdown", + "application.shutdown", + ] + + +@pytest.mark.asyncio +async def test_with_retry_returns_none_when_message_not_modified_network_error(): + with patch( + "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True + ): + platform = _telegram_runtime(bot_token="t") async def _f(): raise NetworkError("Message is not modified") @@ -63,9 +225,7 @@ async def test_with_retry_retries_network_error_then_succeeds(monkeypatch): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t") + platform = _telegram_runtime(bot_token="t") monkeypatch.setattr(asyncio, "sleep", AsyncMock()) @@ -86,9 +246,7 @@ async def test_with_retry_honors_retry_after_timedelta(monkeypatch): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t") + platform = _telegram_runtime(bot_token="t") monkeypatch.setattr(asyncio, "sleep", AsyncMock()) @@ -109,9 +267,7 @@ async def test_with_retry_drops_parse_mode_on_markdown_entity_error(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t") + platform = _telegram_runtime(bot_token="t") calls = [] @@ -130,9 +286,7 @@ async def test_with_retry_can_raise_known_message_errors_for_bulk_fallback(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t") + platform = _telegram_runtime(bot_token="t") async def _f(): raise TelegramError("message can't be deleted") @@ -145,38 +299,45 @@ async def test_with_retry_can_raise_known_message_errors_for_bulk_fallback(): @pytest.mark.asyncio -async def test_queue_send_message_without_limiter_calls_send_message(): +async def test_queue_send_message_uses_required_limiter(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t") - platform._limiter = None + platform = _telegram_runtime(bot_token="t") platform._application = MagicMock() mock_msg = MagicMock() mock_msg.message_id = 1 platform._application.bot = AsyncMock() platform._application.bot.send_message = AsyncMock(return_value=mock_msg) - assert await platform.outbound.queue_send_message("c", "t") == "1" + async def enqueue(operation, dedup_key=None): + return await operation() + + platform._limiter.enqueue = AsyncMock(side_effect=enqueue) + assert ( + await platform.outbound.queue_send_message("c", "t", fire_and_forget=False) + == "1" + ) + platform._limiter.enqueue.assert_awaited_once() platform._application.bot.send_message.assert_awaited_once() @pytest.mark.asyncio -async def test_queue_edit_message_without_limiter_calls_edit_message(): +async def test_queue_edit_message_uses_required_limiter(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t") - platform._limiter = None + platform = _telegram_runtime(bot_token="t") platform._application = MagicMock() platform._application.bot = AsyncMock() platform._application.bot.edit_message_text = AsyncMock() - await platform.outbound.queue_edit_message("c", "1", "t") + async def enqueue(operation, dedup_key=None): + return await operation() + + platform._limiter.enqueue = AsyncMock(side_effect=enqueue) + await platform.outbound.queue_edit_message("c", "1", "t", fire_and_forget=False) + platform._limiter.enqueue.assert_awaited_once() platform._application.bot.edit_message_text.assert_awaited_once() @@ -184,9 +345,7 @@ def test_fire_and_forget_non_coroutine_uses_ensure_future(monkeypatch): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t") + platform = _telegram_runtime(bot_token="t") ef = MagicMock() monkeypatch.setattr(asyncio, "ensure_future", ef) @@ -200,9 +359,7 @@ async def test_on_start_command_replies_and_forwards(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t") + platform = _telegram_runtime(bot_token="t") with patch.object( platform, "_on_telegram_message", new_callable=AsyncMock ) as mock_msg: @@ -219,9 +376,7 @@ async def test_on_telegram_message_handler_error_sends_error_message(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t", allowed_user_id="123") + platform = _telegram_runtime(bot_token="t", allowed_user_id="123") with patch.object( platform.outbound, "send_message", new_callable=AsyncMock ) as mock_send: @@ -247,19 +402,11 @@ async def test_telegram_start_retries_on_network_error(monkeypatch): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="token", allowed_user_id=None) + platform = _telegram_runtime(bot_token="token", allowed_user_id=None) monkeypatch.setattr(asyncio, "sleep", AsyncMock()) - with ( - patch("telegram.ext.Application.builder") as mock_builder, - patch( - "free_claude_code.messaging.limiter.MessagingRateLimiter.get_instance", - AsyncMock(), - ), - ): + with patch("telegram.ext.Application.builder") as mock_builder: mock_app = MagicMock() mock_app.initialize = AsyncMock(side_effect=[NetworkError("no"), None]) mock_app.start = AsyncMock() @@ -269,6 +416,38 @@ async def test_telegram_start_retries_on_network_error(monkeypatch): await platform.start() assert platform.is_connected is True + assert mock_app.initialize.await_count == 2 + mock_app.start.assert_awaited_once_with() + platform._limiter.start.assert_called_once_with() + + +@pytest.mark.asyncio +async def test_telegram_polling_retry_does_not_restart_running_application( + monkeypatch, +): + with patch( + "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True + ): + platform = _telegram_runtime(bot_token="token", allowed_user_id=None) + + monkeypatch.setattr(asyncio, "sleep", AsyncMock()) + + with patch("telegram.ext.Application.builder") as mock_builder: + mock_app = MagicMock() + mock_app.initialize = AsyncMock() + mock_app.start = AsyncMock() + mock_app.updater.start_polling = AsyncMock( + side_effect=[NetworkError("temporary polling failure"), None] + ) + mock_builder.return_value.token.return_value.request.return_value.build.return_value = mock_app + + await platform.start() + + assert platform.is_connected is True + mock_app.initialize.assert_awaited_once_with() + mock_app.start.assert_awaited_once_with() + assert mock_app.updater.start_polling.await_count == 2 + platform._limiter.start.assert_called_once_with() @pytest.mark.asyncio @@ -277,9 +456,7 @@ async def test_edit_message_with_text_exceeding_4096_raises(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t") + platform = _telegram_runtime(bot_token="t") platform._application = MagicMock() platform._application.bot = AsyncMock() platform._application.bot.edit_message_text = AsyncMock( @@ -296,9 +473,7 @@ async def test_edit_message_empty_string(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t") + platform = _telegram_runtime(bot_token="t") platform._application = MagicMock() platform._application.bot = AsyncMock() platform._application.bot.edit_message_text = AsyncMock() @@ -315,9 +490,7 @@ async def test_send_message_empty_string(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t") + platform = _telegram_runtime(bot_token="t") platform._application = MagicMock() mock_msg = MagicMock() mock_msg.message_id = 1 @@ -335,9 +508,7 @@ async def test_on_telegram_message_non_text_update_ignored(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t", allowed_user_id="123") + platform = _telegram_runtime(bot_token="t", allowed_user_id="123") handler = AsyncMock() platform.on_message(handler) @@ -359,9 +530,7 @@ async def test_with_retry_message_not_found_returns_none(): with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - from free_claude_code.messaging.platforms.telegram import TelegramRuntime - - platform = TelegramRuntime(bot_token="t") + platform = _telegram_runtime(bot_token="t") async def _f(): raise TelegramError("message to edit not found") diff --git a/tests/messaging/test_transcription.py b/tests/messaging/test_transcription.py index 0542ecb99a3481fef3bcedfbf627a2d054a05015..2b8ebf32b37e03411e94885de328167698f47f7f 100644 --- a/tests/messaging/test_transcription.py +++ b/tests/messaging/test_transcription.py @@ -1,143 +1,262 @@ -"""Tests for voice note transcription.""" +"""Tests for the instance-owned local Whisper transcriber.""" -import tempfile +import asyncio +import threading +import time from pathlib import Path +from types import SimpleNamespace from unittest.mock import MagicMock, patch import pytest -from free_claude_code.messaging.transcription import ( - MAX_AUDIO_SIZE_BYTES, - transcribe_audio, -) - - -def test_transcribe_file_not_found_raises(): - """Non-existent file raises FileNotFoundError.""" - with pytest.raises(FileNotFoundError, match="not found"): - transcribe_audio(Path("/nonexistent/file.ogg"), "audio/ogg") - - -def test_transcribe_file_too_large_raises(): - """File exceeding max size raises ValueError.""" - with tempfile.NamedTemporaryFile(suffix=".ogg", delete=False) as f: - f.write(b"x" * (MAX_AUDIO_SIZE_BYTES + 1)) - path = Path(f.name) - try: - with pytest.raises(ValueError, match="too large"): - transcribe_audio(path, "audio/ogg", whisper_device="cpu") - finally: - path.unlink(missing_ok=True) - - -def test_transcribe_local_success(): - """Local backend returns transcribed text.""" - with tempfile.NamedTemporaryFile(suffix=".ogg", delete=False) as f: - f.write(b"fake ogg content") - path = Path(f.name) - try: - mock_pipe = MagicMock() - mock_pipe.return_value = {"text": "Hello world"} - fake_audio = {"array": [0.0], "sampling_rate": 16000} - - with ( - patch( - "free_claude_code.messaging.transcription._load_audio", - return_value=fake_audio, - ), - patch( - "free_claude_code.messaging.transcription._get_pipeline", - return_value=mock_pipe, - ), - ): - result = transcribe_audio(path, "audio/ogg", whisper_model="base") - - assert result == "Hello world" - mock_pipe.assert_called_once_with( - fake_audio, generate_kwargs={"language": "en", "task": "transcribe"} +from free_claude_code.messaging.transcription import TranscriptionService + + +def _service(*, api_key: str = "") -> TranscriptionService: + return TranscriptionService( + model="base", + device="cpu", + huggingface_api_key=api_key, + ) + + +def _fake_optional_modules( + pipeline: MagicMock, +) -> tuple[SimpleNamespace, SimpleNamespace, MagicMock, MagicMock]: + torch = SimpleNamespace( + cuda=SimpleNamespace(is_available=lambda: False), + float16=object(), + float32=object(), + ) + + model = MagicMock() + model.to.return_value = model + model_loader = MagicMock() + model_loader.from_pretrained.return_value = model + processor = SimpleNamespace(tokenizer=object(), feature_extractor=object()) + processor_loader = MagicMock() + processor_loader.from_pretrained.return_value = processor + + transformers = SimpleNamespace( + AutoModelForSpeechSeq2Seq=model_loader, + AutoProcessor=processor_loader, + pipeline=MagicMock(return_value=pipeline), + ) + return torch, transformers, model_loader, processor_loader + + +@pytest.mark.asyncio +async def test_transcription_service_transcribes_and_reuses_its_pipeline( + tmp_path: Path, +) -> None: + audio_path = tmp_path / "voice.ogg" + audio_path.write_bytes(b"voice") + pipeline = MagicMock(return_value={"text": " Hello world "}) + torch, transformers, model_loader, processor_loader = _fake_optional_modules( + pipeline + ) + fake_audio = {"array": [0.0], "sampling_rate": 16000} + service = _service(api_key="hf-provider-key") + + with ( + patch.dict( + "sys.modules", + {"torch": torch, "transformers": transformers}, + ), + patch( + "free_claude_code.messaging.transcription._load_audio", + return_value=fake_audio, + ), + ): + first = await service.transcribe(audio_path) + second = await service.transcribe(audio_path) + + assert first == "Hello world" + assert second == "Hello world" + assert transformers.pipeline.call_count == 1 + assert pipeline.call_count == 2 + model_loader.from_pretrained.assert_called_once_with( + "openai/whisper-base", + dtype=torch.float32, + low_cpu_mem_usage=True, + attn_implementation="sdpa", + token="hf-provider-key", + ) + processor_loader.from_pretrained.assert_called_once_with( + "openai/whisper-base", + token="hf-provider-key", + ) + + +@pytest.mark.asyncio +async def test_separate_services_do_not_share_pipeline_instances( + tmp_path: Path, +) -> None: + audio_path = tmp_path / "voice.wav" + audio_path.write_bytes(b"voice") + pipeline = MagicMock(return_value={"text": "ok"}) + torch, transformers, _model_loader, _processor_loader = _fake_optional_modules( + pipeline + ) + first = _service() + second = _service() + + with ( + patch.dict( + "sys.modules", + {"torch": torch, "transformers": transformers}, + ), + patch( + "free_claude_code.messaging.transcription._load_audio", + return_value={"array": [0.0], "sampling_rate": 16000}, + ), + ): + await first.transcribe(audio_path) + await second.transcribe(audio_path) + + assert transformers.pipeline.call_count == 2 + + +@pytest.mark.asyncio +async def test_transcription_service_serializes_concurrent_inference( + tmp_path: Path, +) -> None: + service = _service() + audio_path = tmp_path / "voice.wav" + audio_path.write_bytes(b"voice") + state_lock = threading.Lock() + active = 0 + max_active = 0 + + def transcribe_sync(_path: Path) -> str: + nonlocal active, max_active + with state_lock: + active += 1 + max_active = max(max_active, active) + time.sleep(0.02) + with state_lock: + active -= 1 + return "ok" + + with patch.object(service, "_transcribe_sync", side_effect=transcribe_sync): + results = await asyncio.gather( + service.transcribe(audio_path), + service.transcribe(audio_path), + service.transcribe(audio_path), ) - finally: - path.unlink(missing_ok=True) - - -def test_transcribe_local_empty_segments_returns_no_speech(): - """Local backend with no speech returns placeholder.""" - with tempfile.NamedTemporaryFile(suffix=".ogg", delete=False) as f: - f.write(b"fake ogg") - path = Path(f.name) - try: - mock_pipe = MagicMock() - mock_pipe.return_value = {"text": ""} - fake_audio = {"array": [0.0], "sampling_rate": 16000} - - with ( - patch( - "free_claude_code.messaging.transcription._load_audio", - return_value=fake_audio, - ), - patch( - "free_claude_code.messaging.transcription._get_pipeline", - return_value=mock_pipe, - ), - ): - result = transcribe_audio(path, "audio/ogg", whisper_model="base") - - assert result == "(no speech detected)" - finally: - path.unlink(missing_ok=True) - - -def test_transcribe_invalid_device_raises(): - """Invalid whisper_device raises ValueError.""" - with tempfile.NamedTemporaryFile(suffix=".ogg", delete=False) as f: - f.write(b"fake ogg") - path = Path(f.name) - try: - # Patch _load_audio to avoid ImportError from missing librosa - # Device validation happens in _get_pipeline before torch import - with ( - patch("free_claude_code.messaging.transcription._load_audio"), - pytest.raises(ValueError, match="whisper_device must be 'cpu' or 'cuda'"), - ): - transcribe_audio(path, "audio/ogg", whisper_device="auto") - finally: - path.unlink(missing_ok=True) - - -def test_transcribe_nim_requires_api_key(): - """NIM path rejects empty API key without reading global settings.""" - with tempfile.NamedTemporaryFile(suffix=".ogg", delete=False) as f: - f.write(b"fake ogg") - path = Path(f.name) - try: - with pytest.raises(ValueError, match="non-empty"): - transcribe_audio( - path, - "audio/ogg", - whisper_device="nvidia_nim", - whisper_model="openai/whisper-large-v3", - nvidia_nim_api_key="", - ) - finally: - path.unlink(missing_ok=True) - - -def test_transcribe_local_import_error_raises(): - """Local backend when voice_local extra not installed raises ImportError.""" - with tempfile.NamedTemporaryFile(suffix=".ogg", delete=False) as f: - f.write(b"fake ogg") - path = Path(f.name) - try: - with ( - patch( - "free_claude_code.messaging.transcription._get_pipeline", - side_effect=ImportError( - "Local Whisper requires the voice_local extra. " - "Install with: uv sync --extra voice_local" - ), - ), - pytest.raises(ImportError, match="voice_local extra"), - ): - transcribe_audio(path, "audio/ogg", whisper_device="cpu") - finally: - path.unlink(missing_ok=True) + + assert results == ["ok", "ok", "ok"] + assert max_active == 1 + + +@pytest.mark.asyncio +async def test_transcription_service_close_releases_pipeline_and_is_terminal( + tmp_path: Path, +) -> None: + service = _service() + audio_path = tmp_path / "voice.wav" + audio_path.write_bytes(b"voice") + pipeline = MagicMock(return_value={"text": "ok"}) + + with ( + patch.object(service, "_get_pipeline", return_value=pipeline), + patch( + "free_claude_code.messaging.transcription._load_audio", + return_value={"array": [0.0], "sampling_rate": 16000}, + ), + ): + await service.transcribe(audio_path) + + service._pipeline = pipeline + await service.close() + await service.close() + + assert service._pipeline is None + with pytest.raises(RuntimeError, match="closed"): + await service.transcribe(audio_path) + + +@pytest.mark.asyncio +async def test_cancelled_transcription_keeps_ownership_until_thread_exits( + tmp_path: Path, +) -> None: + service = _service(api_key="hf-secret") + audio_path = tmp_path / "voice.wav" + audio_path.write_bytes(b"voice") + started = threading.Event() + release = threading.Event() + pipeline = object() + + def blocking_transcribe(_path: Path) -> str: + started.set() + if not release.wait(timeout=5): + raise TimeoutError("test did not release active transcription") + service._pipeline = pipeline + return "finished" + + close_task: asyncio.Task[None] | None = None + with patch.object(service, "_transcribe_sync", side_effect=blocking_transcribe): + transcribe_task = asyncio.create_task(service.transcribe(audio_path)) + try: + assert await asyncio.to_thread(started.wait, 2) + transcribe_task.cancel() + await asyncio.sleep(0) + close_task = asyncio.create_task(service.close()) + await asyncio.sleep(0) + + assert not transcribe_task.done() + assert not close_task.done() + assert service._huggingface_api_key == "hf-secret" + finally: + release.set() + + with pytest.raises(asyncio.CancelledError): + await transcribe_task + assert close_task is not None + await close_task + + assert service._pipeline is None + assert service._huggingface_api_key == "" + with pytest.raises(RuntimeError, match="closed"): + await service.transcribe(audio_path) + + +@pytest.mark.asyncio +async def test_transcription_service_returns_no_speech_placeholder( + tmp_path: Path, +) -> None: + service = _service() + audio_path = tmp_path / "voice.ogg" + audio_path.write_bytes(b"voice") + pipeline = MagicMock(return_value={"text": []}) + + with ( + patch.object(service, "_get_pipeline", return_value=pipeline), + patch( + "free_claude_code.messaging.transcription._load_audio", + return_value={"array": [0.0], "sampling_rate": 16000}, + ), + ): + result = await service.transcribe(audio_path) + + assert result == "(no speech detected)" + + +def test_transcription_service_rejects_non_local_device() -> None: + with pytest.raises(ValueError, match="must be 'cpu' or 'cuda'"): + TranscriptionService(model="base", device="nvidia_nim") + + +@pytest.mark.asyncio +async def test_transcription_service_reports_missing_local_extra( + tmp_path: Path, +) -> None: + service = _service() + audio_path = tmp_path / "voice.ogg" + audio_path.write_bytes(b"voice") + + with ( + patch.dict("sys.modules", {"torch": None}), + pytest.raises(ImportError, match="voice_local extra"), + ): + await service.transcribe(audio_path) diff --git a/tests/messaging/test_transcription_nim.py b/tests/messaging/test_transcription_nim.py index 4a7583e3dfd4178584776c18f337ff457c4ef1cb..39e77ea19e2d4782d15c43ed6d7cb1d5d989ed22 100644 --- a/tests/messaging/test_transcription_nim.py +++ b/tests/messaging/test_transcription_nim.py @@ -1,39 +1,244 @@ -"""Tests for NVIDIA NIM voice transcription wiring.""" +"""Tests for the NVIDIA NIM voice transcription adapter.""" +import asyncio from pathlib import Path -from unittest.mock import patch +from threading import Event +from types import SimpleNamespace +from unittest.mock import MagicMock, patch -from free_claude_code.messaging.transcription import transcribe_audio +import pytest +from free_claude_code.providers.nvidia_nim.voice import ( + _NIM_ASR_MODEL_MAP, + NvidiaNimTranscriber, +) -def test_transcribe_audio_nvidia_nim_forwards_api_key(tmp_path: Path) -> None: + +def _fake_riva_client( + transcript: str, +) -> tuple[SimpleNamespace, SimpleNamespace, MagicMock, MagicMock]: + response = SimpleNamespace( + results=[ + SimpleNamespace( + alternatives=[SimpleNamespace(transcript=transcript)], + ) + ] + ) + asr_service = MagicMock() + asr_service.offline_recognize.return_value = response + auth = MagicMock() + + client = SimpleNamespace( + Auth=MagicMock(return_value=auth), + ASRService=MagicMock(return_value=asr_service), + RecognitionConfig=MagicMock(return_value=object()), + ) + riva = SimpleNamespace(__path__=[], client=client) + return riva, client, asr_service, auth + + +@pytest.mark.asyncio +async def test_nvidia_nim_transcriber_calls_riva_with_owned_configuration( + tmp_path: Path, +) -> None: wav = tmp_path / "stub.wav" - wav.write_bytes(b"\x00" * 128) - with patch( - "free_claude_code.messaging.transcription.transcribe_nvidia_nim_audio" - ) as nim_fn: - nim_fn.return_value = "ok" - out = transcribe_audio( - wav, - "audio/wav", - whisper_model="openai/whisper-large-v3", - whisper_device="nvidia_nim", - nvidia_nim_api_key="test-nim-key", - ) - nim_fn.assert_called_once_with( - wav, "openai/whisper-large-v3", api_key="test-nim-key" - ) - assert out == "ok" + wav.write_bytes(b"audio bytes") + transcriber = NvidiaNimTranscriber( + model="openai/whisper-large-v3", + api_key=" test-nim-key ", + ) + riva, client, asr_service, auth = _fake_riva_client(" hello from NIM ") + with patch.dict( + "sys.modules", + {"riva": riva, "riva.client": client}, + ): + result = await transcriber.transcribe(wav) + + assert result == " hello from NIM " + client.Auth.assert_called_once_with( + use_ssl=True, + uri="grpc.nvcf.nvidia.com:443", + metadata_args=[ + ["function-id", "b702f636-f60c-4a3d-a6f4-f3568c13bd7d"], + ["authorization", "Bearer test-nim-key"], + ], + ) + client.RecognitionConfig.assert_called_once_with( + language_code="multi", + max_alternatives=1, + verbatim_transcripts=True, + ) + asr_service.offline_recognize.assert_called_once_with( + b"audio bytes", + client.RecognitionConfig.return_value, + ) + auth.channel.close.assert_called_once_with() -def test_nim_asr_model_map_entries_are_real_function_ids() -> None: - from free_claude_code.providers.nvidia_nim.voice import _NIM_ASR_MODEL_MAP +@pytest.mark.asyncio +async def test_nvidia_nim_transcriber_closes_channel_when_recognition_fails( + tmp_path: Path, +) -> None: + wav = tmp_path / "stub.wav" + wav.write_bytes(b"audio bytes") + transcriber = NvidiaNimTranscriber( + model="openai/whisper-large-v3", + api_key="test-nim-key", + ) + riva, client, asr_service, auth = _fake_riva_client("") + asr_service.offline_recognize.side_effect = RuntimeError("recognition failed") + + with ( + patch.dict( + "sys.modules", + {"riva": riva, "riva.client": client}, + ), + pytest.raises(RuntimeError, match="recognition failed"), + ): + await transcriber.transcribe(wav) + + auth.channel.close.assert_called_once_with() + + +@pytest.mark.asyncio +async def test_nvidia_nim_transcriber_validates_key_and_model_before_import( + tmp_path: Path, +) -> None: + wav = tmp_path / "stub.wav" + wav.write_bytes(b"audio") + + with pytest.raises(ValueError, match="non-empty"): + await NvidiaNimTranscriber( + model="openai/whisper-large-v3", + api_key="", + ).transcribe(wav) + + with pytest.raises(ValueError, match="No NVIDIA NIM config"): + await NvidiaNimTranscriber( + model="unknown/model", + api_key="key", + ).transcribe(wav) + + +@pytest.mark.asyncio +async def test_nvidia_nim_transcriber_close_is_terminal(tmp_path: Path) -> None: + wav = tmp_path / "stub.wav" + wav.write_bytes(b"audio") + transcriber = NvidiaNimTranscriber( + model="openai/whisper-large-v3", + api_key="key", + ) + + await transcriber.close() + await transcriber.close() + + with pytest.raises(RuntimeError, match="closed"): + await transcriber.transcribe(wav) + + +@pytest.mark.asyncio +async def test_nvidia_nim_transcriber_close_waits_for_active_work( + tmp_path: Path, +) -> None: + wav = tmp_path / "stub.wav" + wav.write_bytes(b"audio") + transcriber = NvidiaNimTranscriber( + model="openai/whisper-large-v3", + api_key="secret-key", + ) + started = Event() + release = Event() + + def blocking_transcribe(_file_path: Path) -> str: + started.set() + if not release.wait(timeout=5): + raise TimeoutError("test did not release active transcription") + return "finished" + + close_task: asyncio.Task[None] | None = None + with patch.object( + transcriber, + "_transcribe_sync", + side_effect=blocking_transcribe, + ): + transcribe_task = asyncio.create_task(transcriber.transcribe(wav)) + try: + assert await asyncio.to_thread(started.wait, 2) + close_task = asyncio.create_task(transcriber.close()) + await asyncio.sleep(0) + + assert not close_task.done() + finally: + release.set() + transcript = await transcribe_task + if close_task is not None: + await close_task + + assert transcript == "finished" + assert transcriber._key == "" + with pytest.raises(RuntimeError, match="closed"): + await transcriber.transcribe(wav) + + +@pytest.mark.asyncio +async def test_cancelled_nim_transcription_keeps_key_until_thread_exits( + tmp_path: Path, +) -> None: + wav = tmp_path / "stub.wav" + wav.write_bytes(b"audio") + transcriber = NvidiaNimTranscriber( + model="openai/whisper-large-v3", + api_key="secret-key", + ) + started = Event() + release = Event() + observed_keys: list[str] = [] + + def blocking_transcribe(_file_path: Path) -> str: + observed_keys.append(transcriber._key) + started.set() + if not release.wait(timeout=5): + raise TimeoutError("test did not release active transcription") + observed_keys.append(transcriber._key) + return "finished" + + close_task: asyncio.Task[None] | None = None + with patch.object( + transcriber, + "_transcribe_sync", + side_effect=blocking_transcribe, + ): + transcribe_task = asyncio.create_task(transcriber.transcribe(wav)) + try: + assert await asyncio.to_thread(started.wait, 2) + transcribe_task.cancel() + await asyncio.sleep(0) + close_task = asyncio.create_task(transcriber.close()) + await asyncio.sleep(0) + + assert not transcribe_task.done() + assert not close_task.done() + assert transcriber._key == "secret-key" + finally: + release.set() + + with pytest.raises(asyncio.CancelledError): + await transcribe_task + assert close_task is not None + await close_task + + assert observed_keys == ["secret-key", "secret-key"] + assert transcriber._key == "" + with pytest.raises(RuntimeError, match="closed"): + await transcriber.transcribe(wav) + + +def test_nim_asr_model_map_entries_are_real_function_ids() -> None: for function_id, language_code in _NIM_ASR_MODEL_MAP.values(): assert function_id assert function_id.strip().lower() != "none" - # Hosted NIM function-id is a lowercase UUID string. parts = function_id.split("-") assert len(parts) == 5 - assert all(p for p in parts) + assert all(parts) assert language_code is not None diff --git a/tests/messaging/test_voice_handlers.py b/tests/messaging/test_voice_handlers.py index c1b3e16e754973d8953ff1dd03e39db63838a0b8..a33d77a8341179e4b4db1fb4375e7c906a6bfd5f 100644 --- a/tests/messaging/test_voice_handlers.py +++ b/tests/messaging/test_voice_handlers.py @@ -1,6 +1,5 @@ """Tests for voice note handling in Telegram and Discord platforms.""" -import tempfile from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch @@ -16,10 +15,19 @@ from free_claude_code.messaging.platforms.telegram import TelegramRuntime @pytest.fixture def telegram_platform(): + transcriber = MagicMock() + transcriber.transcribe = AsyncMock(return_value="Hello from voice") + transcriber.close = AsyncMock() with patch( "free_claude_code.messaging.platforms.telegram.TELEGRAM_AVAILABLE", True ): - return TelegramRuntime(bot_token="test_token", allowed_user_id="12345") + platform = TelegramRuntime( + bot_token="test_token", + allowed_user_id="12345", + limiter=MagicMock(), + transcriber=transcriber, + ) + return platform, transcriber @pytest.mark.asyncio @@ -31,7 +39,8 @@ async def test_telegram_voice_disabled_sends_reply(): telegram_platform = TelegramRuntime( bot_token="test_token", allowed_user_id="12345", - voice_note_enabled=False, + limiter=MagicMock(), + transcriber=None, ) mock_update = MagicMock() mock_update.message.voice = MagicMock(file_id="f1", mime_type="audio/ogg") @@ -47,21 +56,24 @@ async def test_telegram_voice_disabled_sends_reply(): @pytest.mark.asyncio async def test_telegram_voice_unauthorized_ignored(telegram_platform): """Voice from unauthorized user is ignored (no reply).""" + platform, transcriber = telegram_platform mock_update = MagicMock() mock_update.message.voice = MagicMock(file_id="f1", mime_type="audio/ogg") mock_update.effective_user.id = 99999 # Not 12345 mock_update.message.reply_text = AsyncMock() - await telegram_platform._on_telegram_voice(mock_update, MagicMock()) + await platform._on_telegram_voice(mock_update, MagicMock()) mock_update.message.reply_text.assert_not_called() + transcriber.transcribe.assert_not_awaited() @pytest.mark.asyncio async def test_telegram_voice_success_invokes_handler(telegram_platform): """Successful transcription invokes message handler with transcribed text.""" + platform, transcriber = telegram_platform handler = AsyncMock() - telegram_platform.on_message(handler) + platform.on_message(handler) mock_update = MagicMock() mock_voice = MagicMock(file_id="f1", mime_type="audio/ogg") @@ -76,47 +88,34 @@ async def test_telegram_voice_success_invokes_handler(telegram_platform): mock_context = MagicMock() mock_context.bot.get_file = AsyncMock(return_value=mock_file) - with tempfile.NamedTemporaryFile(suffix=".ogg", delete=False) as f: - f.write(b"fake") - tmp_path = Path(f.name) - - try: - - async def fake_download(custom_path=None): - if custom_path: - Path(custom_path).write_bytes(b"fake ogg") - - mock_file.download_to_drive = fake_download - - mock_queue_send = AsyncMock(return_value="999") - with ( - patch( - "free_claude_code.messaging.transcription.transcribe_audio", - return_value="Hello from voice", - ), - patch.object( - telegram_platform.outbound, - "queue_send_message", - mock_queue_send, - ), - ): - await telegram_platform._on_telegram_voice(mock_update, mock_context) - - mock_queue_send.assert_called_once() - call_args, call_kw = mock_queue_send.call_args - assert "Transcribing voice note" in call_args[1] - assert call_kw["reply_to"] == "42" - assert call_kw["fire_and_forget"] is False - - handler.assert_called_once() - incoming = handler.call_args[0][0] - assert incoming.text == "Hello from voice" - assert incoming.chat_id == "6789" - assert incoming.user_id == "12345" - assert incoming.platform == "telegram" - assert incoming.status_message_id == "999" - finally: - tmp_path.unlink(missing_ok=True) + async def fake_download(custom_path=None): + if custom_path: + Path(custom_path).write_bytes(b"fake ogg") + + mock_file.download_to_drive = fake_download + + mock_queue_send = AsyncMock(return_value="999") + with patch.object( + platform.outbound, + "queue_send_message", + mock_queue_send, + ): + await platform._on_telegram_voice(mock_update, mock_context) + + mock_queue_send.assert_called_once() + call_args, call_kw = mock_queue_send.call_args + assert "Transcribing voice note" in call_args[1] + assert call_kw["reply_to"] == "42" + assert call_kw["fire_and_forget"] is False + transcriber.transcribe.assert_awaited_once() + + handler.assert_called_once() + incoming = handler.call_args[0][0] + assert incoming.text == "Hello from voice" + assert incoming.chat_id == "6789" + assert incoming.user_id == "12345" + assert incoming.platform == "telegram" + assert incoming.status_message_id == "999" @pytest.mark.skipif(not DISCORD_AVAILABLE, reason="discord.py not installed") @@ -160,7 +159,8 @@ async def test_discord_voice_disabled_sends_reply(): platform = DiscordRuntime( bot_token="token", allowed_channel_ids="123", - voice_note_enabled=False, + limiter=MagicMock(), + transcriber=None, ) platform._message_handler = None diff --git a/tests/messaging/test_voice_services.py b/tests/messaging/test_voice_services.py index 1d6496f20cf93193f87f5a83e0cf7756f5f2f00a..b97fe730a2d734e03f31e8fb8ef955d2f9940468 100644 --- a/tests/messaging/test_voice_services.py +++ b/tests/messaging/test_voice_services.py @@ -1,12 +1,6 @@ -from pathlib import Path -from unittest.mock import patch - import pytest -from free_claude_code.messaging.voice import ( - PendingVoiceRegistry, - VoiceTranscriptionService, -) +from free_claude_code.messaging.voice import PendingVoiceRegistry @pytest.mark.asyncio @@ -28,22 +22,3 @@ async def test_pending_voice_registry_complete_removes_entries(): await registry.complete("chat", "voice-1", "status-1") assert await registry.cancel("chat", "voice-1") is None - - -@pytest.mark.asyncio -async def test_voice_transcription_service_runs_backend(): - service = VoiceTranscriptionService(huggingface_api_key="hf-provider-key") - - with patch( - "free_claude_code.messaging.transcription.transcribe_audio", - return_value="hello", - ) as run: - text = await service.transcribe( - Path("audio.ogg"), - "audio/ogg", - whisper_model="base", - whisper_device="cpu", - ) - - assert text == "hello" - assert run.call_args.kwargs["huggingface_api_key"] == "hf-provider-key" diff --git a/tests/providers/support.py b/tests/providers/support.py new file mode 100644 index 0000000000000000000000000000000000000000..8b8cbf8ed2fcf83b2aeeabd3627cad7d0ad4afb1 --- /dev/null +++ b/tests/providers/support.py @@ -0,0 +1,39 @@ +"""Provider test doubles with explicit limiter ownership.""" + +from collections.abc import Callable +from typing import Any + +from free_claude_code.providers.rate_limit import ProviderRateLimiter + + +class PassthroughProviderRateLimiter(ProviderRateLimiter): + """Skip retry timing while retaining the real concurrency context manager.""" + + def __init__(self) -> None: + super().__init__( + rate_limit=1_000_000, + rate_window=1.0, + max_concurrency=1_000, + ) + + async def execute_with_retry( + self, + fn: Callable[..., Any], + *args: Any, + **kwargs: Any, + ) -> Any: + return await fn(*args, **kwargs) + + +def passthrough_rate_limiter() -> ProviderRateLimiter: + """Return a fresh limiter test double for one provider instance.""" + return PassthroughProviderRateLimiter() + + +def retrying_rate_limiter() -> ProviderRateLimiter: + """Return a fresh real limiter for provider retry-policy tests.""" + return ProviderRateLimiter( + rate_limit=1_000_000, + rate_window=1.0, + max_concurrency=1_000, + ) diff --git a/tests/providers/test_anthropic_messages.py b/tests/providers/test_anthropic_messages.py index b5e3aecb9aa01e40254df0e9effea9380075a742..fa52dedafb46a10a4539c6e2f3a1b8247d31ef8b 100644 --- a/tests/providers/test_anthropic_messages.py +++ b/tests/providers/test_anthropic_messages.py @@ -16,20 +16,23 @@ from free_claude_code.core.anthropic.streaming import ( ) from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.exceptions import ProviderError +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.anthropic_messages import ( AnthropicMessagesTransport, ) from free_claude_code.providers.transports.anthropic_messages.recovery import ( AnthropicMessagesRecovery, ) +from tests.providers.support import passthrough_rate_limiter class NativeProvider(AnthropicMessagesTransport): - def __init__(self, config: ProviderConfig): + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): super().__init__( config, provider_name="TEST_NATIVE", default_base_url="https://example.test/v1", + rate_limiter=rate_limiter, ) def _request_headers(self) -> dict[str, str]: @@ -122,28 +125,28 @@ def provider_config(): ) -@pytest.fixture(autouse=True) +@pytest.fixture def mock_rate_limiter(): @asynccontextmanager async def _slot(): yield - with patch( - "free_claude_code.providers.transports.anthropic_messages.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value + instance = MagicMock(spec=ProviderRateLimiter) - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) + async def _passthrough(fn, *args, **kwargs): + return await fn(*args, **kwargs) - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance + instance.execute_with_retry = AsyncMock(side_effect=_passthrough) + instance.concurrency_slot.side_effect = _slot + yield instance def test_init_configures_httpx_client(provider_config): with patch("httpx.AsyncClient") as mock_client: - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) assert provider._provider_name == "TEST_NATIVE" assert provider._api_key == "test-key" @@ -158,7 +161,10 @@ def test_init_configures_httpx_client(provider_config): def test_default_request_body_strips_internal_fields(provider_config): - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) body = provider._build_request_body(MockRequest()) @@ -169,7 +175,10 @@ def test_default_request_body_strips_internal_fields(provider_config): def test_default_request_body_preserves_thinking_budget(provider_config): - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest( body={ "model": "test-model", @@ -185,7 +194,10 @@ def test_default_request_body_preserves_thinking_budget(provider_config): @pytest.mark.asyncio async def test_send_stream_request_forces_upstream_streaming(provider_config): - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) request_obj = httpx.Request("POST", "https://custom.test/v1/messages") response = FakeResponse() body = {"model": "test-model", "stream": False} @@ -212,7 +224,7 @@ async def test_stream_uses_retry_builds_request_and_closes_response( provider_config, mock_rate_limiter, ): - provider = NativeProvider(provider_config) + provider = NativeProvider(provider_config, rate_limiter=mock_rate_limiter) req = MockRequest() request_obj = httpx.Request("POST", "https://custom.test/v1/messages") response = FakeResponse( @@ -259,7 +271,7 @@ async def test_late_error_after_native_message_stop_keeps_successful_lifecycle( provider_config, mock_rate_limiter, ): - provider = NativeProvider(provider_config) + provider = NativeProvider(provider_config, rate_limiter=mock_rate_limiter) req = MockRequest() lines = [ "event: message_start", @@ -301,7 +313,10 @@ async def test_late_error_after_native_message_stop_keeps_successful_lifecycle( async def test_stream_maps_pre_start_non_200_to_provider_error_and_closes_response( provider_config, ): - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() response = FakeResponse(status_code=500, text="Internal Server Error") @@ -328,7 +343,10 @@ async def test_precommit_native_error_raises_without_leaking_open_block( provider_config, ): """A native error before holdback commit raises instead of sending HTTP 200 SSE.""" - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() mid = "msg_midstream_err" msg_start = format_sse_event( @@ -380,7 +398,10 @@ async def test_midstream_error_after_native_message_delta_does_not_duplicate_ter provider_config, ): """If native upstream emitted message_delta before cutoff, recovery cannot append content.""" - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() msg_start = format_sse_event( "message_start", @@ -519,7 +540,10 @@ async def test_clean_eof_after_complete_native_tool_call_salvages_tool_use( provider_config, ): """Native stream EOF after complete tool args gets a deterministic tool_use tail.""" - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() msg_start = format_sse_event( "message_start", @@ -589,7 +613,10 @@ async def test_clean_eof_after_native_text_continues_with_overlap_trim( provider_config, ): """Native text truncation is continued and overlap-trimmed.""" - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() msg_start = format_sse_event( "message_start", @@ -664,7 +691,10 @@ async def test_clean_eof_after_native_text_continues_with_overlap_trim( @pytest.mark.asyncio async def test_native_recovery_collect_text_requires_message_stop(provider_config): """Native recovery collectors reject truncated continuation streams.""" - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) text_delta = format_sse_event( "content_block_delta", { @@ -696,7 +726,10 @@ async def test_native_recovery_collect_text_requires_message_stop(provider_confi @pytest.mark.asyncio async def test_native_recovery_collect_text_accepts_message_stop(provider_config): """Native recovery collectors return text only after message_stop.""" - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) text_delta = format_sse_event( "content_block_delta", { @@ -729,7 +762,10 @@ async def test_native_recovery_collect_text_accepts_message_stop(provider_config @pytest.mark.asyncio async def test_native_recovery_collect_text_reads_eager_start_content(provider_config): """Native recovery reads text/thinking carried on content_block_start.""" - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) text_start = format_sse_event( "content_block_start", { @@ -791,7 +827,10 @@ async def test_truncated_native_recovery_stream_falls_back_to_error_tail( provider_config, ): """Partial native recovery bytes are not converted into a success tail.""" - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() msg_start = format_sse_event( "message_start", @@ -873,7 +912,10 @@ async def test_precommit_native_holdback_retries_without_leaking_partial( provider_config, ): """A retryable early cutoff before holdback commit is retried invisibly.""" - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() msg_start = format_sse_event( diff --git a/tests/providers/test_anthropic_messages_429_retry.py b/tests/providers/test_anthropic_messages_429_retry.py index 1baa1ad1a75df5d67619e22b9777b1d5b451bf91..e1fc286eb0e41329e654b0250f74b23957628fa6 100644 --- a/tests/providers/test_anthropic_messages_429_retry.py +++ b/tests/providers/test_anthropic_messages_429_retry.py @@ -1,6 +1,5 @@ """Native Anthropic transport: HTTP 429 and upstream 5xx are retried inside execute_with_retry.""" -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -9,7 +8,7 @@ import pytest from free_claude_code.core.anthropic.stream_contracts import event_names, parse_sse_text from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.exceptions import ProviderError -from free_claude_code.providers.rate_limit import GlobalRateLimiter +from tests.providers.support import retrying_rate_limiter from tests.providers.test_anthropic_messages import ( FakeResponse, MockRequest, @@ -40,51 +39,47 @@ def provider_config(): @pytest.mark.asyncio async def test_native_stream_retries_on_http_429_then_streams(provider_config): """First response 429 (closed), second 200 streams; send is called twice.""" - GlobalRateLimiter.reset_instance() - try: - provider = NativeProvider(provider_config) - req = MockRequest() - request_obj = httpx.Request("POST", "https://custom.test/v1/messages") - ok_lines = [ - "event: message_start", - 'data: {"type":"message_start"}', - "", - "event: message_stop", - 'data: {"type":"message_stop"}', - "", - ] - ok_response = FakeResponse(lines=ok_lines) - too_many = FakeResponse(status_code=429, text="rate limited") - - send_calls = {"n": 0} - - async def send_side_effect(*_a, **_kw): - send_calls["n"] += 1 - if send_calls["n"] == 1: - return too_many - return ok_response - - with ( - patch.object(provider._client, "build_request", return_value=request_obj), - patch.object( - provider._client, - "send", - new_callable=AsyncMock, - side_effect=send_side_effect, - ), - patch( - "asyncio.sleep", - new_callable=AsyncMock, - ), - ): - events = [e async for e in provider.stream_response(req)] - - assert send_calls["n"] == 2 - assert too_many.is_closed - assert ok_response.is_closed - _assert_minimal_success_stream(events) - finally: - GlobalRateLimiter.reset_instance() + provider = NativeProvider(provider_config, rate_limiter=retrying_rate_limiter()) + req = MockRequest() + request_obj = httpx.Request("POST", "https://custom.test/v1/messages") + ok_lines = [ + "event: message_start", + 'data: {"type":"message_start"}', + "", + "event: message_stop", + 'data: {"type":"message_stop"}', + "", + ] + ok_response = FakeResponse(lines=ok_lines) + too_many = FakeResponse(status_code=429, text="rate limited") + + send_calls = {"n": 0} + + async def send_side_effect(*_a, **_kw): + send_calls["n"] += 1 + if send_calls["n"] == 1: + return too_many + return ok_response + + with ( + patch.object(provider._client, "build_request", return_value=request_obj), + patch.object( + provider._client, + "send", + new_callable=AsyncMock, + side_effect=send_side_effect, + ), + patch( + "asyncio.sleep", + new_callable=AsyncMock, + ), + ): + events = [e async for e in provider.stream_response(req)] + + assert send_calls["n"] == 2 + assert too_many.is_closed + assert ok_response.is_closed + _assert_minimal_success_stream(events) @pytest.mark.parametrize("status_code", [500, 502, 503, 504]) @@ -93,51 +88,47 @@ async def test_native_stream_retries_on_http_5xx_then_streams( provider_config, status_code ): """First response is retryable 5xx (closed); second 200 streams; send twice.""" - GlobalRateLimiter.reset_instance() - try: - provider = NativeProvider(provider_config) - req = MockRequest() - request_obj = httpx.Request("POST", "https://custom.test/v1/messages") - ok_lines = [ - "event: message_start", - 'data: {"type":"message_start"}', - "", - "event: message_stop", - 'data: {"type":"message_stop"}', - "", - ] - ok_response = FakeResponse(lines=ok_lines) - bad = FakeResponse(status_code=status_code, text="upstream error") - - send_calls = {"n": 0} - - async def send_side_effect(*_a, **_kw): - send_calls["n"] += 1 - if send_calls["n"] == 1: - return bad - return ok_response - - with ( - patch.object(provider._client, "build_request", return_value=request_obj), - patch.object( - provider._client, - "send", - new_callable=AsyncMock, - side_effect=send_side_effect, - ), - patch( - "asyncio.sleep", - new_callable=AsyncMock, - ), - ): - events = [e async for e in provider.stream_response(req)] - - assert send_calls["n"] == 2 - assert bad.is_closed - assert ok_response.is_closed - _assert_minimal_success_stream(events) - finally: - GlobalRateLimiter.reset_instance() + provider = NativeProvider(provider_config, rate_limiter=retrying_rate_limiter()) + req = MockRequest() + request_obj = httpx.Request("POST", "https://custom.test/v1/messages") + ok_lines = [ + "event: message_start", + 'data: {"type":"message_start"}', + "", + "event: message_stop", + 'data: {"type":"message_stop"}', + "", + ] + ok_response = FakeResponse(lines=ok_lines) + bad = FakeResponse(status_code=status_code, text="upstream error") + + send_calls = {"n": 0} + + async def send_side_effect(*_a, **_kw): + send_calls["n"] += 1 + if send_calls["n"] == 1: + return bad + return ok_response + + with ( + patch.object(provider._client, "build_request", return_value=request_obj), + patch.object( + provider._client, + "send", + new_callable=AsyncMock, + side_effect=send_side_effect, + ), + patch( + "asyncio.sleep", + new_callable=AsyncMock, + ), + ): + events = [e async for e in provider.stream_response(req)] + + assert send_calls["n"] == 2 + assert bad.is_closed + assert ok_response.is_closed + _assert_minimal_success_stream(events) @pytest.mark.asyncio @@ -145,46 +136,42 @@ async def test_native_stream_retries_on_pre_send_connection_error_then_streams( provider_config, ): """Pre-response HTTPX transport errors retry through execute_with_retry.""" - GlobalRateLimiter.reset_instance() - try: - provider = NativeProvider(provider_config) - req = MockRequest() - request_obj = httpx.Request("POST", "https://custom.test/v1/messages") - ok_lines = [ - "event: message_start", - 'data: {"type":"message_start"}', - "", - "event: message_stop", - 'data: {"type":"message_stop"}', - "", - ] - ok_response = FakeResponse(lines=ok_lines) - - send_calls = {"n": 0} - - async def send_side_effect(*_a, **_kw): - send_calls["n"] += 1 - if send_calls["n"] == 1: - raise httpx.ConnectError("connect failed", request=request_obj) - return ok_response - - with ( - patch.object(provider._client, "build_request", return_value=request_obj), - patch.object( - provider._client, - "send", - new_callable=AsyncMock, - side_effect=send_side_effect, - ), - patch("asyncio.sleep", new_callable=AsyncMock), - ): - events = [e async for e in provider.stream_response(req)] - - assert send_calls["n"] == 2 - assert ok_response.is_closed - _assert_minimal_success_stream(events) - finally: - GlobalRateLimiter.reset_instance() + provider = NativeProvider(provider_config, rate_limiter=retrying_rate_limiter()) + req = MockRequest() + request_obj = httpx.Request("POST", "https://custom.test/v1/messages") + ok_lines = [ + "event: message_start", + 'data: {"type":"message_start"}', + "", + "event: message_stop", + 'data: {"type":"message_stop"}', + "", + ] + ok_response = FakeResponse(lines=ok_lines) + + send_calls = {"n": 0} + + async def send_side_effect(*_a, **_kw): + send_calls["n"] += 1 + if send_calls["n"] == 1: + raise httpx.ConnectError("connect failed", request=request_obj) + return ok_response + + with ( + patch.object(provider._client, "build_request", return_value=request_obj), + patch.object( + provider._client, + "send", + new_callable=AsyncMock, + side_effect=send_side_effect, + ), + patch("asyncio.sleep", new_callable=AsyncMock), + ): + events = [e async for e in provider.stream_response(req)] + + assert send_calls["n"] == 2 + assert ok_response.is_closed + _assert_minimal_success_stream(events) @pytest.mark.parametrize( @@ -199,162 +186,95 @@ async def test_native_stream_retries_on_pre_send_connection_error_then_streams( @pytest.mark.asyncio async def test_native_stream_5xx_retry_exhausted(provider_config, status_code, substr): """Repeated upstream 5xx exhausts execute_with_retry; user message matches mapping.""" - GlobalRateLimiter.reset_instance() - try: - - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.anthropic_messages.transport.GlobalRateLimiter" - ) as mock_gl: - instance = mock_gl.get_scoped_instance.return_value - real = GlobalRateLimiter( - rate_limit=100, - rate_window=60, - max_concurrency=5, - ) - instance.wait_if_blocked = real.wait_if_blocked - instance.execute_with_retry = real.execute_with_retry - instance.set_blocked = real.set_blocked - instance.concurrency_slot.side_effect = _slot - - provider = NativeProvider(provider_config) - req = MockRequest() - - bad = FakeResponse(status_code=status_code, text="upstream error") - - with ( - patch.object( - provider._client, "build_request", return_value=MagicMock() - ), - patch.object( - provider._client, - "send", - new_callable=AsyncMock, - return_value=bad, - ) as mock_send, - patch("asyncio.sleep", new_callable=AsyncMock), - pytest.raises(ProviderError) as exc_info, - ): - [e async for e in provider.stream_response(req)] - - assert mock_send.await_count == 5 - assert bad.is_closed - assert substr in exc_info.value.message - finally: - GlobalRateLimiter.reset_instance() + + provider = NativeProvider( + provider_config, + rate_limiter=retrying_rate_limiter(), + ) + req = MockRequest() + + bad = FakeResponse(status_code=status_code, text="upstream error") + + with ( + patch.object(provider._client, "build_request", return_value=MagicMock()), + patch.object( + provider._client, + "send", + new_callable=AsyncMock, + return_value=bad, + ) as mock_send, + patch("asyncio.sleep", new_callable=AsyncMock), + pytest.raises(ProviderError) as exc_info, + ): + [e async for e in provider.stream_response(req)] + + assert mock_send.await_count == 5 + assert bad.is_closed + assert substr in exc_info.value.message @pytest.mark.asyncio async def test_native_stream_connection_error_retry_exhausted(provider_config): """Repeated pre-response connection failures exhaust at 5 attempts.""" - GlobalRateLimiter.reset_instance() - try: - - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.anthropic_messages.transport.GlobalRateLimiter" - ) as mock_gl: - instance = mock_gl.get_scoped_instance.return_value - real = GlobalRateLimiter( - rate_limit=100, - rate_window=60, - max_concurrency=5, - ) - instance.wait_if_blocked = real.wait_if_blocked - instance.execute_with_retry = real.execute_with_retry - instance.set_blocked = real.set_blocked - instance.concurrency_slot.side_effect = _slot - - provider = NativeProvider(provider_config) - req = MockRequest() - request_obj = httpx.Request("POST", "https://custom.test/v1/messages") - - with ( - patch.object( - provider._client, "build_request", return_value=request_obj - ), - patch.object( - provider._client, - "send", - new_callable=AsyncMock, - side_effect=httpx.ConnectError( - "connect failed", request=request_obj - ), - ) as mock_send, - patch("asyncio.sleep", new_callable=AsyncMock), - patch( - "free_claude_code.providers.transports.anthropic_messages.stream.trace_event" - ) as trace, - pytest.raises(ProviderError) as exc_info, - ): - [ - e - async for e in provider.stream_response( - req, request_id="req_native_conn" - ) - ] - - assert mock_send.await_count == 5 - error_traces = [ - call.kwargs - for call in trace.call_args_list - if call.kwargs.get("event") == "provider.response.error" - ] - assert error_traces[-1]["request_id"] == "req_native_conn" - assert error_traces[-1]["exc_type"] == "ConnectError" - assert "error_message" not in error_traces[-1] - assert "Provider exception:\nconnect failed" in exc_info.value.message - finally: - GlobalRateLimiter.reset_instance() + + provider = NativeProvider( + provider_config, + rate_limiter=retrying_rate_limiter(), + ) + req = MockRequest() + request_obj = httpx.Request("POST", "https://custom.test/v1/messages") + + with ( + patch.object(provider._client, "build_request", return_value=request_obj), + patch.object( + provider._client, + "send", + new_callable=AsyncMock, + side_effect=httpx.ConnectError("connect failed", request=request_obj), + ) as mock_send, + patch("asyncio.sleep", new_callable=AsyncMock), + patch( + "free_claude_code.providers.transports.anthropic_messages.stream.trace_event" + ) as trace, + pytest.raises(ProviderError) as exc_info, + ): + [e async for e in provider.stream_response(req, request_id="req_native_conn")] + + assert mock_send.await_count == 5 + error_traces = [ + call.kwargs + for call in trace.call_args_list + if call.kwargs.get("event") == "provider.response.error" + ] + assert error_traces[-1]["request_id"] == "req_native_conn" + assert error_traces[-1]["exc_type"] == "ConnectError" + assert "error_message" not in error_traces[-1] + assert "Provider exception:\nconnect failed" in exc_info.value.message @pytest.mark.asyncio async def test_non_retryable_4xx_http_error_not_retried(provider_config): """HTTP 400 from upstream is not retried; single send (passthrough limiter).""" - GlobalRateLimiter.reset_instance() - try: - - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.anthropic_messages.transport.GlobalRateLimiter" - ) as mock_gl: - instance = mock_gl.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - - provider = NativeProvider(provider_config) - req = MockRequest() - err = FakeResponse(status_code=400, text="Bad Request") - - with ( - patch.object( - provider._client, "build_request", return_value=MagicMock() - ), - patch.object( - provider._client, - "send", - new_callable=AsyncMock, - return_value=err, - ) as mock_send, - pytest.raises(ProviderError) as exc_info, - ): - [e async for e in provider.stream_response(req)] - - mock_send.assert_awaited_once() - assert err.is_closed - assert "Invalid request sent to provider" in exc_info.value.message - finally: - GlobalRateLimiter.reset_instance() + + provider = NativeProvider( + provider_config, + rate_limiter=retrying_rate_limiter(), + ) + req = MockRequest() + err = FakeResponse(status_code=400, text="Bad Request") + + with ( + patch.object(provider._client, "build_request", return_value=MagicMock()), + patch.object( + provider._client, + "send", + new_callable=AsyncMock, + return_value=err, + ) as mock_send, + pytest.raises(ProviderError) as exc_info, + ): + [e async for e in provider.stream_response(req)] + + mock_send.assert_awaited_once() + assert err.is_closed + assert "Invalid request sent to provider" in exc_info.value.message diff --git a/tests/providers/test_cerebras.py b/tests/providers/test_cerebras.py index 3e2de6558bf4c5604fdbc04121a4ffcaf33101dc..64d01590e606f87315ba03b5b655dbf1aeee8d5b 100644 --- a/tests/providers/test_cerebras.py +++ b/tests/providers/test_cerebras.py @@ -1,12 +1,12 @@ """Tests for Cerebras Inference (OpenAI-compatible) provider.""" -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import pytest from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.cerebras import CEREBRAS_DEFAULT_BASE, CerebrasProvider +from tests.providers.support import passthrough_rate_limiter class MockMessage: @@ -42,30 +42,9 @@ def cerebras_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - """Mock the global rate limiter to prevent waiting.""" - - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def cerebras_provider(cerebras_config): - return CerebrasProvider(cerebras_config) + return CerebrasProvider(cerebras_config, rate_limiter=passthrough_rate_limiter()) def test_init(cerebras_config): @@ -73,7 +52,9 @@ def test_init(cerebras_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_openai: - provider = CerebrasProvider(cerebras_config) + provider = CerebrasProvider( + cerebras_config, rate_limiter=passthrough_rate_limiter() + ) assert provider._api_key == "test_cerebras_key" assert provider._base_url == CEREBRAS_DEFAULT_BASE mock_openai.assert_called_once() @@ -101,7 +82,8 @@ def test_build_request_body_global_disable_blocks_reasoning_mapping(): rate_limit=10, rate_window=60, enable_thinking=False, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) req = MockRequest() body = provider._build_request_body(req) diff --git a/tests/providers/test_cloudflare.py b/tests/providers/test_cloudflare.py index 2e9d0063efe4e1f03a824cdf3f3474f1ed5b5396..10fbcc700c42a56461396e630064bf9426420fd3 100644 --- a/tests/providers/test_cloudflare.py +++ b/tests/providers/test_cloudflare.py @@ -1,7 +1,6 @@ """Tests for Cloudflare Workers AI OpenAI-compatible chat provider.""" from collections.abc import AsyncIterator -from contextlib import asynccontextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, patch @@ -17,6 +16,7 @@ from free_claude_code.providers.cloudflare import ( cloudflare_ai_base_url, ) from free_claude_code.providers.exceptions import AuthenticationError +from tests.providers.support import passthrough_rate_limiter _ACCOUNT_ID = "account-123" _BASE_URL = f"{CLOUDFLARE_AI_REST_ROOT}/accounts/{_ACCOUNT_ID}/ai/v1" @@ -34,28 +34,13 @@ def cloudflare_config() -> ProviderConfig: ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def cloudflare_provider(cloudflare_config: ProviderConfig) -> CloudflareProvider: - return CloudflareProvider(cloudflare_config, account_id=_ACCOUNT_ID) + return CloudflareProvider( + cloudflare_config, + account_id=_ACCOUNT_ID, + rate_limiter=passthrough_rate_limiter(), + ) def _request(model: str = "@cf/moonshotai/kimi-k2.6") -> MessagesRequest: @@ -88,7 +73,9 @@ def test_missing_account_id_raises_authentication_error( cloudflare_config: ProviderConfig, ) -> None: with pytest.raises(AuthenticationError, match="CLOUDFLARE_ACCOUNT_ID"): - CloudflareProvider(cloudflare_config, account_id=" ") + CloudflareProvider( + cloudflare_config, account_id=" ", rate_limiter=passthrough_rate_limiter() + ) def test_init_composes_account_scoped_openai_chat_base_url( @@ -100,7 +87,11 @@ def test_init_composes_account_scoped_openai_chat_base_url( ) as mock_openai, patch("httpx.AsyncClient") as mock_httpx_client, ): - provider = CloudflareProvider(cloudflare_config, account_id=_ACCOUNT_ID) + provider = CloudflareProvider( + cloudflare_config, + account_id=_ACCOUNT_ID, + rate_limiter=passthrough_rate_limiter(), + ) assert provider._api_key == "test-cloudflare-token" assert provider._base_url == _BASE_URL diff --git a/tests/providers/test_codestral.py b/tests/providers/test_codestral.py index df565f6f74d7e64e294fb5186197bb010796ccc4..ddec523c3f36cf45f45510c2a99b8f7c7b7f1f73 100644 --- a/tests/providers/test_codestral.py +++ b/tests/providers/test_codestral.py @@ -1,6 +1,5 @@ """Tests for Mistral Codestral provider.""" -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -10,6 +9,7 @@ from free_claude_code.providers.codestral import ( CODESTRAL_DEFAULT_BASE, CodestralProvider, ) +from tests.providers.support import passthrough_rate_limiter class MockMessage: @@ -45,30 +45,9 @@ def codestral_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - """Mock the global rate limiter to prevent waiting.""" - - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def codestral_provider(codestral_config): - return CodestralProvider(codestral_config) + return CodestralProvider(codestral_config, rate_limiter=passthrough_rate_limiter()) def test_init(codestral_config): @@ -76,7 +55,9 @@ def test_init(codestral_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_openai: - provider = CodestralProvider(codestral_config) + provider = CodestralProvider( + codestral_config, rate_limiter=passthrough_rate_limiter() + ) assert provider._api_key == "test_codestral_key" assert provider._base_url == CODESTRAL_DEFAULT_BASE mock_openai.assert_called_once() @@ -104,7 +85,8 @@ def test_build_request_body_global_disable_blocks_reasoning_mapping(): rate_limit=10, rate_window=60, enable_thinking=False, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) req = MockRequest() body = provider._build_request_body(req) diff --git a/tests/providers/test_cohere.py b/tests/providers/test_cohere.py index 42100fa21248d5eb7116ca16fadd788d1d7bf422..264b90abba5371ce6d10606052da67da4a1efc0e 100644 --- a/tests/providers/test_cohere.py +++ b/tests/providers/test_cohere.py @@ -1,6 +1,5 @@ """Tests for Cohere Compatibility API provider.""" -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -8,6 +7,7 @@ import pytest from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.cohere import COHERE_DEFAULT_BASE, CohereProvider from free_claude_code.providers.exceptions import InvalidRequestError +from tests.providers.support import passthrough_rate_limiter class MockMessage: @@ -43,28 +43,9 @@ def cohere_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def cohere_provider(cohere_config): - return CohereProvider(cohere_config) + return CohereProvider(cohere_config, rate_limiter=passthrough_rate_limiter()) def test_default_base_url_constant(): @@ -75,7 +56,9 @@ def test_init_uses_default_base_url_and_api_key(cohere_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_openai: - provider = CohereProvider(cohere_config) + provider = CohereProvider( + cohere_config, rate_limiter=passthrough_rate_limiter() + ) assert provider._api_key == "test_cohere_key" assert provider._base_url == COHERE_DEFAULT_BASE @@ -88,7 +71,7 @@ def test_init_strips_trailing_slash(cohere_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ): - provider = CohereProvider(config) + provider = CohereProvider(config, rate_limiter=passthrough_rate_limiter()) assert provider._base_url == COHERE_DEFAULT_BASE @@ -174,7 +157,8 @@ def test_build_request_body_maps_thinking_disabled_to_reasoning_none(): rate_limit=10, rate_window=60, enable_thinking=False, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) body = provider._build_request_body(MockRequest()) diff --git a/tests/providers/test_deepseek.py b/tests/providers/test_deepseek.py index 0783c3619d78faf7ad0b260f2307393868648284..b85358edf1a3bce2c30a71ce32ce0df972f0ed13 100644 --- a/tests/providers/test_deepseek.py +++ b/tests/providers/test_deepseek.py @@ -1,7 +1,6 @@ """Tests for DeepSeek OpenAI-compatible Chat Completions provider.""" import logging -from contextlib import asynccontextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, patch @@ -21,6 +20,7 @@ from free_claude_code.providers.deepseek import ( DeepSeekProvider, ) from free_claude_code.providers.exceptions import InvalidRequestError +from tests.providers.support import passthrough_rate_limiter @pytest.fixture @@ -34,28 +34,9 @@ def deepseek_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def deepseek_provider(deepseek_config): - return DeepSeekProvider(deepseek_config) + return DeepSeekProvider(deepseek_config, rate_limiter=passthrough_rate_limiter()) def test_default_base_url_alias(): @@ -66,7 +47,9 @@ def test_init(deepseek_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_client: - provider = DeepSeekProvider(deepseek_config) + provider = DeepSeekProvider( + deepseek_config, rate_limiter=passthrough_rate_limiter() + ) assert provider._api_key == "test_deepseek_key" assert provider._base_url == "https://api.deepseek.com" assert mock_client.called @@ -181,7 +164,8 @@ def test_build_request_body_respects_global_thinking_disable(): rate_limit=1, rate_window=1, enable_thinking=False, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) request = MessagesRequest.model_validate( { @@ -480,7 +464,8 @@ def test_thinking_off_strips_thinking_history(): rate_limit=1, rate_window=1, enable_thinking=False, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) request = MessagesRequest.model_validate( { @@ -561,7 +546,8 @@ def test_preflight_strips_user_image(): base_url=DEEPSEEK_DEFAULT_BASE, rate_limit=1, rate_window=1, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) # Should not raise; image is stripped. provider.preflight_stream(request, thinking_enabled=True) @@ -583,7 +569,8 @@ def test_preflight_rejects_mcp_servers(): base_url=DEEPSEEK_DEFAULT_BASE, rate_limit=1, rate_window=1, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) with pytest.raises(InvalidRequestError, match="mcp_servers"): provider.preflight_stream(request) @@ -601,7 +588,8 @@ def test_preflight_rejects_listed_server_tools_in_tools_list(): base_url=DEEPSEEK_DEFAULT_BASE, rate_limit=1, rate_window=1, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) with pytest.raises(InvalidRequestError, match="web_search"): provider.preflight_stream(request) @@ -637,7 +625,8 @@ def test_preflight_rejects_server_tool_result_blocks(): base_url=DEEPSEEK_DEFAULT_BASE, rate_limit=1, rate_window=1, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) with pytest.raises(InvalidRequestError, match=r"web_search_tool_result|server"): provider.preflight_stream(request) diff --git a/tests/providers/test_error_mapping.py b/tests/providers/test_error_mapping.py index 3efa91103e3c748b795f852c2ac18fc2cdfb745a..aae25536d8f65b2486b9cdc66ca39e6e93166d01 100644 --- a/tests/providers/test_error_mapping.py +++ b/tests/providers/test_error_mapping.py @@ -1,7 +1,7 @@ """Tests for provider error mapping and core error formatting.""" from pathlib import Path -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock import openai import pytest @@ -28,6 +28,7 @@ from free_claude_code.providers.exceptions import ( OverloadedError, RateLimitError, ) +from free_claude_code.providers.rate_limit import ProviderRateLimiter def _make_openai_error(cls, message="test error", status_code=None): @@ -48,33 +49,35 @@ def _make_statusless_openai_api_error( return openai.APIError(message, request=Request("POST", "http://test"), body=body) +def _rate_limiter() -> MagicMock: + return MagicMock(spec=ProviderRateLimiter) + + class TestMapError: """Tests for map_error function.""" def test_authentication_error(self): """openai.AuthenticationError -> AuthenticationError.""" exc = _make_openai_error(openai.AuthenticationError, status_code=401) - result = map_error(exc) + result = map_error(exc, rate_limiter=_rate_limiter()) assert isinstance(result, AuthenticationError) assert result.status_code == 401 def test_rate_limit_error(self): """openai.RateLimitError -> RateLimitError and triggers global block.""" exc = _make_openai_error(openai.RateLimitError, status_code=429) - with patch( - "free_claude_code.providers.error_mapping.GlobalRateLimiter" - ) as mock_rl: - mock_instance = MagicMock() - mock_rl.get_instance.return_value = mock_instance - result = map_error(exc) - assert isinstance(result, RateLimitError) - assert result.status_code == 429 - mock_instance.set_blocked.assert_called_once_with(60) + limiter = _rate_limiter() + + result = map_error(exc, rate_limiter=limiter) + + assert isinstance(result, RateLimitError) + assert result.status_code == 429 + limiter.set_blocked.assert_called_once_with(60) def test_bad_request_error(self): """openai.BadRequestError -> InvalidRequestError.""" exc = _make_openai_error(openai.BadRequestError, status_code=400) - result = map_error(exc) + result = map_error(exc, rate_limiter=_rate_limiter()) assert isinstance(result, InvalidRequestError) assert result.status_code == 400 @@ -88,7 +91,7 @@ class TestMapError: exc = _make_openai_error( openai.InternalServerError, message=message, status_code=500 ) - result = map_error(exc) + result = map_error(exc, rate_limiter=_rate_limiter()) assert isinstance(result, OverloadedError) assert result.status_code == 529 @@ -97,7 +100,7 @@ class TestMapError: exc = _make_openai_error( openai.InternalServerError, message="Unknown error", status_code=500 ) - result = map_error(exc) + result = map_error(exc, rate_limiter=_rate_limiter()) assert isinstance(result, APIError) assert result.status_code == 500 @@ -120,7 +123,7 @@ class TestMapError: message=f"upstream {status_code}", status_code=status_code, ) - result = map_error(exc) + result = map_error(exc, rate_limiter=_rate_limiter()) assert isinstance(result, APIError) assert result.status_code == status_code assert expect_substr in result.message.lower() @@ -130,7 +133,7 @@ class TestMapError: exc = _make_openai_error( openai.APIError, message="Bad gateway", status_code=502 ) - result = map_error(exc) + result = map_error(exc, rate_limiter=_rate_limiter()) assert isinstance(result, APIError) def test_statusless_api_error_resource_exhausted_maps_to_overloaded(self): @@ -140,7 +143,7 @@ class TestMapError: {"error": {"message": "ResourceExhausted: limit reached", "code": 500}}, ) - result = map_error(exc) + result = map_error(exc, rate_limiter=_rate_limiter()) assert isinstance(result, OverloadedError) assert result.status_code == 529 @@ -166,7 +169,7 @@ class TestMapError: {"error": {"message": "unknown provider failure"}}, ) - result = map_error(exc) + result = map_error(exc, rate_limiter=_rate_limiter()) assert isinstance(result, APIError) assert result.status_code == 500 @@ -175,14 +178,14 @@ class TestMapError: def test_unmapped_exception_passthrough(self): """Non-openai exceptions are returned as-is.""" exc = RuntimeError("unexpected") - result = map_error(exc) + result = map_error(exc, rate_limiter=_rate_limiter()) assert result is exc assert isinstance(result, RuntimeError) def test_value_error_passthrough(self): """ValueError passes through unchanged.""" exc = ValueError("bad value") - result = map_error(exc) + result = map_error(exc, rate_limiter=_rate_limiter()) assert result is exc @@ -228,7 +231,7 @@ def test_openai_bad_request_body_is_user_visible(): message="Thinking mode does not support this tool_choice", status_code=400, ) - mapped = map_error(exc) + mapped = map_error(exc, rate_limiter=_rate_limiter()) msg = format_provider_error_message( mapped, extract_provider_error_detail(exc), @@ -259,7 +262,7 @@ def test_auth_status_with_model_error_body_is_not_only_check_api_key(): response=Response(status_code=401, request=Request("POST", "http://test")), body=body, ) - mapped = map_error(exc) + mapped = map_error(exc, rate_limiter=_rate_limiter()) msg = format_provider_error_message( mapped, extract_provider_error_detail(exc), @@ -282,7 +285,7 @@ def test_http_status_error_json_body_is_compact_and_visible(): json={"error": {"type": "BadRequest", "message": "bad field"}}, ) exc = HTTPStatusError("Bad Request", request=response.request, response=response) - mapped = map_error(exc) + mapped = map_error(exc, rate_limiter=_rate_limiter()) msg = user_visible_message_for_mapped_provider_error( mapped, provider_name="LOCAL", @@ -327,7 +330,7 @@ def test_empty_http_error_body_is_explicitly_reported(): content=b"", ) exc = HTTPStatusError("Server Error", request=response.request, response=response) - mapped = map_error(exc) + mapped = map_error(exc, rate_limiter=_rate_limiter()) msg = format_provider_error_message( mapped, extract_provider_error_detail(exc), @@ -346,7 +349,7 @@ def test_connection_error_without_response_includes_sanitized_cause_chain(): "connect failed authorization: Bearer SECRET token=ALSO_SECRET", request=request, ) - mapped = map_error(exc) + mapped = map_error(exc, rate_limiter=_rate_limiter()) detail = extract_provider_error_detail(exc) msg = format_provider_error_message( mapped, @@ -389,7 +392,7 @@ def test_attached_provider_error_body_is_capped_for_display(): ) exc = HTTPStatusError("Server Error", request=response.request, response=response) attach_provider_error_body(exc, "x" * (PROVIDER_ERROR_BODY_DISPLAY_CAP_BYTES + 10)) - mapped = map_error(exc) + mapped = map_error(exc, rate_limiter=_rate_limiter()) msg = format_provider_error_message( mapped, extract_provider_error_detail(exc), @@ -422,4 +425,4 @@ def test_streaming_transports_pass_scoped_rate_limiter_to_map_error(): ): text = path.read_text(encoding="utf-8") assert "map_error(" in text, str(path) - assert "rate_limiter=self._global_rate_limiter" in text, str(path) + assert "rate_limiter=self._rate_limiter" in text, str(path) diff --git a/tests/providers/test_fireworks.py b/tests/providers/test_fireworks.py index 4596f2cdf49dd30788218e55e1206ce8ef257f30..4615e1d079d579eebb7520fd9e757bf2d33ac1aa 100644 --- a/tests/providers/test_fireworks.py +++ b/tests/providers/test_fireworks.py @@ -1,7 +1,6 @@ """Tests for the Fireworks AI OpenAI-chat provider.""" -from contextlib import asynccontextmanager -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock import pytest @@ -11,25 +10,7 @@ from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.exceptions import InvalidRequestError from free_claude_code.providers.fireworks import FIREWORKS_BASE_URL, FireworksProvider from free_claude_code.providers.transports.openai_chat import OpenAIChatTransport - - -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance +from tests.providers.support import passthrough_rate_limiter @pytest.fixture @@ -41,7 +22,8 @@ def fireworks_provider(): rate_limit=10, rate_window=60, enable_thinking=True, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) @@ -92,7 +74,8 @@ def test_build_request_body_global_disable_blocks_thinking(): rate_limit=1, rate_window=1, enable_thinking=False, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) request = MessagesRequest.model_validate( { diff --git a/tests/providers/test_gemini.py b/tests/providers/test_gemini.py index 89f2a4cb6ec81ea223e9f2b92e59e34abcd2a252..2725fdbd8ec2d955b65efa00537794b2e6416c06 100644 --- a/tests/providers/test_gemini.py +++ b/tests/providers/test_gemini.py @@ -1,6 +1,5 @@ """Tests for Google AI Studio Gemini (OpenAI-compatible) provider.""" -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -10,6 +9,7 @@ from free_claude_code.providers.gemini import GEMINI_DEFAULT_BASE, GeminiProvide from free_claude_code.providers.gemini.quirks import ( GEMINI_SKIP_THOUGHT_SIGNATURE_VALIDATOR, ) +from tests.providers.support import passthrough_rate_limiter class MockMessage: @@ -53,30 +53,9 @@ def gemini_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - """Mock the global rate limiter to prevent waiting.""" - - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def gemini_provider(gemini_config): - return GeminiProvider(gemini_config) + return GeminiProvider(gemini_config, rate_limiter=passthrough_rate_limiter()) def test_init(gemini_config): @@ -84,7 +63,9 @@ def test_init(gemini_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_openai: - provider = GeminiProvider(gemini_config) + provider = GeminiProvider( + gemini_config, rate_limiter=passthrough_rate_limiter() + ) assert provider._api_key == "test_gemini_key" assert ( provider._base_url @@ -146,7 +127,8 @@ def test_build_request_body_global_disable_sets_reasoning_none(): rate_limit=10, rate_window=60, enable_thinking=False, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) req = MockRequest() body = provider._build_request_body(req) diff --git a/tests/providers/test_github_models.py b/tests/providers/test_github_models.py index c65e00e588bed5ef91343f8515282db4529d57b1..b9deedb6ab42d6f7e86f3fa8e2f4902634f63d03 100644 --- a/tests/providers/test_github_models.py +++ b/tests/providers/test_github_models.py @@ -1,7 +1,6 @@ """Tests for GitHub Models OpenAI-compatible provider.""" from collections.abc import AsyncIterator -from contextlib import asynccontextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, patch @@ -17,6 +16,7 @@ from free_claude_code.providers.github_models import ( GitHubModelsProvider, ) from free_claude_code.providers.github_models.client import GITHUB_MODELS_CATALOG_URL +from tests.providers.support import passthrough_rate_limiter @pytest.fixture @@ -30,30 +30,13 @@ def github_models_config() -> ProviderConfig: ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def github_models_provider( github_models_config: ProviderConfig, ) -> GitHubModelsProvider: - return GitHubModelsProvider(github_models_config) + return GitHubModelsProvider( + github_models_config, rate_limiter=passthrough_rate_limiter() + ) def _request(model: str = "openai/gpt-4.1") -> MessagesRequest: @@ -94,7 +77,9 @@ def test_init_uses_default_base_url_api_key_and_github_headers( with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_openai: - provider = GitHubModelsProvider(github_models_config) + provider = GitHubModelsProvider( + github_models_config, rate_limiter=passthrough_rate_limiter() + ) assert provider._api_key == "test-github-models-token" assert provider._base_url == GITHUB_MODELS_DEFAULT_BASE @@ -115,7 +100,7 @@ def test_init_strips_trailing_slash(github_models_config: ProviderConfig) -> Non with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ): - provider = GitHubModelsProvider(config) + provider = GitHubModelsProvider(config, rate_limiter=passthrough_rate_limiter()) assert provider._base_url == GITHUB_MODELS_DEFAULT_BASE diff --git a/tests/providers/test_groq.py b/tests/providers/test_groq.py index 55445ab8ed8e05617efe7f3e2acfd76a398606c6..db0e9f5c82b28ab7d4de7658044e28ec90385fc7 100644 --- a/tests/providers/test_groq.py +++ b/tests/providers/test_groq.py @@ -1,12 +1,12 @@ """Tests for Groq (OpenAI-compatible) provider.""" -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import pytest from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.groq import GROQ_DEFAULT_BASE, GroqProvider +from tests.providers.support import passthrough_rate_limiter class MockMessage: @@ -42,30 +42,9 @@ def groq_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - """Mock the global rate limiter to prevent waiting.""" - - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def groq_provider(groq_config): - return GroqProvider(groq_config) + return GroqProvider(groq_config, rate_limiter=passthrough_rate_limiter()) def test_init(groq_config): @@ -73,7 +52,7 @@ def test_init(groq_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_openai: - provider = GroqProvider(groq_config) + provider = GroqProvider(groq_config, rate_limiter=passthrough_rate_limiter()) assert provider._api_key == "test_groq_key" assert provider._base_url == GROQ_DEFAULT_BASE mock_openai.assert_called_once() @@ -101,7 +80,8 @@ def test_build_request_body_global_disable_blocks_reasoning_mapping(): rate_limit=10, rate_window=60, enable_thinking=False, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) req = MockRequest() body = provider._build_request_body(req) diff --git a/tests/providers/test_huggingface.py b/tests/providers/test_huggingface.py index 7851e6eca0e26e1a9f3778d695996f821dc23f4a..57ef35c92faadce9c3aa481cfca45e97537a1479 100644 --- a/tests/providers/test_huggingface.py +++ b/tests/providers/test_huggingface.py @@ -1,6 +1,5 @@ """Tests for Hugging Face Inference Providers.""" -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -11,6 +10,7 @@ from free_claude_code.providers.huggingface import ( HUGGINGFACE_DEFAULT_BASE, HuggingFaceProvider, ) +from tests.providers.support import passthrough_rate_limiter class MockMessage: @@ -53,28 +53,11 @@ def huggingface_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def huggingface_provider(huggingface_config): - return HuggingFaceProvider(huggingface_config) + return HuggingFaceProvider( + huggingface_config, rate_limiter=passthrough_rate_limiter() + ) def test_default_base_url_constant(): @@ -85,7 +68,9 @@ def test_init_uses_default_base_url_and_api_key(huggingface_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_openai: - provider = HuggingFaceProvider(huggingface_config) + provider = HuggingFaceProvider( + huggingface_config, rate_limiter=passthrough_rate_limiter() + ) assert provider._api_key == "test_hf_key" assert provider._base_url == HUGGINGFACE_DEFAULT_BASE @@ -100,7 +85,7 @@ def test_init_strips_trailing_slash(huggingface_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ): - provider = HuggingFaceProvider(config) + provider = HuggingFaceProvider(config, rate_limiter=passthrough_rate_limiter()) assert provider._base_url == HUGGINGFACE_DEFAULT_BASE diff --git a/tests/providers/test_kimi.py b/tests/providers/test_kimi.py index 23a577d884ade4925906d8c8c1a5177646c6da05..fd9c85bd756fdbe877a77863c3055c6d6570cc5a 100644 --- a/tests/providers/test_kimi.py +++ b/tests/providers/test_kimi.py @@ -1,8 +1,7 @@ """Tests for the Kimi OpenAI-chat provider.""" -from contextlib import asynccontextmanager from types import SimpleNamespace -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock import pytest @@ -13,25 +12,7 @@ from free_claude_code.providers.defaults import KIMI_DEFAULT_BASE from free_claude_code.providers.exceptions import InvalidRequestError from free_claude_code.providers.kimi import KimiProvider from free_claude_code.providers.transports.openai_chat import OpenAIChatTransport - - -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance +from tests.providers.support import passthrough_rate_limiter @pytest.fixture @@ -43,7 +24,8 @@ def kimi_provider(): rate_limit=10, rate_window=60, enable_thinking=True, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) diff --git a/tests/providers/test_llamacpp.py b/tests/providers/test_llamacpp.py index d674ff5129fe86534b5fc03f7c3950dd26b2b4a4..ff6d712029a5bd2f9c27809d1b4e51bb92842f9f 100644 --- a/tests/providers/test_llamacpp.py +++ b/tests/providers/test_llamacpp.py @@ -10,6 +10,7 @@ from free_claude_code.core.anthropic.stream_contracts import parse_sse_text from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.exceptions import ProviderError from free_claude_code.providers.llamacpp import LlamaCppProvider +from tests.providers.support import passthrough_rate_limiter class MockMessage: @@ -64,31 +65,17 @@ def llamacpp_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - """Mock the global rate limiter to prevent waiting.""" - with patch( - "free_claude_code.providers.transports.anthropic_messages.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - instance.wait_if_blocked = AsyncMock(return_value=False) - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - yield instance - - @pytest.fixture def llamacpp_provider(llamacpp_config): - return LlamaCppProvider(llamacpp_config) + return LlamaCppProvider(llamacpp_config, rate_limiter=passthrough_rate_limiter()) def test_init(llamacpp_config): """Test provider initialization.""" with patch("httpx.AsyncClient"): - provider = LlamaCppProvider(llamacpp_config) + provider = LlamaCppProvider( + llamacpp_config, rate_limiter=passthrough_rate_limiter() + ) assert provider._base_url == "http://localhost:8080/v1" assert provider._provider_name == "LLAMACPP" @@ -103,7 +90,7 @@ def test_init_uses_configurable_timeouts(): http_connect_timeout=5.0, ) with patch("httpx.AsyncClient") as mock_client: - LlamaCppProvider(config) + LlamaCppProvider(config, rate_limiter=passthrough_rate_limiter()) call_kwargs = mock_client.call_args[1] timeout = call_kwargs["timeout"] assert timeout.read == 600.0 @@ -120,14 +107,15 @@ def test_init_base_url_strips_trailing_slash(): rate_window=60, ) with patch("httpx.AsyncClient"): - provider = LlamaCppProvider(config) + provider = LlamaCppProvider(config, rate_limiter=passthrough_rate_limiter()) assert provider._base_url == "http://localhost:8080/v1" @pytest.mark.asyncio async def test_stream_response_omits_thinking_when_globally_disabled(llamacpp_config): provider = LlamaCppProvider( - llamacpp_config.model_copy(update={"enable_thinking": False}) + llamacpp_config.model_copy(update={"enable_thinking": False}), + rate_limiter=passthrough_rate_limiter(), ) req = MockRequest() @@ -338,7 +326,7 @@ def test_build_request_body_disabled_thinking_strips_native_thinking_history( ): """With thinking disabled, prior assistant thinking/redacted blocks are omitted.""" config = llamacpp_config.model_copy(update={"enable_thinking": False}) - provider = LlamaCppProvider(config) + provider = LlamaCppProvider(config, rate_limiter=passthrough_rate_limiter()) messages = [ MockMessage("user", "Hi"), MockMessage( diff --git a/tests/providers/test_lmstudio.py b/tests/providers/test_lmstudio.py index ab3005e327541a6f846d53b0d4bb1f40584392e0..4c8293a22c9311e0bd777cddd17acc953dd0a947 100644 --- a/tests/providers/test_lmstudio.py +++ b/tests/providers/test_lmstudio.py @@ -1,6 +1,5 @@ """Tests for LM Studio (OpenAI-compatible chat completions) provider.""" -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -10,6 +9,7 @@ from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.exceptions import InvalidRequestError from free_claude_code.providers.lmstudio import LMStudioProvider from free_claude_code.providers.lmstudio.client import LMSTUDIO_DEFAULT_BASE +from tests.providers.support import passthrough_rate_limiter class MockMessage: @@ -44,30 +44,9 @@ def lmstudio_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - """Mock the global rate limiter to prevent waiting.""" - - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def lmstudio_provider(lmstudio_config): - return LMStudioProvider(lmstudio_config) + return LMStudioProvider(lmstudio_config, rate_limiter=passthrough_rate_limiter()) def test_init(lmstudio_config): @@ -75,7 +54,9 @@ def test_init(lmstudio_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_openai: - provider = LMStudioProvider(lmstudio_config) + provider = LMStudioProvider( + lmstudio_config, rate_limiter=passthrough_rate_limiter() + ) assert provider._api_key == "lm-studio" assert provider._base_url == LMSTUDIO_DEFAULT_BASE assert provider._provider_name == "LMSTUDIO" diff --git a/tests/providers/test_minimax.py b/tests/providers/test_minimax.py index 034b8c3e1ef11cbbf4272392efb2cc186f0f33a0..3168f4cfabfb4ba44a6343408dac32578e76cf43 100644 --- a/tests/providers/test_minimax.py +++ b/tests/providers/test_minimax.py @@ -1,6 +1,5 @@ """Tests for the MiniMax OpenAI-chat provider.""" -from contextlib import asynccontextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch @@ -16,6 +15,7 @@ from free_claude_code.core.anthropic.stream_contracts import ( from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.minimax import MINIMAX_DEFAULT_BASE, MiniMaxProvider from free_claude_code.providers.transports.openai_chat import OpenAIChatTransport +from tests.providers.support import passthrough_rate_limiter class AsyncStream: @@ -34,25 +34,6 @@ class AsyncStream: self.closed = True -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def minimax_provider(): return MiniMaxProvider( @@ -62,7 +43,8 @@ def minimax_provider(): rate_limit=10, rate_window=60, enable_thinking=True, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) diff --git a/tests/providers/test_mistral.py b/tests/providers/test_mistral.py index c6a011d8781ad41bb877f2fde4b2cde4a332a383..1b9e86112a7e3ef25458c5b72056dd9a16331f98 100644 --- a/tests/providers/test_mistral.py +++ b/tests/providers/test_mistral.py @@ -1,6 +1,5 @@ """Tests for Mistral La Plateforme provider.""" -from contextlib import asynccontextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch @@ -11,6 +10,7 @@ from httpx import Request, Response from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.exceptions import ProviderError from free_claude_code.providers.mistral import MISTRAL_DEFAULT_BASE, MistralProvider +from tests.providers.support import passthrough_rate_limiter class MockMessage: @@ -65,30 +65,9 @@ def mistral_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - """Mock the global rate limiter to prevent waiting.""" - - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def mistral_provider(mistral_config): - return MistralProvider(mistral_config) + return MistralProvider(mistral_config, rate_limiter=passthrough_rate_limiter()) def test_init(mistral_config): @@ -96,7 +75,9 @@ def test_init(mistral_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_openai: - provider = MistralProvider(mistral_config) + provider = MistralProvider( + mistral_config, rate_limiter=passthrough_rate_limiter() + ) assert provider._api_key == "test_mistral_key" assert provider._base_url == MISTRAL_DEFAULT_BASE mock_openai.assert_called_once() @@ -188,7 +169,8 @@ def test_build_request_body_global_disable_blocks_reasoning_mapping(): rate_limit=10, rate_window=60, enable_thinking=False, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) req = MockRequest() body = provider._build_request_body(req) @@ -205,7 +187,8 @@ def test_build_request_body_thinking_disabled_strips_prior_mistral_thinking(): rate_limit=10, rate_window=60, enable_thinking=False, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) req = MockRequest( system=None, diff --git a/tests/providers/test_model_validation.py b/tests/providers/test_model_validation.py index 3b3d2f3d6f23be723b71e93e1f20989dcc3fc26d..2342a904aad88697d9cac66190b8125a79a00ebb 100644 --- a/tests/providers/test_model_validation.py +++ b/tests/providers/test_model_validation.py @@ -24,6 +24,7 @@ from free_claude_code.providers.runtime import ProviderRuntime from free_claude_code.providers.runtime.model_cache import ProviderModelCache from free_claude_code.providers.wafer import WaferProvider from free_claude_code.runtime.provider_manager import ProviderRuntimeManager +from tests.providers.support import passthrough_rate_limiter def _settings( @@ -79,7 +80,9 @@ async def test_nim_lists_openai_compatible_model_ids() -> None: with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ): - provider = NvidiaNimProvider(config, nim_settings=NimSettings()) + provider = NvidiaNimProvider( + config, nim_settings=NimSettings(), rate_limiter=passthrough_rate_limiter() + ) with patch.object( provider._client.models, @@ -93,7 +96,8 @@ async def test_nim_lists_openai_compatible_model_ids() -> None: @pytest.mark.asyncio async def test_native_anthropic_messages_provider_lists_model_ids() -> None: provider = LlamaCppProvider( - ProviderConfig(api_key="llamacpp", base_url="http://localhost:8080/v1") + ProviderConfig(api_key="llamacpp", base_url="http://localhost:8080/v1"), + rate_limiter=passthrough_rate_limiter(), ) with patch.object( provider._client, @@ -108,7 +112,9 @@ async def test_native_anthropic_messages_provider_lists_model_ids() -> None: @pytest.mark.asyncio async def test_deepseek_lists_models_from_root_endpoint() -> None: - provider = DeepSeekProvider(ProviderConfig(api_key="deepseek-key")) + provider = DeepSeekProvider( + ProviderConfig(api_key="deepseek-key"), rate_limiter=passthrough_rate_limiter() + ) with patch.object( provider._client.models, "list", @@ -122,7 +128,9 @@ async def test_deepseek_lists_models_from_root_endpoint() -> None: @pytest.mark.asyncio async def test_wafer_lists_models_from_default_models_endpoint() -> None: - provider = WaferProvider(ProviderConfig(api_key="wafer-key")) + provider = WaferProvider( + ProviderConfig(api_key="wafer-key"), rate_limiter=passthrough_rate_limiter() + ) with patch.object( provider._client.models, "list", @@ -136,7 +144,10 @@ async def test_wafer_lists_models_from_default_models_endpoint() -> None: @pytest.mark.asyncio async def test_openrouter_lists_only_tool_capable_models() -> None: - provider = OpenRouterProvider(ProviderConfig(api_key="open-router-key")) + provider = OpenRouterProvider( + ProviderConfig(api_key="open-router-key"), + rate_limiter=passthrough_rate_limiter(), + ) with patch.object( provider._client.models, "list", @@ -168,7 +179,10 @@ async def test_openrouter_lists_only_tool_capable_models() -> None: @pytest.mark.asyncio async def test_openrouter_lists_tool_metadata_with_thinking_support() -> None: - provider = OpenRouterProvider(ProviderConfig(api_key="open-router-key")) + provider = OpenRouterProvider( + ProviderConfig(api_key="open-router-key"), + rate_limiter=passthrough_rate_limiter(), + ) with patch.object( provider._client.models, "list", @@ -206,7 +220,10 @@ async def test_openrouter_lists_tool_metadata_with_thinking_support() -> None: @pytest.mark.asyncio async def test_openrouter_lists_empty_set_when_no_tool_capable_models() -> None: - provider = OpenRouterProvider(ProviderConfig(api_key="open-router-key")) + provider = OpenRouterProvider( + ProviderConfig(api_key="open-router-key"), + rate_limiter=passthrough_rate_limiter(), + ) with patch.object( provider._client.models, "list", @@ -223,7 +240,10 @@ async def test_openrouter_lists_empty_set_when_no_tool_capable_models() -> None: @pytest.mark.asyncio async def test_openrouter_model_metadata_rejects_malformed_ids() -> None: - provider = OpenRouterProvider(ProviderConfig(api_key="open-router-key")) + provider = OpenRouterProvider( + ProviderConfig(api_key="open-router-key"), + rate_limiter=passthrough_rate_limiter(), + ) with ( patch.object( provider._client.models, @@ -241,7 +261,8 @@ async def test_openrouter_model_metadata_rejects_malformed_ids() -> None: @pytest.mark.asyncio async def test_ollama_lists_native_tag_model_ids() -> None: provider = OllamaProvider( - ProviderConfig(api_key="ollama", base_url="http://localhost:11434") + ProviderConfig(api_key="ollama", base_url="http://localhost:11434"), + rate_limiter=passthrough_rate_limiter(), ) with patch.object( provider._client, @@ -267,7 +288,8 @@ async def test_ollama_lists_native_tag_model_ids() -> None: @pytest.mark.asyncio async def test_model_listing_rejects_malformed_payload() -> None: provider = LlamaCppProvider( - ProviderConfig(api_key="llamacpp", base_url="http://localhost:8080/v1") + ProviderConfig(api_key="llamacpp", base_url="http://localhost:8080/v1"), + rate_limiter=passthrough_rate_limiter(), ) with ( patch.object( @@ -284,7 +306,8 @@ async def test_model_listing_rejects_malformed_payload() -> None: @pytest.mark.asyncio async def test_model_listing_raises_http_status_errors() -> None: provider = LlamaCppProvider( - ProviderConfig(api_key="llamacpp", base_url="http://localhost:8080/v1") + ProviderConfig(api_key="llamacpp", base_url="http://localhost:8080/v1"), + rate_limiter=passthrough_rate_limiter(), ) with ( patch.object( diff --git a/tests/providers/test_nvidia_nim.py b/tests/providers/test_nvidia_nim.py index b75825d028df15a1f68a851ca6b4cf4f018c63f7..cff23c252a31ba67633fa48b322c66aa8e8e1189 100644 --- a/tests/providers/test_nvidia_nim.py +++ b/tests/providers/test_nvidia_nim.py @@ -12,6 +12,7 @@ from free_claude_code.providers.nvidia_nim import NvidiaNimProvider from free_claude_code.providers.nvidia_nim.tool_schema import ( NIM_TOOL_ARGUMENT_ALIASES_KEY, ) +from tests.providers.support import passthrough_rate_limiter # Mock data classes @@ -115,30 +116,17 @@ def _make_internal_server_error(message: str) -> openai.InternalServerError: return openai.InternalServerError(message, response=response, body=body) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - """Mock the global rate limiter to prevent waiting.""" - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - instance.wait_if_blocked = AsyncMock(return_value=False) - - # execute_with_retry should call through to the actual function - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - yield instance - - @pytest.mark.asyncio async def test_init(provider_config): """Test provider initialization.""" with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_openai: - provider = NvidiaNimProvider(provider_config, nim_settings=NimSettings()) + provider = NvidiaNimProvider( + provider_config, + nim_settings=NimSettings(), + rate_limiter=passthrough_rate_limiter(), + ) assert provider._api_key == "test_key" assert provider._base_url == "https://test.api.nvidia.com/v1" mock_openai.assert_called_once() @@ -159,7 +147,9 @@ async def test_init_uses_configurable_timeouts(): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_openai: - NvidiaNimProvider(config, nim_settings=NimSettings()) + NvidiaNimProvider( + config, nim_settings=NimSettings(), rate_limiter=passthrough_rate_limiter() + ) call_kwargs = mock_openai.call_args[1] timeout = call_kwargs["timeout"] assert timeout.read == 600.0 @@ -170,7 +160,11 @@ async def test_init_uses_configurable_timeouts(): @pytest.mark.asyncio async def test_build_request_body(provider_config): """Test request body construction.""" - provider = NvidiaNimProvider(provider_config, nim_settings=NimSettings()) + provider = NvidiaNimProvider( + provider_config, + nim_settings=NimSettings(), + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() body = provider._build_request_body(req) @@ -195,6 +189,7 @@ async def test_build_request_body_omits_reasoning_when_globally_disabled( provider = NvidiaNimProvider( provider_config.model_copy(update={"enable_thinking": False}), nim_settings=NimSettings(), + rate_limiter=passthrough_rate_limiter(), ) req = MockRequest() body = provider._build_request_body(req) @@ -208,7 +203,11 @@ async def test_build_request_body_omits_reasoning_when_globally_disabled( async def test_build_request_body_omits_reasoning_when_request_disables_thinking( provider_config, ): - provider = NvidiaNimProvider(provider_config, nim_settings=NimSettings()) + provider = NvidiaNimProvider( + provider_config, + nim_settings=NimSettings(), + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() req.thinking.enabled = False body = provider._build_request_body(req) @@ -355,6 +354,7 @@ async def test_stream_response_suppresses_thinking_when_disabled(provider_config provider = NvidiaNimProvider( provider_config.model_copy(update={"enable_thinking": False}), nim_settings=NimSettings(), + rate_limiter=passthrough_rate_limiter(), ) req = MockRequest() @@ -397,6 +397,7 @@ async def test_stream_response_retries_without_chat_template(provider_config): provider = NvidiaNimProvider( provider_config, nim_settings=NimSettings(chat_template="custom_template"), + rate_limiter=passthrough_rate_limiter(), ) req = MockRequest(model="mistralai/mixtral-8x7b-instruct-v0.1") @@ -449,7 +450,11 @@ async def test_stream_response_retries_without_chat_template(provider_config): async def test_stream_response_retries_without_chat_template_kwargs_issue_993( provider_config, ): - provider = NvidiaNimProvider(provider_config, nim_settings=NimSettings()) + provider = NvidiaNimProvider( + provider_config, + nim_settings=NimSettings(), + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest(model="mistralai/mistral-small-4-119b-2603") mock_chunk = MagicMock() @@ -500,6 +505,7 @@ async def test_stream_response_does_not_retry_unrelated_bad_request(provider_con provider = NvidiaNimProvider( provider_config, nim_settings=NimSettings(chat_template="custom_template"), + rate_limiter=passthrough_rate_limiter(), ) req = MockRequest(model="mistralai/mixtral-8x7b-instruct-v0.1") diff --git a/tests/providers/test_ollama.py b/tests/providers/test_ollama.py index 0b3da2bae0de1c44ca066309cf94c505a1b9c395..fe515cdaad0a9dd5dc4923ebd11b981fa0a7b80b 100644 --- a/tests/providers/test_ollama.py +++ b/tests/providers/test_ollama.py @@ -9,6 +9,7 @@ from free_claude_code.core.anthropic.stream_contracts import parse_sse_text from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.exceptions import ProviderError from free_claude_code.providers.ollama import OLLAMA_DEFAULT_BASE, OllamaProvider +from tests.providers.support import passthrough_rate_limiter class MockMessage: @@ -61,31 +62,17 @@ def ollama_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - """Mock the global rate limiter to prevent waiting.""" - with patch( - "free_claude_code.providers.transports.anthropic_messages.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - instance.wait_if_blocked = AsyncMock(return_value=False) - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - yield instance - - @pytest.fixture def ollama_provider(ollama_config): - return OllamaProvider(ollama_config) + return OllamaProvider(ollama_config, rate_limiter=passthrough_rate_limiter()) def test_init(ollama_config): """Test provider initialization.""" with patch("httpx.AsyncClient"): - provider = OllamaProvider(ollama_config) + provider = OllamaProvider( + ollama_config, rate_limiter=passthrough_rate_limiter() + ) assert provider._base_url == "http://localhost:11434" assert provider._provider_name == "OLLAMA" assert provider._api_key == "ollama" @@ -95,7 +82,7 @@ def test_init_uses_default_base_url(): """Test that provider uses default root URL when not configured.""" config = ProviderConfig(api_key="ollama", base_url=None) with patch("httpx.AsyncClient"): - provider = OllamaProvider(config) + provider = OllamaProvider(config, rate_limiter=passthrough_rate_limiter()) assert provider._base_url == OLLAMA_DEFAULT_BASE @@ -109,7 +96,7 @@ def test_init_uses_configurable_timeouts(): http_connect_timeout=5.0, ) with patch("httpx.AsyncClient") as mock_client: - OllamaProvider(config) + OllamaProvider(config, rate_limiter=passthrough_rate_limiter()) call_kwargs = mock_client.call_args[1] timeout = call_kwargs["timeout"] assert timeout.read == 600.0 @@ -126,7 +113,7 @@ def test_init_base_url_strips_trailing_slash(): rate_window=60, ) with patch("httpx.AsyncClient"): - provider = OllamaProvider(config) + provider = OllamaProvider(config, rate_limiter=passthrough_rate_limiter()) assert provider._base_url == "http://localhost:11434" @@ -139,7 +126,7 @@ def test_init_uses_default_api_key(): rate_window=60, ) with patch("httpx.AsyncClient"): - provider = OllamaProvider(config) + provider = OllamaProvider(config, rate_limiter=passthrough_rate_limiter()) assert provider._api_key == "ollama" @@ -197,7 +184,8 @@ async def test_stream_response(ollama_provider): async def test_build_request_body_omits_thinking_when_disabled(ollama_config): """Global disable suppresses provider-side thinking.""" provider = OllamaProvider( - ollama_config.model_copy(update={"enable_thinking": False}) + ollama_config.model_copy(update={"enable_thinking": False}), + rate_limiter=passthrough_rate_limiter(), ) req = MockRequest() @@ -212,7 +200,8 @@ def test_build_request_body_disabled_thinking_strips_assistant_thinking_blocks( ): """Prior assistant thinking/redacted blocks are removed when policy is off.""" provider = OllamaProvider( - ollama_config.model_copy(update={"enable_thinking": False}) + ollama_config.model_copy(update={"enable_thinking": False}), + rate_limiter=passthrough_rate_limiter(), ) req = MockRequest( system=None, diff --git a/tests/providers/test_open_router.py b/tests/providers/test_open_router.py index 9cf43b9de67c188cc6a32fd1948e57dfb60c21d9..df4a0429e6922bc2b209387b07f0b13460bc9410 100644 --- a/tests/providers/test_open_router.py +++ b/tests/providers/test_open_router.py @@ -1,6 +1,5 @@ """Tests for the OpenRouter OpenAI-chat provider.""" -from contextlib import asynccontextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch @@ -16,6 +15,7 @@ from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.exceptions import InvalidRequestError from free_claude_code.providers.open_router import OpenRouterProvider from free_claude_code.providers.transports.openai_chat import OpenAIChatTransport +from tests.providers.support import passthrough_rate_limiter class AsyncStream: @@ -58,25 +58,6 @@ class MockRequest: setattr(self, key, value) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def open_router_provider(): return OpenRouterProvider( @@ -85,7 +66,8 @@ def open_router_provider(): base_url="https://openrouter.ai/api/v1", rate_limit=10, rate_window=60, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) diff --git a/tests/providers/test_openai_chat_output_cap.py b/tests/providers/test_openai_chat_output_cap.py index 6e08814f58512a409a7b5f02d1ae06dbc0826953..be9b61a3ca63b2b2e874c4f82f253f582f27e172 100644 --- a/tests/providers/test_openai_chat_output_cap.py +++ b/tests/providers/test_openai_chat_output_cap.py @@ -5,7 +5,6 @@ Covers the pure parse/clamp helpers and the transport behavior that clamps and learns the cap so later requests clamp proactively. """ -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -16,6 +15,7 @@ from free_claude_code.providers.transports.openai_chat.output_cap import ( clamp_output_tokens, parse_output_token_cap, ) +from tests.providers.support import passthrough_rate_limiter class _BadRequest(Exception): @@ -129,25 +129,6 @@ class MockRequest: self.thinking.enabled = False -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def groq_provider(): return GroqProvider( @@ -157,7 +138,8 @@ def groq_provider(): rate_limit=10, rate_window=60, enable_thinking=False, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) diff --git a/tests/providers/test_openai_chat_usage.py b/tests/providers/test_openai_chat_usage.py index 2aec6e57b971c33d9b7586fdb7c748e3f3f5ecba..9d052a677842a821ae34d64477c4bf91be885690 100644 --- a/tests/providers/test_openai_chat_usage.py +++ b/tests/providers/test_openai_chat_usage.py @@ -10,7 +10,6 @@ from httpx import Request, Response from free_claude_code.core.anthropic.stream_contracts import parse_sse_text from free_claude_code.providers.base import ProviderConfig -from free_claude_code.providers.rate_limit import GlobalRateLimiter from free_claude_code.providers.transports.openai_chat import OpenAIChatTransport from free_claude_code.providers.transports.openai_chat.usage import ( clone_without_stream_usage, @@ -18,6 +17,7 @@ from free_claude_code.providers.transports.openai_chat.usage import ( request_stream_usage, usage_int, ) +from tests.providers.support import passthrough_rate_limiter class _UsageTestProvider(OpenAIChatTransport): @@ -32,6 +32,7 @@ class _UsageTestProvider(OpenAIChatTransport): provider_name="USAGE_TEST", base_url="https://provider.example/v1", api_key="test_key", + rate_limiter=passthrough_rate_limiter(), ) def _build_request_body( @@ -148,71 +149,60 @@ def test_stream_usage_rejection_does_not_match_unrelated_400(): @pytest.mark.asyncio async def test_openai_chat_stream_requests_usage_and_uses_provider_prompt_tokens(): - GlobalRateLimiter.reset_instance() - try: - provider = _UsageTestProvider() - request = SimpleNamespace(model="m") - usage = SimpleNamespace(prompt_tokens=22, completion_tokens=4) - create = AsyncMock( - return_value=_stream( - [ - _chunk(content="hello"), - _chunk(finish_reason="stop"), - _chunk(usage=usage), - ] - ) - ) - - with patch.object(provider._client.chat.completions, "create", create): - events = [ - event - async for event in provider.stream_response(request, input_tokens=7) + provider = _UsageTestProvider() + request = SimpleNamespace(model="m") + usage = SimpleNamespace(prompt_tokens=22, completion_tokens=4) + create = AsyncMock( + return_value=_stream( + [ + _chunk(content="hello"), + _chunk(finish_reason="stop"), + _chunk(usage=usage), ] - - create.assert_awaited_once() - await_args = create.await_args - assert await_args is not None - assert await_args.kwargs["stream_options"] == {"include_usage": True} - parsed = parse_sse_text("".join(events)) - start_usage = next( - event.data["message"]["usage"] - for event in parsed - if event.event == "message_start" ) - final_usage = next( - event.data["usage"] for event in parsed if event.event == "message_delta" - ) - assert start_usage["input_tokens"] == 7 - assert final_usage == {"input_tokens": 22, "output_tokens": 4} - finally: - GlobalRateLimiter.reset_instance() + ) + + with patch.object(provider._client.chat.completions, "create", create): + events = [ + event async for event in provider.stream_response(request, input_tokens=7) + ] + + create.assert_awaited_once() + await_args = create.await_args + assert await_args is not None + assert await_args.kwargs["stream_options"] == {"include_usage": True} + parsed = parse_sse_text("".join(events)) + start_usage = next( + event.data["message"]["usage"] + for event in parsed + if event.event == "message_start" + ) + final_usage = next( + event.data["usage"] for event in parsed if event.event == "message_delta" + ) + assert start_usage["input_tokens"] == 7 + assert final_usage == {"input_tokens": 22, "output_tokens": 4} @pytest.mark.asyncio async def test_openai_chat_stream_retries_without_usage_when_option_is_rejected(): - GlobalRateLimiter.reset_instance() - try: - provider = _UsageTestProvider() - body = {"model": "m", "messages": [{"role": "user", "content": "x"}]} - request_stream_usage(body) - create = AsyncMock( - side_effect=[ - _bad_request( - "stream_options is unsupported", - {"error": {"message": "stream_options is unsupported"}}, - ), - object(), - ] - ) + provider = _UsageTestProvider() + body = {"model": "m", "messages": [{"role": "user", "content": "x"}]} + request_stream_usage(body) + create = AsyncMock( + side_effect=[ + _bad_request( + "stream_options is unsupported", + {"error": {"message": "stream_options is unsupported"}}, + ), + object(), + ] + ) + + with patch.object(provider._client.chat.completions, "create", create): + _stream_obj, used_body = await provider._create_stream(body) - with patch.object(provider._client.chat.completions, "create", create): - _stream_obj, used_body = await provider._create_stream(body) - - assert create.await_count == 2 - assert create.await_args_list[0].kwargs["stream_options"] == { - "include_usage": True - } - assert "stream_options" not in create.await_args_list[1].kwargs - assert "stream_options" not in used_body - finally: - GlobalRateLimiter.reset_instance() + assert create.await_count == 2 + assert create.await_args_list[0].kwargs["stream_options"] == {"include_usage": True} + assert "stream_options" not in create.await_args_list[1].kwargs + assert "stream_options" not in used_body diff --git a/tests/providers/test_openai_compat_5xx_retry.py b/tests/providers/test_openai_compat_5xx_retry.py index 0576ff043cebd606cc25f9cd6686d305531862ab..4a5067d38c9430b8e056e445e393142938379546 100644 --- a/tests/providers/test_openai_compat_5xx_retry.py +++ b/tests/providers/test_openai_compat_5xx_retry.py @@ -11,7 +11,7 @@ from free_claude_code.config.nim import NimSettings from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.exceptions import ProviderError from free_claude_code.providers.nvidia_nim import NvidiaNimProvider -from free_claude_code.providers.rate_limit import GlobalRateLimiter +from tests.providers.support import retrying_rate_limiter from tests.providers.test_nvidia_nim import MockRequest @@ -34,140 +34,138 @@ def _connection_error(message: str = "connect failed") -> openai.APIConnectionEr @pytest.mark.parametrize("status_code", [500, 502, 503, 504]) @pytest.mark.asyncio async def test_nim_stream_retries_on_openai_5xx_then_streams(status_code): - GlobalRateLimiter.reset_instance() - try: - config = ProviderConfig( - 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, + config = ProviderConfig( + 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, + ) + provider = NvidiaNimProvider( + config, + nim_settings=NimSettings(), + rate_limiter=retrying_rate_limiter(), + ) + req = MockRequest() + + mock_chunk = MagicMock() + mock_chunk.choices = [ + MagicMock( + delta=MagicMock(content="Hi", reasoning_content=""), + finish_reason="stop", ) - provider = NvidiaNimProvider(config, nim_settings=NimSettings()) - req = MockRequest() - - mock_chunk = MagicMock() - mock_chunk.choices = [ - MagicMock( - delta=MagicMock(content="Hi", reasoning_content=""), - finish_reason="stop", - ) - ] - mock_chunk.usage = None - - async def mock_stream(): - yield mock_chunk - - with ( - patch.object( - provider._client.chat.completions, - "create", - new_callable=AsyncMock, - ) as mock_create, - patch("asyncio.sleep", new_callable=AsyncMock), - ): - mock_create.side_effect = [_internal_5xx(status_code), mock_stream()] - events = [e async for e in provider.stream_response(req)] - - assert mock_create.await_count == 2 - assert any("Hi" in e for e in events) - finally: - GlobalRateLimiter.reset_instance() + ] + mock_chunk.usage = None + + async def mock_stream(): + yield mock_chunk + + with ( + patch.object( + provider._client.chat.completions, + "create", + new_callable=AsyncMock, + ) as mock_create, + patch("asyncio.sleep", new_callable=AsyncMock), + ): + mock_create.side_effect = [_internal_5xx(status_code), mock_stream()] + events = [e async for e in provider.stream_response(req)] + + assert mock_create.await_count == 2 + assert any("Hi" in e for e in events) @pytest.mark.asyncio async def test_nim_stream_retries_on_pre_stream_connection_error_then_streams(): - GlobalRateLimiter.reset_instance() - try: - config = ProviderConfig( - 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, + config = ProviderConfig( + 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, + ) + provider = NvidiaNimProvider( + config, + nim_settings=NimSettings(), + rate_limiter=retrying_rate_limiter(), + ) + req = MockRequest() + + mock_chunk = MagicMock() + mock_chunk.choices = [ + MagicMock( + delta=MagicMock(content="Recovered", reasoning_content=""), + finish_reason="stop", ) - provider = NvidiaNimProvider(config, nim_settings=NimSettings()) - req = MockRequest() - - mock_chunk = MagicMock() - mock_chunk.choices = [ - MagicMock( - delta=MagicMock(content="Recovered", reasoning_content=""), - finish_reason="stop", - ) - ] - mock_chunk.usage = None - - async def mock_stream(): - yield mock_chunk - - with ( - patch.object( - provider._client.chat.completions, - "create", - new_callable=AsyncMock, - ) as mock_create, - patch("asyncio.sleep", new_callable=AsyncMock), - ): - mock_create.side_effect = [_connection_error(), mock_stream()] - events = [e async for e in provider.stream_response(req)] - - assert mock_create.await_count == 2 - assert any("Recovered" in e for e in events) - finally: - GlobalRateLimiter.reset_instance() + ] + mock_chunk.usage = None + + async def mock_stream(): + yield mock_chunk + + with ( + patch.object( + provider._client.chat.completions, + "create", + new_callable=AsyncMock, + ) as mock_create, + patch("asyncio.sleep", new_callable=AsyncMock), + ): + mock_create.side_effect = [_connection_error(), mock_stream()] + events = [e async for e in provider.stream_response(req)] + + assert mock_create.await_count == 2 + assert any("Recovered" in e for e in events) @pytest.mark.asyncio async def test_nim_stream_connection_error_exhausted_emits_cause_chain(): - GlobalRateLimiter.reset_instance() - try: - config = ProviderConfig( - 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, - ) - provider = NvidiaNimProvider(config, nim_settings=NimSettings()) - req = MockRequest() - error = _connection_error("upstream disconnected") - - with ( - patch.object( - provider._client.chat.completions, - "create", - new_callable=AsyncMock, - side_effect=error, - ) as mock_create, - patch("asyncio.sleep", new_callable=AsyncMock), - patch( - "free_claude_code.providers.transports.openai_chat.stream.trace_event" - ) as trace, - pytest.raises(ProviderError) as exc_info, - ): - [e async for e in provider.stream_response(req, request_id="req_conn")] - - assert mock_create.await_count == 5 - error_traces = [ - call.kwargs - for call in trace.call_args_list - if call.kwargs.get("event") == "provider.response.error" - ] - assert error_traces[-1]["request_id"] == "req_conn" - assert error_traces[-1]["exc_type"] == "APIConnectionError" - assert "error_message" not in error_traces[-1] - assert ( - "Caused by:\nConnectError: upstream disconnected" in exc_info.value.message - ) - finally: - GlobalRateLimiter.reset_instance() + config = ProviderConfig( + 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, + ) + provider = NvidiaNimProvider( + config, + nim_settings=NimSettings(), + rate_limiter=retrying_rate_limiter(), + ) + req = MockRequest() + error = _connection_error("upstream disconnected") + + with ( + patch.object( + provider._client.chat.completions, + "create", + new_callable=AsyncMock, + side_effect=error, + ) as mock_create, + patch("asyncio.sleep", new_callable=AsyncMock), + patch( + "free_claude_code.providers.transports.openai_chat.stream.trace_event" + ) as trace, + pytest.raises(ProviderError) as exc_info, + ): + [e async for e in provider.stream_response(req, request_id="req_conn")] + + assert mock_create.await_count == 5 + error_traces = [ + call.kwargs + for call in trace.call_args_list + if call.kwargs.get("event") == "provider.response.error" + ] + assert error_traces[-1]["request_id"] == "req_conn" + assert error_traces[-1]["exc_type"] == "APIConnectionError" + assert "error_message" not in error_traces[-1] + assert "Caused by:\nConnectError: upstream disconnected" in exc_info.value.message @pytest.mark.parametrize( @@ -184,33 +182,33 @@ async def test_nim_stream_openai_5xx_exhausted_emits_user_message( status_code, expect_substr, ): - GlobalRateLimiter.reset_instance() - try: - config = ProviderConfig( - 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, - ) - provider = NvidiaNimProvider(config, nim_settings=NimSettings()) - req = MockRequest() - - with ( - patch.object( - provider._client.chat.completions, - "create", - new_callable=AsyncMock, - ) as mock_create, - patch("asyncio.sleep", new_callable=AsyncMock), - ): - mock_create.side_effect = _internal_5xx(status_code) - with pytest.raises(ProviderError) as exc_info: - [e async for e in provider.stream_response(req)] - - assert mock_create.await_count == 5 - assert expect_substr in exc_info.value.message.lower() - finally: - GlobalRateLimiter.reset_instance() + config = ProviderConfig( + 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, + ) + provider = NvidiaNimProvider( + config, + nim_settings=NimSettings(), + rate_limiter=retrying_rate_limiter(), + ) + req = MockRequest() + + with ( + patch.object( + provider._client.chat.completions, + "create", + new_callable=AsyncMock, + ) as mock_create, + patch("asyncio.sleep", new_callable=AsyncMock), + ): + mock_create.side_effect = _internal_5xx(status_code) + with pytest.raises(ProviderError) as exc_info: + [e async for e in provider.stream_response(req)] + + assert mock_create.await_count == 5 + assert expect_substr in exc_info.value.message.lower() diff --git a/tests/providers/test_opencode.py b/tests/providers/test_opencode.py index ec7d08cf9d30876acd1d9f41148af504657065d1..9b93b9f809846c1cca0f7d1cba017ff81c003b44 100644 --- a/tests/providers/test_opencode.py +++ b/tests/providers/test_opencode.py @@ -3,6 +3,7 @@ from free_claude_code.api.models.anthropic import MessagesRequest from free_claude_code.providers.base import ProviderConfig from free_claude_code.providers.opencode import OpenCodeProvider +from tests.providers.support import passthrough_rate_limiter def test_build_request_body_preserves_empty_reasoning_content() -> None: @@ -13,7 +14,8 @@ def test_build_request_body_preserves_empty_reasoning_content() -> None: rate_limit=1, rate_window=1, enable_thinking=True, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) request = MessagesRequest.model_validate( { diff --git a/tests/providers/test_provider_rate_limit.py b/tests/providers/test_provider_rate_limit.py index 17f6c6a824255e22b0a8f67c8ec8a6b2a114b773..caa7fcb0b5f4105ac7c1b823f556c1cdf1fceaf9 100644 --- a/tests/providers/test_provider_rate_limit.py +++ b/tests/providers/test_provider_rate_limit.py @@ -5,13 +5,12 @@ from unittest.mock import AsyncMock, patch import httpx import openai import pytest -import pytest_asyncio from httpx import Request from free_claude_code.providers.rate_limit import ( DEFAULT_UPSTREAM_MAX_RETRIES, UPSTREAM_TRANSIENT_TOTAL_ATTEMPTS, - GlobalRateLimiter, + ProviderRateLimiter, retryable_upstream_status, retryable_upstream_transport_error, ) @@ -79,14 +78,7 @@ def test_retryable_upstream_transport_error_rejects_request_errors() -> None: class TestProviderRateLimiter: - """Tests for providers.rate_limit.GlobalRateLimiter.""" - - @pytest_asyncio.fixture(autouse=True) - async def reset_limiter(self): - """Reset singleton before each test.""" - GlobalRateLimiter.reset_instance() - yield - GlobalRateLimiter.reset_instance() + """Tests for providers.rate_limit.ProviderRateLimiter.""" @pytest.mark.asyncio async def test_proactive_throttling(self): @@ -95,8 +87,7 @@ class TestProviderRateLimiter: Logic ported from verify_provider_limiter.py """ # Re-init with tight limits: 1 request per 0.25 second - GlobalRateLimiter.reset_instance() - limiter = GlobalRateLimiter.get_instance(rate_limit=1, rate_window=0.25) + limiter = ProviderRateLimiter(rate_limit=1, rate_window=0.25) start_time = time.monotonic() @@ -124,8 +115,7 @@ class TestProviderRateLimiter: Test reactive blocking when set_blocked is called. Logic ported from verify_provider_limiter.py """ - GlobalRateLimiter.reset_instance() - limiter = GlobalRateLimiter.get_instance() + limiter = ProviderRateLimiter() start_time = time.monotonic() @@ -153,7 +143,7 @@ class TestProviderRateLimiter: @pytest.mark.asyncio async def test_set_blocked_zero_immediately_unblocks(self): """set_blocked(0) should not actually block.""" - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) limiter.set_blocked(0) # Should not be blocked since 0 seconds from now is already past @@ -164,13 +154,13 @@ class TestProviderRateLimiter: @pytest.mark.asyncio async def test_remaining_wait_when_not_blocked(self): """remaining_wait() should return 0 when not blocked.""" - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) assert limiter.remaining_wait() == 0 @pytest.mark.asyncio async def test_remaining_wait_decreases(self): """remaining_wait() should decrease over time.""" - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) limiter.set_blocked(2.0) wait1 = limiter.remaining_wait() @@ -183,14 +173,13 @@ class TestProviderRateLimiter: @pytest.mark.asyncio async def test_is_blocked_false_initially(self): """is_blocked() should be False for a fresh limiter.""" - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) assert limiter.is_blocked() is False @pytest.mark.asyncio async def test_high_rate_limit_no_throttling(self): """Very high rate limit should not cause throttling.""" - GlobalRateLimiter.reset_instance() - limiter = GlobalRateLimiter.get_instance(rate_limit=10000, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=10000, rate_window=60) start = time.monotonic() for _ in range(20): @@ -201,24 +190,21 @@ class TestProviderRateLimiter: assert duration < 1.0, f"High rate limit caused throttling: {duration:.2f}s" @pytest.mark.asyncio - async def test_singleton_pattern(self): - """get_instance should return the same object.""" - limiter1 = GlobalRateLimiter.get_instance(rate_limit=10, rate_window=1) - limiter2 = GlobalRateLimiter.get_instance() - assert limiter1 is limiter2 + async def test_instances_are_independent(self): + """Each provider limiter owns independent reactive state.""" + limiter1 = ProviderRateLimiter(rate_limit=10, rate_window=1) + limiter2 = ProviderRateLimiter(rate_limit=10, rate_window=1) + + limiter1.set_blocked(1) - @pytest.mark.asyncio - async def test_reset_instance(self): - """reset_instance should allow creating a new instance.""" - limiter1 = GlobalRateLimiter.get_instance(rate_limit=10, rate_window=1) - GlobalRateLimiter.reset_instance() - limiter2 = GlobalRateLimiter.get_instance(rate_limit=20, rate_window=2) assert limiter1 is not limiter2 + assert limiter1.is_blocked() is True + assert limiter2.is_blocked() is False @pytest.mark.asyncio async def test_wait_if_blocked_returns_false_when_not_blocked(self): """wait_if_blocked should return False when not reactively blocked.""" - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) result = await limiter.wait_if_blocked() assert result is False @@ -228,12 +214,9 @@ class TestProviderRateLimiter: Proactive limiter should enforce a strict rolling window: for any i, t[i+rate_limit] - t[i] >= rate_window (within tolerance). """ - GlobalRateLimiter.reset_instance() rate_limit = 2 rate_window = 0.5 - limiter = GlobalRateLimiter.get_instance( - rate_limit=rate_limit, rate_window=rate_window - ) + limiter = ProviderRateLimiter(rate_limit=rate_limit, rate_window=rate_window) acquired: list[float] = [] @@ -257,16 +240,14 @@ class TestProviderRateLimiter: @pytest.mark.asyncio async def test_init_rate_limit_zero_raises(self): """rate_limit <= 0 raises ValueError.""" - GlobalRateLimiter.reset_instance() with pytest.raises(ValueError, match="rate_limit must be > 0"): - GlobalRateLimiter(rate_limit=0, rate_window=60) + ProviderRateLimiter(rate_limit=0, rate_window=60) @pytest.mark.asyncio async def test_init_rate_window_zero_raises(self): """rate_window <= 0 raises ValueError.""" - GlobalRateLimiter.reset_instance() with pytest.raises(ValueError, match="rate_window must be > 0"): - GlobalRateLimiter(rate_limit=10, rate_window=0) + ProviderRateLimiter(rate_limit=10, rate_window=0) @pytest.mark.asyncio async def test_execute_with_retry_exhaust_retries_raises(self): @@ -274,8 +255,7 @@ class TestProviderRateLimiter: import openai from httpx import Request, Response - GlobalRateLimiter.reset_instance() - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) def make_429(): return openai.RateLimitError( @@ -298,8 +278,7 @@ class TestProviderRateLimiter: import openai from httpx import Request, Response - GlobalRateLimiter.reset_instance() - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) def make_429(): return openai.RateLimitError( @@ -329,7 +308,7 @@ class TestProviderRateLimiter: import httpx from httpx import Request, Response - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) call_count = 0 @@ -358,8 +337,7 @@ class TestProviderRateLimiter: import openai from httpx import Request, Response - GlobalRateLimiter.reset_instance() - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) def make_upstream_error(): return openai.InternalServerError( @@ -390,7 +368,7 @@ class TestProviderRateLimiter: import httpx from httpx import Request, Response - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) call_count = 0 @@ -415,7 +393,7 @@ class TestProviderRateLimiter: @pytest.mark.asyncio async def test_execute_with_retry_succeeds_on_statusless_transient_api_error(self): """Status-less SDK APIError transient markers participate in backoff retry.""" - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) call_count = 0 @@ -439,7 +417,7 @@ class TestProviderRateLimiter: @pytest.mark.asyncio async def test_execute_with_retry_succeeds_on_openai_connection_retry(self): """Pre-response OpenAI SDK connection errors retry then succeed.""" - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) call_count = 0 request = Request("POST", "http://x") @@ -471,7 +449,7 @@ class TestProviderRateLimiter: @pytest.mark.asyncio async def test_execute_with_retry_succeeds_on_httpx_transport_retry(self): """Pre-response HTTPX transport errors retry then succeed.""" - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) call_count = 0 @@ -502,7 +480,7 @@ class TestProviderRateLimiter: @pytest.mark.asyncio async def test_execute_with_retry_exhaust_transport_error_attempts(self): """Transport retries exhaust after the shared 5 total attempts.""" - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) call_count = 0 @@ -534,8 +512,7 @@ class TestProviderRateLimiter: import openai from httpx import Request, Response - GlobalRateLimiter.reset_instance() - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) exc = openai.InternalServerError( "unavailable", @@ -557,8 +534,7 @@ class TestProviderRateLimiter: import httpx from httpx import Request, Response - GlobalRateLimiter.reset_instance() - limiter = GlobalRateLimiter.get_instance(rate_limit=100, rate_window=60) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60) call_count = 0 @@ -582,16 +558,14 @@ class TestProviderRateLimiter: @pytest.mark.asyncio async def test_max_concurrency_zero_raises(self): """max_concurrency <= 0 raises ValueError.""" - GlobalRateLimiter.reset_instance() with pytest.raises(ValueError, match="max_concurrency must be > 0"): - GlobalRateLimiter(rate_limit=10, rate_window=60, max_concurrency=0) + ProviderRateLimiter(rate_limit=10, rate_window=60, max_concurrency=0) @pytest.mark.asyncio async def test_concurrency_slot_limits_simultaneous_streams(self): """At most max_concurrency streams can hold a slot simultaneously.""" - GlobalRateLimiter.reset_instance() max_concurrency = 2 - limiter = GlobalRateLimiter.get_instance( + limiter = ProviderRateLimiter( rate_limit=100, rate_window=60, max_concurrency=max_concurrency ) @@ -620,10 +594,7 @@ class TestProviderRateLimiter: @pytest.mark.asyncio async def test_concurrency_slot_releases_on_exception(self): """Slot is released even when the body raises an exception.""" - GlobalRateLimiter.reset_instance() - limiter = GlobalRateLimiter.get_instance( - rate_limit=100, rate_window=60, max_concurrency=1 - ) + limiter = ProviderRateLimiter(rate_limit=100, rate_window=60, max_concurrency=1) assert limiter._concurrency_sem is not None with pytest.raises(RuntimeError): @@ -634,25 +605,17 @@ class TestProviderRateLimiter: assert limiter._concurrency_sem._value == 1 @pytest.mark.asyncio - async def test_get_instance_passes_max_concurrency(self): - """get_instance forwards max_concurrency to the singleton.""" - GlobalRateLimiter.reset_instance() - limiter = GlobalRateLimiter.get_instance( - rate_limit=10, rate_window=60, max_concurrency=3 - ) + async def test_constructor_sets_max_concurrency(self): + """Constructor applies max_concurrency to an independent limiter.""" + limiter = ProviderRateLimiter(rate_limit=10, rate_window=60, max_concurrency=3) assert limiter._concurrency_sem is not None assert limiter._concurrency_sem._value == 3 @pytest.mark.asyncio - async def test_scoped_instances_are_isolated(self): - """Provider-scoped limiters do not share reactive block state.""" - GlobalRateLimiter.reset_instance() - nim = GlobalRateLimiter.get_scoped_instance( - "nvidia_nim", rate_limit=10, rate_window=60 - ) - openrouter = GlobalRateLimiter.get_scoped_instance( - "open_router", rate_limit=20, rate_window=30 - ) + async def test_provider_owned_instances_are_isolated(self): + """Independent provider limiters do not share reactive block state.""" + nim = ProviderRateLimiter(rate_limit=10, rate_window=60) + openrouter = ProviderRateLimiter(rate_limit=20, rate_window=30) assert nim is not openrouter nim.set_blocked(1.0) diff --git a/tests/providers/test_provider_runtime.py b/tests/providers/test_provider_runtime.py index 7b9fa93ae1e95da993d96dd91677534db19b36bd..6e78e13f9dfd80d3b7dd2ddd69ce8aa496c1f567 100644 --- a/tests/providers/test_provider_runtime.py +++ b/tests/providers/test_provider_runtime.py @@ -1,3 +1,4 @@ +import asyncio import subprocess import sys from unittest.mock import AsyncMock, MagicMock, patch @@ -36,11 +37,13 @@ from free_claude_code.providers.nvidia_nim import NvidiaNimProvider from free_claude_code.providers.ollama import OllamaProvider from free_claude_code.providers.open_router import OpenRouterProvider from free_claude_code.providers.opencode import OpenCodeProvider +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.runtime import ( ProviderRuntime, build_provider_config, create_provider, ) +from free_claude_code.providers.sambanova import SambaNovaProvider from free_claude_code.providers.vercel import ( VERCEL_AI_GATEWAY_DEFAULT_BASE, VercelProvider, @@ -342,9 +345,14 @@ def test_create_provider_instantiates_each_builtin(): cohere_api_key="test_cohere_key", github_models_token="test_github_models_token", kimi_api_key="test_kimi_key", + provider_rate_limit=7, + provider_rate_window=11, + provider_max_concurrency=3, + sambanova_api_key="test_sambanova_key", ) cases = { "nvidia_nim": NvidiaNimProvider, + "open_router": OpenRouterProvider, "mistral": MistralProvider, "mistral_codestral": CodestralProvider, "deepseek": DeepSeekProvider, @@ -365,17 +373,34 @@ def test_create_provider_instantiates_each_builtin(): "zai": ZaiProvider, "gemini": GeminiProvider, "groq": GroqProvider, + "sambanova": SambaNovaProvider, "cerebras": CerebrasProvider, } + sentinel_limiter = MagicMock(spec=ProviderRateLimiter) with ( patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ), patch("httpx.AsyncClient"), + patch( + "free_claude_code.providers.runtime.factory.ProviderRateLimiter", + return_value=sentinel_limiter, + ) as limiter_factory, ): for provider_id, provider_cls in cases.items(): - assert isinstance(create_provider(provider_id, settings), provider_cls) + provider = create_provider(provider_id, settings) + + assert isinstance(provider, provider_cls) + assert provider._rate_limiter is sentinel_limiter + limiter_factory.assert_called_once_with( + rate_limit=7, + rate_window=11, + max_concurrency=3, + ) + limiter_factory.reset_mock() + + assert set(cases) == set(PROVIDER_CATALOG) def test_provider_runtime_caches_by_provider_id(): @@ -390,6 +415,50 @@ def test_provider_runtime_caches_by_provider_id(): assert first is second +def test_provider_runtime_provider_owns_one_limiter() -> None: + runtime = ProviderRuntime(_make_settings()) + + with patch( + "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" + ): + first = runtime.resolve_provider("nvidia_nim") + second = runtime.resolve_provider("nvidia_nim") + + assert isinstance(first, NvidiaNimProvider) + assert isinstance(second, NvidiaNimProvider) + assert first._rate_limiter is second._rate_limiter + + +def test_separate_provider_runtimes_never_share_limiters() -> None: + first_runtime = ProviderRuntime(_make_settings()) + second_runtime = ProviderRuntime(_make_settings()) + + with patch( + "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" + ): + first = first_runtime.resolve_provider("nvidia_nim") + second = second_runtime.resolve_provider("nvidia_nim") + + assert isinstance(first, NvidiaNimProvider) + assert isinstance(second, NvidiaNimProvider) + assert first is not second + assert first._rate_limiter is not second._rate_limiter + + +def test_different_providers_in_one_runtime_have_independent_limiters() -> None: + runtime = ProviderRuntime(_make_settings()) + + with patch( + "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" + ): + nim = runtime.resolve_provider("nvidia_nim") + open_router = runtime.resolve_provider("open_router") + + assert isinstance(nim, NvidiaNimProvider) + assert isinstance(open_router, OpenRouterProvider) + assert nim._rate_limiter is not open_router._rate_limiter + + def test_unknown_provider_raises_unknown_provider_type_error(): with pytest.raises(UnknownProviderTypeError, match="Unknown provider_type"): create_provider("unknown", _make_settings()) @@ -397,7 +466,7 @@ def test_unknown_provider_raises_unknown_provider_type_error(): @pytest.mark.asyncio async def test_provider_runtime_cleanup_runs_all_even_if_one_fails() -> None: - """Every provider gets cleanup; cache is cleared even when one raises.""" + """Successful providers leave the cache while failed providers remain retryable.""" p1 = MagicMock() p1.cleanup = AsyncMock(side_effect=RuntimeError("first")) p2 = MagicMock() @@ -409,9 +478,60 @@ async def test_provider_runtime_cleanup_runs_all_even_if_one_fails() -> None: p1.cleanup.assert_awaited_once() p2.cleanup.assert_awaited_once() - assert not runtime.is_cached("a") + assert runtime.is_cached("a") assert not runtime.is_cached("b") + p1.cleanup = AsyncMock() + await runtime.cleanup() + + p1.cleanup.assert_awaited_once() + assert not runtime.is_cached("a") + + +@pytest.mark.asyncio +async def test_cancelled_cleanup_retains_current_and_unvisited_providers() -> None: + first = MagicMock() + second = MagicMock() + third = MagicMock() + second_started = asyncio.Event() + second_attempts = 0 + + async def cleanup_second() -> None: + nonlocal second_attempts + second_attempts += 1 + if second_attempts == 1: + second_started.set() + await asyncio.Event().wait() + + first.cleanup = AsyncMock() + second.cleanup = AsyncMock(side_effect=cleanup_second) + third.cleanup = AsyncMock() + runtime = ProviderRuntime( + _make_settings(), + {"first": first, "second": second, "third": third}, + ) + cleanup_task = asyncio.create_task(runtime.cleanup()) + await second_started.wait() + + cleanup_task.cancel() + with pytest.raises(asyncio.CancelledError): + await cleanup_task + + assert runtime.is_cached("first") is False + assert runtime.is_cached("second") is True + assert runtime.is_cached("third") is True + first.cleanup.assert_awaited_once_with() + third.cleanup.assert_not_awaited() + + await runtime.cleanup() + + first.cleanup.assert_awaited_once_with() + assert second.cleanup.await_count == 2 + third.cleanup.assert_awaited_once_with() + assert runtime.is_cached("first") is False + assert runtime.is_cached("second") is False + assert runtime.is_cached("third") is False + @pytest.mark.asyncio async def test_provider_runtime_cleanup_exceptiongroup_on_multiple_failures() -> None: @@ -425,5 +545,12 @@ async def test_provider_runtime_cleanup_exceptiongroup_on_multiple_failures() -> await runtime.cleanup() assert len(exc_info.value.exceptions) == 2 + assert runtime.is_cached("x") + assert runtime.is_cached("y") + + p1.cleanup = AsyncMock() + p2.cleanup = AsyncMock() + await runtime.cleanup() + assert not runtime.is_cached("x") assert not runtime.is_cached("y") diff --git a/tests/providers/test_provider_transport_logging.py b/tests/providers/test_provider_transport_logging.py index 1759ab63b6211f7f8e8267b91a7cf3106a81a92d..d97eccb5b4affd5493388b8f7555c250be8b8d21 100644 --- a/tests/providers/test_provider_transport_logging.py +++ b/tests/providers/test_provider_transport_logging.py @@ -17,6 +17,7 @@ from free_claude_code.providers.transports.anthropic_messages import ( stream as native_stream, ) from tests.provider_request_mocks import make_openai_compat_stream_request +from tests.providers.support import passthrough_rate_limiter from tests.providers.test_anthropic_messages import ( FakeResponse, MockRequest, @@ -38,30 +39,14 @@ def provider_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.anthropic_messages.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.mark.asyncio async def test_native_non_200_logs_exclude_body_text_by_default( caplog, provider_config ): - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() response = FakeResponse(status_code=500, text="SECRET_UPSTREAM_BODY") @@ -87,7 +72,10 @@ async def test_native_non_200_logs_exclude_body_text_by_default( @pytest.mark.asyncio async def test_native_non_200_logs_body_when_verbose(caplog, provider_config): provider_config.log_api_error_tracebacks = True - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() response = FakeResponse(status_code=500, text="SECRET_UPSTREAM_BODY") @@ -114,7 +102,10 @@ async def test_native_non_200_verbose_logs_only_capped_error_body( caplog, provider_config ): provider_config.log_api_error_tracebacks = True - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() tail = "SECRET_TAIL_NOT_LOGGED" huge = f"{'A' * (NATIVE_MESSAGES_ERROR_BODY_LOG_CAP_BYTES + 50)}{tail}" @@ -143,7 +134,10 @@ async def test_native_non_200_verbose_logs_only_capped_error_body( async def test_native_non_200_default_does_not_read_oversized_body( caplog, provider_config ): - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() huge = f"{'Z' * 500_000}LEAK_MARKER" response = FakeResponse(status_code=500, text=huge) @@ -171,7 +165,10 @@ async def test_native_non_200_default_does_not_read_oversized_body( async def test_native_stream_failure_logs_exclude_exception_str_by_default( caplog, provider_config ): - provider = NativeProvider(provider_config) + provider = NativeProvider( + provider_config, + rate_limiter=passthrough_rate_limiter(), + ) req = MockRequest() response = FakeResponse( lines=[ @@ -213,7 +210,9 @@ async def test_openai_compat_stream_failure_default_logs_exclude_exception_str(c base_url="http://localhost:1/v1", log_api_error_tracebacks=False, ) - provider = NvidiaNimProvider(config, nim_settings=NimSettings()) + provider = NvidiaNimProvider( + config, nim_settings=NimSettings(), rate_limiter=passthrough_rate_limiter() + ) req = make_openai_compat_stream_request() @asynccontextmanager @@ -228,7 +227,7 @@ async def test_openai_compat_stream_failure_default_logs_exclude_exception_str(c side_effect=RuntimeError("SECRET_OPENAI_COMPAT"), ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "concurrency_slot", _noop_slot, ), @@ -249,7 +248,9 @@ async def test_openai_compat_stream_failure_default_logs_cause_types_only(caplog base_url="http://localhost:1/v1", log_api_error_tracebacks=False, ) - provider = NvidiaNimProvider(config, nim_settings=NimSettings()) + provider = NvidiaNimProvider( + config, nim_settings=NimSettings(), rate_limiter=passthrough_rate_limiter() + ) req = make_openai_compat_stream_request() error = openai.APIConnectionError( request=httpx.Request("POST", "http://localhost:1/v1/chat/completions") @@ -268,7 +269,7 @@ async def test_openai_compat_stream_failure_default_logs_cause_types_only(caplog side_effect=error, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "concurrency_slot", _noop_slot, ), @@ -290,7 +291,9 @@ async def test_openai_compat_stream_failure_respects_verbose_flag(caplog): base_url="http://localhost:1/v1", log_api_error_tracebacks=True, ) - provider = NvidiaNimProvider(config, nim_settings=NimSettings()) + provider = NvidiaNimProvider( + config, nim_settings=NimSettings(), rate_limiter=passthrough_rate_limiter() + ) req = make_openai_compat_stream_request() @asynccontextmanager @@ -305,7 +308,7 @@ async def test_openai_compat_stream_failure_respects_verbose_flag(caplog): side_effect=RuntimeError("SECRET_OPENAI_COMPAT"), ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "concurrency_slot", _noop_slot, ), diff --git a/tests/providers/test_sambanova.py b/tests/providers/test_sambanova.py index ef3fad69acbf4cef5c8ee3f6147c693edab8e0fa..4ba8dc1ebdf68df6afc655626182ac7e11067af3 100644 --- a/tests/providers/test_sambanova.py +++ b/tests/providers/test_sambanova.py @@ -1,6 +1,5 @@ """Tests for SambaNova Cloud (OpenAI-compatible) provider.""" -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -10,6 +9,7 @@ from free_claude_code.providers.sambanova import ( SAMBANOVA_DEFAULT_BASE, SambaNovaProvider, ) +from tests.providers.support import passthrough_rate_limiter class MockMessage: @@ -45,30 +45,9 @@ def sambanova_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - """Mock the global rate limiter to prevent waiting.""" - - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def sambanova_provider(sambanova_config): - return SambaNovaProvider(sambanova_config) + return SambaNovaProvider(sambanova_config, rate_limiter=passthrough_rate_limiter()) def test_default_base_url_constant(): @@ -79,7 +58,9 @@ def test_init_uses_default_base_url_and_api_key(sambanova_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_openai: - provider = SambaNovaProvider(sambanova_config) + provider = SambaNovaProvider( + sambanova_config, rate_limiter=passthrough_rate_limiter() + ) assert provider._api_key == "test_sambanova_key" assert provider._base_url == SAMBANOVA_DEFAULT_BASE @@ -94,7 +75,7 @@ def test_init_strips_trailing_slash(sambanova_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ): - provider = SambaNovaProvider(config) + provider = SambaNovaProvider(config, rate_limiter=passthrough_rate_limiter()) assert provider._base_url == SAMBANOVA_DEFAULT_BASE diff --git a/tests/providers/test_streaming_errors.py b/tests/providers/test_streaming_errors.py index 7d5555ce09e16ad8ac90224f35f72c82aa9caf49..e491ed96fb78ddb7017105760241ef96e1d8c391 100644 --- a/tests/providers/test_streaming_errors.py +++ b/tests/providers/test_streaming_errors.py @@ -30,6 +30,7 @@ from free_claude_code.providers.transports.openai_chat.tool_calls import ( iter_heuristic_tool_use_sse, ) from tests.provider_request_mocks import make_openai_compat_stream_request +from tests.providers.support import passthrough_rate_limiter class AsyncStreamMock: @@ -68,7 +69,11 @@ def _make_provider(): rate_limit=10, rate_window=60, ) - return NvidiaNimProvider(config, nim_settings=NimSettings()) + return NvidiaNimProvider( + config, + nim_settings=NimSettings(), + rate_limiter=passthrough_rate_limiter(), + ) def _make_tool_assembler(provider: NvidiaNimProvider) -> OpenAIToolCallAssembler: @@ -86,7 +91,11 @@ def _make_provider_with_thinking_enabled(enabled: bool): rate_window=60, enable_thinking=enabled, ) - return NvidiaNimProvider(config, nim_settings=NimSettings()) + return NvidiaNimProvider( + config, + nim_settings=NimSettings(), + rate_limiter=passthrough_rate_limiter(), + ) def _make_request(model="test-model", stream=True): @@ -193,7 +202,7 @@ class TestStreamingExceptionHandling: side_effect=RuntimeError("API failed"), ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -217,7 +226,7 @@ class TestStreamingExceptionHandling: side_effect=httpx.ReadTimeout(""), ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -250,7 +259,7 @@ class TestStreamingExceptionHandling: return_value=stream_mock, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -279,7 +288,7 @@ class TestStreamingExceptionHandling: return_value=stream_mock, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -312,7 +321,7 @@ class TestStreamingExceptionHandling: return_value=stream_mock, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -348,7 +357,7 @@ class TestStreamingExceptionHandling: return_value=stream_mock, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -384,7 +393,7 @@ class TestStreamingExceptionHandling: return_value=stream_mock, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -414,7 +423,7 @@ class TestStreamingExceptionHandling: return_value=stream_mock, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -446,7 +455,7 @@ class TestStreamingExceptionHandling: return_value=stream_mock, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -477,7 +486,7 @@ class TestStreamingExceptionHandling: return_value=stream_mock, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -521,7 +530,7 @@ class TestStreamingExceptionHandling: return_value=stream_mock, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -1043,7 +1052,7 @@ class TestStreamingExceptionHandling: return await fn(*args, **kwargs) with patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "execute_with_retry", new_callable=AsyncMock, side_effect=_passthrough, @@ -1297,7 +1306,7 @@ class TestStreamChunkEdgeCases: return_value=stream_mock, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -1333,7 +1342,7 @@ class TestStreamChunkEdgeCases: return_value=stream_mock, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -1364,7 +1373,7 @@ class TestStreamChunkEdgeCases: return_value=stream_mock, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, @@ -1421,7 +1430,7 @@ async def test_openai_compat_stream_ends_with_contract_when_tool_name_never_arri return_value=stream_mock, ), patch.object( - provider._global_rate_limiter, + provider._rate_limiter, "wait_if_blocked", new_callable=AsyncMock, return_value=False, diff --git a/tests/providers/test_subagent_interception.py b/tests/providers/test_subagent_interception.py index ef9575334b764cd11d7702b1e2552f9a2ae3d523..84d82d00c06321164388fb88c445882b133b40ce 100644 --- a/tests/providers/test_subagent_interception.py +++ b/tests/providers/test_subagent_interception.py @@ -10,13 +10,18 @@ from free_claude_code.providers.nvidia_nim import NvidiaNimProvider from free_claude_code.providers.transports.openai_chat.tool_calls import ( OpenAIToolCallAssembler, ) +from tests.providers.support import passthrough_rate_limiter @pytest.mark.asyncio async def test_task_tool_interception(): # Setup provider config = ProviderConfig(api_key="test") - provider = NvidiaNimProvider(config, nim_settings=NimSettings()) + provider = NvidiaNimProvider( + config, + nim_settings=NimSettings(), + rate_limiter=passthrough_rate_limiter(), + ) # Mock request and stream ledger with real StreamBlockLedger request = MagicMock() diff --git a/tests/providers/test_vercel.py b/tests/providers/test_vercel.py index 795f42311311b4ff0d85e201188d9cdc7c689151..0eb131af0e91c1ab641b6b87add3f3d5c0b7aa49 100644 --- a/tests/providers/test_vercel.py +++ b/tests/providers/test_vercel.py @@ -1,6 +1,5 @@ """Tests for Vercel AI Gateway provider.""" -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -10,6 +9,7 @@ from free_claude_code.providers.vercel import ( VERCEL_AI_GATEWAY_DEFAULT_BASE, VercelProvider, ) +from tests.providers.support import passthrough_rate_limiter class MockMessage: @@ -45,28 +45,12 @@ def vercel_config(): ) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def vercel_provider(vercel_config): - return VercelProvider(vercel_config) + return VercelProvider( + vercel_config, + rate_limiter=passthrough_rate_limiter(), + ) def test_default_base_url_constant(): @@ -77,7 +61,10 @@ def test_init_uses_default_base_url_and_api_key(vercel_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ) as mock_openai: - provider = VercelProvider(vercel_config) + provider = VercelProvider( + vercel_config, + rate_limiter=passthrough_rate_limiter(), + ) assert provider._api_key == "test_vercel_key" assert provider._base_url == VERCEL_AI_GATEWAY_DEFAULT_BASE @@ -92,7 +79,10 @@ def test_init_strips_trailing_slash(vercel_config): with patch( "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI" ): - provider = VercelProvider(config) + provider = VercelProvider( + config, + rate_limiter=passthrough_rate_limiter(), + ) assert provider._base_url == VERCEL_AI_GATEWAY_DEFAULT_BASE diff --git a/tests/providers/test_wafer.py b/tests/providers/test_wafer.py index 0a16a1cf199c26100c41e2d9c7178a9d5208ef8f..50a243755804b4955fbd9c2a2660f84e170a8246 100644 --- a/tests/providers/test_wafer.py +++ b/tests/providers/test_wafer.py @@ -1,21 +1,22 @@ """Tests for the Wafer OpenAI-chat provider.""" -from contextlib import asynccontextmanager from typing import Any -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock import pytest from free_claude_code.api.models.anthropic import Message, MessagesRequest, Tool from free_claude_code.config.constants import ANTHROPIC_DEFAULT_MAX_OUTPUT_TOKENS from free_claude_code.providers.base import ProviderConfig +from free_claude_code.providers.rate_limit import ProviderRateLimiter from free_claude_code.providers.transports.openai_chat import OpenAIChatTransport from free_claude_code.providers.wafer import WAFER_DEFAULT_BASE, WaferProvider +from tests.providers.support import passthrough_rate_limiter class CountingWaferProvider(WaferProvider): - def __init__(self, config: ProviderConfig): - super().__init__(config) + def __init__(self, config: ProviderConfig, *, rate_limiter: ProviderRateLimiter): + super().__init__(config, rate_limiter=rate_limiter) self.thinking_checks = 0 def _is_thinking_enabled( @@ -25,25 +26,6 @@ class CountingWaferProvider(WaferProvider): return super()._is_thinking_enabled(request, thinking_enabled) -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - @pytest.fixture def wafer_config(): return ProviderConfig( @@ -56,7 +38,10 @@ def wafer_config(): @pytest.fixture def wafer_provider(wafer_config): - return WaferProvider(wafer_config) + return WaferProvider( + wafer_config, + rate_limiter=passthrough_rate_limiter(), + ) def test_default_base_url(): @@ -123,7 +108,10 @@ def test_build_request_body_preserves_request_disabled_thinking(wafer_provider): def test_build_request_body_resolves_thinking_once(wafer_config): - provider = CountingWaferProvider(wafer_config) + provider = CountingWaferProvider( + wafer_config, + rate_limiter=passthrough_rate_limiter(), + ) request = MessagesRequest.model_validate( { "model": "DeepSeek-V4-Pro", diff --git a/tests/providers/test_zai.py b/tests/providers/test_zai.py index 5e585baa2616336658d0a15a3f2d6f34fe44ff18..14872c559ac9e8a2ba1ab27562aed140af2d275f 100644 --- a/tests/providers/test_zai.py +++ b/tests/providers/test_zai.py @@ -1,7 +1,6 @@ """Tests for the Z.ai OpenAI-chat Coding Plan provider.""" -from contextlib import asynccontextmanager -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock import pytest @@ -12,25 +11,7 @@ from free_claude_code.providers.defaults import ZAI_DEFAULT_BASE from free_claude_code.providers.exceptions import InvalidRequestError from free_claude_code.providers.transports.openai_chat import OpenAIChatTransport from free_claude_code.providers.zai import ZaiProvider - - -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch( - "free_claude_code.providers.transports.openai_chat.transport.GlobalRateLimiter" - ) as mock: - instance = mock.get_scoped_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance +from tests.providers.support import passthrough_rate_limiter @pytest.fixture @@ -42,7 +23,8 @@ def zai_provider(): rate_limit=10, rate_window=60, enable_thinking=True, - ) + ), + rate_limiter=passthrough_rate_limiter(), ) diff --git a/tests/runtime/test_application_runtime.py b/tests/runtime/test_application_runtime.py index e8c04d38d75616062bd63f93ebe1892257a357f7..8313ee8c34f684a1fcd5d3d97482eb53f74bcd22 100644 --- a/tests/runtime/test_application_runtime.py +++ b/tests/runtime/test_application_runtime.py @@ -1,9 +1,15 @@ -from unittest.mock import AsyncMock, patch +import asyncio +from pathlib import Path +from unittest.mock import AsyncMock, MagicMock, patch import pytest from free_claude_code.config.admin.persistence import PreparedAdminUpdate from free_claude_code.config.settings import Settings +from free_claude_code.messaging.platforms.ports import ( + InboundMessageHandler, + MessagingPlatformComponents, +) from free_claude_code.providers.runtime import ProviderRuntime from free_claude_code.runtime.application import ApplicationRuntime from free_claude_code.runtime.provider_manager import ProviderRuntimeManager @@ -34,6 +40,76 @@ class TrackingFactory: return runtime +class TrackingTranscriber: + def __init__(self, events: list[str]) -> None: + self.events = events + self.close_calls = 0 + + async def transcribe(self, file_path: Path) -> str: + assert isinstance(file_path, Path) + return "transcribed" + + async def close(self) -> None: + self.close_calls += 1 + self.events.append("transcriber.close") + + +class FailingTranscriber(TrackingTranscriber): + async def close(self) -> None: + await super().close() + raise RuntimeError("transcriber close failed") + + +class CancelledTranscriber(TrackingTranscriber): + async def close(self) -> None: + await super().close() + raise asyncio.CancelledError + + +class CancellingOnceTranscriber(TrackingTranscriber): + async def close(self) -> None: + await super().close() + if self.close_calls == 1: + raise asyncio.CancelledError + + +class TrackingMessagingRuntime: + name = "tracking" + + def __init__( + self, + events: list[str], + *, + fail_quiesce_once: bool = False, + fail_close_once: bool = False, + ) -> None: + self.events = events + self.fail_quiesce_once = fail_quiesce_once + self.fail_close_once = fail_close_once + + async def start(self) -> None: + self.events.append("messaging.start") + + async def quiesce(self) -> None: + self.events.append("messaging.quiesce") + if self.fail_quiesce_once: + self.fail_quiesce_once = False + raise RuntimeError("quiesce failed") + + async def close(self) -> None: + self.events.append("messaging.close") + if self.fail_close_once: + self.fail_close_once = False + raise RuntimeError("close failed") + + def on_message(self, handler: InboundMessageHandler) -> None: + assert callable(handler) + + @property + def is_connected(self) -> bool: + return True + + def _settings(model: str, *, port: int = 8082) -> Settings: return Settings().model_copy(update={"model": model, "port": port}) @@ -71,7 +147,7 @@ async def test_provider_apply_constructs_before_commit_then_publishes(tmp_path) _settings("nvidia_nim/old"), runtime_factory=factory, ) - runtime = ApplicationRuntime(manager) + runtime = ApplicationRuntime(manager, transcriber=None) prepared = _prepared(_settings("nvidia_nim/new"), tmp_path) factory.events.clear() @@ -111,7 +187,7 @@ async def test_candidate_failure_never_commits_and_preserves_current(tmp_path) - _settings("nvidia_nim/old"), runtime_factory=factory, ) - runtime = ApplicationRuntime(manager) + runtime = ApplicationRuntime(manager, transcriber=None) prepared = _prepared(_settings("nvidia_nim/new"), tmp_path) factory.fail = True @@ -142,7 +218,7 @@ async def test_persistence_failure_closes_candidate_and_preserves_current( _settings("nvidia_nim/old"), runtime_factory=factory, ) - runtime = ApplicationRuntime(manager) + runtime = ApplicationRuntime(manager, transcriber=None) prepared = _prepared(_settings("nvidia_nim/new"), tmp_path) with ( @@ -172,7 +248,11 @@ async def test_restart_required_apply_commits_without_hot_publication(tmp_path) runtime_factory=factory, ) restart = AsyncMock() - runtime = ApplicationRuntime(manager, restart_callback=restart) + runtime = ApplicationRuntime( + manager, + transcriber=None, + restart_callback=restart, + ) prepared = _prepared( _settings("nvidia_nim/old", port=9090), tmp_path, @@ -204,3 +284,308 @@ async def test_restart_required_apply_commits_without_hot_publication(tmp_path) await runtime.request_restart() restart.assert_awaited_once() await manager.close() + + +@pytest.mark.asyncio +async def test_close_drains_messaging_before_transcriber_and_is_idempotent() -> None: + events: list[str] = [] + manager = ProviderRuntimeManager(_settings("nvidia_nim/model")) + transcriber = TrackingTranscriber(events) + runtime = ApplicationRuntime(manager, transcriber=transcriber) + runtime._messaging_runtime = TrackingMessagingRuntime(events) + workflow = MagicMock() + workflow.stop_all_tasks = AsyncMock( + side_effect=lambda: events.append("workflow.stop_all") + ) + workflow.close.side_effect = lambda: events.append("workflow.close") + runtime._messaging_workflow = workflow + runtime._cli_manager = MagicMock() + + assert await runtime.close() is True + assert await runtime.close() is True + + assert events == [ + "messaging.quiesce", + "workflow.stop_all", + "workflow.close", + "messaging.close", + "transcriber.close", + ] + assert transcriber.close_calls == 1 + assert runtime._transcriber is None + assert runtime._messaging_runtime is None + assert runtime._messaging_workflow is None + + +@pytest.mark.asyncio +async def test_close_retains_transcriber_ownership_when_close_fails() -> None: + events: list[str] = [] + manager = ProviderRuntimeManager(_settings("nvidia_nim/model")) + transcriber = FailingTranscriber(events) + runtime = ApplicationRuntime(manager, transcriber=transcriber) + runtime._messaging_runtime = TrackingMessagingRuntime(events) + + assert await runtime.close() is False + + assert events == [ + "messaging.quiesce", + "messaging.close", + "transcriber.close", + ] + assert transcriber.close_calls == 1 + assert runtime._transcriber is transcriber + assert runtime._closed is False + await manager.close() + + +@pytest.mark.asyncio +async def test_close_retries_runtime_before_closing_later_resources() -> None: + events: list[str] = [] + manager = ProviderRuntimeManager(_settings("nvidia_nim/model")) + transcriber = TrackingTranscriber(events) + runtime = ApplicationRuntime(manager, transcriber=transcriber) + messaging = TrackingMessagingRuntime(events, fail_close_once=True) + runtime._messaging_runtime = messaging + + assert await runtime.close() is False + + assert events == ["messaging.quiesce", "messaging.close"] + assert runtime._messaging_runtime is messaging + assert runtime._transcriber is transcriber + assert runtime._closed is False + + assert await runtime.close() is True + + assert events == [ + "messaging.quiesce", + "messaging.close", + "messaging.quiesce", + "messaging.close", + "transcriber.close", + ] + assert runtime._messaging_runtime is None + assert runtime._transcriber is None + assert runtime._closed is True + + +@pytest.mark.asyncio +async def test_close_retries_workflow_drain_before_closing_delivery() -> None: + events: list[str] = [] + manager = ProviderRuntimeManager(_settings("nvidia_nim/model")) + runtime = ApplicationRuntime(manager, transcriber=None) + messaging = TrackingMessagingRuntime(events) + workflow = MagicMock() + workflow.stop_all_tasks = AsyncMock(side_effect=[RuntimeError("drain failed"), 0]) + workflow.close.side_effect = lambda: events.append("workflow.close") + runtime._messaging_runtime = messaging + runtime._messaging_workflow = workflow + + assert await runtime.close() is False + + assert events == ["messaging.quiesce"] + assert runtime._messaging_runtime is messaging + assert runtime._messaging_workflow is workflow + assert runtime._closed is False + + assert await runtime.close() is True + + assert events == [ + "messaging.quiesce", + "messaging.quiesce", + "workflow.close", + "messaging.close", + ] + assert runtime._closed is True + + +@pytest.mark.asyncio +async def test_close_does_not_drain_workflow_until_ingress_is_quiescent() -> None: + events: list[str] = [] + manager = ProviderRuntimeManager(_settings("nvidia_nim/model")) + runtime = ApplicationRuntime(manager, transcriber=None) + messaging = TrackingMessagingRuntime(events, fail_quiesce_once=True) + workflow = MagicMock() + workflow.stop_all_tasks = AsyncMock( + side_effect=lambda: events.append("workflow.stop_all") + ) + workflow.close.side_effect = lambda: events.append("workflow.close") + runtime._messaging_runtime = messaging + runtime._messaging_workflow = workflow + + assert await runtime.close() is False + + workflow.stop_all_tasks.assert_not_awaited() + assert runtime._messaging_runtime is messaging + assert runtime._messaging_workflow is workflow + + assert await runtime.close() is True + + assert events == [ + "messaging.quiesce", + "messaging.quiesce", + "workflow.stop_all", + "workflow.close", + "messaging.close", + ] + assert runtime._closed is True + + +@pytest.mark.asyncio +async def test_close_retries_failed_persistence_before_closing_delivery() -> None: + events: list[str] = [] + manager = ProviderRuntimeManager(_settings("nvidia_nim/model")) + runtime = ApplicationRuntime(manager, transcriber=None) + messaging = TrackingMessagingRuntime(events) + workflow = MagicMock() + workflow.stop_all_tasks = AsyncMock( + side_effect=lambda: events.append("workflow.stop_all") + ) + close_calls = 0 + + def close_workflow() -> None: + nonlocal close_calls + close_calls += 1 + if close_calls == 1: + raise RuntimeError("flush failed") + events.append("workflow.close") + + workflow.close.side_effect = close_workflow + runtime._messaging_runtime = messaging + runtime._messaging_workflow = workflow + + await runtime.close() + + assert runtime._messaging_workflow is workflow + assert runtime._messaging_runtime is messaging + assert "messaging.close" not in events + + await runtime.close() + + assert events == [ + "messaging.quiesce", + "workflow.stop_all", + "messaging.quiesce", + "workflow.stop_all", + "workflow.close", + "messaging.close", + ] + assert runtime._closed is True + + +@pytest.mark.asyncio +async def test_cancelled_transcriber_close_retains_ownership() -> None: + events: list[str] = [] + manager = ProviderRuntimeManager(_settings("nvidia_nim/model")) + transcriber = CancelledTranscriber(events) + runtime = ApplicationRuntime(manager, transcriber=transcriber) + + with pytest.raises(asyncio.CancelledError): + await runtime._cleanup_transcriber() + + assert transcriber.close_calls == 1 + assert runtime._transcriber is transcriber + await manager.close() + + +@pytest.mark.asyncio +async def test_cancelled_application_close_remains_retryable() -> None: + events: list[str] = [] + manager = ProviderRuntimeManager(_settings("nvidia_nim/model")) + transcriber = CancellingOnceTranscriber(events) + runtime = ApplicationRuntime(manager, transcriber=transcriber) + + with pytest.raises(asyncio.CancelledError): + await runtime.close() + + assert runtime._closed is False + assert runtime._transcriber is transcriber + + await runtime.close() + + assert transcriber.close_calls == 2 + assert runtime._transcriber is None + assert runtime._closed is True + + +@pytest.mark.asyncio +async def test_startup_failure_closes_owned_transcriber() -> None: + events: list[str] = [] + manager = ProviderRuntimeManager(_settings("nvidia_nim/model")) + transcriber = TrackingTranscriber(events) + runtime = ApplicationRuntime(manager, transcriber=transcriber) + + with ( + patch.object( + manager, + "validate_configured_models", + AsyncMock(side_effect=RuntimeError("startup failed")), + ), + pytest.raises(RuntimeError, match="startup failed"), + ): + await runtime.start() + + assert transcriber.close_calls == 1 + + +@pytest.mark.asyncio +async def test_startup_cancellation_cleans_partial_messaging_and_reraises() -> None: + events: list[str] = [] + manager = ProviderRuntimeManager(_settings("nvidia_nim/model")) + transcriber = TrackingTranscriber(events) + runtime = ApplicationRuntime(manager, transcriber=transcriber) + messaging = TrackingMessagingRuntime(events) + entered = asyncio.Event() + + async def start_messaging() -> None: + runtime._messaging_runtime = messaging + entered.set() + await asyncio.Event().wait() + + with patch.object( + runtime, + "_start_messaging_if_configured", + side_effect=start_messaging, + ): + start_task = asyncio.create_task(runtime.start()) + await entered.wait() + start_task.cancel() + with pytest.raises(asyncio.CancelledError): + await start_task + + assert events == [ + "messaging.quiesce", + "messaging.close", + "transcriber.close", + ] + assert runtime._closed is True + assert runtime._messaging_runtime is None + assert runtime._transcriber is None + + +@pytest.mark.asyncio +async def test_composition_records_runtime_before_workspace_setup() -> None: + events: list[str] = [] + manager = ProviderRuntimeManager(_settings("nvidia_nim/model")) + runtime = ApplicationRuntime(manager, transcriber=None) + messaging = TrackingMessagingRuntime(events) + components = MessagingPlatformComponents( + name="tracking", + runtime=messaging, + outbound=MagicMock(), + ) + + with ( + patch( + "free_claude_code.runtime.application.os.makedirs", + side_effect=OSError("workspace failed"), + ), + pytest.raises(OSError, match="workspace failed"), + ): + await runtime._start_messaging_workflow(components) + + assert runtime._messaging_runtime is messaging + + await runtime.close() + + assert events == ["messaging.quiesce", "messaging.close"] + assert runtime._closed is True diff --git a/tests/runtime/test_provider_manager.py b/tests/runtime/test_provider_manager.py index a0a58933e77ed3aa2a8709b2fa9892c2e71bfcd5..2f6e163ca0cb69d19cf26a335456d7e9b7ac50f5 100644 --- a/tests/runtime/test_provider_manager.py +++ b/tests/runtime/test_provider_manager.py @@ -8,6 +8,7 @@ from free_claude_code.config.settings import Settings from free_claude_code.providers.base import BaseProvider from free_claude_code.providers.exceptions import ServiceUnavailableError from free_claude_code.providers.model_listing import ProviderModelInfo +from free_claude_code.providers.nvidia_nim import NvidiaNimProvider from free_claude_code.providers.runtime import ProviderRuntime from free_claude_code.runtime.provider_manager import ProviderRuntimeManager @@ -17,6 +18,8 @@ class FakeRuntime(ProviderRuntime): self.settings = settings self.cleanup_calls = 0 self.cleanup_error: Exception | None = None + self.cleanup_started: asyncio.Event | None = None + self.cleanup_release: asyncio.Event | None = None self.provider = MagicMock() self.provider.list_model_infos = AsyncMock(return_value=frozenset()) @@ -28,6 +31,10 @@ class FakeRuntime(ProviderRuntime): async def cleanup(self) -> None: self.cleanup_calls += 1 + if self.cleanup_started is not None: + self.cleanup_started.set() + if self.cleanup_release is not None: + await self.cleanup_release.wait() if self.cleanup_error is not None: raise self.cleanup_error @@ -98,6 +105,51 @@ async def test_replacement_keeps_leased_generation_open_until_final_release() -> assert factory.runtimes[1].cleanup_calls == 1 +@pytest.mark.asyncio +async def test_real_hot_replacement_owns_a_limiter_per_provider_generation() -> None: + first_settings = _settings("nvidia_nim/one") + second_settings = _settings("nvidia_nim/two") + clients: list[MagicMock] = [] + + def create_client(*_args: object, **_kwargs: object) -> MagicMock: + client = MagicMock() + client.close = AsyncMock() + clients.append(client) + return client + + with patch( + "free_claude_code.providers.transports.openai_chat.transport.AsyncOpenAI", + side_effect=create_client, + ): + manager = ProviderRuntimeManager(first_settings) + old_lease = await manager.acquire() + old_provider = old_lease.resolve_provider("nvidia_nim") + refresh = AsyncMock() + + with patch.object(manager, "_refresh_generation", refresh): + await manager.replace(second_settings, commit=lambda: None) + new_lease = await manager.acquire() + new_provider = new_lease.resolve_provider("nvidia_nim") + await asyncio.sleep(0) + + assert isinstance(old_provider, NvidiaNimProvider) + assert isinstance(new_provider, NvidiaNimProvider) + assert new_provider is not old_provider + assert new_provider._rate_limiter is not old_provider._rate_limiter + assert old_lease.resolve_provider("nvidia_nim") is old_provider + clients[0].close.assert_not_awaited() + + await new_lease.release() + await old_lease.release() + + clients[0].close.assert_awaited_once() + clients[1].close.assert_not_awaited() + refresh.assert_awaited_once() + await manager.close() + + clients[1].close.assert_awaited_once() + + @pytest.mark.asyncio async def test_replacement_closes_unleased_generation_immediately() -> None: factory = RuntimeFactory() @@ -115,6 +167,120 @@ async def test_replacement_closes_unleased_generation_immediately() -> None: await manager.close() +@pytest.mark.asyncio +async def test_cancelled_replacement_does_not_cancel_owned_generation_cleanup() -> None: + factory = RuntimeFactory() + manager = ProviderRuntimeManager( + _settings("nvidia_nim/one"), + runtime_factory=factory, + ) + cleanup_started = asyncio.Event() + cleanup_release = asyncio.Event() + refresh_started = asyncio.Event() + factory.runtimes[0].cleanup_started = cleanup_started + factory.runtimes[0].cleanup_release = cleanup_release + + async def refresh(*_args: object, **_kwargs: object) -> None: + refresh_started.set() + await asyncio.Event().wait() + + with patch.object(manager, "_refresh_generation", side_effect=refresh): + replace_task = asyncio.create_task( + manager.replace( + _settings("nvidia_nim/two"), + commit=lambda: None, + ) + ) + await cleanup_started.wait() + await refresh_started.wait() + + replace_task.cancel() + with pytest.raises(asyncio.CancelledError): + await replace_task + + retired = manager._retired[1] + assert manager.current_generation_id == 2 + assert retired.cleanup_task is not None + assert not retired.cleanup_task.cancelled() + assert factory.runtimes[0].cleanup_calls == 1 + + close_task = asyncio.create_task(manager.close()) + await asyncio.sleep(0) + assert not close_task.done() + cleanup_release.set() + await close_task + + assert factory.runtimes[0].cleanup_calls == 1 + assert factory.runtimes[1].cleanup_calls == 1 + assert manager._retired == {} + + +@pytest.mark.asyncio +async def test_cancelled_final_lease_release_keeps_owned_cleanup_running() -> None: + factory = RuntimeFactory() + manager = ProviderRuntimeManager( + _settings("nvidia_nim/one"), + runtime_factory=factory, + ) + lease = await manager.acquire() + await manager.replace( + _settings("nvidia_nim/two"), + commit=lambda: None, + ) + cleanup_started = asyncio.Event() + cleanup_release = asyncio.Event() + factory.runtimes[0].cleanup_started = cleanup_started + factory.runtimes[0].cleanup_release = cleanup_release + + release_task = asyncio.create_task(lease.release()) + await cleanup_started.wait() + release_task.cancel() + + with pytest.raises(asyncio.CancelledError): + await release_task + + retired = manager._retired[1] + assert retired.active_leases == 0 + assert retired.cleanup_task is not None + assert not retired.cleanup_task.cancelled() + assert factory.runtimes[0].cleanup_calls == 1 + + close_task = asyncio.create_task(manager.close()) + await asyncio.sleep(0) + assert not close_task.done() + cleanup_release.set() + await close_task + + assert factory.runtimes[0].cleanup_calls == 1 + assert manager._retired == {} + + +@pytest.mark.asyncio +async def test_hot_cleanup_failure_keeps_published_replacement() -> None: + factory = RuntimeFactory() + manager = ProviderRuntimeManager( + _settings("nvidia_nim/one"), + runtime_factory=factory, + ) + factory.runtimes[0].cleanup_error = RuntimeError("cleanup failed") + + generation_id = await manager.replace( + _settings("nvidia_nim/two"), + commit=lambda: None, + ) + + assert generation_id == 2 + assert manager.current_generation_id == 2 + assert factory.runtimes[0].cleanup_calls == 1 + assert 1 in manager._retired + + factory.runtimes[0].cleanup_error = None + await manager.close() + + assert factory.runtimes[0].cleanup_calls == 2 + assert manager._retired == {} + + @pytest.mark.asyncio async def test_candidate_construction_failure_preserves_current_generation() -> None: factory = RuntimeFactory() @@ -137,7 +303,7 @@ async def test_candidate_construction_failure_preserves_current_generation() -> @pytest.mark.asyncio -async def test_persistence_failure_closes_candidate_and_preserves_current() -> None: +async def test_failed_candidate_cleanup_is_retried_at_shutdown() -> None: factory = RuntimeFactory() manager = ProviderRuntimeManager( _settings("nvidia_nim/one"), @@ -145,6 +311,7 @@ async def test_persistence_failure_closes_candidate_and_preserves_current() -> N ) def fail_commit() -> None: + factory.runtimes[1].cleanup_error = RuntimeError("private cleanup detail") raise OSError("disk full") with pytest.raises(OSError, match="disk full"): @@ -156,8 +323,97 @@ async def test_persistence_failure_closes_candidate_and_preserves_current() -> N assert manager.current_generation_id == 1 assert factory.runtimes[0].cleanup_calls == 0 assert factory.runtimes[1].cleanup_calls == 1 + assert manager._unpublished == {factory.runtimes[1]} + + manager.cache_model_infos("nvidia_nim", {ProviderModelInfo("cached")}) + with pytest.raises( + RuntimeError, + match="One or more provider runtimes failed to close", + ) as exc_info: + await manager.close() + + assert "private cleanup detail" not in str(exc_info.value) + assert manager._closed is False + assert manager._unpublished == {factory.runtimes[1]} + assert manager.cached_model_ids() == {"nvidia_nim": frozenset({"cached"})} + + factory.runtimes[1].cleanup_error = None await manager.close() + assert factory.runtimes[0].cleanup_calls == 1 + assert factory.runtimes[1].cleanup_calls == 3 + assert manager._unpublished == set() + assert manager.cached_model_ids() == {} + assert manager._closed is True + + +@pytest.mark.asyncio +async def test_later_replacement_retries_failed_unpublished_candidate() -> None: + factory = RuntimeFactory() + manager = ProviderRuntimeManager( + _settings("nvidia_nim/one"), + runtime_factory=factory, + ) + + def fail_commit() -> None: + factory.runtimes[1].cleanup_error = RuntimeError("cleanup failed") + raise OSError("disk full") + + with pytest.raises(OSError, match="disk full"): + await manager.replace( + _settings("nvidia_nim/two"), + commit=fail_commit, + ) + + factory.runtimes[1].cleanup_error = None + generation_id = await manager.replace( + _settings("nvidia_nim/three"), + commit=lambda: None, + ) + + assert generation_id == 2 + assert factory.runtimes[1].cleanup_calls == 2 + assert manager._unpublished == set() + await manager.close() + + +@pytest.mark.asyncio +async def test_cancelled_candidate_cleanup_remains_owned_until_shutdown() -> None: + factory = RuntimeFactory() + manager = ProviderRuntimeManager( + _settings("nvidia_nim/one"), + runtime_factory=factory, + ) + cleanup_started = asyncio.Event() + cleanup_release = asyncio.Event() + + def fail_commit() -> None: + candidate = factory.runtimes[1] + candidate.cleanup_started = cleanup_started + candidate.cleanup_release = cleanup_release + raise OSError("disk full") + + replace_task = asyncio.create_task( + manager.replace( + _settings("nvidia_nim/two"), + commit=fail_commit, + ) + ) + await cleanup_started.wait() + replace_task.cancel() + + with pytest.raises(asyncio.CancelledError): + await replace_task + + assert manager._unpublished == {factory.runtimes[1]} + assert factory.runtimes[1].cleanup_calls == 1 + + cleanup_release.set() + await manager.close() + + assert factory.runtimes[1].cleanup_calls == 2 + assert manager._unpublished == set() + @pytest.mark.asyncio async def test_concurrent_replacements_are_serialized_in_call_order() -> None: @@ -224,6 +480,104 @@ async def test_shutdown_waits_for_active_lease_then_rejects_new_work() -> None: assert factory.runtimes[0].cleanup_calls == 1 +@pytest.mark.asyncio +async def test_cancelled_shutdown_retains_generation_for_retry() -> None: + factory = RuntimeFactory() + manager = ProviderRuntimeManager( + _settings("nvidia_nim/one"), + runtime_factory=factory, + ) + lease = await manager.acquire() + close_task = asyncio.create_task(manager.close()) + await asyncio.sleep(0) + + close_task.cancel() + with pytest.raises(asyncio.CancelledError): + await close_task + + with pytest.raises(ServiceUnavailableError, match="shutting down"): + await manager.acquire() + with pytest.raises(ServiceUnavailableError, match="shutting down"): + await manager.replace(_settings("nvidia_nim/two"), commit=lambda: None) + assert factory.runtimes[0].cleanup_calls == 0 + + await lease.release() + await manager.close() + + assert factory.runtimes[0].cleanup_calls == 1 + assert manager._closed is True + + +@pytest.mark.asyncio +async def test_cancelled_shutdown_reuses_the_same_owned_cleanup_task() -> None: + factory = RuntimeFactory() + manager = ProviderRuntimeManager( + _settings("nvidia_nim/one"), + runtime_factory=factory, + ) + cleanup_started = asyncio.Event() + cleanup_release = asyncio.Event() + cleanup_calls = 0 + + async def cleanup() -> None: + nonlocal cleanup_calls + cleanup_calls += 1 + cleanup_started.set() + await cleanup_release.wait() + + with patch.object(factory.runtimes[0], "cleanup", side_effect=cleanup): + close_task = asyncio.create_task(manager.close()) + await cleanup_started.wait() + + close_task.cancel() + with pytest.raises(asyncio.CancelledError): + await close_task + + assert manager._closed is False + assert manager._retired + + retry_task = asyncio.create_task(manager.close()) + await asyncio.sleep(0) + assert not retry_task.done() + cleanup_release.set() + await retry_task + + assert cleanup_calls == 1 + assert manager._retired == {} + assert manager._closed is True + + +@pytest.mark.asyncio +async def test_failed_shutdown_cleanup_is_retryable() -> None: + factory = RuntimeFactory() + manager = ProviderRuntimeManager( + _settings("nvidia_nim/one"), + runtime_factory=factory, + ) + manager.cache_model_infos("nvidia_nim", {ProviderModelInfo("cached")}) + factory.runtimes[0].cleanup_error = RuntimeError("private provider detail") + + with pytest.raises( + RuntimeError, + match="One or more provider runtimes failed to close", + ) as exc_info: + await manager.close() + + assert "private provider detail" not in str(exc_info.value) + assert manager._closed is False + assert 1 in manager._retired + assert manager.cached_model_ids() == {"nvidia_nim": frozenset({"cached"})} + assert factory.runtimes[0].cleanup_calls == 1 + + factory.runtimes[0].cleanup_error = None + await manager.close() + + assert factory.runtimes[0].cleanup_calls == 2 + assert manager._retired == {} + assert manager.cached_model_ids() == {} + assert manager._closed is True + + @pytest.mark.asyncio async def test_application_catalog_survives_generation_replacement() -> None: factory = RuntimeFactory() diff --git a/uv.lock b/uv.lock index 958fd0b7f5de6cd817f3a66869c58e076e4a83fc..cf3be302c66b8e608befb485c1cab10febeff07b 100644 --- a/uv.lock +++ b/uv.lock @@ -561,7 +561,7 @@ wheels = [ [[package]] name = "free-claude-code" -version = "3.4.17" +version = "3.4.18" source = { editable = "." } dependencies = [ { name = "aiohttp" },