mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
commit
92b380d4bb
19 changed files with 161 additions and 142 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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__),
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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!")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 23267
|
||||
"limit": 23261
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 27434
|
||||
"limit": 27433
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 292
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue