#!/usr/bin/env python3
"""Write response_matrices/A_S1.json, A_S2.json, A_S3.json: the response matrices A(eps) of Eq. (27) for the three PX2 top
sectors at the twelve sampled eps, with ball midpoints truncated to 60 significant digits (the source enclosure radii, at most about 1e-97, are recorded per eps),
row labels (master integrals of the eta-deformed system, exponent vectors over D_1..D_22) and column labels (the solver's
labels of the large-eta boundary integrals).  Also recomputes rank A at each eps from the truncated data (singular values)
as a check of Table 3.  usage: build_response_matrices.py FP_S1.json FP_S2.json FP_S3.json OUTDIR"""
import sys, json, re, os
import mpmath as mp
from common import sha256_file, dump
mp.mp.dps = 70
files = dict(zip(('S1', 'S2', 'S3'), sys.argv[1:4])); outdir = sys.argv[4]
BALL = re.compile(r'\[([-+0-9.e]+) \+/- ([-+0-9.e]+)\]')
def parse(s):
    s = s.strip()
    if s in ('0', '(0)', '(0 + 0*I)'): return mp.mpf(0), mp.mpf(0), 0.0
    parts = BALL.findall(s)
    rad = max(float(r) for _, r in parts) if parts else 0.0
    if '*I' in s and len(parts) == 2:
        return mp.mpf(parts[0][0]), mp.mpf(parts[1][0]), rad
    if len(parts) == 1:
        if s.startswith('(0 +') or s.startswith('(0+'): return mp.mpf(0), mp.mpf(parts[0][0]), rad
        return mp.mpf(parts[0][0]), mp.mpf(0), rad
    raise ValueError(s[:80])
def t60(x):
    return '0' if x == 0 else mp.nstr(x, 60, strip_zeros=False)
summary = {}
for S, path in files.items():
    d = json.load(open(path))
    pref, sub, epsl, A = d['preferred'], d['sub_masters'], d['eps'], d['A']
    rows = [[int(t) for t in p.split('|')[1:]] for p in pref]
    cols = [c.split('|', 1)[1].replace('|', ',') for c in sub]
    out = {'sector': S, 'top_sector_label': {'S1': 1077759, 'S2': 1091071, 'S3': 599551}[S],
           'definition': 'M_i(eps) = sum_j A_ij(eps) b_j: M = master integrals of the eta-deformed system at eta = 0 (rows), b = boundary integrals of the large-eta region (columns); Eq. (27) of the paper',
           'rows_master_integrals': {'note': 'exponent vectors (n_1..n_22) over D_1..D_22 of the PX2 family; row index = response_matrix_row in boundary_constants_PX2.json', 'index_vectors': rows},
           'columns_boundary_integrals': {'note': "the solver's labels for the boundary integrals of the large-eta region (index vectors over that region's own propagator list, which is not reproduced here); ranks and left null vectors do not depend on the column labeling", 'labels': cols},
           'n_rows': len(rows), 'n_columns': len(cols), 'kinematic_point': {'gamma': '5/4', 'q^2': '-1'},
           'entries': 'A[k][i][j] at eps = eps[k]: strings re, im (ball midpoints truncated to 60 significant digits; the enclosure radii of the source are below max_radius[k])',
           'eps': [], 'max_radius': [], 'A': [], 'rank_check': []}
    maxrad_all = 0.0
    for k in range(len(epsl)):
        e_re, e_im, erad = parse(epsl[k])
        out['eps'].append(t60(e_re))
        Mk = []; radk = erad
        rowsA = []
        for i in range(len(rows)):
            r_ = []
            rowv = []
            for j in range(len(cols)):
                re_, im_, rad = parse(A[k][i][j]); radk = max(radk, rad)
                r_.append({'re': t60(re_), 'im': t60(im_)} if (re_ != 0 or im_ != 0) else 0)
                rowv.append(mp.mpc(re_, im_))
            Mk.append(r_); rowsA.append(rowv)
        out['A'].append(Mk); out['max_radius'].append(f'{radk:.1e}'); maxrad_all = max(maxrad_all, radk)
        # rank from singular values of the truncated midpoints
        M = mp.matrix(rowsA)
        U, Sg, V = mp.svd_c(M)
        svals = sorted([abs(Sg[i]) for i in range(len(Sg))], reverse=True)
        # gap: largest ratio between consecutive singular values
        gaps = [(float(mp.log10(svals[i]/svals[i+1])) if svals[i+1] > 0 else 999.0, i+1) for i in range(len(svals)-1)]
        g, rk = max(gaps)
        out['rank_check'].append({'rank': rk, 'gap_orders_of_magnitude': round(g, 1), 'smallest_kept_sv': mp.nstr(svals[rk-1], 5), 'largest_dropped_sv': mp.nstr(svals[rk], 5) if rk < len(svals) else '0'})
    out['source_sha256'] = {'description': f'response matrix of {S} at full precision (auxiliary-mass-flow solver output)', 'sha256': sha256_file(path)}
    fn = os.path.join(outdir, f'A_{S}.json')
    dump(out, fn)
    ranks = sorted(set(r['rank'] for r in out['rank_check']))
    summary[S] = (len(rows), len(cols), ranks, maxrad_all, os.path.getsize(fn))
    print(f"wrote A_{S}.json: {len(rows)} x {len(cols)} at {len(epsl)} eps; rank {ranks} (min gap {min(r['gap_orders_of_magnitude'] for r in out['rank_check'])} orders); max source radius {maxrad_all:.1e}; {os.path.getsize(fn)} bytes")
