diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py new file mode 100644 index 00000000000..1fa67acbce8 --- /dev/null +++ b/tests/e2e/conftest.py @@ -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() diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py new file mode 100644 index 00000000000..6f4ea6a0812 --- /dev/null +++ b/tests/e2e/e2e_config.py @@ -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] diff --git a/tests/e2e/e2e_gateway.py b/tests/e2e/e2e_gateway.py new file mode 100644 index 00000000000..7b3963b7c9d --- /dev/null +++ b/tests/e2e/e2e_gateway.py @@ -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, + ) diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py new file mode 100644 index 00000000000..1fea12e1417 --- /dev/null +++ b/tests/e2e/e2e_http.py @@ -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="", + chunks=chunks, + ) diff --git a/tests/e2e/lifecycle.py b/tests/e2e/lifecycle.py new file mode 100644 index 00000000000..ecbfee2c579 --- /dev/null +++ b/tests/e2e/lifecycle.py @@ -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 diff --git a/tests/e2e/models.py b/tests/e2e/models.py new file mode 100644 index 00000000000..f21e2ab4ded --- /dev/null +++ b/tests/e2e/models.py @@ -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] = {} diff --git a/tests/e2e/proxy_client.py b/tests/e2e/proxy_client.py new file mode 100644 index 00000000000..d90e4ef0bb0 --- /dev/null +++ b/tests/e2e/proxy_client.py @@ -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="", + 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]}" + ) diff --git a/tests/e2e/transport.py b/tests/e2e/transport.py new file mode 100644 index 00000000000..e783c2922ed --- /dev/null +++ b/tests/e2e/transport.py @@ -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, + )