import random
from dataclasses import dataclass
from typing import List, Sequence

from graph import SimpleGraph


@dataclass
class DynamicsParams:
    beta: float
    p_flip: float
    theta: float


def init_states(n: int, p_coop: float, rng: random.Random) -> List[int]:
    states: List[int] = []
    for _ in range(n):
        states.append(1 if rng.random() < p_coop else 0)
    return states


def _pair_payoff(ai: int, aj: int, theta: float) -> float:
    if ai == 1 and aj == 1:
        return 1.0
    if ai == 1 and aj == 0:
        return -float(theta)
    if ai == 0 and aj == 1:
        return 1.0 + float(theta)
    return 0.0


def _node_payoff(i: int, states: Sequence[int], g: SimpleGraph, theta: float) -> float:
    si = 1 if states[i] else 0
    s = 0.0
    for j in g.adj[i]:
        sj = 1 if states[j] else 0
        s += _pair_payoff(si, sj, theta)
    return s


def _sigmoid(x: float) -> float:
    if x >= 0.0:
        z = pow(2.718281828459045, -x)
        return 1.0 / (1.0 + z)
    z = pow(2.718281828459045, x)
    return z / (1.0 + z)


def async_imitation_updates(states: List[int], g: SimpleGraph, params: DynamicsParams, rng: random.Random, n_updates: int) -> None:
    n = g.n
    beta = float(params.beta)
    p_flip = float(params.p_flip)
    theta = float(params.theta)

    for _ in range(int(n_updates)):
        i = rng.randrange(0, n)
        if not g.adj[i]:
            continue
        j = rng.choice(tuple(g.adj[i]))

        ui = _node_payoff(i, states, g, theta)
        uj = _node_payoff(j, states, g, theta)

        p_adopt = _sigmoid(beta * (uj - ui))
        if rng.random() < p_adopt:
            states[i] = 1 if states[j] else 0

        if p_flip > 0.0 and rng.random() < p_flip:
            states[i] = 0 if states[i] else 1


def run_window(states: List[int], g: SimpleGraph, params: DynamicsParams, rng: random.Random, sweeps: int) -> None:
    n = g.n
    for _ in range(int(sweeps)):
        async_imitation_updates(states, g, params=params, rng=rng, n_updates=n)
