mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy): accept router-wide default_litellm_params in required-param validation (#44483)
* fix(proxy): accept router-wide default_litellm_params in required-param validation The required-body-param check added in #43787 only consulted request data and deployment-level litellm_params, so a param supplied solely by router_settings.default_litellm_params (e.g. max_tokens for /v1/messages, documents for /rerank) was rejected with a 400 at route entry even though the router would have injected it during dispatch. Consult the router-wide defaults in the same place, treating None-valued defaults as absent to mirror the setdefault merge. * fix(proxy): consult only the defaults the dispatching router will merge Review follow-up: router-wide defaults are now consulted only where dispatch actually applies them. user_config requests dispatch on their own throwaway Router, so they consult that config's defaults instead of the global router's; search, managed-agent and eval routes, and model-less direct dispatch never pass through the router's defaults merge, so they keep the route-entry 400. Also adds rerank router-default coverage (unit + e2e) and drops the redundant source comment.
This commit is contained in:
parent
02b8b5cc80
commit
efe0715831
3 changed files with 292 additions and 1 deletions
|
|
@ -245,11 +245,13 @@ def _find_missing_required_body_param(
|
|||
if not missing_present_params:
|
||||
return None
|
||||
candidate_litellm_params: Final = _candidate_deployment_litellm_params(data, llm_router)
|
||||
router_default_litellm_params: Final = _router_default_litellm_params(route_type, data, llm_router)
|
||||
missing_param: Final = next(
|
||||
(
|
||||
param
|
||||
for param in missing_present_params
|
||||
if not any(deployment_params.get(param) is not None for deployment_params in candidate_litellm_params)
|
||||
if router_default_litellm_params.get(param) is None
|
||||
and not any(deployment_params.get(param) is not None for deployment_params in candidate_litellm_params)
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
|
@ -258,6 +260,32 @@ def _find_missing_required_body_param(
|
|||
return MissingBodyParam(name=missing_param, model_deployments_loaded=bool(candidate_litellm_params))
|
||||
|
||||
|
||||
_ROUTE_TYPES_WITHOUT_ROUTER_DEFAULTS_MERGE: Final[frozenset[str]] = frozenset(
|
||||
{"asearch", "acreate_agent", "acreate_eval", "acreate_run"}
|
||||
)
|
||||
|
||||
|
||||
def _router_default_litellm_params(
|
||||
route_type: str,
|
||||
data: Mapping[str, object],
|
||||
llm_router: LitellmRouter | None,
|
||||
) -> Mapping[str, object]:
|
||||
# Mirror exactly the defaults the dispatching router will merge at dispatch time:
|
||||
# user_config requests dispatch on their own throwaway Router, and the listed route
|
||||
# types (plus model-less direct dispatch) never pass through the router's merge.
|
||||
user_config: Final[Mapping[str, object] | None] = (
|
||||
data.get("user_config") if isinstance(data.get("user_config"), Mapping) else None
|
||||
)
|
||||
if user_config is not None:
|
||||
defaults: Final[Mapping[str, object] | None] = user_config.get("default_litellm_params")
|
||||
return defaults if isinstance(defaults, Mapping) else {}
|
||||
model_name: Final = data.get("model")
|
||||
if route_type in _ROUTE_TYPES_WITHOUT_ROUTER_DEFAULTS_MERGE or not isinstance(model_name, str) or not model_name:
|
||||
return {}
|
||||
router_defaults: Final[Mapping[str, object] | None] = getattr(llm_router, "default_litellm_params", None)
|
||||
return router_defaults if isinstance(router_defaults, Mapping) else {}
|
||||
|
||||
|
||||
def _candidate_deployment_litellm_params(
|
||||
data: Mapping[str, object],
|
||||
llm_router: LitellmRouter | None,
|
||||
|
|
|
|||
|
|
@ -1025,6 +1025,99 @@ def test_anthropic_messages_uses_deployment_max_tokens_default(gateway: Gateway)
|
|||
assert outbound.get("max_tokens") == 32, provider_requests
|
||||
|
||||
|
||||
def test_anthropic_messages_uses_router_wide_max_tokens_default(gateway: Gateway, tmp_path: Path) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
identity: Final = f"router-default-{uuid.uuid4().hex}"
|
||||
handle: Final = register_scenario(identity, _response("anthropic_messages"))
|
||||
scenario.cleanups.callback(delete_scenario, handle)
|
||||
config: Final = tmp_path / "router-default.yaml"
|
||||
config.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "router-default-anthropic",
|
||||
"litellm_params": {
|
||||
"model": "anthropic/claude-haiku-4-5",
|
||||
"api_base": handle.api_base(),
|
||||
"api_key": identity,
|
||||
},
|
||||
}
|
||||
],
|
||||
"router_settings": {"default_litellm_params": {"max_tokens": 32}},
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=config) as owned:
|
||||
candidate: Final = Gateway(owned.gateway.client, owned.gateway.key, gateway.upstream_url)
|
||||
response: Final = _post(
|
||||
candidate,
|
||||
"/v1/messages",
|
||||
{"model": "router-default-anthropic", "messages": [{"role": "user", "content": "router default"}]},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
_assert_scripted_response("anthropic_messages", JSON_OBJECT.validate_python(response.json()))
|
||||
observations: Final = _Observations(gateway.upstream_url)
|
||||
captured: Final = eventually(
|
||||
observations.read,
|
||||
lambda _items: any(item.get("method") == "POST" for item in observations.for_scenario(identity)),
|
||||
seconds=20,
|
||||
)
|
||||
provider_requests: Final = tuple(
|
||||
item for item in observations.for_scenario(identity) if item.get("method") == "POST"
|
||||
)
|
||||
assert len(provider_requests) == 1, captured
|
||||
outbound: Final = object_value(provider_requests[0]["body"])
|
||||
assert outbound.get("max_tokens") == 32, provider_requests
|
||||
|
||||
|
||||
def test_rerank_uses_router_wide_documents_default(gateway: Gateway, tmp_path: Path) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
identity: Final = f"router-default-rerank-{uuid.uuid4().hex}"
|
||||
handle: Final = register_scenario(identity, _response("arerank"))
|
||||
scenario.cleanups.callback(delete_scenario, handle)
|
||||
config: Final = tmp_path / "router-default-rerank.yaml"
|
||||
config.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "router-default-rerank",
|
||||
"litellm_params": {
|
||||
"model": "cohere/rerank-v4.0",
|
||||
"api_base": handle.api_base(),
|
||||
"api_key": identity,
|
||||
},
|
||||
}
|
||||
],
|
||||
"router_settings": {"default_litellm_params": {"documents": ["router default document"]}},
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=config) as owned:
|
||||
candidate: Final = Gateway(owned.gateway.client, owned.gateway.key, gateway.upstream_url)
|
||||
response: Final = _post(
|
||||
candidate,
|
||||
"/rerank",
|
||||
{"model": "router-default-rerank", "query": "which document?"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
observations: Final = _Observations(gateway.upstream_url)
|
||||
captured: Final = eventually(
|
||||
observations.read,
|
||||
lambda _items: any(item.get("method") == "POST" for item in observations.for_scenario(identity)),
|
||||
seconds=20,
|
||||
)
|
||||
provider_requests: Final = tuple(
|
||||
item for item in observations.for_scenario(identity) if item.get("method") == "POST"
|
||||
)
|
||||
assert len(provider_requests) == 1, captured
|
||||
outbound: Final = object_value(provider_requests[0]["body"])
|
||||
assert outbound.get("documents") == ["router default document"], provider_requests
|
||||
|
||||
|
||||
def test_anthropic_messages_explicit_null_reaches_upstream(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
model, identity, _handle = _register(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import pytest
|
||||
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
|
@ -1200,6 +1201,175 @@ def test_required_present_body_param_without_router_default_still_raises() -> No
|
|||
assert exc_info.value.param == "max_tokens"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type, data, param, default",
|
||||
[
|
||||
("anthropic_messages", {"model": "claude-router-default", "messages": []}, "max_tokens", 16),
|
||||
("arerank", {"model": "rerank-router-default", "query": "hi"}, "documents", ["router default document"]),
|
||||
],
|
||||
)
|
||||
def test_required_present_body_param_uses_router_wide_default(
|
||||
route_type: str, data: dict[str, object], param: str, default: object
|
||||
) -> None:
|
||||
import litellm
|
||||
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": str(data["model"]),
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_key": "test-key",
|
||||
},
|
||||
}
|
||||
],
|
||||
default_litellm_params={param: default},
|
||||
)
|
||||
|
||||
raise_if_required_body_param_missing(route_type=route_type, data=data, llm_router=router)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type, data, param",
|
||||
[
|
||||
("anthropic_messages", {"model": "claude", "messages": []}, "max_tokens"),
|
||||
("arerank", {"model": "rerank-model", "query": "hi"}, "documents"),
|
||||
],
|
||||
)
|
||||
def test_required_present_body_param_with_none_router_wide_default_still_raises(
|
||||
route_type: str, data: dict[str, object], param: str
|
||||
) -> None:
|
||||
import litellm
|
||||
from litellm.proxy.route_llm_request import (
|
||||
ProxyMissingRequiredParamError,
|
||||
raise_if_required_body_param_missing,
|
||||
)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": str(data["model"]),
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_key": "test-key",
|
||||
},
|
||||
}
|
||||
],
|
||||
default_litellm_params={param: None},
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
raise_if_required_body_param_missing(route_type=route_type, data=data, llm_router=router)
|
||||
|
||||
assert exc_info.value.param == param
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type, data, param",
|
||||
[
|
||||
("asearch", {"model": "search-model"}, "query"),
|
||||
("acreate_eval", {"model": "eval-model"}, "data_source_config"),
|
||||
("acreate_run", {"model": "eval-model"}, "data_source"),
|
||||
("acreate_agent", {"model": "agent-model"}, "name"),
|
||||
("avector_store_search", {}, "query"),
|
||||
],
|
||||
)
|
||||
def test_required_present_body_param_ignores_router_wide_default_when_dispatch_skips_router(
|
||||
route_type: str, data: dict[str, object], param: str
|
||||
) -> None:
|
||||
import litellm
|
||||
from litellm.proxy.route_llm_request import (
|
||||
ProxyMissingRequiredParamError,
|
||||
raise_if_required_body_param_missing,
|
||||
)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": str(data.get("model") or "any-model"),
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_key": "test-key",
|
||||
},
|
||||
}
|
||||
],
|
||||
default_litellm_params={param: "supplied-by-router-default"},
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
raise_if_required_body_param_missing(route_type=route_type, data=data, llm_router=router)
|
||||
|
||||
assert exc_info.value.param == param
|
||||
|
||||
|
||||
def test_required_present_body_param_uses_user_config_router_defaults() -> None:
|
||||
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
|
||||
|
||||
raise_if_required_body_param_missing(
|
||||
route_type="anthropic_messages",
|
||||
data={
|
||||
"model": "claude-user-config",
|
||||
"messages": [],
|
||||
"user_config": {"default_litellm_params": {"max_tokens": 16}},
|
||||
},
|
||||
llm_router=None,
|
||||
)
|
||||
|
||||
|
||||
def test_required_present_body_param_ignores_global_router_defaults_for_user_config_requests() -> None:
|
||||
import litellm
|
||||
from litellm.proxy.route_llm_request import (
|
||||
ProxyMissingRequiredParamError,
|
||||
raise_if_required_body_param_missing,
|
||||
)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "claude-user-config",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_key": "test-key",
|
||||
},
|
||||
}
|
||||
],
|
||||
default_litellm_params={"max_tokens": 16},
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
raise_if_required_body_param_missing(
|
||||
route_type="anthropic_messages",
|
||||
data={"model": "claude-user-config", "messages": [], "user_config": {}},
|
||||
llm_router=router,
|
||||
)
|
||||
|
||||
assert exc_info.value.param == "max_tokens"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"llm_router",
|
||||
[
|
||||
pytest.param(MagicMock(), id="magicmock-router"),
|
||||
pytest.param(SimpleNamespace(get_model_list=lambda **_kwargs: ()), id="router-without-defaults-attr"),
|
||||
],
|
||||
)
|
||||
def test_required_present_body_param_ignores_non_mapping_router_defaults(llm_router: object) -> None:
|
||||
from litellm.proxy.route_llm_request import (
|
||||
ProxyMissingRequiredParamError,
|
||||
raise_if_required_body_param_missing,
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
raise_if_required_body_param_missing(
|
||||
route_type="arerank",
|
||||
data={"model": "rerank-model", "query": "hi"},
|
||||
llm_router=llm_router, # pyright: ignore[reportArgumentType] # deliberately duck-typed routers
|
||||
)
|
||||
|
||||
assert exc_info.value.param == "documents"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type, data, param",
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue