import logging
import time

import click
from celery import shared_task
from sqlalchemy import delete, select, update

from core.db.session_factory import session_factory
from core.rag.index_processor.constant.doc_type import DocType
from core.rag.index_processor.constant.index_type import IndexStructureType
from core.rag.index_processor.index_processor_factory import IndexProcessorFactory
from core.rag.models.document import AttachmentDocument, ChildDocument, Document
from extensions.ext_redis import redis_client
from libs.datetime_utils import naive_utc_now
from models.dataset import DatasetAutoDisableLog, DocumentSegment
from models.dataset import Document as DatasetDocument
from models.enums import IndexingStatus, SegmentStatus

logger = logging.getLogger(__name__)


@shared_task(queue="dataset")
def add_document_to_index_task(dataset_document_id: str):
    """
    Async Add document to index
    :param dataset_document_id:

    Usage: add_document_to_index_task.delay(dataset_document_id)
    """
    logger.info(click.style(f"Start add document to index: {dataset_document_id}", fg="green"))
    start_at = time.perf_counter()

    with session_factory.create_session() as session:
        dataset_document = session.scalar(
            select(DatasetDocument).where(DatasetDocument.id == dataset_document_id).limit(1)
        )
        if not dataset_document:
            logger.info(click.style(f"Document not found: {dataset_document_id}", fg="red"))
            return

        if dataset_document.indexing_status != IndexingStatus.COMPLETED:
            return

        indexing_cache_key = f"document_{dataset_document.id}_indexing"

        try:
            dataset = dataset_document.dataset
            if not dataset:
                raise Exception(f"Document {dataset_document.id} dataset {dataset_document.dataset_id} doesn't exist.")

            segments = session.scalars(
                select(DocumentSegment)
                .where(
                    DocumentSegment.document_id == dataset_document.id,
                    DocumentSegment.status == SegmentStatus.COMPLETED,
                )
                .order_by(DocumentSegment.position.asc())
            ).all()

            documents = []
            multimodal_documents = []
            for segment in segments:
                document = Document(
                    page_content=segment.content,
                    metadata={
                        "doc_id": segment.index_node_id,
                        "doc_hash": segment.index_node_hash,
                        "document_id": segment.document_id,
                        "dataset_id": segment.dataset_id,
                    },
                )
                if dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX:
                    child_chunks = segment.get_child_chunks()
                    if child_chunks:
                        child_documents = []
                        for child_chunk in child_chunks:
                            child_document = ChildDocument(
                                page_content=child_chunk.content,
                                metadata={
                                    "doc_id": child_chunk.index_node_id,
                                    "doc_hash": child_chunk.index_node_hash,
                                    "document_id": segment.document_id,
                                    "dataset_id": segment.dataset_id,
                                },
                            )
                            child_documents.append(child_document)
                        document.children = child_documents
                if dataset.is_multimodal:
                    for attachment in segment.attachments:
                        multimodal_documents.append(
                            AttachmentDocument(
                                page_content=attachment["name"],
                                metadata={
                                    "doc_id": attachment["id"],
                                    "doc_hash": "",
                                    "document_id": segment.document_id,
                                    "dataset_id": segment.dataset_id,
                                    "doc_type": DocType.IMAGE,
                                },
                            )
                        )
                documents.append(document)

            index_type = dataset.doc_form
            index_processor = IndexProcessorFactory(index_type).init_index_processor()
            index_processor.load(dataset, documents, multimodal_documents=multimodal_documents)

            # delete auto disable log
            session.execute(
                delete(DatasetAutoDisableLog).where(DatasetAutoDisableLog.document_id == dataset_document.id)
            )

            # update segment to enable
            session.execute(
                update(DocumentSegment)
                .where(DocumentSegment.document_id == dataset_document.id)
                .values(enabled=True, disabled_at=None, disabled_by=None, updated_at=naive_utc_now())
            )
            session.commit()

            # Enable summary indexes for all segments in this document
            from services.summary_index_service import SummaryIndexService

            segment_ids_list = [segment.id for segment in segments]
            if segment_ids_list:
                try:
                    SummaryIndexService.enable_summaries_for_segments(
                        dataset=dataset,
                        segment_ids=segment_ids_list,
                    )
                except Exception as e:
                    logger.warning("Failed to enable summaries for document %s: %s", dataset_document.id, str(e))

            end_at = time.perf_counter()
            logger.info(
                click.style(f"Document added to index: {dataset_document.id} latency: {end_at - start_at}", fg="green")
            )
        except Exception as e:
            logger.exception("add document to index failed")
            dataset_document.enabled = False
            dataset_document.disabled_at = naive_utc_now()
            dataset_document.indexing_status = IndexingStatus.ERROR
            dataset_document.error = str(e)
            session.commit()
        finally:
            redis_client.delete(indexing_cache_key)
