fix(api): wrap pre-paginated list responses
This commit is contained in:
+42
-20
@@ -201,18 +201,39 @@ def _with_header_project_context(endpoint, route_name: str):
|
|||||||
|
|
||||||
def _with_pagination(endpoint):
|
def _with_pagination(endpoint):
|
||||||
signature = inspect.signature(endpoint)
|
signature = inspect.signature(endpoint)
|
||||||
if "limit" in signature.parameters or "offset" in signature.parameters:
|
handler_limit_parameter = "limit" if "limit" in signature.parameters else None
|
||||||
return endpoint
|
handler_offset_parameter = next(
|
||||||
|
(
|
||||||
|
parameter_name
|
||||||
|
for parameter_name in ("offset", "skip")
|
||||||
|
if parameter_name in signature.parameters
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
handler_handles_pagination = bool(
|
||||||
|
handler_limit_parameter or handler_offset_parameter
|
||||||
|
)
|
||||||
|
|
||||||
@wraps(endpoint)
|
@wraps(endpoint)
|
||||||
async def wrapper(*args, **kwargs):
|
async def wrapper(*args, **kwargs):
|
||||||
limit = kwargs.pop("_rest_limit")
|
if handler_handles_pagination:
|
||||||
offset = kwargs.pop("_rest_offset")
|
limit = kwargs.get(handler_limit_parameter, 0)
|
||||||
|
offset = kwargs.get(handler_offset_parameter, 0)
|
||||||
|
else:
|
||||||
|
limit = kwargs.pop("_rest_limit")
|
||||||
|
offset = kwargs.pop("_rest_offset")
|
||||||
result = endpoint(*args, **kwargs)
|
result = endpoint(*args, **kwargs)
|
||||||
if inspect.isawaitable(result):
|
if inspect.isawaitable(result):
|
||||||
result = await result
|
result = await result
|
||||||
if not isinstance(result, list):
|
if not isinstance(result, list):
|
||||||
return result
|
return result
|
||||||
|
if handler_handles_pagination:
|
||||||
|
return Page(
|
||||||
|
items=result,
|
||||||
|
total=offset + len(result),
|
||||||
|
limit=limit or len(result),
|
||||||
|
offset=offset,
|
||||||
|
)
|
||||||
return Page(
|
return Page(
|
||||||
items=result[offset : offset + limit],
|
items=result[offset : offset + limit],
|
||||||
total=len(result),
|
total=len(result),
|
||||||
@@ -221,22 +242,23 @@ def _with_pagination(endpoint):
|
|||||||
)
|
)
|
||||||
|
|
||||||
parameters = list(signature.parameters.values())
|
parameters = list(signature.parameters.values())
|
||||||
parameters.extend(
|
if not handler_handles_pagination:
|
||||||
[
|
parameters.extend(
|
||||||
inspect.Parameter(
|
[
|
||||||
"_rest_limit",
|
inspect.Parameter(
|
||||||
kind=inspect.Parameter.KEYWORD_ONLY,
|
"_rest_limit",
|
||||||
annotation=int,
|
kind=inspect.Parameter.KEYWORD_ONLY,
|
||||||
default=Query(100, ge=1, le=1000, alias="limit"),
|
annotation=int,
|
||||||
),
|
default=Query(100, ge=1, le=1000, alias="limit"),
|
||||||
inspect.Parameter(
|
),
|
||||||
"_rest_offset",
|
inspect.Parameter(
|
||||||
kind=inspect.Parameter.KEYWORD_ONLY,
|
"_rest_offset",
|
||||||
annotation=int,
|
kind=inspect.Parameter.KEYWORD_ONLY,
|
||||||
default=Query(0, ge=0, alias="offset"),
|
annotation=int,
|
||||||
),
|
default=Query(0, ge=0, alias="offset"),
|
||||||
]
|
),
|
||||||
)
|
]
|
||||||
|
)
|
||||||
wrapper.__signature__ = signature.replace(parameters=parameters)
|
wrapper.__signature__ = signature.replace(parameters=parameters)
|
||||||
return wrapper
|
return wrapper
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from pathlib import Path
|
|||||||
from uuid import uuid4
|
from uuid import uuid4
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from fastapi import FastAPI
|
from fastapi import APIRouter, FastAPI, Query
|
||||||
from fastapi.routing import APIRoute
|
from fastapi.routing import APIRoute
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
@@ -254,3 +254,31 @@ def test_rest_runtime_consumes_injected_project_context(monkeypatch) -> None:
|
|||||||
assert response.json()["items"] == [
|
assert response.json()["items"] == [
|
||||||
{"scheme_name": "burst_case", "scheme_type": "burst_analysis"}
|
{"scheme_name": "burst_case", "scheme_type": "burst_analysis"}
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_rest_runtime_wraps_handler_paginated_list() -> None:
|
||||||
|
source_router = APIRouter()
|
||||||
|
|
||||||
|
@source_router.get("/records", response_model=list[int])
|
||||||
|
async def list_records(
|
||||||
|
skip: int = Query(0, ge=0),
|
||||||
|
limit: int = Query(2, ge=1, le=10),
|
||||||
|
) -> list[int]:
|
||||||
|
records = [10, 20, 30, 40]
|
||||||
|
return records[skip : skip + limit]
|
||||||
|
|
||||||
|
app = FastAPI(redirect_slashes=False)
|
||||||
|
app.include_router(build_rest_router(source_router.routes), prefix="/api/v1")
|
||||||
|
|
||||||
|
response = TestClient(app, raise_server_exceptions=False).get(
|
||||||
|
"/api/v1/records",
|
||||||
|
params={"skip": 1, "limit": 2},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert response.json() == {
|
||||||
|
"items": [20, 30],
|
||||||
|
"total": 3,
|
||||||
|
"limit": 2,
|
||||||
|
"offset": 1,
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user