fork download
  1. # @file tree.007.py
  2. # @ingroup experimental
  3. # Recursive red & black tree.
  4. # @date 01/05/2023
  5.  
  6. class Color:
  7. RED = 0
  8. BLACK = 1
  9.  
  10. class Node:
  11. def __init__(self, data):
  12. self.child = [None] * 2
  13. self.color = Color.RED
  14. self.data = data
  15.  
  16. def __setitem__(self, i, v):
  17. self.child[i] = v
  18.  
  19. def __getitem__(self, i):
  20. return self.child[i]
  21.  
  22. def is_red(x):
  23. return x and x.color == Color.RED
  24.  
  25. def try_setcolor(x, v):
  26. if not x:
  27. return False
  28. x.color = v
  29. return True
  30.  
  31. def rotate(x, R):
  32. L = 1-R
  33. y = x[L]
  34. x[L] = y[R]
  35. y[R] = x
  36. return y
  37.  
  38. class Tree:
  39. def __init__(self):
  40. self.root = None
  41.  
  42. ## Insert. ##
  43.  
  44. def insert(self, key):
  45. self.root = Tree._insert(self.root, key)
  46. self.root.color = Color.BLACK
  47.  
  48. def _insert(root, key):
  49. if not root:
  50. return Node(key)
  51. elif key < root.data:
  52. root[0] = Tree._insert(root[0], key)
  53. elif key > root.data:
  54. root[1] = Tree._insert(root[1], key)
  55. else:
  56. raise KeyError('Key already exists!')
  57. return Tree._insert_balance(root)
  58.  
  59. def _insert_balance(root):
  60. if Node.is_red(root):
  61. pass
  62. elif Node.is_red(root[0]):
  63. root = Tree._insert_balance_lower(root, 0)
  64. elif Node.is_red(root[1]):
  65. root = Tree._insert_balance_lower(root, 1)
  66. return root
  67.  
  68. def _insert_balance_lower(root, R):
  69. L = 1-R
  70. if Node.is_red(root[R][L]):
  71. # 4-node normalize (1).
  72. root[R] = Node.rotate(root[R], R)
  73. if Node.is_red(root[R][R]):
  74. # 4-node normalize (2).
  75. root[R].color = Color.BLACK
  76. root.color = Color.RED
  77. root = Node.rotate(root, L)
  78. if Node.is_red(root[L]):
  79. # 4-node split.
  80. root[L].color = Color.BLACK
  81. root[R].color = Color.BLACK
  82. root.color = Color.RED
  83. return root
  84.  
  85. ## Remove. ##
  86.  
  87. def remove(self, key):
  88. self.root = Tree._remove(self.root, key, [None])
  89.  
  90. def _remove(root, key, last):
  91. if not root:
  92. raise KeyError('Key does not exist!')
  93. elif key < root.data:
  94. root[0] = Tree._remove(root[0], key, last)
  95. elif key > root.data:
  96. root[1] = Tree._remove(root[1], key, last)
  97. elif root[1]:
  98. # Interior node; delete inorder successor.
  99. root[1] = Tree._delete_minimum(root[1], root, last)
  100. else:
  101. return Tree._delete_node(root)
  102. return Tree._remove_balance(root, last)
  103.  
  104. def _delete_minimum(root, interior, last):
  105. if root[0]:
  106. root[0] = Tree._delete_minimum(root[0], interior, last)
  107. else:
  108. interior.data = root.data
  109. return Tree._delete_node(root)
  110. return Tree._remove_balance(root, last)
  111.  
  112. def _delete_node(root):
  113. if Node.try_setcolor(root[0], Color.BLACK):
  114. return root[0]
  115. if Node.try_setcolor(root[1], Color.BLACK):
  116. return root[1]
  117. return None
  118.  
  119. def _remove_balance(root, last):
  120. is_short = False
  121. if not (root[0] or root[1]):
  122. # Root is terminal node.
  123. pass
  124. elif root[0] is last[0]:
  125. root, is_short = Tree._remove_balance_lower(root, 0)
  126. elif root[1] is last[0]:
  127. root, is_short = Tree._remove_balance_lower(root, 1)
  128. # Signal parent.
  129. last[0] = root if is_short else None
  130. return root
  131.  
  132. def _remove_balance_lower(root, R):
  133. L = 1-R
  134. if Node.is_red(root[L]):
  135. # 3-node parent; sibling is 'far' middle.
  136. root[L].color = Color.BLACK
  137. root.color = Color.RED
  138. root = Node.rotate(root, R)
  139. root[R], _ = Tree._remove_balance_close(root[R], R, L)
  140. return root, False
  141. return Tree._remove_balance_close(root, R, L)
  142.  
  143. def _remove_balance_close(root, R, L):
  144. is_short = False
  145. if Node.is_red(root[L][R]):
  146. # Borrow sibling (1).
  147. root[L][R].color = Color.BLACK
  148. root[L].color = Color.RED
  149. root[L] = Node.rotate(root[L], L)
  150. if Node.is_red(root[L][L]):
  151. # Borrow sibling (2).
  152. root[L][L].color = Color.BLACK
  153. root[L].color = root.color
  154. root.color = Color.BLACK
  155. root = Node.rotate(root, R)
  156. else:
  157. # Borrow parent/shorten sibling.
  158. is_short = not Node.is_red(root)
  159. root[L].color = Color.RED
  160. root.color = Color.BLACK
  161. return root, is_short
  162.  
  163. ## Utility. ##
  164.  
  165. def inorder(self):
  166. return Tree._inorder(self.root)
  167.  
  168. def _inorder(root):
  169. if root:
  170. yield from Tree._inorder(root[0])
  171. yield root.data
  172. yield from Tree._inorder(root[1])
  173.  
  174. def levelorder(self):
  175. q = [self.root]
  176. while q.count(None) != len(q):
  177. nodes = q
  178. q = []
  179. r = []
  180. for node in nodes:
  181. assert node, 'Tree unbalanced!'
  182. assert node.color == Color.BLACK, 'Color violation!'
  183. if Node.is_red(node[0]):
  184. # 3-node; left-leaning.
  185. r.append([node[0].data, node.data])
  186. q.extend((node[0][0], node[0][1], node[1]))
  187. elif Node.is_red(node[1]):
  188. # 3-node; right-leaning.
  189. r.append([node.data, node[1].data])
  190. q.extend((node[0], node[1][0], node[1][1]))
  191. else:
  192. # 2-node.
  193. r.append([node.data])
  194. q.extend((node[0], node[1]))
  195. yield r
  196.  
  197. ### Testing. ###
  198.  
  199. import random
  200.  
  201. def pretty_print_tree(t):
  202. lines = (''.join(map(str, x)) for x in t.levelorder())
  203. print('>', next(lines, None))
  204. for line in lines:
  205. print(' ', line)
  206.  
  207. n = 2**3-1
  208. a = list(range(1, 1+n))
  209. t = Tree()
  210.  
  211. b = random.sample(a, k=n)
  212. print('Insert:', b)
  213. for i in b:
  214. t.insert(i)
  215. pretty_print_tree(t)
  216.  
  217. b = list(t.inorder())
  218. assert a == b, 'Bad order!'
  219.  
  220. b = random.sample(a, k=n)
  221. print('Remove:', b)
  222. for i in b:
  223. t.remove(i)
  224. pretty_print_tree(t)
Success #stdin #stdout 0.07s 14272KB
stdin
Standard input is empty
stdout
Insert: [6, 7, 5, 4, 3, 1, 2]
> [6]
> [6, 7]
> [6]
  [5][7]
> [6]
  [4, 5][7]
> [4, 6]
  [3][5][7]
> [4, 6]
  [1, 3][5][7]
> [4]
  [2][6]
  [1][3][5][7]
Remove: [7, 3, 5, 2, 6, 4, 1]
> [2, 4]
  [1][3][5, 6]
> [4]
  [1, 2][5, 6]
> [4]
  [1, 2][6]
> [4]
  [1][6]
> [1, 4]
> [1]
> None