litellm/tests/rust-python-harness/strategies/trace_parity/gateway/execution.py
2026-09-14 14:01:28 -07:00

197 lines
7.4 KiB
Python

from __future__ import annotations
import json
import subprocess
from functools import cache
from pathlib import Path
from typing import Final, Protocol, cast
import httpx
from pydantic import BaseModel, ConfigDict
from ....shared.parity.replay import replay_server
from ....shared.tracing.native import TraceResponsePayload, native_trace_events
from ....shared.tracing.profiler import FunctionTraceEvent, profile_python
from ....shared.tracing.steps import Engine, PipelineProjection, pipeline_projection
from ..models import GatewayRouteSpec, RouteFixture, TraceEngine, TraceExecutionFailure, TraceScenario
from ..reporting import TraceArtifact
class _GatewayResponsePayload(BaseModel):
model_config = ConfigDict(strict=True, extra="forbid")
status: int
body: object
class _GatewayClient(Protocol):
def post(self, url: str, *, json: object, headers: dict[str, str]) -> httpx.Response: ...
_ROUTE_PATHS: Final = {
"messages": "/v1/messages",
"chat_completions": "/v1/chat/completions",
"responses": "/v1/responses",
}
def _collect_python(fixture: RouteFixture, route: GatewayRouteSpec) -> tuple[FunctionTraceEvent, ...]:
from fastapi.testclient import TestClient
import litellm
from litellm.proxy import proxy_server
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.anthropic_endpoints.endpoints import user_api_key_auth
provider_model: Final = cast(str, fixture.kwargs["provider_model"])
model_alias: Final = cast(str, fixture.kwargs["model_alias"])
old_router: Final = proxy_server.llm_router
old_override: Final = proxy_server.app.dependency_overrides.get(user_api_key_auth)
async def authorize() -> UserAPIKeyAuth:
return UserAPIKeyAuth(api_key="trace-key")
proxy_server.llm_router = litellm.Router(
model_list=[
{
"model_name": model_alias,
"litellm_params": {
"model": provider_model,
"api_key": "trace-provider-key",
"api_base": fixture.kwargs["api_base"],
},
}
]
)
proxy_server.app.dependency_overrides[user_api_key_auth] = authorize
try:
with profile_python(Path(litellm.__file__).parent, threads=True) as profiler:
client: Final = cast(_GatewayClient, TestClient(proxy_server.app))
response: Final = client.post(
_ROUTE_PATHS[route.route],
json=fixture.kwargs["body"],
headers={"authorization": "Bearer trace-key"},
)
if response.status_code != 200:
raise RuntimeError(f"Python gateway returned {response.status_code}: {response.text}")
return tuple(profiler.events)
finally:
proxy_server.llm_router = old_router
if old_override is None:
proxy_server.app.dependency_overrides.pop(user_api_key_auth, None)
else:
proxy_server.app.dependency_overrides[user_api_key_auth] = old_override
def _collect_rust(fixture: RouteFixture, route: GatewayRouteSpec) -> tuple[FunctionTraceEvent, ...]:
payload: Final = json.dumps(
{
"path": _ROUTE_PATHS[route.route],
"model_alias": fixture.kwargs["model_alias"],
"provider_model": fixture.kwargs["provider_model"],
"api_base": fixture.kwargs["api_base"],
"body": fixture.kwargs["body"],
}
)
completed: Final = subprocess.run(
(_gateway_trace_binary(),),
input=payload,
capture_output=True,
text=True,
check=False,
)
if completed.returncode != 0:
raise RuntimeError(f"Rust gateway trace failed: {completed.stderr.strip()}")
result: Final = json.loads(completed.stdout)
payload: Final = TraceResponsePayload.model_validate(result)
response: Final = _GatewayResponsePayload.model_validate(payload.response)
if response.status != 200:
raise RuntimeError(f"Rust gateway returned {response.status}: {response.body}")
return native_trace_events(payload)
@cache
def _gateway_trace_binary() -> Path:
repo_root: Final = next(parent for parent in Path(__file__).resolve().parents if (parent / "litellm-rust").is_dir())
rust_root: Final = repo_root / "litellm-rust"
completed: Final = subprocess.run(
(
"cargo",
"build",
"--quiet",
"--package",
"litellm-ai-gateway",
"--features",
"trace-parity",
"--bin",
"trace-parity-gateway",
"--target-dir",
rust_root / "target",
),
cwd=rust_root,
capture_output=True,
text=True,
check=False,
)
if completed.returncode != 0:
raise RuntimeError(f"Rust gateway trace build failed: {completed.stderr.strip()}")
return rust_root / "target" / "debug" / "trace-parity-gateway"
def _collect(
route: GatewayRouteSpec, scenario: TraceScenario, engine: Engine
) -> tuple[FunctionTraceEvent, ...] | TraceExecutionFailure:
try:
with replay_server() as provider:
base_fixture: Final = scenario.fixture(engine, provider.url)
fixture: Final = RouteFixture(
kwargs={**base_fixture.kwargs, "api_base": provider.url},
provider_responses=base_fixture.provider_responses,
)
for response in fixture.provider_responses:
provider.enqueue_response(response)
events: Final = _collect_python(fixture, route) if engine == "python" else _collect_rust(fixture, route)
provider.take_requests(len(fixture.provider_responses))
return events
except Exception as error:
return TraceExecutionFailure(engine, f"{type(error).__name__}: {error}")
def _projections(
python_events: tuple[FunctionTraceEvent, ...],
rust_events: tuple[FunctionTraceEvent, ...],
) -> tuple[PipelineProjection, PipelineProjection, str | None]:
try:
return (
pipeline_projection("python", python_events),
pipeline_projection("rust", rust_events),
None,
)
except ValueError as error:
return PipelineProjection(), PipelineProjection(), f"harness: {error}"
def execute_gateway_trace(
route: GatewayRouteSpec,
scenario: TraceScenario,
engine: TraceEngine = "both",
) -> TraceArtifact:
effective_engine: Final[TraceEngine] = "python" if engine == "both" and not route.rust_supported else engine
python_trace: Final = _collect(route, scenario, "python") if effective_engine != "rust" else ()
rust_trace: Final = _collect(route, scenario, "rust") if effective_engine != "python" else ()
collection_python_error: Final = None if isinstance(python_trace, tuple) else f"python: {python_trace.message}"
rust_error: Final = None if isinstance(rust_trace, tuple) else f"rust: {rust_trace.message}"
python_events: Final = python_trace if isinstance(python_trace, tuple) else ()
rust_events: Final = rust_trace if isinstance(rust_trace, tuple) else ()
python, rust, projection_error = _projections(python_events, rust_events)
python_error: Final = projection_error or collection_python_error
return TraceArtifact.from_traces(
engine=effective_engine,
surface="gateway",
sdk_function=route.route,
scenario=scenario.name,
python=python.steps,
rust=rust.steps,
python_error=python_error,
rust_error=rust_error,
)