import asyncio
from langchain_core.documents import Document
from langchain_openai import OpenAIEmbeddings
from app.core.config import settings
from app.services.pinecone_service import hotel_pinecone_service, hotel_private_pinecone_service


class HotelRAGService:
    def _get_embeddings(self) -> OpenAIEmbeddings:
        if not settings.OPENAI_API_KEY:
            raise ValueError("OpenAI API key is not configured.")

        return OpenAIEmbeddings(
            model=settings.EMBEDDING_MODEL,
            openai_api_key=settings.OPENAI_API_KEY,
            dimensions=512,
        )

    def _available_indexes(self) -> list:
        indexes = []
        for service in (hotel_pinecone_service, hotel_private_pinecone_service):
            if service.pc and service.index_name:
                indexes.append(service.get_index())
        if not indexes:
            raise ValueError(
                "No hotel Pinecone indexes configured. "
                "Set message-center and/or private index credentials."
            )
        return indexes

    async def _safe_retrieve_batch(self, queries: list[str], k: int = 4) -> list[Document]:
        """Batch-embed queries, then search all configured hotel Pinecone indexes."""
        cleaned_queries = [q.strip() for q in queries if isinstance(q, str) and q.strip()]
        if not cleaned_queries:
            return []

        indexes = self._available_indexes()
        embeddings = self._get_embeddings()
        query_vectors = await embeddings.aembed_documents(cleaned_queries)

        async def _query_index(index, vector: list[float]):
            return await asyncio.to_thread(
                index.query,
                vector=vector,
                top_k=k,
                include_metadata=True,
                include_values=False,
            )

        tasks = [
            _query_index(index, vector)
            for vector in query_vectors
            for index in indexes
        ]
        results = await asyncio.gather(*tasks)

        documents: list[Document] = []
        seen_keys: set[tuple[str, str]] = set()

        for result in results:
            matches = getattr(result, "matches", []) or []
            for match in matches:
                metadata = getattr(match, "metadata", None) or {}
                content = metadata.get("content") or metadata.get("text")
                metadata.pop("content", None)
                if not content:
                    continue

                key = (content, metadata.get("source") or metadata.get("url") or "")
                if key in seen_keys:
                    continue

                seen_keys.add(key)
                documents.append(Document(page_content=content, metadata=metadata))

        return documents

    def format_documents(self, docs: list[Document]) -> str:
        multiline_string = ""

        for doc in docs:
            metadata: dict = doc.metadata or {}

            title = metadata.get("title", "N/A")
            content = (doc.page_content or "").strip()
            url = metadata.get("source", None) or metadata.get("url", None)

            multiline_string += (
                "\n<source>\n"
                f"<title>{title}</title>\n"
                f"<content>{content}</content>\n"
                f"<url>{url}</url>\n"
                "</source>\n"
            )

        return multiline_string


def get_hotel_rag_service():
    return HotelRAGService()
