80 lines
2.2 KiB
Python
80 lines
2.2 KiB
Python
from typing import Any
|
|
|
|
from sqlalchemy import func, select
|
|
from sqlalchemy.orm import Session
|
|
|
|
from .database import Base
|
|
from .errors import APIError
|
|
|
|
|
|
def list_rows(
|
|
db: Session,
|
|
model: type[Base],
|
|
*,
|
|
limit: int,
|
|
offset: int,
|
|
filters: dict[str, Any] | None = None,
|
|
order_by: Any | None = None,
|
|
) -> list[Base]:
|
|
stmt = select(model)
|
|
for field, value in (filters or {}).items():
|
|
if value is not None:
|
|
stmt = stmt.where(getattr(model, field) == value)
|
|
if order_by is not None:
|
|
stmt = stmt.order_by(order_by)
|
|
stmt = stmt.limit(limit).offset(offset)
|
|
return list(db.execute(stmt).scalars().all())
|
|
|
|
|
|
def count_rows(db: Session, model: type[Base], filters: dict[str, Any] | None = None) -> int:
|
|
stmt = select(func.count()).select_from(model)
|
|
for field, value in (filters or {}).items():
|
|
if value is not None:
|
|
stmt = stmt.where(getattr(model, field) == value)
|
|
return int(db.execute(stmt).scalar_one())
|
|
|
|
|
|
def get_row_or_404(db: Session, model: type[Base], pk: Any, pk_field: str = "id") -> Base:
|
|
obj = db.execute(select(model).where(getattr(model, pk_field) == pk)).scalar_one_or_none()
|
|
if obj is None:
|
|
raise APIError(status_code=404, detail=f"{model.__tablename__} '{pk}' not found")
|
|
return obj
|
|
|
|
|
|
def create_row(db: Session, model: type[Base], data: dict[str, Any]) -> Base:
|
|
obj = model(**data)
|
|
db.add(obj)
|
|
_commit(db)
|
|
db.refresh(obj)
|
|
return obj
|
|
|
|
|
|
def update_row(db: Session, obj: Base, data: dict[str, Any]) -> Base:
|
|
for field, value in data.items():
|
|
setattr(obj, field, value)
|
|
_commit(db)
|
|
db.refresh(obj)
|
|
return obj
|
|
|
|
|
|
def delete_row(db: Session, obj: Base) -> None:
|
|
db.delete(obj)
|
|
_commit(db)
|
|
|
|
|
|
def _commit(db: Session) -> None:
|
|
from sqlalchemy.exc import IntegrityError
|
|
|
|
try:
|
|
db.commit()
|
|
except IntegrityError as exc:
|
|
db.rollback()
|
|
raise APIError(status_code=409, detail=_integrity_message(exc)) from exc
|
|
|
|
|
|
def _integrity_message(exc: Exception) -> str:
|
|
msg = str(getattr(exc, "orig", exc))
|
|
if "Duplicate entry" in msg:
|
|
return "Duplicate entry: a record with these unique values already exists"
|
|
return "Database integrity error"
|