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