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