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

195 lines
5.6 KiB
Python

"""Phase-0 workflow protocol contract tests.
Locks graph schema, event envelope, error codes, and shipped sample fixtures.
"""
from __future__ import annotations
import json
from pathlib import Path
import pytest
from pydantic import ValidationError
from deerflow.workflows.errors import (
ALL_WORKFLOW_ERROR_CODES,
WorkflowError,
WorkflowErrorBody,
)
from deerflow.workflows.events import (
ALL_WORKFLOW_EVENT_TYPES,
TERMINAL_EVENT_TYPES,
WorkflowEventEnvelope,
)
from deerflow.workflows.samples import list_node_samples, load_node_sample, load_sample_graph
from deerflow.workflows.schemas import (
ALL_NODE_TYPES,
SCHEMA_VERSION,
AgentNodeConfig,
NodeResult,
SqlReadNodeConfig,
StartRunRequest,
WorkflowGraph,
WorkflowNode,
)
def test_schema_version_frozen() -> None:
assert SCHEMA_VERSION == "1.0"
def test_all_node_types_have_samples() -> None:
samples = set(list_node_samples())
assert samples == set(ALL_NODE_TYPES)
for node_type in ALL_NODE_TYPES:
raw = load_node_sample(node_type)
node = WorkflowNode.model_validate(raw)
assert node.type == node_type
def test_report_generation_sample_loads() -> None:
graph = load_sample_graph("report_generation")
assert graph.id == "wf_report_generation"
assert graph.schema_version == "1.0"
types = {n.type for n in graph.nodes}
assert "sql_read" in types
assert "agent" in types
assert "condition" in types
assert "loop" in types
assert "merge" in types
assert "output" in types
digest = graph.graph_hash()
assert len(digest) == 64
assert digest == graph.graph_hash() # deterministic
def test_start_output_sample_loads() -> None:
graph = load_sample_graph("start_output")
assert len(graph.nodes) == 2
assert len(graph.edges) == 1
def test_duplicate_node_id_rejected() -> None:
with pytest.raises(ValidationError):
WorkflowGraph.model_validate(
{
"schemaVersion": "1.0",
"id": "wf_bad",
"nodes": [
{"id": "a", "type": "start"},
{"id": "a", "type": "output"},
],
"edges": [],
}
)
def test_edge_target_must_exist() -> None:
with pytest.raises(ValidationError):
WorkflowGraph.model_validate(
{
"schemaVersion": "1.0",
"id": "wf_bad",
"nodes": [{"id": "start_1", "type": "start"}],
"edges": [{"id": "e1", "source": "start_1", "target": "missing"}],
}
)
def test_camel_case_roundtrip() -> None:
graph = load_sample_graph("start_output")
wire = json.loads(graph.model_dump_json(by_alias=True))
assert "schemaVersion" in wire
assert "inputSchema" in wire
assert "runTimeoutSeconds" in wire["settings"]
again = WorkflowGraph.model_validate(wire)
assert again.graph_hash() == graph.graph_hash()
def test_node_result_and_typed_configs() -> None:
result = NodeResult.model_validate(
{
"data": {"answer": "ok"},
"artifacts": [
{
"artifactId": "a1",
"name": "r.md",
"mimeType": "text/markdown",
"path": "/mnt/data/r.md",
}
],
"warnings": [],
}
)
assert result.artifacts[0].artifact_id == "a1"
agent = AgentNodeConfig.model_validate(
{
"agentId": "x",
"promptTemplate": "hi",
"responseMode": "json",
"threadMode": "isolated_per_node",
}
)
assert agent.agent_id == "x"
sql = SqlReadNodeConfig.model_validate(
{
"dataSourceId": "ds_1",
"statement": "SELECT 1",
"parameters": {},
"maxRows": 10,
}
)
assert sql.data_source_id == "ds_1"
def test_start_run_request() -> None:
req = StartRunRequest.model_validate(
{
"versionId": "wv_03",
"inputs": {"topic": "市场分析"},
"idempotencyKey": "client-uuid",
}
)
assert req.version_id == "wv_03"
assert req.execution_mode == "normal"
def test_event_envelope_fixture() -> None:
path = Path(__file__).resolve().parents[1] / "packages/harness/deerflow/workflows/samples/event_envelope_example.json"
raw = json.loads(path.read_text(encoding="utf-8"))
env = WorkflowEventEnvelope.model_validate(raw)
assert env.event_type == "node.output.delta"
assert env.seq == 42
sse = env.to_sse_dict()
assert sse["event"] == "node.output.delta"
assert sse["runId"] == "run_01"
frame = env.to_sse_frame()
assert frame.startswith("id: 42\n")
assert "event: node.output.delta\n" in frame
assert '"schemaVersion":"1.0"' in frame.replace(" ", "")
def test_event_type_set_frozen() -> None:
assert "run.completed" in ALL_WORKFLOW_EVENT_TYPES
assert "run.failed" in TERMINAL_EVENT_TYPES
assert "run.cancelled" in TERMINAL_EVENT_TYPES
assert len(ALL_WORKFLOW_EVENT_TYPES) == 18
def test_error_body_and_exception() -> None:
assert "WORKFLOW_SQL_POLICY_DENIED" in ALL_WORKFLOW_ERROR_CODES
err = WorkflowError(
"WORKFLOW_SQL_POLICY_DENIED",
"SQL 节点只允许单条只读查询",
node_id="sql_1",
details={"rule": "single_read_statement"},
)
body = err.to_body()
assert isinstance(body, WorkflowErrorBody)
wire = body.model_dump(by_alias=True)
assert wire["nodeId"] == "sql_1"
assert wire["code"] == "WORKFLOW_SQL_POLICY_DENIED"
assert wire["retryable"] is False