#!/usr/bin/env python3
"""
Phase 3 — Verify: 核对迁移结果.

对比 manifest.json（源）和 v0.25.6 API（目标）：
  1. KB 数量
  2. 每个 KB 的文档数量
  3. 解析状态分布
  4. 差异清单

Usage:
    python verify_migration.py -c config.yaml
"""

import argparse
import json
import logging
from collections import defaultdict
from pathlib import Path

import requests
from ruamel.yaml import YAML

logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
log = logging.getLogger("verify")


def load_config(path: str) -> dict:
    yaml = YAML(typ="safe", pure=True)
    with open(path) as f:
        return yaml.load(f)


class RAGFlowAPI:
    def __init__(self, base_url, api_key, timeout=60):
        self.base = f"{base_url.rstrip('/')}/api/v1"
        self.s = requests.Session()
        self.s.headers["Authorization"] = f"Bearer {api_key}"
        self.timeout = timeout

    def _ok(self, resp, action):
        d = resp.json()
        if d.get("code") != 0:
            raise RuntimeError(f"{action}: {d.get('message', d)}")
        return d

    def list_datasets(self):
        r = self.s.get(f"{self.base}/datasets",
                       params={"page": 1, "page_size": 1000}, timeout=self.timeout)
        return self._ok(r, "list")["data"]

    def list_documents(self, dataset_id):
        docs, page = [], 1
        while True:
            r = self.s.get(f"{self.base}/datasets/{dataset_id}/documents",
                           params={"page": page, "page_size": 100}, timeout=self.timeout)
            batch = self._ok(r, "list")["data"].get("docs", [])
            docs.extend(batch)
            if len(batch) < 100:
                break
            page += 1
        return docs


def run(cfg: dict):
    temp_dir = Path(cfg["temp_dir"])

    # 源
    with open(temp_dir / "manifest.json", encoding="utf-8") as f:
        manifest = json.load(f)

    # 目标
    api = RAGFlowAPI(cfg["target"]["base_url"], cfg["target"]["api_key"])
    tgt_kbs = {d["name"]: d for d in api.list_datasets()}

    # 逐 KB 对比
    rows = []
    src_total, tgt_total = 0, 0
    status_summary: dict[str, int] = defaultdict(int)

    for entry in manifest:
        kb_name = entry["kb_name"]
        src_n = len(entry["files"])
        src_total += src_n

        tgt_ds = tgt_kbs.get(kb_name)
        if not tgt_ds:
            rows.append((kb_name, src_n, 0, 0, "❌ 缺失"))
            continue

        tgt_docs = api.list_documents(tgt_ds["id"])
        tgt_n = len(tgt_docs)
        tgt_total += tgt_n

        for d in tgt_docs:
            status_summary[d.get("run", "UNKNOWN")] += 1

        diff = tgt_n - src_n
        flag = "✅" if diff == 0 else f"❌ {diff:+d}"
        rows.append((kb_name, src_n, tgt_n, diff, flag))

    # 打印报告
    print()
    print("=" * 70)
    print("  RAGFlow v0.17.2 → v0.25.6  Migration Verification")
    print("=" * 70)

    print(f"  KBs:      源={len(manifest)}, 目标={len(tgt_kbs)}")
    print(f"  Documents: 源={src_total}, 目标={tgt_total} "
          f"({'✅' if src_total == tgt_total else '❌'})")

    print(f"\n  解析状态:")
    for s, n in sorted(status_summary.items()):
        icon = {"DONE": "✅", "FAIL": "❌", "CANCEL": "⚠️", "0": "⏳", "RUNNING": "🔄"}.get(s, "❓")
        label = {"0": "PENDING"}.get(s, s)
        print(f"    {icon} {label}: {n}")

    print(f"\n  {'KB Name':<32s} {'源':>5s} {'目标':>5s} {'差异':>5s}")
    print(f"  {'─' * 50}")
    for name, src, tgt, diff, flag in rows:
        name = name[:30] if len(name) <= 30 else name[:29] + "…"
        print(f"  {flag} {name:<30s} {src:>5d} {tgt:>5d} {diff:>+5d}")

    print(f"\n{'=' * 70}\n")


def main():
    p = argparse.ArgumentParser()
    p.add_argument("-c", "--config", required=True)
    args = p.parse_args()
    run(load_config(args.config))


if __name__ == "__main__":
    main()
