deerflow-code/offline-backend-20260512/backend/tests/test_deep_research_persistence.py
2026-09-07 18:24:55 +08:00

591 lines
23 KiB
Python

"""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()