100 lines
3.6 KiB
Python
100 lines
3.6 KiB
Python
"""DAO 基类:把「取一条、取一页、软删、计数」这些重复动作收口。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from typing import Any, Generic, Sequence, TypeVar
|
||
|
||
from sqlalchemy import Select, func, select
|
||
from sqlalchemy.orm import Session
|
||
|
||
from app.core.exceptions import NotFoundError
|
||
from app.core.utils import page_count, paginate_params
|
||
from app.model.base import DEL_FLAG_NORMAL, Base
|
||
|
||
T = TypeVar("T", bound=Base)
|
||
|
||
|
||
class BaseDao(Generic[T]):
|
||
model: type[T]
|
||
|
||
# ------------------------------------------------------------ 单条
|
||
@classmethod
|
||
def get(cls, db: Session, pk: Any, with_deleted: bool = False) -> T | None:
|
||
stmt = select(cls.model).where(cls.model.id == pk)
|
||
if not with_deleted:
|
||
stmt = stmt.where(cls.model.is_del == DEL_FLAG_NORMAL)
|
||
return db.scalars(stmt).unique().first()
|
||
|
||
@classmethod
|
||
def get_or_404(cls, db: Session, pk: Any, label: str = "数据") -> T:
|
||
obj = cls.get(db, pk)
|
||
if obj is None:
|
||
raise NotFoundError(f"{label}不存在或已被删除(id={pk})")
|
||
return obj
|
||
|
||
@classmethod
|
||
def first_by(cls, db: Session, **filters: Any) -> T | None:
|
||
stmt = select(cls.model).filter_by(**filters).where(cls.model.is_del == DEL_FLAG_NORMAL)
|
||
return db.scalars(stmt).unique().first()
|
||
|
||
# ------------------------------------------------------------ 一页
|
||
@classmethod
|
||
def paginate(
|
||
cls, db: Session, stmt: Select, page: int = 1, page_size: int = 10
|
||
) -> tuple[list[T], int, int, int]:
|
||
"""返回 (items, total, page, pages),total 由子查询算,不受 limit 影响。"""
|
||
page, page_size = paginate_params(page, page_size)
|
||
count_stmt = select(func.count()).select_from(stmt.order_by(None).subquery())
|
||
total = int(db.scalar(count_stmt) or 0)
|
||
items = db.scalars(stmt.limit(page_size).offset((page - 1) * page_size)).unique().all()
|
||
return list(items), total, page, page_count(total, page_size)
|
||
|
||
@classmethod
|
||
def all(cls, db: Session, stmt: Select, limit: int | None = None) -> list[T]:
|
||
if limit:
|
||
stmt = stmt.limit(limit)
|
||
return list(db.scalars(stmt).unique().all())
|
||
|
||
@classmethod
|
||
def count(cls, db: Session, stmt: Select | None = None) -> int:
|
||
if stmt is None:
|
||
stmt = select(cls.model)
|
||
return int(db.scalar(select(func.count()).select_from(stmt.order_by(None).subquery())) or 0)
|
||
|
||
# ------------------------------------------------------------ 写
|
||
@classmethod
|
||
def add(cls, db: Session, obj: T, flush: bool = True) -> T:
|
||
db.add(obj)
|
||
if flush:
|
||
db.flush()
|
||
return obj
|
||
|
||
@classmethod
|
||
def update(cls, db: Session, obj: T, data: dict[str, Any]) -> T:
|
||
for key, value in data.items():
|
||
if value is not None and hasattr(obj, key):
|
||
setattr(obj, key, value)
|
||
db.flush()
|
||
return obj
|
||
|
||
@classmethod
|
||
def soft_delete(cls, db: Session, obj: T) -> T:
|
||
obj.soft_delete()
|
||
db.flush()
|
||
return obj
|
||
|
||
@classmethod
|
||
def exists(cls, db: Session, **filters: Any) -> bool:
|
||
stmt = select(cls.model.id).filter_by(**filters).limit(1)
|
||
return db.scalar(stmt) is not None
|
||
|
||
# ------------------------------------------------------------ 辅助
|
||
@staticmethod
|
||
def group_count(db: Session, stmt: Select) -> dict[Any, int]:
|
||
"""把 (key, count) 结果转成 dict,用于给列表批量补统计字段。"""
|
||
return {row[0]: row[1] for row in db.execute(stmt).all()}
|
||
|
||
@staticmethod
|
||
def scalar_list(db: Session, stmt: Select) -> Sequence[Any]:
|
||
return db.scalars(stmt).all()
|