mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(interactions): durable cross-pod settlement for background interaction billing (#41955)
* feat(interactions): durable cross-pod settlement for background interaction billing Background interaction billing lived only in the creating replica's memory, so a DELETE routed to another replica, or a restart of the creating one, never billed the completed provider work and the budget reservation was refunded at the poll timeout. The create now registers the billing context in a settlement store before returning, the proxy installs a Prisma-backed store at boot (LiteLLM_BackgroundInteractionSettlement, schema-only migration), any replica claims the row once through a conditional update before billing or releasing, startup resumes every unclaimed row with its remaining timeout, and a give-up records an unsettled outcome instead of silently reconciling to zero. The SDK keeps an in-memory store and behaves as before. * fix(interactions): survive a settlement install failure at boot and stop carrying request headers * fix(interactions): drop the stored request context once a settlement row is settled * fix(interactions): bill the completed response a poll already saw when its claim only answers at the deadline * fix(interactions): carry a missing model through the settlement context for agent-only background creates An interaction created with an agent and no model reaches the poll with no model name, exactly as on main. The settlement context now stores that None instead of rejecting the create, which answered the client with a 500 after the provider had already accepted it. * fix(interactions): leave an unfetchable background interaction to its creating poll when a delete lands elsewhere The remote pre-delete path fetches with only the delete's credentials, so a fetch it cannot make says nothing about the interaction. It used to claim the settlement row and release the reservation anyway, which stopped the creating replica's poll and lost the bill when the delete then failed the same way. It now returns without claiming; the in-process path keeps releasing on an unfetchable state, since its context carries the create's own credentials. * fix(interactions): fail a cross-replica delete when its pre-delete fetch fails so the creating poll keeps the bill * fix(interactions): keep the stored settlement gate when registration raises after landing, and fail resumed-poll deletes closed A registration that raised after its row committed moved the poll to a private in-memory gate, so the creating worker billed while the stored row stayed unclaimed for another replica's delete or the next boot to bill again. The row is now read back once and, when it landed, the poll claims through it like every other settler. A worker that resumed the poll after a restart is not the creator, so its delete on a failed pre-delete fetch now fails with the fetch's error instead of releasing and deleting. After a fleet restart every worker holds resumed polls, which left the fail-closed path applying nowhere. * fix(interactions): settle an unverified registration through the durable claim A create whose settlement-store write raised no longer bills through a private in-memory gate that a later boot's resume cannot see. The claim asks the durable store first and falls back to the local gate only when the store answers that no row exists, and a missing settlement table reads as no rows so a replica without the migration still settles in process. * test(proxy): keep the settlement test where the proxy-infra shard collects it The merge of main moved test_background_interaction_settlement.py under tests/unit/proxy/spend_tracking, but the proxy-db shards claim tests/unit/proxy files one by one in .circleci/scripts/unit_selection.sh, so no CI shard ran it and codecov/patch dropped. tests/test_litellm/proxy/spend_tracking is collected whole by the proxy-infra shard, which is where the test ran before the merge. * fix(interactions): raise on a non-2xx Gemini interaction fetch AsyncHTTPHandler.get never raises for status and the Gemini GET transform only raised when the body was not JSON, so a 500 or 404 carrying Gemini's JSON error body parsed as an interaction with no status. A delete on a replica other than the creator then claimed the settlement as released and forwarded the delete instead of failing closed, and the bill was lost. The transform now raises GeminiError with the vendor's status, as the delete transform already does; the in-process poll already retries a fetch that raises * 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 * test(integration): refuse a repeated delete in the scripted upstream and pin the settlement budget below one estimate * chore(ui): regenerate dashboard API types after merging main * test(integration): accept the 422 budget refusal and a respawned worker's resume The budget cell pinned a 400 that the proxy stopped answering when budget refusals moved to 422, so it now asserts the status and the budget_exceeded error type the sibling budget tests pin. The later-booting replica cell accepts a claimer that is any worker started after the creates, since uvicorn's supervisor can respawn the creator's worker under load and the respawned worker's boot resume claims the rows by design; the single spend row check is unchanged * test(integration): delete the pinned key's interaction with a second key A key whose budget is filled by its own reservation is refused on every route, the DELETE included, so the cell now asserts that 422 and sends the delete with a second key, which is what the reservation release on another replica needs in order to be observable at all --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
ad566c90dd
commit
8596fe954d
16 changed files with 2629 additions and 106 deletions
|
|
@ -0,0 +1,16 @@
|
|||
-- CreateTable
|
||||
CREATE TABLE IF NOT EXISTS "LiteLLM_BackgroundInteractionSettlement" (
|
||||
"interaction_id" TEXT NOT NULL,
|
||||
"custom_llm_provider" TEXT NOT NULL,
|
||||
"create_context" JSONB NOT NULL,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
"claimed_at" TIMESTAMP(3),
|
||||
"claimed_by" TEXT,
|
||||
"settled_at" TIMESTAMP(3),
|
||||
"outcome" TEXT,
|
||||
|
||||
CONSTRAINT "LiteLLM_BackgroundInteractionSettlement_pkey" PRIMARY KEY ("interaction_id")
|
||||
);
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX IF NOT EXISTS "idx_background_interaction_settlement_claimed_at" ON "LiteLLM_BackgroundInteractionSettlement"("claimed_at");
|
||||
|
|
@ -1916,6 +1916,22 @@ model LiteLLM_WorkflowMessage {
|
|||
@@index([run_id])
|
||||
}
|
||||
|
||||
// Pending billing settlements for background interactions, keyed by the
|
||||
// interaction id so any replica can settle one that another replica created.
|
||||
// `claimed_at` is the exactly-once gate: the first conditional update wins.
|
||||
model LiteLLM_BackgroundInteractionSettlement {
|
||||
interaction_id String @id
|
||||
custom_llm_provider String
|
||||
create_context Json
|
||||
created_at DateTime @default(now())
|
||||
claimed_at DateTime?
|
||||
claimed_by String?
|
||||
settled_at DateTime?
|
||||
outcome String?
|
||||
|
||||
@@index([claimed_at], map: "idx_background_interaction_settlement_claimed_at")
|
||||
}
|
||||
|
||||
model LiteLLM_Lens {
|
||||
id String @id
|
||||
version Int @default(0)
|
||||
|
|
|
|||
|
|
@ -23,16 +23,30 @@ caller retrieve the completed output themselves and then delete it before the
|
|||
poll task settles, leaving the work unbilled and the budget reservation
|
||||
refunded at the poll timeout. ``adelete`` therefore settles any pending poll
|
||||
for the interaction before dispatching the delete: it fetches the current
|
||||
state with the create's credentials, bills it if it is terminal with usage,
|
||||
and releases the reservation otherwise. A settlement gate on the create's
|
||||
logging object makes the poll task and the delete path mutually exclusive, so
|
||||
the interaction is billed exactly once no matter who settles first.
|
||||
state, bills it if it is terminal with usage, and releases the reservation
|
||||
otherwise.
|
||||
|
||||
The poll task lives in the process that served the create, so a delete
|
||||
served by another replica, or by the same replica after a restart, finds no
|
||||
task to settle. A ``BackgroundSettlementStore`` makes the pending settlement
|
||||
durable across processes: the create registers the request context that
|
||||
billing needs (never provider credentials), the settlement is claimed
|
||||
exactly once through the store, and a delete on any replica rebuilds the
|
||||
billing context from the store when the poll task is not local. Rows left
|
||||
unclaimed by a process that died are resumed at startup. The default store is
|
||||
in-memory, which keeps the SDK and single-process behavior unchanged; the
|
||||
proxy installs a database-backed one.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Awaitable, Callable, Iterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Final, TypeAlias
|
||||
from collections.abc import Awaitable, Callable, Iterable, Iterator, Mapping, Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError
|
||||
from pydantic_core import PydanticSerializationError, to_jsonable_python
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import (
|
||||
|
|
@ -43,6 +57,7 @@ from litellm.constants import (
|
|||
)
|
||||
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
|
||||
from litellm.types.interactions import InteractionsAPIResponse
|
||||
from litellm.types.utils import CustomPricingLiteLLMParams
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
|
@ -55,6 +70,98 @@ _POLLABLE_STATUSES: Final = frozenset({"in_progress", "queued"})
|
|||
|
||||
_STATUSES_THAT_PRODUCED_OUTPUT: Final = frozenset({"completed", "requires_action"})
|
||||
|
||||
SettlementOutcome: TypeAlias = Literal["billed", "released", "unsettled"]
|
||||
|
||||
|
||||
class BackgroundInteractionCreateContext(BaseModel):
|
||||
"""
|
||||
The part of a create's logging state that billing its settled result needs,
|
||||
in a shape any replica can store and rebuild a logging object from. Provider
|
||||
credentials are deliberately absent: the replica that settles fetches the
|
||||
interaction with its own, exactly as it would serve the delete itself.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
model: str | None
|
||||
call_type: str
|
||||
litellm_call_id: str
|
||||
function_id: str
|
||||
litellm_trace_id: str
|
||||
start_time: datetime
|
||||
custom_llm_provider: str
|
||||
metadata: Mapping[str, JsonValue]
|
||||
custom_pricing: Mapping[str, JsonValue]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PendingBackgroundInteraction:
|
||||
interaction_id: str
|
||||
custom_llm_provider: str
|
||||
create_context: BackgroundInteractionCreateContext
|
||||
created_at: datetime
|
||||
|
||||
|
||||
class BackgroundSettlementStore(Protocol):
|
||||
async def register(self, pending: PendingBackgroundInteraction) -> None: ...
|
||||
|
||||
async def pending(self, interaction_id: str) -> PendingBackgroundInteraction | None: ...
|
||||
|
||||
async def is_claimed(self, interaction_id: str) -> bool: ...
|
||||
|
||||
async def claim(self, interaction_id: str) -> bool: ...
|
||||
|
||||
async def record_outcome(self, interaction_id: str, outcome: SettlementOutcome) -> None: ...
|
||||
|
||||
async def unclaimed(self) -> Sequence[PendingBackgroundInteraction]: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class InMemoryBackgroundSettlementStore:
|
||||
"""
|
||||
Per-process store: a registered interaction maps to its pending row until
|
||||
it is claimed, after which it maps to ``None``. Claiming an interaction the
|
||||
store never saw succeeds once, which is what a poll built without a
|
||||
registration relies on.
|
||||
"""
|
||||
|
||||
_rows: dict[str, PendingBackgroundInteraction | None] = field( # mutable-ok: the registry every settler shares
|
||||
default_factory=dict
|
||||
)
|
||||
|
||||
async def register(self, pending: PendingBackgroundInteraction) -> None:
|
||||
self._rows[pending.interaction_id] = pending
|
||||
|
||||
async def pending(self, interaction_id: str) -> PendingBackgroundInteraction | None:
|
||||
return self._rows.get(interaction_id)
|
||||
|
||||
async def is_claimed(self, interaction_id: str) -> bool:
|
||||
return interaction_id in self._rows and self._rows[interaction_id] is None
|
||||
|
||||
async def claim(self, interaction_id: str) -> bool:
|
||||
if await self.is_claimed(interaction_id):
|
||||
return False
|
||||
self._rows[interaction_id] = None
|
||||
return True
|
||||
|
||||
async def record_outcome(self, interaction_id: str, outcome: SettlementOutcome) -> None:
|
||||
return None
|
||||
|
||||
async def unclaimed(self) -> Sequence[PendingBackgroundInteraction]:
|
||||
return tuple(row for row in self._rows.values() if row is not None)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _StoreSlot:
|
||||
store: BackgroundSettlementStore
|
||||
|
||||
|
||||
_STORE: Final = _StoreSlot(store=InMemoryBackgroundSettlementStore())
|
||||
|
||||
|
||||
def configure_background_settlement_store(store: BackgroundSettlementStore) -> None:
|
||||
_STORE.store = store
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BackgroundInteractionPollContext:
|
||||
|
|
@ -66,12 +173,14 @@ class BackgroundInteractionPollContext:
|
|||
initial_interval_seconds: float = BACKGROUND_INTERACTION_COST_POLL_INITIAL_INTERVAL_SECONDS
|
||||
max_interval_seconds: float = BACKGROUND_INTERACTION_COST_POLL_MAX_INTERVAL_SECONDS
|
||||
timeout_seconds: float = BACKGROUND_INTERACTION_COST_POLL_TIMEOUT_SECONDS
|
||||
store: BackgroundSettlementStore = field(default_factory=InMemoryBackgroundSettlementStore)
|
||||
resumed: bool = False
|
||||
|
||||
|
||||
FetchInteraction: TypeAlias = Callable[[BackgroundInteractionPollContext], Awaitable[InteractionsAPIResponse]]
|
||||
|
||||
|
||||
async def _fetch_interaction(context: BackgroundInteractionPollContext) -> InteractionsAPIResponse:
|
||||
async def fetch_background_interaction(context: BackgroundInteractionPollContext) -> InteractionsAPIResponse:
|
||||
from litellm.interactions import aget
|
||||
|
||||
return await aget(
|
||||
|
|
@ -84,46 +193,159 @@ async def _fetch_interaction(context: BackgroundInteractionPollContext) -> Inter
|
|||
|
||||
|
||||
def _poll_intervals(initial: float, maximum: float, timeout: float) -> Iterator[float]:
|
||||
elapsed = 0.0
|
||||
interval = initial
|
||||
elapsed = 0.0 # rebind-ok: the schedule accumulates the time it has already yielded
|
||||
interval = initial # rebind-ok: the schedule doubles the interval up to the cap
|
||||
while interval > 0 and elapsed + interval <= timeout:
|
||||
yield interval
|
||||
elapsed += interval
|
||||
interval = min(interval * 2, maximum)
|
||||
|
||||
|
||||
_SETTLED_KEY = "background_interaction_settled"
|
||||
_CUSTOM_PRICING_KEYS: Final = frozenset(CustomPricingLiteLLMParams.model_fields.keys())
|
||||
|
||||
_CARRIED_METADATA_KEYS: Final = frozenset(
|
||||
{
|
||||
"model_info",
|
||||
"model_group",
|
||||
"deployment",
|
||||
"tags",
|
||||
"spend_logs_metadata",
|
||||
"requester_metadata",
|
||||
"requester_ip_address",
|
||||
"user_agent",
|
||||
"agent_id",
|
||||
"session_id",
|
||||
"endpoint",
|
||||
"team_alias",
|
||||
"team_id",
|
||||
"applied_guardrails",
|
||||
"prompt_management_metadata",
|
||||
}
|
||||
)
|
||||
|
||||
_CARRIED_METADATA_PREFIX: Final = "user_api_"
|
||||
|
||||
_UNCARRIED_METADATA_KEY: Final = "user_api_key_auth"
|
||||
|
||||
_JSON_VALUE: Final = TypeAdapter(JsonValue)
|
||||
_STRING: Final = TypeAdapter(str)
|
||||
_OBJECT_MAPPING: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
def _is_settled(logging_obj: "LiteLLMLoggingObj") -> bool:
|
||||
return logging_obj.model_call_details.get(_SETTLED_KEY) is True
|
||||
|
||||
|
||||
def _claim_settlement(logging_obj: "LiteLLMLoggingObj") -> bool:
|
||||
"""
|
||||
Exactly-once gate between the poll task and the delete-time settlement:
|
||||
both run on the same event loop and neither awaits between reading and
|
||||
setting the flag, so whichever claims first owns billing or release.
|
||||
"""
|
||||
if _is_settled(logging_obj):
|
||||
def _carries(key: str) -> bool:
|
||||
if key == _UNCARRIED_METADATA_KEY:
|
||||
return False
|
||||
logging_obj.model_call_details[_SETTLED_KEY] = True # rebind-ok: both settlers must see the same settlement flag
|
||||
return True
|
||||
return key in _CARRIED_METADATA_KEYS or key.startswith(_CARRIED_METADATA_PREFIX)
|
||||
|
||||
|
||||
def _json_value(value: object) -> tuple[JsonValue, ...]:
|
||||
try:
|
||||
return (_JSON_VALUE.validate_python(to_jsonable_python(value)),)
|
||||
except (PydanticSerializationError, ValidationError):
|
||||
verbose_logger.debug("Dropping a background interaction metadata value that has no JSON form: %r", type(value))
|
||||
return ()
|
||||
|
||||
|
||||
def _json_values(items: Iterable[tuple[str, object]]) -> Mapping[str, JsonValue]:
|
||||
parsed: Final = ((key, _json_value(value)) for key, value in items)
|
||||
return MappingProxyType({key: values[0] for key, values in parsed if values})
|
||||
|
||||
|
||||
def _as_datetime(start_time: datetime | float) -> datetime:
|
||||
return start_time if isinstance(start_time, datetime) else datetime.fromtimestamp(start_time, tz=timezone.utc)
|
||||
|
||||
|
||||
def _create_context(logging_obj: "LiteLLMLoggingObj", custom_llm_provider: str) -> BackgroundInteractionCreateContext:
|
||||
metadata: Final = _OBJECT_MAPPING.validate_python(
|
||||
get_litellm_metadata_from_kwargs(kwargs=logging_obj.model_call_details)
|
||||
)
|
||||
litellm_params: Final = _OBJECT_MAPPING.validate_python(logging_obj.litellm_params)
|
||||
model: Final = logging_obj.model_call_details.get("model")
|
||||
return BackgroundInteractionCreateContext(
|
||||
model=model if isinstance(model, str) else logging_obj.model,
|
||||
call_type=_STRING.validate_python(logging_obj.call_type),
|
||||
litellm_call_id=logging_obj.litellm_call_id,
|
||||
function_id=logging_obj.function_id,
|
||||
litellm_trace_id=logging_obj.litellm_trace_id,
|
||||
start_time=_as_datetime(logging_obj.start_time),
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
metadata=_json_values((key, value) for key, value in metadata.items() if _carries(key)),
|
||||
custom_pricing=_json_values(
|
||||
(key, value) for key, value in litellm_params.items() if key in _CUSTOM_PRICING_KEYS and value is not None
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _rebuild_logging_obj(create_context: BackgroundInteractionCreateContext) -> "LiteLLMLoggingObj":
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
logging_obj: Final = Logging(
|
||||
model=create_context.model, # pyright: ignore[reportArgumentType] # function_setup builds the live object with the same None for an agent-only create
|
||||
messages=None,
|
||||
stream=False,
|
||||
call_type=create_context.call_type,
|
||||
start_time=create_context.start_time,
|
||||
litellm_call_id=create_context.litellm_call_id,
|
||||
function_id=create_context.function_id,
|
||||
litellm_trace_id=create_context.litellm_trace_id,
|
||||
)
|
||||
litellm_params: Final = {
|
||||
"metadata": dict(create_context.metadata),
|
||||
**create_context.custom_pricing,
|
||||
}
|
||||
logging_obj.update_environment_variables(
|
||||
litellm_params=litellm_params,
|
||||
optional_params={},
|
||||
model=create_context.model,
|
||||
custom_llm_provider=create_context.custom_llm_provider,
|
||||
)
|
||||
return logging_obj
|
||||
|
||||
|
||||
async def _settled_elsewhere(context: BackgroundInteractionPollContext) -> bool:
|
||||
try:
|
||||
return await context.store.is_claimed(context.interaction_id)
|
||||
except Exception as e: # noqa: BLE001 # an unreadable store must not stop the poll; the claim below decides
|
||||
verbose_logger.debug(
|
||||
"Could not read the settlement state of background interaction %s: %s", context.interaction_id, e
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
async def _claim(context: BackgroundInteractionPollContext) -> bool | None:
|
||||
"""
|
||||
Exactly-once gate between every settler of one interaction, on every
|
||||
replica: whoever claims first owns billing or release. ``None`` means the
|
||||
store could not answer, so nothing is owned and the caller retries later.
|
||||
"""
|
||||
try:
|
||||
return await context.store.claim(context.interaction_id)
|
||||
except Exception: # noqa: BLE001 # an unanswerable claim is retried on the next poll rather than billed twice
|
||||
verbose_logger.exception("Could not claim the settlement of background interaction %s", context.interaction_id)
|
||||
return None
|
||||
|
||||
|
||||
async def _record(context: BackgroundInteractionPollContext, outcome: SettlementOutcome) -> SettlementOutcome:
|
||||
try:
|
||||
await context.store.record_outcome(context.interaction_id, outcome)
|
||||
except Exception: # noqa: BLE001 # the outcome is an audit trail; the claim already made the settlement exclusive
|
||||
verbose_logger.exception("Could not record the settlement of background interaction %s", context.interaction_id)
|
||||
return outcome
|
||||
|
||||
|
||||
async def poll_and_log_background_interaction_cost(
|
||||
context: BackgroundInteractionPollContext,
|
||||
fetch_interaction: FetchInteraction = _fetch_interaction,
|
||||
) -> None:
|
||||
last_seen_status: str | None = None
|
||||
fetch_interaction: FetchInteraction = fetch_background_interaction,
|
||||
) -> SettlementOutcome | None:
|
||||
last_response: InteractionsAPIResponse | None = None # rebind-ok: the give-up path settles from the last poll
|
||||
for interval in _poll_intervals(
|
||||
initial=context.initial_interval_seconds,
|
||||
maximum=context.max_interval_seconds,
|
||||
timeout=context.timeout_seconds,
|
||||
):
|
||||
await asyncio.sleep(interval)
|
||||
if _is_settled(context.logging_obj):
|
||||
return
|
||||
if await _settled_elsewhere(context):
|
||||
return None
|
||||
try:
|
||||
response = await fetch_interaction(context)
|
||||
except Exception as e: # noqa: BLE001 # any fetch error must not kill the billing poll loop
|
||||
|
|
@ -133,26 +355,26 @@ async def poll_and_log_background_interaction_cost(
|
|||
e,
|
||||
)
|
||||
continue
|
||||
last_seen_status = response.status
|
||||
last_response = response
|
||||
if response.status not in _TERMINAL_STATUSES:
|
||||
continue
|
||||
if not _claim_settlement(context.logging_obj):
|
||||
return
|
||||
if response.usage is not None:
|
||||
await _bill_settled_interaction(logging_obj=context.logging_obj, response=response)
|
||||
else:
|
||||
await _release_open_budget_reservation(logging_obj=context.logging_obj)
|
||||
return
|
||||
if not _claim_settlement(context.logging_obj):
|
||||
return
|
||||
if last_seen_status is not None and last_seen_status not in _POLLABLE_STATUSES:
|
||||
if (claimed := await _claim(context)) is None:
|
||||
continue
|
||||
if not claimed:
|
||||
return None
|
||||
return await _record(context, await _settle_terminal(logging_obj=context.logging_obj, response=response))
|
||||
if not await _claim(context):
|
||||
return None
|
||||
if last_response is not None and last_response.status in _TERMINAL_STATUSES:
|
||||
return await _record(context, await _settle_terminal(logging_obj=context.logging_obj, response=last_response))
|
||||
if last_response is not None and last_response.status not in _POLLABLE_STATUSES:
|
||||
verbose_logger.error(
|
||||
"Gave up cost polling for background interaction %s after %ss: its last status %r is in neither "
|
||||
"the pollable nor the terminal set, so this proxy never learned how to settle it and its usage "
|
||||
"will not be tracked",
|
||||
context.interaction_id,
|
||||
context.timeout_seconds,
|
||||
last_seen_status,
|
||||
last_response.status,
|
||||
)
|
||||
else:
|
||||
verbose_logger.warning(
|
||||
|
|
@ -161,6 +383,15 @@ async def poll_and_log_background_interaction_cost(
|
|||
context.timeout_seconds,
|
||||
)
|
||||
await _release_open_budget_reservation(logging_obj=context.logging_obj)
|
||||
return await _record(context, "unsettled")
|
||||
|
||||
|
||||
async def _settle_terminal(logging_obj: "LiteLLMLoggingObj", response: InteractionsAPIResponse) -> SettlementOutcome:
|
||||
if response.status in _TERMINAL_STATUSES and response.usage is not None:
|
||||
await _bill_settled_interaction(logging_obj=logging_obj, response=response)
|
||||
return "billed"
|
||||
await _release_open_budget_reservation(logging_obj=logging_obj)
|
||||
return "released"
|
||||
|
||||
|
||||
async def _release_open_budget_reservation(logging_obj: "LiteLLMLoggingObj") -> None:
|
||||
|
|
@ -173,8 +404,8 @@ async def _release_open_budget_reservation(logging_obj: "LiteLLMLoggingObj") ->
|
|||
settlement must release the reservation here or the spend counters stay
|
||||
pinned at the estimated cost.
|
||||
"""
|
||||
metadata = get_litellm_metadata_from_kwargs(kwargs=logging_obj.model_call_details)
|
||||
budget_reservation = metadata.get("user_api_key_budget_reservation")
|
||||
metadata: Final = get_litellm_metadata_from_kwargs(kwargs=logging_obj.model_call_details)
|
||||
budget_reservation: Final = metadata.get("user_api_key_budget_reservation")
|
||||
if not isinstance(budget_reservation, dict):
|
||||
return
|
||||
|
||||
|
|
@ -234,24 +465,88 @@ def missing_usage_is_expected(response: InteractionsAPIResponse) -> bool:
|
|||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ActiveBackgroundPoll:
|
||||
task: "asyncio.Task[None]"
|
||||
task: "asyncio.Task[SettlementOutcome | None]"
|
||||
context: BackgroundInteractionPollContext
|
||||
|
||||
|
||||
_ACTIVE_POLLS: dict[str, _ActiveBackgroundPoll] = {} # mutable-ok: asyncio needs strong refs to running poll tasks
|
||||
_ACTIVE_POLLS: Final[dict[str, _ActiveBackgroundPoll]] = {} # mutable-ok: asyncio needs strong refs to poll tasks
|
||||
|
||||
|
||||
def _discard_poll(interaction_id: str, task: "asyncio.Task[None]") -> None:
|
||||
entry = _ACTIVE_POLLS.get(interaction_id)
|
||||
def _discard_poll(interaction_id: str, task: "asyncio.Task[SettlementOutcome | None]") -> None:
|
||||
entry: Final = _ACTIVE_POLLS.get(interaction_id)
|
||||
if entry is not None and entry.task is task:
|
||||
del _ACTIVE_POLLS[interaction_id]
|
||||
|
||||
|
||||
def maybe_schedule_background_interaction_cost_polling(
|
||||
def _track_poll(
|
||||
context: BackgroundInteractionPollContext, fetch_interaction: FetchInteraction
|
||||
) -> "asyncio.Task[SettlementOutcome | None]":
|
||||
task: Final = asyncio.create_task(poll_and_log_background_interaction_cost(context, fetch_interaction))
|
||||
_ACTIVE_POLLS[context.interaction_id] = _ActiveBackgroundPoll(task=task, context=context)
|
||||
task.add_done_callback(
|
||||
lambda finished, interaction_id=context.interaction_id: _discard_poll(interaction_id, finished)
|
||||
)
|
||||
return task
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _UnverifiedRegistrationStore:
|
||||
"""
|
||||
Store of a create whose registration raised, so whether its row landed is
|
||||
unknown until the durable store answers. The settlement claim asks it
|
||||
first, and only an interaction it reports as never stored settles through
|
||||
the local gate, which no other process can reach.
|
||||
"""
|
||||
|
||||
durable: BackgroundSettlementStore
|
||||
local: InMemoryBackgroundSettlementStore = field(default_factory=InMemoryBackgroundSettlementStore)
|
||||
|
||||
async def register(self, pending: PendingBackgroundInteraction) -> None:
|
||||
await self.durable.register(pending)
|
||||
|
||||
async def pending(self, interaction_id: str) -> PendingBackgroundInteraction | None:
|
||||
return await self.durable.pending(interaction_id)
|
||||
|
||||
async def is_claimed(self, interaction_id: str) -> bool:
|
||||
return await self.local.is_claimed(interaction_id) or await self.durable.is_claimed(interaction_id)
|
||||
|
||||
async def claim(self, interaction_id: str) -> bool:
|
||||
if await self.durable.claim(interaction_id):
|
||||
return True
|
||||
if await self.durable.is_claimed(interaction_id):
|
||||
return False
|
||||
return await self.local.claim(interaction_id)
|
||||
|
||||
async def record_outcome(self, interaction_id: str, outcome: SettlementOutcome) -> None:
|
||||
if await self.local.is_claimed(interaction_id):
|
||||
return
|
||||
await self.durable.record_outcome(interaction_id, outcome)
|
||||
|
||||
async def unclaimed(self) -> Sequence[PendingBackgroundInteraction]:
|
||||
return await self.durable.unclaimed()
|
||||
|
||||
|
||||
async def _registered_store(
|
||||
store: BackgroundSettlementStore, pending: PendingBackgroundInteraction
|
||||
) -> BackgroundSettlementStore:
|
||||
try:
|
||||
await store.register(pending)
|
||||
except Exception: # noqa: BLE001 # a store outage must not fail the create; the claim learns if the row landed
|
||||
verbose_logger.exception(
|
||||
"Could not durably register background interaction %s; its settlement claim decides whether the row landed",
|
||||
pending.interaction_id,
|
||||
)
|
||||
return _UnverifiedRegistrationStore(durable=store)
|
||||
return store
|
||||
|
||||
|
||||
async def maybe_schedule_background_interaction_cost_polling(
|
||||
response: object,
|
||||
create_kwargs: Mapping[str, object],
|
||||
custom_llm_provider: str,
|
||||
) -> "asyncio.Task[None] | None":
|
||||
store: BackgroundSettlementStore | None = None,
|
||||
fetch_interaction: FetchInteraction = fetch_background_interaction,
|
||||
) -> "asyncio.Task[SettlementOutcome | None] | None":
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
if not BACKGROUND_INTERACTION_COST_POLLING_ENABLED:
|
||||
|
|
@ -260,52 +555,141 @@ def maybe_schedule_background_interaction_cost_polling(
|
|||
return None
|
||||
if not is_pollable_background_interaction(response):
|
||||
return None
|
||||
logging_obj = create_kwargs.get("litellm_logging_obj")
|
||||
logging_obj: Final = create_kwargs.get("litellm_logging_obj")
|
||||
if not isinstance(logging_obj, Logging):
|
||||
return None
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return None
|
||||
api_key = create_kwargs.get("api_key")
|
||||
api_base = create_kwargs.get("api_base")
|
||||
context = BackgroundInteractionPollContext(
|
||||
api_key: Final = create_kwargs.get("api_key")
|
||||
api_base: Final = create_kwargs.get("api_base")
|
||||
pending: Final = PendingBackgroundInteraction(
|
||||
interaction_id=response.id,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
create_context=_create_context(logging_obj, custom_llm_provider),
|
||||
created_at=datetime.now(timezone.utc),
|
||||
)
|
||||
context: Final = BackgroundInteractionPollContext(
|
||||
interaction_id=response.id,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key if isinstance(api_key, str) else None,
|
||||
api_base=api_base if isinstance(api_base, str) else None,
|
||||
store=await _registered_store(store or _STORE.store, pending),
|
||||
)
|
||||
task = asyncio.create_task(poll_and_log_background_interaction_cost(context))
|
||||
_ACTIVE_POLLS[context.interaction_id] = _ActiveBackgroundPoll(task=task, context=context)
|
||||
task.add_done_callback(
|
||||
lambda finished, interaction_id=context.interaction_id: _discard_poll(interaction_id, finished)
|
||||
)
|
||||
return task
|
||||
return _track_poll(context, fetch_interaction)
|
||||
|
||||
|
||||
async def _pending(store: BackgroundSettlementStore, interaction_id: str) -> PendingBackgroundInteraction | None:
|
||||
try:
|
||||
return await store.pending(interaction_id)
|
||||
except Exception: # noqa: BLE001 # an unreadable store leaves the interaction to its poll or the counter TTL
|
||||
verbose_logger.exception("Could not look up background interaction %s before its delete", interaction_id)
|
||||
return None
|
||||
|
||||
|
||||
async def _fetch_before_delete(
|
||||
context: BackgroundInteractionPollContext, fetch_interaction: FetchInteraction
|
||||
) -> InteractionsAPIResponse | None:
|
||||
try:
|
||||
return await fetch_interaction(context)
|
||||
except Exception as e: # noqa: BLE001 # the caller decides what an unfetchable pre-delete state means
|
||||
verbose_logger.debug(
|
||||
"Could not fetch background interaction %s before its delete: %s", context.interaction_id, e
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
async def _settle_before_delete(
|
||||
context: BackgroundInteractionPollContext, response: InteractionsAPIResponse | None
|
||||
) -> SettlementOutcome | None:
|
||||
if not await _claim(context):
|
||||
return None
|
||||
if response is None:
|
||||
await _release_open_budget_reservation(logging_obj=context.logging_obj)
|
||||
return await _record(context, "released")
|
||||
return await _record(context, await _settle_terminal(logging_obj=context.logging_obj, response=response))
|
||||
|
||||
|
||||
async def maybe_settle_background_interaction_before_delete(
|
||||
interaction_id: str,
|
||||
fetch_interaction: FetchInteraction = _fetch_interaction,
|
||||
) -> None:
|
||||
entry = _ACTIVE_POLLS.get(interaction_id)
|
||||
if entry is None:
|
||||
return
|
||||
context = entry.context
|
||||
delete_kwargs: Mapping[str, object],
|
||||
fetch_interaction: FetchInteraction = fetch_background_interaction,
|
||||
store: BackgroundSettlementStore | None = None,
|
||||
) -> SettlementOutcome | None:
|
||||
entry: Final = _ACTIVE_POLLS.get(interaction_id)
|
||||
if entry is not None and not entry.context.resumed:
|
||||
return await _settle_before_delete(entry.context, await _fetch_before_delete(entry.context, fetch_interaction))
|
||||
settlement_store: Final = store or _STORE.store
|
||||
pending: Final = await _pending(settlement_store, interaction_id)
|
||||
if pending is None:
|
||||
return None
|
||||
api_key: Final = delete_kwargs.get("api_key")
|
||||
api_base: Final = delete_kwargs.get("api_base")
|
||||
context: Final = BackgroundInteractionPollContext(
|
||||
interaction_id=interaction_id,
|
||||
custom_llm_provider=pending.custom_llm_provider,
|
||||
logging_obj=_rebuild_logging_obj(pending.create_context),
|
||||
api_key=api_key if isinstance(api_key, str) else None,
|
||||
api_base=api_base if isinstance(api_base, str) else None,
|
||||
store=settlement_store,
|
||||
)
|
||||
try:
|
||||
response = await fetch_interaction(context)
|
||||
except Exception as e: # noqa: BLE001 # unfetchable pre-delete state settles by releasing the reservation
|
||||
response: Final = await fetch_interaction(context)
|
||||
except Exception:
|
||||
verbose_logger.debug(
|
||||
"Could not fetch background interaction %s before delete, releasing its reservation: %s",
|
||||
"Failing the delete of background interaction %s: this process could not fetch it with the delete's "
|
||||
"credentials, so the poll that created it keeps the bill",
|
||||
interaction_id,
|
||||
e,
|
||||
)
|
||||
if _claim_settlement(context.logging_obj):
|
||||
await _release_open_budget_reservation(logging_obj=context.logging_obj)
|
||||
return
|
||||
if not _claim_settlement(context.logging_obj):
|
||||
return
|
||||
if response.status in _TERMINAL_STATUSES and response.usage is not None:
|
||||
await _bill_settled_interaction(logging_obj=context.logging_obj, response=response)
|
||||
return
|
||||
await _release_open_budget_reservation(logging_obj=context.logging_obj)
|
||||
raise
|
||||
return await _settle_before_delete(context, response)
|
||||
|
||||
|
||||
async def _unclaimed(store: BackgroundSettlementStore) -> Sequence[PendingBackgroundInteraction]:
|
||||
try:
|
||||
return await store.unclaimed()
|
||||
except Exception: # noqa: BLE001 # an unreadable store at startup leaves its rows for the next boot
|
||||
verbose_logger.exception("Could not list the unsettled background interactions")
|
||||
return ()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PollSchedule:
|
||||
initial_interval_seconds: float = BACKGROUND_INTERACTION_COST_POLL_INITIAL_INTERVAL_SECONDS
|
||||
max_interval_seconds: float = BACKGROUND_INTERACTION_COST_POLL_MAX_INTERVAL_SECONDS
|
||||
timeout_seconds: float = BACKGROUND_INTERACTION_COST_POLL_TIMEOUT_SECONDS
|
||||
|
||||
|
||||
DEFAULT_POLL_SCHEDULE: Final = PollSchedule()
|
||||
|
||||
|
||||
def _resumed_context(
|
||||
row: PendingBackgroundInteraction, store: BackgroundSettlementStore, schedule: PollSchedule
|
||||
) -> BackgroundInteractionPollContext:
|
||||
age_seconds: Final = (datetime.now(timezone.utc) - row.created_at).total_seconds()
|
||||
return BackgroundInteractionPollContext(
|
||||
interaction_id=row.interaction_id,
|
||||
custom_llm_provider=row.custom_llm_provider,
|
||||
logging_obj=_rebuild_logging_obj(row.create_context),
|
||||
initial_interval_seconds=schedule.initial_interval_seconds,
|
||||
max_interval_seconds=schedule.max_interval_seconds,
|
||||
timeout_seconds=max(schedule.timeout_seconds - age_seconds, schedule.initial_interval_seconds),
|
||||
store=store,
|
||||
resumed=True,
|
||||
)
|
||||
|
||||
|
||||
async def resume_unsettled_background_interactions(
|
||||
store: BackgroundSettlementStore,
|
||||
fetch_interaction: FetchInteraction = fetch_background_interaction,
|
||||
schedule: PollSchedule = DEFAULT_POLL_SCHEDULE,
|
||||
) -> tuple["asyncio.Task[SettlementOutcome | None]", ...]:
|
||||
"""
|
||||
Pick up every settlement no process has claimed, which is what a replica
|
||||
that died mid-poll leaves behind. Each resumed poll keeps the remaining
|
||||
share of the original timeout and gets at least one fetch, so a completed
|
||||
interaction is still billed however late the resume comes.
|
||||
"""
|
||||
return tuple(
|
||||
_track_poll(_resumed_context(row, store, schedule), fetch_interaction)
|
||||
for row in await _unclaimed(store)
|
||||
if row.interaction_id not in _ACTIVE_POLLS
|
||||
)
|
||||
|
|
|
|||
|
|
@ -175,7 +175,7 @@ async def acreate(
|
|||
else:
|
||||
response = init_response
|
||||
|
||||
maybe_schedule_background_interaction_cost_polling(
|
||||
await maybe_schedule_background_interaction_cost_polling(
|
||||
response=response,
|
||||
create_kwargs=kwargs,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -464,7 +464,7 @@ async def adelete(
|
|||
extra_headers: dict[str, Any] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
**kwargs: object,
|
||||
) -> DeleteInteractionResult:
|
||||
"""Async: Delete an interaction by its ID."""
|
||||
local_vars: Final = locals()
|
||||
|
|
@ -472,7 +472,7 @@ async def adelete(
|
|||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["adelete_interaction"] = True
|
||||
|
||||
await maybe_settle_background_interaction_before_delete(interaction_id=interaction_id)
|
||||
await maybe_settle_background_interaction_before_delete(interaction_id=interaction_id, delete_kwargs=kwargs)
|
||||
|
||||
func: Final = partial(
|
||||
delete,
|
||||
|
|
|
|||
|
|
@ -313,6 +313,12 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
|
|||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> InteractionsAPIResponse:
|
||||
if not 200 <= raw_response.status_code < 300:
|
||||
raise GeminiError(
|
||||
message=raw_response.text,
|
||||
status_code=raw_response.status_code,
|
||||
headers=dict(raw_response.headers),
|
||||
)
|
||||
try:
|
||||
raw_json: Final = _interaction_body(raw_response)
|
||||
except Exception:
|
||||
|
|
|
|||
|
|
@ -777,6 +777,9 @@ from litellm.proxy.shutdown.scheduled_jobs import (
|
|||
pause_scheduled_jobs,
|
||||
stop_in_flight_scheduler_jobs,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.background_interaction_settlement import (
|
||||
install_background_interaction_settlement,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
get_budget_window_start,
|
||||
release_unbound_budget_reservation,
|
||||
|
|
@ -1394,6 +1397,7 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState
|
|||
await asyncio.sleep(5)
|
||||
|
||||
asyncio.create_task(_run_agent_grant_id_migration())
|
||||
await install_background_interaction_settlement(prisma_client)
|
||||
|
||||
## A coordination_redis block saved from the admin UI lives in the database,
|
||||
## which is only reachable once the prisma client exists. Apply it here, before
|
||||
|
|
|
|||
|
|
@ -1916,6 +1916,22 @@ model LiteLLM_WorkflowMessage {
|
|||
@@index([run_id])
|
||||
}
|
||||
|
||||
// Pending billing settlements for background interactions, keyed by the
|
||||
// interaction id so any replica can settle one that another replica created.
|
||||
// `claimed_at` is the exactly-once gate: the first conditional update wins.
|
||||
model LiteLLM_BackgroundInteractionSettlement {
|
||||
interaction_id String @id
|
||||
custom_llm_provider String
|
||||
create_context Json
|
||||
created_at DateTime @default(now())
|
||||
claimed_at DateTime?
|
||||
claimed_by String?
|
||||
settled_at DateTime?
|
||||
outcome String?
|
||||
|
||||
@@index([claimed_at], map: "idx_background_interaction_settlement_claimed_at")
|
||||
}
|
||||
|
||||
model LiteLLM_Lens {
|
||||
id String @id
|
||||
version Int @default(0)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,211 @@
|
|||
import asyncio
|
||||
import os
|
||||
import socket
|
||||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from itertools import chain
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Protocol, TypeVar
|
||||
|
||||
from pydantic import ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import BACKGROUND_INTERACTION_COST_POLLING_ENABLED
|
||||
from litellm.interactions.background_cost_polling import (
|
||||
DEFAULT_POLL_SCHEDULE,
|
||||
BackgroundInteractionCreateContext,
|
||||
FetchInteraction,
|
||||
PendingBackgroundInteraction,
|
||||
PollSchedule,
|
||||
SettlementOutcome,
|
||||
configure_background_settlement_store,
|
||||
fetch_background_interaction,
|
||||
resume_unsettled_background_interactions,
|
||||
)
|
||||
from litellm.repositories.table_repositories import BackgroundInteractionSettlementRepository
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
class _SettlementRow(Protocol):
|
||||
@property
|
||||
def interaction_id(self) -> str: ...
|
||||
@property
|
||||
def custom_llm_provider(self) -> str: ...
|
||||
@property
|
||||
def create_context(self) -> object: ...
|
||||
@property
|
||||
def created_at(self) -> datetime: ...
|
||||
@property
|
||||
def claimed_at(self) -> datetime | None: ...
|
||||
|
||||
|
||||
class _NewSettlementRow(TypedDict):
|
||||
interaction_id: ReadOnly[str]
|
||||
custom_llm_provider: ReadOnly[str]
|
||||
create_context: ReadOnly[object]
|
||||
created_at: ReadOnly[datetime]
|
||||
|
||||
|
||||
class _RowKey(TypedDict):
|
||||
interaction_id: ReadOnly[str]
|
||||
|
||||
|
||||
class _UnclaimedRowKey(TypedDict):
|
||||
interaction_id: ReadOnly[str]
|
||||
claimed_at: ReadOnly[None]
|
||||
|
||||
|
||||
class _UnclaimedRows(TypedDict):
|
||||
claimed_at: ReadOnly[None]
|
||||
|
||||
|
||||
class _Claim(TypedDict):
|
||||
claimed_at: ReadOnly[datetime]
|
||||
claimed_by: ReadOnly[str]
|
||||
|
||||
|
||||
class _Outcome(TypedDict):
|
||||
settled_at: ReadOnly[datetime]
|
||||
outcome: ReadOnly[SettlementOutcome]
|
||||
create_context: ReadOnly[object]
|
||||
|
||||
|
||||
class _SettlementTableActions(Protocol):
|
||||
def create(self, *, data: _NewSettlementRow) -> Awaitable[_SettlementRow]: ...
|
||||
|
||||
def find_unique(self, *, where: _RowKey) -> Awaitable[_SettlementRow | None]: ...
|
||||
|
||||
def find_many(self, *, where: _UnclaimedRows) -> Awaitable[Sequence[_SettlementRow]]: ...
|
||||
|
||||
def update_many(self, *, data: _Claim | _Outcome, where: _RowKey | _UnclaimedRowKey) -> Awaitable[int]: ...
|
||||
|
||||
|
||||
def _settlement_table(prisma_client: "PrismaClient") -> _SettlementTableActions:
|
||||
return BackgroundInteractionSettlementRepository(prisma_client).table
|
||||
|
||||
|
||||
_CLEARED_CREATE_CONTEXT: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
_T = TypeVar("_T")
|
||||
|
||||
|
||||
async def _read_from_a_table_that_may_not_exist(query: Awaitable[_T], when_missing: _T) -> _T:
|
||||
from prisma.errors import TableNotFoundError # noqa: PLC0415 # local import: prisma may be ungenerated at load
|
||||
|
||||
try:
|
||||
return await query
|
||||
except TableNotFoundError:
|
||||
return when_missing
|
||||
|
||||
|
||||
def _json(data: Mapping[str, object]) -> object:
|
||||
from prisma import Json # noqa: PLC0415 # local import: prisma may be ungenerated at module load in some tools
|
||||
|
||||
return Json.keys(**data)
|
||||
|
||||
|
||||
def _pending_rows(rows: Sequence[_SettlementRow]) -> tuple[PendingBackgroundInteraction, ...]:
|
||||
return tuple(chain.from_iterable(_pending_row(row) for row in rows))
|
||||
|
||||
|
||||
def _pending_row(row: _SettlementRow) -> tuple[PendingBackgroundInteraction, ...]:
|
||||
try:
|
||||
create_context: Final = BackgroundInteractionCreateContext.model_validate(row.create_context)
|
||||
except ValidationError:
|
||||
verbose_proxy_logger.exception(
|
||||
"Background interaction %s has a settlement row this version cannot read; leaving it unsettled",
|
||||
row.interaction_id,
|
||||
)
|
||||
return ()
|
||||
return (
|
||||
PendingBackgroundInteraction(
|
||||
interaction_id=row.interaction_id,
|
||||
custom_llm_provider=row.custom_llm_provider,
|
||||
create_context=create_context,
|
||||
created_at=row.created_at,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PrismaBackgroundSettlementStore:
|
||||
table: _SettlementTableActions
|
||||
claimed_by: str
|
||||
|
||||
async def register(self, pending: PendingBackgroundInteraction) -> None:
|
||||
await self.table.create(
|
||||
data=_NewSettlementRow(
|
||||
interaction_id=pending.interaction_id,
|
||||
custom_llm_provider=pending.custom_llm_provider,
|
||||
create_context=_json(pending.create_context.model_dump(mode="json")),
|
||||
created_at=pending.created_at,
|
||||
)
|
||||
)
|
||||
|
||||
async def pending(self, interaction_id: str) -> PendingBackgroundInteraction | None:
|
||||
row: Final = await self._row(interaction_id)
|
||||
if row is None or row.claimed_at is not None:
|
||||
return None
|
||||
return next(iter(_pending_row(row)), None)
|
||||
|
||||
async def is_claimed(self, interaction_id: str) -> bool:
|
||||
row: Final = await self._row(interaction_id)
|
||||
return row is not None and row.claimed_at is not None
|
||||
|
||||
async def claim(self, interaction_id: str) -> bool:
|
||||
claimed_rows: Final = await _read_from_a_table_that_may_not_exist(
|
||||
self.table.update_many(
|
||||
data=_Claim(claimed_at=datetime.now(timezone.utc), claimed_by=self.claimed_by),
|
||||
where=_UnclaimedRowKey(interaction_id=interaction_id, claimed_at=None),
|
||||
),
|
||||
when_missing=0,
|
||||
)
|
||||
return claimed_rows == 1
|
||||
|
||||
async def _row(self, interaction_id: str) -> _SettlementRow | None:
|
||||
return await _read_from_a_table_that_may_not_exist(
|
||||
self.table.find_unique(where=_RowKey(interaction_id=interaction_id)), when_missing=None
|
||||
)
|
||||
|
||||
async def record_outcome(self, interaction_id: str, outcome: SettlementOutcome) -> None:
|
||||
await self.table.update_many(
|
||||
data=_Outcome(
|
||||
settled_at=datetime.now(timezone.utc), outcome=outcome, create_context=_json(_CLEARED_CREATE_CONTEXT)
|
||||
),
|
||||
where=_RowKey(interaction_id=interaction_id),
|
||||
)
|
||||
|
||||
async def unclaimed(self) -> Sequence[PendingBackgroundInteraction]:
|
||||
return _pending_rows(await self.table.find_many(where=_UnclaimedRows(claimed_at=None)))
|
||||
|
||||
|
||||
async def configure_background_interaction_settlement(
|
||||
table: _SettlementTableActions,
|
||||
claimed_by: str,
|
||||
fetch_interaction: FetchInteraction = fetch_background_interaction,
|
||||
schedule: PollSchedule = DEFAULT_POLL_SCHEDULE,
|
||||
) -> tuple["asyncio.Task[SettlementOutcome | None]", ...]:
|
||||
if not BACKGROUND_INTERACTION_COST_POLLING_ENABLED:
|
||||
return ()
|
||||
store: Final = PrismaBackgroundSettlementStore(table=table, claimed_by=claimed_by)
|
||||
configure_background_settlement_store(store)
|
||||
resumed: Final = await resume_unsettled_background_interactions(store, fetch_interaction, schedule)
|
||||
if resumed:
|
||||
verbose_proxy_logger.info("Resumed cost polling for %s unsettled background interactions", len(resumed))
|
||||
return resumed
|
||||
|
||||
|
||||
async def install_background_interaction_settlement(prisma_client: "PrismaClient") -> None:
|
||||
try:
|
||||
await configure_background_interaction_settlement(
|
||||
table=BackgroundInteractionSettlementRepository(prisma_client).table,
|
||||
claimed_by=f"{socket.gethostname()}:{os.getpid()}",
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # a boot step must survive any DB error; billing then settles in-process as before
|
||||
verbose_proxy_logger.warning(
|
||||
"Durable background interaction settlement is off on this replica, so billing settles in-process only: %s",
|
||||
e,
|
||||
)
|
||||
|
|
@ -267,5 +267,11 @@ class AdaptiveRouterSessionRepository(PrismaTableRepository["prisma_models.LiteL
|
|||
table_name = "litellm_adaptiveroutersession"
|
||||
|
||||
|
||||
class BackgroundInteractionSettlementRepository(
|
||||
PrismaTableRepository["prisma_models.LiteLLM_BackgroundInteractionSettlement"]
|
||||
):
|
||||
table_name = "litellm_backgroundinteractionsettlement"
|
||||
|
||||
|
||||
class RetiredAgentRepository(PrismaTableRepository["prisma_models.LiteLLM_RetiredAgent"]):
|
||||
table_name = "litellm_retiredagent"
|
||||
|
|
|
|||
|
|
@ -1916,6 +1916,22 @@ model LiteLLM_WorkflowMessage {
|
|||
@@index([run_id])
|
||||
}
|
||||
|
||||
// Pending billing settlements for background interactions, keyed by the
|
||||
// interaction id so any replica can settle one that another replica created.
|
||||
// `claimed_at` is the exactly-once gate: the first conditional update wins.
|
||||
model LiteLLM_BackgroundInteractionSettlement {
|
||||
interaction_id String @id
|
||||
custom_llm_provider String
|
||||
create_context Json
|
||||
created_at DateTime @default(now())
|
||||
claimed_at DateTime?
|
||||
claimed_by String?
|
||||
settled_at DateTime?
|
||||
outcome String?
|
||||
|
||||
@@index([claimed_at], map: "idx_background_interaction_settlement_claimed_at")
|
||||
}
|
||||
|
||||
model LiteLLM_Lens {
|
||||
id String @id
|
||||
version Int @default(0)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -219,3 +219,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):
|
||||
|
|
@ -137,6 +150,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())
|
||||
|
|
@ -218,7 +232,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
|
||||
]
|
||||
}
|
||||
)
|
||||
|
|
@ -264,7 +285,14 @@ class Provider:
|
|||
if raw_body:
|
||||
body: Final = JSON_OBJECT.validate_json(raw_body)
|
||||
if isinstance(body, dict):
|
||||
self.observations.put(Observation(request.url.path, request.headers.get("authorization", ""), body))
|
||||
self.observations.put(
|
||||
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:])}"
|
||||
route: Final = next(
|
||||
|
|
@ -280,6 +308,58 @@ 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":
|
||||
if self.interactions.pop(interaction_id, None) is None:
|
||||
return JSONResponse(_interaction_not_found(interaction_id), status_code=404)
|
||||
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)
|
||||
|
|
@ -392,6 +472,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),
|
||||
|
|
@ -402,6 +495,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
|
||||
|
|
@ -411,9 +519,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,
|
||||
|
|
@ -421,19 +529,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,791 @@
|
|||
import math
|
||||
import socket
|
||||
import time
|
||||
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 estimates 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. A key budget below that
|
||||
is filled by the first create's reservation, so the next create is refused until a settlement releases it."""
|
||||
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 _booted_after(pid: int, moment: float) -> bool:
|
||||
try:
|
||||
return psutil.Process(pid).create_time() > moment
|
||||
except psutil.NoSuchProcess:
|
||||
return False
|
||||
|
||||
|
||||
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_after: Final = time.time()
|
||||
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)
|
||||
claimer: Final = _claimer_pid(_await_outcome(item, "billed"))
|
||||
assert claimer in pids or _booted_after(claimer, created_after), (claimer, 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=0.5 * _reservation_pin(rig.creator, model))
|
||||
janitor: Final = scenario.key()
|
||||
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 == 422 and pinned.json()["error"]["type"] == "budget_exceeded", pinned.text
|
||||
_state(rig, first, _completed())
|
||||
still_pinned: Final = _delete(rig.settler, first, key)
|
||||
assert still_pinned.status_code == 422 and still_pinned.json()["error"]["type"] == "budget_exceeded", (
|
||||
still_pinned.text
|
||||
)
|
||||
deleted: Final = _delete(rig.settler, first, janitor)
|
||||
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
|
||||
|
|
@ -1,18 +1,25 @@
|
|||
import asyncio
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
from itertools import islice
|
||||
from typing import Optional
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.interactions.background_cost_polling import (
|
||||
_SETTLED_KEY,
|
||||
_create_context,
|
||||
_poll_intervals,
|
||||
_rebuild_logging_obj,
|
||||
BackgroundInteractionPollContext,
|
||||
InMemoryBackgroundSettlementStore,
|
||||
maybe_schedule_background_interaction_cost_polling,
|
||||
maybe_settle_background_interaction_before_delete,
|
||||
PendingBackgroundInteraction,
|
||||
poll_and_log_background_interaction_cost,
|
||||
PollSchedule,
|
||||
resume_unsettled_background_interactions,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging
|
||||
from litellm.types.interactions import InteractionsAPIResponse
|
||||
|
||||
|
|
@ -63,7 +70,11 @@ async def _raise_on_billing(result: InteractionsAPIResponse) -> None:
|
|||
raise RuntimeError("cost calculation failed for a settled background interaction")
|
||||
|
||||
|
||||
def _context(logging_obj: LitellmLogging, timeout_seconds: float = 1.0) -> BackgroundInteractionPollContext:
|
||||
def _context(
|
||||
logging_obj: LitellmLogging,
|
||||
timeout_seconds: float = 1.0,
|
||||
store: Optional[InMemoryBackgroundSettlementStore] = None,
|
||||
) -> BackgroundInteractionPollContext:
|
||||
return BackgroundInteractionPollContext(
|
||||
interaction_id="interactions/bg-abc",
|
||||
custom_llm_provider="gemini",
|
||||
|
|
@ -71,6 +82,7 @@ def _context(logging_obj: LitellmLogging, timeout_seconds: float = 1.0) -> Backg
|
|||
initial_interval_seconds=0.001,
|
||||
max_interval_seconds=0.002,
|
||||
timeout_seconds=timeout_seconds,
|
||||
store=store if store is not None else InMemoryBackgroundSettlementStore(),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -246,10 +258,11 @@ async def test_poller_retries_after_fetch_error_and_still_bills():
|
|||
@pytest.mark.asyncio
|
||||
async def test_schedule_creates_poll_task_for_in_progress_create():
|
||||
logging_obj = _logging_obj()
|
||||
task = maybe_schedule_background_interaction_cost_polling(
|
||||
task = await maybe_schedule_background_interaction_cost_polling(
|
||||
response=_response("in_progress", with_usage=False),
|
||||
create_kwargs={"litellm_logging_obj": logging_obj},
|
||||
custom_llm_provider="gemini",
|
||||
store=InMemoryBackgroundSettlementStore(),
|
||||
)
|
||||
|
||||
assert isinstance(task, asyncio.Task)
|
||||
|
|
@ -258,6 +271,37 @@ async def test_schedule_creates_poll_task_for_in_progress_create():
|
|||
await task
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_schedule_registers_an_agent_only_create_that_names_no_model():
|
||||
logging_obj = LitellmLogging(
|
||||
model=None,
|
||||
messages=None,
|
||||
stream=False,
|
||||
call_type="acreate_interaction",
|
||||
start_time=time.time(),
|
||||
litellm_call_id="bg-agent-call-id",
|
||||
function_id="bg-agent-fn-id",
|
||||
)
|
||||
logging_obj.update_environment_variables(litellm_params={}, optional_params={}, custom_llm_provider="gemini")
|
||||
store = InMemoryBackgroundSettlementStore()
|
||||
|
||||
task = await maybe_schedule_background_interaction_cost_polling(
|
||||
response=_response("in_progress", with_usage=False),
|
||||
create_kwargs={"litellm_logging_obj": logging_obj},
|
||||
custom_llm_provider="gemini",
|
||||
store=store,
|
||||
)
|
||||
|
||||
assert isinstance(task, asyncio.Task)
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
pending = await store.pending("interactions/bg-abc")
|
||||
assert pending is not None
|
||||
assert pending.create_context.model is None
|
||||
assert _rebuild_logging_obj(pending.create_context).model is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"response,create_kwargs",
|
||||
|
|
@ -271,21 +315,22 @@ async def test_schedule_skips_non_pollable_results(response, create_kwargs):
|
|||
if create_kwargs.get("litellm_logging_obj") == "placeholder":
|
||||
create_kwargs = {"litellm_logging_obj": _logging_obj()}
|
||||
|
||||
task = maybe_schedule_background_interaction_cost_polling(
|
||||
task = await maybe_schedule_background_interaction_cost_polling(
|
||||
response=response,
|
||||
create_kwargs=create_kwargs,
|
||||
custom_llm_provider="gemini",
|
||||
store=InMemoryBackgroundSettlementStore(),
|
||||
)
|
||||
|
||||
assert task is None
|
||||
|
||||
|
||||
def _register_poll(logging_obj: LitellmLogging, poll_fetch=None) -> asyncio.Task:
|
||||
def _register_poll(logging_obj: LitellmLogging, poll_fetch=None, store=None) -> asyncio.Task:
|
||||
import litellm.interactions.background_cost_polling as bg
|
||||
|
||||
if poll_fetch is None:
|
||||
poll_fetch, _ = _fetch_sequence(_response("in_progress", with_usage=False))
|
||||
context = _context(logging_obj)
|
||||
context = _context(logging_obj, store=store)
|
||||
task = asyncio.create_task(poll_and_log_background_interaction_cost(context, fetch_interaction=poll_fetch))
|
||||
bg._ACTIVE_POLLS[context.interaction_id] = bg._ActiveBackgroundPoll(task=task, context=context)
|
||||
task.add_done_callback(lambda finished: bg._discard_poll(context.interaction_id, finished))
|
||||
|
|
@ -300,6 +345,7 @@ async def test_delete_settlement_bills_an_interaction_paused_for_a_tool_result()
|
|||
|
||||
await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc",
|
||||
delete_kwargs={},
|
||||
fetch_interaction=fetch,
|
||||
)
|
||||
|
||||
|
|
@ -316,6 +362,7 @@ async def test_delete_settlement_bills_pending_background_interaction():
|
|||
|
||||
await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc",
|
||||
delete_kwargs={},
|
||||
fetch_interaction=fetch,
|
||||
)
|
||||
|
||||
|
|
@ -334,6 +381,7 @@ async def test_delete_settlement_releases_reservation_when_still_in_progress():
|
|||
|
||||
await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc",
|
||||
delete_kwargs={},
|
||||
fetch_interaction=fetch,
|
||||
)
|
||||
|
||||
|
|
@ -351,6 +399,7 @@ async def test_delete_settlement_releases_reservation_when_prefetch_fails():
|
|||
|
||||
await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc",
|
||||
delete_kwargs={},
|
||||
fetch_interaction=fetch,
|
||||
)
|
||||
|
||||
|
|
@ -370,6 +419,7 @@ async def test_delete_settlement_releases_reservation_when_billing_raises():
|
|||
with pytest.raises(RuntimeError):
|
||||
await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc",
|
||||
delete_kwargs={},
|
||||
fetch_interaction=fetch,
|
||||
)
|
||||
|
||||
|
|
@ -383,6 +433,7 @@ async def test_delete_settlement_ignores_interactions_without_pending_poll():
|
|||
|
||||
await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/never-polled",
|
||||
delete_kwargs={},
|
||||
fetch_interaction=fetch,
|
||||
)
|
||||
|
||||
|
|
@ -400,6 +451,7 @@ async def test_delete_settlement_noop_after_poll_task_finished():
|
|||
settle_fetch, settle_calls = _fetch_sequence(_response("completed", with_usage=True))
|
||||
await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc",
|
||||
delete_kwargs={},
|
||||
fetch_interaction=settle_fetch,
|
||||
)
|
||||
|
||||
|
|
@ -409,12 +461,14 @@ async def test_delete_settlement_noop_after_poll_task_finished():
|
|||
@pytest.mark.asyncio
|
||||
async def test_delete_settlement_does_not_rebill_when_gate_already_claimed():
|
||||
logging_obj = _logging_obj()
|
||||
logging_obj.model_call_details[_SETTLED_KEY] = True
|
||||
task = _register_poll(logging_obj)
|
||||
store = InMemoryBackgroundSettlementStore()
|
||||
assert await store.claim("interactions/bg-abc")
|
||||
task = _register_poll(logging_obj, store=store)
|
||||
fetch, calls = _fetch_sequence(_response("completed", with_usage=True))
|
||||
|
||||
await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc",
|
||||
delete_kwargs={},
|
||||
fetch_interaction=fetch,
|
||||
)
|
||||
|
||||
|
|
@ -426,10 +480,11 @@ async def test_delete_settlement_does_not_rebill_when_gate_already_claimed():
|
|||
@pytest.mark.asyncio
|
||||
async def test_poller_exits_without_billing_once_settled_elsewhere():
|
||||
logging_obj = _logging_obj()
|
||||
logging_obj.model_call_details[_SETTLED_KEY] = True
|
||||
store = InMemoryBackgroundSettlementStore()
|
||||
assert await store.claim("interactions/bg-abc")
|
||||
fetch, calls = _fetch_sequence(_response("completed", with_usage=True))
|
||||
|
||||
await poll_and_log_background_interaction_cost(_context(logging_obj), fetch_interaction=fetch)
|
||||
await poll_and_log_background_interaction_cost(_context(logging_obj, store=store), fetch_interaction=fetch)
|
||||
|
||||
assert calls == []
|
||||
assert logging_obj.model_call_details.get("response_cost") is None
|
||||
|
|
@ -441,10 +496,11 @@ async def test_schedule_respects_kill_switch(monkeypatch):
|
|||
|
||||
monkeypatch.setattr(module, "BACKGROUND_INTERACTION_COST_POLLING_ENABLED", False)
|
||||
|
||||
task = maybe_schedule_background_interaction_cost_polling(
|
||||
task = await maybe_schedule_background_interaction_cost_polling(
|
||||
response=_response("in_progress", with_usage=False),
|
||||
create_kwargs={"litellm_logging_obj": _logging_obj()},
|
||||
custom_llm_provider="gemini",
|
||||
store=InMemoryBackgroundSettlementStore(),
|
||||
)
|
||||
|
||||
assert task is None
|
||||
|
|
@ -480,10 +536,11 @@ async def test_schedule_creates_poll_task_for_queued_create():
|
|||
without a poll task it is never charged at all.
|
||||
"""
|
||||
logging_obj = _logging_obj()
|
||||
task = maybe_schedule_background_interaction_cost_polling(
|
||||
task = await maybe_schedule_background_interaction_cost_polling(
|
||||
response=_response("queued", with_usage=False),
|
||||
create_kwargs={"litellm_logging_obj": logging_obj},
|
||||
custom_llm_provider="gemini",
|
||||
store=InMemoryBackgroundSettlementStore(),
|
||||
)
|
||||
|
||||
assert isinstance(task, asyncio.Task)
|
||||
|
|
@ -543,3 +600,492 @@ async def test_giving_up_on_an_unrecognized_status_says_which_status_it_was(monk
|
|||
|
||||
assert len(errors) == 1
|
||||
assert "halted_for_review" in errors[0]
|
||||
|
||||
|
||||
KEY_HASH = "0123456789abcdef" * 4
|
||||
|
||||
FAST_SCHEDULE = PollSchedule(initial_interval_seconds=0.001, max_interval_seconds=0.002, timeout_seconds=1.0)
|
||||
|
||||
|
||||
def _capturing_fetch(response: InteractionsAPIResponse):
|
||||
captured = []
|
||||
|
||||
async def fetch(context):
|
||||
captured.append(context)
|
||||
return response
|
||||
|
||||
return fetch, captured
|
||||
|
||||
|
||||
def _create_metadata(**extra) -> dict:
|
||||
return {
|
||||
"user_api_key": KEY_HASH,
|
||||
"user_api_key_team_id": "team-1",
|
||||
"user_api_key_auth": object(),
|
||||
**extra,
|
||||
}
|
||||
|
||||
|
||||
async def _create_on_a_replica_that_then_dies(logging_obj: LitellmLogging, store) -> None:
|
||||
import litellm.interactions.background_cost_polling as bg
|
||||
|
||||
poll_fetch, _ = _fetch_sequence(_response("in_progress", with_usage=False))
|
||||
task = await maybe_schedule_background_interaction_cost_polling(
|
||||
response=_response("in_progress", with_usage=False),
|
||||
create_kwargs={"litellm_logging_obj": logging_obj},
|
||||
custom_llm_provider="gemini",
|
||||
store=store,
|
||||
fetch_interaction=poll_fetch,
|
||||
)
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
await asyncio.sleep(0)
|
||||
assert "interactions/bg-abc" not in bg._ACTIVE_POLLS
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_on_another_replica_bills_the_create_from_the_store():
|
||||
"""
|
||||
The regression: the replica that served the create owns the poll task, so
|
||||
a delete served by any other replica used to find nothing to settle and
|
||||
the work went unbilled. The store carries the create's attribution, never
|
||||
its auth object, to whichever replica settles.
|
||||
"""
|
||||
store = InMemoryBackgroundSettlementStore()
|
||||
logging_obj = _logging_obj(litellm_params={"metadata": _create_metadata()})
|
||||
await _create_on_a_replica_that_then_dies(logging_obj, store)
|
||||
fetch, captured = _capturing_fetch(_response("completed", with_usage=True))
|
||||
|
||||
outcome = await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc",
|
||||
delete_kwargs={},
|
||||
fetch_interaction=fetch,
|
||||
store=store,
|
||||
)
|
||||
|
||||
assert outcome == "billed"
|
||||
settled = captured[0].logging_obj
|
||||
assert settled is not logging_obj
|
||||
assert settled.model_call_details["response_cost"] > 0
|
||||
payload_metadata = settled.model_call_details["standard_logging_object"]["metadata"]
|
||||
assert payload_metadata["user_api_key_hash"] == KEY_HASH
|
||||
assert payload_metadata["user_api_key_team_id"] == "team-1"
|
||||
assert "user_api_key_auth" not in get_litellm_metadata_from_kwargs(kwargs=settled.model_call_details)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_on_another_replica_releases_the_create_reservation():
|
||||
store = InMemoryBackgroundSettlementStore()
|
||||
logging_obj = _logging_obj(
|
||||
litellm_params={"metadata": _create_metadata(user_api_key_budget_reservation=_reservation())}
|
||||
)
|
||||
await _create_on_a_replica_that_then_dies(logging_obj, store)
|
||||
fetch, captured = _capturing_fetch(_response("in_progress", with_usage=False))
|
||||
|
||||
outcome = await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc",
|
||||
delete_kwargs={},
|
||||
fetch_interaction=fetch,
|
||||
store=store,
|
||||
)
|
||||
|
||||
assert outcome == "released"
|
||||
settled_metadata = get_litellm_metadata_from_kwargs(kwargs=captured[0].logging_obj.model_call_details)
|
||||
assert settled_metadata["user_api_key_budget_reservation"]["finalized"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_on_another_replica_fails_when_it_cannot_fetch_and_leaves_the_bill_to_the_creating_poll():
|
||||
"""
|
||||
The settling replica fetches with the delete's credentials, never the
|
||||
create's, so a fetch it cannot make (a key only the deployment carries)
|
||||
says nothing about the interaction. Deleting anyway would strand the bill
|
||||
behind a deleted interaction, so the delete fails with the fetch's error
|
||||
and the poll on the creating replica still owns the bill.
|
||||
"""
|
||||
store = InMemoryBackgroundSettlementStore()
|
||||
logging_obj = _logging_obj(litellm_params={"metadata": _create_metadata()})
|
||||
await _create_on_a_replica_that_then_dies(logging_obj, store)
|
||||
fetch, _ = _fetch_sequence(RuntimeError("Google API key is required"))
|
||||
|
||||
with pytest.raises(RuntimeError, match="Google API key is required"):
|
||||
await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc", delete_kwargs={}, fetch_interaction=fetch, store=store
|
||||
)
|
||||
|
||||
assert await store.is_claimed("interactions/bg-abc") is False
|
||||
poll_fetch, _ = _fetch_sequence(_response("completed", with_usage=True))
|
||||
await asyncio.wait_for(_register_poll(logging_obj, poll_fetch=poll_fetch, store=store), timeout=5)
|
||||
assert logging_obj.model_call_details["response_cost"] > 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_settles_once_however_many_replicas_try():
|
||||
store = InMemoryBackgroundSettlementStore()
|
||||
await _create_on_a_replica_that_then_dies(_logging_obj(litellm_params={"metadata": _create_metadata()}), store)
|
||||
first_fetch, first_calls = _capturing_fetch(_response("completed", with_usage=True))
|
||||
second_fetch, second_calls = _capturing_fetch(_response("completed", with_usage=True))
|
||||
|
||||
first = await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc", delete_kwargs={}, fetch_interaction=first_fetch, store=store
|
||||
)
|
||||
second = await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc", delete_kwargs={}, fetch_interaction=second_fetch, store=store
|
||||
)
|
||||
|
||||
assert (first, second) == ("billed", None)
|
||||
assert len(first_calls) == 1
|
||||
assert second_calls == []
|
||||
|
||||
|
||||
def test_create_context_carries_no_request_headers():
|
||||
logging_obj = _logging_obj(
|
||||
litellm_params={
|
||||
"metadata": _create_metadata(
|
||||
requester_custom_headers={"x-api-key": "sk-customer-secret"},
|
||||
proxy_server_request={"headers": {"x-api-key": "sk-customer-secret"}},
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
carried = _create_context(logging_obj, "gemini").metadata
|
||||
|
||||
assert carried["user_api_key_team_id"] == "team-1"
|
||||
assert "requester_custom_headers" not in carried
|
||||
assert "proxy_server_request" not in carried
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_restart_resumes_only_the_rows_no_replica_claimed():
|
||||
store = InMemoryBackgroundSettlementStore()
|
||||
create_context = _create_context(_logging_obj(litellm_params={"metadata": _create_metadata()}), "gemini")
|
||||
for interaction_id in ("interactions/bg-orphaned", "interactions/bg-settled"):
|
||||
await store.register(
|
||||
PendingBackgroundInteraction(
|
||||
interaction_id=interaction_id,
|
||||
custom_llm_provider="gemini",
|
||||
create_context=create_context,
|
||||
created_at=datetime.now(timezone.utc),
|
||||
)
|
||||
)
|
||||
assert await store.claim("interactions/bg-settled")
|
||||
fetch, captured = _capturing_fetch(_response("completed", with_usage=True))
|
||||
|
||||
resumed = await resume_unsettled_background_interactions(store, fetch, schedule=FAST_SCHEDULE)
|
||||
|
||||
assert len(resumed) == 1
|
||||
assert await asyncio.wait_for(resumed[0], timeout=5) == "billed"
|
||||
assert [context.interaction_id for context in captured] == ["interactions/bg-orphaned"]
|
||||
assert captured[0].logging_obj.model_call_details["response_cost"] > 0
|
||||
assert await store.is_claimed("interactions/bg-orphaned")
|
||||
|
||||
|
||||
class _ClaimAnswersOnlyAfterTheLastFetch:
|
||||
def __init__(self):
|
||||
self.store = InMemoryBackgroundSettlementStore()
|
||||
self.fetches = 0
|
||||
self.fetches_at_last_claim = -1
|
||||
|
||||
async def fetch(self, context):
|
||||
self.fetches += 1
|
||||
return _response("completed", with_usage=True)
|
||||
|
||||
async def register(self, pending):
|
||||
await self.store.register(pending)
|
||||
|
||||
async def pending(self, interaction_id):
|
||||
return await self.store.pending(interaction_id)
|
||||
|
||||
async def is_claimed(self, interaction_id):
|
||||
return await self.store.is_claimed(interaction_id)
|
||||
|
||||
async def claim(self, interaction_id):
|
||||
if self.fetches != self.fetches_at_last_claim:
|
||||
self.fetches_at_last_claim = self.fetches
|
||||
raise RuntimeError("database unavailable")
|
||||
return await self.store.claim(interaction_id)
|
||||
|
||||
async def record_outcome(self, interaction_id, outcome):
|
||||
return None
|
||||
|
||||
async def unclaimed(self):
|
||||
return await self.store.unclaimed()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_poller_bills_the_completed_response_it_saw_when_the_claim_only_answers_at_the_deadline():
|
||||
logging_obj = _logging_obj()
|
||||
store = _ClaimAnswersOnlyAfterTheLastFetch()
|
||||
|
||||
outcome = await poll_and_log_background_interaction_cost(
|
||||
_context(logging_obj, timeout_seconds=0.01, store=store),
|
||||
fetch_interaction=store.fetch,
|
||||
)
|
||||
|
||||
assert store.fetches >= 2
|
||||
assert outcome == "billed"
|
||||
assert logging_obj.model_call_details["response_cost"] > 0
|
||||
|
||||
|
||||
class _DownStore:
|
||||
async def register(self, pending):
|
||||
raise RuntimeError("database unavailable")
|
||||
|
||||
async def pending(self, interaction_id):
|
||||
raise RuntimeError("database unavailable")
|
||||
|
||||
async def is_claimed(self, interaction_id):
|
||||
raise RuntimeError("database unavailable")
|
||||
|
||||
async def claim(self, interaction_id):
|
||||
raise RuntimeError("database unavailable")
|
||||
|
||||
async def record_outcome(self, interaction_id, outcome):
|
||||
raise RuntimeError("database unavailable")
|
||||
|
||||
async def unclaimed(self):
|
||||
raise RuntimeError("database unavailable")
|
||||
|
||||
|
||||
class _RegistersThenRaises:
|
||||
def __init__(self):
|
||||
self.store = InMemoryBackgroundSettlementStore()
|
||||
|
||||
async def register(self, pending):
|
||||
await self.store.register(pending)
|
||||
raise RuntimeError("connection reset after the row was committed")
|
||||
|
||||
async def pending(self, interaction_id):
|
||||
return await self.store.pending(interaction_id)
|
||||
|
||||
async def is_claimed(self, interaction_id):
|
||||
return await self.store.is_claimed(interaction_id)
|
||||
|
||||
async def claim(self, interaction_id):
|
||||
return await self.store.claim(interaction_id)
|
||||
|
||||
async def record_outcome(self, interaction_id, outcome):
|
||||
return None
|
||||
|
||||
async def unclaimed(self):
|
||||
return await self.store.unclaimed()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_whose_registration_raised_after_landing_still_claims_the_stored_row():
|
||||
"""
|
||||
A registration that raises after its row committed used to move the poll
|
||||
to a private in-memory gate, so the creating worker billed while the
|
||||
stored row stayed unclaimed for another replica's delete or the next boot
|
||||
to bill again. The row that landed is the gate every settler shares.
|
||||
"""
|
||||
store = _RegistersThenRaises()
|
||||
logging_obj = _logging_obj(litellm_params={"metadata": _create_metadata()})
|
||||
poll_fetch, _ = _fetch_sequence(_response("in_progress", with_usage=False))
|
||||
task = await maybe_schedule_background_interaction_cost_polling(
|
||||
response=_response("in_progress", with_usage=False),
|
||||
create_kwargs={"litellm_logging_obj": logging_obj},
|
||||
custom_llm_provider="gemini",
|
||||
store=store,
|
||||
fetch_interaction=poll_fetch,
|
||||
)
|
||||
fetch, _ = _capturing_fetch(_response("completed", with_usage=True))
|
||||
|
||||
outcome = await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc", delete_kwargs={}, fetch_interaction=fetch, store=store
|
||||
)
|
||||
|
||||
assert outcome == "billed"
|
||||
assert await store.is_claimed("interactions/bg-abc")
|
||||
assert await store.unclaimed() == ()
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_on_a_worker_that_resumed_the_poll_fails_when_it_cannot_fetch():
|
||||
"""
|
||||
After a restart every worker resumes the unclaimed rows, so none of them
|
||||
is the creator whose delete may release and delete on a failed fetch. A
|
||||
resumed worker's delete fails like any other replica's, and its own poll
|
||||
still bills the interaction once it completes.
|
||||
"""
|
||||
store = InMemoryBackgroundSettlementStore()
|
||||
await store.register(
|
||||
PendingBackgroundInteraction(
|
||||
interaction_id="interactions/bg-abc",
|
||||
custom_llm_provider="gemini",
|
||||
create_context=_create_context(_logging_obj(litellm_params={"metadata": _create_metadata()}), "gemini"),
|
||||
created_at=datetime.now(timezone.utc),
|
||||
)
|
||||
)
|
||||
responses = [_response("in_progress", with_usage=False)]
|
||||
|
||||
async def poll_fetch(context):
|
||||
return responses[-1]
|
||||
|
||||
(resumed,) = await resume_unsettled_background_interactions(store, poll_fetch, schedule=FAST_SCHEDULE)
|
||||
fetch, _ = _fetch_sequence(RuntimeError("Google API key is required"))
|
||||
|
||||
with pytest.raises(RuntimeError, match="Google API key is required"):
|
||||
await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc", delete_kwargs={}, fetch_interaction=fetch, store=store
|
||||
)
|
||||
|
||||
assert await store.is_claimed("interactions/bg-abc") is False
|
||||
responses.append(_response("completed", with_usage=True))
|
||||
assert await asyncio.wait_for(resumed, timeout=5) == "billed"
|
||||
|
||||
|
||||
class _LandsThenGoesDown:
|
||||
"""Register commits the row and loses its acknowledgement; every read fails until the store recovers."""
|
||||
|
||||
def __init__(self):
|
||||
self.store = InMemoryBackgroundSettlementStore()
|
||||
self.down = True
|
||||
|
||||
async def register(self, pending):
|
||||
await self.store.register(pending)
|
||||
raise RuntimeError("connection reset after the row was committed")
|
||||
|
||||
async def pending(self, interaction_id):
|
||||
self._answer()
|
||||
return await self.store.pending(interaction_id)
|
||||
|
||||
async def is_claimed(self, interaction_id):
|
||||
self._answer()
|
||||
return await self.store.is_claimed(interaction_id)
|
||||
|
||||
async def claim(self, interaction_id):
|
||||
self._answer()
|
||||
return await self.store.claim(interaction_id)
|
||||
|
||||
async def record_outcome(self, interaction_id, outcome):
|
||||
return None
|
||||
|
||||
async def unclaimed(self):
|
||||
self._answer()
|
||||
return await self.store.unclaimed()
|
||||
|
||||
def _answer(self):
|
||||
if self.down:
|
||||
raise RuntimeError("database unavailable")
|
||||
|
||||
|
||||
class _TableLessStore:
|
||||
"""A replica whose database never got the settlement table: writes fail and reads see no rows."""
|
||||
|
||||
async def register(self, pending):
|
||||
raise RuntimeError("the settlement table does not exist")
|
||||
|
||||
async def pending(self, interaction_id):
|
||||
return None
|
||||
|
||||
async def is_claimed(self, interaction_id):
|
||||
return False
|
||||
|
||||
async def claim(self, interaction_id):
|
||||
return False
|
||||
|
||||
async def record_outcome(self, interaction_id, outcome):
|
||||
raise RuntimeError("the settlement table does not exist")
|
||||
|
||||
async def unclaimed(self):
|
||||
raise RuntimeError("the settlement table does not exist")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_whose_registration_and_read_back_both_failed_bills_once_through_the_landed_row():
|
||||
"""
|
||||
A registration that raised and could not be read back used to give the
|
||||
creator a private in-memory gate, so it billed while the stored row stayed
|
||||
unclaimed for the next boot to resume and bill again. With the durable
|
||||
state unknown, the claim waits for the store and settles through the row.
|
||||
"""
|
||||
store = _LandsThenGoesDown()
|
||||
logging_obj = _logging_obj(litellm_params={"metadata": _create_metadata()})
|
||||
poll_fetch, _ = _fetch_sequence(_response("in_progress", with_usage=False))
|
||||
task = await maybe_schedule_background_interaction_cost_polling(
|
||||
response=_response("in_progress", with_usage=False),
|
||||
create_kwargs={"litellm_logging_obj": logging_obj},
|
||||
custom_llm_provider="gemini",
|
||||
store=store,
|
||||
fetch_interaction=poll_fetch,
|
||||
)
|
||||
fetch, calls = _fetch_sequence(_response("completed", with_usage=True), _response("completed", with_usage=True))
|
||||
|
||||
while_down = await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc", delete_kwargs={}, fetch_interaction=fetch, store=store
|
||||
)
|
||||
store.down = False
|
||||
recovered = await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc", delete_kwargs={}, fetch_interaction=fetch, store=store
|
||||
)
|
||||
|
||||
assert (while_down, recovered) == (None, "billed")
|
||||
assert len(calls) == 2
|
||||
assert await store.is_claimed("interactions/bg-abc")
|
||||
assert await store.unclaimed() == ()
|
||||
assert await resume_unsettled_background_interactions(store, poll_fetch, schedule=FAST_SCHEDULE) == ()
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_whose_store_never_answers_is_not_billed_through_a_private_gate():
|
||||
logging_obj = _logging_obj()
|
||||
poll_fetch, _ = _fetch_sequence(_response("in_progress", with_usage=False))
|
||||
task = await maybe_schedule_background_interaction_cost_polling(
|
||||
response=_response("in_progress", with_usage=False),
|
||||
create_kwargs={"litellm_logging_obj": logging_obj},
|
||||
custom_llm_provider="gemini",
|
||||
store=_DownStore(),
|
||||
fetch_interaction=poll_fetch,
|
||||
)
|
||||
fetch, calls = _fetch_sequence(_response("completed", with_usage=True))
|
||||
|
||||
outcome = await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc",
|
||||
delete_kwargs={},
|
||||
fetch_interaction=fetch,
|
||||
store=_DownStore(),
|
||||
)
|
||||
|
||||
assert outcome is None
|
||||
assert len(calls) == 1
|
||||
assert "response_cost" not in logging_obj.model_call_details
|
||||
assert not task.done()
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_on_a_replica_without_the_settlement_table_still_settles_in_process():
|
||||
logging_obj = _logging_obj()
|
||||
poll_fetch, _ = _fetch_sequence(_response("in_progress", with_usage=False))
|
||||
task = await maybe_schedule_background_interaction_cost_polling(
|
||||
response=_response("in_progress", with_usage=False),
|
||||
create_kwargs={"litellm_logging_obj": logging_obj},
|
||||
custom_llm_provider="gemini",
|
||||
store=_TableLessStore(),
|
||||
fetch_interaction=poll_fetch,
|
||||
)
|
||||
fetch, calls = _fetch_sequence(_response("completed", with_usage=True), _response("completed", with_usage=True))
|
||||
|
||||
outcome = await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc", delete_kwargs={}, fetch_interaction=fetch, store=_TableLessStore()
|
||||
)
|
||||
again = await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-abc", delete_kwargs={}, fetch_interaction=fetch, store=_TableLessStore()
|
||||
)
|
||||
|
||||
assert (outcome, again) == ("billed", None)
|
||||
assert len(calls) == 2
|
||||
assert logging_obj.model_call_details["response_cost"] > 0
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
|
|
|||
|
|
@ -10,12 +10,13 @@ Covers:
|
|||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
|
||||
from litellm.interactions.litellm_responses_transformation.streaming_iterator import (
|
||||
LiteLLMResponsesInteractionsStreamingIterator,
|
||||
)
|
||||
from litellm.llms.gemini.common_utils import GeminiError
|
||||
from litellm.llms.gemini.interactions.transformation import (
|
||||
GoogleAIStudioInteractionsConfig,
|
||||
)
|
||||
|
|
@ -464,6 +465,32 @@ class TestInteractionOperationUrls:
|
|||
)
|
||||
|
||||
|
||||
class TestGetInteractionResponse:
|
||||
@pytest.mark.parametrize("status_code", [404, 500])
|
||||
def test_non_2xx_raises_even_when_the_error_body_is_json(
|
||||
self, config: GoogleAIStudioInteractionsConfig, status_code: int
|
||||
) -> None:
|
||||
raw_response = httpx.Response(
|
||||
status_code,
|
||||
json={"error": {"code": status_code, "message": "boom", "status": "INTERNAL"}},
|
||||
request=httpx.Request("GET", "https://generativelanguage.googleapis.com/v1beta/interactions/x"),
|
||||
)
|
||||
with pytest.raises(GeminiError) as raised:
|
||||
config.transform_get_interaction_response(raw_response=raw_response, logging_obj=MagicMock())
|
||||
assert raised.value.status_code == status_code
|
||||
assert "boom" in str(raised.value)
|
||||
|
||||
def test_2xx_parses_the_interaction(self, config: GoogleAIStudioInteractionsConfig) -> None:
|
||||
raw_response = httpx.Response(
|
||||
200,
|
||||
json={"id": "interaction-1", "object": "interaction", "status": "completed", "steps": []},
|
||||
request=httpx.Request("GET", "https://generativelanguage.googleapis.com/v1beta/interactions/x"),
|
||||
)
|
||||
response = config.transform_get_interaction_response(raw_response=raw_response, logging_obj=MagicMock())
|
||||
assert response.id == "interaction-1"
|
||||
assert response.status == "completed"
|
||||
|
||||
|
||||
class TestTransformRequestSchemaCoalescing:
|
||||
"""Test new-schema request coalescing (Api-Revision: 2026-05-20)."""
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,298 @@
|
|||
import asyncio
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm.interactions.background_cost_polling as bg
|
||||
from litellm.interactions.background_cost_polling import (
|
||||
_create_context,
|
||||
configure_background_settlement_store,
|
||||
maybe_settle_background_interaction_before_delete,
|
||||
PendingBackgroundInteraction,
|
||||
PollSchedule,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging
|
||||
from litellm.proxy.spend_tracking.background_interaction_settlement import (
|
||||
configure_background_interaction_settlement,
|
||||
install_background_interaction_settlement,
|
||||
PrismaBackgroundSettlementStore,
|
||||
)
|
||||
from litellm.types.interactions import InteractionsAPIResponse
|
||||
|
||||
USAGE_BLOCK = {
|
||||
"total_tokens": 175,
|
||||
"total_input_tokens": 100,
|
||||
"input_tokens_by_modality": [{"modality": "text", "tokens": 100}],
|
||||
"total_cached_tokens": 0,
|
||||
"total_output_tokens": 50,
|
||||
"output_tokens_by_modality": [{"modality": "text", "tokens": 50}],
|
||||
"total_tool_use_tokens": 0,
|
||||
"total_thought_tokens": 25,
|
||||
}
|
||||
|
||||
FAST_SCHEDULE = PollSchedule(initial_interval_seconds=0.001, max_interval_seconds=0.002, timeout_seconds=1.0)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Row:
|
||||
interaction_id: str
|
||||
custom_llm_provider: str
|
||||
create_context: object
|
||||
created_at: datetime
|
||||
claimed_at: Optional[datetime] = None
|
||||
claimed_by: Optional[str] = None
|
||||
settled_at: Optional[datetime] = None
|
||||
outcome: Optional[str] = None
|
||||
|
||||
|
||||
class _FakeSettlementTable:
|
||||
"""Just enough of prisma's per-model actions: Json is stored as the data it wraps and read back parsed."""
|
||||
|
||||
def __init__(self, rows: tuple[_Row, ...] = ()):
|
||||
self.rows = {row.interaction_id: row for row in rows}
|
||||
|
||||
async def create(self, *, data):
|
||||
row = _Row(
|
||||
interaction_id=data["interaction_id"],
|
||||
custom_llm_provider=data["custom_llm_provider"],
|
||||
create_context=data["create_context"].data,
|
||||
created_at=data["created_at"],
|
||||
)
|
||||
self.rows[row.interaction_id] = row
|
||||
return row
|
||||
|
||||
async def find_unique(self, *, where):
|
||||
return self.rows.get(where["interaction_id"])
|
||||
|
||||
async def find_many(self, *, where):
|
||||
return self._matching(where)
|
||||
|
||||
async def update_many(self, *, data, where):
|
||||
matched = self._matching(where)
|
||||
for row in matched:
|
||||
for column, value in data.items():
|
||||
setattr(row, column, getattr(value, "data", value) if column == "create_context" else value)
|
||||
return len(matched)
|
||||
|
||||
def _matching(self, where) -> list:
|
||||
return [row for row in self.rows.values() if all(getattr(row, column) == value for column, value in where.items())]
|
||||
|
||||
|
||||
def _logging_obj(metadata: Optional[dict] = None) -> LitellmLogging:
|
||||
logging_obj = LitellmLogging(
|
||||
model="gemini-2.5-flash",
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="acreate_interaction",
|
||||
start_time=time.time(),
|
||||
litellm_call_id="bg-settlement-call-id",
|
||||
function_id="bg-settlement-fn-id",
|
||||
)
|
||||
logging_obj.update_environment_variables(
|
||||
litellm_params={"metadata": metadata or {"user_api_key": "0123456789abcdef" * 4}},
|
||||
optional_params={},
|
||||
model="gemini-2.5-flash",
|
||||
custom_llm_provider="gemini",
|
||||
input="hi",
|
||||
)
|
||||
return logging_obj
|
||||
|
||||
|
||||
def _pending(interaction_id: str) -> PendingBackgroundInteraction:
|
||||
return PendingBackgroundInteraction(
|
||||
interaction_id=interaction_id,
|
||||
custom_llm_provider="gemini",
|
||||
create_context=_create_context(_logging_obj(), "gemini"),
|
||||
created_at=datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
|
||||
def _stored_row(interaction_id: str, claimed: bool = False, create_context: Optional[object] = None) -> _Row:
|
||||
return _Row(
|
||||
interaction_id=interaction_id,
|
||||
custom_llm_provider="gemini",
|
||||
create_context=(
|
||||
create_context
|
||||
if create_context is not None
|
||||
else _create_context(_logging_obj(), "gemini").model_dump(mode="json")
|
||||
),
|
||||
created_at=datetime.now(timezone.utc),
|
||||
claimed_at=datetime.now(timezone.utc) if claimed else None,
|
||||
claimed_by="replica-a:1" if claimed else None,
|
||||
)
|
||||
|
||||
|
||||
def _completed(interaction_id: str) -> InteractionsAPIResponse:
|
||||
return InteractionsAPIResponse(
|
||||
id=interaction_id, model="gemini-2.5-flash", status="completed", steps=[], usage=dict(USAGE_BLOCK)
|
||||
)
|
||||
|
||||
|
||||
def _capturing_fetch():
|
||||
captured = []
|
||||
|
||||
async def fetch(context):
|
||||
captured.append(context)
|
||||
return _completed(context.interaction_id)
|
||||
|
||||
return fetch, captured
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_registered_row_reads_back_as_the_same_pending_interaction():
|
||||
table = _FakeSettlementTable()
|
||||
store = PrismaBackgroundSettlementStore(table=table, claimed_by="replica-a:1")
|
||||
pending = _pending("interactions/bg-1")
|
||||
|
||||
await store.register(pending)
|
||||
|
||||
assert await store.pending("interactions/bg-1") == pending
|
||||
assert await store.unclaimed() == (pending,)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_claim_is_won_by_exactly_one_settler():
|
||||
table = _FakeSettlementTable()
|
||||
replica_a = PrismaBackgroundSettlementStore(table=table, claimed_by="replica-a:1")
|
||||
replica_b = PrismaBackgroundSettlementStore(table=table, claimed_by="replica-b:1")
|
||||
await replica_a.register(_pending("interactions/bg-1"))
|
||||
|
||||
assert await replica_b.claim("interactions/bg-1") is True
|
||||
assert await replica_a.claim("interactions/bg-1") is False
|
||||
assert await replica_a.is_claimed("interactions/bg-1") is True
|
||||
assert await replica_a.pending("interactions/bg-1") is None
|
||||
assert table.rows["interactions/bg-1"].claimed_by == "replica-b:1"
|
||||
|
||||
|
||||
class _MissingSettlementTable:
|
||||
"""Prisma's per-model actions against a database whose migration for this table was held back."""
|
||||
|
||||
async def create(self, *, data):
|
||||
raise self._missing()
|
||||
|
||||
async def find_unique(self, *, where):
|
||||
raise self._missing()
|
||||
|
||||
async def find_many(self, *, where):
|
||||
raise self._missing()
|
||||
|
||||
async def update_many(self, *, data, where):
|
||||
raise self._missing()
|
||||
|
||||
def _missing(self):
|
||||
from prisma.errors import TableNotFoundError
|
||||
|
||||
return TableNotFoundError(
|
||||
{
|
||||
"user_facing_error": {
|
||||
"error_code": "P2021",
|
||||
"meta": {"table": "public.LiteLLM_BackgroundInteractionSettlement"},
|
||||
"message": "The table does not exist in the current database.",
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_missing_table_holds_no_rows_and_takes_no_registration():
|
||||
from prisma.errors import TableNotFoundError
|
||||
|
||||
store = PrismaBackgroundSettlementStore(table=_MissingSettlementTable(), claimed_by="replica-a:1")
|
||||
|
||||
with pytest.raises(TableNotFoundError):
|
||||
await store.register(_pending("interactions/bg-1"))
|
||||
assert await store.pending("interactions/bg-1") is None
|
||||
assert await store.is_claimed("interactions/bg-1") is False
|
||||
assert await store.claim("interactions/bg-1") is False
|
||||
with pytest.raises(TableNotFoundError):
|
||||
await store.unclaimed()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unclaimed_skips_claimed_and_unreadable_rows():
|
||||
table = _FakeSettlementTable(
|
||||
rows=(
|
||||
_stored_row("interactions/bg-orphaned"),
|
||||
_stored_row("interactions/bg-settled", claimed=True),
|
||||
_stored_row("interactions/bg-from-the-future", create_context={"schema": "unknown"}),
|
||||
)
|
||||
)
|
||||
store = PrismaBackgroundSettlementStore(table=table, claimed_by="replica-b:1")
|
||||
|
||||
unclaimed = await store.unclaimed()
|
||||
|
||||
assert [row.interaction_id for row in unclaimed] == ["interactions/bg-orphaned"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_record_outcome_keeps_the_audit_trail_and_drops_the_stored_request_context():
|
||||
table = _FakeSettlementTable()
|
||||
store = PrismaBackgroundSettlementStore(table=table, claimed_by="replica-a:1")
|
||||
await store.register(_pending("interactions/bg-1"))
|
||||
assert await store.claim("interactions/bg-1")
|
||||
assert table.rows["interactions/bg-1"].create_context
|
||||
|
||||
await store.record_outcome("interactions/bg-1", "billed")
|
||||
|
||||
row = table.rows["interactions/bg-1"]
|
||||
assert row.outcome == "billed"
|
||||
assert row.settled_at is not None
|
||||
assert row.claimed_at <= row.settled_at
|
||||
assert row.create_context == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_configure_installs_the_store_and_resumes_the_orphaned_rows():
|
||||
table = _FakeSettlementTable(
|
||||
rows=(_stored_row("interactions/bg-orphaned"), _stored_row("interactions/bg-settled", claimed=True))
|
||||
)
|
||||
fetch, captured = _capturing_fetch()
|
||||
previous_store = bg._STORE.store
|
||||
try:
|
||||
resumed = await configure_background_interaction_settlement(
|
||||
table=table, claimed_by="replica-b:1", fetch_interaction=fetch, schedule=FAST_SCHEDULE
|
||||
)
|
||||
|
||||
assert len(resumed) == 1
|
||||
assert await asyncio.wait_for(resumed[0], timeout=5) == "billed"
|
||||
assert [context.interaction_id for context in captured] == ["interactions/bg-orphaned"]
|
||||
assert table.rows["interactions/bg-orphaned"].claimed_by == "replica-b:1"
|
||||
assert table.rows["interactions/bg-orphaned"].outcome == "billed"
|
||||
|
||||
await table.create(
|
||||
data={
|
||||
"interaction_id": "interactions/bg-created-elsewhere",
|
||||
"custom_llm_provider": "gemini",
|
||||
"create_context": _JsonLike(_create_context(_logging_obj(), "gemini").model_dump(mode="json")),
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
}
|
||||
)
|
||||
outcome = await maybe_settle_background_interaction_before_delete(
|
||||
interaction_id="interactions/bg-created-elsewhere", delete_kwargs={}, fetch_interaction=fetch
|
||||
)
|
||||
|
||||
assert outcome == "billed"
|
||||
assert table.rows["interactions/bg-created-elsewhere"].claimed_by == "replica-b:1"
|
||||
finally:
|
||||
configure_background_settlement_store(previous_store)
|
||||
|
||||
|
||||
class _PrismaClientWithoutSettlementTable:
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_install_keeps_booting_when_the_settlement_table_is_unreachable():
|
||||
previous_store = bg._STORE.store
|
||||
|
||||
await install_background_interaction_settlement(_PrismaClientWithoutSettlementTable())
|
||||
|
||||
assert bg._STORE.store is previous_store
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _JsonLike:
|
||||
data: object
|
||||
Loading…
Add table
Reference in a new issue