PY-76922 Add a barebone implementation of PyIntersectionType

For the time being, it's possible to declare an intersection type
as "A & B", though it should be either wrapped in a string literal,
used with `from __future__ import annotations` or with Python 3.14+
to avoid runtime errors (and there still be warnings about unresolved
`type.__and__`) and it's not possible to generate a type hint for them.
At least until intersection types are added to the typing specification,
or we adopt something similar to "ty_extensions".

Specific use cases, such as using intersections for type narrowing,
`super()`, overloads, etc. will be addressed separately.

Space-RevId: 16c28a020f0b102f54f56ae25b5c25dc143fd76d

GitOrigin-RevId: d2581559197819fbbc49e4f75b62aa2962924938
This commit is contained in:
Mikhail Golubev
2026-01-26 17:03:32 +00:00
committed by intellij-monorepo-bot
parent c003940340
commit abe26b050d
28 changed files with 359 additions and 9 deletions
@@ -316,6 +316,7 @@ public final class PyTypeHintGenerationUtil {
isNoneType(type) ||
// Will be rendered as just Any
type instanceof PyUnsafeUnionType ||
type instanceof PyIntersectionType ||
type instanceof PyTypeParameterType) {
return;
}
@@ -921,6 +921,10 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
if (unionType != null) {
return unionType;
}
final Ref<PyType> intersectionType = getIntersectionType(resolved, context);
if (intersectionType != null) {
return intersectionType;
}
final PyType concatenateType = getConcatenateType(resolved, context);
if (concatenateType != null) {
return Ref.create(concatenateType);
@@ -1071,6 +1075,22 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
return Ref.create(new PyModuleType(firstModuleInitFile));
}
private static Ref<PyType> getIntersectionType(@NotNull PsiElement resolved, @NotNull PyTypingTypeProvider.Context context) {
if (resolved instanceof PyBinaryExpression expression && expression.getOperator() == PyTokenTypes.AND) {
PyExpression left = expression.getLeftExpression();
PyExpression right = expression.getRightExpression();
if (left == null || right == null) return null;
Ref<PyType> leftTypeRef = getType(left, context);
Ref<PyType> rightTypeRef = getType(right, context);
if (leftTypeRef == null || rightTypeRef == null) return null;
PyType intersection = PyIntersectionType.intersection(leftTypeRef.get(), rightTypeRef.get());
return intersection != null ? Ref.create(intersection) : null;
}
return null;
}
private static @Nullable Ref<PyType> getNoneType(@NotNull PyExpression typeHint, @NotNull PsiElement resolved) {
if (typeHint instanceof PyNoneLiteralExpression ||
typeHint instanceof PyReferenceExpression && PyNames.NONE.equals(typeHint.getText())) {
@@ -198,6 +198,12 @@ public abstract class PyTypeRenderer extends PyTypeVisitorExt<@NotNull HtmlChunk
return visitUnknownType();
}
@Override
public @NotNull HtmlChunk visitPyIntersectionType(com.jetbrains.python.psi.types.@NotNull PyIntersectionType intersectionType) {
// There is no way to represent intersections through the standard type hints at the moment
return visitUnknownType();
}
@Override
public @NotNull HtmlChunk visitPySelfType(@NotNull PySelfType selfType) {
HtmlChunk selfTypeRender = className(isRenderingFqn() ? "typing.Self" : "Self"); //NON-NLS
@@ -394,6 +400,11 @@ public abstract class PyTypeRenderer extends PyTypeVisitorExt<@NotNull HtmlChunk
return result.toFragment();
}
@Override
public @NotNull HtmlChunk visitPyIntersectionType(@NotNull PyIntersectionType intersectionType) {
return renderList(ContainerUtil.map(intersectionType.getMembers(), this::render), " & ");
}
private static @Nullable Pair<@NotNull List<PyLiteralType>, @NotNull List<PyType>> extractLiterals(@NotNull PyUnionType type) {
final Collection<PyType> members = type.getMembers();
@@ -504,6 +515,9 @@ public abstract class PyTypeRenderer extends PyTypeVisitorExt<@NotNull HtmlChunk
case " | " -> {
yield styled(separator, PyHighlighter.PY_OPERATION_SIGN);
}
case " & " -> {
yield styled(separator, PyHighlighter.PY_OPERATION_SIGN);
}
default -> {
yield escaped(separator);
}
@@ -431,6 +431,9 @@ public abstract class PyUnresolvedReferencesVisitor extends PyInspectionVisitor
if (type instanceof PyUnsafeUnionType weakUnionType) {
return ContainerUtil.exists(weakUnionType.getMembers(), member -> ignoreUnresolvedMemberForType(member, reference, name));
}
if (type instanceof PyIntersectionType intersectionType) {
return ContainerUtil.exists(intersectionType.getMembers(), member -> ignoreUnresolvedMemberForType(member, reference, name));
}
if (PyTypeChecker.isUnknown(type, myTypeEvalContext)) {
// this almost always means that we don't know the type, so don't show an error in this case
return true;
@@ -60,7 +60,7 @@ public class PyFileElementType extends IStubFileElementType<PyFileStub> {
@Override
public int getStubVersion() {
// Don't forget to update versions of indexes that use the updated stub-based elements
return 106;
return 107;
}
@Override
@@ -53,8 +53,8 @@ public final class PyTypingAliasStubType extends CustomTargetExpressionStubType<
(?x)
\\s*
\\S+(\\[.*])? # initial type like: "list[int]"
(\\s*\\|\\s* # union operator: " | "
\\S+(\\[.*])? # type between union operator
(\\s*[|&]\\s* # union or intersection operators: " | " or " & "
\\S+(\\[.*])? # operand types
)* # repeating
\\s*
""",
@@ -174,7 +174,7 @@ public final class PyTypingAliasStubType extends CustomTargetExpressionStubType<
@Override
public void visitPyBinaryExpression(@NotNull PyBinaryExpression node) {
if (node.getOperator() != PyTokenTypes.OR) {
if (!(node.getOperator() == PyTokenTypes.OR || node.getOperator() == PyTokenTypes.AND)) {
illegal[0] = true;
return;
}
@@ -136,6 +136,9 @@ public final class PyABCUtil {
if (type instanceof PyUnsafeUnionType) {
return PyTypeUtil.toStream(type).nonNull().anyMatch(it -> isSubtype(it, superClassName, context));
}
if (type instanceof PyIntersectionType) {
return PyTypeUtil.toStream(type).nonNull().anyMatch(it -> isSubtype(it, superClassName, context));
}
return false;
}
@@ -149,6 +149,11 @@ public abstract class PyCloningTypeVisitor extends PyTypeVisitorExt<PyType> {
return PyUnsafeUnionType.unsafeUnion(ContainerUtil.map(unsafeUnionType.getMembers(), type -> clone(type)));
}
@Override
public PyType visitPyIntersectionType(@NotNull PyIntersectionType intersectionType) {
return PyIntersectionType.intersection(ContainerUtil.map(intersectionType.getMembers(), type -> clone(type)));
}
@Override
public PyType visitPyTypingNewType(@NotNull PyTypingNewType typingNewType) {
return typingNewType;
@@ -0,0 +1,103 @@
package com.jetbrains.python.psi.types
import com.intellij.openapi.util.NlsSafe
import com.intellij.psi.PsiElement
import com.intellij.util.ProcessingContext
import com.intellij.util.SmartList
import com.jetbrains.python.psi.AccessDirection
import com.jetbrains.python.psi.PyExpression
import com.jetbrains.python.psi.resolve.PyResolveContext
import com.jetbrains.python.psi.resolve.RatedResolveResult
import org.jetbrains.annotations.ApiStatus
import java.util.Collections
@ApiStatus.Experimental
class PyIntersectionType private constructor(members: Collection<PyType?>) : PyType {
val members: Set<PyType?> = Collections.unmodifiableSet<PyType?>(LinkedHashSet(members))
override fun resolveMember(
name: String,
location: PyExpression?,
direction: AccessDirection,
resolveContext: PyResolveContext,
): List<RatedResolveResult?>? {
val ret = SmartList<RatedResolveResult?>()
var allNulls = true
for (member in members) {
if (member != null) {
val result = member.resolveMember(name, location, direction, resolveContext)
if (result != null) {
allNulls = false
ret.addAll(result)
}
}
}
return if (allNulls) null else ret
}
override fun getCompletionVariants(completionPrefix: String?, location: PsiElement?, context: ProcessingContext?): Array<out Any> {
return members.flatMap { it?.getCompletionVariants(completionPrefix, location, context)?.asList() ?: emptyList() }
.distinct()
.toTypedArray()
}
override fun getName(): @NlsSafe String {
return members.joinToString(separator = " & ") { it?.name ?: "Any" }
}
override fun isBuiltin(): Boolean {
return members.all { it != null && it.isBuiltin }
}
override fun assertValid(message: String?) {
for (member in members) {
member?.assertValid(message)
}
}
override fun <T> acceptTypeVisitor(visitor: PyTypeVisitor<T>): T {
if (visitor is PyTypeVisitorExt<T>) {
return visitor.visitPyIntersectionType(this)
}
return visitor.visitPyType(this)
}
override fun equals(other: Any?): Boolean {
if (this === other) return true
if (javaClass != other?.javaClass) return false
other as PyIntersectionType
return members == other.members
}
override fun hashCode(): Int {
return members.hashCode()
}
override fun toString(): String {
return "PyIntersectionType: $name"
}
companion object {
@JvmStatic
fun intersection(vararg types: PyType?): PyType? {
return intersection(types.toList())
}
@JvmStatic
fun intersection(types: Collection<PyType?>): PyType? {
val newMembers = buildSet {
for (member in types) {
if (member is PyIntersectionType) {
addAll(member.members)
}
else {
add(member)
}
}
}
return if (newMembers.size > 1) PyIntersectionType(newMembers) else newMembers.firstOrNull()
}
}
}
@@ -189,6 +189,11 @@ public final class PyRecursiveTypeVisitor extends PyTypeVisitorExt<PyRecursiveTy
return Collections.unmodifiableList(new ArrayList<>(unsafeUnionType.getMembers()));
}
@Override
public @NotNull List<@Nullable PyType> visitPyIntersectionType(@NotNull PyIntersectionType intersectionType) {
return Collections.unmodifiableList(new ArrayList<>(intersectionType.getMembers()));
}
@Override
public @NotNull List<@Nullable PyType> visitPyUnpackedTupleType(@NotNull PyUnpackedTupleType unpackedTupleType) {
return unpackedTupleType.getElementTypes();
@@ -189,6 +189,14 @@ public final class PyTypeChecker {
return Optional.of(match(weakUnionType, actual, context));
}
if (actual instanceof PyIntersectionType intersectionType) {
return Optional.of(match(expected, intersectionType, context));
}
if (expected instanceof PyIntersectionType intersectionType) {
return Optional.of(match(intersectionType, actual, context));
}
if (expected instanceof PyClassType && actual instanceof PyClassType) {
Optional<Boolean> match = match((PyClassType)expected, (PyClassType)actual, context);
if (match.isPresent()) {
@@ -521,10 +529,6 @@ public final class PyTypeChecker {
return ContainerUtil.and(actual.getMembers(), type -> match(expected, type, context).orElse(false));
}
private static boolean match(@NotNull PyType expected, @NotNull PyUnsafeUnionType actual, @NotNull MatchContext context) {
return ContainerUtil.or(actual.getMembers(), type -> match(expected, type, context).orElse(false));
}
private static @NotNull Optional<Boolean> match(@NotNull PyTupleType expected,
@NotNull PyUnionType actual,
@NotNull MatchContext context) {
@@ -547,6 +551,18 @@ public final class PyTypeChecker {
return ContainerUtil.or(expected.getMembers(), type -> match(type, actual, context).orElse(true));
}
private static boolean match(@NotNull PyType expected, @NotNull PyIntersectionType actual, @NotNull MatchContext context) {
return ContainerUtil.or(actual.getMembers(), type -> match(expected, type, context).orElse(false));
}
private static boolean match(@NotNull PyIntersectionType expected, @NotNull PyType actual, @NotNull MatchContext context) {
return ContainerUtil.all(expected.getMembers(), type -> match(type, actual, context).orElse(true));
}
private static boolean match(@NotNull PyType expected, @NotNull PyUnsafeUnionType actual, @NotNull MatchContext context) {
return ContainerUtil.or(actual.getMembers(), type -> match(expected, type, context).orElse(false));
}
private static boolean match(@NotNull PyUnsafeUnionType expected, @NotNull PyType actual, @NotNull MatchContext context) {
if (expected.getMembers().contains(actual)) {
return true;
@@ -1141,6 +1157,9 @@ public final class PyTypeChecker {
if (type instanceof PyUnsafeUnionType weakUnion) {
return ContainerUtil.exists(weakUnion.getMembers(), member -> isUnknown(member, genericsAreUnknown, context));
}
if (type instanceof PyIntersectionType intersectionType) {
return ContainerUtil.exists(intersectionType.getMembers(), member -> isUnknown(member, genericsAreUnknown, context));
}
return false;
}
@@ -160,6 +160,9 @@ public final class PyTypeUtil {
if (type instanceof PyUnsafeUnionType weakUnionType) {
return StreamEx.of(weakUnionType.getMembers());
}
if (type instanceof PyIntersectionType intersectionType) {
return StreamEx.of(intersectionType.getMembers());
}
return StreamEx.of(type);
}
@@ -221,6 +224,11 @@ public final class PyTypeUtil {
return Collectors.collectingAndThen(Collectors.toList(), PyUnsafeUnionType::unsafeUnion);
}
@ApiStatus.Experimental
public static @NotNull Collector<@Nullable PyType, ?, @Nullable PyType> toIntersection() {
return Collectors.collectingAndThen(Collectors.toList(), PyIntersectionType::intersection);
}
public static @NotNull Collector<@Nullable PyType, ?, @Nullable PyType> toUnion(@Nullable PyType streamSource) {
return toUnion(streamSource instanceof PyUnsafeUnionType ? PyUnsafeUnionType::unsafeUnion : PyUnionType::union);
}
@@ -47,6 +47,10 @@ public abstract class PyTypeVisitorExt<T> extends PyTypeVisitor<T> {
return visitPyType(unsafeUnionType);
}
public T visitPyIntersectionType(@NotNull PyIntersectionType intersectionType) {
return visitPyType(intersectionType);
}
public T visitPyTypingNewType(@NotNull PyTypingNewType typingNewType) {
return visitPyClassType(typingNewType);
}
@@ -0,0 +1,51 @@
from typing import Any
class A:
def __iter__(self):
return self
def __next__(self):
return 42
class B:
def __iter__(self):
return self
def __next__(self):
return 42
class C:
pass
def all_intersection_members_match_no_any(iterable: "A & B"):
for _ in iterable:
pass
def some_intersection_members_match_no_any(iterable: "A & B & None"):
for _ in iterable:
pass
def all_intersection_members_dont_match_no_any(iterable: "C & None"):
for _ in <warning descr="Expected type 'collections.Iterable', got 'C & None' instead">iterable</warning>:
pass
def all_intersection_members_match_with_any(iterable: "A & B & Any"):
for _ in iterable:
pass
def some_intersection_members_match_with_any(iterable: "A & B & None & Any"):
for _ in iterable:
pass
def all_intersection_members_dont_match_with_any(iterable: "C & None & Any"):
for _ in iterable:
pass
@@ -0,0 +1,25 @@
from typing import Any
class A:
def method(self):
pass
class B:
def method(self):
pass
# Using string annotations to suppress the warnings about unresolved '__and__' in type
def intersection_with_all_compatible_types(x: 'A <warning descr="Class 'type' does not define '__and__', so the '&' operator cannot be used on its instances">&</warning> B'):
x.method()
def intersection_with_some_incompatible_types(x: 'A <warning descr="Class 'type' does not define '__and__', so the '&' operator cannot be used on its instances">&</warning> None'):
x.method()
def intersection_with_all_incompatible_types(x: 'object <warning descr="Class 'type' does not define '__and__', so the '&' operator cannot be used on its instances">&</warning> None'):
x.<warning descr="Cannot find reference 'method' in 'object & None'">method</warning>()
def intersection_with_incompatible_types_and_any(x: 'Any & None'):
x.method()
@@ -0,0 +1,3 @@
from typing import Any, <warning descr="Unused import statement 'Optional'">Optional</warning>
x: "str <warning descr="Class 'type' does not define '__and__', so the '&' operator cannot be used on its instances">&</warning> Any"
@@ -0,0 +1,2 @@
x: str & int
v<caret>ar = x
@@ -0,0 +1,4 @@
from typing import Any
x: str & int
var: [Any]<caret> = x
@@ -0,0 +1 @@
<html><body><div class="bottom"><icon src="AllIcons.Nodes.Package"/>&nbsp;<code><a href="psi_element://#module#IntersectionType">IntersectionType</a></code></div><div class="definition"><pre><span style="color:#000000;">var</span><span style="">: </span><span style="color:#000000;"><span style="color:#000080;"><a href="psi_element://#typename#int">int</a></span><span style=""> &amp; </span><span style="color:#000080;"><a href="psi_element://#typename#str">str</a></span></span></pre></div></body></html>
@@ -0,0 +1 @@
va<the_ref>r: int & str
+1 -1
View File
@@ -19,7 +19,7 @@ S6_notOk = b"int"
bin1_ok = int | str
bin2_ok = int | str | bool | None
bin3_ok = Union[str, bool] | None
bin4_notOk = str & int
bin4_ok = str & int
list_notOk = [int, str]
bin5_notOk: int | str = "foo"
@@ -0,0 +1,3 @@
Alias = list['int & str']
x: Alias
@@ -879,6 +879,11 @@ public class Py3QuickDocTest extends LightMarkedTestCase {
});
}
// PY-76922
public void testIntersectionType() {
checkHTMLOnly();
}
@Override
protected String getTestDataPath() {
return super.getTestDataPath() + "/quickdoc/";
@@ -6877,6 +6877,22 @@ public class PyTypingTest extends PyTestCase {
""");
}
// PY-76922
public void testIntersectionTypeParsing() {
doTest("int & str", """
expr: int & str
""");
}
// PY-76922
public void testLegacyTypeAliasesWithQuotedIntersectionTypesPreservedInStubs() {
doMultiFileStubAwareTest("list[int & str]", """
from mod import x
expr = x
""");
}
private void doTestNoInjectedText(@NotNull String text) {
myFixture.configureByText(PythonFileType.INSTANCE, text);
final InjectedLanguageManager languageManager = InjectedLanguageManager.getInstance(myFixture.getProject());
@@ -3139,6 +3139,11 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
doTest();
}
// PY-76922
public void testIntersectionImplicitProtocolMatching() {
doTest();
}
// PY-76822
public void testProtocolWithAssignedPropertyInMethod() {
doTestByText("""
@@ -4174,5 +4179,39 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
async_for(<warning descr="Type 'list[int]' doesn't have expected attribute '__aiter__'">[1, 2, 3]</warning>)
""");
}
// PY-76922
public void testIntersectionType() {
doTestByText("""
int_and_str: int & str
str_and_int: int & str
int_or_str: int | str
n: int = int_and_str
s: str = int_and_str
int_and_str = <warning descr="Expected type 'int & str', got 'int' instead">n</warning>
int_and_str = <warning descr="Expected type 'int & str', got 'str' instead">s</warning>
int_or_str = int_and_str
int_and_str = <warning descr="Expected type 'int & str', got 'int | str' instead">int_or_str</warning>
str_and_int = int_and_str
int_and_str = str_and_int
class A: pass
class B: pass
class C(A, B): pass
a_and_b: A & B
a_and_b = <warning descr="Expected type 'A & B', got 'A' instead">A()</warning>
a_and_b = <warning descr="Expected type 'A & B', got 'B' instead">B()</warning>
a_and_b = C()
a: A = a_and_b
b: B = a_and_b
c: C = <warning descr="Expected type 'C', got 'A & B' instead">a_and_b</warning>
""");
}
}
@@ -561,4 +561,9 @@ public class Py3UnresolvedReferencesInspectionTest extends PyInspectionTestCase
super().<warning descr="Cannot find reference 'non_existing' in 'A | ABC'">non_existing</warning>()
""");
}
// PY-76922
public void testIntersectionMemberAttributeAccess() {
doTest();
}
}
@@ -98,6 +98,11 @@ public class PyUnusedImportTest extends PyTestCase {
runWithLanguageLevel(LanguageLevel.PYTHON34, this::doTest);
}
// PY-76922
public void testAnnotationWithQuotedIntersectionUsesImportsFromTyping() {
doTest();
}
public void testSuppressedForUnreachableCode() {
doTest();
}
@@ -336,6 +336,11 @@ public class PyAnnotateVariableTypeIntentionTest extends PyIntentionTestCase {
doAnnotationTest();
}
// PY-76922
public void testIntersectionTypeIsNotDenotable() {
doAnnotationTest();
}
private void doAnnotationTest() {
doTest(LanguageLevel.getLatest());
}