// Tests for `fluxrec train`: reproducibility, metrics, and model envelope // consumption by the `score` command. package cmds import ( "bytes" "encoding/json" "os" "path/filepath" "strings" "testing" "time" "fluxrec/core" ) // syntheticLabels builds a labels.jsonl file interleaving positives and // negatives across time (so the newest validation window contains both // classes), with varied feed titles so the feed-identity token is exercised. // Timestamps are sequential by position: oldest first, newest last, so the // time split's "newest val-frac" is a clean tail slice with mixed labels. func syntheticLabels(t *testing.T, path string, nPos, nNeg int) { t.Helper() now := time.Now().UTC() total := nPos + nNeg rows := make([]core.LabelRow, total) posTitles := []string{ "New evidence on the rate of productivity growth in services", "Why headline-only models are fine actually, a reanalysis", "Field notes from sixty days of careful RSS triage at home", "The economics of attention and the case for reading slower", "A long retrospective on building small personal recommendation tools", "Decade-scale trends in academic publishing, a meta-review", "How I learned to stop worrying and love the title-only classifier", "Notes on serendipity, exploration budgets, and ranked recall", "Reflections on a year of minimal-information inference methods", "The quiet resurgence of logistic regression in small data settings", } negTitles := []string{ "Roundup of this week's product launches and minor updates", "Local weather forecast for the coming weekend and beyond", "Five ways to organize your desk for slightly better focus", "Celebrity interview about their new streaming television show", "Sports scores and highlights from last night's important games", "Markets close mostly unchanged as investors await fresh data", "Travel deals for the autumn season across several destinations", "Cooking tips for weeknight dinners under thirty minutes flat", "Opinion piece on city traffic and the eternal parking problem", "Best of the app store this month, ten small utilities reviewed", } feeds := []string{"Marginal Revolution", "ACM Tech News", "Some Other Blog", "Hacker News", "The Atlantic"} pi, ni := 0, 0 for i := 0; i < total; i++ { // Place a positive on even indices while positives remain, so // labels interleave across the whole time range including the tail. label := core.LabelNeg title := "" if i%2 == 0 && pi < nPos { label, title = core.LabelPos, posTitles[pi%len(posTitles)] pi++ } else { label, title = core.LabelNeg, negTitles[ni%len(negTitles)] ni++ } rows[i] = core.LabelRow{ EntryID: int64(i + 1), URL: "https://example.com/" + label + "/" + itoa(i), Title: title, FeedID: int64(i%len(feeds) + 1), FeedTitle: feeds[i%len(feeds)], Label: label, CapturedAt: now.Add(-time.Duration(total-i) * 24 * time.Hour), } } if err := core.AppendLabels(path, rows); err != nil { t.Fatalf("AppendLabels: %v", err) } } func itoa(i int) string { b := []byte{} if i == 0 { return "0" } for i > 0 { b = append([]byte{byte('0' + i%10)}, b...) i /= 10 } return string(b) } func runTrain(t *testing.T, labelsPath, dir string) (modelPath, reportPath string) { t.Helper() modelPath = filepath.Join(dir, "model.json") reportPath = filepath.Join(dir, "report.json") cmd := &TrainCommand{} if err := cmd.Init([]string{labelsPath, "--model-out", modelPath, "--report-out", reportPath, "--val-frac", "0.3"}); err != nil { t.Fatalf("Init: %v", err) } var out bytes.Buffer if err := cmd.Run(nil, &out); err != nil { t.Fatalf("Run: %v\nstdout: %s", err, out.String()) } return } // B4: same labels → byte-identical model.json (seeds make training deterministic). func TestTrainReproducible(t *testing.T) { dir := t.TempDir() labelsPath := filepath.Join(dir, "labels.jsonl") // Enough rows that a 0.3 val split has >=5 positives in val. syntheticLabels(t, labelsPath, 50, 50) m1, r1 := runTrain(t, labelsPath, dir) body1, err := os.ReadFile(m1) if err != nil { t.Fatalf("read model1: %v", err) } report1, err := os.ReadFile(r1) if err != nil { t.Fatalf("read report1: %v", err) } // Second train to a second pair of paths. dir2 := t.TempDir() m2, r2 := runTrain(t, labelsPath, dir2) body2, err := os.ReadFile(m2) if err != nil { t.Fatalf("read model2: %v", err) } report2, err := os.ReadFile(r2) if err != nil { t.Fatalf("read report2: %v", err) } // Strip CreatedAt (wall clock) before comparing models — the rest must // be byte-identical. Compare reports fully. var env1, env2 core.ModelEnvelope if err := json.Unmarshal(body1, &env1); err != nil { t.Fatal(err) } if err := json.Unmarshal(body2, &env2); err != nil { t.Fatal(err) } env1.CreatedAt = time.Time{} env2.CreatedAt = time.Time{} c1, _ := json.Marshal(env1) c2, _ := json.Marshal(env2) if string(c1) != string(c2) { t.Fatal("model not reproducible: weights/vocab differ across runs with same labels") } if string(report1) != string(report2) { t.Fatal("report not reproducible: differs across runs with same labels") } } // B5: report.json parses and precision@15 is a valid probability. func TestTrainReportValid(t *testing.T) { dir := t.TempDir() labelsPath := filepath.Join(dir, "labels.jsonl") syntheticLabels(t, labelsPath, 50, 50) modelPath, reportPath := runTrain(t, labelsPath, dir) var report trainReport body, err := os.ReadFile(reportPath) if err != nil { t.Fatalf("read report: %v", err) } if err := json.Unmarshal(body, &report); err != nil { t.Fatalf("report unparseable: %v", err) } if report.PrecisionAt15 < 0 || report.PrecisionAt15 > 1 { t.Errorf("precision@15 = %v, want in [0,1]", report.PrecisionAt15) } if report.VocabSize <= 0 { t.Errorf("vocab size = %d, want > 0", report.VocabSize) } if report.TrainPos == 0 || report.TrainNeg == 0 { t.Errorf("train should have both classes: pos=%d neg=%d", report.TrainPos, report.TrainNeg) } if report.ValPos < 5 { t.Errorf("val pos = %d, guard should have required >=5", report.ValPos) } if report.Seed != trainSeed { t.Errorf("seed = %d, want %d", report.Seed, trainSeed) } // Model file exists and has non-empty vocab + weights. mbody, err := os.ReadFile(modelPath) if err != nil { t.Fatalf("read model: %v", err) } var model core.ModelEnvelope if err := json.Unmarshal(mbody, &model); err != nil { t.Fatalf("model unparseable: %v", err) } if len(model.Weights) == 0 || len(model.OrderedVocab) == 0 { t.Fatal("model missing weights or vocab") } } // B6: the produced model.json is consumable by the score command's loader // (same ModelEnvelope type). Round-trips through CreateVectorizerFromModel. func TestTrainModelConsumedByScore(t *testing.T) { dir := t.TempDir() labelsPath := filepath.Join(dir, "labels.jsonl") syntheticLabels(t, labelsPath, 50, 50) modelPath, _ := runTrain(t, labelsPath, dir) model, err := loadModel(modelPath) // same loader the score command uses if err != nil { t.Fatalf("loadModel (score's loader): %v", err) } vec := core.CreateVectorizerFromModel(model) title := "A long retrospective on building small personal recommendation tools" out := vec.Transform([]string{core.ScoringText(title, "Marginal Revolution")}) score, err := core.PredictScore(out[0], model.Weights) if err != nil { t.Fatalf("PredictScore: %v", err) } if score < 0 || score > 1 { t.Errorf("score = %v, want in [0,1]", score) } } // Guard: too few validation positives must error loudly. func TestTrainGuardFewValPositives(t *testing.T) { dir := t.TempDir() labelsPath := filepath.Join(dir, "labels.jsonl") // 5 positives total, val-frac 0.3 → 1-2 val positives < 5. syntheticLabels(t, labelsPath, 5, 20) cmd := &TrainCommand{} if err := cmd.Init([]string{labelsPath, "--val-frac", "0.3", "--model-out", filepath.Join(dir, "m.json"), "--report-out", filepath.Join(dir, "r.json")}); err != nil { t.Fatalf("Init: %v", err) } var out bytes.Buffer err := cmd.Run(nil, &out) if err == nil { t.Fatal("expected error for <5 validation positives, got nil") } if !strings.Contains(err.Error(), "positives in validation") { t.Errorf("error should mention validation positives: %v", err) } } // Guard: empty labels file errors. func TestTrainEmptyLabels(t *testing.T) { dir := t.TempDir() labelsPath := filepath.Join(dir, "empty.jsonl") if err := os.WriteFile(labelsPath, []byte{}, 0o644); err != nil { t.Fatal(err) } cmd := &TrainCommand{} if err := cmd.Init([]string{labelsPath}); err != nil { t.Fatalf("Init: %v", err) } if err := cmd.Run(nil, &bytes.Buffer{}); err == nil { t.Fatal("expected error for empty labels") } }