Files
wickra/bindings/java/gen_golden_test.py
T

294 lines
12 KiB
Python
Raw Normal View History

"""Generate src/test/java/org/wickra/GoldenAllTest.java: a reflection-driven
value-parity test that replays the shared golden input through every one of the
514 Java indicators and checks output bit-for-bit against the Rust reference
fixtures g_<Canonical>.csv. The per-indicator spec is embedded so the test has
no JSON dependency; a single reflective runner covers all archetypes.
Run from repo root: python bindings/java/gen_golden_test.py
"""
import json
import os
ROOT = os.path.normpath(os.path.join(os.path.dirname(__file__), "..", ".."))
MAN = json.load(open(os.path.join(ROOT, "testdata", "golden", "golden_manifest.json")))
def params_lit(ps):
if not ps:
return "new double[]{}"
return "new double[]{" + ", ".join(repr(float(p)) for p in ps) + "}"
specs = []
for e in MAN:
n = e.get("n", e.get("width", 0))
specs.append(f' new Spec("{e["canonical"]}", "{e["arch"]}", {params_lit(e["params"])}, {n}),')
specs_block = "\n".join(specs)
# Canonical core Indicator::name() per indicator, shared across every binding.
NAMES = json.load(open(os.path.join(ROOT, "testdata", "golden", "names.json")))
names_block = "\n".join(
f" NAMES.put({json.dumps(c)}, {json.dumps(NAMES[c])});" for c in sorted(NAMES)
)
TEMPLATE = '''// Code generated by gen_golden_test.py. DO NOT EDIT.
package org.wickra;
import org.junit.jupiter.api.DynamicTest;
import org.junit.jupiter.api.TestFactory;
import java.io.IOException;
import java.lang.reflect.Array;
import java.lang.reflect.Constructor;
import java.lang.reflect.Method;
import java.lang.reflect.RecordComponent;
import java.nio.file.Files;
import java.nio.file.Path;
import java.util.ArrayList;
import java.util.List;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.junit.jupiter.api.DynamicTest.dynamicTest;
/**
* Reflection-driven value parity for the whole 514-indicator catalogue: each
* indicator is reconstructed by its class name, fed the synthetic stream derived
* from the shared golden input (identical to gen_golden's Rust construction) and
* checked bit-for-bit against testdata/golden/g_&lt;Canonical&gt;.csv. One runner
* flattens scalar, multi-output records, profiles and bar arrays by reflection.
*/
class GoldenAllTest {
private static final double TOL = 1e-6;
private record Spec(String canonical, String arch, double[] params, int width) {}
private static final Spec[] SPECS = {
%SPECS%
};
private static final java.util.Map<String, String> NAMES = new java.util.HashMap<>();
static {
%NAMES%
}
private static Path goldenDir() {
java.io.File d = new java.io.File("").getAbsoluteFile();
while (d != null) {
java.io.File g = new java.io.File(d, "testdata/golden");
if (g.isDirectory()) {
return g.toPath();
}
d = d.getParentFile();
}
throw new IllegalStateException("testdata/golden not found");
}
private static double cell(String s) {
return switch (s) {
case "nan" -> Double.NaN;
case "inf" -> Double.POSITIVE_INFINITY;
case "-inf" -> Double.NEGATIVE_INFINITY;
default -> Double.parseDouble(s);
};
}
private static double[][] input() throws IOException {
List<String> lines = Files.readAllLines(goldenDir().resolve("input.csv"));
List<double[]> rows = new ArrayList<>();
for (int i = 1; i < lines.size(); i++) {
if (lines.get(i).isEmpty()) continue;
String[] p = lines.get(i).split(",");
double[] r = new double[p.length];
for (int j = 0; j < p.length; j++) r[j] = Double.parseDouble(p[j]);
rows.add(r);
}
return rows.toArray(new double[0][]);
}
// Keep blank lines (a candle on which no bar closed) so rows stay aligned.
private static double[][] fixture(String name) throws IOException {
List<String> lines = Files.readAllLines(goldenDir().resolve("g_" + name + ".csv"));
List<double[]> rows = new ArrayList<>();
for (int i = 1; i < lines.size(); i++) {
String ln = lines.get(i);
if (ln.isEmpty()) {
rows.add(new double[0]);
continue;
}
String[] p = ln.split(",");
double[] r = new double[p.length];
for (int j = 0; j < p.length; j++) r[j] = cell(p[j]);
rows.add(r);
}
return rows.toArray(new double[0][]);
}
private static double[] nanRow(int n) {
double[] r = new double[n];
java.util.Arrays.fill(r, Double.NaN);
return r;
}
private static double[] derivFields(double[] r) {
double o = r[0], h = r[1], l = r[2], c = r[3], v = r[4];
return new double[]{
(c - o) / c * 0.01, c, c - 0.5, c + 1.0, v * 10.0, v * 0.6, v * 0.4,
v * 0.55, v * 0.45, h - c, c - l,
};
}
private static double[] crossList(double[] r, int which) {
double o = r[0], c = r[3], v = r[4];
double[] out = new double[5];
for (int j = 0; j < 5; j++) {
out[j] = switch (which) {
case 0 -> (c - o) + j;
case 1 -> v + j * 10.0;
case 2 -> j % 2 == 0 ? 1.0 : 0.0; // newHigh
case 3 -> j % 3 == 0 ? 1.0 : 0.0; // newLow
case 4 -> j % 2 == 0 ? 1.0 : 0.0; // aboveMa
default -> j % 3 == 0 ? 1.0 : 0.0; // onBuySignal
};
}
return out;
}
private static double[] obList(double[] r, int which) {
double c = r[3], v = r[4];
double[] out = new double[5];
for (int k = 0; k < 5; k++) {
double kf = k + 1;
out[k] = switch (which) {
case 0 -> c - 0.1 * kf;
case 1 -> v / kf;
case 2 -> c + 0.1 * kf;
default -> v * 0.9 / kf;
};
}
return out;
}
private static double[] flattenRecord(Object o) throws Exception {
RecordComponent[] rc = o.getClass().getRecordComponents();
List<Double> list = new ArrayList<>();
for (RecordComponent c : rc) {
Object v = c.getAccessor().invoke(o);
if (v instanceof double[] arr) {
for (double d : arr) list.add(d);
} else if (v instanceof Number num) {
list.add(num.doubleValue());
}
}
double[] out = new double[list.size()];
for (int i = 0; i < out.length; i++) out[i] = list.get(i);
return out;
}
private static double[] flattenArray(Object arr) throws Exception {
int len = Array.getLength(arr);
List<Double> list = new ArrayList<>();
for (int i = 0; i < len; i++) {
for (double d : flattenRecord(Array.get(arr, i))) list.add(d);
}
double[] out = new double[list.size()];
for (int i = 0; i < out.length; i++) out[i] = list.get(i);
return out;
}
private static Object construct(Spec s) throws Exception {
Class<?> cls = Class.forName("org.wickra." + s.canonical());
Constructor<?> ctor = cls.getConstructors()[0];
Class<?>[] pt = ctor.getParameterTypes();
Object[] args = new Object[pt.length];
for (int i = 0; i < pt.length; i++) {
double v = s.params()[i];
if (pt[i] == int.class) args[i] = (int) Math.round(v);
else if (pt[i] == long.class) args[i] = (long) Math.round(v);
else args[i] = v;
}
return ctor.newInstance(args);
}
private static Method updateMethod(Object ind) {
for (Method m : ind.getClass().getMethods()) {
if (m.getName().equals("update")) return m;
}
throw new IllegalStateException("no update on " + ind.getClass());
}
private static double[] row(Spec s, Object ind, Method upd, double[] r, int i) throws Exception {
double o = r[0], h = r[1], l = r[2], c = r[3], v = r[4];
Object res = switch (s.arch()) {
case "scalar_f64", "multi_f64" -> upd.invoke(ind, c);
case "pairwise", "multi_pairwise" -> upd.invoke(ind, c, o);
case "scalar_candle", "multi_candle", "profile_bins", "profile_pricebins" ->
upd.invoke(ind, o, h, l, c, v, (long) i);
case "trade" -> upd.invoke(ind, c, v, c >= o, (long) i);
case "trademid" -> upd.invoke(ind, c, v, c >= o, (long) i, (h + l) / 2);
case "ob" -> upd.invoke(ind, obList(r, 0), obList(r, 1), obList(r, 2), obList(r, 3));
case "cross" -> upd.invoke(ind, crossList(r, 0), crossList(r, 1), crossList(r, 2),
crossList(r, 3), crossList(r, 4), crossList(r, 5), (long) i);
case "deriv", "deriv_multi" -> {
double[] d = derivFields(r);
yield upd.invoke(ind, d[0], d[1], d[2], d[3], d[4], d[5], d[6], d[7], d[8], d[9], d[10], (long) i);
}
case "bars_close" -> upd.invoke(ind, c, c, c, c, 1.0, 0L);
case "bars_candle4" -> upd.invoke(ind, o, h, l, c, 1.0, 0L);
case "bars_candle5" -> upd.invoke(ind, o, h, l, c, v, 0L);
case "footprint" -> upd.invoke(ind, c, v, c >= o, (long) i);
default -> throw new IllegalStateException("arch " + s.arch());
};
return switch (s.arch()) {
case "scalar_f64", "scalar_candle", "pairwise", "trade", "trademid", "ob", "cross", "deriv" ->
new double[]{((Number) res).doubleValue()};
case "multi_f64", "multi_candle", "multi_pairwise", "deriv_multi", "profile_pricebins" ->
res == null ? nanRow(s.width()) : flattenRecord(res);
case "profile_bins" -> res == null ? nanRow(s.width()) : (double[]) res;
default -> flattenArray(res); // bars_*, footprint
};
}
@TestFactory
List<DynamicTest> golden() throws Exception {
double[][] rows = input();
List<DynamicTest> tests = new ArrayList<>();
for (Spec s : SPECS) {
tests.add(dynamicTest(s.canonical(), () -> {
Object ind = construct(s);
String gotName = (String) ind.getClass().getMethod("name").invoke(ind);
assertTrue(gotName.equals(NAMES.get(s.canonical())),
s.canonical() + ": name() " + gotName + " != " + NAMES.get(s.canonical()));
Method upd = updateMethod(ind);
double[][] exp = fixture(s.canonical());
assertTrue(exp.length == rows.length, s.canonical() + ": row count " + exp.length + " vs " + rows.length);
for (int i = 0; i < rows.length; i++) {
double[] got = row(s, ind, upd, rows[i], i);
double[] want = exp[i];
assertTrue(got.length == want.length,
s.canonical() + " row " + i + ": arity " + got.length + " vs " + want.length);
for (int k = 0; k < want.length; k++) {
double w = want[k], g = got[k];
if (Double.isNaN(w)) {
assertTrue(Double.isNaN(g), s.canonical() + " row " + i + " col " + k + ": want NaN got " + g);
} else if (Double.isInfinite(w)) {
assertTrue(Double.isInfinite(g) && Math.signum(g) == Math.signum(w),
s.canonical() + " row " + i + " col " + k + ": want " + w + " got " + g);
} else {
double tol = TOL * Math.max(1.0, Math.abs(w));
assertTrue(Math.abs(g - w) <= tol,
s.canonical() + " row " + i + " col " + k + ": got " + g + " want " + w);
}
}
}
}));
}
return tests;
}
}
'''
out = TEMPLATE.replace("%SPECS%", specs_block).replace("%NAMES%", names_block)
dest = os.path.join(ROOT, "bindings", "java", "src", "test", "java", "org", "wickra", "GoldenAllTest.java")
open(dest, "w", encoding="utf-8").write(out)
print("generated GoldenAllTest.java with", len(MAN), "indicators")