Files
2026-09-21 19:03:31 +08:00

100 lines
3.6 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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()