from datetime import UTC, datetime, time, timedelta from uuid import UUID from zoneinfo import ZoneInfo, ZoneInfoNotFoundError from dateutil.relativedelta import relativedelta 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, value: int, unit: str | None, user: User ) -> datetime: zone = _user_zone(user) completed_local = completed_at.astimezone(zone) if unit == "months": target_date = completed_local.date() + relativedelta(months=value) else: target_date = completed_local.date() + timedelta(days=value) 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, recurrence.after_completion_unit, 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