From 9c78edc764d93821136310edaf0b1bb8be564470 Mon Sep 17 00:00:00 2001 From: Ian Robinson Date: Wed, 29 Apr 2026 02:21:09 -0400 Subject: [PATCH] Rule Builder: Fix rule hash collisions (#6169) * fix rule hash collisions * just let the dict do the work --- rule_builder/rules.py | 9 ++++----- test/general/test_rule_builder.py | 9 +++++++++ 2 files changed, 13 insertions(+), 5 deletions(-) diff --git a/rule_builder/rules.py b/rule_builder/rules.py index 07c0607c1f..47f91aff5e 100644 --- a/rule_builder/rules.py +++ b/rule_builder/rules.py @@ -36,7 +36,7 @@ def _create_hash_fn(resolved_rule_cls: "CustomRuleRegister") -> Callable[..., in class CustomRuleRegister(type): """A metaclass to contain world custom rules and automatically convert resolved rules to frozen dataclasses""" - resolved_rules: ClassVar[dict[int, "Rule.Resolved"]] = {} + resolved_rules: ClassVar[dict["Rule.Resolved", "Rule.Resolved"]] = {} """A cached of resolved rules to turn each unique one into a singleton""" custom_rules: ClassVar[dict[str, dict[str, type["Rule[Any]"]]]] = {} @@ -64,10 +64,9 @@ class CustomRuleRegister(type): @override def __call__(cls, *args: Any, **kwds: Any) -> Any: rule = super().__call__(*args, **kwds) - rule_hash = hash(rule) - if rule_hash in cls.resolved_rules: - return cls.resolved_rules[rule_hash] - cls.resolved_rules[rule_hash] = rule + if rule in cls.resolved_rules: + return cls.resolved_rules[rule] + cls.resolved_rules[rule] = rule return rule @classmethod diff --git a/test/general/test_rule_builder.py b/test/general/test_rule_builder.py index 85e239175d..682c043f8e 100644 --- a/test/general/test_rule_builder.py +++ b/test/general/test_rule_builder.py @@ -416,6 +416,15 @@ class TestHashes(RuleBuilderTestCase): rule2 = HasAll("2", "2", "2", "1") self.assertEqual(hash(rule1.resolve(world)), hash(rule2.resolve(world))) + def test_hash_collision(self) -> None: + multiworld = setup_solo_multiworld(self.world_cls, steps=("generate_early",), seed=0) + world = multiworld.worlds[1] + rule1 = Has("A", count=1).resolve(world) + rule2 = Has("A", count=1 << 61).resolve(world) + self.assertEqual(hash(rule1), hash(rule2)) + self.assertNotEqual(rule1, rule2) + self.assertNotEqual(id(rule1), id(rule2)) + class TestCaching(CachedRuleBuilderTestCase): multiworld: MultiWorld # pyright: ignore[reportUninitializedInstanceVariable]