154 lines
8.5 KiB
Python
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": "各模块有效记录;成绩和就业还要求学生有效;平均薪资只计算已填写且非负的薪资。"}
|