#!/usr/bin/env python3
"""Ingest final_finance_dataset.csv into ChromaDB for RAG retrieval."""

import argparse
import json
import sys
from pathlib import Path

import pandas as pd
from langchain_core.documents import Document
from tqdm import tqdm

sys.path.insert(0, str(Path(__file__).resolve().parent.parent))

from app.config import get_settings
from app.rag.chain import RAGService


def row_to_document(row: pd.Series, max_chars: int = 3000) -> Document | None:
    text = str(row.get("text_content", "")).strip()
    if not text:
        return None
    if len(text) > max_chars:
        text = text[:max_chars] + "..."

    metadata = {
        "record_id": str(row.get("record_id", "")),
        "data_category": str(row.get("data_category", "")),
        "source_dataset": str(row.get("source_dataset", "")),
    }

    raw_meta = row.get("metadata_json")
    if pd.notna(raw_meta):
        try:
            parsed = json.loads(raw_meta)
            if isinstance(parsed, dict):
                for key in ("Name", "Industry", "NSE Code"):
                    if key in parsed and parsed[key] not in (None, ""):
                        metadata[key.lower().replace(" ", "_")] = str(parsed[key])
        except json.JSONDecodeError:
            pass

    return Document(page_content=text, metadata=metadata)


def count_rows(csv_path: Path) -> int:
    with open(csv_path, encoding="utf-8") as f:
        return sum(1 for _ in f) - 1


def get_indexed_record_ids(collection) -> set[str]:
    indexed: set[str] = set()
    offset = 0
    page_size = 5000

    while True:
        batch = collection.get(include=["metadatas"], limit=page_size, offset=offset)
        metadatas = batch.get("metadatas") or []
        if not metadatas:
            break

        for meta in metadatas:
            if meta and meta.get("record_id"):
                indexed.add(meta["record_id"])

        if len(metadatas) < page_size:
            break
        offset += page_size

    return indexed


def ingest(
    batch_size: int = 1,
    limit: int | None = None,
    reset: bool = False,
    resume: bool = False,
    log_each: bool = False,
) -> None:
    settings = get_settings()
    csv_path = settings.dataset_path

    if not csv_path.exists():
        raise FileNotFoundError(f"Dataset not found: {csv_path}")

    if reset and settings.chroma_persist_dir.exists():
        import shutil

        shutil.rmtree(settings.chroma_persist_dir)
        print(f"Cleared existing vector store at {settings.chroma_persist_dir}")

    total_rows = count_rows(csv_path)
    target_rows = min(total_rows, limit) if limit else total_rows

    service = RAGService(settings)
    store = service.get_vectorstore()
    collection = store._collection

    already_indexed = get_indexed_record_ids(collection) if resume and not reset else set()
    if already_indexed:
        print(f"Resuming — {len(already_indexed):,} records already indexed, skipping those.")

    mode = "one-by-one" if batch_size == 1 else f"batches of {batch_size}"
    print(f"Indexing up to {target_rows:,} records from {csv_path.name} ({mode})...")

    indexed = 0
    skipped = 0
    read = 0

    with tqdm(total=target_rows, desc="Embedding & indexing", unit="docs") as progress:
        if already_indexed:
            progress.update(min(len(already_indexed), target_rows))

        for chunk in pd.read_csv(csv_path, chunksize=batch_size, low_memory=False):
            if limit and read >= limit:
                break

            if limit and read + len(chunk) > limit:
                chunk = chunk.head(limit - read)

            read += len(chunk)

            for _, row in chunk.iterrows():
                record_id = str(row.get("record_id", ""))
                if resume and record_id in already_indexed:
                    skipped += 1
                    continue

                doc = row_to_document(row)
                if not doc:
                    continue

                store.add_documents([doc])
                indexed += 1
                already_indexed.add(record_id)
                progress.update(1)

                if log_each or batch_size == 1:
                    category = doc.metadata.get("data_category", "")
                    tqdm.write(f"  [{indexed + skipped}/{target_rows}] indexed {record_id} ({category})")

    final_count = collection.count()
    print(f"Done. Added {indexed:,} new documents, skipped {skipped:,} existing.")
    print(f"ChromaDB collection '{settings.chroma_collection}' now has {final_count:,} documents.")


def main() -> None:
    parser = argparse.ArgumentParser(description="Ingest finance CSV into ChromaDB")
    parser.add_argument(
        "--batch-size",
        type=int,
        default=1,
        help="Records per embedding batch (default: 1 = one-by-one)",
    )
    parser.add_argument("--limit", type=int, default=None, help="Max rows to read from CSV")
    parser.add_argument("--reset", action="store_true", help="Clear existing vector store first")
    parser.add_argument(
        "--resume",
        action="store_true",
        help="Skip records already indexed (by record_id)",
    )
    parser.add_argument("--log-each", action="store_true", help="Print every indexed record_id")
    args = parser.parse_args()

    ingest(
        batch_size=args.batch_size,
        limit=args.limit,
        reset=args.reset,
        resume=args.resume,
        log_each=args.log_each,
    )


if __name__ == "__main__":
    main()
