#!/usr/bin/env python3
"""
rag_hybrid.py — Búsqueda híbrida (FTS5 + cosine similarity) para la base de conocimiento RAG.

Combina lo mejor de dos mundos:
  - FTS5: búsqueda textual exacta con stemming y diacríticos
  - Embeddings: búsqueda semántica por cosine similarity

La fórmula de combinación es:
  score_total = alpha * sigmoid(fts_score/10) + (1 - alpha) * cosine_sim

Uso:
  python3 rag_hybrid.py "cómo configurar un proxy inverso"
  python3 rag_hybrid.py "hooks de git" --alpha 0.6 --limite 10
"""

import argparse
import math
import os
import sqlite3
import sys
from typing import Any

import numpy as np
from numpy.linalg import norm

DB_PATH = os.path.expanduser("~/.cerebro/rag_conocimiento.db")
EMBEDDING_DIM = 1024
EMBEDDING_BYTES = EMBEDDING_DIM * 4  # 4096 bytes


def blob_to_vector(blob: bytes) -> np.ndarray:
    if len(blob) != EMBEDDING_BYTES:
        raise ValueError(f"BLOB incorrecto: {len(blob)} bytes (esperados {EMBEDDING_BYTES})")
    return np.frombuffer(blob, dtype=np.float32).copy()


def normalizar(vector: np.ndarray) -> np.ndarray:
    norma = norm(vector)
    return vector / norma if norma > 0 else vector


def cosine_similarity(a: np.ndarray, b: np.ndarray) -> float:
    return float(np.dot(a, b))


def sigmoid_fts(fts_score: float) -> float:
    return 1.0 / (1.0 + math.exp(-fts_score / 10.0))


def obtener_embedding_ollama(texto: str) -> np.ndarray:
    import requests
    response = requests.post(
        "http://localhost:11434/api/embeddings",
        json={"model": "bge-m3", "prompt": texto},
        timeout=30,
    )
    response.raise_for_status()
    return np.array(response.json()["embedding"], dtype=np.float32)


def hybrid_search(query: str, alpha: float = 0.4, top_k: int = 10, tag: str | None = None) -> list[dict]:
    query_emb = obtener_embedding_ollama(query)
    query_emb = normalizar(query_emb)

    if not os.path.exists(DB_PATH):
        print(f"Error: BD no encontrada en {DB_PATH}", file=sys.stderr)
        return []

    conn = sqlite3.connect(DB_PATH)
    conn.row_factory = sqlite3.Row

    where_clauses: list[str] = ["chunks_fts MATCH ?"]
    params: list[Any] = [query]

    if tag:
        where_clauses.append("c.tags LIKE ?")
        params.append(f"%{tag}%")

    sql = f"""
        SELECT c.id, c.content, e.vector, c.doc_id,
               bm25(chunks_fts) as fts_score
        FROM chunks_fts
        JOIN chunks c ON chunks_fts.rowid = c.id
        JOIN embeddings e ON e.chunk_id = c.id
        WHERE {" AND ".join(where_clauses)}
        ORDER BY fts_score DESC
        LIMIT ?
    """
    params.append(top_k * 2)

    cursor = conn.execute(sql, params)
    rows = cursor.fetchall()
    conn.close()

    if not rows:
        print(f"Sin resultados para: {query}")
        return []

    results = []
    for row in rows:
        try:
            emb_vec = blob_to_vector(row["vector"])
        except ValueError as e:
            print(f"[aviso] {e}", file=sys.stderr)
            continue

        emb_vec = normalizar(emb_vec)
        cos_sim = cosine_similarity(query_emb, emb_vec)
        fts_norm = sigmoid_fts(row["fts_score"])
        score_total = alpha * fts_norm + (1.0 - alpha) * cos_sim

        results.append({
            "contenido": row["content"],
            "doc_id": row["doc_id"],
            "score": score_total,
            "fts_score": round(row["fts_score"], 4),
            "cos_sim": round(cos_sim, 4),
        })

    results.sort(key=lambda x: x["score"], reverse=True)
    return results[:top_k]


def main():
    parser = argparse.ArgumentParser(description="Búsqueda híbrida FTS5 + embeddings")
    parser.add_argument("consulta", help="Texto de la consulta")
    parser.add_argument("--alpha", type=float, default=0.4, help="Peso de FTS5 (0-1)")
    parser.add_argument("--limite", type=int, default=10, help="Número de resultados")
    parser.add_argument("--tag", help="Filtrar por etiqueta")
    args = parser.parse_args()

    resultados = hybrid_search(query=args.consulta, alpha=args.alpha, top_k=args.limite, tag=args.tag)
    if not resultados:
        sys.exit(0)

    print(f"\n{= * 70}")
    print(f"  Consulta: {args.consulta}  |  Alpha: {args.alpha}  |  Resultados: {len(resultados)}")
    print(f"{= * 70}\n")

    for i, r in enumerate(resultados, 1):
        preview = r["contenido"].replace("\n", " ")[:120]
        print(f"  [{i:2d}] Score: {r[score]:.4f}  (FTS: {r[fts_score]:.4f}  |  Cos: {r[cos_sim]:.4f})")
        print(f"       {preview}\n")

    print(f"{= * 70}")


if __name__ == "__main__":
    main()