init
This commit is contained in:
490
assistant_bot/storage.py
Normal file
490
assistant_bot/storage.py
Normal file
@@ -0,0 +1,490 @@
|
||||
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
|
||||
|
||||
|
||||
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 [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 = 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 [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 int(cursor.lastrowid)
|
||||
|
||||
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 [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 int(cursor.lastrowid)
|
||||
|
||||
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 [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 int(cursor.lastrowid)
|
||||
|
||||
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 [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 [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 int(cursor.lastrowid)
|
||||
|
||||
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 [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
|
||||
Reference in New Issue
Block a user