Compare commits
6 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| b357a5a1b0 | |||
| d653b63fac | |||
| cc4671579f | |||
| 01afdde3e7 | |||
| d6d89c987b | |||
| c2150e8f1a |
@@ -106,7 +106,7 @@ jobs:
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
pip install -e ".[test,graph,vector_stores,llms,extras]"
|
||||
pip install ruff
|
||||
pip install ruff==0.16.0
|
||||
- name: Run Linting
|
||||
if: needs.check_changes.outputs.mem0_changed == 'true'
|
||||
run: make lint
|
||||
|
||||
@@ -11,7 +11,7 @@ install:
|
||||
hatch env create
|
||||
|
||||
install_all:
|
||||
pip install ruff==0.6.9 groq together boto3 litellm ollama chromadb weaviate weaviate-client sentence_transformers vertexai \
|
||||
pip install ruff==0.16.0 groq together boto3 litellm ollama chromadb weaviate weaviate-client sentence_transformers vertexai \
|
||||
google-generativeai elasticsearch opensearch-py vecs "pinecone<7.0.0" pinecone-text faiss-cpu langchain-community \
|
||||
upstash-vector azure-search-documents langchain-memgraph langchain-neo4j langchain-aws rank-bm25 pymochow pymongo psycopg kuzu databricks-sdk valkey
|
||||
|
||||
|
||||
@@ -7,6 +7,18 @@ mode: "wide"
|
||||
<Tabs>
|
||||
<Tab title="Python">
|
||||
|
||||
<Update label="2026-07-25" description="v2.0.14">
|
||||
|
||||
**New Features:**
|
||||
- **Vector Stores:** Add an Oracle AI Vector Search provider (`oracledb`) with connection pooling, `HNSW`/`IVF` indexes, JSON metadata filtering, and six selectable distance metrics ([#5358](https://github.com/mem0ai/mem0/pull/5358))
|
||||
|
||||
**Bug Fixes:**
|
||||
- **Vector Stores:** Translate a `"*"` filter value in OpenSearch into an `exists` query for every key, not just identity keys. It was previously ignored or matched literally against the string `"*"`, so a wildcard filter returned nothing ([#6522](https://github.com/mem0ai/mem0/pull/6522))
|
||||
- **Vector Stores:** Re-raise errors from OpenSearch `search()` instead of returning `[]`, so a transport, auth, or index misconfiguration surfaces instead of looking like zero matches. `keyword_search()` still degrades on failure, since it is a best-effort BM25 signal ([#6519](https://github.com/mem0ai/mem0/pull/6519))
|
||||
- **Vector Stores:** Guard the `text` field in Milvus `update()` behind the `_has_bm25_schema` check, matching `insert()`, so updating a memory in a collection without the BM25 `text`/`sparse` schema no longer fails ([#5705](https://github.com/mem0ai/mem0/pull/5705))
|
||||
|
||||
</Update>
|
||||
|
||||
<Update label="2026-07-22" description="v2.0.13">
|
||||
|
||||
**Bug Fixes:**
|
||||
@@ -1145,6 +1157,18 @@ See the [OSS v2 to v3 migration guide](https://docs.mem0.ai/migration/oss-v2-to-
|
||||
|
||||
<Tab title="TypeScript">
|
||||
|
||||
<Update label="2026-07-25" description="v3.1.2">
|
||||
|
||||
**Bug Fixes:**
|
||||
- **Vector Stores:** Apply every operator in a Cassandra compound field filter (e.g. `{ age: { gte: 10, lte: 20 } }`) instead of stopping after the first, so the remaining bounds are no longer silently ignored ([#6511](https://github.com/mem0ai/mem0/pull/6511))
|
||||
- **Vector Stores:** Stop the Chroma where-clause translator from dropping filter conditions. Same-field ranges (`gte` + `lte`), multi-field conditions inside `$or`, and negated `contains`/`icontains` under `$not` each collapsed to a single clause or vanished, widening the search instead of narrowing it ([#6521](https://github.com/mem0ai/mem0/pull/6521))
|
||||
- **Vector Stores:** Skip `"*"` wildcard filter values in Milvus instead of matching them literally, so a filter like `{ user_id: "*" }` no longer returns zero memories ([#6508](https://github.com/mem0ai/mem0/pull/6508))
|
||||
- **Vector Stores:** Read `textLemmatized` for BM25 keyword search on Milvus, OpenSearch, and MongoDB, matching the field the memory layer actually writes, so hybrid search on those backends no longer loses the keyword signal ([#6497](https://github.com/mem0ai/mem0/pull/6497))
|
||||
- **LLMs:** Forward `responseFormat` to Gemini's `responseMimeType` in `generateResponse()`, so requesting `json_object` returns JSON instead of free-form text ([#6468](https://github.com/mem0ai/mem0/pull/6468))
|
||||
- **LLMs:** Find the Anthropic text block by type instead of indexing `content[0]`, so a thinking-enabled model whose `thinking` block comes first no longer throws `Unexpected response type from Anthropic API` ([#6506](https://github.com/mem0ai/mem0/pull/6506))
|
||||
|
||||
</Update>
|
||||
|
||||
<Update label="2026-07-22" description="v3.1.1">
|
||||
|
||||
**New Features:**
|
||||
|
||||
@@ -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 |
|
||||
|
||||
@@ -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}},
|
||||
)
|
||||
```
|
||||
@@ -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>
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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 |
@@ -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.
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "mem0ai",
|
||||
"version": "3.1.1",
|
||||
"version": "3.1.2",
|
||||
"description": "The Memory Layer For Your AI Apps",
|
||||
"main": "./dist/index.js",
|
||||
"module": "./dist/index.mjs",
|
||||
|
||||
@@ -455,44 +455,64 @@ export class CassandraDB implements VectorStore {
|
||||
return value.includes(payloadValue);
|
||||
}
|
||||
|
||||
// Every operator present in a compound condition must hold (AND), so check
|
||||
// them all instead of returning on the first match. Returning early meant a
|
||||
// range like { gte: 10, lte: 20 } only applied `gte`. Mirrors the databricks
|
||||
// store's matcher.
|
||||
let sawOperator = false;
|
||||
|
||||
if ("eq" in value) {
|
||||
return payloadValue === value.eq;
|
||||
sawOperator = true;
|
||||
if (payloadValue !== value.eq) return false;
|
||||
}
|
||||
if ("ne" in value) {
|
||||
return payloadValue !== value.ne;
|
||||
sawOperator = true;
|
||||
if (payloadValue === value.ne) return false;
|
||||
}
|
||||
if ("gt" in value) {
|
||||
return payloadValue > value.gt;
|
||||
sawOperator = true;
|
||||
if (!(payloadValue > value.gt)) return false;
|
||||
}
|
||||
if ("gte" in value) {
|
||||
return payloadValue >= value.gte;
|
||||
sawOperator = true;
|
||||
if (!(payloadValue >= value.gte)) return false;
|
||||
}
|
||||
if ("lt" in value) {
|
||||
return payloadValue < value.lt;
|
||||
sawOperator = true;
|
||||
if (!(payloadValue < value.lt)) return false;
|
||||
}
|
||||
if ("lte" in value) {
|
||||
return payloadValue <= value.lte;
|
||||
sawOperator = true;
|
||||
if (!(payloadValue <= value.lte)) return false;
|
||||
}
|
||||
if ("in" in value) {
|
||||
return Array.isArray(value.in) && value.in.includes(payloadValue);
|
||||
sawOperator = true;
|
||||
if (!Array.isArray(value.in) || !value.in.includes(payloadValue))
|
||||
return false;
|
||||
}
|
||||
if ("nin" in value) {
|
||||
return !Array.isArray(value.nin) || !value.nin.includes(payloadValue);
|
||||
sawOperator = true;
|
||||
if (Array.isArray(value.nin) && value.nin.includes(payloadValue))
|
||||
return false;
|
||||
}
|
||||
if ("contains" in value) {
|
||||
return (
|
||||
typeof payloadValue === "string" &&
|
||||
payloadValue.includes(value.contains)
|
||||
);
|
||||
sawOperator = true;
|
||||
if (
|
||||
typeof payloadValue !== "string" ||
|
||||
!payloadValue.includes(value.contains)
|
||||
)
|
||||
return false;
|
||||
}
|
||||
if ("icontains" in value) {
|
||||
return (
|
||||
typeof payloadValue === "string" &&
|
||||
payloadValue.toLowerCase().includes(value.icontains.toLowerCase())
|
||||
);
|
||||
sawOperator = true;
|
||||
if (
|
||||
typeof payloadValue !== "string" ||
|
||||
!payloadValue.toLowerCase().includes(value.icontains.toLowerCase())
|
||||
)
|
||||
return false;
|
||||
}
|
||||
|
||||
return payloadValue === value;
|
||||
return sawOperator ? true : payloadValue === value;
|
||||
}
|
||||
|
||||
private filterVector(
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
/// <reference types="jest" />
|
||||
/**
|
||||
* Cassandra vector store — filter matching unit tests.
|
||||
*
|
||||
* Cassandra has no server-side metadata filter, so search()/list() scan rows
|
||||
* and apply filters in-app via matchFieldCondition(). These tests drive that
|
||||
* matcher through the public list() API with an injected fake client.
|
||||
*/
|
||||
import { CassandraDB } from "../src/vector_stores/cassandra";
|
||||
|
||||
type Row = { id: string; payload: Record<string, any> };
|
||||
|
||||
// Minimal fake driver: CREATE statements during initialize() return nothing;
|
||||
// a SELECT returns the seeded rows in one page (pageState undefined => stop).
|
||||
function fakeClient(rows: Row[]) {
|
||||
return {
|
||||
async connect() {},
|
||||
async execute(query: string) {
|
||||
if (/^\s*SELECT/i.test(query)) {
|
||||
return { rows, pageState: undefined };
|
||||
}
|
||||
return { rows: [], pageState: undefined };
|
||||
},
|
||||
async shutdown() {},
|
||||
};
|
||||
}
|
||||
|
||||
function makeStore(rows: Row[]) {
|
||||
return new CassandraDB({
|
||||
keyspace: "mem0",
|
||||
collectionName: "mem0",
|
||||
dimension: 3,
|
||||
client: fakeClient(rows) as any,
|
||||
} as any);
|
||||
}
|
||||
|
||||
describe("CassandraDB filter matching", () => {
|
||||
const rows: Row[] = [
|
||||
{ id: "a", payload: { data: "a", age: 5 } },
|
||||
{ id: "b", payload: { data: "b", age: 15 } },
|
||||
{ id: "c", payload: { data: "c", age: 25 } },
|
||||
];
|
||||
|
||||
it("applies every operator in a compound range filter (not just the first)", async () => {
|
||||
const store = makeStore(rows);
|
||||
// age in [10, 20]: only "b" (15) qualifies. The old matcher returned on the
|
||||
// first operator (gte), so "c" (25) leaked through because lte was ignored.
|
||||
const [results] = await store.list({ age: { gte: 10, lte: 20 } });
|
||||
expect(results.map((r) => r.id)).toEqual(["b"]);
|
||||
});
|
||||
|
||||
it("still matches a single-operator filter", async () => {
|
||||
const store = makeStore(rows);
|
||||
const [results] = await store.list({ age: { gte: 15 } });
|
||||
expect(results.map((r) => r.id).sort()).toEqual(["b", "c"]);
|
||||
});
|
||||
|
||||
it("treats a plain equality filter as before", async () => {
|
||||
const store = makeStore(rows);
|
||||
const [results] = await store.list({ age: 15 });
|
||||
expect(results.map((r) => r.id)).toEqual(["b"]);
|
||||
});
|
||||
});
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -35,6 +35,7 @@ class VectorStoreConfig(BaseModel):
|
||||
"langchain": "LangchainConfig",
|
||||
"s3_vectors": "S3VectorsConfig",
|
||||
"turbopuffer": "TurbopufferConfig",
|
||||
"oracledb": "OracleAIVectorSearchConfig",
|
||||
}
|
||||
|
||||
@model_validator(mode="after")
|
||||
|
||||
@@ -298,10 +298,12 @@ class MilvusDB(VectorStoreBase):
|
||||
if payload is None:
|
||||
payload = existing[0].get("metadata")
|
||||
|
||||
text = ""
|
||||
if payload:
|
||||
text = (payload.get("text_lemmatized") or payload.get("data", ""))[:65535]
|
||||
schema = {"id": vector_id, "vectors": vector, "metadata": payload, "text": text}
|
||||
schema = {"id": vector_id, "vectors": vector, "metadata": payload}
|
||||
if self._has_bm25_schema:
|
||||
text = ""
|
||||
if payload:
|
||||
text = (payload.get("text_lemmatized") or payload.get("data", ""))[:65535]
|
||||
schema["text"] = text
|
||||
self.client.upsert(collection_name=self.collection_name, data=schema)
|
||||
|
||||
def get(self, vector_id) -> Optional[OutputData]:
|
||||
|
||||
@@ -248,7 +248,7 @@ class OpenSearchDB(VectorStoreBase):
|
||||
return results
|
||||
except Exception as e:
|
||||
logger.error(f"Error during search: {e}", exc_info=True)
|
||||
return []
|
||||
raise
|
||||
|
||||
def keyword_search(self, query, top_k=5, filters=None):
|
||||
"""Search for memories using BM25 keyword matching.
|
||||
@@ -293,8 +293,12 @@ class OpenSearchDB(VectorStoreBase):
|
||||
]
|
||||
return results
|
||||
except Exception as e:
|
||||
logger.error(f"Error during keyword search: {e}")
|
||||
return []
|
||||
# Do NOT re-raise here: keyword_search() is a best-effort helper that
|
||||
# search() may call to augment semantic results. Raising would crash
|
||||
# the whole search() call on a keyword-only failure (regression per
|
||||
# maintainer review on #6519). Log with exc_info and degrade to None.
|
||||
logger.error(f"Error during keyword search: {e}", exc_info=True)
|
||||
return None
|
||||
|
||||
def delete(self, vector_id: str) -> None:
|
||||
"""Delete a vector by custom ID."""
|
||||
|
||||
@@ -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
|
||||
+6
-2
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "mem0ai"
|
||||
version = "2.0.13"
|
||||
version = "2.0.14"
|
||||
description = "Long-term memory for AI Agents"
|
||||
authors = [
|
||||
{ name = "Mem0", email = "support@mem0.ai" }
|
||||
@@ -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",
|
||||
@@ -80,7 +81,7 @@ test = [
|
||||
"pytest-asyncio>=0.23.7",
|
||||
]
|
||||
dev = [
|
||||
"ruff>=0.6.5",
|
||||
"ruff==0.16.0",
|
||||
"isort>=5.13.2",
|
||||
"pytest>=8.2.2",
|
||||
]
|
||||
@@ -149,6 +150,9 @@ test = [
|
||||
line-length = 120
|
||||
exclude = ["openmemory/"]
|
||||
|
||||
[tool.ruff.lint]
|
||||
select = ["E4", "E7", "E9", "F"]
|
||||
|
||||
[tool.ruff.lint.isort]
|
||||
known-first-party = ["mem0", "mem0_cli"]
|
||||
|
||||
|
||||
@@ -345,10 +345,44 @@ class TestMilvusDB:
|
||||
assert '(metadata["active"] == True)' in result
|
||||
assert '(metadata["deleted"] == False)' in result
|
||||
|
||||
def test_update_omits_text_field_on_pre_v3_collection(self, mock_milvus_client):
|
||||
"""update() must not include 'text' for collections without BM25 schema."""
|
||||
mock_milvus_client.has_collection.return_value = True
|
||||
mock_milvus_client.describe_collection.return_value = {
|
||||
"fields": [
|
||||
{"name": "id"},
|
||||
{"name": "vectors"},
|
||||
{"name": "metadata"},
|
||||
]
|
||||
}
|
||||
db = MilvusDB(
|
||||
url="http://localhost:19530",
|
||||
token="test_token",
|
||||
collection_name="legacy_collection",
|
||||
embedding_model_dims=1536,
|
||||
metric_type=MetricType.COSINE,
|
||||
db_name="test_db",
|
||||
)
|
||||
assert db._has_bm25_schema is False
|
||||
|
||||
db.update(vector_id="id1", vector=[0.1] * 1536, payload={"data": "hello"})
|
||||
|
||||
upserted = mock_milvus_client.upsert.call_args[1]["data"]
|
||||
assert "text" not in upserted
|
||||
|
||||
def test_update_includes_text_field_on_v3_collection(self, milvus_db, mock_milvus_client):
|
||||
"""update() must include 'text' for collections with BM25 schema."""
|
||||
assert milvus_db._has_bm25_schema is True
|
||||
|
||||
milvus_db.update(vector_id="id1", vector=[0.1] * 1536, payload={"data": "hello"})
|
||||
|
||||
upserted = mock_milvus_client.upsert.call_args[1]["data"]
|
||||
assert upserted["text"] == "hello"
|
||||
|
||||
def test_collection_already_exists(self, mock_milvus_client):
|
||||
"""Test that existing collection is not recreated."""
|
||||
mock_milvus_client.has_collection.return_value = True
|
||||
|
||||
|
||||
MilvusDB(
|
||||
url="http://localhost:19530",
|
||||
token="test_token",
|
||||
@@ -357,7 +391,7 @@ class TestMilvusDB:
|
||||
metric_type=MetricType.L2,
|
||||
db_name="test_db"
|
||||
)
|
||||
|
||||
|
||||
# create_collection should not be called
|
||||
mock_milvus_client.create_collection.assert_not_called()
|
||||
|
||||
|
||||
@@ -424,10 +424,25 @@ class TestOpenSearchDB(unittest.TestCase):
|
||||
|
||||
@patch("mem0.vector_stores.opensearch.logger")
|
||||
def test_search_error_logs_with_exc_info(self, mock_logger):
|
||||
"""Search error logging should include exc_info for full stack trace."""
|
||||
"""Search errors should log with exc_info and re-raise (not swallow as [])."""
|
||||
self.client_mock.search.side_effect = Exception("Search failed")
|
||||
results = self.os_db.search(query="", vectors=[[0.1] * 1536], top_k=5)
|
||||
self.assertEqual(results, [])
|
||||
with self.assertRaises(Exception):
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], top_k=5)
|
||||
mock_logger.error.assert_called_once()
|
||||
call_kwargs = mock_logger.error.call_args
|
||||
self.assertTrue(call_kwargs[1].get("exc_info"), "logger.error must be called with exc_info=True")
|
||||
|
||||
@patch("mem0.vector_stores.opensearch.logger")
|
||||
def test_keyword_search_error_logs_and_degrades(self, mock_logger):
|
||||
"""Keyword search errors should log with exc_info and degrade to None (not raise).
|
||||
|
||||
keyword_search() is a best-effort augmentation for search(); raising here
|
||||
would crash the whole search() call on a keyword-only failure (regression
|
||||
per maintainer review on #6519).
|
||||
"""
|
||||
self.client_mock.search.side_effect = Exception("Keyword search failed")
|
||||
result = self.os_db.keyword_search(query="test", top_k=5)
|
||||
self.assertIsNone(result)
|
||||
mock_logger.error.assert_called_once()
|
||||
call_kwargs = mock_logger.error.call_args
|
||||
self.assertTrue(call_kwargs[1].get("exc_info"), "logger.error must be called with exc_info=True")
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user