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

462 lines
19 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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("……其余问题已省略")