from datetime import UTC, datetime from types import SimpleNamespace from unittest.mock import AsyncMock from uuid import uuid4 import pytest from fastapi import HTTPException from sqlalchemy import select, update from sqlalchemy.dialects import postgresql from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine from backend import recurrence_service from backend.db import Base from backend.models import RecurrenceTemplate, Task, TaskList, User NOW = datetime(2026, 9, 19, 8, 0, tzinfo=UTC) DUE = datetime(2026, 9, 20, 9, 30, tzinfo=UTC) NEXT_DUE = datetime(2026, 9, 21, 9, 30, tzinfo=UTC) def make_user(*, timezone="Asia/Shanghai"): return User(id=uuid4(), username="owner", password_hash="hash", timezone=timezone) def make_task(user, **overrides): values = { "id": uuid4(), "user_id": user.id, "list_id": uuid4(), "title": "task", "completed": False, "due_at": DUE, "due_has_time": True, "version": 3, } values.update(overrides) return Task(**values) def make_recurrence(user, task, **overrides): values = { "id": uuid4(), "user_id": user.id, "task_id": task.id, "rrule": "FREQ=DAILY", "starts_at": DUE, "ends_at": None, "trigger_mode": "scheduled", } values.update(overrides) return RecurrenceTemplate(**values) def result_returning(task): return SimpleNamespace(scalar_one_or_none=lambda: task) @pytest.fixture async def sqlite_db(tmp_path): engine = create_async_engine(f"sqlite+aiosqlite:///{tmp_path / 'recurrence.db'}") async with engine.begin() as connection: await connection.run_sync(Base.metadata.create_all) session_factory = async_sessionmaker(engine, expire_on_commit=False) async with session_factory() as db: yield db await engine.dispose() async def persist_user_and_list(db, username): user = User(username=username, password_hash="hash", timezone="Asia/Shanghai") db.add(user) await db.flush() task_list = TaskList(user_id=user.id, name=f"{username} inbox", is_inbox=True) db.add(task_list) await db.flush() return user, task_list @pytest.mark.asyncio async def test_lock_task_filters_owner_and_soft_deleted_tasks(sqlite_db): owner, owner_list = await persist_user_and_list(sqlite_db, "owner") foreign, foreign_list = await persist_user_and_list(sqlite_db, "foreign") active = Task(user_id=owner.id, list_id=owner_list.id, title="active") deleted = Task( user_id=owner.id, list_id=owner_list.id, title="deleted", deleted_at=NOW, ) foreign_task = Task(user_id=foreign.id, list_id=foreign_list.id, title="foreign") sqlite_db.add_all([active, deleted, foreign_task]) await sqlite_db.commit() assert await recurrence_service.lock_task(sqlite_db, owner.id, active.id) is active assert await recurrence_service.lock_task(sqlite_db, foreign.id, active.id) is None assert await recurrence_service.lock_task(sqlite_db, owner.id, foreign_task.id) is None assert await recurrence_service.lock_task(sqlite_db, owner.id, deleted.id) is None @pytest.mark.asyncio async def test_lock_task_actual_query_compiles_with_for_update(): class CapturingDb: statement = None async def scalar(self, statement): self.statement = statement db = CapturingDb() await recurrence_service.lock_task(db, uuid4(), uuid4()) statement = db.statement assert statement is not None sql = str(statement.compile(dialect=postgresql.dialect())) assert "tasks.user_id" in sql assert "tasks.deleted_at IS NULL" in sql assert sql.endswith(" FOR UPDATE") @pytest.mark.asyncio async def test_apply_task_changes_stale_version_cannot_update_database(sqlite_db): owner, owner_list = await persist_user_and_list(sqlite_db, "owner") task = Task( user_id=owner.id, list_id=owner_list.id, title="original", version=3, ) sqlite_db.add(task) await sqlite_db.commit() task_id = task.id with pytest.raises(HTTPException) as exc_info: await recurrence_service.apply_task_changes( sqlite_db, user=owner, task_id=task_id, expected_version=2, changes={"title": "stale write"}, ) await sqlite_db.rollback() sqlite_db.expire_all() persisted = await sqlite_db.get(Task, task_id) assert exc_info.value.status_code == 409 assert persisted.title == "original" assert persisted.version == 3 @pytest.mark.asyncio async def test_optimistic_update_version_predicate_rejects_race(sqlite_db, monkeypatch): owner, owner_list = await persist_user_and_list(sqlite_db, "owner") task = Task(user_id=owner.id, list_id=owner_list.id, title="original", version=3) sqlite_db.add(task) await sqlite_db.commit() task_id = task.id original_execute = sqlite_db.execute raced = False async def execute_after_concurrent_version_change(statement, *args, **kwargs): nonlocal raced if not raced and getattr(statement, "is_update", False): raced = True await original_execute( update(Task).where(Task.id == task_id).values(title="racer", version=4) ) await sqlite_db.flush() return await original_execute(statement, *args, **kwargs) monkeypatch.setattr(sqlite_db, "execute", execute_after_concurrent_version_change) with pytest.raises(HTTPException) as exc_info: await recurrence_service.apply_task_changes( sqlite_db, user=owner, task_id=task_id, expected_version=3, changes={"title": "loser"}, ) assert raced is True assert exc_info.value.status_code == 409 sqlite_db.expire_all() persisted = await sqlite_db.get(Task, task_id) assert persisted.title == "racer" assert persisted.version == 4 @pytest.mark.asyncio async def test_recurring_completion_resets_only_completed_active_owned_subtasks( sqlite_db, monkeypatch ): owner, owner_list = await persist_user_and_list(sqlite_db, "owner") foreign, foreign_list = await persist_user_and_list(sqlite_db, "foreign") parent = Task( user_id=owner.id, list_id=owner_list.id, title="recurring parent", due_at=DUE, due_has_time=True, version=3, ) other_parent = Task(user_id=owner.id, list_id=owner_list.id, title="other parent") sqlite_db.add_all([parent, other_parent]) await sqlite_db.flush() recurrence = make_recurrence(owner, parent) target = Task( user_id=owner.id, list_id=owner_list.id, parent_id=parent.id, title="target", completed=True, completed_at=NOW, version=2, ) incomplete = Task( user_id=owner.id, list_id=owner_list.id, parent_id=parent.id, title="already incomplete", completed=False, version=2, ) deleted = Task( user_id=owner.id, list_id=owner_list.id, parent_id=parent.id, title="deleted child", completed=True, completed_at=NOW, deleted_at=NOW, version=2, ) other_parent_child = Task( user_id=owner.id, list_id=owner_list.id, parent_id=other_parent.id, title="other parent child", completed=True, completed_at=NOW, version=2, ) foreign_child = Task( user_id=foreign.id, list_id=foreign_list.id, parent_id=parent.id, title="foreign child", completed=True, completed_at=NOW, version=2, ) sqlite_db.add_all( [recurrence, target, incomplete, deleted, other_parent_child, foreign_child] ) await sqlite_db.commit() parent_id = parent.id child_ids = [ target.id, incomplete.id, deleted.id, other_parent_child.id, foreign_child.id, ] monkeypatch.setattr(recurrence_service, "utcnow", lambda: NOW) monkeypatch.setattr(recurrence_service, "occurrences", lambda *args: [NEXT_DUE]) _, changed_fields = await recurrence_service.apply_task_changes( sqlite_db, user=owner, task_id=parent_id, expected_version=3, changes={"completed": True}, ) await sqlite_db.commit() sqlite_db.expire_all() rows = { row.title: row for row in ( await sqlite_db.scalars(select(Task).where(Task.id.in_(child_ids))) ).all() } persisted_parent = await sqlite_db.get(Task, parent_id) assert persisted_parent.completed is False assert persisted_parent.due_at == NEXT_DUE assert changed_fields == {"completed"} assert (rows["target"].completed, rows["target"].completed_at, rows["target"].version) == ( False, None, 3, ) for title in ("deleted child", "other parent child", "foreign child"): assert rows[title].completed is True assert rows[title].completed_at == NOW assert rows[title].version == 2 assert rows["already incomplete"].completed is False assert rows["already incomplete"].version == 2 @pytest.mark.parametrize("timezone", ["Not/A_Zone", ""]) def test_user_zone_rejects_invalid_timezone(timezone): with pytest.raises(HTTPException) as exc_info: recurrence_service._user_zone(make_user(timezone=timezone)) assert exc_info.value.status_code == 422 assert exc_info.value.detail == "用户时区无效" @pytest.mark.asyncio async def test_apply_task_changes_returns_not_found_when_lock_finds_no_task(): db = AsyncMock() db.scalar.return_value = None user = make_user() with pytest.raises(HTTPException) as exc_info: await recurrence_service.apply_task_changes( db, user=user, task_id=uuid4(), expected_version=1, changes={"title": "new"} ) assert exc_info.value.status_code == 404 assert db.scalar.await_count == 1 db.execute.assert_not_awaited() @pytest.mark.asyncio async def test_apply_task_changes_rejects_stale_version_before_loading_recurrence(): user = make_user() task = make_task(user) db = AsyncMock() db.scalar.return_value = task with pytest.raises(HTTPException) as exc_info: await recurrence_service.apply_task_changes( db, user=user, task_id=task.id, expected_version=2, changes={"title": "new"} ) assert exc_info.value.status_code == 409 assert db.scalar.await_count == 1 db.execute.assert_not_awaited() @pytest.mark.asyncio async def test_plain_completion_sets_timestamp_without_resetting_subtasks(monkeypatch): user = make_user() task = make_task(user) updated = make_task(user, id=task.id, completed=True, version=4, completed_at=NOW) db = AsyncMock() db.scalar.side_effect = [task, None] db.execute.return_value = result_returning(updated) monkeypatch.setattr(recurrence_service, "utcnow", lambda: NOW) changes = {"completed": True} returned, changed_fields = await recurrence_service.apply_task_changes( db, user=user, task_id=task.id, expected_version=3, changes=changes ) assert returned is updated assert changed_fields == {"completed"} assert changes["completed_at"] == NOW assert db.execute.await_count == 1 @pytest.mark.asyncio async def test_uncompleting_task_clears_completed_timestamp(monkeypatch): user = make_user() task = make_task(user, completed=True, completed_at=NOW) updated = make_task(user, id=task.id, completed=False, version=4, completed_at=None) db = AsyncMock() db.scalar.side_effect = [task, None] db.execute.return_value = result_returning(updated) monkeypatch.setattr(recurrence_service, "utcnow", lambda: NOW) changes = {"completed": False} await recurrence_service.apply_task_changes( db, user=user, task_id=task.id, expected_version=3, changes=changes ) assert changes["completed_at"] is None @pytest.mark.asyncio async def test_after_completion_reschedules_from_completion_and_records_it(monkeypatch): user = make_user() task = make_task(user) recurrence = make_recurrence( user, task, rrule=None, trigger_mode="after_completion", after_completion_days=2, ) updated = make_task(user, id=task.id, due_at=NEXT_DUE, version=4) db = AsyncMock() db.scalar.side_effect = [task, recurrence] db.execute.side_effect = [result_returning(updated), SimpleNamespace()] monkeypatch.setattr(recurrence_service, "utcnow", lambda: NOW) monkeypatch.setattr(recurrence_service, "_after_completion_due", lambda *args: NEXT_DUE) changes = {"completed": True} await recurrence_service.apply_task_changes( db, user=user, task_id=task.id, expected_version=3, changes=changes ) assert changes == {"completed": False, "due_at": NEXT_DUE, "completed_at": None} assert recurrence.starts_at == NEXT_DUE assert recurrence.last_completed_at == NOW assert db.execute.await_count == 2 @pytest.mark.asyncio async def test_scheduled_completion_advances_due_and_resets_completed_subtasks(monkeypatch): user = make_user() task = make_task(user) recurrence = make_recurrence(user, task) updated = make_task(user, id=task.id, due_at=NEXT_DUE, version=4) db = AsyncMock() db.scalar.side_effect = [task, recurrence] db.execute.side_effect = [result_returning(updated), SimpleNamespace()] monkeypatch.setattr(recurrence_service, "utcnow", lambda: NOW) monkeypatch.setattr(recurrence_service, "occurrences", lambda *args: [NEXT_DUE]) changes = {"completed": True} returned, changed_fields = await recurrence_service.apply_task_changes( db, user=user, task_id=task.id, expected_version=3, changes=changes ) assert returned is updated assert changed_fields == {"completed"} assert changes == {"completed": False, "due_at": NEXT_DUE, "completed_at": None} assert recurrence.starts_at == NEXT_DUE assert db.execute.await_count == 2 @pytest.mark.asyncio async def test_scheduled_completion_stays_completed_when_rule_has_no_future_occurrence(monkeypatch): user = make_user() task = make_task(user) recurrence = make_recurrence(user, task) updated = make_task(user, id=task.id, completed=True, version=4, completed_at=NOW) db = AsyncMock() db.scalar.side_effect = [task, recurrence] db.execute.return_value = result_returning(updated) monkeypatch.setattr(recurrence_service, "utcnow", lambda: NOW) monkeypatch.setattr(recurrence_service, "occurrences", lambda *args: []) changes = {"completed": True} await recurrence_service.apply_task_changes( db, user=user, task_id=task.id, expected_version=3, changes=changes ) assert changes == {"completed": True, "completed_at": NOW} assert recurrence.starts_at == DUE assert db.execute.await_count == 1 @pytest.mark.asyncio async def test_removing_due_date_deletes_recurrence(monkeypatch): user = make_user() task = make_task(user) recurrence = make_recurrence(user, task) updated = make_task(user, id=task.id, due_at=None, version=4) db = AsyncMock() db.scalar.side_effect = [task, recurrence] db.execute.return_value = result_returning(updated) monkeypatch.setattr(recurrence_service, "utcnow", lambda: NOW) await recurrence_service.apply_task_changes( db, user=user, task_id=task.id, expected_version=3, changes={"due_at": None} ) db.delete.assert_awaited_once_with(recurrence) @pytest.mark.asyncio async def test_moving_due_date_keeps_recurrence_anchor_in_sync(monkeypatch): user = make_user() task = make_task(user) recurrence = make_recurrence(user, task) updated = make_task(user, id=task.id, due_at=NEXT_DUE, version=4) db = AsyncMock() db.scalar.side_effect = [task, recurrence] db.execute.return_value = result_returning(updated) monkeypatch.setattr(recurrence_service, "utcnow", lambda: NOW) await recurrence_service.apply_task_changes( db, user=user, task_id=task.id, expected_version=3, changes={"due_at": NEXT_DUE} ) assert recurrence.starts_at == NEXT_DUE db.delete.assert_not_awaited() @pytest.mark.asyncio async def test_completion_request_with_explicit_due_does_not_reanchor_recurrence(monkeypatch): user = make_user() task = make_task(user, completed=True, completed_at=NOW) recurrence = make_recurrence(user, task) updated = make_task(user, id=task.id, completed=True, due_at=NEXT_DUE, version=4) db = AsyncMock() db.scalar.side_effect = [task, recurrence] db.execute.return_value = result_returning(updated) monkeypatch.setattr(recurrence_service, "utcnow", lambda: NOW) await recurrence_service.apply_task_changes( db, user=user, task_id=task.id, expected_version=3, changes={"completed": True, "due_at": NEXT_DUE}, ) assert recurrence.starts_at == DUE @pytest.mark.asyncio async def test_concurrent_update_loss_returns_conflict(monkeypatch): user = make_user() task = make_task(user) db = AsyncMock() db.scalar.side_effect = [task, None] db.execute.return_value = result_returning(None) monkeypatch.setattr(recurrence_service, "utcnow", lambda: NOW) with pytest.raises(HTTPException) as exc_info: await recurrence_service.apply_task_changes( db, user=user, task_id=task.id, expected_version=3, changes={"title": "new"} ) assert exc_info.value.status_code == 409