feat: add sorting and ordering functionality to paginated presets list
This commit is contained in:
parent
670744890c
commit
980f39d52c
4 changed files with 232 additions and 22 deletions
6
API.md
6
API.md
|
|
@ -1771,6 +1771,8 @@ Binary image data with appropriate headers
|
||||||
**Query Parameters**:
|
**Query Parameters**:
|
||||||
- `page` (optional): Page number (1-indexed). Default: `1`.
|
- `page` (optional): Page number (1-indexed). Default: `1`.
|
||||||
- `per_page` (optional): Items per page. Default: `config.default_pagination`.
|
- `per_page` (optional): Items per page. Default: `config.default_pagination`.
|
||||||
|
- `sort` (optional): Comma-separated sort fields. Accepted values: `id`, `name`, `priority`, `default`, `created_at`, `updated_at`. Default: `priority,name`.
|
||||||
|
- `order` (optional): Comma-separated sort directions matching `sort`, or a single direction applied to every requested sort field. Accepted values: `asc`, `desc`. Default: `desc,asc`.
|
||||||
|
|
||||||
**Response**:
|
**Response**:
|
||||||
```json
|
```json
|
||||||
|
|
@ -1800,6 +1802,10 @@ Binary image data with appropriate headers
|
||||||
|
|
||||||
**Notes**:
|
**Notes**:
|
||||||
- `default: true` indicates this is a system default preset (cannot be modified or deleted)
|
- `default: true` indicates this is a system default preset (cannot be modified or deleted)
|
||||||
|
- Default ordering remains `priority desc, name asc`
|
||||||
|
|
||||||
|
**Error Responses**:
|
||||||
|
- `400 Bad Request` - Invalid pagination or sorting query parameters
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -16,7 +16,8 @@ from app.library.Services import Services
|
||||||
from app.library.Singleton import Singleton
|
from app.library.Singleton import Singleton
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from collections.abc import AsyncGenerator
|
from collections.abc import Callable
|
||||||
|
from contextlib import AbstractAsyncContextManager
|
||||||
|
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
from sqlalchemy.engine.result import Result
|
from sqlalchemy.engine.result import Result
|
||||||
|
|
@ -24,13 +25,34 @@ if TYPE_CHECKING:
|
||||||
from sqlalchemy.sql.elements import ColumnElement
|
from sqlalchemy.sql.elements import ColumnElement
|
||||||
from sqlalchemy.sql.selectable import Select
|
from sqlalchemy.sql.selectable import Select
|
||||||
|
|
||||||
|
SessionFactory = Callable[[], AbstractAsyncContextManager[AsyncSession]]
|
||||||
|
|
||||||
LOG: logging.Logger = logging.getLogger(__name__)
|
LOG: logging.Logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class PresetsRepository(metaclass=Singleton):
|
class PresetsRepository(metaclass=Singleton):
|
||||||
def __init__(self, session: AsyncGenerator[AsyncSession] | None = None) -> None:
|
SORT_FIELDS: dict[str, Any] = {
|
||||||
|
"id": PresetModel.id,
|
||||||
|
"name": PresetModel.name,
|
||||||
|
"priority": PresetModel.priority,
|
||||||
|
"default": PresetModel.default,
|
||||||
|
"created_at": PresetModel.created_at,
|
||||||
|
"updated_at": PresetModel.updated_at,
|
||||||
|
}
|
||||||
|
SORT_DIRECTIONS: tuple[str, str] = ("asc", "desc")
|
||||||
|
DEFAULT_SORT_ORDER: tuple[tuple[str, str], ...] = (("priority", "desc"), ("name", "asc"))
|
||||||
|
FIELD_DEFAULT_DIRECTIONS: dict[str, str] = {
|
||||||
|
"id": "asc",
|
||||||
|
"name": "asc",
|
||||||
|
"priority": "desc",
|
||||||
|
"default": "desc",
|
||||||
|
"created_at": "desc",
|
||||||
|
"updated_at": "desc",
|
||||||
|
}
|
||||||
|
|
||||||
|
def __init__(self, session: SessionFactory | None = None) -> None:
|
||||||
self._migrated = False
|
self._migrated = False
|
||||||
self.session: AsyncGenerator[AsyncSession] = session or get_session
|
self.session: SessionFactory = session or get_session
|
||||||
|
|
||||||
async def run_migrations(self) -> None:
|
async def run_migrations(self) -> None:
|
||||||
if self._migrated:
|
if self._migrated:
|
||||||
|
|
@ -82,7 +104,85 @@ class PresetsRepository(metaclass=Singleton):
|
||||||
)
|
)
|
||||||
return list(result.scalars().all())
|
return list(result.scalars().all())
|
||||||
|
|
||||||
async def list_paginated(self, page: int, per_page: int) -> tuple[list[PresetModel], int, int, int]:
|
@classmethod
|
||||||
|
def parse_sorting(cls, sort: str | None = None, order: str | None = None) -> tuple[tuple[str, str], ...]:
|
||||||
|
sort_value = (sort or "").strip()
|
||||||
|
order_value = (order or "").strip()
|
||||||
|
|
||||||
|
if not sort_value and not order_value:
|
||||||
|
return cls.DEFAULT_SORT_ORDER
|
||||||
|
|
||||||
|
if order_value and not sort_value:
|
||||||
|
msg = "sort is required when order is provided."
|
||||||
|
raise ValueError(msg)
|
||||||
|
|
||||||
|
fields: list[str] = []
|
||||||
|
for raw_field in sort_value.split(","):
|
||||||
|
field = raw_field.strip().lower()
|
||||||
|
if not field:
|
||||||
|
continue
|
||||||
|
if field not in cls.SORT_FIELDS:
|
||||||
|
msg = f"sort must use supported fields: {', '.join(cls.SORT_FIELDS)}."
|
||||||
|
raise ValueError(msg)
|
||||||
|
if field not in fields:
|
||||||
|
fields.append(field)
|
||||||
|
|
||||||
|
if not fields:
|
||||||
|
msg = "sort must include at least one field."
|
||||||
|
raise ValueError(msg)
|
||||||
|
|
||||||
|
directions: list[str] = []
|
||||||
|
if order_value:
|
||||||
|
for raw_direction in order_value.split(","):
|
||||||
|
direction = raw_direction.strip().lower()
|
||||||
|
if not direction:
|
||||||
|
continue
|
||||||
|
if direction not in cls.SORT_DIRECTIONS:
|
||||||
|
msg = "order must be 'asc' or 'desc'."
|
||||||
|
raise ValueError(msg)
|
||||||
|
directions.append(direction)
|
||||||
|
|
||||||
|
if not directions:
|
||||||
|
msg = "order must include at least one direction."
|
||||||
|
raise ValueError(msg)
|
||||||
|
|
||||||
|
if len(directions) not in {1, len(fields)}:
|
||||||
|
msg = "order must provide one direction or match the number of sort fields."
|
||||||
|
raise ValueError(msg)
|
||||||
|
|
||||||
|
if not directions:
|
||||||
|
return tuple((field, cls.FIELD_DEFAULT_DIRECTIONS.get(field, "asc")) for field in fields)
|
||||||
|
|
||||||
|
if len(directions) == 1:
|
||||||
|
return tuple((field, directions[0]) for field in fields)
|
||||||
|
|
||||||
|
return tuple((field, directions[index]) for index, field in enumerate(fields))
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _apply_sort_direction(cls, field: str, direction: str) -> Any:
|
||||||
|
column = cls.SORT_FIELDS[field]
|
||||||
|
return column.asc() if direction == "asc" else column.desc()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _build_order_by(cls, sort: str | None = None, order: str | None = None) -> list[Any]:
|
||||||
|
sorting = cls.parse_sorting(sort, order)
|
||||||
|
|
||||||
|
order_by: list[Any] = [cls._apply_sort_direction(field, direction) for field, direction in sorting]
|
||||||
|
|
||||||
|
if all(field != "id" for field, _ in sorting):
|
||||||
|
order_by.append(PresetModel.id.asc())
|
||||||
|
|
||||||
|
return order_by
|
||||||
|
|
||||||
|
async def list_paginated(
|
||||||
|
self,
|
||||||
|
page: int,
|
||||||
|
per_page: int,
|
||||||
|
sort: str | None = None,
|
||||||
|
order: str | None = None,
|
||||||
|
) -> tuple[list[PresetModel], int, int, int]:
|
||||||
|
order_by = self._build_order_by(sort, order)
|
||||||
|
|
||||||
async with self.session() as session:
|
async with self.session() as session:
|
||||||
total: int = await self.count()
|
total: int = await self.count()
|
||||||
total_pages: int = (total + per_page - 1) // per_page if total > 0 else 1
|
total_pages: int = (total + per_page - 1) // per_page if total > 0 else 1
|
||||||
|
|
@ -91,10 +191,7 @@ class PresetsRepository(metaclass=Singleton):
|
||||||
page = total_pages
|
page = total_pages
|
||||||
|
|
||||||
query: Select[tuple[PresetModel]] = (
|
query: Select[tuple[PresetModel]] = (
|
||||||
select(PresetModel)
|
select(PresetModel).order_by(*order_by).limit(per_page).offset((page - 1) * per_page)
|
||||||
.order_by(PresetModel.priority.desc(), PresetModel.name.asc())
|
|
||||||
.limit(per_page)
|
|
||||||
.offset((page - 1) * per_page)
|
|
||||||
)
|
)
|
||||||
result: Result[tuple[PresetModel]] = await session.execute(query)
|
result: Result[tuple[PresetModel]] = await session.execute(query)
|
||||||
return list(result.scalars().all()), total, page, total_pages
|
return list(result.scalars().all()), total, page, total_pages
|
||||||
|
|
@ -136,9 +233,25 @@ class PresetsRepository(metaclass=Singleton):
|
||||||
|
|
||||||
async def create(self, payload: PresetModel | dict) -> PresetModel:
|
async def create(self, payload: PresetModel | dict) -> PresetModel:
|
||||||
async with self.session() as session:
|
async with self.session() as session:
|
||||||
model: PresetModel = PresetModel(**payload) if isinstance(payload, dict) else payload
|
data: dict[str, Any]
|
||||||
if model.id is not None:
|
if isinstance(payload, dict):
|
||||||
model.id = None
|
data = dict(payload)
|
||||||
|
else:
|
||||||
|
data = {
|
||||||
|
"name": payload.name,
|
||||||
|
"description": payload.description,
|
||||||
|
"folder": payload.folder,
|
||||||
|
"template": payload.template,
|
||||||
|
"cookies": payload.cookies,
|
||||||
|
"cli": payload.cli,
|
||||||
|
"default": payload.default,
|
||||||
|
"priority": payload.priority,
|
||||||
|
"created_at": payload.created_at,
|
||||||
|
"updated_at": payload.updated_at,
|
||||||
|
}
|
||||||
|
|
||||||
|
data.pop("id", None)
|
||||||
|
model = PresetModel(**data)
|
||||||
|
|
||||||
model.name = preset_name(model.name)
|
model.name = preset_name(model.name)
|
||||||
|
|
||||||
|
|
@ -163,12 +276,10 @@ class PresetsRepository(metaclass=Singleton):
|
||||||
result: Result[tuple[PresetModel]] = await session.execute(select(PresetModel).where(clause).limit(1))
|
result: Result[tuple[PresetModel]] = await session.execute(select(PresetModel).where(clause).limit(1))
|
||||||
model: PresetModel | None = result.scalar_one_or_none()
|
model: PresetModel | None = result.scalar_one_or_none()
|
||||||
|
|
||||||
if None is model:
|
if model is None:
|
||||||
msg: str = f"Preset '{identifier}' not found."
|
msg: str = f"Preset '{identifier}' not found."
|
||||||
raise KeyError(msg)
|
raise KeyError(msg)
|
||||||
|
|
||||||
assert None is not model
|
|
||||||
|
|
||||||
payload.pop("id", None)
|
payload.pop("id", None)
|
||||||
payload.pop("created_at", None)
|
payload.pop("created_at", None)
|
||||||
payload.pop("updated_at", None)
|
payload.pop("updated_at", None)
|
||||||
|
|
@ -177,9 +288,10 @@ class PresetsRepository(metaclass=Singleton):
|
||||||
if hasattr(model, key):
|
if hasattr(model, key):
|
||||||
setattr(model, key, value)
|
setattr(model, key, value)
|
||||||
|
|
||||||
model.name = preset_name(model.name)
|
normalized_name = preset_name(model.name)
|
||||||
if await self.get_by_name(name=model.name, exclude_id=model.id) is not None:
|
model.name = normalized_name
|
||||||
msg = f"Preset with name '{model.name}' already exists."
|
if await self.get_by_name(name=normalized_name, exclude_id=model.id) is not None:
|
||||||
|
msg = f"Preset with name '{normalized_name}' already exists."
|
||||||
raise ValueError(msg)
|
raise ValueError(msg)
|
||||||
|
|
||||||
await session.commit()
|
await session.commit()
|
||||||
|
|
@ -196,12 +308,11 @@ class PresetsRepository(metaclass=Singleton):
|
||||||
clause = PresetModel.name == preset_name(identifier)
|
clause = PresetModel.name == preset_name(identifier)
|
||||||
|
|
||||||
result: Result[tuple[PresetModel]] = await session.execute(select(PresetModel).where(clause).limit(1))
|
result: Result[tuple[PresetModel]] = await session.execute(select(PresetModel).where(clause).limit(1))
|
||||||
if not (model := result.scalar_one_or_none()):
|
model: PresetModel | None = result.scalar_one_or_none()
|
||||||
|
if model is None:
|
||||||
msg: str = f"Preset '{identifier}' not found."
|
msg: str = f"Preset '{identifier}' not found."
|
||||||
raise KeyError(msg)
|
raise KeyError(msg)
|
||||||
|
|
||||||
assert None is not model
|
|
||||||
|
|
||||||
await session.delete(model)
|
await session.delete(model)
|
||||||
await session.commit()
|
await session.commit()
|
||||||
return model
|
return model
|
||||||
|
|
|
||||||
|
|
@ -23,8 +23,17 @@ def _serialize(model: Any) -> dict:
|
||||||
|
|
||||||
@route("GET", "api/presets/", "presets")
|
@route("GET", "api/presets/", "presets")
|
||||||
async def presets_list(request: Request, encoder: Encoder, repo: PresetsRepository) -> Response:
|
async def presets_list(request: Request, encoder: Encoder, repo: PresetsRepository) -> Response:
|
||||||
page, per_page = normalize_pagination(request)
|
try:
|
||||||
items, total, current_page, total_pages = await repo.list_paginated(page, per_page)
|
page, per_page = normalize_pagination(request)
|
||||||
|
items, total, current_page, total_pages = await repo.list_paginated(
|
||||||
|
page,
|
||||||
|
per_page,
|
||||||
|
sort=request.query.get("sort"),
|
||||||
|
order=request.query.get("order"),
|
||||||
|
)
|
||||||
|
except ValueError as exc:
|
||||||
|
return web.json_response(data={"error": str(exc)}, status=web.HTTPBadRequest.status_code)
|
||||||
|
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
data=PresetList(
|
data=PresetList(
|
||||||
items=[_model(model) for model in items],
|
items=[_model(model) for model in items],
|
||||||
|
|
|
||||||
|
|
@ -1,9 +1,16 @@
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import pytest_asyncio
|
import pytest_asyncio
|
||||||
|
from aiohttp import web
|
||||||
|
from aiohttp.web import Request
|
||||||
|
|
||||||
from app.features.presets.repository import PresetsRepository
|
from app.features.presets.repository import PresetsRepository
|
||||||
|
from app.features.presets.router import presets_list
|
||||||
|
from app.library.encoder import Encoder
|
||||||
from app.library.sqlite_store import SqliteStore
|
from app.library.sqlite_store import SqliteStore
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -69,3 +76,80 @@ class TestPresetsRepository:
|
||||||
assert total == 5, "Should report total count"
|
assert total == 5, "Should report total count"
|
||||||
assert page == 1, "Should be on page 1"
|
assert page == 1, "Should be on page 1"
|
||||||
assert total_pages == 3, "Should have 3 pages total"
|
assert total_pages == 3, "Should have 3 pages total"
|
||||||
|
assert [item.priority for item in items] == [4, 3], "Should keep default priority-desc order"
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_list_paginated_sorts_by_name_desc(self, repo):
|
||||||
|
await repo.create({"name": "Alpha", "priority": 1})
|
||||||
|
await repo.create({"name": "Gamma", "priority": 3})
|
||||||
|
await repo.create({"name": "Beta", "priority": 2})
|
||||||
|
|
||||||
|
items, _, _, _ = await repo.list_paginated(page=1, per_page=10, sort="name", order="desc")
|
||||||
|
|
||||||
|
assert [item.name for item in items] == ["gamma", "beta", "alpha"], "Should sort by requested field"
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_list_paginated_supports_multiple_sort_fields(self, repo):
|
||||||
|
await repo.create({"name": "Charlie", "priority": 2})
|
||||||
|
await repo.create({"name": "Alpha", "priority": 1})
|
||||||
|
await repo.create({"name": "Bravo", "priority": 1})
|
||||||
|
|
||||||
|
items, _, _, _ = await repo.list_paginated(page=1, per_page=10, sort="priority,name", order="asc,desc")
|
||||||
|
|
||||||
|
assert [(item.priority, item.name) for item in items] == [
|
||||||
|
(1, "bravo"),
|
||||||
|
(1, "alpha"),
|
||||||
|
(2, "charlie"),
|
||||||
|
], "Should support multiple sort fields and directions"
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_list_paginated_rejects_invalid_sort_field(self, repo):
|
||||||
|
with pytest.raises(ValueError, match="sort must use supported fields"):
|
||||||
|
await repo.list_paginated(page=1, per_page=10, sort="cli", order="asc")
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_list_paginated_rejects_invalid_sort_direction(self, repo):
|
||||||
|
with pytest.raises(ValueError, match="order must be 'asc' or 'desc'"):
|
||||||
|
await repo.list_paginated(page=1, per_page=10, sort="name", order="sideways")
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_list_paginated_rejects_mismatched_sort_and_order_lengths(self, repo):
|
||||||
|
with pytest.raises(ValueError, match="order must provide one direction or match the number of sort fields"):
|
||||||
|
await repo.list_paginated(page=1, per_page=10, sort="priority,name", order="asc,desc,asc")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
class TestPresetRoutes:
|
||||||
|
async def test_list_route_supports_sort_params(self, repo):
|
||||||
|
await repo.create({"name": "Alpha", "priority": 1})
|
||||||
|
await repo.create({"name": "Bravo", "priority": 1})
|
||||||
|
await repo.create({"name": "Charlie", "priority": 2})
|
||||||
|
|
||||||
|
request = MagicMock(spec=Request)
|
||||||
|
request.query = {"page": "1", "per_page": "10", "sort": "priority,name", "order": "asc,desc"}
|
||||||
|
|
||||||
|
response = await presets_list(request, Encoder(), repo)
|
||||||
|
payload = json.loads(response.text)
|
||||||
|
|
||||||
|
assert response.status == web.HTTPOk.status_code, "Should return 200 for valid sorting"
|
||||||
|
assert [item["name"] for item in payload["items"]] == ["bravo", "alpha", "charlie"], "Should sort response"
|
||||||
|
|
||||||
|
async def test_list_route_rejects_invalid_sort_field(self, repo):
|
||||||
|
request = MagicMock(spec=Request)
|
||||||
|
request.query = {"sort": "cli", "order": "asc"}
|
||||||
|
|
||||||
|
response = await presets_list(request, Encoder(), repo)
|
||||||
|
payload = json.loads(response.text)
|
||||||
|
|
||||||
|
assert response.status == web.HTTPBadRequest.status_code, "Should reject unsupported sort field"
|
||||||
|
assert "sort" in payload["error"], "Should explain invalid sort field"
|
||||||
|
|
||||||
|
async def test_list_route_rejects_invalid_sort_direction(self, repo):
|
||||||
|
request = MagicMock(spec=Request)
|
||||||
|
request.query = {"sort": "name", "order": "sideways"}
|
||||||
|
|
||||||
|
response = await presets_list(request, Encoder(), repo)
|
||||||
|
payload = json.loads(response.text)
|
||||||
|
|
||||||
|
assert response.status == web.HTTPBadRequest.status_code, "Should reject unsupported sort direction"
|
||||||
|
assert "order" in payload["error"], "Should explain invalid sort direction"
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue