#!/usr/bin/env python3
"""抽取 schema 离线校验脚本：与服务端校验规则一致，用于提交前自查。

用法：
    python3 validate_extract_schema.py schema.json
    python3 validate_extract_schema.py --json '{"type":"object","properties":{...}}'
    python3 validate_extract_schema.py schema.json --max-keys 100 --max-depth 6 --max-size-kb 600

校验项（与服务端一致）：
    1. JSON 合法性；
    2. body 大小：序列化后 UTF-8 字节数 <= max-size-kb（默认 600KB）；
    3. 格式：JSON Schema 格式；
    4. 结构完整性：object 必须有非空 properties；array 必须有非空 items；
    5. 嵌套深度：根级字段为第 1 层，object 子字段 +1、array 的 items +1，最多 6 层；
    6. 叶子字段数：object 递归展开、array 算 1 个。
"""
from __future__ import annotations

import argparse
import json
import sys
from typing import Any

MAX_REF_DEPTH = 10          # $ref 链式解析上限（防循环引用）

# 常见 JSON Schema 关键字：误放进 properties 内时给出针对性提示
SCHEMA_KEYWORDS = (
    "$defs", "$ref", "additionalProperties", "anyOf", "description",
    "enum", "items", "properties", "required", "title", "type",
)


class SchemaCheckError(Exception):
    """校验失败：code 与服务端错误码一致，path 为出错字段路径。"""

    def __init__(self, code: int, message: str, path: str = ""):
        super().__init__(message)
        self.code = code
        self.path = path


class SchemaChecker:
    """按服务端规则校验 extract_schema。"""

    def __init__(self, max_keys: int = 100, max_depth: int = 6,
                 max_size_kb: int = 600):
        self.max_keys = max_keys
        self.max_depth = max_depth
        self.max_size_kb = max_size_kb

    def check(self, schema: Any) -> dict:
        self._check_size(schema)
        if not isinstance(schema, dict):
            raise SchemaCheckError(4011, "extract_schema 必须是非空 JSON 对象")
        depth, fields = self._check(schema)
        return {"format": "llama",
                "depth": depth, "leaf_fields": len(fields), "ok": True}

    # ── 校验项 ──

    def _check_size(self, schema: Any) -> None:
        size = len(json.dumps(schema, ensure_ascii=False).encode("utf-8"))
        limit = self.max_size_kb * 1024
        if size > limit:
            raise SchemaCheckError(
                4013, f"schema body 大小 {size} 字节超过上限 {self.max_size_kb}KB（{limit} 字节）")

    # ── llamaparse(JSON Schema) 格式 ──

    def _check(self, schema: dict) -> tuple[int, list]:
        properties = schema.get("properties")
        if not isinstance(properties, dict) or not properties:
            raise SchemaCheckError(4011, "extract_schema 的 properties 缺失或为空")
        defs = schema.get("$defs") if isinstance(schema.get("$defs"), dict) else {}
        fields: list = []
        max_depth: list[int] = [0]
        self._walk(properties, defs, 1, fields, max_depth)
        if len(fields) > self.max_keys:
            raise SchemaCheckError(
                4011, f"extract_schema 字段数 {len(fields)} 超过上限 {self.max_keys}")
        return max_depth[0], fields

    def _walk(self, properties: dict, defs: dict, depth: int,
                    out: list, max_depth: list) -> None:
        """递归校验 + 展平叶子（object 递归展开，array 算 1 个叶子）。"""
        for key, prop in properties.items():
            if not isinstance(prop, dict):
                hint = ""
                if key in SCHEMA_KEYWORDS:
                    hint = (f"；「{key}」是 schema 关键字，可能被误放进了 properties 内，"
                            f"对象级定义应与 properties 平级")
                raise SchemaCheckError(
                    4011,
                    f"字段定义必须是 JSON 对象，实际为 {self._describe(prop)}{hint}",
                    f"properties.{key}")
            self._visit_node(prop, defs, f"properties.{key}", depth, out, max_depth)

    def _visit_node(self, node: dict, defs: dict, path: str, depth: int,
                    out: list, max_depth: list, check_only: bool = False) -> None:
        """校验单个节点：深度、object/array 完整性；object 递归展开，array 算 1 个叶子。

        check_only=True 表示仅做深度/结构校验、不计入叶子数（array 的 items 子树）。
        """
        if depth > self.max_depth:
            raise SchemaCheckError(
                4012, f"extract_schema 嵌套层数超过上限 {self.max_depth}", path)
        max_depth[0] = max(max_depth[0], depth)
        flat = self._flatten(node, defs)
        for t in flat["types"]:
            if t == "object":
                props = flat.get("properties")
                if not isinstance(props, dict) or not props:
                    raise SchemaCheckError(
                        4011, "type 为 object 但 properties 缺失或为空", path)
                for key, val in props.items():
                    if isinstance(val, dict):
                        self._visit_node(val, defs, f"{path}.{key}", depth + 1,
                                        out, max_depth, check_only)
            elif t == "array":
                items = flat.get("items")
                if not isinstance(items, dict) or not items:
                    raise SchemaCheckError(
                        4011, "type 为 array 但 items 缺失或为空", path)
                self._visit_node(items, defs, f"{path}.items", depth + 1,
                                out, max_depth, check_only=True)
        if flat["types"][0] != "object" and not check_only:
            out.append(path)

    def _flatten(self, prop: dict, defs: dict) -> dict:
        """$ref 解析 + anyOf/type 数组归一化，返回与服务端一致的统一形态。"""
        merged = dict(prop)
        for _ in range(MAX_REF_DEPTH):
            ref = merged.get("$ref")
            if not isinstance(ref, str) or not ref.startswith("#/$defs/"):
                break
            target = defs.get(ref.split("/")[-1])
            if not isinstance(target, dict):
                break
            base = dict(target)
            base.update({k: v for k, v in merged.items() if k != "$ref"})
            merged = base
        if isinstance(merged.get("anyOf"), list):
            types: list[str] = []
            properties = items = None
            for branch in merged["anyOf"]:
                if not isinstance(branch, dict):
                    continue
                f = self._flatten(branch, defs)
                for t in f["types"]:
                    if t != "null" and t not in types:
                        types.append(t)
                properties = properties or f.get("properties")
                items = items or f.get("items")
            merged = {k: v for k, v in merged.items() if k != "anyOf"}
            merged.update({"types": types or ["string"],
                           "properties": properties, "items": items})
            return merged
        t = merged.get("type")
        if isinstance(t, list):
            merged["types"] = [str(x) for x in t if x != "null"] or ["string"]
        else:
            merged["types"] = [str(t)] if t else ["string"]
        return merged

    @staticmethod
    def _describe(value: Any) -> str:
        """生成值的人类可读描述，用于错误提示。"""
        if isinstance(value, str):
            text = value if len(value) <= 50 else value[:50] + "..."
            return f"字符串 {json.dumps(text, ensure_ascii=False)}"
        return type(value).__name__


def main() -> int:
    parser = argparse.ArgumentParser(description="抽取 schema 离线校验")
    parser.add_argument("file", nargs="?", help="schema JSON 文件路径")
    parser.add_argument("--json", help="直接传入 schema JSON 字符串")
    parser.add_argument("--max-keys", type=int, default=100, help="叶子字段数上限（默认 100）")
    parser.add_argument("--max-depth", type=int, default=6, help="嵌套深度上限（默认 6）")
    parser.add_argument("--max-size-kb", type=int, default=600, help="body 大小上限 KB（默认 600）")
    args = parser.parse_args()

    raw = args.json
    if raw is None:
        if not args.file:
            parser.error("需要提供 schema JSON 文件路径，或使用 --json 传入 JSON 字符串")
        with open(args.file, encoding="utf-8") as f:
            raw = f.read()
    try:
        schema = json.loads(raw)
    except json.JSONDecodeError as e:
        print(f"[失败] 9004 extract_schema 不是合法 JSON：{e}")
        return 1

    checker = SchemaChecker(args.max_keys, args.max_depth, args.max_size_kb)
    try:
        result = checker.check(schema)
    except SchemaCheckError as e:
        loc = f"（位置：{e.path}）" if e.path else ""
        print(f"[失败] 错误码 {e.code}：{e}{loc}")
        return 1
    print(f"[通过] 格式={result['format']}，叶子字段数={result['leaf_fields']}，"
          f"最大深度={result['depth']}，大小={len(raw.encode('utf-8'))} 字节")
    return 0


if __name__ == "__main__":
    sys.exit(main())
