import gradio as gr
import os
import asyncio
import gc
import torch
import cv2
import numpy as np
from PIL import Image
from transformers import AutoModelForCausalLM, AutoTokenizer
from faster_whisper import WhisperModel
import edge_tts

# --- SYSTEM SETUP ---
TEMP_DIR = "/tmp/jarvis_cache"
os.makedirs(TEMP_DIR, exist_ok=True)

class JarvisEngine:
    def __init__(self):
        # Ears: Faster Whisper (Tiny for CPU speed)
        self.stt = WhisperModel("tiny", device="cpu", compute_type="int8")
        # Brain: Qwen 0.5B (Ultra-low latency Jarvis personality)
        self.model_id = "Qwen/Qwen2.5-0.5B-Instruct"
        self.tokenizer = AutoTokenizer.from_pretrained(self.model_id)
        self.llm = AutoModelForCausalLM.from_pretrained(self.model_id)
        # Eyes: Moondream (Lazy Loaded to save RAM)
        self.vision_model = None
        self.vision_tokenizer = None

    def load_eyes(self):
        if self.vision_model is None:
            self.vision_model = AutoModelForCausalLM.from_pretrained(
                "vikhyatk/moondream2", 
                trust_remote_code=True
            )
            self.vision_tokenizer = AutoTokenizer.from_pretrained("vikhyatk/moondream2")

    async def generate_voice(self, text):
        path = os.path.join(TEMP_DIR, "jarvis_voice.mp3")
        # RyanNeural = Classic British Jarvis tone
        communicate = edge_tts.Communicate(text, "en-GB-RyanNeural")
        await communicate.save(path)
        return path

# Initialize Jarvis
jarvis = JarvisEngine()

# --- FUTURISTIC JARVIS UI STYLING ---
jarvis_css = """
.orb-container { display: flex; justify-content: center; align-items: center; padding: 30px; background: #0b0f19; border-radius: 15px; }
.orb {
    width: 150px; height: 150px; border-radius: 50%;
    background: radial-gradient(circle, #00d2ff 0%, #3a7bd5 100%);
    box-shadow: 0 0 50px #00d2ff;
    animation: pulse 2s infinite ease-in-out;
}
@keyframes pulse {
    0% { transform: scale(1); box-shadow: 0 0 20px #00d2ff; }
    50% { transform: scale(1.1); box-shadow: 0 0 70px #00d2ff; }
    100% { transform: scale(1); box-shadow: 0 0 20px #00d2ff; }
}
.status-active { color: #00ffcc !important; font-weight: bold; text-align: center; font-size: 1.2em; }
.status-sleep { color: #ff4b4b !important; font-weight: bold; text-align: center; font-size: 1.2em; }
footer {visibility: hidden}
"""

# --- CORE ASSISTANT LOGIC ---
async def process_audio(audio_path, history, active_state):
    if not active_state or audio_path is None:
        return history, None

    # 1. Ears (STT)
    segments, _ = jarvis.stt.transcribe(audio_path)
    user_text = " ".join([s.text for s in segments])
    if not user_text.strip():
        return history, None

    # 2. Brain (LLM)
    msgs = [
        {"role": "system", "content": "You are Jarvis, a sophisticated British AI assistant. Be concise, professional, and witty."},
        {"role": "user", "content": user_text}
    ]
    inputs = jarvis.tokenizer(jarvis.tokenizer.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True), return_tensors="pt")
    outputs = jarvis.llm.generate(**inputs, max_new_tokens=128)
    response = jarvis.tokenizer.batch_decode(outputs, skip_special_tokens=True)[0].split("assistant")[-1].strip()

    # 3. Mouth (TTS)
    voice_file = await jarvis.generate_voice(response)
    
    history.append({"role": "user", "content": user_text})
    history.append({"role": "assistant", "content": response})
    return history, voice_file

def analyze_screen(img):
    if img is None: return "Sir, I require a visual feed to analyze the screen."
    jarvis.load_eyes()
    pil_img = Image.fromarray(img)
    desc = jarvis.vision_model.answer_question(
        jarvis.vision_model.encode_image(pil_img), 
        "Describe this screen content in detail and explain what is happening.", 
        jarvis.vision_tokenizer
    )
    return desc

# --- UI LAYOUT ---
with gr.Blocks(css=jarvis_css, theme=gr.themes.Default()) as demo:
    is_active = gr.State(False)
    
    gr.HTML("<div class='orb-container'><div class='orb'></div></div>")
    status_display = gr.Markdown("### 🔴 Assistant Sleeping", elem_classes="status-sleep")
    
    with gr.Row():
        start_btn = gr.Button("🟢 ACTIVATE JARVIS", variant="primary")
        stop_btn = gr.Button("🔴 DEACTIVATE", variant="stop")

    with gr.Tabs():
        with gr.Tab("🎙️ Voice Interface"):
            chatbot = gr.Chatbot(type="messages", label="Jarvis Console", height=400)
            audio_in = gr.Audio(sources=["microphone"], type="filepath", label="Voice Command")
            audio_out = gr.Audio(autoplay=True, visible=False)

        with gr.Tab("🖥️ Screen Understanding"):
            gr.Markdown("Sir, please paste a screenshot (Ctrl+V) or upload an image for analysis.")
            screen_in = gr.Image(label="Visual Input", sources=["upload", "clipboard"])
            screen_btn = gr.Button("Analyze Visuals", variant="primary")
            screen_out = gr.Textbox(label="Jarvis Insight", lines=5)

    # --- EVENT HANDLERS ---
    def set_active():
        return gr.update(value="### 🟢 Jarvis Active", elem_classes="status-active"), True
    
    def set_inactive():
        return gr.update(value="### 🔴 Assistant Sleeping", elem_classes="status-sleep"), False

    start_btn.click(set_active, None, [status_display, is_active])
    stop_btn.click(set_inactive, None, [status_display, is_active])
    
    # Voice interaction triggers when user stops recording
    audio_in.stop_recording(process_audio, [audio_in, chatbot, is_active], [chatbot, audio_out])
    
    # Vision interaction
    screen_btn.click(analyze_screen, [screen_in], [screen_out])

if __name__ == "__main__":
    # show_api=False is critical to bypass the schema bug
    demo.launch(server_name="0.0.0.0", server_port=7860, show_api=False, share=False)