333 lines
11 KiB
Python
333 lines
11 KiB
Python
|
|
"""
|
||
|
|
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
|