From d6d89c987bddf580870db14c69db974edfc5263c Mon Sep 17 00:00:00 2001 From: Elif Sema Balcioglu Date: Thu, 23 Jul 2026 19:00:01 +0300 Subject: [PATCH] Add Oracle Vector Store Integration (#5358) Co-authored-by: kartik-mem0 --- docs/components/vectordbs/config.mdx | 7 +- docs/components/vectordbs/dbs/oracledb.mdx | 134 +++ docs/components/vectordbs/overview.mdx | 3 +- docs/docs.json | 1 + docs/images/provider-icons/oracle.svg | 1 + docs/llms.txt | 1 + mem0/configs/vector_stores/oracledb.py | 113 +++ mem0/utils/factory.py | 1 + mem0/vector_stores/configs.py | 1 + mem0/vector_stores/oracledb.py | 592 +++++++++++ pyproject.toml | 1 + tests/vector_stores/test_oracledb.py | 1055 ++++++++++++++++++++ 12 files changed, 1908 insertions(+), 2 deletions(-) create mode 100644 docs/components/vectordbs/dbs/oracledb.mdx create mode 100644 docs/images/provider-icons/oracle.svg create mode 100644 mem0/configs/vector_stores/oracledb.py create mode 100644 mem0/vector_stores/oracledb.py create mode 100644 tests/vector_stores/test_oracledb.py diff --git a/docs/components/vectordbs/config.mdx b/docs/components/vectordbs/config.mdx index 28676d177..2f1ae75d7 100644 --- a/docs/components/vectordbs/config.mdx +++ b/docs/components/vectordbs/config.mdx @@ -7,7 +7,7 @@ description: "Reference for vector database configuration options in Mem0, inclu The `config` is defined as an object with two main keys: - `vector_store`: Specifies the vector database provider and its configuration - - `provider`: The name of the vector database (e.g., "chroma", "pgvector", "qdrant", "milvus", "upstash_vector", "azure_ai_search", "vertex_ai_vector_search", "valkey") + - `provider`: The name of the vector database (e.g., "chroma", "pgvector", "qdrant", "milvus", "upstash_vector", "azure_ai_search", "vertex_ai_vector_search", "valkey", "oracledb") - `config`: A nested dictionary containing provider-specific settings @@ -95,6 +95,11 @@ Here's a comprehensive list of all parameters that can be used across different | `connection_string` | PostgreSQL connection string (for Supabase/PGVector) | | `index_method` | Vector index method (for Supabase) | | `index_measure` | Distance measure for similarity search (for Supabase) | +| `connection_params` | Connection settings for Oracle AI Vector Search | +| `use_connection_pool` | Create an Oracle connection pool from `connection_params` | +| `distance_metric` | Distance metric for Oracle vector indexing and search | +| `index_type` | Oracle vector index type: `HNSW` or `IVF` | +| `index_parameters` | Oracle vector-index parameters for the selected index type | | Parameter | Description | diff --git a/docs/components/vectordbs/dbs/oracledb.mdx b/docs/components/vectordbs/dbs/oracledb.mdx new file mode 100644 index 000000000..0f8367bc6 --- /dev/null +++ b/docs/components/vectordbs/dbs/oracledb.mdx @@ -0,0 +1,134 @@ +--- +title: "Oracle AI Vector Search" +description: "Use Oracle Database AI Vector Search as a vector store in Mem0 for semantic and relational queries." +--- + +[Oracle AI Vector Search](https://www.oracle.com/database/ai-vector-search/) stores embeddings in an Oracle table using the native `VECTOR` data type, so you can combine semantic search over unstructured data with relational queries over business data in a single database. + +### Requirements + +- Oracle Database 23.4 or later, with a user that can create tables and vector indexes +- The `python-oracledb` driver. In thick mode, Oracle Client 23.4 or later is also required. + +```bash +pip install oracledb +``` + +### Usage + + +```python Python +import os +from mem0 import Memory + +os.environ["OPENAI_API_KEY"] = "sk-xx" + +config = { + "vector_store": { + "provider": "oracledb", + "config": { + "collection_name": "mem0", + "embedding_model_dims": 1536, + "connection_params": { + "user": "mem0_user", + "password": "your-password", + "dsn": "localhost:1521/FREEPDB1", + }, + } + } +} + +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"}) +``` + + +To reuse a connection or pool you already manage, pass it as `client` instead of `connection_params`: + +```python +import oracledb + +pool = oracledb.create_pool(user="mem0_user", password="your-password", dsn="localhost:1521/FREEPDB1") + +config = { + "vector_store": { + "provider": "oracledb", + "config": {"client": pool}, + } +} +``` + +### Config + +Here are the parameters available for configuring Oracle AI Vector Search: + +| Parameter | Description | Default Value | +| --- | --- | --- | +| `connection_params` | Connection settings passed to `python-oracledb`, such as `user`, `password` and `dsn`. See the [connection handling guide](https://python-oracledb.readthedocs.io/en/latest/user_guide/connection_handling.html). | `None` | +| `use_connection_pool` | Create a connection pool from `connection_params` instead of a single connection | `True` | +| `client` | An existing `oracledb.Connection` or `oracledb.ConnectionPool` to use instead of building one from `connection_params` | `None` | +| `collection_name` | Name of the Oracle table that stores vectors and payloads | `mem0` | +| `embedding_model_dims` | Dimension of your embedding vectors, must be greater than 0 | `1536` | +| `distance_metric` | Distance function used for indexing and search: `COSINE`, `EUCLIDEAN`, `EUCLIDEAN_SQUARED`, `DOT`, `HAMMING` or `MANHATTAN` | `COSINE` | +| `do_create_index` | Whether to create a vector index on the collection | `True` | +| `index_type` | Vector index type: `HNSW` or `IVF` | `HNSW` | +| `index_name` | Name of the vector index | `_VEC_IDX` | +| `index_parameters` | Index tuning parameters. For `HNSW`: `neighbors`, `efconstruction`. For `IVF`: `neighbor partitions`, `samples_per_partition`, `min_vectors_per_partition`. | `None` | +| `index_accuracy` | Target index accuracy from 1 to 100, applied as `WITH TARGET ACCURACY ` | `None` | + + + When you pass a pre-built `client`, Mem0 uses it as-is and ignores `connection_params` and `use_connection_pool`. Mem0 does not close a client it did not create. + + +### Vector indexes + +Set the index type with `index_type` and tune it with `index_parameters`: + +```python +config = { + "vector_store": { + "provider": "oracledb", + "config": { + "connection_params": {"user": "mem0_user", "password": "your-password", "dsn": "localhost:1521/FREEPDB1"}, + "index_type": "HNSW", + "index_parameters": {"neighbors": 32, "efconstruction": 200}, + "index_accuracy": 95, + } + } +} +``` + +For the full list of supported options, see the Oracle [`CREATE VECTOR INDEX`](https://docs.oracle.com/en/database/oracle/oracle-database/26/sqlrf/create-vector-index.html) reference. + +### Search scores + +Oracle returns a distance from `VECTOR_DISTANCE`, which Mem0 converts to a `score` where higher means more similar. `COSINE` and the other non-negative metrics produce scores in the range `[0, 1]`. `DOT` returns the inner product, which can fall outside that range. + +### Metadata filters + +Filters run against the JSON `payload` column and support: + +| Filter type | Examples | +| --- | --- | +| Scalar equality | `{"user_id": "alice"}` | +| Field existence | `{"agent_id": "*"}` | +| Comparison | `{"score": {"gte": 0.5}}`, also `eq`, `ne`, `gt`, `lt`, `lte` | +| Membership | `{"category": {"in": ["movies", "books"]}}`, also `nin` | +| String matching | `{"title": {"contains": "sci-fi"}}`, also `icontains` for case-insensitive | +| Logical groups | `{"AND": [...]}`, `{"OR": [...]}`, `{"NOT": [...]}` | + +Multiple fields at the top level are combined with `AND`: + +```python +m.search( + "movie recommendations", + user_id="alice", + filters={"category": {"in": ["movies", "books"]}, "rating": {"gte": 4}}, +) +``` diff --git a/docs/components/vectordbs/overview.mdx b/docs/components/vectordbs/overview.mdx index 18d2209fe..1367f2cf9 100644 --- a/docs/components/vectordbs/overview.mdx +++ b/docs/components/vectordbs/overview.mdx @@ -1,6 +1,6 @@ --- title: Overview -description: "Overview of all supported vector databases in Mem0, including Qdrant, Chroma, PGVector, Pinecone, and more." +description: "Overview of all supported vector databases in Mem0, including Qdrant, Chroma, PGVector, Pinecone, Oracle, and more." --- Mem0 includes built-in support for various popular databases. Memory can utilize the database provided by the user, ensuring efficient use for specific needs. @@ -21,6 +21,7 @@ See the list of supported vector databases below. + diff --git a/docs/docs.json b/docs/docs.json index 8b7bef8a4..2c3b4bfd6 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -207,6 +207,7 @@ "components/vectordbs/dbs/milvus", "components/vectordbs/dbs/pinecone", "components/vectordbs/dbs/mongodb", + "components/vectordbs/dbs/oracledb", "components/vectordbs/dbs/azure", "components/vectordbs/dbs/azure_mysql", "components/vectordbs/dbs/redis", diff --git a/docs/images/provider-icons/oracle.svg b/docs/images/provider-icons/oracle.svg new file mode 100644 index 000000000..b111d0892 --- /dev/null +++ b/docs/images/provider-icons/oracle.svg @@ -0,0 +1 @@ +Oracle diff --git a/docs/llms.txt b/docs/llms.txt index 9b4f15e65..bf56f0751 100644 --- a/docs/llms.txt +++ b/docs/llms.txt @@ -472,6 +472,7 @@ Everything below is OSS-only provider configuration. Skip this entire section wh - [Milvus](https://docs.mem0.ai/components/vectordbs/dbs/milvus) [OSS]: Use for large-scale Milvus deployments. - [Pinecone](https://docs.mem0.ai/components/vectordbs/dbs/pinecone) [OSS]: Use when the user is on Pinecone managed. - [MongoDB](https://docs.mem0.ai/components/vectordbs/dbs/mongodb) [OSS]: Use when Mongo Atlas Vector Search is the backing store. +- [Oracle AI Vector Search](https://docs.mem0.ai/components/vectordbs/dbs/oracledb) [OSS]: Use when Oracle Database AI Vector Search is the backing store. - [Azure AI Search](https://docs.mem0.ai/components/vectordbs/dbs/azure) [OSS]: Use when the user is on Azure AI Search. - [Azure MySQL](https://docs.mem0.ai/components/vectordbs/dbs/azure_mysql) [OSS]: Use when vector search runs on Azure Database for MySQL. - [Redis](https://docs.mem0.ai/components/vectordbs/dbs/redis) [OSS]: Use when Redis Stack is the backing store. diff --git a/mem0/configs/vector_stores/oracledb.py b/mem0/configs/vector_stores/oracledb.py new file mode 100644 index 000000000..b9ad6be99 --- /dev/null +++ b/mem0/configs/vector_stores/oracledb.py @@ -0,0 +1,113 @@ +"""Pydantic configuration for the Oracle AI Vector Search integration.""" + +import re +from typing import Any, Dict, Literal, Optional + +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator + + +def _quote_identifier(name: str) -> str: + name = name.strip() + reg = r'^(?:"[^"]+"|[^".]+)(?:\.(?:"[^"]+"|[^".]+))*$' + pattern_validate = re.compile(reg) + + if not pattern_validate.match(name): + raise ValueError(f"Identifier name {name} is not valid.") + + pattern_match = r'"([^"]+)"|([^".]+)' + groups = re.findall(pattern_match, name) + groups = [m[0] or m[1] for m in groups] + groups = [f'"{g}"' for g in groups] + + return ".".join(groups) + + +class HnswParams(BaseModel): + model_config = ConfigDict(extra="forbid", strict=True) + + neighbors: Optional[int] = Field(None, ge=2, le=2048) + efconstruction: Optional[int] = Field(None, ge=1, le=65535) + + +class IvfParams(BaseModel): + model_config = ConfigDict(extra="forbid", strict=True) + + neighbor_partitions: Optional[int] = Field(None, alias="neighbor partitions", ge=1, le=10_000_000) + samples_per_partition: Optional[int] = Field(None, ge=1) + min_vectors_per_partition: Optional[int] = Field(None, ge=0) + + +class OracleAIVectorSearchConfig(BaseModel): + """Configuration required to connect to an Oracle database with vector search enabled.""" + + connection_params: Optional[dict] = Field(None, description="Database connection parameters, including auth.") + use_connection_pool: bool = Field( + True, + description="Create a ConnectionPool instead of a single Connection when no client is provided", + ) + + client: Optional[Any] = Field( + None, description="Oracle Connection or ConnectionPool (overrides connection string and individual parameters)" + ) + + collection_name: str = Field("mem0", description="Default name for the collection") + embedding_model_dims: int = Field(1536, description="Dimension of the embedding vectors") + distance_metric: Literal["EUCLIDEAN", "EUCLIDEAN_SQUARED", "COSINE", "DOT", "HAMMING", "MANHATTAN"] = Field( + "COSINE", + description="Similarity metric: EUCLIDEAN, EUCLIDEAN_SQUARED, COSINE, DOT, HAMMING or MANHATTAN. Defaults to COSINE", + ) + + do_create_index: Optional[bool] = Field(True, description="Optional whether to create index") + index_type: Literal["HNSW", "IVF"] = Field("HNSW", description="Optional index type, HNSW or IVF") + index_name: Optional[str] = Field(None, description="Optional custom name for the vector index") + index_parameters: Optional[dict] = Field( + None, + description="Optional structured CREATE VECTOR INDEX parameters", + ) + index_accuracy: Optional[int] = Field(None, description="Optional index accuracy") + + @field_validator("distance_metric", "index_type", mode="before") + @classmethod + def _normalize_uppercase(cls, value: Any) -> Any: + return value.upper() if isinstance(value, str) else value + + @model_validator(mode="after") + def _validate_model(self): + """Normalise attributes and validate identifiers/metrics.""" + + if not self.connection_params and not self.client: + raise ValueError("Must provide at least one of `connection_params` and `client`") + + if self.index_name is None: + self.index_name = f"{self.collection_name}_VEC_IDX" + + self.index_name = _quote_identifier(self.index_name) + self.collection_name = _quote_identifier(self.collection_name) + + if self.index_parameters is not None: + parameter_model = HnswParams if self.index_type == "HNSW" else IvfParams + self.index_parameters = parameter_model.model_validate(self.index_parameters).model_dump( + by_alias=True, + exclude_none=True, + ) + + if self.index_accuracy and not (0 < self.index_accuracy <= 100): + raise ValueError("`index_accuracy` must be between 1 and 100") + + if not (0 < self.embedding_model_dims): + raise ValueError("`embedding_model_dims` must be bigger than 0") + + return self + + @model_validator(mode="before") + @classmethod + def validate_extra_fields(cls, values: Dict[str, Any]) -> Dict[str, Any]: + allowed_fields = set(cls.model_fields.keys()) + extra_fields = set(values.keys()) - allowed_fields + if extra_fields: + raise ValueError( + "Extra fields not allowed: {}. Please input only the following fields: {}".format( + ", ".join(sorted(extra_fields)), ", ".join(sorted(allowed_fields)) + ) + ) + return values diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py index cf01ccafc..866a4da73 100644 --- a/mem0/utils/factory.py +++ b/mem0/utils/factory.py @@ -201,6 +201,7 @@ class VectorStoreFactory: "cassandra": "mem0.vector_stores.cassandra.CassandraDB", "neptune": "mem0.vector_stores.neptune_analytics.NeptuneAnalyticsVector", "turbopuffer": "mem0.vector_stores.turbopuffer.TurbopufferDB", + "oracledb": "mem0.vector_stores.oracledb.OracleAIVectorSearch", } @classmethod diff --git a/mem0/vector_stores/configs.py b/mem0/vector_stores/configs.py index 32459dddc..990086c7d 100644 --- a/mem0/vector_stores/configs.py +++ b/mem0/vector_stores/configs.py @@ -35,6 +35,7 @@ class VectorStoreConfig(BaseModel): "langchain": "LangchainConfig", "s3_vectors": "S3VectorsConfig", "turbopuffer": "TurbopufferConfig", + "oracledb": "OracleAIVectorSearchConfig", } @model_validator(mode="after") diff --git a/mem0/vector_stores/oracledb.py b/mem0/vector_stores/oracledb.py new file mode 100644 index 000000000..815baaed4 --- /dev/null +++ b/mem0/vector_stores/oracledb.py @@ -0,0 +1,592 @@ +"""Oracle AI Vector Search vector store integration for mem0.""" + +import array +import json +import logging +import math +import re +import uuid +from contextlib import contextmanager +from typing import Any, Dict, List, Optional + +try: + import oracledb +except ImportError as exc: # pragma: no cover - dependency guard + raise ImportError("Oracle AI Vector Search requires the 'oracledb' package.") from exc + +from pydantic import BaseModel + +from mem0.configs.vector_stores.oracledb import OracleAIVectorSearchConfig +from mem0.vector_stores.base import VectorStoreBase + +logger = logging.getLogger(__name__) + + +class OutputData(BaseModel): + """Standard output structure returned from vector operations.""" + + id: Optional[str] + score: Optional[float] + payload: Optional[Dict[str, Any]] + + +# Allow letters, digits, underscore, dot, brackets, comma, *, space (for 'to') +METADATA_PATTERN = re.compile(r"[a-zA-Z0-9_\.\[\],\s\*]+") + + +def _validate_metadata_key(metadata_key: str) -> None: + if not METADATA_PATTERN.fullmatch(metadata_key): + raise ValueError( + f"Invalid metadata key '{metadata_key}'. " + "Only letters, numbers, underscores, nesting via '.', " + "and array wildcards '[*]' are allowed." + ) + + +_SCORE_FROM_DISTANCE = { + "COSINE": lambda d: max(0.0, min(1.0, 1.0 - d)), + "EUCLIDEAN": lambda d: 1.0 / (1.0 + max(0.0, d)), + "EUCLIDEAN_SQUARED": lambda d: 1.0 / (1.0 + math.sqrt(max(0.0, d))), + "HAMMING": lambda d: 1.0 / (1.0 + max(0.0, d)), + "MANHATTAN": lambda d: 1.0 / (1.0 + max(0.0, d)), + "DOT": lambda d: -d, +} + + +def _convert_distance_to_score(distance: float, metric: str) -> float: + try: + return _SCORE_FROM_DISTANCE[metric.upper()](distance) + except KeyError: + raise ValueError(f"Unsupported distance metric: {metric}") from None + + +_FIELD_OPERATORS = {"eq", "ne", "gt", "gte", "lt", "lte", "in", "nin", "contains", "icontains"} +_COMPARISON_OPERATORS = { + "eq": "==", + "ne": "!=", + "gt": ">", + "gte": ">=", + "lt": "<", + "lte": "<=", +} +_LOGICAL_OPERATORS = { + "$and": "and", + "$or": "or", + "$not": "not", + "AND": "and", + "OR": "or", + "NOT": "not", +} + + +def _json_path(metadata_key: str) -> str: + _validate_metadata_key(metadata_key) + + path_parts: List[str] = [] + for part in metadata_key.split("."): + if part.endswith("[*]"): + path_parts.append(f'."{part[:-3]}"[*]') + else: + path_parts.append(f'."{part}"') + return "".join(path_parts) + + +def _bind_filter_value(value: Any, params: Dict[str, Any]) -> tuple[str, str]: + param = f"f_{len(params)}" + params[param] = value + return f"${param}", f':{param} AS "{param}"' + + +def _json_exists(json_path: str, predicate: str, passings: List[str]) -> str: + passing_clause = f" PASSING {', '.join(passings)}" if passings else "" + return f"JSON_EXISTS(payload, '${json_path}?({predicate})'{passing_clause})" + + +def _validate_scalar_operand(operator: str, value: Any) -> None: + if isinstance(value, (dict, list, tuple, set)): + raise ValueError(f"Oracle filter operator {operator!r} requires a scalar value") + + +def _build_field_condition(metadata_key: str, value: Any, params: Dict[str, Any]) -> str: + json_path = _json_path(metadata_key) + + if value == "*": + return f"JSON_EXISTS(payload, '${json_path}')" + + if not isinstance(value, dict): + _validate_scalar_operand("eq", value) + if value is None: + return _json_exists(json_path, "@ == null", []) + variable, passing = _bind_filter_value(value, params) + return _json_exists(json_path, f"@ == {variable}", [passing]) + + if not value: + raise ValueError(f"Operator filter for field {metadata_key!r} must not be empty") + + unsupported = set(value) - _FIELD_OPERATORS + if unsupported: + raise ValueError( + f"Unsupported Oracle filter operator(s) for field {metadata_key!r}: " + f"{', '.join(sorted(map(str, unsupported)))}" + ) + + predicates: List[str] = [] + passings: List[str] = [] + additional_clauses: List[str] = [] + + for operator, operand in value.items(): + if operator in _COMPARISON_OPERATORS: + _validate_scalar_operand(operator, operand) + if operand is None: + if operator not in {"eq", "ne"}: + raise ValueError(f"Oracle filter operator {operator!r} does not support null") + predicates.append(f"@ {_COMPARISON_OPERATORS[operator]} null") + continue + variable, passing = _bind_filter_value(operand, params) + predicates.append(f"@ {_COMPARISON_OPERATORS[operator]} {variable}") + passings.append(passing) + continue + + if operator in {"in", "nin"}: + if not isinstance(operand, (list, tuple)) or not operand: + raise ValueError(f"Oracle filter operator {operator!r} requires a non-empty list") + + variables: List[str] = [] + list_passings: List[str] = [] + for item in operand: + _validate_scalar_operand(operator, item) + if item is None: + variables.append("null") + continue + variable, passing = _bind_filter_value(item, params) + variables.append(variable) + list_passings.append(passing) + + membership = _json_exists(json_path, f"@ in ({', '.join(variables)})", list_passings) + if operator == "in": + additional_clauses.append(membership) + else: + additional_clauses.append(f"NOT ({membership})") + continue + + if not isinstance(operand, str): + raise ValueError(f"Oracle filter operator {operator!r} requires a string value") + + if operator == "contains": + variable, passing = _bind_filter_value(operand, params) + predicates.append(f"@ has substring {variable}") + passings.append(passing) + else: + variable, passing = _bind_filter_value(operand.lower(), params) + predicates.append(f"@.lower() has substring {variable}") + passings.append(passing) + + clauses = list(additional_clauses) + if predicates: + clauses.insert(0, _json_exists(json_path, " && ".join(predicates), passings)) + + if len(clauses) == 1: + return clauses[0] + return "(" + " AND ".join(clauses) + ")" + + +def _build_filter_group(filters: Dict[str, Any], params: Dict[str, Any]) -> str: + if not isinstance(filters, dict) or not filters: + raise ValueError("Oracle filter groups must be non-empty dictionaries") + + clauses: List[str] = [] + for key, value in filters.items(): + if key in _LOGICAL_OPERATORS: + if not isinstance(value, list) or not value: + raise ValueError(f"Logical filter operator {key!r} requires a non-empty list") + + nested = [_build_filter_group(condition, params) for condition in value] + logical_operator = _LOGICAL_OPERATORS[key] + if logical_operator == "not": + clauses.append(f"NOT ({' OR '.join(nested)})") + else: + joiner = " AND " if logical_operator == "and" else " OR " + clauses.append("(" + joiner.join(nested) + ")") + continue + + if key.startswith("$"): + raise ValueError(f"Unsupported Oracle logical filter operator: {key}") + + clauses.append(_build_field_condition(key, value, params)) + + if len(clauses) == 1: + return clauses[0] + return "(" + " AND ".join(clauses) + ")" + + +class OracleAIVectorSearch(VectorStoreBase): + """Oracle AI Vector Search backend for mem0.""" + + def __init__(self, **kwargs: Any) -> None: + self.config = OracleAIVectorSearchConfig(**kwargs) + self.collection_name = self.config.collection_name + + if self.config.client: + logger.debug("Using Oracle connection pool: %s", self.config.client) + self.client = self.config.client + self._owns_client = False + elif self.config.use_connection_pool: + pool_kwargs = { + "min": 1, + "max": 4, + } + pool_kwargs.update(self.config.connection_params) + + logger.debug("Creating Oracle connection pool") + self.client = oracledb.create_pool(**pool_kwargs) + self._owns_client = True + else: + logger.debug("Creating Oracle connection") + self.client = oracledb.connect(**self.config.connection_params) + self._owns_client = True + + if not (hasattr(self.client, "thin") and self.client.thin): + if oracledb.clientversion()[:2] < (23, 4): + raise RuntimeError( + f"Oracle DB client driver version {'.'.join(map(str, oracledb.clientversion()))} " + "not supported, must be >=23.4 for vector support" + ) + + if isinstance(self.client, oracledb.Connection): + db_version = tuple([int(v) for v in self.client.version.split(".")]) + else: + with self.client.acquire() as conn: + db_version = tuple([int(v) for v in conn.version.split(".")]) + + if db_version < (23, 4): + raise ValueError( + f"Oracle DB version {'.'.join(map(str, db_version))} not supported, must be >=23.4 for vector support" + ) + + self.create_col() + + @contextmanager + def _get_cursor(self, commit: bool = False): + if isinstance(self.client, oracledb.ConnectionPool): + with self.client.acquire() as connection: + with connection.cursor() as cursor: + try: + yield cursor + if commit: + connection.commit() + except Exception: + connection.rollback() + raise + else: + with self.client.cursor() as cursor: + try: + yield cursor + if commit: + self.client.commit() + except Exception: + self.client.rollback() + raise + + # Utility helpers -------------------------------------------------- + @staticmethod + def _load_payload(value: Any) -> Dict[str, Any]: + if value is None: + return {} + if isinstance(value, dict): + return value + if hasattr(value, "read"): + value = value.read() + if isinstance(value, bytes): + value = value.decode("utf-8") + try: + return json.loads(value) + except json.JSONDecodeError: + logger.debug("Failed to decode payload JSON") + raise + + @staticmethod + def _catalog_name(name: str) -> str: + return name.replace('"', "") + + def _create_index_ddl(self) -> str: + accuracy_str = "" + if self.config.index_accuracy: + accuracy_str = f"WITH TARGET ACCURACY {self.config.index_accuracy}" + + parameters = self._index_parameters() + parameters_str = f"PARAMETERS ({parameters})" if parameters else "" + + distance_metric = self.config.distance_metric + + create_index = ( + f"CREATE VECTOR INDEX IF NOT EXISTS {self.config.index_name} ON {self.collection_name} (vector) " + f"ORGANIZATION {'INMEMORY NEIGHBOR GRAPH' if self.config.index_type == 'HNSW' else 'NEIGHBOR PARTITIONS'}" + f" DISTANCE {distance_metric} {accuracy_str} {parameters_str}" + ) + + return create_index + + def _index_parameters(self) -> str: + index_parameters = self.config.index_parameters + if not index_parameters: + return "" + + parameters = [f"type {self.config.index_type}"] + parameters.extend(f"{key} {value}" for key, value in index_parameters.items()) + + return ", ".join(parameters) + + # Vector store API ------------------------------------------------- + def create_col(self) -> None: + """ + Create a new collection (table in Oracle). + Will also initialize vector search index if specified. + """ + with self._get_cursor(commit=True) as cursor: + cursor.execute( + f""" + CREATE TABLE IF NOT EXISTS {self.collection_name} ( + id VARCHAR2(36) PRIMARY KEY, + vector VECTOR({self.config.embedding_model_dims}), + payload JSON + ) + """ + ) + + if self.config.do_create_index: + ddl = self._create_index_ddl() + cursor.execute(ddl) + + def insert( + self, + vectors: List[List[float]], + payloads: Optional[List[Dict[str, Any]]] = None, + ids: Optional[List[str]] = None, + ) -> None: + logger.info(f"Inserting {len(vectors)} vectors into collection {self.collection_name}") + + if payloads is not None and len(payloads) != len(vectors): + raise ValueError(f"Payload count must match vector count. Expected {len(vectors)} got {len(payloads)}.") + if ids is not None and len(ids) != len(vectors): + raise ValueError(f"ID count must match vector count. Expected {len(vectors)} got {len(ids)}.") + + ids = ids or [str(uuid.uuid4()) for _ in vectors] + data = [ + {"id": _id, "vector": array.array("f", vector), "payload": payload} + for vector, payload, _id in zip(vectors, payloads or [{}] * len(vectors), ids) + ] + + with self._get_cursor(commit=True) as cursor: + cursor.setinputsizes( + vector=oracledb.DB_TYPE_VECTOR, + payload=oracledb.DB_TYPE_JSON, + ) + cursor.executemany( + f"INSERT INTO {self.collection_name} (id, vector, payload) VALUES (:id, :vector, :payload)", data + ) + + def search( + self, + query: str, + vectors: List[float], + top_k: int = 5, + filters: Optional[Dict[str, Any]] = None, + ) -> List[OutputData]: + """ + Search for similar vectors using the vector search index. + + Args: + query (str): Query string + vectors (List[float]): Query vector. + top_k (int, optional): Number of results to return. Defaults to 5. + filters (Dict, optional): Filters to apply to the search. + + Returns: + List[OutputData]: Search results. + """ + filter_clause, params = self._build_filters(filters) + + distance_metric = self.config.distance_metric + + sql = ( + f"SELECT id, payload, VECTOR_DISTANCE(vector, :query_vec, {distance_metric}) distance " + f"FROM {self.collection_name} {filter_clause} ORDER BY distance FETCH APPROX FIRST :limit ROWS ONLY" + ) + + with self._get_cursor() as cursor: + cursor.execute(sql, query_vec=array.array("f", vectors), limit=top_k, **params) + rows = cursor.fetchall() + + return [ + OutputData( + id=row[0], + payload=self._load_payload(row[1]), + score=_convert_distance_to_score(float(row[2]), distance_metric), + ) + for row in rows + ] + + def _build_filters(self, filters: Optional[Dict[str, Any]]) -> tuple[str, Dict[str, Any]]: + if not filters: + return "", {} + + params: Dict[str, Any] = {} + return "WHERE " + _build_filter_group(filters, params), params + + def delete(self, vector_id: str) -> None: + """ + Delete a vector by ID. + + Args: + vector_id (str): ID of the vector to delete. + """ + with self._get_cursor(commit=True) as cursor: + cursor.execute(f"DELETE FROM {self.collection_name} WHERE id = :id", id=vector_id) + + def update( + self, + vector_id: str, + vector: Optional[List[float]] = None, + payload: Optional[Dict[str, Any]] = None, + ) -> 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. + """ + if vector is None and payload is None: + return + + with self._get_cursor(commit=True) as cursor: + sets, params = [], {"vector_id": vector_id} + if vector is not None: + sets.append("vector = :vector") + params["vector"] = array.array("f", vector) + cursor.setinputsizes(vector=oracledb.DB_TYPE_VECTOR) + if payload is not None: + sets.append("payload = :payload") + params["payload"] = payload + cursor.setinputsizes(payload=oracledb.DB_TYPE_JSON) + cursor.execute(f"UPDATE {self.collection_name} SET {', '.join(sets)} WHERE id = :vector_id", params) + + 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. + """ + with self._get_cursor() as cursor: + cursor.execute( + f"SELECT id, payload FROM {self.collection_name} WHERE id = :vector_id", + vector_id=vector_id, + ) + row = cursor.fetchone() + if row is None: + return None + return OutputData(id=row[0], score=None, payload=self._load_payload(row[1])) + + def list_cols(self) -> List[str]: + """ + List all collections. + + Returns: + List[str]: List of collection names. + """ + with self._get_cursor() as cursor: + cursor.execute("SELECT table_name FROM user_tables") + tables = [row[0] for row in cursor.fetchall()] + return tables + + def delete_col(self) -> None: + """Delete a collection.""" + with self._get_cursor(commit=True) as cursor: + cursor.execute(f"DROP TABLE {self.collection_name} PURGE") + + def col_info(self) -> Dict[str, Any]: + """ + Get information about a collection. + + Returns: + Dict[str, Any]: Collection information. + """ + owner, table_name = self._split_collection_name() + + sql = f""" + SELECT + table_name, + (SELECT COUNT(*) FROM {self.collection_name}) AS row_count, + (SELECT + ROUND(SUM(bytes) / 1024 / 1024, 2) || ' MB' + FROM user_segments + WHERE segment_name = :table_name + AND segment_type = 'TABLE' + ) AS total_size + FROM all_tables + WHERE table_name = :table_name + AND owner = NVL(:owner, USER) + """ + + with self._get_cursor() as cursor: + cursor.execute(sql, table_name=table_name, owner=owner) + result = cursor.fetchone() + + if result is None: + raise ValueError(f"Collection {self.collection_name} not found") + + return {"name": result[0], "count": result[1], "size": result[2]} + + def _split_collection_name(self) -> tuple[Optional[str], str]: + """Split the quoted collection name into its optional owner and table parts.""" + segments = re.findall(r'"([^"]+)"', self.collection_name) + if len(segments) > 1: + return segments[-2], segments[-1] + return None, segments[-1] + + def list(self, filters: Optional[Dict[str, Any]] = None, top_k: Optional[int] = 100) -> List[List[OutputData]]: + """ + List all vectors in a collection. + + Args: + filters (Dict, optional): Filters to apply to the list. + top_k (int, optional): Number of vectors to return. Defaults to 100. + + Returns: + List[List[OutputData]]: A single-element list holding the list of vectors. + """ + filter_clause, params = self._build_filters(filters) + + limit_clause = "" + if top_k is not None: + limit_clause = " FETCH FIRST :limit ROWS ONLY" + params["limit"] = top_k + + sql = f"SELECT id, payload FROM {self.collection_name} {filter_clause} {limit_clause}" + + with self._get_cursor() as cursor: + cursor.execute(sql, **params) + rows = cursor.fetchall() + + return [[OutputData(id=row[0], score=None, payload=self._load_payload(row[1])) for row in rows]] + + def reset(self) -> None: + """Reset the index by deleting and recreating it.""" + logger.warning("Resetting collection %s", self.collection_name) + self.delete_col() + self.create_col() + + def __del__(self) -> None: + """ + Close the database connection pool when the object is deleted. + """ + try: + if getattr(self, "_owns_client", False): + self.client.close() + except Exception: + pass diff --git a/pyproject.toml b/pyproject.toml index f22bfb8ec..4c440c5ab 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -52,6 +52,7 @@ vector-stores = [ "elasticsearch>=8.0.0,<9.0.0", "pymilvus>=2.4.0,<2.6.0", "langchain-aws>=0.2.23,<0.3.0", + "oracledb>=2.2.0", ] llms = [ "groq>=0.3.0", diff --git a/tests/vector_stores/test_oracledb.py b/tests/vector_stores/test_oracledb.py new file mode 100644 index 000000000..25ff9e25a --- /dev/null +++ b/tests/vector_stores/test_oracledb.py @@ -0,0 +1,1055 @@ +import os +import uuid +from contextlib import nullcontext +from types import SimpleNamespace +from typing import Any, Dict +from unittest.mock import MagicMock + +import oracledb +import pytest + +from mem0.configs.vector_stores.oracledb import ( + OracleAIVectorSearchConfig, + _quote_identifier, +) +from mem0.vector_stores.oracledb import OracleAIVectorSearch, _convert_distance_to_score + +# Global Oracle connection settings (override via env to run in different environments) +ORACLE_USER = os.environ.get("ORACLE_USER") or "" +ORACLE_PASSWORD = os.environ.get("ORACLE_PASSWORD") or "" +ORACLE_DSN = os.environ.get("ORACLE_DSN") or "" + +requires_oracle_credentials = pytest.mark.skipif( + not (ORACLE_USER and ORACLE_DSN), + reason="Oracle credentials not configured", +) + +DIM = 128 + + +def _unique_collection_name() -> str: + # Keep under Oracle's 30-char identifier limit + return f"TEST_MEM0_{uuid.uuid4().hex[:8]}" + + +# Representative coverage of the old matrix. Every option value from the previous +# grid appears in at least one case, without creating a 1924-test DDL-heavy suite. +REPRESENTATIVE_CASES = [ + { + "name": "params-cosine-hnsw-default-noacc-noparams", + "use_connection_pool": False, + "distance_metric": "COSINE", + "index_type": "HNSW", + "custom_index_name": False, + "index_accuracy": None, + "index_parameters": False, + }, + { + "name": "params-euclidean-ivf-custom-acc90-params", + "use_connection_pool": False, + "distance_metric": "EUCLIDEAN", + "index_type": "IVF", + "custom_index_name": True, + "index_accuracy": 90, + "index_parameters": True, + }, + { + "name": "pool-cosine-ivf-default-acc90-noparams", + "use_connection_pool": True, + "distance_metric": "COSINE", + "index_type": "IVF", + "custom_index_name": False, + "index_accuracy": 90, + "index_parameters": False, + }, + { + "name": "pool-euclidean-hnsw-custom-noacc-params", + "use_connection_pool": True, + "distance_metric": "EUCLIDEAN", + "index_type": "HNSW", + "custom_index_name": True, + "index_accuracy": None, + "index_parameters": True, + }, +] + + +def _build_oracle_db(case: Dict[str, Any], *, do_create_index: bool) -> OracleAIVectorSearch: + collection_name = _unique_collection_name() + conn_params = {"user": ORACLE_USER, "password": ORACLE_PASSWORD, "dsn": ORACLE_DSN} + config_kwargs: Dict[str, Any] = { + "collection_name": collection_name, + "embedding_model_dims": DIM, + "distance_metric": case["distance_metric"], + "index_type": case["index_type"], + "do_create_index": do_create_index, + "use_connection_pool": case["use_connection_pool"], + } + + if case.get("custom_index_name"): + config_kwargs["index_name"] = f"{collection_name}_IDX" + if case.get("index_accuracy") is not None: + config_kwargs["index_accuracy"] = case["index_accuracy"] + if case.get("index_parameters"): + config_kwargs["index_parameters"] = ( + {"neighbors": 40, "efconstruction": 64} if case["index_type"] == "HNSW" else {"neighbor partitions": 10} + ) + + if case.get("use_connection_pool"): + config_kwargs["client"] = oracledb.create_pool(min=1, max=4, **conn_params) + else: + config_kwargs["connection_params"] = conn_params + + return OracleAIVectorSearch(**config_kwargs) + + +@pytest.fixture( + params=[REPRESENTATIVE_CASES[0]], + ids=lambda p: p["name"], +) +def oracle_db(request): + """ + Stable Oracle fixture for CRUD/search/list behavior. + Uses a single representative config and skips vector-index creation to avoid + repeated DDL lock contention on the shared Oracle instance. + """ + if not (ORACLE_USER and ORACLE_DSN): + pytest.skip("Oracle credentials not configured") + + db = _build_oracle_db(request.param, do_create_index=False) + + try: + yield db + finally: + try: + db.delete_col() + except Exception: + # Ignore failures (e.g., already dropped) + pass + + +@requires_oracle_credentials +@pytest.mark.parametrize("case", REPRESENTATIVE_CASES, ids=lambda case: case["name"]) +def test_initialize_create_col(case: Dict[str, Any]): + oracle_db = _build_oracle_db(case, do_create_index=False) + + try: + # Verify config normalization and DDL generation for each representative case + collection_name = oracle_db.collection_name.strip('"') + expected_index_name = ( + f'"{collection_name}_IDX"' if case["custom_index_name"] else f'"{collection_name}_VEC_IDX"' + ) + assert oracle_db.config.embedding_model_dims == DIM + assert oracle_db.config.distance_metric in ("COSINE", "EUCLIDEAN") + assert oracle_db.config.index_type in ("HNSW", "IVF") + assert oracle_db.config.index_name == expected_index_name + assert oracle_db.config.index_accuracy == case["index_accuracy"] + assert bool(case["index_parameters"]) == bool(oracle_db.config.index_parameters) + ddl = oracle_db._create_index_ddl() + assert oracle_db.config.index_name in ddl + assert oracle_db.collection_name in ddl + if case["index_type"] == "HNSW": + assert "INMEMORY NEIGHBOR GRAPH" in ddl + else: + assert "NEIGHBOR PARTITIONS" in ddl + if case["index_accuracy"] is not None: + assert f"WITH TARGET ACCURACY {case['index_accuracy']}" in ddl + if case["index_parameters"]: + assert "PARAMETERS (" in ddl + assert f"type {case['index_type']}" in ddl + else: + assert "PARAMETERS (" not in ddl + + tables = oracle_db.list_cols() + target = oracle_db.collection_name.strip('"').upper() + assert target in [t.upper() for t in tables] + finally: + try: + oracle_db.delete_col() + except Exception: + pass + + +@requires_oracle_credentials +def test_create_col_with_index_smoke(): + case = REPRESENTATIVE_CASES[1] + oracle_db = _build_oracle_db(case, do_create_index=True) + + try: + tables = oracle_db.list_cols() + target = oracle_db.collection_name.strip('"').upper() + assert target in [t.upper() for t in tables] + finally: + try: + oracle_db.delete_col() + except Exception: + pass + + +@requires_oracle_credentials +def test_index_parameters_are_structured_and_allowlisted(): + conn_params = {"user": ORACLE_USER, "password": ORACLE_PASSWORD, "dsn": ORACLE_DSN} + collection_name = _unique_collection_name() + oracle_db = OracleAIVectorSearch( + collection_name=collection_name, + embedding_model_dims=DIM, + connection_params=conn_params, + do_create_index=False, + index_type="HNSW", + index_parameters={"neighbors": 40, "efconstruction": 64}, + ) + + try: + ddl = oracle_db._create_index_ddl() + assert "PARAMETERS (type HNSW, neighbors 40, efconstruction 64)" in ddl + finally: + oracle_db.delete_col() + + +@requires_oracle_credentials +def test_ivf_index_parameters_are_structured_and_allowlisted(): + conn_params = {"user": ORACLE_USER, "password": ORACLE_PASSWORD, "dsn": ORACLE_DSN} + collection_name = _unique_collection_name() + oracle_db = OracleAIVectorSearch( + collection_name=collection_name, + embedding_model_dims=DIM, + connection_params=conn_params, + do_create_index=False, + index_type="IVF", + index_parameters={ + "neighbor partitions": 10, + "samples_per_partition": 4, + "min_vectors_per_partition": 2, + }, + ) + + try: + ddl = oracle_db._create_index_ddl() + assert ( + "PARAMETERS (type IVF, neighbor partitions 10, samples_per_partition 4, min_vectors_per_partition 2)" + ) in ddl + finally: + oracle_db.delete_col() + + +def test_index_parameters_reject_unsupported_fragments(): + conn_params = {"user": ORACLE_USER, "password": ORACLE_PASSWORD, "dsn": ORACLE_DSN} + + with pytest.raises(ValueError, match="Extra inputs are not permitted"): + OracleAIVectorSearch( + collection_name=_unique_collection_name(), + embedding_model_dims=DIM, + connection_params=conn_params, + do_create_index=False, + index_type="HNSW", + index_parameters={"parallel": "8 NOLOGGING"}, + ) + + with pytest.raises(ValueError, match="Input should be a valid integer"): + OracleAIVectorSearch( + collection_name=_unique_collection_name(), + embedding_model_dims=DIM, + connection_params=conn_params, + do_create_index=False, + index_type="IVF", + index_parameters={"neighbor partitions": "10) PARALLEL 8"}, + ) + + +def test_index_parameters_reject_non_string_keys(): + with pytest.raises(ValueError, match="Keys should be strings"): + OracleAIVectorSearchConfig( + collection_name=_unique_collection_name(), + embedding_model_dims=DIM, + client=object(), + index_type="HNSW", + index_parameters={1: 10}, + ) + + +def test_index_parameters_canonicalize_int_subclasses(): + class FormattedInt(int): + def __format__(self, format_spec): + return "40) PARALLEL 8" + + config = OracleAIVectorSearchConfig( + collection_name=_unique_collection_name(), + embedding_model_dims=DIM, + client=object(), + index_type="HNSW", + index_parameters={"neighbors": FormattedInt(40)}, + ) + oracle_db = object.__new__(OracleAIVectorSearch) + oracle_db.config = config + oracle_db.collection_name = config.collection_name + + ddl = oracle_db._create_index_ddl() + assert type(config.index_parameters["neighbors"]) is int + assert "PARALLEL 8" not in ddl + assert "PARAMETERS (type HNSW, neighbors 40)" in ddl + + +def test_ivf_index_parameters_use_oracle_ddl_names(): + config = OracleAIVectorSearchConfig( + collection_name=_unique_collection_name(), + embedding_model_dims=DIM, + client=object(), + index_type="ivf", + index_parameters={ + "neighbor partitions": 10, + "samples_per_partition": 4, + "min_vectors_per_partition": 2, + }, + ) + oracle_db = object.__new__(OracleAIVectorSearch) + oracle_db.config = config + oracle_db.collection_name = config.collection_name + + assert config.index_type == "IVF" + assert config.index_parameters == { + "neighbor partitions": 10, + "samples_per_partition": 4, + "min_vectors_per_partition": 2, + } + assert ( + oracle_db._index_parameters() + == "type IVF, neighbor partitions 10, samples_per_partition 4, min_vectors_per_partition 2" + ) + + +@pytest.mark.parametrize( + ("field", "value"), + [ + ("distance_metric", None), + ("index_type", None), + ("use_connection_pool", None), + ("collection_name", None), + ], +) +def test_config_rejects_none_for_non_optional_fields(field, value): + with pytest.raises(ValueError): + OracleAIVectorSearchConfig(client=object(), **{field: value}) + + +@pytest.mark.parametrize( + ("metric", "distance", "expected_score"), + [ + ("COSINE", 0.25, 0.75), + ("cosine", -0.01, 1.0), + ("COSINE", 1.25, 0.0), + ("EUCLIDEAN", 3.0, 0.25), + ("EUCLIDEAN", -0.01, 1.0), + ("EUCLIDEAN_SQUARED", 9.0, 0.25), + ("HAMMING", 3.0, 0.25), + ("MANHATTAN", 3.0, 0.25), + ("DOT", -0.75, 0.75), + ("DOT", 0.25, -0.25), + ], +) +def test_convert_distance_to_score(metric, distance, expected_score): + assert _convert_distance_to_score(distance, metric) == pytest.approx(expected_score) + + +def test_convert_distance_to_score_rejects_unknown_metric(): + with pytest.raises(ValueError, match="Unsupported distance metric: UNKNOWN"): + _convert_distance_to_score(0.5, "UNKNOWN") + + +def test_search_and_list_follow_base_contract(): + search_cursor = MagicMock() + search_cursor.fetchall.return_value = [ + ("close", '{"label": "close"}', 0.1), + ("far", '{"label": "far"}', 0.8), + ] + list_cursor = MagicMock() + list_cursor.fetchall.return_value = [ + ("listed", '{"name": "listed"}'), + ] + + store = object.__new__(OracleAIVectorSearch) + store.collection_name = '"MEM0"' + store.config = SimpleNamespace(distance_metric="COSINE") + store._get_cursor = MagicMock( + side_effect=[ + nullcontext(search_cursor), + nullcontext(list_cursor), + ] + ) + + search_results = store.search( + query="unused", + vectors=[1.0, 0.0], + top_k=2, + filters={"score": {"gte": 5}}, + ) + list_results = store.list(top_k=2) + + assert [result.score for result in search_results] == pytest.approx([0.9, 0.2]) + assert search_results[0].score > search_results[1].score + search_sql = search_cursor.execute.call_args.args[0] + assert "@ >= $f_0" in search_sql + assert search_cursor.execute.call_args.kwargs["f_0"] == 5 + assert isinstance(list_results[0], list) + assert list_results[0][0].payload["name"] == "listed" + + +def test_build_filters_wildcard_requires_field_existence(): + store = object.__new__(OracleAIVectorSearch) + + clause, params = store._build_filters({"run_id": "*"}) + + assert clause == """WHERE JSON_EXISTS(payload, '$."run_id"')""" + assert params == {} + + +def test_build_filters_rejects_empty_metadata_key(): + store = object.__new__(OracleAIVectorSearch) + + with pytest.raises(ValueError, match="Invalid metadata key"): + store._build_filters({"": "alice"}) + + +@pytest.mark.parametrize( + ("collection_name", "expected"), + [ + ("MEM0", (None, "MEM0")), + ("SCHEMA.MEM0", ("SCHEMA", "MEM0")), + ('"my.table"', (None, "my.table")), + ], +) +def test_split_collection_name(collection_name, expected): + store = object.__new__(OracleAIVectorSearch) + store.collection_name = _quote_identifier(collection_name) + + assert store._split_collection_name() == expected + + +def test_col_info_looks_up_the_unqualified_table_name(): + cursor = MagicMock() + cursor.fetchone.return_value = ("MEM0", 7, "1.5 MB") + + store = object.__new__(OracleAIVectorSearch) + store.collection_name = _quote_identifier("SCHEMA.MEM0") + store._get_cursor = MagicMock(return_value=nullcontext(cursor)) + + info = store.col_info() + + assert info == {"name": "MEM0", "count": 7, "size": "1.5 MB"} + assert cursor.execute.call_args.kwargs == {"table_name": "MEM0", "owner": "SCHEMA"} + + +def test_col_info_raises_when_collection_is_missing(): + cursor = MagicMock() + cursor.fetchone.return_value = None + + store = object.__new__(OracleAIVectorSearch) + store.collection_name = _quote_identifier("MEM0") + store._get_cursor = MagicMock(return_value=nullcontext(cursor)) + + with pytest.raises(ValueError, match="not found"): + store.col_info() + + +def test_build_filters_combines_wildcard_and_scalar_equality(): + store = object.__new__(OracleAIVectorSearch) + + clause, params = store._build_filters({"user_id": "alice", "run_id": "*"}) + + assert """JSON_EXISTS(payload, '$."user_id"?(@ == $f_0)' PASSING :f_0 AS "f_0")""" in clause + assert """JSON_EXISTS(payload, '$."run_id"')""" in clause + assert " AND " in clause + assert params == {"f_0": "alice"} + + +@pytest.mark.parametrize( + ("operator", "predicate"), + [ + ("eq", "@ == $f_0"), + ("ne", "@ != $f_0"), + ("gt", "@ > $f_0"), + ("gte", "@ >= $f_0"), + ("lt", "@ < $f_0"), + ("lte", "@ <= $f_0"), + ], +) +def test_build_filters_comparison_operators(operator, predicate): + store = object.__new__(OracleAIVectorSearch) + + clause, params = store._build_filters({"score": {operator: 5}}) + + assert predicate in clause + assert params == {"f_0": 5} + + +def test_build_filters_combines_comparisons_for_same_field(): + store = object.__new__(OracleAIVectorSearch) + + clause, params = store._build_filters({"score": {"gte": 5, "lt": 10}}) + + assert '$."score"?(@ >= $f_0 && @ < $f_1)' in clause + assert params == {"f_0": 5, "f_1": 10} + + +@pytest.mark.parametrize( + ("operator", "expected"), + [ + ("in", "JSON_EXISTS"), + ("nin", "NOT (JSON_EXISTS"), + ], +) +def test_build_filters_membership_operators_expand_binds(operator, expected): + store = object.__new__(OracleAIVectorSearch) + + clause, params = store._build_filters({"category": {operator: ["work", "personal"]}}) + + assert expected in clause + assert "@ in ($f_0, $f_1)" in clause + assert params == {"f_0": "work", "f_1": "personal"} + + +def test_build_filters_string_operators(): + store = object.__new__(OracleAIVectorSearch) + + contains_clause, contains_params = store._build_filters({"title": {"contains": "Meeting"}}) + icontains_clause, icontains_params = store._build_filters({"title": {"icontains": "Meet.ing"}}) + + assert "@ has substring $f_0" in contains_clause + assert contains_params == {"f_0": "Meeting"} + assert "@.lower() has substring $f_0" in icontains_clause + assert icontains_params == {"f_0": "meet.ing"} + + +def test_build_filters_operator_eq_treats_asterisk_as_literal(): + store = object.__new__(OracleAIVectorSearch) + + clause, params = store._build_filters({"status": {"eq": "*"}}) + + assert "@ == $f_0" in clause + assert params == {"f_0": "*"} + + +@pytest.mark.parametrize( + ("filters", "predicate", "params"), + [ + ({"nullable": None}, "@ == null", {}), + ({"nullable": {"eq": None}}, "@ == null", {}), + ({"nullable": {"ne": None}}, "@ != null", {}), + ({"nullable": {"in": [None, "set"]}}, "@ in (null, $f_0)", {"f_0": "set"}), + ], +) +def test_build_filters_supports_json_null(filters, predicate, params): + store = object.__new__(OracleAIVectorSearch) + + clause, actual_params = store._build_filters(filters) + + assert predicate in clause + assert actual_params == params + + +def test_build_filters_nested_logical_operators_share_bind_namespace(): + store = object.__new__(OracleAIVectorSearch) + filters = { + "user_id": "alice", + "$or": [ + {"score": {"gte": 5}}, + { + "$and": [ + {"status": {"eq": "active"}}, + {"category": "work"}, + ] + }, + ], + "$not": [{"archived": {"eq": "yes"}}], + } + + clause, params = store._build_filters(filters) + + assert " OR " in clause + assert " AND " in clause + assert "NOT (" in clause + assert params == { + "f_0": "alice", + "f_1": 5, + "f_2": "active", + "f_3": "work", + "f_4": "yes", + } + + +def test_build_filters_accepts_unprocessed_logical_operator_names(): + store = object.__new__(OracleAIVectorSearch) + + clause, params = store._build_filters( + { + "AND": [ + {"score": {"gte": 5}}, + {"OR": [{"category": "work"}, {"category": "personal"}]}, + ] + } + ) + + assert " AND " in clause + assert " OR " in clause + assert params == {"f_0": 5, "f_1": "work", "f_2": "personal"} + + +@pytest.mark.parametrize( + ("filters", "message"), + [ + ({"score": {"between": [1, 2]}}, "Unsupported Oracle filter operator"), + ({"score": {}}, "must not be empty"), + ({"score": {"in": []}}, "requires a non-empty list"), + ({"title": {"contains": 5}}, "requires a string value"), + ({"score": {"gte": [5]}}, "requires a scalar value"), + ({"score": {"gte": None}}, "does not support null"), + ({"$xor": [{"score": 5}]}, "Unsupported Oracle logical filter operator"), + ({"$or": []}, "requires a non-empty list"), + ({"bad-key": "value"}, "Invalid metadata key"), + ], +) +def test_build_filters_rejects_invalid_filter_shapes(filters, message): + store = object.__new__(OracleAIVectorSearch) + + with pytest.raises(ValueError, match=message): + store._build_filters(filters) + + +def test_insert_and_get(oracle_db: OracleAIVectorSearch): + vectors = [[0.1] * DIM, [0.2] * DIM] + payloads = [{"name": "vector1"}, {"name": "vector2"}] + + oracle_db.insert(vectors, payloads=payloads) + + listed = oracle_db.list(top_k=10)[0] + assert len(listed) >= 2 + seen_names = {item.payload.get("name") for item in listed} + assert {"vector1", "vector2"}.issubset(seen_names) + + # Fetch one by id (Oracle RAW(16) id is generated by DB) + some_id = listed[0].id + got = oracle_db.get(vector_id=some_id) + assert got is not None + assert got.id == some_id + assert isinstance(got.payload, dict) + + +def test_search(oracle_db: OracleAIVectorSearch): + # Create predictable geometry; works for COSINE or EUCLIDEAN + pos_vec = [1.0] * DIM + neg_vec = [-1.0] * DIM + mid_vec = [1.0 if i % 2 == 0 else 0.0 for i in range(DIM)] + payloads = [ + {"name": "pos", "user_id": "u1"}, + {"name": "neg", "user_id": "u2"}, + {"name": "mid", "user_id": "u3"}, + ] + oracle_db.insert([pos_vec, neg_vec, mid_vec], payloads=payloads) + + results = oracle_db.search("unused", vectors=pos_vec, top_k=3) + assert isinstance(results, list) + assert len(results) >= 1 + + names = {r.payload.get("name") for r in results} + assert "pos" in names # closest to query + + +def test_search_with_filters(oracle_db: OracleAIVectorSearch): + vec = [0.5] * DIM + payloads = [ + {"name": "a", "user_id": "alice", "agent_id": "agent1", "run_id": "run1"}, + {"name": "b", "user_id": "bob", "agent_id": "agent2", "run_id": "run2"}, + ] + oracle_db.insert([vec, vec], payloads=payloads) + + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + results = oracle_db.search("unused", vectors=vec, top_k=5, filters=filters) + + assert len(results) >= 1 + for r in results: + assert r.payload.get("user_id") == "alice" + assert r.payload.get("agent_id") == "agent1" + assert r.payload.get("run_id") == "run1" + + +def test_search_with_single_filter(oracle_db: OracleAIVectorSearch): + vec = [0.7] * DIM + payloads = [ + {"name": "x", "user_id": "alice"}, + {"name": "y", "user_id": "bob"}, + ] + oracle_db.insert([vec, vec], payloads=payloads) + + results = oracle_db.search("unused", vectors=vec, top_k=5, filters={"user_id": "alice"}) + assert len(results) >= 1 + for r in results: + assert r.payload.get("user_id") == "alice" + + +def test_search_with_no_filters(oracle_db: OracleAIVectorSearch): + vec = [0.33] * DIM + oracle_db.insert([vec], payloads=[{"k": "v"}]) + + results = oracle_db.search("unused", vectors=vec, top_k=1, filters=None) + assert len(results) == 1 + + +def test_extended_filtering(oracle_db: OracleAIVectorSearch): + vector = [0.42] * DIM + oracle_db.insert( + [vector] * 4, + payloads=[ + { + "name": "Alpha Meeting", + "score": 10, + "category": "work", + "status": "active", + "run_id": "r1", + "enabled": True, + "nullable": None, + "ratio": 1.25, + "created_at": "2025-01-15", + "profile": {"department": "Engineering", "skills": ["Python", "SQL"]}, + }, + { + "name": "beta meeting", + "score": 5, + "category": "personal", + "status": "inactive", + "enabled": False, + "nullable": "set", + "ratio": 2.5, + "created_at": "2024-12-31", + "profile": {"department": "Engineering", "skills": ["Java"]}, + }, + { + "name": "Gamma", + "score": 20, + "category": "work", + "status": "active", + "run_id": "r3", + "enabled": True, + "ratio": 3.75, + "created_at": "2025-06-01", + "profile": {"department": "Sales", "skills": ["Python"]}, + }, + { + "name": "Literal", + "score": 12, + "category": "other", + "status": "*", + "enabled": False, + "nullable": "value", + "ratio": 4.0, + "created_at": "2026-01-01", + "profile": {"department": "Support", "skills": []}, + }, + ], + ids=["alpha", "beta", "gamma", "literal"], + ) + + def matching_names(filters): + results = oracle_db.list(filters=filters, top_k=10)[0] + return {result.payload["name"] for result in results} + + assert matching_names({"score": {"gte": 6, "lt": 20}}) == {"Alpha Meeting", "Literal"} + assert matching_names({"category": {"eq": "work"}}) == {"Alpha Meeting", "Gamma"} + assert matching_names({"category": {"ne": "work"}}) == {"beta meeting", "Literal"} + assert matching_names({"score": {"lte": 10}}) == {"Alpha Meeting", "beta meeting"} + assert matching_names({"category": {"in": ["work", "personal"]}}) == { + "Alpha Meeting", + "beta meeting", + "Gamma", + } + assert matching_names({"category": {"nin": ["work", "personal"]}}) == {"Literal"} + assert matching_names({"name": {"contains": "Meeting"}}) == {"Alpha Meeting"} + assert matching_names({"name": {"icontains": "meeting"}}) == {"Alpha Meeting", "beta meeting"} + assert matching_names({"run_id": "*"}) == {"Alpha Meeting", "Gamma"} + assert matching_names({"status": {"eq": "*"}}) == {"Literal"} + assert matching_names({"profile.department": {"eq": "Engineering"}}) == { + "Alpha Meeting", + "beta meeting", + } + assert matching_names({"profile.skills[*]": {"eq": "Python"}}) == {"Alpha Meeting", "Gamma"} + assert matching_names({"enabled": {"eq": True}}) == {"Alpha Meeting", "Gamma"} + assert matching_names({"nullable": {"eq": None}}) == {"Alpha Meeting"} + assert matching_names({"nullable": {"ne": None}}) == {"beta meeting", "Literal"} + assert matching_names({"nullable": {"in": [None, "set"]}}) == {"Alpha Meeting", "beta meeting"} + assert matching_names({"created_at": {"gte": "2025-01-01", "lt": "2026-01-01"}}) == { + "Alpha Meeting", + "Gamma", + } + assert matching_names({"ratio": {"gt": 1.25, "lte": 3.75}}) == {"beta meeting", "Gamma"} + assert matching_names({"score": {"gte": 10, "in": [10, 12]}}) == {"Alpha Meeting", "Literal"} + assert matching_names( + { + "$or": [ + {"score": {"lt": 6}}, + {"score": {"gt": 15}}, + ] + } + ) == {"beta meeting", "Gamma"} + assert matching_names({"$not": [{"category": {"eq": "personal"}}]}) == { + "Alpha Meeting", + "Gamma", + "Literal", + } + assert matching_names( + { + "AND": [ + {"score": {"gte": 10}}, + { + "OR": [ + {"category": "work"}, + {"status": {"eq": "*"}}, + ] + }, + ] + } + ) == {"Alpha Meeting", "Gamma", "Literal"} + assert matching_names( + { + "AND": [ + {"enabled": {"eq": True}}, + { + "OR": [ + {"profile.department": {"eq": "Engineering"}}, + { + "AND": [ + {"score": {"gt": 15}}, + {"NOT": [{"category": {"eq": "personal"}}]}, + ] + }, + ] + }, + ] + } + ) == {"Alpha Meeting", "Gamma"} + + +def test_delete(oracle_db: OracleAIVectorSearch): + vec = [0.9] * DIM + oracle_db.insert([vec], payloads=[{"name": "to_delete"}]) + + listed = oracle_db.list(top_k=10)[0] + assert len(listed) >= 1 + target_id = listed[0].id + + oracle_db.delete(vector_id=target_id) + got = oracle_db.get(vector_id=target_id) + assert got is None + + +def test_reset_recreates_empty_usable_collection(oracle_db: OracleAIVectorSearch): + vector = [0.15] * DIM + oracle_db.insert([vector], ids=["before-reset"], payloads=[{"name": "before"}]) + assert oracle_db.get("before-reset") is not None + + oracle_db.reset() + + assert oracle_db.get("before-reset") is None + assert oracle_db.list(top_k=10) == [[]] + + oracle_db.insert([vector], ids=["after-reset"], payloads=[{"name": "after"}]) + result = oracle_db.get("after-reset") + assert result is not None + assert result.payload["name"] == "after" + + +def test_update(oracle_db: OracleAIVectorSearch): + vec = [0.01] * DIM + oracle_db.insert([vec], payloads=[{"name": "old"}]) + + listed = oracle_db.list(top_k=10)[0] + assert len(listed) >= 1 + target_id = listed[0].id + + updated_vec = [0.02] * DIM + updated_payload = {"name": "new"} + oracle_db.update(vector_id=target_id, vector=updated_vec, payload=updated_payload) + + got = oracle_db.get(vector_id=target_id) + assert got is not None + assert got.payload.get("name") == "new" + + +def test_list_cols(oracle_db: OracleAIVectorSearch): + tables = oracle_db.list_cols() + target = oracle_db.collection_name.strip('"').upper() + assert target in [t.upper() for t in tables] + + +def test_delete_col_isolated(oracle_db: OracleAIVectorSearch): + # Use a separate, isolated collection to test drop; reuse current fixture's metric/index options + collection_name = _unique_collection_name() + cfg: Dict[str, Any] = { + "collection_name": collection_name, + "embedding_model_dims": DIM, + "distance_metric": oracle_db.config.distance_metric, + "index_type": oracle_db.config.index_type, + "do_create_index": False, + "connection_params": {"user": ORACLE_USER, "password": ORACLE_PASSWORD, "dsn": ORACLE_DSN}, + } + + # If the fixture used a pool object, pass it as well. + if getattr(oracle_db.config, "client", None): + cfg["client"] = oracle_db.config.client + + tmp_db = OracleAIVectorSearch(**cfg) + + tgt = tmp_db.collection_name.strip('"').upper() + tables_before = [t.upper() for t in tmp_db.list_cols()] + assert tgt in tables_before + + tmp_db.delete_col() + + tables_after = [t.upper() for t in tmp_db.list_cols()] + assert tgt not in tables_after + + +def test_col_info(oracle_db: OracleAIVectorSearch): + info = oracle_db.col_info() + # Structure sanity checks; exact values depend on DB state + assert isinstance(info, dict) + assert "name" in info and "count" in info and "size" in info + + +def test_list(oracle_db: OracleAIVectorSearch): + v1, v2 = [0.11] * DIM, [0.22] * DIM + oracle_db.insert([v1, v2], payloads=[{"key": "value1"}, {"key": "value2"}]) + + results = oracle_db.list(top_k=2) + assert isinstance(results[0], list) + listed = results[0] + assert len(listed) <= 2 + # Both inserted might be returned if table had no prior rows + if len(listed) == 2: + payloads = [r.payload for r in listed] + keys = {p.get("key") for p in payloads} + assert keys.issubset({"value1", "value2"}) + + +def test_list_with_filters(oracle_db: OracleAIVectorSearch): + v = [0.44] * DIM + oracle_db.insert( + [v, v], + payloads=[ + {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"}, + {"user_id": "bob", "agent_id": "agent2", "run_id": "run2"}, + ], + ) + + filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"} + results = oracle_db.list(filters=filters, top_k=10)[0] + assert len(results) >= 1 + for r in results: + assert r.payload.get("user_id") == "alice" + assert r.payload.get("agent_id") == "agent1" + assert r.payload.get("run_id") == "run1" + + +def test_list_with_single_filter(oracle_db: OracleAIVectorSearch): + v = [0.55] * DIM + oracle_db.insert( + [v, v], + payloads=[ + {"user_id": "alice"}, + {"user_id": "bob"}, + ], + ) + + results = oracle_db.list(filters={"user_id": "alice"}, top_k=10)[0] + assert len(results) >= 1 + for r in results: + assert r.payload.get("user_id") == "alice" + + +def test_list_with_no_filters(oracle_db: OracleAIVectorSearch): + v = [0.66] * DIM + oracle_db.insert([v], payloads=[{"k": "v"}]) + + results = oracle_db.list(filters=None, top_k=10)[0] + assert len(results) >= 1 + + +def test_list_returns_nested_output(oracle_db: OracleAIVectorSearch): + oracle_db.insert([[0.12] * DIM], payloads=[{"name": "nested"}]) + + results = oracle_db.list(top_k=10) + + assert isinstance(results, list) + assert results + assert isinstance(results[0], list) + assert results[0][0].payload["name"] == "nested" + + +def test_update_accepts_empty_payload(oracle_db: OracleAIVectorSearch): + oracle_db.insert([[0.21] * DIM], payloads=[{"name": "before"}], ids=["row-1"]) + + oracle_db.update("row-1", payload={}) + + result = oracle_db.get("row-1") + assert result is not None + assert result.payload == {} + + +@requires_oracle_credentials +def test_does_not_close_caller_supplied_pool(): + pool = oracledb.create_pool( + min=1, + max=2, + user=ORACLE_USER, + password=ORACLE_PASSWORD, + dsn=ORACLE_DSN, + ) + db = OracleAIVectorSearch( + collection_name=_unique_collection_name(), + embedding_model_dims=4, + do_create_index=False, + client=pool, + ) + + try: + db.__del__() + with pool.acquire() as conn: + with conn.cursor() as cursor: + cursor.execute("SELECT 1 FROM dual") + assert cursor.fetchone()[0] == 1 + finally: + try: + db.delete_col() + finally: + pool.close() + + +@requires_oracle_credentials +def test_documentation(): + from mem0 import Memory + + if not os.environ.get("OPENAI_API_KEY"): + pytest.skip("OPENAI_API_KEY is required for the end-to-end documentation test") + + config = { + "vector_store": { + "provider": "oracledb", + "config": { + "connection_params": {"user": ORACLE_USER, "password": ORACLE_PASSWORD, "dsn": ORACLE_DSN}, + "do_create_index": False, + }, + }, + } + + 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"}) + results = m.search("What movie to watch?", user_id="alice", limit=2)["results"] + assert len(results) == 2 + assert all(res["user_id"] == "alice" for res in results) + m.reset()