PY-53693 PyCharm adds _: KW_ONLY to dataclass' init signature

GitOrigin-RevId: 471e10b4e12a36a640da2f3006719372797b5cf1
This commit is contained in:
Petr
2024-07-08 17:53:13 +00:00
committed by intellij-monorepo-bot
parent 303c01c1af
commit a6fc2c09c6
5 changed files with 172 additions and 11 deletions
@@ -31,6 +31,7 @@ object PyDataclassNames {
const val DATACLASSES_ASDICT = "dataclasses.asdict"
const val DATACLASSES_FIELD = "dataclasses.field"
const val DATACLASSES_REPLACE = "dataclasses.replace"
const val DATACLASSES_KW_ONLY = "dataclasses.KW_ONLY"
const val DUNDER_POST_INIT = "__post_init__"
val DECORATOR_PARAMETERS = listOf("init", "repr", "eq", "order", "unsafe_hash", "frozen", "kw_only")
val HELPER_FUNCTIONS = setOf(DATACLASSES_FIELDS, DATACLASSES_ASDICT, "dataclasses.astuple", DATACLASSES_REPLACE)
@@ -140,24 +140,31 @@ class PyDataclassTypeProvider : PyTypeProviderBase() {
seenKeywordOnlyClass = seenKeywordOnlyClass || parameters.kwOnly
if (seenInit) {
current
val fieldsInfo = current
.classAttributes
.asReversed()
.asSequence()
.filterNot { PyTypingTypeProvider.isClassVar(it, context) }
.mapNotNull { fieldToParameter(current, it, parameters.type, ellipsis, context) }
.filterNot { it.first in seenNames }
.forEach { (name, kwOnly, parameter) ->
// note: attributes are visited from inheritors to ancestors, in reversed order for every of them
.toList()
if ((seenKeywordOnlyClass || kwOnly) && name !in collected) {
keywordOnly += name
}
val indexOfKeywordOnlyAttribute = fieldsInfo.indexOfLast { (_, _, parameter) ->
parameter != null && isKwOnlyMarkerField(parameter, context)
}
if (parameter == null) {
seenNames.add(name)
}
else if (parameters.type.asPredefinedType == PyDataclassParameters.PredefinedType.STD) {
fieldsInfo.forEachIndexed { index, (name, kwOnly, parameter) ->
// note: attributes are visited from inheritors to ancestors, in reversed order for every of them
if ((seenKeywordOnlyClass || index < indexOfKeywordOnlyAttribute || kwOnly) && name !in collected) {
keywordOnly += name
}
if (parameter == null) {
seenNames.add(name)
}
else if (!isKwOnlyMarkerField(parameter, context)) {
if (parameters.type.asPredefinedType == PyDataclassParameters.PredefinedType.STD) {
// std: attribute that overrides ancestor's attribute does not change the order but updates type
collected[name] = collected.remove(name) ?: parameter
}
@@ -166,12 +173,20 @@ class PyDataclassTypeProvider : PyTypeProviderBase() {
collected[name] = parameter
}
}
}
}
}
return if (seenInit) PyCallableTypeImpl(buildParameters(elementGenerator, collected, keywordOnly), clsType.toInstance()) else null
}
private fun isKwOnlyMarkerField(parameter: PyCallableParameter, context: TypeEvalContext): Boolean {
val psi = parameter.declarationElement
if (psi !is PyTargetExpression) return false
val typeHint = PyTypingTypeProvider.getAnnotationValue(psi, context) as? PyReferenceExpression ?: return false
val type = Ref.deref(PyTypingTypeProvider.getType(typeHint, context))
return type is PyClassType && type.classQName == Dataclasses.DATACLASSES_KW_ONLY
}
private fun buildParameters(elementGenerator: PyElementGenerator,
fields: Map<String, PyCallableParameter>,
keywordOnly: Set<String>): List<PyCallableParameter> {
@@ -0,0 +1,46 @@
from dataclasses import dataclass, KW_ONLY
@dataclass
class Base:
a: int
qq: KW_ONLY
b: int
Base(<arg1>)
@dataclass
class Derived1(Base):
b: int
Derived1(<arg2>)
@dataclass
class Derived2(Base):
qq: int
Derived2(<arg3>)
@dataclass
class Derived3(Base):
c: int
ww: KW_ONLY
d: int
Derived3(<arg4>)
@dataclass
class Derived4(Base):
ww: KW_ONLY
qq: int
Derived4(<arg5>)
@dataclass
class Base2:
a: str
@dataclass
class Derived5(Base2):
a: KW_ONLY
Derived5(<arg6>)
@@ -1205,6 +1205,18 @@ public class PyParameterInfoTest extends LightMarkedTestCase {
feignCtrlP(marks.get("<arg4>").getTextOffset()).check("*, a: int", new String[]{"*, a: int"});
}
// PY-53693
public void testInitializingDataclassKwOnlyAttribute() {
final Map<String, PsiElement> marks = loadTest(6);
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, qq: int, *, b: int", new String[]{"a: int, "});
feignCtrlP(marks.get("<arg4>").getTextOffset()).check("a: int, c: int, *, b: int, d: int", new String[]{"a: int, "});
feignCtrlP(marks.get("<arg5>").getTextOffset()).check("a: int, *, b: int, qq: int", new String[]{"a: int, "});
feignCtrlP(marks.get("<arg6>").getTextOffset()).check("a: str", new String[]{"a: str"});
}
// PY-49946
public void testInitializingDataclassKwOnlyOnClassOverridingHierarchy() {
final Map<String, PsiElement> marks = loadTest(3);
@@ -224,4 +224,91 @@ public class Py3ArgumentListInspectionTest extends PyInspectionTestCase {
foo(title='Blade Runner', year=1982)
""");
}
// PY-53693
public void testInitializingDataclassWithKwOnlyAttribute() {
doTestByText("""
from dataclasses import dataclass, KW_ONLY
@dataclass
class MyClass:
a: int
qq: KW_ONLY
b: int
MyClass(0, b=0)
MyClass(0, <warning descr="Unexpected argument">0</warning><warning descr="Parameter 'b' unfilled">)</warning>
""");
}
// PY-53693
public void testInitializingDerivedDataclassWithKwOnlyAttribute() {
doTestByText("""
from dataclasses import dataclass, KW_ONLY
@dataclass
class Base:
a: int
qq: KW_ONLY
b: int
@dataclass
class Derived(Base):
c: int
ww: KW_ONLY
d: int
Derived(0, 0, b=0, d=0)
Derived(0, 0, <warning descr="Unexpected argument">0</warning>, b=0<warning descr="Parameter 'd' unfilled">)</warning>
Derived(0, 0, <warning descr="Unexpected argument">0</warning>, d=0<warning descr="Parameter 'b' unfilled">)</warning>
""");
}
// PY-53693
public void testInitializingDerivedDataclassWithOverridenAttribute() {
doTestByText("""
from dataclasses import dataclass, KW_ONLY
@dataclass
class Base:
a: int
qq: KW_ONLY
b: int
@dataclass
class Derived(Base):
b: int
Derived(0, 0)
""");
}
// PY-53693
public void testInitializingDerivedDataclassWithOverridenKwOnlyAttribute() {
doTestByText("""
from dataclasses import dataclass, KW_ONLY
@dataclass
class Base:
a: int
qq: KW_ONLY
b: int
@dataclass
class Derived1(Base):
qq: int
Derived1(0, 0, b=0)
Derived1(0, 0, <warning descr="Unexpected argument">0</warning><warning descr="Parameter 'b' unfilled">)</warning>
@dataclass
class Derived2(Base):
ww: KW_ONLY
qq: int
Derived2(0, b=0, qq=0)
Derived2(0, <warning descr="Unexpected argument">0</warning>, qq=0<warning descr="Parameter 'b' unfilled">)</warning>
Derived2(0, <warning descr="Unexpected argument">0</warning>, b=0<warning descr="Parameter 'qq' unfilled">)</warning>
""");
}
}