From e2e51f055d9a66fd363dd9e56dd32d180e8539a6 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Wed, 22 Jul 2026 16:31:03 -0700 Subject: [PATCH 1/3] refactor(proxy): type the PATCH /team/{team_id} request body (#34195) * feat(ui): add react-hook-form + zod form infrastructure Introduce the shared form layer the dashboard's antd forms will migrate onto, with no user-visible change yet. - pin react-hook-form, @hookform/resolvers, and zod (kept on 3.25.76 and imported via the zod/v4 entrypoint so openai's optional zod ^3 peer still resolves and npm ci stays clean) - vendor the base-vega Field family into components/shared/form as forwardRef components on the repo's cva.config, since base-vega ships no form primitive and its field source imports class-variance-authority and is React 19 style - add a FormField bridge that binds a react-hook-form Controller to the Field layer and wires label, description, and error ids into aria attributes - add pickDirty, which narrows a submitted body to the top-level keys the user actually touched so a partial update stops re-sending untouched fields pickDirty reads dirtiness at the top level because react-hook-form tracks it per leaf, so an edited array arrives as [true, false] and a cleared list as an empty array that still carries its default-length dirty markers; the falsy clear tokens (null, [], {}, 0, false) all survive. Tests cover the Field primitives, the FormField aria wiring against a live zod resolver, and pickDirty both as a unit and driven through a real react-hook-form instance. * test(ui): lock pickDirty behavior on a pure field-array reorder react-hook-form compares each array element to its default positionally by value, so useFieldArray move/swap and a reordered scalar array all mark the moved indices dirty and pickDirty sends the whole array; a swap of two equal elements is a value-level no-op and is correctly omitted. Covers the reorder case a review flagged as untested. * feat(proxy): publish a typed request body for PATCH /team/{team_id} The route validated its body into UpdateTeamRequest but read it off the raw request, so the OpenAPI spec carried no requestBody and the dashboard's generated client could not type the call at all. - add PatchTeamRequest, UpdateTeamRequest with an optional team_id, since PATCH takes the id from the path; a body team_id is still accepted when it matches - validate the body through PatchTeamRequest before delegating to update_team - declare the request body on the route and regenerate schema.d.ts The handler keeps reading the raw body rather than declaring a typed parameter. FastAPI validates a declared body before the handler runs, which would replace the 400 for a non-object body with a 422 and move absent-vs-null out of reach of the RFC 7386 metadata merge; those are pinned by existing tests, so the schema is declared on the route instead and every error path is unchanged. Validation is shape-preserving: the body is dumped with exclude_unset so an omitted field never reaches the write, an explicit null still clears, and a partial object_permission does not gain sibling sub-keys, which would wipe them given the column merges rather than replaces. Tests extend the existing patch harness rather than replacing it. * refactor(proxy): declare the PATCH /team/{team_id} body as a typed parameter Replaces the hand-written OpenAPI declaration added earlier in this branch. The route now takes data: PatchTeamRequest, so FastAPI generates the request body itself and emits a $ref to the model instead of an inlined copy that would go stale as fields are added. The earlier approach was a workaround built on a wrong premise. Declaring the body does not cost absent-vs-null: model_fields_set preserves it, which is how POST /team/update already gets its tri-state, and a nested null inside metadata survives validation untouched, so the RFC 7386 merge is unaffected. The one real change is the status code for a malformed body. The route answered 400 for a non-object body and 500 for a wrongly typed field, reporting a caller mistake as a server fault; both are now 422, matching POST /team/update and the other typed management endpoints. The two tests that pinned the old parse-level errors are replaced by one that pins the 422 through the ASGI stack, and the handler drops its manual parsing entirely. --- litellm/proxy/_types.py | 11 ++ .../management_endpoints/team_endpoints.py | 33 ++-- .../test_team_endpoints.py | 151 +++++++++++++++--- ui/litellm-dashboard/src/lib/http/schema.d.ts | 113 ++++++++++++- 4 files changed, 262 insertions(+), 46 deletions(-) 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/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/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index d109fc1783d..e612bddaca1 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -13860,11 +13860,12 @@ export interface paths { * Patch Team * @description 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' --header 'Authorization: Bearer sk-1234' --header 'Content-Type: application/json' --data-raw '{ @@ -28807,6 +28808,102 @@ export interface components { litellm_params?: components["schemas"]["PromptLiteLLMParams"] | null; prompt_info?: components["schemas"]["PromptInfo"] | null; }; + /** + * PatchTeamRequest + * @description 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. + */ + PatchTeamRequest: { + /** Access Group Ids */ + access_group_ids?: string[] | null; + /** Allowed Passthrough Routes */ + allowed_passthrough_routes?: unknown[] | null; + /** Allowed Vector Store Indexes */ + allowed_vector_store_indexes?: components["schemas"]["AllowedVectorStoreIndexItem"][] | null; + /** Blocked */ + blocked?: boolean | null; + /** Budget Duration */ + budget_duration?: string | null; + /** Budget Limits */ + budget_limits?: components["schemas"]["BudgetLimitEntry"][] | null; + /** Default Team Member Models */ + default_team_member_models?: string[] | null; + /** Disable Global Guardrails */ + disable_global_guardrails?: boolean | null; + /** Enforced Batch Output Expires After */ + enforced_batch_output_expires_after?: { + [key: string]: unknown; + } | null; + /** Enforced File Expires After */ + enforced_file_expires_after?: { + [key: string]: unknown; + } | null; + /** Guardrails */ + guardrails?: string[] | null; + /** Max Budget */ + max_budget?: number | null; + /** Mcp Rpm Limit */ + mcp_rpm_limit?: { + [key: string]: number; + } | null; + /** Metadata */ + metadata?: { + [key: string]: unknown; + } | null; + /** Model Aliases */ + model_aliases?: { + [key: string]: unknown; + } | null; + /** Model Rpm Limit */ + model_rpm_limit?: { + [key: string]: number; + } | null; + /** Model Tpm Limit */ + model_tpm_limit?: { + [key: string]: number; + } | null; + /** Models */ + models?: unknown[] | null; + object_permission?: components["schemas"]["LiteLLM_ObjectPermissionBase"] | null; + /** Organization Id */ + organization_id?: string | null; + /** Policies */ + policies?: string[] | null; + /** Prompts */ + prompts?: string[] | null; + /** Router Settings */ + router_settings?: { + [key: string]: unknown; + } | null; + /** Rpm Limit */ + rpm_limit?: number | null; + /** Secret Manager Settings */ + secret_manager_settings?: { + [key: string]: unknown; + } | null; + /** Soft Budget */ + soft_budget?: number | null; + /** Tags */ + tags?: unknown[] | null; + /** Team Alias */ + team_alias?: string | null; + /** Team Id */ + team_id?: string | null; + /** Team Member Budget */ + team_member_budget?: number | null; + /** Team Member Budget Duration */ + team_member_budget_duration?: string | null; + /** Team Member Key Duration */ + team_member_key_duration?: string | null; + /** Team Member Rpm Limit */ + team_member_rpm_limit?: number | null; + /** Team Member Tpm Limit */ + team_member_tpm_limit?: number | null; + /** Tpm Limit */ + tpm_limit?: number | null; + }; /** * PerTestingCriteriaResult * @description Results for a specific testing criteria @@ -50613,7 +50710,11 @@ export interface operations { }; cookie?: never; }; - requestBody?: never; + requestBody: { + content: { + "application/json": components["schemas"]["PatchTeamRequest"]; + }; + }; responses: { /** @description Successful Response */ 200: { From 1ae406953cfd254b5fcbf601a15ff4f8cdd2babf Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 22 Jul 2026 16:31:33 -0700 Subject: [PATCH 2/3] feat(ui): edit fallback chains from router settings (#32841) * feat(ui): edit fallback chains from router settings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(ui): address review nits on edit fallbacks modal Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(ui): fetch models via react-query in edit fallbacks modal Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Mubashir Osmani Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../Fallbacks/EditFallbacks.test.tsx | 87 ++++++++++++++ .../Fallbacks/EditFallbacks.tsx | 106 ++++++++++++++++++ .../Fallbacks/FallbackGroupConfig.tsx | 13 ++- .../Fallbacks/Fallbacks.test.tsx | 56 ++++++--- .../RouterSettings/Fallbacks/Fallbacks.tsx | 34 +++++- 5 files changed, 278 insertions(+), 18 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.test.tsx create mode 100644 ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.tsx 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