diff --git a/litellm/models/team.py b/litellm/models/team.py index 8edf10703b1..32512373648 100644 --- a/litellm/models/team.py +++ b/litellm/models/team.py @@ -7,14 +7,76 @@ budget-window value types and the team-model alias table). Re-exported from """ import json +from collections.abc import Mapping from datetime import datetime -from typing import Final, Literal, Optional +from typing import Annotated, Final, Literal, Optional -from pydantic import BaseModel, ConfigDict, Field, model_validator +from pydantic import AfterValidator, BaseModel, ConfigDict, Field, model_validator from litellm.models.object_permission import LiteLLM_ObjectPermissionTable from litellm.types.llms.base import LiteLLMPydanticObjectBase +TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY: Final = "team_member_max_budget_alert_emails" + + +def _parse_team_member_budget_alert_threshold(raw: object) -> str: + if isinstance(raw, str) and raw.isdigit() and len(raw) <= 3 and 1 <= int(raw) <= 100: + return str(int(raw)) + raise ValueError( + f"metadata.{TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY} thresholds must be whole-number percentages " + f"from 1 to 100, got {raw!r}" + ) + + +def _is_plausible_email(raw: object) -> bool: + if not isinstance(raw, str): + return False + local, at, domain = raw.strip().partition("@") + return bool(local and at and domain) and not any(c.isspace() or c == "@" for c in local + domain) + + +def _parse_team_member_budget_alert_recipients(threshold: str, raw: object) -> list[str]: + if raw is None: + return [] + if not isinstance(raw, list): + raise ValueError( + f"metadata.{TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY}[{threshold!r}] must be a list of email addresses" + ) + invalid: Final = [email for email in raw if not _is_plausible_email(email)] + if invalid: + raise ValueError( + f"metadata.{TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY}[{threshold!r}] has invalid email addresses: {invalid!r}" + ) + return list(dict.fromkeys(email.strip() for email in raw)) + + +def validate_team_request_metadata(metadata: dict) -> dict: + """ + Reject a malformed team_member_max_budget_alert_emails on write and store it canonically, + e.g. {"50": [], "100": ["finance@x.com"]}, so a bad threshold is a 422 instead of an alert that never fires. + """ + raw: Final = metadata.get(TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY) + if raw is None: + return metadata + if not isinstance(raw, Mapping): + raise ValueError( + f"metadata.{TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY} must map percentages to email lists, " + 'e.g. {"50": [], "100": ["finance@example.com"]}' + ) + parsed: Final[dict[str, list[str]]] = {} + for raw_threshold, raw_recipients in raw.items(): + threshold = _parse_team_member_budget_alert_threshold(raw_threshold) + if threshold in parsed: + raise ValueError( + f"metadata.{TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY} lists the {threshold}% threshold more than once" + ) + parsed[threshold] = _parse_team_member_budget_alert_recipients(threshold, raw_recipients) + return {**metadata, TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY: parsed} + + +# Request-side only: LiteLLM_TeamTable keeps plain `dict` so reading a stored row never fails validation. +TeamRequestMetadata = Annotated[dict, AfterValidator(validate_team_request_metadata)] + class MemberBase(LiteLLMPydanticObjectBase): user_id: str | None = Field( diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 54574ed64e3..4a52cbdd12e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2075,6 +2075,7 @@ class OrgMember(MemberBase): from litellm.models.team import TeamBase as TeamBase # noqa: E402 +from litellm.models.team import TeamRequestMetadata as TeamRequestMetadata # noqa: E402 RouterSettingsDict = Annotated[ dict[str, object], @@ -2083,6 +2084,7 @@ RouterSettingsDict = Annotated[ class NewTeamRequest(TeamBase): + metadata: TeamRequestMetadata | None = None router_settings: RouterSettingsDict | None = None model_aliases: dict | None = None model_max_budget: GenericBudgetConfigType | None = Field( @@ -2158,7 +2160,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase): team_id: str # required team_alias: str | None = None organization_id: str | None = None - metadata: dict | None = None + metadata: TeamRequestMetadata | None = None tpm_limit: int | None = None rpm_limit: int | None = None tpd_limit: int | None = None @@ -2211,6 +2213,9 @@ class PatchTeamRequest(UpdateTeamRequest): """ team_id: str | None = None + # A merge-patch body, not final metadata: `{"50": null}` must survive to delete that threshold. + # patch_team validates the merged result when it re-parses it as UpdateTeamRequest. + metadata: dict | None = None class ResetTeamBudgetRequest(LiteLLMPydanticObjectBase): diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 2c93aad820c..dd658148b47 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -40,6 +40,7 @@ from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.models.project import LiteLLM_ProjectTable +from litellm.models.team import TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY from litellm.proxy._types import ( RBAC_ROLES, CallInfo, @@ -5524,9 +5525,6 @@ async def _virtual_key_max_budget_alert_check( ) -TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY: Final = "team_member_max_budget_alert_emails" - - def _team_member_max_budget_alert_check( team_id: str, team_alias: str | None, 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 7a9b6b66946..c5c512c27ea 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -12947,6 +12947,101 @@ def test_patch_team_request_makes_team_id_optional(): assert set(UpdateTeamRequest.model_fields).issubset(set(PatchTeamRequest.model_fields)) +# --------------------------------------------------------------------------- +# metadata.team_member_max_budget_alert_emails is validated on every team write +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("request_cls_name", ["NewTeamRequest", "UpdateTeamRequest"]) +def test_team_member_budget_alert_emails_stored_canonically(request_cls_name): + import litellm.proxy._types as proxy_types + + parsed = getattr(proxy_types, request_cls_name).model_validate( + { + "team_id": "t", + "metadata": { + "cost_center": "1", + "team_member_max_budget_alert_emails": { + "050": None, + "100": [" finance@example.com", "finance@example.com"], + }, + }, + } + ) + + assert parsed.metadata == { + "cost_center": "1", + "team_member_max_budget_alert_emails": {"50": [], "100": ["finance@example.com"]}, + } + + +@pytest.mark.parametrize("request_cls_name", ["NewTeamRequest", "UpdateTeamRequest"]) +@pytest.mark.parametrize( + "alert_emails", + [ + {"0": []}, + {"101": []}, + {"12.5": []}, + {"-5": []}, + {"fifty": []}, + {"50": [], "050": []}, + {"50": "finance@example.com"}, + {"50": ["not-an-email"]}, + {"50": ["a b@example.com"]}, + {"50": [123]}, + ["50"], + ], +) +def test_team_member_budget_alert_emails_rejects_malformed(request_cls_name, alert_emails): + from pydantic import ValidationError + + import litellm.proxy._types as proxy_types + + with pytest.raises(ValidationError, match="team_member_max_budget_alert_emails"): + getattr(proxy_types, request_cls_name).model_validate( + {"team_id": "t", "metadata": {"team_member_max_budget_alert_emails": alert_emails}} + ) + + +def test_team_row_with_malformed_alert_emails_still_loads(): + """Validation is write-side only; reading a stored team never fails on this key.""" + from litellm.proxy._types import LiteLLM_TeamTable + + row = LiteLLM_TeamTable(team_id="t", metadata={"team_member_max_budget_alert_emails": {"0": "x"}}) + + assert row.metadata == {"team_member_max_budget_alert_emails": {"0": "x"}} + + +@pytest.mark.asyncio +async def test_patch_null_deletes_one_alert_threshold(): + """The PATCH body is a merge patch, so {"50": null} must delete that threshold rather than + being coerced to an empty recipient list before the merge.""" + _, update_mock = await _drive_team_write( + "patch", + existing_metadata={"team_member_max_budget_alert_emails": {"50": [], "100": ["finance@example.com"]}}, + raw_body={"metadata": {"team_member_max_budget_alert_emails": {"50": None}}}, + ) + + written_metadata = update_mock.call_args.kwargs["data"]["metadata"] + if isinstance(written_metadata, str): + written_metadata = json.loads(written_metadata) + assert written_metadata["team_member_max_budget_alert_emails"] == {"100": ["finance@example.com"]} + + +@pytest.mark.asyncio +async def test_patch_rejects_malformed_alert_threshold(): + from litellm.proxy._types import ProxyException + + with pytest.raises(ProxyException) as exc: + await _drive_team_write( + "patch", + existing_metadata={"team_member_max_budget_alert_emails": {"50": []}}, + raw_body={"metadata": {"team_member_max_budget_alert_emails": {"150": []}}}, + ) + + assert "team_member_max_budget_alert_emails" in str(exc.value.message) + + def test_patch_team_route_publishes_its_request_body_schema(): """The dashboard's generated client types this call off the OpenAPI spec, which FastAPI can only emit because the body is a declared parameter."""