mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
feat(native): preserve request identity and capabilities
This commit is contained in:
parent
fba6af26ad
commit
0e1c72d8fe
16 changed files with 246 additions and 97 deletions
|
|
@ -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() {
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
),
|
||||
)
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
40
tests/test_litellm/rust_bridge/test_request_context.py
Normal file
40
tests/test_litellm/rust_bridge/test_request_context.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue