#!/usr/bin/env python3
"""
extract_swagger.py — 本地 Swagger 文档提取与解析（与服务端 parseApiDocs 一比一对齐）

设计目标：
- 一步完成"下载 → 解析 → 按 tag/path 过滤 → 关联 DTO definitions → 输出"
- 避免 AI 在客户端用 curl|head 截断 JSON / 写临时 python 脚本踩花括号坑
- 输出结构与服务端 api_tool(action="parse_swagger_json") 完全一致，可平替

参考实现：szcd-mcp-server/src/lib/swagger.js 的 parseApiDocs + server.js 的 filterApis
        递归深度 5，两层 definitions（summary + detail），引用树传递闭包裁剪。

用法：
  # 方式 A：脚本内部 curl 下载（自动缓存到 /tmp/swagger-cache/，TTL 10 分钟）
  extract_swagger.py --url http://10.x.x.x/svc/v2/api-docs [--cookie 'JSESSIONID=xxx']
  extract_swagger.py --url ... --no-cache         # 强制重新下载
  extract_swagger.py --url ... --cache-ttl 3600   # 自定义 TTL（秒），0 关闭缓存

  # 方式 B：读取本地已下载的 JSON 文件
  extract_swagger.py --file /tmp/swagger.json

  # 过滤（对 url/summary/operationId/tag/description/dtoRef/responseVoName 模糊匹配，大小写不敏感）
  extract_swagger.py --file /tmp/swagger.json --filter tag
  extract_swagger.py --file /tmp/swagger.json --filter tag --filter user

  # 仅查询单个 definition（兜底单 VO 查询，等价于服务端 get_definition）
  extract_swagger.py --file /tmp/swagger.json --definition KnowledgeOperLogVO

  # 输出到文件而非 stdout
  extract_swagger.py --file /tmp/swagger.json --filter tag --out /tmp/parsed.json

容错：
  - 自动修复八进制字面量（如 [011,015] → [11,15]），无需手动预处理
  - JSON 解析失败给出错误位置上下文 + 常见原因提示（HTML 跳转/截断/八进制）

跨平台：
  - Linux/macOS：`python3 extract_swagger.py ...`（脚本 shebang 已带 #!/usr/bin/env python3）
  - Windows：`python extract_swagger.py ...`（python3 命令可能不可用）

退出码：0 成功；1 参数错误；2 网络/IO 错误；3 解析错误
"""

import argparse
import hashlib
import json
import os
import re
import sys
import subprocess
import tempfile
import time
from pathlib import Path

MAX_DEPTH = 5
METHOD_KEYS = ["post", "get", "put", "delete", "patch"]

# 缓存配置（与 Node 版 extract-swagger.mjs 完全一致：共用同一目录 + sha256 key）
CACHE_DIR = Path(tempfile.gettempdir()) / "swagger-cache"
DEFAULT_CACHE_TTL = 600  # 10 分钟

# 八进制字面量：JSON 标准不允许数字以 0 开头（除 0 本身或小数），
# 但部分 Java/旧网关序列化会产出 [011,015,022] 这类数组。
# 匹配 JSON value 位置上的非法前导零数字（不含小数、不含 0 单独存在）。
OCTAL_LITERAL_RE = re.compile(r'(?<=[\s,\[:])0(\d+)(?=[\s,\]}])')


# ==================== 工具函数 ====================

def _ref_name(prop):
    """从 property 定义中提取 $ref/originalRef 名称（不含数组 items 层）。"""
    if not isinstance(prop, dict):
        return None
    if prop.get("originalRef"):
        return prop["originalRef"]
    ref = prop.get("$ref")
    if ref:
        return ref.rsplit("/", 1)[-1]
    return None


def _ref_name_with_items(prop):
    """提取 $ref 名称，包含数组 items 层（对齐 JS 端 expandProperties 第 250-251 行）。"""
    if not isinstance(prop, dict):
        return None
    name = _ref_name(prop)
    if name:
        return name
    items = prop.get("items")
    if isinstance(items, dict):
        return _ref_name(items)
    return None


def _resolve_property_type(prop):
    """对齐 JS 端 resolvePropertyType：数组返回 array<XXX>，否则返回 type 或 $ref 名。"""
    if not isinstance(prop, dict):
        return "string"
    if prop.get("type") == "array":
        items = prop.get("items") or {}
        item_type = (
            items.get("type")
            or items.get("originalRef")
            or (items.get("$ref", "").rsplit("/", 1)[-1] if items.get("$ref") else None)
            or "object"
        )
        return f"array<{item_type}>"
    if prop.get("type"):
        return prop["type"]
    if prop.get("originalRef"):
        return prop["originalRef"]
    if prop.get("$ref"):
        return prop["$ref"].rsplit("/", 1)[-1]
    return "object"


# ==================== 递归展开（对齐 JS 端 expandProperties，深度 5） ====================

def expand_properties(properties, definitions, depth=0, max_depth=MAX_DEPTH):
    if not properties:
        return []
    result = []
    for name, prop in properties.items():
        if not isinstance(prop, dict):
            prop = {}
        entry = {
            "name": name,
            "type": _resolve_property_type(prop),
            "description": prop.get("description", ""),
        }
        if "format" in prop:
            entry["format"] = prop["format"]
        if "example" in prop:
            entry["example"] = prop["example"]
        if "enum" in prop:
            entry["enum"] = prop["enum"]

        ref_name = _ref_name_with_items(prop)
        if ref_name:
            entry["refName"] = ref_name

        # 数组类型：展开 items 的嵌套属性
        if prop.get("type") == "array" and ref_name and depth < max_depth:
            item_def = (definitions or {}).get(ref_name)
            if isinstance(item_def, dict) and item_def.get("properties"):
                entry["itemProperties"] = expand_properties(
                    item_def["properties"], definitions, depth + 1, max_depth
                )

        # 对象类型（非数组）：展开嵌套属性
        if prop.get("type") == "object" and ref_name and depth < max_depth:
            nested_def = (definitions or {}).get(ref_name)
            if isinstance(nested_def, dict) and nested_def.get("properties"):
                entry["nestedProperties"] = expand_properties(
                    nested_def["properties"], definitions, depth + 1, max_depth
                )

        # 无 $ref 的 object 且有自身 properties（内联定义）
        if (
            prop.get("type") == "object"
            and not ref_name
            and prop.get("properties")
            and depth < max_depth
        ):
            entry["nestedProperties"] = expand_properties(
                prop["properties"], definitions, depth + 1, max_depth
            )

        result.append(entry)
    return result


# ==================== 主解析（对齐 JS 端 parseApiDocs） ====================

