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
@@ -77,6 +77,43 @@ await memory.add(messages, { userId: "john" });
```
</CodeGroup>
As an alternative to using an API key, the Azure Identity credential chain can be used to authenticate with [Azure OpenAI role-based security](https://learn.microsoft.com/en-us/azure/ai-foundry/openai/how-to/role-based-access-control).
<Note> If an API key is provided, it will be used for authentication over an Azure Identity </Note>
Below is a sample configuration for using Mem0 with Azure OpenAI and Azure Identity:
```python
import os
from mem0 import Memory
# You can set the values directly in the config dictionary or use environment variables
os.environ["LLM_AZURE_DEPLOYMENT"] = "your-deployment-name"
os.environ["LLM_AZURE_ENDPOINT"] = "your-api-base-url"
os.environ["LLM_AZURE_API_VERSION"] = "version-to-use"
config = {
"llm": {
"provider": "azure_openai_structured",
"config": {
"model": "your-deployment-name",
"temperature": 0.1,
"max_tokens": 2000,
"azure_kwargs": {
"azure_deployment": "<your-deployment-name>",
"api_version": "<version-to-use>",
"azure_endpoint": "<your-api-base-url>",
"default_headers": {
"CustomHeader": "your-custom-header",
}
}
}
}
}
```
Refer to [Azure Identity troubleshooting tips](https://github.com/Azure/azure-sdk-for-python/blob/main/sdk/identity/azure-identity/TROUBLESHOOTING.md#troubleshoot-environmentcredential-authentication-issues) for setting up an Azure Identity credential.
### Config
Here are the parameters available for configuring Azure OpenAI embedder:
@@ -6,6 +6,8 @@ title: Azure OpenAI
To use Azure OpenAI models, you have to set the `LLM_AZURE_OPENAI_API_KEY`, `LLM_AZURE_ENDPOINT`, `LLM_AZURE_DEPLOYMENT` and `LLM_AZURE_API_VERSION` environment variables. You can obtain the Azure API key from the [Azure](https://azure.microsoft.com/).
Optionally, you can use Azure Identity to authenticate with Azure OpenAI, which allows you to use managed identities or service principals for production and Azure CLI login for development instead of an API key. If an Azure Identity is to be used, ***do not*** set the `LLM_AZURE_OPENAI_API_KEY` environment variable or the api_key in the config dictionary.
> **Note**: The following are currently unsupported with reasoning models `Parallel tool calling`,`temperature`, `top_p`, `presence_penalty`, `frequency_penalty`, `logprobs`, `top_logprobs`, `logit_bias`, `max_tokens`
@@ -116,6 +118,44 @@ config = {
}
```
As an alternative to using an API key, the Azure Identity credential chain can be used to authenticate with [Azure OpenAI role-based security](https://learn.microsoft.com/en-us/azure/ai-foundry/openai/how-to/role-based-access-control).
<Note> If an API key is provided, it will be used for authentication over an Azure Identity </Note>
Below is a sample configuration for using Mem0 with Azure OpenAI and Azure Identity:
```python
import os
from mem0 import Memory
# You can set the values directly in the config dictionary or use environment variables
os.environ["LLM_AZURE_DEPLOYMENT"] = "your-deployment-name"
os.environ["LLM_AZURE_ENDPOINT"] = "your-api-base-url"
os.environ["LLM_AZURE_API_VERSION"] = "version-to-use"
config = {
"llm": {
"provider": "azure_openai_structured",
"config": {
"model": "your-deployment-name",
"temperature": 0.1,
"max_tokens": 2000,
"azure_kwargs": {
"azure_deployment": "<your-deployment-name>",
"api_version": "<version-to-use>",
"azure_endpoint": "<your-api-base-url>",
"default_headers": {
"CustomHeader": "your-custom-header",
}
}
}
}
}
```
Refer to [Azure Identity troubleshooting tips](https://github.com/Azure/azure-sdk-for-python/blob/main/sdk/identity/azure-identity/TROUBLESHOOTING.md#troubleshoot-environmentcredential-authentication-issues) for setting up an Azure Identity credential.
## Config
All available parameters for the `azure_openai` config are present in [Master List of All Params in Config](../config).
+87 -7
View File
@@ -16,8 +16,8 @@ config = {
"vector_store": {
"provider": "azure_ai_search",
"config": {
"service_name": "ai-search-test",
"api_key": "*****",
"service_name": "<your-azure-ai-search-service-name>",
"api_key": "<your-api-key>",
"collection_name": "mem0",
"embedding_model_dims": 1536
}
@@ -41,8 +41,8 @@ config = {
"vector_store": {
"provider": "azure_ai_search",
"config": {
"service_name": "ai-search-test",
"api_key": "*****",
"service_name": "<your-azure-ai-search-service-name>",
"api_key": "<your-api-key>",
"collection_name": "mem0",
"embedding_model_dims": 1536,
"compression_type": "binary",
@@ -59,8 +59,8 @@ config = {
"vector_store": {
"provider": "azure_ai_search",
"config": {
"service_name": "ai-search-test",
"api_key": "*****",
"service_name": "<your-azure-ai-search-service-name>",
"api_key": "<your-api-key>",
"collection_name": "mem0",
"embedding_model_dims": 1536,
"hybrid_search": True,
@@ -70,12 +70,92 @@ config = {
}
```
## Using Azure Identity for Authentication
As an alternative to using an API key, the Azure Identity credential chain can be used to authenticate with Azure OpenAI. The list below shows the order of precedence for credential application:
1. **Environment Credential:**
Azure client ID, secret, tenant ID, or certificate in environment variables for service principal authentication.
2. **Workload Identity Credential:**
Utilizes Azure Workload Identity (relevant for Kubernetes and Azure workloads).
3. **Managed Identity Credential:**
Authenticates as a Managed Identity (for apps/services hosted in Azure with Managed Identity enabled), this is the most secure production credential.
4. **Shared Token Cache Credential / Visual Studio Credential (Windows only):**
Uses cached credentials from Visual Studio sign-ins (and sometimes VS Code if SSO is enabled).
5. **Azure CLI Credential:**
Uses the currently logged-in user from the Azure CLI (`az login`), this is the most common development credential.
6. **Azure PowerShell Credential:**
Uses the identity from Azure PowerShell (`Connect-AzAccount`).
7. **Azure Developer CLI Credential:**
Uses the session from Azure Developer CLI (`azd auth login`).
<Note> If an API is provided, it will be used for authentication over an Azure Identity </Note>
To enable Role-Based Access Control (RBAC) for Azure AI Search, follow these steps:
1. In the Azure Portal, navigate to your **Azure AI Search** service.
2. In the left menu, select **Settings** > **Keys**.
3. Change the authentication setting to **Role-based access control**, or **Both** if you need API key compatibility. The default is “Key-based authentication”—you must switch it to use Azure roles.
4. **Go to Access Control (IAM):**
- In the Azure Portal, select your Search service.
- Click **Access Control (IAM)** on the left.
5. **Add a Role Assignment:**
- Click **Add** > **Add role assignment**.
6. **Choose Role:**
- Mem0 requires the **Search Index Data Contributor** and **Search Service Contributor** role.
7. **Choose Member**
- To assign to a User, Group, Service Principle or Managed Identity:
- For production it is recommended to use a service principal or managed identity.
- For a service principal: select **User, group, or service principal** and search for the service principal.
- For a managed identity: select **Managed identity** and choose the managed identity.
- For development, you can assign the role to a user account.
- For development: select ***User, group, or service principal** and pick a Azure Entra ID account (the same used with `az login`).
8. **Complete the Assignment:**
- Click **Review + Assign**.
If you are using Azure Identity, do not set the `api_key` in the configuration.
```python
config = {
"vector_store": {
"provider": "azure_ai_search",
"config": {
"service_name": "<your-azure-ai-search-service-name>",
"collection_name": "mem0",
"embedding_model_dims": 1536,
"compression_type": "binary",
"use_float16": True # Use half precision for storage efficiency
}
}
}
```
### Environment Variables to set to use Azure Identity Credential:
* For an Environment Credential, you will need to setup a Service Principal and set the following environment variables:
- `AZURE_TENANT_ID`: Your Azure Active Directory tenant ID.
- `AZURE_CLIENT_ID`: The client ID of your service principal or managed identity.
- `AZURE_CLIENT_SECRET`: The client secret of your service principal.
* For a User-Assigned Managed Identity, you will need to set the following environment variable:
- `AZURE_CLIENT_ID`: The client ID of the user-assigned managed identity.
* For a System-Assigned Managed Identity, no additional environment variables are needed.
### Developer logins to use for a Azure Identity Credential:
* For an Azure CLI Credential, you need to have the Azure CLI installed and logged in with `az login`.
* For an Azure PowerShell Credential, you need to have the Azure PowerShell module installed and logged in with `Connect-AzAccount`.
* For an Azure Developer CLI Credential, you need to have the Azure Developer CLI installed and logged in with `azd auth login`.
Troubleshooting tips for [Azure Identity](https://github.com/Azure/azure-sdk-for-python/blob/main/sdk/identity/azure-identity/TROUBLESHOOTING.md#troubleshoot-environmentcredential-authentication-issues).
## Configuration Parameters
| Parameter | Description | Default Value | Options |
| --- | --- | --- | --- |
| `service_name` | Azure AI Search service name | Required | - |
| `api_key` | API key of the Azure AI Search service | Required | - |
| `api_key` | API key of the Azure AI Search service | Optional | If not present, the [Azure Identity](#using-azure-identity-for-authentication) credential chain will be used |
| `collection_name` | The name of the collection/index to store vectors | `mem0` | Any valid index name |
| `embedding_model_dims` | Dimensions of the embedding model | `1536` | Any integer value |
| `compression_type` | Type of vector compression to use | `none` | `none`, `scalar`, `binary` |
+15
View File
@@ -1,11 +1,14 @@
import os
from typing import Literal, Optional
from azure.identity import DefaultAzureCredential, get_bearer_token_provider
from openai import AzureOpenAI
from mem0.configs.embeddings.base import BaseEmbedderConfig
from mem0.embeddings.base import EmbeddingBase
SCOPE = "https://cognitiveservices.azure.com/.default"
class AzureOpenAIEmbedding(EmbeddingBase):
def __init__(self, config: Optional[BaseEmbedderConfig] = None):
@@ -17,9 +20,21 @@ class AzureOpenAIEmbedding(EmbeddingBase):
api_version = self.config.azure_kwargs.api_version or os.getenv("EMBEDDING_AZURE_API_VERSION")
default_headers = self.config.azure_kwargs.default_headers
# If the API key is not provided or is a placeholder, use DefaultAzureCredential.
if api_key is None or api_key == "" or api_key == "your-api-key":
self.credential = DefaultAzureCredential()
azure_ad_token_provider = get_bearer_token_provider(
self.credential,
SCOPE,
)
api_key = None
else:
azure_ad_token_provider = None
self.client = AzureOpenAI(
azure_deployment=azure_deployment,
azure_endpoint=azure_endpoint,
azure_ad_token_provider=azure_ad_token_provider,
api_version=api_version,
api_key=api_key,
http_client=self.config.http_client,
+15
View File
@@ -2,6 +2,7 @@ import json
import os
from typing import Dict, List, Optional, Union
from azure.identity import DefaultAzureCredential, get_bearer_token_provider
from openai import AzureOpenAI
from mem0.configs.llms.azure import AzureOpenAIConfig
@@ -9,6 +10,8 @@ from mem0.configs.llms.base import BaseLlmConfig
from mem0.llms.base import LLMBase
from mem0.memory.utils import extract_json
SCOPE = "https://cognitiveservices.azure.com/.default"
class AzureOpenAILLM(LLMBase):
def __init__(self, config: Optional[Union[BaseLlmConfig, AzureOpenAIConfig, Dict]] = None):
@@ -43,9 +46,21 @@ class AzureOpenAILLM(LLMBase):
api_version = self.config.azure_kwargs.api_version or os.getenv("LLM_AZURE_API_VERSION")
default_headers = self.config.azure_kwargs.default_headers
# If the API key is not provided or is a placeholder, use DefaultAzureCredential.
if api_key is None or api_key == "" or api_key == "your-api-key":
self.credential = DefaultAzureCredential()
azure_ad_token_provider = get_bearer_token_provider(
self.credential,
SCOPE,
)
api_key = None
else:
azure_ad_token_provider = None
self.client = AzureOpenAI(
azure_deployment=azure_deployment,
azure_endpoint=azure_endpoint,
azure_ad_token_provider=azure_ad_token_provider,
api_version=api_version,
api_key=api_key,
http_client=self.config.http_client,
+19 -4
View File
@@ -1,11 +1,14 @@
import os
from typing import Dict, List, Optional
from azure.identity import DefaultAzureCredential, get_bearer_token_provider
from openai import AzureOpenAI
from mem0.configs.llms.base import BaseLlmConfig
from mem0.llms.base import LLMBase
SCOPE = "https://cognitiveservices.azure.com/.default"
class AzureOpenAIStructuredLLM(LLMBase):
def __init__(self, config: Optional[BaseLlmConfig] = None):
@@ -15,16 +18,28 @@ class AzureOpenAIStructuredLLM(LLMBase):
if not self.config.model:
self.config.model = "gpt-4o-2024-08-06"
api_key = os.getenv("LLM_AZURE_OPENAI_API_KEY") or self.config.azure_kwargs.api_key
azure_deployment = os.getenv("LLM_AZURE_DEPLOYMENT") or self.config.azure_kwargs.azure_deployment
azure_endpoint = os.getenv("LLM_AZURE_ENDPOINT") or self.config.azure_kwargs.azure_endpoint
api_version = os.getenv("LLM_AZURE_API_VERSION") or self.config.azure_kwargs.api_version
api_key = self.config.azure_kwargs.api_key or os.getenv("LLM_AZURE_OPENAI_API_KEY")
azure_deployment = self.config.azure_kwargs.azure_deployment or os.getenv("LLM_AZURE_DEPLOYMENT")
azure_endpoint = self.config.azure_kwargs.azure_endpoint or os.getenv("LLM_AZURE_ENDPOINT")
api_version = self.config.azure_kwargs.api_version or os.getenv("LLM_AZURE_API_VERSION")
default_headers = self.config.azure_kwargs.default_headers
# If the API key is not provided or is a placeholder, use DefaultAzureCredential.
if api_key is None or api_key == "" or api_key == "your-api-key":
self.credential = DefaultAzureCredential()
azure_ad_token_provider = get_bearer_token_provider(
self.credential,
SCOPE,
)
api_key = None
else:
azure_ad_token_provider = None
# Can display a warning if API version is of model and api-version
self.client = AzureOpenAI(
azure_deployment=azure_deployment,
azure_endpoint=azure_endpoint,
azure_ad_token_provider=azure_ad_token_provider,
api_version=api_version,
api_key=api_key,
http_client=self.config.http_client,
+19 -4
View File
@@ -11,6 +11,7 @@ from mem0.vector_stores.base import VectorStoreBase
try:
from azure.core.credentials import AzureKeyCredential
from azure.core.exceptions import ResourceNotFoundError
from azure.identity import DefaultAzureCredential
from azure.search.documents import SearchClient
from azure.search.documents.indexes import SearchIndexClient
from azure.search.documents.indexes.models import (
@@ -77,14 +78,21 @@ class AzureAISearch(VectorStoreBase):
self.hybrid_search = hybrid_search
self.vector_filter_mode = vector_filter_mode
# If the API key is not provided or is a placeholder, use DefaultAzureCredential.
if self.api_key is None or self.api_key == "" or self.api_key == "your-api-key":
credential = DefaultAzureCredential()
self.api_key = None
else:
credential = AzureKeyCredential(self.api_key)
self.search_client = SearchClient(
endpoint=f"https://{service_name}.search.windows.net",
index_name=self.index_name,
credential=AzureKeyCredential(api_key),
credential=credential,
)
self.index_client = SearchIndexClient(
endpoint=f"https://{service_name}.search.windows.net",
credential=AzureKeyCredential(api_key),
credential=credential,
)
self.search_client._client._config.user_agent_policy.add_user_agent("mem0")
@@ -358,16 +366,23 @@ class AzureAISearch(VectorStoreBase):
# Delete the collection
self.delete_col()
# If the API key is not provided or is a placeholder, use DefaultAzureCredential.
if self.api_key is None or self.api_key == "" or self.api_key == "your-api-key":
credential = DefaultAzureCredential()
self.api_key = None
else:
credential = AzureKeyCredential(self.api_key)
# Reinitialize the clients
service_endpoint = f"https://{self.service_name}.search.windows.net"
self.search_client = SearchClient(
endpoint=service_endpoint,
index_name=self.index_name,
credential=AzureKeyCredential(self.api_key),
credential=credential,
)
self.index_client = SearchIndexClient(
endpoint=service_endpoint,
credential=AzureKeyCredential(self.api_key),
credential=credential,
)
# Add user agent
+1
View File
@@ -42,6 +42,7 @@ vector_stores = [
"pymongo>=4.13.2",
"pymochow>=2.2.9",
"databricks-sdk>=0.63.0",
"azure-identity>=1.24.0",
]
llms = [
"groq>=0.3.0",
@@ -50,3 +50,117 @@ def test_embed_text_with_default_headers(default_headers, expected_header):
assert embedder.client.api_key == "test"
assert embedder.client._api_version == "test_version"
assert embedder.client.default_headers.get("Test") == expected_header
@pytest.fixture
def base_embedder_config():
class DummyAzureKwargs:
api_key = None
azure_deployment = None
azure_endpoint = None
api_version = None
default_headers = None
class DummyConfig(BaseEmbedderConfig):
azure_kwargs = DummyAzureKwargs()
http_client = None
model = "test-model"
return DummyConfig()
def test_init_with_api_key(monkeypatch, base_embedder_config):
base_embedder_config.azure_kwargs.api_key = "test-key"
base_embedder_config.azure_kwargs.azure_deployment = "test-deployment"
base_embedder_config.azure_kwargs.azure_endpoint = "https://test.endpoint"
base_embedder_config.azure_kwargs.api_version = "2024-01-01"
base_embedder_config.azure_kwargs.default_headers = {"X-Test": "Header"}
with (
patch("mem0.embeddings.azure_openai.AzureOpenAI") as mock_azure_openai,
patch("mem0.embeddings.azure_openai.DefaultAzureCredential") as mock_cred,
patch("mem0.embeddings.azure_openai.get_bearer_token_provider") as mock_token_provider,
):
AzureOpenAIEmbedding(base_embedder_config)
mock_azure_openai.assert_called_once_with(
azure_deployment="test-deployment",
azure_endpoint="https://test.endpoint",
azure_ad_token_provider=None,
api_version="2024-01-01",
api_key="test-key",
http_client=None,
default_headers={"X-Test": "Header"},
)
mock_cred.assert_not_called()
mock_token_provider.assert_not_called()
def test_init_with_env_vars(monkeypatch, base_embedder_config):
monkeypatch.setenv("EMBEDDING_AZURE_OPENAI_API_KEY", "env-key")
monkeypatch.setenv("EMBEDDING_AZURE_DEPLOYMENT", "env-deployment")
monkeypatch.setenv("EMBEDDING_AZURE_ENDPOINT", "https://env.endpoint")
monkeypatch.setenv("EMBEDDING_AZURE_API_VERSION", "2024-02-02")
with patch("mem0.embeddings.azure_openai.AzureOpenAI") as mock_azure_openai:
AzureOpenAIEmbedding(base_embedder_config)
mock_azure_openai.assert_called_once_with(
azure_deployment="env-deployment",
azure_endpoint="https://env.endpoint",
azure_ad_token_provider=None,
api_version="2024-02-02",
api_key="env-key",
http_client=None,
default_headers=None,
)
def test_init_with_default_azure_credential(monkeypatch, base_embedder_config):
base_embedder_config.azure_kwargs.api_key = ""
with (
patch("mem0.embeddings.azure_openai.DefaultAzureCredential") as mock_cred,
patch("mem0.embeddings.azure_openai.get_bearer_token_provider") as mock_token_provider,
patch("mem0.embeddings.azure_openai.AzureOpenAI") as mock_azure_openai,
):
mock_cred_instance = Mock()
mock_cred.return_value = mock_cred_instance
mock_token_provider_instance = Mock()
mock_token_provider.return_value = mock_token_provider_instance
AzureOpenAIEmbedding(base_embedder_config)
mock_cred.assert_called_once()
mock_token_provider.assert_called_once_with(mock_cred_instance, "https://cognitiveservices.azure.com/.default")
mock_azure_openai.assert_called_once_with(
azure_deployment=None,
azure_endpoint=None,
azure_ad_token_provider=mock_token_provider_instance,
api_version=None,
api_key=None,
http_client=None,
default_headers=None,
)
def test_init_with_placeholder_api_key(monkeypatch, base_embedder_config):
base_embedder_config.azure_kwargs.api_key = "your-api-key"
with (
patch("mem0.embeddings.azure_openai.DefaultAzureCredential") as mock_cred,
patch("mem0.embeddings.azure_openai.get_bearer_token_provider") as mock_token_provider,
patch("mem0.embeddings.azure_openai.AzureOpenAI") as mock_azure_openai,
):
mock_cred_instance = Mock()
mock_cred.return_value = mock_cred_instance
mock_token_provider_instance = Mock()
mock_token_provider.return_value = mock_token_provider_instance
AzureOpenAIEmbedding(base_embedder_config)
mock_cred.assert_called_once()
mock_token_provider.assert_called_once_with(mock_cred_instance, "https://cognitiveservices.azure.com/.default")
mock_azure_openai.assert_called_once_with(
azure_deployment=None,
azure_endpoint=None,
azure_ad_token_provider=mock_token_provider_instance,
api_version=None,
api_key=None,
http_client=None,
default_headers=None,
)
+128
View File
@@ -124,7 +124,135 @@ def test_generate_with_http_proxies(default_headers):
http_client=mock_http_client_instance,
azure_deployment=None,
azure_endpoint=None,
azure_ad_token_provider=None,
api_version=None,
default_headers=default_headers,
)
mock_http_client.assert_called_once_with(proxies="http://testproxy.mem0.net:8000")
def test_init_with_api_key(monkeypatch):
# Patch environment variables to None to force config usage
monkeypatch.delenv("LLM_AZURE_OPENAI_API_KEY", raising=False)
monkeypatch.delenv("LLM_AZURE_DEPLOYMENT", raising=False)
monkeypatch.delenv("LLM_AZURE_ENDPOINT", raising=False)
monkeypatch.delenv("LLM_AZURE_API_VERSION", raising=False)
config = AzureOpenAIConfig(
model=MODEL,
temperature=TEMPERATURE,
max_tokens=MAX_TOKENS,
top_p=TOP_P,
)
# Set Azure kwargs directly
config.azure_kwargs.api_key = "test-key"
config.azure_kwargs.azure_deployment = "test-deployment"
config.azure_kwargs.azure_endpoint = "https://test-endpoint"
config.azure_kwargs.api_version = "2024-01-01"
config.azure_kwargs.default_headers = {"x-test": "header"}
config.http_client = None
with patch("mem0.llms.azure_openai.AzureOpenAI") as mock_azure_openai:
llm = AzureOpenAILLM(config)
mock_azure_openai.assert_called_once_with(
azure_deployment="test-deployment",
azure_endpoint="https://test-endpoint",
azure_ad_token_provider=None,
api_version="2024-01-01",
api_key="test-key",
http_client=None,
default_headers={"x-test": "header"},
)
assert llm.config.model == MODEL
def test_init_with_env_vars(monkeypatch):
monkeypatch.setenv("LLM_AZURE_OPENAI_API_KEY", "env-key")
monkeypatch.setenv("LLM_AZURE_DEPLOYMENT", "env-deployment")
monkeypatch.setenv("LLM_AZURE_ENDPOINT", "https://env-endpoint")
monkeypatch.setenv("LLM_AZURE_API_VERSION", "2024-02-02")
config = AzureOpenAIConfig(model=None)
config.azure_kwargs.api_key = None
config.azure_kwargs.azure_deployment = None
config.azure_kwargs.azure_endpoint = None
config.azure_kwargs.api_version = None
config.azure_kwargs.default_headers = None
config.http_client = None
with patch("mem0.llms.azure_openai.AzureOpenAI") as mock_azure_openai:
llm = AzureOpenAILLM(config)
mock_azure_openai.assert_called_once_with(
azure_deployment="env-deployment",
azure_endpoint="https://env-endpoint",
azure_ad_token_provider=None,
api_version="2024-02-02",
api_key="env-key",
http_client=None,
default_headers=None,
)
# Should default to "gpt-4o" if model is None
assert llm.config.model == "gpt-4o"
def test_init_with_default_azure_credential(monkeypatch):
# No API key in config or env, triggers DefaultAzureCredential
monkeypatch.delenv("LLM_AZURE_OPENAI_API_KEY", raising=False)
config = AzureOpenAIConfig(model=MODEL)
config.azure_kwargs.api_key = None
config.azure_kwargs.azure_deployment = "dep"
config.azure_kwargs.azure_endpoint = "https://endpoint"
config.azure_kwargs.api_version = "2024-03-03"
config.azure_kwargs.default_headers = None
config.http_client = None
with (
patch("mem0.llms.azure_openai.DefaultAzureCredential") as mock_cred,
patch("mem0.llms.azure_openai.get_bearer_token_provider") as mock_token_provider,
patch("mem0.llms.azure_openai.AzureOpenAI") as mock_azure_openai,
):
mock_cred_instance = mock_cred.return_value
mock_token_provider.return_value = "token-provider"
AzureOpenAILLM(config)
mock_cred.assert_called_once()
mock_token_provider.assert_called_once_with(mock_cred_instance, "https://cognitiveservices.azure.com/.default")
mock_azure_openai.assert_called_once_with(
azure_deployment="dep",
azure_endpoint="https://endpoint",
azure_ad_token_provider="token-provider",
api_version="2024-03-03",
api_key=None,
http_client=None,
default_headers=None,
)
def test_init_with_placeholder_api_key(monkeypatch):
# Placeholder API key should trigger DefaultAzureCredential
config = AzureOpenAIConfig(model=MODEL)
config.azure_kwargs.api_key = "your-api-key"
config.azure_kwargs.azure_deployment = "dep"
config.azure_kwargs.azure_endpoint = "https://endpoint"
config.azure_kwargs.api_version = "2024-04-04"
config.azure_kwargs.default_headers = None
config.http_client = None
with (
patch("mem0.llms.azure_openai.DefaultAzureCredential") as mock_cred,
patch("mem0.llms.azure_openai.get_bearer_token_provider") as mock_token_provider,
patch("mem0.llms.azure_openai.AzureOpenAI") as mock_azure_openai,
):
mock_cred_instance = mock_cred.return_value
mock_token_provider.return_value = "token-provider"
AzureOpenAILLM(config)
mock_cred.assert_called_once()
mock_token_provider.assert_called_once_with(mock_cred_instance, "https://cognitiveservices.azure.com/.default")
mock_azure_openai.assert_called_once_with(
azure_deployment="dep",
azure_endpoint="https://endpoint",
azure_ad_token_provider="token-provider",
api_version="2024-04-04",
api_key=None,
http_client=None,
default_headers=None,
)
+100
View File
@@ -0,0 +1,100 @@
from unittest import mock
from mem0.llms.azure_openai_structured import SCOPE, AzureOpenAIStructuredLLM
class DummyAzureKwargs:
def __init__(
self,
api_key=None,
azure_deployment="test-deployment",
azure_endpoint="https://test-endpoint.openai.azure.com",
api_version="2024-06-01-preview",
default_headers=None,
):
self.api_key = api_key
self.azure_deployment = azure_deployment
self.azure_endpoint = azure_endpoint
self.api_version = api_version
self.default_headers = default_headers
class DummyConfig:
def __init__(
self,
model=None,
azure_kwargs=None,
temperature=0.7,
max_tokens=256,
top_p=1.0,
http_client=None,
):
self.model = model
self.azure_kwargs = azure_kwargs or DummyAzureKwargs()
self.temperature = temperature
self.max_tokens = max_tokens
self.top_p = top_p
self.http_client = http_client
@mock.patch("mem0.llms.azure_openai_structured.AzureOpenAI")
def test_init_with_api_key(mock_azure_openai):
config = DummyConfig(model="test-model", azure_kwargs=DummyAzureKwargs(api_key="real-key"))
llm = AzureOpenAIStructuredLLM(config)
assert llm.config.model == "test-model"
mock_azure_openai.assert_called_once()
args, kwargs = mock_azure_openai.call_args
assert kwargs["api_key"] == "real-key"
assert kwargs["azure_ad_token_provider"] is None
@mock.patch("mem0.llms.azure_openai_structured.AzureOpenAI")
@mock.patch("mem0.llms.azure_openai_structured.get_bearer_token_provider")
@mock.patch("mem0.llms.azure_openai_structured.DefaultAzureCredential")
def test_init_with_default_credential(mock_credential, mock_token_provider, mock_azure_openai):
config = DummyConfig(model=None, azure_kwargs=DummyAzureKwargs(api_key=None))
mock_token_provider.return_value = "token-provider"
llm = AzureOpenAIStructuredLLM(config)
# Should set default model if not provided
assert llm.config.model == "gpt-4o-2024-08-06"
mock_credential.assert_called_once()
mock_token_provider.assert_called_once_with(mock_credential.return_value, SCOPE)
mock_azure_openai.assert_called_once()
args, kwargs = mock_azure_openai.call_args
assert kwargs["api_key"] is None
assert kwargs["azure_ad_token_provider"] == "token-provider"
def test_init_with_env_vars(monkeypatch, mocker):
mock_azure_openai = mocker.patch("mem0.llms.azure_openai_structured.AzureOpenAI")
monkeypatch.setenv("LLM_AZURE_DEPLOYMENT", "test-deployment")
monkeypatch.setenv("LLM_AZURE_ENDPOINT", "https://test-endpoint.openai.azure.com")
monkeypatch.setenv("LLM_AZURE_API_VERSION", "2024-06-01-preview")
config = DummyConfig(model="test-model", azure_kwargs=DummyAzureKwargs(api_key=None))
AzureOpenAIStructuredLLM(config)
mock_azure_openai.assert_called_once()
args, kwargs = mock_azure_openai.call_args
assert kwargs["api_key"] is None
assert kwargs["azure_deployment"] == "test-deployment"
assert kwargs["azure_endpoint"] == "https://test-endpoint.openai.azure.com"
assert kwargs["api_version"] == "2024-06-01-preview"
@mock.patch("mem0.llms.azure_openai_structured.AzureOpenAI")
def test_init_with_placeholder_api_key_uses_default_credential(
mock_azure_openai,
):
with (
mock.patch("mem0.llms.azure_openai_structured.DefaultAzureCredential") as mock_credential,
mock.patch("mem0.llms.azure_openai_structured.get_bearer_token_provider") as mock_token_provider,
):
config = DummyConfig(model=None, azure_kwargs=DummyAzureKwargs(api_key="your-api-key"))
mock_token_provider.return_value = "token-provider"
llm = AzureOpenAIStructuredLLM(config)
assert llm.config.model == "gpt-4o-2024-08-06"
mock_credential.assert_called_once()
mock_token_provider.assert_called_once_with(mock_credential.return_value, SCOPE)
mock_azure_openai.assert_called_once()
args, kwargs = mock_azure_openai.call_args
assert kwargs["api_key"] is None
assert kwargs["azure_ad_token_provider"] == "token-provider"
+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()