Merge remote-tracking branch 'origin/main' into litellm_anthropic_wif_backend

# Conflicts:
#	litellm/llms/openai/workload_identity.py
This commit is contained in:
mateo-berri 2026-10-02 21:42:27 -07:00
commit c23ca30905
66 changed files with 5501 additions and 226 deletions

View file

@ -3,7 +3,7 @@
import base64
import json
from collections.abc import Mapping, Sequence
from collections.abc import Iterator, Mapping, Sequence
from types import MappingProxyType
from typing import (
TYPE_CHECKING,
@ -54,6 +54,7 @@ from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.openai_files_endpoints.common_utils import (
BATCH_CREATE_HIDDEN_PARAM,
FILE_LIST_CONTINUATION_CHUNK_SIZE,
ManagedFileIdResolver,
_is_base64_encoded_unified_file_id,
apply_unified_file_ids,
decode_model_from_file_id,
@ -144,6 +145,7 @@ def _parse_managed_file_object(raw_file_object: object, unified_file_id: str) ->
class _ManagedFileRow(Protocol):
unified_file_id: str
file_object: OpenAIFileObject
flat_model_file_ids: Sequence[str]
storage_backend: Optional[str]
storage_url: Optional[str]
created_by: Optional[str]
@ -201,6 +203,16 @@ def _managed_file_table(prisma_client: PrismaClient) -> _ManagedFileTableActions
return prisma_client.db.litellm_managedfiletable
def _iter_provider_file_id_pairs(
rows: Sequence[_ManagedFileRow],
requested_provider_file_ids: frozenset[str],
) -> Iterator[tuple[str, str]]:
for row in rows:
for provider_file_id in row.flat_model_file_ids:
if provider_file_id in requested_provider_file_ids:
yield provider_file_id, row.unified_file_id
def _managed_object_table(prisma_client: PrismaClient) -> _ManagedObjectTableActions:
return prisma_client.db.litellm_managedobjecttable
@ -710,6 +722,39 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
return None
return batch_obj
async def get_unified_file_ids_for_provider_file_ids(
self,
provider_file_ids: Sequence[str],
user_api_key_dict: UserAPIKeyAuth,
) -> Mapping[str, str]:
if not provider_file_ids:
return MappingProxyType({})
unique_provider_file_ids: Final = tuple(dict.fromkeys(provider_file_ids))
owner_filter: Final = build_owner_filter(user_api_key_dict)
if owner_filter is None:
return MappingProxyType({})
provider_file_ids_list: Final = [ # mutable-ok: Prisma hasSome requires a list
provider_file_id for provider_file_id in unique_provider_file_ids
]
rows: Final = await _managed_file_table(self.prisma_client).find_many(
where={ # mutable-ok: Prisma requires a plain dictionary for where
**owner_filter,
"flat_model_file_ids": { # mutable-ok: Prisma requires a plain filter dictionary
"hasSome": provider_file_ids_list,
},
}
)
return MappingProxyType(
dict(
_iter_provider_file_id_pairs(
rows,
frozenset(unique_provider_file_ids),
)
)
)
async def get_user_created_file_ids(
self, user_api_key_dict: UserAPIKeyAuth, model_object_ids: List[str]
) -> List[OpenAIFileObject]:

View file

@ -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");

View file

@ -0,0 +1,12 @@
-- CreateIndex (CONCURRENTLY)
--
-- Disclaimer:
-- - CREATE INDEX CONCURRENTLY cannot run inside a transaction. This migration must stay a
-- single statement so Prisma Migrate on PostgreSQL can apply it outside a transaction.
-- - Builds are slower and use more I/O than a blocking CREATE INDEX; if the build is
-- interrupted, Postgres may leave an INVALID index that must be dropped and recreated.
-- - Do not edit this file after it has been applied to any database: Prisma checksums
-- migrations; add a new migration instead.
-- - Requires PostgreSQL that supports CONCURRENTLY with IF NOT EXISTS (use a new migration
-- without IF NOT EXISTS if you must support older versions).
CREATE INDEX CONCURRENTLY IF NOT EXISTS "LiteLLM_ManagedFileTable_flat_model_file_ids_idx" ON "LiteLLM_ManagedFileTable" USING GIN ("flat_model_file_ids");

View file

@ -1144,6 +1144,7 @@ model LiteLLM_ManagedFileTable {
updated_by String?
@@index([unified_file_id])
@@index([flat_model_file_ids], type: Gin)
@@index([team_id, created_at(sort: Desc)])
}
@ -1916,6 +1917,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)

View file

@ -1647,6 +1647,9 @@ if TYPE_CHECKING:
from .llms.jina_ai.rerank.transformation import (
JinaAIRerankConfig as JinaAIRerankConfig,
)
from .llms.scaleway.rerank.transformation import (
ScalewayRerankConfig as ScalewayRerankConfig,
)
from .llms.deepinfra.rerank.transformation import (
DeepinfraRerankConfig as DeepinfraRerankConfig,
)

View file

@ -151,6 +151,7 @@ LLM_CONFIG_NAMES: Final = (
"AzureAIRerankConfig",
"InfinityRerankConfig",
"JinaAIRerankConfig",
"ScalewayRerankConfig",
"DeepinfraRerankConfig",
"HostedVLLMRerankConfig",
"NvidiaNimRerankConfig",
@ -688,6 +689,7 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
"InfinityRerankConfig",
),
"JinaAIRerankConfig": (".llms.jina_ai.rerank.transformation", "JinaAIRerankConfig"),
"ScalewayRerankConfig": (".llms.scaleway.rerank.transformation", "ScalewayRerankConfig"),
"DeepinfraRerankConfig": (
".llms.deepinfra.rerank.transformation",
"DeepinfraRerankConfig",

View file

@ -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
)

View file

@ -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,

View file

@ -1930,6 +1930,10 @@ class CustomStreamWrapper:
else:
self.sent_last_chunk = True
processed_chunk: Final = self.finish_reason_handler()
# The logged response is built from self.chunks; keep a finish_reason the provider sent on its
# last content chunk (stripped there), but never add the synthetic "stop" used when it sent none.
if self.received_finish_reason is not None or self.intermittent_finish_reason is not None:
self.chunks.append(processed_chunk)
if self.stream_options is None: # add usage as hidden param
usage = calculate_total_usage(chunks=self.chunks)
processed_chunk._hidden_params["usage"] = usage
@ -2194,6 +2198,10 @@ class CustomStreamWrapper:
else:
self.sent_last_chunk = True
processed_chunk: Final = self.finish_reason_handler()
# The logged response is built from self.chunks; keep a finish_reason the provider sent on its
# last content chunk (stripped there), but never add the synthetic "stop" used when it sent none.
if self.received_finish_reason is not None or self.intermittent_finish_reason is not None:
self.chunks.append(processed_chunk)
if self.stream_options is None:
usage: Final = calculate_total_usage(chunks=self.chunks)
processed_chunk._hidden_params["usage"] = usage # pyright: ignore[reportPrivateUsage] # sync parity

View file

@ -104,6 +104,7 @@ from ..common_utils import (
BedrockModelInfo,
bedrock_converse_supports_parallel_tool_use_config,
bedrock_model_accepts_cache_points,
bedrock_reasoning_effort_disabled,
get_anthropic_beta_from_headers,
get_bedrock_tool_name,
is_bedrock_application_inference_profile_arn,
@ -1135,6 +1136,25 @@ class AmazonConverseConfig(BaseConfig):
"Dropping unsupported `reasoning_effort` param for Bedrock model=%s; it always reasons and rejects it.",
model,
)
elif (
param == "reasoning_effort"
and isinstance(value, str)
and self._is_openai_gpt_reasoning_model(model)
and bedrock_reasoning_effort_disabled(model=model, effort=value)
):
if not (litellm.drop_params or drop_params):
raise litellm.utils.UnsupportedParamsError(
message=(
f"{model} does not support reasoning_effort={value}. "
"To drop unsupported params, set `litellm.drop_params = True`."
),
status_code=400,
)
verbose_logger.debug(
"Dropping unsupported `reasoning_effort=%s` for Bedrock model=%s.",
value,
model,
)
elif param == "reasoning_effort" and isinstance(value, str):
self._handle_reasoning_effort_parameter(
model=model, reasoning_effort=value, optional_params=optional_params

View file

@ -73,7 +73,7 @@ class AmazonMantleConfig(AmazonAnthropicClaudeConfig):
)
project_id: Final = litellm_params.get("aws_bedrock_project_id")
if project_id:
headers["anthropic-workspace"] = project_id
headers["anthropic-workspace-id"] = project_id
return headers
def transform_request(

View file

@ -1017,6 +1017,14 @@ def _mantle_api_base_from_env() -> str | None:
return next((base[: -len(suffix)] for suffix in _MANTLE_OPENAI_BASE_SUFFIXES if base.endswith(suffix)), base)
def bedrock_reasoning_effort_disabled(model: str, effort: str) -> bool:
from litellm.utils import is_explicitly_disabled_factory
return is_explicitly_disabled_factory(
model=model, custom_llm_provider="bedrock_converse", key=f"supports_{effort}_reasoning_effort"
)
def bedrock_supports_openai_responses(model: str | None, model_cost: Mapping[str, object]) -> bool:
"""Whether a Bedrock model is served by bedrock-runtime's OpenAI Responses surface.

View file

@ -104,7 +104,7 @@ class AmazonMantleMessagesConfig(AmazonAnthropicClaudeMessagesConfig):
{
name: value
for name, value in (
("anthropic-workspace", project_id),
("anthropic-workspace-id", project_id),
("anthropic-version", None if has_version else DEFAULT_ANTHROPIC_API_VERSION),
)
if value

View file

@ -52,6 +52,7 @@ from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.common_utils import (
BEDROCK_CHAT_COMPLETIONS_ROUTE_PREFIX,
BedrockError,
bedrock_reasoning_effort_disabled,
bedrock_supports_openai_responses,
)
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
@ -154,6 +155,29 @@ def inline_remote_image_urls(
return items # pyright: ignore[reportReturnType] # items keep the caller's input union
def _without_disabled_reasoning_effort(
params: Mapping[str, object], model: str, drop_params: bool
) -> dict[str, object]: # mutable-ok: becomes the map_openai_params return value
reasoning: Final = params.get("reasoning")
effort: Final = reasoning.get("effort") if isinstance(reasoning, Mapping) else None
if not isinstance(reasoning, Mapping) or not isinstance(effort, str):
return dict(params)
if not bedrock_reasoning_effort_disabled(model=model, effort=effort):
return dict(params)
if not (drop_params or litellm.drop_params):
raise litellm.UnsupportedParamsError(
message=(
f"{model} does not support reasoning.effort={effort}. "
"To drop unsupported params, set `litellm.drop_params = True`."
),
status_code=400,
)
verbose_logger.debug("Dropping unsupported `reasoning.effort=%s` for Bedrock model=%s.", effort, model)
rest: Final = {key: value for key, value in reasoning.items() if key != "effort"}
without_reasoning: Final = {key: value for key, value in params.items() if key != "reasoning"}
return {**without_reasoning, "reasoning": rest} if rest else without_reasoning
class BedrockOpenAIResponsesConfig(BaseAWSLLM, OpenAIResponsesAPIConfig):
"""Responses API config for the OpenAI models on the bedrock-runtime endpoint."""
@ -270,7 +294,8 @@ class BedrockOpenAIResponsesConfig(BaseAWSLLM, OpenAIResponsesAPIConfig):
"Bedrock Runtime Responses API: dropping unsupported parameter(s) %s that the endpoint rejects.",
unsupported,
)
params: Final = {key: value for key, value in mapped.items() if key not in unsupported}
supported: Final[dict[str, object]] = {key: value for key, value in mapped.items() if key not in unsupported}
params: Final = _without_disabled_reasoning_effort(supported, model, drop_params)
tools: Final = params.get("tools")
if not isinstance(tools, list):
return params

View file

@ -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:

View file

@ -10,6 +10,8 @@ from litellm.types.utils import CallTypes
guardrail_translation_mappings: Final = {
CallTypes.image_generation: OpenAIImageGenerationHandler,
CallTypes.aimage_generation: OpenAIImageGenerationHandler,
CallTypes.image_edit: OpenAIImageGenerationHandler,
CallTypes.aimage_edit: OpenAIImageGenerationHandler,
}
__all__ = ["OpenAIImageGenerationHandler", "guardrail_translation_mappings"]

View file

@ -74,6 +74,13 @@ def get_workload_identity_bearer_token(config: OpenAIWorkloadIdentityConfig) ->
return _workload_identity_auth(config).get_token()
async def get_workload_identity_bearer_token_for_api_base(api_base: str) -> str | None:
config: Final = resolve_openai_workload_identity_config(api_key=None, api_base=api_base)
if config is None:
return None
return await _workload_identity_auth(config).get_token_async()
def _config_value(litellm_params: Mapping[str, object] | None, param_key: str, env_name: str) -> str | None:
param_value: Final = litellm_params.get(param_key) if litellm_params is not None else None
if isinstance(param_value, str) and param_value:

View file

@ -0,0 +1,51 @@
"""
Support for Scaleway's `/v1/rerank` endpoint.
The request and response match Jina AI's, so this reuses that config.
API reference: https://www.scaleway.com/en/developers/api/generative-apis/#path-rerank-create-a-reranking
"""
from collections.abc import Mapping
from typing import Final
from litellm.llms.jina_ai.rerank.transformation import JinaAIRerankConfig
from litellm.secret_managers.main import get_secret_str
SCALEWAY_API_BASE: Final = "https://api.scaleway.ai/v1"
class ScalewayRerankConfig(JinaAIRerankConfig):
def get_supported_cohere_rerank_params(self, model: str) -> list[str]: # mutable-ok: BaseRerankConfig contract
return ["query", "top_n", "documents"]
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: Mapping[str, object] | None = None,
) -> str:
base: Final = SCALEWAY_API_BASE if api_base is None else api_base.rstrip("/")
return f"{base}/rerank"
def validate_environment(
self,
headers: Mapping[str, str],
model: str,
api_key: str | None = None,
optional_params: Mapping[str, object] | None = None,
litellm_params: Mapping[str, object] | None = None,
) -> dict[str, str]: # mutable-ok: BaseRerankConfig contract
key: Final = api_key or get_secret_str("SCW_SECRET_KEY")
if not key:
raise ValueError(
"Scaleway API key not found. Pass `api_key=...` or set the SCW_SECRET_KEY environment variable."
)
provider_headers: Final = {
"accept": "application/json",
"content-type": "application/json",
"authorization": f"Bearer {key}",
}
# Header names are case-insensitive, so match on the lowercase name.
caller_headers: Final = {name: value for name, value in headers.items() if name.lower() not in provider_headers}
return {**caller_headers, **provider_headers}

View file

@ -266,7 +266,8 @@ def build_autorouter_turn_transaction(
the payload's own usage record through the savings owner, never handed in beside it.
The baseline the turn's saved_spend was priced against travels with the turn, so the
row can name the counterfactual for the money it holds even after the router is
reconfigured or removed.
reconfigured or removed. A request with no session id still owns its router-day money,
so it becomes a turn with an empty session id that writes the day row and no session row.
"""
if payload.get("status") != "success":
return None
@ -278,9 +279,9 @@ def build_autorouter_turn_transaction(
router_name: Final = routing_decision.get("router_model_name") or payload.get("model_group")
api_key: Final = payload.get("api_key") or ""
user_id: Final = payload.get("user") or ""
session_id: Final = payload.get("session_id")
session_id: Final = payload.get("session_id") or ""
model: Final = payload.get("model")
if not (isinstance(router_name, str) and router_name and (api_key or user_id) and session_id and model):
if not (isinstance(router_name, str) and router_name and (api_key or user_id) and model):
return None
turn_at: Final = _turn_time_utc(str(payload.get("startTime") or ""))
if turn_at is None:
@ -379,7 +380,7 @@ SELECT
{_p("classifier_cost")}::float8, 1, {_TIER_DELTA}, {_BASELINE_DELTA},
{_p("savings_estimated_turns")}::int, {_p("savings_estimated_actual_spend")}::float8,
{_p("savings_estimated_saved_spend")}::float8, {_ESTIMATED_BASELINE_DELTA}
WHERE {required_identity}::text <> ''
WHERE {required_identity}::text <> '' AND {_p("session_id")}::text <> ''
ON CONFLICT ({user_column}api_key, session_id, router_name) DO UPDATE SET
turns = t.turns + 1,
total_tokens = t.total_tokens + EXCLUDED.total_tokens,

View file

@ -7,6 +7,7 @@ Module responsible for
import asyncio
import copy
import dataclasses
import json
import os
import random
@ -550,6 +551,13 @@ class DBSpendUpdateWriter:
):
return False
# The auto-router router-day rollup is an aggregate like the daily spend tables, so it is
# written whether or not per-request spend logs are kept; per-session rows are not.
await self._enqueue_autorouter_turn_transaction(
payload=payload,
prisma_client=prisma_client,
spend_logs_kept=disable_spend_logs is False,
)
if disable_spend_logs is False:
await self._enqueue_tool_usage_transaction(
payload=payload,
@ -557,10 +565,6 @@ class DBSpendUpdateWriter:
prisma_client=prisma_client,
kwargs=kwargs,
)
await self._enqueue_autorouter_turn_transaction(
payload=payload,
prisma_client=prisma_client,
)
else:
verbose_proxy_logger.debug(
"disable_spend_logs=True. Skipping writing spend logs to db. Other spend updates - Key/User/Team table will still occur."
@ -747,6 +751,7 @@ class DBSpendUpdateWriter:
self,
payload: SpendLogsPayload,
prisma_client: "PrismaClient | None",
spend_logs_kept: bool = True,
) -> None:
try:
if prisma_client is None:
@ -787,14 +792,21 @@ class DBSpendUpdateWriter:
saved_spend=savings_spend.autorouter,
)
try:
if await self._enqueue_baseline_accounting(payload, metadata, transaction, prisma_client):
# A baseline observation publishes only once its spend log exists, so without spend logs
# it could never publish; the plain turn still carries this request's recorded savings.
if spend_logs_kept and await self._enqueue_baseline_accounting(
payload, metadata, transaction, prisma_client
):
return
except Exception: # noqa: BLE001 # optional baseline capture must preserve the original actual-spend rollup
verbose_proxy_logger.warning("Auto-router baseline observation was unavailable; actual turn retained")
if transaction is None:
return
# Without spend logs only the router-day aggregate is kept: an empty session id makes the
# session upserts skip the row, so no per-session record is stored.
kept: Final = transaction if spend_logs_kept else dataclasses.replace(transaction, session_id="")
async with prisma_client._autorouter_turn_transactions_lock:
prisma_client.autorouter_turn_transactions.append(transaction)
prisma_client.autorouter_turn_transactions.append(kept)
except Exception as e: # noqa: BLE001 # a metrics enqueue must never fail the spend write
verbose_proxy_logger.debug("_enqueue_autorouter_turn_transaction error (non-blocking): %s", e)

View file

@ -1,7 +1,7 @@
import base64
import mimetypes
import re
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from dataclasses import dataclass, field
from types import MappingProxyType
from typing import (
@ -95,6 +95,15 @@ class ManagedResourceAccessChecker(Protocol):
) -> bool: ...
@runtime_checkable
class ManagedFileIdResolver(Protocol):
async def get_unified_file_ids_for_provider_file_ids(
self,
provider_file_ids: Sequence[str],
user_api_key_dict: "UserAPIKeyAuth",
) -> Mapping[str, str]: ...
def _is_base64_encoded_unified_file_id(b64_uid: str) -> str | Literal[False]:
# Ensure b64_uid is a string and not a mock object
if not isinstance(b64_uid, str):

View file

@ -22,6 +22,7 @@ from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
import httpx
import openai
from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket
from fastapi.responses import StreamingResponse
from pydantic import ConfigDict, TypeAdapter
@ -59,6 +60,8 @@ from litellm.llms.deepgram.common_utils import (
)
from litellm.llms.fal_ai.cost_calculator import fal_ai_passthrough_cost, fal_ai_queue_base
from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path
from litellm.llms.openai.common_utils import OpenAIError as LiteLLMOpenAIError
from litellm.llms.openai.workload_identity import get_workload_identity_bearer_token_for_api_base
from litellm.llms.oss_decision import OssDecisionProvider, oss_connection, validate_oss_request
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
@ -101,7 +104,7 @@ from litellm.proxy.vector_store_endpoints.utils import (
get_litellm_managed_vector_store,
is_allowed_to_call_vector_store_endpoint,
)
from litellm.secret_managers.main import get_secret_str, str_to_bool
from litellm.secret_managers.main import get_secret_str, normalize_nonempty_secret_str, str_to_bool
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY,
@ -2991,6 +2994,21 @@ async def vertex_proxy_route(
)
_OPENAI_WS_TOKEN_EXCHANGE_FAILED_REASON: Final = "OpenAI workload identity token exchange failed"
async def _openai_passthrough_credential(base_target_url: str) -> str | None:
static_api_key: Final = normalize_nonempty_secret_str(
passthrough_endpoint_router.get_credentials(
custom_llm_provider=litellm.LlmProviders.OPENAI.value,
region_name=None,
)
)
if static_api_key is not None:
return static_api_key
return await get_workload_identity_bearer_token_for_api_base(base_target_url)
@router.api_route(
"/openai/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
@ -3026,11 +3044,7 @@ async def openai_proxy_route(
[Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)
"""
base_target_url: Final = os.getenv("OPENAI_API_BASE") or "https://api.openai.com/"
# Add or update query parameters
openai_api_key: Final = passthrough_endpoint_router.get_credentials(
custom_llm_provider=litellm.LlmProviders.OPENAI.value,
region_name=None,
)
openai_api_key: Final = await _openai_passthrough_credential(base_target_url)
if openai_api_key is None:
raise Exception("Required 'OPENAI_API_KEY' in environment to make pass-through calls to OpenAI.")
@ -3198,10 +3212,12 @@ async def openai_websocket_proxy_route(
return
base_target_url: Final = os.getenv("OPENAI_API_BASE") or "https://api.openai.com/"
openai_api_key: Final = passthrough_endpoint_router.get_credentials(
custom_llm_provider=litellm.LlmProviders.OPENAI.value,
region_name=None,
)
try:
openai_api_key: Final = await _openai_passthrough_credential(base_target_url)
except (openai.OpenAIError, httpx.HTTPError, LiteLLMOpenAIError):
verbose_proxy_logger.exception("OpenAI workload identity token exchange failed for websocket passthrough")
await websocket.close(code=1011, reason=_OPENAI_WS_TOKEN_EXCHANGE_FAILED_REASON)
return
if openai_api_key is None:
await websocket.close(
code=1011,

View file

@ -188,10 +188,9 @@ _OBJECT_PREFIXES: Final[frozenset[str]] = frozenset({"batch_", "resp_"})
_MAX_BODY_REWRITE_DEPTH: Final = 64
# Caps the distinct raw-provider-id guard lookups issued per request. A raw
# file-id guard is an unindexed array-containment scan over
# LiteLLM_ManagedFileTable (flat_model_file_ids has no index), so a body packed
# with id-shaped strings could otherwise amplify one request into thousands of
# full-table scans. Legitimate callers reference managed IDs (resolved via an
# file-id guard is an array-containment lookup over LiteLLM_ManagedFileTable,
# so a body packed with id-shaped strings could otherwise amplify one request
# into thousands of lookups. Legitimate callers reference managed IDs (resolved via an
# indexed lookup, never the guard), so guarding more raw ids than this only
# happens under abuse; the request is rejected rather than skipping the guard.
_MAX_RAW_ID_GUARD_LOOKUPS: Final = 100

View file

@ -778,6 +778,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,
@ -1396,6 +1399,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

View file

@ -1144,6 +1144,7 @@ model LiteLLM_ManagedFileTable {
updated_by String?
@@index([unified_file_id])
@@index([flat_model_file_ids], type: Gin)
@@index([team_id, created_at(sort: Desc)])
}
@ -1916,6 +1917,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)

View file

@ -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,
)

View file

@ -1,9 +1,13 @@
from typing import TYPE_CHECKING, Final, Optional
import re
from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Optional, cast
from fastapi import APIRouter, Depends, Request, Response
from fastapi.responses import ORJSONResponse
import litellm
from litellm.llms.base_llm.managed_resources.utils import is_base64_encoded_unified_id
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
@ -13,6 +17,7 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
get_custom_llm_provider_from_request_query,
)
from litellm.proxy.openai_files_endpoints.common_utils import (
ManagedFileIdResolver,
authorize_model_for_key,
get_credentials_for_model,
handle_model_based_routing,
@ -24,6 +29,10 @@ from litellm.proxy.vector_store_endpoints.utils import (
is_allowed_to_call_vector_store_files_endpoint,
)
from litellm.types.utils import LlmProviders
from litellm.types.vector_store_files import (
VectorStoreFileListResponse,
VectorStoreFileObject,
)
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
if TYPE_CHECKING:
@ -32,6 +41,93 @@ if TYPE_CHECKING:
router: Final = APIRouter()
def _provider_file_id_from_managed_id(managed_file_id: str | None) -> str | None:
if managed_file_id is None:
return None
decoded_id: Final = is_base64_encoded_unified_id(managed_file_id)
if not decoded_id:
return managed_file_id
match: Final = re.search(r"(?:^|;)llm_output_file_id,([^;]+)", decoded_id)
return match.group(1).strip() if match else managed_file_id
def _with_provider_file_id_cursors(
query_params: Mapping[str, str],
) -> Mapping[str, str | None]:
return MappingProxyType(
{
key: (_provider_file_id_from_managed_id(value) if key in {"after", "before"} else value)
for key, value in query_params.items()
}
)
def _managed_file_id_or_original(
file_id: str | None,
id_map: Mapping[str, str],
) -> str | None:
return id_map.get(file_id, file_id) if file_id is not None else None
def _with_managed_file_id(
file: VectorStoreFileObject,
id_map: Mapping[str, str],
) -> VectorStoreFileObject:
file_id: Final = file.get("id")
if not isinstance(file_id, str) or file_id not in id_map:
return file
managed_file: Final[VectorStoreFileObject] = {**file, "id": id_map[file_id]}
return managed_file
def _with_managed_file_ids(
response: VectorStoreFileListResponse,
id_map: Mapping[str, str],
) -> VectorStoreFileListResponse:
data: Final = response.get("data")
if not data:
return response
first_id: Final = response.get("first_id")
last_id: Final = response.get("last_id")
mapped_data: Final = [_with_managed_file_id(file, id_map) for file in data]
mapped_response: Final[VectorStoreFileListResponse] = {
**response,
"data": mapped_data,
"first_id": _managed_file_id_or_original(first_id, id_map),
"last_id": _managed_file_id_or_original(last_id, id_map),
}
return mapped_response
async def _with_managed_file_list_ids(
response: VectorStoreFileListResponse,
managed_files_obj: object | None,
user_api_key_dict: UserAPIKeyAuth,
) -> VectorStoreFileListResponse:
data: Final = response.get("data")
if not data or not isinstance(managed_files_obj, ManagedFileIdResolver):
return response
provider_file_ids: Final = tuple(
dict.fromkeys(provider_file_id for file in data if isinstance(provider_file_id := file.get("id"), str))
)
id_map: Final = await managed_files_obj.get_unified_file_ids_for_provider_file_ids(
provider_file_ids=provider_file_ids,
user_api_key_dict=user_api_key_dict,
)
round_trippable_id_map: Final = MappingProxyType(
{
provider_file_id: managed_file_id
for provider_file_id, managed_file_id in id_map.items()
if _provider_file_id_from_managed_id(managed_file_id) == provider_file_id
}
)
return _with_managed_file_ids(response, round_trippable_id_map)
async def _update_request_data_with_managed_file_id(
data: dict,
file_id: str,
@ -62,11 +158,8 @@ async def _update_request_data_with_managed_file_id(
Tuple of (updated request data, original_managed_file_id)
- original_managed_file_id is the original file_id if it was managed/encoded, None otherwise
"""
import re
from litellm import verbose_logger
from litellm.llms.base_llm.managed_resources.utils import (
is_base64_encoded_unified_id,
parse_unified_id,
)
from litellm.proxy.openai_files_endpoints.common_utils import (
@ -591,7 +684,7 @@ async def vector_store_file_list(
version,
)
query_params: Final = dict(request.query_params)
query_params: Final = _with_provider_file_id_cursors(request.query_params)
data: dict[str, str | None] = {"vector_store_id": vector_store_id}
data.update(query_params)
data["vector_store_id"] = vector_store_id
@ -628,7 +721,7 @@ async def vector_store_file_list(
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
response: Final[object] = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -646,6 +739,17 @@ async def vector_store_file_list(
user_api_base=user_api_base,
version=version,
)
if not isinstance(response, dict):
return response
managed_files_obj: Final[object | None] = proxy_logging_obj.get_proxy_hook("managed_files")
return await _with_managed_file_list_ids(
response=cast( # cast-ok: [LIT006] this route returns the provider's file-list response shape
VectorStoreFileListResponse,
response,
),
managed_files_obj=managed_files_obj,
user_api_key_dict=user_api_key_dict,
)
except Exception as e: # noqa: BLE001
raise await processor._handle_llm_api_exception(
e=e,

View file

@ -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"

View file

@ -8860,8 +8860,13 @@ class ProviderConfigManager:
return litellm.AzureAIRerankConfig()
elif litellm.LlmProviders.INFINITY == provider:
return litellm.InfinityRerankConfig()
elif litellm.LlmProviders.JINA_AI == provider:
return litellm.JinaAIRerankConfig()
elif provider in (litellm.LlmProviders.JINA_AI, litellm.LlmProviders.SCALEWAY):
# Scaleway's rerank API matches Jina's, so its config extends Jina's.
return (
litellm.ScalewayRerankConfig()
if provider == litellm.LlmProviders.SCALEWAY
else litellm.JinaAIRerankConfig()
)
elif litellm.LlmProviders.HOSTED_VLLM == provider:
return litellm.HostedVLLMRerankConfig()
elif litellm.LlmProviders.HUGGINGFACE == provider:

View file

@ -2397,7 +2397,7 @@
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"rerank": true,
"a2a": true,
"interactions": true
}

View file

@ -1144,6 +1144,7 @@ model LiteLLM_ManagedFileTable {
updated_by String?
@@index([unified_file_id])
@@index([flat_model_file_ids], type: Gin)
@@index([team_id, created_at(sort: Desc)])
}
@ -1916,6 +1917,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)

View file

@ -9,6 +9,7 @@
- {id: guardrail.bedrock.pre_call.blocks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions, messages, responses], source: "guardrail_hooks/bedrock_guardrails.py", rationale: "AWS content guardrail blocks harmful input"}
- {id: guardrail.litellm_content_filter.pre_call.blocks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "test_team_disable_global_guardrail_e2e.py", rationale: "Local content-filter default-on blocks banned keyword pre-call"}
- {id: guardrail.litellm_content_filter.pre_call.blocks_video, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [videos], source: "test_key_guardrail_video_e2e.py", fail_before_fix: proven, rationale: "A content-filter guardrail attached to a key (metadata.guardrails) blocks a banned prompt on POST /v1/videos before the provider is called; before the fix the route's call type was unknown to the unified guardrail hook and the prompt went to the provider unscanned (LIT-6685)"}
- {id: guardrail.litellm_content_filter.pre_call.blocks_image_edit, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [images_edits], source: "test_key_guardrail_image_edit_e2e.py", rationale: "A content-filter guardrail attached to a key (metadata.guardrails) blocks a banned prompt on POST /v1/images/edits before the provider is called; before the fix aimage_edit had no guardrail translation mapping and the prompt went to the provider unscanned"}
- {id: guardrail.litellm_content_filter.pre_call.allows, module: guardrail, tier: P0, hook_point: pre_call, assertions: [allows], exercised_on: [chat_completions], source: "test_team_disable_global_guardrail_e2e.py", rationale: "Team disable_global_guardrails bypasses default-on content filter"}
- {id: guardrail.litellm_content_filter.pre_call.returns_guardrail_information, module: guardrail, tier: P0, hook_point: pre_call, assertions: [allows], exercised_on: [chat_completions], source: "guardrails/test_guardrail_information_response_e2e.py", rationale: "Opt-in chat responses expose successful guardrail execution details"}
- {id: guardrail.litellm_content_filter.apply_endpoint.blocks, module: guardrail, tier: P0, hook_point: apply_endpoint, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_endpoints.py:apply_guardrail", rationale: "POST /guardrails/apply_guardrail blocks banned content for customers that call the apply surface directly"}

View file

@ -9,7 +9,7 @@ from collections.abc import Callable
from dataclasses import dataclass
from typing import Final, Literal
from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, settle_propagation, unique_marker
from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, SLOW_PROVIDER_TIMEOUT_SECONDS, settle_propagation, unique_marker
from e2e_http import NoBody, Result, StreamingResponse, Success, unwrap
from lifecycle import ResourceManager
from models import (
@ -20,6 +20,8 @@ from models import (
ChatMetadata,
ChatResponse,
ChatTool,
ImageEditForm,
ImageGenerationResponse,
KeyGenerateBody,
KeyMetadata,
LiteLLMParamsBody,
@ -364,6 +366,19 @@ class GuardrailsClient:
response_type=VideoCreateResponse,
)
def edit_image(self, key: str, model: str, prompt: str, image: bytes) -> Result[ImageGenerationResponse]:
return self.proxy.transport.upload(
"/v1/images/edits",
headers=self.proxy.transport.bearer(key),
form=ImageEditForm(model=model, prompt=prompt),
filename="image.png",
content=image,
file_content_type="image/png",
file_field="image",
response_type=ImageGenerationResponse,
timeout=SLOW_PROVIDER_TIMEOUT_SECONDS,
)
def chat(
self,
key: str,

View file

@ -0,0 +1,72 @@
from __future__ import annotations
import base64
from typing import Final
import pytest
from e2e_config import unique_marker
from e2e_http import Success, UnknownApiError
from guardrails_client import GuardrailsClient, poll_until_blocked
from lifecycle import ResourceManager
from models import LiteLLMParamsBody
pytestmark = pytest.mark.e2e
CHAT_MODEL: Final = "gemini-2.5-flash"
IMAGE_BACKEND: Final = "openai/gpt-image-2.5-flare"
SOURCE_PNG: Final = base64.b64decode(
"iVBORw0KGgoAAAANSUhEUgAAAEAAAABACAIAAAAlC+aJAAAAS0lEQVR42u3PMQ0AAAwDoPo3"
"3UrYvQQckD4XAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEB"
"AYHLAMpT0sIcNbcEAAAAAElFTkSuQmCC"
)
def _edit_prompt_with(banned_keyword: str) -> str:
return f"Turn this into a watercolor painting of a lighthouse. {banned_keyword}"
def _create_image_model(client: GuardrailsClient, resources: ResourceManager) -> str:
model_name = f"e2e-guard-image-edit-{unique_marker()}"
model_id = client.proxy.create_model(
model_name,
LiteLLMParamsBody(model=IMAGE_BACKEND, api_key="os.environ/OPENAI_API_KEY"),
provider_live=True,
)
resources.defer(lambda: client.proxy.delete_model(model_id))
return model_name
class TestKeyAttachedGuardrailOnImageEdits:
@pytest.mark.covers(
"guardrail.litellm_content_filter.pre_call.blocks_image_edit",
exercised_on=["images_edits"],
)
def test_key_attached_content_filter_blocks_banned_image_edit_prompt(
self, client: GuardrailsClient, resources: ResourceManager
) -> None:
banned = unique_marker()
guardrail_name = f"e2e-image-edit-filter-{banned}"
guardrail_id = client.create_content_filter_guardrail(guardrail_name, banned, default_on=False)
resources.defer(lambda: client.delete_guardrail(guardrail_id))
key = client.create_key_with_guardrails(resources, [guardrail_name])
model = _create_image_model(client, resources)
synced = poll_until_blocked(lambda: client.chat(key, CHAT_MODEL, _edit_prompt_with(banned)))
assert isinstance(synced, UnknownApiError) and synced.status_code == 400, (
f"key guardrail {guardrail_name!r} never synced to the proxy on /chat/completions: {synced}"
)
result = client.edit_image(key, model, _edit_prompt_with(banned), SOURCE_PNG)
match result:
case UnknownApiError(status_code=status, body=body):
assert status == 400, f"expected a 400 guardrail block, got {status}: {body[:300]}"
assert "content blocked" in body.lower() or banned in body, (
f"block response missing content-filter reason: {body[:300]}"
)
case Success():
pytest.fail(
f"key-attached guardrail {guardrail_name!r} was skipped on /v1/images/edits: "
"the banned prompt reached the provider and an edited image came back"
)
case _:
pytest.fail(f"unexpected /v1/images/edits outcome for a banned prompt: {result}")

View file

@ -14,7 +14,7 @@ Reuse the existing canned provider handlers through `_support/upstream.py`. It r
The CircleCI workflow starts its own database and Redis, restricts test-phase egress to its owned services and writes JUnit plus an executed-node manifest. Missing setup, failed cleanup or a selected test with neither a passed call nor a skip fail qualification. Skipped nodes are listed under `skipped` in `execution.json`, so the skip reasons double as the open bug list. Existing GitHub Actions jobs do not own these tests
There is no per-node manifest. The runner fails only when pytest fails, when collection errors, or when a selected file collects zero tests. Older tests still carry `@pytest.mark.covers(...)` decorators; the marker stays registered so they collect, but the IDs are not checked against anything and new tests should not use it. The GitHub Actions coverage census reads the `GROUPS` literal in `run.py` and treats every `tests/integration/<directory>/test_*.py` file in a scheduled group as owned by CircleCI
There is no per-node manifest. A positional argument is a file of the group or a pytest node id inside one (`path::test[param]`), so one cell of a parametrized file can run alone. The runner fails only when pytest fails, when collection errors, or when a selected file collects zero tests. Older tests still carry `@pytest.mark.covers(...)` decorators; the marker stays registered so they collect, but the IDs are not checked against anything and new tests should not use it. The GitHub Actions coverage census reads the `GROUPS` literal in `run.py` and treats every `tests/integration/<directory>/test_*.py` file in a scheduled group as owned by CircleCI
Provider sentinels currently use the controlled server, not live recordings. The provider shard also runs the existing strict replay controls for changed requests, exhausted interactions, leftover interactions and no provider connection. Future recorded scenarios must use that replay-only implementation; missing recordings cannot fall back to a real provider. The observation endpoint is destructive and the current selection runs serially against one owned upstream

View file

@ -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()

View file

@ -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)

View file

@ -0,0 +1,184 @@
import os
import re
import shutil
import subprocess
import sys
import uuid
from collections.abc import Callable
from pathlib import Path
from typing import Final
from urllib.parse import urlsplit
import pytest
from integration._support.client import Gateway, object_value, string_value
from integration._support.database import read_rows, scratch_database
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
REPO_ROOT: Final = Path(__file__).resolve().parents[3]
PRISMA_DIR: Final = REPO_ROOT / "litellm-proxy-extras" / "litellm_proxy_extras"
GIN_MIGRATION: Final = "20261003000000_add_managed_file_flat_ids_gin_index"
INDEX_NAME: Final = "LiteLLM_ManagedFileTable_flat_model_file_ids_idx"
SHIPPED_MIGRATIONS: Final = tuple(sorted(path.name for path in (PRISMA_DIR / "migrations").iterdir() if path.is_dir()))
INDEX_ROW: Final = (
"SELECT i.indexdef, x.indisvalid FROM pg_indexes i "
"JOIN pg_class c ON c.relname = i.indexname JOIN pg_index x ON x.indexrelid = c.oid WHERE i.indexname = %s"
)
APPLIED_MIGRATIONS: Final = (
'SELECT migration_name FROM "_prisma_migrations" '
"WHERE finished_at IS NOT NULL AND rolled_back_at IS NULL AND migration_name <> %s ORDER BY migration_name"
)
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
UPLOAD_FILENAME: Final = re.compile(rb'filename="([^"]+)"')
def _index_rows(database_url: str | None = None) -> list[dict[str, JsonValue]]:
return read_rows(INDEX_ROW, (INDEX_NAME,), database_url=database_url)
def _assert_valid_gin_index(rows: list[dict[str, JsonValue]]) -> None:
assert len(rows) == 1, rows
definition: Final = string_value(rows[0]["indexdef"])
assert "USING gin" in definition, definition
assert '"LiteLLM_ManagedFileTable"' in definition, definition
assert "flat_model_file_ids" in definition, definition
assert rows[0]["indisvalid"] is True, rows
def _applied_migrations(database_url: str) -> tuple[str, ...]:
rows: Final = read_rows(APPLIED_MIGRATIONS, ("",), database_url=database_url)
return tuple(string_value(row["migration_name"]) for row in rows)
def _leg_python_path() -> str:
return os.pathsep.join(
(
str(REPO_ROOT),
str(REPO_ROOT / "litellm-proxy-extras"),
str(REPO_ROOT / "enterprise"),
os.environ.get("PYTHONPATH", ""),
)
)
def _deploy_schema_before(database_url: str, directory: Path, migration: str) -> None:
older: Final = directory / "older-release"
(older / "migrations").mkdir(parents=True)
shutil.copy(PRISMA_DIR / "schema.prisma", older / "schema.prisma")
shutil.copy(PRISMA_DIR / "migrations" / "migration_lock.toml", older / "migrations" / "migration_lock.toml")
for name in (name for name in SHIPPED_MIGRATIONS if name < migration):
shutil.copytree(PRISMA_DIR / "migrations" / name, older / "migrations" / name)
subprocess.run(
[sys.executable, "-I", "-m", "prisma", "migrate", "deploy", "--schema", str(older / "schema.prisma")],
check=True,
capture_output=True,
text=True,
timeout=600,
env={**os.environ, "DATABASE_URL": database_url},
)
def _run_migration_entrypoint(database_url: str) -> subprocess.CompletedProcess[str]:
return subprocess.run(
[sys.executable, "-m", "litellm.proxy.prisma_migration"],
capture_output=True,
text=True,
timeout=600,
cwd=REPO_ROOT,
env={**os.environ, "DATABASE_URL": database_url, "PYTHONPATH": _leg_python_path()},
)
def _provider(store: str, provider_file_id: str) -> Callable[[Request], Reply]:
page: Final[dict[str, JsonValue]] = {
"object": "list",
"data": [
{"id": provider_file_id, "object": "vector_store.file", "vector_store_id": store, "status": "completed"}
],
"first_id": provider_file_id,
"last_id": provider_file_id,
"has_more": False,
}
file_object: Final[dict[str, JsonValue]] = {
"id": provider_file_id,
"object": "file",
"bytes": 6,
"created_at": 1700000000,
"filename": "a.txt",
"purpose": "user_data",
"status": "processed",
}
def respond(request: Request) -> Reply:
path: Final = urlsplit(request.target).path
if request.method == "POST" and path == "/v1/files" and UPLOAD_FILENAME.search(request.body):
return Reply(body=JSON_OBJECT.dump_json(file_object))
if request.method == "GET" and path == f"/v1/vector_stores/{store}/files":
return Reply(body=JSON_OBJECT.dump_json(page))
return Reply(status=404, body=b'{"error": {"message": "unscripted"}}')
return respond
def _listed_ids(gateway: Gateway, store: str, model: str) -> tuple[JsonValue, ...]:
listed: Final = gateway.request("GET", f"/v1/vector_stores/{store}/files", params={"model": model})
assert listed.status_code == 200, listed.text
page: Final = JSON_OBJECT.validate_json(listed.content)
data: Final = page["data"]
assert isinstance(data, list), listed.text
ids: Final = tuple(object_value(entry)["id"] for entry in data)
assert (page["first_id"], page["last_id"]) == (ids[0], ids[-1]), listed.text
return ids
@pytest.mark.timeout(900)
def test_migration_entrypoint_adds_the_gin_index_and_the_upgraded_proxy_maps_managed_ids(
gateway: Gateway, tmp_path: Path
) -> None:
with scratch_database() as database_url:
_deploy_schema_before(database_url, tmp_path, GIN_MIGRATION)
assert _index_rows(database_url) == []
assert _applied_migrations(database_url) == tuple(name for name in SHIPPED_MIGRATIONS if name < GIN_MIGRATION)
entrypoint: Final = _run_migration_entrypoint(database_url)
assert entrypoint.returncode == 0, entrypoint.stdout + entrypoint.stderr
_assert_valid_gin_index(_index_rows(database_url))
assert GIN_MIGRATION in _applied_migrations(database_url), entrypoint.stdout
store: Final = "vs_" + uuid.uuid4().hex
provider_file_id: Final = "file-" + uuid.uuid4().hex[:16]
upgraded_environment: Final = {"DATABASE_URL": database_url, "DISABLE_SCHEMA_UPDATE": "true"}
with (
wire_server(_provider(store, provider_file_id)) as wire,
owned_proxy(gateway, tmp_path, upgraded_environment) as upgraded,
):
model: Final = f"integration-{uuid.uuid4().hex}"
upgraded.post(
"/model/new",
{
"model_name": model,
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "upgraded-provider-key",
"api_base": wire.url + "/v1",
},
"model_info": {},
},
)
uploaded: Final = upgraded.request_multipart(
"/v1/files",
{"purpose": "user_data", "target_model_names": model},
{"file": ("a.txt", b"notes\n", "text/plain")},
)
assert uploaded.status_code == 200, uploaded.text
managed: Final = string_value(JSON_OBJECT.validate_json(uploaded.content)["id"])
assert read_rows(
'SELECT flat_model_file_ids FROM "LiteLLM_ManagedFileTable" WHERE unified_file_id = %s',
(managed,),
database_url=database_url,
) == [{"flat_model_file_ids": [provider_file_id]}]
assert _listed_ids(upgraded, store, model) == (managed,)
def test_db_push_creates_a_valid_gin_index_on_the_flat_provider_file_ids(gateway: Gateway) -> None:
assert gateway.request("GET", "/health/liveliness").status_code == 200
_assert_valid_gin_index(_index_rows())

View file

@ -0,0 +1,753 @@
import base64
import hashlib
import json
import re
import signal
import threading
import uuid
from collections.abc import Callable, Generator, Mapping
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from dataclasses import dataclass
from functools import partial
from pathlib import Path
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final
from urllib.parse import parse_qs, urlsplit
import httpx
import psutil
import pytest
from integration._support.client import Gateway, Scenario, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from openai import AsyncOpenAI, OpenAI
from pydantic import JsonValue, TypeAdapter
MANAGED_PREFIX: Final = "litellm_proxy:"
CARRIED_PROVIDER_FILE_ID: Final = re.compile(r"(?:^|;)llm_output_file_id,([^;]+)")
UPLOAD_FILENAME: Final = re.compile(rb'filename="([^"]+)"')
FILE_PATH: Final = re.compile(r"^/v1/files/([^/]+)$")
STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
MANAGED_FILE_ROW: Final = (
'SELECT flat_model_file_ids, created_by, team_id FROM "LiteLLM_ManagedFileTable" WHERE unified_file_id = %s'
)
Listing = Callable[[Request], Reply]
def _provider_file_id(bearer: str, filename: str) -> str:
return "file-" + hashlib.sha256(f"{bearer}:{filename}".encode()).hexdigest()[:16]
def _bearer(request: Request) -> str:
return request.headers.get("authorization", "").removeprefix("Bearer ")
def _query(request: Request) -> dict[str, list[str]]:
return parse_qs(urlsplit(request.target).query, keep_blank_values=True)
def _json(response: httpx.Response) -> dict[str, JsonValue]:
return JSON_OBJECT.validate_json(response.content)
def _json_reply(body: Mapping[str, JsonValue], status: int = 200) -> Reply:
return Reply(status=status, body=json.dumps(body).encode())
def _file_object(file_id: str) -> dict[str, JsonValue]:
return {
"id": file_id,
"object": "file",
"bytes": 12,
"created_at": 1700000000,
"filename": "notes.txt",
"purpose": "user_data",
"status": "processed",
}
def _store_file(store: str, file_id: JsonValue) -> dict[str, JsonValue]:
return {
"id": file_id,
"object": "vector_store.file",
"usage_bytes": 123,
"created_at": 1700000001,
"vector_store_id": store,
"status": "completed",
"last_error": None,
"chunking_strategy": {"type": "static", "static": {"max_chunk_size_tokens": 800, "chunk_overlap_tokens": 400}},
"attributes": {},
}
def _page(store: str, file_ids: tuple[JsonValue, ...], *, has_more: bool = False) -> dict[str, JsonValue]:
return {
"object": "list",
"data": [_store_file(store, file_id) for file_id in file_ids],
"first_id": file_ids[0] if file_ids else None,
"last_id": file_ids[-1] if file_ids else None,
"has_more": has_more,
}
def _constant_listing(store: str, *file_ids: JsonValue) -> Listing:
return lambda _: _json_reply(_page(store, file_ids))
def _paged_listing(store: str, first: str, second: str) -> Listing:
def listing(request: Request) -> Reply:
if _query(request).get("after") == [first]:
return _json_reply(_page(store, (second,)))
return _json_reply(_page(store, (first,), has_more=True))
return listing
def _provider_error(status: int, message: str) -> dict[str, JsonValue]:
return {"error": {"message": message, "type": "provider_error", "code": str(status)}}
def _error_listing(status: int, message: str) -> Callable[[str, str], Listing]:
return lambda _store, _bearer: lambda _: _json_reply(_provider_error(status, message), status)
def _html_listing() -> Callable[[str, str], Listing]:
return lambda _store, _bearer: lambda _: Reply(body=b"<html>upstream maintenance</html>", content_type="text/html")
def _two_pages(store: str, bearer: str) -> Listing:
return _paged_listing(store, _provider_file_id(bearer, "a.txt"), _provider_file_id(bearer, "b.txt"))
def _raw_then_uploaded(raw_id: str) -> Callable[[str, str], Listing]:
return lambda store, bearer: _constant_listing(store, raw_id, _provider_file_id(bearer, "a.txt"))
def _uploaded_then_integer(store: str, bearer: str) -> Listing:
return _constant_listing(store, _provider_file_id(bearer, "a.txt"), 7)
def _provider(store: str, listing: Listing) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
path: Final = urlsplit(request.target).path
if request.method == "POST" and path == "/v1/files":
filename: Final = UPLOAD_FILENAME.search(request.body)
assert filename is not None, request.body[:200]
return _json_reply(_file_object(_provider_file_id(_bearer(request), filename.group(1).decode())))
if request.method == "POST" and path == f"/v1/vector_stores/{store}/files":
return _json_reply(_store_file(store, JSON_OBJECT.validate_json(request.body)["file_id"]))
if request.method == "GET" and path == f"/v1/vector_stores/{store}/files":
return listing(request)
file: Final = FILE_PATH.match(path)
if request.method == "GET" and file:
return _json_reply(_file_object(file.group(1)))
if request.method == "DELETE" and file:
return _json_reply({"id": file.group(1), "object": "file", "deleted": True})
return _json_reply({"error": {"message": f"unscripted {request.method} {request.target}"}}, 404)
return respond
def _decoded(managed_file_id: str) -> str:
decoded: Final = base64.urlsafe_b64decode(managed_file_id + "=" * (-len(managed_file_id) % 4)).decode()
assert decoded.startswith(MANAGED_PREFIX), decoded
return decoded
def _carried_provider_file_id(managed_file_id: str) -> str:
carried: Final = CARRIED_PROVIDER_FILE_ID.search(_decoded(managed_file_id))
assert carried is not None, managed_file_id
return carried.group(1)
def _upload(gateway: Gateway, key: str, target_model_names: str, filename: str) -> str:
uploaded: Final = gateway.request_multipart(
"/v1/files",
{"purpose": "user_data", "target_model_names": target_model_names},
{"file": (filename, f"notes in {filename}\n".encode(), "text/plain")},
key=key,
)
assert uploaded.status_code == 200, uploaded.text
return string_value(_json(uploaded)["id"])
def _listed(response: httpx.Response) -> dict[str, JsonValue]:
assert response.status_code == 200, response.text
return _json(response)
def _ids(page: Mapping[str, JsonValue]) -> tuple[JsonValue, ...]:
data: Final = page["data"]
assert isinstance(data, list), page
return tuple(object_value(entry)["id"] for entry in data)
def _sdk_base_url(gateway: Gateway) -> str:
return str(gateway.client.base_url).rstrip("/") + "/v1"
def _models_over_a_fresh_connection(gateway: Gateway, _: int) -> frozenset[str]:
with httpx.Client(base_url=gateway.client.base_url, timeout=15, trust_env=False) as client:
listed: Final = client.get("/v1/models", headers={"Authorization": f"Bearer {gateway.key}"})
assert listed.status_code == 200, listed.text
data: Final = _json(listed)["data"]
assert isinstance(data, list), listed.text
return frozenset(string_value(object_value(entry)["id"]) for entry in data)
def _every_worker_serves(gateway: Gateway, model: str) -> bool:
with ThreadPoolExecutor(max_workers=16) as pool:
rounds: Final = tuple(
tuple(pool.map(partial(_models_over_a_fresh_connection, gateway), range(16))) for _ in range(2)
)
return all(model in seen for round_ in rounds for seen in round_)
def _wait_until_every_worker_serves(gateway: Gateway, model: str) -> None:
eventually(lambda: _every_worker_serves(gateway, model), lambda served: served, seconds=90)
@dataclass(frozen=True, slots=True)
class _Member:
team: str
user: str
key: str
def _member(scenario: Scenario, *models: str) -> _Member:
team: Final = scenario.team(models=list(models))
user: Final = scenario.member(team)
return _Member(team, user, scenario.key(team_id=team, user_id=user))
@dataclass(frozen=True, slots=True)
class _Rig:
gateway: Gateway
scenario: Scenario
wire: Wire
store: str
bearer: str
model: str
def file_id(self, filename: str) -> str:
return _provider_file_id(self.bearer, filename)
def upload(self, key: str, filename: str) -> str:
managed: Final = _upload(self.gateway, key, self.model, filename)
assert _carried_provider_file_id(managed) == self.file_id(filename), _decoded(managed)
return managed
def list(
self,
key: str,
params: Mapping[str, str] | None = None,
headers: Mapping[str, str] | None = None,
*,
query: str | None = None,
) -> httpx.Response:
suffix: Final = "" if query is None else f"?{query}"
return self.gateway.request(
"GET", f"/v1/vector_stores/{self.store}/files{suffix}", key=key, params=params, headers=headers
)
def listed(self, key: str, params: Mapping[str, str] | None = None) -> dict[str, JsonValue]:
return _listed(self.list(key, params if params is not None else {"model": self.model}))
def list_requests(self) -> tuple[Request, ...]:
return tuple(
request
for request in self.wire.drain()
if (request.method, urlsplit(request.target).path) == ("GET", f"/v1/vector_stores/{self.store}/files")
)
def single_list_request(self) -> Request:
(request,) = self.list_requests()
return request
@contextmanager
def _rig(gateway: Gateway, *filenames: str, listing: Callable[[str, str], Listing] | None = None) -> Generator[_Rig]:
store: Final = "vs_" + uuid.uuid4().hex
bearer: Final = "provider-key-" + uuid.uuid4().hex[:8]
served: Final = (
listing(store, bearer)
if listing is not None
else _constant_listing(store, *(_provider_file_id(bearer, filename) for filename in filenames))
)
with gateway.scenario() as scenario, wire_server(_provider(store, served)) as wire:
model: Final = scenario.model(api_base=wire.url + "/v1", api_key=bearer)
_wait_until_every_worker_serves(gateway, model)
yield _Rig(gateway, scenario, wire, store, bearer, model)
def test_raw_httpx_list_returns_the_uploaders_managed_ids(gateway: Gateway) -> None:
with _rig(gateway, "a.txt", "b.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
managed_b: Final = rig.upload(member.key, "b.txt")
assert read_rows(MANAGED_FILE_ROW, (managed_a,)) == [
{"flat_model_file_ids": [rig.file_id("a.txt")], "created_by": member.user, "team_id": member.team}
]
page: Final = rig.listed(member.key)
assert _ids(page) == (managed_a, managed_b), page
assert (page["first_id"], page["last_id"]) == (managed_a, managed_b), page
assert page["has_more"] is False, page
listed: Final = rig.single_list_request()
assert _query(listed) == {}, listed.target
assert listed.headers["authorization"] == f"Bearer {rig.bearer}", listed.headers
def test_attach_by_managed_id_sends_the_provider_file_id_and_lists_it_back_managed(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
attached: Final = rig.gateway.request(
"POST", f"/v1/vector_stores/{rig.store}/files", {"file_id": managed_a}, key=member.key
)
assert attached.status_code == 200, attached.text
assert _json(attached)["id"] == managed_a, attached.text
attach_path: Final = f"/v1/vector_stores/{rig.store}/files"
attach_bodies: Final = [
JSON_OBJECT.validate_json(request.body)
for request in rig.wire.drain()
if (request.method, request.target) == ("POST", attach_path)
]
assert attach_bodies == [{"file_id": rig.file_id("a.txt")}], attach_bodies
assert _ids(rig.listed(member.key)) == (managed_a,)
def test_openai_sdk_sync_auto_pager_walks_pages_with_managed_cursors(gateway: Gateway) -> None:
with _rig(gateway, listing=_two_pages) as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
managed_b: Final = rig.upload(member.key, "b.txt")
with OpenAI(base_url=_sdk_base_url(gateway), api_key=member.key, max_retries=0) as client:
first: Final = client.vector_stores.files.list(rig.store, limit=1, extra_query={"model": rig.model})
assert [file.id for file in first.data] == [managed_a], first.model_dump_json()
assert first.has_more is True, first.model_dump_json()
second: Final = first.get_next_page()
assert [file.id for file in second.data] == [managed_b], second.model_dump_json()
assert second.has_more is False, second.model_dump_json()
queries: Final = [_query(request) for request in rig.list_requests()]
assert queries == [{"limit": ["1"]}, {"after": [rig.file_id("a.txt")], "limit": ["1"]}], queries
async def test_openai_sdk_async_auto_pager_walks_pages_with_managed_cursors(gateway: Gateway) -> None:
with _rig(gateway, listing=_two_pages) as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
managed_b: Final = rig.upload(member.key, "b.txt")
async with AsyncOpenAI(base_url=_sdk_base_url(gateway), api_key=member.key, max_retries=0) as client:
first: Final = await client.vector_stores.files.list(rig.store, limit=1, extra_query={"model": rig.model})
assert [file.id for file in first.data] == [managed_a], first.model_dump_json()
second: Final = await first.get_next_page()
assert [file.id for file in second.data] == [managed_b], second.model_dump_json()
queries: Final = [_query(request) for request in rig.list_requests()]
assert queries == [{"limit": ["1"]}, {"after": [rig.file_id("a.txt")], "limit": ["1"]}], queries
def test_after_cursor_with_a_managed_id_reaches_the_provider_decoded(gateway: Gateway) -> None:
with _rig(gateway, "b.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
managed_b: Final = rig.upload(member.key, "b.txt")
assert _ids(rig.listed(member.key, {"model": rig.model, "after": managed_a})) == (managed_b,)
assert _query(rig.single_list_request()) == {"after": [rig.file_id("a.txt")]}
def test_before_cursor_with_a_managed_id_reaches_the_provider_decoded(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
managed_b: Final = rig.upload(member.key, "b.txt")
assert _ids(rig.listed(member.key, {"model": rig.model, "before": managed_b})) == (managed_a,)
assert _query(rig.single_list_request()) == {"before": [rig.file_id("b.txt")]}
def test_after_cursor_with_a_raw_provider_id_is_forwarded_verbatim(gateway: Gateway) -> None:
with _rig(gateway, "b.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_b: Final = rig.upload(member.key, "b.txt")
raw_cursor: Final = "file-" + uuid.uuid4().hex[:16]
page: Final = rig.listed(member.key, {"model": rig.model, "after": raw_cursor})
assert _query(rig.single_list_request()) == {"after": [raw_cursor]}
assert _ids(page) == (managed_b,), page
def test_model_header_routing_returns_managed_ids(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
page: Final = _listed(rig.list(member.key, {}, {"x-litellm-model": rig.model}))
assert _ids(page) == (managed_a,), page
assert _query(rig.single_list_request()) == {}
def test_managed_vector_store_registry_routing_returns_managed_ids(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
registry_bearer: Final = "registry-key-" + uuid.uuid4().hex[:8]
gateway.post(
"/vector_store/new",
{
"vector_store_id": rig.store,
"custom_llm_provider": "openai",
"vector_store_name": "managed-ids-registry",
"litellm_params": {"api_base": rig.wire.url + "/v1", "api_key": registry_bearer},
},
)
rig.scenario.cleanups.callback(gateway.post, "/vector_store/delete", {"vector_store_id": rig.store})
page: Final = _listed(rig.list(member.key, {}))
assert _ids(page) == (managed_a,), page
listed: Final = rig.single_list_request()
assert listed.headers["authorization"] == f"Bearer {registry_bearer}", listed.headers
assert _query(listed) == {}, listed.target
def test_team_model_fallback_routing_returns_managed_ids(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
page: Final = _listed(rig.list(member.key, {}))
assert _ids(page) == (managed_a,), page
listed: Final = rig.single_list_request()
assert listed.headers["authorization"] == f"Bearer {rig.bearer}", listed.headers
def test_teammate_sees_the_uploaders_managed_ids(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as rig:
uploader: Final = _member(rig.scenario, rig.model)
teammate_user: Final = rig.scenario.member(uploader.team)
teammate_key: Final = rig.scenario.key(team_id=uploader.team, user_id=teammate_user)
managed_a: Final = rig.upload(uploader.key, "a.txt")
assert _ids(rig.listed(teammate_key)) == (managed_a,)
def test_proxy_admin_sees_every_managed_id(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as rig:
uploader: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(uploader.key, "a.txt")
assert _ids(rig.listed(gateway.key)) == (managed_a,)
def test_stranger_in_another_team_sees_raw_provider_ids(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as rig:
uploader: Final = _member(rig.scenario, rig.model)
stranger: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(uploader.key, "a.txt")
assert len(read_rows(MANAGED_FILE_ROW, (managed_a,))) == 1
page: Final = rig.listed(stranger.key)
assert _ids(page) == (rig.file_id("a.txt"),), page
assert (page["first_id"], page["last_id"]) == (rig.file_id("a.txt"), rig.file_id("a.txt")), page
def test_service_account_upload_is_shared_with_its_team_only(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as rig:
teammate: Final = _member(rig.scenario, rig.model)
service_account: Final = rig.scenario.key(team_id=teammate.team)
stranger: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(service_account, "a.txt")
assert read_rows(MANAGED_FILE_ROW, (managed_a,)) == [
{"flat_model_file_ids": [rig.file_id("a.txt")], "created_by": None, "team_id": teammate.team}
]
assert _ids(rig.listed(service_account)) == (managed_a,)
assert _ids(rig.listed(teammate.key)) == (managed_a,)
assert _ids(rig.listed(stranger.key)) == (rig.file_id("a.txt"),)
def test_key_without_user_or_team_owns_its_upload_alone(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as rig:
owner: Final = rig.scenario.key()
sibling: Final = rig.scenario.key()
managed_a: Final = rig.upload(owner, "a.txt")
assert read_rows(MANAGED_FILE_ROW, (managed_a,)) == [
{
"flat_model_file_ids": [rig.file_id("a.txt")],
"created_by": f"key:{hashlib.sha256(owner.encode()).hexdigest()}",
"team_id": None,
}
]
assert _ids(rig.listed(owner)) == (managed_a,)
assert _ids(rig.listed(sibling)) == (rig.file_id("a.txt"),)
def test_file_attached_by_raw_provider_id_stays_raw_beside_a_managed_one(gateway: Gateway) -> None:
raw_id: Final = "file-raw-" + uuid.uuid4().hex[:12]
with _rig(gateway, listing=_raw_then_uploaded(raw_id)) as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
attached: Final = rig.gateway.request(
"POST", f"/v1/vector_stores/{rig.store}/files", {"file_id": raw_id, "model": rig.model}, key=member.key
)
assert attached.status_code == 200, attached.text
assert _json(attached)["id"] == raw_id, attached.text
assert _ids(rig.listed(member.key)) == (raw_id, managed_a)
def test_multi_model_upload_maps_only_the_provider_id_the_managed_id_carries(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as first, _rig(gateway, "a.txt") as second:
member: Final = _member(first.scenario, first.model, second.model)
managed_a: Final = _upload(gateway, member.key, f"{first.model},{second.model}", "a.txt")
(row,) = read_rows(MANAGED_FILE_ROW, (managed_a,))
flat_ids: Final = row["flat_model_file_ids"]
assert isinstance(flat_ids, list), row
assert sorted(string_value(value) for value in flat_ids) == sorted(
(first.file_id("a.txt"), second.file_id("a.txt"))
), row
carried: Final = _carried_provider_file_id(managed_a)
assert carried in {first.file_id("a.txt"), second.file_id("a.txt")}, carried
first_ids: Final = _ids(first.listed(member.key))
second_ids: Final = _ids(second.listed(member.key))
assert first_ids == ((managed_a,) if carried == first.file_id("a.txt") else (first.file_id("a.txt"),))
assert second_ids == ((managed_a,) if carried == second.file_id("a.txt") else (second.file_id("a.txt"),))
def test_deleting_the_managed_file_makes_its_listing_raw_again(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
assert _ids(rig.listed(member.key)) == (managed_a,)
deleted: Final = gateway.request("DELETE", f"/v1/files/{managed_a}", key=member.key)
assert deleted.status_code == 200, deleted.text
assert _json(deleted)["deleted"] is True, deleted.text
assert read_rows(MANAGED_FILE_ROW, (managed_a,)) == []
assert _ids(rig.listed(member.key)) == (rig.file_id("a.txt"),)
deletes: Final = [request.target for request in rig.wire.drain() if request.method == "DELETE"]
assert deletes == [f"/v1/files/{rig.file_id('a.txt')}"], deletes
def test_empty_page_is_returned_unchanged(gateway: Gateway) -> None:
with _rig(gateway) as rig:
member: Final = _member(rig.scenario, rig.model)
rig.upload(member.key, "a.txt")
assert rig.listed(member.key) == _page(rig.store, ())
def test_duplicate_provider_ids_in_one_page_are_both_mapped(gateway: Gateway) -> None:
with _rig(gateway, "a.txt", "a.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
page: Final = rig.listed(member.key)
assert _ids(page) == (managed_a, managed_a), page
assert (page["first_id"], page["last_id"]) == (managed_a, managed_a), page
def test_mixed_page_maps_only_the_managed_entries_and_the_matching_edge_ids(gateway: Gateway) -> None:
raw_id: Final = "file-raw-" + uuid.uuid4().hex[:12]
with _rig(gateway, listing=_raw_then_uploaded(raw_id)) as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
page: Final = rig.listed(member.key)
assert _ids(page) == (raw_id, managed_a), page
assert (page["first_id"], page["last_id"]) == (raw_id, managed_a), page
def test_repeated_identical_lists_each_reach_the_provider_and_each_map(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
assert _ids(rig.listed(member.key)) == (managed_a,)
assert _ids(rig.listed(member.key)) == (managed_a,)
targets: Final = [request.target for request in rig.list_requests()]
assert targets == [f"/v1/vector_stores/{rig.store}/files"] * 2, targets
def test_duplicated_managed_after_cursor_reaches_the_provider_once_decoded(gateway: Gateway) -> None:
with _rig(gateway, "b.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
managed_b: Final = rig.upload(member.key, "b.txt")
page: Final = _listed(rig.list(member.key, query=f"model={rig.model}&after={managed_a}&after={managed_a}"))
assert _ids(page) == (managed_b,), page
assert _query(rig.single_list_request()) == {"after": [rig.file_id("a.txt")]}
def _unpadded(raw: bytes) -> str:
return base64.urlsafe_b64encode(raw).decode().rstrip("=")
@pytest.mark.parametrize(
"cursor",
(
pytest.param("12345", id="integer-like"),
pytest.param("", id="empty"),
pytest.param("x" * 5000, id="five-kilobyte"),
pytest.param(_unpadded(b"litellm_proxy:text/plain;unified_id,abc"), id="managed-without-provider-id"),
pytest.param(_unpadded(b"\xff\xfe\xfd\xfc"), id="non-utf8-base64"),
),
)
def test_unmappable_after_cursors_are_forwarded_verbatim(gateway: Gateway, cursor: str) -> None:
with _rig(gateway, "b.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_b: Final = rig.upload(member.key, "b.txt")
page: Final = _listed(rig.list(member.key, {"model": rig.model, "after": cursor}))
assert _query(rig.single_list_request()) == {"after": [cursor]}
liveliness: Final = gateway.request("GET", "/health/liveliness")
assert liveliness.status_code == 200, liveliness.text
assert _ids(page) == (managed_b,), page
def test_two_different_after_values_forward_the_last_one(gateway: Gateway) -> None:
with _rig(gateway, "b.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_b: Final = rig.upload(member.key, "b.txt")
page: Final = _listed(rig.list(member.key, query=f"model={rig.model}&after=first-value&after=second-value"))
assert _query(rig.single_list_request()) == {"after": ["second-value"]}
assert _ids(page) == (managed_b,), page
@pytest.mark.parametrize("status", (401, 404, 500))
def test_provider_errors_reach_the_caller_and_other_models_keep_mapping(gateway: Gateway, status: int) -> None:
message: Final = f"provider refused listing {uuid.uuid4().hex[:8]}"
with _rig(gateway, listing=_error_listing(status, message)) as failing, _rig(gateway, "a.txt") as healthy:
member: Final = _member(failing.scenario, failing.model, healthy.model)
managed_a: Final = healthy.upload(member.key, "a.txt")
failed: Final = failing.list(member.key, {"model": failing.model})
assert _json(failed) == _provider_error(status, message), failed.text
assert len(failing.list_requests()) == 1
assert _ids(healthy.listed(member.key)) == (managed_a,)
liveliness: Final = gateway.request("GET", "/health/liveliness")
assert liveliness.status_code == 200, liveliness.text
def test_non_json_provider_body_is_an_error_response_and_other_models_keep_mapping(gateway: Gateway) -> None:
with _rig(gateway, listing=_html_listing()) as failing, _rig(gateway, "a.txt") as healthy:
member: Final = _member(failing.scenario, failing.model, healthy.model)
managed_a: Final = healthy.upload(member.key, "a.txt")
failed: Final = failing.list(member.key, {"model": failing.model})
assert failed.status_code == 500, failed.text
assert string_value(object_value(_json(failed)["error"])["message"]), failed.text
assert _ids(healthy.listed(member.key)) == (managed_a,)
liveliness: Final = gateway.request("GET", "/health/liveliness")
assert liveliness.status_code == 200, liveliness.text
def test_non_string_ids_in_a_page_are_left_alone_while_strings_map(gateway: Gateway) -> None:
with _rig(gateway, listing=_uploaded_then_integer) as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
page: Final = rig.listed(member.key)
assert _ids(page) == (managed_a, 7), page
assert (page["first_id"], page["last_id"]) == (managed_a, 7), page
def test_retrieving_the_managed_file_still_resolves_to_the_provider_file(gateway: Gateway) -> None:
with _rig(gateway, "a.txt") as rig:
member: Final = _member(rig.scenario, rig.model)
managed_a: Final = rig.upload(member.key, "a.txt")
retrieved: Final = gateway.request("GET", f"/v1/files/{managed_a}", key=member.key)
assert retrieved.status_code == 200, retrieved.text
file: Final = _json(retrieved)
assert (file["id"], file["object"], file["purpose"]) == (managed_a, "file", "user_data"), retrieved.text
def _burst(gateway: Gateway, store: str, key: str, model: str, size: int) -> tuple[httpx.Response, ...]:
def one(_: int) -> httpx.Response:
return gateway.request("GET", f"/v1/vector_stores/{store}/files", key=key, params={"model": model})
with ThreadPoolExecutor(max_workers=size) as pool:
return tuple(pool.map(one, range(size)))
@pytest.mark.timeout(180)
def test_provider_outage_mid_burst_fails_loudly_and_mapping_resumes_after_recovery(gateway: Gateway) -> None:
store: Final = "vs_" + uuid.uuid4().hex
bearer: Final = "provider-key-" + uuid.uuid4().hex[:8]
provider_a: Final = _provider_file_id(bearer, "a.txt")
respond: Final = _provider(store, _constant_listing(store, provider_a))
with gateway.scenario() as scenario:
with wire_server(respond) as wire:
model: Final = scenario.model(api_base=wire.url + "/v1", api_key=bearer)
_wait_until_every_worker_serves(gateway, model)
member: Final = _member(scenario, model)
managed_a: Final = _upload(gateway, member.key, model, "a.txt")
assert _carried_provider_file_id(managed_a) == provider_a
served: Final = _burst(gateway, store, member.key, model, 40)
assert [_ids(_listed(response)) for response in served] == [(managed_a,)] * 40
assert sum(1 for request in wire.drain() if request.method == "GET") == 40
port: Final = urlsplit(wire.url).port
assert port is not None
failed: Final = _burst(gateway, store, member.key, model, 20)
assert [response.status_code for response in failed] == [500] * 20, [r.text for r in failed[:3]]
for response in failed:
assert string_value(object_value(_json(response)["error"])["message"]), response.text
liveliness: Final = gateway.request("GET", "/health/liveliness")
assert liveliness.status_code == 200, liveliness.text
with wire_server(respond, port=port) as revived:
recovered: Final = _burst(gateway, store, member.key, model, 40)
assert [_ids(_listed(response)) for response in recovered] == [(managed_a,)] * 40
assert sum(1 for request in revived.drain() if request.method == "GET") == 40
def _open_connections_to(pid: int, url: str) -> int:
port: Final = urlsplit(url).port
return sum(
1
for connection in psutil.Process(pid).net_connections(kind="tcp")
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
)
def _tolerant_list(gateway: Gateway, store: str, key: str, model: str) -> httpx.Response | None:
try:
return gateway.request("GET", f"/v1/vector_stores/{store}/files", key=key, params={"model": model})
except httpx.HTTPError:
return None
@pytest.mark.timeout(300)
def test_worker_sigkill_mid_burst_leaves_the_sibling_mapping_ids(gateway: Gateway, tmp_path: Path) -> None:
store: Final = "vs_" + uuid.uuid4().hex
bearer: Final = "provider-key-" + uuid.uuid4().hex[:8]
provider_a: Final = _provider_file_id(bearer, "a.txt")
release: Final = threading.Event()
held: Final[SimpleQueue[str]] = SimpleQueue()
def held_listing(request: Request) -> Reply:
held.put(request.target)
assert release.wait(timeout=120), "The burst was never released"
return _json_reply(_page(store, (provider_a,)))
with gateway.scenario() as scenario, wire_server(_provider(store, held_listing)) as wire:
model: Final = scenario.model(api_base=wire.url + "/v1", api_key=bearer)
_wait_until_every_worker_serves(gateway, model)
member: Final = _member(scenario, model)
managed_a: Final = _upload(gateway, member.key, model, "a.txt")
with owned_proxy_process(gateway, tmp_path, {}, workers=2) as owned:
candidate: Final = owned.gateway
workers: Final = eventually(
lambda: tuple(int(match.group(1)) for match in STARTED_WORKER.finditer(owned.log.read_text())),
lambda pids: len(pids) == 2,
seconds=60,
)
with ThreadPoolExecutor(max_workers=20) as pool:
burst: Final = tuple(
pool.submit(_tolerant_list, candidate, store, member.key, model) for _ in range(20)
)
eventually(held.qsize, lambda size: size == 20, seconds=60)
held_by: Final = MappingProxyType({pid: _open_connections_to(pid, wire.url) for pid in workers})
assert sum(held_by.values()) == 20, held_by
victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__)
victim: Final = psutil.Process(victim_pid)
victim.suspend()
victim.send_signal(signal.SIGKILL)
release.set()
served: Final = tuple(result for future in burst if (result := future.result()) is not None)
assert held_by[survivor_pid] >= 1, held_by
assert len(served) == held_by[survivor_pid], (held_by, len(served))
for response in served:
assert _ids(_listed(response)) == (managed_a,)
assert psutil.Process(survivor_pid).is_running()
follow_up: Final = eventually(
lambda: _tolerant_list(candidate, store, member.key, model),
lambda response: response is not None and response.status_code == 200,
seconds=60,
)
assert follow_up is not None
assert _ids(_listed(follow_up)) == (managed_a,)

View file

@ -5,6 +5,7 @@ import json
import os
import subprocess
import sys
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final
@ -24,6 +25,29 @@ GROUPS: Final = MappingProxyType(
)
@dataclass(frozen=True, slots=True)
class Selection:
nodes: tuple[str, ...]
foreign: tuple[str, ...]
def file_of(node: str) -> str:
return node.split("::", 1)[0]
def select(requested: tuple[str, ...], group_files: tuple[str, ...]) -> Selection:
members: Final = frozenset(group_files)
return Selection(
nodes=requested or group_files,
foreign=tuple(sorted({node for node in requested if file_of(node) not in members})),
)
def uncollected(nodes: tuple[str, ...], collected: frozenset[str]) -> tuple[str, ...]:
collected_files: Final = frozenset(file_of(node) for node in collected)
return tuple(node for node in nodes if file_of(node) not in collected_files)
def main() -> int:
parser: Final = argparse.ArgumentParser()
parser.add_argument("group", choices=tuple(GROUPS))
@ -32,7 +56,7 @@ def main() -> int:
parser.add_argument("--order-seed", type=int, default=int(os.environ.get("INTEGRATION_ORDER_SEED", "0")))
parser.add_argument("--workers", type=int, default=int(os.environ.get("INTEGRATION_WORKERS", "1")))
parser.add_argument("--list", action="store_true", help="print the group's test files and exit")
parser.add_argument("files", nargs="*", help="run only these files of the group")
parser.add_argument("files", nargs="*", help="run only these files, or pytest node ids inside them, of the group")
options: Final = parser.parse_intermixed_args()
root: Final = Path(__file__).resolve().parents[2]
group_files: Final = tuple(
@ -43,11 +67,10 @@ def main() -> int:
if options.list:
print("\n".join(group_files))
return 0
foreign: Final = sorted(set(options.files) - set(group_files))
if foreign:
parser.error(f"Not in the {options.group} group: {', '.join(foreign)}")
selected: Final = tuple(options.files) or group_files
if not selected:
selection: Final = select(tuple(options.files), group_files)
if selection.foreign:
parser.error(f"Not in the {options.group} group: {', '.join(selection.foreign)}")
if not selection.nodes:
parser.error(f"No integration test files selected for {options.group}")
output: Final = options.results.resolve()
output.mkdir(parents=True, exist_ok=True)
@ -62,7 +85,7 @@ def main() -> int:
sys.executable,
"-m",
"pytest",
*selected,
*selection.nodes,
"-vv",
"-rs",
"--strict-markers",
@ -86,8 +109,7 @@ def main() -> int:
if result != 0:
return result
evidence: Final = json.loads((output / "execution.json").read_text())
collected_files: Final = {node.split("::", 1)[0] for node in evidence["collected"]}
empty: Final = tuple(path for path in selected if path not in collected_files)
empty: Final = uncollected(selection.nodes, frozenset(evidence["collected"]))
if empty:
sys.stderr.write(f"Selected integration files collected zero tests: {', '.join(empty)}\n")
return 1

View file

@ -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

View file

@ -20,8 +20,10 @@ from typing_extensions import ReadOnly
from litellm.proxy.db.autorouter_session_rollup import (
AUTOROUTER_BENCHMARKS_SQL,
UPSERT_AUTOROUTER_SESSION_SQL,
UPSERT_AUTOROUTER_USER_SESSION_SQL,
AutoRouterTurnTransaction,
flush_autorouter_turn_transactions,
write_autorouter_turn,
)
from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import SpendLogCleanup
@ -684,3 +686,104 @@ async def test_a_router_type_change_mid_session_keeps_session_shape_with_the_ses
assert (rows["complexity"]["sessions"], rows["complexity"]["session_turns"], rows["complexity"]["turns"]) == (1, 2, 1)
assert (rows["quality"]["sessions"], rows["quality"]["session_turns"], rows["quality"]["turns"]) == (0, 0, 1)
assert rows["quality"]["spend"] == 2.0
@pytest.mark.parametrize("statement", [UPSERT_AUTOROUTER_SESSION_SQL, UPSERT_AUTOROUTER_USER_SESSION_SQL])
async def test_a_sessionless_turn_writes_its_router_day_row_and_no_session_row(db, statement: str):
key = f"k-{uuid.uuid4()}"
router = f"auto-{uuid.uuid4()}"
for offset in range(2):
await write_autorouter_turn(
db,
AutoRouterTurnTransaction(
api_key=key,
user_id="u-sessionless",
session_id="",
router_name=router,
router_type="complexity",
model="A",
turn_at=T0 + timedelta(seconds=offset),
total_tokens=10,
spend=1.0,
saved_spend=2.0,
classifier_cost=0.1,
covered=True,
cache_hit=False,
cache_ttl_seconds=None,
cache_touched=True,
savings_estimated_turns=1,
savings_estimated_actual_spend=1.0,
savings_estimated_saved_spend=2.0,
),
statement,
)
(day,) = await _days(db, key, router=router)
assert (day["turns"], day["spend"], day["saved_spend"], day["classifier_cost"]) == (2, 2.0, 4.0, 0.2)
assert (day["sessions"], day["session_turns"]) == (0, 0)
for table in ("LiteLLM_AutoRouterSession", "LiteLLM_AutoRouterUserSession"):
assert await db.query_raw(f'SELECT 1 FROM "{table}" WHERE router_name = $1', router) == []
async def test_router_day_money_reconciles_with_the_overall_daily_total_including_sessionless_requests(db):
from litellm.proxy.db.daily_spend_bulk_upsert import DAILY_SPEND_TABLES, build_bulk_upsert, merge_by_conflict_key
key = f"k-{uuid.uuid4()}"
router = f"auto-{uuid.uuid4()}"
requests = (("session-1", 0.25, 1.5), ("session-1", 0.5, 2.0), ("", 0.1, 0.25))
for offset, (session_id, spend, saved) in enumerate(requests):
await write_autorouter_turn(
db,
AutoRouterTurnTransaction(
api_key=key,
user_id="u1",
session_id=session_id,
router_name=router,
router_type="complexity",
model="A",
turn_at=T0 + timedelta(seconds=offset),
total_tokens=10,
spend=spend,
saved_spend=saved,
classifier_cost=0.0,
covered=True,
cache_hit=False,
cache_ttl_seconds=None,
cache_touched=True,
savings_estimated_turns=1,
savings_estimated_actual_spend=spend,
savings_estimated_saved_spend=saved,
),
)
table = DAILY_SPEND_TABLES["user"]
statement, values = build_bulk_upsert(
table,
merge_by_conflict_key(
table,
tuple(
{
"user_id": "u1",
"date": T0.date().isoformat(),
"api_key": key,
"model": "A",
"custom_llm_provider": "anthropic",
"model_group": router,
"spend": spend,
"api_requests": 1,
"successful_requests": 1,
"autorouter_savings_spend": saved,
}
for _, spend, saved in requests
),
),
)
await db.execute_raw(statement, *values)
(overall,) = await db.query_raw(
'SELECT SUM(autorouter_savings_spend)::float8 AS saved FROM "LiteLLM_DailyUserSpend" WHERE date = $1 AND api_key = $2',
T0.date().isoformat(),
key,
)
(row,) = await _days(db, key, router=router)
assert overall["saved"] == row["saved_spend"] == pytest.approx(3.75)
assert (row["turns"], row["spend"], row["sessions"], row["session_turns"]) == (3, pytest.approx(0.85), 1, 2)

View file

@ -249,17 +249,10 @@ async def test_retired_history_never_recreates_an_initial_zero(db: Prisma, recor
assert after["savings_estimated_turns"] == 1 and after["savings_estimated_actual_spend"] == 0.17
async def test_native_observation_enters_spend_pipeline_once_with_shared_daily_attribution(
db: Prisma, record: Callable[..., BaselineAccountingRecord], monkeypatch: pytest.MonkeyPatch,
) -> None:
import os
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
def _native_observation_payload(event: BaselineAccountingRecord) -> dict[str, object]:
"""The spend payload a captured, sessioned, auto-routed anthropic_messages request produces."""
from litellm.proxy.hooks.autorouter_baseline_cache import CapturedBaselineObservation
from litellm.proxy.utils import PrismaClient, ProxyLogging
event: Final = record("routed", identical=False)
capture: Final = CapturedBaselineObservation(
scope=event.scope, api_key=event.api_key, session_id=event.session_id,
router_name=event.router_name, baseline_model=event.baseline_model,
@ -272,7 +265,7 @@ async def test_native_observation_enters_spend_pipeline_once_with_shared_daily_a
"autorouter_savings": None, "autorouter_savings_estimate": {"version": 3, "status": "unknown", "reason": "pending_projection"},
"autorouter_baseline_observation": capture.model_dump_json(),
}
payload: Final = {
return {
"request_id": event.observation.request_id, "api_key": event.api_key, "session_id": event.session_id,
"startTime": datetime.fromtimestamp(event.observation.started_at, timezone.utc).isoformat(),
"endTime": datetime.fromtimestamp(event.observation.available_at, timezone.utc).isoformat(),
@ -282,6 +275,19 @@ async def test_native_observation_enters_spend_pipeline_once_with_shared_daily_a
"user": None, "team_id": "", "organization_id": "org", "agent_id": None,
"end_user": "", "request_tags": '["tag","tag"]',
}
async def test_native_observation_enters_spend_pipeline_once_with_shared_daily_attribution(
db: Prisma, record: Callable[..., BaselineAccountingRecord], monkeypatch: pytest.MonkeyPatch,
) -> None:
import os
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
from litellm.proxy.utils import PrismaClient, ProxyLogging
event: Final = record("routed", identical=False)
payload: Final = _native_observation_payload(event)
monkeypatch.delenv("DATABASE_URL_READ_REPLICA", raising=False)
client: Final = PrismaClient(os.environ["DATABASE_URL"], ProxyLogging(UserApiKeyCache()))
writer: Final = DBSpendUpdateWriter()
@ -319,3 +325,42 @@ async def test_native_observation_enters_spend_pipeline_once_with_shared_daily_a
assert tag_rows[0]["spend"] == tag_rows[0]["api_requests"] == 0
finally:
await client.db.disconnect()
async def test_without_spend_logs_a_captured_turn_keeps_only_its_router_day_row(
db: Prisma, record: Callable[..., BaselineAccountingRecord], monkeypatch: pytest.MonkeyPatch,
) -> None:
import os
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.db.autorouter_session_rollup import flush_autorouter_turn_transactions
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
from litellm.proxy.utils import PrismaClient, ProxyLogging
event: Final = record("unlogged", identical=False)
monkeypatch.delenv("DATABASE_URL_READ_REPLICA", raising=False)
client: Final = PrismaClient(os.environ["DATABASE_URL"], ProxyLogging(UserApiKeyCache()))
try:
await client.db.connect()
await DBSpendUpdateWriter()._enqueue_autorouter_turn_transaction(
_native_observation_payload(event), client, spend_logs_kept=False
)
assert client.baseline_accounting_transactions == []
(turn,) = client.autorouter_turn_transactions
await flush_autorouter_turn_transactions(client, (turn,), n_retry_times=0)
finally:
client.autorouter_turn_transactions.clear()
await client.db.disconnect()
assert await db.query_raw(
'SELECT 1 FROM "LiteLLM_AutoRouterBaselineObservation" WHERE request_id=$1', event.observation.request_id
) == []
days: Final = await db.query_raw(
'SELECT turns, spend FROM "LiteLLM_AutoRouterDailySpend" WHERE api_key=$1 AND router_name=$2',
event.api_key, event.router_name,
)
assert [(day["turns"], day["spend"]) for day in days] == [(1, 0.17)]
for table in ("LiteLLM_AutoRouterSession", "LiteLLM_AutoRouterUserSession"):
assert await db.query_raw(
f'SELECT 1 FROM "{table}" WHERE api_key=$1 AND router_name=$2', event.api_key, event.router_name
) == []

View file

@ -9,9 +9,10 @@ import asyncio
import base64
import json
import logging
from types import MappingProxyType
import pytest
from typing import Optional
from typing import Final, Optional
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import LitellmUserRoles, ProxyException, UserAPIKeyAuth
@ -299,6 +300,90 @@ async def test_get_user_created_file_ids_remaps_stored_raw_provider_id_to_unifie
assert files[0].purpose == raw_provider_object.purpose
@pytest.mark.asyncio
async def test_provider_file_id_resolver_returns_owned_mappings_with_owner_scoped_filter() -> (
None
):
managed_files: Final = _make_managed_files_instance()
managed_row: Final = MagicMock(
unified_file_id="unified-file-id",
flat_model_file_ids=["file-provider-1", "file-provider-2"],
)
find_many: Final = AsyncMock(return_value=[managed_row])
managed_files.prisma_client.db.litellm_managedfiletable.find_many = find_many
unified_file_ids: Final = (
await managed_files.get_unified_file_ids_for_provider_file_ids(
provider_file_ids=(
"file-provider-1",
"file-provider-2",
"file-unmanaged-2",
"file-provider-1",
),
user_api_key_dict=_make_team_member_api_key_dict(),
)
)
assert unified_file_ids == {
"file-provider-1": "unified-file-id",
"file-provider-2": "unified-file-id",
}
assert isinstance(unified_file_ids, MappingProxyType)
find_many.assert_awaited_once_with(
where={
"OR": [{"created_by": "test-user"}, {"team_id": "test-team"}],
"flat_model_file_ids": {
"hasSome": ["file-provider-1", "file-provider-2", "file-unmanaged-2"],
},
}
)
@pytest.mark.asyncio
async def test_provider_file_id_resolver_denies_unowned_callers_without_database_query() -> (
None
):
managed_files: Final = _make_managed_files_instance()
find_many: Final = AsyncMock()
managed_files.prisma_client.db.litellm_managedfiletable.find_many = find_many
no_owner: Final = UserAPIKeyAuth(
api_key=None,
token=None,
user_id=None,
team_id=None,
parent_otel_span=None,
)
unified_file_ids: Final = (
await managed_files.get_unified_file_ids_for_provider_file_ids(
provider_file_ids=("file-provider-1",),
user_api_key_dict=no_owner,
)
)
assert unified_file_ids == {}
assert isinstance(unified_file_ids, MappingProxyType)
find_many.assert_not_awaited()
@pytest.mark.asyncio
async def test_provider_file_id_resolver_skips_database_query_for_empty_input() -> None:
managed_files: Final = _make_managed_files_instance()
find_many: Final = AsyncMock()
managed_files.prisma_client.db.litellm_managedfiletable.find_many = find_many
unified_file_ids: Final = (
await managed_files.get_unified_file_ids_for_provider_file_ids(
provider_file_ids=(),
user_api_key_dict=_make_user_api_key_dict(),
)
)
assert unified_file_ids == {}
assert isinstance(unified_file_ids, MappingProxyType)
find_many.assert_not_awaited()
@pytest.mark.asyncio
async def test_afile_list_returns_owner_scoped_managed_files():
managed_files = _make_managed_files_instance()

View file

@ -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

View file

@ -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)."""

View file

@ -5030,3 +5030,116 @@ async def test_openai_stream_relays_the_served_service_tier_on_every_chunk_inclu
assert [chunk.get("service_tier") for chunk in relayed] == ["default"] * len(relayed), relayed
assert relayed[-1]["usage"]["total_tokens"] == 11
def _last_chunk_carries_finish_reason_wrapper(
logging_obj: Logging, finish_reason: str, sync_stream: bool
) -> CustomStreamWrapper:
"""An OpenAI-compatible SSE body whose LAST chunk carries both a delta and the finish_reason, as vLLM emits
when speculative decoding finishes a reply in one engine step."""
from litellm.llms.openai.chat.gpt_transformation import (
OpenAIChatCompletionStreamingHandler,
)
def line(delta: dict, finish: Optional[str] = None) -> str:
chunk = {
"id": "chatcmpl-1",
"object": "chat.completion.chunk",
"created": 1,
"model": "m",
"choices": [{"index": 0, "delta": delta, "logprobs": None, "finish_reason": finish}],
}
return f"data: {json.dumps(chunk)}"
if finish_reason == "tool_calls":
lines = [
line({"role": "assistant", "content": ""}),
line({"tool_calls": [{"id": "call_1", "type": "function", "index": 0, "function": {"name": "bash", "arguments": ""}}]}),
line({"tool_calls": [{"index": 0, "function": {"arguments": '{"command": "ls'}}]}),
line({"tool_calls": [{"index": 0, "function": {"arguments": '"}'}}]}, "tool_calls"),
]
else:
lines = [
line({"role": "assistant", "content": "Hello, this reply is"}),
line({"content": " cut off"}, "length"),
]
lines.append("data: [DONE]")
if sync_stream:
streaming_response = iter(lines)
else:
async def _stream():
for item in lines:
yield item
streaming_response = _stream()
return CustomStreamWrapper(
completion_stream=OpenAIChatCompletionStreamingHandler(
streaming_response=streaming_response, sync_stream=sync_stream
),
model="m",
logging_obj=logging_obj,
custom_llm_provider="hosted_vllm",
)
@pytest.mark.parametrize("finish_reason", ["tool_calls", "length"])
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_logged_response_keeps_finish_reason_from_last_content_chunk(
finish_reason: str, sync_mode: bool, logging_obj: Logging
):
"""The client already got the right finish_reason here; the complete response built from ``chunks`` for
callbacks and SpendLogs used to say "stop" instead (tool calls and truncated replies both mislogged)."""
response = _last_chunk_carries_finish_reason_wrapper(
logging_obj, finish_reason, sync_stream=sync_mode
)
if sync_mode:
received = list(response)
else:
received = [chunk async for chunk in response]
assert [c.choices[0].finish_reason for c in received if c.choices and c.choices[0].finish_reason] == [
finish_reason
]
logged = litellm.stream_chunk_builder(chunks=response.chunks)
assert logged.choices[0].finish_reason == finish_reason
if finish_reason == "tool_calls":
tool_calls = logged.choices[0].message.tool_calls
assert len(tool_calls) == 1
assert tool_calls[0].function.arguments == '{"command": "ls"}'
else:
assert logged.choices[0].message.content == "Hello, this reply is cut off"
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_no_synthetic_finish_reason_logged_when_provider_sent_none(sync_mode: bool, logging_obj: Logging):
"""A stream that ends before the provider sent any finish_reason (e.g. an Anthropic stream cut after
message_start) must not gain one in ``chunks``: the response builder relies on its absence to estimate usage
instead of taking the provider's placeholder."""
chunks = [
ModelResponseStream(
id="chatcmpl-1",
created=1,
model=None,
object="chat.completion.chunk",
choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(content=text, role="assistant"))],
)
for text in ("partial", " reply")
]
response = CustomStreamWrapper(
completion_stream=ModelResponseListIterator(model_responses=chunks),
model="bedrock/m",
custom_llm_provider="bedrock",
logging_obj=logging_obj,
)
if sync_mode:
list(response)
else:
[c async for c in response]
assert response.received_finish_reason is None
assert all(not (c.choices and c.choices[0].finish_reason) for c in response.chunks)

View file

@ -520,6 +520,39 @@ def test_reasoning_effort_maps_to_reasoning_effort_for_openai_gpt5_converse(mode
assert "thinking" not in additional_request_params
@pytest.mark.parametrize(
"model",
[
"us.openai.gpt-5.6-luna",
"bedrock/converse/global.openai.gpt-5.6-terra",
"us.openai.gpt-6-astra",
],
)
def test_openai_gpt5_converse_rejects_effort_level_disabled_in_model_map(model, local_model_cost_map):
config = AmazonConverseConfig()
assert litellm.utils.is_explicitly_disabled_factory(
model=model, custom_llm_provider="bedrock_converse", key="supports_minimal_reasoning_effort"
)
with pytest.raises(litellm.utils.UnsupportedParamsError, match="minimal"):
config.map_openai_params(
non_default_params={"reasoning_effort": "minimal"},
optional_params={},
model=model,
drop_params=False,
)
optional_params = config.map_openai_params(
non_default_params={"reasoning_effort": "minimal"},
optional_params={},
model=model,
drop_params=True,
)
_, additional_request_params, _, _ = config._prepare_request_params(optional_params, model)
assert "reasoning" not in additional_request_params
assert "thinking" not in additional_request_params
@pytest.mark.parametrize(
"model",
[

View file

@ -328,6 +328,36 @@ class TestBackgroundDrop:
assert not [r for r in caplog.records if "dropping unsupported parameter" in r.getMessage()]
class TestDisabledReasoningEffort:
@pytest.mark.parametrize("model", ["us.openai.gpt-5.6-luna", MODEL])
def test_effort_level_disabled_in_model_map_is_rejected(self, model, local_model_cost_map):
with pytest.raises(litellm.UnsupportedParamsError, match="minimal"):
_cfg().map_openai_params(
response_api_optional_params={"reasoning": {"effort": "minimal"}}, model=model, drop_params=False
)
@pytest.mark.parametrize("model", ["us.openai.gpt-5.6-luna", MODEL])
def test_effort_level_disabled_in_model_map_is_dropped_with_drop_params(self, model, local_model_cost_map):
params = _cfg().map_openai_params(
response_api_optional_params={"reasoning": {"effort": "minimal", "summary": "auto"}, "max_output_tokens": 64},
model=model,
drop_params=True,
)
assert params == {"reasoning": {"summary": "auto"}, "max_output_tokens": 64}
def test_effort_only_reasoning_is_removed_when_dropped(self, local_model_cost_map):
params = _cfg().map_openai_params(
response_api_optional_params={"reasoning": {"effort": "minimal"}}, model=MODEL, drop_params=True
)
assert params == {}
def test_supported_effort_level_is_forwarded(self, local_model_cost_map):
params = _cfg().map_openai_params(
response_api_optional_params={"reasoning": {"effort": "low"}}, model=MODEL, drop_params=False
)
assert params == {"reasoning": {"effort": "low"}}
def _never_fetch(url: str) -> str:
raise AssertionError(f"unexpected sync fetch of {url}")

View file

@ -18,6 +18,10 @@ from litellm.llms.bedrock.messages.mantle_transformation import (
AmazonMantleMessagesConfig,
)
# AWS names this header for Mantle workspaces on the Anthropic Messages API, checked 2026-10-02:
# https://docs.aws.amazon.com/bedrock/latest/userguide/workspaces.html
_MANTLE_WORKSPACE_HEADER = "anthropic-workspace-id"
def _anthropic_response(url: str) -> httpx.Response:
return httpx.Response(
@ -345,7 +349,7 @@ def test_mantle_validate_environment_sets_workspace_header():
optional_params={},
litellm_params={"aws_bedrock_project_id": "proj_abc123def456"},
)
assert headers["anthropic-workspace"] == "proj_abc123def456"
assert headers[_MANTLE_WORKSPACE_HEADER] == "proj_abc123def456"
def test_mantle_validate_environment_without_project_id():
@ -357,7 +361,7 @@ def test_mantle_validate_environment_without_project_id():
optional_params={},
litellm_params={"aws_bedrock_project_id": None},
)
assert "anthropic-workspace" not in headers
assert _MANTLE_WORKSPACE_HEADER not in headers
def test_mantle_messages_validate_environment_sets_workspace_header():
@ -370,7 +374,7 @@ def test_mantle_messages_validate_environment_sets_workspace_header():
litellm_params={"aws_bedrock_project_id": "proj_abc123def456"},
api_base="https://bedrock-mantle.us-east-1.api.aws/anthropic/v1/messages",
)
assert headers["anthropic-workspace"] == "proj_abc123def456"
assert headers[_MANTLE_WORKSPACE_HEADER] == "proj_abc123def456"
assert api_base == "https://bedrock-mantle.us-east-1.api.aws/anthropic/v1/messages"
@ -383,7 +387,7 @@ def test_mantle_messages_validate_environment_without_project_id():
optional_params={},
litellm_params={},
)
assert "anthropic-workspace" not in headers
assert _MANTLE_WORKSPACE_HEADER not in headers
def test_mantle_completion_sends_workspace_header_and_clean_body():
@ -409,7 +413,7 @@ def test_mantle_completion_sends_workspace_header_and_clean_body():
assert response.choices[0].message.content == "ok"
assert len(requests) == 1
assert requests[0]["path"] == "/anthropic/v1/messages"
assert requests[0]["headers"]["anthropic-workspace"] == "proj_abc123def456"
assert requests[0]["headers"][_MANTLE_WORKSPACE_HEADER] == "proj_abc123def456"
assert "aws_bedrock_project_id" not in requests[0]["body"]
@ -443,7 +447,7 @@ async def test_mantle_anthropic_messages_sends_workspace_header_and_clean_body()
assert response["content"][0]["text"] == "ok"
assert len(requests) == 1
assert requests[0]["path"] == "/anthropic/v1/messages"
assert requests[0]["headers"]["anthropic-workspace"] == "proj_abc123def456"
assert requests[0]["headers"][_MANTLE_WORKSPACE_HEADER] == "proj_abc123def456"
assert "aws_bedrock_project_id" not in requests[0]["body"]

View file

@ -193,7 +193,8 @@ class TestEnvironment:
assert "anthropic-version" not in merged
def test_project_id_becomes_the_workspace_header(self):
assert self._validate({}, {"aws_bedrock_project_id": "proj_123"})["anthropic-workspace"] == "proj_123"
# header name from https://docs.aws.amazon.com/bedrock/latest/userguide/workspaces.html, checked 2026-10-02
assert self._validate({}, {"aws_bedrock_project_id": "proj_123"})["anthropic-workspace-id"] == "proj_123"
class TestRequestBody:

View file

@ -0,0 +1,136 @@
import json
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
import respx
import litellm
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
SCALEWAY_RERANK_BODY = {
"id": "rerank-a89e6d7b8b97492ea81569c65fbfff49",
"model": "qwen3-embedding-8b",
"usage": {"total_tokens": 99},
"results": [
{
"index": 1,
"document": {"text": "Oceans can be sorted by size: Pacific, Atlantic, Indian", "multi_modal": None},
"relevance_score": 0.6456239223480225,
},
{
"index": 0,
"document": {"text": "The Pacific is approximately 165 million km²", "multi_modal": None},
"relevance_score": 0.6059925556182861,
},
],
}
DOCUMENTS = ["The Pacific is approximately 165 million km²", "Oceans can be sorted by size: Pacific, Atlantic, Indian"]
def test_scaleway_rerank_posts_to_the_documented_endpoint(respx_mock: respx.MockRouter, monkeypatch):
monkeypatch.delenv("SCALEWAY_API_BASE", raising=False)
route = respx_mock.post("https://api.scaleway.ai/v1/rerank")
route.return_value = httpx.Response(200, json=SCALEWAY_RERANK_BODY)
response = litellm.rerank(
model="scaleway/qwen3-embedding-8b",
query="What is the biggest area of water on earth ?",
documents=DOCUMENTS,
top_n=2,
api_key="scw-key",
)
request = route.calls[0].request
assert request.headers["authorization"] == "Bearer scw-key"
assert json.loads(request.content) == {
"model": "qwen3-embedding-8b",
"query": "What is the biggest area of water on earth ?",
"documents": DOCUMENTS,
"top_n": 2,
}
assert [r["index"] for r in response.results] == [1, 0]
assert response.results[0]["relevance_score"] == pytest.approx(0.6456239223480225)
assert response.results[0]["document"]["text"].startswith("Oceans")
assert response.id == SCALEWAY_RERANK_BODY["id"]
assert response.meta["billed_units"]["total_tokens"] == 99
def test_scaleway_rerank_reads_the_key_from_scw_secret_key(respx_mock: respx.MockRouter, monkeypatch):
monkeypatch.setenv("SCW_SECRET_KEY", "env-scw-key")
route = respx_mock.post("https://api.scaleway.ai/v1/rerank")
route.return_value = httpx.Response(200, json=SCALEWAY_RERANK_BODY)
litellm.rerank(model="scaleway/qwen3-embedding-8b", query="q", documents=DOCUMENTS)
assert route.calls[0].request.headers["authorization"] == "Bearer env-scw-key"
def test_scaleway_rerank_honors_api_base(respx_mock: respx.MockRouter):
route = respx_mock.post("https://scw.example/v1/rerank")
route.return_value = httpx.Response(200, json=SCALEWAY_RERANK_BODY)
litellm.rerank(
model="scaleway/qwen3-embedding-8b",
query="q",
documents=DOCUMENTS,
api_key="scw-key",
api_base="https://scw.example/v1/",
)
assert route.called
def test_scaleway_rerank_does_not_send_return_documents(respx_mock: respx.MockRouter):
"""The Scaleway API has no such field, so it must not reach the request body."""
route = respx_mock.post("https://api.scaleway.ai/v1/rerank")
route.return_value = httpx.Response(200, json=SCALEWAY_RERANK_BODY)
litellm.rerank(
model="scaleway/qwen3-embedding-8b",
query="q",
documents=DOCUMENTS,
return_documents=True,
api_key="scw-key",
)
assert "return_documents" not in json.loads(route.calls[0].request.content)
def test_scaleway_rerank_without_a_key_names_the_env_var(monkeypatch):
monkeypatch.delenv("SCW_SECRET_KEY", raising=False)
with pytest.raises(litellm.APIConnectionError, match="SCW_SECRET_KEY"):
litellm.rerank(model="scaleway/qwen3-embedding-8b", query="q", documents=DOCUMENTS)
def test_scaleway_rerank_caller_headers_cannot_replace_the_provider_key(respx_mock: respx.MockRouter):
route = respx_mock.post("https://api.scaleway.ai/v1/rerank")
route.return_value = httpx.Response(200, json=SCALEWAY_RERANK_BODY)
litellm.rerank(
model="scaleway/qwen3-embedding-8b",
query="q",
documents=DOCUMENTS,
api_key="scw-key",
headers={"Authorization": "Bearer caller-key", "x-trace": "abc"},
)
request = route.calls[0].request
assert request.headers["authorization"] == "Bearer scw-key"
assert request.headers["x-trace"] == "abc"
@pytest.mark.asyncio
async def test_scaleway_arerank_posts_to_the_documented_endpoint():
client = MagicMock(spec=AsyncHTTPHandler)
client.post = AsyncMock(return_value=httpx.Response(200, json=SCALEWAY_RERANK_BODY))
response = await litellm.arerank(
model="scaleway/qwen3-embedding-8b", query="q", documents=DOCUMENTS, api_key="scw-key", client=client
)
assert client.post.await_args.kwargs["url"] == "https://api.scaleway.ai/v1/rerank"
assert client.post.await_args.kwargs["headers"]["authorization"] == "Bearer scw-key"
assert [r["index"] for r in response.results] == [1, 0]

View file

@ -111,7 +111,6 @@ class TestBuildTransaction:
[
{"status": "failure"},
{"api_key": ""},
{"session_id": None},
{"model": ""},
{"startTime": "not-a-time"},
],
@ -119,6 +118,12 @@ class TestBuildTransaction:
def test_incomplete_payloads_are_skipped(self, payload_overrides: dict):
assert _build(payload=_payload(**payload_overrides)) is None
@pytest.mark.parametrize("session_id", [None, ""])
def test_a_request_without_a_session_keeps_its_router_day_money(self, session_id: str | None) -> None:
transaction: Final = _build(payload=_payload(session_id=session_id))
assert transaction is not None
assert (transaction.session_id, transaction.router_name, transaction.spend) == ("", "live-auto", 0.01)
@pytest.mark.parametrize("metadata", [{}, {"routing_decision": None}, {"routing_decision": {}}])
def test_requests_without_a_routing_decision_are_skipped(self, metadata: dict):
assert _build(metadata=metadata) is None

View file

@ -289,6 +289,65 @@ async def test_update_database_skips_tool_usage_when_spend_logs_disabled():
assert prisma.tool_usage_transactions == []
@pytest.mark.asyncio
@pytest.mark.parametrize("disable_spend_logs", [True, False])
@pytest.mark.parametrize("session_id", ["session-1", None])
async def test_a_routed_request_reaches_the_auto_router_rollup_whether_or_not_spend_logs_are_kept(
disable_spend_logs: bool, session_id: str | None
) -> None:
db_writer = DBSpendUpdateWriter()
db_writer._insert_spend_log_to_db = AsyncMock()
db_writer._batch_database_updates = AsyncMock()
prisma = _tool_usage_prisma()
prisma.autorouter_turn_transactions = []
prisma._autorouter_turn_transactions_lock = asyncio.Lock()
routed_payload: Final = {
**_minimal_spend_payload(),
"status": "success",
"api_key": "hashed-key",
"user": "u1",
"session_id": session_id,
"model": "claude-haiku-4-5",
"model_group": "smart-router",
"spend": 0.25,
"startTime": "2026-07-25T10:00:00+00:00",
"metadata": json.dumps(
{
"routing_decision": {"router_model_name": "smart-router", "router_type": "complexity"},
"autorouter_savings": 1.5,
}
),
}
with (
patch("litellm.proxy.proxy_server.disable_spend_logs", disable_spend_logs), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"),
patch(
"litellm.proxy.spend_tracking.spend_tracking_utils.get_logging_payload",
return_value=routed_payload,
),
):
await db_writer.update_database(
token="test-token",
user_id="u1",
end_user_id=None,
team_id=None,
org_id=None,
kwargs={"model": "smart-router"},
completion_response=_tool_call_response("get_weather"),
start_time=datetime.now(timezone.utc),
end_time=datetime.now(timezone.utc),
response_cost=0.25,
)
(turn,) = prisma.autorouter_turn_transactions
stored_session: Final = session_id if session_id and not disable_spend_logs else ""
assert (turn.router_name, turn.router_type, turn.session_id) == ("smart-router", "complexity", stored_session)
assert (turn.spend, turn.saved_spend) == (0.25, 1.5)
assert (prisma.tool_usage_transactions == []) is disable_spend_logs
Statement = tuple[str, tuple[object, ...]]

View file

@ -1,5 +1,6 @@
"""Tests for unified guardrail."""
import io
import logging
from types import SimpleNamespace
from typing import TYPE_CHECKING, Final, Literal
@ -373,6 +374,37 @@ class TestUnifiedLLMGuardrails:
assert result["prompt"] == "a paper boat on a stream [GUARDRAILED]"
assert result["seconds"] == "4"
@pytest.mark.asyncio
@pytest.mark.parametrize("call_type", ["aimage_edit", "image_edit"])
async def test_image_edit_routes_scan_prompt_and_keep_rewrite(self, monkeypatch, call_type: str) -> None:
"""/v1/images/edits dispatches call_type="aimage_edit", which had no translation mapping,
so the hook returned the request unscanned. Runs against the discovered handler map."""
_patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings())
handler = UnifiedLLMGuardrails()
guardrail = RewritingGuardrail()
image = io.BytesIO(b"\x89PNG\r\n\x1a\n")
data = {
"guardrail_to_apply": guardrail,
"model": "gemini-3-pro-image",
"prompt": "a watercolor painting of a lighthouse",
"image": [image],
}
result = await handler.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
cache=DualCache(),
data=data,
call_type=call_type,
)
assert guardrail.event_history == [GuardrailEventHooks.pre_call]
assert [call["inputs"]["texts"] for call in guardrail.apply_calls] == [
["a watercolor painting of a lighthouse"]
]
assert guardrail.apply_calls[0]["inputs"]["model"] == "gemini-3-pro-image"
assert result["prompt"] == "a watercolor painting of a lighthouse [GUARDRAILED]"
assert result["image"] == [image]
class TestAsyncModerationHook:
@pytest.mark.asyncio
async def test_uses_mcp_event_type(self):
@ -419,6 +451,29 @@ class TestUnifiedLLMGuardrails:
assert guardrail.event_history == [GuardrailEventHooks.during_call]
@pytest.mark.asyncio
async def test_runs_for_image_edits(self, monkeypatch) -> None:
_patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings())
handler = UnifiedLLMGuardrails()
guardrail = RecordingGuardrail()
data = {
"guardrail_to_apply": guardrail,
"model": "gemini-3-pro-image",
"prompt": "a watercolor painting of a lighthouse",
"image": [io.BytesIO(b"\x89PNG\r\n\x1a\n")],
}
await handler.async_moderation_hook(
data=data,
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
call_type=CallTypes.aimage_edit.value,
)
assert guardrail.event_history == [GuardrailEventHooks.during_call]
assert [call["inputs"]["texts"] for call in guardrail.apply_calls] == [
["a watercolor painting of a lighthouse"]
]
class TestAsyncPostCallStreamingIteratorHook:
@pytest.mark.asyncio
async def test_streaming_content_not_lost_on_sampled_chunks(self, monkeypatch):

View file

@ -3393,6 +3393,72 @@ def test_openai_passthrough_forwards_verbatim_to_openai(
assert route.calls.last.request.headers["authorization"] == "Bearer sk-upstream"
@pytest.fixture
def openai_wif_env(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None:
from litellm.llms.openai.workload_identity import _workload_identity_auth
token_file: Final = tmp_path / "subject_token.jwt"
token_file.write_text("subject-token-from-file")
monkeypatch.delenv("OPENAI_API_BASE", raising=False)
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
monkeypatch.setattr(litellm, "api_base", None)
monkeypatch.setenv("OPENAI_IDENTITY_PROVIDER_ID", "idp_test123")
monkeypatch.setenv("OPENAI_SERVICE_ACCOUNT_ID", "user-test456")
monkeypatch.setenv("OPENAI_IDENTITY_TOKEN_FILE", str(token_file))
_workload_identity_auth.cache_clear()
@pytest.mark.parametrize("static_key", [None, "", " "])
def test_openai_passthrough_uses_workload_identity_token_without_static_key(
openai_passthrough_client: TestClient,
openai_wif_env: None,
monkeypatch: pytest.MonkeyPatch,
static_key: str | None,
) -> None:
if static_key is None:
monkeypatch.delenv("OPENAI_API_KEY")
else:
monkeypatch.setenv("OPENAI_API_KEY", static_key)
with respx.mock(assert_all_called=True) as upstream:
token_exchange = upstream.post("https://auth.openai.com/oauth/token").mock(
return_value=httpx.Response(200, json={"access_token": "wif-bearer", "expires_in": 3600})
)
route = upstream.post("https://api.openai.com/v1/responses").mock(
return_value=httpx.Response(200, json={"id": "upstream_123"})
)
response = openai_passthrough_client.post(
"/openai_passthrough/v1/responses", json={"model": "gpt-5.1", "input": "hi"}
)
assert (response.status_code, response.json()) == (200, {"id": "upstream_123"})
assert route.calls.last.request.headers["authorization"] == "Bearer wif-bearer"
assert json.loads(token_exchange.calls.last.request.content)["subject_token"] == "subject-token-from-file"
@pytest.mark.asyncio
async def test_openai_passthrough_never_sends_workload_identity_token_to_foreign_api_base(
openai_wif_env: None, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
monkeypatch.setenv("OPENAI_API_BASE", "https://my-vllm.internal/")
monkeypatch.setenv("OPENAI_BASE_URL", "https://api.openai.com/v1")
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value=None,
),
respx.mock(assert_all_mocked=True) as upstream,
pytest.raises(Exception, match="Required 'OPENAI_API_KEY'"),
):
await openai_proxy_route(
endpoint="v1/responses",
request=MagicMock(spec=Request),
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=MagicMock(),
)
assert upstream.calls.call_count == 0
class TestCursorProxyRoute:
"""Tests for the Cursor Cloud Agents pass-through route."""

View file

@ -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

View file

@ -7,9 +7,13 @@ from types import MappingProxyType, SimpleNamespace
from typing import Final
from unittest.mock import patch
import httpx
import pytest
import respx
from starlette.routing import WebSocketRoute
import litellm
from litellm.llms.openai.workload_identity import _workload_identity_auth
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
_OPENAI_WS_DISABLED_REFUSAL,
@ -174,6 +178,65 @@ async def test_openai_websocket_accepts_first_client_subprotocol():
assert websocket.closed is None
TOKEN_EXCHANGE_URL: Final = "https://auth.openai.com/oauth/token"
@pytest.fixture
def openai_wif_token_file(monkeypatch, tmp_path):
token_file = tmp_path / "subject_token.jwt"
token_file.write_text("subject-token-from-file")
monkeypatch.delenv("OPENAI_API_BASE", raising=False)
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
monkeypatch.setattr(litellm, "api_base", None)
monkeypatch.setenv("OPENAI_IDENTITY_PROVIDER_ID", "idp_test123")
monkeypatch.setenv("OPENAI_SERVICE_ACCOUNT_ID", "user-test456")
monkeypatch.setenv("OPENAI_IDENTITY_TOKEN_FILE", str(token_file))
_workload_identity_auth.cache_clear()
return token_file
@pytest.mark.asyncio
async def test_openai_websocket_uses_workload_identity_token_without_static_key(openai_wif_token_file):
websocket = _FakeWebSocket("/openai_passthrough/v1/realtime", "model=gpt-realtime")
with patch(GET_CREDENTIALS, return_value=None), respx.mock(assert_all_called=True) as upstream:
upstream.post(TOKEN_EXCHANGE_URL).mock(
return_value=httpx.Response(200, json={"access_token": "wif-bearer", "expires_in": 3600})
)
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), ENABLED)
assert [call.custom_headers for call in served.relay.calls] == [
MappingProxyType({"Authorization": "Bearer wif-bearer"})
]
assert websocket.closed is None
@pytest.mark.asyncio
@pytest.mark.parametrize(
"subject_token_present, exchange_outcome",
[
(True, httpx.Response(401, json={"error": "invalid_grant"})),
(True, httpx.ConnectError("auth.openai.com unreachable")),
(False, httpx.Response(200, json={"access_token": "wif-bearer", "expires_in": 3600})),
],
ids=["rejected", "unreachable", "missing_subject_token"],
)
async def test_openai_websocket_closes_cleanly_when_workload_identity_exchange_fails(
openai_wif_token_file, subject_token_present, exchange_outcome
):
if not subject_token_present:
openai_wif_token_file.unlink()
websocket = _FakeWebSocket("/openai_passthrough/v1/realtime", "model=gpt-realtime")
with patch(GET_CREDENTIALS, return_value=None), respx.mock(assert_all_called=False) as upstream:
upstream.post(TOKEN_EXCHANGE_URL).mock(side_effect=exchange_outcome)
served = await _serve(websocket, "v1/realtime", UserAPIKeyAuth(), ENABLED)
assert websocket.closed == (1011, "OpenAI workload identity token exchange failed")
assert websocket.accepts == []
assert served.relay.calls == []
@pytest.mark.asyncio
async def test_openai_websocket_closes_cleanly_when_provider_credentials_missing():
websocket = _FakeWebSocket("/openai/v1/realtime", "model=gpt-4o-realtime-preview")

View file

@ -1,3 +1,5 @@
import base64
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -5,6 +7,12 @@ from fastapi import HTTPException, Request, Response
import litellm
from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable, UserAPIKeyAuth
from litellm.types.utils import SpecialEnums
from litellm.types.vector_store_files import (
VectorStoreFileListResponse,
VectorStoreFileObject,
VectorStoreFileStatus,
)
def _mock_request() -> MagicMock:
@ -107,18 +115,54 @@ async def test_vector_store_file_create_forces_path_id_over_body_id():
@pytest.mark.asyncio
async def test_vector_store_file_list_resolves_managed_vector_store_before_team_fallback():
import base64
async def test_vector_store_file_list_resolves_managed_ids_and_cursors():
from litellm.proxy.vector_store_files_endpoints.endpoints import (
vector_store_file_list,
)
captured_data = {}
provider_file_id: Final = "file-list-owned"
managed_file_data: Final = (
SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
"application/json",
"unified-file",
"managed-deployment",
provider_file_id,
"managed-deployment-id",
)
)
managed_file_id: Final = (
base64.urlsafe_b64encode(managed_file_data.encode()).decode().rstrip("=")
)
user_api_key_dict: Final = UserAPIKeyAuth(team_models=["team-openai"])
managed_file: Final[VectorStoreFileObject] = {
"id": provider_file_id,
"object": "vector_store.file",
"created_at": 1700000000,
"usage_bytes": 100,
"vector_store_id": "vs_provider_native",
"status": VectorStoreFileStatus.COMPLETED,
"last_error": None,
"chunking_strategy": {"type": "auto"},
"attributes": {"source": "test"},
}
provider_response: Final[VectorStoreFileListResponse] = {
"object": "list",
"data": [managed_file],
"first_id": provider_file_id,
"last_id": provider_file_id,
"has_more": False,
}
expected_response: Final[VectorStoreFileListResponse] = {
**provider_response,
"data": [{**managed_file, "id": managed_file_id}],
"first_id": managed_file_id,
"last_id": managed_file_id,
}
async def fake_base_process(self, **kwargs):
captured_data.update(self.data)
return {"ok": True}
return provider_response
raw_vector_store_id = (
"litellm_proxy:vector_store;"
@ -133,7 +177,7 @@ async def test_vector_store_file_list_resolves_managed_vector_store_before_team_
request = _mock_request()
request.method = "GET"
request.query_params = {"limit": "10"}
request.query_params = {"after": managed_file_id, "limit": "10"}
request.url.path = f"/v1/vector_stores/{vector_store_id}/files"
llm_router = MagicMock()
@ -147,6 +191,11 @@ async def test_vector_store_file_list_resolves_managed_vector_store_before_team_
}
llm_router.get_deployment_credentials_with_provider.side_effect = get_credentials
managed_files_obj = MagicMock()
resolver = AsyncMock(return_value={provider_file_id: managed_file_id})
managed_files_obj.get_unified_file_ids_for_provider_file_ids = resolver
proxy_logging_obj = MagicMock()
proxy_logging_obj.get_proxy_hook.return_value = managed_files_obj
with (
patch(
@ -154,6 +203,7 @@ async def test_vector_store_file_list_resolves_managed_vector_store_before_team_
new=AsyncMock(return_value=None),
),
patch("litellm.proxy.proxy_server.llm_router", llm_router),
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj),
patch(
"litellm.proxy.vector_store_files_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request",
new=fake_base_process,
@ -163,16 +213,22 @@ async def test_vector_store_file_list_resolves_managed_vector_store_before_team_
vector_store_id=vector_store_id,
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(team_models=["team-openai"]),
user_api_key_dict=user_api_key_dict,
)
assert response == {"ok": True}
assert response == expected_response
assert captured_data["after"] == provider_file_id
assert captured_data["vector_store_id"] == "vs_provider_native"
assert captured_data["api_key"] == "sk-managed-deployment"
assert captured_data["model"] == "openai/managed-deployment"
llm_router.get_deployment_credentials_with_provider.assert_called_once_with(
model_id="managed-deployment"
)
proxy_logging_obj.get_proxy_hook.assert_called_once_with("managed_files")
resolver.assert_awaited_once_with(
provider_file_ids=(provider_file_id,),
user_api_key_dict=user_api_key_dict,
)
@pytest.mark.asyncio

View file

@ -10,9 +10,11 @@ is attached to a vector store or read back under shared provider credentials.
"""
import base64
from collections.abc import Mapping, Sequence
from copy import deepcopy
from dataclasses import dataclass
from typing import Literal
from unittest.mock import MagicMock, patch
from typing import Final, Literal
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -23,8 +25,15 @@ import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.vector_store_files_endpoints.endpoints import (
_update_request_data_with_managed_file_id,
_with_managed_file_list_ids,
_with_provider_file_id_cursors,
)
from litellm.types.utils import SpecialEnums
from litellm.types.vector_store_files import (
VectorStoreFileListResponse,
VectorStoreFileObject,
VectorStoreFileStatus,
)
RAW_FILE_ID = "file-victim-abc123"
CALLER = UserAPIKeyAuth(api_key="sk-test", user_id="attacker-user", team_id="team-b")
@ -51,13 +60,46 @@ class ManagedResourceAccessCheckerStub:
return False
def _unified_file_id() -> str:
@dataclass(frozen=True)
class ManagedFileIdResolverStub:
resolver: AsyncMock
async def get_unified_file_ids_for_provider_file_ids(
self,
provider_file_ids: Sequence[str],
user_api_key_dict: UserAPIKeyAuth,
) -> Mapping[str, str]:
return await self.resolver(
provider_file_ids=provider_file_ids,
user_api_key_dict=user_api_key_dict,
)
def _unified_file_id(provider_file_id: str = RAW_FILE_ID) -> str:
unified = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format(
"application/json", "victim-unified-id", "gpt-4o-mini", RAW_FILE_ID, "gpt-4o-mini-id"
"application/json",
"victim-unified-id",
"gpt-4o-mini",
provider_file_id,
"gpt-4o-mini-id",
)
return base64.urlsafe_b64encode(unified.encode()).decode().rstrip("=")
def _vector_store_file_row(file_id: str) -> VectorStoreFileObject:
return {
"id": file_id,
"object": "vector_store.file",
"created_at": 1700000000,
"usage_bytes": 100,
"vector_store_id": "vs-test",
"status": VectorStoreFileStatus.COMPLETED,
"last_error": None,
"chunking_strategy": {"type": "auto"},
"attributes": {"source": "test"},
}
async def _resolve(
file_id: str,
file_access: Literal["allow", "deny", "missing"] = "allow",
@ -72,6 +114,113 @@ async def _resolve(
)
@pytest.mark.parametrize(
"provider_ids",
[
(RAW_FILE_ID, "file-unmanaged-123"),
("file-unmanaged-123", RAW_FILE_ID),
],
)
@pytest.mark.asyncio
async def test_vector_store_file_list_maps_owned_ids_and_preserves_raw_ids(
provider_ids: tuple[str, str],
) -> None:
managed_file_id: Final = _unified_file_id()
expected_provider_ids: Final = tuple(
managed_file_id if provider_file_id == RAW_FILE_ID else provider_file_id
for provider_file_id in provider_ids
)
provider_response: Final[VectorStoreFileListResponse] = {
"object": "list",
"data": [
_vector_store_file_row(provider_file_id)
for provider_file_id in provider_ids
],
"first_id": provider_ids[0],
"last_id": provider_ids[1],
"has_more": True,
}
original_response: Final = deepcopy(provider_response)
resolver: Final = AsyncMock(return_value={RAW_FILE_ID: managed_file_id})
managed_files_obj: Final = ManagedFileIdResolverStub(resolver=resolver)
response: Final = await _with_managed_file_list_ids(
response=provider_response,
managed_files_obj=managed_files_obj,
user_api_key_dict=CALLER,
)
expected_response: Final[VectorStoreFileListResponse] = {
"object": "list",
"data": [
_vector_store_file_row(provider_file_id)
for provider_file_id in expected_provider_ids
],
"first_id": expected_provider_ids[0],
"last_id": expected_provider_ids[1],
"has_more": True,
}
assert response == expected_response
assert provider_response == original_response
resolver.assert_awaited_once_with(
provider_file_ids=tuple(dict.fromkeys(provider_ids)),
user_api_key_dict=CALLER,
)
@pytest.mark.asyncio
async def test_vector_store_file_list_only_maps_round_trippable_ids() -> None:
managed_file_id: Final = _unified_file_id("file-model-a")
provider_response: Final[VectorStoreFileListResponse] = {
"object": "list",
"data": [
_vector_store_file_row("file-model-a"),
_vector_store_file_row("file-model-b"),
],
"first_id": "file-model-a",
"last_id": "file-model-b",
"has_more": False,
}
resolver: Final = AsyncMock(
return_value={
"file-model-a": managed_file_id,
"file-model-b": managed_file_id,
}
)
managed_files_obj: Final = ManagedFileIdResolverStub(resolver=resolver)
response: Final = await _with_managed_file_list_ids(
response=provider_response,
managed_files_obj=managed_files_obj,
user_api_key_dict=CALLER,
)
expected_response: Final[VectorStoreFileListResponse] = {
"object": "list",
"data": [
_vector_store_file_row(managed_file_id),
_vector_store_file_row("file-model-b"),
],
"first_id": managed_file_id,
"last_id": "file-model-b",
"has_more": False,
}
assert response == expected_response
def test_vector_store_file_list_translates_managed_cursors_and_preserves_raw_after() -> (
None
):
managed_file_id: Final = _unified_file_id()
assert _with_provider_file_id_cursors(
{"after": managed_file_id, "before": managed_file_id}
) == {"after": RAW_FILE_ID, "before": RAW_FILE_ID}
assert _with_provider_file_id_cursors({"after": RAW_FILE_ID}) == {
"after": RAW_FILE_ID
}
@pytest.mark.asyncio
async def test_raw_file_id_rejected_when_managed_files_required():
with patch.object(litellm, "require_managed_files", True):

View file

@ -0,0 +1,36 @@
from typing import Final
from tests.integration.run import select, uncollected
_GROUP: Final = (
"tests/integration/cost_calculation/test_cost_tracking.py",
"tests/integration/cost_calculation/test_rollups.py",
)
_CELL: Final = (
"tests/integration/cost_calculation/test_cost_tracking.py"
"::test_case_bills_expected_cost[perplexity/pplx-decider-v1-27b-decisions]"
)
def test_a_node_id_inside_a_group_file_is_selected_as_written() -> None:
selection: Final = select((_CELL,), _GROUP)
assert selection.nodes == (_CELL,)
assert selection.foreign == ()
def test_a_node_id_outside_the_group_is_foreign_by_its_file() -> None:
foreign: Final = "tests/integration/providers/test_decisions_wire.py::test_key_checks_match_chat"
assert select((foreign, _CELL), _GROUP).foreign == (foreign,)
def test_no_request_selects_every_group_file() -> None:
assert select((), _GROUP).nodes == _GROUP
def test_a_node_id_whose_file_collected_tests_is_not_empty() -> None:
collected: Final = frozenset({_CELL, "tests/integration/cost_calculation/test_cost_tracking.py::test_other"})
assert uncollected((_CELL,), collected) == ()
def test_a_selected_file_that_collected_nothing_is_reported() -> None:
assert uncollected(_GROUP, frozenset({_CELL})) == ("tests/integration/cost_calculation/test_rollups.py",)

View file

@ -5,7 +5,7 @@ import { renderWithProviders } from "@/../tests/test-utils";
import { MonitoringSetup } from "./LensOverview";
import { LensSetup } from "./LensSetup";
import { apiClient } from "@/components/networking";
import type { Settings } from "./lensData";
import { initialWatches, watchChecks, type Settings } from "./lensData";
vi.mock("@/components/networking", () => ({ apiClient: { post: vi.fn(), get: vi.fn() } }));
@ -105,12 +105,13 @@ describe("Lens setup", () => {
fireEvent.change(screen.getByRole("combobox", { name: "Metadata key 1" }), { target: { value: "swarm" } });
fireEvent.change(screen.getByRole("combobox", { name: "Metadata value 1" }), { target: { value: "research" } });
await user.click(screen.getByRole("button", { name: "Continue" }));
await user.click(screen.getByRole("button", { name: /Add your own/ }));
await user.type(screen.getByRole("textbox", { name: "Check 1" }), "Find incomplete reports");
await user.click(screen.getByRole("button", { name: "Add check" }));
await user.click(screen.getByRole("button", { name: /Add your own/ }));
fireEvent.change(screen.getByRole("textbox", { name: "Check 2" }), {
target: { value: "Find repeated searches\nInclude retries that add no information" },
});
await user.click(screen.getByRole("button", { name: "Add check" }));
await user.click(screen.getByRole("button", { name: /Add your own/ }));
await user.click(screen.getByRole("button", { name: "Remove check 3" }));
await user.click(screen.getByRole("button", { name: "Continue" }));
await waitFor(() => expect(screen.getByRole("button", { name: "Run investigation" })).toBeEnabled());
@ -129,6 +130,7 @@ describe("Lens setup", () => {
model: "analysis",
monthly_budget: 100,
checks: [
...watchChecks(initialWatches(undefined)),
expect.objectContaining({ instruction: "Find incomplete reports" }),
expect.objectContaining({ instruction: "Find repeated searches\nInclude retries that add no information" }),
],
@ -411,3 +413,65 @@ it("saves a discovered agent independently of the application name", async () =>
expect.objectContaining({ agent_name: "research_agent", service: "shared-service" }),
);
});
describe("Watch for", () => {
const tile = (name: string) => screen.getByRole("button", { name: new RegExp(`^${name}`) });
it("saves exactly the presets the user toggled, by click and by number key", async () => {
const save = vi.fn().mockResolvedValue(undefined);
const user = userEvent.setup();
renderWithProviders(
<LensSetup models={["analysis"]} defaultModel="analysis" accessToken="test" onClose={vi.fn()} onSave={save} />,
);
await user.click(screen.getByRole("button", { name: "Continue" }));
await user.click(tile("unhappy"));
tile("unsolved").focus();
await user.keyboard("6");
await user.keyboard("{ArrowRight}{ArrowRight}");
expect(tile("unsafe")).toHaveFocus();
expect(tile("unhappy")).toHaveAttribute("aria-pressed", "false");
expect(tile("looping")).toHaveAttribute("aria-pressed", "true");
await user.click(screen.getByRole("button", { name: "Continue" }));
await waitFor(() => expect(screen.getByRole("button", { name: "Run investigation" })).toBeEnabled());
await user.click(screen.getByRole("button", { name: "Run investigation" }));
const saved = (save.mock.calls[0][0] as Settings).checks.map((check) => check.id);
expect(saved).toEqual(["watch_unsolved", "watch_blocked", "watch_looping"]);
});
it("keeps an edited investigation's preset choices and custom checks apart", async () => {
const save = vi.fn().mockResolvedValue(undefined);
const user = userEvent.setup();
const initial: Settings = {
...settings,
checks: [{ id: "watch_invented", instruction: "old wording", enabled: true }, settings.checks[1]],
};
renderWithProviders(
<LensSetup initial={initial} models={["analysis"]} accessToken="test" onClose={vi.fn()} onSave={save} />,
);
await user.click(screen.getByRole("button", { name: "Continue" }));
expect(tile("invented")).toHaveAttribute("aria-pressed", "true");
expect(tile("unsolved")).toHaveAttribute("aria-pressed", "false");
expect(screen.getByRole("textbox", { name: "Check 1" })).toHaveValue("Find incomplete reports");
await user.click(screen.getByRole("button", { name: "Continue" }));
await waitFor(() => expect(screen.getByRole("button", { name: "Save changes" })).toBeEnabled());
await user.click(screen.getByRole("button", { name: "Save changes" }));
expect((save.mock.calls[0][0] as Settings).checks).toEqual([
...watchChecks(new Set(["watch_invented"])),
settings.checks[1],
]);
});
it("lets a run start from presets alone and blocks it once nothing is selected", async () => {
const user = userEvent.setup();
renderWithProviders(
<LensSetup models={["analysis"]} defaultModel="analysis" accessToken="test" onClose={vi.fn()} onSave={vi.fn()} />,
);
await user.click(screen.getByRole("button", { name: "Continue" }));
await user.click(screen.getByRole("button", { name: "Continue" }));
expect(await screen.findByRole("button", { name: "Run investigation" })).toBeInTheDocument();
await user.click(screen.getByRole("button", { name: "Back" }));
for (const name of ["unsolved", "blocked", "unhappy"]) await user.click(tile(name));
await user.click(screen.getByRole("button", { name: "Continue" }));
expect(screen.getByRole("alert")).toHaveTextContent("pick something to watch for");
});
});

View file

@ -1,7 +1,7 @@
"use client";
import { useState } from "react";
import { Plus, X } from "lucide-react";
import { X } from "lucide-react";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { Textarea } from "@/components/ui/textarea";
@ -16,7 +16,16 @@ import {
import { SearchSelect } from "@/components/shared/SearchSelect";
import { DurationInput } from "./DurationInput";
import { ActivityScope, type ActivitySelection } from "./ActivityScope";
import { analysisModelOptions, normalizeFilters, type AnalysisModelInfo, type Settings } from "./lensData";
import { WatchPicker } from "./WatchPicker";
import {
analysisModelOptions,
initialWatches,
isWatch,
normalizeFilters,
watchChecks,
type AnalysisModelInfo,
type Settings,
} from "./lensData";
function validateSample(selection: ActivitySelection) {
const hours = selection.lookback_hours ?? 24;
@ -77,7 +86,8 @@ export function LensSetup({
};
const [selection, setSelection] = useState(initialSelection);
const [context, setContext] = useState(initial?.context ?? "");
const [questions, setQuestions] = useState(() => (initial?.checks?.length ? initial.checks : [newCheck()]));
const [watching, setWatching] = useState<ReadonlySet<string>>(() => initialWatches(initial?.checks));
const [questions, setQuestions] = useState(() => (initial?.checks ?? []).filter((check) => !isWatch(check)));
const [selectedModel, setModel] = useState<string | null>(initial?.model ?? null);
const model = selectedModel ?? defaultModel ?? "";
const [budget, setBudget] = useState(initial?.monthly_budget ?? 100);
@ -102,8 +112,8 @@ export function LensSetup({
if (step >= 2 && manualSelection && !selection.execution_ids?.length)
throw new Error("Choose at least one run or turn off individual selection");
validateSample(selection);
if (step >= 1 && !context.trim() && !filledChecks.length)
throw new Error("Describe the expected behavior or what to look out for");
const nothingToCheck = !context.trim() && !filledChecks.length && !watching.size;
if (step >= 1 && nothingToCheck) throw new Error("Describe the expected behavior or pick something to watch for");
if (filledChecks.some((check) => check.instruction.trim().length < 3))
throw new Error("Use at least three characters for each check");
};
@ -132,7 +142,10 @@ export function LensSetup({
interval_minutes: interval,
concurrency: initial?.concurrency ?? 8,
filters: normalizeFilters(selection.filters ?? []),
checks: filledChecks.map((check) => ({ ...check, instruction: check.instruction.trim() })),
checks: [
...watchChecks(watching),
...filledChecks.map((check) => ({ ...check, instruction: check.instruction.trim() })),
],
};
await onSave(settings);
} catch (cause) {
@ -172,7 +185,7 @@ export function LensSetup({
}}
>
<DialogContent
className={`flex max-h-[90dvh] flex-col gap-6 overflow-hidden ${step === 2 ? "sm:max-w-3xl" : "sm:max-w-xl"}`}
className={`flex max-h-[90dvh] flex-col gap-6 overflow-hidden ${["sm:max-w-xl", "sm:max-w-2xl", "sm:max-w-3xl"][step]}`}
>
<DialogHeader>
<DialogTitle className="text-xl">{headings[step]}</DialogTitle>
@ -239,8 +252,13 @@ export function LensSetup({
placeholder="Answer the customer's question using verified sources and explain when information is missing."
/>
</label>
<fieldset className="space-y-3">
<legend className="mb-2 text-sm font-medium">What should we look out for?</legend>
<WatchPicker
selected={watching}
onChange={setWatching}
onAddCustom={() => setQuestions([...questions, newCheck()])}
/>
<fieldset className="space-y-2">
<legend className="sr-only">Custom checks</legend>
{questions.map((check, index) => (
<div key={check.id} className="flex items-start gap-2">
<Textarea
@ -254,7 +272,7 @@ export function LensSetup({
)
}
rows={2}
placeholder="e.g. Repeated searches that add no useful information"
placeholder="e.g. Quotes a price without checking the pricing tool"
/>
<Button
variant="ghost"
@ -266,29 +284,7 @@ export function LensSetup({
</Button>
</div>
))}
<Button variant="outline" size="sm" onClick={() => setQuestions([...questions, newCheck()])}>
<Plus className="size-3.5" /> Add check
</Button>
</fieldset>
{(!context.trim() || !filledChecks.length) && (
<Button
variant="link"
className="h-auto px-0"
onClick={() => {
if (!context.trim())
setContext(
"Answer the user's question using verified sources. Explain when information is missing.",
);
if (!filledChecks.length)
setQuestions([
newCheck("Find repeated work that adds no useful information."),
newCheck("Find claims that contradict the available evidence."),
]);
}}
>
Use an example
</Button>
)}
</>
)}
{step === 2 && (

View file

@ -51,6 +51,7 @@ import {
type Finding,
type Settings,
type Job,
watches,
} from "./lensData";
const money = (n: number) =>
@ -581,13 +582,13 @@ export function LensView({
</div>
)}
<div className="pt-2">
<h3 className="text-base font-semibold">What should we look out for?</h3>
<h3 className="text-base font-semibold">Watch for</h3>
<p className="mt-1 text-xs text-muted-foreground">Specific problems or patterns to investigate.</p>
</div>
{batchSettings?.checks.map((c, index) => (
<div key={c.id} className="flex items-start gap-3 border-b py-4">
<span className="mt-0.5 text-xs tabular-nums text-muted-foreground">{index + 1}.</span>
<p className="text-sm leading-6 flex-1">{c.instruction}</p>
<CheckSummary check={c} />
{!c.enabled && <span className="text-xs text-muted-foreground">Disabled</span>}
</div>
))}
@ -771,3 +772,14 @@ export function LensView({
</section>
);
}
function CheckSummary({ check }: { check: Settings["checks"][number] }) {
const watch = watches.find((item) => item.id === check.id);
if (!watch) return <p className="flex-1 text-sm leading-6">{check.instruction}</p>;
return (
<p className="grid flex-1 gap-0.5">
<span className="text-sm font-medium">{watch.name}</span>
<span className="text-xs text-muted-foreground">{watch.summary}</span>
</p>
);
}

View file

@ -0,0 +1,177 @@
"use client";
import { useEffect, useRef, useState, type KeyboardEvent } from "react";
import { watches } from "./lensData";
const dotColors = ["#8b5cf6", "#22b3e8", "#e3a32b", "#eb6b93", "#22b3e8", "#8b5cf6", "#e3a32b", "#eb6b93"];
const lensBlue = { light: "#0011b3", dark: "#8b9bff" };
const columns = 120;
const rows = 7;
const cell = 6;
function dotColor(lit: boolean, pastLens: boolean, incoming: string, blue: string): string {
if (!lit) return "#94a3b8";
return pastLens ? blue : incoming;
}
function DotFlow({ active }: { active: readonly string[] }) {
const canvas = useRef<HTMLCanvasElement>(null);
useEffect(() => {
const node = canvas.current;
const context = node?.getContext("2d");
if (!node || !context) return;
const colors = active.length ? active : ["#94a3b8"];
const still = window.matchMedia("(prefers-reduced-motion: reduce)").matches;
const blue = document.documentElement.classList.contains("dark") ? lensBlue.dark : lensBlue.light;
const lensColumn = Math.floor(columns * 0.62);
const draw = (time: number) => {
context.clearRect(0, 0, node.width, node.height);
for (let row = 0; row < rows; row++) {
for (let column = 0; column < columns; column++) {
const x = column * cell + cell / 2;
const y = row * cell + cell / 2;
const center = (rows - 1) / 2;
const funnel =
column < lensColumn ? Math.abs(row - center) <= center * (1 - column / lensColumn) + 0.6 : row === center;
const wave = Math.sin(column * 0.55 - time / 260 + row * 1.7);
const lit = funnel && wave > 0.35;
context.globalAlpha = lit ? 0.9 : 0.12;
context.fillStyle = dotColor(lit, column >= lensColumn, colors[(row + column) % colors.length], blue);
context.beginPath();
context.arc(x, y, lit ? 1.6 : 1, 0, Math.PI * 2);
context.fill();
}
}
context.globalAlpha = 1;
context.strokeStyle = blue;
context.lineWidth = 1.5;
const lx = lensColumn * cell - 1;
context.beginPath();
context.moveTo(lx + 3, 1);
context.lineTo(lx, 1);
context.lineTo(lx, rows * cell - 1);
context.lineTo(lx + 3, rows * cell - 1);
context.stroke();
};
if (still) {
draw(0);
return;
}
let frame = requestAnimationFrame(function loop(time) {
draw(time);
frame = requestAnimationFrame(loop);
});
return () => cancelAnimationFrame(frame);
}, [active]);
return (
<canvas
ref={canvas}
aria-hidden="true"
width={columns * cell}
height={rows * cell}
className="h-10 w-full opacity-80"
/>
);
}
export function WatchPicker({
selected,
onChange,
onAddCustom,
}: {
selected: ReadonlySet<string>;
onChange: (next: ReadonlySet<string>) => void;
onAddCustom: () => void;
}) {
const [cursor, setCursor] = useState(0);
const items = useRef<(HTMLButtonElement | null)[]>([]);
const toggle = (id: string) =>
onChange(new Set(selected.has(id) ? [...selected].filter((item) => item !== id) : [...selected, id]));
const move = (index: number) => {
const next = (index + watches.length) % watches.length;
setCursor(next);
items.current[next]?.focus();
};
const onKey = (event: KeyboardEvent<HTMLDivElement>) => {
const digit = Number(event.key);
if (event.key === "ArrowRight" || event.key === "l") move(cursor + 1);
else if (event.key === "ArrowLeft" || event.key === "h") move(cursor - 1);
else if (event.key === "ArrowDown" || event.key === "j") move(cursor + 4);
else if (event.key === "ArrowUp" || event.key === "k") move(cursor - 4);
else if (digit >= 1 && digit <= watches.length) {
move(digit - 1);
toggle(watches[digit - 1].id);
} else return;
event.preventDefault();
};
const activeColors = watches.flatMap((watch, index) => (selected.has(watch.id) ? [dotColors[index]] : []));
return (
<fieldset className="space-y-2.5">
<div className="flex items-end justify-between gap-3">
<legend className="text-sm font-medium">Watch for</legend>
<span className="text-xs tabular-nums text-muted-foreground">
{selected.size} of {watches.length} selected
</span>
</div>
<DotFlow active={activeColors} />
<div role="group" aria-label="Watch for" onKeyDown={onKey} className="grid grid-cols-2 gap-2.5 sm:grid-cols-4">
{watches.map((watch, index) => {
const on = selected.has(watch.id);
return (
<button
key={watch.id}
ref={(node) => {
items.current[index] = node;
}}
type="button"
aria-pressed={on}
title={watch.summary}
tabIndex={index === cursor ? 0 : -1}
onFocus={() => setCursor(index)}
onClick={() => toggle(watch.id)}
className={`flex h-[5.25rem] flex-col justify-start gap-1 rounded-xl px-3.5 py-3 text-left outline-none transition-all duration-200 ease-out focus-visible:outline-2 focus-visible:outline-offset-2 focus-visible:outline-ring active:scale-[0.97] ${
on
? "bg-background text-foreground ring-[1.5px] ring-inset ring-foreground"
: "bg-muted/60 text-muted-foreground hover:bg-muted hover:text-foreground"
}`}
>
<span className="flex items-center justify-between gap-1">
<span className="text-sm font-medium">{watch.name}</span>
<svg
viewBox="0 0 16 16"
aria-hidden="true"
className={`size-3 transition-all duration-200 ${on ? "scale-100 opacity-100" : "scale-50 opacity-0"}`}
>
<path
d="M3 8.5l3.2 3.2L13 5"
fill="none"
stroke="currentColor"
strokeWidth="2.2"
strokeLinecap="round"
strokeLinejoin="round"
/>
</svg>
</span>
<span className="line-clamp-2 text-xs leading-snug text-muted-foreground">{watch.summary}</span>
</button>
);
})}
</div>
<button
type="button"
onClick={onAddCustom}
className="flex h-11 w-full items-center gap-2.5 rounded-xl bg-muted/60 px-3.5 text-left text-sm outline-none transition-colors hover:bg-muted focus-visible:ring-2 focus-visible:ring-ring focus-visible:ring-offset-2"
>
<span
aria-hidden="true"
className="flex size-5 items-center justify-center rounded-full bg-background text-sm leading-none"
>
+
</span>
<span className="font-medium">Add your own</span>
<span className="text-muted-foreground">describe anything else in plain English</span>
</button>
</fieldset>
);
}

View file

@ -20,12 +20,97 @@ export function scopeLabel(settings: Partial<Pick<Settings, "service" | "agent_n
);
}
export const starterQuestions = [
"Find repeated work or tool calls that add no useful information.",
"Find tool failures or retries that the agent does not recover from.",
"Identify recurring user needs and successful ways the agent handles them.",
export interface Watch {
id: string;
name: string;
summary: string;
instruction: string;
defaultOn: boolean;
}
export const watches: readonly Watch[] = [
{
id: "watch_unsolved",
name: "unsolved",
summary: "didn't finish what was asked",
instruction:
"Find runs where the agent failed to solve what the user asked for: wrong or partial answers, giving up, or stopping mid-task.",
defaultOn: true,
},
{
id: "watch_blocked",
name: "blocked",
summary: "missing a tool, data or skill",
instruction:
"Find runs where the agent could not do a step because it lacked a tool, data or capability, including when it tells the user it cannot help.",
defaultOn: true,
},
{
id: "watch_permissions",
name: "permissions",
summary: "denied, unapproved or overstepped",
instruction:
"Find runs with permission problems: access denied, an approval or confirmation the agent skipped or mishandled, or the agent acting on resources it was not granted.",
defaultOn: false,
},
{
id: "watch_unhappy",
name: "unhappy",
summary: "user annoyed or had to repeat",
instruction:
"Find runs where the user seems dissatisfied: repeating or rephrasing the same request, correcting the agent, or expressing annoyance.",
defaultOn: true,
},
{
id: "watch_swallowed",
name: "swallowed",
summary: "ignored a failed tool call",
instruction:
"Find runs where a tool call failed or returned an error and the agent continued as if it had succeeded, without retrying or telling the user.",
defaultOn: false,
},
{
id: "watch_looping",
name: "looping",
summary: "repeats steps without progress",
instruction:
"Find runs where the agent repeats the same tool call, search or step several times without getting new information or making progress.",
defaultOn: false,
},
{
id: "watch_invented",
name: "invented",
summary: "claims no tool ever returned",
instruction:
"Find runs where the agent states facts, identifiers, numbers or results that do not appear in any tool output or source it had.",
defaultOn: false,
},
{
id: "watch_unsafe",
name: "unsafe",
summary: "harmful, deceptive or rule-bending",
instruction:
"Find runs with malicious or unsafe behavior from the agent or the user: destructive or irreversible actions, deception, leaking secrets or private data, or attempts to bypass instructions or safeguards.",
defaultOn: false,
},
];
export function watchChecks(enabled: ReadonlySet<string>): Settings["checks"] {
return watches
.filter((watch) => enabled.has(watch.id))
.map(({ id, instruction }) => ({ id, instruction, enabled: true }));
}
export function initialWatches(checks: Settings["checks"] | undefined): ReadonlySet<string> {
if (!checks?.length) return new Set(watches.filter((watch) => watch.defaultOn).map((watch) => watch.id));
const ids = new Set(watches.map((watch) => watch.id));
return new Set(checks.filter((check) => ids.has(check.id) && check.enabled).map((check) => check.id));
}
export function isWatch(check: Settings["checks"][number]): boolean {
return watches.some((watch) => watch.id === check.id);
}
export function normalizeFilters(filters: NonNullable<Settings["filters"]>): Settings["filters"] {
return filters.map((f) => {
if (!f.key.trim() || !f.value.trim()) throw new Error("Choose a key and value for every condition, or remove it");