fix: p0 issues, added types and shared functions for each test suite

This commit is contained in:
mubashir1osmani 2026-06-18 22:35:55 -07:00
parent 156ca9ff13
commit 76affd0165
No known key found for this signature in database
GPG key ID: AB055FF67D0B4D9A
8 changed files with 1373 additions and 0 deletions

70
tests/e2e/conftest.py Normal file
View file

@ -0,0 +1,70 @@
"""Shared fixtures for all live e2e suites under tests/e2e/.
Design rule: skip on environment, fail on behavior. If the proxy is unreachable
the whole session skips; once a request reaches the proxy, behavior is asserted.
Lifecycle: the `resources` fixture maps the init -> run -> teardown contract
(lifecycle.E2ECase) onto pytest - setup is init(), the test body is run(), and
teardown deletes every resource the test created on the long-lived proxy.
Each suite provides its own `client` fixture (a lifecycle.ResourceClient); these
shared fixtures build on it.
"""
import sys
from pathlib import Path
from typing import Iterator
import pytest
import requests
from e2e_config import PROXY_BASE_URL
from lifecycle import GatewayProvider, ResourceManager
def pytest_configure(config: pytest.Config) -> None:
config.addinivalue_line(
"markers",
"e2e: live test that requires a running proxy and real provider keys",
)
def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None:
"""Once the whole e2e session is done (all suites), truncate the spend logs so
the DB doesn't accumulate test rows. Best-effort: a cleanup failure (no DB
reachable) must not fail the run."""
sys.path.insert(0, str(Path(__file__).parent / "spend_tracking"))
try:
from spend_e2e_client import reset_spend_logs # pyright: ignore
reset_spend_logs()
except Exception as exc: # noqa: BLE001 - cleanup is best-effort
print(f"spend-log cleanup skipped: {exc}")
@pytest.fixture(scope="session", autouse=True)
def _require_live_proxy() -> None:
"""Skip the entire session unless a proxy answers its liveness probe."""
try:
resp = requests.get(f"{PROXY_BASE_URL}/health/liveliness", timeout=5)
except requests.RequestException as exc:
pytest.skip(f"No live proxy at {PROXY_BASE_URL}: {exc}")
return
if resp.status_code >= 500:
pytest.skip(f"Proxy at {PROXY_BASE_URL} returned {resp.status_code}")
@pytest.fixture
def resources(client: GatewayProvider) -> Iterator[ResourceManager]:
"""init -> run -> teardown: create a manager, run the test, release resources.
Cleanup goes through the shared Gateway, whatever the suite's client adds."""
manager = ResourceManager(client=client.gateway)
manager.init()
yield manager
manager.teardown()
@pytest.fixture
def scoped_key(resources: ResourceManager) -> str:
"""A fresh all-models key per test, auto-deleted by the resources teardown."""
return resources.key()

23
tests/e2e/e2e_config.py Normal file
View file

@ -0,0 +1,23 @@
"""Generic configuration for live e2e tests against a running LiteLLM proxy.
Shared by every e2e suite under tests/e2e/. Values come from the
environment so the same tests run against localhost or a deployed proxy.
"""
import os
import uuid
PROXY_BASE_URL = os.environ.get("LITELLM_PROXY_URL", "http://localhost:4000").rstrip("/")
MASTER_KEY = os.environ.get("LITELLM_MASTER_KEY", "sk-1234")
# Writes on the proxy are eventually consistent (e.g. spend rows flush on
# proxy_batch_write_at, ~60s). Read-backs poll to this deadline, never sleep-once.
POLL_TIMEOUT = float(os.environ.get("E2E_POLL_TIMEOUT", "120"))
POLL_INTERVAL = float(os.environ.get("E2E_POLL_INTERVAL", "5"))
REQUEST_TIMEOUT = float(os.environ.get("E2E_REQUEST_TIMEOUT", "60"))
def unique_marker() -> str:
"""A short unique token per call/run, so concurrent runs and the shared
response cache never collide on prompts, tags, or customer ids."""
return uuid.uuid4().hex[:12]

187
tests/e2e/e2e_gateway.py Normal file
View file

