#!/usr/bin/env python3
"""
build_index.py - One-time indexing for CIE-10 semantic search.

Uses OpenAI text-embedding-3-large at 1024 dims for high-quality
Spanish medical embeddings.

Usage:
    OPENAI_API_KEY=$(cat /etc/cie10/openai.key) \
        python3 build_index.py Buscador_CIE10_Optimizado.html cie10_embeddings.json

Cost: ~$0.05 USD for 12,811 descriptions.
Runtime: ~5-8 minutes.
"""

import json
import os
import re
import sys
import time
from pathlib import Path

from openai import OpenAI

MODEL = "text-embedding-3-large"
DIMENSIONS = 1024
BATCH_SIZE = 200


def extract_data_array(html_path: Path) -> list[dict]:
    html = html_path.read_text(encoding="utf-8")
    match = re.search(r"const\s+data\s*=\s*(\[.*?\]);", html, re.DOTALL)
    if not match:
        raise RuntimeError("Could not locate `const data = [...]` block in HTML.")
    return json.loads(match.group(1))


def embed_batch(client: OpenAI, texts: list[str]) -> list[list[float]]:
    for attempt in range(5):
        try:
            response = client.embeddings.create(
                model=MODEL,
                input=texts,
                dimensions=DIMENSIONS,
            )
            return [item.embedding for item in response.data]
        except Exception as exc:
            wait = 2 ** attempt
            print(f"  ! batch failed ({exc}), retrying in {wait}s...", file=sys.stderr)
            time.sleep(wait)
    raise RuntimeError("Batch failed after 5 retries.")


def main() -> None:
    if len(sys.argv) != 3:
        print("Usage: build_index.py <input.html> <output.json>", file=sys.stderr)
        sys.exit(1)

    input_html = Path(sys.argv[1])
    output_json = Path(sys.argv[2])

    api_key = os.environ.get("OPENAI_API_KEY")
    if not api_key:
        print("ERROR: set OPENAI_API_KEY environment variable.", file=sys.stderr)
        sys.exit(1)

    client = OpenAI(api_key=api_key)

    print(f"Extracting data array from {input_html}...")
    entries = extract_data_array(input_html)
    print(f"  found {len(entries)} entries")

    print(f"Embedding with {MODEL} ({DIMENSIONS} dims), batch size {BATCH_SIZE}...")
    start = time.time()
    embedded: list[dict] = []

    for i in range(0, len(entries), BATCH_SIZE):
        batch = entries[i : i + BATCH_SIZE]
        texts = [item["descripcion"] for item in batch]
        vectors = embed_batch(client, texts)
        for item, vector in zip(batch, vectors):
            embedded.append(
                {
                    "codigo": item["codigo"],
                    "descripcion": item["descripcion"],
                    "embedding": [round(float(x), 5) for x in vector],
                }
            )
        done = min(i + BATCH_SIZE, len(entries))
        elapsed = time.time() - start
        rate = done / elapsed if elapsed else 0
        eta = (len(entries) - done) / rate if rate else 0
        print(f"  {done}/{len(entries)}  ({rate:.0f}/s, ETA {eta:.0f}s)")

    print(f"Writing {output_json}...")
    output_json.write_text(json.dumps(embedded, ensure_ascii=False, separators=(",", ":")))
    size_mb = output_json.stat().st_size / 1024 / 1024
    print(f"Done. {len(embedded)} entries, {size_mb:.1f} MB, {time.time() - start:.0f}s total.")


if __name__ == "__main__":
    main()
