#!/usr/bin/env python3
# Author: Tom Sapletta · https://tom.sapletta.com
# Part of the ifURI solution.

"""Use host compute for OCR while files live on a remote node.

Flow:

1. Host crawls the node sandbox with ``fs://host/dir/query/list``.
2. Host asks its local OCR connector to run ``document/query/text_from_uri``.
3. OCR connector fetches every file through ``fs://host/file/query/blob``.
4. OCR runs on the host and reports are written locally.
"""

from __future__ import annotations

import argparse
import csv
import json
import sys
import time
import urllib.error
import urllib.request
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable

REPO = Path(__file__).resolve().parents[2]
for candidate in (REPO / "urirun-connector-ocr", REPO / "urirun/adapters/python"):
    if str(candidate) not in sys.path:
        sys.path.insert(0, str(candidate))

from urirun_connector_ocr import document_text_from_uri  # noqa: E402

DEFAULT_NODE_URL = "http://192.168.188.201:8765"
DEFAULT_EXTENSIONS = "pdf,png,jpg,jpeg,txt,md"


@dataclass
class NodeClient:
    base_url: str
    timeout: int = 120

    def run(self, uri: str, payload: dict[str, Any]) -> dict[str, Any]:
        body = json.dumps({"uri": uri, "payload": payload}).encode("utf-8")
        request = urllib.request.Request(
            self.base_url.rstrip("/") + "/run",
            data=body,
            headers={"Content-Type": "application/json"},
            method="POST",
        )
        try:
            with urllib.request.urlopen(request, timeout=self.timeout) as response:
                return json.loads(response.read().decode("utf-8") or "{}")
        except urllib.error.HTTPError as exc:
            raw = exc.read().decode("utf-8", "replace") if exc.fp else ""
            try:
                return json.loads(raw or "{}")
            except json.JSONDecodeError:
                return {"ok": False, "status": exc.code, "error": raw}


def route_value(envelope: dict[str, Any]) -> dict[str, Any]:
    result = envelope.get("result") or {}
    if isinstance(result, dict) and isinstance(result.get("value"), dict):
        return result["value"]
    return result if isinstance(result, dict) else {}


def split_extensions(raw: str) -> set[str]:
    return {
        item if item.startswith(".") else f".{item}"
        for item in (part.strip().lower() for part in raw.replace(";", ",").split(","))
        if item
    }


def join_path(base: str, name: str) -> str:
    return name if base in {"", "."} else f"{base.rstrip('/')}/{name}"


def crawl_files(
    client: NodeClient,
    *,
    start: str = ".",
    extensions: str = DEFAULT_EXTENSIONS,
    max_files: int = 100,
    max_depth: int = 8,
) -> list[dict[str, Any]]:
    allowed = split_extensions(extensions)
    files: list[dict[str, Any]] = []
    queue: list[tuple[str, int]] = [(start, 0)]
    while queue and len(files) < max_files:
        current, depth = queue.pop(0)
        env = client.run("fs://host/dir/query/list", {"path": current})
        data = route_value(env)
        if not data.get("ok"):
            raise RuntimeError(f"dir list failed for {current}: {data.get('error') or env}")
        for entry in data.get("entries") or []:
            name = str(entry.get("name") or "")
            path = join_path(current, name)
            kind = entry.get("type")
            if kind == "dir" and depth < max_depth:
                queue.append((path, depth + 1))
            elif kind == "file" and Path(name).suffix.lower() in allowed:
                files.append({"path": path, "size": entry.get("size")})
                if len(files) >= max_files:
                    break
    return files


def run_host_compute(
    args: argparse.Namespace,
    *,
    client: NodeClient | None = None,
    ocr_func: Callable[..., dict[str, Any]] = document_text_from_uri,
) -> dict[str, Any]:
    client = client or NodeClient(args.node_url, timeout=args.timeout)
    started = time.time()
    files = crawl_files(
        client,
        start=args.start_path,
        extensions=args.extensions,
        max_files=args.max_files,
        max_depth=args.max_depth,
    )
    rows: list[dict[str, Any]] = []
    for item in files:
        path = item["path"]
        result = ocr_func(
            source_node_url=args.node_url,
            source_uri="fs://host/file/query/blob",
            source_payload_json=json.dumps({"path": path, "max_bytes": args.max_blob_bytes}),
            backend=args.backend,
            max_chars=args.max_chars,
            max_input_bytes=args.max_blob_bytes,
            timeout=args.timeout,
        )
        rows.append({
            "path": path,
            "size": item.get("size", ""),
            "ok": bool(result.get("ok")),
            "backend": result.get("backend", ""),
            "chars": result.get("chars", 0),
            "sha256": (result.get("source") or {}).get("sha256", ""),
            "error": result.get("error", ""),
            "preview": " ".join(str(result.get("text") or "").split())[:240],
        })

    ok_count = sum(1 for row in rows if row["ok"])
    payload = {
        "ok": True,
        "node_url": args.node_url,
        "start_path": args.start_path,
        "elapsed_seconds": round(time.time() - started, 3),
        "count": len(rows),
        "ok_count": ok_count,
        "failed_count": len(rows) - ok_count,
        "rows": rows,
    }
    payload["reports"] = write_reports(Path(args.output_dir), payload)
    return payload


def write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    fields = ["path", "size", "ok", "backend", "chars", "sha256", "error", "preview"]
    with path.open("w", encoding="utf-8", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=fields)
        writer.writeheader()
        for row in rows:
            writer.writerow({field: row.get(field, "") for field in fields})


def write_reports(output_dir: Path, payload: dict[str, Any]) -> dict[str, str]:
    output_dir.mkdir(parents=True, exist_ok=True)
    raw = output_dir / "host_compute_ocr_raw.json"
    csv_path = output_dir / "host_compute_ocr.csv"
    md_path = output_dir / "host_compute_ocr.md"
    raw.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")
    write_csv(csv_path, payload["rows"])
    md_path.write_text(
        "\n".join([
            "# Host Compute OCR",
            "",
            f"- node: `{payload['node_url']}`",
            f"- start path: `{payload['start_path']}`",
            f"- files: {payload['count']}",
            f"- OCR ok: {payload['ok_count']}",
            f"- OCR failed: {payload['failed_count']}",
            "",
        ]),
        encoding="utf-8",
    )
    return {"raw_json": str(raw), "csv": str(csv_path), "markdown": str(md_path)}


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description="Run OCR on the host for files fetched from a node over URI.")
    parser.add_argument("--node-url", default=DEFAULT_NODE_URL)
    parser.add_argument("--start-path", default=".")
    parser.add_argument("--extensions", default=DEFAULT_EXTENSIONS)
    parser.add_argument("--max-files", type=int, default=25)
    parser.add_argument("--max-depth", type=int, default=8)
    parser.add_argument("--max-chars", type=int, default=1200)
    parser.add_argument("--max-blob-bytes", type=int, default=10 * 1024 * 1024)
    parser.add_argument("--backend", default="auto")
    parser.add_argument("--timeout", type=int, default=180)
    parser.add_argument("--output-dir", default=".state")
    return parser


def main(argv: list[str] | None = None) -> int:
    args = build_parser().parse_args(argv)
    result = run_host_compute(args)
    print(json.dumps({
        "ok": True,
        "count": result["count"],
        "ok_count": result["ok_count"],
        "failed_count": result["failed_count"],
        "reports": result["reports"],
    }, ensure_ascii=False, indent=2))
    return 0


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