268 lines
8.4 KiB
Python
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",
|
|
]
|