import os
import httpx
from typing import Dict, List, Optional, Any
from dataclasses import dataclass, field
import json


@dataclass
class ZAiMessage:
    role: str
    content: str


@dataclass
class ZAiResponse:
    content: str
    reasoning_content: Optional[str] = None
    model: str = ""
    usage: Dict[str, int] = field(default_factory=dict)
    finish_reason: str = ""


class ZAiClient:
    """
    Z.AI API Client for GLM models.
    Endpoint: https://api.z.ai/api/paas/v4/chat/completions
    """

    BASE_URL = "https://api.z.ai/api/paas/v4"

    def __init__(self, api_key: Optional[str] = None, timeout: float = 60.0):
        self.api_key = api_key or os.environ.get("ZAI_API_KEY")
        if not self.api_key:
            raise ValueError("ZAI_API_KEY not provided and not in environment")

        self.client = httpx.Client(timeout=timeout)

    def chat_completion(
        self,
        model: str,
        messages: List[Dict[str, str]],
        thinking: Optional[Dict[str, str]] = None,
        max_tokens: int = 4096,
        temperature: float = 0.7,
        stream: bool = False,
        **kwargs,
    ) -> ZAiResponse:
        """
        Make a chat completion request.

        Args:
            model: glm-4.7, glm-4.7-flash, glm-4.7-flashx, glm-5
            messages: List of {role, content} dicts
            thinking: {"type": "enabled"} or {"type": "disabled"}
            max_tokens: Maximum output tokens
            temperature: Sampling temperature
            stream: Enable streaming (not implemented in this method)
        """
        url = f"{self.BASE_URL}/chat/completions"

        payload = {
            "model": model,
            "messages": messages,
            "max_tokens": max_tokens,
            "temperature": temperature,
            "stream": stream,
        }

        if thinking:
            payload["thinking"] = thinking

        payload.update(kwargs)

        headers = {
            "Content-Type": "application/json",
            "Authorization": f"Bearer {self.api_key}",
        }

        response = self.client.post(url, json=payload, headers=headers)

        if response.status_code != 200:
            raise RuntimeError(
                f"Z.AI API error: {response.status_code} - {response.text}"
            )

        data = response.json()

        choice = data.get("choices", [{}])[0]
        message = choice.get("message", {})

        return ZAiResponse(
            content=message.get("content", ""),
            reasoning_content=message.get("reasoning_content"),
            model=data.get("model", model),
            usage=data.get("usage", {}),
            finish_reason=choice.get("finish_reason", ""),
        )

    def chat_stream(
        self,
        model: str,
        messages: List[Dict[str, str]],
        thinking: Optional[Dict[str, str]] = None,
        max_tokens: int = 4096,
        temperature: float = 0.7,
    ):
        """
        Streaming chat completion. Yields (reasoning_chunk, content_chunk) tuples.
        """
        url = f"{self.BASE_URL}/chat/completions"

        payload = {
            "model": model,
            "messages": messages,
            "max_tokens": max_tokens,
            "temperature": temperature,
            "stream": True,
        }

        if thinking:
            payload["thinking"] = thinking

        headers = {
            "Content-Type": "application/json",
            "Authorization": f"Bearer {self.api_key}",
        }

        with self.client.stream("POST", url, json=payload, headers=headers) as response:
            if response.status_code != 200:
                raise RuntimeError(f"Z.AI API error: {response.status_code}")

            for line in response.iter_lines():
                if not line or line == "data: [DONE]":
                    continue

                if line.startswith("data: "):
                    try:
                        data = json.loads(line[6:])
                        delta = data.get("choices", [{}])[0].get("delta", {})
                        reasoning = delta.get("reasoning_content", "")
                        content = delta.get("content", "")
                        yield reasoning, content
                    except json.JSONDecodeError:
                        continue

    def close(self):
        self.client.close()

    def __enter__(self):
        return self

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


MODEL_CAPABILITIES = {
    "glm-4.7-flash": {
        "thinking": False,
        "context": 200000,
        "max_output": 128000,
        "tier": "small",
        "description": "Lightweight, completely free",
    },
    "glm-4.7-flashx": {
        "thinking": False,
        "context": 200000,
        "max_output": 128000,
        "tier": "small",
        "description": "Lightweight, high-speed, affordable",
    },
    "glm-4.7": {
        "thinking": True,
        "context": 200000,
        "max_output": 128000,
        "tier": "hybrid",
        "description": "Full model with optional thinking mode",
    },
    "glm-5": {
        "thinking": True,
        "context": 200000,
        "max_output": 128000,
        "tier": "powerhouse",
        "description": "Flagship model for complex tasks",
    },
}
