From 76cf3bf6acaf900c0e00ba93d53c6745930e5cbd Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 30 Jul 2026 13:48:43 +0000 Subject: [PATCH] chore(typing): clear basedpyright Any errors in proxy auth, repositories, and openai transforms Replace `Model(**untyped_dict)` construction with `Model.model_validate(...)` at the hot Any seams, and give the repository layer a real record type instead of `Any`. reportAny 22710 -> 21448, reportExplicitAny 7283 -> 7269, with every other rule at or below its baseline repo-wide. --- basedpyright-code-budget.json | 10 ++-- litellm/llms/openai/openai.py | 24 ++++---- .../llms/openai/responses/transformation.py | 8 +-- .../mcp_server/mcp_server_manager.py | 2 +- litellm/proxy/auth/auth_checks.py | 4 +- litellm/proxy/auth/oauth2_proxy_hook.py | 19 +++--- litellm/proxy/auth/resolvers/store.py | 2 +- .../management_endpoints/common_utils.py | 4 +- .../mcp_management_endpoints.py | 4 +- litellm/proxy/management_helpers/utils.py | 8 +-- litellm/proxy/proxy_server.py | 8 +-- litellm/repositories/base_repository.py | 54 ++++++++++------- .../repositories/organization_repository.py | 6 +- litellm/repositories/project_repository.py | 9 +-- litellm/repositories/team_repository.py | 56 +++++++++--------- .../verification_token_repository.py | 58 +++++++++---------- ruff-strict-budget.json | 4 +- .../repositories/test_repositories.py | 19 ++++-- type-discipline-budget.json | 4 +- 19 files changed, 161 insertions(+), 142 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 3f89ff179eb..f4e506afcde 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,6 +1,6 @@ { "reportAny": { - "limit": 33171 + "limit": 31909 }, "reportArgumentType": { "limit": 2645 @@ -24,7 +24,7 @@ "limit": 42 }, "reportExplicitAny": { - "limit": 10228 + "limit": 10214 }, "reportFunctionMemberAccess": { "limit": 11 @@ -90,7 +90,7 @@ "limit": 12 }, "reportReturnType": { - "limit": 221 + "limit": 219 }, "reportTypedDictNotRequiredAccess": { "limit": 27 @@ -99,7 +99,7 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 45522 + "limit": 45366 }, "reportUnknownLambdaType": { "limit": 113 @@ -111,7 +111,7 @@ "limit": 20341 }, "reportUnknownVariableType": { - "limit": 32052 + "limit": 32051 }, "reportUnnecessaryCast": { "limit": 177 diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 6b191144a11..8fc7e6d0ebd 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -1626,7 +1626,7 @@ class OpenAIFilesAPI(BaseLLM): openai_client: AsyncOpenAI, ) -> OpenAIFileObject: response = await openai_client.files.create(**create_file_data) # type: ignore[arg-type] - return OpenAIFileObject(**response.model_dump()) + return OpenAIFileObject.model_validate(response.model_dump()) def create_file( self, @@ -1662,7 +1662,7 @@ class OpenAIFilesAPI(BaseLLM): create_file_data=create_file_data, openai_client=openai_client ) response = cast(OpenAI, openai_client).files.create(**create_file_data) # type: ignore[arg-type] - return OpenAIFileObject(**response.model_dump()) + return OpenAIFileObject.model_validate(response.model_dump()) async def afile_content( self, @@ -1986,7 +1986,7 @@ class OpenAIBatchesAPI(BaseLLM): openai_client: AsyncOpenAI, ) -> LiteLLMBatch: response = await openai_client.batches.create(**create_batch_data) # type: ignore[arg-type] - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) def create_batch( self, @@ -2023,7 +2023,7 @@ class OpenAIBatchesAPI(BaseLLM): ) response = cast(OpenAI, openai_client).batches.create(**create_batch_data) # type: ignore[arg-type] - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) async def aretrieve_batch( self, @@ -2032,7 +2032,7 @@ class OpenAIBatchesAPI(BaseLLM): ) -> LiteLLMBatch: verbose_logger.debug("retrieving batch, args= %s", retrieve_batch_data) response = await openai_client.batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type] - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) def retrieve_batch( self, @@ -2068,7 +2068,7 @@ class OpenAIBatchesAPI(BaseLLM): retrieve_batch_data=retrieve_batch_data, openai_client=openai_client ) response = cast(OpenAI, openai_client).batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type] - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) async def acancel_batch( self, @@ -2077,7 +2077,7 @@ class OpenAIBatchesAPI(BaseLLM): ) -> LiteLLMBatch: verbose_logger.debug("async cancelling batch, args= %s", cancel_batch_data) response = await openai_client.batches.cancel(**cancel_batch_data) - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) def cancel_batch( self, @@ -2117,7 +2117,7 @@ class OpenAIBatchesAPI(BaseLLM): if not isinstance(openai_client, OpenAI): raise ValueError("OpenAI client is not an instance of OpenAI. Make sure you passed a sync OpenAI client.") response = openai_client.batches.cancel(**cancel_batch_data) - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) async def alist_batches( self, @@ -2477,9 +2477,9 @@ class OpenAIAssistantsAPI(BaseLLM): response_obj: Optional[OpenAIMessage] = None if getattr(thread_message, "status", None) is None: thread_message.status = "completed" - response_obj = OpenAIMessage(**thread_message.dict()) + response_obj = OpenAIMessage.model_validate(thread_message.dict()) else: - response_obj = OpenAIMessage(**thread_message.dict()) + response_obj = OpenAIMessage.model_validate(thread_message.dict()) return response_obj # fmt: off @@ -2556,9 +2556,9 @@ class OpenAIAssistantsAPI(BaseLLM): response_obj: Optional[OpenAIMessage] = None if getattr(thread_message, "status", None) is None: thread_message.status = "completed" - response_obj = OpenAIMessage(**thread_message.dict()) + response_obj = OpenAIMessage.model_validate(thread_message.dict()) else: - response_obj = OpenAIMessage(**thread_message.dict()) + response_obj = OpenAIMessage.model_validate(thread_message.dict()) return response_obj async def async_get_messages( diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 3c2ae238a0b..dc4e98e6216 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -280,7 +280,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): raw_response_headers = dict(raw_response.headers) processed_headers = process_response_headers(raw_response_headers) try: - response = ResponsesAPIResponse(**raw_response_json) + response = ResponsesAPIResponse.model_validate(raw_response_json) except Exception: verbose_logger.debug(f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct") response = ResponsesAPIResponse.model_construct(**raw_response_json) @@ -506,7 +506,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code) raw_response_headers = dict(raw_response.headers) processed_headers = process_response_headers(raw_response_headers) - response = ResponsesAPIResponse(**raw_response_json) + response = ResponsesAPIResponse.model_validate(raw_response_json) response._hidden_params["additional_headers"] = processed_headers response._hidden_params["headers"] = raw_response_headers @@ -588,7 +588,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): raw_response_headers = dict(raw_response.headers) processed_headers = process_response_headers(raw_response_headers) - response = ResponsesAPIResponse(**raw_response_json) + response = ResponsesAPIResponse.model_validate(raw_response_json) response._hidden_params["additional_headers"] = processed_headers response._hidden_params["headers"] = raw_response_headers @@ -647,7 +647,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): processed_headers = process_response_headers(raw_response_headers) try: - response = ResponsesAPIResponse(**raw_response_json) + response = ResponsesAPIResponse.model_validate(raw_response_json) except Exception: verbose_logger.debug(f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct") response = ResponsesAPIResponse.model_construct(**raw_response_json) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 3e0775ac09e..f61ac4866b0 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -5322,7 +5322,7 @@ class MCPServerManager: ] } ) - db_mcp_servers = [LiteLLM_MCPServerTable(**r.model_dump()) for r in raw_rows] + db_mcp_servers = [LiteLLM_MCPServerTable.model_validate(r.model_dump()) for r in raw_rows] verbose_logger.info(f"Found {len(db_mcp_servers)} MCP servers in database") previous_registry = self.registry diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index c46bc110ca8..9d023292074 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2434,7 +2434,7 @@ class ExperimentalUIJWTToken: if decrypted_token is None: return None try: - return UserAPIKeyAuth(**json.loads(decrypted_token)) + return UserAPIKeyAuth.model_validate(json.loads(decrypted_token)) except Exception as e: raise Exception(f"Invalid hash key. Hash key={hashed_token}. Decrypted token={decrypted_token}. Error: {e}") @@ -2553,7 +2553,7 @@ async def get_key_object( code=status.HTTP_401_UNAUTHORIZED, ) - _response = UserAPIKeyAuth(**_valid_token.model_dump(exclude_none=True)) + _response = UserAPIKeyAuth.model_validate(_valid_token.model_dump(exclude_none=True)) # Load object_permission if object_permission_id exists but object_permission is not loaded if _response.object_permission_id and not _response.object_permission: diff --git a/litellm/proxy/auth/oauth2_proxy_hook.py b/litellm/proxy/auth/oauth2_proxy_hook.py index 2b0593d3618..ca6a7ee4b1d 100644 --- a/litellm/proxy/auth/oauth2_proxy_hook.py +++ b/litellm/proxy/auth/oauth2_proxy_hook.py @@ -1,4 +1,5 @@ -from typing import Any, Dict, FrozenSet +from collections.abc import Mapping +from typing import Dict, FrozenSet, List, Union from fastapi import Request @@ -83,21 +84,17 @@ async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth: "(signature-validated) instead of header-trust." ) - auth_data: Dict[str, Any] = {} - for key, header in oauth2_config_mappings.items(): - value = request.headers.get(header) - if not value: - continue - if key == "models": - auth_data[key] = [model.strip() for model in value.split(",")] - else: - auth_data[key] = value + auth_data: Mapping[str, Union[str, List[str]]] = { + key: [model.strip() for model in value.split(",")] if key == "models" else value + for key, header in oauth2_config_mappings.items() + if (value := request.headers.get(header)) + } verbose_proxy_logger.debug( "Auth data before creating UserAPIKeyAuth object: keys=%s", list(auth_data.keys()), ) - user_api_key_auth = UserAPIKeyAuth(**auth_data) + user_api_key_auth = UserAPIKeyAuth.model_validate(auth_data) verbose_proxy_logger.debug( "UserAPIKeyAuth object created with keys: %s", list(user_api_key_auth.__fields_set__), diff --git a/litellm/proxy/auth/resolvers/store.py b/litellm/proxy/auth/resolvers/store.py index ad9fd234163..7c2bd324064 100644 --- a/litellm/proxy/auth/resolvers/store.py +++ b/litellm/proxy/auth/resolvers/store.py @@ -118,7 +118,7 @@ class IdentityStore: if from_db is None: raise KeyNotFoundError(hashed_token) - key = UserAPIKeyAuth(**from_db.model_dump(exclude_none=True)) + key = UserAPIKeyAuth.model_validate(from_db.model_dump(exclude_none=True)) if key.object_permission_id and not key.object_permission: try: diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index 8162babef40..877130c2066 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -216,7 +216,7 @@ async def _user_has_admin_privileges( teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_obj.teams}}) for team in teams: - team_obj = LiteLLM_TeamTable(**team.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team.model_dump()) if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): return True @@ -288,7 +288,7 @@ async def _team_admin_can_invite_user( for team in teams if _is_user_team_admin( user_api_key_dict=user_api_key_dict, - team_obj=LiteLLM_TeamTable(**team.model_dump()), + team_obj=LiteLLM_TeamTable.model_validate(team.model_dump()), ) ] if not admin_team_ids: diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 1205d23ce02..64cc13a5543 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -459,7 +459,7 @@ if MCP_AVAILABLE: payload_dict: dict[str, Any] = loaded try: - return MCPServer(**payload_dict) + return MCPServer.model_validate(payload_dict) except Exception as e: verbose_proxy_logger.debug(f"Invalid temporary MCP server payload in Redis cache: {str(e)}") return None @@ -704,7 +704,7 @@ if MCP_AVAILABLE: except AttributeError: payload_dict = payload.dict() # type: ignore[attr-defined] payload_dict["credentials"] = inherited_credentials - return NewMCPServerRequest(**payload_dict) + return NewMCPServerRequest.model_validate(payload_dict) def _build_temporary_mcp_server_record( payload: NewMCPServerRequest, diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index 7c52b04c4eb..14ba1dbfde1 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -308,7 +308,7 @@ async def add_new_member( ) await _append_team_id_if_absent(prisma_client, new_member.user_id, team_id) if _returned_user is not None: - returned_user = LiteLLM_UserTable(**_returned_user.model_dump()) + returned_user = LiteLLM_UserTable.model_validate(_returned_user.model_dump()) elif new_member.user_email is not None: new_user_defaults = get_new_internal_user_defaults(user_id=str(uuid.uuid4()), user_email=new_member.user_email) ## user email is not unique acc. to prisma schema -> future improvement @@ -323,11 +323,11 @@ async def add_new_member( _returned_user = await prisma_client.insert_data(data=new_user_defaults, table_name="user") # type: ignore if _returned_user is not None: - returned_user = LiteLLM_UserTable(**_returned_user.model_dump()) + returned_user = LiteLLM_UserTable.model_validate(_returned_user.model_dump()) elif len(existing_user_row) == 1: user_info = existing_user_row[0] await _append_team_id_if_absent(prisma_client, user_info.user_id, team_id) - returned_user = LiteLLM_UserTable(**user_info.model_dump()) + returned_user = LiteLLM_UserTable.model_validate(user_info.model_dump()) elif len(existing_user_row) > 1: raise HTTPException( status_code=400, @@ -354,7 +354,7 @@ async def add_new_member( include={"litellm_budget_table": True}, ) - returned_team_membership = LiteLLM_TeamMembership(**_returned_team_membership.model_dump()) + returned_team_membership = LiteLLM_TeamMembership.model_validate(_returned_team_membership.model_dump()) if returned_user is None: raise Exception("Unable to update user table with membership information!") diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 18a927e7a44..ba41083e9b8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -5398,7 +5398,7 @@ class ProxyConfig: # decrypt values for k, v in _litellm_params.items(): _litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v) - _litellm_params = LiteLLM_Params(**_litellm_params) + _litellm_params = LiteLLM_Params.model_validate(_litellm_params) else: verbose_proxy_logger.error( @@ -5429,7 +5429,7 @@ class ProxyConfig: # decrypt values for k, v in _litellm_params.items(): _litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v) - _litellm_params = LiteLLM_Params(**_litellm_params) + _litellm_params = LiteLLM_Params.model_validate(_litellm_params) else: verbose_proxy_logger.error( f"Invalid model added to proxy db. Invalid litellm params. litellm_params={_litellm_params}" @@ -13063,7 +13063,7 @@ def _get_model_group_info( _model_group_info = llm_router.get_model_group_info(model_group=model) if _model_group_info is not None: - model_groups.append(ModelGroupInfoProxy(**_model_group_info.model_dump())) + model_groups.append(ModelGroupInfoProxy.model_validate(_model_group_info.model_dump())) else: model_group_info = ModelGroupInfoProxy( model_group=model, @@ -14782,7 +14782,7 @@ async def update_config_general_settings( ) try: - ConfigGeneralSettings(**{data.field_name: data.field_value}) + ConfigGeneralSettings.model_validate({data.field_name: data.field_value}) except Exception: raise HTTPException( status_code=400, diff --git a/litellm/repositories/base_repository.py b/litellm/repositories/base_repository.py index 40aeb6df3de..755e4595c01 100644 --- a/litellm/repositories/base_repository.py +++ b/litellm/repositories/base_repository.py @@ -3,38 +3,58 @@ Base repository class with common functionality. """ from abc import ABC, abstractmethod -from typing import Any, Dict, Generic, List, Optional, Type, TypeVar +from collections.abc import Iterable, Mapping, Sequence +from typing import Any, Dict, Generic, List, Optional, Protocol, Tuple, Type, TypeVar, Union, runtime_checkable from pydantic import BaseModel T = TypeVar("T", bound=BaseModel) -def _record_to_dict(record: Any) -> Dict[str, Any]: - if isinstance(record, dict): - return record - if hasattr(record, "model_dump") and callable(record.model_dump): +@runtime_checkable +class SupportsModelDump(Protocol): + def model_dump(self) -> Dict[str, object]: ... + + +@runtime_checkable +class SupportsDict(Protocol): + def dict(self) -> Dict[str, object]: ... + + +DbRecord = Union[ + Mapping[str, object], + SupportsModelDump, + SupportsDict, + Sequence[Tuple[str, object]], +] + + +def record_to_dict(record: DbRecord) -> Mapping[str, object]: + """Project a database record into a mapping of column name to value.""" + if isinstance(record, SupportsModelDump): return record.model_dump() - if hasattr(record, "dict") and callable(record.dict): + if isinstance(record, SupportsDict): return record.dict() - return dict(record) + if isinstance(record, Mapping): + return record + return {key: value for key, value in record} class BaseRepository(ABC, Generic[T]): """Abstract base class for all repositories.""" - def __init__(self, prisma_client: Any): + def __init__(self, prisma_client: Any): # any-ok: PrismaClient is an untyped runtime wrapper self._prisma_client = prisma_client @property - def prisma_client(self) -> Any: + def prisma_client(self) -> Any: # any-ok: PrismaClient is an untyped runtime wrapper if self._prisma_client is None: raise RuntimeError("No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys") return self._prisma_client @property @abstractmethod - def table(self) -> Any: + def table(self) -> Any: # any-ok: Prisma table actions are reached through the untyped client wrapper """Return the Prisma table for this repository.""" ... @@ -44,21 +64,15 @@ class BaseRepository(ABC, Generic[T]): """Return the domain model class for this repository.""" ... - def _to_model(self, record: Any) -> Optional[T]: + def _to_model(self, record: Optional[DbRecord]) -> Optional[T]: """Convert a database record to a domain model.""" if record is None: return None - return self.model_class(**_record_to_dict(record)) + return self.model_class.model_validate(record_to_dict(record)) - def _to_model_list(self, records: List[Any]) -> List[T]: + def _to_model_list(self, records: Iterable[Optional[DbRecord]]) -> List[T]: """Convert a list of database records to domain models.""" - result: List[T] = [] - for r in records: - if r is not None: - model = self._to_model(r) - if model is not None: - result.append(model) - return result + return [model for record in records if record is not None and (model := self._to_model(record)) is not None] async def find_by_id(self, id_value: str, id_field: str = "id") -> Optional[T]: """Find a record by its primary key.""" diff --git a/litellm/repositories/organization_repository.py b/litellm/repositories/organization_repository.py index d5f8c990001..99c4a881736 100644 --- a/litellm/repositories/organization_repository.py +++ b/litellm/repositories/organization_repository.py @@ -26,10 +26,8 @@ class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]): async def find_by_alias(self, organization_alias: str) -> Optional[LiteLLM_OrganizationTable]: """Find an organization by alias.""" - records = await self.table.find_many(where={"organization_alias": organization_alias}) - if records: - return self._to_model(records[0]) - return None + organizations = await self.find_many(where={"organization_alias": organization_alias}) + return organizations[0] if organizations else None async def create_organization( self, diff --git a/litellm/repositories/project_repository.py b/litellm/repositories/project_repository.py index 86faaf2e13c..27cb346e1b1 100644 --- a/litellm/repositories/project_repository.py +++ b/litellm/repositories/project_repository.py @@ -24,15 +24,12 @@ class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]): async def find_by_alias(self, project_alias: str) -> Optional[LiteLLM_ProjectTable]: """Find a project by alias.""" - records = await self.table.find_many(where={"project_alias": project_alias}) - if records: - return self._to_model(records[0]) - return None + projects = await self.find_many(where={"project_alias": project_alias}) + return projects[0] if projects else None async def find_by_team_id(self, team_id: str) -> List[LiteLLM_ProjectTable]: """Find all projects belonging to a team.""" - records = await self.table.find_many(where={"team_id": team_id}) - return self._to_model_list(records) + return await self.find_many(where={"team_id": team_id}) async def create_project( self, diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index 68875bd7972..25437cfe49a 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -3,55 +3,59 @@ Team repository for database operations on LiteLLM_TeamTable. """ import json +from collections.abc import Mapping from datetime import datetime from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type from pydantic import TypeAdapter from litellm.models.team import LiteLLM_TeamTable, Member -from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.base_repository import ( + BaseRepository, + DbRecord, + record_to_dict, +) if TYPE_CHECKING: from prisma import Prisma _MEMBERS_WITH_ROLES_ADAPTER = TypeAdapter(list[Member]) +_JSON_ENCODED_TEAM_FIELDS = ( + "metadata", + "model_spend", + "model_max_budget", + "router_settings", + "budget_limits", + "members_with_roles", +) class TeamRepository(BaseRepository[LiteLLM_TeamTable]): """Repository for team database operations.""" @property - def table(self) -> Any: + def table(self) -> Any: # any-ok: PrismaClient.db is an untyped runtime wrapper return self.prisma_client.db.litellm_teamtable @property - def deleted_table(self) -> Any: + def deleted_table(self) -> Any: # any-ok: PrismaClient.db is an untyped runtime wrapper return self.prisma_client.db.litellm_deletedteamtable @property def model_class(self) -> Type[LiteLLM_TeamTable]: return LiteLLM_TeamTable - def _to_model(self, record: Any) -> Optional[LiteLLM_TeamTable]: + def _to_model(self, record: Optional[DbRecord]) -> Optional[LiteLLM_TeamTable]: """Convert a database record to a Team model.""" if record is None: return None - data = record.dict() if hasattr(record, "dict") else dict(record) + data = { + field: json.loads(value) if field in _JSON_ENCODED_TEAM_FIELDS and isinstance(value, str) else value + for field, value in record_to_dict(record).items() + } - json_fields = [ - "metadata", - "model_spend", - "model_max_budget", - "router_settings", - "budget_limits", - "members_with_roles", - ] - for field in json_fields: - if isinstance(data.get(field), str): - data[field] = json.loads(data[field]) - - return LiteLLM_TeamTable(**data) + return LiteLLM_TeamTable.model_validate(data) async def get_members_with_roles_locked(self, tx: "Prisma", team_id: str) -> List[Member]: """Return the team's members_with_roles, locking the row FOR UPDATE. @@ -103,8 +107,8 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): organization_id: Optional[str] = None, admins: Optional[List[str]] = None, members: Optional[List[str]] = None, - members_with_roles: Optional[Dict[str, Any]] = None, - metadata: Optional[Dict[str, Any]] = None, + members_with_roles: Optional[Mapping[str, object]] = None, + metadata: Optional[Mapping[str, object]] = None, max_budget: Optional[float] = None, soft_budget: Optional[float] = None, models: Optional[List[str]] = None, @@ -115,7 +119,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): object_permission_id: Optional[str] = None, ) -> LiteLLM_TeamTable: """Create a new team.""" - data: Dict[str, Any] = {"team_id": team_id} + data: Dict[str, object] = {"team_id": team_id} if team_alias is not None: data["team_alias"] = team_alias if organization_id is not None: @@ -154,8 +158,8 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): organization_id: Optional[str] = None, admins: Optional[List[str]] = None, members: Optional[List[str]] = None, - members_with_roles: Optional[Dict[str, Any]] = None, - metadata: Optional[Dict[str, Any]] = None, + members_with_roles: Optional[Mapping[str, object]] = None, + metadata: Optional[Mapping[str, object]] = None, max_budget: Optional[float] = None, soft_budget: Optional[float] = None, models: Optional[List[str]] = None, @@ -167,7 +171,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): object_permission_id: Optional[str] = None, ) -> Optional[LiteLLM_TeamTable]: """Update a team.""" - data: Dict[str, Any] = {} + data: Dict[str, object] = {} if team_alias is not None: data["team_alias"] = team_alias if organization_id is not None: @@ -228,9 +232,9 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): return team - def _build_archive_data(self, team: LiteLLM_TeamTable) -> Dict[str, Any]: + def _build_archive_data(self, team: LiteLLM_TeamTable) -> Dict[str, object]: """Build archive data dict with only columns that exist in LiteLLM_DeletedTeamTable.""" - data: Dict[str, Any] = {"team_id": team.team_id} + data: Dict[str, object] = {"team_id": team.team_id} if team.team_alias is not None: data["team_alias"] = team.team_alias if team.organization_id is not None: diff --git a/litellm/repositories/verification_token_repository.py b/litellm/repositories/verification_token_repository.py index 19352c1b3c4..f7795f15fd5 100644 --- a/litellm/repositories/verification_token_repository.py +++ b/litellm/repositories/verification_token_repository.py @@ -3,14 +3,18 @@ VerificationToken repository for database operations on LiteLLM_VerificationToke """ import json -from collections.abc import Iterator, Mapping +from collections.abc import Mapping from datetime import datetime -from typing import TYPE_CHECKING, Any, Protocol +from typing import TYPE_CHECKING, Any from litellm.models.verification_token import ( LiteLLM_VerificationToken, ) -from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.base_repository import ( + BaseRepository, + DbRecord, + record_to_dict, +) if TYPE_CHECKING: from prisma.models import ( @@ -19,11 +23,17 @@ if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient - -class _DictConvertible(Protocol): - def dict(self) -> dict[str, object]: ... - - def __iter__(self) -> Iterator[tuple[str, object]]: ... +_JSON_ENCODED_TOKEN_FIELDS = ( + "aliases", + "config", + "permissions", + "metadata", + "model_spend", + "model_max_budget", + "router_settings", + "budget_limits", + "litellm_budget_table", +) class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): @@ -46,31 +56,21 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): def model_class(self) -> type[LiteLLM_VerificationToken]: return LiteLLM_VerificationToken - def _to_model(self, record: _DictConvertible | None) -> LiteLLM_VerificationToken | None: + def _to_model(self, record: DbRecord | None) -> LiteLLM_VerificationToken | None: """Convert a database record to a VerificationToken model.""" if record is None: return None - data = record.dict() if hasattr(record, "dict") else dict(record) - - json_fields = [ - "aliases", - "config", - "permissions", - "metadata", - "model_spend", - "model_max_budget", - "router_settings", - "budget_limits", - "litellm_budget_table", - ] - for field in json_fields: - value = data.get(field) - if isinstance(value, str): - data[field] = json.loads(value) - - if data.get("org_id") is None and data.get("organization_id") is not None: - data["org_id"] = data["organization_id"] + decoded = { + field: json.loads(value) if field in _JSON_ENCODED_TOKEN_FIELDS and isinstance(value, str) else value + for field, value in record_to_dict(record).items() + } + organization_id = decoded.get("organization_id") + data = ( + decoded + if decoded.get("org_id") is not None or organization_id is None + else {**decoded, "org_id": organization_id} + ) return LiteLLM_VerificationToken.model_validate(data) diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 3b5ec5b0dee..f1c205c2425 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -24,7 +24,7 @@ "limit": 130 }, "ANN401": { - "limit": 2013 + "limit": 2010 }, "ASYNC230": { "limit": 14 @@ -324,7 +324,7 @@ "limit": 883 }, "UP006": { - "limit": 12146 + "limit": 12142 }, "UP007": { "limit": 2526 diff --git a/tests/test_litellm/repositories/test_repositories.py b/tests/test_litellm/repositories/test_repositories.py index 6308faf8fc7..c923b722991 100644 --- a/tests/test_litellm/repositories/test_repositories.py +++ b/tests/test_litellm/repositories/test_repositories.py @@ -203,23 +203,32 @@ class TestBaseRepository: assert len(budgets) == 1 def test_record_to_dict_branches(self): - from litellm.repositories.base_repository import _record_to_dict + from litellm.repositories.base_repository import record_to_dict - assert _record_to_dict({"a": 1}) == {"a": 1} + assert record_to_dict({"a": 1}) == {"a": 1} class WithModelDump: def model_dump(self): return {"src": "model_dump"} - assert _record_to_dict(WithModelDump()) == {"src": "model_dump"} + assert record_to_dict(WithModelDump()) == {"src": "model_dump"} class WithDict: def dict(self): return {"src": "dict"} - assert _record_to_dict(WithDict()) == {"src": "dict"} + assert record_to_dict(WithDict()) == {"src": "dict"} - assert _record_to_dict([("k", "v")]) == {"k": "v"} + assert record_to_dict([("k", "v")]) == {"k": "v"} + + class WithBoth: + def model_dump(self): + return {"src": "model_dump"} + + def dict(self): + return {"src": "dict"} + + assert record_to_dict(WithBoth()) == {"src": "model_dump"} class TestBudgetRepository: diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 25cc1621d54..69b2506dc08 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 23267 + "limit": 23261 }, "LIT002": { - "limit": 27434 + "limit": 27433 }, "LIT003": { "limit": 292