PY-90901 memoize composite type hashCode and short-circuit equals

Deeply nested composite types (unions, unsafe unions, intersections)
whose structural equals/hashCode were recomputed on every call caused a
quadratic/exponential storm in the type-eval HashMap caches: a Find
Usages run on a large project (netflix/dispatch) wedged the FJP workers
in recursive PyUnionType/PyClassTypeImpl equality and starved the EDT of
the write-intent lock, freezing the UI.

Introduce PyCompositeTypeBase as the single home for the three
set-of-members composites. It memoizes hashCode (the member set is
immutable) and adds a hashCode-mismatch fast path to equals, so unequal
composites are rejected without the O(n) member-set walk. PyClassTypeImpl
gets the same fast path, ordered after the cheap myClass/isDefinition
checks so a trivial mismatch never forces the deep type-argument hash.

PyCompositeTypeEqualityPerformanceTest guards this with a shared-subtree
DAG whose root is reachable via 2^depth paths: a deterministic call
counter stays linear with the fix and blows past the bound without it.

(cherry picked from commit 8fdc2ee131c12842750f8242cee6177522cbf5a5)

GitOrigin-RevId: 1c3d041a4922e63592fa15550a0ede9188fa38d5
This commit is contained in:
Daniil Kalinin
2026-08-07 15:21:46 +00:00
committed by intellij-monorepo-bot
parent 913d7f0ee2
commit 69b528d076
6 changed files with 211 additions and 47 deletions
@@ -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;
@@ -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<PyType?>
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
}
}
@@ -12,8 +12,11 @@ import org.jetbrains.annotations.ApiStatus
import java.util.Collections
@ApiStatus.Experimental
class PyIntersectionType private constructor(members: Collection<PyType?>) : PyCompositeType {
override val members: Set<PyType?> = Collections.unmodifiableSet<PyType?>(LinkedHashSet(members))
class PyIntersectionType private constructor(members: Collection<PyType?>) : PyCompositeTypeBase() {
override val memberSet: Set<PyType?> = Collections.unmodifiableSet(LinkedHashSet(members))
override val members: Set<PyType?>
get() = memberSet
override fun resolveMember(
name: String,
@@ -58,19 +61,6 @@ class PyIntersectionType private constructor(members: Collection<PyType?>) : 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"
}
@@ -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<? extends RatedResolveResult> 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<Object> 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();
@@ -51,9 +51,11 @@ import java.util.Collections
* @see PyUnionType.createWeakType
*/
@ApiStatus.Experimental
class PyUnsafeUnionType private constructor(members: Collection<PyType?>) : PyCompositeType {
override val members: Set<PyType?> = LinkedHashSet(members)
get() = Collections.unmodifiableSet<PyType?>(field)
class PyUnsafeUnionType private constructor(members: Collection<PyType?>) : PyCompositeTypeBase() {
override val memberSet: Set<PyType?> = Collections.unmodifiableSet(LinkedHashSet(members))
override val members: Set<PyType?>
get() = memberSet
override fun resolveMember(
name: String,
@@ -98,19 +100,6 @@ class PyUnsafeUnionType private constructor(members: Collection<PyType?>) : 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"
}
@@ -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<Pair<String, (List<PyType?>) -> 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<RatedResolveResult>? = null
override fun getCompletionVariants(completionPrefix: String?, location: PsiElement, context: ProcessingContext): Array<out Any> = 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<RatedResolveResult>? = null
override fun getCompletionVariants(completionPrefix: String?, location: PsiElement, context: ProcessingContext): Array<out Any> = 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<RatedResolveResult>? = null
override fun getCompletionVariants(completionPrefix: String?, location: PsiElement, context: ProcessingContext): Array<out Any> = 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): 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
}
}