"""RAG retrieval: similarity search over user messages and history lookups.""" from __future__ import annotations import logging import numpy as np from vibe_bot import llm_client from vibe_bot.config import ( EMBEDDING_ENDPOINT, EMBEDDING_ENDPOINT_KEY, EMBEDDING_MODEL, SIMILARITY_THRESHOLD, TOP_K_RESULTS, ) from vibe_bot.db.connection import connect logger = logging.getLogger(__name__) def search_similar_messages( db_path: str, query: str, top_k: int = TOP_K_RESULTS, min_similarity: float = SIMILARITY_THRESHOLD, ) -> list[tuple[str, str, float]]: """Search for messages similar to the query using embeddings. A single JOIN pulls every user row, its stored embedding, the stored L2 norm, and its ``_response`` companion. Similarities are one matrix multiply over the stored norms — no per-vector renormalization. Rows with a missing or zero norm score 0 instead of dividing by zero. """ query_embedding = llm_client.embedding( text=query, model=EMBEDDING_MODEL, url=EMBEDDING_ENDPOINT, api_key=EMBEDDING_ENDPOINT_KEY, ) if not query_embedding: return [] query_vector = np.array(query_embedding, dtype=np.float32) query_norm = float(np.linalg.norm(query_vector)) if query_norm == 0: return [] conn = connect(db_path) try: cursor = conn.cursor() cursor.execute( """ SELECT cm.content, r.content, me.embedding, me.norm FROM chat_messages cm JOIN message_embeddings me ON me.message_id = cm.message_id LEFT JOIN chat_messages r ON r.message_id = cm.message_id || '_response' WHERE cm.role = 'user' """, ) rows = cursor.fetchall() finally: conn.close() if not rows: return [] n_rows = len(rows) blobs = [embedding_blob for _c, _r, embedding_blob, _n in rows] dim = len(blobs[0]) // 4 if sum(len(blob) for blob in blobs) == n_rows * dim * 4: vectors = np.frombuffer(b"".join(blobs), dtype=np.float32).reshape(n_rows, dim) else: # Mixed blob lengths (e.g. EMBEDDING_MODEL changed mid-life) can't be # batched into one reshape; reconstruct per row, zero-padded (or # truncated) to the query dim so the single matrix multiply still works. vectors = np.zeros((n_rows, query_vector.size), dtype=np.float32) for i, blob in enumerate(blobs): row = np.frombuffer(blob, dtype=np.float32) k = min(row.size, query_vector.size) vectors[i, :k] = row[:k] norms = np.array( [ stored_norm if stored_norm is not None else 0.0 for _c, _r, _b, stored_norm in rows ], dtype=np.float64, ) safe_norms = np.where(norms > 0, norms, 1.0) similarities = vectors @ query_vector / (safe_norms * query_norm) similarities = np.where(norms > 0, similarities, 0.0) results: list[tuple[str, str, float]] = [] for (content, response, _blob, _norm), similarity in zip( rows, similarities, strict=True ): if response is None or similarity < min_similarity: continue results.append((str(content), str(response), float(similarity))) results.sort(key=lambda item: item[2], reverse=True) return results[:top_k] def get_bot_history( db_path: str, bot_name: str, limit: int = 20 ) -> list[tuple[str, str]]: """Get message history for a specific custom bot. Args: bot_name: The name of the custom bot. limit: Maximum number of messages to retrieve. Returns: List of (user_message, bot_response) tuples. """ conn = connect(db_path) cursor = conn.cursor() logger.debug( "Fetching last %d messages for bot %r", limit, bot_name, ) cursor.execute( """ SELECT message_id, content FROM chat_messages WHERE bot_name = ? AND message_id NOT LIKE '%%_response' ORDER BY timestamp DESC LIMIT ? """, (bot_name, limit), ) conversations: list[tuple[str, str]] = [] try: for message_id, msg_content in cursor.fetchall(): logger.debug("Finding response for message_id=%s", message_id) cursor.execute( """ SELECT content FROM chat_messages WHERE message_id = ? ORDER BY timestamp DESC """, (f"{message_id}_response",), ) response_row = cursor.fetchone() if response_row: logger.debug("Found response for message_id=%s", message_id) conversations.append((str(msg_content), str(response_row[0]))) else: logger.debug("No response found") finally: conn.close() return conversations def get_user_history( db_path: str, user_id: str, limit: int = 20 ) -> list[tuple[str, str]]: """Get message history for a specific user.""" conn = connect(db_path) cursor = conn.cursor() logger.debug("Fetching last %d user messages", limit) cursor.execute( """ SELECT message_id, content FROM chat_messages WHERE user_id = ? AND role = 'user' ORDER BY timestamp DESC LIMIT ? """, (user_id, limit), ) # Format is [user message, bot response] conversations: list[tuple[str, str]] = [] try: for message_id, msg_content in cursor.fetchall(): logger.debug("Finding response for message_id=%s", message_id) cursor.execute( """ SELECT content FROM chat_messages WHERE message_id = ? ORDER BY timestamp DESC """, (f"{message_id}_response",), ) response_row = cursor.fetchone() if response_row: logger.debug("Found response for message_id=%s", message_id) conversations.append((str(msg_content), str(response_row[0]))) else: logger.debug("No response found") finally: conn.close() return conversations