feat: repeat tasks after completion
ci / gitleaks (push) Successful in 7s
ci / docker (push) Successful in 3m31s

This commit is contained in:
2026-09-10 21:21:28 +08:00
parent 84592db098
commit 64b8525720
14 changed files with 849 additions and 114 deletions
+65 -63
View File
@@ -7,7 +7,7 @@ import shutil
import time
from collections import defaultdict, deque
from contextlib import asynccontextmanager
from datetime import UTC, datetime, timedelta
from datetime import UTC, datetime
from pathlib import Path, PureWindowsPath
from uuid import UUID, uuid4
@@ -43,8 +43,9 @@ from .models import (
User,
utcnow,
)
from .mvp import audit, occurrences
from .mvp import audit
from .mvp import router as mvp_router
from .recurrence_service import apply_task_changes, lock_task
from .schemas import (
BatchResult,
BatchTaskUpdate,
@@ -66,6 +67,7 @@ from .schemas import (
TaskReorder,
TaskUpdate,
UserOut,
UserUpdate,
)
@@ -193,6 +195,18 @@ async def me(user: User = Depends(current_user)):
return user
@app.patch("/api/v1/me", response_model=UserOut)
async def update_me(
payload: UserUpdate,
user: User = Depends(current_user),
db: AsyncSession = Depends(get_db),
):
user.timezone = payload.timezone
await db.commit()
await db.refresh(user)
return user
@app.post("/api/v1/auth/change-password", status_code=204)
async def change_password(
payload: ChangePasswordRequest,
@@ -958,12 +972,12 @@ async def create_task(
)
if parent is None:
raise HTTPException(status_code=400, detail="父任务必须是同一清单的顶层任务")
if payload.rrule:
if not payload.due_at:
raise HTTPException(status_code=422, detail="重复任务需要截止时间")
recurrence_requested = payload.rrule or payload.trigger_mode
if recurrence_requested:
from .mvp import parse_rrule
parse_rrule(payload.rrule)
data = payload.model_dump(exclude={"rrule"})
if payload.rrule:
parse_rrule(payload.rrule)
data = payload.model_dump(exclude={"rrule", "trigger_mode", "after_completion_days"})
parent_filter = Task.parent_id == payload.parent_id if payload.parent_id else Task.parent_id.is_(None)
max_position = await db.scalar(select(func.max(Task.position)).where(
Task.user_id == user.id,
@@ -974,8 +988,17 @@ async def create_task(
task = Task(user_id=user.id, position=(max_position if max_position is not None else -1) + 1, **data)
db.add(task)
await db.flush()
if payload.rrule:
db.add(RecurrenceTemplate(user_id=user.id, task_id=task.id, rrule=payload.rrule.upper(), starts_at=task.due_at))
if recurrence_requested:
db.add(
RecurrenceTemplate(
user_id=user.id,
task_id=task.id,
rrule=payload.rrule.upper() if payload.rrule else None,
starts_at=task.due_at,
trigger_mode=payload.trigger_mode or "scheduled",
after_completion_days=payload.after_completion_days,
)
)
audit(db, user.id, "create", "task", task.id)
await db.commit()
await db.refresh(task)
@@ -1142,30 +1165,11 @@ async def update_task(
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)))
task = await lock_task(db, user.id, task_id)
if task is None:
raise HTTPException(status_code=404, detail="任务不存在")
data = payload.model_dump(exclude_unset=True)
expected_version = data.pop("version")
recurrence = None
reset_subtasks = False
if data.get("completed") is True:
recurrence = await db.scalar(select(RecurrenceTemplate).where(
RecurrenceTemplate.task_id == task_id, RecurrenceTemplate.user_id == user.id
))
if recurrence:
next_items = occurrences(
recurrence.rrule,
recurrence.starts_at,
recurrence.starts_at + timedelta(microseconds=1),
recurrence.starts_at + timedelta(days=3660),
recurrence.ends_at,
)
if next_items:
data["completed"] = False
data["due_at"] = next_items[0]
recurrence.starts_at = next_items[0]
reset_subtasks = True
if "list_id" in data:
await _owned_list(db, user.id, data["list_id"])
if task.parent_id:
@@ -1174,46 +1178,19 @@ async def update_task(
)
if parent_list != data["list_id"]:
raise HTTPException(status_code=400, detail="子任务必须与父任务属于同一清单")
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, changed = await apply_task_changes(
db,
user=user,
task_id=task_id,
expected_version=expected_version,
changes=data,
)
task = result.scalar_one_or_none()
if task is None:
exists_id = await db.scalar(
select(Task.id).where(Task.id == task_id, Task.user_id == user.id, Task.deleted_at.is_(None))
)
if exists_id:
raise HTTPException(status_code=409, detail="任务已被更新,请刷新后重试")
raise HTTPException(status_code=404, detail="任务不存在")
if reset_subtasks:
reset_at = utcnow()
await db.execute(
update(Task)
.where(
Task.parent_id == task.id,
Task.user_id == user.id,
Task.deleted_at.is_(None),
Task.completed.is_(True),
)
.values(completed=False, version=Task.version + 1, updated_at=reset_at)
)
if "list_id" in data and task.parent_id is None:
await db.execute(
update(Task)
.where(Task.parent_id == task.id, Task.user_id == user.id, Task.deleted_at.is_(None))
.values(list_id=task.list_id, version=Task.version + 1, updated_at=utcnow())
)
changed = {k for k in data if k not in {"version", "updated_at"}}
if changed & {"title", "description", "priority", "due_at", "list_id", "completed"}:
action = "complete" if data.get("completed") is True else "update"
audit(db, user.id, action, "task", task.id, fields=sorted(changed))
@@ -1337,7 +1314,32 @@ async def batch_update_tasks(
standalone_children = [task for task in tasks if task.parent_id is not None]
if standalone_children:
raise HTTPException(status_code=400, detail="子任务不能脱离父任务单独移动")
changes = payload.model_dump(exclude_unset=True, exclude={"task_ids", "soft_delete"})
changes = payload.model_dump(exclude_unset=True, exclude={"task_ids", "soft_delete", "versions"})
if payload.completed is True:
versions = payload.versions
other_changes = {key: value for key, value in changes.items() if key != "completed"}
for task_id in task_ids:
await apply_task_changes(
db,
user=user,
task_id=task_id,
expected_version=versions[task_id],
changes={"completed": True, **other_changes},
)
if payload.list_id is not None:
parent_ids = [task.id for task in tasks if task.parent_id is None]
if parent_ids:
await db.execute(
update(Task)
.where(
Task.parent_id.in_(parent_ids),
Task.user_id == user.id,
Task.deleted_at.is_(None),
)
.values(list_id=payload.list_id, version=Task.version + 1, updated_at=utcnow())
)
await db.commit()
return BatchResult(updated=len(task_ids))
if payload.soft_delete:
changes["deleted_at"] = utcnow()
if changes:
+4 -1
View File
@@ -148,9 +148,12 @@ class RecurrenceTemplate(Base):
id: Mapped[UUID] = mapped_column(primary_key=True, default=new_id)
user_id: Mapped[UUID] = mapped_column(ForeignKey("users.id", ondelete="CASCADE"), index=True)
task_id: Mapped[UUID] = mapped_column(ForeignKey("tasks.id", ondelete="CASCADE"), unique=True)
rrule: Mapped[str] = mapped_column(Text)
rrule: Mapped[str | None] = mapped_column(Text, nullable=True)
starts_at: Mapped[datetime] = mapped_column(UTCDateTime())
ends_at: Mapped[datetime | None] = mapped_column(UTCDateTime(), nullable=True)
trigger_mode: Mapped[str] = mapped_column(String(32), default="scheduled")
after_completion_days: Mapped[int | None] = mapped_column(Integer, nullable=True)
last_completed_at: Mapped[datetime | None] = mapped_column(UTCDateTime(), nullable=True)
created_at: Mapped[datetime] = mapped_column(UTCDateTime(), default=utcnow)
+105 -11
View File
@@ -50,13 +50,27 @@ def audit(db: AsyncSession, user_id: UUID, action: str, entity_type: str, entity
class RecurrenceCreate(BaseModel):
task_id: UUID
rrule: str = Field(min_length=5, max_length=1000)
rrule: str | None = Field(default=None, min_length=5, max_length=1000)
trigger_mode: str = Field(default="scheduled", pattern="^(scheduled|after_completion)$")
after_completion_days: int | None = Field(default=None, ge=1, le=3650)
@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 == "after_completion" and (
self.after_completion_days is None or self.rrule is not None
):
raise ValueError("after_completion requires days and no rrule")
return self
class RecurrenceChange(BaseModel):
title: str | None = Field(None, min_length=1, max_length=500)
due_at: datetime | None = None
rrule: str | None = None
trigger_mode: str | None = Field(default=None, pattern="^(scheduled|after_completion)$")
after_completion_days: int | None = Field(default=None, ge=1, le=3650)
class OccurrenceComplete(BaseModel):
@@ -208,21 +222,49 @@ async def get_task_recurrence(task_id: UUID, user: User = Depends(current_user),
))
if row is None:
return None
return {"id": row.id, "task_id": row.task_id, "rrule": row.rrule, "starts_at": row.starts_at, "ends_at": row.ends_at}
return {
"id": row.id,
"task_id": row.task_id,
"rrule": row.rrule,
"starts_at": row.starts_at,
"ends_at": row.ends_at,
"trigger_mode": row.trigger_mode,
"after_completion_days": row.after_completion_days,
"last_completed_at": row.last_completed_at,
}
@router.post("/recurrences", status_code=201)
async def create_recurrence(payload: RecurrenceCreate, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
task = await owned_task(db, user.id, payload.task_id)
if task.parent_id is not None:
raise HTTPException(422, "仅顶层任务可设置重复")
if not task.due_at:
raise HTTPException(422, "重复任务需要截止时间")
parse_rrule(payload.rrule)
if payload.rrule:
parse_rrule(payload.rrule)
if await db.scalar(select(RecurrenceTemplate.id).where(RecurrenceTemplate.task_id == task.id)):
raise HTTPException(409, "任务已有重复规则")
row = RecurrenceTemplate(user_id=user.id, task_id=task.id, rrule=payload.rrule.upper(), starts_at=task.due_at)
row = RecurrenceTemplate(
user_id=user.id,
task_id=task.id,
rrule=payload.rrule.upper() if payload.rrule else None,
starts_at=task.due_at,
trigger_mode=payload.trigger_mode,
after_completion_days=payload.after_completion_days,
)
db.add(row)
await db.commit(); await db.refresh(row)
return {"id": row.id, "task_id": row.task_id, "rrule": row.rrule, "starts_at": row.starts_at}
return {
"id": row.id,
"task_id": row.task_id,
"rrule": row.rrule,
"starts_at": row.starts_at,
"ends_at": row.ends_at,
"trigger_mode": row.trigger_mode,
"after_completion_days": row.after_completion_days,
"last_completed_at": row.last_completed_at,
}
async def upsert_exception(db, template_id, at):
@@ -238,6 +280,9 @@ async def upsert_exception(db, template_id, at):
async def edit_recurrence(recurrence_id: UUID, payload: RecurrenceChange, scope: str = Query("all", pattern="^(this|this-and-future|all)$"), occurrence_at: datetime | None = None, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
template = await owned_recurrence(db, user.id, recurrence_id)
task = await owned_task(db, user.id, template.task_id)
requested_mode = payload.trigger_mode or template.trigger_mode
if requested_mode == "after_completion" and scope != "all":
raise HTTPException(422, "完成后重复仅支持修改全部规则")
if scope == "this":
if not occurrence_at: raise HTTPException(422, "需要 occurrence_at")
ensure_real_occurrence(template, occurrence_at)
@@ -254,16 +299,49 @@ async def edit_recurrence(recurrence_id: UUID, payload: RecurrenceChange, scope:
elif payload.title:
task.title = payload.title
else:
if payload.rrule: parse_rrule(payload.rrule); template.rrule = payload.rrule.upper()
if payload.title is not None: task.title = payload.title
if payload.due_at is not None: task.due_at = payload.due_at; template.starts_at = payload.due_at
requested_days = (
payload.after_completion_days
if "after_completion_days" in payload.model_fields_set
else template.after_completion_days
)
requested_rrule = payload.rrule if payload.rrule is not None else template.rrule
if requested_mode == "after_completion":
if requested_days is None:
raise HTTPException(422, "完成后重复需要间隔天数")
template.trigger_mode = requested_mode
template.after_completion_days = requested_days
template.rrule = None
else:
if requested_rrule is None:
raise HTTPException(422, "定期重复需要 RRULE")
parse_rrule(requested_rrule)
template.trigger_mode = requested_mode
template.after_completion_days = None
template.rrule = requested_rrule.upper()
if payload.title is not None:
task.title = payload.title
if payload.due_at is not None:
task.due_at = payload.due_at
template.starts_at = payload.due_at
await db.commit()
return {"id": template.id, "scope": scope}
await db.refresh(template)
return {
"id": template.id,
"task_id": template.task_id,
"rrule": template.rrule,
"starts_at": template.starts_at,
"ends_at": template.ends_at,
"trigger_mode": template.trigger_mode,
"after_completion_days": template.after_completion_days,
"last_completed_at": template.last_completed_at,
}
@router.post("/recurrences/{recurrence_id}/complete")
async def complete_occurrence(recurrence_id: UUID, payload: OccurrenceComplete, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
template = await owned_recurrence(db, user.id, recurrence_id)
if template.trigger_mode != "scheduled":
raise HTTPException(422, "完成后重复不支持按发生时刻完成")
ensure_real_occurrence(template, payload.occurrence_at)
row = await upsert_exception(db, template.id, payload.occurrence_at); row.completed = True
audit(db, user.id, "complete", "task", template.task_id, occurrence_at=payload.occurrence_at.isoformat())
@@ -273,6 +351,8 @@ async def complete_occurrence(recurrence_id: UUID, payload: OccurrenceComplete,
@router.delete("/recurrences/{recurrence_id}", status_code=204)
async def delete_recurrence(recurrence_id: UUID, scope: str = Query("all", pattern="^(this|this-and-future|all)$"), occurrence_at: datetime | None = None, user: User = Depends(current_user), db: AsyncSession = Depends(get_db)):
template = await owned_recurrence(db, user.id, recurrence_id)
if template.trigger_mode == "after_completion" and scope != "all":
raise HTTPException(422, "完成后重复仅支持取消全部规则")
if scope == "this":
if not occurrence_at: raise HTTPException(422, "需要 occurrence_at")
ensure_real_occurrence(template, occurrence_at)
@@ -1031,7 +1111,7 @@ def _export_payload(folders, lists, tasks, recurrences, habits, countdowns):
"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],
"tasks": [serialize(x, ["id", "list_id", "parent_id", "title", "description", "priority", "completed", "due_at", "due_has_time", "external_id", "deleted_at"]) for x in tasks],
"recurrences": [serialize(x, ["id", "task_id", "rrule", "starts_at", "ends_at"]) for x in recurrences],
"recurrences": [serialize(x, ["id", "task_id", "rrule", "starts_at", "ends_at", "trigger_mode", "after_completion_days", "last_completed_at"]) for x in recurrences],
"habits": [serialize(x, ["id", "name", "kind", "target", "max_value", "schedule_type", "weekdays", "month_days", "interval_days", "start_date", "archived_at", "position"]) for x in habits],
"countdowns": [serialize(x, ["id", "title", "event_date", "calendar_mode", "lunar_month", "lunar_day", "ignore_year", "kind", "repeat_rule", "icon", "pinned", "archived_at", "created_at", "updated_at"]) for x in countdowns],
}
@@ -1228,12 +1308,26 @@ async def restore_json(payload: dict, mode: str = Query("merge", pattern="^(merg
task_id = task_id_map.get(raw.get("task_id"))
if not task_id:
continue
trigger_mode = raw.get("trigger_mode", "scheduled")
days = raw.get("after_completion_days")
rrule = raw.get("rrule")
if trigger_mode not in {"scheduled", "after_completion"}:
raise HTTPException(422, "无效的重复触发模式")
if trigger_mode == "after_completion":
if not isinstance(days, int) or isinstance(days, bool) or not 1 <= days <= 3650 or rrule is not None:
raise HTTPException(422, "无效的完成后重复备份")
elif not isinstance(rrule, str):
raise HTTPException(422, "定期重复缺少 RRULE")
db.add(RecurrenceTemplate(
user_id=user.id,
task_id=task_id,
rrule=raw["rrule"],
rrule=rrule,
starts_at=datetime.fromisoformat(raw["starts_at"]),
ends_at=datetime.fromisoformat(raw["ends_at"]) if raw.get("ends_at") else None,
trigger_mode=trigger_mode,
after_completion_days=days,
last_completed_at=datetime.fromisoformat(raw["last_completed_at"])
if raw.get("last_completed_at") else None,
))
for raw in payload.get("habits", []):
row = Habit(
+125
View File
@@ -0,0 +1,125 @@
from datetime import UTC, datetime, time, timedelta
from uuid import UUID
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
from fastapi import HTTPException
from sqlalchemy import select, update
from sqlalchemy.ext.asyncio import AsyncSession
from .models import RecurrenceTemplate, Task, User, utcnow
from .mvp import occurrences
async def lock_task(db: AsyncSession, user_id: UUID, task_id: UUID) -> Task | None:
return await db.scalar(
select(Task)
.where(Task.id == task_id, Task.user_id == user_id, Task.deleted_at.is_(None))
.with_for_update()
)
def _user_zone(user: User) -> ZoneInfo:
try:
return ZoneInfo(user.timezone)
except (ZoneInfoNotFoundError, ValueError) as exc:
raise HTTPException(422, "用户时区无效") from exc
def _after_completion_due(task: Task, completed_at: datetime, days: int, user: User) -> datetime:
zone = _user_zone(user)
completed_local = completed_at.astimezone(zone)
target_date = completed_local.date() + timedelta(days=days)
if task.due_has_time:
due_local = task.due_at.astimezone(zone)
wall_time = due_local.timetz().replace(tzinfo=None)
target_wall = datetime.combine(target_date, wall_time)
candidate = target_wall.replace(tzinfo=zone, fold=0)
# Normalize through UTC: DST gaps roll forward by their gap (02:30 -> 03:30),
# while ambiguous wall times deterministically keep the first occurrence (fold=0).
target_local = candidate.astimezone(UTC).astimezone(zone)
else:
target_local = datetime.combine(target_date, time(23, 59, 59), tzinfo=zone)
return target_local.astimezone(UTC)
async def apply_task_changes(
db: AsyncSession,
*,
user: User,
task_id: UUID,
expected_version: int,
changes: dict,
) -> tuple[Task, set[str]]:
task = await lock_task(db, user.id, task_id)
if task is None:
raise HTTPException(404, "任务不存在")
if task.version != expected_version:
raise HTTPException(409, "任务已被更新,请刷新后重试")
changed = set(changes)
recurrence = await db.scalar(
select(RecurrenceTemplate)
.where(RecurrenceTemplate.task_id == task.id, RecurrenceTemplate.user_id == user.id)
.with_for_update()
)
reset_subtasks = False
if changes.get("completed") is True and not task.completed and recurrence:
if recurrence.trigger_mode == "after_completion":
completed_at = utcnow()
next_due = _after_completion_due(
task, completed_at, recurrence.after_completion_days, user
)
changes["completed"] = False
changes["due_at"] = next_due
recurrence.starts_at = next_due
recurrence.last_completed_at = completed_at
reset_subtasks = True
else:
next_items = occurrences(
recurrence.rrule,
recurrence.starts_at,
recurrence.starts_at + timedelta(microseconds=1),
recurrence.starts_at + timedelta(days=3660),
recurrence.ends_at,
)
if next_items:
changes["completed"] = False
changes["due_at"] = next_items[0]
recurrence.starts_at = next_items[0]
reset_subtasks = True
if "due_at" in changes:
if changes["due_at"] is None and recurrence is not None:
await db.delete(recurrence)
recurrence = None
elif recurrence is not None and changes.get("completed") is not True:
recurrence.starts_at = changes["due_at"]
now = 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(**changes, version=Task.version + 1, updated_at=now)
.returning(Task)
)
updated_task = result.scalar_one_or_none()
if updated_task is None:
raise HTTPException(409, "任务已被更新,请刷新后重试")
if reset_subtasks:
await db.execute(
update(Task)
.where(
Task.parent_id == task.id,
Task.user_id == user.id,
Task.deleted_at.is_(None),
Task.completed.is_(True),
)
.values(completed=False, version=Task.version + 1, updated_at=now)
)
return updated_task, changed
+41
View File
@@ -1,5 +1,6 @@
from datetime import datetime
from uuid import UUID
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
@@ -29,6 +30,20 @@ class UserOut(BaseModel):
model_config = ConfigDict(from_attributes=True)
id: UUID
username: str
timezone: str
class UserUpdate(BaseModel):
timezone: str = Field(min_length=1, max_length=64)
@field_validator("timezone")
@classmethod
def validate_timezone(cls, value: str) -> str:
try:
ZoneInfo(value)
except (ZoneInfoNotFoundError, ValueError) as exc:
raise ValueError("timezone must be a valid IANA timezone") from exc
return value
class SessionOut(BaseModel):
@@ -108,6 +123,8 @@ class TaskCreate(BaseModel):
due_has_time: bool = False
parent_id: UUID | None = None
rrule: str | None = Field(default=None, min_length=5, max_length=1000)
trigger_mode: str | None = Field(default=None, pattern="^(scheduled|after_completion)$")
after_completion_days: int | None = Field(default=None, ge=1, le=3650)
@field_validator("title")
@classmethod
@@ -117,6 +134,22 @@ class TaskCreate(BaseModel):
raise ValueError("title cannot be blank")
return value
@model_validator(mode="after")
def validate_recurrence(self):
has_recurrence = self.rrule is not None or self.trigger_mode is not None or self.after_completion_days is not None
if has_recurrence and self.due_at is None:
raise ValueError("recurrence requires due_at")
if has_recurrence and self.parent_id is not None:
raise ValueError("only top-level tasks can recur")
if self.trigger_mode == "after_completion":
if self.after_completion_days is None or self.rrule is not None:
raise ValueError("after_completion requires days and no rrule")
elif self.trigger_mode == "scheduled" and self.rrule is None:
raise ValueError("scheduled recurrence requires rrule")
elif self.trigger_mode is None and self.after_completion_days is not None:
raise ValueError("after_completion_days requires after_completion mode")
return self
class TaskUpdate(BaseModel):
title: str | None = Field(default=None, min_length=1, max_length=500)
@@ -188,6 +221,7 @@ class BatchTaskUpdate(BaseModel):
list_id: UUID | None = None
due_at: datetime | None = None
soft_delete: bool | None = None
versions: dict[UUID, int] | None = None
@model_validator(mode="after")
def require_operation(self):
@@ -196,6 +230,13 @@ class BatchTaskUpdate(BaseModel):
raise ValueError("at least one batch operation is required")
if self.soft_delete is False:
raise ValueError("soft_delete can only be true")
if self.completed is True and self.versions is None:
raise ValueError("completion requires versions")
if self.versions is not None:
if set(self.versions) != set(self.task_ids):
raise ValueError("versions must cover every task")
if any(version < 1 for version in self.versions.values()):
raise ValueError("versions must be positive")
return self