#!/usr/bin/env python3
"""Extract historical bidding documents into Markdown knowledge and upsert to WeKnora."""

from __future__ import annotations

import argparse
import csv
import datetime as dt
import json
import os
import re
import subprocess
import sys
import tempfile
import urllib.parse
from collections import Counter
from pathlib import Path
from xml.etree import ElementTree as ET
from zipfile import ZipFile


REPO_ROOT = Path(__file__).resolve().parents[2]
NOTEBOOKLM_SCRIPTS = Path("/Users/jack/code/099-github/notebooklm-mcp/scripts")
if str(NOTEBOOKLM_SCRIPTS) not in sys.path:
    sys.path.insert(0, str(NOTEBOOKLM_SCRIPTS))

from sync_daily_inputs_to_weknora import (  # noqa: E402
    DEFAULT_BASE_URL,
    DEFAULT_EMBEDDING_MODEL_DIMENSION,
    DEFAULT_EMBEDDING_MODEL_NAME,
    KnowledgeDoc,
    WeKnoraClient,
)


KB_NAME = "ZHCT 历史招投标标书知识库"
KB_DESCRIPTION = (
    "2025-2026 历史招投标标书、采购文件、商务技术文件、报价资料、证书附件和项目材料的本地 RAG 知识库。"
    "原始文件保留在企业微信微盘，本知识库保存文本抽取、来源路径、项目年份和证据边界。"
)
TAG_NAME = "source:historical-bidding"


TEXT_EXTENSIONS = {".txt", ".md", ".csv"}
DOCUMENT_EXTENSIONS = {".pdf", ".docx", ".doc", ".xlsx", ".pptx"} | TEXT_EXTENSIONS
SKIP_TEXT_EXTENSIONS = {".png", ".jpg", ".jpeg", ".dwg", ".zip", ".ebid", ".tbj", ".qtb"}


def clean_text(text: str) -> str:
    text = text.replace("\x00", "")
    text = re.sub(r"[ \t]+", " ", text)
    text = re.sub(r"\n{4,}", "\n\n\n", text)
    return text.strip()


def safe_slug(value: str, limit: int = 120) -> str:
    value = re.sub(r"[^\w\u4e00-\u9fff.-]+", "-", value, flags=re.UNICODE).strip("-")
    return value[:limit] or "document"


def read_text_file(path: Path) -> str:
    return path.read_text(encoding="utf-8", errors="ignore")


def extract_pdf(path: Path) -> str:
    cmd = ["pdftotext", "-layout", str(path), "-"]
    result = subprocess.run(cmd, check=False, capture_output=True, text=True, timeout=120)
    if result.returncode != 0:
        return ""
    return result.stdout


def extract_pdf_ocr(path: Path, *, max_pages: int = 0) -> str:
    try:
        info = subprocess.run(["pdfinfo", str(path)], check=False, capture_output=True, text=True, timeout=15).stdout
    except Exception:
        info = ""
    pages = 0
    for line in info.splitlines():
        if line.startswith("Pages:"):
            try:
                pages = int(line.split(":", 1)[1].strip())
            except ValueError:
                pages = 0
    if max_pages and pages and pages > max_pages:
        pages = max_pages

    with tempfile.TemporaryDirectory() as tmp:
        prefix = Path(tmp) / "page"
        cmd = ["pdftoppm", "-r", "130", "-png"]
        if pages:
            cmd.extend(["-f", "1", "-l", str(pages)])
        cmd.extend([str(path), str(prefix)])
        rendered = subprocess.run(cmd, check=False, capture_output=True, text=True, timeout=max(180, pages * 20 if pages else 600))
        if rendered.returncode != 0:
            return ""
        texts: list[str] = []
        images = sorted(Path(tmp).glob("page-*.png"))
        for index, image in enumerate(images, start=1):
            ocr = subprocess.run(
                ["tesseract", str(image), "stdout", "-l", "chi_sim+eng", "--psm", "6"],
                check=False,
                capture_output=True,
                text=True,
                timeout=90,
            )
            if ocr.stdout.strip():
                texts.append(f"\n\n## OCR page {index}\n\n{ocr.stdout}")
        return "\n".join(texts)


def extract_docx(path: Path) -> str:
    try:
        import docx  # type: ignore
    except Exception:
        return extract_with_textutil(path)
    try:
        doc = docx.Document(str(path))
    except Exception:
        return extract_with_textutil(path)
    parts: list[str] = []
    for para in doc.paragraphs:
        if para.text.strip():
            parts.append(para.text)
    for table in doc.tables:
        for row in table.rows:
            cells = [cell.text.strip().replace("\n", " / ") for cell in row.cells]
            if any(cells):
                parts.append(" | ".join(cells))
    return "\n".join(parts)


def extract_with_textutil(path: Path) -> str:
    with tempfile.TemporaryDirectory() as tmp:
        out_dir = Path(tmp)
        result = subprocess.run(
            ["textutil", "-convert", "txt", "-output", str(out_dir / "out.txt"), str(path)],
            check=False,
            capture_output=True,
            text=True,
            timeout=120,
        )
        out = out_dir / "out.txt"
        if result.returncode == 0 and out.exists():
            return out.read_text(encoding="utf-8", errors="ignore")
    return ""


def extract_xlsx(path: Path) -> str:
    try:
        import openpyxl  # type: ignore
    except Exception:
        return ""
    try:
        wb = openpyxl.load_workbook(path, data_only=True, read_only=True)
    except Exception:
        return ""
    parts: list[str] = []
    for sheet in wb.worksheets:
        parts.append(f"## Sheet: {sheet.title}")
        row_count = 0
        for row in sheet.iter_rows(values_only=True):
            values = [str(v).strip() for v in row if v is not None and str(v).strip()]
            if values:
                parts.append(" | ".join(values))
                row_count += 1
            if row_count >= 5000:
                parts.append("[sheet truncated at 5000 non-empty rows]")
                break
    return "\n".join(parts)


def extract_pptx(path: Path) -> str:
    ns = {
        "a": "http://schemas.openxmlformats.org/drawingml/2006/main",
    }
    parts: list[str] = []
    try:
        with ZipFile(path) as zf:
            slide_names = sorted(name for name in zf.namelist() if re.match(r"ppt/slides/slide\d+\.xml$", name))
            for slide_name in slide_names:
                xml = zf.read(slide_name)
                root = ET.fromstring(xml)
                texts = [node.text.strip() for node in root.findall(".//a:t", ns) if node.text and node.text.strip()]
                if texts:
                    parts.append(f"## {slide_name}")
                    parts.extend(texts)
    except Exception:
        return ""
    return "\n".join(parts)


def extract_document(path: Path, extension: str, *, ocr_pdf: bool = False, ocr_max_pages: int = 0) -> tuple[str, str]:
    try:
        if extension == ".pdf":
            text = extract_pdf(path)
            if text.strip():
                return text, "pdftotext"
            if ocr_pdf:
                return extract_pdf_ocr(path, max_pages=ocr_max_pages), "pdftoppm+tesseract"
            return "", "pdftotext"
        if extension == ".docx":
            return extract_docx(path), "python-docx"
        if extension == ".doc":
            return extract_with_textutil(path), "textutil"
        if extension == ".xlsx":
            return extract_xlsx(path), "openpyxl"
        if extension == ".pptx":
            return extract_pptx(path), "pptx-xml"
        if extension in TEXT_EXTENSIONS:
            return read_text_file(path), "text"
    except Exception:
        return "", "failed"
    return "", "unsupported"


def ensure_kb(client: WeKnoraClient, name: str) -> str:
    response = client.request("GET", "/knowledge-bases")
    items = response.get("data") if isinstance(response, dict) else response
    if isinstance(items, dict):
        items = items.get("items") or items.get("list") or items.get("data")
    if isinstance(items, list):
        for item in items:
            if item.get("name") == name:
                return str(item.get("id") or item.get("knowledge_base_id") or item.get("uuid"))
    payload = {
        "name": name,
        "description": KB_DESCRIPTION,
        "embedding_model_name": DEFAULT_EMBEDDING_MODEL_NAME,
        "embedding_model_dimension": DEFAULT_EMBEDDING_MODEL_DIMENSION,
    }
    created = client.request("POST", "/knowledge-bases", payload)
    data = created.get("data") if isinstance(created, dict) else created
    if isinstance(data, dict):
        return str(data.get("id") or data.get("knowledge_base_id") or data.get("uuid"))
    raise RuntimeError(f"failed to create knowledge base: {created}")


def ensure_tag(client: WeKnoraClient) -> dict[str, str]:
    try:
        response = client.request("GET", "/tags")
        items = response.get("data") if isinstance(response, dict) else response
        if isinstance(items, dict):
            items = items.get("items") or items.get("list") or items.get("data")
        if isinstance(items, list):
            for item in items:
                if item.get("name") == TAG_NAME:
                    return {TAG_NAME: str(item.get("id") or item.get("uuid"))}
        created = client.request("POST", "/tags", {"name": TAG_NAME, "color": "#8a5138"})
        data = created.get("data") if isinstance(created, dict) else created
        if isinstance(data, dict):
            return {TAG_NAME: str(data.get("id") or data.get("uuid"))}
    except Exception:
        pass
    return {}


