diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 536c58df65a..58995ee3a8f 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -423,6 +423,29 @@ RouteType = Literal[ ] +# Settings that the Router accepts as per-request kwargs. These override the +# global router settings for this specific request. +_PER_REQUEST_ROUTER_SETTINGS: Final = [ + "fallbacks", + "context_window_fallbacks", + "content_policy_fallbacks", + "num_retries", + "timeout", + "model_group_retry_policy", + "routing_strategy", + "enable_tag_filtering", +] + + +def _apply_router_settings_override(data: dict, override_settings: Any) -> None: + """Merge key/team router settings into ``data`` (request values win).""" + if not isinstance(override_settings, dict): + return + for key in _PER_REQUEST_ROUTER_SETTINGS: + if key in override_settings and key not in data: + data[key] = override_settings[key] + + async def route_request( data: dict, llm_router: LitellmRouter | None, @@ -476,6 +499,12 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr data.pop("enable_tag_filtering", None) + # Always remove ``router_settings_override`` from the body so it can't leak + # to the provider, and apply its settings on every routing branch. + has_router_settings_override: Final = "router_settings_override" in data + if has_router_settings_override: + _apply_router_settings_override(data, data.pop("router_settings_override")) + team_id: Final = get_team_id_from_data(data) router_model_names: Final = llm_router.model_names if llm_router is not None else [] is_proxy_admin_without_team: Final = team_id is None and _is_proxy_admin_request(data) @@ -508,31 +537,8 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr elif "user_config" in data: return _route_user_config_request(data, route_type) - elif "router_settings_override" in data: - # Apply per-request router settings overrides from key/team config - # Instead of creating a new Router (expensive), merge settings into kwargs - # The Router already supports per-request overrides for these settings - override_settings: Final = data.pop("router_settings_override") - - # Settings that the Router accepts as per-request kwargs - # These override the global router settings for this specific request - per_request_settings: Final = [ - "fallbacks", - "context_window_fallbacks", - "content_policy_fallbacks", - "num_retries", - "timeout", - "model_group_retry_policy", - "routing_strategy", - "enable_tag_filtering", - ] - - # Merge override settings into data (only if not already set in request) - for key in per_request_settings: - if key in override_settings and key not in data: - data[key] = override_settings[key] - - # Use main router with overridden kwargs + elif has_router_settings_override: + # Settings were already merged into ``data`` above; use the main router if llm_router is not None: return getattr(llm_router, f"{route_type}")(**data) else: diff --git a/tests/test_litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py index 0b51062dd66..246dfe56825 100644 --- a/tests/test_litellm/proxy/test_route_llm_request.py +++ b/tests/test_litellm/proxy/test_route_llm_request.py @@ -1325,3 +1325,28 @@ def test_proxy_model_not_found_error_keeps_the_raw_model_only_in_the_client_resp assert raw_model in error.detail["error"] assert raw_model not in error.spend_log_error_message assert error.spend_log_error_message.startswith("/chat/completions: Invalid model name passed in") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("extra", [{"api_key": "sk-test"}, {"api_base": "http://localhost:1234"}]) +async def test_route_request_router_settings_override_stripped_with_api_key_or_base(extra): + """ + router_settings_override must not leak to the provider when the request + carries api_key/api_base, and its settings must still reach the Router. + """ + data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "Hello"}], + "router_settings_override": {"num_retries": 2}, + **extra, + } + + llm_router = MagicMock() + llm_router.acompletion.return_value = "success" + + response = await route_request(data, llm_router, None, "acompletion") + + assert response == "success" + call_kwargs = llm_router.acompletion.call_args[1] + assert "router_settings_override" not in call_kwargs + assert call_kwargs["num_retries"] == 2