#!/usr/bin/env python3
"""
Phase 1 — Export: 从 v0.17.2 按知识库导出文件.

1. 连接 v0.17.2 MySQL，读取 Knowledgebase → Document → File 关联
2. 从 MinIO 下载文件到本地
3. 输出 manifest.json（kb_id 为主键，避免同名冲突）

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

import argparse
import json
import logging
import os
import sys
from pathlib import Path

import pymysql
from minio import Minio
from ruamel.yaml import YAML

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


# ── config ──────────────────────────────────────────────────────────────────

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


# ── helpers ─────────────────────────────────────────────────────────────────

def connect_mysql(cfg: dict):
    mc = cfg["mysql"]
    conn = pymysql.connect(
        host=mc["host"], port=mc["port"], user=mc["user"],
        password=mc["password"], database=mc["database"],
        charset="utf8mb4", cursorclass=pymysql.cursors.DictCursor,
    )
    log.info("MySQL: %s:%d/%s", mc["host"], mc["port"], mc["database"])
    return conn


def fetch(conn, sql, *params):
    with conn.cursor() as cur:
        cur.execute(sql, params)
        return cur.fetchall()


def connect_minio(cfg: dict) -> Minio:
    sc = cfg["storage"]
    client = Minio(sc["host"], access_key=sc["access_key"],
                   secret_key=sc["secret_key"], secure=sc.get("secure", False))
    log.info("MinIO: %s", sc["host"])
    return client


# ── main ────────────────────────────────────────────────────────────────────

def run(cfg: dict):
    temp_dir = Path(cfg["temp_dir"])
    file_dir = temp_dir / "files"
    file_dir.mkdir(parents=True, exist_ok=True)

    conn = connect_mysql(cfg["source"])
    minio_client = connect_minio(cfg["source"])

    # 1. 读取所有 KB
    kbs = fetch(conn, """
        SELECT id, name, description
        FROM knowledgebase WHERE status = '1'
    """)
    log.info("KB count: %d", len(kbs))

    manifest = []
    total_files = 0

    for kb in kbs:
        kb_id = kb["id"]
        kb_name = kb["name"]

        # 2. 该 KB 下所有 document
        docs = fetch(conn, """
            SELECT id, name, location, type
            FROM document WHERE kb_id = %s AND status = '1'
        """, kb_id)

        if not docs:
            manifest.append({"kb_id": kb_id, "kb_name": kb_name,
                           "kb_description": kb.get("description", ""), "files": []})
            continue

        doc_ids = [d["id"] for d in docs]
        doc_map = {d["id"]: d for d in docs}

        # 3. file2document 关联
        placeholders = ",".join(["%s"] * len(doc_ids))
        f2d = fetch(conn, f"""
            SELECT document_id, file_id
            FROM file2document WHERE document_id IN ({placeholders})
        """, *doc_ids)

        file_ids = list({r["file_id"] for r in f2d})
        file_map = {}
        if file_ids:
            fp = ",".join(["%s"] * len(file_ids))
            files = fetch(conn, f"""
                SELECT id, name FROM file WHERE id IN ({fp})
            """, *file_ids)
            file_map = {f["id"]: f for f in files}

        # doc_id → file_name
        doc_file_names: dict[str, str] = {}
        for r in f2d:
            f = file_map.get(r["file_id"])
            if f:
                doc_file_names[r["document_id"]] = f["name"]

        # 4. 下载文件，目录名: {kb_name}_{kb_id前8位}
        safe_dir = f"{kb_name}_{kb_id[:8]}"
        kb_file_dir = file_dir / safe_dir
        kb_file_dir.mkdir(parents=True, exist_ok=True)

        files_manifest = []
        for doc in docs:
            doc_id = doc["id"]
            # 优先用 file 表的 name，其次用 document.name
            doc_name = doc_file_names.get(doc_id) or doc.get("name") or f"{doc_id}.bin"
            local_path = kb_file_dir / doc_name

            # 处理重名
            counter = 1
            stem, suffix = os.path.splitext(doc_name)
            while local_path.exists():
                local_path = kb_file_dir / f"{stem}_{counter}{suffix}"
                counter += 1

            # 从 MinIO 下载
            bucket = kb_id
            key = doc.get("location")
            if not key:
                log.warning("doc %s has no location, skip", doc_id)
                continue

            try:
                minio_client.fget_object(bucket, key, str(local_path))
                files_manifest.append({
                    "name": local_path.name,
                    "local_path": str(local_path.relative_to(temp_dir)),
                })
                total_files += 1
            except Exception as e:
                log.error("download failed %s/%s: %s", bucket, key, e)

        manifest.append({
            "kb_id": kb_id,
            "kb_name": kb_name,
            "kb_description": kb.get("description", ""),
            "files": files_manifest,
        })
        log.info("KB %s: %d files", kb_name, len(files_manifest))

    conn.close()

    # 5. 写 manifest
    manifest_path = temp_dir / "manifest.json"
    with open(manifest_path, "w", encoding="utf-8") as f:
        json.dump(manifest, f, ensure_ascii=False, indent=2)
    log.info("manifest → %s", manifest_path)
    log.info("Done: %d KBs, %d files", len(manifest), total_files)


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()
