Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 8 additions & 30 deletions pyrefly/lib/alt/expr.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -873,23 +872,14 @@ 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,
HintRef::with_ty_opt(hint, elem_hint.as_ref()),
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(
Expand Down Expand Up @@ -3222,37 +3212,25 @@ impl<'ctx, 'answer, Ans: LookupAnswer> AnswersSolver<'ctx, 'answer, Ans> {
fn list_literal_infer(
&self,
list: &ExprList,
elt_hint: Option<ListElementHint>,
elt_hint: Option<Type>,
hint: Option<HintRef>,
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)))
}
}

Expand Down
38 changes: 2 additions & 36 deletions pyrefly/lib/alt/unwrap.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<Type>, Option<Type>) {
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.
Expand Down Expand Up @@ -333,27 +313,13 @@ impl<'ctx, 'answer, Ans: LookupAnswer> AnswersSolver<'ctx, 'answer, Ans> {
}
}

pub(crate) fn decompose_list(&self, hint: &Type) -> Option<ListElementHint> {
pub fn decompose_list(&self, hint: &Type) -> Option<Type> {
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
}
Expand Down
52 changes: 25 additions & 27 deletions pyrefly/lib/solver/solver.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -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.
Expand All @@ -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)
Expand Down Expand Up @@ -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 {
Expand All @@ -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
Expand All @@ -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);
Expand Down
20 changes: 16 additions & 4 deletions pyrefly/lib/solver/subset.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 => {
Expand Down Expand Up @@ -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.
Expand All @@ -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.
Expand Down
7 changes: 6 additions & 1 deletion pyrefly/lib/test/calls.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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])
"#,
);

Expand Down
37 changes: 33 additions & 4 deletions pyrefly/lib/test/contextual.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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])
"#,
);

Expand Down
Loading