class BinHeap: def __init__(self, nodes): self.heap = [] self.pos = [-1] * nodes def sift_up(self, idx): heap = self.heap pos = self.pos backup = heap[idx] cur = idx while cur > 0: par = (cur - 1) // 2 if backup[1] >= heap[par][1]: break heap[cur] = heap[par] pos[heap[cur][0]] = cur cur = par heap[cur] = backup pos[backup[0]] = cur def sift_down(self, idx): heap = self.heap pos = self.pos n = len(heap) backup = heap[idx] cur = idx while True: left = cur * 2 + 1 if left >= n: break right = left + 1 tar = right if right < n and heap[right][1] < heap[left][1] else left if backup[1] <= heap[tar][1]: break heap[cur] = heap[tar] pos[heap[cur][0]] = cur cur = tar heap[cur] = backup pos[backup[0]] = cur def add(self, key, dist): idx = len(self.heap) self.heap.append([key, dist]) self.pos[key] = idx self.sift_up(idx) def extract_min(self): heap = self.heap pos = self.pos if not heap: return None min_node, min_dist = heap[0] pos[min_node] = -1 last = heap.pop() if heap: heap[0] = last pos[last[0]] = 0 self.sift_down(0) return (min_node, min_dist) def decrease_key(self, key, new_dist): idx = self.pos[key] if idx == -1: return None node = self.heap[idx] if new_dist >= node[1]: return None node[1] = new_dist self.sift_up(idx)