@ -0,0 +1,187 @@
"""Gateway: the shared proxy operations, DI'd into every client (composition).
A frozen-slots dataclass holding a Transport plus poll config. Clients hold a
Gateway and add their own route methods; the lifecycle ResourceManager uses the
Gateway's key/customer methods for cleanup. Read-backs are eventually consistent
(proxy_batch_write_at ~60s) so they poll to a deadline.
"""
from __future__ import annotations
import time
from collections.abc import Callable
from dataclasses import dataclass
from e2e_http import (
NoBody,
ProbeResult,
Result,
StreamingResponse,
Success,
unwrap,
)
from models import (
ChatBody,
ChatResponse,
CustomerDeleteBody,
EmbedBody,
EmbedResponse,
KeyDeleteBody,
KeyGenerateBody,
KeyGenerateResponse,
KeyInfo,
KeyInfoParams,
KeyInfoResponse,
SpendLogRow,
SpendLogs,
SpendLogsParams,
)
from e2e_config import (
MASTER_KEY,
POLL_INTERVAL,
POLL_TIMEOUT,
PROXY_BASE_URL,
REQUEST_TIMEOUT,
)
from transport import HttpTransport, Transport
RowsPredicate = Callable[[list[SpendLogRow]], bool]
@dataclass(frozen=True, slots=True)
class Gateway:
transport: Transport
poll_timeout: float = 120.0
poll_interval: float = 5.0
# ---- keys / customers (satisfies lifecycle.ResourceClient) ----------
def generate_key(self, body: KeyGenerateBody) -> str:
return unwrap(
self.transport.post(
"/key/generate",
headers=self.transport.master,
json=body,
response_type=KeyGenerateResponse,
)
).key
def delete_key(self, key: str) -> None:
_ = self.transport.post(
"/key/delete",
headers=self.transport.master,
json=KeyDeleteBody(keys=[key]),
response_type=NoBody,
)
def delete_customers(self, user_ids: list[str]) -> None:
if not user_ids:
return
_ = self.transport.post(
"/customer/delete",
headers=self.transport.master,
json=CustomerDeleteBody(user_ids=user_ids),
response_type=NoBody,
)
def key_info(self, key: str) -> KeyInfo:
return unwrap(
self.transport.get(
"/key/info",
headers=self.transport.master,
params=KeyInfoParams(key=key),
response_type=KeyInfoResponse,
)
).info
# ---- LLM calls ------------------------------------------------------
def chat(self, key: str, body: ChatBody) -> Result[ChatResponse]:
return self.transport.post(
"/chat/completions",
headers=self.transport.bearer(key),
json=body,
response_type=ChatResponse,
)
def chat_stream(self, key: str, body: ChatBody) -> StreamingResponse:
return self.transport.stream(
"/chat/completions", headers=self.transport.bearer(key), json=body
)
def embed(self, key: str, body: EmbedBody) -> Result[EmbedResponse]:
return self.transport.post(
"/embeddings",
headers=self.transport.bearer(key),
json=body,
response_type=EmbedResponse,
)
# ---- spend read-back ------------------------------------------------
def spend_logs(self, params: SpendLogsParams) -> list[SpendLogRow]:
result = self.transport.get(
"/spend/logs",
headers=self.transport.master,
params=params,
response_type=SpendLogs,
)
match result:
case Success(data=logs):
return logs.root
case _:
return []
def poll_logs_for_key(
self, key: str, *, min_rows: int = 1, predicate: RowsPredicate | None = None
) -> list[SpendLogRow]:
return self._poll(
lambda: self.spend_logs(SpendLogsParams(api_key=key)), min_rows, predicate
)
def poll_logs_for_request_id(
self,
request_id: str,
*,
min_rows: int = 1,
predicate: RowsPredicate | None = None,
) -> list[SpendLogRow]:
return self._poll(
lambda: self.spend_logs(SpendLogsParams(request_id=request_id)),
min_rows,
predicate,
)
def _poll(
self,
fetch: Callable[[], list[SpendLogRow]],
min_rows: int,
predicate: RowsPredicate | None,
) -> list[SpendLogRow]:
deadline = time.monotonic() + self.poll_timeout
rows: list[SpendLogRow] = []
while time.monotonic() < deadline:
rows = fetch()
if len(rows) >= min_rows and (predicate is None or predicate(rows)):
return rows
time.sleep(self.poll_interval)
return rows
# ---- route probe ----------------------------------------------------
def probe(self, path: str, *, params: NoBody) -> ProbeResult:
return self.transport.probe(path, params=params)
def build_gateway() -> Gateway:
"""The Gateway every suite's client is built from: an HttpTransport pointed at
the configured proxy, with the shared poll budget."""
return Gateway(
transport=HttpTransport(
base_url=PROXY_BASE_URL,
master_key=MASTER_KEY,
request_timeout=REQUEST_TIMEOUT,
),
poll_timeout=POLL_TIMEOUT,
poll_interval=POLL_INTERVAL,
)

267
tests/e2e/e2e_http.py Normal file
View file

@ -0,0 +1,267 @@
"""The ONLY module permitted to call ``requests.*``.
Enforced by tests/code_coverage_tests/check_e2e_no_raw_requests.py. Every request
body / query / header / response is a pydantic model; outcomes are a tagged union
(``Result[R]``) so callers ``match`` on them instead of catching exceptions.
Named e2e_http (not http) so it does not shadow the stdlib ``http`` package that
requests itself imports.
"""
from __future__ import annotations
from typing import Generic, Iterator, Literal, NewType, TypeVar, cast
import requests
from pydantic import BaseModel, ConfigDict, Field
URL = NewType("URL", str)
class Headers(BaseModel):
"""Base for header models. Subclasses may alias to hyphenated header names
(e.g. ``x-litellm-api-key``); serialization uses by_alias."""
model_config = ConfigDict(populate_by_name=True)
class AuthHeaders(Headers):
# litellm accepts either; set whichever the call needs, leave the other None.
authorization: str | None = None
x_litellm_api_key: str | None = Field(default=None, alias="x-litellm-api-key")
class NoBody(BaseModel):
"""Empty body/query for routes that take none."""
# ---------- Result types ----------
R = TypeVar("R", bound=BaseModel)
class Success(BaseModel, Generic[R]):
kind: Literal["success"] = "success"
data: R
class NetworkError(BaseModel):
kind: Literal["network"] = "network"
message: str
class UnauthorizedError(BaseModel):
kind: Literal["unauthorized"] = "unauthorized"
class RateLimitedError(BaseModel):
kind: Literal["rate_limited"] = "rate_limited"
retry_after_seconds: int | None = None
# litellm overloads 429 for budget_exceeded too, so keep the body to tell them apart.
body: str = ""
class ValidationError(BaseModel):
kind: Literal["validation"] = "validation"
message: str
class UnknownApiError(BaseModel):
kind: Literal["unknown"] = "unknown"
status_code: int
body: str
type Result[R: BaseModel] = (
Success[R]
| NetworkError
| UnauthorizedError
| RateLimitedError
| ValidationError
| UnknownApiError
)
class ProbeResult(BaseModel):
"""A route's reachability: status + body, no schema validation. Healthy ==
route exists (not 404) and the handler did not crash (not 5xx)."""
status_code: int
body: str
@property
def healthy(self) -> bool:
return 200 <= self.status_code != 404 and self.status_code < 500
class StreamingResponse(BaseModel):
"""Raw outcome for calls whose body is provider-native or streamed: status, the
x-litellm-call-id header (== SpendLogs.request_id), the content-type (which
tells streaming `text/event-stream` from non-streaming `application/json`), and
the body. Used by passthrough and streaming, where one validated JSON model
does not fit."""
status_code: int
call_id: str | None = None # x-litellm-call-id header
content_type: str | None = None
body: str
chunks: int = 0 # streamed events (0 for non-streaming)
@property
def ok(self) -> bool:
return 200 <= self.status_code < 300
@property
def is_streaming(self) -> bool:
return "text/event-stream" in (self.content_type or "")
def _hdr(resp: requests.Response, name: str) -> str | None:
value = resp.headers.get(name)
return value if isinstance(value, str) else None
def unwrap[R: BaseModel](result: Result[R]) -> R:
match result:
case Success(data=data):
return data
case _:
raise AssertionError(result)
def is_ok[R: BaseModel](result: Result[R]) -> bool:
match result:
case Success():
return True
case _:
return False
def _headers(headers: BaseModel) -> dict[str, str]:
dumped: dict[str, object] = headers.model_dump(by_alias=True, exclude_none=True)
return {key: str(value) for key, value in dumped.items()}
def _classify[R: BaseModel](
resp: requests.Response, response_type: type[R]
) -> Result[R]:
if resp.status_code == 401:
return UnauthorizedError()
if resp.status_code == 429:
return RateLimitedError(body=resp.text)
if not resp.ok:
return UnknownApiError(status_code=resp.status_code, body=resp.text)
try:
return Success(data=response_type.model_validate(resp.json()))
except Exception as exc: # noqa: BLE001 - any parse/validation failure is a value
return ValidationError(message=str(exc))
def post[R: BaseModel](
url: URL,
*,
headers: BaseModel,
json: BaseModel,
response_type: type[R],
timeout: float = 30.0,
) -> Result[R]:
try:
resp = requests.post(
str(url),
headers=_headers(headers),
json=json.model_dump(by_alias=True, exclude_none=True),
timeout=timeout,
)
except requests.RequestException as exc:
return NetworkError(message=str(exc))
return _classify(resp, response_type)
def get[R: BaseModel](
url: URL,
*,
headers: BaseModel,
params: BaseModel,
response_type: type[R],
timeout: float = 30.0,
) -> Result[R]:
try:
resp = requests.get(
str(url),
headers=_headers(headers),
params=params.model_dump(by_alias=True, exclude_none=True),
timeout=timeout,
)
except requests.RequestException as exc:
return NetworkError(message=str(exc))
return _classify(resp, response_type)
def delete[R: BaseModel](
url: URL,
*,
headers: BaseModel,
json: BaseModel,
response_type: type[R],
timeout: float = 30.0,
) -> Result[R]:
try:
resp = requests.delete(
str(url),
headers=_headers(headers),
json=json.model_dump(by_alias=True, exclude_none=True),
timeout=timeout,
)
except requests.RequestException as exc:
return NetworkError(message=str(exc))
return _classify(resp, response_type)
def probe(
url: URL, *, headers: BaseModel, params: BaseModel, timeout: float = 30.0
) -> ProbeResult:
try:
resp = requests.get(
str(url),
headers=_headers(headers),
params=params.model_dump(by_alias=True, exclude_none=True),
timeout=timeout,
)
except requests.RequestException as exc:
return ProbeResult(status_code=-1, body=str(exc))
return ProbeResult(status_code=resp.status_code, body=resp.text)
def stream(
url: URL, *, headers: BaseModel, json: BaseModel, timeout: float = 60.0
) -> StreamingResponse:
"""Streaming (SSE) call: consumes the stream counting events, and captures the
x-litellm-call-id + content-type headers. Body is elided."""
try:
resp = requests.post(
str(url),
headers=_headers(headers),
json=json.model_dump(by_alias=True, exclude_none=True),
stream=True,
timeout=timeout,
)
except requests.RequestException as exc:
return StreamingResponse(status_code=-1, body=str(exc))
call_id = _hdr(resp, "x-litellm-call-id")
content_type = _hdr(resp, "content-type")
if not (200 <= resp.status_code < 300):
return StreamingResponse(
status_code=resp.status_code,
call_id=call_id,
content_type=content_type,
body=resp.text,
)
lines = cast("Iterator[bytes]", resp.iter_lines())
chunks = sum(1 for line in lines if line)
return StreamingResponse(
status_code=resp.status_code,
call_id=call_id,
content_type=content_type,
body="<streamed>",
chunks=chunks,
)

113
tests/e2e/lifecycle.py Normal file
View file

@ -0,0 +1,113 @@
"""Lifecycle contract and resource cleanup for stateful e2e tests.
Shared by every e2e suite under tests/e2e/. The proxy under test is
long-lived and never reset between tests, so anything a test creates (keys,
customers, teams, orgs, users, guardrails, budgets, ...) persists unless
explicitly deleted. Every check follows an init -> run -> teardown lifecycle;
teardown releases each resource init() created, even when run() raises.
In pytest terms (see conftest.py): the `resources` fixture's setup is init(),
the test body is run(), and the fixture's teardown is teardown().
"""
from dataclasses import dataclass, field
from typing import Callable, List, Protocol, runtime_checkable
from e2e_gateway import Gateway
from models import KeyGenerateBody
@runtime_checkable
class E2ECase(Protocol):
"""A stateful e2e check run against a long-lived proxy.
init() acquires resources, run() exercises behaviour and asserts, teardown()
releases everything init() created. teardown() must run even if run() raises.
"""
def init(self) -> None: ...
def run(self) -> None: ...
def teardown(self) -> None: ...
def run_case(case: E2ECase) -> None:
"""Drive a case through its lifecycle: init -> run -> teardown.
teardown always runs, even if run() raises (or skips), so resources the case
created on the long-lived proxy are released.
"""
case.init()
try:
case.run()
finally:
case.teardown()
@runtime_checkable
class ResourceClient(Protocol):
"""Proxy operations the convenience creators use. Resource types without a
creator here are handled generically via ResourceManager.defer(). The Gateway
satisfies this."""
def generate_key(self, body: KeyGenerateBody) -> str: ...
def delete_key(self, key: str) -> None: ...
def delete_customers(self, user_ids: List[str]) -> None: ...
@runtime_checkable
class GatewayProvider(Protocol):
"""Every suite's client exposes the shared Gateway, which the resources fixture
uses for cleanup. The client adds its own route methods on top."""
@property
def gateway(self) -> Gateway: ...
@dataclass
class ResourceManager:
"""Registry of teardown actions for resources a test creates on the stateful
proxy.
Not limited to any resource type: register a cleanup with ``defer()`` for a
key, customer, team, org, user, guardrail, budget, MCP server - anything with
a delete. The two most common resources have sugar (``key``, ``customer``);
everything else is ``resources.defer(lambda: client.delete_team(team_id))``.
Cleanups run LIFO (so a resource is removed before whatever it depends on) and
best-effort (one failing cleanup never blocks the rest).
"""
client: ResourceClient
_cleanups: List[Callable[[], None]] = field(
default_factory=list
) # mutable-ok: append-only teardown registry
def init(self) -> None:
"""No global setup needed today; present for lifecycle symmetry."""
return None
def defer(self, cleanup: Callable[[], None]) -> None:
"""Register a teardown action for any resource the test just created."""
self._cleanups.append(cleanup)
def key(self) -> str:
"""Create an all-models virtual key; delete it on teardown."""
key = self.client.generate_key(KeyGenerateBody(models=[]))
self.defer(lambda: self.client.delete_key(key))
return key
def customer(self, customer_id: str) -> str:
"""Track an end-user id (from the `user` param); delete it on teardown."""
self.defer(lambda: self.client.delete_customers([customer_id]))
return customer_id
def teardown(self) -> None:
for cleanup in reversed(self._cleanups):
try:
cleanup()
except Exception:
pass # best-effort: a failed cleanup must not block the rest

200
tests/e2e/models.py Normal file
View file

@ -0,0 +1,200 @@
"""Shared pydantic request/response models for the e2e gateway.
Only the fields the tests read are modelled; pydantic ignores the rest, so a
response validates without mirroring every proxy field. No untyped dicts.
"""
from __future__ import annotations
from pydantic import BaseModel, RootModel
# ---------- keys ----------
class ModelBudgetEntry(BaseModel):
budget_limit: float
time_period: str
class BudgetWindow(BaseModel):
budget_duration: str
max_budget: float
class KeyGenerateBody(BaseModel):
models: list[str] = []
duration: str | None = None
max_budget: float | None = None
soft_budget: float | None = None
budget_duration: str | None = None
user_id: str | None = None
team_id: str | None = None
budget_id: str | None = None
model_max_budget: dict[str, ModelBudgetEntry] | None = None
budget_limits: list[BudgetWindow] | None = None
class KeyGenerateResponse(BaseModel):
key: str
class KeyDeleteBody(BaseModel):
keys: list[str]
class KeyInfoParams(BaseModel):
key: str
class LiteLLMBudgetTable(BaseModel):
max_budget: float | None = None
soft_budget: float | None = None
budget_duration: str | None = None
budget_reset_at: str | None = None
class KeyInfo(BaseModel):
spend: float | None = None
max_budget: float | None = None
budget_reset_at: str | None = None
budget_id: str | None = None
litellm_budget_table: LiteLLMBudgetTable | None = None
class KeyInfoResponse(BaseModel):
info: KeyInfo
# ---------- customers ----------
class CustomerDeleteBody(BaseModel):
user_ids: list[str]
# ---------- chat / embeddings ----------
class ChatMetadata(BaseModel):
tags: list[str] | None = None
class ChatMessage(BaseModel):
role: str
content: str
class ChatBody(BaseModel):
model: str
messages: list[ChatMessage]
stream: bool = False
max_tokens: int | None = None
user: str | None = None
metadata: ChatMetadata | None = None
class OutMessage(BaseModel):
content: str | None = None
class ChatChoice(BaseModel):
message: OutMessage | None = None
class Usage(BaseModel):
prompt_tokens: int | None = None
completion_tokens: int | None = None
total_tokens: int | None = None
class ChatResponse(BaseModel):
id: str | None = None
model: str | None = None
choices: list[ChatChoice] = []
usage: Usage | None = None
class EmbedBody(BaseModel):
model: str
input: str
class EmbedResponse(BaseModel):
model: str | None = None
# ---------- spend logs ----------
class SpendLogRow(BaseModel):
request_id: str | None = None
model: str | None = None
spend: float | None = None
status: str | None = None
cache_hit: str | None = None
call_type: str | None = None
custom_llm_provider: str | None = None
end_user: str | None = None
prompt_tokens: int | None = None
completion_tokens: int | None = None
total_tokens: int | None = None
request_tags: list[str] | None = None
class SpendLogs(RootModel[list[SpendLogRow]]):
pass
class SpendLogsParams(BaseModel):
request_id: str | None = None
api_key: str | None = None
# ---------- spend calculate ----------
class SpendCalculateBody(BaseModel):
model: str
messages: list[ChatMessage]
class SpendCalculateResponse(BaseModel):
cost: float
# ---------- spend tags ----------
class TagSpend(BaseModel):
individual_request_tag: str
log_count: int | None = None
total_spend: float | None = None
class TagSpends(RootModel[list[TagSpend]]):
pass
class SpendTagsResponse(BaseModel):
spend_per_tag: list[TagSpend] | None = None
# ---------- route probing ----------
class DateRangeParams(BaseModel):
start_date: str
end_date: str
class RouteSpec(RootModel[dict[str, object]]):
"""One /openapi.json path entry: a map of HTTP method -> operation. Only the
method names are read, so the operation specs stay opaque."""
@property
def methods(self) -> frozenset[str]:
return frozenset(method.lower() for method in self.root)
class OpenAPISchema(BaseModel):
paths: dict[str, RouteSpec] = {}

397
tests/e2e/proxy_client.py Normal file
View file

@ -0,0 +1,397 @@
"""Generic HTTP client for live e2e tests against a running LiteLLM proxy.
Shared by every e2e suite under tests/e2e/. Covers the proxy operations any
suite needs: key/customer management (so the shared ResourceManager can clean up),
OpenAI-compatible calls, route probing, and SpendLogs read-back. Suite-specific
clients subclass ProxyClient (see spend_tracking/, llm_translation/, budgets/).
Talks to the proxy over real HTTP so every test sees what a real client sees:
the x-litellm-call-id header, the response body, and the rows the proxy writes.
Writes are eventually consistent (proxy_batch_write_at ~60s), so read-backs poll
to a deadline rather than sleeping once.
"""
import json
import time
from dataclasses import dataclass
from typing import Callable, Dict, List, Optional, Protocol, runtime_checkable
import pytest
import requests
from e2e_config import (
MASTER_KEY,
POLL_INTERVAL,
POLL_TIMEOUT,
PROXY_BASE_URL,
REQUEST_TIMEOUT,
unique_marker,
)
from e2e_gateway import Gateway
from transport import HttpTransport
__all__ = [
"CallResult",
"CallOutcome",
"ProbeResult",
"ProxyClient",
"SpendLogRow",
"auth_headers",
"proxy_client_kwargs",
"require_successful_call",
"unique_marker",
]
SpendLogRow = Dict[str, object]
@runtime_checkable
class CallOutcome(Protocol):
"""Anything with an HTTP status and body that can pass the skip/fail boundary.
Read-only members so frozen dataclasses (CallResult, PassthroughResult) match.
"""
@property
def status_code(self) -> int: ...
@property
def body(self) -> str: ...
@property
def ok(self) -> bool: ...
@dataclass(frozen=True, slots=True)
class ProbeResult:
"""Outcome of a single route probe: enough to see *why* it (mis)behaved."""
url: str
status_code: int
body: str
@property
def healthy(self) -> bool:
# Route exists (not 404), handler did not crash (not 5xx), request
# completed (not -1). A 4xx (missing params/auth) still means it ran.
return 200 <= self.status_code != 404 and self.status_code < 500
def __str__(self) -> str:
return f"GET {self.url} -> {self.status_code}\n{self.body[:600]}"
@dataclass(frozen=True, slots=True)
class CallResult:
"""Outcome of a single OpenAI-compatible call made through the proxy."""
status_code: int
call_id: Optional[str] # x-litellm-call-id response header
response_id: Optional[str] # body "id"; SpendLogs.request_id is derived from this
response_cost_header: Optional[str] # x-litellm-response-cost header
body: str
content: Optional[str]
@property
def ok(self) -> bool:
return 200 <= self.status_code < 300
def auth_headers(key: str) -> Dict[str, str]:
return {"Authorization": f"Bearer {key}", "Content-Type": "application/json"}
class ProxyClient:
def __init__(
self,
base_url: str,
master_key: str,
*,
request_timeout: float,
poll_timeout: float,
poll_interval: float,
) -> None:
self._base_url = base_url.rstrip("/")
self._master_key = master_key
self._request_timeout = request_timeout
self._poll_timeout = poll_timeout
self._poll_interval = poll_interval
@property
def gateway(self) -> Gateway:
"""The shared typed Gateway over this client's proxy; the resources fixture
cleans up through it while suites not yet migrated keep their own methods."""
return Gateway(
transport=HttpTransport(
base_url=self._base_url,
master_key=self._master_key,
request_timeout=self._request_timeout,
),
poll_timeout=self._poll_timeout,
poll_interval=self._poll_interval,
)
# ---- key / customer management (satisfies lifecycle.ResourceClient) ----
def generate_key(
self,
*,
models: Optional[List[str]] = None,
max_budget: Optional[float] = None,
metadata: Optional[Dict[str, object]] = None,
extra_params: Optional[Dict[str, object]] = None,
) -> str:
payload: Dict[str, object] = {"models": models or [], "duration": None}
if max_budget is not None:
payload["max_budget"] = max_budget
if metadata is not None:
payload["metadata"] = metadata
if extra_params:
payload.update(extra_params)
resp = requests.post(
f"{self._base_url}/key/generate",
headers=auth_headers(self._master_key),
json=payload,
timeout=self._request_timeout,
)
resp.raise_for_status()
return str(resp.json()["key"])
def key_info(self, key: str) -> Dict[str, object]:
resp = requests.get(
f"{self._base_url}/key/info",
headers=auth_headers(self._master_key),
params={"key": key},
timeout=self._request_timeout,
)
resp.raise_for_status()
return dict(resp.json().get("info", {}))
def delete_key(self, key: str) -> None:
"""Best-effort teardown; a failed cleanup must not fail the test."""
try:
requests.post(
f"{self._base_url}/key/delete",
headers=auth_headers(self._master_key),
json={"keys": [key]},
timeout=self._request_timeout,
)
except requests.RequestException:
pass
def delete_customers(self, user_ids: List[str]) -> None:
"""Best-effort teardown of end-user/customer rows the `user` param creates."""
if not user_ids:
return
try:
requests.post(
f"{self._base_url}/customer/delete",
headers=auth_headers(self._master_key),
json={"user_ids": user_ids},
timeout=self._request_timeout,
)
except requests.RequestException:
pass
# ---- OpenAI-compatible calls ----------------------------------------
def chat(
self,
key: str,
model: str,
content: str,
*,
stream: bool = False,
metadata: Optional[Dict[str, object]] = None,
extra_body: Optional[Dict[str, object]] = None,
) -> CallResult:
body: Dict[str, object] = {
"model": model,
"messages": [{"role": "user", "content": content}],
"stream": stream,
}
if metadata is not None:
body["metadata"] = metadata
if extra_body is not None:
body.update(extra_body)
url = f"{self._base_url}/chat/completions"
if stream:
return self._chat_stream(url, key, body)
resp = requests.post(
url, headers=auth_headers(key), json=body, timeout=self._request_timeout
)
parsed = resp.json() if resp.content else {}
choices = parsed.get("choices") or [{}]
message_content = (choices[0].get("message") or {}).get("content")
return CallResult(
status_code=resp.status_code,
call_id=resp.headers.get("x-litellm-call-id"),
response_id=parsed.get("id"),
response_cost_header=resp.headers.get("x-litellm-response-cost"),
body=resp.text,
content=message_content,
)
def _chat_stream(self, url: str, key: str, body: Dict[str, object]) -> CallResult:
resp = requests.post(
url,
headers=auth_headers(key),
json=body,
stream=True,
timeout=self._request_timeout,
)
if not (200 <= resp.status_code < 300):
return CallResult(
status_code=resp.status_code,
call_id=resp.headers.get("x-litellm-call-id"),
response_id=None,
response_cost_header=resp.headers.get("x-litellm-response-cost"),
body=resp.text,
content=None,
)
response_id: Optional[str] = None
parts: List[str] = []
for raw in resp.iter_lines():
if not raw:
continue
line = raw.decode("utf-8")
if not line.startswith("data:"):
continue
data = line[len("data:") :].strip()
if data == "[DONE]":
break
chunk = json.loads(data)
response_id = chunk.get("id", response_id)
for choice in chunk.get("choices", []):
piece = (choice.get("delta") or {}).get("content")
if piece:
parts.append(piece)
return CallResult(
status_code=resp.status_code,
call_id=resp.headers.get("x-litellm-call-id"),
response_id=response_id,
response_cost_header=resp.headers.get("x-litellm-response-cost"),
body="<streamed>",
content="".join(parts) or None,
)
def embed(self, key: str, model: str, text: str) -> CallResult:
resp = requests.post(
f"{self._base_url}/embeddings",
headers=auth_headers(key),
json={"model": model, "input": text},
timeout=self._request_timeout,
)
parsed = resp.json() if resp.content else {}
return CallResult(
status_code=resp.status_code,
call_id=resp.headers.get("x-litellm-call-id"),
response_id=parsed.get("id"),
response_cost_header=resp.headers.get("x-litellm-response-cost"),
body=resp.text,
content=None,
)
# ---- route discovery -------------------------------------------------
def get_openapi(self) -> Dict[str, object]:
"""The proxy's live route schema from /openapi.json."""
resp = requests.get(
f"{self._base_url}/openapi.json", timeout=self._request_timeout
)
resp.raise_for_status()
return dict(resp.json())
def probe(self, path: str, params: Optional[Dict[str, str]] = None) -> ProbeResult:
"""GET a route with master-key auth; capture status + body to show why."""
url = f"{self._base_url}{path}"
try:
resp = requests.get(
url,
headers=auth_headers(self._master_key),
params=params or {},
timeout=self._request_timeout,
)
except requests.RequestException as exc:
return ProbeResult(url=url, status_code=-1, body=f"request error: {exc}")
return ProbeResult(url=url, status_code=resp.status_code, body=resp.text)
# ---- SpendLogs read-back --------------------------------------------
def _get_logs(
self, *, request_id: Optional[str] = None, api_key: Optional[str] = None
) -> List[SpendLogRow]:
params: Dict[str, str] = {}
if request_id is not None:
params["request_id"] = request_id
if api_key is not None:
params["api_key"] = api_key
resp = requests.get(
f"{self._base_url}/spend/logs",
headers=auth_headers(self._master_key),
params=params,
timeout=self._request_timeout,
)
if resp.status_code != 200:
return []
data = resp.json()
return [dict(row) for row in data] if isinstance(data, list) else []
def poll_logs_for_key(
self,
key: str,
*,
min_rows: int = 1,
predicate: Optional[Callable[[List[SpendLogRow]], bool]] = None,
) -> List[SpendLogRow]:
return self._poll(lambda: self._get_logs(api_key=key), min_rows, predicate)
def poll_logs_for_request_id(
self,
request_id: str,
*,
min_rows: int = 1,
predicate: Optional[Callable[[List[SpendLogRow]], bool]] = None,
) -> List[SpendLogRow]:
return self._poll(
lambda: self._get_logs(request_id=request_id), min_rows, predicate
)
def _poll(
self,
fetch: Callable[[], List[SpendLogRow]],
min_rows: int,
predicate: Optional[Callable[[List[SpendLogRow]], bool]],
) -> List[SpendLogRow]:
deadline = time.monotonic() + self._poll_timeout
rows: List[SpendLogRow] = []
while time.monotonic() < deadline:
rows = fetch()
satisfied = len(rows) >= min_rows and (
predicate is None or predicate(rows)
)
if satisfied:
return rows
time.sleep(self._poll_interval)
return rows
def proxy_client_kwargs() -> Dict[str, object]:
"""Constructor kwargs shared by every ProxyClient subclass."""
return {
"base_url": PROXY_BASE_URL,
"master_key": MASTER_KEY,
"request_timeout": REQUEST_TIMEOUT,
"poll_timeout": POLL_TIMEOUT,
"poll_interval": POLL_INTERVAL,
}
def require_successful_call(result: CallOutcome) -> None:
"""A call that should have succeeded but didn't is a hard failure, never a
skip - if the proxy can't make a call it's expected to, the test must fail."""
if result.ok:
return
pytest.fail(
f"upstream call failed (status {result.status_code}); "
f"body={result.body[:300]}"
)

