diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 7df725bf965..20dc14aa5a3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1856,6 +1856,17 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase): default_team_member_models: Optional[List[str]] = None # default allowed_models seeded onto new team members +class PatchTeamRequest(UpdateTeamRequest): + """ + Body of PATCH /team/{team_id}. + + Identical to UpdateTeamRequest except team_id is optional, because PATCH takes it + from the path. A team_id in the body is still accepted when it matches the path. + """ + + team_id: str | None = None + + class ResetTeamBudgetRequest(LiteLLMPydanticObjectBase): """ internal type used to reset the budget on a team diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 3a7e89aa20d..59b0cbc4ae7 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -47,6 +47,7 @@ from litellm.proxy._types import ( Member, NewTeamRequest, OrgMember, + PatchTeamRequest, ProxyErrorTypes, ProxyException, SpecialManagementEndpointEnums, @@ -1956,6 +1957,7 @@ async def update_team( ) async def patch_team( team_id: str, + data: PatchTeamRequest, http_request: Request, user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], litellm_changed_by: Annotated[ @@ -1968,11 +1970,12 @@ async def patch_team( """ Partially update a team using RFC 7386 JSON Merge Patch semantics. - `team_id` is taken from the path. `metadata` is merged with the team's stored - metadata rather than replacing it: an omitted key is preserved, `key: null` - deletes it, and any other value overwrites (recursing into nested objects). - Every other field behaves exactly like `POST /team/update` (omitted preserves, - a value overwrites). Returns the full updated team. + `team_id` is taken from the path; a `team_id` in the body is accepted only when it + matches. `metadata` is merged with the team's stored metadata rather than replacing + it: an omitted key is preserved, `key: null` deletes it, and any other value + overwrites (recursing into nested objects). Every other field behaves exactly like + `POST /team/update` (omitted preserves, a value overwrites). Returns the full + updated team. ``` curl --location --request PATCH 'http://0.0.0.0:4000/team/8d916b1c-510d-4894-a334-1c16a93344f5' \ @@ -1992,21 +1995,15 @@ async def patch_team( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - try: - body = await http_request.json() - except (json.JSONDecodeError, ValueError): - raise HTTPException(status_code=400, detail={"error": "Request body must be a JSON object"}) - if not isinstance(body, dict): - raise HTTPException(status_code=400, detail={"error": "Request body must be a JSON object"}) - - body_team_id = body.pop("team_id", None) - if body_team_id is not None and body_team_id != team_id: + if data.team_id is not None and data.team_id != team_id: raise HTTPException( status_code=400, - detail={"error": f"team_id in body ({body_team_id}) does not match team_id in path ({team_id})"}, + detail={"error": f"team_id in body ({data.team_id}) does not match team_id in path ({team_id})"}, ) - if "metadata" in body: + patch_fields = data.model_dump(exclude_unset=True, exclude={"team_id"}) + + if "metadata" in patch_fields: existing_team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) if existing_team_row is None: raise HTTPException( @@ -2014,9 +2011,9 @@ async def patch_team( detail={"error": f"Team not found, passed team_id={team_id}"}, ) existing_metadata = existing_team_row.metadata if isinstance(existing_team_row.metadata, dict) else {} - body["metadata"] = apply_json_merge_patch(existing_metadata, body["metadata"]) + patch_fields["metadata"] = apply_json_merge_patch(existing_metadata, patch_fields["metadata"]) - update_request = UpdateTeamRequest(team_id=team_id, **body) + update_request = UpdateTeamRequest(team_id=team_id, **patch_fields) result = await update_team( data=update_request, diff --git a/litellm/proxy/mcp_registry.json b/litellm/proxy/mcp_registry.json index 84431634e24..f37fc39813e 100644 --- a/litellm/proxy/mcp_registry.json +++ b/litellm/proxy/mcp_registry.json @@ -198,13 +198,9 @@ "icon_url": "https://cdn.simpleicons.org/googledrive", "category": "Productivity", "registry_url": null, - "transport": "stdio", - "command": "npx", - "args": ["-y", "@modelcontextprotocol/server-gdrive"], - "env_vars": [ - {"name": "GOOGLE_CLIENT_ID", "description": "Google OAuth Client ID", "secret": false}, - {"name": "GOOGLE_CLIENT_SECRET", "description": "Google OAuth Client Secret", "secret": true} - ] + "transport": "http", + "url": "https://drivemcp.googleapis.com/mcp/v1", + "env_vars": [] }, { "name": "google_calendar", 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 d5d61341c93..5202c8cbfc0 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -9813,7 +9813,6 @@ async def _drive_team_write( raw_body=None, user=None, find_returns_none=False, - json_side_effect=None, ): """Drive POST ``update_team`` or PATCH ``patch_team`` against a mocked team. @@ -9828,6 +9827,7 @@ async def _drive_team_write( from litellm.proxy._types import ( LiteLLM_TeamTable, LitellmUserRoles, + PatchTeamRequest, UpdateTeamRequest, UserAPIKeyAuth, ) @@ -9873,14 +9873,10 @@ async def _drive_team_write( litellm_changed_by=None, ) else: - if json_side_effect is not None: - req.json = AsyncMock(side_effect=json_side_effect) - else: - req.json = AsyncMock( - return_value=raw_body if raw_body is not None else dict(payload or {}) - ) + body = raw_body if raw_body is not None else dict(payload or {}) result = await patch_team( team_id=_PATCH_TEAM_ID, + data=PatchTeamRequest.model_validate(body), http_request=req, user_api_key_dict=auth, litellm_changed_by=None, @@ -10028,25 +10024,36 @@ async def test_patch_strips_system_managed_metadata_key_like_post(): assert patch_meta == {"cost_center": "9999"} -@pytest.mark.asyncio -@pytest.mark.parametrize("raw_body", [["not", "an", "object"], "a-string", 42, True]) -async def test_patch_rejects_non_object_body(raw_body): - from litellm.proxy._types import ProxyException +@pytest.mark.parametrize( + "kwargs", + [ + {"json": ["not", "an", "object"]}, + {"json": "a-string"}, + {"json": 42}, + {"content": b"{not json"}, + {"json": {"tpm_limit": "not-an-int"}}, + ], + ids=["list", "string", "number", "malformed-json", "wrong-field-type"], +) +def test_patch_rejects_a_malformed_body_with_422(kwargs): + """The body is a declared parameter, so FastAPI rejects a malformed one before the + handler runs. This is the same 422 POST /team/update already returns; the route + previously answered 400 here and 500 for a wrongly typed field, reporting a caller + mistake as a server fault.""" + from fastapi import FastAPI + from fastapi.testclient import TestClient - with pytest.raises(ProxyException) as exc: - await _drive_team_write("patch", existing_metadata={"a": 1}, raw_body=raw_body) - assert exc.value.code == "400" or exc.value.code == 400 + from litellm.proxy._types import PatchTeamRequest + app = FastAPI() -@pytest.mark.asyncio -async def test_patch_rejects_invalid_json_body(): - from litellm.proxy._types import ProxyException + @app.patch("/team/{team_id}") + async def _route(team_id: str, data: PatchTeamRequest): # pragma: no cover - schema only + return {} - with pytest.raises(ProxyException) as exc: - await _drive_team_write( - "patch", existing_metadata={"a": 1}, json_side_effect=ValueError("no body") - ) - assert exc.value.code == "400" or exc.value.code == 400 + response = TestClient(app).patch("/team/abc", **kwargs) + + assert response.status_code == 422 @pytest.mark.asyncio @@ -10116,3 +10123,103 @@ async def test_patch_returns_full_team_object_not_wrapper(): ) assert isinstance(result, LiteLLM_TeamTable) assert result.team_id == _PATCH_TEAM_ID + + +# --------------------------------------------------------------------------- +# PATCH body is validated through PatchTeamRequest before it is handed to +# update_team. The write below must stay byte-identical to what the untyped +# **body construction produced, or a partial update starts writing columns the +# caller never mentioned. +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_patch_writes_only_the_keys_the_caller_sent(): + """An omitted field must not reach the DB write at all. If validation ever + materialises defaults, every unmentioned column gets overwritten with null.""" + _, update_mock = await _drive_team_write("patch", raw_body={"tpm_limit": 5}) + written = update_mock.call_args.kwargs["data"] + + assert written["tpm_limit"] == 5 + for untouched in ("rpm_limit", "max_budget", "models", "blocked", "budget_duration"): + assert untouched not in written, f"{untouched} was written despite not being sent" + + +@pytest.mark.asyncio +async def test_patch_preserves_explicit_null_as_a_clear(): + """null is a clear, not an omission: it has to survive validation and reach the write.""" + _, update_mock = await _drive_team_write("patch", raw_body={"max_budget": None}) + written = update_mock.call_args.kwargs["data"] + + assert "max_budget" in written + assert written["max_budget"] is None + + +def _patch_body_to_update_request(body: dict): + """The exact reshaping patch_team performs between the raw body and update_team.""" + from litellm.proxy._types import PatchTeamRequest, UpdateTeamRequest + + parsed = PatchTeamRequest.model_validate(body) + return UpdateTeamRequest( + team_id=_PATCH_TEAM_ID, + **parsed.model_dump(exclude_unset=True, exclude={"team_id"}), + ) + + +@pytest.mark.parametrize( + "body", + [ + {"tpm_limit": 5}, + {"max_budget": None}, + {"object_permission": {"vector_stores": []}}, + {"metadata": {"a": 1, "b": None}}, + {"models": ["gpt-4"], "blocked": False}, + ], + ids=["scalar", "explicit-null", "partial-nested", "metadata-with-null", "list-and-false"], +) +def test_patch_body_reshaping_adds_no_keys_the_caller_did_not_send(body): + """Validating through PatchTeamRequest must be shape-preserving. If it ever + materialises defaults, a partial update silently overwrites untouched columns, + and for the merge-only object_permission it would wipe sibling sub-keys.""" + reshaped = _patch_body_to_update_request(body) + dumped = reshaped.model_dump(exclude_unset=True, exclude={"team_id"}) + + assert dumped == body + assert reshaped.model_fields_set == set(body) | {"team_id"} + + +@pytest.mark.asyncio +async def test_patch_ignores_unknown_body_keys(): + """Unknown keys were silently dropped by the previous construction; keep that.""" + _, update_mock = await _drive_team_write( + "patch", raw_body={"tpm_limit": 5, "not_a_team_field": "x"} + ) + written = update_mock.call_args.kwargs["data"] + + assert written["tpm_limit"] == 5 + assert "not_a_team_field" not in written + + +def test_patch_team_request_makes_team_id_optional(): + """PATCH takes team_id from the path, so the body model must not require it, + while still inheriting every UpdateTeamRequest field.""" + from litellm.proxy._types import PatchTeamRequest, UpdateTeamRequest + + parsed = PatchTeamRequest.model_validate({"tpm_limit": 5}) + + assert parsed.team_id is None + assert parsed.model_fields_set == {"tpm_limit"} + assert set(UpdateTeamRequest.model_fields).issubset(set(PatchTeamRequest.model_fields)) + + +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.""" + from litellm.proxy.proxy_server import app + + operation = app.openapi()["paths"]["/team/{team_id}"]["patch"] + schema = operation["requestBody"]["content"]["application/json"]["schema"] + + assert schema == {"$ref": "#/components/schemas/PatchTeamRequest"} + properties = app.openapi()["components"]["schemas"]["PatchTeamRequest"]["properties"] + assert "tpm_limit" in properties and "metadata" in properties diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.test.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.test.tsx new file mode 100644 index 00000000000..225c308af91 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.test.tsx @@ -0,0 +1,87 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import EditFallbacks, { Fallbacks } from "./EditFallbacks"; +import * as fetchModelsModule from "@/components/llm_calls/fetch_models"; + +vi.mock("@/components/llm_calls/fetch_models", () => ({ + fetchAvailableModels: vi.fn(), +})); + +const renderWithQueryClient = (ui: React.ReactElement) => { + const queryClient = new QueryClient({ + defaultOptions: { queries: { retry: false } }, + }); + return render({ui}); +}; + +describe("EditFallbacks", () => { + const accessToken = "test-token"; + const fallbackEntry = { "gpt-4": ["gpt-3.5-turbo", "claude-3-opus"] }; + const value: Fallbacks = [{ "gpt-4": ["gpt-3.5-turbo", "claude-3-opus"] }, { "claude-3-opus": ["gpt-4"] }]; + + const setup = (overrides: Partial> = {}) => { + const onChange = overrides.onChange ?? vi.fn().mockResolvedValue(undefined); + const onClose = overrides.onClose ?? vi.fn(); + renderWithQueryClient( + , + ); + return { onChange, onClose }; + }; + + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(fetchModelsModule.fetchAvailableModels).mockResolvedValue([ + { model_group: "gpt-4", mode: "chat" }, + { model_group: "gpt-3.5-turbo", mode: "chat" }, + { model_group: "claude-3-opus", mode: "chat" }, + { model_group: "gemini-pro", mode: "chat" }, + ]); + }); + + it("prefills the existing fallback chain for the primary model", async () => { + setup(); + await waitFor(() => { + expect(screen.getByText("gpt-3.5-turbo")).toBeInTheDocument(); + expect(screen.getByText("claude-3-opus")).toBeInTheDocument(); + }); + }); + + it("removes a fallback model and saves only the edited entry", async () => { + const user = userEvent.setup(); + const onChange = vi.fn().mockResolvedValue(undefined); + const onClose = vi.fn(); + setup({ onChange, onClose }); + + await screen.findByText("gpt-3.5-turbo"); + await user.click(screen.getByTestId("remove-fallback-gpt-3.5-turbo")); + + await user.click(screen.getByRole("button", { name: /save changes/i })); + + await waitFor(() => { + expect(onChange).toHaveBeenCalledWith([{ "gpt-4": ["claude-3-opus"] }, { "claude-3-opus": ["gpt-4"] }]); + }); + await waitFor(() => expect(onClose).toHaveBeenCalled()); + }); + + it("blocks saving with an empty fallback chain", async () => { + const user = userEvent.setup(); + const onChange = vi.fn().mockResolvedValue(undefined); + setup({ fallbackEntry: { "gpt-4": ["gpt-3.5-turbo"] }, onChange }); + + await screen.findByText("gpt-3.5-turbo"); + await user.click(screen.getByTestId("remove-fallback-gpt-3.5-turbo")); + + const saveButton = screen.getByRole("button", { name: /save changes/i }); + expect(saveButton).toBeDisabled(); + expect(onChange).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.tsx new file mode 100644 index 00000000000..938e1104301 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.tsx @@ -0,0 +1,106 @@ +/** + * Modal for editing an existing fallback entry + * Lets the user add/remove models from a primary model's fallback chain + * Reuses FallbackGroupConfig with the primary model locked + */ + +import { Button } from "antd"; +import { useQuery } from "@tanstack/react-query"; +import { Pencil } from "lucide-react"; +import React, { useMemo, useState } from "react"; +import { fetchAvailableModels } from "@/components/llm_calls/fetch_models"; +import NotificationManager from "../../../molecules/notifications_manager"; +import { AddFallbacksModal } from "./AddFallbacksModal"; +import { FallbackGroup, FallbackGroupConfig } from "./FallbackGroupConfig"; + +export type FallbackEntry = { [modelName: string]: string[] }; +export type Fallbacks = FallbackEntry[]; + +interface EditFallbacksProps { + accessToken: string; + fallbackEntry: FallbackEntry; + value: Fallbacks; + onChange: (fallbacks: Fallbacks) => Promise; + onClose: () => void; + maxFallbacks?: number; +} + +const toGroup = (entry: FallbackEntry): FallbackGroup => { + const primaryModel = Object.keys(entry)[0] ?? null; + return { + id: "edit", + primaryModel, + fallbackModels: primaryModel ? [...(entry[primaryModel] ?? [])] : [], + }; +}; + +export default function EditFallbacks({ + accessToken, + fallbackEntry, + value, + onChange, + onClose, + maxFallbacks = 10, +}: EditFallbacksProps) { + const [group, setGroup] = useState(() => toGroup(fallbackEntry)); + const [isSaving, setIsSaving] = useState(false); + + const { data: modelGroups = [] } = useQuery({ + queryKey: ["availableModels", "fallbacks"], + queryFn: () => fetchAvailableModels(accessToken), + enabled: Boolean(accessToken), + }); + + const availableModels = useMemo( + () => Array.from(new Set(modelGroups.map((option) => option.model_group))).sort(), + [modelGroups], + ); + + const handleSave = async () => { + const primaryModel = group.primaryModel; + if (!primaryModel) { + return; + } + + const updatedFallbacks = (value || []).map((entry) => + primaryModel in entry ? { ...entry, [primaryModel]: group.fallbackModels } : entry, + ); + + setIsSaving(true); + try { + await onChange(updatedFallbacks); + NotificationManager.success(`Fallbacks for ${primaryModel} updated successfully!`); + onClose(); + } catch (error) { + console.error("Error updating fallbacks:", error); + } finally { + setIsSaving(false); + } + }; + + return ( + + +
+ + +
+
+ ); +} diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx index e91c938b87a..381818a53f5 100644 --- a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx @@ -18,9 +18,16 @@ interface FallbackGroupConfigProps { onChange: (updatedGroup: FallbackGroup) => void; availableModels: string[]; maxFallbacks: number; + disablePrimaryModel?: boolean; } -export function FallbackGroupConfig({ group, onChange, availableModels, maxFallbacks }: FallbackGroupConfigProps) { +export function FallbackGroupConfig({ + group, + onChange, + availableModels, + maxFallbacks, + disablePrimaryModel = false, +}: FallbackGroupConfigProps) { // Filter available options for fallbacks (exclude primary only, allow already selected to be shown for deselection) const availableFallbackOptions = availableModels.filter((m) => m !== group.primaryModel); @@ -70,12 +77,13 @@ export function FallbackGroupConfig({ group, onChange, availableModels, maxFallb placeholder="Select primary model" value={group.primaryModel} onChange={handlePrimaryChange} + disabled={disablePrimaryModel} showSearch getPopupContainer={(trigger) => trigger.parentElement || document.body} filterOption={(input, option) => (option?.label ?? "").toLowerCase().includes(input.toLowerCase())} options={availableModels.map((m) => ({ label: m, value: m }))} /> - {!group.primaryModel && ( + {!disablePrimaryModel && !group.primaryModel && (
Select a model to begin configuring fallbacks @@ -176,6 +184,7 @@ export function FallbackGroupConfig({ group, onChange, availableModels, maxFallb