"""CLIP image embeddings (lazy singleton). Shared by match API and scripts."""

from __future__ import annotations

import io
import logging
from threading import Lock
from typing import Optional

import numpy as np
import torch
from PIL import Image
from transformers import CLIPModel, CLIPProcessor

logger = logging.getLogger(__name__)

_DEFAULT_MODEL = "openai/clip-vit-base-patch32"

_lock = Lock()
_model: Optional[CLIPModel] = None
_processor: Optional[CLIPProcessor] = None
_device: Optional[str] = None
_model_name: Optional[str] = None


def _ensure_loaded(model_name: str = _DEFAULT_MODEL) -> tuple[CLIPModel, CLIPProcessor, str]:
    global _model, _processor, _device, _model_name
    with _lock:
        if _model is not None and _processor is not None and _device is not None:
            if _model_name == model_name:
                return _model, _processor, _device
        device = "cuda" if torch.cuda.is_available() else "cpu"
        logger.info("Loading CLIP model %s on %s", model_name, device)
        model = CLIPModel.from_pretrained(model_name).to(device)
        processor = CLIPProcessor.from_pretrained(model_name)
        model.eval()
        _model = model
        _processor = processor
        _device = device
        _model_name = model_name
        return _model, _processor, _device


def _features_to_tensor(outputs: object) -> torch.Tensor:
    # transformers>=5: get_image_features may return BaseModelOutputWithPooling.
    if hasattr(outputs, "pooler_output") and outputs.pooler_output is not None:
        return outputs.pooler_output
    if torch.is_tensor(outputs):
        return outputs
    if isinstance(outputs, (tuple, list)):
        return outputs[1] if len(outputs) > 1 else outputs[0]
    raise TypeError(f"Unexpected CLIP output type: {type(outputs)!r}")


def embed_image_bytes(image_bytes: bytes, model_name: str = _DEFAULT_MODEL) -> np.ndarray:
    model, processor, device = _ensure_loaded(model_name)
    image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
    inputs = processor(images=image, return_tensors="pt")
    inputs = {k: v.to(device) for k, v in inputs.items()}

    with torch.no_grad():
        outputs = model.get_image_features(**inputs)
        image_features = _features_to_tensor(outputs)
        image_features = image_features / image_features.norm(dim=-1, keepdim=True)

    return image_features.cpu().numpy().reshape(-1).astype(np.float64)


def embed_image_path(path: str, model_name: str = _DEFAULT_MODEL) -> np.ndarray:
    with open(path, "rb") as f:
        return embed_image_bytes(f.read(), model_name=model_name)
