mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
* refactor(rust): extract litellm-host-native as the shared Rust host driver Move service and hook dispatch out of host-http into a Driver that owns the machine and Rust handlers, returning at completion or a stream boundary and holding the demand reply until the consumer advances. Move the in-process runner onto the same driver. host-http now layers encoding, SSE, body polling and lifecycle observation over it. host-python keeps driving litellm-host directly Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(rust): interrupt the machine when the in-process stream consumer fails Restores the pre-refactor interruption path for StreamConsumer errors via Driver::fail and ports the generic run lifecycle tests into host-native. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(rust): separate the machine contract from coroutine execution * auth update * refactor(rust): use standard flow control for host requests * style(rust): keep host driver imports formatted * chores * mostly relocation * refactor(rust): separate interceptors from queued observers * refactor(rust): centralize legacy callback mappings and lifecycle * docs: define Python host boundaries and migration plan * refactor: enforce Python host and bridge boundaries * refactor(rust): separate operations from callback composition * refactor(rust): compose SDK policy through call hooks --------- Co-authored-by: Yujong Lee <yujong@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
353 lines
15 KiB
Python
353 lines
15 KiB
Python
import asyncio
|
|
from collections.abc import Awaitable, Coroutine, Mapping
|
|
from typing import Final, Literal, TypeAlias
|
|
|
|
import pytest
|
|
from pydantic import JsonValue, TypeAdapter
|
|
|
|
import litellm
|
|
from litellm import RateLimitError
|
|
from litellm.integrations.custom_logger import CustomLogger
|
|
from litellm.models.credentials import CredentialItem
|
|
from litellm.responses.utils import ResponsesAPIRequestUtils
|
|
from litellm.rust_bridge import _native
|
|
from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest
|
|
from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest
|
|
from litellm.types.llms.openai import ResponsesAPIResponse
|
|
from litellm.types.utils import CallTypes, ModelResponse
|
|
from tests.test_litellm_rust.support.callback_recorder import RecordingLogger
|
|
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
|
|
from tests.test_litellm_rust.support.requests import MESSAGES, MESSAGES_MODEL, MESSAGES_RESPONSE, request_body
|
|
|
|
pytestmark = pytest.mark.requires_rust_extension
|
|
Route: TypeAlias = Literal["chat", "responses"]
|
|
_OBJECT: Final = TypeAdapter(dict[str, object])
|
|
NativeResult: TypeAlias = (
|
|
ModelResponse
|
|
| ResponsesAPIResponse
|
|
| Coroutine[object, object, ModelResponse]
|
|
| Coroutine[object, object, ResponsesAPIResponse]
|
|
)
|
|
RESPONSES_MODEL: Final = "openai/gpt-6-sol"
|
|
RESPONSES_RESPONSE: Final[dict[str, JsonValue]] = {
|
|
"id": "resp_native",
|
|
"object": "response",
|
|
"created_at": 1,
|
|
"model": RESPONSES_MODEL.removeprefix("openai/"),
|
|
"status": "completed",
|
|
"output": [
|
|
{
|
|
"type": "message",
|
|
"id": "msg_native",
|
|
"role": "assistant",
|
|
"status": "completed",
|
|
"content": [{"type": "output_text", "text": "native response", "annotations": []}],
|
|
}
|
|
],
|
|
"usage": {"input_tokens": 5, "output_tokens": 4, "total_tokens": 9},
|
|
}
|
|
|
|
|
|
@pytest.fixture(params=("chat", "responses"))
|
|
def route(request: pytest.FixtureRequest) -> Route:
|
|
return TypeAdapter(Route).validate_python(request.param)
|
|
|
|
|
|
def native_call(
|
|
route: Route, asynchronous: bool, server: RecordingServer, options: Mapping[str, object]
|
|
) -> NativeResult:
|
|
server.default_response = ResponseSpec(body=MESSAGES_RESPONSE if route == "chat" else RESPONSES_RESPONSE)
|
|
if route == "chat":
|
|
kwargs: Final = {
|
|
"model": MESSAGES_MODEL,
|
|
"messages": list(MESSAGES),
|
|
"api_key": "test-key",
|
|
"api_base": server.base_url,
|
|
"max_tokens": 32,
|
|
**options,
|
|
}
|
|
request: Final = LiteLLMChatCompletionsRequest(
|
|
MESSAGES_MODEL, list(MESSAGES), None, "test-key", server.base_url, None, None, kwargs
|
|
)
|
|
return (_native.acompletion if asynchronous else _native.completion)(request, (), kwargs)
|
|
response_kwargs: Final = {
|
|
"model": RESPONSES_MODEL,
|
|
"input": "hello",
|
|
"api_key": "test-key",
|
|
"api_base": server.base_url,
|
|
"max_output_tokens": 32,
|
|
**options,
|
|
}
|
|
response_request: Final = LiteLLMResponsesRequest(
|
|
RESPONSES_MODEL, "hello", None, "test-key", server.base_url, "openai", None, response_kwargs
|
|
)
|
|
return (_native.aresponses if asynchronous else _native.responses)(response_request, (), response_kwargs)
|
|
|
|
|
|
async def execute(route: Route, asynchronous: bool, server: RecordingServer, options: Mapping[str, object]) -> object:
|
|
if not asynchronous:
|
|
return await asyncio.to_thread(native_call, route, False, server, options)
|
|
result: Final = native_call(route, True, server, options)
|
|
assert isinstance(result, Awaitable)
|
|
return await result
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("asynchronous", (False, True))
|
|
async def test_native_inference_returns_public_models_and_logs_once(
|
|
route: Route,
|
|
asynchronous: bool,
|
|
recording_server: RecordingServer,
|
|
) -> None:
|
|
recorder: Final = RecordingLogger()
|
|
result: Final = await execute(route, asynchronous, recording_server, {"callbacks": [recorder], "temperature": 0.25})
|
|
assert len(recording_server.requests) == 1
|
|
sent: Final = recording_server.requests[0]
|
|
body: Final = _OBJECT.validate_python(sent.body)
|
|
assert body["temperature"] == 0.25
|
|
assert request_body(_OBJECT.validate_python(recorder.wait_for("log_pre_api_call")[0].kwargs)) == body
|
|
if route == "chat":
|
|
assert isinstance(result, ModelResponse)
|
|
assert result.choices[0].message.content == "Hello from native Messages"
|
|
assert sent.path == "/v1/messages"
|
|
else:
|
|
assert isinstance(result, ResponsesAPIResponse)
|
|
assert result.output_text == "native response"
|
|
assert sent.path == "/responses"
|
|
success: Final = await recorder.wait_for_async("async_log_success_event" if asynchronous else "log_success_event")
|
|
assert len(success) == 1
|
|
if isinstance(result, ModelResponse):
|
|
assert success[0].response is result
|
|
else:
|
|
logged: Final = success[0].response
|
|
assert isinstance(logged, ResponsesAPIResponse)
|
|
assert isinstance(result, ResponsesAPIResponse)
|
|
assert logged.id == result.id
|
|
assert logged.output_text == result.output_text
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_native_inference_pre_call_edits_reach_the_provider(
|
|
route: Route, recording_server: RecordingServer
|
|
) -> None:
|
|
class Edit(CustomLogger):
|
|
def log_pre_api_call(self, model: object, messages: object, kwargs: dict[str, object]) -> None:
|
|
request_body(kwargs)["temperature"] = 0.75
|
|
|
|
await execute(route, True, recording_server, {"callbacks": [Edit()]})
|
|
assert _OBJECT.validate_python(recording_server.requests[0].body)["temperature"] == 0.75
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("from_credentials", (False, True))
|
|
async def test_native_resource_setup_uses_deployment_hook_arguments(
|
|
route: Route,
|
|
recording_server: RecordingServer,
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
from_credentials: bool,
|
|
) -> None:
|
|
invalid_settings: Final = {"ssl_verify": object()}
|
|
credential: Final = CredentialItem(
|
|
credential_name="resource-settings", credential_info={}, credential_values=invalid_settings
|
|
)
|
|
monkeypatch.setattr(litellm, "credential_list", [credential])
|
|
|
|
class Prepare(CustomLogger):
|
|
async def async_pre_call_deployment_hook(
|
|
self, kwargs: dict[str, object], call_type: CallTypes | None
|
|
) -> dict[str, object]:
|
|
return {
|
|
**kwargs,
|
|
**({"litellm_credential_name": credential.credential_name} if from_credentials else invalid_settings),
|
|
}
|
|
|
|
litellm.callbacks.append(Prepare())
|
|
recorder: Final = RecordingLogger()
|
|
recording_server.expected_requests = 0
|
|
with pytest.raises(ValueError, match=r"request\.ssl_verify") as caught:
|
|
await execute(route, True, recording_server, {"callbacks": [recorder]})
|
|
failure: Final = await recorder.wait_for_async("async_log_failure_event")
|
|
assert len(failure) == 1
|
|
assert _OBJECT.validate_python(failure[0].kwargs)["exception"] is caught.value
|
|
assert not recording_server.requests
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("asynchronous", (False, True))
|
|
async def test_sdk_policy_rejection_precedes_resource_setup_and_is_logged_once(
|
|
route: Route,
|
|
asynchronous: bool,
|
|
recording_server: RecordingServer,
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
) -> None:
|
|
monkeypatch.setattr(litellm, "max_budget", 1.0)
|
|
monkeypatch.setattr(litellm, "_current_cost", 2.0)
|
|
monkeypatch.setattr(litellm, "ssl_verify", object())
|
|
recorder: Final = RecordingLogger()
|
|
recording_server.expected_requests = 0
|
|
with pytest.raises(litellm.BudgetExceededError) as caught:
|
|
await execute(route, asynchronous, recording_server, {"callbacks": [recorder]})
|
|
failure: Final = await recorder.wait_for_async("async_log_failure_event" if asynchronous else "log_failure_event")
|
|
assert len(failure) == 1
|
|
assert _OBJECT.validate_python(failure[0].kwargs)["exception"] is caught.value
|
|
assert not recording_server.requests
|
|
assert not any("success" in name for name in recorder.names)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("asynchronous", (False, True))
|
|
async def test_native_inference_provider_failure_is_terminal_and_shared_with_callbacks(
|
|
route: Route,
|
|
asynchronous: bool,
|
|
recording_server: RecordingServer,
|
|
) -> None:
|
|
recorder: Final = RecordingLogger()
|
|
recording_server.enqueue(
|
|
ResponseSpec(body={"error": {"message": "slow down", "type": "rate_limit_error"}}, status=429)
|
|
)
|
|
with pytest.raises(RateLimitError) as caught:
|
|
await execute(route, asynchronous, recording_server, {"callbacks": [recorder]})
|
|
assert getattr(caught.value, "status_code", None) == 429
|
|
assert len(recording_server.requests) == 1
|
|
failure: Final = await recorder.wait_for_async("async_log_failure_event" if asynchronous else "log_failure_event")
|
|
assert len(failure) == 1
|
|
assert _OBJECT.validate_python(failure[0].kwargs)["exception"] is caught.value
|
|
assert not any("success" in name for name in recorder.names)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_unstarted_native_inference_has_no_provider_or_callback_effects(
|
|
route: Route,
|
|
recording_server: RecordingServer,
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
) -> None:
|
|
recorder: Final = RecordingLogger()
|
|
recording_server.expected_requests = 0
|
|
monkeypatch.setattr(litellm, "ssl_verify", object())
|
|
monkeypatch.setattr(litellm, "max_budget", 1.0)
|
|
monkeypatch.setattr(litellm, "_current_cost", 2.0)
|
|
pending: Final = native_call(route, True, recording_server, {"callbacks": [recorder]})
|
|
assert asyncio.iscoroutine(pending)
|
|
pending.close()
|
|
assert not recording_server.requests
|
|
assert not recorder.events
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize(
|
|
"options",
|
|
(
|
|
{"stream": True},
|
|
{"extra_body": {"provider_option": True}},
|
|
{"mock_response": "mock"},
|
|
{"num_retries": 1},
|
|
{"use_chat_completions_api": True},
|
|
{"model_list": []},
|
|
),
|
|
)
|
|
async def test_native_responses_declines_unsupported_requests_before_callbacks(
|
|
recording_server: RecordingServer,
|
|
options: Mapping[str, object],
|
|
) -> None:
|
|
recorder: Final = RecordingLogger()
|
|
recording_server.expected_requests = 0
|
|
with pytest.raises(_native.RustBridgeDeclined):
|
|
native_call("responses", True, recording_server, {**options, "callbacks": [recorder]})
|
|
assert not recording_server.requests
|
|
assert not recorder.events
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("asynchronous", (False, True))
|
|
async def test_native_chat_validation_failure_is_terminal(
|
|
asynchronous: bool, recording_server: RecordingServer
|
|
) -> None:
|
|
recording_server.expected_requests = 0
|
|
with pytest.raises(Exception, match="chat completions requires at least one message") as failure:
|
|
await execute("chat", asynchronous, recording_server, {"messages": []})
|
|
assert not isinstance(failure.value, _native.RustBridgeDeclined)
|
|
assert not recording_server.requests
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_native_projection_reads_positional_parameters(route: Route, recording_server: RecordingServer) -> None:
|
|
from litellm.chat_completions.dispatch import (
|
|
_DISPATCH as chat_dispatch, # pyright: ignore[reportPrivateUsage] # exercise the request passed to the native boundary
|
|
)
|
|
from litellm.responses.dispatch import (
|
|
_DISPATCH as responses_dispatch, # pyright: ignore[reportPrivateUsage] # exercise the request passed to the native boundary
|
|
)
|
|
|
|
recording_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE if route == "chat" else RESPONSES_RESPONSE)
|
|
kwargs: Final = {"api_key": "test-key", "api_base": recording_server.base_url}
|
|
if route == "chat":
|
|
args: Final = (MESSAGES_MODEL, list(MESSAGES), 12.0, 0.25)
|
|
request: Final = chat_dispatch.request(args, kwargs)
|
|
assert request is not None
|
|
await asyncio.to_thread(_native.completion, request, args, kwargs)
|
|
assert _OBJECT.validate_python(recording_server.requests[0].body)["temperature"] == 0.25
|
|
else:
|
|
response_args: Final = ("hello", RESPONSES_MODEL, None, "Be brief", 16)
|
|
response_request: Final = responses_dispatch.request(response_args, kwargs)
|
|
assert response_request is not None
|
|
await asyncio.to_thread(_native.responses, response_request, response_args, kwargs)
|
|
body: Final = _OBJECT.validate_python(recording_server.requests[0].body)
|
|
assert body["instructions"] == "Be brief"
|
|
assert body["max_output_tokens"] == 16
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("asynchronous", (False, True))
|
|
@pytest.mark.parametrize("source", ("explicit", "base_url", "global", "provider", "environment", "empty"))
|
|
async def test_native_connection_settings_reach_the_provider(
|
|
route: Route,
|
|
asynchronous: bool,
|
|
source: str,
|
|
recording_server: RecordingServer,
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
) -> None:
|
|
key: Final = "selected-key"
|
|
explicit: Final = source in ("explicit", "base_url")
|
|
monkeypatch.setattr(litellm, "api_key", key if source in ("global", "empty") else ("unused" if explicit else None))
|
|
monkeypatch.setattr(litellm, "openai_key", key if source == "provider" else ("unused" if explicit else None))
|
|
monkeypatch.setattr(litellm, "anthropic_key", key if source == "provider" else ("unused" if explicit else None))
|
|
monkeypatch.setattr(
|
|
litellm,
|
|
"api_base",
|
|
None if source == "environment" else ("http://127.0.0.1:1" if explicit else recording_server.base_url),
|
|
)
|
|
monkeypatch.setenv(
|
|
"OPENAI_API_KEY" if route == "responses" else "ANTHROPIC_API_KEY", key if source == "environment" else "unused"
|
|
)
|
|
for name in ("OPENAI_BASE_URL", "OPENAI_API_BASE", "ANTHROPIC_BASE_URL", "ANTHROPIC_API_BASE"):
|
|
monkeypatch.setenv(name, recording_server.base_url if source == "environment" else "http://127.0.0.1:1")
|
|
result: Final = await execute(
|
|
route,
|
|
asynchronous,
|
|
recording_server,
|
|
{
|
|
"api_key": key if explicit else ("" if source == "empty" else None),
|
|
"api_base": recording_server.base_url if source == "explicit" else ("" if source == "empty" else None),
|
|
**({"base_url": recording_server.base_url} if source == "base_url" else {}),
|
|
},
|
|
)
|
|
assert isinstance(result, ModelResponse | ResponsesAPIResponse)
|
|
assert len(recording_server.requests) == 1
|
|
headers: Final = recording_server.requests[0].headers
|
|
assert headers["x-api-key" if route == "chat" else "authorization"] == (key if route == "chat" else f"Bearer {key}")
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("asynchronous", (False, True))
|
|
@pytest.mark.parametrize("encoded", (False, True))
|
|
async def test_native_responses_decode_continuation_ids(
|
|
asynchronous: bool, encoded: bool, recording_server: RecordingServer
|
|
) -> None:
|
|
original: Final = "resp_upstream"
|
|
previous: Final = (
|
|
ResponsesAPIRequestUtils._build_responses_api_response_id("openai", "deployment", original)
|
|
if encoded
|
|
else original
|
|
)
|
|
await execute("responses", asynchronous, recording_server, {"previous_response_id": previous})
|
|
assert _OBJECT.validate_python(recording_server.requests[0].body)["previous_response_id"] == original
|