import typing from typing import List, Dict from fastapi import FastAPI, Depends from pydantic import BaseModel from starlette.middleware.cors import CORSMiddleware from starlette.requests import Request from starlette.responses import Response, JSONResponse import httpx from anki import AnkiClient, CookieStorage app = FastAPI( title='anki', docs_url='/') app.add_middleware(CORSMiddleware, allow_origins=['*'], allow_methods=['*'], allow_headers=['*']) class SessionCookieStorage(CookieStorage): def __init__(self, req: Request, res: Response): self.res = res self.req = req def save_cookies(self, cookies: dict): for k, v in cookies.items(): self.res.set_cookie(k, v, max_age=24 * 60 * 60) def load_cookies(self) -> typing.Optional[dict]: if 'ankiweb' not in self.req.cookies: return None return {**self.req.cookies} def get_anki_client(req: Request, res: Response) -> AnkiClient: return AnkiClient(SessionCookieStorage(req, res)) @app.post('/login', response_model=Dict[str, str]) async def login(username: str, password: str, anki: AnkiClient = Depends(get_anki_client)): anki.login(username, password) return anki.session @app.get('/info', summary='Get note types and their fields', response_model=List[Dict]) async def info(anki: AnkiClient = Depends(get_anki_client)): return anki.get_editor_context() class CreateNote(BaseModel): deck: str note_type: str note_fields: Dict[str, str] tags: str = '' @app.post('/notes', summary='Create a note') async def create_note(input: CreateNote, anki: AnkiClient = Depends(get_anki_client)): anki.create_note(note_type=input.note_type, deck=input.deck, fields=input.note_fields, tags=input.tags) @app.exception_handler(Exception) async def handle_errors(req: Request, error: Exception): if isinstance(error, httpx.HTTPStatusError): message = 'anki returned error' status = error.response.status_code elif isinstance(error, PermissionError): message = str(error) status = 401 else: message = 'oops' status = 500 return JSONResponse(content={ 'message': message }, status_code=status) if __name__ == '__main__': import uvicorn uvicorn.run(app, host='0.0.0.0')