219 lines
8.1 KiB
Python
219 lines
8.1 KiB
Python
from contextlib import asynccontextmanager
|
|
from pathlib import Path
|
|
from uuid import UUID
|
|
|
|
from fastapi import Depends, FastAPI, HTTPException, Response
|
|
from fastapi.responses import FileResponse
|
|
from fastapi.staticfiles import StaticFiles
|
|
from sqlalchemy import func, select, update
|
|
from sqlalchemy.exc import IntegrityError
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from .auth import (
|
|
COOKIE_NAME,
|
|
current_user,
|
|
hash_password,
|
|
hash_token,
|
|
issue_session,
|
|
session_token,
|
|
verify_password,
|
|
)
|
|
from .db import create_schema, get_db
|
|
from .models import AppState, Folder, Session, Task, TaskList, User, utcnow
|
|
from .schemas import (
|
|
FolderCreate,
|
|
FolderOut,
|
|
InitializeRequest,
|
|
ListCreate,
|
|
ListOut,
|
|
LoginRequest,
|
|
TaskCreate,
|
|
TaskOut,
|
|
TaskPage,
|
|
TaskUpdate,
|
|
UserOut,
|
|
)
|
|
|
|
|
|
@asynccontextmanager
|
|
async def lifespan(app: FastAPI):
|
|
from .config import get_settings
|
|
|
|
if get_settings().auto_create_schema:
|
|
await create_schema()
|
|
yield
|
|
|
|
|
|
app = FastAPI(title="dodo", version="0.1.0", lifespan=lifespan, docs_url="/api/docs", openapi_url="/api/openapi.json")
|
|
|
|
|
|
@app.get("/health/live")
|
|
async def live():
|
|
return {"status": "ok"}
|
|
|
|
|
|
@app.get("/health/ready")
|
|
async def ready(db: AsyncSession = Depends(get_db)):
|
|
await db.execute(select(1))
|
|
return {"status": "ok"}
|
|
|
|
|
|
@app.get("/api/v1/setup/status")
|
|
async def setup_status(db: AsyncSession = Depends(get_db)):
|
|
count = await db.scalar(select(func.count()).select_from(User))
|
|
return {"initialized": bool(count)}
|
|
|
|
|
|
@app.post("/api/v1/setup/initialize", response_model=UserOut, status_code=201)
|
|
async def initialize(payload: InitializeRequest, response: Response, db: AsyncSession = Depends(get_db)):
|
|
count = await db.scalar(select(func.count()).select_from(User))
|
|
if count:
|
|
raise HTTPException(status_code=409, detail="系统已经初始化")
|
|
db.add(AppState(key="initialized"))
|
|
user = User(username=payload.username, password_hash=hash_password(payload.password))
|
|
db.add(user)
|
|
try:
|
|
await db.flush()
|
|
except IntegrityError as exc:
|
|
await db.rollback()
|
|
raise HTTPException(status_code=409, detail="系统已经初始化") from exc
|
|
db.add(TaskList(user_id=user.id, name="收集箱", is_inbox=True))
|
|
await db.commit()
|
|
await db.refresh(user)
|
|
await issue_session(db, response, user)
|
|
return user
|
|
|
|
|
|
@app.post("/api/v1/auth/login", response_model=UserOut)
|
|
async def login(payload: LoginRequest, response: Response, db: AsyncSession = Depends(get_db)):
|
|
user = await db.scalar(select(User).where(User.username == payload.username))
|
|
if user is None or not verify_password(user.password_hash, payload.password):
|
|
raise HTTPException(status_code=401, detail="用户名或密码错误")
|
|
await issue_session(db, response, user)
|
|
return user
|
|
|
|
|
|
@app.get("/api/v1/me", response_model=UserOut)
|
|
async def me(user: User = Depends(current_user)):
|
|
return user
|
|
|
|
|
|
@app.post("/api/v1/auth/logout", status_code=204)
|
|
async def logout(
|
|
response: Response,
|
|
token: str = Depends(session_token),
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
session = await db.scalar(select(Session).where(Session.token_hash == hash_token(token)))
|
|
if session is not None:
|
|
await db.delete(session)
|
|
await db.commit()
|
|
response.delete_cookie(COOKIE_NAME, path="/")
|
|
return Response(status_code=204, headers=response.headers)
|
|
|
|
|
|
@app.post("/api/v1/folders", response_model=FolderOut, status_code=201)
|
|
async def create_folder(payload: FolderCreate, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
|
folder = Folder(user_id=user.id, name=payload.name)
|
|
db.add(folder); await db.commit(); await db.refresh(folder)
|
|
return folder
|
|
|
|
|
|
@app.get("/api/v1/folders", response_model=list[FolderOut])
|
|
async def list_folders(user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
|
return list((await db.scalars(select(Folder).where(Folder.user_id == user.id).order_by(Folder.position, Folder.created_at))).all())
|
|
|
|
|
|
@app.post("/api/v1/lists", response_model=ListOut, status_code=201)
|
|
async def create_list(payload: ListCreate, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
|
if payload.folder_id and not await db.scalar(select(Folder.id).where(Folder.id == payload.folder_id, Folder.user_id == user.id)):
|
|
raise HTTPException(status_code=404, detail="文件夹不存在")
|
|
item = TaskList(user_id=user.id, folder_id=payload.folder_id, name=payload.name)
|
|
db.add(item); await db.commit(); await db.refresh(item)
|
|
return item
|
|
|
|
|
|
@app.get("/api/v1/lists", response_model=list[ListOut])
|
|
async def list_lists(user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
|
return list((await db.scalars(select(TaskList).where(TaskList.user_id == user.id).order_by(TaskList.is_inbox.desc(), TaskList.position, TaskList.created_at))).all())
|
|
|
|
|
|
@app.post("/api/v1/tasks", response_model=TaskOut, status_code=201)
|
|
async def create_task(payload: TaskCreate, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
|
if not await db.scalar(select(TaskList.id).where(TaskList.id == payload.list_id, TaskList.user_id == user.id)):
|
|
raise HTTPException(status_code=404, detail="清单不存在")
|
|
if payload.parent_id:
|
|
parent = await db.scalar(
|
|
select(Task).where(
|
|
Task.id == payload.parent_id,
|
|
Task.user_id == user.id,
|
|
Task.list_id == payload.list_id,
|
|
Task.parent_id.is_(None),
|
|
Task.deleted_at.is_(None),
|
|
)
|
|
)
|
|
if parent is None:
|
|
raise HTTPException(status_code=400, detail="父任务必须属于同一清单")
|
|
task = Task(user_id=user.id, **payload.model_dump())
|
|
db.add(task); await db.commit(); await db.refresh(task)
|
|
return task
|
|
|
|
|
|
@app.get("/api/v1/tasks", response_model=TaskPage)
|
|
async def list_tasks(user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
|
items = list((await db.scalars(select(Task).where(Task.user_id == user.id, Task.deleted_at.is_(None), Task.parent_id.is_(None)).order_by(Task.completed, Task.position, Task.created_at))).all())
|
|
return TaskPage(items=items)
|
|
|
|
|
|
@app.patch("/api/v1/tasks/{task_id}", response_model=TaskOut)
|
|
async def update_task(task_id: UUID, payload: TaskUpdate, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
|
data = payload.model_dump(exclude_unset=True)
|
|
expected_version = data.pop("version")
|
|
data["version"] = Task.version + 1
|
|
data["updated_at"] = utcnow()
|
|
result = await db.execute(
|
|
update(Task)
|
|
.where(
|
|
Task.id == task_id,
|
|
Task.user_id == user.id,
|
|
Task.deleted_at.is_(None),
|
|
Task.version == expected_version,
|
|
)
|
|
.values(**data)
|
|
.returning(Task)
|
|
)
|
|
task = result.scalar_one_or_none()
|
|
if task is None:
|
|
exists = await db.scalar(
|
|
select(Task.id).where(Task.id == task_id, Task.user_id == user.id, Task.deleted_at.is_(None))
|
|
)
|
|
if exists:
|
|
raise HTTPException(status_code=409, detail="任务已被更新,请刷新后重试")
|
|
raise HTTPException(status_code=404, detail="任务不存在")
|
|
await db.commit()
|
|
return task
|
|
|
|
|
|
@app.delete("/api/v1/tasks/{task_id}", status_code=204)
|
|
async def delete_task(task_id: UUID, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
|
task = await db.scalar(select(Task).where(Task.id == task_id, Task.user_id == user.id, Task.deleted_at.is_(None)))
|
|
if task is None:
|
|
raise HTTPException(status_code=404, detail="任务不存在")
|
|
task.deleted_at = utcnow()
|
|
task.version += 1
|
|
await db.commit()
|
|
return Response(status_code=204)
|
|
|
|
|
|
static_dir = Path(__file__).parent / "static"
|
|
if static_dir.exists():
|
|
app.mount("/assets", StaticFiles(directory=static_dir / "assets"), name="assets")
|
|
|
|
@app.get("/{path:path}", include_in_schema=False)
|
|
async def spa(path: str):
|
|
root = static_dir.resolve()
|
|
target = (root / path).resolve()
|
|
if target.is_file() and target.is_relative_to(root):
|
|
return FileResponse(target)
|
|
return FileResponse(root / "index.html")
|