diff --git a/python/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java b/python/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java index 44e1b49b56cc..0288dbea7055 100644 --- a/python/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java +++ b/python/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java @@ -410,10 +410,10 @@ public class PyClassTypeImpl extends UserDataHolderBase implements PyClassType { if (Objects.equals(t1, t2)) { return 0; } - else if (t2 == null || t1 != null && Sets.newHashSet(t1.getAncestorTypes(context)).contains(t2)) { + else if (t2 == null || t1 != null && t1.getAncestorTypes(context).contains(t2)) { return 1; } - else if (t1 == null || Sets.newHashSet(t2.getAncestorTypes(context)).contains(t1)) { + else if (t1 == null || t2.getAncestorTypes(context).contains(t1)) { return -1; } else { @@ -466,19 +466,15 @@ public class PyClassTypeImpl extends UserDataHolderBase implements PyClassType { @Nullable @Override public List getParameters(@NotNull TypeEvalContext context) { - if (isDefinition()) { - List params = getParametersOfMethod(PyNames.INIT, context); - if (params == null) { - // TODO better way to resolve the constructor method here - params = getParametersOfMethod(PyNames.NEW, context); - } - if (params != null) { - // Skip "self" for __init__ and "cls" for __new__ - return params.subList(1, params.size()); - } - return null; - } - return getParametersOfMethod(PyNames.CALL, context); + final List methodNames = isDefinition() ? Arrays.asList(PyNames.INIT, PyNames.NEW) : Collections.singletonList(PyNames.CALL); + + return StreamEx + .of(methodNames) + .map(name -> getParametersOfMethod(name, context)) + .findFirst(Objects::nonNull) + // Skip "self" for __init__/__call__ and "cls" for __new__ + .map(parameters -> ContainerUtil.subList(parameters, 1)) + .orElse(null); } @Nullable diff --git a/python/testData/inspections/PyTypeCheckerInspection/CallableInstanceAgainstCallable.py b/python/testData/inspections/PyTypeCheckerInspection/CallableInstanceAgainstCallable.py new file mode 100644 index 000000000000..6f53981d0b05 --- /dev/null +++ b/python/testData/inspections/PyTypeCheckerInspection/CallableInstanceAgainstCallable.py @@ -0,0 +1,10 @@ +from typing import Dict + + +class Key: + def __call__(self, obj): + pass + + +def foo(d: Dict[str, int]): + print(sorted(d.items(), key=Key())) diff --git a/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java index 3125a8f86cbd..0f8ef889d30c 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java @@ -451,4 +451,8 @@ public class PyTypeCheckerInspectionTest extends PyInspectionTestCase { public void testClassMetaAttrsAgainstStructural() { runWithLanguageLevel(LanguageLevel.PYTHON30, this::doTest); } + + public void testCallableInstanceAgainstCallable() { + runWithLanguageLevel(LanguageLevel.PYTHON35, this::doTest); + } }