"""Concurrency regressions from docs/WORKFLOW_STUDIO_BACKEND_REMEDIATION_ZH.md. Covers the P1-01/02/03 rework with the races the old two-step implementations lost: atomic human-input resume, draft optimistic locking, publish version-number allocation and event ``seq`` allocation — plus the sink's drop-not-publish rule and the shared JSON Schema validator. The SQL paths run against a real ``aiosqlite`` database (not mocks) because the conditional-UPDATE semantics only exist in SQL. ``connect_args={"timeout": 30}`` gives busy-waiting writers room so the interleavings exercise CAS rather than lock errors. """ from __future__ import annotations import asyncio import uuid from datetime import UTC, datetime, timedelta from typing import Any import pytest import pytest_asyncio from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine from deerflow.persistence.base import Base from deerflow.persistence.workflow_events import ( MemoryWorkflowEventStore, SqlWorkflowEventStore, WorkflowRunEventRow, ) from deerflow.persistence.workflow_runs import MemoryWorkflowRunStore, SqlWorkflowRunStore from deerflow.persistence.workflows import ( MemoryWorkflowStore, WorkflowDraftConflictError, WorkflowRepository, ) from deerflow.workflows.runtime.sink import PersistingWorkflowEventSink from deerflow.workflows.schema_validation import schema_issues def _graph(marker: str) -> dict[str, Any]: return { "schemaVersion": "1.0", "id": f"wf_{marker}", "name": marker, "nodes": [{"id": "start_1", "type": "start", "config": {}}, {"id": "end_1", "type": "output", "config": {"mapping": {}}}], "edges": [{"id": "e1", "source": "start_1", "target": "end_1"}], } async def _paused_run(store: Any, *, token: str = "tok-resume") -> str: """Drive a fresh run to ``awaiting_input`` the way the executor does.""" run, _created = await store.create_run( {"workflow_id": "wf-1", "workflow_version_id": "v-1", "owner_id": "user-a", "context": {"depth": 0}} ) run_id = run["id"] claimed = await store.claim_run(run_id, lease_owner="worker-1", lease_until=datetime.now(UTC) + timedelta(seconds=60)) assert claimed is not None and claimed["status"] == "running" paused = await store.set_awaiting_input( run_id, lease_owner="worker-1", pending_input={"nodeId": "ask", "prompt": "确认"}, resume_token=token, context={"loop_state": {"step": 2}, "resume_payload": None}, ) assert paused is not None and paused["status"] == "awaiting_input" return run_id # ── SQL fixture ───────────────────────────────────────────────────────── @pytest_asyncio.fixture async def session_factory(tmp_path): engine = create_async_engine( f"sqlite+aiosqlite:///{tmp_path / 'workflow-concurrency.db'}", connect_args={"timeout": 30}, ) async with engine.begin() as conn: import deerflow.persistence.workflow_events.model # noqa: F401 import deerflow.persistence.workflow_runs.model # noqa: F401 import deerflow.persistence.workflows.model # noqa: F401 await conn.run_sync(Base.metadata.create_all) yield async_sessionmaker(engine, expire_on_commit=False) await engine.dispose() # ── P1-01 atomic resume ───────────────────────────────────────────────── @pytest.mark.asyncio async def test_memory_resume_writes_payload_and_context_in_one_step() -> None: store = MemoryWorkflowRunStore() run_id = await _paused_run(store) resumed = await store.resume_with_payload( run_id, "tok-resume", {"nodeId": "ask", "action": "submit", "values": {"comment": "同意"}}, ) assert resumed is not None and resumed["status"] == "queued" # The checkpoint survived the resume; only the payload slot was written. assert resumed["context"]["loop_state"] == {"step": 2} assert resumed["context"]["resume_payload"]["values"] == {"comment": "同意"} assert resumed["resume_token"] is None and resumed["pending_input"] is None # One-shot token: a replay of the same token can never resume again. assert await store.resume_with_payload(run_id, "tok-resume", {"values": {"comment": "再次"}}) is None @pytest.mark.asyncio async def test_memory_concurrent_resume_single_winner() -> None: store = MemoryWorkflowRunStore() run_id = await _paused_run(store) results = await asyncio.gather( store.resume_with_payload(run_id, "tok-resume", {"values": {"comment": "A"}}), store.resume_with_payload(run_id, "tok-resume", {"values": {"comment": "B"}}), ) winners = [r for r in results if r is not None] assert len(winners) == 1, results final = await store.get_run(run_id) assert final is not None and final["status"] == "queued" assert final["context"]["resume_payload"]["values"]["comment"] in ("A", "B") assert final["context"]["resume_payload"] == winners[0]["context"]["resume_payload"] @pytest.mark.asyncio async def test_memory_resume_rejects_wrong_token() -> None: store = MemoryWorkflowRunStore() run_id = await _paused_run(store) assert await store.resume_with_payload(run_id, "not-the-token", {"values": {}}) is None still = await store.get_run(run_id) assert still is not None and still["status"] == "awaiting_input" and still["resume_token"] == "tok-resume" @pytest.mark.asyncio async def test_sql_concurrent_resume_single_winner(session_factory) -> None: store = SqlWorkflowRunStore(session_factory) run_id = await _paused_run(store) results = await asyncio.gather( store.resume_with_payload(run_id, "tok-resume", {"nodeId": "ask", "action": "submit", "values": {"comment": "A"}}), store.resume_with_payload(run_id, "tok-resume", {"nodeId": "ask", "action": "submit", "values": {"comment": "B"}}), ) winners = [r for r in results if r is not None] assert len(winners) == 1, results final = await store.get_run(run_id) assert final is not None and final["status"] == "queued" assert final["resume_token"] is None and final["pending_input"] is None assert final["context"]["loop_state"] == {"step": 2} assert final["context"]["resume_payload"] == winners[0]["context"]["resume_payload"] # ── P1-02 draft optimistic locking + publish counter ─────────────────── @pytest.mark.asyncio async def test_sql_stale_draft_save_is_rejected_without_overwrite(session_factory) -> None: store = WorkflowRepository(session_factory) created = await store.create_definition({"name": "draft", "owner_id": "user-a", "draft_graph": _graph("v1")}) wf_id = created["id"] saved = await store.save_draft(wf_id, expected_revision=0, graph=_graph("v2")) assert saved["draft_revision"] == 1 # The editor that still holds revision 0 must lose, not overwrite. with pytest.raises(WorkflowDraftConflictError) as exc: await store.save_draft(wf_id, expected_revision=0, graph=_graph("stale")) assert exc.value.current_revision == 1 fresh = await store.get_definition(wf_id, include_draft=True) assert fresh is not None and fresh["draft_graph"]["name"] == "v2" assert fresh["draft_revision"] == 1 @pytest.mark.asyncio async def test_sql_concurrent_draft_saves_yield_exactly_one_winner(session_factory) -> None: store = WorkflowRepository(session_factory) created = await store.create_definition({"name": "race", "owner_id": "user-a", "draft_graph": _graph("base")}) wf_id = created["id"] results = await asyncio.gather( store.save_draft(wf_id, expected_revision=0, graph=_graph("editor-a")), store.save_draft(wf_id, expected_revision=0, graph=_graph("editor-b")), return_exceptions=True, ) winners = [r for r in results if not isinstance(r, BaseException)] conflicts = [r for r in results if isinstance(r, WorkflowDraftConflictError)] assert len(winners) == 1 and len(conflicts) == 1, results fresh = await store.get_definition(wf_id, include_draft=True) assert fresh is not None and fresh["draft_revision"] == 1 assert fresh["draft_graph"]["name"] == winners[0]["draft_graph"]["name"] @pytest.mark.asyncio async def test_sql_concurrent_publishes_get_adjacent_unique_version_numbers(session_factory) -> None: from deerflow.persistence.workflows.model import WorkflowDefinitionRow store = WorkflowRepository(session_factory) created = await store.create_definition({"name": "publish-race", "owner_id": "user-a", "draft_graph": _graph("p")}) wf_id = created["id"] async def _publish(n: int) -> dict[str, Any]: return await store.publish_version( wf_id, graph=_graph(f"p{n}"), graph_hash=f"hash-{n}", published_by="user-a", change_note=f"v{n}", ) publishers = 5 versions = await asyncio.gather(*[_publish(n) for n in range(publishers)]) numbers = sorted(v["version_number"] for v in versions) assert numbers == [1, 2, 3, 4, 5] listed = await store.list_versions(wf_id) assert sorted(v["version_number"] for v in listed) == [1, 2, 3, 4, 5] # The counter must sit past every allocated number. async with session_factory() as session: row = await session.get(WorkflowDefinitionRow, wf_id) assert row is not None and row.next_version_number == 6 @pytest.mark.asyncio async def test_sql_publish_after_legacy_versions_never_reuses_a_number(session_factory) -> None: """A legacy DB backfilled the counter to 1 while versions already existed.""" from sqlalchemy import update from deerflow.persistence.workflows.model import WorkflowDefinitionRow store = WorkflowRepository(session_factory) created = await store.create_definition({"name": "legacy", "owner_id": "user-a", "draft_graph": _graph("l")}) wf_id = created["id"] for n in (1, 2, 3): await store.publish_version(wf_id, graph=_graph(f"l{n}"), graph_hash=f"h{n}", published_by="user-a") async with session_factory() as session: await session.execute( update(WorkflowDefinitionRow).where(WorkflowDefinitionRow.id == wf_id).values(next_version_number=1) ) await session.commit() version = await store.publish_version(wf_id, graph=_graph("l4"), graph_hash="h4", published_by="user-a") assert version["version_number"] == 4 @pytest.mark.asyncio async def test_memory_publish_numbers_survive_a_stale_counter() -> None: store = MemoryWorkflowStore() created = await store.create_definition({"name": "mem", "owner_id": "user-a", "draft_graph": _graph("m")}) wf_id = created["id"] for n in (1, 2): await store.publish_version(wf_id, graph=_graph(f"m{n}"), graph_hash=f"h{n}", published_by="user-a") # Simulate the legacy backfill: counter behind the existing versions. store._definitions[wf_id]["next_version_number"] = 1 version = await store.publish_version(wf_id, graph=_graph("m3"), graph_hash="h3", published_by="user-a") assert version["version_number"] == 3 # ── P1-03 event seq counter ───────────────────────────────────────────── @pytest.mark.asyncio async def test_sql_concurrent_appends_get_unique_monotonic_seqs(session_factory) -> None: runs = SqlWorkflowRunStore(session_factory) events = SqlWorkflowEventStore(session_factory) run, _created = await runs.create_run( {"workflow_id": "wf-1", "workflow_version_id": "v-1", "owner_id": "user-a"} ) appended = await asyncio.gather( *[ events.append(run_id=run["id"], workflow_id="wf-1", version_id="v-1", event_type="node.progress", payload={"i": i}) for i in range(10) ] ) seqs = sorted(int(row["seq"]) for row in appended) assert seqs == list(range(1, 11)) stored = await events.list_after(run["id"]) assert [e["seq"] for e in stored] == list(range(1, 11)) assert await events.max_seq(run["id"]) == 10 @pytest.mark.asyncio async def test_sql_seq_counter_self_heals_legacy_rows(session_factory) -> None: """Legacy DBs got the counter column backfilled to 0 with events present.""" runs = SqlWorkflowRunStore(session_factory) events = SqlWorkflowEventStore(session_factory) run, _created = await runs.create_run( {"workflow_id": "wf-1", "workflow_version_id": "v-1", "owner_id": "user-a"} ) async with session_factory() as session: for seq in (1, 2, 3): session.add( WorkflowRunEventRow( run_id=run["id"], workflow_id="wf-1", version_id="v-1", seq=seq, event_type="node.started", payload_json="{}", ) ) await session.commit() row = await events.append(run_id=run["id"], workflow_id="wf-1", version_id="v-1", event_type="node.completed") assert row["seq"] == 4 assert await events.max_seq(run["id"]) == 4 # The healed counter keeps allocating past the recovered maximum. nxt = await events.append(run_id=run["id"], workflow_id="wf-1", version_id="v-1", event_type="run.completed") assert nxt["seq"] == 5 @pytest.mark.asyncio async def test_sql_append_without_run_row_falls_back_to_max_seq(session_factory) -> None: """Event store used standalone (no run row) still allocates from MAX(seq)+1.""" events = SqlWorkflowEventStore(session_factory) orphan = f"run-{uuid.uuid4().hex[:8]}" first = await events.append(run_id=orphan, workflow_id="wf-1", version_id="v-1", event_type="run.created") second = await events.append(run_id=orphan, workflow_id="wf-1", version_id="v-1", event_type="run.started") assert (first["seq"], second["seq"]) == (1, 2) @pytest.mark.asyncio async def test_sql_event_replay_keeps_full_native_message_frame(session_factory) -> None: """A reconnect must receive the same large DeerFlow tool frame as live SSE.""" events = SqlWorkflowEventStore(session_factory) run_id = f"run-{uuid.uuid4().hex[:8]}" payload = { "messages": [ { "type": "tool", "name": "web_search", "content": {"results": [{"snippet": "x" * 90_000}]}, }, {"langgraph_node": "research"}, ] } await events.append( run_id=run_id, workflow_id="wf-1", version_id="v-1", event_type="node.message", payload=payload, ) replayed = await events.list_after(run_id) assert replayed[0]["payload"] == payload # ── sink: persist-then-publish, never ghost events ────────────────────── class _ExplodingEventStore(MemoryWorkflowEventStore): async def append(self, **_kwargs: Any) -> dict[str, Any]: raise RuntimeError("event log unavailable") @pytest.mark.asyncio async def test_sink_drops_unpersisted_events_instead_of_publishing_ghosts() -> None: published: list[dict[str, Any]] = [] async def _publish(frame: dict[str, Any]) -> None: published.append(frame) sink = PersistingWorkflowEventSink( _ExplodingEventStore(), run_id="run-1", workflow_id="wf-1", version_id="v-1", publish=_publish, ) envelope = await sink.emit("node.started", data={"chunk": "hi"}, node_id="n1") assert envelope.seq == 0 # no fabricated seq assert envelope.data["eventPersistenceDegraded"] is True assert envelope.data["chunk"] == "hi" assert published == [] # nothing went live: no ghost frame on the wire assert sink.degraded_count == 1 @pytest.mark.asyncio async def test_sink_publishes_only_after_successful_persist() -> None: published: list[dict[str, Any]] = [] async def _publish(frame: dict[str, Any]) -> None: published.append(frame) store = MemoryWorkflowEventStore() sink = PersistingWorkflowEventSink(store, run_id="run-1", workflow_id="wf-1", version_id="v-1", publish=_publish) envelope = await sink.emit("run.started", data={"attempt": 1}) assert envelope.seq == 1 assert len(published) == 1 assert published[0]["seq"] == 1 and published[0]["data"] == {"attempt": 1} assert "eventPersistenceDegraded" not in published[0]["data"] assert [e["seq"] for e in await store.list_after("run-1")] == [1] # ── P2 shared JSON Schema validator ───────────────────────────────────── def test_schema_issues_skips_empty_schemas() -> None: assert schema_issues({"anything": 1}, {}) == [] assert schema_issues({"anything": 1}, None) == [] assert schema_issues({"topic": "ok"}, {"type": "object", "properties": {}}) == [] def test_schema_issues_reports_nesting_paths_in_chinese() -> None: schema = { "type": "object", "required": ["topic", "options"], "properties": { "topic": {"type": "string"}, "mode": {"enum": ["fast", "slow"]}, "options": { "type": "array", "items": {"type": "object", "required": ["name"], "properties": {"name": {"type": "string"}}}, }, "meta": {"type": "object", "additionalProperties": False}, }, } issues = schema_issues( { "topic": 123, "mode": "turbo", "options": [{"name": "ok"}, {"title": "x"}], "meta": {"extra": "b"}, }, schema, ) by_path = {tuple(i["path"]): i for i in issues} assert "类型不正确" in by_path[("topic",)]["message"] assert "不在允许范围" in by_path[("mode",)]["message"] # Array element violations keep their index in the path. assert ("options", 1) in by_path and "缺少必填字段:name" in by_path[("options", 1)]["message"] # additionalProperties errors point at the offending object/field. meta_issues = [i for i in issues if i["path"] and i["path"][0] == "meta"] assert meta_issues and any("额外" in i["message"] for i in meta_issues) def test_schema_issues_reports_missing_top_level_keys() -> None: issues = schema_issues({}, {"type": "object", "required": ["topic"]}) assert issues == [{"path": [], "message": "输入 缺少必填字段:topic"}] def test_schema_issues_tolerates_malformed_schemas() -> None: # A non-string pattern only explodes during validation, not construction. issues = schema_issues("abc", {"pattern": 123}) assert issues == [{"path": [], "message": "schema 定义本身无效,无法用于校验"}] def test_schema_issues_caps_reported_errors() -> None: schema = {"type": "object", "properties": {f"f{i}": {"type": "string"} for i in range(30)}} value = {f"f{i}": i for i in range(30)} issues = schema_issues(value, schema) assert len(issues) == 21 # 20 concrete issues + the ellipsis marker assert issues[-1]["message"].startswith("……其余问题已省略")