Explorar o código

script/research/event_graph: major bugs fix in the algo

ghassmo %!s(int64=3) %!d(string=hai) anos
pai
achega
68d4ac7691
Modificáronse 1 ficheiros con 98 adicións e 40 borrados
  1. 98 40
      script/research/event_graph/main.py

+ 98 - 40
script/research/event_graph/main.py

@@ -14,7 +14,7 @@ EventIds = list[EventId]
 class Event:
 class Event:
     def __init__(self, parents: EventIds):
     def __init__(self, parents: EventIds):
         self.timestamp = datetime.now().timestamp
         self.timestamp = datetime.now().timestamp
-        self.parents = parents
+        self.parents = sorted(parents)
 
 
     def set_timestamp(self, timestamp):
     def set_timestamp(self, timestamp):
         self.timestamp = timestamp
         self.timestamp = timestamp
@@ -48,17 +48,12 @@ class Event:
  E7: [E4]
  E7: [E4]
  E8: [E2]
  E8: [E2]
 """
 """
+
+
 class Graph:
 class Graph:
     def __init__(self):
     def __init__(self):
         self.events = dict()
         self.events = dict()
 
 
-    def heads(self):
-        # NOTE: we will need to keep track of heads for creating new events.
-        #       Not needed for this demo though.
-
-        # XXX this for testing purpose
-        return [list(self.events.keys())[-1]]
-
     def add_event(self, event: Event):
     def add_event(self, event: Event):
         self.events[event.hash()] = event
         self.events[event.hash()] = event
 
 
@@ -72,7 +67,7 @@ class Graph:
         missing_events = []
         missing_events = []
 
 
         for e in events:
         for e in events:
-            if e not in self.events:
+            if self.events.get(e) == None:
                 missing_events.append(e)
                 missing_events.append(e)
 
 
         return missing_events
         return missing_events
@@ -94,24 +89,44 @@ class Node:
         genesis_event = Event([])
         genesis_event = Event([])
         genesis_event.set_timestamp(0.0)
         genesis_event.set_timestamp(0.0)
 
 
+        # make the root node as head
+        self.heads = [genesis_event.hash()]
+
         self.genesis_event = genesis_event
         self.genesis_event = genesis_event
         self.active_pool.add_event(genesis_event)
         self.active_pool.add_event(genesis_event)
 
 
-    def last_event(self):
-        return self.active_pool.heads()
+    def remove_heads(self, event):
+        for p in event.parents:
+            if p in self.heads:
+                self.heads.remove(p)
+
+    def update_heads(self, event):
+        event_hash = event.hash()
+        self.remove_heads(event)
+        self.heads.append(event_hash)
 
 
     def receive_new_event(self, event: Event):
     def receive_new_event(self, event: Event):
+        event_hash = event.hash()
 
 
         # reject events with no parents
         # reject events with no parents
-        if len(event.parents) == 0:
+        if not event.parents:
+            return
+
+        # reject event already exist in active pool
+        if not self.active_pool.check([event_hash]):
+            return
+
+        # reject event already exist in orphan pool
+        if not self.orphan_pool.check([event_hash]):
             return
             return
 
 
         missing_parents = self.active_pool.check(event.parents)
         missing_parents = self.active_pool.check(event.parents)
 
 
-        if len(missing_parents) == 0:
+        if not missing_parents:
             # if there are no missing parents
             # if there are no missing parents
             # add the event to active pool
             # add the event to active pool
             self.active_pool.add_event(event)
             self.active_pool.add_event(event)
+            self.update_heads(event)
 
 
             # events list to be removed from orphan pool
             # events list to be removed from orphan pool
             # after relink
             # after relink
@@ -132,34 +147,44 @@ class Node:
             request_list = []
             request_list = []
             self.check_parents(request_list, missing_parents)
             self.check_parents(request_list, missing_parents)
 
 
+            print(f"{self.name} request from the network: {request_list}")
+
             # XXX
             # XXX
             # send all the missing parents in request_list
             # send all the missing parents in request_list
             # to the node who send this event
             # to the node who send this event
 
 
-    def check_parents(self, request_list, parents: EventIds):
+    def check_parents(self, request_list, parents: EventIds, visited=[]):
         for parent_hash in parents:
         for parent_hash in parents:
+            if parent_hash in visited:
+                continue
+
+            visited.append(parent_hash)
+
             if parent_hash in self.orphan_pool.events:
             if parent_hash in self.orphan_pool.events:
                 parent = self.orphan_pool.events[parent_hash]
                 parent = self.orphan_pool.events[parent_hash]
 
 
                 # recursive call
                 # recursive call
-                self.check_parents(request_list, parent.parents)
+                self.check_parents(request_list, parent.parents, visited)
             else:
             else:
                 request_list.append(parent_hash)
                 request_list.append(parent_hash)
 
 
     def relink(self, event: Event, remove_list=[]):
     def relink(self, event: Event, remove_list=[]):
+        event_hash = event.hash()
+
         # check if the orphan pool has an event linked
         # check if the orphan pool has an event linked
         # to the new added event
         # to the new added event
         for (orphan_hash, orphan) in self.orphan_pool.events.items():
         for (orphan_hash, orphan) in self.orphan_pool.events.items():
             if orphan_hash in remove_list:
             if orphan_hash in remove_list:
                 continue
                 continue
 
 
-            if event.hash() not in orphan.parents:
+            if event_hash not in orphan.parents:
                 continue
                 continue
 
 
             missing_parents = self.active_pool.check(orphan.parents)
             missing_parents = self.active_pool.check(orphan.parents)
 
 
-            if len(missing_parents) == 0:
+            if not missing_parents:
                 self.active_pool.add_event(orphan)
                 self.active_pool.add_event(orphan)
+                self.update_heads(orphan)
                 remove_list.append(orphan_hash)
                 remove_list.append(orphan_hash)
 
 
                 # recursive call
                 # recursive call
@@ -167,41 +192,40 @@ class Node:
 
 
     def __str__(self):
     def __str__(self):
         return f"""------
         return f"""------
-            \n Name: {self.name}
-            \n Active Pool: {self.active_pool}
-            \n Orphan Pool: {self.orphan_pool}"""
-
+	        \n Name: {self.name}
+	        \n Active Pool: {self.active_pool}
+	        \n Orphan Pool: {self.orphan_pool}"""
 
 
-NODES_N = 10
-BROADCAST_ATTEMPT = 3
 
 
 MAX_BROADCAST_DELAY = 2
 MAX_BROADCAST_DELAY = 2
 MIN_BROADCAST_DELAY = 0
 MIN_BROADCAST_DELAY = 0
 
 
+NODES_N = 15
+BROADCAST_ATTEMPT = 3
+BROADCAST_TIMEOUT = NODES_N * BROADCAST_ATTEMPT * MAX_BROADCAST_DELAY
+
 
 
 async def recv_loop(node, peer, queue):
 async def recv_loop(node, peer, queue):
     while True:
     while True:
         event = await queue.get()
         event = await queue.get()
-        queue.task_done()
         node.receive_new_event(event)
         node.receive_new_event(event)
-        print(f"{node.name} receive: {event.hash()} from {peer}")
+        queue.task_done()
+        print(f"{node.name} receive event from {peer}: \n {event}")
 
 
 
 
 async def send_loop(node, queue):
 async def send_loop(node, queue):
     for _ in range(BROADCAST_ATTEMPT):
     for _ in range(BROADCAST_ATTEMPT):
 
 
         await asyncio.sleep(randint(MIN_BROADCAST_DELAY, MAX_BROADCAST_DELAY))
         await asyncio.sleep(randint(MIN_BROADCAST_DELAY, MAX_BROADCAST_DELAY))
-        event = Event(node.last_event())
+        event = Event(node.heads)
 
 
+        print(f"{node.name} broadcast event: \n {event}")
         for _ in range(NODES_N):
         for _ in range(NODES_N):
             await queue.put(event)
             await queue.put(event)
             await queue.join()
             await queue.join()
 
 
-        print(f"{node.name} broadcast: {event.hash()}")
-
 
 
 async def main():
 async def main():
-
     send_tasks = []
     send_tasks = []
     recv_tasks = []
     recv_tasks = []
 
 
@@ -225,29 +249,63 @@ async def main():
             recv_task = asyncio.create_task(recv_loop(node, peer, queue))
             recv_task = asyncio.create_task(recv_loop(node, peer, queue))
             recv_tasks.append(recv_task)
             recv_tasks.append(recv_task)
 
 
-    g = asyncio.gather(*recv_tasks)
+    try:
+        # run recv tasks
+        r_g = asyncio.gather(*send_tasks)
+
+        # run and wait for send tasks
+        s_g = asyncio.gather(*send_tasks)
+        await asyncio.wait_for(s_g, BROADCAST_TIMEOUT)
+
+        # cancel recv tasks
+        r_g.cancel()
+
+        # assert if all nodes share the same active pool graph
+        assert (all(n.active_pool.events.keys() ==
+                nodes[0].active_pool.events.keys() for n in nodes))
+
+        # assert if all nodes share the same orphan pool graph
+        assert (all(n.orphan_pool.events.keys() ==
+                nodes[0].orphan_pool.events.keys() for n in nodes))
 
 
-    await asyncio.gather(*send_tasks)
+        print_graph([nodes[0]])
 
 
+    except asyncio.exceptions.TimeoutError:
+        print("Broadcast TimeoutError")
+
+
+def print_nodes(nodes):
     for node in nodes:
     for node in nodes:
         print(node.name)
         print(node.name)
         print(len(node.active_pool.events))
         print(len(node.active_pool.events))
         print(len(node.orphan_pool.events))
         print(len(node.orphan_pool.events))
         print("###############")
         print("###############")
 
 
-    graph = nx.Graph()
 
 
-    node_to_draw = nodes[0]
-    print(node_to_draw)
+def print_graph(nodes):
+    for (i, node) in enumerate(nodes):
+        graph = nx.Graph()
 
 
-    for (h, ev) in node_to_draw.active_pool.events.items():
-        graph.add_node(h[:5])
-        graph.add_edges_from([(h[:5], p[:5]) for p in ev.parents])
+        for (h, ev) in node.active_pool.events.items():
+            graph.add_node(h[:5])
+            graph.add_edges_from([(h[:5], p[:5]) for p in ev.parents])
 
 
-    nx.draw(graph, with_labels=True, node_color="#69aaff", node_size=400)
-    plt.show()
+        colors = []
+
+        node_heads = [h[:5] for h in node.heads]
 
 
-    await g
+        for n in graph.nodes():
+            if n == "8aed6":
+                colors.append("red")
+            elif n in node_heads:
+                colors.append("yellow")
+            else:
+                colors.append("blue")
+
+        plt.figure(i)
+        nx.draw_networkx(graph, with_labels=True, node_color=colors)
+
+    plt.show()
 
 
 
 
 def test_node():
 def test_node():