Files
test/app/dao/student_dao.py
2026-09-21 19:03:31 +08:00

128 lines
4.7 KiB
Python

"""学生 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}