test: improve recurrence and floating button coverage
This commit is contained in:
@@ -0,0 +1,530 @@
|
||||
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
|
||||
Reference in New Issue
Block a user