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"}