import os
import sys
import code
import atexit
import litert_lm
import numpy as np

try:
    import readline
except ImportError:
    readline = None

MODEL_PATH = os.path.expanduser("~/.litert-lm/models/gemma-4-E2B-it.litertlm/model.litertlm")

print("[Initialization] Loading optimized LiteRT-LM Context on CPU...")
engine = litert_lm.Engine(MODEL_PATH, backend=litert_lm.Backend.CPU())

# 1. RIGID BRIEF SYSTEM CONSTRAINTS: This cuts off verbose token outputs early
system_instruction = litert_lm.Message.system(
    "You are an embedded Python co-interpreter. Analyze the state variables provided. "
    "CRITICAL: Be extremely brief. Answer in 1 sentence or a direct parameter assignment. "
    "Do not give introductory text, conversational fluff, or long bullet points."
)

# 2. PERSISTENT CONVERSATION MANIFOLD: Reusing this object keeps the KV-Cache hot in RAM
conversation = engine.create_conversation(messages=[system_instruction])

# Graceful context cleanup handler
def shutdown_manifold():
    print("\n[Shutdown] Safely unlinking memory caches and closing conversation...")
    try:
        conversation.close()
    except Exception:
        pass
atexit.register(shutdown_manifold)

# Shared global variable workspace
shared_globals = {
    "np": np,
    "current_coordinates": np.array([5.0, -3.5]),
    "matrix_z": np.array([[1.0, 2.0], [3.0, np.nan]])
}

def ai(query_string):
    """
    High-speed, memory-retaining state inspector.
    """
    # Build light runtime snapshot mapping 
    namespace_snapshot = "RAM State Snapshot:\n"
    for var_name, var_val in list(shared_globals.items()):
        if var_name in ["np", "litert_lm", "shared_globals", "ai", "readline", "atexit"]: 
            continue
        if isinstance(var_val, np.ndarray):
            namespace_snapshot += f"- {var_name}:\n{var_val}\n"
        else:
            namespace_snapshot += f"- {var_name}: {var_val}\n"
            
    full_prompt = (
        f"{namespace_snapshot}\n"
        f"Query: {query_string}"
    )
    
    print("\n[AI Co-Interpreter]: ", end="")
    sys.stdout.flush()
    
    # 3. HIGH-SPEED SAMPLER TWEAKS: 
    # Calling standard send_message_async with strict constraint arguments.
    # Setting k=1 forces deterministic top-1 decoding, which is significantly faster on a CPU.
    stream = conversation.send_message_async(
        full_prompt,
        sampler_params={"type": "TOP_K", "k": 1, "temperature": 0.2}
    )
    
    for chunk in stream:
        try:
            text_piece = chunk["content"][0]["text"]
            sys.stdout.write(text_piece)
            sys.stdout.flush()
        except (KeyError, IndexError):
            pass
    print("\n")

shared_globals["ai"] = ai

if readline:
    readline.parse_and_bind("tab: complete")
    readline.set_history_length(1000)

banner_msg = """
====================================================================
     XYFLOW HYPER-SPEED DUAL INTERPRETER (v6.0-KV-CACHE)
====================================================================
System status: OPERATIONAL (Persistent KV-cache + Top-K optimization)
Brevity lock:  ACTIVE (The AI will provide ultra-short responses)

Try standard sequential test:
>>> ai("Remember this string code value: 42")
>>> ai("What was the numeric string value I told you to remember?")
====================================================================
"""

console = code.InteractiveConsole(locals=shared_globals)
console.interact(banner=banner_msg, exitmsg="")