#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
AICraft Web Search MCP Server
=============================
为 Codex++ / Codex CLI 提供真实的联网搜索能力。

背景:Codex 内置的 web_search 是"本地自定义工具",但执行器未实现,
每次调用都返回 `unsupported custom tool call: web_search`(卡死重试)。
本 MCP server 把搜索接回 AICraft 平台的博查搜索端点(/v1/tools/search),
让 Codex 通过 MCP 工具真正完成联网检索——与模型服务商无关。

协议:Model Context Protocol (MCP) over stdio,newline-delimited JSON-RPC 2.0。

用法(在 Codex config.toml 中注册):
    [mcp_servers.web_search]
    command = "python"
    args = ["C:/Users/dell/.codex/mcp/web_search_mcp.py"]
    env = { "PYTHONUTF8" = "1" }

环境变量(可选):
    SEARCH_API_URL  默认 https://aicraftapi.com/v1/tools/search
    SEARCH_API_KEY  默认取 AICRAFT_MCP_KEY,再取 ~/.codex-session-delete/settings.json 的 relayApiKey
"""
import json
import os
import sys
import time
import urllib.request
import urllib.error

# ── 配置 ──────────────────────────────────────────────────────────────
DEFAULT_API_URL = "https://aicraftapi.com/v1/tools/search"
LOG_PATH = os.path.join(os.path.dirname(os.path.abspath(__file__)), "web_search_mcp.log")
TOOL_NAME = "web_search"
TOOL_DESCRIPTION = (
    "实时联网搜索。用于查询最新资讯、实时行情、新闻、文档等需要联网获取的信息。"
    "输入自然语言查询词(建议含日期),返回带标题/链接/摘要的搜索结果列表。"
)
TOOL_SCHEMA = {
    "type": "object",
    "properties": {
        "query": {
            "type": "string",
            "description": "搜索查询词,建议包含日期与关键限定词,例如 '2026-08-12 A股 上证指数 收盘'",
        },
        "max_results": {
            "type": "integer",
            "description": "最多返回条数(1-8,默认5)",
            "minimum": 1,
            "maximum": 8,
        },
    },
    "required": ["query"],
}


def _log(msg: str) -> None:
    """追加一行日志,便于排障。"""
    try:
        with open(LOG_PATH, "a", encoding="utf-8") as f:
            f.write(f"{time.strftime('%Y-%m-%dT%H:%M:%S')} {msg}\n")
    except Exception:
        pass


def _load_api_key() -> str:
    """解析搜索 API key: 环境变量 > settings.json(顶层 relayApiKey > profile 级 > authContents)。"""
    for var in ("AICRAFT_MCP_KEY", "SEARCH_API_KEY"):
        v = os.environ.get(var, "").strip()
        if v:
            return v
    try:
        sp = os.path.join(os.path.dirname(os.path.abspath(__file__)),
                          "..", "..", ".codex-session-delete", "settings.json")
        sp = os.path.abspath(sp)
        with open(sp, encoding="utf-8") as f:
            data = json.load(f)
        # 1) 顶层 relayApiKey(新版 Codex++ 存这里)
        k = str(data.get("relayApiKey") or "").strip()
        if k:
            return k
        # 2) profile 级 relayApiKey
        for prof in data.get("relayProfiles", []):
            k = str(prof.get("relayApiKey") or "").strip()
            if k:
                return k
        # 3) profile authContents 里的 OPENAI_API_KEY
        for prof in data.get("relayProfiles", []):
            ac = prof.get("authContents") or ""
            if isinstance(ac, str):
                try:
                    ac = json.loads(ac)
                except Exception:
                    ac = {}
            k = str((ac or {}).get("OPENAI_API_KEY") or "").strip()
            if k:
                return k
    except Exception:
        pass
    return ""


def _clean(s: str) -> str:
    """去掉代理字符(不合法 Unicode),防止编码异常导致崩溃。"""
    try:
        return s.encode("utf-8", errors="ignore").decode("utf-8")
    except Exception:
        return ""


def _gbk_safe(s: str) -> str:
    """把文本过滤成 GBK 可编码字符(防止中文 Windows 中继解析失败,同 video_mcp)。"""
    try:
        return s.encode("gbk", errors="ignore").decode("gbk")
    except Exception:
        return s


def _search(query: str, max_results: int) -> dict:
    """调用 AICraft 博查搜索端点,返回结构化结果。"""
    url = os.environ.get("SEARCH_API_URL", DEFAULT_API_URL)
    key = _load_api_key()
    query = _clean(query)
    payload = json.dumps({"query": query, "max": max_results},
                         ensure_ascii=False).encode("utf-8", errors="ignore")
    headers = {"Content-Type": "application/json"}
    if key:
        headers["Authorization"] = "Bearer " + key
    req = urllib.request.Request(url, data=payload, headers=headers, method="POST")
    try:
        with urllib.request.urlopen(req, timeout=30) as resp:
            body = resp.read().decode("utf-8", errors="replace")
            return {"ok": True, "status": resp.status, "data": json.loads(body)}
    except urllib.error.HTTPError as e:
        err = e.read().decode("utf-8", errors="replace")
        return {"ok": False, "status": e.code, "error": err[:500]}
    except Exception as e:
        return {"ok": False, "status": None, "error": f"{type(e).__name__}: {e}"}


def _results_to_text(data: dict) -> str:
    """把搜索返回的 results 列表拼成给模型看的文本。"""
    results = data.get("results") or []
    provider = data.get("provider", "?")
    cached = data.get("cached", False)
    if not results:
        return f"[Web search · {provider} · cached={cached}] 未找到相关结果。"
    lines = [f"[Web search · {provider} · cached={cached} · {len(results)} 条结果]"]
    for i, r in enumerate(results, 1):
        title = r.get("title", "").strip()
        url = r.get("url", "").strip()
        snippet = (r.get("snippet") or r.get("summary") or "").strip()
        date = r.get("date", "")
        lines.append(f"\n{i}. {title}\n   URL: {url}")
        if date:
            lines.append(f"   日期: {date}")
        if snippet:
            lines.append(f"   摘要: {snippet[:400]}")
    return "\n".join(lines)


# ── MCP 协议处理 ─────────────────────────────────────────────────────
def _rpc(id_, result=None, error=None) -> str:
    msg = {"jsonrpc": "2.0", "id": id_}
    if error is not None:
        msg["error"] = error
    else:
        msg["result"] = result
    return json.dumps(msg, ensure_ascii=False)


def main() -> None:
    # Windows 下强制 UTF-8 输出,避免中文乱码
    if sys.platform == "win32":
        sys.stdout.reconfigure(encoding="utf-8")
        sys.stderr.reconfigure(encoding="utf-8")
    os.environ.setdefault("PYTHONUTF8", "1")

    _log("MCP server 启动")
    initialized = False
    for raw in sys.stdin:
        raw = raw.strip()
        if not raw:
            continue
        try:
            msg = json.loads(raw)
        except json.JSONDecodeError:
            _log(f"非法 JSON: {raw[:200]}")
            continue

        method = msg.get("method")
        mid = msg.get("id")
        params = msg.get("params") or {}

        if mid is None:
            # 通知类消息
            if method == "notifications/initialized":
                initialized = True
            elif method == "notifications/cancelled":
                pass
            continue

        if method == "initialize":
            protocol = params.get("protocolVersion", "2024-11-05")
            out = {
                "protocolVersion": protocol,
                "capabilities": {"tools": {}},
                "serverInfo": {"name": "aicraft-web-search", "version": "1.0.0"},
            }
            print(_rpc(mid, out), flush=True)
        elif method == "ping":
            print(_rpc(mid, {}), flush=True)
        elif method == "tools/list":
            out = {
                "tools": [{
                    "name": TOOL_NAME,
                    "description": TOOL_DESCRIPTION,
                    "inputSchema": TOOL_SCHEMA,
                }]
            }
            print(_rpc(mid, out), flush=True)
        elif method == "tools/call":
            name = params.get("name", "")
            args = params.get("arguments") or {}
            if name != TOOL_NAME:
                print(_rpc(mid, error={
                    "code": -32602, "message": f"未知工具: {name}"}), flush=True)
                continue
            query = str(args.get("query", "")).strip()
            try:
                max_results = max(1, min(8, int(args.get("max_results", 5))))
            except (TypeError, ValueError):
                max_results = 5
            if not query:
                print(_rpc(mid, error={
                    "code": -32602, "message": "query 不能为空"}), flush=True)
                continue
            _log(f"tools/call web_search query={query[:60]} max={max_results}")
            res = _search(query, max_results)
            if not res.get("ok"):
                text = (f"[Web search 失败 · HTTP {res.get('status')}] "
                        f"{res.get('error', '未知错误')}")
                _log(f"  失败: {res.get('status')} {res.get('error', '')[:120]}")
            else:
                text = _results_to_text(res.get("data") or {})
                _log(f"  成功: {len(text)} 字符")
            print(_rpc(mid, {"content": [{"type": "text", "text": _gbk_safe(text)}]}), flush=True)
        else:
            print(_rpc(mid, error={
                "code": -32601, "message": f"未知方法: {method}"}), flush=True)


if __name__ == "__main__":
    main()
