Add Oracle Vector Store Integration (#5358)

Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
Elif Sema Balcioglu
2026-07-23 19:00:01 +03:00
committed by GitHub
parent c2150e8f1a
commit d6d89c987b
12 changed files with 1908 additions and 2 deletions
+6 -1
View File
@@ -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 |
</Tab>
<Tab title="TypeScript">
| Parameter | Description |
+134
View File
@@ -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
<CodeGroup>
```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"})
```
</CodeGroup>
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 | `<collection_name>_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 <n>` | `None` |
<Note>
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.
</Note>
### 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}},
)
```
+2 -1
View File
@@ -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.
<Card title="Milvus" icon="/images/provider-icons/milvus.svg" href="/components/vectordbs/dbs/milvus"></Card>
<Card title="Pinecone" icon="/images/provider-icons/pinecone.svg" href="/components/vectordbs/dbs/pinecone"></Card>
<Card title="MongoDB" icon="/images/provider-icons/mongodb.svg" href="/components/vectordbs/dbs/mongodb"></Card>
<Card title="Oracle AI Vector Search" icon="/images/provider-icons/oracle.svg" href="/components/vectordbs/dbs/oracledb"></Card>
<Card title="Azure" icon="/images/provider-icons/azure-color.svg" href="/components/vectordbs/dbs/azure"></Card>
<Card title="Redis" icon="/images/provider-icons/redis.svg" href="/components/vectordbs/dbs/redis"></Card>
<Card title="Valkey" icon="/images/provider-icons/valkey.svg" href="/components/vectordbs/dbs/valkey"></Card>
+1
View File
@@ -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",
+1
View File
@@ -0,0 +1 @@
<svg fill="#8F74E0" role="img" viewBox="0 0 93.9 59.4" xmlns="http://www.w3.org/2000/svg"><title>Oracle</title><path d="M30.5,59.4H65c16.4-0.4,29.3-14.1,28.9-30.4C93.5,13.1,80.7,0.4,65,0H30.5C14.1-0.4,0.4,12.5,0,28.9s12.5,30,28.9,30.4C29.4,59.4,29.9,59.4,30.5,59.4 M64.2,48.9h-33c-10.6-0.3-18.9-9.2-18.6-19.8C13,19,21.1,10.8,31.2,10.5h33c10.6-0.3,19.5,8,19.8,18.6c0.3,10.6-8,19.5-18.6,19.8C65,48.9,64.6,48.9,64.2,48.9"/></svg>

After

Width:  |  Height:  |  Size: 427 B

+1
View File
@@ -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.
+113
View File
@@ -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
+1
View File
@@ -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
+1
View File
@@ -35,6 +35,7 @@ class VectorStoreConfig(BaseModel):
"langchain": "LangchainConfig",
"s3_vectors": "S3VectorsConfig",
"turbopuffer": "TurbopufferConfig",
"oracledb": "OracleAIVectorSearchConfig",
}
@model_validator(mode="after")
+592
View File
@@ -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
+1
View File
@@ -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",
File diff suppressed because it is too large Load Diff