Files
stu_teacher/dao/statistics_dao.py
T

462 lines
18 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# dao/statistics_dao.py
# 统计分析模块的数据访问层:
# - 全部使用 SQLAlchemy 表达式动态拼装查询(避免原生 SQL 字符串拼接带来的注入风险)
# - 覆盖需求 2.6(动态年龄/班级统计/成绩统计/就业统计)与 2.7(高级筛选器/聚合统计)
from typing import List, Tuple
from sqlalchemy import and_, or_, func, case
from sqlalchemy.orm import Session
from model.students import Student
from model.classes import Classinfo
from model.scores import Score
from model.employment import EmploymentBase
from scheme.statistics import FilterRule, FilterGroup
# 高级筛选器允许的字段白名单 -> (SQLAlchemy 列, 是否需要外联 employment_base)
_FILTER_FIELDS = {
"stu_id": (Student.stu_id, False),
"stu_name": (Student.stu_name, False),
"age": (Student.age, False),
"gender": (Student.gender, False),
"education": (Student.education, False),
"major": (Student.major, False),
"native_place": (Student.native_place, False),
"status": (Student.status, False),
"class_id": (Student.class_id, False),
"class_name": (Classinfo.class_name, False),
"salary": (EmploymentBase.salary, True),
"company_name": (EmploymentBase.company_name, True),
}
class StatisticsDAO:
# ==================== 2.6.1 动态年龄范围查询 ====================
@staticmethod
def students_by_age(
db: Session,
op: str,
value: int = None,
min_value: int = None,
max_value: int = None,
) -> List[Student]:
"""
动态年龄查询
:param op: 比较条件 gt/lt/eq/gte/lte/between
"""
query = (
db.query(Student)
.join(Classinfo, Student.class_id == Classinfo.class_id)
.filter(Student.is_deleted == 0)
)
if op == "gt":
query = query.filter(Student.age > value)
elif op == "lt":
query = query.filter(Student.age < value)
elif op == "eq":
query = query.filter(Student.age == value)
elif op == "gte":
query = query.filter(Student.age >= value)
elif op == "lte":
query = query.filter(Student.age <= value)
elif op == "between":
query = query.filter(Student.age >= min_value, Student.age <= max_value)
else:
raise ValueError(f"不支持的年龄比较条件: {op}")
return query.order_by(Student.age).all()
# ==================== 2.6.1 多维度班级统计 ====================
@staticmethod
def class_gender_stats(db: Session) -> List[dict]:
"""统计每个班级总人数及男女分布"""
male_cnt = func.sum(case((Student.gender == "男", 1), else_=0))
female_cnt = func.sum(case((Student.gender == "女", 1), else_=0))
rows = (
db.query(
Classinfo.class_id,
Classinfo.class_name,
func.count(Student.stu_id).label("total"),
male_cnt.label("male"),
female_cnt.label("female"),
)
.join(Student, Student.class_id == Classinfo.class_id, isouter=True)
.filter(Classinfo.is_deleted == 0, Student.is_deleted == 0)
.group_by(Classinfo.class_id, Classinfo.class_name)
.order_by(Classinfo.class_id)
.all()
)
return [
{
"class_id": r.class_id,
"class_name": r.class_name,
"total": r.total or 0,
"male": r.male or 0,
"female": r.female or 0,
}
for r in rows
]
# ==================== 2.6.2 每次考试都在分数线以上的学生 ====================
@staticmethod
def students_all_above(db: Session, line: float) -> List[dict]:
"""查询每次考试成绩都在 line 分以上的学生(按最低分聚合判断)"""
rows = (
db.query(
Student.stu_id,
Student.stu_name,
Classinfo.class_name,
func.count(Score.exam_id).label("exam_count"),
func.min(Score.score).label("min_score"),
)
.join(Score, Score.stu_id == Student.stu_id)
.join(Classinfo, Student.class_id == Classinfo.class_id)
.filter(Score.is_deleted == 0, Student.is_deleted == 0)
.group_by(Student.stu_id, Student.stu_name, Classinfo.class_name)
.having(func.min(Score.score) >= line)
.all()
)
result = []
for r in rows:
details = (
db.query(Score.exam_id, Score.score)
.filter(Score.stu_id == r.stu_id, Score.is_deleted == 0)
.order_by(Score.exam_id)
.all()
)
result.append(
{
"stu_id": r.stu_id,
"stu_name": r.stu_name,
"class_name": r.class_name,
"exam_count": r.exam_count,
"min_score": r.min_score,
"scores": [{"exam_id": d.exam_id, "score": d.score} for d in details],
}
)
return result
# ==================== 2.6.2 不及格次数 >= N 的学生 ====================
@staticmethod
def fail_students(db: Session, min_times: int, line: float = 60.0) -> List[dict]:
"""查询不及格(< line)次数 >= min_times 的学生及其不及格明细"""
fail_cond = and_(Score.score < line, Score.is_deleted == 0)
rows = (
db.query(
Student.stu_id,
Student.stu_name,
Classinfo.class_name,
func.count(Score.exam_id).label("fail_count"),
)
.join(Score, Score.stu_id == Student.stu_id)
.join(Classinfo, Student.class_id == Classinfo.class_id)
.filter(Student.is_deleted == 0, fail_cond)
.group_by(Student.stu_id, Student.stu_name, Classinfo.class_name)
.having(func.count(Score.exam_id) >= min_times)
.all()
)
result = []
for r in rows:
details = (
db.query(Score.exam_id, Score.score)
.filter(fail_cond, Score.stu_id == r.stu_id)
.order_by(Score.exam_id)
.all()
)
result.append(
{
"stu_id": r.stu_id,
"stu_name": r.stu_name,
"class_name": r.class_name,
"fail_count": r.fail_count,
"fail_details": [{"exam_id": d.exam_id, "score": d.score} for d in details],
}
)
return result
# ==================== 2.6.2 每次考试每个班级的平均分(动态排序) ====================
@staticmethod
def class_exam_avg(db: Session, exam_id: int = None, order: str = "desc") -> List[dict]:
"""统计每次考试每个班级的平均分,order: asc/desc"""
avg_expr = func.round(func.avg(Score.score), 2)
query = (
db.query(
Score.exam_id,
Classinfo.class_id,
Classinfo.class_name,
avg_expr.label("avg_score"),
)
.join(Student, Score.stu_id == Student.stu_id)
.join(Classinfo, Student.class_id == Classinfo.class_id)
.filter(Score.is_deleted == 0, Student.is_deleted == 0)
.group_by(Score.exam_id, Classinfo.class_id, Classinfo.class_name)
)
if exam_id is not None:
query = query.filter(Score.exam_id == exam_id)
query = query.order_by(avg_expr.desc() if order == "desc" else avg_expr.asc())
return [
{
"exam_id": r.exam_id,
"class_id": r.class_id,
"class_name": r.class_name,
"avg_score": float(r.avg_score),
}
for r in query.all()
]
# ==================== 2.6.3 就业薪资 Top N ====================
@staticmethod
def top_salary(db: Session, n: int) -> List[dict]:
"""薪资排名 Top N(从就业基础表取最新薪资)"""
rows = (
db.query(
EmploymentBase.stu_id,
EmploymentBase.stu_name,
EmploymentBase.class_name,
EmploymentBase.job_time,
EmploymentBase.company_name,
EmploymentBase.salary,
)
.join(Student, EmploymentBase.stu_id == Student.stu_id)
.filter(
EmploymentBase.is_deleted == 0,
Student.is_deleted == 0,
EmploymentBase.salary > 0,
)
.order_by(EmploymentBase.salary.desc())
.limit(n)
.all()
)
return [
{
"stu_id": r.stu_id,
"stu_name": r.stu_name,
"class_name": r.class_name,
"job_time": r.job_time,
"company_name": r.company_name,
"salary": r.salary,
}
for r in rows
]
# ==================== 2.6.3 每个学生的就业时长 ====================
@staticmethod
def employment_durations(db: Session) -> List[dict]:
"""就业时长 = offer下发时间(job_time) - 就业开放时间(employment_open_time),单位天"""
duration_days = func.datediff(EmploymentBase.job_time, EmploymentBase.employment_open_time)
rows = (
db.query(
EmploymentBase.stu_id,
EmploymentBase.stu_name,
EmploymentBase.class_name,
EmploymentBase.employment_open_time,
EmploymentBase.job_time,
duration_days.label("duration_days"),
)
.filter(EmploymentBase.is_deleted == 0)
.order_by(EmploymentBase.stu_id)
.all()
)
return [
{
"stu_id": r.stu_id,
"stu_name": r.stu_name,
"class_name": r.class_name,
"employment_open_time": r.employment_open_time,
"job_time": r.job_time,
# 未拿到 offer 记为 -1,前端展示为"未就业"
"duration_days": int(r.duration_days) if r.duration_days is not None else -1,
}
for r in rows
]
# ==================== 2.6.3 每个班级平均就业时长 ====================
@staticmethod
def class_avg_duration(db: Session) -> List[dict]:
"""平均就业时长:仅统计进入就业阶段(有就业开放时间)的学生;
平均值仅对已拿到 offer 的学生计算"""
opened = func.count(EmploymentBase.stu_id)
offered = func.sum(case((EmploymentBase.job_time.isnot(None), 1), else_=0))
avg_days = func.round(
func.avg(
case(
(
EmploymentBase.job_time.isnot(None),
func.datediff(EmploymentBase.job_time, EmploymentBase.employment_open_time),
)
)
),
1,
)
rows = (
db.query(
Student.class_id,
Classinfo.class_name,
opened.label("opened_count"),
offered.label("offered_count"),
avg_days.label("avg_duration_days"),
)
.join(EmploymentBase, EmploymentBase.stu_id == Student.stu_id)
.join(Classinfo, Student.class_id == Classinfo.class_id)
.filter(Student.is_deleted == 0, EmploymentBase.is_deleted == 0)
.group_by(Student.class_id, Classinfo.class_name)
.order_by(Student.class_id)
.all()
)
return [
{
"class_id": r.class_id,
"class_name": r.class_name,
"opened_count": r.opened_count or 0,
"offered_count": int(r.offered_count or 0),
"avg_duration_days": float(r.avg_duration_days) if r.avg_duration_days is not None else 0.0,
}
for r in rows
]
# ==================== 2.7.2 成绩波动分析(最大分差 Top N) ====================
@staticmethod
def score_volatility(db: Session, top_n: int = 5) -> List[dict]:
"""成绩波动最大 Top N(最大分差 = 最高分 - 最低分,SQL 聚合计算)"""
diff_expr = (func.max(Score.score) - func.min(Score.score)).label("diff")
rows = (
db.query(
Student.stu_id,
Student.stu_name,
Classinfo.class_name,
func.max(Score.score).label("max_score"),
func.min(Score.score).label("min_score"),
diff_expr,
)
.join(Score, Score.stu_id == Student.stu_id)
.join(Classinfo, Student.class_id == Classinfo.class_id)
.filter(Score.is_deleted == 0, Student.is_deleted == 0)
.group_by(Student.stu_id, Student.stu_name, Classinfo.class_name)
.order_by(diff_expr.desc())
.limit(top_n)
.all()
)
return [
{
"stu_id": r.stu_id,
"stu_name": r.stu_name,
"class_name": r.class_name,
"max_score": r.max_score,
"min_score": r.min_score,
"diff": float(r.diff),
}
for r in rows
]
# ==================== 2.7.2 班级就业漏斗 ====================
@staticmethod
def employment_funnel(db: Session, high_salary_line: float = 10000.0) -> List[dict]:
"""每个班级:总人数 -> 已就业人数 -> 高薪人数(>10k) -> 就业率,按就业率降序"""
employed_cnt = func.sum(case((EmploymentBase.stu_id.isnot(None), 1), else_=0))
high_salary_cnt = func.sum(
case((and_(EmploymentBase.stu_id.isnot(None), EmploymentBase.salary > high_salary_line), 1), else_=0)
)
rows = (
db.query(
Classinfo.class_id,
Classinfo.class_name,
func.count(Student.stu_id).label("total"),
employed_cnt.label("employed"),
high_salary_cnt.label("high_salary"),
)
.join(Student, Student.class_id == Classinfo.class_id)
.join(
EmploymentBase,
and_(
EmploymentBase.stu_id == Student.stu_id,
EmploymentBase.is_deleted == 0,
),
isouter=True,
)
.filter(Student.is_deleted == 0, Classinfo.is_deleted == 0)
.group_by(Classinfo.class_id, Classinfo.class_name)
.all()
)
result = []
for r in rows:
total = r.total or 0
employed = int(r.employed or 0)
rate = round(employed / total * 100, 2) if total else 0.0
result.append(
{
"class_id": r.class_id,
"class_name": r.class_name,
"total": total,
"employed": employed,
"high_salary": int(r.high_salary or 0),
"employment_rate": rate,
}
)
result.sort(key=lambda x: x["employment_rate"], reverse=True)
return result
# ============================================================
# 2.7.1 通用高级筛选器:把规则树递归翻译为 SQLAlchemy 表达式
# ============================================================
class FilterBuilder:
@staticmethod
def _rule_to_expr(rule: FilterRule):
"""把单条规则翻译为 SQLAlchemy 比较表达式"""
if rule.field not in _FILTER_FIELDS:
raise ValueError(f"不支持筛选的字段: {rule.field},允许的字段: {sorted(_FILTER_FIELDS)}")
column, _ = _FILTER_FIELDS[rule.field]
op = rule.operator
if op == ">":
return column > rule.value
if op == "<":
return column < rule.value
if op == "=":
return column == rule.value
if op == "!=":
return column != rule.value
if op == ">=":
return column >= rule.value
if op == "<=":
return column <= rule.value
if op == "like":
return column.like(f"%{rule.value}%")
if op == "in":
if not isinstance(rule.value, (list, tuple)):
raise ValueError("operator=in 时 value 必须是列表")
return column.in_(list(rule.value))
raise ValueError(f"不支持的操作符: {op}")
@classmethod
def to_expr(cls, rules: list):
"""把规则列表(顶层默认 AND)翻译为一个组合表达式"""
if not rules:
raise ValueError("筛选规则不能为空")
exprs = []
for r in rules:
exprs.append(cls._node_to_expr(r))
return and_(*exprs) if len(exprs) > 1 else exprs[0]
@classmethod
def _node_to_expr(cls, node):
"""递归处理规则节点:FilterGroup 组合子规则,FilterRule 直接翻译"""
if isinstance(node, FilterGroup):
sub = [cls._node_to_expr(r) for r in node.sub_rules]
return or_(*sub) if node.logic == "OR" else and_(*sub)
if isinstance(node, FilterRule):
return cls._rule_to_expr(node)
raise ValueError(f"无法识别的筛选规则节点: {type(node)}")
@staticmethod
def query_students(db: Session, rules: list) -> Tuple[int, list]:
"""执行高级筛选查询(student 模型,自动关联班级与就业表)"""
expr = FilterBuilder.to_expr(rules)
query = (
db.query(Student)
.join(Classinfo, Student.class_id == Classinfo.class_id)
.outerjoin(EmploymentBase, EmploymentBase.stu_id == Student.stu_id)
.filter(Student.is_deleted == 0, expr)
)
total = query.count()
items = query.order_by(Student.stu_id).all()
return total, items