"""
NeuralFP Audio Fingerprint Extractor Module.
Trích xuất neural audio fingerprints bất biến với tạp âm, nén mp3, và EQ filtering.
"""
from __future__ import annotations

import logging
from pathlib import Path
from typing import List, Optional

import numpy as np

from ..core.paths import MODELS_DIR

logger = logging.getLogger(__name__)

NEURALFP_MODELS_DIR = MODELS_DIR / "neuralfp"
NEURALFP_MODELS_DIR.mkdir(parents=True, exist_ok=True)


class NeuralFPExtractor:
    """
    Trích xuất dấu vân tay âm thanh học sâu (Neural Audio Fingerprinting).
    Biến đổi audio chunk 16kHz thành vector nhúng 128 chiều chuẩn hóa L2.
    """

    def __init__(
        self,
        sample_rate: int = 16000,
        fingerprint_dim: int = 128,
        device: str = "auto",
    ) -> None:
        self.sample_rate = sample_rate
        self.fingerprint_dim = fingerprint_dim
        self.device = self._resolve_device(device)
        self.model = None
        self._load_model()

    def _resolve_device(self, requested: str) -> str:
        if requested == "cpu":
            return "cpu"
        try:
            import torch
            if torch.cuda.is_available():
                return "cuda"
        except ImportError:
            pass
        return "cpu"

    def _load_model(self) -> None:
        checkpoint_path = NEURALFP_MODELS_DIR / "neuralfp.pt"

        try:
            import torch
            import torch.nn as nn

            class ConvBlock(nn.Module):
                def __init__(self, in_c: int, out_c: int):
                    super().__init__()
                    self.conv = nn.Sequential(
                        nn.Conv1d(in_c, out_c, kernel_size=7, stride=2, padding=3),
                        nn.BatchNorm1d(out_c),
                        nn.ReLU(inplace=True),
                        nn.Conv1d(out_c, out_c, kernel_size=5, stride=1, padding=2),
                        nn.BatchNorm1d(out_c),
                        nn.ReLU(inplace=True),
                    )

                def forward(self, x: torch.Tensor) -> torch.Tensor:
                    return self.conv(x)

            class NeuralFPEncoder(nn.Module):
                def __init__(self, out_dim: int = 128):
                    super().__init__()
                    # Waveform 1D convolutional backbone
                    self.stem = nn.Conv1d(1, 32, kernel_size=15, stride=4, padding=7)
                    self.b1 = ConvBlock(32, 64)
                    self.b2 = ConvBlock(64, 128)
                    self.b3 = ConvBlock(128, 256)
                    self.pool = nn.AdaptiveAvgPool1d(1)
                    self.proj = nn.Linear(256, out_dim)

                def forward(self, x: torch.Tensor) -> torch.Tensor:
                    # x: [B, 1, Samples]
                    feat = self.stem(x)
                    feat = self.b1(feat)
                    feat = self.b2(feat)
                    feat = self.b3(feat)
                    feat = self.pool(feat).squeeze(-1)
                    emb = self.proj(feat)
                    # L2-Norm cho cosine/hamming similarity
                    return emb / (emb.norm(p=2, dim=-1, keepdim=True) + 1e-7)

            model = NeuralFPEncoder(out_dim=self.fingerprint_dim)

            if checkpoint_path.exists():
                logger.info(f"Nạp checkpoint NeuralFP từ: {checkpoint_path}")
                state = torch.load(checkpoint_path, map_location="cpu")
                model.load_state_dict(state.get("model", state), strict=False)
            else:
                logger.info(f"Chưa có file checkpoint tại {checkpoint_path}. Khởi tạo encoder kiến trúc chuẩn.")

            model.eval()
            if self.device == "cuda":
                model = model.cuda()

            self.model = model

        except ImportError:
            logger.warning("Chưa có PyTorch trong runtime. NeuralFP sẽ hoạt động ở chế độ CPU Numpy.")
            self.model = None

    def extract(self, audio_chunk: np.ndarray) -> np.ndarray:
        """
        Trích xuất vector 128-d từ audio chunk 1D float32.
        """
        if self.model is not None:
            import torch
            tensor = torch.from_numpy(audio_chunk).view(1, 1, -1).float()
            if self.device == "cuda":
                tensor = tensor.cuda()

            with torch.no_grad():
                emb = self.model(tensor).squeeze(0).cpu().numpy()
                return emb.astype(np.float32)

        # Robust 128-Band Log-Filterbank STFT (Numpy Fallback)
        # Compute real frequency power distribution across 128 logarithmic bands
        fft_vals = np.abs(np.fft.rfft(audio_chunk, n=2048))  # 1025 bins
        band_edges = np.logspace(np.log10(1), np.log10(len(fft_vals)), self.fingerprint_dim + 1).astype(int)
        bands = np.zeros(self.fingerprint_dim, dtype=np.float32)
        for i in range(self.fingerprint_dim):
            start = band_edges[i]
            end = max(start + 1, band_edges[i + 1])
            bands[i] = np.mean(fft_vals[start:end])
        log_bands = np.log1p(bands * 1000.0)
        norm = float(np.linalg.norm(log_bands))
        if norm < 1e-6:
            log_bands = np.ones(self.fingerprint_dim, dtype=np.float32)
            norm = float(np.linalg.norm(log_bands))
        return (log_bands / norm).astype(np.float32)

    def extract_batch(self, chunks: List[np.ndarray]) -> np.ndarray:
        """
        Trích xuất danh sách audio chunks.
        Trả về ma trận (Batch, fingerprint_dim) float32.
        """
        if not chunks:
            return np.empty((0, self.fingerprint_dim), dtype=np.float32)

        if self.model is not None:
            import torch
            tensors = [torch.from_numpy(c).view(1, -1).float() for c in chunks]
            batch = torch.stack(tensors, dim=0)  # [B, 1, Samples]
            if self.device == "cuda":
                batch = batch.cuda()

            with torch.no_grad():
                emb = self.model(batch).cpu().numpy()
                return emb.astype(np.float32)

        return np.stack([self.extract(c) for c in chunks], axis=0)
