# main.py

import logging
import asyncio
import os
from typing import Optional, List

from fastapi import FastAPI, Depends, HTTPException, Request, Query
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse, FileResponse

from langchain_ollama import OllamaLLM
from langchain_community.utilities import SQLDatabase
import torch
from transformers import pipeline

# Local imports
from config.settings import Config
from models.schemas import (
    AIResponse, HealthResponse, ConversationCreateRequest,
    ConversationResponse, AskConversationRequest, MessagePair
)
from services.auth_service import AuthService
from services.sql_agent import SQLAgent
from services.general_agent import GeneralAgent
from services.ai_router import AIRouter
from services.cache_service import CacheService
from services.docx_agent import DocxAgent
from services.conversation_service import ConversationService
from pydantic import BaseModel
# --- 1. CONFIGURATION AND LOGGING ---
logging.basicConfig(
    level=getattr(logging, Config.LOG_LEVEL.upper(), logging.INFO),
    format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
logger = logging.getLogger(__name__)

# --- 2. SERVICE INITIALIZATION ---
logger.info("--- Initializing Eltrion AI System v5.0 ---")
try:
    Config.validate()
    db_engine_args = {"connect_args": {"sslmode": "prefer"}}
    db = SQLDatabase.from_uri(Config.DATABASE_URI, engine_args=db_engine_args)

    llm = OllamaLLM(model=Config.OLLAMA_MODEL)
    sql_llm = OllamaLLM(model=Config.SQL_MODEL)
    

    translator = pipeline("translation", model= Config.TRANSLATE_MODEL, src_lang="ara", tgt_lang="eng_Latn", device=-1)

    cache_service = CacheService()
    docx_agent = DocxAgent()
    auth_service = AuthService(db)
    conversation_service = ConversationService(database_uri=Config.DATABASE_URI)
    sql_agent = SQLAgent(db, sql_llm, llm, translator, cache_service)
    general_agent = GeneralAgent(llm)
    ai_router = AIRouter(sql_agent, general_agent, docx_agent, cache_service, conversation_service)

    logger.info("--- Services initialized successfully ---")
except Exception as e:
    logger.error(f"Fatal error during service initialization: {e}", exc_info=True)
    raise RuntimeError("Could not initialize services.") from e

# --- 3. FASTAPI APP ---
app = FastAPI(
    title=Config.API_TITLE,
    description="Intelligent AI assistant with conversational memory.",
    version="5.0.0",
    docs_url="/docs",
    redoc_url="/redoc"
)

app.add_middleware(
    CORSMiddleware,
    allow_origins=Config.CORS_ORIGINS,
    allow_credentials=Config.CORS_CREDENTIALS,
    allow_methods=Config.CORS_METHODS,
    allow_headers=Config.CORS_HEADERS,
)


# --- 4. API ENDPOINTS ---

@app.get("/health", response_model=HealthResponse, tags=["Monitoring"])
async def health_check():
    return HealthResponse(status="healthy", service=Config.API_TITLE, version=Config.API_VERSION)


@app.post("/conversations", response_model=ConversationResponse, tags=["Conversations"])
async def create_new_conversation(
        req: ConversationCreateRequest,
        company_id: int = Depends(auth_service.get_current_company_id)
):
    """Starts a new conversation thread."""
    title = req.title or "New Conversation"
    conversation_id = conversation_service.create_conversation(title, company_id)
    conv_list = conversation_service.get_conversations_for_company(company_id, page=1, page_size=1)
    if not conv_list:
        raise HTTPException(status_code=404, detail="Could not retrieve newly created conversation.")
    return conv_list[0]


@app.get("/conversations", response_model=List[ConversationResponse], tags=["Conversations"])
async def get_all_conversations(
        company_id: int = Depends(auth_service.get_current_company_id),
        page: int = Query(1, gt=0),
        page_size: int = Query(20, gt=0, le=100)
):
    """Retrieves a paginated list of all conversations for the company."""
    return conversation_service.get_conversations_for_company(company_id, page, page_size)


@app.get("/conversations/{conversation_id}", response_model=List[MessagePair], tags=["Conversations"])
async def get_conversation_details(
        conversation_id: int,
        company_id: int = Depends(auth_service.get_current_company_id)
):
    """Retrieves the full message history for a single conversation."""
    return conversation_service.get_conversation_history(conversation_id)


@app.post("/conversations/{conversation_id}/ask", response_model=AIResponse, tags=["Conversations"])
async def ask_in_conversation(
        conversation_id: int,
        req: AskConversationRequest,
        fastapi_req: Request,
        company_id: int = Depends(auth_service.get_current_company_id)
):
    """Asks a new question within the context of an existing conversation."""
    result = await asyncio.to_thread(
        ai_router.route_query,
        question=req.question,
        company_id=str(company_id),
        request=fastapi_req,
        conversation_id=conversation_id
    )
    return AIResponse(**result)


@app.get("/download/{token}", tags=["Documents"])
async def download_document(token: str):
    """Downloads a generated .docx file using a one-time token."""
    logger.info(f"Received download request for token: {token}")
    cache_key = f"download:{token}"
    filepath = cache_service.get(cache_key)

    if not filepath or not isinstance(filepath, str) or not os.path.exists(filepath):
        logger.warning(f"Invalid or expired download token: {token}")
        raise HTTPException(status_code=404, detail="Document not found or link expired.")

    return FileResponse(
        path=filepath,
        media_type='application/vnd.openxmlformats-officedocument.wordprocessingml.document',
        filename=os.path.basename(filepath)
    )
class FeedbackRequest(BaseModel):
    expected_answer: str

@app.post("/feedback/{conversation_id}", tags=["Conversations"])
async def submit_feedback(
        conversation_id: int,
        req: FeedbackRequest,
        company_id: int = Depends(auth_service.get_current_company_id)
):
    """Collect corrected answers for nightly fine‑tune."""
    cache_service.xadd(
        "feedback",
        {
            "conversation_id": conversation_id,
            "expected_answer": req.expected_answer,
        }
    )
    return {"status": "received"}