This commit is contained in:
Yujong Lee 2026-09-16 14:38:59 -07:00
parent 64cd6538a6
commit 62c862796a
5 changed files with 141 additions and 172 deletions

View file

@ -1,6 +1,4 @@
import json
from collections.abc import Mapping
from types import MappingProxyType
from typing import Any, Final
import httpx
@ -16,8 +14,6 @@ from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
get_async_httpx_client,
)
from litellm.rust_bridge.chat_completions import native as rust_chat_completions_bridge
from litellm.rust_bridge.chat_completions.native import rust_chat_completions_accepts
from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
@ -26,22 +22,6 @@ from ..common_utils import BedrockError, _get_all_bedrock_regions, error_respons
from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call
def _sigv4_principal(credentials: Credentials | None) -> Mapping[str, str]:
if credentials is None:
return MappingProxyType({})
return MappingProxyType(
{
key: value
for key, value in (
("aws_access_key_id", credentials.access_key),
("aws_secret_access_key", credentials.secret_key),
("aws_session_token", credentials.token),
)
if value is not None
}
)
def make_sync_call(
client: HTTPHandler | None,
api_base: str,
@ -401,87 +381,6 @@ class BedrockConverseLLM(BaseAWSLLM):
# Filter beta headers in HTTP headers before making the request
headers = update_headers_with_filtered_beta(headers=headers, provider="bedrock_converse")
# The Rust core owns the whole call for the subset it accepts. Ask
# before transforming so whichever path runs emits pre_call once, and
# hand down the credentials, region and endpoint this handler already
# resolved so both paths sign as the same principal. Bearer-token auth
# resolves no SigV4 principal at all, and each path reads that token
# itself.
rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy
**optional_params,
**_sigv4_principal(credentials),
"aws_region_name": aws_region_name,
}
serves_via_rust: Final = rust_chat_completions_accepts(
model=model,
messages=messages,
optional_params=rust_optional_params,
custom_llm_provider="bedrock",
litellm_params=litellm_params,
stream=stream,
)
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(
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,
)
if rust_response is not None:
return rust_response
### ROUTING (ASYNC, STREAMING, SYNC)
if acompletion:
if isinstance(client, HTTPHandler):
@ -548,21 +447,15 @@ class BedrockConverseLLM(BaseAWSLLM):
)
## 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,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": proxy_endpoint_url,
"headers": prepped.headers,
},
)
logging_obj.pre_call(
input=messages,
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:

View file

@ -1,4 +1,4 @@
"""Declarative Rust/Python selection matrix for every public LiteLLM route.
"""Declarative Rust/Python selection for routes with Rust integration.
Rules are static data matched top to bottom; the first match wins and a
context with no matching rule stays on Python. Whether the Rust core can serve
@ -20,13 +20,7 @@ class Route(str, Enum):
CHAT_COMPLETIONS = "chat_completions"
MESSAGES = "messages"
RESPONSES = "responses"
EMBEDDING = "embedding"
RERANK = "rerank"
IMAGE_GENERATION = "image_generation"
IMAGE_EDIT = "image_edit"
SPEECH = "speech"
TRANSCRIPTION = "transcription"
MODERATION = "moderation"
OCR = "ocr"
@ -66,16 +60,6 @@ Rules: TypeAlias = tuple[Rule, ...]
RULES: Final[Rules] = (
Rule(Route.OCR, Rollout.RUST_OPT_OUT),
Rule(Route.TRANSCRIPTION, Rollout.RUST_REQUIRED, providers=frozenset({"bedrock"})),
Rule(Route.TRANSCRIPTION, Rollout.PYTHON_ONLY),
Rule(Route.CHAT_COMPLETIONS, Rollout.PYTHON_ONLY),
Rule(Route.MESSAGES, Rollout.PYTHON_ONLY),
Rule(Route.RESPONSES, Rollout.PYTHON_ONLY),
Rule(Route.EMBEDDING, Rollout.PYTHON_ONLY),
Rule(Route.RERANK, Rollout.PYTHON_ONLY),
Rule(Route.IMAGE_GENERATION, Rollout.PYTHON_ONLY),
Rule(Route.IMAGE_EDIT, Rollout.PYTHON_ONLY),
Rule(Route.SPEECH, Rollout.PYTHON_ONLY),
Rule(Route.MODERATION, Rollout.PYTHON_ONLY),
)

View file

@ -9,6 +9,7 @@ from __future__ import annotations
import asyncio
from datetime import datetime, timedelta, timezone
from typing import Final
from unittest.mock import MagicMock, patch
import boto3
@ -93,7 +94,11 @@ CONVERSE_RESPONSE = {
async def _drive_async_completion(
*, skip_pre_call_logging: bool, logging_obj, credentials: Credentials = RESOLVED_CREDENTIALS
*,
skip_pre_call_logging: bool,
logging_obj,
credentials: Credentials = RESOLVED_CREDENTIALS,
outer_dispatch: bool = False,
):
"""Run the real `async_completion` with a stubbed transport."""
import httpx as _httpx
@ -110,6 +115,9 @@ async def _drive_async_completion(
client.post = post
client.__class__ = AsyncHTTPHandler
if outer_dispatch:
return await _run(credentials=credentials, acompletion=True, client=client, logging_obj=logging_obj)
return await BedrockConverseLLM().async_completion(
model="anthropic.claude-sonnet-4-5-v1:0",
messages=[{"role": "user", "content": "hi"}],
@ -160,6 +168,26 @@ async def test_async_completion_signs_off_the_event_loop(monkeypatch):
assert probe.served_during_refresh is True
@pytest.mark.asyncio
@pytest.mark.parametrize("rust_enabled", (False, True))
async def test_python_only_async_dispatch_refreshes_credentials_off_the_event_loop(
monkeypatch: pytest.MonkeyPatch, rust_enabled: bool
) -> None:
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
monkeypatch.setenv("LITELLM_RUST", "1" if rust_enabled else "0")
configuration.rust(rust_enabled)
probe: Final = EventLoopProbe()
release: Final = asyncio.create_task(probe.release_refresh_from_the_loop())
response: Final = await _drive_async_completion(
skip_pre_call_logging=False, logging_obj=MagicMock(), credentials=probe.credentials(), outer_dispatch=True
)
await release
assert response.choices[0].message.content == "hi"
assert probe.served_during_refresh is True
def _sync_client_returning_converse_response():
client = MagicMock()
client.post.side_effect = lambda **_kwargs: httpx.Response(

View file

@ -1,55 +1,83 @@
from __future__ import annotations
from collections.abc import Generator
from typing import Final
import pytest
from litellm.rust_bridge import catalog
from litellm.rust_bridge import catalog, configuration
from litellm.rust_bridge.catalog import Context, Delivery, Route, Rule
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.configuration import Decision, Rollout
def test_every_route_has_an_explicit_default_rule() -> None:
declared: Final = frozenset(
rule.route for rule in catalog.RULES if rule.providers is None and rule.deliveries is None
)
assert declared == frozenset(Route)
@pytest.fixture(autouse=True)
def isolated_configuration(monkeypatch: pytest.MonkeyPatch) -> Generator[None]:
monkeypatch.delenv("LITELLM_RUST", raising=False)
configuration.reset_rust_configuration()
yield
configuration.reset_rust_configuration()
@pytest.mark.parametrize("route", tuple(Route))
@pytest.mark.parametrize("provider", (None, "bedrock", "mistral", "anthropic", "openai", "azure_ai", "unknown"))
@pytest.mark.parametrize("delivery", tuple(Delivery))
@pytest.mark.parametrize("process", (None, False, True))
@pytest.mark.parametrize("environment", (None, "0", "1"))
def test_shipped_decisions(
monkeypatch: pytest.MonkeyPatch,
route: Route,
provider: str | None,
delivery: Delivery,
process: bool | None,
environment: str | None,
) -> None:
configuration.rust(process)
if environment is not None:
monkeypatch.setenv("LITELLM_RUST", environment)
context: Final = Context(route, provider=provider, model="test-model", delivery=delivery)
if route is Route.OCR:
enabled: Final = environment == "1" if environment is not None else process is not False
assert catalog.rollout(context) is Rollout.RUST_OPT_OUT
assert catalog.decision(context) is (Decision.RUST_WITH_FALLBACK if enabled else Decision.PYTHON)
elif route is Route.TRANSCRIPTION and provider == "bedrock":
assert catalog.rollout(context) is Rollout.RUST_REQUIRED
assert catalog.decision(context) is Decision.RUST_REQUIRED
else:
assert catalog.rollout(context) is Rollout.PYTHON_ONLY
assert catalog.decision(context) is Decision.PYTHON
@pytest.mark.parametrize("route", tuple(Route))
def test_missing_rule_stays_on_python_even_when_rust_is_enabled(monkeypatch: pytest.MonkeyPatch, route: Route) -> None:
configuration.rust(True)
monkeypatch.setenv("LITELLM_RUST", "1")
assert catalog.rollout(Context(route), rules=()) is Rollout.PYTHON_ONLY
assert catalog.decision(Context(route), rules=()) is Decision.PYTHON
@pytest.mark.parametrize(
("context", "expected"),
(
(Context(Route.OCR), Rollout.RUST_OPT_OUT),
(Context(Route.OCR, provider="mistral", model="mistral-ocr-latest"), Rollout.RUST_OPT_OUT),
(Context(Route.TRANSCRIPTION, provider="bedrock"), Rollout.RUST_REQUIRED),
(Context(Route.TRANSCRIPTION, provider="openai"), Rollout.PYTHON_ONLY),
(Context(Route.TRANSCRIPTION), Rollout.PYTHON_ONLY),
(Context(Route.CHAT_COMPLETIONS, provider="anthropic"), Rollout.PYTHON_ONLY),
(Context(Route.CHAT_COMPLETIONS, provider="bedrock"), Rollout.PYTHON_ONLY),
(Context(Route.MESSAGES, provider="anthropic"), Rollout.PYTHON_ONLY),
(Context(Route.MESSAGES, provider="azure_ai"), Rollout.PYTHON_ONLY),
(Context(Route.RESPONSES, provider="openai", delivery=Delivery.WEBSOCKET), Rollout.PYTHON_ONLY),
(Context(Route.RESPONSES, provider="openai", model="m", delivery=Delivery.WEBSOCKET), Decision.RUST_REQUIRED),
(Context(Route.RESPONSES, provider="openai", model="m"), Decision.PYTHON),
(Context(Route.RESPONSES, provider="openai", model="m", delivery=Delivery.STREAMING), Decision.PYTHON),
(Context(Route.RESPONSES, provider="openai", model="other", delivery=Delivery.WEBSOCKET), Decision.PYTHON),
(Context(Route.RESPONSES, provider="anthropic", model="m", delivery=Delivery.WEBSOCKET), Decision.PYTHON),
(Context(Route.MESSAGES, provider="openai", model="m", delivery=Delivery.WEBSOCKET), Decision.PYTHON),
),
)
def test_shipped_rules(context: Context, expected: Rollout) -> None:
assert catalog.rollout(context) is expected
def test_only_ocr_and_bedrock_transcription_can_reach_rust() -> None:
rust_capable: Final = frozenset(
(rule.route, rule.providers) for rule in catalog.RULES if rule.rollout is not Rollout.PYTHON_ONLY
)
assert rust_capable == frozenset({(Route.OCR, None), (Route.TRANSCRIPTION, frozenset({"bedrock"}))})
def test_first_matching_rule_wins() -> None:
def test_first_matching_rule_respects_every_constraint(context: Context, expected: Decision) -> None:
rules: Final = (
Rule(Route.EMBEDDING, Rollout.RUST_REQUIRED, providers=frozenset({"openai"}), models=frozenset({"m"})),
Rule(Route.EMBEDDING, Rollout.RUST_OPT_IN, providers=frozenset({"openai"})),
Rule(Route.EMBEDDING, Rollout.PYTHON_ONLY),
Rule(
Route.RESPONSES,
Rollout.RUST_REQUIRED,
providers=frozenset({"openai"}),
models=frozenset({"m"}),
deliveries=frozenset({Delivery.WEBSOCKET}),
),
Rule(Route.RESPONSES, Rollout.PYTHON_ONLY),
)
assert catalog.rollout(Context(Route.EMBEDDING, provider="openai", model="m"), rules) is Rollout.RUST_REQUIRED
assert catalog.rollout(Context(Route.EMBEDDING, provider="openai", model="other"), rules) is Rollout.RUST_OPT_IN
assert catalog.rollout(Context(Route.EMBEDDING, provider="cohere", model="m"), rules) is Rollout.PYTHON_ONLY
assert catalog.rollout(Context(Route.RERANK, provider="openai", model="m"), rules) is Rollout.PYTHON_ONLY
assert catalog.decision(context, rules) is expected

View file

@ -8,7 +8,7 @@ import pytest
from litellm.exceptions import APIError
from litellm.rust_bridge import bindings, configuration, runtime
from litellm.rust_bridge.catalog import Context, Route, Rule
from litellm.rust_bridge.catalog import Context, Delivery, Route, Rule
from litellm.rust_bridge.configuration import Rollout
@ -145,7 +145,43 @@ def test_context_outside_rule_stays_on_python() -> None:
configuration.rust(True)
assert run(Rollout.RUST_REQUIRED, calls, context=Context(Route.MESSAGES, provider="openai")) == "python"
assert run(Rollout.RUST_REQUIRED, calls, context=Context(Route.EMBEDDING, provider="anthropic")) == "python"
assert run(Rollout.RUST_REQUIRED, calls, context=Context(Route.RESPONSES, provider="anthropic")) == "python"
assert calls.calls == (PYTHON, PYTHON)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"context",
(
Context(Route.CHAT_COMPLETIONS, provider="anthropic"),
Context(Route.CHAT_COMPLETIONS, provider="bedrock"),
Context(Route.MESSAGES, provider="anthropic"),
Context(Route.RESPONSES, provider="openai"),
Context(Route.TRANSCRIPTION, provider="openai"),
),
)
@pytest.mark.parametrize("delivery", tuple(Delivery))
async def test_shipped_python_routes_never_load_native(
monkeypatch: pytest.MonkeyPatch, context: Context, delivery: Delivery
) -> None:
monkeypatch.setenv("LITELLM_RUST", "1")
configuration.rust(True)
calls: Final = recorder()
request: Final = Context(context.route, provider=context.provider, delivery=delivery)
def reject_load(value: object) -> NativeFn | None:
pytest.fail("Python-only dispatch must not load a native binding")
bound: Final = bindings.NativeBinding("_messages", validate=reject_load)
async def native(fn: NativeFn) -> str:
return fn()
async def python() -> str:
return calls.python()
assert runtime.run(request, binding=bound, native=lambda fn: fn(), python=calls.python) == PYTHON
assert await runtime.arun(request, binding=bound, native=native, python=python) == PYTHON
assert calls.calls == (PYTHON, PYTHON)