#!/usr/bin/env python3
"""
Comprehensive Diagnosis of Novelty Sensor Issue

HYPOTHESIS: Novelty sensor stuck at zero because state construction doesn't
provide the surprisal values that the sensor expects.
"""

import sys
from pathlib import Path
import numpy as np

sys.path.insert(0, str(Path(__file__).parent))

from salience_os_seed.conversation.session import ConversationSession, ConversationConfig
from salience_os_seed.core.sensors.novelty import NoveltySensor, _extract_surprisal, _extract_tokens
from salience_os_seed.core.sensors.base import MedianMADNormalizer

print("="*100)
print("NOVELTY SENSOR DIAGNOSTIC")
print("="*100)
print()

# ==============================================================================
# PHASE 1: Inspect what state is being built
# ==============================================================================
print("PHASE 1: Inspecting state construction...")
print()

session = ConversationSession(ConversationConfig())

# Build state like the session does
test_text = "Hello MONIKA this is a test"
state = session._build_state(test_text, speaker="user")

print(f"State keys: {list(state.keys())}")
print()

print("Prediction dict:")
for key, val in state.get("prediction", {}).items():
    if isinstance(val, np.ndarray):
        print(f"  {key}: array shape {val.shape}, sample: {val.flatten()[:5]}")
    else:
        print(f"  {key}: {val}")
print()

print("Context dict:")
for key, val in state.get("context", {}).items():
    print(f"  {key}: {val}")
print()

# ==============================================================================
# PHASE 2: Test what novelty sensor extracts
# ==============================================================================
print("PHASE 2: Testing novelty sensor extraction...")
print()

surprisal = _extract_surprisal(state)
tokens = _extract_tokens(state)

print(f"Extracted surprisal: {surprisal}")
print(f"Extracted tokens: {tokens}")
print()

# ==============================================================================
# PHASE 3: Test novelty sensor directly
# ==============================================================================
print("PHASE 3: Testing novelty sensor measurement...")
print()

normalizer = MedianMADNormalizer()
normalizer.register_baseline("novelty", [0.1, 0.3, 0.6, 0.9])
sensor = NoveltySensor(normalizer)

# Measure with current state
raw_novelty = sensor._measure(state, {}, {})
print(f"Raw novelty: {raw_novelty}")
print()

# ==============================================================================
# PHASE 4: Identify the issue
# ==============================================================================
print("PHASE 4: Root cause analysis...")
print()

print("Issue identified:")
print("  - State provides: prediction.token_logits, prediction.entropy_estimate")
print("  - Sensor expects: prediction.surprisal OR prediction.token_logprob")
print("  - Sensor fallback: state.fallback_surprisal (defaults to 0.0)")
print()
print("Result:")
print(f"  - Surprisal extracted: {surprisal} (should be > 0 for novelty)")
print(f"  - Baseline surprisal: 3.5")
print(f"  - Delta: max({surprisal} - 3.5, 0.0) = {max(surprisal - 3.5, 0.0)}")
print(f"  - Novelty: sqrt(tanh({max(surprisal - 3.5, 0.0)}) * freshness) = {raw_novelty}")
print()

# ==============================================================================
# PHASE 5: Test with corrected state
# ==============================================================================
print("PHASE 5: Testing with corrected state...")
print()

# Manually add surprisal
corrected_state = dict(state)
# Calculate surprisal from entropy_estimate
entropy = state["prediction"]["entropy_estimate"]
# Surprisal should be higher for uncertain/novel predictions
corrected_state["prediction"]["surprisal"] = float(entropy * 2.0 + 1.0)  # Heuristic

print(f"Added surprisal: {corrected_state['prediction']['surprisal']}")

# Reset sensor to clear n-gram history
sensor2 = NoveltySensor(normalizer)
raw_novelty_corrected = sensor2._measure(corrected_state, {}, {})

print(f"Raw novelty (corrected): {raw_novelty_corrected}")
print()

# ==============================================================================
# PHASE 6: Test multiple novel inputs
# ==============================================================================
print("PHASE 6: Testing with multiple novel inputs...")
print()

sensor3 = NoveltySensor(normalizer)

test_inputs = [
    "Hello MONIKA",
    "How are you",
    "What is your name",
    "Tell me about yourself",
    "Never seen before xyz quantum",
]

for text in test_inputs:
    test_state = session._build_state(text, speaker="user")
    # Add surprisal based on text uniqueness
    entropy = test_state["prediction"]["entropy_estimate"]
    test_state["prediction"]["surprisal"] = float(entropy * 2.0 + 2.0)
    
    novelty = sensor3._measure(test_state, {}, {})
    print(f"  '{text}' → surprisal={test_state['prediction']['surprisal']:.2f}, novelty={novelty:.3f}")

print()

# ==============================================================================
# PHASE 7: Solution summary
# ==============================================================================
print("="*100)
print("SOLUTION SUMMARY")
print("="*100)
print()

print("ROOT CAUSE:")
print("  State construction in session.py:_build_state() does NOT provide")
print("  the 'surprisal' or 'token_logprob' fields that novelty sensor expects.")
print()

print("CONSEQUENCE:")
print("  - Sensor falls back to fallback_surprisal=0.0")
print("  - Surprisal delta = max(0.0 - 3.5, 0.0) = 0.0")
print("  - Novelty = sqrt(tanh(0.0) * freshness) = 0.0")
print("  - ALL novelty readings become 0.0")
print()

print("FIX OPTIONS:")
print()
print("  Option 1: Add surprisal to state construction")
print("    - Calculate from model's actual prediction uncertainty")
print("    - Use cross-entropy or perplexity from proto_lm")
print()
print("  Option 2: Use entropy_estimate as fallback")
print("    - Modify novelty sensor to check entropy_estimate")
print("    - Already available in prediction dict")
print()
print("  Option 3: Calculate from token_logits")
print("    - Convert logits to probabilities")
print("    - Calculate entropy/surprisal")
print("    - Use as prediction.surprisal")
print()

print("RECOMMENDATION: Option 3")
print("  Calculate surprisal from token_logits in _build_state()")
print("  This keeps the sensor interface unchanged and provides accurate surprisal")
print()

print("="*100)
