Support kw_only in attr.s (PY-34374)

GitOrigin-RevId: a8e78103f3373c967cd429b21e0ced0d530b468d
This commit is contained in:
Semyon Proshev
2019-12-04 12:33:00 +00:00
committed by intellij-monorepo-bot
parent 9ef2a6f638
commit 1efa876ad0
19 changed files with 478 additions and 28 deletions
@@ -137,8 +137,8 @@ private fun parseDataclassParametersFromAST(cls: PyClass, context: TypeEvalConte
private fun parseDataclassParametersFromStub(stub: PyDataclassStub?): PyDataclassParameters? {
return stub?.let {
PyDataclassParameters(
it.initValue(), it.reprValue(), it.eqValue(), it.orderValue(), it.unsafeHashValue(), it.frozenValue(),
null, null, null, null, null, null,
it.initValue(), it.reprValue(), it.eqValue(), it.orderValue(), it.unsafeHashValue(), it.frozenValue(), it.kwOnly(),
null, null, null, null, null, null, null,
Type.valueOf(it.type), emptyMap()
)
}
@@ -164,12 +164,14 @@ data class PyDataclassParameters(val init: Boolean,
val order: Boolean,
val unsafeHash: Boolean,
val frozen: Boolean,
val kwOnly: Boolean,
val initArgument: PyExpression?,
val reprArgument: PyExpression?,
val eqArgument: PyExpression?,
val orderArgument: PyExpression?,
val unsafeHashArgument: PyExpression?,
val frozenArgument: PyExpression?,
val kwOnlyArgument: PyExpression?,
val type: Type,
val others: Map<String, PyExpression>) {
@@ -189,6 +191,7 @@ private class PyDataclassParametersBuilder(private val type: Type,
private const val DEFAULT_ORDER = false
private const val DEFAULT_UNSAFE_HASH = false
private const val DEFAULT_FROZEN = false
private const val DEFAULT_KW_ONLY = false
}
private var init = DEFAULT_INIT
@@ -197,6 +200,7 @@ private class PyDataclassParametersBuilder(private val type: Type,
private var order = if (type == Type.ATTRS) DEFAULT_EQ else DEFAULT_ORDER
private var unsafeHash = DEFAULT_UNSAFE_HASH
private var frozen = DEFAULT_FROZEN
private var kwOnly = DEFAULT_KW_ONLY
private var initArgument: PyExpression? = null
private var reprArgument: PyExpression? = null
@@ -204,6 +208,7 @@ private class PyDataclassParametersBuilder(private val type: Type,
private var orderArgument: PyExpression? = null
private var unsafeHashArgument: PyExpression? = null
private var frozenArgument: PyExpression? = null
private var kwOnlyArgument: PyExpression? = null
private val others = mutableMapOf<String, PyExpression>()
@@ -285,6 +290,11 @@ private class PyDataclassParametersBuilder(private val type: Type,
unsafeHashArgument = argument
return
}
"kw_only" -> {
kwOnly = PyEvaluator.evaluateAsBoolean(value, DEFAULT_KW_ONLY)
kwOnlyArgument = argument
return
}
}
}
@@ -293,7 +303,9 @@ private class PyDataclassParametersBuilder(private val type: Type,
}
}
fun build() = PyDataclassParameters(init, repr, eq, order, unsafeHash, frozen,
initArgument, reprArgument, eqArgument, orderArgument, unsafeHashArgument, frozenArgument,
type, others)
fun build() = PyDataclassParameters(
init, repr, eq, order, unsafeHash, frozen, kwOnly,
initArgument, reprArgument, eqArgument, orderArgument, unsafeHashArgument, frozenArgument, kwOnlyArgument,
type, others
)
}
@@ -62,7 +62,7 @@ public class PyFileElementType extends IStubFileElementType<PyFileStub> {
@Override
public int getStubVersion() {
// Don't forget to update versions of indexes that use the updated stub-based elements
return 79;
return 80;
}
@Nullable
@@ -26,14 +26,14 @@ class PyDataclassStubImpl private constructor(private val type: String,
private val eq: Boolean,
private val order: Boolean,
private val unsafeHash: Boolean,
private val frozen: Boolean) : PyDataclassStub {
private val frozen: Boolean,
private val kwOnly: Boolean) : PyDataclassStub {
companion object {
fun create(cls: PyClass): PyDataclassStub? {
return parseDataclassParametersForStub(cls)?.let {
PyDataclassStubImpl(it.type.toString(), it.init, it.repr, it.eq, it.order, it.unsafeHash,
it.frozen)
PyDataclassStubImpl(it.type.toString(), it.init, it.repr, it.eq, it.order, it.unsafeHash, it.frozen, it.kwOnly)
}
}
@@ -47,8 +47,9 @@ class PyDataclassStubImpl private constructor(private val type: String,
val order = stream.readBoolean()
val unsafeHash = stream.readBoolean()
val frozen = stream.readBoolean()
val kwOnly = stream.readBoolean()
return PyDataclassStubImpl(type, init, repr, eq, order, unsafeHash, frozen)
return PyDataclassStubImpl(type, init, repr, eq, order, unsafeHash, frozen, kwOnly)
}
}
@@ -62,13 +63,15 @@ class PyDataclassStubImpl private constructor(private val type: String,
stream.writeBoolean(order)
stream.writeBoolean(unsafeHash)
stream.writeBoolean(frozen)
stream.writeBoolean(kwOnly)
}
override fun getType() = type
override fun initValue() = init
override fun reprValue() = repr
override fun eqValue() = eq
override fun orderValue() = order
override fun unsafeHashValue() = unsafeHash
override fun frozenValue() = frozen
override fun getType(): String = type
override fun initValue(): Boolean = init
override fun reprValue(): Boolean = repr
override fun eqValue(): Boolean = eq
override fun orderValue(): Boolean = order
override fun unsafeHashValue(): Boolean = unsafeHash
override fun frozenValue(): Boolean = frozen
override fun kwOnly(): Boolean = kwOnly
}
@@ -47,4 +47,10 @@ public interface PyDataclassStub extends PyCustomClassStub {
* its default value if it is not specified or could not be evaluated.
*/
boolean frozenValue();
/**
* @return value of `kw_only` (attrs) parameter or
* its default value if it is not specified or could not be evaluated.
*/
boolean kwOnly();
}
@@ -110,10 +110,13 @@ class PyDataclassTypeProvider : PyTypeProviderBase() {
val clsType = (context.getType(cls) as? PyClassLikeType) ?: return null
val resolveContext = PyResolveContext.defaultContext().withTypeEvalContext(context)
val ellipsis = PyElementGenerator.getInstance(cls.project).createEllipsis()
val elementGenerator = PyElementGenerator.getInstance(cls.project)
val ellipsis = elementGenerator.createEllipsis()
val collected = linkedMapOf<String, PyCallableParameter>()
var seenInit = false
val keywordOnly = linkedSetOf<String>()
var seenKeywordOnlyClass = false
for (currentType in StreamEx.of(clsType).append(cls.getAncestorTypes(context))) {
if (currentType == null ||
@@ -128,6 +131,7 @@ class PyDataclassTypeProvider : PyTypeProviderBase() {
val parameters = parseDataclassParameters(current, context) ?: continue
seenInit = seenInit || parameters.init
seenKeywordOnlyClass = seenKeywordOnlyClass || parameters.kwOnly
if (seenInit) {
current
@@ -140,6 +144,10 @@ class PyDataclassTypeProvider : PyTypeProviderBase() {
parameter.name?.let {
// note: attributes are visited from inheritors to ancestors, in reversed order for every of them
if (seenKeywordOnlyClass && it !in collected) {
keywordOnly += it
}
if (parameters.type == PyDataclassParameters.Type.STD) {
// std: attribute that overrides ancestor's attribute does not change the order but updates type
collected[it] = collected.remove(it) ?: parameter
@@ -153,7 +161,28 @@ class PyDataclassTypeProvider : PyTypeProviderBase() {
}
}
return if (seenInit) PyCallableTypeImpl(collected.values.reversed(), clsType.toInstance()) else null
return if (seenInit) PyCallableTypeImpl(buildParameters(elementGenerator, collected, keywordOnly), clsType.toInstance()) else null
}
private fun buildParameters(elementGenerator: PyElementGenerator,
fields: Map<String, PyCallableParameter>,
keywordOnly: Set<String>): List<PyCallableParameter> {
if (keywordOnly.isEmpty()) return fields.values.reversed()
val positionalOrKeyword = mutableListOf<PyCallableParameter>()
val keyword = mutableListOf<PyCallableParameter>()
for ((name, value) in fields.entries.reversed()) {
if (name !in keywordOnly) {
positionalOrKeyword += value
}
else {
keyword += value
}
}
val singleStarParameter = elementGenerator.createSingleStarParameter()
return positionalOrKeyword + listOf(PyCallableParameterImpl.psi(singleStarParameter)) + keyword
}
private fun fieldToParameter(cls: PyClass,
@@ -109,7 +109,10 @@ class PyDataclassInspection : PyInspection() {
PyNamedTupleInspection.inspectFieldsOrder(
node,
{ parseDataclassParameters(it, myTypeEvalContext) != null },
{
val parameters = parseDataclassParameters(it, myTypeEvalContext)
parameters != null && !parameters.kwOnly
},
dataclassParameters.type == PyDataclassParameters.Type.STD,
myTypeEvalContext,
this::registerProblem,
@@ -525,8 +528,7 @@ class PyDataclassInspection : PyInspection() {
val default = call.getKeywordArgument("default")
val factory = call.getKeywordArgument("factory")
if (default != null && factory != null && !resolvesToOmittedDefault(default,
PyDataclassParameters.Type.ATTRS)) {
if (default != null && factory != null && !resolvesToOmittedDefault(default, PyDataclassParameters.Type.ATTRS)) {
registerProblem(call.argumentList, "Cannot specify both 'default' and 'factory'", ProblemHighlightType.GENERIC_ERROR)
}
}
@@ -19,18 +19,20 @@ class PyNamedTupleInspection : PyInspection() {
companion object {
fun inspectFieldsOrder(cls: PyClass,
ancestorsFilter: (PyClass) -> Boolean,
classFieldsFilter: (PyClass) -> Boolean,
checkInheritedOrder: Boolean,
context: TypeEvalContext,
callback: (PsiElement, String, ProblemHighlightType) -> Unit,
fieldsFilter: (PyTargetExpression) -> Boolean = { true },
hasAssignedValue: (PyTargetExpression) -> Boolean = PyTargetExpression::hasAssignedValue) {
val fieldsProcessor = processFields(cls, fieldsFilter, hasAssignedValue)
val fieldsProcessor = if (classFieldsFilter(cls)) processFields(cls, fieldsFilter, hasAssignedValue) else null
if (fieldsProcessor?.lastFieldWithoutDefaultValue == null && !checkInheritedOrder) return
val ancestors = cls.getAncestorClasses(context)
val ancestorsFields = ancestors.map {
when {
!ancestorsFilter(it) -> Ancestor.FILTERED
!classFieldsFilter(it) -> Ancestor.FILTERED
processFields(it, fieldsFilter, hasAssignedValue).fieldsWithDefaultValue.isNotEmpty() -> Ancestor.HAS_FIELD_WITH_DEFAULT_VALUE
else -> Ancestor.HAS_NOT_FIELD_WITH_DEFAULT_VALUE
}
@@ -53,7 +55,7 @@ class PyNamedTupleInspection : PyInspection() {
}
}
val lastFieldWithoutDefaultValue = fieldsProcessor.lastFieldWithoutDefaultValue
val lastFieldWithoutDefaultValue = fieldsProcessor?.lastFieldWithoutDefaultValue
if (lastFieldWithoutDefaultValue != null) {
if (ancestorsFields.contains(Ancestor.HAS_FIELD_WITH_DEFAULT_VALUE)) {
cls.nameIdentifier?.let { name ->
@@ -103,7 +105,7 @@ class PyNamedTupleInspection : PyInspection() {
if (node != null &&
LanguageLevel.forElement(node).isAtLeast(LanguageLevel.PYTHON36) &&
PyNamedTupleTypeProvider.isTypingNamedTupleDirectInheritor(node, myTypeEvalContext)) {
inspectFieldsOrder(node, { false }, false, myTypeEvalContext, this::registerProblem)
inspectFieldsOrder(node, { it == node }, false, myTypeEvalContext, this::registerProblem)
}
}
}
@@ -0,0 +1,29 @@
import attr
@attr.dataclass(kw_only=True)
class Base:
a: int = 1
# no kw_only, no default
@attr.dataclass
class Derived1(Base):
b: str
# kw_only, no default
@attr.dataclass(kw_only=True)
class Derived2(Base):
b: str
# no kw_only, default
@attr.dataclass
class Derived3(Base):
b: str = "b"
# kw_only, default
@attr.dataclass(kw_only=True)
class Derived4(Base):
b: str = "b"
@@ -0,0 +1,40 @@
import attr
# no kw_only, no default
@attr.dataclass
class Base1:
a: int
@attr.dataclass(kw_only=True)
class Derived1(Base1):
b: str = "b"
# kw_only, no default
@attr.dataclass(kw_only=True)
class Base2:
a: int
@attr.dataclass(kw_only=True)
class Derived2(Base2):
b: str = "b"
# no kw_only, default
@attr.dataclass
class Base3:
a: int = 1
@attr.dataclass(kw_only=True)
class Derived3(Base3):
b: str = "b"
# kw_only, default
@attr.dataclass(kw_only=True)
class Base4:
a: int = 1
@attr.dataclass(kw_only=True)
class Derived4(Base4):
b: str = "b"
@@ -0,0 +1,29 @@
import attr
@attr.dataclass(kw_only=True)
class Base:
a: int
# no kw_only, no default
@attr.dataclass
class Derived1(Base):
b: str
# kw_only, no default
@attr.dataclass(kw_only=True)
class Derived2(Base):
b: str
# no kw_only, default
@attr.dataclass
class Derived3(Base):
b: str = "b"
# kw_only, default
@attr.dataclass(kw_only=True)
class Derived4(Base):
b: str = "b"
@@ -0,0 +1,40 @@
import attr
# no kw_only, no default
@attr.dataclass
class Base1:
a: int
@attr.dataclass(kw_only=True)
class Derived1(Base1):
b: str
# kw_only, no default
@attr.dataclass(kw_only=True)
class Base2:
a: int
@attr.dataclass(kw_only=True)
class Derived2(Base2):
b: str
# no kw_only, default
@attr.dataclass
class Base3:
a: int = 1
@attr.dataclass(kw_only=True)
class Derived3(Base3):
b: str
# kw_only, default
@attr.dataclass(kw_only=True)
class Base4:
a: int = 1
@attr.dataclass(kw_only=True)
class Derived4(Base4):
b: str
@@ -0,0 +1,45 @@
import attr
@attr.s(kw_only=True)
class Base1:
a = attr.ib(type=int)
@attr.s
class Derived1(Base1):
b = attr.ib(type=int)
Derived1(<arg1>)
@attr.s(kw_only=True)
class Base2:
a = attr.ib(type=int)
@attr.s
class Derived2(Base2):
b = attr.ib(type=int, default=1)
Derived2(<arg2>)
@attr.s(kw_only=True)
class Base3:
a = attr.ib(type=int, default=1)
@attr.s
class Derived3(Base3):
b = attr.ib(type=int)
Derived3(<arg3>)
@attr.s(kw_only=True)
class Base4:
a = attr.ib(type=int, default=1)
@attr.s
class Derived4(Base4):
b = attr.ib(type=int, default=1)
Derived4(<arg4>)
@@ -0,0 +1,53 @@
import attr
@attr.s(kw_only=True)
class A1:
a = attr.ib(type=int)
A1(<arg1>)
@attr.s(kw_only=True)
class Base1:
a = attr.ib(type=int)
@attr.s(kw_only=True)
class Derived1(Base1):
b = attr.ib(type=int)
Derived1(<arg2>)
@attr.s(kw_only=True)
class Base2:
a = attr.ib(type=int)
@attr.s(kw_only=True)
class Derived2(Base2):
b = attr.ib(type=int, default=1)
Derived2(<arg3>)
@attr.s(kw_only=True)
class Base3:
a = attr.ib(type=int, default=1)
@attr.s(kw_only=True)
class Derived3(Base3):
b = attr.ib(type=int)
Derived3(<arg4>)
@attr.s(kw_only=True)
class Base4:
a = attr.ib(type=int, default=1)
@attr.s(kw_only=True)
class Derived4(Base4):
b = attr.ib(type=int, default=1)
Derived4(<arg5>)
@@ -0,0 +1,34 @@
import attr
@attr.s
class Base1:
a = attr.ib(type=int)
@attr.s(kw_only=True)
class Derived1(Base1):
a = attr.ib(type=int)
Derived1(<arg1>)
@attr.s(kw_only=True)
class Base2:
a = attr.ib(type=int)
@attr.s
class Derived2(Base2):
a = attr.ib(type=int)
Derived2(<arg2>)
@attr.s(kw_only=True)
class Base3:
a = attr.ib(type=int)
@attr.s(kw_only=True)
class Derived3(Base3):
a = attr.ib(type=int)
Derived3(<arg3>)
@@ -0,0 +1,45 @@
import attr
@attr.s
class Base1:
a = attr.ib(type=int)
@attr.s(kw_only=True)
class Derived1(Base1):
b = attr.ib(type=int)
Derived1(<arg1>)
@attr.s
class Base2:
a = attr.ib(type=int)
@attr.s(kw_only=True)
class Derived2(Base2):
b = attr.ib(type=int, default=1)
Derived2(<arg2>)
@attr.s
class Base3:
a = attr.ib(type=int, default=1)
@attr.s(kw_only=True)
class Derived3(Base3):
b = attr.ib(type=int)
Derived3(<arg3>)
@attr.s
class Base4:
a = attr.ib(type=int, default=1)
@attr.s(kw_only=True)
class Derived4(Base4):
b = attr.ib(type=int, default=1)
Derived4(<arg4>)
@@ -0,0 +1,10 @@
import attr
@attr.s(kw_only=True)
class Foo1:
bar = attr.ib(type=str)
@attr.s(kw_only=False)
class Foo2:
bar = attr.ib(type=str)
@attr.s
class Foo3:
bar = attr.ib(type=str)
@@ -878,6 +878,46 @@ public class PyParameterInfoTest extends LightMarkedTestCase {
);
}
// PY-34374
public void testInitializingAttrsKwOnlyOnClass() {
final Map<String, PsiElement> marks = loadTest(5);
feignCtrlP(marks.get("<arg1>").getTextOffset()).check("*, a: int", new String[]{"*, a: int"});
feignCtrlP(marks.get("<arg2>").getTextOffset()).check("*, a: int, b: int", new String[]{"*, a: int"});
feignCtrlP(marks.get("<arg3>").getTextOffset()).check("*, a: int, b: int=...", new String[]{"*, a: int"});
feignCtrlP(marks.get("<arg4>").getTextOffset()).check("*, a: int=..., b: int", new String[]{"*, a: int"});
feignCtrlP(marks.get("<arg5>").getTextOffset()).check("*, a: int=..., b: int=...", new String[]{"*, a: int"});
}
// PY-34374
public void testInitializingAttrsKwOnlyOnBaseClass() {
final Map<String, PsiElement> marks = loadTest(4);
feignCtrlP(marks.get("<arg1>").getTextOffset()).check("b: int, *, a: int", new String[]{"b: int, "});
feignCtrlP(marks.get("<arg2>").getTextOffset()).check("b: int=..., *, a: int", new String[]{"b: int=..., "});
feignCtrlP(marks.get("<arg3>").getTextOffset()).check("b: int, *, a: int=...", new String[]{"b: int, "});
feignCtrlP(marks.get("<arg4>").getTextOffset()).check("b: int=..., *, a: int=...", new String[]{"b: int=..., "});
}
// PY-34374
public void testInitializingAttrsKwOnlyOnDerivedClass() {
final Map<String, PsiElement> marks = loadTest(4);
feignCtrlP(marks.get("<arg1>").getTextOffset()).check("*, a: int, b: int", new String[]{"*, a: int"});
feignCtrlP(marks.get("<arg2>").getTextOffset()).check("*, a: int, b: int=...", new String[]{"*, a: int"});
feignCtrlP(marks.get("<arg3>").getTextOffset()).check("*, a: int=..., b: int", new String[]{"*, a: int"});
feignCtrlP(marks.get("<arg4>").getTextOffset()).check("*, a: int=..., b: int=...", new String[]{"*, a: int"});
}
// PY-34374
public void testInitializingAttrsKwOnlyOnClassOverridingHierarchy() {
final Map<String, PsiElement> marks = loadTest(3);
feignCtrlP(marks.get("<arg1>").getTextOffset()).check("*, a: int", new String[]{"*, a: int"});
feignCtrlP(marks.get("<arg2>").getTextOffset()).check("a: int", new String[]{"a: int"});
feignCtrlP(marks.get("<arg3>").getTextOffset()).check("*, a: int", new String[]{"*, a: int"});
}
// PY-28957
public void testDataclassesReplace() {
runWithLanguageLevel(
@@ -966,6 +966,17 @@ public class PyStubsTest extends PyTestCase {
);
}
// PY-34374
public void testAttrsKwOnlyOnClass() {
final PyFile file = getTestFile();
assertTrue(file.findTopLevelClass("Foo1").getStub().getCustomStub(PyDataclassStub.class).kwOnly());
assertFalse(file.findTopLevelClass("Foo2").getStub().getCustomStub(PyDataclassStub.class).kwOnly());
assertFalse(file.findTopLevelClass("Foo3").getStub().getCustomStub(PyDataclassStub.class).kwOnly());
assertNotParsed(file);
}
// PY-27398
public void testTypingNewType() {
runWithLanguageLevel(
@@ -66,6 +66,26 @@ public class PyDataclassInspectionTest extends PyInspectionTestCase {
doTest();
}
// PY-34374
public void testFieldsOrderInInheritanceKwOnlyNoDefaultBase() {
doTest();
}
// PY-34374
public void testFieldsOrderInInheritanceKwOnlyDefaultBase() {
doTest();
}
// PY-34374
public void testFieldsOrderInInheritanceKwOnlyNoDefaultDerived() {
doTest();
}
// PY-34374
public void testFieldsOrderInInheritanceKwOnlyDefaultDerived() {
doTest();
}
// PY-27398
public void testComparisonForOrdered() {
doTest();
@@ -279,7 +299,7 @@ public class PyDataclassInspectionTest extends PyInspectionTestCase {
@Override
protected void doTest() {
runWithLanguageLevel(
LanguageLevel.PYTHON37,
LanguageLevel.getLatest(),
() -> {
myFixture.copyFileToProject(getTestCaseDirectory() + "/dataclasses.py", "dataclasses.py");
super.doTest();