feature: add Azure Identity for Azure OpenAI and Azure AI Search authentication (#3262)
This commit is contained in:
@@ -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).
|
||||
|
||||
@@ -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` |
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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"
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user