diff --git a/src/backend/core/tests/test_reranking.py b/src/backend/core/tests/test_reranking.py new file mode 100644 index 0000000..66e1a7f --- /dev/null +++ b/src/backend/core/tests/test_reranking.py @@ -0,0 +1,99 @@ +""" +Test suite for reranking service +""" + +import logging +from unittest.mock import MagicMock, patch + +import pytest + +from core.services.reranking import rerank + +pytestmark = pytest.mark.django_db + +logger = logging.getLogger(__name__) + + +ORIGINAL_HITS = [ + { + "_id": "doc-0", + "_score": 1.5, + "_source": { + "title.en": "Dogs", + "content.en": "Dogs are domesticated wolves", + }, + }, + { + "_id": "doc-1", + "_score": 1.2, + "_source": { + "title.en": "Cats", + "content.en": "Cats are popular pets", + }, + }, + { + "_id": "doc-2", + "_score": 1.0, + "_source": { + "title.en": "Birds", + "content.en": "Birds fly in the sky", + }, + }, +] + + +@patch("core.services.reranking.get_reranker") +def test_return_original_results_if_reranker_import_fails(mocked_get_reranker, caplog): + """Test that original results are returned if reranker import fails""" + mocked_get_reranker.return_value = None + + with caplog.at_level(logging.WARNING): + result = rerank("test query", ORIGINAL_HITS) + + assert result == ORIGINAL_HITS + assert any( + "Could not import reranker, returning original results" in message + for message in caplog.messages + ) + + +@patch("core.services.reranking.get_reranker") +def test_return_original_results_if_reranking_fails(mocked_get_reranker, caplog): + """Test that original results are returned if reranking fails""" + mock_reranker = MagicMock() + mock_reranker.rerank.side_effect = Exception("Reranking service error") + mocked_get_reranker.return_value = mock_reranker + + q = "test query" + with caplog.at_level(logging.ERROR): + result = rerank(q, ORIGINAL_HITS) + + assert result == ORIGINAL_HITS + assert any( + "Reranking failed: Reranking service error, returning original results" + in message + for message in caplog.messages + ) + + +def test_reranking_success(caplog, settings): + """Test successful reranking of search results""" + settings.RERANKER_MODEL_NAME = "ms-marco-MiniLM-L-12-v2" + + q = "cats" + with caplog.at_level(logging.INFO): + reranked_hits = rerank(q, ORIGINAL_HITS.copy()) + + assert any( + f"Reranking 3 results for query: {q}" in message for message in caplog.messages + ) + assert any( + "Reranking completed, returned 3 results" in message + for message in caplog.messages + ) + + # if the reranking is not too bad it should definitely rerank doc-1 at top position + assert reranked_hits[0]["_id"] == ORIGINAL_HITS[1]["_id"] + # reranked_hits should be sorted by _reranked_score + assert reranked_hits[0]["_reranked_score"] > reranked_hits[1]["_reranked_score"] + assert reranked_hits[1]["_reranked_score"] > reranked_hits[2]["_reranked_score"] diff --git a/src/backend/core/tests/test_search.py b/src/backend/core/tests/test_search.py index 4d91337..6018546 100644 --- a/src/backend/core/tests/test_search.py +++ b/src/backend/core/tests/test_search.py @@ -5,6 +5,7 @@ Test suite for opensearch search service import logging import operator from json import dumps as json_dumps +from unittest.mock import patch import pytest import responses @@ -712,3 +713,41 @@ def test_search_filtering_by_query_path_and_tag(): assert result["hits"]["total"]["value"] == len(documents_to_search) assert returned_ids == expected_ids + + +@patch("core.services.search.rerank") +def test_reranker_disabled(mocked_rerank, settings): + """Test that reranker is not called when disabled""" + settings.RERANKER_ENABLED = False + documents = bulk_create_documents( + [ + {"title": "wolf", "content": "wolves live in packs and hunt together"}, + {"title": "dog", "content": "dogs are loyal domestic animals"}, + {"title": "cat", "content": "cats are curious and independent pets"}, + ] + ) + service = factories.ServiceFactory(name=SERVICE_NAME) + prepare_index(service.index_name, documents) + + search(q="canine pet", **search_params(service)) + + mocked_rerank.assert_not_called() + + +@patch("core.services.search.rerank") +def test_reranker_enabled(mocked_rerank, settings): + """Test that reranker is called when enabled""" + settings.RERANKER_ENABLED = True + documents = bulk_create_documents( + [ + {"title": "wolf", "content": "wolves live in packs and hunt together"}, + {"title": "dog", "content": "dogs are loyal domestic animals"}, + {"title": "cat", "content": "cats are curious and independent pets"}, + ] + ) + service = factories.ServiceFactory(name=SERVICE_NAME) + prepare_index(service.index_name, documents) + + search(q="canine pet", **search_params(service)) + + mocked_rerank.assert_called_once()