# Auto-generated Stone Cube framework (see notebook for example usage)
import numpy as np
from dataclasses import dataclass

@dataclass
class XBridgeCell:
    mode: str = "latch"
    state: float = 0.0
    c_mid: float = 0.0
    def step(self, inp: float, alpha: float = 0.2, beta: float = 0.98):
        if self.mode == "latch":
            self.state = (1 - alpha) * self.state + alpha * inp
        elif self.mode == "osc":
            self.state = -self.state * beta + (1 - beta) * inp
        elif self.mode == "staged":
            self.c_mid = 0.9 * self.c_mid + 0.1 * inp
            self.state = 0.9 * self.state + 0.1 * self.c_mid
        self.state = float(np.clip(self.state, -1.0, 1.0))
        return self.state

class StoneCube3D:
    def __init__(self, n=1, seed=0, eta=0.8, sigma=1.0, r=2.0, 
                 kappa=(0.05, 0.05, 0.05), phase="base"):
        self.n = n
        self.size = 2*n + 1
        self.eta = eta
        self.sigma = sigma
        self.r = r
        self.kx, self.ky, self.kz = kappa
        self.phase = phase
        rng = np.random.default_rng(seed)
        self.u = rng.uniform(-0.5, 0.5, size=(self.size, self.size, self.size))
        self.cells = np.empty((self.size, self.size, self.size), dtype=object)
        for ix in range(self.size):
            for iy in range(self.size):
                for iz in range(self.size):
                    mode = ("latch" if phase == "base" else 
                            "osc" if phase == "reversed" else
                            "staged")
                    self.cells[ix, iy, iz] = XBridgeCell(mode=mode, state=float(self.u[ix,iy,iz]))
    def G_inf(self, u):
        return u + 0.75 * u
    def g_local(self, u):
        return u + self.r * u * (1.0 - u)
    def S_sigma(self, x):
        return self.sigma * np.tanh(x / max(1e-8, self.sigma))
    def laplacian_aniso(self, U):
        xp = np.roll(U, -1, axis=0); xm = np.roll(U, 1, axis=0)
        yp = np.roll(U, -1, axis=1); ym = np.roll(U, 1, axis=1)
        zp = np.roll(U, -1, axis=2); zm = np.roll(U, 1, axis=2)
        xp[-1,:,:] = U[-1,:,:]; xm[0,:,:] = U[0,:,:]
        yp[:, -1,:] = U[:, -1,:]; ym[:, 0,:] = U[:, 0,:]
        zp[:, :, -1] = U[:, :, -1]; zm[:, :, 0] = U[:, :, 0]
        lap = self.kx * (xp + xm - 2*U) +               self.ky * (yp + ym - 2*U) +               self.kz * (zp + zm - 2*U)
        return lap
    def step(self):
        U = self.u
        if self.phase == "base":
            eta = min(1.0, max(0.8, self.eta)); r = 2.0
            kx, ky, kz = self.kx, self.ky, self.kz
        elif self.phase == "reversed":
            eta = min(0.6, max(0.3, self.eta)); r = 3.5
            kx, ky, kz = 0.5*self.kx, 0.5*self.ky, 0.5*self.kz
        else:
            eta = 1.0; r = 0.0
            kx, ky, kz = self.kx, self.ky, self.kz
        gU = self.g_local(U) if r != 0.0 else U
        lap = self.laplacian_aniso(U)
        bounded = self.S_sigma(gU + lap)
        unbounded = self.G_inf(U)
        U_next = (1.0 - eta) * unbounded + eta * bounded
        for ix in range(self.size):
            for iy in range(self.size):
                for iz in range(self.size):
                    self.cells[ix,iy,iz].state = float(U_next[ix,iy,iz])
                    U_next[ix,iy,iz] = self.cells[ix,iy,iz].step(U_next[ix,iy,iz])
        self.u = np.clip(U_next, -1.0, 1.0)
        return self.u
    def snapshot_points(self):
        n = self.n
        xs, ys, zs, vs = [], [], [], []
        for ix in range(self.size):
            for iy in range(self.size):
                for iz in range(self.size):
                    xs.append(ix - n); ys.append(iy - n); zs.append(iz - n)
                    vs.append(self.u[ix,iy,iz])
        import numpy as np
        return np.array(xs), np.array(ys), np.array(zs), np.array(vs)
