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()
|