def parse_api_docs(api_docs, referenced_defs=None):
    paths = api_docs.get("paths") or {}
    tags = api_docs.get("tags") or []
    definitions = api_docs.get("definitions") or {}
    info = api_docs.get("info") or {}
    api_base_path = api_docs.get("basePath") or ""

    parsed_tags = [
        {"name": t.get("name"), "description": t.get("description", "")} for t in tags
    ]

    parsed_apis = []
    auto_referenced_defs = set()

    for url, path_value in paths.items():
        if not isinstance(path_value, dict):
            continue
        for method in METHOD_KEYS:
            api_def = path_value.get(method)
            if not isinstance(api_def, dict):
                continue
            tag = (api_def.get("tags") or ["default"])[0]
            parameters = api_def.get("parameters") or []

            params = []
            for p in parameters:
                schema = p.get("schema") or {}
                param_info = {
                    "name": p.get("name"),
                    "in": p.get("in"),
                    "required": p.get("required", False),
                    "type": p.get("type") or schema.get("type") or "string",
                    "description": p.get("description", ""),
                }
                # body 参数引用的 DTO
                ref_name = schema.get("originalRef") or (
                    schema.get("$ref", "").rsplit("/", 1)[-1] if schema.get("$ref") else None
                )
                if ref_name:
                    param_info["dtoRef"] = ref_name
                    auto_referenced_defs.add(ref_name)
                    dto = definitions.get(ref_name)
                    if isinstance(dto, dict) and dto.get("properties"):
                        dto_props = expand_properties(dto["properties"], definitions, 0, MAX_DEPTH)
                        required_fields = dto.get("required") or []
                        for prop in dto_props:
                            if prop["name"] in required_fields:
                                prop["required"] = True
                        param_info["dtoProperties"] = dto_props
                params.append(param_info)

            # 响应（递归展开嵌套 VO）
            response_info = None
            responses = api_def.get("responses") or {}
            r200 = responses.get("200") or {}
            r_schema = r200.get("schema") if isinstance(r200, dict) else None
            if isinstance(r_schema, dict):
                ref_name = r_schema.get("originalRef") or (
                    r_schema.get("$ref", "").rsplit("/", 1)[-1] if r_schema.get("$ref") else None
                )
                if ref_name:
                    auto_referenced_defs.add(ref_name)
                    vo = definitions.get(ref_name)
                    if isinstance(vo, dict) and vo.get("properties"):
                        response_info = {
                            "voName": ref_name,
                            "properties": expand_properties(
                                vo["properties"], definitions, 0, MAX_DEPTH
                            ),
                        }
                    else:
                        response_info = {"voName": ref_name}

            parsed_apis.append({
                "url": url,
                "method": method.upper(),
                "tag": tag,
                "summary": api_def.get("summary", ""),
                "description": api_def.get("description", ""),
                "deprecated": api_def.get("deprecated", False),
                "parameters": params,
                "response": response_info,
                "operationId": api_def.get("operationId"),
            })

    # 收集被引用 definitions 的传递闭包
    collected_refs = set(auto_referenced_defs)
    queue = list(auto_referenced_defs)
    while queue:
        name = queue.pop(0)
        d = definitions.get(name)
        if not isinstance(d, dict) or not d.get("properties"):
            continue
        for prop in d["properties"].values():
            ref_name = _ref_name_with_items(prop)
            if ref_name and ref_name in definitions and ref_name not in collected_refs:
                collected_refs.add(ref_name)
                queue.append(ref_name)

    # definitions 摘要层：所有 VO/DTO 的 name + type + description
    definitions_summary = [
        {
            "name": name,
            "type": d.get("type", "object") if isinstance(d, dict) else "object",
            "description": (d.get("description") or d.get("title") or "") if isinstance(d, dict) else "",
        }
        for name, d in definitions.items()
    ]

    # definitionsDetail 详情层：仅展开被引用的（含传递闭包），递归深度 5
    effective_refs = referenced_defs if referenced_defs is not None else collected_refs
    definitions_detail = []
    for name in effective_refs:
        d = definitions.get(name)
        if isinstance(d, dict):
            definitions_detail.append({
                "name": name,
                "type": d.get("type", "object"),
                "description": d.get("description") or d.get("title") or "",
                "properties": expand_properties(
                    d.get("properties") or {}, definitions, 0, MAX_DEPTH
                ),
            })

    return {
        "info": {
            "title": info.get("title"),
            "version": info.get("version"),
            "description": info.get("description"),
        } if info else None,
        "basePath": api_base_path,
        "tags": parsed_tags,
        "apis": parsed_apis,
        "definitions": definitions_summary,
        "definitionsDetail": definitions_detail,
    }


# ==================== 过滤（对齐 JS 端 filterApis + apiFilter 裁剪 definitionsDetail） ====================

def filter_apis(apis, keywords):
    """对齐 JS 端 filterApis：对 url/summary/operationId/tag/description/dtoRef/voName 模糊匹配。"""
    if not keywords or not apis:
        return apis
    lower_keywords = [k.lower() for k in keywords]
    out = []
    for api in apis:
        dto_refs = " ".join(p.get("dtoRef", "") for p in (api.get("parameters") or []) if p.get("dtoRef"))
        response_vo = (api.get("response") or {}).get("voName", "") if api.get("response") else ""
        searchable = " ".join([
            api.get("url", ""),
            api.get("summary", ""),
            api.get("operationId") or "",
            api.get("tag", ""),
            api.get("description", ""),
            dto_refs,
            response_vo,
        ]).lower()
        if any(kw in searchable for kw in lower_keywords):
            out.append(api)
    return out


def trim_definitions_detail(result, keywords):
    """apiFilter 生效时，definitionsDetail 仅保留匹配 API 引用的（含传递闭包）。
    对齐 JS 端 server.js 第 1581-1604 行的裁剪逻辑。"""
    apis = result.get("apis") or []
    matched_refs = set()
    for api in apis:
        for p in api.get("parameters") or []:
            if p.get("dtoRef"):
                matched_refs.add(p["dtoRef"])
        if api.get("response") and api["response"].get("voName"):
            matched_refs.add(api["response"]["voName"])

    detail_index = {d["name"]: d for d in result.get("definitionsDetail") or []}
    detail_names = set(matched_refs)
    queue = list(matched_refs)
    while queue:
        name = queue.pop(0)
        detail = detail_index.get(name)
        if not detail or not detail.get("properties"):
            continue
        for prop in detail["properties"]:
            ref = prop.get("refName")
            if ref and ref not in detail_names:
                detail_names.add(ref)
                queue.append(ref)

    result["definitionsDetail"] = [d for d in result.get("definitionsDetail") or [] if d["name"] in detail_names]
    result["_filter"] = {
        "keywords": keywords,
        "matched": len(apis),
    }
    return result


# ==================== get_definition（单 VO 查询兜底） ====================

def get_definition(api_docs, definition_name):
    """对齐 JS 端 get_definition action：按名查询单个 VO，找不到给模糊匹配建议。"""
    definitions = api_docs.get("definitions") or {}
    if definition_name in definitions:
        result = parse_api_docs(api_docs, referenced_defs={definition_name})
        target = next((d for d in result["definitionsDetail"] if d["name"] == definition_name), None)
        return {"found": True, "definition": target}

    # 模糊匹配
    lower = definition_name.lower()
    suggestions = [n for n in definitions.keys() if lower in n.lower()][:10]
    return {
        "found": False,
        "definitionName": definition_name,
        "suggestions": suggestions,
        "message": f"未找到 definition '{definition_name}'。" + (
            f"相近建议：{', '.join(suggestions)}" if suggestions else "无相近匹配。"
        ),
    }


# ==================== 输入：curl 下载 / 读文件 / 缓存 ====================

def _diagnose_json_error(text, err):
    """根据错误位置给出 ±80 字符上下文 + 常见原因建议。"""
    pos = getattr(err, "pos", None)
    if pos is None:
        # 用 lineno/colno 反推 pos
        lines = text.split("\n")
        pos = sum(len(l) + 1 for l in lines[: err.lineno - 1]) + (err.colno - 1)
    start, end = max(0, pos - 80), min(len(text), pos + 80)
    snippet = text[start:end].replace("\n", "\\n")
    marker = " " * (pos - start) + "^"

    # 常见原因诊断
    hints = []
    low = text[:500].lstrip()
    if low.startswith("<"):
        hints.append("响应以 '<' 开头，可能是 HTML 错误页/登录跳转页，而非 JSON。检查 cookie/auth 是否过期")
    if OCTAL_LITERAL_RE.search(text):
        hints.append("检测到八进制字面量（如 [011,015]），脚本应已自动修复——若仍报错请提交 issue")
    if len(text) < 1024 and text.rstrip().endswith(("}", "]")) is False:
        hints.append("响应可能被截断（curl 管道/终端缓冲），建议改用 --file 模式直接读本地文件")

    msg = f"JSON 解析失败 at line {err.lineno} col {err.colno}: {err.msg}\n"
    msg += f"  上下文: ...{snippet}...\n"
    msg += f"          {' ' * 11}{marker}\n"
    if hints:
        msg += "  可能原因:\n" + "\n".join(f"    - {h}" for h in hints)
    return msg


def _fix_octal_literals(text):
    """把 [011,015] 这类八进制字面量改成 [11,15]（仅在 array/value 位置）。"""
    return OCTAL_LITERAL_RE.sub(lambda m: m.group(1), text)


def load_json_text(text):
    """解析 JSON；首次失败则尝试八进制修复 fallback，再失败给出诊断信息。"""
    try:
        return json.loads(text)
    except json.JSONDecodeError as e:
        # Fallback: 尝试修复八进制字面量
        if OCTAL_LITERAL_RE.search(text):
            fixed = _fix_octal_literals(text)
            if fixed != text:
                try:
                    parsed = json.loads(fixed)
                    print(f"[INFO] 自动修复八进制字面量后解析成功", file=sys.stderr)
                    return parsed
                except json.JSONDecodeError as e2:
                    raise RuntimeError(_diagnose_json_error(fixed, e2)) from None
        raise RuntimeError(_diagnose_json_error(text, e)) from None


def _cache_key(url, cookie, auth):
    h = hashlib.sha256()
    h.update(url.encode("utf-8"))
    if cookie:
        h.update(b"\x00" + cookie.encode("utf-8"))
    if auth:
        h.update(b"\x00" + auth.encode("utf-8"))
    return h.hexdigest()[:16]


def fetch_via_curl(url, cookie=None, auth=None, timeout=30):
    """脚本内部走 curl 下载 Swagger JSON。失败抛 RuntimeError。"""
    cmd = ["curl", "-sS", "--connect-timeout", "10", "--max-time", str(timeout),
           "-H", "Accept: application/json"]
    if cookie:
        cmd += ["-H", f"Cookie: {cookie}"]
    if auth:
        cmd += ["-H", f"Authorization: Basic {auth}"]
    cmd.append(url)
    try:
        r = subprocess.run(cmd, capture_output=True, text=True, check=False)
    except FileNotFoundError:
        raise RuntimeError("找不到 curl 命令，请确认本地已安装 curl") from None
    if r.returncode != 0:
        raise RuntimeError(f"curl 失败 (exit {r.returncode}): {r.stderr.strip() or '无错误输出'}")
    if not r.stdout.strip():
        raise RuntimeError("curl 返回空响应")
    return r.stdout


def fetch_with_cache(url, cookie=None, auth=None, timeout=30, ttl=DEFAULT_CACHE_TTL, no_cache=False):
    """带 TTL 缓存的 fetch：缓存命中直接读 /tmp/swagger-cache/<key>.json，否则 curl 拉取并写入。"""
    if no_cache or ttl <= 0:
        return fetch_via_curl(url, cookie=cookie, auth=auth, timeout=timeout)

    key = _cache_key(url, cookie, auth)
    CACHE_DIR.mkdir(parents=True, exist_ok=True)
    cache_file = CACHE_DIR / f"{key}.json"
    meta_file = CACHE_DIR / f"{key}.meta.json"

    if cache_file.exists() and meta_file.exists():
        try:
            meta = json.loads(meta_file.read_text(encoding="utf-8"))
            age = time.time() - meta.get("cached_at", 0)
            if age < ttl:
                print(f"[CACHE] 命中 {key} (age={int(age)}s, ttl={ttl}s)", file=sys.stderr)
                return cache_file.read_text(encoding="utf-8")
        except Exception:
            pass  # 缓存损坏 → 重新拉

    raw = fetch_via_curl(url, cookie=cookie, auth=auth, timeout=timeout)
    try:
        cache_file.write_text(raw, encoding="utf-8")
        meta_file.write_text(json.dumps({
            "url": url, "cached_at": time.time(), "ttl": ttl, "size": len(raw),
        }), encoding="utf-8")
        print(f"[CACHE] 已写入 {key} ({len(raw)} bytes)", file=sys.stderr)
    except Exception as e:
        print(f"[WARN] 缓存写入失败: {e}", file=sys.stderr)
    return raw


# ==================== CLI ====================

def main():
    parser = argparse.ArgumentParser(
        description="本地 Swagger 文档提取与解析（与服务端 parseApiDocs 一比一对齐）",
        formatter_class=argparse.RawDescriptionHelpFormatter,
        epilog=__doc__,
    )
    src = parser.add_mutually_exclusive_group(required=True)
    src.add_argument("--url", help="Swagger api-docs URL（脚本内部走 curl 下载）")
    src.add_argument("--file", help="本地 Swagger JSON 文件路径")

    parser.add_argument("--cookie", help="--url 模式可选，鉴权后的 Cookie 值（如 'JSESSIONID=xxx'）")
    parser.add_argument("--auth", help="--url 模式可选，Basic 鉴权 base64(user:pass)")
    parser.add_argument("--timeout", type=int, default=30, help="--url 模式 curl 超时秒数，默认 30")
    parser.add_argument("--cache-ttl", type=int, default=DEFAULT_CACHE_TTL,
                        help=f"--url 模式缓存 TTL 秒数，默认 {DEFAULT_CACHE_TTL}（10 分钟），设为 0 关闭")
    parser.add_argument("--no-cache", action="store_true",
                        help="--url 模式强制重新下载，忽略缓存")

    parser.add_argument(
        "--filter", action="append", default=[],
        help="按关键词过滤 api（可多次指定，对 url/summary/operationId/tag/description/dtoRef/voName 模糊匹配）",
    )
    parser.add_argument(
        "--definition",
        help="仅查询单个 VO/DTO 的完整字段定义（等价于服务端 get_definition action）",
    )
    parser.add_argument("--out", help="输出到文件（默认 stdout）")
    parser.add_argument(
        "--indent", type=int, default=2,
        help="JSON 输出缩进，默认 2；设为 0 输出紧凑 JSON",
    )

    args = parser.parse_args()

    # 1. 读取原始 Swagger JSON
    try:
        if args.url:
            raw_text = fetch_with_cache(
                args.url, cookie=args.cookie, auth=args.auth, timeout=args.timeout,
                ttl=args.cache_ttl, no_cache=args.no_cache,
            )
        else:
            raw_text = Path(args.file).read_text(encoding="utf-8")
    except Exception as e:
        print(f"[ERROR] 读取 Swagger 失败: {e}", file=sys.stderr)
        sys.exit(2)

    # 2. 解析 JSON
    try:
        api_docs = load_json_text(raw_text)
    except Exception as e:
        print(f"[ERROR] {e}", file=sys.stderr)
        sys.exit(3)

    # 3. 路径分支：--definition 单查 / 默认完整解析
    try:
        if args.definition:
            result = get_definition(api_docs, args.definition)
        else:
            result = parse_api_docs(api_docs)
            if args.filter:
                total = len(result["apis"])
                result["apis"] = filter_apis(result["apis"], args.filter)
                result = trim_definitions_detail(result, args.filter)
                result["_filter"]["total"] = total
    except Exception as e:
        print(f"[ERROR] 解析失败: {e}", file=sys.stderr)
        sys.exit(3)

    # 4. 输出
    indent = args.indent if args.indent > 0 else None
    out_text = json.dumps(result, ensure_ascii=False, indent=indent)
    if args.out:
        Path(args.out).write_text(out_text, encoding="utf-8")
        print(f"[OK] 已输出到 {args.out} ({len(out_text)} bytes)", file=sys.stderr)
    else:
        sys.stdout.write(out_text)
        sys.stdout.write("\n")


if __name__ == "__main__":
    main()
