diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index c35c17aa359..108e4266902 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -14,17 +14,32 @@ import json import math import traceback from datetime import datetime, timezone -from typing import Annotated, Any, Dict, List, Mapping, Optional, Tuple, Union, cast +from typing import ( + AbstractSet, + Annotated, + Any, + Dict, + List, + Mapping, + Optional, + Protocol, + Sequence, + Tuple, + Union, + cast, +) import fastapi from fastapi import APIRouter, Depends, Header, HTTPException, Request, status -from pydantic import BaseModel +from pydantic import BaseModel, JsonValue, TypeAdapter import litellm from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid from litellm.integrations.prometheus import PrometheusLogger from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.models.access_group import LiteLLM_AccessGroupTable +from litellm.models.budget import LiteLLM_BudgetTable, LiteLLM_BudgetTableFull from litellm.proxy._types import ( UI_TEAM_ID, BlockTeamRequest, @@ -78,6 +93,7 @@ from litellm.proxy.auth.auth_checks import ( from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars from litellm.proxy.common_utils.json_merge_patch import apply_json_merge_patch +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.management_endpoints.common_utils import ( _check_passthrough_routes_caller_permission, _is_user_org_admin_for_team, @@ -106,7 +122,7 @@ from litellm.proxy.management_helpers.utils import ( add_new_member, management_endpoint_wrapper, ) -from litellm.proxy.utils import PrismaClient, handle_exception_on_proxy +from litellm.proxy.utils import PrismaClient, ProxyLogging, handle_exception_on_proxy from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.organization_repository import OrganizationRepository from litellm.repositories.table_repositories import ( @@ -151,9 +167,9 @@ def _sanitize_for_log(value: object) -> str: async def _refresh_cached_team( - team_row: Any, - user_api_key_cache: Any, - proxy_logging_obj: Any, + team_row: LiteLLM_TeamTable, + user_api_key_cache: UserApiKeyCache, + proxy_logging_obj: ProxyLogging | None, ) -> None: """ Refresh the in-memory cached team object after a DB write. @@ -274,7 +290,7 @@ class TeamMemberBudgetHandler: if team_member_budget_duration is not None: budget_request.budget_duration = team_member_budget_duration - team_member_budget_table = await new_budget( + team_member_budget_table: LiteLLM_BudgetTable = await new_budget( budget_obj=budget_request, user_api_key_dict=user_api_key_dict, ) @@ -322,7 +338,7 @@ class TeamMemberBudgetHandler: if team_member_budget_duration is not None: budget_request.budget_duration = team_member_budget_duration - budget_row = await update_budget( + budget_row: LiteLLM_BudgetTable = await update_budget( budget_obj=budget_request, user_api_key_dict=user_api_key_dict, ) @@ -395,7 +411,7 @@ class TeamMemberBudgetHandler: @staticmethod async def backfill_team_member_budget_entries( team_id: str, - members_with_roles: List[Union[Member, dict]], + members_with_roles: Sequence[Member | dict], team_member_budget_id: str, prisma_client: PrismaClient, ) -> None: @@ -414,7 +430,9 @@ class TeamMemberBudgetHandler: return # Batch-fetch existing memberships for this team (avoids N+1 queries) - existing_memberships = await TeamMembershipRepository(prisma_client).table.find_many(where={"team_id": team_id}) + existing_memberships: Sequence[LiteLLM_TeamMembership] = await TeamMembershipRepository( + prisma_client + ).table.find_many(where={"team_id": team_id}) existing_user_ids = {m.user_id for m in existing_memberships} # Identify members with no existing membership row. @@ -447,7 +465,7 @@ class TeamMemberBudgetHandler: # Heal existing membership rows that predate the team_member_budget # configuration: populate budget_id where it is currently NULL. # Rows with an explicit budget_id (per-member override) are left alone. - updated = await TeamMembershipRepository(prisma_client).table.update_many( + updated: int = await TeamMembershipRepository(prisma_client).table.update_many( where={"team_id": team_id, "budget_id": None}, data={"budget_id": team_member_budget_id}, ) @@ -460,7 +478,11 @@ class TeamMemberBudgetHandler: ) -def _get_default_team_param(field: str) -> Any: +_DefaultTeamParamValue = float | int | str | Sequence[str] +_DEFAULT_TEAM_PARAM_ADAPTER: TypeAdapter[_DefaultTeamParamValue | None] = TypeAdapter(_DefaultTeamParamValue | None) + + +def _get_default_team_param(field: str) -> _DefaultTeamParamValue | None: """ Returns a default value for the given field from litellm.default_team_params config. Returns None if no default is configured. @@ -470,16 +492,14 @@ def _get_default_team_param(field: str) -> Any: default_params = litellm.default_team_params if default_params is None: return None - if isinstance(default_params, dict): - value = default_params.get(field) - else: - value = getattr(default_params, field, None) - if value is None: + raw_value: object = ( + default_params.get(field) if isinstance(default_params, dict) else getattr(default_params, field, None) + ) + if raw_value is None: return None - # Convert enum values in lists to strings - if isinstance(value, list): - return [v.value if hasattr(v, "value") else v for v in value] - return value + if isinstance(raw_value, list): + return tuple(v.value if hasattr(v, "value") else v for v in raw_value) + return _DEFAULT_TEAM_PARAM_ADAPTER.validate_python(raw_value) def _is_available_team(team_id: str, user_api_key_dict: UserAPIKeyAuth) -> bool: @@ -503,16 +523,12 @@ async def get_all_team_memberships( # else: # where_obj = {"user_id": str(user_id), "team_id": {"in": team_id}} - team_memberships = await TeamMembershipRepository(prisma_client).table.find_many( + team_memberships: Sequence[LiteLLM_TeamMembership] = await TeamMembershipRepository(prisma_client).table.find_many( where=where_obj, include={"litellm_budget_table": True}, ) - returned_tm: List[LiteLLM_TeamMembership] = [] - for tm in team_memberships: - returned_tm.append(LiteLLM_TeamMembership.model_validate(tm.model_dump())) - - return returned_tm + return [LiteLLM_TeamMembership.model_validate(tm.model_dump()) for tm in team_memberships] def _check_team_model_specific_limits( @@ -765,14 +781,12 @@ async def _check_org_team_limits( # calculate allocated tpm/rpm limit # check if specified tpm/rpm limit is greater than allocated tpm/rpm limit - teams = await TeamRepository(prisma_client).table.find_many( + teams: Sequence[LiteLLM_TeamTable] = await TeamRepository(prisma_client).table.find_many( where={"organization_id": org_table.organization_id}, ) # Convert teams to LiteLLM_TeamTable objects - team_objs: List[LiteLLM_TeamTable] = [] - for team in teams: - team_objs.append(LiteLLM_TeamTable.model_validate(team.model_dump())) + team_objs = [LiteLLM_TeamTable.model_validate(team.model_dump()) for team in teams] check_org_team_model_specific_limits( teams=team_objs, @@ -790,7 +804,7 @@ async def _check_user_team_limits( data: Union[NewTeamRequest, UpdateTeamRequest], user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, - user_api_key_cache: Any, + user_api_key_cache: UserApiKeyCache, ) -> None: """ Enforce the caller's personal limits when CREATING a standalone team. @@ -1051,7 +1065,7 @@ async def new_team( ) # Check if license is over limit - total_teams = await TeamRepository(prisma_client).table.count() + total_teams = await TeamRepository(prisma_client).count() if total_teams and _license_check.is_team_count_over_limit(team_count=total_teams): raise HTTPException( status_code=403, @@ -1153,7 +1167,7 @@ async def new_team( created_by=user_api_key_dict.user_id or litellm_proxy_admin_name, updated_by=user_api_key_dict.user_id or litellm_proxy_admin_name, ) - model_dict = await ModelTableRepository(prisma_client).table.create( + model_dict: LiteLLM_ModelTable = await ModelTableRepository(prisma_client).table.create( {**litellm_modeltable.json(exclude_none=True)} # type: ignore ) # type: ignore @@ -1357,11 +1371,11 @@ async def _create_team_update_audit_log( async def _update_model_table( data: UpdateTeamRequest, - model_id: Optional[str], + model_id: int | None, prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, -) -> Optional[str]: +) -> int | None: """ Upsert model table and return the model id """ @@ -1374,7 +1388,7 @@ async def _update_model_table( updated_by=user_api_key_dict.user_id or litellm_proxy_admin_name, ) if model_id is None: - model_dict = await ModelTableRepository(prisma_client).table.create( + model_dict: LiteLLM_ModelTable = await ModelTableRepository(prisma_client).table.create( data={**litellm_modeltable.json(exclude_none=True)} # type: ignore ) else: @@ -1394,7 +1408,7 @@ async def _update_model_table( async def _auto_add_team_members_to_organization( team: LiteLLM_TeamTable, organization: LiteLLM_OrganizationTableWithMembers, - prisma_client: Any, + prisma_client: PrismaClient, ) -> None: """ When moving a team to an org, ensure all team members are also org members. @@ -1432,11 +1446,11 @@ async def _auto_add_team_members_to_organization( async def fetch_and_validate_organization( organization_id: str, - existing_team_row: Any, + existing_team_row: LiteLLM_TeamTable, llm_router: Optional[Router], - prisma_client: Any, + prisma_client: PrismaClient, user_api_key_dict: Optional[UserAPIKeyAuth] = None, -) -> Any: +) -> LiteLLM_OrganizationTable: """ Fetch and validate an organization for team update operations. @@ -1455,7 +1469,7 @@ async def fetch_and_validate_organization( if llm_router is None: raise HTTPException(status_code=500, detail={"error": CommonProxyErrors.no_llm_router.value}) - organization_row = await OrganizationRepository(prisma_client).table.find_unique( + organization_row: LiteLLM_OrganizationTable | None = await OrganizationRepository(prisma_client).table.find_unique( where={"organization_id": organization_id}, include={"litellm_budget_table": True, "members": True, "teams": True}, ) @@ -1704,7 +1718,9 @@ async def update_team( detail={"error": f"soft_budget must be a non-negative finite number. Received: {data.soft_budget}"}, ) - existing_team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": data.team_id}) + existing_team_row: LiteLLM_TeamTable | None = await TeamRepository(prisma_client).table.find_unique( + where={"team_id": data.team_id} + ) if existing_team_row is None: raise HTTPException( @@ -1732,9 +1748,9 @@ async def update_team( ) if data.max_budget is not None: - existing_soft_budget = getattr(existing_team_row, "soft_budget", None) - soft_budget_to_check = data.soft_budget if data.soft_budget is not None else existing_soft_budget - if soft_budget_to_check is not None and isinstance(soft_budget_to_check, (int, float)): + soft_budget_to_check = data.soft_budget if data.soft_budget is not None else existing_team_row.soft_budget + is_numeric_soft_budget = isinstance(soft_budget_to_check, (int, float)) # pyright: ignore[reportUnnecessaryIsInstance] # guards mocked non-float values in tests + if soft_budget_to_check is not None and is_numeric_soft_budget: if data.max_budget <= soft_budget_to_check: raise HTTPException( status_code=400, @@ -1750,7 +1766,7 @@ async def update_team( # so without this gate an org-admin could hand their team to any # other org (or capture a team from another org they once # administered into a new destination). - current_org_id = getattr(existing_team_row, "organization_id", None) + current_org_id = existing_team_row.organization_id if ( data.organization_id != current_org_id and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value @@ -1794,7 +1810,8 @@ async def update_team( org_id_to_check = ( data.organization_id if data.organization_id is not None else existing_team_row.organization_id ) - if org_id_to_check is not None and isinstance(org_id_to_check, str) and prisma_client is not None: + is_str_org_id = isinstance(org_id_to_check, str) # pyright: ignore[reportUnnecessaryIsInstance] # guards mocked non-str values in tests + if org_id_to_check is not None and is_str_org_id and prisma_client is not None: org_table = await get_org_object( org_id=org_id_to_check, user_api_key_cache=user_api_key_cache, @@ -1820,8 +1837,9 @@ async def update_team( # Drop server-owned metadata keys from caller input so they can only # be written by the same code path that creates the underlying rows. - if isinstance(updated_kv.get("metadata"), dict): - TeamMemberBudgetHandler.strip_system_managed_metadata_keys(updated_kv["metadata"]) + updated_kv_metadata = updated_kv.get("metadata") + if isinstance(updated_kv_metadata, dict): + TeamMemberBudgetHandler.strip_system_managed_metadata_keys(updated_kv_metadata) # Check budget_duration and budget_reset_at _set_budget_reset_at(data, updated_kv) @@ -1908,7 +1926,7 @@ async def update_team( updated_kv["router_settings"] = safe_dumps(updated_kv["router_settings"]) updated_kv = prisma_client.jsonify_team_object(db_data=updated_kv) - team_row: Optional[LiteLLM_TeamTable] = await TeamRepository(prisma_client).table.update( + team_row: LiteLLM_TeamTable | None = await TeamRepository(prisma_client).table.update( where={"team_id": data.team_id}, data=updated_kv, # `object_permission` is included so `_refresh_cached_team` @@ -2004,14 +2022,19 @@ async def patch_team( patch_fields = data.model_dump(exclude_unset=True, exclude={"team_id"}) if "metadata" in patch_fields: - existing_team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) + existing_team_row: LiteLLM_TeamTable | None = await TeamRepository(prisma_client).table.find_unique( + where={"team_id": team_id} + ) if existing_team_row is None: raise HTTPException( status_code=404, detail={"error": f"Team not found, passed team_id={team_id}"}, ) - existing_metadata = existing_team_row.metadata if isinstance(existing_team_row.metadata, dict) else {} - patch_fields["metadata"] = apply_json_merge_patch(existing_metadata, patch_fields["metadata"]) + existing_metadata: JsonValue = ( + existing_team_row.metadata if isinstance(existing_team_row.metadata, dict) else {} + ) + metadata_patch: JsonValue = patch_fields["metadata"] + patch_fields["metadata"] = apply_json_merge_patch(existing_metadata, metadata_patch) update_request = UpdateTeamRequest.model_validate({"team_id": team_id, **patch_fields}) @@ -2259,7 +2282,7 @@ async def _process_team_members( # Resolve allowed_models: explicit request value, or fall back to team's default_team_member_models member_allowed_models = data.allowed_models - team_default_member_models = getattr(complete_team_data, "default_team_member_models", None) + team_default_member_models = complete_team_data.default_team_member_models if member_allowed_models is None and team_default_member_models: member_allowed_models = team_default_member_models @@ -2470,10 +2493,10 @@ async def _validate_and_populate_member_user_info( ) # Get the single user - user_by_email = users_by_email[0] + matched_user_by_email = users_by_email[0] # Verify the user_id matches - if user_by_email.user_id != member.user_id: + if matched_user_by_email.user_id != member.user_id: raise HTTPException( status_code=400, detail={ @@ -2486,7 +2509,7 @@ async def _validate_and_populate_member_user_info( # Case 2: Only user_email provided - populate user_id from DB if member.user_email is not None and member.user_id is None: - user_by_email = await UserRepository(prisma_client).table.find_first( + user_by_email: LiteLLM_UserTable | None = await UserRepository(prisma_client).table.find_first( where={"user_email": {"equals": member.user_email, "mode": "insensitive"}} ) @@ -2515,7 +2538,9 @@ async def _validate_and_populate_member_user_info( # Case 3: Only user_id provided - populate user_email from DB if user exists if member.user_id is not None and member.user_email is None: - user_by_id = await UserRepository(prisma_client).table.find_unique(where={"user_id": member.user_id}) + user_by_id: LiteLLM_UserTable | None = await UserRepository(prisma_client).table.find_unique( + where={"user_id": member.user_id} + ) if user_by_id is None: # User doesn't exist yet - allow it to pass with user_email as None @@ -2706,7 +2731,9 @@ async def team_member_delete( detail={"error": "Either user_id or user_email needs to be passed in"}, ) - _existing_team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": data.team_id}) + _existing_team_row: LiteLLM_TeamTable | None = await TeamRepository(prisma_client).table.find_unique( + where={"team_id": data.team_id} + ) if _existing_team_row is None: raise HTTPException( @@ -2760,7 +2787,7 @@ async def team_member_delete( key_val["user_id"] = data.user_id elif data.user_email is not None: key_val["user_email"] = data.user_email - existing_user_rows = await UserRepository(prisma_client).table.find_many( + existing_user_rows: Sequence[LiteLLM_UserTable] | None = await UserRepository(prisma_client).table.find_many( where=key_val # type: ignore ) @@ -2783,7 +2810,7 @@ async def team_member_delete( user_ids_to_delete.add(data.user_id) if existing_user_rows is not None and isinstance(existing_user_rows, list): for existing_user in existing_user_rows: - if getattr(existing_user, "user_id", None): + if existing_user.user_id: user_ids_to_delete.add(existing_user.user_id) for _uid in user_ids_to_delete: @@ -2834,7 +2861,9 @@ _MEMBER_BUDGET_PATCH_FIELDS = { } -def _build_member_budget_patch(data: TeamMemberUpdateRequest) -> Dict[str, Any]: +def _build_member_budget_patch( + data: TeamMemberUpdateRequest, +) -> Dict[str, Any]: # any-ok: heterogeneous patch values; consumer _upsert_budget_and_membership is Dict[str, Any] """Map the budget fields the request actually set (merge-patch: a sent value updates, an explicit null clears, an absent field is left untouched) to their budget-table columns.""" @@ -2910,7 +2939,9 @@ async def team_member_update( _validate_budget_duration(data.budget_duration) - _existing_team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": data.team_id}) + _existing_team_row: LiteLLM_TeamTable | None = await TeamRepository(prisma_client).table.find_unique( + where={"team_id": data.team_id} + ) if _existing_team_row is None: raise HTTPException( @@ -2977,7 +3008,7 @@ async def team_member_update( ### upsert new budget budget_patch = _build_member_budget_patch(data) - async with prisma_client.db.tx() as tx: + async with prisma_client.tx() as tx: await _upsert_budget_and_membership( tx=tx, team_id=data.team_id, @@ -3137,7 +3168,9 @@ async def bulk_team_member_add( }, ) # get all users from the database - all_users_in_db = await UserRepository(prisma_client).table.find_many(order={"created_at": "desc"}) + all_users_in_db: Sequence[LiteLLM_UserTable] = await UserRepository(prisma_client).table.find_many( + order={"created_at": "desc"} + ) data.members = [ Member( user_id=user.user_id, @@ -3378,7 +3411,7 @@ def _transform_teams_to_deleted_records( teams: List[LiteLLM_TeamTable], user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str] = None, -) -> List[Dict[str, Any]]: +) -> List[Dict[str, Any]]: # any-ok: per-field JSON-serialized Prisma create_many payload; values are heterogeneous """Transform teams into deleted team records ready for persistence.""" if not teams: return [] @@ -3423,7 +3456,7 @@ def _transform_teams_to_deleted_records( async def _save_deleted_team_records( - records: List[Dict[str, Any]], + records: List[Dict[str, Any]], # any-ok: heterogeneous JSON-serialized Prisma create_many payload prisma_client: PrismaClient, ) -> None: """Save deleted team records to the database.""" @@ -3505,7 +3538,7 @@ async def _add_team_member_budget_table( team_info_response_object: TeamInfoResponseObjectTeamTable, ) -> TeamInfoResponseObjectTeamTable: try: - team_budget = await BudgetRepository(prisma_client).table.find_unique( + team_budget: LiteLLM_BudgetTableFull | None = await BudgetRepository(prisma_client).table.find_unique( where={"budget_id": team_member_budget_id} ) team_info_response_object.team_member_budget_table = team_budget @@ -3517,7 +3550,7 @@ async def _add_team_member_budget_table( return team_info_response_object -async def _resolve_team_access_group_resources(_team_info: Any) -> None: +async def _resolve_team_access_group_resources(_team_info: TeamInfoResponseObjectTeamTable) -> None: """Populate access_group_models / mcp_server_ids / agent_ids on the team info response by resolving inherited resources from its access groups.""" if not _team_info.access_group_ids: @@ -3818,7 +3851,9 @@ async def block_team( if prisma_client is None: raise Exception("No DB Connected.") - existing_team = await TeamRepository(prisma_client).table.find_unique(where={"team_id": data.team_id}) + existing_team: LiteLLM_TeamTable | None = await TeamRepository(prisma_client).table.find_unique( + where={"team_id": data.team_id} + ) if existing_team is None: raise HTTPException( status_code=404, @@ -3831,7 +3866,7 @@ async def block_team( user_api_key_dict=user_api_key_dict, ) - record = await TeamRepository(prisma_client).table.update( + record: LiteLLM_TeamTable | None = await TeamRepository(prisma_client).table.update( where={"team_id": data.team_id}, data={"blocked": True}, # type: ignore ) @@ -3867,7 +3902,9 @@ async def unblock_team( if prisma_client is None: raise Exception("No DB Connected.") - existing_team = await TeamRepository(prisma_client).table.find_unique(where={"team_id": data.team_id}) + existing_team: LiteLLM_TeamTable | None = await TeamRepository(prisma_client).table.find_unique( + where={"team_id": data.team_id} + ) if existing_team is None: raise HTTPException( status_code=404, @@ -3880,7 +3917,7 @@ async def unblock_team( user_api_key_dict=user_api_key_dict, ) - record = await TeamRepository(prisma_client).table.update( + record: LiteLLM_TeamTable | None = await TeamRepository(prisma_client).table.update( where={"team_id": data.team_id}, data={"blocked": False}, # type: ignore ) @@ -3914,7 +3951,9 @@ async def list_available_teams( return [] # filter out teams that the user is already a member of - user_info = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_api_key_dict.user_id}) + user_info: LiteLLM_UserTable | None = await UserRepository(prisma_client).table.find_unique( + where={"user_id": user_api_key_dict.user_id} + ) if user_info is None: raise HTTPException( status_code=404, @@ -3924,7 +3963,9 @@ async def list_available_teams( available_teams = [team for team in available_teams if team not in user_info_correct_type.teams] - available_teams_db = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": available_teams}}) + available_teams_db: Sequence[LiteLLM_TeamTable] = await TeamRepository(prisma_client).table.find_many( + where={"team_id": {"in": available_teams}} + ) available_teams_correct_type = [LiteLLM_TeamTable.model_validate(team.model_dump()) for team in available_teams_db] @@ -3933,9 +3974,9 @@ async def list_available_teams( async def _get_org_admin_org_ids( user_id: str, - prisma_client: Any, - user_api_key_cache: Any, - proxy_logging_obj: Any, + prisma_client: PrismaClient, + user_api_key_cache: UserApiKeyCache, + proxy_logging_obj: ProxyLogging | None, ) -> Optional[List[str]]: """ Return the list of organization IDs where the user is an org admin. @@ -3972,18 +4013,18 @@ async def _build_team_list_where_conditions( organization_id: Optional[str], user_id: Optional[str], use_deleted_table: bool, + user_api_key_cache: UserApiKeyCache, + proxy_logging_obj: ProxyLogging | None, search: Optional[str] = None, org_admin_org_ids: Optional[List[str]] = None, - user_api_key_cache: Optional[Any] = None, - proxy_logging_obj: Optional[Any] = None, -) -> Optional[Dict[str, Any]]: +) -> Dict[str, Any] | None: # any-ok: heterogeneous Prisma where-clause built incrementally across branches """ Build where conditions for team list query. Returns None when the query is guaranteed to yield no results (e.g. user has no team memberships), allowing the caller to skip the DB round-trip. """ - where_conditions: Dict[str, Any] = {} + where_conditions: Dict[str, Any] = {} # any-ok: same heterogeneous Prisma where-clause shape as the return type if team_id: where_conditions["team_id"] = team_id @@ -4011,7 +4052,7 @@ async def _build_team_list_where_conditions( user_object_correct_type = await get_user_object( user_id=user_id, prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, # type: ignore[arg-type] + user_api_key_cache=user_api_key_cache, user_id_upsert=False, proxy_logging_obj=proxy_logging_obj, ) @@ -4065,7 +4106,7 @@ async def _batch_resolve_access_group_resources( return {} unique_ids = list(set(all_access_group_ids)) - rows = await AccessGroupRepository(_prisma_client).table.find_many( + rows: Sequence[LiteLLM_AccessGroupTable] = await AccessGroupRepository(_prisma_client).table.find_many( where={"access_group_id": {"in": unique_ids}}, ) @@ -4112,9 +4153,18 @@ def _convert_teams_to_response_models( return team_list +def _extract_team_id_count(row: Mapping[str, object]) -> int: + count_obj = row.get("_count") + if isinstance(count_obj, Mapping): + count = count_obj.get("team_id", 0) + if isinstance(count, int): + return count + return 0 + + async def _get_keys_count_by_team( - prisma_client: Any, - teams: list, + prisma_client: PrismaClient, + teams: Sequence[LiteLLM_TeamTable | LiteLLM_DeletedTeamTable], ) -> Dict[str, int]: """Aggregate virtual-key counts per team for the given page of teams. @@ -4122,25 +4172,25 @@ async def _get_keys_count_by_team( bounded by page_size and uses the existing @@index([team_id]), so this is one DB round-trip per page. Returns an empty map when the page has no teams. """ - page_team_ids = [getattr(t, "team_id", None) for t in teams if getattr(t, "team_id", None)] + page_team_ids = [t.team_id for t in teams if t.team_id] if not page_team_ids: return {} - grouped = await VerificationTokenRepository(prisma_client).table.group_by( + grouped: Sequence[Mapping[str, object]] = await VerificationTokenRepository(prisma_client).table.group_by( by=["team_id"], where={"team_id": {"in": page_team_ids}}, count={"team_id": True}, ) - return {row["team_id"]: row.get("_count", {}).get("team_id", 0) for row in grouped if row.get("team_id")} + return {str(row["team_id"]): _extract_team_id_count(row) for row in grouped if row.get("team_id")} async def _enforce_list_team_v2_access( user_api_key_dict: UserAPIKeyAuth, user_id: Optional[str], organization_id: Optional[str], - prisma_client: Any, - user_api_key_cache: Any, - proxy_logging_obj: Any, + prisma_client: PrismaClient, + user_api_key_cache: UserApiKeyCache, + proxy_logging_obj: ProxyLogging | None, ) -> Tuple[Optional[str], Optional[List[str]]]: """Enforce access control for list_team_v2. @@ -4332,14 +4382,16 @@ async def list_team_v2( # Get teams with pagination if use_deleted_table: - teams = await DeletedTeamRepository(prisma_client).table.find_many( + teams: Sequence[LiteLLM_TeamTable | LiteLLM_DeletedTeamTable] = await DeletedTeamRepository( + prisma_client + ).table.find_many( where=where_conditions, skip=skip, take=page_size, order=order_by if order_by else {"created_at": "desc"}, # Default sort ) # Get total count for pagination - total_count = await DeletedTeamRepository(prisma_client).table.count(where=where_conditions) + total_count: int = await DeletedTeamRepository(prisma_client).table.count(where=where_conditions) else: teams = await TeamRepository(prisma_client).table.find_many( where=where_conditions, @@ -4348,7 +4400,7 @@ async def list_team_v2( order=order_by if order_by else {"created_at": "desc"}, # Default sort ) # Get total count for pagination - total_count = await TeamRepository(prisma_client).table.count(where=where_conditions) + total_count = await TeamRepository(prisma_client).count(where=where_conditions) # Calculate total pages total_pages = -(-total_count // page_size) # Ceiling division @@ -4360,7 +4412,7 @@ async def list_team_v2( keys_count_by_team = await _get_keys_count_by_team(prisma_client, teams) # Convert Prisma models to response models with members_count and keys_count - team_list = _convert_teams_to_response_models(teams, use_deleted_table, keys_count_by_team=keys_count_by_team) + team_list = _convert_teams_to_response_models(list(teams), use_deleted_table, keys_count_by_team=keys_count_by_team) # Resolve resources inherited from access groups (single batch query) if not use_deleted_table: @@ -4388,12 +4440,22 @@ async def list_team_v2( } +class _RawTeamRow(Protocol): + """Shape of an un-validated Prisma team row: JSON columns (e.g. + ``members_with_roles``) decode to raw dicts, not our pydantic models.""" + + team_id: str + members_with_roles: Sequence[Mapping[str, object]] | None + + def model_dump(self) -> Mapping[str, object]: ... + + async def _authorize_and_filter_teams( user_api_key_dict: UserAPIKeyAuth, user_id: Optional[str], - prisma_client: Any, - user_api_key_cache: Any, - proxy_logging_obj: Any, + prisma_client: PrismaClient, + user_api_key_cache: UserApiKeyCache, + proxy_logging_obj: ProxyLogging | None, ) -> list: """ Authorize the /team/list request and return filtered teams. @@ -4441,7 +4503,7 @@ async def _authorize_and_filter_teams( if allowed_org_ids is not None: # Org admin: query DB for teams in their orgs - org_teams = await TeamRepository(prisma_client).table.find_many( + org_teams: Sequence[_RawTeamRow] = await TeamRepository(prisma_client).table.find_many( where={"organization_id": {"in": allowed_org_ids}}, include={"litellm_model_table": True}, ) @@ -4455,7 +4517,9 @@ async def _authorize_and_filter_teams( ] elif user_id: # Regular user: fetch all and filter by membership (Prisma can't filter JSON arrays) - response = await TeamRepository(prisma_client).table.find_many(include={"litellm_model_table": True}) + response: Sequence[_RawTeamRow] = await TeamRepository(prisma_client).table.find_many( + include={"litellm_model_table": True} + ) return [ team for team in response @@ -4517,7 +4581,9 @@ async def list_team( _team_memberships.append(tm) # add all keys that belong to the team - keys = await VerificationTokenRepository(prisma_client).table.find_many(where={"team_id": team.team_id}) + keys: List[LiteLLM_VerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many( + where={"team_id": team.team_id} + ) try: returned_responses.append( @@ -4565,10 +4631,10 @@ async def get_paginated_teams( # Calculate skip for pagination skip = (page - 1) * page_size # Get total count - total_count = await TeamRepository(prisma_client).table.count() + total_count = await TeamRepository(prisma_client).count() # Get paginated teams - teams = await TeamRepository(prisma_client).table.find_many( + teams: List[LiteLLM_TeamTable] = await TeamRepository(prisma_client).table.find_many( skip=skip, take=page_size, order={"team_alias": "asc"}, # Sort by team_alias @@ -4633,7 +4699,7 @@ async def ui_view_teams( } # Query users with pagination and filters - teams = await TeamRepository(prisma_client).table.find_many( + teams: Sequence[LiteLLM_TeamTable] = await TeamRepository(prisma_client).table.find_many( where=where_conditions, skip=skip, take=page_size, @@ -4701,7 +4767,9 @@ async def team_model_add( raise HTTPException(status_code=500, detail={"error": "No db connected"}) # Get existing team - team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": data.team_id}) + team_row: LiteLLM_TeamTable | None = await TeamRepository(prisma_client).table.find_unique( + where={"team_id": data.team_id} + ) if team_row is None: raise HTTPException( @@ -4747,7 +4815,7 @@ async def team_model_add( # the writer and lets Prisma bump updated_at. # `include` mirrors the relations the auth path consumes off the cached # team object so that `_refresh_cached_team` doesn't null them out. - updated_team = await TeamRepository(prisma_client).table.update( + updated_team: LiteLLM_TeamTable = await TeamRepository(prisma_client).table.update( where={"team_id": data.team_id}, data={"updated_at": datetime.now(timezone.utc)}, include={"object_permission": True}, # type: ignore @@ -4801,7 +4869,9 @@ async def team_model_delete( raise HTTPException(status_code=500, detail={"error": "No db connected"}) # Get existing team - team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": data.team_id}) + team_row: LiteLLM_TeamTable | None = await TeamRepository(prisma_client).table.find_unique( + where={"team_id": data.team_id} + ) if team_row is None: raise HTTPException( @@ -4829,7 +4899,7 @@ async def team_model_delete( updated_models = [m for m in current_models if m not in data.models] # Update team. See team_model_add for the rationale on `include`. - updated_team = await TeamRepository(prisma_client).table.update( + updated_team: LiteLLM_TeamTable = await TeamRepository(prisma_client).table.update( where={"team_id": data.team_id}, data={"models": updated_models}, include={"object_permission": True}, # type: ignore @@ -4964,7 +5034,7 @@ async def update_team_member_permissions( }, ) # Update the team member permissions - updated_team = await TeamRepository(prisma_client).table.update( + updated_team: LiteLLM_TeamTable = await TeamRepository(prisma_client).table.update( where={"team_id": data.team_id}, data={"team_member_permissions": data.team_member_permissions}, ) @@ -5034,15 +5104,17 @@ async def bulk_update_team_member_permissions( } -async def _compute_and_batch_updates(prisma_client, teams, permissions_to_add: set) -> int: +async def _compute_and_batch_updates( + prisma_client: PrismaClient, + teams: Sequence[LiteLLM_TeamTable], + permissions_to_add: AbstractSet[str], +) -> int: """Compute merged permissions and batch-write updates. Returns count of teams updated.""" - updates = [] - for team in teams: - existing = set(team.team_member_permissions or []) - if permissions_to_add <= existing: - continue - merged = sorted(existing | permissions_to_add) # normalise to alphabetical order - updates.append((team.team_id, merged)) + updates = [ + (team.team_id, sorted(set(team.team_member_permissions or []) | permissions_to_add)) + for team in teams + if not permissions_to_add <= set(team.team_member_permissions or []) + ] if updates: batcher = prisma_client.db.batch_() @@ -5056,9 +5128,13 @@ async def _compute_and_batch_updates(prisma_client, teams, permissions_to_add: s return len(updates) -async def _append_permissions_to_specific_teams(prisma_client, team_ids: List[str], permissions_to_add: set) -> int: +async def _append_permissions_to_specific_teams( + prisma_client: PrismaClient, + team_ids: List[str], + permissions_to_add: AbstractSet[str], +) -> int: """Fetch specific teams by ID and append permissions.""" - teams = await TeamRepository(prisma_client).table.find_many( + teams: Sequence[LiteLLM_TeamTable] = await TeamRepository(prisma_client).table.find_many( where={"team_id": {"in": team_ids}}, ) @@ -5073,22 +5149,20 @@ async def _append_permissions_to_specific_teams(prisma_client, team_ids: List[st return await _compute_and_batch_updates(prisma_client, teams, permissions_to_add) -async def _append_permissions_to_all_teams(prisma_client, permissions_to_add: set) -> int: +async def _append_permissions_to_all_teams(prisma_client: PrismaClient, permissions_to_add: AbstractSet[str]) -> int: """Paginated read + batched write across all teams.""" teams_updated = 0 - cursor = None + cursor: str | None = None BATCH_SIZE = 500 while True: - find_args: dict = { - "take": BATCH_SIZE, - "order": {"team_id": "asc"}, - } - if cursor is not None: - find_args["cursor"] = {"team_id": cursor} - find_args["skip"] = 1 + find_args: Mapping[str, object] = ( + {"take": BATCH_SIZE, "order": {"team_id": "asc"}, "cursor": {"team_id": cursor}, "skip": 1} + if cursor is not None + else {"take": BATCH_SIZE, "order": {"team_id": "asc"}} + ) - teams = await TeamRepository(prisma_client).table.find_many(**find_args) + teams: Sequence[LiteLLM_TeamTable] = await TeamRepository(prisma_client).table.find_many(**find_args) if not teams: break @@ -5188,8 +5262,12 @@ async def get_team_daily_activity( where_condition = {} if team_ids_list: where_condition["team_id"] = {"in": list(team_ids_list)} - team_aliases = await TeamRepository(prisma_client).table.find_many(where=where_condition) - team_alias_metadata = {t.team_id: {"team_alias": t.team_alias} for t in team_aliases} + team_aliases: Sequence[LiteLLM_TeamTable] = await TeamRepository(prisma_client).table.find_many( + where=where_condition + ) + team_alias_metadata: Mapping[str, Dict[str, object]] = { + t.team_id: {"team_alias": t.team_alias} for t in team_aliases + } # Check if user is team admin or has /team/daily/activity permission # If not, filter by user's API keys. @@ -5219,9 +5297,9 @@ async def get_team_daily_activity( # If user does not have full team view, filter by their API keys if not has_full_team_view: # Get all API keys for this user - user_keys = await VerificationTokenRepository(prisma_client).table.find_many( - where={"user_id": user_api_key_dict.user_id} - ) + user_keys: Sequence[LiteLLM_VerificationToken] = await VerificationTokenRepository( + prisma_client + ).table.find_many(where={"user_id": user_api_key_dict.user_id}) user_api_keys = [key.token for key in user_keys if key.token] # If user has no API keys, return empty result if not user_api_keys: