462 lines
19 KiB
Python
462 lines
19 KiB
Python
"""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("……其余问题已省略")
|