Skip to content
Open
Show file tree
Hide file tree
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
194 changes: 175 additions & 19 deletions burr/core/action.py
Original file line number Diff line number Diff line change
Expand Up @@ -431,6 +431,23 @@ def tags(self) -> list[str]:
ast.IsNot,
)

# Resource bounds for safe_expr. The grammar above is side-effect free, but a
# handful of operators can still produce values whose size is not bounded by
# the size of the expression (``10 ** 10 ** 10``, ``"a" * 10000000000``). We
# keep those operators and bound them instead. Each bound is enforced at run
# time by the interpreter; the ones that involve only literals are also
# enforced up front at ``safe_expr()`` call time.
_SAFE_EXPR_MAX_SOURCE_CHARS = 10_000 # length of the expression string, checked before parsing
_SAFE_EXPR_MAX_NODES = 500 # total AST nodes in one expression
_SAFE_EXPR_MAX_INT_BITS = 4096 # width of any int produced by arithmetic (or written as a literal)
_SAFE_EXPR_MAX_POW_EXPONENT = 64 # |exponent| for int ** int
_SAFE_EXPR_MAX_POW_BASE_BITS = 64 # width of the base for int ** int (64 bits * 64 = 4096 bits)
_SAFE_EXPR_MAX_SEQ_REPEAT = 10_000 # n in ``sequence * n``
_SAFE_EXPR_MAX_SEQ_LEN = 1_000_000 # length of a sequence produced by ``*`` or ``+``

# Sequence types whose ``*`` / ``+`` allocate proportionally to their operands.
_SAFE_EXPR_SEQUENCE_TYPES = (str, bytes, bytearray, list, tuple)


class _SafeExprValidator(ast.NodeVisitor):
"""Walks an AST and raises ``ValueError`` on any node not on the allowlist.
Expand All @@ -451,6 +468,8 @@ def visit_Expression(self, node: ast.Expression) -> None:

def visit_Constant(self, node: ast.Constant) -> None:
if isinstance(node.value, (int, float, str, bool)) or node.value is None:
if isinstance(node.value, int) and node.value.bit_length() > _SAFE_EXPR_MAX_INT_BITS:
self._reject(node, f"integer literal wider than {_SAFE_EXPR_MAX_INT_BITS} bits")
return
self._reject(node, f"constant of type {type(node.value).__name__}")

Expand Down Expand Up @@ -654,22 +673,70 @@ def _eval_BinOp(self, node: ast.BinOp):
right = self.eval(node.right)
op = node.op
if isinstance(op, ast.Add):
return left + right
if isinstance(op, ast.Sub):
return left - right
if isinstance(op, ast.Mult):
result = self._add(left, right)
elif isinstance(op, ast.Sub):
result = left - right
elif isinstance(op, ast.Mult):
result = self._mult(left, right)
elif isinstance(op, ast.Div):
result = left / right
elif isinstance(op, ast.FloorDiv):
result = left // right
elif isinstance(op, ast.Mod):
if isinstance(left, (str, bytes, bytearray)):
# ``"%0999999999d" % x`` allocates according to the format string, so
# printf-style formatting is not offered; use str() / comparisons instead.
raise ValueError("safe_expr: % formatting of strings is not supported")
result = left % right
elif isinstance(op, ast.Pow):
result = self._pow(left, right)
else: # pragma: no cover - validator catches this
raise ValueError(f"safe_expr: unsupported binary op {type(op).__name__}")
# General backstop: no arithmetic step may produce an int wider than the
# bound, regardless of which operator produced it.
if isinstance(result, int) and result.bit_length() > _SAFE_EXPR_MAX_INT_BITS:
raise ValueError(f"safe_expr: integer result wider than {_SAFE_EXPR_MAX_INT_BITS} bits")
return result

@staticmethod
def _add(left, right):
if isinstance(left, _SAFE_EXPR_SEQUENCE_TYPES) and isinstance(
right, _SAFE_EXPR_SEQUENCE_TYPES
):
if len(left) + len(right) > _SAFE_EXPR_MAX_SEQ_LEN:
raise ValueError(f"safe_expr: sequence result longer than {_SAFE_EXPR_MAX_SEQ_LEN}")
return left + right

@staticmethod
def _mult(left, right):
# ``sequence * n`` (either order) allocates len(sequence) * n; bound both the
# repeat count and the resulting length. Plain numeric products fall through.
if isinstance(left, _SAFE_EXPR_SEQUENCE_TYPES) and isinstance(right, int):
seq, n = left, right
elif isinstance(right, _SAFE_EXPR_SEQUENCE_TYPES) and isinstance(left, int):
seq, n = right, left
else:
return left * right
if isinstance(op, ast.Div):
return left / right
if isinstance(op, ast.FloorDiv):
return left // right
if isinstance(op, ast.Mod):
return left % right
if isinstance(op, ast.Pow):
return left**right
raise ValueError( # pragma: no cover
f"safe_expr: unsupported binary op {type(op).__name__}"
)
if n > _SAFE_EXPR_MAX_SEQ_REPEAT:
raise ValueError(
f"safe_expr: sequence repeat count exceeds {_SAFE_EXPR_MAX_SEQ_REPEAT}"
)
if len(seq) * n > _SAFE_EXPR_MAX_SEQ_LEN:
raise ValueError(f"safe_expr: sequence result longer than {_SAFE_EXPR_MAX_SEQ_LEN}")
return seq * n

@staticmethod
def _pow(left, right):
if not isinstance(left, (int, float)) or not isinstance(right, (int, float)):
raise ValueError("safe_expr: ** requires int or float operands")
if isinstance(left, int) and isinstance(right, int):
# int ** int is the only path whose result grows without limit; floats
# either stay within float range or raise OverflowError immediately.
if abs(right) > _SAFE_EXPR_MAX_POW_EXPONENT:
raise ValueError("safe_expr: exponent out of range")
if left.bit_length() > _SAFE_EXPR_MAX_POW_BASE_BITS:
raise ValueError("safe_expr: base of ** out of range")
return left**right

def _eval_Tuple(self, node: ast.Tuple):
return tuple(self.eval(e) for e in node.elts)
Expand All @@ -694,6 +761,64 @@ def _eval_Call(self, node: ast.Call):
return func(*args)


def _check_safe_expr_size(tree: ast.AST) -> None:
"""Rejects expressions with more than ``_SAFE_EXPR_MAX_NODES`` AST nodes.

Uses the iterative :func:`ast.walk`, so this runs safely *before* the
recursive validator / interpreter ever see a very deep tree.
"""
count = 0
for _ in ast.walk(tree):
count += 1
if count > _SAFE_EXPR_MAX_NODES:
raise ValueError(f"safe_expr: expression has more than {_SAFE_EXPR_MAX_NODES} nodes")


# Node types that make up a literal-only subtree: no names, no state, no calls.
_SAFE_EXPR_LITERAL_NODES = (
ast.Constant,
ast.BinOp,
ast.UnaryOp,
ast.Tuple,
ast.List,
ast.Set,
ast.Dict,
ast.operator,
ast.unaryop,
ast.expr_context,
)


def _check_safe_expr_literal_arithmetic(tree: ast.AST) -> None:
"""Evaluates literal-only arithmetic subtrees now, so that out-of-range
expressions such as ``10 ** 10 ** 10`` are rejected at ``safe_expr()``
call time rather than at run time.

A literal-only subtree has the same value regardless of state, so a bound
it trips here is one it would trip on every run. Only the interpreter's
``ValueError`` bounds are surfaced; anything else (e.g. ``ZeroDivisionError``
from ``1 / 0``) is left to run time, as before.
"""
interpreter = _SafeExprInterpreter({})

def is_literal_only(node: ast.AST) -> bool:
return all(isinstance(n, _SAFE_EXPR_LITERAL_NODES) for n in ast.walk(node))

def visit(node: ast.AST) -> None:
if isinstance(node, ast.BinOp) and is_literal_only(node):
try:
interpreter.eval(node)
except ValueError:
raise
except Exception:
pass
return
for child in ast.iter_child_nodes(node):
visit(child)

visit(tree)


class Condition(Function):
KEY = "PROCEED"

Expand Down Expand Up @@ -820,17 +945,48 @@ def safe_expr(expr: str) -> "Condition":
comprehensions and generator expressions, the walrus operator,
``await``, ``yield``, imports, and any ``Call`` not on the builtin allowlist.

Resource use is bounded as well, so a single expression cannot allocate
without limit. The expression string may be at most 10,000 characters and
may have at most 500 AST nodes; an expression nested too deeply for the
parser is rejected the same way. ``**`` accepts
only ``int`` / ``float`` operands; for ``int ** int`` the exponent must be
within +/-64 and the base must fit in 64 bits. ``sequence * n`` requires
``n <= 10_000``, and ``*`` / ``+`` on sequences may not produce more than
1,000,000 elements. ``%`` is not applied to strings (no printf-style
formatting). No arithmetic step (and no integer literal) may produce an
``int`` wider than 4096 bits. Bounds that depend only on literals are
checked at ``safe_expr()`` call time; the rest raise ``ValueError`` from
``run()``.

:param expr: Expression to evaluate
:return: A condition that evaluates the given expression
:raises ValueError: if the expression contains any disallowed construct
(raised at call time -- the condition is rejected before it ever runs).
:raises ValueError: if the expression contains any disallowed construct, or a
literal-only subexpression exceeds the bounds above (raised at call time --
the condition is rejected before it ever runs). Also raised from ``run()``
when a state-dependent value exceeds the bounds.
:raises SyntaxError: if the expression is not syntactically valid Python.
"""
# Parse first. This raises SyntaxError for malformed input, which is fine.
tree = ast.parse(expr, mode="eval")
# Bound the source length before parsing: the parser's own cost (and its
# nesting limit, which varies by Python version) scales with the input.
if len(expr) > _SAFE_EXPR_MAX_SOURCE_CHARS:
raise ValueError(
f"safe_expr: expression exceeds {_SAFE_EXPR_MAX_SOURCE_CHARS} characters"
)
# Parse. This raises SyntaxError for malformed input, which is fine. A
# deeply nested expression can instead exhaust the parser itself on some
# Python versions; surface that as the same kind of rejection.
try:
tree = ast.parse(expr, mode="eval")
except (RecursionError, MemoryError):
raise ValueError("safe_expr: expression is too deeply nested") from None
# Bound the size of the tree before handing it to anything recursive.
_check_safe_expr_size(tree)
# Validate the whole tree against the allowlist *now*, at call time. If any
# disallowed node exists we raise here, before constructing the Condition.
_SafeExprValidator().visit(tree)
# Literal-only arithmetic has a state-independent value, so its bounds can
# be (and are) checked here too.
_check_safe_expr_literal_arithmetic(tree)

