307 lines
10 KiB
Python
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)
|