Full type support for namedtuple._replace (PY-27148)

Provide correct parameters for typed and untyped NT.
Clarify return type for untyped case.
This commit is contained in:
Semyon Proshev
2018-01-17 22:06:35 +03:00
parent a81ded8e78
commit bafc005066
13 changed files with 418 additions and 68 deletions
@@ -157,7 +157,7 @@ public abstract class PyElementGenerator {
* @param qualifier from where {@code name} will be imported (module name)
* @param name text of the reference in import element
* @param alias optional alias for {@code as alias} part
* @return created {@link com.jetbrains.python.psi.PyFromImportStatement}
* @return created {@link PyFromImportStatement}
*/
@NotNull
public abstract PyFromImportStatement createFromImportStatement(@NotNull LanguageLevel languageLevel,
@@ -171,10 +171,16 @@ public abstract class PyElementGenerator {
* @param languageLevel language level for created element
* @param name text of the reference in import element (module name)
* @param alias optional alias for {@code as alias} part
* @return created {@link com.jetbrains.python.psi.PyImportStatement}
* @return created {@link PyImportStatement}
*/
@NotNull
public abstract PyImportStatement createImportStatement(@NotNull LanguageLevel languageLevel,
@NotNull String name,
@Nullable String alias);
@NotNull
public abstract PyNoneLiteralExpression createEllipsis();
@NotNull
public abstract PySingleStarParameter createSingleStarParameter();
}
@@ -33,11 +33,14 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType {
@NotNull
private final DefinitionLevel myDefinitionLevel;
private final boolean myTyped;
public PyNamedTupleType(@NotNull PyClass tupleClass,
@NotNull PsiElement declaration,
@NotNull String name,
@NotNull Map<String, FieldTypeAndDefaultValue> fields,
@NotNull DefinitionLevel definitionLevel) {
@NotNull DefinitionLevel definitionLevel,
boolean typed) {
super(tupleClass,
Collections.unmodifiableList(ContainerUtil.map(fields.values(), typeAndValue -> typeAndValue.getType())),
false,
@@ -47,6 +50,7 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType {
myFields = Collections.unmodifiableMap(fields);
myName = name;
myDefinitionLevel = definitionLevel;
myTyped = typed;
}
@Override
@@ -74,11 +78,11 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType {
@Override
public PyType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite) {
if (myDefinitionLevel == DefinitionLevel.NT_FUNCTION) {
return new PyNamedTupleType(myClass, myDeclaration, myName, myFields, DefinitionLevel.NEW_TYPE);
return new PyNamedTupleType(myClass, myDeclaration, myName, myFields, DefinitionLevel.NEW_TYPE, myTyped);
}
else if (myDefinitionLevel == DefinitionLevel.NEW_TYPE) {
final Map<String, FieldTypeAndDefaultValue> fields = takeFieldsTypesFromCallSiteIfNeeded(context, callSite);
return new PyNamedTupleType(myClass, myDeclaration, myName, fields, DefinitionLevel.INSTANCE);
return new PyNamedTupleType(myClass, myDeclaration, myName, fields, DefinitionLevel.INSTANCE, myTyped);
}
return null;
@@ -88,7 +92,7 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType {
@Override
public PyClassType toInstance() {
return myDefinitionLevel == DefinitionLevel.NEW_TYPE
? new PyNamedTupleType(myClass, myDeclaration, myName, myFields, DefinitionLevel.INSTANCE)
? new PyNamedTupleType(myClass, myDeclaration, myName, myFields, DefinitionLevel.INSTANCE, myTyped)
: this;
}
@@ -97,7 +101,7 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType {
public PyClassLikeType toClass() {
return myDefinitionLevel == DefinitionLevel.INSTANCE
? this
: new PyNamedTupleType(myClass, myDeclaration, myName, myFields, DefinitionLevel.NEW_TYPE);
: new PyNamedTupleType(myClass, myDeclaration, myName, myFields, DefinitionLevel.NEW_TYPE, myTyped);
}
@Override
@@ -149,10 +153,33 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType {
: null;
}
public boolean isTyped() {
return myTyped;
}
@NotNull
public PyNamedTupleType clarifyFields(@NotNull Map<String, PyType> fieldNameToType) {
if (!myTyped) {
final LinkedHashMap<String, FieldTypeAndDefaultValue> newFields = new LinkedHashMap<>(myFields);
for (Map.Entry<String, PyType> entry : fieldNameToType.entrySet()) {
final String fieldName = entry.getKey();
if (newFields.containsKey(fieldName)) {
newFields.put(fieldName, new FieldTypeAndDefaultValue(entry.getValue(), null));
}
}
return new PyNamedTupleType(myClass, myDeclaration, myName, newFields, myDefinitionLevel, false);
}
return this;
}
@NotNull
private Map<String, FieldTypeAndDefaultValue> takeFieldsTypesFromCallSiteIfNeeded(@NotNull TypeEvalContext context,
@NotNull PyCallSiteExpression callSite) {
if (StreamEx.of(getElementTypes()).allMatch(Objects::isNull)) {
if (!myTyped) {
final List<PyExpression> arguments = callSite.getArguments(null);
if (arguments.size() == myFields.size()) {
@@ -24,6 +24,7 @@ import com.jetbrains.python.psi.types.TypeEvalContext
class PyStdlibOverridingTypeProvider : PyTypeProviderBase(), PyOverridingTypeProvider {
override fun getReferenceType(referenceTarget: PsiElement, context: TypeEvalContext, anchor: PsiElement?): PyType? {
return PyStdlibTypeProvider.getNamedTupleTypeForResolvedCallee(referenceTarget, context, anchor)
return PyStdlibTypeProvider.getNamedTupleTypeForResolvedCallee(referenceTarget, context, anchor) ?:
PyStdlibTypeProvider.getNamedTupleReplaceType(referenceTarget, context, anchor)
}
}
@@ -26,6 +26,7 @@ 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;
@@ -94,6 +95,11 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase {
return namedTupleTypeForCallee;
}
final PyCallableType namedTupleReplaceType = getNamedTupleReplaceType(referenceExpression, context);
if (namedTupleReplaceType != null) {
return namedTupleReplaceType;
}
return null;
}
@@ -122,6 +128,23 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase {
return null;
}
@Nullable
static PyCallableType getNamedTupleReplaceType(@NotNull PsiElement referenceTarget,
@NotNull TypeEvalContext context,
@Nullable PsiElement anchor) {
if (referenceTarget instanceof PyFunction && anchor instanceof PyCallExpression) {
final PyClass containingClass = ((PyFunction)referenceTarget).getContainingClass();
if (containingClass != null && PyTypingTypeProvider.NAMEDTUPLE.equals(containingClass.getQualifiedName())) {
final PyExpression callee = ((PyCallExpression)anchor).getCallee();
if (callee instanceof PyReferenceExpression) {
return getNamedTupleReplaceType((PyReferenceExpression)callee, context);
}
}
}
return null;
}
@Nullable
private static PyType getEnumType(@NotNull PsiElement referenceTarget, @NotNull TypeEvalContext context,
@Nullable PsiElement anchor) {
@@ -252,8 +275,7 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase {
final PyClassLikeType classLikeType = as(firstArgument != null ? context.getType(firstArgument) : null, PyClassLikeType.class);
return classLikeType != null ? Ref.create(classLikeType.toInstance()) : null;
}
else if (callSite != null &&
ArrayUtil.contains(qname, PyTypingTypeProvider.NAMEDTUPLE + "._make", PyTypingTypeProvider.NAMEDTUPLE + "._replace")) {
else if (callSite != null && qname.equals(PyTypingTypeProvider.NAMEDTUPLE + "._make")) {
final PyExpression receiver = callSite.getReceiver(function);
if (receiver != null) {
final PyType receiverType = context.getType(receiver);
@@ -357,6 +379,16 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase {
return type instanceof PyClassLikeType && ContainerUtil.exists(((PyClassLikeType)type).getAncestorTypes(context), isNT);
}
public static boolean isTypingNamedTupleDirectInheritor(@NotNull PyClass cls, @NotNull TypeEvalContext context) {
final Condition<PyClassLikeType> isTypingNT =
type ->
type != null &&
!(type instanceof PyNamedTupleType) &&
PyTypingTypeProvider.NAMEDTUPLE.equals(type.getClassQName());
return ContainerUtil.exists(cls.getSuperClassTypes(context), isTypingNT);
}
@Nullable
@Override
public PyType getContextManagerVariableType(@NotNull PyClass contextManager,
@@ -409,14 +441,48 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase {
}
@Nullable
private static PyNamedTupleType getNamedTupleTypeForTypingNTInheritorAsCallee(@NotNull PyClass cls, @NotNull TypeEvalContext context) {
final Condition<PyClassLikeType> isTypingNT =
type ->
type != null &&
!(type instanceof PyNamedTupleType) &&
PyTypingTypeProvider.NAMEDTUPLE.equals(type.getClassQName());
private static PyCallableType getNamedTupleReplaceType(@NotNull PyReferenceExpression referenceExpression,
@NotNull TypeEvalContext context) {
final PyCallExpression call = PyCallExpressionNavigator.getPyCallExpressionByCallee(referenceExpression);
if (call == null) {
return null;
}
if (ContainerUtil.exists(cls.getSuperClassTypes(context), isTypingNT)) {
final PyExpression qualifier = referenceExpression.getQualifier();
if (qualifier != null && "_replace".equals(referenceExpression.getReferencedName())) {
final PyType qualifierType = context.getType(qualifier);
if (qualifierType instanceof PyClassLikeType) {
final PyNamedTupleType namedTupleType = StreamEx
.of(qualifierType)
.append(((PyClassLikeType)qualifierType).getSuperClassTypes(context))
.select(PyNamedTupleType.class)
.findFirst()
.orElse(null);
if (namedTupleType != null) {
if (namedTupleType.isTyped()) {
return createTypedNamedTupleReplaceType(referenceExpression, namedTupleType.getFields(), qualifierType);
}
else {
return createUntypedNamedTupleReplaceType(call, namedTupleType.getFields(), qualifierType, context);
}
}
if (qualifierType instanceof PyClassType) {
final PyClass cls = ((PyClassType)qualifierType).getPyClass();
if (isTypingNamedTupleDirectInheritor(cls, context)) {
return createTypedNamedTupleReplaceType(referenceExpression, collectTypingNTInheritorFields(cls, context), qualifierType);
}
}
}
}
return null;
}
@Nullable
private static PyNamedTupleType getNamedTupleTypeForTypingNTInheritorAsCallee(@NotNull PyClass cls, @NotNull TypeEvalContext context) {
if (isTypingNamedTupleDirectInheritor(cls, context)) {
final String name = cls.getName();
if (name != null) {
final PsiElement typingNT =
@@ -425,35 +491,12 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase {
final PyClass tupleClass = as(typingNT, PyClass.class);
if (tupleClass != null) {
final Set<PyTargetExpression> fields = new TreeSet<>(Comparator.comparingInt(PyTargetExpression::getTextOffset));
cls.processClassLevelDeclarations(
new PsiScopeProcessor() {
@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;
}
}
);
final Collector<PyTargetExpression, ?, LinkedHashMap<String, PyNamedTupleType.FieldTypeAndDefaultValue>> toNTFields =
Collectors.toMap(PyTargetExpression::getName,
field -> new PyNamedTupleType.FieldTypeAndDefaultValue(context.getType(field), field.findAssignedValue()),
(v1, v2) -> v2,
LinkedHashMap::new);
return new PyNamedTupleType(tupleClass,
cls,
name,
fields.stream().collect(toNTFields),
PyNamedTupleType.DefinitionLevel.NEW_TYPE);
collectTypingNTInheritorFields(cls, context),
PyNamedTupleType.DefinitionLevel.NEW_TYPE,
true);
}
}
}
@@ -491,11 +534,14 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase {
return null;
}
final Map<String, Optional<String>> fields = stub.getFields();
return new PyNamedTupleType(tupleClass,
referenceTarget,
stub.getName(),
parseNamedTupleFields(referenceTarget, stub.getFields(), context),
definitionLevel);
parseNamedTupleFields(referenceTarget, fields, context),
definitionLevel,
ContainerUtil.find(fields.values(), Optional::isPresent) != null);
}
@Nullable
@@ -509,6 +555,92 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase {
return null;
}
@NotNull
private static PyCallableType createTypedNamedTupleReplaceType(@NotNull PsiElement anchor,
@NotNull Map<String, PyNamedTupleType.FieldTypeAndDefaultValue> fields,
@NotNull PyType resultType) {
final List<PyCallableParameter> parameters = new ArrayList<>();
final PyElementGenerator elementGenerator = PyElementGenerator.getInstance(anchor.getProject());
parameters.add(PyCallableParameterImpl.psi(elementGenerator.createSingleStarParameter()));
final PyNoneLiteralExpression ellipsis = elementGenerator.createEllipsis();
for (Map.Entry<String, PyNamedTupleType.FieldTypeAndDefaultValue> entry : fields.entrySet()) {
final PyExpression value = entry.getValue().getDefaultValue();
parameters.add(PyCallableParameterImpl.nonPsi(entry.getKey(), entry.getValue().getType(), value == null ? ellipsis : value));
}
return new PyCallableTypeImpl(parameters, resultType);
}
@NotNull
private static PyCallableType createUntypedNamedTupleReplaceType(@NotNull PyCallExpression call,
@NotNull Map<String, PyNamedTupleType.FieldTypeAndDefaultValue> fields,
@NotNull PyType resultType,
@NotNull TypeEvalContext context) {
final List<PyCallableParameter> parameters = new ArrayList<>();
final PyElementGenerator elementGenerator = PyElementGenerator.getInstance(call.getProject());
parameters.add(PyCallableParameterImpl.psi(elementGenerator.createSingleStarParameter()));
final PyNoneLiteralExpression ellipsis = elementGenerator.createEllipsis();
for (String name : fields.keySet()) {
parameters.add(PyCallableParameterImpl.nonPsi(name, null, ellipsis));
}
if (resultType instanceof PyNamedTupleType) {
final Map<String, PyType> newFields = new HashMap<>();
for (PyExpression argument : call.getArguments()) {
if (argument instanceof PyKeywordArgument) {
final PyKeywordArgument keywordArgument = (PyKeywordArgument)argument;
final PyExpression value = keywordArgument.getValueExpression();
if (value != null) {
newFields.put(keywordArgument.getKeyword(), context.getType(value));
}
}
}
return new PyCallableTypeImpl(parameters, ((PyNamedTupleType)resultType).clarifyFields(newFields));
}
else {
return new PyCallableTypeImpl(parameters, resultType);
}
}
@NotNull
private static LinkedHashMap<String, PyNamedTupleType.FieldTypeAndDefaultValue> collectTypingNTInheritorFields(@NotNull PyClass cls,
@NotNull TypeEvalContext context) {
final List<PyTargetExpression> fields = new ArrayList<>();
cls.processClassLevelDeclarations(
new PsiScopeProcessor() {
@Override
public boolean execute(@NotNull PsiElement element, @NotNull ResolveState substitutor) {
if (element instanceof PyTargetExpression) {
final PyTargetExpression target = (PyTargetExpression)element;
if (target.getAnnotationValue() != null) {
fields.add(target);
}
}
return true;
}
}
);
final Collector<PyTargetExpression, ?, LinkedHashMap<String, PyNamedTupleType.FieldTypeAndDefaultValue>> toNTFields =
Collectors.toMap(PyTargetExpression::getName,
field -> new PyNamedTupleType.FieldTypeAndDefaultValue(context.getType(field), field.findAssignedValue()),
(v1, v2) -> v2,
LinkedHashMap::new);
return fields.stream().collect(toNTFields);
}
@Nullable
private static PyNamedTupleType getNamedTupleTypeFromAST(@NotNull PyCallExpression expression,
@NotNull TypeEvalContext context,
@@ -8,12 +8,11 @@ import com.intellij.psi.PsiElement
import com.intellij.psi.PsiElementVisitor
import com.intellij.psi.ResolveState
import com.intellij.psi.scope.PsiScopeProcessor
import com.jetbrains.python.codeInsight.stdlib.PyNamedTupleType
import com.jetbrains.python.codeInsight.stdlib.PyStdlibTypeProvider
import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider
import com.jetbrains.python.psi.LanguageLevel
import com.jetbrains.python.psi.PyClass
import com.jetbrains.python.psi.PyTargetExpression
import com.jetbrains.python.psi.types.PyClassLikeType
import com.jetbrains.python.psi.types.TypeEvalContext
import java.util.*
@@ -52,17 +51,12 @@ class PyNamedTupleInspection : PyInspection() {
override fun visitPyClass(node: PyClass?) {
super.visitPyClass(node)
if (node != null && LanguageLevel.forElement(node).isAtLeast(LanguageLevel.PYTHON36) && isTypingNTInheritor(node)) {
if (node != null &&
LanguageLevel.forElement(node).isAtLeast(LanguageLevel.PYTHON36) &&
PyStdlibTypeProvider.isTypingNamedTupleDirectInheritor(node, myTypeEvalContext)) {
inspectFieldsOrder(node, myTypeEvalContext, this::registerProblem)
}
}
private fun isTypingNTInheritor(cls: PyClass): Boolean {
val isTypingNT: (PyClassLikeType?) -> Boolean =
{ it != null && it !is PyNamedTupleType && PyTypingTypeProvider.NAMEDTUPLE == it.classQName }
return cls.getSuperClassTypes(myTypeEvalContext).find(isTypingNT) != null
}
}
private class FieldsProcessor(private val context: TypeEvalContext) : PsiScopeProcessor {
@@ -58,6 +58,7 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
myProject = project;
}
@Override
public ASTNode createNameIdentifier(String name, LanguageLevel languageLevel) {
final PsiFile dummyFile = createDummyFile(languageLevel, name);
final PyExpressionStatement expressionStatement = (PyExpressionStatement)dummyFile.getFirstChild();
@@ -92,6 +93,7 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
return "dummy." + PythonFileType.INSTANCE.getDefaultExtension();
}
@Override
public PyStringLiteralExpression createStringLiteralAlreadyEscaped(String str) {
final PsiFile dummyFile = createDummyFile(LanguageLevel.getDefault(), "a=(" + str + ")");
final PyAssignmentStatement expressionStatement = (PyAssignmentStatement)dummyFile.getFirstChild();
@@ -183,12 +185,14 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
return createStringLiteralAlreadyEscaped(buf.toString());
}
@Override
public PyListLiteralExpression createListLiteral() {
final PsiFile dummyFile = createDummyFile(LanguageLevel.getDefault(), "[]");
final PyExpressionStatement expressionStatement = (PyExpressionStatement)dummyFile.getFirstChild();
return (PyListLiteralExpression)expressionStatement.getFirstChild();
}
@Override
public ASTNode createComma() {
final PsiFile dummyFile = createDummyFile(LanguageLevel.getDefault(), "[0,]");
final PyExpressionStatement expressionStatement = (PyExpressionStatement)dummyFile.getFirstChild();
@@ -196,6 +200,7 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
return zero.getTreeNext().copyElement();
}
@Override
public ASTNode createDot() {
final PsiFile dummyFile = createDummyFile(LanguageLevel.getDefault(), "a.b");
final PyExpressionStatement expressionStatement = (PyExpressionStatement)dummyFile.getFirstChild();
@@ -227,6 +232,7 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
// TODO: Adds comma to empty list: adding "foo" to () will create (foo,). That is why "insertItemIntoListRemoveRedundantCommas" was created.
// We probably need to fix this method and delete insertItemIntoListRemoveRedundantCommas
@Override
public PsiElement insertItemIntoList(PyElement list, @Nullable PyExpression afterThis, PyExpression toInsert)
throws IncorrectOperationException {
ASTNode add = toInsert.getNode().copyElement();
@@ -265,6 +271,7 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
return add.getPsi();
}
@Override
public PyBinaryExpression createBinaryExpression(String s, PyExpression expr, PyExpression listLiteral) {
final PsiFile dummyFile = createDummyFile(LanguageLevel.getDefault(), "a " + s + " b");
final PyExpressionStatement expressionStatement = (PyExpressionStatement)dummyFile.getFirstChild();
@@ -275,10 +282,12 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
return binExpr;
}
@Override
public PyExpression createExpressionFromText(final String text) {
return createExpressionFromText(LanguageLevel.getDefault(), text);
}
@Override
@NotNull
public PyExpression createExpressionFromText(final LanguageLevel languageLevel, final String text) {
final PsiFile dummyFile = createDummyFile(languageLevel, text);
@@ -289,6 +298,7 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
throw new IncorrectOperationException("could not parse text as expression: " + text);
}
@Override
@NotNull
public PyCallExpression createCallExpression(final LanguageLevel langLevel, String functionName) {
final PsiFile dummyFile = createDummyFile(langLevel, functionName + "()");
@@ -328,6 +338,7 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
static final int[] FROM_ROOT = new int[]{0};
@Override
@NotNull
public <T> T createFromText(LanguageLevel langLevel, Class<T> aClass, final String text) {
return createFromText(langLevel, aClass, text, FROM_ROOT);
@@ -341,6 +352,7 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
static int[] PATH_PARAMETER = {0, 3, 1};
@Override
public PyNamedParameter createParameter(@NotNull String name) {
return createParameter(name, null, null, LanguageLevel.getDefault());
}
@@ -358,6 +370,7 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
}
@Override
public PyNamedParameter createParameter(@NotNull String name, @Nullable String defaultValue, @Nullable String annotation,
@NotNull LanguageLevel languageLevel) {
String parameterText = name;
@@ -377,6 +390,7 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
return (PyKeywordArgument)callExpression.getArguments()[0];
}
@Override
@NotNull
public <T> T createFromText(LanguageLevel langLevel, Class<T> aClass, final String text, final int[] path) {
return createFromText(langLevel, aClass, text, path, false);
@@ -441,6 +455,7 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
return function.getStatementList();
}
@Override
public PyExpressionStatement createDocstring(String content) {
return createFromText(LanguageLevel.getDefault(),
PyExpressionStatement.class, content + "\n");
@@ -469,6 +484,18 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
return createFromText(languageLevel, PyImportStatement.class, statement);
}
@NotNull
@Override
public PyNoneLiteralExpression createEllipsis() {
return createFromText(LanguageLevel.PYTHON30, PyNoneLiteralExpression.class, "...", new int[]{0, 0});
}
@NotNull
@Override
public PySingleStarParameter createSingleStarParameter() {
return createFromText(LanguageLevel.PYTHON30, PySingleStarParameter.class, "def foo(*): pass", new int[]{0, 3, 1});
}
private static class CommasOnly extends NotNullPredicate<LeafPsiElement> {
@Override
protected boolean applyNotNull(@NotNull final LeafPsiElement input) {
@@ -0,0 +1,48 @@
from collections import namedtuple
MyTup1 = namedtuple("MyTup1", "bar baz")
mt1 = MyTup1(1, 2)
# empty
mt1._replace()
# one
mt1._replace(bar=2)
mt1._replace(baz=1)
mt1._replace(<warning descr="Unexpected argument">foo=1</warning>)
mt1._replace(<warning descr="Unexpected argument">1</warning>)
# two
mt1._replace(bar=1, baz=2)
mt1._replace(baz=2, bar=1)
mt1._replace(baz=2, <warning descr="Unexpected argument">foo=1</warning>)
mt1._replace(<warning descr="Unexpected argument">2</warning>, <warning descr="Unexpected argument">1</warning>)
# two
mt1._replace(bar=1, baz=2, <warning descr="Unexpected argument">foo=3</warning>)
mt1._replace(<warning descr="Unexpected argument">1</warning>, <warning descr="Unexpected argument">2</warning>, <warning descr="Unexpected argument">3</warning>)
class MyTup2(namedtuple("MyTup2", "bar baz")):
pass
mt2 = MyTup2(1, 2)
# empty
mt2._replace()
# one
mt2._replace(bar=2)
mt2._replace(baz=1)
mt2._replace(<warning descr="Unexpected argument">foo=1</warning>)
mt2._replace(<warning descr="Unexpected argument">1</warning>)
# two
mt2._replace(bar=1, baz=2)
mt2._replace(baz=2, bar=1)
mt2._replace(baz=2, <warning descr="Unexpected argument">foo=1</warning>)
mt2._replace(<warning descr="Unexpected argument">2</warning>, <warning descr="Unexpected argument">1</warning>)
# two
mt2._replace(bar=1, baz=2, <warning descr="Unexpected argument">foo=3</warning>)
mt2._replace(<warning descr="Unexpected argument">1</warning>, <warning descr="Unexpected argument">2</warning>, <warning descr="Unexpected argument">3</warning>)
@@ -0,0 +1,49 @@
import typing
MyTup1 = typing.NamedTuple("MyTup1", bar=int, baz=int)
mt1 = MyTup1(1, 2)
# empty
mt1._replace()
# one
mt1._replace(bar=2)
mt1._replace(baz=1)
mt1._replace(<warning descr="Unexpected argument">foo=1</warning>)
mt1._replace(<warning descr="Unexpected argument">1</warning>)
# two
mt1._replace(bar=1, baz=2)
mt1._replace(baz=2, bar=1)
mt1._replace(baz=2, <warning descr="Unexpected argument">foo=1</warning>)
mt1._replace(<warning descr="Unexpected argument">2</warning>, <warning descr="Unexpected argument">1</warning>)
# two
mt1._replace(bar=1, baz=2, <warning descr="Unexpected argument">foo=3</warning>)
mt1._replace(<warning descr="Unexpected argument">1</warning>, <warning descr="Unexpected argument">2</warning>, <warning descr="Unexpected argument">3</warning>)
class MyTup2(typing.NamedTuple):
bar: int
baz: int
mt2 = MyTup2(1, 2)
# empty
mt2._replace()
# one
mt2._replace(bar=2)
mt2._replace(baz=1)
mt2._replace(<warning descr="Unexpected argument">foo=1</warning>)
mt2._replace(<warning descr="Unexpected argument">1</warning>)
# two
mt2._replace(bar=1, baz=2)
mt2._replace(baz=2, bar=1)
mt2._replace(baz=2, <warning descr="Unexpected argument">foo=1</warning>)
mt2._replace(<warning descr="Unexpected argument">2</warning>, <warning descr="Unexpected argument">1</warning>)
# two
mt2._replace(bar=1, baz=2, <warning descr="Unexpected argument">foo=3</warning>)
mt2._replace(<warning descr="Unexpected argument">1</warning>, <warning descr="Unexpected argument">2</warning>, <warning descr="Unexpected argument">3</warning>)
@@ -0,0 +1,12 @@
from collections import namedtuple
MyTup1 = namedtuple("MyTup1", "bar baz")
class MyTup2(namedtuple("MyTup2", "bar baz")):
pass
MyTup1(1, 2)._replace(<arg1>)
MyTup2(1, 2)._replace(<arg2>)
@@ -0,0 +1,13 @@
import typing
MyTup1 = typing.NamedTuple("MyTup2", bar=int, baz=str)
class MyTup2(typing.NamedTuple):
bar: int
baz: str
MyTup1(1, "")._replace(<arg1>)
MyTup2(1, "")._replace(<arg2>)
@@ -641,24 +641,19 @@ public class PyParameterInfoTest extends LightMarkedTestCase {
// PY-22249
public void testInitializingCollectionsNamedTuple() {
runWithLanguageLevel(
LanguageLevel.PYTHON35,
() -> {
final Map<String, PsiElement> test = loadTest(2);
final Map<String, PsiElement> test = loadTest(2);
for (int offset : StreamEx.of(test.values()).map(PsiElement::getTextOffset)) {
final List<String> texts = Collections.singletonList("bar, baz");
final List<String[]> highlighted = Collections.singletonList(new String[]{"bar, "});
for (int offset : StreamEx.of(test.values()).map(PsiElement::getTextOffset)) {
final List<String> texts = Collections.singletonList("bar, baz");
final List<String[]> highlighted = Collections.singletonList(new String[]{"bar, "});
feignCtrlP(offset).check(texts, highlighted, Collections.singletonList(ArrayUtil.EMPTY_STRING_ARRAY));
}
}
);
feignCtrlP(offset).check(texts, highlighted, Collections.singletonList(ArrayUtil.EMPTY_STRING_ARRAY));
}
}
public void testInitializingTypingNamedTuple() {
runWithLanguageLevel(
LanguageLevel.PYTHON35,
LanguageLevel.PYTHON36,
() -> {
final Map<String, PsiElement> test = loadTest(7);
@@ -696,6 +691,29 @@ public class PyParameterInfoTest extends LightMarkedTestCase {
);
}
// PY-27148
public void testCollectionsNamedTupleReplace() {
final Map<String, PsiElement> test = loadTest(2);
for (int offset : StreamEx.of(test.values()).map(PsiElement::getTextOffset)) {
feignCtrlP(offset).check("*, bar=..., baz=...", ArrayUtil.EMPTY_STRING_ARRAY);
}
}
// PY-27148
public void testTypingNamedTupleReplace() {
runWithLanguageLevel(
LanguageLevel.PYTHON36,
() -> {
final Map<String, PsiElement> test = loadTest(2);
for (int offset : StreamEx.of(test.values()).map(PsiElement::getTextOffset)) {
feignCtrlP(offset).check("*, bar: int=..., baz: str=...", ArrayUtil.EMPTY_STRING_ARRAY);
}
}
);
}
// PY-26582
public void testStructuralType() {
runWithLanguageLevel(
@@ -2670,6 +2670,11 @@ public class PyTypeTest extends PyTestCase {
"class Cat(namedtuple(\"Cat\", \"name age\")):\n" +
" pass\n" +
"expr = Cat(\"name\", 5)._replace(name=\"newname\")");
doTest("str",
"from collections import namedtuple\n" +
"Cat = namedtuple(\"Cat\", \"name age\")\n" +
"expr = Cat(\"name\", 5)._replace(age=\"five\").age");
}
// PY-27148
@@ -2691,6 +2696,14 @@ public class PyTypeTest extends PyTestCase {
"Cat = NamedTuple(\"Cat\", name=str, age=int)\n" +
"expr = Cat(\"name\", 5)._replace(name=\"newname\")")
);
runWithLanguageLevel(
LanguageLevel.PYTHON36,
() -> doTest("int",
"from typing import NamedTuple\n" +
"Cat = NamedTuple(\"Cat\", name=str, age=int)\n" +
"expr = Cat(\"name\", 5)._replace(age=\"give\").age")
);
}
private static List<TypeEvalContext> getTypeEvalContexts(@NotNull PyExpression element) {
@@ -319,6 +319,16 @@ public class PyArgumentListInspectionTest extends PyInspectionTestCase {
runWithLanguageLevel(LanguageLevel.PYTHON30, this::doTest);
}
// PY-27148
public void testCollectionsNamedTupleReplace() {
doTest();
}
// PY-27148
public void testTypingNamedTupleReplace() {
runWithLanguageLevel(LanguageLevel.PYTHON36, this::doTest);
}
// PY-27398
public void testInitializingDataclass() {
runWithLanguageLevel(LanguageLevel.PYTHON37, this::doMultiFileTest);