diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 0b0a61192e6..ff5e881761b 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,9 +1,9 @@ { "reportAny": { - "limit": 13429 + "limit": 13426 }, "reportArgumentType": { - "limit": 2198 + "limit": 2192 }, "reportAssignmentType": { "limit": 319 @@ -24,13 +24,13 @@ "limit": 19 }, "reportExplicitAny": { - "limit": 3369 + "limit": 3368 }, "reportFunctionMemberAccess": { "limit": 7 }, "reportGeneralTypeIssues": { - "limit": 101 + "limit": 100 }, "reportIncompatibleMethodOverride": { "limit": 56 @@ -57,7 +57,7 @@ "limit": 5570 }, "reportMissingTypeArgument": { - "limit": 15281 + "limit": 15279 }, "reportMissingTypeStubs": { "limit": 40 @@ -90,7 +90,7 @@ "limit": 8 }, "reportReturnType": { - "limit": 180 + "limit": 178 }, "reportTypedDictNotRequiredAccess": { "limit": 22 @@ -99,31 +99,31 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 44358 + "limit": 44019 }, "reportUnknownLambdaType": { "limit": 109 }, "reportUnknownMemberType": { - "limit": 38271 + "limit": 38266 }, "reportUnknownParameterType": { - "limit": 19584 + "limit": 19583 }, "reportUnknownVariableType": { - "limit": 29814 + "limit": 29810 }, "reportUnnecessaryCast": { "limit": 110 }, "reportUnnecessaryComparison": { - "limit": 687 + "limit": 686 }, "reportUnnecessaryContains": { "limit": 4 }, "reportUnnecessaryIsInstance": { - "limit": 816 + "limit": 815 }, "reportUntypedBaseClass": { "limit": 0 diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index c82be07a5c5..dc3a4d179b7 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -27,6 +27,8 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts +from litellm.rust_bridge.dispatch import anative_first, native_first +from litellm.rust_bridge.runtime import DispatchResult from litellm.types.llms.anthropic import ( ContentBlockDelta, ContentBlockStart, @@ -368,15 +370,7 @@ class AnthropicChatCompletion(BaseLLM): if config is None: raise ValueError(f"Provider config not found for model: {model} and provider: {custom_llm_provider}") - def build_request() -> tuple[dict, dict]: # mutable-ok: rewritten in place downstream - """Translate the request the Python way, returning `(headers, data)`. - - The pair stays mutable because the streaming path rewrites it in - place (`data["stream"] = True`) before sending. - - Shared by the normal path and by the Rust path's fallback, which - builds it only when the Rust call did not serve the request. - """ + def prepare_python() -> tuple[dict[str, str], dict[str, object]]: # mutable-ok: stream mutates data request_data: Final = config.transform_request( model=model, messages=messages, @@ -384,12 +378,29 @@ class AnthropicChatCompletion(BaseLLM): litellm_params=litellm_params, headers=headers, ) - return update_request_with_filtered_beta( + python_headers, data = update_request_with_filtered_beta( headers=headers, request_data=request_data, provider=custom_llm_provider, ) + ## LOGGING + # Reaching here with `serves_via_rust` set means the Rust attempt + # declined at call time, before the provider was called, and already + # logged this request. That is the same attempt continuing. + if not serves_via_rust: + logging_obj.pre_call( + input=messages, + api_key=api_key, + additional_args={ + "complete_input_dict": data, + "api_base": api_base, + "headers": python_headers, + }, + ) + print_verbose(f"_is_function_call: {_is_function_call}") + return python_headers, data + # The Rust core owns the whole call for the subset it accepts, so ask # before transforming: whichever path runs emits pre_call exactly once. # `get_config` merges the class-level defaults (Anthropic's required @@ -406,67 +417,26 @@ class AnthropicChatCompletion(BaseLLM): litellm_params=litellm_params, stream=stream, ) + rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict + "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent + "model": model, + "messages": messages, + **rust_optional_params, + }, + "api_base": api_base, + "headers": headers, + } if serves_via_rust: - rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict - "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent - "model": model, - "messages": messages, - **rust_optional_params, - }, - "api_base": api_base, - "headers": headers, - } logging_obj.pre_call(input=messages, api_key=api_key, additional_args=rust_logging_args) - log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( - logging_obj=logging_obj, - messages=messages, - api_key=api_key, - additional_args=rust_logging_args, - ) - if acompletion is True: + log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( + logging_obj=logging_obj, + messages=messages, + api_key=api_key, + additional_args=rust_logging_args, + ) - async def python_fallback() -> "ModelResponse | CustomStreamWrapper": - # pre_call already fired for this request above. The Rust - # path only declines before the provider is called, so this - # is the same attempt continuing, not a second one. - fallback_headers, fallback_data = build_request() - return await self.acompletion_function( - model=model, - messages=messages, - data=fallback_data, - api_base=api_base, - custom_prompt_dict=custom_prompt_dict, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - api_key=api_key, - provider_config=config, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, - _is_function_call=_is_function_call, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=fallback_headers, - client=client, - json_mode=json_mode, - timeout=timeout, - ) - - return rust_chat_completions_bridge.achat_completions_or_fallback( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, - python_fallback=python_fallback, - ) - rust_response: Final = rust_chat_completions_bridge.chat_completions( + def native_completion() -> DispatchResult[ModelResponse]: + return rust_chat_completions_bridge.chat_completions( model=model, messages=messages, optional_params=rust_optional_params, @@ -477,34 +447,37 @@ class AnthropicChatCompletion(BaseLLM): extra_headers=headers, timeout=timeout, on_response=log_rust_post_call, + eligible=serves_via_rust, ) - if rust_response is not None: - return rust_response - headers, data = build_request() - - ## LOGGING - # Reaching here with `serves_via_rust` set means the Rust attempt - # declined at call time, before the provider was called, and already - # logged this request. That is the same attempt continuing. - if not serves_via_rust: - logging_obj.pre_call( - input=messages, + async def native_acompletion() -> DispatchResult[ModelResponse]: + return await rust_chat_completions_bridge.achat_completions( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, api_key=api_key, - additional_args={ - "complete_input_dict": data, - "api_base": api_base, - "headers": headers, - }, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + eligible=serves_via_rust, ) - print_verbose(f"_is_function_call: {_is_function_call}") - if acompletion is True: + + @anative_first( + native=native_acompletion, + route="chat_completions", + errors=lambda: rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model), + ) + async def execute_async() -> ModelResponse | CustomStreamWrapper: + headers, data = prepare_python() if ( stream is True ): # if function call - fake the streaming (need complete blocks for output parsing in openai format) print_verbose("makes async anthropic streaming POST request") data["stream"] = stream - return self.acompletion_stream_function( + return await self.acompletion_stream_function( model=model, messages=messages, data=data, @@ -526,7 +499,7 @@ class AnthropicChatCompletion(BaseLLM): client=(client if client is not None and isinstance(client, AsyncHTTPHandler) else None), ) else: - return self.acompletion_function( + return await self.acompletion_function( model=model, messages=messages, data=data, @@ -548,7 +521,14 @@ class AnthropicChatCompletion(BaseLLM): json_mode=json_mode, timeout=timeout, ) - else: + + @native_first( + native=native_completion, + route="chat_completions", + errors=lambda: rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model), + ) + def execute_sync() -> ModelResponse | CustomStreamWrapper: + headers, data = prepare_python() ## COMPLETION CALL if ( stream is True @@ -580,13 +560,12 @@ class AnthropicChatCompletion(BaseLLM): ) else: - if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client(params={"timeout": timeout}) - else: - client = client + python_client: Final = ( + client if isinstance(client, HTTPHandler) else _get_httpx_client(params={"timeout": timeout}) + ) try: - response: Final = client.post( + response: Final = python_client.post( api_base, headers=headers, data=json.dumps(data), @@ -607,20 +586,21 @@ class AnthropicChatCompletion(BaseLLM): status_code=status_code, headers=error_headers, ) + return config.transform_response( + model=model, + raw_response=response, + model_response=model_response, + logging_obj=logging_obj, + api_key=api_key, + request_data=data, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + encoding=encoding, + json_mode=json_mode, + ) - return config.transform_response( - model=model, - raw_response=response, - model_response=model_response, - logging_obj=logging_obj, - api_key=api_key, - request_data=data, - messages=messages, - optional_params=optional_params, - litellm_params=litellm_params, - encoding=encoding, - json_mode=json_mode, - ) + return execute_async() if acompletion else execute_sync() def embedding(self): # logic for parsing in - calling - parsing out model embedding calls diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py index b1f8c957ff4..6d6be069f99 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -1,13 +1,23 @@ import base64 -from typing import Final +from typing import Final, NoReturn import httpx from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.rust_bridge import transcription as rust_transcription_bridge +from litellm.rust_bridge.dispatch import PROPAGATE, anative_first, native_first +from litellm.rust_bridge.runtime import DispatchResult, adapt_result from litellm.types.utils import FileTypes, TranscriptionResponse +def _unavailable() -> NoReturn: + raise RuntimeError("Rust audio transcription bridge is unavailable") + + +async def _aunavailable() -> NoReturn: + _unavailable() + + class BedrockAudioTranscriptionRustDispatch: @staticmethod def _audio_payload(audio_file: FileTypes) -> dict[str, object]: @@ -31,6 +41,37 @@ class BedrockAudioTranscriptionRustDispatch: "filename": processed_audio.filename, } + def _attempt_audio_transcriptions( + self, + *, + model: str, + audio_file: FileTypes, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout: float | httpx.Timeout | None, + ) -> DispatchResult[TranscriptionResponse]: + result: Final = rust_transcription_bridge.transcription( + model=model, + audio=self._audio_payload(audio_file), + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout=timeout, + ) + return adapt_result(result, lambda response: TranscriptionResponse(**response)) + + @native_first( + native=_attempt_audio_transcriptions, + route="audio transcription", + errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout: ( + PROPAGATE + ), + ) def audio_transcriptions( self, *, @@ -43,7 +84,21 @@ class BedrockAudioTranscriptionRustDispatch: optional_params: dict[str, object], timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: - rust_response: Final = rust_transcription_bridge.transcription( + _unavailable() + + async def _attempt_async_audio_transcriptions( + self, + *, + model: str, + audio_file: FileTypes, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout: float | httpx.Timeout | None, + ) -> DispatchResult[TranscriptionResponse]: + result: Final = await rust_transcription_bridge.atranscription( model=model, audio=self._audio_payload(audio_file), api_key=api_key, @@ -53,10 +108,15 @@ class BedrockAudioTranscriptionRustDispatch: optional_params=optional_params, timeout=timeout, ) - if rust_response is None: - raise RuntimeError("Rust audio transcription bridge is unavailable") - return TranscriptionResponse(**rust_response) + return adapt_result(result, lambda response: TranscriptionResponse(**response)) + @anative_first( + native=_attempt_async_audio_transcriptions, + route="audio transcription", + errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout: ( + PROPAGATE + ), + ) async def async_audio_transcriptions( self, *, @@ -69,16 +129,4 @@ class BedrockAudioTranscriptionRustDispatch: optional_params: dict[str, object], timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: - rust_response: Final = await rust_transcription_bridge.atranscription( - model=model, - audio=self._audio_payload(audio_file), - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout=timeout, - ) - if rust_response is None: - raise RuntimeError("Rust audio transcription bridge is unavailable") - return TranscriptionResponse(**rust_response) + await _aunavailable() diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index a75124325ae..2918ea96b8c 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -18,6 +18,8 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts +from litellm.rust_bridge.dispatch import anative_first, native_first +from litellm.rust_bridge.runtime import DispatchResult from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper @@ -406,54 +408,25 @@ class BedrockConverseLLM(BaseAWSLLM): litellm_params=litellm_params, stream=stream, ) + rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict + "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent + "messages": messages, + **optional_params, + }, + "api_base": proxy_endpoint_url, + "headers": headers, + } if serves_via_rust: - rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict - "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent - "messages": messages, - **optional_params, - }, - "api_base": proxy_endpoint_url, - "headers": headers, - } logging_obj.pre_call(input=messages, api_key="", additional_args=rust_logging_args) - log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( - logging_obj=logging_obj, - messages=messages, - api_key="", - additional_args=rust_logging_args, - ) - if acompletion: - return rust_chat_completions_bridge.achat_completions_or_fallback( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=proxy_endpoint_url, - custom_llm_provider="bedrock", - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, - python_fallback=lambda: self.async_completion( - model=model, - messages=messages, - api_base=proxy_endpoint_url, - model_response=model_response, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=headers, - timeout=timeout, - client=client, - credentials=credentials, - api_key=api_key, - skip_pre_call_logging=True, - ), - ) - rust_response: Final = rust_chat_completions_bridge.chat_completions( + log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( + logging_obj=logging_obj, + messages=messages, + api_key="", + additional_args=rust_logging_args, + ) + + def native_completion() -> DispatchResult[ModelResponse]: + return rust_chat_completions_bridge.chat_completions( model=model, messages=messages, optional_params=rust_optional_params, @@ -464,16 +437,33 @@ class BedrockConverseLLM(BaseAWSLLM): extra_headers=headers, timeout=timeout, on_response=log_rust_post_call, + eligible=serves_via_rust, ) - if rust_response is not None: - return rust_response - ### ROUTING (ASYNC, STREAMING, SYNC) - if acompletion: - if isinstance(client, HTTPHandler): - client = None + async def native_acompletion() -> DispatchResult[ModelResponse]: + return await rust_chat_completions_bridge.achat_completions( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, + api_key=api_key, + api_base=proxy_endpoint_url, + custom_llm_provider="bedrock", + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + eligible=serves_via_rust, + ) + + @anative_first( + native=native_acompletion, + route="chat_completions", + errors=lambda: rust_chat_completions_bridge.error_handling("bedrock", model), + ) + async def execute_async() -> ModelResponse | CustomStreamWrapper: + python_client: Final = None if isinstance(client, HTTPHandler) else client if stream is True: - return self.async_streaming( + return await self.async_streaming( model=model, messages=messages, api_base=proxy_endpoint_url, @@ -486,7 +476,7 @@ class BedrockConverseLLM(BaseAWSLLM): logger_fn=logger_fn, headers=headers, timeout=timeout, - client=client, + client=python_client, json_mode=json_mode, fake_stream=fake_stream, credentials=credentials, @@ -494,7 +484,7 @@ class BedrockConverseLLM(BaseAWSLLM): stream_chunk_size=stream_chunk_size, ) ### ASYNC COMPLETION - return self.async_completion( + return await self.async_completion( model=model, messages=messages, api_base=proxy_endpoint_url, @@ -507,108 +497,112 @@ class BedrockConverseLLM(BaseAWSLLM): logger_fn=logger_fn, headers=headers, timeout=timeout, - client=client, + client=python_client, credentials=credentials, api_key=api_key, + skip_pre_call_logging=serves_via_rust, + ) + + @native_first( + native=native_completion, + route="chat_completions", + errors=lambda: rust_chat_completions_bridge.error_handling("bedrock", model), + ) + def execute_sync() -> ModelResponse | CustomStreamWrapper: + ## TRANSFORMATION ## + + _data: Final = litellm.AmazonConverseConfig()._transform_request( + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=extra_headers, + ) + data: Final = json.dumps(_data) + + prepped: Final = self.get_request_headers( + credentials=credentials, + aws_region_name=aws_region_name, + extra_headers=extra_headers, + endpoint_url=proxy_endpoint_url, + data=data, + headers=headers, + api_key=api_key, ) - ## TRANSFORMATION ## + ## LOGGING + # Reaching here with `serves_via_rust` set means the synchronous Rust + # attempt declined at call time, before the provider was called, and + # already logged this request. That is the same attempt continuing. + if not serves_via_rust: + logging_obj.pre_call( + input=messages, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": proxy_endpoint_url, + "headers": prepped.headers, + }, + ) + resolved_timeout: Final = httpx.Timeout(timeout) if isinstance(timeout, (float, int)) else timeout + python_client: Final = ( + _get_httpx_client({"timeout": resolved_timeout} if resolved_timeout is not None else None) + if client is None or isinstance(client, AsyncHTTPHandler) + else client + ) - _data: Final = litellm.AmazonConverseConfig()._transform_request( - model=model, - messages=messages, - optional_params=optional_params, - litellm_params=litellm_params, - headers=extra_headers, - ) - data: Final = json.dumps(_data) + if stream is not None and stream is True: + completion_stream, response_headers = make_sync_call( + client=python_client, + api_base=proxy_endpoint_url, + headers=prepped.headers, + data=data, + model=model, + messages=messages, + logging_obj=logging_obj, + json_mode=json_mode, + fake_stream=fake_stream, + stream_chunk_size=stream_chunk_size, + ) + streaming_response: Final = CustomStreamWrapper( + completion_stream=completion_stream, + model=model, + custom_llm_provider="bedrock", + logging_obj=logging_obj, + _response_headers=response_headers, + ) - prepped: Final = self.get_request_headers( - credentials=credentials, - aws_region_name=aws_region_name, - extra_headers=extra_headers, - endpoint_url=proxy_endpoint_url, - data=data, - headers=headers, - api_key=api_key, - ) + return streaming_response - ## LOGGING - # Reaching here with `serves_via_rust` set means the synchronous Rust - # attempt declined at call time, before the provider was called, and - # already logged this request. That is the same attempt continuing. - # The asynchronous branch above returns before this point, and hands - # its own fallback `skip_pre_call_logging=True` for the same reason. - if not serves_via_rust: - logging_obj.pre_call( - input=messages, + ### COMPLETION + + try: + response: Final = python_client.post( + url=proxy_endpoint_url, + headers=prepped.headers, + data=data, + logging_obj=logging_obj, + ) + response.raise_for_status() + except httpx.HTTPStatusError as err: + error_code: Final = err.response.status_code + raise BedrockError(status_code=error_code, message=err.response.text) + except httpx.TimeoutException: + raise BedrockError(status_code=408, message="Timeout error occurred.") + + sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response( + model=model, + response=response, + model_response=model_response, + stream=stream if isinstance(stream, bool) else False, + logging_obj=logging_obj, api_key="", - additional_args={ - "complete_input_dict": data, - "api_base": proxy_endpoint_url, - "headers": prepped.headers, - }, - ) - if client is None or isinstance(client, AsyncHTTPHandler): - _params: Final = {} - if timeout is not None: - if isinstance(timeout, float) or isinstance(timeout, int): - timeout = httpx.Timeout(timeout) - _params["timeout"] = timeout - client = _get_httpx_client(_params) - else: - client = client - - if stream is not None and stream is True: - completion_stream, response_headers = make_sync_call( - client=(client if client is not None and isinstance(client, HTTPHandler) else None), - api_base=proxy_endpoint_url, - headers=prepped.headers, data=data, - model=model, messages=messages, - logging_obj=logging_obj, - json_mode=json_mode, - fake_stream=fake_stream, - stream_chunk_size=stream_chunk_size, - ) - streaming_response: Final = CustomStreamWrapper( - completion_stream=completion_stream, - model=model, - custom_llm_provider="bedrock", - logging_obj=logging_obj, - _response_headers=response_headers, + optional_params=optional_params, + encoding=encoding, ) + sync_transformed_response.set_provider_response_headers(response.headers) + return sync_transformed_response - return streaming_response - - ### COMPLETION - - try: - response: Final = client.post( - url=proxy_endpoint_url, - headers=prepped.headers, - data=data, - logging_obj=logging_obj, - ) - response.raise_for_status() - except httpx.HTTPStatusError as err: - error_code: Final = err.response.status_code - raise BedrockError(status_code=error_code, message=err.response.text) - except httpx.TimeoutException: - raise BedrockError(status_code=408, message="Timeout error occurred.") - - sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response( - model=model, - response=response, - model_response=model_response, - stream=stream if isinstance(stream, bool) else False, - logging_obj=logging_obj, - api_key="", - data=data, - messages=messages, - optional_params=optional_params, - encoding=encoding, - ) - sync_transformed_response.set_provider_response_headers(response.headers) - return sync_transformed_response + return execute_async() if acompletion else execute_sync() diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 552b1549db8..4d6a427e00e 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1,8 +1,8 @@ import asyncio import json import ssl -from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence -from contextlib import asynccontextmanager +from collections.abc import AsyncGenerator, AsyncIterator, Coroutine, Iterator, Mapping, Sequence +from contextlib import AbstractAsyncContextManager, asynccontextmanager from functools import lru_cache from types import MappingProxyType, ModuleType from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, TypeVar, Union, cast, get_type_hints @@ -92,6 +92,8 @@ from litellm.responses.streaming_iterator import ( ResponsesWebSocketStreaming, SyncResponsesAPIStreamingIterator, ) +from litellm.rust_bridge.dispatch import PYTHON_ON_ERROR, anative_context, anative_first +from litellm.rust_bridge.runtime import DispatchResult, NativeSkipped, NativeSkipReason, adapt_result from litellm.types.containers.main import ( ContainerFileListResponse, ContainerListResponse, @@ -2225,116 +2227,111 @@ class BaseLLMHTTPHandler: }, ) - rust_messages_response: Final = await self._maybe_rust_anthropic_messages( - custom_llm_provider=custom_llm_provider, - litellm_params=litellm_params, - has_agentic_hook=self._has_agentic_completion_hook(logging_obj), - model=model, - api_key=api_key, - api_base=api_base, - headers=headers, - request_body=request_body, - timeout=self._resolve_anthropic_messages_timeout( + async def native_messages() -> DispatchResult[AnthropicMessagesResponse | AsyncIterator[object]]: + result: Final = await self._attempt_rust_anthropic_messages( + custom_llm_provider=custom_llm_provider, litellm_params=litellm_params, - stream=stream or False, - custom_llm_provider=custom_llm_provider, - ), - ) - if rust_messages_response is not None: - if stream: - return self._rust_anthropic_messages_fake_stream(rust_messages_response) - return await self._finalize_anthropic_messages_response( - initial_response=rust_messages_response, + has_agentic_hook=self._has_agentic_completion_hook(logging_obj), model=model, - messages=messages, - anthropic_messages_provider_config=anthropic_messages_provider_config, - anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, - logging_obj=logging_obj, - custom_llm_provider=custom_llm_provider, api_key=api_key, - kwargs=kwargs, - ) - - response: Final = await self._async_post_anthropic_messages_with_http_error_retry( - async_httpx_client=async_httpx_client, - request_url=request_url, - headers=headers, - signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json), - request_body=request_body, - stream=stream or False, - logging_obj=logging_obj, - provider_config=anthropic_messages_provider_config, - litellm_params=litellm_params, - api_key=api_key, - model=model, - timeout=self._resolve_anthropic_messages_timeout( - litellm_params=litellm_params, - stream=stream or False, - custom_llm_provider=custom_llm_provider, - ), - ) - - # used for logging + cost tracking - logging_obj.model_call_details["httpx_response"] = response - - initial_response: AsyncIterator | AnthropicMessagesResponse - if stream: - from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( - AnthropicMessagesStreamingResponse, - anthropic_messages_stream_hidden_params, - ) - - completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator( - model=model, - httpx_response=response, + api_base=api_base, + headers=headers, request_body=request_body, - litellm_logging_obj=logging_obj, + timeout=self._resolve_anthropic_messages_timeout( + litellm_params=litellm_params, + stream=stream or False, + custom_llm_provider=custom_llm_provider, + ), ) - stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers) + return adapt_result(result, self._rust_anthropic_messages_fake_stream) if stream else result - if not self._has_agentic_completion_hook(logging_obj): - # No callback overrides async_should_run_agentic_loop, so the - # agentic wrapper's only effect would be buffering every chunk - # and rebuilding the response from SSE at end-of-stream to call - # hooks that all return (False, {}). Stream through directly and - # skip that per-chunk + end-of-stream overhead. - return AnthropicMessagesStreamingResponse( - completion_stream=completion_stream, - hidden_params=stream_hidden_params, + @anative_first(native=native_messages, route="messages", errors=lambda: PYTHON_ON_ERROR) + async def execute_messages() -> AnthropicMessagesResponse | AsyncIterator[object]: + response: Final = await self._async_post_anthropic_messages_with_http_error_retry( + async_httpx_client=async_httpx_client, + request_url=request_url, + headers=headers, + signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json), + request_body=request_body, + stream=stream or False, + logging_obj=logging_obj, + provider_config=anthropic_messages_provider_config, + litellm_params=litellm_params, + api_key=api_key, + model=model, + timeout=self._resolve_anthropic_messages_timeout( + litellm_params=litellm_params, + stream=stream or False, + custom_llm_provider=custom_llm_provider, + ), + ) + + # used for logging + cost tracking + logging_obj.model_call_details["httpx_response"] = response + + initial_response: AsyncIterator | AnthropicMessagesResponse + if stream: + from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( + AnthropicMessagesStreamingResponse, + anthropic_messages_stream_hidden_params, ) - from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import ( - AgenticAnthropicStreamingIterator, - ) + completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator( + model=model, + httpx_response=response, + request_body=request_body, + litellm_logging_obj=logging_obj, + ) + stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers) - held_back_tool_names: Final = self._server_fulfilled_tools_in_request( - logging_obj=logging_obj, - tools=anthropic_messages_optional_request_params.get("tools"), - ) - initial_response = AgenticAnthropicStreamingIterator( - completion_stream=completion_stream, - http_handler=self, - model=model, - messages=messages, - anthropic_messages_provider_config=anthropic_messages_provider_config, - anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, - logging_obj=logging_obj, - custom_llm_provider=custom_llm_provider, - kwargs={**kwargs, "api_key": api_key} if api_key else kwargs, - hold_back=bool(held_back_tool_names), - server_fulfilled_tool_names=held_back_tool_names, - ) - return AnthropicMessagesStreamingResponse( - completion_stream=initial_response, - hidden_params=stream_hidden_params, - ) - else: - initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response( - model=model, - raw_response=response, - logging_obj=logging_obj, - ) + if not self._has_agentic_completion_hook(logging_obj): + # No callback overrides async_should_run_agentic_loop, so the + # agentic wrapper's only effect would be buffering every chunk + # and rebuilding the response from SSE at end-of-stream to call + # hooks that all return (False, {}). Stream through directly and + # skip that per-chunk + end-of-stream overhead. + return AnthropicMessagesStreamingResponse( + completion_stream=completion_stream, + hidden_params=stream_hidden_params, + ) + from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import ( + AgenticAnthropicStreamingIterator, + ) + + held_back_tool_names: Final = self._server_fulfilled_tools_in_request( + logging_obj=logging_obj, + tools=anthropic_messages_optional_request_params.get("tools"), + ) + initial_response = AgenticAnthropicStreamingIterator( + completion_stream=completion_stream, + http_handler=self, + model=model, + messages=messages, + anthropic_messages_provider_config=anthropic_messages_provider_config, + anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, + logging_obj=logging_obj, + custom_llm_provider=custom_llm_provider, + kwargs={**kwargs, "api_key": api_key} if api_key else kwargs, + hold_back=bool(held_back_tool_names), + server_fulfilled_tool_names=held_back_tool_names, + ) + return AnthropicMessagesStreamingResponse( + completion_stream=initial_response, + hidden_params=stream_hidden_params, + ) + else: + initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response( + model=model, + raw_response=response, + logging_obj=logging_obj, + ) + + return initial_response + + initial_response: Final = await execute_messages() + if stream: + return initial_response return await self._finalize_anthropic_messages_response( initial_response=initial_response, model=model, @@ -2384,7 +2381,7 @@ class BaseLLMHTTPHandler: ) @staticmethod - async def _maybe_rust_anthropic_messages( + async def _attempt_rust_anthropic_messages( *, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, @@ -2395,41 +2392,36 @@ class BaseLLMHTTPHandler: headers: dict, request_body: dict, timeout: float | httpx.Timeout | None, - ) -> AnthropicMessagesResponse | None: + ) -> DispatchResult[AnthropicMessagesResponse]: if custom_llm_provider not in ("azure_ai", "anthropic"): - return None + return NativeSkipped(NativeSkipReason.INELIGIBLE) from litellm.rust_bridge.configuration import rust_enabled if not rust_enabled(): - return None + return NativeSkipped(NativeSkipReason.DISABLED) if has_agentic_hook: - return None + return NativeSkipped(NativeSkipReason.INELIGIBLE) from litellm.rust_bridge import messages as rust_messages_bridge upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"} - try: - rust_response: Final = await rust_messages_bridge.amessages( - model=model, - body=upstream_body, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=headers, - timeout=timeout, - ) - except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path - verbose_logger.debug( - "Rust Anthropic messages bridge raised %s; falling back to Python path", - type(rust_error).__name__, - ) - return None - if rust_response is None: - return None + result: Final = await rust_messages_bridge.amessages( + model=model, + body=upstream_body, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=headers, + timeout=timeout, + ) - response_obj: Final = cast(AnthropicMessagesResponse, dict(rust_response)) - response_obj["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}} - return response_obj + def adapt(rust_response: dict[str, object]) -> AnthropicMessagesResponse: + return cast( + AnthropicMessagesResponse, + {**rust_response, "_hidden_params": {"additional_headers": {"x-litellm-rust": "true"}}}, + ) + + return adapt_result(result, adapt) @staticmethod def _rust_anthropic_messages_fake_stream( @@ -6507,20 +6499,22 @@ class BaseLLMHTTPHandler: }, ) + from litellm.rust_bridge import responses_websocket as rust_responses_websocket + + async def attempt_connection() -> DispatchResult[ + AbstractAsyncContextManager[rust_responses_websocket.ConnectionAdapter] + ]: + if not _rust_responses_websocket_enabled(custom_llm_provider): + return NativeSkipped(NativeSkipReason.INELIGIBLE) + return await rust_responses_websocket.managed_connect( + url=ws_url, + headers={str(key): str(value) for key, value in headers.items()}, + timeout=timeout, + ) + + @anative_context(native=attempt_connection, route="responses_websocket", errors=lambda: PYTHON_ON_ERROR) @asynccontextmanager - async def _backend_connection(): - if _rust_responses_websocket_enabled(custom_llm_provider): - from litellm.rust_bridge import responses_websocket as rust_responses_websocket - - rust_backend: Final = await rust_responses_websocket.connect( - url=ws_url, - headers={str(key): str(value) for key, value in headers.items()}, - timeout=timeout, - ) - if rust_backend is not None: - yield rust_backend - return - + async def _backend_connection() -> AsyncGenerator[ClientConnection, None]: async with websockets.connect( ws_url, additional_headers=headers, diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 56c2292d00c..7a3209b9eb4 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -8,7 +8,6 @@ import mimetypes import os import re from collections.abc import Callable, Coroutine, Mapping -from dataclasses import dataclass from io import IOBase from typing import Any, Final, cast @@ -18,18 +17,16 @@ import litellm from litellm._logging import verbose_logger from litellm.constants import request_timeout from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.llms.azure_ai.ocr.common_utils import ( - is_azure_document_intelligence_model, -) +from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model from litellm.llms.base_llm.ocr.transformation import ( OCR_REQUEST_FORMAT_PARAM, - BaseOCRConfig, OCRResponse, parse_ocr_request_format, ) from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge import ocr as rust_ocr_bridge -from litellm.rust_bridge.configuration import rust_enabled +from litellm.rust_bridge.dispatch import PROPAGATE, anative_first, native_first +from litellm.rust_bridge.runtime import DispatchResult from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager, client @@ -38,36 +35,6 @@ base_llm_http_handler = BaseLLMHTTPHandler() ################################################# -@dataclass -class _PreparedOCRRequest: - model: str - document: dict[str, Any] - api_key: str | None - api_base: str | None - custom_llm_provider: str - extra_headers: dict[str, object] | None - provider_config: BaseOCRConfig - optional_params: dict[str, object] - litellm_params: dict[str, object] - effective_timeout: float | httpx.Timeout - litellm_logging_obj: LiteLLMLoggingObj - - -@dataclass -class _PreparedRustOCRCall: - api_key: str | None - api_base: str | None - headers: dict[str, object] - optional_params: dict[str, object] - - -_RUST_OCR_PROVIDERS: Final = { - "mistral", - "azure_ai", - "vertex_ai", -} - - def _prepare_ocr_request( model: str, document: Mapping[str, object], @@ -77,7 +44,7 @@ def _prepare_ocr_request( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, kwargs: dict[str, object], -) -> _PreparedOCRRequest: +) -> rust_ocr_bridge.PreparedOCRRequest: litellm_logging_obj: Final = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj")) litellm_call_id: Final = cast(str | None, kwargs.get("litellm_call_id", None)) @@ -174,7 +141,7 @@ def _prepare_ocr_request( custom_llm_provider=custom_llm_provider, ) - return _PreparedOCRRequest( + return rust_ocr_bridge.PreparedOCRRequest( model=model, document=document, api_key=api_key, @@ -189,146 +156,67 @@ def _prepare_ocr_request( ) -def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool: - if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native": - return False - if not prepared_request.provider_config.supports_rust_bridge(): - return False - return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS - - -def _rust_bridge_optional_params( - prepared_request: _PreparedOCRRequest, - resolve_secret: Callable[[str], str | None], -) -> dict[str, object]: - optional_params: Final = dict(prepared_request.optional_params) - if prepared_request.custom_llm_provider == "vertex_ai": - vertex_project: Final = ( - prepared_request.litellm_params.get("vertex_project") - or prepared_request.litellm_params.get("vertex_ai_project") - or litellm.vertex_project - or resolve_secret("VERTEXAI_PROJECT") - ) - vertex_location: Final = ( - prepared_request.litellm_params.get("vertex_location") - or prepared_request.litellm_params.get("vertex_ai_location") - or litellm.vertex_location - or resolve_secret("VERTEXAI_LOCATION") - or resolve_secret("VERTEX_LOCATION") - ) - if vertex_project is not None: - optional_params["vertex_project"] = vertex_project - if vertex_location is not None: - optional_params["vertex_location"] = vertex_location - return optional_params - - -def _rust_bridge_api_base( - prepared_request: _PreparedOCRRequest, - resolve_secret: Callable[[str], str | None], -) -> str | None: - if prepared_request.api_base is not None: - return prepared_request.api_base - if prepared_request.custom_llm_provider == "azure_ai": - if is_azure_document_intelligence_model(prepared_request.model): - return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT") - return resolve_secret("AZURE_AI_API_BASE") - return None - - -def _prepare_rust_ocr_call( - prepared_request: _PreparedOCRRequest, +@anative_first( + native=rust_ocr_bridge.aattempt_ocr, + route="ocr", + errors=lambda prepared_request, resolve_api_key: PROPAGATE, +) +async def _execute_aocr( + prepared_request: rust_ocr_bridge.PreparedOCRRequest, resolve_api_key: Callable[[str], str | None], -) -> _PreparedRustOCRCall: - provider_config: Final = prepared_request.provider_config - api_key_env_var: Final = provider_config.get_api_key_env_var() - resolved_api_key: Final = prepared_request.api_key or ( - resolve_api_key(api_key_env_var) if api_key_env_var is not None else None - ) - resolved_headers: Final = provider_config.validate_environment( - headers=prepared_request.extra_headers or {}, - model=prepared_request.model, - api_key=resolved_api_key, - api_base=prepared_request.api_base, - litellm_params=prepared_request.litellm_params, - ) - resolved_complete_url: Final = provider_config.get_complete_url( - api_base=prepared_request.api_base, +) -> OCRResponse: + pending: Final = base_llm_http_handler.ocr( model=prepared_request.model, + document=prepared_request.document, optional_params=prepared_request.optional_params, + timeout=prepared_request.effective_timeout, + logging_obj=prepared_request.litellm_logging_obj, + api_key=prepared_request.api_key, + api_base=prepared_request.api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + aocr=True, + headers=prepared_request.extra_headers, + provider_config=prepared_request.provider_config, litellm_params=prepared_request.litellm_params, ) - rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key) - rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key) - prepared_request.litellm_logging_obj.pre_call( - input="OCR document processing", - api_key=resolved_api_key, - additional_args={ - "complete_input_dict": { - "model": prepared_request.model, - "document": prepared_request.document, - **rust_optional_params, - }, - "api_base": resolved_complete_url, - "headers": resolved_headers, - }, - ) - return _PreparedRustOCRCall( - api_key=resolved_api_key, - api_base=rust_api_base, - headers=cast(dict[str, object], resolved_headers), - optional_params=rust_optional_params, - ) + response: Final = await pending if asyncio.iscoroutine(pending) else pending + if response is None: + raise ValueError(f"Got an unexpected None response from the OCR API: {response}") + return response -def _run_rust_ocr( - prepared_request: _PreparedOCRRequest, +def _attempt_ocr( + prepared_request: rust_ocr_bridge.PreparedOCRRequest, resolve_api_key: Callable[[str], str | None], -) -> OCRResponse | None: - if rust_ocr_bridge.load_rust_ocr() is None: - return None - prepared: Final = _prepare_rust_ocr_call( - prepared_request=prepared_request, - resolve_api_key=resolve_api_key, - ) - rust_response: Final = rust_ocr_bridge.ocr( + is_async: bool, +) -> DispatchResult[OCRResponse]: + return rust_ocr_bridge.attempt_ocr(prepared_request=prepared_request, resolve_api_key=resolve_api_key) + + +@native_first( + native=_attempt_ocr, + route="ocr", + errors=lambda prepared_request, resolve_api_key, is_async: PROPAGATE, +) +def _execute_ocr( + prepared_request: rust_ocr_bridge.PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], + is_async: bool, +) -> OCRResponse | Coroutine[object, object, OCRResponse]: + return base_llm_http_handler.ocr( model=prepared_request.model, document=prepared_request.document, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared_request.custom_llm_provider, - extra_headers=prepared.headers, - optional_params=prepared.optional_params, + optional_params=prepared_request.optional_params, timeout=prepared_request.effective_timeout, - ) - if rust_response is None: - return None - return OCRResponse.model_validate(rust_response) - - -async def _run_rust_aocr( - prepared_request: _PreparedOCRRequest, - resolve_api_key: Callable[[str], str | None], -) -> OCRResponse | None: - if rust_ocr_bridge.load_rust_aocr() is None: - return None - prepared: Final = _prepare_rust_ocr_call( - prepared_request=prepared_request, - resolve_api_key=resolve_api_key, - ) - rust_response: Final = await rust_ocr_bridge.aocr( - model=prepared_request.model, - document=prepared_request.document, - api_key=prepared.api_key, - api_base=prepared.api_base, + logging_obj=prepared_request.litellm_logging_obj, + api_key=prepared_request.api_key, + api_base=prepared_request.api_base, custom_llm_provider=prepared_request.custom_llm_provider, - extra_headers=prepared.headers, - optional_params=prepared.optional_params, - timeout=prepared_request.effective_timeout, + aocr=is_async, + headers=prepared_request.extra_headers, + provider_config=prepared_request.provider_config, + litellm_params=prepared_request.litellm_params, ) - if rust_response is None: - return None - return OCRResponse.model_validate(rust_response) @client @@ -425,40 +313,9 @@ async def aocr( custom_llm_provider = prepared.custom_llm_provider completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider}) - if _rust_ocr_supported(prepared) and rust_enabled(): - from litellm.secret_managers.main import get_secret_str + from litellm.secret_managers.main import get_secret_str - rust_response: Final = await _run_rust_aocr( - prepared_request=prepared, - resolve_api_key=get_secret_str, - ) - if rust_response is None: - verbose_logger.debug("Async Rust OCR bridge unavailable; falling back to Python path") - else: - return rust_response - - response = base_llm_http_handler.ocr( - model=prepared.model, - document=prepared.document, - optional_params=prepared.optional_params, - timeout=prepared.effective_timeout, - logging_obj=prepared.litellm_logging_obj, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared.custom_llm_provider, - aocr=True, - headers=prepared.extra_headers, - provider_config=prepared.provider_config, - litellm_params=prepared.litellm_params, - ) - - if asyncio.iscoroutine(response): - response = await response - - if response is None: - raise ValueError(f"Got an unexpected None response from the OCR API: {response}") - - return response + return await _execute_aocr(prepared_request=prepared, resolve_api_key=get_secret_str) except Exception as e: raise litellm.exception_type( model=model, @@ -697,34 +554,9 @@ def ocr( custom_llm_provider = prepared.custom_llm_provider completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider}) - if _rust_ocr_supported(prepared) and rust_enabled(): - from litellm.secret_managers.main import get_secret_str + from litellm.secret_managers.main import get_secret_str - rust_response: Final = _run_rust_ocr( - prepared_request=prepared, - resolve_api_key=get_secret_str, - ) - if rust_response is None: - verbose_logger.debug("Rust OCR bridge unavailable; falling back to Python path") - else: - return rust_response - - response: Final = base_llm_http_handler.ocr( - model=prepared.model, - document=prepared.document, - optional_params=prepared.optional_params, - timeout=prepared.effective_timeout, - logging_obj=prepared.litellm_logging_obj, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared.custom_llm_provider, - aocr=_is_async, - headers=prepared.extra_headers, - provider_config=prepared.provider_config, - litellm_params=prepared.litellm_params, - ) - - return response + return _execute_ocr(prepared_request=prepared, resolve_api_key=get_secret_str, is_async=_is_async) except Exception as e: raise litellm.exception_type( model=model, diff --git a/litellm/rust_bridge/bindings.py b/litellm/rust_bridge/bindings.py index d16f150a2aa..ab7b7c296b9 100644 --- a/litellm/rust_bridge/bindings.py +++ b/litellm/rust_bridge/bindings.py @@ -1,9 +1,11 @@ from __future__ import annotations from collections.abc import Callable -from typing import Final, Generic, TypeVar +from types import ModuleType +from typing import Final, Generic, TypeVar, cast # noqa: TID251 # PyO3 module boundary from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.protocols import NativeModule BindingT = TypeVar("BindingT") @@ -15,21 +17,38 @@ class _Unset: _UNSET: Final = _Unset() +class Unchanged: + pass + + +UNCHANGED: Final = Unchanged() + + class NativeBinding(Generic[BindingT]): """Resolve one native attribute with an explicit, resettable test override.""" - def __init__(self, attribute: str, *, validate: Callable[[object], BindingT | None]) -> None: - self._attribute: Final = attribute - self._validate: Final = validate + def __init__( + self, + select: Callable[[NativeModule], BindingT], + *, + module_loader: Callable[[], ModuleType | None] | None = None, + ) -> None: + self._select: Final = select + self._module_loader: Final = module_loader self._override: BindingT | None | _Unset = _UNSET def load(self) -> BindingT | None: if not isinstance(self._override, _Unset): return self._override - native: Final = get_native_bridge() + native: Final = self._module_loader() if self._module_loader is not None else get_native_bridge() if native is None: return None - return self._validate(getattr(native, self._attribute, None)) + module: Final = cast(NativeModule, native) # cast-ok: PyO3 exports are validated individually below + try: + value: Final = self._select(module) + except AttributeError: + return None + return value if callable(value) else None def override(self, value: BindingT | None) -> None: self._override = value @@ -38,12 +57,21 @@ class NativeBinding(Generic[BindingT]): self._override = _UNSET -def native_exception_types() -> tuple[type[BaseException], type[BaseException]] | None: - native: Final = get_native_bridge() - if native is None: - return None - declined: Final = getattr(native, "RustBridgeDeclined", None) - upstream: Final = getattr(native, "RustUpstreamError", None) - if not isinstance(declined, type) or not isinstance(upstream, type): - return None - return declined, upstream +_DECLINED: Final = NativeBinding(lambda native: native.RustBridgeDeclined) +_UPSTREAM: Final = NativeBinding(lambda native: native.RustUpstreamError) + + +def _exception_class(value: object) -> type[BaseException] | None: + if isinstance(value, type) and issubclass(value, BaseException): + return value + return None + + +def native_upstream_types() -> tuple[type[BaseException], ...]: + upstream: Final = _exception_class(_UPSTREAM.load()) + return () if upstream is None else (upstream,) + + +def native_declined_types() -> tuple[type[BaseException], ...]: + declined: Final = _exception_class(_DECLINED.load()) + return () if declined is None else (declined,) diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index 674bd8847f7..d9f06098903 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -4,30 +4,31 @@ The Rust core owns the conversation translation, the provider call, and the response normalization for the subset of `/chat/completions` requests it accepts. This module only marshals inputs and hands the normalized result to LiteLLM's existing `ModelResponse` builder. - -``None`` means the provider was never called, so the caller is free to serve the -request on the Python path. A failure after the call was issued raises instead: -retrying it there would bill the customer for the same work twice. """ from __future__ import annotations import json -from collections.abc import Awaitable, Callable, Mapping, Sequence -from dataclasses import dataclass +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Final, Protocol import httpx from pydantic import TypeAdapter, ValidationError from litellm._logging import verbose_logger -from litellm.exceptions import APIError from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( convert_to_model_response_object, ) from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged from litellm.rust_bridge.configuration import rust_enabled -from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.dispatch import APIErrorMapping, ErrorAction, ErrorHandling +from litellm.rust_bridge.protocols import ( + RustAchatCompletions, + RustChatCompletions, + RustChatCompletionsDecline, +) +from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.types.utils import ModelResponse @@ -45,47 +46,6 @@ _LITELLM_METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object]) RUST_RESPONSE_HEADER: Final = "x-litellm-rust" -class RustChatCompletions(Protocol): - def __call__( - self, - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object] | None, - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout_seconds: float | None, - ) -> Mapping[str, object]: - raise NotImplementedError - - -class RustAchatCompletions(Protocol): - def __call__( - self, - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object] | None, - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout_seconds: float | None, - ) -> Awaitable[Mapping[str, object]]: - raise NotImplementedError - - -class RustChatCompletionsDecline(Protocol): - def __call__( - self, - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object] | None, - custom_llm_provider: str | None, - ) -> str | None: - raise NotImplementedError - - class ResponseObserver(Protocol): """Invoked with the payload the core returned, on success only. @@ -126,67 +86,36 @@ def response_logger( return log -class _Unset: - pass - - -_UNSET: Final[_Unset] = _Unset() - - -@dataclass(slots=True) -class _RustChatCompletionsState: - chat_completions: RustChatCompletions | None = None - achat_completions: RustAchatCompletions | None = None - decline: RustChatCompletionsDecline | None = None - - -_STATE: Final[_RustChatCompletionsState] = _RustChatCompletionsState() +_CHAT: Final[NativeBinding[RustChatCompletions]] = NativeBinding(lambda native: native.chat_completions) +_ACHAT: Final[NativeBinding[RustAchatCompletions]] = NativeBinding(lambda native: native.achat_completions) +_CHAT_PREFLIGHT: Final[NativeBinding[RustChatCompletionsDecline]] = NativeBinding( + lambda native: native.chat_completions_decline +) def set_rust_chat_completions( *, - chat_completions: RustChatCompletions | None | _Unset = _UNSET, - achat_completions: RustAchatCompletions | None | _Unset = _UNSET, - decline: RustChatCompletionsDecline | None | _Unset = _UNSET, + chat_completions: RustChatCompletions | None | Unchanged = UNCHANGED, + achat_completions: RustAchatCompletions | None | Unchanged = UNCHANGED, + decline: RustChatCompletionsDecline | None | Unchanged = UNCHANGED, ) -> None: """Inject the native callables, so tests can supply a double instead of patching module attributes.""" - if not isinstance(chat_completions, _Unset): - _STATE.chat_completions = chat_completions - if not isinstance(achat_completions, _Unset): - _STATE.achat_completions = achat_completions - if not isinstance(decline, _Unset): - _STATE.decline = decline - - -def load_rust_chat_completions() -> RustChatCompletions | None: - if _STATE.chat_completions is not None: - return _STATE.chat_completions - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - loaded: RustChatCompletions | None = getattr(native_bridge, "chat_completions", None) - return loaded - - -def load_rust_achat_completions() -> RustAchatCompletions | None: - if _STATE.achat_completions is not None: - return _STATE.achat_completions - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - loaded: RustAchatCompletions | None = getattr(native_bridge, "achat_completions", None) - return loaded - - -def _load_rust_decline() -> RustChatCompletionsDecline | None: - if _STATE.decline is not None: - return _STATE.decline - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - loaded: RustChatCompletionsDecline | None = getattr(native_bridge, "chat_completions_decline", None) - return loaded + if not isinstance(chat_completions, Unchanged): + if chat_completions is None: + _CHAT.reset() + else: + _CHAT.override(chat_completions) + if not isinstance(achat_completions, Unchanged): + if achat_completions is None: + _ACHAT.reset() + else: + _ACHAT.override(achat_completions) + if not isinstance(decline, Unchanged): + if decline is None: + _CHAT_PREFLIGHT.reset() + else: + _CHAT_PREFLIGHT.override(decline) def _anthropic_user_id_reaches_the_body(litellm_params: Mapping[str, object] | None) -> bool: @@ -247,12 +176,12 @@ def rust_chat_completions_accepts( return False if stream: return False - if not rust_enabled(): - return False if _litellm_metadata_reaches_the_provider(custom_llm_provider, litellm_params): verbose_logger.debug("Rust chat completions declined (litellm metadata user_id); using the Python path") return False - decline: Final = _load_rust_decline() + if not rust_enabled(): + return False + decline: Final = _CHAT_PREFLIGHT.load() if decline is None: return False try: @@ -262,67 +191,12 @@ def rust_chat_completions_accepts( optional_params=optional_params, custom_llm_provider=custom_llm_provider, ) - except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path - verbose_logger.debug( - "Rust chat completions gate raised %s; staying on the Python path", - type(rust_error).__name__, - ) + except Exception as error: # noqa: BLE001 # capability checks perform no provider I/O + verbose_logger.debug("Native chat acceptance check failed: %s", error) return False if reason is not None: - verbose_logger.debug("Rust chat completions declined (%s); using the Python path", reason) - return False - return True - - -def _rust_bridge_exceptions() -> tuple[type[BaseException], type[BaseException]] | None: - """`(declined, upstream_failed)` from the native module, or None when absent.""" - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - declined: Final = getattr(native_bridge, "RustBridgeDeclined", None) - upstream: Final = getattr(native_bridge, "RustUpstreamError", None) - if declined is None or upstream is None: - return None - return declined, upstream - - -def _reraise_or_decline( - rust_error: BaseException, - *, - model: str, - custom_llm_provider: str | None, -) -> None: - """Re-raise a failure the provider already saw, or return so the caller declines. - - A request that never reached the provider is safe to serve on the Python - path. One that did is not: the provider has already done the work, so a - second attempt bills for it twice. Those surface as an `APIError` carrying - the upstream status, which LiteLLM's exception mapping already understands. - """ - exceptions: Final = _rust_bridge_exceptions() - if exceptions is None: - verbose_logger.debug( - "Rust chat completions bridge raised %s; falling back to Python path", - type(rust_error).__name__, - ) - return - declined, upstream_failed = exceptions - if isinstance(rust_error, upstream_failed): - args: Final = rust_error.args - status: Final = args[0] if args else 0 - message: Final = args[1] if len(args) > 1 else "" - raise APIError( - status_code=int(status) or 500, - message=f"litellm rust chat completions: {message}", - llm_provider=custom_llm_provider or "", - model=model, - ) - if not isinstance(rust_error, declined): - raise rust_error - verbose_logger.debug( - "Rust chat completions declined before calling the provider (%s); using the Python path", - rust_error, - ) + verbose_logger.debug("Native chat request is ineligible: %s", reason) + return reason is None def _build_model_response( @@ -351,12 +225,14 @@ def chat_completions( extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, on_response: ResponseObserver, -) -> ModelResponse | None: - rust_chat_completions: Final = load_rust_chat_completions() - if rust_chat_completions is None: - return None - try: - rust_response: Final = rust_chat_completions( + eligible: bool = True, +) -> DispatchResult[ModelResponse]: + def adapt(rust_response: Mapping[str, object]) -> ModelResponse: + on_response(rust_response) + return _build_model_response(rust_response, model_response) + + def call(native: RustChatCompletions, timeout_seconds: float | None) -> Mapping[str, object]: + return native( model=model, messages=messages, optional_params=optional_params, @@ -364,13 +240,17 @@ def chat_completions( api_base=api_base, custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), + timeout_seconds=timeout_seconds, ) - except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw - _reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider) - return None - on_response(rust_response) - return _build_model_response(rust_response, model_response) + + return attempt( + load=_CHAT.load, + enabled=rust_enabled(), + eligible=eligible, + prepare=lambda: timeout_to_seconds(timeout), + call=call, + adapt=adapt, + ) async def achat_completions( @@ -385,12 +265,14 @@ async def achat_completions( extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, on_response: ResponseObserver, -) -> ModelResponse | None: - rust_achat_completions: Final = load_rust_achat_completions() - if rust_achat_completions is None: - return None - try: - rust_response: Final = await rust_achat_completions( + eligible: bool = True, +) -> DispatchResult[ModelResponse]: + def adapt(rust_response: Mapping[str, object]) -> ModelResponse: + on_response(rust_response) + return _build_model_response(rust_response, model_response) + + async def call(native: RustAchatCompletions, timeout_seconds: float | None) -> Mapping[str, object]: + return await native( model=model, messages=messages, optional_params=optional_params, @@ -398,49 +280,22 @@ async def achat_completions( api_base=api_base, custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), + timeout_seconds=timeout_seconds, ) - except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw - _reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider) - return None - on_response(rust_response) - return _build_model_response(rust_response, model_response) - -async def achat_completions_or_fallback( - *, - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object], - model_response: ModelResponse, - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout: float | httpx.Timeout | None, - on_response: ResponseObserver, - python_fallback: Callable[[], Awaitable[object]], -) -> object: - """Await the Rust path, falling back to the caller's own Python path when - the bridge is unavailable or the call fails. - - The caller supplies the fallback, so the bridge stays free of provider - dispatch. This exists because a caller that dispatches asynchronously has - already returned a coroutine by the time a Rust failure surfaces, and so - cannot fall back on its own. - """ - response: Final = await achat_completions( - model=model, - messages=messages, - optional_params=optional_params, - model_response=model_response, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout=timeout, - on_response=on_response, + return await aattempt( + load=_ACHAT.load, + enabled=rust_enabled(), + eligible=eligible, + prepare=lambda: timeout_to_seconds(timeout), + call=call, + adapt=adapt, + ) + + +def error_handling(provider: str, model: str) -> ErrorHandling: + return ErrorHandling( + declined=ErrorAction.SKIP, + upstream=APIErrorMapping(provider=provider, model=model), + missing_metadata=ErrorAction.SKIP, ) - if response is not None: - return response - return await python_fallback() diff --git a/litellm/rust_bridge/dispatch.py b/litellm/rust_bridge/dispatch.py new file mode 100644 index 00000000000..8c275db3533 --- /dev/null +++ b/litellm/rust_bridge/dispatch.py @@ -0,0 +1,186 @@ +from __future__ import annotations + +from collections.abc import AsyncGenerator, Awaitable, Callable +from contextlib import AbstractAsyncContextManager, asynccontextmanager +from dataclasses import dataclass +from enum import Enum +from functools import wraps +from typing import Final, ParamSpec, TypeAlias, TypeVar + +from litellm._logging import verbose_logger +from litellm.exceptions import APIError +from litellm.rust_bridge.bindings import native_declined_types, native_upstream_types +from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, NativeSkipped, NativeSkipReason + +NativeT = TypeVar("NativeT") +PythonT = TypeVar("PythonT") +P = ParamSpec("P") + + +class ErrorAction(Enum): + RAISE = "raise" + SKIP = "skip" + + +@dataclass(frozen=True, slots=True) +class APIErrorMapping: + provider: str + model: str + + +FailureAction: TypeAlias = ErrorAction | APIErrorMapping + + +@dataclass(frozen=True, slots=True) +class ErrorHandling: + declined: FailureAction = ErrorAction.RAISE + upstream: FailureAction = ErrorAction.RAISE + unknown: FailureAction = ErrorAction.RAISE + missing_metadata: FailureAction = ErrorAction.RAISE + unexpected: FailureAction = ErrorAction.RAISE + + +PROPAGATE: Final = ErrorHandling() +PYTHON_ON_ERROR: Final = ErrorHandling( + declined=ErrorAction.SKIP, + upstream=ErrorAction.SKIP, + unknown=ErrorAction.SKIP, + missing_metadata=ErrorAction.SKIP, + unexpected=ErrorAction.SKIP, +) + + +def _handle_error(error: Exception, action: FailureAction, route: str, reason: NativeSkipReason) -> NativeSkipped: + match action: + case ErrorAction.SKIP: + return NativeSkipped(reason, str(error)) + case ErrorAction.RAISE: + raise error + case APIErrorMapping(provider, model): + args: Final[tuple[object, ...]] = error.args + attribute_status: Final = getattr(error, "status_code", None) + attribute_message: Final = getattr(error, "message", None) + status_value: Final = attribute_status if isinstance(attribute_status, int) else (args[0] if args else 0) + message_value: Final = ( + attribute_message if isinstance(attribute_message, str) else (args[1] if len(args) > 1 else str(error)) + ) + status: Final = status_value if isinstance(status_value, int) else 0 + message: Final = message_value if isinstance(message_value, str) else str(message_value) + raise APIError( + status_code=status or 500, + message=f"litellm rust {route}: {message}", + llm_provider=provider, + model=model, + ) from error + + +def _resolve(result: DispatchResult[NativeT], errors: ErrorHandling, route: str) -> Handled[NativeT] | NativeSkipped: + if not isinstance(result, NativeFailed): + return result + declined: Final = native_declined_types() + upstream: Final = native_upstream_types() + if not declined or not upstream: + return _handle_error(result.error, errors.missing_metadata, route, NativeSkipReason.FAILED) + if isinstance(result.error, declined): + return _handle_error(result.error, errors.declined, route, NativeSkipReason.DECLINED) + if isinstance(result.error, upstream): + return _handle_error(result.error, errors.upstream, route, NativeSkipReason.FAILED) + return _handle_error(result.error, errors.unknown, route, NativeSkipReason.FAILED) + + +def _log_skip(route: str, skipped: NativeSkipped) -> None: + verbose_logger.debug("Native %s skipped (%s): %s", route, skipped.reason.value, skipped.detail or "") + + +def native_first( + *, + native: Callable[P, DispatchResult[NativeT]], + route: str, + errors: Callable[P, ErrorHandling], +) -> Callable[[Callable[P, PythonT]], Callable[P, NativeT | PythonT]]: + def wrap(implementation: Callable[P, PythonT]) -> Callable[P, NativeT | PythonT]: + @wraps(implementation) + def run( + *args: P.args, + **kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature + ) -> NativeT | PythonT: + rules: Final = errors(*args, **kwargs) + try: + attempted: Final = native(*args, **kwargs) + except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures + skipped: Final = _handle_error(error, rules.unexpected, route, NativeSkipReason.FAILED) + _log_skip(route, skipped) + else: + result: Final = _resolve(attempted, rules, route) + if isinstance(result, Handled): + return result.value + _log_skip(route, result) + return implementation(*args, **kwargs) + + return run + + return wrap + + +def anative_first( + *, + native: Callable[P, Awaitable[DispatchResult[NativeT]]], + route: str, + errors: Callable[P, ErrorHandling], +) -> Callable[[Callable[P, Awaitable[PythonT]]], Callable[P, Awaitable[NativeT | PythonT]]]: + def wrap(implementation: Callable[P, Awaitable[PythonT]]) -> Callable[P, Awaitable[NativeT | PythonT]]: + @wraps(implementation) + async def run( + *args: P.args, + **kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature + ) -> NativeT | PythonT: + rules: Final = errors(*args, **kwargs) + try: + attempted: Final = await native(*args, **kwargs) + except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures + skipped: Final = _handle_error(error, rules.unexpected, route, NativeSkipReason.FAILED) + _log_skip(route, skipped) + else: + result: Final = _resolve(attempted, rules, route) + if isinstance(result, Handled): + return result.value + _log_skip(route, result) + return await implementation(*args, **kwargs) + + return run + + return wrap + + +def anative_context( + *, + native: Callable[P, Awaitable[DispatchResult[AbstractAsyncContextManager[NativeT]]]], + route: str, + errors: Callable[P, ErrorHandling], +) -> Callable[ + [Callable[P, AbstractAsyncContextManager[PythonT]]], + Callable[P, AbstractAsyncContextManager[NativeT | PythonT]], +]: + def wrap( + implementation: Callable[P, AbstractAsyncContextManager[PythonT]], + ) -> Callable[P, AbstractAsyncContextManager[NativeT | PythonT]]: + @anative_first(native=native, route=route, errors=errors) + async def acquire( + *args: P.args, + **kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature + ) -> AbstractAsyncContextManager[PythonT]: + return implementation(*args, **kwargs) + + @wraps(implementation) + @asynccontextmanager + async def run( + *args: P.args, + **kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature + ) -> AsyncGenerator[NativeT | PythonT, None]: + manager: Final = await acquire(*args, **kwargs) + async with manager as connection: + yield connection + + return run + + return wrap diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index 40d0ddf622b..160e6f0f743 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -2,90 +2,42 @@ from __future__ import annotations -from collections.abc import Awaitable -from dataclasses import dataclass -from typing import Final, Protocol, cast +from typing import Final import httpx +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged +from litellm.rust_bridge.protocols import RustAmessages, RustMessages +from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity from litellm.rust_bridge.timeouts import timeout_to_seconds - -class RustMessages(Protocol): - def __call__( - self, - model: str, - body: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - timeout_seconds: float | None, - ) -> dict[str, object]: - raise NotImplementedError - - -class RustAmessages(Protocol): - def __call__( - self, - model: str, - body: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - timeout_seconds: float | None, - ) -> Awaitable[dict[str, object]]: - raise NotImplementedError - - -class _Unset: - pass - - -_UNSET: Final[_Unset] = _Unset() - - -@dataclass(slots=True) -class _RustMessagesState: - messages: RustMessages | None = None - amessages: RustAmessages | None = None - - -_STATE: Final[_RustMessagesState] = _RustMessagesState() +_MESSAGES: Final[NativeBinding[RustMessages]] = NativeBinding(lambda native: native.messages) +_AMESSAGES: Final[NativeBinding[RustAmessages]] = NativeBinding(lambda native: native.amessages) def set_rust_messages( *, - messages: RustMessages | None | _Unset = _UNSET, - amessages: RustAmessages | None | _Unset = _UNSET, + messages: RustMessages | None | Unchanged = UNCHANGED, + amessages: RustAmessages | None | Unchanged = UNCHANGED, ) -> None: - if not isinstance(messages, _Unset): - _STATE.messages = messages - if not isinstance(amessages, _Unset): - _STATE.amessages = amessages + if not isinstance(messages, Unchanged): + if messages is None: + _MESSAGES.reset() + else: + _MESSAGES.override(messages) + if not isinstance(amessages, Unchanged): + if amessages is None: + _AMESSAGES.reset() + else: + _AMESSAGES.override(amessages) def load_rust_messages() -> RustMessages | None: - if _STATE.messages is not None: - return _STATE.messages - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - return cast(RustMessages, getattr(native_bridge, "messages", None)) + return _MESSAGES.load() def load_rust_amessages() -> RustAmessages | None: - if _STATE.amessages is not None: - return _STATE.amessages - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - return cast(RustAmessages, getattr(native_bridge, "amessages", None)) + return _AMESSAGES.load() def messages( @@ -97,18 +49,22 @@ def messages( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - rust_messages: Final = load_rust_messages() - if rust_messages is None: - return None - return rust_messages( - model=model, - body=body, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), +) -> DispatchResult[dict[str, object]]: + return attempt( + load=_MESSAGES.load, + enabled=True, + eligible=True, + prepare=lambda: timeout_to_seconds(timeout), + call=lambda rust_messages, timeout_seconds: rust_messages( + model=model, + body=body, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_seconds, + ), + adapt=identity, ) @@ -121,16 +77,20 @@ async def amessages( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - rust_amessages: Final = load_rust_amessages() - if rust_amessages is None: - return None - return await rust_amessages( - model=model, - body=body, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), +) -> DispatchResult[dict[str, object]]: + return await aattempt( + load=_AMESSAGES.load, + enabled=True, + eligible=True, + prepare=lambda: timeout_to_seconds(timeout), + call=lambda rust_amessages, timeout_seconds: rust_amessages( + model=model, + body=body, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_seconds, + ), + adapt=identity, ) diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index b7fdb5a98ef..c04645cb29d 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -1,56 +1,57 @@ -"""Thin Python wrapper for the native Rust OCR bridge.""" - from __future__ import annotations -from collections.abc import Awaitable -from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables +from collections.abc import Callable +from dataclasses import dataclass +from typing import Final import httpx +from pydantic import TypeAdapter +import litellm +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model +from litellm.llms.base_llm.ocr.transformation import OCR_REQUEST_FORMAT_PARAM, BaseOCRConfig, OCRResponse +from litellm.rust_bridge import configuration as _configuration from litellm.rust_bridge.bindings import NativeBinding -from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds +from litellm.rust_bridge.protocols import RustAocr, RustOcr +from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt +from litellm.rust_bridge.timeouts import timeout_to_seconds + +_OCR: Final[NativeBinding[RustOcr]] = NativeBinding(lambda native: native.ocr) +_AOCR: Final[NativeBinding[RustAocr]] = NativeBinding(lambda native: native.aocr) +_HEADERS: Final = TypeAdapter(dict[str, object]) -class RustOcr(Protocol): - def __call__( - self, - model: str, - document: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, - ) -> dict[str, object]: - raise NotImplementedError +@dataclass(frozen=True, slots=True) +class PreparedOCRRequest: + model: str + document: dict[str, object] + api_key: str | None + api_base: str | None + custom_llm_provider: str + extra_headers: dict[str, object] | None + provider_config: BaseOCRConfig + optional_params: dict[str, object] + litellm_params: dict[str, object] + effective_timeout: float | httpx.Timeout + litellm_logging_obj: LiteLLMLoggingObj -class RustAocr(Protocol): - def __call__( - self, - model: str, - document: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, - ) -> Awaitable[dict[str, object]]: - raise NotImplementedError +@dataclass(frozen=True, slots=True) +class _PreparedRustOCRCall: + api_key: str | None + api_base: str | None + headers: dict[str, object] + optional_params: dict[str, object] -def _as_ocr(value: object) -> RustOcr | None: - return cast(RustOcr, value) if callable(value) else None - - -def _as_aocr(value: object) -> RustAocr | None: - return cast(RustAocr, value) if callable(value) else None - - -_OCR: Final = NativeBinding("ocr", validate=_as_ocr) -_AOCR: Final = NativeBinding("aocr", validate=_as_aocr) +_RUST_OCR_PROVIDERS: Final = frozenset( + { + "mistral", + "azure_ai", + "vertex_ai", + } +) def load_rust_ocr() -> RustOcr | None: @@ -61,53 +62,150 @@ def load_rust_aocr() -> RustAocr | None: return _AOCR.load() -def ocr( - *, - model: str, - document: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - rust_ocr: Final = load_rust_ocr() - if rust_ocr is None: - return None - return rust_ocr( - model=model, - document=document, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=_timeout_to_seconds(timeout), +def _rust_ocr_supported(prepared_request: PreparedOCRRequest) -> bool: + if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native": + return False + if not prepared_request.provider_config.supports_rust_bridge(): + return False + return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS + + +def _rust_bridge_optional_params( + prepared_request: PreparedOCRRequest, + resolve_secret: Callable[[str], str | None], +) -> dict[str, object]: + if prepared_request.custom_llm_provider != "vertex_ai": + return prepared_request.optional_params + vertex_project: Final = ( + prepared_request.litellm_params.get("vertex_project") + or prepared_request.litellm_params.get("vertex_ai_project") + or litellm.vertex_project + or resolve_secret("VERTEXAI_PROJECT") + ) + vertex_location: Final = ( + prepared_request.litellm_params.get("vertex_location") + or prepared_request.litellm_params.get("vertex_ai_location") + or litellm.vertex_location + or resolve_secret("VERTEXAI_LOCATION") + or resolve_secret("VERTEX_LOCATION") + ) + return { + **prepared_request.optional_params, + **{ + name: value + for name, value in (("vertex_project", vertex_project), ("vertex_location", vertex_location)) + if value is not None + }, + } + + +def _rust_bridge_api_base( + prepared_request: PreparedOCRRequest, + resolve_secret: Callable[[str], str | None], +) -> str | None: + if prepared_request.api_base is not None: + return prepared_request.api_base + if prepared_request.custom_llm_provider == "azure_ai": + if is_azure_document_intelligence_model(prepared_request.model): + return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT") + return resolve_secret("AZURE_AI_API_BASE") + return None + + +def _prepare_rust_ocr_call( + prepared_request: PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], +) -> _PreparedRustOCRCall: + provider_config: Final = prepared_request.provider_config + api_key_env_var: Final = provider_config.get_api_key_env_var() + resolved_api_key: Final = prepared_request.api_key or ( + resolve_api_key(api_key_env_var) if api_key_env_var is not None else None + ) + resolved_headers: Final = _HEADERS.validate_python( + provider_config.validate_environment( + headers=prepared_request.extra_headers or {}, + model=prepared_request.model, + api_key=resolved_api_key, + api_base=prepared_request.api_base, + litellm_params=prepared_request.litellm_params, + ) + ) + resolved_complete_url: Final = provider_config.get_complete_url( + api_base=prepared_request.api_base, + model=prepared_request.model, + optional_params=prepared_request.optional_params, + litellm_params=prepared_request.litellm_params, + ) + rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key) + rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key) + prepared_request.litellm_logging_obj.pre_call( + input="OCR document processing", + api_key=resolved_api_key, + additional_args={ + "complete_input_dict": { + "model": prepared_request.model, + "document": prepared_request.document, + **rust_optional_params, + }, + "api_base": resolved_complete_url, + "headers": resolved_headers, + }, + ) + return _PreparedRustOCRCall( + api_key=resolved_api_key, + api_base=rust_api_base, + headers=resolved_headers, + optional_params=rust_optional_params, ) -async def aocr( - *, - model: str, - document: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - rust_aocr: Final = load_rust_aocr() - if rust_aocr is None: - return None - return await rust_aocr( - model=model, - document=document, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=_timeout_to_seconds(timeout), +def attempt_ocr( + prepared_request: PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], +) -> DispatchResult[OCRResponse]: + return attempt( + load=_OCR.load, + enabled=_configuration.rust_enabled(), + prepare=lambda: _prepare_rust_ocr_call( + prepared_request=prepared_request, + resolve_api_key=resolve_api_key, + ), + call=lambda native, prepared: native( + model=prepared_request.model, + document=prepared_request.document, + api_key=prepared.api_key, + api_base=prepared.api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + extra_headers=prepared.headers, + optional_params=prepared.optional_params, + timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), + ), + adapt=OCRResponse.model_validate, + eligible=_rust_ocr_supported(prepared_request), + ) + + +async def aattempt_ocr( + prepared_request: PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], +) -> DispatchResult[OCRResponse]: + return await aattempt( + load=_AOCR.load, + enabled=_configuration.rust_enabled(), + prepare=lambda: _prepare_rust_ocr_call( + prepared_request=prepared_request, + resolve_api_key=resolve_api_key, + ), + call=lambda native, prepared: native( + model=prepared_request.model, + document=prepared_request.document, + api_key=prepared.api_key, + api_base=prepared.api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + extra_headers=prepared.headers, + optional_params=prepared.optional_params, + timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), + ), + adapt=OCRResponse.model_validate, + eligible=_rust_ocr_supported(prepared_request), ) diff --git a/litellm/rust_bridge/protocols.py b/litellm/rust_bridge/protocols.py new file mode 100644 index 00000000000..b08fcbf81aa --- /dev/null +++ b/litellm/rust_bridge/protocols.py @@ -0,0 +1,180 @@ +from __future__ import annotations + +from collections.abc import Awaitable, Mapping, Sequence +from typing import Protocol + + +class RustChatCompletions(Protocol): + def __call__( + self, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object] | None, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: Mapping[str, object] | None, + timeout_seconds: float | None, + ) -> Mapping[str, object]: ... + + +class RustAchatCompletions(Protocol): + def __call__( + self, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object] | None, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: Mapping[str, object] | None, + timeout_seconds: float | None, + ) -> Awaitable[Mapping[str, object]]: ... + + +class RustChatCompletionsDecline(Protocol): + def __call__( + self, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object] | None, + custom_llm_provider: str | None, + ) -> str | None: ... + + +class RustResponsesWebSocket(Protocol): + async def send_text(self, text: str) -> None: ... + + async def recv_text(self) -> str | None: ... + + async def close(self) -> None: ... + + +class RustResponsesWebSocketConnection(Protocol): + @classmethod + async def connect( + cls, + url: str, + headers: dict[str, str], + timeout_seconds: float | None, + ) -> RustResponsesWebSocket: ... + + +class NativeModule(Protocol): + @property + def chat_completions(self) -> RustChatCompletions: ... + + @property + def achat_completions(self) -> RustAchatCompletions: ... + + @property + def chat_completions_decline(self) -> RustChatCompletionsDecline: ... + + @property + def ResponsesWebSocketConnection(self) -> type[RustResponsesWebSocketConnection]: ... + + @property + def RustBridgeDeclined(self) -> type[BaseException]: ... + + @property + def RustUpstreamError(self) -> type[BaseException]: ... + + @property + def messages(self) -> RustMessages: ... + + @property + def amessages(self) -> RustAmessages: ... + + @property + def ocr(self) -> RustOcr: ... + + @property + def aocr(self) -> RustAocr: ... + + @property + def transcription(self) -> RustTranscription: ... + + @property + def atranscription(self) -> RustAtranscription: ... + + +class RustMessages(Protocol): + def __call__( + self, + model: str, + body: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None, + timeout_seconds: float | None, + ) -> dict[str, object]: ... + + +class RustAmessages(Protocol): + def __call__( + self, + model: str, + body: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None, + timeout_seconds: float | None, + ) -> Awaitable[dict[str, object]]: ... + + +class RustOcr(Protocol): + def __call__( + self, + model: str, + document: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout_seconds: float | None, + ) -> dict[str, object]: ... + + +class RustAocr(Protocol): + def __call__( + self, + model: str, + document: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout_seconds: float | None, + ) -> Awaitable[dict[str, object]]: ... + + +class RustTranscription(Protocol): + def __call__( + self, + model: str, + audio: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout_seconds: float | None, + ) -> dict[str, object]: ... + + +class RustAtranscription(Protocol): + def __call__( + self, + model: str, + audio: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout_seconds: float | None, + ) -> Awaitable[dict[str, object]]: ... diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index 0634867af1c..d673ab8431c 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -2,70 +2,39 @@ from __future__ import annotations -from dataclasses import dataclass -from typing import Final, Protocol +from collections.abc import AsyncGenerator +from contextlib import AbstractAsyncContextManager, asynccontextmanager +from typing import Final import httpx from websockets.exceptions import ConnectionClosedOK -from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged +from litellm.rust_bridge.configuration import rust_enabled +from litellm.rust_bridge.protocols import ( + RustResponsesWebSocket, + RustResponsesWebSocketConnection, +) +from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result from litellm.rust_bridge.timeouts import timeout_to_seconds - -class RustResponsesWebSocket(Protocol): - async def send_text(self, text: str) -> None: ... - - async def recv_text(self) -> str | None: ... - - async def close(self) -> None: ... - - -class RustResponsesWebSocketConnection(Protocol): - @classmethod - async def connect( - cls, - url: str, - headers: dict[str, str], - timeout_seconds: float | None, - ) -> RustResponsesWebSocket: ... - - -class _Unset: - pass - - -_UNSET: Final[_Unset] = _Unset() - - -@dataclass(slots=True) -class _RustResponsesWebSocketState: - connection: RustResponsesWebSocketConnection | None = None - - -_STATE: Final[_RustResponsesWebSocketState] = _RustResponsesWebSocketState() +_RESPONSES_WEBSOCKET: Final[NativeBinding[RustResponsesWebSocketConnection]] = NativeBinding( + lambda native: native.ResponsesWebSocketConnection, +) def set_rust_responses_websocket( *, - connection: RustResponsesWebSocketConnection | None | _Unset = _UNSET, + connection: RustResponsesWebSocketConnection | None | Unchanged = UNCHANGED, ) -> None: - if not isinstance(connection, _Unset): - _STATE.connection = connection + if not isinstance(connection, Unchanged): + if connection is None: + _RESPONSES_WEBSOCKET.reset() + else: + _RESPONSES_WEBSOCKET.override(connection) -def load_rust_responses_websocket() -> RustResponsesWebSocketConnection | None: - if _STATE.connection is not None: - return _STATE.connection - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - connection_type: Final[RustResponsesWebSocketConnection | None] = getattr( - native_bridge, "ResponsesWebSocketConnection", None - ) - return connection_type - - -class _ConnectionAdapter: +class ConnectionAdapter: def __init__(self, connection: RustResponsesWebSocket): self._connection: Final[RustResponsesWebSocket] = connection @@ -87,16 +56,34 @@ async def connect( url: str, headers: dict[str, str], timeout: float | httpx.Timeout | None, -) -> _ConnectionAdapter | None: - connection_type: Final = load_rust_responses_websocket() - if connection_type is None: - return None - try: - connection: Final = await connection_type.connect( +) -> DispatchResult[ConnectionAdapter]: + return await aattempt( + load=_RESPONSES_WEBSOCKET.load, + enabled=rust_enabled(), + eligible=True, + prepare=lambda: timeout_to_seconds(timeout), + call=lambda connection_type, timeout_seconds: connection_type.connect( url=url, headers=headers, - timeout_seconds=timeout_to_seconds(timeout), - ) - except Exception: # noqa: BLE001 # bridge failures must fall back to Python - return None - return _ConnectionAdapter(connection) + timeout_seconds=timeout_seconds, + ), + adapt=ConnectionAdapter, + ) + + +@asynccontextmanager +async def _connection_context(connection: ConnectionAdapter) -> AsyncGenerator[ConnectionAdapter, None]: + try: + yield connection + finally: + await connection.close() + + +async def managed_connect( + *, + url: str, + headers: dict[str, str], + timeout: float | httpx.Timeout | None, +) -> DispatchResult[AbstractAsyncContextManager[ConnectionAdapter]]: + result: Final = await connect(url=url, headers=headers, timeout=timeout) + return adapt_result(result, _connection_context) diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py index d411673439f..b908a32e879 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -3,148 +3,93 @@ from __future__ import annotations from collections.abc import Awaitable, Callable from dataclasses import dataclass from enum import Enum -from typing import Final, Generic, NoReturn, TypeAlias, TypeVar - -from litellm.exceptions import APIError -from litellm.rust_bridge.bindings import native_exception_types +from typing import Final, Generic, TypeAlias, TypeVar +BindingT = TypeVar("BindingT") NativeT = TypeVar("NativeT") +RequestT = TypeVar("RequestT") ResultT = TypeVar("ResultT") -class FallbackMode(Enum): - PYTHON = "python" - RUST_REQUIRED = "rust_required" +class NativeSkipReason(Enum): + DISABLED = "disabled" + INELIGIBLE = "ineligible" + UNAVAILABLE = "unavailable" + DECLINED = "declined" + FAILED = "failed" @dataclass(frozen=True, slots=True) -class RustHandled(Generic[ResultT]): +class Handled(Generic[ResultT]): value: ResultT @dataclass(frozen=True, slots=True) -class RustDeclined: - reason: str +class NativeSkipped: + reason: NativeSkipReason + detail: str | None = None @dataclass(frozen=True, slots=True) -class RustUnavailable: - pass +class NativeFailed: + error: Exception -RustAttempt: TypeAlias = RustHandled[ResultT] | RustDeclined | RustUnavailable +DispatchResult: TypeAlias = Handled[ResultT] | NativeSkipped | NativeFailed -@dataclass(frozen=True, slots=True) -class BridgeErrorContext: - route: str - provider: str - model: str - - -def invoke( - *, - native_call: Callable[[], NativeT] | None, - fallback: Callable[[], ResultT], - adapt: Callable[[NativeT], ResultT], - mode: FallbackMode, - context: BridgeErrorContext, -) -> ResultT: - result: Final = attempt(native_call=native_call, adapt=adapt, context=context) - if isinstance(result, RustHandled): - return result.value - if mode is FallbackMode.PYTHON: - return fallback() - _raise_required(result, context) - - -async def ainvoke( - *, - native_call: Callable[[], Awaitable[NativeT]] | None, - fallback: Callable[[], Awaitable[ResultT]], - adapt: Callable[[NativeT], ResultT], - mode: FallbackMode, - context: BridgeErrorContext, -) -> ResultT: - result: Final = await aattempt(native_call=native_call, adapt=adapt, context=context) - if isinstance(result, RustHandled): - return result.value - if mode is FallbackMode.PYTHON: - return await fallback() - _raise_required(result, context) +def _select(load: Callable[[], BindingT | None], enabled: bool, eligible: bool) -> BindingT | NativeSkipped: + if not enabled: + return NativeSkipped(NativeSkipReason.DISABLED) + if not eligible: + return NativeSkipped(NativeSkipReason.INELIGIBLE) + binding: Final = load() + return NativeSkipped(NativeSkipReason.UNAVAILABLE) if binding is None else binding def attempt( *, - native_call: Callable[[], NativeT] | None, + load: Callable[[], BindingT | None], + enabled: bool, + eligible: bool, + prepare: Callable[[], RequestT], + call: Callable[[BindingT, RequestT], NativeT], adapt: Callable[[NativeT], ResultT], - context: BridgeErrorContext, -) -> RustAttempt[ResultT]: - if native_call is None: - return RustUnavailable() - exceptions: Final = native_exception_types() - if exceptions is None: - return RustHandled(adapt(native_call())) - declined, upstream = exceptions +) -> DispatchResult[ResultT]: + binding: Final = _select(load, enabled, eligible) + if isinstance(binding, NativeSkipped): + return binding try: - value: Final = native_call() - except declined as error: - return RustDeclined(reason=_decline_reason(error)) - except upstream as error: - _raise_upstream(error, context) - return RustHandled(adapt(value)) + value: Final = call(binding, prepare()) + except Exception as error: # noqa: BLE001 # orchestration applies the endpoint's declared error policy + return NativeFailed(error) + return Handled(adapt(value)) async def aattempt( *, - native_call: Callable[[], Awaitable[NativeT]] | None, + load: Callable[[], BindingT | None], + enabled: bool, + eligible: bool, + prepare: Callable[[], RequestT], + call: Callable[[BindingT, RequestT], Awaitable[NativeT]], adapt: Callable[[NativeT], ResultT], - context: BridgeErrorContext, -) -> RustAttempt[ResultT]: - if native_call is None: - return RustUnavailable() - exceptions: Final = native_exception_types() - if exceptions is None: - return RustHandled(adapt(await native_call())) - declined, upstream = exceptions +) -> DispatchResult[ResultT]: + binding: Final = _select(load, enabled, eligible) + if isinstance(binding, NativeSkipped): + return binding try: - value: Final = await native_call() - except declined as error: - return RustDeclined(reason=_decline_reason(error)) - except upstream as error: - _raise_upstream(error, context) - return RustHandled(adapt(value)) + value: Final = await call(binding, prepare()) + except Exception as error: # noqa: BLE001 # orchestration applies the endpoint's declared error policy + return NativeFailed(error) + return Handled(adapt(value)) -def _decline_reason(error: BaseException) -> str: - reason: Final[object] = error.args[0] if error.args else str(error) - return reason if isinstance(reason, str) else str(reason) +def identity(value: ResultT) -> ResultT: + return value -def _raise_required( - result: RustDeclined | RustUnavailable, - context: BridgeErrorContext, -) -> NoReturn: - raise RuntimeError(f"Rust {context.route} bridge {_required_reason(result)}") - - -def _required_reason(result: RustDeclined | RustUnavailable) -> str: - match result: - case RustUnavailable(): - return "is unavailable" - case RustDeclined(reason=reason): - return f"declined the request: {reason}" - - -def _raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn: - args: Final[tuple[object, ...]] = error.args - status_value: Final = args[0] if args else 0 - message_value: Final = args[1] if len(args) > 1 else str(error) - status: Final = status_value if isinstance(status_value, int) else 0 - message: Final = message_value if isinstance(message_value, str) else str(message_value) - raise APIError( - status_code=status or 500, - message=f"litellm rust {context.route}: {message}", - llm_provider=context.provider, - model=context.model, - ) from error +def adapt_result(result: DispatchResult[NativeT], adapt: Callable[[NativeT], ResultT]) -> DispatchResult[ResultT]: + if isinstance(result, Handled): + return Handled(adapt(result.value)) + return result diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py index 6c81786accd..4970b86b18d 100644 --- a/litellm/rust_bridge/transcription.py +++ b/litellm/rust_bridge/transcription.py @@ -1,99 +1,41 @@ from __future__ import annotations -from collections.abc import Awaitable -from dataclasses import dataclass -from typing import Final, Protocol, cast +from typing import Final import httpx +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged +from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription +from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity from litellm.rust_bridge.timeouts import timeout_to_seconds - -class RustTranscription(Protocol): - def __call__( - self, - model: str, - audio: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, - ) -> dict[str, object]: - raise NotImplementedError - - -class RustAtranscription(Protocol): - def __call__( - self, - model: str, - audio: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, - ) -> Awaitable[dict[str, object]]: - raise NotImplementedError - - -class _Unset: - pass - - -_UNSET: Final[_Unset] = _Unset() - - -@dataclass -class _RustTranscriptionState: - transcription: RustTranscription | None = None - atranscription: RustAtranscription | None = None - - -_STATE: Final = _RustTranscriptionState() +_TRANSCRIPTION: Final[NativeBinding[RustTranscription]] = NativeBinding(lambda native: native.transcription) +_ATRANSCRIPTION: Final[NativeBinding[RustAtranscription]] = NativeBinding(lambda native: native.atranscription) def configure_rust_transcription( *, - transcription: RustTranscription | None | _Unset = _UNSET, - atranscription: RustAtranscription | None | _Unset = _UNSET, + transcription: RustTranscription | None | Unchanged = UNCHANGED, + atranscription: RustAtranscription | None | Unchanged = UNCHANGED, ) -> None: - if not isinstance(transcription, _Unset): - _STATE.transcription = transcription - if not isinstance(atranscription, _Unset): - _STATE.atranscription = atranscription + if not isinstance(transcription, Unchanged): + if transcription is None: + _TRANSCRIPTION.reset() + else: + _TRANSCRIPTION.override(transcription) + if not isinstance(atranscription, Unchanged): + if atranscription is None: + _ATRANSCRIPTION.reset() + else: + _ATRANSCRIPTION.override(atranscription) def load_rust_transcription() -> RustTranscription | None: - if _STATE.transcription is not None: - return _STATE.transcription - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - return ( - None - if native_bridge is None - else cast( # cast-ok: native extension protocol is runtime-defined - RustTranscription, getattr(native_bridge, "transcription", None) - ) - ) + return _TRANSCRIPTION.load() def load_rust_atranscription() -> RustAtranscription | None: - if _STATE.atranscription is not None: - return _STATE.atranscription - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - return ( - None - if native_bridge is None - else cast( # cast-ok: native extension protocol is runtime-defined - RustAtranscription, getattr(native_bridge, "atranscription", None) - ) - ) + return _ATRANSCRIPTION.load() def transcription( @@ -106,19 +48,23 @@ def transcription( extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - rust_transcription: Final = load_rust_transcription() - if rust_transcription is None: - return None - return rust_transcription( - model=model, - audio=audio, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=timeout_to_seconds(timeout), +) -> DispatchResult[dict[str, object]]: + return attempt( + load=_TRANSCRIPTION.load, + enabled=True, + eligible=True, + prepare=lambda: timeout_to_seconds(timeout), + call=lambda rust_transcription, timeout_seconds: rust_transcription( + model=model, + audio=audio, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout_seconds=timeout_seconds, + ), + adapt=identity, ) @@ -132,17 +78,21 @@ async def atranscription( extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - rust_atranscription: Final = load_rust_atranscription() - if rust_atranscription is None: - return None - return await rust_atranscription( - model=model, - audio=audio, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=timeout_to_seconds(timeout), +) -> DispatchResult[dict[str, object]]: + return await aattempt( + load=_ATRANSCRIPTION.load, + enabled=True, + eligible=True, + prepare=lambda: timeout_to_seconds(timeout), + call=lambda rust_atranscription, timeout_seconds: rust_atranscription( + model=model, + audio=audio, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout_seconds=timeout_seconds, + ), + adapt=identity, ) diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index fd7b30bc314..613526ff4d5 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -12,7 +12,7 @@ "limit": 1979 }, "ANN202": { - "limit": 831 + "limit": 830 }, "ANN204": { "limit": 683 @@ -150,13 +150,13 @@ "limit": 253 }, "PLW0127": { - "limit": 57 + "limit": 55 }, "PLW0602": { "limit": 215 }, "PLW0603": { - "limit": 190 + "limit": 184 }, "PLW1508": { "limit": 190 @@ -195,7 +195,7 @@ "limit": 22 }, "SIM101": { - "limit": 56 + "limit": 55 }, "SIM102": { "limit": 310 @@ -231,7 +231,7 @@ "limit": 5 }, "TID251": { - "limit": 1035 + "limit": 1029 }, "TRY002": { "limit": 524 @@ -246,7 +246,7 @@ "limit": 109 }, "TRY300": { - "limit": 852 + "limit": 846 }, "UP028": { "limit": 2 diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index a30474245c6..d2096bc515b 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -9,6 +9,7 @@ import pytest import litellm from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge import configuration +from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, ) @@ -133,9 +134,9 @@ def test_load_rust_amessages_returns_injected_impl(): assert rust_messages.load_rust_amessages() is bridge -def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch): +def test_messages_wrapper_reports_unavailable(monkeypatch): monkeypatch.setattr( - importlib.import_module("litellm.rust_bridge"), + importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: None, ) @@ -150,7 +151,7 @@ def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch): extra_headers={}, timeout=30.0, ) - assert result is None + assert result == NativeSkipped(NativeSkipReason.UNAVAILABLE) def test_messages_wrapper_forwards_args_and_converts_timeout(): @@ -168,7 +169,7 @@ def test_messages_wrapper_forwards_args_and_converts_timeout(): timeout=httpx.Timeout(600.0, read=42.0), ) - assert response == FAKE_MESSAGES_RESPONSE + assert response == Handled(FAKE_MESSAGES_RESPONSE) assert bridge.calls[0] == { "model": "claude-sonnet-4-5", "body": REQUEST_BODY, @@ -196,7 +197,7 @@ async def test_amessages_wrapper_forwards_args(): timeout=12.5, ) - assert response == FAKE_MESSAGES_RESPONSE + assert response == Handled(FAKE_MESSAGES_RESPONSE) assert bridge.calls[0]["model"] == "claude-sonnet-4-5" assert bridge.calls[0]["timeout_seconds"] == 12.5 @@ -214,7 +215,7 @@ def _gate(**overrides): "timeout": 30.0, } kwargs.update(overrides) - return BaseLLMHTTPHandler._maybe_rust_anthropic_messages(**kwargs) + return BaseLLMHTTPHandler._attempt_rust_anthropic_messages(**kwargs) @pytest.mark.asyncio @@ -225,7 +226,8 @@ async def test_gate_invokes_rust_and_marks_response_header(): response = await _gate() - assert response is not None + assert isinstance(response, Handled) + response = response.value assert response["id"] == "msg_123" assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} call = bridge.calls[0] @@ -238,14 +240,13 @@ async def test_gate_invokes_rust_and_marks_response_header(): @pytest.mark.asyncio -async def test_gate_falls_back_to_python_when_bridge_raises(): +async def test_gate_reports_failure_to_harness(): bridge = RaisingAsyncMessages() litellm.rust(True) rust_messages.set_rust_messages(amessages=bridge) response = await _gate() - - assert response is None + assert isinstance(response, NativeFailed) assert bridge.calls == 1 @@ -256,7 +257,7 @@ async def test_gate_skips_rust_when_flag_absent(): response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure")) - assert response is None + assert isinstance(response, NativeSkipped) assert bridge.calls == 0 @@ -268,7 +269,8 @@ async def test_gate_uses_process_enable_without_request_override(): response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure")) - assert response is not None + assert isinstance(response, Handled) + response = response.value assert bridge.calls[0]["custom_llm_provider"] == "azure_ai" @@ -286,7 +288,8 @@ async def test_gate_invokes_rust_for_native_anthropic_provider(): headers={"x-api-key": "sk-ant", "anthropic-version": "2023-06-01"}, ) - assert response is not None + assert isinstance(response, Handled) + response = response.value assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} assert bridge.calls[0]["custom_llm_provider"] == "anthropic" assert bridge.calls[0]["api_key"] == "sk-ant" @@ -303,7 +306,8 @@ async def test_gate_invokes_rust_when_env_var_set(monkeypatch): litellm_params=GenericLiteLLMParams(api_key="sk-ant"), ) - assert response is not None + assert isinstance(response, Handled) + response = response.value assert bridge.calls[0]["custom_llm_provider"] == "anthropic" @@ -318,7 +322,7 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch): litellm_params=GenericLiteLLMParams(api_key="sk-ant"), ) - assert response is None + assert isinstance(response, NativeSkipped) assert bridge.calls == 0 @@ -330,7 +334,7 @@ async def test_gate_skips_rust_for_unsupported_provider(): response = await _gate(custom_llm_provider="openai") - assert response is None + assert isinstance(response, NativeSkipped) assert bridge.calls == 0 @@ -342,7 +346,7 @@ async def test_gate_skips_rust_for_agentic_hook(): response = await _gate(has_agentic_hook=True) - assert response is None + assert isinstance(response, NativeSkipped) assert bridge.calls == 0 @@ -358,7 +362,8 @@ async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag(): request_body=streaming_body, ) - assert response is not None + assert isinstance(response, Handled) + response = response.value assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} assert "stream" not in bridge.calls[0]["body"] assert bridge.calls[0]["body"] == REQUEST_BODY @@ -383,7 +388,7 @@ async def test_fake_stream_wraps_rust_response_as_anthropic_sse(): @pytest.mark.asyncio async def test_gate_falls_back_when_bridge_unavailable(monkeypatch): monkeypatch.setattr( - importlib.import_module("litellm.rust_bridge"), + importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: None, ) @@ -391,4 +396,56 @@ async def test_gate_falls_back_when_bridge_unavailable(monkeypatch): response = await _gate() - assert response is None + assert isinstance(response, NativeSkipped) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("selection", ("native", "disabled", "failed")) +async def test_messages_handler_runs_selected_backend_once(selection: str) -> None: + from datetime import datetime + + import httpx + + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.llms.anthropic.experimental_pass_through.messages.transformation import AnthropicMessagesConfig + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + bridge = RaisingAsyncMessages() if selection == "failed" else RecordingAsyncMessages() + rust_messages.set_rust_messages(amessages=bridge) + litellm.rust(selection != "disabled") + requests: list[httpx.Request] = [] + + def respond(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(200, json=FAKE_MESSAGES_RESPONSE) + + logging_obj = Logging( + model=FAKE_MESSAGES_RESPONSE["model"], + messages=[], + stream=False, + call_type="anthropic_messages", + start_time=datetime.now(), + litellm_call_id="harness-test", + function_id="harness-test", + ) + client = AsyncHTTPHandler() + await client.client.aclose() + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as transport: + client.client = transport + response = await BaseLLMHTTPHandler().async_anthropic_messages_handler( + model=FAKE_MESSAGES_RESPONSE["model"], + messages=[{"role": "user", "content": "hello"}], + anthropic_messages_provider_config=AnthropicMessagesConfig(), + anthropic_messages_optional_request_params={"max_tokens": 10}, + custom_llm_provider="anthropic", + litellm_params=GenericLiteLLMParams(), + logging_obj=logging_obj, + api_key="sk-test", + api_base="https://example.test", + client=client, + ) + assert response["id"] == FAKE_MESSAGES_RESPONSE["id"] + assert len(requests) == (0 if selection == "native" else 1) + assert (bridge.calls if isinstance(bridge, RaisingAsyncMessages) else len(bridge.calls)) == ( + 0 if selection == "disabled" else 1 + ) diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py index b4b173b20c3..da70f422f44 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py @@ -2467,7 +2467,7 @@ class TestRustChatCompletionsHook: def declining_native(**_kwargs): raise _Declined("blank message text") - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) bridge.set_rust_chat_completions( decline=lambda **_kwargs: None, chat_completions=declining_native ) @@ -2499,7 +2499,7 @@ class TestRustChatCompletionsHook: RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) async def declining_native(**_kwargs): raise _Declined("blank message text") @@ -2559,7 +2559,7 @@ class TestRustChatCompletionsHook: RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) def declining_native(**_kwargs): raise _Declined("blank message text") diff --git a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py index f34b8eb1fb9..becd6ecb832 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py +++ b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py @@ -205,7 +205,7 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch): RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) async def declining_native(**_kwargs): raise _Declined("blank message text") @@ -282,7 +282,7 @@ async def test_pre_call_logging_fires_once_even_when_the_rust_path_declines(): return ModelResponse() with ( - patch.object(bridge, "get_native_bridge", lambda: _FakeNative()), + patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()), patch.object( BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS ), @@ -389,7 +389,7 @@ def test_pre_call_logging_fires_once_when_the_sync_rust_path_declines(): logging_obj = MagicMock() - with patch.object(bridge, "get_native_bridge", lambda: _FakeNative()): + with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()): bridge.set_rust_chat_completions( decline=lambda **_kwargs: None, chat_completions=declining_native ) @@ -475,7 +475,7 @@ def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines(): logging_obj, calls = _recording_logging_obj() - with patch.object(bridge, "get_native_bridge", lambda: _FakeNative()): + with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()): bridge.set_rust_chat_completions( decline=lambda **_kwargs: None, chat_completions=declining_native ) diff --git a/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py b/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py index 0c8b1cc2836..c25e52e4421 100644 --- a/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py +++ b/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py @@ -11,7 +11,8 @@ supplied api_base is always honoured. from litellm.llms.azure_ai.ocr.common_utils import ( is_azure_document_intelligence_model, ) -from litellm.ocr.main import _prepare_ocr_request, _rust_bridge_api_base +from litellm.ocr.main import _prepare_ocr_request +from litellm.rust_bridge.ocr import _rust_bridge_api_base _DOC = {"type": "document_url", "document_url": "https://example.com/doc.pdf"} _DOC_INTELLIGENCE_ENDPOINT = "https://di.cognitiveservices.azure.com" diff --git a/tests/test_litellm/ocr/test_ocr_native_format.py b/tests/test_litellm/ocr/test_ocr_native_format.py index 249fbda713e..1ab740c4267 100644 --- a/tests/test_litellm/ocr/test_ocr_native_format.py +++ b/tests/test_litellm/ocr/test_ocr_native_format.py @@ -12,13 +12,13 @@ import pytest import litellm from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig from litellm.llms.cohere.ocr.transformation import CohereParseConfig -from litellm.ocr.main import _PreparedOCRRequest, _rust_ocr_supported +from litellm.rust_bridge.ocr import PreparedOCRRequest, _rust_ocr_supported DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"} -def _prepared(optional_params: dict[str, object]) -> _PreparedOCRRequest: - return _PreparedOCRRequest( +def _prepared(optional_params: dict[str, object]) -> PreparedOCRRequest: + return PreparedOCRRequest( model="doc-intelligence/prebuilt-layout", document=dict(DOCUMENT), api_key="fake-key", diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index d28184b96a6..7e04a4f0f4b 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -3,7 +3,8 @@ import builtins import importlib import types -from typing import Any +from typing import Any, Final +from unittest.mock import AsyncMock, Mock import httpx import pytest @@ -11,13 +12,14 @@ import pytest import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.rust_bridge import configuration +from litellm.rust_bridge.runtime import Handled +from litellm.rust_bridge.timeouts import timeout_to_seconds # `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr` # function onto `litellm.ocr` and shadows the submodule, so import the modules # explicitly via importlib rather than attribute traversal. ocr_main = importlib.import_module("litellm.ocr.main") rust_bridge = importlib.import_module("litellm.rust_bridge.ocr") -rust_bridge_bindings = importlib.import_module("litellm.rust_bridge.bindings") rust_bridge_loader = importlib.import_module("litellm.rust_bridge.loader") MODEL = "mistral/mistral-ocr-latest" @@ -162,6 +164,9 @@ class FakeOCRConfig: def get_api_key_env_var(self) -> str: return self.api_key_env_var + def supports_rust_bridge(self) -> bool: + return True + def validate_environment( self, *, @@ -198,7 +203,7 @@ def build_prepared_request( litellm_params: dict[str, object] | None = None, timeout: float | httpx.Timeout | None = 12.5, ) -> Any: - return ocr_main._PreparedOCRRequest( + return rust_bridge.PreparedOCRRequest( model=model, document=document, api_key=api_key, @@ -334,7 +339,7 @@ def test_toggle_without_ocr_arg_preserves_injected_impl(): def test_explicit_ocr_none_clears_injected_impl(monkeypatch): monkeypatch.setattr( - rust_bridge_bindings, + importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: None, ) @@ -354,7 +359,7 @@ def test_load_rust_ocr_none_when_extension_absent(monkeypatch): """With no injected impl and no compiled wheel, the loader returns None so the caller degrades to the Python path instead of raising ImportError.""" monkeypatch.setattr( - rust_bridge_bindings, + importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: None, ) @@ -371,7 +376,7 @@ def test_load_rust_ocr_uses_compiled_extension(monkeypatch): fake_module.ocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined] fake_module.aocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined] monkeypatch.setattr( - rust_bridge_bindings, + importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: fake_module, ) @@ -382,74 +387,9 @@ def test_load_rust_ocr_uses_compiled_extension(monkeypatch): def test_timeout_to_seconds_handles_float_timeout_and_none(): - assert rust_bridge._timeout_to_seconds(12.5) == 12.5 - assert rust_bridge._timeout_to_seconds(None) is None - assert rust_bridge._timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0 - - -def test_bridge_wrapper_forwards_prepared_args_and_wraps_response(): - bridge = RecordingBridge() - - litellm.rust(True) - - rust_bridge._OCR.override(bridge) - response = rust_bridge.ocr( - model="mistral-ocr-latest", - document=DOCUMENT, - api_key="sk-test", - api_base="https://proxy.internal", - custom_llm_provider="mistral", - extra_headers={"Authorization": "Bearer sk-test", "x-trace-id": "trace-1"}, - optional_params={"include_image_base64": True, "pages": [0]}, - timeout=12.5, - ) - - assert response == FAKE_OCR_RESPONSE - call = bridge.calls[0] - assert call == { - "model": "mistral-ocr-latest", - "document": DOCUMENT, - "api_key": "sk-test", - "api_base": "https://proxy.internal", - "custom_llm_provider": "mistral", - "extra_headers": { - "Authorization": "Bearer sk-test", - "x-trace-id": "trace-1", - }, - "optional_params": {"include_image_base64": True, "pages": [0]}, - "timeout_seconds": 12.5, - } - - -@pytest.mark.asyncio -async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response(): - bridge = RecordingAsyncBridge() - - litellm.rust(True) - - rust_bridge._AOCR.override(bridge) - response = await rust_bridge.aocr( - model="mistral-ocr-maas", - document=DOCUMENT, - api_key=None, - api_base=None, - custom_llm_provider="vertex_ai", - extra_headers=None, - optional_params={"vertex_project": "project-1"}, - timeout=httpx.Timeout(30.0, read=42.0), - ) - - assert response == FAKE_OCR_RESPONSE - assert bridge.calls[0] == { - "model": "mistral-ocr-maas", - "document": DOCUMENT, - "api_key": None, - "api_base": None, - "custom_llm_provider": "vertex_ai", - "extra_headers": None, - "optional_params": {"vertex_project": "project-1"}, - "timeout_seconds": 42.0, - } + assert timeout_to_seconds(12.5) == 12.5 + assert timeout_to_seconds(None) is None + assert timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0 def test_run_rust_ocr_prepares_request_and_wraps_response(): @@ -458,7 +398,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): litellm.rust(True) rust_bridge._OCR.override(bridge) - response = ocr_main._run_rust_ocr( + response = rust_bridge.attempt_ocr( prepared_request=build_prepared_request( logging_obj=logging_obj, api_base="https://proxy.internal", @@ -469,6 +409,8 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): resolve_api_key=lambda _name: None, ) + assert isinstance(response, Handled) + response = response.value assert isinstance(response, OCRResponse) assert response.pages[0].markdown == "hello world" assert bridge.calls[0] == { @@ -491,7 +433,7 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( + rust_bridge.attempt_ocr( prepared_request=build_prepared_request(api_key=None, timeout=None), resolve_api_key=lambda name: "sk-from-vault" if name == "MISTRAL_API_KEY" else None, ) @@ -507,7 +449,7 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver(): def _resolver(name: str) -> str | None: raise AssertionError(f"resolver should not be called for {name}") - ocr_main._run_rust_ocr( + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( api_key="sk-explicit", timeout=None, @@ -528,7 +470,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var(): resolver_calls.append(name) return "sk-provider-env" - ocr_main._run_rust_ocr( + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( provider_config=FakeOCRConfig(api_key_env_var="PROVIDER_OCR_API_KEY"), model="provider-ocr-model", @@ -547,7 +489,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", @@ -580,7 +522,7 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana "VERTEXAI_LOCATION": "us-east5", }.get(name) - ocr_main._run_rust_ocr( + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", @@ -598,7 +540,7 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( custom_llm_provider="azure_ai", model="pixtral-12b-2409", @@ -616,7 +558,7 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( custom_llm_provider="azure_ai", model="doc-intelligence/prebuilt-layout", @@ -637,7 +579,7 @@ def test_run_rust_ocr_runs_pre_call_logging(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( logging_obj=logging_obj, api_base="https://api.mistral.ai/v1", @@ -661,30 +603,6 @@ def test_run_rust_ocr_runs_pre_call_logging(): } -def test_ocr_routes_to_rust_when_enabled(fake_bridge): - response = litellm.ocr( - model=MODEL, - document=DOCUMENT, - api_key="sk-test", - extra_headers={"x-trace-id": "trace-1"}, - include_image_base64=True, - ) - - assert isinstance(response, OCRResponse) - assert response.pages[0].markdown == "hello world" - assert len(fake_bridge.calls) == 1 - call = fake_bridge.calls[0] - assert call["model"] == "mistral-ocr-latest" - assert call["document"] == DOCUMENT - assert call["api_key"] == "sk-test" - assert call["custom_llm_provider"] == "mistral" - assert call["extra_headers"] == { - "Authorization": "Bearer sk-test", - "x-trace-id": "trace-1", - } - assert call["optional_params"].get("include_image_base64") is True - - def test_ocr_routes_azure_ai_to_rust_when_enabled(fake_bridge): response = litellm.ocr( model="azure_ai/pixtral-12b-2409", @@ -794,34 +712,50 @@ def test_ocr_passes_default_request_timeout_to_rust(fake_bridge): assert fake_bridge.calls[0]["timeout_seconds"] == float(request_timeout) -def test_ocr_does_not_route_to_rust_when_disabled(): - """With the flag off, the bridge must not be consulted even if an impl exists.""" - bridge = RecordingBridge() - litellm.rust(False) - rust_bridge._OCR.override(bridge) - # The impl stays available for injection, but the disabled flag gates usage, - # so ocr() never reaches the Rust path (asserted via the enabled-path test). - assert bridge.calls == [] +@pytest.mark.asyncio +@pytest.mark.parametrize("enabled", (False, True)) +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_ocr_fallback_skips_native_preparation( + monkeypatch: pytest.MonkeyPatch, enabled: bool, asynchronous: bool +) -> None: + monkeypatch.setattr(importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: None) + litellm.rust(enabled) + expected: Final = OCRResponse(pages=[], model="mistral-ocr-latest", object="ocr") + fallback: Final = AsyncMock(return_value=expected) if asynchronous else Mock(return_value=expected) + + def unexpected_preparation(*_args: object, **_kwargs: object) -> None: + pytest.fail("Python fallback must not resolve native credentials or emit native pre_call") + + monkeypatch.setattr(rust_bridge, "_prepare_rust_ocr_call", unexpected_preparation) + monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fallback) + + response: Final = ( + await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + if asynchronous + else litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + ) + + assert response is expected + fallback.assert_called_once() -def test_ocr_falls_back_to_python_when_bridge_unavailable(monkeypatch): - """Rust enabled but no bridge available (no injected impl, no compiled wheel): - ocr() must degrade to the Python HTTP handler instead of raising.""" - monkeypatch.setattr(rust_bridge, "load_rust_ocr", lambda: None) - litellm.rust(True) # enabled, but load_rust_ocr() returns None in CI +@pytest.mark.asyncio +async def test_aocr_rejects_empty_python_fallback_response(monkeypatch: pytest.MonkeyPatch) -> None: + captured: dict[str, object] = {} - captured = {} + def fake_exception_type(**kwargs: object) -> CapturedException: + captured.update(kwargs) + return CapturedException("wrapped") - def fake_handler_ocr(**kwargs): - captured["called"] = True - return OCRResponse(pages=[], model="mistral-ocr-latest", object="ocr") + monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) + monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", AsyncMock(return_value=None)) - monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fake_handler_ocr) + with pytest.raises(CapturedException, match="wrapped"): + await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test") - response = litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") - - assert captured.get("called") is True # Python path was used - assert isinstance(response, OCRResponse) + original: Final = captured["original_exception"] + assert isinstance(original, ValueError) + assert str(original) == "Got an unexpected None response from the OCR API: None" def test_ocr_provider_configs_expose_api_key_env_vars(): diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index 74d96bda336..e502689797b 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -4,6 +4,7 @@ import pytest from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled from litellm.rust_bridge import configuration, responses_websocket +from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason class _FakeNativeConnection: @@ -57,31 +58,29 @@ def test_rust_websocket_bridge_uses_process_enablement() -> None: @pytest.mark.asyncio async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None: - adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection()) + adapter = responses_websocket.ConnectionAdapter(_ClosedNativeConnection()) with pytest.raises(responses_websocket.ConnectionClosedOK): await adapter.recv() @pytest.mark.asyncio -async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(responses_websocket, "_STATE", responses_websocket._RustResponsesWebSocketState()) - monkeypatch.setattr(responses_websocket, "get_native_bridge", lambda: None) +async def test_bridge_reports_unavailable(monkeypatch: pytest.MonkeyPatch) -> None: + configuration.rust(True) + responses_websocket._RESPONSES_WEBSOCKET.override(None) - assert ( - await responses_websocket.connect( - url="wss://example.test/responses", - headers={}, - timeout=None, - ) - is None - ) + assert await responses_websocket.connect( + url="wss://example.test/responses", + headers={}, + timeout=None, + ) == NativeSkipped(NativeSkipReason.UNAVAILABLE) @pytest.mark.asyncio async def test_enabled_bridge_connects_and_adapts_socket( monkeypatch: pytest.MonkeyPatch, ) -> None: + configuration.rust(True) responses_websocket.set_rust_responses_websocket(connection=_FakeNativeBridge) connection = await responses_websocket.connect( @@ -90,7 +89,56 @@ async def test_enabled_bridge_connects_and_adapts_socket( timeout=1.0, ) - assert connection is not None + assert isinstance(connection, Handled) + connection = connection.value await connection.send("response.create") assert await connection.recv() == "response.completed" await connection.close() + + +class _FailingNativeBridge: + @classmethod + async def connect( + cls, + *, + url: str, + headers: dict[str, str], + timeout_seconds: float | None, + ) -> _FakeNativeConnection: + raise RuntimeError("connection failed") + + +@pytest.mark.asyncio +async def test_connection_failure_is_reported_to_orchestration() -> None: + configuration.rust(True) + responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge) + result = await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None) + assert isinstance(result, NativeFailed) + assert str(result.error) == "connection failed" + + +@pytest.mark.asyncio +async def test_managed_connection_closes_native_socket_on_consumer_failure() -> None: + configuration.rust(True) + socket = _FakeNativeConnection() + + class Bridge: + @classmethod + async def connect( + cls, *, url: str, headers: dict[str, str], timeout_seconds: float | None + ) -> _FakeNativeConnection: + return socket + + responses_websocket.set_rust_responses_websocket(connection=Bridge) + result = await responses_websocket.managed_connect(url="wss://example.test/responses", headers={}, timeout=1.0) + assert isinstance(result, Handled) + + async def use_connection() -> None: + async with result.value as connection: + await connection.send("hello") + raise ValueError("consumer failed") + + with pytest.raises(ValueError, match="consumer failed"): + await use_connection() + assert socket.sent == ["hello"] + assert socket.closed diff --git a/tests/test_litellm/rust_bridge/test_bindings.py b/tests/test_litellm/rust_bridge/test_bindings.py index 88036a5a556..cd562ee91c1 100644 --- a/tests/test_litellm/rust_bridge/test_bindings.py +++ b/tests/test_litellm/rust_bridge/test_bindings.py @@ -1,3 +1,7 @@ +import json +import subprocess +import sys +from pathlib import Path from types import SimpleNamespace from typing import Final @@ -6,30 +10,104 @@ import pytest from litellm.rust_bridge import bindings -def test_binding_distinguishes_disable_from_reset(monkeypatch) -> None: - native = SimpleNamespace(route=lambda: "native") +def test_binding_distinguishes_disable_from_reset(monkeypatch: pytest.MonkeyPatch) -> None: + native: Final = SimpleNamespace(chat_completions=lambda: "native") monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) - binding: bindings.NativeBinding[object] = bindings.NativeBinding("route", validate=lambda value: value) - - assert binding.load() is native.route + binding: Final = bindings.NativeBinding(lambda module: module.chat_completions) + assert binding.load() is native.chat_completions binding.override(None) assert binding.load() is None - - replacement = object() - binding.override(replacement) - assert binding.load() is replacement - + replacement: Final = SimpleNamespace(chat_completions=lambda: "replacement") + binding.override(replacement.chat_completions) + assert binding.load() is replacement.chat_completions binding.reset() - assert binding.load() is native.route + assert binding.load() is native.chat_completions -@pytest.mark.parametrize(("value", "expected"), ((3, 3), ("invalid", None), (None, None))) -def test_binding_validates_native_attribute( - monkeypatch: pytest.MonkeyPatch, value: object, expected: int | None -) -> None: - native: Final = SimpleNamespace(route=value) +@pytest.mark.parametrize("native", (None, SimpleNamespace(), SimpleNamespace(chat_completions=3))) +def test_missing_or_invalid_export_is_unavailable(monkeypatch: pytest.MonkeyPatch, native: object) -> None: monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) - binding: Final = bindings.NativeBinding("route", validate=lambda item: item if isinstance(item, int) else None) + binding: Final = bindings.NativeBinding(lambda module: module.chat_completions) - assert binding.load() == expected + assert binding.load() is None + + +def test_selection_is_lazy_and_preserves_other_exports(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(bindings, "get_native_bridge", lambda: pytest.fail("must not load during construction")) + binding: Final = bindings.NativeBinding(lambda module: module.chat_completions) + native: Final = SimpleNamespace(chat_completions=lambda: "native", achat_completions=None) + monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) + + assert binding.load() is native.chat_completions + assert bindings.NativeBinding(lambda module: module.achat_completions).load() is None + + +@pytest.mark.parametrize("invalid", (None, str, lambda: None)) +def test_native_exception_types_reject_non_exception_classes(monkeypatch: pytest.MonkeyPatch, invalid: object) -> None: + native: Final = SimpleNamespace(RustBridgeDeclined=invalid, RustUpstreamError=RuntimeError) + monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) + + assert bindings.native_declined_types() == () + assert bindings.native_upstream_types() == (RuntimeError,) + + +@pytest.mark.parametrize( + ("expression", "expected_rule"), + ( + ("NativeBinding(lambda native: native.chat_completion)", "reportAttributeAccessIssue"), + ( + "wrong: NativeBinding[RustAchatCompletions] = NativeBinding(lambda native: native.chat_completions)", + "reportAssignmentType", + ), + ("NativeBinding(lambda native: native.ocrr)", "reportAttributeAccessIssue"), + ( + "wrong: NativeBinding[RustAmessages] = NativeBinding(lambda native: native.messages)", + "reportAssignmentType", + ), + ( + "wrong: NativeBinding[RustAocr] = NativeBinding(lambda native: native.ocr)", + "reportAssignmentType", + ), + ( + "wrong: NativeBinding[RustAtranscription] = NativeBinding(lambda native: native.transcription)", + "reportAssignmentType", + ), + ), +) +def test_selectors_are_checked_by_type_checker(tmp_path: Path, expression: str, expected_rule: str) -> None: + source: Final = tmp_path / "binding_contract.py" + source.write_text( + "from typing_extensions import assert_type\n" + "from litellm.rust_bridge.bindings import NativeBinding\n" + "from litellm.rust_bridge.protocols import RustChatCompletions, RustAchatCompletions, " + "RustMessages, RustAmessages, RustOcr, RustAocr, RustTranscription, RustAtranscription\n" + "binding = NativeBinding(lambda native: native.chat_completions)\n" + "assert_type(binding, NativeBinding[RustChatCompletions])\n" + "assert_type(NativeBinding(lambda native: native.messages), NativeBinding[RustMessages])\n" + "assert_type(NativeBinding(lambda native: native.ocr), NativeBinding[RustOcr])\n" + "assert_type(NativeBinding(lambda native: native.transcription), NativeBinding[RustTranscription])\n" + + expression + + "\n" + ) + config: Final = tmp_path / "pyrightconfig.json" + config.write_text( + json.dumps( + { + "include": [str(source)], + "extraPaths": [str(Path(__file__).resolve().parents[3])], + "typeCheckingMode": "basic", + } + ) + ) + result: Final = subprocess.run( + [sys.executable, "-m", "basedpyright", "--project", str(config), "--outputjson"], + capture_output=True, + text=True, + check=False, + ) + diagnostics: Final = json.loads(result.stdout)["generalDiagnostics"] + assert result.returncode == 1, result.stdout + result.stderr + assert [(item["rule"], item["range"]["start"]["line"]) for item in diagnostics] == [ + (expected_rule, len(source.read_text().splitlines()) - 1) + ] diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index b2fd2e6dcc0..f3f2ac651bf 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -10,8 +10,9 @@ from __future__ import annotations import pytest import litellm -from litellm.rust_bridge import configuration +from litellm.rust_bridge import bindings, configuration from litellm.rust_bridge import chat_completions as bridge +from litellm.rust_bridge.runtime import Handled, NativeSkipped, NativeFailed from litellm.types.utils import ModelResponse RUST_RESPONSE = { @@ -54,7 +55,7 @@ class _FakeNative: def _fake_native_bridge(monkeypatch): """Expose the bridge's exception classes without the compiled extension.""" - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative()) def _hide_native_bridge(monkeypatch): @@ -63,7 +64,7 @@ def _hide_native_bridge(monkeypatch): There is no injection seam for "the .so is absent", so the loader itself is replaced; every other case here uses `set_rust_chat_completions`. """ - monkeypatch.setattr(bridge, "get_native_bridge", lambda: None) + monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) @pytest.fixture(autouse=True) @@ -250,7 +251,8 @@ class TestSyncCall: result = bridge.chat_completions(**_call_kwargs(model_response)) - assert result is not None + assert isinstance(result, Handled) + result = result.value assert result.choices[0].message.content == "hello from rust" assert result.choices[0].finish_reason == "stop" assert result.model == "claude-sonnet-4-5-20260101" @@ -266,14 +268,14 @@ class TestSyncCall: bridge.chat_completions(**_call_kwargs(ModelResponse())) assert native.calls[0]["timeout_seconds"] == 30.0 - def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): + def test_reports_unavailable_bridge(self, monkeypatch): _hide_native_bridge(monkeypatch) - assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + assert isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeSkipped) - def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch): + def test_reports_native_decline_to_orchestration(self, monkeypatch): _fake_native_bridge(monkeypatch) bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming"))) - assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + assert isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeFailed) class TestAsyncCall: @@ -281,115 +283,18 @@ class TestAsyncCall: async def test_builds_a_model_response(self): bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall()) result = await bridge.achat_completions(**_call_kwargs(ModelResponse())) - assert result is not None + assert isinstance(result, Handled) + result = result.value assert result.choices[0].message.content == "hello from rust" assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} @pytest.mark.asyncio - async def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): + async def test_reports_unavailable_bridge(self, monkeypatch): _hide_native_bridge(monkeypatch) - assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None + assert isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeSkipped) @pytest.mark.asyncio - async def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch): + async def test_reports_native_decline_to_orchestration(self, monkeypatch): _fake_native_bridge(monkeypatch) bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming"))) - assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None - - -class TestAsyncFallbackWrapper: - @pytest.mark.asyncio - async def test_returns_the_rust_response_without_running_the_fallback(self): - bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall()) - ran = [] - - async def fallback(): - ran.append(True) - return "python" - - result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert result.choices[0].message.content == "hello from rust" - assert ran == [] - - @pytest.mark.asyncio - async def test_runs_the_fallback_when_the_core_declines(self, monkeypatch): - _fake_native_bridge(monkeypatch) - bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming"))) - - async def fallback(): - return "python" - - result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert result == "python" - - @pytest.mark.asyncio - async def test_runs_the_fallback_when_the_bridge_is_unavailable(self, monkeypatch): - _hide_native_bridge(monkeypatch) - - async def fallback(): - return "python" - - result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert result == "python" - - -class TestFailureClassification: - """A failure the provider already saw must not be retried on the Python - path: it would bill the customer for the same work twice.""" - - @pytest.fixture(autouse=True) - def _native_exceptions(self, monkeypatch): - _fake_native_bridge(monkeypatch) - - def test_a_decline_falls_back_because_nothing_was_sent(self): - bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming"))) - assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None - - def test_an_upstream_failure_is_surfaced_with_its_status(self): - from litellm.exceptions import APIError - - bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(429, "429: rate limited"))) - with pytest.raises(APIError) as raised: - bridge.chat_completions(**_call_kwargs(ModelResponse())) - assert raised.value.status_code == 429 - assert "rate limited" in str(raised.value) - - def test_a_transport_failure_with_no_response_surfaces_as_a_500(self): - from litellm.exceptions import APIError - - bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(0, "connection reset"))) - with pytest.raises(APIError) as raised: - bridge.chat_completions(**_call_kwargs(ModelResponse())) - assert raised.value.status_code == 500 - - def test_an_unrecognized_error_is_not_swallowed(self): - bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=RuntimeError("something else"))) - with pytest.raises(RuntimeError): - bridge.chat_completions(**_call_kwargs(ModelResponse())) - - @pytest.mark.asyncio - async def test_the_async_wrapper_does_not_fall_back_on_an_upstream_failure(self): - from litellm.exceptions import APIError - - bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeUpstream(500, "500: boom"))) - ran = [] - - async def fallback(): - ran.append(True) - return "python" - - with pytest.raises(APIError): - await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert ran == [], "a request the provider already served must not be re-issued" - - @pytest.mark.asyncio - async def test_the_async_wrapper_falls_back_on_a_decline(self): - bridge.set_rust_chat_completions( - achat_completions=_RecordingAsyncCall(error=_FakeDeclined("blank message text")) - ) - - async def fallback(): - return "python" - - result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert result == "python" + assert isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeFailed) diff --git a/tests/test_litellm/rust_bridge/test_dispatch.py b/tests/test_litellm/rust_bridge/test_dispatch.py new file mode 100644 index 00000000000..9258372fb93 --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_dispatch.py @@ -0,0 +1,307 @@ +from __future__ import annotations + +import asyncio +import logging +from types import SimpleNamespace +from typing import Final + +import pytest + +from litellm.exceptions import APIError +from litellm.rust_bridge import bindings +from litellm.rust_bridge.chat_completions import error_handling +from litellm.rust_bridge.dispatch import PROPAGATE, PYTHON_ON_ERROR, anative_first, native_first +from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, NativeSkipped, NativeSkipReason + + +class Declined(Exception): + pass + + +class Upstream(Exception): + pass + + +@pytest.fixture(autouse=True) +def native_metadata(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + bindings, "get_native_bridge", lambda: SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream) + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +@pytest.mark.parametrize("reason", tuple(NativeSkipReason)) +async def test_shared_dispatch_calls_python_once_and_logs_skip( + asynchronous: bool, reason: NativeSkipReason, caplog: pytest.LogCaptureFixture +) -> None: + caplog.set_level(logging.DEBUG, logger="LiteLLM") + calls: Final[list[str]] = [] + + def native() -> DispatchResult[str]: + calls.append("native") + return NativeSkipped(reason, "diagnostic detail") + + async def anative() -> DispatchResult[str]: + return native() + + def python() -> str: + calls.append("python") + return "python response" + + async def apython() -> str: + return python() + + result: Final = ( + await anative_first(native=anative, route="test", errors=lambda: PROPAGATE)(apython)() + if asynchronous + else native_first(native=native, route="test", errors=lambda: PROPAGATE)(python)() + ) + assert result == "python response" + assert calls == ["native", "python"] + assert f"Native test skipped ({reason.value}): diagnostic detail" in caplog.text + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_native_success_does_not_run_python_even_when_value_is_none(asynchronous: bool) -> None: + + async def native() -> DispatchResult[None]: + return Handled(None) + + def python() -> str: + pytest.fail("handled results must not run Python") + + async def apython() -> str: + return python() + + result: Final = ( + await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(apython)() + if asynchronous + else native_first(native=lambda: Handled(None), route="test", errors=lambda: PYTHON_ON_ERROR)(python)() + ) + assert result is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +@pytest.mark.parametrize("policy", ("chat", "propagate", "python")) +@pytest.mark.parametrize("kind", ("declined", "upstream", "unknown", "unexpected", "missing")) +async def test_declarations_preserve_endpoint_error_behavior( + monkeypatch: pytest.MonkeyPatch, asynchronous: bool, policy: str, kind: str +) -> None: + if kind == "missing": + monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) + error: Final = ( + Declined("unsupported") + if kind == "declined" + else Upstream(429, "rate limited") + if kind == "upstream" + else RuntimeError("failed") + ) + rules: Final = ( + error_handling("anthropic", "model") + if policy == "chat" + else PYTHON_ON_ERROR + if policy == "python" + else PROPAGATE + ) + calls: Final[list[str]] = [] + + def native() -> DispatchResult[str]: + if kind == "unexpected": + raise error + return NativeFailed(error) + + async def anative() -> DispatchResult[str]: + return native() + + def python() -> str: + calls.append("python") + return "python response" + + async def apython() -> str: + return python() + + async def run() -> str: + if asynchronous: + return await anative_first(native=anative, route="chat_completions", errors=lambda: rules)(apython)() + return native_first(native=native, route="chat_completions", errors=lambda: rules)(python)() + + if policy == "python" or (policy == "chat" and kind in ("declined", "missing")): + assert await run() == "python response" + assert calls == ["python"] + elif policy == "chat" and kind == "upstream": + with pytest.raises(APIError) as caught: + await run() + assert caught.value.status_code == 429 + assert caught.value.model == "model" + assert caught.value.llm_provider == "anthropic" + assert caught.value.__cause__ is error + assert calls == [] + else: + with pytest.raises(type(error)) as caught_original: + await run() + assert caught_original.value is error + assert calls == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_python_failure_is_never_reclassified_as_native_failure(asynchronous: bool) -> None: + error: Final = RuntimeError("Python failed") + calls: Final[list[str]] = [] + + async def native() -> DispatchResult[str]: + return NativeSkipped(NativeSkipReason.UNAVAILABLE) + + def python() -> str: + calls.append("python") + raise error + + async def apython() -> str: + return python() + + async def run() -> str: + if asynchronous: + return await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(apython)() + return native_first( + native=lambda: NativeSkipped(NativeSkipReason.UNAVAILABLE), route="test", errors=lambda: PYTHON_ON_ERROR + )(python)() + + with pytest.raises(RuntimeError) as caught: + await run() + assert caught.value is error + assert calls == ["python"] + + +@pytest.mark.asyncio +async def test_cancellation_does_not_run_python() -> None: + + async def native() -> DispatchResult[str]: + raise asyncio.CancelledError + + async def python() -> str: + pytest.fail("cancellation must not dispatch Python") + + with pytest.raises(asyncio.CancelledError): + await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(python)() + + +@pytest.mark.parametrize("status", (0, 401, 403, 429, 500, 503)) +def test_chat_upstream_mapping_preserves_status_message_and_context(status: int) -> None: + error: Final = Upstream(status, "upstream failed") + with pytest.raises(APIError, match="upstream failed") as caught: + native_first( + native=lambda: NativeFailed(error), + route="chat_completions", + errors=lambda: error_handling("anthropic", "model"), + )(lambda: pytest.fail("upstream errors must not run Python"))() + assert caught.value.status_code == (status or 500) + assert caught.value.model == "model" + assert caught.value.llm_provider == "anthropic" + assert caught.value.__cause__ is error + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_registered_wrapper_preserves_arguments_and_request_error_context(asynchronous: bool) -> None: + calls: Final[list[tuple[str, str, str]]] = [] + + def native(provider: str, *, model: str) -> DispatchResult[str]: + calls.append(("native", provider, model)) + return ( + NativeFailed(Upstream(429, "limited")) + if model == "limited" + else NativeSkipped(NativeSkipReason.UNAVAILABLE) + ) + + async def anative(provider: str, *, model: str) -> DispatchResult[str]: + return native(provider, model=model) + + def rules(provider: str, *, model: str): + return error_handling(provider, model) + + @native_first(native=native, route="chat_completions", errors=rules) + def execute(provider: str, *, model: str) -> str: + calls.append(("python", provider, model)) + return model + + @anative_first(native=anative, route="chat_completions", errors=rules) + async def aexecute(provider: str, *, model: str) -> str: + calls.append(("python", provider, model)) + return model + + assert (await aexecute("first", model="ok") if asynchronous else execute("first", model="ok")) == "ok" + + async def fail() -> None: + if asynchronous: + await aexecute("second", model="limited") + else: + execute("second", model="limited") + + with pytest.raises(APIError) as caught: + await fail() + assert caught.value.llm_provider == "second" + assert caught.value.model == "limited" + assert calls == [("native", "first", "ok"), ("python", "first", "ok"), ("native", "second", "limited")] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("selection", ("native", "unavailable", "failed")) +@pytest.mark.parametrize("failure", ("none", "body", "cleanup", "cancel")) +async def test_context_selection_and_lifetime_are_separate(selection: str, failure: str) -> None: + from collections.abc import AsyncGenerator + from contextlib import AbstractAsyncContextManager, asynccontextmanager + + from litellm.rust_bridge.dispatch import anative_context + + events: Final[list[str]] = [] + error: Final = RuntimeError("connection use failed") + + @asynccontextmanager + async def connection(name: str) -> AsyncGenerator[str, None]: + events.append(f"{name}:enter") + try: + yield name + finally: + events.append(f"{name}:exit") + if failure == "cleanup": + raise error + + async def native() -> DispatchResult[AbstractAsyncContextManager[str]]: + events.append("attempt") + if selection == "failed": + raise RuntimeError("connect failed") + if selection == "unavailable": + return NativeSkipped(NativeSkipReason.UNAVAILABLE) + return Handled(connection("native")) + + @anative_context(native=native, route="websocket", errors=lambda: PYTHON_ON_ERROR) + def execute() -> AbstractAsyncContextManager[str]: + events.append("python") + return connection("python") + + async def run() -> None: + async with execute() as name: + assert name == ("native" if selection == "native" else "python") + if failure == "body": + raise error + if failure == "cancel": + raise asyncio.CancelledError + + if failure == "none": + await run() + elif failure == "cancel": + with pytest.raises(asyncio.CancelledError): + await run() + else: + with pytest.raises(RuntimeError) as caught: + await run() + assert caught.value is error + expected: Final = ( + ["attempt", "native:enter", "native:exit"] + if selection == "native" + else ["attempt", "python", "python:enter", "python:exit"] + ) + assert events == expected diff --git a/tests/test_litellm/rust_bridge/test_runtime.py b/tests/test_litellm/rust_bridge/test_runtime.py index b0fa510069b..a882e23fb58 100644 --- a/tests/test_litellm/rust_bridge/test_runtime.py +++ b/tests/test_litellm/rust_bridge/test_runtime.py @@ -1,95 +1,117 @@ from __future__ import annotations -from types import SimpleNamespace +from typing import Final import pytest -from litellm.exceptions import APIError -from litellm.rust_bridge import bindings, runtime - - -class RustBridgeDeclined(Exception): - pass - - -class RustUpstreamError(Exception): - pass - - -@pytest.fixture(autouse=True) -def native_exceptions(monkeypatch: pytest.MonkeyPatch) -> None: - native = SimpleNamespace( - RustBridgeDeclined=RustBridgeDeclined, - RustUpstreamError=RustUpstreamError, - ) - monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) - - -def context() -> runtime.BridgeErrorContext: - return runtime.BridgeErrorContext(route="messages", provider="anthropic", model="model") - - -def test_invoke_tags_native_decline_before_running_fallback() -> None: - calls: list[str] = [] - - def decline() -> object: - calls.append("rust") - raise RustBridgeDeclined("unsupported") - - value = runtime.invoke( - native_call=decline, - fallback=lambda: calls.append("python") or "fallback", - adapt=str, - mode=runtime.FallbackMode.PYTHON, - context=context(), - ) - - assert value == "fallback" - assert calls == ["rust", "python"] - - -def test_invoke_translates_upstream_without_fallback() -> None: - def fail() -> object: - raise RustUpstreamError(429, "rate limited") - - with pytest.raises(APIError, match="rate limited") as caught: - runtime.invoke( - native_call=fail, - fallback=lambda: pytest.fail("fallback must not run"), - adapt=str, - mode=runtime.FallbackMode.PYTHON, - context=context(), - ) - - assert caught.value.status_code == 429 +from litellm.rust_bridge import runtime @pytest.mark.asyncio -async def test_ainvoke_handles_native_success() -> None: - async def native() -> int: +@pytest.mark.parametrize("asynchronous", (False, True)) +@pytest.mark.parametrize("state", ("disabled", "ineligible", "unavailable", "handled")) +async def test_attempt_only_prepares_selected_requests(asynchronous: bool, state: str) -> None: + events: Final[list[str]] = [] + + def load() -> object | None: + events.append("load") + return None if state == "unavailable" else object() + + def prepare() -> int: + events.append("prepare") return 3 - async def fallback() -> str: - pytest.fail("fallback must not run") + def call(_binding: object, request: int) -> int: + events.append("call") + return request * 2 - assert ( - await runtime.ainvoke( - native_call=native, - fallback=fallback, - adapt=str, - mode=runtime.FallbackMode.PYTHON, - context=context(), + async def acall(binding: object, request: int) -> int: + return call(binding, request) + + def adapt(value: int) -> str: + events.append("adapt") + return str(value) + + result: Final = ( + await runtime.aattempt( + load=load, + enabled=state != "disabled", + eligible=state != "ineligible", + prepare=prepare, + call=acall, + adapt=adapt, + ) + if asynchronous + else runtime.attempt( + load=load, + enabled=state != "disabled", + eligible=state != "ineligible", + prepare=prepare, + call=call, + adapt=adapt, ) - == "3" ) + if state == "handled": + assert result == runtime.Handled("6") + assert events == ["load", "prepare", "call", "adapt"] + else: + assert result == runtime.NativeSkipped(runtime.NativeSkipReason(state)) + assert events == (["load"] if state == "unavailable" else []) -def test_required_mode_rejects_unavailable_bridge() -> None: - with pytest.raises(RuntimeError, match="is unavailable"): - runtime.invoke( - native_call=None, - fallback=lambda: pytest.fail("fallback must not run"), - adapt=str, - mode=runtime.FallbackMode.RUST_REQUIRED, - context=context(), - ) +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +@pytest.mark.parametrize("phase", ("prepare", "call")) +async def test_attempt_reports_failure_without_deciding_retry(asynchronous: bool, phase: str) -> None: + error: Final = RuntimeError("native failure") + + def prepare() -> int: + if phase == "prepare": + raise error + return 3 + + def call(_binding: object, request: int) -> int: + raise error + + async def acall(binding: object, request: int) -> int: + return call(binding, request) + + def adapt(value: int) -> str: + pytest.fail("failed attempts cannot be adapted") + + result: Final = ( + await runtime.aattempt(load=object, enabled=True, eligible=True, prepare=prepare, call=acall, adapt=adapt) + if asynchronous + else runtime.attempt(load=object, enabled=True, eligible=True, prepare=prepare, call=call, adapt=adapt) + ) + assert isinstance(result, runtime.NativeFailed) + assert result.error is error + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_adaptation_failure_remains_distinct_from_native_failure(asynchronous: bool) -> None: + error: Final = ValueError("invalid response") + + async def acall(_binding: object, request: int) -> int: + return request + + def adapt(value: int) -> str: + raise error + + async def run() -> None: + if asynchronous: + await runtime.aattempt(load=object, enabled=True, eligible=True, prepare=lambda: 3, call=acall, adapt=adapt) + else: + runtime.attempt( + load=object, + enabled=True, + eligible=True, + prepare=lambda: 3, + call=lambda binding, request: request, + adapt=adapt, + ) + + with pytest.raises(ValueError, match="invalid response") as caught: + await run() + assert caught.value is error diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index 112464bda22..0cbe0bf5277 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -4,6 +4,7 @@ import pytest import litellm from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch +from litellm.rust_bridge.runtime import Handled rust_bridge = importlib.import_module("litellm.rust_bridge.transcription") @@ -55,7 +56,8 @@ def test_enabled_sync_bridge_receives_audio() -> None: optional_params={"temperature": 0}, timeout=5.0, ) - assert result == {"text": "hello"} + assert isinstance(result, Handled) + assert result.value == {"text": "hello"} assert bridge.calls[0]["audio"] == {"data": "AQI=", "format": "wav", "filename": "audio.wav"} @@ -72,18 +74,19 @@ async def test_enabled_async_bridge() -> None: optional_params={}, timeout=None, ) - assert result == {"text": "async"} + assert result == Handled({"text": "async"}) def test_loader_returns_none_without_native_extension(monkeypatch: pytest.MonkeyPatch) -> None: rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) - monkeypatch.setattr("litellm.rust_bridge.get_native_bridge", lambda: None) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) assert rust_bridge.load_rust_transcription() is None assert rust_bridge.load_rust_atranscription() is None def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(rust_bridge, "transcription", lambda **_: None) + rust_bridge.configure_rust_transcription(transcription=None) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) with pytest.raises(RuntimeError, match="bridge is unavailable"): BedrockAudioTranscriptionRustDispatch().audio_transcriptions( @@ -100,10 +103,8 @@ def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> @pytest.mark.asyncio async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: - async def unavailable(**_: object) -> None: - return None - - monkeypatch.setattr(rust_bridge, "atranscription", unavailable) + rust_bridge.configure_rust_transcription(atranscription=None) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) with pytest.raises(RuntimeError, match="bridge is unavailable"): await BedrockAudioTranscriptionRustDispatch().async_audio_transcriptions( diff --git a/type-discipline-budget.json b/type-discipline-budget.json index e7186dfe186..78b28d2d775 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22180 + "limit": 22165 }, "LIT002": { "limit": 26729 @@ -15,7 +15,7 @@ "limit": 0 }, "LIT006": { - "limit": 1035 + "limit": 1022 }, "LIT007": { "limit": 0 @@ -27,10 +27,10 @@ "limit": 0 }, "LIT010": { - "limit": 16426 + "limit": 16419 }, "LIT011": { - "limit": 5506 + "limit": 5497 }, "LIT012": { "limit": 4486