from typing import List from fastapi import FastAPI, UploadFile, File, Depends, Form from sqlalchemy.ext.asyncio import AsyncSession 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 app = FastAPI(title="AI Study Service") app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_methods=["*"], 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)