feat(native): preserve request identity and capabilities

This commit is contained in:
Yujong Lee 2026-09-05 23:04:57 -07:00 committed by yujonglee
parent fba6af26ad
commit 0e1c72d8fe
16 changed files with 246 additions and 97 deletions

View file

@ -46,7 +46,7 @@ pub(super) fn resolve_request(
) -> Result<ResolvedChatCompletionsRequest, Error> {
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() {

View file

@ -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(

View file

@ -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()

View file

@ -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(

View file

@ -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(

View file

@ -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(

View file

@ -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,

View file

@ -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,

View file

@ -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"
),
)
),
),
),
),

View file

@ -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]

View file

@ -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)

View file

@ -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,

View file

@ -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()

View file

@ -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(

View file

@ -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()

View file

@ -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