Files
------/dao/workspace_dao.py

154 lines
8.5 KiB
Python

"""受控 ORM 查询,共享给前端和 AI;筛选、统计、排序在数据库内完成。"""
from datetime import datetime
from enum import Enum
from decimal import Decimal
from math import isfinite
from sqlalchemy import func, case, literal_column, Float, Integer, Numeric, String
from model import Student, Classes, Teacher, Advisor, Score, Employment
from model.teacher_model import teacher_class, teacher_student
from schema.workspace_schema import DataQuery
MODELS = {"students": Student, "classes": Classes, "teachers": Teacher,
"advisors": Advisor, "scores": Score, "employment": Employment}
ACTIVE = {"students": Student.flag, "classes": Classes.is_del,
"teachers": Teacher.is_deleted, "advisors": Advisor.flag,
"scores": Score.flag, "employment": Employment.flag}
def query_definition(db, module):
if module in ("teacher_classes", "teacher_students"):
if module == "teacher_classes":
fields = {"tid": Teacher.tid, "t_name": Teacher.t_name,
"cid": Classes.cid, "class_name": Classes.class_name}
query = db.query(Teacher).join(teacher_class, Teacher.tid == teacher_class.c.tid).join(Classes, Classes.cid == teacher_class.c.cid).filter(Teacher.is_deleted == 1, Classes.is_del == 1)
else:
fields = {"tid": Teacher.tid, "t_name": Teacher.t_name,
"sid": Student.sid, "student_name": Student.student_name}
query = db.query(Teacher).join(teacher_student, Teacher.tid == teacher_student.c.tid).join(Student, Student.sid == teacher_student.c.sid).filter(Teacher.is_deleted == 1, Student.flag == 1)
return query, fields
model = MODELS[module]
fields = {column.name: getattr(model, column.name) for column in model.__table__.columns
if column.name not in ("flag", "is_del", "is_deleted")}
query = db.query(model).filter(ACTIVE[module] == 1)
if module in ("scores", "employment"):
query = query.join(Student, model.student_id == Student.sid).filter(Student.flag == 1)
fields.update(student_name=Student.student_name, student_no=Student.student_no, class_id=Student.class_id)
if module in ("students", "scores", "employment"):
query = query.outerjoin(Classes, Student.class_id == Classes.cid).outerjoin(Advisor, Student.advisor_id == Advisor.id)
fields.update(class_name=Classes.class_name, advisor_name=Advisor.advisor_name)
if module != "students":
fields["advisor_id"] = Student.advisor_id
if module == "employment":
# 异常或缺失日期不作为 0 天参与统计。
if db.bind.dialect.name == "sqlite":
duration = func.julianday(Employment.employment_start_time) - func.julianday(Employment.offer_time)
else:
duration = func.timestampdiff(literal_column("SECOND"), Employment.offer_time, Employment.employment_start_time) / 86400.0
fields["duration_days"] = case((Employment.employment_start_time >= Employment.offer_time, duration), else_=None)
return query, fields
def public_value(value):
if isinstance(value, Enum):
return value.value
if isinstance(value, datetime):
return value.isoformat(timespec="seconds")
if hasattr(value, "isoformat"):
return value.isoformat()
if isinstance(value, Decimal):
return float(value)
return value
def execute_query(db, request: DataQuery):
query, fields = query_definition(db, request.module)
def field(name):
if name not in fields:
raise ValueError(f"{request.module} 不支持字段 {name}")
return fields[name]
for condition in request.filters:
column = field(condition.field)
value = condition.value
if condition.op in ("is_null", "not_null"):
expression = column.is_(None) if condition.op == "is_null" else column.isnot(None)
else:
if value is None:
raise ValueError("筛选条件缺少值")
if isinstance(column.type, (Integer, Float, Numeric)):
try:
value = float(value)
if not isfinite(value):
raise ValueError("数值必须有限")
except (ValueError, TypeError):
raise ValueError(f"{condition.field} 必须为数值") from None
if condition.op == "contains":
if not isinstance(column.type, String):
raise ValueError("包含匹配仅支持文本字段")
expression = column.contains(str(value), autoescape=True)
else:
operators = {"eq": column.__eq__, "gt": column.__gt__, "ge": column.__ge__, "lt": column.__lt__, "le": column.__le__}
expression = operators[condition.op](value)
query = query.filter(expression)
if request.aggregate:
group_fields = {name: field(name) for name in dict.fromkeys(request.group_by)}
if request.aggregate == "count":
metric = func.count()
else:
target = field(request.aggregate_field)
if request.aggregate_field != "duration_days" and not isinstance(target.type, (Integer, Float, Numeric)):
raise ValueError("该字段不支持数值统计")
if request.aggregate_field == "salary":
query = query.filter(target >= 0)
metric = getattr(func, request.aggregate)(target)
query = query.with_entities(*[col.label(name) for name, col in group_fields.items()], metric.label("value"))
if group_fields:
query = query.group_by(*group_fields.values())
if request.having_min is not None:
query = query.having(metric >= request.having_min)
order_fields = {**group_fields, "value": metric}
else:
query = query.with_entities(*[column.label(name) for name, column in fields.items()])
order_fields = fields
sort_name = request.sort_by or ("value" if request.aggregate else next(iter(fields)))
if sort_name not in order_fields:
raise ValueError(f"当前查询不能按 {sort_name} 排序")
total = query.count()
sort_column = order_fields[sort_name]
query = query.order_by(sort_column.desc() if request.descending else sort_column.asc())
# 同值时按其余键稳定排序,保证分页结果一致。
tie_fields = list(order_fields) if request.aggregate else [next(iter(fields))]
for name in tie_fields:
if name != sort_name:
query = query.order_by(order_fields[name].asc())
rows = []
for row in query.offset((request.page - 1) * request.page_size).limit(request.page_size).all():
result = {key: public_value(value) for key, value in row._mapping.items()}
if request.module == "advisors" and "phone" in result and result["phone"]:
phone = result["phone"]
result["phone"] = phone[:3] + "****" + phone[7:]
rows.append(result)
return {"items": rows, "total": total, "page": request.page, "page_size": request.page_size,
"pages": max(1, (total + request.page_size - 1) // request.page_size),
"truncated": total > len(rows), "query": request.model_dump(exclude_none=True),
"queried_at": datetime.now().astimezone().isoformat(timespec="seconds"),
"scope": "仅当前有效记录;成绩和就业同时要求学生有效。列表按页返回,聚合在完整筛选结果上计算。关联班级、顾问名称保留历史关联。空值代表未填写;就业时长为就业开放时间减 Offer 下发时间,相同时间为0天,负值不参与统计;顾问电话脱敏。"}
def overview(db):
counts = {module: db.query(model).filter(ACTIVE[module] == 1).count() for module, model in MODELS.items()}
jobs = db.query(Employment).join(Student, Employment.student_id == Student.sid).filter(Employment.flag == 1, Student.flag == 1)
counts['employment'] = jobs.count()
counts['scores'] = db.query(Score).join(Student, Score.student_id == Student.sid).filter(Score.flag == 1, Student.flag == 1).count()
sample = jobs.filter(Employment.salary >= 0)
avg = sample.with_entities(func.avg(Employment.salary)).scalar()
return {"counts": counts, "offers": jobs.filter(Employment.offer_time.isnot(None)).count(),
"average_salary": round(float(avg), 2) if avg is not None else None,
"salary_samples": sample.count(),
"queried_at": datetime.now().astimezone().isoformat(timespec="seconds"),
"scope": "各模块有效记录;成绩和就业还要求学生有效;平均薪资只计算已填写且非负的薪资。"}