Added azure mysql for mem0 (#3531)
This commit is contained in:
@@ -0,0 +1,128 @@
|
||||
---
|
||||
title: Azure MySQL
|
||||
---
|
||||
|
||||
[Azure Database for MySQL](https://azure.microsoft.com/products/mysql) is a fully managed relational database service that provides enterprise-grade reliability and security. It supports JSON-based vector storage for semantic search capabilities in AI applications.
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
|
||||
os.environ["OPENAI_API_KEY"] = "sk-xx"
|
||||
|
||||
config = {
|
||||
"vector_store": {
|
||||
"provider": "azure_mysql",
|
||||
"config": {
|
||||
"host": "your-server.mysql.database.azure.com",
|
||||
"port": 3306,
|
||||
"user": "your_username",
|
||||
"password": "your_password",
|
||||
"database": "mem0_db",
|
||||
"collection_name": "memories",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
m = Memory.from_config(config)
|
||||
messages = [
|
||||
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
|
||||
{"role": "assistant", "content": "How about thriller movies? They can be quite engaging."},
|
||||
{"role": "user", "content": "I'm not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
|
||||
]
|
||||
m.add(messages, user_id="alice", metadata={"category": "movies"})
|
||||
```
|
||||
|
||||
#### Using Azure Managed Identity
|
||||
|
||||
For production deployments, use Azure Managed Identity instead of passwords:
|
||||
|
||||
```python
|
||||
config = {
|
||||
"vector_store": {
|
||||
"provider": "azure_mysql",
|
||||
"config": {
|
||||
"host": "your-server.mysql.database.azure.com",
|
||||
"user": "your_username",
|
||||
"database": "mem0_db",
|
||||
"collection_name": "memories",
|
||||
"use_azure_credential": True, # Uses DefaultAzureCredential
|
||||
"ssl_disabled": False
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
<Note>
|
||||
When `use_azure_credential` is enabled, the password is obtained via Azure DefaultAzureCredential (supports Managed Identity, Azure CLI, etc.)
|
||||
</Note>
|
||||
|
||||
### Config
|
||||
|
||||
Here are the parameters available for configuring Azure MySQL:
|
||||
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `host` | MySQL server hostname | Required |
|
||||
| `port` | MySQL server port | `3306` |
|
||||
| `user` | Database user | Required |
|
||||
| `password` | Database password (optional with Azure credential) | `None` |
|
||||
| `database` | Database name | Required |
|
||||
| `collection_name` | Table name for storing vectors | `"mem0"` |
|
||||
| `embedding_model_dims` | Dimensions of embedding vectors | `1536` |
|
||||
| `use_azure_credential` | Use Azure DefaultAzureCredential | `False` |
|
||||
| `ssl_ca` | Path to SSL CA certificate | `None` |
|
||||
| `ssl_disabled` | Disable SSL (not recommended) | `False` |
|
||||
| `minconn` | Minimum connections in pool | `1` |
|
||||
| `maxconn` | Maximum connections in pool | `5` |
|
||||
|
||||
### Setup
|
||||
|
||||
#### Create MySQL Flexible Server using Azure CLI:
|
||||
|
||||
```bash
|
||||
# Create resource group
|
||||
az group create --name mem0-rg --location eastus
|
||||
|
||||
# Create MySQL Flexible Server
|
||||
az mysql flexible-server create \
|
||||
--resource-group mem0-rg \
|
||||
--name mem0-mysql-server \
|
||||
--location eastus \
|
||||
--admin-user myadmin \
|
||||
--admin-password <YourPassword> \
|
||||
--version 8.0.21
|
||||
|
||||
# Create database
|
||||
az mysql flexible-server db create \
|
||||
--resource-group mem0-rg \
|
||||
--server-name mem0-mysql-server \
|
||||
--database-name mem0_db
|
||||
|
||||
# Configure firewall
|
||||
az mysql flexible-server firewall-rule create \
|
||||
--resource-group mem0-rg \
|
||||
--name mem0-mysql-server \
|
||||
--rule-name AllowMyIP \
|
||||
--start-ip-address <YourIP> \
|
||||
--end-ip-address <YourIP>
|
||||
```
|
||||
|
||||
#### Enable Azure AD Authentication:
|
||||
|
||||
1. In Azure Portal, navigate to your MySQL Flexible Server
|
||||
2. Go to **Security** > **Authentication** and enable Azure AD
|
||||
3. Add your application's managed identity as a MySQL user:
|
||||
|
||||
```sql
|
||||
CREATE AADUSER 'your-app-identity' IDENTIFIED BY 'your-client-id';
|
||||
GRANT ALL PRIVILEGES ON mem0_db.* TO 'your-app-identity'@'%';
|
||||
FLUSH PRIVILEGES;
|
||||
```
|
||||
|
||||
<Tip>
|
||||
For production, use [Managed Identity](https://learn.microsoft.com/azure/active-directory/managed-identities-azure-resources/) to eliminate password management.
|
||||
</Tip>
|
||||
@@ -156,6 +156,7 @@
|
||||
"components/vectordbs/dbs/pinecone",
|
||||
"components/vectordbs/dbs/mongodb",
|
||||
"components/vectordbs/dbs/azure",
|
||||
"components/vectordbs/dbs/azure_mysql",
|
||||
"components/vectordbs/dbs/redis",
|
||||
"components/vectordbs/dbs/valkey",
|
||||
"components/vectordbs/dbs/elasticsearch",
|
||||
@@ -496,6 +497,7 @@
|
||||
"v0x/components/vectordbs/dbs/pinecone",
|
||||
"v0x/components/vectordbs/dbs/mongodb",
|
||||
"v0x/components/vectordbs/dbs/azure",
|
||||
"v0x/components/vectordbs/dbs/azure_mysql",
|
||||
"v0x/components/vectordbs/dbs/redis",
|
||||
"v0x/components/vectordbs/dbs/valkey",
|
||||
"v0x/components/vectordbs/dbs/elasticsearch",
|
||||
|
||||
@@ -0,0 +1,128 @@
|
||||
---
|
||||
title: Azure MySQL
|
||||
---
|
||||
|
||||
[Azure Database for MySQL](https://azure.microsoft.com/products/mysql) is a fully managed relational database service that provides enterprise-grade reliability and security. It supports JSON-based vector storage for semantic search capabilities in AI applications.
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
|
||||
os.environ["OPENAI_API_KEY"] = "sk-xx"
|
||||
|
||||
config = {
|
||||
"vector_store": {
|
||||
"provider": "azure_mysql",
|
||||
"config": {
|
||||
"host": "your-server.mysql.database.azure.com",
|
||||
"port": 3306,
|
||||
"user": "your_username",
|
||||
"password": "your_password",
|
||||
"database": "mem0_db",
|
||||
"collection_name": "memories",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
m = Memory.from_config(config)
|
||||
messages = [
|
||||
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
|
||||
{"role": "assistant", "content": "How about thriller movies? They can be quite engaging."},
|
||||
{"role": "user", "content": "I'm not a big fan of thriller movies but I love sci-fi movies."},
|
||||
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
|
||||
]
|
||||
m.add(messages, user_id="alice", metadata={"category": "movies"})
|
||||
```
|
||||
|
||||
#### Using Azure Managed Identity
|
||||
|
||||
For production deployments, use Azure Managed Identity instead of passwords:
|
||||
|
||||
```python
|
||||
config = {
|
||||
"vector_store": {
|
||||
"provider": "azure_mysql",
|
||||
"config": {
|
||||
"host": "your-server.mysql.database.azure.com",
|
||||
"user": "your_username",
|
||||
"database": "mem0_db",
|
||||
"collection_name": "memories",
|
||||
"use_azure_credential": True, # Uses DefaultAzureCredential
|
||||
"ssl_disabled": False
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
<Note>
|
||||
When `use_azure_credential` is enabled, the password is obtained via Azure DefaultAzureCredential (supports Managed Identity, Azure CLI, etc.)
|
||||
</Note>
|
||||
|
||||
### Config
|
||||
|
||||
Here are the parameters available for configuring Azure MySQL:
|
||||
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `host` | MySQL server hostname | Required |
|
||||
| `port` | MySQL server port | `3306` |
|
||||
| `user` | Database user | Required |
|
||||
| `password` | Database password (optional with Azure credential) | `None` |
|
||||
| `database` | Database name | Required |
|
||||
| `collection_name` | Table name for storing vectors | `"mem0"` |
|
||||
| `embedding_model_dims` | Dimensions of embedding vectors | `1536` |
|
||||
| `use_azure_credential` | Use Azure DefaultAzureCredential | `False` |
|
||||
| `ssl_ca` | Path to SSL CA certificate | `None` |
|
||||
| `ssl_disabled` | Disable SSL (not recommended) | `False` |
|
||||
| `minconn` | Minimum connections in pool | `1` |
|
||||
| `maxconn` | Maximum connections in pool | `5` |
|
||||
|
||||
### Setup
|
||||
|
||||
#### Create MySQL Flexible Server using Azure CLI:
|
||||
|
||||
```bash
|
||||
# Create resource group
|
||||
az group create --name mem0-rg --location eastus
|
||||
|
||||
# Create MySQL Flexible Server
|
||||
az mysql flexible-server create \
|
||||
--resource-group mem0-rg \
|
||||
--name mem0-mysql-server \
|
||||
--location eastus \
|
||||
--admin-user myadmin \
|
||||
--admin-password <YourPassword> \
|
||||
--version 8.0.21
|
||||
|
||||
# Create database
|
||||
az mysql flexible-server db create \
|
||||
--resource-group mem0-rg \
|
||||
--server-name mem0-mysql-server \
|
||||
--database-name mem0_db
|
||||
|
||||
# Configure firewall
|
||||
az mysql flexible-server firewall-rule create \
|
||||
--resource-group mem0-rg \
|
||||
--name mem0-mysql-server \
|
||||
--rule-name AllowMyIP \
|
||||
--start-ip-address <YourIP> \
|
||||
--end-ip-address <YourIP>
|
||||
```
|
||||
|
||||
#### Enable Azure AD Authentication:
|
||||
|
||||
1. In Azure Portal, navigate to your MySQL Flexible Server
|
||||
2. Go to **Security** > **Authentication** and enable Azure AD
|
||||
3. Add your application's managed identity as a MySQL user:
|
||||
|
||||
```sql
|
||||
CREATE AADUSER 'your-app-identity' IDENTIFIED BY 'your-client-id';
|
||||
GRANT ALL PRIVILEGES ON mem0_db.* TO 'your-app-identity'@'%';
|
||||
FLUSH PRIVILEGES;
|
||||
```
|
||||
|
||||
<Tip>
|
||||
For production, use [Managed Identity](https://learn.microsoft.com/azure/active-directory/managed-identities-azure-resources/) to eliminate password management.
|
||||
</Tip>
|
||||
@@ -0,0 +1,84 @@
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
|
||||
class AzureMySQLConfig(BaseModel):
|
||||
"""Configuration for Azure MySQL vector database."""
|
||||
|
||||
host: str = Field(..., description="MySQL server host (e.g., myserver.mysql.database.azure.com)")
|
||||
port: int = Field(3306, description="MySQL server port")
|
||||
user: str = Field(..., description="Database user")
|
||||
password: Optional[str] = Field(None, description="Database password (not required if using Azure credential)")
|
||||
database: str = Field(..., description="Database name")
|
||||
collection_name: str = Field("mem0", description="Collection/table name")
|
||||
embedding_model_dims: int = Field(1536, description="Dimensions of the embedding model")
|
||||
use_azure_credential: bool = Field(
|
||||
False,
|
||||
description="Use Azure DefaultAzureCredential for authentication instead of password"
|
||||
)
|
||||
ssl_ca: Optional[str] = Field(None, description="Path to SSL CA certificate")
|
||||
ssl_disabled: bool = Field(False, description="Disable SSL connection (not recommended for production)")
|
||||
minconn: int = Field(1, description="Minimum number of connections in the pool")
|
||||
maxconn: int = Field(5, description="Maximum number of connections in the pool")
|
||||
connection_pool: Optional[Any] = Field(
|
||||
None,
|
||||
description="Pre-configured connection pool object (overrides other connection parameters)"
|
||||
)
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def check_auth(cls, values: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Validate authentication parameters."""
|
||||
# If connection_pool is provided, skip validation
|
||||
if values.get("connection_pool") is not None:
|
||||
return values
|
||||
|
||||
use_azure_credential = values.get("use_azure_credential", False)
|
||||
password = values.get("password")
|
||||
|
||||
# Either password or Azure credential must be provided
|
||||
if not use_azure_credential and not password:
|
||||
raise ValueError(
|
||||
"Either 'password' must be provided or 'use_azure_credential' must be set to True"
|
||||
)
|
||||
|
||||
return values
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def check_required_fields(cls, values: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Validate required fields."""
|
||||
# If connection_pool is provided, skip validation of individual parameters
|
||||
if values.get("connection_pool") is not None:
|
||||
return values
|
||||
|
||||
required_fields = ["host", "user", "database"]
|
||||
missing_fields = [field for field in required_fields if not values.get(field)]
|
||||
|
||||
if missing_fields:
|
||||
raise ValueError(
|
||||
f"Missing required fields: {', '.join(missing_fields)}. "
|
||||
f"These fields are required when not using a pre-configured connection_pool."
|
||||
)
|
||||
|
||||
return values
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def validate_extra_fields(cls, values: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Validate that no extra fields are provided."""
|
||||
allowed_fields = set(cls.model_fields.keys())
|
||||
input_fields = set(values.keys())
|
||||
extra_fields = input_fields - allowed_fields
|
||||
|
||||
if extra_fields:
|
||||
raise ValueError(
|
||||
f"Extra fields not allowed: {', '.join(extra_fields)}. "
|
||||
f"Please input only the following fields: {', '.join(allowed_fields)}"
|
||||
)
|
||||
|
||||
return values
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
@@ -162,6 +162,7 @@ class VectorStoreFactory:
|
||||
"milvus": "mem0.vector_stores.milvus.MilvusDB",
|
||||
"upstash_vector": "mem0.vector_stores.upstash_vector.UpstashVector",
|
||||
"azure_ai_search": "mem0.vector_stores.azure_ai_search.AzureAISearch",
|
||||
"azure_mysql": "mem0.vector_stores.azure_mysql.AzureMySQL",
|
||||
"pinecone": "mem0.vector_stores.pinecone.PineconeDB",
|
||||
"mongodb": "mem0.vector_stores.mongodb.MongoDB",
|
||||
"redis": "mem0.vector_stores.redis.RedisDB",
|
||||
|
||||
@@ -0,0 +1,463 @@
|
||||
import json
|
||||
import logging
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
try:
|
||||
import pymysql
|
||||
from pymysql.cursors import DictCursor
|
||||
from dbutils.pooled_db import PooledDB
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"Azure MySQL vector store requires PyMySQL and DBUtils. "
|
||||
"Please install them using 'pip install pymysql dbutils'"
|
||||
)
|
||||
|
||||
try:
|
||||
from azure.identity import DefaultAzureCredential
|
||||
AZURE_IDENTITY_AVAILABLE = True
|
||||
except ImportError:
|
||||
AZURE_IDENTITY_AVAILABLE = False
|
||||
|
||||
from mem0.vector_stores.base import VectorStoreBase
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class OutputData(BaseModel):
|
||||
id: Optional[str]
|
||||
score: Optional[float]
|
||||
payload: Optional[dict]
|
||||
|
||||
|
||||
class AzureMySQL(VectorStoreBase):
|
||||
def __init__(
|
||||
self,
|
||||
host: str,
|
||||
port: int,
|
||||
user: str,
|
||||
password: Optional[str],
|
||||
database: str,
|
||||
collection_name: str,
|
||||
embedding_model_dims: int,
|
||||
use_azure_credential: bool = False,
|
||||
ssl_ca: Optional[str] = None,
|
||||
ssl_disabled: bool = False,
|
||||
minconn: int = 1,
|
||||
maxconn: int = 5,
|
||||
connection_pool: Optional[Any] = None,
|
||||
):
|
||||
"""
|
||||
Initialize the Azure MySQL vector store.
|
||||
|
||||
Args:
|
||||
host (str): MySQL server host
|
||||
port (int): MySQL server port
|
||||
user (str): Database user
|
||||
password (str, optional): Database password (not required if using Azure credential)
|
||||
database (str): Database name
|
||||
collection_name (str): Collection/table name
|
||||
embedding_model_dims (int): Dimension of the embedding vector
|
||||
use_azure_credential (bool): Use Azure DefaultAzureCredential for authentication
|
||||
ssl_ca (str, optional): Path to SSL CA certificate
|
||||
ssl_disabled (bool): Disable SSL connection
|
||||
minconn (int): Minimum number of connections in the pool
|
||||
maxconn (int): Maximum number of connections in the pool
|
||||
connection_pool (Any, optional): Pre-configured connection pool
|
||||
"""
|
||||
self.host = host
|
||||
self.port = port
|
||||
self.user = user
|
||||
self.password = password
|
||||
self.database = database
|
||||
self.collection_name = collection_name
|
||||
self.embedding_model_dims = embedding_model_dims
|
||||
self.use_azure_credential = use_azure_credential
|
||||
self.ssl_ca = ssl_ca
|
||||
self.ssl_disabled = ssl_disabled
|
||||
self.connection_pool = connection_pool
|
||||
|
||||
# Handle Azure authentication
|
||||
if use_azure_credential:
|
||||
if not AZURE_IDENTITY_AVAILABLE:
|
||||
raise ImportError(
|
||||
"Azure Identity is required for Azure credential authentication. "
|
||||
"Please install it using 'pip install azure-identity'"
|
||||
)
|
||||
self._setup_azure_auth()
|
||||
|
||||
# Setup connection pool
|
||||
if self.connection_pool is None:
|
||||
self._setup_connection_pool(minconn, maxconn)
|
||||
|
||||
# Create collection if it doesn't exist
|
||||
collections = self.list_cols()
|
||||
if collection_name not in collections:
|
||||
self.create_col(name=collection_name, vector_size=embedding_model_dims, distance="cosine")
|
||||
|
||||
def _setup_azure_auth(self):
|
||||
"""Setup Azure authentication using DefaultAzureCredential."""
|
||||
try:
|
||||
credential = DefaultAzureCredential()
|
||||
# Get access token for Azure Database for MySQL
|
||||
token = credential.get_token("https://ossrdbms-aad.database.windows.net/.default")
|
||||
# Use token as password
|
||||
self.password = token.token
|
||||
logger.info("Successfully authenticated using Azure DefaultAzureCredential")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to authenticate with Azure: {e}")
|
||||
raise
|
||||
|
||||
def _setup_connection_pool(self, minconn: int, maxconn: int):
|
||||
"""Setup MySQL connection pool."""
|
||||
connect_kwargs = {
|
||||
"host": self.host,
|
||||
"port": self.port,
|
||||
"user": self.user,
|
||||
"password": self.password,
|
||||
"database": self.database,
|
||||
"charset": "utf8mb4",
|
||||
"cursorclass": DictCursor,
|
||||
"autocommit": False,
|
||||
}
|
||||
|
||||
# SSL configuration
|
||||
if not self.ssl_disabled:
|
||||
ssl_config = {"ssl_verify_cert": True}
|
||||
if self.ssl_ca:
|
||||
ssl_config["ssl_ca"] = self.ssl_ca
|
||||
connect_kwargs["ssl"] = ssl_config
|
||||
|
||||
try:
|
||||
self.connection_pool = PooledDB(
|
||||
creator=pymysql,
|
||||
mincached=minconn,
|
||||
maxcached=maxconn,
|
||||
maxconnections=maxconn,
|
||||
blocking=True,
|
||||
**connect_kwargs
|
||||
)
|
||||
logger.info("Successfully created MySQL connection pool")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to create connection pool: {e}")
|
||||
raise
|
||||
|
||||
@contextmanager
|
||||
def _get_cursor(self, commit: bool = False):
|
||||
"""
|
||||
Context manager to get a cursor from the connection pool.
|
||||
Auto-commits or rolls back based on exception.
|
||||
"""
|
||||
conn = self.connection_pool.connection()
|
||||
cur = conn.cursor()
|
||||
try:
|
||||
yield cur
|
||||
if commit:
|
||||
conn.commit()
|
||||
except Exception as exc:
|
||||
conn.rollback()
|
||||
logger.error(f"Database error: {exc}", exc_info=True)
|
||||
raise
|
||||
finally:
|
||||
cur.close()
|
||||
conn.close()
|
||||
|
||||
def create_col(self, name: str = None, vector_size: int = None, distance: str = "cosine"):
|
||||
"""
|
||||
Create a new collection (table in MySQL).
|
||||
Enables vector extension and creates appropriate indexes.
|
||||
|
||||
Args:
|
||||
name (str, optional): Collection name (uses self.collection_name if not provided)
|
||||
vector_size (int, optional): Vector dimension (uses self.embedding_model_dims if not provided)
|
||||
distance (str): Distance metric (cosine, euclidean, dot_product)
|
||||
"""
|
||||
table_name = name or self.collection_name
|
||||
dims = vector_size or self.embedding_model_dims
|
||||
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
# Create table with vector column
|
||||
cur.execute(f"""
|
||||
CREATE TABLE IF NOT EXISTS `{table_name}` (
|
||||
id VARCHAR(255) PRIMARY KEY,
|
||||
vector JSON,
|
||||
payload JSON,
|
||||
INDEX idx_payload_keys ((CAST(payload AS CHAR(255)) ARRAY))
|
||||
)
|
||||
""")
|
||||
logger.info(f"Created collection '{table_name}' with vector dimension {dims}")
|
||||
|
||||
def insert(self, vectors: List[List[float]], payloads: Optional[List[Dict]] = None, ids: Optional[List[str]] = None):
|
||||
"""
|
||||
Insert vectors into the collection.
|
||||
|
||||
Args:
|
||||
vectors (List[List[float]]): List of vectors to insert
|
||||
payloads (List[Dict], optional): List of payloads corresponding to vectors
|
||||
ids (List[str], optional): List of IDs corresponding to vectors
|
||||
"""
|
||||
logger.info(f"Inserting {len(vectors)} vectors into collection {self.collection_name}")
|
||||
|
||||
if payloads is None:
|
||||
payloads = [{}] * len(vectors)
|
||||
if ids is None:
|
||||
import uuid
|
||||
ids = [str(uuid.uuid4()) for _ in range(len(vectors))]
|
||||
|
||||
data = []
|
||||
for vector, payload, vec_id in zip(vectors, payloads, ids):
|
||||
data.append((vec_id, json.dumps(vector), json.dumps(payload)))
|
||||
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
cur.executemany(
|
||||
f"INSERT INTO `{self.collection_name}` (id, vector, payload) VALUES (%s, %s, %s) "
|
||||
f"ON DUPLICATE KEY UPDATE vector = VALUES(vector), payload = VALUES(payload)",
|
||||
data
|
||||
)
|
||||
|
||||
def _cosine_distance(self, vec1_json: str, vec2: List[float]) -> str:
|
||||
"""Generate SQL for cosine distance calculation."""
|
||||
# For MySQL, we need to calculate cosine similarity manually
|
||||
# This is a simplified version - in production, you'd use stored procedures or UDFs
|
||||
return """
|
||||
1 - (
|
||||
(SELECT SUM(a.val * b.val) /
|
||||
(SQRT(SUM(a.val * a.val)) * SQRT(SUM(b.val * b.val))))
|
||||
FROM (
|
||||
SELECT JSON_EXTRACT(vector, CONCAT('$[', idx, ']')) as val
|
||||
FROM (SELECT @row := @row + 1 as idx FROM (SELECT 0 UNION ALL SELECT 1 UNION ALL SELECT 2 UNION ALL SELECT 3) t1, (SELECT 0 UNION ALL SELECT 1 UNION ALL SELECT 2 UNION ALL SELECT 3) t2) indices
|
||||
WHERE idx < JSON_LENGTH(vector)
|
||||
) a,
|
||||
(
|
||||
SELECT JSON_EXTRACT(%s, CONCAT('$[', idx, ']')) as val
|
||||
FROM (SELECT @row := @row + 1 as idx FROM (SELECT 0 UNION ALL SELECT 1 UNION ALL SELECT 2 UNION ALL SELECT 3) t1, (SELECT 0 UNION ALL SELECT 1 UNION ALL SELECT 2 UNION ALL SELECT 3) t2) indices
|
||||
WHERE idx < JSON_LENGTH(%s)
|
||||
) b
|
||||
WHERE a.idx = b.idx
|
||||
)
|
||||
"""
|
||||
|
||||
def search(
|
||||
self,
|
||||
query: str,
|
||||
vectors: List[float],
|
||||
limit: int = 5,
|
||||
filters: Optional[Dict] = None,
|
||||
) -> List[OutputData]:
|
||||
"""
|
||||
Search for similar vectors using cosine similarity.
|
||||
|
||||
Args:
|
||||
query (str): Query string (not used in vector search)
|
||||
vectors (List[float]): Query vector
|
||||
limit (int): Number of results to return
|
||||
filters (Dict, optional): Filters to apply to the search
|
||||
|
||||
Returns:
|
||||
List[OutputData]: Search results
|
||||
"""
|
||||
filter_conditions = []
|
||||
filter_params = []
|
||||
|
||||
if filters:
|
||||
for k, v in filters.items():
|
||||
filter_conditions.append("JSON_EXTRACT(payload, %s) = %s")
|
||||
filter_params.extend([f"$.{k}", json.dumps(v)])
|
||||
|
||||
filter_clause = "WHERE " + " AND ".join(filter_conditions) if filter_conditions else ""
|
||||
|
||||
# For simplicity, we'll compute cosine similarity in Python
|
||||
# In production, you'd want to use MySQL stored procedures or UDFs
|
||||
with self._get_cursor() as cur:
|
||||
query_sql = f"""
|
||||
SELECT id, vector, payload
|
||||
FROM `{self.collection_name}`
|
||||
{filter_clause}
|
||||
"""
|
||||
cur.execute(query_sql, filter_params)
|
||||
results = cur.fetchall()
|
||||
|
||||
# Calculate cosine similarity in Python
|
||||
import numpy as np
|
||||
query_vec = np.array(vectors)
|
||||
scored_results = []
|
||||
|
||||
for row in results:
|
||||
vec = np.array(json.loads(row['vector']))
|
||||
# Cosine similarity
|
||||
similarity = np.dot(query_vec, vec) / (np.linalg.norm(query_vec) * np.linalg.norm(vec))
|
||||
distance = 1 - similarity
|
||||
scored_results.append((row['id'], distance, row['payload']))
|
||||
|
||||
# Sort by distance and limit
|
||||
scored_results.sort(key=lambda x: x[1])
|
||||
scored_results = scored_results[:limit]
|
||||
|
||||
return [
|
||||
OutputData(id=r[0], score=float(r[1]), payload=json.loads(r[2]) if isinstance(r[2], str) else r[2])
|
||||
for r in scored_results
|
||||
]
|
||||
|
||||
def delete(self, vector_id: str):
|
||||
"""
|
||||
Delete a vector by ID.
|
||||
|
||||
Args:
|
||||
vector_id (str): ID of the vector to delete
|
||||
"""
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
cur.execute(f"DELETE FROM `{self.collection_name}` WHERE id = %s", (vector_id,))
|
||||
|
||||
def update(
|
||||
self,
|
||||
vector_id: str,
|
||||
vector: Optional[List[float]] = None,
|
||||
payload: Optional[Dict] = None,
|
||||
):
|
||||
"""
|
||||
Update a vector and its payload.
|
||||
|
||||
Args:
|
||||
vector_id (str): ID of the vector to update
|
||||
vector (List[float], optional): Updated vector
|
||||
payload (Dict, optional): Updated payload
|
||||
"""
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
if vector is not None:
|
||||
cur.execute(
|
||||
f"UPDATE `{self.collection_name}` SET vector = %s WHERE id = %s",
|
||||
(json.dumps(vector), vector_id),
|
||||
)
|
||||
if payload is not None:
|
||||
cur.execute(
|
||||
f"UPDATE `{self.collection_name}` SET payload = %s WHERE id = %s",
|
||||
(json.dumps(payload), vector_id),
|
||||
)
|
||||
|
||||
def get(self, vector_id: str) -> Optional[OutputData]:
|
||||
"""
|
||||
Retrieve a vector by ID.
|
||||
|
||||
Args:
|
||||
vector_id (str): ID of the vector to retrieve
|
||||
|
||||
Returns:
|
||||
OutputData: Retrieved vector or None if not found
|
||||
"""
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(
|
||||
f"SELECT id, vector, payload FROM `{self.collection_name}` WHERE id = %s",
|
||||
(vector_id,),
|
||||
)
|
||||
result = cur.fetchone()
|
||||
if not result:
|
||||
return None
|
||||
return OutputData(
|
||||
id=result['id'],
|
||||
score=None,
|
||||
payload=json.loads(result['payload']) if isinstance(result['payload'], str) else result['payload']
|
||||
)
|
||||
|
||||
def list_cols(self) -> List[str]:
|
||||
"""
|
||||
List all collections (tables).
|
||||
|
||||
Returns:
|
||||
List[str]: List of collection names
|
||||
"""
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute("SHOW TABLES")
|
||||
return [row[f"Tables_in_{self.database}"] for row in cur.fetchall()]
|
||||
|
||||
def delete_col(self):
|
||||
"""Delete the collection (table)."""
|
||||
with self._get_cursor(commit=True) as cur:
|
||||
cur.execute(f"DROP TABLE IF EXISTS `{self.collection_name}`")
|
||||
logger.info(f"Deleted collection '{self.collection_name}'")
|
||||
|
||||
def col_info(self) -> Dict[str, Any]:
|
||||
"""
|
||||
Get information about the collection.
|
||||
|
||||
Returns:
|
||||
Dict[str, Any]: Collection information
|
||||
"""
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute("""
|
||||
SELECT
|
||||
TABLE_NAME as name,
|
||||
TABLE_ROWS as count,
|
||||
ROUND(((DATA_LENGTH + INDEX_LENGTH) / 1024 / 1024), 2) as size_mb
|
||||
FROM information_schema.TABLES
|
||||
WHERE TABLE_SCHEMA = %s AND TABLE_NAME = %s
|
||||
""", (self.database, self.collection_name))
|
||||
result = cur.fetchone()
|
||||
|
||||
if result:
|
||||
return {
|
||||
"name": result['name'],
|
||||
"count": result['count'],
|
||||
"size": f"{result['size_mb']} MB"
|
||||
}
|
||||
return {}
|
||||
|
||||
def list(
|
||||
self,
|
||||
filters: Optional[Dict] = None,
|
||||
limit: int = 100
|
||||
) -> List[List[OutputData]]:
|
||||
"""
|
||||
List all vectors in the collection.
|
||||
|
||||
Args:
|
||||
filters (Dict, optional): Filters to apply
|
||||
limit (int): Number of vectors to return
|
||||
|
||||
Returns:
|
||||
List[List[OutputData]]: List of vectors
|
||||
"""
|
||||
filter_conditions = []
|
||||
filter_params = []
|
||||
|
||||
if filters:
|
||||
for k, v in filters.items():
|
||||
filter_conditions.append("JSON_EXTRACT(payload, %s) = %s")
|
||||
filter_params.extend([f"$.{k}", json.dumps(v)])
|
||||
|
||||
filter_clause = "WHERE " + " AND ".join(filter_conditions) if filter_conditions else ""
|
||||
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(
|
||||
f"""
|
||||
SELECT id, vector, payload
|
||||
FROM `{self.collection_name}`
|
||||
{filter_clause}
|
||||
LIMIT %s
|
||||
""",
|
||||
(*filter_params, limit)
|
||||
)
|
||||
results = cur.fetchall()
|
||||
|
||||
return [[
|
||||
OutputData(
|
||||
id=r['id'],
|
||||
score=None,
|
||||
payload=json.loads(r['payload']) if isinstance(r['payload'], str) else r['payload']
|
||||
) for r in results
|
||||
]]
|
||||
|
||||
def reset(self):
|
||||
"""Reset the collection by deleting and recreating it."""
|
||||
logger.warning(f"Resetting collection {self.collection_name}...")
|
||||
self.delete_col()
|
||||
self.create_col(name=self.collection_name, vector_size=self.embedding_model_dims)
|
||||
|
||||
def __del__(self):
|
||||
"""Close the connection pool when the object is deleted."""
|
||||
try:
|
||||
if hasattr(self, 'connection_pool') and self.connection_pool:
|
||||
self.connection_pool.close()
|
||||
except Exception:
|
||||
pass
|
||||
@@ -21,6 +21,7 @@ class VectorStoreConfig(BaseModel):
|
||||
"neptune": "NeptuneAnalyticsConfig",
|
||||
"upstash_vector": "UpstashVectorConfig",
|
||||
"azure_ai_search": "AzureAISearchConfig",
|
||||
"azure_mysql": "AzureMySQLConfig",
|
||||
"redis": "RedisDBConfig",
|
||||
"valkey": "ValkeyConfig",
|
||||
"databricks": "DatabricksConfig",
|
||||
|
||||
+3
-1
@@ -27,6 +27,7 @@ dependencies = [
|
||||
graph = [
|
||||
"langchain-neo4j>=0.4.0",
|
||||
"langchain-aws>=0.2.23",
|
||||
"langchain-memgraph>=0.1.0",
|
||||
"neo4j>=5.23.1",
|
||||
"rank-bm25>=0.2.2",
|
||||
"kuzu>=0.11.0",
|
||||
@@ -44,6 +45,8 @@ vector_stores = [
|
||||
"psycopg-pool>=3.2.6,<4.0.0",
|
||||
"pymongo>=4.13.2",
|
||||
"pymochow>=2.2.9",
|
||||
"pymysql>=1.1.0",
|
||||
"dbutils>=3.0.3",
|
||||
"valkey>=6.0.0",
|
||||
"databricks-sdk>=0.63.0",
|
||||
"azure-identity>=1.24.0",
|
||||
@@ -69,7 +72,6 @@ extras = [
|
||||
"sentence-transformers>=5.0.0",
|
||||
"elasticsearch>=8.0.0,<9.0.0",
|
||||
"opensearch-py>=2.0.0",
|
||||
"langchain-memgraph>=0.1.0",
|
||||
]
|
||||
test = [
|
||||
"pytest>=8.2.2",
|
||||
|
||||
@@ -0,0 +1,269 @@
|
||||
import json
|
||||
import pytest
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from mem0.vector_stores.azure_mysql import AzureMySQL, OutputData
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_connection_pool():
|
||||
"""Create a mock connection pool."""
|
||||
pool = Mock()
|
||||
conn = Mock()
|
||||
cursor = Mock()
|
||||
|
||||
# Setup cursor mock
|
||||
cursor.fetchall = Mock(return_value=[])
|
||||
cursor.fetchone = Mock(return_value=None)
|
||||
cursor.execute = Mock()
|
||||
cursor.executemany = Mock()
|
||||
cursor.close = Mock()
|
||||
|
||||
# Setup connection mock
|
||||
conn.cursor = Mock(return_value=cursor)
|
||||
conn.commit = Mock()
|
||||
conn.rollback = Mock()
|
||||
conn.close = Mock()
|
||||
|
||||
# Setup pool mock
|
||||
pool.connection = Mock(return_value=conn)
|
||||
pool.close = Mock()
|
||||
|
||||
return pool
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def azure_mysql_instance(mock_connection_pool):
|
||||
"""Create an AzureMySQL instance with mocked connection pool."""
|
||||
with patch('mem0.vector_stores.azure_mysql.PooledDB') as mock_pooled_db:
|
||||
mock_pooled_db.return_value = mock_connection_pool
|
||||
|
||||
instance = AzureMySQL(
|
||||
host="test-server.mysql.database.azure.com",
|
||||
port=3306,
|
||||
user="testuser",
|
||||
password="testpass",
|
||||
database="testdb",
|
||||
collection_name="test_collection",
|
||||
embedding_model_dims=128,
|
||||
use_azure_credential=False,
|
||||
ssl_disabled=True,
|
||||
)
|
||||
instance.connection_pool = mock_connection_pool
|
||||
return instance
|
||||
|
||||
|
||||
def test_azure_mysql_init(mock_connection_pool):
|
||||
"""Test AzureMySQL initialization."""
|
||||
with patch('mem0.vector_stores.azure_mysql.PooledDB') as mock_pooled_db:
|
||||
mock_pooled_db.return_value = mock_connection_pool
|
||||
|
||||
instance = AzureMySQL(
|
||||
host="test-server.mysql.database.azure.com",
|
||||
port=3306,
|
||||
user="testuser",
|
||||
password="testpass",
|
||||
database="testdb",
|
||||
collection_name="test_collection",
|
||||
embedding_model_dims=128,
|
||||
)
|
||||
|
||||
assert instance.host == "test-server.mysql.database.azure.com"
|
||||
assert instance.port == 3306
|
||||
assert instance.user == "testuser"
|
||||
assert instance.database == "testdb"
|
||||
assert instance.collection_name == "test_collection"
|
||||
assert instance.embedding_model_dims == 128
|
||||
|
||||
|
||||
def test_create_col(azure_mysql_instance):
|
||||
"""Test collection creation."""
|
||||
azure_mysql_instance.create_col(name="new_collection", vector_size=256)
|
||||
|
||||
# Verify that execute was called (table creation)
|
||||
conn = azure_mysql_instance.connection_pool.connection()
|
||||
cursor = conn.cursor()
|
||||
assert cursor.execute.called
|
||||
|
||||
|
||||
def test_insert(azure_mysql_instance):
|
||||
"""Test vector insertion."""
|
||||
vectors = [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]
|
||||
payloads = [{"text": "test1"}, {"text": "test2"}]
|
||||
ids = ["id1", "id2"]
|
||||
|
||||
azure_mysql_instance.insert(vectors=vectors, payloads=payloads, ids=ids)
|
||||
|
||||
conn = azure_mysql_instance.connection_pool.connection()
|
||||
cursor = conn.cursor()
|
||||
assert cursor.executemany.called
|
||||
|
||||
|
||||
def test_search(azure_mysql_instance):
|
||||
"""Test vector search."""
|
||||
# Mock the database response
|
||||
conn = azure_mysql_instance.connection_pool.connection()
|
||||
cursor = conn.cursor()
|
||||
cursor.fetchall = Mock(return_value=[
|
||||
{
|
||||
'id': 'id1',
|
||||
'vector': json.dumps([0.1, 0.2, 0.3]),
|
||||
'payload': json.dumps({"text": "test1"})
|
||||
},
|
||||
{
|
||||
'id': 'id2',
|
||||
'vector': json.dumps([0.4, 0.5, 0.6]),
|
||||
'payload': json.dumps({"text": "test2"})
|
||||
}
|
||||
])
|
||||
|
||||
query_vector = [0.2, 0.3, 0.4]
|
||||
results = azure_mysql_instance.search(query="test", vectors=query_vector, limit=5)
|
||||
|
||||
assert isinstance(results, list)
|
||||
assert cursor.execute.called
|
||||
|
||||
|
||||
def test_delete(azure_mysql_instance):
|
||||
"""Test vector deletion."""
|
||||
azure_mysql_instance.delete(vector_id="test_id")
|
||||
|
||||
conn = azure_mysql_instance.connection_pool.connection()
|
||||
cursor = conn.cursor()
|
||||
assert cursor.execute.called
|
||||
|
||||
|
||||
def test_update(azure_mysql_instance):
|
||||
"""Test vector update."""
|
||||
new_vector = [0.7, 0.8, 0.9]
|
||||
new_payload = {"text": "updated"}
|
||||
|
||||
azure_mysql_instance.update(vector_id="test_id", vector=new_vector, payload=new_payload)
|
||||
|
||||
conn = azure_mysql_instance.connection_pool.connection()
|
||||
cursor = conn.cursor()
|
||||
assert cursor.execute.called
|
||||
|
||||
|
||||
def test_get(azure_mysql_instance):
|
||||
"""Test retrieving a vector by ID."""
|
||||
# Mock the database response
|
||||
conn = azure_mysql_instance.connection_pool.connection()
|
||||
cursor = conn.cursor()
|
||||
cursor.fetchone = Mock(return_value={
|
||||
'id': 'test_id',
|
||||
'vector': json.dumps([0.1, 0.2, 0.3]),
|
||||
'payload': json.dumps({"text": "test"})
|
||||
})
|
||||
|
||||
result = azure_mysql_instance.get(vector_id="test_id")
|
||||
|
||||
assert result is not None
|
||||
assert isinstance(result, OutputData)
|
||||
assert result.id == "test_id"
|
||||
|
||||
|
||||
def test_list_cols(azure_mysql_instance):
|
||||
"""Test listing collections."""
|
||||
# Mock the database response
|
||||
conn = azure_mysql_instance.connection_pool.connection()
|
||||
cursor = conn.cursor()
|
||||
cursor.fetchall = Mock(return_value=[
|
||||
{"Tables_in_testdb": "collection1"},
|
||||
{"Tables_in_testdb": "collection2"}
|
||||
])
|
||||
|
||||
collections = azure_mysql_instance.list_cols()
|
||||
|
||||
assert isinstance(collections, list)
|
||||
assert len(collections) == 2
|
||||
|
||||
|
||||
def test_delete_col(azure_mysql_instance):
|
||||
"""Test collection deletion."""
|
||||
azure_mysql_instance.delete_col()
|
||||
|
||||
conn = azure_mysql_instance.connection_pool.connection()
|
||||
cursor = conn.cursor()
|
||||
assert cursor.execute.called
|
||||
|
||||
|
||||
def test_col_info(azure_mysql_instance):
|
||||
"""Test getting collection information."""
|
||||
# Mock the database response
|
||||
conn = azure_mysql_instance.connection_pool.connection()
|
||||
cursor = conn.cursor()
|
||||
cursor.fetchone = Mock(return_value={
|
||||
'name': 'test_collection',
|
||||
'count': 100,
|
||||
'size_mb': 1.5
|
||||
})
|
||||
|
||||
info = azure_mysql_instance.col_info()
|
||||
|
||||
assert isinstance(info, dict)
|
||||
assert cursor.execute.called
|
||||
|
||||
|
||||
def test_list(azure_mysql_instance):
|
||||
"""Test listing vectors."""
|
||||
# Mock the database response
|
||||
conn = azure_mysql_instance.connection_pool.connection()
|
||||
cursor = conn.cursor()
|
||||
cursor.fetchall = Mock(return_value=[
|
||||
{
|
||||
'id': 'id1',
|
||||
'vector': json.dumps([0.1, 0.2, 0.3]),
|
||||
'payload': json.dumps({"text": "test1"})
|
||||
}
|
||||
])
|
||||
|
||||
results = azure_mysql_instance.list(limit=10)
|
||||
|
||||
assert isinstance(results, list)
|
||||
assert len(results) > 0
|
||||
|
||||
|
||||
def test_reset(azure_mysql_instance):
|
||||
"""Test resetting the collection."""
|
||||
azure_mysql_instance.reset()
|
||||
|
||||
conn = azure_mysql_instance.connection_pool.connection()
|
||||
cursor = conn.cursor()
|
||||
# Should call execute at least twice (drop and create)
|
||||
assert cursor.execute.call_count >= 2
|
||||
|
||||
|
||||
@pytest.mark.skipif(True, reason="Requires Azure credentials")
|
||||
def test_azure_credential_authentication():
|
||||
"""Test Azure DefaultAzureCredential authentication."""
|
||||
with patch('mem0.vector_stores.azure_mysql.DefaultAzureCredential') as mock_cred:
|
||||
mock_token = Mock()
|
||||
mock_token.token = "test_token"
|
||||
mock_cred.return_value.get_token.return_value = mock_token
|
||||
|
||||
instance = AzureMySQL(
|
||||
host="test-server.mysql.database.azure.com",
|
||||
port=3306,
|
||||
user="testuser",
|
||||
password=None,
|
||||
database="testdb",
|
||||
collection_name="test_collection",
|
||||
embedding_model_dims=128,
|
||||
use_azure_credential=True,
|
||||
)
|
||||
|
||||
assert instance.password == "test_token"
|
||||
|
||||
|
||||
def test_output_data_model():
|
||||
"""Test OutputData model."""
|
||||
data = OutputData(
|
||||
id="test_id",
|
||||
score=0.95,
|
||||
payload={"text": "test"}
|
||||
)
|
||||
|
||||
assert data.id == "test_id"
|
||||
assert data.score == 0.95
|
||||
assert data.payload == {"text": "test"}
|
||||
Reference in New Issue
Block a user