#!/usr/bin/env python3
"""Embed every image and every query with one model and save the vectors.

Output: <out>/<provider>.npz with arrays `img` (N x D), `txt` (N x D), plus
<out>/<provider>.json with timing and version metadata. Rows are in the order
of data/queries.json.

Local models run on CPU. Hosted models read credentials from the environment:
  AWS credentials (boto3 default chain)  amazon.titan-embed-image-v1,
                                         cohere.embed-v4:0,
                                         amazon.nova-2-multimodal-embeddings-v1:0
  JINA_API_KEY                           jina-clip-v2 via api.jina.ai
No credential is written to the output.

Usage:
  python run.py --data data --out results/2026-09-10 --provider clip-vit-b-32
"""
import argparse
import base64
import json
import os
import time

import numpy as np
from PIL import Image

BATCH = int(os.environ.get("EMB_BATCH", "16"))


def load_items(data: str) -> list[dict]:
    return json.load(open(os.path.join(data, "queries.json")))["items"]


def img_path(data: str, it: dict) -> str:
    return os.path.join(data, "images", it["file"])


# --- local open-weight models -------------------------------------------------


def run_open_clip(model_name: str, pretrained: str, data: str, items: list[dict]):
    import open_clip
    import torch

    model, _, preprocess = open_clip.create_model_and_transforms(
        model_name, pretrained=pretrained
    )
    tok = open_clip.get_tokenizer(model_name)
    model.eval()
    img_vecs, txt_vecs = [], []
    t0 = time.perf_counter()
    with torch.no_grad():
        for i in range(0, len(items), BATCH):
            batch = items[i : i + BATCH]
            ims = torch.stack(
                [preprocess(Image.open(img_path(data, it)).convert("RGB")) for it in batch]
            )
            img_vecs.append(model.encode_image(ims).float().numpy())
    t_img = time.perf_counter() - t0
    t0 = time.perf_counter()
    with torch.no_grad():
        for i in range(0, len(items), BATCH):
            batch = items[i : i + BATCH]
            txt_vecs.append(
                model.encode_text(tok([it["query"] for it in batch])).float().numpy()
            )
    t_txt = time.perf_counter() - t0
    return (
        np.concatenate(img_vecs),
        np.concatenate(txt_vecs),
        t_img,
        t_txt,
        f"open_clip {open_clip.__version__} {model_name}/{pretrained}, CPU fp32",
    )


def _feats(out):
    # transformers >=5 returns a model output; older versions return the tensor
    return out if hasattr(out, "numpy") else out.pooler_output


def run_siglip(repo: str, data: str, items: list[dict]):
    import torch
    import transformers
    from transformers import AutoModel, AutoProcessor

    model = AutoModel.from_pretrained(repo).eval()
    proc = AutoProcessor.from_pretrained(repo)
    img_vecs, txt_vecs = [], []
    t0 = time.perf_counter()
    with torch.no_grad():
        for i in range(0, len(items), BATCH):
            batch = items[i : i + BATCH]
            px = proc(
                images=[Image.open(img_path(data, it)).convert("RGB") for it in batch],
                return_tensors="pt",
            )
            img_vecs.append(_feats(model.get_image_features(**px)).float().numpy())
    t_img = time.perf_counter() - t0
    t0 = time.perf_counter()
    with torch.no_grad():
        for i in range(0, len(items), BATCH):
            batch = items[i : i + BATCH]
            tx = proc(
                text=[it["query"] for it in batch],
                padding="max_length",
                truncation=True,
                max_length=64,
                return_tensors="pt",
            )
            txt_vecs.append(_feats(model.get_text_features(**tx)).float().numpy())
    t_txt = time.perf_counter() - t0
    return (
        np.concatenate(img_vecs),
        np.concatenate(txt_vecs),
        t_img,
        t_txt,
        f"transformers {transformers.__version__} {repo}, CPU fp32",
    )


def run_jina_local(data: str, items: list[dict]):
    import torch
    import transformers
    from transformers import AutoModel

    repo = "jinaai/jina-clip-v2"
    model = AutoModel.from_pretrained(repo, trust_remote_code=True).eval()
    img_vecs, txt_vecs = [], []
    t0 = time.perf_counter()
    with torch.no_grad():
        for i in range(0, len(items), BATCH):
            batch = items[i : i + BATCH]
            img_vecs.append(
                np.asarray(
                    model.encode_image(
                        [img_path(data, it) for it in batch], truncate_dim=1024
                    )
                )
            )
    t_img = time.perf_counter() - t0
    t0 = time.perf_counter()
    with torch.no_grad():
        for i in range(0, len(items), BATCH):
            batch = items[i : i + BATCH]
            txt_vecs.append(
                np.asarray(
                    model.encode_text(
                        [it["query"] for it in batch], task="retrieval.query", truncate_dim=1024
                    )
                )
            )
    t_txt = time.perf_counter() - t0
    return (
        np.concatenate(img_vecs),
        np.concatenate(txt_vecs),
        t_img,
        t_txt,
        f"transformers {transformers.__version__} {repo} (trust_remote_code), 1024-d, CPU fp32",
    )


# --- hosted models --------------------------------------------------------------


def _b64(path: str) -> str:
    return base64.b64encode(open(path, "rb").read()).decode()


def _bedrock():
    import boto3

    return boto3.client("bedrock-runtime", region_name="us-east-1")


def run_titan(data: str, items: list[dict]):
    br = _bedrock()
    mid = "amazon.titan-embed-image-v1"
    img_vecs, txt_vecs = [], []
    t0 = time.perf_counter()
    for it in items:
        r = br.invoke_model(
            modelId=mid,
            body=json.dumps(
                {
                    "inputImage": _b64(img_path(data, it)),
                    "embeddingConfig": {"outputEmbeddingLength": 1024},
                }
            ),
        )
        img_vecs.append(json.loads(r["body"].read())["embedding"])
    t_img = time.perf_counter() - t0
    t0 = time.perf_counter()
    for it in items:
        r = br.invoke_model(
            modelId=mid,
            body=json.dumps(
                {"inputText": it["query"], "embeddingConfig": {"outputEmbeddingLength": 1024}}
            ),
        )
        txt_vecs.append(json.loads(r["body"].read())["embedding"])
    t_txt = time.perf_counter() - t0
    return np.array(img_vecs), np.array(txt_vecs), t_img, t_txt, f"Bedrock {mid}, 1024-d, us-east-1"


def run_cohere_v4(data: str, items: list[dict]):
    br = _bedrock()
    mid = "cohere.embed-v4:0"
    img_vecs, txt_vecs = [], []
    t0 = time.perf_counter()
    for it in items:
        r = br.invoke_model(
            modelId=mid,
            body=json.dumps(
                {
                    "images": [f"data:image/jpeg;base64,{_b64(img_path(data, it))}"],
                    "input_type": "image",
                    "embedding_types": ["float"],
                    "output_dimension": 1024,
                }
            ),
        )
        img_vecs.append(json.loads(r["body"].read())["embeddings"]["float"][0])
    t_img = time.perf_counter() - t0
    t0 = time.perf_counter()
    for i in range(0, len(items), 96):
        batch = items[i : i + 96]
        r = br.invoke_model(
            modelId=mid,
            body=json.dumps(
                {
                    "texts": [it["query"] for it in batch],
                    "input_type": "search_query",
                    "embedding_types": ["float"],
                    "output_dimension": 1024,
                }
            ),
        )
        txt_vecs.extend(json.loads(r["body"].read())["embeddings"]["float"])
    t_txt = time.perf_counter() - t0
    return np.array(img_vecs), np.array(txt_vecs), t_img, t_txt, f"Bedrock {mid}, 1024-d, us-east-1"


