mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
refactor: migrate to gateway client
This commit is contained in:
parent
6843ca0325
commit
ff7c3c55c7
15 changed files with 566 additions and 686 deletions
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
|
|
|||
|
|
@ -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), (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]}"
|
||||
)
|
||||
|
|
@ -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),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue