from typing import Optional, Union from pathlib import Path from fastapi import APIRouter, Depends, File, HTTPException, UploadFile from sqlalchemy.orm import Session from auth import get_current_user from database import get_db from models import User from schemas import ( MemoryTransformerPredictResponse, MemoryVisualizationResponse, WordBatchCreateRequest, WordBatchCreateResponse, WordCreate, WordListPageOut, WordMemoryDetailResponse, WordOut, WordScanItem, WordScanResponse, WordScanTextRequest, WordUpdate, ) from services.memory_transformer import memory_transformer_service from services.memory_visual_service import memory_visual_service from services.word_scan_service import word_scan_service from services.word_service import word_service router = APIRouter(prefix="/api/words", tags=["words"]) @router.post("", response_model=WordOut) def create_word( data: WordCreate, db: Session = Depends(get_db), current_user: User = Depends(get_current_user), ): word = word_service.create_word(db, current_user, data.model_dump()) return word @router.post("/scan-text", response_model=WordScanResponse) def scan_word_text( data: WordScanTextRequest, current_user: User = Depends(get_current_user), ): items = word_scan_service.parse_scan_text(data.text) return WordScanResponse(items=[WordScanItem(**item) for item in items]) @router.post("/scan-image", response_model=WordScanResponse) async def scan_word_image( file: UploadFile = File(...), current_user: User = Depends(get_current_user), ): data = await file.read() if len(data) > 10 * 1024 * 1024: raise HTTPException(status_code=400, detail="图片不能超过 10MB") suffix = Path(file.filename or "scan.png").suffix or ".png" items = word_scan_service.scan_image(data, suffix=suffix) return WordScanResponse(items=[WordScanItem(**item) for item in items]) @router.post("/batch", response_model=WordBatchCreateResponse) def batch_create_words( data: WordBatchCreateRequest, db: Session = Depends(get_db), current_user: User = Depends(get_current_user), ): result = word_service.batch_create_words( db, current_user, [item.model_dump() for item in data.items], ) return WordBatchCreateResponse( created=result["created"], skipped=result["skipped"], items=result["items"], ) @router.get("", response_model=Union[list[WordOut], WordListPageOut]) def list_words( status: Optional[str] = None, book_id: Optional[int] = None, page: Optional[int] = None, page_size: int = 20, db: Session = Depends(get_db), current_user: User = Depends(get_current_user), ): if page is not None: items, total = word_service.list_words_page( db, current_user, status=status, book_id=book_id, page=page, page_size=page_size, ) return WordListPageOut( items=items, total=total, page=max(1, page), page_size=max(1, min(page_size, 100)), ) return word_service.list_words(db, current_user, status, book_id=book_id) @router.get("/memory-viz", response_model=MemoryVisualizationResponse) def memory_visualization( book_id: int = 0, db: Session = Depends(get_db), current_user: User = Depends(get_current_user), ): return memory_visual_service.get_visualization(db, current_user, book_id=book_id) @router.get("/{word_id}/memory", response_model=WordMemoryDetailResponse) def word_memory( word_id: int, db: Session = Depends(get_db), current_user: User = Depends(get_current_user), ): return memory_visual_service.get_word_memory(db, current_user, word_id) @router.get("/{word_id}/memory-model", response_model=MemoryTransformerPredictResponse) def word_memory_model( word_id: int, db: Session = Depends(get_db), current_user: User = Depends(get_current_user), ): word = word_service.get_word(db, current_user, word_id) return memory_transformer_service.predict_for_word(db, current_user, word) @router.get("/{word_id}", response_model=WordOut) def get_word( word_id: int, db: Session = Depends(get_db), current_user: User = Depends(get_current_user), ): return word_service.get_word(db, current_user, word_id) @router.delete("/{word_id}") def delete_word( word_id: int, db: Session = Depends(get_db), current_user: User = Depends(get_current_user), ): word_service.delete_word(db, current_user, word_id) return {"ok": True} @router.patch("/{word_id}", response_model=WordOut) def update_word( word_id: int, data: WordUpdate, db: Session = Depends(get_db), current_user: User = Depends(get_current_user), ): return word_service.update_word(db, current_user, word_id, data.model_dump(exclude_unset=True))