diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index d07ac0c5586..6c9c51770bb 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -955,17 +955,14 @@ async def get_default_end_user_budget( # Fetch from database try: - budget_record: Final = await BudgetRepository(prisma_client).table.find_unique( - where={"budget_id": litellm.max_end_user_budget_id} - ) + _budget_obj: Final = await BudgetRepository(prisma_client).find_by_id(litellm.max_end_user_budget_id) - if budget_record is None: + if _budget_obj is None: verbose_proxy_logger.warning( "Default end user budget not found in database: %s", litellm.max_end_user_budget_id ) return None - _budget_obj: Final = LiteLLM_BudgetTable.model_validate(budget_record.dict()) # Cache the budget for 60 seconds await user_api_key_cache.async_set_cache( key=cache_key, @@ -1014,7 +1011,7 @@ async def get_team_member_default_budget( return LiteLLM_BudgetTable.model_validate(cached_budget) try: - budget_record: Final = await BudgetRepository(prisma_client).table.find_unique(where={"budget_id": budget_id}) + budget_record: Final = await BudgetRepository(prisma_client).find_by_id(budget_id) if budget_record is None: verbose_proxy_logger.warning("Team-default member budget not found in database: %s", budget_id) @@ -1022,11 +1019,11 @@ async def get_team_member_default_budget( await user_api_key_cache.async_set_cache( key=cache_key, - value=budget_record.dict(), + value=budget_record.model_dump(), ttl=get_management_object_ttl(user_api_key_cache), ) - return LiteLLM_BudgetTable.model_validate(budget_record.dict()) + return budget_record except Exception: verbose_proxy_logger.exception("Error fetching team-default member budget %s", budget_id) @@ -1171,16 +1168,14 @@ async def get_end_user_object( # Fetch from database try: - response: Final = await EndUserRepository(prisma_client).table.find_unique( + _response = await EndUserRepository(prisma_client).find_unique_model( where={"user_id": end_user_id}, include={"litellm_budget_table": True, "object_permission": True}, ) - if response is None: + if _response is None: raise Exception - # Convert to LiteLLM_EndUserTable object - _response = LiteLLM_EndUserTable.model_validate(response.dict()) # Apply default budget if needed _response = await _apply_default_budget_to_end_user( @@ -1345,8 +1340,8 @@ async def get_tag_objects_batch( if not tag_names: return {} - tag_objects: Final = {} - uncached_tags: Final = [] + tag_objects: Final[dict[str, LiteLLM_TagTable]] = {} + uncached_tags: Final[list[str]] = [] # Try to get all tags from cache first for tag_name in tag_names: @@ -1363,7 +1358,7 @@ async def get_tag_objects_batch( # Batch fetch uncached tags from DB in one query if uncached_tags: try: - db_tags: Final = await TagRepository(prisma_client).table.find_many( + db_tags: Final = await TagRepository(prisma_client).find_many_models( where={"tag_name": {"in": uncached_tags}}, include={"litellm_budget_table": True}, ) @@ -1372,13 +1367,12 @@ async def get_tag_objects_batch( for db_tag in db_tags: tag_name = db_tag.tag_name cache_key = f"tag:{tag_name}" - _tag_obj = LiteLLM_TagTable.model_validate(db_tag.dict()) await user_api_key_cache.async_set_cache( key=cache_key, - value=_tag_obj, + value=db_tag, model_type=LiteLLM_TagTable, ) - tag_objects[tag_name] = _tag_obj + tag_objects[tag_name] = db_tag except Exception as e: verbose_proxy_logger.debug("Error batch fetching tags from database: %s", e) @@ -1457,15 +1451,14 @@ async def get_team_membership( # else, check db try: - response: Final = await TeamMembershipRepository(prisma_client).table.find_unique( + _response: Final = await TeamMembershipRepository(prisma_client).find_unique_model( where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, include={"litellm_budget_table": True}, ) - if response is None: + if _response is None: return None - _response: Final = LiteLLM_TeamMembership.model_validate(response.dict()) await user_api_key_cache.async_set_cache( key=_key, value=_response, @@ -1768,11 +1761,10 @@ async def get_user_object( ] response.organization_memberships = _dumped_memberships - _response = LiteLLM_UserTable.model_validate(dict(response)) - _response = await _backfill_null_user_email( + _response: Final = await _backfill_null_user_email( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, - user_row=_response, + user_row=LiteLLM_UserTable.model_validate(dict(response)), user_email=user_email, ) response_dict: Final = _response.model_dump() @@ -2148,17 +2140,16 @@ async def get_access_object( # Not in cache - fetch from DB try: - response: Final = await AccessGroupRepository(prisma_client).table.find_unique( + _response: Final = await AccessGroupRepository(prisma_client).find_unique_model( where={"access_group_id": access_group_id} ) - if response is None: + if _response is None: raise HTTPException( status_code=404, detail={"error": f"Access group doesn't exist in db. Access group={access_group_id}."}, ) - _response: Final = LiteLLM_AccessGroupTable.model_validate(response.dict()) # Save to cache await _cache_access_object( @@ -2224,7 +2215,7 @@ async def get_team_object_by_alias( # Query database by team_alias try: - teams: Final = await TeamRepository(prisma_client).table.find_many(where={"team_alias": team_alias}) + teams: Final = await TeamRepository(prisma_client).find_many(where={"team_alias": team_alias}) if not teams: raise HTTPException( @@ -2329,7 +2320,7 @@ async def get_org_object_by_alias( # Query database by organization_alias try: - orgs = await OrganizationRepository(prisma_client).table.find_many(where={"organization_alias": org_alias}) + orgs = await OrganizationRepository(prisma_client).find_many(where={"organization_alias": org_alias}) if not orgs: raise HTTPException( @@ -2546,16 +2537,10 @@ async def get_jwt_key_mapping_object( Returns the hashed token (str) if a matching active mapping is found, else None. """ - mapping: Final = await JWTKeyMappingRepository(prisma_client).table.find_first( - where={ - "jwt_claim_name": jwt_claim_name, - "jwt_claim_value": jwt_claim_value, - "is_active": True, - } + return await JWTKeyMappingRepository(prisma_client).find_active_token( + jwt_claim_name=jwt_claim_name, + jwt_claim_value=jwt_claim_value, ) - if mapping is not None: - return mapping.token - return None @log_db_metrics @@ -2674,14 +2659,11 @@ async def get_object_permission( # else, check db try: - response: Final = await ObjectPermissionRepository(prisma_client).table.find_unique( - where={"object_permission_id": object_permission_id} - ) + _perm_obj: Final = await ObjectPermissionRepository(prisma_client).find_by_id(object_permission_id) - if response is None: + if _perm_obj is None: return None - _perm_obj: Final = LiteLLM_ObjectPermissionTable.model_validate(response.dict()) await user_api_key_cache.async_set_cache( key=key, value=_perm_obj, @@ -2804,11 +2786,14 @@ async def get_org_object( return deserialized_org # else, check db try: - query_kwargs: Final[dict[str, Any]] = {"where": {"organization_id": org_id}} - if include_budget_table: - query_kwargs["include"] = {"litellm_budget_table": True} - - response: Final = await OrganizationRepository(prisma_client).table.find_unique(**query_kwargs) + organization_repository: Final = OrganizationRepository(prisma_client) + _org_response: Final = ( + await organization_repository.find_by_id_with_relations( + org_id, include={"litellm_budget_table": True}, id_field="organization_id" + ) + if include_budget_table + else await organization_repository.find_by_id(org_id) + ) except Exception: # An operational failure (DB down, timeout, cache fault) is NOT the same fact as a confirmed # missing row, and relabelling it as "doesn't exist" made every caller unable to tell them @@ -2817,12 +2802,12 @@ async def get_org_object( # Exception are unaffected. raise - if response is None: + if _org_response is None: raise OrganizationNotFoundError( f"Organization doesn't exist in db. Organization={org_id}. Create organization via `/organization/new` call." ) - _org_obj: Final = LiteLLM_OrganizationTable.model_validate(response.model_dump()) + _org_obj: Final = _org_response # Cache the result await user_api_key_cache.async_set_cache( key=cache_key, @@ -3763,7 +3748,7 @@ async def _virtual_key_soft_budget_check( ) -def _parse_email_list(raw: Any) -> list[str]: +def _parse_email_list(raw: object) -> list[str]: """Parse emails from a list or comma-separated string.""" if isinstance(raw, list): return [e.strip() for e in raw if isinstance(e, str) and e.strip()] @@ -4294,9 +4279,8 @@ async def get_project_object( return deserialized_project # Fetch from DB - project_row: Final = await ProjectRepository(prisma_client).table.find_unique( - where={"project_id": project_id}, - include={"litellm_budget_table": True}, + project_row: Final = await ProjectRepository(prisma_client).find_by_id_with_relations( + project_id, include={"litellm_budget_table": True}, id_field="project_id" ) if project_row is None: return None diff --git a/litellm/repositories/base_repository.py b/litellm/repositories/base_repository.py index 7008099fe8c..ce6673c3b4f 100644 --- a/litellm/repositories/base_repository.py +++ b/litellm/repositories/base_repository.py @@ -8,17 +8,19 @@ from typing import Any, Final, Generic, Protocol, TypeVar, runtime_checkable from pydantic import BaseModel +from litellm.repositories.prisma_protocols import PrismaTableActions + T = TypeVar("T", bound=BaseModel) @runtime_checkable class SupportsModelDump(Protocol): - def model_dump(self) -> dict[str, object]: ... + def model_dump(self) -> Mapping[str, object]: ... @runtime_checkable class SupportsDict(Protocol): - def dict(self) -> dict[str, object]: ... + def dict(self) -> Mapping[str, object]: ... DbRecord = Mapping[str, object] | SupportsModelDump | SupportsDict | Sequence[tuple[str, object]] @@ -53,6 +55,11 @@ class BaseRepository(ABC, Generic[T]): """Return the Prisma table for this repository.""" ... + @property + def typed_table(self) -> PrismaTableActions: # any-ok: Prisma table actions are untyped at runtime + """The repository table, narrowed to the action surface the repositories use.""" + return self.table + @property @abstractmethod def model_class(self) -> type[T]: @@ -71,7 +78,14 @@ class BaseRepository(ABC, Generic[T]): async def find_by_id(self, id_value: str, id_field: str = "id") -> T | None: """Find a record by its primary key.""" - record: Final = await self.table.find_unique(where={id_field: id_value}) + record: Final = await self.typed_table.find_unique(where={id_field: id_value}) + return self._to_model(record) + + async def find_by_id_with_relations( + self, id_value: str, *, include: Mapping[str, object], id_field: str = "id" + ) -> T | None: + """Find a record by its primary key, eagerly loading the given Prisma relations.""" + record: Final = await self.typed_table.find_unique(where={id_field: id_value}, include=include) return self._to_model(record) async def find_many( @@ -82,41 +96,36 @@ class BaseRepository(ABC, Generic[T]): order: dict[str, str] | None = None, ) -> list[T]: """Find multiple records matching the criteria.""" - kwargs: Final[dict[str, Any]] = {} - if where: - kwargs["where"] = where - if skip is not None: - kwargs["skip"] = skip - if take is not None: - kwargs["take"] = take - if order: - kwargs["order"] = order - - records: Final = await self.table.find_many(**kwargs) + records: Final = await self.typed_table.find_many( + where=where or None, + skip=skip, + take=take, + order=order or None, + ) return self._to_model_list(records) async def create(self, data: dict[str, Any]) -> T: """Create a new record.""" - record: Final = await self.table.create(data=data) + record: Final = await self.typed_table.create(data=data) model: Final = self._to_model(record) assert model is not None return model async def update(self, id_value: str, data: dict[str, Any], id_field: str = "id") -> T | None: """Update an existing record.""" - record: Final = await self.table.update(where={id_field: id_value}, data=data) + record: Final = await self.typed_table.update(where={id_field: id_value}, data=data) return self._to_model(record) async def delete(self, id_value: str, id_field: str = "id") -> T | None: """Delete a record by its primary key.""" - record: Final = await self.table.delete(where={id_field: id_value}) + record: Final = await self.typed_table.delete(where={id_field: id_value}) return self._to_model(record) async def count(self, where: dict[str, Any] | None = None) -> int: """Count records matching the criteria.""" - return await self.table.count(where=where) + return await self.typed_table.count(where=where) async def exists(self, id_value: str, id_field: str = "id") -> bool: """Check if a record exists.""" - record: Final = await self.table.find_unique(where={id_field: id_value}) + record: Final = await self.typed_table.find_unique(where={id_field: id_value}) return record is not None diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index 6aff196ff10..3ca531a8b89 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -16,6 +16,40 @@ class PrismaRecord(Protocol): def dict(self) -> Mapping[str, object]: ... +class PrismaTableActions(Protocol): + """The subset of prisma-client-py table actions the repositories rely on.""" + + async def find_unique( + self, *, where: Mapping[str, object], include: Mapping[str, object] | None = None + ) -> PrismaRecord | None: ... + + async def find_first( + self, *, where: Mapping[str, object], include: Mapping[str, object] | None = None + ) -> PrismaRecord | None: ... + + async def find_many( + self, + *, + where: Mapping[str, object] | None = None, + include: Mapping[str, object] | None = None, + skip: int | None = None, + take: int | None = None, + order: Mapping[str, str] | None = None, + ) -> Sequence[PrismaRecord]: ... + + async def create( + self, *, data: Mapping[str, object], include: Mapping[str, object] | None = None + ) -> PrismaRecord: ... + + async def update( + self, *, where: Mapping[str, object], data: Mapping[str, object] + ) -> PrismaRecord | None: ... + + async def delete(self, *, where: Mapping[str, object]) -> PrismaRecord | None: ... + + async def count(self, *, where: Mapping[str, object] | None = None) -> int: ... + + class ReadOnlyTable(Protocol): async def find_many(self, *, where: Mapping[str, object]) -> Sequence[PrismaRecord]: ... diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py index be19f290ba6..c497deb79ca 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -7,9 +7,21 @@ These are thin wrappers for tables that do not (yet) need domain-specific query methods; richer repositories live in their own modules. """ -from typing import Any +from collections.abc import Mapping +from typing import Any, Final, Generic, TypeVar +from pydantic import BaseModel, TypeAdapter + +from litellm.models.access_group import LiteLLM_AccessGroupTable +from litellm.models.end_user import LiteLLM_EndUserTable +from litellm.models.tag import LiteLLM_TagTable +from litellm.models.team_membership import LiteLLM_TeamMembership from litellm.proxy.common_utils.config_sync_pubsub import wrap_table_actions_for_config_sync +from litellm.repositories.prisma_protocols import PrismaTableActions + +ModelT = TypeVar("ModelT", bound=BaseModel) + +_TOKEN_ADAPTER: Final = TypeAdapter(str) class PrismaTableRepository: @@ -33,6 +45,42 @@ class PrismaTableRepository: table_name=self.table_name, ) + @property + def typed_table(self) -> PrismaTableActions: # any-ok: Prisma table actions are untyped at runtime + """The repository table, narrowed to the action surface the repositories use.""" + return self.table + + +class ModelBackedTable(Generic[ModelT]): + """Mixin that validates rows of a prisma table into a pydantic model before handing them out.""" + + model_class: type[ModelT] + + @property + def typed_table(self) -> PrismaTableActions: + raise NotImplementedError + + async def find_unique_model( + self, *, where: Mapping[str, object], include: Mapping[str, object] | None = None + ) -> ModelT | None: + """Load a single row by a unique constraint and validate it into the domain model.""" + record: Final = await self.typed_table.find_unique(where=where, include=include) + return None if record is None else self.model_class.model_validate(record.dict()) + + async def find_first_model( + self, *, where: Mapping[str, object], include: Mapping[str, object] | None = None + ) -> ModelT | None: + """Load the first row matching the filter and validate it into the domain model.""" + record: Final = await self.typed_table.find_first(where=where, include=include) + return None if record is None else self.model_class.model_validate(record.dict()) + + async def find_many_models( + self, *, where: Mapping[str, object], include: Mapping[str, object] | None = None + ) -> tuple[ModelT, ...]: + """Load every row matching the filter and validate them into domain models.""" + records: Final = await self.typed_table.find_many(where=where, include=include) + return tuple(self.model_class.model_validate(record.dict()) for record in records) + class PolicyRepository(PrismaTableRepository): table_name = "litellm_policytable" @@ -70,12 +118,14 @@ class ClaudeCodePluginRepository(PrismaTableRepository): table_name = "litellm_claudecodeplugintable" -class TeamMembershipRepository(PrismaTableRepository): +class TeamMembershipRepository(PrismaTableRepository, ModelBackedTable[LiteLLM_TeamMembership]): table_name = "litellm_teammembership" + model_class = LiteLLM_TeamMembership -class EndUserRepository(PrismaTableRepository): +class EndUserRepository(PrismaTableRepository, ModelBackedTable[LiteLLM_EndUserTable]): table_name = "litellm_endusertable" + model_class = LiteLLM_EndUserTable class ManagedVectorStoresRepository(PrismaTableRepository): @@ -94,8 +144,9 @@ class PromptRepository(PrismaTableRepository): table_name = "litellm_prompttable" -class TagRepository(PrismaTableRepository): +class TagRepository(PrismaTableRepository, ModelBackedTable[LiteLLM_TagTable]): table_name = "litellm_tagtable" + model_class = LiteLLM_TagTable class InvitationLinkRepository(PrismaTableRepository): @@ -105,6 +156,17 @@ class InvitationLinkRepository(PrismaTableRepository): class JWTKeyMappingRepository(PrismaTableRepository): table_name = "litellm_jwtkeymapping" + async def find_active_token(self, *, jwt_claim_name: str, jwt_claim_value: str) -> str | None: + """Return the key token mapped to an active JWT claim pair.""" + record: Final = await self.typed_table.find_first( + where={ + "jwt_claim_name": jwt_claim_name, + "jwt_claim_value": jwt_claim_value, + "is_active": True, + } + ) + return None if record is None else _TOKEN_ADAPTER.validate_python(record.dict()["token"]) + class ManagedFileRepository(PrismaTableRepository): table_name = "litellm_managedfiletable" @@ -142,8 +204,9 @@ class ModelTableRepository(PrismaTableRepository): table_name = "litellm_modeltable" -class AccessGroupRepository(PrismaTableRepository): +class AccessGroupRepository(PrismaTableRepository, ModelBackedTable[LiteLLM_AccessGroupTable]): table_name = "litellm_accessgrouptable" + model_class = LiteLLM_AccessGroupTable class SSOConfigRepository(PrismaTableRepository):