mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
cleanup
This commit is contained in:
parent
64cd6538a6
commit
62c862796a
5 changed files with 141 additions and 172 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue