"""班级 DAO(需求 2.4)。""" from __future__ import annotations from sqlalchemy import Select, func, or_, select from sqlalchemy.orm import Session from app.core.utils import age_expression # noqa: F401 (保留给后续扩展) from app.dao.base_dao import BaseDao from app.model import Clazz, Student, class_teachers class ClazzDao(BaseDao[Clazz]): model = Clazz @classmethod def build_stmt( cls, keyword: str | None = None, status: int | None = None, advisor_id: int | None = None, head_teacher_id: int | None = None, teacher_id: int | None = None, order_by: str = "id", order: str = "desc", ) -> Select: stmt = select(Clazz).where(Clazz.alive()) if keyword: like = f"%{keyword.strip()}%" stmt = stmt.where( or_(Clazz.name.like(like), Clazz.class_no.like(like), Clazz.direction.like(like)) ) if status: stmt = stmt.where(Clazz.status == status) if advisor_id: stmt = stmt.where(Clazz.advisor_id == advisor_id) if head_teacher_id: stmt = stmt.where(Clazz.head_teacher_id == head_teacher_id) if teacher_id: stmt = stmt.where( Clazz.id.in_( select(class_teachers.c.class_id).where(class_teachers.c.teacher_id == teacher_id) ) ) sortable = { "id": Clazz.id, "class_no": Clazz.class_no, "name": Clazz.name, "open_date": Clazz.open_date, "close_date": Clazz.close_date, "status": Clazz.status, "capacity": Clazz.capacity, } column = sortable.get(order_by or "id", Clazz.id) stmt = stmt.order_by(column.desc() if (order or "desc").lower() == "desc" else column.asc()) return stmt @classmethod def get_by_class_no(cls, db: Session, class_no: str, with_deleted: bool = False) -> Clazz | None: stmt = select(Clazz).where(Clazz.class_no == class_no) if not with_deleted: stmt = stmt.where(Clazz.alive()) return db.scalars(stmt).unique().first() @classmethod def name_map(cls, db: Session) -> dict[int, str]: stmt = select(Clazz.id, Clazz.name).where(Clazz.alive()) return {row[0]: row[1] for row in db.execute(stmt).all()} @classmethod def no_map(cls, db: Session) -> dict[int, str]: stmt = select(Clazz.id, Clazz.class_no).where(Clazz.alive()) return {row[0]: row[1] for row in db.execute(stmt).all()} @classmethod def next_class_no(cls, db: Session, direction_prefix: str, year: int) -> str: """班级编号规则:方向缩写 + 年份 + 两位序号,如 JAVA202601。""" prefix = f"{direction_prefix}{year}" stmt = select(func.max(Clazz.class_no)).where(Clazz.class_no.like(f"{prefix}%")) last = db.scalar(stmt) seq = int(last[len(prefix):]) + 1 if last and last[len(prefix):].isdigit() else 1 return f"{prefix}{seq:02d}" @classmethod def student_counts(cls, db: Session, only_alive: bool = True) -> dict[int, int]: stmt = select(Student.class_id, func.count(Student.id)).where(Student.class_id.is_not(None)) if only_alive: stmt = stmt.where(Student.alive()) return {row[0]: row[1] for row in db.execute(stmt.group_by(Student.class_id)).all()}