diff --git a/litellm/proxy/management_endpoints/management_v1/common.py b/litellm/proxy/management_endpoints/management_v1/common.py index c0e7f49f2e9..caa34d75f3d 100644 --- a/litellm/proxy/management_endpoints/management_v1/common.py +++ b/litellm/proxy/management_endpoints/management_v1/common.py @@ -3,7 +3,7 @@ from urllib.parse import urlencode from fastapi import Request -from fastapi.dependencies.utils import get_flat_dependant +from fastapi.dependencies.models import Dependant from fastapi.responses import JSONResponse from litellm.types.proxy.management_endpoints.management_v1 import ( @@ -35,12 +35,19 @@ def problem_response(problem: ProblemDetail) -> JSONResponse: ) +def _flat_query_params(dependant: Dependant) -> tuple[str, ...]: + return ( + *(field.alias for field in dependant.query_params), + *(alias for dependency in dependant.dependencies for alias in _flat_query_params(dependency)), + ) + + def _declared_query_params(request: Request) -> frozenset[str]: route = request.scope.get("route") dependant = getattr(route, "dependant", None) - if dependant is None: + if not isinstance(dependant, Dependant): return frozenset() - return frozenset(field.alias for field in get_flat_dependant(dependant, skip_repeats=True).query_params) + return frozenset(_flat_query_params(dependant)) async def reject_unknown_query_params(request: Request) -> None: diff --git a/tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py b/tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py new file mode 100644 index 00000000000..1933cdf77fb --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py @@ -0,0 +1,45 @@ +from typing import Optional + +from fastapi import Depends, FastAPI, Query, Request +from fastapi.testclient import TestClient + +from litellm.proxy.management_endpoints.management_v1.common import ( + ManagementProblem, + problem_response, + reject_unknown_query_params, +) + + +def _team_scope(team_id: Optional[str] = Query(default=None)) -> Optional[str]: + return team_id + + +app = FastAPI() + + +@app.exception_handler(ManagementProblem) +async def _handle(request: Request, exc: ManagementProblem): + return problem_response(exc.problem) + + +@app.get("/things", dependencies=[Depends(reject_unknown_query_params)]) +def _things(page: int = Query(default=1), _scope: Optional[str] = Depends(_team_scope)) -> dict: + return {"ok": True} + + +client = TestClient(app) + + +def test_accepts_query_params_declared_on_the_route_and_nested_dependencies(): + """`team_id` is declared on a nested dependency, not the route signature, so a + shallow scan of the route's own params would wrongly reject it.""" + assert client.get("/things?page=2&team_id=t1").status_code == 200 + + +def test_rejects_unknown_param_and_reports_nested_declared_params_as_allowed(): + response = client.get("/things?bogus=1") + + assert response.status_code == 400 + body = response.json() + assert "bogus" in body["detail"] + assert {"page", "team_id"}.issubset(set(body["allowed"]))