#!/usr/bin/env python3
"""
Aria Cognitive-Health-Probe (KAR-685)
=====================================
Routine "cognitive health check" wie vom LLM-Brain-Rot-Paper empfohlen
(Xing et al. 2025, arXiv:2510.13928): misst REAL die zwei Faehigkeiten, die
das Paper als am staerksten degradierend identifiziert hat:

  1. Reasoning      (ARC-Challenge-Stil, Multiple-Choice mit Chain-of-Thought)
  2. Long-Context   (RULER-CWE-Stil, Needle-Retrieval aus langem Kontext)

Die Probe ruft das Modell WIRKLICH auf und scort echte Antworten — es werden
KEINE Zahlen geschaetzt. Ergebnis (model, scores, timestamp) wird nach
state/cognitive-health/probe-log.jsonl angehaengt, damit Drift ueber Zeit
sichtbar wird. Designed als Baustein fuer den Eval-Harness (KAR-680/681).

Gated auf ANTHROPIC_API_KEY (fehlt er -> sauberer Exit, kein Crash).

Run:
  python3 aria-cognitive-health-probe.py [--model MODEL] [--needles N] [--filler-chars N] [--json]
Default-Modell: $ARIA_HEALTH_PROBE_MODEL oder triage_model (Haiku, guenstig).
"""
from __future__ import annotations
import argparse
import json
import os
import sys
from datetime import datetime, timezone
from pathlib import Path

CONFIG_PATH = Path("/root/aria/brain/youtube/.config.yaml")
LOG_DIR = Path("/root/aria/state/cognitive-health")
LOG_PATH = LOG_DIR / "probe-log.jsonl"

# --- Reasoning-Probe: ARC-Challenge-Stil, handverifizierte Antworten ---------
# Bewusst Reasoning (Deduktion/Arithmetik/Kausalitaet), nicht Trivia, damit
# "Thought-Skipping" (Paper-Failure-Mode) als Fehler sichtbar wird.
REASONING_ITEMS = [
    {"id": "r1", "q": "Alle Bloops sind Razzies. Alle Razzies sind Lazzies. Welche Aussage folgt zwingend?",
     "choices": {"A": "Alle Lazzies sind Bloops", "B": "Alle Bloops sind Lazzies",
                 "C": "Manche Razzies sind keine Lazzies", "D": "Keine Bloops sind Lazzies"}, "answer": "B"},
    {"id": "r2", "q": "Ein Zug faehrt 60 km in 45 Minuten. Wie weit kommt er bei gleicher Geschwindigkeit in 1 Stunde?",
     "choices": {"A": "75 km", "B": "80 km", "C": "90 km", "D": "60 km"}, "answer": "B"},
    {"id": "r3", "q": "Wenn es regnet, ist die Strasse nass. Die Strasse ist nass. Was folgt logisch zwingend?",
     "choices": {"A": "Es regnet", "B": "Es regnet nicht",
                 "C": "Nichts davon folgt zwingend", "D": "Es wird regnen"}, "answer": "C"},
    {"id": "r4", "q": "Anna ist aelter als Bea. Bea ist aelter als Cara. Cara ist aelter als Dora. Wer ist am juengsten?",
     "choices": {"A": "Anna", "B": "Bea", "C": "Cara", "D": "Dora"}, "answer": "D"},
    {"id": "r5", "q": "Ein Hemd kostet nach 20% Rabatt 40 Euro. Was war der urspruengliche Preis?",
     "choices": {"A": "48 Euro", "B": "50 Euro", "C": "52 Euro", "D": "60 Euro"}, "answer": "B"},
    {"id": "r6", "q": "Folge: 2, 6, 12, 20, 30, ? Welche Zahl kommt als naechstes?",
     "choices": {"A": "36", "B": "40", "C": "42", "D": "44"}, "answer": "C"},
    {"id": "r7", "q": "Kein Vogel ist ein Saeugetier. Eine Fledermaus ist ein Saeugetier. Was folgt?",
     "choices": {"A": "Eine Fledermaus ist ein Vogel", "B": "Eine Fledermaus ist kein Vogel",
                 "C": "Manche Voegel sind Fledermaeuse", "D": "Nichts folgt"}, "answer": "B"},
    {"id": "r8", "q": "5 Maschinen brauchen 5 Minuten fuer 5 Teile. Wie lange brauchen 100 Maschinen fuer 100 Teile?",
     "choices": {"A": "100 Minuten", "B": "20 Minuten", "C": "5 Minuten", "D": "1 Minute"}, "answer": "C"},
    {"id": "r9", "q": "Wenn A>B und B>C und C>D, und D=10, A=40, welche Aussage MUSS wahr sein?",
     "choices": {"A": "B=30", "B": "C=20", "C": "B liegt zwischen 10 und 40", "D": "B+C=50"}, "answer": "C"},
    {"id": "r10", "q": "Ein Seerosenfeld verdoppelt sich taeglich und bedeckt den See an Tag 48 vollstaendig. An welchem Tag war der See halb bedeckt?",
     "choices": {"A": "Tag 24", "B": "Tag 47", "C": "Tag 46", "D": "Tag 12"}, "answer": "B"},
    {"id": "r11", "q": "Alle Premium-Kunden bekommen Support. Kai bekommt keinen Support. Was folgt zwingend?",
     "choices": {"A": "Kai ist Premium-Kunde", "B": "Kai ist kein Premium-Kunde",
                 "C": "Kai bekommt spaeter Support", "D": "Nichts folgt"}, "answer": "B"},
    {"id": "r12", "q": "Drei Freunde teilen eine Rechnung von 90 Euro. Einer zahlt doppelt so viel wie jeder der anderen zwei (die gleich viel zahlen). Wie viel zahlt der, der am meisten zahlt?",
     "choices": {"A": "30 Euro", "B": "36 Euro", "C": "45 Euro", "D": "60 Euro"}, "answer": "C"},
]

