Merge pull request #35240 from BerriAI/litellm_daily_any_cleanup_07_30_2026_run2

chore(typing): clear basedpyright Any errors in proxy auth, repositories, and openai transforms
This commit is contained in:
Mateo Wang 2026-07-30 10:41:18 -07:00 • committed by GitHub
commit 92b380d4bb
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
19 changed files with 161 additions and 142 deletions

View file

@ -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

View file

@ -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(

View file

@ -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)

View file

@ -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

View file

@ -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:

View file

@ -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__),

View file

@ -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:

View file

@ -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:

View file

@ -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,

View file

@ -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!")

View file

@ -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,

View file

@ -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."""

View file

@ -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,

View file

@ -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,

View file

@ -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:

View file

@ -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)

View file

@ -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

View file

@ -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:

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 23267
"limit": 23261
},
"LIT002": {
"limit": 27434
"limit": 27433
},
"LIT003": {
"limit": 292