mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
test(integration): audit durable background interaction settlement across replicas
Twenty-six deterministic cells drive a one-worker creator and a two-worker settler against an owned scripted Gemini upstream: cross-replica deletes bill once, failed and cancelled interactions release, a later replica resumes unclaimed rows, custom deployment pricing bills at the deployment rate, a fetch the settler cannot make fails the delete closed, odd ids are refused, a missing settlement table keeps in-process billing, polling disabled registers nothing, the budget reservation is released by the settler, an upstream outage mid-burst fails closed and recovers, killed workers hand their polls to the respawned ones, and concurrent deletes on a slow upstream settle exactly once. The support upstream gains a scripted interaction store with per-id GET status and delay, and the process helper gains an owned upstream a test can stop and restart
This commit is contained in:
parent
44c3026aeb
commit
063c3dec77
3 changed files with 989 additions and 20 deletions
|
|
@ -5,7 +5,7 @@ import subprocess
|
|||
import sys
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping
|
||||
from collections.abc import Generator, Iterator, Mapping
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
|
@ -220,3 +220,65 @@ def owned_proxy_process(
|
|||
yield OwnedProxy(Gateway(client, gateway.key, gateway.upstream_url), process, launch.log)
|
||||
finally:
|
||||
_stop(process)
|
||||
|
||||
|
||||
_UPSTREAM_READY_SECONDS: Final = 60
|
||||
|
||||
|
||||
class UpstreamSlot:
|
||||
"""A scripted upstream a test module owns on a fixed port, so a cell can take it down and bring it back."""
|
||||
|
||||
__slots__ = ("directory", "port", "process", "root")
|
||||
|
||||
def __init__(self, directory: Path, port: int, root: Path) -> None:
|
||||
self.directory = directory
|
||||
self.port = port
|
||||
self.root = root
|
||||
self.process: subprocess.Popen[bytes] | None = None
|
||||
|
||||
@property
|
||||
def url(self) -> str:
|
||||
return f"http://127.0.0.1:{self.port}"
|
||||
|
||||
def start(self) -> None:
|
||||
assert self.process is None, "Owned upstream is already running"
|
||||
output: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR") or self.directory)
|
||||
log_path: Final = output / f"owned-upstream-{self.port}-{uuid.uuid4().hex}.log"
|
||||
with log_path.open("w") as log:
|
||||
process: Final = subprocess.Popen(
|
||||
[sys.executable, "-m", "integration._support.upstream", "--port", str(self.port)],
|
||||
cwd=self.root,
|
||||
env=dict(os.environ),
|
||||
stdout=log,
|
||||
stderr=subprocess.STDOUT,
|
||||
start_new_session=True,
|
||||
)
|
||||
self.process = process
|
||||
deadline: Final = time.monotonic() + _UPSTREAM_READY_SECONDS
|
||||
while process.poll() is None:
|
||||
try:
|
||||
if httpx.get(f"{self.url}/health", timeout=2, trust_env=False).status_code == 200:
|
||||
return
|
||||
except httpx.TransportError:
|
||||
pass
|
||||
assert time.monotonic() < deadline, f"Owned upstream readiness deadline exceeded: {log_path}"
|
||||
time.sleep(0.1)
|
||||
raise AssertionError(f"Owned upstream exited before readiness: {log_path}")
|
||||
|
||||
def stop(self) -> None:
|
||||
process: Final = self.process
|
||||
assert process is not None, "Owned upstream is not running"
|
||||
self.process = None
|
||||
_stop(process)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def owned_upstream(directory: Path) -> Generator[UpstreamSlot]:
|
||||
root: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3])
|
||||
slot: Final = UpstreamSlot(directory, _free_port(), root)
|
||||
slot.start()
|
||||
try:
|
||||
yield slot
|
||||
finally:
|
||||
if slot.process is not None:
|
||||
slot.stop()
|
||||
|
|
|
|||
|
|
@ -29,7 +29,7 @@ from integration.cost_calculation.cost_tracking_case import (
|
|||
StoredResponse,
|
||||
TextResponse,
|
||||
)
|
||||
from pydantic import BaseModel, JsonValue, TypeAdapter, ValidationError
|
||||
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError
|
||||
from starlette.applications import Starlette
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import JSONResponse, Response, StreamingResponse
|
||||
|
|
@ -64,6 +64,19 @@ class Observation:
|
|||
path: str
|
||||
authorization: str
|
||||
body: dict[str, JsonValue]
|
||||
method: str = "POST"
|
||||
api_key: str = ""
|
||||
|
||||
|
||||
class InteractionState(BaseModel):
|
||||
"""What the scripted Interactions API answers for one interaction id until a DELETE drops it."""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
status: str
|
||||
usage: dict[str, JsonValue] | None = None
|
||||
get_status: int = 200
|
||||
delay_seconds: float = 0
|
||||
|
||||
|
||||
class _ScenarioRegistration(BaseModel):
|
||||
|
|
@ -89,9 +102,12 @@ def _aws_event_frame(
|
|||
scenario_id: str,
|
||||
unique_id: str,
|
||||
) -> bytes:
|
||||
payload_bytes: Final = json.dumps(payload, separators=(",", ":")).replace(
|
||||
"$REQUEST_ID", scenario_id
|
||||
).replace("$UNIQUE_ID", unique_id).encode()
|
||||
payload_bytes: Final = (
|
||||
json.dumps(payload, separators=(",", ":"))
|
||||
.replace("$REQUEST_ID", scenario_id)
|
||||
.replace("$UNIQUE_ID", unique_id)
|
||||
.encode()
|
||||
)
|
||||
headers_bytes: Final = (
|
||||
_aws_str_header(":event-type", event_type)
|
||||
+ _aws_str_header(":content-type", "application/json")
|
||||
|
|
@ -123,6 +139,7 @@ class Provider:
|
|||
observations: SimpleQueue[Observation] = field(default_factory=SimpleQueue)
|
||||
scripts: dict[str, deque[int]] = field(default_factory=dict)
|
||||
scenario_store: ScenarioStore = field(default_factory=ScenarioStore)
|
||||
interactions: dict[str, InteractionState] = field(default_factory=dict)
|
||||
|
||||
async def chat(self, request: Request) -> Response:
|
||||
body: Final = JSON_OBJECT.validate_json(await request.body())
|
||||
|
|
@ -147,7 +164,13 @@ class Provider:
|
|||
status: Final = script.popleft()
|
||||
if status != 200:
|
||||
return JSONResponse(
|
||||
{"error": {"message": "Controlled provider failure", "type": error_type(status), "code": str(status)}},
|
||||
{
|
||||
"error": {
|
||||
"message": "Controlled provider failure",
|
||||
"type": error_type(status),
|
||||
"code": str(status),
|
||||
}
|
||||
},
|
||||
status_code=status,
|
||||
)
|
||||
return await chat_completions(request)
|
||||
|
|
@ -198,7 +221,14 @@ class Provider:
|
|||
return JSONResponse(
|
||||
{
|
||||
"requests": [
|
||||
{"path": value.path, "authorization": value.authorization, "body": value.body} for value in values
|
||||
{
|
||||
"path": value.path,
|
||||
"authorization": value.authorization,
|
||||
"body": value.body,
|
||||
"method": value.method,
|
||||
"api_key": value.api_key,
|
||||
}
|
||||
for value in values
|
||||
]
|
||||
}
|
||||
)
|
||||
|
|
@ -245,7 +275,12 @@ class Provider:
|
|||
body: Final = JSON_OBJECT.validate_json(raw_body)
|
||||
if isinstance(body, dict):
|
||||
self.observations.put(
|
||||
Observation(request.url.path, request.headers.get("authorization", ""), body)
|
||||
Observation(
|
||||
request.url.path,
|
||||
request.headers.get("authorization", ""),
|
||||
body,
|
||||
api_key=request.headers.get("x-goog-api-key", ""),
|
||||
)
|
||||
)
|
||||
if isinstance(response, RoutedResponse):
|
||||
route_key: Final = f"{request.method} /{'/'.join(segments[1:])}"
|
||||
|
|
@ -262,6 +297,57 @@ class Provider:
|
|||
return self._response(route, scenario_id)
|
||||
return self._response(response, scenario_id)
|
||||
|
||||
async def interaction_state(self, request: Request) -> Response:
|
||||
interaction_id: Final = cast(str, request.path_params["interaction_id"])
|
||||
if request.method == "DELETE":
|
||||
self.interactions.pop(interaction_id, None)
|
||||
return JSONResponse({"interaction_id": interaction_id, "registered": False})
|
||||
try:
|
||||
state: Final = InteractionState.model_validate_json(await request.body())
|
||||
except ValidationError as exc:
|
||||
return JSONResponse({"error": str(exc)}, status_code=400)
|
||||
self.interactions[interaction_id] = state
|
||||
return JSONResponse({"interaction_id": interaction_id, "registered": True})
|
||||
|
||||
def _observe_interaction(self, request: Request) -> None:
|
||||
self.observations.put(
|
||||
Observation(
|
||||
request.url.path,
|
||||
request.headers.get("authorization", ""),
|
||||
{},
|
||||
method=request.method,
|
||||
api_key=request.headers.get("x-goog-api-key", ""),
|
||||
)
|
||||
)
|
||||
|
||||
async def interaction(self, request: Request) -> Response:
|
||||
self._observe_interaction(request)
|
||||
interaction_id: Final = cast(str, request.path_params["interaction_id"])
|
||||
state: Final = self.interactions.get(interaction_id)
|
||||
if state is None:
|
||||
return JSONResponse(_interaction_not_found(interaction_id), status_code=404)
|
||||
if state.delay_seconds:
|
||||
await asyncio.sleep(state.delay_seconds)
|
||||
if request.method == "DELETE":
|
||||
del self.interactions[interaction_id]
|
||||
return JSONResponse({})
|
||||
if state.get_status != 200:
|
||||
return JSONResponse(
|
||||
{"error": {"code": state.get_status, "message": "Scripted interaction fetch failure"}},
|
||||
status_code=state.get_status,
|
||||
)
|
||||
return JSONResponse(_interaction_body(interaction_id, state))
|
||||
|
||||
async def cancel_interaction(self, request: Request) -> Response:
|
||||
self._observe_interaction(request)
|
||||
interaction_id: Final = cast(str, request.path_params["interaction_id"])
|
||||
state: Final = self.interactions.get(interaction_id)
|
||||
if state is None:
|
||||
return JSONResponse(_interaction_not_found(interaction_id), status_code=404)
|
||||
cancelled: Final = InteractionState(status="cancelled", usage=state.usage, get_status=state.get_status)
|
||||
self.interactions[interaction_id] = cancelled
|
||||
return JSONResponse(_interaction_body(interaction_id, cancelled))
|
||||
|
||||
async def realtime(self, websocket: WebSocket) -> None:
|
||||
scenario_id: Final = websocket.headers.get("authorization", "").removeprefix("Bearer ")
|
||||
response: Final = self.scenario_store.get(scenario_id)
|
||||
|
|
@ -300,11 +386,10 @@ class Provider:
|
|||
match response:
|
||||
case JsonResponse():
|
||||
return Response(
|
||||
content=json.dumps(response.body, separators=(",", ":")).replace(
|
||||
"$REQUEST_ID", scenario_id
|
||||
).replace(
|
||||
"$UNIQUE_ID", unique_id
|
||||
).encode(),
|
||||
content=json.dumps(response.body, separators=(",", ":"))
|
||||
.replace("$REQUEST_ID", scenario_id)
|
||||
.replace("$UNIQUE_ID", unique_id)
|
||||
.encode(),
|
||||
media_type=response.content_type,
|
||||
status_code=response.status,
|
||||
)
|
||||
|
|
@ -321,6 +406,7 @@ class Provider:
|
|||
)
|
||||
case SseResponse():
|
||||
if response.frame_delay_ms > 0:
|
||||
|
||||
async def stream() -> AsyncIterator[bytes]:
|
||||
for frame in response.frames:
|
||||
yield (
|
||||
|
|
@ -329,9 +415,11 @@ class Provider:
|
|||
await asyncio.sleep(response.frame_delay_ms / 1000)
|
||||
|
||||
return StreamingResponse(stream(), media_type=response.content_type)
|
||||
stream_body: Final = ("\n\n".join(response.frames) + "\n\n").replace(
|
||||
"$REQUEST_ID", scenario_id
|
||||
).replace("$UNIQUE_ID", unique_id)
|
||||
stream_body: Final = (
|
||||
("\n\n".join(response.frames) + "\n\n")
|
||||
.replace("$REQUEST_ID", scenario_id)
|
||||
.replace("$UNIQUE_ID", unique_id)
|
||||
)
|
||||
return Response(content=stream_body.encode(), media_type=response.content_type)
|
||||
case EventStreamResponse():
|
||||
events: Final = (
|
||||
|
|
@ -372,6 +460,19 @@ class Provider:
|
|||
Route("/v1/embeddings", embeddings, methods=["POST"]),
|
||||
Route("/v1/moderations", moderations, methods=["POST"]),
|
||||
Route("/vector_stores/{vector_store_id}/search", self.vector_store_search, methods=["POST"]),
|
||||
Route("/__interactions/{interaction_id}", self.interaction_state, methods=["PUT", "DELETE"]),
|
||||
Route("/v1beta/interactions/{interaction_id}:cancel", self.cancel_interaction, methods=["POST"]),
|
||||
Route("/v1beta/interactions/{interaction_id}", self.interaction, methods=["GET", "DELETE"]),
|
||||
Route(
|
||||
"/{prefix:path}/v1beta/interactions/{interaction_id}:cancel",
|
||||
self.cancel_interaction,
|
||||
methods=["POST"],
|
||||
),
|
||||
Route(
|
||||
"/{prefix:path}/v1beta/interactions/{interaction_id}",
|
||||
self.interaction,
|
||||
methods=["GET", "DELETE"],
|
||||
),
|
||||
Route("/{path:path}", self.scripted, methods=["POST"]),
|
||||
Route("/{path:path}", self.scripted, methods=["GET"]),
|
||||
WebSocketRoute("/v1/realtime", self.realtime),
|
||||
|
|
@ -382,6 +483,21 @@ class Provider:
|
|||
CONTROL_URL: Final = os.environ.get("INTEGRATION_UPSTREAM_URL", "http://127.0.0.1:8190").rstrip("/")
|
||||
|
||||
|
||||
def _interaction_not_found(interaction_id: str) -> dict[str, JsonValue]:
|
||||
return {"error": {"code": 404, "message": f"Interaction {interaction_id} not found", "status": "NOT_FOUND"}}
|
||||
|
||||
|
||||
def _interaction_body(interaction_id: str, state: InteractionState) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"id": interaction_id,
|
||||
"object": "interaction",
|
||||
"model": "gemini-3.8-flash",
|
||||
"status": state.status,
|
||||
"steps": [],
|
||||
"usage": state.usage,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ScenarioHandle:
|
||||
scenario_id: str
|
||||
|
|
@ -391,9 +507,9 @@ class ScenarioHandle:
|
|||
return f"{self.control_url}/{self.scenario_id}"
|
||||
|
||||
|
||||
def register_scenario(scenario_id: str, response: StoredResponse) -> ScenarioHandle:
|
||||
def register_scenario(scenario_id: str, response: StoredResponse, *, control_url: str = CONTROL_URL) -> ScenarioHandle:
|
||||
http_response: Final = httpx.post(
|
||||
f"{CONTROL_URL}/__scenarios",
|
||||
f"{control_url}/__scenarios",
|
||||
json={"scenario_id": scenario_id, "response": response.model_dump(mode="json")},
|
||||
trust_env=False,
|
||||
timeout=15,
|
||||
|
|
@ -401,19 +517,35 @@ def register_scenario(scenario_id: str, response: StoredResponse) -> ScenarioHan
|
|||
http_response.raise_for_status()
|
||||
return ScenarioHandle(
|
||||
scenario_id=scenario_id,
|
||||
control_url=CONTROL_URL,
|
||||
control_url=control_url,
|
||||
)
|
||||
|
||||
|
||||
def delete_scenario(handle: ScenarioHandle) -> None:
|
||||
response: Final = httpx.delete(
|
||||
f"{CONTROL_URL}/__scenarios/{handle.scenario_id}",
|
||||
f"{handle.control_url}/__scenarios/{handle.scenario_id}",
|
||||
trust_env=False,
|
||||
timeout=15,
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
|
||||
def set_interaction_state(control_url: str, interaction_id: str, state: InteractionState) -> None:
|
||||
response: Final = httpx.put(
|
||||
f"{control_url}/__interactions/{interaction_id}",
|
||||
content=state.model_dump_json(),
|
||||
headers={"content-type": "application/json"},
|
||||
trust_env=False,
|
||||
timeout=15,
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
|
||||
def clear_interaction_state(control_url: str, interaction_id: str) -> None:
|
||||
response: Final = httpx.delete(f"{control_url}/__interactions/{interaction_id}", trust_env=False, timeout=15)
|
||||
response.raise_for_status()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser: Final = argparse.ArgumentParser()
|
||||
parser.add_argument("--port", type=int, default=8190)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,775 @@
|
|||
import math
|
||||
import socket
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import psutil
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import (
|
||||
JSON_OBJECT,
|
||||
Gateway,
|
||||
eventually,
|
||||
gateway_from_environment,
|
||||
object_value,
|
||||
string_value,
|
||||
)
|
||||
from integration._support.database import read_rows, write_rows
|
||||
from integration._support.process import (
|
||||
UpstreamSlot,
|
||||
group_members,
|
||||
owned_proxy,
|
||||
owned_proxy_process,
|
||||
owned_upstream,
|
||||
)
|
||||
from integration._support.upstream import (
|
||||
InteractionState,
|
||||
clear_interaction_state,
|
||||
register_scenario,
|
||||
set_interaction_state,
|
||||
)
|
||||
from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse
|
||||
from pydantic import JsonValue
|
||||
|
||||
from litellm.proxy.spend_tracking.budget_reservation import DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(900)
|
||||
|
||||
_MODEL: Final = "gemini/gemini-3.8-flash"
|
||||
_INPUT_TOKENS: Final = 300
|
||||
_OUTPUT_TOKENS: Final = 41
|
||||
_USAGE: Final[dict[str, JsonValue]] = {
|
||||
"total_input_tokens": _INPUT_TOKENS,
|
||||
"total_output_tokens": _OUTPUT_TOKENS,
|
||||
"total_tool_use_tokens": 0,
|
||||
"total_reasoning_tokens": 0,
|
||||
}
|
||||
_CUSTOM_INPUT_RATE: Final = 2e-06
|
||||
_CUSTOM_OUTPUT_RATE: Final = 4e-05
|
||||
_ENV_KEY: Final = "integration-gemini-env-key"
|
||||
_DEPLOYMENT_KEY: Final = "integration-gemini-deployment-key"
|
||||
_CREATOR_POLL: Final = {"BACKGROUND_INTERACTION_COST_POLL_INITIAL_INTERVAL_SECONDS": "300"}
|
||||
_SETTLER_POLL: Final = {
|
||||
"BACKGROUND_INTERACTION_COST_POLL_INITIAL_INTERVAL_SECONDS": "1",
|
||||
"BACKGROUND_INTERACTION_COST_POLL_MAX_INTERVAL_SECONDS": "1",
|
||||
"BACKGROUND_INTERACTION_COST_POLL_TIMEOUT_SECONDS": "8",
|
||||
}
|
||||
_RESUMER_POLL: Final = {
|
||||
"BACKGROUND_INTERACTION_COST_POLL_INITIAL_INTERVAL_SECONDS": "1",
|
||||
"BACKGROUND_INTERACTION_COST_POLL_MAX_INTERVAL_SECONDS": "1",
|
||||
"BACKGROUND_INTERACTION_COST_POLL_TIMEOUT_SECONDS": "120",
|
||||
}
|
||||
_SPEND_QUERY: Final = (
|
||||
"SELECT request_id, spend, call_type, status, model, prompt_tokens, completion_tokens "
|
||||
'FROM "LiteLLM_SpendLogs" WHERE request_id = %s'
|
||||
)
|
||||
_SETTLEMENT_QUERY: Final = (
|
||||
"SELECT interaction_id, claimed_by, outcome, claimed_at IS NOT NULL AS claimed, "
|
||||
'settled_at IS NOT NULL AS settled, create_context FROM "LiteLLM_BackgroundInteractionSettlement" '
|
||||
"WHERE interaction_id = %s"
|
||||
)
|
||||
_SETTLEMENT_TABLE_PRESENT_QUERY: Final = "SELECT to_regclass(%s) IS NOT NULL AS present"
|
||||
_SETTLEMENT_TABLE: Final = '"LiteLLM_BackgroundInteractionSettlement"'
|
||||
_SETTLEMENT_BY_CALL_QUERY: Final = (
|
||||
'SELECT interaction_id FROM "LiteLLM_BackgroundInteractionSettlement" WHERE create_context->>%s = %s'
|
||||
)
|
||||
_OUTAGE_RENAME: Final = (
|
||||
'ALTER TABLE IF EXISTS "LiteLLM_BackgroundInteractionSettlement" '
|
||||
'RENAME TO "LiteLLM_BackgroundInteractionSettlement_outage"'
|
||||
)
|
||||
_OUTAGE_RESTORE: Final = (
|
||||
'ALTER TABLE IF EXISTS "LiteLLM_BackgroundInteractionSettlement_outage" '
|
||||
'RENAME TO "LiteLLM_BackgroundInteractionSettlement"'
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Deployments:
|
||||
"""Config deployments every replica boots with, so no worker ever misses a model added at run time."""
|
||||
|
||||
in_progress: str
|
||||
completed_at_once: str
|
||||
failing_create: str
|
||||
custom_priced: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Rig:
|
||||
gateway: Gateway
|
||||
upstream: UpstreamSlot
|
||||
config: Path
|
||||
models: Deployments
|
||||
creator: Gateway
|
||||
settler: Gateway
|
||||
settler_pid: int
|
||||
directory: Path
|
||||
|
||||
def environment(self, **poll: str) -> dict[str, str]:
|
||||
return {"GEMINI_API_BASE": self.upstream.url, "GEMINI_API_KEY": _ENV_KEY, **poll}
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]:
|
||||
directory: Final = tmp_path_factory.mktemp("settlement")
|
||||
with gateway_from_environment() as gateway, owned_upstream(directory) as upstream:
|
||||
models: Final = _register_deployments(upstream.url)
|
||||
config: Final = _write_config(directory, upstream.url, models)
|
||||
environment: Final = {"GEMINI_API_BASE": upstream.url, "GEMINI_API_KEY": _ENV_KEY}
|
||||
with (
|
||||
owned_proxy(gateway, directory, {**environment, **_CREATOR_POLL}, config=config, workers=1) as creator,
|
||||
owned_proxy_process(
|
||||
gateway, directory, {**environment, **_SETTLER_POLL}, config=config, workers=2
|
||||
) as settler,
|
||||
):
|
||||
yield Rig(gateway, upstream, config, models, creator, settler.gateway, settler.process.pid, directory)
|
||||
|
||||
|
||||
def _register_deployments(upstream_url: str) -> Deployments:
|
||||
suffix: Final = uuid.uuid4().hex[:8]
|
||||
models: Final = Deployments(
|
||||
in_progress=f"settle-in-progress-{suffix}",
|
||||
completed_at_once=f"settle-completed-at-once-{suffix}",
|
||||
failing_create=f"settle-failing-create-{suffix}",
|
||||
custom_priced=f"settle-custom-priced-{suffix}",
|
||||
)
|
||||
_register_scenarios(upstream_url, models)
|
||||
return models
|
||||
|
||||
|
||||
def _register_scenarios(upstream_url: str, models: Deployments) -> None:
|
||||
scripted: Final = {
|
||||
models.in_progress: _interaction("in_progress", None),
|
||||
models.completed_at_once: _interaction("completed", _USAGE),
|
||||
models.failing_create: JsonResponse(
|
||||
content_type="application/json", body={"error": {"message": "boom"}}, status=500
|
||||
),
|
||||
models.custom_priced: _interaction("in_progress", None),
|
||||
}
|
||||
for name, response in scripted.items():
|
||||
register_scenario(
|
||||
name,
|
||||
RoutedResponse(content_type="application/x-routed", routes={"POST /v1beta/interactions": response}),
|
||||
control_url=upstream_url,
|
||||
)
|
||||
|
||||
|
||||
def _write_config(directory: Path, upstream_url: str, models: Deployments) -> Path:
|
||||
base: Final = JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
|
||||
custom_pricing: Final = {"input_cost_per_token": _CUSTOM_INPUT_RATE, "output_cost_per_token": _CUSTOM_OUTPUT_RATE}
|
||||
model_list: Final = [
|
||||
{
|
||||
"model_name": name,
|
||||
"litellm_params": {
|
||||
"model": _MODEL,
|
||||
"api_base": f"{upstream_url}/{name}",
|
||||
"api_key": _DEPLOYMENT_KEY,
|
||||
**(custom_pricing if name == models.custom_priced else {}),
|
||||
},
|
||||
}
|
||||
for name in (models.in_progress, models.completed_at_once, models.failing_create, models.custom_priced)
|
||||
]
|
||||
path: Final = directory / "settlement_config.yaml"
|
||||
path.write_text(yaml.safe_dump({**base, "model_list": model_list}))
|
||||
return path
|
||||
|
||||
|
||||
def _interaction(status: str, usage: dict[str, JsonValue] | None, http_status: int = 200) -> JsonResponse:
|
||||
return JsonResponse(
|
||||
content_type="application/json",
|
||||
body={
|
||||
"id": "$UNIQUE_ID",
|
||||
"object": "interaction",
|
||||
"model": "gemini-3.8-flash",
|
||||
"status": status,
|
||||
"steps": [],
|
||||
"usage": usage,
|
||||
},
|
||||
status=http_status,
|
||||
)
|
||||
|
||||
|
||||
def _completed() -> InteractionState:
|
||||
return InteractionState(status="completed", usage=_USAGE)
|
||||
|
||||
|
||||
def _create(
|
||||
replica: Gateway,
|
||||
model: str,
|
||||
key: str,
|
||||
*,
|
||||
path: str = "/v1beta/interactions",
|
||||
background: bool = True,
|
||||
text: str | None = None,
|
||||
) -> str:
|
||||
response: Final = replica.request(
|
||||
"POST",
|
||||
path,
|
||||
{"model": model, "input": text or f"settle {uuid.uuid4().hex}", "background": background},
|
||||
key=key,
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
return string_value(JSON_OBJECT.validate_json(response.content)["id"])
|
||||
|
||||
|
||||
def _state(rig: Rig, interaction_id: str, state: InteractionState) -> None:
|
||||
set_interaction_state(rig.upstream.url, interaction_id, state)
|
||||
|
||||
|
||||
def _delete(replica: Gateway, interaction_id: str, key: str, *, path: str = "/v1beta/interactions") -> httpx.Response:
|
||||
return replica.request("DELETE", f"{path}/{interaction_id}", key=key)
|
||||
|
||||
|
||||
def _delete_ok(replica: Gateway, interaction_id: str, key: str) -> None:
|
||||
deleted: Final = _delete(replica, interaction_id, key)
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
|
||||
|
||||
def _delete_concurrently(replica: Gateway, interaction_ids: Sequence[str], key: str) -> tuple[int, ...]:
|
||||
def status(interaction_id: str) -> int:
|
||||
return _delete(replica, interaction_id, key).status_code
|
||||
|
||||
with ThreadPoolExecutor(max_workers=8) as pool:
|
||||
return tuple(pool.map(status, interaction_ids))
|
||||
|
||||
|
||||
def _assert_unclaimed(interaction_id: str) -> None:
|
||||
row: Final = _settlement(interaction_id)
|
||||
assert row is not None and row["claimed"] is False and row["outcome"] is None, row
|
||||
|
||||
|
||||
def _spend_rows(request_id: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(_SPEND_QUERY, (request_id,))
|
||||
|
||||
|
||||
def _settlement(interaction_id: str) -> dict[str, JsonValue] | None:
|
||||
rows: Final = read_rows(_SETTLEMENT_QUERY, (interaction_id,))
|
||||
return rows[0] if rows else None
|
||||
|
||||
|
||||
def _settlement_table_present() -> bool:
|
||||
return read_rows(_SETTLEMENT_TABLE_PRESENT_QUERY, (_SETTLEMENT_TABLE,))[0]["present"] is True
|
||||
|
||||
|
||||
def _settlement_if_stored(interaction_id: str) -> dict[str, JsonValue] | None:
|
||||
return _settlement(interaction_id) if _settlement_table_present() else None
|
||||
|
||||
|
||||
def _settlements_by_call_if_stored(call_id: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(_SETTLEMENT_BY_CALL_QUERY, ("litellm_call_id", call_id)) if _settlement_table_present() else []
|
||||
|
||||
|
||||
def _await_spend_row(interaction_id: str, seconds: float = 30) -> dict[str, JsonValue]:
|
||||
return eventually(lambda: _spend_rows(interaction_id), lambda rows: len(rows) == 1, seconds=seconds)[0]
|
||||
|
||||
|
||||
def _await_outcome(interaction_id: str, outcome: str, seconds: float = 30) -> dict[str, JsonValue]:
|
||||
row: Final = eventually(
|
||||
lambda: _settlement(interaction_id),
|
||||
lambda value: value is not None and value["outcome"] == outcome,
|
||||
seconds=seconds,
|
||||
)
|
||||
assert row is not None
|
||||
return row
|
||||
|
||||
|
||||
def _model_info(replica: Gateway, model: str) -> Mapping[str, JsonValue]:
|
||||
entries: Final = replica.get("/model/info")["data"]
|
||||
assert isinstance(entries, list), entries
|
||||
return object_value(
|
||||
next(object_value(entry)["model_info"] for entry in entries if object_value(entry)["model_name"] == model)
|
||||
)
|
||||
|
||||
|
||||
def _rates(replica: Gateway, model: str) -> tuple[float, float]:
|
||||
info: Final = _model_info(replica, model)
|
||||
input_rate: Final = info["input_cost_per_token"]
|
||||
output_rate: Final = info["output_cost_per_token"]
|
||||
assert isinstance(input_rate, float) and isinstance(output_rate, float), info
|
||||
return input_rate, output_rate
|
||||
|
||||
|
||||
def _reservation_pin(replica: Gateway, model: str) -> float:
|
||||
"""What one background create reserves on the key before its usage is known: the output tokens the
|
||||
estimator assumes, at the deployment's output rate, with the prompt's few input tokens left as slack."""
|
||||
info: Final = _model_info(replica, model)
|
||||
max_output: Final = info["max_output_tokens"]
|
||||
output_rate: Final = info["output_cost_per_token"]
|
||||
assert isinstance(max_output, int) and isinstance(output_rate, float), info
|
||||
return min(max_output, DEFAULT_MAX_OUTPUT_TOKENS_FALLBACK) * output_rate
|
||||
|
||||
|
||||
def _assert_billed(row: Mapping[str, JsonValue], rates: tuple[float, float]) -> float:
|
||||
expected: Final = _INPUT_TOKENS * rates[0] + _OUTPUT_TOKENS * rates[1]
|
||||
spend: Final = row["spend"]
|
||||
assert isinstance(spend, float) and math.isclose(spend, expected, rel_tol=1e-9), (row, expected)
|
||||
assert row["call_type"] == "acreate_interaction", row
|
||||
assert row["status"] == "success", row
|
||||
assert row["prompt_tokens"] == _INPUT_TOKENS and row["completion_tokens"] == _OUTPUT_TOKENS, row
|
||||
return spend
|
||||
|
||||
|
||||
def _key_spend(replica: Gateway, key: str) -> float:
|
||||
spend: Final = object_value(replica.get("/key/info", {"key": key})["info"])["spend"]
|
||||
assert isinstance(spend, float | int), spend
|
||||
return float(spend)
|
||||
|
||||
|
||||
def _await_key_spend(replica: Gateway, key: str, expected: float) -> None:
|
||||
eventually(lambda: _key_spend(replica, key), lambda spend: math.isclose(spend, expected, rel_tol=1e-9), seconds=30)
|
||||
|
||||
|
||||
def _drain(rig: Rig) -> list[JsonValue]:
|
||||
observed: Final = httpx.get(f"{rig.upstream.url}/__observations", trust_env=False, timeout=15)
|
||||
observed.raise_for_status()
|
||||
requests: Final = JSON_OBJECT.validate_json(observed.content)["requests"]
|
||||
assert isinstance(requests, list), requests
|
||||
return requests
|
||||
|
||||
|
||||
def _calls(rig: Rig, interaction_id: str) -> tuple[tuple[str, str], ...]:
|
||||
suffix: Final = f"/v1beta/interactions/{interaction_id}"
|
||||
return tuple(
|
||||
(string_value(object_value(entry)["method"]), string_value(object_value(entry)["api_key"]))
|
||||
for entry in _drain(rig)
|
||||
if string_value(object_value(entry)["path"]).endswith(suffix)
|
||||
)
|
||||
|
||||
|
||||
def _claimer_pid(row: Mapping[str, JsonValue]) -> int:
|
||||
claimed_by: Final = string_value(row["claimed_by"])
|
||||
host, _, pid = claimed_by.rpartition(":")
|
||||
assert host == socket.gethostname(), claimed_by
|
||||
return int(pid)
|
||||
|
||||
|
||||
def _worker_pids(root_pid: int) -> frozenset[int]:
|
||||
return frozenset(
|
||||
process.pid for process in group_members(root_pid) if process.pid != root_pid and _is_spawned_worker(process)
|
||||
)
|
||||
|
||||
|
||||
def _is_spawned_worker(process: psutil.Process) -> bool:
|
||||
try:
|
||||
return process.name().lower().startswith("python") and "resource_tracker" not in " ".join(process.cmdline())
|
||||
except psutil.Error:
|
||||
return False
|
||||
|
||||
|
||||
def _readiness(replica: Gateway) -> int:
|
||||
try:
|
||||
return replica.request("GET", "/health/readiness").status_code
|
||||
except httpx.TransportError:
|
||||
return 0
|
||||
|
||||
|
||||
def test_creator_poll_bills_a_completed_background_interaction_once(rig: Rig) -> None:
|
||||
with rig.settler.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(rig.settler, model, key)
|
||||
_state(rig, created, _completed())
|
||||
spend: Final = _assert_billed(_await_spend_row(created), _rates(rig.settler, model))
|
||||
_await_key_spend(rig.settler, key, spend)
|
||||
assert len(_spend_rows(created)) == 1
|
||||
|
||||
|
||||
def test_creator_poll_records_its_settlement_durably(rig: Rig) -> None:
|
||||
with rig.settler.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
created: Final = _create(rig.settler, model, scenario.key())
|
||||
_state(rig, created, _completed())
|
||||
_await_spend_row(created)
|
||||
row: Final = _await_outcome(created, "billed")
|
||||
assert row["claimed"] is True and row["settled"] is True, row
|
||||
assert row["create_context"] == {}, row
|
||||
_claimer_pid(row)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path", ["/v1beta/interactions", "/interactions"])
|
||||
def test_delete_on_another_replica_bills_the_creators_interaction_once(rig: Rig, path: str) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(rig.creator, model, key, path=path)
|
||||
_state(rig, created, _completed())
|
||||
_drain(rig)
|
||||
deleted: Final = _delete(rig.settler, created, key, path=path)
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
spend: Final = _assert_billed(_await_spend_row(created), _rates(rig.creator, model))
|
||||
row: Final = _await_outcome(created, "billed")
|
||||
assert _claimer_pid(row) in _worker_pids(rig.settler_pid), row
|
||||
assert _calls(rig, created) == (("GET", _ENV_KEY), ("DELETE", _ENV_KEY))
|
||||
_await_key_spend(rig.creator, key, spend)
|
||||
assert len(_spend_rows(created)) == 1
|
||||
|
||||
|
||||
def test_delete_of_a_failed_interaction_releases_without_a_spend_row(rig: Rig) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(rig.creator, model, key)
|
||||
_state(rig, created, InteractionState(status="failed", usage=None))
|
||||
deleted: Final = _delete(rig.settler, created, key)
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
_await_outcome(created, "released")
|
||||
assert _spend_rows(created) == []
|
||||
assert _key_spend(rig.creator, key) == 0
|
||||
|
||||
|
||||
def test_delete_of_a_requires_action_interaction_bills_its_usage(rig: Rig) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(rig.creator, model, key)
|
||||
_state(rig, created, InteractionState(status="requires_action", usage=_USAGE))
|
||||
deleted: Final = _delete(rig.settler, created, key)
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
_assert_billed(_await_spend_row(created), _rates(rig.creator, model))
|
||||
_await_outcome(created, "billed")
|
||||
|
||||
|
||||
def test_a_replica_booting_later_resumes_and_bills_unclaimed_interactions(rig: Rig) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = tuple(_create(rig.creator, model, key) for _ in range(3))
|
||||
for item in created:
|
||||
_state(rig, item, _completed())
|
||||
rates: Final = _rates(rig.creator, model)
|
||||
with owned_proxy_process(
|
||||
rig.gateway, rig.directory, rig.environment(**_RESUMER_POLL), config=rig.config, workers=2
|
||||
) as resumer:
|
||||
pids: Final = _worker_pids(resumer.process.pid)
|
||||
assert len(pids) == 2, pids
|
||||
for item in created:
|
||||
_assert_billed(_await_spend_row(item, seconds=90), rates)
|
||||
assert _claimer_pid(_await_outcome(item, "billed")) in pids
|
||||
for item in created:
|
||||
assert len(_spend_rows(item)) == 1
|
||||
|
||||
|
||||
def test_deletes_on_the_creating_proxy_bill_each_interaction_once(rig: Rig) -> None:
|
||||
with rig.settler.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = tuple(_create(rig.settler, model, key) for _ in range(8))
|
||||
for item in created:
|
||||
_state(rig, item, _completed())
|
||||
assert _delete_concurrently(rig.settler, created, key) == (200,) * 8
|
||||
rates: Final = _rates(rig.settler, model)
|
||||
for item in created:
|
||||
_assert_billed(_await_spend_row(item), rates)
|
||||
_await_outcome(item, "billed")
|
||||
_await_key_spend(rig.settler, key, 8 * (_INPUT_TOKENS * rates[0] + _OUTPUT_TOKENS * rates[1]))
|
||||
for item in created:
|
||||
assert len(_spend_rows(item)) == 1
|
||||
|
||||
|
||||
def test_custom_deployment_pricing_bills_at_the_deployment_rate_on_another_replica(rig: Rig) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.custom_priced
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(rig.creator, model, key)
|
||||
_state(rig, created, _completed())
|
||||
deleted: Final = _delete(rig.settler, created, key)
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
_assert_billed(_await_spend_row(created), (_CUSTOM_INPUT_RATE, _CUSTOM_OUTPUT_RATE))
|
||||
_await_outcome(created, "billed")
|
||||
|
||||
|
||||
def test_cancel_then_delete_on_another_replica_releases_without_a_spend_row(rig: Rig) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(rig.creator, model, key)
|
||||
_state(rig, created, InteractionState(status="in_progress"))
|
||||
cancelled: Final = rig.settler.request("POST", f"/v1beta/interactions/{created}/cancel", {}, key=key)
|
||||
assert cancelled.status_code == 200, cancelled.text
|
||||
before_delete: Final = _settlement(created)
|
||||
assert before_delete is not None and before_delete["claimed"] is False, before_delete
|
||||
deleted: Final = _delete(rig.settler, created, key)
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
_await_outcome(created, "released")
|
||||
assert _spend_rows(created) == []
|
||||
|
||||
|
||||
def test_delete_fails_closed_when_the_settling_replica_cannot_fetch(rig: Rig) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(rig.creator, model, key)
|
||||
_state(rig, created, InteractionState(status="completed", usage=_USAGE, get_status=500))
|
||||
_drain(rig)
|
||||
refused: Final = _delete(rig.settler, created, key)
|
||||
assert refused.status_code >= 500, refused.text
|
||||
assert "Scripted interaction fetch failure" in refused.text, refused.text
|
||||
assert _calls(rig, created) == (("GET", _ENV_KEY),)
|
||||
_assert_unclaimed(created)
|
||||
assert _spend_rows(created) == []
|
||||
_state(rig, created, _completed())
|
||||
deleted: Final = _delete(rig.settler, created, key)
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
_assert_billed(_await_spend_row(created), _rates(rig.creator, model))
|
||||
_await_outcome(created, "billed")
|
||||
|
||||
|
||||
def test_delete_of_an_interaction_the_vendor_purged_sends_no_delete_and_keeps_the_row(rig: Rig) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(rig.creator, model, key)
|
||||
clear_interaction_state(rig.upstream.url, created)
|
||||
_drain(rig)
|
||||
deleted: Final = _delete(rig.settler, created, key)
|
||||
assert deleted.status_code == 404, deleted.text
|
||||
assert _calls(rig, created) == (("GET", _ENV_KEY),)
|
||||
row: Final = _settlement(created)
|
||||
assert row is not None and row["claimed"] is False, row
|
||||
assert _spend_rows(created) == []
|
||||
|
||||
|
||||
def test_reading_an_interaction_never_bills_it(rig: Rig) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(rig.creator, model, key)
|
||||
_state(rig, created, InteractionState(status="in_progress"))
|
||||
read_ids: Final = tuple(str(uuid.uuid4()) for _ in range(2))
|
||||
first: Final = rig.settler.request(
|
||||
"GET", f"/v1beta/interactions/{created}", key=key, headers={"x-litellm-call-id": read_ids[0]}
|
||||
)
|
||||
assert first.status_code == 200 and JSON_OBJECT.validate_json(first.content)["status"] == "in_progress", (
|
||||
first.text
|
||||
)
|
||||
_state(rig, created, _completed())
|
||||
second: Final = rig.settler.request(
|
||||
"GET", f"/v1beta/interactions/{created}", key=key, headers={"x-litellm-call-id": read_ids[1]}
|
||||
)
|
||||
assert second.status_code == 200 and JSON_OBJECT.validate_json(second.content)["usage"] == _USAGE, second.text
|
||||
assert _key_spend(rig.creator, key) == 0
|
||||
assert _spend_rows(created) == []
|
||||
deleted: Final = _delete(rig.settler, created, key)
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
spend: Final = _assert_billed(_await_spend_row(created), _rates(rig.creator, model))
|
||||
_await_key_spend(rig.creator, key, spend)
|
||||
for read_id in read_ids:
|
||||
assert all(row["spend"] == 0 for row in _spend_rows(read_id)), _spend_rows(read_id)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"interaction_id",
|
||||
[f"missing-{uuid.uuid4().hex}", "x" * 5000, "a.b:c", "%2F..%2Fup"],
|
||||
ids=["unknown", "five-kilobytes", "punctuation", "encoded-traversal"],
|
||||
)
|
||||
def test_delete_of_an_odd_or_unknown_id_is_refused_and_the_proxy_keeps_serving(rig: Rig, interaction_id: str) -> None:
|
||||
with rig.settler.scenario() as scenario:
|
||||
key: Final = scenario.key()
|
||||
deleted: Final = rig.settler.request("DELETE", f"/v1beta/interactions/{interaction_id}", key=key)
|
||||
assert 400 <= deleted.status_code < 500, deleted.text
|
||||
assert _readiness(rig.settler) == 200
|
||||
assert _key_spend(rig.settler, key) == 0
|
||||
|
||||
|
||||
def test_a_missing_settlement_table_leaves_in_process_billing_intact(rig: Rig) -> None:
|
||||
write_rows(_OUTAGE_RENAME, ())
|
||||
try:
|
||||
with rig.settler.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(rig.settler, model, key)
|
||||
_state(rig, created, _completed())
|
||||
spend: Final = _assert_billed(_await_spend_row(created), _rates(rig.settler, model))
|
||||
_await_key_spend(rig.settler, key, spend)
|
||||
deleted: Final = _delete(rig.creator, created, key)
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
assert len(_spend_rows(created)) == 1
|
||||
finally:
|
||||
write_rows(_OUTAGE_RESTORE, ())
|
||||
|
||||
|
||||
def test_a_failed_create_registers_nothing(rig: Rig) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.failing_create
|
||||
key: Final = scenario.key()
|
||||
call_id: Final = str(uuid.uuid4())
|
||||
response: Final = rig.creator.request(
|
||||
"POST",
|
||||
"/v1beta/interactions",
|
||||
{"model": model, "input": f"settle {uuid.uuid4().hex}", "background": True},
|
||||
key=key,
|
||||
headers={"x-litellm-call-id": call_id},
|
||||
)
|
||||
assert response.status_code >= 500, response.text
|
||||
assert _key_spend(rig.creator, key) == 0
|
||||
assert all(row["spend"] == 0 for row in _spend_rows(call_id)), _spend_rows(call_id)
|
||||
assert _settlements_by_call_if_stored(call_id) == []
|
||||
|
||||
|
||||
def test_polling_disabled_replica_registers_nothing_and_never_bills(rig: Rig) -> None:
|
||||
disabled: Final = rig.environment(BACKGROUND_INTERACTION_COST_POLLING_ENABLED="false")
|
||||
with (
|
||||
owned_proxy(rig.gateway, rig.directory, disabled, config=rig.config, workers=1) as quiet,
|
||||
quiet.scenario() as scenario,
|
||||
):
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(quiet, model, key)
|
||||
_state(rig, created, _completed())
|
||||
deleted: Final = _delete(rig.settler, created, key)
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
assert _settlement_if_stored(created) is None
|
||||
assert _key_spend(quiet, key) == 0
|
||||
assert _spend_rows(created) == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize("background", [False, True], ids=["synchronous", "background"])
|
||||
def test_a_create_that_completes_at_once_is_billed_by_the_create_alone(rig: Rig, background: bool) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.completed_at_once
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(rig.creator, model, key, background=background)
|
||||
spend: Final = _assert_billed(_await_spend_row(created), _rates(rig.creator, model))
|
||||
_await_key_spend(rig.creator, key, spend)
|
||||
assert _settlement_if_stored(created) is None
|
||||
_state(rig, created, _completed())
|
||||
deleted: Final = _delete(rig.settler, created, key)
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
_await_key_spend(rig.creator, key, spend)
|
||||
assert len(_spend_rows(created)) == 1
|
||||
|
||||
|
||||
def test_identical_creates_settle_as_separate_interactions(rig: Rig) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
text: Final = f"settle {uuid.uuid4().hex}"
|
||||
created: Final = tuple(_create(rig.creator, model, key, text=text) for _ in range(3))
|
||||
assert len({item for item in created}) == 3, created
|
||||
for item in created:
|
||||
_state(rig, item, _completed())
|
||||
_delete_ok(rig.settler, item, key)
|
||||
rates: Final = _rates(rig.creator, model)
|
||||
for item in created:
|
||||
_assert_billed(_await_spend_row(item), rates)
|
||||
_await_outcome(item, "billed")
|
||||
_await_key_spend(rig.creator, key, 3 * (_INPUT_TOKENS * rates[0] + _OUTPUT_TOKENS * rates[1]))
|
||||
|
||||
|
||||
def test_settlement_on_another_replica_releases_the_creators_budget_reservation(rig: Rig) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key(max_budget=1.5 * _reservation_pin(rig.creator, model))
|
||||
first: Final = _create(rig.creator, model, key)
|
||||
pinned: Final = rig.creator.request(
|
||||
"POST", "/v1beta/interactions", {"model": model, "input": "settle pinned", "background": True}, key=key
|
||||
)
|
||||
assert pinned.status_code == 400, pinned.text
|
||||
_state(rig, first, _completed())
|
||||
deleted: Final = _delete(rig.settler, first, key)
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
spend: Final = _assert_billed(_await_spend_row(first), _rates(rig.creator, model))
|
||||
_await_key_spend(rig.creator, key, spend)
|
||||
released: Final = eventually(
|
||||
lambda: (
|
||||
rig.creator.request(
|
||||
"POST",
|
||||
"/v1beta/interactions",
|
||||
{"model": model, "input": "settle released", "background": True},
|
||||
key=key,
|
||||
).status_code
|
||||
),
|
||||
lambda status: status == 200,
|
||||
seconds=20,
|
||||
return_last_on_timeout=True,
|
||||
)
|
||||
assert released == 200
|
||||
|
||||
|
||||
def test_a_poll_that_never_sees_a_terminal_status_records_unsettled_and_releases(rig: Rig) -> None:
|
||||
with rig.settler.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(rig.settler, model, key)
|
||||
_state(rig, created, InteractionState(status="in_progress"))
|
||||
row: Final = _await_outcome(created, "unsettled", seconds=40)
|
||||
assert row["create_context"] == {}, row
|
||||
assert _spend_rows(created) == []
|
||||
assert _key_spend(rig.settler, key) == 0
|
||||
deleted: Final = _delete(rig.creator, created, key)
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
assert _spend_rows(created) == []
|
||||
|
||||
|
||||
def test_an_upstream_outage_fails_deletes_closed_and_every_interaction_bills_once_after_recovery(rig: Rig) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = tuple(_create(rig.creator, model, key) for _ in range(16))
|
||||
rates: Final = _rates(rig.creator, model)
|
||||
rig.upstream.stop()
|
||||
try:
|
||||
refused: Final = _delete_concurrently(rig.settler, created, key)
|
||||
assert all(status >= 500 for status in refused), refused
|
||||
for item in created:
|
||||
_assert_unclaimed(item)
|
||||
assert _readiness(rig.creator) == 200 and _readiness(rig.settler) == 200
|
||||
finally:
|
||||
rig.upstream.start()
|
||||
_register_scenarios(rig.upstream.url, rig.models)
|
||||
for item in created:
|
||||
_state(rig, item, _completed())
|
||||
assert _delete_concurrently(rig.settler, created, key) == (200,) * 16
|
||||
for item in created:
|
||||
_assert_billed(_await_spend_row(item), rates)
|
||||
_await_outcome(item, "billed")
|
||||
_await_key_spend(rig.creator, key, 16 * (_INPUT_TOKENS * rates[0] + _OUTPUT_TOKENS * rates[1]))
|
||||
for item in created:
|
||||
assert len(_spend_rows(item)) == 1
|
||||
|
||||
|
||||
def test_killed_workers_leave_their_polls_to_the_respawned_workers(rig: Rig) -> None:
|
||||
with (
|
||||
owned_proxy_process(
|
||||
rig.gateway, rig.directory, rig.environment(**_RESUMER_POLL), config=rig.config, workers=2
|
||||
) as resumer,
|
||||
resumer.gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = tuple(_create(resumer.gateway, model, key) for _ in range(16))
|
||||
for item in created:
|
||||
_state(rig, item, InteractionState(status="in_progress"))
|
||||
rates: Final = _rates(resumer.gateway, model)
|
||||
killed: Final = _worker_pids(resumer.process.pid)
|
||||
assert len(killed) == 2, killed
|
||||
victims: Final = tuple(psutil.Process(pid) for pid in killed)
|
||||
for victim in victims:
|
||||
victim.kill()
|
||||
psutil.wait_procs(victims, timeout=15)
|
||||
for item in created:
|
||||
_state(rig, item, _completed())
|
||||
for item in created:
|
||||
_assert_billed(_await_spend_row(item, seconds=150), rates)
|
||||
assert _claimer_pid(_await_outcome(item, "billed")) not in killed
|
||||
assert eventually(lambda: _readiness(resumer.gateway), lambda status: status == 200, seconds=60) == 200
|
||||
for item in created:
|
||||
assert len(_spend_rows(item)) == 1
|
||||
|
||||
|
||||
def test_concurrent_deletes_on_a_slow_upstream_settle_exactly_once(rig: Rig) -> None:
|
||||
with rig.creator.scenario() as scenario:
|
||||
model: Final = rig.models.in_progress
|
||||
key: Final = scenario.key()
|
||||
created: Final = _create(rig.creator, model, key)
|
||||
_state(rig, created, InteractionState(status="completed", usage=_USAGE, delay_seconds=1.5))
|
||||
statuses: Final = _delete_concurrently(rig.settler, (created, created), key)
|
||||
assert sorted(statuses) == [200, 404], statuses
|
||||
spend: Final = _assert_billed(_await_spend_row(created), _rates(rig.creator, model))
|
||||
_await_outcome(created, "billed")
|
||||
_await_key_spend(rig.creator, key, spend)
|
||||
assert len(_spend_rows(created)) == 1
|
||||
Loading…
Add table
Reference in a new issue