summary refs log tree commit diff
path: root/synapse/storage/roommember.py
diff options
context:
space:
mode:
authorErik Johnston <erik@matrix.org>2016-08-26 10:59:40 +0100
committerErik Johnston <erik@matrix.org>2016-08-26 10:59:40 +0100
commit1ccdc1e93a5ae854fa89751a78c9103940a9f9e6 (patch)
tree4eaf547dd315e5d8fc70585cf2d1951ffd0ccaf2 /synapse/storage/roommember.py
parentAdd measure on check_host_in_room (diff)
downloadsynapse-1ccdc1e93a5ae854fa89751a78c9103940a9f9e6.tar.xz
Cache check_host_in_room
Diffstat (limited to 'synapse/storage/roommember.py')
-rw-r--r--synapse/storage/roommember.py35
1 files changed, 35 insertions, 0 deletions
diff --git a/synapse/storage/roommember.py b/synapse/storage/roommember.py
index 2cab065bca..5ce5e8da37 100644
--- a/synapse/storage/roommember.py
+++ b/synapse/storage/roommember.py
@@ -400,3 +400,38 @@ class RoomMemberStore(SQLBaseStore):
         )
 
         defer.returnValue(set(row["user_id"] for row in rows))
+
+    def is_host_joined(self, room_id, host, state_group, state_ids):
+        if not state_group:
+            # If state_group is None it means it has yet to be assigned a
+            # state group, i.e. we need to make sure that calls with a state_group
+            # of None don't hit previous cached calls with a None state_group.
+            # To do this we set the state_group to a new object as object() != object()
+            state_group = object()
+
+        return self._get_joined_users_from_context(
+            room_id, state_group, state_ids
+        )
+
+    @cachedInlineCallbacks(num_args=3)
+    def _is_host_joined(self, room_id, host, state_group, current_state_ids):
+        # We don't use `state_group`, its there so that we can cache based
+        # on it. However, its important that its never None, since two current_state's
+        # with a state_group of None are likely to be different.
+        # See bulk_get_push_rules_for_room for how we work around this.
+        assert state_group is not None
+
+        for (etype, state_key), event_id in current_state_ids.items():
+            if etype == EventTypes.Member:
+                try:
+                    if get_domain_from_id(state_key) != host:
+                        continue
+                except:
+                    logger.warn("state_key not user_id: %s", state_key)
+                    continue
+
+                event = yield self.store.get_event(event_id, allow_none=True)
+                if event and event.content["membership"] == Membership.JOIN:
+                    defer.returnValue(True)
+
+        defer.returnValue(False)