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
10 changes: 5 additions & 5 deletions crates/vm/src/stdlib/_ast.rs
Original file line number Diff line number Diff line change
Expand Up @@ -122,7 +122,7 @@ fn get_required_node_field<T: Node>(
if vm.is_none(&value) {
return Err(vm.new_value_error(format!("field '{field}' is required for {typ}")));
}
let recursion_context = format!(" while traversing '{typ}' node");
let recursion_context = format!("while traversing '{typ}' node");
vm.with_recursion(&recursion_context, || {
Node::ast_from_object(vm, source_file, value)
})
Expand Down Expand Up @@ -177,7 +177,7 @@ fn convert_node_list_field<T: Node>(
) -> PyResult<Vec<T>> {
let len = list.borrow_vec().len();
let mut result = Vec::with_capacity(len);
let recursion_context = format!(" while traversing '{typ}' node");
let recursion_context = format!("while traversing '{typ}' node");
for i in 0..len {
let item = {
let items = list.borrow_vec();
Expand Down Expand Up @@ -470,7 +470,7 @@ fn scan_ast_source_extent(
for i in 0..len {
let field = fields.get_item(i as isize, vm)?;
if let Some(value) = get_attribute_from_field(vm, object, field)? {
vm.with_recursion(" while scanning AST node", || {
vm.with_recursion("while scanning AST node", || {
scan_ast_source_extent(vm, &value, extent)
})?;
}
Expand All @@ -479,13 +479,13 @@ fn scan_ast_source_extent(
} else if let Some(list) = object.downcast_ref::<PyList>() {
let items = list.borrow_vec().to_vec();
for item in items {
vm.with_recursion(" while scanning AST node", || {
vm.with_recursion("while scanning AST node", || {
scan_ast_source_extent(vm, &item, extent)
})?;
}
} else if let Some(tuple) = object.downcast_ref::<PyTuple>() {
for item in tuple.as_slice() {
vm.with_recursion(" while scanning AST node", || {
vm.with_recursion("while scanning AST node", || {
scan_ast_source_extent(vm, item, extent)
})?;
}
Expand Down
8 changes: 4 additions & 4 deletions crates/vm/src/stdlib/_ast/constant.rs
Original file line number Diff line number Diff line change
Expand Up @@ -355,7 +355,7 @@ fn first_invalid_constant_type(vm: &VirtualMachine, value_object: &PyObject) ->
let cls = value_object.class();
let class_name = cls.name().to_owned();
if cls.is(vm.ctx.types.tuple_type) {
vm.with_recursion(" during compilation", || {
vm.with_recursion("during compilation", || {
let tuple = value_object
.to_owned()
.downcast::<PyTuple>()
Expand All @@ -374,7 +374,7 @@ fn first_invalid_constant_type(vm: &VirtualMachine, value_object: &PyObject) ->
Ok(class_name)
})
} else if cls.is(vm.ctx.types.frozenset_type) {
vm.with_recursion(" during compilation", || {
vm.with_recursion("during compilation", || {
let set = value_object.to_owned().downcast::<PyFrozenSet>().unwrap();
for item in set.elements() {
if let Some(invalid_type) = first_invalid_constant_type_opt(vm, &item)? {
Expand Down Expand Up @@ -575,7 +575,7 @@ impl Node for ConstantLiteral {
.into_iter()
.map(|object| {
let object = object.clone();
vm.with_recursion(" during compilation", || {
vm.with_recursion("during compilation", || {
Node::ast_from_object(vm, source_file, object)
})
})
Expand All @@ -587,7 +587,7 @@ impl Node for ConstantLiteral {
.elements()
.into_iter()
.map(|object| {
vm.with_recursion(" during compilation", || {
vm.with_recursion("during compilation", || {
Node::ast_from_object(vm, source_file, object)
})
})
Expand Down
4 changes: 2 additions & 2 deletions crates/vm/src/stdlib/_ast/statement.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1321,7 +1321,7 @@ fn except_handler_list_from_field(
let len = list.borrow_vec().len();
let mut result = Vec::with_capacity(len);
let mut runtime_values = Vec::with_capacity(len);
let recursion_context = format!(" while traversing '{typ}' node");
let recursion_context = format!("while traversing '{typ}' node");
for i in 0..len {
let item = {
let items = list.borrow_vec();
Expand Down Expand Up @@ -1544,7 +1544,7 @@ fn import_from_level_from_field(
let Some(value) = get_node_field_opt(vm, object, "level")? else {
return Ok((0, None));
};
let level = vm.with_recursion(" while traversing 'ImportFrom' node", || {
let level = vm.with_recursion("while traversing 'ImportFrom' node", || {
node_object_to_i32(vm, &value)
})?;
if level < 0 {
Expand Down
23 changes: 16 additions & 7 deletions crates/vm/src/vm/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2603,9 +2603,7 @@ impl VirtualMachine {
let counted_too_deep = false;

if counted_too_deep || self.check_c_stack_overflow() {
return Err(
self.new_recursion_error(format!("maximum recursion depth exceeded {_where}"))
);
return Err(self.new_recursion_depth_error(_where));
}

#[cfg(any(miri, target_env = "musl"))]
Expand All @@ -2632,7 +2630,7 @@ impl VirtualMachine {
// code -- an `__add__` chain, a sort key that sorts -- takes more than
// the margin in that many.
if self.check_c_stack_overflow() {
return Err(self.new_recursion_error(String::new()));
return Err(self.new_recursion_depth_error(""));
}

self.recursion_depth.update(|d| d + 1);
Expand Down Expand Up @@ -2739,7 +2737,7 @@ impl VirtualMachine {
iframe: &mut crate::frame::InterpreterFrame,
) -> PyResult<IframeEntryState> {
if self.check_c_stack_overflow() {
return Err(self.new_recursion_error(String::new()));
return Err(self.new_recursion_depth_error(""));
}

self.recursion_depth.update(|d| d + 1);
Expand Down Expand Up @@ -2880,7 +2878,7 @@ impl VirtualMachine {
) -> PyResult<GenFrameLink> {
self.check_recursive_call("")?;
if self.check_c_stack_overflow() {
return Err(self.new_recursion_error(String::new()));
return Err(self.new_recursion_depth_error(""));
}
self.recursion_depth.update(|d| d + 1);

Expand Down Expand Up @@ -3059,10 +3057,21 @@ impl VirtualMachine {
self.tracing_depth.get() != 0
}

#[cold]
pub fn new_recursion_depth_error(&self, _where: &str) -> PyBaseExceptionRef {
let _where = _where.trim();
let msg = if _where.is_empty() {
"maximum recursion depth exceeded".to_string()
} else {
format!("maximum recursion depth exceeded {_where}")
};
self.new_recursion_error(msg)
}

// To be called right before raising the recursion depth.
fn check_recursive_call(&self, _where: &str) -> PyResult<()> {
if self.recursion_depth.get() >= self.recursion_limit.get() {
Err(self.new_recursion_error(format!("maximum recursion depth exceeded {_where}")))
Err(self.new_recursion_depth_error(_where))
} else {
Ok(())
}
Expand Down
2 changes: 1 addition & 1 deletion extra_tests/snippets/builtin_hash.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,7 @@ def __hash__(self):
# Deep enough to reach the native stack guard; CPython, which also runs
# this snippet, dies on the same value.
deep_tuple = ()
for _ in range(100_000):
for _ in range(500_000):
deep_tuple = (deep_tuple,)
with assert_raises(RecursionError):
hash(deep_tuple)
Expand Down
30 changes: 30 additions & 0 deletions extra_tests/snippets/recursion.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,3 +44,33 @@ class Descr:
pass
else:
raise AssertionError("descr.x should not resolve")


def recurse():
recurse()


try:
recurse()
except RecursionError as e:
assert str(e) == "maximum recursion depth exceeded", f"unexpected message: {e!r}"
assert e.args == ("maximum recursion depth exceeded",), (
f"unexpected args: {e.args!r}"
)

import sys

prev_limit = sys.getrecursionlimit()
try:
sys.setrecursionlimit(50)
try:
recurse()
except RecursionError as e:
assert str(e) == "maximum recursion depth exceeded", (
f"unexpected message: {e!r}"
)
assert e.args == ("maximum recursion depth exceeded",), (
f"unexpected args: {e.args!r}"
)
finally:
sys.setrecursionlimit(prev_limit)
4 changes: 2 additions & 2 deletions extra_tests/snippets/stdlib_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,14 +54,14 @@ def _run_missing_type_params_regression():
list[self_referential]

nested = [0]
for _ in range(100_000):
for _ in range(500_000):
nested = [nested]
with assert_raises(RecursionError):
list[nested]

# hashing an alias walks the same shape
deep_alias = int
for _ in range(100_000):
for _ in range(500_000):
deep_alias = list[deep_alias]
with assert_raises(RecursionError):
hash(deep_alias)
Loading