"""LLM-free stand-in for a 'decision model' routing task: given a snippet of text, pick one
of 5 doc sections (a Choice-shaped decision, the same shape TypeSafe's Jev returns for its
Choice primitive). TF-IDF + multinomial logistic regression, trained on the FastAPI docs corpus.

No network calls: no OpenRouter/TypeSafe key is set in this environment, so this measures the
classic-ML baseline only, not Jev itself. See the post for what Jev's own vendor numbers say.

Split is by FILE, not by snippet, so no snippet from a test file is seen in training.
"""
import glob, os, re, time, statistics as st, collections
import numpy as np
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score, f1_score, confusion_matrix

CATS = {"tutorial", "advanced", "reference", "how-to", "deployment"}

def label_of(path):
    name = os.path.basename(path)
    prefix = name.split("_")[0].removesuffix(".md")
    return prefix if prefix in CATS else None

def paragraphs(text, min_len=200, max_len=800):
    """Split on blank lines, merge short paragraphs up to max_len, drop anything under min_len."""
    parts = [p.strip() for p in re.split(r"\n\s*\n", text) if p.strip()]
    out, cur = [], ""
    for p in parts:
        cur = f"{cur}\n\n{p}" if cur else p
        if len(cur) >= min_len:
            out.append(cur[:max_len])
            cur = ""
    if len(cur) >= min_len:
        out.append(cur[:max_len])
    return out

files = sorted(glob.glob("../corpus/*.md"))
by_cat = collections.defaultdict(list)
for f in files:
    lab = label_of(f)
    if lab:
        by_cat[lab].append(f)

print("files per category:", {k: len(v) for k, v in sorted(by_cat.items())})

train_files, test_files = [], []
for cat, fs in by_cat.items():
    tr, te = train_test_split(sorted(fs), test_size=0.25, random_state=0)
    train_files += [(f, cat) for f in tr]
    test_files += [(f, cat) for f in te]
print(f"train files: {len(train_files)}, test files: {len(test_files)}")

def snippets_for(file_label_list):
    X, y, file_of = [], [], []
    for f, cat in file_label_list:
        text = open(f, encoding="utf-8").read()
        for s in paragraphs(text):
            X.append(s); y.append(cat); file_of.append(f)
    return X, y, file_of

X_train, y_train, _ = snippets_for(train_files)
X_test, y_test, file_of_test = snippets_for(test_files)
print(f"train snippets: {len(X_train)}, test snippets: {len(X_test)}")
print("train snippets per category:", dict(collections.Counter(y_train)))
print("test snippets per category:", dict(collections.Counter(y_test)))

vec = TfidfVectorizer(max_features=5000, ngram_range=(1, 1), stop_words="english", min_df=2)
Xtr = vec.fit_transform(X_train)
Xte = vec.transform(X_test)

clf = LogisticRegression(max_iter=2000, C=5.0, class_weight="balanced")
clf.fit(Xtr, y_train)

pred = clf.predict(Xte)
proba = clf.predict_proba(Xte)
classes = list(clf.classes_)
conf = proba.max(axis=1)

acc = accuracy_score(y_test, pred)
macro_f1 = f1_score(y_test, pred, average="macro")
print(f"\naccuracy={acc:.3f} macro_f1={macro_f1:.3f}  (baseline if always predicting majority class 'tutorial': "
      f"{y_test.count('tutorial')/len(y_test):.3f})")

print("\nconfusion matrix (rows=true, cols=predicted), classes:", classes)
cm = confusion_matrix(y_test, pred, labels=classes)
print("        " + " ".join(f"{c[:6]:>7}" for c in classes))
for c, row in zip(classes, cm):
    print(f"{c[:7]:>8}" + " ".join(f"{v:>7}" for v in row))

# --- calibration: does predicted confidence match observed accuracy? ---
print("\ncalibration (predicted max-probability bucket vs observed accuracy):")
bins = [0.2, 0.4, 0.6, 0.8, 1.01]
lo = 0.0
correct = (pred == np.array(y_test))
for hi in bins:
    mask = (conf >= lo) & (conf < hi)
    n = mask.sum()
    if n:
        print(f"  confidence [{lo:.1f},{hi:.1f}): n={n:4d} mean_confidence={conf[mask].mean():.3f} "
              f"observed_accuracy={correct[mask].mean():.3f}")
    lo = hi

# --- latency: single-snippet transform+predict, this machine, this corpus size ---
sample = X_test[:300] if len(X_test) >= 300 else X_test
t = []
for s in sample:
    t0 = time.perf_counter()
    v = vec.transform([s])
    clf.predict_proba(v)
    t.append(time.perf_counter() - t0)
print(f"\nmedian single-snippet latency (tfidf transform + logreg predict_proba): {st.median(t)*1000:.3f}ms "
      f"over {len(sample)} calls, min={min(t)*1000:.3f}ms max={max(t)*1000:.3f}ms")

# --- cost illustration: apply Jev's *published* per-token rate to *our own* token count ---
# Word count is an approximation of tokens (no tokenizer downloaded); Jev's real token count would differ.
words = sum(len(s.split()) for s in X_test)
approx_tokens = words * 1.3  # rough words->tokens fudge factor, English text
jev_rate_per_m = 0.042  # USD per 1M input tokens, https://openrouter.ai/typesafe/jev-1.13, checked 2026-09-27
jev_cost = approx_tokens / 1e6 * jev_rate_per_m
print(f"\nfor reference only (not measured, arithmetic on published price): routing these {len(X_test)} test "
      f"snippets through Jev at its published input rate would cost about ${jev_cost:.4f} for ~{approx_tokens:,.0f} "
      f"approx tokens (word-count-based estimate, not a real tokenizer). This local classifier: $0, no network call.")
