diff --git a/pyrefly/lib/alt/expr.rs b/pyrefly/lib/alt/expr.rs index a9cff8b3dd..53b45834c3 100644 --- a/pyrefly/lib/alt/expr.rs +++ b/pyrefly/lib/alt/expr.rs @@ -97,7 +97,6 @@ use crate::alt::shape_extension::is_int_tuple_bound; use crate::alt::solve::TypeFormContext; use crate::alt::solve::UntypeContext; use crate::alt::unwrap::HintRef; -use crate::alt::unwrap::ListElementHint; use crate::binding::binding::Binding; use crate::binding::binding::Key; use crate::binding::binding::KeyYield; @@ -873,9 +872,6 @@ impl<'ctx, 'answer, Ans: LookupAnswer> AnswersSolver<'ctx, 'answer, Ans> { hint, |hint| self.decompose_list(hint), |elem_hint, hint| { - let (elem_hint, partial_fallback) = elem_hint - .map(ListElementHint::into_parts) - .unwrap_or_default(); self.ifs_infer(&x.generators, errors); let elem_ty = self.expr_infer_with_hint_promote( &x.elt, @@ -883,13 +879,7 @@ impl<'ctx, 'answer, Ans: LookupAnswer> AnswersSolver<'ctx, 'answer, Ans> { errors, HintCoercion::BestEffort, ); - let ty = self.heap.mk_class_type(self.stdlib.list(elem_ty)); - if let Some(partial_fallback) = partial_fallback { - self.solver() - .replace_unresolved_partials(ty, &partial_fallback) - } else { - ty - } + self.heap.mk_class_type(self.stdlib.list(elem_ty)) }, ), Expr::SetComp(x) => self.infer_with_decomposed_hint( @@ -3222,37 +3212,25 @@ impl<'ctx, 'answer, Ans: LookupAnswer> AnswersSolver<'ctx, 'answer, Ans> { fn list_literal_infer( &self, list: &ExprList, - elt_hint: Option, + elt_hint: Option, hint: Option, errors: &ErrorCollector, ) -> Type { if list.is_empty() { - let elem_ty = match elt_hint { - Some(ListElementHint::Hint(elem_hint)) => elem_hint, - Some(ListElementHint::UninformativeAny(_)) | None => self - .solver() + let elem_ty = elt_hint.unwrap_or_else(|| { + self.solver() .fresh_partial_contained(self.uniques, list.range) - .to_type(self.heap), - }; + .to_type(self.heap) + }); self.heap.mk_class_type(self.stdlib.list(elem_ty)) } else { - let (elt_hint, partial_fallback) = elt_hint - .map(ListElementHint::into_parts) - .unwrap_or_default(); let elem_tys = self.elts_infer( &list.elts, HintRef::with_ty_opt(hint, elt_hint.as_ref()), errors, ); - let ty = self - .heap - .mk_class_type(self.stdlib.list(self.unions(elem_tys))); - if let Some(partial_fallback) = partial_fallback { - self.solver() - .replace_unresolved_partials(ty, &partial_fallback) - } else { - ty - } + self.heap + .mk_class_type(self.stdlib.list(self.unions(elem_tys))) } } diff --git a/pyrefly/lib/alt/unwrap.rs b/pyrefly/lib/alt/unwrap.rs index e25decd060..b47e8078ed 100644 --- a/pyrefly/lib/alt/unwrap.rs +++ b/pyrefly/lib/alt/unwrap.rs @@ -26,26 +26,6 @@ use crate::types::types::Var; /// individually, as doing so would be prohibitively expensive. pub const MAX_HINT_WIDTH: usize = 32; -/// A contextual element hint for list literals and comprehensions. -pub(crate) enum ListElementHint { - /// A hint that should be applied while inferring the element. - Hint(Type), - /// An `Any` hint that is ignored for inference but retained as a fallback. - /// - /// The contained type is always `Any`, preserving its original `AnyStyle`. - UninformativeAny(Type), -} - -impl ListElementHint { - /// Split the hint into an inference hint and a fallback. - pub(crate) fn into_parts(self) -> (Option, Option) { - match self { - Self::Hint(ty) => (Some(ty), None), - Self::UninformativeAny(ty) => (None, Some(ty)), - } - } -} - // The error collector is None for a "soft" type hint, where we try to // match an expression against a hint, but fall back to the inferred type // without any errors if the hint is incompatible. @@ -333,27 +313,13 @@ impl<'ctx, 'answer, Ans: LookupAnswer> AnswersSolver<'ctx, 'answer, Ans> { } } - pub(crate) fn decompose_list(&self, hint: &Type) -> Option { + pub fn decompose_list(&self, hint: &Type) -> Option { let elem = self.fresh_var(); let list_type = self .heap .mk_class_type(self.stdlib.list(elem.to_type(self.heap))); if self.is_subset_eq(&list_type, hint) { - match self.resolve_var_opt(hint, elem) { - Some(elem_hint) - if elem_hint.is_any() - && hint - .collect_maybe_placeholder_vars() - .into_iter() - .any(|var| self.solver().var_is_quantified(var)) => - { - // An `Any` element hint obtained while the container hint still has an unsolved - // generic variable carries no information about the literal's elements. - Some(ListElementHint::UninformativeAny(elem_hint)) - } - Some(elem_hint) => Some(ListElementHint::Hint(elem_hint)), - None => None, - } + self.resolve_var_opt(hint, elem) } else { None } diff --git a/pyrefly/lib/solver/solver.rs b/pyrefly/lib/solver/solver.rs index ab6f15cbf4..bb6fba3987 100644 --- a/pyrefly/lib/solver/solver.rs +++ b/pyrefly/lib/solver/solver.rs @@ -793,30 +793,6 @@ impl Solver { ) } - /// Replace unresolved empty-container element types with `fallback` in a copy of `ty`. - pub(crate) fn replace_unresolved_partials(&self, mut ty: Type, fallback: &Type) -> Type { - self.expand_mut(&mut ty); - let partials: SmallSet<_> = { - let variables = self.variables.lock(); - ty.collect_maybe_placeholder_vars() - .into_iter() - .filter(|var| { - matches!( - &*variables.get(*var), - Variable::PartialQuantified(_) | Variable::PartialContained(_) - ) - }) - .collect() - }; - ty.transform_mut(&mut |ty| { - if matches!(ty, Type::Var(var) if partials.contains(var)) { - *ty = fallback.clone(); - } - }); - self.simplify_mut(&mut ty); - ty - } - /// Returns true if the given type is a Var that points to a partial variable. pub fn is_partial(&self, ty: &Type) -> bool { if let Type::Var(v) = ty { @@ -2752,6 +2728,7 @@ impl Solver { solver: self, type_order, gas: INITIAL_GAS, + checking_typevar_bound: false, active_call_context: CallContext::outside(), subset_cache: SmallMap::new(), class_protocol_assumptions: SmallSet::new(), @@ -3481,6 +3458,10 @@ pub struct Subset<'solver, 'subset, Ans: LookupAnswer> { pub(crate) solver: &'solver Solver, pub type_order: TypeOrder<'solver, Ans>, gas: Gas, + /// Bound validation must not use `Any` to solve decomposition variables. + /// These variables represent contextual hints, and `Any` in a generic bound + /// permits arbitrary element types rather than requiring an `Any` hint. + pub(crate) checking_typevar_bound: bool, /// Invariant: there is a single active call context for a subset query. /// Nested work is recursive subset checking inside the same call, not a /// nested full call pipeline with independent call-scoped solving. @@ -3504,7 +3485,7 @@ pub struct Subset<'solver, 'subset, Ans: LookupAnswer> { /// must be discarded. Only entries added during the failing computation are /// removed; entries from earlier (independent) computations are preserved. /// This works because `SmallMap` preserves insertion order. - pub subset_cache: SmallMap<(Type, Type, SubsetCacheContext), SubsetCacheEntry>, + pub subset_cache: SmallMap<(Type, Type, SubsetCacheContext, bool), SubsetCacheEntry>, /// Class-level recursive assumptions for protocol checks. /// When checking `got <: protocol` where got's type arguments contain Vars /// (indicating we're in a recursive pattern), we track (got_class, protocol_class) @@ -3919,10 +3900,10 @@ impl<'solver, 'subset, Ans: LookupAnswer> Subset<'solver, 'subset, Ans> { (answer, specialization_error) } Restriction::ShapeExtension(_) | Restriction::Bound(_) | Restriction::Unrestricted => { - if self.is_subset_eq(&t1_p, &bound).is_err() { + if self.is_subset_eq_typevar_bound(&t1_p, &bound).is_err() { // If the promoted type fails, try again with the original type, in case the bound itself is literal. // This could be more optimized, but errors are rare, so this code path should not be hot. - if self.is_subset_eq(t1, &bound).is_err() { + if self.is_subset_eq_typevar_bound(t1, &bound).is_err() { // If the original type is also an error, use the promoted type. let specialization_error = TypeVarSpecializationError::BadBoundSpecialization { @@ -3948,6 +3929,14 @@ impl<'solver, 'subset, Ans: LookupAnswer> Subset<'solver, 'subset, Ans> { .is_shape_extension_binding_source(v) } + /// Validate a type variable's bound without deriving contextual `Any` hints. + fn is_subset_eq_typevar_bound(&mut self, got: &Type, bound: &Type) -> Result<(), SubsetError> { + let previous = mem::replace(&mut self.checking_typevar_bound, true); + let result = self.is_subset_eq(got, bound); + self.checking_typevar_bound = previous; + result + } + /// Implementation of Var subset cases, calling onward to solve non-Var cases. /// /// This function does two things: it checks that got <: want, and it solves free variables assuming that @@ -3961,6 +3950,15 @@ impl<'solver, 'subset, Ans: LookupAnswer> Subset<'solver, 'subset, Ans> { fn is_subset_eq_var(&mut self, got: &Type, want: &Type) -> Result<(), SubsetError> { match (got, want) { _ if got == want => Ok(()), + (Type::Var(var), Type::Any(_)) | (Type::Any(_), Type::Var(var)) + if self.checking_typevar_bound + && matches!( + &*self.solver.variables.lock().get(*var), + Variable::Unwrap(_) + ) => + { + Ok(()) + } (Type::Var(v1), Type::Var(v2)) => { self.record_deferred_residual_target_vars(*v1, want); self.record_deferred_residual_target_vars(*v2, got); diff --git a/pyrefly/lib/solver/subset.rs b/pyrefly/lib/solver/subset.rs index 1570ba0b39..4f8710c6be 100644 --- a/pyrefly/lib/solver/subset.rs +++ b/pyrefly/lib/solver/subset.rs @@ -1760,10 +1760,14 @@ impl<'solver, 'subset, Ans: LookupAnswer> Subset<'solver, 'subset, Ans> { pub fn is_subset_eq_impl(&mut self, got: &Type, want: &Type) -> Result<(), SubsetError> { let context_key = self.active_call_context.subset_cache_context(); let cache_key = if self.can_be_recursive(got, want) { - // Cache keys include which argument is being matched, so argument-scoped comparisons - // do not suppress context-sensitive side effects. - // The vast majority of checks run under `Default` context. - let key = (got.clone(), want.clone(), context_key); + // The matched argument and bound validation affect inference side effects, + // so checks in different contexts must not share cached results. + let key = ( + got.clone(), + want.clone(), + context_key, + self.checking_typevar_bound, + ); if let Some(entry) = self.subset_cache.get(&key) { return match entry { SubsetCacheEntry::InProgress => { @@ -1815,6 +1819,11 @@ impl<'solver, 'subset, Ans: LookupAnswer> Subset<'solver, 'subset, Ans> { match (got, want) { (Type::Any(_), _) => { all(want.collect_maybe_placeholder_vars().iter(), |var| { + // Bound validation leaves decomposition variables unconstrained, + // but still resolves actual empty-container partials to `Any`. + if self.checking_typevar_bound { + return self.is_subset_eq(got, &var.to_type(&self.solver.heap)); + } // Variables in `want` now have `Any` as a lower bound. // TODO(https://github.com/facebook/pyrefly/issues/105): whether to add a lower // or upper bound should depend on variance. @@ -1826,6 +1835,9 @@ impl<'solver, 'subset, Ans: LookupAnswer> Subset<'solver, 'subset, Ans> { } (_, Type::Any(_)) => { all(got.collect_maybe_placeholder_vars().iter(), |var| { + if self.checking_typevar_bound { + return self.is_subset_eq(&var.to_type(&self.solver.heap), want); + } // Variables in `got` now have `Any` as an upper bound. // TODO(https://github.com/facebook/pyrefly/issues/105): whether to add a lower // or upper bound should depend on variance. diff --git a/pyrefly/lib/test/calls.rs b/pyrefly/lib/test/calls.rs index 10a8bc4548..3bf578ed83 100644 --- a/pyrefly/lib/test/calls.rs +++ b/pyrefly/lib/test/calls.rs @@ -325,11 +325,16 @@ reduce(max, [1,2]) ); testcase!( - test_iter_list_literal, + test_iter_container_literal, r#" from typing import Iterator, assert_type assert_type(iter([0]), Iterator[int]) +assert_type(iter({0}), Iterator[int]) +assert_type(iter({0: 0}), Iterator[int]) +assert_type(iter([x for x in [0]]), Iterator[int]) +assert_type(iter({x for x in [0]}), Iterator[int]) +assert_type(iter({x: x for x in [0]}), Iterator[int]) "#, ); diff --git a/pyrefly/lib/test/contextual.rs b/pyrefly/lib/test/contextual.rs index 6d9979db6e..66b2216781 100644 --- a/pyrefly/lib/test/contextual.rs +++ b/pyrefly/lib/test/contextual.rs @@ -252,20 +252,49 @@ reveal_type(keep([{}])) # E: revealed type: list[dict[Any, Any]] "#, ); +testcase!( + test_any_bounded_typevar_container_hints, + r#" +from typing import Any, Iterator, Protocol, assert_type + +def keep_set[T: set[Any]](x: T) -> T: ... +def keep_dict[T: dict[Any, Any]](x: T) -> T: ... +assert_type(keep_set({0}), set[int]) +assert_type(keep_set({x for x in [0]}), set[int]) +assert_type(keep_dict({0: "a"}), dict[int, str]) +assert_type(keep_dict({x: str(x) for x in [0]}), dict[int, str]) +assert_type(keep_dict({}), dict[Any, Any]) +assert_type(keep_dict({0: []}), dict[int, list[Any]]) + +class ReturnsIterator[T](Protocol): + def __iter__(self) -> T: ... + +def iterator[I: Iterator[Any]](x: ReturnsIterator[I]) -> I: ... +assert_type(iterator([0]), Iterator[int]) +assert_type(iterator({0}), Iterator[int]) +assert_type(iterator({0: "a"}), Iterator[int]) + +# Concrete bounds and explicit Any annotations still supply contextual types. +def floats[T: list[float]](x: T) -> T: ... +def explicit(x: list[Any]) -> list[Any]: ... +assert_type(floats([0]), list[float]) +assert_type(explicit([0]), list[Any]) + "#, +); + // An unconstrained generic return should not narrow permanently from its first mutation. testcase!( - bug = "Repeated append over-narrows an empty generic list return", test_any_bounded_typevar_empty_list_repeated_append, r#" -from typing import Any, reveal_type +from typing import Any, assert_type def keep[T: list[Any]](x: T) -> T: ... def f(): xs = keep([]) xs.append(1) - xs.append(2) # E: Argument `Literal[2]` is not assignable to parameter `object` with type `Literal[1]` - reveal_type(xs) # E: revealed type: list[Literal[1]] | list[Any] + xs.append(2) + assert_type(xs, list[Any]) "#, );