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...")
# FIXED: Sampler parameter constraints are declared directly at engine initialization
# Setting k=1 creates a highly deterministic, lightning-fast CPU inference step
engine = litert_lm.Engine(
    MODEL_PATH, 
    backend=litert_lm.Backend.CPU(),
    sampler_params={"type": "TOP_K", "k": 1, "temperature": 0.2}
)

# RIGID BRIEF SYSTEM CONSTRAINTS: 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."
)

# PERSISTENT CONVERSATION MANIFOLD: 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 loop.
    """
    # Build clean active global snapshot tracking
    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()
    
    # FIXED: Cleared 'sampler_params' from argument list to respect API signature
    stream = conversation.send_message_async(full_prompt)
    
    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.1-FIXED)
====================================================================
System status: OPERATIONAL (Deterministic top-1 decoding active)
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="")