PY-82818 PyCharm infers None for dict.get(..., None) if dict key/value types are unknown

GitOrigin-RevId: d12ed4598d91f1ab4824f6e5953e275672027064
This commit is contained in:
Aleksandr.Govenko
2025-09-26 19:05:33 +00:00
committed by intellij-monorepo-bot
parent ab81c86fe5
commit 039326217e
4 changed files with 59 additions and 24 deletions
@@ -2,7 +2,6 @@
package com.jetbrains.python.psi.types;
import com.intellij.openapi.util.*;
import com.intellij.openapi.util.registry.Registry;
import com.intellij.psi.PsiElement;
import com.intellij.psi.PsiFile;
import com.intellij.util.ArrayUtil;
@@ -271,16 +270,16 @@ public final class PyTypeChecker {
return false;
}
PyType substituted = context.mySubstitutions.typeVars.get(expected);
Ref<PyType> substitutedRef = context.mySubstitutions.typeVars.get(expected);
Ref<? extends PyType> defaultTypeRef = expected.getDefaultType();
if (defaultTypeRef != null) {
PyType defaultType = defaultTypeRef.get();
// Skip default substitution
if (defaultType != null && defaultType.equals(substituted)) {
substituted = null;
if (defaultType != null && defaultType.equals(Ref.deref(substitutedRef))) {
substitutedRef = null;
}
}
final PyType substitution = substituted;
final PyType substitution = Ref.deref(substitutedRef);
PyType bound = expected.getBound();
List<@Nullable PyType> constraints = expected.getConstraints();
int matchedConstraintIndex = -1;
@@ -306,8 +305,8 @@ public final class PyTypeChecker {
}
}
if (substitution != null) {
if (expected.equals(safeActual) || substitution.equals(expected)) {
if (substitutedRef != null) {
if (expected.equals(safeActual) || expected.equals(substitution)) {
return true;
}
@@ -321,12 +320,12 @@ public final class PyTypeChecker {
if (safeActual != null) {
PyType type = constraints.isEmpty() ? safeActual : constraints.get(matchedConstraintIndex);
context.mySubstitutions.typeVars.put(expected, type);
context.mySubstitutions.typeVars.put(expected, Ref.create(type));
}
else {
PyType effectiveBound = PyTypeUtil.getEffectiveBound(expected);
if (effectiveBound != null) {
context.mySubstitutions.typeVars.put(expected, PyUnionType.createWeakType(effectiveBound));
context.mySubstitutions.typeVars.put(expected, Ref.create(PyUnionType.createWeakType(effectiveBound)));
}
}
@@ -934,7 +933,7 @@ public final class PyTypeChecker {
Map<PyType, PyType> substitutionsFromClassDefinition = provider.getGenericSubstitutions(classType.getPyClass(), context);
for (Map.Entry<PyType, PyType> entry : substitutionsFromClassDefinition.entrySet()) {
if (entry.getKey() instanceof PyTypeVarType typeVarType) {
result.typeVars.put(typeVarType, entry.getValue());
result.typeVars.put(typeVarType, Ref.create(entry.getValue()));
}
else if (entry.getKey() instanceof PyTypeVarTupleType typeVarTuple) {
assert entry.getValue() instanceof PyPositionalVariadicType;
@@ -1092,7 +1091,7 @@ public final class PyTypeChecker {
boolean isAlreadyBound = existingSubstitutions.typeVars.containsKey(returnTypeParam) ||
existingSubstitutions.typeVars.containsKey(invert(returnTypeParam));
if (canGetBoundFromArguments && !isAlreadyBound) {
existingSubstitutions.typeVars.put(returnTypeParam, Ref.deref(returnTypeParam.getDefaultType()));
existingSubstitutions.typeVars.put(returnTypeParam, (Ref<PyType>)returnTypeParam.getDefaultType());
}
}
for (PyParamSpecType paramSpecType : typeParamsFromReturnType.paramSpecs) {
@@ -1224,10 +1223,11 @@ public final class PyTypeChecker {
}
return typeVarType;
}
PyType substitution = substitutions.typeVars.get(typeVarType);
if (substitution == null) {
Ref<PyType> substitutionRef = substitutions.typeVars.get(typeVarType);
PyType substitution = Ref.deref(substitutionRef);
if (substitutionRef == null) {
final PyInstantiableType<?> invertedTypeVar = invert(typeVarType);
final PyInstantiableType<?> invertedSubstitution = as(substitutions.typeVars.get(invertedTypeVar), PyInstantiableType.class);
final PyInstantiableType<?> invertedSubstitution = as(Ref.deref(substitutions.typeVars.get(invertedTypeVar)), PyInstantiableType.class);
if (invertedSubstitution != null) {
substitution = invert(invertedSubstitution);
}
@@ -1548,7 +1548,7 @@ public final class PyTypeChecker {
.select(PyClassType.class)
.map(type -> collectTypeSubstitutions(type, context))
.forEach(newSubstitutions -> {
for (Map.Entry<PyTypeVarType, PyType> typeVarMapping : newSubstitutions.typeVars.entrySet()) {
for (Map.Entry<PyTypeVarType, Ref<PyType>> typeVarMapping : newSubstitutions.typeVars.entrySet()) {
substitutions.typeVars.putIfAbsent(typeVarMapping.getKey(), typeVarMapping.getValue());
}
for (Map.Entry<PyTypeVarTupleType, PyPositionalVariadicType> typeVarMapping : newSubstitutions.typeVarTuples.entrySet()) {
@@ -1738,7 +1738,7 @@ public final class PyTypeChecker {
if (mapping != null) {
for (Couple<PyType> pair : mapping.getMappedTypes()) {
if (pair.getFirst() instanceof PyTypeVarType typeVar) {
substitutions.typeVars.put(typeVar, pair.getSecond());
substitutions.typeVars.put(typeVar, Ref.create(pair.getSecond()));
}
else if (pair.getFirst() instanceof PyTypeVarTupleType typeVarTuple) {
substitutions.typeVarTuples.put(typeVarTuple, as(pair.getSecond(), PyPositionalVariadicType.class));
@@ -1806,11 +1806,13 @@ public final class PyTypeChecker {
@ApiStatus.Experimental
public static class GenericSubstitutions {
private final @NotNull Map<PyTypeVarType, PyType> typeVars;
// Nullable-Nullable because of com.jetbrains.python.psi.types.PyTypeChecker.collectTypeSubstitutions
private final @NotNull Map<PyTypeVarType, @Nullable Ref<@Nullable PyType>> typeVars;
private final @NotNull Map<PyTypeVarTupleType, PyPositionalVariadicType> typeVarTuples;
private final @NotNull Map<PyTypeVarTupleType, @Nullable PyPositionalVariadicType> typeVarTuples;
private final @NotNull Map<PyParamSpecType, PyCallableParameterVariadicType> paramSpecs;
private final @NotNull Map<PyParamSpecType, @Nullable PyCallableParameterVariadicType> paramSpecs;
private @Nullable PyType qualifierType;
@@ -1818,6 +1820,7 @@ public final class PyTypeChecker {
this(
EntryStream.of(typeParameters)
.selectKeys(PyTypeVarType.class)
.mapValues(Ref::create)
.toCustomMap(LinkedHashMap::new),
EntryStream.of(typeParameters)
.selectKeys(PyTypeVarTupleType.class)
@@ -1835,7 +1838,7 @@ public final class PyTypeChecker {
this(new LinkedHashMap<>(), new LinkedHashMap<>(), new LinkedHashMap<>(), null);
}
private GenericSubstitutions(@NotNull Map<PyTypeVarType, PyType> typeVars,
private GenericSubstitutions(@NotNull Map<PyTypeVarType, Ref<PyType>> typeVars,
@NotNull Map<PyTypeVarTupleType, PyPositionalVariadicType> typeVarTuples,
@NotNull Map<PyParamSpecType, PyCallableParameterVariadicType> paramSpecs,
@Nullable PyType qualifierType) {
@@ -1849,7 +1852,7 @@ public final class PyTypeChecker {
return Collections.unmodifiableMap(paramSpecs);
}
public @NotNull Map<PyTypeVarType, PyType> getTypeVars() {
public @NotNull Map<PyTypeVarType, Ref<PyType>> getTypeVars() {
return Collections.unmodifiableMap(typeVars);
}
@@ -2,6 +2,7 @@
package com.jetbrains.python.psi.types;
import com.intellij.openapi.extensions.ExtensionPointName;
import com.intellij.openapi.util.Ref;
import org.jetbrains.annotations.ApiStatus;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
@@ -23,9 +24,9 @@ public interface PyTypeCheckerExtension {
@NotNull TypeEvalContext context,
@NotNull PyTypeChecker.GenericSubstitutions substitutions) {
Map<PyGenericType, PyType> legacyTypeVarSubs = new HashMap<>();
for (Map.Entry<PyTypeVarType, PyType> entry : substitutions.getTypeVars().entrySet()) {
for (Map.Entry<PyTypeVarType, Ref<PyType>> entry : substitutions.getTypeVars().entrySet()) {
if (entry.getKey() instanceof PyGenericType legacyTypeVar) {
legacyTypeVarSubs.put(legacyTypeVar, entry.getValue());
legacyTypeVarSubs.put(legacyTypeVar, Ref.deref(entry.getValue()));
}
}
return match(expected, actual, context, legacyTypeVarSubs);
@@ -18,6 +18,37 @@ import java.util.Map;
public class Py3TypeTest extends PyTestCase {
public static final String TEST_DIRECTORY = "/types/";
/**
Overload signatures for dict.get and dict.pop in builtins.pyi differ slightly,
dict.get has default value for "default" parameter. This affect the logic of overload resolution.
Therefore it makes sense to test both.
<p>
<pre>{@code
@overload
def get(self, key: _KT, default: None = None, /) -> _VT | None: ...
# mode overloads...
}</pre>
<p>
<pre>{@code
@overload
def pop(self, key: _KT, /) -> _VT: ...
# mode overloads...
}</pre>
*/
// PY-82818
public void testGetFromDictWithDefaultNoneValue() {
doTest("Any | None", """
d = {}
expr = d.get("abc", None)""");
}
// PY-82818
public void testPopFromDictWithDefaultNoneValue() {
doTest("Any", """
d = {}
expr = d.pop("abc", None)""");
}
public void testYieldInsideLambda() {
// Checks that foo is not a generator
doTest("int", """
@@ -1775,7 +1775,7 @@ public class PyTypeTest extends PyTestCase {
// PY-20409
public void testGetFromDictWithDefaultNoneValue() {
doTest("Optional[Any]",
doTest("Any",
"d = {}\n" +
"expr = d.get(\"abc\", None)");
}