litellm/tests/test_litellm_rust/test_inference.py
devin-ai-integration[bot] e4190d86a6
refactor(rust): centralize host execution and compose callbacks (#43515)
* 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>
2026-09-28 19:20:27 +00:00

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