#!/usr/bin/env python3
"""Export readable Weknora document chunks into size-bounded Coze Markdown batches."""

from __future__ import annotations

import argparse
import json
import re
import shutil
import sqlite3
from collections import defaultdict
from dataclasses import asdict, dataclass
from pathlib import Path


DEFAULT_DB = Path("/Users/jack/code/099-github/WeKnora/data/weknora.db")
EXCLUDED_KB = "DeltaFStation 研报一手信息库"


@dataclass
class Batch:
    knowledge_base: str
    path: str
    documents: int
    bytes: int


def safe_name(value: str) -> str:
    value = re.sub(r"[^0-9A-Za-z\u4e00-\u9fff]+", "-", value).strip("-")
    return value[:80] or "knowledge-base"


def write_batch(output: Path, name: str, index: int, parts: list[str], document_count: int) -> Batch:
    file_name = f"{safe_name(name)}-{index:03d}.md"
    path = output / file_name
    text = f"# {name}\n\n" + "\n\n".join(parts) + "\n"
    path.write_text(text, encoding="utf-8")
    return Batch(name, file_name, document_count, path.stat().st_size)


def export_knowledge_base(
    connection: sqlite3.Connection,
    output: Path,
    kb_id: str,
    kb_name: str,
    max_bytes: int,
) -> tuple[list[Batch], list[dict[str, str]]]:
    rows = connection.execute(
        """
        SELECT k.id, k.title, k.parse_status, c.chunk_index, c.content
        FROM knowledges k
        LEFT JOIN chunks c ON c.knowledge_id = k.id AND c.deleted_at IS NULL
        WHERE k.knowledge_base_id = ? AND k.deleted_at IS NULL
        ORDER BY k.created_at, c.chunk_index
        """,
        (kb_id,),
    ).fetchall()

    documents: dict[str, dict[str, object]] = {}
    for knowledge_id, title, status, chunk_index, content in rows:
        item = documents.setdefault(
            knowledge_id,
            {"title": title, "status": status, "chunks": []},
        )
        if content:
            item["chunks"].append((chunk_index, content))

    batches: list[Batch] = []
    unavailable: list[dict[str, str]] = []
    current_parts: list[str] = []
    current_bytes = 0
    current_documents = 0
    batch_index = 1

    for item in documents.values():
        chunks = item["chunks"]
        if not chunks:
            unavailable.append({"knowledge_base": kb_name, "title": str(item["title"]), "status": str(item["status"])})
            continue
        body = "\n\n".join(content for _, content in sorted(chunks))
        document = f"## {item['title']}\n\n{body}"
        document_bytes = len(document.encode("utf-8")) + 2
        if current_parts and current_bytes + document_bytes > max_bytes:
            batches.append(write_batch(output, kb_name, batch_index, current_parts, current_documents))
            batch_index += 1
            current_parts, current_bytes, current_documents = [], 0, 0
        current_parts.append(document)
        current_bytes += document_bytes
        current_documents += 1

    if current_parts:
        batches.append(write_batch(output, kb_name, batch_index, current_parts, current_documents))
    return batches, unavailable


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--database", type=Path, default=DEFAULT_DB)
    parser.add_argument("--output", type=Path, required=True)
    parser.add_argument("--max-bytes", type=int, default=6 * 1024 * 1024)
    args = parser.parse_args()

    if not args.database.is_file():
        raise SystemExit(f"Database not found: {args.database}")
    if args.output.exists():
        shutil.rmtree(args.output)
    args.output.mkdir(parents=True)

    connection = sqlite3.connect(f"file:{args.database}?mode=ro", uri=True)
    knowledge_bases = connection.execute(
        """
        SELECT id, name FROM knowledge_bases
        WHERE deleted_at IS NULL AND name <> ?
        ORDER BY created_at
        """,
        (EXCLUDED_KB,),
    ).fetchall()
    batches: list[Batch] = []
    unavailable: list[dict[str, str]] = []
    for kb_id, kb_name in knowledge_bases:
        current_batches, current_unavailable = export_knowledge_base(
            connection, args.output, kb_id, kb_name, args.max_bytes
        )
        batches.extend(current_batches)
        unavailable.extend(current_unavailable)
    connection.close()

    manifest = {
        "source_database": str(args.database),
        "excluded_knowledge_base": EXCLUDED_KB,
        "batch_count": len(batches),
        "batches": [asdict(batch) for batch in batches],
        "unavailable_document_count": len(unavailable),
        "unavailable_documents": unavailable,
    }
    (args.output / "manifest.json").write_text(
        json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8"
    )
    print(json.dumps(manifest, ensure_ascii=False, indent=2))


if __name__ == "__main__":
    main()
