diff --git a/models/library.py b/models/library.py index f162b46..515cd5f 100644 --- a/models/library.py +++ b/models/library.py @@ -5,46 +5,50 @@ from sqlalchemy.orm import relationship from models.base import AuditBase +class Subject(AuditBase): + name = Column(String(45), nullable=False) + + modules = relationship("Module", back_populates="subject") + library_files = relationship("LibraryFile", back_populates="subject") + records = relationship("Record", back_populates="subject") + mistakes = relationship("Mistake", back_populates="subject") + + +class Module(AuditBase): + subject_id = Column(BigInteger, ForeignKey("subject.id", ondelete="RESTRICT"), nullable=False, index=True) + name = Column(String(45), nullable=False) + + subject = relationship("Subject", back_populates="modules") + library_files = relationship("LibraryFile", back_populates="module") + + class LibraryFile(AuditBase): - subject = Column(String(45), nullable=False) - module = Column(String(45), nullable=False) + subject_id = Column(BigInteger, ForeignKey("subject.id", ondelete="RESTRICT"), nullable=False, index=True) + module_id = Column(BigInteger, ForeignKey("module.id", ondelete="RESTRICT"), nullable=False, index=True) name = Column(String(45), nullable=False) type = Column(String(45), nullable=False) size = Column(Integer, nullable=False) content = Column(Text, nullable=False) - practice_record = relationship( - "PracticeRecord", - back_populates="library_file", - cascade="all, delete-orphan" - ) + subject = relationship("Subject", back_populates="library_files") + module = relationship("Module", back_populates="library_files") -class PracticeRecord(AuditBase): - library_file_id = Column(BigInteger, ForeignKey("library_file.id"), nullable=False) - +class Record(AuditBase): + type = Column(String(45), nullable=False) + subject_id = Column(BigInteger, ForeignKey("subject.id", ondelete="RESTRICT"), nullable=False, index=True) total_count = Column(Integer, nullable=False) correct_count = Column(Integer, nullable=False) wrong_count = Column(Integer, nullable=False) - library_file = relationship( - "LibraryFile", - back_populates="practice_record", - ) - - -class ExamRecord(AuditBase): - subject = Column(String(45), nullable=False) - total_count = Column(Integer, nullable=False) - correct_count = Column(Integer, nullable=False) - wrong_count = Column(Integer, nullable=False) + subject = relationship("Subject", back_populates="records") class Mistake(AuditBase): - subject = Column(String(45), nullable=False) - module = Column(String(45), nullable=False) - name = Column(String(45), nullable=False) + subject_id = Column(BigInteger, ForeignKey("subject.id", ondelete="RESTRICT"), nullable=False, index=True) type = Column(String(45), nullable=False) question = Column(Text, nullable=False) options = Column(JSONB, nullable=False) answer = Column(JSONB, nullable=False) + + subject = relationship("Subject", back_populates="mistakes") diff --git a/routers/__init__.py b/routers/__init__.py index 68e3c1d..924caae 100644 --- a/routers/__init__.py +++ b/routers/__init__.py @@ -1,13 +1,11 @@ from .agent import router as agent_router from .library import router as library_router -from .practice import router as practice_router -from .exam import router as exam_router +from .record import router as record_router from .mistake import router as mistake_router routers = [ agent_router, library_router, - practice_router, - exam_router, + record_router, mistake_router, ] diff --git a/routers/exam.py b/routers/exam.py deleted file mode 100644 index d391b78..0000000 --- a/routers/exam.py +++ /dev/null @@ -1,13 +0,0 @@ -from fastapi import APIRouter, Depends -from sqlalchemy.ext.asyncio import AsyncSession - -from database import get_db -from schemas.exam import ExamRecordRequest -from service import exam_service - -router = APIRouter(prefix="/exam", tags=["Exam"]) - - -@router.post("/record", response_model=bool) -async def add_record(data: ExamRecordRequest, db: AsyncSession = Depends(get_db), ): - return await exam_service.add_record(db, data) diff --git a/routers/mistake.py b/routers/mistake.py index 15089c9..f284860 100644 --- a/routers/mistake.py +++ b/routers/mistake.py @@ -9,14 +9,14 @@ from service import mistake_service router = APIRouter(prefix="/mistake", tags=["Mistake"]) -@router.get("", response_model=List[MistakeResponse]) -async def list_mistakes(db: AsyncSession = Depends(get_db)): - return await mistake_service.list_mistakes(db) +@router.get("/{subject}", response_model=List[MistakeResponse]) +async def list_mistakes(subject: str, db: AsyncSession = Depends(get_db)): + return await mistake_service.list_mistakes(subject, db) -@router.post("", response_model=bool) -async def add_mistakes(data: list[MistakeRequest], db: AsyncSession = Depends(get_db)): - return await mistake_service.add_mistakes(db, data) +@router.post("/{subject}", response_model=bool) +async def add_mistakes(subject: str, data: list[MistakeRequest], db: AsyncSession = Depends(get_db)): + return await mistake_service.add_mistakes(db, subject, data) @router.delete("/{mistake_id}", response_model=bool) diff --git a/routers/practice.py b/routers/practice.py deleted file mode 100644 index c67e41d..0000000 --- a/routers/practice.py +++ /dev/null @@ -1,19 +0,0 @@ -from typing import List -from fastapi import APIRouter, Depends -from sqlalchemy.ext.asyncio import AsyncSession - -from database import get_db -from schemas.practice import PracticeRecordResponse, PracticeRecordRequest -from service import practice_service - -router = APIRouter(prefix="/practice", tags=["Practice"]) - - -@router.get("/record/{user_id}", response_model=List[PracticeRecordResponse]) -async def get_record(user_id: int, db: AsyncSession = Depends(get_db)): - return await practice_service.get_record(db, user_id) - - -@router.post("/record", response_model=bool) -async def add_record(data: PracticeRecordRequest, db: AsyncSession = Depends(get_db), ): - return await practice_service.add_record(db, data) diff --git a/routers/record.py b/routers/record.py new file mode 100644 index 0000000..77ce903 --- /dev/null +++ b/routers/record.py @@ -0,0 +1,19 @@ +from typing import List +from fastapi import APIRouter, Depends +from sqlalchemy.ext.asyncio import AsyncSession + +from database import get_db +from schemas.record import RecordResponse, RecordRequest +from service import record_service + +router = APIRouter(prefix="/record", tags=["Record"]) + + +@router.get("/{subject}", response_model=List[RecordResponse]) +async def get_record(subject: str, db: AsyncSession = Depends(get_db)): + return await record_service.get_record(db, subject) + + +@router.post("/{subject}", response_model=bool) +async def add_record(subject: str, data: RecordRequest, db: AsyncSession = Depends(get_db), ): + return await record_service.add_record(db, subject, data) diff --git a/schemas/exam.py b/schemas/exam.py deleted file mode 100644 index a9bc3f3..0000000 --- a/schemas/exam.py +++ /dev/null @@ -1,16 +0,0 @@ -from pydantic import BaseModel - - -class ExamRecordRequest(BaseModel): - subject: str - totalCount: int - correctCount: int - wrongCount: int - - -class ExamRecordResponse(ExamRecordRequest): - id: int - createTime: str - - class Config: - from_attributes = True diff --git a/schemas/mistake.py b/schemas/mistake.py index f0b7f99..f84b614 100644 --- a/schemas/mistake.py +++ b/schemas/mistake.py @@ -4,9 +4,6 @@ from pydantic import BaseModel class MistakeRequest(BaseModel): - subject: str - module: str - name: str type: str question: str options: List[str] diff --git a/schemas/practice.py b/schemas/record.py similarity index 68% rename from schemas/practice.py rename to schemas/record.py index 3f277dd..4f48f50 100644 --- a/schemas/practice.py +++ b/schemas/record.py @@ -3,14 +3,14 @@ from typing import List from pydantic import BaseModel -class PracticeRecordRequest(BaseModel): - libraryFileId: int +class RecordRequest(BaseModel): + type: str totalCount: int correctCount: int wrongCount: int -class PracticeRecordResponse(PracticeRecordRequest): +class RecordResponse(RecordRequest): id: int createTime: str @@ -18,7 +18,7 @@ class PracticeRecordResponse(PracticeRecordRequest): from_attributes = True -class PracticeCheckResult(BaseModel): +class CheckResult(BaseModel): id: int answer: List[str] select: List[str] diff --git a/service/exam_service.py b/service/exam_service.py deleted file mode 100644 index 058dcf1..0000000 --- a/service/exam_service.py +++ /dev/null @@ -1,18 +0,0 @@ -from sqlalchemy.ext.asyncio import AsyncSession - -from models.library import ExamRecord -from schemas.exam import ExamRecordRequest - - -async def add_record(db: AsyncSession, data: ExamRecordRequest) -> bool: - record = ExamRecord( - subject=data.subject, - total_count=data.totalCount, - correct_count=data.correctCount, - wrong_count=data.wrongCount, - ) - - db.add(record) - await db.commit() - - return True diff --git a/service/library_service.py b/service/library_service.py index 0704a78..187ca2b 100644 --- a/service/library_service.py +++ b/service/library_service.py @@ -1,19 +1,30 @@ +from fastapi import HTTPException from sqlalchemy import select +from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.orm import joinedload -from models.library import LibraryFile +from models.library import LibraryFile, Subject, Module from schemas.library import LibraryFileRequest, LibraryFileResponse async def list_files(db: AsyncSession) -> list[LibraryFileResponse]: - stmt = select(LibraryFile).order_by(LibraryFile.create_time.desc()) - files = (await db.execute(stmt)).scalars().all() + stmt = ( + select(LibraryFile) + .options( + joinedload(LibraryFile.subject), + joinedload(LibraryFile.module), + ) + .order_by(LibraryFile.create_time.desc()) + ) + + files = (await db.execute(stmt)).scalars().unique().all() return [ LibraryFileResponse( id=file.id, - subject=file.subject, - module=file.module, + subject=file.subject.name, + module=file.module.name, name=file.name, type=file.type, size=file.size, @@ -25,60 +36,108 @@ async def list_files(db: AsyncSession) -> list[LibraryFileResponse]: async def add_file(db: AsyncSession, data: LibraryFileRequest) -> bool: - """ - 创建题库文件 - """ - file = LibraryFile( - subject=data.subject, - module=data.module, - name=data.name, - size=data.size, - type=data.type, - content=data.content, - ) + try: + subject_id = await get_or_create_subject(db, data.subject) + module_id = await get_or_create_module(db, subject_id, data.module) - db.add(file) - await db.commit() + file = LibraryFile( + subject_id=subject_id, + module_id=module_id, + name=data.name, + size=data.size, + type=data.type, + content=data.content, + ) - return True + db.add(file) + await db.commit() + return True + + except SQLAlchemyError: + await db.rollback() + raise async def edit_file(db: AsyncSession, file_id: int, data: LibraryFileRequest) -> bool: - """ - 编辑题库文件 - """ + try: + result = await db.execute( + select(LibraryFile).where(LibraryFile.id == file_id) + ) + file = result.scalar_one_or_none() + if not file: + return False - result = await db.execute( - select(LibraryFile).where(LibraryFile.id == file_id) - ) - file = result.scalar_one_or_none() - if not file: - return False + subject_id = await get_or_create_subject(db, data.subject) + module_id = await get_or_create_module(db, subject_id, data.module) - file.subject = data.subject - file.module = data.module - file.name = data.name - file.size = data.size - file.type = data.type - file.content = data.content + file.subject_id = subject_id + file.module_id = module_id + file.name = data.name + file.size = data.size + file.type = data.type + file.content = data.content - await db.commit() + await db.commit() + return True - return True + except SQLAlchemyError: + await db.rollback() + raise async def get_file(db: AsyncSession, file_id: int) -> LibraryFileResponse: - stmt = (select(LibraryFile).where(LibraryFile.id == file_id)) - result = await db.execute(stmt) - file = result.scalar_one_or_none() + stmt = ( + select(LibraryFile) + .options( + joinedload(LibraryFile.subject), + joinedload(LibraryFile.module), + ) + .where(LibraryFile.id == file_id) + ) + + file = (await db.execute(stmt)).scalar_one_or_none() + if not file: + raise HTTPException(status_code=404, detail="File not found") return LibraryFileResponse( id=file.id, - subject=file.subject, - module=file.module, + subject=file.subject.name, + module=file.module.name, name=file.name, type=file.type, size=file.size, content=file.content, uploadTime=file.create_time.strftime("%Y-%m-%d %H:%M:%S"), ) + + +async def get_or_create_subject(db: AsyncSession, name: str) -> int: + stmt = select(Subject).where(Subject.name == name) + result = await db.execute(stmt) + subject = result.scalar_one_or_none() + + if subject: + return subject.id + + subject = Subject(name=name) + db.add(subject) + await db.flush() + + return subject.id + + +async def get_or_create_module(db: AsyncSession, subject_id: int, name: str) -> int: + stmt = select(Module).where( + Module.subject_id == subject_id, + Module.name == name, + ) + module = (await db.execute(stmt)).scalar_one_or_none() + + if module: + return module.id + + module = Module(subject_id=subject_id, name=name) + db.add(module) + await db.flush() + + return module.id diff --git a/service/mistake_service.py b/service/mistake_service.py index e32ef54..f25d18d 100644 --- a/service/mistake_service.py +++ b/service/mistake_service.py @@ -1,47 +1,67 @@ from sqlalchemy import select, delete +from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.orm import joinedload -from models.library import Mistake +from models.library import Mistake, Subject from schemas.mistake import MistakeResponse, MistakeRequest +from service import subject_service -async def list_mistakes(db: AsyncSession) -> list[MistakeResponse]: - stmt = select(Mistake).order_by(Mistake.id.desc()) - mistakes = (await db.execute(stmt)).scalars().all() +async def list_mistakes(subject: str, db: AsyncSession) -> list[MistakeResponse]: + subject_id = await subject_service.get_subject_by_name(db, subject) + stmt = ( + select(Mistake) + .options(joinedload(Mistake.subject)) + .where(Mistake.subject_id == subject_id) + .order_by(Mistake.id.desc()) + ) + + mistakes = (await db.execute(stmt)).scalars().unique().all() return [ MistakeResponse( id=mistake.id, - subject=mistake.subject, - module=mistake.module, - name=mistake.name, type=mistake.type, question=mistake.question, options=mistake.options, - answer=mistake.answer + answer=mistake.answer, ) for mistake in mistakes ] -async def add_mistakes(db: AsyncSession, data: list[MistakeRequest]): - db.add_all( - [ - Mistake( - subject=item.subject, - module=item.module, - name=item.name, - type=item.type, - question=item.question, - options=item.options, - answer=item.answer, - ) - for item in data - ] - ) +async def add_mistakes(db: AsyncSession, subject: str, data: list[MistakeRequest]): + subject_id = await subject_service.get_subject_by_name(db, subject) - await db.commit() - return True + subject = ( + await db.execute( + select(Subject).where(Subject.id == subject_id) + ) + ).scalar_one_or_none() + + if not subject: + raise ValueError("Subject not found") + + mistakes = [ + Mistake( + subject_id=subject_id, + type=item.type, + question=item.question, + options=item.options, + answer=item.answer, + ) + for item in data + ] + + try: + db.add_all(mistakes) + await db.flush() + await db.commit() + return True + except SQLAlchemyError: + await db.rollback() + raise async def remove_mistake(db: AsyncSession, mistake_id: int) -> bool: diff --git a/service/practice_service.py b/service/practice_service.py deleted file mode 100644 index 5ecba50..0000000 --- a/service/practice_service.py +++ /dev/null @@ -1,38 +0,0 @@ -from typing import List - -from sqlalchemy import select -from sqlalchemy.ext.asyncio import AsyncSession - -from models.library import PracticeRecord -from schemas.practice import PracticeRecordRequest, PracticeRecordResponse - - -async def get_record(db: AsyncSession, library_file_id: int) -> List[PracticeRecordResponse]: - stmt = select(PracticeRecord).where(PracticeRecord.library_file_id == library_file_id).order_by(PracticeRecord.create_time.desc()) - records = (await db.execute(stmt)).scalars().all() - - return [ - PracticeRecordResponse( - id=record.id, - libraryFileId=record.library_file_id, - totalCount=record.total_count, - correctCount=record.correct_count, - wrongCount=record.wrong_count, - createTime=record.create_time.strftime("%Y-%m-%d %H:%M:%S"), - ) - for record in records - ] - - -async def add_record(db: AsyncSession, data: PracticeRecordRequest): - record = PracticeRecord( - library_file_id=data.libraryFileId, - total_count=data.totalCount, - correct_count=data.correctCount, - wrong_count=data.wrongCount, - ) - - db.add(record) - await db.commit() - - return True diff --git a/service/record_service.py b/service/record_service.py new file mode 100644 index 0000000..962cc99 --- /dev/null +++ b/service/record_service.py @@ -0,0 +1,53 @@ +from typing import List + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.orm import joinedload + +from models.library import Record +from schemas.record import RecordResponse, RecordRequest +from service import subject_service + + +async def get_record(db: AsyncSession, subject: str) -> List[RecordResponse]: + subject_id = await subject_service.get_subject_by_name(db, subject) + + stmt = ( + select(Record) + .options(joinedload(Record.subject)) + .where(Record.subject_id == subject_id) + .order_by(Record.create_time.desc()) + ) + + records = (await db.execute(stmt)).scalars().unique().all() + + return [ + RecordResponse( + id=record.id, + type=record.type, + totalCount=record.total_count, + correctCount=record.correct_count, + wrongCount=record.wrong_count, + createTime=record.create_time.strftime("%Y-%m-%d %H:%M:%S"), + ) + for record in records + ] + + +async def add_record(db: AsyncSession, subject: str, data: RecordRequest): + subject_id = await subject_service.get_subject_by_name(db, subject) + + record = Record( + subject_id=subject_id, + type=data.type, + total_count=data.totalCount, + correct_count=data.correctCount, + wrong_count=data.wrongCount, + ) + + db.add(record) + await db.flush() + await db.refresh(record) + await db.commit() + + return True diff --git a/service/subject_service.py b/service/subject_service.py new file mode 100644 index 0000000..0c7deda --- /dev/null +++ b/service/subject_service.py @@ -0,0 +1,13 @@ +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from models.library import Subject + + +async def get_subject_by_name(db: AsyncSession, subject: str) -> int: + stmt = select(Subject).where(Subject.name == subject) + subject = (await db.execute(stmt)).scalar_one_or_none() + if not subject: + raise ValueError("Subject not found") + + return subject.id