Files
study-agent-service/main.py
2026-07-30 13:07:26 +08:00

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)