import random
from dataclasses import dataclass
from typing import Any, Dict, List, Mapping, Sequence, Set, Tuple


@dataclass
class SimpleGraph:
    n: int
    adj: List[Set[int]]
    edges: List[Tuple[int, int]]

    @staticmethod
    def from_edges(n: int, edges: Sequence[Tuple[int, int]]) -> "SimpleGraph":
        adj: List[Set[int]] = [set() for _ in range(n)]
        uniq: List[Tuple[int, int]] = []
        for (u0, v0) in edges:
            u = int(u0)
            v = int(v0)
            if u == v:
                continue
            if u < 0 or v < 0 or u >= n or v >= n:
                continue
            a, b = (u, v) if u < v else (v, u)
            if b in adj[a]:
                continue
            adj[a].add(b)
            adj[b].add(a)
            uniq.append((a, b))
        return SimpleGraph(n=n, adj=adj, edges=uniq)

    @staticmethod
    def erdos_renyi(n: int, p: float, rng: random.Random) -> "SimpleGraph":
        edges: List[Tuple[int, int]] = []
        for u in range(n):
            for v in range(u + 1, n):
                if rng.random() < p:
                    edges.append((u, v))
        return SimpleGraph.from_edges(n, edges)

    @staticmethod
    def stochastic_block_model(
        n: int,
        sizes: Sequence[int],
        p_matrix: Sequence[Sequence[float]],
        rng: random.Random,
        *,
        shuffle_nodes: bool = False,
    ) -> "SimpleGraph":
        if n <= 1:
            raise ValueError("stochastic_block_model requires n > 1")

        sizes_i = [int(x) for x in sizes]
        if not sizes_i:
            raise ValueError("stochastic_block_model requires at least one block")
        if any(int(s) <= 0 for s in sizes_i):
            raise ValueError("stochastic_block_model requires all sizes > 0")
        if int(sum(sizes_i)) != int(n):
            raise ValueError("stochastic_block_model requires sum(sizes) == n")

        k = int(len(sizes_i))
        if len(p_matrix) != k:
            raise ValueError("stochastic_block_model requires p_matrix to be k x k")
        for row in p_matrix:
            if len(row) != k:
                raise ValueError("stochastic_block_model requires p_matrix to be k x k")

        p_mat: List[List[float]] = []
        for i in range(k):
            prow: List[float] = []
            for j in range(k):
                p = float(p_matrix[i][j])
                if p < 0.0 or p > 1.0:
                    raise ValueError("stochastic_block_model requires p_matrix values in [0, 1]")
                prow.append(float(p))
            p_mat.append(prow)

        block_of_node: List[int] = []
        for bi, s in enumerate(sizes_i):
            block_of_node.extend([int(bi)] * int(s))
        if shuffle_nodes:
            rng.shuffle(block_of_node)

        edges: List[Tuple[int, int]] = []
        for u in range(n):
            bu = int(block_of_node[int(u)])
            for v in range(u + 1, n):
                bv = int(block_of_node[int(v)])
                if rng.random() < float(p_mat[bu][bv]):
                    edges.append((int(u), int(v)))

        return SimpleGraph.from_edges(n, edges)

    @staticmethod
    def ring_lattice(n: int, k: int) -> "SimpleGraph":
        if k <= 0 or k >= n:
            raise ValueError("ring_lattice requires 0 < k < n")
        if k % 2 != 0:
            raise ValueError("ring_lattice requires even k")
        half = k // 2
        edges: List[Tuple[int, int]] = []
        for u in range(n):
            for d in range(1, half + 1):
                edges.append((u, (u + d) % n))
        return SimpleGraph.from_edges(n, edges)

    @staticmethod
    def barabasi_albert(n: int, m: int, rng: random.Random) -> "SimpleGraph":
        if m <= 0 or m >= n:
            raise ValueError("barabasi_albert requires 0 < m < n")

        m0 = m + 1
        edges: List[Tuple[int, int]] = []
        for u in range(m0):
            for v in range(u + 1, m0):
                edges.append((u, v))

        degrees = [0 for _ in range(n)]
        for (u, v) in edges:
            degrees[u] += 1
            degrees[v] += 1

        stubs: List[int] = []
        for i in range(m0):
            stubs.extend([i] * degrees[i])

        for new_node in range(m0, n):
            targets = set()
            if not stubs:
                while len(targets) < m:
                    targets.add(rng.randrange(0, new_node))
            else:
                tries = 0
                max_tries = 10 * m * m + 100
                while len(targets) < m and tries < max_tries:
                    targets.add(stubs[rng.randrange(0, len(stubs))])
                    tries += 1
                while len(targets) < m:
                    targets.add(rng.randrange(0, new_node))

            for t in targets:
                tt = int(t)
                edges.append((new_node, tt))
                degrees[new_node] += 1
                degrees[tt] += 1
                stubs.append(tt)
                stubs.append(new_node)

        return SimpleGraph.from_edges(n, edges)


def detect_communities(g: SimpleGraph, cfg: Mapping[str, Any], rng: random.Random) -> Tuple[List[int], List[List[int]]]:
    typ = str(cfg.get("type", "label_propagation")).lower()
    if typ != "label_propagation":
        raise ValueError(f"Unsupported meso.community_detection.type: {typ}")

    max_iters = int(cfg.get("max_iters", 25))
    if max_iters <= 0:
        raise ValueError("meso.community_detection.max_iters must be > 0")
    shuffle = bool(cfg.get("shuffle", True))

    labels = [int(i) for i in range(g.n)]
    nodes = [int(i) for i in range(g.n)]

    for _ in range(max_iters):
        changed = False
        if shuffle:
            rng.shuffle(nodes)

        for u in nodes:
            nbrs = g.adj[u]
            if not nbrs:
                continue

            counts: Dict[int, int] = {}
            for v in nbrs:
                lv = int(labels[int(v)])
                counts[lv] = int(counts.get(lv, 0)) + 1

            best_count = max(int(x) for x in counts.values())
            cands = [int(lab) for (lab, cnt) in counts.items() if int(cnt) == int(best_count)]
            cands.sort()
            best_label = int(rng.choice(cands)) if len(cands) > 1 else int(cands[0])

            if int(best_label) != int(labels[u]):
                labels[u] = int(best_label)
                changed = True

        if not changed:
            break

    uniq = sorted({int(x) for x in labels})
    remap = {int(old): int(i) for i, old in enumerate(uniq)}
    norm = [int(remap[int(x)]) for x in labels]

    comms: List[List[int]] = [[] for _ in range(len(uniq))]
    for i, lab in enumerate(norm):
        comms[int(lab)].append(int(i))

    return norm, comms


def _add_edge(g: SimpleGraph, u: int, v: int) -> None:
    if u == v:
        return
    a, b = (u, v) if u < v else (v, u)
    if b in g.adj[a]:
        return
    g.adj[a].add(b)
    g.adj[b].add(a)
    g.edges.append((a, b))


def _remove_edge(g: SimpleGraph, u: int, v: int) -> None:
    a, b = (u, v) if u < v else (v, u)
    if b not in g.adj[a]:
        return
    g.adj[a].remove(b)
    g.adj[b].remove(a)
    try:
        g.edges.remove((a, b))
    except ValueError:
        g.edges = [(x, y) for x, y in g.edges if not (x == a and y == b)]


def add_edge(g: SimpleGraph, u: int, v: int) -> None:
    _add_edge(g, u, v)


def remove_edge(g: SimpleGraph, u: int, v: int) -> None:
    _remove_edge(g, u, v)


def rewire_edges(g: SimpleGraph, prob: float, rng: random.Random, max_tries: int) -> int:
    if prob <= 0:
        return 0
    rewired = 0
    snapshot = list(g.edges)
    for (u, v) in snapshot:
        if rng.random() >= prob:
            continue
        _remove_edge(g, u, v)
        for _ in range(max_tries):
            a = rng.randrange(0, g.n)
            b = rng.randrange(0, g.n)
            if a == b:
                continue
            x, y = (a, b) if a < b else (b, a)
            if y in g.adj[x]:
                continue
            _add_edge(g, x, y)
            rewired += 1
            break
    return rewired
