#!/usr/bin/env python3
r"""eval_row33.py (vendored subset) -- the host module the row-33 eps^0 engine (row33_eps0.py) looks
up under this name: the exact-rational differential systems of row 33 and the fixed-eps Taylor-step
transport class the engine drives at eps = 0 on its block (Laurent-layer) systems.

Kept from the research script eval_row33.py, byte for byte: the data load (row33_data.json beside
this file: the 4-master self-energy w-system `sigvp_de`, the 8-master one-loop box w-system
`box1eq_de`, and the recorded reference strings under `oracles`) and the class RatDE (exact-rational
DE -> fixed-eps rational functions of w via sympy; Taylor-series stepping `step`, path transport
`path`/`to` with targets read off the local series).  Everything else in that script -- the fixed-eps
dispersive value, the pole-layer demo, the retired grid extrapolation -- is served by
lbl3vp-evaluate.py or not needed by this leg, and is not here.
Dependencies: mpmath, sympy.
"""
import os, json
import mpmath as mp
import sympy as sp

HERE = os.path.dirname(os.path.abspath(__file__))

DATA = json.load(open(os.path.join(HERE, 'row33_data.json')))
ORC = DATA['oracles']


# ---------------------------------------------------------------------
# Fixed-eps rational DE transport (Taylor-step ODE integrator).  The engine
# instantiates it at eps = 0 on the block systems of the Laurent layers.
# ---------------------------------------------------------------------
class RatDE:
    def __init__(self, de_json, eps_rat, sings):
        self.masters = [tuple(m) for m in de_json['masters']]
        self.n = len(self.masters)
        W, D = sp.symbols('w d')
        d_sym = sp.Rational(4) - 2 * eps_rat
        self.funcs = [[None] * self.n for _ in range(self.n)]
        for i in range(self.n):
            for j in range(self.n):
                e = sp.sympify(de_json['A'][i][j], locals={'w': W, 'd': D})
                if e == 0:
                    continue
                e = sp.cancel(e.subs(D, d_sym))
                num, den = sp.fraction(e)
                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.funcs[i][j] = (nc, dc)
        self.sings = [mp.mpc(s) for s in sings]

    def nearest(self, w):
        w = mp.mpc(w)
        return min(abs(w - s) for s in self.sings)

    def _shift(self, coefs, w0):
        p = [mp.mpc(coefs[0])]
        for c in coefs[1:]:
            pn = [mp.mpc(0)] * (len(p) + 1)
            for m, a in enumerate(p):
                pn[m] += a * w0
                pn[m + 1] += a
            pn[0] += c
            p = pn
        return p

    def step(self, M, w0, h, N):
        n = self.n
        ser = [[] for _ in range(N + 1)]
        for i in range(n):
            for j in range(n):
                cd = self.funcs[i][j]
                if cd is None:
                    continue
                nc, dc = cd
                ph = self._shift(nc, w0)
                qh = self._shift(dc, w0)
                q0 = qh[0]
                r = [mp.mpc(1) / q0]
                for m in range(1, N + 1):
                    s = mp.mpc(0)
                    for l in range(1, min(m, len(qh) - 1) + 1):
                        s += qh[l] * r[m - l]
                    r.append(-s / q0)
                for m in range(N + 1):
                    s = mp.mpc(0)
                    for l in range(min(m, len(ph) - 1) + 1):
                        s += ph[l] * r[m - l]
                    if s != 0:
                        ser[m].append((i, j, s))
        C = [list(M)]
        for m in range(N):
            s = [mp.mpc(0)] * n
            for l in range(m + 1):
                cv = C[m - l]
                for (ri, ci, a) in ser[l]:
                    s[ri] += a * cv[ci]
            inv = mp.mpf(1) / (m + 1)
            C.append([x * inv for x in s])
        v = [mp.mpc(0)] * n
        hp = mp.mpc(1)
        for m in range(N + 1):
            cm = C[m]
            for i in range(n):
                v[i] += cm[i] * hp
            hp *= h
        return v, C

    def path(self, M0, w_from, targets, N, sf, deflect=None):
        n = self.n
        M = list(M0)
        w = mp.mpc(w_from)
        out = {}
        ti = 0
        nT = len(targets)
        end = mp.mpc(targets[-1])
        sgn = 1 if end.real > w.real else -1
        wps = []
        for s in sorted(deflect or [], key=lambda x: sgn * x):
            if min(w.real, end.real) < s < max(w.real, end.real):
                r = mp.mpf('0.35')
                wps += [mp.mpc(s - sgn * r, 0), mp.mpc(s, r), mp.mpc(s + sgn * r, 0)]
        wps.append(end)
        for wp in wps:
            while abs(w - wp) > mp.mpf('1e-200'):
                d = self.nearest(w)
                hmax = sf * d
                dirn = wp - w
                stp = dirn if abs(dirn) <= hmax else dirn / abs(dirn) * hmax
                Mnew, C = self.step(M, w, stp, N)
                w_next = w + stp
                if abs(w.imag) < mp.mpf('1e-50') and abs(w_next.imag) < mp.mpf('1e-50') and stp != 0:
                    while ti < nT:
                        t = mp.mpc(targets[ti])
                        frac = ((t - w) / stp).real
                        if frac < -mp.mpf('1e-50') or frac > 1 + mp.mpf('1e-50'):
                            break
                        hh = t - w
                        v = [mp.mpc(0)] * n
                        hp = mp.mpc(1)
                        for m in range(len(C)):
                            cm = C[m]
                            for i in range(n):
                                v[i] += cm[i] * hp
                            hp *= hh
                        out[targets[ti]] = v
                        ti += 1
                M = Mnew
                w = w_next
            w = wp
        while ti < nT:
            out[targets[ti]] = list(M)
            ti += 1
        return out, M

    def to(self, M0, w_from, w_to, N, sf, deflect=None):
        _, M = self.path(M0, w_from, [w_to], N, sf, deflect)
        return M
