main.py 8.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335
  1. from hashlib import sha256
  2. from datetime import datetime
  3. from random import randint
  4. import asyncio
  5. import matplotlib.pyplot as plt
  6. import networkx as nx
  7. EventId = str
  8. EventIds = list[EventId]
  9. class Event:
  10. def __init__(self, parents: EventIds):
  11. self.timestamp = datetime.now().timestamp
  12. self.parents = sorted(parents)
  13. def set_timestamp(self, timestamp):
  14. self.timestamp = timestamp
  15. def hash(self) -> str:
  16. m = sha256()
  17. m.update(str.encode(str(self.timestamp)))
  18. for p in self.parents:
  19. m.update(str.encode(str(p)))
  20. return m.digest().hex()
  21. def __str__(self):
  22. res = f"{self.hash()}"
  23. for p in self.parents:
  24. res += f"\n |"
  25. res += f"\n - {p}"
  26. res += f"\n"
  27. return res
  28. """
  29. ## Graph Example
  30. E1: []
  31. E2: [E1]
  32. E3: [E1]
  33. E4: [E3]
  34. E5: [E3]
  35. E6: [E4, E5]
  36. E7: [E4]
  37. E8: [E2]
  38. """
  39. class Graph:
  40. def __init__(self):
  41. self.events = dict()
  42. def add_event(self, event: Event):
  43. self.events[event.hash()] = event
  44. def remove_event(self, event_id: EventId):
  45. if event_id in self.events:
  46. del self.events[event_id]
  47. # check if given events are exist in the graph
  48. # return a list of missing events
  49. def check(self, events: EventIds) -> EventIds:
  50. missing_events = []
  51. for e in events:
  52. if self.events.get(e) == None:
  53. missing_events.append(e)
  54. return missing_events
  55. def __str__(self):
  56. res = ""
  57. for event in self.events.values():
  58. res += f"\n {event}"
  59. return res
  60. class Node:
  61. def __init__(self, name: str):
  62. self.name = name
  63. self.orphan_pool = Graph()
  64. self.active_pool = Graph()
  65. # the active pool should always start with one event
  66. genesis_event = Event([])
  67. genesis_event.set_timestamp(0.0)
  68. # make the root node as head
  69. self.heads = [genesis_event.hash()]
  70. self.genesis_event = genesis_event
  71. self.active_pool.add_event(genesis_event)
  72. def remove_heads(self, event):
  73. for p in event.parents:
  74. if p in self.heads:
  75. self.heads.remove(p)
  76. def update_heads(self, event):
  77. event_hash = event.hash()
  78. self.remove_heads(event)
  79. self.heads.append(event_hash)
  80. def receive_new_event(self, event: Event):
  81. event_hash = event.hash()
  82. # reject events with no parents
  83. if not event.parents:
  84. return
  85. # reject event already exist in active pool
  86. if not self.active_pool.check([event_hash]):
  87. return
  88. # reject event already exist in orphan pool
  89. if not self.orphan_pool.check([event_hash]):
  90. return
  91. missing_parents = self.active_pool.check(event.parents)
  92. if not missing_parents:
  93. # if there are no missing parents
  94. # add the event to active pool
  95. self.active_pool.add_event(event)
  96. self.update_heads(event)
  97. # events list to be removed from orphan pool
  98. # after relink
  99. remove_list: EventIds = []
  100. self.relink(event, remove_list)
  101. # clean up orphan pool
  102. for ev in remove_list:
  103. self.orphan_pool.remove_event(ev)
  104. else:
  105. # add the received event to the orphan pool
  106. self.orphan_pool.add_event(event)
  107. # check if all the missing parents are in orphan pool
  108. # if the missing parents and their links not in orphan pool, request
  109. # them from the network
  110. request_list = []
  111. self.check_parents(request_list, missing_parents)
  112. print(f"{self.name} request from the network: {request_list}")
  113. # XXX
  114. # send all the missing parents in request_list
  115. # to the node who send this event
  116. def check_parents(self, request_list, parents: EventIds, visited=[]):
  117. for parent_hash in parents:
  118. if parent_hash in visited:
  119. continue
  120. visited.append(parent_hash)
  121. if parent_hash in self.orphan_pool.events:
  122. parent = self.orphan_pool.events[parent_hash]
  123. # recursive call
  124. self.check_parents(request_list, parent.parents, visited)
  125. else:
  126. request_list.append(parent_hash)
  127. def relink(self, event: Event, remove_list=[]):
  128. event_hash = event.hash()
  129. # check if the orphan pool has an event linked
  130. # to the new added event
  131. for (orphan_hash, orphan) in self.orphan_pool.events.items():
  132. if orphan_hash in remove_list:
  133. continue
  134. if event_hash not in orphan.parents:
  135. continue
  136. missing_parents = self.active_pool.check(orphan.parents)
  137. if not missing_parents:
  138. self.active_pool.add_event(orphan)
  139. self.update_heads(orphan)
  140. remove_list.append(orphan_hash)
  141. # recursive call
  142. self.relink(orphan, remove_list)
  143. def __str__(self):
  144. return f"""------
  145. \n Name: {self.name}
  146. \n Active Pool: {self.active_pool}
  147. \n Orphan Pool: {self.orphan_pool}"""
  148. MAX_BROADCAST_DELAY = 2
  149. MIN_BROADCAST_DELAY = 0
  150. NODES_N = 15
  151. BROADCAST_ATTEMPT = 3
  152. BROADCAST_TIMEOUT = NODES_N * BROADCAST_ATTEMPT * MAX_BROADCAST_DELAY
  153. async def recv_loop(node, peer, queue):
  154. while True:
  155. event = await queue.get()
  156. node.receive_new_event(event)
  157. queue.task_done()
  158. print(f"{node.name} receive event from {peer}: \n {event}")
  159. async def send_loop(node, queue):
  160. for _ in range(BROADCAST_ATTEMPT):
  161. await asyncio.sleep(randint(MIN_BROADCAST_DELAY, MAX_BROADCAST_DELAY))
  162. event = Event(node.heads)
  163. print(f"{node.name} broadcast event: \n {event}")
  164. for _ in range(NODES_N):
  165. await queue.put(event)
  166. await queue.join()
  167. async def main():
  168. send_tasks = []
  169. recv_tasks = []
  170. nodes = []
  171. queues = dict()
  172. print(f"Run {NODES_N} Nodes")
  173. for i in range(NODES_N):
  174. node = Node(f"Node{i}")
  175. nodes.append(node)
  176. queue = asyncio.Queue()
  177. queues[node.name] = queue
  178. send_task = asyncio.create_task(send_loop(node, queue))
  179. send_tasks.append(send_task)
  180. for node in nodes:
  181. for (peer, queue) in queues.items():
  182. recv_task = asyncio.create_task(recv_loop(node, peer, queue))
  183. recv_tasks.append(recv_task)
  184. try:
  185. # run recv tasks
  186. r_g = asyncio.gather(*send_tasks)
  187. # run and wait for send tasks
  188. s_g = asyncio.gather(*send_tasks)
  189. await asyncio.wait_for(s_g, BROADCAST_TIMEOUT)
  190. # cancel recv tasks
  191. r_g.cancel()
  192. # assert if all nodes share the same active pool graph
  193. assert (all(n.active_pool.events.keys() ==
  194. nodes[0].active_pool.events.keys() for n in nodes))
  195. # assert if all nodes share the same orphan pool graph
  196. assert (all(n.orphan_pool.events.keys() ==
  197. nodes[0].orphan_pool.events.keys() for n in nodes))
  198. print_graph([nodes[0]])
  199. except asyncio.exceptions.TimeoutError:
  200. print("Broadcast TimeoutError")
  201. def print_nodes(nodes):
  202. for node in nodes:
  203. print(node.name)
  204. print(len(node.active_pool.events))
  205. print(len(node.orphan_pool.events))
  206. print("###############")
  207. def print_graph(nodes):
  208. for (i, node) in enumerate(nodes):
  209. graph = nx.Graph()
  210. for (h, ev) in node.active_pool.events.items():
  211. graph.add_node(h[:5])
  212. graph.add_edges_from([(h[:5], p[:5]) for p in ev.parents])
  213. colors = []
  214. node_heads = [h[:5] for h in node.heads]
  215. for n in graph.nodes():
  216. if n == "8aed6":
  217. colors.append("red")
  218. elif n in node_heads:
  219. colors.append("yellow")
  220. else:
  221. colors.append("blue")
  222. plt.figure(i)
  223. nx.draw_networkx(graph, with_labels=True, node_color=colors)
  224. plt.show()
  225. def test_node():
  226. node_a = Node("NodeA")
  227. event0 = node_a.genesis_event
  228. event1 = Event([event0.hash()])
  229. event2 = Event([event1.hash()])
  230. event3 = Event([event2.hash(), event0.hash()])
  231. event4 = Event([event1.hash(), event3.hash()])
  232. event5 = Event([event4.hash(), "FAKEHASH"])
  233. event6 = Event([event5.hash(), event3.hash()])
  234. node_a.receive_new_event(event3)
  235. node_a.receive_new_event(event2)
  236. node_a.receive_new_event(event1)
  237. node_a.receive_new_event(event5)
  238. node_a.receive_new_event(event6)
  239. node_a.receive_new_event(event4)
  240. print(node_a)
  241. if __name__ == "__main__":
  242. # test_node()
  243. asyncio.run(main())