use std::collections::{HashMap, HashSet};
use crate::{ExprType, Statement, StatementType};
pub struct ScopeBindings {
pub assigned: Vec<String>,
pub needs_mut: HashSet<String>,
pub optional: HashSet<String>,
}
#[derive(Clone, Copy, PartialEq)]
enum Init {
No,
Maybe,
Yes,
}
pub(crate) const MUTATING_METHODS: &[&str] = &[
"append",
"extend",
"insert",
"remove",
"pop",
"popitem",
"clear",
"sort",
"reverse",
"update",
"add",
"discard",
"setdefault",
"intersection_update",
"difference_update",
"symmetric_difference_update",
"push",
"read",
"readline",
"readlines",
"write",
"writelines",
"close",
"writerow",
"writerows",
];
struct Analysis<'r> {
assigned: Vec<String>,
needs_mut: HashSet<String>,
optional: HashSet<String>,
state: HashMap<String, Init>,
resolve_call: &'r dyn Fn(&crate::Call) -> Option<bool>,
}
impl Analysis<'_> {
fn init_of(&self, name: &str) -> Init {
self.state.get(name).copied().unwrap_or(Init::No)
}
fn record_store(&mut self, name: &str, multi: bool) {
if !self.assigned.contains(&name.to_string()) {
self.assigned.push(name.to_string());
}
if multi || self.init_of(name) != Init::No {
self.needs_mut.insert(name.to_string());
}
self.state.insert(name.to_string(), Init::Yes);
}
fn record_mutation(&mut self, name: &str) {
self.needs_mut.insert(name.to_string());
}
}
fn merge_states(
branches: Vec<HashMap<String, Init>>,
) -> HashMap<String, Init> {
let mut merged: HashMap<String, Init> = HashMap::new();
let keys: HashSet<String> = branches
.iter()
.flat_map(|b| b.keys().cloned())
.collect();
for key in keys {
let states: Vec<Init> = branches
.iter()
.map(|b| b.get(&key).copied().unwrap_or(Init::No))
.collect();
let all_yes = states.iter().all(|s| *s == Init::Yes);
let all_no = states.iter().all(|s| *s == Init::No);
let v = if all_yes {
Init::Yes
} else if all_no {
Init::No
} else {
Init::Maybe
};
merged.insert(key, v);
}
merged
}
pub fn analyze_scope(body: &[Statement], initialized: &[String]) -> ScopeBindings {
analyze_scope_with(body, initialized, &|_| None)
}
pub(crate) fn analyze_scope_with(
body: &[Statement],
initialized: &[String],
resolve_call: &dyn Fn(&crate::Call) -> Option<bool>,
) -> ScopeBindings {
let mut a = Analysis {
assigned: Vec::new(),
needs_mut: HashSet::new(),
optional: HashSet::new(),
state: initialized
.iter()
.map(|n| (n.clone(), Init::Yes))
.collect(),
resolve_call,
};
walk_stmts(body, &mut a, false);
let assigned = a
.assigned
.into_iter()
.filter(|n| !initialized.contains(n))
.collect();
ScopeBindings {
assigned,
needs_mut: a.needs_mut,
optional: a.optional,
}
}
fn record_target(target: &ExprType, a: &mut Analysis<'_>, multi: bool) {
match target {
ExprType::Name(name) => a.record_store(&name.id, multi),
ExprType::Tuple(tuple) => {
for elt in &tuple.elts {
if let ExprType::Name(name) = elt {
a.record_store(&name.id, multi);
a.record_mutation(&name.id);
} else {
record_target(elt, a, multi);
}
}
}
ExprType::Subscript(sub) => {
if let Some(name) = chain_base_name(&sub.value) {
a.record_mutation(name);
}
walk_subscript_kind(&sub.kind, a);
}
ExprType::Attribute(attr) => {
if let Some(name) = chain_base_name(&attr.value) {
a.record_mutation(name);
}
}
_ => {}
}
}
pub(crate) fn chain_base_name(expr: &ExprType) -> Option<&str> {
match expr {
ExprType::Name(name) => Some(&name.id),
ExprType::Subscript(sub) => chain_base_name(&sub.value),
ExprType::Attribute(attr) => chain_base_name(&attr.value),
_ => None,
}
}
pub(crate) fn class_call_resolver<'a>(
ctx: &'a crate::CodeGenContext,
symbols: &'a crate::SymbolTableScopes,
) -> impl Fn(&crate::Call) -> Option<bool> + 'a {
move |call| {
let ExprType::Attribute(attr) = call.func.as_ref() else {
return None;
};
let class = crate::receiver_class(&attr.value, ctx, symbols)?;
if !class.methods().any(|m| m.name == attr.attr) {
return None;
}
Some(class.method_needs_mut_self(&attr.attr, symbols))
}
}
fn walk_stmts(body: &[Statement], a: &mut Analysis<'_>, multi: bool) {
for stmt in body {
match &stmt.statement {
StatementType::Assign(assign) => {
walk_expr(&assign.value, a);
let value_is_none = crate::is_none_expr(&assign.value);
for target in &assign.targets {
record_target(target, a, multi);
if value_is_none {
if let ExprType::Name(name) = target {
a.optional.insert(name.id.clone());
}
}
}
}
StatementType::AugAssign(aug) => {
walk_expr(&aug.value, a);
if let ExprType::Name(name) = &aug.target {
a.record_store(&name.id, multi);
a.record_mutation(&name.id);
} else {
record_target(&aug.target, a, multi);
}
}
StatementType::Expr(e) => walk_expr(&e.value, a),
StatementType::Call(call) => walk_call(call, a),
StatementType::Return(Some(e)) => walk_expr(&e.value, a),
StatementType::Assert { test, msg } => {
walk_expr(test, a);
if let Some(m) = msg {
walk_expr(m, a);
}
}
StatementType::Raise(raise) => {
if let Some(exc) = &raise.exc {
walk_expr(exc, a);
}
if let Some(cause) = &raise.cause {
walk_expr(cause, a);
}
}
StatementType::If(s) => {
walk_expr(&s.test, a);
let before = a.state.clone();
walk_stmts(&s.body, a, multi);
let after_body = std::mem::replace(&mut a.state, before);
walk_stmts(&s.orelse, a, multi);
let after_else = std::mem::replace(&mut a.state, HashMap::new());
a.state = merge_states(vec![after_body, after_else]);
}
StatementType::While(s) => {
walk_expr(&s.test, a);
walk_loop(&s.body, &s.orelse, a, multi);
}
StatementType::For(s) => {
walk_expr(&s.iter, a);
walk_loop(&s.body, &s.orelse, a, multi);
}
StatementType::AsyncFor(s) => {
walk_expr(&s.iter, a);
walk_loop(&s.body, &s.orelse, a, multi);
}
StatementType::Try(s) => {
let before = a.state.clone();
walk_stmts(&s.body, a, true);
let after_body = a.state.clone();
let handler_entry =
merge_states(vec![before, after_body.clone()]);
let mut exits = Vec::new();
for handler in &s.handlers {
a.state = handler_entry.clone();
walk_stmts(&handler.body, a, multi);
exits.push(std::mem::replace(&mut a.state, HashMap::new()));
}
a.state = after_body;
walk_stmts(&s.orelse, a, multi);
exits.push(std::mem::replace(&mut a.state, HashMap::new()));
a.state = merge_states(exits);
walk_stmts(&s.finalbody, a, multi);
}
StatementType::With(s) => {
for item in &s.items {
walk_expr(&item.context_expr, a);
}
walk_stmts(&s.body, a, multi);
}
StatementType::AsyncWith(s) => {
for item in &s.items {
walk_expr(&item.context_expr, a);
}
walk_stmts(&s.body, a, multi);
}
_ => {}
}
}
}
fn walk_loop(
body: &[Statement],
orelse: &[Statement],
a: &mut Analysis<'_>,
outer_multi: bool,
) {
let before = a.state.clone();
walk_stmts(body, a, true);
let after_body = std::mem::replace(&mut a.state, HashMap::new());
a.state = merge_states(vec![before, after_body]);
walk_stmts(orelse, a, outer_multi);
}
const FIRST_ARG_MUTATORS: &[&str] = &[
"heappush",
"heappop",
"heapify",
"heappushpop",
"heapreplace",
"writer",
];
fn walk_call(call: &crate::Call, a: &mut Analysis<'_>) {
if let ExprType::Attribute(attr) = call.func.as_ref() {
if let ExprType::Name(m) = attr.value.as_ref() {
if matches!(m.id.as_str(), "heapq" | "csv")
&& FIRST_ARG_MUTATORS.contains(&attr.attr.as_str())
{
if let Some(first) = call.args.first() {
if let Some(name) = chain_base_name(first) {
a.record_mutation(name);
}
}
}
}
let mutates = match (a.resolve_call)(call) {
Some(verdict) => verdict,
None => MUTATING_METHODS.contains(&attr.attr.as_str()),
};
if mutates {
if let Some(name) = chain_base_name(&attr.value) {
a.record_mutation(name);
}
}
walk_expr(&attr.value, a);
} else {
if let ExprType::Name(n) = call.func.as_ref() {
if FIRST_ARG_MUTATORS.contains(&n.id.as_str()) {
if let Some(first) = call.args.first() {
if let Some(name) = chain_base_name(first) {
a.record_mutation(name);
}
}
}
}
walk_expr(&call.func, a);
}
for arg in &call.args {
walk_expr(arg, a);
}
for kw in &call.keywords {
walk_expr(&kw.value, a);
}
}
fn walk_subscript_kind(kind: &crate::SubscriptKind, a: &mut Analysis<'_>) {
match kind {
crate::SubscriptKind::Index(i) => walk_expr(i, a),
crate::SubscriptKind::Slice { lower, upper, step } => {
for bound in [lower, upper, step].into_iter().flatten() {
walk_expr(bound, a);
}
}
}
}
fn walk_expr(expr: &ExprType, a: &mut Analysis<'_>) {
match expr {
ExprType::Call(call) => walk_call(call, a),
ExprType::BinOp(op) => {
walk_expr(&op.left, a);
walk_expr(&op.right, a);
}
ExprType::BoolOp(op) => {
for v in &op.values {
walk_expr(v, a);
}
}
ExprType::UnaryOp(op) => walk_expr(&op.operand, a),
ExprType::Compare(cmp) => {
walk_expr(&cmp.left, a);
for c in &cmp.comparators {
walk_expr(c, a);
}
}
ExprType::IfExp(e) => {
walk_expr(&e.test, a);
walk_expr(&e.body, a);
walk_expr(&e.orelse, a);
}
ExprType::NamedExpr(e) => {
walk_expr(&e.left, a);
walk_expr(&e.right, a);
}
ExprType::Dict(d) => {
for k in d.keys.iter().flatten() {
walk_expr(k, a);
}
for v in &d.values {
walk_expr(v, a);
}
}
ExprType::Set(s) => {
for e in &s.elts {
walk_expr(e, a);
}
}
ExprType::List(elts) => {
for e in elts {
walk_expr(e, a);
}
}
ExprType::Tuple(t) => {
for e in &t.elts {
walk_expr(e, a);
}
}
ExprType::Attribute(attr) => walk_expr(&attr.value, a),
ExprType::Subscript(sub) => {
walk_expr(&sub.value, a);
walk_subscript_kind(&sub.kind, a);
}
ExprType::Starred(s) => walk_expr(&s.value, a),
ExprType::Await(e) => walk_expr(&e.value, a),
ExprType::Yield(y) => {
if let Some(v) = &y.value {
walk_expr(v, a);
}
}
ExprType::YieldFrom(y) => walk_expr(&y.value, a),
ExprType::FormattedValue(f) => walk_expr(&f.value, a),
ExprType::JoinedStr(j) => {
for v in &j.values {
walk_expr(v, a);
}
}
_ => {}
}
}