# Collect Name references for keys, mirroring expr().
all_builtins = builtins.__dict__
Expand Down
136 changes: 136 additions & 0 deletions tests/core/test_action.py
Original file line number Diff line number Diff line change
Expand Up @@ -563,6 +563,142 @@ def test_safe_expr_does_not_eval_during_validation():
cond.run(State({"x": 0}))


# --- resource bounds --------------------------------------------------------


def test_safe_expr_pow_small_works():
assert _eval_value(Condition.safe_expr("2 ** 10"), {}) == 1024
assert _eval_value(Condition.safe_expr("x ** y"), {"x": 2, "y": 10}) == 1024
assert _eval_value(Condition.safe_expr("2 ** 0.5"), {}) == 2**0.5
# At the ceiling: a 64-bit base raised to 64 gives a 4033-bit result, still allowed.
assert _eval_value(Condition.safe_expr("x ** 64"), {"x": 2**63}) == (2**63) ** 64


def test_safe_expr_pow_literal_out_of_range_rejected_at_call_time():
with pytest.raises(ValueError, match="exponent out of range"):
Condition.safe_expr("10 ** 10 ** 10")
with pytest.raises(ValueError, match="exponent out of range"):
Condition.safe_expr("2 ** 65")
with pytest.raises(ValueError, match=r"base of \*\* out of range"):
Condition.safe_expr("18446744073709551616 ** 2") # 2 ** 64 is 65 bits wide
# A literal subexpression nested under a state-dependent one is still caught here.
with pytest.raises(ValueError, match="exponent out of range"):
Condition.safe_expr("x > 10 ** 10 ** 10")


def test_safe_expr_pow_state_out_of_range_rejected_at_run_time():
cond = Condition.safe_expr("x ** y")
assert cond.run(State({"x": 2, "y": 10})) == {Condition.KEY: True}
with pytest.raises(ValueError, match="exponent out of range"):
cond.run(State({"x": 10, "y": 10**10}))
with pytest.raises(ValueError, match="exponent out of range"):
cond.run(State({"x": 2, "y": -65}))
with pytest.raises(ValueError, match=r"base of \*\* out of range"):
cond.run(State({"x": 2**64, "y": 2}))
with pytest.raises(ValueError, match="int or float operands"):
cond.run(State({"x": "a", "y": 2}))


def test_safe_expr_pow_float_path_left_to_python():
# Floats cannot grow in memory, so float ** is not width-bounded: Python either
# returns a float or raises OverflowError straight away.
assert _eval_value(Condition.safe_expr("x ** y"), {"x": 2.0, "y": 100.0}) == 2.0**100
assert _eval_value(Condition.safe_expr("x ** y"), {"x": 0.5, "y": 10**10}) == 0.0
with pytest.raises(OverflowError):
_eval_value(Condition.safe_expr("x ** y"), {"x": 1e308, "y": 2})


