init
This commit is contained in:
1
tests/__init__.py
Normal file
1
tests/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Project tests."""
|
||||
71
tests/test_config.py
Normal file
71
tests/test_config.py
Normal file
@@ -0,0 +1,71 @@
|
||||
import os
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
from assistant_bot import config
|
||||
|
||||
|
||||
class WhisperConfigTests(unittest.TestCase):
|
||||
def test_gpu_int8_defaults(self) -> None:
|
||||
variable_names = (
|
||||
config.WHISPER_MODEL_ENV_NAME,
|
||||
config.WHISPER_DEVICE_ENV_NAME,
|
||||
config.WHISPER_COMPUTE_TYPE_ENV_NAME,
|
||||
)
|
||||
with patch("assistant_bot.config.load_env_file"), patch.dict(
|
||||
os.environ,
|
||||
{},
|
||||
clear=False,
|
||||
):
|
||||
for variable_name in variable_names:
|
||||
os.environ.pop(variable_name, None)
|
||||
|
||||
self.assertEqual(config.get_whisper_model(), "large-v3")
|
||||
self.assertEqual(config.get_whisper_device(), "cuda")
|
||||
self.assertEqual(config.get_whisper_compute_type(), "int8")
|
||||
|
||||
|
||||
class AssistantModeConfigTests(unittest.TestCase):
|
||||
def test_local_mode_is_default(self) -> None:
|
||||
with patch("assistant_bot.config.load_env_file"), patch.dict(
|
||||
os.environ,
|
||||
{},
|
||||
clear=False,
|
||||
):
|
||||
os.environ.pop(config.MODE_ENV_NAME, None)
|
||||
self.assertEqual(config.get_assistant_mode(), "local")
|
||||
|
||||
def test_yandex_mode_is_supported(self) -> None:
|
||||
with patch("assistant_bot.config.load_env_file"), patch.dict(
|
||||
os.environ,
|
||||
{config.MODE_ENV_NAME: "YANDEX"},
|
||||
clear=False,
|
||||
):
|
||||
self.assertEqual(config.get_assistant_mode(), "yandex")
|
||||
|
||||
def test_unknown_mode_is_rejected(self) -> None:
|
||||
with patch("assistant_bot.config.load_env_file"), patch.dict(
|
||||
os.environ,
|
||||
{config.MODE_ENV_NAME: "cloud"},
|
||||
clear=False,
|
||||
):
|
||||
with self.assertRaises(RuntimeError):
|
||||
config.get_assistant_mode()
|
||||
|
||||
def test_yandex_model_is_full_gpt_uri(self) -> None:
|
||||
with patch("assistant_bot.config.load_env_file"), patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
config.YANDEX_CLOUD_FOLDER_ENV_NAME: "folder-id",
|
||||
config.YANDEX_CLOUD_MODEL_ENV_NAME: "yandexgpt/latest",
|
||||
},
|
||||
clear=False,
|
||||
):
|
||||
self.assertEqual(
|
||||
config.get_default_yandex_model(),
|
||||
"gpt://folder-id/yandexgpt/latest",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
148
tests/test_core.py
Normal file
148
tests/test_core.py
Normal file
@@ -0,0 +1,148 @@
|
||||
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 execute_agent_tool, parse_agent_decision
|
||||
from assistant_bot.reminders import parse_reminder, unit_to_timedelta
|
||||
from assistant_bot.storage import AssistantStorage
|
||||
|
||||
|
||||
class ReminderParserTests(unittest.TestCase):
|
||||
def test_supported_relative_unit(self) -> None:
|
||||
self.assertEqual(unit_to_timedelta(2, "часа").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 StorageTests(unittest.TestCase):
|
||||
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 = AssistantStorage(database_path)
|
||||
storage.set_user_model(42, "yandexgpt", "yandex")
|
||||
|
||||
self.assertEqual(storage.get_user_model(42, "yandex"), "yandexgpt")
|
||||
|
||||
def test_provider_models_are_stored_separately(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
storage = AssistantStorage(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 = AssistantStorage(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 = AssistantStorage(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()
|
||||
68
tests/test_speech.py
Normal file
68
tests/test_speech.py
Normal file
@@ -0,0 +1,68 @@
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from assistant_bot.speech import SpeechRecognitionError, SpeechRecognizer
|
||||
|
||||
|
||||
class FakeSegment:
|
||||
def __init__(self, text: str) -> None:
|
||||
self.text = text
|
||||
|
||||
|
||||
class SpeechRecognizerTests(unittest.TestCase):
|
||||
def test_transcribe_joins_non_empty_segments(self) -> None:
|
||||
factory_calls = []
|
||||
transcribe_calls = []
|
||||
|
||||
class FakeModel:
|
||||
def transcribe(self, audio_path, **kwargs):
|
||||
transcribe_calls.append((audio_path, kwargs))
|
||||
return iter(
|
||||
[FakeSegment(" Привет "), FakeSegment(""), FakeSegment("мир")]
|
||||
), None
|
||||
|
||||
def model_factory(*args, **kwargs):
|
||||
factory_calls.append((args, kwargs))
|
||||
return FakeModel()
|
||||
|
||||
recognizer = SpeechRecognizer(
|
||||
model_name="small",
|
||||
device="cpu",
|
||||
compute_type="int8",
|
||||
language="ru",
|
||||
model_factory=model_factory,
|
||||
)
|
||||
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
audio_path = Path(directory) / "voice.ogg"
|
||||
result = recognizer.transcribe(audio_path)
|
||||
|
||||
self.assertEqual(result, "Привет мир")
|
||||
self.assertEqual(
|
||||
factory_calls[0],
|
||||
(("small",), {"device": "cpu", "compute_type": "int8"}),
|
||||
)
|
||||
self.assertEqual(
|
||||
transcribe_calls[0][1],
|
||||
{"language": "ru", "beam_size": 5, "vad_filter": True},
|
||||
)
|
||||
|
||||
def test_transcribe_wraps_model_errors(self) -> None:
|
||||
def failing_factory(*_args, **_kwargs):
|
||||
raise RuntimeError("model is unavailable")
|
||||
|
||||
recognizer = SpeechRecognizer(
|
||||
model_name="small",
|
||||
device="cpu",
|
||||
compute_type="int8",
|
||||
language=None,
|
||||
model_factory=failing_factory,
|
||||
)
|
||||
|
||||
with self.assertRaises(SpeechRecognitionError):
|
||||
recognizer.transcribe("voice.ogg")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
75
tests/test_voice_handler.py
Normal file
75
tests/test_voice_handler.py
Normal file
@@ -0,0 +1,75 @@
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from assistant_bot.handlers import private_voice
|
||||
|
||||
|
||||
class VoiceHandlerTests(unittest.IsolatedAsyncioTestCase):
|
||||
@staticmethod
|
||||
def make_update_and_context(duration: int, transcript: str = "Напомни позвонить"):
|
||||
status_message = SimpleNamespace(edit_text=AsyncMock())
|
||||
message = SimpleNamespace(
|
||||
chat_id=100,
|
||||
voice=SimpleNamespace(duration=duration, file_id="voice-file-id"),
|
||||
reply_text=AsyncMock(return_value=status_message),
|
||||
)
|
||||
update = SimpleNamespace(message=message)
|
||||
telegram_file = SimpleNamespace(download_to_drive=AsyncMock())
|
||||
bot = SimpleNamespace(
|
||||
get_file=AsyncMock(return_value=telegram_file),
|
||||
send_chat_action=AsyncMock(),
|
||||
)
|
||||
recognizer = MagicMock()
|
||||
recognizer.transcribe.return_value = transcript
|
||||
context = SimpleNamespace(
|
||||
bot=bot,
|
||||
application=SimpleNamespace(
|
||||
bot_data={
|
||||
"speech_recognizer": recognizer,
|
||||
"voice_max_duration": 120,
|
||||
}
|
||||
),
|
||||
)
|
||||
return update, context, status_message, recognizer, telegram_file
|
||||
|
||||
async def test_rejects_voice_message_over_duration_limit(self) -> None:
|
||||
update, context, _status, recognizer, _telegram_file = (
|
||||
self.make_update_and_context(duration=121)
|
||||
)
|
||||
|
||||
await private_voice(update, context)
|
||||
|
||||
update.message.reply_text.assert_awaited_once_with(
|
||||
"Голосовое сообщение слишком длинное. Максимум: 120 сек."
|
||||
)
|
||||
recognizer.transcribe.assert_not_called()
|
||||
|
||||
async def test_transcribes_voice_and_passes_text_to_agent(self) -> None:
|
||||
update, context, status, recognizer, telegram_file = (
|
||||
self.make_update_and_context(duration=10)
|
||||
)
|
||||
|
||||
with patch(
|
||||
"assistant_bot.handlers.run_agent_prompt",
|
||||
new_callable=AsyncMock,
|
||||
) as run_agent:
|
||||
await private_voice(update, context)
|
||||
|
||||
recognizer.transcribe.assert_called_once()
|
||||
temporary_path = Path(recognizer.transcribe.call_args.args[0])
|
||||
self.assertFalse(temporary_path.exists())
|
||||
telegram_file.download_to_drive.assert_awaited_once_with(
|
||||
custom_path=temporary_path
|
||||
)
|
||||
status.edit_text.assert_awaited_once_with("Распознано: Напомни позвонить")
|
||||
run_agent.assert_awaited_once_with(
|
||||
update,
|
||||
context,
|
||||
"Напомни позвонить",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
132
tests/test_yandex_ai.py
Normal file
132
tests/test_yandex_ai.py
Normal file
@@ -0,0 +1,132 @@
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from assistant_bot.yandex_ai import YandexAIClient, YandexSpeechRecognizer
|
||||
|
||||
|
||||
class FakeCompletion:
|
||||
def __init__(self) -> None:
|
||||
self.configuration = None
|
||||
self.messages = None
|
||||
self.timeout = None
|
||||
|
||||
def configure(self, **kwargs):
|
||||
self.configuration = kwargs
|
||||
return self
|
||||
|
||||
async def run(self, messages, timeout):
|
||||
self.messages = messages
|
||||
self.timeout = timeout
|
||||
return [SimpleNamespace(text=" готово ")]
|
||||
|
||||
|
||||
class YandexAIClientTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_chat_uses_sdk_messages_and_json_mode(self) -> None:
|
||||
completion = FakeCompletion()
|
||||
requested_models = []
|
||||
|
||||
def completions(model):
|
||||
requested_models.append(model)
|
||||
return completion
|
||||
|
||||
sdk = SimpleNamespace(
|
||||
models=SimpleNamespace(completions=completions),
|
||||
chat=SimpleNamespace(
|
||||
completions=SimpleNamespace(
|
||||
list=self._list_models,
|
||||
)
|
||||
),
|
||||
)
|
||||
client = YandexAIClient(folder_id="folder-id", sdk=sdk)
|
||||
|
||||
answer = await client.chat(
|
||||
"yandexgpt",
|
||||
[
|
||||
{"role": "system", "content": "Отвечай кратко"},
|
||||
{"role": "user", "content": "Привет"},
|
||||
],
|
||||
json_mode=True,
|
||||
)
|
||||
|
||||
self.assertEqual(answer, "готово")
|
||||
self.assertEqual(requested_models, ["gpt://folder-id/yandexgpt"])
|
||||
self.assertEqual(completion.configuration, {"response_format": "json"})
|
||||
self.assertEqual(
|
||||
completion.messages,
|
||||
[
|
||||
{"role": "system", "text": "Отвечай кратко"},
|
||||
{"role": "user", "text": "Привет"},
|
||||
],
|
||||
)
|
||||
self.assertEqual(completion.timeout, 180)
|
||||
|
||||
async def test_list_models_returns_sorted_uris(self) -> None:
|
||||
sdk = SimpleNamespace(
|
||||
chat=SimpleNamespace(
|
||||
completions=SimpleNamespace(
|
||||
list=self._list_models,
|
||||
)
|
||||
)
|
||||
)
|
||||
client = YandexAIClient(folder_id="folder-id", sdk=sdk)
|
||||
|
||||
self.assertEqual(
|
||||
await client.list_models(),
|
||||
["gpt://folder/alice-ai/latest", "gpt://folder/yandexgpt/latest"],
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _list_models():
|
||||
return [
|
||||
SimpleNamespace(uri="gpt://folder/yandexgpt/latest"),
|
||||
SimpleNamespace(uri="gpt://folder/alice-ai/latest"),
|
||||
]
|
||||
|
||||
|
||||
class YandexSpeechRecognizerTests(unittest.TestCase):
|
||||
def test_transcribe_uses_speechkit_ogg_opus(self) -> None:
|
||||
calls = []
|
||||
|
||||
class FakeRecognizer:
|
||||
def run(self, audio, timeout):
|
||||
calls.append((audio, timeout))
|
||||
return SimpleNamespace(text=" Привет, мир ")
|
||||
|
||||
format_value = object()
|
||||
|
||||
def speech_to_text(**kwargs):
|
||||
calls.append(kwargs)
|
||||
return FakeRecognizer()
|
||||
|
||||
sdk = SimpleNamespace(
|
||||
speechkit=SimpleNamespace(
|
||||
AudioFormat=SimpleNamespace(OGG_OPUS=format_value),
|
||||
speech_to_text=speech_to_text,
|
||||
)
|
||||
)
|
||||
recognizer = YandexSpeechRecognizer(
|
||||
folder_id="folder-id",
|
||||
language="ru-RU",
|
||||
model="general",
|
||||
sdk=sdk,
|
||||
)
|
||||
|
||||
with patch.object(Path, "read_bytes", return_value=b"ogg-data"):
|
||||
result = recognizer.transcribe("voice.ogg")
|
||||
|
||||
self.assertEqual(result, "Привет, мир")
|
||||
self.assertEqual(
|
||||
calls[0],
|
||||
{
|
||||
"audio_format": format_value,
|
||||
"language_codes": "ru-RU",
|
||||
"model": "general",
|
||||
},
|
||||
)
|
||||
self.assertEqual(calls[1], (b"ogg-data", 180))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user