mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
Merge branch 'union-type-inference'
This commit is contained in:
@@ -68,6 +68,7 @@ import com.jetbrains.python.magicLiteral.PyMagicLiteralTools;
|
||||
import com.jetbrains.python.psi.impl.PyBuiltinCache;
|
||||
import com.jetbrains.python.psi.impl.PyPsiUtils;
|
||||
import com.jetbrains.python.psi.impl.PythonLanguageLevelPusher;
|
||||
import com.jetbrains.python.psi.resolve.RatedResolveResult;
|
||||
import com.jetbrains.python.psi.types.*;
|
||||
import com.jetbrains.python.refactoring.classes.PyDependenciesComparator;
|
||||
import com.jetbrains.python.refactoring.classes.extractSuperclass.PyExtractSuperclassHelper;
|
||||
@@ -750,6 +751,40 @@ public class PyUtil {
|
||||
return currentElement;
|
||||
}
|
||||
|
||||
@NotNull
|
||||
public static List<PsiElement> multiResolveTopPriority(@NotNull PsiPolyVariantReference reference) {
|
||||
return filterTopPriorityResults(reference.multiResolve(false));
|
||||
}
|
||||
|
||||
@NotNull
|
||||
private static List<PsiElement> filterTopPriorityResults(@NotNull ResolveResult[] resolveResults) {
|
||||
if (resolveResults.length == 0) {
|
||||
return Collections.emptyList();
|
||||
}
|
||||
final List<PsiElement> filtered = new ArrayList<PsiElement>();
|
||||
final int maxRate = getMaxRate(resolveResults);
|
||||
for (ResolveResult resolveResult : resolveResults) {
|
||||
final int rate = resolveResult instanceof RatedResolveResult ? ((RatedResolveResult)resolveResult).getRate() : 0;
|
||||
if (rate >= maxRate) {
|
||||
filtered.add(resolveResult.getElement());
|
||||
}
|
||||
}
|
||||
return filtered;
|
||||
}
|
||||
|
||||
private static int getMaxRate(@NotNull ResolveResult[] resolveResults) {
|
||||
int maxRate = Integer.MIN_VALUE;
|
||||
for (ResolveResult resolveResult : resolveResults) {
|
||||
if (resolveResult instanceof RatedResolveResult) {
|
||||
final int rate = ((RatedResolveResult)resolveResult).getRate();
|
||||
if (rate > maxRate) {
|
||||
maxRate = rate;
|
||||
}
|
||||
}
|
||||
}
|
||||
return maxRate;
|
||||
}
|
||||
|
||||
/**
|
||||
* Gets class init method
|
||||
*
|
||||
|
||||
@@ -17,9 +17,10 @@ package com.jetbrains.python.psi.impl;
|
||||
|
||||
import com.intellij.codeInsight.completion.CompletionUtil;
|
||||
import com.intellij.openapi.util.Pair;
|
||||
import com.intellij.openapi.util.Ref;
|
||||
import com.intellij.psi.PsiElement;
|
||||
import com.intellij.psi.PsiPolyVariantReference;
|
||||
import com.intellij.psi.PsiReference;
|
||||
import com.intellij.psi.ResolveResult;
|
||||
import com.intellij.psi.util.PsiTreeUtil;
|
||||
import com.jetbrains.python.FunctionParameter;
|
||||
import com.jetbrains.python.PyNames;
|
||||
@@ -429,57 +430,18 @@ public class PyCallExpressionHelper {
|
||||
}
|
||||
// normal cases
|
||||
final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context);
|
||||
ResolveResult[] targets = ((PyReferenceExpression)callee).getReference(resolveContext).multiResolve(false);
|
||||
if (targets.length > 0) {
|
||||
PsiElement target = targets[0].getElement();
|
||||
if (target == null) {
|
||||
return null;
|
||||
}
|
||||
PyClass cls = null;
|
||||
PyFunction init = null;
|
||||
if (target instanceof PyClass) {
|
||||
cls = (PyClass)target;
|
||||
init = cls.findInitOrNew(true);
|
||||
}
|
||||
else if (target instanceof PyFunction) {
|
||||
final PyFunction f = (PyFunction)target;
|
||||
if (PyNames.INIT.equals(f.getName())) {
|
||||
init = f;
|
||||
cls = f.getContainingClass();
|
||||
final PsiPolyVariantReference reference = ((PyReferenceExpression)callee).getReference(resolveContext);
|
||||
final List<PyType> members = new ArrayList<PyType>();
|
||||
for (PsiElement target : PyUtil.multiResolveTopPriority(reference)) {
|
||||
if (target != null) {
|
||||
final Ref<? extends PyType> typeRef = getCallTargetReturnType(call, target, context);
|
||||
if (typeRef != null) {
|
||||
members.add(typeRef.get());
|
||||
}
|
||||
}
|
||||
if (init != null) {
|
||||
final PyType t = init.getCallType(context, call);
|
||||
if (cls != null) {
|
||||
if (init.getContainingClass() != cls) {
|
||||
if (t instanceof PyCollectionType) {
|
||||
final PyType elementType = ((PyCollectionType)t).getElementType(context);
|
||||
return new PyCollectionTypeImpl(cls, false, elementType);
|
||||
}
|
||||
return new PyClassTypeImpl(cls, false);
|
||||
}
|
||||
}
|
||||
if (t != null && !(t instanceof PyNoneType)) {
|
||||
return t;
|
||||
}
|
||||
if (cls != null && t == null) {
|
||||
final PyFunction newMethod = cls.findMethodByName(PyNames.NEW, true);
|
||||
if (newMethod != null && !PyBuiltinCache.getInstance(call).isBuiltin(newMethod)) {
|
||||
return PyUnionType.createWeakType(new PyClassTypeImpl(cls, false));
|
||||
}
|
||||
}
|
||||
}
|
||||
if (cls != null) {
|
||||
return new PyClassTypeImpl(cls, false);
|
||||
}
|
||||
final PyType providedType = PyReferenceExpressionImpl.getReferenceTypeFromProviders(target, context, call);
|
||||
if (providedType instanceof PyCallableType) {
|
||||
return ((PyCallableType)providedType).getCallType(context, call);
|
||||
}
|
||||
if (target instanceof Callable) {
|
||||
final Callable callable = (Callable)target;
|
||||
return callable.getCallType(context, call);
|
||||
}
|
||||
}
|
||||
if (!members.isEmpty()) {
|
||||
return PyUnionType.union(members);
|
||||
}
|
||||
}
|
||||
if (callee == null) {
|
||||
@@ -499,6 +461,57 @@ public class PyCallExpressionHelper {
|
||||
}
|
||||
}
|
||||
|
||||
@Nullable
|
||||
private static Ref<? extends PyType> getCallTargetReturnType(@NotNull PyCallExpression call, @NotNull PsiElement target,
|
||||
@NotNull TypeEvalContext context) {
|
||||
PyClass cls = null;
|
||||
PyFunction init = null;
|
||||
if (target instanceof PyClass) {
|
||||
cls = (PyClass)target;
|
||||
init = cls.findInitOrNew(true);
|
||||
}
|
||||
else if (target instanceof PyFunction) {
|
||||
final PyFunction f = (PyFunction)target;
|
||||
if (PyNames.INIT.equals(f.getName())) {
|
||||
init = f;
|
||||
cls = f.getContainingClass();
|
||||
}
|
||||
}
|
||||
if (init != null) {
|
||||
final PyType t = init.getCallType(context, call);
|
||||
if (cls != null) {
|
||||
if (init.getContainingClass() != cls) {
|
||||
if (t instanceof PyCollectionType) {
|
||||
final PyType elementType = ((PyCollectionType)t).getElementType(context);
|
||||
return Ref.create(new PyCollectionTypeImpl(cls, false, elementType));
|
||||
}
|
||||
return Ref.create(new PyClassTypeImpl(cls, false));
|
||||
}
|
||||
}
|
||||
if (t != null && !(t instanceof PyNoneType)) {
|
||||
return Ref.create(t);
|
||||
}
|
||||
if (cls != null && t == null) {
|
||||
final PyFunction newMethod = cls.findMethodByName(PyNames.NEW, true);
|
||||
if (newMethod != null && !PyBuiltinCache.getInstance(call).isBuiltin(newMethod)) {
|
||||
return Ref.create(PyUnionType.createWeakType(new PyClassTypeImpl(cls, false)));
|
||||
}
|
||||
}
|
||||
}
|
||||
if (cls != null) {
|
||||
return Ref.create(new PyClassTypeImpl(cls, false));
|
||||
}
|
||||
final PyType providedType = PyReferenceExpressionImpl.getReferenceTypeFromProviders(target, context, call);
|
||||
if (providedType instanceof PyCallableType) {
|
||||
return Ref.create(((PyCallableType)providedType).getCallType(context, call));
|
||||
}
|
||||
if (target instanceof Callable) {
|
||||
final Callable callable = (Callable)target;
|
||||
return Ref.create(callable.getCallType(context, call));
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
@NotNull
|
||||
private static Maybe<PyType> getSuperCallType(@NotNull PyCallExpression call, TypeEvalContext context) {
|
||||
final PyExpression callee = call.getCallee();
|
||||
|
||||
@@ -216,12 +216,14 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere
|
||||
return typeOfProperty.get();
|
||||
}
|
||||
}
|
||||
ResolveResult[] targets = getReference(PyResolveContext.noImplicits().withTypeEvalContext(context)).multiResolve(false);
|
||||
if (targets.length == 0) {
|
||||
final PsiPolyVariantReference reference = getReference(PyResolveContext.noImplicits().withTypeEvalContext(context));
|
||||
final List<PsiElement> targets = PyUtil.multiResolveTopPriority(reference);
|
||||
if (targets.isEmpty()) {
|
||||
return getQualifiedReferenceTypeByControlFlow(context);
|
||||
}
|
||||
for (ResolveResult resolveResult : targets) {
|
||||
PsiElement target = resolveResult.getElement();
|
||||
|
||||
final List<PyType> members = new ArrayList<PyType>();
|
||||
for (PsiElement target : targets) {
|
||||
if (target == this || target == null) {
|
||||
continue;
|
||||
}
|
||||
@@ -229,12 +231,10 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere
|
||||
LOG.error("Reference " + this + " resolved to invalid element " + target + " (text=" + target.getText() + ")");
|
||||
continue;
|
||||
}
|
||||
type = getTypeFromTarget(target, context, this);
|
||||
if (type != null) {
|
||||
return type;
|
||||
}
|
||||
members.add(getTypeFromTarget(target, context, this));
|
||||
}
|
||||
return null;
|
||||
|
||||
return PyUnionType.union(members);
|
||||
}
|
||||
finally {
|
||||
TypeEvalStack.evaluated(this);
|
||||
|
||||
@@ -30,6 +30,9 @@ import com.jetbrains.python.psi.types.*;
|
||||
import org.jetbrains.annotations.NotNull;
|
||||
import org.jetbrains.annotations.Nullable;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
|
||||
/**
|
||||
* @author yole
|
||||
*/
|
||||
@@ -66,29 +69,32 @@ public class PySubscriptionExpressionImpl extends PyElementImpl implements PySub
|
||||
@Nullable
|
||||
@Override
|
||||
public PyType getType(@NotNull TypeEvalContext context, @NotNull TypeEvalContext.Key key) {
|
||||
PyType res = null;
|
||||
final PsiReference ref = getReference(PyResolveContext.noImplicits().withTypeEvalContext(context));
|
||||
final PsiElement resolved = ref.resolve();
|
||||
if (resolved instanceof Callable) {
|
||||
res = ((Callable)resolved).getCallType(context, this);
|
||||
}
|
||||
if (PyTypeChecker.isUnknown(res) || res instanceof PyNoneType) {
|
||||
final PyExpression indexExpression = getIndexExpression();
|
||||
if (indexExpression != null) {
|
||||
final PyType type = context.getType(getOperand());
|
||||
final PyClass cls = (type instanceof PyClassType) ? ((PyClassType)type).getPyClass() : null;
|
||||
if (cls != null && PyABCUtil.isSubclass(cls, PyNames.MAPPING)) {
|
||||
return res;
|
||||
}
|
||||
if (type instanceof PySubscriptableType) {
|
||||
res = ((PySubscriptableType)type).getElementType(indexExpression, context);
|
||||
}
|
||||
else if (type instanceof PyCollectionType) {
|
||||
res = ((PyCollectionType)type).getElementType(context);
|
||||
final PsiPolyVariantReference reference = getReference(PyResolveContext.noImplicits().withTypeEvalContext(context));
|
||||
final List<PyType> members = new ArrayList<PyType>();
|
||||
for (PsiElement resolved : PyUtil.multiResolveTopPriority(reference)) {
|
||||
PyType res = null;
|
||||
if (resolved instanceof Callable) {
|
||||
res = ((Callable)resolved).getCallType(context, this);
|
||||
}
|
||||
if (PyTypeChecker.isUnknown(res) || res instanceof PyNoneType) {
|
||||
final PyExpression indexExpression = getIndexExpression();
|
||||
if (indexExpression != null) {
|
||||
final PyType type = context.getType(getOperand());
|
||||
final PyClass cls = (type instanceof PyClassType) ? ((PyClassType)type).getPyClass() : null;
|
||||
if (cls != null && PyABCUtil.isSubclass(cls, PyNames.MAPPING)) {
|
||||
return res;
|
||||
}
|
||||
if (type instanceof PySubscriptableType) {
|
||||
res = ((PySubscriptableType)type).getElementType(indexExpression, context);
|
||||
}
|
||||
else if (type instanceof PyCollectionType) {
|
||||
res = ((PyCollectionType)type).getElementType(context);
|
||||
}
|
||||
}
|
||||
}
|
||||
members.add(res);
|
||||
}
|
||||
return res;
|
||||
return PyUnionType.union(members);
|
||||
}
|
||||
|
||||
@Override
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
def test():
|
||||
xs = map(lambda x: x + 1, [1, 2, 3])
|
||||
print('foo' + <warning descr="Expected type 'str | unicode', got 'int' instead">xs[0]</warning>)
|
||||
ys = map(str, iter([1, 2, 3]))
|
||||
print(1 + <warning descr="Expected type 'Number', got 'str' instead">ys[0]</warning>, 'bar' + ys[1])
|
||||
print('foo' + xs[0]) # Can be a str since map returns list[V] | str | unicode
|
||||
ys = map(tuple, iter([1, 2, 3]))
|
||||
print(1 + <warning descr="Expected type 'Number', got 'tuple | str | unicode' instead">ys[0]</warning>, 'bar' + ys[1])
|
||||
|
||||
@@ -621,12 +621,25 @@ public class PyTypeTest extends PyTestCase {
|
||||
}
|
||||
|
||||
public void testFunctionTypeAsUnificationArgument() {
|
||||
doTest("int",
|
||||
doTest("list[int] | str | unicode",
|
||||
"def map2(f, xs):\n" +
|
||||
" '''\n" +
|
||||
" :type f: (T) -> V | None\n" +
|
||||
" :type xs: collections.Iterable[T] | bytes | unicode\n" +
|
||||
" :rtype: list[V] | bytes | unicode\n" +
|
||||
" :type xs: collections.Iterable[T] | str | unicode\n" +
|
||||
" :rtype: list[V] | str | unicode\n" +
|
||||
" '''\n" +
|
||||
" pass\n" +
|
||||
"\n" +
|
||||
"expr = map2(lambda x: 10, ['1', '2', '3'])\n");
|
||||
}
|
||||
|
||||
public void testFunctionTypeAsUnificationArgumentWithSubscription() {
|
||||
doTest("int | str | unicode",
|
||||
"def map2(f, xs):\n" +
|
||||
" '''\n" +
|
||||
" :type f: (T) -> V | None\n" +
|
||||
" :type xs: collections.Iterable[T] | str | unicode\n" +
|
||||
" :rtype: list[V] | str | unicode\n" +
|
||||
" '''\n" +
|
||||
" pass\n" +
|
||||
"\n" +
|
||||
@@ -873,6 +886,60 @@ public class PyTypeTest extends PyTestCase {
|
||||
"expr = np.ones(10) * 2\n");
|
||||
}
|
||||
|
||||
public void testUnionTypeAttributeOfDifferentTypes() {
|
||||
doTest("list | int",
|
||||
"class Foo:\n" +
|
||||
" x = []\n" +
|
||||
"\n" +
|
||||
"class Bar:\n" +
|
||||
" x = 42\n" +
|
||||
"\n" +
|
||||
"def f(c):\n" +
|
||||
" o = Foo() if c else Bar()\n" +
|
||||
" expr = o.x\n");
|
||||
}
|
||||
|
||||
// PY-11364
|
||||
public void testUnionTypeAttributeCallOfDifferentTypes() {
|
||||
doTest("C1 | C2",
|
||||
"class C1:\n" +
|
||||
" def foo(self):\n" +
|
||||
" return self\n" +
|
||||
"\n" +
|
||||
"class C2:\n" +
|
||||
" def foo(self):\n" +
|
||||
" return self\n" +
|
||||
"\n" +
|
||||
"def f():\n" +
|
||||
" '''\n" +
|
||||
" :rtype: C1 | C2\n" +
|
||||
" '''\n" +
|
||||
" pass\n" +
|
||||
"\n" +
|
||||
"expr = f().foo()\n");
|
||||
}
|
||||
|
||||
// PY-12862
|
||||
public void testUnionTypeAttributeSubscriptionOfDifferentTypes() {
|
||||
doTest("C1 | C2",
|
||||
"class C1:\n" +
|
||||
" def __getitem__(self, item):\n" +
|
||||
" return self\n" +
|
||||
"\n" +
|
||||
"class C2:\n" +
|
||||
" def __getitem__(self, item):\n" +
|
||||
" return self\n" +
|
||||
"\n" +
|
||||
"def f():\n" +
|
||||
" '''\n" +
|
||||
" :rtype: C1 | C2\n" +
|
||||
" '''\n" +
|
||||
" pass\n" +
|
||||
"\n" +
|
||||
"expr = f()[0]\n" +
|
||||
"print(expr)\n");
|
||||
}
|
||||
|
||||
private static TypeEvalContext getTypeEvalContext(@NotNull PyExpression element) {
|
||||
return TypeEvalContext.userInitiated(element.getProject(), element.getContainingFile()).withTracing();
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user