From 46b6eae799b8ee6fb7b63d6007876aaa77971827 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Mon, 3 Aug 2026 12:57:12 -0700 Subject: [PATCH] feat(teams): apply default organization to new teams from default team settings (#35540) * feat(teams): apply default organization to new teams from default team settings Adds organization_id to DefaultTeamSSOParams so proxy admins can pick a default organization in Default Team Settings. new_team applies it before org validation whenever a team is created without an explicit organization_id, so API, Admin UI, SCIM, SSO, and team upsert creations all inherit it and go through the same existence and org-limit checks. Explicit organization selections win and existing teams are untouched. The default is validated at save time (PATCH /update/default_team_settings returns 400 for an unknown org) and at create time, where a missing org now surfaces as a clean 400 instead of a 500 by routing OrganizationNotFoundError into the previously dead org_table None guard. The Admin UI Default Team Settings tab gets a Default Organization row backed by the shared OrganizationDropdown. * fix(teams): validate org limits against final team state including defaults Applies default_team_params and the legacy max_budget fallback before the organization validation block, so _check_org_team_limits sees the values the team will actually be persisted with. Also loads the org's budget table in the lookup; without include_budget_table every budget comparison in _check_org_team_limits was skipped because litellm_budget_table was None. * test(proxy_behavior): pin org team limits as enforced on /team/new The dead-code pins existed to turn red when include_budget_table went live; that happened, so the scenarios now assert the 400 rejections plus within-cap acceptance, and the unknown-org pin asserts the handler's 400 instead of the surfaced 500. --- .../management_endpoints/team_endpoints.py | 46 ++-- .../proxy_setting_endpoints.py | 34 +++ .../proxy/management_endpoints/ui_sso.py | 4 + .../management/test_team_budget_limits.py | 79 +++--- .../management/test_team_new.py | 29 +- .../test_team_default_params.py | 259 +++++++++++++----- .../proxy/management_endpoints/test_ui_sso.py | 49 ++++ .../test_proxy_setting_endpoints.py | 88 ++++++ .../src/components/TeamSSOSettings.test.tsx | 172 +++++++++++- .../src/components/TeamSSOSettings.tsx | 35 ++- .../OrganizationDropdown.tsx | 4 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 5 + 12 files changed, 656 insertions(+), 148 deletions(-) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 54ef697d16e..4dd86e5769d 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -76,6 +76,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.auth_checks import ( + OrganizationNotFoundError, _cache_team_object, allowed_route_check_inside_route, can_org_access_model, @@ -1210,24 +1211,10 @@ async def new_team( detail={"error": f"Team id = {data.team_id} already exists. Please use a different team id."}, ) - # check org key limits - done here to handle inheriting org id from team - if data.organization_id is not None and prisma_client is not None: - org_table = await get_org_object( - org_id=data.organization_id, - user_api_key_cache=user_api_key_cache, - prisma_client=prisma_client, - ) - if org_table is None: - raise HTTPException( - status_code=400, - detail=f"Organization not found for organization_id={data.organization_id}", - ) - - await _check_org_team_limits( - org_table=org_table, - data=data, - prisma_client=prisma_client, - ) + if data.organization_id is None: + default_organization_id = _get_default_team_param("organization_id") + if isinstance(default_organization_id, str): + data.organization_id = default_organization_id # Apply defaults from litellm.default_team_params for any fields # not explicitly provided in the request. @@ -1255,6 +1242,29 @@ async def new_team( if default_budget is not None: data.max_budget = default_budget + # check org key limits - done here to handle inheriting org id from team + if data.organization_id is not None and prisma_client is not None: + try: + org_table = await get_org_object( + org_id=data.organization_id, + user_api_key_cache=user_api_key_cache, + prisma_client=prisma_client, + include_budget_table=True, + ) + except OrganizationNotFoundError: + org_table = None + if org_table is None: + raise HTTPException( + status_code=400, + detail=f"Organization not found for organization_id={data.organization_id}", + ) + + await _check_org_team_limits( + org_table=org_table, + data=data, + prisma_client=prisma_client, + ) + if ( user_api_key_dict.user_role is None or user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN ): # don't restrict proxy admin diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index 60c88c0c371..8ed848ac1bf 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -23,6 +23,7 @@ from litellm.proxy.config_resolvers.sso import ( ) from litellm.proxy.utils import invalidate_config_param from litellm.repositories.config_repository import ConfigRepository +from litellm.repositories.organization_repository import OrganizationRepository from litellm.repositories.table_repositories import ( SSOConfigRepository, UISettingsRepository, @@ -636,6 +637,36 @@ async def _validate_default_teams_exist(teams: list[str] | list[NewUserRequestTe ) +async def _validate_default_organization_exists(organization_id: str) -> None: + """Reject a default organization that cannot be assigned. + + Teams are created from these settings long after they are saved, and an unknown + organization id would fail every future team creation instead of here, where the + admin who typed it can still fix it. + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={ # mutable-ok: HTTPException detail must be a plain dict for FastAPI JSON serialization + "error": "Database not connected. Please connect a database." + }, + ) + + organization_exists = await OrganizationRepository(prisma_client).exists( + organization_id, id_field="organization_id" + ) + if not organization_exists: + raise HTTPException( + status_code=400, + detail={ # mutable-ok: HTTPException detail must be a plain dict for FastAPI JSON serialization + "error": f"Organization not found: {organization_id}. " + "An organization must exist before it can be set as the default organization for new teams." + }, + ) + + async def update_default_team_member_budget(teams: list[NewUserRequestTeam], user_api_key_dict: UserAPIKeyAuth): """ 1. Update the max member budget for the team @@ -774,6 +805,9 @@ async def update_default_team_settings( Update the default team parameters for SSO users. These settings will be applied to new teams created from SSO. """ + if settings.organization_id is not None: + await _validate_default_organization_exists(settings.organization_id) + return await _update_litellm_setting( settings=settings, settings_key="default_team_params", diff --git a/litellm/types/proxy/management_endpoints/ui_sso.py b/litellm/types/proxy/management_endpoints/ui_sso.py index d4b1d98f957..f68d818d991 100644 --- a/litellm/types/proxy/management_endpoints/ui_sso.py +++ b/litellm/types/proxy/management_endpoints/ui_sso.py @@ -229,3 +229,7 @@ class DefaultTeamSSOParams(LiteLLMPydanticObjectBase): default=None, description="Default permissions granted to members of newly created teams (e.g. /key/generate, /key/update, /key/delete). /key/info and /key/health are always included.", ) + organization_id: str | None = Field( + default=None, + description="Default organization for new teams created without an explicit organization", + ) diff --git a/tests/proxy_behavior/management/test_team_budget_limits.py b/tests/proxy_behavior/management/test_team_budget_limits.py index dad775370ad..96a6fe7234a 100644 --- a/tests/proxy_behavior/management/test_team_budget_limits.py +++ b/tests/proxy_behavior/management/test_team_budget_limits.py @@ -10,14 +10,14 @@ Pins the five helpers Driven through /team/new + /team/update. -Structural finding pinned here, identical in shape to F1's org aggregate: -both call sites (lines 985 + 1751) load the org via `get_org_object` -WITHOUT `include_budget_table=True`, so `org_table.litellm_budget_table` -is `None` and the org max_budget / org tpm / org rpm guards inside -`_check_org_team_limits` (lines 641–694, 670–694) silently no-op. The -`models` subset guard (lines 654–667) IS reachable because it reads -`org_table.models` directly. The `_check_user_team_limits` guards reach -all branches through `user_api_key_dict`, no relation include needed. +Structural finding, updated: /team/new loads the org via `get_org_object` +WITH `include_budget_table=True`, so the org max_budget / org tpm / org rpm +guards inside `_check_org_team_limits` are live there and are pinned as +enforced below. /team/update still loads the org without the budget +relation, so its budget guards remain no-ops. The `models` subset guard IS +reachable on both because it reads `org_table.models` directly. The +`_check_user_team_limits` guards reach all branches through +`user_api_key_dict`, no relation include needed. """ import uuid @@ -132,48 +132,67 @@ async def test_check_org_team_limits_models_subset( headers={"Authorization": f"Bearer {seeder}"}, json=body, ) - assert ( - resp.status_code == expected_status - ), f"{body!r} → {resp.status_code}: {resp.text}" + assert resp.status_code == expected_status, f"{body!r} → {resp.status_code}: {resp.text}" rows = await prisma.db.litellm_teamtable.find_many(where={"team_id": team_id}) assert len(rows) == (1 if expected_status == 200 else 0) # --------------------------------------------------------------------------- -# _check_org_team_limits — budget / tpm / rpm structurally unreachable -# (org_table.litellm_budget_table is None at guard time). Pin the -# no-op behavior so a future change that flips include_budget_table=True -# turns these into reds. +# _check_org_team_limits — budget / tpm / rpm live on /team/new since its +# get_org_object call passes include_budget_table=True. (/team/update still +# loads the org without the budget relation, so its guards remain no-ops.) # --------------------------------------------------------------------------- -_ORG_BUDGET_DEAD_SCENARIOS = [ +_ORG_BUDGET_ENFORCED_SCENARIOS = [ ( - "org_budget/over_max_budget_unenforced", + "org_budget/over_max_budget_rejected", {"max_budget": 100, "tpm_limit": None, "rpm_limit": None}, {"max_budget": 999_999}, + 400, ), ( - "org_tpm/over_unenforced", + "org_budget/within_max_budget_accepted", + {"max_budget": 100, "tpm_limit": None, "rpm_limit": None}, + {"max_budget": 50}, + 200, + ), + ( + "org_tpm/over_rejected", {"max_budget": None, "tpm_limit": 100, "rpm_limit": None}, {"tpm_limit": 999_999}, + 400, ), ( - "org_rpm/over_unenforced", + "org_tpm/within_accepted", + {"max_budget": None, "tpm_limit": 100, "rpm_limit": None}, + {"tpm_limit": 50}, + 200, + ), + ( + "org_rpm/over_rejected", {"max_budget": None, "tpm_limit": None, "rpm_limit": 100}, {"rpm_limit": 999_999}, + 400, + ), + ( + "org_rpm/within_accepted", + {"max_budget": None, "tpm_limit": None, "rpm_limit": 100}, + {"rpm_limit": 50}, + 200, ), ] @pytest.mark.parametrize( - "org_budget,body_extras", - [(b, c) for (_id, b, c) in _ORG_BUDGET_DEAD_SCENARIOS], - ids=[s[0] for s in _ORG_BUDGET_DEAD_SCENARIOS], + "org_budget,body_extras,expected_status", + [(b, c, d) for (_id, b, c, d) in _ORG_BUDGET_ENFORCED_SCENARIOS], + ids=[s[0] for s in _ORG_BUDGET_ENFORCED_SCENARIOS], ) -async def test_check_org_team_limits_budget_dead_code_pin( +async def test_check_org_team_limits_budget_enforced( org_budget, body_extras: Dict[str, Any], + expected_status: int, proxy_client, prisma, scratch, @@ -192,9 +211,9 @@ async def test_check_org_team_limits_budget_dead_code_pin( **body_extras, }, ) - assert resp.status_code == 200, resp.text + assert resp.status_code == expected_status, f"{body_extras!r} → {resp.status_code}: {resp.text}" rows = await prisma.db.litellm_teamtable.find_many(where={"team_id": team_id}) - assert len(rows) == 1 + assert len(rows) == (1 if expected_status == 200 else 0) # --------------------------------------------------------------------------- @@ -279,9 +298,9 @@ async def test_check_user_team_limits( **body_extras, }, ) - assert ( - resp.status_code == expected_status - ), f"caps={actor_caps} body={body_extras} → {resp.status_code}: {resp.text}" + assert resp.status_code == expected_status, ( + f"caps={actor_caps} body={body_extras} → {resp.status_code}: {resp.text}" + ) rows = await prisma.db.litellm_teamtable.find_many(where={"team_id": team_id}) assert len(rows) == (1 if expected_status == 200 else 0) @@ -376,9 +395,7 @@ async def test_proxy_admin_raise_budget_allowed(proxy_client, prisma, scratch): async def test_team_admin_remove_budget_cap_blocked(proxy_client, prisma, scratch): """A team admin cannot strip the team's cap (max_budget=null); removing the ceiling is the strongest possible raise -> proxy-admin only.""" - caller_cleartext = await _seed_scratch_actor_with_caps( - prisma, scratch.prefix, max_budget=100000.0 - ) + caller_cleartext = await _seed_scratch_actor_with_caps(prisma, scratch.prefix, max_budget=100000.0) team_id = await create_scratch_team( prisma, team_id=scratch.tag("team"), diff --git a/tests/proxy_behavior/management/test_team_new.py b/tests/proxy_behavior/management/test_team_new.py index 7b07f259641..9846566d0b9 100644 --- a/tests/proxy_behavior/management/test_team_new.py +++ b/tests/proxy_behavior/management/test_team_new.py @@ -72,13 +72,9 @@ async def test_team_new_authz_matrix( headers={"Authorization": f"Bearer {caller.cleartext}"}, json=body, ) - assert ( - resp.status_code == expected_status - ), f"{actor.value} org={org_target}: {resp.status_code} {resp.text}" + assert resp.status_code == expected_status, f"{actor.value} org={org_target}: {resp.status_code} {resp.text}" - row = await prisma.db.litellm_teamtable.find_unique( - where={"team_id": scratch.prefix} - ) + row = await prisma.db.litellm_teamtable.find_unique(where={"team_id": scratch.prefix}) if expected_status == 200: assert row is not None assert row.organization_id == org_id @@ -94,9 +90,7 @@ async def test_team_new_rejects_negative_budget(proxy_client, prisma, scratch, w json={"team_id": scratch.prefix, "max_budget": -1}, ) assert resp.status_code == 400, resp.text - row = await prisma.db.litellm_teamtable.find_unique( - where={"team_id": scratch.prefix} - ) + row = await prisma.db.litellm_teamtable.find_unique(where={"team_id": scratch.prefix}) assert row is None @@ -118,12 +112,10 @@ async def test_team_new_rejects_duplicate_team_id(proxy_client, prisma, scratch, assert second.status_code == 400, second.text -async def test_team_new_unknown_organization_is_500( - proxy_client, prisma, scratch, world -): - """SURFACED, NOT ENDORSED: a /team/new with an organization_id that does - not exist currently fails 500 (the role-resolution layer raises before - the handler's own 400 'Organization not found' check is reached).""" +async def test_team_new_unknown_organization_is_400(proxy_client, prisma, scratch, world): + """A /team/new with an organization_id that does not exist fails 400: + OrganizationNotFoundError is routed into the handler's own + 'Organization not found' guard instead of escaping as a 500.""" resp = await proxy_client.post( "/team/new", headers={"Authorization": f"Bearer {world.keys[Actor.PROXY_ADMIN].cleartext}"}, @@ -132,8 +124,7 @@ async def test_team_new_unknown_organization_is_500( "organization_id": scratch.tag("no-such-org"), }, ) - assert resp.status_code == 500, resp.text - row = await prisma.db.litellm_teamtable.find_unique( - where={"team_id": scratch.prefix} - ) + assert resp.status_code == 400, resp.text + assert "Organization not found" in resp.text + row = await prisma.db.litellm_teamtable.find_unique(where={"team_id": scratch.prefix}) assert row is None diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_default_params.py b/tests/test_litellm/proxy/management_endpoints/test_team_default_params.py index e0b90332ca0..a485d95db06 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_default_params.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_default_params.py @@ -10,13 +10,14 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException -sys.path.insert( - 0, os.path.abspath("../../../") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../../")) # Adds the parent directory to the system path import litellm from litellm.proxy._types import ( + LiteLLM_BudgetTable, + LiteLLM_OrganizationTable, NewTeamRequest, + ProxyException, UserAPIKeyAuth, LitellmUserRoles, ) @@ -76,9 +77,7 @@ class TestConfigFieldsDefaultTeamParams: db_param_value=db_settings, ) - assert result["litellm_settings"]["default_team_params"] == { - "max_budget": 100.0 - } + assert result["litellm_settings"]["default_team_params"] == {"max_budget": 100.0} # Existing keys preserved assert result["litellm_settings"]["cache"] is False @@ -172,6 +171,22 @@ class TestNewTeamDefaultParamsApplied: user_role=LitellmUserRoles.PROXY_ADMIN, ) + def _make_org(self, organization_id: str, max_budget: float | None = None) -> LiteLLM_OrganizationTable: + return LiteLLM_OrganizationTable( + organization_id=organization_id, + budget_id="budget-id", + created_by="admin-user", + updated_by="admin-user", + litellm_budget_table=None if max_budget is None else LiteLLM_BudgetTable(max_budget=max_budget), + ) + + def _patch_org_lookup(self, monkeypatch, **mock_kwargs) -> AsyncMock: + from litellm.proxy.management_endpoints import team_endpoints + + lookup = AsyncMock(**mock_kwargs) + monkeypatch.setattr(team_endpoints, "get_org_object", lookup) + return lookup + @pytest.mark.asyncio async def test_all_defaults_applied_when_not_provided(self, monkeypatch): """When no budget/rate/permission fields are in the request, all defaults apply.""" @@ -312,6 +327,7 @@ class TestNewTeamDefaultParamsApplied: assert data.tpm_limit is None assert data.rpm_limit is None assert data.team_member_permissions is None + assert data.organization_id is None @pytest.mark.asyncio async def test_legacy_default_team_settings_fallback(self, monkeypatch): @@ -370,6 +386,144 @@ class TestNewTeamDefaultParamsApplied: # default_team_params wins (100.0), legacy fallback (999.0) not used assert data.max_budget == 100.0 + @pytest.mark.asyncio + async def test_default_organization_applied_and_validated(self, monkeypatch): + """The default org must land before the org-validation block, so a defaulted + org goes through the same existence + org-limit checks as an explicit one.""" + from litellm.proxy.management_endpoints.team_endpoints import new_team + + monkeypatch.setattr(litellm, "default_team_params", {"organization_id": "default-org"}) + org_lookup = self._patch_org_lookup(monkeypatch, return_value=self._make_org("default-org")) + + data = NewTeamRequest(team_alias="my-team") + + try: + await new_team( + data=data, + user_api_key_dict=self._make_admin_auth(), + http_request=MagicMock(), + ) + except Exception: + pass + + assert data.organization_id == "default-org" + org_lookup.assert_awaited_once() + assert org_lookup.await_args.kwargs["org_id"] == "default-org" + + @pytest.mark.asyncio + async def test_explicit_organization_wins_over_default(self, monkeypatch): + """An organization_id in the request must not be replaced by the default.""" + from litellm.proxy.management_endpoints.team_endpoints import new_team + + monkeypatch.setattr(litellm, "default_team_params", {"organization_id": "default-org"}) + org_lookup = self._patch_org_lookup(monkeypatch, return_value=self._make_org("explicit-org")) + + data = NewTeamRequest(team_alias="my-team", organization_id="explicit-org") + + try: + await new_team( + data=data, + user_api_key_dict=self._make_admin_auth(), + http_request=MagicMock(), + ) + except Exception: + pass + + assert data.organization_id == "explicit-org" + assert org_lookup.await_args.kwargs["org_id"] == "explicit-org" + + @pytest.mark.asyncio + async def test_nonexistent_default_organization_returns_400(self, monkeypatch): + """get_org_object raises instead of returning None, so an org that no longer + exists surfaced as a 500; team creation must report a 400 instead.""" + from litellm.proxy.auth.auth_checks import OrganizationNotFoundError + from litellm.proxy.management_endpoints.team_endpoints import new_team + + monkeypatch.setattr(litellm, "default_team_params", {"organization_id": "deleted-org"}) + self._patch_org_lookup( + monkeypatch, + side_effect=OrganizationNotFoundError("Organization doesn't exist in db. Organization=deleted-org"), + ) + + with pytest.raises(ProxyException) as exc_info: + await new_team( + data=NewTeamRequest(team_alias="my-team"), + user_api_key_dict=self._make_admin_auth(), + http_request=MagicMock(), + ) + + assert exc_info.value.code == "400" + assert "deleted-org" in exc_info.value.message + + @pytest.mark.asyncio + async def test_defaulted_max_budget_validated_against_org_budget(self, monkeypatch): + """Defaults must be applied BEFORE _check_org_team_limits runs, or a default + max_budget above the org's cap is persisted unchecked.""" + from litellm.proxy.management_endpoints.team_endpoints import new_team + + monkeypatch.setattr( + litellm, + "default_team_params", + {"organization_id": "capped-org", "max_budget": 500.0}, + ) + self._patch_org_lookup(monkeypatch, return_value=self._make_org("capped-org", max_budget=100.0)) + + with pytest.raises(ProxyException) as exc_info: + await new_team( + data=NewTeamRequest(team_alias="my-team"), + user_api_key_dict=self._make_admin_auth(), + http_request=MagicMock(), + ) + + assert exc_info.value.code == "400" + assert "exceeds organization's max_budget" in exc_info.value.message + + @pytest.mark.asyncio + async def test_explicit_budget_validated_against_default_org_budget(self, monkeypatch): + """The org lookup must load the budget table (include_budget_table=True); + without it litellm_budget_table is None and every budget comparison is skipped.""" + from litellm.proxy.management_endpoints.team_endpoints import new_team + + monkeypatch.setattr(litellm, "default_team_params", {"organization_id": "capped-org"}) + org_lookup = self._patch_org_lookup(monkeypatch, return_value=self._make_org("capped-org", max_budget=100.0)) + + with pytest.raises(ProxyException) as exc_info: + await new_team( + data=NewTeamRequest(team_alias="my-team", max_budget=500.0), + user_api_key_dict=self._make_admin_auth(), + http_request=MagicMock(), + ) + + assert exc_info.value.code == "400" + assert "exceeds organization's max_budget" in exc_info.value.message + assert org_lookup.await_args.kwargs["include_budget_table"] is True + + @pytest.mark.asyncio + async def test_defaults_within_org_budget_still_created(self, monkeypatch): + """A default budget under the org cap must not be rejected by the reordered check.""" + from litellm.proxy.management_endpoints.team_endpoints import new_team + + monkeypatch.setattr( + litellm, + "default_team_params", + {"organization_id": "capped-org", "max_budget": 50.0}, + ) + self._patch_org_lookup(monkeypatch, return_value=self._make_org("capped-org", max_budget=100.0)) + + data = NewTeamRequest(team_alias="my-team") + + try: + await new_team( + data=data, + user_api_key_dict=self._make_admin_auth(), + http_request=MagicMock(), + ) + except Exception: + pass + + assert data.organization_id == "capped-org" + assert data.max_budget == 50.0 + # --------------------------------------------------------------------------- # _update_litellm_setting: setattr ordering @@ -536,18 +690,12 @@ class TestBulkUpdateTeamMemberPermissions: mock_batcher.commit = AsyncMock(return_value=None) mock_prisma = MagicMock() - mock_prisma.db.litellm_teamtable.find_many = AsyncMock( - return_value=[team_a, team_b] - ) + mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_a, team_b]) mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) - data = BulkUpdateTeamMemberPermissionsRequest( - permissions=["/team/daily/activity"], apply_to_all_teams=True - ) - result = await bulk_update_team_member_permissions( - data=data, user_api_key_dict=self._admin_key_dict() - ) + data = BulkUpdateTeamMemberPermissionsRequest(permissions=["/team/daily/activity"], apply_to_all_teams=True) + result = await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict()) assert result["teams_updated"] == 2 calls = mock_batcher.litellm_teamtable.update.call_args_list @@ -555,19 +703,14 @@ class TestBulkUpdateTeamMemberPermissions: team_a_call = [c for c in calls if c.kwargs["where"]["team_id"] == "team-a"][0] assert "/key/generate" in team_a_call.kwargs["data"]["team_member_permissions"] - assert ( - "/team/daily/activity" - in team_a_call.kwargs["data"]["team_member_permissions"] - ) + assert "/team/daily/activity" in team_a_call.kwargs["data"]["team_member_permissions"] team_b_call = [c for c in calls if c.kwargs["where"]["team_id"] == "team-b"][0] assert "/key/delete" in team_b_call.kwargs["data"]["team_member_permissions"] assert "/key/update" in team_b_call.kwargs["data"]["team_member_permissions"] @pytest.mark.asyncio - async def test_all_teams_skips_teams_that_already_have_permission( - self, monkeypatch - ): + async def test_all_teams_skips_teams_that_already_have_permission(self, monkeypatch): """apply_to_all_teams: teams that already have the permission are skipped.""" from litellm.proxy.management_endpoints.team_endpoints import ( bulk_update_team_member_permissions, @@ -583,18 +726,12 @@ class TestBulkUpdateTeamMemberPermissions: mock_batcher.commit = AsyncMock(return_value=None) mock_prisma = MagicMock() - mock_prisma.db.litellm_teamtable.find_many = AsyncMock( - return_value=[team_has, team_missing] - ) + mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_has, team_missing]) mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) - data = BulkUpdateTeamMemberPermissionsRequest( - permissions=["/team/daily/activity"], apply_to_all_teams=True - ) - result = await bulk_update_team_member_permissions( - data=data, user_api_key_dict=self._admin_key_dict() - ) + data = BulkUpdateTeamMemberPermissionsRequest(permissions=["/team/daily/activity"], apply_to_all_teams=True) + result = await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict()) assert result["teams_updated"] == 1 calls = mock_batcher.litellm_teamtable.update.call_args_list @@ -618,18 +755,12 @@ class TestBulkUpdateTeamMemberPermissions: mock_batcher.commit = AsyncMock(return_value=None) mock_prisma = MagicMock() - mock_prisma.db.litellm_teamtable.find_many = AsyncMock( - side_effect=[page1, page2] - ) + mock_prisma.db.litellm_teamtable.find_many = AsyncMock(side_effect=[page1, page2]) mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) - data = BulkUpdateTeamMemberPermissionsRequest( - permissions=["/team/daily/activity"], apply_to_all_teams=True - ) - result = await bulk_update_team_member_permissions( - data=data, user_api_key_dict=self._admin_key_dict() - ) + data = BulkUpdateTeamMemberPermissionsRequest(permissions=["/team/daily/activity"], apply_to_all_teams=True) + result = await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict()) assert result["teams_updated"] == 502 find_calls = mock_prisma.db.litellm_teamtable.find_many.call_args_list @@ -656,18 +787,14 @@ class TestBulkUpdateTeamMemberPermissions: mock_batcher.commit = AsyncMock(return_value=None) mock_prisma = MagicMock() - mock_prisma.db.litellm_teamtable.find_many = AsyncMock( - return_value=[team_a, team_b] - ) + mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_a, team_b]) mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) data = BulkUpdateTeamMemberPermissionsRequest( permissions=["/team/daily/activity"], team_ids=["team-a", "team-b"] ) - result = await bulk_update_team_member_permissions( - data=data, user_api_key_dict=self._admin_key_dict() - ) + result = await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict()) assert result["teams_updated"] == 2 @@ -692,18 +819,14 @@ class TestBulkUpdateTeamMemberPermissions: mock_batcher.commit = AsyncMock(return_value=None) mock_prisma = MagicMock() - mock_prisma.db.litellm_teamtable.find_many = AsyncMock( - return_value=[team_has, team_missing] - ) + mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_has, team_missing]) mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) data = BulkUpdateTeamMemberPermissionsRequest( permissions=["/team/daily/activity"], team_ids=["team-has", "team-missing"] ) - result = await bulk_update_team_member_permissions( - data=data, user_api_key_dict=self._admin_key_dict() - ) + result = await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict()) assert result["teams_updated"] == 1 calls = mock_batcher.litellm_teamtable.update.call_args_list @@ -731,9 +854,7 @@ class TestBulkUpdateTeamMemberPermissions: ) with pytest.raises(HTTPException) as exc_info: - await bulk_update_team_member_permissions( - data=data, user_api_key_dict=self._admin_key_dict() - ) + await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict()) assert exc_info.value.status_code == 404 assert "team-b" in str(exc_info.value.detail) @@ -753,14 +874,10 @@ class TestBulkUpdateTeamMemberPermissions: mock_prisma = MagicMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) - data = BulkUpdateTeamMemberPermissionsRequest( - permissions=["/team/daily/activity"] - ) + data = BulkUpdateTeamMemberPermissionsRequest(permissions=["/team/daily/activity"]) with pytest.raises(HTTPException) as exc_info: - await bulk_update_team_member_permissions( - data=data, user_api_key_dict=self._admin_key_dict() - ) + await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict()) assert exc_info.value.status_code == 400 @@ -784,9 +901,7 @@ class TestBulkUpdateTeamMemberPermissions: ) with pytest.raises(HTTPException) as exc_info: - await bulk_update_team_member_permissions( - data=data, user_api_key_dict=self._admin_key_dict() - ) + await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict()) assert exc_info.value.status_code == 400 @@ -804,9 +919,7 @@ class TestBulkUpdateTeamMemberPermissions: monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) data = BulkUpdateTeamMemberPermissionsRequest(permissions=[]) - result = await bulk_update_team_member_permissions( - data=data, user_api_key_dict=self._admin_key_dict() - ) + result = await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._admin_key_dict()) assert result["teams_updated"] == 0 mock_prisma.db.litellm_teamtable.find_many.assert_not_called() @@ -824,14 +937,10 @@ class TestBulkUpdateTeamMemberPermissions: mock_prisma = MagicMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) - data = BulkUpdateTeamMemberPermissionsRequest( - permissions=["/team/daily/activity"], apply_to_all_teams=True - ) + data = BulkUpdateTeamMemberPermissionsRequest(permissions=["/team/daily/activity"], apply_to_all_teams=True) with pytest.raises(HTTPException) as exc_info: - await bulk_update_team_member_permissions( - data=data, user_api_key_dict=self._non_admin_key_dict() - ) + await bulk_update_team_member_permissions(data=data, user_api_key_dict=self._non_admin_key_dict()) assert exc_info.value.status_code == 403 @@ -844,6 +953,4 @@ class TestBulkUpdateTeamMemberPermissions: ) with pytest.raises(ValidationError): - BulkUpdateTeamMemberPermissionsRequest( - permissions=["/not/a/real/permission"] - ) + BulkUpdateTeamMemberPermissionsRequest(permissions=["/not/a/real/permission"]) diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 795b7cd5a9e..979eb09d7db 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -606,6 +606,55 @@ async def test_default_team_params(team_params): assert create_call_args["models"] == ["special-gpt-5"] +@pytest.mark.asyncio +@pytest.mark.parametrize( + "team_params", + [ + DefaultTeamSSOParams(max_budget=10, budget_duration="1d", organization_id="default-org"), + {"max_budget": 10, "budget_duration": "1d", "organization_id": "default-org"}, + ], +) +async def test_default_team_params_organization_id_reaches_sso_created_team(team_params): + """The SSO auto-team path builds NewTeamRequest straight from default_team_params, + so a default organization_id must land on the created team row and be validated.""" + from litellm.proxy._types import LiteLLM_OrganizationTable + + litellm.default_team_params = team_params + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.find_first = AsyncMock(return_value=None) + mock_prisma.db.litellm_teamtable.create = AsyncMock() + mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) + mock_prisma.get_data = AsyncMock(return_value=None) + mock_prisma.jsonify_team_object = MagicMock(side_effect=lambda db_data: db_data) + + mock_org = LiteLLM_OrganizationTable( + organization_id="default-org", + budget_id="budget-id", + created_by="admin", + updated_by="admin", + ) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( + "litellm.proxy.management_endpoints.team_endpoints.get_org_object", + AsyncMock(return_value=mock_org), + ) as mock_get_org: + team_id = str(uuid.uuid4()) + await MicrosoftSSOHandler.create_litellm_teams_from_service_principal_team_ids( + service_principal_teams=[ + MicrosoftServicePrincipalTeam( + principalId=team_id, + principalDisplayName="Test Team", + ) + ] + ) + + mock_prisma.db.litellm_teamtable.create.assert_called_once() + create_call_args = mock_prisma.db.litellm_teamtable.create.call_args.kwargs["data"] + assert create_call_args["organization_id"] == "default-org" + assert mock_get_org.call_args.kwargs["org_id"] == "default-org" + + @pytest.mark.asyncio async def test_create_team_without_default_params(): """ diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index d4fd5bc2dce..1075bffbeb2 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -2816,6 +2816,94 @@ def test_update_internal_user_settings_without_teams_skips_team_lookup(mock_prox assert mock_proxy_config["save_call_count"]() == 1 +@pytest.fixture +def mock_organization_lookup(monkeypatch): + """Back /update/default_team_settings with a fake organization table. + + Yields the set of organization ids that exist; the test mutates it before the call. + """ + from unittest.mock import AsyncMock, MagicMock + + import litellm + import litellm.proxy.proxy_server as proxy_server_module + + existing_organization_ids: set = set() + + async def _find_unique(where): + organization_id = where["organization_id"] + if organization_id not in existing_organization_ids: + return None + return {"organization_id": organization_id} + + find_unique = AsyncMock(side_effect=_find_unique) + fake_prisma = MagicMock() + fake_prisma.db.litellm_organizationtable.find_unique = find_unique + + monkeypatch.setattr(proxy_server_module, "prisma_client", fake_prisma) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) + monkeypatch.setattr(litellm, "default_team_params", {}) + + return { + "existing_organization_ids": existing_organization_ids, + "find_unique": find_unique, + } + + +def test_update_default_team_settings_rejects_unknown_organization( + mock_proxy_config, mock_auth, mock_organization_lookup +): + """Regression: an unknown default org saved fine here and then failed every + future team creation, far from the admin who typed it.""" + mock_organization_lookup["existing_organization_ids"].add("real-org") + + resp = client.patch( + "/update/default_team_settings", + json={"max_budget": 10.0, "organization_id": "ghost-org"}, + ) + + assert resp.status_code == 400, resp.text + assert "ghost-org" in resp.json()["detail"]["error"] + assert mock_proxy_config["save_call_count"]() == 0 + + import litellm + + assert litellm.default_team_params == {} + + +def test_update_default_team_settings_saves_when_organization_exists( + mock_proxy_config, mock_auth, mock_organization_lookup +): + """A real organization id still saves and reaches the in-memory settings.""" + mock_organization_lookup["existing_organization_ids"].add("real-org") + + resp = client.patch( + "/update/default_team_settings", + json={"max_budget": 10.0, "organization_id": "real-org"}, + ) + + assert resp.status_code == 200, resp.text + assert resp.json()["settings"]["organization_id"] == "real-org" + assert mock_proxy_config["save_call_count"]() == 1 + + import litellm + + assert litellm.default_team_params["organization_id"] == "real-org" + + +def test_update_default_team_settings_without_organization_skips_lookup( + mock_proxy_config, mock_auth, mock_organization_lookup +): + """Settings changes that don't set an organization must not pay for a DB round trip.""" + resp = client.patch( + "/update/default_team_settings", + json={"max_budget": 10.0}, + ) + + assert resp.status_code == 200, resp.text + mock_organization_lookup["find_unique"].assert_not_awaited() + assert mock_proxy_config["save_call_count"]() == 1 + + def test_update_mcp_semantic_filter_settings_requires_proxy_admin(monkeypatch): """Non-admin callers must not mutate global MCP semantic filter settings.""" from litellm.proxy._types import UserAPIKeyAuth diff --git a/ui/litellm-dashboard/src/components/TeamSSOSettings.test.tsx b/ui/litellm-dashboard/src/components/TeamSSOSettings.test.tsx index dd2dc42fe88..431931eb575 100644 --- a/ui/litellm-dashboard/src/components/TeamSSOSettings.test.tsx +++ b/ui/litellm-dashboard/src/components/TeamSSOSettings.test.tsx @@ -1,8 +1,8 @@ import React from "react"; -import { screen, waitFor } from "@testing-library/react"; +import { screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import { renderWithProviders } from "../../tests/test-utils"; +import { renderWithProviders, testQueryClient } from "../../tests/test-utils"; import TeamSSOSettings from "./TeamSSOSettings"; import * as networking from "./networking"; import NotificationsManager from "./molecules/notifications_manager"; @@ -37,6 +37,46 @@ vi.mock("./key_team_helpers/fetch_available_models_team_key", () => ({ getModelDisplayName: vi.fn((model: string) => model), })); +vi.mock("./common_components/OrganizationDropdown", () => ({ + default: ({ + organizations, + value, + onChange, + placeholder, + loading, + }: { + organizations?: { organization_id: string; organization_alias: string }[] | null; + value?: string; + onChange?: (value: string) => void; + placeholder?: string; + loading?: boolean; + }) => ( +
+ + +
+ ), +})); + vi.mock("./ModelSelect/ModelSelect", () => { const ModelSelect = ({ value, onChange }: { value: string[]; onChange: (value: string[]) => void }) => (