591 lines
23 KiB
Python
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()
|