#!/usr/bin/env python3
"""
Overthinking Detector (prototype)
================================
Analyzes a message history and flags when the author appears to be
"spiraling": going in circles on one topic, repeating themselves, or
escalating emotionally without new information.

Input: a JSON array of messages, e.g.
  [{"id":1,"text":"...","ts":"2026-09-01T10:00:00Z"}, ...]
or a plain text file with one message per line (chat export).

Output: per-window spiraling score (0-100) + human-readable flags.

Run:
  python overthinking_detector.py messages.json
  python overthinking_detector.py chat.txt
"""
from __future__ import annotations
import json, re, sys, hashlib
from collections import Counter
from datetime import datetime, timezone
from pathlib import Path
from typing import Any


def _norm(s: str) -> str:
    s = s.lower()
    s = re.sub(r"https?://\S+", " ", s)
    s = re.sub(r"[^a-z0-9 ]+", " ", s)
    return re.sub(r"\s+", " ", s).strip()


def _tokenize(s: str) -> set[str]:
    return set(_norm(s).split()) - _STOP

_STOP = set("the a an and or but if then is are was were be been to of in on for with at by from that this it as i you he she we they my your our their me him her us them".split())

ESCALATION = re.compile(r"\b(hate|angry|furious|frustrat\w*|stuck|trapped|spiral\w*|panic\w*|anxious|anxiety|overwhelm\w*|can'?t stop|keep going|again and again|obsess\w*|why can'?t|what if|should i|what do i do)\b", re.I)
CIRCULAR = re.compile(r"\b(still|again|repeating|same thing|back to|as i said|like i said|going in circles|round and round|rehash\w*)\b", re.I)


def _jaccard(a: set[str], b: set[str]) -> float:
    if not a and not b: return 0.0
    inter = len(a & b); union = len(a | b)
    return inter / union if union else 0.0


def _sha(s: str) -> str:
    return hashlib.md5(s.encode()).hexdigest()[:10]


def analyze(messages: list[dict[str, Any]], window: int = 10) -> dict[str, Any]:
    """Score windows of `window` messages for spiraling signals."""
    msgs = sorted(messages, key=lambda m: m.get("ts", ""))
    results = []
    for i in range(0, max(1, len(msgs) - window + 1), 1):
        chunk = msgs[i:i + window]
        texts = [m.get("text", "") for m in chunk]
        norms = [_norm(t) for t in texts]

        # 1. Repetition: pair-wise similarity of consecutive messages
        sims = [_jaccard(set(_norm(t).split()), set(_norm(texts[j + 1]).split()))
                for j, t in enumerate(texts[:-1])]
        rep = sum(1 for s in sims if s > 0.45) / max(1, len(sims))

        # 2. Near-duplicate messages (same core content re-sent)
        seen = Counter(_sha(n) for n in norms)
        dup = sum(c - 1 for c in seen.values() if c > 1) / max(1, len(norms))

        # 3. Escalation: emotional keyword density
        esc = sum(len(ESCALATION.findall(t)) for t in texts) / max(1, len(texts))

        # 4. Circularity: loop-back phrases
        circ = sum(len(CIRCULAR.findall(t)) for t in texts) / max(1, len(texts))

        # 5. Message velocity (if timestamps available)
        vel = 0.0
        ts = [m.get("ts") for m in chunk]
        if all(ts):
            try:
                dts = [datetime.fromisoformat(t.replace("Z", "+00:00")) for t in ts]
                if len(dts) > 1:
                    span = (dts[-1] - dts[0]).total_seconds()
                    if span > 0:
                        vel = (len(chunk) - 1) / span * 3600  # msgs/hour
            except Exception:
                pass

        score = round(100 * min(1.0, 0.35 * rep + 0.20 * dup + 0.25 * min(1.0, esc) + 0.15 * min(1.0, circ) + 0.05 * min(1.0, vel / 30.0)), 1)
        flags = []
        if rep > 0.5: flags.append("repeating earlier messages")
        if dup > 0.2: flags.append("near-duplicate messages (re-sending same content)")
        if esc > 0.4: flags.append("emotional escalation")
        if circ > 0.3: flags.append("circular rehashing")
        if vel > 20: flags.append(f"high message velocity ({vel:.0f}/hr)")

        results.append({
            "window": [m.get("id", idx) for idx, m in enumerate(chunk)],
            "score": score,
            "spiraling": score >= 40,
            "flags": flags,
            "suggestion": ("Pause and write one concrete next action." if score >= 40
                           else "No strong spiraling signal."),
        })
    top = sorted(results, key=lambda r: -r["score"])[:5]
    return {
        "messages_analyzed": len(msgs),
        "windows": len(results),
        "overall_risk": round(sum(r["score"] for r in results) / max(1, len(results)), 1),
        "top_spiraling_windows": top,
    }


def load(path: str) -> list[dict[str, Any]]:
    p = Path(path)
    if p.suffix == ".json":
        data = json.loads(p.read_text())
        if isinstance(data, dict) and "messages" in data:
            data = data["messages"]
        return [{"id": i, "text": (m.get("text") or m.get("content") or ""), "ts": m.get("ts") or m.get("timestamp") or m.get("time") or ""}
                for i, m in enumerate(data)]
    lines = [l.strip() for l in p.read_text(errors="replace").splitlines() if l.strip()]
    return [{"id": i, "text": l, "ts": ""} for i, l in enumerate(lines)]


if __name__ == "__main__":
    if len(sys.argv) < 2:
        print(__doc__); sys.exit(1)
    msgs = load(sys.argv[1])
    print(json.dumps(analyze(msgs), indent=2))
