PY-86223 Fix context for type hint code fragments

Use the containing PSI element as the context for the type hint code fragment instead of the whole file. This allows resolving symbols within the local scopes.

When looking for a start instruction in `PyDefUseUtil.getLatestDefs()`, use the parent element's instruction if the element itself is missing from the CFG.

Fix resolution of string type hints and comments by allowing forward references in PyExpressionCodeFragment.

Space-RevId: b44603176d5b50b20f9a5964714d3661d71c8a08

GitOrigin-RevId: 715bcd011543c060e0dff965d94fbd8cc6f7e591
This commit is contained in:
Petr
2026-01-22 17:32:16 +00:00
committed by intellij-monorepo-bot
parent 5228aa38cb
commit b0697d6b99
7 changed files with 140 additions and 16 deletions
@@ -1421,7 +1421,7 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
private static @Nullable PyExpression toExpression(@NotNull String contents, @NotNull PsiElement anchor) {
final PsiFile file = FileContextUtil.getContextFile(anchor);
if (file == null) return null;
PyExpression fragment = PyUtil.createExpressionFromFragment(contents, file);
PyExpression fragment = PyUtil.createExpressionFromFragment(contents, anchor);
if (fragment != null) {
fragment.getContainingFile().putUserData(FRAGMENT_OWNER, anchor);
}
@@ -426,13 +426,13 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere
if (target instanceof PyElement && context.allowDataFlow(anchor)) {
final ScopeOwner scopeOwner = ScopeUtil.getScopeOwner(anchor);
final String name = ((PyElement)target).getName();
if (scopeOwner != null &&
name != null &&
(!ScopeUtil.getElementsOfAccessType(name, scopeOwner, ReadWriteInstruction.ACCESS.ASSERTTYPE).isEmpty()
|| target instanceof PyTargetExpression || target instanceof PyNamedParameter)) {
final PyType type = getTypeByControlFlow(name, context, anchor, scopeOwner);
if (type != null) {
return type;
if (scopeOwner != null && name != null) {
if (!ScopeUtil.getElementsOfAccessType(name, scopeOwner, ReadWriteInstruction.ACCESS.ASSERTTYPE).isEmpty() ||
(target instanceof PyTargetExpression || target instanceof PyNamedParameter) && ScopeUtil.getScopeOwner(target) == scopeOwner) {
final PyType type = getTypeByControlFlow(name, context, anchor, scopeOwner);
if (type != null) {
return type;
}
}
}
}
@@ -421,8 +421,11 @@ public final class PyResolveUtil {
if (PyiUtil.isInsideStub(element)) {
return true;
}
// Forward references are allowed in annotations according to PEP 563
PsiFile file = element.getContainingFile();
if (file instanceof PyExpressionCodeFragment) {
return true;
}
// Forward references are allowed in annotations according to PEP 563
if (file instanceof PyFile pyFile) {
boolean nonEagerEvaluationEnabled = pyFile.hasImportFromFuture(FutureFeature.ANNOTATIONS) ||
pyFile.getLanguageLevel().isAtLeast(LanguageLevel.PYTHON314);
@@ -51,18 +51,20 @@ public final class PyDefUseUtil {
boolean acceptTypeAssertions,
boolean acceptImplicitImports,
@NotNull TypeEvalContext context) {
return getLatestDefs(ControlFlowCache.getControlFlow(block), varName, anchor, acceptTypeAssertions, acceptImplicitImports, context);
return getLatestDefs(ControlFlowCache.getControlFlow(block), block, varName, anchor, acceptTypeAssertions, acceptImplicitImports,
context);
}
public static @NotNull List<Instruction> getLatestDefs(@NotNull PyControlFlow controlFlow,
@NotNull ScopeOwner scopeOwner,
@NotNull String varName,
@NotNull PsiElement anchor,
boolean acceptTypeAssertions,
boolean acceptImplicitImports,
@NotNull TypeEvalContext context) {
final Instruction[] instructions = controlFlow.getInstructions();
int startNum = findStartInstructionId(anchor, controlFlow);
int startNum = findStartInstructionId(anchor, controlFlow, scopeOwner);
if (startNum < 0) {
return Collections.emptyList();
}
@@ -152,13 +154,19 @@ public final class PyDefUseUtil {
return varQname.getComponentCount() > elementQname.getComponentCount() && varQname.matchesPrefix(elementQname);
}
private static int findStartInstructionId(@NotNull PsiElement startAnchor, @NotNull PyControlFlow flow) {
private static int findStartInstructionId(@NotNull PsiElement startAnchor, @NotNull PyControlFlow flow, @NotNull ScopeOwner scopeOwner) {
PsiElement realCfgAnchor = startAnchor;
final PyAugAssignmentStatement augAssignment = PyAugAssignmentStatementNavigator.getStatementByTarget(startAnchor);
if (augAssignment != null) {
realCfgAnchor = augAssignment;
}
int instr = flow.getInstruction(realCfgAnchor);
int instr = -1;
for (PsiElement element = realCfgAnchor; element != null && element != scopeOwner; element = element.getParent()) {
instr = flow.getInstruction(element);
if (instr >= 0) {
break;
}
}
if (instr < 0) {
return instr;
}
@@ -983,7 +983,7 @@ public class Py3TypeTest extends PyTestCase {
if (a := input()) in ("abba", False):
expr = a
""");
// PY-83625
doTest("Literal[\"b\", \"c\"]",
"""
@@ -4341,8 +4341,16 @@ public class Py3TypeTest extends PyTestCase {
// PY-74257
public void testNotProperlyImportedQualifiedNameInTypeHint() {
doMultiFileTest("Any", """
from lib import f
// TODO lib.py can be unstubbed
//doMultiFileTest("Any", """
// from lib import f
//
// expr = f()
// """);
doTest("Any", """
import pkg
def f() -> "pkg.subpkg.mod.MyClass": ...
expr = f()
""");
@@ -4709,6 +4717,88 @@ public class Py3TypeTest extends PyTestCase {
""");
}
// PY-86223
public void testQuotedTypeParameterInTypeHint() {
doTest("T", """
def foo[T](p: "T"):
expr = p
"""
);
}
// PY-86223
public void testGenericTypeWithQuotedTypeParameterInTypeHint() {
doTest("list[T]", """
def foo[T](p: list["T"]):
expr = p
"""
);
}
// PY-86223
public void testQuotedGenericTypeWithTypeParameterInTypeHint() {
doTest("list[T]", """
def foo[T](p: "list[T]"):
expr = p
"""
);
}
// PY-86223
public void testQuotedReferenceToLocalClassInTypeHint() {
doTest("tuple[A, B]", """
def outer():
class A: ...
def inner(a: "A", b: "B"):
expr = (a, b)
class B: ...
"""
);
}
public void testQuotedForwardReferenceInTypeHint() {
doTest("MyClass", """
def foo(x: "MyClass"):
expr = x
class MyClass: ...
"""
);
}
public void testGenericTypeWithQuotedForwardReferenceInTypeHint() {
doTest("list[MyClass]", """
def foo(x: list["MyClass"]):
expr = x
class MyClass: ...
"""
);
}
public void testQuotedGenericTypeWithForwardReferenceInTypeHint() {
doTest("list[MyClass]", """
def foo(x: "list[MyClass]"):
expr = x
class MyClass: ...
"""
);
}
public void testIncompleteQualifiedNameClashesWithLocalVariable() {
doTest("str", """
class MyClass:
foo = 'spam'
def f(foo):
_ = foo.illegal
expr = MyClass.foo
""");
}
private void doTest(final String expectedType, final String text) {
myFixture.configureByText(PythonFileType.INSTANCE, text);
final PyExpression expr = myFixture.findElementByText("expr", PyExpression.class);
@@ -4243,6 +4243,17 @@ public class PyTypeTest extends PyTestCase {
""");
}
public void testQuotedForwardReferenceInTypeComment() {
doTest("MyClass", """
def foo(x):
# type: (MyClass) -> None
expr = x
class MyClass: ...
"""
);
}
private static List<TypeEvalContext> getTypeEvalContexts(@NotNull PyExpression element) {
return ImmutableList.of(TypeEvalContext.codeAnalysis(element.getProject(), element.getContainingFile()).withTracing(),
TypeEvalContext.userInitiated(element.getProject(), element.getContainingFile()).withTracing());
@@ -3252,6 +3252,18 @@ public class PyTypeHintsInspectionTest extends PyInspectionTestCase {
""");
}
// PY-86223
public void testGenericTypeWithQuotedTypeParameterInTypeHint() {
doTestByText("""
from typing import assert_type
def foo[T](x: list["T"]):
assert_type(x, list[T])
assert_type(x, list["T"])
""");
}
@NotNull
@Override
protected Class<? extends PyInspection> getInspectionClass() {