"""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.`` 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", ]