mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: p0 issues, added types and shared functions for each test suite
This commit is contained in:
parent
156ca9ff13
commit
76affd0165
8 changed files with 1373 additions and 0 deletions
70
tests/e2e/conftest.py
Normal file
70
tests/e2e/conftest.py
Normal 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
23
tests/e2e/e2e_config.py
Normal 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
187
tests/e2e/e2e_gateway.py
Normal 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
267
tests/e2e/e2e_http.py
Normal 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
113
tests/e2e/lifecycle.py
Normal 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
200
tests/e2e/models.py
Normal 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
397
tests/e2e/proxy_client.py
Normal 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
116
tests/e2e/transport.py
Normal 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,
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue