""" Unit tests for src.utils.rag_utils module Tests RAG (Retrieval-Augmented Generation) utility functions including: - RAGConfig configuration - Client initialization and caching - Vectorstore creation - Hybrid retriever creation - Document reranking """ import pytest from unittest.mock import Mock, patch, MagicMock from langchain.schema import Document from src.utils.rag_utils import ( RAGConfig, HSC_CONFIG, RERANKER_CONFIGS, get_bedrock_client, get_reranker_client, get_embeddings, create_vectorstore, create_hybrid_retriever, rerank_documents, ) class TestRAGConfig: """Test RAGConfig dataclass""" def test_default_config(self): """Test that HSC_CONFIG has expected default values""" assert HSC_CONFIG.CHUNK_SIZE == 2000 assert HSC_CONFIG.CHUNK_OVERLAP == 200 assert HSC_CONFIG.SEMANTIC_TOP_K == 35 assert HSC_CONFIG.KEYWORD_TOP_K == 25 assert HSC_CONFIG.ENSEMBLE_TOP_K == 40 assert HSC_CONFIG.RERANKER_TOP_K == 15 assert HSC_CONFIG.BM25_WEIGHT == 0.40 assert HSC_CONFIG.VECTOR_WEIGHT == 0.60 assert HSC_CONFIG.RERANKER == "amazon" assert HSC_CONFIG.EMBEDDING_MODEL_ID == "amazon.titan-embed-text-v2:0" assert HSC_CONFIG.BEDROCK_REGION == "us-east-1" def test_custom_config(self): """Test creating custom RAGConfig""" custom_config = RAGConfig( CHUNK_SIZE=1000, CHUNK_OVERLAP=100, RERANKER="cohere" ) assert custom_config.CHUNK_SIZE == 1000 assert custom_config.CHUNK_OVERLAP == 100 assert custom_config.RERANKER == "cohere" # Defaults should still apply assert custom_config.SEMANTIC_TOP_K == 35 def test_reranker_configs(self): """Test RERANKER_CONFIGS dictionary structure""" assert "amazon" in RERANKER_CONFIGS assert "cohere" in RERANKER_CONFIGS # Amazon config assert RERANKER_CONFIGS["amazon"]["model_id"] == "amazon.rerank-v1:0" assert RERANKER_CONFIGS["amazon"]["name"] == "Amazon Rerank 1.0" assert RERANKER_CONFIGS["amazon"]["region"] == "us-west-2" # Cohere config assert RERANKER_CONFIGS["cohere"]["model_id"] == "cohere.rerank-v3-5:0" assert RERANKER_CONFIGS["cohere"]["name"] == "Cohere Rerank 3.5" assert RERANKER_CONFIGS["cohere"]["region"] == "us-east-1" class TestClientFunctions: """Test client initialization and caching functions""" @patch('src.utils.rag_utils.boto3.client') def test_get_bedrock_client(self, mock_boto_client): """Test Bedrock client creation and caching""" # Reset module-level cache import src.utils.rag_utils as rag_utils rag_utils._bedrock_client = None mock_client = Mock() mock_boto_client.return_value = mock_client # First call should create client client1 = get_bedrock_client() assert client1 == mock_client mock_boto_client.assert_called_once_with("bedrock-runtime", region_name="us-east-1") # Second call should return cached client client2 = get_bedrock_client() assert client2 == mock_client assert mock_boto_client.call_count == 1 # Should not create new client @patch('src.utils.rag_utils.boto3.client') def test_get_reranker_client_amazon(self, mock_boto_client): """Test Amazon reranker client creation""" # Reset module-level cache import src.utils.rag_utils as rag_utils rag_utils._reranker_clients = {} mock_client = Mock() mock_boto_client.return_value = mock_client client = get_reranker_client("amazon") assert client == mock_client mock_boto_client.assert_called_with("bedrock-runtime", region_name="us-west-2") @patch('src.utils.rag_utils.boto3.client') def test_get_reranker_client_cohere(self, mock_boto_client): """Test Cohere reranker client creation""" # Reset module-level cache import src.utils.rag_utils as rag_utils rag_utils._reranker_clients = {} mock_client = Mock() mock_boto_client.return_value = mock_client client = get_reranker_client("cohere") assert client == mock_client mock_boto_client.assert_called_with("bedrock-runtime", region_name="us-east-1") @patch('src.utils.rag_utils.get_bedrock_client') @patch('src.utils.rag_utils.BedrockEmbeddings') def test_get_embeddings(self, mock_embeddings_class, mock_get_client): """Test embeddings model creation and caching""" # Reset module-level cache import src.utils.rag_utils as rag_utils rag_utils._embeddings = None mock_client = Mock() mock_get_client.return_value = mock_client mock_embeddings = Mock() mock_embeddings_class.return_value = mock_embeddings # First call should create embeddings embeddings1 = get_embeddings() assert embeddings1 == mock_embeddings mock_embeddings_class.assert_called_once_with( model_id="amazon.titan-embed-text-v2:0", client=mock_client ) # Second call should return cached embeddings embeddings2 = get_embeddings() assert embeddings2 == mock_embeddings assert mock_embeddings_class.call_count == 1 class TestVectorstoreFunctions: """Test vectorstore and retriever creation functions""" @patch('src.utils.rag_utils.FAISS.from_documents') def test_create_vectorstore(self, mock_faiss): """Test FAISS vectorstore creation""" mock_vectorstore = Mock() mock_faiss.return_value = mock_vectorstore mock_embeddings = Mock() chunks = [ Document(page_content="test content 1"), Document(page_content="test content 2") ] result = create_vectorstore(chunks, mock_embeddings) assert result == mock_vectorstore mock_faiss.assert_called_once() # Verify chunks and embeddings were passed call_args = mock_faiss.call_args assert call_args[0][0] == chunks assert call_args[0][1] == mock_embeddings @patch('src.utils.rag_utils.BM25Retriever.from_documents') @patch('src.utils.rag_utils.EnsembleRetriever') def test_create_hybrid_retriever(self, mock_ensemble, mock_bm25): """Test hybrid retriever creation with BM25 and semantic search""" # Setup mocks mock_bm25_retriever = Mock() mock_bm25.return_value = mock_bm25_retriever mock_ensemble_retriever = Mock() mock_ensemble.return_value = mock_ensemble_retriever mock_vectorstore = Mock() mock_vector_retriever = Mock() mock_vectorstore.as_retriever.return_value = mock_vector_retriever chunks = [Document(page_content="test")] config = RAGConfig() result = create_hybrid_retriever(chunks, mock_vectorstore, config) assert result == mock_ensemble_retriever # Verify BM25 was created mock_bm25.assert_called_once() # Verify vectorstore retriever was created mock_vectorstore.as_retriever.assert_called_once() # Verify ensemble was created with correct weights call_args = mock_ensemble.call_args[1] assert call_args['weights'] == [config.VECTOR_WEIGHT, config.BM25_WEIGHT] class TestRerankDocuments: """Test document reranking functionality""" def test_rerank_empty_docs(self): """Test reranking with empty document list""" result = rerank_documents("query", [], Mock()) assert result == [] @patch('src.utils.rag_utils.json.loads') def test_rerank_amazon_model(self, mock_json_loads): """Test reranking with Amazon model""" # Setup mock_client = Mock() mock_response = { "body": Mock(read=Mock(return_value=b'{}')) } mock_client.invoke_model.return_value = mock_response mock_json_loads.return_value = { "results": [ {"index": 0, "relevance_score": 0.9}, {"index": 1, "relevance_score": 0.7} ] } docs = [ Document(page_content="doc1", metadata={"source": "test"}), Document(page_content="doc2", metadata={"source": "test"}) ] result = rerank_documents( query="test query", docs=docs, client=mock_client, top_k=2, model="amazon" ) # Verify results assert len(result) == 2 assert result[0].page_content == "doc1" assert result[0].metadata["rerank_score"] == 0.9 assert result[1].page_content == "doc2" assert result[1].metadata["rerank_score"] == 0.7 @patch('src.utils.rag_utils.json.loads') def test_rerank_cohere_model(self, mock_json_loads): """Test reranking with Cohere model""" mock_client = Mock() mock_response = { "body": Mock(read=Mock(return_value=b'{}')) } mock_client.invoke_model.return_value = mock_response mock_json_loads.return_value = { "results": [ {"index": 1, "relevance_score": 0.8}, {"index": 0, "relevance_score": 0.6} ] } docs = [ Document(page_content="doc1"), Document(page_content="doc2") ] result = rerank_documents( query="test query", docs=docs, client=mock_client, top_k=2, model="cohere" ) # Verify reranked order (doc2 should be first due to higher score) assert len(result) == 2 assert result[0].page_content == "doc2" assert result[1].page_content == "doc1" @patch('src.utils.rag_utils.json.loads', side_effect=Exception("API Error")) @patch('src.utils.rag_utils.logging') def test_rerank_error_handling(self, mock_logging, mock_json_loads): """Test that reranking failures return top_k documents""" docs = [ Document(page_content=f"doc{i}") for i in range(10) ] mock_client = Mock() result = rerank_documents( query="test", docs=docs, client=mock_client, top_k=3, model="amazon" ) # Should return first top_k docs on error assert len(result) == 3 assert result[0].page_content == "doc0" # Verify error was logged mock_logging.error.assert_called_once() def test_rerank_top_k_limit(self): """Test that reranking respects top_k limit""" mock_client = Mock() mock_response = { "body": Mock(read=Mock(return_value=b'{"results": [{"index": 0, "relevance_score": 0.9}, {"index": 1, "relevance_score": 0.8}, {"index": 2, "relevance_score": 0.7}]}')) } mock_client.invoke_model.return_value = mock_response docs = [Document(page_content=f"doc{i}") for i in range(3)] with patch('src.utils.rag_utils.json.loads') as mock_json: mock_json.return_value = { "results": [ {"index": 0, "relevance_score": 0.9}, {"index": 1, "relevance_score": 0.8}, {"index": 2, "relevance_score": 0.7} ] } result = rerank_documents( query="test", docs=docs, client=mock_client, top_k=2, model="amazon" ) assert len(result) == 2