diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 961087ef445..2f5c38b0131 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1376,7 +1376,7 @@ async def generate_key_fn( - auto_rotate: Optional[bool] - Whether this key should be automatically rotated (regenerated) - rotation_interval: Optional[str] - How often to auto-rotate this key (e.g., '30s', '30m', '30h', '30d'). Required if auto_rotate=True. - allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint. - - router_settings: Optional[UpdateRouterConfig] - key-specific router settings. Example - {"model_group_retry_policy": {"max_retries": 5}}. IF null or {} then no router settings. + - router_settings: Optional[UpdateRouterConfig] - key-specific router settings. Example - {"model_group_retry_policy": {"gpt-4": {"RateLimitErrorRetries": 5}}}. IF null or {} then no router settings. - access_group_ids: Optional[List[str]] - List of access group IDs to associate with the key. Access groups define which models a key can access. Example - ["access_group_1", "access_group_2"]. - budget_limits: Optional[list] - List of concurrent budget windows for the key. Each window specifies a budget_limit, time_period, and optional budget_duration. Example - [{"budget_limit": 10.0, "time_period": "1d"}, {"budget_limit": 50.0, "time_period": "7d"}]. @@ -2376,7 +2376,7 @@ async def update_key_fn( - auto_rotate: Optional[bool] - Whether this key should be automatically rotated - rotation_interval: Optional[str] - How often to rotate this key (e.g., '30d', '90d'). Required if auto_rotate=True - allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint. - - router_settings: Optional[UpdateRouterConfig] - key-specific router settings. Example - {"model_group_retry_policy": {"max_retries": 5}}. IF null or {} then no router settings. + - router_settings: Optional[UpdateRouterConfig] - key-specific router settings. Example - {"model_group_retry_policy": {"gpt-4": {"RateLimitErrorRetries": 5}}}. IF null or {} then no router settings. - access_group_ids: Optional[List[str]] - List of access group IDs to associate with the key. Access groups define which models a key can access. Example - ["access_group_1", "access_group_2"]. - budget_limits: Optional[list] - List of concurrent budget windows for the key. Each window specifies a budget_limit, time_period, and optional budget_duration. Example - [{"budget_limit": 10.0, "time_period": "1d"}, {"budget_limit": 50.0, "time_period": "7d"}]. diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index fb835535af8..0ea2b9e05f9 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -949,7 +949,7 @@ async def new_team( - allowed_passthrough_routes: Optional[List[str]] - List of allowed pass through routes for the team. - allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint. - secret_manager_settings: Optional[dict] - Secret manager settings for the team. [Docs](https://docs.litellm.ai/docs/secret_managers/overview) - - router_settings: Optional[UpdateRouterConfig] - team-specific router settings. Example - {"model_group_retry_policy": {"max_retries": 5}}. IF null or {} then no router settings. + - router_settings: Optional[UpdateRouterConfig] - team-specific router settings. Example - {"model_group_retry_policy": {"gpt-4": {"RateLimitErrorRetries": 5}}}. IF null or {} then no router settings. - access_group_ids: Optional[List[str]] - List of access group IDs to associate with the team. Access groups define which models the team can access. Example - ["access_group_1", "access_group_2"]. - enforced_file_expires_after: Optional[dict] - Enforced file expiration policy for the team. Keys created under this team will inherit this policy for file uploads. Example - {"anchor": "created_at", "days": 30}. - enforced_batch_output_expires_after: Optional[dict] - Enforced batch output file expiration policy for the team. Keys created under this team will inherit this policy for batch output files. Example - {"anchor": "created_at", "days": 30}. @@ -1608,7 +1608,7 @@ async def update_team( Example - update team TPM Limit - allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint. - secret_manager_settings: Optional[dict] - Secret manager settings for the team. [Docs](https://docs.litellm.ai/docs/secret_managers/overview) - - router_settings: Optional[UpdateRouterConfig] - team-specific router settings. Example - {"model_group_retry_policy": {"max_retries": 5}}. IF null or {} then no router settings. + - router_settings: Optional[UpdateRouterConfig] - team-specific router settings. Example - {"model_group_retry_policy": {"gpt-4": {"RateLimitErrorRetries": 5}}}. IF null or {} then no router settings. - access_group_ids: Optional[List[str]] - List of access group IDs to associate with the team. Access groups define which models the team can access. Example - ["access_group_1", "access_group_2"]. - enforced_file_expires_after: Optional[dict] - Enforced file expiration policy for the team. Keys created under this team will inherit this policy for file uploads. Example - {"anchor": "created_at", "days": 30}. - enforced_batch_output_expires_after: Optional[dict] - Enforced batch output file expiration policy for the team. Keys created under this team will inherit this policy for batch output files. Example - {"anchor": "created_at", "days": 30}. diff --git a/litellm/router.py b/litellm/router.py index 1d23e838563..cdb45de66db 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -9588,6 +9588,7 @@ class Router: "retry_after", "fallbacks", "context_window_fallbacks", + "retry_policy", "model_group_retry_policy", "model_group_alias", "enable_weighted_failover", @@ -9612,6 +9613,12 @@ class Router: elif var == "routing_groups": self._routing_groups_input = kwargs[var] rebuild_routing_groups = True + elif var == "retry_policy": + value = kwargs[var] + if isinstance(value, dict): + value = RetryPolicy(**value) + if value is None or isinstance(value, RetryPolicy): + setattr(self, var, value) else: value = kwargs[var] # only run routing strategy init if it has changed diff --git a/litellm/types/router.py b/litellm/types/router.py index e74bd4ceeca..0c3485deae7 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -81,6 +81,22 @@ class RouterConfig(BaseModel): model_config = ConfigDict(protected_namespaces=()) +class RetryPolicy(BaseModel): + """ + Use this to set a custom number of retries per exception type + If RateLimitErrorRetries = 3, then 3 retries will be made for RateLimitError + Mapping of Exception type to number of retries + https://docs.litellm.ai/docs/exception_mapping + """ + + BadRequestErrorRetries: Optional[int] = None + AuthenticationErrorRetries: Optional[int] = None + TimeoutErrorRetries: Optional[int] = None + RateLimitErrorRetries: Optional[int] = None + ContentPolicyViolationErrorRetries: Optional[int] = None + InternalServerErrorRetries: Optional[int] = None + + class UpdateRouterConfig(BaseModel): """ Set of params that you can modify via `router.update_settings()`. @@ -89,7 +105,8 @@ class UpdateRouterConfig(BaseModel): routing_strategy_args: Optional[dict] = None routing_strategy: Optional[str] = None routing_groups: Optional[List[RoutingGroup]] = None - model_group_retry_policy: Optional[dict] = None + retry_policy: Optional[RetryPolicy] = None + model_group_retry_policy: Optional[Dict[str, RetryPolicy]] = None model_group_affinity_config: Optional[Dict[str, List[str]]] = None allowed_fails: Optional[int] = None cooldown_time: Optional[float] = None @@ -495,22 +512,6 @@ class AllowedFailsPolicy(BaseModel): InternalServerErrorAllowedFails: Optional[int] = None -class RetryPolicy(BaseModel): - """ - Use this to set a custom number of retries per exception type - If RateLimitErrorRetries = 3, then 3 retries will be made for RateLimitError - Mapping of Exception type to number of retries - https://docs.litellm.ai/docs/exception_mapping - """ - - BadRequestErrorRetries: Optional[int] = None - AuthenticationErrorRetries: Optional[int] = None - TimeoutErrorRetries: Optional[int] = None - RateLimitErrorRetries: Optional[int] = None - ContentPolicyViolationErrorRetries: Optional[int] = None - InternalServerErrorRetries: Optional[int] = None - - class AlertingConfig(BaseModel): """ Use this configure alerting for the router. Receive alerts on the following events diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 02c9a249ae8..5c68cdc32dc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -5890,13 +5890,12 @@ async def test_generate_key_with_router_settings(monkeypatch): generate_key_fn, ) - # Test router_settings with sample data - # Using valid UpdateRouterConfig fields (retry_policy is not a valid field, - # but model_group_retry_policy is, which also tests nested dict serialization) + # model_group_retry_policy maps a model group to a RetryPolicy, exercising + # nested-model serialization through the key record router_settings_data = { "routing_strategy": "usage-based", "num_retries": 3, - "model_group_retry_policy": {"max_retries": 5}, + "model_group_retry_policy": {"gpt-4": {"RateLimitErrorRetries": 5}}, } request_data = GenerateKeyRequest( diff --git a/tests/test_litellm/test_router_retry_policy_update.py b/tests/test_litellm/test_router_retry_policy_update.py new file mode 100644 index 00000000000..450391fd503 --- /dev/null +++ b/tests/test_litellm/test_router_retry_policy_update.py @@ -0,0 +1,279 @@ +""" +Tests for the retry_policy fix on Router.update_settings (LIT-3152). + +Bug: the Admin UI Model Retry Settings tab posts a global ``retry_policy`` +through ``POST /config/update`` -> ``UpdateRouterConfig`` -> +``Router.update_settings``. Both the pydantic schema and +``update_settings`` were dropping the field silently: + +- ``UpdateRouterConfig`` had no ``retry_policy`` attribute, so + ``model_dump(exclude_none=True)`` returned ``{}`` for that key. +- ``Router.update_settings`` had no ``"retry_policy"`` entry in + ``_allowed_settings``, so even when fed directly the call was a no-op + (``Setting {} is not allowed`` debug log). + +The net effect was that after saving retry counts in the UI and +reloading, every value snapped back to ``defaultRetry = num_retries`` +(2 by default), exactly matching the ticket repro. + +This file pins both halves of the fix. +""" + +import json +import os +import sys +from dataclasses import dataclass +from unittest.mock import AsyncMock, MagicMock + +import pytest +from pydantic import ValidationError + +sys.path.insert(0, os.path.abspath("../..")) + +import litellm +from litellm.types.router import RetryPolicy, UpdateRouterConfig + +# --------------------------------------------------------------------------- +# UpdateRouterConfig schema membership (LIT-3152 part 1) +# --------------------------------------------------------------------------- + + +def test_update_router_config_exposes_retry_policy_field(): + """retry_policy must be a declared field on UpdateRouterConfig. + + Without it, Pydantic silently strips the key from the /config/update + payload before the proxy even calls llm_router.update_settings. + """ + assert "retry_policy" in UpdateRouterConfig.model_fields + + +def test_update_router_config_accepts_retry_policy_payload(): + """The exact payload the Admin UI Model Retry Settings tab sends must + round-trip through the schema's ``dict(exclude_none=True)`` form, since + that is what /config/update writes to the LiteLLM_Config row.""" + payload = { + "retry_policy": { + "BadRequestErrorRetries": 5, + "RateLimitErrorRetries": 7, + "TimeoutErrorRetries": 3, + } + } + cfg = UpdateRouterConfig(**payload) + dumped = cfg.model_dump(exclude_none=True) + assert "retry_policy" in dumped + assert dumped["retry_policy"]["BadRequestErrorRetries"] == 5 + assert dumped["retry_policy"]["RateLimitErrorRetries"] == 7 + assert dumped["retry_policy"]["TimeoutErrorRetries"] == 3 + + +def test_update_router_config_rejects_malformed_retry_policy(): + """The field is typed as RetryPolicy, so /config/update validates the + payload at the boundary and rejects non-numeric counts with a 422 instead + of silently persisting garbage the apply path would later have to drop.""" + with pytest.raises(ValidationError): + UpdateRouterConfig(retry_policy={"BadRequestErrorRetries": "not-an-int"}) + + +def test_update_router_config_rejects_malformed_model_group_retry_policy(): + """model_group_retry_policy is Dict[str, RetryPolicy], so each per-group + policy is validated the same way.""" + with pytest.raises(ValidationError): + UpdateRouterConfig( + model_group_retry_policy={"gpt-4": {"RateLimitErrorRetries": "x"}} + ) + + +# --------------------------------------------------------------------------- +# Router.update_settings retry_policy path (LIT-3152 part 2) +# --------------------------------------------------------------------------- + + +def _build_router() -> litellm.Router: + return litellm.Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": { + "model": "openai/gpt-4", + "api_key": "sk-fake", + "api_base": "http://localhost:9999", + }, + } + ] + ) + + +def test_update_settings_persists_retry_policy_dict(): + """When the proxy's ``_add_router_settings_from_db_config`` calls + ``llm_router.update_settings(retry_policy={...})`` after reading the + DB row, the dict must land on ``self.retry_policy`` as a typed + ``RetryPolicy`` (mirroring ``Router.__init__`` semantics).""" + router = _build_router() + assert router.retry_policy is None # baseline + + router.update_settings( + retry_policy={ + "BadRequestErrorRetries": 5, + "RateLimitErrorRetries": 7, + "TimeoutErrorRetries": 3, + } + ) + + assert isinstance(router.retry_policy, RetryPolicy) + assert router.retry_policy.BadRequestErrorRetries == 5 + assert router.retry_policy.RateLimitErrorRetries == 7 + assert router.retry_policy.TimeoutErrorRetries == 3 + + +def test_update_settings_accepts_retry_policy_object_unchanged(): + """A pre-built ``RetryPolicy`` instance must pass through verbatim so + callers that already constructed one (e.g. tests or programmatic + callers) keep working.""" + router = _build_router() + + policy = RetryPolicy(BadRequestErrorRetries=2) + router.update_settings(retry_policy=policy) + + assert router.retry_policy is policy + + +def test_update_settings_ignores_malformed_retry_policy(): + """A non-dict, non-``RetryPolicy`` value (e.g. a YAML typo like + ``retry_policy: 5`` reaching ``update_settings``) must not land on + ``self.retry_policy``. ``Router.__init__`` already drops such inputs; + the update path must match so a malformed config can't store garbage + that ``get_num_retries_from_retry_policy`` would only choke on at + request time.""" + router = _build_router() + + existing = RetryPolicy(BadRequestErrorRetries=4) + router.update_settings(retry_policy=existing) + assert router.retry_policy is existing + + for bad_value in (5, "RateLimitErrorRetries=7", ["BadRequestErrorRetries"]): + router.update_settings(retry_policy=bad_value) + assert router.retry_policy is existing + + +def test_update_settings_get_settings_round_trip_for_retry_policy(): + """``GET /get/config/callbacks`` serializes ``llm_router.get_settings()`` + back to the UI. After updating, the round-trip must reflect the new + values rather than the pre-update sentinel.""" + router = _build_router() + pre = router.get_settings().get("retry_policy") + assert pre is None + + router.update_settings( + retry_policy={ + "BadRequestErrorRetries": 5, + "RateLimitErrorRetries": 7, + } + ) + post = router.get_settings().get("retry_policy") + assert post is not None + assert post.BadRequestErrorRetries == 5 + assert post.RateLimitErrorRetries == 7 + + +def test_update_settings_unrelated_kwargs_still_skipped(): + """Regression guard: the new branch must not relax the + ``_allowed_settings`` allowlist for unrelated keys. An unknown + setting should still be dropped silently as before.""" + router = _build_router() + router.update_settings(this_is_not_a_router_setting=123) + assert not hasattr(router, "this_is_not_a_router_setting") + + +# --------------------------------------------------------------------------- +# End-to-end persist -> apply -> read-back (LIT-3152 part 3) +# +# The tests above pin each layer in isolation, so they would all still pass +# if a regression flipped ``ConfigYAML.router_settings`` back to a loose +# ``dict`` (silently dropping retry_policy on the DB write) or if +# ``_add_router_settings_from_db_config`` stopped pushing the stored row onto +# the live router. This drives the real handler chain an Admin UI save +# triggers — ``POST /config/update`` writes the LiteLLM_Config row, +# ``add_deployment`` applies it to ``llm_router``, and +# ``GET /get/config/callbacks`` serializes it back — so the round trip is +# pinned, not just the pieces. +# --------------------------------------------------------------------------- + + +@dataclass(frozen=True, slots=True) +class _FakeConfigRow: + param_name: str + param_value: dict + + +class _FakeConfigTable: + """Stand-in for ``prisma_client.db.litellm_config``. + + Reproduces the one behavior the apply path relies on: a value written as a + JSON string by ``/config/update`` reads back as a parsed dict, which is + what ``_add_router_settings_from_db_config``'s + ``isinstance(param_value, dict)`` branch requires to forward the settings. + """ + + def __init__(self): + self.rows = {} + + async def find_first(self, where): + return self.rows.get(where["param_name"]) + + async def upsert(self, where, data): + name = where["param_name"] + raw = (data["update"] if name in self.rows else data["create"])["param_value"] + self.rows[name] = _FakeConfigRow(name, json.loads(raw) if isinstance(raw, str) else raw) + + +@pytest.mark.asyncio +async def test_config_update_persists_and_reads_back_retry_policy(monkeypatch): + """The exact global retry_policy save the UI performs must survive the + real ``/config/update`` -> DB -> apply -> ``/get/config/callbacks`` path, + not snap back to the ``num_retries`` fallback the ticket reported.""" + import litellm.proxy.proxy_server as proxy_server + from litellm.proxy._types import ConfigYAML, LitellmUserRoles, UserAPIKeyAuth + + router = _build_router() + assert router.retry_policy is None + + fake_table = _FakeConfigTable() + prisma_client = MagicMock() + prisma_client.db.litellm_config = fake_table + + async def _apply_router_settings(*args, **kwargs): + await proxy_server.proxy_config._add_router_settings_from_db_config( + config_data={}, llm_router=router, prisma_client=prisma_client + ) + + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server.proxy_config, "add_deployment", _apply_router_settings) + monkeypatch.setattr(proxy_server.proxy_config, "get_config", AsyncMock(return_value={})) + + posted = UpdateRouterConfig( + retry_policy=RetryPolicy( + BadRequestErrorRetries=5, + TimeoutErrorRetries=3, + RateLimitErrorRetries=7, + ) + ) + await proxy_server.update_config( + config_info=ConfigYAML(router_settings=posted), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234"), + ) + + persisted = fake_table.rows["router_settings"].param_value["retry_policy"] + assert persisted == { + "BadRequestErrorRetries": 5, + "TimeoutErrorRetries": 3, + "RateLimitErrorRetries": 7, + } + + assert isinstance(router.retry_policy, RetryPolicy) + assert router.retry_policy.RateLimitErrorRetries == 7 + + read_back = (await proxy_server.get_config())["router_settings"]["retry_policy"] + assert read_back.BadRequestErrorRetries == 5 + assert read_back.TimeoutErrorRetries == 3 + assert read_back.RateLimitErrorRetries == 7 diff --git a/ui/litellm-dashboard/eslint-metrics.json b/ui/litellm-dashboard/eslint-metrics.json index 28f9fc3a6af..01c5a241562 100644 --- a/ui/litellm-dashboard/eslint-metrics.json +++ b/ui/litellm-dashboard/eslint-metrics.json @@ -1,5 +1,5 @@ { - "@typescript-eslint/no-explicit-any": 2027, + "@typescript-eslint/no-explicit-any": 2026, "complexity": 128, "max-depth": 61 } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/routerSettings/useUpdateRetryPolicy.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/routerSettings/useUpdateRetryPolicy.ts new file mode 100644 index 00000000000..f01e71cab6c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/routerSettings/useUpdateRetryPolicy.ts @@ -0,0 +1,17 @@ +import { setCallbacksCall } from "@/components/networking"; +import { useMutation } from "@tanstack/react-query"; + +export interface RetryPolicyPayload { + retry_policy?: Record | null; + model_group_retry_policy?: Record | undefined> | null; +} + +export const useUpdateRetryPolicy = (accessToken: string | null) => + useMutation({ + mutationFn: async (policy: RetryPolicyPayload) => { + if (!accessToken) { + throw new Error("Access token is required"); + } + return setCallbacksCall(accessToken, { router_settings: policy }); + }, + }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx index 2f8f7350db9..85793ab7d15 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx @@ -2,13 +2,14 @@ import { useCredentials } from "@/app/(dashboard)/hooks/credentials/useCredentia import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap"; import { useModelsInfo } from "@/app/(dashboard)/hooks/models/useModels"; import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; +import { useUpdateRetryPolicy } from "@/app/(dashboard)/hooks/routerSettings/useUpdateRetryPolicy"; import AllModelsTab from "@/app/(dashboard)/models-and-endpoints/components/AllModelsTab"; import ModelRetrySettingsTab from "@/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab"; import PriceDataManagementTab from "@/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab"; import { handleAddModelSubmit } from "@/components/add_model/handle_add_model_submit"; import { Team } from "@/components/key_team_helpers/key_list"; import CredentialsPanel from "@/components/model_add/credentials"; -import { getCallbacksCall, setCallbacksCall } from "@/components/networking"; +import { getCallbacksCall } from "@/components/networking"; import { Providers, getPlaceholder, getProviderModels } from "@/components/provider_info_helpers"; import { getDisplayModelName } from "@/components/view_model/model_name_display"; import { transformModelData } from "./utils/modelDataTransformer"; @@ -19,7 +20,7 @@ import { Col, Grid, Icon, Tab, TabGroup, TabList, TabPanel, TabPanels } from "@t import type { UploadProps } from "antd"; import { Form } from "antd"; import { PlusCircleOutlined } from "@ant-design/icons"; -import React, { useEffect, useMemo, useState } from "react"; +import React, { useCallback, useEffect, useMemo, useState } from "react"; import AddModelTab from "../../../components/add_model/add_model_tab"; import HealthCheckComponent from "../../../components/model_dashboard/HealthCheckComponent"; import ModelGroupAliasSettings from "../../../components/model_group_alias_settings"; @@ -42,6 +43,13 @@ interface GlobalRetryPolicyObject { [retryPolicyKey: string]: number; } +interface RouterSettings { + model_group_retry_policy?: RetryPolicyObject | null; + retry_policy?: GlobalRetryPolicyObject | null; + num_retries?: number | null; + model_group_alias?: { [key: string]: string } | null; +} + const HEALTH_PAGE_SIZE = 50; const ModelsAndEndpointsView: React.FC = ({ premiumUser, teams }) => { @@ -52,6 +60,7 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te const [selectedProvider, setSelectedProvider] = useState(Providers.Anthropic); const [selectedModelGroup, setSelectedModelGroup] = useState(null); + const [retryScope, setRetryScope] = useState("global"); const [modelGroupRetryPolicy, setModelGroupRetryPolicy] = useState(null); const [globalRetryPolicy, setGlobalRetryPolicy] = useState(null); const [defaultRetry, setDefaultRetry] = useState(0); @@ -78,6 +87,7 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te const { data: credentialsResponse, isLoading: isLoadingCredentials } = useCredentials(); const credentialsList = credentialsResponse?.credentials || []; const { data: uiSettings, isLoading: isLoadingUISettings } = useUISettings(); + const updateRetryPolicy = useUpdateRetryPolicy(accessToken); const availableModelGroups = useMemo(() => { if (!modelDataResponse?.data) return []; @@ -189,61 +199,66 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te refetchModels(); }; - const handleSaveRetrySettings = async () => { - if (!accessToken) { - return; + const fetchRouterSettings = useCallback(async (): Promise => { + if (!accessToken || !userID || !userRole) { + return null; } - try { - const payload: any = { - router_settings: {}, - }; - - if (selectedModelGroup === "global") { - if (globalRetryPolicy) { - payload.router_settings.retry_policy = globalRetryPolicy; - } - NotificationsManager.success("Global retry settings saved successfully"); - } else { - if (modelGroupRetryPolicy) { - payload.router_settings.model_group_retry_policy = modelGroupRetryPolicy; - } - NotificationsManager.success(`Retry settings saved successfully for ${selectedModelGroup}`); - } - - await setCallbacksCall(accessToken, payload); + const routerSettingsInfo = await getCallbacksCall(accessToken, userID, userRole); + return routerSettingsInfo.router_settings; } catch (error) { - NotificationsManager.fromBackend("Failed to save retry settings"); + console.error("Error fetching model data:", error); + return null; } + }, [accessToken, userID, userRole]); + + const applyRouterSettings = useCallback((routerSettings: RouterSettings) => { + setModelGroupRetryPolicy(routerSettings.model_group_retry_policy ?? null); + setGlobalRetryPolicy(routerSettings.retry_policy ?? null); + setDefaultRetry(routerSettings.num_retries ?? 2); + setModelGroupAlias(routerSettings.model_group_alias || {}); + }, []); + + const loadRetrySettings = useCallback(async () => { + const routerSettings = await fetchRouterSettings(); + if (routerSettings) { + applyRouterSettings(routerSettings); + } + }, [fetchRouterSettings, applyRouterSettings]); + + const handleSaveRetrySettings = () => { + updateRetryPolicy.mutate( + { + retry_policy: globalRetryPolicy, + model_group_retry_policy: modelGroupRetryPolicy, + }, + { + onSuccess: () => { + NotificationsManager.success("Retry settings saved successfully"); + loadRetrySettings(); + }, + onError: () => { + NotificationsManager.fromBackend("Failed to save retry settings"); + }, + }, + ); }; useEffect(() => { if (!accessToken || !token || !userRole || !userID || !modelDataResponse) { return; } - const fetchData = async () => { - try { - const routerSettingsInfo = await getCallbacksCall(accessToken, userID, userRole); - let router_settings = routerSettingsInfo.router_settings; - - let model_group_retry_policy = router_settings.model_group_retry_policy; - let default_retries = router_settings.num_retries; - - setModelGroupRetryPolicy(model_group_retry_policy); - setGlobalRetryPolicy(router_settings.retry_policy); - setDefaultRetry(default_retries); - - const model_group_alias = router_settings.model_group_alias || {}; - setModelGroupAlias(model_group_alias); - } catch (error) { - console.error("Error fetching model data:", error); + let active = true; + void (async () => { + const routerSettings = await fetchRouterSettings(); + if (active && routerSettings) { + applyRouterSettings(routerSettings); } + })(); + return () => { + active = false; }; - - if (accessToken && token && userRole && userID && modelDataResponse) { - fetchData(); - } - }, [accessToken, token, userRole, userID, modelDataResponse]); + }, [accessToken, token, userRole, userID, modelDataResponse, fetchRouterSettings, applyRouterSettings]); const isLoading = isLoadingModels || isLoadingModelCostMap || isLoadingCredentials || isLoadingUISettings; @@ -483,8 +498,8 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te panel: ( = ({ premiumUser, te modelGroupRetryPolicy={modelGroupRetryPolicy} setModelGroupRetryPolicy={setModelGroupRetryPolicy} handleSaveRetrySettings={handleSaveRetrySettings} + isSaving={updateRetryPolicy.isPending} /> ), }, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.test.tsx index 6f89a41034b..13862094ba1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.test.tsx @@ -15,7 +15,7 @@ vi.mock("@tremor/react", async (importOriginal) => { return { ...actual, TabPanel: ({ children }: { children: React.ReactNode }) => React.createElement("div", null, children), - Button: React.forwardRef(({ children, ...props }, ref) => + Button: React.forwardRef(({ children, variant, size, loading, ...props }, ref) => React.createElement("button", { ...props, ref }, children), ), Tooltip: ({ children }: { children?: React.ReactNode }) => React.createElement(React.Fragment, null, children), @@ -90,7 +90,7 @@ describe("ModelRetrySettingsTab", () => { expect(inputs[0]).toHaveValue("0"); }); - it("should fall back to globalRetryPolicy when no model-specific value is set (model scope)", () => { + it("should leave model-scope rows empty with the inherited value as placeholder when there is no override", () => { const globalRetryPolicy: GlobalRetryPolicy = { TimeoutErrorRetries: 7, }; @@ -105,12 +105,45 @@ describe("ModelRetrySettingsTab", () => { />, ); - // The TimeoutError row is 3rd (index 2) + // No override exists, so the input is empty and the inherited value is only + // a placeholder -- this is what keeps 0 ("zero retries") distinct from + // "inherit the global value". const inputs = screen.getAllByRole("spinbutton"); - expect(inputs[2]).toHaveValue("7"); + expect(inputs[2]).toHaveValue(""); // TimeoutError row + expect(inputs[2]).toHaveAttribute("placeholder", "7"); // inherited from global + expect(inputs[0]).toHaveValue(""); // BadRequestError row (no global) + expect(inputs[0]).toHaveAttribute("placeholder", "1"); // inherited from defaultRetry + }); - // Rows without a global value fall back to defaultRetry - expect(inputs[0]).toHaveValue("1"); + it("should clear a model-group override when Reset is clicked", async () => { + const user = userEvent.setup(); + const setModelGroupRetryPolicy = vi.fn(); + render( + , + ); + + // Reset only renders for rows that actually have an override + const resetButtons = screen.getAllByRole("button", { name: /reset/i }); + expect(resetButtons).toHaveLength(1); + + await user.click(resetButtons[0]); + + const updater = setModelGroupRetryPolicy.mock.calls.at(-1)![0]; + const result = updater({ "gpt-4": { BadRequestErrorRetries: 5 } }); + expect(result["gpt-4"]).not.toHaveProperty("BadRequestErrorRetries"); + }); + + it("should disable the Save button while a save is in flight", () => { + render(); + + expect(screen.getByRole("button", { name: /save/i })).toBeDisabled(); }); it("should prefer model-specific retry count over the global value (model scope)", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.tsx index 753aeedce54..5ff761b0663 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.tsx @@ -20,6 +20,7 @@ interface ModelRetrySettingsTabProps { modelGroupRetryPolicy: RetryPolicyObject | null; setModelGroupRetryPolicy: React.Dispatch>; handleSaveRetrySettings: () => void; + isSaving?: boolean; } const retryPolicyMap: Record = { @@ -41,8 +42,26 @@ const ModelRetrySettingsTab = ({ modelGroupRetryPolicy, setModelGroupRetryPolicy, handleSaveRetrySettings, + isSaving = false, }: ModelRetrySettingsTabProps) => { - // const [modelGroupRetryPolicy, setModelGroupRetryPolicy] = useState(null); + const isGlobalScope = selectedModelGroup === "global"; + + const setGlobalValue = (retryPolicyKey: string, value: number | null) => { + if (value == null) return; + setGlobalRetryPolicy((prev) => ({ ...(prev ?? {}), [retryPolicyKey]: value })); + }; + + const setModelOverride = (retryPolicyKey: string, value: number | null) => { + setModelGroupRetryPolicy((prev) => { + const groupPolicy = { ...(prev?.[selectedModelGroup!] ?? {}) }; + if (value == null) { + delete groupPolicy[retryPolicyKey]; + } else { + groupPolicy[retryPolicyKey] = value; + } + return { ...(prev ?? {}), [selectedModelGroup!]: groupPolicy }; + }); + }; return ( @@ -51,13 +70,12 @@ const ModelRetrySettingsTab = ({ Retry Policy Scope: