#!/usr/bin/env python3
r"""Row 33 symbolic spectral density rho_VP(w;eps) — fully de-AMFlowed.

rho_VP(w;eps) = rho_exact(w;eps) - c(eps) * psi_0(w;eps) * theta(w-9)

  rho_exact  : exact 2-body density (closed Gamma form; analytic on (1,inf))
  psi        : Frobenius solution of the homogeneous {0,1} 2x2 block at w=9,
               exponent alpha = d-2, leading vector v0 = (R9[0,1], alpha-R9[0,0]),
               exact-rational recursion (series radius 8)
  c(eps)     : ANALYTIC (closed form): c = A0/v0[0],
               A0 = -(pi/16) N2^2 9^{-(d-2)/2} 48^{(d-3)/2} 4^{-(d-2)/2}
                    * Beta((d-1)/2,(d-1)/2) * (2/3)^{d-2},  N2 = G(1-e)/G(2-2e)

Node coverage:
  1 < w < 9      rho_exact closed form (exact)
  9 <= w <= WSW  series form rho_exact - c*psi_0 (Frobenius series, K terms)
  w  > WSW       Im-4-vector transport of the EXACT sigvp DE from w=WSW with the
                 SYMBOLIC seed (ImM0, ImM1, 0, ImTB): the Im parts satisfy the
                 same real rational DE on (9,inf); ImM1 is closed-form via DE
                 row 0 (analytic d/dw of rho_exact) + c*psi_1.
No AMFlow anywhere. The exact-rational DE in the data file is the only input.
"""
import mpmath as mp
import sympy as sp


class Rho33:
    def __init__(self, VP, eps_rat, DPS, K=None, wsw=None):
        """VP = imported lbl3vp-evaluate module (for RatDE + DATA only —
        no stored boundary values are used)."""
        self.VP = VP
        self.eps_rat = sp.Rational(eps_rat)
        mp.mp.dps = DPS + 15
        self.dps = DPS
        self.eps = mp.mpf(self.eps_rat.p) / self.eps_rat.q
        e = self.eps
        self.d = 4 - 2 * e
        self.alpha = self.d - 2
        self.G = mp.gamma(1 + e) * mp.gamma(1 - e) / (e * mp.gamma(2 - 2 * e))
        self.T = mp.gamma(e) / (1 - e)             # tadpole m^2=1  (= -Gamma(eps-1))
        self.N2 = mp.gamma(1 - e) / mp.gamma(2 - 2 * e)
        self.wsw = mp.mpf(wsw if wsw is not None else 10)
        self.K = K if K is not None else int(1.11 * (DPS + 8)) + 4

        # --- fixed-eps A matrix as mpf coefficient arrays (num, den) ---
        W, D = sp.symbols('w d')
        dval = sp.Rational(4) - 2 * self.eps_rat
        self.A = [[sp.sympify(VP.DATA['sigvp_de']['A'][i][j], locals={'w': W, 'd': D})
                   for j in range(4)] for i in range(4)]
        self.Afx = [[None] * 4 for _ in range(4)]
        for i in range(4):
            for j in range(4):
                el = self.A[i][j]
                if el == 0:
                    continue
                num, den = sp.fraction(sp.cancel(el.subs(D, dval)))
                nc = [mp.mpf(sp.Rational(c).p) / mp.mpf(sp.Rational(c).q)
                      for c in sp.Poly(num, W).all_coeffs()]
                dc = [mp.mpf(sp.Rational(c).p) / mp.mpf(sp.Rational(c).q)
                      for c in sp.Poly(den, W).all_coeffs()]
                self.Afx[i][j] = (nc, dc)

        self._build_psi()
        self._c_analytic()
        self._de = None

    # ---------- exact 2-body pieces ----------
    def rho_exact(self, w):
        e, G = self.eps, self.G
        w = mp.mpf(w)
        u = w - 1
        P = -(1 - 2 * e) * (w + 1) / u + e * u / 6
        return -mp.pi * G * P * u ** (-2 * e) * w ** (e - 1)

    def drho_exact(self, w):
        """analytic d/dw of rho_exact (elementary)."""
        e, G = self.eps, self.G
        w = mp.mpf(w)
        u = w - 1
        P = -(1 - 2 * e) * (w + 1) / u + e * u / 6
        Pp = 2 * (1 - 2 * e) / u ** 2 + e / 6
        return -mp.pi * G * (Pp + P * (-2 * e / u + (e - 1) / w)) * u ** (-2 * e) * w ** (e - 1)

    def ImTB(self, w):
        e = self.eps
        w = mp.mpf(w)
        u = w - 1
        ImB = mp.pi * self.N2 * u ** (1 - 2 * e) * w ** (e - 1)
        return self.T * ImB

    def _Aat(self, i, j, w):
        cd = self.Afx[i][j]
        if cd is None:
            return mp.mpf(0)
        nc, dc = cd
        nv = mp.mpf(0)
        for c in nc:
            nv = nv * w + c
        dv = mp.mpf(0)
        for c in dc:
            dv = dv * w + c
        return nv / dv

    def D2b(self, w):
        """(ImM0, ImM1) of the analytically-continued 2-body disc (valid w>1,
        w != 9 only through the DE row-0 solve; formulas analytic at 9)."""
        w = mp.mpf(w)
        Im0 = -self.rho_exact(w)
        dIm0 = -self.drho_exact(w)
        Im1 = (dIm0 - self._Aat(0, 0, w) * Im0 - self._Aat(0, 3, w) * self.ImTB(w)) \
            / self._Aat(0, 1, w)
        return Im0, Im1

    # ---------- Frobenius psi at w=9 (exact-rational recursion in mpf) ----------
    def _taylor_at9(self, i, j, K):
        """Taylor coefficients (length K+2) of the ANALYTIC part of A[i][j] at
        w=9, plus the residue R (simple pole).  Power-series division in mpf."""
        cd = self.Afx[i][j]
        if cd is None:
            return mp.mpf(0), [mp.mpf(0)] * (K + 2)
        nc, dc = cd
        M = K + 4

        def shift(coefs):
            # p(w) -> Taylor coeffs in uu = w-9, ascending, length M
            p = [coefs[0]]
            for c in coefs[1:]:
                pn = [mp.mpf(0)] * (len(p) + 1)
                for m, a in enumerate(p):
                    pn[m] += a * 9
                    pn[m + 1] += a
                pn[0] += c
                p = pn
            p = p + [mp.mpf(0)] * M
            return p[:M]

        ns = shift(nc)
        ds = shift(dc)
        if abs(ds[0]) > mp.mpf('1e-30'):        # no pole at 9
            r = [ns[0] / ds[0]]
            for m in range(1, M):
                s = ns[m]
                for l in range(1, m + 1):
                    s -= ds[l] * r[m - l] if l < len(ds) else 0
                r.append(s / ds[0])
            return mp.mpf(0), r[:K + 2]
        # simple pole: den = uu * den1
        d1 = ds[1:] + [mp.mpf(0)]
        r = [ns[0] / d1[0]]                     # (num/den1) series
        for m in range(1, M):
            s = ns[m]
            for l in range(1, m + 1):
                s -= d1[l] * r[m - l] if l < len(d1) else 0
            r.append(s / d1[0])
        # entry = r0/uu + sum_{k>=0} r_{k+1} uu^k
        return r[0], r[1:K + 3][:K + 2]

    def _build_psi(self):
        K = self.K
        R = [[None] * 2 for _ in range(2)]
        S = [[None] * 2 for _ in range(2)]
        for i in range(2):
            for j in range(2):
                R[i][j], S[i][j] = self._taylor_at9(i, j, K)
        alpha = self.alpha
        v0 = [R[0][1], alpha - R[0][0]]
        ps = [v0]
        for k in range(1, K + 1):
            rhs = [mp.mpf(0), mp.mpf(0)]
            for l in range(k):
                for i in range(2):
                    rhs[i] += S[i][0][k - 1 - l] * ps[l][0] + S[i][1][k - 1 - l] * ps[l][1]
            a = (alpha + k) - R[0][0]
            b = -R[0][1]
            cc = -R[1][0]
            dd = (alpha + k) - R[1][1]
            det = a * dd - b * cc
            ps.append([(dd * rhs[0] - b * rhs[1]) / det,
                       (-cc * rhs[0] + a * rhs[1]) / det])
        self.R9 = R
        self.v0 = v0
        self.ps = ps

    def psi(self, w):
        """(psi_0, psi_1)(w) from the Frobenius series (use w-9 <= ~1.5)."""
        uu = mp.mpf(w) - 9
        a0 = mp.mpf(0)
        a1 = mp.mpf(0)
        up = mp.mpf(1)
        for k in range(self.K + 1):
            a0 += self.ps[k][0] * up
            a1 += self.ps[k][1] * up
            up *= uu
        f = uu ** self.alpha
        return f * a0, f * a1

    def _c_analytic(self):
        e, d = self.eps, self.d
        A0 = -mp.pi * self.N2 ** 2 * mp.mpf(9) ** (-(d - 2) / 2) * mp.mpf(48) ** ((d - 3) / 2) \
            * mp.mpf(4) ** (-(d - 2) / 2) * mp.beta((d - 1) / 2, (d - 1) / 2) \
            * (mp.mpf(2) / 3) ** (d - 2) / 16
        self.c = A0 / self.v0[0]

    # ---------- assembled density ----------
    def rho_series(self, w):
        """rho_VP for 9 <= w <= wsw+ (series window)."""
        p0, _ = self.psi(w)
        return self.rho_exact(w) - self.c * p0

    def seed_at(self, w0):
        """Symbolic Im-4-vector seed (ImM0, ImM1, 0, ImTB) at w0 > 9."""
        Im0, Im1 = self.D2b(w0)
        p0, p1 = self.psi(w0)
        return [Im0 + self.c * p0, Im1 + self.c * p1, mp.mpf(0), self.ImTB(w0)]

    def transport_tail(self, targets, N_TAY, SF):
        """rho_VP at sorted targets > wsw via exact-DE transport of the
        symbolic Im-vector seed from wsw.  Returns {w: rho}."""
        if self._de is None:
            self._de = self.VP.RatDE(self.VP.DATA['sigvp_de'], self.eps_rat, [0, 1, 9, -3])
        M0 = self.seed_at(self.wsw)
        res, _ = self._de.path(M0, self.wsw, sorted(targets), N_TAY, SF)
        # the transported vector IS the Im-vector (real DE, real seed):
        # rho_VP = -Im M0 = -(component 0);  mpc imaginary parts are numerically 0
        return {t: -mp.re(res[t][0]) for t in targets}

    def rho(self, w, tail_cache=None):
        w = mp.mpf(w)
        if w <= 1:
            return mp.mpf(0)
        if w < 9:
            return self.rho_exact(w)
        if w <= self.wsw:
            return self.rho_series(w)
        if tail_cache is not None and w in tail_cache:
            return tail_cache[w]
        raise KeyError(f"tail node {mp.nstr(w, 20)} not in cache")
