diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java index e36fe0b744dc..c3c2bab63836 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java @@ -876,8 +876,10 @@ public class PyClassTypeImpl extends UserDataHolderBase implements PyClassType { PyClassTypeImpl classType = (PyClassTypeImpl)o; + // Cheap fields first, then the memoized-hashCode fast path guarding the deep myTypeArguments walk. if (myIsDefinition != classType.myIsDefinition) return false; if (!myClass.equals(classType.myClass)) return false; + if (hashCode() != classType.hashCode()) return false; if (!myTypeArguments.equals(classType.myTypeArguments)) return false; return true; diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCompositeTypeBase.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCompositeTypeBase.kt new file mode 100644 index 000000000000..6b805eb70ee9 --- /dev/null +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCompositeTypeBase.kt @@ -0,0 +1,26 @@ +// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +package com.jetbrains.python.psi.types + +import org.jetbrains.annotations.ApiStatus + +/** + * Shared base for set-of-members [PyCompositeType]s ([PyUnionType], [PyUnsafeUnionType], [PyIntersectionType]). + * Provides the single memoized `hashCode` + fast-path `equals`, avoiding the structural-equality + * storm on deeply nested composite types. + */ +@ApiStatus.Internal +abstract class PyCompositeTypeBase : PyCompositeType { + /** Members compared for equality/hashing; must be effectively immutable (the hash code is memoized from it). */ + protected abstract val memberSet: Set + + private val cachedHashCode: Int by lazy(LazyThreadSafetyMode.PUBLICATION) { memberSet.hashCode() } + + final override fun hashCode(): Int = cachedHashCode + + final override fun equals(other: Any?): Boolean { + if (this === other) return true + if (other == null || javaClass != other.javaClass) return false // different kinds are never equal + other as PyCompositeTypeBase + return cachedHashCode == other.cachedHashCode && memberSet == other.memberSet + } +} diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyIntersectionType.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyIntersectionType.kt index 358c1cc50998..07eb5c0dd7b1 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyIntersectionType.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyIntersectionType.kt @@ -12,8 +12,11 @@ import org.jetbrains.annotations.ApiStatus import java.util.Collections @ApiStatus.Experimental -class PyIntersectionType private constructor(members: Collection) : PyCompositeType { - override val members: Set = Collections.unmodifiableSet(LinkedHashSet(members)) +class PyIntersectionType private constructor(members: Collection) : PyCompositeTypeBase() { + override val memberSet: Set = Collections.unmodifiableSet(LinkedHashSet(members)) + + override val members: Set + get() = memberSet override fun resolveMember( name: String, @@ -58,19 +61,6 @@ class PyIntersectionType private constructor(members: Collection) : PyC 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" } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyUnionType.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyUnionType.java index 1ffcf7f1158d..cfef49843fb5 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyUnionType.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyUnionType.java @@ -28,11 +28,10 @@ import java.util.Set; import java.util.function.Function; import java.util.stream.Collectors; -import static com.jetbrains.python.psi.types.PyTypeUtilKt.isAnyOrUnknown; import static com.jetbrains.python.psi.types.PyTypeUtilKt.isUnknown; -public class PyUnionType implements PyCompositeType { +public class PyUnionType extends PyCompositeTypeBase { @ApiStatus.Internal public static boolean isStrictSemanticsEnabled() { @@ -46,6 +45,11 @@ public class PyUnionType implements PyCompositeType { myMembers = new LinkedHashSet<>(members); } + @Override + protected @NotNull Set<@Nullable PyType> getMemberSet() { + return Collections.unmodifiableSet(myMembers); + } + @Override public @Nullable List resolveMember(@NotNull String name, @Nullable PyExpression location, @@ -66,7 +70,9 @@ public class PyUnionType implements PyCompositeType { } @Override - public Object[] getCompletionVariants(String completionPrefix, PsiElement location, ProcessingContext context) { + public Object @NotNull [] getCompletionVariants(String completionPrefix, + @NotNull PsiElement location, + @NotNull ProcessingContext context) { Set variants = new HashSet<>(); for (PyType member : myMembers) { if (member != null) { @@ -248,19 +254,6 @@ public class PyUnionType implements PyCompositeType { return union(ContainerUtil.filter(getMembers(), it -> !isUnknown(it))); } - @Override - public boolean equals(Object other) { - if (other instanceof PyUnionType otherType) { - return myMembers.equals(otherType.myMembers); - } - return false; - } - - @Override - public int hashCode() { - return myMembers.hashCode(); - } - @Override public String toString() { return "PyUnionType: " + getName(); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyUnsafeUnionType.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyUnsafeUnionType.kt index 614b08b15de2..dcf9cf14449c 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyUnsafeUnionType.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyUnsafeUnionType.kt @@ -51,9 +51,11 @@ import java.util.Collections * @see PyUnionType.createWeakType */ @ApiStatus.Experimental -class PyUnsafeUnionType private constructor(members: Collection) : PyCompositeType { - override val members: Set = LinkedHashSet(members) - get() = Collections.unmodifiableSet(field) +class PyUnsafeUnionType private constructor(members: Collection) : PyCompositeTypeBase() { + override val memberSet: Set = Collections.unmodifiableSet(LinkedHashSet(members)) + + override val members: Set + get() = memberSet override fun resolveMember( name: String, @@ -98,19 +100,6 @@ class PyUnsafeUnionType private constructor(members: Collection) : PyCo return visitor.visitPyType(this) } - override fun equals(other: Any?): Boolean { - if (this === other) return true - if (javaClass != other?.javaClass) return false - - other as PyUnsafeUnionType - - return members == other.members - } - - override fun hashCode(): Int { - return members.hashCode() - } - override fun toString(): String { return "PyUnsafeUnionType: $name" } diff --git a/python/testSrc/com/jetbrains/python/types/PyCompositeTypeEqualityPerformanceTest.kt b/python/testSrc/com/jetbrains/python/types/PyCompositeTypeEqualityPerformanceTest.kt new file mode 100644 index 000000000000..7fc47d21d0ae --- /dev/null +++ b/python/testSrc/com/jetbrains/python/types/PyCompositeTypeEqualityPerformanceTest.kt @@ -0,0 +1,164 @@ +// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +package com.jetbrains.python.types + +import com.intellij.idea.TestFor +import com.intellij.psi.PsiElement +import com.intellij.util.ProcessingContext +import com.jetbrains.python.allure.Components +import com.jetbrains.python.allure.Layers +import com.jetbrains.python.allure.Subsystems +import com.jetbrains.python.fixtures.PyCodeInsightTestCase +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 com.jetbrains.python.psi.types.PyClassTypeImpl +import com.jetbrains.python.psi.types.PyCompositeTypeBase +import com.jetbrains.python.psi.types.PyIntersectionType +import com.jetbrains.python.psi.types.PyType +import com.jetbrains.python.psi.types.PyUnionType +import com.jetbrains.python.psi.types.PyUnsafeUnionType +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertFalse +import org.junit.jupiter.api.Assertions.assertNotEquals +import org.junit.jupiter.api.Assertions.assertTimeoutPreemptively +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Test +import java.time.Duration +import java.util.concurrent.atomic.AtomicLong + +/** + * Guards [PyCompositeTypeBase]'s memoized `hashCode` and fast-path `equals` against the structural-equality + * storm on deeply nested composite types (unions, unsafe unions, intersections) that froze the UI (PY-90901). + * + * How it works: the fake [CountingLeaf]/[Wrapper] types tally every `hashCode`/`equals` call they receive. + * [buildSharedDag] builds a DAG (Directed Acyclic Graph), not a tree: each level wraps the *same* child in two + * distinct wrappers, so the root has only `depth` distinct nodes yet is reachable via 2^depth paths. A naive + * `hashCode` re-walks every path — O(2^depth), ~2^23 calls at [DEPTH] = 22 — while a memoized one visits each + * node once — O(depth). The tests assert the counter stays under [LINEAR_BOUND]: the O(depth) fix clears it with + * a huge margin, the O(2^depth) original blows past it. + * Assertions use this deterministic counter and run against every composite kind. + */ +@Layers.Functional +@TestFor(issues = ["PY-90901", "PY-89956"], + classes = [PyCompositeTypeBase::class, PyUnionType::class, PyUnsafeUnionType::class, PyIntersectionType::class, PyClassTypeImpl::class]) +class PyCompositeTypeEqualityPerformanceTest : PyCodeInsightTestCase() { + + private val counter = AtomicLong() + + private val composites: List) -> PyType>> = listOf( + "PyUnionType" to { members -> PyUnionType.union(members)!! }, + "PyUnsafeUnionType" to { members -> PyUnsafeUnionType.unsafeUnion(members)!! }, + "PyIntersectionType" to { members -> PyIntersectionType.intersection(members)!! }, + ) + + /** Leaf type with stable identity; counts every `hashCode`/`equals` it receives. */ + private inner class CountingLeaf(private val tag: String) : PyType { + override fun hashCode(): Int { counter.incrementAndGet(); return tag.hashCode() } + override fun equals(other: Any?): Boolean { counter.incrementAndGet(); return other is CountingLeaf && other.tag == tag } + override fun resolveMember(name: String, location: PyExpression?, direction: AccessDirection, resolveContext: PyResolveContext): List? = null + override fun getCompletionVariants(completionPrefix: String?, location: PsiElement, context: ProcessingContext): Array = emptyArray() + override val name: String get() = tag + override val isBuiltin: Boolean get() = false + override fun assertValid(message: String?) {} + } + + /** Non-memoizing wrapper delegating into [child]; distinct [tag]s over a shared child give the DAG its branching. */ + private inner class Wrapper(private val tag: Int, private val child: PyType) : PyType { + override fun hashCode(): Int { counter.incrementAndGet(); return 31 * tag + child.hashCode() } + override fun equals(other: Any?): Boolean { counter.incrementAndGet(); return other is Wrapper && other.tag == tag && other.child == child } + override fun resolveMember(name: String, location: PyExpression?, direction: AccessDirection, resolveContext: PyResolveContext): List? = null + override fun getCompletionVariants(completionPrefix: String?, location: PsiElement, context: ProcessingContext): Array = emptyArray() + override val name: String get() = "W$tag" + override val isBuiltin: Boolean get() = false + override fun assertValid(message: String?) {} + } + + /** Leaf with a fixed hash but identity-by-[tag] equals, to force a hash collision between distinct members. */ + private class CollidingLeaf(private val tag: String) : PyType { + override fun hashCode(): Int = 0 + override fun equals(other: Any?): Boolean = other is CollidingLeaf && other.tag == tag + override fun resolveMember(name: String, location: PyExpression?, direction: AccessDirection, resolveContext: PyResolveContext): List? = null + override fun getCompletionVariants(completionPrefix: String?, location: PsiElement, context: ProcessingContext): Array = emptyArray() + override val name: String get() = tag + override val isBuiltin: Boolean get() = false + override fun assertValid(message: String?) {} + } + + private fun buildSharedDag(depth: Int, leafA: PyType, leafB: PyType, make: (List) -> PyType): PyType { + var node: PyType = make(listOf(leafA, leafB)) + repeat(depth) { node = make(listOf(Wrapper(0, node), Wrapper(1, node))) } + return node + } + + @Test + fun `building a shared-subtree composite DAG does not blow up hashCode`() { + assertTimeoutPreemptively(Duration.ofSeconds(60)) { + for ((name, make) in composites) { + counter.set(0) + buildSharedDag(DEPTH, CountingLeaf("A"), CountingLeaf("B"), make).hashCode() + assertTrue(counter.get() < LINEAR_BOUND) { "$name: hashCode blew up: ${counter.get()} (want < $LINEAR_BOUND at depth $DEPTH)" } + } + } + } + + @Test + fun `equals of two distinct nested composites short-circuits on hashCode`() { + assertTimeoutPreemptively(Duration.ofSeconds(60)) { + for ((name, make) in composites) { + val a = buildSharedDag(DEPTH, CountingLeaf("A"), CountingLeaf("B"), make) + val b = buildSharedDag(DEPTH, CountingLeaf("A"), CountingLeaf("X"), make) + counter.set(0) + assertFalse(a == b) { "$name: DAGs differing at the deepest leaf must not be equal" } + assertTrue(counter.get() < LINEAR_BOUND) { "$name: unequal-equals walked members: ${counter.get()} (want < $LINEAR_BOUND)" } + } + } + } + + @Test + fun `equal nested composites stay equal and hash-consistent`() { + for ((name, make) in composites) { + val a = buildSharedDag(DEPTH, CountingLeaf("A"), CountingLeaf("B"), make) + val b = buildSharedDag(DEPTH, CountingLeaf("A"), CountingLeaf("B"), make) + assertEquals(a, b) { "$name: structurally identical DAGs must be equal" } + assertEquals(a.hashCode(), b.hashCode()) { "$name: equal composites must have equal hash codes" } + } + } + + @Test + fun `composite equality is order-independent`() { + for ((name, make) in composites) { + val a = CountingLeaf("A") + val b = CountingLeaf("B") + assertEquals(make(listOf(a, b)), make(listOf(b, a))) { "$name: member order must not affect equality" } + } + } + + @Test + fun `hash-colliding members do not produce false equality`() { + for ((name, make) in composites) { + val common = CountingLeaf("C") + val x = make(listOf(CollidingLeaf("P"), common)) + val y = make(listOf(CollidingLeaf("Q"), common)) + assertEquals(x.hashCode(), y.hashCode()) { "$name: colliding leaves must give equal composite hash codes (precondition)" } + assertNotEquals(x, y) { "$name: the hashCode fast path must fall through to the member-set comparison" } + } + } + + @Test + fun `composites of different kinds are never equal`() { + val a = CountingLeaf("A") + val b = CountingLeaf("B") + val union = PyUnionType.union(listOf(a, b))!! + val intersection = PyIntersectionType.intersection(listOf(a, b))!! + // Same member set -> identical hashCode, so this also checks the exact-class gate precedes the fast path. + assertEquals(union.hashCode(), intersection.hashCode()) + assertNotEquals(union, intersection) + assertNotEquals(intersection, union) + } + + private companion object { + const val DEPTH = 22 + const val LINEAR_BOUND = 100_000L + } +}