352 lines
12 KiB
Python
352 lines
12 KiB
Python
import sqlite3
|
|
import tempfile
|
|
import unittest
|
|
from contextlib import closing
|
|
from datetime import datetime, timezone
|
|
from pathlib import Path
|
|
from zoneinfo import ZoneInfo
|
|
|
|
from assistant_bot.agent import (
|
|
build_agent_tool_prompt,
|
|
execute_agent_tool,
|
|
parse_agent_decision,
|
|
)
|
|
from assistant_bot.agent_tools import AGENT_TOOLS
|
|
from assistant_bot.migrations import LATEST_SCHEMA_VERSION
|
|
from assistant_bot.reminders import parse_reminder, unit_to_timedelta
|
|
from assistant_bot.storage import AssistantStorage
|
|
|
|
|
|
DEFAULT_MODELS = {
|
|
"local": "qwen3.5:9b",
|
|
"yandex": "gpt://folder/yandexgpt/latest",
|
|
}
|
|
|
|
|
|
def create_storage(path: Path) -> AssistantStorage:
|
|
return AssistantStorage(path, default_models=DEFAULT_MODELS)
|
|
|
|
|
|
class ReminderParserTests(unittest.TestCase):
|
|
def test_supported_relative_unit(self) -> None:
|
|
delta = unit_to_timedelta(2, "часа")
|
|
self.assertIsNotNone(delta)
|
|
assert delta is not None
|
|
self.assertEqual(delta.total_seconds(), 7200)
|
|
|
|
def test_absolute_datetime(self) -> None:
|
|
result = parse_reminder(
|
|
"2030-01-02 10:30 проверить отчет",
|
|
ZoneInfo("UTC"),
|
|
)
|
|
|
|
self.assertIsNotNone(result)
|
|
assert result is not None
|
|
self.assertEqual(
|
|
result.remind_at_utc,
|
|
datetime(2030, 1, 2, 10, 30, tzinfo=timezone.utc),
|
|
)
|
|
self.assertEqual(result.text, "проверить отчет")
|
|
|
|
|
|
class AgentDecisionTests(unittest.TestCase):
|
|
def test_tool_call_json(self) -> None:
|
|
decision = parse_agent_decision(
|
|
'{"tool_calls":[{"name":"create_note","arguments":{"text":"идея"}}]}'
|
|
)
|
|
|
|
self.assertIsNone(decision.final)
|
|
self.assertEqual(decision.tool_calls[0]["name"], "create_note")
|
|
self.assertEqual(decision.tool_calls[0]["arguments"]["text"], "идея")
|
|
|
|
def test_context_reset_flag(self) -> None:
|
|
decision = parse_agent_decision(
|
|
'{"final":"Перейдем к новой теме","reset_context":true}'
|
|
)
|
|
|
|
self.assertEqual(decision.final, "Перейдем к новой теме")
|
|
self.assertTrue(decision.reset_context)
|
|
|
|
|
|
class AgentToolContractTests(unittest.TestCase):
|
|
TOOL_NAMES = (
|
|
"get_current_datetime",
|
|
"remember",
|
|
"list_memory",
|
|
"delete_memory",
|
|
"create_note",
|
|
"list_notes",
|
|
"delete_note",
|
|
"create_reminder",
|
|
"list_reminders",
|
|
"cancel_reminder",
|
|
"create_status",
|
|
"list_statuses",
|
|
"update_status",
|
|
"delete_status",
|
|
"search_conversation",
|
|
)
|
|
|
|
def test_prompt_exposes_all_supported_tool_names(self) -> None:
|
|
prompt = build_agent_tool_prompt(ZoneInfo("UTC"))
|
|
|
|
self.assertEqual(
|
|
tuple(tool.name for tool in AGENT_TOOLS),
|
|
self.TOOL_NAMES,
|
|
)
|
|
for tool_name in self.TOOL_NAMES:
|
|
with self.subTest(tool_name=tool_name):
|
|
self.assertIn(tool_name, prompt)
|
|
|
|
def test_crud_tool_result_shapes_remain_stable(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
storage = create_storage(Path(directory) / "assistant.sqlite3")
|
|
common = {
|
|
"storage": storage,
|
|
"user_id": 42,
|
|
"chat_id": 100,
|
|
"tz": ZoneInfo("UTC"),
|
|
}
|
|
|
|
memory = execute_agent_tool(
|
|
"remember",
|
|
{"text": "короткие ответы"},
|
|
**common,
|
|
)
|
|
self.assertEqual(
|
|
memory,
|
|
{"ok": True, "id": memory["id"], "text": "короткие ответы"},
|
|
)
|
|
self.assertEqual(
|
|
execute_agent_tool("list_memory", {"limit": 10}, **common)[
|
|
"items"
|
|
][0]["id"],
|
|
memory["id"],
|
|
)
|
|
self.assertEqual(
|
|
execute_agent_tool(
|
|
"delete_memory",
|
|
{"id": memory["id"]},
|
|
**common,
|
|
),
|
|
{"ok": True, "id": memory["id"]},
|
|
)
|
|
|
|
note = execute_agent_tool(
|
|
"create_note",
|
|
{"text": "идея"},
|
|
**common,
|
|
)
|
|
self.assertEqual(
|
|
note,
|
|
{"ok": True, "id": note["id"], "text": "идея"},
|
|
)
|
|
self.assertEqual(
|
|
execute_agent_tool("list_notes", {}, **common)["items"][0][
|
|
"id"
|
|
],
|
|
note["id"],
|
|
)
|
|
self.assertEqual(
|
|
execute_agent_tool(
|
|
"delete_note",
|
|
{"id": note["id"]},
|
|
**common,
|
|
),
|
|
{"ok": True, "id": note["id"]},
|
|
)
|
|
|
|
status = execute_agent_tool(
|
|
"create_status",
|
|
{"title": "паспорт", "status": "ожидание"},
|
|
**common,
|
|
)
|
|
self.assertEqual(
|
|
status,
|
|
{
|
|
"ok": True,
|
|
"id": status["id"],
|
|
"title": "паспорт",
|
|
"status": "ожидание",
|
|
},
|
|
)
|
|
self.assertEqual(
|
|
execute_agent_tool(
|
|
"update_status",
|
|
{"id": status["id"], "status": "готово"},
|
|
**common,
|
|
),
|
|
{"ok": True, "id": status["id"], "status": "готово"},
|
|
)
|
|
self.assertEqual(
|
|
execute_agent_tool("list_statuses", {}, **common)["items"][0][
|
|
"status"
|
|
],
|
|
"готово",
|
|
)
|
|
self.assertEqual(
|
|
execute_agent_tool(
|
|
"delete_status",
|
|
{"id": status["id"]},
|
|
**common,
|
|
),
|
|
{"ok": True, "id": status["id"]},
|
|
)
|
|
|
|
def test_unknown_tool_error_remains_stable(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
result = execute_agent_tool(
|
|
"missing",
|
|
{},
|
|
create_storage(Path(directory) / "assistant.sqlite3"),
|
|
user_id=42,
|
|
chat_id=100,
|
|
tz=ZoneInfo("UTC"),
|
|
)
|
|
|
|
self.assertEqual(
|
|
result,
|
|
{"ok": False, "error": "unknown tool: missing"},
|
|
)
|
|
|
|
|
|
class StorageTests(unittest.TestCase):
|
|
def test_new_database_uses_latest_schema_version(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
database_path = Path(directory) / "assistant.sqlite3"
|
|
create_storage(database_path)
|
|
|
|
with closing(sqlite3.connect(database_path)) as connection:
|
|
version = connection.execute(
|
|
"PRAGMA user_version"
|
|
).fetchone()[0]
|
|
|
|
self.assertEqual(version, LATEST_SCHEMA_VERSION)
|
|
|
|
def test_user_authorization_is_persisted(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
database_path = Path(directory) / "assistant.sqlite3"
|
|
storage = create_storage(database_path)
|
|
|
|
self.assertFalse(storage.is_user_authorized(42))
|
|
storage.authorize_user(42)
|
|
|
|
reopened_storage = create_storage(database_path)
|
|
self.assertTrue(reopened_storage.is_user_authorized(42))
|
|
self.assertFalse(reopened_storage.is_user_authorized(43))
|
|
|
|
def test_legacy_user_settings_gets_yandex_model_column(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
database_path = Path(directory) / "assistant.sqlite3"
|
|
with closing(sqlite3.connect(database_path)) as connection:
|
|
connection.execute(
|
|
"""
|
|
CREATE TABLE user_settings (
|
|
user_id INTEGER PRIMARY KEY,
|
|
ollama_model TEXT,
|
|
updated_at TEXT NOT NULL
|
|
)
|
|
"""
|
|
)
|
|
connection.commit()
|
|
|
|
storage = create_storage(database_path)
|
|
storage.set_user_model(42, "yandexgpt", "yandex")
|
|
|
|
self.assertEqual(storage.get_user_model(42, "yandex"), "yandexgpt")
|
|
with closing(sqlite3.connect(database_path)) as connection:
|
|
version = connection.execute(
|
|
"PRAGMA user_version"
|
|
).fetchone()[0]
|
|
self.assertEqual(version, LATEST_SCHEMA_VERSION)
|
|
|
|
def test_default_models_are_injected(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
storage = create_storage(Path(directory) / "assistant.sqlite3")
|
|
|
|
self.assertEqual(
|
|
storage.get_user_model(42, "local"),
|
|
DEFAULT_MODELS["local"],
|
|
)
|
|
self.assertEqual(
|
|
storage.get_user_model(42, "yandex"),
|
|
DEFAULT_MODELS["yandex"],
|
|
)
|
|
|
|
def test_provider_models_are_stored_separately(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
storage = create_storage(Path(directory) / "assistant.sqlite3")
|
|
|
|
storage.set_user_model(42, "qwen3.5:9b", "local")
|
|
storage.set_user_model(42, "yandexgpt", "yandex")
|
|
|
|
self.assertEqual(storage.get_user_model(42, "local"), "qwen3.5:9b")
|
|
self.assertEqual(storage.get_user_model(42, "yandex"), "yandexgpt")
|
|
|
|
def test_memory_crud(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
storage = create_storage(Path(directory) / "assistant.sqlite3")
|
|
memory_id = storage.add_memory(42, "короткие ответы")
|
|
|
|
self.assertEqual(storage.list_memories(42)[0]["id"], memory_id)
|
|
self.assertTrue(storage.delete_memory(42, memory_id))
|
|
self.assertEqual(storage.list_memories(42), [])
|
|
|
|
def test_new_context_preserves_searchable_archive(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
storage = create_storage(Path(directory) / "assistant.sqlite3")
|
|
storage.add_conversation_exchange(
|
|
42,
|
|
100,
|
|
"Обсудим проект Альфа",
|
|
"Какой аспект проекта интересует?",
|
|
)
|
|
storage.start_new_conversation(42, 100)
|
|
storage.add_conversation_exchange(
|
|
42,
|
|
100,
|
|
"Снова обсуждаем проект Альфа",
|
|
"Продолжаем обсуждение проекта.",
|
|
)
|
|
storage.add_conversation_exchange(42, 200, "Другой чат", "Другой ответ")
|
|
|
|
messages = storage.list_conversation_messages(42, 100)
|
|
matches = storage.search_conversation_messages(42, 100, "проект Альфа")
|
|
|
|
self.assertEqual(
|
|
[(row["role"], row["content"]) for row in messages],
|
|
[
|
|
("user", "Снова обсуждаем проект Альфа"),
|
|
("assistant", "Продолжаем обсуждение проекта."),
|
|
],
|
|
)
|
|
self.assertEqual(matches[0]["content"], "Снова обсуждаем проект Альфа")
|
|
self.assertIn("Обсудим проект Альфа", [row["content"] for row in matches])
|
|
self.assertEqual(len(storage.list_conversation_messages(42, 200)), 2)
|
|
|
|
archive_window = storage.conversation_message_window(
|
|
42,
|
|
100,
|
|
int(matches[-1]["id"]),
|
|
)
|
|
self.assertIn(
|
|
"Какой аспект проекта интересует?",
|
|
[row["content"] for row in archive_window],
|
|
)
|
|
|
|
tool_result = execute_agent_tool(
|
|
"search_conversation",
|
|
{"query": "проект Альфа", "limit": 2},
|
|
storage,
|
|
user_id=42,
|
|
chat_id=100,
|
|
tz=ZoneInfo("UTC"),
|
|
)
|
|
self.assertTrue(tool_result["ok"])
|
|
self.assertTrue(tool_result["discussions"])
|
|
self.assertTrue(tool_result["discussions"][0]["messages"])
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|