Suppress "Unbound local variable" for targets defined in assignment expressions inside comprehensions under conditions (PEP 572) (PY-33886)

GitOrigin-RevId: ef7edfd58d6c36037903c54ed2c832696c4ff0b6
This commit is contained in:
Semyon Proshev
2019-07-02 06:52:16 +03:00
committed by intellij-monorepo-bot
parent aa1da541de
commit b600a9bebf
2 changed files with 46 additions and 1 deletions
@@ -129,7 +129,7 @@ public class PyUnboundLocalVariableInspection extends PyInspection {
return;
}
}
if (resolvedUnderWithStatement(node, resolved)) {
if (resolvedUnderWithStatement(node, resolved) || resolvedUnderAssignmentExpressionAndCondition(node, resolved)) {
return;
}
if (PyUnreachableCodeInspection.hasAnyInterruptedControlFlowPaths(node)) {
@@ -159,6 +159,17 @@ public class PyUnboundLocalVariableInspection extends PyInspection {
PsiTreeUtil.getParentOfType(resolved, PyWithStatement.class, true, ScopeOwner.class) != null;
}
private static boolean resolvedUnderAssignmentExpressionAndCondition(@NotNull PyReferenceExpression node,
@Nullable PsiElement resolved) {
return resolved instanceof PyTargetExpression &&
PyUtil.inSameFile(node, resolved) &&
resolved.getParent() instanceof PyAssignmentExpression &&
PsiTreeUtil.getParentOfType(
PsiTreeUtil.getParentOfType(resolved.getParent(), PyComprehensionElement.class, true, PyConditionalStatementPart.class),
PyConditionalStatementPart.class
) != null;
}
private static boolean isFirstUnboundRead(@NotNull PyReferenceExpression node, @NotNull ScopeOwner owner) {
final String nodeName = node.getReferencedName();
final Scope scope = ControlFlowCache.getScope(owner);
@@ -290,6 +290,40 @@ public class PyUnboundLocalVariableInspectionTest extends PyInspectionTestCase {
" print(val)");
}
// PY-33886
public void testAssignmentExpressions() {
runWithLanguageLevel(
LanguageLevel.getLatest(),
() -> doTestByText(
"def foo():\n" +
" if any((comment := line).startswith('#') for line in lines):\n" +
" print(\"First comment:\", comment)\n" +
" else:\n" +
" print(\"There are no comments\")\n" +
"\n" +
" if all((nonblank := line).strip() == '' for line in lines):\n" +
" print(\"All lines are blank\")\n" +
" else:\n" +
" print(\"First non-blank line:\", nonblank)\n" +
"\n" +
"\n" +
"def bar():\n" +
" [(comment := line).startswith('#') for line in lines]\n" +
" print(<warning descr=\"Local variable 'comment' might be referenced before assignment\">comment</warning>)\n" +
"\n" +
"\n" +
"def baz():\n" +
" while (line := input()) and any((first_digit := c).isdigit() for c in line):\n" +
" print(line, first_digit)\n" +
"\n" +
"\n" +
"def more():\n" +
" if (x := True) or (y := 'spam'):\n" +
" print(<warning descr=\"Local variable 'y' might be referenced before assignment\">y</warning>)\n"
)
);
}
@NotNull
@Override
protected Class<? extends PyInspection> getInspectionClass() {