#!/usr/bin/env python3
"""
Phase 2 — Import: 通过 v0.25.6 REST API 创建知识库并上传文件.

读取 manifest.json，对每个 KB：
  1. POST /api/v1/datasets         创建知识库
  2. POST /api/v1/datasets/{id}/documents  上传文件（自动触发解析）

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

import argparse
import json
import logging
import os
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("import")


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


class RAGFlowAPI:
    """v0.25.6 REST API 最小封装."""

    def __init__(self, base_url: str, api_key: str, timeout: int = 300):
        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):
        try:
            d = resp.json()
        except json.JSONDecodeError:
            raise RuntimeError(f"{action}: bad JSON ({resp.status_code})")
        if d.get("code") != 0:
            raise RuntimeError(f"{action}: {d.get('message', d)}")
        return d

    def create_dataset(self, name: str, description: str = "") -> dict:
        r = self.s.post(f"{self.base}/datasets",
                        json={"name": name, "description": description},
                        timeout=self.timeout)
        return self._ok(r, "create_dataset")["data"]

    def upload_documents(self, dataset_id: str, files: list[tuple[str, bytes]]) -> list[dict]:
        """files: [(filename, blob), ...]"""
        mpart = [("file", (name, blob)) for name, blob in files]
        r = self.s.post(f"{self.base}/datasets/{dataset_id}/documents",
                        files=mpart, timeout=self.timeout)
        return self._ok(r, "upload_documents")["data"]

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

    def list_documents(self, dataset_id: str) -> list[dict]:
        """获取某个 dataset 下所有文档."""
        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)
            data = self._ok(r, "list_documents").get("data", {})
            batch = data.get("docs", [])
            docs.extend(batch)
            if len(batch) < 100:
                break
            page += 1
        return docs

    def delete_documents(self, dataset_id: str, doc_ids: list[str]):
        self.s.delete(f"{self.base}/datasets/{dataset_id}/documents",
                      json={"ids": doc_ids}, timeout=self.timeout)


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

    # 1. 读 manifest
    with open(temp_dir / "manifest.json", encoding="utf-8") as f:
        manifest = json.load(f)
    log.info("Loaded manifest: %d KBs", len(manifest))

    # 2. API 客户端
    api = RAGFlowAPI(cfg["target"]["base_url"], cfg["target"]["api_key"],
                     cfg["target"].get("timeout", 300))

    # 已有 dataset（按 name 索引，避免重复创建）
    existing = {ds["name"]: ds for ds in api.list_datasets()}
    log.info("Existing datasets: %d", len(existing))

    total_uploaded = 0

    for entry in manifest:
        kb_id = entry["kb_id"]
        kb_name = entry["kb_name"]
        files = entry["files"]

        log.info("KB: %s (%d files)", kb_name, len(files))

        # 3. 创建或复用 dataset
        if kb_name in existing:
            ds = existing[kb_name]
            log.info("  reuse dataset: %s", ds["id"])
        else:
            ds = api.create_dataset(kb_name, entry.get("kb_description", ""))
            log.info("  created dataset: %s", ds["id"])

        if not files:
            continue

        # 4. 上传文件（分批）
        batch_size = cfg.get("batch_size", 10)
        for i in range(0, len(files), batch_size):
            batch = files[i:i + batch_size]
            uploads = []
            for f in batch:
                local = temp_dir / f["local_path"]
                if not local.exists():
                    log.warning("  file missing: %s", local)
                    continue
                uploads.append((f["name"], local.read_bytes()))

            if not uploads:
                continue

            try:
                api.upload_documents(ds["id"], uploads)
                total_uploaded += len(uploads)
                log.info("  batch %d/%d: %d docs",
                         i // batch_size + 1,
                         (len(files) + batch_size - 1) // batch_size,
                         len(uploads))
            except Exception as e:
                log.error("  upload failed: %s", e)

    log.info("Done: %d files uploaded", total_uploaded)


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