Files
test/app/dao/base_dao.py
T

100 lines
3.6 KiB
Python
Raw Normal View History

2026-09-21 19:03:31 +08:00
"""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()