Compare commits
16 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| b357a5a1b0 | |||
| d653b63fac | |||
| cc4671579f | |||
| 01afdde3e7 | |||
| d6d89c987b | |||
| c2150e8f1a | |||
| 19c7bb84a2 | |||
| e6281ab724 | |||
| a71d7bdbe3 | |||
| 56ec7d20f1 | |||
| ca2abca2b8 | |||
| 0e582adc6c | |||
| 9caffeaa7b | |||
| a58e0586ad | |||
| a9cb4bb644 | |||
| 7bf84b8d38 |
@@ -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>
|
||||
|
||||
+2
-1
@@ -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",
|
||||
@@ -633,7 +634,7 @@
|
||||
},
|
||||
{
|
||||
"source": "/platform/features/expiration-date",
|
||||
"destination": "/"
|
||||
"destination": "/platform/features/memory-expiration"
|
||||
},
|
||||
{
|
||||
"source": "/cookbooks/essentials/memory-expiration-short-and-long-term",
|
||||
|
||||
@@ -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,9 +1,10 @@
|
||||
---
|
||||
title: Memory Expiration
|
||||
description: "Give a memory a shelf life: set an expiration date and it stops surfacing in search once that date passes, without deleting the record."
|
||||
title: "Memory Expiration in Mem0"
|
||||
sidebarTitle: "Memory Expiration"
|
||||
description: "Set an expiration date on a Mem0 memory and it stops surfacing in search once that date passes. Nothing is deleted. Works on Platform and Open Source."
|
||||
---
|
||||
|
||||
# Memory Expiration
|
||||
## Why Use Memory Expiration?
|
||||
|
||||
Some facts are only true for a while. A trial plan ends, a seasonal preference goes stale, a support ticket ages past its retention window. Set an `expiration_date` on a memory and Mem0 stops surfacing it once that date passes, so you don't need a cleanup job hunting for rows to delete.
|
||||
|
||||
@@ -18,7 +19,7 @@ Some facts are only true for a while. A trial plan ends, a seasonal preference g
|
||||
- **No expiration date means never expires.** That is the default for every memory.
|
||||
- **Malformed dates fail open**: a stored value Mem0 can't parse is treated as *not* expired. A bad date never makes a memory silently vanish.
|
||||
|
||||
## Set an expiration date
|
||||
## How do you set an expiration date on a memory?
|
||||
|
||||
Set it when you add the memory:
|
||||
|
||||
@@ -97,7 +98,7 @@ The same parameter and spelling work on the OSS `Memory` class. See <Link href="
|
||||
Expired memories are dropped *before* your `top_k` is applied, so Mem0 widens the internal candidate pool first and short result sets are rare. They are not impossible: if nearly every memory in a scope has expired, a call can still return fewer than `top_k` results. Pass `show_expired: true` to get the full set back.
|
||||
</Note>
|
||||
|
||||
## Clear an expiration date
|
||||
## How do you clear or remove an expiration date?
|
||||
|
||||
Pass an explicit `None` (Python) or `null` (TypeScript) to make the memory permanent again. The SDKs deliberately preserve that null instead of treating it as "argument not supplied".
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -100,12 +100,16 @@ export class AnthropicLLM implements LLM {
|
||||
return { content, role: "assistant", toolCalls };
|
||||
}
|
||||
|
||||
const firstBlock = response.content[0];
|
||||
if (firstBlock.type === "text") {
|
||||
return firstBlock.text;
|
||||
} else {
|
||||
throw new Error("Unexpected response type from Anthropic API");
|
||||
// Thinking-enabled responses put a thinking block before the text block,
|
||||
// and a response can carry no text block at all, so find the text block
|
||||
// like the tools branch above instead of indexing content[0]. Mirrors the
|
||||
// Python provider (#6481).
|
||||
for (const block of response.content) {
|
||||
if (block.type === "text") {
|
||||
return block.text;
|
||||
}
|
||||
}
|
||||
return "";
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
|
||||
@@ -59,6 +59,15 @@ export class GoogleLLM implements LLM {
|
||||
];
|
||||
}
|
||||
|
||||
// Honor a requested JSON response format (parity with the Python SDK's
|
||||
// mem0/llms/gemini.py). Gemini's structured output is opt-in via
|
||||
// responseMimeType — without it the model is never told to emit JSON, so
|
||||
// callers passing json_object silently get free-form text and depend on a
|
||||
// fragile markdown-fence strip downstream.
|
||||
if (responseFormat?.type === "json_object") {
|
||||
config.responseMimeType = "application/json";
|
||||
}
|
||||
|
||||
const completion = await this.google.models.generateContent({
|
||||
contents,
|
||||
model: this.model,
|
||||
|
||||
@@ -440,4 +440,52 @@ describe("generateWhereClause", () => {
|
||||
ChromaDB.generateWhereClause({ $not: [{ age: { gt: 18 } }] }),
|
||||
).toEqual({ age: { $lte: 18 } });
|
||||
});
|
||||
|
||||
// Regression tests for the three where-clause translation bugs fixed in
|
||||
// Python by #6452 and tracked for the TS SDK in #6513. ChromaDB allows
|
||||
// exactly one operator or field per dict level.
|
||||
it("keeps both bounds of a same-field range as $and-combined clauses", () => {
|
||||
expect(ChromaDB.generateWhereClause({ age: { gte: 18, lte: 65 } })).toEqual(
|
||||
{ $and: [{ age: { $gte: 18 } }, { age: { $lte: 65 } }] },
|
||||
);
|
||||
});
|
||||
|
||||
it("keeps same-field ranges inside $or branches", () => {
|
||||
expect(
|
||||
ChromaDB.generateWhereClause({
|
||||
$or: [{ age: { gte: 18, lte: 65 } }, { vip: true }],
|
||||
}),
|
||||
).toEqual({
|
||||
$or: [
|
||||
{ $and: [{ age: { $gte: 18 } }, { age: { $lte: 65 } }] },
|
||||
{ vip: { $eq: true } },
|
||||
],
|
||||
});
|
||||
});
|
||||
|
||||
it("wraps multi-field conditions inside $or in $and", () => {
|
||||
expect(
|
||||
ChromaDB.generateWhereClause({
|
||||
$or: [{ age: { gte: 18 }, vip: true }, { city: "sh" }],
|
||||
}),
|
||||
).toEqual({
|
||||
$or: [
|
||||
{ $and: [{ age: { $gte: 18 } }, { vip: { $eq: true } }] },
|
||||
{ city: { $eq: "sh" } },
|
||||
],
|
||||
});
|
||||
});
|
||||
|
||||
it("negates $not contains/icontains instead of dropping the clause", () => {
|
||||
expect(
|
||||
ChromaDB.generateWhereClause({
|
||||
$not: [{ title: { contains: "draft" } }],
|
||||
}),
|
||||
).toEqual({ title: { $ne: "draft" } });
|
||||
expect(
|
||||
ChromaDB.generateWhereClause({
|
||||
$not: [{ title: { icontains: "draft" } }],
|
||||
}),
|
||||
).toEqual({ title: { $ne: "draft" } });
|
||||
});
|
||||
});
|
||||
|
||||
@@ -267,6 +267,36 @@ describe("Milvus vector store (TS OSS SDK)", () => {
|
||||
});
|
||||
});
|
||||
|
||||
it("skips a wildcard '*' filter value and keeps the rest", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
client.searchResponse = { results: [] };
|
||||
const store = makeStore(client, { metricType: "COSINE" });
|
||||
await store.initialize();
|
||||
|
||||
// "*" means match-any: it must be dropped, not emitted as `== "*"` (which
|
||||
// matches nothing), leaving only the real agent_id clause.
|
||||
await store.search([0.1, 0.2, 0.3], 5, {
|
||||
user_id: "*",
|
||||
agent_id: "a1",
|
||||
});
|
||||
|
||||
const searchCall = client.calls.find((c) => c.method === "search")!;
|
||||
expect(searchCall.args.filter).toBe('(metadata["agent_id"] == "a1")');
|
||||
});
|
||||
|
||||
it("omits the filter entirely when every value is a wildcard", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
const store = makeStore(client);
|
||||
await store.initialize();
|
||||
|
||||
await store.list({ user_id: "*" });
|
||||
|
||||
// All clauses dropped, so list() falls back to its match-all "" filter
|
||||
// rather than a literal `(metadata["user_id"] == "*")` that matches nothing.
|
||||
const queryCall = client.calls.filter((c) => c.method === "query").pop()!;
|
||||
expect(queryCall.args.filter).toBe("");
|
||||
});
|
||||
|
||||
it("normalises L2 distances into a 0..1 similarity score", async () => {
|
||||
const client = new FakeMilvusClient({ existing: ["mem0"] });
|
||||
client.searchResponse = {
|
||||
@@ -432,12 +462,18 @@ describe("Milvus vector store (TS OSS SDK)", () => {
|
||||
["a"],
|
||||
[{ data: "hello world", text_lemmatized: "hello world lemma" }],
|
||||
);
|
||||
await store.insert(
|
||||
[[0.2, 0.3, 0.4]],
|
||||
["c"],
|
||||
[{ data: "hello world", textLemmatized: "hello world camel" }],
|
||||
);
|
||||
// Falls back to raw data when there is no lemmatized text.
|
||||
await store.insert([[0.4, 0.5, 0.6]], ["b"], [{ data: "just data" }]);
|
||||
|
||||
const insertCalls = client.calls.filter((c) => c.method === "insert");
|
||||
expect(insertCalls[0].args.data[0].text).toBe("hello world lemma");
|
||||
expect(insertCalls[1].args.data[0].text).toBe("just data");
|
||||
expect(insertCalls[1].args.data[0].text).toBe("hello world camel");
|
||||
expect(insertCalls[2].args.data[0].text).toBe("just data");
|
||||
});
|
||||
|
||||
it("writes the BM25 text field on update for a BM25 collection", async () => {
|
||||
|
||||
@@ -5,6 +5,7 @@ const mockFindOne = jest.fn();
|
||||
const mockUpdateOne = jest.fn();
|
||||
const mockListSearchIndexes = jest.fn();
|
||||
const mockCreateSearchIndex = jest.fn();
|
||||
const mockDropSearchIndex = jest.fn();
|
||||
const mockDrop = jest.fn();
|
||||
const mockToArray = jest.fn();
|
||||
const mockLimit = jest.fn().mockReturnThis();
|
||||
@@ -26,6 +27,7 @@ const mockCollection = {
|
||||
updateOne: mockUpdateOne,
|
||||
listSearchIndexes: mockListSearchIndexes,
|
||||
createSearchIndex: mockCreateSearchIndex,
|
||||
dropSearchIndex: mockDropSearchIndex,
|
||||
drop: mockDrop,
|
||||
find: mockFind,
|
||||
aggregate: mockAggregate,
|
||||
@@ -74,6 +76,25 @@ describe("MongoDB Vector Store", () => {
|
||||
await store.close();
|
||||
});
|
||||
|
||||
const expectedTextSearchIndexDefinition = {
|
||||
name: "test_col_text_search_index",
|
||||
definition: {
|
||||
mappings: {
|
||||
dynamic: false,
|
||||
fields: {
|
||||
payload: {
|
||||
type: "document",
|
||||
fields: {
|
||||
data: { type: "string" },
|
||||
textLemmatized: { type: "string" },
|
||||
text_lemmatized: { type: "string" },
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
it("should initialize client and check/create collection and indexes", async () => {
|
||||
await store.initialize();
|
||||
|
||||
@@ -84,6 +105,83 @@ describe("MongoDB Vector Store", () => {
|
||||
});
|
||||
expect(mockCollection.deleteOne).toHaveBeenCalledWith({ _id: 0 });
|
||||
expect(mockCreateSearchIndex).toHaveBeenCalledTimes(2);
|
||||
expect(mockCreateSearchIndex).toHaveBeenCalledWith(
|
||||
expectedTextSearchIndexDefinition,
|
||||
);
|
||||
});
|
||||
|
||||
it("should drop and recreate a stale text search index on upgrade", async () => {
|
||||
mockListCollections.mockReturnValue({
|
||||
toArray: jest.fn().mockResolvedValue([{ name: "test_col" }]),
|
||||
});
|
||||
mockListSearchIndexes.mockReturnValue({
|
||||
toArray: jest.fn().mockResolvedValue([
|
||||
{ name: "test_col_vector_index" },
|
||||
{
|
||||
name: "test_col_text_search_index",
|
||||
definition: {
|
||||
mappings: {
|
||||
dynamic: false,
|
||||
fields: {
|
||||
payload: {
|
||||
type: "document",
|
||||
fields: {
|
||||
data: { type: "string" },
|
||||
text_lemmatized: { type: "string" },
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
]),
|
||||
});
|
||||
|
||||
await store.initialize();
|
||||
|
||||
expect(mockDropSearchIndex).toHaveBeenCalledWith(
|
||||
"test_col_text_search_index",
|
||||
);
|
||||
expect(mockCreateSearchIndex).toHaveBeenCalledTimes(1);
|
||||
expect(mockCreateSearchIndex).toHaveBeenCalledWith(
|
||||
expectedTextSearchIndexDefinition,
|
||||
);
|
||||
});
|
||||
|
||||
it("should not recreate a text search index that already has textLemmatized", async () => {
|
||||
mockListCollections.mockReturnValue({
|
||||
toArray: jest.fn().mockResolvedValue([{ name: "test_col" }]),
|
||||
});
|
||||
mockListSearchIndexes.mockReturnValue({
|
||||
toArray: jest.fn().mockResolvedValue([
|
||||
{ name: "test_col_vector_index" },
|
||||
{
|
||||
name: "test_col_text_search_index",
|
||||
latestDefinition: expectedTextSearchIndexDefinition.definition,
|
||||
},
|
||||
]),
|
||||
});
|
||||
|
||||
await store.initialize();
|
||||
|
||||
expect(mockDropSearchIndex).not.toHaveBeenCalled();
|
||||
expect(mockCreateSearchIndex).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should map payload.textLemmatized in the text search index", async () => {
|
||||
await store.initialize();
|
||||
|
||||
const textIndexCall = mockCreateSearchIndex.mock.calls.find(
|
||||
([arg]: any[]) => arg.name === "test_col_text_search_index",
|
||||
);
|
||||
expect(textIndexCall).toBeDefined();
|
||||
expect(textIndexCall![0].definition.mappings.fields.payload.fields).toEqual(
|
||||
{
|
||||
data: { type: "string" },
|
||||
text_lemmatized: { type: "string" },
|
||||
textLemmatized: { type: "string" },
|
||||
},
|
||||
);
|
||||
});
|
||||
|
||||
it("should insert documents correctly", async () => {
|
||||
@@ -197,7 +295,11 @@ describe("MongoDB Vector Store", () => {
|
||||
index: "test_col_text_search_index",
|
||||
text: {
|
||||
query: "test",
|
||||
path: ["payload.data", "payload.text_lemmatized"],
|
||||
path: [
|
||||
"payload.data",
|
||||
"payload.text_lemmatized",
|
||||
"payload.textLemmatized",
|
||||
],
|
||||
},
|
||||
},
|
||||
},
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -273,14 +273,14 @@ export class ChromaDB implements VectorStore {
|
||||
private static convertCondition(
|
||||
key: string,
|
||||
value: any,
|
||||
): Record<string, any> | null {
|
||||
): Array<Record<string, any>> {
|
||||
// Wildcard - ChromaDB has no direct wildcard, so skip this filter.
|
||||
if (value === "*") {
|
||||
return null;
|
||||
return [];
|
||||
}
|
||||
|
||||
if (Array.isArray(value)) {
|
||||
return { [key]: { $in: value } };
|
||||
return [{ [key]: { $in: value } }];
|
||||
}
|
||||
|
||||
if (value !== null && typeof value === "object") {
|
||||
@@ -294,19 +294,31 @@ export class ChromaDB implements VectorStore {
|
||||
in: "$in",
|
||||
nin: "$nin",
|
||||
};
|
||||
const condition: Record<string, any> = {};
|
||||
for (const [op, val] of Object.entries(value)) {
|
||||
if (op in opMap) {
|
||||
condition[key] = { [opMap[op]]: val };
|
||||
} else {
|
||||
// contains/icontains and unknown operators fall back to equality.
|
||||
condition[key] = { $eq: val };
|
||||
}
|
||||
}
|
||||
return condition;
|
||||
// ChromaDB allows exactly one operator per field expression, so each
|
||||
// operator becomes its own clause (combined with $and by the caller).
|
||||
// Previously each operator overwrote the last, silently dropping range
|
||||
// bounds. contains/icontains and unknown operators fall back to
|
||||
// equality.
|
||||
return Object.entries(value).map(([op, val]) => ({
|
||||
[key]: { [opMap[op] ?? "$eq"]: val },
|
||||
}));
|
||||
}
|
||||
|
||||
return { [key]: { $eq: value } };
|
||||
return [{ [key]: { $eq: value } }];
|
||||
}
|
||||
|
||||
/** Combine clauses under a logical operator, unwrapping singletons. */
|
||||
private static combineClauses(
|
||||
clauses: Array<Record<string, any>>,
|
||||
operator: "$and" | "$or",
|
||||
): Record<string, any> | null {
|
||||
if (clauses.length === 0) {
|
||||
return null;
|
||||
}
|
||||
if (clauses.length === 1) {
|
||||
return clauses[0];
|
||||
}
|
||||
return { [operator]: clauses };
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -337,66 +349,55 @@ export class ChromaDB implements VectorStore {
|
||||
if (key === "$or" || key === "OR") {
|
||||
const orConditions: any[] = [];
|
||||
for (const condition of value as any[]) {
|
||||
const built: Record<string, any> = {};
|
||||
const subClauses: Array<Record<string, any>> = [];
|
||||
for (const [subKey, subValue] of Object.entries(condition)) {
|
||||
const converted = ChromaDB.convertCondition(subKey, subValue);
|
||||
if (converted) Object.assign(built, converted);
|
||||
subClauses.push(...ChromaDB.convertCondition(subKey, subValue));
|
||||
}
|
||||
if (Object.keys(built).length > 0) orConditions.push(built);
|
||||
}
|
||||
if (orConditions.length > 1) {
|
||||
processed.push({ $or: orConditions });
|
||||
} else if (orConditions.length === 1) {
|
||||
processed.push(orConditions[0]);
|
||||
// Multi-field conditions must be wrapped in $and — ChromaDB rejects
|
||||
// flat objects with more than one field per level.
|
||||
const combined = ChromaDB.combineClauses(subClauses, "$and");
|
||||
if (combined) orConditions.push(combined);
|
||||
}
|
||||
const combinedOr = ChromaDB.combineClauses(orConditions, "$or");
|
||||
if (combinedOr) processed.push(combinedOr);
|
||||
} else if (key === "$not" || key === "NOT") {
|
||||
// De Morgan: NOT(a AND b) is (NOT a) OR (NOT b), so the negated fields
|
||||
// within one condition are combined with $or, and separate conditions
|
||||
// are combined with $and. This mirrors the Python SDK's ChromaDB port.
|
||||
const negatedPerGroup: any[] = [];
|
||||
for (const condition of value as any[]) {
|
||||
const negatedFields: any[] = [];
|
||||
const negatedFields: Array<Record<string, any>> = [];
|
||||
for (const [subKey, subValue] of Object.entries(condition)) {
|
||||
if (subValue !== null && typeof subValue === "object") {
|
||||
for (const [op, val] of Object.entries(subValue as any)) {
|
||||
const neg = negateOp[op];
|
||||
if (neg) {
|
||||
const converted = ChromaDB.convertCondition(subKey, {
|
||||
[neg]: val,
|
||||
});
|
||||
if (converted) negatedFields.push(converted);
|
||||
}
|
||||
// Unknown operators mirror the positive-path equality
|
||||
// fallback as inequality (previously they were silently
|
||||
// dropped, which could erase the entire NOT clause).
|
||||
const neg = negateOp[op] ?? "ne";
|
||||
negatedFields.push(
|
||||
...ChromaDB.convertCondition(subKey, { [neg]: val }),
|
||||
);
|
||||
}
|
||||
} else {
|
||||
const converted = ChromaDB.convertCondition(subKey, {
|
||||
ne: subValue,
|
||||
});
|
||||
if (converted) negatedFields.push(converted);
|
||||
negatedFields.push(
|
||||
...ChromaDB.convertCondition(subKey, { ne: subValue }),
|
||||
);
|
||||
}
|
||||
}
|
||||
if (negatedFields.length > 1) {
|
||||
negatedPerGroup.push({ $or: negatedFields });
|
||||
} else if (negatedFields.length === 1) {
|
||||
negatedPerGroup.push(negatedFields[0]);
|
||||
}
|
||||
}
|
||||
if (negatedPerGroup.length > 1) {
|
||||
processed.push({ $and: negatedPerGroup });
|
||||
} else if (negatedPerGroup.length === 1) {
|
||||
processed.push(negatedPerGroup[0]);
|
||||
const combined = ChromaDB.combineClauses(negatedFields, "$or");
|
||||
if (combined) negatedPerGroup.push(combined);
|
||||
}
|
||||
const combinedNot = ChromaDB.combineClauses(negatedPerGroup, "$and");
|
||||
if (combinedNot) processed.push(combinedNot);
|
||||
} else {
|
||||
const converted = ChromaDB.convertCondition(key, value);
|
||||
if (converted) processed.push(converted);
|
||||
const combined = ChromaDB.combineClauses(
|
||||
ChromaDB.convertCondition(key, value),
|
||||
"$and",
|
||||
);
|
||||
if (combined) processed.push(combined);
|
||||
}
|
||||
}
|
||||
|
||||
if (processed.length === 0) {
|
||||
return undefined;
|
||||
}
|
||||
if (processed.length === 1) {
|
||||
return processed[0];
|
||||
}
|
||||
return { $and: processed };
|
||||
return ChromaDB.combineClauses(processed, "$and") ?? undefined;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -235,6 +235,13 @@ export class Milvus implements VectorStore {
|
||||
if (!Milvus.SAFE_FILTER_KEY.test(key)) {
|
||||
throw new Error(`Invalid filter key: ${JSON.stringify(key)}`);
|
||||
}
|
||||
if (value === "*") {
|
||||
// Wildcard - match any value. Milvus has no direct wildcard, so skip
|
||||
// the clause rather than emitting a literal `== "*"` that matches
|
||||
// nothing. Mirrors the Python provider (#6187) and the chroma/pinecone
|
||||
// stores.
|
||||
continue;
|
||||
}
|
||||
if (typeof value === "string") {
|
||||
// Escape backslashes before quotes so a value can't break out of the
|
||||
// string literal (order matters, exactly as in the Python provider).
|
||||
@@ -252,13 +259,13 @@ export class Milvus implements VectorStore {
|
||||
}
|
||||
|
||||
/**
|
||||
* Text fed to the BM25 sparse index for a payload. Prefers the lemmatized
|
||||
* text, falls back to the raw memory `data`, and truncates to the VarChar
|
||||
* limit (mirrors the Python provider).
|
||||
* Text fed to the BM25 sparse index for a payload. Prefers `textLemmatized`,
|
||||
* then `text_lemmatized`, then raw `data`; truncates to the VarChar limit.
|
||||
*/
|
||||
private bm25Text(payload?: Record<string, any>): string {
|
||||
if (!payload) return "";
|
||||
const raw = payload.text_lemmatized || payload.data || "";
|
||||
const raw =
|
||||
payload.textLemmatized || payload.text_lemmatized || payload.data || "";
|
||||
return String(raw).slice(0, 65535);
|
||||
}
|
||||
|
||||
|
||||
@@ -62,6 +62,46 @@ export class MongoDB implements VectorStore {
|
||||
return this._initPromise;
|
||||
}
|
||||
|
||||
private textSearchIndexDefinition(textIndexName: string) {
|
||||
return {
|
||||
name: textIndexName,
|
||||
definition: {
|
||||
mappings: {
|
||||
dynamic: false,
|
||||
fields: {
|
||||
payload: {
|
||||
type: "document",
|
||||
fields: {
|
||||
data: { type: "string" },
|
||||
textLemmatized: { type: "string" },
|
||||
text_lemmatized: { type: "string" },
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
private textSearchIndexMappingIsCurrent(
|
||||
index: Record<string, unknown>,
|
||||
): boolean {
|
||||
const definition =
|
||||
(index.latestDefinition as Record<string, unknown> | undefined) ??
|
||||
(index.definition as Record<string, unknown> | undefined);
|
||||
const payloadFields = (
|
||||
(definition?.mappings as Record<string, unknown> | undefined)?.fields as
|
||||
| Record<string, unknown>
|
||||
| undefined
|
||||
)?.payload as Record<string, unknown> | undefined;
|
||||
const fields =
|
||||
(payloadFields?.fields as Record<string, unknown> | undefined) ?? {};
|
||||
const textLemmatized = fields.textLemmatized as
|
||||
| { type?: string }
|
||||
| undefined;
|
||||
return textLemmatized?.type === "string";
|
||||
}
|
||||
|
||||
private async _doInitialize(): Promise<void> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
@@ -111,32 +151,24 @@ export class MongoDB implements VectorStore {
|
||||
// Create Text Search Index for keywordSearch
|
||||
const textIndexName = `${this.collectionName}_text_search_index`;
|
||||
try {
|
||||
let foundTextIndex = false;
|
||||
let existingTextIndex: Record<string, unknown> | null = null;
|
||||
try {
|
||||
const indexes = await this.collection.listSearchIndexes().toArray();
|
||||
foundTextIndex = indexes.some((idx) => idx.name === textIndexName);
|
||||
existingTextIndex =
|
||||
(indexes.find((idx) => idx.name === textIndexName) as
|
||||
| Record<string, unknown>
|
||||
| undefined) ?? null;
|
||||
} catch (e) {
|
||||
// ignore
|
||||
}
|
||||
|
||||
if (!foundTextIndex) {
|
||||
await this.collection.createSearchIndex({
|
||||
name: textIndexName,
|
||||
definition: {
|
||||
mappings: {
|
||||
dynamic: false,
|
||||
fields: {
|
||||
payload: {
|
||||
type: "document",
|
||||
fields: {
|
||||
data: { type: "string" },
|
||||
text_lemmatized: { type: "string" },
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
});
|
||||
const textSearchIndex = this.textSearchIndexDefinition(textIndexName);
|
||||
|
||||
if (!existingTextIndex) {
|
||||
await this.collection.createSearchIndex(textSearchIndex);
|
||||
} else if (!this.textSearchIndexMappingIsCurrent(existingTextIndex)) {
|
||||
await this.collection.dropSearchIndex(textIndexName);
|
||||
await this.collection.createSearchIndex(textSearchIndex);
|
||||
}
|
||||
} catch (e: any) {
|
||||
console.warn(
|
||||
@@ -286,7 +318,11 @@ export class MongoDB implements VectorStore {
|
||||
index: textIndexName,
|
||||
text: {
|
||||
query: query,
|
||||
path: ["payload.data", "payload.text_lemmatized"],
|
||||
path: [
|
||||
"payload.data",
|
||||
"payload.text_lemmatized",
|
||||
"payload.textLemmatized",
|
||||
],
|
||||
},
|
||||
},
|
||||
},
|
||||
|
||||
@@ -276,6 +276,7 @@ export class OpenSearchDB implements VectorStore {
|
||||
should: [
|
||||
{ match: { "payload.data": query } },
|
||||
{ match: { "payload.text_lemmatized": query } },
|
||||
{ match: { "payload.textLemmatized": query } },
|
||||
],
|
||||
minimum_should_match: 1,
|
||||
};
|
||||
|
||||
@@ -72,6 +72,37 @@ describe("AnthropicLLM (unit)", () => {
|
||||
expect(callArgs.tool_choice).toBeUndefined();
|
||||
});
|
||||
|
||||
// Regression: thinking-enabled models emit a thinking block before the text
|
||||
// block. Indexing content[0] threw "Unexpected response type"; the text block
|
||||
// must be found by type instead (TS parity with #6481).
|
||||
it("returns the text block when a thinking block precedes it (no tools)", async () => {
|
||||
mockCreate.mockResolvedValueOnce({
|
||||
content: [
|
||||
{ type: "thinking", thinking: "Let me reason about this." },
|
||||
{ type: "text", text: '{"facts": ["fact1"]}' },
|
||||
],
|
||||
});
|
||||
|
||||
const llm = new AnthropicLLM({ apiKey: "test-key" });
|
||||
const result = await llm.generateResponse([
|
||||
{ role: "user", content: "Hello" },
|
||||
]);
|
||||
|
||||
expect(result).toBe('{"facts": ["fact1"]}');
|
||||
});
|
||||
|
||||
// A response carrying no text block at all must resolve to "" rather than throw.
|
||||
it("returns an empty string when no text block is present (no tools)", async () => {
|
||||
mockCreate.mockResolvedValueOnce({
|
||||
content: [{ type: "thinking", thinking: "Thinking only." }],
|
||||
});
|
||||
|
||||
const llm = new AnthropicLLM({ apiKey: "test-key" });
|
||||
await expect(
|
||||
llm.generateResponse([{ role: "user", content: "Hello" }]),
|
||||
).resolves.toBe("");
|
||||
});
|
||||
|
||||
// Bug #1 regression: bare string "auto" must NOT be sent; object form required
|
||||
it("forwards tool_choice as { type: 'auto' } (not bare string) when tools are provided", async () => {
|
||||
mockCreate.mockResolvedValueOnce({
|
||||
|
||||
@@ -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"]);
|
||||
});
|
||||
});
|
||||
@@ -185,6 +185,37 @@ describe("GoogleLLM (unit)", () => {
|
||||
expect(response.toolCalls[1].name).toBe("add_graph_memory");
|
||||
});
|
||||
|
||||
// Regression: generateResponse accepted a responseFormat argument but never
|
||||
// forwarded it, so Gemini was never told to emit JSON (parity with the Python
|
||||
// SDK's mem0/llms/gemini.py, which sets response_mime_type/response_schema).
|
||||
it("forwards json_object responseFormat as responseMimeType", async () => {
|
||||
mockGenerateContent.mockResolvedValueOnce({
|
||||
text: '{"facts": ["fact1"]}',
|
||||
functionCalls: null,
|
||||
});
|
||||
|
||||
const llm = new GoogleLLM({ apiKey: "test-key" });
|
||||
await llm.generateResponse([{ role: "user", content: "Extract facts" }], {
|
||||
type: "json_object",
|
||||
});
|
||||
|
||||
const callArgs = mockGenerateContent.mock.calls[0][0];
|
||||
expect(callArgs.config.responseMimeType).toBe("application/json");
|
||||
});
|
||||
|
||||
it("does not set responseMimeType when no responseFormat is given", async () => {
|
||||
mockGenerateContent.mockResolvedValueOnce({
|
||||
text: "plain text",
|
||||
functionCalls: null,
|
||||
});
|
||||
|
||||
const llm = new GoogleLLM({ apiKey: "test-key" });
|
||||
await llm.generateResponse([{ role: "user", content: "Hello" }]);
|
||||
|
||||
const callArgs = mockGenerateContent.mock.calls[0][0];
|
||||
expect(callArgs.config.responseMimeType).toBeUndefined();
|
||||
});
|
||||
|
||||
it("formats generateChat messages and joins Gemini response parts", async () => {
|
||||
mockGenerateContent.mockResolvedValueOnce({
|
||||
candidates: [
|
||||
|
||||
@@ -153,4 +153,26 @@ describe("OpenSearchDB", () => {
|
||||
|
||||
await expect(store.get("missing")).resolves.toBeNull();
|
||||
});
|
||||
|
||||
it("keywordSearch queries lemmatized payload fields", async () => {
|
||||
const client = createClient();
|
||||
const store = await createStore(client);
|
||||
|
||||
await store.keywordSearch("stud french", 5, { user_id: "alice" });
|
||||
|
||||
const searchCall = client.search.mock.calls.find(
|
||||
([arg]: any[]) => arg.index === collectionName,
|
||||
);
|
||||
expect(searchCall).toBeDefined();
|
||||
const should = searchCall![0].body.query.bool.should;
|
||||
expect(should).toContainEqual({
|
||||
match: { "payload.textLemmatized": "stud french" },
|
||||
});
|
||||
expect(should).toContainEqual({
|
||||
match: { "payload.text_lemmatized": "stud french" },
|
||||
});
|
||||
expect(should).toContainEqual({
|
||||
match: { "payload.data": "stud french" },
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -254,110 +254,111 @@ class ChromaDB(VectorStoreBase):
|
||||
def _generate_where_clause(where: dict[str, any]) -> dict[str, any]:
|
||||
"""
|
||||
Generate a properly formatted where clause for ChromaDB.
|
||||
|
||||
|
||||
ChromaDB's where grammar allows exactly one field or one logical
|
||||
operator per dict level, so multiple operators on the same field and
|
||||
multi-field conditions must be combined with an explicit ``$and``.
|
||||
|
||||
Args:
|
||||
where (dict[str, any]): The filter conditions.
|
||||
|
||||
|
||||
Returns:
|
||||
dict[str, any]: Properly formatted where clause for ChromaDB.
|
||||
"""
|
||||
if where is None:
|
||||
return None
|
||||
|
||||
def convert_condition(key: str, value: any) -> dict:
|
||||
"""Convert universal filter format to ChromaDB format."""
|
||||
|
||||
op_map = {
|
||||
"eq": "$eq",
|
||||
"ne": "$ne",
|
||||
"gt": "$gt",
|
||||
"gte": "$gte",
|
||||
"lt": "$lt",
|
||||
"lte": "$lte",
|
||||
"in": "$in",
|
||||
"nin": "$nin",
|
||||
}
|
||||
# Negation of each operator. contains/icontains fall back to equality
|
||||
# on the positive path (ChromaDB has no substring match), so their
|
||||
# negation falls back to inequality for consistency.
|
||||
negate_map = {
|
||||
"eq": "$ne",
|
||||
"ne": "$eq",
|
||||
"gt": "$lte",
|
||||
"gte": "$lt",
|
||||
"lt": "$gte",
|
||||
"lte": "$gt",
|
||||
"in": "$nin",
|
||||
"nin": "$in",
|
||||
}
|
||||
|
||||
def convert_condition(key: str, value: any) -> list:
|
||||
"""Convert one field condition to a list of single-field ChromaDB clauses."""
|
||||
if value == "*":
|
||||
# Wildcard - match any value (ChromaDB doesn't have direct wildcard, so we skip this filter)
|
||||
return []
|
||||
if isinstance(value, dict):
|
||||
# One clause per operator: ChromaDB rejects field expressions
|
||||
# with more than one operator, so a range like
|
||||
# {"gte": 18, "lte": 65} must become two clauses combined
|
||||
# with $and by the caller (previously each operator
|
||||
# overwrote the last, silently dropping bounds).
|
||||
# contains/icontains and unknown operators fall back to equality.
|
||||
return [{key: {op_map.get(op, "$eq"): val}} for op, val in value.items()]
|
||||
# Simple equality
|
||||
return [{key: {"$eq": value}}]
|
||||
|
||||
def combine(clauses: list, operator: str):
|
||||
"""Combine clauses under a logical operator, unwrapping singletons."""
|
||||
if not clauses:
|
||||
return None
|
||||
elif isinstance(value, dict):
|
||||
# Handle comparison operators
|
||||
chroma_condition = {}
|
||||
for op, val in value.items():
|
||||
if op == "eq":
|
||||
chroma_condition[key] = {"$eq": val}
|
||||
elif op == "ne":
|
||||
chroma_condition[key] = {"$ne": val}
|
||||
elif op == "gt":
|
||||
chroma_condition[key] = {"$gt": val}
|
||||
elif op == "gte":
|
||||
chroma_condition[key] = {"$gte": val}
|
||||
elif op == "lt":
|
||||
chroma_condition[key] = {"$lt": val}
|
||||
elif op == "lte":
|
||||
chroma_condition[key] = {"$lte": val}
|
||||
elif op == "in":
|
||||
chroma_condition[key] = {"$in": val}
|
||||
elif op == "nin":
|
||||
chroma_condition[key] = {"$nin": val}
|
||||
elif op in ["contains", "icontains"]:
|
||||
# ChromaDB doesn't support contains, fallback to equality
|
||||
chroma_condition[key] = {"$eq": val}
|
||||
else:
|
||||
# Unknown operator, treat as equality
|
||||
chroma_condition[key] = {"$eq": val}
|
||||
return chroma_condition
|
||||
else:
|
||||
# Simple equality
|
||||
return {key: {"$eq": value}}
|
||||
|
||||
if len(clauses) == 1:
|
||||
return clauses[0]
|
||||
return {operator: clauses}
|
||||
|
||||
processed_filters = []
|
||||
|
||||
|
||||
for key, value in where.items():
|
||||
if key == "$or":
|
||||
# Handle OR conditions
|
||||
or_conditions = []
|
||||
for condition in value:
|
||||
or_condition = {}
|
||||
sub_clauses = []
|
||||
for sub_key, sub_value in condition.items():
|
||||
converted = convert_condition(sub_key, sub_value)
|
||||
if converted:
|
||||
or_condition.update(converted)
|
||||
if or_condition:
|
||||
or_conditions.append(or_condition)
|
||||
|
||||
if len(or_conditions) > 1:
|
||||
processed_filters.append({"$or": or_conditions})
|
||||
elif len(or_conditions) == 1:
|
||||
processed_filters.append(or_conditions[0])
|
||||
|
||||
sub_clauses.extend(convert_condition(sub_key, sub_value))
|
||||
combined = combine(sub_clauses, "$and")
|
||||
if combined:
|
||||
or_conditions.append(combined)
|
||||
combined_or = combine(or_conditions, "$or")
|
||||
if combined_or:
|
||||
processed_filters.append(combined_or)
|
||||
|
||||
elif key == "$not":
|
||||
negate_op = {
|
||||
"eq": "$ne", "ne": "$eq",
|
||||
"gt": "$lte", "gte": "$lt",
|
||||
"lt": "$gte", "lte": "$gt",
|
||||
"in": "$nin", "nin": "$in",
|
||||
}
|
||||
negated_per_group = []
|
||||
for condition in value:
|
||||
negated_fields = []
|
||||
for sub_key, sub_value in condition.items():
|
||||
if isinstance(sub_value, dict):
|
||||
for op, val in sub_value.items():
|
||||
neg = negate_op.get(op)
|
||||
if neg:
|
||||
negated_fields.append({sub_key: {neg: val}})
|
||||
# Unknown operators mirror the positive-path
|
||||
# equality fallback as inequality (previously
|
||||
# they were silently dropped, which could
|
||||
# erase the entire NOT clause).
|
||||
negated_fields.append({sub_key: {negate_map.get(op, "$ne"): val}})
|
||||
else:
|
||||
negated_fields.append({sub_key: {"$ne": sub_value}})
|
||||
if len(negated_fields) > 1:
|
||||
negated_per_group.append({"$or": negated_fields})
|
||||
elif len(negated_fields) == 1:
|
||||
negated_per_group.append(negated_fields[0])
|
||||
# NOT(a AND b) == (NOT a) OR (NOT b)
|
||||
combined = combine(negated_fields, "$or")
|
||||
if combined:
|
||||
negated_per_group.append(combined)
|
||||
combined_not = combine(negated_per_group, "$and")
|
||||
if combined_not:
|
||||
processed_filters.append(combined_not)
|
||||
|
||||
if len(negated_per_group) > 1:
|
||||
processed_filters.append({"$and": negated_per_group})
|
||||
elif len(negated_per_group) == 1:
|
||||
processed_filters.append(negated_per_group[0])
|
||||
|
||||
else:
|
||||
# Regular condition
|
||||
converted = convert_condition(key, value)
|
||||
if converted:
|
||||
processed_filters.append(converted)
|
||||
|
||||
combined = combine(convert_condition(key, value), "$and")
|
||||
if combined:
|
||||
processed_filters.append(combined)
|
||||
|
||||
# Return appropriate format based on number of conditions
|
||||
if len(processed_filters) == 0:
|
||||
return None
|
||||
elif len(processed_filters) == 1:
|
||||
return processed_filters[0]
|
||||
else:
|
||||
return {"$and": processed_filters}
|
||||
return combine(processed_filters, "$and")
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -16,6 +16,7 @@ from mem0.vector_stores.base import VectorStoreBase
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_SAFE_FILTER_KEY = re.compile(r"^[a-zA-Z_][a-zA-Z0-9_.]*$")
|
||||
_IDENTITY_FILTER_KEYS = ("user_id", "agent_id", "run_id")
|
||||
|
||||
|
||||
def _validate_filter(key: str, value) -> None:
|
||||
@@ -28,6 +29,29 @@ def _validate_filter(key: str, value) -> None:
|
||||
)
|
||||
|
||||
|
||||
def _build_filter_clauses(filters):
|
||||
"""Build term clauses from every filter key, not just the identity keys."""
|
||||
filter_clauses = []
|
||||
for key, value in (filters or {}).items():
|
||||
if value is None:
|
||||
continue
|
||||
if value == "*":
|
||||
# "Any value" wildcard (a documented Platform pattern): match
|
||||
# documents where the field exists — as opensearch.ts already
|
||||
# does for every key — instead of a literal, near-always-empty
|
||||
# term match on the string "*".
|
||||
_validate_filter(key, value)
|
||||
filter_clauses.append({"exists": {"field": f"payload.{key}"}})
|
||||
continue
|
||||
if key not in _IDENTITY_FILTER_KEYS and not isinstance(value, (str, int, float, bool)):
|
||||
logger.debug(f"Ignoring non-scalar filter value for key {key!r}")
|
||||
continue
|
||||
_validate_filter(key, value)
|
||||
field = f"payload.{key}.keyword" if isinstance(value, str) else f"payload.{key}"
|
||||
filter_clauses.append({"term": {field: value}})
|
||||
return filter_clauses
|
||||
|
||||
|
||||
class OutputData(BaseModel):
|
||||
id: str
|
||||
score: float
|
||||
@@ -204,13 +228,7 @@ class OpenSearchDB(VectorStoreBase):
|
||||
query_body = {"size": top_k * 2, "query": None}
|
||||
|
||||
# Prepare filter conditions if applicable
|
||||
filter_clauses = []
|
||||
if filters:
|
||||
for key in ["user_id", "run_id", "agent_id"]:
|
||||
value = filters.get(key)
|
||||
if value:
|
||||
_validate_filter(key, value)
|
||||
filter_clauses.append({"term": {f"payload.{key}.keyword": value}})
|
||||
filter_clauses = _build_filter_clauses(filters)
|
||||
|
||||
# Combine knn with filters if needed
|
||||
if filter_clauses:
|
||||
@@ -230,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.
|
||||
@@ -255,13 +273,7 @@ class OpenSearchDB(VectorStoreBase):
|
||||
}
|
||||
|
||||
# Apply filters consistently with the existing search() method
|
||||
filter_clauses = []
|
||||
if filters:
|
||||
for key in ["user_id", "run_id", "agent_id"]:
|
||||
value = filters.get(key)
|
||||
if value:
|
||||
_validate_filter(key, value)
|
||||
filter_clauses.append({"term": {f"payload.{key}.keyword": value}})
|
||||
filter_clauses = _build_filter_clauses(filters)
|
||||
|
||||
if filter_clauses:
|
||||
bool_query["filter"] = filter_clauses
|
||||
@@ -281,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."""
|
||||
@@ -370,13 +386,7 @@ class OpenSearchDB(VectorStoreBase):
|
||||
"""List all memories with optional filters."""
|
||||
query: Dict = {"query": {"match_all": {}}}
|
||||
|
||||
filter_clauses = []
|
||||
if filters:
|
||||
for key in ["user_id", "run_id", "agent_id"]:
|
||||
value = filters.get(key)
|
||||
if value:
|
||||
_validate_filter(key, value)
|
||||
filter_clauses.append({"term": {f"payload.{key}.keyword": value}})
|
||||
filter_clauses = _build_filter_clauses(filters)
|
||||
|
||||
if filter_clauses:
|
||||
query["query"] = {"bool": {"filter": filter_clauses}}
|
||||
|
||||
@@ -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
|
||||
@@ -459,13 +459,13 @@ class PGVector(VectorStoreBase):
|
||||
self._ensure_collection()
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(
|
||||
sql.SQL("SELECT id, vector, payload FROM {} WHERE id = %s").format(self._col()),
|
||||
sql.SQL("SELECT id, payload FROM {} WHERE id = %s").format(self._col()),
|
||||
(vector_id,),
|
||||
)
|
||||
result = cur.fetchone()
|
||||
if not result:
|
||||
return None
|
||||
return OutputData(id=str(result[0]), score=None, payload=result[2])
|
||||
return OutputData(id=str(result[0]), score=None, payload=result[1])
|
||||
|
||||
def list_cols(self) -> List[str]:
|
||||
"""
|
||||
@@ -528,7 +528,7 @@ class PGVector(VectorStoreBase):
|
||||
with self._get_cursor() as cur:
|
||||
cur.execute(
|
||||
sql.SQL("""
|
||||
SELECT id, vector, payload
|
||||
SELECT id, payload
|
||||
FROM {}
|
||||
{}
|
||||
LIMIT %s
|
||||
@@ -536,7 +536,7 @@ class PGVector(VectorStoreBase):
|
||||
(*filter_params, top_k),
|
||||
)
|
||||
results = cur.fetchall()
|
||||
return [[OutputData(id=str(r[0]), score=None, payload=r[2]) for r in results]]
|
||||
return [[OutputData(id=str(r[0]), score=None, payload=r[1]) for r in results]]
|
||||
|
||||
def __del__(self) -> None:
|
||||
"""
|
||||
|
||||
+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"]
|
||||
|
||||
|
||||
@@ -128,7 +128,10 @@ export default function ConfigurationPage() {
|
||||
<Label className="text-xs">Provider</Label>
|
||||
<Select
|
||||
value={llmProvider}
|
||||
onValueChange={setLlmProvider}
|
||||
onValueChange={(value) => {
|
||||
setLlmProvider(value);
|
||||
setLlmApiKey("");
|
||||
}}
|
||||
disabled={!isAdmin || !providers}
|
||||
>
|
||||
<SelectTrigger>
|
||||
|
||||
@@ -419,7 +419,10 @@ export default function SetupPage() {
|
||||
<Label htmlFor="setup-llm-provider">LLM Provider</Label>
|
||||
<Select
|
||||
value={llmProvider}
|
||||
onValueChange={setLlmProvider}
|
||||
onValueChange={(value) => {
|
||||
setLlmProvider(value);
|
||||
setLlmApiKey("");
|
||||
}}
|
||||
disabled={!providers}
|
||||
>
|
||||
<SelectTrigger id="setup-llm-provider">
|
||||
|
||||
@@ -344,3 +344,53 @@ def test_chroma_config_rejects_no_config():
|
||||
"""Test that ChromaDbConfig rejects when no connection config is provided."""
|
||||
with pytest.raises(ValueError):
|
||||
ChromaDbConfig()
|
||||
|
||||
|
||||
def test_generate_where_clause_same_field_range_keeps_both_bounds():
|
||||
"""A same-field range must produce both bounds, each as its own single-operator
|
||||
clause combined with $and (ChromaDB rejects multi-operator field expressions).
|
||||
|
||||
Regression test: each operator previously overwrote the previous one, so
|
||||
{"gte": 18, "lte": 65} silently degraded to {"$lte": 65} and returned rows
|
||||
the caller explicitly excluded.
|
||||
"""
|
||||
result = ChromaDB._generate_where_clause({"age": {"gte": 18, "lte": 65}})
|
||||
assert result == {"$and": [{"age": {"$gte": 18}}, {"age": {"$lte": 65}}]}
|
||||
|
||||
|
||||
def test_generate_where_clause_or_with_same_field_range():
|
||||
"""Same-field ranges inside $or branches must also keep both bounds."""
|
||||
result = ChromaDB._generate_where_clause({"$or": [{"age": {"gte": 18, "lte": 65}}, {"vip": True}]})
|
||||
assert result == {
|
||||
"$or": [
|
||||
{"$and": [{"age": {"$gte": 18}}, {"age": {"$lte": 65}}]},
|
||||
{"vip": {"$eq": True}},
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def test_generate_where_clause_or_with_multi_field_condition():
|
||||
"""Multi-field conditions inside $or must be wrapped in $and — ChromaDB
|
||||
rejects flat dicts with more than one field per level."""
|
||||
result = ChromaDB._generate_where_clause({"$or": [{"age": {"gte": 18}, "vip": True}, {"city": "sh"}]})
|
||||
assert result == {
|
||||
"$or": [
|
||||
{"$and": [{"age": {"$gte": 18}}, {"vip": {"$eq": True}}]},
|
||||
{"city": {"$eq": "sh"}},
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def test_generate_where_clause_not_contains_negates_instead_of_vanishing():
|
||||
"""$not with contains/icontains must produce a negated clause.
|
||||
|
||||
Regression test: operators missing from the negation map were silently
|
||||
dropped, which could erase the whole where clause and return unfiltered
|
||||
results. contains falls back to equality on the positive path, so its
|
||||
negation falls back to inequality.
|
||||
"""
|
||||
result = ChromaDB._generate_where_clause({"$not": [{"title": {"contains": "draft"}}]})
|
||||
assert result == {"title": {"$ne": "draft"}}
|
||||
|
||||
result = ChromaDB._generate_where_clause({"$not": [{"title": {"icontains": "draft"}}]})
|
||||
assert result == {"title": {"$ne": "draft"}}
|
||||
|
||||
@@ -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")
|
||||
@@ -594,3 +609,98 @@ class TestOpenSearchFilterValidation(unittest.TestCase):
|
||||
results = self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"user_id": "alice"})
|
||||
self.assertEqual(results, [])
|
||||
self.client_mock.search.assert_called_once()
|
||||
|
||||
|
||||
class TestOpenSearchCustomFilters(TestOpenSearchFilterValidation):
|
||||
"""Custom (non-identity) filter keys must be honored, not silently dropped."""
|
||||
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
self.client_mock.search.return_value = {"hits": {"hits": []}}
|
||||
|
||||
def _search_body(self):
|
||||
return self.client_mock.search.call_args[1]["body"]
|
||||
|
||||
def test_search_honors_custom_filter_key(self):
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"user_id": "alice", "category": "billing"})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertIn({"term": {"payload.user_id.keyword": "alice"}}, clauses)
|
||||
self.assertIn({"term": {"payload.category.keyword": "billing"}}, clauses)
|
||||
|
||||
def test_keyword_search_honors_custom_filter_key(self):
|
||||
self.os_db.keyword_search(query="report", filters={"category": "billing"})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertIn({"term": {"payload.category.keyword": "billing"}}, clauses)
|
||||
|
||||
def test_list_honors_custom_filter_key(self):
|
||||
self.os_db.list(filters={"category": "billing"})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertIn({"term": {"payload.category.keyword": "billing"}}, clauses)
|
||||
|
||||
def test_non_string_filter_values_use_plain_field(self):
|
||||
"""Non-string identity/scalar values must match against the plain payload field, not .keyword."""
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"age": 30, "archived": False})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertIn({"term": {"payload.age": 30}}, clauses)
|
||||
self.assertIn({"term": {"payload.archived": False}}, clauses)
|
||||
|
||||
def test_custom_filter_keys_still_validated(self):
|
||||
with self.assertRaises(ValueError):
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"bad key!": "x"})
|
||||
|
||||
def test_search_ignores_or_operator_filter(self):
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"user_id": "alice", "$or": [{"a": 1}]})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertEqual(clauses, [{"term": {"payload.user_id.keyword": "alice"}}])
|
||||
|
||||
def test_search_ignores_operator_shaped_filter_value(self):
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"user_id": "alice", "score": {"gte": 5}})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertEqual(clauses, [{"term": {"payload.user_id.keyword": "alice"}}])
|
||||
|
||||
def test_search_translates_wildcard_filter_value_to_exists(self):
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"user_id": "alice", "category": "*"})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertEqual(
|
||||
clauses,
|
||||
[
|
||||
{"term": {"payload.user_id.keyword": "alice"}},
|
||||
{"exists": {"field": "payload.category"}},
|
||||
],
|
||||
)
|
||||
|
||||
def test_list_ignores_or_operator_filter(self):
|
||||
self.os_db.list(filters={"user_id": "alice", "$or": [{"a": 1}]})
|
||||
self.client_mock.search.assert_called_once()
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertEqual(clauses, [{"term": {"payload.user_id.keyword": "alice"}}])
|
||||
|
||||
|
||||
class TestOpenSearchWildcardFilters(TestOpenSearchFilterValidation):
|
||||
"""The "*" wildcard means "any value" (a documented Platform pattern) and
|
||||
must build an exists query — as opensearch.ts does — for every key,
|
||||
including the identity keys, instead of a literal term match on "*"."""
|
||||
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
self.client_mock.search.return_value = {"hits": {"hits": []}}
|
||||
|
||||
def _search_body(self):
|
||||
return self.client_mock.search.call_args[1]["body"]
|
||||
|
||||
def test_identity_key_wildcard_builds_exists_query(self):
|
||||
self.os_db.search(query="", vectors=[[0.1] * 1536], filters={"agent_id": "*"})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertIn({"exists": {"field": "payload.agent_id"}}, clauses)
|
||||
self.assertNotIn({"term": {"payload.agent_id.keyword": "*"}}, clauses)
|
||||
|
||||
def test_custom_key_wildcard_builds_exists_query(self):
|
||||
self.os_db.list(filters={"category": "*"})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertIn({"exists": {"field": "payload.category"}}, clauses)
|
||||
|
||||
def test_wildcard_combines_with_term_filters(self):
|
||||
self.os_db.keyword_search(query="q", filters={"user_id": "u1", "topic": "*"})
|
||||
clauses = self._search_body()["query"]["bool"]["filter"]
|
||||
self.assertIn({"term": {"payload.user_id.keyword": "u1"}}, clauses)
|
||||
self.assertIn({"exists": {"field": "payload.topic"}}, clauses)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -711,7 +711,7 @@ class TestPGVector(unittest.TestCase):
|
||||
mock_get_cursor.return_value.__exit__.return_value = None
|
||||
|
||||
self.mock_cursor.fetchall.return_value = [] # No existing collections
|
||||
self.mock_cursor.fetchone.return_value = (self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"})
|
||||
self.mock_cursor.fetchone.return_value = (self.test_ids[0], {"key": "value1"})
|
||||
|
||||
pgvector = PGVector(
|
||||
dbname="test_db",
|
||||
@@ -734,7 +734,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify get query was executed
|
||||
get_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "SELECT id, vector, payload" in str(call)]
|
||||
if "SELECT id, payload" in str(call)]
|
||||
self.assertTrue(len(get_calls) > 0)
|
||||
|
||||
# Verify result
|
||||
@@ -756,7 +756,7 @@ class TestPGVector(unittest.TestCase):
|
||||
mock_get_cursor.return_value.__exit__.return_value = None
|
||||
|
||||
self.mock_cursor.fetchall.return_value = [] # No existing collections
|
||||
self.mock_cursor.fetchone.return_value = (self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"})
|
||||
self.mock_cursor.fetchone.return_value = (self.test_ids[0], {"key": "value1"})
|
||||
|
||||
pgvector = PGVector(
|
||||
dbname="test_db",
|
||||
@@ -779,7 +779,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify get query was executed
|
||||
get_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "SELECT id, vector, payload" in str(call)]
|
||||
if "SELECT id, payload" in str(call)]
|
||||
self.assertTrue(len(get_calls) > 0)
|
||||
|
||||
# Verify result
|
||||
@@ -1050,8 +1050,8 @@ class TestPGVector(unittest.TestCase):
|
||||
mock_get_cursor.return_value.__exit__.return_value = None
|
||||
|
||||
self.mock_cursor.fetchall.return_value = [
|
||||
(self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}),
|
||||
(self.test_ids[1], [0.4, 0.5, 0.6], {"key": "value2"}),
|
||||
(self.test_ids[0], {"key": "value1"}),
|
||||
(self.test_ids[1], {"key": "value2"}),
|
||||
]
|
||||
|
||||
pgvector = PGVector(
|
||||
@@ -1075,7 +1075,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify list query was executed
|
||||
list_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "SELECT id, vector, payload" in str(call)]
|
||||
if "SELECT id, payload" in str(call)]
|
||||
self.assertTrue(len(list_calls) > 0)
|
||||
|
||||
# Verify result
|
||||
@@ -1098,8 +1098,8 @@ class TestPGVector(unittest.TestCase):
|
||||
mock_get_cursor.return_value.__exit__.return_value = None
|
||||
|
||||
self.mock_cursor.fetchall.return_value = [
|
||||
(self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}),
|
||||
(self.test_ids[1], [0.4, 0.5, 0.6], {"key": "value2"}),
|
||||
(self.test_ids[0], {"key": "value1"}),
|
||||
(self.test_ids[1], {"key": "value2"}),
|
||||
]
|
||||
|
||||
pgvector = PGVector(
|
||||
@@ -1123,7 +1123,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify list query was executed
|
||||
list_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "SELECT id, vector, payload" in str(call)]
|
||||
if "SELECT id, payload" in str(call)]
|
||||
self.assertTrue(len(list_calls) > 0)
|
||||
|
||||
# Verify result
|
||||
@@ -1440,7 +1440,7 @@ class TestPGVector(unittest.TestCase):
|
||||
mock_get_cursor.return_value.__exit__.return_value = None
|
||||
|
||||
self.mock_cursor.fetchall.return_value = [
|
||||
(self.test_ids[0], [0.1, 0.2, 0.3], {"user_id": "alice", "agent_id": "agent1"}),
|
||||
(self.test_ids[0], {"user_id": "alice", "agent_id": "agent1"}),
|
||||
]
|
||||
|
||||
pgvector = PGVector(
|
||||
@@ -1465,7 +1465,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify list query was executed with filters
|
||||
list_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "SELECT id, vector, payload" in str(call) and "WHERE" in str(call)]
|
||||
if "SELECT id, payload" in str(call) and "WHERE" in str(call)]
|
||||
self.assertTrue(len(list_calls) > 0)
|
||||
|
||||
# Verify results
|
||||
@@ -1489,7 +1489,7 @@ class TestPGVector(unittest.TestCase):
|
||||
mock_get_cursor.return_value.__exit__.return_value = None
|
||||
|
||||
self.mock_cursor.fetchall.return_value = [
|
||||
(self.test_ids[0], [0.1, 0.2, 0.3], {"user_id": "alice", "agent_id": "agent1"}),
|
||||
(self.test_ids[0], {"user_id": "alice", "agent_id": "agent1"}),
|
||||
]
|
||||
|
||||
pgvector = PGVector(
|
||||
@@ -1514,7 +1514,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify list query was executed with filters
|
||||
list_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "SELECT id, vector, payload" in str(call) and "WHERE" in str(call)]
|
||||
if "SELECT id, payload" in str(call) and "WHERE" in str(call)]
|
||||
self.assertTrue(len(list_calls) > 0)
|
||||
|
||||
# Verify results
|
||||
@@ -1538,7 +1538,7 @@ class TestPGVector(unittest.TestCase):
|
||||
mock_get_cursor.return_value.__exit__.return_value = None
|
||||
|
||||
self.mock_cursor.fetchall.return_value = [
|
||||
(self.test_ids[0], [0.1, 0.2, 0.3], {"user_id": "alice"}),
|
||||
(self.test_ids[0], {"user_id": "alice"}),
|
||||
]
|
||||
|
||||
pgvector = PGVector(
|
||||
@@ -1563,7 +1563,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify list query was executed with single filter
|
||||
list_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "SELECT id, vector, payload" in str(call) and "WHERE" in str(call)]
|
||||
if "SELECT id, payload" in str(call) and "WHERE" in str(call)]
|
||||
self.assertTrue(len(list_calls) > 0)
|
||||
|
||||
# Verify results
|
||||
@@ -1586,7 +1586,7 @@ class TestPGVector(unittest.TestCase):
|
||||
mock_get_cursor.return_value.__exit__.return_value = None
|
||||
|
||||
self.mock_cursor.fetchall.return_value = [
|
||||
(self.test_ids[0], [0.1, 0.2, 0.3], {"user_id": "alice"}),
|
||||
(self.test_ids[0], {"user_id": "alice"}),
|
||||
]
|
||||
|
||||
pgvector = PGVector(
|
||||
@@ -1611,7 +1611,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify list query was executed with single filter
|
||||
list_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "SELECT id, vector, payload" in str(call) and "WHERE" in str(call)]
|
||||
if "SELECT id, payload" in str(call) and "WHERE" in str(call)]
|
||||
self.assertTrue(len(list_calls) > 0)
|
||||
|
||||
# Verify results
|
||||
@@ -1634,8 +1634,8 @@ class TestPGVector(unittest.TestCase):
|
||||
mock_get_cursor.return_value.__exit__.return_value = None
|
||||
|
||||
self.mock_cursor.fetchall.return_value = [
|
||||
(self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}),
|
||||
(self.test_ids[1], [0.4, 0.5, 0.6], {"key": "value2"}),
|
||||
(self.test_ids[0], {"key": "value1"}),
|
||||
(self.test_ids[1], {"key": "value2"}),
|
||||
]
|
||||
|
||||
pgvector = PGVector(
|
||||
@@ -1659,7 +1659,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify list query was executed without WHERE clause
|
||||
list_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "SELECT id, vector, payload" in str(call) and "WHERE" not in str(call)]
|
||||
if "SELECT id, payload" in str(call) and "WHERE" not in str(call)]
|
||||
self.assertTrue(len(list_calls) > 0)
|
||||
|
||||
# Verify results
|
||||
@@ -1682,8 +1682,8 @@ class TestPGVector(unittest.TestCase):
|
||||
mock_get_cursor.return_value.__exit__.return_value = None
|
||||
|
||||
self.mock_cursor.fetchall.return_value = [
|
||||
(self.test_ids[0], [0.1, 0.2, 0.3], {"key": "value1"}),
|
||||
(self.test_ids[1], [0.4, 0.5, 0.6], {"key": "value2"}),
|
||||
(self.test_ids[0], {"key": "value1"}),
|
||||
(self.test_ids[1], {"key": "value2"}),
|
||||
]
|
||||
|
||||
pgvector = PGVector(
|
||||
@@ -1707,7 +1707,7 @@ class TestPGVector(unittest.TestCase):
|
||||
|
||||
# Verify list query was executed without WHERE clause
|
||||
list_calls = [call for call in self.mock_cursor.execute.call_args_list
|
||||
if "SELECT id, vector, payload" in str(call) and "WHERE" not in str(call)]
|
||||
if "SELECT id, payload" in str(call) and "WHERE" not in str(call)]
|
||||
self.assertTrue(len(list_calls) > 0)
|
||||
|
||||
# Verify results
|
||||
|
||||
Reference in New Issue
Block a user