diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyNamedParameterImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyNamedParameterImpl.java index bb0da26b4d07..37077b976a56 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyNamedParameterImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyNamedParameterImpl.java @@ -205,36 +205,43 @@ public class PyNamedParameterImpl extends PyBaseElementImpl types = new ArrayList<>(); - final PyResolveContext resolveContext = PyResolveContext.defaultContext(context); - final PyCallableParameter parameter = PyCallableParameterImpl.psi(this); + // Guess the type under an assumed type to prevent recursion + final PyType assumedResult = context.assumeType(this, null, ctx -> { + // Guess the type from file-local calls + if (ctx.allowCallContext(this)) { + final List types = new ArrayList<>(); + final PyResolveContext resolveContext = PyResolveContext.defaultContext(ctx); + final PyCallableParameter parameter = PyCallableParameterImpl.psi(this); - processLocalCalls( - func, call -> { - StreamEx - .of(call.multiMapArguments(resolveContext)) - .flatCollection(mapping -> mapping.getMappedParameters().entrySet()) - .filter(entry -> parameter.equals(entry.getValue())) - .map(Map.Entry::getKey) - .nonNull() - .map(context::getType) - .nonNull() - .forEach(types::add); - return true; + processLocalCalls( + func, call -> { + StreamEx + .of(call.multiMapArguments(resolveContext)) + .flatCollection(mapping -> mapping.getMappedParameters().entrySet()) + .filter(entry -> parameter.equals(entry.getValue())) + .map(Map.Entry::getKey) + .nonNull() + .map(ctx::getType) + .nonNull() + .forEach(types::add); + return true; + } + ); + + if (!types.isEmpty()) { + return PyUnionType.createWeakType(PyUnionType.union(types)); } - ); - - if (!types.isEmpty()) { - return PyUnionType.createWeakType(PyUnionType.union(types)); } - } - if (context.maySwitchToAST(this)) { - final PyType typeFromUsages = getTypeFromUsages(context); - if (typeFromUsages != null) { - return typeFromUsages; + if (ctx.maySwitchToAST(this)) { + final PyType typeFromUsages = getTypeFromUsages(ctx); + if (typeFromUsages != null) { + return typeFromUsages; + } } + return null; + }); + if (assumedResult != null) { + return assumedResult; } } } diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index 40d1c79b6866..4122c35dcda9 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -2,6 +2,8 @@ package com.jetbrains.python; import com.google.common.collect.ImmutableList; +import com.intellij.openapi.util.RecursionManager; +import com.intellij.openapi.util.registry.Registry; import com.intellij.testFramework.LightProjectDescriptor; import com.jetbrains.python.documentation.docstrings.DocStringFormat; import com.jetbrains.python.fixtures.PyTestCase; @@ -783,6 +785,8 @@ public class PyTypeTest extends PyTestCase { } public void testParameterFromUsages() { + RecursionManager.assertOnRecursionPrevention(myFixture.getTestRootDisposable()); + Registry.get("python.use.better.control.flow.type.inference").setValue(true); final String text = """ def foo(bar): expr = bar