#!/usr/bin/env python3
"""
COMPREHENSIVE MCP SERVER FIX
Fixes persistent state caching, checkpoint loading, and gradient issues
"""
import sys
import shutil
from pathlib import Path

print("=" * 60)
print("COMPREHENSIVE MCP FIX")
print("=" * 60)

# 1. Clear ALL cached state files
state_files = [
    "conversation_state.json",
    "storage/proto_lm/checkpoint.pt",
    Path.home() / "AppData/Local/Programs/Windsurf/conversation_state.json",
]

for f in state_files:
    f = Path(f)
    if f.exists():
        f.unlink()
        print(f"✓ Deleted cached state: {f}")

# 2. Clear checkpoint repository
ckpt_dir = Path("storage/proto_lm/checkpoints")
if ckpt_dir.exists():
    for item in ckpt_dir.iterdir():
        if item.is_dir():
            shutil.rmtree(item)
        else:
            item.unlink()
    print(f"✓ Cleared checkpoint repository")

# 3. Reset checkpoint index
index_file = ckpt_dir / "index.json"
index_file.parent.mkdir(parents=True, exist_ok=True)
index_file.write_text('{"records": [], "active": null}')
print(f"✓ Reset checkpoint index")

# 4. Create fresh model with working vocab
sys.path.insert(0, 'c:/MONIKA')
from salience_os_seed.proto_lm._torch_impl import ProtoLanguageModel, TrainingConfig
import torch

print("\n" + "=" * 60)
print("Creating fresh model...")
print("=" * 60)

m = ProtoLanguageModel(TrainingConfig(checkpoint_path=None))

print(f"Vocab size: {m.vocab.size()}")
print(f"Sample tokens [116:126]: {m.vocab.tokens[116:126]}")

# 5. Train with minimal quality data
training = [
    "Hello",
    " I am Monika",
    " I am learning",
    " Thank you",
    " Hello I am Monika",
] * 50

print(f"\nTraining {len(training)} examples...")
for i, text in enumerate(training):
    m.training_step(text)
    if (i+1) % 50 == 0:
        print(f"  {i+1}/{len(training)}: loss={m._latest_loss:.4f}, grads={m._latest_grad_health['frac_nonzero']:.2%}")

print(f"\n✓ Training complete:")
print(f"  Final loss: {m._latest_loss:.4f}")
print(f"  Grad health: {m._latest_grad_health['frac_nonzero']:.2%} active")
print(f"  Step: {m.step}")

# 6. Test generation
result = m.sample(" Hello I", max_tokens=3, temperature=0.2, repetition_penalty=1.0)
print(f"\n✓ Generation test:")
print(f"  ' Hello I' -> '{result}'")

# 7. Save to DEFAULT checkpoint location
default_ckpt = Path("storage/proto_lm/checkpoint.pt")
default_ckpt.parent.mkdir(parents=True, exist_ok=True)

checkpoint = {
    'model': m.state_dict(),
    'optimizer': m.optimizer.state_dict(),
    'step': m.step,
    'vocab': {'tokens': m.vocab.tokens, 'merges': m.vocab.merges},
    'scheduler': None
}

torch.save(checkpoint, default_ckpt)
print(f"\n✓ Saved to: {default_ckpt}")

print("\n" + "=" * 60)
print("MCP SERVER FIX COMPLETE")
print("=" * 60)
print("Now restart the MCP server")