REASONING_PROMPT = """Beantworte jede Frage Schritt fuer Schritt (denke kurz nach),
gib dann fuer JEDE Frage den Buchstaben der richtigen Antwort.

Antworte AUSSCHLIESSLICH mit einem JSON-Objekt: {{"r1": "A", "r2": "B", ...}}
Keine Erklaerung im JSON, nur id->Buchstabe.

Fragen:
{questions}
"""


def build_reasoning_prompt() -> str:
    lines = []
    for it in REASONING_ITEMS:
        ch = "  ".join(f"{k}) {v}" for k, v in it["choices"].items())
        lines.append(f'{it["id"]}: {it["q"]}\n   {ch}')
    return REASONING_PROMPT.format(questions="\n".join(lines))


# --- Long-Context-Probe: RULER-CWE-Stil Needle-Retrieval ---------------------
FILLER_SENTENCE = ("Das Wissenssystem verarbeitet Notizen sorgfaeltig und haelt "
                   "den Kontext sauber, um zuverlaessige Antworten zu liefern. ")
# Feste, deterministische Needles (token -> code). Reproduzierbar ueber Laeufe.
NEEDLES = [
    ("orchid", "7392"), ("basalt", "1845"), ("nimbus", "9061"),
    ("verdigris", "5573"), ("quokka", "2208"), ("zephyr", "4417"),
    ("marlin", "6630"), ("cobalt", "8124"),
]


