from __future__ import annotations

from typing import Any

from app.agents.base import AgentExecutionContext
from app.agents.personas import get_persona

from .models import VALID_TASK_TYPES, RouteResult


class OrchestratorRoutingMixin:
    """Task routing and execution logic across specialist agents."""

    async def route(
        self,
        task_type: str,
        query: str,
        top_k: int = 3,
        agent_code: str | None = None,
        retrieval_filters: dict[str, Any] | None = None,
    ) -> RouteResult:
        """Route a task to the matching logical agent."""

        resolved_task_type = self._resolve_task_type(
            task_type=task_type, agent_code=agent_code
        )
        agent = self._agents[resolved_task_type]
        persona = get_persona(agent_code or "minh")
        effective_filters = self._merge_retrieval_filters(
            task_type=resolved_task_type,
            agent_code=persona["agent_code"],
            retrieval_filters=retrieval_filters,
        )
        retrieved_chunks = await self._retrieve(
            task_type=resolved_task_type,
            query=query,
            top_k=top_k,
            retrieval_filters=effective_filters,
        )
        context = AgentExecutionContext(
            query=query,
            task_type=resolved_task_type,
            retrieved_chunks=retrieved_chunks,
            agent_code=persona["agent_code"],
            persona_name=persona["display_name"],
            persona_profile=persona,
        )
        answer = await agent.run(context)
        return RouteResult(
            selected_agent=agent.agent_name,
            answer=answer,
            retrieved_chunks_count=len(retrieved_chunks),
            task_type=resolved_task_type,
            agent_code=persona["agent_code"],
        )

    async def _retrieve(
        self,
        task_type: str,
        query: str,
        top_k: int,
        retrieval_filters: dict[str, Any] | None = None,
    ) -> list[dict[str, Any]]:
        """Retrieve relevant knowledge using Hybrid Search (BM25 + Dense Qdrant + RRF)."""

        filters = retrieval_filters or {"task_type": task_type}
        try:
            return await self._hybrid_search_service.search(
                query=query,
                top_k=top_k,
                task_type=task_type,
                retrieval_filters=filters,
                mode="hybrid",
            )
        except Exception:
            return []

    def _resolve_task_type(self, task_type: str, agent_code: str | None) -> str:
        """Resolve task type from explicit input or persona defaults."""

        if task_type:
            self._validate_task_type(task_type)
            return task_type
        if agent_code:
            persona = get_persona(agent_code)
            default_task_type = persona["primary_capabilities"][0]
            self._validate_task_type(default_task_type)
            return default_task_type
        raise ValueError(
            "task_type is required when no agent_code default is available"
        )

    def _merge_retrieval_filters(
        self,
        task_type: str,
        agent_code: str | None,
        retrieval_filters: dict[str, Any] | None,
    ) -> dict[str, Any]:
        filters = {"task_type": task_type}
        for key, value in (retrieval_filters or {}).items():
            if value is None or value == "":
                continue
            filters[key] = value
        if (
            "agent_scope" not in filters
            and agent_code
            and (retrieval_filters and "agent_scope" in retrieval_filters)
        ):
            filters["agent_scope"] = agent_code
        return filters

    def _validate_task_type(self, task_type: str) -> None:
        """Raise a clear error for unsupported task types."""

        if task_type not in VALID_TASK_TYPES:
            expected = ", ".join(sorted(VALID_TASK_TYPES))
            raise ValueError(f"Unsupported task_type. Expected one of: {expected}")
