Check parameters passed to Callable and suggest quick fixes (PY-20530)

This commit is contained in:
Semyon Proshev
2018-06-13 23:12:46 +03:00
parent 0d9bf8ebff
commit 4e69ba030f
11 changed files with 155 additions and 7 deletions
@@ -4,6 +4,7 @@ package com.jetbrains.python.inspections
import com.intellij.codeInsight.controlflow.ControlFlowUtil
import com.intellij.codeInspection.*
import com.intellij.openapi.project.Project
import com.intellij.openapi.util.TextRange
import com.intellij.psi.PsiElement
import com.intellij.psi.PsiElementVisitor
import com.intellij.psi.PsiFileFactory
@@ -64,15 +65,18 @@ class PyTypeHintsInspection : PyInspection() {
}
}
override fun visitPySubscriptionExpression(node: PySubscriptionExpression?) {
override fun visitPySubscriptionExpression(node: PySubscriptionExpression) {
super.visitPySubscriptionExpression(node)
if (node != null) {
val callee = node.operand as? PyReferenceExpression
val calleeQName = callee?.let { PyResolveUtil.resolveImportedElementQNameLocally(it) } ?: emptyList()
val operand = node.operand as? PyReferenceExpression ?: return
val index = node.indexExpression ?: return
if (genericQName in calleeQName) {
checkGenericParameters(node.indexExpression)
val callableQName = QualifiedName.fromDottedString(PyTypingTypeProvider.CALLABLE)
PyResolveUtil.resolveImportedElementQNameLocally(operand).forEach {
when (it) {
genericQName -> checkGenericParameters(index)
callableQName -> checkCallableParameters(index)
}
}
}
@@ -386,7 +390,7 @@ class PyTypeHintsInspection : PyInspection() {
return Pair(if (seenGeneric) genericTypeVars else null, nonGenericTypeVars)
}
private fun checkGenericParameters(index: PyExpression?) {
private fun checkGenericParameters(index: PyExpression) {
val parameters = (index as? PyTupleExpression)?.elements ?: arrayOf(index)
val typeVars = mutableSetOf<PsiElement>()
@@ -411,6 +415,41 @@ class PyTypeHintsInspection : PyInspection() {
}
}
private fun checkCallableParameters(index: PyExpression) {
val message = "'Callable' must be used as 'Callable[[arg, ...], result]'"
if (index !is PyTupleExpression) {
registerProblem(index, message, ProblemHighlightType.GENERIC_ERROR)
return
}
val parameters = index.elements
if (parameters.size > 2) {
val possiblyLastParameter = parameters[parameters.size - 2]
registerProblem(index,
message,
ProblemHighlightType.GENERIC_ERROR,
null,
TextRange.create(0, possiblyLastParameter.startOffsetInParent + possiblyLastParameter.textLength),
SurroundElementsWithSquareBracketsQuickFix())
}
else if (parameters.size < 2) {
registerProblem(index, message, ProblemHighlightType.GENERIC_ERROR)
}
else {
val first = parameters.first()
if (first !is PyListLiteralExpression && !(first is PyNoneLiteralExpression && first.isEllipsis)) {
registerProblem(first,
message,
ProblemHighlightType.GENERIC_ERROR,
null,
if (first is PyParenthesizedExpression) ReplaceWithListQuickFix() else SurroundElementWithSquareBracketsQuickFix())
}
}
}
private fun followNotTypingOpaque(target: PyTargetExpression): Boolean {
return !PyTypingTypeProvider.OPAQUE_NAMES.contains(target.qualifiedName)
}
@@ -472,5 +511,51 @@ class PyTypeHintsInspection : PyInspection() {
?.let { element.replace(it) }
}
}
private class SurroundElementsWithSquareBracketsQuickFix : LocalQuickFix {
override fun getFamilyName() = "Surround with square brackets"
override fun applyFix(project: Project, descriptor: ProblemDescriptor) {
val element = descriptor.psiElement as? PyTupleExpression ?: return
val list = PyElementGenerator.getInstance(project).createListLiteral()
val originalElements = element.elements
originalElements.dropLast(1).forEach { list.add(it) }
originalElements.dropLast(2).forEach { it.delete() }
element.elements.first().replace(list)
}
}
private class SurroundElementWithSquareBracketsQuickFix : LocalQuickFix {
override fun getFamilyName() = "Surround with square brackets"
override fun applyFix(project: Project, descriptor: ProblemDescriptor) {
val element = descriptor.psiElement
val list = PyElementGenerator.getInstance(project).createListLiteral()
list.add(element)
element.replace(list)
}
}
private class ReplaceWithListQuickFix : LocalQuickFix {
override fun getFamilyName() = "Replace with square brackets"
override fun applyFix(project: Project, descriptor: ProblemDescriptor) {
val element = descriptor.psiElement
val expression = (element as? PyParenthesizedExpression)?.containedExpression ?: return
val elements = expression.let { if (it is PyTupleExpression) it.elements else arrayOf(it) }
val list = PyElementGenerator.getInstance(project).createListLiteral()
elements.forEach { list.add(it) }
element.replace(list)
}
}
}
}
@@ -0,0 +1,3 @@
from typing import Callable
f: Callable[<error descr="'Callable' must be used as 'Callable[[arg, ...], result]'">int<caret>, str</error>, str]
@@ -0,0 +1,3 @@
from typing import Callable
f: Callable[[int, str], str]
@@ -0,0 +1,3 @@
from typing import Callable
e: Callable[<error descr="'Callable' must be used as 'Callable[[arg, ...], result]'">i<caret>nt</error>, str]
@@ -0,0 +1,3 @@
from typing import Callable
g: Callable[<error descr="'Callable' must be used as 'Callable[[arg, ...], result]'">(int<caret>)</error>, str]
@@ -0,0 +1,3 @@
from typing import Callable
g: Callable[[int], str]
@@ -0,0 +1,3 @@
from typing import Callable
g: Callable[<error descr="'Callable' must be used as 'Callable[[arg, ...], result]'">(int<caret>, str)</error>, str]
@@ -0,0 +1,3 @@
from typing import Callable
g: Callable[[int, str], str]
@@ -0,0 +1,3 @@
from typing import Callable
e: Callable[[int], str]
@@ -534,6 +534,25 @@ public class PyTypeHintsInspectionTest extends PyInspectionTestCase {
);
}
// PY-20530
public void testCallableParameters() {
runWithLanguageLevel(
LanguageLevel.PYTHON36,
() -> doTestByText("from typing import Callable\n" +
"\n" +
"a: Callable[..., str]\n" +
"b: Callable[[int], str]\n" +
"c: Callable[[int, str], str]\n" +
"\n" +
"d: Callable[<error descr=\"'Callable' must be used as 'Callable[[arg, ...], result]'\">...</error>]\n" +
"e: Callable[<error descr=\"'Callable' must be used as 'Callable[[arg, ...], result]'\">int</error>, str]\n" +
"f: Callable[<error descr=\"'Callable' must be used as 'Callable[[arg, ...], result]'\">int, str</error>, str]\n" +
"g: Callable[<error descr=\"'Callable' must be used as 'Callable[[arg, ...], result]'\">(int, str)</error>, str]\n" +
"h: Callable[<error descr=\"'Callable' must be used as 'Callable[[arg, ...], result]'\">int</error>]\n" +
"h: Callable[<error descr=\"'Callable' must be used as 'Callable[[arg, ...], result]'\">(int)</error>, str]")
);
}
@NotNull
@Override
protected Class<? extends PyInspection> getInspectionClass() {
@@ -65,4 +65,24 @@ class PyTypeHintsQuickFixTest : PyQuickFixTestCase() {
fun testInstanceCheckOnOperandReference() {
doQuickFixTest(PyTypeHintsInspection::class.java, "Remove generic parameter(s)")
}
// PY-20530
fun testCallableMoreThanTwoElements() {
doQuickFixTest(PyTypeHintsInspection::class.java, "Surround with square brackets", LanguageLevel.PYTHON37)
}
// PY-20530
fun testCallableTwoElements() {
doQuickFixTest(PyTypeHintsInspection::class.java, "Surround with square brackets", LanguageLevel.PYTHON37)
}
// PY-20530
fun testCallableTwoElementsTupleAsFirst() {
doQuickFixTest(PyTypeHintsInspection::class.java, "Replace with square brackets", LanguageLevel.PYTHON37)
}
// PY-20530
fun testCallableTwoElementsParenthesizedAsFirst() {
doQuickFixTest(PyTypeHintsInspection::class.java, "Replace with square brackets", LanguageLevel.PYTHON37)
}
}