from typing import Dict, Any, List
from sqlalchemy.orm import Session
import aiohttp
import json
from src.utils.pdf_parser import get_pdf_parser
from src.core.logging import logger
from src.core.config import settings
from src.modules.models import BnkParsedDocument, BnkDocumentChunk, BnkConsolidatedLLMOutput, BnkAIReqRes
from src.utils.azure_openai import AzureOpenAIClient
from src.modules.prompt.controller import build_response_from_openai_result


# ---------------- Helper Functions -----------------
def estimate_tokens(text: str) -> int:
    """Estimate token count (rough: 1 token ≈ 4 chars)"""
    return len(text) // 4


def chunk_text_with_overlap(text: str, max_tokens: int = 120000, overlap_tokens: int = 500) -> List[Dict]:
    """
    Split text into chunks with overlapping for context continuity
    
    Args:
        text: Full extracted text
        max_tokens: Max tokens per chunk
        overlap_tokens: Number of tokens to overlap between chunks
    
    Returns:
        List of chunk dicts with chunk_index, chunk_text, token_count
    """
    total_tokens = estimate_tokens(text)
    
    if total_tokens <= max_tokens:
        return [{
            "chunk_index": 0,
            "chunk_text": text,
            "token_count": total_tokens
        }]
    
    chunks = []
    max_chars = max_tokens * 4
    overlap_chars = overlap_tokens * 4
    
    start = 0
    chunk_idx = 0
    
    while start < len(text):
        end = min(start + max_chars, len(text))
        
        # Find sentence boundary for clean break
        if end < len(text):
            for separator in ['. ', '.\n', '!\n', '?\n', '\n\n']:
                last_sep = text[start:end].rfind(separator)
                if last_sep != -1 and last_sep > max_chars // 2:
                    end = start + last_sep + len(separator)
                    break
        
        chunk_text = text[start:end].strip()
        chunks.append({
            "chunk_index": chunk_idx,
            "chunk_text": chunk_text,
            "token_count": estimate_tokens(chunk_text)
        })
        
        # Move start with overlap (go back overlap_chars for next chunk)
        if end >= len(text):
            break
        start = end - overlap_chars
        if start < 0:
            start = 0
        chunk_idx += 1
    
    logger.info(f"[Chunk] Created {len(chunks)} chunks with {overlap_tokens} token overlap")
    return chunks





# ==================== TASK A: Extract & Chunk ====================
async def task_extract_and_chunk(
    db: Session,
    ai_req_res_id: int,
    file_urls: List[str],
    max_tokens: int = 120000,
    overlap_tokens: int = 500
) -> Dict[str, Any]:
    """
    Task A: Extract text from documents and store chunks in DB
    Works for 1 or many documents
    
    Args:
        db: Database session
        ai_req_res_id: Primary key ID from bnk_ai_req_res table
        file_urls: List of document URLs (can be single item)
        max_tokens: Max tokens per chunk
        overlap_tokens: Overlap between chunks
    
    Returns:
        Dict with all parsed_document_ids and chunk_ids
    """
    all_parsed_docs = []
    all_chunk_ids = []
    
    for file_url in file_urls:
        # Extract text using LlamaParse
        try:
            parser = get_pdf_parser()
            documents = parser.parse_url(file_url)
            extracted_text = parser.get_text(documents)
            logger.info(f"[Task A] Extracted {len(extracted_text)} chars from {file_url}")
        except Exception as e:
            logger.error(f"[Task A] Extraction failed for {file_url}: {e}")
            continue
        
        if not extracted_text:
            logger.warning(f"[Task A] No text extracted from {file_url}")
            continue
        
        # Chunk with overlap
        chunks = chunk_text_with_overlap(extracted_text, max_tokens, overlap_tokens)
        
        # Save to BnkParsedDocument
        parsed_doc = BnkParsedDocument(
            ai_req_res_id=ai_req_res_id,
            source_url=file_url,
            extracted_text=extracted_text,
            token_count=estimate_tokens(extracted_text),
            chunked=len(chunks)
        )
        db.add(parsed_doc)
        db.commit()
        db.refresh(parsed_doc)
        all_parsed_docs.append(parsed_doc.id)
        logger.info(f"[Task A] Saved BnkParsedDocument id={parsed_doc.id}")
        
        # Save each chunk
        for chunk_data in chunks:
            chunk_record = BnkDocumentChunk(
                parsed_document_id=parsed_doc.id,
                chunk_index=chunk_data["chunk_index"],
                chunk_text=chunk_data["chunk_text"],
                token_count=chunk_data["token_count"]
            )
            db.add(chunk_record)
            db.commit()
            db.refresh(chunk_record)
            all_chunk_ids.append(chunk_record.id)
    
    return {
        "success": True,
        "ai_req_res_id": ai_req_res_id,
        "parsed_document_ids": all_parsed_docs,
        "chunk_ids": all_chunk_ids,
        "documents_processed": len(all_parsed_docs),
        "total_chunks": len(all_chunk_ids)
    }


# ==================== TASK B: Process Chunks with LLM ====================
async def task_process_chunks_llm(
    db: Session,
    ai_req_res_id: int,
    file_type: str,
    data: Dict = None,
    request_id: str = None
) -> Dict[str, Any]:
    """
    Task B: Read unprocessed chunks from DB and send to LLM
    Uses AzureOpenAIClient.verify_document with pre-extracted chunk text
    
    Args:
        db: Database session
        ai_req_res_id: Primary key ID from bnk_ai_req_res table
        file_type: Document type for prompt selection (same as original flow)
        data: Payload data from request (for verification matching)
        request_id: Request ID for Pusher notifications (jailbreak errors only)
    
    Returns:
        Dict with processed chunk count
    """
    # Get all parsed documents for this request
    parsed_docs = db.query(BnkParsedDocument).filter(
        BnkParsedDocument.ai_req_res_id == ai_req_res_id
    ).all()
    
    if not parsed_docs:
        return {"success": False, "error": f"No parsed documents for ai_req_res_id={ai_req_res_id}"}
    
    # Use existing AzureOpenAIClient
    client = AzureOpenAIClient()
    processed_count = 0
    jailbreak_detected = False  # Track if any document has jailbreak
    
    for parsed_doc in parsed_docs:
        # Get unprocessed chunks (llm_output is NULL)
        chunks = db.query(BnkDocumentChunk).filter(
            BnkDocumentChunk.parsed_document_id == parsed_doc.id,
            BnkDocumentChunk.llm_output == None
        ).order_by(BnkDocumentChunk.chunk_index).all()
        
        document_has_jailbreak = False  # Flag per document
        
        for chunk in chunks:
            try:
                # Log which prompt will be used based on file_type
                logger.info(f"[Task B] Sending chunk id={chunk.id} to LLM with file_type='{file_type}' (prompt selected based on this type)")
                
                # Use verify_document with pre-extracted chunk text (skips LlamaParse)
                llm_output = await client.verify_document(
                    file_type=file_type,
                    payload=data or {},
                    file_url=parsed_doc.source_url or "",
                    extracted_text=chunk.chunk_text,
                    request_id=request_id
                )
                
                logger.info(f"[Task B] Got response for chunk {chunk.id}: type={type(llm_output)}, has_error={'error' in llm_output if isinstance(llm_output, dict) else 'N/A'}")
                
                # Check if jailbreak error - skip remaining chunks of THIS document
                if isinstance(llm_output, dict) and llm_output.get("error") and llm_output.get("error_type") == "content_filter":
                    logger.warning(f"[Task B] Jailbreak detected for chunk {chunk.id} of document {parsed_doc.id} - notification sent, skipping remaining chunks of this document")
                    chunk.llm_output = llm_output
                    db.commit()
                    document_has_jailbreak = True
                    jailbreak_detected = True  # Flag for celery to skip Task 3
                    break  # Skip remaining chunks of this document
                
                chunk.llm_output = llm_output
                db.commit()
                processed_count += 1
                logger.info(f"[Task B] Stored llm_output for chunk id={chunk.id}")
                
            except Exception as e:
                logger.error(f"[Task B] Chunk id={chunk.id} failed with exception: {e}")
                chunk.llm_output = {"error": str(e)}
                db.commit()
        
        if document_has_jailbreak:
            logger.info(f"[Task B] Skipped remaining chunks of document {parsed_doc.id} due to jailbreak")
            # Mark all remaining unprocessed chunks of this document with jailbreak error
            remaining_chunks = db.query(BnkDocumentChunk).filter(
                BnkDocumentChunk.parsed_document_id == parsed_doc.id,
                BnkDocumentChunk.llm_output == None
            ).all()
            for remaining_chunk in remaining_chunks:
                remaining_chunk.llm_output = {
                    "error": True,
                    "error_type": "content_filter",
                    "message": "Skipped due to jailbreak in previous chunk of same document"
                }
            db.commit()
    
    # Note: Jailbreak errors are stored in chunk.llm_output and will be
    # processed by Task C consolidation, then formatted by controller.py
    
    return {
        "success": True,
        "ai_req_res_id": ai_req_res_id,
        "chunks_processed": processed_count,
        "jailbreak_detected": jailbreak_detected
    }


# ==================== TASK C: Consolidate LLM Outputs ====================
async def task_consolidate_outputs(
    db: Session,
    ai_req_res_id: int,
    file_type: str = "unknown"
) -> Dict[str, Any]:
    """
    Task C: Consolidate all chunk LLM outputs into final response
    
    Args:
        db: Database session
        ai_req_res_id: Primary key ID from bnk_ai_req_res table
        file_type: Document type for response formatting
    
    Returns:
        Dict with consolidated response
    """
    # Get all chunk outputs for this request
    parsed_docs = db.query(BnkParsedDocument).filter(
        BnkParsedDocument.ai_req_res_id == ai_req_res_id
    ).all()
    
    if not parsed_docs:
        return {"success": False, "error": f"No parsed documents for ai_req_res_id={ai_req_res_id}"}
    
    all_outputs = []
    
    for parsed_doc in parsed_docs:
        chunks = db.query(BnkDocumentChunk).filter(
            BnkDocumentChunk.parsed_document_id == parsed_doc.id
        ).order_by(BnkDocumentChunk.chunk_index).all()
        
        logger.info(f"[Task C] Found {len(chunks)} chunks for parsed_doc_id={parsed_doc.id}")
        
        for chunk in chunks:
            logger.info(f"[Task C] Chunk {chunk.id}: llm_output exists={chunk.llm_output is not None}, type={type(chunk.llm_output)}")
            
            if chunk.llm_output:
                # Log the actual output for debugging
                logger.info(f"[Task C] Chunk {chunk.id} output keys: {chunk.llm_output.keys() if isinstance(chunk.llm_output, dict) else 'not dict'}")
                
                # Include only successful responses and jailbreak errors (user-fixable)
                # Skip API errors, system errors, exceptions (these get retried)
                if isinstance(chunk.llm_output, dict):
                    has_error = chunk.llm_output.get("error", False)
                    error_type = chunk.llm_output.get("error_type")
                    
                    logger.info(f"[Task C] Chunk {chunk.id}: has_error={has_error}, error_type={error_type}")
                    
                    # Skip if error but NOT jailbreak
                    if has_error and error_type != "content_filter":
                        logger.info(f"[Task C] Skipping chunk {chunk.id} - non-jailbreak error")
                        continue
                
                logger.info(f"[Task C] Including chunk {chunk.id} in outputs")
                all_outputs.append({
                    "chunk_id": chunk.id,
                    "chunk_index": chunk.chunk_index,
                    "source_url": parsed_doc.source_url,
                    "output": chunk.llm_output
                })
    
    logger.info(f"[Task C] Total outputs collected: {len(all_outputs)}")
    
    if not all_outputs:
        return {"success": False, "error": "No valid LLM outputs to consolidate"}
    
    # Consolidate - merge chunk outputs while preserving the expected response format
    if len(all_outputs) == 1:
        raw_response = all_outputs[0]["output"]
        source_urls = [all_outputs[0].get("source_url", "")]
        logger.info("[Task C] Single chunk - no consolidation needed")
    else:
        # Multiple chunks: merge analysis_results from all chunks
        # Use the first chunk as base and merge analysis_results
        base_output = all_outputs[0]["output"]
        
        # Collect all analysis_results from all chunks
        merged_analysis_results = []
        seen_fields = set()
        
        for chunk_data in all_outputs:
            chunk_output = chunk_data.get("output", {})
            if isinstance(chunk_output, dict):
                analysis_results = chunk_output.get("analysis_results", [])
                for result in analysis_results:
                    field_name = result.get("field_name", "")
                    # Avoid duplicates - keep first occurrence or best match
                    if field_name not in seen_fields:
                        merged_analysis_results.append(result)
                        seen_fields.add(field_name)
                    else:
                        # If this chunk has a match and previous didn't, prefer this one
                        if result.get("is_match", False):
                            # Replace the existing one
                            for i, existing in enumerate(merged_analysis_results):
                                if existing.get("field_name") == field_name:
                                    merged_analysis_results[i] = result
                                    break
        
        # Build raw response preserving the expected format
        raw_response = base_output.copy() if isinstance(base_output, dict) else {}
        raw_response["analysis_results"] = merged_analysis_results
        
        # Recalculate overall_match based on merged results
        all_matched = all(r.get("is_match", False) for r in merged_analysis_results) if merged_analysis_results else False
        raw_response["overall_match"] = all_matched
        
        # Collect source URLs
        source_urls = list(set(c.get("source_url", "") for c in all_outputs if c.get("source_url")))
        
        logger.info(f"[Task C] Merged {len(all_outputs)} chunk outputs, {len(merged_analysis_results)} analysis fields")
    
    # Format using controller's response builder (same pipeline as original flow)
    file_url = ",".join(source_urls) if source_urls else ""
    formatted_response = build_response_from_openai_result(raw_response, file_type, file_url)
    
    # Convert Pydantic model to dict for storage
    consolidated_response = formatted_response.dict() if hasattr(formatted_response, 'dict') else formatted_response
    
    # Save consolidated output
    consolidated_record = BnkConsolidatedLLMOutput(
        ai_req_res_id=ai_req_res_id,
        consolidated_response=consolidated_response
    )
    db.add(consolidated_record)
    db.commit()
    logger.info(f"[Task C] Saved BnkConsolidatedLLMOutput")
    
    # Update BnkAIReqRes.response
    job = db.query(BnkAIReqRes).filter(BnkAIReqRes.id == ai_req_res_id).first()
    if job:
        job.response = consolidated_response
        db.commit()
        logger.info(f"[Task C] Updated BnkAIReqRes.response")
    
    return {
        "success": True,
        "ai_req_res_id": ai_req_res_id,
        "chunk_count": len(all_outputs),
        "consolidated_response": consolidated_response
    }
