Files
doczyai-pipelines/fieldExtraction/tests/test_rag_utils.py
T

333 lines
11 KiB
Python
Raw Normal View History

"""
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