"""
Dijkstra's Algorithm — MAT5EJ302(1) Module III
Single-source shortest paths in weighted graphs with non-negative edge weights.
Uses Python's heapq (binary min-heap) for the priority queue.
Time: O((V + E) log V)
"""

import heapq


def dijkstra(graph, source):
    """
    Dijkstra's shortest path algorithm.

    Args:
        graph: dict mapping vertex -> list of (neighbour, weight) tuples
        source: starting vertex

    Returns:
        dist:  dict[vertex] -> shortest distance from source
        prev:  dict[vertex] -> previous vertex on shortest path
    """
    dist = {v: float('inf') for v in graph}
    prev = {v: None for v in graph}
    dist[source] = 0

    pq = [(0, source)]     # (distance, vertex)
    visited = set()

    while pq:
        d, u = heapq.heappop(pq)

        if u in visited:
            continue        # stale entry in heap
        visited.add(u)

        for v, weight in graph.get(u, []):
            if v not in visited:
                new_dist = dist[u] + weight
                if new_dist < dist[v]:
                    dist[v] = new_dist
                    prev[v] = u
                    heapq.heappush(pq, (new_dist, v))

    return dist, prev


def dijkstra_verbose(graph, source):
    """Dijkstra with step-by-step output."""
    dist = {v: float('inf') for v in graph}
    prev = {v: None for v in graph}
    dist[source] = 0
    pq = [(0, source)]
    visited = set()

    step = 1
    print(f"\nDijkstra from {source}:")
    print(f"Initial dist: { {v: (d if d < float('inf') else '∞') for v, d in dist.items()} }")

    while pq:
        d, u = heapq.heappop(pq)
        if u in visited:
            continue
        visited.add(u)
        print(f"\n  Step {step}: Extract {u} (dist={d}). Finalised: {sorted(visited)}")

        for v, weight in graph.get(u, []):
            if v not in visited:
                new_dist = dist[u] + weight
                if new_dist < dist[v]:
                    old = dist[v]
                    dist[v] = new_dist
                    prev[v] = u
                    heapq.heappush(pq, (new_dist, v))
                    old_str = '∞' if old == float('inf') else str(old)
                    print(f"    Relax edge {u}→{v} (w={weight}): dist[{v}] = {old_str} → {new_dist}")
        step += 1

    return dist, prev


def get_path(prev, source, target):
    """Reconstruct shortest path from source to target."""
    if prev.get(target) is None and target != source:
        return None   # unreachable
    path = []
    v = target
    while v is not None:
        path.append(v)
        v = prev.get(v)
    path.reverse()
    return path


if __name__ == "__main__":
    print("=== Dijkstra Example 1 (Module III trace) ===")
    graph1 = {
        'A': [('B', 4), ('C', 2)],
        'B': [('C', 1), ('D', 5)],
        'C': [('B', 1), ('D', 8), ('E', 10)],
        'D': [('E', 2)],
        'E': []
    }
    dist, prev = dijkstra(graph1, 'A')

    print(f"  {'Vertex':>8}  {'Distance':>10}  {'Path'}")
    for v in sorted(dist):
        path = get_path(prev, 'A', v)
        path_str = ' → '.join(str(x) for x in path) if path else 'unreachable'
        print(f"  {v:>8}  {dist[v]:>10}  {path_str}")

    print()
    print("=== Dijkstra Step-by-Step (Small Graph) ===")
    graph2 = {
        'S': [('A', 7), ('B', 3)],
        'A': [('C', 4)],
        'B': [('A', 1), ('C', 2)],
        'C': [('D', 5)],
        'D': []
    }
    dist2, prev2 = dijkstra_verbose(graph2, 'S')

    print(f"\nFinal distances from S:")
    for v in sorted(dist2):
        path = get_path(prev2, 'S', v)
        path_str = ' → '.join(str(x) for x in path) if path else 'unreachable'
        d = dist2[v] if dist2[v] < float('inf') else '∞'
        print(f"  S → {v}: dist = {d}, path = {path_str}")

    print()
    print("=== Dijkstra on Larger Numeric Graph ===")
    graph3 = {
        1: [(2, 10), (3, 3)],
        2: [(4, 2)],
        3: [(2, 4), (4, 8), (5, 2)],
        4: [(5, 5)],
        5: [(4, 1)]
    }
    dist3, prev3 = dijkstra(graph3, 1)
    print(f"  Shortest distances from vertex 1:")
    for v in sorted(dist3):
        path = get_path(prev3, 1, v)
        path_str = ' → '.join(str(x) for x in path) if path else 'unreachable'
        print(f"  1 → {v}: dist = {dist3[v]}, path = {path_str}")
