diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index be81357e3d9..ff5e881761b 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -3,7 +3,7 @@ "limit": 13426 }, "reportArgumentType": { - "limit": 2194 + "limit": 2192 }, "reportAssignmentType": { "limit": 319 @@ -30,7 +30,7 @@ "limit": 7 }, "reportGeneralTypeIssues": { - "limit": 101 + "limit": 100 }, "reportIncompatibleMethodOverride": { "limit": 56 @@ -57,7 +57,7 @@ "limit": 5570 }, "reportMissingTypeArgument": { - "limit": 15281 + "limit": 15279 }, "reportMissingTypeStubs": { "limit": 40 @@ -99,31 +99,31 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 44025 + "limit": 44019 }, "reportUnknownLambdaType": { "limit": 109 }, "reportUnknownMemberType": { - "limit": 38269 + "limit": 38266 }, "reportUnknownParameterType": { - "limit": 19584 + "limit": 19583 }, "reportUnknownVariableType": { - "limit": 29813 + "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 318bb043270..dc3a4d179b7 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -27,7 +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 adispatch, dispatch +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, @@ -369,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, @@ -385,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 @@ -407,114 +417,67 @@ 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, + 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, + ) + + def native_completion() -> DispatchResult[ModelResponse]: + return rust_chat_completions_bridge.chat_completions( + model=model, messages=messages, + optional_params=rust_optional_params, + model_response=model_response, api_key=api_key, - additional_args=rust_logging_args, + 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, ) - if acompletion is True: - 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 adispatch( - native=lambda: 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=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, - ), - python=python_fallback, - route="chat_completions", - errors=rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model), - ) - rust_response: Final = dispatch( - native=lambda: rust_chat_completions_bridge.chat_completions( - 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=lambda: None, - route="chat_completions", - errors=rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model), - ) - 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, @@ -536,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, @@ -558,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 @@ -590,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), @@ -617,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 a1452602523..6d6be069f99 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -5,7 +5,8 @@ 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, adispatch, dispatch +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 @@ -40,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, *, @@ -52,23 +84,39 @@ class BedrockAudioTranscriptionRustDispatch: optional_params: dict[str, object], timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: - rust_response: Final = dispatch( - native=lambda: 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, - ), - python=_unavailable, - route="audio transcription", - errors=PROPAGATE, - ) - return TranscriptionResponse(**rust_response) + _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, + 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)) + + @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, *, @@ -81,19 +129,4 @@ class BedrockAudioTranscriptionRustDispatch: optional_params: dict[str, object], timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: - rust_response: Final = await adispatch( - native=lambda: 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, - ), - python=_aunavailable, - route="audio transcription", - errors=PROPAGATE, - ) - 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 e3ee89a2455..2918ea96b8c 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -18,7 +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 adispatch, dispatch +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 @@ -407,83 +408,62 @@ 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 adispatch( - native=lambda: 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, - ), - python=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, - ), - route="chat_completions", - errors=rust_chat_completions_bridge.error_handling("bedrock", model), - ) - rust_response: Final = dispatch( - native=lambda: rust_chat_completions_bridge.chat_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, - ), - python=lambda: None, - route="chat_completions", - errors=rust_chat_completions_bridge.error_handling("bedrock", model), - ) - if rust_response is not None: - return rust_response + log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( + logging_obj=logging_obj, + messages=messages, + api_key="", + additional_args=rust_logging_args, + ) - ### ROUTING (ASYNC, STREAMING, SYNC) - if acompletion: - if isinstance(client, HTTPHandler): - client = None + def native_completion() -> DispatchResult[ModelResponse]: + return rust_chat_completions_bridge.chat_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, + ) + + 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, @@ -496,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, @@ -504,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, @@ -517,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 e361581ec2a..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"} - from litellm.rust_bridge.dispatch import PYTHON_ON_ERROR, adispatch, async_none - - rust_response: Final = await adispatch( - native=lambda: 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, - ), - python=async_none, - route="messages", - errors=PYTHON_ON_ERROR, + 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, ) - if rust_response is None: - return None - 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,26 +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 - from litellm.rust_bridge.dispatch import PYTHON_ON_ERROR, adispatch, async_none - - rust_backend: Final = await adispatch( - native=lambda: rust_responses_websocket.connect( - url=ws_url, - headers={str(key): str(value) for key, value in headers.items()}, - timeout=timeout, - ), - python=async_none, - route="responses_websocket", - errors=PYTHON_ON_ERROR, - ) - 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 03ed110ba34..7a3209b9eb4 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -7,7 +7,7 @@ import base64 import mimetypes import os import re -from collections.abc import Coroutine, Mapping +from collections.abc import Callable, Coroutine, Mapping from io import IOBase from typing import Any, Final, cast @@ -25,7 +25,8 @@ from litellm.llms.base_llm.ocr.transformation import ( ) from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge import ocr as rust_ocr_bridge -from litellm.rust_bridge.dispatch import PROPAGATE, adispatch, dispatch +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 @@ -155,6 +156,69 @@ def _prepare_ocr_request( ) +@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], +) -> 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, + ) + 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 _attempt_ocr( + prepared_request: rust_ocr_bridge.PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], + 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, + 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=is_async, + headers=prepared_request.extra_headers, + provider_config=prepared_request.provider_config, + litellm_params=prepared_request.litellm_params, + ) + + @client async def aocr( model: str, @@ -251,32 +315,7 @@ async def aocr( from litellm.secret_managers.main import get_secret_str - async def python_fallback() -> OCRResponse: - pending: 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=True, - headers=prepared.extra_headers, - provider_config=prepared.provider_config, - litellm_params=prepared.litellm_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 - - return await adispatch( - native=lambda: rust_ocr_bridge.aattempt_ocr(prepared_request=prepared, resolve_api_key=get_secret_str), - python=python_fallback, - route="ocr", - errors=PROPAGATE, - ) + return await _execute_aocr(prepared_request=prepared, resolve_api_key=get_secret_str) except Exception as e: raise litellm.exception_type( model=model, @@ -517,28 +556,7 @@ def ocr( from litellm.secret_managers.main import get_secret_str - def python_fallback() -> OCRResponse | Coroutine[object, object, OCRResponse]: - return 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 dispatch( - native=lambda: rust_ocr_bridge.attempt_ocr(prepared_request=prepared, resolve_api_key=get_secret_str), - python=python_fallback, - route="ocr", - errors=PROPAGATE, - ) + 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/chat_completions.py b/litellm/rust_bridge/chat_completions.py index 24e97993a1f..d9f06098903 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -225,6 +225,7 @@ def chat_completions( extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, on_response: ResponseObserver, + eligible: bool = True, ) -> DispatchResult[ModelResponse]: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) @@ -245,7 +246,7 @@ def chat_completions( return attempt( load=_CHAT.load, enabled=rust_enabled(), - eligible=True, + eligible=eligible, prepare=lambda: timeout_to_seconds(timeout), call=call, adapt=adapt, @@ -264,6 +265,7 @@ async def achat_completions( extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, on_response: ResponseObserver, + eligible: bool = True, ) -> DispatchResult[ModelResponse]: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) @@ -284,7 +286,7 @@ async def achat_completions( return await aattempt( load=_ACHAT.load, enabled=rust_enabled(), - eligible=True, + eligible=eligible, prepare=lambda: timeout_to_seconds(timeout), call=call, adapt=adapt, diff --git a/litellm/rust_bridge/dispatch.py b/litellm/rust_bridge/dispatch.py index 7572e31a8d2..8c275db3533 100644 --- a/litellm/rust_bridge/dispatch.py +++ b/litellm/rust_bridge/dispatch.py @@ -1,9 +1,11 @@ from __future__ import annotations -from collections.abc import Awaitable, Callable +from collections.abc import AsyncGenerator, Awaitable, Callable +from contextlib import AbstractAsyncContextManager, asynccontextmanager from dataclasses import dataclass from enum import Enum -from typing import Final, TypeAlias, TypeVar +from functools import wraps +from typing import Final, ParamSpec, TypeAlias, TypeVar from litellm._logging import verbose_logger from litellm.exceptions import APIError @@ -12,6 +14,7 @@ from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, N NativeT = TypeVar("NativeT") PythonT = TypeVar("PythonT") +P = ParamSpec("P") class ErrorAction(Enum): @@ -89,49 +92,95 @@ def _log_skip(route: str, skipped: NativeSkipped) -> None: verbose_logger.debug("Native %s skipped (%s): %s", route, skipped.reason.value, skipped.detail or "") -def dispatch( +def native_first( *, - native: Callable[[], DispatchResult[NativeT]], - python: Callable[[], PythonT], + native: Callable[P, DispatchResult[NativeT]], route: str, - errors: ErrorHandling, -) -> NativeT | PythonT: - try: - attempted: Final = native() - except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures - unexpected: Final = _handle_error(error, errors.unexpected, route, NativeSkipReason.FAILED) - _log_skip(route, unexpected) - return python() - result: Final = _resolve(attempted, errors, route) - match result: - case Handled(value): - return value - case NativeSkipped(): - _log_skip(route, result) - return python() + 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 -async def adispatch( +def anative_first( *, - native: Callable[[], Awaitable[DispatchResult[NativeT]]], - python: Callable[[], Awaitable[PythonT]], + native: Callable[P, Awaitable[DispatchResult[NativeT]]], route: str, - errors: ErrorHandling, -) -> NativeT | PythonT: - try: - attempted: Final = await native() - except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures - unexpected: Final = _handle_error(error, errors.unexpected, route, NativeSkipReason.FAILED) - _log_skip(route, unexpected) - return await python() - result: Final = _resolve(attempted, errors, route) - match result: - case Handled(value): - return value - case NativeSkipped(): - _log_skip(route, result) - return await python() + 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 -async def async_none() -> None: - return None +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/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index 8b95729bfa9..d673ab8431c 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -2,6 +2,8 @@ from __future__ import annotations +from collections.abc import AsyncGenerator +from contextlib import AbstractAsyncContextManager, asynccontextmanager from typing import Final import httpx @@ -13,7 +15,7 @@ from litellm.rust_bridge.protocols import ( RustResponsesWebSocket, RustResponsesWebSocketConnection, ) -from litellm.rust_bridge.runtime import DispatchResult, aattempt +from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result from litellm.rust_bridge.timeouts import timeout_to_seconds _RESPONSES_WEBSOCKET: Final[NativeBinding[RustResponsesWebSocketConnection]] = NativeBinding( @@ -32,7 +34,7 @@ def set_rust_responses_websocket( _RESPONSES_WEBSOCKET.override(connection) -class _ConnectionAdapter: +class ConnectionAdapter: def __init__(self, connection: RustResponsesWebSocket): self._connection: Final[RustResponsesWebSocket] = connection @@ -54,7 +56,7 @@ async def connect( url: str, headers: dict[str, str], timeout: float | httpx.Timeout | None, -) -> DispatchResult[_ConnectionAdapter]: +) -> DispatchResult[ConnectionAdapter]: return await aattempt( load=_RESPONSES_WEBSOCKET.load, enabled=rust_enabled(), @@ -65,5 +67,23 @@ async def connect( headers=headers, timeout_seconds=timeout_seconds, ), - adapt=_ConnectionAdapter, + 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 48f1ecd23b4..b908a32e879 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -87,3 +87,9 @@ async def aattempt( def identity(value: ResultT) -> ResultT: return value + + +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/ruff-strict-budget.json b/ruff-strict-budget.json index fd3bc01321c..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,7 +150,7 @@ "limit": 253 }, "PLW0127": { - "limit": 57 + "limit": 55 }, "PLW0602": { "limit": 215 @@ -195,7 +195,7 @@ "limit": 22 }, "SIM101": { - "limit": 56 + "limit": 55 }, "SIM102": { "limit": 310 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 1502a47a7d9..bb7d68864ce 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -9,7 +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, NativeSkipped, NativeSkipReason +from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, ) @@ -127,6 +127,16 @@ def test_load_rust_messages_returns_injected_impl(): assert rust_messages.load_rust_messages() is bridge +def test_bare_rust_still_toggles_ocr(): + from litellm.rust_bridge.ocr import rust_ocr_enabled + + litellm.rust(True) + assert rust_ocr_enabled() is True + + litellm.rust(False) + assert rust_ocr_enabled() is False + + def test_load_rust_amessages_returns_injected_impl(): bridge = RecordingAsyncMessages() litellm.rust(True) @@ -215,7 +225,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 @@ -226,7 +236,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] @@ -239,13 +250,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 +267,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,10 +279,24 @@ 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" +@pytest.mark.asyncio +async def test_gate_ignores_request_flag_when_process_enabled(): + bridge = RecordingAsyncMessages() + litellm.rust(True) + rust_messages.set_rust_messages(amessages=bridge) + + response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False)) + + assert isinstance(response, Handled) + response = response.value + assert len(bridge.calls) == 1 + + @pytest.mark.asyncio async def test_gate_invokes_rust_for_native_anthropic_provider(): bridge = RecordingAsyncMessages() @@ -286,7 +311,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 +329,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 +345,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 +357,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 +369,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 +385,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 @@ -391,4 +419,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/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/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index d00ef5c8127..e502689797b 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -4,7 +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, NativeSkipped, NativeSkipReason, NativeFailed +from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason class _FakeNativeConnection: @@ -58,7 +58,7 @@ 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() @@ -115,3 +115,30 @@ async def test_connection_failure_is_reported_to_orchestration() -> None: 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_dispatch.py b/tests/test_litellm/rust_bridge/test_dispatch.py index 3641729e521..9258372fb93 100644 --- a/tests/test_litellm/rust_bridge/test_dispatch.py +++ b/tests/test_litellm/rust_bridge/test_dispatch.py @@ -10,7 +10,7 @@ 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, adispatch, dispatch +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 @@ -53,9 +53,9 @@ async def test_shared_dispatch_calls_python_once_and_logs_skip( return python() result: Final = ( - await adispatch(native=anative, python=apython, route="test", errors=PROPAGATE) + await anative_first(native=anative, route="test", errors=lambda: PROPAGATE)(apython)() if asynchronous - else dispatch(native=native, python=python, route="test", errors=PROPAGATE) + else native_first(native=native, route="test", errors=lambda: PROPAGATE)(python)() ) assert result == "python response" assert calls == ["native", "python"] @@ -65,6 +65,7 @@ async def test_shared_dispatch_calls_python_once_and_logs_skip( @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) @@ -75,9 +76,9 @@ async def test_native_success_does_not_run_python_even_when_value_is_none(asynch return python() result: Final = ( - await adispatch(native=native, python=apython, route="test", errors=PYTHON_ON_ERROR) + await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(apython)() if asynchronous - else dispatch(native=lambda: Handled(None), python=python, route="test", errors=PYTHON_ON_ERROR) + else native_first(native=lambda: Handled(None), route="test", errors=lambda: PYTHON_ON_ERROR)(python)() ) assert result is None @@ -124,8 +125,8 @@ async def test_declarations_preserve_endpoint_error_behavior( async def run() -> str: if asynchronous: - return await adispatch(native=anative, python=apython, route="chat_completions", errors=rules) - return dispatch(native=native, python=python, route="chat_completions", errors=rules) + 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" @@ -163,13 +164,10 @@ async def test_python_failure_is_never_reclassified_as_native_failure(asynchrono async def run() -> str: if asynchronous: - return await adispatch(native=native, python=apython, route="test", errors=PYTHON_ON_ERROR) - return dispatch( - native=lambda: NativeSkipped(NativeSkipReason.UNAVAILABLE), - python=python, - route="test", - errors=PYTHON_ON_ERROR, - ) + 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() @@ -179,6 +177,7 @@ async def test_python_failure_is_never_reclassified_as_native_failure(asynchrono @pytest.mark.asyncio async def test_cancellation_does_not_run_python() -> None: + async def native() -> DispatchResult[str]: raise asyncio.CancelledError @@ -186,20 +185,123 @@ async def test_cancellation_does_not_run_python() -> None: pytest.fail("cancellation must not dispatch Python") with pytest.raises(asyncio.CancelledError): - await adispatch(native=native, python=python, route="test", errors=PYTHON_ON_ERROR) + 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: - dispatch( + native_first( native=lambda: NativeFailed(error), - python=lambda: pytest.fail("upstream errors must not run Python"), route="chat_completions", - errors=error_handling("anthropic", "model"), - ) + 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/type-discipline-budget.json b/type-discipline-budget.json index 144732cfd47..78b28d2d775 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -30,7 +30,7 @@ "limit": 16419 }, "LIT011": { - "limit": 5506 + "limit": 5497 }, "LIT012": { "limit": 4486