98 lines
3.7 KiB
Python
98 lines
3.7 KiB
Python
"""成绩 DAO(需求 2.2)。"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from sqlalchemy import Select, func, or_, select
|
|
from sqlalchemy.orm import Session
|
|
|
|
from app.dao.base_dao import BaseDao
|
|
from app.model import Clazz, Score, Student
|
|
|
|
|
|
class ScoreDao(BaseDao[Score]):
|
|
model = Score
|
|
|
|
@classmethod
|
|
def build_stmt(
|
|
cls,
|
|
stu_id: int | None = None,
|
|
exam_seq: int | None = None,
|
|
class_id: int | None = None,
|
|
flag: int | None = None,
|
|
min_score: float | None = None,
|
|
max_score: float | None = None,
|
|
keyword: str | None = None,
|
|
order_by: str = "id",
|
|
order: str = "desc",
|
|
) -> Select:
|
|
stmt = select(Score).join(Student, Student.id == Score.stu_id).where(
|
|
Score.alive(), Student.alive()
|
|
)
|
|
|
|
if stu_id:
|
|
stmt = stmt.where(Score.stu_id == stu_id)
|
|
if exam_seq:
|
|
stmt = stmt.where(Score.exam_seq == exam_seq)
|
|
if class_id:
|
|
stmt = stmt.where(Student.class_id == class_id)
|
|
if flag is not None:
|
|
stmt = stmt.where(Score.flag == flag)
|
|
if min_score is not None:
|
|
stmt = stmt.where(Score.score >= min_score)
|
|
if max_score is not None:
|
|
stmt = stmt.where(Score.score <= max_score)
|
|
if keyword:
|
|
like = f"%{keyword.strip()}%"
|
|
stmt = stmt.where(or_(Student.name.like(like), Student.stu_no.like(like)))
|
|
|
|
sortable = {
|
|
"id": Score.id,
|
|
"exam_seq": Score.exam_seq,
|
|
"score": Score.score,
|
|
"exam_date": Score.exam_date,
|
|
"student_name": Student.name,
|
|
"class_name": Clazz.name,
|
|
"stu_no": Student.stu_no,
|
|
}
|
|
column = sortable.get(order_by or "id", Score.id)
|
|
stmt = stmt.order_by(column.desc() if (order or "desc").lower() == "desc" else column.asc())
|
|
return stmt
|
|
|
|
@classmethod
|
|
def get_by_stu_seq(cls, db: Session, stu_id: int, exam_seq: int, with_deleted: bool = False) -> Score | None:
|
|
"""按 (学生, 序次) 取记录 —— 唯一键保证最多一条。"""
|
|
stmt = select(Score).where(Score.stu_id == stu_id, Score.exam_seq == exam_seq)
|
|
if not with_deleted:
|
|
stmt = stmt.where(Score.alive())
|
|
return db.scalars(stmt).unique().first()
|
|
|
|
@classmethod
|
|
def list_by_student(cls, db: Session, stu_id: int) -> list[Score]:
|
|
stmt = select(Score).where(Score.alive(), Score.stu_id == stu_id).order_by(Score.exam_seq.asc())
|
|
return list(db.scalars(stmt).unique().all())
|
|
|
|
@classmethod
|
|
def max_exam_seq(cls, db: Session) -> int:
|
|
return int(db.scalar(select(func.max(Score.exam_seq)).where(Score.alive())) or 0)
|
|
|
|
@classmethod
|
|
def student_score_map(cls, db: Session) -> dict[int, list[float]]:
|
|
"""一次性取出所有学生的成绩列表,给波动分析/均分用,避免 N+1。"""
|
|
stmt = select(Score.stu_id, Score.exam_seq, Score.score).where(Score.alive()).order_by(
|
|
Score.stu_id, Score.exam_seq
|
|
)
|
|
result: dict[int, list[float]] = {}
|
|
for stu_id, _seq, score in db.execute(stmt).all():
|
|
result.setdefault(stu_id, []).append(float(score))
|
|
return result
|
|
|
|
@classmethod
|
|
def exam_seq_list(cls, db: Session) -> list[int]:
|
|
stmt = select(Score.exam_seq).where(Score.alive()).distinct().order_by(Score.exam_seq)
|
|
return [int(x) for x in db.scalars(stmt).all()]
|
|
|
|
@classmethod
|
|
def count_by_student(cls, db: Session) -> dict[int, int]:
|
|
stmt = select(Score.stu_id, func.count(Score.id)).where(Score.alive()).group_by(Score.stu_id)
|
|
return {row[0]: row[1] for row in db.execute(stmt).all()}
|