195 lines
5.6 KiB
Python
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
|