diff --git a/python/src/com/jetbrains/python/codeInsight/intentions/SpecifyTypeInDocstringIntention.java b/python/src/com/jetbrains/python/codeInsight/intentions/SpecifyTypeInDocstringIntention.java index 42c7231c08ec..5f101669c4fc 100644 --- a/python/src/com/jetbrains/python/codeInsight/intentions/SpecifyTypeInDocstringIntention.java +++ b/python/src/com/jetbrains/python/codeInsight/intentions/SpecifyTypeInDocstringIntention.java @@ -20,6 +20,7 @@ import com.jetbrains.python.PyTokenTypes; import com.jetbrains.python.documentation.PyDocumentationSettings; import com.jetbrains.python.documentation.PythonDocumentationProvider; import com.jetbrains.python.psi.*; +import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.types.PyReturnTypeReference; import com.jetbrains.python.psi.types.PyType; import com.jetbrains.python.psi.types.TypeEvalContext; @@ -50,12 +51,25 @@ public class SpecifyTypeInDocstringIntention implements IntentionAction { if (elementAt != null && !(elementAt.getNode().getElementType() == PyTokenTypes.IDENTIFIER)) elementAt = file.findElementAt(editor.getCaretModel().getOffset()); - PyFunction parentFunction = PsiTreeUtil.getParentOfType(elementAt, PyFunction.class); - if (parentFunction != null) { - final ASTNode nameNode = parentFunction.getNameNode(); - if (nameNode != null && nameNode.getPsi() == elementAt) { - myText = PyBundle.message("INTN.specify.return.type"); - return true; + PyCallExpression callExpression = PsiTreeUtil.getParentOfType(elementAt, PyCallExpression.class); + if (callExpression != null ) { + PyAssignmentStatement assignmentStatement = PsiTreeUtil.getParentOfType(elementAt, PyAssignmentStatement.class); + if (assignmentStatement != null) { + PyType type = assignmentStatement.getAssignedValue().getType(TypeEvalContext.slow()); + if (type == null || type instanceof PyReturnTypeReference) { + myText = PyBundle.message("INTN.specify.return.type"); + return true; + } + } + } + else { + PyFunction parentFunction = PsiTreeUtil.getParentOfType(elementAt, PyFunction.class); + if (parentFunction != null) { + final ASTNode nameNode = parentFunction.getNameNode(); + if (nameNode != null && nameNode.getPsi() == elementAt) { + myText = PyBundle.message("INTN.specify.return.type"); + return true; + } } } @@ -104,7 +118,20 @@ public class SpecifyTypeInDocstringIntention implements IntentionAction { String type = "type"; String name = ""; + PyCallExpression callExpression = PsiTreeUtil.getParentOfType(elementAt, PyCallExpression.class); PyFunction pyFunction = PsiTreeUtil.getParentOfType(elementAt, PyFunction.class); + PyExpression problemElement = PyUtil.findProblemElement(editor, file, PyNamedParameter.class, PyQualifiedExpression.class); + if (callExpression != null ) { + PyAssignmentStatement assignmentStatement = PsiTreeUtil.getParentOfType(elementAt, PyAssignmentStatement.class); + if (assignmentStatement != null) { + PyType pyType = assignmentStatement.getAssignedValue().getType(TypeEvalContext.slow()); + if (pyType == null || pyType instanceof PyReturnTypeReference) { + pyFunction = (PyFunction)callExpression.resolveCalleeFunction(PyResolveContext.defaultContext()); + problemElement = null; + type = "rtype"; + } + } + } if (pyFunction != null) { final ASTNode nameNode = pyFunction.getNameNode(); if (nameNode != null && nameNode.getPsi() == elementAt) { @@ -119,7 +146,7 @@ public class SpecifyTypeInDocstringIntention implements IntentionAction { } PsiReference reference = null; - PyExpression problemElement = PyUtil.findProblemElement(editor, file, PyNamedParameter.class, PyQualifiedExpression.class); + PyElementGenerator elementGenerator = PyElementGenerator.getInstance(project); if (problemElement != null) { name = problemElement.getName(); @@ -137,7 +164,7 @@ public class SpecifyTypeInDocstringIntention implements IntentionAction { final ASTNode nameNode = pyFunction.getNameNode(); if ((pyFunction != null && (problemElement instanceof PyParameter || reference != null && reference.resolve() instanceof PyParameter)) || - elementAt == nameNode.getPsi()) { + elementAt == nameNode.getPsi() || callExpression != null) { PyStringLiteralExpression docStringExpression = pyFunction.getDocStringExpression(); int startOffset; int endOffset; diff --git a/python/src/com/jetbrains/python/codeInsight/intentions/SpecifyTypeInPy3AnnotationsIntention.java b/python/src/com/jetbrains/python/codeInsight/intentions/SpecifyTypeInPy3AnnotationsIntention.java index 4f91fe89dff8..d33c9e28d47a 100644 --- a/python/src/com/jetbrains/python/codeInsight/intentions/SpecifyTypeInPy3AnnotationsIntention.java +++ b/python/src/com/jetbrains/python/codeInsight/intentions/SpecifyTypeInPy3AnnotationsIntention.java @@ -50,12 +50,25 @@ public class SpecifyTypeInPy3AnnotationsIntention implements IntentionAction { if (elementAt != null && !(elementAt.getNode().getElementType() == PyTokenTypes.IDENTIFIER)) elementAt = file.findElementAt(editor.getCaretModel().getOffset()); - PyFunction parentFunction = PsiTreeUtil.getParentOfType(elementAt, PyFunction.class); - if (parentFunction != null) { - final ASTNode nameNode = parentFunction.getNameNode(); - if (nameNode != null && nameNode.getPsi() == elementAt) { - myText = PyBundle.message("INTN.specify.returt.type.in.annotation"); - return true; + PyCallExpression callExpression = PsiTreeUtil.getParentOfType(elementAt, PyCallExpression.class); + if (callExpression != null ) { + PyAssignmentStatement assignmentStatement = PsiTreeUtil.getParentOfType(elementAt, PyAssignmentStatement.class); + if (assignmentStatement != null) { + PyType type = assignmentStatement.getAssignedValue().getType(TypeEvalContext.slow()); + if (type == null || type instanceof PyReturnTypeReference) { + myText = PyBundle.message("INTN.specify.returt.type.in.annotation"); + return true; + } + } + } + else { + PyFunction parentFunction = PsiTreeUtil.getParentOfType(elementAt, PyFunction.class); + if (parentFunction != null) { + final ASTNode nameNode = parentFunction.getNameNode(); + if (nameNode != null && nameNode.getPsi() == elementAt) { + myText = PyBundle.message("INTN.specify.returt.type.in.annotation"); + return true; + } } } @@ -159,10 +172,16 @@ public class SpecifyTypeInPy3AnnotationsIntention implements IntentionAction { } else { PsiElement elementAt = file.findElementAt(editor.getCaretModel().getOffset() - 1); + PyCallExpression callExpression = PsiTreeUtil.getParentOfType(elementAt, PyCallExpression.class); if (elementAt != null && !(elementAt.getNode().getElementType() == PyTokenTypes.IDENTIFIER)) elementAt = file.findElementAt(editor.getCaretModel().getOffset()); - callable = PsiTreeUtil.getParentOfType(elementAt, PyFunction.class); + + if (callExpression != null) { + callable = callExpression.resolveCalleeFunction(PyResolveContext.defaultContext()); + } + else + callable = PsiTreeUtil.getParentOfType(elementAt, PyFunction.class); } if (callable instanceof PyFunction && ((PyFunction)callable).getAnnotation() == null) { final String functionSignature = "def " + callable.getName() + callable.getParameterList().getText(); diff --git a/python/testData/intentions/afterTypeInDocstring2.py b/python/testData/intentions/afterTypeInDocstring2.py new file mode 100644 index 000000000000..d38dfb10211b --- /dev/null +++ b/python/testData/intentions/afterTypeInDocstring2.py @@ -0,0 +1,9 @@ +def func1(x): + """ + :rtype : object + """ + return x + +def func2(x): + y = func1(x.keys()) + return y.startswith('foo') \ No newline at end of file diff --git a/python/testData/intentions/beforeTypeInDocstring2.py b/python/testData/intentions/beforeTypeInDocstring2.py new file mode 100644 index 000000000000..7d5fc9e8d4b8 --- /dev/null +++ b/python/testData/intentions/beforeTypeInDocstring2.py @@ -0,0 +1,6 @@ +def func1(x): + return x + +def func2(x): + y = func1(x.keys()) + return y.startswith('foo') \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/PyIntentionTest.java b/python/testSrc/com/jetbrains/python/PyIntentionTest.java index af4120b4deb1..2ecf33912659 100644 --- a/python/testSrc/com/jetbrains/python/PyIntentionTest.java +++ b/python/testSrc/com/jetbrains/python/PyIntentionTest.java @@ -241,6 +241,10 @@ public class PyIntentionTest extends PyTestCase { doTest(PyBundle.message("INTN.specify.return.type")); } + public void testTypeInDocstring2() { + doTest(PyBundle.message("INTN.specify.return.type")); + } + public void testTypeInPy3Annotation() { //PY-7045 doTest(PyBundle.message("INTN.specify.type.in.annotation"), LanguageLevel.PYTHON32); }