Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/standard-lists-api-d1dc4a

This commit is contained in:
Yuneng Jiang 2026-08-10 15:19:46 -07:00
commit a9857bb362
No known key found for this signature in database
43 changed files with 2641 additions and 1613 deletions

View file

@ -99,19 +99,19 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 45004
"limit": 44996
},
"reportUnknownLambdaType": {
"limit": 113
},
"reportUnknownMemberType": {
"limit": 39649
"limit": 39643
},
"reportUnknownParameterType": {
"limit": 20132
},
"reportUnknownVariableType": {
"limit": 31156
"limit": 31153
},
"reportUnnecessaryCast": {
"limit": 118

View file

@ -1493,6 +1493,8 @@ SPEND_LOG_QUEUE_POLL_INTERVAL: Final = float(os.getenv("SPEND_LOG_QUEUE_POLL_INT
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000))
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute
PROXY_BUDGET_RESCHEDULER_MIN_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597))
RESET_BUDGET_JOB_BATCH_SIZE: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_BATCH_SIZE", "500")))
RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN", "100")))
PROXY_BATCH_POLLING_INTERVAL: Final = int(os.getenv("PROXY_BATCH_POLLING_INTERVAL", 3600))
MAX_OBJECTS_PER_POLL_CYCLE: Final = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50)))
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS: Final = max(1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7)))

View file

@ -3,8 +3,11 @@ Per-feature OpenAPI snapshot for lazy-loaded routers.
The committed JSON is generated by `python -m litellm.proxy._lazy_openapi_snapshot`
and consumed at runtime so /openapi.json can show full route info for unloaded
features without importing them. CI verifies the file is current and surfaces
any drift as a neutral check.
features without importing them. No CI job regenerates this file; drift surfaces
only indirectly through check-ui-api-types.yml, which rebuilds schema.d.ts from
app.openapi() with the committed snapshot injected. After changing any lazily
loaded route or this generator, rerun the module and commit the JSON, then run
`npm run gen:api` in ui/litellm-dashboard and commit schema.d.ts.
"""
import json
@ -89,8 +92,6 @@ def generate_snapshot() -> dict[str, dict]:
from litellm.proxy.proxy_server import app, ensure_unique_openapi_operation_ids
for feat in LAZY_FEATURES:
if feat.module_path in sys.modules:
continue
try:
module = importlib.import_module(feat.module_path)
feat.register_fn(app, module)
@ -100,7 +101,7 @@ def generate_snapshot() -> dict[str, dict]:
fragments: Final[dict[str, dict]] = {}
used_operation_ids: Final[set[str]] = set()
for feat in LAZY_FEATURES:
feat_routes = [r for r in app.routes if any(getattr(r, "path", "").startswith(p) for p in feat.path_prefixes)]
feat_routes = [r for r in app.routes if feat.matches(getattr(r, "path", ""))]
if not feat_routes:
continue
_stabilize_multi_method_route_ids(feat_routes)

View file

@ -1,14 +1,21 @@
import asyncio
import json
import time
from collections.abc import Callable, Sequence
from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Final, Literal, Protocol, TypeVar
from types import MappingProxyType
from typing import Final, Literal, Protocol, TypeVar, assert_never
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching.dual_cache import DualCache
from litellm.constants import GLOBAL_PROXY_SPEND_CACHE_KEY, LITELLM_PROXY_BUDGET_NAME
from litellm.constants import (
GLOBAL_PROXY_SPEND_CACHE_KEY,
LITELLM_PROXY_BUDGET_NAME,
RESET_BUDGET_JOB_BATCH_SIZE,
RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN,
)
from litellm.proxy._types import (
LiteLLM_BudgetTableFull,
LiteLLM_EndUserTable,
@ -30,7 +37,10 @@ from litellm.repositories.table_repositories import (
TeamMembershipRepository,
)
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.unit_of_work import spend_reset_unit_of_work
from litellm.repositories.unit_of_work import (
budget_cascade_unit_of_work,
spend_reset_unit_of_work,
)
from litellm.repositories.verification_token_repository import (
VerificationTokenRepository,
)
@ -38,6 +48,9 @@ from litellm.types.services import ServiceTypes
_RowT = TypeVar("_RowT")
_LINKED_KEYS_WHERE: Final[Mapping[str, object]] = MappingProxyType({"budget_duration": None, "spend": {"gt": 0}})
_SPENT_ROWS_WHERE: Final[Mapping[str, object]] = MappingProxyType({"spend": {"gt": 0}})
class _TeamMembershipRow(Protocol):
@property
@ -62,39 +75,130 @@ class _TagRow(Protocol):
def tag_name(self) -> str: ...
class _EndUserRow(Protocol):
@property
def user_id(self) -> str: ...
def _team_membership_counter_key(row: _TeamMembershipRow) -> str:
return f"spend:team_member:{row.user_id}:{row.team_id}"
def _team_membership_cache_key(row: _TeamMembershipRow) -> str:
return f"{row.team_id}_{row.user_id}"
def _team_membership_cache_keys(row: _TeamMembershipRow) -> tuple[str, ...]:
return (f"{row.team_id}_{row.user_id}",)
def _key_counter_key(row: _KeyRow) -> str:
return f"spend:key:{row.token}"
def _key_cache_key(row: _KeyRow) -> str:
return row.token
def _key_cache_keys(row: _KeyRow) -> tuple[str, ...]:
return (row.token,)
def _org_counter_key(row: _OrgRow) -> str:
return f"spend:org:{row.organization_id}"
def _org_cache_keys(row: _OrgRow) -> Sequence[str]:
return [
def _org_cache_keys(row: _OrgRow) -> tuple[str, ...]:
return (
f"org_id:{row.organization_id}",
f"org_id:{row.organization_id}:with_budget",
]
)
def _tag_counter_key(row: _TagRow) -> str:
return f"spend:tag:{row.tag_name}"
def _tag_cache_key(row: _TagRow) -> str:
return f"tag:{row.tag_name}"
def _tag_cache_keys(row: _TagRow) -> tuple[str, ...]:
return (f"tag:{row.tag_name}",)
def _budget_link_where(
budget_ids: Sequence[str],
extra: Mapping[str, object] = MappingProxyType({}),
) -> dict[str, object]:
return {"budget_id": {"in": list(budget_ids)}, **extra}
@dataclass(frozen=True, slots=True)
class _BudgetCascade:
"""Everything one budget-tier reset touches, resolved before any write."""
budgets: tuple[LiteLLM_BudgetTableFull, ...] = ()
budget_ids: tuple[str, ...] = ()
budget_resets: tuple[tuple[str, datetime], ...] = ()
endusers: tuple[_EndUserRow, ...] = ()
counter_keys: tuple[str, ...] = ()
cache_keys: tuple[str, ...] = ()
@dataclass(frozen=True, slots=True)
class _BudgetCascadeCommitted:
cascade: _BudgetCascade
advanced: int
@dataclass(frozen=True, slots=True)
class _BudgetCascadeFailed:
cascade: _BudgetCascade
error: Exception
_EMPTY_CASCADE: Final = _BudgetCascade()
@dataclass(frozen=True, slots=True)
class _ChunkOutcome:
"""One chunk of a reset phase: rows read, and rows whose new budget_reset_at
cleared the due cutoff. Anything else is still due and would come straight
back on the next fetch, so it is not progress."""
fetched: int
advanced: int
_NO_PROGRESS: Final = _ChunkOutcome(fetched=0, advanced=0)
def _as_utc(moment: datetime) -> datetime:
return moment if moment.tzinfo is not None else moment.replace(tzinfo=timezone.utc)
def _count_advanced(reset_ats: Iterable[object], cutoff: datetime) -> int:
"""How many rows the write actually moved past the due cutoff.
A budget_duration of "0s" (or one the parser cannot read) resolves to the
current time, so the row is written and stays due. Counting it as progress
would re-read the same chunk until the per-run cap on every tick.
"""
utc_cutoff: Final = _as_utc(cutoff)
return sum(1 for reset_at in reset_ats if isinstance(reset_at, datetime) and _as_utc(reset_at) > utc_cutoff)
def _phase_is_drained(outcome: _ChunkOutcome) -> bool:
"""A short chunk means the due rows ran out. A full chunk that advanced
nothing would be re-read unchanged forever, so it ends the phase too and
those rows wait for the next tick."""
return outcome.fetched < RESET_BUDGET_JOB_BATCH_SIZE or outcome.advanced == 0
async def _run_phase_in_chunks(process_chunk: Callable[[], Awaitable[_ChunkOutcome]]) -> None:
"""Drive one reset phase a chunk at a time, capped so a single run cannot
spin unbounded: leftovers are picked up by the next tick."""
for _ in range(RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN):
if _phase_is_drained(await process_chunk()):
return
def _budget_cascade_event_metadata(cascade: _BudgetCascade) -> dict[str, object]:
return {
"num_budgets_found": len(cascade.budgets),
"budgets_found": json.dumps(cascade.budgets, indent=4, default=str),
"num_endusers_found": len(cascade.endusers),
"endusers_found": json.dumps(cascade.endusers, indent=4, default=str),
}
class ResetBudgetJob:
@ -122,21 +226,14 @@ class ResetBudgetJob:
Updates db
"""
if self.prisma_client is not None:
### RESET KEY BUDGET ###
await self.reset_budget_for_litellm_keys()
if self.prisma_client is None:
return
### RESET USER BUDGET ###
await self.reset_budget_for_litellm_users()
## Reset Team Budget
await self.reset_budget_for_litellm_teams()
### RESET ENDUSER (Customer) BUDGET and corresponding Budget duration ###
await self.reset_budget_for_litellm_budget_table()
### RESET MULTI-WINDOW BUDGETS ###
await self.reset_budget_windows()
await self.reset_budget_for_litellm_keys()
await self.reset_budget_for_litellm_users()
await self.reset_budget_for_litellm_teams()
await self.reset_budget_for_litellm_budget_table()
await self.reset_budget_windows()
@staticmethod
async def _invalidate_spend_counter(counter_key: str) -> None:
@ -194,238 +291,195 @@ class ResetBudgetJob:
e,
)
async def _cascade_reset_spend_for_budget_link(
async def _fetch_linked_rows(
self,
budgets_to_reset: list[LiteLLM_BudgetTableFull],
table: SpendLinkedTable[_RowT],
counter_key_fn: Callable[[_RowT], str],
where: Mapping[str, object],
log_subject: str,
extra_where: dict[str, object] | None = None,
cache_key_fn: Callable[[_RowT], str | Sequence[str]] | None = None,
):
"""
Generic cascade: zero spend on rows whose budget_id is in the reset set.
) -> tuple[_RowT, ...]:
"""Read the rows the cascade will zero, so their counters can be
invalidated once the transaction commits."""
try:
return tuple(await table.find_many(where=where))
except Exception as e:
verbose_proxy_logger.warning("Failed to fetch %s for counter invalidation: %s", log_subject, e)
return ()
``cache_key_fn`` is optional: when provided, after the DB update each
matching row's entry or entries in ``user_api_key_cache`` are dropped so
cached spend cannot stay pinned above the zeroed DB row after a reset.
async def _collect_endusers_to_reset(self, budget_ids: Sequence[str]) -> tuple[_EndUserRow, ...]:
linked: Final[Sequence[_EndUserRow] | None] = await self.prisma_client.get_data(
table_name="enduser",
query_type="find_all",
budget_id_list=list(budget_ids),
)
if litellm.max_end_user_budget_id is None or litellm.max_end_user_budget_id not in budget_ids:
return tuple(linked or ())
return (*(linked or ()), *await self._get_endusers_with_no_budget_id())
async def _collect_budget_cascade(self, budgets_to_reset: Sequence[LiteLLM_BudgetTableFull]) -> _BudgetCascade:
"""Resolve every row the expiring budget tiers gate, before any write.
Keys carrying their own budget_duration are left out: they run on their
own schedule via reset_budget_for_litellm_keys(), so sweeping them here
would reset them twice.
"""
budget_ids: Final = [b.budget_id for b in budgets_to_reset if b.budget_id is not None]
budget_ids: Final = tuple(b.budget_id for b in budgets_to_reset if b.budget_id is not None)
if not budget_ids:
return _EMPTY_CASCADE
team_memberships: Final[tuple[_TeamMembershipRow, ...]] = await self._fetch_linked_rows(
table=TeamMembershipRepository(self.prisma_client).table,
where=_budget_link_where(budget_ids),
log_subject="team memberships",
)
keys: Final[tuple[_KeyRow, ...]] = await self._fetch_linked_rows(
table=VerificationTokenRepository(self.prisma_client).table,
where=_budget_link_where(budget_ids, _LINKED_KEYS_WHERE),
log_subject="keys",
)
orgs: Final[tuple[_OrgRow, ...]] = await self._fetch_linked_rows(
table=OrganizationRepository(self.prisma_client).table,
where=_budget_link_where(budget_ids, _SPENT_ROWS_WHERE),
log_subject="orgs",
)
tags: Final[tuple[_TagRow, ...]] = await self._fetch_linked_rows(
table=TagRepository(self.prisma_client).table,
where=_budget_link_where(budget_ids, _SPENT_ROWS_WHERE),
log_subject="tags",
)
return _BudgetCascade(
budgets=tuple(budgets_to_reset),
budget_ids=budget_ids,
budget_resets=tuple(
(
b.budget_id,
compute_budget_reset_at(budget_duration=b.budget_duration, settings=self.reset_settings),
)
for b in budgets_to_reset
if b.budget_id is not None and b.budget_duration is not None
),
endusers=await self._collect_endusers_to_reset(budget_ids),
counter_keys=(
*(_team_membership_counter_key(row) for row in team_memberships),
*(_key_counter_key(row) for row in keys),
*(_org_counter_key(row) for row in orgs),
*(_tag_counter_key(row) for row in tags),
),
cache_keys=(
*(key for row in team_memberships for key in _team_membership_cache_keys(row)),
*(key for row in keys for key in _key_cache_keys(row)),
*(key for row in orgs for key in _org_cache_keys(row)),
*(key for row in tags for key in _tag_cache_keys(row)),
),
)
async def _commit_budget_cascade(self, cascade: _BudgetCascade) -> None:
"""Zero the gated spend and advance ``budget_reset_at`` in one transaction.
Advancing the window on its own hides the tier from every later tick
while its dependents stay pinned at the cap for the whole window;
batching both means a mid-cascade failure persists nothing and the rows
stay due for the next run.
"""
if not cascade.budget_ids:
return
where: Final[dict[str, object]] = {"budget_id": {"in": budget_ids}}
if extra_where:
where.update(extra_where)
enduser_ids: Final = tuple(row.user_id for row in cascade.endusers)
async with budget_cascade_unit_of_work(self.prisma_client.db.batch_) as uow:
uow.team_memberships.queue_spend_zero(where=_budget_link_where(cascade.budget_ids))
uow.keys.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _LINKED_KEYS_WHERE))
uow.organizations.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _SPENT_ROWS_WHERE))
uow.tags.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _SPENT_ROWS_WHERE))
if enduser_ids:
uow.endusers.queue_spend_zero(where={"user_id": {"in": list(enduser_ids)}})
for budget_id, budget_reset_at in cascade.budget_resets:
uow.budgets.queue_window_advance(budget_id=budget_id, budget_reset_at=budget_reset_at)
try:
rows: Sequence[_RowT] = await table.find_many(where=where)
except Exception as e:
rows = ()
verbose_proxy_logger.warning("Failed to fetch %s for counter invalidation: %s", log_subject, e)
update_result: Final = await table.update_many(where=where, data={"spend": 0})
for row in rows:
await self._invalidate_spend_counter(counter_key_fn(row))
if cache_key_fn is not None:
cache_keys = cache_key_fn(row)
if isinstance(cache_keys, str):
cache_keys = [cache_keys]
for cache_key in cache_keys:
await self._invalidate_user_api_key_cache_entry(cache_key)
return update_result
async def reset_budget_for_litellm_team_members(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
"""
Resets the budget for all LiteLLM Team Members if their budget has expired
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=TeamMembershipRepository(self.prisma_client).table,
counter_key_fn=_team_membership_counter_key,
log_subject="team memberships",
cache_key_fn=_team_membership_cache_key,
)
async def reset_budget_for_keys_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
"""
Resets the spend for keys linked to budget tiers that are being reset.
Excludes keys with their own budget_duration; those are reset by
reset_budget_for_litellm_keys() to avoid double-resetting.
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=VerificationTokenRepository(self.prisma_client).table,
counter_key_fn=_key_counter_key,
log_subject="keys",
extra_where={"budget_duration": None, "spend": {"gt": 0}},
cache_key_fn=_key_cache_key,
)
async def reset_budget_for_orgs_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
"""
Resets the spend for orgs linked to budget tiers that are being reset.
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=OrganizationRepository(self.prisma_client).table,
counter_key_fn=_org_counter_key,
log_subject="orgs",
extra_where={"spend": {"gt": 0}},
cache_key_fn=_org_cache_keys,
)
async def reset_budget_for_tags_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
"""
Resets the spend for tags linked to budget tiers that are being reset.
Also drops each tag's ``user_api_key_cache`` entry so the next
``_tag_max_budget_check`` reloads the zeroed row from the DB.
``SpendCounterReseed.from_db`` intentionally returns ``None`` for
tags, so the budget check falls back to the cached
``LiteLLM_TagTable.spend`` once the spend counter expires; without
this invalidation, that stale ``.spend`` keeps the tag over-budget
indefinitely.
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=TagRepository(self.prisma_client).table,
counter_key_fn=_tag_counter_key,
log_subject="tags",
extra_where={"spend": {"gt": 0}},
cache_key_fn=_tag_cache_key,
)
async def reset_budget_for_litellm_budget_table(self):
"""
Resets the budget for all LiteLLM End-Users (Customers), and Team Members if their budget has expired
The corresponding Budget duration is also updated.
"""
async def _invalidate_budget_cascade_caches(self, cascade: _BudgetCascade) -> None:
for counter_key in cascade.counter_keys:
await self._invalidate_spend_counter(counter_key)
for cache_key in cascade.cache_keys:
await self._invalidate_user_api_key_cache_entry(cache_key)
async def _reset_expired_budget_cascade(self) -> _BudgetCascadeCommitted | _BudgetCascadeFailed:
now: Final = datetime.now(timezone.utc)
start_time: Final = time.time()
endusers_to_reset: list[LiteLLM_EndUserTable] | None = None
budgets_to_reset: list[LiteLLM_BudgetTableFull] | None = None
updated_endusers: Final[list[LiteLLM_EndUserTable]] = []
failed_endusers: Final = []
try:
budgets_to_reset = await self.prisma_client.get_data(
table_name="budget", query_type="find_all", reset_at=now
)
if budgets_to_reset is not None and len(budgets_to_reset) > 0:
for budget in budgets_to_reset:
budget = await ResetBudgetJob._reset_budget_reset_at_date(budget, now, self.reset_settings)
await self.prisma_client.update_data(
query_type="update_many",
data_list=budgets_to_reset,
table_name="budget",
)
budget_ids_to_reset = [budget.budget_id for budget in budgets_to_reset if budget.budget_id is not None]
endusers_to_reset = await self.prisma_client.get_data(
table_name="enduser",
query_type="find_all",
budget_id_list=budget_ids_to_reset,
)
# Also reset end users with no budget_id (NULL) who use the
# default budget via litellm.max_end_user_budget_id. These
# users are enforced in-memory but never had budget_id
# persisted, so the query above misses them.
if litellm.max_end_user_budget_id is not None and litellm.max_end_user_budget_id in budget_ids_to_reset:
default_budget_endusers: Final = await self._get_endusers_with_no_budget_id()
if default_budget_endusers:
if endusers_to_reset is None:
endusers_to_reset = default_budget_endusers
else:
endusers_to_reset.extend(default_budget_endusers)
await self.reset_budget_for_litellm_team_members(budgets_to_reset=budgets_to_reset)
await self.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=budgets_to_reset)
await self.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=budgets_to_reset)
await self.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=budgets_to_reset)
if endusers_to_reset is not None and len(endusers_to_reset) > 0:
for enduser in endusers_to_reset:
try:
updated_enduser = await ResetBudgetJob._reset_budget_for_enduser(enduser=enduser)
if updated_enduser is not None:
updated_endusers.append(updated_enduser)
else:
failed_endusers.append(
{
"enduser": enduser,
"error": "Returned None without exception",
}
)
except Exception as e:
failed_endusers.append({"enduser": enduser, "error": str(e)})
verbose_proxy_logger.exception("Failed to reset budget for enduser: %s", enduser)
verbose_proxy_logger.debug(
"Updated users %s",
json.dumps(updated_endusers, indent=4, default=str),
)
await self.prisma_client.update_data(
query_type="update_many",
data_list=updated_endusers,
table_name="enduser",
)
end_time = time.time()
if len(failed_endusers) > 0: # If any endusers failed to reset
raise Exception(
f"Failed to reset {len(failed_endusers)} endusers: {json.dumps(failed_endusers, default=str)}"
)
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
call_type="reset_budget_budget_table",
start_time=start_time,
end_time=end_time,
event_metadata={
"num_budgets_found": (len(budgets_to_reset) if budgets_to_reset else 0),
"budgets_found": json.dumps(budgets_to_reset, indent=4, default=str),
"num_endusers_found": (len(endusers_to_reset) if endusers_to_reset else 0),
"endusers_found": json.dumps(endusers_to_reset, indent=4, default=str),
"num_endusers_updated": len(updated_endusers),
"endusers_updated": json.dumps(updated_endusers, indent=4, default=str),
"num_endusers_failed": len(failed_endusers),
"endusers_failed": json.dumps(failed_endusers, indent=4, default=str),
},
)
budgets_to_reset: Final[Sequence[LiteLLM_BudgetTableFull] | None] = await self.prisma_client.get_data(
table_name="budget",
query_type="find_all",
reset_at=now,
limit=RESET_BUDGET_JOB_BATCH_SIZE,
)
cascade: Final = await self._collect_budget_cascade(budgets_to_reset or ())
except Exception as e:
end_time = time.time()
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_failure_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
error=e,
call_type="reset_budget_endusers",
start_time=start_time,
end_time=end_time,
event_metadata={
"num_budgets_found": (len(budgets_to_reset) if budgets_to_reset else 0),
"budgets_found": json.dumps(budgets_to_reset, indent=4, default=str),
"num_endusers_found": (len(endusers_to_reset) if endusers_to_reset else 0),
"endusers_found": json.dumps(endusers_to_reset, indent=4, default=str),
},
return _BudgetCascadeFailed(cascade=_EMPTY_CASCADE, error=e)
try:
await self._commit_budget_cascade(cascade)
except Exception as e:
return _BudgetCascadeFailed(cascade=cascade, error=e)
await self._invalidate_budget_cascade_caches(cascade)
return _BudgetCascadeCommitted(
cascade=cascade,
advanced=_count_advanced(
(reset_at for _, reset_at in cascade.budget_resets),
cutoff=datetime.now(timezone.utc),
),
)
async def reset_budget_for_litellm_budget_table(self) -> None:
"""
Resets the spend a budget tier gates (end users, team members, keys,
orgs, tags) and advances the tier's budget_reset_at, atomically.
Caches are invalidated only after the transaction commits, so a failed
run cannot leave a zeroed counter in front of an un-reset DB row.
"""
await _run_phase_in_chunks(self._reset_budget_for_litellm_budget_table_chunk)
async def _reset_budget_for_litellm_budget_table_chunk(self) -> _ChunkOutcome:
start_time: Final = time.time()
outcome: Final = await self._reset_expired_budget_cascade()
end_time: Final = time.time()
match outcome:
case _BudgetCascadeCommitted(cascade=cascade, advanced=advanced):
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
call_type="reset_budget_budget_table",
start_time=start_time,
end_time=end_time,
event_metadata={
**_budget_cascade_event_metadata(cascade),
"num_endusers_updated": len(cascade.endusers),
"num_endusers_failed": 0,
},
)
)
)
verbose_proxy_logger.exception("Failed to reset budget for endusers: %s", e)
return _ChunkOutcome(fetched=len(cascade.budgets), advanced=advanced)
case _BudgetCascadeFailed(cascade=cascade, error=error):
verbose_proxy_logger.exception(
"Failed to reset the budget table cascade (team member, enduser, org and tag spend, plus "
"budget_reset_at); nothing was committed and the budgets stay due for the next run: %s",
error,
exc_info=error,
)
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_failure_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
error=error,
call_type="reset_budget_endusers",
start_time=start_time,
end_time=end_time,
event_metadata=_budget_cascade_event_metadata(cascade),
)
)
return _NO_PROGRESS
case _:
assert_never(outcome)
async def _get_endusers_with_no_budget_id(
self,
@ -486,18 +540,50 @@ class ResetBudgetJob:
for t in updated_teams:
uow.teams.queue_spend_reset(team_id=t.team_id, budget_reset_at=t.budget_reset_at)
async def reset_budget_for_litellm_keys(self):
def _emit_phase_failure(
self,
call_type: str,
error: Exception,
start_time: float,
end_time: float,
event_metadata: dict[str, object],
) -> None:
"""Report rows that could not be reset without failing the chunk: the
rows that did reset are already committed, and raising here would cost
the phase every remaining chunk this tick.
"""
verbose_proxy_logger.error("%s: %s", call_type, error)
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_failure_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
error=error,
call_type=call_type,
start_time=start_time,
end_time=end_time,
event_metadata=event_metadata,
)
)
async def reset_budget_for_litellm_keys(self) -> None:
"""
Resets the budget for all the litellm keys
Catches Exceptions and logs them
"""
await _run_phase_in_chunks(self._reset_budget_for_litellm_keys_chunk)
async def _reset_budget_for_litellm_keys_chunk(self) -> _ChunkOutcome:
now: Final = datetime.utcnow()
start_time: Final = time.time()
keys_to_reset: list[LiteLLM_VerificationToken] | None = None
try:
keys_to_reset = await self.prisma_client.get_data(
table_name="key", query_type="find_all", expires=now, reset_at=now
table_name="key",
query_type="find_all",
expires=now,
reset_at=now,
limit=RESET_BUDGET_JOB_BATCH_SIZE,
)
verbose_proxy_logger.debug("Keys to reset %s", json.dumps(keys_to_reset, indent=4, default=str))
updated_keys: Final[list[LiteLLM_VerificationToken]] = []
@ -528,8 +614,25 @@ class ResetBudgetJob:
await self._invalidate_spend_counter(f"spend:key:{token}")
end_time = time.time()
if len(failed_keys) > 0: # If any keys failed to reset
raise Exception(f"Failed to reset {len(failed_keys)} keys: {json.dumps(failed_keys, default=str)}")
outcome: Final = _ChunkOutcome(
fetched=len(keys_to_reset) if keys_to_reset else 0,
advanced=_count_advanced(
(k.budget_reset_at for k in updated_keys),
cutoff=datetime.now(timezone.utc),
),
)
if len(failed_keys) > 0:
self._emit_phase_failure(
call_type="reset_budget_keys",
error=Exception(f"Failed to reset {len(failed_keys)} keys: {json.dumps(failed_keys, default=str)}"),
start_time=start_time,
end_time=end_time,
event_metadata={
"num_keys_found": len(keys_to_reset) if keys_to_reset else 0,
"keys_found": json.dumps(keys_to_reset, indent=4, default=str),
},
)
return outcome
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
@ -565,16 +668,27 @@ class ResetBudgetJob:
)
)
verbose_proxy_logger.exception("Failed to reset budget for keys: %s", e)
return _NO_PROGRESS
else:
return outcome
async def reset_budget_for_litellm_users(self):
async def reset_budget_for_litellm_users(self) -> None:
"""
Resets the budget for all LiteLLM Internal Users if their budget has expired
"""
await _run_phase_in_chunks(self._reset_budget_for_litellm_users_chunk)
async def _reset_budget_for_litellm_users_chunk(self) -> _ChunkOutcome:
now: Final = datetime.utcnow()
start_time: Final = time.time()
users_to_reset: list[LiteLLM_UserTable] | None = None
try:
users_to_reset = await self.prisma_client.get_data(table_name="user", query_type="find_all", reset_at=now)
users_to_reset = await self.prisma_client.get_data(
table_name="user",
query_type="find_all",
reset_at=now,
limit=RESET_BUDGET_JOB_BATCH_SIZE,
)
updated_users: Final[list[LiteLLM_UserTable]] = []
failed_users: Final = []
if users_to_reset is not None and len(users_to_reset) > 0:
@ -609,8 +723,27 @@ class ResetBudgetJob:
await self._invalidate_global_proxy_spend_cache()
end_time = time.time()
if len(failed_users) > 0: # If any users failed to reset
raise Exception(f"Failed to reset {len(failed_users)} users: {json.dumps(failed_users, default=str)}")
outcome: Final = _ChunkOutcome(
fetched=len(users_to_reset) if users_to_reset else 0,
advanced=_count_advanced(
(u.budget_reset_at for u in updated_users),
cutoff=datetime.now(timezone.utc),
),
)
if len(failed_users) > 0:
self._emit_phase_failure(
call_type="reset_budget_users",
error=Exception(
f"Failed to reset {len(failed_users)} users: {json.dumps(failed_users, default=str)}"
),
start_time=start_time,
end_time=end_time,
event_metadata={
"num_users_found": len(users_to_reset) if users_to_reset else 0,
"users_found": json.dumps(users_to_reset, indent=4, default=str),
},
)
return outcome
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
@ -646,16 +779,27 @@ class ResetBudgetJob:
)
)
verbose_proxy_logger.exception("Failed to reset budget for users: %s", e)
return _NO_PROGRESS
else:
return outcome
async def reset_budget_for_litellm_teams(self):
async def reset_budget_for_litellm_teams(self) -> None:
"""
Resets the budget for all LiteLLM Internal Teams if their budget has expired
"""
await _run_phase_in_chunks(self._reset_budget_for_litellm_teams_chunk)
async def _reset_budget_for_litellm_teams_chunk(self) -> _ChunkOutcome:
now: Final = datetime.utcnow()
start_time: Final = time.time()
teams_to_reset: list[LiteLLM_TeamTable] | None = None
try:
teams_to_reset = await self.prisma_client.get_data(table_name="team", query_type="find_all", reset_at=now)
teams_to_reset = await self.prisma_client.get_data(
table_name="team",
query_type="find_all",
reset_at=now,
limit=RESET_BUDGET_JOB_BATCH_SIZE,
)
updated_teams: Final[list[LiteLLM_TeamTable]] = []
failed_teams: Final = []
if teams_to_reset is not None and len(teams_to_reset) > 0:
@ -688,8 +832,27 @@ class ResetBudgetJob:
await self._invalidate_spend_counter(f"spend:team:{team_id}")
end_time = time.time()
if len(failed_teams) > 0: # If any teams failed to reset
raise Exception(f"Failed to reset {len(failed_teams)} teams: {json.dumps(failed_teams, default=str)}")
outcome: Final = _ChunkOutcome(
fetched=len(teams_to_reset) if teams_to_reset else 0,
advanced=_count_advanced(
(t.budget_reset_at for t in updated_teams),
cutoff=datetime.now(timezone.utc),
),
)
if len(failed_teams) > 0:
self._emit_phase_failure(
call_type="reset_budget_teams",
error=Exception(
f"Failed to reset {len(failed_teams)} teams: {json.dumps(failed_teams, default=str)}"
),
start_time=start_time,
end_time=end_time,
event_metadata={
"num_teams_found": len(teams_to_reset) if teams_to_reset else 0,
"teams_found": json.dumps(teams_to_reset, indent=4, default=str),
},
)
return outcome
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
@ -725,6 +888,9 @@ class ResetBudgetJob:
)
)
verbose_proxy_logger.exception("Failed to reset budget for teams: %s", e)
return _NO_PROGRESS
else:
return outcome
@staticmethod
async def _reset_expired_window(
@ -882,33 +1048,6 @@ class ResetBudgetJob:
)
return user
@staticmethod
async def _reset_budget_for_enduser(
enduser: LiteLLM_EndUserTable,
) -> LiteLLM_EndUserTable | None:
try:
enduser.spend = 0.0
except Exception as e:
verbose_proxy_logger.exception("Error resetting budget for enduser: %s. Item: %s", e, enduser)
raise e
return enduser
@staticmethod
async def _reset_budget_reset_at_date(
budget: LiteLLM_BudgetTableFull,
current_time: datetime,
reset_settings: BudgetResetSettings,
) -> LiteLLM_BudgetTableFull:
try:
if budget.budget_duration is not None:
budget.budget_reset_at = compute_budget_reset_at(
budget_duration=budget.budget_duration, settings=reset_settings
)
except Exception as e:
verbose_proxy_logger.exception("Error resetting budget_reset_at for budget: %s. Item: %s", e, budget)
raise e
return budget
@staticmethod
async def _reset_budget_for_key(
key: LiteLLM_VerificationToken,

View file

@ -20,7 +20,10 @@ from fastapi import APIRouter, Depends, HTTPException
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.proxy.management_endpoints.common_utils import (
_user_has_admin_view,
validate_budget_duration,
)
from litellm.proxy.utils import jsonify_object
from litellm.repositories.budget_repository import BudgetRepository
@ -72,6 +75,8 @@ async def new_budget(
detail={"error": f"soft_budget must be a non-negative finite number. Received: {budget_obj.soft_budget}"},
)
validate_budget_duration(budget_obj.budget_duration)
# Validate model_max_budget if present
if budget_obj.model_max_budget is not None and len(budget_obj.model_max_budget) > 0:
from litellm.proxy.management_endpoints.key_management_endpoints import (
@ -153,6 +158,8 @@ async def update_budget(
detail={"error": f"soft_budget must be a non-negative finite number. Received: {budget_obj.soft_budget}"},
)
validate_budget_duration(budget_obj.budget_duration)
# Validate model_max_budget if present in update
if budget_obj.model_max_budget is not None and len(budget_obj.model_max_budget) > 0:
from litellm.proxy.management_endpoints.key_management_endpoints import (

View file

@ -22,6 +22,35 @@ def validate_finite_spend(spend: float | None) -> None:
)
def validate_budget_duration(budget_duration: str | None) -> None:
"""Reject budget durations that can't be parsed, are non-positive, or
overflow date math, so a bad value can't be persisted and later crash the
budget reset job.
A non-positive duration also resolves to a reset time of "now", which leaves
the row permanently due: the reset job re-reads it every tick and, once
enough of them exist, they fill each batch and starve every other tenant's
reset.
"""
if budget_duration is None:
return
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
try:
if duration_in_seconds(budget_duration) <= 0:
raise ValueError("budget_duration must be positive")
get_budget_reset_time(budget_duration=budget_duration)
except (ValueError, OverflowError):
raise HTTPException(
status_code=400,
detail={
"error": f"Invalid budget_duration '{budget_duration}'. Use a format like '1h', '24h', '7d', or '30d'."
},
)
from litellm._logging import verbose_proxy_logger
from litellm.caching import DualCache
from litellm.proxy._types import (

View file

@ -23,6 +23,7 @@ from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity
from litellm.proxy.management_endpoints.common_utils import validate_budget_duration
from litellm.proxy.management_helpers.object_permission_utils import (
_set_object_permission,
handle_update_object_permission_common,
@ -184,6 +185,7 @@ def new_budget_request(data: NewCustomerRequest) -> BudgetNewRequest | None:
if budget_kv_pairs:
budget_request: Final = BudgetNewRequest(**budget_kv_pairs)
validate_budget_duration(budget_request.budget_duration)
if budget_request.budget_reset_at is None and budget_request.budget_duration is not None:
budget_request.budget_reset_at = datetime.utcnow() + timedelta(
seconds=duration_in_seconds(duration=budget_request.budget_duration)

View file

@ -42,6 +42,7 @@ from litellm.proxy.management_endpoints.common_utils import (
_is_user_team_admin,
_user_has_admin_view,
require_caller_user_id_for_non_admin,
validate_budget_duration,
validate_finite_spend,
)
from litellm.proxy.management_endpoints.key_management_endpoints import (
@ -506,6 +507,8 @@ async def new_user(
status_code=500,
detail=CommonProxyErrors.db_not_connected_error.value,
)
validate_budget_duration(data.budget_duration)
# Check for duplicate user_id or email
await _check_duplicate_user_id(data.user_id, prisma_client)
await _check_duplicate_user_email(data.user_email, prisma_client)
@ -1185,6 +1188,7 @@ def _update_internal_user_params(data_json: dict, data: UpdateUserRequest | Upda
if "budget_duration" in non_default_values:
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
validate_budget_duration(non_default_values["budget_duration"])
non_default_values["budget_reset_at"] = get_budget_reset_time(
budget_duration=non_default_values["budget_duration"]
)

View file

@ -82,6 +82,7 @@ from litellm.proxy.management_endpoints.common_utils import (
_set_object_metadata_field,
_team_member_has_permission,
_user_has_admin_view,
validate_budget_duration,
validate_finite_spend,
)
from litellm.proxy.management_endpoints.model_management_endpoints import (
@ -844,6 +845,8 @@ async def _common_key_generation_helper(
premium_user=premium_user,
)
validate_budget_duration(data.budget_duration)
if data.throttle_on_budget_exceeded is True and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value:
raise HTTPException(
status_code=403,
@ -1024,7 +1027,7 @@ async def _common_key_generation_helper(
# Only set budget_duration on key when explicitly provided. Keys with budget_id
# but no explicit budget_duration follow their linked budget tier's schedule;
# reset_budget_for_keys_linked_to_budgets() resets them when the tier resets.
# reset_budget_for_litellm_budget_table() resets them when the tier resets.
# This avoids duplicating budget_duration on keys so tier updates apply automatically.
if "budget_duration" in data_json:
data_json["key_budget_duration"] = data_json.pop("budget_duration", None)
@ -2401,6 +2404,7 @@ async def _validate_update_key_data(
"""Validate permissions and constraints for key update."""
# Reject NaN/±inf spend before it can reach the DB / spend counter.
validate_finite_spend(data.spend)
validate_budget_duration(data.budget_duration)
_is_proxy_admin: Final = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value

View file

@ -96,6 +96,7 @@ from litellm.proxy.management_endpoints.common_utils import (
_update_metadata_fields,
_upsert_budget_and_membership,
_user_has_admin_view,
validate_budget_duration,
)
from litellm.proxy.management_endpoints.organization_endpoints import (
add_member_to_organization,
@ -1258,6 +1259,9 @@ async def new_team(
detail={"error": f"soft_budget must be a non-negative finite number. Received: {data.soft_budget}"},
)
validate_budget_duration(data.budget_duration)
validate_budget_duration(data.team_member_budget_duration)
if data.soft_budget is not None:
if data.max_budget is not None:
# If max_budget is set, soft_budget must be strictly lower than max_budget
@ -1947,6 +1951,9 @@ async def update_team(
detail={"error": f"soft_budget must be a non-negative finite number. Received: {data.soft_budget}"},
)
validate_budget_duration(data.budget_duration)
validate_budget_duration(data.team_member_budget_duration)
existing_team_row = await _team_db(prisma_client).find_unique(where={"team_id": data.team_id})
if existing_team_row is None:
@ -2979,7 +2986,7 @@ async def team_member_add(
except HTTPException as e:
raise e
_validate_budget_duration(data.budget_duration)
validate_budget_duration(data.budget_duration)
prisma_client = cast(PrismaClient, prisma_client)
@ -3282,29 +3289,6 @@ def _build_member_budget_patch(data: TeamMemberUpdateRequest) -> dict[str, objec
}
def _validate_budget_duration(budget_duration: str | None) -> None:
"""Reject budget durations that can't be parsed, are non-positive, or
overflow date math, so a bad value can't be persisted and later crash the
budget reset job."""
if budget_duration is None:
return
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
try:
if duration_in_seconds(budget_duration) <= 0:
raise ValueError("budget_duration must be positive")
get_budget_reset_time(budget_duration=budget_duration)
except (ValueError, OverflowError):
raise HTTPException(
status_code=400,
detail={
"error": f"Invalid budget_duration '{budget_duration}'. Use a format like '1h', '24h', '7d', or '30d'."
},
)
@router.post(
"/team/member_update",
tags=["team management"],
@ -3342,7 +3326,7 @@ async def team_member_update(
detail={"error": "Either user_id or user_email needs to be passed in"},
)
_validate_budget_duration(data.budget_duration)
validate_budget_duration(data.budget_duration)
_existing_team_row: Final = await _team_db(prisma_client).find_unique(where={"team_id": data.team_id})

View file

@ -6860,9 +6860,19 @@ class ProxyConfig:
guardrail_id = guardrail.get("guardrail_id")
if guardrail_id:
db_guardrail_ids.add(guardrail_id)
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
guardrail=cast(Guardrail, guardrail),
)
try:
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
guardrail=cast(Guardrail, guardrail),
)
except Exception as e: # noqa: BLE001 # one unloadable row must not stop the remaining guardrails
verbose_proxy_logger.error(
"litellm.proxy.proxy_server.py::ProxyConfig:_init_guardrails_in_db - "
"skipping guardrail '%s' (ID: %s): %s: %s",
guardrail.get("guardrail_name"),
guardrail_id,
type(e).__name__,
e,
)
# Drop in-memory DB-backed entries whose row was deleted on another
# pod. Config-loaded entries are never touched.

View file

@ -3486,13 +3486,15 @@ class PrismaClient:
r.expires = r.expires.isoformat()
elif query_type == "find_all" and expires is not None and reset_at is not None:
response = await VerificationTokenRepository(self).table.find_many(
take=limit,
where={
"OR": [
{"expires": None},
{"expires": {"gt": expires}},
],
"budget_reset_at": {"lt": reset_at},
}
"NOT": {"budget_duration": None},
},
)
if response is not None and len(response) > 0:
for r in response:
@ -3542,6 +3544,7 @@ class PrismaClient:
response = await UserRepository(self).table.find_many(where=key_val)
elif query_type == "find_all" and reset_at is not None:
response = await UserRepository(self).table.find_many(
take=limit,
where={
# A user seeded from default_internal_user_params
# (or created via /user/new without an explicit
@ -3552,16 +3555,12 @@ class PrismaClient:
# of the row, silently exceeding max_budget. Treat a
# NULL budget_reset_at with a non-NULL budget_duration
# as due, matching the budget-table query below.
"NOT": {"budget_duration": None},
"OR": [
{
"AND": [
{"budget_reset_at": None},
{"NOT": {"budget_duration": None}},
]
},
{"budget_reset_at": None},
{"budget_reset_at": {"lt": reset_at}},
],
}
},
)
elif query_type == "find_all" and user_id_list is not None:
response = await UserRepository(self).table.find_many(where={"user_id": {"in": user_id_list}})
@ -3617,17 +3616,14 @@ class PrismaClient:
elif table_name == "budget" and reset_at is not None:
if query_type == "find_all":
response = await BudgetRepository(self).table.find_many(
take=limit,
where={
"NOT": {"budget_duration": None},
"OR": [
{
"AND": [
{"budget_reset_at": None},
{"NOT": {"budget_duration": None}},
]
},
{"budget_reset_at": None},
{"budget_reset_at": {"lt": reset_at}},
]
}
],
},
)
return response
@ -3645,20 +3641,17 @@ class PrismaClient:
)
elif query_type == "find_all" and reset_at is not None:
response = await TeamRepository(self).table.find_many(
take=limit,
where={
# Same NULL budget_reset_at gap as the user query
# above: a team with a budget_duration but no
# initialized budget_reset_at would never be reset.
"NOT": {"budget_duration": None},
"OR": [
{
"AND": [
{"budget_reset_at": None},
{"NOT": {"budget_duration": None}},
]
},
{"budget_reset_at": None},
{"budget_reset_at": {"lt": reset_at}},
],
}
},
)
elif query_type == "find_all" and user_id is not None:
response = await TeamRepository(self).table.find_many(

View file

@ -1,4 +1,4 @@
from typing import Any, Final
from typing import Annotated, Any, Final
from fastapi import APIRouter, Depends, HTTPException, Request, Response
@ -18,7 +18,8 @@ from litellm.proxy.vector_store_endpoints.utils import (
get_litellm_managed_vector_store,
)
from litellm.repositories.table_repositories import ManagedVectorStoreIndexRepository
from litellm.types.vector_stores import IndexCreateRequest
from litellm.types.vector_stores import IndexCreateRequest, IndexListResponse
from litellm.vector_stores.vector_store_registry import VectorStoreIndexRegistry
router: Final = APIRouter()
########################################################
@ -549,14 +550,15 @@ async def index_create(
Create an index. Just writes the index to the database.
```bash
curl -L -X POST 'http://0.0.0.0:4000/indexes/create' \
curl -L -X POST 'http://0.0.0.0:4000/v1/indexes' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer sk-1234' \
-H 'LiteLLM-Beta: indexes_beta=v1' \
-d '{
-d '{
"index_name": "dall-e-3",
"vector_store_index": "real-index-name",
"vector_store_name": "azure-ai-search"
"litellm_params": {
"vector_store_index": "real-index-name",
"vector_store_name": "azure-ai-search"
}
}'
```
"""
@ -592,3 +594,36 @@ async def index_create(
new_index = await ManagedVectorStoreIndexRepository(prisma_client).table.create(data=jsonify_object(index_data))
return new_index.model_dump()
@router.get(
"/v1/indexes",
dependencies=[Depends(user_api_key_auth)],
response_model=IndexListResponse,
)
async def index_list(
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
) -> IndexListResponse:
"""
List all vector store indexes. Proxy admin only.
```bash
curl -L -X GET 'http://0.0.0.0:4000/v1/indexes' \
-H 'Authorization: Bearer sk-1234'
```
"""
from litellm.proxy.proxy_server import prisma_client
assert_proxy_admin_for_vector_store_index_management(
user_api_key_dict,
operation="list",
)
if prisma_client is None:
raise HTTPException(
status_code=500,
detail=CommonProxyErrors.db_not_connected_error.value,
)
indexes: Final = await VectorStoreIndexRegistry._get_vector_store_indexes_from_db(prisma_client)
return IndexListResponse(data=indexes)

View file

@ -41,7 +41,7 @@ def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
def assert_proxy_admin_for_vector_store_index_management(
user_api_key_dict: UserAPIKeyAuth,
*,
operation: Literal["create", "delete", "update"] = "create",
operation: Literal["create", "delete", "update", "list"] = "create",
) -> None:
"""Raise 403 unless the caller is a proxy admin."""
if _is_proxy_admin(user_api_key_dict):

View file

@ -70,10 +70,14 @@ from litellm.repositories.table_repositories import (
)
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.unit_of_work import (
BudgetCascadeUnitOfWork,
BudgetWindowWrites,
KeySpendResetWrites,
LinkedSpendResetWrites,
SpendResetUnitOfWork,
TeamSpendResetWrites,
UserSpendResetWrites,
budget_cascade_unit_of_work,
spend_reset_unit_of_work,
)
from litellm.repositories.user_repository import UserRepository
@ -88,7 +92,9 @@ __all__ = [
"AgentsRepository",
"AuditLogRepository",
"BatchTable",
"BudgetCascadeUnitOfWork",
"BudgetRepository",
"BudgetWindowWrites",
"CacheConfigRepository",
"ClaudeCodePluginRepository",
"ConfigOverridesRepository",
@ -107,6 +113,7 @@ __all__ = [
"InvitationLinkRepository",
"JWTKeyMappingRepository",
"KeySpendResetWrites",
"LinkedSpendResetWrites",
"MCPServerRepository",
"MCPToolsetRepository",
"MCPUserCredentialsRepository",
@ -149,5 +156,6 @@ __all__ = [
"WorkflowEventRepository",
"WorkflowMessageRepository",
"WorkflowRunRepository",
"budget_cascade_unit_of_work",
"spend_reset_unit_of_work",
]

View file

@ -29,6 +29,8 @@ class SpendLinkedTable(Protocol[RowT_co]):
class BatchTable(Protocol):
def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ...
def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ...
class PrismaBatch(Protocol):
@property
@ -40,4 +42,19 @@ class PrismaBatch(Protocol):
@property
def litellm_teamtable(self) -> BatchTable: ...
@property
def litellm_budgettable(self) -> BatchTable: ...
@property
def litellm_teammembership(self) -> BatchTable: ...
@property
def litellm_organizationtable(self) -> BatchTable: ...
@property
def litellm_tagtable(self) -> BatchTable: ...
@property
def litellm_endusertable(self) -> BatchTable: ...
async def commit(self) -> None: ...

View file

@ -1,17 +1,21 @@
"""
Unit of work over a single Prisma batch.
Units of work over a single Prisma batch.
``spend_reset_unit_of_work`` opens one ``db.batch_()`` and binds a typed write
Each context manager here opens one ``db.batch_()`` and binds a typed write
repository per table to it, so every update queued through the yielded object
lands in the same transaction. The batch commits when the block exits cleanly
and is abandoned, writing nothing, when the block raises.
Each write repository queues narrow ``{spend, budget_reset_at}`` updates
``spend_reset_unit_of_work`` covers the per-row key/user/team resets;
``budget_cascade_unit_of_work`` covers a budget tier's reset, where the
dependent spend and the tier's next window have to move together.
Each write repository queues narrow ``{spend}`` / ``{budget_reset_at}`` updates
instead of full-model writes, which trip ``prisma.errors.DataError`` on rows
carrying fields the update input type rejects (see #27730).
"""
from collections.abc import AsyncGenerator, Callable
from collections.abc import AsyncGenerator, Callable, Mapping
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import datetime
@ -43,6 +47,24 @@ class TeamSpendResetWrites:
self.table.update(where={"team_id": team_id}, data={"spend": 0, "budget_reset_at": budget_reset_at})
@dataclass(frozen=True, slots=True)
class LinkedSpendResetWrites:
table: BatchTable
def queue_spend_zero(self, where: Mapping[str, object]) -> None:
self.table.update_many(where=where, data={"spend": 0})
@dataclass(frozen=True, slots=True)
class BudgetWindowWrites:
table: BatchTable
def queue_window_advance(self, budget_id: str, budget_reset_at: datetime) -> None:
"""``update_many`` so a tier deleted between the read and the commit is a
no-op row count instead of a P2025 that aborts the whole chunk."""
self.table.update_many(where={"budget_id": budget_id}, data={"budget_reset_at": budget_reset_at})
@dataclass(frozen=True, slots=True)
class SpendResetUnitOfWork:
keys: KeySpendResetWrites
@ -50,6 +72,23 @@ class SpendResetUnitOfWork:
teams: TeamSpendResetWrites
@dataclass(frozen=True, slots=True)
class BudgetCascadeUnitOfWork:
"""Every write a budget-tier reset performs, bound to one batch.
The dependent spend rows and the budget rows' ``budget_reset_at`` advance
must land together: advancing the window without zeroing the spend it
gates leaves the dependents pinned at their cap until the next window.
"""
team_memberships: LinkedSpendResetWrites
keys: LinkedSpendResetWrites
organizations: LinkedSpendResetWrites
tags: LinkedSpendResetWrites
endusers: LinkedSpendResetWrites
budgets: BudgetWindowWrites
@asynccontextmanager
async def spend_reset_unit_of_work(new_batch: Callable[[], PrismaBatch]) -> AsyncGenerator[SpendResetUnitOfWork, None]:
batch = new_batch()
@ -59,3 +98,19 @@ async def spend_reset_unit_of_work(new_batch: Callable[[], PrismaBatch]) -> Asyn
teams=TeamSpendResetWrites(table=batch.litellm_teamtable),
)
await batch.commit()
@asynccontextmanager
async def budget_cascade_unit_of_work(
new_batch: Callable[[], PrismaBatch],
) -> AsyncGenerator[BudgetCascadeUnitOfWork, None]:
batch = new_batch()
yield BudgetCascadeUnitOfWork(
team_memberships=LinkedSpendResetWrites(table=batch.litellm_teammembership),
keys=LinkedSpendResetWrites(table=batch.litellm_verificationtoken),
organizations=LinkedSpendResetWrites(table=batch.litellm_organizationtable),
tags=LinkedSpendResetWrites(table=batch.litellm_tagtable),
endusers=LinkedSpendResetWrites(table=batch.litellm_endusertable),
budgets=BudgetWindowWrites(table=batch.litellm_budgettable),
)
await batch.commit()

View file

@ -277,6 +277,11 @@ class LiteLLM_ManagedVectorStoreIndex(BaseModel):
updated_by: str | None = None
class IndexListResponse(BaseModel):
object: Literal["list"] = "list"
data: tuple[LiteLLM_ManagedVectorStoreIndex, ...]
class VectorStoreIndexType(str, Enum):
"""Type of vector store index"""

View file

@ -9,10 +9,10 @@
"limit": 832
},
"ANN201": {
"limit": 2031
"limit": 2023
},
"ANN202": {
"limit": 861
"limit": 860
},
"ANN204": {
"limit": 713
@ -237,13 +237,13 @@
"limit": 1226
},
"TRY002": {
"limit": 528
"limit": 524
},
"TRY004": {
"limit": 96
},
"TRY201": {
"limit": 407
"limit": 405
},
"TRY203": {
"limit": 113

View file

@ -252,26 +252,40 @@ def scan_comments(path: Path, source: str) -> tuple[Comments, tuple[Violation, .
# --------------------------------------------------------------------------- #
def mutable_names_in(annotation: ast.expr) -> Iterator[str]:
def _is_literal_subscript(node: ast.AST) -> bool:
if not isinstance(node, ast.Subscript):
return False
base: Final = node.value
return (isinstance(base, ast.Name) and base.id == "Literal") or (
isinstance(base, ast.Attribute) and base.attr == "Literal"
)
def mutable_names_in(annotation: ast.AST) -> Iterator[str]:
"""Yield mutable-collection names anywhere inside an annotation expression.
Matches bare names (`dict`, `MutableMapping`) and dotted access (`typing.Dict`,
`collections.deque`, `collections.abc.MutableMapping`), descends through nesting
(`Mapping[str, list[int]]`, `tuple[set[int], ...]`) and string forward references.
Skips `Literal[...]` subtrees: their string arguments are values, not forward
references, so `Literal["list"]` is not the `list` type.
"""
for node in ast.walk(annotation):
if isinstance(node, ast.Name) and node.id in MUTABLE_COLLECTIONS:
yield node.id
elif isinstance(node, ast.Attribute) and node.attr in MUTABLE_COLLECTIONS:
yield node.attr
elif isinstance(node, ast.Constant):
value: object = node.value # forward references arrive as string constants
if isinstance(value, str):
try:
inner = ast.parse(value, mode="eval").body
except SyntaxError:
continue
yield from mutable_names_in(inner)
if _is_literal_subscript(annotation):
return
if isinstance(annotation, ast.Name) and annotation.id in MUTABLE_COLLECTIONS:
yield annotation.id
elif isinstance(annotation, ast.Attribute) and annotation.attr in MUTABLE_COLLECTIONS:
yield annotation.attr
elif isinstance(annotation, ast.Constant):
value: object = annotation.value # forward references arrive as string constants
if isinstance(value, str):
try:
inner = ast.parse(value, mode="eval").body
except SyntaxError:
return
yield from mutable_names_in(inner)
for child in ast.iter_child_nodes(annotation):
yield from mutable_names_in(child)
def _mutable_ann(path: Path, line: int, name: str, where: str) -> Violation:

View file

@ -44,39 +44,87 @@ def _attrify(d: dict):
return _AttrDict(d)
def _wire_batcher_for_test(prisma_client):
def _wire_batcher_for_test(prisma_client, fail_commit=False):
"""
Wire prisma_client.db.batch_() to return a mock batcher whose .commit() is
awaitable and whose per-table .update() calls get captured. The reset job
writes key/user/team resets via prisma.db.batch_().<table>.update not via
prisma_client.update_data so tests must let that batch path complete.
awaitable and whose per-table .update()/.update_many() calls get captured.
The reset job writes every reset through prisma.db.batch_() key/user/team
rows one by one, and the budget tier's cascade as a single transaction — so
tests must let that batch path complete.
Returns the list that will accumulate {table, where, data} dicts from
each captured update call.
Only committed batches contribute to the returned list, mirroring prisma:
with fail_commit=True the transaction blows up and must persist nothing.
Returns the list that will accumulate {table, op, where, data} dicts from
each captured write.
"""
batch_calls = []
def make_batcher():
queued = []
class _Table:
def __init__(self, table_name):
self._table_name = table_name
def update(self, where=None, data=None):
batch_calls.append(
{"table": self._table_name, "where": where, "data": data}
queued.append(
{
"table": self._table_name,
"op": "update",
"where": where,
"data": data,
}
)
def update_many(self, where=None, data=None):
queued.append(
{
"table": self._table_name,
"op": "update_many",
"where": where,
"data": data,
}
)
async def commit():
if fail_commit:
raise RuntimeError("simulated Postgres failure committing the batch")
batch_calls.extend(queued)
batcher = MagicMock()
batcher.litellm_verificationtoken = _Table("key")
batcher.litellm_usertable = _Table("user")
batcher.litellm_teamtable = _Table("team")
batcher.commit = AsyncMock(return_value=None)
batcher.litellm_budgettable = _Table("budget")
batcher.litellm_teammembership = _Table("team_membership")
batcher.litellm_organizationtable = _Table("org")
batcher.litellm_tagtable = _Table("tag")
batcher.litellm_endusertable = _Table("enduser")
batcher.commit = commit
return batcher
prisma_client.db.batch_ = MagicMock(side_effect=make_batcher)
return batch_calls
def _wire_cascade_reads_for_test(prisma_client):
"""
The budget tier's cascade reads the rows it is about to zero, so their
spend counters can be invalidated after the commit. Give each of those
tables an awaitable find_many so the reads resolve instead of falling into
the job's warn-and-continue path.
"""
for table in (
"litellm_teammembership",
"litellm_verificationtoken",
"litellm_organizationtable",
"litellm_tagtable",
"litellm_endusertable",
):
getattr(prisma_client.db, table).find_many = AsyncMock(return_value=[])
@pytest.mark.asyncio
async def test_reset_budget_keys_partial_failure():
"""
@ -250,41 +298,18 @@ async def test_reset_budget_users_partial_failure():
@pytest.mark.asyncio
async def test_reset_budget_endusers_partial_failure():
async def test_reset_budget_endusers_cascade_failure_is_all_or_nothing():
"""
Test that if one enduser fails to reset, the reset loop still processes the other endusers.
We simulate six endsers where the first fails and the others are updated.
A failure anywhere in the budget-tier cascade must persist nothing, so the
tier stays due and the next scheduler tick retries it. Before the fix the
job committed the new budget_reset_at first and zeroed the dependent spend
afterwards, so a failure here left the tier stamped for the next window
while every end user stayed at the cap.
"""
user1 = {
"user_id": "user1",
"spend": 20.0,
"budget_id": "budget1",
} # Will trigger simulated failure
user2 = {
"user_id": "user2",
"spend": 25.0,
"budget_id": "budget1",
} # Should be updated
user3 = {
"user_id": "user3",
"spend": 30.0,
"budget_id": "budget1",
} # Should be updated
user4 = {
"user_id": "user4",
"spend": 35.0,
"budget_id": "budget1",
} # Should be updated
user5 = {
"user_id": "user5",
"spend": 40.0,
"budget_id": "budget1",
} # Should be updated
user6 = {
"user_id": "user6",
"spend": 45.0,
"budget_id": "budget1",
} # Should be updated
endusers = [
_attrify({"user_id": f"user{i}", "spend": 20.0 + i, "budget_id": "budget1"})
for i in range(1, 7)
]
budget1 = LiteLLM_BudgetTableFull(
**{
@ -301,23 +326,13 @@ async def test_reset_budget_endusers_partial_failure():
if table_name == "budget":
return [budget1]
elif table_name == "enduser":
return [user1, user2, user3, user4, user5, user6]
return endusers
return []
prisma_client.get_data = AsyncMock()
prisma_client.get_data.side_effect = get_data_mock
prisma_client.update_data = AsyncMock()
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
batch_calls = _wire_batcher_for_test(prisma_client, fail_commit=True)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -326,41 +341,13 @@ async def test_reset_budget_endusers_partial_failure():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_enduser(enduser):
if enduser["user_id"] == "user1":
raise Exception("Simulated failure for user1")
enduser["spend"] = 0.0
return enduser
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
async def fake_reset_team_members(budgets_to_reset):
return 1
with (
patch.object(
ResetBudgetJob,
"_reset_budget_for_enduser",
side_effect=fake_reset_enduser,
) as mock_reset_enduser,
patch.object(
ResetBudgetJob,
"reset_budget_for_litellm_team_members",
side_effect=fake_reset_team_members,
) as mock_reset_team_members,
):
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
assert mock_reset_enduser.call_count == 6
assert prisma_client.update_data.await_count == 2
update_call = prisma_client.update_data.call_args
assert update_call.kwargs.get("table_name") == "enduser"
updated_users = update_call.kwargs.get("data_list", [])
assert len(updated_users) == 5
assert updated_users[0]["user_id"] == "user2"
assert updated_users[1]["user_id"] == "user3"
assert updated_users[2]["user_id"] == "user4"
assert updated_users[3]["user_id"] == "user5"
assert updated_users[4]["user_id"] == "user6"
assert batch_calls == [], "a failed cascade must not persist any write"
assert (
prisma_client.update_data.await_count == 0
), "budget_reset_at must not be advanced outside the cascade transaction"
failure_hook_calls = (
proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args_list
@ -369,6 +356,66 @@ async def test_reset_budget_endusers_partial_failure():
call.kwargs.get("call_type") == "reset_budget_endusers"
for call in failure_hook_calls
)
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_not_called()
@pytest.mark.asyncio
async def test_reset_budget_endusers_are_zeroed_with_the_budget_window_advance():
"""
The happy path: every end user the tier gates is zeroed and the tier's
budget_reset_at advances, all inside one transaction.
"""
endusers = [
_attrify({"user_id": f"user{i}", "spend": 20.0 + i, "budget_id": "budget1"})
for i in range(1, 7)
]
budget1 = LiteLLM_BudgetTableFull(
**{
"budget_id": "budget1",
"max_budget": 65.0,
"budget_duration": "2d",
"created_at": datetime.now(timezone.utc) - timedelta(days=3),
}
)
prisma_client = MagicMock()
async def get_data_mock(table_name, *args, **kwargs):
if table_name == "budget":
return [budget1]
elif table_name == "enduser":
return endusers
return []
prisma_client.get_data = AsyncMock()
prisma_client.get_data.side_effect = get_data_mock
prisma_client.update_data = AsyncMock()
batch_calls = _wire_batcher_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock()
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
assert prisma_client.db.batch_.call_count == 1, "the cascade must be one transaction"
enduser_writes = [c for c in batch_calls if c["table"] == "enduser"]
assert len(enduser_writes) == 1
assert enduser_writes[0]["where"]["user_id"]["in"] == [f"user{i}" for i in range(1, 7)]
assert enduser_writes[0]["data"] == {"spend": 0}
budget_writes = [c for c in batch_calls if c["table"] == "budget"]
assert len(budget_writes) == 1
assert budget_writes[0]["where"] == {"budget_id": "budget1"}
assert budget_writes[0]["data"]["budget_reset_at"] > datetime.now(timezone.utc)
proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_not_called()
@pytest.mark.asyncio
@ -500,16 +547,8 @@ async def test_reset_budget_continues_other_categories_on_failure():
key1, key2 = _attrify(key1), _attrify(key2)
user1, user2 = _attrify(user1), _attrify(user2)
team1, team2 = _attrify(team1), _attrify(team2)
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
enduser1 = _attrify(enduser1)
_wire_cascade_reads_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -541,13 +580,6 @@ async def test_reset_budget_continues_other_categories_on_failure():
).isoformat()
return team
async def fake_reset_enduser(enduser):
enduser["spend"] = 0.0
return enduser
async def fake_reset_team_members(budgets_to_reset):
return 1
with (
patch.object(
ResetBudgetJob, "_reset_budget_for_key", side_effect=fake_reset_key
@ -558,14 +590,6 @@ async def test_reset_budget_continues_other_categories_on_failure():
patch.object(
ResetBudgetJob, "_reset_budget_for_team", side_effect=fake_reset_team
) as mock_reset_team,
patch.object(
ResetBudgetJob, "_reset_budget_for_enduser", side_effect=fake_reset_enduser
) as mock_reset_enduser,
patch.object(
ResetBudgetJob,
"reset_budget_for_litellm_team_members",
side_effect=fake_reset_team_members,
) as mock_reset_team_members,
):
# Call the overall reset_budget method.
await job.reset_budget()
@ -575,29 +599,22 @@ async def test_reset_budget_continues_other_categories_on_failure():
called_tables = {
call.kwargs.get("table_name") for call in prisma_client.get_data.await_args_list
}
if mock_reset_team_members.call_count > 0:
called_tables.add("team_membership")
assert called_tables == {
"key",
"user",
"team",
"budget",
"enduser",
"team_membership",
}
assert called_tables == {"key", "user", "team", "budget", "enduser"}
# After the fix, keys/users/teams write via prisma.db.batch_().<table>.update,
# so only budget + enduser still go through update_data.
calls = prisma_client.update_data.await_args_list
update_data_tables = [c.kwargs.get("table_name") for c in calls]
assert sorted(update_data_tables) == ["budget", "enduser"]
# Every category writes through the batch path now, so update_data is unused.
prisma_client.update_data.assert_not_awaited()
# Check enduser update: enduser succeed.
enduser_call = next(c for c in calls if c.kwargs.get("table_name") == "enduser")
assert len(enduser_call.kwargs.get("data_list", [])) == 1
# The budget tier's cascade still ran despite the failing user category.
assert len([c for c in batch_calls if c["table"] == "team_membership"]) == 1
enduser_writes = [c for c in batch_calls if c["table"] == "enduser"]
assert len(enduser_writes) == 1
assert enduser_writes[0]["where"] == {"user_id": {"in": ["user1"]}}
assert enduser_writes[0]["data"] == {"spend": 0}
# Check the new batch write path: 2 keys + 1 user (user1 failed) + 2 teams.
key_writes = [c for c in batch_calls if c["table"] == "key"]
# `op` separates the per-row resets from the cascade sweep, which also
# targets the key table.
key_writes = [c for c in batch_calls if c["table"] == "key" and c["op"] == "update"]
user_writes = [c for c in batch_calls if c["table"] == "user"]
team_writes = [c for c in batch_calls if c["table"] == "team"]
assert len(key_writes) == 2
@ -974,12 +991,12 @@ async def test_service_logger_teams_failure():
@pytest.mark.asyncio
async def test_service_logger_endusers_success():
"""
Test that when resetting endusers succeeds the service logger success hook is called with
the correct metadata and no exception is logged.
Test that when the budget-tier cascade commits, the service logger success
hook is called with the correct metadata and no exception is logged.
"""
endusers = [
{"user_id": "user1", "spend": 25.0, "budget_id": "budget1"},
{"user_id": "user2", "spend": 25.0, "budget_id": "budget1"},
_attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}),
_attrify({"user_id": "user2", "spend": 25.0, "budget_id": "budget1"}),
]
budgets = [
LiteLLM_BudgetTableFull(
@ -1002,16 +1019,8 @@ async def test_service_logger_endusers_success():
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
prisma_client.update_data = AsyncMock()
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
batch_calls = _wire_batcher_for_test(prisma_client)
_wire_cascade_reads_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -1020,31 +1029,16 @@ async def test_service_logger_endusers_success():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_enduser(enduser):
enduser["spend"] = 0.0
return enduser
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
mock_verbose_exc.assert_not_called()
async def fake_reset_team_members(budgets_to_reset):
return 1
with (
patch.object(
ResetBudgetJob,
"_reset_budget_for_enduser",
side_effect=fake_reset_enduser,
) as mock_reset_enduser,
patch.object(
ResetBudgetJob,
"reset_budget_for_litellm_team_members",
side_effect=fake_reset_team_members,
) as mock_reset_team_members,
):
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
mock_verbose_exc.assert_not_called()
enduser_writes = [c for c in batch_calls if c["table"] == "enduser"]
assert len(enduser_writes) == 1
assert enduser_writes[0]["where"] == {"user_id": {"in": ["user1", "user2"]}}
proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_called_once()
(
@ -1062,12 +1056,12 @@ async def test_service_logger_endusers_success():
@pytest.mark.asyncio
async def test_service_logger_endusers_failure():
"""
Test that a failure during enduser reset calls the failure hook with appropriate metadata,
logs the exception, and does not call the success hook.
Test that a failed cascade calls the failure hook with the rows it had
found, logs the exception, and does not call the success hook.
"""
endusers = [
{"user_id": "user1", "spend": 25.0, "budget_id": "budget1"},
{"user_id": "user2", "spend": 25.0, "budget_id": "budget1"},
_attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}),
_attrify({"user_id": "user2", "spend": 25.0, "budget_id": "budget1"}),
]
budgets = [
LiteLLM_BudgetTableFull(
@ -1090,16 +1084,8 @@ async def test_service_logger_endusers_failure():
prisma_client = MagicMock()
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
prisma_client.update_data = AsyncMock()
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_organizationtable.update_many (used by reset_budget_for_orgs_linked_to_budgets)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
# Mock db.litellm_tagtable.update_many (used by reset_budget_for_tags_linked_to_budgets)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
_wire_batcher_for_test(prisma_client, fail_commit=True)
_wire_cascade_reads_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -1108,39 +1094,16 @@ async def test_service_logger_endusers_failure():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_enduser(enduser):
if enduser["user_id"] == "user1":
raise Exception("Simulated failure for user1")
enduser["spend"] = 0.0
return enduser
async def fake_reset_team_members(budgets_to_reset):
return 1
with (
patch.object(
ResetBudgetJob,
"_reset_budget_for_enduser",
side_effect=fake_reset_enduser,
) as mock_reset_enduser,
patch.object(
ResetBudgetJob,
"reset_budget_for_litellm_team_members",
side_effect=fake_reset_team_members,
) as mock_reset_team_members,
):
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
# Verify exception logging
assert mock_verbose_exc.call_count >= 1
# Verify exception was logged with correct message
assert any(
"Failed to reset budget for enduser" in str(call.args)
for call in mock_verbose_exc.call_args_list
)
with patch(
"litellm.proxy.common_utils.reset_budget_job.verbose_proxy_logger.exception"
) as mock_verbose_exc:
await job.reset_budget_for_litellm_budget_table()
await asyncio.sleep(0.1)
# The log must name the whole cascade, not just end users: the write
# that failed could have been any of team member / enduser / org / tag
# spend or the budget_reset_at advance.
assert mock_verbose_exc.call_count == 1
assert "budget table cascade" in str(mock_verbose_exc.call_args.args[0])
proxy_logging_obj.service_logging_obj.async_service_failure_hook.assert_called_once()
(
@ -1158,8 +1121,8 @@ async def test_service_logger_endusers_failure():
@pytest.mark.asyncio
async def test_reset_budget_for_litellm_team_members_called():
"""
Test that when reset_budget_for_litellm_budget_table is called,
team members' budgets are also reset via reset_budget_for_litellm_team_members
Test that when reset_budget_for_litellm_budget_table is called, team
members' spend is zeroed as part of the cascade transaction.
"""
# Arrange
budget1 = LiteLLM_BudgetTableFull(
@ -1171,7 +1134,7 @@ async def test_reset_budget_for_litellm_team_members_called():
}
)
enduser1 = {"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}
enduser1 = _attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"})
prisma_client = MagicMock()
@ -1184,20 +1147,9 @@ async def test_reset_budget_for_litellm_team_members_called():
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
prisma_client.update_data = AsyncMock()
# Mock the db.litellm_teammembership.update_many call
prisma_client.db = MagicMock()
prisma_client.db.litellm_teammembership = MagicMock()
prisma_client.db.litellm_teammembership.update_many = AsyncMock(
return_value={"count": 2}
)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 0}
)
prisma_client.db.litellm_organizationtable.update_many = AsyncMock(
return_value={"count": 0}
)
prisma_client.db.litellm_tagtable.update_many = AsyncMock(return_value={"count": 0})
batch_calls = _wire_batcher_for_test(prisma_client)
_wire_cascade_reads_for_test(prisma_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj.service_logging_obj = MagicMock()
@ -1206,23 +1158,11 @@ async def test_reset_budget_for_litellm_team_members_called():
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
async def fake_reset_enduser(enduser):
enduser["spend"] = 0.0
return enduser
with patch.object(
ResetBudgetJob,
"_reset_budget_for_enduser",
side_effect=fake_reset_enduser,
):
# Act
await job.reset_budget_for_litellm_budget_table()
# Act
await job.reset_budget_for_litellm_budget_table()
# Assert
# Verify that the team membership update was called
prisma_client.db.litellm_teammembership.update_many.assert_called_once()
# Verify the call was made with correct parameters
call_args = prisma_client.db.litellm_teammembership.update_many.call_args
assert call_args.kwargs["where"]["budget_id"]["in"] == ["budget1"]
assert call_args.kwargs["data"]["spend"] == 0
team_member_writes = [c for c in batch_calls if c["table"] == "team_membership"]
assert len(team_member_writes) == 1
assert team_member_writes[0]["where"]["budget_id"]["in"] == ["budget1"]
assert team_member_writes[0]["data"] == {"spend": 0}

File diff suppressed because it is too large Load diff

View file

@ -72,6 +72,39 @@ async def test_new_budget_success(client_and_mocks):
mock_table.create.assert_awaited_once()
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
@pytest.mark.asyncio
async def test_new_budget_rejects_a_duration_that_never_advances(
client_and_mocks, bad_duration
):
"""A zero-length window resets to "now", so the row is due again the moment
it is written and the reset job re-reads it on every tick forever."""
client, _, mock_table = client_and_mocks
resp = client.post(
"/budget/new",
json={"budget_id": "budget_bad", "max_budget": 10.0, "budget_duration": bad_duration},
)
assert resp.status_code == 400, resp.text
assert "Invalid budget_duration" in resp.json()["detail"]["error"]
mock_table.create.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_budget_rejects_a_duration_that_never_advances(client_and_mocks):
client, _, mock_table = client_and_mocks
resp = client.post(
"/budget/update",
json={"budget_id": "budget_456", "budget_duration": "0s"},
)
assert resp.status_code == 400, resp.text
assert "Invalid budget_duration" in resp.json()["detail"]["error"]
mock_table.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_new_budget_db_not_connected(client_and_mocks, monkeypatch):
client, mock_prisma, mock_table = client_and_mocks

View file

@ -628,6 +628,58 @@ class TestValidateFiniteSpendErrorDetail:
}
class TestValidateBudgetDuration:
"""`validate_budget_duration` keeps durations that never advance out of the
database.
A duration of "0s" resolves to a reset time of now, so the row is due again
the instant it is written. The reset job re-reads such rows on every tick
and, once one tenant owns enough of them, they fill each batch and starve
every other tenant's reset.
"""
def test_none_is_allowed(self):
from litellm.proxy.management_endpoints.common_utils import (
validate_budget_duration,
)
assert validate_budget_duration(None) is None
@pytest.mark.parametrize("duration", ["30s", "5m", "1h", "1d", "7d", "30d", "1mo"])
def test_positive_durations_are_allowed(self, duration):
from litellm.proxy.management_endpoints.common_utils import (
validate_budget_duration,
)
assert validate_budget_duration(duration) is None
@pytest.mark.parametrize("duration", ["0s", "0m", "0h", "0d", "-5m", "abc", ""])
def test_non_advancing_durations_are_rejected(self, duration):
from fastapi import HTTPException
from litellm.proxy.management_endpoints.common_utils import (
validate_budget_duration,
)
with pytest.raises(HTTPException) as exc_info:
validate_budget_duration(duration)
assert exc_info.value.status_code == 400
def test_rejection_detail_is_exact(self):
from fastapi import HTTPException
from litellm.proxy.management_endpoints.common_utils import (
validate_budget_duration,
)
with pytest.raises(HTTPException) as exc_info:
validate_budget_duration("0s")
assert exc_info.value.detail == {
"error": "Invalid budget_duration '0s'. Use a format like '1h', '24h', '7d', or '30d'."
}
class TestRequireCallerUserIdErrorDetail:
"""The 403 for a service-account key must carry the exact error body."""

View file

@ -749,6 +749,40 @@ def test_char_new_body(mock_prisma_client, mock_user_api_key_auth):
assert response.json() == _EXPECTED_CUSTOMER
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
def test_customer_new_rejects_a_duration_that_never_advances(
mock_prisma_client, mock_user_api_key_auth, bad_duration
):
"""A zero-length window resets to "now", leaving the customer's budget row
permanently due for the reset job to re-read every tick."""
mock_prisma_client.db.litellm_endusertable.create = AsyncMock(return_value=_row(_FULL_DB_ROW))
response = client.post(
"/customer/new",
json={"user_id": "c1", "max_budget": 10.0, "budget_duration": bad_duration},
headers={"Authorization": "Bearer k"},
)
assert response.status_code == 400, response.text
assert "Invalid budget_duration" in response.text
mock_prisma_client.db.litellm_endusertable.create.assert_not_awaited()
def test_customer_new_accepts_a_normal_duration(mock_prisma_client, mock_user_api_key_auth):
mock_prisma_client.db.litellm_endusertable.create = AsyncMock(return_value=_row(_FULL_DB_ROW))
mock_prisma_client.db.litellm_budgettable.create = AsyncMock(
return_value=_row({"budget_id": "b1", "max_budget": 10.0})
)
response = client.post(
"/customer/new",
json={"user_id": "c1", "max_budget": 10.0, "budget_duration": "30d"},
headers={"Authorization": "Bearer k"},
)
assert response.status_code == 200, response.text
def test_char_update_body(mock_prisma_client, mock_user_api_key_auth):
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(
return_value=_row({"user_id": "c1", "blocked": False})

View file

@ -788,6 +788,68 @@ def test_update_internal_user_params_reset_spend_and_max_budget():
assert "budget_duration" not in non_default_values # Should not add default values
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
def test_update_internal_user_params_rejects_a_duration_that_never_advances(bad_duration):
"""A zero-length window resets to "now", so the user row is due again the
moment it is written and the reset job re-reads it on every tick. Enough of
them fill each batch and starve other tenants' resets.
"""
from fastapi import HTTPException
from litellm.proxy._types import UpdateUserRequest
from litellm.proxy.management_endpoints.internal_user_endpoints import (
_update_internal_user_params,
)
data = UpdateUserRequest(user_id="test_user_id", budget_duration=bad_duration)
with pytest.raises(HTTPException) as exc_info:
_update_internal_user_params(data_json=data.model_dump(exclude_unset=True), data=data)
assert exc_info.value.status_code == 400
assert "Invalid budget_duration" in str(exc_info.value.detail)
def test_update_internal_user_params_accepts_a_normal_duration():
from litellm.proxy._types import UpdateUserRequest
from litellm.proxy.management_endpoints.internal_user_endpoints import (
_update_internal_user_params,
)
data = UpdateUserRequest(user_id="test_user_id", budget_duration="30d")
non_default_values = _update_internal_user_params(data_json=data.model_dump(exclude_unset=True), data=data)
assert non_default_values["budget_duration"] == "30d"
assert non_default_values["budget_reset_at"] is not None
@pytest.mark.asyncio
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
async def test_new_user_rejects_a_duration_that_never_advances(mocker, bad_duration):
"""/user/new must reject the same never-advancing durations /user/update does."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
mocker.patch("litellm.proxy.proxy_server.prisma_client", MagicMock())
duplicate_check = mocker.patch(
"litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_id",
new=AsyncMock(),
)
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
with pytest.raises(ProxyException) as exc_info:
await new_user(
data=NewUserRequest(budget_duration=bad_duration),
user_api_key_dict=admin,
)
assert str(exc_info.value.code) == "400"
assert "Invalid budget_duration" in str(exc_info.value.message)
duplicate_check.assert_not_awaited()
@pytest.mark.asyncio
async def test_new_user_license_over_limit(mocker):
"""

View file

@ -2527,6 +2527,72 @@ def _setup_update_key_mocks(monkeypatch, mock_prisma_client):
monkeypatch.setattr("litellm.store_audit_logs", False)
@pytest.mark.asyncio
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
async def test_update_key_rejects_a_duration_that_never_advances(monkeypatch, bad_duration):
"""A zero-length window resets to "now", so the key row is due again the
moment it is written. The reset job re-reads such rows on every tick, and a
tenant with enough of them fills each batch and starves other tenants.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
hashed_token = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b"
key_in_db = LiteLLM_VerificationToken(token=hashed_token, user_id="test-user")
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=key_in_db
)
mock_prisma_client.update_data = AsyncMock()
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
with pytest.raises(ProxyException) as exc_info:
await update_key_fn(
request=MagicMock(),
data=UpdateKeyRequest(key=hashed_token, budget_duration=bad_duration),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user"
),
litellm_changed_by=None,
)
assert str(exc_info.value.code) == "400"
assert "Invalid budget_duration" in str(exc_info.value.message)
mock_prisma_client.update_data.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
async def test_generate_key_rejects_a_duration_that_never_advances(monkeypatch, bad_duration):
"""/key/generate must reject the same never-advancing durations /key/update does."""
from litellm.proxy.management_endpoints.key_management_endpoints import (
generate_key_fn,
)
mock_prisma_client = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
new=AsyncMock(),
) as mock_generate:
with pytest.raises(ProxyException) as exc_info:
await generate_key_fn(
data=GenerateKeyRequest(budget_duration=bad_duration),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234", user_id="1234"
),
)
assert str(exc_info.value.code) == "400"
assert "Invalid budget_duration" in str(exc_info.value.message)
mock_generate.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_key_by_alias_only(monkeypatch):
"""
@ -8081,7 +8147,7 @@ async def test_key_with_budget_id_does_not_store_budget_duration():
budget_duration, the key does NOT get budget_duration stored on it.
Keys with budget_id follow their linked budget tier's reset schedule;
reset_budget_for_keys_linked_to_budgets() resets them when the tier resets.
reset_budget_for_litellm_budget_table() resets them when the tier resets.
This avoids duplicating budget_duration on keys so tier updates apply
automatically to all linked keys.
"""

View file

@ -380,6 +380,66 @@ async def test_update_team_permissions_success(mock_db_client, mock_admin_auth):
app.dependency_overrides = {}
@pytest.mark.asyncio
@pytest.mark.parametrize("field", ["budget_duration", "team_member_budget_duration"])
@pytest.mark.parametrize("bad_duration", ["0s", "-5m"])
async def test_new_team_rejects_a_duration_that_never_advances(
mock_db_client, mock_admin_auth, field, bad_duration
):
"""A zero-length window resets to "now", so the team row is due again the
moment it is written. The reset job re-reads such rows on every tick, and a
tenant with enough of them fills each batch and starves other tenants.
"""
from fastapi import Request
from litellm.proxy._types import NewTeamRequest
from litellm.proxy.management_endpoints.team_endpoints import new_team
mock_db_client.db = MagicMock()
mock_team_create = AsyncMock()
mock_db_client.db.litellm_teamtable = MagicMock()
mock_db_client.db.litellm_teamtable.create = mock_team_create
with pytest.raises(ProxyException) as exc_info:
await new_team(
data=NewTeamRequest(team_alias="my-team", **{field: bad_duration}),
http_request=MagicMock(spec=Request),
user_api_key_dict=mock_admin_auth,
)
assert str(exc_info.value.code) == "400"
assert "Invalid budget_duration" in str(exc_info.value.message)
mock_team_create.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("field", ["budget_duration", "team_member_budget_duration"])
async def test_update_team_rejects_a_duration_that_never_advances(
mock_db_client, mock_admin_auth, field
):
"""/team/update must reject the same never-advancing durations /team/new does."""
from fastapi import Request
from litellm.proxy._types import UpdateTeamRequest
from litellm.proxy.management_endpoints.team_endpoints import update_team
mock_db_client.db = MagicMock()
mock_find_unique = AsyncMock(return_value=None)
mock_db_client.db.litellm_teamtable = MagicMock()
mock_db_client.db.litellm_teamtable.find_unique = mock_find_unique
with pytest.raises(ProxyException) as exc_info:
await update_team(
data=UpdateTeamRequest(team_id="team-1", **{field: "0s"}),
http_request=MagicMock(spec=Request),
user_api_key_dict=mock_admin_auth,
)
assert str(exc_info.value.code) == "400"
assert "Invalid budget_duration" in str(exc_info.value.message)
mock_find_unique.assert_not_awaited()
@pytest.mark.asyncio
async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth):
"""

View file

@ -2639,3 +2639,67 @@ async def test_ProxyConfig__init_non_llm_configs_empty_agents_key_clears_remembe
assert clean_agent_registry.config_agents == ()
clean_agent_registry.load_agents_from_db_and_config(db_agents=None)
assert clean_agent_registry.get_agent_list() == ()
# ---------------------------------------------------------------------------
# _init_guardrails_in_db
# ---------------------------------------------------------------------------
def _db_guardrail_row(guardrail_id: str, guardrail_type: str) -> dict[str, object]:
return {
"guardrail_id": guardrail_id,
"guardrail_name": f"name-{guardrail_id}",
"litellm_params": {"guardrail": guardrail_type, "mode": "pre_call"},
"guardrail_info": None,
"team_id": None,
}
@pytest.mark.asyncio
async def test_ProxyConfig__init_guardrails_in_db_skips_only_the_unloadable_row(monkeypatch):
"""
A single DB row that fails to initialize used to abort the whole loop, so one
typo'd guardrail type left the proxy running with zero guardrails loaded.
The failing row's id must still reach reconcile_db_guardrails so that eviction
pass cannot treat a row that is alive in the DB as one that was deleted.
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy.guardrails import guardrail_registry as registry_module
from litellm.types.guardrails import Guardrail, GuardrailEventHooks, LitellmParams
class _RecordingHandler(registry_module.InMemoryGuardrailHandler):
def __init__(self) -> None:
super().__init__()
self.reconciled_with: list[set[str]] = []
def reconcile_db_guardrails(self, db_guardrail_ids: set[str]) -> list[str]:
self.reconciled_with.append(set(db_guardrail_ids))
return super().reconcile_db_guardrails(db_guardrail_ids)
handler = _RecordingHandler()
monkeypatch.setattr(registry_module, "IN_MEMORY_GUARDRAIL_HANDLER", handler)
def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail:
return CustomGuardrail(
guardrail_name=guardrail["guardrail_name"],
event_hook=GuardrailEventHooks.pre_call,
default_on=False,
)
monkeypatch.setitem(registry_module.guardrail_initializer_registry, "lit5367_ok", _initializer)
prisma_client = MagicMock()
prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(
return_value=[
_db_guardrail_row("first", "lit5367_ok"),
_db_guardrail_row("broken", "litellm_tool_permission"),
_db_guardrail_row("last", "lit5367_ok"),
]
)
await ProxyConfig()._init_guardrails_in_db(prisma_client=prisma_client)
assert sorted(handler.IN_MEMORY_GUARDRAILS) == ["first", "last"]
assert handler.reconciled_with == [{"first", "broken", "last"}]

View file

@ -1,6 +1,7 @@
import sys
from types import ModuleType, SimpleNamespace
from litellm.proxy._lazy_features import LazyFeature
from litellm.proxy._lazy_openapi_snapshot import _normalize_operation_ids
@ -22,22 +23,20 @@ def test_generate_snapshot_uses_shared_operation_id_reservations(monkeypatch):
fake_lazy_features_module = ModuleType("litellm.proxy._lazy_features")
fake_lazy_features_module.LAZY_FEATURES = [
SimpleNamespace(
LazyFeature(
name="feature-a",
module_path="fake_feature_a",
path_prefixes=("/feature-a",),
register_fn=lambda app, module: None,
),
SimpleNamespace(
LazyFeature(
name="feature-b",
module_path="fake_feature_b",
path_prefixes=("/feature-b",),
register_fn=lambda app, module: None,
),
]
monkeypatch.setitem(
sys.modules, "litellm.proxy._lazy_features", fake_lazy_features_module
)
monkeypatch.setitem(sys.modules, "litellm.proxy._lazy_features", fake_lazy_features_module)
def fake_get_openapi(title, version, routes):
path = routes[0].path
@ -58,30 +57,59 @@ def test_generate_snapshot_uses_shared_operation_id_reservations(monkeypatch):
fake_proxy_server_module = ModuleType("litellm.proxy.proxy_server")
fake_proxy_server_module.app = fake_app
fake_proxy_server_module.ensure_unique_openapi_operation_ids = (
fake_ensure_unique_openapi_operation_ids
)
monkeypatch.setitem(
sys.modules, "litellm.proxy.proxy_server", fake_proxy_server_module
)
fake_proxy_server_module.ensure_unique_openapi_operation_ids = fake_ensure_unique_openapi_operation_ids
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy_server_module)
monkeypatch.setattr("fastapi.openapi.utils.get_openapi", fake_get_openapi)
fragments = _lazy_openapi_snapshot.generate_snapshot()
assert (
fragments["feature-a"]["paths"]["/feature-a/items"]["get"]["operationId"]
== "shared_operation_id_get"
)
assert (
fragments["feature-b"]["paths"]["/feature-b/items"]["get"]["operationId"]
== "shared_operation_id_get_2"
)
assert fragments["feature-a"]["paths"]["/feature-a/items"]["get"]["tags"] == [
"feature-a"
]
assert fragments["feature-b"]["paths"]["/feature-b/items"]["get"]["tags"] == [
"feature-b"
assert fragments["feature-a"]["paths"]["/feature-a/items"]["get"]["operationId"] == "shared_operation_id_get"
assert fragments["feature-b"]["paths"]["/feature-b/items"]["get"]["operationId"] == "shared_operation_id_get_2"
assert fragments["feature-a"]["paths"]["/feature-a/items"]["get"]["tags"] == ["feature-a"]
assert fragments["feature-b"]["paths"]["/feature-b/items"]["get"]["tags"] == ["feature-b"]
def test_generate_snapshot_registers_transitively_imported_modules(monkeypatch):
"""A feature module already in sys.modules (pulled in transitively by an
earlier feature) must still get register_fn called, else its routes never
mount and its fragment silently vanishes from the snapshot. Fragment
collection must also honor path_suffixes, not just prefixes."""
from litellm.proxy import _lazy_openapi_snapshot
fake_app = SimpleNamespace(title="LiteLLM test", version="0.0.0", routes=[])
fake_module = ModuleType("fake_transitive_feature")
monkeypatch.setitem(sys.modules, "fake_transitive_feature", fake_module)
def register_fn(app, module):
app.routes.append(SimpleNamespace(path="/transitive/items"))
app.routes.append(SimpleNamespace(path="/v1/{param}/deep/leaf"))
fake_lazy_features_module = ModuleType("litellm.proxy._lazy_features")
fake_lazy_features_module.LAZY_FEATURES = [
LazyFeature(
name="transitive",
module_path="fake_transitive_feature",
path_prefixes=("/transitive",),
path_suffixes=("/deep/leaf",),
register_fn=register_fn,
)
]
monkeypatch.setitem(sys.modules, "litellm.proxy._lazy_features", fake_lazy_features_module)
def fake_get_openapi(title, version, routes):
return {"paths": {route.path: {"get": {"operationId": f"op{i}_get"}} for i, route in enumerate(routes)}}
fake_proxy_server_module = ModuleType("litellm.proxy.proxy_server")
fake_proxy_server_module.app = fake_app
fake_proxy_server_module.ensure_unique_openapi_operation_ids = lambda schema, reserved_operation_ids: schema
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy_server_module)
monkeypatch.setattr("fastapi.openapi.utils.get_openapi", fake_get_openapi)
fragments = _lazy_openapi_snapshot.generate_snapshot()
assert fragments["transitive"]["paths"]["/transitive/items"]["get"]["tags"] == ["transitive"]
assert "/v1/{param}/deep/leaf" in fragments["transitive"]["paths"]
def test_normalize_operation_ids_uses_each_http_method():

View file

@ -16,10 +16,11 @@ import litellm
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
LiteLLM_ManagedVectorStore,
)
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.vector_store_endpoints.endpoints import (
_update_request_data_with_litellm_managed_vector_store_registry,
index_create,
index_list,
)
from litellm.proxy.vector_store_files_endpoints.endpoints import (
_update_request_data_with_model_routing_hint,
@ -37,7 +38,7 @@ from litellm.proxy.vector_store_endpoints.utils import (
is_allowed_to_call_vector_store_endpoint,
is_allowed_to_call_vector_store_files_endpoint,
)
from litellm.types.vector_stores import IndexCreateRequest
from litellm.types.vector_stores import IndexCreateRequest, IndexListResponse
from litellm.types.utils import LlmProviders
@ -1316,6 +1317,93 @@ class TestIndexCreate:
mock_prisma.db.litellm_managedvectorstoreindextable.create.assert_awaited_once()
class TestIndexList:
def _admin(self) -> UserAPIKeyAuth:
return UserAPIKeyAuth(
token="sk-test",
key_name="sk-...test",
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id="admin-user",
)
def _index_row(self, index_id: str, index_name: str) -> dict:
return {
"id": index_id,
"index_name": index_name,
"litellm_params": {
"vector_store_index": f"real-{index_name}",
"vector_store_name": "azure-ai-search",
},
"index_info": None,
"created_at": datetime(2026, 1, 2, tzinfo=timezone.utc),
"created_by": "admin-user",
"updated_at": datetime(2026, 1, 2, tzinfo=timezone.utc),
"updated_by": "admin-user",
}
@pytest.mark.asyncio
async def test_index_list_requires_admin(self):
"""Index topology must never reach non-admins, not even via a DB read."""
mock_prisma = MagicMock()
mock_prisma.db.litellm_managedvectorstoreindextable.find_many = AsyncMock()
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
with pytest.raises(HTTPException) as exc_info:
await index_list(
user_api_key_dict=UserAPIKeyAuth(
token="sk-test",
key_name="sk-...test",
user_role=LitellmUserRoles.INTERNAL_USER,
)
)
assert exc_info.value.status_code == 403
assert "Only proxy admins can list" in exc_info.value.detail
mock_prisma.db.litellm_managedvectorstoreindextable.find_many.assert_not_awaited()
@pytest.mark.asyncio
async def test_index_list_requires_db_connection(self):
with patch("litellm.proxy.proxy_server.prisma_client", None):
with pytest.raises(HTTPException) as exc_info:
await index_list(user_api_key_dict=self._admin())
assert exc_info.value.status_code == 500
assert CommonProxyErrors.db_not_connected_error.value in exc_info.value.detail
@pytest.mark.asyncio
async def test_index_list_returns_db_rows_newest_first(self):
"""Rows round-trip into typed models and DB ordering (created_at desc) is requested."""
rows = [
self._index_row("idx-2", "index-b"),
self._index_row("idx-1", "index-a"),
]
mock_prisma = MagicMock()
mock_prisma.db.litellm_managedvectorstoreindextable.find_many = AsyncMock(return_value=rows)
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
result = await index_list(user_api_key_dict=self._admin())
assert isinstance(result, IndexListResponse)
assert result.object == "list"
assert [index.index_name for index in result.data] == ["index-b", "index-a"]
assert result.data[0].litellm_params.vector_store_index == "real-index-b"
assert result.data[0].litellm_params.vector_store_name == "azure-ai-search"
assert result.data[1].litellm_params.vector_store_index == "real-index-a"
mock_prisma.db.litellm_managedvectorstoreindextable.find_many.assert_awaited_once_with(
order={"created_at": "desc"}
)
def test_get_v1_indexes_route_registered(self):
from litellm.proxy.vector_store_endpoints.endpoints import router
routes = [
(method, getattr(route, "path", None))
for route in router.routes
for method in (getattr(route, "methods", None) or ())
]
assert ("GET", "/v1/indexes") in routes
class TestIsAllowedToCallVectorStoreFilesEndpoint:
def _mock_provider_config(self):
provider_config = MagicMock()

View file

@ -3,7 +3,10 @@ from typing import Any, Dict, List, Mapping, Tuple
import pytest
from litellm.repositories.unit_of_work import spend_reset_unit_of_work
from litellm.repositories.unit_of_work import (
budget_cascade_unit_of_work,
spend_reset_unit_of_work,
)
class FakeBatchTable:
@ -14,6 +17,9 @@ class FakeBatchTable:
def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> None:
self._calls.append((self._table_name, dict(where), dict(data)))
def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> None:
self._calls.append((f"{self._table_name}.update_many", dict(where), dict(data)))
class FakeBatch:
def __init__(self):
@ -22,6 +28,11 @@ class FakeBatch:
self.litellm_verificationtoken = FakeBatchTable("litellm_verificationtoken", self.calls)
self.litellm_usertable = FakeBatchTable("litellm_usertable", self.calls)
self.litellm_teamtable = FakeBatchTable("litellm_teamtable", self.calls)
self.litellm_budgettable = FakeBatchTable("litellm_budgettable", self.calls)
self.litellm_teammembership = FakeBatchTable("litellm_teammembership", self.calls)
self.litellm_organizationtable = FakeBatchTable("litellm_organizationtable", self.calls)
self.litellm_tagtable = FakeBatchTable("litellm_tagtable", self.calls)
self.litellm_endusertable = FakeBatchTable("litellm_endusertable", self.calls)
async def commit(self) -> None:
self.commit_count += 1
@ -64,3 +75,53 @@ async def test_empty_block_still_commits_the_batch():
assert batch.commit_count == 1
assert batch.calls == []
async def test_budget_cascade_dependents_and_window_advance_share_one_batch():
batch = FakeBatch()
reset_at = datetime(2026, 8, 3, 12, 0, tzinfo=timezone.utc)
linked = {"budget_id": {"in": ["budget-1"]}}
async with budget_cascade_unit_of_work(lambda: batch) as uow:
uow.team_memberships.queue_spend_zero(where=linked)
uow.keys.queue_spend_zero(where=linked)
uow.organizations.queue_spend_zero(where=linked)
uow.tags.queue_spend_zero(where=linked)
uow.endusers.queue_spend_zero(where={"user_id": {"in": ["enduser-1"]}})
uow.budgets.queue_window_advance(budget_id="budget-1", budget_reset_at=reset_at)
assert batch.commit_count == 0
assert batch.commit_count == 1
assert batch.calls == [
("litellm_teammembership.update_many", linked, {"spend": 0}),
("litellm_verificationtoken.update_many", linked, {"spend": 0}),
("litellm_organizationtable.update_many", linked, {"spend": 0}),
("litellm_tagtable.update_many", linked, {"spend": 0}),
("litellm_endusertable.update_many", {"user_id": {"in": ["enduser-1"]}}, {"spend": 0}),
("litellm_budgettable.update_many", {"budget_id": "budget-1"}, {"budget_reset_at": reset_at}),
]
async def test_budget_window_advance_tolerates_a_tier_deleted_mid_chunk():
"""A tier deleted between the read and the commit must not abort the batch:
``update`` raises P2025 on a missing row and takes every other write in the
chunk down with it, while ``update_many`` just matches nothing."""
batch = FakeBatch()
async with budget_cascade_unit_of_work(lambda: batch) as uow:
uow.budgets.queue_window_advance(budget_id="budget-1", budget_reset_at=datetime.now(timezone.utc))
assert [call[0] for call in batch.calls] == ["litellm_budgettable.update_many"]
async def test_budget_cascade_raising_inside_block_skips_commit():
"""A failure part-way through must leave budget_reset_at where it was, so
the tier is still due on the next tick."""
batch = FakeBatch()
with pytest.raises(RuntimeError, match="boom"):
async with budget_cascade_unit_of_work(lambda: batch) as uow:
uow.team_memberships.queue_spend_zero(where={"budget_id": {"in": ["budget-1"]}})
raise RuntimeError("boom")
assert batch.commit_count == 0

View file

@ -117,6 +117,17 @@ def test_typing_alias_and_forward_ref_annotations_are_flagged(tmp_path):
assert "LIT001" in _codes(tmp_path, 'x: "dict[str, int]"\n')
def test_literal_string_args_are_values_not_forward_refs(tmp_path):
assert "LIT001" not in _codes(tmp_path, 'from typing import Literal\nx: Literal["list"] = "list"\n')
assert "LIT001" not in _codes(
tmp_path,
'from typing import Literal\ndef f(op: Literal["create", "list"] = "create") -> None:\n return None\n',
)
assert "LIT001" not in _codes(tmp_path, 'import typing\nx: typing.Literal["dict"] = "dict"\n')
assert "LIT001" in _codes(tmp_path, 'from typing import Literal\nx: dict[str, Literal["a"]]\n')
assert "LIT001" in _codes(tmp_path, "x: \"Literal['x'] | list[int]\"\n")
def test_readonly_annotations_are_clean(tmp_path):
for ann in ("Mapping[str, int]", "Sequence[int]", "tuple[int, ...]", "frozenset[int]"):
assert "LIT001" not in _codes(tmp_path, f"from typing import Mapping, Sequence\nx: {ann}\n")

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 23064
"limit": 23057
},
"LIT002": {
"limit": 27166
"limit": 27156
},
"LIT003": {
"limit": 269
@ -27,9 +27,9 @@
"limit": 0
},
"LIT010": {
"limit": 16753
"limit": 16744
},
"LIT011": {
"limit": 5598
"limit": 5596
}
}

View file

@ -1,4 +1,4 @@
import { act, fireEvent, render, screen, waitFor, within } from "@testing-library/react";
import { act, cleanup, fireEvent, render, screen, waitFor, within } from "@testing-library/react";
import { beforeAll, beforeEach, describe, expect, it, vi } from "vitest";
import * as networking from "@/components/networking";
import EntityUsage from "./EntityUsage";
@ -39,11 +39,21 @@ vi.mock("../EndpointUsage/EndpointUsage", () => ({
}));
vi.mock("@/components/UsagePage/components/EntityUsage/TopKeyView", () => ({
default: () => <div>Top Keys</div>,
default: ({ topKeys }: { topKeys: { api_key: string; spend: number }[] }) => (
<div>
<span>Top Keys</span>
<span>{`top-keys:${topKeys.map((row) => `${row.api_key}=${row.spend}`).join("|")}`}</span>
</div>
),
}));
vi.mock("./TopModelView", () => ({
default: () => <div>Top Models</div>,
default: ({ topModels }: { topModels: { key: string; spend: number }[] }) => (
<div>
<span>Top Models</span>
<span>{`top-models:${topModels.map((row) => `${row.key}=${row.spend}`).join("|")}`}</span>
</div>
),
}));
vi.mock("@/components/EntityUsageExport/EntityUsageExportModal", () => ({
@ -856,6 +866,47 @@ describe("EntityUsage", () => {
expect(logo.getAttribute("src")).toContain("openai_small");
});
describe("capability gating", () => {
it.each([
["organization", () => mockOrganizationDailyActivityCall, "Organization Spend Overview"],
["agent", () => mockAgentDailyActivityCall, "Agent Spend Overview"],
] as const)("fetches %s activity for an admin but not for an internal user", async (entityType, call, heading) => {
render(<EntityUsage {...defaultProps} entityType={entityType} />);
await waitFor(() => {
expect(call()).toHaveBeenCalled();
});
cleanup();
call().mockClear();
render(<EntityUsage {...defaultProps} entityType={entityType} userRole="Internal User" />);
expect(await screen.findByText(heading)).toBeInTheDocument();
expect(call()).not.toHaveBeenCalled();
});
it("keeps the team breakdown but drops its agent sub-fetch for an internal user", async () => {
render(<EntityUsage {...defaultProps} entityType="team" userRole="Internal User" />);
await waitFor(() => {
expect(mockTeamDailyActivityCall).toHaveBeenCalled();
});
expect(screen.getByText("Team Spend Overview")).toBeInTheDocument();
expect(mockAgentDailyActivityCall).not.toHaveBeenCalled();
expect(screen.queryByText("Agent Activity")).not.toBeInTheDocument();
expect(screen.queryByText("Top Agents Driving Spend")).not.toBeInTheDocument();
});
it("keeps the tag breakdown for an internal user", async () => {
render(<EntityUsage {...defaultProps} entityType="tag" userRole="Internal User" />);
await waitFor(() => {
expect(mockTagDailyActivityCall).toHaveBeenCalled();
});
expect(screen.getByText("Tag Spend Overview")).toBeInTheDocument();
});
});
it("renders a letter avatar instead of an img for an unknown provider slug", async () => {
const spendDataUnknownProvider = {
...mockSpendData,
@ -881,4 +932,39 @@ describe("EntityUsage", () => {
expect(screen.queryByAltText("zzz-internal logo")).not.toBeInTheDocument();
expect(screen.getByText("z")).toBeInTheDocument();
});
it("feeds the key, model and agent tables from their own breakdowns", async () => {
const usageMetrics = {
spend: 30.75,
api_requests: 300,
successful_requests: 290,
failed_requests: 10,
total_tokens: 15000,
prompt_tokens: 9000,
completion_tokens: 6000,
cache_read_input_tokens: 0,
cache_creation_input_tokens: 0,
};
mockTeamDailyActivityCall.mockResolvedValue({
...mockSpendData,
results: [
{
...mockSpendData.results[0],
breakdown: {
...mockSpendData.results[0].breakdown,
model_groups: { "gpt-4o": { metrics: { ...usageMetrics, spend: 70.25 }, metadata: {} } },
api_keys: { "sk-abc": { metrics: usageMetrics, metadata: { key_alias: "prod-key", team_id: null } } },
},
},
],
});
render(<EntityUsage {...defaultProps} entityType="team" />);
await waitFor(() => {
expect(screen.getByText("top-keys:sk-abc=30.75")).toBeInTheDocument();
});
expect(screen.getByText("top-models:gpt-4o=70.25")).toBeInTheDocument();
expect(screen.getByText(/^top-models:Code Review Agent=/)).toBeInTheDocument();
});
});

View file

@ -1,8 +1,16 @@
import useTeams from "@/app/(dashboard)/hooks/useTeams";
import { BarChart, DonutChart } from "@/components/shared/charts";
import {
getProviderSpend,
getTopAgents,
getTopAPIKeys,
getTopModels,
type ExtendedDailyData,
} from "./entityUsageAggregations";
import { buildCostBreakdownTiles, buildSummaryTiles, hasFlatCost, type SummaryTile } from "./entityUsageSummary";
import { MoneyCell } from "@/components/shared/table_cells";
import { Card as ShadcnCard, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
import { hasCapability, type Capability } from "@/utils/capabilities";
import { formatNumberWithCommas } from "@/utils/dataUtils";
import {
Card,
@ -41,13 +49,7 @@ import {
} from "@/components/networking";
import { Logo } from "@/components/molecules/logo/Logo";
import { usePaginatedDailyActivity } from "../../hooks/usePaginatedDailyActivity";
import {
BreakdownMetrics,
DailyData,
EntityMetricWithMetadata,
KeyMetricWithMetadata,
TagUsage,
} from "@/components/UsagePage/types";
import { EntityMetricWithMetadata } from "@/components/UsagePage/types";
import { valueFormatterSpend } from "@/components/UsagePage/utils/value_formatters";
import EndpointUsage from "../EndpointUsage/EndpointUsage";
import ModelViewToggle, { ModelViewType } from "../ModelViewToggle";
@ -69,10 +71,6 @@ interface EntityMetrics {
metadata: Record<string, any>;
}
type ExtendedDailyData = DailyData & {
breakdown: BreakdownMetrics;
};
interface EntitySpendData {
results: ExtendedDailyData[];
metadata: {
@ -110,7 +108,19 @@ const ENTITY_FETCH_FNS: Record<EntityType, (...args: any[]) => Promise<any>> = {
user: userDailyActivityCall,
};
const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, entityId, entityList, dateValue }) => {
const ENTITY_CAPABILITIES: Partial<Record<EntityType, Capability>> = {
organization: "viewOrganizationUsage",
agent: "viewAgentUsage",
};
const EntityUsage: React.FC<EntityUsageProps> = ({
accessToken,
entityType,
entityId,
entityList,
userRole,
dateValue,
}) => {
const { teams } = useTeams();
const [selectedTags, setSelectedTags] = useState<string[]>([]);
const [modelViewType, setModelViewType] = useState<ModelViewType>("groups");
@ -128,7 +138,11 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
}, [entityType, selectedTags]);
const fetchFn = ENTITY_FETCH_FNS[entityType];
const enabled = !!accessToken && !!startTime && !!endTime;
const entityCapability = ENTITY_CAPABILITIES[entityType];
const canViewEntity = entityCapability === undefined || hasCapability(userRole, entityCapability);
const showAgentBreakdown = entityType === "team" && hasCapability(userRole, "viewAgentUsage");
const hasRequestWindow = !!accessToken && !!startTime && !!endTime;
const enabled = hasRequestWindow && canViewEntity;
const {
data: spendDataRaw,
@ -153,7 +167,7 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
} = usePaginatedDailyActivity({
fetchFn: agentDailyActivityCall,
args: [accessToken, startTime, endTime, null],
enabled: enabled && entityType === "team",
enabled: enabled && showAgentBreakdown,
});
const agentSpendData = agentSpendDataRaw as unknown as EntitySpendData;
@ -161,164 +175,7 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
const modelBreakdownKey = modelViewType === "groups" ? "model_groups" : "models";
const modelMetrics = processActivityData(spendData, modelBreakdownKey, teams || []);
const keyMetrics = processActivityData(spendData, "api_keys", teams || []);
const agentMetrics = entityType === "team" ? processActivityData(agentSpendData, "entities", teams || []) : {};
const getTopModels = () => {
const modelSpend: { [key: string]: any } = {};
spendData.results.forEach((day) => {
Object.entries(day.breakdown[modelBreakdownKey] || {}).forEach(([model, metrics]) => {
if (!modelSpend[model]) {
modelSpend[model] = {
spend: 0,
requests: 0,
successful_requests: 0,
failed_requests: 0,
tokens: 0,
};
}
try {
modelSpend[model].spend += metrics.metrics.spend;
} catch (e) {
console.error(`Error adding spend for ${model}: ${e}, got metrics: ${JSON.stringify(metrics)}`);
}
modelSpend[model].requests += metrics.metrics.api_requests;
modelSpend[model].successful_requests += metrics.metrics.successful_requests;
modelSpend[model].failed_requests += metrics.metrics.failed_requests;
modelSpend[model].tokens += metrics.metrics.total_tokens;
});
});
return Object.entries(modelSpend)
.map(([model, metrics]) => ({
key: model,
...metrics,
}))
.sort((a, b) => b.spend - a.spend)
.slice(0, topModelsLimit);
};
const getTopAgents = () => {
const agentSpend: { [key: string]: any } = {};
agentSpendData.results.forEach((day) => {
Object.entries(day.breakdown.entities || {}).forEach(([agentId, data]) => {
if (!agentSpend[agentId]) {
agentSpend[agentId] = {
spend: 0,
requests: 0,
successful_requests: 0,
failed_requests: 0,
tokens: 0,
agent_name: (data.metadata as any)?.agent_name || agentId,
};
}
agentSpend[agentId].spend += data.metrics.spend;
agentSpend[agentId].requests += data.metrics.api_requests;
agentSpend[agentId].successful_requests += data.metrics.successful_requests;
agentSpend[agentId].failed_requests += data.metrics.failed_requests;
agentSpend[agentId].tokens += data.metrics.total_tokens;
});
});
return Object.entries(agentSpend)
.map(([agentId, metrics]) => ({
key: metrics.agent_name,
...metrics,
}))
.sort((a, b) => b.spend - a.spend)
.slice(0, topAgentsLimit);
};
const getTopAPIKeys = () => {
const keySpend: { [key: string]: KeyMetricWithMetadata } = {};
spendData.results.forEach((day) => {
const { breakdown } = day;
const { entities } = breakdown;
const tagDictionary = Object.keys(entities).reduce((acc: { [key: string]: TagUsage[] }, entity) => {
const { api_key_breakdown } = entities[entity];
Object.keys(api_key_breakdown).forEach((key) => {
const tagUsage = { tag: entity, usage: api_key_breakdown[key].metrics.spend };
if (acc[key]) {
acc[key].push(tagUsage);
} else {
acc[key] = [tagUsage];
}
});
return acc;
}, {});
Object.entries(day.breakdown.api_keys || {}).forEach(([key, metrics]) => {
if (!keySpend[key]) {
keySpend[key] = {
metrics: {
spend: 0,
prompt_tokens: 0,
completion_tokens: 0,
total_tokens: 0,
api_requests: 0,
successful_requests: 0,
failed_requests: 0,
cache_read_input_tokens: 0,
cache_creation_input_tokens: 0,
},
metadata: {
key_alias: metrics.metadata.key_alias,
team_id: metrics.metadata.team_id || null,
tags: tagDictionary[key] || [],
},
};
}
keySpend[key].metrics.spend += metrics.metrics.spend;
keySpend[key].metrics.prompt_tokens += metrics.metrics.prompt_tokens;
keySpend[key].metrics.completion_tokens += metrics.metrics.completion_tokens;
keySpend[key].metrics.total_tokens += metrics.metrics.total_tokens;
keySpend[key].metrics.api_requests += metrics.metrics.api_requests;
keySpend[key].metrics.successful_requests += metrics.metrics.successful_requests;
keySpend[key].metrics.failed_requests += metrics.metrics.failed_requests;
keySpend[key].metrics.cache_read_input_tokens += metrics.metrics.cache_read_input_tokens || 0;
keySpend[key].metrics.cache_creation_input_tokens += metrics.metrics.cache_creation_input_tokens || 0;
});
});
return Object.entries(keySpend)
.map(([api_key, metrics]) => ({
api_key,
key_alias: metrics.metadata.key_alias || "-", // Using truncated key as alias
tags: metrics.metadata.tags || "-",
spend: metrics.metrics.spend,
}))
.sort((a, b) => b.spend - a.spend)
.slice(0, topKeysLimit);
};
const getProviderSpend = () => {
const providerSpend: { [key: string]: any } = {};
spendData.results.forEach((day) => {
Object.entries(day.breakdown.providers || {}).forEach(([provider, metrics]) => {
if (!providerSpend[provider]) {
providerSpend[provider] = {
provider,
spend: 0,
requests: 0,
successful_requests: 0,
failed_requests: 0,
tokens: 0,
};
}
try {
providerSpend[provider].spend += metrics.metrics.spend;
providerSpend[provider].requests += metrics.metrics.api_requests;
providerSpend[provider].successful_requests += metrics.metrics.successful_requests;
providerSpend[provider].failed_requests += metrics.metrics.failed_requests;
providerSpend[provider].tokens += metrics.metrics.total_tokens;
} catch (e) {
console.error(`Error processing provider ${provider}: ${e}`);
}
});
});
return Object.values(providerSpend)
.filter((provider) => provider.spend > 0)
.sort((a, b) => b.spend - a.spend);
};
const agentMetrics = showAgentBreakdown ? processActivityData(agentSpendData, "entities", teams || []) : {};
const getAllTags = () => {
if (entityList) {
@ -616,7 +473,7 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
<Card>
<Title>Top Virtual Keys</Title>
<TopKeyView
topKeys={getTopAPIKeys()}
topKeys={getTopAPIKeys(spendData.results, topKeysLimit)}
teams={null}
showTags={entityType === "tag"}
topKeysLimit={topKeysLimit}
@ -633,20 +490,19 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
<ModelViewToggle value={modelViewType} onChange={setModelViewType} />
</div>
<TopModelView
topModels={getTopModels()}
topModels={getTopModels(spendData.results, modelBreakdownKey, topModelsLimit)}
topModelsLimit={topModelsLimit}
setTopModelsLimit={setTopModelsLimit}
/>
</Card>
</Col>
{/* Top Agents - only for team entity type */}
{entityType === "team" && (
{showAgentBreakdown && (
<Col numColSpan={2}>
<Card>
<Title>Top Agents Driving Spend</Title>
<TopModelView
topModels={getTopAgents()}
topModels={getTopAgents(agentSpendData.results, topAgentsLimit)}
topModelsLimit={topAgentsLimit}
setTopModelsLimit={setTopAgentsLimit}
/>
@ -663,7 +519,7 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
<Col numColSpan={1}>
<DonutChart
className="mt-4 h-40"
data={getProviderSpend()}
data={getProviderSpend(spendData.results)}
index="provider"
category="spend"
valueFormatter={(value) => `$${formatNumberWithCommas(value, 2)}`}
@ -685,7 +541,7 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
</TableRow>
</TableHead>
<TableBody>
{getProviderSpend().map((provider) => (
{getProviderSpend(spendData.results).map((provider) => (
<TableRow key={provider.provider}>
<TableCell>
<div className="flex items-center space-x-2">
@ -727,7 +583,7 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
</>
),
},
...(entityType === "team"
...(showAgentBreakdown
? [{ key: "agents", label: "Agent Activity", content: <ActivityMetrics modelMetrics={agentMetrics} /> }]
: []),
{
@ -776,7 +632,7 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
}
/>
)}
{agentIsFetchingMore && entityType === "team" && (
{agentIsFetchingMore && showAgentBreakdown && (
<Alert
banner
type="warning"
@ -800,7 +656,7 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
}
/>
)}
{agentCancelled && entityType === "team" && (
{agentCancelled && showAgentBreakdown && (
<Alert
banner
type="info"

View file

@ -0,0 +1,168 @@
import { BreakdownMetrics, DailyData, KeyMetricWithMetadata, TagUsage } from "@/components/UsagePage/types";
export type ExtendedDailyData = DailyData & {
breakdown: BreakdownMetrics;
};
export type ModelBreakdownKey = "models" | "model_groups";
export const getTopModels = (
results: ExtendedDailyData[],
modelBreakdownKey: ModelBreakdownKey,
topModelsLimit: number,
) => {
const modelSpend: { [key: string]: any } = {};
results.forEach((day) => {
Object.entries(day.breakdown[modelBreakdownKey] || {}).forEach(([model, metrics]) => {
if (!modelSpend[model]) {
modelSpend[model] = {
spend: 0,
requests: 0,
successful_requests: 0,
failed_requests: 0,
tokens: 0,
};
}
try {
modelSpend[model].spend += metrics.metrics.spend;
} catch (e) {
console.error(`Error adding spend for ${model}: ${e}, got metrics: ${JSON.stringify(metrics)}`);
}
modelSpend[model].requests += metrics.metrics.api_requests;
modelSpend[model].successful_requests += metrics.metrics.successful_requests;
modelSpend[model].failed_requests += metrics.metrics.failed_requests;
modelSpend[model].tokens += metrics.metrics.total_tokens;
});
});
return Object.entries(modelSpend)
.map(([model, metrics]) => ({
key: model,
...metrics,
}))
.sort((a, b) => b.spend - a.spend)
.slice(0, topModelsLimit);
};
export const getTopAgents = (results: ExtendedDailyData[], topAgentsLimit: number) => {
const agentSpend: { [key: string]: any } = {};
results.forEach((day) => {
Object.entries(day.breakdown.entities || {}).forEach(([agentId, data]) => {
if (!agentSpend[agentId]) {
agentSpend[agentId] = {
spend: 0,
requests: 0,
successful_requests: 0,
failed_requests: 0,
tokens: 0,
agent_name: (data.metadata as any)?.agent_name || agentId,
};
}
agentSpend[agentId].spend += data.metrics.spend;
agentSpend[agentId].requests += data.metrics.api_requests;
agentSpend[agentId].successful_requests += data.metrics.successful_requests;
agentSpend[agentId].failed_requests += data.metrics.failed_requests;
agentSpend[agentId].tokens += data.metrics.total_tokens;
});
});
return Object.entries(agentSpend)
.map(([agentId, metrics]) => ({
key: metrics.agent_name,
...metrics,
}))
.sort((a, b) => b.spend - a.spend)
.slice(0, topAgentsLimit);
};
export const getTopAPIKeys = (results: ExtendedDailyData[], topKeysLimit: number) => {
const keySpend: { [key: string]: KeyMetricWithMetadata } = {};
results.forEach((day) => {
const { breakdown } = day;
const { entities } = breakdown;
const tagDictionary = Object.keys(entities).reduce((acc: { [key: string]: TagUsage[] }, entity) => {
const { api_key_breakdown } = entities[entity];
Object.keys(api_key_breakdown).forEach((key) => {
const tagUsage = { tag: entity, usage: api_key_breakdown[key].metrics.spend };
if (acc[key]) {
acc[key].push(tagUsage);
} else {
acc[key] = [tagUsage];
}
});
return acc;
}, {});
Object.entries(day.breakdown.api_keys || {}).forEach(([key, metrics]) => {
if (!keySpend[key]) {
keySpend[key] = {
metrics: {
spend: 0,
prompt_tokens: 0,
completion_tokens: 0,
total_tokens: 0,
api_requests: 0,
successful_requests: 0,
failed_requests: 0,
cache_read_input_tokens: 0,
cache_creation_input_tokens: 0,
},
metadata: {
key_alias: metrics.metadata.key_alias,
team_id: metrics.metadata.team_id || null,
tags: tagDictionary[key] || [],
},
};
}
keySpend[key].metrics.spend += metrics.metrics.spend;
keySpend[key].metrics.prompt_tokens += metrics.metrics.prompt_tokens;
keySpend[key].metrics.completion_tokens += metrics.metrics.completion_tokens;
keySpend[key].metrics.total_tokens += metrics.metrics.total_tokens;
keySpend[key].metrics.api_requests += metrics.metrics.api_requests;
keySpend[key].metrics.successful_requests += metrics.metrics.successful_requests;
keySpend[key].metrics.failed_requests += metrics.metrics.failed_requests;
keySpend[key].metrics.cache_read_input_tokens += metrics.metrics.cache_read_input_tokens || 0;
keySpend[key].metrics.cache_creation_input_tokens += metrics.metrics.cache_creation_input_tokens || 0;
});
});
return Object.entries(keySpend)
.map(([api_key, metrics]) => ({
api_key,
key_alias: metrics.metadata.key_alias || "-", // Using truncated key as alias
tags: metrics.metadata.tags || "-",
spend: metrics.metrics.spend,
}))
.sort((a, b) => b.spend - a.spend)
.slice(0, topKeysLimit);
};
export const getProviderSpend = (results: ExtendedDailyData[]) => {
const providerSpend: { [key: string]: any } = {};
results.forEach((day) => {
Object.entries(day.breakdown.providers || {}).forEach(([provider, metrics]) => {
if (!providerSpend[provider]) {
providerSpend[provider] = {
provider,
spend: 0,
requests: 0,
successful_requests: 0,
failed_requests: 0,
tokens: 0,
};
}
try {
providerSpend[provider].spend += metrics.metrics.spend;
providerSpend[provider].requests += metrics.metrics.api_requests;
providerSpend[provider].successful_requests += metrics.metrics.successful_requests;
providerSpend[provider].failed_requests += metrics.metrics.failed_requests;
providerSpend[provider].tokens += metrics.metrics.total_tokens;
} catch (e) {
console.error(`Error processing provider ${provider}: ${e}`);
}
});
});
return Object.values(providerSpend)
.filter((provider) => provider.spend > 0)
.sort((a, b) => b.spend - a.spend);
};

View file

@ -502,6 +502,8 @@ describe("UsagePage", () => {
userId: "user-123",
userEmail: "test@example.com",
userRole: "Internal User",
userRoleLabel: "Internal User",
isViewOnly: false,
premiumUser: true,
disabledPersonalKeyCreation: false,
showSSOBanner: false,
@ -861,6 +863,27 @@ describe("UsagePage", () => {
});
});
it.each(["organization", "agent"])("should not render the %s usage view for an internal user", async (usageView) => {
mockUseAuthorized.mockReturnValue(nonAdminSession);
renderWithProviders(<UsagePage {...defaultProps} organizations={mockOrganizations} />);
await waitFor(() => {
expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled();
});
const usageSelect = screen.getByTestId("usage-view-select");
act(() => {
fireEvent.change(usageSelect, { target: { value: "team" } });
});
expect(screen.getAllByText("Entity Usage").length).toBeGreaterThan(0);
act(() => {
fireEvent.change(usageSelect, { target: { value: usageView } });
});
expect(screen.queryByText("Entity Usage")).not.toBeInTheDocument();
});
describe("admin user selector", () => {
it("should render user selector for admin users in global view", async () => {
renderWithProviders(<UsagePage {...defaultProps} />);

View file

@ -33,6 +33,7 @@ import { useCustomers } from "@/app/(dashboard)/hooks/customers/useCustomers";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { useCurrentUser } from "@/app/(dashboard)/hooks/users/useCurrentUser";
import { useInfiniteUsers } from "@/app/(dashboard)/hooks/users/useUsers";
import { hasCapability } from "@/utils/capabilities";
import { formatNumberWithCommas } from "@/utils/dataUtils";
import { all_admin_roles, internalUserRoles } from "@/utils/roles";
import { ActivityMetrics, processActivityData } from "@/components/activity_metrics";
@ -109,6 +110,8 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
const { data: currentUser } = useCurrentUser();
const isAdmin = all_admin_roles.includes(userRole || "");
const canViewTagUsage = isAdmin || internalUserRoles.includes(userRole || "");
const canViewOrganizationUsage = hasCapability(userRole, "viewOrganizationUsage");
const canViewAgentUsage = hasCapability(userRole, "viewAgentUsage");
// Debounced search for user selector
const [userSearchInput, setUserSearchInput] = useState("");
@ -513,7 +516,7 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
<UsageViewSelect
value={usageView}
onChange={(value) => setUsageView(value)}
isAdmin={isAdmin}
userRole={userRole}
canViewTagUsage={canViewTagUsage}
/>
<AdvancedDatePicker value={dateValue} onValueChange={handleDateChange} />
@ -950,7 +953,7 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
)}
{/* Organization Usage Panel */}
{usageView === "organization" && (
{usageView === "organization" && canViewOrganizationUsage && (
<EntityUsage
accessToken={accessToken}
entityType="organization"
@ -1033,7 +1036,7 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
/>
</>
)}
{usageView === "agent" && (
{usageView === "agent" && canViewAgentUsage && (
<EntityUsage
accessToken={accessToken}
entityType="agent"

View file

@ -90,15 +90,16 @@ describe("UsageViewSelect", () => {
});
it("should render", () => {
render(<UsageViewSelect value="global" onChange={mockOnChange} isAdmin={false} />);
render(<UsageViewSelect value="global" onChange={mockOnChange} userRole="Internal User" />);
expect(screen.getByText("Usage View")).toBeInTheDocument();
expect(screen.getByText("Select the usage data you want to view")).toBeInTheDocument();
expect(screen.getByRole("combobox")).toBeInTheDocument();
expect(screen.getByRole("option", { name: "Your Usage" })).toBeInTheDocument();
});
it("should call onChange when value changes", () => {
render(<UsageViewSelect value="global" onChange={mockOnChange} isAdmin={true} />);
render(<UsageViewSelect value="global" onChange={mockOnChange} userRole="Admin" />);
const select = screen.getByRole("combobox");
act(() => {
@ -109,14 +110,32 @@ describe("UsageViewSelect", () => {
});
it("should show Tag Usage for non-admin users with tag usage permission", () => {
render(<UsageViewSelect value="global" onChange={mockOnChange} isAdmin={false} canViewTagUsage={true} />);
render(<UsageViewSelect value="global" onChange={mockOnChange} userRole="Internal User" canViewTagUsage={true} />);
expect(screen.getByRole("option", { name: "Tag Usage" })).toBeInTheDocument();
});
it("should hide Tag Usage for non-admin users without tag usage permission", () => {
render(<UsageViewSelect value="global" onChange={mockOnChange} isAdmin={false} />);
render(<UsageViewSelect value="global" onChange={mockOnChange} userRole="Internal User" />);
expect(screen.queryByRole("option", { name: "Tag Usage" })).not.toBeInTheDocument();
});
it.each(["Organization Usage", "Agent Usage (A2A)"])("should show %s to an admin", (optionName) => {
render(<UsageViewSelect value="global" onChange={mockOnChange} userRole="Admin" />);
expect(screen.getByRole("option", { name: optionName })).toBeInTheDocument();
});
it.each(["Organization Usage", "Agent Usage (A2A)"])("should hide %s from an internal user", (optionName) => {
render(<UsageViewSelect value="global" onChange={mockOnChange} userRole="Internal User" canViewTagUsage={true} />);
expect(screen.queryByRole("option", { name: optionName })).not.toBeInTheDocument();
});
it.each(["Team Usage", "Tag Usage"])("should keep %s available to an internal user", (optionName) => {
render(<UsageViewSelect value="global" onChange={mockOnChange} userRole="Internal User" canViewTagUsage={true} />);
expect(screen.getByRole("option", { name: optionName })).toBeInTheDocument();
});
});

View file

@ -11,6 +11,8 @@ import {
} from "@ant-design/icons";
import { Badge, Select } from "antd";
import React from "react";
import { hasCapability, type Capability } from "@/utils/capabilities";
import { all_admin_roles } from "@/utils/roles";
export type UsageOption =
| "global"
| "my-usage"
@ -24,7 +26,7 @@ export type UsageOption =
export interface UsageViewSelectProps {
value: UsageOption;
onChange: (value: UsageOption) => void;
isAdmin: boolean;
userRole: string | null;
canViewTagUsage?: boolean;
title?: string;
description?: string;
@ -35,6 +37,7 @@ interface OptionConfig {
label: string;
description: string;
icon: React.ReactNode;
capability?: Capability;
adminOnly?: boolean;
showForAdmin?: string;
showForNonAdmin?: string;
@ -63,12 +66,9 @@ const OPTIONS: OptionConfig[] = [
{
value: "organization",
label: "Organization Usage",
showForAdmin: "Organization Usage",
showForNonAdmin: "Your Organization Usage",
description: "View organization-level usage",
descriptionForAdmin: "View usage across all organizations",
descriptionForNonAdmin: "View your organization's usage",
description: "View usage across all organizations",
icon: <BankOutlined style={{ fontSize: "16px" }} />,
capability: "viewOrganizationUsage",
},
{
value: "team",
@ -95,7 +95,7 @@ const OPTIONS: OptionConfig[] = [
label: "Agent Usage (A2A)",
description: "View usage by AI agents",
icon: <RobotOutlined style={{ fontSize: "16px" }} />,
adminOnly: true,
capability: "viewAgentUsage",
},
{
value: "user",
@ -115,14 +115,18 @@ const OPTIONS: OptionConfig[] = [
export const UsageViewSelect: React.FC<UsageViewSelectProps> = ({
value,
onChange,
isAdmin,
userRole,
canViewTagUsage = false,
title = "Usage View",
description = "Select the usage data you want to view",
"data-id": dataId,
}) => {
const isAdmin = all_admin_roles.includes(userRole ?? "");
const getFilteredOptions = () => {
return OPTIONS.filter((option) => {
if (option.capability) {
return hasCapability(userRole, option.capability);
}
if (option.value === "tag" && canViewTagUsage) {
return true;
}

View file

@ -1,70 +1,40 @@
import { describe, expect, it } from "vitest";
import { hasCapability, rolesWithCapability } from "./capabilities";
import { hasCapability, rolesWithCapability, type Capability } from "./capabilities";
const ADMIN_ROLES = ["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"];
const NON_ADMIN_ROLES = [
"Internal User",
"Internal Viewer",
"internal_user",
"App User",
"Org Admin",
"Unknown Role",
"",
null,
undefined,
];
const ADMIN_ONLY_CAPABILITIES: Capability[] = [
"viewToolPolicies",
"viewAuditLogs",
"viewDeletedTeams",
"viewPolicies",
"viewPrompts",
"viewOrganizationUsage",
"viewAgentUsage",
];
describe("hasCapability", () => {
it.each(["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"])(
"should grant viewToolPolicies to %s",
(role) => {
expect(hasCapability(role, "viewToolPolicies")).toBe(true);
},
);
describe.each(ADMIN_ONLY_CAPABILITIES)("%s", (capability) => {
it.each(ADMIN_ROLES)("should grant it to %s", (role) => {
expect(hasCapability(role, capability)).toBe(true);
});
it.each(["Internal User", "Internal Viewer", "App User", "Org Admin", "Unknown Role", "", null, undefined])(
"should deny viewToolPolicies to %s",
(role) => {
expect(hasCapability(role, "viewToolPolicies")).toBe(false);
},
);
it.each(["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"])("should grant viewPolicies to %s", (role) => {
expect(hasCapability(role, "viewPolicies")).toBe(true);
});
it.each([
"Internal User",
"Internal Viewer",
"internal_user",
"App User",
"Org Admin",
"Unknown Role",
"",
null,
undefined,
])("should deny viewPolicies to %s", (role) => {
expect(hasCapability(role, "viewPolicies")).toBe(false);
});
it.each(["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"])("should grant viewPrompts to %s", (role) => {
expect(hasCapability(role, "viewPrompts")).toBe(true);
});
it.each([
"Internal User",
"Internal Viewer",
"internal_user",
"App User",
"Org Admin",
"Unknown Role",
"",
null,
undefined,
])("should deny viewPrompts to %s", (role) => {
expect(hasCapability(role, "viewPrompts")).toBe(false);
});
});
describe.each(["viewAuditLogs", "viewDeletedTeams"] as const)("hasCapability - %s", (capability) => {
it.each(["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"])("should grant it to %s", (role) => {
expect(hasCapability(role, capability)).toBe(true);
});
it.each(["Internal User", "Internal Viewer", "App User", "Org Admin", "Unknown Role", "", null, undefined])(
"should deny it to %s",
(role) => {
it.each(NON_ADMIN_ROLES)("should deny it to %s", (role) => {
expect(hasCapability(role, capability)).toBe(false);
},
);
});
});
});
describe("rolesWithCapability", () => {

View file

@ -6,6 +6,8 @@ const CAPABILITY_ROLES = {
viewDeletedTeams: all_admin_roles,
viewPolicies: all_admin_roles,
viewPrompts: all_admin_roles,
viewOrganizationUsage: all_admin_roles,
viewAgentUsage: all_admin_roles,
} as const satisfies Record<string, readonly string[]>;
export type Capability = keyof typeof CAPABILITY_ROLES;