from typing import Dict, List, Optional, Any
from dataclasses import dataclass, field
import time

from core.config import get_api_key
from core.model_pool import ModelConfig, ThinkingMode, get_thinking_payload, get_config
from core.zai_client import ZAiClient, ZAiResponse


@dataclass
class ExecutionResult:
    success: bool
    content: str
    reasoning: Optional[str] = None
    model_used: str = ""
    latency_ms: int = 0
    tokens_used: int = 0
    error: Optional[str] = None
    metadata: Dict[str, Any] = field(default_factory=dict)


class ModelExecutor:
    """
    Executes model calls based on role configuration.
    Handles fallback, retry, and thinking mode.
    """

    def __init__(self, api_key: Optional[str] = None):
        self._api_key = api_key or get_api_key("zai")
        self._client: Optional[ZAiClient] = None

    @property
    def client(self) -> ZAiClient:
        if self._client is None:
            if not self._api_key:
                raise ValueError(
                    "No ZAI_API_KEY configured. Set ZAI_API_KEY env var or call set_api_key()."
                )
            self._client = ZAiClient(self._api_key)
        return self._client

    def execute(
        self,
        role: str,
        messages: List[Dict[str, str]],
        max_tokens: int = 4096,
        temperature: float = 0.7,
        **kwargs,
    ) -> ExecutionResult:
        """
        Execute a model call for a given role.

        Args:
            role: Role identifier (e.g., "router_small", "verifier_v1")
            messages: Chat messages
            max_tokens: Max output tokens
            temperature: Sampling temperature
        """
        config = get_config(role)
        if not config:
            return ExecutionResult(
                success=False, content="", error=f"Unknown role: {role}"
            )

        return self._execute_with_config(
            config, messages, max_tokens, temperature, **kwargs
        )

    def _execute_with_config(
        self,
        config: ModelConfig,
        messages: List[Dict[str, str]],
        max_tokens: int,
        temperature: float,
        **kwargs,
    ) -> ExecutionResult:
        """Execute with retry and fallback logic."""

        models_to_try = [config.model_id]
        if config.fallback_model_id:
            models_to_try.append(config.fallback_model_id)

        thinking = get_thinking_payload(config.thinking_mode)

        last_error = None

        for attempt in range(config.max_retries):
            for model_id in models_to_try:
                start_time = time.time()

                try:
                    response = self.client.chat_completion(
                        model=model_id,
                        messages=messages,
                        thinking=thinking,
                        max_tokens=max_tokens,
                        temperature=temperature,
                    )

                    latency_ms = int((time.time() - start_time) * 1000)

                    return ExecutionResult(
                        success=True,
                        content=response.content,
                        reasoning=response.reasoning_content,
                        model_used=response.model,
                        latency_ms=latency_ms,
                        tokens_used=response.usage.get("total_tokens", 0),
                        metadata={
                            "role": config.role,
                            "tier": config.tier.value,
                            "thinking": config.thinking_mode.value,
                        },
                    )

                except Exception as e:
                    last_error = str(e)
                    continue

        return ExecutionResult(
            success=False,
            content="",
            error=f"All attempts failed. Last error: {last_error}",
        )

    def execute_simple(
        self, role: str, prompt: str, system: Optional[str] = None, **kwargs
    ) -> ExecutionResult:
        """Convenience method for single-prompt execution."""
        messages = []
        if system:
            messages.append({"role": "system", "content": system})
        messages.append({"role": "user", "content": prompt})

        return self.execute(role, messages, **kwargs)

    def close(self):
        if self._client:
            self._client.close()
            self._client = None

    def __enter__(self):
        return self

    def __exit__(self, exc_type, exc_val, exc_tb):
        self.close()
