This commit is contained in:
Astra Logical 2026-07-23 12:10:13 -04:00
parent d382b084d2
commit 4d258bb395

View file

@ -61,15 +61,22 @@ class CustomMemoryStateStore(MemoryStateStore):
class CustomMemoryCryptoStore(MemoryCryptoStore):
async def put_cross_signing_key(self, user_id: UserID, usage: str, key) -> None:
"""Fixes the mautrix MemoryCryptoStore AttributeError when updating cross-signing keys."""
try:
await super().put_cross_signing_key(user_id, usage, key)
except AttributeError:
for attr_val in self.__dict__.values():
if isinstance(attr_val, dict):
attr_val[(user_id, usage)] = key
attr_val[user_id] = key
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
if not hasattr(self, "_cross_signing_keys"):
self._cross_signing_keys = {}
async def put_cross_signing_key(self, user_id: UserID, usage, key) -> None:
usage_str = usage.value if hasattr(usage, "value") else str(usage)
if user_id not in self._cross_signing_keys:
self._cross_signing_keys[user_id] = {}
self._cross_signing_keys[user_id][usage_str] = key
self._cross_signing_keys[user_id][usage] = key
async def get_cross_signing_key(self, user_id: UserID, usage):
usage_str = usage.value if hasattr(usage, "value") else str(usage)
user_dict = self._cross_signing_keys.get(user_id, {})
return user_dict.get(usage_str) or user_dict.get(usage)
# ------------------------------------------------------------------------------