diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index 47eaeeca3d6..bf817dd4015 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -27,6 +27,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo convert_to_model_response_object, ) from litellm.llms.bedrock.request_metadata import get_bedrock_request_metadata_fields +from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge.bindings import UNCHANGED, Unchanged from litellm.rust_bridge.configuration import rust_enabled from litellm.rust_bridge.protocols import ( @@ -113,6 +114,8 @@ _CHAT: Final[EndpointDispatch[RustChatCompletions, RustAchatCompletions]] = Endp asynchronous=lambda native: native.achat_completions, enabled=rust_enabled, ) + + def set_rust_chat_completions( *, chat_completions: RustChatCompletions | None | Unchanged = UNCHANGED, diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 3dd5f672ba9..ae6b7d362ca 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -224,24 +224,11 @@ def _reset_rust_flag(): """Keep the global toggle isolated between tests.""" rust_bridge._OCR.sync.reset() rust_bridge._OCR.asynchronous.reset() - rust_bridge._PREFLIGHT.reset() configuration.reset_rust_configuration() rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL - rust_bridge._PREFLIGHT.override( - lambda model, custom_llm_provider, *, context: ( - "unsupported feature" - if any(getattr(context.capabilities, key) for key in ("stream", "has_agentic_hook", "has_custom_client")) - or ( - context.capabilities.request_format == "native" - and not (custom_llm_provider == "azure_ai" and "doc-intelligence" in model) - ) - else None - ) - ) yield rust_bridge._OCR.sync.reset() rust_bridge._OCR.asynchronous.reset() - rust_bridge._PREFLIGHT.reset() configuration.reset_rust_configuration() rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL @@ -378,7 +365,6 @@ def test_explicit_ocr_none_clears_injected_impl(monkeypatch): rust_bridge._OCR.sync.reset() rust_bridge._OCR.asynchronous.reset() - rust_bridge._PREFLIGHT.reset() assert rust_bridge.load_rust_ocr() is None assert rust_bridge.load_rust_aocr() is None diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index 1ec1a8a9621..c1caf99b9a2 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -44,18 +44,10 @@ class _FakeNativeBridge: @pytest.fixture(autouse=True) def reset_responses_websocket(): - responses_websocket.set_rust_responses_websocket(connection=None, decline=None) + responses_websocket.set_rust_responses_websocket(connection=None) configuration.reset_rust_configuration() - responses_websocket.set_rust_responses_websocket( - decline=lambda model, custom_llm_provider, *, context: ( - "unsupported feature" - if any(getattr(context.capabilities, key) for key in ("stream", "has_agentic_hook", "has_custom_client")) - or context.capabilities.request_format == "native" - else None - ) - ) yield - responses_websocket.set_rust_responses_websocket(connection=None, decline=None) + responses_websocket.set_rust_responses_websocket(connection=None) configuration.reset_rust_configuration() @@ -181,11 +173,7 @@ async def test_connection_dispatch_cleans_up_without_reconnecting(native, sessio finally: await python_socket.close() - responses_websocket.set_rust_responses_websocket(connection=Native) - if not native: - responses_websocket.set_rust_responses_websocket( - decline=lambda model, custom_llm_provider, **features: "declined" - ) + responses_websocket.set_rust_responses_websocket(connection=Native if native else None) async def run(): async with responses_websocket.open_connection( @@ -210,7 +198,7 @@ async def test_connection_dispatch_cleans_up_without_reconnecting(native, sessio @pytest.mark.asyncio -async def test_missing_acceptance_export_uses_python_connection_once(): +async def test_missing_acceptance_export_keeps_native_failure_terminal(): from contextlib import asynccontextmanager calls = [] @@ -226,10 +214,10 @@ async def test_missing_acceptance_export_uses_python_connection_once(): configuration.rust(True) responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge) - responses_websocket._PREFLIGHT.override(None) - async with responses_websocket.open_connection( - url="wss://example.test", headers={}, timeout=1, model="model", provider="openai", fallback=python - ) as connection: - assert connection is socket - assert calls == ["python"] - assert socket.closed + with pytest.raises(RuntimeError, match="connection failed"): + async with responses_websocket.open_connection( + url="wss://example.test", headers={}, timeout=1, model="model", provider="openai", fallback=python + ): + pass + assert calls == [] + assert not socket.closed diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index 875b9ea5a8c..d95ba09d08d 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -351,6 +351,7 @@ def test_typed_capability_and_provider_metadata_facts_are_isolated(): assert anthropic_options({"metadata": {"user_id": None}}).has_user_id is False + @pytest.mark.parametrize("provider", ["anthropic", "bedrock", "openai"]) @pytest.mark.parametrize("asynchronous", [False, True]) @pytest.mark.asyncio @@ -375,9 +376,7 @@ async def test_public_completion_discovers_any_provider(provider, asynchronous): @pytest.mark.parametrize("asynchronous", [False, True]) -@pytest.mark.parametrize( - "failure", ["decline", "unavailable", "error", "malformed", "cancelled"] -) +@pytest.mark.parametrize("failure", ["decline", "unavailable", "error", "malformed", "cancelled"]) @pytest.mark.asyncio async def test_public_completion_fallback_contract(monkeypatch, asynchronous, failure): import asyncio diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index 87ed412d379..88fc84a736c 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -15,13 +15,9 @@ rust_bridge = importlib.import_module("litellm.rust_bridge.transcription") @pytest.fixture(autouse=True) def reset_rust_transcription() -> None: - rust_bridge.configure_rust_transcription( - transcription=None, - atranscription=None, - decline=lambda model, custom_llm_provider, *, context: None, - ) + rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) yield - rust_bridge.configure_rust_transcription(transcription=None, atranscription=None, decline=None) + rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) class SyncBridge: