diff --git a/src/embedding_client.py b/src/embedding_client.py index 1301e279..08a4ff04 100644 --- a/src/embedding_client.py +++ b/src/embedding_client.py @@ -1,16 +1,15 @@ +from __future__ import annotations + import asyncio import logging import threading import time from collections import defaultdict from collections.abc import Awaitable, Callable -from typing import Any, Literal, NamedTuple, TypeVar +from typing import TYPE_CHECKING, Any, Literal, NamedTuple, TypeVar, cast import tiktoken -from google import genai -from google.genai import types as genai_types from nanoid import generate as generate_nanoid -from openai import AsyncOpenAI from .config import ( EmbeddingEncodingFormat, @@ -19,6 +18,10 @@ from .config import ( settings, ) +if TYPE_CHECKING: + from google import genai + from openai import AsyncOpenAI + logger = logging.getLogger(__name__) _T = TypeVar("_T") @@ -189,6 +192,9 @@ class _EmbeddingClient: if self.transport == "gemini": if not config.api_key: raise ValueError("Gemini API key is required") + from google import genai + from google.genai import types as genai_types + # 10-minute HTTP timeout, in lockstep with the LLM registry's Gemini # client (`src/llm/registry.py:_build_gemini_http_options`). Without # this, a stalled Gemini embedding socket wedges the deriver worker @@ -208,6 +214,8 @@ class _EmbeddingClient: else: # openai if not config.api_key: raise ValueError("OpenAI API key is required") + from openai import AsyncOpenAI + self.client = AsyncOpenAI( api_key=config.api_key, base_url=config.base_url, @@ -264,11 +272,11 @@ class _EmbeddingClient: f"Query exceeds maximum token limit of {self.max_embedding_tokens} tokens (got {token_count} tokens)" ) - # Bind the typed client at the dispatch site so pyright can narrow it - # for the closures without needing `assert isinstance(...)` (bandit - # B101). The closures close over the narrowed local, not `self.client`. - if isinstance(self.client, genai.Client): - gemini_client = self.client + # Dispatch on transport rather than isinstance so this module never + # needs the SDK types at runtime; the cast gives the closures a typed + # local to close over. + if self.transport == "gemini": + gemini_client = cast("genai.Client", self.client) async def _call_gemini() -> list[float]: response = await gemini_client.aio.models.embed_content( @@ -290,7 +298,7 @@ class _EmbeddingClient: fn=_call_gemini, ) - openai_client = self.client + openai_client = cast("AsyncOpenAI", self.client) async def _call_openai() -> list[float]: openai_kwargs: dict[str, Any] = {"model": self.model, "input": [query]} @@ -482,8 +490,11 @@ class _EmbeddingClient: attempt is a distinct provider hit and shows up as its own line item in analytics.""" result: dict[str, dict[int, list[float]]] = defaultdict(dict) - if isinstance(self.client, genai.Client): - response = await self.client.aio.models.embed_content( + if self.transport == "gemini": + from google.genai import types as genai_types + + gemini_client = cast("genai.Client", self.client) + response = await gemini_client.aio.models.embed_content( model=self.model, # One Content per item: a list of bare strings is folded # into a single document by gemini-embedding-2*, which @@ -508,7 +519,8 @@ class _EmbeddingClient: self._apply_encoding_format(openai_kwargs) if self.send_dimensions: openai_kwargs["dimensions"] = self.vector_dimensions - response = await self.client.embeddings.create(**openai_kwargs) + openai_client = cast("AsyncOpenAI", self.client) + response = await openai_client.embeddings.create(**openai_kwargs) self._validate_embedding_count(len(batch), len(response.data)) for item, embedding_data in zip(batch, response.data, strict=True): result[item.text_id][item.chunk_index] = ( @@ -617,10 +629,10 @@ class EmbeddingClient: and allowing the application to start even if API keys are not yet configured. """ - _instance: "_EmbeddingClient | None" = None + _instance: _EmbeddingClient | None = None _instance_signature: tuple[object, ...] | None = None _lock: threading.Lock = threading.Lock() - _wrapper_instance: "EmbeddingClient | None" = None + _wrapper_instance: EmbeddingClient | None = None def __new__(cls): """Ensure only one instance of EmbeddingClient exists.""" diff --git a/src/llm/__init__.py b/src/llm/__init__.py index ae47bc53..e0913685 100644 --- a/src/llm/__init__.py +++ b/src/llm/__init__.py @@ -15,6 +15,7 @@ from .registry import ( CLIENTS, backend_for_provider, client_for_model_config, + default_client, get_anthropic_client, get_anthropic_override_client, get_backend, @@ -51,6 +52,7 @@ __all__ = [ "VerbosityType", "backend_for_provider", "client_for_model_config", + "default_client", "default_transport_api_key", "get_anthropic_client", "get_anthropic_override_client", diff --git a/src/llm/executor.py b/src/llm/executor.py index 956bc673..eaa79c1d 100644 --- a/src/llm/executor.py +++ b/src/llm/executor.py @@ -26,7 +26,7 @@ from .backend import CompletionResult as BackendCompletionResult from .backend import StreamChunk as BackendStreamChunk from .backend import ToolCallResult from .capture import build_captured_call, dispatch_captured_call, has_exporters -from .registry import CLIENTS, backend_for_provider +from .registry import backend_for_provider, default_client from .request_builder import execute_completion, execute_stream from .runtime import ( AttemptPlan, @@ -439,7 +439,7 @@ async def honcho_llm_call_inner( post-stream at this layer; aggregate envelopes (DialecticCompletedEvent etc.) carry the accurate totals. """ - client = client_override or CLIENTS.get(provider) + client = client_override or default_client(provider) if client is None: raise ValueError(f"Missing client for {provider}") diff --git a/src/llm/registry.py b/src/llm/registry.py index 48c8191d..9b18bcb8 100644 --- a/src/llm/registry.py +++ b/src/llm/registry.py @@ -9,20 +9,12 @@ history adapter selection) lives here now. from __future__ import annotations from functools import lru_cache -from typing import assert_never - -from anthropic import AsyncAnthropic -from google import genai -from google.genai import types as genai_types -from openai import AsyncOpenAI +from typing import TYPE_CHECKING, assert_never from src.config import ModelConfig, ModelTransport, settings from src.exceptions import ValidationException from .backend import ProviderBackend -from .backends.anthropic import AnthropicBackend -from .backends.gemini import GeminiBackend -from .backends.openai import OpenAIBackend from .credentials import default_transport_api_key from .history_adapters import ( AnthropicHistoryAdapter, @@ -32,6 +24,15 @@ from .history_adapters import ( ) from .types import ProviderClient +if TYPE_CHECKING: + from anthropic import AsyncAnthropic + from google import genai + from google.genai import types as genai_types + from openai import AsyncOpenAI + +# Provider SDKs are imported lazily inside the client factories below so a +# process only pays the import-time memory cost of the providers it uses. + # Default client-level HTTP timeouts. Anthropic accepts seconds (float); # google-genai's HttpOptions.timeout is an int in milliseconds, so the Gemini # value is kept separately. Both default to 10 minutes to match the existing @@ -71,6 +72,8 @@ def _build_gemini_http_options(base_url: str | None) -> genai_types.HttpOptions: timeout even when no ``base_url`` is configured — that's the path the default ``get_gemini_client`` takes and it's the one that was hanging. """ + from google.genai import types as genai_types + return genai_types.HttpOptions( base_url=base_url, timeout=_GEMINI_TIMEOUT_MS, @@ -80,6 +83,8 @@ def _build_gemini_http_options(base_url: str | None) -> genai_types.HttpOptions: @lru_cache(maxsize=1) def get_anthropic_client() -> AsyncAnthropic: """Default Anthropic client built from settings.LLM.ANTHROPIC_API_KEY.""" + from anthropic import AsyncAnthropic + return AsyncAnthropic( api_key=settings.LLM.ANTHROPIC_API_KEY, base_url=settings.LLM.ANTHROPIC_BASE_URL, @@ -90,6 +95,8 @@ def get_anthropic_client() -> AsyncAnthropic: @lru_cache(maxsize=1) def get_openai_client() -> AsyncOpenAI: """Default OpenAI client built from settings.LLM.OPENAI_API_KEY.""" + from openai import AsyncOpenAI + return AsyncOpenAI( api_key=settings.LLM.OPENAI_API_KEY, base_url=settings.LLM.OPENAI_BASE_URL, @@ -100,6 +107,8 @@ def get_openai_client() -> AsyncOpenAI: @lru_cache(maxsize=1) def get_gemini_client() -> genai.Client: """Default Gemini client built from settings.LLM.GEMINI_API_KEY.""" + from google import genai + return genai.Client( api_key=settings.LLM.GEMINI_API_KEY, http_options=_build_gemini_http_options(settings.LLM.GEMINI_BASE_URL), @@ -113,6 +122,8 @@ def get_openai_override_client( base_url: str | None, api_key: str | None ) -> AsyncOpenAI: """OpenAI client for a specific (base_url, api_key) pair. Cached by key.""" + from openai import AsyncOpenAI + return AsyncOpenAI( api_key=api_key, base_url=base_url, @@ -126,6 +137,8 @@ def get_anthropic_override_client( api_key: str | None, ) -> AsyncAnthropic: """Anthropic client for a specific (base_url, api_key) pair. Cached by key.""" + from anthropic import AsyncAnthropic + return AsyncAnthropic( api_key=api_key, base_url=base_url, timeout=_ANTHROPIC_TIMEOUT_S ) @@ -136,35 +149,48 @@ def get_gemini_override_client( base_url: str | None, api_key: str | None ) -> genai.Client: """Gemini client for a specific (base_url, api_key) pair. Cached by key.""" + from google import genai + return genai.Client( api_key=api_key, http_options=_build_gemini_http_options(base_url), ) -# Module-level default-client registry, populated at import time. Tests patch -# this dict via `patch.dict(CLIENTS, {...})` to inject mock provider clients. +# Module-level default-client registry, populated lazily on first use so a +# provider's SDK is only imported when that provider is actually called. Tests +# patch this dict via `patch.dict(CLIENTS, {...})` to inject mock provider +# clients; a patched entry always wins because `default_client` checks the +# dict before constructing anything. CLIENTS: dict[ModelTransport, ProviderClient] = {} -if settings.LLM.ANTHROPIC_API_KEY: - CLIENTS["anthropic"] = AsyncAnthropic( - api_key=settings.LLM.ANTHROPIC_API_KEY, - base_url=settings.LLM.ANTHROPIC_BASE_URL, - timeout=_ANTHROPIC_TIMEOUT_S, - ) -if settings.LLM.OPENAI_API_KEY: - CLIENTS["openai"] = AsyncOpenAI( - api_key=settings.LLM.OPENAI_API_KEY, - base_url=settings.LLM.OPENAI_BASE_URL, - default_headers=_default_headers_for(settings.LLM.OPENAI_BASE_URL), - ) +def default_client(provider: ModelTransport) -> ProviderClient | None: + """Default client for ``provider``, built on first use. -if settings.LLM.GEMINI_API_KEY: - CLIENTS["gemini"] = genai.Client( - api_key=settings.LLM.GEMINI_API_KEY, - http_options=_build_gemini_http_options(settings.LLM.GEMINI_BASE_URL), - ) + Returns None when no API key is configured for the provider. + """ + existing = CLIENTS.get(provider) + if existing is not None: + return existing + + if provider == "anthropic": + if not settings.LLM.ANTHROPIC_API_KEY: + return None + client: ProviderClient = get_anthropic_client() + elif provider == "openai": + if not settings.LLM.OPENAI_API_KEY: + return None + client = get_openai_client() + elif provider == "gemini": + if not settings.LLM.GEMINI_API_KEY: + return None + client = get_gemini_client() + else: + assert_never(provider) + + CLIENTS[provider] = client + return client def client_for_model_config( @@ -178,7 +204,7 @@ def client_for_model_config( override factories. """ if model_config.api_key is None and model_config.base_url is None: - existing_client = CLIENTS.get(provider) + existing_client = default_client(provider) if existing_client is not None: return existing_client @@ -202,10 +228,16 @@ def backend_for_provider( ) -> ProviderBackend: """Wrap a raw provider SDK client in the matching ProviderBackend adapter.""" if provider == "anthropic": + from .backends.anthropic import AnthropicBackend + return AnthropicBackend(client) if provider == "openai": + from .backends.openai import OpenAIBackend + return OpenAIBackend(client) if provider == "gemini": + from .backends.gemini import GeminiBackend + return GeminiBackend(client) assert_never(provider) @@ -236,6 +268,7 @@ __all__ = [ "CLIENTS", "backend_for_provider", "client_for_model_config", + "default_client", "get_anthropic_client", "get_anthropic_override_client", "get_backend", diff --git a/src/llm/types.py b/src/llm/types.py index 33058e09..b010f6a8 100644 --- a/src/llm/types.py +++ b/src/llm/types.py @@ -12,12 +12,13 @@ from collections.abc import AsyncIterator, Callable from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Generic, Literal, TypeVar -from anthropic import AsyncAnthropic -from google import genai -from openai import AsyncOpenAI from pydantic import BaseModel, Field if TYPE_CHECKING: + from anthropic import AsyncAnthropic + from google import genai + from openai import AsyncOpenAI + from src.llm.capture import CapturedMessage logger = logging.getLogger(__name__) @@ -30,8 +31,13 @@ ReasoningEffortType = ( ) VerbosityType = Literal["low", "medium", "high"] | None -# Raw SDK client union used by the provider-selection layer. -ProviderClient = AsyncAnthropic | AsyncOpenAI | genai.Client +# Raw SDK client union used by the provider-selection layer. The SDK types are +# only imported for type checking; at runtime this stays Any so importing this +# module doesn't load any provider SDK. +if TYPE_CHECKING: + ProviderClient = AsyncAnthropic | AsyncOpenAI | genai.Client +else: + ProviderClient = Any @dataclass diff --git a/tests/llm/test_embedding_client.py b/tests/llm/test_embedding_client.py index 0769c19d..b50a66f0 100644 --- a/tests/llm/test_embedding_client.py +++ b/tests/llm/test_embedding_client.py @@ -70,7 +70,7 @@ async def test_openai_embedding_client_uses_configured_model_and_dimensions( self.base_url: str | None = base_url self.embeddings: FakeOpenAIEmbeddingsAPI = fake_embeddings - monkeypatch.setattr("src.embedding_client.AsyncOpenAI", FakeOpenAIClient) + monkeypatch.setattr("openai.AsyncOpenAI", FakeOpenAIClient) client = _EmbeddingClient( EmbeddingModelConfig( @@ -107,7 +107,7 @@ async def test_openai_embedding_client_rejects_dimension_mismatch( def __init__(self, *, api_key: str | None, base_url: str | None) -> None: self.embeddings: FakeOpenAIEmbeddingsAPI = fake_embeddings - monkeypatch.setattr("src.embedding_client.AsyncOpenAI", FakeOpenAIClient) + monkeypatch.setattr("openai.AsyncOpenAI", FakeOpenAIClient) client = _EmbeddingClient( EmbeddingModelConfig( @@ -156,7 +156,7 @@ async def test_gemini_embedding_client_uses_output_dimensionality( self.http_options: Any = http_options self.aio: Any = SimpleNamespace(models=FakeGeminiModels()) - monkeypatch.setattr("src.embedding_client.genai.Client", FakeGeminiClient) + monkeypatch.setattr("google.genai.Client", FakeGeminiClient) client = _EmbeddingClient( EmbeddingModelConfig( @@ -201,7 +201,7 @@ async def test_gemini_embedding_client_keeps_timeout_without_base_url( self.http_options: Any = http_options self.aio: Any = SimpleNamespace(models=SimpleNamespace()) - monkeypatch.setattr("src.embedding_client.genai.Client", FakeGeminiClient) + monkeypatch.setattr("google.genai.Client", FakeGeminiClient) client = _EmbeddingClient( EmbeddingModelConfig( @@ -238,7 +238,7 @@ def _build_openai_client( self.base_url: str | None = base_url self.embeddings: FakeOpenAIEmbeddingsAPI = fake_embeddings - monkeypatch.setattr("src.embedding_client.AsyncOpenAI", FakeOpenAIClient) + monkeypatch.setattr("openai.AsyncOpenAI", FakeOpenAIClient) client = _EmbeddingClient( EmbeddingModelConfig( @@ -384,7 +384,7 @@ async def test_gemini_simple_batch_embed_respects_configured_max_batch_size( def __init__(self, *, api_key: str | None, http_options: Any) -> None: self.aio: Any = SimpleNamespace(models=FakeGeminiModels()) - monkeypatch.setattr("src.embedding_client.genai.Client", FakeGeminiClient) + monkeypatch.setattr("google.genai.Client", FakeGeminiClient) client = _EmbeddingClient( EmbeddingModelConfig( @@ -432,7 +432,7 @@ async def test_gemini_simple_batch_embed_defaults_to_100_when_unset( def __init__(self, *, api_key: str | None, http_options: Any) -> None: self.aio: Any = SimpleNamespace(models=FakeGeminiModels()) - monkeypatch.setattr("src.embedding_client.genai.Client", FakeGeminiClient) + monkeypatch.setattr("google.genai.Client", FakeGeminiClient) client = _EmbeddingClient( EmbeddingModelConfig( @@ -710,7 +710,7 @@ async def test_simple_batch_embed_respects_token_budget_per_request( def __init__(self, *, api_key: str | None, base_url: str | None) -> None: self.embeddings: FakeOpenAIEmbeddingsAPI = fake_embeddings - monkeypatch.setattr("src.embedding_client.AsyncOpenAI", FakeOpenAIClient) + monkeypatch.setattr("openai.AsyncOpenAI", FakeOpenAIClient) # max_input_tokens=100 per single input; max_tokens_per_request=120 total, # so two ~80-token inputs must end up in *separate* requests. @@ -748,7 +748,7 @@ async def test_simple_batch_embed_rejects_oversized_input( def __init__(self, *, api_key: str | None, base_url: str | None) -> None: self.embeddings: FakeOpenAIEmbeddingsAPI = fake_embeddings - monkeypatch.setattr("src.embedding_client.AsyncOpenAI", FakeOpenAIClient) + monkeypatch.setattr("openai.AsyncOpenAI", FakeOpenAIClient) client = _EmbeddingClient( EmbeddingModelConfig( @@ -778,7 +778,7 @@ def test_prepare_chunks_returns_ordered_chunks( def __init__(self, *, api_key: str | None, base_url: str | None) -> None: self.embeddings: FakeOpenAIEmbeddingsAPI = fake_embeddings - monkeypatch.setattr("src.embedding_client.AsyncOpenAI", FakeOpenAIClient) + monkeypatch.setattr("openai.AsyncOpenAI", FakeOpenAIClient) client = _EmbeddingClient( EmbeddingModelConfig( @@ -844,7 +844,7 @@ async def test_gemini_process_batch_wraps_contents_as_content_part( self.http_options: Any = http_options self.aio: Any = SimpleNamespace(models=FakeGeminiModels()) - monkeypatch.setattr("src.embedding_client.genai.Client", FakeGeminiClient) + monkeypatch.setattr("google.genai.Client", FakeGeminiClient) client = _EmbeddingClient( EmbeddingModelConfig( diff --git a/tests/llm/test_registry.py b/tests/llm/test_registry.py index 49dc72fd..0e4b304f 100644 --- a/tests/llm/test_registry.py +++ b/tests/llm/test_registry.py @@ -50,7 +50,7 @@ def test_get_gemini_client_sets_http_timeout(monkeypatch: pytest.MonkeyPatch) -> """Default Gemini client must carry an HttpOptions timeout, not None.""" monkeypatch.setattr(app_config.settings.LLM, "GEMINI_BASE_URL", None) - with patch("src.llm.registry.genai.Client") as mock_client: + with patch("google.genai.Client") as mock_client: registry_module.get_gemini_client() assert mock_client.call_count == 1 @@ -68,7 +68,7 @@ def test_get_gemini_client_preserves_custom_base_url( app_config.settings.LLM, "GEMINI_BASE_URL", "https://gemini-proxy.example.com" ) - with patch("src.llm.registry.genai.Client") as mock_client: + with patch("google.genai.Client") as mock_client: registry_module.get_gemini_client() http_options = mock_client.call_args.kwargs["http_options"] @@ -80,7 +80,7 @@ def test_get_gemini_client_preserves_custom_base_url( @pytest.mark.usefixtures("fresh_lru_caches") def test_get_gemini_override_client_sets_http_timeout() -> None: """Override Gemini client must also carry a timeout.""" - with patch("src.llm.registry.genai.Client") as mock_client: + with patch("google.genai.Client") as mock_client: registry_module.get_gemini_override_client( "https://gemini-proxy.example.com", "sk-override" ) @@ -94,7 +94,7 @@ def test_get_gemini_override_client_sets_http_timeout() -> None: @pytest.mark.usefixtures("fresh_lru_caches") def test_get_gemini_override_client_handles_missing_base_url() -> None: """Override Gemini client with no base URL still carries a timeout.""" - with patch("src.llm.registry.genai.Client") as mock_client: + with patch("google.genai.Client") as mock_client: registry_module.get_gemini_override_client(None, "sk-override") http_options = mock_client.call_args.kwargs["http_options"] @@ -109,7 +109,7 @@ def test_get_anthropic_client_keeps_600s_timeout( """Anthropic timeout is the established behavior — lock it.""" monkeypatch.setattr(app_config.settings.LLM, "ANTHROPIC_BASE_URL", None) - with patch("src.llm.registry.AsyncAnthropic") as mock_anthropic: + with patch("anthropic.AsyncAnthropic") as mock_anthropic: registry_module.get_anthropic_client() assert mock_anthropic.call_args.kwargs["timeout"] == _ANTHROPIC_TIMEOUT_S @@ -118,7 +118,7 @@ def test_get_anthropic_client_keeps_600s_timeout( @pytest.mark.usefixtures("fresh_lru_caches") def test_get_anthropic_override_client_keeps_600s_timeout() -> None: """Override Anthropic client also keeps the 600s timeout.""" - with patch("src.llm.registry.AsyncAnthropic") as mock_anthropic: + with patch("anthropic.AsyncAnthropic") as mock_anthropic: registry_module.get_anthropic_override_client(None, "sk-override") assert mock_anthropic.call_args.kwargs["timeout"] == _ANTHROPIC_TIMEOUT_S diff --git a/tests/llm/test_telemetry_llm_call.py b/tests/llm/test_telemetry_llm_call.py index ef3cc21c..56a36191 100644 --- a/tests/llm/test_telemetry_llm_call.py +++ b/tests/llm/test_telemetry_llm_call.py @@ -277,7 +277,7 @@ class TestExecutorEndToEnd: @pytest.mark.asyncio async def test_success_path_emits_one_event(self): - from src.llm import executor + from src.llm import executor, registry emitted: list[BaseEvent] = [] result = BackendCompletionResult( @@ -285,7 +285,7 @@ class TestExecutorEndToEnd: ) with ( - patch.object(executor, "CLIENTS", {"anthropic": object()}), + patch.object(registry, "CLIENTS", {"anthropic": object()}), patch.object( executor, "backend_for_provider", @@ -326,7 +326,7 @@ class TestExecutorEndToEnd: 'error' — client disconnects / shutdowns must not pollute error rates.""" import asyncio - from src.llm import executor + from src.llm import executor, registry emitted: list[BaseEvent] = [] @@ -334,7 +334,7 @@ class TestExecutorEndToEnd: raise asyncio.CancelledError() with ( - patch.object(executor, "CLIENTS", {"anthropic": object()}), + patch.object(registry, "CLIENTS", {"anthropic": object()}), patch.object(executor, "backend_for_provider", return_value=object()), patch.object(executor, "execute_completion", new=_cancel), patch( @@ -364,7 +364,7 @@ class TestExecutorEndToEnd: import asyncio from collections.abc import AsyncIterator - from src.llm import executor + from src.llm import executor, registry emitted: list[BaseEvent] = [] @@ -377,7 +377,7 @@ class TestExecutorEndToEnd: return _cancelling_stream() with ( - patch.object(executor, "CLIENTS", {"anthropic": object()}), + patch.object(registry, "CLIENTS", {"anthropic": object()}), patch.object(executor, "backend_for_provider", return_value=object()), patch.object(executor, "execute_stream", new=_setup_stream), patch.object( @@ -418,7 +418,7 @@ class TestExecutorEndToEnd: generator without awaiting `execute_stream`, hiding setup failures from tenacity. """ - from src.llm import executor + from src.llm import executor, registry emitted: list[BaseEvent] = [] @@ -426,7 +426,7 @@ class TestExecutorEndToEnd: raise RuntimeError("rate limited") with ( - patch.object(executor, "CLIENTS", {"anthropic": object()}), + patch.object(registry, "CLIENTS", {"anthropic": object()}), patch.object(executor, "backend_for_provider", return_value=object()), patch.object(executor, "execute_stream", new=_setup_explodes), patch( @@ -455,7 +455,7 @@ class TestExecutorEndToEnd: @pytest.mark.asyncio async def test_error_path_still_emits_via_finally(self): - from src.llm import executor + from src.llm import executor, registry emitted: list[BaseEvent] = [] @@ -463,7 +463,7 @@ class TestExecutorEndToEnd: raise RuntimeError("backend exploded") with ( - patch.object(executor, "CLIENTS", {"anthropic": object()}), + patch.object(registry, "CLIENTS", {"anthropic": object()}), patch.object( executor, "backend_for_provider", @@ -570,7 +570,7 @@ class TestStreamFinalResponseRetryAttempt: async def test_attempt_index_bumps_across_retries(self): from collections.abc import AsyncIterator - from src.llm import executor, tool_loop + from src.llm import executor, registry, tool_loop emitted: list[BaseEvent] = [] @@ -607,7 +607,7 @@ class TestStreamFinalResponseRetryAttempt: ) with ( - patch.object(executor, "CLIENTS", {"anthropic": object()}), + patch.object(registry, "CLIENTS", {"anthropic": object()}), patch.object(executor, "backend_for_provider", return_value=object()), patch.object(executor, "execute_stream", new=_flaky_setup), patch(