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