deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/workflows/expressions.py
2026-09-07 18:24:55 +08:00

268 lines
8.4 KiB
Python

"""Safe expression / template evaluation for workflow node configs.
Two surfaces, both intentionally tiny:
* ``render_template("… {{ nodes.a.data.title }} …", state)`` — string
interpolation over a whitelisted state tree.
* ``evaluate_condition("nodes.a.data.score > 0.5", state)`` — boolean
predicate over the same tree.
Neither uses :func:`eval`. The condition grammar is a single comparison or a
chain joined by ``and`` / ``or``; operands are either a state path or a JSON
literal. Anything else raises ``WORKFLOW_EXPRESSION_INVALID`` so a malformed
graph fails loudly at publish time instead of silently at run time.
"""
from __future__ import annotations
import json
import re
from typing import Any
from deerflow.workflows.errors import WorkflowError
_TEMPLATE_RE = re.compile(r"\{\{\s*([^{}]+?)\s*\}\}")
_ROOT_KEYS = ("inputs", "nodes", "run", "loop", "env")
# Coze canvas node ids are commonly numeric (for example ``100001``). They
# are still dictionary keys, not Python identifiers, so permit a numeric
# dotted segment while keeping roots and normal field names restricted.
_PATH_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*(?:\.(?:[A-Za-z_][A-Za-z0-9_\-]*|\d+)|\[\d+\])*$")
_COMPARATORS = ("==", "!=", ">=", "<=", ">", "<", " in ", " contains ")
_MAX_RENDER_CHARS = 200_000
def _invalid(expr: str, reason: str) -> WorkflowError:
return WorkflowError(
"WORKFLOW_EXPRESSION_INVALID",
f"表达式不合法:{reason}",
details={"expression": expr[:200]},
)
def resolve_path(path: str, state: dict[str, Any]) -> Any:
"""Read ``a.b[0].c`` from the state tree. Missing segments yield ``None``."""
expr = path.strip()
if not _PATH_RE.match(expr):
raise _invalid(expr, "路径格式不合法")
root = expr.split(".", 1)[0].split("[", 1)[0]
if root not in _ROOT_KEYS:
raise _invalid(expr, f"根节点必须是 {'/'.join(_ROOT_KEYS)} 之一")
current: Any = state
for segment in re.findall(r"[A-Za-z_][A-Za-z0-9_\-]*|\d+|\[\d+\]", expr):
if current is None:
return None
if segment.startswith("["):
index = int(segment[1:-1])
if not isinstance(current, (list, tuple)) or index >= len(current):
return None
current = current[index]
elif isinstance(current, dict):
current = current.get(segment)
else:
return None
return current
def _stringify(value: Any) -> str:
if value is None:
return ""
if isinstance(value, str):
return value
if isinstance(value, bool):
return "true" if value else "false"
if isinstance(value, (int, float)):
return str(value)
return json.dumps(value, ensure_ascii=False, default=str)
def render_template(template: str, state: dict[str, Any]) -> str:
"""Replace every ``{{ path }}`` with its stringified state value."""
if not template:
return ""
def _sub(match: re.Match[str]) -> str:
return _stringify(resolve_path(match.group(1), state))
rendered = _TEMPLATE_RE.sub(_sub, template)
if len(rendered) > _MAX_RENDER_CHARS:
raise WorkflowError(
"WORKFLOW_LIMIT_EXCEEDED",
"模板渲染结果超出长度上限",
details={"limit": _MAX_RENDER_CHARS},
)
return rendered
def render_value(value: Any, state: dict[str, Any]) -> Any:
"""Recursively render templates inside dict / list / str structures.
A string that is exactly one ``{{ path }}`` keeps the resolved value's
native type instead of being stringified.
"""
if isinstance(value, str):
match = _TEMPLATE_RE.fullmatch(value.strip())
if match:
return resolve_path(match.group(1), state)
return render_template(value, state)
if isinstance(value, dict):
return {k: render_value(v, state) for k, v in value.items()}
if isinstance(value, list):
return [render_value(v, state) for v in value]
return value
def _operand(token: str, state: dict[str, Any]) -> Any:
text = token.strip()
if not text:
raise _invalid(token, "操作数为空")
match = _TEMPLATE_RE.fullmatch(text)
if match:
return resolve_path(match.group(1), state)
lowered = text.lower()
if lowered in ("true", "false"):
return lowered == "true"
if lowered in ("null", "none"):
return None
if (text[0], text[-1]) in (('"', '"'), ("'", "'")) and len(text) >= 2:
return text[1:-1]
try:
return json.loads(text)
except json.JSONDecodeError:
pass
if _PATH_RE.match(text) and text.split(".", 1)[0].split("[", 1)[0] in _ROOT_KEYS:
return resolve_path(text, state)
return text
def _truthy(value: Any) -> bool:
if isinstance(value, str):
return value.strip().lower() not in ("", "false", "0", "null", "none")
return bool(value)
def _compare(left: Any, op: str, right: Any) -> bool:
if op == "==":
return left == right
if op == "!=":
return left != right
if op == "in":
try:
return left in right # type: ignore[operator]
except TypeError:
return False
if op == "contains":
try:
return right in left # type: ignore[operator]
except TypeError:
return False
try:
if op == ">":
return left > right
if op == "<":
return left < right
if op == ">=":
return left >= right
if op == "<=":
return left <= right
except TypeError:
return False
raise _invalid(op, f"不支持的比较运算符 {op}")
def _split_top_level(expr: str, keyword: str) -> list[str] | None:
"""Split on ``keyword`` outside quotes / braces; None when absent."""
parts: list[str] = []
depth = 0
quote: str | None = None
buffer: list[str] = []
index = 0
pattern = f" {keyword} "
while index < len(expr):
char = expr[index]
if quote:
buffer.append(char)
if char == quote:
quote = None
index += 1
continue
if char in ('"', "'"):
quote = char
buffer.append(char)
index += 1
continue
if char in "[{(":
depth += 1
elif char in "]})":
depth -= 1
if depth == 0 and expr[index : index + len(pattern)] == pattern:
parts.append("".join(buffer))
buffer = []
index += len(pattern)
continue
buffer.append(char)
index += 1
if not parts:
return None
parts.append("".join(buffer))
return parts
def evaluate_condition(expression: str, state: dict[str, Any]) -> bool:
"""Evaluate a boolean predicate without ``eval``."""
expr = (expression or "").strip()
if not expr:
return False
or_parts = _split_top_level(expr, "or")
if or_parts:
return any(evaluate_condition(part, state) for part in or_parts)
and_parts = _split_top_level(expr, "and")
if and_parts:
return all(evaluate_condition(part, state) for part in and_parts)
if expr.startswith("not "):
return not evaluate_condition(expr[4:], state)
if expr.startswith("(") and expr.endswith(")"):
return evaluate_condition(expr[1:-1], state)
for comparator in _COMPARATORS:
op = comparator.strip()
index = expr.find(comparator)
if index <= 0:
continue
left = _operand(expr[:index], state)
right = _operand(expr[index + len(comparator) :], state)
return _compare(left, op, right)
return _truthy(_operand(expr, state))
_NODES_IDENT_RE = re.compile(r"(?:\{\{\s*)?nodes\.([A-Za-z_0-9][A-Za-z0-9_\-]*)")
def collect_referenced_nodes(value: Any) -> set[str]:
"""Node ids referenced through ``nodes.<id>`` anywhere in a config tree.
Covers both ``{{ nodes.x.data }}`` templates and bare condition paths
such as ``nodes.x.data.score > 0.5``.
"""
found: set[str] = set()
if isinstance(value, str):
found.update(_NODES_IDENT_RE.findall(value))
elif isinstance(value, dict):
for item in value.values():
found |= collect_referenced_nodes(item)
elif isinstance(value, list):
for item in value:
found |= collect_referenced_nodes(item)
return found
__all__ = [
"collect_referenced_nodes",
"evaluate_condition",
"render_template",
"render_value",
"resolve_path",
]