92 lines
3.4 KiB
Python
92 lines
3.4 KiB
Python
"""班级 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()}
|