diff --git a/src/joserfc/_keys.py b/src/joserfc/_keys.py index e33ab7e..10427f3 100644 --- a/src/joserfc/_keys.py +++ b/src/joserfc/_keys.py @@ -124,7 +124,8 @@ def __bool__(self) -> bool: return bool(self.keys) def __eq__(self, other: t.Any) -> bool: - assert isinstance(other, KeySet) + if not isinstance(other, KeySet): + return NotImplemented return self.keys == other.keys def as_dict(self, private: bool = False, **params: t.Any) -> KeySetSerialization: diff --git a/tests/jwk/test_jwk_set.py b/tests/jwk/test_jwk_set.py index 06c4bbc..c4b90a7 100644 --- a/tests/jwk/test_jwk_set.py +++ b/tests/jwk/test_jwk_set.py @@ -103,3 +103,16 @@ def test_key_eq_with_new_keys(self): key_set2 = KeySet([RSAKey.import_key(k.as_dict(private=True)) for k in key_set1]) self.assertIsNot(key_set1, key_set2) self.assertEqual(key_set1, key_set2) + + def test_key_set_eq_with_unrelated_type(self): + key_set = KeySet.generate_key_set("oct", 8, count=1) + self.assertFalse(key_set == "foo") + self.assertNotEqual(key_set, "foo") + + def test_key_set_eq_uses_reflected_comparison(self): + class EqualToKeySet: + def __eq__(self, other): + return isinstance(other, KeySet) + + key_set = KeySet.generate_key_set("oct", 8, count=1) + self.assertTrue(key_set == EqualToKeySet())