mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(management): boot proxy under fastapi that dropped get_flat_dependant
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
440b1bcf65
commit
c668067f4a
2 changed files with 55 additions and 3 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"]))
|
||||
Loading…
Add table
Reference in a new issue