import gradio as gr
import torch
import gc
import os
from PIL import Image

from model_loader import ModelLoader
from utils import process_image, process_video, generate_references

loader = ModelLoader()

def chat(message: str, history, file=None):
    """Clean & Stable Chat Function"""
    if history is None:
        history = []

    # Empty message handling
    if not message or not message.strip():
        history.append((message or "", "Please type a message or upload a file."))
        return history

    # Process uploaded file
    context = ""
    if file:
        try:
            file_path = file.name if hasattr(file, 'name') else str(file)
            ext = os.path.splitext(file_path)[1].lower()
            if ext in [".jpg", ".jpeg", ".png", ".webp", ".gif"]:
                img = Image.open(file_path)
                caption = process_image(img)
                context = f"**Image Description**: {caption}\n\n"
            elif ext in [".mp4", ".mov", ".avi", ".webm"]:
                summary = process_video(file_path)
                context = f"**Video Analysis**:\n{summary}\n\n"
        except:
            context = "**File uploaded.**\n\n"

    # Generate response
    full_prompt = f"{context}User: {message}\n\nAssistant:"

    try:
        model, tokenizer = loader.load_text_model()
        inputs = tokenizer(full_prompt, return_tensors="pt").to(loader.device)
        
        with torch.no_grad():
            outputs = model.generate(
                **inputs,
                max_new_tokens=512,
                temperature=0.7,
                do_sample=True,
                pad_token_id=tokenizer.eos_token_id,
                repetition_penalty=1.1,
            )
        
        response_text = tokenizer.decode(outputs[0][inputs.input_ids.shape[1]:], 
                                       skip_special_tokens=True)
        
        final_response = f"{response_text.strip()}\n\n**Sources / References:**\n{generate_references(message[:100])}"
        
    except Exception as e:
        final_response = f"⚠️ Sorry, I encountered an error while processing your request."

    # Tuple format (most compatible)
    history.append((message, final_response))

    # Memory cleanup
    gc.collect()
    if torch.cuda.is_available():
        torch.cuda.empty_cache()

    return history


# ===================== UI =====================
with gr.Blocks() as demo:
    gr.Markdown("# 🛡️ Private Personal AI Chatbot\n**Fully Local • Private • Open Source**")
    
    chatbot = gr.Chatbot(
        height=650,
        avatar_images=(None, "🧠"),
        label="Chatbot"
    )
    
    with gr.Row():
        msg = gr.Textbox(
            placeholder="Ask anything... (Image & Video supported)",
            scale=4,
            container=False,
            autofocus=True
        )
        submit_btn = gr.Button("Send", variant="primary")

    with gr.Row():
        file_upload = gr.File(
            label="📎 Upload Image or Video",
            file_types=["image", ".mp4", ".mov", ".avi", ".webm"]
        )
        clear_btn = gr.Button("🗑️ Clear Chat", variant="stop")

    gr.Examples(
        examples=[["List of action movies 2025"]],
        inputs=[msg, file_upload]
    )

    gr.Markdown("**Privacy**: Everything runs in memory. No data is stored.")

    # Events
    submit_btn.click(
        chat,
        inputs=[msg, chatbot, file_upload],
        outputs=[chatbot]
    ).then(lambda: "", outputs=[msg]) \
     .then(lambda: None, outputs=[file_upload])

    msg.submit(
        chat,
        inputs=[msg, chatbot, file_upload],
        outputs=[chatbot]
    ).then(lambda: "", outputs=[msg]) \
     .then(lambda: None, outputs=[file_upload])

    clear_btn.click(lambda: [], outputs=[chatbot], queue=False)


if __name__ == "__main__":
    demo.launch(
        server_name="0.0.0.0",
        server_port=7860,
        theme=gr.themes.Soft(),
        share=False
    )