mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
bring x-litellm-rust
This commit is contained in:
parent
6d30006592
commit
c01db259d7
3 changed files with 82 additions and 3 deletions
12
litellm/rust_bridge/response_metadata.py
Normal file
12
litellm/rust_bridge/response_metadata.py
Normal file
|
|
@ -0,0 +1,12 @@
|
|||
from typing import TypeVar
|
||||
|
||||
from litellm.router_utils.add_retry_fallback_headers import (
|
||||
_add_headers_to_response, # pyright: ignore[reportPrivateUsage] # reuse the proxy's identity-preserving response metadata writer
|
||||
)
|
||||
|
||||
ResultT = TypeVar("ResultT")
|
||||
|
||||
|
||||
def mark_rust_response(response: ResultT) -> ResultT:
|
||||
_add_headers_to_response(response, {"x-litellm-rust": "true"})
|
||||
return response
|
||||
|
|
@ -10,6 +10,7 @@ from litellm.exceptions import APIError
|
|||
from litellm.rust_bridge.bindings import NativeBinding, native_exception_types
|
||||
from litellm.rust_bridge.catalog import RULES, Context, Rules, decision
|
||||
from litellm.rust_bridge.configuration import Decision
|
||||
from litellm.rust_bridge.response_metadata import mark_rust_response
|
||||
|
||||
NativeT = TypeVar("NativeT")
|
||||
ResultT = TypeVar("ResultT")
|
||||
|
|
@ -60,7 +61,7 @@ def run(
|
|||
context=_error_context(context),
|
||||
)
|
||||
if isinstance(result, RustHandled):
|
||||
return result.value
|
||||
return mark_rust_response(result.value)
|
||||
if selected is Decision.RUST_REQUIRED:
|
||||
_raise_required(result, _error_context(context))
|
||||
return python()
|
||||
|
|
@ -88,7 +89,7 @@ async def arun(
|
|||
context=_error_context(context),
|
||||
)
|
||||
if isinstance(result, RustHandled):
|
||||
return result.value
|
||||
return mark_rust_response(result.value)
|
||||
if selected is Decision.RUST_REQUIRED:
|
||||
_raise_required(result, _error_context(context))
|
||||
return await python()
|
||||
|
|
|
|||
|
|
@ -1,12 +1,14 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Generator
|
||||
from collections.abc import Callable, Generator
|
||||
from types import SimpleNamespace
|
||||
from typing import Final, Protocol
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
|
||||
from litellm.rust_bridge import bindings, configuration, runtime
|
||||
from litellm.rust_bridge.catalog import Context, Delivery, Route, Rule
|
||||
from litellm.rust_bridge.configuration import Rollout
|
||||
|
|
@ -199,6 +201,70 @@ def test_unavailable_native_falls_back_to_python() -> None:
|
|||
assert calls.calls == (PYTHON,)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("missing", (False, True))
|
||||
async def test_python_fallback_does_not_claim_rust_execution(missing: bool) -> None:
|
||||
calls: Final = recorder(RustBridgeDeclined("unsupported"))
|
||||
bound: Final = binding(None if missing else calls.rust)
|
||||
expected: Final = OCRResponse(pages=[], model="python")
|
||||
|
||||
def native(fn: NativeFn) -> OCRResponse:
|
||||
fn()
|
||||
pytest.fail("native must decline before constructing a response")
|
||||
|
||||
async def anative(fn: NativeFn) -> OCRResponse:
|
||||
return native(fn)
|
||||
|
||||
async def python() -> OCRResponse:
|
||||
return expected
|
||||
|
||||
assert (
|
||||
runtime.run(CONTEXT, binding=bound, native=native, python=lambda: expected, rules=rules(Rollout.RUST_OPT_OUT))
|
||||
is expected
|
||||
)
|
||||
assert (
|
||||
await runtime.arun(CONTEXT, binding=bound, native=anative, python=python, rules=rules(Rollout.RUST_OPT_OUT))
|
||||
is expected
|
||||
)
|
||||
assert get_hidden_params_dict(expected) == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("shape", ("model", "dict"))
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_native_response_marker_reaches_caller_with_existing_metadata(shape: str, asynchronous: bool) -> None:
|
||||
hidden: Final = {"additional_headers": {"x-request-id": "upstream"}, "response_cost": 0.01}
|
||||
response: Final[OCRResponse | dict[str, object]] = (
|
||||
OCRResponse(pages=[], model="native") if shape == "model" else {"content": "native", "_hidden_params": hidden}
|
||||
)
|
||||
if isinstance(response, OCRResponse):
|
||||
response._hidden_params = hidden # pyright: ignore[reportPrivateUsage] # seed SDK metadata to verify it survives native marking
|
||||
bound: Final[bindings.NativeBinding[Callable[[], object]]] = bindings.NativeBinding("ocr", validate=lambda _: None)
|
||||
bound.override(lambda: response)
|
||||
|
||||
def python() -> object:
|
||||
pytest.fail("native success must not fall back")
|
||||
|
||||
async def anative(fn: Callable[[], object]) -> object:
|
||||
return fn()
|
||||
|
||||
async def apython() -> object:
|
||||
return python()
|
||||
|
||||
result: Final = (
|
||||
await runtime.arun(CONTEXT, binding=bound, native=anative, python=apython, rules=rules(Rollout.RUST_REQUIRED))
|
||||
if asynchronous
|
||||
else runtime.run(
|
||||
CONTEXT, binding=bound, native=lambda fn: fn(), python=python, rules=rules(Rollout.RUST_REQUIRED)
|
||||
)
|
||||
)
|
||||
assert result is response
|
||||
assert get_hidden_params_dict(result) == {
|
||||
"response_cost": 0.01,
|
||||
"additional_headers": {"x-request-id": "upstream", "x-litellm-rust": "true"},
|
||||
}
|
||||
|
||||
|
||||
def test_upstream_error_maps_to_api_error_without_fallback() -> None:
|
||||
calls: Final = recorder(RustUpstreamError(429, "rate limited"))
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue