Files

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"