"""学生 DAO(需求 2.1 的查询部分)。""" from __future__ import annotations from datetime import date from sqlalchemy import Select, func, or_, select from sqlalchemy.orm import Session from app.core.utils import age_expression from app.dao.base_dao import BaseDao from app.model import Clazz, Score, Student class StudentDao(BaseDao[Student]): model = Student # ------------------------------------------------------------ 查询条件 @classmethod def build_stmt( cls, keyword: str | None = None, class_id: int | None = None, class_no: str | None = None, status: int | None = None, gender: int | None = None, advisor_id: int | None = None, education: str | None = None, age_min: int | None = None, age_max: int | None = None, order_by: str = "id", order: str = "desc", ) -> Select: stmt = select(Student).where(Student.alive()) if keyword: like = f"%{keyword.strip()}%" stmt = stmt.where( or_( Student.name.like(like), Student.stu_no.like(like), Student.phone.like(like), Student.major.like(like), Student.graduate_school.like(like), Student.native_place.like(like), ) ) if class_id: stmt = stmt.where(Student.class_id == class_id) if class_no: stmt = stmt.where( Student.class_id.in_(select(Clazz.id).where(Clazz.class_no == class_no, Clazz.alive())) ) if status: stmt = stmt.where(Student.status == status) if gender: stmt = stmt.where(Student.gender == gender) if advisor_id: stmt = stmt.where(Student.advisor_id == advisor_id) if education: stmt = stmt.where(Student.education == education) if age_min is not None: stmt = stmt.where(age_expression(Student.birth_date) >= age_min) if age_max is not None: stmt = stmt.where(age_expression(Student.birth_date) <= age_max) # 排序:白名单映射,避免直接把前端字符串塞进 order_by sortable = { "id": Student.id, "stu_no": Student.stu_no, "name": Student.name, "age": age_expression(Student.birth_date), "enroll_date": Student.enroll_date, "graduate_date": Student.graduate_date, "status": Student.status, "class_id": Student.class_id, } column = sortable.get(order_by or "id", Student.id) stmt = stmt.order_by(column.desc() if (order or "desc").lower() == "desc" else column.asc()) return stmt # ------------------------------------------------------------ 专用查询 @classmethod def get_by_stu_no(cls, db: Session, stu_no: str, with_deleted: bool = False) -> Student | None: stmt = select(Student).where(Student.stu_no == stu_no) if not with_deleted: stmt = stmt.where(Student.alive()) return db.scalars(stmt).unique().first() @classmethod def count_by_class(cls, db: Session) -> dict[int, int]: stmt = ( select(Student.class_id, func.count(Student.id)) .where(Student.alive(), Student.class_id.is_not(None)) .group_by(Student.class_id) ) return cls.group_count(db, stmt) @classmethod def count_by_status(cls, db: Session) -> dict[int, int]: stmt = select(Student.status, func.count(Student.id)).where(Student.alive()).group_by(Student.status) return cls.group_count(db, stmt) @classmethod def last_stu_no(cls, db: Session, prefix: str) -> str | None: """取同前缀下最大的学号(含已删除,避免复用导致唯一键撞车)。""" stmt = ( select(func.max(Student.stu_no)) .where(Student.stu_no.like(f"{prefix}%")) ) return db.scalar(stmt) @classmethod def get_with_relations(cls, db: Session, pk: int) -> Student | None: return cls.get(db, pk) # ------------------------------------------------------------ 成绩相关 @classmethod def warning_student_ids(cls, db: Session) -> list[int]: stmt = select(func.distinct(Score.stu_id)).where(Score.alive(), Score.flag == 1) return list(db.scalars(stmt).all()) @classmethod def birthday_stats(cls, db: Session) -> dict[str, date | None]: stmt = select(func.min(Student.birth_date), func.max(Student.birth_date)).where(Student.alive()) row = db.execute(stmt).first() return {"min": row[0] if row else None, "max": row[1] if row else None}