from fastapi import HTTPException
|
from sqlalchemy import false, select
|
from sqlalchemy.orm import Session
|
|
from app.models import GraderSubject, ImportBatch, User
|
|
|
def list_all_subject_names(db: Session) -> list[str]:
|
return list(db.scalars(select(ImportBatch.subject_name).distinct().order_by(ImportBatch.subject_name)))
|
|
|
def get_grader_subjects(db: Session, user_id: int) -> list[str]:
|
return list(
|
db.scalars(
|
select(GraderSubject.subject_name)
|
.where(GraderSubject.user_id == user_id)
|
.order_by(GraderSubject.subject_name)
|
)
|
)
|
|
|
def validate_subject_names(db: Session, names: list[str]) -> list[str]:
|
cleaned = []
|
seen = set()
|
for n in names:
|
s = (n or "").strip()
|
if not s or s in seen:
|
continue
|
seen.add(s)
|
cleaned.append(s)
|
if not cleaned:
|
raise HTTPException(400, detail="请至少选择一个科目")
|
known = set(list_all_subject_names(db))
|
unknown = [s for s in cleaned if s not in known]
|
if unknown:
|
raise HTTPException(400, detail=f"未知科目:{', '.join(unknown)}")
|
return cleaned
|
|
|
def set_grader_subjects(db: Session, user_id: int, names: list[str]) -> None:
|
validated = validate_subject_names(db, names)
|
db.query(GraderSubject).filter(GraderSubject.user_id == user_id).delete()
|
for name in validated:
|
db.add(GraderSubject(user_id=user_id, subject_name=name))
|
|
|
def allowed_subjects_for_user(db: Session, user: User) -> list[str] | None:
|
"""None = 不限制(管理员);列表 = 考评员可见科目。"""
|
if user.role == "admin":
|
return None
|
if user.role == "grader":
|
return get_grader_subjects(db, user.id)
|
return []
|
|
|
def assert_subject_allowed(db: Session, user: User, subject_name: str | None) -> None:
|
if not subject_name:
|
raise HTTPException(404, detail="记录不存在")
|
allowed = allowed_subjects_for_user(db, user)
|
if allowed is None:
|
return
|
if subject_name not in allowed:
|
raise HTTPException(403, detail="无权查看该科目数据")
|
|
|
def apply_subject_scope(q, subject_column, allowed: list[str] | None):
|
if allowed is None:
|
return q
|
if not allowed:
|
return q.where(false())
|
return q.where(subject_column.in_(allowed))
|