diff --git a/backend/auth.py b/backend/auth.py index 27ab83e..2c0bc65 100644 --- a/backend/auth.py +++ b/backend/auth.py @@ -39,7 +39,7 @@ async def issue_session(db: AsyncSession, response: Response, user: User, reques db.add(Session( token_hash=hash_token(token), user_id=user.id, expires_at=expires, ip_address=request.client.host if request and request.client else None, - user_agent=request.headers.get("user-agent", "")[:500] if request else None, + user_agent=request.headers.get("user-agent", "")[:255] if request else None, )) await db.commit() response.set_cookie( diff --git a/backend/main.py b/backend/main.py index 9608830..9c1ba9a 100644 --- a/backend/main.py +++ b/backend/main.py @@ -39,6 +39,7 @@ from .schemas import ( ListOut, LoginRequest, NameUpdate, + SessionOut, TagCreate, TagOut, TaskCreate, @@ -213,7 +214,7 @@ async def logout( return Response(status_code=204, headers=response.headers) -@app.get("/api/v1/sessions") +@app.get("/api/v1/sessions", response_model=list[SessionOut]) async def list_sessions( token: str = Depends(session_token), user: User = Depends(current_user), @@ -559,7 +560,7 @@ async def list_tasks( if page is not None: size = page_size or limit items = list((await db.scalars(query.order_by(*ordering).offset((page - 1) * size).limit(size))).all()) - return TaskPage(items=items, total=total, page=page, page_size=size) + return TaskPage(items=await _task_details(db, items), total=total, page=page, page_size=size) if cursor: created_at, task_id = _decode_cursor(cursor) query = query.where( @@ -569,31 +570,48 @@ async def list_tasks( has_more = len(rows) > limit items = rows[:limit] next_cursor = _encode_cursor(items[-1].created_at, items[-1].id) if has_more else None - return TaskPage(items=items, next_cursor=next_cursor, total=total, page=1, page_size=limit) + return TaskPage(items=await _task_details(db, items), next_cursor=next_cursor, total=total, page=1, page_size=limit) -async def _task_detail(db: AsyncSession, task: Task) -> TaskDetailOut: - tags = list( - ( - await db.scalars( - select(Tag) - .join(TaskTag, TaskTag.tag_id == Tag.id) - .where(TaskTag.task_id == task.id) - .order_by(Tag.name, Tag.id) - ) - ).all() - ) +async def _task_details(db: AsyncSession, tasks: list[Task]) -> list[TaskDetailOut]: + if not tasks: + return [] + task_ids = [task.id for task in tasks] + tag_rows = ( + await db.execute( + select(TaskTag.task_id, Tag) + .join(Tag, Tag.id == TaskTag.tag_id) + .where(TaskTag.task_id.in_(task_ids)) + .order_by(Tag.name, Tag.id) + ) + ).all() subtasks = list( ( await db.scalars( select(Task) - .where(Task.parent_id == task.id, Task.deleted_at.is_(None)) + .where(Task.parent_id.in_(task_ids), Task.deleted_at.is_(None)) .order_by(Task.position, Task.created_at, Task.id) ) ).all() ) - data = TaskOut.model_validate(task).model_dump() - return TaskDetailOut(**data, tags=tags, subtasks=subtasks) + tags_by_task: dict[UUID, list[Tag]] = defaultdict(list) + subtasks_by_task: dict[UUID, list[Task]] = defaultdict(list) + for task_id, tag in tag_rows: + tags_by_task[task_id].append(tag) + for subtask in subtasks: + subtasks_by_task[subtask.parent_id].append(subtask) + return [ + TaskDetailOut( + **TaskOut.model_validate(task).model_dump(), + tags=tags_by_task[task.id], + subtasks=subtasks_by_task[task.id], + ) + for task in tasks + ] + + +async def _task_detail(db: AsyncSession, task: Task) -> TaskDetailOut: + return (await _task_details(db, [task]))[0] @app.get("/api/v1/tasks/{task_id}", response_model=TaskDetailOut) @@ -707,7 +725,7 @@ async def list_trash( if page is not None: size = page_size or limit items = list((await db.scalars(query.order_by(*ordering).offset((page - 1) * size).limit(size))).all()) - return TaskPage(items=items, total=total, page=page, page_size=size) + return TaskPage(items=await _task_details(db, items), total=total, page=page, page_size=size) if cursor: created_at, task_id = _decode_cursor(cursor) query = query.where( @@ -717,7 +735,7 @@ async def list_trash( has_more = len(rows) > limit items = rows[:limit] next_cursor = _encode_cursor(items[-1].created_at, items[-1].id) if has_more else None - return TaskPage(items=items, next_cursor=next_cursor, total=total, page=1, page_size=limit) + return TaskPage(items=await _task_details(db, items), next_cursor=next_cursor, total=total, page=1, page_size=limit) @app.post("/api/v1/tasks/{task_id}/restore", response_model=TaskDetailOut) @@ -791,6 +809,9 @@ async def batch_update_tasks( raise HTTPException(status_code=404, detail="一个或多个任务不存在") if payload.list_id is not None: await _owned_list(db, user.id, payload.list_id) + standalone_children = [task for task in tasks if task.parent_id is not None] + if standalone_children: + raise HTTPException(status_code=400, detail="子任务不能脱离父任务单独移动") tag_ids = await _validate_tags(db, user.id, payload.tag_ids) changes = payload.model_dump(exclude_unset=True, exclude={"task_ids", "tag_ids", "soft_delete"}) if payload.soft_delete: @@ -798,9 +819,25 @@ async def batch_update_tasks( if changes: changes["version"] = Task.version + 1 changes["updated_at"] = utcnow() + target_ids = set(task_ids) + if payload.soft_delete: + parent_ids = [task.id for task in tasks if task.parent_id is None] + if parent_ids: + child_ids = list( + ( + await db.scalars( + select(Task.id).where( + Task.parent_id.in_(parent_ids), + Task.user_id == user.id, + Task.deleted_at.is_(None), + ) + ) + ).all() + ) + target_ids.update(child_ids) await db.execute( update(Task) - .where(Task.id.in_(task_ids), Task.user_id == user.id, Task.deleted_at.is_(None)) + .where(Task.id.in_(target_ids), Task.user_id == user.id, Task.deleted_at.is_(None)) .values(**changes) ) if payload.list_id is not None: diff --git a/backend/models.py b/backend/models.py index eb854e5..7ce95eb 100644 --- a/backend/models.py +++ b/backend/models.py @@ -52,7 +52,7 @@ class Session(Base): created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=utcnow) last_seen_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=utcnow) ip_address: Mapped[str | None] = mapped_column(String(64), nullable=True) - user_agent: Mapped[str | None] = mapped_column(String(500), nullable=True) + user_agent: Mapped[str | None] = mapped_column(String(255), nullable=True) class Folder(Base): diff --git a/backend/mvp.py b/backend/mvp.py index 9e6f387..ce791ca 100644 --- a/backend/mvp.py +++ b/backend/mvp.py @@ -17,7 +17,7 @@ from fastapi import APIRouter, Depends, File, HTTPException, Query, Response, Up from fastapi.responses import FileResponse from icalendar import Calendar from pydantic import BaseModel, Field, model_validator -from sqlalchemy import case, delete, func, select +from sqlalchemy import case, delete, func, select, update from sqlalchemy.ext.asyncio import AsyncSession from .auth import current_user @@ -36,6 +36,7 @@ from .models import ( Tag, Task, TaskList, + TaskTag, User, new_id, utcnow, @@ -718,29 +719,113 @@ async def export_json(user: User = Depends(current_user), db: AsyncSession = Dep def serialize(row, fields): return {f: (str(v) if isinstance((v := getattr(row, f)), UUID) else v.isoformat() if isinstance(v, (date, datetime)) else v) for f in fields} folders = list((await db.scalars(select(Folder).where(Folder.user_id == user.id))).all()); lists = list((await db.scalars(select(TaskList).where(TaskList.user_id == user.id))).all()); tags = list((await db.scalars(select(Tag).where(Tag.user_id == user.id))).all()); tasks = list((await db.scalars(select(Task).where(Task.user_id == user.id))).all()); habits = list((await db.scalars(select(Habit).where(Habit.user_id == user.id))).all()) - return {"version": 1, "exported_at": utcnow(), "folders": [serialize(x,["id","name","position","deleted_at"]) for x in folders], "lists": [serialize(x,["id","folder_id","name","is_inbox","position","deleted_at"]) for x in lists], "tags": [serialize(x,["id","name","color"]) for x in tags], "tasks": [serialize(x,["id","list_id","parent_id","title","description","priority","completed","due_at","external_id","deleted_at"]) for x in tasks], "habits": [serialize(x,["id","name","kind","target","max_value","schedule_type","weekdays","month_days","interval_days","start_date","archived_at"]) for x in habits]} + return { + "version": 1, + "exported_at": utcnow(), + "folders": [serialize(x, ["id", "name", "position", "deleted_at"]) for x in folders], + "lists": [serialize(x, ["id", "folder_id", "name", "is_inbox", "position", "deleted_at"]) for x in lists], + "tags": [serialize(x, ["id", "name", "color"]) for x in tags], + "tasks": [serialize(x, ["id", "list_id", "parent_id", "title", "description", "priority", "completed", "due_at", "external_id", "deleted_at"]) for x in tasks], + "task_tags": [{"task_id": str(x.task_id), "tag_id": str(x.tag_id)} for x in (await db.scalars(select(TaskTag))).all() if x.task_id in {task.id for task in tasks}], + "habits": [serialize(x, ["id", "name", "kind", "target", "max_value", "schedule_type", "weekdays", "month_days", "interval_days", "start_date", "archived_at"]) for x in habits], + } @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 payload.get("version") != 1: raise HTTPException(422, "不支持的备份版本") + if payload.get("version") != 1: + raise HTTPException(422, "不支持的备份版本") if mode == "replace": - await db.execute(delete(Task).where(Task.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)) + await db.execute(delete(TaskTag).where(TaskTag.task_id.in_(select(Task.id).where(Task.user_id == user.id)))) + await db.execute(delete(Task).where(Task.user_id == user.id)) + await db.execute(delete(Tag).where(Tag.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 = {} + tag_id_map = {} for raw in payload.get("folders", []): - old = raw["id"]; row = Folder(user_id=user.id, name=raw["name"], position=raw.get("position",0)); db.add(row); await db.flush(); id_map[old] = row.id + old = raw["id"] + row = Folder(user_id=user.id, name=raw["name"], position=raw.get("position", 0)) + db.add(row) + await db.flush() + id_map[old] = row.id inbox = None for raw in payload.get("lists", []): - row = TaskList(user_id=user.id, folder_id=id_map.get(raw.get("folder_id")), name=raw["name"], is_inbox=raw.get("is_inbox",False), position=raw.get("position",0)); db.add(row); await db.flush(); id_map[raw["id"]] = row.id - if row.is_inbox: inbox = row - if not inbox: inbox = TaskList(user_id=user.id, name="收集箱", is_inbox=True); db.add(inbox); await db.flush() + row = TaskList( + user_id=user.id, + folder_id=id_map.get(raw.get("folder_id")), + name=raw["name"], + is_inbox=raw.get("is_inbox", False), + position=raw.get("position", 0), + ) + db.add(row) + await db.flush() + id_map[raw["id"]] = row.id + if row.is_inbox: + inbox = row + if not inbox: + inbox = TaskList(user_id=user.id, name="收集箱", is_inbox=True) + db.add(inbox) + await db.flush() + for raw in payload.get("tags", []): + row = Tag(user_id=user.id, name=raw["name"], color=raw.get("color", "#f15a29")) + db.add(row) + await db.flush() + tag_id_map[raw["id"]] = row.id restored = 0 + pending_tasks = [] for raw in payload.get("tasks", []): ext = raw.get("external_id") existing = await db.scalar(select(Task).where(Task.user_id == user.id, Task.external_id == ext)) if ext else None - if existing and mode == "merge": continue - row = Task(user_id=user.id, list_id=id_map.get(raw.get("list_id"), inbox.id), title=raw["title"], description=raw.get("description", ""), priority=raw.get("priority",0), completed=raw.get("completed",False), due_at=datetime.fromisoformat(raw["due_at"]) if raw.get("due_at") else None, external_id=ext); db.add(row); restored += 1 - audit(db, user.id, "restore", "backup", count=restored, mode=mode); await db.commit(); return {"restored": restored, "mode": mode} + if existing and mode == "merge": + task_id_map[raw["id"]] = existing.id + continue + row = Task( + user_id=user.id, + list_id=id_map.get(raw.get("list_id"), inbox.id), + title=raw["title"], + description=raw.get("description", ""), + priority=raw.get("priority", 0), + completed=raw.get("completed", False), + due_at=datetime.fromisoformat(raw["due_at"]) if raw.get("due_at") else None, + external_id=ext, + ) + db.add(row) + await db.flush() + task_id_map[raw["id"]] = row.id + pending_tasks.append((row.id, raw.get("parent_id"))) + restored += 1 + for task_id, old_parent_id in pending_tasks: + if old_parent_id and old_parent_id in task_id_map: + await db.execute(update(Task).where(Task.id == task_id, Task.user_id == user.id).values(parent_id=task_id_map[old_parent_id])) + task_tag_rows = [] + for raw in payload.get("task_tags", []): + task_id = task_id_map.get(raw.get("task_id")) + tag_id = tag_id_map.get(raw.get("tag_id")) + if task_id and tag_id: + task_tag_rows.append(TaskTag(task_id=task_id, tag_id=tag_id)) + if task_tag_rows: + db.add_all(task_tag_rows) + for raw in payload.get("habits", []): + row = Habit( + user_id=user.id, + name=raw["name"], + kind=raw.get("kind", "boolean"), + target=raw.get("target", 1), + max_value=raw.get("max_value"), + schedule_type=raw.get("schedule_type", "daily"), + weekdays=raw.get("weekdays"), + month_days=raw.get("month_days"), + interval_days=raw.get("interval_days"), + start_date=date.fromisoformat(raw["start_date"]), + archived_at=datetime.fromisoformat(raw["archived_at"]) if raw.get("archived_at") else None, + ) + db.add(row) + audit(db, user.id, "restore", "backup", count=restored, mode=mode) + await db.commit() + return {"restored": restored, "mode": mode} @router.get("/audit-logs") diff --git a/backend/schemas.py b/backend/schemas.py index 6adecc6..463e010 100644 --- a/backend/schemas.py +++ b/backend/schemas.py @@ -20,6 +20,17 @@ class UserOut(BaseModel): username: str +class SessionOut(BaseModel): + model_config = ConfigDict(from_attributes=True) + id: UUID + current: bool + created_at: datetime + last_seen_at: datetime + expires_at: datetime + ip_address: str | None + user_agent: str | None + + class NameUpdate(BaseModel): name: str = Field(min_length=1, max_length=120) @@ -106,7 +117,7 @@ class TaskDetailOut(TaskOut): class TaskPage(BaseModel): - items: list[TaskOut] + items: list[TaskDetailOut] next_cursor: str | None = None total: int = 0 page: int = 1 diff --git a/migrations/versions/0003_mvp_models.py b/migrations/versions/0003_mvp_models.py index 0696678..6d73d6e 100644 --- a/migrations/versions/0003_mvp_models.py +++ b/migrations/versions/0003_mvp_models.py @@ -13,8 +13,14 @@ depends_on = None def upgrade() -> None: + bind = op.get_bind() + dialect = bind.dialect.name op.add_column("tasks", sa.Column("external_id", sa.String(255), nullable=True)) - op.create_unique_constraint("uq_tasks_external_id", "tasks", ["user_id", "external_id"]) + if dialect == "sqlite": + with op.batch_alter_table("tasks") as batch_op: + batch_op.create_unique_constraint("uq_tasks_external_id", ["user_id", "external_id"]) + else: + op.create_unique_constraint("uq_tasks_external_id", "tasks", ["user_id", "external_id"]) op.add_column("sessions", sa.Column("ip_address", sa.String(64), nullable=True)) op.add_column("sessions", sa.Column("user_agent", sa.String(255), nullable=True)) diff --git a/migrations/versions/0006_remove_duplicate_habit_index.py b/migrations/versions/0006_remove_duplicate_habit_index.py index 3cfa3f6..f175560 100644 --- a/migrations/versions/0006_remove_duplicate_habit_index.py +++ b/migrations/versions/0006_remove_duplicate_habit_index.py @@ -3,6 +3,7 @@ Revision ID: 0006 Revises: 0005 """ +import sqlalchemy as sa from alembic import op revision = "0006" @@ -12,12 +13,18 @@ depends_on = None def upgrade() -> None: - op.drop_index("ix_habit_logs_habit_day", table_name="habit_logs") + bind = op.get_bind() + indexes = {index["name"] for index in sa.inspect(bind).get_indexes("habit_logs")} + if "ix_habit_logs_habit_day" in indexes: + op.drop_index("ix_habit_logs_habit_day", table_name="habit_logs") def downgrade() -> None: - op.create_index( - "ix_habit_logs_habit_day", - "habit_logs", - ["habit_id", "day"], - ) + bind = op.get_bind() + indexes = {index["name"] for index in sa.inspect(bind).get_indexes("habit_logs")} + if "ix_habit_logs_habit_day" not in indexes: + op.create_index( + "ix_habit_logs_habit_day", + "habit_logs", + ["habit_id", "day"], + ) diff --git a/tests/test_app.py b/tests/test_app.py index 941caf9..bc2727e 100644 --- a/tests/test_app.py +++ b/tests/test_app.py @@ -172,7 +172,10 @@ def test_task_search_tags_subtasks_and_recycle_bin(client): assert detail.status_code == 200 assert detail.json()["tags"][0]["name"] == "重要" assert len(detail.json()["subtasks"]) == 1 - assert len(client.get("/api/v1/tasks", params={"q": "咖啡"}).json()["items"]) == 1 + listed = client.get("/api/v1/tasks", params={"q": "咖啡"}).json()["items"] + assert len(listed) == 1 + assert listed[0]["tags"][0]["name"] == "重要" + assert listed[0]["subtasks"][0]["title"] == "比较价格" assert client.delete(f"/api/v1/tasks/{parent['id']}").status_code == 204 trash = client.get("/api/v1/trash").json()["items"] @@ -197,6 +200,29 @@ def test_batch_complete_and_move(client): assert all(row["completed"] and row["list_id"] == other["id"] for row in rows) +def test_batch_move_rejects_standalone_subtask_and_delete_cascades(client): + client = initialized_client(client) + inbox = client.get("/api/v1/lists").json()[0] + other = client.post("/api/v1/lists", json={"name": "稍后"}).json() + parent = client.post("/api/v1/tasks", json={"title": "父", "list_id": inbox["id"]}).json() + child = client.post( + "/api/v1/tasks", + json={"title": "子", "list_id": inbox["id"], "parent_id": parent["id"]}, + ).json() + + moved_child = client.post( + "/api/v1/tasks/batch", json={"task_ids": [child["id"]], "list_id": other["id"]} + ) + assert moved_child.status_code == 400 + + deleted_parent = client.post( + "/api/v1/tasks/batch", json={"task_ids": [parent["id"]], "soft_delete": True} + ) + assert deleted_parent.status_code == 200 + assert client.get(f"/api/v1/tasks/{parent['id']}").status_code == 404 + assert client.get(f"/api/v1/tasks/{child['id']}").status_code == 404 + + def test_inbox_is_protected_and_deleted_collections_are_hidden(client): client = initialized_client(client) inbox = client.get("/api/v1/lists").json()[0] @@ -258,6 +284,37 @@ def test_tags_are_global_searchable_and_validated(client): assert invalid.status_code == 404 +def test_restore_replace_recovers_tags_habits_and_task_links(client): + client = initialized_client(client) + inbox = client.get("/api/v1/lists").json()[0] + tag = client.post("/api/v1/tags", json={"name": "备份标签", "color": "#f15a29"}).json() + client.post( + "/api/v1/tasks", + json={"title": "备份任务", "list_id": inbox["id"], "tag_ids": [tag["id"]]}, + ).json() + client.post( + "/api/v1/habits", + json={"name": "俯卧撑", "kind": "boolean", "schedule_type": "daily"}, + ).json() + exported = client.get("/api/v1/export") + assert exported.status_code == 200 + + client.post("/api/v1/tags", json={"name": "现有标签", "color": "#123abc"}) + client.post( + "/api/v1/habits", + json={"name": "深蹲", "kind": "boolean", "schedule_type": "daily"}, + ) + restored = client.post("/api/v1/restore?mode=replace", json=exported.json()) + assert restored.status_code == 200 + + tags = client.get("/api/v1/tags").json() + assert [row["name"] for row in tags] == ["备份标签"] + habits = client.get("/api/v1/habits").json() + assert [row["name"] for row in habits] == ["俯卧撑"] + listed = client.get("/api/v1/tasks", params={"q": "备份任务"}).json()["items"] + assert listed[0]["tags"][0]["name"] == "备份标签" + + def test_recycle_bin_restore_and_permanent_delete_include_subtasks(client): client = initialized_client(client) inbox = client.get("/api/v1/lists").json()[0]