Bugs
This commit is contained in:
parent
4d258bb395
commit
ef0cedce3d
1 changed files with 30 additions and 3 deletions
|
|
@ -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
|
||||||
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------------------
|
# ------------------------------------------------------------------------------
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue