130 lines
4.7 KiB
Python
130 lines
4.7 KiB
Python
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()
|
|
if changes.get("completed") is True and not task.completed:
|
|
changes["completed_at"] = now
|
|
elif changes.get("completed") is False:
|
|
changes["completed_at"] = None
|
|
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, completed_at=None, version=Task.version + 1, updated_at=now)
|
|
)
|
|
return updated_task, changed
|