Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| cf2d42a7fd |
@@ -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 🎉!
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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 🎉!
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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.
|
||||
|
||||
@@ -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`
|
||||
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
Vendored
+2
-2
@@ -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
@@ -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>
|
||||
```
|
||||
|
||||
@@ -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>
|
||||
@@ -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),
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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.")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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]],
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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"):
|
||||
|
||||
@@ -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 []
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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({
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
@@ -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"
|
||||
|
||||
|
||||
@@ -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"})
|
||||
@@ -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
|
||||
@@ -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
@@ -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!"
|
||||
|
||||
|
||||
|
||||
@@ -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"
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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=[])
|
||||
|
||||
|
||||
Reference in New Issue
Block a user