feat: repeat tasks after completion
This commit is contained in:
+65
-63
@@ -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
@@ -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
@@ -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(
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user