def test_safe_expr_sequence_repeat_small_works():
assert _eval_value(Condition.safe_expr('"a" * 3'), {}) == "aaa"
assert _eval_value(Condition.safe_expr("3 * [0]"), {}) == [0, 0, 0]
assert _eval_value(Condition.safe_expr("s * n"), {"s": "ab", "n": 2}) == "abab"
assert _eval_value(Condition.safe_expr("t * 2"), {"t": (1,)}) == (1, 1)


def test_safe_expr_sequence_repeat_out_of_range_rejected():
with pytest.raises(ValueError, match="repeat count"):
Condition.safe_expr('"a" * 10000000000')
with pytest.raises(ValueError, match="repeat count"):
Condition.safe_expr("10000000000 * [0]")
# Each step within the repeat-count bound is still capped on the resulting length.
with pytest.raises(ValueError, match="longer than"):
Condition.safe_expr('("a" * 10000) * 10000')
cond = Condition.safe_expr("s * n")
assert cond.run(State({"s": "a", "n": 3})) == {Condition.KEY: True}
with pytest.raises(ValueError, match="repeat count"):
cond.run(State({"s": "a", "n": 10000000000}))
with pytest.raises(ValueError, match="longer than"):
cond.run(State({"s": "a" * 1000, "n": 10_000}))


def test_safe_expr_sequence_concat_bounded():
assert _eval_value(Condition.safe_expr("s + t"), {"s": "a", "t": "b"}) == "ab"
assert _eval_value(Condition.safe_expr("s + t"), {"s": [1], "t": [2]}) == [1, 2]
cond = Condition.safe_expr("s + t")
with pytest.raises(ValueError, match="longer than"):
cond.run(State({"s": "a" * 600_000, "t": "b" * 600_000}))


def test_safe_expr_string_formatting_not_supported():
assert _eval_value(Condition.safe_expr("x % 3"), {"x": 10}) == 1
cond = Condition.safe_expr("s % n")
with pytest.raises(ValueError, match="formatting"):
cond.run(State({"s": "%0999d", "n": 1}))


def test_safe_expr_int_result_width_bounded():
big = 2**4000
assert _eval_value(Condition.safe_expr("x + 1"), {"x": big}) == big + 1
cond = Condition.safe_expr("x * y")
assert cond.run(State({"x": 7, "y": 6})) == {Condition.KEY: True}
with pytest.raises(ValueError, match="wider than 4096 bits"):
cond.run(State({"x": big, "y": big}))


def test_safe_expr_int_literal_width_bounded():
Condition.safe_expr(str(2**4095)) # exactly 4096 bits: allowed
with pytest.raises(ValueError, match="integer literal wider"):
Condition.safe_expr(str(2**4096))


def test_safe_expr_node_count_bounded():
# Wide and shallow.
with pytest.raises(ValueError, match="more than 500 nodes"):
Condition.safe_expr("[" + ", ".join(["1"] * 3000) + "]")
# Deep chain that every supported parser accepts: rejected by the node count
# before any recursive walk of the tree is attempted.
with pytest.raises(ValueError, match="more than 500 nodes"):
Condition.safe_expr("1" + " + 1" * 600)
# A moderately long expression is still fine.
cond = Condition.safe_expr(" + ".join(["x"] * 100) + " == 100")
assert cond.run(State({"x": 1})) == {Condition.KEY: True}


def test_safe_expr_very_deep_chain_rejected():
# Depending on the Python version the parser either builds this tree (and the
# node count rejects it) or gives up on nesting depth first (and we surface
# that as a ValueError too). Either way: ValueError, never RecursionError.
with pytest.raises(ValueError):
Condition.safe_expr("1" + " + 1" * 3000)


def test_safe_expr_source_length_bounded():
# Rejected on raw length, before the parser ever sees it.
with pytest.raises(ValueError, match="exceeds 10000 characters"):
Condition.safe_expr("x" + " " * 20_000)
# Just under the cap still parses (and is then judged on its own merits).
cond = Condition.safe_expr("x" + " " * 9_000 + "== 1")
assert cond.run(State({"x": 1})) == {Condition.KEY: True}


def test_safe_expr_literal_bounds_check_defers_other_errors():
# Literal-only arithmetic is bounds-checked at call time, but any other error it
# would raise (here, division by zero) is still deferred to run(), as before.
cond = Condition.safe_expr("1 / 0")
with pytest.raises(ZeroDivisionError):
cond.run(State({}))


def _eval_value(cond: Condition, state: dict):
"""Helper: get the raw (pre-bool-coerce) interpreter result for a safe_expr.

Expand Down
Loading