feat: strengthen backup and mobile workflows
This commit is contained in:
+65
-12
@@ -11,7 +11,7 @@ from zoneinfo import ZoneInfo
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Query, Response, UploadFile
|
||||
from fastapi.responses import FileResponse
|
||||
from pydantic import BaseModel, Field, StrictBool, field_validator, model_validator
|
||||
from sqlalchemy import case, delete, func, select, update
|
||||
from sqlalchemy import case, func, select, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from .auth import current_user
|
||||
@@ -44,6 +44,42 @@ from .models import (
|
||||
|
||||
router = APIRouter(prefix="/api/v1")
|
||||
|
||||
LEGACY_BACKUP_MAX_BYTES = 16 * 1024 * 1024
|
||||
LEGACY_BACKUP_MAX_RECORDS = 10_000
|
||||
LEGACY_BACKUP_MAX_FIELD_BYTES = 1024 * 1024
|
||||
_LEGACY_READ_CHUNK = 64 * 1024
|
||||
_LEGACY_ENTITIES = ("folders", "lists", "tasks", "recurrences", "habits", "countdowns", "memos")
|
||||
|
||||
|
||||
def _legacy_error(code: str, message: str) -> HTTPException:
|
||||
return HTTPException(422, {"code": code, "message": message})
|
||||
|
||||
|
||||
def _validate_legacy_payload_limits(payload: dict) -> None:
|
||||
total = 0
|
||||
for entity in _LEGACY_ENTITIES:
|
||||
rows = payload.get(entity, [])
|
||||
if not isinstance(rows, list):
|
||||
raise _legacy_error("legacy_backup_invalid", "旧版备份实体格式无效")
|
||||
total += len(rows)
|
||||
if total > LEGACY_BACKUP_MAX_RECORDS:
|
||||
raise _legacy_error("legacy_backup_too_many_records", "旧版备份记录过多")
|
||||
for row in rows:
|
||||
if not isinstance(row, dict):
|
||||
raise _legacy_error("legacy_backup_invalid", "旧版备份记录格式无效")
|
||||
for value in row.values():
|
||||
if isinstance(value, str) and len(value.encode("utf-8")) > LEGACY_BACKUP_MAX_FIELD_BYTES:
|
||||
raise _legacy_error("legacy_backup_field_too_large", "旧版备份字段过大")
|
||||
|
||||
|
||||
async def _read_legacy_upload(file: UploadFile) -> bytes:
|
||||
content = bytearray()
|
||||
while chunk := await file.read(_LEGACY_READ_CHUNK):
|
||||
content.extend(chunk)
|
||||
if len(content) > LEGACY_BACKUP_MAX_BYTES:
|
||||
raise _legacy_error("legacy_backup_too_large", "旧版备份文件过大")
|
||||
return bytes(content)
|
||||
|
||||
|
||||
def audit(db: AsyncSession, user_id: UUID, action: str, entity_type: str, entity_id=None, **details):
|
||||
db.add(AuditLog(user_id=user_id, action=action, entity_type=entity_type, entity_id=entity_id, details=details))
|
||||
@@ -223,8 +259,10 @@ class RecurrenceCreate(BaseModel):
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_mode(self):
|
||||
if self.trigger_mode == "scheduled" and self.rrule is None:
|
||||
raise ValueError("scheduled recurrence requires rrule")
|
||||
if self.trigger_mode == "scheduled" and (
|
||||
self.rrule is None or self.after_completion_days is not None
|
||||
):
|
||||
raise ValueError("scheduled recurrence requires rrule and no completion interval")
|
||||
if self.trigger_mode == "after_completion" and (
|
||||
self.after_completion_days is None or self.rrule is not None
|
||||
):
|
||||
@@ -1209,6 +1247,16 @@ async def habit_stats(habit_id: UUID, user: User = Depends(current_user), db: As
|
||||
_ALLOWED_MIME = {"text/plain", "text/csv", "application/pdf", "image/jpeg", "image/png", "image/gif", "application/json", "application/zip"}
|
||||
|
||||
|
||||
@router.get("/tasks/{task_id}/attachments")
|
||||
async def list_attachments(task_id: UUID, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
||||
await owned_task(db, user.id, task_id)
|
||||
rows = (await db.scalars(select(Attachment).where(
|
||||
Attachment.task_id == task_id, Attachment.user_id == user.id
|
||||
).order_by(Attachment.created_at, Attachment.id))).all()
|
||||
return [{"id": row.id, "task_id": row.task_id, "filename": row.filename,
|
||||
"mime_type": row.mime_type, "size": row.size} for row in rows]
|
||||
|
||||
|
||||
@router.post("/tasks/{task_id}/attachments", status_code=201)
|
||||
async def upload_attachment(task_id: UUID, file: UploadFile = File(...), user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
||||
await owned_task(db, user.id, task_id)
|
||||
@@ -1349,7 +1397,11 @@ async def restore_csv(
|
||||
user: User = Depends(current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
text = (await file.read()).decode("utf-8-sig")
|
||||
try:
|
||||
raw_content = await _read_legacy_upload(file)
|
||||
text = raw_content.decode("utf-8-sig")
|
||||
except UnicodeDecodeError as exc:
|
||||
raise _legacy_error("legacy_backup_invalid", "无效的 Dodo CSV 备份") from exc
|
||||
payload = {
|
||||
"version": 1,
|
||||
"folders": [],
|
||||
@@ -1363,9 +1415,14 @@ async def restore_csv(
|
||||
try:
|
||||
for row in csv.DictReader(io.StringIO(text)):
|
||||
entity = row.get("entity", "")
|
||||
data = row.get("data")
|
||||
if entity not in payload or entity == "version":
|
||||
raise ValueError("unknown entity")
|
||||
payload[entity].append(json.loads(row["data"]))
|
||||
if not isinstance(data, str) or len(data.encode("utf-8")) > LEGACY_BACKUP_MAX_FIELD_BYTES:
|
||||
raise ValueError("field too large")
|
||||
payload[entity].append(json.loads(data))
|
||||
if sum(len(payload[name]) for name in _LEGACY_ENTITIES) > LEGACY_BACKUP_MAX_RECORDS:
|
||||
raise _legacy_error("legacy_backup_too_many_records", "旧版备份记录过多")
|
||||
except (csv.Error, json.JSONDecodeError, KeyError, TypeError, ValueError) as exc:
|
||||
raise HTTPException(422, "无效的 Dodo CSV 备份") from exc
|
||||
return await restore_json(payload, mode, user, db)
|
||||
@@ -1436,16 +1493,12 @@ def _validate_countdown_backups(payload: dict) -> list[dict]:
|
||||
|
||||
@router.post("/restore")
|
||||
async def restore_json(payload: dict, mode: str = Query("merge", pattern="^(merge|replace)$"), user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
|
||||
if mode == "replace":
|
||||
raise _legacy_error("legacy_replace_unsupported", "旧版备份仅支持合并恢复")
|
||||
if payload.get("version") != 1:
|
||||
raise HTTPException(422, "不支持的备份版本")
|
||||
_validate_legacy_payload_limits(payload)
|
||||
parsed_countdowns = _validate_countdown_backups(payload)
|
||||
if mode == "replace":
|
||||
await db.execute(delete(Memo).where(Memo.user_id == user.id))
|
||||
await db.execute(delete(Countdown).where(Countdown.user_id == user.id))
|
||||
await db.execute(delete(Task).where(Task.user_id == user.id))
|
||||
await db.execute(delete(Habit).where(Habit.user_id == user.id))
|
||||
await db.execute(delete(TaskList).where(TaskList.user_id == user.id))
|
||||
await db.execute(delete(Folder).where(Folder.user_id == user.id))
|
||||
id_map = {}
|
||||
task_id_map = {}
|
||||
for raw in payload.get("folders", []):
|
||||
|
||||
Reference in New Issue
Block a user