86 lines
2.8 KiB
Python
86 lines
2.8 KiB
Python
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)
|