"""Persistence tests for deep-research sessions, jobs, events, sources, messages. Uses in-memory SQLite + aiosqlite, mirroring ``test_ai_writing_session_repo``. Focuses on the durable-job mechanics that are load-bearing for correctness: - Session CRUD + user isolation. - Job ``try_create_or_get_active`` idempotency (unique dedupe key). - Job lease claim / renew / release (atomic CAS). - Two-phase cancel (``request_cancel`` → ``finalize_cancel``). - Event seq monotonicity + ``list_after`` replay (SSE reconnect correctness). Pattern: a plain ``async def _make_repos()`` helper (not a fixture) — matches the project convention for async DB setup in pytest-asyncio strict mode. """ from __future__ import annotations from datetime import UTC, datetime, timedelta import pytest from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine async def _make_repos(): """Create an in-memory SQLite repo set with all deep-research tables.""" from deerflow.persistence.base import Base from deerflow.persistence.deep_research_events.model import DeepResearchEventRow # noqa: F401 from deerflow.persistence.deep_research_events.sql import DeepResearchEventRepository from deerflow.persistence.deep_research_jobs.model import DeepResearchJobRow # noqa: F401 from deerflow.persistence.deep_research_jobs.sql import DeepResearchJobRepository from deerflow.persistence.deep_research_messages.model import DeepResearchMessageRow # noqa: F401 from deerflow.persistence.deep_research_messages.sql import DeepResearchMessageRepository from deerflow.persistence.deep_research_sessions.model import DeepResearchSessionRow # noqa: F401 from deerflow.persistence.deep_research_sessions.sql import DeepResearchSessionRepository from deerflow.persistence.deep_research_sources.model import DeepResearchSourceRow # noqa: F401 from deerflow.persistence.deep_research_sources.sql import DeepResearchSourceRepository engine = create_async_engine("sqlite+aiosqlite:///:memory:") async with engine.begin() as conn: await conn.run_sync(Base.metadata.create_all) sf = async_sessionmaker(engine, expire_on_commit=False) return { "sessions": DeepResearchSessionRepository(sf), "jobs": DeepResearchJobRepository(sf), "events": DeepResearchEventRepository(sf), "sources": DeepResearchSourceRepository(sf), "messages": DeepResearchMessageRepository(sf), "_engine": engine, } # ── sessions ──────────────────────────────────────────────────────────────── @pytest.mark.asyncio async def test_session_create_and_get(): repos = await _make_repos() try: created = await repos["sessions"].create(id="drs_1", user_id="u1", title="Test", query="AI research", mode="deep") assert created["id"] == "drs_1" assert created["status"] == "draft" got = await repos["sessions"].get("drs_1", user_id="u1") assert got is not None assert got["query"] == "AI research" finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_deep_research_writes_never_depend_on_post_commit_refresh(monkeypatch): """Control-plane writes must survive a MySQL replica-lag readback.""" from sqlalchemy.ext.asyncio import AsyncSession async def fail_refresh(self, instance, *args, **kwargs): # noqa: ANN001, ARG001 raise AssertionError("Deep Research writes must not call session.refresh") monkeypatch.setattr(AsyncSession, "refresh", fail_refresh) repos = await _make_repos() try: session = await repos["sessions"].create( id="drs_1", user_id="u1", title="Test", query="AI research" ) assert session["id"] == "drs_1" assert await repos["sessions"].update("drs_1", user_id="u1", status="running") job, created = await repos["jobs"].try_create_or_get_active( id="drj_1", session_id="drs_1", user_id="u1" ) assert created is True assert job["id"] == "drj_1" source = await repos["sources"].upsert( session_id="drs_1", id="src_1", user_id="u1", title="A", raw_content="source content", source="web", content_hash="hash_1", ) assert source["id"] == "src_1" message = await repos["messages"].append( id="msg_1", session_id="drs_1", user_id="u1", role="assistant", content="已开始撰写", ) assert message["id"] == "msg_1" finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_session_user_isolation(): repos = await _make_repos() try: await repos["sessions"].create(id="drs_1", user_id="u1", query="q1") assert await repos["sessions"].get("drs_1", user_id="u2") is None finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_session_list_strips_heavy_columns(): repos = await _make_repos() try: await repos["sessions"].create(id="drs_1", user_id="u1", query="q", config_snapshot={"mode": "deep"}) items = await repos["sessions"].list_by_user(user_id="u1") assert len(items) == 1 assert "report_markdown" not in items[0] assert "config_snapshot" not in items[0] finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_session_update_and_delete(): repos = await _make_repos() try: await repos["sessions"].create(id="drs_1", user_id="u1", query="q") updated = await repos["sessions"].update("drs_1", user_id="u1", title="New Title") assert updated["title"] == "New Title" assert await repos["sessions"].delete("drs_1", user_id="u1") is True assert await repos["sessions"].get("drs_1", user_id="u1") is None finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_count_active_by_user(): repos = await _make_repos() try: await repos["sessions"].create(id="drs_1", user_id="u1", query="q", status="running") await repos["sessions"].create(id="drs_2", user_id="u1", query="q", status="completed") count = await repos["sessions"].count_active_by_user(user_id="u1") assert count == 1 finally: await repos["_engine"].dispose() # ── jobs: idempotent create ───────────────────────────────────────────────── @pytest.mark.asyncio async def test_job_try_create_returns_created_true(): repos = await _make_repos() try: job, created = await repos["jobs"].try_create_or_get_active(id="drj_1", session_id="drs_1", user_id="u1", request_id="req_1") assert created is True assert job["status"] == "queued" assert job["active_dedupe_key"] is not None finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_job_create_never_reads_back_after_commit(monkeypatch): """A committed job must not fail when a MySQL replica briefly lags.""" from sqlalchemy.ext.asyncio import AsyncSession async def fail_refresh(self, instance, *args, **kwargs): # noqa: ANN001, ARG001 raise AssertionError("DeepResearchJobRepository must not call session.refresh") monkeypatch.setattr(AsyncSession, "refresh", fail_refresh) repos = await _make_repos() try: job, created = await repos["jobs"].try_create_or_get_active( id="drj_1", session_id="drs_1", user_id="u1", request_id="req_1", ) assert created is True assert job["id"] == "drj_1" assert job["status"] == "queued" finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_job_try_create_idempotent_returns_existing(): repos = await _make_repos() try: job1, created1 = await repos["jobs"].try_create_or_get_active(id="drj_1", session_id="drs_1", user_id="u1", request_id="req_1") job2, created2 = await repos["jobs"].try_create_or_get_active(id="drj_2", session_id="drs_1", user_id="u1", request_id="req_2") assert created1 is True assert created2 is False assert job2["id"] == job1["id"] finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_job_dedupe_released_on_terminal(): repos = await _make_repos() try: job, _ = await repos["jobs"].try_create_or_get_active(id="drj_1", session_id="drs_1", user_id="u1") assert job["active_dedupe_key"] is not None await repos["jobs"].update_progress("drj_1", lease_owner=None, user_id="u1", status="completed") job2, created2 = await repos["jobs"].try_create_or_get_active(id="drj_2", session_id="drs_1", user_id="u1") assert created2 is True assert job2["id"] == "drj_2" finally: await repos["_engine"].dispose() # ── jobs: lease mechanics ─────────────────────────────────────────────────── @pytest.mark.asyncio async def test_job_claim_succeeds_for_queued(): repos = await _make_repos() try: await repos["jobs"].try_create_or_get_active(id="drj_1", session_id="drs_1", user_id="u1") lease_until = datetime.now(UTC) + timedelta(seconds=60) claimed = await repos["jobs"].claim_job("drj_1", lease_owner="w1", lease_until=lease_until) assert claimed is not None assert claimed["status"] == "running" assert claimed["lease_owner"] == "w1" assert claimed["attempt"] == 1 finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_dispatcher_claims_queued_job_and_starts_report_worker(): """The claimed snapshot must reach the executor before the job can run.""" from app.gateway.deep_research_job_dispatcher import DeepResearchJobDispatcher class _Executor: def __init__(self) -> None: self.started: list[tuple[dict, str, str | None]] = [] def start_job(self, snapshot: dict, *, job_id: str, lease_owner: str | None) -> bool: self.started.append((snapshot, job_id, lease_owner)) return True repos = await _make_repos() try: await repos["jobs"].try_create_or_get_active( id="drj_1", session_id="drs_1", user_id="u1", input_snapshot={"query": "测试报告"}, ) runner = _Executor() dispatcher = DeepResearchJobDispatcher(repos["jobs"], runner, worker_id="worker_1") started = await dispatcher.dispatch_once() assert started == 1 assert runner.started == [({"query": "测试报告"}, "drj_1", "worker_1")] finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_job_claim_fails_if_already_claimed(): repos = await _make_repos() try: await repos["jobs"].try_create_or_get_active(id="drj_1", session_id="drs_1", user_id="u1") lease_until = datetime.now(UTC) + timedelta(seconds=60) first = await repos["jobs"].claim_job("drj_1", lease_owner="w1", lease_until=lease_until) assert first is not None second = await repos["jobs"].claim_job("drj_1", lease_owner="w2", lease_until=lease_until) assert second is None finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_job_renew_lease_only_by_owner(): repos = await _make_repos() try: await repos["jobs"].try_create_or_get_active(id="drj_1", session_id="drs_1", user_id="u1") lease_until = datetime.now(UTC) + timedelta(seconds=60) await repos["jobs"].claim_job("drj_1", lease_owner="w1", lease_until=lease_until) assert await repos["jobs"].renew_lease("drj_1", lease_owner="w1", lease_until=lease_until) is True assert await repos["jobs"].renew_lease("drj_1", lease_owner="w2", lease_until=lease_until) is False finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_job_expired_lease_reclaimable(): repos = await _make_repos() try: await repos["jobs"].try_create_or_get_active(id="drj_1", session_id="drs_1", user_id="u1") expired = datetime.now(UTC) - timedelta(seconds=10) await repos["jobs"].claim_job("drj_1", lease_owner="w1", lease_until=expired) now = datetime.now(UTC) claimable = await repos["jobs"].list_claimable(now=now) assert "drj_1" in claimable reclaimed = await repos["jobs"].claim_job("drj_1", lease_owner="w2", lease_until=now + timedelta(seconds=60)) assert reclaimed is not None assert reclaimed["lease_owner"] == "w2" finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_dispatcher_drops_expired_running_job_instead_of_reclaiming(): from app.gateway.deep_research_job_dispatcher import DeepResearchJobDispatcher class _Executor: def __init__(self) -> None: self.started: list[str] = [] def start_job(self, snapshot: dict, *, job_id: str, lease_owner: str | None) -> bool: # noqa: ARG002 self.started.append(job_id) return True repos = await _make_repos() try: await repos["jobs"].try_create_or_get_active( id="drj_1", session_id="drs_1", user_id="u1", input_snapshot={"query": "测试报告"}, ) expired = datetime.now(UTC) - timedelta(seconds=10) await repos["jobs"].claim_job("drj_1", lease_owner="w1", lease_until=expired) dispatcher = DeepResearchJobDispatcher(repos["jobs"], _Executor(), worker_id="worker_2") started = await dispatcher.dispatch_once() assert started == 0 row = await repos["jobs"].get_unscoped("drj_1") assert row["status"] == "cancelled" finally: await repos["_engine"].dispose() # ── jobs: two-phase cancel ────────────────────────────────────────────────── @pytest.mark.asyncio async def test_job_cancel_two_phase(): repos = await _make_repos() try: await repos["jobs"].try_create_or_get_active(id="drj_1", session_id="drs_1", user_id="u1") lease_until = datetime.now(UTC) + timedelta(seconds=60) await repos["jobs"].claim_job("drj_1", lease_owner="w1", lease_until=lease_until) requested = await repos["jobs"].request_cancel("drj_1", user_id="u1") assert requested is not None assert requested["cancel_requested"] is True finalized = await repos["jobs"].finalize_cancel("drj_1") assert finalized is not None assert finalized["status"] == "cancelled" assert finalized["active_dedupe_key"] is None finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_job_cancel_idempotent(): repos = await _make_repos() try: await repos["jobs"].try_create_or_get_active(id="drj_1", session_id="drs_1", user_id="u1") first = await repos["jobs"].request_cancel("drj_1", user_id="u1") assert first is not None await repos["jobs"].finalize_cancel("drj_1") second = await repos["jobs"].request_cancel("drj_1", user_id="u1") assert second is None finally: await repos["_engine"].dispose() # ── jobs: progress update with lease gate ─────────────────────────────────── @pytest.mark.asyncio async def test_job_update_progress_lease_gated(): repos = await _make_repos() try: await repos["jobs"].try_create_or_get_active(id="drj_1", session_id="drs_1", user_id="u1") lease_until = datetime.now(UTC) + timedelta(seconds=60) await repos["jobs"].claim_job("drj_1", lease_owner="w1", lease_until=lease_until) updated = await repos["jobs"].update_progress("drj_1", lease_owner="w1", phase="collecting", progress=50) assert updated is not None assert updated["phase"] == "collecting" assert updated["progress"] == 50 rejected = await repos["jobs"].update_progress("drj_1", lease_owner="w2", phase="writing") assert rejected is None finally: await repos["_engine"].dispose() # ── events: monotonic seq + replay ────────────────────────────────────────── @pytest.mark.asyncio async def test_event_seq_monotonic(): repos = await _make_repos() try: repo = repos["events"] e1 = await repo.append(session_id="drs_1", job_id="drj_1", event_type="phase_changed", payload={"to": "planning"}) e2 = await repo.append(session_id="drs_1", job_id="drj_1", event_type="source_added", payload={"id": "s1"}) e3 = await repo.append(session_id="drs_1", job_id="drj_1", event_type="source_added", payload={"id": "s2"}) assert e1["seq"] == 1 assert e2["seq"] == 2 assert e3["seq"] == 3 assert await repo.last_seq("drj_1") == 3 finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_event_append_never_reads_back_after_commit(monkeypatch): """Event persistence must work when a post-commit read hits a lagging replica.""" from sqlalchemy.ext.asyncio import AsyncSession async def fail_refresh(self, instance, *args, **kwargs): # noqa: ANN001, ARG001 raise AssertionError("DeepResearchEventRepository must not call session.refresh") monkeypatch.setattr(AsyncSession, "refresh", fail_refresh) repos = await _make_repos() try: event = await repos["events"].append( session_id="drs_1", job_id="drj_1", event_type="phase_changed", payload={"to": "planning"}, ) assert event["id"] is not None assert event["seq"] == 1 finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_event_list_after_replay(): repos = await _make_repos() try: repo = repos["events"] for i in range(5): await repo.append(session_id="drs_1", job_id="drj_1", event_type="heartbeat", payload={"i": i}) replayed = await repo.list_after("drj_1", after=2) assert len(replayed) == 3 assert [e["seq"] for e in replayed] == [3, 4, 5] finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_event_payload_round_trip(): repos = await _make_repos() try: repo = repos["events"] await repo.append( session_id="drs_1", job_id="drj_1", event_type="plan_created", payload={"subtopics": ["AI", "ML"], "nested": {"a": 1}}, ) events = await repo.list_after("drj_1", after=0) assert events[0]["payload"]["subtopics"] == ["AI", "ML"] assert events[0]["payload"]["nested"]["a"] == 1 finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_event_payload_serializes_native_source_datetime(): """Source-added events share the source model's native ``published_at``.""" repos = await _make_repos() try: repo = repos["events"] published_at = datetime(2026, 9, 2, 10, 9, tzinfo=UTC) await repo.append( session_id="drs_1", job_id="drj_1", event_type="source_added", payload={"source": {"id": "src_1", "published_at": published_at}}, ) events = await repo.list_after("drj_1", after=0) assert events[0]["payload"]["source"]["published_at"] == published_at.isoformat() finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_event_seq_independent_per_job(): repos = await _make_repos() try: repo = repos["events"] await repo.append(session_id="drs_1", job_id="drj_1", event_type="heartbeat") await repo.append(session_id="drs_1", job_id="drj_1", event_type="heartbeat") await repo.append(session_id="drs_1", job_id="drj_2", event_type="heartbeat") assert await repo.last_seq("drj_1") == 2 assert await repo.last_seq("drj_2") == 1 finally: await repos["_engine"].dispose() # ── sources ───────────────────────────────────────────────────────────────── @pytest.mark.asyncio async def test_source_upsert_and_list(): repos = await _make_repos() try: await repos["sources"].upsert( session_id="drs_1", id="src_1", user_id="u1", title="Source A", url="https://example.com/a", raw_content="content A", source="web", source_type="web", content_hash="hash_a", ) items = await repos["sources"].list_by_session("drs_1", user_id="u1") assert len(items) == 1 assert items[0]["title"] == "Source A" assert "raw_content" not in items[0] finally: await repos["_engine"].dispose() @pytest.mark.asyncio async def test_source_set_selected(): repos = await _make_repos() try: await repos["sources"].upsert( session_id="drs_1", id="src_1", user_id="u1", title="A", raw_content="c", source="web", content_hash="h", ) updated = await repos["sources"].set_selected("src_1", selected=True, reason="relevant", citation_key="1", user_id="u1") assert updated["selected"] is True assert updated["citation_key"] == "1" finally: await repos["_engine"].dispose() # ── messages ──────────────────────────────────────────────────────────────── @pytest.mark.asyncio async def test_message_append_and_list(): repos = await _make_repos() try: await repos["messages"].append(id="msg_1", session_id="drs_1", user_id="u1", role="user", content="question?") await repos["messages"].append( id="msg_2", session_id="drs_1", user_id="u1", role="assistant", content="answer", citation_source_ids=["src_1"], ) msgs = await repos["messages"].list_by_session("drs_1", user_id="u1") assert len(msgs) == 2 assert msgs[0]["role"] == "user" assert msgs[1]["citation_source_ids"] == ["src_1"] finally: await repos["_engine"].dispose()