mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(rbac): pass plain dicts to prisma order args so DB roles load
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
ccb4b06cd6
commit
4589729c38
4 changed files with 30 additions and 6 deletions
|
|
@ -18,6 +18,7 @@ from pydantic import TypeAdapter, ValidationError
|
|||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import LiteLLMRoutes, LitellmUserRoles
|
||||
from litellm.repositories.prisma_args import prisma_args
|
||||
from litellm.types.custom_rbac import CustomRBACRole, CustomRBACRoleResponse
|
||||
|
||||
CUSTOM_RBAC_ROLES_CONFIG_KEY: Final = "custom_rbac_roles"
|
||||
|
|
@ -119,7 +120,7 @@ def get_config_custom_rbac_roles() -> tuple[CustomRBACRole, ...]:
|
|||
|
||||
|
||||
async def get_db_custom_rbac_roles(table: _CustomRoleTable) -> tuple[CustomRBACRoleResponse, ...]:
|
||||
records: Final = await table.find_many(order=_ORDER_BY_ROLE_NAME)
|
||||
records: Final = await table.find_many(order=prisma_args(_ORDER_BY_ROLE_NAME))
|
||||
return tuple(CustomRBACRoleResponse.model_validate(record.dict()) for record in records)
|
||||
|
||||
|
||||
|
|
@ -180,8 +181,8 @@ async def validate_assigned_user_role(user_role: LitellmUserRoles | str | None)
|
|||
async def get_active_custom_rbac_engine() -> CustomRBACEngine | None:
|
||||
"""The engine for the currently configured roles, or None when no custom role exists.
|
||||
|
||||
A DB read failure reuses the last known policy so a transient outage cannot silently
|
||||
downgrade a governed role to the built-in role permissions.
|
||||
A DB read failure reuses the last known policy, or the config defined roles alone, so a
|
||||
transient outage cannot silently downgrade a governed role to the built-in permissions.
|
||||
"""
|
||||
cached: Final = _ENGINE_CACHE.get_fresh()
|
||||
if cached is not None:
|
||||
|
|
@ -192,7 +193,11 @@ async def get_active_custom_rbac_engine() -> CustomRBACEngine | None:
|
|||
db_roles: Final = () if table is None else await get_db_custom_rbac_roles(table=table)
|
||||
except Exception as exc: # noqa: BLE001 # any DB failure must keep the last known policy, not drop it
|
||||
verbose_proxy_logger.exception("Failed to load custom RBAC roles from the DB: %s", exc)
|
||||
return _ENGINE_CACHE.get_stale()
|
||||
stale: Final = _ENGINE_CACHE.get_stale()
|
||||
config_only: Final = get_config_custom_rbac_roles()
|
||||
if stale is not None or not config_only:
|
||||
return stale
|
||||
return build_custom_rbac_engine(roles=config_only)
|
||||
|
||||
roles: Final = get_config_custom_rbac_roles() + tuple(
|
||||
CustomRBACRole(
|
||||
|
|
|
|||
|
|
@ -96,7 +96,10 @@ async def _reject_unknown_inherits(
|
|||
return
|
||||
known: Final = (
|
||||
frozenset(role.role_name for role in get_config_custom_rbac_roles())
|
||||
| frozenset(str(record.dict()["role_name"]) for record in await table.find_many(order=_ORDER_BY_ROLE_NAME))
|
||||
| frozenset(
|
||||
str(record.dict()["role_name"])
|
||||
for record in await table.find_many(order=prisma_args(_ORDER_BY_ROLE_NAME))
|
||||
)
|
||||
| frozenset((role_name,))
|
||||
)
|
||||
unknown: Final = tuple(parent for parent in inherits if parent not in known)
|
||||
|
|
@ -190,7 +193,7 @@ async def list_custom_roles(
|
|||
)
|
||||
for role in get_config_custom_rbac_roles()
|
||||
)
|
||||
records: Final = await _role_table().find_many(order=_ORDER_BY_ROLE_NAME)
|
||||
records: Final = await _role_table().find_many(order=prisma_args(_ORDER_BY_ROLE_NAME))
|
||||
return CustomRBACRoleListResponse(roles=config_roles + tuple(_to_response(record) for record in records))
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -166,6 +166,7 @@ class _FakeTable:
|
|||
self.reads = 0
|
||||
|
||||
async def find_many(self, order):
|
||||
assert type(order) is dict, "prisma only serializes plain dicts"
|
||||
self.reads += 1
|
||||
if self._fail:
|
||||
raise RuntimeError("db down")
|
||||
|
|
@ -228,6 +229,19 @@ class TestEngineLoading:
|
|||
|
||||
assert after_failure is engine
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_db_failure_without_cache_still_serves_config_roles(self):
|
||||
settings = {_CONFIG_KEY: [{"role_name": "cfg-role", "allowed_routes": ["/key/info"]}]}
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.general_settings", settings),
|
||||
patch(_TABLE_PATH, return_value=_FakeTable(records=(), fail=True)),
|
||||
):
|
||||
engine = await get_active_custom_rbac_engine()
|
||||
|
||||
assert engine is not None
|
||||
assert engine.is_route_allowed(role_name="cfg-role", route="/key/info") is True
|
||||
assert engine.is_route_allowed(role_name="cfg-role", route="/key/generate") is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_roles_means_no_engine(self):
|
||||
with (
|
||||
|
|
|
|||
|
|
@ -32,10 +32,12 @@ class _FakeTable:
|
|||
self.deleted: list[str] = []
|
||||
|
||||
async def find_unique(self, where):
|
||||
assert type(where) is dict, "prisma only serializes plain dicts"
|
||||
row = self.rows.get(where["role_name"])
|
||||
return None if row is None else _FakeRecord(row)
|
||||
|
||||
async def find_many(self, order):
|
||||
assert type(order) is dict, "prisma only serializes plain dicts"
|
||||
return [_FakeRecord(row) for row in self.rows.values()]
|
||||
|
||||
async def create(self, data):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue