Files
2026-08-19 13:16:43 -04:00

307 lines
10 KiB
Python

"""Chat message store: persistence, cleanup, and recency queries."""
from __future__ import annotations
import logging
import sqlite3
import numpy as np
from vibe_bot import llm_client
from vibe_bot.config import (
DB_PATH,
EMBEDDING_ENDPOINT,
EMBEDDING_ENDPOINT_KEY,
EMBEDDING_MODEL,
MAX_HISTORY_MESSAGES,
SIMILARITY_THRESHOLD,
TOP_K_RESULTS,
)
from vibe_bot.db.connection import connect
from vibe_bot.db.schema import initialize_chat_tables
from vibe_bot.db.search import (
get_bot_history,
get_user_history,
search_similar_messages,
)
from vibe_bot.db.timing import (
get_image_generation_time_estimate,
record_image_generation_time,
)
from vibe_bot.db.vectors import (
bytes_to_vector,
cosine_similarity,
vector_to_bytes,
)
logger = logging.getLogger(__name__)
class ChatDatabase:
"""SQLite store for chat history, embedding-backed RAG, and image timing."""
def __init__(self, db_path: str = DB_PATH) -> None:
"""Initialize the database connection.
Args:
db_path: Path to the SQLite database file.
"""
logger.info("Initializing ChatDatabase with path: %s", db_path)
self.db_path = db_path
initialize_chat_tables(db_path)
def _vector_to_bytes(self, vector: list[float]) -> bytes:
"""Convert vector to bytes for SQLite storage."""
return vector_to_bytes(vector)
def _bytes_to_vector(self, blob: bytes) -> np.ndarray:
"""Convert bytes back to a vector."""
return bytes_to_vector(blob)
def _calculate_similarity(self, vec1: np.ndarray, vec2: np.ndarray) -> float:
"""Calculate cosine similarity between two vectors."""
return cosine_similarity(vec1, vec2)
def add_message(
self,
*,
message_id: str,
user_id: str,
username: str,
content: str,
bot_name: str | None = None,
channel_id: str | None = None,
guild_id: str | None = None,
role: str = "user",
embed: bool = True,
) -> bool:
"""Add a message to the database, optionally storing its embedding.
Args:
role: Either "user" (a human message) or "assistant" (a bot
response). Used to scope RAG retrieval instead of matching a
hard-coded bot username.
embed: Whether to generate and store an embedding for the message.
Response rows pass False: only user rows feed RAG retrieval.
"""
logger.debug("Adding message %s from user %s", message_id, user_id)
conn = connect(self.db_path)
cursor = conn.cursor()
try:
logger.debug(
"Inserting message into chat_messages table: message_id=%s",
message_id,
)
cursor.execute(
"""
INSERT OR REPLACE INTO chat_messages
(message_id, user_id, username, content, bot_name, channel_id,
guild_id, role)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
""",
(
message_id,
user_id,
username,
content,
bot_name,
channel_id,
guild_id,
role,
),
)
logger.debug("Message %s inserted into chat_messages table", message_id)
if embed:
logger.debug("Generating embedding for message %s", message_id)
embedding = llm_client.embedding(
content,
model=EMBEDDING_MODEL,
url=EMBEDDING_ENDPOINT,
api_key=EMBEDDING_ENDPOINT_KEY,
)
if embedding:
logger.debug(
"Embedding generated successfully for message %s, "
"storing in database",
message_id,
)
vector = np.array(embedding, dtype=np.float32)
cursor.execute(
"""
INSERT OR REPLACE INTO message_embeddings
(message_id, embedding, norm)
VALUES (?, ?, ?)
""",
(message_id, vector.tobytes(), float(np.linalg.norm(vector))),
)
logger.debug(
"Embedding stored in message_embeddings table for message %s",
message_id,
)
else:
logger.warning(
"Failed to generate embedding for message %s, "
"skipping embedding storage",
message_id,
)
logger.debug("Checking if cleanup of old messages is needed")
self._cleanup_old_messages(cursor)
conn.commit()
except Exception:
logger.exception("Error adding message %s", message_id)
conn.rollback()
return False
else:
logger.debug("Successfully added message %s to database", message_id)
return True
finally:
conn.close()
def _cleanup_old_messages(self, cursor: sqlite3.Cursor) -> None:
"""Remove old messages to stay within the limit.
The rows to delete are captured up front. Deriving the embedding
message_ids from a fresh subquery *after* the chat_messages delete
would select the next-oldest live rows instead of the ones just
removed, orphaning the deleted rows' embeddings and deleting the
embeddings of rows that should survive.
"""
cursor.execute(
"""
SELECT COUNT(*) FROM chat_messages
""",
)
count = cursor.fetchone()[0]
if count <= MAX_HISTORY_MESSAGES:
return
excess = count - MAX_HISTORY_MESSAGES
cursor.execute(
"""
SELECT id, message_id FROM chat_messages
ORDER BY timestamp ASC
LIMIT ?
""",
(excess,),
)
oldest = cursor.fetchall()
if not oldest:
return
row_ids = [row[0] for row in oldest]
# Include each row's `_response` companion so a deleted user message
# also sheds its response embedding (and vice versa).
message_ids: list[str] = []
for _id, message_id in oldest:
message_ids.append(message_id)
message_ids.append(f"{message_id}_response")
id_placeholders = ", ".join("?" for _ in row_ids)
cursor.execute(
f"DELETE FROM chat_messages WHERE id IN ({id_placeholders})",
row_ids,
)
mid_placeholders = ", ".join("?" for _ in message_ids)
cursor.execute(
f"DELETE FROM message_embeddings WHERE message_id IN ({mid_placeholders})",
message_ids,
)
def search_similar_messages(
self,
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."""
return search_similar_messages(
self.db_path,
query,
top_k=top_k,
min_similarity=min_similarity,
)
def get_bot_history(self, 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.
"""
return get_bot_history(self.db_path, bot_name, limit)
def get_user_history(self, user_id: str, limit: int = 20) -> list[tuple[str, str]]:
"""Get message history for a specific user."""
return get_user_history(self.db_path, user_id, limit)
def get_conversation_context(
self,
user_id: str,
current_message: str,
max_context: int = 5,
) -> list[dict[str, str]]:
"""Get relevant conversation context for RAG."""
recent_messages = get_user_history(self.db_path, user_id, limit=max_context * 2)
similar_messages = search_similar_messages(
self.db_path,
current_message,
top_k=max_context,
)
context_parts: list[dict[str, str]] = []
for user_message, bot_message in recent_messages:
context_parts.append({"role": "assistant", "content": bot_message})
context_parts.append({"role": "user", "content": user_message})
for user_message, bot_message, _similarity in similar_messages:
context_parts.append({"role": "assistant", "content": bot_message})
context_parts.append({"role": "user", "content": user_message})
# Conversation history needs to be delivered in "newest context last" order
context_parts.reverse()
return context_parts
def clear_all_messages(self) -> None:
"""Clear all messages and embeddings from the database."""
conn = connect(self.db_path)
cursor = conn.cursor()
cursor.execute("DELETE FROM message_embeddings")
cursor.execute("DELETE FROM chat_messages")
conn.commit()
conn.close()
def record_image_generation_time(self, duration_seconds: float) -> bool:
"""Record how long an image generation took.
Args:
duration_seconds: Wall-clock seconds the generation took.
"""
return record_image_generation_time(self.db_path, duration_seconds)
def get_image_generation_time_estimate(self) -> float | None:
"""Get a moving-average estimate of image generation time.
Returns:
The average duration in seconds over the most recent generations,
or None if there is no generation history yet.
"""
return get_image_generation_time_estimate(self.db_path)