This commit is contained in:
Astra Logical 2026-07-23 12:12:34 -04:00
parent 4d258bb395
commit ef0cedce3d

View file

@ -60,6 +60,27 @@ class CustomMemoryStateStore(MemoryStateStore):
return shared return shared
class CrossSigningKeyWrapper:
def __init__(self, key_val):
if hasattr(key_val, "key"):
self.key = key_val.key
elif isinstance(key_val, str):
self.key = key_val
else:
self.key = str(key_val)
def __str__(self):
return self.key
def __repr__(self):
return f"CrossSigningKeyWrapper(key={self.key!r})"
def __eq__(self, other):
if hasattr(other, "key"):
return self.key == other.key
return self.key == str(other)
class CustomMemoryCryptoStore(MemoryCryptoStore): class CustomMemoryCryptoStore(MemoryCryptoStore):
def __init__(self, *args, **kwargs): def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs) super().__init__(*args, **kwargs)
@ -70,13 +91,19 @@ class CustomMemoryCryptoStore(MemoryCryptoStore):
usage_str = usage.value if hasattr(usage, "value") else str(usage) usage_str = usage.value if hasattr(usage, "value") else str(usage)
if user_id not in self._cross_signing_keys: if user_id not in self._cross_signing_keys:
self._cross_signing_keys[user_id] = {} self._cross_signing_keys[user_id] = {}
self._cross_signing_keys[user_id][usage_str] = key wrapped = CrossSigningKeyWrapper(key)
self._cross_signing_keys[user_id][usage] = key self._cross_signing_keys[user_id][usage_str] = wrapped
self._cross_signing_keys[user_id][usage] = wrapped
async def get_cross_signing_key(self, user_id: UserID, usage): async def get_cross_signing_key(self, user_id: UserID, usage):
usage_str = usage.value if hasattr(usage, "value") else str(usage) usage_str = usage.value if hasattr(usage, "value") else str(usage)
user_dict = self._cross_signing_keys.get(user_id, {}) user_dict = self._cross_signing_keys.get(user_id, {})
return user_dict.get(usage_str) or user_dict.get(usage) val = user_dict.get(usage_str) or user_dict.get(usage)
if val is not None and not isinstance(val, CrossSigningKeyWrapper):
val = CrossSigningKeyWrapper(val)
user_dict[usage_str] = val
user_dict[usage] = val
return val
# ------------------------------------------------------------------------------ # ------------------------------------------------------------------------------