#!/usr/bin/env bash
set -euo pipefail

# bench-regression.sh — Hard merge gate for model changes
#
# Runs candidate vs champion on real-v1-test (3x alternating),
# writes comparison.json, exits non-zero on failure.
#
# Usage:
#   BK_FINETUNED_MODEL=./models/bge-small-finetuned-onnx scripts/bench-regression.sh
#
# Requires: jq, bun, benchmarks/search/baselines/champion-mean-v1-test.json

SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
PROJECT_ROOT="$(cd "$SCRIPT_DIR/.." && pwd)"
cd "$PROJECT_ROOT"

CHAMPION_BASELINE="benchmarks/search/baselines/champion-mean-v1-test.json"
SPLIT_FILE="benchmarks/search/splits/test-cases.json"
TIMESTAMP="$(date -u +%Y%m%dT%H%M%S)"
RESULTS_DIR="benchmarks/search/regression-results/${TIMESTAMP}"

# Validate prerequisites
if [ ! -f "$CHAMPION_BASELINE" ]; then
  echo "ERROR: Champion baseline not found: $CHAMPION_BASELINE"
  echo "Run Step 8A first."
  exit 1
fi

if [ ! -f "$SPLIT_FILE" ]; then
  echo "ERROR: Test split not found: $SPLIT_FILE"
  echo "Run split-dataset.ts first."
  exit 1
fi

mkdir -p "$RESULTS_DIR"

echo "============================================================"
echo "REGRESSION GATE — bench-regression.sh"
echo "============================================================"
echo "Timestamp: $TIMESTAMP"
echo "Results:   $RESULTS_DIR"
echo "Champion:  $CHAMPION_BASELINE"
echo ""

# --- Run 3 alternating champion/candidate pairs ---

run_champion() {
  local run_num=$1
  local artifact="$RESULTS_DIR/champion-run-${run_num}.json"
  echo "[Run ${run_num}/3] Champion (baseline)..."
  (
    unset BK_FINETUNED_MODEL BK_QUERY_PREFIX
    bun run bench:search --dataset real-v1-test --reuse-index \
      --artifacts "$artifact" 2>&1 | grep -E "Hit@1|SUMMARY|Artifact"
  )
  echo "  -> $artifact"
}

run_candidate() {
  local run_num=$1
  local artifact="$RESULTS_DIR/candidate-run-${run_num}.json"
  echo "[Run ${run_num}/3] Candidate (finetuned)..."
  # Finetuned BGE models still need the BGE query prefix
  BK_QUERY_PREFIX='Represent this sentence for searching relevant passages: ' \
    bun run bench:search --dataset real-v1-test --setup --force \
    --artifacts "$artifact" 2>&1 | grep -E "Hit@1|SUMMARY|Artifact"
  echo "  -> $artifact"
}

# Alternating: C1 -> X1 -> C2 -> X2 -> C3 -> X3
for i in 1 2 3; do
  run_champion "$i"
  run_candidate "$i"
done

echo ""
echo "All 6 runs complete. Analyzing..."
echo ""

# --- Fingerprint parity check ---

echo "Fingerprint parity check..."
CHAMPION_FP=$(jq -S '.configFingerprint | del(.embedding.model, .embedding.modelSha256)' "$RESULTS_DIR/champion-run-1.json")
CANDIDATE_FP=$(jq -S '.configFingerprint | del(.embedding.model, .embedding.modelSha256)' "$RESULTS_DIR/candidate-run-1.json")

if [ "$CHAMPION_FP" != "$CANDIDATE_FP" ]; then
  echo "FAIL: Fingerprint mismatch (non-model fields differ)"
  diff <(echo "$CHAMPION_FP") <(echo "$CANDIDATE_FP") || true
  echo ""
  echo "Only embedding.model and embedding.modelSha256 may differ."
  # Write failed comparison
  jq -n \
    --arg splitHash "$(sha256sum "$SPLIT_FILE" | cut -d' ' -f1)" \
    --arg gitSha "$(git rev-parse --short HEAD)" \
    --arg verdict "FAIL: fingerprint mismatch" \
    '{splitHash: $splitHash, gitSha: $gitSha, verdict: $verdict}' \
    > "$RESULTS_DIR/comparison.json"
  exit 1
fi
echo "  PASS: all non-model fields identical"

# --- Compute integer hits ---
# Weighted integer hits = sum of weight for cases where hitAt1 == true
# Matches aggregation in metrics.ts:216

compute_hits() {
  local artifact=$1
  jq '[.results[] | select(.metrics.hitAt1 == true) | .weight // 1] | add // 0' "$artifact"
}

echo ""
echo "Integer hit counts (weighted):"

CHAMPION_HITS=()
CANDIDATE_HITS=()
for i in 1 2 3; do
  ch=$(compute_hits "$RESULTS_DIR/champion-run-${i}.json")
  cx=$(compute_hits "$RESULTS_DIR/candidate-run-${i}.json")
  CHAMPION_HITS+=("$ch")
  CANDIDATE_HITS+=("$cx")
  echo "  Pair $i: champion=$ch, candidate=$cx (delta=$((cx - ch)))"
done

# --- Stability check ---
echo ""
echo "Stability check..."
if [ "${CHAMPION_HITS[0]}" != "${CHAMPION_HITS[1]}" ] || [ "${CHAMPION_HITS[0]}" != "${CHAMPION_HITS[2]}" ]; then
  echo "  WARNING: Champion hits vary across runs: ${CHAMPION_HITS[*]}"
fi
if [ "${CANDIDATE_HITS[0]}" != "${CANDIDATE_HITS[1]}" ] || [ "${CANDIDATE_HITS[0]}" != "${CANDIDATE_HITS[2]}" ]; then
  echo "  WARNING: Candidate hits vary across runs: ${CANDIDATE_HITS[*]}"
fi
STABLE="true"
if [ "${CANDIDATE_HITS[0]}" != "${CANDIDATE_HITS[1]}" ] || [ "${CANDIDATE_HITS[0]}" != "${CANDIDATE_HITS[2]}" ]; then
  STABLE="false"
fi
echo "  Candidate stable: $STABLE"

# --- Per-query win/loss (using run-1 as reference) ---

echo ""
echo "Per-query analysis (run 1)..."

# Extract hitAt1 per case ID for champion and candidate
CHAMPION_CASES=$(jq -r '[.results[] | {id, hit: .metrics.hitAt1, weight: (.weight // 1)}]' "$RESULTS_DIR/champion-run-1.json")
CANDIDATE_CASES=$(jq -r '[.results[] | {id, hit: .metrics.hitAt1, weight: (.weight // 1)}]' "$RESULTS_DIR/candidate-run-1.json")

WINS=0
LOSSES=0
TIES=0

# Compare per case
CASE_IDS=$(jq -r '.[].id' <<< "$CHAMPION_CASES")
while IFS= read -r case_id; do
  ch_hit=$(jq -r --arg id "$case_id" '.[] | select(.id == $id) | .hit' <<< "$CHAMPION_CASES")
  cx_hit=$(jq -r --arg id "$case_id" '.[] | select(.id == $id) | .hit' <<< "$CANDIDATE_CASES")
  if [ "$ch_hit" = "false" ] && [ "$cx_hit" = "true" ]; then
    ((WINS++)) || true
  elif [ "$ch_hit" = "true" ] && [ "$cx_hit" = "false" ]; then
    ((LOSSES++)) || true
  else
    ((TIES++)) || true
  fi
done <<< "$CASE_IDS"

echo "  Wins: $WINS, Losses: $LOSSES, Ties: $TIES"

# --- Category analysis ---

echo ""
echo "Category analysis (run 1)..."

# Per-category integer hit deltas
CATEGORY_DELTAS=$(jq -n \
  --argjson champion "$CHAMPION_CASES" \
  --argjson candidate "$CANDIDATE_CASES" \
  '[($champion | group_by(.id | split("-")[0:2] | join("-")) | .[] |
    {category: (.[0].id | split("-")[0:2] | join("-")),
     champion_hits: [.[] | select(.hit) | .weight // 1] | add // 0}) ] as $ch_cats |
   [($candidate | group_by(.id | split("-")[0:2] | join("-")) | .[] |
    {category: (.[0].id | split("-")[0:2] | join("-")),
     candidate_hits: [.[] | select(.hit) | .weight // 1] | add // 0}) ] as $cx_cats |
   [$ch_cats[] as $c | {category: $c.category,
     champion: $c.champion_hits,
     candidate: ([$cx_cats[] | select(.category == $c.category) | .candidate_hits][0] // 0),
     delta: (([$cx_cats[] | select(.category == $c.category) | .candidate_hits][0] // 0) - $c.champion_hits)}]')

echo "$CATEGORY_DELTAS" | jq -r '.[] | "  \(.category): champion=\(.champion) candidate=\(.candidate) delta=\(.delta)"'

# Check category guardrail: no category loses more than 1 weighted hit
CATEGORY_FAIL=$(echo "$CATEGORY_DELTAS" | jq '[.[] | select(.delta < -1)] | length')

# --- P95 latency ---

echo ""
echo "P95 latency:"
P95_DELTAS=()
for i in 1 2 3; do
  ch_p95=$(jq '.summary.latency.p95' "$RESULTS_DIR/champion-run-${i}.json")
  cx_p95=$(jq '.summary.latency.p95' "$RESULTS_DIR/candidate-run-${i}.json")
  delta=$(echo "$cx_p95 - $ch_p95" | bc -l 2>/dev/null || echo "0")
  P95_DELTAS+=("$delta")
  printf "  Pair %d: champion=%.1fms candidate=%.1fms delta=%.1fms\n" "$i" "$ch_p95" "$cx_p95" "$delta"
done

# Median P95 check (baseline + 20ms)
CHAMPION_P95_MEDIAN=$(for i in 1 2 3; do jq '.summary.latency.p95' "$RESULTS_DIR/champion-run-${i}.json"; done | sort -n | sed -n '2p')
CANDIDATE_P95_MEDIAN=$(for i in 1 2 3; do jq '.summary.latency.p95' "$RESULTS_DIR/candidate-run-${i}.json"; done | sort -n | sed -n '2p')
LATENCY_LIMIT=$(echo "$CHAMPION_P95_MEDIAN + 20" | bc -l)
LATENCY_OK=$(echo "$CANDIDATE_P95_MEDIAN <= $LATENCY_LIMIT" | bc -l)

printf "  Median P95: champion=%.1fms candidate=%.1fms limit=%.1fms\n" "$CHAMPION_P95_MEDIAN" "$CANDIDATE_P95_MEDIAN" "$LATENCY_LIMIT"

# --- Gate decisions ---

echo ""
echo "============================================================"
echo "GATE EVALUATION"
echo "============================================================"

VERDICT="PASS"
REASONS=()

# Gate 1: Integer-hit gate (raised bar)
HIT_DELTA=$((CANDIDATE_HITS[0] - CHAMPION_HITS[0]))
echo "Gate 1 — Integer-hit gate: delta=$HIT_DELTA weighted hits"

GATE1_PASS="false"
if [ "$HIT_DELTA" -ge 2 ]; then
  echo "  PASS (path a): gain >= 2 weighted hits"
  GATE1_PASS="true"
elif [ "$HIT_DELTA" -ge 1 ] && [ "$LOSSES" -eq 0 ]; then
  echo "  PASS (path b): gain >= 1 AND zero losses"
  GATE1_PASS="true"
else
  echo "  FAIL: delta=$HIT_DELTA, losses=$LOSSES (need >=2 hits, or >=1 with 0 losses)"
  VERDICT="FAIL"
  REASONS+=("integer-hit gate: delta=$HIT_DELTA losses=$LOSSES")
fi

# Gate 2: Category guardrail
echo "Gate 2 — Category guardrail: $CATEGORY_FAIL categories with >1 hit loss"
if [ "$CATEGORY_FAIL" -gt 0 ]; then
  echo "  FAIL: category regression detected"
  VERDICT="FAIL"
  REASONS+=("category guardrail: $CATEGORY_FAIL categories regressed >1 hit")
else
  echo "  PASS"
fi

# Gate 3: Per-query win/loss
echo "Gate 3 — Per-query: wins=$WINS losses=$LOSSES"
if [ "$WINS" -gt "$LOSSES" ]; then
  echo "  PASS"
else
  echo "  FAIL: wins must exceed losses"
  VERDICT="FAIL"
  REASONS+=("per-query: wins=$WINS <= losses=$LOSSES")
fi

# Gate 4: Latency
echo "Gate 4 — Latency: candidate median P95 within baseline + 20ms"
if [ "$LATENCY_OK" -eq 1 ]; then
  echo "  PASS"
else
  echo "  FAIL: candidate P95 exceeds limit"
  VERDICT="FAIL"
  REASONS+=("latency: candidate=$CANDIDATE_P95_MEDIAN > limit=$LATENCY_LIMIT")
fi

# Gate 5: Stability
echo "Gate 5 — Stability: candidate hits identical across 3 runs"
if [ "$STABLE" = "true" ]; then
  echo "  PASS"
else
  echo "  FAIL: candidate hit count varies"
  VERDICT="FAIL"
  REASONS+=("stability: hits vary ${CANDIDATE_HITS[*]}")
fi

# --- Build comparison.json ---

SPLIT_HASH=$(sha256sum "$SPLIT_FILE" 2>/dev/null | cut -d' ' -f1 || shasum -a 256 "$SPLIT_FILE" | cut -d' ' -f1)
GIT_SHA=$(git rev-parse --short HEAD)

# Model SHA (try ONNX model if BK_FINETUNED_MODEL is set)
MODEL_SHA="base-model"
if [ -n "${BK_FINETUNED_MODEL:-}" ] && [ -f "${BK_FINETUNED_MODEL}/onnx/model.onnx" ]; then
  MODEL_SHA=$(sha256sum "${BK_FINETUNED_MODEL}/onnx/model.onnx" 2>/dev/null | cut -d' ' -f1 || shasum -a 256 "${BK_FINETUNED_MODEL}/onnx/model.onnx" | cut -d' ' -f1)
fi

# Training manifest (if exists)
TRAINING_MANIFEST="{}"
if [ -f "training/data/training-manifest.json" ]; then
  TRAINING_MANIFEST=$(cat "training/data/training-manifest.json")
fi

REASON_STR=""
if [ ${#REASONS[@]} -gt 0 ]; then
  REASON_STR=$(printf '%s; ' "${REASONS[@]}")
fi

jq -n \
  --arg splitHash "$SPLIT_HASH" \
  --arg gitSha "$GIT_SHA" \
  --arg modelSha "$MODEL_SHA" \
  --argjson championFP "$(jq '.configFingerprint' "$RESULTS_DIR/champion-run-1.json")" \
  --argjson candidateFP "$(jq '.configFingerprint' "$RESULTS_DIR/candidate-run-1.json")" \
  --argjson championHits "[${CHAMPION_HITS[0]},${CHAMPION_HITS[1]},${CHAMPION_HITS[2]}]" \
  --argjson candidateHits "[${CANDIDATE_HITS[0]},${CANDIDATE_HITS[1]},${CANDIDATE_HITS[2]}]" \
  --argjson categoryDeltas "$CATEGORY_DELTAS" \
  --argjson perQueryWinLoss "{\"wins\":$WINS,\"losses\":$LOSSES,\"ties\":$TIES}" \
  --argjson trainingManifest "$TRAINING_MANIFEST" \
  --arg verdict "$VERDICT" \
  --arg reason "$REASON_STR" \
  '{
    splitHash: $splitHash,
    gitSha: $gitSha,
    modelSha: $modelSha,
    championFingerprint: $championFP,
    candidateFingerprint: $candidateFP,
    integerHits: {champion: $championHits, candidate: $candidateHits},
    categoryDeltas: $categoryDeltas,
    perQueryWinLoss: $perQueryWinLoss,
    p95Deltas: [],
    trainingManifest: $trainingManifest,
    verdict: $verdict,
    reason: (if $reason == "" then null else $reason end)
  }' > "$RESULTS_DIR/comparison.json"

echo ""
echo "============================================================"
echo "VERDICT: $VERDICT"
if [ -n "$REASON_STR" ]; then
  echo "Reason: $REASON_STR"
fi
echo "Comparison: $RESULTS_DIR/comparison.json"
echo "============================================================"

if [ "$VERDICT" = "PASS" ]; then
  exit 0
else
  exit 1
fi
