Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 38 additions & 1 deletion dsl_runtime/lang/expression.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from __future__ import annotations

import ast
from collections.abc import Mapping
import operator
import time
from typing import Any, Dict
Expand Down Expand Up @@ -28,6 +29,29 @@
ast.UAdd: operator.pos,
}

_BLOCKED_ATTRS = frozenset(
{
"__class__",
"__dict__",
"__globals__",
"__mro__",
"__subclasses__",
"__bases__",
"__base__",
"__code__",
"__closure__",
"__func__",
"__self__",
"__module__",
}
)


def _validate_attr_name(attr: str) -> None:
"""Reject private/dunder attributes and non-identifier attribute names."""
if not attr.isidentifier() or attr.startswith("_") or attr in _BLOCKED_ATTRS:
raise ValueError(f"不允许访问属性: {attr}")


class SafeEvaluator(ast.NodeVisitor):
"""安全表达式求值器,支持算术/比较/逻辑/变量。"""
Expand Down Expand Up @@ -86,8 +110,21 @@ def visit_IfExp(self, node: ast.IfExp) -> Any: # type: ignore[override]
return self.visit(node.body) if self.visit(node.test) else self.visit(node.orelse)

def visit_Attribute(self, node: ast.Attribute) -> Any: # type: ignore[override]
attr = node.attr
_validate_attr_name(attr)
value = self.visit(node.value)
return getattr(value, node.attr)

if isinstance(value, Mapping):
if attr in value:
return value[attr]
raise KeyError(attr)

# Only expose data attributes that are explicitly present on an instance.
# This avoids reaching descriptors/classes such as __class__, mro, or methods.
instance_attrs = getattr(value, "__dict__", {})
if isinstance(instance_attrs, dict) and attr in instance_attrs:
return instance_attrs[attr]
raise ValueError(f"不允许访问属性: {attr}")

def visit_Subscript(self, node: ast.Subscript) -> Any: # type: ignore[override]
value = self.visit(node.value)
Expand Down
Loading