Files
kandrusyak_bot/assistant_bot/storage.py
kandrusyak c04fb9480b fix syntax
2026-07-25 15:47:36 +03:00

502 lines
18 KiB
Python

import re
import sqlite3
from contextlib import contextmanager
from datetime import datetime
from pathlib import Path
from typing import Any
from .config import get_default_model
from .time_utils import to_utc_iso, utc_now
def row_to_dict(row: sqlite3.Row) -> dict[str, Any]:
return {str(key): row[key] for key in row.keys()}
def required_lastrowid(cursor: sqlite3.Cursor) -> int:
lastrowid = cursor.lastrowid
if lastrowid is None:
raise RuntimeError("SQLite insert did not return a row ID")
return lastrowid
class AssistantStorage:
def __init__(self, path: Path) -> None:
self.path = path
self.path.parent.mkdir(parents=True, exist_ok=True)
self._init_db()
def _connect(self) -> sqlite3.Connection:
connection = sqlite3.connect(self.path)
connection.row_factory = sqlite3.Row
return connection
@contextmanager
def _connection(self):
connection = self._connect()
try:
yield connection
connection.commit()
except Exception:
connection.rollback()
raise
finally:
connection.close()
def _init_db(self) -> None:
with self._connection() as connection:
connection.execute("PRAGMA journal_mode=WAL")
connection.executescript(
"""
CREATE TABLE IF NOT EXISTS user_settings (
user_id INTEGER PRIMARY KEY,
ollama_model TEXT,
yandex_model TEXT,
updated_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS memories (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
text TEXT NOT NULL,
created_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS notes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
text TEXT NOT NULL,
created_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS reminders (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
chat_id INTEGER NOT NULL,
text TEXT NOT NULL,
remind_at TEXT NOT NULL,
status TEXT NOT NULL DEFAULT 'pending',
created_at TEXT NOT NULL,
sent_at TEXT
);
CREATE TABLE IF NOT EXISTS tracked_items (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
title TEXT NOT NULL,
status TEXT NOT NULL,
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS conversation_messages (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
chat_id INTEGER NOT NULL,
role TEXT NOT NULL CHECK(role IN ('user', 'assistant')),
content TEXT NOT NULL,
created_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS conversation_contexts (
user_id INTEGER NOT NULL,
chat_id INTEGER NOT NULL,
started_after_id INTEGER NOT NULL DEFAULT 0,
updated_at TEXT NOT NULL,
PRIMARY KEY (user_id, chat_id)
);
CREATE INDEX IF NOT EXISTS idx_memories_user_created
ON memories(user_id, created_at DESC);
CREATE INDEX IF NOT EXISTS idx_notes_user_created
ON notes(user_id, created_at DESC);
CREATE INDEX IF NOT EXISTS idx_reminders_due
ON reminders(status, remind_at);
CREATE INDEX IF NOT EXISTS idx_tracked_user_updated
ON tracked_items(user_id, updated_at DESC);
CREATE INDEX IF NOT EXISTS idx_conversation_chat_id
ON conversation_messages(user_id, chat_id, id DESC);
"""
)
columns = {
str(row["name"])
for row in connection.execute("PRAGMA table_info(user_settings)")
}
if "yandex_model" not in columns:
connection.execute(
"ALTER TABLE user_settings ADD COLUMN yandex_model TEXT"
)
def add_conversation_exchange(
self,
user_id: int,
chat_id: int,
user_content: str,
assistant_content: str,
) -> None:
now = to_utc_iso(utc_now())
with self._connection() as connection:
connection.executemany(
"""
INSERT INTO conversation_messages(user_id, chat_id, role, content, created_at)
VALUES (?, ?, ?, ?, ?)
""",
(
(user_id, chat_id, "user", user_content, now),
(user_id, chat_id, "assistant", assistant_content, now),
),
)
def list_conversation_messages(
self,
user_id: int,
chat_id: int,
limit: int = 12,
) -> list[dict[str, Any]]:
with self._connection() as connection:
context_row = connection.execute(
"""
SELECT started_after_id
FROM conversation_contexts
WHERE user_id = ? AND chat_id = ?
""",
(user_id, chat_id),
).fetchone()
started_after_id = int(context_row["started_after_id"]) if context_row else 0
rows = connection.execute(
"""
SELECT id, role, content, created_at
FROM conversation_messages
WHERE user_id = ? AND chat_id = ? AND id > ?
ORDER BY id DESC
LIMIT ?
""",
(user_id, chat_id, started_after_id, limit),
).fetchall()
return [row_to_dict(row) for row in reversed(rows)]
def start_new_conversation(self, user_id: int, chat_id: int) -> None:
now = to_utc_iso(utc_now())
with self._connection() as connection:
row = connection.execute(
"""
SELECT COALESCE(MAX(id), 0) AS last_id
FROM conversation_messages
WHERE user_id = ? AND chat_id = ?
""",
(user_id, chat_id),
).fetchone()
last_id = int(row["last_id"])
connection.execute(
"""
INSERT INTO conversation_contexts(user_id, chat_id, started_after_id, updated_at)
VALUES (?, ?, ?, ?)
ON CONFLICT(user_id, chat_id) DO UPDATE SET
started_after_id = excluded.started_after_id,
updated_at = excluded.updated_at
""",
(user_id, chat_id, last_id, now),
)
def search_conversation_messages(
self,
user_id: int,
chat_id: int,
query: str,
limit: int = 10,
) -> list[dict[str, Any]]:
normalized_query = query.casefold().strip()
terms = list(
dict.fromkeys(
part for part in re.findall(r"\w+", normalized_query) if len(part) > 1
)
)
if not terms:
return []
with self._connection() as connection:
rows = connection.execute(
"""
SELECT id, role, content, created_at
FROM conversation_messages
WHERE user_id = ? AND chat_id = ?
ORDER BY id DESC
""",
(user_id, chat_id),
).fetchall()
if not rows:
return []
newest_id = int(rows[0]["id"])
matches: list[tuple[float, dict[str, Any]]] = []
for row in rows:
content = str(row["content"]).casefold()
matched_terms = sum(term in content for term in terms)
if not matched_terms:
continue
coverage = matched_terms / len(terms)
phrase_bonus = 1.0 if len(normalized_query) > 2 and normalized_query in content else 0.0
age = max(0, newest_id - int(row["id"]))
recency_weight = 1.0 / (1.0 + age / 50.0)
score = (coverage + phrase_bonus) * (0.5 + 0.5 * recency_weight)
item = row_to_dict(row)
item["score"] = round(score, 4)
matches.append((score, item))
matches.sort(key=lambda item: (item[0], int(item[1]["id"])), reverse=True)
return [item for _score, item in matches[: max(1, limit)]]
def conversation_message_window(
self,
user_id: int,
chat_id: int,
message_id: int,
radius: int = 2,
) -> list[dict[str, Any]]:
with self._connection() as connection:
before = connection.execute(
"""
SELECT id, role, content, created_at
FROM conversation_messages
WHERE user_id = ? AND chat_id = ? AND id <= ?
ORDER BY id DESC
LIMIT ?
""",
(user_id, chat_id, message_id, radius + 1),
).fetchall()
after = connection.execute(
"""
SELECT id, role, content, created_at
FROM conversation_messages
WHERE user_id = ? AND chat_id = ? AND id > ?
ORDER BY id ASC
LIMIT ?
""",
(user_id, chat_id, message_id, radius),
).fetchall()
rows = [*reversed(before), *after]
return [row_to_dict(row) for row in rows]
def get_user_model(self, user_id: int, provider: str = "local") -> str:
column = self._model_column(provider)
with self._connection() as connection:
row = connection.execute(
f"SELECT {column} AS model FROM user_settings WHERE user_id = ?",
(user_id,),
).fetchone()
if row and row["model"]:
return str(row["model"])
return get_default_model(provider)
def set_user_model(
self,
user_id: int,
model: str,
provider: str = "local",
) -> None:
column = self._model_column(provider)
now = to_utc_iso(utc_now())
with self._connection() as connection:
connection.execute(
f"""
INSERT INTO user_settings(user_id, {column}, updated_at)
VALUES (?, ?, ?)
ON CONFLICT(user_id) DO UPDATE SET
{column} = excluded.{column},
updated_at = excluded.updated_at
""",
(user_id, model, now),
)
@staticmethod
def _model_column(provider: str) -> str:
try:
return {
"local": "ollama_model",
"yandex": "yandex_model",
}[provider]
except KeyError as exc:
raise ValueError(f"Unknown AI provider: {provider}") from exc
def add_memory(self, user_id: int, text: str) -> int:
with self._connection() as connection:
cursor = connection.execute(
"INSERT INTO memories(user_id, text, created_at) VALUES (?, ?, ?)",
(user_id, text, to_utc_iso(utc_now())),
)
return required_lastrowid(cursor)
def list_memories(self, user_id: int, limit: int = 20) -> list[dict[str, Any]]:
with self._connection() as connection:
rows = connection.execute(
"""
SELECT id, text, created_at
FROM memories
WHERE user_id = ?
ORDER BY created_at DESC, id DESC
LIMIT ?
""",
(user_id, limit),
).fetchall()
return [row_to_dict(row) for row in rows]
def delete_memory(self, user_id: int, memory_id: int) -> bool:
with self._connection() as connection:
cursor = connection.execute(
"DELETE FROM memories WHERE user_id = ? AND id = ?",
(user_id, memory_id),
)
return cursor.rowcount > 0
def add_note(self, user_id: int, text: str) -> int:
with self._connection() as connection:
cursor = connection.execute(
"INSERT INTO notes(user_id, text, created_at) VALUES (?, ?, ?)",
(user_id, text, to_utc_iso(utc_now())),
)
return required_lastrowid(cursor)
def list_notes(self, user_id: int, limit: int = 20) -> list[dict[str, Any]]:
with self._connection() as connection:
rows = connection.execute(
"""
SELECT id, text, created_at
FROM notes
WHERE user_id = ?
ORDER BY created_at DESC, id DESC
LIMIT ?
""",
(user_id, limit),
).fetchall()
return [row_to_dict(row) for row in rows]
def delete_note(self, user_id: int, note_id: int) -> bool:
with self._connection() as connection:
cursor = connection.execute(
"DELETE FROM notes WHERE user_id = ? AND id = ?",
(user_id, note_id),
)
return cursor.rowcount > 0
def add_reminder(
self,
user_id: int,
chat_id: int,
text: str,
remind_at_utc: datetime,
) -> int:
with self._connection() as connection:
cursor = connection.execute(
"""
INSERT INTO reminders(user_id, chat_id, text, remind_at, created_at)
VALUES (?, ?, ?, ?, ?)
""",
(
user_id,
chat_id,
text,
to_utc_iso(remind_at_utc),
to_utc_iso(utc_now()),
),
)
return required_lastrowid(cursor)
def list_reminders(self, user_id: int, limit: int = 20) -> list[dict[str, Any]]:
with self._connection() as connection:
rows = connection.execute(
"""
SELECT id, text, remind_at, status, sent_at
FROM reminders
WHERE user_id = ? AND status = 'pending'
ORDER BY remind_at ASC, id ASC
LIMIT ?
""",
(user_id, limit),
).fetchall()
return [row_to_dict(row) for row in rows]
def cancel_reminder(self, user_id: int, reminder_id: int) -> bool:
with self._connection() as connection:
cursor = connection.execute(
"""
UPDATE reminders
SET status = 'cancelled'
WHERE user_id = ? AND id = ? AND status = 'pending'
""",
(user_id, reminder_id),
)
return cursor.rowcount > 0
def due_reminders(self, now_utc: datetime, limit: int = 20) -> list[dict[str, Any]]:
with self._connection() as connection:
rows = connection.execute(
"""
SELECT id, user_id, chat_id, text, remind_at
FROM reminders
WHERE status = 'pending' AND remind_at <= ?
ORDER BY remind_at ASC, id ASC
LIMIT ?
""",
(to_utc_iso(now_utc), limit),
).fetchall()
return [row_to_dict(row) for row in rows]
def mark_reminder_sent(self, reminder_id: int) -> None:
with self._connection() as connection:
connection.execute(
"""
UPDATE reminders
SET status = 'sent', sent_at = ?
WHERE id = ? AND status = 'pending'
""",
(to_utc_iso(utc_now()), reminder_id),
)
def add_tracked_item(self, user_id: int, title: str, status: str) -> int:
now = to_utc_iso(utc_now())
with self._connection() as connection:
cursor = connection.execute(
"""
INSERT INTO tracked_items(user_id, title, status, created_at, updated_at)
VALUES (?, ?, ?, ?, ?)
""",
(user_id, title, status, now, now),
)
return required_lastrowid(cursor)
def list_tracked_items(self, user_id: int, limit: int = 30) -> list[dict[str, Any]]:
with self._connection() as connection:
rows = connection.execute(
"""
SELECT id, title, status, created_at, updated_at
FROM tracked_items
WHERE user_id = ?
ORDER BY updated_at DESC, id DESC
LIMIT ?
""",
(user_id, limit),
).fetchall()
return [row_to_dict(row) for row in rows]
def set_tracked_status(self, user_id: int, item_id: int, status: str) -> bool:
with self._connection() as connection:
cursor = connection.execute(
"""
UPDATE tracked_items
SET status = ?, updated_at = ?
WHERE user_id = ? AND id = ?
""",
(status, to_utc_iso(utc_now()), user_id, item_id),
)
return cursor.rowcount > 0
def delete_tracked_item(self, user_id: int, item_id: int) -> bool:
with self._connection() as connection:
cursor = connection.execute(
"DELETE FROM tracked_items WHERE user_id = ? AND id = ?",
(user_id, item_id),
)
return cursor.rowcount > 0