tests
This commit is contained in:
@@ -0,0 +1,132 @@
|
||||
import pytest
|
||||
from uuid import uuid4
|
||||
from unittest.mock import AsyncMock
|
||||
try:
|
||||
from backend.src.application.use_cases.collection_use_cases import CollectionUseCases
|
||||
from backend.src.shared.exceptions import NotFoundError, ForbiddenError
|
||||
from backend.src.domain.repositories.collection_access_repository import ICollectionAccessRepository
|
||||
except ImportError:
|
||||
from src.application.use_cases.collection_use_cases import CollectionUseCases
|
||||
from src.shared.exceptions import NotFoundError, ForbiddenError
|
||||
from src.domain.repositories.collection_access_repository import ICollectionAccessRepository
|
||||
|
||||
|
||||
class TestCollectionUseCases:
|
||||
|
||||
@pytest.fixture
|
||||
def collection_use_cases(self, mock_collection_repository, mock_user_repository):
|
||||
mock_access_repository = AsyncMock()
|
||||
mock_access_repository.get_by_user_and_collection = AsyncMock(return_value=None)
|
||||
mock_access_repository.create = AsyncMock()
|
||||
mock_access_repository.delete_by_user_and_collection = AsyncMock(return_value=True)
|
||||
mock_access_repository.list_by_user = AsyncMock(return_value=[])
|
||||
|
||||
return CollectionUseCases(
|
||||
collection_repository=mock_collection_repository,
|
||||
access_repository=mock_access_repository,
|
||||
user_repository=mock_user_repository
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_collection_success(self, collection_use_cases, mock_user,
|
||||
mock_collection_repository, mock_user_repository):
|
||||
owner_id = uuid4()
|
||||
mock_user_repository.get_by_id = AsyncMock(return_value=mock_user)
|
||||
mock_collection_repository.create = AsyncMock(return_value=mock_user)
|
||||
|
||||
result = await collection_use_cases.create_collection(
|
||||
name="Тестовая коллекция",
|
||||
owner_id=owner_id,
|
||||
description="Описание",
|
||||
is_public=False
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
mock_user_repository.get_by_id.assert_called_once_with(owner_id)
|
||||
mock_collection_repository.create.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_collection_user_not_found(self, collection_use_cases, mock_user_repository):
|
||||
owner_id = uuid4()
|
||||
mock_user_repository.get_by_id = AsyncMock(return_value=None)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
await collection_use_cases.create_collection(
|
||||
name="Коллекция",
|
||||
owner_id=owner_id
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_collection_success(self, collection_use_cases, mock_collection, mock_collection_repository):
|
||||
collection_id = uuid4()
|
||||
mock_collection_repository.get_by_id = AsyncMock(return_value=mock_collection)
|
||||
|
||||
result = await collection_use_cases.get_collection(collection_id)
|
||||
|
||||
assert result == mock_collection
|
||||
mock_collection_repository.get_by_id.assert_called_once_with(collection_id)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_collection_not_found(self, collection_use_cases, mock_collection_repository):
|
||||
collection_id = uuid4()
|
||||
mock_collection_repository.get_by_id = AsyncMock(return_value=None)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
await collection_use_cases.get_collection(collection_id)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_collection_success(self, collection_use_cases, mock_collection, mock_collection_repository):
|
||||
collection_id = uuid4()
|
||||
user_id = uuid4()
|
||||
mock_collection.owner_id = user_id
|
||||
mock_collection_repository.get_by_id = AsyncMock(return_value=mock_collection)
|
||||
mock_collection_repository.update = AsyncMock(return_value=mock_collection)
|
||||
|
||||
result = await collection_use_cases.update_collection(
|
||||
collection_id=collection_id,
|
||||
user_id=user_id,
|
||||
name="Обновленное название"
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert mock_collection.name == "Обновленное название"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_collection_forbidden(self, collection_use_cases, mock_collection, mock_collection_repository):
|
||||
collection_id = uuid4()
|
||||
user_id = uuid4()
|
||||
owner_id = uuid4()
|
||||
mock_collection.owner_id = owner_id
|
||||
mock_collection_repository.get_by_id = AsyncMock(return_value=mock_collection)
|
||||
|
||||
with pytest.raises(ForbiddenError):
|
||||
await collection_use_cases.update_collection(
|
||||
collection_id=collection_id,
|
||||
user_id=user_id,
|
||||
name="Название"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_access_owner(self, collection_use_cases, mock_collection, mock_collection_repository):
|
||||
collection_id = uuid4()
|
||||
user_id = uuid4()
|
||||
mock_collection.owner_id = user_id
|
||||
mock_collection.is_public = False
|
||||
mock_collection_repository.get_by_id = AsyncMock(return_value=mock_collection)
|
||||
|
||||
result = await collection_use_cases.check_access(collection_id, user_id)
|
||||
|
||||
assert result is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_access_public(self, collection_use_cases, mock_collection, mock_collection_repository):
|
||||
collection_id = uuid4()
|
||||
user_id = uuid4()
|
||||
owner_id = uuid4()
|
||||
mock_collection.owner_id = owner_id
|
||||
mock_collection.is_public = True
|
||||
mock_collection_repository.get_by_id = AsyncMock(return_value=mock_collection)
|
||||
|
||||
result = await collection_use_cases.check_access(collection_id, user_id)
|
||||
|
||||
assert result is True
|
||||
@@ -0,0 +1,114 @@
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, patch, MagicMock
|
||||
import httpx
|
||||
from tg_bot.infrastructure.external.deepseek_client import DeepSeekClient
|
||||
from tg_bot.infrastructure.external.deepseek_client import DeepSeekAPIError
|
||||
|
||||
|
||||
class TestDeepSeekClient:
|
||||
|
||||
@pytest.fixture
|
||||
def deepseek_client(self):
|
||||
return DeepSeekClient(api_key="test_key", api_url="https://api.test.com/v1/chat/completions")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completion_success(self, deepseek_client):
|
||||
messages = [
|
||||
{"role": "user", "content": "Тестовый вопрос"}
|
||||
]
|
||||
|
||||
mock_response_data = {
|
||||
"choices": [{
|
||||
"message": {
|
||||
"content": "Тестовый ответ от DeepSeek"
|
||||
}
|
||||
}],
|
||||
"usage": {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 20,
|
||||
"total_tokens": 30
|
||||
}
|
||||
}
|
||||
|
||||
with patch('httpx.AsyncClient') as mock_client:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = mock_response_data
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client_instance = AsyncMock()
|
||||
mock_client_instance.__aenter__ = AsyncMock(return_value=mock_client_instance)
|
||||
mock_client_instance.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_client_instance.post = AsyncMock(return_value=mock_response)
|
||||
mock_client.return_value = mock_client_instance
|
||||
|
||||
result = await deepseek_client.chat_completion(messages)
|
||||
|
||||
assert "content" in result
|
||||
assert result["content"] == "Тестовый ответ от DeepSeek"
|
||||
assert "usage" in result
|
||||
assert result["usage"]["total_tokens"] == 30
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completion_no_api_key(self):
|
||||
client = DeepSeekClient(api_key=None)
|
||||
messages = [{"role": "user", "content": "Вопрос"}]
|
||||
|
||||
result = await client.chat_completion(messages)
|
||||
|
||||
assert "content" in result
|
||||
assert "DEEPSEEK_API_KEY" in result["content"] or "не установлен" in result["content"]
|
||||
assert result["usage"]["total_tokens"] == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completion_api_error(self, deepseek_client):
|
||||
import httpx
|
||||
messages = [{"role": "user", "content": "Вопрос"}]
|
||||
|
||||
with patch('httpx.AsyncClient') as mock_client:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 401
|
||||
mock_response.raise_for_status.side_effect = httpx.HTTPStatusError(
|
||||
"Unauthorized", request=MagicMock(), response=mock_response
|
||||
)
|
||||
|
||||
mock_client_instance = AsyncMock()
|
||||
mock_client_instance.__aenter__ = AsyncMock(return_value=mock_client_instance)
|
||||
mock_client_instance.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_client_instance.post = AsyncMock(return_value=mock_response)
|
||||
mock_client.return_value = mock_client_instance
|
||||
|
||||
with pytest.raises(DeepSeekAPIError):
|
||||
await deepseek_client.chat_completion(messages)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completion_with_parameters(self, deepseek_client):
|
||||
messages = [{"role": "user", "content": "Вопрос"}]
|
||||
|
||||
mock_response_data = {
|
||||
"choices": [{"message": {"content": "Ответ"}}],
|
||||
"usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15}
|
||||
}
|
||||
|
||||
with patch('httpx.AsyncClient') as mock_client:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = mock_response_data
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client_instance = AsyncMock()
|
||||
mock_client_instance.__aenter__ = AsyncMock(return_value=mock_client_instance)
|
||||
mock_client_instance.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_client_instance.post = AsyncMock(return_value=mock_response)
|
||||
mock_client.return_value = mock_client_instance
|
||||
|
||||
result = await deepseek_client.chat_completion(
|
||||
messages,
|
||||
model="deepseek-chat",
|
||||
temperature=0.7,
|
||||
max_tokens=100
|
||||
)
|
||||
|
||||
assert result["content"] == "Ответ"
|
||||
call_args = mock_client_instance.post.call_args
|
||||
assert call_args is not None
|
||||
@@ -0,0 +1,141 @@
|
||||
import pytest
|
||||
from uuid import uuid4
|
||||
from unittest.mock import AsyncMock
|
||||
try:
|
||||
from backend.src.application.use_cases.document_use_cases import DocumentUseCases
|
||||
from backend.src.shared.exceptions import NotFoundError, ForbiddenError
|
||||
from backend.src.application.services.document_parser_service import DocumentParserService
|
||||
except ImportError:
|
||||
from src.application.use_cases.document_use_cases import DocumentUseCases
|
||||
from src.shared.exceptions import NotFoundError, ForbiddenError
|
||||
from src.application.services.document_parser_service import DocumentParserService
|
||||
|
||||
|
||||
class TestDocumentUseCases:
|
||||
|
||||
@pytest.fixture
|
||||
def document_use_cases(self, mock_document_repository, mock_collection_repository):
|
||||
mock_parser = AsyncMock()
|
||||
mock_parser.parse_pdf = AsyncMock(return_value=("Парсенный документ", "Содержание"))
|
||||
|
||||
return DocumentUseCases(
|
||||
document_repository=mock_document_repository,
|
||||
collection_repository=mock_collection_repository,
|
||||
parser_service=mock_parser
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_document_success(self, document_use_cases, mock_collection, mock_document_repository, mock_collection_repository):
|
||||
collection_id = uuid4()
|
||||
mock_collection_repository.get_by_id = AsyncMock(return_value=mock_collection)
|
||||
mock_document_repository.create = AsyncMock(return_value=mock_collection)
|
||||
|
||||
result = await document_use_cases.create_document(
|
||||
collection_id=collection_id,
|
||||
title="Тестовый документ",
|
||||
content="Содержание",
|
||||
metadata={"type": "law"}
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
mock_collection_repository.get_by_id.assert_called_once_with(collection_id)
|
||||
mock_document_repository.create.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_document_collection_not_found(self, document_use_cases, mock_collection_repository):
|
||||
collection_id = uuid4()
|
||||
mock_collection_repository.get_by_id = AsyncMock(return_value=None)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
await document_use_cases.create_document(
|
||||
collection_id=collection_id,
|
||||
title="Документ",
|
||||
content="Содержание"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_document_success(self, document_use_cases, mock_document, mock_document_repository):
|
||||
document_id = uuid4()
|
||||
mock_document_repository.get_by_id = AsyncMock(return_value=mock_document)
|
||||
|
||||
result = await document_use_cases.get_document(document_id)
|
||||
|
||||
assert result == mock_document
|
||||
mock_document_repository.get_by_id.assert_called_once_with(document_id)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_document_not_found(self, document_use_cases, mock_document_repository):
|
||||
document_id = uuid4()
|
||||
mock_document_repository.get_by_id = AsyncMock(return_value=None)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
await document_use_cases.get_document(document_id)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_document_success(self, document_use_cases, mock_document, mock_collection,
|
||||
mock_document_repository, mock_collection_repository):
|
||||
document_id = uuid4()
|
||||
user_id = uuid4()
|
||||
mock_document.collection_id = uuid4()
|
||||
mock_collection.owner_id = user_id
|
||||
|
||||
mock_document_repository.get_by_id = AsyncMock(return_value=mock_document)
|
||||
mock_collection_repository.get_by_id = AsyncMock(return_value=mock_collection)
|
||||
mock_document_repository.update = AsyncMock(return_value=mock_document)
|
||||
|
||||
result = await document_use_cases.update_document(
|
||||
document_id=document_id,
|
||||
user_id=user_id,
|
||||
title="Обновленное название"
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert mock_document.title == "Обновленное название"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_document_forbidden(self, document_use_cases, mock_document, mock_collection,
|
||||
mock_document_repository, mock_collection_repository):
|
||||
document_id = uuid4()
|
||||
user_id = uuid4()
|
||||
owner_id = uuid4()
|
||||
mock_document.collection_id = uuid4()
|
||||
mock_collection.owner_id = owner_id
|
||||
|
||||
mock_document_repository.get_by_id = AsyncMock(return_value=mock_document)
|
||||
mock_collection_repository.get_by_id = AsyncMock(return_value=mock_collection)
|
||||
|
||||
with pytest.raises(ForbiddenError):
|
||||
await document_use_cases.update_document(
|
||||
document_id=document_id,
|
||||
user_id=user_id,
|
||||
title="Название"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_document_success(self, document_use_cases, mock_document, mock_collection,
|
||||
mock_document_repository, mock_collection_repository):
|
||||
document_id = uuid4()
|
||||
user_id = uuid4()
|
||||
mock_document.collection_id = uuid4()
|
||||
mock_collection.owner_id = user_id
|
||||
|
||||
mock_document_repository.get_by_id = AsyncMock(return_value=mock_document)
|
||||
mock_collection_repository.get_by_id = AsyncMock(return_value=mock_collection)
|
||||
mock_document_repository.delete = AsyncMock(return_value=True)
|
||||
|
||||
result = await document_use_cases.delete_document(document_id, user_id)
|
||||
|
||||
assert result is True
|
||||
mock_document_repository.delete.assert_called_once_with(document_id)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_collection_documents(self, document_use_cases, mock_collection, mock_documents_list,
|
||||
mock_collection_repository, mock_document_repository):
|
||||
collection_id = uuid4()
|
||||
mock_collection_repository.get_by_id = AsyncMock(return_value=mock_collection)
|
||||
mock_document_repository.list_by_collection = AsyncMock(return_value=mock_documents_list)
|
||||
|
||||
result = await document_use_cases.list_collection_documents(collection_id, skip=0, limit=10)
|
||||
|
||||
assert len(result) == len(mock_documents_list)
|
||||
mock_document_repository.list_by_collection.assert_called_once_with(collection_id, skip=0, limit=10)
|
||||
@@ -0,0 +1,171 @@
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, patch, MagicMock
|
||||
from uuid import uuid4
|
||||
from tg_bot.application.services.rag_service import RAGService
|
||||
|
||||
|
||||
class TestRAGService:
|
||||
|
||||
@pytest.fixture
|
||||
def rag_service(self):
|
||||
service = RAGService()
|
||||
from tg_bot.infrastructure.external.deepseek_client import DeepSeekClient
|
||||
service.deepseek_client = DeepSeekClient()
|
||||
return service
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_documents_in_collections_success(self, rag_service):
|
||||
user_telegram_id = "123456789"
|
||||
query = "трудовой договор"
|
||||
|
||||
mock_documents = [
|
||||
{
|
||||
"document_id": str(uuid4()),
|
||||
"title": "Трудовой кодекс РФ",
|
||||
"content": "Содержание о трудовых договорах",
|
||||
"collection_name": "Законы"
|
||||
},
|
||||
{
|
||||
"document_id": str(uuid4()),
|
||||
"title": "Правила оформления",
|
||||
"content": "Как оформить трудовой договор",
|
||||
"collection_name": "Инструкции"
|
||||
}
|
||||
]
|
||||
|
||||
with patch('aiohttp.ClientSession') as mock_session:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status = 200
|
||||
mock_response.json = AsyncMock(return_value={
|
||||
"user_id": str(uuid4())
|
||||
})
|
||||
|
||||
mock_collections_response = AsyncMock()
|
||||
mock_collections_response.status = 200
|
||||
mock_collections_response.json = AsyncMock(return_value=[
|
||||
{"collection_id": str(uuid4()), "name": "Законы"}
|
||||
])
|
||||
|
||||
mock_search_response = AsyncMock()
|
||||
mock_search_response.status = 200
|
||||
mock_search_response.json = AsyncMock(return_value=mock_documents)
|
||||
|
||||
mock_session_instance = MagicMock()
|
||||
mock_session_instance.__aenter__ = AsyncMock(return_value=mock_session_instance)
|
||||
mock_session_instance.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_session_instance.get = AsyncMock(side_effect=[
|
||||
mock_response,
|
||||
mock_collections_response,
|
||||
mock_search_response
|
||||
])
|
||||
mock_session.return_value = mock_session_instance
|
||||
|
||||
result = await rag_service.search_documents_in_collections(
|
||||
user_telegram_id, query, limit_per_collection=5
|
||||
)
|
||||
|
||||
assert len(result) > 0
|
||||
assert result[0]["title"] == "Трудовой кодекс РФ"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_documents_empty_result(self, rag_service):
|
||||
user_telegram_id = "123456789"
|
||||
query = "несуществующий запрос"
|
||||
|
||||
with patch('aiohttp.ClientSession') as mock_session:
|
||||
mock_response = AsyncMock()
|
||||
mock_response.status = 200
|
||||
mock_response.json = AsyncMock(return_value={
|
||||
"user_id": str(uuid4())
|
||||
})
|
||||
|
||||
mock_collections_response = AsyncMock()
|
||||
mock_collections_response.status = 200
|
||||
mock_collections_response.json = AsyncMock(return_value=[])
|
||||
|
||||
mock_session_instance = MagicMock()
|
||||
mock_session_instance.__aenter__ = AsyncMock(return_value=mock_session_instance)
|
||||
mock_session_instance.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_session_instance.get = AsyncMock(side_effect=[
|
||||
mock_response,
|
||||
mock_collections_response
|
||||
])
|
||||
mock_session.return_value = mock_session_instance
|
||||
|
||||
result = await rag_service.search_documents_in_collections(
|
||||
user_telegram_id, query
|
||||
)
|
||||
|
||||
assert result == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_answer_with_rag_success(self, rag_service, mock_rag_response):
|
||||
question = "Какие права имеет работник?"
|
||||
user_telegram_id = "123456789"
|
||||
|
||||
with patch.object(rag_service, 'search_documents_in_collections') as mock_search, \
|
||||
patch.object(rag_service, 'deepseek_client') as mock_client:
|
||||
|
||||
mock_search.return_value = [
|
||||
{
|
||||
"document_id": str(uuid4()),
|
||||
"title": "Трудовой кодекс",
|
||||
"content": "Работник имеет право на...",
|
||||
"collection_name": "Законы"
|
||||
}
|
||||
]
|
||||
|
||||
mock_client.chat_completion = AsyncMock(return_value={
|
||||
"content": "Работник имеет следующие права...",
|
||||
"usage": {"prompt_tokens": 100, "completion_tokens": 200, "total_tokens": 300}
|
||||
})
|
||||
|
||||
result = await rag_service.generate_answer_with_rag(question, user_telegram_id)
|
||||
|
||||
assert "answer" in result
|
||||
assert "sources" in result
|
||||
assert "usage" in result
|
||||
assert len(result["sources"]) <= 5
|
||||
assert result["answer"] != ""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_answer_limits_to_top5(self, rag_service):
|
||||
question = "Тестовый вопрос"
|
||||
user_telegram_id = "123456789"
|
||||
|
||||
many_documents = [
|
||||
{
|
||||
"document_id": str(uuid4()),
|
||||
"title": f"Документ {i}",
|
||||
"content": f"Содержание {i}",
|
||||
"collection_name": "Коллекция"
|
||||
}
|
||||
for i in range(20)
|
||||
]
|
||||
|
||||
with patch.object(rag_service, 'search_documents_in_collections') as mock_search, \
|
||||
patch.object(rag_service, 'deepseek_client') as mock_client:
|
||||
|
||||
mock_search.return_value = many_documents
|
||||
mock_client.chat_completion = AsyncMock(return_value={
|
||||
"content": "Ответ",
|
||||
"usage": {}
|
||||
})
|
||||
|
||||
result = await rag_service.generate_answer_with_rag(question, user_telegram_id)
|
||||
|
||||
assert len(result["sources"]) == 5
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_answer_no_documents(self, rag_service):
|
||||
question = "Вопрос без документов"
|
||||
user_telegram_id = "123456789"
|
||||
|
||||
with patch.object(rag_service, 'search_documents_in_collections') as mock_search:
|
||||
mock_search.return_value = []
|
||||
|
||||
result = await rag_service.generate_answer_with_rag(question, user_telegram_id)
|
||||
|
||||
assert result["sources"] == []
|
||||
assert "Релевантные документы не найдены" in result.get("answer", "") or \
|
||||
result["answer"] == "No relevant documents found"
|
||||
@@ -0,0 +1,193 @@
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from datetime import datetime, timedelta
|
||||
from tg_bot.domain.services.user_service import UserService
|
||||
from tg_bot.infrastructure.database.models import UserModel
|
||||
|
||||
|
||||
class TestUserService:
|
||||
|
||||
@pytest.fixture
|
||||
def mock_session(self):
|
||||
session = AsyncMock()
|
||||
session.execute = AsyncMock()
|
||||
session.add = MagicMock()
|
||||
session.commit = AsyncMock()
|
||||
session.rollback = AsyncMock()
|
||||
return session
|
||||
|
||||
@pytest.fixture
|
||||
def user_service(self, mock_session):
|
||||
return UserService(mock_session)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_user_by_telegram_id_success(self, user_service, mock_session):
|
||||
telegram_id = 123456789
|
||||
mock_user = UserModel(
|
||||
telegram_id=str(telegram_id),
|
||||
username="test_user",
|
||||
first_name="Test",
|
||||
last_name="User"
|
||||
)
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none = MagicMock(return_value=mock_user)
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
result = await user_service.get_user_by_telegram_id(telegram_id)
|
||||
|
||||
assert result == mock_user
|
||||
assert result.telegram_id == str(telegram_id)
|
||||
mock_session.execute.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_user_by_telegram_id_not_found(self, user_service, mock_session):
|
||||
telegram_id = 999999999
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none = MagicMock(return_value=None)
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
result = await user_service.get_user_by_telegram_id(telegram_id)
|
||||
|
||||
assert result is None
|
||||
mock_session.execute.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_or_create_user_new_user(self, user_service, mock_session):
|
||||
telegram_id = 123456789
|
||||
username = "new_user"
|
||||
first_name = "New"
|
||||
last_name = "User"
|
||||
|
||||
mock_result_not_found = MagicMock()
|
||||
mock_result_not_found.scalar_one_or_none = MagicMock(return_value=None)
|
||||
|
||||
mock_result_found = MagicMock()
|
||||
created_user = UserModel(
|
||||
telegram_id=str(telegram_id),
|
||||
username=username,
|
||||
first_name=first_name,
|
||||
last_name=last_name
|
||||
)
|
||||
mock_result_found.scalar_one_or_none = MagicMock(return_value=created_user)
|
||||
|
||||
mock_session.execute.side_effect = [mock_result_not_found, mock_result_found]
|
||||
|
||||
result = await user_service.get_or_create_user(telegram_id, username, first_name, last_name)
|
||||
|
||||
assert result is not None
|
||||
assert result.telegram_id == str(telegram_id)
|
||||
assert result.username == username
|
||||
mock_session.add.assert_called_once()
|
||||
mock_session.commit.assert_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_or_create_user_existing_user(self, user_service, mock_session):
|
||||
telegram_id = 123456789
|
||||
existing_user = UserModel(
|
||||
telegram_id=str(telegram_id),
|
||||
username="old_username",
|
||||
first_name="Old",
|
||||
last_name="Name"
|
||||
)
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none = MagicMock(return_value=existing_user)
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
result = await user_service.get_or_create_user(
|
||||
telegram_id, "new_username", "New", "Name"
|
||||
)
|
||||
|
||||
assert result == existing_user
|
||||
assert result.username == "new_username"
|
||||
assert result.first_name == "New"
|
||||
assert result.last_name == "Name"
|
||||
mock_session.commit.assert_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_user_questions_success(self, user_service, mock_session):
|
||||
telegram_id = 123456789
|
||||
user = UserModel(
|
||||
telegram_id=str(telegram_id),
|
||||
questions_used=5
|
||||
)
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none = MagicMock(return_value=user)
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
result = await user_service.update_user_questions(telegram_id)
|
||||
|
||||
assert result is True
|
||||
assert user.questions_used == 6
|
||||
mock_session.commit.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_user_questions_user_not_found(self, user_service, mock_session):
|
||||
telegram_id = 999999999
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none = MagicMock(return_value=None)
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
result = await user_service.update_user_questions(telegram_id)
|
||||
|
||||
assert result is False
|
||||
mock_session.commit.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_activate_premium_success(self, user_service, mock_session):
|
||||
telegram_id = 123456789
|
||||
user = UserModel(
|
||||
telegram_id=str(telegram_id),
|
||||
is_premium=False,
|
||||
premium_until=None
|
||||
)
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none = MagicMock(return_value=user)
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
result = await user_service.activate_premium(telegram_id)
|
||||
|
||||
assert result is True
|
||||
assert user.is_premium is True
|
||||
assert user.premium_until is not None
|
||||
assert user.premium_until > datetime.now()
|
||||
mock_session.commit.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_activate_premium_extend_existing(self, user_service, mock_session):
|
||||
telegram_id = 123456789
|
||||
existing_premium_until = datetime.now() + timedelta(days=10)
|
||||
user = UserModel(
|
||||
telegram_id=str(telegram_id),
|
||||
is_premium=True,
|
||||
premium_until=existing_premium_until
|
||||
)
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none = MagicMock(return_value=user)
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
result = await user_service.activate_premium(telegram_id)
|
||||
|
||||
assert result is True
|
||||
assert user.is_premium is True
|
||||
assert user.premium_until > existing_premium_until
|
||||
mock_session.commit.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_activate_premium_user_not_found(self, user_service, mock_session):
|
||||
telegram_id = 999999999
|
||||
|
||||
mock_result = MagicMock()
|
||||
mock_result.scalar_one_or_none = MagicMock(return_value=None)
|
||||
mock_session.execute.return_value = mock_result
|
||||
|
||||
result = await user_service.activate_premium(telegram_id)
|
||||
|
||||
assert result is False
|
||||
mock_session.commit.assert_not_called()
|
||||
Reference in New Issue
Block a user