Compare commits

..

1 Commits

43 changed files with 232 additions and 4233 deletions
+32 -47
View File
@@ -1,56 +1,41 @@
name: Bug Report
description: Report a bug in mem0
labels: ["bug"]
name: 🐛 Bug Report
description: Create a report to help us reproduce and fix the bug
body:
- type: dropdown
id: component
attributes:
label: Component
description: Which part of mem0 is affected?
options:
- Core / Python SDK
- TypeScript SDK
- Vector Store (Qdrant, PGVector, Redis, Chroma, etc.)
- Graph Memory (Neo4j, Memgraph, etc.)
- Ollama / Local Models
- OpenMemory (MCP Server)
- OpenClaw
- REST API
- Other
validations:
required: true
- type: markdown
attributes:
value: >
#### Before submitting a bug, please make sure the issue hasn't been already addressed by searching through [the existing and past issues](https://github.com/embedchain/embedchain/issues?q=is%3Aissue+sort%3Acreated-desc+).
- type: textarea
attributes:
label: 🐛 Describe the bug
description: |
Please provide a clear and concise description of what the bug is.
- type: textarea
id: description
attributes:
label: Description
value: |
### Summary
If relevant, add a minimal example so that we can reproduce the error by running the code. It is very important for the snippet to be as succinct (minimal) as possible, so please take time to trim down any irrelevant code to help us debug efficiently. We are going to copy-paste your code and we expect to get the same result as you did: avoid any external data, and include the relevant imports, etc. For example:
A clear summary of the bug.
```python
# All necessary imports at the beginning
import embedchain as ec
# Your code goes here
### Steps to Reproduce
```python
from mem0 import Memory
```
m = Memory()
# Your code here...
```
Please also paste or describe the results you observe instead of the expected results. If you observe an error, please paste the error message including the **full** traceback of the exception. It may be relevant to wrap error messages in ```` ```triple quotes blocks``` ````.
placeholder: |
A clear and concise description of what the bug is.
### Expected Behavior
```python
Sample code to reproduce the problem
```
What you expected to happen.
### Actual Behavior
What actually happened. Paste the full error traceback if applicable.
### Environment
- mem0 version:
- Python/Node version:
- OS:
validations:
required: true
```
The error message you got, with the full traceback.
````
validations:
required: true
- type: markdown
attributes:
value: >
Thanks for contributing 🎉!
+5 -5
View File
@@ -1,8 +1,8 @@
blank_issues_enabled: true
contact_links:
- name: Discord Community
- name: 1-on-1 Session
url: https://cal.com/taranjeetio/ec
about: Speak directly with Taranjeet, the founder, to discuss issues, share feedback, or explore improvements for Embedchain
- name: Discord
url: https://discord.gg/6PzXDgEjG5
about: Ask questions and discuss with the community
- name: Documentation
url: https://docs.mem0.ai
about: Read the official mem0 documentation
about: General community discussions
+9 -21
View File
@@ -1,23 +1,11 @@
name: Documentation Issue
description: Report an issue or suggest an improvement to the mem0 docs
labels: ["documentation"]
name: Documentation
description: Report an issue related to the Embedchain docs.
title: "DOC: <Please write a comprehensive title after the 'DOC: ' prefix>"
body:
- type: textarea
id: description
attributes:
label: Description
value: |
### Page
Link to the docs page: https://docs.mem0.ai/...
### What's Wrong or Missing
Describe what's incorrect, unclear, or missing.
### Suggested Fix
How should the docs be improved?
validations:
required: true
- type: textarea
attributes:
label: "Issue with current documentation:"
description: >
Please make sure to leave a reference to the document/code you're
referring to.
+21 -40
View File
@@ -1,42 +1,23 @@
name: Feature Request
description: Suggest a new feature or improvement for mem0
labels: ["enhancement"]
name: 🚀 Feature request
description: Submit a proposal/request for a new Embedchain feature
body:
- type: dropdown
id: component
attributes:
label: Component
description: Which part of mem0 does this relate to?
options:
- Core / Python SDK
- TypeScript SDK
- Vector Store (Qdrant, PGVector, Redis, Chroma, etc.)
- Graph Memory (Neo4j, Memgraph, etc.)
- Ollama / Local Models
- OpenMemory (MCP Server)
- OpenClaw
- REST API
- Benchmarks / Evals
- Other
validations:
required: true
- type: textarea
id: description
attributes:
label: Description
value: |
### Use Case
What problem are you trying to solve?
### Proposed Solution
How should this work? Include API examples or pseudocode if helpful.
### Alternatives Considered
Any workarounds you've tried or other approaches considered.
validations:
required: true
- type: textarea
id: feature-request
attributes:
label: 🚀 The feature
description: >
A clear and concise description of the feature proposal
validations:
required: true
- type: textarea
attributes:
label: Motivation, pitch
description: >
Please outline the motivation for the proposal. Is your feature request related to a specific problem? e.g., *"I'm working on X and would like Y to be possible"*. If this is related to another GitHub issue, please link here too.
validations:
required: true
- type: markdown
attributes:
value: >
Thanks for contributing 🎉!
+28 -25
View File
@@ -1,38 +1,41 @@
## Linked Issue
Closes #<!-- issue number -->
## Description
<!-- What does this PR do? Why is it needed? -->
Please include a summary of the change and which issue is fixed. Please also include relevant motivation and context. List any dependencies that are required for this change.
## Type of Change
Fixes # (issue)
- [ ] Bug fix (non-breaking change that fixes an issue)
- [ ] New feature (non-breaking change that adds functionality)
- [ ] Breaking change (fix or feature that would cause existing functionality to change)
- [ ] Refactor (no functional changes)
## Type of change
Please delete options that are not relevant.
- [ ] Bug fix (non-breaking change which fixes an issue)
- [ ] New feature (non-breaking change which adds functionality)
- [ ] Breaking change (fix or feature that would cause existing functionality to not work as expected)
- [ ] Refactor (does not change functionality, e.g. code style improvements, linting)
- [ ] Documentation update
## Breaking Changes
## How Has This Been Tested?
<!-- If this is a breaking change, describe what breaks and the migration path. Delete this section if not applicable. -->
Please describe the tests that you ran to verify your changes. Provide instructions so we can reproduce. Please also list any relevant details for your test configuration
N/A
Please delete options that are not relevant.
## Test Coverage
- [ ] Unit Test
- [ ] Test Script (please provide)
- [ ] I added/updated unit tests
- [ ] I added/updated integration tests
- [ ] I tested manually (describe below)
- [ ] No tests needed (explain why)
## Checklist:
<!-- Describe how you tested this, or link to CI results. -->
- [ ] My code follows the style guidelines of this project
- [ ] I have performed a self-review of my own code
- [ ] I have commented my code, particularly in hard-to-understand areas
- [ ] I have made corresponding changes to the documentation
- [ ] My changes generate no new warnings
- [ ] I have added tests that prove my fix is effective or that my feature works
- [ ] New and existing unit tests pass locally with my changes
- [ ] Any dependent changes have been merged and published in downstream modules
- [ ] I have checked my code and corrected any misspellings
## Checklist
## Maintainer Checklist
- [ ] My code follows the project's style guidelines
- [ ] I have performed a self-review of my code
- [ ] I have added tests that prove my fix/feature works
- [ ] New and existing tests pass locally
- [ ] I have updated documentation if needed
- [ ] closes #xxxx (Replace xxxx with the GitHub issue number)
- [ ] Made sure Checks passed
-20
View File
@@ -1,20 +0,0 @@
# Maps dropdown selections to GitHub labels
# Used by the advanced-issue-labeler GitHub Action
component:
- label: "sdk-python"
matcher: "Core / Python SDK"
- label: "sdk-typescript"
matcher: "TypeScript SDK"
- label: "vector-store"
matcher: "Vector Store"
- label: "graph-memory"
matcher: "Graph Memory"
- label: "ollama"
matcher: "Ollama"
- label: "openmemory"
matcher: "OpenMemory"
- label: "openclaw"
matcher: "OpenClaw"
- label: "rest-api"
matcher: "REST API"
-39
View File
@@ -1,39 +0,0 @@
name: Auto-label issues
on:
issues:
types: [opened]
permissions:
contents: read
issues: write
jobs:
label:
runs-on: ubuntu-latest
steps:
- uses: stefanbuck/github-issue-parser@v3
id: issue-parser
with:
template-path: .github/ISSUE_TEMPLATE/bug_report.yml
- uses: redhat-plumbers-in-action/advanced-issue-labeler@v3
with:
issue-form: ${{ steps.issue-parser.outputs.jsonString }}
section: component
token: ${{ secrets.GITHUB_TOKEN }}
config-path: .github/advanced-issue-labeler.yml
- uses: stefanbuck/github-issue-parser@v3
id: feature-parser
if: contains(github.event.issue.labels.*.name, 'enhancement')
with:
template-path: .github/ISSUE_TEMPLATE/feature_request.yml
- uses: redhat-plumbers-in-action/advanced-issue-labeler@v3
if: contains(github.event.issue.labels.*.name, 'enhancement')
with:
issue-form: ${{ steps.feature-parser.outputs.jsonString }}
section: component
token: ${{ secrets.GITHUB_TOKEN }}
config-path: .github/advanced-issue-labeler.yml
-48
View File
@@ -1,48 +0,0 @@
name: Close stale issues
on:
schedule:
- cron: '0 0 * * *'
workflow_dispatch:
permissions:
issues: write
pull-requests: write
jobs:
stale:
runs-on: ubuntu-latest
steps:
- uses: actions/stale@v9
with:
# Issue settings
days-before-issue-stale: 90
days-before-issue-close: 14
stale-issue-label: 'stale'
stale-issue-message: >
This issue has been automatically marked as stale because it has not
had any activity in 90 days. It will be closed in 14 days if no
further activity occurs. If this is still relevant, please leave a
comment or remove the `stale` label.
close-issue-message: >
This issue has been closed due to inactivity. If this is still
relevant, feel free to reopen it or create a new issue.
# PR settings — mark stale but never auto-close
days-before-pr-stale: 90
days-before-pr-close: -1
stale-pr-label: 'stale'
stale-pr-message: >
This pull request has been automatically marked as stale because it
has not had any activity in 90 days. Please update your branch and
address any review comments, or it may be closed in the future.
# Exempt these labels from stale processing
exempt-issue-labels: 'P0-critical,P1-high,good first issue,security'
exempt-pr-labels: 'P0-critical,P1-high'
# Remove stale label when there is new activity
remove-stale-when-updated: true
# Process up to 100 issues per run to stay within API limits
operations-per-run: 100
+20 -29
View File
@@ -17,10 +17,8 @@ config = {
"workspace_url": "https://your-workspace.databricks.com",
"access_token": "your-access-token",
"endpoint_name": "your-vector-search-endpoint",
"catalog": "your_catalog",
"schema": "your_schema",
"table_name": "your_table",
"collection_name": "your_index_name",
"index_name": "catalog.schema.index_name",
"source_table_name": "catalog.schema.source_table",
"embedding_dimension": 1536
}
}
@@ -44,22 +42,17 @@ Here are the parameters available for configuring Databricks Vector Search:
| --- | --- | --- |
| `workspace_url` | The URL of your Databricks workspace | **Required** |
| `access_token` | Personal Access Token for authentication | `None` |
| `client_id` | Service principal client ID (alternative to access_token) | `None` |
| `client_secret` | Service principal client secret (required with client_id) | `None` |
| `azure_client_id` | Azure AD application client ID (for Azure Databricks) | `None` |
| `azure_client_secret` | Azure AD application client secret (for Azure Databricks) | `None` |
| `service_principal_client_id` | Service principal client ID (alternative to access_token) | `None` |
| `service_principal_client_secret` | Service principal client secret (required with client_id) | `None` |
| `endpoint_name` | Name of the Vector Search endpoint | **Required** |
| `catalog` | Unity Catalog catalog name | **Required** |
| `schema` | Unity Catalog schema name | **Required** |
| `table_name` | Source Delta table name | **Required** |
| `collection_name` | Vector search index name | `mem0` |
| `index_type` | Index type: `DELTA_SYNC` or `DIRECT_ACCESS` | `DELTA_SYNC` |
| `embedding_model_endpoint_name` | Databricks serving endpoint for embeddings | `None` |
| `index_name` | Name of the vector index (Unity Catalog format: catalog.schema.index) | **Required** |
| `source_table_name` | Name of the source Delta table (Unity Catalog format: catalog.schema.table) | **Required** |
| `embedding_dimension` | Dimension of self-managed embeddings | `1536` |
| `embedding_source_column` | Column name for text when using Databricks-computed embeddings | `None` |
| `embedding_model_endpoint_name` | Databricks serving endpoint for embeddings | `None` |
| `embedding_vector_column` | Column name for self-managed embedding vectors | `embedding` |
| `endpoint_type` | Type of endpoint (`STANDARD` or `STORAGE_OPTIMIZED`) | `STANDARD` |
| `pipeline_type` | Sync pipeline type: `TRIGGERED` or `CONTINUOUS` | `TRIGGERED` |
| `warehouse_name` | Databricks SQL warehouse name (if using SQL warehouse) | `None` |
| `query_type` | Query type: `ANN` or `HYBRID` | `ANN` |
| `sync_computed_embeddings` | Whether to sync computed embeddings automatically | `True` |
### Authentication
@@ -72,13 +65,11 @@ config = {
"provider": "databricks",
"config": {
"workspace_url": "https://your-workspace.databricks.com",
"client_id": "your-service-principal-id",
"client_secret": "your-service-principal-secret",
"service_principal_client_id": "your-service-principal-id",
"service_principal_client_secret": "your-service-principal-secret",
"endpoint_name": "your-endpoint",
"catalog": "your_catalog",
"schema": "your_schema",
"table_name": "your_table",
"collection_name": "your_index_name",
"index_name": "catalog.schema.index_name",
"source_table_name": "catalog.schema.source_table"
}
}
}
@@ -93,10 +84,8 @@ config = {
"workspace_url": "https://your-workspace.databricks.com",
"access_token": "your-personal-access-token",
"endpoint_name": "your-endpoint",
"catalog": "your_catalog",
"schema": "your_schema",
"table_name": "your_table",
"collection_name": "your_index_name",
"index_name": "catalog.schema.index_name",
"source_table_name": "catalog.schema.source_table"
}
}
}
@@ -114,6 +103,7 @@ config = {
"config": {
# ... authentication config ...
"embedding_dimension": 768, # Match your embedding model
"embedding_vector_column": "embedding"
}
}
}
@@ -128,6 +118,7 @@ config = {
"provider": "databricks",
"config": {
# ... authentication config ...
"embedding_source_column": "text",
"embedding_model_endpoint_name": "e5-small-v2"
}
}
@@ -136,8 +127,8 @@ config = {
### Important Notes
- **Index Types**: This implementation supports both `DELTA_SYNC` (auto-syncs with source Delta table) and `DIRECT_ACCESS` (manage vectors directly) index types.
- **Unity Catalog**: The source table and index are created under the specified `catalog.schema` namespace.
- **Delta Sync Index**: This implementation uses Delta Sync Index, which automatically syncs with your source Delta table. Direct vector insertion/deletion/update operations will log warnings as they're not supported with Delta Sync.
- **Unity Catalog**: Both the source table and index must be in Unity Catalog format (`catalog.schema.table_name`).
- **Endpoint Auto-Creation**: If the specified endpoint doesn't exist, it will be created automatically.
- **Index Auto-Creation**: If the specified index doesn't exist, it will be created automatically with the provided configuration.
- **Filter Support**: Supports filtering by metadata fields, with different syntax for STANDARD vs STORAGE_OPTIMIZED endpoints.
+1 -3
View File
@@ -48,7 +48,7 @@ const config = {
password: '123',
host: '127.0.0.1',
port: 5432,
dbname: 'vector_store', // Optional; TypeScript OSS defaults to `vector_store` when omitted
dbname: 'vector_store', // Optional, defaults to 'postgres'
diskann: false, // Optional, requires pgvectorscale extension
hnsw: false, // Optional, for HNSW indexing
},
@@ -85,8 +85,6 @@ Here are the parameters available for configuring pgvector:
| `connection_string` | PostgreSQL connection string (overrides individual connection parameters) | `None` |
| `connection_pool` | psycopg2 connection pool object (overrides connection string and individual parameters) | `None` |
**Note (TypeScript OSS):** If you omit `dbname`, the TypeScript client uses the database name `vector_store`. Python defaults to `postgres` for `dbname`, as in the table above.
**Note**: The connection parameters have the following priority:
1. `connection_pool` (highest priority)
2. `connection_string`
-3
View File
@@ -40,7 +40,6 @@
"icon": "rocket",
"pages": [
"platform/overview",
"vibecoding",
"platform/mem0-mcp",
"platform/platform-vs-oss",
"platform/quickstart"
@@ -143,7 +142,6 @@
"icon": "rocket",
"pages": [
"open-source/overview",
"vibecoding",
"open-source/python-quickstart",
"open-source/node-quickstart"
]
@@ -304,7 +302,6 @@
"icon": "square-terminal",
"pages": [
"openmemory/overview",
"vibecoding",
"openmemory/quickstart",
"openmemory/integrations"
]
+2 -2
View File
@@ -166,13 +166,13 @@ Call out the most common mistake or edge case for this layer.
title="[Related cookbook / deep dive]"
description="[Why this pairs well with the current guide]"
icon="arrow-right"
href="#related-link"
href="/[related-link]"
/>
<Card
title="[Next cookbook in journey]"
description="[Set expectation for the next step]"
icon="rocket"
href="#next-link"
href="/[next-link]"
/>
</CardGroup>
```
+2 -2
View File
@@ -145,13 +145,13 @@ npm install mem0ai@[version]
title="[Deep dive reference]"
description="[Why this reference matters post-migration]"
icon="book"
href="#reference-link"
href="/[reference-link]"
/>
<Card
title="[Applied example or next step]"
description="[What readers can build now]"
icon="rocket"
href="#example-link"
href="/[example-link]"
/>
</CardGroup>
```
-181
View File
@@ -1,181 +0,0 @@
---
title: "Vibecoding with Mem0"
sidebarTitle: "Vibecoding"
description: "Agent skills, starter prompts, and setup for building with Mem0 using AI coding tools."
icon: "wand-magic-sparkles"
---
These docs are designed to be easily consumable by LLMs. Each page has a button that lets you copy the page as Markdown or paste directly into ChatGPT, Claude, or any AI coding tool.
We follow the llms.txt standard:
- [llms.txt](https://docs.mem0.ai/llms.txt)
<CardGroup cols={2}>
<Card title="Get an API Key" icon="key" href="https://app.mem0.ai">
Sign up for Mem0 Platform and start building
</Card>
<Card title="Quickstart" icon="rocket" href="/platform/quickstart">
Store your first memory in under 5 minutes
</Card>
</CardGroup>
## Agent Skills
Teach your coding assistant how to build with Mem0:
```bash
npx skills add https://github.com/mem0ai/mem0 --skill mem0
```
Works with Claude Code, Cursor, Windsurf, and any assistant that supports skills. Once installed, your assistant understands Mem0's full API, framework integrations, and common patterns.
## Claude Code Plugin
The [OpenMemory plugin](https://github.com/mem0ai/claude-code-plugin) gives Claude Code **persistent memory across sessions, projects, and teams** — automatically.
<Steps>
<Step title="Get your API key">
Sign up at [app.openmemory.dev](https://app.openmemory.dev).
</Step>
<Step title="Install the plugin">
```bash
/plugin add mem0ai/claude-code-plugin
```
</Step>
<Step title="Set your environment variable">
```bash
export OPENMEMORY_API_KEY="your-key-here"
```
</Step>
<Step title="Start coding">
The plugin activates automatically. It captures decisions at session end, preserves context during compaction, and retrieves relevant memories at session start.
</Step>
</Steps>
## MCP Server Setup
Connect Cursor, Windsurf, Claude Desktop, or any MCP-compatible client to Mem0.
<Tabs>
<Tab title="OpenMemory Hosted">
Sign up at [app.openmemory.dev](https://app.openmemory.dev), then pick your client:
<CodeGroup>
```bash Claude Desktop
npx @openmemory/install --client claude --env OPENMEMORY_API_KEY=your-key
```
```bash Cursor
npx @openmemory/install --client cursor --env OPENMEMORY_API_KEY=your-key
```
```bash Windsurf
npx @openmemory/install --client windsurf --env OPENMEMORY_API_KEY=your-key
```
</CodeGroup>
For full setup options, see [OpenMemory Quickstart](/openmemory/quickstart).
</Tab>
<Tab title="Mem0 Platform MCP">
Get your API key from [app.mem0.ai](https://app.mem0.ai), then add to your MCP config:
```json
{
"mcpServers": {
"mem0": {
"command": "uvx",
"args": ["mem0-mcp-server"],
"env": {
"MEM0_API_KEY": "m0-...",
"MEM0_DEFAULT_USER_ID": "your-handle"
}
}
}
}
```
For Docker, Smithery, and advanced options, see [Mem0 MCP Setup](/platform/mem0-mcp).
</Tab>
</Tabs>
## Universal Starter Prompt
Copy this into any AI tool to start building with Mem0:
```text
I want to start building with Mem0 — a self-improving memory layer for LLM
applications that gives agents persistent context across sessions.
## Mem0 Resources
**Documentation:**
- Main docs: https://docs.mem0.ai
- Platform Quickstart: https://docs.mem0.ai/platform/quickstart
- OSS Python Quickstart: https://docs.mem0.ai/open-source/python-quickstart
- OSS Node.js Quickstart: https://docs.mem0.ai/open-source/node-quickstart
- API Reference: https://docs.mem0.ai/api-reference
- Full LLM-friendly docs: https://docs.mem0.ai/llms.txt
**Code & Examples:**
- Core repo: https://github.com/mem0ai/mem0
- Python SDK: pip install mem0ai
- TypeScript SDK: npm install mem0ai
- Cookbooks: https://docs.mem0.ai/cookbooks/overview
**What Mem0 Does:**
Mem0 is a memory layer for AI apps — managed (Mem0 Platform) or self-hosted
(Open Source). It stores, retrieves, and manages user memories so agents
remember preferences, learn from interactions, and personalize over time.
Sub-50ms retrieval. Dual storage: vector embeddings + graph databases.
**Architecture Overview:**
- Memory is scoped by user_id, agent_id, or run_id
- Core operations: add, search, update, delete
- Memory types: factual (preferences, facts), episodic (past interactions),
semantic (concept relationships), working (session state)
- Integration pattern: retrieve relevant memories → generate response → store
new memories
**Quick Usage (Python Platform):**
from mem0 import MemoryClient
client = MemoryClient(api_key="m0-xxx")
client.add("I prefer dark mode and use VS Code.", user_id="user1")
results = client.search("What editor do they use?", user_id="user1")
**Quick Usage (JavaScript Platform):**
import MemoryClient from 'mem0ai';
const client = new MemoryClient({ apiKey: 'm0-xxx' });
await client.add([{ role: "user", content: "I prefer dark mode." }], { user_id: "user1" });
const results = await client.search("What editor?", { user_id: "user1" });
**Quick Usage (Python Open Source):**
from mem0 import Memory
m = Memory()
m.add("I prefer dark mode and use VS Code.", user_id="user1")
results = m.search("What editor do they use?", user_id="user1")
Help me integrate Mem0 into my project. Start by asking what I'm building,
what language/framework I'm using, and whether I want managed or self-hosted.
```
## Go Deeper
<CardGroup cols={2}>
<Card title="Platform Quickstart" icon="cloud" href="/platform/quickstart">
Get started with the managed API
</Card>
<Card title="Open Source" icon="code-branch" href="/open-source/overview">
Self-host with full control
</Card>
<Card title="Cookbooks" icon="book" href="/cookbooks/overview">
Production-ready tutorials and examples
</Card>
<Card title="API Reference" icon="code" href="/api-reference">
Explore every REST endpoint
</Card>
</CardGroup>
-1
View File
@@ -31,7 +31,6 @@ export const MemoryUpdateSchema = z.object({
old_memory: z
.string()
.optional()
.nullable()
.describe(
"The previous content of the memory item if the event was UPDATE.",
),
@@ -1,4 +1,4 @@
import { Client } from "pg";
import { Client, Pool } from "pg";
import { VectorStore } from "./base";
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
@@ -35,7 +35,6 @@ export class PGVector implements VectorStore {
host: config.host,
port: config.port,
});
this.initialize().catch(console.error);
}
async initialize(): Promise<void> {
@@ -118,11 +118,6 @@ jest.mock("../src/vector_stores/azure_ai_search", () => ({
.fn()
.mockImplementation((config) => ({ type: "azure-ai-search", config })),
}));
jest.mock("../src/vector_stores/pgvector", () => ({
PGVector: jest
.fn()
.mockImplementation((config) => ({ type: "pgvector", config })),
}));
jest.mock("../src/storage/SupabaseHistoryManager", () => ({
SupabaseHistoryManager: jest
.fn()
@@ -241,7 +236,6 @@ describe("VectorStoreFactory", () => {
["langchain"],
["vectorize"],
["azure-ai-search"],
["pgvector"],
])("creates vector store for provider '%s'", (provider) => {
expect(() =>
VectorStoreFactory.create(provider, dummyVSConfig),
+2 -2
View File
@@ -16,7 +16,7 @@ class AnthropicConfig(BaseLlmConfig):
temperature: float = 0.1,
api_key: Optional[str] = None,
max_tokens: int = 2000,
top_p: Optional[float] = None,
top_p: float = 0.1,
top_k: int = 1,
enable_vision: bool = False,
vision_details: Optional[str] = "auto",
@@ -32,7 +32,7 @@ class AnthropicConfig(BaseLlmConfig):
temperature: Controls randomness, defaults to 0.1
api_key: Anthropic API key, defaults to None
max_tokens: Maximum tokens to generate, defaults to 2000
top_p: Nucleus sampling parameter, defaults to None (omitted to avoid conflict with temperature)
top_p: Nucleus sampling parameter, defaults to 0.1
top_k: Top-k sampling parameter, defaults to 1
enable_vision: Enable vision capabilities, defaults to False
vision_details: Vision detail level, defaults to "auto"
+2 -2
View File
@@ -32,8 +32,8 @@ class ChromaDbConfig(BaseModel):
values.pop("path", None)
return values
# Check if local/server configuration is provided
local_config = bool(path) or bool(host and port)
# Check if local/server configuration is provided (excluding default tmp path for cloud config)
local_config = bool(path and path != "/tmp/chroma") or bool(host and port)
if not cloud_config and not local_config:
raise ValueError("Either ChromaDB Cloud configuration (api_key, tenant) or local configuration (path or host/port) must be provided.")
+1 -1
View File
@@ -16,7 +16,7 @@ class QdrantConfig(BaseModel):
path: Optional[str] = Field("/tmp/qdrant", description="Path for local Qdrant database")
url: Optional[str] = Field(None, description="Full URL for Qdrant server")
api_key: Optional[str] = Field(None, description="API key for Qdrant server")
on_disk: Optional[bool] = Field(False,description="Enables persistent storage. Vectors are kept on disk (True) or in memory (False). Does not delete the local database path.")
on_disk: Optional[bool] = Field(False, description="Enables persistent storage")
@model_validator(mode="before")
@classmethod
-22
View File
@@ -409,28 +409,6 @@ class NeptuneBase(ABC):
"""
pass
def delete(self, data, filters):
"""
Delete graph entities associated with the given memory text.
Extracts entities and relationships from the memory text using the same
pipeline as add(), then deletes the matching relationships in the graph.
Args:
data (str): The memory text whose graph entities should be removed.
filters (dict): Scope filters (user_id, agent_id, run_id).
"""
try:
entity_type_map = self._retrieve_nodes_from_data(data, filters)
if not entity_type_map:
logger.debug("No entities found in memory text, skipping graph cleanup")
return
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
if to_be_deleted:
self._delete_entities(to_be_deleted, filters["user_id"])
except Exception as e:
logger.error(f"Error during graph cleanup for memory delete: {e}")
def delete_all(self, filters):
cypher, params = self._delete_all_cypher(filters)
self.graph.query(cypher, params=params)
-25
View File
@@ -40,31 +40,6 @@ class AnthropicLLM(LLMBase):
api_key = self.config.api_key or os.getenv("ANTHROPIC_API_KEY")
self.client = anthropic.Anthropic(api_key=api_key)
def _get_common_params(self, **kwargs) -> Dict:
"""Get common parameters, avoiding sending both temperature and top_p together.
Anthropic rejects requests that include both temperature and top_p.
When both are set, we keep temperature and drop top_p.
"""
params = {}
if self.config.max_tokens is not None:
params["max_tokens"] = self.config.max_tokens
has_temperature = self.config.temperature is not None
has_top_p = self.config.top_p is not None
if has_temperature and has_top_p:
# Anthropic forbids both; prefer temperature
params["temperature"] = self.config.temperature
elif has_temperature:
params["temperature"] = self.config.temperature
elif has_top_p:
params["top_p"] = self.config.top_p
params.update(kwargs)
return params
def generate_response(
self,
messages: List[Dict[str, str]],
-22
View File
@@ -246,28 +246,6 @@ class MemoryGraph:
logger.info(f"Returned {len(search_results)} search results")
return search_results
def delete(self, data, filters):
"""
Delete graph entities associated with the given memory text.
Extracts entities and relationships from the memory text using the same
pipeline as add(), then deletes the matching relationships in the graph.
Args:
data (str): The memory text whose graph entities should be removed.
filters (dict): Scope filters (user_id, agent_id, run_id).
"""
try:
entity_type_map = self._retrieve_nodes_from_data(data, filters)
if not entity_type_map:
logger.debug("No entities found in memory text, skipping graph cleanup")
return
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
if to_be_deleted:
self._delete_entities(to_be_deleted, filters)
except Exception as e:
logger.error(f"Error during graph cleanup for memory delete: {e}")
def delete_all(self, filters):
"""Delete all nodes and relationships for a user or specific agent."""
where_parts = ["n.user_id = %s"]
-22
View File
@@ -129,28 +129,6 @@ class MemoryGraph:
return search_results
def delete(self, data, filters):
"""
Delete graph entities associated with the given memory text.
Extracts entities and relationships from the memory text using the same
pipeline as add(), then soft-deletes the matching relationships in the graph.
Args:
data (str): The memory text whose graph entities should be removed.
filters (dict): Scope filters (user_id, agent_id, run_id).
"""
try:
entity_type_map = self._retrieve_nodes_from_data(data, filters)
if not entity_type_map:
logger.debug("No entities found in memory text, skipping graph cleanup")
return
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
if to_be_deleted:
self._delete_entities(to_be_deleted, filters)
except Exception as e:
logger.error(f"Error during graph cleanup for memory delete: {e}")
def delete_all(self, filters):
# Build node properties for filtering
node_props = ["user_id: $user_id"]
-22
View File
@@ -149,28 +149,6 @@ class MemoryGraph:
return search_results
def delete(self, data, filters):
"""
Delete graph entities associated with the given memory text.
Extracts entities and relationships from the memory text using the same
pipeline as add(), then deletes the matching relationships in the graph.
Args:
data (str): The memory text whose graph entities should be removed.
filters (dict): Scope filters (user_id, agent_id, run_id).
"""
try:
entity_type_map = self._retrieve_nodes_from_data(data, filters)
if not entity_type_map:
logger.debug("No entities found in memory text, skipping graph cleanup")
return
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
if to_be_deleted:
self._delete_entities(to_be_deleted, filters)
except Exception as e:
logger.error(f"Error during graph cleanup for memory delete: {e}")
def delete_all(self, filters):
# Build node properties for filtering
node_props = ["user_id: $user_id"]
+28 -72
View File
@@ -1101,27 +1101,7 @@ class Memory(MemoryBase):
memory_id (str): ID of the memory to delete.
"""
capture_event("mem0.delete", self, {"memory_id": memory_id, "sync_type": "sync"})
existing_memory = self.vector_store.get(vector_id=memory_id)
if existing_memory is None:
raise ValueError(f"Memory with id {memory_id} not found")
# Clean up graph entities before deleting from vector store
if self.enable_graph:
try:
memory_text = existing_memory.payload.get("data", "")
if memory_text:
filters = {}
for key in ("user_id", "agent_id", "run_id"):
val = existing_memory.payload.get(key)
if val:
filters[key] = val
if filters.get("user_id"):
self.graph.delete(memory_text, filters)
except Exception as e:
logger.error(f"Error cleaning up graph for memory {memory_id}: {e}")
self._delete_memory(memory_id, existing_memory)
self._delete_memory(memory_id)
return {"message": "Memory deleted successfully!"}
def delete_all(self, user_id: Optional[str] = None, agent_id: Optional[str] = None, run_id: Optional[str] = None):
@@ -1180,24 +1160,24 @@ class Memory(MemoryBase):
else:
embeddings = self.embedding_model.embed(data, memory_action="add")
memory_id = str(uuid.uuid4())
new_metadata = deepcopy(metadata) if metadata is not None else {}
new_metadata["data"] = data
new_metadata["hash"] = hashlib.md5(data.encode()).hexdigest()
new_metadata["created_at"] = datetime.now(timezone.utc).isoformat()
metadata = metadata or {}
metadata["data"] = data
metadata["hash"] = hashlib.md5(data.encode()).hexdigest()
metadata["created_at"] = datetime.now(timezone.utc).isoformat()
self.vector_store.insert(
vectors=[embeddings],
ids=[memory_id],
payloads=[new_metadata],
payloads=[metadata],
)
self.db.add_history(
memory_id,
None,
data,
"ADD",
created_at=new_metadata.get("created_at"),
actor_id=new_metadata.get("actor_id"),
role=new_metadata.get("role"),
created_at=metadata.get("created_at"),
actor_id=metadata.get("actor_id"),
role=metadata.get("role"),
)
return memory_id
@@ -1231,10 +1211,9 @@ class Memory(MemoryBase):
if metadata is None:
raise ValueError("Metadata cannot be done for procedural memory.")
new_metadata = deepcopy(metadata)
new_metadata["memory_type"] = MemoryType.PROCEDURAL.value
metadata["memory_type"] = MemoryType.PROCEDURAL.value
embeddings = self.embedding_model.embed(procedural_memory, memory_action="add")
memory_id = self._create_memory(procedural_memory, {procedural_memory: embeddings}, metadata=new_metadata)
memory_id = self._create_memory(procedural_memory, {procedural_memory: embeddings}, metadata=metadata)
capture_event("mem0._create_procedural_memory", self, {"memory_id": memory_id, "sync_type": "sync"})
result = {"results": [{"id": memory_id, "memory": procedural_memory, "event": "ADD"}]}
@@ -1298,12 +1277,11 @@ class Memory(MemoryBase):
)
return memory_id
def _delete_memory(self, memory_id, existing_memory=None):
def _delete_memory(self, memory_id):
logger.info(f"Deleting memory with {memory_id=}")
existing_memory = self.vector_store.get(vector_id=memory_id)
if existing_memory is None:
existing_memory = self.vector_store.get(vector_id=memory_id)
if existing_memory is None:
raise ValueError(f"Memory with id {memory_id} not found")
raise ValueError(f"Memory with id {memory_id} not found")
prev_value = existing_memory.payload.get("data", "")
self.vector_store.delete(vector_id=memory_id)
self.db.add_history(
@@ -2196,27 +2174,7 @@ class AsyncMemory(MemoryBase):
memory_id (str): ID of the memory to delete.
"""
capture_event("mem0.delete", self, {"memory_id": memory_id, "sync_type": "async"})
existing_memory = await asyncio.to_thread(self.vector_store.get, vector_id=memory_id)
if existing_memory is None:
raise ValueError(f"Memory with id {memory_id} not found")
# Clean up graph entities before deleting from vector store
if self.enable_graph:
try:
memory_text = existing_memory.payload.get("data", "")
if memory_text:
filters = {}
for key in ("user_id", "agent_id", "run_id"):
val = existing_memory.payload.get(key)
if val:
filters[key] = val
if filters.get("user_id"):
await asyncio.to_thread(self.graph.delete, memory_text, filters)
except Exception as e:
logger.error(f"Error cleaning up graph for memory {memory_id}: {e}")
await self._delete_memory(memory_id, existing_memory)
await self._delete_memory(memory_id)
return {"message": "Memory deleted successfully!"}
async def delete_all(self, user_id=None, agent_id=None, run_id=None):
@@ -2279,16 +2237,16 @@ class AsyncMemory(MemoryBase):
embeddings = await asyncio.to_thread(self.embedding_model.embed, data, memory_action="add")
memory_id = str(uuid.uuid4())
new_metadata = deepcopy(metadata) if metadata is not None else {}
new_metadata["data"] = data
new_metadata["hash"] = hashlib.md5(data.encode()).hexdigest()
new_metadata["created_at"] = datetime.now(timezone.utc).isoformat()
metadata = metadata or {}
metadata["data"] = data
metadata["hash"] = hashlib.md5(data.encode()).hexdigest()
metadata["created_at"] = datetime.now(timezone.utc).isoformat()
await asyncio.to_thread(
self.vector_store.insert,
vectors=[embeddings],
ids=[memory_id],
payloads=[new_metadata],
payloads=[metadata],
)
await asyncio.to_thread(
@@ -2297,9 +2255,9 @@ class AsyncMemory(MemoryBase):
None,
data,
"ADD",
created_at=new_metadata.get("created_at"),
actor_id=new_metadata.get("actor_id"),
role=new_metadata.get("role"),
created_at=metadata.get("created_at"),
actor_id=metadata.get("actor_id"),
role=metadata.get("role"),
)
return memory_id
@@ -2348,10 +2306,9 @@ class AsyncMemory(MemoryBase):
if metadata is None:
raise ValueError("Metadata cannot be done for procedural memory.")
new_metadata = deepcopy(metadata)
new_metadata["memory_type"] = MemoryType.PROCEDURAL.value
metadata["memory_type"] = MemoryType.PROCEDURAL.value
embeddings = await asyncio.to_thread(self.embedding_model.embed, procedural_memory, memory_action="add")
memory_id = await self._create_memory(procedural_memory, {procedural_memory: embeddings}, metadata=new_metadata)
memory_id = await self._create_memory(procedural_memory, {procedural_memory: embeddings}, metadata=metadata)
capture_event("mem0._create_procedural_memory", self, {"memory_id": memory_id, "sync_type": "async"})
result = {"results": [{"id": memory_id, "memory": procedural_memory, "event": "ADD"}]}
@@ -2418,12 +2375,11 @@ class AsyncMemory(MemoryBase):
)
return memory_id
async def _delete_memory(self, memory_id, existing_memory=None):
async def _delete_memory(self, memory_id):
logger.info(f"Deleting memory with {memory_id=}")
existing_memory = await asyncio.to_thread(self.vector_store.get, vector_id=memory_id)
if existing_memory is None:
existing_memory = await asyncio.to_thread(self.vector_store.get, vector_id=memory_id)
if existing_memory is None:
raise ValueError(f"Memory with id {memory_id} not found")
raise ValueError(f"Memory with id {memory_id} not found")
prev_value = existing_memory.payload.get("data", "")
await asyncio.to_thread(self.vector_store.delete, vector_id=memory_id)
-22
View File
@@ -134,28 +134,6 @@ class MemoryGraph:
return search_results
def delete(self, data, filters):
"""
Delete graph entities associated with the given memory text.
Extracts entities and relationships from the memory text using the same
pipeline as add(), then deletes the matching relationships in the graph.
Args:
data (str): The memory text whose graph entities should be removed.
filters (dict): Scope filters (user_id, agent_id).
"""
try:
entity_type_map = self._retrieve_nodes_from_data(data, filters)
if not entity_type_map:
logger.debug("No entities found in memory text, skipping graph cleanup")
return
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
if to_be_deleted:
self._delete_entities(to_be_deleted, filters)
except Exception as e:
logger.error(f"Error during graph cleanup for memory delete: {e}")
def delete_all(self, filters):
"""Delete all nodes and relationships for a user or specific agent."""
if filters.get("agent_id"):
+42 -60
View File
@@ -65,7 +65,7 @@ class Databricks(VectorStoreBase):
catalog (str): Unity Catalog catalog name.
schema (str): Unity Catalog schema name.
table_name (str): Source Delta table name.
collection_name (str, optional): Vector search index name (default: "mem0").
index_name (str, optional): Vector search index name (default: "mem0").
index_type (str, optional): Index type, either "DELTA_SYNC" or "DIRECT_ACCESS" (default: "DELTA_SYNC").
embedding_model_endpoint_name (str, optional): Embedding model endpoint for Databricks-computed embeddings.
embedding_dimension (int, optional): Vector embedding dimensions (default: 1536).
@@ -85,7 +85,7 @@ class Databricks(VectorStoreBase):
self.fully_qualified_index_name = f"{self.catalog}.{self.schema}.{self.index_name}"
# Configuration
self.index_type = VectorIndexType(index_type) if isinstance(index_type, str) else index_type
self.index_type = index_type
self.embedding_model_endpoint_name = embedding_model_endpoint_name
self.embedding_dimension = embedding_dimension
self.endpoint_type = endpoint_type
@@ -261,11 +261,11 @@ class Databricks(VectorStoreBase):
)
logger.info(f"Successfully created source table '{self.fully_qualified_table_name}'")
self.client.table_constraints.create(
full_name_arg=self.fully_qualified_table_name,
full_name_arg="logistics_dev.ai.dev_memory",
constraint=TableConstraint(
primary_key_constraint=PrimaryKeyConstraint(
name=f"pk_{self.table_name}",
child_columns=["memory_id"],
name="pk_dev_memory", # Name of the primary key constraint
child_columns=["memory_id"], # Columns that make up the primary key
)
),
)
@@ -439,29 +439,29 @@ class Databricks(VectorStoreBase):
try:
filters_json = json.dumps(filters) if filters else None
# Choose query mode per Databricks SDK contract:
# - query_text: for Delta Sync Index with model endpoint
# - query_vector: for Direct Access Index and Delta Sync Index with self-managed vectors
query_kwargs = {
"index_name": self.fully_qualified_index_name,
"columns": self.column_names,
"num_results": limit,
"query_type": self.query_type,
"filters_json": filters_json,
}
uses_model_endpoint = (
self.index_type == VectorIndexType.DELTA_SYNC and self.embedding_model_endpoint_name
)
if uses_model_endpoint:
if not query:
raise ValueError("Query text is required for Delta Sync Index with model endpoint.")
query_kwargs["query_text"] = query
elif vectors:
query_kwargs["query_vector"] = vectors
# Choose query type
if self.index_type == VectorIndexType.DELTA_SYNC and query:
# Text-based search
sdk_results = self.client.vector_search_indexes.query_index(
index_name=self.fully_qualified_index_name,
columns=self.column_names,
query_text=query,
num_results=limit,
query_type=self.query_type,
filters_json=filters_json,
)
elif self.index_type == VectorIndexType.DIRECT_ACCESS and vectors:
# Vector-based search
sdk_results = self.client.vector_search_indexes.query_index(
index_name=self.fully_qualified_index_name,
columns=self.column_names,
query_vector=vectors,
num_results=limit,
query_type=self.query_type,
filters_json=filters_json,
)
else:
raise ValueError("Must provide vectors for search.")
sdk_results = self.client.vector_search_indexes.query_index(**query_kwargs)
raise ValueError("Must provide query text for DELTA_SYNC or vectors for DIRECT_ACCESS.")
# Parse results
result_data = sdk_results.result if hasattr(sdk_results, "result") else sdk_results
@@ -572,23 +572,14 @@ class Databricks(VectorStoreBase):
filters = {"memory_id": vector_id}
filters_json = json.dumps(filters)
# Use query_text for Delta Sync with model endpoint, query_vector otherwise
query_kwargs = {
"index_name": self.fully_qualified_index_name,
"columns": self.column_names,
"num_results": 1,
"query_type": self.query_type,
"filters_json": filters_json,
}
uses_model_endpoint = (
self.index_type == VectorIndexType.DELTA_SYNC and self.embedding_model_endpoint_name
results = self.client.vector_search_indexes.query_index(
index_name=self.fully_qualified_index_name,
columns=self.column_names,
query_text=" ", # Empty query, rely on filters
num_results=1,
query_type=self.query_type,
filters_json=filters_json,
)
if uses_model_endpoint:
query_kwargs["query_text"] = " "
else:
query_kwargs["query_vector"] = [0.0] * self.embedding_dimension
results = self.client.vector_search_indexes.query_index(**query_kwargs)
# Process results
result_data = results.result if hasattr(results, "result") else results
@@ -598,7 +589,7 @@ class Databricks(VectorStoreBase):
raise KeyError(f"Vector with ID {vector_id} not found")
result = data_array[0]
columns = [col.name for col in results.manifest.columns] if results.manifest and results.manifest.columns else []
columns = columns = [col.name for col in results.manifest.columns] if results.manifest and results.manifest.columns else []
row_data = dict(zip(columns, result))
# Build payload following the standard schema
@@ -695,23 +686,14 @@ class Databricks(VectorStoreBase):
filters_json = json.dumps(filters) if filters else None
num_results = limit or 100
columns = self.column_names
# Use query_text for Delta Sync with model endpoint, query_vector otherwise
query_kwargs = {
"index_name": self.fully_qualified_index_name,
"columns": columns,
"num_results": num_results,
"query_type": self.query_type,
"filters_json": filters_json,
}
uses_model_endpoint = (
self.index_type == VectorIndexType.DELTA_SYNC and self.embedding_model_endpoint_name
sdk_results = self.client.vector_search_indexes.query_index(
index_name=self.fully_qualified_index_name,
columns=columns,
query_text=" ",
num_results=num_results,
query_type=self.query_type,
filters_json=filters_json,
)
if uses_model_endpoint:
query_kwargs["query_text"] = " "
else:
query_kwargs["query_vector"] = [0.0] * self.embedding_dimension
sdk_results = self.client.vector_search_indexes.query_index(**query_kwargs)
result_data = sdk_results.result if hasattr(sdk_results, "result") else sdk_results
data_array = result_data.data_array if hasattr(result_data, "data_array") else []
+13 -13
View File
@@ -26,7 +26,7 @@ class OutputData(BaseModel):
class MongoDB(VectorStoreBase):
VECTOR_TYPE = "vector"
VECTOR_TYPE = "knnVector"
SIMILARITY_METRIC = "cosine"
def __init__(self, db_name: str, collection_name: str, embedding_model_dims: int, mongo_uri: str):
@@ -69,16 +69,17 @@ class MongoDB(VectorStoreBase):
else:
search_index_model = SearchIndexModel(
name=self.index_name,
type="vectorSearch",
definition={
"fields": [
{
"type": self.VECTOR_TYPE,
"path": "embedding",
"numDimensions": self.embedding_model_dims,
"similarity": self.SIMILARITY_METRIC,
}
]
"mappings": {
"dynamic": False,
"fields": {
"embedding": {
"type": self.VECTOR_TYPE,
"dimensions": self.embedding_model_dims,
"similarity": self.SIMILARITY_METRIC,
}
},
}
},
)
collection.create_search_index(search_index_model)
@@ -140,7 +141,7 @@ class MongoDB(VectorStoreBase):
"$vectorSearch": {
"index": self.index_name,
"limit": limit,
"numCandidates": min(limit * 20, 10000),
"numCandidates": limit,
"queryVector": vectors,
"path": "embedding",
}
@@ -197,8 +198,7 @@ class MongoDB(VectorStoreBase):
if vector is not None:
update_fields["embedding"] = vector
if payload is not None:
for key, value in payload.items():
update_fields[f"payload.{key}"] = value
update_fields["payload"] = payload
if update_fields:
try:
+6 -2
View File
@@ -1,4 +1,6 @@
import logging
import os
import shutil
from typing import Optional
from qdrant_client import QdrantClient
@@ -46,8 +48,7 @@ class Qdrant(VectorStoreBase):
path (str, optional): Path for local Qdrant database. Defaults to None.
url (str, optional): Full URL for Qdrant server. Defaults to None.
api_key (str, optional): API key for Qdrant server. Defaults to None.
on_disk (bool, optional): Enables persistent storage. Vectors are stored on disk (True) or in memory (False).
Does not delete the local database path. Defaults to False.
on_disk (bool, optional): Enables persistent storage. Defaults to False.
"""
if client:
self.client = client
@@ -65,6 +66,9 @@ class Qdrant(VectorStoreBase):
if not params:
params["path"] = path
self.is_local = True
if not on_disk:
if os.path.exists(path) and os.path.isdir(path):
shutil.rmtree(path)
else:
self.is_local = False
+1 -1
View File
@@ -184,7 +184,7 @@ async def search_memory(query: str) -> str:
for h in hits:
# All vector db search functions return OutputData class
id, score, payload = h.id, h.score, h.payload
if allowed and (h.id is None or h.id not in allowed):
if allowed and h.id is None or h.id not in allowed:
continue
results.append({
-5
View File
@@ -119,9 +119,6 @@ class MemoryCreate(BaseModel):
agent_id: Optional[str] = None
run_id: Optional[str] = None
metadata: Optional[Dict[str, Any]] = None
infer: Optional[bool] = Field(None, description="Whether to extract facts from messages. Defaults to True.")
memory_type: Optional[str] = Field(None, description="Type of memory to store (e.g. 'core').")
prompt: Optional[str] = Field(None, description="Custom prompt to use for fact extraction.")
class SearchRequest(BaseModel):
@@ -130,8 +127,6 @@ class SearchRequest(BaseModel):
run_id: Optional[str] = None
agent_id: Optional[str] = None
filters: Optional[Dict[str, Any]] = None
limit: Optional[int] = Field(None, description="Maximum number of results to return.")
threshold: Optional[float] = Field(None, description="Minimum similarity score for results.")
@app.post("/configure", summary="Configure Mem0")
-100
View File
@@ -1,100 +0,0 @@
from unittest.mock import Mock, patch
import pytest
pytest.importorskip("anthropic", reason="anthropic package not installed")
from mem0.configs.llms.anthropic import AnthropicConfig
from mem0.configs.llms.base import BaseLlmConfig
from mem0.llms.anthropic import AnthropicLLM
@pytest.fixture
def mock_anthropic_client():
with patch("mem0.llms.anthropic.anthropic") as mock_anthropic:
mock_client = Mock()
mock_anthropic.Anthropic.return_value = mock_client
yield mock_client
def test_default_config_omits_top_p(mock_anthropic_client):
"""Default AnthropicConfig should not set top_p to avoid conflict with temperature."""
config = AnthropicConfig(model="claude-3-5-sonnet-20240620", api_key="test-key")
assert config.top_p is None
assert config.temperature == 0.1
def test_generate_response_does_not_send_top_p_by_default(mock_anthropic_client):
"""Anthropic API rejects temperature and top_p together; top_p must be omitted by default."""
config = AnthropicConfig(model="claude-3-5-sonnet-20240620", api_key="test-key")
llm = AnthropicLLM(config)
mock_response = Mock()
mock_response.content = [Mock(text="Hello!")]
mock_anthropic_client.messages.create.return_value = mock_response
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hi"},
]
llm.generate_response(messages)
call_kwargs = mock_anthropic_client.messages.create.call_args[1]
assert "top_p" not in call_kwargs
assert call_kwargs["temperature"] == 0.1
def test_generate_response_sends_top_p_alone_when_no_temperature(mock_anthropic_client):
"""When user sets only top_p (no temperature), top_p should be sent."""
config = AnthropicConfig(model="claude-3-5-sonnet-20240620", api_key="test-key", top_p=0.9, temperature=None)
llm = AnthropicLLM(config)
mock_response = Mock()
mock_response.content = [Mock(text="Hello!")]
mock_anthropic_client.messages.create.return_value = mock_response
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hi"},
]
llm.generate_response(messages)
call_kwargs = mock_anthropic_client.messages.create.call_args[1]
assert call_kwargs["top_p"] == 0.9
assert "temperature" not in call_kwargs
def test_both_set_prefers_temperature_over_top_p(mock_anthropic_client):
"""When both temperature and top_p are set, temperature wins and top_p is dropped."""
config = AnthropicConfig(model="claude-3-5-sonnet-20240620", api_key="test-key", top_p=0.9, temperature=0.5)
llm = AnthropicLLM(config)
mock_response = Mock()
mock_response.content = [Mock(text="Hello!")]
mock_anthropic_client.messages.create.return_value = mock_response
messages = [{"role": "user", "content": "Hi"}]
llm.generate_response(messages)
call_kwargs = mock_anthropic_client.messages.create.call_args[1]
assert call_kwargs["temperature"] == 0.5
assert "top_p" not in call_kwargs
def test_base_config_conversion_does_not_send_both(mock_anthropic_client):
"""BaseLlmConfig defaults both temperature=0.1 and top_p=0.1; Anthropic must not send both."""
base_config = BaseLlmConfig(model="claude-3-5-sonnet-20240620", api_key="test-key")
llm = AnthropicLLM(base_config)
mock_response = Mock()
mock_response.content = [Mock(text="Hello!")]
mock_anthropic_client.messages.create.return_value = mock_response
messages = [{"role": "user", "content": "Hi"}]
llm.generate_response(messages)
call_kwargs = mock_anthropic_client.messages.create.call_args[1]
assert "temperature" in call_kwargs
assert "top_p" not in call_kwargs
-195
View File
@@ -188,201 +188,6 @@ async def test_async_update_memory_uses_utc_timestamps(mocker):
_assert_utc_timestamp(payload["updated_at"])
class TestMetadataNotMutated:
"""Tests that metadata dicts passed to memory methods are not mutated in-place (issue #2648)."""
def test_create_memory_does_not_mutate_metadata(self, mocker):
memory = _build_memory_instance(mocker, Memory)
original_metadata = {"user_id": "test_user", "category": "sports"}
metadata_copy = original_metadata.copy()
memory._create_memory("test data", {"test data": [0.1, 0.2, 0.3]}, metadata=original_metadata)
assert original_metadata == metadata_copy, (
f"_create_memory mutated the caller's metadata dict: {original_metadata} != {metadata_copy}"
)
def test_create_memory_stores_correct_payload(self, mocker):
memory = _build_memory_instance(mocker, Memory)
metadata = {"user_id": "test_user", "category": "sports"}
memory._create_memory("test data", {"test data": [0.1, 0.2, 0.3]}, metadata=metadata)
payload = memory.vector_store.insert.call_args.kwargs["payloads"][0]
assert payload["data"] == "test data"
assert payload["user_id"] == "test_user"
assert payload["category"] == "sports"
assert "hash" in payload
assert "created_at" in payload
def test_create_memory_with_none_metadata(self, mocker):
memory = _build_memory_instance(mocker, Memory)
memory._create_memory("test data", {"test data": [0.1, 0.2, 0.3]}, metadata=None)
payload = memory.vector_store.insert.call_args.kwargs["payloads"][0]
assert payload["data"] == "test data"
assert "hash" in payload
def test_create_memory_shared_metadata_across_calls(self, mocker):
"""Verify that sharing a metadata dict between multiple _create_memory calls is safe."""
memory = _build_memory_instance(mocker, Memory)
shared_metadata = {"user_id": "test_user"}
memory._create_memory("first memory", {"first memory": [0.1, 0.2, 0.3]}, metadata=shared_metadata)
memory._create_memory("second memory", {"second memory": [0.4, 0.5, 0.6]}, metadata=shared_metadata)
assert shared_metadata == {"user_id": "test_user"}, "shared metadata was mutated across calls"
# Verify each call got the correct data
first_payload = memory.vector_store.insert.call_args_list[0].kwargs["payloads"][0]
second_payload = memory.vector_store.insert.call_args_list[1].kwargs["payloads"][0]
assert first_payload["data"] == "first memory"
assert second_payload["data"] == "second memory"
def test_create_memory_preserves_role_and_actor_id_in_history(self, mocker):
"""Verify that role and actor_id from metadata flow through to add_history after deepcopy."""
memory = _build_memory_instance(mocker, Memory)
metadata = {"user_id": "test_user", "role": "assistant", "actor_id": "bot-1"}
memory._create_memory("test data", {"test data": [0.1, 0.2, 0.3]}, metadata=metadata)
# Verify the payload stored in vector store has all fields
payload = memory.vector_store.insert.call_args.kwargs["payloads"][0]
assert payload["role"] == "assistant"
assert payload["actor_id"] == "bot-1"
assert payload["user_id"] == "test_user"
assert payload["data"] == "test data"
# Verify add_history received the correct role and actor_id
history_call = memory.db.add_history.call_args
assert history_call.kwargs["role"] == "assistant"
assert history_call.kwargs["actor_id"] == "bot-1"
# And the original metadata is still untouched
assert metadata == {"user_id": "test_user", "role": "assistant", "actor_id": "bot-1"}
def test_create_memory_with_nested_metadata_not_mutated(self, mocker):
"""Verify deepcopy protects nested structures in metadata."""
memory = _build_memory_instance(mocker, Memory)
metadata = {"user_id": "test_user", "tags": ["important", "urgent"], "config": {"key": "val"}}
import copy
metadata_snapshot = copy.deepcopy(metadata)
memory._create_memory("test data", {"test data": [0.1, 0.2, 0.3]}, metadata=metadata)
assert metadata == metadata_snapshot, "Nested metadata structures were mutated"
def test_update_memory_does_not_mutate_metadata(self, mocker):
memory = _build_memory_instance(mocker, Memory)
memory.vector_store.get.return_value = MagicMock(
payload={"data": "old data", "user_id": "test_user", "created_at": "2026-01-01T00:00:00+00:00"}
)
original_metadata = {"category": "updated"}
metadata_copy = original_metadata.copy()
memory._update_memory("mem-id", "new data", {"new data": [0.1, 0.2, 0.3]}, metadata=original_metadata)
assert original_metadata == metadata_copy, (
f"_update_memory mutated the caller's metadata dict: {original_metadata} != {metadata_copy}"
)
def test_add_to_vector_store_no_infer_does_not_mutate_metadata(self, mocker):
"""Verify _add_to_vector_store with infer=False doesn't leak metadata between messages."""
memory = _build_memory_instance(mocker, Memory)
memory.embedding_model.embed.return_value = [0.1, 0.2, 0.3]
original_metadata = {"user_id": "test_user"}
metadata_copy = original_metadata.copy()
messages = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi there", "name": "bot-1"},
]
result = memory._add_to_vector_store(messages, original_metadata, filters={}, infer=False)
# Metadata should not be mutated
assert original_metadata == metadata_copy, (
f"_add_to_vector_store mutated the caller's metadata: {original_metadata}"
)
# Should have created 2 memories
assert len(result) == 2
assert result[0]["role"] == "user"
assert result[1]["role"] == "assistant"
assert result[1]["actor_id"] == "bot-1"
# Verify each insert got distinct payloads with correct roles
insert_calls = memory.vector_store.insert.call_args_list
first_payload = insert_calls[0].kwargs["payloads"][0]
second_payload = insert_calls[1].kwargs["payloads"][0]
assert first_payload["role"] == "user"
assert "actor_id" not in first_payload # user message has no name
assert second_payload["role"] == "assistant"
assert second_payload["actor_id"] == "bot-1"
@pytest.mark.asyncio
async def test_async_create_memory_does_not_mutate_metadata(self, mocker):
memory = _build_memory_instance(mocker, AsyncMemory)
original_metadata = {"user_id": "test_user", "category": "sports"}
metadata_copy = original_metadata.copy()
await memory._create_memory("test data", {"test data": [0.1, 0.2, 0.3]}, metadata=original_metadata)
assert original_metadata == metadata_copy, (
f"async _create_memory mutated the caller's metadata dict: {original_metadata} != {metadata_copy}"
)
@pytest.mark.asyncio
async def test_async_create_memory_shared_metadata_across_calls(self, mocker):
memory = _build_memory_instance(mocker, AsyncMemory)
shared_metadata = {"user_id": "test_user"}
await memory._create_memory("first memory", {"first memory": [0.1, 0.2, 0.3]}, metadata=shared_metadata)
await memory._create_memory("second memory", {"second memory": [0.4, 0.5, 0.6]}, metadata=shared_metadata)
assert shared_metadata == {"user_id": "test_user"}, "shared metadata was mutated across async calls"
@pytest.mark.asyncio
async def test_async_add_to_vector_store_no_infer_does_not_mutate_metadata(self, mocker):
"""Verify async _add_to_vector_store with infer=False doesn't leak metadata between messages."""
memory = _build_memory_instance(mocker, AsyncMemory)
memory.embedding_model.embed.return_value = [0.1, 0.2, 0.3]
original_metadata = {"user_id": "test_user"}
metadata_copy = original_metadata.copy()
messages = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi there", "name": "bot-1"},
]
result = await memory._add_to_vector_store(messages, original_metadata, effective_filters={}, infer=False)
assert original_metadata == metadata_copy, (
f"async _add_to_vector_store mutated the caller's metadata: {original_metadata}"
)
assert len(result) == 2
assert result[0]["role"] == "user"
assert result[1]["role"] == "assistant"
@pytest.mark.asyncio
async def test_async_update_memory_does_not_mutate_metadata(self, mocker):
memory = _build_memory_instance(mocker, AsyncMemory)
memory.vector_store.get.return_value = MagicMock(
payload={"data": "old data", "user_id": "test_user", "created_at": "2026-01-01T00:00:00+00:00"}
)
original_metadata = {"category": "updated"}
metadata_copy = original_metadata.copy()
await memory._update_memory("mem-id", "new data", {"new data": [0.1, 0.2, 0.3]}, metadata=original_metadata)
assert original_metadata == metadata_copy, (
f"async _update_memory mutated the caller's metadata dict: {original_metadata} != {metadata_copy}"
)
def test_normalize_iso_timestamp_to_utc_preserves_naive_values():
assert _normalize_iso_timestamp_to_utc("2026-03-18T00:00:00") == "2026-03-18T00:00:00"
-517
View File
@@ -1,517 +0,0 @@
"""Tests for graph cleanup on memory deletion (issue #3245)."""
from unittest.mock import MagicMock, patch
import pytest
from mem0.configs.base import MemoryConfig
class MockVectorMemory:
def __init__(self, memory_id, payload, score=0.8):
self.id = memory_id
self.payload = payload
self.score = score
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_calls_graph_cleanup_when_graph_enabled(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""When graph is enabled, delete() should call graph.delete() with memory text and filters."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
# Enable graph with a mock
memory.enable_graph = True
memory.graph = MagicMock()
# Set up vector store to return a memory with graph-relevant data
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1",
{
"data": "Alice likes Bob",
"user_id": "user-1",
"agent_id": "agent-1",
"hash": "abc",
},
)
memory.delete("mem-1")
# graph.delete should have been called with the memory text and filters
memory.graph.delete.assert_called_once_with(
"Alice likes Bob", {"user_id": "user-1", "agent_id": "agent-1"}
)
# _delete_memory should still have been called (vector store + history cleanup)
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_skips_graph_when_not_enabled(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""When graph is not enabled, delete() should not attempt graph cleanup."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
assert memory.enable_graph is False
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"}
)
result = memory.delete("mem-1")
assert result == {"message": "Memory deleted successfully!"}
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_continues_if_graph_cleanup_fails(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""If graph cleanup raises an exception, delete() should still succeed."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
memory.graph.delete.side_effect = RuntimeError("Neo4j connection lost")
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"}
)
# Should not raise
result = memory.delete("mem-1")
assert result == {"message": "Memory deleted successfully!"}
# Vector store deletion should still proceed
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_skips_graph_when_no_user_id(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""Graph cleanup should be skipped if the memory has no user_id."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
# Memory with no user_id
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"data": "Some data", "hash": "abc"}
)
memory.delete("mem-1")
# graph.delete should NOT have been called since there's no user_id
memory.graph.delete.assert_not_called()
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_skips_graph_when_no_memory_text(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""Graph cleanup should be skipped if the memory has no text data."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"user_id": "user-1", "hash": "abc"}
)
memory.delete("mem-1")
memory.graph.delete.assert_not_called()
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_passes_all_filters_to_graph(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""Graph cleanup should include all available filters (user_id, agent_id, run_id)."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1",
{
"data": "Alice likes Bob",
"user_id": "user-1",
"agent_id": "agent-1",
"run_id": "run-1",
"hash": "abc",
},
)
memory.delete("mem-1")
memory.graph.delete.assert_called_once_with(
"Alice likes Bob",
{"user_id": "user-1", "agent_id": "agent-1", "run_id": "run-1"},
)
@pytest.mark.asyncio
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
async def test_async_delete_calls_graph_cleanup(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""Async delete() should also perform graph cleanup."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import AsyncMemory
config = MemoryConfig()
memory = AsyncMemory(config)
memory.enable_graph = True
memory.graph = MagicMock()
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1",
{
"data": "Alice likes Bob",
"user_id": "user-1",
"hash": "abc",
},
)
result = await memory.delete("mem-1")
assert result == {"message": "Memory deleted successfully!"}
memory.graph.delete.assert_called_once_with("Alice likes Bob", {"user_id": "user-1"})
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@pytest.mark.asyncio
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
async def test_async_delete_continues_if_graph_cleanup_fails(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""Async delete() should continue even if graph cleanup fails."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import AsyncMemory
config = MemoryConfig()
memory = AsyncMemory(config)
memory.enable_graph = True
memory.graph = MagicMock()
memory.graph.delete.side_effect = RuntimeError("Graph error")
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"}
)
result = await memory.delete("mem-1")
assert result == {"message": "Memory deleted successfully!"}
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_raises_for_nonexistent_memory_with_graph_enabled(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""delete() should raise ValueError for non-existent memory even with graph enabled."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
mock_vector_store.get.return_value = None
with pytest.raises(ValueError, match="Memory with id non-existent not found"):
memory.delete("non-existent")
memory.graph.delete.assert_not_called()
mock_vector_store.delete.assert_not_called()
@pytest.mark.asyncio
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
async def test_async_delete_raises_for_nonexistent_memory_with_graph_enabled(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""Async delete() should raise ValueError for non-existent memory even with graph enabled."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_store.get.return_value = None
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import AsyncMemory
config = MemoryConfig()
memory = AsyncMemory(config)
memory.enable_graph = True
memory.graph = MagicMock()
with pytest.raises(ValueError, match="Memory with id non-existent not found"):
await memory.delete("non-existent")
memory.graph.delete.assert_not_called()
mock_vector_store.delete.assert_not_called()
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_all_does_not_trigger_per_memory_graph_cleanup(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""delete_all() should use graph.delete_all(), not per-memory graph.delete()."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
mem1 = MockVectorMemory("mem-1", {"data": "Alice likes Bob", "user_id": "user-1"})
mem2 = MockVectorMemory("mem-2", {"data": "Bob likes Charlie", "user_id": "user-1"})
mock_vector_store.list.return_value = ([mem1, mem2], 2)
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"data": "Alice likes Bob", "user_id": "user-1"}
)
memory.delete_all(user_id="user-1")
# graph.delete (per-memory) should NOT be called
memory.graph.delete.assert_not_called()
# graph.delete_all (bulk) SHOULD be called
memory.graph.delete_all.assert_called_once_with({"user_id": "user-1"})
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_internal_delete_memory_does_not_trigger_graph_cleanup(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""_delete_memory() should NOT call graph.delete() — only the public delete() does.
This ensures that the DELETE branch inside _add_to_vector_store() (which calls
_delete_memory directly) does not interfere with the parallel graph pipeline
running in _add_to_graph().
"""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"}
)
# Call _delete_memory directly (as _add_to_vector_store does for DELETE events)
memory._delete_memory("mem-1")
# graph.delete should NOT have been called — graph cleanup is only in delete()
memory.graph.delete.assert_not_called()
# But vector store deletion should proceed
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
def test_graph_memory_delete_calls_internal_methods():
"""Test that MemoryGraph.delete() calls the expected internal pipeline methods."""
from unittest.mock import patch as _patch
# We need to mock the Neo4j import
with _patch.dict("sys.modules", {"langchain_neo4j": MagicMock(), "rank_bm25": MagicMock()}):
from mem0.memory.graph_memory import MemoryGraph
with _patch.object(MemoryGraph, "__init__", return_value=None):
graph = MemoryGraph.__new__(MemoryGraph)
# Mock the internal methods
graph._retrieve_nodes_from_data = MagicMock(
return_value={"alice": "person", "bob": "person"}
)
graph._establish_nodes_relations_from_data = MagicMock(
return_value=[
{"source": "alice", "destination": "bob", "relationship": "likes"}
]
)
graph._delete_entities = MagicMock(return_value=[])
filters = {"user_id": "user-1"}
graph.delete("Alice likes Bob", filters)
graph._retrieve_nodes_from_data.assert_called_once_with("Alice likes Bob", filters)
graph._establish_nodes_relations_from_data.assert_called_once_with(
"Alice likes Bob", filters, {"alice": "person", "bob": "person"}
)
graph._delete_entities.assert_called_once_with(
[{"source": "alice", "destination": "bob", "relationship": "likes"}],
filters,
)
def test_graph_memory_delete_skips_when_no_entities():
"""Test that MemoryGraph.delete() does nothing when no entities are extracted."""
from unittest.mock import patch as _patch
with _patch.dict("sys.modules", {"langchain_neo4j": MagicMock(), "rank_bm25": MagicMock()}):
from mem0.memory.graph_memory import MemoryGraph
with _patch.object(MemoryGraph, "__init__", return_value=None):
graph = MemoryGraph.__new__(MemoryGraph)
graph._retrieve_nodes_from_data = MagicMock(return_value={})
graph._establish_nodes_relations_from_data = MagicMock()
graph._delete_entities = MagicMock()
graph.delete("Some text", {"user_id": "user-1"})
graph._retrieve_nodes_from_data.assert_called_once()
graph._establish_nodes_relations_from_data.assert_not_called()
graph._delete_entities.assert_not_called()
def test_graph_memory_delete_handles_exception():
"""Test that MemoryGraph.delete() catches exceptions without raising."""
from unittest.mock import patch as _patch
with _patch.dict("sys.modules", {"langchain_neo4j": MagicMock(), "rank_bm25": MagicMock()}):
from mem0.memory.graph_memory import MemoryGraph
with _patch.object(MemoryGraph, "__init__", return_value=None):
graph = MemoryGraph.__new__(MemoryGraph)
graph._retrieve_nodes_from_data = MagicMock(
side_effect=RuntimeError("LLM error")
)
# Should not raise
graph.delete("Some text", {"user_id": "user-1"})
-821
View File
@@ -1,821 +0,0 @@
"""
End-to-end tests for graph cleanup on memory deletion against real
Neo4j, Memgraph, and Apache AGE instances running in Docker.
Requires:
docker run -d --name mem0-neo4j-test -p 7687:7687 -e NEO4J_AUTH=neo4j/testpassword neo4j:5.23
docker run -d --name mem0-memgraph-test -p 7688:7687 memgraph/memgraph:latest
docker run -d --name mem0-age-test -p 5432:5432 -e POSTGRES_USER=postgres -e POSTGRES_PASSWORD=testpassword -e POSTGRES_DB=testdb apache/age:latest
Tests are skipped automatically if the databases or required Python
packages are not available.
"""
import hashlib
import sys
import warnings
from unittest.mock import MagicMock
import pytest
warnings.filterwarnings("ignore")
EMBEDDING_DIMS = 64
# ---------------------------------------------------------------------------
# Deterministic embedding helper (shared across backends)
# ---------------------------------------------------------------------------
def _make_deterministic_embedder():
cache = {}
counter = [0]
def embed(text, *args, **kwargs):
t = text.lower().strip()
if t not in cache:
vec = [0.0] * EMBEDDING_DIMS
idx = counter[0] % EMBEDDING_DIMS
vec[idx] = 1.0
h = hashlib.sha256(t.encode()).digest()
for i in range(EMBEDDING_DIMS):
vec[i] += float(h[i % len(h)]) / 25500.0
norm = sum(v * v for v in vec) ** 0.5
cache[t] = [v / norm for v in vec]
counter[0] += 1
return cache[t]
mock = MagicMock()
mock.embed.side_effect = embed
mock.config.embedding_dims = EMBEDDING_DIMS
return mock
def _make_mock_llm(entities, relations):
"""Create an LLM mock that returns specific entities and relations."""
mock = MagicMock()
def generate_response(messages, tools):
tool_names = []
for t in tools:
if isinstance(t, dict):
fn = t.get("function", t)
tool_names.append(fn.get("name", ""))
else:
tool_names.append(getattr(t, "name", str(t)))
if any("extract_entities" in n for n in tool_names):
return {
"tool_calls": [
{"name": "extract_entities", "arguments": {"entities": entities}}
]
}
elif any("establish" in n or "relation" in n for n in tool_names):
return {
"tool_calls": [
{"name": "establish_nodes_relations", "arguments": {"entities": relations}}
]
}
elif any("delete" in n for n in tool_names):
return {"tool_calls": []}
return {"tool_calls": []}
mock.generate_response.side_effect = generate_response
return mock
# ===========================================================================
# NEO4J
# ===========================================================================
def _port_open(host, port, timeout=1):
"""Quick TCP check — avoids slow driver-level timeouts."""
import socket
try:
with socket.create_connection((host, port), timeout=timeout):
return True
except OSError:
return False
def _neo4j_available():
if not _port_open("localhost", 7687):
return False
try:
from langchain_neo4j import Neo4jGraph
g = Neo4jGraph(
url="bolt://localhost:7687",
username="neo4j",
password="testpassword",
refresh_schema=False,
driver_config={"notifications_min_severity": "OFF"},
)
g.query("RETURN 1")
return True
except Exception:
return False
requires_neo4j = pytest.mark.skipif(not _neo4j_available(), reason="Neo4j not available")
@pytest.fixture
def neo4j_graph():
"""Create a Neo4j-backed MemoryGraph with mocked LLM/embedder."""
from langchain_neo4j import Neo4jGraph
from mem0.memory.graph_memory import MemoryGraph
mg = MemoryGraph.__new__(MemoryGraph)
mg.graph = Neo4jGraph(
url="bolt://localhost:7687",
username="neo4j",
password="testpassword",
refresh_schema=False,
driver_config={"notifications_min_severity": "OFF"},
)
mg.graph.query("MATCH (n) DETACH DELETE n")
mg.node_label = ":`__Entity__`"
mg.llm_provider = "openai"
mg.user_id = None
mg.threshold = 0.99
mg.embedding_model = _make_deterministic_embedder()
mg.llm = MagicMock()
mg.config = MagicMock()
mg.config.graph_store.custom_prompt = None
mg.config.graph_store.config.base_label = True
yield mg
mg.graph.query("MATCH (n) DETACH DELETE n")
@requires_neo4j
class TestNeo4jDeleteE2E:
def _node_count(self, mg):
return mg.graph.query("MATCH (n) RETURN count(n) AS cnt")[0]["cnt"]
def _valid_edge_count(self, mg):
return mg.graph.query(
"MATCH ()-[r]->() WHERE r.valid IS NULL OR r.valid = true RETURN count(r) AS cnt"
)[0]["cnt"]
def _invalid_edge_count(self, mg):
return mg.graph.query(
"MATCH ()-[r]->() WHERE r.valid = false RETURN count(r) AS cnt"
)[0]["cnt"]
def test_add_creates_graph_data(self, neo4j_graph):
mg = neo4j_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._node_count(mg) == 2
assert self._valid_edge_count(mg) == 1
def test_delete_soft_deletes_relationships(self, neo4j_graph):
"""Neo4j delete() should set r.valid=false (soft-delete), not hard-delete."""
mg = neo4j_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 1
assert self._invalid_edge_count(mg) == 0
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 0
assert self._invalid_edge_count(mg) == 1 # soft-deleted, not removed
assert self._node_count(mg) == 2 # nodes preserved
def test_delete_preserves_other_relationships(self, neo4j_graph):
mg = neo4j_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Charlie", "entity_type": "person"}],
[{"source": "Alice", "destination": "Charlie", "relationship": "knows"}],
)
mg.add("Alice knows Charlie", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 2
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 1
assert self._invalid_edge_count(mg) == 1
def test_delete_user_isolation(self, neo4j_graph):
mg = neo4j_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
mg.add("Alice likes Bob", {"user_id": "u2"})
assert self._valid_edge_count(mg) == 2
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 1
def test_delete_all_hard_deletes(self, neo4j_graph):
mg = neo4j_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._node_count(mg) == 2
mg.delete_all({"user_id": "u1"})
assert self._node_count(mg) == 0
def test_add_delete_add_cycle(self, neo4j_graph):
mg = neo4j_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 0
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 1
# ===========================================================================
# MEMGRAPH
# ===========================================================================
def _memgraph_available():
if not _port_open("localhost", 7688):
return False
try:
from langchain_memgraph.graphs.memgraph import Memgraph
g = Memgraph("bolt://localhost:7688", "memgraph", "memgraph")
g.query("RETURN 1")
return True
except Exception:
return False
requires_memgraph = pytest.mark.skipif(
not _memgraph_available(), reason="Memgraph not available"
)
@pytest.fixture
def memgraph_graph():
"""Create a Memgraph-backed MemoryGraph with mocked LLM/embedder."""
from langchain_memgraph.graphs.memgraph import Memgraph
from mem0.memory.memgraph_memory import MemoryGraph
mg = MemoryGraph.__new__(MemoryGraph)
mg.graph = Memgraph("bolt://localhost:7688", "memgraph", "memgraph")
mg.graph.query("MATCH (n) DETACH DELETE n")
try:
mg.graph.query("DROP VECTOR INDEX memzero;")
except Exception:
pass
mg.graph.query(
f"CREATE VECTOR INDEX memzero ON :Entity(embedding) "
f"WITH CONFIG {{'dimension': {EMBEDDING_DIMS}, 'capacity': 1000, 'metric': 'cos'}};"
)
try:
mg.graph.query("CREATE INDEX ON :Entity(user_id);")
except Exception:
pass
try:
mg.graph.query("CREATE INDEX ON :Entity;")
except Exception:
pass
mg.llm_provider = "openai"
mg.user_id = None
mg.threshold = 0.99
mg.embedding_model = _make_deterministic_embedder()
mg.llm = MagicMock()
mg.config = MagicMock()
mg.config.graph_store.custom_prompt = None
mg.config.embedder.config = {"embedding_dims": EMBEDDING_DIMS}
yield mg
mg.graph.query("MATCH (n) DETACH DELETE n")
@requires_memgraph
class TestMemgraphDeleteE2E:
def _node_count(self, mg):
return mg.graph.query("MATCH (n:Entity) RETURN count(n) AS cnt")[0]["cnt"]
def _edge_count(self, mg):
return mg.graph.query("MATCH ()-[r]->() RETURN count(r) AS cnt")[0]["cnt"]
def test_add_creates_graph_data(self, memgraph_graph):
mg = memgraph_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._node_count(mg) == 2
assert self._edge_count(mg) == 1
def test_delete_hard_deletes_relationships(self, memgraph_graph):
"""Memgraph delete() should hard-delete the relationship (DELETE r)."""
mg = memgraph_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 0
assert self._node_count(mg) == 2
def test_delete_preserves_other_relationships(self, memgraph_graph):
mg = memgraph_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Charlie", "entity_type": "person"}],
[{"source": "Alice", "destination": "Charlie", "relationship": "knows"}],
)
mg.add("Alice knows Charlie", {"user_id": "u1"})
assert self._edge_count(mg) == 2
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 1
def test_delete_user_isolation(self, memgraph_graph):
mg = memgraph_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
mg.add("Alice likes Bob", {"user_id": "u2"})
assert self._edge_count(mg) == 2
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 1
def test_delete_all_hard_deletes(self, memgraph_graph):
mg = memgraph_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._node_count(mg) == 2
mg.delete_all({"user_id": "u1"})
assert self._node_count(mg) == 0
assert self._edge_count(mg) == 0
def test_add_delete_add_cycle(self, memgraph_graph):
mg = memgraph_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 0
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 1
# ===========================================================================
# APACHE AGE
# ===========================================================================
def _age_available():
if not _port_open("localhost", 5432):
return False
try:
import age
ag = age.connect(
host="localhost",
port=5432,
dbname="testdb",
user="postgres",
password="testpassword",
)
with ag.connection.cursor() as cur:
cur.execute("CREATE EXTENSION IF NOT EXISTS age;")
cur.execute("SET search_path = ag_catalog, '$user', public;")
ag.connection.commit()
ag.close()
return True
except Exception:
return False
requires_age = pytest.mark.skipif(not _age_available(), reason="Apache AGE not available")
@pytest.fixture
def age_graph():
"""Create an Apache AGE-backed MemoryGraph with mocked LLM/embedder."""
import age
from mem0.memory.apache_age_memory import MemoryGraph
graph_name = "mem0_test_delete"
ag = age.connect(
graph=graph_name,
host="localhost",
port=5432,
dbname="testdb",
user="postgres",
password="testpassword",
)
with ag.connection.cursor() as cur:
cur.execute("CREATE EXTENSION IF NOT EXISTS age;")
cur.execute("SET search_path = ag_catalog, '$user', public;")
ag.connection.commit()
age.setUpAge(ag.connection, graph_name)
ag.connection.commit()
mg = MemoryGraph.__new__(MemoryGraph)
mg.ag = ag
mg.graph_name = graph_name
mg.llm_provider = "openai"
mg.user_id = None
mg.threshold = 0.99
mg.embedding_model = _make_deterministic_embedder()
mg.llm = MagicMock()
mg.config = MagicMock()
mg.config.graph_store.custom_prompt = None
try:
ag.execCypher("MATCH (n) DETACH DELETE n")
ag.commit()
except Exception:
ag.rollback()
yield mg
try:
ag.execCypher("MATCH (n) DETACH DELETE n")
ag.commit()
except Exception:
ag.rollback()
ag.close()
def _age_node_count(mg):
cursor = mg.ag.execCypher("MATCH (n) RETURN count(n)", cols=["cnt"])
rows = cursor.fetchall()
return rows[0][0] if rows else 0
def _age_edge_count(mg):
cursor = mg.ag.execCypher("MATCH ()-[r]->() RETURN count(r)", cols=["cnt"])
rows = cursor.fetchall()
return rows[0][0] if rows else 0
@requires_age
class TestApacheAgeDeleteE2E:
def test_add_creates_graph_data(self, age_graph):
mg = age_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert _age_node_count(mg) == 2
assert _age_edge_count(mg) == 1
def test_delete_hard_deletes_relationships(self, age_graph):
"""Apache AGE delete() should hard-delete the relationship."""
mg = age_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 0
assert _age_node_count(mg) == 2
def test_delete_preserves_other_relationships(self, age_graph):
mg = age_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Charlie", "entity_type": "person"}],
[{"source": "Alice", "destination": "Charlie", "relationship": "knows"}],
)
mg.add("Alice knows Charlie", {"user_id": "u1"})
assert _age_edge_count(mg) == 2
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 1
def test_delete_user_isolation(self, age_graph):
mg = age_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
mg.add("Alice likes Bob", {"user_id": "u2"})
assert _age_edge_count(mg) == 2
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 1
def test_delete_all_hard_deletes(self, age_graph):
mg = age_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert _age_node_count(mg) == 2
mg.delete_all({"user_id": "u1"})
assert _age_node_count(mg) == 0
assert _age_edge_count(mg) == 0
def test_add_delete_add_cycle(self, age_graph):
mg = age_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 0
mg.add("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 1
# ===========================================================================
# NEPTUNE (tested via Neo4j OpenCypher — same query language)
# ===========================================================================
def _neptune_test_available():
"""Neptune uses OpenCypher — we test NeptuneBase.delete() against Neo4j."""
if not _port_open("localhost", 7687):
return False
try:
# Mock langchain_aws so NeptuneBase can be imported without AWS deps
sys.modules.setdefault("langchain_aws", MagicMock())
sys.modules.setdefault("botocore", MagicMock())
sys.modules.setdefault("botocore.config", MagicMock())
from mem0.graphs.neptune.base import NeptuneBase # noqa: F401
from langchain_neo4j import Neo4jGraph
g = Neo4jGraph(
url="bolt://localhost:7687",
username="neo4j",
password="testpassword",
refresh_schema=False,
driver_config={"notifications_min_severity": "OFF"},
)
g.query("RETURN 1")
return True
except Exception:
return False
requires_neptune_test = pytest.mark.skipif(
not _neptune_test_available(),
reason="Neo4j not available (used as OpenCypher backend for Neptune tests)",
)
def _make_concrete_neptune_subclass():
"""Create a concrete NeptuneBase subclass for testing, backed by Neo4j."""
# Ensure mocks are in place for import
sys.modules.setdefault("langchain_aws", MagicMock())
sys.modules.setdefault("botocore", MagicMock())
sys.modules.setdefault("botocore.config", MagicMock())
from mem0.graphs.neptune.base import NeptuneBase
class TestableNeptune(NeptuneBase):
def __init__(self):
pass
def _delete_entities_cypher(self, source, destination, relationship, user_id):
cypher = f"""
MATCH (n:`__Entity__` {{name: $source_name, user_id: $user_id}})
-[r:{relationship}]->
(m:`__Entity__` {{name: $dest_name, user_id: $user_id}})
DELETE r
RETURN n.name AS source, m.name AS target, type(r) AS relationship
"""
return cypher, {"source_name": source, "dest_name": destination, "user_id": user_id}
def _delete_all_cypher(self, filters):
return (
"MATCH (n:`__Entity__` {user_id: $user_id}) DETACH DELETE n",
{"user_id": filters["user_id"]},
)
# Stubs for abstract methods not used in delete path
def _add_entities_by_source_cypher(self, *a, **kw): pass
def _add_entities_by_destination_cypher(self, *a, **kw): pass
def _add_relationship_entities_cypher(self, *a, **kw): pass
def _add_new_entities_cypher(self, *a, **kw): pass
def _search_source_node_cypher(self, *a, **kw): pass
def _search_destination_node_cypher(self, *a, **kw): pass
def _get_all_cypher(self, *a, **kw): pass
def _search_graph_db_cypher(self, *a, **kw): pass
return TestableNeptune
@pytest.fixture
def neptune_graph():
"""NeptuneBase subclass backed by a real Neo4j container."""
from langchain_neo4j import Neo4jGraph
cls = _make_concrete_neptune_subclass()
mg = cls()
mg.graph = Neo4jGraph(
url="bolt://localhost:7687",
username="neo4j",
password="testpassword",
refresh_schema=False,
driver_config={"notifications_min_severity": "OFF"},
)
mg.graph.query("MATCH (n) DETACH DELETE n")
mg.node_label = ":`__Entity__`"
mg.llm_provider = "openai"
mg.user_id = None
mg.threshold = 0.99
mg.embedding_model = _make_deterministic_embedder()
mg.llm = MagicMock()
mg.config = MagicMock()
mg.config.graph_store.custom_prompt = None
yield mg
mg.graph.query("MATCH (n) DETACH DELETE n")
def _neptune_node_count(mg):
return mg.graph.query("MATCH (n) RETURN count(n) AS cnt")[0]["cnt"]
def _neptune_edge_count(mg):
return mg.graph.query("MATCH ()-[r]->() RETURN count(r) AS cnt")[0]["cnt"]
def _neptune_create_entities(mg, user_id):
"""Create test entities directly via Cypher."""
mg.graph.query(f"""
CREATE (a:`__Entity__` {{name: 'alice', user_id: '{user_id}'}})
CREATE (b:`__Entity__` {{name: 'bob', user_id: '{user_id}'}})
CREATE (a)-[:likes]->(b)
""")
@requires_neptune_test
class TestNeptuneDeleteE2E:
"""Test NeptuneBase.delete() using Neo4j as the OpenCypher backend.
Neptune uses standard OpenCypher, the same query language as Neo4j.
This validates that:
- NeptuneBase.delete() correctly calls _delete_entities(to_be_deleted, user_id) with a string
- The generated Cypher from _delete_entities_cypher runs correctly
- User isolation works
"""
def test_delete_removes_relationship(self, neptune_graph):
mg = neptune_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
_neptune_create_entities(mg, "u1")
assert _neptune_edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _neptune_edge_count(mg) == 0
assert _neptune_node_count(mg) == 2
def test_delete_user_isolation(self, neptune_graph):
mg = neptune_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
_neptune_create_entities(mg, "u1")
_neptune_create_entities(mg, "u2")
assert _neptune_edge_count(mg) == 2
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _neptune_edge_count(mg) == 1
def test_delete_passes_user_id_string_not_dict(self, neptune_graph):
"""Verify NeptuneBase.delete() passes filters['user_id'] (string) to _delete_entities."""
mg = neptune_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
original = mg._delete_entities
call_args = []
def spy(to_be_deleted, user_id):
call_args.append(("to_be_deleted", to_be_deleted, "user_id", user_id))
return original(to_be_deleted, user_id)
mg._delete_entities = spy
_neptune_create_entities(mg, "u1")
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert len(call_args) == 1
assert call_args[0][3] == "u1"
assert isinstance(call_args[0][3], str)
def test_delete_all(self, neptune_graph):
mg = neptune_graph
_neptune_create_entities(mg, "u1")
assert _neptune_node_count(mg) == 2
mg.delete_all({"user_id": "u1"})
assert _neptune_node_count(mg) == 0
def test_add_delete_add_cycle(self, neptune_graph):
mg = neptune_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
_neptune_create_entities(mg, "u1")
assert _neptune_edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _neptune_edge_count(mg) == 0
_neptune_create_entities(mg, "u1")
assert _neptune_edge_count(mg) == 1
-736
View File
@@ -1,736 +0,0 @@
"""
End-to-end tests for graph cleanup on memory deletion (issue #3245).
Uses a real Kuzu embedded database to verify that graph entities are
correctly cleaned up when memories are deleted. LLM and embedding calls
are mocked to provide deterministic entity extraction.
Tests are skipped automatically if kuzu is not installed.
"""
import shutil
import tempfile
from unittest.mock import MagicMock, patch
import pytest
from mem0.configs.base import MemoryConfig
try:
import kuzu # noqa: F401
_kuzu_available = True
except ImportError:
_kuzu_available = False
requires_kuzu = pytest.mark.skipif(not _kuzu_available, reason="kuzu is not installed")
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _node_count(kuzu_graph):
"""Return total node count in the Kuzu graph."""
result = kuzu_graph.execute("MATCH (n:Entity) RETURN count(n) AS cnt")
rows = list(result.rows_as_dict())
return int(rows[0]["cnt"])
def _edge_count(kuzu_graph):
"""Return total edge count in the Kuzu graph."""
result = kuzu_graph.execute("MATCH ()-[r:CONNECTED_TO]->() RETURN count(r) AS cnt")
rows = list(result.rows_as_dict())
return int(rows[0]["cnt"])
def _get_edges(kuzu_graph):
"""Return all edges as list of (source, relationship, destination) tuples."""
result = kuzu_graph.execute(
"MATCH (s:Entity)-[r:CONNECTED_TO]->(d:Entity) "
"RETURN s.name AS src, r.name AS rel, d.name AS dst"
)
return [(row["src"], row["rel"], row["dst"]) for row in result.rows_as_dict()]
def _get_nodes(kuzu_graph):
"""Return all node names."""
result = kuzu_graph.execute("MATCH (n:Entity) RETURN n.name AS name, n.user_id AS uid")
return [(row["name"], row["uid"]) for row in result.rows_as_dict()]
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
class MockVectorMemory:
"""Mimics the object returned by vector_store.get()."""
def __init__(self, memory_id, payload, score=0.8):
self.id = memory_id
self.payload = payload
self.score = score
@pytest.fixture
def kuzu_graph_memory():
"""
Create a real Kuzu-backed MemoryGraph with mocked LLM and embedder.
Yields (graph_memory_instance, kuzu_connection) then cleans up.
"""
import os
import kuzu
tmpdir = tempfile.mkdtemp()
db_path = os.path.join(tmpdir, "test.kuzu")
db = kuzu.Database(db_path)
conn = kuzu.Connection(db)
# We'll construct the MemoryGraph by bypassing __init__ and setting up manually
from mem0.memory.kuzu_memory import MemoryGraph
mg = MemoryGraph.__new__(MemoryGraph)
# Real Kuzu connection
mg.db = db
mg.graph = conn
mg.node_label = ":Entity"
mg.rel_label = ":CONNECTED_TO"
mg.kuzu_create_schema()
# Deterministic embedding: use one-hot-style vectors per entity name
# to avoid accidental cosine similarity matches between different entities
embedding_dims = 64
mg.embedding_dims = embedding_dims
_embed_cache = {}
_embed_counter = [0]
def deterministic_embed(text):
"""Generate a deterministic, near-orthogonal embedding for each unique text."""
text_lower = text.lower().strip()
if text_lower not in _embed_cache:
# Create a sparse vector — set a unique dimension to 1.0
vec = [0.0] * embedding_dims
idx = _embed_counter[0] % embedding_dims
vec[idx] = 1.0
# Add small noise to other dims so it's not exactly zero
import hashlib
h = hashlib.sha256(text_lower.encode()).digest()
for i in range(embedding_dims):
vec[i] += float(h[i % len(h)]) / 25500.0 # tiny noise
norm = sum(v * v for v in vec) ** 0.5
_embed_cache[text_lower] = [v / norm for v in vec]
_embed_counter[0] += 1
return _embed_cache[text_lower]
mock_embedder = MagicMock()
mock_embedder.embed.side_effect = deterministic_embed
mock_embedder.config.embedding_dims = embedding_dims
mg.embedding_model = mock_embedder
# Mock LLM — configured per-test via mock_embedder
mg.llm = MagicMock()
mg.llm_provider = "openai"
mg.user_id = None
# High threshold so only identical entity names merge, not similar ones
mg.threshold = 0.99
mg.config = MagicMock()
mg.config.graph_store.custom_prompt = None
yield mg, conn
# Cleanup
conn.close()
shutil.rmtree(tmpdir, ignore_errors=True)
def _setup_llm_for_entities(mg, entities, relations):
"""
Configure the mock LLM to return specific entities and relations.
entities: list of {"entity": str, "entity_type": str}
relations: list of {"source": str, "destination": str, "relationship": str}
"""
def generate_response(messages, tools):
# Detect which tool is being called based on tool definition names
tool_names = []
for t in tools:
if isinstance(t, dict):
fn = t.get("function", t)
tool_names.append(fn.get("name", ""))
else:
tool_names.append(getattr(t, "name", str(t)))
if any("extract_entities" in n for n in tool_names):
return {
"tool_calls": [
{
"name": "extract_entities",
"arguments": {"entities": entities},
}
]
}
elif any("establish" in n or "relation" in n for n in tool_names):
return {
"tool_calls": [
{
"name": "establish_nodes_relations",
"arguments": {"entities": relations},
}
]
}
elif any("delete" in n for n in tool_names):
# For _get_delete_entities_from_search_output during add() — return nothing to delete
return {"tool_calls": []}
return {"tool_calls": []}
mg.llm.generate_response.side_effect = generate_response
# ---------------------------------------------------------------------------
# End-to-end tests
# ---------------------------------------------------------------------------
@requires_kuzu
class TestKuzuGraphDeleteE2E:
"""End-to-end tests using a real Kuzu database."""
def test_add_creates_nodes_and_edges(self, kuzu_graph_memory):
"""Baseline: verify add() actually creates graph data."""
mg, conn = kuzu_graph_memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
filters = {"user_id": "test_user"}
mg.add("Alice likes Bob", filters)
assert _node_count(conn) == 2
assert _edge_count(conn) == 1
edges = _get_edges(conn)
assert ("alice", "likes", "bob") in edges
def test_delete_removes_edges_created_by_add(self, kuzu_graph_memory):
"""Core test: delete() should remove the relationships that add() created."""
mg, conn = kuzu_graph_memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
filters = {"user_id": "test_user"}
mg.add("Alice likes Bob", filters)
assert _edge_count(conn) == 1
# Now delete using the same text — should remove the relationship
mg.delete("Alice likes Bob", filters)
assert _edge_count(conn) == 0
# Nodes remain (we don't delete nodes on single memory delete)
assert _node_count(conn) == 2
def test_delete_only_removes_matching_edges(self, kuzu_graph_memory):
"""delete() should only remove edges matching the extracted relationships."""
mg, conn = kuzu_graph_memory
# First add: Alice likes Bob
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
filters = {"user_id": "test_user"}
mg.add("Alice likes Bob", filters)
# Second add: Alice knows Charlie
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Charlie", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Charlie", "relationship": "knows"},
],
)
mg.add("Alice knows Charlie", filters)
assert _edge_count(conn) == 2
# Delete only the "Alice likes Bob" memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
mg.delete("Alice likes Bob", filters)
assert _edge_count(conn) == 1
edges = _get_edges(conn)
assert ("alice", "knows", "charlie") in edges
assert ("alice", "likes", "bob") not in edges
def test_delete_with_different_user_id_does_not_affect_other_users(self, kuzu_graph_memory):
"""delete() scoped to user_id should not touch another user's graph data."""
mg, conn = kuzu_graph_memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
# Add for user1
mg.add("Alice likes Bob", {"user_id": "user1"})
# Add same data for user2
mg.add("Alice likes Bob", {"user_id": "user2"})
assert _edge_count(conn) == 2
# Delete only user1's data
mg.delete("Alice likes Bob", {"user_id": "user1"})
assert _edge_count(conn) == 1
# Remaining edge belongs to user2
nodes = _get_nodes(conn)
user2_nodes = [n for n in nodes if n[1] == "user2"]
assert len(user2_nodes) == 2
def test_delete_nonexistent_relationship_is_safe(self, kuzu_graph_memory):
"""delete() on data that doesn't exist in the graph should be a no-op."""
mg, conn = kuzu_graph_memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "hates"},
],
)
filters = {"user_id": "test_user"}
# Nothing in the graph yet
assert _edge_count(conn) == 0
assert _node_count(conn) == 0
# Should not raise
mg.delete("Alice hates Bob", filters)
assert _edge_count(conn) == 0
assert _node_count(conn) == 0
def test_delete_with_llm_failure_does_not_raise(self, kuzu_graph_memory):
"""If LLM fails during entity extraction, delete() should not raise."""
mg, conn = kuzu_graph_memory
# Make LLM raise
mg.llm.generate_response.side_effect = RuntimeError("LLM service down")
filters = {"user_id": "test_user"}
# Should not raise
mg.delete("Alice likes Bob", filters)
def test_delete_with_empty_entity_extraction(self, kuzu_graph_memory):
"""If LLM returns no entities, delete() should be a no-op."""
mg, conn = kuzu_graph_memory
# Add real data
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
filters = {"user_id": "test_user"}
mg.add("Alice likes Bob", filters)
assert _edge_count(conn) == 1
# Now delete but LLM returns no entities
_setup_llm_for_entities(mg, entities=[], relations=[])
mg.delete("some text", filters)
# Data should still be there
assert _edge_count(conn) == 1
def test_delete_all_removes_everything_for_user(self, kuzu_graph_memory):
"""delete_all() should remove all nodes/edges for a user (baseline behavior)."""
mg, conn = kuzu_graph_memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
filters = {"user_id": "test_user"}
mg.add("Alice likes Bob", filters)
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Bob", "entity_type": "person"},
{"entity": "Charlie", "entity_type": "person"},
],
relations=[
{"source": "Bob", "destination": "Charlie", "relationship": "knows"},
],
)
mg.add("Bob knows Charlie", filters)
assert _node_count(conn) >= 3
assert _edge_count(conn) == 2
mg.delete_all(filters)
assert _node_count(conn) == 0
assert _edge_count(conn) == 0
def test_add_delete_add_cycle(self, kuzu_graph_memory):
"""Verify that add → delete → re-add works correctly."""
mg, conn = kuzu_graph_memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
filters = {"user_id": "test_user"}
# Add
mg.add("Alice likes Bob", filters)
assert _edge_count(conn) == 1
# Delete
mg.delete("Alice likes Bob", filters)
assert _edge_count(conn) == 0
# Re-add
mg.add("Alice likes Bob", filters)
assert _edge_count(conn) == 1
edges = _get_edges(conn)
assert ("alice", "likes", "bob") in edges
@requires_kuzu
class TestMemoryDeleteWithGraphE2E:
"""
End-to-end tests for Memory.delete() with graph enabled.
Uses a real Kuzu database for the graph store and mocks for
the vector store, LLM, and embedder.
"""
@pytest.fixture
def memory_with_graph(self):
"""Create a Memory instance with a real Kuzu graph backend."""
import os
import kuzu
tmpdir = tempfile.mkdtemp()
with (
patch("mem0.utils.factory.EmbedderFactory.create") as mock_embedder_factory,
patch("mem0.utils.factory.VectorStoreFactory.create") as mock_vector_factory,
patch("mem0.utils.factory.LlmFactory.create") as mock_llm_factory,
patch("mem0.memory.storage.SQLiteManager") as mock_sqlite,
):
_mem_embed_cache = {}
_mem_embed_counter = [0]
def _mem_deterministic_embed(text, *args, **kwargs):
text_lower = text.lower().strip()
if text_lower not in _mem_embed_cache:
import hashlib
vec = [0.0] * 64
idx = _mem_embed_counter[0] % 64
vec[idx] = 1.0
h = hashlib.sha256(text_lower.encode()).digest()
for i in range(64):
vec[i] += float(h[i % len(h)]) / 25500.0
norm = sum(v * v for v in vec) ** 0.5
_mem_embed_cache[text_lower] = [v / norm for v in vec]
_mem_embed_counter[0] += 1
return _mem_embed_cache[text_lower]
mock_embedder = MagicMock()
mock_embedder.embed.side_effect = _mem_deterministic_embed
mock_embedder.config.embedding_dims = 64
mock_embedder_factory.return_value = mock_embedder
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm = MagicMock()
mock_llm_factory.return_value = mock_llm
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
# Now wire up a real Kuzu graph
db_path = os.path.join(tmpdir, "test.kuzu")
db = kuzu.Database(db_path)
conn = kuzu.Connection(db)
from mem0.memory.kuzu_memory import MemoryGraph as KuzuMemoryGraph
graph = KuzuMemoryGraph.__new__(KuzuMemoryGraph)
graph.db = db
graph.graph = conn
graph.node_label = ":Entity"
graph.rel_label = ":CONNECTED_TO"
graph.kuzu_create_schema()
graph.embedding_dims = 64
graph.embedding_model = mock_embedder
graph.llm = mock_llm
graph.llm_provider = "openai"
graph.user_id = None
graph.threshold = 0.99
graph.config = MagicMock()
graph.config.graph_store.custom_prompt = None
memory.graph = graph
memory.enable_graph = True
yield memory, mock_vector_store, mock_llm, conn
conn.close()
shutil.rmtree(tmpdir, ignore_errors=True)
def test_memory_delete_triggers_graph_cleanup(self, memory_with_graph):
"""
Full integration: Memory.delete() should clean up both vector store and graph.
"""
memory, mock_vs, mock_llm, conn = memory_with_graph
# 1. Manually add entities to the graph (simulating what add() would do)
_setup_llm_for_memory_graph(
mock_llm,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
memory.graph.add("Alice likes Bob", {"user_id": "user-1"})
assert _edge_count(conn) == 1
# 2. Set up mock vector store to return this memory
mock_vs.get.return_value = MockVectorMemory(
"mem-1",
{"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"},
)
# 3. Delete the memory
result = memory.delete("mem-1")
assert result == {"message": "Memory deleted successfully!"}
# 4. Verify graph was cleaned up
assert _edge_count(conn) == 0
# 5. Verify vector store was also cleaned up
mock_vs.delete.assert_called_once_with(vector_id="mem-1")
def test_memory_delete_with_graph_preserves_other_users_data(self, memory_with_graph):
"""Deleting user1's memory should not affect user2's graph data."""
memory, mock_vs, mock_llm, conn = memory_with_graph
_setup_llm_for_memory_graph(
mock_llm,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
# Add data for two users
memory.graph.add("Alice likes Bob", {"user_id": "user-1"})
memory.graph.add("Alice likes Bob", {"user_id": "user-2"})
assert _edge_count(conn) == 2
# Delete only user-1's memory
mock_vs.get.return_value = MockVectorMemory(
"mem-1",
{"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"},
)
memory.delete("mem-1")
# user-2's data should be intact
assert _edge_count(conn) == 1
nodes = _get_nodes(conn)
remaining_user_ids = set(uid for _, uid in nodes)
assert "user-2" in remaining_user_ids
def test_memory_delete_graph_failure_still_deletes_vector(self, memory_with_graph):
"""If graph cleanup fails, vector store deletion should still proceed."""
memory, mock_vs, mock_llm, conn = memory_with_graph
# Make LLM raise during entity extraction (graph cleanup will fail)
mock_llm.generate_response.side_effect = RuntimeError("LLM exploded")
mock_vs.get.return_value = MockVectorMemory(
"mem-1",
{"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"},
)
result = memory.delete("mem-1")
assert result == {"message": "Memory deleted successfully!"}
mock_vs.delete.assert_called_once_with(vector_id="mem-1")
def test_memory_delete_all_uses_bulk_not_per_memory(self, memory_with_graph):
"""delete_all() should use delete_all() on graph, not per-memory delete()."""
memory, mock_vs, mock_llm, conn = memory_with_graph
_setup_llm_for_memory_graph(
mock_llm,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
memory.graph.add("Alice likes Bob", {"user_id": "user-1"})
assert _edge_count(conn) == 1
# Set up vector store to return memories for deletion
mem1 = MockVectorMemory("mem-1", {"data": "Alice likes Bob", "user_id": "user-1"})
mock_vs.list.return_value = ([mem1], 1)
mock_vs.get.return_value = mem1
memory.delete_all(user_id="user-1")
# After delete_all, graph should be empty (via graph.delete_all)
assert _edge_count(conn) == 0
assert _node_count(conn) == 0
def test_memory_delete_nonexistent_raises_without_graph_side_effects(self, memory_with_graph):
"""Deleting a non-existent memory should raise ValueError without touching graph."""
memory, mock_vs, mock_llm, conn = memory_with_graph
# Add some graph data that should NOT be affected
_setup_llm_for_memory_graph(
mock_llm,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
memory.graph.add("Alice likes Bob", {"user_id": "user-1"})
assert _edge_count(conn) == 1
# Memory doesn't exist in vector store
mock_vs.get.return_value = None
with pytest.raises(ValueError, match="Memory with id non-existent not found"):
memory.delete("non-existent")
# Graph data should be untouched
assert _edge_count(conn) == 1
def _setup_llm_for_memory_graph(mock_llm, entities, relations):
"""Configure mock LLM for the Memory-level graph operations."""
def generate_response(messages, tools):
tool_names = []
for t in tools:
if isinstance(t, dict):
fn = t.get("function", t)
tool_names.append(fn.get("name", ""))
else:
tool_names.append(getattr(t, "name", str(t)))
if any("extract_entities" in n for n in tool_names):
return {
"tool_calls": [
{
"name": "extract_entities",
"arguments": {"entities": entities},
}
]
}
elif any("establish" in n or "relation" in n for n in tool_names):
return {
"tool_calls": [
{
"name": "establish_nodes_relations",
"arguments": {"entities": relations},
}
]
}
elif any("delete" in n for n in tool_names):
return {"tool_calls": []}
return {"tool_calls": []}
mock_llm.generate_response.side_effect = generate_response
+1 -3
View File
@@ -186,9 +186,7 @@ def test_delete(memory_instance):
result = memory_instance.delete("test_id")
# delete() now fetches the memory first and passes it to _delete_memory
existing_memory = memory_instance.vector_store.get.return_value
memory_instance._delete_memory.assert_called_once_with("test_id", existing_memory)
memory_instance._delete_memory.assert_called_once_with("test_id")
assert result["message"] == "Memory deleted successfully!"
-507
View File
@@ -1,507 +0,0 @@
"""Tests for REST API parameter forwarding.
Verifies that the Pydantic request models in server/main.py correctly accept
and forward all parameters supported by the underlying Memory class methods,
including limit, threshold, infer, memory_type, and prompt — which were
previously silently dropped by Pydantic v2's default extra='ignore' behavior.
"""
import importlib
import os
from unittest.mock import MagicMock, patch
import pytest
pytest.importorskip("fastapi", reason="fastapi not installed")
from fastapi.testclient import TestClient
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture
def _mock_memory():
"""Patch Memory.from_config so the server imports without a real backend."""
mock_instance = MagicMock()
mock_instance.add.return_value = {"results": [{"id": "mem-1", "event": "ADD", "memory": "test"}]}
mock_instance.search.return_value = [{"id": "mem-1", "memory": "test", "score": 0.9}]
mock_instance.get.return_value = {"id": "mem-1", "memory": "test memory"}
mock_instance.get_all.return_value = [{"id": "mem-1", "memory": "test memory"}]
mock_instance.update.return_value = {"message": "Memory updated"}
mock_instance.history.return_value = [{"id": "mem-1", "old_memory": "a", "new_memory": "b"}]
mock_instance.delete.return_value = None
mock_instance.delete_all.return_value = {"message": "Memories deleted"}
mock_instance.reset.return_value = None
with patch.dict(os.environ, {"OPENAI_API_KEY": "fake-key", "ADMIN_API_KEY": ""}):
with patch("mem0.Memory.from_config", return_value=mock_instance):
yield mock_instance
@pytest.fixture
def client(_mock_memory):
"""Return a TestClient wired to the server app with mocked Memory."""
import server.main as server_main
with patch.dict(os.environ, {"ADMIN_API_KEY": ""}):
importlib.reload(server_main)
return TestClient(server_main.app)
@pytest.fixture
def mock_memory(_mock_memory):
return _mock_memory
# ===========================================================================
# SearchRequest: limit parameter
# ===========================================================================
class TestSearchLimit:
"""Verify that the limit parameter is accepted and forwarded to Memory.search()."""
def test_limit_forwarded(self, client, mock_memory):
resp = client.post("/search", json={"query": "food", "user_id": "u1", "limit": 5})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["limit"] == 5
def test_limit_one(self, client, mock_memory):
resp = client.post("/search", json={"query": "food", "user_id": "u1", "limit": 1})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["limit"] == 1
def test_limit_omitted_uses_memory_default(self, client, mock_memory):
"""When limit is not sent, it should not appear in the kwargs,
allowing Memory.search() to use its own default (100)."""
resp = client.post("/search", json={"query": "food", "user_id": "u1"})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert "limit" not in kwargs
# ===========================================================================
# SearchRequest: threshold parameter
# ===========================================================================
class TestSearchThreshold:
"""Verify that the threshold parameter is accepted and forwarded."""
def test_threshold_forwarded(self, client, mock_memory):
resp = client.post("/search", json={"query": "food", "user_id": "u1", "threshold": 0.8})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["threshold"] == 0.8
def test_threshold_zero(self, client, mock_memory):
"""threshold=0.0 is a valid falsy value that must not be filtered out."""
resp = client.post("/search", json={"query": "food", "user_id": "u1", "threshold": 0.0})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["threshold"] == 0.0
def test_threshold_omitted_uses_memory_default(self, client, mock_memory):
resp = client.post("/search", json={"query": "food", "user_id": "u1"})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert "threshold" not in kwargs
# ===========================================================================
# SearchRequest: limit + threshold together
# ===========================================================================
class TestSearchLimitAndThreshold:
def test_both_forwarded(self, client, mock_memory):
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "limit": 10, "threshold": 0.5
})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["limit"] == 10
assert kwargs["threshold"] == 0.5
# ===========================================================================
# MemoryCreate: infer parameter
# ===========================================================================
class TestAddInfer:
"""Verify that the infer parameter is accepted and forwarded to Memory.add()."""
def test_infer_false_forwarded(self, client, mock_memory):
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "Store this exactly"}],
"user_id": "u1",
"infer": False,
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert kwargs["infer"] is False
def test_infer_true_forwarded(self, client, mock_memory):
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "I like pizza"}],
"user_id": "u1",
"infer": True,
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert kwargs["infer"] is True
def test_infer_omitted_uses_memory_default(self, client, mock_memory):
"""When infer is not sent, it should not appear in kwargs,
allowing Memory.add() to use its own default (True)."""
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "hello"}],
"user_id": "u1",
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert "infer" not in kwargs
# ===========================================================================
# MemoryCreate: memory_type parameter
# ===========================================================================
class TestAddMemoryType:
"""Verify that the memory_type parameter is accepted and forwarded."""
def test_memory_type_forwarded(self, client, mock_memory):
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "I like pizza"}],
"user_id": "u1",
"memory_type": "core",
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert kwargs["memory_type"] == "core"
def test_memory_type_omitted(self, client, mock_memory):
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "hello"}],
"user_id": "u1",
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert "memory_type" not in kwargs
# ===========================================================================
# MemoryCreate: prompt parameter
# ===========================================================================
class TestAddPrompt:
"""Verify that the prompt parameter is accepted and forwarded."""
def test_prompt_forwarded(self, client, mock_memory):
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "I like pizza"}],
"user_id": "u1",
"prompt": "Extract food preferences only.",
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert kwargs["prompt"] == "Extract food preferences only."
def test_prompt_omitted(self, client, mock_memory):
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "hello"}],
"user_id": "u1",
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert "prompt" not in kwargs
# ===========================================================================
# MemoryCreate: all new params together
# ===========================================================================
class TestAddAllNewParams:
def test_infer_memory_type_and_prompt_together(self, client, mock_memory):
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "I like pizza"}],
"user_id": "u1",
"infer": False,
"memory_type": "core",
"prompt": "Custom extraction prompt.",
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert kwargs["infer"] is False
assert kwargs["memory_type"] == "core"
assert kwargs["prompt"] == "Custom extraction prompt."
# ===========================================================================
# Edge cases: falsy-but-valid values must not be filtered out
# ===========================================================================
class TestFalsyValues:
"""The handler filters with `v is not None`. Falsy values like False, 0,
0.0, and empty string must still be forwarded."""
def test_infer_false_not_filtered(self, client, mock_memory):
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "test"}],
"user_id": "u1",
"infer": False,
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert kwargs["infer"] is False
def test_threshold_zero_not_filtered(self, client, mock_memory):
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "threshold": 0.0,
})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["threshold"] == 0.0
# ===========================================================================
# Extra/unknown fields are still silently ignored (existing Pydantic behavior)
# ===========================================================================
class TestUnknownFieldsIgnored:
def test_unknown_search_field_ignored(self, client, mock_memory):
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "bogus_field": "xyz",
})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert "bogus_field" not in kwargs
def test_unknown_add_field_ignored(self, client, mock_memory):
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "test"}],
"user_id": "u1",
"unknown_param": 42,
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert "unknown_param" not in kwargs
# ===========================================================================
# Backward compatibility: existing params still work
# ===========================================================================
class TestExistingParamsUnchanged:
def test_search_filters_still_forwarded(self, client, mock_memory):
resp = client.post("/search", json={
"query": "food",
"user_id": "u1",
"agent_id": "a1",
"filters": {"category": "food"},
})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["user_id"] == "u1"
assert kwargs["agent_id"] == "a1"
assert kwargs["filters"] == {"category": "food"}
def test_add_metadata_still_forwarded(self, client, mock_memory):
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "test"}],
"user_id": "u1",
"agent_id": "a1",
"metadata": {"source": "test"},
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert kwargs["user_id"] == "u1"
assert kwargs["agent_id"] == "a1"
assert kwargs["metadata"] == {"source": "test"}
# ===========================================================================
# OpenAPI schema: new fields are documented
# ===========================================================================
class TestOpenAPISchema:
"""Verify the new fields appear in the auto-generated OpenAPI schema."""
def test_search_schema_includes_limit(self, client):
schema = client.get("/openapi.json").json()
search_props = schema["components"]["schemas"]["SearchRequest"]["properties"]
assert "limit" in search_props
assert search_props["limit"]["description"] == "Maximum number of results to return."
def test_search_schema_includes_threshold(self, client):
schema = client.get("/openapi.json").json()
search_props = schema["components"]["schemas"]["SearchRequest"]["properties"]
assert "threshold" in search_props
def test_add_schema_includes_infer(self, client):
schema = client.get("/openapi.json").json()
add_props = schema["components"]["schemas"]["MemoryCreate"]["properties"]
assert "infer" in add_props
def test_add_schema_includes_memory_type(self, client):
schema = client.get("/openapi.json").json()
add_props = schema["components"]["schemas"]["MemoryCreate"]["properties"]
assert "memory_type" in add_props
def test_add_schema_includes_prompt(self, client):
schema = client.get("/openapi.json").json()
add_props = schema["components"]["schemas"]["MemoryCreate"]["properties"]
assert "prompt" in add_props
# ===========================================================================
# Pydantic type validation: invalid types return 422
# ===========================================================================
class TestTypeValidation:
"""Verify FastAPI/Pydantic rejects invalid types with 422."""
def test_limit_string_rejected(self, client):
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "limit": "not_a_number",
})
assert resp.status_code == 422
def test_threshold_string_rejected(self, client):
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "threshold": "high",
})
assert resp.status_code == 422
def test_infer_string_coerced_by_pydantic(self, client, mock_memory):
"""Pydantic v2 coerces truthy strings like 'yes' to True for bool fields."""
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "test"}],
"user_id": "u1",
"infer": "yes",
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert kwargs["infer"] is True
def test_infer_invalid_value_rejected(self, client):
"""A value that cannot be coerced to bool should be rejected."""
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "test"}],
"user_id": "u1",
"infer": [1, 2, 3],
})
assert resp.status_code == 422
def test_limit_float_rejected(self, client):
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "limit": 5.7,
})
assert resp.status_code == 422
def test_memory_type_int_rejected(self, client):
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "test"}],
"user_id": "u1",
"memory_type": 123,
})
assert resp.status_code == 422
# ===========================================================================
# Explicit null values: treated as omitted (filtered by `is not None`)
# ===========================================================================
class TestExplicitNull:
"""When a client sends null for an optional field, it should be treated
as omitted — the Memory class default should be used."""
def test_limit_null_uses_memory_default(self, client, mock_memory):
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "limit": None,
})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert "limit" not in kwargs
def test_infer_null_uses_memory_default(self, client, mock_memory):
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "test"}],
"user_id": "u1",
"infer": None,
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert "infer" not in kwargs
def test_prompt_null_uses_memory_default(self, client, mock_memory):
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "test"}],
"user_id": "u1",
"prompt": None,
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
assert "prompt" not in kwargs
# ===========================================================================
# Verify exact call signatures match Memory method params
# ===========================================================================
class TestCallSignatureMatch:
"""Ensure forwarded params exactly match Memory.add() and Memory.search()
keyword argument names — a typo here would cause a TypeError at runtime."""
def test_search_kwargs_are_valid(self, client, mock_memory):
"""All kwargs forwarded to Memory.search() must be in its signature."""
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "agent_id": "a1",
"run_id": "r1", "filters": {"k": "v"},
"limit": 10, "threshold": 0.5,
})
assert resp.status_code == 200
# The handler passes query= as a keyword arg, so it appears in kwargs too
_, kwargs = mock_memory.search.call_args
valid_params = {"query", "user_id", "agent_id", "run_id", "limit", "filters", "threshold", "rerank"}
for key in kwargs:
assert key in valid_params, f"Unexpected kwarg '{key}' forwarded to Memory.search()"
def test_add_kwargs_are_valid(self, client, mock_memory):
"""All kwargs forwarded to Memory.add() must be in its signature."""
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "hi"}],
"user_id": "u1", "agent_id": "a1", "run_id": "r1",
"metadata": {"k": "v"},
"infer": False, "memory_type": "core", "prompt": "custom",
})
assert resp.status_code == 200
# The handler passes messages= as a keyword arg, so it appears in kwargs too
_, kwargs = mock_memory.add.call_args
valid_params = {"messages", "user_id", "agent_id", "run_id", "metadata", "infer", "memory_type", "prompt"}
for key in kwargs:
assert key in valid_params, f"Unexpected kwarg '{key}' forwarded to Memory.add()"
def test_messages_excluded_from_params_dict(self, client, mock_memory):
"""messages is passed separately via messages= kwarg, not duplicated from model_dump."""
resp = client.post("/memories", json={
"messages": [{"role": "user", "content": "hi"}],
"user_id": "u1",
})
assert resp.status_code == 200
_, kwargs = mock_memory.add.call_args
# messages should be present (passed explicitly) and be a list of dicts
assert "messages" in kwargs
assert isinstance(kwargs["messages"], list)
assert kwargs["messages"][0] == {"role": "user", "content": "hi"}
def test_query_passed_explicitly(self, client, mock_memory):
"""query is passed as an explicit keyword arg to Memory.search()."""
resp = client.post("/search", json={"query": "food", "user_id": "u1"})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["query"] == "food"
-13
View File
@@ -3,7 +3,6 @@ from unittest.mock import Mock, patch
import pytest
from mem0.vector_stores.chroma import ChromaDB
from mem0.configs.vector_stores.chroma import ChromaDbConfig
@pytest.fixture
@@ -250,15 +249,3 @@ def test_generate_where_clause_non_string_values():
# ChromaDB accepts non-string values in filters
expected = {"$and": [{"user_id": {"$eq": "alice"}}, {"count": {"$eq": 5}}, {"active": {"$eq": True}}]}
assert result == expected
def test_chroma_config_accepts_default_tmp_path():
"""Test that ChromaDbConfig accepts the default /tmp/chroma path."""
config = ChromaDbConfig(path="/tmp/chroma")
assert config.path == "/tmp/chroma"
def test_chroma_config_rejects_no_config():
"""Test that ChromaDbConfig rejects when no connection config is provided."""
with pytest.raises(ValueError):
ChromaDbConfig()
+2 -499
View File
@@ -1,12 +1,8 @@
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
pytest.importorskip("databricks", reason="databricks-sdk package not installed")
from databricks.sdk.service.vectorsearch import VectorIndexType, QueryVectorIndexResponse, ResultManifest, ResultData, ColumnInfo
from mem0.vector_stores.databricks import Databricks
import pytest
# ---------------------- Fixtures ---------------------- #
@@ -209,34 +205,9 @@ def test_search_direct_access_vector(db_instance_direct, mock_workspace_client):
assert results[0].score == 0.77
def test_search_delta_sync_self_managed_vectors(mock_workspace_client):
"""DELTA_SYNC without embedding model endpoint should use query_vector, not query_text."""
mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=True)
# DELTA_SYNC without embedding_model_endpoint_name = self-managed vectors
inst = Databricks(
workspace_url="https://test",
access_token="tok",
endpoint_name="vs-endpoint",
catalog="catalog",
schema="schema",
table_name="table",
warehouse_name="test-warehouse",
index_type=VectorIndexType.DELTA_SYNC,
embedding_dimension=4,
# NOTE: no embedding_model_endpoint_name
)
mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace(
result=SimpleNamespace(data_array=[])
)
inst.search(query="ignored", vectors=[0.1, 0.2, 0.3, 0.4], limit=5)
call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert "query_vector" in call_kwargs
assert "query_text" not in call_kwargs
def test_search_missing_params_raises(db_instance_delta):
with pytest.raises(ValueError):
db_instance_delta.search(query="", vectors=[0.1, 0.2]) # DELTA_SYNC with model endpoint requires query text
db_instance_delta.search(query="", vectors=[0.1, 0.2]) # DELTA_SYNC requires query text
# ---------------------- Delete Tests ---------------------- #
@@ -304,54 +275,6 @@ def test_get_vector(db_instance_delta, mock_workspace_client):
assert res.id == "id-get"
assert res.payload["data"] == "some memory"
assert res.payload["tag"] == "x"
# DELTA_SYNC should use query_text, not query_vector
call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert "query_text" in call_kwargs
assert "query_vector" not in call_kwargs
def test_get_vector_direct_access(db_instance_direct, mock_workspace_client):
"""get() on a DIRECT_ACCESS index must use query_vector instead of query_text."""
mock_workspace_client.vector_search_indexes.query_index.return_value = QueryVectorIndexResponse(
manifest=ResultManifest(columns=[
ColumnInfo(name="memory_id"),
ColumnInfo(name="hash"),
ColumnInfo(name="agent_id"),
ColumnInfo(name="run_id"),
ColumnInfo(name="user_id"),
ColumnInfo(name="memory"),
ColumnInfo(name="metadata"),
ColumnInfo(name="created_at"),
ColumnInfo(name="updated_at"),
ColumnInfo(name="embedding"),
ColumnInfo(name="score"),
]),
result=ResultData(
data_array=[
[
"id-get-da",
"h",
"a",
"r",
"u",
"direct access memory",
'{"tag":"da"}',
"2024-01-01T00:00:00",
"2024-01-01T00:00:00",
[0.1, 0.2, 0.3, 0.4],
"0.88",
]
]
)
)
res = db_instance_direct.get("id-get-da")
assert res.id == "id-get-da"
assert res.payload["data"] == "direct access memory"
# DIRECT_ACCESS should use query_vector, not query_text
call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert "query_vector" in call_kwargs
assert "query_text" not in call_kwargs
assert call_kwargs["query_vector"] == [0.0] * 4 # embedding_dimension=4
# ---------------------- Collection Info / Listing Tests ---------------------- #
@@ -407,185 +330,6 @@ def test_list_memories(db_instance_delta, mock_workspace_client):
assert isinstance(res, list)
assert len(res[0]) == 1
assert res[0][0].id == "id-get"
# DELTA_SYNC should use query_text
call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert "query_text" in call_kwargs
assert "query_vector" not in call_kwargs
def test_list_memories_direct_access(db_instance_direct, mock_workspace_client):
"""list() on a DIRECT_ACCESS index must use query_vector instead of query_text."""
mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace(
result=SimpleNamespace(
data_array=[
[
"id-da-list",
"h",
"a",
"r",
"u",
"direct memory",
None,
"2024-01-01T00:00:00",
"2024-01-01T00:00:00",
[0.1, 0.2, 0.3, 0.4],
]
]
)
)
res = db_instance_direct.list(limit=5)
assert isinstance(res, list)
assert len(res[0]) == 1
assert res[0][0].id == "id-da-list"
assert res[0][0].payload["data"] == "direct memory"
# DIRECT_ACCESS should use query_vector
call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert "query_vector" in call_kwargs
assert "query_text" not in call_kwargs
assert call_kwargs["query_vector"] == [0.0] * 4
def test_get_vector_delta_sync_self_managed(mock_workspace_client):
"""get() on DELTA_SYNC without model endpoint should use query_vector."""
mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=True)
inst = Databricks(
workspace_url="https://test",
access_token="tok",
endpoint_name="vs-endpoint",
catalog="catalog",
schema="schema",
table_name="table",
warehouse_name="test-warehouse",
index_type=VectorIndexType.DELTA_SYNC,
embedding_dimension=4,
# NOTE: no embedding_model_endpoint_name
)
mock_workspace_client.vector_search_indexes.query_index.return_value = QueryVectorIndexResponse(
manifest=ResultManifest(columns=[
ColumnInfo(name="memory_id"), ColumnInfo(name="hash"),
ColumnInfo(name="agent_id"), ColumnInfo(name="run_id"),
ColumnInfo(name="user_id"), ColumnInfo(name="memory"),
ColumnInfo(name="metadata"), ColumnInfo(name="created_at"),
ColumnInfo(name="updated_at"),
]),
result=ResultData(data_array=[["id-sm", "h", None, None, None, "self-managed mem", None, None, None]]),
)
res = inst.get("id-sm")
assert res.id == "id-sm"
call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert "query_vector" in call_kwargs
assert "query_text" not in call_kwargs
assert call_kwargs["query_vector"] == [0.0] * 4
def test_list_memories_delta_sync_self_managed(mock_workspace_client):
"""list() on DELTA_SYNC without model endpoint should use query_vector."""
mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=True)
inst = Databricks(
workspace_url="https://test",
access_token="tok",
endpoint_name="vs-endpoint",
catalog="catalog",
schema="schema",
table_name="table",
warehouse_name="test-warehouse",
index_type=VectorIndexType.DELTA_SYNC,
embedding_dimension=4,
# NOTE: no embedding_model_endpoint_name
)
mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace(
result=SimpleNamespace(data_array=[])
)
inst.list(limit=5)
call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert "query_vector" in call_kwargs
assert "query_text" not in call_kwargs
def test_list_memories_default_limit(db_instance_delta, mock_workspace_client):
"""list() with no limit should default to 100."""
mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace(
result=SimpleNamespace(data_array=[])
)
db_instance_delta.list(limit=None)
call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert call_kwargs["num_results"] == 100
# ---------------------- Table Creation Tests ---------------------- #
def test_ensure_source_table_uses_dynamic_names(mock_workspace_client):
"""Verify _ensure_source_table_exists uses self.fully_qualified_table_name and
self.table_name for the PK constraint, not hardcoded values."""
mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=False)
Databricks(
workspace_url="https://test",
access_token="tok",
endpoint_name="vs-endpoint",
catalog="my_catalog",
schema="my_schema",
table_name="my_memories",
collection_name="my_index",
warehouse_name="test-warehouse",
index_type=VectorIndexType.DELTA_SYNC,
embedding_model_endpoint_name="embedding-endpoint",
)
# _ensure_source_table_exists was called during __init__ via create_col
constraint_call = mock_workspace_client.table_constraints.create.call_args
assert constraint_call.kwargs["full_name_arg"] == "my_catalog.my_schema.my_memories"
pk_name = constraint_call.kwargs["constraint"].primary_key_constraint.name
assert pk_name == "pk_my_memories"
# ---------------------- Config Validation Tests ---------------------- #
def test_config_rejects_old_doc_params():
"""Config should reject the old documentation parameter names like index_name and source_table_name."""
from mem0.configs.vector_stores.databricks import DatabricksConfig
with pytest.raises(ValueError, match="Extra fields not allowed"):
DatabricksConfig(
workspace_url="https://test",
access_token="tok",
endpoint_name="ep",
catalog="cat",
schema="sch",
table_name="tbl",
index_name="catalog.schema.index", # old param from docs
)
def test_config_rejects_source_table_name():
"""Config should reject source_table_name which was in old docs."""
from mem0.configs.vector_stores.databricks import DatabricksConfig
with pytest.raises(ValueError, match="Extra fields not allowed"):
DatabricksConfig(
workspace_url="https://test",
access_token="tok",
endpoint_name="ep",
catalog="cat",
schema="sch",
table_name="tbl",
source_table_name="catalog.schema.table", # old param from docs
)
def test_config_accepts_correct_params():
"""Config should accept all the correct parameter names."""
from mem0.configs.vector_stores.databricks import DatabricksConfig
config = DatabricksConfig(
workspace_url="https://test",
access_token="tok",
endpoint_name="ep",
catalog="cat",
schema="sch",
table_name="tbl",
collection_name="my_index",
embedding_dimension=768,
)
assert config.collection_name == "my_index"
assert config.embedding_dimension == 768
# ---------------------- Reset Tests ---------------------- #
@@ -597,244 +341,3 @@ def test_reset(db_instance_delta, mock_workspace_client):
with patch.object(db_instance_delta, "create_col", wraps=db_instance_delta.create_col) as create_spy:
db_instance_delta.reset()
assert create_spy.called
# ---------------------- End-to-End Config → Factory → CRUD Tests ---------------------- #
def test_e2e_config_to_factory_delta_sync(mock_workspace_client):
"""End-to-end: VectorStoreConfig validates docs-correct params, factory creates Databricks instance."""
from mem0.vector_stores.configs import VectorStoreConfig
from mem0.utils.factory import VectorStoreFactory
# Step 1: Config validation (simulates what Memory.from_config does)
vs_config = VectorStoreConfig(
provider="databricks",
config={
"workspace_url": "https://my-workspace.databricks.com",
"access_token": "my-token",
"endpoint_name": "my-endpoint",
"catalog": "prod_catalog",
"schema": "ai_schema",
"table_name": "memories_table",
"collection_name": "my_index",
"embedding_dimension": 768,
"warehouse_name": "test-warehouse",
},
)
assert vs_config.config.collection_name == "my_index"
assert vs_config.config.catalog == "prod_catalog"
# Step 2: Factory instantiation (same as MemoryBase.__init__)
instance = VectorStoreFactory.create("databricks", vs_config.config)
assert isinstance(instance, Databricks)
assert instance.fully_qualified_table_name == "prod_catalog.ai_schema.memories_table"
assert instance.fully_qualified_index_name == "prod_catalog.ai_schema.my_index"
assert instance.embedding_dimension == 768
def test_e2e_config_to_factory_direct_access(mock_workspace_client):
"""End-to-end: DIRECT_ACCESS via config → factory creates correct instance."""
from mem0.vector_stores.configs import VectorStoreConfig
from mem0.utils.factory import VectorStoreFactory
mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=True)
vs_config = VectorStoreConfig(
provider="databricks",
config={
"workspace_url": "https://my-workspace.databricks.com",
"access_token": "my-token",
"endpoint_name": "my-endpoint",
"catalog": "cat",
"schema": "sch",
"table_name": "tbl",
"index_type": "DIRECT_ACCESS",
"embedding_dimension": 4,
"warehouse_name": "test-warehouse",
},
)
instance = VectorStoreFactory.create("databricks", vs_config.config)
assert isinstance(instance, Databricks)
assert "embedding" in instance.column_names
def test_e2e_old_docs_config_rejected():
"""End-to-end: Config from old docs (with index_name, source_table_name) is rejected at validation."""
from mem0.vector_stores.configs import VectorStoreConfig
with pytest.raises(ValueError, match="Extra fields not allowed"):
VectorStoreConfig(
provider="databricks",
config={
"workspace_url": "https://my-workspace.databricks.com",
"access_token": "my-token",
"endpoint_name": "my-endpoint",
"index_name": "catalog.schema.index_name",
"source_table_name": "catalog.schema.source_table",
"embedding_dimension": 1536,
},
)
def test_e2e_crud_lifecycle_delta_sync(mock_workspace_client):
"""End-to-end CRUD lifecycle: insert → search → get → list → update → delete."""
from mem0.vector_stores.configs import VectorStoreConfig
from mem0.utils.factory import VectorStoreFactory
vs_config = VectorStoreConfig(
provider="databricks",
config={
"workspace_url": "https://test",
"access_token": "tok",
"endpoint_name": "ep",
"catalog": "cat",
"schema": "sch",
"table_name": "tbl",
"warehouse_name": "test-warehouse",
"embedding_model_endpoint_name": "emb-ep",
},
)
db = VectorStoreFactory.create("databricks", vs_config.config)
# INSERT
db.insert(
vectors=[[0.1, 0.2]],
payloads=[{"data": "test memory", "user_id": "u1", "hash": "h1"}],
ids=["mem-001"],
)
insert_sql = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs["statement"]
assert "INSERT INTO cat.sch.tbl" in insert_sql
assert "mem-001" in insert_sql
# SEARCH
mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace(
result=SimpleNamespace(
data_array=[["mem-001", "h1", None, None, "u1", "test memory", None, None, None, 0.95]]
)
)
results = db.search(query="test", vectors=None, limit=5)
assert len(results) == 1
assert results[0].id == "mem-001"
assert results[0].payload["data"] == "test memory"
search_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert search_kwargs["query_text"] == "test"
# GET
mock_workspace_client.vector_search_indexes.query_index.return_value = QueryVectorIndexResponse(
manifest=ResultManifest(columns=[
ColumnInfo(name="memory_id"), ColumnInfo(name="hash"),
ColumnInfo(name="agent_id"), ColumnInfo(name="run_id"),
ColumnInfo(name="user_id"), ColumnInfo(name="memory"),
ColumnInfo(name="metadata"), ColumnInfo(name="created_at"),
ColumnInfo(name="updated_at"),
]),
result=ResultData(data_array=[["mem-001", "h1", None, None, "u1", "test memory", None, None, None]]),
)
got = db.get("mem-001")
assert got.id == "mem-001"
assert got.payload["data"] == "test memory"
get_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert "query_text" in get_kwargs
assert "query_vector" not in get_kwargs
# LIST
mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace(
result=SimpleNamespace(
data_array=[["mem-001", "h1", None, None, "u1", "test memory", None, None, None]]
)
)
listed = db.list(filters={"user_id": "u1"}, limit=10)
assert len(listed[0]) == 1
assert listed[0][0].id == "mem-001"
list_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert "query_text" in list_kwargs
assert list_kwargs["num_results"] == 10
# UPDATE
db.update(vector_id="mem-001", payload={"memory": "updated memory"})
update_sql = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs["statement"]
assert "UPDATE cat.sch.tbl" in update_sql
assert "mem-001" in update_sql
assert "updated memory" in update_sql
# DELETE
db.delete("mem-001")
delete_sql = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs["statement"]
assert "DELETE FROM cat.sch.tbl" in delete_sql
assert "mem-001" in delete_sql
def test_e2e_crud_lifecycle_direct_access(mock_workspace_client):
"""End-to-end CRUD lifecycle for DIRECT_ACCESS: insert → search → get → list."""
from mem0.vector_stores.configs import VectorStoreConfig
from mem0.utils.factory import VectorStoreFactory
mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=True)
vs_config = VectorStoreConfig(
provider="databricks",
config={
"workspace_url": "https://test",
"access_token": "tok",
"endpoint_name": "ep",
"catalog": "cat",
"schema": "sch",
"table_name": "tbl",
"index_type": "DIRECT_ACCESS",
"embedding_dimension": 4,
"warehouse_name": "test-warehouse",
"embedding_model_endpoint_name": "emb-ep",
},
)
db = VectorStoreFactory.create("databricks", vs_config.config)
assert "embedding" in db.column_names
# INSERT with vector
db.insert(
vectors=[[0.1, 0.2, 0.3, 0.4]],
payloads=[{"data": "direct memory", "user_id": "u1", "hash": "h1"}],
ids=["mem-da-001"],
)
insert_sql = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs["statement"]
assert "array(0.1, 0.2, 0.3, 0.4)" in insert_sql
# SEARCH with vector
mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace(
result=SimpleNamespace(
data_array=[["mem-da-001", "h1", None, None, "u1", "direct memory", None, None, None, [0.1, 0.2, 0.3, 0.4], 0.9]]
)
)
results = db.search(query="", vectors=[0.1, 0.2, 0.3, 0.4], limit=5)
assert len(results) == 1
search_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert "query_vector" in search_kwargs
assert "query_text" not in search_kwargs
# GET — must use query_vector for DIRECT_ACCESS
mock_workspace_client.vector_search_indexes.query_index.return_value = QueryVectorIndexResponse(
manifest=ResultManifest(columns=[
ColumnInfo(name="memory_id"), ColumnInfo(name="hash"),
ColumnInfo(name="agent_id"), ColumnInfo(name="run_id"),
ColumnInfo(name="user_id"), ColumnInfo(name="memory"),
ColumnInfo(name="metadata"), ColumnInfo(name="created_at"),
ColumnInfo(name="updated_at"), ColumnInfo(name="embedding"),
]),
result=ResultData(data_array=[["mem-da-001", "h1", None, None, "u1", "direct memory", None, None, None, [0.1, 0.2, 0.3, 0.4]]]),
)
got = db.get("mem-da-001")
assert got.id == "mem-da-001"
get_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert "query_vector" in get_kwargs
assert get_kwargs["query_vector"] == [0.0] * 4
# LIST — must use query_vector for DIRECT_ACCESS
mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace(
result=SimpleNamespace(
data_array=[["mem-da-001", "h1", None, None, "u1", "direct memory", None, None, None, [0.1, 0.2, 0.3, 0.4]]]
)
)
db.list(limit=5)
list_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs
assert "query_vector" in list_kwargs
assert "query_text" not in list_kwargs
+12 -55
View File
@@ -48,16 +48,17 @@ def test_initalize_create_col(mongo_vector_fixture):
search_index_model = args[0].document
assert search_index_model == {
"name": "test_collection_vector_index",
"type": "vectorSearch",
"definition": {
"fields": [
{
"type": "vector",
"path": "embedding",
"numDimensions": 1536,
"similarity": "cosine",
}
]
"mappings": {
"dynamic": False,
"fields": {
"embedding": {
"type": "knnVector",
"dimensions": 1536,
"similarity": "cosine",
}
},
}
},
}
assert mongo_vector.collection == mock_collection
@@ -94,7 +95,7 @@ def test_search(mongo_vector_fixture):
"$vectorSearch": {
"index": "test_collection_vector_index",
"limit": 2,
"numCandidates": 40,
"numCandidates": 2,
"queryVector": query_vector,
"path": "embedding",
},
@@ -203,10 +204,6 @@ def test_delete(mongo_vector_fixture):
def test_update(mongo_vector_fixture):
"""
Test that update() uses dot notation for payload fields instead of replacing
the entire payload document.
"""
mongo_vector, mock_collection, _ = mongo_vector_fixture
vector_id = "id1"
updated_vector = [0.3] * 1536
@@ -216,51 +213,11 @@ def test_update(mongo_vector_fixture):
mongo_vector.update(vector_id=vector_id, vector=updated_vector, payload=updated_payload)
# Should use dot notation (payload.name) instead of full replacement (payload)
mock_collection.update_one.assert_called_once_with(
{"_id": vector_id}, {"$set": {"embedding": updated_vector, "payload.name": "updated_vector"}}
{"_id": vector_id}, {"$set": {"embedding": updated_vector, "payload": updated_payload}}
)
def test_update_payload_only_uses_dot_notation(mongo_vector_fixture):
"""
Test that updating only the payload uses dot notation for each field,
preserving existing metadata fields not included in the update.
"""
mongo_vector, mock_collection, _ = mongo_vector_fixture
mock_collection.update_one.return_value = MagicMock(matched_count=1)
mongo_vector.update(
vector_id="id1",
payload={"data": "updated text", "hash": "def456", "updated_at": "2025-06-01"},
)
set_arg = mock_collection.update_one.call_args[0][1]["$set"]
# Only the specified fields should be in $set, using dot notation
assert set_arg == {
"payload.data": "updated text",
"payload.hash": "def456",
"payload.updated_at": "2025-06-01",
}
# "payload" key itself should NOT appear (that would replace the whole document)
assert "payload" not in set_arg
def test_update_vector_only_does_not_touch_payload(mongo_vector_fixture):
"""Test that updating only the vector does not touch any payload fields."""
mongo_vector, mock_collection, _ = mongo_vector_fixture
mock_collection.update_one.return_value = MagicMock(matched_count=1)
mongo_vector.update(vector_id="id1", vector=[0.7, 0.8, 0.9])
set_arg = mock_collection.update_one.call_args[0][1]["$set"]
assert set_arg == {"embedding": [0.7, 0.8, 0.9]}
assert not any(k.startswith("payload.") for k in set_arg)
def test_get(mongo_vector_fixture):
mongo_vector, mock_collection, _ = mongo_vector_fixture
vector_id = "id1"
+1 -20
View File
@@ -1,8 +1,6 @@
import os
import tempfile
import unittest
import uuid
from unittest.mock import MagicMock, patch
from unittest.mock import MagicMock
from qdrant_client import QdrantClient
from qdrant_client.models import (
@@ -33,23 +31,6 @@ class TestQdrant(unittest.TestCase):
on_disk=True,
)
def test_local_path_on_disk_false_preserves_existing_directory(self):
"""#4473: local path must not be removed when on_disk is False."""
with tempfile.TemporaryDirectory() as tmp:
sentinel = os.path.join(tmp, "sentinel")
with open(sentinel, "w", encoding="utf-8") as f:
f.write("keep")
mock_client = MagicMock()
mock_client.get_collections.return_value = MagicMock(collections=[])
with patch("mem0.vector_stores.qdrant.QdrantClient", return_value=mock_client):
Qdrant(
collection_name="c",
embedding_model_dims=128,
path=tmp,
on_disk=False,
)
self.assertTrue(os.path.isfile(sentinel))
def test_create_col(self):
self.client_mock.get_collections.return_value = MagicMock(collections=[])