From 44e091aedb3f7877e1b058c0480963621d4d6e0f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 09:16:18 +0000 Subject: [PATCH 1/2] chore(typing): clear basedpyright Any errors in proxy management endpoints Convert pydantic table-model construction from Cls(**row.model_dump()) kwargs-unpacking to Cls.model_validate(...) across the management endpoint hotspot files (team, key, internal user, scim, model management, spend tracking, auth checks, proxy_server). Unpacking an untyped dict reports one Any-typed argument per matched model field, so each converted site clears 10-35 diagnostics while running the exact same pydantic validation. Conversions were limited to models verified to use pydantic's default __init__; UserAPIKeyAuth and LiteLLM_VerificationTokenView keep their custom kwargs-rewriting __init__ and are untouched. Two locally-verified helper params move from Any to object. Whole-tree basedpyright, measured against the branch point in the same environment: reportAny 24,431 -> 22,741 (-1,690), reportArgumentType 2,189 -> 2,136 (-53), reportUnknownArgumentType 34,370 -> 34,067 (-303), reportExplicitAny 7,285 -> 7,283 (-2); total 154,882 -> 152,834 (-2,048) with no rule increasing anywhere and no per-file increases. No casts, no suppressions, no behavior changes. Budgets ratcheted: basedpyright -2,048 across 4 rules, ruff ANN401 -2. --- basedpyright-code-budget.json | 8 +- litellm/proxy/auth/auth_checks.py | 34 ++++----- .../internal_user_endpoints.py | 24 +++--- .../key_management_endpoints.py | 36 +++++---- .../model_management_endpoints.py | 4 +- .../management_endpoints/scim/scim_v2.py | 8 +- .../management_endpoints/team_endpoints.py | 76 ++++++++++--------- litellm/proxy/proxy_server.py | 12 ++- .../spend_management_endpoints.py | 4 +- ruff-strict-budget.json | 2 +- tests/test_litellm/proxy/test_proxy_server.py | 24 +++--- 11 files changed, 125 insertions(+), 107 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 28602fc235f..db3c2502e94 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,9 +1,9 @@ { "reportAny": { - "limit": 34906 + "limit": 33216 }, "reportArgumentType": { - "limit": 2701 + "limit": 2648 }, "reportAssignmentType": { "limit": 330 @@ -24,7 +24,7 @@ "limit": 42 }, "reportExplicitAny": { - "limit": 10230 + "limit": 10228 }, "reportFunctionMemberAccess": { "limit": 11 @@ -99,7 +99,7 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 45870 + "limit": 45567 }, "reportUnknownLambdaType": { "limit": 113 diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 07dfdc4fb43..d02fc02d4bf 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -953,7 +953,7 @@ async def get_default_end_user_budget( ) return None - _budget_obj = LiteLLM_BudgetTable(**budget_record.dict()) + _budget_obj = LiteLLM_BudgetTable.model_validate(budget_record.dict()) # Cache the budget for 60 seconds await user_api_key_cache.async_set_cache( key=cache_key, @@ -999,7 +999,7 @@ async def get_team_member_default_budget( if isinstance(cached_budget, LiteLLM_BudgetTable): return cached_budget if isinstance(cached_budget, dict): - return LiteLLM_BudgetTable(**cached_budget) + return LiteLLM_BudgetTable.model_validate(cached_budget) try: budget_record = await BudgetRepository(prisma_client).table.find_unique(where={"budget_id": budget_id}) @@ -1014,7 +1014,7 @@ async def get_team_member_default_budget( ttl=get_management_object_ttl(user_api_key_cache), ) - return LiteLLM_BudgetTable(**budget_record.dict()) + return LiteLLM_BudgetTable.model_validate(budget_record.dict()) except Exception: verbose_proxy_logger.exception(f"Error fetching team-default member budget {budget_id}") @@ -1168,7 +1168,7 @@ async def get_end_user_object( raise Exception # Convert to LiteLLM_EndUserTable object - _response = LiteLLM_EndUserTable(**response.dict()) + _response = LiteLLM_EndUserTable.model_validate(response.dict()) # Apply default budget if needed _response = await _apply_default_budget_to_end_user( @@ -1360,7 +1360,7 @@ async def get_tag_objects_batch( for db_tag in db_tags: tag_name = db_tag.tag_name cache_key = f"tag:{tag_name}" - _tag_obj = LiteLLM_TagTable(**db_tag.dict()) + _tag_obj = LiteLLM_TagTable.model_validate(db_tag.dict()) await user_api_key_cache.async_set_cache( key=cache_key, value=_tag_obj, @@ -1453,7 +1453,7 @@ async def get_team_membership( if response is None: return None - _response = LiteLLM_TeamMembership(**response.dict()) + _response = LiteLLM_TeamMembership.model_validate(response.dict()) await user_api_key_cache.async_set_cache( key=_key, value=_response, @@ -1719,13 +1719,13 @@ async def get_user_object( if response.organization_memberships is not None and len(response.organization_memberships) > 0: # dump each organization membership to type LiteLLM_OrganizationMembershipTable _dumped_memberships = [ - LiteLLM_OrganizationMembershipTable(**membership.model_dump()) + LiteLLM_OrganizationMembershipTable.model_validate(membership.model_dump()) for membership in response.organization_memberships if membership is not None ] response.organization_memberships = _dumped_memberships - _response = LiteLLM_UserTable(**dict(response)) + _response = LiteLLM_UserTable.model_validate(dict(response)) response_dict = _response.model_dump() # save the user object to cache @@ -1862,7 +1862,7 @@ async def _get_team_db_check(team_id: str, prisma_client: PrismaClient, team_id_ http_request=mock_request, user_api_key_dict=system_admin_user, ) - response = LiteLLM_TeamTable(**created_team_dict) + response = LiteLLM_TeamTable.model_validate(created_team_dict) return response @@ -1894,7 +1894,7 @@ async def _get_team_object_from_user_api_key_cache( if response is None: raise Exception - _response = LiteLLM_TeamTableCachedObj(**response.dict()) + _response = LiteLLM_TeamTableCachedObj.model_validate(response.dict()) # Load object_permission if object_permission_id exists but object_permission is not loaded if _response.object_permission_id and not _response.object_permission: @@ -2085,7 +2085,7 @@ async def get_access_object( detail={"error": f"Access group doesn't exist in db. Access group={access_group_id}."}, ) - _response = LiteLLM_AccessGroupTable(**response.dict()) + _response = LiteLLM_AccessGroupTable.model_validate(response.dict()) # Save to cache await _cache_access_object( @@ -2170,7 +2170,7 @@ async def get_team_object_by_alias( ) team = teams[0] - team_obj = LiteLLM_TeamTableCachedObj(**team.model_dump()) + team_obj = LiteLLM_TeamTableCachedObj.model_validate(team.model_dump()) # Load object_permission if object_permission_id exists but object_permission is not loaded if team_obj.object_permission_id and not team_obj.object_permission: @@ -2272,7 +2272,7 @@ async def get_org_object_by_alias( ) org = orgs[0] - org_obj = LiteLLM_OrganizationTable(**org.model_dump()) + org_obj = LiteLLM_OrganizationTable.model_validate(org.model_dump()) # Cache the result await user_api_key_cache.async_set_cache( @@ -2605,7 +2605,7 @@ async def get_object_permission( if response is None: return None - _perm_obj = LiteLLM_ObjectPermissionTable(**response.dict()) + _perm_obj = LiteLLM_ObjectPermissionTable.model_validate(response.dict()) await user_api_key_cache.async_set_cache( key=key, value=_perm_obj, @@ -2665,7 +2665,7 @@ async def get_managed_vector_store_rows_by_uuids( row_dict = dict(row) if hasattr(row, "__dict__") else {} if not row_dict: continue - cached_obj = LiteLLM_ManagedVectorStoresTable(**row_dict) + cached_obj = LiteLLM_ManagedVectorStoresTable.model_validate(row_dict) key = "managed_vector_store_id:{}".format(cached_obj.vector_store_id) await user_api_key_cache.async_set_cache( key=key, @@ -2746,7 +2746,7 @@ async def get_org_object( f"Organization doesn't exist in db. Organization={org_id}. Create organization via `/organization/new` call." ) - _org_obj = LiteLLM_OrganizationTable(**response.model_dump()) + _org_obj = LiteLLM_OrganizationTable.model_validate(response.model_dump()) # Cache the result await user_api_key_cache.async_set_cache( key=cache_key, @@ -4221,7 +4221,7 @@ async def get_project_object( if project_row is None: return None - project_obj = LiteLLM_ProjectTableCachedObj(**project_row.model_dump()) + project_obj = LiteLLM_ProjectTableCachedObj.model_validate(project_row.model_dump()) # Cache with TTL following _cache_management_object pattern project_obj.last_refreshed_at = time.time() diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index a2c16e88839..1bd0a19bfb3 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -513,7 +513,7 @@ async def new_user( response_dict["key"] = response.get("token", "") - new_user_response = NewUserResponse(**response_dict) + new_user_response = NewUserResponse.model_validate(response_dict) ######################################################### ########## USER CREATED HOOK ################ @@ -879,7 +879,7 @@ async def _check_user_info_v2_access( # Get all teams the caller belongs to teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": caller_user.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): # Check if target user is in this team if team.team_id in (target_user.teams or []): @@ -1013,11 +1013,11 @@ async def _get_user_info_for_proxy_admin(user_api_key_dict: UserAPIKeyAuth): for key in _keys_in_db: if key.get("models") is None: key["models"] = [] - keys_in_db.append(LiteLLM_VerificationToken(**key)) + keys_in_db.append(LiteLLM_VerificationToken.model_validate(key)) # cast all teams to LiteLLM_TeamTable _teams_in_db: list = results[0]["teams"] or [] - _teams_in_db = [LiteLLM_TeamTable(**team) for team in _teams_in_db] + _teams_in_db = [LiteLLM_TeamTable.model_validate(team) for team in _teams_in_db] _teams_in_db.sort(key=lambda x: getattr(x, "team_alias", "") or "") returned_keys = _process_keys_for_user_info(keys=keys_in_db, all_teams=_teams_in_db) @@ -1146,7 +1146,7 @@ async def _schedule_user_update_audit_log( try: updated_user_row = await UserRepository(prisma_client).table.find_first(where={"user_id": response["user_id"]}) if updated_user_row: - user_row_typed = LiteLLM_UserTable(**updated_user_row.model_dump(exclude_none=True)) + user_row_typed = LiteLLM_UserTable.model_validate(updated_user_row.model_dump(exclude_none=True)) asyncio.create_task( UserManagementEventHooks.create_internal_user_audit_log( user_id=user_row_typed.user_id, @@ -1172,7 +1172,7 @@ def _check_user_update_authz( raise HTTPException(status_code=403, detail="Only proxy admins can modify user roles.") if existing_user_row is not None: - typed_row = LiteLLM_UserTable(**existing_user_row.model_dump(exclude_none=True)) + typed_row = LiteLLM_UserTable.model_validate(existing_user_row.model_dump(exclude_none=True)) if not can_user_call_user_update(user_api_key_dict=user_api_key_dict, user_info=typed_row): raise HTTPException( status_code=403, @@ -1248,7 +1248,7 @@ async def _update_single_user_helper( _check_user_update_authz(user_request, user_api_key_dict, existing_user_row) if existing_user_row is not None: - existing_user_row = LiteLLM_UserTable(**existing_user_row.model_dump(exclude_none=True)) + existing_user_row = LiteLLM_UserTable.model_validate(existing_user_row.model_dump(exclude_none=True)) # Prevent budget self-escalation (GHSA-wvg4-6222-3q4r): non-admin callers # must not be able to raise their own budget/spend fields. @@ -1998,7 +1998,11 @@ async def get_users( for user in users: user_dump = user.model_dump() user_dump["metadata"] = _redact_scim_enterprise_metadata(user_dump.get("metadata")) - user_list.append(LiteLLM_UserTableWithKeyCount(**user_dump, key_count=user_key_counts.get(user.user_id, 0))) + user_list.append( + LiteLLM_UserTableWithKeyCount.model_validate( + {**user_dump, "key_count": user_key_counts.get(user.user_id, 0)} + ) + ) else: user_list = [] @@ -2157,7 +2161,7 @@ async def delete_user( teams_to_update = [] for team in fetch_all_teams: is_member_in_team, new_team_members = _cleanup_members_with_roles( - existing_team_row=LiteLLM_TeamTable(**team.model_dump()), + existing_team_row=LiteLLM_TeamTable.model_validate(team.model_dump()), data=TeamMemberDeleteRequest( team_id=team.team_id, user_id=user_row.user_id, @@ -2438,7 +2442,7 @@ async def ui_view_users( if not users: return [] - return [LiteLLM_UserTableFiltered(**user.model_dump()) for user in users] + return [LiteLLM_UserTableFiltered.model_validate(user.model_dump()) for user in users] except HTTPException: raise diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index ac6a2a4a7db..e7ee5ffa849 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1078,7 +1078,7 @@ async def _common_key_generation_helper( response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response - response = GenerateKeyResponse(**response) + response = GenerateKeyResponse.model_validate(response) response.token = response.token_id # remap token to use the hash, and leave the key in the `key` field [TODO]: clean up generate_key_helper_fn to do this @@ -3047,10 +3047,12 @@ async def bulk_update_team_keys( ) # team_id from validated scope, never user payload — drives _check_team_key_limits. - update_key_request = UpdateKeyRequest( - key=token, - team_id=data.team_id, - **update_field_dict, + update_key_request = UpdateKeyRequest.model_validate( + { + "key": token, + "team_id": data.team_id, + **update_field_dict, + } ) updated_key_info = await _process_single_key_update( update_key_request=update_key_request, @@ -4048,12 +4050,14 @@ def _transform_verification_tokens_to_deleted_records( records = [] for key in keys: key_payload = key.model_dump() - deleted_record = LiteLLM_DeletedVerificationToken( - **key_payload, - deleted_at=deleted_at, - deleted_by=user_api_key_dict.user_id, - deleted_by_api_key=user_api_key_dict.api_key, - litellm_changed_by=litellm_changed_by, + deleted_record = LiteLLM_DeletedVerificationToken.model_validate( + { + **key_payload, + "deleted_at": deleted_at, + "deleted_by": user_api_key_dict.user_id, + "deleted_by_api_key": user_api_key_dict.api_key, + "litellm_changed_by": litellm_changed_by, + } ) record = deleted_record.model_dump() @@ -4535,7 +4539,7 @@ async def _execute_virtual_key_regeneration( proxy_logging_obj=proxy_logging_obj, ) - response = GenerateKeyResponse(**updated_token_dict) + response = GenerateKeyResponse.model_validate(updated_token_dict) asyncio.create_task( KeyManagementEventHooks.async_key_rotated_hook( data=data, @@ -4853,7 +4857,7 @@ async def _check_proxy_or_team_admin_for_key( ) -def _validate_reset_spend_value(reset_to: Any, key_in_db: LiteLLM_VerificationToken) -> float: +def _validate_reset_spend_value(reset_to: object, key_in_db: LiteLLM_VerificationToken) -> float: if not isinstance(reset_to, (int, float)): raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -5029,7 +5033,7 @@ async def validate_key_list_check( code=status.HTTP_403_FORBIDDEN, ) - complete_user_info = LiteLLM_UserTable(**complete_user_info_db_obj.model_dump()) + complete_user_info = LiteLLM_UserTable.model_validate(complete_user_info_db_obj.model_dump()) # internal user can only see their own keys if user_id: @@ -5102,7 +5106,7 @@ async def _fetch_user_team_objects( if teams is None: return [] - return [LiteLLM_TeamTable(**team.model_dump()) for team in teams] + return [LiteLLM_TeamTable.model_validate(team.model_dump()) for team in teams] def _get_admin_team_ids_from_objects( @@ -5851,7 +5855,7 @@ async def _list_key_helper( if return_full_object is True or (expand and "user" in expand): if use_deleted_table: # Use deleted key type to preserve deleted_at, deleted_by, etc. - key_list.append(LiteLLM_DeletedVerificationToken(**key_dict)) + key_list.append(LiteLLM_DeletedVerificationToken.model_validate(key_dict)) else: key_list.append(UserAPIKeyAuth(**key_dict)) # Return full key object else: diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 1c0e7211493..b6422d7f5ae 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -1051,7 +1051,7 @@ class ModelManagementAuthChecks: status_code=400, detail={"error": "Team id={} does not exist in db".format(model_params.model_info.team_id)}, ) - existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump()) + existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump()) ModelManagementAuthChecks.can_user_make_team_model_call( team_id=model_params.model_info.team_id, @@ -1089,7 +1089,7 @@ class ModelManagementAuthChecks: status_code=400, detail={"error": "Team id={} does not exist in db".format(model_params.model_info.team_id)}, ) - team_obj = LiteLLM_TeamTable(**team_obj_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_obj_row.model_dump()) return ModelManagementAuthChecks.can_user_make_team_model_call( team_id=model_params.model_info.team_id, diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 90eae5bbb21..582d34dcec8 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -1317,7 +1317,7 @@ async def delete_user( where={"team_id": team.team_id}, data={"members": new_members} ) - team_row = LiteLLM_TeamTable(**team.model_dump()) + team_row = LiteLLM_TeamTable.model_validate(team.model_dump()) if any(member.user_id == user_id for member in team_row.members_with_roles or []): await team_member_delete( data=TeamMemberDeleteRequest(team_id=team_row.team_id, user_id=user_id), @@ -2145,7 +2145,9 @@ async def patch_group( refreshed_team = await TeamRepository(prisma_client).table.find_unique(where={"team_id": group_id}) refreshed_current = ( - set(await _get_team_member_user_ids_from_team(LiteLLM_TeamTable(**refreshed_team.model_dump()))) + set( + await _get_team_member_user_ids_from_team(LiteLLM_TeamTable.model_validate(refreshed_team.model_dump())) + ) if refreshed_team else snapshot_members ) @@ -2173,7 +2175,7 @@ async def patch_group( # Convert to SCIM format and return scim_group = await ScimTransformations.transform_litellm_team_to_scim_group( - LiteLLM_TeamTable(**updated_team.model_dump()) + LiteLLM_TeamTable.model_validate(updated_team.model_dump()) ) return scim_group diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 59b0cbc4ae7..c35c17aa359 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -141,7 +141,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import ( router = APIRouter() -def _sanitize_for_log(value: Any) -> str: +def _sanitize_for_log(value: object) -> str: """Strip CR/LF from user-controlled values to prevent log injection.""" try: text = str(value) @@ -171,7 +171,7 @@ async def _refresh_cached_team( """ await _cache_team_object( team_id=team_row.team_id, - team_table=LiteLLM_TeamTableCachedObj(**team_row.model_dump()), + team_table=LiteLLM_TeamTableCachedObj.model_validate(team_row.model_dump()), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) @@ -510,7 +510,7 @@ async def get_all_team_memberships( returned_tm: List[LiteLLM_TeamMembership] = [] for tm in team_memberships: - returned_tm.append(LiteLLM_TeamMembership(**tm.model_dump())) + returned_tm.append(LiteLLM_TeamMembership.model_validate(tm.model_dump())) return returned_tm @@ -772,7 +772,7 @@ async def _check_org_team_limits( # Convert teams to LiteLLM_TeamTable objects team_objs: List[LiteLLM_TeamTable] = [] for team in teams: - team_objs.append(LiteLLM_TeamTable(**team.model_dump())) + team_objs.append(LiteLLM_TeamTable.model_validate(team.model_dump())) check_org_team_model_specific_limits( teams=team_objs, @@ -1467,9 +1467,9 @@ async def fetch_and_validate_organization( ) is_proxy_admin = user_api_key_dict is not None and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN - organization = LiteLLM_OrganizationTableWithMembers(**organization_row.model_dump()) + organization = LiteLLM_OrganizationTableWithMembers.model_validate(organization_row.model_dump()) validate_team_org_change( - team=LiteLLM_TeamTable(**existing_team_row.model_dump()), + team=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()), organization=organization, llm_router=llm_router, is_proxy_admin=is_proxy_admin, @@ -1477,7 +1477,7 @@ async def fetch_and_validate_organization( if is_proxy_admin: await _auto_add_team_members_to_organization( - team=LiteLLM_TeamTable(**existing_team_row.model_dump()), + team=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()), organization=organization, prisma_client=prisma_client, ) @@ -1714,7 +1714,7 @@ async def update_team( # Verify caller has access to manage this team await _verify_team_access( - team_obj=LiteLLM_TeamTable(**existing_team_row.model_dump()), + team_obj=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()), user_api_key_dict=user_api_key_dict, ) @@ -2013,7 +2013,7 @@ async def patch_team( 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"]) - update_request = UpdateTeamRequest(team_id=team_id, **patch_fields) + update_request = UpdateTeamRequest.model_validate({"team_id": team_id, **patch_fields}) result = await update_team( data=update_request, @@ -2591,7 +2591,7 @@ async def team_member_add( detail={"error": f"Team not found for team_id={getattr(data, 'team_id', None)}"}, ) - complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump()) + complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()) team_member_add_duplication_check( data=data, @@ -2636,10 +2636,12 @@ async def team_member_add( _emit_team_members_metric(complete_team_data) - return TeamAddMemberResponse( - **updated_team.model_dump(), - updated_users=updated_users, - updated_team_memberships=updated_team_memberships, + return TeamAddMemberResponse.model_validate( + { + **updated_team.model_dump(), + "updated_users": updated_users, + "updated_team_memberships": updated_team_memberships, + } ) @@ -2711,7 +2713,7 @@ async def team_member_delete( status_code=400, detail={"error": "Team id={} does not exist in db".format(data.team_id)}, ) - existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump()) + existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump()) ## CHECK IF USER IS PROXY ADMIN OR TEAM ADMIN OR ORG ADMIN @@ -2915,7 +2917,7 @@ async def team_member_update( status_code=400, detail={"error": "Team id={} does not exist in db".format(data.team_id)}, ) - existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump()) + existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump()) ## CHECK IF USER IS PROXY ADMIN OR TEAM ADMIN OR ORG ADMIN @@ -3261,7 +3263,7 @@ async def delete_team( status_code=404, detail={"error": f"Team not found, passed team_id={team_id}"}, ) - team_row_pydantic = LiteLLM_TeamTable(**team_row_base.model_dump()) + team_row_pydantic = LiteLLM_TeamTable.model_validate(team_row_base.model_dump()) # Verify caller has access to manage this team await _verify_team_access( @@ -3385,12 +3387,14 @@ def _transform_teams_to_deleted_records( records = [] for team in teams: team_payload = team.model_dump() - deleted_record = LiteLLM_DeletedTeamTable( - **team_payload, - deleted_at=deleted_at, - deleted_by=user_api_key_dict.user_id, - deleted_by_api_key=user_api_key_dict.api_key, - litellm_changed_by=litellm_changed_by, + deleted_record = LiteLLM_DeletedTeamTable.model_validate( + { + **team_payload, + "deleted_at": deleted_at, + "deleted_by": user_api_key_dict.user_id, + "deleted_by_api_key": user_api_key_dict.api_key, + "litellm_changed_by": litellm_changed_by, + } ) record = deleted_record.model_dump() @@ -3580,7 +3584,7 @@ async def team_info( ) await validate_membership( user_api_key_dict=user_api_key_dict, - team_table=LiteLLM_TeamTable(**team_info.model_dump()), + team_table=LiteLLM_TeamTable.model_validate(team_info.model_dump()), ) ## GET ALL KEYS ## @@ -3615,9 +3619,9 @@ async def team_info( returned_tm = await get_all_team_memberships(prisma_client, [team_id], user_id=None) if isinstance(team_info, dict): - _team_info = TeamInfoResponseObjectTeamTable(**team_info) + _team_info = TeamInfoResponseObjectTeamTable.model_validate(team_info) elif isinstance(team_info, BaseModel): - _team_info = TeamInfoResponseObjectTeamTable(**team_info.model_dump()) + _team_info = TeamInfoResponseObjectTeamTable.model_validate(team_info.model_dump()) else: _team_info = TeamInfoResponseObjectTeamTable() @@ -3823,7 +3827,7 @@ async def block_team( # Verify caller has access to manage this team await _verify_team_access( - team_obj=LiteLLM_TeamTable(**existing_team.model_dump()), + team_obj=LiteLLM_TeamTable.model_validate(existing_team.model_dump()), user_api_key_dict=user_api_key_dict, ) @@ -3872,7 +3876,7 @@ async def unblock_team( # Verify caller has access to manage this team await _verify_team_access( - team_obj=LiteLLM_TeamTable(**existing_team.model_dump()), + team_obj=LiteLLM_TeamTable.model_validate(existing_team.model_dump()), user_api_key_dict=user_api_key_dict, ) @@ -3916,13 +3920,13 @@ async def list_available_teams( status_code=404, detail={"error": "User not found"}, ) - user_info_correct_type = LiteLLM_UserTable(**user_info.model_dump()) + user_info_correct_type = LiteLLM_UserTable.model_validate(user_info.model_dump()) 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_correct_type = [LiteLLM_TeamTable(**team.model_dump()) for team in available_teams_db] + available_teams_correct_type = [LiteLLM_TeamTable.model_validate(team.model_dump()) for team in available_teams_db] return available_teams_correct_type @@ -4090,7 +4094,7 @@ def _convert_teams_to_response_models( team_dict = team.dict() if use_deleted_table: - team_list.append(LiteLLM_DeletedTeamTable(**team_dict)) + team_list.append(LiteLLM_DeletedTeamTable.model_validate(team_dict)) else: members_with_roles = team_dict.get("members_with_roles") if not isinstance(members_with_roles, list): @@ -4705,7 +4709,7 @@ async def team_model_add( detail={"error": f"Team not found, passed team_id={data.team_id}"}, ) - team_obj = LiteLLM_TeamTable(**team_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) # Authorization check - only proxy admin, team admin, or org admin can add models if ( @@ -4805,7 +4809,7 @@ async def team_model_delete( detail={"error": f"Team not found, passed team_id={data.team_id}"}, ) - team_obj = LiteLLM_TeamTable(**team_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) # Authorization check - only proxy admin, team admin, or org admin can remove models if ( @@ -4873,7 +4877,7 @@ async def team_member_permissions( check_db_only=True, ) - complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump()) + complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()) # Admin Viewer follows the read-parity rule: see team permissions like # a Proxy Admin would. Team / org admins keep their existing scope. @@ -4940,7 +4944,7 @@ async def update_team_member_permissions( check_db_only=True, ) - complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump()) + complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()) # Available-team self-join must NOT grant write access to team-wide # permission policies; only proxy/team/org admins can update them. @@ -5201,7 +5205,7 @@ async def get_team_daily_activity( if not _user_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases: has_full_team_view = True for team_alias in team_aliases: - team_obj = LiteLLM_TeamTable(**team_alias.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_alias.model_dump()) is_admin = _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj) has_perm = _team_member_has_permission( user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4486cd7de59..70484eb1e4e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -11214,11 +11214,15 @@ async def get_all_team_models( if user_teams == "*": team_db_objects = await TeamRepository(prisma_client).table.find_many() - team_db_objects_typed = [LiteLLM_TeamTable(**team_db_object.model_dump()) for team_db_object in team_db_objects] + team_db_objects_typed = [ + LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) for team_db_object in team_db_objects + ] else: team_db_objects = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_teams}}) - team_db_objects_typed = [LiteLLM_TeamTable(**team_db_object.model_dump()) for team_db_object in team_db_objects] + team_db_objects_typed = [ + LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) for team_db_object in team_db_objects + ] team_models = _add_team_models_to_all_models( team_db_objects_typed=team_db_objects_typed, @@ -11292,7 +11296,7 @@ async def _populate_team_access_on_models( where={"user_id": user_api_key_dict.user_id} ) if user_db_object is not None: - user_object = LiteLLM_UserTable(**user_db_object.model_dump()) + user_object = LiteLLM_UserTable.model_validate(user_db_object.model_dump()) user_teams = user_object.teams or [] direct_access_models = get_direct_access_models( user_db_object=user_object, @@ -11827,7 +11831,7 @@ async def _load_team_object_for_model_filter(team_id: str, prisma_client: Prisma if team_db_object is None: verbose_proxy_logger.warning(f"Team {team_id} not found in database") return None - return LiteLLM_TeamTable(**team_db_object.model_dump()) + return LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) except Exception as e: verbose_proxy_logger.exception(f"Error fetching team {team_id}: {str(e)}") return None diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 0c525ee9466..9aae9ca2875 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -3548,7 +3548,7 @@ async def _can_team_member_view_log( team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) if team_row is None: return False - team_obj = LiteLLM_TeamTable(**team_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): return True return _team_member_has_permission( @@ -3640,7 +3640,7 @@ async def _get_permitted_team_ids_for_spend_logs( permitted: List[str] = [] for team_row in team_rows: - team_obj = LiteLLM_TeamTable(**team_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): permitted.append(team_obj.team_id) elif _team_member_has_permission( diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index addee5fc68a..f3b4fce97d3 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -24,7 +24,7 @@ "limit": 130 }, "ANN401": { - "limit": 2015 + "limit": 2013 }, "ASYNC230": { "limit": 14 diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 536d24d4b4e..62f5ced7a39 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -1571,14 +1571,14 @@ async def test_get_all_team_models(): with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class: # Configure the mock class to return proper instances - def mock_team_table_constructor(**kwargs): + def mock_team_table_constructor(data): mock_instance = MagicMock() - mock_instance.team_id = kwargs["team_id"] - mock_instance.models = kwargs["models"] - mock_instance.access_group_ids = kwargs.get("access_group_ids") + mock_instance.team_id = data["team_id"] + mock_instance.models = data["models"] + mock_instance.access_group_ids = data.get("access_group_ids") return mock_instance - mock_team_table_class.side_effect = mock_team_table_constructor + mock_team_table_class.model_validate.side_effect = mock_team_table_constructor result = await get_all_team_models( user_teams="*", @@ -1607,7 +1607,7 @@ async def test_get_all_team_models(): mock_litellm_teamtable.find_many.return_value = [mock_team1] with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class: - mock_team_table_class.side_effect = mock_team_table_constructor + mock_team_table_class.model_validate.side_effect = mock_team_table_constructor result = await get_all_team_models( user_teams=["team1"], @@ -1658,7 +1658,7 @@ async def test_get_all_team_models(): mock_router.get_model_list.side_effect = mock_get_model_list_with_none with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class: - mock_team_table_class.side_effect = mock_team_table_constructor + mock_team_table_class.model_validate.side_effect = mock_team_table_constructor result = await get_all_team_models( user_teams=["team1"], @@ -2373,14 +2373,14 @@ async def test_get_all_team_models_with_access_groups(): with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_tt_class: - def mock_team_table_constructor(**kwargs): + def mock_team_table_constructor(data): mock_instance = MagicMock() - mock_instance.team_id = kwargs["team_id"] - mock_instance.models = kwargs["models"] - mock_instance.access_group_ids = kwargs.get("access_group_ids") + mock_instance.team_id = data["team_id"] + mock_instance.models = data["models"] + mock_instance.access_group_ids = data.get("access_group_ids") return mock_instance - mock_tt_class.side_effect = mock_team_table_constructor + mock_tt_class.model_validate.side_effect = mock_team_table_constructor result = await get_all_team_models( user_teams=["team1"], From 095364fd046bc746cbb58a7808d8b249126d867f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 10:04:58 +0000 Subject: [PATCH 2/2] test: cover the model_validate conversion sites flagged by codecov Add regression tests for the db-fetch paths whose converted construction lines were uncovered: the auth_checks getters (default end user budget, end user, team membership, access group, team by alias, org by alias, object permission, managed vector stores, project), get_all_team_memberships and list_available_teams in team_endpoints, and the proxy admin user info helper. Each test feeds a mocked prisma row through the real function and asserts the validated model's fields, so a bad model_validate conversion on any of these paths now fails a test instead of only dropping coverage. --- .../proxy/auth/test_auth_checks.py | 251 ++++++++++++++++++ .../test_internal_user_endpoints.py | 35 +++ .../test_team_endpoints.py | 65 +++++ 3 files changed, 351 insertions(+) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index ccb20976df9..a5bf5e280d8 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -4762,3 +4762,254 @@ async def test_skip_user_budget_on_team_key_flag_restores_old_behavior(): request=MagicMock(spec=Request), ) assert result is True + + +@pytest.mark.asyncio +async def test_get_default_end_user_budget_db_fetch_returns_validated_budget(monkeypatch): + from litellm.proxy.auth.auth_checks import get_default_end_user_budget + + monkeypatch.setattr(litellm, "max_end_user_budget_id", "budget-default-1") + + budget_row = MagicMock() + budget_row.dict = lambda: {"budget_id": "budget-default-1", "max_budget": 12.5, "tpm_limit": 100} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_default_end_user_budget( + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_BudgetTable) + assert result.max_budget == 12.5 + assert result.tpm_limit == 100 + mock_cache.async_set_cache.assert_awaited_once() + assert mock_cache.async_set_cache.call_args.kwargs["value"] is result + + +@pytest.mark.asyncio +async def test_get_end_user_object_db_fetch_returns_validated_end_user(): + from litellm.proxy.auth.auth_checks import get_end_user_object + + end_user_row = MagicMock() + end_user_row.dict = lambda: {"user_id": "eu-1", "blocked": False, "spend": 3.0} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock(return_value=end_user_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_end_user_object( + end_user_id="eu-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_EndUserTable) + assert result.user_id == "eu-1" + assert result.blocked is False + assert result.spend == 3.0 + + +@pytest.mark.asyncio +async def test_get_team_membership_db_fetch_returns_validated_membership(): + from litellm.proxy._types import LiteLLM_TeamMembership + from litellm.proxy.auth.auth_checks import get_team_membership + + membership_row = MagicMock() + membership_row.dict = lambda: {"user_id": "u-1", "team_id": "t-1", "spend": 1.5} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_team_membership( + user_id="u-1", + team_id="t-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_TeamMembership) + assert result.user_id == "u-1" + assert result.team_id == "t-1" + assert result.spend == 1.5 + + +@pytest.mark.asyncio +async def test_get_access_object_db_fetch_returns_validated_access_group(): + from litellm.proxy._types import LiteLLM_AccessGroupTable + from litellm.proxy.auth.auth_checks import get_access_object + + access_row = MagicMock() + access_row.dict = lambda: { + "access_group_id": "ag-1", + "access_group_name": "group one", + "access_model_names": ["gpt-4"], + } + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=access_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_access_object( + access_group_id="ag-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + proxy_logging_obj=None, + ) + + assert isinstance(result, LiteLLM_AccessGroupTable) + assert result.access_group_id == "ag-1" + assert result.access_model_names == ["gpt-4"] + + +@pytest.mark.asyncio +async def test_get_team_object_by_alias_db_fetch_returns_cached_obj(): + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.auth.auth_checks import get_team_object_by_alias + + team_row = MagicMock() + team_row.model_dump = lambda: {"team_id": "t-9", "team_alias": "alias-9", "models": ["gpt-4"]} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row]) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_team_object_by_alias( + team_alias="alias-9", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_TeamTableCachedObj) + assert result.team_id == "t-9" + assert result.team_alias == "alias-9" + assert result.models == ["gpt-4"] + + +@pytest.mark.asyncio +async def test_get_org_object_by_alias_db_fetch_returns_validated_org(): + from litellm.proxy._types import LiteLLM_OrganizationTable + from litellm.proxy.auth.auth_checks import get_org_object_by_alias + + org_row = MagicMock() + org_row.model_dump = lambda: { + "organization_id": "org-1", + "organization_alias": "org-alias", + "budget_id": "b-1", + "created_by": "admin", + "updated_by": "admin", + "models": [], + } + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[org_row]) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_org_object_by_alias( + org_alias="org-alias", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_OrganizationTable) + assert result.organization_id == "org-1" + assert result.budget_id == "b-1" + + +@pytest.mark.asyncio +async def test_get_object_permission_db_fetch_returns_validated_permission(): + from litellm.proxy.auth.auth_checks import get_object_permission + + perm_row = MagicMock() + perm_row.dict = lambda: {"object_permission_id": "op-1", "vector_stores": ["vs-1"]} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=perm_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_object_permission( + object_permission_id="op-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_ObjectPermissionTable) + assert result.object_permission_id == "op-1" + assert result.vector_stores == ["vs-1"] + + +@pytest.mark.asyncio +async def test_get_managed_vector_store_rows_by_uuids_db_fetch_validates_rows(): + from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable + from litellm.proxy.auth.auth_checks import get_managed_vector_store_rows_by_uuids + + vs_row = MagicMock() + vs_row.model_dump = lambda: {"vector_store_id": "vs-7", "custom_llm_provider": "openai"} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=[vs_row]) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_managed_vector_store_rows_by_uuids( + uuids=["vs-7"], + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert len(result) == 1 + assert isinstance(result[0], LiteLLM_ManagedVectorStoresTable) + assert result[0].vector_store_id == "vs-7" + assert result[0].custom_llm_provider == "openai" + + +@pytest.mark.asyncio +async def test_get_project_object_db_fetch_returns_cached_obj(): + from litellm.proxy._types import LiteLLM_ProjectTableCachedObj + from litellm.proxy.auth.auth_checks import get_project_object + + project_row = MagicMock() + project_row.model_dump = lambda: {"project_id": "p-1", "project_alias": "proj", "team_id": "t-1"} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock(return_value=project_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_project_object( + project_id="p-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_ProjectTableCachedObj) + assert result.project_id == "p-1" + assert result.project_alias == "proj" diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 5cbc3e72d83..8de42ca89da 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -3667,3 +3667,38 @@ async def test_add_user_to_team_keeps_already_a_member_quiet(mocker, caplog): ) assert [r.getMessage() for r in caplog.records if r.levelno >= logging.ERROR] == [] + + +@pytest.mark.asyncio +async def test_get_user_info_for_proxy_admin_validates_keys_and_teams(): + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.proxy._types import LiteLLM_TeamTable, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _get_user_info_for_proxy_admin, + ) + + raw_rows = [ + { + "teams": [ + {"team_id": "team-b", "team_alias": "beta"}, + {"team_id": "team-a", "team_alias": "alpha"}, + ], + "keys": [ + {"token": "hashed-token-1", "team_id": "team-a", "models": None, "spend": 1.0}, + ], + } + ] + + mock_prisma_client = MagicMock() + mock_prisma_client.db.query_raw = AsyncMock(return_value=raw_rows) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): + result = await _get_user_info_for_proxy_admin(user_api_key_dict=UserAPIKeyAuth(user_id=None)) + + assert all(isinstance(team, LiteLLM_TeamTable) for team in result.teams) + assert [team.team_alias for team in result.teams] == ["alpha", "beta"] + assert len(result.keys) == 1 + returned_key = result.keys[0] + assert returned_key["team_id"] == "team-a" + assert returned_key["models"] == [] diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 5202c8cbfc0..1e4d1759062 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -10223,3 +10223,68 @@ def test_patch_team_route_publishes_its_request_body_schema(): assert schema == {"$ref": "#/components/schemas/PatchTeamRequest"} properties = app.openapi()["components"]["schemas"]["PatchTeamRequest"]["properties"] assert "tpm_limit" in properties and "metadata" in properties + + +@pytest.mark.asyncio +async def test_get_all_team_memberships_validates_rows(): + from litellm.proxy._types import LiteLLM_TeamMembership + from litellm.proxy.management_endpoints.team_endpoints import ( + get_all_team_memberships, + ) + + membership_row = MagicMock() + membership_row.model_dump = lambda: { + "user_id": "member-1", + "team_id": "team-1", + "spend": 2.5, + } + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[membership_row]) + + result = await get_all_team_memberships(mock_prisma_client, ["team-1"], user_id="member-1") + + assert len(result) == 1 + assert isinstance(result[0], LiteLLM_TeamMembership) + assert result[0].user_id == "member-1" + assert result[0].team_id == "team-1" + assert result[0].spend == 2.5 + find_many_kwargs = mock_prisma_client.db.litellm_teammembership.find_many.call_args.kwargs + assert find_many_kwargs["where"] == {"team_id": {"in": ["team-1"]}, "user_id": {"in": ["member-1"]}} + + +@pytest.mark.asyncio +async def test_list_available_teams_filters_joined_and_validates_rows(monkeypatch): + from fastapi import Request + + import litellm + from litellm.proxy.management_endpoints.team_endpoints import list_available_teams + + monkeypatch.setattr( + litellm, + "default_internal_user_params", + {"available_teams": ["team-open", "team-joined"]}, + ) + + user_row = MagicMock() + user_row.model_dump = lambda: {"user_id": "u-1", "teams": ["team-joined"]} + + open_team_row = MagicMock() + open_team_row.model_dump = lambda: {"team_id": "team-open", "team_alias": "open team"} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[open_team_row]) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): + result = await list_available_teams( + http_request=MagicMock(spec=Request), + user_api_key_dict=UserAPIKeyAuth(user_id="u-1"), + ) + + assert len(result) == 1 + assert isinstance(result[0], LiteLLM_TeamTable) + assert result[0].team_id == "team-open" + assert result[0].team_alias == "open team" + find_many_kwargs = mock_prisma_client.db.litellm_teamtable.find_many.call_args.kwargs + assert find_many_kwargs["where"] == {"team_id": {"in": ["team-open"]}}