import json
from collections.abc import AsyncIterator
from typing import Any

import httpx
from openai import AsyncOpenAI

import config
from tools.web_search import execute_tool, get_tools

API_TIMEOUT = httpx.Timeout(60.0, connect=10.0)
MAX_TOOL_ROUNDS = 2

def should_use_web_search(_user_text: str = "") -> bool:
    """用户已在设置中开启联网时，始终向模型提供搜索工具。"""
    return config.settings.web_search_enabled


def _assistant_message_dict(message: Any) -> dict[str, Any]:
    data: dict[str, Any] = {"role": "assistant", "content": message.content or ""}
    if message.tool_calls:
        data["tool_calls"] = [
            {
                "id": tc.id,
                "type": "function",
                "function": {
                    "name": tc.function.name,
                    "arguments": tc.function.arguments,
                },
            }
            for tc in message.tool_calls
        ]
    return data


class DeepSeekClient:
    def _get_client(self) -> AsyncOpenAI:
        return AsyncOpenAI(
            api_key=config.settings.deepseek_api_key,
            base_url="https://api.deepseek.com",
            timeout=API_TIMEOUT,
        )

    async def _run_tool_rounds(
        self,
        messages: list[dict[str, Any]],
        model: str | None = None,
        on_search: Any | None = None,
    ) -> tuple[list[dict[str, Any]], list[str], str | None]:
        client = self._get_client()
        tools = get_tools()
        search_queries: list[str] = []
        current_messages = list(messages)

        for _ in range(MAX_TOOL_ROUNDS):
            response = await client.chat.completions.create(
                model=model or config.settings.deepseek_model,
                messages=current_messages,
                tools=tools,
                tool_choice="auto",
                max_tokens=config.settings.max_tokens,
            )
            message = response.choices[0].message

            if message.tool_calls:
                current_messages.append(_assistant_message_dict(message))
                for tool_call in message.tool_calls:
                    fn = tool_call.function
                    if fn.name == "web_search":
                        try:
                            args = json.loads(fn.arguments or "{}")
                            q = (args.get("query") or "").strip()
                        except json.JSONDecodeError:
                            q = ""
                        if q:
                            search_queries.append(q)
                            if on_search:
                                await on_search(q)
                        result = await execute_tool(fn.name, fn.arguments)
                    else:
                        result = await execute_tool(fn.name, fn.arguments)
                    current_messages.append(
                        {
                            "role": "tool",
                            "tool_call_id": tool_call.id,
                            "content": result,
                        }
                    )
                continue

            if message.content:
                return current_messages, search_queries, message.content
            break

        return current_messages, search_queries, None

    async def stream_chat(
        self,
        messages: list[dict[str, Any]],
        model: str | None = None,
        use_tools: bool = False,
        on_search: Any | None = None,
    ) -> AsyncIterator[str]:
        if not config.settings.deepseek_api_key:
            raise RuntimeError("未配置 DEEPSEEK_API_KEY")

        working_messages = list(messages)
        if use_tools:
            working_messages, _, final_text = await self._run_tool_rounds(
                working_messages, model, on_search=on_search
            )
            if final_text:
                yield final_text
                return

        stream = await self._get_client().chat.completions.create(
            model=model or config.settings.deepseek_model,
            messages=working_messages,
            stream=True,
            max_tokens=config.settings.max_tokens,
        )
        async for chunk in stream:
            delta = chunk.choices[0].delta.content
            if delta:
                yield delta

    async def chat(
        self,
        messages: list[dict[str, Any]],
        model: str | None = None,
        use_tools: bool = False,
    ) -> str:
        if not config.settings.deepseek_api_key:
            raise RuntimeError("未配置 DEEPSEEK_API_KEY")

        working_messages = list(messages)
        if use_tools:
            working_messages, _, final_text = await self._run_tool_rounds(
                working_messages, model
            )
            if final_text:
                return final_text

        response = await self._get_client().chat.completions.create(
            model=model or config.settings.deepseek_model,
            messages=working_messages,
            stream=False,
            max_tokens=config.settings.max_tokens,
        )
        return response.choices[0].message.content or ""
