#!/usr/bin/env python3
"""
Proper training using SalienceOS architecture.

This script demonstrates the CORRECT way to train MONIKA:
1. Use ConversationSession (not direct proto_lm calls)
2. Enable salience filtering (conditional learning)
3. Run through full runtime orchestration
4. Track adaptive coordinator
"""

import sys
from pathlib import Path

# Add repo to path
sys.path.insert(0, str(Path(__file__).parent))

from salience_os_seed.conversation.session import (
    ConversationSession,
    ConversationConfig,
    IngestionConfig,
)
from salience_os_seed.conversation.filters import IngestionThresholds
from salience_os_seed.proto_lm.trainer import TrainingConfig


def train_with_salience_gating(
    examples: list[str],
    *,
    checkpoint_path: str = "storage/proto_lm/checkpoint.pt",
    enable_filter: bool = True,
    min_uncertainty: float = 0.05,
    min_novelty: float = 0.05,
    max_drag: float = 0.9,
    verbose: bool = True,
):
    """
    Train using proper salience-gated pipeline.
    
    Args:
        examples: List of training text examples
        checkpoint_path: Where to save checkpoint
        enable_filter: Whether to use salience filtering
        min_uncertainty: Minimum uncertainty to accept (0-1)
        min_novelty: Minimum novelty to accept (0-1)
        max_drag: Maximum drag to accept (0-1)
        verbose: Print detailed metrics
    """
    
    # Configure training
    training_cfg = TrainingConfig()
    training_cfg.checkpoint_path = checkpoint_path
    training_cfg.device = "cuda"  # or "cpu"
    
    # Configure ingestion with salience filtering
    ingestion_cfg = IngestionConfig(
        thresholds=IngestionThresholds(
            enabled=enable_filter,
            min_uncertainty=min_uncertainty,
            min_novelty=min_novelty,
            max_drag=max_drag,
        ),
        chunk_size=2048,
        dedupe_enabled=True,  # Skip duplicates
        allow_reingest_duplicates=False,
    )
    
    # Create session (includes runtime + proto_lm + adaptive coordinator)
    session = ConversationSession(
        ConversationConfig(
            lm=training_cfg,
            learning_enabled=True,
            ingestion=ingestion_cfg,
            auto_save_path=checkpoint_path,
        )
    )
    
    print(f"=== Starting Training ===")
    print(f"Device: {session.proto_lm.device}")
    print(f"Initial step: {session.proto_lm.step}")
    print(f"Initial vocab: {session.proto_lm.vocab.size()}")
    print(f"Salience filter: {'ENABLED' if enable_filter else 'DISABLED'}")
    if enable_filter:
        print(f"  - min_uncertainty: {min_uncertainty}")
        print(f"  - min_novelty: {min_novelty}")
        print(f"  - max_drag: {max_drag}")
    print()
    
    accepted_total = 0
    rejected_total = 0
    
    for idx, text in enumerate(examples, 1):
        # This does the full pipeline:
        # 1. Salience evaluation (if filter enabled)
        # 2. Training step (if accepted)
        # 3. Runtime orchestration
        # 4. Adaptive tracking
        processed, metrics = session.ingest_text(
            text,
            source=f"example_{idx}",
            allow_duplicates=False,
        )
        
        if processed > 0:
            accepted_total += processed
            if verbose:
                print(f"✓ Example {idx} ACCEPTED")
                print(f"  Step: {session.proto_lm.step}")
                print(f"  Loss: {session.proto_lm._latest_loss:.4f}")
                print(f"  Salience: {metrics.salience_raw}")
                print(f"  Decision: {metrics.decision.action.operator.name}")
                print(f"  Meta: {metrics.meta_report}")
                print()
        else:
            rejected_total += 1
            if verbose:
                print(f"✗ Example {idx} REJECTED (low salience)")
                print()
    
    # Save final checkpoint
    final_path = session.proto_lm.save_checkpoint(
        checkpoint_path,
        reason="training_complete",
        metadata={"accepted": accepted_total, "rejected": rejected_total},
    )
    
    print(f"=== Training Complete ===")
    print(f"Final step: {session.proto_lm.step}")
    print(f"Final vocab: {session.proto_lm.vocab.size()}")
    print(f"Accepted: {accepted_total}")
    print(f"Rejected: {rejected_total}")
    print(f"Checkpoint: {final_path}")
    
    return session


def main():
    """Example usage."""
    
    # Example training data (complex philosophical sentences)
    training_examples = [
        "Systems exhibit emergent properties from component interactions following simple local rules",
        "Rules constrain behavior creating order and enabling coordination among distributed agents",
        "Agents pursue objectives through planned action sequences adapted to current conditions",
        "Conditions determine what actions prove feasible and effective given available resources",
        "Resources include time attention energy knowledge skills and material assets",
        "Hello I am MONIKA learning to communicate through language",
        "Language enables humans to communicate complex abstract ideas precisely and efficiently",
        "Efficiently structured communication reduces misunderstanding by establishing shared meanings",
        "Meanings emerge through repeated usage patterns in social interaction contexts",
        "Thank you for teaching me how to speak and understand",
    ]
    
    print("="*80)
    print("PROPER SALIENCE-GATED TRAINING")
    print("="*80)
    print()
    
    session = train_with_salience_gating(
        training_examples,
        checkpoint_path="storage/proto_lm/properly_trained.pt",
        enable_filter=True,  # Use salience gating
        min_uncertainty=0.05,  # Require some uncertainty
        min_novelty=0.05,      # Require some novelty
        max_drag=0.9,          # Limit computational drag
        verbose=True,
    )
    
    print()
    print("="*80)
    print("TESTING GENERATION")
    print("="*80)
    print()
    
    # Test generation
    test_prompts = [
        "Hello I am MONIKA",
        "Thank you",
        "Systems exhibit",
    ]
    
    for prompt in test_prompts:
        response = session.proto_lm.sample(
            prompt,
            max_tokens=20,
            temperature=0.8,
            repetition_penalty=1.5,
        )
        print(f"Prompt: '{prompt}'")
        print(f"Output: '{response}'")
        print()


if __name__ == "__main__":
    main()