def run_nova_mm(data: str, items: list[dict]):
    br = _bedrock()
    mid = "amazon.nova-2-multimodal-embeddings-v1:0"
    img_vecs, txt_vecs = [], []
    t0 = time.perf_counter()
    for it in items:
        r = br.invoke_model(
            modelId=mid,
            body=json.dumps(
                {
                    "taskType": "SINGLE_EMBEDDING",
                    "singleEmbeddingParams": {
                        "embeddingPurpose": "GENERIC_INDEX",
                        "embeddingDimension": 1024,
                        "image": {
                            "format": "jpeg",
                            "source": {"bytes": _b64(img_path(data, it))},
                        },
                    },
                }
            ),
        )
        img_vecs.append(json.loads(r["body"].read())["embeddings"][0]["embedding"])
    t_img = time.perf_counter() - t0
    t0 = time.perf_counter()
    for it in items:
        r = br.invoke_model(
            modelId=mid,
            body=json.dumps(
                {
                    "taskType": "SINGLE_EMBEDDING",
                    "singleEmbeddingParams": {
                        "embeddingPurpose": "TEXT_RETRIEVAL",
                        "embeddingDimension": 1024,
                        "text": {"truncationMode": "END", "value": it["query"]},
                    },
                }
            ),
        )
        txt_vecs.append(json.loads(r["body"].read())["embeddings"][0]["embedding"])
    t_txt = time.perf_counter() - t0
    return np.array(img_vecs), np.array(txt_vecs), t_img, t_txt, f"Bedrock {mid}, 1024-d, us-east-1"


def run_jina_api(data: str, items: list[dict]):
    import requests

    key = os.environ["JINA_API_KEY"]
    h = {"Authorization": f"Bearer {key}", "Content-Type": "application/json"}
    url = "https://api.jina.ai/v1/embeddings"
    img_vecs, txt_vecs = [], []
    t0 = time.perf_counter()
    for i in range(0, len(items), 32):
        batch = items[i : i + 32]
        r = requests.post(
            url,
            headers=h,
            json={
                "model": "jina-clip-v2",
                "dimensions": 1024,
                "input": [{"image": _b64(img_path(data, it))} for it in batch],
            },
            timeout=300,
        )
        r.raise_for_status()
        img_vecs.extend(d["embedding"] for d in r.json()["data"])
    t_img = time.perf_counter() - t0
    t0 = time.perf_counter()
    for i in range(0, len(items), 32):
        batch = items[i : i + 32]
        r = requests.post(
            url,
            headers=h,
            json={
                "model": "jina-clip-v2",
                "dimensions": 1024,
                "task": "retrieval.query",
                "input": [{"text": it["query"]} for it in batch],
            },
            timeout=300,
        )
        r.raise_for_status()
        txt_vecs.extend(d["embedding"] for d in r.json()["data"])
    t_txt = time.perf_counter() - t0
    return np.array(img_vecs), np.array(txt_vecs), t_img, t_txt, "Jina API jina-clip-v2, 1024-d"


PROVIDERS = {
    "clip-vit-b-32": lambda d, it: run_open_clip("ViT-B-32", "openai", d, it),
    "clip-vit-l-14": lambda d, it: run_open_clip("ViT-L-14", "openai", d, it),
    "siglip2-base": lambda d, it: run_siglip("google/siglip2-base-patch16-256", d, it),
    "siglip2-so400m": lambda d, it: run_siglip("google/siglip2-so400m-patch16-256", d, it),
    "jina-clip-v2": run_jina_local,
    "jina-clip-v2-api": run_jina_api,
    "titan-multimodal": run_titan,
    "cohere-embed-v4": run_cohere_v4,
    "nova-multimodal": run_nova_mm,
}


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--data", default="data")
    ap.add_argument("--out", required=True)
    ap.add_argument("--provider", required=True, choices=sorted(PROVIDERS))
    a = ap.parse_args()
    items = load_items(a.data)
    os.makedirs(a.out, exist_ok=True)
    img, txt, t_img, t_txt, version = PROVIDERS[a.provider](a.data, items)
    np.savez_compressed(os.path.join(a.out, f"{a.provider}.npz"), img=img, txt=txt)
    meta = {
        "provider": a.provider,
        "version": version,
        "n_items": len(items),
        "dim": int(img.shape[1]),
        "image_encode_seconds": round(t_img, 2),
        "text_encode_seconds": round(t_txt, 2),
        "run_at": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
    }
    json.dump(meta, open(os.path.join(a.out, f"{a.provider}.json"), "w"), indent=1)
    print(json.dumps(meta))


if __name__ == "__main__":
    main()
