mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/standard-lists-api-d1dc4a
This commit is contained in:
commit
a9857bb362
43 changed files with 2641 additions and 1613 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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: ...
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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"}]
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
};
|
||||
|
|
@ -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} />);
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue