#!/usr/bin/env python3
"""Test the newly added MONIKA capabilities."""

import asyncio
import json
from pathlib import Path
from mcp import ClientSession, StdioServerParameters
from mcp.client.stdio import stdio_client


async def test_new_capabilities():
    """Test all newly wired capabilities."""
    
    server_script = Path(__file__).parent / "start_mcp_server.py"
    server_params = StdioServerParameters(
        command="python",
        args=[str(server_script)],
        env=None,
    )
    
    async with stdio_client(server_params) as (read, write):
        async with ClientSession(read, write) as session:
            await session.initialize()
            
            print("=" * 70)
            print("TESTING NEW MONIKA CAPABILITIES")
            print("=" * 70)
            
            # List tools to verify new ones are present
            tools_result = await session.list_tools()
            tools = {tool.name: tool for tool in tools_result.tools}
            print(f"\nTotal tools: {len(tools)}")
            
            new_tools = [
                "scratchpad_read",
                "scratchpad_history", 
                "scratchpad_4d_path",
                "sass_query",
                "meta_state_report",
                "verification_suite_run",
                "action_scores_detailed",
            ]
            
            print("\nNew tools status:")
            for tool in new_tools:
                status = "✓" if tool in tools else "✗"
                print(f"  {status} {tool}")
            
            # Test 1: Execute a runtime step to populate scratchpad
            print("\n" + "=" * 70)
            print("TEST 1: Runtime Step + Scratchpad Population")
            print("=" * 70)
            
            result = await session.call_tool("runtime_step", {
                "text": "Let's reason through a complex problem involving recursive algorithms and memory optimization."
            })
            step = json.loads(result.content[0].text)
            print(f"✓ Step {step['step']} executed")
            print(f"  Budget left: {step['budget_left']}")
            
            # Test 2: Read scratchpad
            print("\n" + "=" * 70)
            print("TEST 2: Scratchpad Read")
            print("=" * 70)
            
            result = await session.call_tool("scratchpad_read", {})
            scratch = json.loads(result.content[0].text)
            print(f"✓ Current trace steps: {len(scratch['current_trace'])}")
            print(f"  Token usage: {scratch['current_tokens']}/{scratch['max_tokens']}")
            if scratch['current_trace']:
                print(f"  Last thought: {scratch['current_trace'][-1][:80]}...")
            
            # Test 3: Scratchpad history
            print("\n" + "=" * 70)
            print("TEST 3: Scratchpad History")
            print("=" * 70)
            
            result = await session.call_tool("scratchpad_history", {"max_traces": 3})
            history = json.loads(result.content[0].text)
            print(f"✓ Retrieved {len(history['traces'])} traces")
            print(f"  Rolling token mean: {history['rolling_token_mean']:.1f}")
            for i, trace in enumerate(history['traces']):
                print(f"  Trace {i+1}: {trace['outcome']} - {trace['summary']}")
            
            # Test 4: 4D Path visualization
            print("\n" + "=" * 70)
            print("TEST 4: 4D Reasoning Path")
            print("=" * 70)
            
            result = await session.call_tool("scratchpad_4d_path", {})
            path_data = json.loads(result.content[0].text)
            if path_data.get('path'):
                print(f"✓ 4D Path found!")
                print(f"  Summary: {path_data['summary']}")
                print("\n  ASCII Visualization:")
                print(path_data['ascii_viz'])
                
                # Show path details
                path = path_data['path']
                if 'points' in path and path['points']:
                    print(f"\n  Path coordinates ({len(path['points'])} points):")
                    for i, pt in enumerate(path['points'][:3]):
                        print(f"    {i}: x={pt['x']:.2f} y={pt['y']:.2f} z={pt['z']:.2f} w={pt['w']:.2f}")
                    print(f"  Centroid: {path['centroid']}")
            else:
                print("  No 4D path yet:", path_data.get('message'))
            
            # Test 5: SASS Query
            print("\n" + "=" * 70)
            print("TEST 5: SASS Semantic Search")
            print("=" * 70)
            
            result = await session.call_tool("sass_query", {
                "query": "algorithms recursive optimization"
            })
            sass_result = json.loads(result.content[0].text)
            print(f"✓ SASS query: '{sass_result['query']}'")
            print(f"  Results: {len(sass_result['results'])}")
            for i, res in enumerate(sass_result['results'][:3]):
                print(f"  {i+1}. [score={res['score']:.3f}] {res['text'][:80]}...")
            
            # Test 6: Meta State Report
            print("\n" + "=" * 70)
            print("TEST 6: Meta State Self-Report")
            print("=" * 70)
            
            result = await session.call_tool("meta_state_report", {})
            meta = json.loads(result.content[0].text)
            print(f"✓ Meta state:")
            print(f"  Confidence: {meta['confidence']:.3f}")
            print(f"  ROI: {meta['roi']:.3f}")
            print(f"  History length: {meta['history_length']}")
            print(f"  Snapshot keys: {list(meta['snapshot'].keys())}")
            
            # Test 7: Verification Suite
            print("\n" + "=" * 70)
            print("TEST 7: Verification Suite")
            print("=" * 70)
            
            result = await session.call_tool("verification_suite_run", {
                "context": "Testing MONIKA verification systems"
            })
            verify = json.loads(result.content[0].text)
            print(f"✓ Verification outcomes: {len(verify['outcomes'])}")
            passed = sum(1 for o in verify['outcomes'] if o['passed'])
            print(f"  Passed: {passed}/{len(verify['outcomes'])}")
            for i, outcome in enumerate(verify['outcomes'][:3]):
                status = "PASS" if outcome['passed'] else "FAIL"
                print(f"  {i+1}. [{status}] {outcome['evidence'][:60]}...")
            
            # Test 8: Detailed Action Scores
            print("\n" + "=" * 70)
            print("TEST 8: Detailed Action Scores & Salience")
            print("=" * 70)
            
            result = await session.call_tool("action_scores_detailed", {})
            actions = json.loads(result.content[0].text)
            print(f"✓ Action candidates: {len(actions['scores'])}")
            
            if actions['scores']:
                print("  Top 5 actions:")
                for i, action in enumerate(actions['scores'][:5]):
                    act = action['action']
                    print(f"    {i+1}. {act['operator']} (depth={act['depth']}, patch={act['patch']})")
                    print(f"       Score: {action['score']:.3f}")
                    print(f"       Rationale: {action['rationale'][:60]}...")
            else:
                print("  No action scores yet (controller not initialized)")
            
            # Show top salience signals
            salience = actions.get('salience_vector', {})
            if salience:
                sorted_sal = sorted(salience.items(), key=lambda x: abs(x[1]), reverse=True)
                print(f"\n  Top 5 salience signals:")
                for key, val in sorted_sal[:5]:
                    print(f"    {key}: {val:.4f}")
            
            # Final summary
            print("\n" + "=" * 70)
            print("CAPABILITY TEST COMPLETE")
            print("=" * 70)
            
            working = sum(1 for tool in new_tools if tool in tools)
            print(f"✓ {working}/{len(new_tools)} new capabilities operational")
            print("\nMONIKA's internal state is now fully introspectable!")
            print("You can now:")
            print("  • View 4D reasoning paths in real-time")
            print("  • Execute SASS semantic queries")
            print("  • Monitor scratchpad working memory")
            print("  • Read meta-state reports")
            print("  • Run verification suites")
            print("  • Inspect detailed action scoring")


if __name__ == "__main__":
    asyncio.run(test_new_capabilities())
