feature: add Azure Identity for Azure OpenAI and Azure AI Search authentication (#3262)

This commit is contained in:
David A. Torres
2025-08-20 13:30:27 -07:00
committed by GitHub
parent e4c5582808
commit 4487785cec
12 changed files with 686 additions and 15 deletions
+111
View File
@@ -553,3 +553,114 @@ def test_search_basic(azure_ai_search_instance):
assert results[0].id == "doc1"
assert results[0].score == 0.95
assert results[0].payload == {"content": "Test content"}
def test_init_with_valid_api_key(mock_clients):
"""Test __init__ with a valid API key and all required parameters."""
mock_search_client, mock_index_client, mock_azure_key_credential = mock_clients
instance = AzureAISearch(
service_name="test-service",
collection_name="test-index",
api_key="test-api-key",
embedding_model_dims=128,
compression_type="scalar",
use_float16=True,
hybrid_search=True,
vector_filter_mode="preFilter",
)
# Check attributes
assert instance.service_name == "test-service"
assert instance.api_key == "test-api-key"
assert instance.index_name == "test-index"
assert instance.collection_name == "test-index"
assert instance.embedding_model_dims == 128
assert instance.compression_type == "scalar"
assert instance.use_float16 is True
assert instance.hybrid_search is True
assert instance.vector_filter_mode == "preFilter"
# Check that AzureKeyCredential was used
mock_azure_key_credential.assert_called_with("test-api-key")
# Check that user agent was set
mock_search_client._client._config.user_agent_policy.add_user_agent.assert_called_with("mem0")
mock_index_client._client._config.user_agent_policy.add_user_agent.assert_called_with("mem0")
# Check that create_col was called if collection does not exist
mock_index_client.create_or_update_index.assert_called_once()
def test_init_with_default_api_key_triggers_default_credential(monkeypatch, mock_clients):
"""Test __init__ uses DefaultAzureCredential if api_key is None or placeholder."""
mock_search_client, mock_index_client, mock_azure_key_credential = mock_clients
# Patch DefaultAzureCredential to a mock so we can check if it's called
with patch("mem0.vector_stores.azure_ai_search.DefaultAzureCredential") as mock_default_cred:
# Test with api_key=None
AzureAISearch(
service_name="test-service",
collection_name="test-index",
api_key=None,
embedding_model_dims=64,
)
mock_default_cred.assert_called_once()
# Test with api_key=""
AzureAISearch(
service_name="test-service",
collection_name="test-index",
api_key="",
embedding_model_dims=64,
)
assert mock_default_cred.call_count == 2
# Test with api_key="your-api-key"
AzureAISearch(
service_name="test-service",
collection_name="test-index",
api_key="your-api-key",
embedding_model_dims=64,
)
assert mock_default_cred.call_count == 3
def test_init_sets_compression_type_to_none_if_unspecified(mock_clients):
"""Test __init__ sets compression_type to 'none' if not specified."""
mock_search_client, mock_index_client, _ = mock_clients
instance = AzureAISearch(
service_name="test-service",
collection_name="test-index",
api_key="test-api-key",
embedding_model_dims=32,
)
assert instance.compression_type == "none"
def test_init_does_not_create_col_if_collection_exists(mock_clients):
"""Test __init__ does not call create_col if collection already exists."""
mock_search_client, mock_index_client, _ = mock_clients
# Simulate collection already exists
mock_index_client.list_index_names.return_value = ["test-index"]
AzureAISearch(
service_name="test-service",
collection_name="test-index",
api_key="test-api-key",
embedding_model_dims=16,
)
# create_or_update_index should not be called since collection exists
mock_index_client.create_or_update_index.assert_not_called()
def test_init_calls_create_col_if_collection_missing(mock_clients):
"""Test __init__ calls create_col if collection does not exist."""
mock_search_client, mock_index_client, _ = mock_clients
# Simulate collection does not exist
mock_index_client.list_index_names.return_value = []
AzureAISearch(
service_name="test-service",
collection_name="missing-index",
api_key="test-api-key",
embedding_model_dims=16,
)
mock_index_client.create_or_update_index.assert_called_once()