"""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