diff --git a/python/src/com/jetbrains/python/codeInsight/stdlib/PyNamedTupleType.java b/python/src/com/jetbrains/python/codeInsight/stdlib/PyNamedTupleType.java index 818790416894..c0b94078f4eb 100644 --- a/python/src/com/jetbrains/python/codeInsight/stdlib/PyNamedTupleType.java +++ b/python/src/com/jetbrains/python/codeInsight/stdlib/PyNamedTupleType.java @@ -19,6 +19,7 @@ import com.intellij.codeInsight.lookup.LookupElementBuilder; import com.intellij.psi.PsiElement; import com.intellij.util.ArrayUtil; import com.intellij.util.ProcessingContext; +import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.psi.AccessDirection; import com.jetbrains.python.psi.PyCallSiteExpression; import com.jetbrains.python.psi.PyClass; @@ -155,6 +156,17 @@ public class PyNamedTupleType extends PyClassTypeImpl implements PyCallableType return Collections.unmodifiableList(myFields); } + @Override + public boolean isCallable() { + return myDefinitionLevel == DefinitionLevel.NEW_TYPE; + } + + @Nullable + @Override + public List getParameters(@NotNull TypeEvalContext context) { + return isCallable() ? ContainerUtil.map(myFields, field -> new PyCallableParameterImpl(field, null)) : null; + } + public enum DefinitionLevel { AS_SUPERCLASS, diff --git a/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java index be3e84a30e6e..089b0e697cda 100644 --- a/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java @@ -19,6 +19,7 @@ import com.google.common.collect.ImmutableSet; import com.intellij.openapi.extensions.Extensions; import com.intellij.openapi.util.Ref; import com.intellij.psi.PsiElement; +import com.intellij.psi.ResolveResult; import com.intellij.psi.util.QualifiedName; import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.PyNames; @@ -28,19 +29,19 @@ import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.impl.PyBuiltinCache; import com.jetbrains.python.psi.impl.PyCallExpressionHelper; +import com.jetbrains.python.psi.impl.PyCallExpressionNavigator; import com.jetbrains.python.psi.impl.PyTypeProvider; import com.jetbrains.python.psi.impl.stubs.PyNamedTupleStubImpl; +import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.resolve.PyResolveImportUtil; import com.jetbrains.python.psi.stubs.PyNamedTupleStub; import com.jetbrains.python.psi.stubs.PyTargetExpressionStub; import com.jetbrains.python.psi.types.*; +import one.util.streamex.StreamEx; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; -import java.util.Arrays; -import java.util.List; -import java.util.Map; -import java.util.Set; +import java.util.*; import static com.jetbrains.python.psi.PyUtil.as; @@ -94,6 +95,15 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { return PyBuiltinCache.getInstance(referenceExpression).getBoolType(); } } + + final PyCallableType typeFromNTInheritorInitializing = Optional + .ofNullable(getTypeFromCollectionsNTInheritorInitializing(referenceExpression, context)) + .orElseGet(() -> getTypeFromTypingNTInheritorInitializing(referenceExpression, context)); + + if (typeFromNTInheritorInitializing != null) { + return typeFromNTInheritorInitializing; + } + return null; } @@ -157,6 +167,46 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { return null; } + @Nullable + private static PyCallableType getTypeFromCollectionsNTInheritorInitializing(@NotNull PyReferenceExpression referenceExpression, + @NotNull TypeEvalContext context) { + if (PyCallExpressionNavigator.getPyCallExpressionByCallee(referenceExpression) == null) { + return null; + } + + final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); + final ResolveResult[] resolveResults = referenceExpression.getReference(resolveContext).multiResolve(false); + + return StreamEx + .of(PyUtil.filterTopPriorityResults(resolveResults)) + .select(PyTypedElement.class) + .map(context::getType) + .select(PyClassLikeType.class) + .flatCollection(type -> type.getAncestorTypes(context)) + .select(PyNamedTupleType.class) + .findFirst() + .orElse(null); + } + + @Nullable + private static PyCallableType getTypeFromTypingNTInheritorInitializing(@NotNull PyReferenceExpression referenceExpression, + @NotNull TypeEvalContext context) { + if (PyCallExpressionNavigator.getPyCallExpressionByCallee(referenceExpression) == null) { + return null; + } + + final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); + final ResolveResult[] resolveResults = referenceExpression.getReference(resolveContext).multiResolve(false); + + return StreamEx + .of(PyUtil.filterTopPriorityResults(resolveResults)) + .select(PyTypedElement.class) + .map(context::getType) + .select(PyNamedTupleType.class) + .findFirst() + .orElse(null); + } + @Nullable @Override public Ref getCallType(@NotNull PyFunction function, @Nullable PyCallSiteExpression callSite, @NotNull TypeEvalContext context) { @@ -178,7 +228,7 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { else if ("object.__new__".equals(qname) && callSite instanceof PyCallExpression) { final PyExpression firstArgument = ((PyCallExpression)callSite).getArgument(0, PyExpression.class); final PyClassLikeType classLikeType = as(firstArgument != null ? context.getType(firstArgument) : null, PyClassLikeType.class); - return classLikeType != null ? Ref.create(classLikeType.toInstance()) : null; + return classLikeType != null ? Ref.create(classLikeType.toInstance()) : null; } } @@ -186,7 +236,8 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { } @Nullable - private static Ref getTupleMultiplicationResultType(@NotNull PyBinaryExpression multiplication, @NotNull TypeEvalContext context) { + private static Ref getTupleMultiplicationResultType(@NotNull PyBinaryExpression multiplication, + @NotNull TypeEvalContext context) { final PyTupleType leftTupleType = as(context.getType(multiplication.getLeftExpression()), PyTupleType.class); if (leftTupleType == null) { return null; @@ -265,7 +316,9 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { @Nullable @Override - public PyType getContextManagerVariableType(@NotNull PyClass contextManager, @NotNull PyExpression withExpression, @NotNull TypeEvalContext context) { + public PyType getContextManagerVariableType(@NotNull PyClass contextManager, + @NotNull PyExpression withExpression, + @NotNull TypeEvalContext context) { if ("contextlib.closing".equals(contextManager.getQualifiedName()) && withExpression instanceof PyCallExpression) { PyExpression closee = ((PyCallExpression)withExpression).getArgument(0, PyExpression.class); if (closee != null) { @@ -289,7 +342,8 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { if (stub != null) { return getNamedTupleTypeFromStub(target, stub.getCustomStub(PyNamedTupleStub.class), PyNamedTupleType.DefinitionLevel.NEW_TYPE); - } else { + } + else { return getNamedTupleTypeFromAST(target, context, PyNamedTupleType.DefinitionLevel.NEW_TYPE); } } @@ -338,8 +392,10 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { return null; } - final PyClass tupleClass = as(PyResolveImportUtil.resolveTopLevelMember(QualifiedName.fromDottedString(PyTypingTypeProvider.NAMEDTUPLE), - PyResolveImportUtil.fromFoothold(referenceTarget)), PyClass.class); + final PsiElement typingNT = PyResolveImportUtil.resolveTopLevelMember(QualifiedName.fromDottedString(PyTypingTypeProvider.NAMEDTUPLE), + PyResolveImportUtil.fromFoothold(referenceTarget)); + + final PyClass tupleClass = as(typingNT, PyClass.class); if (tupleClass == null) { return null; } diff --git a/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java b/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java index 950e3c554557..31c47d433f59 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java +++ b/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java @@ -16,6 +16,7 @@ package com.jetbrains.python.psi.impl; import com.intellij.codeInsight.completion.CompletionUtil; +import com.intellij.openapi.extensions.Extensions; import com.intellij.openapi.util.Pair; import com.intellij.openapi.util.Ref; import com.intellij.psi.PsiElement; @@ -118,6 +119,13 @@ public class PyCallExpressionHelper { public static List multiResolveRatedCallee(@NotNull PyCallExpression call, @NotNull PyResolveContext resolveContext, int implicitOffset) { + final PyExpression callee = call.getCallee(); + + final List calleesFromProviders = getCalleesFromProviders(callee, resolveContext.getTypeEvalContext()); + if (calleesFromProviders != null) { + return calleesFromProviders; + } + final TypeEvalContext context = resolveContext.getTypeEvalContext(); final List ratedMarkedCallees = new ArrayList<>(); @@ -133,7 +141,9 @@ public class PyCallExpressionHelper { return forEveryScopeTakeOverloadsOtherwiseImplementations(ratedMarkedCallees, PyCallExpression.PyRatedMarkedCallee::getElement, context) // while clarifying resolve results we could get duplicate callable types so we have to group them and select result with highest rate - .collect(Collectors.groupingBy(markedCallee -> markedCallee.getMarkedCallee().getCallableType(), LinkedHashMap::new, Collectors.toList())) + .collect( + Collectors.groupingBy(markedCallee -> markedCallee.getMarkedCallee().getCallableType(), LinkedHashMap::new, Collectors.toList()) + ) .entrySet() .stream() .map(entry -> entry.getValue().stream().max(Comparator.comparingInt(PyCallExpression.PyRatedMarkedCallee::getRate)).orElse(null)) @@ -141,6 +151,27 @@ public class PyCallExpressionHelper { .collect(Collectors.toList()); } + @Nullable + private static List getCalleesFromProviders(@Nullable PyExpression callee, @NotNull TypeEvalContext context) { + if (callee instanceof PyReferenceExpression) { + final PyReferenceExpression referenceExpression = (PyReferenceExpression)callee; + + final List callees = StreamEx + .of(Extensions.getExtensions(PyTypeProvider.EP_NAME)) + .map(provider -> provider.getReferenceExpressionType(referenceExpression, context)) + .select(PyCallableType.class) + .map(type -> new PyCallExpression.PyMarkedCallee(type, null, null, 0, false)) + .map(markedCallee -> new PyCallExpression.PyRatedMarkedCallee(markedCallee, RatedResolveResult.RATE_NORMAL)) + .toList(); + + if (!callees.isEmpty()) { + return callees; + } + } + + return null; + } + @NotNull private static List multiResolveCallee(@Nullable PyExpression callee, @NotNull PyResolveContext resolveContext) { diff --git a/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java b/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java index 448a53da9856..78ce47f4d11c 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java @@ -19,9 +19,11 @@ import com.intellij.codeInsight.completion.CompletionUtil; import com.intellij.lang.ASTNode; import com.intellij.navigation.ItemPresentation; import com.intellij.openapi.util.Comparing; +import com.intellij.openapi.util.Condition; import com.intellij.openapi.util.NotNullLazyValue; import com.intellij.openapi.util.Ref; import com.intellij.psi.*; +import com.intellij.psi.scope.BaseScopeProcessor; import com.intellij.psi.scope.PsiScopeProcessor; import com.intellij.psi.search.LocalSearchScope; import com.intellij.psi.search.SearchScope; @@ -38,6 +40,8 @@ import com.jetbrains.python.PythonDialectsTokenSetProvider; import com.jetbrains.python.codeInsight.controlflow.ControlFlowCache; import com.jetbrains.python.codeInsight.controlflow.ScopeOwner; import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil; +import com.jetbrains.python.codeInsight.stdlib.PyNamedTupleType; +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; import com.jetbrains.python.documentation.docstrings.DocStringUtil; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.impl.stubs.PyClassElementType; @@ -127,7 +131,55 @@ public class PyClassImpl extends PyBaseElementImpl implements PyCla @Override public PyType getType(@NotNull TypeEvalContext context, @NotNull TypeEvalContext.Key key) { - return new PyClassTypeImpl(this, true); + return Optional + .ofNullable(getTypeForTypingNTInheritor(context)) + .orElseGet(() -> new PyClassTypeImpl(this, true)); + } + + @Nullable + private PyNamedTupleType getTypeForTypingNTInheritor(@NotNull TypeEvalContext context) { + final Condition isTypingNT = + type -> + type != null && + !(type instanceof PyNamedTupleType) && + PyTypingTypeProvider.NAMEDTUPLE.equals(type.getClassQName()); + + if (ContainerUtil.exists(getSuperClassTypes(context), isTypingNT)) { + final String name = getName(); + if (name != null) { + final PsiElement typingNT = resolveTopLevelMember(QualifiedName.fromDottedString(PyTypingTypeProvider.NAMEDTUPLE), + fromFoothold(this)); + + final PyClass tupleClass = as(typingNT, PyClass.class); + if (tupleClass != null) { + final Set fields = new TreeSet<>(Comparator.comparingInt(PyTargetExpression::getTextOffset)); + + processClassLevelDeclarations( + new BaseScopeProcessor() { + @Override + public boolean execute(@NotNull PsiElement element, @NotNull ResolveState substitutor) { + if (element instanceof PyTargetExpression) { + final PyTargetExpression target = (PyTargetExpression)element; + if (target.getAnnotation() != null) { + fields.add(target); + } + } + + return true; + } + } + ); + + return new PyNamedTupleType(tupleClass, + this, + name, + ContainerUtil.map(fields, PyTargetExpression::getName), + PyNamedTupleType.DefinitionLevel.NEW_TYPE); + } + } + } + + return null; } private class NewStyleCachedValueProvider implements ParameterizedCachedValueProvider { diff --git a/python/testData/inspections/PyArgumentListInspection/initializingCollectionsNamedTuple.py b/python/testData/inspections/PyArgumentListInspection/initializingCollectionsNamedTuple.py new file mode 100644 index 000000000000..d1784c11cbdb --- /dev/null +++ b/python/testData/inspections/PyArgumentListInspection/initializingCollectionsNamedTuple.py @@ -0,0 +1,47 @@ +from collections import namedtuple + + +MyTup1 = namedtuple(bar='') +MyTup2 = namedtuple("MyTup2", "bar baz") + + +class MyTup3(namedtuple(bar='')): + pass + + +class MyTup4(namedtuple("MyTup4", "bar baz")): + pass + + +# empty +MyTup2() + +# one +MyTup2(bar='') +MyTup2(baz='') + +# two +MyTup2('', '') +MyTup2(bar='', baz='') +MyTup2(baz='', bar='') + +# three +MyTup2(bar='', baz='', foo='') +MyTup2('', '', '') + + +# empty +MyTup4() + +# one +MyTup4(bar='') +MyTup4(baz='') + +# two +MyTup4('', '') +MyTup4(bar='', baz='') +MyTup4(baz='', bar='') + +# three +MyTup4(bar='', baz='', foo='') +MyTup4('', '', '') diff --git a/python/testData/inspections/PyArgumentListInspection/initializingTypingNamedTuple.py b/python/testData/inspections/PyArgumentListInspection/initializingTypingNamedTuple.py new file mode 100644 index 000000000000..e7abe15e59d4 --- /dev/null +++ b/python/testData/inspections/PyArgumentListInspection/initializingTypingNamedTuple.py @@ -0,0 +1,108 @@ +import typing + + +MyTup1 = typing.NamedTuple(bar='') +MyTup2 = typing.NamedTuple("MyTup2", bar=int, baz=str) +MyTup3 = typing.NamedTuple("MyTup2", [("bar", int), ("baz", str)]) + + +class MyTup4(typing.NamedTuple): + bar: int + baz: str + + +class MyTup5(typing.NamedTuple): + bar: int + baz: str + foo = 5 + + +class MyTup6(typing.NamedTuple): + bar: int + baz: str + foo: int + + +# empty +MyTup2() + +# one +MyTup2(bar='') +MyTup2(baz='') + +# two +MyTup2('', '') +MyTup2(bar='', baz='') +MyTup2(baz='', bar='') + +# three +MyTup2(bar='', baz='', foo='') +MyTup2('', '', '') + + +# empty +MyTup3() + +# one +MyTup3(bar='') +MyTup3(baz='') + +# two +MyTup3('', '') +MyTup3(bar='', baz='') +MyTup3(baz='', bar='') + +# three +MyTup3(bar='', baz='', foo='') +MyTup3('', '', '') + + +# empty +MyTup4() + +# one +MyTup4(bar='') +MyTup4(baz='') + +# two +MyTup4('', '') +MyTup4(bar='', baz='') +MyTup4(baz='', bar='') + +# three +MyTup4(bar='', baz='', foo='') +MyTup4('', '', '') + + +# empty +MyTup5() + +# one +MyTup5(bar='') +MyTup5(baz='') + +# two +MyTup5('', '') +MyTup5(bar='', baz='') +MyTup5(baz='', bar='') + +# three +MyTup5(bar='', baz='', foo='') +MyTup5('', '', '') + + +# empty +MyTup6() + +# one +MyTup6(bar='') +MyTup6(baz='') + +# two +MyTup6('', '') +MyTup6(bar='', baz='') +MyTup6(baz='', bar='') + +# three +MyTup6(bar='', baz='', foo='') +MyTup6('', '', '') diff --git a/python/testData/paramInfo/InitializingCollectionsNamedTuple.py b/python/testData/paramInfo/InitializingCollectionsNamedTuple.py new file mode 100644 index 000000000000..fbdb7f42ae99 --- /dev/null +++ b/python/testData/paramInfo/InitializingCollectionsNamedTuple.py @@ -0,0 +1,12 @@ +from collections import namedtuple + + +MyTup1 = namedtuple("MyTup1", "bar baz") + + +class MyTup2(namedtuple("MyTup2", "bar baz")): + pass + + +MyTup1() +MyTup2() \ No newline at end of file diff --git a/python/testData/paramInfo/InitializingTypingNamedTuple.py b/python/testData/paramInfo/InitializingTypingNamedTuple.py new file mode 100644 index 000000000000..dc02b83641ae --- /dev/null +++ b/python/testData/paramInfo/InitializingTypingNamedTuple.py @@ -0,0 +1,29 @@ +import typing + + +MyTup2 = typing.NamedTuple("MyTup2", bar=int, baz=str) +MyTup3 = typing.NamedTuple("MyTup2", [("bar", int), ("baz", str)]) + + +class MyTup4(typing.NamedTuple): + bar: int + baz: str + + +class MyTup5(typing.NamedTuple): + bar: int + baz: str + foo = 5 + + +class MyTup6(typing.NamedTuple): + bar: int + baz: str + foo: int + + +MyTup2() +MyTup3() +MyTup4() +MyTup5() +MyTup6() diff --git a/python/testSrc/com/jetbrains/python/PyParameterInfoTest.java b/python/testSrc/com/jetbrains/python/PyParameterInfoTest.java index 5164fd8a81c3..4da351d6125d 100644 --- a/python/testSrc/com/jetbrains/python/PyParameterInfoTest.java +++ b/python/testSrc/com/jetbrains/python/PyParameterInfoTest.java @@ -32,6 +32,7 @@ import com.jetbrains.python.fixtures.LightMarkedTestCase; import com.jetbrains.python.psi.LanguageLevel; import com.jetbrains.python.psi.PyArgumentList; import com.jetbrains.python.psi.PyCallExpression; +import one.util.streamex.StreamEx; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -650,6 +651,44 @@ public class PyParameterInfoTest extends LightMarkedTestCase { ); } + // PY-22249 + public void testInitializingCollectionsNamedTuple() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> { + final Map test = loadTest(2); + + for (int offset : StreamEx.of(test.values()).map(PsiElement::getTextOffset)) { + final List texts = Collections.singletonList("bar, baz"); + final List highlighted = Collections.singletonList(new String[]{"bar, "}); + + feignCtrlP(offset).check(texts, highlighted, Collections.singletonList(ArrayUtil.EMPTY_STRING_ARRAY)); + } + } + ); + } + + public void testInitializingTypingNamedTuple() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> { + final Map test = loadTest(5); + + for (int offset : StreamEx.of(1, 2, 3, 4).map(number -> test.get("").getTextOffset())) { + final List texts = Collections.singletonList("bar, baz"); + final List highlighted = Collections.singletonList(new String[]{"bar, "}); + + feignCtrlP(offset).check(texts, highlighted, Collections.singletonList(ArrayUtil.EMPTY_STRING_ARRAY)); + } + + final List texts = Collections.singletonList("bar, baz, foo"); + final List highlighted = Collections.singletonList(new String[]{"bar, "}); + + feignCtrlP(test.get("").getTextOffset()).check(texts, highlighted, Collections.singletonList(ArrayUtil.EMPTY_STRING_ARRAY)); + } + ); + } + /** * Imitates pressing of Ctrl+P; fails if results are not as expected. * @param offset offset of 'cursor' where Ctrl+P is pressed. diff --git a/python/testSrc/com/jetbrains/python/inspections/PyArgumentListInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyArgumentListInspectionTest.java index 3d3e9ca16b8a..1d2fe98221fb 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyArgumentListInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyArgumentListInspectionTest.java @@ -273,6 +273,16 @@ public class PyArgumentListInspectionTest extends PyTestCase { doTest(); } + // PY-19293, PY-22102 + public void testInitializingTypingNamedTuple() { + runWithLanguageLevel(LanguageLevel.PYTHON36, this::doTest); + } + + // PY-4344, PY-8422, PY-22269, PY-22740 + public void testInitializingCollectionsNamedTuple() { + doTest(); + } + // PY-22971 public void testOverloadsAndImplementationInClass() { runWithLanguageLevel(LanguageLevel.PYTHON35, this::doTest);