116
tests/e2e/transport.py Normal file
View file

@ -0,0 +1,116 @@
"""Transport: the typed request primitives clients use, behind a Protocol.
`Transport` is what each client depends on (composition + DI); `HttpTransport` is
the concrete frozen-slots dataclass that fulfils it via the e2e_http wrapper. No
client touches requests.* or builds raw dicts; they pass pydantic models here.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Protocol
from pydantic import BaseModel
import e2e_http
from e2e_http import URL, AuthHeaders, ProbeResult, Result, StreamingResponse
class Transport(Protocol):
def post[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]: ...
def stream(
self, path: str, *, headers: BaseModel, json: BaseModel
) -> StreamingResponse: ...
def get[R: BaseModel](
self,
path: str,
*,
headers: BaseModel,
params: BaseModel,
response_type: type[R],
) -> Result[R]: ...
def delete[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]: ...
def probe(self, path: str, *, params: BaseModel) -> ProbeResult: ...
def bearer(self, key: str) -> AuthHeaders: ...
@property
def master(self) -> AuthHeaders: ...
@dataclass(frozen=True, slots=True)
class HttpTransport:
base_url: str
master_key: str
request_timeout: float = 60.0
def _url(self, path: str) -> URL:
return URL(f"{self.base_url.rstrip('/')}{path}")
def bearer(self, key: str) -> AuthHeaders:
return AuthHeaders(authorization=f"Bearer {key}")
@property
def master(self) -> AuthHeaders:
return self.bearer(self.master_key)
def post[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]:
return e2e_http.post(
self._url(path),
headers=headers,
json=json,
response_type=response_type,
timeout=self.request_timeout,
)
def get[R: BaseModel](
self,
path: str,
*,
headers: BaseModel,
params: BaseModel,
response_type: type[R],
) -> Result[R]:
return e2e_http.get(
self._url(path),
headers=headers,
params=params,
response_type=response_type,
timeout=self.request_timeout,
)
def delete[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]:
return e2e_http.delete(
self._url(path),
headers=headers,
json=json,
response_type=response_type,
timeout=self.request_timeout,
)
def stream(
self, path: str, *, headers: BaseModel, json: BaseModel
) -> StreamingResponse:
return e2e_http.stream(
self._url(path), headers=headers, json=json, timeout=self.request_timeout
)
def probe(self, path: str, *, params: BaseModel) -> ProbeResult:
return e2e_http.probe(
self._url(path),
headers=self.master,
params=params,
timeout=self.request_timeout,
)