def build_markdown(row: dict[str, str], text: str, method: str) -> str:
    rel = row["relative_to_dest"]
    parts = rel.split("/")
    year = parts[0] if parts else row["source_label"]
    project = parts[1] if len(parts) > 1 else ""
    return "\n".join(
        [
            f"# {Path(rel).name}",
            "",
            "## Source",
            "",
            f"- year: {year}",
            f"- project: {project}",
            f"- relative_path: `{rel}`",
            f"- extracted_path: `{row['extracted_path']}`",
            f"- extension: `{row['extension']}`",
            f"- file_size: {row['file_size']}",
            f"- sha256: `{row['sha256']}`",
            f"- extraction_method: `{method}`",
            "",
            "## Extracted Text",
            "",
            text,
        ]
    )


def main() -> int:
    parser = argparse.ArgumentParser()
    parser.add_argument("--file-index", required=True)
    parser.add_argument("--out-dir", required=True)
    parser.add_argument("--manifest-csv", required=True)
    parser.add_argument("--upsert", action="store_true")
    parser.add_argument("--ocr-failed-pdf", action="store_true")
    parser.add_argument("--ocr-max-pages", type=int, default=0)
    parser.add_argument("--base-url", default=os.getenv("WEKNORA_BASE_URL", DEFAULT_BASE_URL))
    parser.add_argument("--kb-name", default=KB_NAME)
    parser.add_argument("--batch-size", type=int, default=10)
    parser.add_argument("--limit", type=int, default=0)
    args = parser.parse_args()

    out_dir = Path(args.out_dir)
    out_dir.mkdir(parents=True, exist_ok=True)
    rows = list(csv.DictReader(open(args.file_index, encoding="utf-8")))
    manifest_rows: list[dict[str, str]] = []
    docs: list[KnowledgeDoc] = []
    counters: Counter[str] = Counter()

    for idx, row in enumerate(rows, start=1):
        ext = row["extension"].lower()
        path = Path(row["extracted_path"])
        status = "skipped"
        method = "n/a"
        text = ""
        md_rel = ""
        if ext in DOCUMENT_EXTENSIONS and path.exists():
            text, method = extract_document(path, ext, ocr_pdf=args.ocr_failed_pdf, ocr_max_pages=args.ocr_max_pages)
            text = clean_text(text)
            if text:
                slug = safe_slug(f"{idx:04d}-{row['relative_to_dest']}")
                md_path = out_dir / f"{slug}.md"
                md_path.write_text(build_markdown(row, text, method), encoding="utf-8")
                md_rel = str(md_path)
                status = "extracted"
                title = f"[历史标书] {row['relative_to_dest']}"
                docs.append(
                    KnowledgeDoc(
                        title=title,
                        content=md_path.read_text(encoding="utf-8"),
                        tag=TAG_NAME,
                        channel="historical-bidding",
                        source_kind="enterprise-wechat-wedrive",
                        item_key=f"historical-bidding:{row['sha256']}",
                    )
                )
            else:
                status = "empty_or_failed"
        elif ext in SKIP_TEXT_EXTENSIONS:
            status = "metadata_only"
        else:
            status = "unsupported"

        counters[status] += 1
        manifest_rows.append(
            {
                **row,
                "knowledge_status": status,
                "extraction_method": method,
                "text_chars": str(len(text)),
                "markdown_path": md_rel,
            }
        )
        if args.limit and len(docs) >= args.limit:
            break

    fieldnames = list(manifest_rows[0].keys()) if manifest_rows else []
    manifest_path = Path(args.manifest_csv)
    manifest_path.parent.mkdir(parents=True, exist_ok=True)
    with manifest_path.open("w", newline="", encoding="utf-8") as f:
        writer = csv.DictWriter(f, fieldnames=fieldnames)
        writer.writeheader()
        writer.writerows(manifest_rows)

    result: dict[str, object] = {
        "generated_at": dt.datetime.now().isoformat(timespec="seconds"),
        "documents_for_rag": len(docs),
        "status_counts": dict(counters),
        "manifest_csv": str(manifest_path),
        "markdown_dir": str(out_dir),
        "upsert": False,
    }

    if args.upsert and docs:
        client = WeKnoraClient(args.base_url)
        client.auto_setup()
        kb_id = ensure_kb(client, args.kb_name)
        tag_ids = ensure_tag(client)
        upserted = 0
        errors: list[dict[str, str]] = []
        for doc in docs:
            try:
                client.upsert_manual(kb_id, doc, tag_ids)
                upserted += 1
            except Exception as exc:
                errors.append({"title": doc.title, "error": str(exc)[:500]})
        result.update({"upsert": True, "knowledge_base": args.kb_name, "knowledge_base_id": kb_id, "upserted": upserted, "errors": errors})

    result_path = manifest_path.with_suffix(".result.json")
    result_path.write_text(json.dumps(result, ensure_ascii=False, indent=2), encoding="utf-8")
    print(json.dumps(result, ensure_ascii=False, indent=2))
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
