mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #35719 from BerriAI/litellm_daily_any_cleanup_08_03_2026
chore(typing): clear basedpyright Any errors in budget reset, access groups, and cache settings
This commit is contained in:
commit
41722b1cbc
6 changed files with 240 additions and 86 deletions
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 29811
|
||||
"limit": 29682
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2645
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 42
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 9471
|
||||
"limit": 9440
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 11
|
||||
|
|
@ -54,10 +54,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportMissingParameterType": {
|
||||
"limit": 5855
|
||||
"limit": 5848
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15852
|
||||
"limit": 15850
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 41
|
||||
|
|
@ -99,19 +99,19 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 45324
|
||||
"limit": 45297
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 113
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 40452
|
||||
"limit": 40411
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 20309
|
||||
"limit": 20301
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 31978
|
||||
"limit": 31968
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 177
|
||||
|
|
|
|||
|
|
@ -1,12 +1,13 @@
|
|||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Literal
|
||||
from typing import Literal, Protocol, TypeVar
|
||||
|
||||
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.proxy._types import (
|
||||
LiteLLM_BudgetTableFull,
|
||||
|
|
@ -33,6 +34,98 @@ from litellm.repositories.verification_token_repository import (
|
|||
)
|
||||
from litellm.types.services import ServiceTypes
|
||||
|
||||
_RowT = TypeVar("_RowT")
|
||||
_RowT_co = TypeVar("_RowT_co", covariant=True)
|
||||
|
||||
|
||||
class _PrismaRecord(Protocol):
|
||||
def dict(self) -> Mapping[str, object]: ...
|
||||
|
||||
|
||||
class _BatchTable(Protocol):
|
||||
def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> None: ...
|
||||
|
||||
|
||||
class _ResetBatcher(Protocol):
|
||||
@property
|
||||
def litellm_verificationtoken(self) -> _BatchTable: ...
|
||||
|
||||
@property
|
||||
def litellm_usertable(self) -> _BatchTable: ...
|
||||
|
||||
@property
|
||||
def litellm_teamtable(self) -> _BatchTable: ...
|
||||
|
||||
async def commit(self) -> None: ...
|
||||
|
||||
|
||||
class _EndUserTable(Protocol):
|
||||
async def find_many(self, where: Mapping[str, object]) -> Sequence[_PrismaRecord]: ...
|
||||
|
||||
|
||||
class _SpendLinkedTable(Protocol[_RowT_co]):
|
||||
async def find_many(self, where: Mapping[str, object]) -> Sequence[_RowT_co]: ...
|
||||
|
||||
async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
|
||||
|
||||
|
||||
class _TeamMembershipRow(Protocol):
|
||||
@property
|
||||
def user_id(self) -> str: ...
|
||||
|
||||
@property
|
||||
def team_id(self) -> str: ...
|
||||
|
||||
|
||||
class _KeyRow(Protocol):
|
||||
@property
|
||||
def token(self) -> str: ...
|
||||
|
||||
|
||||
class _OrgRow(Protocol):
|
||||
@property
|
||||
def organization_id(self) -> str: ...
|
||||
|
||||
|
||||
class _TagRow(Protocol):
|
||||
@property
|
||||
def tag_name(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 _key_counter_key(row: _KeyRow) -> str:
|
||||
return f"spend:key:{row.token}"
|
||||
|
||||
|
||||
def _key_cache_key(row: _KeyRow) -> 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 [
|
||||
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}"
|
||||
|
||||
|
||||
class ResetBudgetJob:
|
||||
"""
|
||||
|
|
@ -134,11 +227,11 @@ class ResetBudgetJob:
|
|||
async def _cascade_reset_spend_for_budget_link(
|
||||
self,
|
||||
budgets_to_reset: list[LiteLLM_BudgetTableFull],
|
||||
table: Any,
|
||||
counter_key_fn: Callable[[Any], str],
|
||||
table: "_SpendLinkedTable[_RowT]",
|
||||
counter_key_fn: Callable[[_RowT], str],
|
||||
log_subject: str,
|
||||
extra_where: dict | None = None,
|
||||
cache_key_fn: Callable[[Any], str | list[str]] | None = None,
|
||||
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.
|
||||
|
|
@ -151,14 +244,14 @@ class ResetBudgetJob:
|
|||
if not budget_ids:
|
||||
return
|
||||
|
||||
where: dict = {"budget_id": {"in": budget_ids}}
|
||||
where: dict[str, object] = {"budget_id": {"in": budget_ids}}
|
||||
if extra_where:
|
||||
where.update(extra_where)
|
||||
|
||||
try:
|
||||
rows = await table.find_many(where=where)
|
||||
rows: Sequence[_RowT] = await table.find_many(where=where)
|
||||
except Exception as e:
|
||||
rows = []
|
||||
rows = ()
|
||||
verbose_proxy_logger.warning("Failed to fetch %s for counter invalidation: %s", log_subject, e)
|
||||
|
||||
update_result = await table.update_many(where=where, data={"spend": 0})
|
||||
|
|
@ -181,9 +274,9 @@ class ResetBudgetJob:
|
|||
return await self._cascade_reset_spend_for_budget_link(
|
||||
budgets_to_reset=budgets_to_reset,
|
||||
table=TeamMembershipRepository(self.prisma_client).table,
|
||||
counter_key_fn=lambda m: f"spend:team_member:{m.user_id}:{m.team_id}",
|
||||
counter_key_fn=_team_membership_counter_key,
|
||||
log_subject="team memberships",
|
||||
cache_key_fn=lambda m: f"{m.team_id}_{m.user_id}",
|
||||
cache_key_fn=_team_membership_cache_key,
|
||||
)
|
||||
|
||||
async def reset_budget_for_keys_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
|
||||
|
|
@ -196,10 +289,10 @@ class ResetBudgetJob:
|
|||
return await self._cascade_reset_spend_for_budget_link(
|
||||
budgets_to_reset=budgets_to_reset,
|
||||
table=VerificationTokenRepository(self.prisma_client).table,
|
||||
counter_key_fn=lambda k: f"spend:key:{k.token}",
|
||||
counter_key_fn=_key_counter_key,
|
||||
log_subject="keys",
|
||||
extra_where={"budget_duration": None, "spend": {"gt": 0}},
|
||||
cache_key_fn=lambda k: k.token,
|
||||
cache_key_fn=_key_cache_key,
|
||||
)
|
||||
|
||||
async def reset_budget_for_orgs_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
|
||||
|
|
@ -209,13 +302,10 @@ class ResetBudgetJob:
|
|||
return await self._cascade_reset_spend_for_budget_link(
|
||||
budgets_to_reset=budgets_to_reset,
|
||||
table=OrganizationRepository(self.prisma_client).table,
|
||||
counter_key_fn=lambda o: f"spend:org:{o.organization_id}",
|
||||
counter_key_fn=_org_counter_key,
|
||||
log_subject="orgs",
|
||||
extra_where={"spend": {"gt": 0}},
|
||||
cache_key_fn=lambda o: [
|
||||
f"org_id:{o.organization_id}",
|
||||
f"org_id:{o.organization_id}:with_budget",
|
||||
],
|
||||
cache_key_fn=_org_cache_keys,
|
||||
)
|
||||
|
||||
async def reset_budget_for_tags_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
|
||||
|
|
@ -233,10 +323,10 @@ class ResetBudgetJob:
|
|||
return await self._cascade_reset_spend_for_budget_link(
|
||||
budgets_to_reset=budgets_to_reset,
|
||||
table=TagRepository(self.prisma_client).table,
|
||||
counter_key_fn=lambda t: f"spend:tag:{t.tag_name}",
|
||||
counter_key_fn=_tag_counter_key,
|
||||
log_subject="tags",
|
||||
extra_where={"spend": {"gt": 0}},
|
||||
cache_key_fn=lambda t: f"tag:{t.tag_name}",
|
||||
cache_key_fn=_tag_cache_key,
|
||||
)
|
||||
|
||||
async def reset_budget_for_litellm_budget_table(self):
|
||||
|
|
@ -376,13 +466,14 @@ class ResetBudgetJob:
|
|||
rely on the default budget (litellm.max_end_user_budget_id) applied
|
||||
in-memory during auth checks.
|
||||
"""
|
||||
rows = await EndUserRepository(self.prisma_client).table.find_many(
|
||||
table: _EndUserTable = EndUserRepository(self.prisma_client).table
|
||||
rows = await table.find_many(
|
||||
where={
|
||||
"budget_id": None,
|
||||
"spend": {"gt": 0},
|
||||
},
|
||||
)
|
||||
return [LiteLLM_EndUserTable(**row.dict()) for row in rows]
|
||||
return [LiteLLM_EndUserTable.model_validate(row.dict()) for row in rows]
|
||||
|
||||
async def _write_key_reset_updates(self, updated_keys: list[LiteLLM_VerificationToken]) -> None:
|
||||
"""
|
||||
|
|
@ -395,7 +486,7 @@ class ResetBudgetJob:
|
|||
aborts the entire batch — silently leaving spend over the cap and
|
||||
budget_reset_at unchanged forever.
|
||||
"""
|
||||
batcher = self.prisma_client.db.batch_()
|
||||
batcher: _ResetBatcher = self.prisma_client.db.batch_()
|
||||
for k in updated_keys:
|
||||
token = getattr(k, "token", None)
|
||||
if token is None:
|
||||
|
|
@ -414,7 +505,7 @@ class ResetBudgetJob:
|
|||
that trips Prisma's DataError on rows carrying unrecognised fields
|
||||
(see #27730).
|
||||
"""
|
||||
batcher = self.prisma_client.db.batch_()
|
||||
batcher: _ResetBatcher = self.prisma_client.db.batch_()
|
||||
for u in updated_users:
|
||||
user_id = getattr(u, "user_id", None)
|
||||
if user_id is None:
|
||||
|
|
@ -433,7 +524,7 @@ class ResetBudgetJob:
|
|||
that trips Prisma's DataError on rows carrying unrecognised fields
|
||||
(see #27730).
|
||||
"""
|
||||
batcher = self.prisma_client.db.batch_()
|
||||
batcher: _ResetBatcher = self.prisma_client.db.batch_()
|
||||
for t in updated_teams:
|
||||
team_id = getattr(t, "team_id", None)
|
||||
if team_id is None:
|
||||
|
|
@ -688,7 +779,7 @@ class ResetBudgetJob:
|
|||
async def _reset_expired_window(
|
||||
window: dict,
|
||||
counter_key: str,
|
||||
spend_counter_cache: Any,
|
||||
spend_counter_cache: DualCache,
|
||||
now: datetime,
|
||||
reset_settings: BudgetResetSettings,
|
||||
) -> bool:
|
||||
|
|
|
|||
|
|
@ -1,3 +1,6 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from typing import Protocol
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -29,6 +32,74 @@ router = APIRouter(
|
|||
)
|
||||
|
||||
|
||||
class _AccessGroupRecord(Protocol):
|
||||
@property
|
||||
def access_group_id(self) -> str: ...
|
||||
|
||||
@property
|
||||
def assigned_team_ids(self) -> Sequence[str] | None: ...
|
||||
|
||||
@property
|
||||
def assigned_key_ids(self) -> Sequence[str] | None: ...
|
||||
|
||||
def dict(self) -> Mapping[str, object]: ...
|
||||
|
||||
|
||||
class _TeamRecord(Protocol):
|
||||
@property
|
||||
def team_id(self) -> str: ...
|
||||
|
||||
@property
|
||||
def access_group_ids(self) -> Sequence[str] | None: ...
|
||||
|
||||
|
||||
class _KeyRecord(Protocol):
|
||||
@property
|
||||
def token(self) -> str: ...
|
||||
|
||||
@property
|
||||
def access_group_ids(self) -> Sequence[str] | None: ...
|
||||
|
||||
|
||||
class _AccessGroupTable(Protocol):
|
||||
async def find_unique(self, where: Mapping[str, object]) -> _AccessGroupRecord | None: ...
|
||||
|
||||
async def find_many(self, order: Mapping[str, object]) -> Sequence[_AccessGroupRecord]: ...
|
||||
|
||||
async def create(self, data: Mapping[str, object]) -> _AccessGroupRecord: ...
|
||||
|
||||
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _AccessGroupRecord: ...
|
||||
|
||||
async def delete(self, where: Mapping[str, object]) -> object: ...
|
||||
|
||||
|
||||
class _TeamTable(Protocol):
|
||||
async def find_unique(self, where: Mapping[str, object]) -> _TeamRecord | None: ...
|
||||
|
||||
async def find_many(self, where: Mapping[str, object]) -> Sequence[_TeamRecord]: ...
|
||||
|
||||
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> object: ...
|
||||
|
||||
|
||||
class _KeyTable(Protocol):
|
||||
async def find_unique(self, where: Mapping[str, object]) -> _KeyRecord | None: ...
|
||||
|
||||
async def find_many(self, where: Mapping[str, object]) -> Sequence[_KeyRecord]: ...
|
||||
|
||||
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> object: ...
|
||||
|
||||
|
||||
class _AccessGroupTx(Protocol):
|
||||
@property
|
||||
def litellm_accessgrouptable(self) -> _AccessGroupTable: ...
|
||||
|
||||
@property
|
||||
def litellm_teamtable(self) -> _TeamTable: ...
|
||||
|
||||
@property
|
||||
def litellm_verificationtoken(self) -> _KeyTable: ...
|
||||
|
||||
|
||||
def _require_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
raise HTTPException(
|
||||
|
|
@ -48,29 +119,16 @@ def _require_admin_view(user_api_key_dict: UserAPIKeyAuth) -> None:
|
|||
)
|
||||
|
||||
|
||||
def _record_to_response(record) -> AccessGroupResponse:
|
||||
return AccessGroupResponse(
|
||||
access_group_id=record.access_group_id,
|
||||
access_group_name=record.access_group_name,
|
||||
description=record.description,
|
||||
access_model_names=record.access_model_names,
|
||||
access_mcp_server_ids=record.access_mcp_server_ids,
|
||||
access_agent_ids=record.access_agent_ids,
|
||||
assigned_team_ids=record.assigned_team_ids,
|
||||
assigned_key_ids=record.assigned_key_ids,
|
||||
created_at=record.created_at,
|
||||
created_by=record.created_by,
|
||||
updated_at=record.updated_at,
|
||||
updated_by=record.updated_by,
|
||||
)
|
||||
def _record_to_response(record: _AccessGroupRecord) -> AccessGroupResponse:
|
||||
return AccessGroupResponse.model_validate(record.dict())
|
||||
|
||||
|
||||
def _record_to_access_group_table(record) -> LiteLLM_AccessGroupTable:
|
||||
def _record_to_access_group_table(record: _AccessGroupRecord) -> LiteLLM_AccessGroupTable:
|
||||
"""Convert a Prisma record to a LiteLLM_AccessGroupTable pydantic object for caching."""
|
||||
return LiteLLM_AccessGroupTable(**record.dict())
|
||||
return LiteLLM_AccessGroupTable.model_validate(record.dict())
|
||||
|
||||
|
||||
async def _cache_access_group_record(record) -> None:
|
||||
async def _cache_access_group_record(record: _AccessGroupRecord) -> None:
|
||||
"""
|
||||
Cache an access group Prisma record in the user_api_key_cache.
|
||||
|
||||
|
|
@ -109,7 +167,7 @@ async def _invalidate_cache_access_group(access_group_id: str) -> None:
|
|||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def _sync_add_access_group_to_teams(tx, team_ids: list[str], access_group_id: str) -> None:
|
||||
async def _sync_add_access_group_to_teams(tx: _AccessGroupTx, team_ids: list[str], access_group_id: str) -> None:
|
||||
"""Add access_group_id to each team's access_group_ids (idempotent)."""
|
||||
for team_id in team_ids:
|
||||
team = await tx.litellm_teamtable.find_unique(where={"team_id": team_id})
|
||||
|
|
@ -120,18 +178,18 @@ async def _sync_add_access_group_to_teams(tx, team_ids: list[str], access_group_
|
|||
)
|
||||
|
||||
|
||||
async def _sync_remove_access_group_from_teams(tx, team_ids: list[str], access_group_id: str) -> None:
|
||||
async def _sync_remove_access_group_from_teams(tx: _AccessGroupTx, team_ids: list[str], access_group_id: str) -> None:
|
||||
"""Remove access_group_id from each team's access_group_ids (idempotent)."""
|
||||
for team_id in team_ids:
|
||||
team = await tx.litellm_teamtable.find_unique(where={"team_id": team_id})
|
||||
if team is not None and access_group_id in (team.access_group_ids or []):
|
||||
await tx.litellm_teamtable.update(
|
||||
where={"team_id": team_id},
|
||||
data={"access_group_ids": [ag for ag in team.access_group_ids if ag != access_group_id]},
|
||||
data={"access_group_ids": [ag for ag in (team.access_group_ids or ()) if ag != access_group_id]},
|
||||
)
|
||||
|
||||
|
||||
async def _sync_add_access_group_to_keys(tx, key_tokens: list[str], access_group_id: str) -> None:
|
||||
async def _sync_add_access_group_to_keys(tx: _AccessGroupTx, key_tokens: list[str], access_group_id: str) -> None:
|
||||
"""Add access_group_id to each key's access_group_ids (idempotent)."""
|
||||
for token in key_tokens:
|
||||
key = await tx.litellm_verificationtoken.find_unique(where={"token": token})
|
||||
|
|
@ -142,14 +200,14 @@ async def _sync_add_access_group_to_keys(tx, key_tokens: list[str], access_group
|
|||
)
|
||||
|
||||
|
||||
async def _sync_remove_access_group_from_keys(tx, key_tokens: list[str], access_group_id: str) -> None:
|
||||
async def _sync_remove_access_group_from_keys(tx: _AccessGroupTx, key_tokens: list[str], access_group_id: str) -> None:
|
||||
"""Remove access_group_id from each key's access_group_ids (idempotent)."""
|
||||
for token in key_tokens:
|
||||
key = await tx.litellm_verificationtoken.find_unique(where={"token": token})
|
||||
if key is not None and access_group_id in (key.access_group_ids or []):
|
||||
await tx.litellm_verificationtoken.update(
|
||||
where={"token": token},
|
||||
data={"access_group_ids": [ag for ag in key.access_group_ids if ag != access_group_id]},
|
||||
data={"access_group_ids": [ag for ag in (key.access_group_ids or ()) if ag != access_group_id]},
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -280,6 +338,7 @@ async def create_access_group(
|
|||
prisma_client = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
try:
|
||||
tx: _AccessGroupTx
|
||||
async with prisma_client.db.tx() as tx:
|
||||
existing = await tx.litellm_accessgrouptable.find_unique(
|
||||
where={"access_group_name": data.access_group_name}
|
||||
|
|
@ -347,7 +406,8 @@ async def list_access_groups(
|
|||
_require_admin_view(user_api_key_dict)
|
||||
prisma_client = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
records = await AccessGroupRepository(prisma_client).table.find_many(order={"created_at": "desc"})
|
||||
table: _AccessGroupTable = AccessGroupRepository(prisma_client).table
|
||||
records = await table.find_many(order={"created_at": "desc"})
|
||||
return [_record_to_response(r) for r in records]
|
||||
|
||||
|
||||
|
|
@ -362,7 +422,8 @@ async def get_access_group(
|
|||
_require_admin_view(user_api_key_dict)
|
||||
prisma_client = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
record = await AccessGroupRepository(prisma_client).table.find_unique(where={"access_group_id": access_group_id})
|
||||
table: _AccessGroupTable = AccessGroupRepository(prisma_client).table
|
||||
record = await table.find_unique(where={"access_group_id": access_group_id})
|
||||
if record is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
|
|
@ -408,6 +469,7 @@ async def update_access_group(
|
|||
keys_to_remove: list[str] = []
|
||||
|
||||
try:
|
||||
tx: _AccessGroupTx
|
||||
async with prisma_client.db.tx() as tx:
|
||||
# Read inside the transaction so delta computation is consistent with the write,
|
||||
# avoiding a TOCTOU race where a concurrent update could make deltas stale.
|
||||
|
|
@ -480,6 +542,7 @@ async def delete_access_group(
|
|||
affected_team_ids: list[str] = []
|
||||
affected_key_tokens: list[str] = []
|
||||
|
||||
tx: _AccessGroupTx
|
||||
async with prisma_client.db.tx() as tx:
|
||||
existing = await tx.litellm_accessgrouptable.find_unique(where={"access_group_id": access_group_id})
|
||||
if existing is None:
|
||||
|
|
@ -512,7 +575,7 @@ async def delete_access_group(
|
|||
for team in teams_with_group:
|
||||
await tx.litellm_teamtable.update(
|
||||
where={"team_id": team.team_id},
|
||||
data={"access_group_ids": [ag for ag in (team.access_group_ids or []) if ag != access_group_id]},
|
||||
data={"access_group_ids": [ag for ag in (team.access_group_ids or ()) if ag != access_group_id]},
|
||||
)
|
||||
# Use _sync_remove only for out-of-sync teams not found by the hasSome query.
|
||||
out_of_sync_team_ids = set(existing.assigned_team_ids or []) - {t.team_id for t in teams_with_group}
|
||||
|
|
@ -522,7 +585,7 @@ async def delete_access_group(
|
|||
for key in keys_with_group:
|
||||
await tx.litellm_verificationtoken.update(
|
||||
where={"token": key.token},
|
||||
data={"access_group_ids": [ag for ag in (key.access_group_ids or []) if ag != access_group_id]},
|
||||
data={"access_group_ids": [ag for ag in (key.access_group_ids or ()) if ag != access_group_id]},
|
||||
)
|
||||
# Use _sync_remove only for out-of-sync keys not found by the hasSome query.
|
||||
out_of_sync_key_tokens = set(existing.assigned_key_ids or []) - {k.token for k in keys_with_group}
|
||||
|
|
|
|||
|
|
@ -63,7 +63,7 @@ _REDACTED_VALUE = "***REDACTED***"
|
|||
_URL_OVERRIDDEN_CONNECTION_FIELDS: frozenset = frozenset({"host", "port", "db", "password", "username"})
|
||||
|
||||
|
||||
def _resolve_cache_url_precedence(settings: Mapping[str, Any]) -> dict[str, Any]:
|
||||
def _resolve_cache_url_precedence(settings: Mapping[str, object]) -> dict[str, Any]:
|
||||
"""Return cache settings with the url-vs-discrete-fields ambiguity resolved.
|
||||
|
||||
When a full ``url`` is supplied it wins: the discrete
|
||||
|
|
@ -80,7 +80,7 @@ def _resolve_cache_url_precedence(settings: Mapping[str, Any]) -> dict[str, Any]
|
|||
return {k: v for k, v in settings.items() if k not in _URL_OVERRIDDEN_CONNECTION_FIELDS}
|
||||
|
||||
|
||||
def _parse_stored_settings(cache_settings_value: object) -> dict[str, Any]:
|
||||
def _parse_stored_settings(cache_settings_value: object) -> dict[str, object]:
|
||||
"""Normalize a stored cache_settings blob to a dict.
|
||||
|
||||
The prisma column comes back as either a JSON string or an already-parsed
|
||||
|
|
@ -91,7 +91,7 @@ def _parse_stored_settings(cache_settings_value: object) -> dict[str, Any]:
|
|||
return parsed if isinstance(parsed, dict) else {}
|
||||
|
||||
|
||||
def _overlay_environment(stored: Mapping[str, Any]) -> dict[str, Any]:
|
||||
def _overlay_environment(stored: Mapping[str, object]) -> dict[str, object]:
|
||||
"""Fill connection fields from the REDIS_* environment the cache actually reads.
|
||||
|
||||
A response cache pointed at Redis resolves host/port/password/etc. from the
|
||||
|
|
@ -113,7 +113,7 @@ def _overlay_environment(stored: Mapping[str, Any]) -> dict[str, Any]:
|
|||
return effective
|
||||
|
||||
|
||||
def _redact_credentials(settings: Mapping[str, Any]) -> dict[str, Any]:
|
||||
def _redact_credentials(settings: Mapping[str, object]) -> dict[str, object]:
|
||||
"""Replace credential-bearing values with a fixed marker, keeping the rest.
|
||||
|
||||
The marker is unambiguous on the way back in: an admin who edits an
|
||||
|
|
@ -164,7 +164,7 @@ def _target_repr(value: object) -> str:
|
|||
return str(value)
|
||||
|
||||
|
||||
def _saved_secret_is_reusable(incoming: Mapping[str, Any], saved: Mapping[str, Any]) -> bool:
|
||||
def _saved_secret_is_reusable(incoming: Mapping[str, object], saved: Mapping[str, object]) -> bool:
|
||||
"""Whether a stored credential may be restored for this request.
|
||||
|
||||
A stored secret belongs to the stored connection target, so it is reused only
|
||||
|
|
@ -197,7 +197,7 @@ def _saved_secret_is_reusable(incoming: Mapping[str, Any], saved: Mapping[str, A
|
|||
return True
|
||||
|
||||
|
||||
def _merge_over_saved(incoming: Mapping[str, Any], saved: Mapping[str, Any]) -> dict[str, Any]:
|
||||
def _merge_over_saved(incoming: Mapping[str, object], saved: Mapping[str, object]) -> dict[str, Any]:
|
||||
"""Keep the stored secret behind any credential the caller echoed back redacted or omitted.
|
||||
|
||||
GET returns credentials as the marker and the form never re-prefills a
|
||||
|
|
@ -239,7 +239,7 @@ def _merge_over_saved(incoming: Mapping[str, Any], saved: Mapping[str, Any]) ->
|
|||
return merged
|
||||
|
||||
|
||||
def _redact_settings(settings: Mapping[str, Any] | None) -> dict[str, Any]:
|
||||
def _redact_settings(settings: Mapping[str, object] | None) -> dict[str, object]:
|
||||
"""Replace every value in a settings map with a fixed marker.
|
||||
|
||||
Cache config carries Redis credentials (passwords, connection strings).
|
||||
|
|
@ -268,8 +268,8 @@ def _log_audit_task_exception(task: "asyncio.Task[None]") -> None:
|
|||
async def _emit_cache_settings_audit_log(
|
||||
*,
|
||||
action: AUDIT_ACTIONS,
|
||||
before_settings: Mapping[str, Any] | None,
|
||||
after_settings: Mapping[str, Any] | None,
|
||||
before_settings: Mapping[str, object] | None,
|
||||
after_settings: Mapping[str, object] | None,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
litellm_changed_by: str | None,
|
||||
) -> None:
|
||||
|
|
@ -313,17 +313,17 @@ class CacheSettingsManager:
|
|||
Tracks last cache params to avoid unnecessary reinitialization.
|
||||
"""
|
||||
|
||||
_last_cache_params: dict[str, Any] | None = None
|
||||
_last_cache_params: dict[str, object] | None = None
|
||||
|
||||
@staticmethod
|
||||
def _cache_params_equal(params1: dict[str, Any], params2: dict[str, Any]) -> bool:
|
||||
def _cache_params_equal(params1: dict[str, object], params2: dict[str, object]) -> bool:
|
||||
"""
|
||||
Compare two cache parameter dictionaries for equality.
|
||||
Normalizes values and filters out UI-only fields.
|
||||
"""
|
||||
|
||||
# Normalize by removing None values and UI-only fields
|
||||
def normalize(params: dict[str, Any]) -> dict[str, Any]:
|
||||
def normalize(params: dict[str, object]) -> dict[str, object]:
|
||||
normalized = {}
|
||||
for k, v in params.items():
|
||||
if k == "redis_type": # Skip UI-only field
|
||||
|
|
@ -390,7 +390,7 @@ class CacheSettingsManager:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def update_cache_params(cache_params: dict[str, Any]):
|
||||
def update_cache_params(cache_params: dict[str, object]):
|
||||
"""
|
||||
Update the last cache params after initialization.
|
||||
Called after cache settings are updated via the API.
|
||||
|
|
@ -400,12 +400,12 @@ class CacheSettingsManager:
|
|||
|
||||
class CacheSettingsResponse(BaseModel):
|
||||
fields: list[CacheSettingsField] = Field(description="List of all configurable cache settings with metadata")
|
||||
current_values: dict[str, Any] = Field(description="Current values of cache settings")
|
||||
current_values: dict[str, object] = Field(description="Current values of cache settings")
|
||||
redis_type_descriptions: dict[str, str] = Field(description="Descriptions for each Redis type option")
|
||||
|
||||
|
||||
class CacheTestRequest(BaseModel):
|
||||
cache_settings: dict[str, Any] = Field(description="Cache settings to test connection with")
|
||||
cache_settings: dict[str, object] = Field(description="Cache settings to test connection with")
|
||||
|
||||
|
||||
class CacheTestResponse(BaseModel):
|
||||
|
|
@ -415,7 +415,7 @@ class CacheTestResponse(BaseModel):
|
|||
|
||||
|
||||
class CacheSettingsUpdateRequest(BaseModel):
|
||||
cache_settings: dict[str, Any] = Field(description="Cache settings to save")
|
||||
cache_settings: dict[str, object] = Field(description="Cache settings to save")
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
@ -441,7 +441,7 @@ async def get_cache_settings(
|
|||
cache_fields = [field.model_copy(deep=True) for field in CACHE_SETTINGS_FIELDS]
|
||||
|
||||
# Read the stored settings (decrypted); an env-only cache has none.
|
||||
stored: dict[str, Any] = {}
|
||||
stored: dict[str, object] = {}
|
||||
if prisma_client is not None:
|
||||
cache_config = await CacheConfigRepository(prisma_client).table.find_unique(where={"id": "cache_config"})
|
||||
if cache_config is not None and cache_config.cache_settings:
|
||||
|
|
@ -507,7 +507,7 @@ async def test_cache_connection(
|
|||
# A credential the form left untouched arrives redacted; resolve it back
|
||||
# to the stored secret so the test connects with the real password. A
|
||||
# lookup failure must not block the test, so fall back to no stored row.
|
||||
saved_settings: dict[str, Any] = {}
|
||||
saved_settings: dict[str, object] = {}
|
||||
if prisma_client is not None:
|
||||
try:
|
||||
existing_row = await CacheConfigRepository(prisma_client).table.find_unique(
|
||||
|
|
@ -590,8 +590,8 @@ async def update_cache_settings(
|
|||
# Read the stored row first: its decrypted values back any credential the
|
||||
# caller echoed back redacted, and its key set drives the audit diff.
|
||||
existing_row = await CacheConfigRepository(prisma_client).table.find_unique(where={"id": "cache_config"})
|
||||
before_settings: dict[str, Any] | None = None
|
||||
saved_settings: dict[str, Any] = {}
|
||||
before_settings: dict[str, object] | None = None
|
||||
saved_settings: dict[str, object] = {}
|
||||
if existing_row is not None and existing_row.cache_settings:
|
||||
before_settings = _parse_stored_settings(existing_row.cache_settings)
|
||||
saved_settings = proxy_config._decrypt_db_variables(variables_dict=before_settings)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"ANN001": {
|
||||
"limit": 3104
|
||||
"limit": 3097
|
||||
},
|
||||
"ANN002": {
|
||||
"limit": 69
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 130
|
||||
},
|
||||
"ANN401": {
|
||||
"limit": 1850
|
||||
"limit": 1848
|
||||
},
|
||||
"ASYNC230": {
|
||||
"limit": 14
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 23350
|
||||
"limit": 23349
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 27255
|
||||
"limit": 27252
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 292
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue