from typing import Any, cast

from pydantic import BaseModel, Field
from sqlalchemy import select

from core.app.app_config.entities import DatasetRetrieveConfigEntity, ModelConfig
from core.rag.datasource.retrieval_service import DefaultRetrievalModelDict, RetrievalService
from core.rag.entities import DocumentContext, RetrievalSourceMetadata
from core.rag.index_processor.constant.index_type import IndexTechniqueType
from core.rag.models.document import Document as RetrievalDocument
from core.rag.retrieval.dataset_retrieval import DatasetRetrieval
from core.rag.retrieval.retrieval_methods import RetrievalMethod
from core.tools.utils.dataset_retriever.dataset_retriever_base_tool import DatasetRetrieverBaseTool
from extensions.ext_database import db
from models.dataset import Dataset
from models.dataset import Document as DatasetDocument
from services.external_knowledge_service import ExternalDatasetService

default_retrieval_model: DefaultRetrievalModelDict = {
    "search_method": RetrievalMethod.SEMANTIC_SEARCH,
    "reranking_enable": False,
    "reranking_model": {"reranking_provider_name": "", "reranking_model_name": ""},
    "reranking_mode": "reranking_model",
    "top_k": 2,
    "score_threshold_enabled": False,
}


class DatasetRetrieverToolInput(BaseModel):
    query: str = Field(..., description="Query for the dataset to be used to retrieve the dataset.")


class DatasetRetrieverTool(DatasetRetrieverBaseTool):
    """Tool for querying a Dataset."""

    name: str = "dataset"
    args_schema: type[BaseModel] = DatasetRetrieverToolInput
    description: str = "use this to retrieve a dataset. "
    dataset_id: str
    user_id: str | None = None
    retrieve_config: DatasetRetrieveConfigEntity
    inputs: dict[str, Any]

    @classmethod
    def from_dataset(cls, dataset: Dataset, **kwargs):
        description = dataset.description
        if not description:
            description = "useful for when you want to answer queries about the " + dataset.name

        description = description.replace("\n", "").replace("\r", "")
        return cls(
            name=f"dataset_{dataset.id.replace('-', '_')}",
            tenant_id=dataset.tenant_id,
            dataset_id=dataset.id,
            description=description,
            **kwargs,
        )

    def _run(self, query: str) -> str:
        dataset_stmt = select(Dataset).where(Dataset.tenant_id == self.tenant_id, Dataset.id == self.dataset_id)
        dataset = db.session.scalar(dataset_stmt)

        if not dataset:
            return ""
        for hit_callback in self.hit_callbacks:
            hit_callback.on_query(query, dataset.id)
        dataset_retrieval = DatasetRetrieval()
        metadata_filter_document_ids, metadata_condition = dataset_retrieval.get_metadata_filter_condition(
            [dataset.id],
            query,
            self.tenant_id,
            self.user_id or "unknown",
            cast(str, self.retrieve_config.metadata_filtering_mode),
            cast(ModelConfig, self.retrieve_config.metadata_model_config),
            self.retrieve_config.metadata_filtering_conditions,
            self.inputs,
        )
        if metadata_filter_document_ids:
            document_ids_filter = metadata_filter_document_ids.get(dataset.id, [])
        else:
            document_ids_filter = None
        if dataset.provider == "external":
            results: list[RetrievalDocument] = []
            external_documents = ExternalDatasetService.fetch_external_knowledge_retrieval(
                tenant_id=dataset.tenant_id,
                dataset_id=dataset.id,
                query=query,
                external_retrieval_parameters=dataset.retrieval_model,
                metadata_condition=metadata_condition,
            )
            for external_document in external_documents:
                document = RetrievalDocument(
                    page_content=external_document.get("content"),
                    metadata=external_document.get("metadata"),
                    provider="external",
                )
                if document.metadata is not None:
                    document.metadata["score"] = external_document.get("score")
                    document.metadata["title"] = external_document.get("title")
                    document.metadata["dataset_id"] = dataset.id
                    document.metadata["dataset_name"] = dataset.name
                    results.append(document)
            # deal with external documents
            context_list: list[RetrievalSourceMetadata] = []
            for position, item in enumerate(results, start=1):
                if item.metadata is not None:
                    source = RetrievalSourceMetadata(
                        position=position,
                        dataset_id=item.metadata.get("dataset_id"),
                        dataset_name=item.metadata.get("dataset_name"),
                        document_id=item.metadata.get("document_id") or item.metadata.get("title"),
                        document_name=item.metadata.get("title"),
                        data_source_type="external",
                        retriever_from=self.retriever_from,
                        score=item.metadata.get("score"),
                        title=item.metadata.get("title"),
                        content=item.page_content,
                    )
                    context_list.append(source)
            for hit_callback in self.hit_callbacks:
                hit_callback.return_retriever_resource_info(context_list)

            return str("\n".join([item.page_content for item in results]))
        else:
            if metadata_condition and not document_ids_filter:
                return ""
            # get retrieval model , if the model is not setting , using default
            retrieval_model = dataset.retrieval_model or default_retrieval_model
            retrieval_resource_list: list[RetrievalSourceMetadata] = []
            if dataset.indexing_technique == IndexTechniqueType.ECONOMY:
                # use keyword table query
                documents = RetrievalService.retrieve(
                    retrieval_method=RetrievalMethod.KEYWORD_SEARCH,
                    dataset_id=dataset.id,
                    query=query,
                    top_k=self.top_k,
                    document_ids_filter=document_ids_filter,
                )
                return str("\n".join([document.page_content for document in documents]))
            else:
                if self.top_k > 0:
                    # retrieval source
                    documents = RetrievalService.retrieve(
                        retrieval_method=retrieval_model.get("search_method", "semantic_search"),
                        dataset_id=dataset.id,
                        query=query,
                        top_k=self.top_k,
                        score_threshold=retrieval_model.get("score_threshold", 0.0)
                        if retrieval_model["score_threshold_enabled"]
                        else 0.0,
                        reranking_model=retrieval_model.get("reranking_model")
                        if retrieval_model["reranking_enable"]
                        else None,
                        reranking_mode=retrieval_model.get("reranking_mode") or "reranking_model",
                        weights=retrieval_model.get("weights"),
                        document_ids_filter=document_ids_filter,
                    )
                else:
                    documents = []
                for hit_callback in self.hit_callbacks:
                    hit_callback.on_tool_end(documents)
                document_score_list = {}
                if dataset.indexing_technique != IndexTechniqueType.ECONOMY:
                    for item in documents:
                        if item.metadata is not None and item.metadata.get("score"):
                            document_score_list[item.metadata["doc_id"]] = item.metadata["score"]
                document_context_list: list[DocumentContext] = []
                records = RetrievalService.format_retrieval_documents(documents)
                if records:
                    for record in records:
                        segment = record.segment
                        # Build content: if summary exists, add it before the segment content
                        if segment.answer:
                            segment_content = f"question:{segment.get_sign_content()} answer:{segment.answer}"
                        else:
                            segment_content = segment.get_sign_content()

                        # If summary exists, prepend it to the content
                        if record.summary:
                            final_content = f"{record.summary}\n{segment_content}"
                        else:
                            final_content = segment_content

                        document_context_list.append(
                            DocumentContext(
                                content=final_content,
                                score=record.score,
                            )
                        )

                    if self.return_resource:
                        for record in records:
                            segment = record.segment
                            dataset = db.session.get(Dataset, segment.dataset_id)
                            dataset_document_stmt = select(DatasetDocument).where(
                                DatasetDocument.id == segment.document_id,
                                DatasetDocument.enabled == True,
                                DatasetDocument.archived == False,
                            )
                            document = db.session.scalar(dataset_document_stmt)
                            if dataset and document:
                                source = RetrievalSourceMetadata(
                                    dataset_id=dataset.id,
                                    dataset_name=dataset.name,
                                    document_id=document.id,
                                    document_name=document.name,
                                    data_source_type=document.data_source_type,
                                    segment_id=segment.id,
                                    retriever_from=self.retriever_from,
                                    score=record.score or 0.0,
                                    doc_metadata=document.doc_metadata,
                                )

                                if self.retriever_from == "dev":
                                    source.hit_count = segment.hit_count
                                    source.word_count = segment.word_count
                                    source.segment_position = segment.position
                                    source.index_node_hash = segment.index_node_hash
                                if segment.answer:
                                    source.content = f"question:{segment.content} \nanswer:{segment.answer}"
                                else:
                                    source.content = segment.content
                                # Add summary if this segment was retrieved via summary
                                if hasattr(record, "summary") and record.summary:
                                    source.summary = record.summary
                                retrieval_resource_list.append(source)

            if self.return_resource and retrieval_resource_list:
                retrieval_resource_list = sorted(
                    retrieval_resource_list,
                    key=lambda x: x.score or 0.0,
                    reverse=True,
                )
                for position, item in enumerate(retrieval_resource_list, start=1):  # type: ignore
                    item.position = position  # type: ignore
                for hit_callback in self.hit_callbacks:
                    hit_callback.return_retriever_resource_info(retrieval_resource_list)
            if document_context_list:
                document_context_list = sorted(document_context_list, key=lambda x: x.score or 0.0, reverse=True)
                return str("\n".join([document_context.content for document_context in document_context_list]))
            return ""
