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

1124 lines
42 KiB
Python

"""Phase 2/3 tests: expressions, scheduler semantics, policies."""
from __future__ import annotations
import asyncio
from types import SimpleNamespace
from typing import Any
import pytest
from app.gateway.workflow_deep_research_adapter import WorkflowDeepResearchAdapter
from deerflow.config.workflow_config import WorkflowConfig
from deerflow.persistence.workflow_events import MemoryWorkflowEventStore
from deerflow.persistence.workflow_runs import MemoryWorkflowRunStore
from deerflow.workflows.errors import WorkflowError
from deerflow.workflows.expressions import collect_referenced_nodes, evaluate_condition, render_template, render_value
from deerflow.workflows.nodes import build_default_registry
from deerflow.workflows.runtime import (
CancelToken,
PersistingWorkflowEventSink,
RunContext,
WorkflowEngine,
WorkflowPaused,
WorkflowRuntimeDeps,
)
from deerflow.workflows.schemas import WorkflowGraph
from deerflow.workflows.security.http_policy import enforce_http_target
from deerflow.workflows.security.sql_policy import assert_read_only, assert_tables_allowed, enforce_limit
from deerflow.workflows.validator import validate_workflow_graph
def _graph(nodes: list[dict[str, Any]], edges: list[dict[str, Any]], **kwargs: Any) -> WorkflowGraph:
return WorkflowGraph.model_validate(
{
"schemaVersion": "1.0",
"id": "wf-test",
"name": "test",
"nodes": nodes,
"edges": edges,
**kwargs,
}
)
async def _run(
graph: WorkflowGraph,
inputs: dict[str, Any] | None = None,
*,
deps: WorkflowRuntimeDeps | None = None,
resume_payload: dict[str, Any] | None = None,
env: dict[str, Any] | None = None,
config: WorkflowConfig | None = None,
) -> tuple[Any, list[dict[str, Any]]]:
run_store = MemoryWorkflowRunStore()
event_store = MemoryWorkflowEventStore()
published: list[dict[str, Any]] = []
async def publish(event: dict[str, Any]) -> None:
published.append(event)
sink = PersistingWorkflowEventSink(event_store, run_id="run-1", workflow_id=graph.id, version_id="v1", publish=publish)
ctx = RunContext(
run_id="run-1",
workflow_id=graph.id,
version_id="v1",
owner_id="user-a",
graph=graph,
inputs=inputs or {},
deps=deps or WorkflowRuntimeDeps(),
cancel=CancelToken(),
emit=sink.emit,
resume_payload=resume_payload,
env=env or {},
)
engine = WorkflowEngine(build_default_registry(config or WorkflowConfig()), run_store=run_store)
result = await engine.run(ctx)
return result, published
# ── expressions ─────────────────────────────────────────────────────────
def test_expression_rendering_and_paths() -> None:
state = {
"inputs": {"topic": "储能"},
"nodes": {"n1": {"data": {"items": [{"title": "A"}, {"title": "B"}], "score": 0.9}}},
"run": {"id": "r1"},
"loop": {},
"env": {},
}
assert render_template("主题:{{ inputs.topic }}", state) == "主题:储能"
assert render_template("{{ nodes.n1.data.items[1].title }}", state) == "B"
# A lone placeholder keeps the native type.
assert render_value("{{ nodes.n1.data.score }}", state) == 0.9
assert render_value({"a": ["{{ inputs.topic }}"]}, state) == {"a": ["储能"]}
# Missing paths resolve to empty rather than exploding.
assert render_template("{{ nodes.nope.data.x }}", state) == ""
def test_expression_supports_numeric_canvas_node_ids() -> None:
state = {
"inputs": {},
"nodes": {"100002": {"data": {"text": "流式回复"}}},
"run": {},
"loop": {},
"env": {},
}
assert render_template("{{ nodes.100002.data.text }}", state) == "流式回复"
assert collect_referenced_nodes({"answer": "{{ nodes.100002.data.text }}"}) == {"100002"}
def test_expression_rejects_unknown_roots_and_code() -> None:
state = {"inputs": {}, "nodes": {}, "run": {}, "loop": {}, "env": {}}
with pytest.raises(WorkflowError) as exc:
render_template("{{ __import__('os').system('id') }}", state)
assert exc.value.code == "WORKFLOW_EXPRESSION_INVALID"
with pytest.raises(WorkflowError):
render_template("{{ os.environ.SECRET }}", state)
def test_condition_evaluation() -> None:
state = {
"inputs": {"count": 5, "tags": ["a", "b"], "name": "储能报告"},
"nodes": {"n1": {"data": {"ok": True}}},
"run": {},
"loop": {"n1": 2},
"env": {},
}
assert evaluate_condition("inputs.count > 3", state) is True
assert evaluate_condition("inputs.count > 3 and nodes.n1.data.ok", state) is True
assert evaluate_condition("inputs.count > 9 or nodes.n1.data.ok == true", state) is True
assert evaluate_condition("'a' in inputs.tags", state) is True
assert evaluate_condition("inputs.name contains '储能'", state) is True
assert evaluate_condition("not nodes.n1.data.ok", state) is False
assert evaluate_condition("loop.n1 < 5", state) is True
# ── scheduler ───────────────────────────────────────────────────────────
@pytest.mark.asyncio
async def test_start_output_skeleton_emits_ordered_events() -> None:
graph = _graph(
[
{"id": "start", "type": "start", "name": "开始"},
{"id": "end", "type": "output", "name": "输出", "config": {"mapping": {"echo": "{{ inputs.topic }}"}}},
],
[{"id": "e1", "source": "start", "target": "end"}],
)
result, events = await _run(graph, {"topic": "储能"})
assert result.output == {"echo": "储能"}
seqs = [e["seq"] for e in events]
assert seqs == sorted(seqs) and seqs[0] == 1
kinds = [e["event"] for e in events]
assert kinds == [
"node.started",
"node.completed",
"node.started",
"node.completed",
]
@pytest.mark.asyncio
async def test_condition_prunes_untaken_branch() -> None:
graph = _graph(
[
{"id": "start", "type": "start"},
{
"id": "gate",
"type": "condition",
"config": {
"branches": [{"name": "big", "expression": "inputs.n > 10"}],
"defaultBranch": "small",
},
},
{"id": "big", "type": "transform", "config": {"operations": [{"op": "set", "target": "path", "value": "big"}]}},
{"id": "small", "type": "transform", "config": {"operations": [{"op": "set", "target": "path", "value": "small"}]}},
{"id": "end", "type": "output", "config": {"mapping": {"path": "{{ nodes.small.data.path }}"}}},
],
[
{"id": "e1", "source": "start", "target": "gate"},
{"id": "e2", "source": "gate", "target": "big", "sourcePort": "big"},
{"id": "e3", "source": "gate", "target": "small", "sourcePort": "small"},
{"id": "e4", "source": "small", "target": "end"},
],
)
result, _ = await _run(graph, {"n": 1})
assert result.output == {"path": "small"}
assert "big" not in result.node_results
@pytest.mark.asyncio
async def test_parallel_fanout_joins_at_merge() -> None:
graph = _graph(
[
{"id": "start", "type": "start"},
{"id": "a", "type": "transform", "config": {"operations": [{"op": "set", "target": "v", "value": "A"}]}},
{"id": "b", "type": "transform", "config": {"operations": [{"op": "set", "target": "v", "value": "B"}]}},
{"id": "join", "type": "merge", "config": {"strategy": "all"}},
{"id": "end", "type": "output", "config": {"mapping": {"joined": "{{ nodes.join.data.succeeded }}"}}},
],
[
{"id": "e1", "source": "start", "target": "a"},
{"id": "e2", "source": "start", "target": "b"},
{"id": "e3", "source": "a", "target": "join"},
{"id": "e4", "source": "b", "target": "join"},
{"id": "e5", "source": "join", "target": "end"},
],
)
result, _ = await _run(graph)
assert sorted(result.output["joined"]) == ["a", "b"]
@pytest.mark.asyncio
async def test_loop_reruns_body_and_respects_cap() -> None:
graph = _graph(
[
{"id": "start", "type": "start"},
{"id": "body", "type": "transform", "config": {"operations": [{"op": "set", "target": "tick", "value": "1"}]}},
{
"id": "gate",
"type": "loop",
"config": {"maxIterations": 3, "continueWhen": "true", "bodyEntry": "body"},
},
{"id": "end", "type": "output", "config": {"mapping": {"iterations": "{{ nodes.gate.data.iteration }}"}}},
],
[
{"id": "e1", "source": "start", "target": "body"},
{"id": "e2", "source": "body", "target": "gate"},
{"id": "e3", "source": "gate", "target": "body", "sourcePort": "continue"},
{"id": "e4", "source": "gate", "target": "end", "sourcePort": "done"},
],
settings={"maxLoopIterations": 5},
)
result, _ = await _run(graph)
# maxIterations=3 means the gate stops handing out the continue port at 2.
assert result.output["iterations"] == 2
@pytest.mark.asyncio
async def test_max_steps_limit_is_enforced() -> None:
graph = _graph(
[
{"id": "start", "type": "start"},
{"id": "a", "type": "transform"},
{"id": "end", "type": "output"},
],
[
{"id": "e1", "source": "start", "target": "a"},
{"id": "e2", "source": "a", "target": "end"},
],
settings={"maxSteps": 2},
)
with pytest.raises(WorkflowError) as exc:
await _run(graph)
assert exc.value.code == "WORKFLOW_MAX_STEPS"
@pytest.mark.asyncio
async def test_human_input_pauses_then_resumes() -> None:
nodes = [
{"id": "start", "type": "start"},
{
"id": "ask",
"type": "human_input",
"config": {"prompt": "请确认 {{ inputs.topic }}", "formSchema": {"required": ["comment"]}, "actions": ["submit"]},
},
{"id": "end", "type": "output", "config": {"mapping": {"comment": "{{ nodes.ask.data.values.comment }}"}}},
]
edges = [
{"id": "e1", "source": "start", "target": "ask"},
{"id": "e2", "source": "ask", "target": "end"},
]
graph = _graph(nodes, edges)
with pytest.raises(WorkflowPaused) as paused:
await _run(graph, {"topic": "储能"})
assert paused.value.node_id == "ask"
assert paused.value.pending_input["prompt"] == "请确认 储能"
assert paused.value.pending_input["formSchema"]["required"] == ["comment"]
result, _ = await _run(
graph,
{"topic": "储能"},
resume_payload={"nodeId": "ask", "action": "submit", "values": {"comment": "同意"}},
)
assert result.output == {"comment": "同意"}
@pytest.mark.asyncio
async def test_human_input_rejects_missing_required_field() -> None:
graph = _graph(
[
{"id": "start", "type": "start"},
{"id": "ask", "type": "human_input", "config": {"formSchema": {"required": ["comment"]}}},
{"id": "end", "type": "output"},
],
[
{"id": "e1", "source": "start", "target": "ask"},
{"id": "e2", "source": "ask", "target": "end"},
],
)
with pytest.raises(WorkflowError) as exc:
await _run(graph, resume_payload={"nodeId": "ask", "action": "submit", "values": {}})
assert exc.value.code == "WORKFLOW_HUMAN_INPUT_INVALID"
@pytest.mark.asyncio
async def test_human_input_exposes_natural_language_reply_to_downstream_nodes() -> None:
graph = _graph(
[
{"id": "start", "type": "start"},
{"id": "ask", "type": "human_input", "config": {"actions": ["submit"]}},
{"id": "end", "type": "output", "config": {"mapping": {"reply": "{{ nodes.ask.data.values.message }}"}}},
],
[
{"id": "e1", "source": "start", "target": "ask"},
{"id": "e2", "source": "ask", "target": "end"},
],
)
result, _ = await _run(
graph,
resume_payload={"nodeId": "ask", "action": "submit", "values": {}, "message": "目标客户是连锁零售商"},
)
assert result.output == {"reply": "目标客户是连锁零售商"}
@pytest.mark.asyncio
async def test_evidence_normalizer_builds_a_bounded_traceable_evidence_pack() -> None:
graph = _graph(
[
{"id": "start", "type": "start"},
{
"id": "sales_research",
"type": "transform",
"config": {"operations": [{"op": "set", "target": "text", "value": "近四周新客转化率下降 12%"}]},
},
{
"id": "voc_research",
"type": "transform",
"config": {"operations": [{"op": "set", "target": "text", "value": "用户反馈显示结算流程复杂"}]},
},
{
"id": "evidence",
"type": "evidence_normalizer",
"config": {
"sources": [
{"nodeId": "sales_research", "label": "销售数据"},
{"nodeId": "voc_research", "label": "用户反馈"},
]
},
},
{
"id": "end",
"type": "output",
"config": {"mapping": {"evidence": "{{ nodes.evidence.data.evidencePack }}"}},
},
],
[
{"id": "e1", "source": "start", "target": "sales_research"},
{"id": "e2", "source": "start", "target": "voc_research"},
{"id": "e3", "source": "sales_research", "target": "evidence"},
{"id": "e4", "source": "voc_research", "target": "evidence"},
{"id": "e5", "source": "evidence", "target": "end"},
],
)
result, events = await _run(graph)
pack = result.output["evidence"]
assert [item["role"] for item in pack["items"]] == ["销售数据", "用户反馈"]
assert pack["items"][0]["claim"] == "近四周新客转化率下降 12%"
assert pack["items"][0]["source"]["uri"] == "workflow://run-1/nodes/sales_research"
assert pack["assumptions"] == [] and pack["gaps"] == []
assert any(
event["event"] == "node.progress"
and event["data"].get("phase") == "evidence.normalized"
for event in events
)
@pytest.mark.asyncio
async def test_deep_research_writer_streams_report_and_registers_markdown_artifact() -> None:
async def run_deep_research(**kwargs: Any) -> dict[str, Any]:
await kwargs["on_event"](
{"type": "phase_changed", "phase": "writing", "payload": {}}
)
await kwargs["on_event"](
{"type": "report_delta", "phase": "writing", "payload": {"delta": "# 储能报告"}}
)
assert kwargs["topic"] == "储能"
assert kwargs["evidence_pack"]["items"][0]["role"] == "行业研究"
return {
"session_id": "drs_wf_1",
"job_id": "drj_wf_1",
"report_markdown": "# 储能报告\n\n基于已提供材料。",
"source_count": 1,
}
async def save_artifact(**kwargs: Any) -> dict[str, Any]:
return {"id": "artifact-report", **kwargs}
graph = _graph(
[
{"id": "start", "type": "start"},
{
"id": "report",
"type": "deep_research_write",
"config": {
"topicTemplate": "{{ inputs.topic }}",
"evidenceBinding": {
"items": [{"claim": "市场需求上升", "role": "行业研究"}],
},
"reportInstructionTemplate": "写成正式报告",
},
},
{"id": "end", "type": "output", "config": {"mapping": {"text": "{{ nodes.report.data.text }}"}}},
],
[
{"id": "s-r", "source": "start", "target": "report"},
{"id": "r-e", "source": "report", "target": "end"},
],
)
result, events = await _run(
graph,
{"topic": "储能"},
deps=WorkflowRuntimeDeps(
run_deep_research=run_deep_research,
save_artifact=save_artifact,
),
)
assert result.output["text"].startswith("# 储能报告")
assert result.node_results["report"]["data"]["deepResearchJobId"] == "drj_wf_1"
assert result.node_results["report"]["artifacts"][0]["artifact_id"] == "artifact-report"
assert any(
event["event"] == "node.output.delta" and event["data"].get("delta") == "# 储能报告"
for event in events
)
def test_deep_research_writer_requires_evidence_pack_and_rejects_nested_multi_agent() -> None:
graph = _graph(
[
{"id": "start", "type": "start"},
{"id": "report", "type": "deep_research_write", "config": {"researchConfig": {"mode": "multi_agent"}}},
{"id": "end", "type": "output"},
],
[
{"id": "s-r", "source": "start", "target": "report"},
{"id": "r-e", "source": "report", "target": "end"},
],
)
issues = validate_workflow_graph(graph)
assert {issue.code for issue in issues} >= {"WORKFLOW_SCHEMA_INVALID"}
assert any("Evidence Pack" in issue.message for issue in issues)
assert any("multi_agent" in issue.message for issue in issues)
@pytest.mark.asyncio
async def test_workflow_deep_research_adapter_reuses_durable_job_contract() -> None:
class Sessions:
def __init__(self) -> None:
self.rows: dict[str, dict[str, Any]] = {}
async def get(self, session_id: str, *, user_id: str) -> dict[str, Any] | None:
return self.rows.get(session_id)
async def create(self, **fields: Any) -> dict[str, Any]:
row = {"status": "draft", "active_job_id": None, **fields}
self.rows[row["id"]] = row
return row
async def update(self, session_id: str, *, user_id: str, **fields: Any) -> dict[str, Any] | None:
row = self.rows.get(session_id)
if row is not None:
row.update(fields)
return row
async def count_active_by_user(self, *, user_id: str) -> int:
return 0
class Sources:
def __init__(self) -> None:
self.rows: list[dict[str, Any]] = []
async def upsert(self, **fields: Any) -> dict[str, Any]:
self.rows.append(fields)
return fields
class Jobs:
def __init__(self, sessions: Sessions) -> None:
self.sessions = sessions
self.row: dict[str, Any] | None = None
async def get_active_for_session(self, session_id: str, *, user_id: str) -> dict[str, Any] | None:
return self.row if self.row and self.row["status"] not in {"completed", "failed", "cancelled"} else None
async def try_create_or_get_active(self, **fields: Any) -> tuple[dict[str, Any], bool]:
self.row = {"id": fields["id"], "status": "queued", "phase": "initializing", **fields}
return self.row, True
async def get(self, job_id: str, *, user_id: str) -> dict[str, Any] | None:
assert self.row is not None and self.row["id"] == job_id
self.row["status"] = "completed"
session = self.sessions.rows[self.row["session_id"]]
session.update({"status": "completed", "report_markdown": "# 可追溯报告"})
return self.row
async def request_cancel(self, job_id: str, *, user_id: str) -> None:
raise AssertionError("test should not cancel")
class Events:
async def list_after(self, job_id: str, *, after: int, limit: int) -> list[dict[str, Any]]:
return [{"seq": 1, "event_type": "report_chunk", "phase": "writing", "payload": {"delta": "# 可追溯报告"}}] if after == 0 else []
class Dispatcher:
nudges = 0
def nudge(self) -> None:
self.nudges += 1
sessions = Sessions()
sources = Sources()
jobs = Jobs(sessions)
dispatcher = Dispatcher()
app = SimpleNamespace(
state=SimpleNamespace(
deep_research_session_store=sessions,
deep_research_job_store=jobs,
deep_research_event_store=Events(),
deep_research_source_store=sources,
deep_research_dispatcher=dispatcher,
deep_research_live_hub=None,
)
)
events: list[dict[str, Any]] = []
async def on_event(event: dict[str, Any]) -> None:
events.append(event)
outcome = await WorkflowDeepResearchAdapter(app).run(
run_id="run-1",
node_id="report",
owner_id="user-a",
topic="储能产业趋势",
evidence_pack={
"items": [
{
"claim": "上游研究指出储能装机需求增长。",
"role": "行业研究",
"source": {"nodeId": "research", "uri": "workflow://run-1/nodes/research"},
}
]
},
research_config={"mode": "detailed"},
report_instruction="写成正式报告",
cancel=CancelToken(),
on_event=on_event,
)
assert outcome["report_markdown"] == "# 可追溯报告"
assert jobs.row is not None and jobs.row["input_snapshot"]["entry"] == "regenerate"
assert jobs.row["input_snapshot"]["config"]["collection_mode"] == "legacy"
assert sources.rows[0]["selected"] is True
assert "工作流 Evidence Pack" in sources.rows[0]["raw_content"]
assert dispatcher.nudges == 1
assert any(event["type"] == "report_chunk" for event in events)
def test_evidence_normalizer_rejects_non_upstream_or_unknown_sources() -> None:
graph = _graph(
[
{"id": "start", "type": "start"},
{"id": "evidence", "type": "evidence_normalizer", "config": {"sources": [{"nodeId": "end"}, {"nodeId": "missing"}]}},
{"id": "end", "type": "output"},
],
[
{"id": "e1", "source": "start", "target": "evidence"},
{"id": "e2", "source": "evidence", "target": "end"},
],
)
issues = validate_workflow_graph(graph)
assert {issue.code for issue in issues} >= {
"WORKFLOW_EXPRESSION_INVALID",
"WORKFLOW_NODE_NOT_FOUND",
}
@pytest.mark.asyncio
async def test_node_failure_can_be_tolerated_with_on_error_continue() -> None:
graph = _graph(
[
{"id": "start", "type": "start"},
{
"id": "boom",
"type": "transform",
"config": {"onError": "continue", "operations": [{"op": "not_a_real_op"}]},
},
{"id": "end", "type": "output", "config": {"mapping": {"ok": "true"}}},
],
[
{"id": "e1", "source": "start", "target": "boom"},
{"id": "e2", "source": "boom", "target": "end"},
],
)
result, events = await _run(graph)
assert result.output == {"ok": "true"}
assert result.node_results["boom"]["metadata"]["failed"] is True
assert any(e["event"] == "node.failed" for e in events)
@pytest.mark.asyncio
async def test_agent_node_streams_deltas_and_parses_json() -> None:
async def run_agent(**kwargs: Any) -> dict[str, Any]:
on_delta = kwargs["on_delta"]
await on_delta("思考")
await on_delta("完成")
assert "储能" in kwargs["prompt"]
return {"text": '```json\n{"title": "储能报告"}\n```', "usage": {"tokens": 12}}
graph = _graph(
[
{"id": "start", "type": "start"},
{
"id": "writer",
"type": "agent",
"config": {
"agentId": "a1",
"promptTemplate": "写一份关于 {{ inputs.topic }} 的报告",
"responseMode": "json",
"responseSchema": {"required": ["title"]},
"parallelGroupId": "round-1",
"roundIndex": 1,
"agentName": "报告撰写智能体",
"agentAvatarUrl": "https://assets.example.com/writer.png",
},
},
{"id": "end", "type": "output", "config": {"mapping": {"title": "{{ nodes.writer.data.json.title }}"}}},
],
[
{"id": "e1", "source": "start", "target": "writer"},
{"id": "e2", "source": "writer", "target": "end"},
],
)
result, events = await _run(graph, {"topic": "储能"}, deps=WorkflowRuntimeDeps(run_agent=run_agent))
assert result.output == {"title": "储能报告"}
deltas = [(e["data"].get("delta") or e["data"].get("text")) for e in events if e["event"] == "node.output.delta"]
assert "".join(deltas) == "思考完成"
assert all(e["data"].get("channel") == "answer" for e in events if e["event"] == "node.output.delta")
started = next(event for event in events if event["event"] == "node.started" and event["nodeId"] == "writer")
assert started["data"]["nodeKind"] == "agent"
assert started["data"]["parallelGroupId"] == "round-1"
assert started["data"]["roundIndex"] == 1
assert started["data"]["agent"] == {
"agentId": "a1",
"name": "报告撰写智能体",
"avatarUrl": "https://assets.example.com/writer.png",
}
@pytest.mark.asyncio
async def test_agent_clarification_pauses_and_resume_is_routed_to_the_same_node() -> None:
seen_resume_payloads: list[dict[str, Any] | None] = []
async def run_agent(**kwargs: Any) -> dict[str, Any]:
resume_payload = kwargs.get("resume_payload")
seen_resume_payloads.append(resume_payload)
if resume_payload is None:
return {
"awaiting_input": {
"toolCallId": "clarify-1",
"prompt": "请确认分析的地域范围。",
"actions": ["submit"],
}
}
assert resume_payload["values"] == {"answer": "中国市场"}
return {"text": "已按中国市场范围完成分析。"}
graph = _graph(
[
{"id": "start", "type": "start"},
{
"id": "research",
"type": "agent",
"config": {"agentId": "researcher", "promptTemplate": "分析 {{ inputs.topic }}"},
},
{"id": "end", "type": "output", "config": {"mapping": {"text": "{{ nodes.research.data.text }}"}}},
],
[
{"id": "e1", "source": "start", "target": "research"},
{"id": "e2", "source": "research", "target": "end"},
],
)
deps = WorkflowRuntimeDeps(run_agent=run_agent)
with pytest.raises(WorkflowPaused) as paused:
await _run(graph, {"topic": "新能源车"}, deps=deps)
assert paused.value.node_id == "research"
assert paused.value.pending_input["toolCallId"] == "clarify-1"
result, _events = await _run(
graph,
{"topic": "新能源车"},
deps=deps,
resume_payload={
"nodeId": "research",
"values": {"answer": "中国市场"},
"message": "只分析中国市场",
},
)
assert result.output == {"text": "已按中国市场范围完成分析。"}
assert seen_resume_payloads == [None, {"nodeId": "research", "values": {"answer": "中国市场"}, "message": "只分析中国市场"}]
@pytest.mark.asyncio
async def test_agent_node_includes_rendered_workflow_input_bindings() -> None:
seen: dict[str, Any] = {}
async def run_agent(**kwargs: Any) -> dict[str, Any]:
seen.update(kwargs)
return {"text": "已按任务完成"}
graph = _graph(
[
{"id": "start", "type": "start"},
{
"id": "writer",
"type": "agent",
"config": {
"agentId": "a1",
"promptTemplate": "请完成用户的请求。",
"inputBindings": {
"task": "{{ inputs.query }}",
"allInputs": "{{ inputs }}",
},
},
},
{"id": "end", "type": "output", "config": {"mapping": {"text": "{{ nodes.writer.data.text }}"}}},
],
[
{"id": "e1", "source": "start", "target": "writer"},
{"id": "e2", "source": "writer", "target": "end"},
],
)
result, _ = await _run(
graph,
{"query": "分析本周销售数据", "region": "华东"},
deps=WorkflowRuntimeDeps(run_agent=run_agent),
)
assert result.output == {"text": "已按任务完成"}
assert '"task": "分析本周销售数据"' in seen["prompt"]
assert '"region": "华东"' in seen["prompt"]
assert "<workflow-input>" in seen["prompt"]
@pytest.mark.asyncio
async def test_agent_node_receives_feedback_only_when_its_revision_targets_it() -> None:
prompts: list[str] = []
async def run_agent(**kwargs: Any) -> dict[str, Any]:
prompts.append(kwargs["prompt"])
return {"text": "已修订"}
graph = _graph(
[
{"id": "start", "type": "start"},
{"id": "research", "type": "agent", "config": {"agentId": "a1", "promptTemplate": "分析 {{ inputs.topic }}"}},
{"id": "other", "type": "agent", "config": {"agentId": "a2", "promptTemplate": "不应接收反馈"}},
{"id": "end", "type": "output", "config": {"mapping": {"text": "{{ nodes.research.data.text }}"}}},
],
[
{"id": "s-r", "source": "start", "target": "research"},
{"id": "s-o", "source": "start", "target": "other"},
{"id": "r-e", "source": "research", "target": "end"},
],
)
await _run(
graph,
{"topic": "储能"},
deps=WorkflowRuntimeDeps(run_agent=run_agent),
env={"nodeFeedback": {"research": {"message": "请改从成本和风险角度分析"}}},
)
feedback_prompt = next(prompt for prompt in prompts if "分析" in prompt)
other_prompt = next(prompt for prompt in prompts if "不应接收反馈" in prompt)
assert "<workflow-user-feedback>" in feedback_prompt
assert "成本和风险" in feedback_prompt
assert "<workflow-user-feedback>" not in other_prompt
@pytest.mark.asyncio
async def test_missing_capability_becomes_typed_error() -> None:
graph = _graph(
[
{"id": "start", "type": "start"},
{"id": "writer", "type": "agent", "config": {"agentId": "a1", "promptTemplate": "hi"}},
],
[{"id": "e1", "source": "start", "target": "writer"}],
)
with pytest.raises(WorkflowError) as exc:
await _run(graph)
assert exc.value.code == "WORKFLOW_RESOURCE_MISSING"
@pytest.mark.asyncio
async def test_cancellation_stops_the_run() -> None:
run_store = MemoryWorkflowRunStore()
event_store = MemoryWorkflowEventStore()
graph = _graph(
[{"id": "start", "type": "start"}, {"id": "end", "type": "output"}],
[{"id": "e1", "source": "start", "target": "end"}],
)
sink = PersistingWorkflowEventSink(event_store, run_id="run-x", workflow_id=graph.id, version_id="v1")
cancel = CancelToken()
cancel.cancel()
ctx = RunContext(
run_id="run-x",
workflow_id=graph.id,
version_id="v1",
owner_id="u",
graph=graph,
inputs={},
deps=WorkflowRuntimeDeps(),
cancel=cancel,
emit=sink.emit,
)
engine = WorkflowEngine(build_default_registry(), run_store=run_store)
with pytest.raises(WorkflowError) as exc:
await engine.run(ctx)
assert exc.value.code == "WORKFLOW_CANCELLED"
# ── policies ────────────────────────────────────────────────────────────
def test_sql_policy_blocks_writes_and_stacked_statements() -> None:
assert assert_read_only("SELECT id FROM t WHERE x = :x").startswith("SELECT")
assert assert_read_only("WITH c AS (SELECT 1) SELECT * FROM c").startswith("WITH")
for bad in (
"DELETE FROM t",
"SELECT 1; DROP TABLE t",
"SELECT 1 -- ;\n; DROP TABLE t",
"SELECT * FROM t INTO OUTFILE '/tmp/x'",
"SELECT pg_sleep(10)",
"UPDATE t SET a = 1",
):
with pytest.raises(WorkflowError) as exc:
assert_read_only(bad)
assert exc.value.code == "WORKFLOW_SQL_POLICY_DENIED"
def test_sql_limit_is_appended_once() -> None:
assert enforce_limit("SELECT 1", 10) == "SELECT 1 LIMIT 10"
assert enforce_limit("SELECT 1 LIMIT 5", 10) == "SELECT 1 LIMIT 5"
def test_sql_table_allowlist_blocks_unlisted_tables() -> None:
assert_tables_allowed("SELECT id FROM report_source", ["report_source"])
with pytest.raises(WorkflowError) as exc:
assert_tables_allowed("SELECT id FROM secrets", ["report_source"])
assert exc.value.code == "WORKFLOW_SQL_POLICY_DENIED"
def test_downstream_node_reference_is_rejected_at_publish() -> None:
issues = validate_workflow_graph(
{
"schemaVersion": "1.0",
"id": "wf_ref",
"nodes": [
{"id": "start", "type": "start"},
{
"id": "early",
"type": "transform",
"config": {"operations": [{"op": "set", "target": "x", "value": "{{ nodes.late.data.x }}"}]},
},
{
"id": "late",
"type": "transform",
"config": {"operations": [{"op": "set", "target": "x", "value": "1"}]},
},
{"id": "end", "type": "output", "config": {"mapping": {"x": "{{ nodes.late.data.x }}"}}},
],
"edges": [
{"id": "e1", "source": "start", "target": "early"},
{"id": "e2", "source": "early", "target": "late"},
{"id": "e3", "source": "late", "target": "end"},
],
}
)
assert any(i.code == "WORKFLOW_EXPRESSION_INVALID" and i.node_id == "early" for i in issues)
@pytest.mark.asyncio
async def test_http_policy_blocks_private_and_metadata_targets() -> None:
policy = WorkflowConfig().http
for url in (
"http://127.0.0.1:8001/admin",
"http://169.254.169.254/latest/meta-data/",
"http://localhost/",
"file:///etc/passwd",
"http://example.com:22/",
):
with pytest.raises(WorkflowError) as exc:
await enforce_http_target(url, policy)
assert exc.value.code in ("WORKFLOW_HTTP_POLICY_DENIED", "WORKFLOW_HTTP_FAILED")
@pytest.mark.asyncio
async def test_http_policy_allows_private_when_configured() -> None:
policy = WorkflowConfig().http.model_copy(update={"allow_private_networks": True, "allowed_ports": [8001]})
target = await enforce_http_target("http://127.0.0.1:8001/ping", policy)
assert target["host"] == "127.0.0.1"
@pytest.mark.asyncio
async def test_event_sink_persists_before_publishing() -> None:
store = MemoryWorkflowEventStore()
seen: list[int] = []
async def publish(event: dict[str, Any]) -> None:
# Persisted rows must already be readable when the live copy arrives.
rows = await store.list_after("r1", after_seq=0)
seen.append(len(rows))
sink = PersistingWorkflowEventSink(store, run_id="r1", workflow_id="w1", version_id="v1", publish=publish)
await sink.emit("run.started")
await sink.emit("run.completed")
assert seen == [1, 2]
assert [r["seq"] for r in await store.list_after("r1")] == [1, 2]
@pytest.mark.asyncio
async def test_run_store_lease_and_idempotency() -> None:
from datetime import UTC, datetime, timedelta
store = MemoryWorkflowRunStore()
payload = {
"workflow_id": "w1",
"workflow_version_id": "v1",
"owner_id": "u1",
"idempotency_key": "k1",
"input": {"a": 1},
}
run, created = await store.create_run(payload)
assert created is True
again, created_again = await store.create_run(payload)
assert created_again is False and again["id"] == run["id"]
until = datetime.now(UTC) + timedelta(seconds=30)
claimed = await store.claim_run(run["id"], lease_owner="w-a", lease_until=until)
assert claimed is not None and claimed["status"] == "running"
# A second worker cannot steal a live lease.
assert await store.claim_run(run["id"], lease_owner="w-b", lease_until=until) is None
assert await store.renew_lease(run["id"], lease_owner="w-a", lease_until=until) is True
assert await store.renew_lease(run["id"], lease_owner="w-b", lease_until=until) is False
# Expired lease is reclaimable (crash recovery).
await store.renew_lease(run["id"], lease_owner="w-a", lease_until=datetime.now(UTC) - timedelta(seconds=1))
assert await store.claim_run(run["id"], lease_owner="w-b", lease_until=until) is not None
paused = await store.set_awaiting_input(run["id"], lease_owner="w-b", pending_input={"nodeId": "ask"}, resume_token="t1")
assert paused is not None and paused["status"] == "awaiting_input"
assert await store.consume_resume_token(run["id"], "wrong") is None
resumed = await store.consume_resume_token(run["id"], "t1")
assert resumed is not None and resumed["status"] == "queued"
# One-shot: the same token cannot be replayed.
assert await store.consume_resume_token(run["id"], "t1") is None
@pytest.mark.asyncio
async def test_code_node_disabled_by_default() -> None:
graph = _graph(
[
{"id": "start", "type": "start"},
{"id": "calc", "type": "code", "config": {"source": "output['x'] = 1"}},
],
[{"id": "e1", "source": "start", "target": "calc"}],
)
with pytest.raises(WorkflowError) as exc:
await _run(graph)
assert exc.value.code == "WORKFLOW_CODE_SANDBOX_FAILED"
@pytest.mark.asyncio
async def test_code_node_executes_when_enabled() -> None:
config = WorkflowConfig()
config.code.enabled = True
graph = _graph(
[
{"id": "start", "type": "start"},
{
"id": "calc",
"type": "code",
"config": {
"source": "output['doubled'] = inputs['n'] * 2",
"inputs": {"n": "{{ inputs.n }}"},
},
},
{"id": "end", "type": "output", "config": {"mapping": {"doubled": "{{ nodes.calc.data.doubled }}"}}},
],
[
{"id": "e1", "source": "start", "target": "calc"},
{"id": "e2", "source": "calc", "target": "end"},
],
)
result, _ = await _run(graph, {"n": 21}, config=config)
assert result.output == {"doubled": 42}
@pytest.mark.asyncio
async def test_code_node_cannot_open_sockets() -> None:
config = WorkflowConfig()
config.code.enabled = True
graph = _graph(
[
{"id": "start", "type": "start"},
{
"id": "calc",
"type": "code",
"config": {"source": "import socket\nsocket.socket()"},
},
],
[{"id": "e1", "source": "start", "target": "calc"}],
)
with pytest.raises(WorkflowError) as exc:
await _run(graph, config=config)
assert exc.value.code == "WORKFLOW_CODE_SANDBOX_FAILED"
def test_asyncio_smoke() -> None:
# Guards against an event-loop policy regression on Windows runners.
assert asyncio.get_event_loop_policy() is not None
@pytest.mark.asyncio
async def test_skill_node_invokes_agent_with_restricted_skills() -> None:
seen: dict[str, Any] = {}
async def run_agent(**kwargs: Any) -> dict[str, Any]:
seen.update(kwargs)
return {"text": "已完成检索", "thread_id": "thr-1"}
graph = _graph(
[
{"id": "start", "type": "start"},
{
"id": "skill",
"type": "skill",
"config": {
"agentId": "collector",
"skillNames": ["web_search"],
"promptTemplate": "查 {{ inputs.topic }}",
"mode": "agent_skill",
},
},
{"id": "end", "type": "output", "config": {"mapping": {"text": "{{ nodes.skill.data.text }}"}}},
],
[
{"id": "e1", "source": "start", "target": "skill"},
{"id": "e2", "source": "skill", "target": "end"},
],
)
result, _published = await _run(graph, {"topic": "储能"}, deps=WorkflowRuntimeDeps(run_agent=run_agent))
assert result.output == {"text": "已完成检索"}
assert seen["agent_id"] == "collector"
assert seen["skill_names"] == ["web_search"]
assert "储能" in seen["prompt"]
@pytest.mark.asyncio
async def test_callable_skill_mode_is_rejected() -> None:
graph = _graph(
[
{"id": "start", "type": "start"},
{
"id": "skill",
"type": "skill",
"config": {"agentId": "collector", "skillNames": ["web_search"], "mode": "callable_skill"},
},
],
[{"id": "e1", "source": "start", "target": "skill"}],
)
with pytest.raises(WorkflowError) as exc:
await _run(graph)
assert exc.value.code == "WORKFLOW_SKILL_FAILED"
@pytest.mark.asyncio
async def test_subworkflow_node_folds_child_output() -> None:
async def run_subworkflow(**kwargs: Any) -> dict[str, Any]:
assert kwargs["workflow_id"] == "child-wf"
assert kwargs["version_id"] == "child-v1"
assert kwargs["inputs"] == {"topic": "储能"}
return {"run_id": "child-run", "status": "completed", "output": {"title": "子报告"}}
graph = _graph(
[
{"id": "start", "type": "start"},
{
"id": "child",
"type": "subworkflow",
"config": {
"workflowId": "child-wf",
"versionId": "child-v1",
"inputMapping": {"topic": "{{ inputs.topic }}"},
},
},
{"id": "end", "type": "output", "config": {"mapping": {"title": "{{ nodes.child.data.title }}"}}},
],
[
{"id": "e1", "source": "start", "target": "child"},
{"id": "e2", "source": "child", "target": "end"},
],
)
result, _ = await _run(graph, {"topic": "储能"}, deps=WorkflowRuntimeDeps(run_subworkflow=run_subworkflow))
assert result.output == {"title": "子报告"}