197 lines
6.1 KiB
Python
197 lines
6.1 KiB
Python
import unittest
|
|
from pathlib import Path
|
|
from types import SimpleNamespace
|
|
from unittest.mock import patch
|
|
|
|
from assistant_bot.yandex_ai import (
|
|
YandexAIClient,
|
|
YandexSpeechRecognizer,
|
|
describe_yandex_error,
|
|
prepare_chat_messages,
|
|
)
|
|
|
|
|
|
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 FakeChatCompletions:
|
|
def __init__(self, completion, list_models) -> None:
|
|
self.completion = completion
|
|
self.list = list_models
|
|
self.requested_models = []
|
|
|
|
def __call__(self, model):
|
|
self.requested_models.append(model)
|
|
return self.completion
|
|
|
|
|
|
class YandexAIClientTests(unittest.IsolatedAsyncioTestCase):
|
|
def test_multiple_system_messages_are_merged_at_the_beginning(self) -> None:
|
|
self.assertEqual(
|
|
prepare_chat_messages(
|
|
[
|
|
{"role": "system", "content": "Primary instructions"},
|
|
{"role": "user", "content": "Earlier question"},
|
|
{"role": "system", "content": "Runtime context"},
|
|
{"role": "assistant", "content": "Earlier answer"},
|
|
{"role": "user", "content": "Current question"},
|
|
]
|
|
),
|
|
[
|
|
{
|
|
"role": "system",
|
|
"content": "Primary instructions\n\nRuntime context",
|
|
},
|
|
{"role": "user", "content": "Earlier question"},
|
|
{"role": "assistant", "content": "Earlier answer"},
|
|
{"role": "user", "content": "Current question"},
|
|
],
|
|
)
|
|
|
|
def test_http_error_includes_safe_response_body(self) -> None:
|
|
error = RuntimeError("forbidden")
|
|
error.response = SimpleNamespace(
|
|
status_code=403,
|
|
text='{"error":{"message":"folder mismatch"}}',
|
|
)
|
|
|
|
self.assertEqual(
|
|
describe_yandex_error(error),
|
|
'HTTP 403: {"error":{"message":"folder mismatch"}}',
|
|
)
|
|
|
|
async def test_chat_uses_sdk_messages_and_json_mode(self) -> None:
|
|
completion = FakeCompletion()
|
|
completions = FakeChatCompletions(completion, self._list_models)
|
|
|
|
sdk = SimpleNamespace(
|
|
chat=SimpleNamespace(
|
|
completions=completions,
|
|
),
|
|
)
|
|
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(
|
|
completions.requested_models,
|
|
["gpt://folder-id/yandexgpt"],
|
|
)
|
|
self.assertEqual(completion.configuration, {"response_format": "json"})
|
|
self.assertEqual(
|
|
completion.messages,
|
|
[
|
|
{"role": "system", "content": "Отвечай кратко"},
|
|
{"role": "user", "content": "Привет"},
|
|
],
|
|
)
|
|
self.assertEqual(completion.timeout, 180)
|
|
|
|
async def test_chat_repairs_legacy_uri_without_folder_id(self) -> None:
|
|
completion = FakeCompletion()
|
|
completions = FakeChatCompletions(completion, self._list_models)
|
|
sdk = SimpleNamespace(
|
|
chat=SimpleNamespace(completions=completions),
|
|
)
|
|
client = YandexAIClient(folder_id="folder-id", sdk=sdk)
|
|
|
|
await client.chat(
|
|
"gpt://qwen3.6-35b-a3b/latest",
|
|
[{"role": "user", "content": "Привет"}],
|
|
)
|
|
|
|
self.assertEqual(
|
|
completions.requested_models,
|
|
["gpt://folder-id/qwen3.6-35b-a3b/latest"],
|
|
)
|
|
|
|
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()
|