diff --git a/tests/e2e/budgets/budget_client.py b/tests/e2e/budgets/budget_client.py index 1f2c6039973..cfe70a605f9 100644 --- a/tests/e2e/budgets/budget_client.py +++ b/tests/e2e/budgets/budget_client.py @@ -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()) diff --git a/tests/e2e/budgets/conftest.py b/tests/e2e/budgets/conftest.py index 88c4cae478f..236822f4309 100644 --- a/tests/e2e/budgets/conftest.py +++ b/tests/e2e/budgets/conftest.py @@ -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 diff --git a/tests/e2e/budgets/test_budget_crud_e2e.py b/tests/e2e/budgets/test_budget_crud_e2e.py index fe4104328ad..7e0c614638d 100644 --- a/tests/e2e/budgets/test_budget_crud_e2e.py +++ b/tests/e2e/budgets/test_budget_crud_e2e.py @@ -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 diff --git a/tests/e2e/budgets/test_budget_enforcement_e2e.py b/tests/e2e/budgets/test_budget_enforcement_e2e.py index c68cdf872ee..ea2fba47331 100644 --- a/tests/e2e/budgets/test_budget_enforcement_e2e.py +++ b/tests/e2e/budgets/test_budget_enforcement_e2e.py @@ -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)) diff --git a/tests/e2e/budgets/test_budget_reset_e2e.py b/tests/e2e/budgets/test_budget_reset_e2e.py index 35923e71410..4727bd8f9cf 100644 --- a/tests/e2e/budgets/test_budget_reset_e2e.py +++ b/tests/e2e/budgets/test_budget_reset_e2e.py @@ -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 diff --git a/tests/e2e/budgets/test_model_max_budget_e2e.py b/tests/e2e/budgets/test_model_max_budget_e2e.py index 7ed3e0c2834..34616624c12 100644 --- a/tests/e2e/budgets/test_model_max_budget_e2e.py +++ b/tests/e2e/budgets/test_model_max_budget_e2e.py @@ -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)) diff --git a/tests/e2e/budgets/test_multi_window_budget_e2e.py b/tests/e2e/budgets/test_multi_window_budget_e2e.py index b9200a1ea32..3a3720302c8 100644 --- a/tests/e2e/budgets/test_multi_window_budget_e2e.py +++ b/tests/e2e/budgets/test_multi_window_budget_e2e.py @@ -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)) diff --git a/tests/e2e/budgets/test_soft_budget_e2e.py b/tests/e2e/budgets/test_soft_budget_e2e.py index 51b1a99a1dc..85ff0c25792 100644 --- a/tests/e2e/budgets/test_soft_budget_e2e.py +++ b/tests/e2e/budgets/test_soft_budget_e2e.py @@ -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), ( diff --git a/tests/e2e/budgets/test_tag_budget_e2e.py b/tests/e2e/budgets/test_tag_budget_e2e.py index 77a3ff35eca..6f71b1b204b 100644 --- a/tests/e2e/budgets/test_tag_budget_e2e.py +++ b/tests/e2e/budgets/test_tag_budget_e2e.py @@ -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 diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py index 3a5973cd8b0..7458f316852 100644 --- a/tests/e2e/e2e_http.py +++ b/tests/e2e/e2e_http.py @@ -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="", 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) diff --git a/tests/e2e/llm_translation/conftest.py b/tests/e2e/llm_translation/conftest.py index e747a9bf3b8..fbf008cf085 100644 --- a/tests/e2e/llm_translation/conftest.py +++ b/tests/e2e/llm_translation/conftest.py @@ -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 diff --git a/tests/e2e/llm_translation/passthrough_client.py b/tests/e2e/llm_translation/passthrough_client.py index 557a337dee5..fff4064a328 100644 --- a/tests/e2e/llm_translation/passthrough_client.py +++ b/tests/e2e/llm_translation/passthrough_client.py @@ -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, "", 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, "", 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()) diff --git a/tests/e2e/llm_translation/test_passthrough_e2e.py b/tests/e2e/llm_translation/test_passthrough_e2e.py index 9c5010b637e..37d55c665b3 100644 --- a/tests/e2e/llm_translation/test_passthrough_e2e.py +++ b/tests/e2e/llm_translation/test_passthrough_e2e.py @@ -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) diff --git a/tests/e2e/proxy_client.py b/tests/e2e/proxy_client.py deleted file mode 100644 index 151aa37a594..00000000000 --- a/tests/e2e/proxy_client.py +++ /dev/null @@ -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="", - 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 index e783c2922ed..65e093993c4 100644 --- a/tests/e2e/transport.py +++ b/tests/e2e/transport.py @@ -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),