from __future__ import annotations

import argparse
import hashlib
import json
import os
from pathlib import Path

from langchain_community.vectorstores import FAISS

from app.soat_crosswalk.infrastructure.catalogs import JsonCupsSoatCrosswalk
from app.soat_tariffs.infrastructure.json_catalog import JsonSoatTariffCatalog
from modules.processing.embeddings import get_hf_embeddings


PROJECT_ROOT = Path(__file__).resolve().parents[3]
DEFAULT_OUTPUT = PROJECT_ROOT / "soat_procedures_faiss"


def build_soat_procedures_index(output_dir: str | Path = DEFAULT_OUTPUT) -> Path:
    tariff_catalog = JsonSoatTariffCatalog()
    catalog = tariff_catalog.load(2026)
    crosswalk = JsonCupsSoatCrosswalk().load()
    aliases: dict[str, list[str]] = {}
    for relationship in crosswalk.relationships:
        aliases.setdefault(relationship.soat_code, []).append(relationship.cups_code)

    texts: list[str] = []
    metadatas: list[dict[str, object]] = []
    for entry in catalog.entries:
        cups_aliases = sorted(set(aliases.get(entry.code, [])))
        alias_text = f" Alias CUPS confirmados: {', '.join(cups_aliases)}." if cups_aliases else ""
        texts.append(
            f"SOAT {entry.code}. {entry.description}. Grupo quirúrgico {entry.surgical_group}."
            f"{alias_text}"
        )
        metadatas.append(
            {
                "codigo_soat": entry.code,
                "descripcion": entry.description,
                "grupo_quirurgico": entry.surgical_group,
                "cups_aliases": cups_aliases,
                "catalog_version": catalog.version,
                "crosswalk_version": crosswalk.version,
            }
        )
    embeddings = get_hf_embeddings(os.getenv("EMBEDDINGS_MODEL", "intfloat/e5-small-v2"))
    vectorstore = FAISS.from_texts(texts, embedding=embeddings, metadatas=metadatas)
    output = Path(output_dir)
    output.mkdir(parents=True, exist_ok=True)
    vectorstore.save_local(str(output))
    artifacts = {
        filename: hashlib.sha256((output / filename).read_bytes()).hexdigest()
        for filename in ("index.faiss", "index.pkl")
    }
    (output / "manifest.json").write_text(
        json.dumps(
            {
                "version": "soat-procedures-2026-v1",
                "tariff_catalog_version": catalog.version,
                "crosswalk_version": crosswalk.version,
                "documents": len(texts),
                "artifacts": artifacts,
            },
            indent=2,
            sort_keys=True,
        )
        + "\n",
        encoding="utf-8",
    )
    return output


def main() -> None:
    parser = argparse.ArgumentParser(description="Construye el FAISS granular de procedimientos SOAT")
    parser.add_argument("--out", type=Path, default=DEFAULT_OUTPUT)
    args = parser.parse_args()
    output = build_soat_procedures_index(args.out)
    print(f"Índice granular SOAT generado en: {output}")


if __name__ == "__main__":
    main()
