#!python
#
#    JKS - Measurement database system
#    Copyright (C) 2013-2026  Christoph Lehner (christoph.lehner@ur.de, https://github.com/lehner/jks)
#
#    This program is free software; you can redistribute it and/or modify
#    it under the terms of the GNU General Public License as published by
#    the Free Software Foundation; either version 2 of the License, or
#    (at your option) any later version.
#
#    This program is distributed in the hope that it will be useful,
#    but WITHOUT ANY WARRANTY; without even the implied warranty of
#    MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE.  See the
#    GNU General Public License for more details.
#
#    You should have received a copy of the GNU General Public License along
#    with this program; if not, write to the Free Software Foundation, Inc.,
#    51 Franklin Street, Fifth Floor, Boston, MA 02110-1301 USA.
#
import jks, sys, numpy as np
if len(sys.argv) != 6:
    print("%s database.jks tag1 t1 tag2 t2" % sys.argv[0])
    print("")
    print("Prints the correlation between tag1[t1] and tag2[t2]:")
    print("- stat: from the jackknife blocks")
    print("- each tagged systematic '!v' is a single shifted evaluation, i.e. a fully")
    print("  correlated shift d_v = block_v - mean; its correlation is sign(d1 d2), so")
    print("  the two shifts are printed instead")
    print("- sys:   sum_v d1 d2 / sqrt(sum_v d1^2 sum_v d2^2), all variations combined")
    print("- total: from tcov() = stat + sum_v d_v d_v^T, as used downstream")
    sys.exit(0)

db, tag1, t1, tag2, t2 = sys.argv[1], sys.argv[2], int(sys.argv[3]), sys.argv[4], int(sys.argv[5])

res = jks.resamples(db)
for tag in (tag1, tag2):
    assert tag in res.set, f"tag {tag!r} not in {db}"

# one joint 2-vector, so that stat blocks and variations are aligned
jk = res.apply(lambda r: [np.atleast_1d(r[tag1])[t1], np.atleast_1d(r[tag2])[t2]])
mean = np.asarray(jk.mean(), float)

def corr(c):
    c = np.asarray(c, float)
    d = c[0, 0] * c[1, 1]
    return c[0, 1] / np.sqrt(d) if d > 0.0 else float("nan")

cs = np.asarray(jk.cov(), float)
print(f"{tag1}[{t1}] = {mean[0]:.10g} +- {np.sqrt(cs[0, 0]):.3e} (stat)")
print(f"{tag2}[{t2}] = {mean[1]:.10g} +- {np.sqrt(cs[1, 1]):.3e} (stat)")
print("")
print(f"stat correlation:   {corr(cs): .6f}")

# variations of the database that shift neither quantity are not listed
shifts = [(v, np.asarray(jk.blocks[jk.tags.index("!" + v)], float) - mean) for v in jk.vars()]
shifts = [(v, d) for v, d in shifts if d.any()]
if shifts:
    print("")
    print(f" variation              shift {tag1}[{t1}]   shift {tag2}[{t2}]   sign")
    for v, d in shifts:
        sg = "+1" if d[0] * d[1] > 0 else ("-1" if d[0] * d[1] < 0 else " -")
        print(f" {v:20s}  {d[0]: .3e}          {d[1]: .3e}         {sg}")
    csys = sum(np.outer(d, d) for _, d in shifts)   # = sum_v cov(v)
    print("")
    if csys[0, 0] > 0.0 and csys[1, 1] > 0.0:
        print(f"sys correlation:    {corr(csys): .6f}   (all variations combined)")
    else:
        which = f"{tag1}[{t1}]" if csys[0, 0] == 0.0 else f"{tag2}[{t2}]"
        print(f"sys correlation:    undefined ({which} has no systematic shift)")
    print(f"total correlation:  {corr(jk.tcov()): .6f}   (tcov: stat + sys)")
