use std::{collections::BTreeMap, sync::Arc};
use rustpython_parser::ast::{self};
use super::params::evaluate_param_defaults;
use crate::{
error::{EvalError, EvalResult, InterpreterError},
eval::eval_expr,
state::InterpreterState,
tools::Tools,
value::{FunctionDef, FunctionParams, LambdaDef, Param, Value},
};
#[inline(never)]
fn resolve_closure_cells(
state: &mut InterpreterState,
node: &ast::StmtFunctionDef,
closure: &BTreeMap<String, Value>,
) -> (Vec<String>, Option<u64>, Vec<(String, u64)>) {
let nonlocal_names = collect_nonlocal_names(&node.body);
let nonlocal_cell_id = if nonlocal_names.is_empty() {
None
} else {
let existing = state
.frame_cell_owners
.last()
.and_then(|owners| nonlocal_names.iter().find_map(|n| owners.get(n).copied()));
let cell_id = existing.unwrap_or_else(|| {
let id = state.next_nonlocal_cell_id;
state.next_nonlocal_cell_id = state.next_nonlocal_cell_id.wrapping_add(1);
id
});
let already: std::collections::HashSet<String> = state
.nonlocal_cells
.get(&cell_id)
.map(|c| c.keys().cloned().collect())
.unwrap_or_default();
let seeds: Vec<(String, Value)> = nonlocal_names
.iter()
.filter(|n| !already.contains(n.as_str()))
.filter_map(|n| state.variables.get(n).map(|v| (n.clone(), v.clone())))
.collect();
let cell = state.nonlocal_cells.entry(cell_id).or_default();
for (n, v) in seeds {
cell.insert(n, v);
}
if let Some(owners) = state.frame_cell_owners.last_mut() {
for n in &nonlocal_names {
owners.insert(n.clone(), cell_id);
}
}
Some(cell_id)
};
let cell_refreshes: Vec<(String, u64)> =
state.frame_cell_owners.last().map_or_else(Vec::new, |owners| {
closure
.keys()
.filter(|name| !nonlocal_names.contains(name))
.filter_map(|name| owners.get(name).map(|&id| (name.clone(), id)))
.collect()
});
(nonlocal_names, nonlocal_cell_id, cell_refreshes)
}
pub async fn eval_function_def(
state: &mut InterpreterState,
node: &ast::StmtFunctionDef,
tools: &Tools,
) -> EvalResult {
eval_function_def_with(state, node, false, tools).await
}
pub async fn eval_function_def_with(
state: &mut InterpreterState,
node: &ast::StmtFunctionDef,
is_async: bool,
tools: &Tools,
) -> EvalResult {
let name = node.name.as_str();
crate::security::validator::validate_name(
crate::security::validator::NameContext::FunctionDefinition,
name,
)?;
if tools.contains_key(name) {
return Err(InterpreterError::Security(format!(
"'{name}' is not allowed to be overridden"
))
.into());
}
let mut params = build_function_params(&node.args)?;
evaluate_param_defaults(state, &mut params, tools).await?;
let body_key = format!("{name}#{}", state.next_cursor_id);
state.next_cursor_id = state.next_cursor_id.wrapping_add(1);
state.function_bodies.insert(body_key.clone(), Arc::new(node.body.clone()));
let source = extract_function_source(&state.current_source, node);
let closure: BTreeMap<String, Value> =
state.variables.iter().map(|(k, v)| (k.clone(), v.clone())).collect();
let mut bound = param_names(¶ms);
bound.extend(collect_assigned_names(&node.body).0);
let free = collect_free_names(&bound, &node.body);
ensure_capture_cells(state, &free);
let self_recursion_cell = if state.call_depth > 0
&& free.iter().any(|n| n == name)
&& !collect_nonlocal_names(&node.body).iter().any(|n| n == name)
{
let cell_id = state
.frame_cell_owners
.last()
.and_then(|owners| owners.get(name).copied())
.unwrap_or_else(|| {
let id = state.next_nonlocal_cell_id;
state.next_nonlocal_cell_id = state.next_nonlocal_cell_id.wrapping_add(1);
id
});
if let Some(owners) = state.frame_cell_owners.last_mut() {
owners.insert(name.to_string(), cell_id);
}
state
.nonlocal_cells
.entry(cell_id)
.or_default()
.entry(name.to_string())
.or_insert(Value::None);
Some((name.to_string(), cell_id))
} else {
None
};
let (nonlocal_names, nonlocal_cell_id, mut cell_refreshes) =
resolve_closure_cells(state, node, &closure);
if let Some(sc) = self_recursion_cell {
if !cell_refreshes.iter().any(|(n, _)| *n == sc.0) {
cell_refreshes.push(sc);
}
}
let (mut assigned_names, global_names) = collect_assigned_names(&node.body);
assigned_names.retain(|n| !nonlocal_names.contains(n) && !global_names.contains(n));
let is_module_level = state.call_depth == 0;
let is_generator = contains_yield_stmts(&node.body);
let docstring = extract_docstring(&node.body);
let qualname = state.qualname_for(name);
let mut ann_specs: Vec<(&str, Option<&ast::Expr>)> = Vec::new();
for awd in node.args.posonlyargs.iter().chain(node.args.args.iter()) {
ann_specs.push((awd.def.arg.as_str(), awd.def.annotation.as_deref()));
}
if let Some(a) = node.args.vararg.as_deref() {
ann_specs.push((a.arg.as_str(), a.annotation.as_deref()));
}
for awd in &node.args.kwonlyargs {
ann_specs.push((awd.def.arg.as_str(), awd.def.annotation.as_deref()));
}
if let Some(a) = node.args.kwarg.as_deref() {
ann_specs.push((a.arg.as_str(), a.annotation.as_deref()));
}
let mut annotations: Vec<(String, Value)> = Vec::new();
for (pname, ann) in ann_specs {
if let Some(expr) = ann {
if let Ok(v) = eval_expr(state, expr, tools).await {
annotations.push((pname.to_string(), v));
}
}
}
if let Some(ret) = node.returns.as_deref() {
if let Ok(v) = eval_expr(state, ret, tools).await {
annotations.push(("return".to_string(), v));
}
}
let mut func = Value::Function(std::sync::Arc::new(FunctionDef {
name: name.to_string(),
body_key,
wraps_name: None,
params,
closure,
source,
nonlocal_names,
is_generator,
nonlocal_cell_id,
assigned_names,
global_names,
is_module_level,
docstring,
cell_refreshes,
qualname,
annotations,
is_async,
}));
for decorator in node.decorator_list.iter().rev() {
let dec_val = eval_expr(state, decorator, tools).await?;
func = crate::eval::classes::apply_decorator(state, &dec_val, func, tools).await?;
}
state.set_variable(name, func).map_err(EvalError::Interpreter)?;
Ok(Value::None)
}
pub(super) fn apply_nonlocal_cell(
state: &mut InterpreterState,
func_def: &FunctionDef,
local_scope: &rustc_hash::FxHashMap<String, Value>,
) -> Result<(), EvalError> {
let Some(cell_id) = func_def.nonlocal_cell_id else { return Ok(()) };
let Some(cell) = state.nonlocal_cells.get(&cell_id).cloned() else { return Ok(()) };
for (name, value) in cell {
if !local_scope.contains_key(&name) {
state.set_variable(&name, value).map_err(EvalError::Interpreter)?;
}
}
Ok(())
}
#[inline(never)]
pub(super) fn apply_function_scope(
state: &mut InterpreterState,
func_def: &FunctionDef,
local_scope: &rustc_hash::FxHashMap<String, Value>,
) -> Result<(), EvalError> {
for (name, value) in &func_def.closure {
if local_scope.contains_key(name) || func_def.global_names.contains(name) {
continue;
}
if func_def.cell_refreshes.iter().any(|(n, _)| n == name) {
continue;
}
if func_def.is_module_level && state.variables.contains_key(name) {
continue;
}
if let Some(live) = state.variables.get(name) {
if live == value {
continue;
}
}
state.set_variable(name, value.clone()).map_err(EvalError::Interpreter)?;
}
apply_nonlocal_cell(state, func_def, local_scope)?;
apply_cell_refreshes(state, &func_def.cell_refreshes, local_scope)?;
for (name, value) in local_scope {
state.set_variable(name, value.clone()).map_err(EvalError::Interpreter)?;
}
Ok(())
}
fn apply_cell_refreshes(
state: &mut InterpreterState,
refreshes: &[(String, u64)],
local_scope: &rustc_hash::FxHashMap<String, Value>,
) -> Result<(), EvalError> {
for (name, cell_id) in refreshes {
if local_scope.contains_key(name) {
continue;
}
if let Some(value) = state.nonlocal_cells.get(cell_id).and_then(|c| c.get(name)).cloned() {
state.set_variable(name, value).map_err(EvalError::Interpreter)?;
}
}
Ok(())
}
pub(super) fn apply_lambda_scope(
state: &mut InterpreterState,
lambda_def: &LambdaDef,
local_scope: &rustc_hash::FxHashMap<String, Value>,
) -> Result<(), EvalError> {
for (name, value) in &lambda_def.closure {
if local_scope.contains_key(name) {
continue;
}
if lambda_def.cell_refreshes.iter().any(|(n, _)| n == name) {
continue;
}
if lambda_def.is_module_level && state.variables.contains_key(name) {
continue;
}
if let Some(live) = state.variables.get(name) {
if live == value {
continue;
}
}
state.set_variable(name, value.clone()).map_err(EvalError::Interpreter)?;
}
apply_cell_refreshes(state, &lambda_def.cell_refreshes, local_scope)?;
for (name, value) in local_scope {
state.set_variable(name, value.clone()).map_err(EvalError::Interpreter)?;
}
Ok(())
}
pub(super) fn writeback_nonlocal_cell(state: &mut InterpreterState, func_def: &FunctionDef) {
let Some(cell_id) = func_def.nonlocal_cell_id else { return };
let writeback: Vec<(String, Value)> = func_def
.nonlocal_names
.iter()
.filter_map(|n| state.variables.get(n).map(|v| (n.clone(), v.clone())))
.collect();
if let Some(cell) = state.nonlocal_cells.get_mut(&cell_id) {
for (n, v) in writeback {
cell.insert(n, v);
}
}
}
fn collect_nonlocal_names(body: &[ast::Stmt]) -> Vec<String> {
let mut names = Vec::new();
collect_nonlocal_names_inner(body, &mut names);
names
}
fn collect_nonlocal_names_inner(body: &[ast::Stmt], out: &mut Vec<String>) {
for stmt in body {
match stmt {
ast::Stmt::Nonlocal(node) => {
for ident in &node.names {
let n = ident.as_str().to_string();
if !out.contains(&n) {
out.push(n);
}
}
}
ast::Stmt::If(node) => {
collect_nonlocal_names_inner(&node.body, out);
collect_nonlocal_names_inner(&node.orelse, out);
}
ast::Stmt::For(node) => {
collect_nonlocal_names_inner(&node.body, out);
collect_nonlocal_names_inner(&node.orelse, out);
}
ast::Stmt::While(node) => {
collect_nonlocal_names_inner(&node.body, out);
collect_nonlocal_names_inner(&node.orelse, out);
}
ast::Stmt::With(node) => {
collect_nonlocal_names_inner(&node.body, out);
}
ast::Stmt::Try(node) => {
collect_nonlocal_names_inner(&node.body, out);
collect_nonlocal_names_inner(&node.orelse, out);
collect_nonlocal_names_inner(&node.finalbody, out);
for handler in &node.handlers {
let ast::ExceptHandler::ExceptHandler(h) = handler;
collect_nonlocal_names_inner(&h.body, out);
}
}
_ => {}
}
}
}
fn param_names(params: &FunctionParams) -> Vec<String> {
let mut names: Vec<String> = params.args.iter().map(|p| p.name.clone()).collect();
names.extend(params.kwonlyargs.iter().map(|p| p.name.clone()));
if let Some(v) = ¶ms.vararg {
names.push(v.clone());
}
if let Some(k) = ¶ms.kwarg {
names.push(k.clone());
}
names
}
#[must_use]
pub(crate) fn collect_free_names(bound: &[String], body: &[ast::Stmt]) -> Vec<String> {
let mut reads = Vec::new();
for stmt in body {
collect_reads_stmt(stmt, &mut reads);
}
reads.retain(|n| !bound.contains(n));
reads
}
fn collect_reads_stmt(stmt: &ast::Stmt, out: &mut Vec<String>) {
use ast::Stmt;
let body_of = |b: &[Stmt], out: &mut Vec<String>| {
for s in b {
collect_reads_stmt(s, out);
}
};
match stmt {
Stmt::Expr(n) => collect_reads_expr(&n.value, out),
Stmt::Return(n) => {
if let Some(v) = &n.value {
collect_reads_expr(v, out);
}
}
Stmt::Assign(n) => collect_reads_expr(&n.value, out),
Stmt::AugAssign(n) => {
collect_reads_expr(&n.target, out);
collect_reads_expr(&n.value, out);
}
Stmt::AnnAssign(n) => {
if let Some(v) = &n.value {
collect_reads_expr(v, out);
}
}
Stmt::For(n) => {
collect_reads_expr(&n.iter, out);
body_of(&n.body, out);
body_of(&n.orelse, out);
}
Stmt::While(n) => {
collect_reads_expr(&n.test, out);
body_of(&n.body, out);
body_of(&n.orelse, out);
}
Stmt::If(n) => {
collect_reads_expr(&n.test, out);
body_of(&n.body, out);
body_of(&n.orelse, out);
}
Stmt::With(n) => {
for item in &n.items {
collect_reads_expr(&item.context_expr, out);
}
body_of(&n.body, out);
}
Stmt::Try(n) => {
body_of(&n.body, out);
body_of(&n.orelse, out);
body_of(&n.finalbody, out);
for h in &n.handlers {
let ast::ExceptHandler::ExceptHandler(eh) = h;
body_of(&eh.body, out);
}
}
Stmt::FunctionDef(n) => body_of(&n.body, out),
Stmt::ClassDef(n) => body_of(&n.body, out),
Stmt::Delete(n) => {
for t in &n.targets {
collect_reads_expr(t, out);
}
}
Stmt::Raise(n) => {
if let Some(e) = &n.exc {
collect_reads_expr(e, out);
}
if let Some(c) = &n.cause {
collect_reads_expr(c, out);
}
}
Stmt::Assert(n) => {
collect_reads_expr(&n.test, out);
if let Some(m) = &n.msg {
collect_reads_expr(m, out);
}
}
_ => {}
}
}
#[allow(clippy::too_many_lines)]
fn collect_reads_expr(expr: &ast::Expr, out: &mut Vec<String>) {
use ast::Expr;
match expr {
Expr::Name(n) if matches!(n.ctx, ast::ExprContext::Load) => push_unique(out, n.id.as_str()),
Expr::Name(_) => {}
Expr::BoolOp(n) => {
for v in &n.values {
collect_reads_expr(v, out);
}
}
Expr::BinOp(n) => {
collect_reads_expr(&n.left, out);
collect_reads_expr(&n.right, out);
}
Expr::UnaryOp(n) => collect_reads_expr(&n.operand, out),
Expr::Compare(n) => {
collect_reads_expr(&n.left, out);
for c in &n.comparators {
collect_reads_expr(c, out);
}
}
Expr::IfExp(n) => {
collect_reads_expr(&n.test, out);
collect_reads_expr(&n.body, out);
collect_reads_expr(&n.orelse, out);
}
Expr::Call(n) => {
collect_reads_expr(&n.func, out);
for a in &n.args {
collect_reads_expr(a, out);
}
for k in &n.keywords {
collect_reads_expr(&k.value, out);
}
}
Expr::Attribute(n) => collect_reads_expr(&n.value, out),
Expr::Subscript(n) => {
collect_reads_expr(&n.value, out);
collect_reads_expr(&n.slice, out);
}
Expr::Starred(n) => collect_reads_expr(&n.value, out),
Expr::Slice(n) => {
if let Some(l) = &n.lower {
collect_reads_expr(l, out);
}
if let Some(u) = &n.upper {
collect_reads_expr(u, out);
}
if let Some(s) = &n.step {
collect_reads_expr(s, out);
}
}
Expr::Lambda(n) => collect_reads_expr(&n.body, out),
Expr::Tuple(n) => {
for e in &n.elts {
collect_reads_expr(e, out);
}
}
Expr::List(n) => {
for e in &n.elts {
collect_reads_expr(e, out);
}
}
Expr::Set(n) => {
for e in &n.elts {
collect_reads_expr(e, out);
}
}
Expr::Dict(n) => {
for k in n.keys.iter().flatten() {
collect_reads_expr(k, out);
}
for v in &n.values {
collect_reads_expr(v, out);
}
}
Expr::ListComp(n) => collect_reads_comp(&n.elt, None, &n.generators, out),
Expr::SetComp(n) => collect_reads_comp(&n.elt, None, &n.generators, out),
Expr::GeneratorExp(n) => collect_reads_comp(&n.elt, None, &n.generators, out),
Expr::DictComp(n) => collect_reads_comp(&n.key, Some(&n.value), &n.generators, out),
Expr::JoinedStr(n) => {
for v in &n.values {
collect_reads_expr(v, out);
}
}
Expr::FormattedValue(n) => {
collect_reads_expr(&n.value, out);
if let Some(spec) = &n.format_spec {
collect_reads_expr(spec, out);
}
}
Expr::NamedExpr(n) => collect_reads_expr(&n.value, out),
Expr::Await(n) => collect_reads_expr(&n.value, out),
Expr::Yield(n) => {
if let Some(v) = &n.value {
collect_reads_expr(v, out);
}
}
Expr::YieldFrom(n) => collect_reads_expr(&n.value, out),
_ => {}
}
}
fn collect_reads_comp(
elt: &ast::Expr,
value: Option<&ast::Expr>,
generators: &[ast::Comprehension],
out: &mut Vec<String>,
) {
collect_reads_expr(elt, out);
if let Some(v) = value {
collect_reads_expr(v, out);
}
for comp in generators {
collect_reads_expr(&comp.iter, out);
for cond in &comp.ifs {
collect_reads_expr(cond, out);
}
}
}
fn ensure_capture_cells(state: &mut InterpreterState, free_names: &[String]) -> Vec<(String, u64)> {
if state.frame_cell_owners.is_empty() {
return Vec::new();
}
let mut refreshes = Vec::new();
for name in free_names {
let Some(current) = state.variables.get(name).cloned() else {
continue;
};
let existing = state.frame_cell_owners.last().and_then(|owners| owners.get(name).copied());
let cell_id = existing.unwrap_or_else(|| {
let id = state.next_nonlocal_cell_id;
state.next_nonlocal_cell_id = state.next_nonlocal_cell_id.wrapping_add(1);
id
});
if existing.is_none() {
state.nonlocal_cells.entry(cell_id).or_default().insert(name.clone(), current);
if let Some(owners) = state.frame_cell_owners.last_mut() {
owners.insert(name.clone(), cell_id);
}
}
refreshes.push((name.clone(), cell_id));
}
refreshes
}
pub(crate) fn collect_assigned_names(body: &[ast::Stmt]) -> (Vec<String>, Vec<String>) {
let mut assigned = Vec::new();
let mut globals = Vec::new();
collect_assigned_names_inner(body, &mut assigned, &mut globals);
(assigned, globals)
}
fn push_unique(out: &mut Vec<String>, name: &str) {
let s = name.to_string();
if !out.contains(&s) {
out.push(s);
}
}
fn collect_target_names(target: &ast::Expr, out: &mut Vec<String>) {
match target {
ast::Expr::Name(n) => push_unique(out, n.id.as_str()),
ast::Expr::Tuple(t) => {
for elt in &t.elts {
collect_target_names(elt, out);
}
}
ast::Expr::List(l) => {
for elt in &l.elts {
collect_target_names(elt, out);
}
}
ast::Expr::Starred(s) => collect_target_names(&s.value, out),
_ => {}
}
}
fn collect_walrus_targets(expr: &ast::Expr, out: &mut Vec<String>) {
match expr {
ast::Expr::NamedExpr(node) => {
collect_target_names(&node.target, out);
collect_walrus_targets(&node.value, out);
}
ast::Expr::BoolOp(node) => {
for v in &node.values {
collect_walrus_targets(v, out);
}
}
ast::Expr::BinOp(node) => {
collect_walrus_targets(&node.left, out);
collect_walrus_targets(&node.right, out);
}
ast::Expr::UnaryOp(node) => collect_walrus_targets(&node.operand, out),
ast::Expr::IfExp(node) => {
collect_walrus_targets(&node.test, out);
collect_walrus_targets(&node.body, out);
collect_walrus_targets(&node.orelse, out);
}
ast::Expr::Compare(node) => {
collect_walrus_targets(&node.left, out);
for c in &node.comparators {
collect_walrus_targets(c, out);
}
}
ast::Expr::Call(node) => {
collect_walrus_targets(&node.func, out);
for a in &node.args {
collect_walrus_targets(a, out);
}
for kw in &node.keywords {
collect_walrus_targets(&kw.value, out);
}
}
ast::Expr::Attribute(node) => collect_walrus_targets(&node.value, out),
ast::Expr::Subscript(node) => {
collect_walrus_targets(&node.value, out);
collect_walrus_targets(&node.slice, out);
}
ast::Expr::Starred(node) => collect_walrus_targets(&node.value, out),
ast::Expr::Tuple(node) => {
for e in &node.elts {
collect_walrus_targets(e, out);
}
}
ast::Expr::List(node) => {
for e in &node.elts {
collect_walrus_targets(e, out);
}
}
ast::Expr::Set(node) => {
for e in &node.elts {
collect_walrus_targets(e, out);
}
}
ast::Expr::Dict(node) => {
for k in node.keys.iter().flatten() {
collect_walrus_targets(k, out);
}
for v in &node.values {
collect_walrus_targets(v, out);
}
}
ast::Expr::FormattedValue(node) => {
collect_walrus_targets(&node.value, out);
if let Some(fmt) = &node.format_spec {
collect_walrus_targets(fmt, out);
}
}
ast::Expr::JoinedStr(node) => {
for v in &node.values {
collect_walrus_targets(v, out);
}
}
ast::Expr::Slice(node) => {
if let Some(l) = &node.lower {
collect_walrus_targets(l, out);
}
if let Some(u) = &node.upper {
collect_walrus_targets(u, out);
}
if let Some(s) = &node.step {
collect_walrus_targets(s, out);
}
}
ast::Expr::Yield(node) => {
if let Some(v) = &node.value {
collect_walrus_targets(v, out);
}
}
ast::Expr::YieldFrom(node) => collect_walrus_targets(&node.value, out),
ast::Expr::Await(node) => collect_walrus_targets(&node.value, out),
ast::Expr::ListComp(node) => {
collect_walrus_targets(&node.elt, out);
for g in &node.generators {
collect_walrus_targets(&g.iter, out);
for c in &g.ifs {
collect_walrus_targets(c, out);
}
}
}
ast::Expr::SetComp(node) => {
collect_walrus_targets(&node.elt, out);
for g in &node.generators {
collect_walrus_targets(&g.iter, out);
for c in &g.ifs {
collect_walrus_targets(c, out);
}
}
}
ast::Expr::DictComp(node) => {
collect_walrus_targets(&node.key, out);
collect_walrus_targets(&node.value, out);
for g in &node.generators {
collect_walrus_targets(&g.iter, out);
for c in &g.ifs {
collect_walrus_targets(c, out);
}
}
}
ast::Expr::GeneratorExp(node) => {
collect_walrus_targets(&node.elt, out);
for g in &node.generators {
collect_walrus_targets(&g.iter, out);
for c in &g.ifs {
collect_walrus_targets(c, out);
}
}
}
_ => {}
}
}
fn collect_assigned_names_inner(
body: &[ast::Stmt],
assigned: &mut Vec<String>,
globals: &mut Vec<String>,
) {
for stmt in body {
match stmt {
ast::Stmt::Global(node) => {
for ident in &node.names {
push_unique(globals, ident.as_str());
}
}
ast::Stmt::Assign(node) => {
for target in &node.targets {
collect_target_names(target, assigned);
}
collect_walrus_targets(&node.value, assigned);
}
ast::Stmt::AugAssign(node) => {
collect_target_names(&node.target, assigned);
collect_walrus_targets(&node.value, assigned);
}
ast::Stmt::AnnAssign(node) => {
collect_target_names(&node.target, assigned);
if let Some(v) = &node.value {
collect_walrus_targets(v, assigned);
}
}
ast::Stmt::Delete(node) => {
for target in &node.targets {
collect_target_names(target, assigned);
}
}
ast::Stmt::Expr(node) => collect_walrus_targets(&node.value, assigned),
ast::Stmt::Return(node) => {
if let Some(v) = &node.value {
collect_walrus_targets(v, assigned);
}
}
ast::Stmt::Raise(node) => {
if let Some(exc) = &node.exc {
collect_walrus_targets(exc, assigned);
}
if let Some(cause) = &node.cause {
collect_walrus_targets(cause, assigned);
}
}
ast::Stmt::Assert(node) => {
collect_walrus_targets(&node.test, assigned);
if let Some(msg) = &node.msg {
collect_walrus_targets(msg, assigned);
}
}
ast::Stmt::For(node) => {
collect_target_names(&node.target, assigned);
collect_walrus_targets(&node.iter, assigned);
collect_assigned_names_inner(&node.body, assigned, globals);
collect_assigned_names_inner(&node.orelse, assigned, globals);
}
ast::Stmt::AsyncFor(node) => {
collect_target_names(&node.target, assigned);
collect_walrus_targets(&node.iter, assigned);
collect_assigned_names_inner(&node.body, assigned, globals);
collect_assigned_names_inner(&node.orelse, assigned, globals);
}
ast::Stmt::While(node) => {
collect_walrus_targets(&node.test, assigned);
collect_assigned_names_inner(&node.body, assigned, globals);
collect_assigned_names_inner(&node.orelse, assigned, globals);
}
ast::Stmt::If(node) => {
collect_walrus_targets(&node.test, assigned);
collect_assigned_names_inner(&node.body, assigned, globals);
collect_assigned_names_inner(&node.orelse, assigned, globals);
}
ast::Stmt::With(node) => {
for item in &node.items {
collect_walrus_targets(&item.context_expr, assigned);
if let Some(target) = &item.optional_vars {
collect_target_names(target, assigned);
}
}
collect_assigned_names_inner(&node.body, assigned, globals);
}
ast::Stmt::AsyncWith(node) => {
for item in &node.items {
collect_walrus_targets(&item.context_expr, assigned);
if let Some(target) = &item.optional_vars {
collect_target_names(target, assigned);
}
}
collect_assigned_names_inner(&node.body, assigned, globals);
}
ast::Stmt::Try(node) => {
collect_assigned_names_inner(&node.body, assigned, globals);
collect_assigned_names_inner(&node.orelse, assigned, globals);
collect_assigned_names_inner(&node.finalbody, assigned, globals);
for handler in &node.handlers {
let ast::ExceptHandler::ExceptHandler(h) = handler;
if let Some(name) = &h.name {
push_unique(assigned, name.as_str());
}
if let Some(t) = &h.type_ {
collect_walrus_targets(t, assigned);
}
collect_assigned_names_inner(&h.body, assigned, globals);
}
}
ast::Stmt::Import(node) => {
for alias in &node.names {
let name = alias.asname.as_ref().map_or_else(
|| {
alias
.name
.as_str()
.split('.')
.next()
.unwrap_or(alias.name.as_str())
.to_string()
},
|asname| asname.as_str().to_string(),
);
push_unique(assigned, &name);
}
}
ast::Stmt::ImportFrom(node) => {
for alias in &node.names {
if alias.name.as_str() == "*" {
continue;
}
let name = alias.asname.as_ref().map_or_else(
|| alias.name.as_str().to_string(),
|a| a.as_str().to_string(),
);
push_unique(assigned, &name);
}
}
ast::Stmt::FunctionDef(node) => {
push_unique(assigned, node.name.as_str());
}
ast::Stmt::AsyncFunctionDef(node) => {
push_unique(assigned, node.name.as_str());
}
ast::Stmt::ClassDef(node) => {
push_unique(assigned, node.name.as_str());
}
_ => {}
}
}
}
pub(crate) struct VariableCheckpoint {
snapshots: Vec<(String, Option<Value>)>,
}
impl VariableCheckpoint {
pub(crate) fn capture<I, S>(state: &InterpreterState, names: I) -> Self
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let snapshots: Vec<(String, Option<Value>)> = names
.into_iter()
.map(|n| {
let name = n.as_ref();
let prev = state.variables.get(name).cloned();
(name.to_string(), prev)
})
.collect();
Self { snapshots }
}
pub(crate) fn restore(self, state: &mut InterpreterState) {
for (name, prev) in self.snapshots {
match prev {
Some(v) => {
state.variables.insert(name, v);
}
None => {
state.variables.remove(&name);
}
}
}
}
}
pub(crate) fn extract_docstring(body: &[ast::Stmt]) -> Option<String> {
if let Some(ast::Stmt::Expr(e)) = body.first() {
if let ast::Expr::Constant(ast::ExprConstant { value: ast::Constant::Str(s), .. }) =
e.value.as_ref()
{
return Some(s.to_string());
}
}
None
}
pub(crate) fn extract_function_source(source: &str, node: &ast::StmtFunctionDef) -> String {
use rustpython_parser::text_size::TextRange;
let range: TextRange = node.range;
let start = range.start().to_usize();
let end = range.end().to_usize();
if start < source.len() && end <= source.len() && start < end {
source[start..end].to_string()
} else {
format!("def {}(): pass", node.name)
}
}
pub fn build_function_params(args: &ast::Arguments) -> Result<FunctionParams, EvalError> {
let positional: Vec<Param> = args
.posonlyargs
.iter()
.chain(args.args.iter())
.map(|awd| Param {
name: awd.def.arg.as_str().to_string(),
annotation: awd.def.annotation.as_deref().and_then(|ann| match ann {
ast::Expr::Name(n) => Some(n.id.as_str().to_string()),
_ => None,
}),
})
.collect();
let all_args_with_default: Vec<&ast::ArgWithDefault> =
args.posonlyargs.iter().chain(args.args.iter()).collect();
let mut defaults: Vec<String> = Vec::new();
for awd in &all_args_with_default {
if let Some(ref default_expr) = awd.default {
defaults.push(unparse_expr(default_expr)?);
}
}
let kwonlyargs: Vec<Param> = args
.kwonlyargs
.iter()
.map(|awd| Param { name: awd.def.arg.as_str().to_string(), annotation: None })
.collect();
let mut kw_defaults: Vec<Option<String>> = Vec::with_capacity(args.kwonlyargs.len());
for awd in &args.kwonlyargs {
kw_defaults.push(match &awd.default {
Some(d) => Some(unparse_expr(d)?),
None => None,
});
}
let vararg = args.vararg.as_ref().map(|a| a.arg.as_str().to_string());
let kwarg = args.kwarg.as_ref().map(|a| a.arg.as_str().to_string());
Ok(FunctionParams {
args: positional,
defaults,
default_values: Vec::new(),
vararg,
kwonlyargs,
kw_defaults,
kw_default_values: Vec::new(),
kwarg,
posonly_count: args.posonlyargs.len(),
})
}
fn unparse_comprehensions(gens: &[ast::Comprehension]) -> Result<String, EvalError> {
let mut parts = Vec::with_capacity(gens.len());
for g in gens {
let mut clause = format!("for {} in {}", unparse_expr(&g.target)?, unparse_expr(&g.iter)?);
for cond in &g.ifs {
clause.push_str(&format!(" if {}", unparse_expr(cond)?));
}
parts.push(clause);
}
Ok(parts.join(" "))
}
fn unparse_format_spec(spec: &ast::Expr) -> Result<String, EvalError> {
let ast::Expr::JoinedStr(js) = spec else {
return unparse_expr(spec);
};
let mut out = String::new();
for value in &js.values {
match value {
ast::Expr::Constant(ast::ExprConstant { value: ast::Constant::Str(s), .. }) => {
out.push_str(s);
}
ast::Expr::FormattedValue(fv) => {
out.push('{');
out.push_str(&unparse_expr(&fv.value)?);
out.push('}');
}
_ => {}
}
}
Ok(out)
}
fn unparse_expr(expr: &ast::Expr) -> Result<String, EvalError> {
let join = |exprs: &[ast::Expr], sep: &str| -> Result<String, EvalError> {
Ok(exprs.iter().map(unparse_expr).collect::<Result<Vec<_>, _>>()?.join(sep))
};
Ok(match expr {
ast::Expr::Constant(c) => match &c.value {
ast::Constant::None => "None".to_string(),
ast::Constant::Bool(true) => "True".to_string(),
ast::Constant::Bool(false) => "False".to_string(),
ast::Constant::Int(i) => format!("{i}"),
ast::Constant::Float(f) => {
if f.fract() == 0.0 && f.is_finite() {
format!("{f:.1}")
} else {
format!("{f}")
}
}
ast::Constant::Str(s) => format!("'{}'", s.replace('\\', "\\\\").replace('\'', "\\'")),
ast::Constant::Bytes(b) => format!("b'{}'", String::from_utf8_lossy(b)),
ast::Constant::Ellipsis => "...".to_string(),
ast::Constant::Tuple(items) => {
let parts: Vec<String> = items
.iter()
.map(|c| {
unparse_expr(&ast::Expr::Constant(ast::ExprConstant {
range: rustpython_parser::text_size::TextRange::default(),
value: c.clone(),
kind: None,
}))
})
.collect::<Result<_, _>>()?;
format!("({})", parts.join(", "))
}
ast::Constant::Complex { real, imag } => format!("complex({real}, {imag})"),
},
ast::Expr::Name(n) => n.id.to_string(),
ast::Expr::JoinedStr(js) => {
let mut out = String::from("f\"");
for value in &js.values {
match value {
ast::Expr::Constant(ast::ExprConstant {
value: ast::Constant::Str(s), ..
}) => {
out.push_str(&s.replace('{', "{{").replace('}', "}}").replace('"', "\\\""));
}
ast::Expr::FormattedValue(fv) => {
out.push('{');
out.push_str(&unparse_expr(&fv.value)?);
match fv.conversion {
ast::ConversionFlag::Str => out.push_str("!s"),
ast::ConversionFlag::Repr => out.push_str("!r"),
ast::ConversionFlag::Ascii => out.push_str("!a"),
ast::ConversionFlag::None => {}
}
if let Some(spec) = &fv.format_spec {
out.push(':');
out.push_str(&unparse_format_spec(spec)?);
}
out.push('}');
}
other => {
return Err(InterpreterError::TypeError(format!(
"unsupported f-string default component (see CONFORMANCE.md#unsupported-language-features): {:?}",
std::mem::discriminant(other)
))
.into());
}
}
}
out.push('"');
out
}
ast::Expr::List(l) => format!("[{}]", join(&l.elts, ", ")?),
ast::Expr::ListComp(c) => {
format!("[{} {}]", unparse_expr(&c.elt)?, unparse_comprehensions(&c.generators)?)
}
ast::Expr::SetComp(c) => {
format!("{{{} {}}}", unparse_expr(&c.elt)?, unparse_comprehensions(&c.generators)?)
}
ast::Expr::GeneratorExp(c) => {
format!("({} {})", unparse_expr(&c.elt)?, unparse_comprehensions(&c.generators)?)
}
ast::Expr::DictComp(c) => format!(
"{{{}: {} {}}}",
unparse_expr(&c.key)?,
unparse_expr(&c.value)?,
unparse_comprehensions(&c.generators)?
),
ast::Expr::Set(s) => {
format!("{{{}}}", join(&s.elts, ", ")?)
}
ast::Expr::Tuple(t) => {
let parts: Vec<String> = t.elts.iter().map(unparse_expr).collect::<Result<_, _>>()?;
if parts.len() == 1 {
format!("({},)", parts[0])
} else {
format!("({})", parts.join(", "))
}
}
ast::Expr::Dict(d) => {
let mut parts = Vec::with_capacity(d.keys.len());
for (k, v) in d.keys.iter().zip(d.values.iter()) {
parts.push(match k {
None => format!("**{}", unparse_expr(v)?),
Some(key) => format!("{}: {}", unparse_expr(key)?, unparse_expr(v)?),
});
}
format!("{{{}}}", parts.join(", "))
}
ast::Expr::UnaryOp(u) => {
let op = match u.op {
ast::UnaryOp::USub => "-",
ast::UnaryOp::UAdd => "+",
ast::UnaryOp::Not => "not ",
ast::UnaryOp::Invert => "~",
};
format!("{op}{}", unparse_expr(&u.operand)?)
}
ast::Expr::BinOp(b) => {
let op = match b.op {
ast::Operator::Add => "+",
ast::Operator::Sub => "-",
ast::Operator::Mult => "*",
ast::Operator::Div => "/",
ast::Operator::FloorDiv => "//",
ast::Operator::Mod => "%",
ast::Operator::Pow => "**",
ast::Operator::LShift => "<<",
ast::Operator::RShift => ">>",
ast::Operator::BitOr => "|",
ast::Operator::BitXor => "^",
ast::Operator::BitAnd => "&",
ast::Operator::MatMult => "@",
};
format!("({} {op} {})", unparse_expr(&b.left)?, unparse_expr(&b.right)?)
}
ast::Expr::BoolOp(b) => {
let op = match b.op {
ast::BoolOp::And => " and ",
ast::BoolOp::Or => " or ",
};
format!("({})", join(&b.values, op)?)
}
ast::Expr::Compare(c) => {
let mut out = format!("({}", unparse_expr(&c.left)?);
for (op, comparator) in c.ops.iter().zip(c.comparators.iter()) {
let op = match op {
ast::CmpOp::Eq => "==",
ast::CmpOp::NotEq => "!=",
ast::CmpOp::Lt => "<",
ast::CmpOp::LtE => "<=",
ast::CmpOp::Gt => ">",
ast::CmpOp::GtE => ">=",
ast::CmpOp::Is => "is",
ast::CmpOp::IsNot => "is not",
ast::CmpOp::In => "in",
ast::CmpOp::NotIn => "not in",
};
out.push_str(&format!(" {op} {}", unparse_expr(comparator)?));
}
out.push(')');
out
}
ast::Expr::IfExp(f) => format!(
"({} if {} else {})",
unparse_expr(&f.body)?,
unparse_expr(&f.test)?,
unparse_expr(&f.orelse)?,
),
ast::Expr::Attribute(a) => format!("{}.{}", unparse_expr(&a.value)?, a.attr),
ast::Expr::Subscript(s) => {
format!("{}[{}]", unparse_expr(&s.value)?, unparse_expr(&s.slice)?)
}
ast::Expr::Slice(s) => {
let part = |o: &Option<Box<ast::Expr>>| -> Result<String, EvalError> {
match o {
Some(e) => unparse_expr(e),
None => Ok(String::new()),
}
};
match &s.step {
Some(step) => {
format!("{}:{}:{}", part(&s.lower)?, part(&s.upper)?, unparse_expr(step)?)
}
None => format!("{}:{}", part(&s.lower)?, part(&s.upper)?),
}
}
ast::Expr::Starred(s) => format!("*{}", unparse_expr(&s.value)?),
ast::Expr::Lambda(l) => {
let params = build_function_params(&l.args)?;
let mut names = params.args.iter().map(|p| p.name.clone()).collect::<Vec<_>>();
if let Some(v) = ¶ms.vararg {
names.push(format!("*{v}"));
}
for kw in ¶ms.kwonlyargs {
names.push(kw.name.clone());
}
if let Some(kw) = ¶ms.kwarg {
names.push(format!("**{kw}"));
}
format!("lambda {}: {}", names.join(", "), unparse_expr(&l.body)?)
}
ast::Expr::Call(c) => {
let func = unparse_expr(&c.func)?;
let mut arg_strs: Vec<String> =
c.args.iter().map(unparse_expr).collect::<Result<_, _>>()?;
for kw in &c.keywords {
match &kw.arg {
Some(name) => arg_strs.push(format!("{name}={}", unparse_expr(&kw.value)?)),
None => arg_strs.push(format!("**{}", unparse_expr(&kw.value)?)),
}
}
format!("{func}({})", arg_strs.join(", "))
}
other => {
return Err(InterpreterError::TypeError(format!(
"unsupported default argument expression (see CONFORMANCE.md#unsupported-language-features): {:?}",
std::mem::discriminant(other)
))
.into());
}
})
}
pub async fn eval_lambda_def(
state: &mut InterpreterState,
node: &ast::ExprLambda,
tools: &Tools,
) -> EvalResult {
let mut params = build_lambda_params(&node.args)?;
evaluate_param_defaults(state, &mut params, tools).await?;
let lambda_id = format!("__lambda_{}", state.lambda_bodies.len());
state.lambda_bodies.insert(lambda_id.clone(), Arc::new((*node.body).clone()));
let source = extract_lambda_source(&state.current_source, node);
let closure: BTreeMap<String, Value> =
state.variables.iter().map(|(k, v)| (k.clone(), v.clone())).collect();
let mut assigned_names = Vec::new();
collect_walrus_targets(&node.body, &mut assigned_names);
let is_module_level = state.call_depth == 0;
let mut bound = param_names(¶ms);
bound.extend(assigned_names.iter().cloned());
let mut free = Vec::new();
collect_reads_expr(&node.body, &mut free);
free.retain(|n| !bound.contains(n));
let cell_refreshes = ensure_capture_cells(state, &free);
let qualname = state.qualname_for("<lambda>");
Ok(Value::Lambda(std::sync::Arc::new(LambdaDef {
params,
lambda_id,
source,
closure,
assigned_names,
is_module_level,
cell_refreshes,
qualname,
})))
}
fn extract_lambda_source(source: &str, node: &ast::ExprLambda) -> String {
use rustpython_parser::text_size::TextRange;
let range: TextRange = node.range;
let start = range.start().to_usize();
let end = range.end().to_usize();
if start < source.len() && end <= source.len() && start < end {
source[start..end].to_string()
} else {
"lambda: None".to_string()
}
}
fn build_lambda_params(args: &ast::Arguments) -> Result<FunctionParams, EvalError> {
build_function_params(args)
}
pub(crate) fn contains_yield_stmts(stmts: &[ast::Stmt]) -> bool {
stmts.iter().any(contains_yield_stmt)
}
fn contains_yield_stmt(stmt: &ast::Stmt) -> bool {
use ast::Stmt;
match stmt {
Stmt::Expr(e) => contains_yield_expr(&e.value),
Stmt::Assign(a) => {
contains_yield_expr(&a.value) || a.targets.iter().any(contains_yield_expr)
}
Stmt::AugAssign(a) => contains_yield_expr(&a.value) || contains_yield_expr(&a.target),
Stmt::AnnAssign(a) => a.value.as_deref().is_some_and(contains_yield_expr),
Stmt::Return(r) => r.value.as_deref().is_some_and(contains_yield_expr),
Stmt::If(node) => {
contains_yield_expr(&node.test)
|| contains_yield_stmts(&node.body)
|| contains_yield_stmts(&node.orelse)
}
Stmt::For(node) => {
contains_yield_expr(&node.iter)
|| contains_yield_stmts(&node.body)
|| contains_yield_stmts(&node.orelse)
}
Stmt::While(node) => {
contains_yield_expr(&node.test)
|| contains_yield_stmts(&node.body)
|| contains_yield_stmts(&node.orelse)
}
Stmt::Try(node) => {
contains_yield_stmts(&node.body)
|| contains_yield_stmts(&node.orelse)
|| contains_yield_stmts(&node.finalbody)
|| node.handlers.iter().any(|h| match h {
ast::ExceptHandler::ExceptHandler(eh) => contains_yield_stmts(&eh.body),
})
}
Stmt::With(node) => contains_yield_stmts(&node.body),
Stmt::Match(node) => node.cases.iter().any(|c| contains_yield_stmts(&c.body)),
Stmt::Raise(node) => {
node.exc.as_deref().is_some_and(contains_yield_expr)
|| node.cause.as_deref().is_some_and(contains_yield_expr)
}
_ => false,
}
}
pub(super) fn contains_yield_expr(expr: &ast::Expr) -> bool {
use ast::Expr;
match expr {
Expr::Yield(_) | Expr::YieldFrom(_) => true,
Expr::BoolOp(node) => node.values.iter().any(contains_yield_expr),
Expr::BinOp(node) => contains_yield_expr(&node.left) || contains_yield_expr(&node.right),
Expr::UnaryOp(node) => contains_yield_expr(&node.operand),
Expr::IfExp(node) => {
contains_yield_expr(&node.test)
|| contains_yield_expr(&node.body)
|| contains_yield_expr(&node.orelse)
}
Expr::Compare(node) => {
contains_yield_expr(&node.left) || node.comparators.iter().any(contains_yield_expr)
}
Expr::Call(node) => {
contains_yield_expr(&node.func)
|| node.args.iter().any(contains_yield_expr)
|| node.keywords.iter().any(|kw| contains_yield_expr(&kw.value))
}
Expr::Attribute(node) => contains_yield_expr(&node.value),
Expr::Subscript(node) => {
contains_yield_expr(&node.value) || contains_yield_expr(&node.slice)
}
Expr::Starred(node) => contains_yield_expr(&node.value),
Expr::Tuple(node) => node.elts.iter().any(contains_yield_expr),
Expr::List(node) => node.elts.iter().any(contains_yield_expr),
Expr::Set(node) => node.elts.iter().any(contains_yield_expr),
Expr::Dict(node) => {
node.values.iter().any(contains_yield_expr)
|| node.keys.iter().any(|k| k.as_ref().is_some_and(contains_yield_expr))
}
Expr::JoinedStr(node) => node.values.iter().any(contains_yield_expr),
Expr::FormattedValue(node) => contains_yield_expr(&node.value),
Expr::NamedExpr(node) => contains_yield_expr(&node.value),
_ => false,
}
}