Support kw_only in attr.ib (PY-33189)

GitOrigin-RevId: 7354162dc40909842f95f9eb5c9e6b507c544ac8
This commit is contained in:
Semyon Proshev
2020-06-16 02:00:57 +03:00
committed by intellij-monorepo-bot
parent 416fc4085e
commit 327f9cbffd
14 changed files with 242 additions and 53 deletions
@@ -168,10 +168,10 @@ class PyDataclassTypeProvider : PyTypeProviderBase() {
.filterNot { PyTypingTypeProvider.isClassVar(it, context) }
.mapNotNull { fieldToParameter(current, it, parameters.type, ellipsis, context) }
.filterNot { it.first in seenNames }
.forEach { (name, parameter) ->
.forEach { (name, kwOnly, parameter) ->
// note: attributes are visited from inheritors to ancestors, in reversed order for every of them
if (seenKeywordOnlyClass && name !in collected) {
if ((seenKeywordOnlyClass || kwOnly) && name !in collected) {
keywordOnly += name
}
@@ -218,12 +218,12 @@ class PyDataclassTypeProvider : PyTypeProviderBase() {
field: PyTargetExpression,
dataclassType: PyDataclassParameters.Type,
ellipsis: PyNoneLiteralExpression,
context: TypeEvalContext): Pair<String, PyCallableParameter?>? {
context: TypeEvalContext): Triple<String, Boolean, PyCallableParameter?>? {
val fieldName = field.name ?: return null
val stub = field.stub
val fieldStub = if (stub == null) PyDataclassFieldStubImpl.create(field) else stub.getCustomStub(PyDataclassFieldStub::class.java)
if (fieldStub != null && !fieldStub.initValue()) return fieldName to null
if (fieldStub != null && !fieldStub.initValue()) return Triple(fieldName, false, null)
if (fieldStub == null && field.annotationValue == null) return null // skip fields that are not annotated
val parameterName =
@@ -234,11 +234,13 @@ class PyDataclassTypeProvider : PyTypeProviderBase() {
else it
}
val parameter = PyCallableParameterImpl.nonPsi(parameterName,
getTypeForParameter(cls, field, dataclassType, context),
getDefaultValueForParameter(cls, field, fieldStub, dataclassType, ellipsis, context))
val parameter = PyCallableParameterImpl.nonPsi(
parameterName,
getTypeForParameter(cls, field, dataclassType, context),
getDefaultValueForParameter(cls, field, fieldStub, dataclassType, ellipsis, context)
)
return parameterName to parameter
return Triple(parameterName, fieldStub?.kwOnly() == true, parameter)
}
private fun getTypeForParameter(cls: PyClass,
@@ -122,7 +122,7 @@ class PyDataclassInspection : PyInspection() {
val fieldStub = if (stub == null) PyDataclassFieldStubImpl.create(it)
else stub.getCustomStub(PyDataclassFieldStub::class.java)
(fieldStub == null || fieldStub.initValue()) &&
(fieldStub == null || fieldStub.initValue() && !fieldStub.kwOnly()) &&
!(fieldStub == null && it.annotationValue == null) && // skip fields that are not annotated
!PyTypingTypeProvider.isClassVar(it, myTypeEvalContext) // skip classvars
},
@@ -63,7 +63,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 80;
return 81;
}
@Nullable
@@ -18,18 +18,16 @@ import com.jetbrains.python.psi.stubs.PyDataclassFieldStub
import java.io.IOException
class PyDataclassFieldStubImpl private constructor(private val calleeName: QualifiedName,
private val hasDefault: Boolean,
private val hasDefaultFactory: Boolean,
private val initValue: Boolean) : PyDataclassFieldStub {
private val parameters: FieldParameters) : PyDataclassFieldStub {
companion object {
fun create(expression: PyTargetExpression): PyDataclassFieldStub? {
val value = expression.findAssignedValue() as? PyCallExpression ?: return null
val callee = value.callee as? PyReferenceExpression ?: return null
val calleeNameAndType = calculateCalleeNameAndType(callee) ?: return null
val arguments = analyzeArguments(value, calleeNameAndType.second) ?: return null
val parameters = analyzeArguments(value, calleeNameAndType.second) ?: return null
return PyDataclassFieldStubImpl(calleeNameAndType.first, arguments.first, arguments.second, arguments.third)
return PyDataclassFieldStubImpl(calleeNameAndType.first, parameters)
}
@Throws(IOException::class)
@@ -38,8 +36,12 @@ class PyDataclassFieldStubImpl private constructor(private val calleeName: Quali
val hasDefault = stream.readBoolean()
val hasDefaultFactory = stream.readBoolean()
val initValue = stream.readBoolean()
val kwOnly = stream.readBoolean()
return PyDataclassFieldStubImpl(QualifiedName.fromDottedString(calleeName), hasDefault, hasDefaultFactory, initValue)
return PyDataclassFieldStubImpl(
QualifiedName.fromDottedString(calleeName),
FieldParameters(hasDefault, hasDefaultFactory, initValue, kwOnly)
)
}
private fun calculateCalleeNameAndType(callee: PyReferenceExpression): Pair<QualifiedName, PyDataclassParameters.PredefinedType>? {
@@ -60,20 +62,22 @@ class PyDataclassFieldStubImpl private constructor(private val calleeName: Quali
return null
}
private fun analyzeArguments(call: PyCallExpression, type: PyDataclassParameters.PredefinedType): Triple<Boolean, Boolean, Boolean>? {
private fun analyzeArguments(call: PyCallExpression, type: PyDataclassParameters.PredefinedType): FieldParameters? {
val initValue = PyEvaluator.evaluateAsBooleanNoResolve(call.getKeywordArgument("init"), true)
if (type == PyDataclassParameters.PredefinedType.STD) {
val default = call.getKeywordArgument("default")
val defaultFactory = call.getKeywordArgument("default_factory")
return Triple(default != null && !resolvesToOmittedDefault(default, type),
defaultFactory != null && !resolvesToOmittedDefault(defaultFactory, type),
initValue)
return FieldParameters(default != null && !resolvesToOmittedDefault(default, type),
defaultFactory != null && !resolvesToOmittedDefault(defaultFactory, type),
initValue,
false)
}
else if (type == PyDataclassParameters.PredefinedType.ATTRS) {
val default = call.getKeywordArgument("default")
val hasFactory = call.getKeywordArgument("factory").let { it != null && it.text != PyNames.NONE }
val kwOnly = PyEvaluator.evaluateAsBooleanNoResolve(call.getKeywordArgument("kw_only"), false)
if (default != null && !resolvesToOmittedDefault(default, type)) {
val callee = (default as? PyCallExpression)?.callee as? PyReferenceExpression
@@ -81,10 +85,10 @@ class PyDataclassFieldStubImpl private constructor(private val calleeName: Quali
callee != null &&
QualifiedName.fromComponents("attr", "Factory") in PyResolveUtil.resolveImportedElementQNameLocally(callee)
return Triple(!hasFactoryInDefault, hasFactory || hasFactoryInDefault, initValue)
return FieldParameters(!hasFactoryInDefault, hasFactory || hasFactoryInDefault, initValue, kwOnly)
}
return Triple(false, hasFactory, initValue)
return FieldParameters(false, hasFactory, initValue, kwOnly)
}
return null
@@ -97,13 +101,20 @@ class PyDataclassFieldStubImpl private constructor(private val calleeName: Quali
override fun serialize(stream: StubOutputStream) {
stream.writeName(calleeName.toString())
stream.writeBoolean(hasDefault)
stream.writeBoolean(hasDefaultFactory)
stream.writeBoolean(initValue)
stream.writeBoolean(parameters.hasDefault)
stream.writeBoolean(parameters.hasDefaultFactory)
stream.writeBoolean(parameters.initValue)
stream.writeBoolean(parameters.kwOnly)
}
override fun getCalleeName(): QualifiedName = calleeName
override fun hasDefault(): Boolean = hasDefault
override fun hasDefaultFactory(): Boolean = hasDefaultFactory
override fun initValue(): Boolean = initValue
override fun hasDefault(): Boolean = parameters.hasDefault
override fun hasDefaultFactory(): Boolean = parameters.hasDefaultFactory
override fun initValue(): Boolean = parameters.initValue
override fun kwOnly(): Boolean = parameters.kwOnly
private data class FieldParameters(val hasDefault: Boolean,
val hasDefaultFactory: Boolean,
val initValue: Boolean,
val kwOnly: Boolean)
}
@@ -21,4 +21,9 @@ public interface PyDataclassFieldStub extends CustomTargetExpressionStub {
* @return true if field is used in `__init__`.
*/
boolean initValue();
/**
* @return true if field is used in `__init__` as a keyword-only parameter (attrs).
*/
boolean kwOnly();
}
@@ -1,29 +1,58 @@
import attr
@attr.dataclass(kw_only=True)
class Base:
class Base1:
a: int = 1
# no kw_only, no default
@attr.dataclass
class Derived1(Base):
class Derived11(Base1):
b: str
# kw_only, no default
@attr.dataclass(kw_only=True)
class Derived2(Base):
class Derived12(Base1):
b: str
# no kw_only, default
@attr.dataclass
class Derived3(Base):
class Derived13(Base1):
b: str = "b"
# kw_only, default
@attr.dataclass(kw_only=True)
class Derived4(Base):
class Derived14(Base1):
b: str = "b"
@attr.s
class Base2:
a = attr.ib(type=int, kw_only=True)
# no kw_only, no default
@attr.dataclass
class Derived21(Base2):
b: str
# kw_only, no default
@attr.dataclass(kw_only=True)
class Derived22(Base2):
b: str
# no kw_only, default
@attr.dataclass
class Derived23(Base2):
b: str = "b"
# kw_only, default
@attr.dataclass(kw_only=True)
class Derived24(Base2):
b: str = "b"
@@ -6,9 +6,13 @@ class Base1:
a: int
@attr.dataclass(kw_only=True)
class Derived1(Base1):
class Derived11(Base1):
b: str = "b"
@attr.s
class Derived12(Base1):
b = attr.ib(type=str, default="b", kw_only=True)
# kw_only, no default
@attr.dataclass(kw_only=True)
@@ -16,9 +20,13 @@ class Base2:
a: int
@attr.dataclass(kw_only=True)
class Derived2(Base2):
class Derived21(Base2):
b: str = "b"
@attr.s
class Derived22(Base2):
b = attr.ib(type=str, default="b", kw_only=True)
# no kw_only, default
@attr.dataclass
@@ -26,9 +34,13 @@ class Base3:
a: int = 1
@attr.dataclass(kw_only=True)
class Derived3(Base3):
class Derived31(Base3):
b: str = "b"
@attr.s
class Derived32(Base3):
b = attr.ib(type=str, default="b", kw_only=True)
# kw_only, default
@attr.dataclass(kw_only=True)
@@ -36,5 +48,9 @@ class Base4:
a: int = 1
@attr.dataclass(kw_only=True)
class Derived4(Base4):
b: str = "b"
class Derived41(Base4):
b: str = "b"
@attr.s
class Derived42(Base4):
b = attr.ib(type=str, default="b", kw_only=True)
@@ -1,29 +1,58 @@
import attr
@attr.dataclass(kw_only=True)
class Base:
class Base1:
a: int
# no kw_only, no default
@attr.dataclass
class Derived1(Base):
class Derived11(Base1):
b: str
# kw_only, no default
@attr.dataclass(kw_only=True)
class Derived2(Base):
class Derived12(Base1):
b: str
# no kw_only, default
@attr.dataclass
class Derived3(Base):
class Derived13(Base1):
b: str = "b"
# kw_only, default
@attr.dataclass(kw_only=True)
class Derived4(Base):
class Derived14(Base1):
b: str = "b"
@attr.s
class Base2:
a = attr.ib(type=int, kw_only=True)
# no kw_only, no default
@attr.dataclass
class Derived21(Base2):
b: str
# kw_only, no default
@attr.dataclass(kw_only=True)
class Derived22(Base2):
b: str
# no kw_only, default
@attr.dataclass
class Derived23(Base2):
b: str = "b"
# kw_only, default
@attr.dataclass(kw_only=True)
class Derived24(Base2):
b: str = "b"
@@ -6,9 +6,13 @@ class Base1:
a: int
@attr.dataclass(kw_only=True)
class Derived1(Base1):
class Derived11(Base1):
b: str
@attr.s
class Derived12(Base1):
b = attr.ib(type=str, kw_only=True)
# kw_only, no default
@attr.dataclass(kw_only=True)
@@ -16,9 +20,13 @@ class Base2:
a: int
@attr.dataclass(kw_only=True)
class Derived2(Base2):
class Derived21(Base2):
b: str
@attr.s
class Derived22(Base2):
b = attr.ib(type=str, kw_only=True)
# no kw_only, default
@attr.dataclass
@@ -26,9 +34,13 @@ class Base3:
a: int = 1
@attr.dataclass(kw_only=True)
class Derived3(Base3):
class Derived31(Base3):
b: str
@attr.s
class Derived32(Base3):
b = attr.ib(type=str, kw_only=True)
# kw_only, default
@attr.dataclass(kw_only=True)
@@ -36,5 +48,9 @@ class Base4:
a: int = 1
@attr.dataclass(kw_only=True)
class Derived4(Base4):
b: str
class Derived41(Base4):
b: str
@attr.s
class Derived42(Base4):
b = attr.ib(type=str, kw_only=True)
@@ -0,0 +1,53 @@
import attr
@attr.s
class A:
a = attr.ib(kw_only=True)
b = attr.ib(type=int)
A(<arg1>)
@attr.s
class B1:
a = attr.ib(kw_only=True)
@attr.s
class D1(B1):
b = attr.ib(type=int)
D1(<arg2>)
@attr.s
class B2:
a = attr.ib(type=int)
@attr.s
class D2(B2):
b = attr.ib(kw_only=True)
D2(<arg3>)
@attr.s
class B3:
a = attr.ib(type=int)
@attr.s
class D3(B3):
a = attr.ib(kw_only=True)
D3(<arg4>)
@attr.s
class B4:
a = attr.ib(kw_only=True)
@attr.s
class D4(B4):
a = attr.ib(type=int)
D4(<arg5>)
@@ -0,0 +1,6 @@
import attr
@attr.s
class Foo:
bar1 = attr.ib(type=str)
bar2 = attr.ib(type=str, kw_only=True)
bar3 = attr.ib(type=str, kw_only=False)
@@ -909,6 +909,17 @@ public class PyParameterInfoTest extends LightMarkedTestCase {
feignCtrlP(marks.get("<arg3>").getTextOffset()).check("*, a: int", new String[]{"*, a: int"});
}
// PY-33189
public void testInitializingAttrsKwOnlyOnFields() {
final var marks = loadTest(5);
feignCtrlP(marks.get("<arg1>").getTextOffset()).check("b: int, *, a", new String[]{"b: int, "});
feignCtrlP(marks.get("<arg2>").getTextOffset()).check("b: int, *, a", new String[]{"b: int, "});
feignCtrlP(marks.get("<arg3>").getTextOffset()).check("a: int, *, b", new String[]{"a: int, "});
feignCtrlP(marks.get("<arg4>").getTextOffset()).check("*, a", ArrayUtil.EMPTY_STRING_ARRAY);
feignCtrlP(marks.get("<arg5>").getTextOffset()).check("a: int", new String[]{"a: int"});
}
// PY-28957
public void testDataclassesReplace() {
runWithLanguageLevel(
@@ -343,7 +343,6 @@ public class PyStubsTest extends PyTestCase {
private PyFile getTestFile(final String fileName) {
VirtualFile sourceFile = myFixture.copyFileToProject(fileName);
assert sourceFile != null;
PsiFile psiFile = myFixture.getPsiManager().findFile(sourceFile);
return (PyFile)psiFile;
}
@@ -977,6 +976,18 @@ public class PyStubsTest extends PyTestCase {
assertNotParsed(file);
}
// PY-33189
public void testAttrsKwOnlyOnField() {
final PyFile file = getTestFile();
final PyClass cls = file.findTopLevelClass("Foo");
assertFalse(cls.findClassAttribute("bar1", false, null).getStub().getCustomStub(PyDataclassFieldStub.class).kwOnly());
assertTrue(cls.findClassAttribute("bar2", false, null).getStub().getCustomStub(PyDataclassFieldStub.class).kwOnly());
assertFalse(cls.findClassAttribute("bar3", false, null).getStub().getCustomStub(PyDataclassFieldStub.class).kwOnly());
assertNotParsed(file);
}
// PY-27398
public void testTypingNewType() {
runWithLanguageLevel(
@@ -66,22 +66,22 @@ public class PyDataclassInspectionTest extends PyInspectionTestCase {
doTest();
}
// PY-34374
// PY-34374, PY-33189
public void testFieldsOrderInInheritanceKwOnlyNoDefaultBase() {
doTest();
}
// PY-34374
// PY-34374, PY-33189
public void testFieldsOrderInInheritanceKwOnlyDefaultBase() {
doTest();
}
// PY-34374
// PY-34374, PY-33189
public void testFieldsOrderInInheritanceKwOnlyNoDefaultDerived() {
doTest();
}
// PY-34374
// PY-34374, PY-33189
public void testFieldsOrderInInheritanceKwOnlyDefaultDerived() {
doTest();
}