diff --git a/agent/prompt.py b/agent/prompt.py index 6665995..d1e27a9 100644 --- a/agent/prompt.py +++ b/agent/prompt.py @@ -1,4 +1,4 @@ -study_prompt = """ +question_prompt = """ 你是一位专业的出题专家,擅长根据学习材料生成高质量的客观题(单选题和多选题)。 请根据下方【学习材料】生成 5~10 道题目,严格遵守以下要求: @@ -50,4 +50,35 @@ study_prompt = """ "select": [] } ] +""" + +knowledge_prompt = """ +你是一位资深教学与教研专家,擅长将各类专业知识系统化梳理,转化为结构清晰、易懂好记的学习资料,适用于不同基础的学习者。 + +请根据下方【学科】·【模块】·【名称】生成一份高质量学习资料,**以 Markdown 格式输出**,用于系统学习与复习。 + +要求如下: +1. 全文控制在 **500~800字**(可根据内容复杂度±10%),语言准确、条理清晰,适合打印或电子阅读。 +2. Markdown 结构必须包含以下一级或二级标题: + - ## 模块概述 + (定义、作用、学习意义,可简要说明该模块在主题中的地位) + - ## 核心知识点 + (分点讲解,重点概念用 **粗体** 标注) + - ## 对比与辨析 + (对比易混淆概念) + - ## 记忆口诀 / 核心规律 + (简洁押韵或逻辑清晰的归纳,便于记忆) + - ## 本章小结 + (总结本模块最核心的要点) +3. 内容需适配当前主题的属性: + - 理工类:突出公式、定理、推导过程、解题思路 + - 医学类:突出机制、适应症、禁忌症、注意事项 + - 文科 / 法商 / 管理类:突出概念界定、逻辑框架、典型案例 +4. 写作风格: + - 严谨、客观,贴近国内主流教材表述 + - 避免口语化(如“同学们”“大家注意啦”) + - 不假设学习者是在校大学生,兼顾自学者与在职进修人员 +5. 输出要求: + - 仅输出 Markdown 正文,不要包含 ``、JSON 或其他格式 + - 不要在开头重复本提示词内容 """ \ No newline at end of file diff --git a/agent/study.py b/agent/study.py index d21b9e9..cfbd2f2 100644 --- a/agent/study.py +++ b/agent/study.py @@ -1,9 +1,14 @@ from langchain.agents import create_agent from agent.model import model -from agent.prompt import study_prompt +from agent.prompt import question_prompt, knowledge_prompt -study_agent = create_agent( +question_agent = create_agent( model=model, - system_prompt=study_prompt + system_prompt=question_prompt +) + +knowledge_agent = create_agent( + model=model, + system_prompt=knowledge_prompt ) diff --git a/main.py b/main.py index 34e85c4..33d7351 100644 --- a/main.py +++ b/main.py @@ -1,19 +1,7 @@ -from typing import List - -from fastapi import FastAPI, UploadFile, File, Depends, Form -from sqlalchemy.ext.asyncio import AsyncSession +from fastapi import FastAPI from starlette.middleware.cors import CORSMiddleware -from starlette.responses import StreamingResponse -from database import get_db -from schemas.agent import QueryRequest -from schemas.library import LibraryFileRequest, LibraryFileResponse -from schemas.mistake import MistakeResponse, MistakeRequest -from schemas.practice import PracticeRecordResponse, PracticeRecordRequest -from service import library_service, practice_service, mistake_service -from service.agent_service import query_agent -from storage import upload_rustfs -from ocr import ocr_from_url +from routers import routers app = FastAPI(title="AI Study Service") @@ -24,62 +12,5 @@ app.add_middleware( allow_headers=["*"], ) - -@app.post("/library/upload") -def upload(file: UploadFile = File(...), filename: str = Form(...)): - url = upload_rustfs(file.file, filename) - return ocr_from_url(url) - - -@app.post("/library/question") -def generate(query: QueryRequest): - """流式对话""" - return StreamingResponse( - query_agent(query), - media_type="text/event-stream" - ) - - -@app.get("/library/files", response_model=List[LibraryFileResponse]) -async def list_files(db: AsyncSession = Depends(get_db)): - return await library_service.list_files(db) - - -@app.get("/library/files/{id}", response_model=LibraryFileResponse) -async def get_file(id: int, db: AsyncSession = Depends(get_db)): - return await library_service.get_file(db, id) - - -@app.post("/library/files", response_model=bool) -async def add_file(data: LibraryFileRequest, db: AsyncSession = Depends(get_db)): - return await library_service.add_file(db, data) - - -@app.put("/library/files/{id}", response_model=bool) -async def edit_file(id: int, data: LibraryFileRequest, db: AsyncSession = Depends(get_db)): - return await library_service.edit_file(db, id, data) - - -@app.get("/practice/record/{id}", response_model=List[PracticeRecordResponse]) -async def get_record(id: int, db: AsyncSession = Depends(get_db)): - return await practice_service.get_record(db, id) - - -@app.post("/practice/record", response_model=bool) -async def add_record(data: PracticeRecordRequest, db: AsyncSession = Depends(get_db)): - return await practice_service.add_record(db, data) - - -@app.get("/mistake", response_model=List[MistakeResponse]) -async def list_mistakes(db: AsyncSession = Depends(get_db)): - return await mistake_service.list_mistakes(db) - - -@app.post("/mistake", response_model=bool) -async def list_mistakes(data: list[MistakeRequest], db: AsyncSession = Depends(get_db)): - return await mistake_service.add_mistakes(db, data) - - -@app.delete("/mistake/{id}", response_model=bool) -async def remove_mistake(id: int, db: AsyncSession = Depends(get_db)): - return await mistake_service.remove_mistake(db, id) +for router in routers: + app.include_router(router) diff --git a/routers/__init__.py b/routers/__init__.py new file mode 100644 index 0000000..0fed23e --- /dev/null +++ b/routers/__init__.py @@ -0,0 +1,9 @@ +from .library import router as library_router +from .practice import router as practice_router +from .mistake import router as mistake_router + +routers = [ + library_router, + practice_router, + mistake_router, +] diff --git a/routers/library.py b/routers/library.py new file mode 100644 index 0000000..b1c8a27 --- /dev/null +++ b/routers/library.py @@ -0,0 +1,55 @@ +from typing import List +from fastapi import APIRouter, Depends, File, Form, UploadFile +from sqlalchemy.ext.asyncio import AsyncSession +from starlette.responses import StreamingResponse + +from database import get_db +from schemas.library import LibraryFileRequest, LibraryFileResponse +from schemas.agent import QueryQuestionRequest, QueryKnowledgeRequest +from service import library_service, agent_service +from storage import upload_rustfs +from ocr import ocr_from_url + +router = APIRouter(prefix="/library", tags=["Library"]) + + +@router.post("/upload") +def upload(file: UploadFile = File(...), filename: str = Form(...)): + url = upload_rustfs(file.file, filename) + return ocr_from_url(url) + + +@router.post("/question") +def generate_question(query: QueryQuestionRequest): + return StreamingResponse( + agent_service.query_question_agent(query), + media_type="text/event-stream", + ) + + +@router.post("/knowledge") +def generate_knowledge(query: QueryKnowledgeRequest): + return StreamingResponse( + agent_service.query_knowledge_agent(query), + media_type="text/event-stream", + ) + + +@router.get("/files", response_model=List[LibraryFileResponse]) +async def list_files(db: AsyncSession = Depends(get_db)): + return await library_service.list_files(db) + + +@router.get("/files/{file_id}", response_model=LibraryFileResponse) +async def get_file(file_id: int, db: AsyncSession = Depends(get_db)): + return await library_service.get_file(db, file_id) + + +@router.post("/files", response_model=bool) +async def add_file(data: LibraryFileRequest, db: AsyncSession = Depends(get_db)): + return await library_service.add_file(db, data) + + +@router.put("/files/{file_id}", response_model=bool) +async def edit_file(file_id: int, data: LibraryFileRequest, db: AsyncSession = Depends(get_db)): + return await library_service.edit_file(db, file_id, data) diff --git a/routers/mistake.py b/routers/mistake.py new file mode 100644 index 0000000..15089c9 --- /dev/null +++ b/routers/mistake.py @@ -0,0 +1,24 @@ +from typing import List +from fastapi import APIRouter, Depends +from sqlalchemy.ext.asyncio import AsyncSession + +from database import get_db +from schemas.mistake import MistakeResponse, MistakeRequest +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.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.delete("/{mistake_id}", response_model=bool) +async def remove_mistake(mistake_id: int, db: AsyncSession = Depends(get_db)): + return await mistake_service.remove_mistake(db, mistake_id) diff --git a/routers/practice.py b/routers/practice.py new file mode 100644 index 0000000..c67e41d --- /dev/null +++ b/routers/practice.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.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/schemas/agent.py b/schemas/agent.py index 4f39ab8..c0b46ec 100644 --- a/schemas/agent.py +++ b/schemas/agent.py @@ -1,5 +1,11 @@ from pydantic import BaseModel -class QueryRequest(BaseModel): +class QueryQuestionRequest(BaseModel): message: str + + +class QueryKnowledgeRequest(BaseModel): + subject: str + module: str + name: str diff --git a/service/agent_service.py b/service/agent_service.py index f19e5d8..82f1d4e 100644 --- a/service/agent_service.py +++ b/service/agent_service.py @@ -1,15 +1,32 @@ from langchain_core.messages import AIMessageChunk, HumanMessage -from agent.study import study_agent -from schemas.agent import QueryRequest +from agent.study import question_agent, knowledge_agent +from schemas.agent import QueryQuestionRequest, QueryKnowledgeRequest -def query_agent(query: QueryRequest): +def query_question_agent(query: QueryQuestionRequest): try: user_msg = f"学习内容:\n{query.message}" # 流式调用Agent - for chunk, metadata in study_agent.stream( + for chunk, metadata in question_agent.stream( + {"messages": [HumanMessage(content=user_msg)]}, + stream_mode="messages" + ): + if isinstance(chunk, AIMessageChunk): + if isinstance(chunk, AIMessageChunk) and chunk.content: + yield chunk.content + except Exception as e: + print(f"\n[错误]: {str(e)}") + yield "信息检索失败,请重新输入问题提问" + + +def query_knowledge_agent(query: QueryKnowledgeRequest): + try: + user_msg = f"学科:{query.subject},模块:{query.module}, 名称:{query.name}" + + # 流式调用Agent + for chunk, metadata in knowledge_agent.stream( {"messages": [HumanMessage(content=user_msg)]}, stream_mode="messages" ):