From bda42a2bda506ecb40b7e14ccd06a6a5e43966fd Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 03:50:25 -0700 Subject: [PATCH] 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. --- .../migration.sql | 16 + .../litellm_proxy_extras/schema.prisma | 16 + .../interactions/background_cost_polling.py | 504 +++++++++++++++--- litellm/interactions/main.py | 4 +- litellm/proxy/proxy_server.py | 4 + litellm/proxy/schema.prisma | 16 + .../background_interaction_settlement.py | 175 ++++++ litellm/repositories/table_repositories.py | 6 + schema.prisma | 16 + .../test_background_cost_polling.py | 234 +++++++- .../test_background_interaction_settlement.py | 238 +++++++++ 11 files changed, 1130 insertions(+), 99 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260915000000_add_background_interaction_settlement/migration.sql create mode 100644 litellm/proxy/spend_tracking/background_interaction_settlement.py create mode 100644 tests/test_litellm/proxy/spend_tracking/test_background_interaction_settlement.py diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260915000000_add_background_interaction_settlement/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260915000000_add_background_interaction_settlement/migration.sql new file mode 100644 index 00000000000..94d5e98f2a7 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260915000000_add_background_interaction_settlement/migration.sql @@ -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"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index dd7967aafe3..333d6ebf650 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1672,3 +1672,19 @@ model LiteLLM_WorkflowMessage { @@unique([run_id, sequence_number]) @@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") +} diff --git a/litellm/interactions/background_cost_polling.py b/litellm/interactions/background_cost_polling.py index 51325354e7d..5fc2a443ce0 100644 --- a/litellm/interactions/background_cost_polling.py +++ b/litellm/interactions/background_cost_polling.py @@ -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 + 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,13 @@ 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) 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( @@ -79,53 +187,161 @@ async def _fetch_interaction(context: BackgroundInteractionPollContext) -> Inter custom_llm_provider=context.custom_llm_provider, api_key=context.api_key, api_base=context.api_base, - **{ - "no-log": True - }, # mutable-ok: "no-log" is not a valid identifier, so it can only be passed through a mapping + **{"no-log": True}, # mutable-ok: "no-log" is not an identifier, so it only passes through a mapping ) 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", + "requester_custom_headers", + "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) -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]: + return MappingProxyType({key: parsed for key, value in items for parsed in _json_value(value)}) + + +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 = get_litellm_metadata_from_kwargs(kwargs=logging_obj.model_call_details) + model: Final = logging_obj.model_call_details.get("model") + return BackgroundInteractionCreateContext( + model=model if isinstance(model, str) else logging_obj.model, + call_type=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 logging_obj.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, + 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 = { # mutable-ok: Logging scrubs and merges the params it is handed in place + "metadata": dict(create_context.metadata), # mutable-ok: Logging pops keys from the metadata it is handed + **create_context.custom_pricing, + } + logging_obj.update_environment_variables( + litellm_params=litellm_params, + optional_params={}, # mutable-ok: Logging stores the optional params it is handed and updates them in place + 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_seen_status: str | None = None # rebind-ok: the give-up log names the status the poll last saw 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 @@ -138,15 +354,13 @@ async def poll_and_log_background_interaction_cost( last_seen_status = response.status 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 (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_seen_status is not None and last_seen_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 " @@ -163,6 +377,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: @@ -175,8 +398,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 @@ -236,24 +459,53 @@ 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 + + +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 poll settles from this process + verbose_logger.exception( + "Could not durably register background interaction %s; only this process can settle it", + pending.interaction_id, + ) + fallback: Final = InMemoryBackgroundSettlementStore() + await fallback.register(pending) + return fallback + 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: @@ -262,52 +514,134 @@ 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 # unfetchable pre-delete state settles by releasing the reservation + verbose_logger.debug( + "Could not fetch background interaction %s before delete, releasing its reservation: %s", + context.interaction_id, + e, + ) + return None + + +async def _settle_before_delete( + context: BackgroundInteractionPollContext, fetch_interaction: FetchInteraction +) -> SettlementOutcome | None: + response: Final = await _fetch_before_delete(context, fetch_interaction) + 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: + return await _settle_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, + ) + return await _settle_before_delete(context, fetch_interaction) + + +async def _unclaimed(store: BackgroundSettlementStore) -> Sequence[PendingBackgroundInteraction]: try: - response = await fetch_interaction(context) - except Exception as e: # noqa: BLE001 # unfetchable pre-delete state settles by releasing the reservation - verbose_logger.debug( - "Could not fetch background interaction %s before delete, releasing its reservation: %s", - 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) + 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, + ) + + +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 + ) diff --git a/litellm/interactions/main.py b/litellm/interactions/main.py index 8a33e9b39c5..209acd5e0f9 100644 --- a/litellm/interactions/main.py +++ b/litellm/interactions/main.py @@ -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, @@ -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, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f81a3166a28..cce31b7ddc7 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -661,6 +661,9 @@ from litellm.proxy.route_llm_request import route_request from litellm.proxy.route_priority import hot_routes_first from litellm.proxy.search_endpoints.endpoints import router as search_router from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager +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 from litellm.proxy.spend_tracking.spend_counter_batch import ( PendingSpendIncrement, @@ -1221,6 +1224,7 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: 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 diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index dd7967aafe3..333d6ebf650 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1672,3 +1672,19 @@ model LiteLLM_WorkflowMessage { @@unique([run_id, sequence_number]) @@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") +} diff --git a/litellm/proxy/spend_tracking/background_interaction_settlement.py b/litellm/proxy/spend_tracking/background_interaction_settlement.py new file mode 100644 index 00000000000..5d8a0bdfa3f --- /dev/null +++ b/litellm/proxy/spend_tracking/background_interaction_settlement.py @@ -0,0 +1,175 @@ +import asyncio +import os +import socket +from collections.abc import Awaitable, Sequence +from dataclasses import dataclass +from datetime import datetime, timezone +from typing import TYPE_CHECKING, Final, Protocol + +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] + + +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 + + +def _pending_rows(rows: Sequence[_SettlementRow]) -> tuple[PendingBackgroundInteraction, ...]: + return tuple(pending for row in rows for pending in _pending_row(row)) + + +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: + from prisma import Json # noqa: PLC0415 # local import: prisma may be ungenerated at module load in some tools + + 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.table.find_unique(where=_RowKey(interaction_id=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.table.find_unique(where=_RowKey(interaction_id=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 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), + ) + return claimed_rows == 1 + + 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), + 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: + await configure_background_interaction_settlement( + table=BackgroundInteractionSettlementRepository(prisma_client).table, + claimed_by=f"{socket.gethostname()}:{os.getpid()}", + ) diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py index 1ad7a735d96..ceaae2715fb 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -246,3 +246,9 @@ class AuditLogRepository(PrismaTableRepository["prisma_models.LiteLLM_AuditLog"] class AdaptiveRouterSessionRepository(PrismaTableRepository["prisma_models.LiteLLM_AdaptiveRouterSession"]): table_name = "litellm_adaptiveroutersession" + + +class BackgroundInteractionSettlementRepository( + PrismaTableRepository["prisma_models.LiteLLM_BackgroundInteractionSettlement"] +): + table_name = "litellm_backgroundinteractionsettlement" diff --git a/schema.prisma b/schema.prisma index dd7967aafe3..333d6ebf650 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1672,3 +1672,19 @@ model LiteLLM_WorkflowMessage { @@unique([run_id, sequence_number]) @@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") +} diff --git a/tests/test_litellm/interactions/test_background_cost_polling.py b/tests/test_litellm/interactions/test_background_cost_polling.py index 97f09de1b52..15943d1e695 100644 --- a/tests/test_litellm/interactions/test_background_cost_polling.py +++ b/tests/test_litellm/interactions/test_background_cost_polling.py @@ -1,18 +1,24 @@ 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, 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 +69,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 +81,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 +257,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) @@ -271,21 +283,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 +313,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 +330,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 +349,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 +367,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 +387,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 +401,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 +419,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 +429,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 +448,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 +464,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 +504,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 +568,188 @@ 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_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 == [] + + +@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 _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") + + +@pytest.mark.asyncio +async def test_create_still_settles_on_its_own_replica_when_the_store_is_down(): + 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 == "billed" + assert len(calls) == 1 + assert logging_obj.model_call_details["response_cost"] > 0 + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task diff --git a/tests/test_litellm/proxy/spend_tracking/test_background_interaction_settlement.py b/tests/test_litellm/proxy/spend_tracking/test_background_interaction_settlement.py new file mode 100644 index 00000000000..38963f58ad2 --- /dev/null +++ b/tests/test_litellm/proxy/spend_tracking/test_background_interaction_settlement.py @@ -0,0 +1,238 @@ +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, + 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, 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" + + +@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_on_the_row(): + 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") + + 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 + + +@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) + + +@dataclass(frozen=True) +class _JsonLike: + data: object