from __future__ import annotations

import argparse
import logging
import re
import sys
from pathlib import Path

import pdfplumber
import tiktoken
from dotenv import load_dotenv

ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT))
load_dotenv(ROOT / ".env")

from app.config import get_settings  # noqa: E402
from app.db import close_pool, get_connection, init_pool  # noqa: E402
from app.embeddings import embed_texts  # noqa: E402
from app.vertex import init_vertex  # noqa: E402

log = logging.getLogger("ingest")

EMBED_BATCH = 16


def extract_pdf_text(pdf_path: Path) -> str:
    pages: list[str] = []
    with pdfplumber.open(pdf_path) as pdf:
        for i, page in enumerate(pdf.pages, start=1):
            raw = page.extract_text() or ""
            cleaned = re.sub(r"[ \t]+", " ", raw)
            cleaned = re.sub(r"\n{3,}", "\n\n", cleaned).strip()
            if cleaned:
                pages.append(f"[Page {i}]\n{cleaned}")
    if not pages:
        raise RuntimeError(f"No extractable text in {pdf_path}")
    return "\n\n".join(pages)


def chunk_text(text: str, chunk_tokens: int, overlap: int) -> list[str]:
    enc = tiktoken.get_encoding("cl100k_base")
    token_ids = enc.encode(text)
    if chunk_tokens <= overlap:
        raise ValueError("CHUNK_TOKENS must be greater than CHUNK_OVERLAP.")

    chunks: list[str] = []
    start = 0
    n = len(token_ids)
    while start < n:
        end = min(start + chunk_tokens, n)
        piece = enc.decode(token_ids[start:end]).strip()
        if piece:
            chunks.append(piece)
        if end == n:
            break
        start += chunk_tokens - overlap
    return chunks


def replace_chunks(source_file: str, chunks: list[str], embeddings: list[list[float]]) -> None:
    with get_connection() as conn:
        with conn.cursor() as cur:
            cur.execute(
                "DELETE FROM document_chunks WHERE source_file = %s",
                (source_file,),
            )
            for index, (text, vector) in enumerate(zip(chunks, embeddings, strict=True)):
                cur.execute(
                    """
                    INSERT INTO document_chunks (text_chunk, embedding, source_file, chunk_index)
                    VALUES (%s, %s, %s, %s)
                    """,
                    (text, vector, source_file, index),
                )
        conn.commit()


def batched(items: list[str], size: int) -> list[list[str]]:
    return [items[i : i + size] for i in range(0, len(items), size)]


def main() -> None:
    logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
    parser = argparse.ArgumentParser(description="Ingest a local PDF into PartnerLogic.")
    parser.add_argument("--pdf", default=None, help="Path to the source PDF.")
    args = parser.parse_args()

    settings = get_settings()
    pdf_path = Path(args.pdf or settings.pdf_path).expanduser().resolve()
    if not pdf_path.exists():
        raise SystemExit(
            f"PDF not found: {pdf_path}\n"
            "Place the document at data/shareholder_agreement.pdf or run:\n"
            "  python scripts/generate_sample_pdf.py"
        )

    log.info("Extracting %s", pdf_path)
    text = extract_pdf_text(pdf_path)
    chunks = chunk_text(text, settings.chunk_tokens, settings.chunk_overlap)
    log.info("Split into %s chunks (%s-token window, %s overlap)", len(chunks), settings.chunk_tokens, settings.chunk_overlap)

    init_vertex()
    init_pool()
    try:
        embeddings: list[list[float]] = []
        for batch_no, batch in enumerate(batched(chunks, EMBED_BATCH), start=1):
            log.info("Embedding batch %s (%s chunks)", batch_no, len(batch))
            embeddings.extend(embed_texts(batch, task_type="RETRIEVAL_DOCUMENT"))

        source_file = pdf_path.name
        replace_chunks(source_file, chunks, embeddings)
        log.info("Wrote %s chunks to document_chunks (source_file=%s)", len(chunks), source_file)
    finally:
        close_pool()


if __name__ == "__main__":
    main()
