got a salsa cycle

This commit is contained in:
Jack O'Connor
2026-01-13 23:39:41 -08:00
parent 4bcf624e9e
commit 8effe4e84c
2 changed files with 209 additions and 77 deletions

View File

@@ -66,35 +66,75 @@ impl Loop {
}
}
/// Visitor that collects places with augmented assignments within a loop body.
/// Visitor that collects all places (symbols and members) that are assigned within a loop body.
///
/// This is used for cyclic control flow analysis. Only augmented assignments (like `i += 1`)
/// create value cycles where the new value depends on the old value. Simple assignments
/// (like `x = 2`) don't create cycles and don't need loop header definitions.
/// This is used for cyclic control flow analysis. Any place assigned in a loop may have
/// a different value on subsequent iterations, so we create loop header definitions
/// to represent the possible values at loop entry.
///
/// The visitor also tracks which places have augmented assignments (like `i += 1`), since
/// these create unbounded value growth that requires type widening.
///
/// The visitor does NOT recurse into nested function/class definitions since those are different scopes.
#[derive(Debug, Default)]
struct LoopBindingCollector {
/// The places that have augmented assignments within the loop body.
/// All places that are assigned within the loop body.
bound_places: Vec<PlaceExpr>,
/// Places that have augmented assignments (subset of `bound_places`).
augmented_places: Vec<PlaceExpr>,
}
impl LoopBindingCollector {
/// Collect all places with augmented assignments in the given statements.
fn collect(body: &[ast::Stmt]) -> Vec<PlaceExpr> {
/// Collect all places assigned in the given statements.
/// Returns `(all_bound_places, augmented_places)`.
fn collect(body: &[ast::Stmt]) -> (Vec<PlaceExpr>, Vec<PlaceExpr>) {
let mut collector = Self::default();
collector.visit_body(body);
collector.augmented_places
(collector.bound_places, collector.augmented_places)
}
/// Add a place from an expression target (assignment LHS).
fn add_place_from_target(&mut self, target: &ast::Expr) {
match target {
// Simple name assignment: x = ...
ast::Expr::Name(name) => {
self.bound_places.push(PlaceExpr::from_expr_name(name));
}
// Attribute/subscript assignment: x.y = ..., x[0] = ...
ast::Expr::Attribute(_) | ast::Expr::Subscript(_) => {
if let Some(place) = PlaceExpr::try_from_expr(target) {
self.bound_places.push(place);
}
}
// Unpacking: x, y = ... or (x, y) = ... or [x, y] = ...
ast::Expr::Tuple(tuple) => {
for elt in &tuple.elts {
self.add_place_from_target(elt);
}
}
ast::Expr::List(list) => {
for elt in &list.elts {
self.add_place_from_target(elt);
}
}
// Starred in unpacking: *x = ...
ast::Expr::Starred(starred) => {
self.add_place_from_target(&starred.value);
}
_ => {}
}
}
/// Add a place from an augmented assignment target.
fn add_augmented_place(&mut self, target: &ast::Expr) {
// First add to bound_places
self.add_place_from_target(target);
// Then also track as augmented
match target {
// Simple name: i += ...
ast::Expr::Name(name) => {
self.augmented_places.push(PlaceExpr::from_expr_name(name));
}
// Attribute/subscript: obj.x += ..., arr[0] += ...
ast::Expr::Attribute(_) | ast::Expr::Subscript(_) => {
if let Some(place) = PlaceExpr::try_from_expr(target) {
self.augmented_places.push(place);
@@ -108,47 +148,86 @@ impl LoopBindingCollector {
impl<'ast> Visitor<'ast> for LoopBindingCollector {
fn visit_stmt(&mut self, stmt: &'ast ast::Stmt) {
match stmt {
// Only augmented assignments create value cycles
// Assignment statements
ast::Stmt::Assign(node) => {
for target in &node.targets {
self.add_place_from_target(target);
}
}
ast::Stmt::AugAssign(node) => {
self.add_augmented_place(&node.target);
}
ast::Stmt::AnnAssign(node) => {
// Only a binding if there's a value
if node.value.is_some() {
self.add_place_from_target(&node.target);
}
}
// Simple assignments don't create cycles - skip them
ast::Stmt::Assign(_) | ast::Stmt::AnnAssign(_) => {}
// For loop - recurse into body
// For loop target is bound
ast::Stmt::For(node) => {
self.add_place_from_target(&node.target);
self.visit_body(&node.body);
self.visit_body(&node.orelse);
}
// With statement - recurse into body
// With statement binds the `as` target
ast::Stmt::With(node) => {
for item in &node.items {
if let Some(vars) = &item.optional_vars {
self.add_place_from_target(vars);
}
}
self.visit_body(&node.body);
}
// Try/except - recurse into all branches
// Exception handler binds the `as` name
ast::Stmt::Try(node) => {
self.visit_body(&node.body);
for handler in &node.handlers {
let ast::ExceptHandler::ExceptHandler(h) = handler;
if let Some(name) = &h.name {
self.bound_places
.push(PlaceExpr::Symbol(Symbol::new(name.id.clone())));
}
self.visit_body(&h.body);
}
self.visit_body(&node.orelse);
self.visit_body(&node.finalbody);
}
// Import/function/class definitions don't involve augmented assignments
ast::Stmt::Import(_)
| ast::Stmt::ImportFrom(_)
| ast::Stmt::FunctionDef(_)
| ast::Stmt::ClassDef(_) => {
// Do NOT recurse into function/class definitions - different scope
// Import statements bind names
ast::Stmt::Import(node) => {
for alias in &node.names {
let name = alias.asname.as_ref().unwrap_or(&alias.name);
self.bound_places
.push(PlaceExpr::Symbol(Symbol::new(name.id.clone())));
}
}
ast::Stmt::ImportFrom(node) => {
for alias in &node.names {
if &*alias.name != "*" {
let name = alias.asname.as_ref().unwrap_or(&alias.name);
self.bound_places
.push(PlaceExpr::Symbol(Symbol::new(name.id.clone())));
}
}
}
// Match statement - recurse into case bodies only
// Function/class definitions bind their name (but we don't recurse into their body)
ast::Stmt::FunctionDef(node) => {
self.bound_places
.push(PlaceExpr::Symbol(Symbol::new(node.name.id.clone())));
}
ast::Stmt::ClassDef(node) => {
self.bound_places
.push(PlaceExpr::Symbol(Symbol::new(node.name.id.clone())));
}
// Match statement can bind names in patterns
ast::Stmt::Match(node) => {
for case in &node.cases {
self.collect_pattern_bindings(&case.pattern);
self.visit_body(&case.body);
}
}
@@ -159,11 +238,65 @@ impl<'ast> Visitor<'ast> for LoopBindingCollector {
}
fn visit_expr(&mut self, expr: &'ast ast::Expr) {
// Continue walking to find augmented assignments in nested expressions
// Named expressions (walrus operator): x := ...
if let ast::Expr::Named(node) = expr {
self.add_place_from_target(&node.target);
}
walk_expr(self, expr);
}
}
impl LoopBindingCollector {
/// Collect bindings from match patterns.
fn collect_pattern_bindings(&mut self, pattern: &ast::Pattern) {
match pattern {
ast::Pattern::MatchAs(p) => {
if let Some(name) = &p.name {
self.bound_places
.push(PlaceExpr::Symbol(Symbol::new(name.id.clone())));
}
if let Some(pat) = &p.pattern {
self.collect_pattern_bindings(pat);
}
}
ast::Pattern::MatchStar(p) => {
if let Some(name) = &p.name {
self.bound_places
.push(PlaceExpr::Symbol(Symbol::new(name.id.clone())));
}
}
ast::Pattern::MatchMapping(p) => {
for pat in &p.patterns {
self.collect_pattern_bindings(pat);
}
if let Some(rest) = &p.rest {
self.bound_places
.push(PlaceExpr::Symbol(Symbol::new(rest.id.clone())));
}
}
ast::Pattern::MatchSequence(p) => {
for pat in &p.patterns {
self.collect_pattern_bindings(pat);
}
}
ast::Pattern::MatchClass(p) => {
for pat in &p.arguments.patterns {
self.collect_pattern_bindings(pat);
}
for kw in &p.arguments.keywords {
self.collect_pattern_bindings(&kw.pattern);
}
}
ast::Pattern::MatchOr(p) => {
for pat in &p.patterns {
self.collect_pattern_bindings(pat);
}
}
ast::Pattern::MatchValue(_) | ast::Pattern::MatchSingleton(_) => {}
}
}
}
struct ScopeInfo {
file_scope_id: FileScopeId,
/// Current loop state; None if we are not currently visiting a loop
@@ -2038,16 +2171,16 @@ impl<'ast> Visitor<'ast> for SemanticIndexBuilder<'_, 'ast> {
let predicate = self.record_expression_narrowing_constraint(test);
self.record_reachability_constraint(predicate);
// Collect places with augmented assignments (like `i += 1`) in the loop body.
// These create value cycles where the new value depends on the old value,
// requiring loop header definitions for proper type inference.
let augmented_places = LoopBindingCollector::collect(body);
// Collect all places assigned in the loop body.
// Loop headers are needed for all bindings to track values that could
// flow back from previous iterations.
let (bound_places, _augmented_places) = LoopBindingCollector::collect(body);
// Create loop header definitions for each place with an augmented assignment.
// This makes the widened type visible to uses inside the loop body.
for place_expr in augmented_places {
// Create loop header definitions for each bound place.
// The loop header type will be the union of the seed type (before loop)
// and all types assigned in the loop body.
for place_expr in bound_places {
let place_id = self.add_place(place_expr);
// Use push_additional_definition to allow multiple headers for the same while loop
let loop_header_ref = LoopHeaderDefinitionNodeRef {
while_stmt,
place: place_id,

View File

@@ -3814,67 +3814,66 @@ impl<'db, 'ast> TypeInferenceBuilder<'db, 'ast> {
/// Infer the type for a loop header definition.
///
/// Loop headers represent the fixed-point type of a place at loop entry. The type is computed
/// by widening the "seed" type (the type visible before the loop) to account for values that
/// could be assigned during loop iterations.
/// Loop headers represent the fixed-point type of a place at loop entry.
/// The type is the union of:
/// 1. The seed type (visible before the loop)
/// 2. Types from all bindings within the loop body
fn infer_loop_header_definition(
&mut self,
_loop_header: &LoopHeaderDefinitionKind,
loop_header: &LoopHeaderDefinitionKind,
definition: Definition<'db>,
) {
let db = self.db();
let module = self.module();
// Get bindings visible before this definition (seed bindings)
// Get seed type (visible before loop)
let file_scope = self.scope().file_scope_id(db);
let use_def = self.index.use_def_map(file_scope);
let seed_bindings = use_def.bindings_at_definition(definition);
// Compute seed type from the bindings
let seed_place = place_from_bindings(db, seed_bindings);
let seed_ty = match seed_place.place {
Place::Defined(defined) => defined.ty,
Place::Undefined => Type::unknown(),
};
// Widen the type for loop iteration
let widened_ty = self.widen_for_loop(seed_ty);
// Get the while loop body range
let while_stmt = loop_header.while_stmt(&module);
let body_range = while_stmt.body.first().map(|first| {
let last = while_stmt.body.last().unwrap();
ruff_text_size::TextRange::new(first.range().start(), last.range().end())
});
// Find all bindings for this place within the loop body and collect their types
let place = definition.place(db);
let mut all_types = vec![seed_ty];
if let Some(body_range) = body_range {
let all_bindings = use_def.reachable_bindings(place);
for binding_with_constraints in all_bindings {
let DefinitionState::Defined(binding) = binding_with_constraints.binding else {
continue;
};
// Skip the loop header itself
if binding == definition {
continue;
}
// Check if this binding is within the loop body
let binding_range = binding.kind(db).full_range(&module);
if body_range.contains_range(binding_range) {
let binding_ty = binding_type(db, binding);
all_types.push(binding_ty);
}
}
}
// Union all types together
let final_ty = UnionType::from_elements(db, all_types);
self.bindings
.insert(definition, widened_ty, self.multi_inference_state);
}
/// Widen a type for loop iteration.
///
/// For literals, this widens to their base type to handle cases where values
/// change across loop iterations:
/// ```python
/// i = 0
/// while ...:
/// i += 1 # i could be any int after multiple iterations
/// ```
fn widen_for_loop(&self, ty: Type<'db>) -> Type<'db> {
let db = self.db();
match ty {
// Widen integer literals to int
Type::IntLiteral(_) => KnownClass::Int.to_instance(db),
// Widen string literals to LiteralString
Type::StringLiteral(_) => Type::LiteralString,
// Widen unions element-wise
Type::Union(union) => {
let widened_elements: Vec<_> = union
.elements(db)
.iter()
.map(|&elem| self.widen_for_loop(elem))
.collect();
UnionType::from_elements(db, widened_elements)
}
// For other types, keep as-is for now
_ => ty,
}
.insert(definition, final_ty, self.multi_inference_state);
}
fn infer_match_statement(&mut self, match_statement: &ast::StmtMatch) {