From 0e1c72d8fed1b3cba7ff9abb22c085b068aa2181 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 23:04:57 -0700 Subject: [PATCH] feat(native): preserve request identity and capabilities --- .../core/src/chat_completions/prepare.rs | 2 +- litellm/llms/anthropic/chat/handler.py | 9 +++- .../bedrock/audio_transcription/__init__.py | 34 ++++++++++++- litellm/llms/bedrock/chat/converse_handler.py | 9 +++- litellm/llms/custom_httpx/llm_http_handler.py | 14 ++++++ litellm/main.py | 2 + litellm/rust_bridge/chat_completions.py | 21 +++++--- litellm/rust_bridge/messages.py | 17 ++++--- litellm/rust_bridge/ocr.py | 49 ++++++++----------- litellm/rust_bridge/request.py | 44 +++++++++++++++-- litellm/rust_bridge/responses_websocket.py | 12 +++-- litellm/rust_bridge/transcription.py | 21 +++++--- .../test_rust_bridge_messages.py | 23 --------- .../chat/test_anthropic_chat_handler.py | 21 +++++--- .../chat/test_bedrock_converse_handler.py | 25 +++++++--- .../rust_bridge/test_request_context.py | 40 +++++++++++++++ 16 files changed, 246 insertions(+), 97 deletions(-) create mode 100644 tests/test_litellm/rust_bridge/test_request_context.py diff --git a/litellm-rust/crates/core/src/chat_completions/prepare.rs b/litellm-rust/crates/core/src/chat_completions/prepare.rs index d22a66ed6c0..92012e9cf4a 100644 --- a/litellm-rust/crates/core/src/chat_completions/prepare.rs +++ b/litellm-rust/crates/core/src/chat_completions/prepare.rs @@ -46,7 +46,7 @@ pub(super) fn resolve_request( ) -> Result { let (model, config) = resolve_provider_config(request.model, options.custom_llm_provider.as_deref()) - .map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?; + .map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?; let messages = parse_messages(request.messages).map_err(|_| Error::Declined("unreadable message list"))?; if messages.is_empty() { diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index a0c84615224..4bb9d0ba338 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -28,7 +28,7 @@ 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, provider_errors -from litellm.rust_bridge.request import anthropic_options +from litellm.rust_bridge.request import anthropic_options, request_context from litellm.rust_bridge.runtime import DispatchResult from litellm.types.llms.anthropic import ( ContentBlockDelta, @@ -435,6 +435,11 @@ class AnthropicChatCompletion(BaseLLM): api_key=api_key, additional_args=rust_logging_args, ) + rust_context: Final = request_context( + logging_obj=logging_obj, + request_model=logging_obj.model, + litellm_params=litellm_params, + ) def native_completion() -> DispatchResult[ModelResponse]: return rust_chat_completions_bridge.chat_completions( @@ -452,6 +457,7 @@ class AnthropicChatCompletion(BaseLLM): stream=bool(stream), has_custom_client=client is not None, eligible=serves_via_rust, + context=rust_context, ) async def native_acompletion() -> DispatchResult[ModelResponse]: @@ -470,6 +476,7 @@ class AnthropicChatCompletion(BaseLLM): stream=bool(stream), has_custom_client=client is not None, eligible=serves_via_rust, + context=rust_context, ) @anative_first( diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py index 6fc1a82d2c6..87847953a8f 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -1,11 +1,14 @@ import base64 +from io import IOBase from typing import Final, NoReturn import httpx from litellm.litellm_core_utils.audio_utils.utils import process_audio_file +from litellm.litellm_core_utils.litellm_logging import Logging from litellm.rust_bridge import transcription as rust_transcription_bridge from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors +from litellm.rust_bridge.request import request_context from litellm.rust_bridge.runtime import DispatchResult, adapt_result from litellm.types.utils import FileTypes, TranscriptionResponse @@ -19,6 +22,17 @@ async def _aunavailable() -> NoReturn: class BedrockAudioTranscriptionRustDispatch: + @staticmethod + def _input_source_kind(audio_file: FileTypes) -> str: + content: Final = audio_file[1] if isinstance(audio_file, tuple) else audio_file + if isinstance(content, (bytes, bytearray, memoryview)): + return "bytes" + if isinstance(content, IOBase): + return "file" + if isinstance(content, str): + return "path" + return "opaque" + @staticmethod def _audio_payload(audio_file: FileTypes) -> dict[str, object]: processed_audio: Final = process_audio_file(audio_file) @@ -52,6 +66,7 @@ class BedrockAudioTranscriptionRustDispatch: extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, + logging_obj: Logging | None = None, ) -> DispatchResult[TranscriptionResponse]: result: Final = rust_transcription_bridge.transcription( model=model, @@ -62,13 +77,19 @@ class BedrockAudioTranscriptionRustDispatch: extra_headers=extra_headers, optional_params=optional_params, timeout=timeout, + input_source_kind=self._input_source_kind(audio_file), + context=request_context( + logging_obj=logging_obj, + request_model=logging_obj.model if logging_obj is not None else model, + litellm_params=logging_obj.litellm_params if logging_obj is not None else None, + ), ) 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: ( + errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout, logging_obj=None: ( provider_errors(custom_llm_provider, model) ), ) @@ -83,6 +104,7 @@ class BedrockAudioTranscriptionRustDispatch: extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, + logging_obj: Logging | None = None, ) -> TranscriptionResponse: _unavailable() @@ -97,6 +119,7 @@ class BedrockAudioTranscriptionRustDispatch: extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, + logging_obj: Logging | None = None, ) -> DispatchResult[TranscriptionResponse]: result: Final = await rust_transcription_bridge.atranscription( model=model, @@ -107,13 +130,19 @@ class BedrockAudioTranscriptionRustDispatch: extra_headers=extra_headers, optional_params=optional_params, timeout=timeout, + input_source_kind=self._input_source_kind(audio_file), + context=request_context( + logging_obj=logging_obj, + request_model=logging_obj.model if logging_obj is not None else model, + litellm_params=logging_obj.litellm_params if logging_obj is not None else None, + ), ) 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: ( + errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout, logging_obj=None: ( provider_errors(custom_llm_provider, model) ), ) @@ -128,5 +157,6 @@ class BedrockAudioTranscriptionRustDispatch: extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, + logging_obj: Logging | None = None, ) -> TranscriptionResponse: await _aunavailable() diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 982af97a6c4..2d2db6c17a5 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -19,7 +19,7 @@ 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, provider_errors -from litellm.rust_bridge.request import bedrock_options +from litellm.rust_bridge.request import bedrock_options, request_context from litellm.rust_bridge.runtime import DispatchResult from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper @@ -425,6 +425,11 @@ class BedrockConverseLLM(BaseAWSLLM): api_key="", additional_args=rust_logging_args, ) + rust_context: Final = request_context( + logging_obj=logging_obj, + request_model=logging_obj.model, + litellm_params=litellm_params, + ) def native_completion() -> DispatchResult[ModelResponse]: return rust_chat_completions_bridge.chat_completions( @@ -442,6 +447,7 @@ class BedrockConverseLLM(BaseAWSLLM): stream=bool(stream), has_custom_client=client is not None, eligible=serves_via_rust, + context=rust_context, ) async def native_acompletion() -> DispatchResult[ModelResponse]: @@ -460,6 +466,7 @@ class BedrockConverseLLM(BaseAWSLLM): stream=bool(stream), has_custom_client=client is not None, eligible=serves_via_rust, + context=rust_context, ) @anative_first( diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 5aa580ae47a..4e7cc94ae8c 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2244,6 +2244,7 @@ class BaseLLMHTTPHandler: stream=stream or False, custom_llm_provider=custom_llm_provider, ), + logging_obj=logging_obj, ) return adapt_result(result, self._rust_anthropic_messages_fake_stream) if stream else result @@ -2398,6 +2399,7 @@ class BaseLLMHTTPHandler: headers: dict, request_body: dict, timeout: float | httpx.Timeout | None, + logging_obj: LiteLLMLoggingObj | None = None, ) -> DispatchResult[AnthropicMessagesResponse]: if custom_llm_provider not in ("azure_ai", "anthropic"): return NativeSkipped(NativeSkipReason.INELIGIBLE) @@ -2409,6 +2411,7 @@ class BaseLLMHTTPHandler: return NativeSkipped(NativeSkipReason.INELIGIBLE) from litellm.rust_bridge import messages as rust_messages_bridge + from litellm.rust_bridge.request import request_context upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"} result: Final = await rust_messages_bridge.amessages( @@ -2422,6 +2425,11 @@ class BaseLLMHTTPHandler: stream=stream, has_custom_client=has_custom_client, has_agentic_hook=has_agentic_hook, + context=request_context( + logging_obj=logging_obj, + request_model=logging_obj.model if logging_obj is not None else model, + litellm_params=litellm_params.model_dump(), + ), ) def adapt(rust_response: dict[str, object]) -> AnthropicMessagesResponse: @@ -6509,6 +6517,7 @@ class BaseLLMHTTPHandler: ) from litellm.rust_bridge import responses_websocket as rust_responses_websocket + from litellm.rust_bridge.request import request_context async def attempt_connection() -> DispatchResult[ AbstractAsyncContextManager[rust_responses_websocket.ConnectionAdapter] @@ -6519,6 +6528,11 @@ class BaseLLMHTTPHandler: url=ws_url, headers={str(key): str(value) for key, value in headers.items()}, timeout=timeout, + context=request_context( + logging_obj=logging_obj, + request_model=logging_obj.model, + litellm_params=litellm_params.model_dump(), + ), ) @anative_context( diff --git a/litellm/main.py b/litellm/main.py index d4da18e8f6f..eca4bb50649 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -7895,6 +7895,7 @@ def transcription( extra_headers=extra_headers, optional_params=optional_params, timeout=timeout, + logging_obj=litellm_logging_obj, ) else: response = dispatch.audio_transcriptions( @@ -7906,6 +7907,7 @@ def transcription( extra_headers=extra_headers, optional_params=optional_params, timeout=timeout, + logging_obj=litellm_logging_obj, ) elif provider_config is not None: response = base_llm_http_handler.audio_transcriptions( diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index 6af8471e2b3..553b8b60e6a 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -36,6 +36,7 @@ from litellm.rust_bridge.request import ( NativeRequestOptions, PreparedNativeCall, call_native, + with_capabilities, ) from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -239,12 +240,15 @@ def chat_completions( stream: bool = False, has_custom_client: bool = False, eligible: bool = True, + context: NativeRequestContext | None = None, ) -> 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, prepared: PreparedNativeCall[NativeChatCompletionsRequest]) -> Mapping[str, object]: + def call( + native: RustChatCompletions, prepared: PreparedNativeCall[NativeChatCompletionsRequest] + ) -> Mapping[str, object]: return call_native(native, prepared) return attempt( @@ -262,12 +266,13 @@ def chat_completions( bedrock=bedrock, anthropic=anthropic, ), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( execution_mode="sync", stream=stream, has_custom_client=has_custom_client, - ) + ), ), ), call=call, @@ -292,6 +297,7 @@ async def achat_completions( stream: bool = False, has_custom_client: bool = False, eligible: bool = True, + context: NativeRequestContext | None = None, ) -> DispatchResult[ModelResponse]: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) @@ -318,12 +324,13 @@ async def achat_completions( bedrock=bedrock, anthropic=anthropic, ), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( execution_mode="async", stream=stream, has_custom_client=has_custom_client, - ) + ), ), ), call=call, diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index 445009e4cfb..97ea14f2aa8 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -15,6 +15,7 @@ from litellm.rust_bridge.request import ( NativeRequestOptions, PreparedNativeCall, call_native, + with_capabilities, ) from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -60,6 +61,7 @@ def messages( stream: bool = False, has_custom_client: bool = False, has_agentic_hook: bool = False, + context: NativeRequestContext | None = None, ) -> DispatchResult[dict[str, object]]: return attempt( load=_MESSAGES.load, @@ -74,13 +76,14 @@ def messages( extra_headers=extra_headers, timeout_seconds=timeout_to_seconds(timeout), ), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( execution_mode="sync", stream=stream, has_custom_client=has_custom_client, has_agentic_hook=has_agentic_hook, - ) + ), ), ), call=call_native, @@ -100,6 +103,7 @@ async def amessages( stream: bool = False, has_custom_client: bool = False, has_agentic_hook: bool = False, + context: NativeRequestContext | None = None, ) -> DispatchResult[dict[str, object]]: return await aattempt( load=_AMESSAGES.load, @@ -114,13 +118,14 @@ async def amessages( extra_headers=extra_headers, timeout_seconds=timeout_to_seconds(timeout), ), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( execution_mode="async", stream=stream, has_custom_client=has_custom_client, has_agentic_hook=has_agentic_hook, - ) + ), ), ), call=call_native, diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index e2236317853..516a79e385b 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -12,23 +12,20 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging 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 UNCHANGED, NativeBinding, Unchanged +from litellm.rust_bridge.bindings import NativeBinding from litellm.rust_bridge.protocols import RustAocr, RustOcr from litellm.rust_bridge.request import ( NativeOCRRequest, NativeRequestCapabilities, - NativeRequestContext, NativeRequestOptions, PreparedNativeCall, call_native, + request_context, vertex_options, ) from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt from litellm.rust_bridge.timeouts import timeout_to_seconds -rust: Final = _configuration.rust -rust_ocr_enabled: Final = _configuration.rust_ocr_enabled - _OCR: Final[NativeBinding[RustOcr]] = NativeBinding(lambda native: native.ocr) _AOCR: Final[NativeBinding[RustAocr]] = NativeBinding(lambda native: native.aocr) _HEADERS: Final = TypeAdapter(dict[str, object]) @@ -66,23 +63,6 @@ _RUST_OCR_PROVIDERS: Final = frozenset( ) -def set_rust_ocr( - *, - ocr: RustOcr | None | Unchanged = UNCHANGED, - aocr: RustAocr | None | Unchanged = UNCHANGED, -) -> None: - if not isinstance(ocr, Unchanged): - if ocr is None: - _OCR.reset() - else: - _OCR.override(ocr) - if not isinstance(aocr, Unchanged): - if aocr is None: - _AOCR.reset() - else: - _AOCR.override(aocr) - - def load_rust_ocr() -> RustOcr | None: return _OCR.load() @@ -109,6 +89,11 @@ def _ocr_input_source_kind(document: dict[str, object]) -> str: return "inline" +def _ocr_request_format(optional_params: dict[str, object]) -> str | None: + value = optional_params.get(OCR_REQUEST_FORMAT_PARAM) + return value if isinstance(value, str) else None + + def _rust_bridge_optional_params( prepared_request: PreparedOCRRequest, resolve_secret: Callable[[str], str | None], @@ -204,7 +189,7 @@ def attempt_ocr( ) -> DispatchResult[OCRResponse]: return attempt( load=_OCR.load, - enabled=rust_ocr_enabled(), + enabled=_configuration.rust_enabled(), prepare=lambda: _prepare_rust_ocr_call( prepared_request=prepared_request, resolve_api_key=resolve_api_key, @@ -225,14 +210,18 @@ def attempt_ocr( timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), vertex=vertex_options(prepared.optional_params), ), - context=NativeRequestContext( + context=request_context( + logging_obj=prepared_request.litellm_logging_obj, + request_model=prepared_request.model, + litellm_params=prepared_request.litellm_params, capabilities=NativeRequestCapabilities( execution_mode="sync", input_source_kind=_ocr_input_source_kind(prepared_request.document), + request_format=_ocr_request_format(prepared_request.optional_params), native_response_format=( prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" ), - ) + ), ), ), ), @@ -247,7 +236,7 @@ async def aattempt_ocr( ) -> DispatchResult[OCRResponse]: return await aattempt( load=_AOCR.load, - enabled=rust_ocr_enabled(), + enabled=_configuration.rust_enabled(), prepare=lambda: _prepare_rust_ocr_call( prepared_request=prepared_request, resolve_api_key=resolve_api_key, @@ -268,14 +257,18 @@ async def aattempt_ocr( timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), vertex=vertex_options(prepared.optional_params), ), - context=NativeRequestContext( + context=request_context( + logging_obj=prepared_request.litellm_logging_obj, + request_model=prepared_request.model, + litellm_params=prepared_request.litellm_params, capabilities=NativeRequestCapabilities( execution_mode="async", input_source_kind=_ocr_input_source_kind(prepared_request.document), + request_format=_ocr_request_format(prepared_request.optional_params), native_response_format=( prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" ), - ) + ), ), ), ), diff --git a/litellm/rust_bridge/request.py b/litellm/rust_bridge/request.py index e81d10f587f..562ab1f4950 100644 --- a/litellm/rust_bridge/request.py +++ b/litellm/rust_bridge/request.py @@ -1,7 +1,8 @@ from __future__ import annotations from collections.abc import Mapping, Sequence -from dataclasses import dataclass +from dataclasses import dataclass, replace +from types import MappingProxyType from typing import Generic, Protocol, TypeVar @@ -110,6 +111,43 @@ class NativeRequestContext: capabilities: NativeRequestCapabilities = NativeRequestCapabilities() +def request_context( + *, + logging_obj: object | None, + request_model: str, + litellm_params: Mapping[str, object] | None = None, + capabilities: NativeRequestCapabilities | None = None, +) -> NativeRequestContext: + params = litellm_params if litellm_params is not None else MappingProxyType({}) + metadata_value = params.get("metadata") or params.get("litellm_metadata") + metadata = metadata_value if isinstance(metadata_value, Mapping) else MappingProxyType({}) + + def string(name: str) -> str | None: + value = params.get(name, metadata.get(name)) + return value if isinstance(value, str) else None + + call_id = getattr(logging_obj, "litellm_call_id", None) + trace_id = getattr(logging_obj, "litellm_trace_id", None) + return NativeRequestContext( + litellm_call_id=call_id if isinstance(call_id, str) else None, + trace_id=trace_id if isinstance(trace_id, str) else None, + request_model=request_model, + attribution=RequestAttribution( + user_api_key_hash=string("user_api_key_hash"), + user_api_key_user_id=string("user_api_key_user_id"), + user_api_key_team_id=string("user_api_key_team_id"), + ), + capabilities=capabilities or NativeRequestCapabilities(), + ) + + +def with_capabilities( + context: NativeRequestContext, + capabilities: NativeRequestCapabilities, +) -> NativeRequestContext: + return replace(context, capabilities=capabilities) + + RequestT = TypeVar("RequestT") RequestContraT = TypeVar("RequestContraT", contravariant=True) ResultT = TypeVar("ResultT", covariant=True) @@ -152,14 +190,14 @@ class NativeMessagesRequest: @dataclass(frozen=True, slots=True) class NativeOCRRequest: model: str - document: dict[str, object] + document: object optional_params: dict[str, object] @dataclass(frozen=True, slots=True) class NativeTranscriptionRequest: model: str - audio: dict[str, object] + audio: object optional_params: dict[str, object] diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index 2dc3e18bd2c..a1ffe107cda 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -22,6 +22,7 @@ from litellm.rust_bridge.request import ( NativeResponsesWebSocketRequest, PreparedNativeCall, call_native, + with_capabilities, ) from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -66,6 +67,7 @@ async def connect( timeout: float | httpx.Timeout | None, websocket_mode: str = "native", requires_connection: bool = True, + context: NativeRequestContext | None = None, ) -> DispatchResult[ConnectionAdapter]: return await aattempt( load=_RESPONSES_WEBSOCKET.load, @@ -74,11 +76,13 @@ async def connect( prepare=lambda: PreparedNativeCall( request=NativeResponsesWebSocketRequest(url=url), options=NativeRequestOptions(extra_headers=headers, timeout_seconds=timeout_to_seconds(timeout)), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( + execution_mode="async", websocket_mode=websocket_mode, requires_connection=requires_connection, - ) + ), ), ), call=lambda connection_type, prepared: call_native(connection_type.connect, prepared), @@ -101,6 +105,7 @@ async def managed_connect( timeout: float | httpx.Timeout | None, websocket_mode: str = "managed", requires_connection: bool = True, + context: NativeRequestContext | None = None, ) -> DispatchResult[AbstractAsyncContextManager[ConnectionAdapter]]: result: Final = await connect( url=url, @@ -108,5 +113,6 @@ async def managed_connect( timeout=timeout, websocket_mode=websocket_mode, requires_connection=requires_connection, + context=context, ) return adapt_result(result, _connection_context) diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py index 97f4eb485d1..8d1997e565c 100644 --- a/litellm/rust_bridge/transcription.py +++ b/litellm/rust_bridge/transcription.py @@ -14,6 +14,7 @@ from litellm.rust_bridge.request import ( PreparedNativeCall, bedrock_options, call_native, + with_capabilities, ) from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -50,7 +51,7 @@ def load_rust_atranscription() -> RustAtranscription | None: def transcription( *, model: str, - audio: dict[str, object], + audio: object, api_key: str | None, api_base: str | None, custom_llm_provider: str | None, @@ -60,6 +61,7 @@ def transcription( stream: bool = False, has_custom_client: bool = False, input_source_kind: str | None = None, + context: NativeRequestContext | None = None, ) -> DispatchResult[dict[str, object]]: return attempt( load=_TRANSCRIPTION.load, @@ -75,13 +77,14 @@ def transcription( timeout_seconds=timeout_to_seconds(timeout), bedrock=bedrock_options(optional_params), ), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( execution_mode="sync", stream=stream, has_custom_client=has_custom_client, input_source_kind=input_source_kind, - ) + ), ), ), call=call_native, @@ -92,7 +95,7 @@ def transcription( async def atranscription( *, model: str, - audio: dict[str, object], + audio: object, api_key: str | None, api_base: str | None, custom_llm_provider: str | None, @@ -102,6 +105,7 @@ async def atranscription( stream: bool = False, has_custom_client: bool = False, input_source_kind: str | None = None, + context: NativeRequestContext | None = None, ) -> DispatchResult[dict[str, object]]: return await aattempt( load=_ATRANSCRIPTION.load, @@ -117,13 +121,14 @@ async def atranscription( timeout_seconds=timeout_to_seconds(timeout), bedrock=bedrock_options(optional_params), ), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( execution_mode="async", stream=stream, has_custom_client=has_custom_client, input_source_kind=input_source_kind, - ) + ), ), ), call=call_native, 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 7af9a080da6..6327d3d914a 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -126,16 +126,6 @@ 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) @@ -310,19 +300,6 @@ async def test_gate_uses_process_enable_without_request_override(): 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() 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 da70f422f44..db73a349b79 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 @@ -2309,8 +2309,17 @@ class TestRustChatCompletionsHook: seen["gate"].append(kwargs) return decline_reason - def native(**kwargs): - seen["call"].append(kwargs) + def native(request, *, options, context): + seen["call"].append( + { + "model": request.model, + "messages": request.messages, + "optional_params": request.optional_params, + "api_key": options.api_key, + "api_base": options.api_base, + "context": context, + } + ) if sync_error is not None: raise sync_error return dict(sync_result if sync_result is not None else self.RUST_RESPONSE) @@ -2464,7 +2473,7 @@ class TestRustChatCompletionsHook: RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - def declining_native(**_kwargs): + def declining_native(_request, *, options, context): raise _Declined("blank message text") monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) @@ -2501,7 +2510,7 @@ class TestRustChatCompletionsHook: monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) - async def declining_native(**_kwargs): + async def declining_native(_request, *, options, context): raise _Declined("blank message text") bridge.set_rust_chat_completions( @@ -2528,7 +2537,7 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion from litellm.rust_bridge import chat_completions as bridge - async def native(**_kwargs): + async def native(_request, *, options, context): return dict(self.RUST_RESPONSE) bridge.set_rust_chat_completions( @@ -2561,7 +2570,7 @@ class TestRustChatCompletionsHook: monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) - def declining_native(**_kwargs): + def declining_native(_request, *, options, context): raise _Declined("blank message text") bridge.set_rust_chat_completions( 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 becd6ecb832..bd1d93e7a3a 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 @@ -65,8 +65,17 @@ def _inject(*, decline_reason=None, error: Exception | None = None): seen["gate"].append(kwargs) return decline_reason - def native(**kwargs): - seen["call"].append(kwargs) + def native(request, *, options, context): + seen["call"].append( + { + "model": request.model, + "messages": request.messages, + "optional_params": request.optional_params, + "api_key": options.api_key, + "api_base": options.api_base, + "context": context, + } + ) if error is not None: raise error return dict(RUST_RESPONSE) @@ -207,7 +216,7 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch): monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) - async def declining_native(**_kwargs): + async def declining_native(_request, **_kwargs): raise _Declined("blank message text") bridge.set_rust_chat_completions( @@ -237,7 +246,7 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch): @pytest.mark.asyncio async def test_the_async_path_serves_the_rust_response_without_the_fallback(): - async def native(**_kwargs): + async def native(_request, **_kwargs): return dict(RUST_RESPONSE) bridge.set_rust_chat_completions( @@ -271,7 +280,7 @@ async def test_pre_call_logging_fires_once_even_when_the_rust_path_declines(): RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - async def declining_native(**_kwargs): + async def declining_native(_request, **_kwargs): raise _Declined("blank message text") logging_obj = MagicMock() @@ -384,7 +393,7 @@ def test_pre_call_logging_fires_once_when_the_sync_rust_path_declines(): RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - def declining_native(**_kwargs): + def declining_native(_request, **_kwargs): raise _Declined("blank message text") logging_obj = MagicMock() @@ -438,7 +447,7 @@ async def test_post_call_logging_fires_on_the_async_rust_path(): cannot drift apart the way the pre_call suppression once did.""" import json - async def native(**_kwargs): + async def native(_request, **_kwargs): return dict(RUST_RESPONSE) bridge.set_rust_chat_completions( @@ -470,7 +479,7 @@ def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines(): RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - def declining_native(**_kwargs): + def declining_native(_request, **_kwargs): raise _Declined("blank message text") logging_obj, calls = _recording_logging_obj() diff --git a/tests/test_litellm/rust_bridge/test_request_context.py b/tests/test_litellm/rust_bridge/test_request_context.py new file mode 100644 index 00000000000..1717fac035c --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_request_context.py @@ -0,0 +1,40 @@ +from types import SimpleNamespace + +from litellm.rust_bridge.request import NativeRequestCapabilities, request_context + + +def test_request_context_preserves_identity_attribution_and_capabilities() -> None: + capabilities = NativeRequestCapabilities(execution_mode="async", stream=True) + + context = request_context( + logging_obj=SimpleNamespace(litellm_call_id="call-1", litellm_trace_id="trace-1"), + request_model="router-alias", + litellm_params={ + "metadata": { + "user_api_key_hash": "hash-1", + "user_api_key_user_id": "user-1", + "user_api_key_team_id": "team-1", + } + }, + capabilities=capabilities, + ) + + assert context.litellm_call_id == "call-1" + assert context.trace_id == "trace-1" + assert context.request_model == "router-alias" + assert context.attribution.user_api_key_hash == "hash-1" + assert context.attribution.user_api_key_user_id == "user-1" + assert context.attribution.user_api_key_team_id == "team-1" + assert context.capabilities is capabilities + + +def test_request_context_ignores_untyped_identity_values() -> None: + context = request_context( + logging_obj=SimpleNamespace(litellm_call_id=1, litellm_trace_id=[]), + request_model="model", + litellm_params={"user_api_key_user_id": 42}, + ) + + assert context.litellm_call_id is None + assert context.trace_id is None + assert context.attribution.user_api_key_user_id is None