def build_longctx(filler_chars: int, n_needles: int) -> tuple[str, list[tuple[str, str]]]:
    needles = NEEDLES[:n_needles]
    base = (FILLER_SENTENCE * (filler_chars // len(FILLER_SENTENCE) + 1))[:filler_chars]
    # Needles deterministisch ueber den Text verteilt einstreuen
    chunk = max(1, len(base) // (len(needles) + 1))
    parts, pos = [], 0
    for i, (tok, code) in enumerate(needles):
        seg_end = pos + chunk
        parts.append(base[pos:seg_end])
        parts.append(f" MERKE: Der geheime Code fuer '{tok}' ist {code}. ")
        pos = seg_end
    parts.append(base[pos:])
    context = "".join(parts)
    return context, needles


def build_longctx_prompt(context: str, needles: list[tuple[str, str]]) -> str:
    asked = ", ".join(f"'{tok}'" for tok, _ in needles)
    return (
        "Im folgenden langen Text sind mehrere geheime Codes versteckt "
        "(Format: \"Der geheime Code fuer 'X' ist NNNN\").\n"
        f"Finde die Codes fuer diese Tokens: {asked}.\n"
        "Antworte AUSSCHLIESSLICH mit JSON: {\"token\": \"code\", ...}.\n\n"
        "=== TEXT START ===\n" + context + "\n=== TEXT ENDE ==="
    )


# --- Model-Call + Parsing ----------------------------------------------------
def call_model(client, model: str, prompt: str, max_tokens: int = 1024) -> tuple[dict | None, str]:
    try:
        resp = client.messages.create(
            model=model, max_tokens=max_tokens,
            messages=[{"role": "user", "content": prompt}],
        )
    except Exception as e:  # noqa
        return None, f"api_error: {e}"
    text = "".join(b.text for b in resp.content if hasattr(b, "text"))
    raw = text
    try:
        if "```" in text:
            text = text.split("```")[1]
            if text.startswith("json"):
                text = text[4:]
        first, last = text.find("{"), text.rfind("}")
        if first != -1 and last != -1:
            return json.loads(text[first:last + 1]), raw
    except (json.JSONDecodeError, ValueError):
        pass
    return None, raw


def score_reasoning(parsed: dict | None) -> tuple[int, int, list[str]]:
    if not parsed:
        return 0, len(REASONING_ITEMS), [it["id"] for it in REASONING_ITEMS]
    correct, wrong = 0, []
    for it in REASONING_ITEMS:
        got = str(parsed.get(it["id"], "")).strip().upper()[:1]
        if got == it["answer"]:
            correct += 1
        else:
            wrong.append(it["id"])
    return correct, len(REASONING_ITEMS), wrong


def score_longctx(parsed: dict | None, needles: list[tuple[str, str]]) -> tuple[int, int, list[str]]:
    if not parsed:
        return 0, len(needles), [t for t, _ in needles]
    # case-insensitive key-match
    norm = {str(k).strip().lower().strip("'\""): str(v) for k, v in parsed.items()}
    correct, missed = 0, []
    for tok, code in needles:
        got = norm.get(tok.lower(), "")
        if code in got:
            correct += 1
        else:
            missed.append(tok)
    return correct, len(needles), missed


def main(argv: list[str]) -> int:
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", default=None)
    ap.add_argument("--needles", type=int, default=8)
    ap.add_argument("--filler-chars", type=int, default=12000)
    ap.add_argument("--json", action="store_true", help="nur JSON-Ergebnis ausgeben")
    args = ap.parse_args(argv)

    # Modell aufloesen
    model = args.model or os.environ.get("ARIA_HEALTH_PROBE_MODEL")
    if not model:
        try:
            import yaml
            with CONFIG_PATH.open() as f:
                model = yaml.safe_load(f)["models"]["triage_model"]
        except Exception:  # noqa
            model = "claude-haiku-4-5-20251001"

    api_key = os.environ.get("ANTHROPIC_API_KEY")
    if not api_key:
        print("[health-probe] ANTHROPIC_API_KEY not set -> skip (no crash).")
        return 0
    try:
        from anthropic import Anthropic
    except ImportError:
        print("[health-probe] anthropic SDK missing. pip install anthropic")
        return 2
    client = Anthropic(api_key=api_key)

    # 1) Reasoning
    r_parsed, r_raw = call_model(client, model, build_reasoning_prompt())
    r_correct, r_total, r_wrong = score_reasoning(r_parsed)

    # 2) Long-Context
    ctx, needles = build_longctx(args.filler_chars, args.needles)
    l_parsed, l_raw = call_model(client, model, build_longctx_prompt(ctx, needles))
    l_correct, l_total, l_missed = score_longctx(l_parsed, needles)

    reasoning_pct = round(100.0 * r_correct / r_total, 1)
    longctx_pct = round(100.0 * l_correct / l_total, 1)
    record = {
        "ts": datetime.now(timezone.utc).isoformat(),
        "model": model,
        "reasoning": {"correct": r_correct, "total": r_total, "pct": reasoning_pct,
                      "wrong": r_wrong, "parse_ok": r_parsed is not None},
        "longctx": {"correct": l_correct, "total": l_total, "pct": longctx_pct,
                    "missed": l_missed, "filler_chars": args.filler_chars,
                    "parse_ok": l_parsed is not None},
    }

    LOG_DIR.mkdir(parents=True, exist_ok=True)
    with LOG_PATH.open("a") as f:
        f.write(json.dumps(record, ensure_ascii=False) + "\n")

    if args.json:
        print(json.dumps(record, ensure_ascii=False, indent=2))
    else:
        print(f"[health-probe] model={model}")
        print(f"  Reasoning (ARC-Stil):     {r_correct}/{r_total} = {reasoning_pct}%"
              + (f"  wrong={r_wrong}" if r_wrong else ""))
        print(f"  Long-Context (RULER-Stil):{l_correct}/{l_total} = {longctx_pct}%"
              + (f"  missed={l_missed}" if l_missed else ""))
        print(f"  -> {LOG_PATH}")
    return 0


if __name__ == "__main__":
    sys.exit(main(sys.argv[1:]))
