from sqlalchemy.orm import Session from typing import Optional, List, Tuple from model.classes_model import Classes from model.teachers import Teacher from schemas.class_scheme import ( ClassCreate, ClassUpdateBase, ClassHeadTeacherUpdate, ClassTeachersUpdate, ClassDelete, ) class ClassDAO: # ---------- 创建 ---------- @staticmethod def create(db: Session, req: ClassCreate) -> Classes: # 1) 班级名唯一 exists = ( db.query(Classes) .filter(Classes.class_name == req.class_name) .first() ) if exists: raise ValueError("班级名称已存在") # 2) 班主任存在 head_teacher = ( db.query(Teacher).filter(Teacher.t_id == req.head_teacher_id).first() ) if not head_teacher: raise ValueError("班主任不存在") # 3) 授课老师全部存在 teachers: List[Teacher] = [] if req.teacher_ids: teacher_ids = set(req.teacher_ids) teachers = ( db.query(Teacher).filter(Teacher.t_id.in_(teacher_ids)).all() ) if len(teachers) != len(teacher_ids): raise ValueError("部分授课老师不存在") data = req.model_dump( exclude={"teacher_ids"}, exclude_none=True, ) db_class = Classes(**data) db_class.teachers = teachers db.add(db_class) db.commit() db.refresh(db_class) return db_class @staticmethod def get_by_id(db: Session, class_id: int) -> Optional[Classes]: return ( db.query(Classes) .filter( Classes.class_id == class_id, Classes.class_delete == False, # noqa: E712 ) .first() ) # ---------- 分页 ---------- @staticmethod def list_page( db: Session, page_num: int, page_size: int, class_name: Optional[str] = None, head_teacher_id: Optional[int] = None, ) -> Tuple[int, List[Classes]]: query = db.query(Classes).filter(Classes.class_delete == False) # noqa: E712 if class_name: query = query.filter(Classes.class_name.like(f"%{class_name}%")) if head_teacher_id is not None: query = query.filter(Classes.head_teacher_id == head_teacher_id) total = query.count() rows = ( query.order_by(Classes.class_id.desc()) .offset((page_num - 1) * page_size) .limit(page_size) .all() ) return total, rows @staticmethod def update_base( db: Session, class_id: int, req: ClassUpdateBase ) -> Optional[Classes]: db_class = ClassDAO.get_by_id(db, class_id) if not db_class: return None data = req.model_dump(exclude_unset=True, exclude_none=True) # 改班级名要查重 new_name = data.get("class_name") if new_name and new_name != db_class.class_name: exists = ( db.query(Classes) .filter( Classes.class_name == new_name, Classes.class_id != class_id, ) .first() ) if exists: raise ValueError("班级名称已存在") for k, v in data.items(): setattr(db_class, k, v) db.commit() db.refresh(db_class) return db_class @staticmethod def update_head_teacher( db: Session, class_id: int, req: ClassHeadTeacherUpdate ) -> Optional[Classes]: db_class = ClassDAO.get_by_id(db, class_id) if not db_class: return None teacher = ( db.query(Teacher).filter(Teacher.t_id == req.head_teacher_id).first() ) if not teacher: raise ValueError("班主任不存在") db_class.head_teacher_id = req.head_teacher_id db.commit() db.refresh(db_class) return db_class @staticmethod def update_teachers( db: Session, class_id: int, req: ClassTeachersUpdate ) -> Optional[Classes]: db_class = ClassDAO.get_by_id(db, class_id) if not db_class: return None teacher_ids = set(req.teacher_ids) teachers = ( db.query(Teacher).filter(Teacher.t_id.in_(teacher_ids)).all() ) if teacher_ids else [] if len(teachers) != len(teacher_ids): raise ValueError("部分授课老师不存在") db_class.teachers = teachers db.commit() db.refresh(db_class) return db_class @staticmethod def batch_toggle_delete(db: Session, req: ClassDelete) -> int: count = ( db.query(Classes) .filter(Classes.class_id.in_(req.class_ids)) .update( {"class_delete": req.deleted}, synchronize_session=False, ) ) db.commit() return count