#!/usr/bin/env python3
"""
NUCLEAR OPTION: Completely rebuild MCP initialization to force fresh model loading
"""
import sys
sys.path.insert(0, 'c:/MONIKA')

from pathlib import Path
import torch
import shutil

print("=" * 70)
print("NUCLEAR FIX: Complete MCP Rebuild")
print("=" * 70)

# 1. NUKE everything
paths_to_delete = [
    "storage/proto_lm",
    "conversation_state.json",
    Path.home() / "AppData/Local/Programs/Windsurf/conversation_state.json",
]

for p in paths_to_delete:
    p = Path(p)
    if p.exists():
        if p.is_dir():
            shutil.rmtree(p)
        else:
            p.unlink()
        print(f"✓ DELETED: {p}")

# 2. Create fresh directories
Path("storage/proto_lm/checkpoints").mkdir(parents=True, exist_ok=True)

# 3. Create checkpoint index
import json
index = {"records": [], "active": None}
Path("storage/proto_lm/checkpoints/index.json").write_text(json.dumps(index, indent=2))
print("✓ Created empty checkpoint index")

# 4. Build FRESH model with NO cached state
from salience_os_seed.proto_lm._torch_impl import ProtoLanguageModel, TrainingConfig

config = TrainingConfig(
    checkpoint_path=None,  # CRITICAL: No checkpoint
    checkpoint_repository="storage/proto_lm/checkpoints",
)

print("\nCreating absolutely fresh model...")
model = ProtoLanguageModel(config)

print(f"Vocab size: {model.vocab.size()}")
print(f"Has word tokens: {' Hello' in model.vocab.token_to_id and ' I' in model.vocab.token_to_id}")
print(f"Step: {model.step}")
print(f"Grad health: {model._latest_grad_health}")

# 5. Quick train to initialize gradients properly
print("\nQuick training to initialize gradients...")
for i, text in enumerate([" Hello", " I am Monika", " Thank you"] * 30):
    model.training_step(text)
    if (i+1) % 30 == 0:
        print(f"  {i+1}/90: loss={model._latest_loss:.4f}, grads={model._latest_grad_health['frac_nonzero']:.2%}")

print(f"\n✓ Final state:")
print(f"  Step: {model.step}")
print(f"  Loss: {model._latest_loss:.4f}")
print(f"  Gradients: {model._latest_grad_health['frac_nonzero']:.2%} active")

# 6. Test generation
result = model.sample(" Hello I", max_tokens=3, temperature=0.3)
print(f"  Generation: ' Hello I' -> '{result}'")

# 7. Save as DEFAULT checkpoint with proper metadata
ckpt_path = Path("storage/proto_lm/checkpoint.pt")
checkpoint = {
    'model': model.state_dict(),
    'optimizer': model.optimizer.state_dict(),
    'step': model.step,
    'vocab': {'tokens': model.vocab.tokens, 'merges': model.vocab.merges},
    'scheduler': None,
    '_metadata': {
        'created': 'NUCLEAR_FIX',
        'vocab_size': model.vocab.size(),
        'has_gradients': True,
    }
}
torch.save(checkpoint, ckpt_path)
print(f"\n✓ Saved to: {ckpt_path}")

# 8. Verify it loads correctly
print("\nVerifying checkpoint loads correctly...")
test_model = ProtoLanguageModel(TrainingConfig(checkpoint_path=str(ckpt_path)))
print(f"  Loaded step: {test_model.step}")
print(f"  Loaded vocab size: {test_model.vocab.size()}")
print(f"  Has ' Hello': {' Hello' in test_model.vocab.token_to_id}")

print("\n" + "=" * 70)
print("NUCLEAR FIX COMPLETE - MCP Server should now load fresh model")
print("=" * 70)
