diff --git a/docs/components/vectordbs/dbs/azure_mysql.mdx b/docs/components/vectordbs/dbs/azure_mysql.mdx
new file mode 100644
index 000000000..bfcca4892
--- /dev/null
+++ b/docs/components/vectordbs/dbs/azure_mysql.mdx
@@ -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
+ }
+ }
+}
+```
+
+
+When `use_azure_credential` is enabled, the password is obtained via Azure DefaultAzureCredential (supports Managed Identity, Azure CLI, etc.)
+
+
+### 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 \
+ --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 \
+ --end-ip-address
+```
+
+#### 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;
+```
+
+
+For production, use [Managed Identity](https://learn.microsoft.com/azure/active-directory/managed-identities-azure-resources/) to eliminate password management.
+
diff --git a/docs/docs.json b/docs/docs.json
index e14429538..1ce43f4d7 100644
--- a/docs/docs.json
+++ b/docs/docs.json
@@ -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",
diff --git a/docs/v0x/components/vectordbs/dbs/azure_mysql.mdx b/docs/v0x/components/vectordbs/dbs/azure_mysql.mdx
new file mode 100644
index 000000000..bfcca4892
--- /dev/null
+++ b/docs/v0x/components/vectordbs/dbs/azure_mysql.mdx
@@ -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
+ }
+ }
+}
+```
+
+
+When `use_azure_credential` is enabled, the password is obtained via Azure DefaultAzureCredential (supports Managed Identity, Azure CLI, etc.)
+
+
+### 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 \
+ --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 \
+ --end-ip-address
+```
+
+#### 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;
+```
+
+
+For production, use [Managed Identity](https://learn.microsoft.com/azure/active-directory/managed-identities-azure-resources/) to eliminate password management.
+
diff --git a/mem0/configs/vector_stores/azure_mysql.py b/mem0/configs/vector_stores/azure_mysql.py
new file mode 100644
index 000000000..e5d468612
--- /dev/null
+++ b/mem0/configs/vector_stores/azure_mysql.py
@@ -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
diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py
index c70c1bee1..8a61ac0fe 100644
--- a/mem0/utils/factory.py
+++ b/mem0/utils/factory.py
@@ -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",
diff --git a/mem0/vector_stores/azure_mysql.py b/mem0/vector_stores/azure_mysql.py
new file mode 100644
index 000000000..2d9ab373b
--- /dev/null
+++ b/mem0/vector_stores/azure_mysql.py
@@ -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
diff --git a/mem0/vector_stores/configs.py b/mem0/vector_stores/configs.py
index ff9fd995f..42edf53a6 100644
--- a/mem0/vector_stores/configs.py
+++ b/mem0/vector_stores/configs.py
@@ -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",
diff --git a/pyproject.toml b/pyproject.toml
index 99dbef01a..c166187de 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -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",
diff --git a/tests/vector_stores/test_azure_mysql.py b/tests/vector_stores/test_azure_mysql.py
new file mode 100644
index 000000000..1b6dd6ef3
--- /dev/null
+++ b/tests/vector_stores/test_azure_mysql.py
@@ -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"}