From 8effe4e84ce4d68ced6be4adc2afe585100bf41c Mon Sep 17 00:00:00 2001 From: Jack O'Connor Date: Tue, 13 Jan 2026 23:39:41 -0800 Subject: [PATCH] got a salsa cycle --- .../src/semantic_index/builder.rs | 199 +++++++++++++++--- .../src/types/infer/builder.rs | 87 ++++---- 2 files changed, 209 insertions(+), 77 deletions(-) diff --git a/crates/ty_python_semantic/src/semantic_index/builder.rs b/crates/ty_python_semantic/src/semantic_index/builder.rs index 568cb04996..bcc6a2ec24 100644 --- a/crates/ty_python_semantic/src/semantic_index/builder.rs +++ b/crates/ty_python_semantic/src/semantic_index/builder.rs @@ -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, + /// Places that have augmented assignments (subset of `bound_places`). augmented_places: Vec, } impl LoopBindingCollector { - /// Collect all places with augmented assignments in the given statements. - fn collect(body: &[ast::Stmt]) -> Vec { + /// Collect all places assigned in the given statements. + /// Returns `(all_bound_places, augmented_places)`. + fn collect(body: &[ast::Stmt]) -> (Vec, Vec) { 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, diff --git a/crates/ty_python_semantic/src/types/infer/builder.rs b/crates/ty_python_semantic/src/types/infer/builder.rs index 70e2a59bd1..1713419121 100644 --- a/crates/ty_python_semantic/src/types/infer/builder.rs +++ b/crates/ty_python_semantic/src/types/infer/builder.rs @@ -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) {