refactor: migrate to gateway client

This commit is contained in:
mubashir1osmani 2026-06-19 00:27:13 -07:00
parent 6843ca0325
commit ff7c3c55c7
No known key found for this signature in database
GPG key ID: AB055FF67D0B4D9A
15 changed files with 566 additions and 686 deletions

View file

@ -1,96 +1,250 @@
"""Client for budget e2e tests: the shared ProxyClient plus budget-bearing entity
"""Client for budget e2e tests: the shared Gateway plus budget-bearing entity
management (user / team / team-member / org / customer / tag / budget-table) and
info reads.
Over-budget surfaces as a ``budget_exceeded`` error; ``is_budget_block`` detects it
on a CallResult. Create methods return the new id and raise on failure; tests
register the matching delete with ``resources.defer(...)`` for cleanup.
on a chat outcome. Create methods return the new id and raise on failure; tests
register the matching delete with ``resources.defer(...)`` for cleanup. The request
and response models are co-located here because only this suite uses them.
"""
from typing import Dict, Optional
from __future__ import annotations
import requests
from pydantic import TypeAdapter
from pydantic.dataclasses import dataclass
from dataclasses import dataclass
from proxy_client import CallResult, ProxyClient, auth_headers, proxy_client_kwargs
from pydantic import BaseModel, RootModel
from e2e_gateway import Gateway, build_gateway
from e2e_http import NoBody, StreamingResponse, Success, unwrap
from models import (
BudgetWindow,
ChatBody,
ChatMessage,
ChatMetadata,
KeyGenerateBody,
ModelBudgetEntry,
)
def is_budget_block(result: CallResult) -> bool:
class UserNewBody(BaseModel):
max_budget: float
class UserNewResponse(BaseModel):
user_id: str
class UserDeleteBody(BaseModel):
user_ids: list[str]
class CustomerNewBody(BaseModel):
user_id: str
max_budget: float
class OrgNewBody(BaseModel):
organization_alias: str
max_budget: float
class OrgNewResponse(BaseModel):
organization_id: str
class OrgDeleteBody(BaseModel):
organization_ids: list[str]
class TeamMember(BaseModel):
role: str
user_id: str
class TeamNewBody(BaseModel):
team_alias: str
max_budget: float | None = None
organization_id: str | None = None
class TeamNewResponse(BaseModel):
team_id: str
class TeamDeleteBody(BaseModel):
team_ids: list[str]
class TeamMemberAddBody(BaseModel):
team_id: str
member: TeamMember
max_budget_in_team: float | None = None
class TagNewBody(BaseModel):
name: str
max_budget: float
class TagDeleteBody(BaseModel):
name: str
class BudgetNewBody(BaseModel):
max_budget: float
soft_budget: float | None = None
budget_duration: str | None = None
class BudgetNewResponse(BaseModel):
budget_id: str
class BudgetDeleteBody(BaseModel):
id: str
class BudgetInfoBody(BaseModel):
budgets: list[str]
class BudgetRow(BaseModel):
budget_id: str | None = None
max_budget: float | None = None
soft_budget: float | None = None
budget_duration: str | None = None
budget_reset_at: str | None = None
class BudgetInfoResponse(RootModel[list[BudgetRow]]):
pass
def is_budget_block(result: StreamingResponse) -> bool:
"""True if the call was rejected for being over budget (vs a provider error)."""
return not result.ok and "budget_exceeded" in result.body
def model_budget(model: str, limit: float, period: str = "30d") -> dict:
"""A model_max_budget dict entry: per-model cap with a reset window."""
return {model: {"budget_limit": limit, "time_period": period}}
def model_budget(
model: str, limit: float, period: str = "30d"
) -> dict[str, ModelBudgetEntry]:
"""A model_max_budget entry: per-model cap with a reset window."""
return {model: ModelBudgetEntry(budget_limit=limit, time_period=period)}
@dataclass(frozen=True, slots=True)
class BudgetRow:
"""A /budget/info row: only the fields tests assert on, pydantic ignores the rest."""
class BudgetClient:
gateway: Gateway
budget_id: Optional[str] = None
max_budget: Optional[float] = None
soft_budget: Optional[float] = None
budget_duration: Optional[str] = None
budget_reset_at: Optional[str] = None
# ---- generic key ops (delegate to the shared Gateway) ---------------
_BUDGET_ROWS = TypeAdapter(tuple[BudgetRow, ...])
class BudgetClient(ProxyClient):
def _post(self, path: str, body: Dict[str, object]) -> Dict[str, object]:
resp = requests.post(
f"{self._base_url}{path}",
headers=auth_headers(self._master_key),
json=body,
timeout=self._request_timeout,
def generate_key(
self,
*,
models: list[str] | None = None,
max_budget: float | None = None,
soft_budget: float | None = None,
budget_duration: str | None = None,
budget_id: str | None = None,
user_id: str | None = None,
team_id: str | None = None,
model_max_budget: dict[str, ModelBudgetEntry] | None = None,
budget_limits: list[BudgetWindow] | None = None,
) -> str:
return self.gateway.generate_key(
KeyGenerateBody(
models=models or [],
max_budget=max_budget,
soft_budget=soft_budget,
budget_duration=budget_duration,
budget_id=budget_id,
user_id=user_id,
team_id=team_id,
model_max_budget=model_max_budget,
budget_limits=budget_limits,
)
)
resp.raise_for_status()
data = resp.json() if resp.text else {}
# delete routes return a bare count, not an object; callers ignore it.
return data if isinstance(data, dict) else {}
def _delete(self, path: str, body: Dict[str, object]) -> None:
resp = requests.delete(
f"{self._base_url}{path}",
headers=auth_headers(self._master_key),
json=body,
timeout=self._request_timeout,
def delete_key(self, key: str) -> None:
self.gateway.delete_key(key)
def delete_customers(self, user_ids: list[str]) -> None:
self.gateway.delete_customers(user_ids)
# ---- chat (raw HTTP outcome: a budget block surfaces as a non-2xx) --
def chat(
self,
key: str,
model: str,
content: str,
*,
max_tokens: int | None = None,
user: str | None = None,
tags: list[str] | None = None,
) -> StreamingResponse:
return self.gateway.transport.send(
"/chat/completions",
headers=self.gateway.transport.bearer(key),
json=ChatBody(
model=model,
messages=[ChatMessage(role="user", content=content)],
max_tokens=max_tokens,
user=user,
metadata=ChatMetadata(tags=tags) if tags else None,
),
)
resp.raise_for_status()
# ---- internal user --------------------------------------------------
def create_user(self, *, max_budget: float, budget_duration: Optional[str] = None) -> str:
body: Dict[str, object] = {"max_budget": max_budget}
if budget_duration is not None:
body["budget_duration"] = budget_duration
return str(self._post("/user/new", body)["user_id"])
def create_user(self, *, max_budget: float) -> str:
return unwrap(
self.gateway.transport.post(
"/user/new",
headers=self.gateway.transport.master,
json=UserNewBody(max_budget=max_budget),
response_type=UserNewResponse,
)
).user_id
def delete_user(self, user_id: str) -> None:
self._post("/user/delete", {"user_ids": [user_id]})
_ = self.gateway.transport.post(
"/user/delete",
headers=self.gateway.transport.master,
json=UserDeleteBody(user_ids=[user_id]),
response_type=NoBody,
)
# ---- customer / end-user -------------------------------------------
def create_customer(self, customer_id: str, *, max_budget: float) -> str:
self._post("/customer/new", {"user_id": customer_id, "max_budget": max_budget})
resp = self.gateway.transport.send(
"/customer/new",
headers=self.gateway.transport.master,
json=CustomerNewBody(user_id=customer_id, max_budget=max_budget),
)
assert resp.ok, resp.body
return customer_id
# ---- organization ---------------------------------------------------
def create_org(self, *, max_budget: float, alias: str) -> str:
return str(
self._post(
return unwrap(
self.gateway.transport.post(
"/organization/new",
{"organization_alias": alias, "max_budget": max_budget},
)["organization_id"]
)
headers=self.gateway.transport.master,
json=OrgNewBody(organization_alias=alias, max_budget=max_budget),
response_type=OrgNewResponse,
)
).organization_id
def delete_org(self, org_id: str) -> None:
self._delete("/organization/delete", {"organization_ids": [org_id]})
_ = self.gateway.transport.delete(
"/organization/delete",
headers=self.gateway.transport.master,
json=OrgDeleteBody(organization_ids=[org_id]),
response_type=NoBody,
)
# ---- team -----------------------------------------------------------
@ -98,39 +252,62 @@ class BudgetClient(ProxyClient):
self,
*,
alias: str,
max_budget: Optional[float] = None,
organization_id: Optional[str] = None,
extra: Optional[Dict[str, object]] = None,
max_budget: float | None = None,
organization_id: str | None = None,
) -> str:
body: Dict[str, object] = {"team_alias": alias}
if max_budget is not None:
body["max_budget"] = max_budget
if organization_id is not None:
body["organization_id"] = organization_id
if extra:
body.update(extra)
return str(self._post("/team/new", body)["team_id"])
return unwrap(
self.gateway.transport.post(
"/team/new",
headers=self.gateway.transport.master,
json=TeamNewBody(
team_alias=alias,
max_budget=max_budget,
organization_id=organization_id,
),
response_type=TeamNewResponse,
)
).team_id
def delete_team(self, team_id: str) -> None:
self._post("/team/delete", {"team_ids": [team_id]})
_ = self.gateway.transport.post(
"/team/delete",
headers=self.gateway.transport.master,
json=TeamDeleteBody(team_ids=[team_id]),
response_type=NoBody,
)
def add_team_member(self, team_id: str, user_id: str, *, max_budget_in_team: Optional[float] = None) -> None:
body: Dict[str, object] = {
"team_id": team_id,
"member": {"role": "user", "user_id": user_id},
}
if max_budget_in_team is not None:
body["max_budget_in_team"] = max_budget_in_team
self._post("/team/member_add", body)
def add_team_member(
self, team_id: str, user_id: str, *, max_budget_in_team: float | None = None
) -> None:
resp = self.gateway.transport.send(
"/team/member_add",
headers=self.gateway.transport.master,
json=TeamMemberAddBody(
team_id=team_id,
member=TeamMember(role="user", user_id=user_id),
max_budget_in_team=max_budget_in_team,
),
)
assert resp.ok, resp.body
# ---- tag ------------------------------------------------------------
def create_tag(self, name: str, *, max_budget: float) -> str:
self._post("/tag/new", {"name": name, "max_budget": max_budget})
resp = self.gateway.transport.send(
"/tag/new",
headers=self.gateway.transport.master,
json=TagNewBody(name=name, max_budget=max_budget),
)
assert resp.ok, resp.body
return name
def delete_tag(self, name: str) -> None:
self._post("/tag/delete", {"name": name})
_ = self.gateway.transport.post(
"/tag/delete",
headers=self.gateway.transport.master,
json=TagDeleteBody(name=name),
response_type=NoBody,
)
# ---- budget table ---------------------------------------------------
@ -138,29 +315,43 @@ class BudgetClient(ProxyClient):
self,
*,
max_budget: float,
soft_budget: Optional[float] = None,
budget_duration: Optional[str] = None,
soft_budget: float | None = None,
budget_duration: str | None = None,
) -> str:
body: Dict[str, object] = {"max_budget": max_budget}
if soft_budget is not None:
body["soft_budget"] = soft_budget
if budget_duration is not None:
body["budget_duration"] = budget_duration
return str(self._post("/budget/new", body)["budget_id"])
return unwrap(
self.gateway.transport.post(
"/budget/new",
headers=self.gateway.transport.master,
json=BudgetNewBody(
max_budget=max_budget,
soft_budget=soft_budget,
budget_duration=budget_duration,
),
response_type=BudgetNewResponse,
)
).budget_id
def delete_budget(self, budget_id: str) -> None:
self._post("/budget/delete", {"id": budget_id})
_ = self.gateway.transport.post(
"/budget/delete",
headers=self.gateway.transport.master,
json=BudgetDeleteBody(id=budget_id),
response_type=NoBody,
)
def budget_info(self, budget_id: str) -> tuple[BudgetRow, ...]:
resp = requests.post(
f"{self._base_url}/budget/info",
headers=auth_headers(self._master_key),
json={"budgets": [budget_id]},
timeout=self._request_timeout,
result = self.gateway.transport.post(
"/budget/info",
headers=self.gateway.transport.master,
json=BudgetInfoBody(budgets=[budget_id]),
response_type=BudgetInfoResponse,
)
resp.raise_for_status()
return _BUDGET_ROWS.validate_python(resp.json())
match result:
case Success(data=data):
return tuple(data.root)
case _:
return ()
def build_client() -> BudgetClient:
return BudgetClient(**proxy_client_kwargs())
return BudgetClient(gateway=build_gateway())

View file

@ -1,9 +1,9 @@
"""Budgets suite's `client` fixture.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
live in the parent tests/e2e/conftest.py. BudgetClient subclasses
ProxyClient, so it satisfies lifecycle.ResourceClient and the shared `resources`
fixture cleans up keys; tests register entity deletes via `resources.defer(...)`.
live in the parent tests/e2e/conftest.py. BudgetClient holds the shared Gateway,
so the `resources` fixture cleans up keys through it; tests register entity deletes
via `resources.defer(...)`.
"""
import pytest

View file

@ -27,7 +27,7 @@ def test_budget_crud_roundtrip(client: BudgetClient, resources: ResourceManager)
assert row.budget_reset_at, "budget_duration did not schedule a reset"
# Attach the budget to a key and confirm the key reflects it.
key = client.generate_key(extra_params={"budget_id": budget_id})
key = client.generate_key(budget_id=budget_id)
resources.defer(lambda: client.delete_key(key))
info = client.gateway.key_info(key)
linked = info.litellm_budget_table
@ -43,7 +43,7 @@ def test_budget_delete_removes_it(client: BudgetClient, resources: ResourceManag
def test_budget_duration_schedules_reset_on_key(client: BudgetClient, resources: ResourceManager) -> None:
key = client.generate_key(max_budget=10.0, extra_params={"budget_duration": "30d"})
key = client.generate_key(max_budget=10.0, budget_duration="30d")
resources.defer(lambda: client.delete_key(key))
reset_at = client.gateway.key_info(key).budget_reset_at

View file

@ -12,13 +12,14 @@ enforcement is broken -> fail.
import time
from dataclasses import dataclass, field
from typing import Callable, Dict, List, Type
from typing import Callable, List, Type
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import run_case
from proxy_client import require_successful_call, unique_marker
pytestmark = pytest.mark.e2e
@ -27,11 +28,14 @@ def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") ->
block within a couple calls off real-time reservation counters; the end-user
budget enforces off table spend that lands on the batch write, so it takes a
few more. A non-budget error fails hard (never a skip)."""
extra: Dict[str, object] = {"max_tokens": 16}
if user:
extra["user"] = user
for _ in range(40):
result = client.chat(key, "claude-haiku-4-5", f"spend {unique_marker()}", extra_body=extra)
result = client.chat(
key,
"claude-haiku-4-5",
f"spend {unique_marker()}",
max_tokens=16,
user=user or None,
)
if is_budget_block(result):
return
require_successful_call(result)
@ -75,7 +79,7 @@ class InternalUserBudgetCase(_BudgetCase):
user_id = self.client.create_user(max_budget=3e-6)
self._undo.append(lambda: self.client.delete_user(user_id))
# personal key (no team) -> the user budget governs
self.key = self.client.generate_key(extra_params={"user_id": user_id})
self.key = self.client.generate_key(user_id=user_id)
self._undo.append(lambda: self.client.delete_key(self.key))
@ -104,7 +108,7 @@ class OrganizationBudgetCase(_BudgetCase):
alias=f"e2e-budget-team-{unique_marker()}", organization_id=org_id
)
self._undo.append(lambda: self.client.delete_team(team_id))
self.key = self.client.generate_key(extra_params={"team_id": team_id})
self.key = self.client.generate_key(team_id=team_id)
self._undo.append(lambda: self.client.delete_key(self.key))
@ -119,9 +123,7 @@ class TeamMemberBudgetCase(_BudgetCase):
user_id = self.client.create_user(max_budget=100.0)
self._undo.append(lambda: self.client.delete_user(user_id))
self.client.add_team_member(team_id, user_id, max_budget_in_team=3e-6)
self.key = self.client.generate_key(
extra_params={"team_id": team_id, "user_id": user_id}
)
self.key = self.client.generate_key(team_id=team_id, user_id=user_id)
self._undo.append(lambda: self.client.delete_key(self.key))

View file

@ -13,22 +13,23 @@ import time
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
from proxy_client import require_successful_call, unique_marker
pytestmark = pytest.mark.e2e
def _call(client: BudgetClient, key: str):
return client.chat(
key, "claude-haiku-4-5", f"reset {unique_marker()}", extra_body={"max_tokens": 16}
key, "claude-haiku-4-5", f"reset {unique_marker()}", max_tokens=16
)
def test_key_budget_resets_after_duration(
client: BudgetClient, resources: ResourceManager
) -> None:
key = client.generate_key(max_budget=3e-6, extra_params={"budget_duration": "30s"})
key = client.generate_key(max_budget=3e-6, budget_duration="30s")
resources.defer(lambda: client.delete_key(key))
# 1. exceed the budget -> litellm returns budget_exceeded

View file

@ -11,8 +11,9 @@ import time
import pytest
from budget_client import BudgetClient, is_budget_block, model_budget
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
from proxy_client import require_successful_call, unique_marker
pytestmark = pytest.mark.e2e
@ -21,9 +22,7 @@ FREE_MODEL = "gemini-2.5-flash"
def _call(client: BudgetClient, key: str, model: str):
result = client.chat(
key, model, f"hi {unique_marker()}", extra_body={"max_tokens": 16}
)
result = client.chat(key, model, f"hi {unique_marker()}", max_tokens=16)
if not result.ok and not is_budget_block(result):
require_successful_call(result) # non-budget error -> skip
return result
@ -33,11 +32,9 @@ def test_model_max_budget_isolates_per_model(
client: BudgetClient, resources: ResourceManager
) -> None:
key = client.generate_key(
extra_params={
"model_max_budget": {
**model_budget(CAPPED_MODEL, 1e-6),
**model_budget(FREE_MODEL, 1000.0),
}
model_max_budget={
**model_budget(CAPPED_MODEL, 1e-6),
**model_budget(FREE_MODEL, 1000.0),
}
)
resources.defer(lambda: client.delete_key(key))

View file

@ -13,8 +13,10 @@ import time
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
from proxy_client import require_successful_call, unique_marker
from models import BudgetWindow
pytestmark = pytest.mark.e2e
@ -23,7 +25,7 @@ WINDOW_SECONDS = 30 # the tight window; calls succeed again only after it elaps
def _call(client: BudgetClient, key: str):
return client.chat(
key, "claude-haiku-4-5", f"window {unique_marker()}", extra_body={"max_tokens": 16}
key, "claude-haiku-4-5", f"window {unique_marker()}", max_tokens=16
)
@ -31,12 +33,10 @@ def test_short_window_blocks_then_resets(
client: BudgetClient, resources: ResourceManager
) -> None:
key = client.generate_key(
extra_params={
"budget_limits": [
{"budget_duration": f"{WINDOW_SECONDS}s", "max_budget": 3e-6},
{"budget_duration": "1m", "max_budget": 1.0}, # roomy: never blocks
]
}
budget_limits=[
BudgetWindow(budget_duration=f"{WINDOW_SECONDS}s", max_budget=3e-6),
BudgetWindow(budget_duration="1m", max_budget=1.0), # roomy: never blocks
]
)
resources.defer(lambda: client.delete_key(key))

View file

@ -10,8 +10,9 @@ load-bearing behavior: soft != block.
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
from proxy_client import require_successful_call, unique_marker
pytestmark = pytest.mark.e2e
@ -20,14 +21,12 @@ def test_soft_budget_does_not_block(
client: BudgetClient, resources: ResourceManager
) -> None:
# soft far below max: spend crosses soft immediately, stays under max.
key = client.generate_key(
max_budget=1000.0, extra_params={"soft_budget": 1e-9}
)
key = client.generate_key(max_budget=1000.0, soft_budget=1e-9)
resources.defer(lambda: client.delete_key(key))
for _ in range(3):
result = client.chat(
key, "claude-haiku-4-5", f"hi {unique_marker()}", extra_body={"max_tokens": 16}
key, "claude-haiku-4-5", f"hi {unique_marker()}", max_tokens=16
)
require_successful_call(result) # skip if provider unavailable
assert not is_budget_block(result), (

View file

@ -11,8 +11,9 @@ import time
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
from proxy_client import require_successful_call, unique_marker
pytestmark = pytest.mark.e2e
@ -24,8 +25,8 @@ def _tagged_call(client: BudgetClient, key: str, tag: str):
key,
"claude-haiku-4-5",
f"hi {unique_marker()}",
metadata={"tags": [tag]},
extra_body={"max_tokens": 16},
tags=[tag],
max_tokens=16,
)
if not result.ok and not is_budget_block(result):
require_successful_call(result) # non-budget error -> skip

View file

@ -12,6 +12,7 @@ from __future__ import annotations
from typing import Generic, Iterator, Literal, NewType, TypeVar, cast
import pytest
import requests
from pydantic import BaseModel, ConfigDict, Field
@ -137,11 +138,28 @@ def is_ok[R: BaseModel](result: Result[R]) -> bool:
return False
def require_successful_call(result: StreamingResponse) -> 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}); body={result.body[:300]}"
)
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 _params(params: BaseModel | None) -> dict[str, str]:
if params is None:
return {}
dumped: dict[str, object] = params.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]:
@ -232,24 +250,10 @@ def probe(
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))
def _streaming_outcome(resp: requests.Response, stream: bool) -> StreamingResponse:
call_id = _hdr(resp, "x-litellm-call-id")
content_type = _hdr(resp, "content-type")
if not (200 <= resp.status_code < 300):
if not stream or not (200 <= resp.status_code < 300):
return StreamingResponse(
status_code=resp.status_code,
call_id=call_id,
@ -265,3 +269,38 @@ def stream(
body="<streamed>",
chunks=chunks,
)
def send(
url: URL,
*,
headers: BaseModel,
json: BaseModel,
params: BaseModel | None = None,
stream: bool = False,
timeout: float = 60.0,
) -> StreamingResponse:
"""Raw POST returning the unparsed HTTP outcome: status, full body, and the
x-litellm-call-id header. For native/passthrough bodies and for calls judged by
status rather than a typed JSON model (e.g. a budget block is a non-2xx). With
``stream=True`` the SSE body is consumed and its events counted instead."""
try:
resp = requests.post(
str(url),
headers=_headers(headers),
params=_params(params),
json=json.model_dump(by_alias=True, exclude_none=True),
stream=stream,
timeout=timeout,
)
except requests.RequestException as exc:
return StreamingResponse(status_code=-1, body=str(exc))
return _streaming_outcome(resp, stream)
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."""
return send(url, headers=headers, json=json, stream=True, timeout=timeout)

View file

@ -1,9 +1,8 @@
"""LLM-translation suite's `client` fixture.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
live in the parent tests/e2e/conftest.py. PassthroughClient (via ProxyClient)
exposes the shared Gateway, so the `resources` fixture cleans up keys this suite
creates.
live in the parent tests/e2e/conftest.py. PassthroughClient holds the shared
Gateway, so the `resources` fixture cleans up keys this suite creates.
"""
import pytest

View file

@ -1,49 +1,105 @@
"""Client for LLM-translation e2e tests over the proxy's passthrough endpoints.
Extends the shared ProxyClient with native provider passthrough calls. A
passthrough request is sent in the PROVIDER's native format (Gemini
A passthrough request is sent in the PROVIDER's native format (Gemini
generateContent, Anthropic /v1/messages) to the proxy, which forwards it to the
provider and still logs a SpendLogs row (call_type="pass_through_endpoint"). The
litellm virtual key is passed as the provider key; the proxy swaps in the real
env credential. SpendLogs.request_id == the x-litellm-call-id response header.
litellm virtual key is passed as the provider key; the proxy swaps in the real env
credential. SpendLogs.request_id == the x-litellm-call-id response header. The
native request models are co-located here because only this suite uses them.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Dict, Iterator, List, Optional, cast
import requests
from pydantic import BaseModel, Field
from e2e_config import (
MASTER_KEY,
POLL_INTERVAL,
POLL_TIMEOUT,
PROXY_BASE_URL,
REQUEST_TIMEOUT,
)
from proxy_client import ProxyClient
from e2e_gateway import Gateway, build_gateway
from e2e_http import Headers, StreamingResponse
from models import ChatMessage
Tools = List[Dict[str, object]]
class JsonSchemaProperty(BaseModel):
type: str
class JsonSchema(BaseModel):
type: str
properties: dict[str, JsonSchemaProperty]
required: list[str]
class GeminiHeaders(Headers):
x_goog_api_key: str = Field(serialization_alias="x-goog-api-key")
content_type: str = Field(
default="application/json", serialization_alias="Content-Type"
)
tags: str | None = None
class AnthropicHeaders(Headers):
x_api_key: str = Field(serialization_alias="x-api-key")
anthropic_version: str = Field(
default="2023-06-01", serialization_alias="anthropic-version"
)
content_type: str = Field(
default="application/json", serialization_alias="Content-Type"
)
tags: str | None = None
class AltSseParams(BaseModel):
alt: str = "sse"
class GeminiPart(BaseModel):
text: str
class GeminiContent(BaseModel):
role: str = "user"
parts: list[GeminiPart]
class GeminiFunctionDeclaration(BaseModel):
name: str
description: str
parameters: JsonSchema
class GeminiTool(BaseModel):
function_declarations: list[GeminiFunctionDeclaration] = Field(
serialization_alias="functionDeclarations"
)
class GeminiGenerateBody(BaseModel):
contents: list[GeminiContent]
tools: list[GeminiTool] | None = None
class AnthropicTool(BaseModel):
name: str
description: str
input_schema: JsonSchema
class AnthropicMessageBody(BaseModel):
model: str
max_tokens: int
messages: list[ChatMessage]
tools: list[AnthropicTool] | None = None
stream: bool = False
def _tags_header(tags: list[str] | None) -> str | None:
return ",".join(tags) if tags else None
@dataclass(frozen=True, slots=True)
class PassthroughResult:
"""Outcome of a native passthrough call. ``call_id`` correlates to the row."""
class PassthroughClient:
gateway: Gateway
status_code: int
call_id: Optional[str] # x-litellm-call-id -> SpendLogs.request_id
body: str
chunks: int = 0 # number of streamed events (0 for non-streaming)
@property
def ok(self) -> bool:
return 200 <= self.status_code < 300
def _tag_header(tags: Optional[List[str]]) -> Dict[str, str]:
return {"tags": ",".join(tags)} if tags else {}
class PassthroughClient(ProxyClient):
# ---- Gemini native passthrough (/gemini/v1beta/...) -----------------
def gemini_generate(
@ -52,48 +108,29 @@ class PassthroughClient(ProxyClient):
model: str,
text: str,
*,
tools: Optional[Tools] = None,
tags: Optional[List[str]] = None,
) -> PassthroughResult:
body: Dict[str, object] = {"contents": [{"role": "user", "parts": [{"text": text}]}]}
if tools is not None:
body["tools"] = tools
headers = {
"x-goog-api-key": key,
"Content-Type": "application/json",
**_tag_header(tags),
}
resp = requests.post(
f"{self._base_url}/gemini/v1beta/models/{model}:generateContent",
headers=headers,
json=body,
timeout=self._request_timeout,
)
return PassthroughResult(
resp.status_code, resp.headers.get("x-litellm-call-id"), resp.text
tools: list[GeminiTool] | None = None,
tags: list[str] | None = None,
) -> StreamingResponse:
return self.gateway.transport.send(
f"/gemini/v1beta/models/{model}:generateContent",
headers=GeminiHeaders(x_goog_api_key=key, tags=_tags_header(tags)),
json=GeminiGenerateBody(
contents=[GeminiContent(parts=[GeminiPart(text=text)])], tools=tools
),
)
def gemini_stream(
self, key: str, model: str, text: str, *, tags: Optional[List[str]] = None
) -> PassthroughResult:
headers = {
"x-goog-api-key": key,
"Content-Type": "application/json",
**_tag_header(tags),
}
resp = requests.post(
f"{self._base_url}/gemini/v1beta/models/{model}:streamGenerateContent",
headers=headers,
params={"alt": "sse"},
json={"contents": [{"role": "user", "parts": [{"text": text}]}]},
self, key: str, model: str, text: str, *, tags: list[str] | None = None
) -> StreamingResponse:
return self.gateway.transport.send(
f"/gemini/v1beta/models/{model}:streamGenerateContent",
headers=GeminiHeaders(x_goog_api_key=key, tags=_tags_header(tags)),
json=GeminiGenerateBody(
contents=[GeminiContent(parts=[GeminiPart(text=text)])]
),
params=AltSseParams(),
stream=True,
timeout=self._request_timeout,
)
call_id = resp.headers.get("x-litellm-call-id")
if not (200 <= resp.status_code < 300):
return PassthroughResult(resp.status_code, call_id, resp.text)
chunks = sum(1 for line in cast("Iterator[bytes]", resp.iter_lines()) if line)
return PassthroughResult(resp.status_code, call_id, "<streamed>", chunks)
# ---- Anthropic native passthrough (/anthropic/v1/messages) ----------
@ -104,48 +141,23 @@ class PassthroughClient(ProxyClient):
text: str,
*,
max_tokens: int = 64,
tools: Optional[Tools] = None,
tools: list[AnthropicTool] | None = None,
stream: bool = False,
tags: Optional[List[str]] = None,
) -> PassthroughResult:
body: Dict[str, object] = {
"model": model,
"max_tokens": max_tokens,
"messages": [{"role": "user", "content": text}],
}
if tools is not None:
body["tools"] = tools
if stream:
body["stream"] = True
headers = {
"x-api-key": key,
"anthropic-version": "2023-06-01",
"Content-Type": "application/json",
**_tag_header(tags),
}
url = f"{self._base_url}/anthropic/v1/messages"
if not stream:
resp = requests.post(
url, headers=headers, json=body, timeout=self._request_timeout
)
return PassthroughResult(
resp.status_code, resp.headers.get("x-litellm-call-id"), resp.text
)
resp = requests.post(
url, headers=headers, json=body, stream=True, timeout=self._request_timeout
tags: list[str] | None = None,
) -> StreamingResponse:
return self.gateway.transport.send(
"/anthropic/v1/messages",
headers=AnthropicHeaders(x_api_key=key, tags=_tags_header(tags)),
json=AnthropicMessageBody(
model=model,
max_tokens=max_tokens,
messages=[ChatMessage(role="user", content=text)],
tools=tools,
stream=stream,
),
stream=stream,
)
call_id = resp.headers.get("x-litellm-call-id")
if not (200 <= resp.status_code < 300):
return PassthroughResult(resp.status_code, call_id, resp.text)
chunks = sum(1 for line in cast("Iterator[bytes]", resp.iter_lines()) if line)
return PassthroughResult(resp.status_code, call_id, "<streamed>", chunks)
def build_client() -> PassthroughClient:
return PassthroughClient(
base_url=PROXY_BASE_URL,
master_key=MASTER_KEY,
request_timeout=REQUEST_TIMEOUT,
poll_timeout=POLL_TIMEOUT,
poll_interval=POLL_INTERVAL,
)
return PassthroughClient(gateway=build_gateway())

View file

@ -13,14 +13,22 @@ A passthrough call returning non-2xx fails hard (never a skip); once it returns
import pytest
from e2e_config import unique_marker
from e2e_http import StreamingResponse, require_successful_call
from models import SpendLogRow
from passthrough_client import PassthroughClient, PassthroughResult
from proxy_client import require_successful_call, unique_marker
from passthrough_client import (
AnthropicTool,
GeminiFunctionDeclaration,
GeminiTool,
JsonSchema,
JsonSchemaProperty,
PassthroughClient,
)
pytestmark = pytest.mark.e2e
def _fetch_cost_breakdown(client: PassthroughClient, result: PassthroughResult) -> SpendLogRow:
def _fetch_cost_breakdown(client: PassthroughClient, result: StreamingResponse) -> SpendLogRow:
"""The passthrough call's logged row, polled until it carries a cost.
Asserts (not skips) that a 2xx passthrough call produced a costed row - the
@ -76,19 +84,19 @@ def test_gemini_passthrough_tool_call_logs_cost(
"gemini-2.5-flash",
"What is the weather in Paris? Use the get_weather tool.",
tools=[
{
"functionDeclarations": [
{
"name": "get_weather",
"description": "Get the weather for a city",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
}
GeminiTool(
function_declarations=[
GeminiFunctionDeclaration(
name="get_weather",
description="Get the weather for a city",
parameters=JsonSchema(
type="object",
properties={"city": JsonSchemaProperty(type="string")},
required=["city"],
),
)
]
}
)
],
)
require_successful_call(result)
@ -133,15 +141,15 @@ def test_anthropic_passthrough_tool_call_logs_cost(
"claude-haiku-4-5",
"What is the weather in Paris? Use the get_weather tool.",
tools=[
{
"name": "get_weather",
"description": "Get the weather for a city",
"input_schema": {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
}
AnthropicTool(
name="get_weather",
description="Get the weather for a city",
input_schema=JsonSchema(
type="object",
properties={"city": JsonSchemaProperty(type="string")},
required=["city"],
),
)
],
)
require_successful_call(result)

View file

@ -1,397 +0,0 @@
"""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 < 500 and self.status_code != 404
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]}"
)

View file

@ -25,6 +25,16 @@ class Transport(Protocol):
self, path: str, *, headers: BaseModel, json: BaseModel
) -> StreamingResponse: ...
def send(
self,
path: str,
*,
headers: BaseModel,
json: BaseModel,
params: BaseModel | None = None,
stream: bool = False,
) -> StreamingResponse: ...
def get[R: BaseModel](
self,
path: str,
@ -107,6 +117,24 @@ class HttpTransport:
self._url(path), headers=headers, json=json, timeout=self.request_timeout
)
def send(
self,
path: str,
*,
headers: BaseModel,
json: BaseModel,
params: BaseModel | None = None,
stream: bool = False,
) -> StreamingResponse:
return e2e_http.send(
self._url(path),
headers=headers,
json=json,
params=params,
stream=stream,
timeout=self.request_timeout,
)
def probe(self, path: str, *, params: BaseModel) -> ProbeResult:
return e2e_http.probe(
self._url(path),