mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
refactor(native): wrap endpoint execution in a shared harness
This commit is contained in:
parent
3fe0d809d1
commit
57c21b6898
16 changed files with 899 additions and 619 deletions
|
|
@ -3,7 +3,7 @@
|
|||
"limit": 13426
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2194
|
||||
"limit": 2192
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 319
|
||||
|
|
@ -30,7 +30,7 @@
|
|||
"limit": 7
|
||||
},
|
||||
"reportGeneralTypeIssues": {
|
||||
"limit": 101
|
||||
"limit": 100
|
||||
},
|
||||
"reportIncompatibleMethodOverride": {
|
||||
"limit": 56
|
||||
|
|
@ -57,7 +57,7 @@
|
|||
"limit": 5570
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15281
|
||||
"limit": 15279
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -99,31 +99,31 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 44025
|
||||
"limit": 44019
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38269
|
||||
"limit": 38266
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19584
|
||||
"limit": 19583
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 29813
|
||||
"limit": 29810
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 110
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"limit": 687
|
||||
"limit": 686
|
||||
},
|
||||
"reportUnnecessaryContains": {
|
||||
"limit": 4
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 816
|
||||
"limit": 815
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 0
|
||||
|
|
|
|||
|
|
@ -27,7 +27,8 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
)
|
||||
from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge
|
||||
from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
|
||||
from litellm.rust_bridge.dispatch import adispatch, dispatch
|
||||
from litellm.rust_bridge.dispatch import anative_first, native_first
|
||||
from litellm.rust_bridge.runtime import DispatchResult
|
||||
from litellm.types.llms.anthropic import (
|
||||
ContentBlockDelta,
|
||||
ContentBlockStart,
|
||||
|
|
@ -369,15 +370,7 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
if config is None:
|
||||
raise ValueError(f"Provider config not found for model: {model} and provider: {custom_llm_provider}")
|
||||
|
||||
def build_request() -> tuple[dict, dict]: # mutable-ok: rewritten in place downstream
|
||||
"""Translate the request the Python way, returning `(headers, data)`.
|
||||
|
||||
The pair stays mutable because the streaming path rewrites it in
|
||||
place (`data["stream"] = True`) before sending.
|
||||
|
||||
Shared by the normal path and by the Rust path's fallback, which
|
||||
builds it only when the Rust call did not serve the request.
|
||||
"""
|
||||
def prepare_python() -> tuple[dict[str, str], dict[str, object]]: # mutable-ok: stream mutates data
|
||||
request_data: Final = config.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
|
|
@ -385,12 +378,29 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
return update_request_with_filtered_beta(
|
||||
python_headers, data = update_request_with_filtered_beta(
|
||||
headers=headers,
|
||||
request_data=request_data,
|
||||
provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
# Reaching here with `serves_via_rust` set means the Rust attempt
|
||||
# declined at call time, before the provider was called, and already
|
||||
# logged this request. That is the same attempt continuing.
|
||||
if not serves_via_rust:
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": api_base,
|
||||
"headers": python_headers,
|
||||
},
|
||||
)
|
||||
print_verbose(f"_is_function_call: {_is_function_call}")
|
||||
return python_headers, data
|
||||
|
||||
# The Rust core owns the whole call for the subset it accepts, so ask
|
||||
# before transforming: whichever path runs emits pre_call exactly once.
|
||||
# `get_config` merges the class-level defaults (Anthropic's required
|
||||
|
|
@ -407,114 +417,67 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
litellm_params=litellm_params,
|
||||
stream=stream,
|
||||
)
|
||||
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
|
||||
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
**rust_optional_params,
|
||||
},
|
||||
"api_base": api_base,
|
||||
"headers": headers,
|
||||
}
|
||||
if serves_via_rust:
|
||||
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
|
||||
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
**rust_optional_params,
|
||||
},
|
||||
"api_base": api_base,
|
||||
"headers": headers,
|
||||
}
|
||||
logging_obj.pre_call(input=messages, api_key=api_key, additional_args=rust_logging_args)
|
||||
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
|
||||
logging_obj=logging_obj,
|
||||
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
api_key=api_key,
|
||||
additional_args=rust_logging_args,
|
||||
)
|
||||
|
||||
def native_completion() -> DispatchResult[ModelResponse]:
|
||||
return rust_chat_completions_bridge.chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
additional_args=rust_logging_args,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
eligible=serves_via_rust,
|
||||
)
|
||||
if acompletion is True:
|
||||
|
||||
async def python_fallback() -> "ModelResponse | CustomStreamWrapper":
|
||||
# pre_call already fired for this request above. The Rust
|
||||
# path only declines before the provider is called, so this
|
||||
# is the same attempt continuing, not a second one.
|
||||
fallback_headers, fallback_data = build_request()
|
||||
return await self.acompletion_function(
|
||||
model=model,
|
||||
messages=messages,
|
||||
data=fallback_data,
|
||||
api_base=api_base,
|
||||
custom_prompt_dict=custom_prompt_dict,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
encoding=encoding,
|
||||
api_key=api_key,
|
||||
provider_config=config,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
stream=stream,
|
||||
_is_function_call=_is_function_call,
|
||||
litellm_params=litellm_params,
|
||||
logger_fn=logger_fn,
|
||||
headers=fallback_headers,
|
||||
client=client,
|
||||
json_mode=json_mode,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
return adispatch(
|
||||
native=lambda: rust_chat_completions_bridge.achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
),
|
||||
python=python_fallback,
|
||||
route="chat_completions",
|
||||
errors=rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model),
|
||||
)
|
||||
rust_response: Final = dispatch(
|
||||
native=lambda: rust_chat_completions_bridge.chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
),
|
||||
python=lambda: None,
|
||||
route="chat_completions",
|
||||
errors=rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model),
|
||||
)
|
||||
if rust_response is not None:
|
||||
return rust_response
|
||||
|
||||
headers, data = build_request()
|
||||
|
||||
## LOGGING
|
||||
# Reaching here with `serves_via_rust` set means the Rust attempt
|
||||
# declined at call time, before the provider was called, and already
|
||||
# logged this request. That is the same attempt continuing.
|
||||
if not serves_via_rust:
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
async def native_acompletion() -> DispatchResult[ModelResponse]:
|
||||
return await rust_chat_completions_bridge.achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": api_base,
|
||||
"headers": headers,
|
||||
},
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
eligible=serves_via_rust,
|
||||
)
|
||||
print_verbose(f"_is_function_call: {_is_function_call}")
|
||||
if acompletion is True:
|
||||
|
||||
@anative_first(
|
||||
native=native_acompletion,
|
||||
route="chat_completions",
|
||||
errors=lambda: rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model),
|
||||
)
|
||||
async def execute_async() -> ModelResponse | CustomStreamWrapper:
|
||||
headers, data = prepare_python()
|
||||
if (
|
||||
stream is True
|
||||
): # if function call - fake the streaming (need complete blocks for output parsing in openai format)
|
||||
print_verbose("makes async anthropic streaming POST request")
|
||||
data["stream"] = stream
|
||||
return self.acompletion_stream_function(
|
||||
return await self.acompletion_stream_function(
|
||||
model=model,
|
||||
messages=messages,
|
||||
data=data,
|
||||
|
|
@ -536,7 +499,7 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
client=(client if client is not None and isinstance(client, AsyncHTTPHandler) else None),
|
||||
)
|
||||
else:
|
||||
return self.acompletion_function(
|
||||
return await self.acompletion_function(
|
||||
model=model,
|
||||
messages=messages,
|
||||
data=data,
|
||||
|
|
@ -558,7 +521,14 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
json_mode=json_mode,
|
||||
timeout=timeout,
|
||||
)
|
||||
else:
|
||||
|
||||
@native_first(
|
||||
native=native_completion,
|
||||
route="chat_completions",
|
||||
errors=lambda: rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model),
|
||||
)
|
||||
def execute_sync() -> ModelResponse | CustomStreamWrapper:
|
||||
headers, data = prepare_python()
|
||||
## COMPLETION CALL
|
||||
if (
|
||||
stream is True
|
||||
|
|
@ -590,13 +560,12 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
)
|
||||
|
||||
else:
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
client = _get_httpx_client(params={"timeout": timeout})
|
||||
else:
|
||||
client = client
|
||||
python_client: Final = (
|
||||
client if isinstance(client, HTTPHandler) else _get_httpx_client(params={"timeout": timeout})
|
||||
)
|
||||
|
||||
try:
|
||||
response: Final = client.post(
|
||||
response: Final = python_client.post(
|
||||
api_base,
|
||||
headers=headers,
|
||||
data=json.dumps(data),
|
||||
|
|
@ -617,20 +586,21 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
status_code=status_code,
|
||||
headers=error_headers,
|
||||
)
|
||||
return config.transform_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
model_response=model_response,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
request_data=data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
encoding=encoding,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
|
||||
return config.transform_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
model_response=model_response,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
request_data=data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
encoding=encoding,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
return execute_async() if acompletion else execute_sync()
|
||||
|
||||
def embedding(self):
|
||||
# logic for parsing in - calling - parsing out model embedding calls
|
||||
|
|
|
|||
|
|
@ -5,7 +5,8 @@ import httpx
|
|||
|
||||
from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
|
||||
from litellm.rust_bridge import transcription as rust_transcription_bridge
|
||||
from litellm.rust_bridge.dispatch import PROPAGATE, adispatch, dispatch
|
||||
from litellm.rust_bridge.dispatch import PROPAGATE, anative_first, native_first
|
||||
from litellm.rust_bridge.runtime import DispatchResult, adapt_result
|
||||
from litellm.types.utils import FileTypes, TranscriptionResponse
|
||||
|
||||
|
||||
|
|
@ -40,6 +41,37 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
"filename": processed_audio.filename,
|
||||
}
|
||||
|
||||
def _attempt_audio_transcriptions(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
audio_file: FileTypes,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> DispatchResult[TranscriptionResponse]:
|
||||
result: Final = rust_transcription_bridge.transcription(
|
||||
model=model,
|
||||
audio=self._audio_payload(audio_file),
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
)
|
||||
return adapt_result(result, lambda response: TranscriptionResponse(**response))
|
||||
|
||||
@native_first(
|
||||
native=_attempt_audio_transcriptions,
|
||||
route="audio transcription",
|
||||
errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout: (
|
||||
PROPAGATE
|
||||
),
|
||||
)
|
||||
def audio_transcriptions(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -52,23 +84,39 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> TranscriptionResponse:
|
||||
rust_response: Final = dispatch(
|
||||
native=lambda: rust_transcription_bridge.transcription(
|
||||
model=model,
|
||||
audio=self._audio_payload(audio_file),
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
),
|
||||
python=_unavailable,
|
||||
route="audio transcription",
|
||||
errors=PROPAGATE,
|
||||
)
|
||||
return TranscriptionResponse(**rust_response)
|
||||
_unavailable()
|
||||
|
||||
async def _attempt_async_audio_transcriptions(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
audio_file: FileTypes,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> DispatchResult[TranscriptionResponse]:
|
||||
result: Final = await rust_transcription_bridge.atranscription(
|
||||
model=model,
|
||||
audio=self._audio_payload(audio_file),
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
)
|
||||
return adapt_result(result, lambda response: TranscriptionResponse(**response))
|
||||
|
||||
@anative_first(
|
||||
native=_attempt_async_audio_transcriptions,
|
||||
route="audio transcription",
|
||||
errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout: (
|
||||
PROPAGATE
|
||||
),
|
||||
)
|
||||
async def async_audio_transcriptions(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -81,19 +129,4 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> TranscriptionResponse:
|
||||
rust_response: Final = await adispatch(
|
||||
native=lambda: rust_transcription_bridge.atranscription(
|
||||
model=model,
|
||||
audio=self._audio_payload(audio_file),
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
),
|
||||
python=_aunavailable,
|
||||
route="audio transcription",
|
||||
errors=PROPAGATE,
|
||||
)
|
||||
return TranscriptionResponse(**rust_response)
|
||||
await _aunavailable()
|
||||
|
|
|
|||
|
|
@ -18,7 +18,8 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
)
|
||||
from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge
|
||||
from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
|
||||
from litellm.rust_bridge.dispatch import adispatch, dispatch
|
||||
from litellm.rust_bridge.dispatch import anative_first, native_first
|
||||
from litellm.rust_bridge.runtime import DispatchResult
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
|
|
@ -407,83 +408,62 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
litellm_params=litellm_params,
|
||||
stream=stream,
|
||||
)
|
||||
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
|
||||
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
|
||||
"messages": messages,
|
||||
**optional_params,
|
||||
},
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": headers,
|
||||
}
|
||||
if serves_via_rust:
|
||||
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
|
||||
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
|
||||
"messages": messages,
|
||||
**optional_params,
|
||||
},
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": headers,
|
||||
}
|
||||
logging_obj.pre_call(input=messages, api_key="", additional_args=rust_logging_args)
|
||||
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
api_key="",
|
||||
additional_args=rust_logging_args,
|
||||
)
|
||||
if acompletion:
|
||||
return adispatch(
|
||||
native=lambda: rust_chat_completions_bridge.achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=proxy_endpoint_url,
|
||||
custom_llm_provider="bedrock",
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
),
|
||||
python=lambda: self.async_completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=proxy_endpoint_url,
|
||||
model_response=model_response,
|
||||
encoding=encoding,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
stream=stream,
|
||||
litellm_params=litellm_params,
|
||||
logger_fn=logger_fn,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
credentials=credentials,
|
||||
api_key=api_key,
|
||||
skip_pre_call_logging=True,
|
||||
),
|
||||
route="chat_completions",
|
||||
errors=rust_chat_completions_bridge.error_handling("bedrock", model),
|
||||
)
|
||||
rust_response: Final = dispatch(
|
||||
native=lambda: rust_chat_completions_bridge.chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=proxy_endpoint_url,
|
||||
custom_llm_provider="bedrock",
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
),
|
||||
python=lambda: None,
|
||||
route="chat_completions",
|
||||
errors=rust_chat_completions_bridge.error_handling("bedrock", model),
|
||||
)
|
||||
if rust_response is not None:
|
||||
return rust_response
|
||||
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
api_key="",
|
||||
additional_args=rust_logging_args,
|
||||
)
|
||||
|
||||
### ROUTING (ASYNC, STREAMING, SYNC)
|
||||
if acompletion:
|
||||
if isinstance(client, HTTPHandler):
|
||||
client = None
|
||||
def native_completion() -> DispatchResult[ModelResponse]:
|
||||
return rust_chat_completions_bridge.chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=proxy_endpoint_url,
|
||||
custom_llm_provider="bedrock",
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
eligible=serves_via_rust,
|
||||
)
|
||||
|
||||
async def native_acompletion() -> DispatchResult[ModelResponse]:
|
||||
return await rust_chat_completions_bridge.achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=proxy_endpoint_url,
|
||||
custom_llm_provider="bedrock",
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
eligible=serves_via_rust,
|
||||
)
|
||||
|
||||
@anative_first(
|
||||
native=native_acompletion,
|
||||
route="chat_completions",
|
||||
errors=lambda: rust_chat_completions_bridge.error_handling("bedrock", model),
|
||||
)
|
||||
async def execute_async() -> ModelResponse | CustomStreamWrapper:
|
||||
python_client: Final = None if isinstance(client, HTTPHandler) else client
|
||||
if stream is True:
|
||||
return self.async_streaming(
|
||||
return await self.async_streaming(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=proxy_endpoint_url,
|
||||
|
|
@ -496,7 +476,7 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
logger_fn=logger_fn,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
client=python_client,
|
||||
json_mode=json_mode,
|
||||
fake_stream=fake_stream,
|
||||
credentials=credentials,
|
||||
|
|
@ -504,7 +484,7 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
stream_chunk_size=stream_chunk_size,
|
||||
)
|
||||
### ASYNC COMPLETION
|
||||
return self.async_completion(
|
||||
return await self.async_completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=proxy_endpoint_url,
|
||||
|
|
@ -517,108 +497,112 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
logger_fn=logger_fn,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
client=python_client,
|
||||
credentials=credentials,
|
||||
api_key=api_key,
|
||||
skip_pre_call_logging=serves_via_rust,
|
||||
)
|
||||
|
||||
@native_first(
|
||||
native=native_completion,
|
||||
route="chat_completions",
|
||||
errors=lambda: rust_chat_completions_bridge.error_handling("bedrock", model),
|
||||
)
|
||||
def execute_sync() -> ModelResponse | CustomStreamWrapper:
|
||||
## TRANSFORMATION ##
|
||||
|
||||
_data: Final = litellm.AmazonConverseConfig()._transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=extra_headers,
|
||||
)
|
||||
data: Final = json.dumps(_data)
|
||||
|
||||
prepped: Final = self.get_request_headers(
|
||||
credentials=credentials,
|
||||
aws_region_name=aws_region_name,
|
||||
extra_headers=extra_headers,
|
||||
endpoint_url=proxy_endpoint_url,
|
||||
data=data,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
## TRANSFORMATION ##
|
||||
## LOGGING
|
||||
# Reaching here with `serves_via_rust` set means the synchronous Rust
|
||||
# attempt declined at call time, before the provider was called, and
|
||||
# already logged this request. That is the same attempt continuing.
|
||||
if not serves_via_rust:
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": prepped.headers,
|
||||
},
|
||||
)
|
||||
resolved_timeout: Final = httpx.Timeout(timeout) if isinstance(timeout, (float, int)) else timeout
|
||||
python_client: Final = (
|
||||
_get_httpx_client({"timeout": resolved_timeout} if resolved_timeout is not None else None)
|
||||
if client is None or isinstance(client, AsyncHTTPHandler)
|
||||
else client
|
||||
)
|
||||
|
||||
_data: Final = litellm.AmazonConverseConfig()._transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=extra_headers,
|
||||
)
|
||||
data: Final = json.dumps(_data)
|
||||
if stream is not None and stream is True:
|
||||
completion_stream, response_headers = make_sync_call(
|
||||
client=python_client,
|
||||
api_base=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=data,
|
||||
model=model,
|
||||
messages=messages,
|
||||
logging_obj=logging_obj,
|
||||
json_mode=json_mode,
|
||||
fake_stream=fake_stream,
|
||||
stream_chunk_size=stream_chunk_size,
|
||||
)
|
||||
streaming_response: Final = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
logging_obj=logging_obj,
|
||||
_response_headers=response_headers,
|
||||
)
|
||||
|
||||
prepped: Final = self.get_request_headers(
|
||||
credentials=credentials,
|
||||
aws_region_name=aws_region_name,
|
||||
extra_headers=extra_headers,
|
||||
endpoint_url=proxy_endpoint_url,
|
||||
data=data,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
)
|
||||
return streaming_response
|
||||
|
||||
## LOGGING
|
||||
# Reaching here with `serves_via_rust` set means the synchronous Rust
|
||||
# attempt declined at call time, before the provider was called, and
|
||||
# already logged this request. That is the same attempt continuing.
|
||||
# The asynchronous branch above returns before this point, and hands
|
||||
# its own fallback `skip_pre_call_logging=True` for the same reason.
|
||||
if not serves_via_rust:
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
### COMPLETION
|
||||
|
||||
try:
|
||||
response: Final = python_client.post(
|
||||
url=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=data,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as err:
|
||||
error_code: Final = err.response.status_code
|
||||
raise BedrockError(status_code=error_code, message=err.response.text)
|
||||
except httpx.TimeoutException:
|
||||
raise BedrockError(status_code=408, message="Timeout error occurred.")
|
||||
|
||||
sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response(
|
||||
model=model,
|
||||
response=response,
|
||||
model_response=model_response,
|
||||
stream=stream if isinstance(stream, bool) else False,
|
||||
logging_obj=logging_obj,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": prepped.headers,
|
||||
},
|
||||
)
|
||||
if client is None or isinstance(client, AsyncHTTPHandler):
|
||||
_params: Final = {}
|
||||
if timeout is not None:
|
||||
if isinstance(timeout, float) or isinstance(timeout, int):
|
||||
timeout = httpx.Timeout(timeout)
|
||||
_params["timeout"] = timeout
|
||||
client = _get_httpx_client(_params)
|
||||
else:
|
||||
client = client
|
||||
|
||||
if stream is not None and stream is True:
|
||||
completion_stream, response_headers = make_sync_call(
|
||||
client=(client if client is not None and isinstance(client, HTTPHandler) else None),
|
||||
api_base=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=data,
|
||||
model=model,
|
||||
messages=messages,
|
||||
logging_obj=logging_obj,
|
||||
json_mode=json_mode,
|
||||
fake_stream=fake_stream,
|
||||
stream_chunk_size=stream_chunk_size,
|
||||
)
|
||||
streaming_response: Final = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
logging_obj=logging_obj,
|
||||
_response_headers=response_headers,
|
||||
optional_params=optional_params,
|
||||
encoding=encoding,
|
||||
)
|
||||
sync_transformed_response.set_provider_response_headers(response.headers)
|
||||
return sync_transformed_response
|
||||
|
||||
return streaming_response
|
||||
|
||||
### COMPLETION
|
||||
|
||||
try:
|
||||
response: Final = client.post(
|
||||
url=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=data,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as err:
|
||||
error_code: Final = err.response.status_code
|
||||
raise BedrockError(status_code=error_code, message=err.response.text)
|
||||
except httpx.TimeoutException:
|
||||
raise BedrockError(status_code=408, message="Timeout error occurred.")
|
||||
|
||||
sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response(
|
||||
model=model,
|
||||
response=response,
|
||||
model_response=model_response,
|
||||
stream=stream if isinstance(stream, bool) else False,
|
||||
logging_obj=logging_obj,
|
||||
api_key="",
|
||||
data=data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
encoding=encoding,
|
||||
)
|
||||
sync_transformed_response.set_provider_response_headers(response.headers)
|
||||
return sync_transformed_response
|
||||
return execute_async() if acompletion else execute_sync()
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
import asyncio
|
||||
import json
|
||||
import ssl
|
||||
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence
|
||||
from contextlib import asynccontextmanager
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Coroutine, Iterator, Mapping, Sequence
|
||||
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType, ModuleType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, TypeVar, Union, cast, get_type_hints
|
||||
|
|
@ -92,6 +92,8 @@ from litellm.responses.streaming_iterator import (
|
|||
ResponsesWebSocketStreaming,
|
||||
SyncResponsesAPIStreamingIterator,
|
||||
)
|
||||
from litellm.rust_bridge.dispatch import PYTHON_ON_ERROR, anative_context, anative_first
|
||||
from litellm.rust_bridge.runtime import DispatchResult, NativeSkipped, NativeSkipReason, adapt_result
|
||||
from litellm.types.containers.main import (
|
||||
ContainerFileListResponse,
|
||||
ContainerListResponse,
|
||||
|
|
@ -2225,116 +2227,111 @@ class BaseLLMHTTPHandler:
|
|||
},
|
||||
)
|
||||
|
||||
rust_messages_response: Final = await self._maybe_rust_anthropic_messages(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
has_agentic_hook=self._has_agentic_completion_hook(logging_obj),
|
||||
model=model,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
headers=headers,
|
||||
request_body=request_body,
|
||||
timeout=self._resolve_anthropic_messages_timeout(
|
||||
async def native_messages() -> DispatchResult[AnthropicMessagesResponse | AsyncIterator[object]]:
|
||||
result: Final = await self._attempt_rust_anthropic_messages(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
stream=stream or False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
),
|
||||
)
|
||||
if rust_messages_response is not None:
|
||||
if stream:
|
||||
return self._rust_anthropic_messages_fake_stream(rust_messages_response)
|
||||
return await self._finalize_anthropic_messages_response(
|
||||
initial_response=rust_messages_response,
|
||||
has_agentic_hook=self._has_agentic_completion_hook(logging_obj),
|
||||
model=model,
|
||||
messages=messages,
|
||||
anthropic_messages_provider_config=anthropic_messages_provider_config,
|
||||
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_key=api_key,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
response: Final = await self._async_post_anthropic_messages_with_http_error_retry(
|
||||
async_httpx_client=async_httpx_client,
|
||||
request_url=request_url,
|
||||
headers=headers,
|
||||
signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json),
|
||||
request_body=request_body,
|
||||
stream=stream or False,
|
||||
logging_obj=logging_obj,
|
||||
provider_config=anthropic_messages_provider_config,
|
||||
litellm_params=litellm_params,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
timeout=self._resolve_anthropic_messages_timeout(
|
||||
litellm_params=litellm_params,
|
||||
stream=stream or False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
),
|
||||
)
|
||||
|
||||
# used for logging + cost tracking
|
||||
logging_obj.model_call_details["httpx_response"] = response
|
||||
|
||||
initial_response: AsyncIterator | AnthropicMessagesResponse
|
||||
if stream:
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
AnthropicMessagesStreamingResponse,
|
||||
anthropic_messages_stream_hidden_params,
|
||||
)
|
||||
|
||||
completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator(
|
||||
model=model,
|
||||
httpx_response=response,
|
||||
api_base=api_base,
|
||||
headers=headers,
|
||||
request_body=request_body,
|
||||
litellm_logging_obj=logging_obj,
|
||||
timeout=self._resolve_anthropic_messages_timeout(
|
||||
litellm_params=litellm_params,
|
||||
stream=stream or False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
),
|
||||
)
|
||||
stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers)
|
||||
return adapt_result(result, self._rust_anthropic_messages_fake_stream) if stream else result
|
||||
|
||||
if not self._has_agentic_completion_hook(logging_obj):
|
||||
# No callback overrides async_should_run_agentic_loop, so the
|
||||
# agentic wrapper's only effect would be buffering every chunk
|
||||
# and rebuilding the response from SSE at end-of-stream to call
|
||||
# hooks that all return (False, {}). Stream through directly and
|
||||
# skip that per-chunk + end-of-stream overhead.
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
completion_stream=completion_stream,
|
||||
hidden_params=stream_hidden_params,
|
||||
@anative_first(native=native_messages, route="messages", errors=lambda: PYTHON_ON_ERROR)
|
||||
async def execute_messages() -> AnthropicMessagesResponse | AsyncIterator[object]:
|
||||
response: Final = await self._async_post_anthropic_messages_with_http_error_retry(
|
||||
async_httpx_client=async_httpx_client,
|
||||
request_url=request_url,
|
||||
headers=headers,
|
||||
signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json),
|
||||
request_body=request_body,
|
||||
stream=stream or False,
|
||||
logging_obj=logging_obj,
|
||||
provider_config=anthropic_messages_provider_config,
|
||||
litellm_params=litellm_params,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
timeout=self._resolve_anthropic_messages_timeout(
|
||||
litellm_params=litellm_params,
|
||||
stream=stream or False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
),
|
||||
)
|
||||
|
||||
# used for logging + cost tracking
|
||||
logging_obj.model_call_details["httpx_response"] = response
|
||||
|
||||
initial_response: AsyncIterator | AnthropicMessagesResponse
|
||||
if stream:
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
AnthropicMessagesStreamingResponse,
|
||||
anthropic_messages_stream_hidden_params,
|
||||
)
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
|
||||
AgenticAnthropicStreamingIterator,
|
||||
)
|
||||
completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator(
|
||||
model=model,
|
||||
httpx_response=response,
|
||||
request_body=request_body,
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers)
|
||||
|
||||
held_back_tool_names: Final = self._server_fulfilled_tools_in_request(
|
||||
logging_obj=logging_obj,
|
||||
tools=anthropic_messages_optional_request_params.get("tools"),
|
||||
)
|
||||
initial_response = AgenticAnthropicStreamingIterator(
|
||||
completion_stream=completion_stream,
|
||||
http_handler=self,
|
||||
model=model,
|
||||
messages=messages,
|
||||
anthropic_messages_provider_config=anthropic_messages_provider_config,
|
||||
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs={**kwargs, "api_key": api_key} if api_key else kwargs,
|
||||
hold_back=bool(held_back_tool_names),
|
||||
server_fulfilled_tool_names=held_back_tool_names,
|
||||
)
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
completion_stream=initial_response,
|
||||
hidden_params=stream_hidden_params,
|
||||
)
|
||||
else:
|
||||
initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
if not self._has_agentic_completion_hook(logging_obj):
|
||||
# No callback overrides async_should_run_agentic_loop, so the
|
||||
# agentic wrapper's only effect would be buffering every chunk
|
||||
# and rebuilding the response from SSE at end-of-stream to call
|
||||
# hooks that all return (False, {}). Stream through directly and
|
||||
# skip that per-chunk + end-of-stream overhead.
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
completion_stream=completion_stream,
|
||||
hidden_params=stream_hidden_params,
|
||||
)
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
|
||||
AgenticAnthropicStreamingIterator,
|
||||
)
|
||||
|
||||
held_back_tool_names: Final = self._server_fulfilled_tools_in_request(
|
||||
logging_obj=logging_obj,
|
||||
tools=anthropic_messages_optional_request_params.get("tools"),
|
||||
)
|
||||
initial_response = AgenticAnthropicStreamingIterator(
|
||||
completion_stream=completion_stream,
|
||||
http_handler=self,
|
||||
model=model,
|
||||
messages=messages,
|
||||
anthropic_messages_provider_config=anthropic_messages_provider_config,
|
||||
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs={**kwargs, "api_key": api_key} if api_key else kwargs,
|
||||
hold_back=bool(held_back_tool_names),
|
||||
server_fulfilled_tool_names=held_back_tool_names,
|
||||
)
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
completion_stream=initial_response,
|
||||
hidden_params=stream_hidden_params,
|
||||
)
|
||||
else:
|
||||
initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
return initial_response
|
||||
|
||||
initial_response: Final = await execute_messages()
|
||||
if stream:
|
||||
return initial_response
|
||||
return await self._finalize_anthropic_messages_response(
|
||||
initial_response=initial_response,
|
||||
model=model,
|
||||
|
|
@ -2384,7 +2381,7 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
async def _maybe_rust_anthropic_messages(
|
||||
async def _attempt_rust_anthropic_messages(
|
||||
*,
|
||||
custom_llm_provider: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
|
|
@ -2395,41 +2392,36 @@ class BaseLLMHTTPHandler:
|
|||
headers: dict,
|
||||
request_body: dict,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> AnthropicMessagesResponse | None:
|
||||
) -> DispatchResult[AnthropicMessagesResponse]:
|
||||
if custom_llm_provider not in ("azure_ai", "anthropic"):
|
||||
return None
|
||||
return NativeSkipped(NativeSkipReason.INELIGIBLE)
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
|
||||
if not rust_enabled():
|
||||
return None
|
||||
return NativeSkipped(NativeSkipReason.DISABLED)
|
||||
if has_agentic_hook:
|
||||
return None
|
||||
return NativeSkipped(NativeSkipReason.INELIGIBLE)
|
||||
|
||||
from litellm.rust_bridge import messages as rust_messages_bridge
|
||||
|
||||
upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"}
|
||||
from litellm.rust_bridge.dispatch import PYTHON_ON_ERROR, adispatch, async_none
|
||||
|
||||
rust_response: Final = await adispatch(
|
||||
native=lambda: rust_messages_bridge.amessages(
|
||||
model=model,
|
||||
body=upstream_body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
),
|
||||
python=async_none,
|
||||
route="messages",
|
||||
errors=PYTHON_ON_ERROR,
|
||||
result: Final = await rust_messages_bridge.amessages(
|
||||
model=model,
|
||||
body=upstream_body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
if rust_response is None:
|
||||
return None
|
||||
|
||||
response_obj: Final = cast(AnthropicMessagesResponse, dict(rust_response))
|
||||
response_obj["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}}
|
||||
return response_obj
|
||||
def adapt(rust_response: dict[str, object]) -> AnthropicMessagesResponse:
|
||||
return cast(
|
||||
AnthropicMessagesResponse,
|
||||
{**rust_response, "_hidden_params": {"additional_headers": {"x-litellm-rust": "true"}}},
|
||||
)
|
||||
|
||||
return adapt_result(result, adapt)
|
||||
|
||||
@staticmethod
|
||||
def _rust_anthropic_messages_fake_stream(
|
||||
|
|
@ -6507,26 +6499,22 @@ class BaseLLMHTTPHandler:
|
|||
},
|
||||
)
|
||||
|
||||
from litellm.rust_bridge import responses_websocket as rust_responses_websocket
|
||||
|
||||
async def attempt_connection() -> DispatchResult[
|
||||
AbstractAsyncContextManager[rust_responses_websocket.ConnectionAdapter]
|
||||
]:
|
||||
if not _rust_responses_websocket_enabled(custom_llm_provider):
|
||||
return NativeSkipped(NativeSkipReason.INELIGIBLE)
|
||||
return await rust_responses_websocket.managed_connect(
|
||||
url=ws_url,
|
||||
headers={str(key): str(value) for key, value in headers.items()},
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
@anative_context(native=attempt_connection, route="responses_websocket", errors=lambda: PYTHON_ON_ERROR)
|
||||
@asynccontextmanager
|
||||
async def _backend_connection():
|
||||
if _rust_responses_websocket_enabled(custom_llm_provider):
|
||||
from litellm.rust_bridge import responses_websocket as rust_responses_websocket
|
||||
from litellm.rust_bridge.dispatch import PYTHON_ON_ERROR, adispatch, async_none
|
||||
|
||||
rust_backend: Final = await adispatch(
|
||||
native=lambda: rust_responses_websocket.connect(
|
||||
url=ws_url,
|
||||
headers={str(key): str(value) for key, value in headers.items()},
|
||||
timeout=timeout,
|
||||
),
|
||||
python=async_none,
|
||||
route="responses_websocket",
|
||||
errors=PYTHON_ON_ERROR,
|
||||
)
|
||||
if rust_backend is not None:
|
||||
yield rust_backend
|
||||
return
|
||||
|
||||
async def _backend_connection() -> AsyncGenerator[ClientConnection, None]:
|
||||
async with websockets.connect(
|
||||
ws_url,
|
||||
additional_headers=headers,
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import base64
|
|||
import mimetypes
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Coroutine, Mapping
|
||||
from collections.abc import Callable, Coroutine, Mapping
|
||||
from io import IOBase
|
||||
from typing import Any, Final, cast
|
||||
|
||||
|
|
@ -25,7 +25,8 @@ from litellm.llms.base_llm.ocr.transformation import (
|
|||
)
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge import ocr as rust_ocr_bridge
|
||||
from litellm.rust_bridge.dispatch import PROPAGATE, adispatch, dispatch
|
||||
from litellm.rust_bridge.dispatch import PROPAGATE, anative_first, native_first
|
||||
from litellm.rust_bridge.runtime import DispatchResult
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import ProviderConfigManager, client
|
||||
|
||||
|
|
@ -155,6 +156,69 @@ def _prepare_ocr_request(
|
|||
)
|
||||
|
||||
|
||||
@anative_first(
|
||||
native=rust_ocr_bridge.aattempt_ocr,
|
||||
route="ocr",
|
||||
errors=lambda prepared_request, resolve_api_key: PROPAGATE,
|
||||
)
|
||||
async def _execute_aocr(
|
||||
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
) -> OCRResponse:
|
||||
pending: Final = base_llm_http_handler.ocr(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
optional_params=prepared_request.optional_params,
|
||||
timeout=prepared_request.effective_timeout,
|
||||
logging_obj=prepared_request.litellm_logging_obj,
|
||||
api_key=prepared_request.api_key,
|
||||
api_base=prepared_request.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
aocr=True,
|
||||
headers=prepared_request.extra_headers,
|
||||
provider_config=prepared_request.provider_config,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
response: Final = await pending if asyncio.iscoroutine(pending) else pending
|
||||
if response is None:
|
||||
raise ValueError(f"Got an unexpected None response from the OCR API: {response}")
|
||||
return response
|
||||
|
||||
|
||||
def _attempt_ocr(
|
||||
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
is_async: bool,
|
||||
) -> DispatchResult[OCRResponse]:
|
||||
return rust_ocr_bridge.attempt_ocr(prepared_request=prepared_request, resolve_api_key=resolve_api_key)
|
||||
|
||||
|
||||
@native_first(
|
||||
native=_attempt_ocr,
|
||||
route="ocr",
|
||||
errors=lambda prepared_request, resolve_api_key, is_async: PROPAGATE,
|
||||
)
|
||||
def _execute_ocr(
|
||||
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
is_async: bool,
|
||||
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
|
||||
return base_llm_http_handler.ocr(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
optional_params=prepared_request.optional_params,
|
||||
timeout=prepared_request.effective_timeout,
|
||||
logging_obj=prepared_request.litellm_logging_obj,
|
||||
api_key=prepared_request.api_key,
|
||||
api_base=prepared_request.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
aocr=is_async,
|
||||
headers=prepared_request.extra_headers,
|
||||
provider_config=prepared_request.provider_config,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
async def aocr(
|
||||
model: str,
|
||||
|
|
@ -251,32 +315,7 @@ async def aocr(
|
|||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
async def python_fallback() -> OCRResponse:
|
||||
pending: Final = base_llm_http_handler.ocr(
|
||||
model=prepared.model,
|
||||
document=prepared.document,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared.effective_timeout,
|
||||
logging_obj=prepared.litellm_logging_obj,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared.custom_llm_provider,
|
||||
aocr=True,
|
||||
headers=prepared.extra_headers,
|
||||
provider_config=prepared.provider_config,
|
||||
litellm_params=prepared.litellm_params,
|
||||
)
|
||||
response: Final = await pending if asyncio.iscoroutine(pending) else pending
|
||||
if response is None:
|
||||
raise ValueError(f"Got an unexpected None response from the OCR API: {response}")
|
||||
return response
|
||||
|
||||
return await adispatch(
|
||||
native=lambda: rust_ocr_bridge.aattempt_ocr(prepared_request=prepared, resolve_api_key=get_secret_str),
|
||||
python=python_fallback,
|
||||
route="ocr",
|
||||
errors=PROPAGATE,
|
||||
)
|
||||
return await _execute_aocr(prepared_request=prepared, resolve_api_key=get_secret_str)
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
|
|
@ -517,28 +556,7 @@ def ocr(
|
|||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
def python_fallback() -> OCRResponse | Coroutine[object, object, OCRResponse]:
|
||||
return base_llm_http_handler.ocr(
|
||||
model=prepared.model,
|
||||
document=prepared.document,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared.effective_timeout,
|
||||
logging_obj=prepared.litellm_logging_obj,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared.custom_llm_provider,
|
||||
aocr=_is_async,
|
||||
headers=prepared.extra_headers,
|
||||
provider_config=prepared.provider_config,
|
||||
litellm_params=prepared.litellm_params,
|
||||
)
|
||||
|
||||
return dispatch(
|
||||
native=lambda: rust_ocr_bridge.attempt_ocr(prepared_request=prepared, resolve_api_key=get_secret_str),
|
||||
python=python_fallback,
|
||||
route="ocr",
|
||||
errors=PROPAGATE,
|
||||
)
|
||||
return _execute_ocr(prepared_request=prepared, resolve_api_key=get_secret_str, is_async=_is_async)
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -225,6 +225,7 @@ def chat_completions(
|
|||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
on_response: ResponseObserver,
|
||||
eligible: bool = True,
|
||||
) -> DispatchResult[ModelResponse]:
|
||||
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
|
||||
on_response(rust_response)
|
||||
|
|
@ -245,7 +246,7 @@ def chat_completions(
|
|||
return attempt(
|
||||
load=_CHAT.load,
|
||||
enabled=rust_enabled(),
|
||||
eligible=True,
|
||||
eligible=eligible,
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
|
|
@ -264,6 +265,7 @@ async def achat_completions(
|
|||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
on_response: ResponseObserver,
|
||||
eligible: bool = True,
|
||||
) -> DispatchResult[ModelResponse]:
|
||||
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
|
||||
on_response(rust_response)
|
||||
|
|
@ -284,7 +286,7 @@ async def achat_completions(
|
|||
return await aattempt(
|
||||
load=_ACHAT.load,
|
||||
enabled=rust_enabled(),
|
||||
eligible=True,
|
||||
eligible=eligible,
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
|
|
|
|||
|
|
@ -1,9 +1,11 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable
|
||||
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Final, TypeAlias, TypeVar
|
||||
from functools import wraps
|
||||
from typing import Final, ParamSpec, TypeAlias, TypeVar
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.exceptions import APIError
|
||||
|
|
@ -12,6 +14,7 @@ from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, N
|
|||
|
||||
NativeT = TypeVar("NativeT")
|
||||
PythonT = TypeVar("PythonT")
|
||||
P = ParamSpec("P")
|
||||
|
||||
|
||||
class ErrorAction(Enum):
|
||||
|
|
@ -89,49 +92,95 @@ def _log_skip(route: str, skipped: NativeSkipped) -> None:
|
|||
verbose_logger.debug("Native %s skipped (%s): %s", route, skipped.reason.value, skipped.detail or "")
|
||||
|
||||
|
||||
def dispatch(
|
||||
def native_first(
|
||||
*,
|
||||
native: Callable[[], DispatchResult[NativeT]],
|
||||
python: Callable[[], PythonT],
|
||||
native: Callable[P, DispatchResult[NativeT]],
|
||||
route: str,
|
||||
errors: ErrorHandling,
|
||||
) -> NativeT | PythonT:
|
||||
try:
|
||||
attempted: Final = native()
|
||||
except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures
|
||||
unexpected: Final = _handle_error(error, errors.unexpected, route, NativeSkipReason.FAILED)
|
||||
_log_skip(route, unexpected)
|
||||
return python()
|
||||
result: Final = _resolve(attempted, errors, route)
|
||||
match result:
|
||||
case Handled(value):
|
||||
return value
|
||||
case NativeSkipped():
|
||||
_log_skip(route, result)
|
||||
return python()
|
||||
errors: Callable[P, ErrorHandling],
|
||||
) -> Callable[[Callable[P, PythonT]], Callable[P, NativeT | PythonT]]:
|
||||
def wrap(implementation: Callable[P, PythonT]) -> Callable[P, NativeT | PythonT]:
|
||||
@wraps(implementation)
|
||||
def run(
|
||||
*args: P.args,
|
||||
**kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature
|
||||
) -> NativeT | PythonT:
|
||||
rules: Final = errors(*args, **kwargs)
|
||||
try:
|
||||
attempted: Final = native(*args, **kwargs)
|
||||
except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures
|
||||
skipped: Final = _handle_error(error, rules.unexpected, route, NativeSkipReason.FAILED)
|
||||
_log_skip(route, skipped)
|
||||
else:
|
||||
result: Final = _resolve(attempted, rules, route)
|
||||
if isinstance(result, Handled):
|
||||
return result.value
|
||||
_log_skip(route, result)
|
||||
return implementation(*args, **kwargs)
|
||||
|
||||
return run
|
||||
|
||||
return wrap
|
||||
|
||||
|
||||
async def adispatch(
|
||||
def anative_first(
|
||||
*,
|
||||
native: Callable[[], Awaitable[DispatchResult[NativeT]]],
|
||||
python: Callable[[], Awaitable[PythonT]],
|
||||
native: Callable[P, Awaitable[DispatchResult[NativeT]]],
|
||||
route: str,
|
||||
errors: ErrorHandling,
|
||||
) -> NativeT | PythonT:
|
||||
try:
|
||||
attempted: Final = await native()
|
||||
except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures
|
||||
unexpected: Final = _handle_error(error, errors.unexpected, route, NativeSkipReason.FAILED)
|
||||
_log_skip(route, unexpected)
|
||||
return await python()
|
||||
result: Final = _resolve(attempted, errors, route)
|
||||
match result:
|
||||
case Handled(value):
|
||||
return value
|
||||
case NativeSkipped():
|
||||
_log_skip(route, result)
|
||||
return await python()
|
||||
errors: Callable[P, ErrorHandling],
|
||||
) -> Callable[[Callable[P, Awaitable[PythonT]]], Callable[P, Awaitable[NativeT | PythonT]]]:
|
||||
def wrap(implementation: Callable[P, Awaitable[PythonT]]) -> Callable[P, Awaitable[NativeT | PythonT]]:
|
||||
@wraps(implementation)
|
||||
async def run(
|
||||
*args: P.args,
|
||||
**kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature
|
||||
) -> NativeT | PythonT:
|
||||
rules: Final = errors(*args, **kwargs)
|
||||
try:
|
||||
attempted: Final = await native(*args, **kwargs)
|
||||
except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures
|
||||
skipped: Final = _handle_error(error, rules.unexpected, route, NativeSkipReason.FAILED)
|
||||
_log_skip(route, skipped)
|
||||
else:
|
||||
result: Final = _resolve(attempted, rules, route)
|
||||
if isinstance(result, Handled):
|
||||
return result.value
|
||||
_log_skip(route, result)
|
||||
return await implementation(*args, **kwargs)
|
||||
|
||||
return run
|
||||
|
||||
return wrap
|
||||
|
||||
|
||||
async def async_none() -> None:
|
||||
return None
|
||||
def anative_context(
|
||||
*,
|
||||
native: Callable[P, Awaitable[DispatchResult[AbstractAsyncContextManager[NativeT]]]],
|
||||
route: str,
|
||||
errors: Callable[P, ErrorHandling],
|
||||
) -> Callable[
|
||||
[Callable[P, AbstractAsyncContextManager[PythonT]]],
|
||||
Callable[P, AbstractAsyncContextManager[NativeT | PythonT]],
|
||||
]:
|
||||
def wrap(
|
||||
implementation: Callable[P, AbstractAsyncContextManager[PythonT]],
|
||||
) -> Callable[P, AbstractAsyncContextManager[NativeT | PythonT]]:
|
||||
@anative_first(native=native, route=route, errors=errors)
|
||||
async def acquire(
|
||||
*args: P.args,
|
||||
**kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature
|
||||
) -> AbstractAsyncContextManager[PythonT]:
|
||||
return implementation(*args, **kwargs)
|
||||
|
||||
@wraps(implementation)
|
||||
@asynccontextmanager
|
||||
async def run(
|
||||
*args: P.args,
|
||||
**kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature
|
||||
) -> AsyncGenerator[NativeT | PythonT, None]:
|
||||
manager: Final = await acquire(*args, **kwargs)
|
||||
async with manager as connection:
|
||||
yield connection
|
||||
|
||||
return run
|
||||
|
||||
return wrap
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncGenerator
|
||||
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
|
@ -13,7 +15,7 @@ from litellm.rust_bridge.protocols import (
|
|||
RustResponsesWebSocket,
|
||||
RustResponsesWebSocketConnection,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import DispatchResult, aattempt
|
||||
from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
_RESPONSES_WEBSOCKET: Final[NativeBinding[RustResponsesWebSocketConnection]] = NativeBinding(
|
||||
|
|
@ -32,7 +34,7 @@ def set_rust_responses_websocket(
|
|||
_RESPONSES_WEBSOCKET.override(connection)
|
||||
|
||||
|
||||
class _ConnectionAdapter:
|
||||
class ConnectionAdapter:
|
||||
def __init__(self, connection: RustResponsesWebSocket):
|
||||
self._connection: Final[RustResponsesWebSocket] = connection
|
||||
|
||||
|
|
@ -54,7 +56,7 @@ async def connect(
|
|||
url: str,
|
||||
headers: dict[str, str],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> DispatchResult[_ConnectionAdapter]:
|
||||
) -> DispatchResult[ConnectionAdapter]:
|
||||
return await aattempt(
|
||||
load=_RESPONSES_WEBSOCKET.load,
|
||||
enabled=rust_enabled(),
|
||||
|
|
@ -65,5 +67,23 @@ async def connect(
|
|||
headers=headers,
|
||||
timeout_seconds=timeout_seconds,
|
||||
),
|
||||
adapt=_ConnectionAdapter,
|
||||
adapt=ConnectionAdapter,
|
||||
)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _connection_context(connection: ConnectionAdapter) -> AsyncGenerator[ConnectionAdapter, None]:
|
||||
try:
|
||||
yield connection
|
||||
finally:
|
||||
await connection.close()
|
||||
|
||||
|
||||
async def managed_connect(
|
||||
*,
|
||||
url: str,
|
||||
headers: dict[str, str],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> DispatchResult[AbstractAsyncContextManager[ConnectionAdapter]]:
|
||||
result: Final = await connect(url=url, headers=headers, timeout=timeout)
|
||||
return adapt_result(result, _connection_context)
|
||||
|
|
|
|||
|
|
@ -87,3 +87,9 @@ async def aattempt(
|
|||
|
||||
def identity(value: ResultT) -> ResultT:
|
||||
return value
|
||||
|
||||
|
||||
def adapt_result(result: DispatchResult[NativeT], adapt: Callable[[NativeT], ResultT]) -> DispatchResult[ResultT]:
|
||||
if isinstance(result, Handled):
|
||||
return Handled(adapt(result.value))
|
||||
return result
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@
|
|||
"limit": 1979
|
||||
},
|
||||
"ANN202": {
|
||||
"limit": 831
|
||||
"limit": 830
|
||||
},
|
||||
"ANN204": {
|
||||
"limit": 683
|
||||
|
|
@ -150,7 +150,7 @@
|
|||
"limit": 253
|
||||
},
|
||||
"PLW0127": {
|
||||
"limit": 57
|
||||
"limit": 55
|
||||
},
|
||||
"PLW0602": {
|
||||
"limit": 215
|
||||
|
|
@ -195,7 +195,7 @@
|
|||
"limit": 22
|
||||
},
|
||||
"SIM101": {
|
||||
"limit": 56
|
||||
"limit": 55
|
||||
},
|
||||
"SIM102": {
|
||||
"limit": 310
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ import pytest
|
|||
import litellm
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge.runtime import Handled, NativeSkipped, NativeSkipReason
|
||||
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
)
|
||||
|
|
@ -127,6 +127,16 @@ def test_load_rust_messages_returns_injected_impl():
|
|||
assert rust_messages.load_rust_messages() is bridge
|
||||
|
||||
|
||||
def test_bare_rust_still_toggles_ocr():
|
||||
from litellm.rust_bridge.ocr import rust_ocr_enabled
|
||||
|
||||
litellm.rust(True)
|
||||
assert rust_ocr_enabled() is True
|
||||
|
||||
litellm.rust(False)
|
||||
assert rust_ocr_enabled() is False
|
||||
|
||||
|
||||
def test_load_rust_amessages_returns_injected_impl():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.rust(True)
|
||||
|
|
@ -215,7 +225,7 @@ def _gate(**overrides):
|
|||
"timeout": 30.0,
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return BaseLLMHTTPHandler._maybe_rust_anthropic_messages(**kwargs)
|
||||
return BaseLLMHTTPHandler._attempt_rust_anthropic_messages(**kwargs)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -226,7 +236,8 @@ async def test_gate_invokes_rust_and_marks_response_header():
|
|||
|
||||
response = await _gate()
|
||||
|
||||
assert response is not None
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert response["id"] == "msg_123"
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
call = bridge.calls[0]
|
||||
|
|
@ -239,13 +250,13 @@ async def test_gate_invokes_rust_and_marks_response_header():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_falls_back_to_python_when_bridge_raises():
|
||||
async def test_gate_reports_failure_to_harness():
|
||||
bridge = RaisingAsyncMessages()
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate()
|
||||
assert response is None
|
||||
assert isinstance(response, NativeFailed)
|
||||
assert bridge.calls == 1
|
||||
|
||||
|
||||
|
|
@ -256,7 +267,7 @@ async def test_gate_skips_rust_when_flag_absent():
|
|||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
|
||||
|
||||
assert response is None
|
||||
assert isinstance(response, NativeSkipped)
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
|
|
@ -268,10 +279,24 @@ async def test_gate_uses_process_enable_without_request_override():
|
|||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
|
||||
|
||||
assert response is not None
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert bridge.calls[0]["custom_llm_provider"] == "azure_ai"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_ignores_request_flag_when_process_enabled():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False))
|
||||
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert len(bridge.calls) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_invokes_rust_for_native_anthropic_provider():
|
||||
bridge = RecordingAsyncMessages()
|
||||
|
|
@ -286,7 +311,8 @@ async def test_gate_invokes_rust_for_native_anthropic_provider():
|
|||
headers={"x-api-key": "sk-ant", "anthropic-version": "2023-06-01"},
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert bridge.calls[0]["custom_llm_provider"] == "anthropic"
|
||||
assert bridge.calls[0]["api_key"] == "sk-ant"
|
||||
|
|
@ -303,7 +329,8 @@ async def test_gate_invokes_rust_when_env_var_set(monkeypatch):
|
|||
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert bridge.calls[0]["custom_llm_provider"] == "anthropic"
|
||||
|
||||
|
||||
|
|
@ -318,7 +345,7 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch):
|
|||
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
|
||||
)
|
||||
|
||||
assert response is None
|
||||
assert isinstance(response, NativeSkipped)
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
|
|
@ -330,7 +357,7 @@ async def test_gate_skips_rust_for_unsupported_provider():
|
|||
|
||||
response = await _gate(custom_llm_provider="openai")
|
||||
|
||||
assert response is None
|
||||
assert isinstance(response, NativeSkipped)
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
|
|
@ -342,7 +369,7 @@ async def test_gate_skips_rust_for_agentic_hook():
|
|||
|
||||
response = await _gate(has_agentic_hook=True)
|
||||
|
||||
assert response is None
|
||||
assert isinstance(response, NativeSkipped)
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
|
|
@ -358,7 +385,8 @@ async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag():
|
|||
request_body=streaming_body,
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert "stream" not in bridge.calls[0]["body"]
|
||||
assert bridge.calls[0]["body"] == REQUEST_BODY
|
||||
|
|
@ -391,4 +419,56 @@ async def test_gate_falls_back_when_bridge_unavailable(monkeypatch):
|
|||
|
||||
response = await _gate()
|
||||
|
||||
assert response is None
|
||||
assert isinstance(response, NativeSkipped)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("selection", ("native", "disabled", "failed"))
|
||||
async def test_messages_handler_runs_selected_backend_once(selection: str) -> None:
|
||||
from datetime import datetime
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import AnthropicMessagesConfig
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
bridge = RaisingAsyncMessages() if selection == "failed" else RecordingAsyncMessages()
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
litellm.rust(selection != "disabled")
|
||||
requests: list[httpx.Request] = []
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
requests.append(request)
|
||||
return httpx.Response(200, json=FAKE_MESSAGES_RESPONSE)
|
||||
|
||||
logging_obj = Logging(
|
||||
model=FAKE_MESSAGES_RESPONSE["model"],
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="anthropic_messages",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="harness-test",
|
||||
function_id="harness-test",
|
||||
)
|
||||
client = AsyncHTTPHandler()
|
||||
await client.client.aclose()
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as transport:
|
||||
client.client = transport
|
||||
response = await BaseLLMHTTPHandler().async_anthropic_messages_handler(
|
||||
model=FAKE_MESSAGES_RESPONSE["model"],
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
anthropic_messages_provider_config=AnthropicMessagesConfig(),
|
||||
anthropic_messages_optional_request_params={"max_tokens": 10},
|
||||
custom_llm_provider="anthropic",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
logging_obj=logging_obj,
|
||||
api_key="sk-test",
|
||||
api_base="https://example.test",
|
||||
client=client,
|
||||
)
|
||||
assert response["id"] == FAKE_MESSAGES_RESPONSE["id"]
|
||||
assert len(requests) == (0 if selection == "native" else 1)
|
||||
assert (bridge.calls if isinstance(bridge, RaisingAsyncMessages) else len(bridge.calls)) == (
|
||||
0 if selection == "disabled" else 1
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import pytest
|
|||
|
||||
from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled
|
||||
from litellm.rust_bridge import configuration, responses_websocket
|
||||
from litellm.rust_bridge.runtime import Handled, NativeSkipped, NativeSkipReason, NativeFailed
|
||||
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
|
||||
|
||||
|
||||
class _FakeNativeConnection:
|
||||
|
|
@ -58,7 +58,7 @@ def test_rust_websocket_bridge_uses_process_enablement() -> None:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None:
|
||||
adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection())
|
||||
adapter = responses_websocket.ConnectionAdapter(_ClosedNativeConnection())
|
||||
|
||||
with pytest.raises(responses_websocket.ConnectionClosedOK):
|
||||
await adapter.recv()
|
||||
|
|
@ -115,3 +115,30 @@ async def test_connection_failure_is_reported_to_orchestration() -> None:
|
|||
result = await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None)
|
||||
assert isinstance(result, NativeFailed)
|
||||
assert str(result.error) == "connection failed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_managed_connection_closes_native_socket_on_consumer_failure() -> None:
|
||||
configuration.rust(True)
|
||||
socket = _FakeNativeConnection()
|
||||
|
||||
class Bridge:
|
||||
@classmethod
|
||||
async def connect(
|
||||
cls, *, url: str, headers: dict[str, str], timeout_seconds: float | None
|
||||
) -> _FakeNativeConnection:
|
||||
return socket
|
||||
|
||||
responses_websocket.set_rust_responses_websocket(connection=Bridge)
|
||||
result = await responses_websocket.managed_connect(url="wss://example.test/responses", headers={}, timeout=1.0)
|
||||
assert isinstance(result, Handled)
|
||||
|
||||
async def use_connection() -> None:
|
||||
async with result.value as connection:
|
||||
await connection.send("hello")
|
||||
raise ValueError("consumer failed")
|
||||
|
||||
with pytest.raises(ValueError, match="consumer failed"):
|
||||
await use_connection()
|
||||
assert socket.sent == ["hello"]
|
||||
assert socket.closed
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ import pytest
|
|||
from litellm.exceptions import APIError
|
||||
from litellm.rust_bridge import bindings
|
||||
from litellm.rust_bridge.chat_completions import error_handling
|
||||
from litellm.rust_bridge.dispatch import PROPAGATE, PYTHON_ON_ERROR, adispatch, dispatch
|
||||
from litellm.rust_bridge.dispatch import PROPAGATE, PYTHON_ON_ERROR, anative_first, native_first
|
||||
from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, NativeSkipped, NativeSkipReason
|
||||
|
||||
|
||||
|
|
@ -53,9 +53,9 @@ async def test_shared_dispatch_calls_python_once_and_logs_skip(
|
|||
return python()
|
||||
|
||||
result: Final = (
|
||||
await adispatch(native=anative, python=apython, route="test", errors=PROPAGATE)
|
||||
await anative_first(native=anative, route="test", errors=lambda: PROPAGATE)(apython)()
|
||||
if asynchronous
|
||||
else dispatch(native=native, python=python, route="test", errors=PROPAGATE)
|
||||
else native_first(native=native, route="test", errors=lambda: PROPAGATE)(python)()
|
||||
)
|
||||
assert result == "python response"
|
||||
assert calls == ["native", "python"]
|
||||
|
|
@ -65,6 +65,7 @@ async def test_shared_dispatch_calls_python_once_and_logs_skip(
|
|||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_native_success_does_not_run_python_even_when_value_is_none(asynchronous: bool) -> None:
|
||||
|
||||
async def native() -> DispatchResult[None]:
|
||||
return Handled(None)
|
||||
|
||||
|
|
@ -75,9 +76,9 @@ async def test_native_success_does_not_run_python_even_when_value_is_none(asynch
|
|||
return python()
|
||||
|
||||
result: Final = (
|
||||
await adispatch(native=native, python=apython, route="test", errors=PYTHON_ON_ERROR)
|
||||
await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(apython)()
|
||||
if asynchronous
|
||||
else dispatch(native=lambda: Handled(None), python=python, route="test", errors=PYTHON_ON_ERROR)
|
||||
else native_first(native=lambda: Handled(None), route="test", errors=lambda: PYTHON_ON_ERROR)(python)()
|
||||
)
|
||||
assert result is None
|
||||
|
||||
|
|
@ -124,8 +125,8 @@ async def test_declarations_preserve_endpoint_error_behavior(
|
|||
|
||||
async def run() -> str:
|
||||
if asynchronous:
|
||||
return await adispatch(native=anative, python=apython, route="chat_completions", errors=rules)
|
||||
return dispatch(native=native, python=python, route="chat_completions", errors=rules)
|
||||
return await anative_first(native=anative, route="chat_completions", errors=lambda: rules)(apython)()
|
||||
return native_first(native=native, route="chat_completions", errors=lambda: rules)(python)()
|
||||
|
||||
if policy == "python" or (policy == "chat" and kind in ("declined", "missing")):
|
||||
assert await run() == "python response"
|
||||
|
|
@ -163,13 +164,10 @@ async def test_python_failure_is_never_reclassified_as_native_failure(asynchrono
|
|||
|
||||
async def run() -> str:
|
||||
if asynchronous:
|
||||
return await adispatch(native=native, python=apython, route="test", errors=PYTHON_ON_ERROR)
|
||||
return dispatch(
|
||||
native=lambda: NativeSkipped(NativeSkipReason.UNAVAILABLE),
|
||||
python=python,
|
||||
route="test",
|
||||
errors=PYTHON_ON_ERROR,
|
||||
)
|
||||
return await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(apython)()
|
||||
return native_first(
|
||||
native=lambda: NativeSkipped(NativeSkipReason.UNAVAILABLE), route="test", errors=lambda: PYTHON_ON_ERROR
|
||||
)(python)()
|
||||
|
||||
with pytest.raises(RuntimeError) as caught:
|
||||
await run()
|
||||
|
|
@ -179,6 +177,7 @@ async def test_python_failure_is_never_reclassified_as_native_failure(asynchrono
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancellation_does_not_run_python() -> None:
|
||||
|
||||
async def native() -> DispatchResult[str]:
|
||||
raise asyncio.CancelledError
|
||||
|
||||
|
|
@ -186,20 +185,123 @@ async def test_cancellation_does_not_run_python() -> None:
|
|||
pytest.fail("cancellation must not dispatch Python")
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await adispatch(native=native, python=python, route="test", errors=PYTHON_ON_ERROR)
|
||||
await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(python)()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("status", (0, 401, 403, 429, 500, 503))
|
||||
def test_chat_upstream_mapping_preserves_status_message_and_context(status: int) -> None:
|
||||
error: Final = Upstream(status, "upstream failed")
|
||||
with pytest.raises(APIError, match="upstream failed") as caught:
|
||||
dispatch(
|
||||
native_first(
|
||||
native=lambda: NativeFailed(error),
|
||||
python=lambda: pytest.fail("upstream errors must not run Python"),
|
||||
route="chat_completions",
|
||||
errors=error_handling("anthropic", "model"),
|
||||
)
|
||||
errors=lambda: error_handling("anthropic", "model"),
|
||||
)(lambda: pytest.fail("upstream errors must not run Python"))()
|
||||
assert caught.value.status_code == (status or 500)
|
||||
assert caught.value.model == "model"
|
||||
assert caught.value.llm_provider == "anthropic"
|
||||
assert caught.value.__cause__ is error
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_registered_wrapper_preserves_arguments_and_request_error_context(asynchronous: bool) -> None:
|
||||
calls: Final[list[tuple[str, str, str]]] = []
|
||||
|
||||
def native(provider: str, *, model: str) -> DispatchResult[str]:
|
||||
calls.append(("native", provider, model))
|
||||
return (
|
||||
NativeFailed(Upstream(429, "limited"))
|
||||
if model == "limited"
|
||||
else NativeSkipped(NativeSkipReason.UNAVAILABLE)
|
||||
)
|
||||
|
||||
async def anative(provider: str, *, model: str) -> DispatchResult[str]:
|
||||
return native(provider, model=model)
|
||||
|
||||
def rules(provider: str, *, model: str):
|
||||
return error_handling(provider, model)
|
||||
|
||||
@native_first(native=native, route="chat_completions", errors=rules)
|
||||
def execute(provider: str, *, model: str) -> str:
|
||||
calls.append(("python", provider, model))
|
||||
return model
|
||||
|
||||
@anative_first(native=anative, route="chat_completions", errors=rules)
|
||||
async def aexecute(provider: str, *, model: str) -> str:
|
||||
calls.append(("python", provider, model))
|
||||
return model
|
||||
|
||||
assert (await aexecute("first", model="ok") if asynchronous else execute("first", model="ok")) == "ok"
|
||||
|
||||
async def fail() -> None:
|
||||
if asynchronous:
|
||||
await aexecute("second", model="limited")
|
||||
else:
|
||||
execute("second", model="limited")
|
||||
|
||||
with pytest.raises(APIError) as caught:
|
||||
await fail()
|
||||
assert caught.value.llm_provider == "second"
|
||||
assert caught.value.model == "limited"
|
||||
assert calls == [("native", "first", "ok"), ("python", "first", "ok"), ("native", "second", "limited")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("selection", ("native", "unavailable", "failed"))
|
||||
@pytest.mark.parametrize("failure", ("none", "body", "cleanup", "cancel"))
|
||||
async def test_context_selection_and_lifetime_are_separate(selection: str, failure: str) -> None:
|
||||
from collections.abc import AsyncGenerator
|
||||
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
||||
|
||||
from litellm.rust_bridge.dispatch import anative_context
|
||||
|
||||
events: Final[list[str]] = []
|
||||
error: Final = RuntimeError("connection use failed")
|
||||
|
||||
@asynccontextmanager
|
||||
async def connection(name: str) -> AsyncGenerator[str, None]:
|
||||
events.append(f"{name}:enter")
|
||||
try:
|
||||
yield name
|
||||
finally:
|
||||
events.append(f"{name}:exit")
|
||||
if failure == "cleanup":
|
||||
raise error
|
||||
|
||||
async def native() -> DispatchResult[AbstractAsyncContextManager[str]]:
|
||||
events.append("attempt")
|
||||
if selection == "failed":
|
||||
raise RuntimeError("connect failed")
|
||||
if selection == "unavailable":
|
||||
return NativeSkipped(NativeSkipReason.UNAVAILABLE)
|
||||
return Handled(connection("native"))
|
||||
|
||||
@anative_context(native=native, route="websocket", errors=lambda: PYTHON_ON_ERROR)
|
||||
def execute() -> AbstractAsyncContextManager[str]:
|
||||
events.append("python")
|
||||
return connection("python")
|
||||
|
||||
async def run() -> None:
|
||||
async with execute() as name:
|
||||
assert name == ("native" if selection == "native" else "python")
|
||||
if failure == "body":
|
||||
raise error
|
||||
if failure == "cancel":
|
||||
raise asyncio.CancelledError
|
||||
|
||||
if failure == "none":
|
||||
await run()
|
||||
elif failure == "cancel":
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await run()
|
||||
else:
|
||||
with pytest.raises(RuntimeError) as caught:
|
||||
await run()
|
||||
assert caught.value is error
|
||||
expected: Final = (
|
||||
["attempt", "native:enter", "native:exit"]
|
||||
if selection == "native"
|
||||
else ["attempt", "python", "python:enter", "python:exit"]
|
||||
)
|
||||
assert events == expected
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@
|
|||
"limit": 16419
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5506
|
||||
"limit": 5497
|
||||
},
|
||||
"LIT012": {
|
||||
"limit": 4486
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue