complete restructure
This commit is contained in:
@@ -0,0 +1,167 @@
|
||||
"""Schema creation and column migrations for the chat and custom-bot tables."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import sqlite3
|
||||
|
||||
import numpy as np
|
||||
|
||||
from vibe_bot.db.connection import connect
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def initialize_chat_tables(db_path: str) -> None:
|
||||
"""Create (and migrate) the chat history and embedding tables."""
|
||||
logger.info("Initializing SQLite database at %s", db_path)
|
||||
conn = connect(db_path)
|
||||
cursor = conn.cursor()
|
||||
|
||||
logger.info("Creating chat_messages table if not exists")
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS chat_messages (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
message_id TEXT UNIQUE,
|
||||
user_id TEXT,
|
||||
username TEXT,
|
||||
content TEXT,
|
||||
timestamp TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
channel_id TEXT,
|
||||
guild_id TEXT
|
||||
)
|
||||
""",
|
||||
)
|
||||
logger.info("chat_messages table initialized successfully")
|
||||
_migrate_chat_messages(cursor)
|
||||
|
||||
logger.info("Creating message_embeddings table if not exists")
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS message_embeddings (
|
||||
message_id TEXT PRIMARY KEY,
|
||||
embedding BLOB,
|
||||
norm REAL,
|
||||
FOREIGN KEY (message_id) REFERENCES chat_messages(message_id)
|
||||
)
|
||||
""",
|
||||
)
|
||||
logger.info("message_embeddings table initialized successfully")
|
||||
_migrate_message_embeddings(cursor)
|
||||
|
||||
logger.info("Creating idx_timestamp index if not exists")
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE INDEX IF NOT EXISTS idx_timestamp ON chat_messages(timestamp)
|
||||
""",
|
||||
)
|
||||
logger.info("idx_timestamp index created successfully")
|
||||
|
||||
logger.info("Creating idx_user_id index if not exists")
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE INDEX IF NOT EXISTS idx_user_id ON chat_messages(user_id)
|
||||
""",
|
||||
)
|
||||
logger.info("idx_user_id index created successfully")
|
||||
|
||||
logger.info("Creating image_generation_times table if not exists")
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS image_generation_times (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
duration_seconds REAL NOT NULL,
|
||||
generated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||
)
|
||||
""",
|
||||
)
|
||||
logger.info("image_generation_times table initialized successfully")
|
||||
|
||||
conn.commit()
|
||||
logger.info("Database initialization completed successfully")
|
||||
conn.close()
|
||||
|
||||
|
||||
def _migrate_chat_messages(cursor: sqlite3.Cursor) -> None:
|
||||
"""Add the bot_name and role columns to pre-existing databases."""
|
||||
logger.info("Checking for chat_messages column migrations")
|
||||
cursor.execute("PRAGMA table_info(chat_messages)")
|
||||
columns = {row[1] for row in cursor.fetchall()}
|
||||
|
||||
if "bot_name" not in columns:
|
||||
logger.info("Adding bot_name column to chat_messages table")
|
||||
cursor.execute("ALTER TABLE chat_messages ADD COLUMN bot_name TEXT")
|
||||
logger.info("bot_name column added successfully")
|
||||
|
||||
# role replaces the old convention of identifying bot responses by a
|
||||
# hard-coded bot username.
|
||||
if "role" not in columns:
|
||||
logger.info("Adding role column to chat_messages table")
|
||||
cursor.execute("ALTER TABLE chat_messages ADD COLUMN role TEXT")
|
||||
cursor.execute(
|
||||
"UPDATE chat_messages SET role = 'assistant' "
|
||||
"WHERE message_id LIKE '%_response' AND role IS NULL",
|
||||
)
|
||||
cursor.execute(
|
||||
"UPDATE chat_messages SET role = 'user' WHERE role IS NULL",
|
||||
)
|
||||
logger.info("role column added and backfilled")
|
||||
|
||||
|
||||
# Backfill in batches so a large legacy table does not build one huge
|
||||
# executemany parameter list in memory.
|
||||
NORM_BACKFILL_BATCH = 500
|
||||
|
||||
|
||||
def _migrate_message_embeddings(cursor: sqlite3.Cursor) -> None:
|
||||
"""Add the norm column to pre-existing databases and backfill it.
|
||||
|
||||
The norm is the L2 norm of the stored float32 blob, so search can score
|
||||
candidates with one matrix multiply and no per-vector renormalization.
|
||||
"""
|
||||
logger.info("Checking for message_embeddings column migrations")
|
||||
cursor.execute("PRAGMA table_info(message_embeddings)")
|
||||
columns = {row[1] for row in cursor.fetchall()}
|
||||
|
||||
if "norm" in columns:
|
||||
return
|
||||
|
||||
logger.info("Adding norm column to message_embeddings table")
|
||||
cursor.execute("ALTER TABLE message_embeddings ADD COLUMN norm REAL")
|
||||
|
||||
cursor.execute(
|
||||
"SELECT message_id, embedding FROM message_embeddings "
|
||||
"WHERE embedding IS NOT NULL",
|
||||
)
|
||||
rows = cursor.fetchall()
|
||||
for start in range(0, len(rows), NORM_BACKFILL_BATCH):
|
||||
cursor.executemany(
|
||||
"UPDATE message_embeddings SET norm = ? WHERE message_id = ?",
|
||||
[
|
||||
(
|
||||
float(np.linalg.norm(np.frombuffer(blob, dtype=np.float32))),
|
||||
message_id,
|
||||
)
|
||||
for message_id, blob in rows[start : start + NORM_BACKFILL_BATCH]
|
||||
],
|
||||
)
|
||||
logger.info("norm column added and backfilled for %d rows", len(rows))
|
||||
|
||||
|
||||
def initialize_custom_bots_table(db_path: str) -> None:
|
||||
"""Create the custom bots table in SQLite."""
|
||||
conn = connect(db_path)
|
||||
conn.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS custom_bots (
|
||||
bot_name TEXT PRIMARY KEY,
|
||||
system_prompt TEXT NOT NULL,
|
||||
created_by TEXT NOT NULL,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
is_active INTEGER DEFAULT 1
|
||||
)
|
||||
""",
|
||||
)
|
||||
conn.commit()
|
||||
conn.close()
|
||||
Reference in New Issue
Block a user