Compare commits
23 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 5e7adc4d12 | |||
| 14393b5962 | |||
| f590c9596a | |||
| dd5f7e39a8 | |||
| 70ab76a053 | |||
| 2af3a72f73 | |||
| 8005e18aec | |||
| 39551145b8 | |||
| c2bc28e589 | |||
| fec2fe6a2c | |||
| d0c23a5950 | |||
| b05cce581b | |||
| 8a57967c8a | |||
| 756b0b1b6d | |||
| 726bcc80b2 | |||
| 9383e9a255 | |||
| ddaa655edf | |||
| 739534c0a3 | |||
| 633b035342 | |||
| ccbe5861a1 | |||
| 50c3cf44f1 | |||
| d6d2588ef5 | |||
| 6c1741e3a4 |
@@ -12,7 +12,7 @@
|
||||
"name": "mem0",
|
||||
"source": "./integrations/mem0-plugin",
|
||||
"description": "Mem0 memory layer for AI applications. Add persistent memory, personalization, and semantic search to Claude workflows.",
|
||||
"version": "0.2.12"
|
||||
"version": "0.2.13"
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
"name": "mem0",
|
||||
"source": "./integrations/mem0-plugin",
|
||||
"description": "Mem0 memory layer for AI applications. Add persistent memory, personalization, and semantic search.",
|
||||
"version": "0.2.12"
|
||||
"version": "0.2.13"
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
@@ -9,12 +9,10 @@ body:
|
||||
label: Component
|
||||
description: Which part of mem0 is affected?
|
||||
options:
|
||||
- Core / Python SDK
|
||||
- Python SDK
|
||||
- TypeScript SDK
|
||||
- Vector Store (Qdrant, PGVector, Redis, Chroma, etc.)
|
||||
- Graph Memory (Neo4j, Memgraph, etc.)
|
||||
- Ollama / Local Models
|
||||
- OpenClaw
|
||||
- Vector Store
|
||||
- Plugin
|
||||
- REST API
|
||||
- Other
|
||||
validations:
|
||||
|
||||
@@ -9,14 +9,11 @@ body:
|
||||
label: Component
|
||||
description: Which part of mem0 does this relate to?
|
||||
options:
|
||||
- Core / Python SDK
|
||||
- Python SDK
|
||||
- TypeScript SDK
|
||||
- Vector Store (Qdrant, PGVector, Redis, Chroma, etc.)
|
||||
- Graph Memory (Neo4j, Memgraph, etc.)
|
||||
- Ollama / Local Models
|
||||
- OpenClaw
|
||||
- Vector Store
|
||||
- Plugin
|
||||
- REST API
|
||||
- Benchmarks / Evals
|
||||
- Other
|
||||
validations:
|
||||
required: true
|
||||
|
||||
@@ -1,18 +1,15 @@
|
||||
# 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: "openclaw"
|
||||
matcher: "OpenClaw"
|
||||
- label: "rest-api"
|
||||
matcher: "REST API"
|
||||
policy:
|
||||
- section:
|
||||
- id: ['component']
|
||||
block-list: ['Other']
|
||||
label:
|
||||
- name: 'sdk-python'
|
||||
keys: ['Python SDK']
|
||||
- name: 'sdk-typescript'
|
||||
keys: ['TypeScript SDK']
|
||||
- name: 'vector-store'
|
||||
keys: ['Vector Store']
|
||||
- name: 'plugin'
|
||||
keys: ['Plugin']
|
||||
- name: 'rest-api'
|
||||
keys: ['REST API']
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
{
|
||||
"language": {
|
||||
"sdk-python": [
|
||||
"python", "pip install", "pypi", "pyproject", "requirements.txt",
|
||||
"from mem0", "import mem0", "traceback", "pydantic", "asyncmemory",
|
||||
"poetry", "virtualenv", "venv", "conda", "pytest", "async def"
|
||||
],
|
||||
"sdk-typescript": [
|
||||
"typescript", "javascript", "pnpm", "yarn", "node.js", "nodejs",
|
||||
"mem0-ts", "mem0ai/oss", "tsconfig", "await import",
|
||||
"=> {", "undefined is not"
|
||||
]
|
||||
},
|
||||
"area": {
|
||||
"plugin": [
|
||||
"openclaw", "openclaw-mem0", "openclaw.json", "openclaw plugin",
|
||||
"claude code", "opencode", "pi agent", "mem0-plugin",
|
||||
"cursor plugin", "codex plugin", "editor plugin"
|
||||
],
|
||||
"openmemory": [
|
||||
"openmemory", "open memory", "localhost:8765", "localhost:3000",
|
||||
"openmemory ui", "openmemory/api", "openmemory/ui"
|
||||
],
|
||||
"cli": ["mem0-cli", "@mem0/cli", "npx mem0", "command line"],
|
||||
"vector-store": [
|
||||
"pgvector", "pinecone", "chroma", "chromadb", "weaviate",
|
||||
"milvus", "faiss", "vector store", "vectorstore",
|
||||
"elasticsearch", "supabase", "azure ai search",
|
||||
"s3 vectors", "mongodb"
|
||||
],
|
||||
"integrations": [
|
||||
"vercel ai", "vercel-ai-sdk", "@mem0/vercel-ai-provider",
|
||||
"llamaindex", "crewai", "autogen", "langgraph"
|
||||
],
|
||||
"rest-api": [
|
||||
"rest api", "fastapi", "docker-compose", "/v1/memories",
|
||||
"localhost:8000", "localhost:8888", "curl -x", "http endpoint"
|
||||
],
|
||||
"documentation": [
|
||||
"docs.mem0.ai", "documentation", "typo", "readme", "docstring", "broken link",
|
||||
"issue on docs", "docs:", "link to the docs page", "issue with current documentation"
|
||||
]
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
sdk-python:
|
||||
- changed-files:
|
||||
- any-glob-to-any-file:
|
||||
- 'mem0/**'
|
||||
- 'tests/**'
|
||||
- 'cli/python/**'
|
||||
- 'pyproject.toml'
|
||||
- 'poetry.lock'
|
||||
|
||||
sdk-typescript:
|
||||
- changed-files:
|
||||
- any-glob-to-any-file:
|
||||
- 'mem0-ts/**'
|
||||
- 'cli/node/**'
|
||||
|
||||
vector-store:
|
||||
- changed-files:
|
||||
- any-glob-to-any-file:
|
||||
- 'mem0/vector_stores/**'
|
||||
- 'mem0-ts/src/oss/src/vector_stores/**'
|
||||
|
||||
rest-api:
|
||||
- changed-files:
|
||||
- any-glob-to-any-file: 'server/**'
|
||||
|
||||
openmemory:
|
||||
- changed-files:
|
||||
- any-glob-to-any-file: 'openmemory/**'
|
||||
|
||||
integrations:
|
||||
- changed-files:
|
||||
- any-glob-to-any-file: 'integrations/**'
|
||||
|
||||
plugin:
|
||||
- changed-files:
|
||||
- all-globs-to-any-file:
|
||||
- 'integrations/**'
|
||||
- '!integrations/vercel-ai-sdk/**'
|
||||
- any-glob-to-any-file:
|
||||
- 'skills/**'
|
||||
- '.agents/**'
|
||||
- '.claude-plugin/**'
|
||||
- '.codex-plugin/**'
|
||||
- '.cursor-plugin/**'
|
||||
- 'marketplace.json'
|
||||
|
||||
cli:
|
||||
- changed-files:
|
||||
- any-glob-to-any-file: 'cli/**'
|
||||
|
||||
documentation:
|
||||
- changed-files:
|
||||
- any-glob-to-any-file:
|
||||
- 'docs/**'
|
||||
- 'examples/**'
|
||||
- '*.md'
|
||||
|
||||
ci:
|
||||
- changed-files:
|
||||
- any-glob-to-any-file:
|
||||
- '.github/**'
|
||||
- 'scripts/**'
|
||||
- '.pre-commit-config.yaml'
|
||||
@@ -0,0 +1,44 @@
|
||||
const fs = require('fs');
|
||||
|
||||
function componentLabels(keywords) {
|
||||
return Object.values(keywords).flatMap(Object.keys);
|
||||
}
|
||||
|
||||
function toMatcher(term) {
|
||||
const escaped = term.replace(/[.*+?^${}()|[\]\\]/g, '\\$&');
|
||||
const prefix = /^[a-z0-9]/i.test(term) ? '\\b' : '';
|
||||
return new RegExp(prefix + escaped, 'i');
|
||||
}
|
||||
|
||||
function scoreGroup(text, group) {
|
||||
let winner = null;
|
||||
let best = 0;
|
||||
for (const [label, terms] of Object.entries(group)) {
|
||||
const score = terms.reduce((n, term) => n + (toMatcher(term).test(text) ? 1 : 0), 0);
|
||||
if (score > best) {
|
||||
winner = label;
|
||||
best = score;
|
||||
}
|
||||
}
|
||||
return winner;
|
||||
}
|
||||
|
||||
const UMBRELLA = { plugin: 'integrations' };
|
||||
|
||||
function inferComponentLabels(text, keywords) {
|
||||
if (!text) return [];
|
||||
const labels = [scoreGroup(text, keywords.language), scoreGroup(text, keywords.area)].filter(
|
||||
Boolean,
|
||||
);
|
||||
for (const label of labels.slice()) {
|
||||
const parent = UMBRELLA[label];
|
||||
if (parent && !labels.includes(parent)) labels.push(parent);
|
||||
}
|
||||
return labels;
|
||||
}
|
||||
|
||||
function loadKeywords(file) {
|
||||
return JSON.parse(fs.readFileSync(file, 'utf8'));
|
||||
}
|
||||
|
||||
module.exports = { componentLabels, inferComponentLabels, loadKeywords };
|
||||
@@ -0,0 +1,108 @@
|
||||
const assert = require('assert');
|
||||
const path = require('path');
|
||||
const { inferComponentLabels, loadKeywords } = require('./infer-component-labels.js');
|
||||
|
||||
const keywords = loadKeywords(path.join(__dirname, '..', 'component-keywords.json'));
|
||||
|
||||
const cases = [
|
||||
{
|
||||
number: 6210,
|
||||
title: "but(anthropic): sampling parameters returns 400 error for new model",
|
||||
body: "### Component\n\nCore / Python SDK\n\n### Description\n\n### Summary\n\nWhen using Anthropic latest models such as `claude-opus-4-7`, `claude-opus-4-8`, or `claude-sonnet-5`, Mem0 still sends sampling parameters like `temperature` / `top_p`. These models do not support those parameters, causing Anthropic API requests to fail.\n\nSee https://platform.claude.com/docs/en/about-claude/models/migration-guide\n\n### Steps to Reproduce\n\n```python\n from mem0 import Memory\n\n m = Memory.from_config({\n \"llm\": {\n \"provider\": \"anthropic\",\n \"config\": {\n \"model\": \"claude-opus-4-8\",\n \"api_key\": \"your-anthropic-api-key\"\n },\n },\n ...\n })\n```\n\n### Expected Behavior\n\nMem0 should detect Anthropic models that do not support sampling parameters and omit temperature and top_p from the request.\n\nFor models that still support sampling parameters, such as claude-opus-4-6, claude-sonnet-4-6, and claude-haiku-4-5, Mem0 should continue sending supported sampling parameters till they're deprecated.\n\n### Actual Behavior\n\nMem0 includes temperature by default for Anthropic requests. With newer Anthropic models that do not support sampling parameters, the API request fails because unsupported parameters are sent.\n\n### Environment\n\n - mem0 version: 2.0.11\n - Python/Node version: Python 3.11\n - OS: macOS\n",
|
||||
expected: ["sdk-python"],
|
||||
},
|
||||
{
|
||||
number: 5770,
|
||||
title: "feat(ts-sdk): add FastEmbed embedding provider",
|
||||
body: "## Summary\n\nThe Python SDK supports **FastEmbed** as an embedding provider, but the TypeScript OSS SDK (`mem0ai/oss`) does not. Add it to bring the TS SDK to parity.\n\n| | |\n|---|---|\n| Python reference | `mem0/embeddings/fastembed.py` |\n| Registered in (Python) | `mem0/utils/factory.py` (EmbedderFactory) |\n| Target file (TypeScript) | `mem0-ts/src/oss/src/embeddings/fastembed.ts` |\n| Suggested implementation | Use the `fastembed` npm package (ONNX local embeddings). |\n\n## Requirements\n\n- [ ] Implement `FastEmbedEmbedder` in `mem0-ts/src/oss/src/embeddings/fastembed.ts`, extending `Embedder` (`mem0-ts/src/oss/src/embeddings/base.ts`) and mirroring the Python provider's behavior (embed / embedBatch).\n- [ ] Register the `\"fastembed\"` provider in `mem0-ts/src/oss/src/utils/factory.ts` (EmbedderFactory).\n- [ ] Add config typing in `mem0-ts/src/oss/src/types/`.\n- [ ] Add a unit test under `mem0-ts/src/oss/src/tests/`.\n- [ ] Add `fastembed` to `mem0-ts/package.json` (optional/peer dependency, lazy-imported like other providers).\n- [ ] Update docs under `docs/` if this provider is user-facing.\n\n## Reference pattern\n\nMirror an existing TS provider: `embeddings/openai.ts`.\n\n## Notes\n\n`fastembed` (v2.x) is the JS port of Qdrant's FastEmbed — local/offline embeddings. Mirror the default model in `mem0/embeddings/fastembed.py`.\n\n---\n_Part of the TypeScript ↔ Python SDK provider-parity effort. One provider per issue (atomic)._\n",
|
||||
expected: ["sdk-typescript"],
|
||||
},
|
||||
{
|
||||
number: 3940,
|
||||
title: "Milvus database will return distance not similarity score",
|
||||
body: "### 🐛 Describe the bug\n\nMilvus database will return distance not similarity score\n\n## in milvus.py\n\ndef _parse_output(self, data: list):\n \"\"\"\n Parse the output data.\n\n Args:\n data (Dict): Output data.\n\n Returns:\n List[OutputData]: Parsed output data.\n \"\"\"\n memory = []\n\n for value in data:\n uid, score, metadata = (\n value.get(\"id\"),\n value.get(\"distance\"), # here\n value.get(\"entity\", {}).get(\"metadata\"),\n )\n\n memory_obj = OutputData(id=uid, score=score, payload=metadata)\n memory.append(memory_obj)\n\n return memory\n",
|
||||
expected: ["vector-store"],
|
||||
},
|
||||
{
|
||||
number: 5290,
|
||||
title: "Recall search failed: Bad Request Using OpenAI Embedding Model",
|
||||
body: "### Component\n\nOpenClaw\n\n### Description\n\n### Summary\nuse openclaw.json config:\n\n```json\n...\n\"embedder\": {\n \"provider\": \"openai\",\n \"config\": {\n \"model\": \"bge-base-zh-v1.5\",\n \"embedding_dims\": 1024,\n \"embeddingDims\": 1024,\n \"url\": \"https://xxxxxxxxx/v1\",\n \"apiKey\": \"xxxxxxxxxxxx\"\n }\n },\n\"vectorStore\": {\n \"provider\": \"qdrant\",\n \"config\": {\n \"url\": \"http://qdrant:6333\",\n \"apiKey\": \"${QDRANT_API_KEY}\",\n \"collectionName\": \"mem0\",\n \"embeddingModelDims\": 1024\n }\n }\n```\n```\n\nopenclaw log info is:\n\n```\n23:14:20 Api key is used with unsecure connection.\n23:14:21 [mem0] Recall search failed: Bad Request\n23:14:21 [plugins] openclaw-mem0: skills-mode recall (strategy=smart) injecting 0 memories (~20 tokens)\n23:14:22 [ws] ⇄ res ✓ sessions.list 256ms conn=d1eb9bc4…17da id=201b8113…c9dc\n23:14:22 [ws] ⇄ res ✓ sessions.list 264ms conn=d1eb9bc4…17da id=4939f962…2f16\n23:14:34 [ws] ⇄ res ✓ sessions.list 250ms conn=d1eb9bc4…17da id=f7ad503f…baa6\n23:15:12 [mem0] **Recall search failed: Bad Request**\n23:15:12 [plugins] openclaw-mem0: skills-mode recall (strategy=smart) injecting 0 memories (~20 tokens)\n23:15:12 [ws] ⇄ res ✓ sessions.list 288ms conn=d1eb9bc4…17da id=9e20bb86…371e\n23:15:13 [ws] ⇄ res ✓ sessions.list 268ms conn=d1eb9bc4…17da id=3b49a2ad…7ada\n23:15:20 [ws] ⇄ res ✓ sessions.list 235ms conn=d1eb9bc4…17da id=a192da30…069f\n```\n\n### Actual Behavior\n\nembedding model response ok,response message has 1024 vectors,but the vectors are submitted to vector-db:qdrant with all zero vectors,and vectors has only 256 size.\n\n```http\nPOST /collections/mem0/points/search HTTP/1.1\nhost: qdrant:6333\nconnection: keep-alive\nuser-agent: qdrant-js/1.13.0\napi-key: xxxxxxxxxxxxxxxxxxxxxxxxxxxx\nContent-Type: application/json\nAccept: application/json\naccept-language: *\nsec-fetch-mode: cors\naccept-encoding: gzip, deflate\ncontent-length: 651\n\n{\"vector\":[0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0],\"limit\":120,\"offset\":0,\"filter\":{\"must\":[{\"key\":\"user_id\",\"match\":{\"value\":\"agent\"}}]},\"with_payload\":true,\"with_vector\":false}\n\n**HTTP/1.1 400 Bad Request**\ntransfer-encoding: chunked\ncontent-type: application/json\nvary: accept-encoding, Origin, Access-Control-Request-Method, Access-Control-Request-Headers\ncontent-encoding: gzip\n\n```\n\n### Expected Behavior\n\nembedding model response ok by tcpdump, response message has 1024 vectors,and this vectors are submitted to vector-db:qdrant with the same vectors,and vectors has also 1024 size.\n\n\n### Environment\n\n- openclaw-mem0 version: 1.0.11\n- qdrant: 1.13.6\n",
|
||||
expected: ["plugin", "integrations"],
|
||||
},
|
||||
{
|
||||
number: 3696,
|
||||
title: "Cannot set expiration_date for memory in REST API server (Docker Compose)",
|
||||
body: "### 🐛 Describe the bug\n\nI'm using docker compose to deploy a REST API server. When adding memory, I'm unable to set the expiration_date. Is this feature not supported?",
|
||||
expected: ["rest-api"],
|
||||
},
|
||||
{
|
||||
number: 3444,
|
||||
title: "Fix: Openmemory run.sh non-existent vector-store route",
|
||||
body: "### 🐛 Describe the bug\n\n# Vector_store not implemented\nThere is many references to ` ${NEXT_PUBLIC_API_URL}/api/v1/config/mem0/vector_store` in lines 280, 293, 306, 319, 332, 345, 358, and 371. \n```bash\ncurl -fsS -X PUT \"${NEXT_PUBLIC_API_URL}/api/v1/config/mem0/vector_store\" # Line 280 and for each vector store\n```\nBut the api route is not implemented in `api/app/routers/config.py`.\n# Suggested solution\nI would implement `vector_store` route or remove and use `update_configuration` for all config updates. Also Create class with all config keys for vector_store",
|
||||
expected: ["openmemory"],
|
||||
},
|
||||
{
|
||||
number: 6252,
|
||||
title: "cursor: on_file_read_cursor.sh ignores auto_search / MEM0_AUTO_SEARCH",
|
||||
body: "### Component\n\nCursor / mem0-plugin\n\n### Description\n\n`on_file_read_cursor.sh` never checks `MEM0_AUTO_SEARCH`. In Claude Code, #6065/#6071 added a guard on `on_file_read.sh`, but the Cursor PreToolUse variant still always calls `file_context.py` (and thus Platform search) once `MEM0_API_KEY` is set.\n\n### Expected\n\nWhen `auto_search: false` / `MEM0_AUTO_SEARCH=false`, `on_file_read_cursor.sh` should exit 0 without searching.\n\n### Actual\n\nTimeline search still runs.\n\n### Related\n\n#6065, #6071, #6250\n",
|
||||
expected: ["plugin", "integrations"],
|
||||
},
|
||||
{
|
||||
number: 6032,
|
||||
title: "docs: fix typos and punctuation errors across docs",
|
||||
body: "### Description\n\n### Page\nMultiple pages — see list below.\n\n### What's Wrong or Missing\n1. https://docs.mem0.ai/components/llms/overview — \"a llm\" should be \"an LLM\"\n2. https://docs.mem0.ai/components/vectordbs/dbs/azure — 2 comma splices + \"setup\" used as a verb (should be \"set up\")\n3. https://docs.mem0.ai/components/embedders/models/azure_openai — \"from the Azure.\" is an incomplete sentence\n4. https://docs.mem0.ai/components/llms/models/azure_openai — same incomplete \"from the Azure\" phrasing\n5. https://docs.mem0.ai/cookbooks/companions/voice-companion-openai — \"an important information\" (uncountable noun)\n6. https://docs.mem0.ai/cookbooks/essentials/exporting-memories — comma splice\n7. https://docs.mem0.ai/cookbooks/integrations/tavily-search — \"usecase\" should be \"use case\"\n8. https://docs.mem0.ai/cookbooks/overview — broken parallelism in bullet list\n9. README.md — \"Github App\" should be \"GitHub App\"\n10. https://docs.mem0.ai/platform/overview — table cell not capitalized like other rows\n\n### Suggested Fix\nApply the corrections listed above for each page. I will submit a PR soon addressing all of the issues mentioned.",
|
||||
expected: ["documentation"],
|
||||
},
|
||||
];
|
||||
|
||||
const cliRegressionCase = {
|
||||
number: 3144,
|
||||
title: "Bug Report: Memory Score Does Not Match Expected Relevance in Local Search",
|
||||
body: "### 🐛 Describe the bug\n\n#### Description\n\nWhen using the locally deployed `mem0` server, the returned memory `score` from the `search` interface does not align with the expected semantic relevance. In particular, irrelevant or less relevant memories sometimes receive higher scores than directly related ones.\n\n#### Reproduction Steps\n\n```python\nmem0 = mem0_client(mode=\"local\")\nprint(\"Mem0 client initialized successfully.\")\n\nprint(\"Adding memories...\")\nresult = mem0.add(messages=[\n {\"role\": \"user\", \"content\": \"I like drinking coffee in the morning\"},\n {\"role\": \"user\", \"content\": \"I enjoy reading books at night\"}\n], user_id=\"alice\")\nprint(\"Memory added:\", result)\n\nprint(\"Searching memories...\")\nsearch_result = mem0.search(query=\"coffee\", user_id=\"alice\", top_k=2)\nprint(\"Search results:\", search_result)\n```\n\n#### Actual Output\n\n```json\n{\n \"results\": [\n {\n \"id\": \"5099b5be-c673-4f09-99de-a196f43b6476\",\n \"memory\": \"Likes drinking coffee in the morning\",\n \"score\": 0.5115111920687857\n },\n {\n \"id\": \"08df5c51-c52b-4c45-a5b6-b3f864ea149a\",\n \"memory\": \"Enjoys reading books at night\",\n \"score\": 0.7755568273863331\n }\n ],\n \"relations\": [\n {\"source\": \"coffee\", \"relationship\": \"consumed_in\", \"destination\": \"morning\"},\n {\"source\": \"user_id:_alice\", \"relationship\": \"likes\", \"destination\": \"coffee\"},\n {\"source\": \"user_id:_alice\", \"relationship\": \"likes_drinking\", \"destination\": \"coffee\"},\n {\"source\": \"user_id:_alice\", \"relationship\": \"in_time\", \"destination\": \"morning\"},\n {\"source\": \"user_id:_alice\", \"relationship\": \"drinks_in\", \"destination\": \"morning\"}\n ]\n}\n```\n\n#### Expected Behavior\n\nThe memory `\"Likes drinking coffee in the morning\"` should have a **higher score** than `\"Enjoys reading books at night\"` when querying for `\"coffee\"`, since it is directly semantically related.",
|
||||
};
|
||||
|
||||
let failures = 0;
|
||||
|
||||
function run(name, fn) {
|
||||
try {
|
||||
fn();
|
||||
console.log(`PASS ${name}`);
|
||||
} catch (err) {
|
||||
failures++;
|
||||
console.error(`FAIL ${name}: ${err.message}`);
|
||||
}
|
||||
}
|
||||
|
||||
for (const { number, title, body, expected } of cases) {
|
||||
const text = `${title}
|
||||
|
||||
${body}`;
|
||||
run(`#${number}`, () => {
|
||||
assert.deepStrictEqual(inferComponentLabels(text, keywords), expected);
|
||||
});
|
||||
}
|
||||
|
||||
run('#3144 cliKeywordPrefixSubstringRegression', () => {
|
||||
const text = `${cliRegressionCase.title}
|
||||
|
||||
${cliRegressionCase.body}`;
|
||||
const inferred = inferComponentLabels(text, keywords);
|
||||
assert.ok(!inferred.includes('cli'), `expected 'cli' absent (body contains 'Mem0 client', a substring of the removed 'mem0 cli' term), got ${JSON.stringify(inferred)}`);
|
||||
});
|
||||
|
||||
run('noKeywordMatchReturnsEmptyArray', () => {
|
||||
const text = 'The weather today is sunny and I went for a walk in the park with my dog.';
|
||||
assert.deepStrictEqual(inferComponentLabels(text, keywords), []);
|
||||
});
|
||||
|
||||
run('emptyStringReturnsEmptyArray', () => {
|
||||
assert.deepStrictEqual(inferComponentLabels('', keywords), []);
|
||||
});
|
||||
|
||||
if (failures > 0) {
|
||||
console.error(`
|
||||
${failures} test(s) failed.`);
|
||||
process.exit(1);
|
||||
}
|
||||
console.log(`
|
||||
All ${cases.length + 3} tests passed.`);
|
||||
@@ -12,28 +12,56 @@ jobs:
|
||||
label:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: stefanbuck/github-issue-parser@v3
|
||||
id: issue-parser
|
||||
continue-on-error: true
|
||||
with:
|
||||
template-path: .github/ISSUE_TEMPLATE/bug_report.yml
|
||||
|
||||
- uses: redhat-plumbers-in-action/advanced-issue-labeler@v3
|
||||
continue-on-error: true
|
||||
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')
|
||||
- name: Infer component from text when the form was not used
|
||||
uses: actions/github-script@v7
|
||||
with:
|
||||
template-path: .github/ISSUE_TEMPLATE/feature_request.yml
|
||||
script: |
|
||||
const {
|
||||
componentLabels,
|
||||
inferComponentLabels,
|
||||
loadKeywords,
|
||||
} = require(`${process.env.GITHUB_WORKSPACE}/.github/scripts/infer-component-labels.js`);
|
||||
|
||||
- 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
|
||||
const { data: issue } = await github.rest.issues.get({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
issue_number: context.issue.number,
|
||||
});
|
||||
|
||||
const keywords = loadKeywords(`${process.env.GITHUB_WORKSPACE}/.github/component-keywords.json`);
|
||||
const known = componentLabels(keywords);
|
||||
const existing = issue.labels.map((label) => label.name || label);
|
||||
if (existing.some((name) => known.includes(name))) {
|
||||
core.info(`Component label already present: ${existing.join(', ')}`);
|
||||
return;
|
||||
}
|
||||
|
||||
const labels = inferComponentLabels(`${issue.title}\n\n${issue.body || ''}`, keywords);
|
||||
|
||||
if (labels.length === 0) {
|
||||
core.info('No component could be inferred from the issue text');
|
||||
return;
|
||||
}
|
||||
|
||||
core.info(`Inferred: ${labels.join(', ')}`);
|
||||
await github.rest.issues.addLabels({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
issue_number: context.issue.number,
|
||||
labels,
|
||||
});
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
name: PR Labeler
|
||||
|
||||
on:
|
||||
pull_request_target:
|
||||
types: [opened, synchronize, reopened, edited]
|
||||
|
||||
concurrency:
|
||||
group: pr-labeler-${{ github.event.pull_request.number }}
|
||||
cancel-in-progress: true
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: write
|
||||
issues: read
|
||||
|
||||
jobs:
|
||||
label:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/labeler@v5
|
||||
with:
|
||||
repo-token: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Propagate labels from linked issues
|
||||
uses: actions/github-script@v7
|
||||
with:
|
||||
script: |
|
||||
const allowed = new Set([
|
||||
'sdk-python', 'sdk-typescript', 'vector-store', 'plugin',
|
||||
'rest-api', 'openmemory', 'documentation', 'ci', 'cli', 'integrations',
|
||||
]);
|
||||
const umbrella = { plugin: 'integrations' };
|
||||
const { repository } = await github.graphql(
|
||||
`query ($owner: String!, $repo: String!, $number: Int!) {
|
||||
repository(owner: $owner, name: $repo) {
|
||||
pullRequest(number: $number) {
|
||||
closingIssuesReferences(first: 20) {
|
||||
nodes { labels(first: 50) { nodes { name } } }
|
||||
}
|
||||
}
|
||||
}
|
||||
}`,
|
||||
{ owner: context.repo.owner, repo: context.repo.repo, number: context.issue.number },
|
||||
);
|
||||
|
||||
const labels = new Set();
|
||||
for (const issue of repository.pullRequest.closingIssuesReferences.nodes) {
|
||||
for (const label of issue.labels.nodes) {
|
||||
if (allowed.has(label.name)) labels.add(label.name);
|
||||
}
|
||||
}
|
||||
|
||||
for (const label of [...labels]) {
|
||||
if (umbrella[label]) labels.add(umbrella[label]);
|
||||
}
|
||||
|
||||
if (labels.size > 0) {
|
||||
await github.rest.issues.addLabels({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
issue_number: context.issue.number,
|
||||
labels: [...labels],
|
||||
});
|
||||
}
|
||||
@@ -462,6 +462,7 @@ Publishing is routed through a single entry point: **`release.yml` (Release Rout
|
||||
| Workflow | File | Purpose |
|
||||
|----------|------|---------|
|
||||
| Issue Labeler | `issue-labeler.yml` | Automatic issue labeling |
|
||||
| PR Labeler | `pr-labeler.yml` | Path-based PR labeling plus propagating labels from linked issues |
|
||||
| Stale Bot | `stale.yml` | Marks stale issues and PRs |
|
||||
| llms.txt Check | `docs-llms-txt-check.yml` | Blocks PRs touching `docs/**/*.mdx` when `docs/llms.txt` is out of sync. Fix locally with `python scripts/check-llms-txt-coverage.py --write`. |
|
||||
|
||||
|
||||
@@ -79,7 +79,7 @@ new_project = client.project.create(
|
||||
|
||||
### Update Project Settings
|
||||
|
||||
Modify project configuration including custom instructions, categories, language preferences, retrieval criteria, and memory decay:
|
||||
Modify project configuration including custom instructions, categories, language preferences, and memory decay:
|
||||
|
||||
```python
|
||||
# Update project with custom categories
|
||||
@@ -98,14 +98,6 @@ client.project.update(
|
||||
# Use the input language for memory storage and retrieval
|
||||
client.project.update(multilingual=True)
|
||||
|
||||
# Set retrieval criteria to control which memories are surfaced in search
|
||||
client.project.update(
|
||||
retrieval_criteria=[
|
||||
{"name": "relevance", "description": "How directly relevant this memory is to the current topic or user query", "weight": 3},
|
||||
{"name": "access_frequency", "description": "How often this memory has been accessed or surfaced recently", "weight": 1}
|
||||
]
|
||||
)
|
||||
|
||||
# Enable Memory Decay (boosts recently-accessed memories at search time)
|
||||
client.project.update(decay=True)
|
||||
|
||||
@@ -120,34 +112,6 @@ client.project.update(
|
||||
)
|
||||
```
|
||||
|
||||
#### Set Retrieval Criteria
|
||||
|
||||
`retrieval_criteria` is a per-project list of dictionaries (`List[Dict]`) that shapes how memories are ranked and filtered during search. Each dictionary has three fields: `name` (identifier), `description` (interpreted by the LLM to score each memory), and `weight` (relative influence on the final score). Use this to focus retrieval on intent-aligned or signal-specific memories:
|
||||
|
||||
```python
|
||||
client.project.update(
|
||||
retrieval_criteria=[
|
||||
{
|
||||
"name": "joy",
|
||||
"description": "Measure the intensity of positive emotions such as happiness, excitement, or amusement expressed in the memory. A higher score reflects greater joy.",
|
||||
"weight": 3
|
||||
},
|
||||
{
|
||||
"name": "curiosity",
|
||||
"description": "Assess the extent to which the memory reflects inquisitiveness or interest in exploring new information. A higher score reflects stronger curiosity.",
|
||||
"weight": 2
|
||||
},
|
||||
{
|
||||
"name": "access_frequency",
|
||||
"description": "How often this memory has been accessed or surfaced recently.",
|
||||
"weight": 1
|
||||
}
|
||||
]
|
||||
)
|
||||
```
|
||||
|
||||
Pass an empty list to clear all criteria and restore default retrieval behaviour.
|
||||
|
||||
#### Toggle Memory Decay
|
||||
|
||||
`decay` is a per-project boolean that turns on [Memory Decay](/platform/features/memory-decay): a search-time ranking bias that reinforces recently-accessed memories and gently dampens stale ones. The flag is `false` by default; set it via the same project-update endpoint:
|
||||
|
||||
@@ -7,6 +7,23 @@ mode: "wide"
|
||||
<Tabs>
|
||||
<Tab title="Python">
|
||||
|
||||
<Update label="2026-07-22" description="v2.0.13">
|
||||
|
||||
**Bug Fixes:**
|
||||
- **Vector Stores:** Fix `reset()` silently leaving stale vectors behind on local (on-disk) Qdrant when the old collection directory could not be removed, for example an open file handle on Windows or NFS ([#6412](https://github.com/mem0ai/mem0/pull/6412))
|
||||
- **Core:** Stop `update()` metadata from overwriting or injecting `user_id`, `agent_id`, `run_id`, or `actor_id`. These identity fields are immutable after creation, so passing them in `metadata` can no longer move a memory into a different tenant's scope ([#6278](https://github.com/mem0ai/mem0/pull/6278))
|
||||
- **Vector Stores:** Scope Pinecone `delete_col()`/`reset()` to the configured namespace instead of deleting the whole index, so resetting a namespaced Pinecone store no longer wipes out the other namespaces sharing that index ([#6287](https://github.com/mem0ai/mem0/pull/6287))
|
||||
- **Vector Stores:** Convert Baidu Mochow's raw L2 distance into a similarity score in `search()` (`1 / (1 + distance)`), so closer matches rank higher instead of lower, matching the Milvus provider and the rest of the `VectorStoreBase` contract ([#6435](https://github.com/mem0ai/mem0/pull/6435))
|
||||
- **LLMs:** Read `OPENAI_BASE_URL` (was `OPENAI_API_BASE`) in `OpenAIStructuredLLM`, matching the official OpenAI SDK's environment variable and the rest of the OpenAI-compatible providers ([#6322](https://github.com/mem0ai/mem0/pull/6322))
|
||||
|
||||
**Improvements:**
|
||||
- **LLMs:** Remove a dead, no-op `api_key` attribute check from `LLMBase.__init__` ([#6460](https://github.com/mem0ai/mem0/pull/6460))
|
||||
|
||||
**Changes:**
|
||||
- **Client:** Remove the `retrieval_criteria` parameter from `MemoryClient.update_project()`/`AsyncMemoryClient.update_project()` and `Project.update()`/`AsyncProject.update()`. It was accepted and forwarded but never affected retrieval, so removing it is not a behavior change ([#6313](https://github.com/mem0ai/mem0/pull/6313))
|
||||
|
||||
</Update>
|
||||
|
||||
<Update label="2026-07-13" description="v2.0.12">
|
||||
|
||||
**New Features:**
|
||||
@@ -1128,6 +1145,23 @@ See the [OSS v2 to v3 migration guide](https://docs.mem0.ai/migration/oss-v2-to-
|
||||
|
||||
<Tab title="TypeScript">
|
||||
|
||||
<Update label="2026-07-22" description="v3.1.1">
|
||||
|
||||
**New Features:**
|
||||
- **Embeddings:** Add an AWS Bedrock embedding provider ([#6185](https://github.com/mem0ai/mem0/pull/6185))
|
||||
|
||||
**Bug Fixes:**
|
||||
- **Packaging:** Finish the lazy-loading work started in v3.1.0. The remaining LLMs (Anthropic, Google, Groq, LangChain, Mistral, Ollama), embedders (Google, LangChain, Ollama, Vertex AI), vector stores (Azure AI Search, Azure MySQL, Baidu, LangChain, Qdrant, Redis, Supabase, Valkey, Vectorize), and the Supabase history store still imported their SDKs at module load, so importing `mem0ai/oss` required every provider package to be installed ([#6389](https://github.com/mem0ai/mem0/pull/6389))
|
||||
- **Vector Stores:** Convert Baidu Mochow's raw L2 distance into a similarity score in `search()` (`1 / (1 + distance)`), so closer matches rank higher instead of lower. A row the backend returns without a score is now left `undefined` instead of being treated as the closest match ([#6485](https://github.com/mem0ai/mem0/pull/6485))
|
||||
- **Memory (OSS):** Coerce non-string entity IDs (e.g. a numeric `user_id`) to strings instead of crashing on `.trim()` ([#6263](https://github.com/mem0ai/mem0/pull/6263))
|
||||
- **Memory (OSS):** Stop `update()` metadata from overwriting or injecting `user_id`, `agent_id`, `run_id`, or `actor_id` (in either snake_case or camelCase). These identity fields are immutable after creation, so passing them in `metadata` can no longer move a memory into a different tenant's scope ([#6343](https://github.com/mem0ai/mem0/pull/6343))
|
||||
- **Vector Stores:** Scope Pinecone `deleteCol()`/`reset()` to the configured namespace instead of deleting the whole index, so resetting a namespaced Pinecone store no longer wipes out the other namespaces sharing that index ([#6287](https://github.com/mem0ai/mem0/pull/6287))
|
||||
|
||||
**Changes:**
|
||||
- **Client:** Remove the unused `retrievalCriteria` field from `PromptUpdatePayload`. It was accepted and forwarded but never affected retrieval, so removing it is not a behavior change ([#6313](https://github.com/mem0ai/mem0/pull/6313))
|
||||
|
||||
</Update>
|
||||
|
||||
<Update label="2026-07-13" description="v3.1.0">
|
||||
|
||||
The largest provider release for the TypeScript OSS SDK so far: 17 new vector stores, 5 new LLM providers, 4 new embedders, and reranking support. Importing `mem0ai/oss` no longer pulls in any provider SDK, so you only install what you actually configure.
|
||||
@@ -1825,6 +1859,15 @@ A full-featured command-line interface for Mem0, available in both Python and No
|
||||
<Tabs>
|
||||
<Tab title="Mem0 Plugin">
|
||||
|
||||
<Update label="2026-07-14" description="mem0-plugin v0.2.13">
|
||||
|
||||
**Fixes:**
|
||||
- **Assistant messages no longer stored as your own:** The session-summary hook (fires at the end of every assistant turn) and the post-compaction hook were sending the assistant's own message to Mem0 tagged `role: "user"`. Because Mem0 extracts *facts about the user* from each message and uses `role` to decide who spoke, the assistant's first-person prose was being saved as the human's stated preferences — "I recommend we drop Redis" became `User prefers dropping Redis entirely`. Both hooks now send `role: "assistant"`, so the same session is stored as `Assistant recommended...`. Affects Claude Code, Cursor, Codex, and Antigravity, which share these hooks.
|
||||
|
||||
Existing memories written by the previous versions are not rewritten. If your memories contain preferences you never expressed, delete them — the plugin will not recreate them.
|
||||
|
||||
</Update>
|
||||
|
||||
<Update label="2026-06-30" description="mem0-plugin v0.2.12">
|
||||
|
||||
**New Features:**
|
||||
@@ -2075,6 +2118,13 @@ Initial release of the Mem0 plugin for Claude Code and Cursor, followed by Codex
|
||||
|
||||
<Tab title="OpenCode">
|
||||
|
||||
<Update label="2026-07-22" description="OpenCode plugin v0.2.2">
|
||||
|
||||
**Fixes:**
|
||||
- **Shell-profile API key recovery:** When `MEM0_API_KEY` isn't set in the process environment, the plugin now falls back to reading it from `.zshrc`, `.bashrc`, `.zprofile`, `.bash_profile`, or `.profile`, fixing startup failures on clients (e.g. Desktop) that launch without shell-exported environment variables.
|
||||
|
||||
</Update>
|
||||
|
||||
<Update label="2026-06-30" description="OpenCode plugin v0.2.1">
|
||||
|
||||
**Improvements:**
|
||||
@@ -2148,6 +2198,15 @@ Initial release of the Mem0 plugin for Claude Code and Cursor, followed by Codex
|
||||
|
||||
<Tab title="Antigravity">
|
||||
|
||||
<Update label="2026-07-14" description="Antigravity plugin v0.1.5">
|
||||
|
||||
**Fixes:**
|
||||
- **Assistant messages no longer stored as your own:** The session-summary hook (fires at the end of every assistant turn) and the post-compaction hook were sending the assistant's own message to Mem0 tagged `role: "user"`. Because Mem0 extracts *facts about the user* from each message and uses `role` to decide who spoke, the assistant's first-person prose was being saved as the human's stated preferences — "I recommend we drop Redis" became `User prefers dropping Redis entirely`. Both hooks now send `role: "assistant"`, so the same session is stored as `Assistant recommended...`.
|
||||
|
||||
Existing memories written by the previous versions are not rewritten. If your memories contain preferences you never expressed, delete them — the plugin will not recreate them.
|
||||
|
||||
</Update>
|
||||
|
||||
<Update label="2026-06-30" description="Antigravity plugin v0.1.4">
|
||||
|
||||
**New Features:**
|
||||
|
||||
@@ -3,11 +3,27 @@ title: AWS Bedrock
|
||||
description: "Configure AWS Bedrock as an embedding provider in Mem0 with IAM credentials and boto3 authentication."
|
||||
---
|
||||
|
||||
To use AWS Bedrock embedding models, you need to have the appropriate AWS credentials and permissions. The embeddings implementation relies on the `boto3` library.
|
||||
To use AWS Bedrock embedding models, you need the appropriate AWS credentials and permissions. Python uses `boto3`, and TypeScript uses `@aws-sdk/client-bedrock-runtime`.
|
||||
|
||||
Both SDKs support the Amazon Titan and Cohere embedding model families.
|
||||
|
||||
### Setup
|
||||
- Ensure you have model access from the [AWS Bedrock Console](https://us-east-1.console.aws.amazon.com/bedrock/home?region=us-east-1#/modelaccess)
|
||||
- Authenticate the boto3 client using a method described in the [AWS documentation](https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html)
|
||||
|
||||
- Model access is automatic: Bedrock enables serverless foundation models on first invocation in AWS commercial regions, and the [Model access page has been retired](https://docs.aws.amazon.com/bedrock/latest/userguide/model-access.html). Cohere models are served from AWS Marketplace, so an account's first invocation must come from a principal with the `aws-marketplace:Subscribe` permission; after that, any user in the account can invoke them. Browse the models available to you in the [Bedrock model catalog](https://console.aws.amazon.com/bedrock/).
|
||||
- Install the AWS client for your language:
|
||||
|
||||
<CodeGroup>
|
||||
```bash Python
|
||||
pip install boto3
|
||||
```
|
||||
|
||||
```bash TypeScript
|
||||
npm install @aws-sdk/client-bedrock-runtime
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
In TypeScript this package is an optional peer dependency, so it is only required when you actually use the Bedrock embedder.
|
||||
|
||||
- Set up environment variables for authentication:
|
||||
```bash
|
||||
export AWS_REGION=us-east-1
|
||||
@@ -15,6 +31,8 @@ To use AWS Bedrock embedding models, you need to have the appropriate AWS creden
|
||||
export AWS_SECRET_ACCESS_KEY=your-secret-key
|
||||
```
|
||||
|
||||
Both SDKs fall back to the standard AWS credential chain (environment variables, shared config, SSO, or an instance role) when you do not pass credentials in the config, so you rarely need to hardcode keys. See the [boto3 credentials guide](https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html) for the Python resolution order.
|
||||
|
||||
### Usage
|
||||
|
||||
<CodeGroup>
|
||||
@@ -48,8 +66,46 @@ messages = [
|
||||
]
|
||||
m.add(messages, user_id="alice")
|
||||
```
|
||||
|
||||
```typescript TypeScript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
// Credentials are read from the AWS default chain (AWS_REGION,
|
||||
// AWS_ACCESS_KEY_ID, AWS_SECRET_ACCESS_KEY, SSO, or an instance role).
|
||||
const memory = new Memory({
|
||||
embedder: {
|
||||
provider: "aws_bedrock",
|
||||
config: {
|
||||
model: "amazon.titan-embed-text-v2:0",
|
||||
awsRegion: "us-west-2",
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
const messages = [
|
||||
{ role: "user", content: "I'm planning to watch a movie tonight. Any recommendations?" },
|
||||
{ role: "assistant", content: "How about thriller movies? They can be quite engaging." },
|
||||
{ role: "user", content: "I'm not a big fan of thriller movies but I love sci-fi movies." },
|
||||
{ role: "assistant", content: "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future." },
|
||||
];
|
||||
await memory.add(messages, { userId: "alice" });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
### Choosing a model
|
||||
|
||||
| Model | Notes |
|
||||
| --- | --- |
|
||||
| `amazon.titan-embed-text-v1` | Default. Fixed 1536-dimension output. |
|
||||
| `amazon.titan-embed-text-v2:0` | Supports a configurable output size of 256, 512, or 1024. |
|
||||
| `cohere.embed-english-v3` | English text. Embeds up to 96 texts per request. |
|
||||
| `cohere.embed-multilingual-v3` | Multilingual text. Embeds up to 96 texts per request. |
|
||||
| `cohere.embed-v4:0` | Text. Embeds up to 96 texts per request. Supports a configurable output size of 256, 512, 1024, or 1536. TypeScript only. |
|
||||
|
||||
Custom output sizes are model specific. In Python, only Titan Text Embeddings V2 accepts one. In TypeScript, Titan Text Embeddings V2 and Cohere Embed v4 both do, and `embeddingDims` is ignored on Titan V1 and on Cohere v3, which have no such parameter. When you do set it, make sure your vector store dimension matches, otherwise inserts will fail.
|
||||
|
||||
Bedrock caps a Cohere embedding call at 96 texts. The TypeScript SDK splits larger batches into multiple requests for you, so a 200 text batch becomes 3 calls.
|
||||
|
||||
### Config
|
||||
|
||||
Here are the parameters available for configuring AWS Bedrock embedder:
|
||||
@@ -64,4 +120,16 @@ Here are the parameters available for configuring AWS Bedrock embedder:
|
||||
| `aws_secret_access_key` | AWS secret access key for authentication | `None` |
|
||||
| `aws_session_token` | AWS session token for temporary credentials | `None` |
|
||||
</Tab>
|
||||
<Tab title="TypeScript">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `model` | The name of the embedding model to use | `amazon.titan-embed-text-v1` |
|
||||
| `awsRegion` | AWS region for the Bedrock client. Falls back to the `AWS_REGION` environment variable | `us-west-2` |
|
||||
| `embeddingDims` | Output vector size. Titan Text Embeddings V2 (256, 512, or 1024) and Cohere Embed v4 (256, 512, 1024, or 1536) only | `undefined` |
|
||||
| `awsAccessKeyId` | AWS access key ID for authentication | `undefined` |
|
||||
| `awsSecretAccessKey` | AWS secret access key for authentication | `undefined` |
|
||||
| `awsSessionToken` | AWS session token for temporary credentials | `undefined` |
|
||||
|
||||
Omit the three credential fields to use the AWS default credential chain. If you do pass them, `awsAccessKeyId` and `awsSecretAccessKey` are both required.
|
||||
</Tab>
|
||||
</Tabs>
|
||||
|
||||
@@ -10,7 +10,7 @@ Mem0 offers support for various embedding models, allowing users to choose the o
|
||||
See the list of supported embedders below.
|
||||
|
||||
<Note>
|
||||
All embedders listed below are supported in the Python implementation. The TypeScript implementation supports: **OpenAI**, **Azure OpenAI**, **FastEmbed**, **Google AI**, **Langchain**, **LM Studio**, **Ollama**, and **Together**.
|
||||
All embedders listed below are supported in the Python implementation. The TypeScript implementation supports: **OpenAI**, **Azure OpenAI**, **AWS Bedrock**, **FastEmbed**, **Google AI**, **Hugging Face**, **Langchain**, **LM Studio**, **Ollama**, **Together**, and **Vertex AI**.
|
||||
</Note>
|
||||
|
||||
<CardGroup cols={4}>
|
||||
|
||||
@@ -83,6 +83,50 @@ await client.add(messages, {
|
||||
Expect a `status: "PENDING"` response with an `event_id`. Poll `GET /v1/event/{event_id}/` to confirm completion.
|
||||
</Info>
|
||||
|
||||
### Automatic conversation context
|
||||
|
||||
On the Platform, you only send new messages. Mem0 automatically pulls the earlier messages that share the same identifiers (`user_id`, and `run_id` if you use one) and uses them as context when extracting memories, so you never need to resend conversation history.
|
||||
|
||||
This means a follow-up turn is understood against what came before it:
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
# First interaction
|
||||
client.add(
|
||||
[{"role": "user", "content": "My dog's name is Biscuit. He's a golden retriever."}],
|
||||
user_id="alice",
|
||||
)
|
||||
|
||||
# Later — send only the new turn, no history
|
||||
client.add(
|
||||
[{"role": "user", "content": "He turned 5 today, and I'm taking him to the vet on Friday."}],
|
||||
user_id="alice",
|
||||
)
|
||||
# Stored as: "User's dog Biscuit turned 5" — "He" is resolved against the earlier turn.
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
// First interaction
|
||||
await client.add(
|
||||
[{ role: "user", content: "My dog's name is Biscuit. He's a golden retriever." }],
|
||||
{ userId: "alice" },
|
||||
);
|
||||
|
||||
// Later — send only the new turn, no history
|
||||
await client.add(
|
||||
[{ role: "user", content: "He turned 5 today, and I'm taking him to the vet on Friday." }],
|
||||
{ userId: "alice" },
|
||||
);
|
||||
// Stored as: "User's dog Biscuit turned 5" — "He" is resolved against the earlier turn.
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
Without that earlier turn, the same message can only be stored as "User's male pet turned 5", because there is nothing to resolve "He" against. Scope each conversation with a consistent `user_id` (plus `run_id` for a distinct session) and Mem0 handles the rest.
|
||||
|
||||
<Info>
|
||||
This is default behavior and needs no configuration. Earlier SDK versions gated it behind a `version="v2"` argument on `add`; that argument no longer exists and is ignored if sent.
|
||||
</Info>
|
||||
|
||||
## Add with Mem0 Open Source
|
||||
|
||||
<CodeGroup>
|
||||
|
||||
+8
-2
@@ -84,8 +84,6 @@
|
||||
"pages": [
|
||||
"platform/features/advanced-retrieval",
|
||||
"platform/advanced-memory-operations",
|
||||
"platform/features/criteria-retrieval",
|
||||
"platform/features/contextual-add",
|
||||
"platform/features/custom-instructions",
|
||||
"platform/features/memory-decay"
|
||||
]
|
||||
@@ -609,6 +607,10 @@
|
||||
]
|
||||
},
|
||||
"redirects": [
|
||||
{
|
||||
"source": "/platform/features/contextual-add",
|
||||
"destination": "/core-concepts/memory-operations/add"
|
||||
},
|
||||
{
|
||||
"source": "/changelog/openclaw",
|
||||
"destination": "/changelog/sdk"
|
||||
@@ -1228,6 +1230,10 @@
|
||||
{
|
||||
"source": "/open-source/multimodal-support",
|
||||
"destination": "/open-source/features/multimodal-support"
|
||||
},
|
||||
{
|
||||
"source": "/platform/features/criteria-retrieval",
|
||||
"destination": "/platform/features/advanced-retrieval"
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
@@ -65,7 +65,9 @@ The plugin uses the same shell scripts as Claude Code, Cursor, and Codex: hooks
|
||||
| **User prompt** | `UserPromptSubmit` | Searches relevant memories before each message |
|
||||
| **Pre-tool** | `PreToolUse` | Blocks MEMORY.md writes, enforces `user_id`/`app_id` on mem0 tools |
|
||||
| **Post-tool** | `PostToolUse` | Tracks stats, scans bash errors for related memories |
|
||||
| **Stop** | `Stop` | Stores a session summary when the session ends |
|
||||
| **Stop** | `Stop` | Stores a session summary at the end of every assistant turn (not just at session end) |
|
||||
|
||||
What you type is stored as yours. What the agent produces — session summaries and compaction summaries — is stored as the assistant's, so its suggestions never become your stated preferences.
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
|
||||
@@ -44,14 +44,14 @@ Install the full plugin including MCP server, lifecycle hooks, and SDK skill.
|
||||
|
||||
1. Add the Mem0 marketplace:
|
||||
|
||||
```
|
||||
/plugin marketplace add mem0ai/mem0
|
||||
```bash
|
||||
claude plugin marketplace add mem0ai/mem0
|
||||
```
|
||||
|
||||
2. Install the plugin:
|
||||
|
||||
```
|
||||
/plugin install mem0@mem0-plugins
|
||||
```bash
|
||||
claude plugin install mem0@mem0-plugins
|
||||
```
|
||||
|
||||
**Claude Cowork desktop app:** Open the Cowork tab, click **Customize** in the sidebar, click **Browse plugins**, and install Mem0.
|
||||
@@ -88,6 +88,15 @@ Add to your Claude Code MCP config (`.mcp.json`):
|
||||
}
|
||||
```
|
||||
|
||||
### Managing the Plugin
|
||||
|
||||
```bash
|
||||
claude plugin update mem0@mem0-plugins # update the plugin to the latest version (restart to apply)
|
||||
claude plugin marketplace update mem0-plugins # refresh the marketplace catalog
|
||||
claude plugin uninstall mem0@mem0-plugins # uninstall the plugin (keeps the marketplace)
|
||||
claude plugin marketplace remove mem0-plugins # unregister the marketplace entirely
|
||||
```
|
||||
|
||||
<Info icon="check">
|
||||
Start a new session and ask: *"List my mem0 entities"* or *"Search my memories for hello"*. If the `mem0` tools appear and respond, you're all set.
|
||||
</Info>
|
||||
@@ -143,9 +152,11 @@ When installed via the plugin marketplace, Mem0 hooks into Claude Code's lifecyc
|
||||
| **User prompt** | `UserPromptSubmit` | Searches relevant memories before each message; skips short prompts |
|
||||
| **Pre-tool (3 handlers)** | `PreToolUse` | Blocks MEMORY.md writes; enforces `user_id`/`app_id` on mem0 tool calls; scans files being read for relevant memory context |
|
||||
| **Post-tool** | `PostToolUse` | Tracks stats, scans bash errors for related memories |
|
||||
| **Stop** | `Stop` | Stores a session summary when the session ends |
|
||||
| **Stop** | `Stop` | Stores a session summary at the end of every assistant turn (not just at session end) |
|
||||
| **Pre-compact** | `PreCompact` | Stores a summary before the context is compacted |
|
||||
|
||||
What you type is stored as yours. What Claude produces — session summaries and compaction summaries — is stored as the assistant's, so its suggestions never become your stated preferences.
|
||||
|
||||
## Example Workflow
|
||||
|
||||
```text
|
||||
@@ -153,16 +164,17 @@ When installed via the plugin marketplace, Mem0 hooks into Claude Code's lifecyc
|
||||
You: Let's refactor the auth module to use JWT tokens instead of sessions.
|
||||
|
||||
# Claude searches memories, finds nothing relevant, proceeds with the work.
|
||||
# After completing the task, Mem0 stores:
|
||||
# Mem0 stores what you said as yours:
|
||||
# - Your preference: "Prefers TypeScript, uses ESLint"
|
||||
# ...and what Claude did as the assistant's, in the session summary:
|
||||
# - Decision: "Migrated auth from sessions to JWT tokens"
|
||||
# - Files modified: auth/middleware.ts, auth/token.ts
|
||||
# - User preference: "Prefers TypeScript, uses ESLint"
|
||||
|
||||
# Session 2 (days later): Related work
|
||||
You: Add refresh token rotation to the auth system.
|
||||
|
||||
# Claude searches memories, retrieves the JWT migration context.
|
||||
# Knows the file structure, decisions made, and user preferences.
|
||||
# Knows the file structure, decisions made, and your stated preferences.
|
||||
# Continues seamlessly without re-explaining the codebase.
|
||||
```
|
||||
|
||||
|
||||
@@ -125,20 +125,23 @@ When installed via the plugin marketplace, Mem0 hooks into Codex's lifecycle to
|
||||
| **User prompt** | `UserPromptSubmit` | Searches relevant memories before each message |
|
||||
| **Pre-tool (3 handlers)** | `PreToolUse` | Blocks MEMORY.md writes; enforces `user_id`/`app_id` on mem0 tool calls; scans files being read for relevant memory context |
|
||||
| **Post-tool** | `PostToolUse` | Tracks stats, scans bash errors for related memories |
|
||||
| **Stop** | `Stop` | Stores a session summary when the session ends |
|
||||
| **Stop** | `Stop` | Stores a session summary at the end of every assistant turn (not just at session end) |
|
||||
| **Pre-compact** | `PreCompact` | Stores a summary before the context is compacted |
|
||||
|
||||
What you type is stored as yours. What Codex produces — session summaries and compaction summaries — is stored as the assistant's, so its suggestions never become your stated preferences.
|
||||
|
||||
## Example Workflow
|
||||
|
||||
```text
|
||||
# Task 1: Setting up a new service
|
||||
You: Create a REST API for the notifications service using Express and TypeScript.
|
||||
|
||||
# Codex searches memories, finds user preferences from prior tasks.
|
||||
# After completing the task, Mem0 stores:
|
||||
# Codex searches memories, finds your preferences from prior tasks.
|
||||
# Mem0 stores what you said as yours:
|
||||
# - Your preference: "Prefers explicit error types over generic catch-all"
|
||||
# ...and what Codex did as the assistant's, in the session summary:
|
||||
# - Decision: "Notifications service uses Express + TypeScript + Zod validation"
|
||||
# - Convention: "All API routes follow /api/v1/{resource} pattern"
|
||||
# - Preference: "User prefers explicit error types over generic catch-all"
|
||||
|
||||
# Task 2 (days later): Extending the service
|
||||
You: Add WebSocket support for real-time notification delivery.
|
||||
|
||||
@@ -96,10 +96,11 @@ Once installed, the following tools are available in every Cursor session:
|
||||
You: The API endpoint /users is taking 3 seconds. Help me optimize it.
|
||||
|
||||
# Cursor agent searches memories, proceeds with investigation.
|
||||
# After completing the task, Mem0 stores:
|
||||
# Mem0 stores what you said as yours:
|
||||
# - Your preference: "Prefers query-level fixes over caching"
|
||||
# ...and what the agent did as the assistant's:
|
||||
# - Learning: "N+1 query in UserService.getAll(): fixed with eager loading"
|
||||
# - Decision: "Added database index on users.email column"
|
||||
# - Preference: "User prefers query-level fixes over caching"
|
||||
|
||||
# Session 2 (next week): Similar issue
|
||||
You: The /orders endpoint is also slow, same pattern as before.
|
||||
|
||||
@@ -204,9 +204,7 @@ If the user is on a pre-current major (Python < 2, TS < 3, or Platform `output_f
|
||||
|
||||
### Features - Advanced Retrieval
|
||||
- [Advanced Retrieval](https://docs.mem0.ai/platform/features/advanced-retrieval) [Platform]: Use when the user needs keyword search, reranking, or hybrid retrieval.
|
||||
- [Criteria-Based Retrieval](https://docs.mem0.ai/platform/features/criteria-retrieval) [Platform]: Use when targeting memories by custom criteria, not just semantic similarity.
|
||||
- [Temporal Reasoning](https://docs.mem0.ai/platform/features/temporal-reasoning) [Platform]: Use when time-aware searches like last week, upcoming, or right now need better result ordering.
|
||||
- [Contextual Add](https://docs.mem0.ai/platform/features/contextual-add) [Platform]: Use when `add()` should consider the surrounding conversation, not just the latest turn.
|
||||
- [Custom Instructions](https://docs.mem0.ai/platform/features/custom-instructions) [Platform]: Use when tailoring what Mem0 extracts and stores on Platform.
|
||||
- [Memory Decay](https://docs.mem0.ai/platform/features/memory-decay) [Platform]: Use when search results should boost recently-reinforced memories and dampen stale ones. Opt in per project; applies at search time and never filters candidates out.
|
||||
- [Advanced Memory Operations](https://docs.mem0.ai/platform/advanced-memory-operations) [Platform]: Use when basic CRUD is not enough - batch ops, complex filters, workflows.
|
||||
|
||||
@@ -148,7 +148,7 @@ See the full catalog in <Link href="/components/llms/overview">Components</Link>
|
||||
- Qdrant connection errors: confirm port `6333` is exposed and the API key (if set) matches.
|
||||
- Empty search results: verify the embedder model name. A mismatch causes dimension errors.
|
||||
- `Unknown reranker` (Python): upgrade the SDK with `pip install --upgrade mem0ai` to load the latest provider registry.
|
||||
- `Cannot find module` (Node): import from the OSS entry point, `import { Memory } from "mem0ai/oss"`, not `"mem0ai"`.
|
||||
- `Cannot find module` (Node): two common causes. First, import from the OSS entry point, `import { Memory } from "mem0ai/oss"`, not `"mem0ai"`. Second, provider SDKs are optional peer dependencies loaded on demand, so install the one for the provider you configured (for example `npm install @qdrant/js-client-rest` for Qdrant). Installing `mem0ai` alone only pulls in the providers used by default; you do not need SDKs for providers you never select.
|
||||
|
||||
<CardGroup cols={2}>
|
||||
<Card
|
||||
|
||||
@@ -1,251 +0,0 @@
|
||||
---
|
||||
title: Contextual Memory Creation
|
||||
description: "Add messages with automatic context management - no manual history tracking required"
|
||||
---
|
||||
|
||||
## What is Contextual Memory Creation?
|
||||
|
||||
Contextual memory creation automatically manages message history, allowing you to focus on building AI experiences without manually tracking interactions. Simply send new messages, and Mem0 handles the context automatically.
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
# Just send new messages - Mem0 handles the context
|
||||
messages = [
|
||||
{"role": "user", "content": "I love Italian food, especially pasta"},
|
||||
{"role": "assistant", "content": "Great! I'll remember your preference for Italian cuisine."}
|
||||
]
|
||||
|
||||
client.add(messages, user_id="user123")
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
// Just send new messages - Mem0 handles the context
|
||||
const messages = [
|
||||
{"role": "user", "content": "I love Italian food, especially pasta"},
|
||||
{"role": "assistant", "content": "Great! I'll remember your preference for Italian cuisine."}
|
||||
];
|
||||
|
||||
await client.add(messages, { userId: "user123" });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
## Why Use Contextual Memory Creation?
|
||||
|
||||
- **Simple**: Send only new messages, no manual history tracking
|
||||
- **Efficient**: Smaller payloads and faster processing
|
||||
- **Automatic**: Context management handled by Mem0
|
||||
- **Reliable**: No risk of missing interaction history
|
||||
- **Scalable**: Works seamlessly as your application grows
|
||||
|
||||
## How It Works
|
||||
|
||||
### Basic Usage
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
# First interaction
|
||||
messages1 = [
|
||||
{"role": "user", "content": "Hi, I'm Sarah from New York"},
|
||||
{"role": "assistant", "content": "Hello Sarah! Nice to meet you."}
|
||||
]
|
||||
client.add(messages1, user_id="sarah")
|
||||
|
||||
# Later interaction - just send new messages
|
||||
messages2 = [
|
||||
{"role": "user", "content": "I'm planning a trip to Italy next month"},
|
||||
{"role": "assistant", "content": "How exciting! Italy is beautiful this time of year."}
|
||||
]
|
||||
client.add(messages2, user_id="sarah")
|
||||
# Mem0 automatically knows Sarah is from New York and can use this context
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
// First interaction
|
||||
const messages1 = [
|
||||
{"role": "user", "content": "Hi, I'm Sarah from New York"},
|
||||
{"role": "assistant", "content": "Hello Sarah! Nice to meet you."}
|
||||
];
|
||||
await client.add(messages1, { userId: "sarah" });
|
||||
|
||||
// Later interaction - just send new messages
|
||||
const messages2 = [
|
||||
{"role": "user", "content": "I'm planning a trip to Italy next month"},
|
||||
{"role": "assistant", "content": "How exciting! Italy is beautiful this time of year."}
|
||||
];
|
||||
await client.add(messages2, { userId: "sarah" });
|
||||
// Mem0 automatically knows Sarah is from New York and can use this context
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
## Organization Strategies
|
||||
|
||||
Choose the right approach based on your application's needs:
|
||||
|
||||
### User-Level Memories (`user_id` only)
|
||||
|
||||
**Best for:** Personal preferences, profile information, long-term user data
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
# Persistent user memories across all interactions
|
||||
messages = [
|
||||
{"role": "user", "content": "I'm allergic to nuts and dairy"},
|
||||
{"role": "assistant", "content": "I've noted your allergies for future reference."}
|
||||
]
|
||||
|
||||
client.add(messages, user_id="user123")
|
||||
# This allergy info will be available in ALL future interactions
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
// Persistent user memories across all interactions
|
||||
const messages = [
|
||||
{"role": "user", "content": "I'm allergic to nuts and dairy"},
|
||||
{"role": "assistant", "content": "I've noted your allergies for future reference."}
|
||||
];
|
||||
|
||||
await client.add(messages, { userId: "user123" });
|
||||
// This allergy info will be available in ALL future interactions
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
### Session-Specific Memories (`user_id` + `run_id`)
|
||||
|
||||
**Best for:** Task-specific context, separate interaction threads, project-based sessions
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
# Trip planning session
|
||||
messages1 = [
|
||||
{"role": "user", "content": "I want to plan a 5-day trip to Tokyo"},
|
||||
{"role": "assistant", "content": "Perfect! Let's plan your Tokyo adventure."}
|
||||
]
|
||||
client.add(messages1, user_id="user123", run_id="tokyo-trip-2024")
|
||||
|
||||
# Later in the same trip planning session
|
||||
messages2 = [
|
||||
{"role": "user", "content": "I prefer staying near Shibuya"},
|
||||
{"role": "assistant", "content": "Great choice! Shibuya is very convenient."}
|
||||
]
|
||||
client.add(messages2, user_id="user123", run_id="tokyo-trip-2024")
|
||||
|
||||
# Different session for work project (separate context)
|
||||
work_messages = [
|
||||
{"role": "user", "content": "Let's discuss the Q4 marketing strategy"},
|
||||
{"role": "assistant", "content": "Sure! What are your main goals for Q4?"}
|
||||
]
|
||||
client.add(work_messages, user_id="user123", run_id="q4-marketing")
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
// Trip planning session
|
||||
const messages1 = [
|
||||
{"role": "user", "content": "I want to plan a 5-day trip to Tokyo"},
|
||||
{"role": "assistant", "content": "Perfect! Let's plan your Tokyo adventure."}
|
||||
];
|
||||
await client.add(messages1, { userId: "user123", runId: "tokyo-trip-2024" });
|
||||
|
||||
// Later in the same trip planning session
|
||||
const messages2 = [
|
||||
{"role": "user", "content": "I prefer staying near Shibuya"},
|
||||
{"role": "assistant", "content": "Great choice! Shibuya is very convenient."}
|
||||
];
|
||||
await client.add(messages2, { userId: "user123", runId: "tokyo-trip-2024" });
|
||||
|
||||
// Different session for work project (separate context)
|
||||
const workMessages = [
|
||||
{"role": "user", "content": "Let's discuss the Q4 marketing strategy"},
|
||||
{"role": "assistant", "content": "Sure! What are your main goals for Q4?"}
|
||||
];
|
||||
await client.add(workMessages, { userId: "user123", runId: "q4-marketing" });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
## Real-World Use Cases
|
||||
|
||||
<Tabs>
|
||||
<Tab title="Customer Support">
|
||||
```python Python
|
||||
# Support ticket context - keeps interaction focused
|
||||
messages = [
|
||||
{"role": "user", "content": "My subscription isn't working"},
|
||||
{"role": "assistant", "content": "I can help with that. What specific issue are you experiencing?"},
|
||||
{"role": "user", "content": "I can't access premium features even though I paid"}
|
||||
]
|
||||
|
||||
# Each support ticket gets its own run_id
|
||||
client.add(messages,
|
||||
user_id="customer123",
|
||||
run_id="ticket-2024-001"
|
||||
)
|
||||
```
|
||||
</Tab>
|
||||
<Tab title="Personal AI Assistant">
|
||||
```python Python
|
||||
# Personal preferences (persistent across all interactions)
|
||||
preference_messages = [
|
||||
{"role": "user", "content": "I prefer morning workouts and vegetarian meals"},
|
||||
{"role": "assistant", "content": "Got it! I'll keep your fitness and dietary preferences in mind."}
|
||||
]
|
||||
|
||||
client.add(preference_messages, user_id="user456")
|
||||
|
||||
# Daily planning session (session-specific)
|
||||
planning_messages = [
|
||||
{"role": "user", "content": "Help me plan tomorrow's schedule"},
|
||||
{"role": "assistant", "content": "Of course! I'll consider your morning workout preference."}
|
||||
]
|
||||
|
||||
client.add(planning_messages,
|
||||
user_id="user456",
|
||||
run_id="daily-plan-2024-01-15"
|
||||
)
|
||||
```
|
||||
</Tab>
|
||||
<Tab title="Educational Platform">
|
||||
```python Python
|
||||
# Student profile (persistent)
|
||||
profile_messages = [
|
||||
{"role": "user", "content": "I'm studying computer science and struggle with math"},
|
||||
{"role": "assistant", "content": "I'll tailor explanations to help with math concepts."}
|
||||
]
|
||||
|
||||
client.add(profile_messages, user_id="student789")
|
||||
|
||||
# Specific lesson session
|
||||
lesson_messages = [
|
||||
{"role": "user", "content": "Can you explain algorithms?"},
|
||||
{"role": "assistant", "content": "Sure! I'll explain algorithms with math-friendly examples."}
|
||||
]
|
||||
|
||||
client.add(lesson_messages,
|
||||
user_id="student789",
|
||||
run_id="algorithms-lesson-1"
|
||||
)
|
||||
```
|
||||
</Tab>
|
||||
</Tabs>
|
||||
|
||||
## Best Practices
|
||||
|
||||
### ✅ Do
|
||||
- **Organize by context scope**: Use `user_id` only for persistent data, add `run_id` for session-specific context
|
||||
- **Keep messages focused** on the current interaction
|
||||
- **Test with real interaction flows** to ensure context works as expected
|
||||
|
||||
### ❌ Don't
|
||||
- Send duplicate messages or interaction history
|
||||
- Skip identifiers like `user_id` or `run_id` that scope the memory
|
||||
- Mix contextual and non-contextual approaches in the same application
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
| Issue | Solution |
|
||||
|-------|----------|
|
||||
| **Context not working** | Ensure each call uses the same `user_id` / `run_id` combo; version is automatic |
|
||||
| **Wrong context retrieved** | Check if you need separate `run_id` values for different interaction topics |
|
||||
| **Missing interaction history** | Verify all messages in the interaction thread use the same `user_id` and `run_id` |
|
||||
| **Too much irrelevant context** | Use more specific `run_id` values to separate different interaction types |
|
||||
|
||||
|
||||
<Snippet file="get-help.mdx" />
|
||||
@@ -1,201 +0,0 @@
|
||||
---
|
||||
title: Criteria Retrieval
|
||||
description: "Rank and retrieve memories based on custom-defined criteria like emotional tone, intent, and behavioral signals."
|
||||
---
|
||||
|
||||
Mem0's Criteria Retrieval feature allows you to retrieve memories based on your defined criteria. It goes beyond generic semantic relevance and ranks memories based on what matters to your application: emotional tone, intent, behavioral signals, or other custom traits.
|
||||
|
||||
Instead of just searching for "how similar a memory is to this query," you can define what relevance truly means for your project. For example:
|
||||
|
||||
- Prioritize joyful memories when building a wellness assistant
|
||||
- Downrank negative memories in a productivity-focused agent
|
||||
- Highlight curiosity in a tutoring agent
|
||||
|
||||
You define criteria: custom attributes like "joy", "negativity", "confidence", or "urgency", and assign weights to control how they influence scoring. When you search, Mem0 uses these to re-rank semantically relevant memories, favoring those that better match your intent.
|
||||
|
||||
This gives you nuanced, intent-aware memory search that adapts to your use case.
|
||||
|
||||
|
||||
|
||||
## When to Use Criteria Retrieval
|
||||
|
||||
Use Criteria Retrieval if:
|
||||
|
||||
- You’re building an agent that should react to **emotions** or **behavioral signals**
|
||||
- You want to guide memory selection based on **context**, not just content
|
||||
- You have domain-specific signals like "risk", "positivity", "confidence", etc. that shape recall
|
||||
|
||||
|
||||
|
||||
## Setting Up Criteria Retrieval
|
||||
|
||||
Let’s walk through how to configure and use Criteria Retrieval step by step.
|
||||
|
||||
### Initialize the Client
|
||||
|
||||
Before defining any criteria, make sure to initialize the `MemoryClient` with your credentials and project ID:
|
||||
|
||||
```python
|
||||
from mem0 import MemoryClient
|
||||
|
||||
client = MemoryClient(api_key="your_mem0_api_key")
|
||||
```
|
||||
|
||||
### Define Your Criteria
|
||||
|
||||
Each criterion includes:
|
||||
- A `name` (used in scoring)
|
||||
- A `description` (interpreted by the LLM)
|
||||
- A `weight` (how much it influences the final score)
|
||||
|
||||
```python
|
||||
retrieval_criteria = [
|
||||
{
|
||||
"name": "joy",
|
||||
"description": "Measure the intensity of positive emotions such as happiness, excitement, or amusement expressed in the sentence. A higher score reflects greater joy.",
|
||||
"weight": 3
|
||||
},
|
||||
{
|
||||
"name": "curiosity",
|
||||
"description": "Assess the extent to which the sentence reflects inquisitiveness, interest in exploring new information, or asking questions. A higher score reflects stronger curiosity.",
|
||||
"weight": 2
|
||||
},
|
||||
{
|
||||
"name": "emotion",
|
||||
"description": "Evaluate the presence and depth of sadness or negative emotional tone, including expressions of disappointment, frustration, or sorrow. A higher score reflects greater sadness.",
|
||||
"weight": 1
|
||||
}
|
||||
]
|
||||
```
|
||||
|
||||
### Apply Criteria to Your Project
|
||||
|
||||
Once defined, register the criteria to your project:
|
||||
|
||||
```python
|
||||
client.project.update(retrieval_criteria=retrieval_criteria)
|
||||
```
|
||||
|
||||
Criteria apply project-wide. Once set, they affect all searches automatically.
|
||||
|
||||
|
||||
## Example Walkthrough
|
||||
|
||||
After setting up your criteria, you can use them to filter and retrieve memories. Here's an example:
|
||||
|
||||
### Add Memories
|
||||
|
||||
```python
|
||||
messages = [
|
||||
{"role": "user", "content": "What a beautiful sunny day! I feel so refreshed and ready to take on anything!"},
|
||||
{"role": "user", "content": "I've always wondered how storms form, what triggers them in the atmosphere?"},
|
||||
{"role": "user", "content": "It's been raining for days, and it just makes everything feel heavier."},
|
||||
{"role": "user", "content": "Finally I get time to draw something today, after a long time!! I am super happy today."}
|
||||
]
|
||||
|
||||
client.add(messages, user_id="alice")
|
||||
```
|
||||
|
||||
### Run Standard vs. Criteria-Based Search
|
||||
|
||||
```python
|
||||
# Search with criteria enabled
|
||||
filters = {"user_id": "alice"}
|
||||
results_with_criteria = client.search(
|
||||
query="Why I am feeling happy today?",
|
||||
filters=filters
|
||||
)
|
||||
|
||||
# To disable criteria for a specific search
|
||||
results_without_criteria = client.search(
|
||||
query="Why I am feeling happy today?",
|
||||
filters=filters,
|
||||
use_criteria=False # Disable criteria-based scoring
|
||||
)
|
||||
```
|
||||
|
||||
### Compare Results
|
||||
|
||||
### Search Results (with Criteria)
|
||||
```text
|
||||
[
|
||||
{"memory": "User feels refreshed and ready to take on anything on a beautiful sunny day", "score": 0.666, ...},
|
||||
{"memory": "User finally has time to draw something after a long time", "score": 0.616, ...},
|
||||
{"memory": "User is happy today", "score": 0.500, ...},
|
||||
{"memory": "User is curious about how storms form and what triggers them in the atmosphere.", "score": 0.400, ...},
|
||||
{"memory": "It has been raining for days, making everything feel heavier.", "score": 0.116, ...}
|
||||
]
|
||||
```
|
||||
|
||||
### Search Results (without Criteria)
|
||||
```text
|
||||
[
|
||||
{"memory": "User is happy today", "score": 0.607, ...},
|
||||
{"memory": "User feels refreshed and ready to take on anything on a beautiful sunny day", "score": 0.512, ...},
|
||||
{"memory": "It has been raining for days, making everything feel heavier.", "score": 0.4617, ...},
|
||||
{"memory": "User is curious about how storms form and what triggers them in the atmosphere.", "score": 0.340, ...},
|
||||
{"memory": "User finally has time to draw something after a long time", "score": 0.336, ...},
|
||||
]
|
||||
```
|
||||
|
||||
## Search Results Comparison
|
||||
|
||||
1. **Memory Ordering**: With criteria, memories with high joy scores (like feeling refreshed and drawing) are ranked higher. Without criteria, the most relevant memory ("User is happy today") comes first.
|
||||
2. **Score Distribution**: With criteria, scores are more spread out (0.116 to 0.666) and reflect the criteria weights. Without criteria, scores are more clustered (0.336 to 0.607) and based purely on relevance.
|
||||
3. **Trait Sensitivity**: "Rainy day" content is penalized due to negative tone, while "Storm curiosity" is recognized and scored accordingly.
|
||||
|
||||
|
||||
|
||||
## Key Differences vs. Standard Search
|
||||
|
||||
| Aspect | Standard Search | Criteria Retrieval |
|
||||
|-------------------------|--------------------------------------|-------------------------------------------------|
|
||||
| Ranking Logic | Semantic similarity only | Semantic + LLM-based criteria scoring |
|
||||
| Control Over Relevance | None | Fully customizable with weighted criteria |
|
||||
| Memory Reordering | Static based on similarity | Dynamically re-ranked by intent alignment |
|
||||
| Emotional Sensitivity | No tone or trait awareness | Incorporates emotion, tone, or custom behaviors |
|
||||
| Activation | Default (no criteria defined) | Enabled when criteria are defined in project |
|
||||
|
||||
<Note>
|
||||
If no criteria are defined for a project, search behaves normally based on semantic similarity only.
|
||||
</Note>
|
||||
|
||||
|
||||
|
||||
## Best Practices
|
||||
|
||||
- Choose 3-5 criteria that reflect your application's intent
|
||||
- Make descriptions clear and distinct; these are interpreted by an LLM
|
||||
- Use stronger weights to amplify the impact of important traits
|
||||
- Avoid redundant or ambiguous criteria (e.g., "positivity" and "joy")
|
||||
- Always handle empty result sets in your application logic
|
||||
|
||||
|
||||
|
||||
## How It Works
|
||||
|
||||
1. **Criteria Definition**: Define custom criteria with a name, description, and weight. These describe what matters in a memory (e.g., joy, urgency, empathy).
|
||||
2. **Project Configuration**: Register these criteria using `project.update()`. They apply at the project level and automatically influence all searches.
|
||||
3. **Memory Retrieval**: When you perform a search, Mem0 first retrieves relevant memories based on the query.
|
||||
4. **Weighted Scoring**: Each retrieved memory is evaluated and scored against your defined criteria and weights.
|
||||
|
||||
This lets you prioritize memories that align with your agent's goals and not just those that look similar to the query.
|
||||
|
||||
<Note>
|
||||
Criteria retrieval is automatically enabled when criteria are defined in your project. Use `use_criteria=False` in search to temporarily disable it for a specific query. `use_criteria` is a server-side parameter passed through to the Platform API: it is not a typed option in the SDK's `SearchMemoryOptions` interface, but the server accepts and processes it when included in the request body.
|
||||
</Note>
|
||||
|
||||
|
||||
|
||||
## Summary
|
||||
|
||||
- Define what "relevant" means using criteria
|
||||
- Apply them per project via `project.update()`
|
||||
- Criteria-aware search activates automatically when criteria are configured
|
||||
- Build agents that reason not just with relevance, but **contextual importance**
|
||||
|
||||
---
|
||||
|
||||
Need help designing or tuning your criteria?
|
||||
|
||||
<Snippet file="get-help.mdx" />
|
||||
@@ -60,7 +60,6 @@ Mem0 offers two powerful ways to add memory to your AI applications. Choose base
|
||||
| **Multimodal support** | ✅ | ✅ |
|
||||
| **Custom categories** | ✅ | Limited |
|
||||
| **Advanced retrieval** | ✅ | ✅ |
|
||||
| **Criteria retrieval** | ✅ | ❌ |
|
||||
| **Temporal reasoning** | ✅ (v3) | ❌ |
|
||||
| **Memory decay** | ✅ (v3) | ❌ |
|
||||
| **Graph memory** | ✅ Built-in | ✅ External graph store |
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "mem0",
|
||||
"version": "0.2.12",
|
||||
"version": "0.2.13",
|
||||
"description": "Persistent memory for Claude Code. Remembers decisions, patterns, and preferences across sessions.",
|
||||
"author": {
|
||||
"name": "Mem0",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "mem0",
|
||||
"version": "0.2.12",
|
||||
"version": "0.2.13",
|
||||
"description": "Persistent memory for Codex. Remembers decisions, patterns, and preferences across sessions.",
|
||||
"author": {
|
||||
"name": "Mem0",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "mem0",
|
||||
"version": "0.2.12",
|
||||
"version": "0.2.13",
|
||||
"description": "Mem0 memory layer for AI applications. Add persistent memory, personalization, and semantic search using the Mem0 Platform MCP server.",
|
||||
"author": {
|
||||
"name": "Mem0",
|
||||
|
||||
@@ -0,0 +1,153 @@
|
||||
import {afterEach, describe, expect, test} from "bun:test";
|
||||
import {mkdtempSync, mkdirSync, writeFileSync} from "fs";
|
||||
import {tmpdir} from "os";
|
||||
import {join} from "path";
|
||||
import {parseApiKeyLine, resolveApiKey} from "./api-key";
|
||||
import Mem0Plugin from "./opencode-mem0";
|
||||
|
||||
const originalKey = process.env.MEM0_API_KEY;
|
||||
const originalHome = process.env.HOME;
|
||||
const originalUserProfile = process.env.USERPROFILE;
|
||||
const originalTelemetry = process.env.MEM0_TELEMETRY;
|
||||
const originalFetch = globalThis.fetch;
|
||||
const testCleanups: Array<() => Promise<void> | void> = [];
|
||||
|
||||
afterEach(async () => {
|
||||
while (testCleanups.length > 0) {
|
||||
await testCleanups.pop()?.();
|
||||
}
|
||||
if (originalKey === undefined) delete process.env.MEM0_API_KEY;
|
||||
else process.env.MEM0_API_KEY = originalKey;
|
||||
if (originalHome === undefined) delete process.env.HOME;
|
||||
else process.env.HOME = originalHome;
|
||||
if (originalUserProfile === undefined) delete process.env.USERPROFILE;
|
||||
else process.env.USERPROFILE = originalUserProfile;
|
||||
if (originalTelemetry === undefined) delete process.env.MEM0_TELEMETRY;
|
||||
else process.env.MEM0_TELEMETRY = originalTelemetry;
|
||||
globalThis.fetch = originalFetch;
|
||||
});
|
||||
|
||||
function home(): string {
|
||||
return mkdtempSync(join(tmpdir(), "mem0-api-key-"));
|
||||
}
|
||||
|
||||
function pluginContext(logs: unknown[]) {
|
||||
return {
|
||||
client: {app: {log: async (entry: unknown) => logs.push(entry)}},
|
||||
$: () => ({quiet: async () => ({stdout: ""})}),
|
||||
} as any;
|
||||
}
|
||||
|
||||
function stubFetch(): void {
|
||||
globalThis.fetch = (async (input) => {
|
||||
const url = typeof input === "string" ? input : input instanceof URL ? input.toString() : input.url;
|
||||
const body =
|
||||
url.includes("/v1/ping/")
|
||||
? {status: "ok", userEmail: "plugin-test@mem0.dev"}
|
||||
: url.includes("/v1/projects/")
|
||||
? {customCategories: []}
|
||||
: {};
|
||||
return new Response(JSON.stringify(body), {
|
||||
status: 200,
|
||||
headers: {"content-type": "application/json"},
|
||||
});
|
||||
}) as typeof fetch;
|
||||
}
|
||||
|
||||
function captureDeferredPluginCleanup(): () => Promise<void> {
|
||||
const baseline = new Set(process.listeners("beforeExit"));
|
||||
return async () => {
|
||||
await Promise.resolve();
|
||||
for (const listener of process.listeners("beforeExit")) {
|
||||
if (!baseline.has(listener)) {
|
||||
process.off("beforeExit", listener as (...args: any[]) => void);
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
describe("parseApiKeyLine", () => {
|
||||
test("accepts literal export and assignment forms", () => {
|
||||
expect(parseApiKeyLine('export MEM0_API_KEY="m0-quoted" # comment')).toBe("m0-quoted");
|
||||
expect(parseApiKeyLine("MEM0_API_KEY=m0-literal")).toBe("m0-literal");
|
||||
});
|
||||
|
||||
test("rejects unrelated, empty, and executable-looking values", () => {
|
||||
expect(parseApiKeyLine("OTHER_KEY=value")).toBeUndefined();
|
||||
expect(parseApiKeyLine("MEM0_API_KEY=")).toBeUndefined();
|
||||
expect(parseApiKeyLine("MEM0_API_KEY=$MEM0_API_KEY")).toBeUndefined();
|
||||
expect(parseApiKeyLine("MEM0_API_KEY=$(cat secret)")).toBeUndefined();
|
||||
});
|
||||
});
|
||||
|
||||
describe("resolveApiKey", () => {
|
||||
test("explicit environment value wins over profiles", () => {
|
||||
const dir = home();
|
||||
writeFileSync(join(dir, ".zshrc"), "MEM0_API_KEY=profile\n");
|
||||
expect(resolveApiKey({MEM0_API_KEY: " explicit "}, dir)).toBe("explicit");
|
||||
});
|
||||
|
||||
test("uses the first valid allowlisted profile", () => {
|
||||
const dir = home();
|
||||
writeFileSync(join(dir, ".zshrc"), "MEM0_API_KEY=$UNSET\n");
|
||||
writeFileSync(join(dir, ".bashrc"), "export MEM0_API_KEY='m0-from-bashrc'\n");
|
||||
writeFileSync(join(dir, ".profile"), "MEM0_API_KEY=late\n");
|
||||
expect(resolveApiKey({}, dir)).toBe("m0-from-bashrc");
|
||||
});
|
||||
|
||||
test("continues after an unreadable or absent earlier profile", () => {
|
||||
const dir = home();
|
||||
mkdirSync(join(dir, ".zshrc"));
|
||||
writeFileSync(join(dir, ".profile"), "MEM0_API_KEY=m0-later\n");
|
||||
expect(resolveApiKey({}, dir)).toBe("m0-later");
|
||||
});
|
||||
|
||||
test("ignores unsupported files and invalid assignments", () => {
|
||||
const dir = home();
|
||||
writeFileSync(join(dir, ".env"), "MEM0_API_KEY=unsupported\n");
|
||||
writeFileSync(join(dir, ".zshrc"), "MEM0_API_KEY= # empty\nMEM0_API_KEY=$(unsafe)\n");
|
||||
expect(resolveApiKey({}, dir)).toBe("");
|
||||
});
|
||||
|
||||
test("keeps the missing-key guard when no source yields a key", async () => {
|
||||
const dir = home();
|
||||
writeFileSync(join(dir, ".zshrc"), "MEM0_API_KEY= # empty\nMEM0_API_KEY=$UNSET\n");
|
||||
delete process.env.MEM0_API_KEY;
|
||||
process.env.HOME = dir;
|
||||
process.env.USERPROFILE = dir;
|
||||
process.env.MEM0_TELEMETRY = "false";
|
||||
stubFetch();
|
||||
testCleanups.push(captureDeferredPluginCleanup());
|
||||
|
||||
const logs: unknown[] = [];
|
||||
const plugin = await Mem0Plugin(pluginContext(logs));
|
||||
|
||||
expect(logs).toHaveLength(1);
|
||||
expect(logs[0]).toMatchObject({
|
||||
body: {
|
||||
level: "error",
|
||||
message: "MEM0_API_KEY environment variable not set. Get one at https://app.mem0.ai/dashboard/api-keys",
|
||||
},
|
||||
});
|
||||
expect(plugin).toEqual({});
|
||||
});
|
||||
|
||||
test("recovers the issue's shell-profile startup path", async () => {
|
||||
const dir = home();
|
||||
// Problem 2 in issue #6003: Desktop has no process key, but .zshrc does.
|
||||
writeFileSync(join(dir, ".zshrc"), 'export MEM0_API_KEY="m0-from-profile"\n');
|
||||
delete process.env.MEM0_API_KEY;
|
||||
process.env.HOME = dir;
|
||||
process.env.USERPROFILE = dir;
|
||||
process.env.MEM0_TELEMETRY = "false";
|
||||
stubFetch();
|
||||
testCleanups.push(captureDeferredPluginCleanup());
|
||||
|
||||
const logs: unknown[] = [];
|
||||
const plugin = await Mem0Plugin(pluginContext(logs));
|
||||
|
||||
expect(logs).toHaveLength(0);
|
||||
expect(Object.keys(plugin)).toContain("chat.message");
|
||||
expect(plugin).toHaveProperty("tool");
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,30 @@
|
||||
import {readFileSync} from "fs";
|
||||
import {homedir} from "os";
|
||||
import {join} from "path";
|
||||
|
||||
const PROFILE_FILES = [".zshrc", ".bashrc", ".zprofile", ".bash_profile", ".profile"];
|
||||
|
||||
export function parseApiKeyLine(line: string): string | undefined {
|
||||
const match = line.match(/^\s*(?:export\s+)?MEM0_API_KEY=(.*)$/);
|
||||
if (!match) return undefined;
|
||||
|
||||
const value = match[1].replace(/#.*$/, "").trim().replace(/^("|')(.*)\1$/, "$2").trim();
|
||||
return value && !value.startsWith("$") ? value : undefined;
|
||||
}
|
||||
|
||||
export function resolveApiKey(env: NodeJS.ProcessEnv = process.env, homeDir = homedir()): string {
|
||||
const explicit = env.MEM0_API_KEY?.trim();
|
||||
if (explicit) return explicit;
|
||||
|
||||
for (const profile of PROFILE_FILES) {
|
||||
try {
|
||||
for (const line of readFileSync(join(homeDir, profile), "utf8").split(/\r?\n/)) {
|
||||
const key = parseApiKeyLine(line);
|
||||
if (key) return key;
|
||||
}
|
||||
} catch {
|
||||
}
|
||||
}
|
||||
|
||||
return "";
|
||||
}
|
||||
@@ -25,6 +25,7 @@ import {
|
||||
} from "./dream";
|
||||
import {asScope, scopeSearchFilters, scopeWriteParams, resolveDefaultScope, SCOPE_GUIDANCE, type Scope} from "./scope";
|
||||
import {parseProjectFromRemote} from "./project";
|
||||
import {resolveApiKey} from "./api-key";
|
||||
|
||||
async function getUserId(): Promise<string> {
|
||||
if (process.env.MEM0_USER_ID) return process.env.MEM0_USER_ID;
|
||||
@@ -258,7 +259,7 @@ function extractUserText(input: any, output: any): string {
|
||||
const Mem0Plugin: Plugin = async (ctx) => {
|
||||
const {$, client} = ctx;
|
||||
|
||||
const apiKey = process.env.MEM0_API_KEY;
|
||||
const apiKey = resolveApiKey();
|
||||
|
||||
if (!apiKey) {
|
||||
try {
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@mem0/opencode-plugin",
|
||||
"version": "0.2.1",
|
||||
"version": "0.2.2",
|
||||
"type": "module",
|
||||
"description": "Mem0 persistent memory plugin for OpenCode — add, search, and manage memories across sessions",
|
||||
"main": "dist/index.js",
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"id": "mem0",
|
||||
"name": "mem0",
|
||||
"version": "0.1.4",
|
||||
"version": "0.1.5",
|
||||
"description": "Persistent semantic memory for Antigravity agents. Cross-session, user-level recall via the Mem0 Platform MCP server. 16 slash commands, lifecycle hooks for auto-capture and metadata enforcement.",
|
||||
"author": { "name": "Mem0", "email": "support@mem0.ai" },
|
||||
"publisher": "mem0ai",
|
||||
|
||||
@@ -104,8 +104,11 @@ def store_summary(api_key: str, summary: str, user_id: str, session_id: str, pro
|
||||
}
|
||||
if branch:
|
||||
metadata["branch"] = branch
|
||||
# The compact summary is model-authored prose, in the first person and with no
|
||||
# framing to mark it as such. Under role="user" mem0 reads "I recommend X" as
|
||||
# the human saying it and stores "User recommends X".
|
||||
body = {
|
||||
"messages": [{"role": "user", "content": summary}],
|
||||
"messages": [{"role": "assistant", "content": summary}],
|
||||
"user_id": user_id,
|
||||
"app_id": project_id,
|
||||
"metadata": metadata,
|
||||
|
||||
@@ -175,8 +175,12 @@ def store_summary(
|
||||
if files:
|
||||
metadata["files_touched"] = files[:20]
|
||||
|
||||
# summary_prompt wraps the assistant's own last message. Mem0 extracts "facts
|
||||
# about the user" from each message and role is the only signal telling it who
|
||||
# spoke, so role="user" here turns Claude's opinions into the human's stated
|
||||
# preferences ("User prefers dropping Redis...").
|
||||
body = {
|
||||
"messages": [{"role": "user", "content": summary_prompt}],
|
||||
"messages": [{"role": "assistant", "content": summary_prompt}],
|
||||
"user_id": user_id,
|
||||
"app_id": project_id,
|
||||
"run_id": session_id,
|
||||
|
||||
@@ -8,7 +8,6 @@ Additional platform capabilities beyond core CRUD operations.
|
||||
- [Entity Linking](#entity-linking)
|
||||
- [Custom Categories](#custom-categories)
|
||||
- [Custom Instructions](#custom-instructions)
|
||||
- [Criteria Retrieval](#criteria-retrieval)
|
||||
- [Feedback Mechanism](#feedback-mechanism)
|
||||
- [Memory Export](#memory-export)
|
||||
- [Group Chat](#group-chat)
|
||||
@@ -165,44 +164,6 @@ await client.updateProject({ customInstructions: "Your guidelines here..." });
|
||||
|
||||
---
|
||||
|
||||
## Criteria Retrieval
|
||||
|
||||
Custom attribute-based memory ranking using LLM-evaluated criteria with weights. Goes beyond semantic similarity to prioritize memories based on domain-specific signals.
|
||||
|
||||
### Configuration
|
||||
|
||||
```python
|
||||
# Define criteria at project level
|
||||
retrieval_criteria = [
|
||||
{"name": "joy", "description": "Positive emotions like happiness and excitement", "weight": 3},
|
||||
{"name": "curiosity", "description": "Inquisitiveness and desire to learn", "weight": 2},
|
||||
{"name": "urgency", "description": "Time-sensitive or high-priority items", "weight": 4},
|
||||
]
|
||||
client.project.update(retrieval_criteria=retrieval_criteria)
|
||||
```
|
||||
|
||||
```typescript
|
||||
await client.updateProject({
|
||||
retrievalCriteria: [
|
||||
{ name: 'joy', description: 'Positive emotions', weight: 3 },
|
||||
{ name: 'urgency', description: 'Time-sensitive items', weight: 4 },
|
||||
],
|
||||
});
|
||||
```
|
||||
|
||||
### Usage
|
||||
|
||||
Once configured, `client.search()` automatically applies criteria ranking:
|
||||
|
||||
```python
|
||||
# Criteria-weighted results returned automatically
|
||||
results = client.search("Why am I feeling happy?", filters={"user_id": "alice"})
|
||||
```
|
||||
|
||||
**Best for:** Wellness assistants, tutoring platforms, productivity tools — any app needing intent-aware retrieval.
|
||||
|
||||
---
|
||||
|
||||
## Feedback Mechanism
|
||||
|
||||
Provide feedback on extracted memories to improve system quality over time.
|
||||
|
||||
@@ -0,0 +1,144 @@
|
||||
"""Regression tests: assistant-authored text must never be posted as role="user".
|
||||
|
||||
The Stop hook (capture_session_summary) and the post-compact hook
|
||||
(capture_compact_summary) both ship *model-authored* prose to
|
||||
POST /v3/memories/add/. Mem0's fact extractor renders each message as
|
||||
"{role}: {content}" and is instructed to extract "facts and preferences about
|
||||
the user" — so role is the only signal separating what the human said from what
|
||||
Claude said.
|
||||
|
||||
Posting Claude's own words under role="user" made the extractor read Claude's
|
||||
first-person prose ("I recommend pgvector", "I found the bug in auth.py") as the
|
||||
*human's* statements and store them under their user_id. The Stop hook fires on
|
||||
every assistant turn, so this corrupted memory on nearly every message.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
|
||||
class _FakeResp:
|
||||
status = 200
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *_):
|
||||
return False
|
||||
|
||||
|
||||
def _capture(monkeypatch, module):
|
||||
"""Patch urlopen so store_summary posts nowhere; capture the request body."""
|
||||
captured: dict = {}
|
||||
|
||||
def fake_urlopen(req, timeout=0):
|
||||
captured["body"] = json.loads(req.data.decode("utf-8"))
|
||||
return _FakeResp()
|
||||
|
||||
monkeypatch.setattr(module.urllib.request, "urlopen", fake_urlopen)
|
||||
return captured
|
||||
|
||||
|
||||
# Claude's own voice — first-person prose that must never be attributed to the human.
|
||||
ASSISTANT_PROSE = (
|
||||
"I traced the root cause to auth.py and I recommend we switch to pgvector "
|
||||
"for the vector store. I'll refactor the session handler next."
|
||||
)
|
||||
|
||||
|
||||
def test_session_summary_posts_assistant_prose_as_assistant(monkeypatch):
|
||||
"""Stop hook: the last assistant message must be tagged role="assistant"."""
|
||||
import capture_session_summary as css
|
||||
|
||||
captured = _capture(monkeypatch, css)
|
||||
|
||||
css.store_summary(
|
||||
api_key="test-key",
|
||||
summary_prompt=css.build_summary_prompt(ASSISTANT_PROSE, []),
|
||||
user_id="u1",
|
||||
session_id="s1",
|
||||
project_id="p1",
|
||||
branch="main",
|
||||
files=[],
|
||||
)
|
||||
|
||||
messages = captured["body"]["messages"]
|
||||
for msg in messages:
|
||||
if ASSISTANT_PROSE in msg["content"]:
|
||||
assert msg["role"] == "assistant", (
|
||||
"Claude's own words were posted as role='user' — mem0 will extract "
|
||||
"them as facts about the human. Got role=%r" % msg["role"]
|
||||
)
|
||||
break
|
||||
else:
|
||||
raise AssertionError("assistant prose never made it into the payload")
|
||||
|
||||
|
||||
def test_compact_summary_posts_assistant_prose_as_assistant(monkeypatch):
|
||||
"""Post-compact hook: the compact summary is model-authored, not user-authored."""
|
||||
import capture_compact_summary as ccs
|
||||
|
||||
captured = _capture(monkeypatch, ccs)
|
||||
|
||||
ccs.store_summary(
|
||||
api_key="test-key",
|
||||
summary=ASSISTANT_PROSE,
|
||||
user_id="u1",
|
||||
session_id="s1",
|
||||
project_id="p1",
|
||||
branch="main",
|
||||
)
|
||||
|
||||
messages = captured["body"]["messages"]
|
||||
for msg in messages:
|
||||
if ASSISTANT_PROSE in msg["content"]:
|
||||
assert msg["role"] == "assistant", (
|
||||
"Compact summary (written by Claude) was posted as role='user'. Got role=%r" % msg["role"]
|
||||
)
|
||||
break
|
||||
else:
|
||||
raise AssertionError("assistant prose never made it into the payload")
|
||||
|
||||
|
||||
def test_no_user_role_message_carries_assistant_prose(monkeypatch):
|
||||
"""Belt and braces: no user-role message may contain the assistant's words."""
|
||||
import capture_session_summary as css
|
||||
|
||||
captured = _capture(monkeypatch, css)
|
||||
|
||||
css.store_summary(
|
||||
api_key="test-key",
|
||||
summary_prompt=css.build_summary_prompt(ASSISTANT_PROSE, ["auth.py"]),
|
||||
user_id="u1",
|
||||
session_id="s1",
|
||||
project_id="p1",
|
||||
branch="main",
|
||||
files=["auth.py"],
|
||||
)
|
||||
|
||||
for msg in captured["body"]["messages"]:
|
||||
if msg["role"] == "user":
|
||||
assert ASSISTANT_PROSE not in msg["content"], (
|
||||
"A user-role message carries Claude's prose — this is the misattribution bug."
|
||||
)
|
||||
|
||||
|
||||
def test_auto_capture_preserves_real_roles():
|
||||
"""auto_capture is the reference: it must pass roles through untouched."""
|
||||
import auto_capture
|
||||
|
||||
lines = [
|
||||
json.dumps({"type": "user", "message": {"role": "user", "content": "why is the build failing on main?"}}),
|
||||
json.dumps(
|
||||
{
|
||||
"type": "assistant",
|
||||
"message": {"role": "assistant", "content": [{"type": "text", "text": ASSISTANT_PROSE}]},
|
||||
}
|
||||
),
|
||||
]
|
||||
|
||||
messages = auto_capture.extract_recent_exchanges(lines)
|
||||
|
||||
assert [m["role"] for m in messages] == ["user", "assistant"]
|
||||
assert ASSISTANT_PROSE in messages[1]["content"]
|
||||
+37
-1
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "mem0ai",
|
||||
"version": "3.1.0",
|
||||
"version": "3.1.1",
|
||||
"description": "The Memory Layer For Your AI Apps",
|
||||
"main": "./dist/index.js",
|
||||
"module": "./dist/index.mjs",
|
||||
@@ -151,6 +151,42 @@
|
||||
"@aws-sdk/client-bedrock-runtime": ">=3.0.0 <3.968.0"
|
||||
},
|
||||
"peerDependenciesMeta": {
|
||||
"@qdrant/js-client-rest": {
|
||||
"optional": true
|
||||
},
|
||||
"redis": {
|
||||
"optional": true
|
||||
},
|
||||
"@supabase/supabase-js": {
|
||||
"optional": true
|
||||
},
|
||||
"cloudflare": {
|
||||
"optional": true
|
||||
},
|
||||
"@azure/search-documents": {
|
||||
"optional": true
|
||||
},
|
||||
"@azure/identity": {
|
||||
"optional": true
|
||||
},
|
||||
"@langchain/core": {
|
||||
"optional": true
|
||||
},
|
||||
"@anthropic-ai/sdk": {
|
||||
"optional": true
|
||||
},
|
||||
"@google/genai": {
|
||||
"optional": true
|
||||
},
|
||||
"groq-sdk": {
|
||||
"optional": true
|
||||
},
|
||||
"@mistralai/mistralai": {
|
||||
"optional": true
|
||||
},
|
||||
"ollama": {
|
||||
"optional": true
|
||||
},
|
||||
"mysql2": {
|
||||
"optional": true
|
||||
},
|
||||
|
||||
@@ -59,7 +59,6 @@ export interface ProjectOptions {
|
||||
export interface PromptUpdatePayload {
|
||||
customInstructions?: string;
|
||||
customCategories?: custom_categories[];
|
||||
retrievalCriteria?: any[];
|
||||
version?: string;
|
||||
memoryDepth?: string | null;
|
||||
usecaseSetting?: string | number;
|
||||
|
||||
@@ -0,0 +1,292 @@
|
||||
import { Embedder } from "./base";
|
||||
import { EmbeddingConfig } from "../types";
|
||||
|
||||
const DEFAULT_MODEL = "amazon.titan-embed-text-v1";
|
||||
const DEFAULT_REGION = "us-west-2";
|
||||
|
||||
// Cohere's Bedrock embed API rejects an InvokeModel call carrying more than 96
|
||||
// texts, so `embedBatch` chunks at that boundary.
|
||||
const COHERE_MAX_BATCH = 96;
|
||||
|
||||
// Titan has no server-side batch endpoint -- one InvokeModel call per text --
|
||||
// so without a cap a large embedBatch() would fan out one request per text.
|
||||
// Bounds concurrency the same way COHERE_MAX_BATCH bounds the Cohere path.
|
||||
const TITAN_MAX_CONCURRENCY = 4;
|
||||
|
||||
// Cohere wants to know whether a text is being embedded for storage or for a
|
||||
// retrieval query; embedding a search query in document mode silently
|
||||
// degrades retrieval. Titan ignores this and has no equivalent parameter.
|
||||
const COHERE_INPUT_TYPES: Record<"add" | "update" | "search", string> = {
|
||||
add: "search_document",
|
||||
update: "search_document",
|
||||
search: "search_query",
|
||||
};
|
||||
|
||||
type BedrockRuntimeModule = typeof import("@aws-sdk/client-bedrock-runtime");
|
||||
|
||||
interface BedrockCredentials {
|
||||
accessKeyId: string;
|
||||
secretAccessKey: string;
|
||||
sessionToken?: string;
|
||||
}
|
||||
|
||||
interface BedrockEmbeddingResponse {
|
||||
// Titan returns a single vector. Cohere v3 returns a flat array of vectors;
|
||||
// Cohere v4, when `embedding_types` is requested, nests it as `{ float }`.
|
||||
embedding?: number[];
|
||||
embeddings?: number[][] | { float?: number[][] };
|
||||
}
|
||||
|
||||
/**
|
||||
* Runs `fn` over `items` with at most `limit` calls in flight at once,
|
||||
* returning results in input order regardless of completion order.
|
||||
*/
|
||||
async function mapWithConcurrencyLimit<T, R>(
|
||||
items: T[],
|
||||
limit: number,
|
||||
fn: (item: T) => Promise<R>,
|
||||
): Promise<R[]> {
|
||||
const results: R[] = new Array(items.length);
|
||||
let next = 0;
|
||||
|
||||
async function worker(): Promise<void> {
|
||||
while (next < items.length) {
|
||||
const index = next++;
|
||||
results[index] = await fn(items[index]);
|
||||
}
|
||||
}
|
||||
|
||||
await Promise.all(
|
||||
Array.from({ length: Math.min(limit, items.length) }, worker),
|
||||
);
|
||||
return results;
|
||||
}
|
||||
|
||||
/**
|
||||
* AWS Bedrock embedder, mirroring `mem0/embeddings/aws_bedrock.py`.
|
||||
*
|
||||
* Supports the Amazon Titan and Cohere embedding model families. The
|
||||
* `@aws-sdk/client-bedrock-runtime` dependency is lazily imported so the
|
||||
* package stays optional: importing this module never forces the SDK to be
|
||||
* installed until a Bedrock embedder actually embeds something.
|
||||
*/
|
||||
export class AWSBedrockEmbedder implements Embedder {
|
||||
private readonly model: string;
|
||||
private readonly region: string;
|
||||
private readonly embeddingDims?: number;
|
||||
private readonly credentials?: BedrockCredentials;
|
||||
private clientPromise?: Promise<{
|
||||
sdk: BedrockRuntimeModule;
|
||||
client: { send: (command: any) => Promise<{ body?: Uint8Array }> };
|
||||
}>;
|
||||
|
||||
constructor(config: EmbeddingConfig) {
|
||||
this.model = config.model || DEFAULT_MODEL;
|
||||
this.region = config.awsRegion || process.env.AWS_REGION || DEFAULT_REGION;
|
||||
this.embeddingDims = config.embeddingDims;
|
||||
|
||||
const hasKeyPair = Boolean(
|
||||
config.awsAccessKeyId && config.awsSecretAccessKey,
|
||||
);
|
||||
const hasAnyCredential = Boolean(
|
||||
config.awsAccessKeyId ||
|
||||
config.awsSecretAccessKey ||
|
||||
config.awsSessionToken,
|
||||
);
|
||||
|
||||
// Partially configured credentials would silently fall back to the default
|
||||
// chain, embedding under an identity the caller never chose.
|
||||
if (hasAnyCredential && !hasKeyPair) {
|
||||
throw new Error(
|
||||
"AWS Bedrock requires both awsAccessKeyId and awsSecretAccessKey when any explicit credential is configured. " +
|
||||
"Omit all credential fields to use the AWS default credential chain.",
|
||||
);
|
||||
}
|
||||
|
||||
// Leaving `credentials` unset lets the AWS SDK resolve them from its
|
||||
// default chain: environment, shared config, SSO, or the instance role.
|
||||
if (hasKeyPair) {
|
||||
this.credentials = {
|
||||
accessKeyId: config.awsAccessKeyId!,
|
||||
secretAccessKey: config.awsSecretAccessKey!,
|
||||
...(config.awsSessionToken && { sessionToken: config.awsSessionToken }),
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
private async loadSdk(): Promise<BedrockRuntimeModule> {
|
||||
try {
|
||||
return await import("@aws-sdk/client-bedrock-runtime");
|
||||
} catch (error) {
|
||||
// Only a genuine module-resolution failure gets the friendly install
|
||||
// hint. Node's native ESM loader raises ERR_MODULE_NOT_FOUND; Jest's
|
||||
// and bundlers' CJS-style resolvers raise MODULE_NOT_FOUND. Anything
|
||||
// else (e.g. the package is installed but throws while loading, such
|
||||
// as on a Node version older than the SDK's own engines requirement)
|
||||
// rethrows unchanged instead of being misreported as "not installed".
|
||||
const code = (error as { code?: string } | undefined)?.code;
|
||||
if (code === "ERR_MODULE_NOT_FOUND" || code === "MODULE_NOT_FOUND") {
|
||||
throw Object.assign(
|
||||
new Error(
|
||||
"The '@aws-sdk/client-bedrock-runtime' package is required to use the AWS Bedrock embedder. " +
|
||||
"Install it with: npm install @aws-sdk/client-bedrock-runtime",
|
||||
),
|
||||
{ cause: error },
|
||||
);
|
||||
}
|
||||
throw error;
|
||||
}
|
||||
}
|
||||
|
||||
private async createClient(): Promise<{
|
||||
sdk: BedrockRuntimeModule;
|
||||
client: { send: (command: any) => Promise<{ body?: Uint8Array }> };
|
||||
}> {
|
||||
const sdk = await this.loadSdk();
|
||||
return {
|
||||
sdk,
|
||||
client: new sdk.BedrockRuntimeClient({
|
||||
region: this.region,
|
||||
...(this.credentials && { credentials: this.credentials }),
|
||||
}),
|
||||
};
|
||||
}
|
||||
|
||||
private getClient() {
|
||||
// Memoized so concurrent embed() calls share one client instead of each
|
||||
// racing to build their own. Cleared on rejection so a transient failure
|
||||
// (e.g. a network blip while resolving credentials) doesn't permanently
|
||||
// disable Bedrock for the rest of this embedder's lifetime.
|
||||
if (!this.clientPromise) {
|
||||
this.clientPromise = this.createClient().catch((err) => {
|
||||
this.clientPromise = undefined;
|
||||
throw err;
|
||||
});
|
||||
}
|
||||
return this.clientPromise;
|
||||
}
|
||||
|
||||
private isCohereModel(): boolean {
|
||||
return this.model.startsWith("cohere.");
|
||||
}
|
||||
|
||||
private isCohereV4Model(): boolean {
|
||||
return this.model.includes("embed-v4");
|
||||
}
|
||||
|
||||
private buildRequestBody(
|
||||
texts: string[],
|
||||
memoryAction?: "add" | "update" | "search",
|
||||
): Record<string, unknown> {
|
||||
if (this.isCohereModel()) {
|
||||
const body: Record<string, unknown> = {
|
||||
texts,
|
||||
input_type: memoryAction
|
||||
? COHERE_INPUT_TYPES[memoryAction]
|
||||
: "search_document",
|
||||
};
|
||||
|
||||
// Only Embed v4 understands embedding_types / output_dimension; v3
|
||||
// rejects unknown fields, so they're guarded to the v4 model family.
|
||||
if (this.isCohereV4Model()) {
|
||||
body.embedding_types = ["float"];
|
||||
if (this.embeddingDims !== undefined) {
|
||||
body.output_dimension = this.embeddingDims;
|
||||
}
|
||||
}
|
||||
return body;
|
||||
}
|
||||
|
||||
// Titan accepts one text per call. Only Titan Text Embeddings V2 supports
|
||||
// a caller-chosen output size (256/512/1024), so the field is guarded the
|
||||
// same way the Python provider guards it.
|
||||
return {
|
||||
inputText: texts[0],
|
||||
...(this.embeddingDims !== undefined &&
|
||||
this.model.includes("titan-embed-text-v2") && {
|
||||
dimensions: this.embeddingDims,
|
||||
}),
|
||||
};
|
||||
}
|
||||
|
||||
private async invoke(
|
||||
texts: string[],
|
||||
memoryAction?: "add" | "update" | "search",
|
||||
): Promise<number[][]> {
|
||||
const { sdk, client } = await this.getClient();
|
||||
|
||||
let payload: BedrockEmbeddingResponse;
|
||||
try {
|
||||
const response = await client.send(
|
||||
new sdk.InvokeModelCommand({
|
||||
modelId: this.model,
|
||||
contentType: "application/json",
|
||||
accept: "application/json",
|
||||
body: new TextEncoder().encode(
|
||||
JSON.stringify(this.buildRequestBody(texts, memoryAction)),
|
||||
),
|
||||
}),
|
||||
);
|
||||
payload = JSON.parse(new TextDecoder().decode(response.body));
|
||||
} catch (error) {
|
||||
const message = error instanceof Error ? error.message : String(error);
|
||||
throw new Error(
|
||||
`Error getting embedding from AWS Bedrock model ${this.model}: ${message}`,
|
||||
);
|
||||
}
|
||||
|
||||
// Validated outside the try so this message is not re-wrapped by the catch.
|
||||
// Cohere v3 replies with a flat `embeddings` array; v4 (when
|
||||
// embedding_types is requested) nests it under `.float`.
|
||||
const embeddings = this.isCohereModel()
|
||||
? Array.isArray(payload.embeddings)
|
||||
? payload.embeddings
|
||||
: payload.embeddings?.float
|
||||
: payload.embedding && [payload.embedding];
|
||||
|
||||
// `[]` is truthy, so a lone zero-length vector must be checked for
|
||||
// explicitly -- otherwise it passes the length check and hands the
|
||||
// caller an empty embedding instead of an error.
|
||||
if (
|
||||
!embeddings ||
|
||||
embeddings.length !== texts.length ||
|
||||
embeddings.some((embedding) => embedding.length === 0)
|
||||
) {
|
||||
throw new Error(
|
||||
`AWS Bedrock model ${this.model} returned no embedding for one or more inputs`,
|
||||
);
|
||||
}
|
||||
return embeddings;
|
||||
}
|
||||
|
||||
async embed(
|
||||
text: string,
|
||||
memoryAction?: "add" | "update" | "search",
|
||||
): Promise<number[]> {
|
||||
return (await this.invoke([text], memoryAction))[0];
|
||||
}
|
||||
|
||||
async embedBatch(
|
||||
texts: string[],
|
||||
memoryAction?: "add" | "update" | "search",
|
||||
): Promise<number[][]> {
|
||||
if (texts.length === 0) return [];
|
||||
|
||||
if (!this.isCohereModel()) {
|
||||
return mapWithConcurrencyLimit(texts, TITAN_MAX_CONCURRENCY, (text) =>
|
||||
this.embed(text, memoryAction),
|
||||
);
|
||||
}
|
||||
|
||||
const embeddings: number[][] = [];
|
||||
for (let i = 0; i < texts.length; i += COHERE_MAX_BATCH) {
|
||||
embeddings.push(
|
||||
...(await this.invoke(
|
||||
texts.slice(i, i + COHERE_MAX_BATCH),
|
||||
memoryAction,
|
||||
)),
|
||||
);
|
||||
}
|
||||
return embeddings;
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { FlagEmbedding } from "fastembed";
|
||||
import { Embedder } from "./base";
|
||||
import { EmbeddingConfig } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
// FastEmbed only ships a fixed set of ONNX models (fastembed's `EmbeddingModel`
|
||||
// enum, minus CUSTOM). Mirrored here as literals so an invalid model name can
|
||||
@@ -54,14 +55,11 @@ export class FastEmbedEmbedder implements Embedder {
|
||||
* consumers that never touch FastEmbed don't need it installed.
|
||||
*/
|
||||
private async initEmbeddingModel(): Promise<FlagEmbedding> {
|
||||
let sdk: any;
|
||||
try {
|
||||
sdk = await import("fastembed");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The 'fastembed' package is required to use the FastEmbed embedder. Install it with: npm install fastembed",
|
||||
);
|
||||
}
|
||||
const sdk = await loadPeer(
|
||||
"fastembed",
|
||||
"FastEmbed embedder",
|
||||
() => import("fastembed"),
|
||||
);
|
||||
|
||||
return sdk.FlagEmbedding.init({ model: this.modelName });
|
||||
}
|
||||
|
||||
@@ -1,21 +1,32 @@
|
||||
import { GoogleGenAI } from "@google/genai";
|
||||
import type { GoogleGenAI } from "@google/genai";
|
||||
import { Embedder } from "./base";
|
||||
import { EmbeddingConfig } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class GoogleEmbedder implements Embedder {
|
||||
private google: GoogleGenAI;
|
||||
private google!: GoogleGenAI;
|
||||
private model: string;
|
||||
private embeddingDims: number | undefined;
|
||||
private readonly apiKey: string | undefined;
|
||||
|
||||
constructor(config: EmbeddingConfig) {
|
||||
this.google = new GoogleGenAI({
|
||||
apiKey: config.apiKey || process.env.GOOGLE_API_KEY,
|
||||
});
|
||||
this.apiKey = config.apiKey || process.env.GOOGLE_API_KEY;
|
||||
this.model = config.model || "gemini-embedding-001";
|
||||
this.embeddingDims = config.embeddingDims;
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.google) return;
|
||||
const sdk = await loadPeer(
|
||||
"@google/genai",
|
||||
"Google embedder",
|
||||
() => import("@google/genai"),
|
||||
);
|
||||
this.google = new sdk.GoogleGenAI({ apiKey: this.apiKey });
|
||||
}
|
||||
|
||||
async embed(text: string): Promise<number[]> {
|
||||
await this.ensureClient();
|
||||
const response = await this.google.models.embedContent({
|
||||
model: this.model,
|
||||
contents: text,
|
||||
@@ -27,6 +38,7 @@ export class GoogleEmbedder implements Embedder {
|
||||
}
|
||||
|
||||
async embedBatch(texts: string[]): Promise<number[][]> {
|
||||
await this.ensureClient();
|
||||
const response = await this.google.models.embedContent({
|
||||
model: this.model,
|
||||
contents: texts,
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { Embeddings } from "@langchain/core/embeddings";
|
||||
import type { Embeddings } from "@langchain/core/embeddings";
|
||||
import { Embedder } from "./base";
|
||||
import { EmbeddingConfig } from "../types";
|
||||
|
||||
|
||||
@@ -1,19 +1,19 @@
|
||||
import { Ollama } from "ollama";
|
||||
import type { Ollama } from "ollama";
|
||||
import { Embedder } from "./base";
|
||||
import { EmbeddingConfig } from "../types";
|
||||
import { logger } from "../utils/logger";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class OllamaEmbedder implements Embedder {
|
||||
private ollama: Ollama;
|
||||
private ollama!: Ollama;
|
||||
private model: string;
|
||||
private embeddingDims?: number;
|
||||
private readonly host: string;
|
||||
// Using this variable to avoid calling the Ollama server multiple times
|
||||
private initialized: boolean = false;
|
||||
|
||||
constructor(config: EmbeddingConfig) {
|
||||
this.ollama = new Ollama({
|
||||
host: config.url || config.baseURL || "http://localhost:11434",
|
||||
});
|
||||
this.host = config.url || config.baseURL || "http://localhost:11434";
|
||||
this.model = config.model || "nomic-embed-text:latest";
|
||||
this.embeddingDims = config.embeddingDims || 768;
|
||||
this.ensureModelExists().catch((err) => {
|
||||
@@ -21,7 +21,18 @@ export class OllamaEmbedder implements Embedder {
|
||||
});
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.ollama) return;
|
||||
const sdk = await loadPeer(
|
||||
"ollama",
|
||||
"Ollama embedder",
|
||||
() => import("ollama"),
|
||||
);
|
||||
this.ollama = new sdk.Ollama({ host: this.host });
|
||||
}
|
||||
|
||||
async embed(text: string): Promise<number[]> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
await this.ensureModelExists();
|
||||
} catch (err) {
|
||||
@@ -54,6 +65,7 @@ export class OllamaEmbedder implements Embedder {
|
||||
if (this.initialized) {
|
||||
return true;
|
||||
}
|
||||
await this.ensureClient();
|
||||
const local_models = await this.ollama.list();
|
||||
const target = OllamaEmbedder.normalizeModelName(this.model);
|
||||
if (
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { PredictionServiceClient } from "@google-cloud/aiplatform";
|
||||
import { Embedder } from "./base";
|
||||
import { VertexAIConfig } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
type AIPlatform = typeof import("@google-cloud/aiplatform");
|
||||
type ClientOptions = NonNullable<
|
||||
@@ -105,15 +106,11 @@ export class VertexAIEmbedder implements Embedder {
|
||||
}
|
||||
|
||||
private async createClient(): Promise<void> {
|
||||
let aiplatform: AIPlatform;
|
||||
try {
|
||||
aiplatform = await import("@google-cloud/aiplatform");
|
||||
} catch (err) {
|
||||
throw new Error(
|
||||
"Failed to import '@google-cloud/aiplatform'. Please install it to use the Vertex AI embedding provider: " +
|
||||
(err as Error).message,
|
||||
);
|
||||
}
|
||||
const aiplatform: AIPlatform = await loadPeer(
|
||||
"@google-cloud/aiplatform",
|
||||
"Vertex AI embedding provider",
|
||||
() => import("@google-cloud/aiplatform"),
|
||||
);
|
||||
|
||||
const client = new aiplatform.PredictionServiceClient(this.clientOptions);
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ export * from "./memory";
|
||||
export * from "./memory/memory.types";
|
||||
export * from "./types";
|
||||
export * from "./embeddings/base";
|
||||
export * from "./embeddings/aws_bedrock";
|
||||
export * from "./embeddings/huggingface";
|
||||
export * from "./embeddings/openai";
|
||||
export * from "./embeddings/ollama";
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
import Anthropic from "@anthropic-ai/sdk";
|
||||
import type Anthropic from "@anthropic-ai/sdk";
|
||||
import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class AnthropicLLM implements LLM {
|
||||
private client: Anthropic;
|
||||
private client!: Anthropic;
|
||||
private readonly clientArgs: { apiKey: string; baseURL?: string };
|
||||
private model: string;
|
||||
private maxTokens: number;
|
||||
private temperature?: number;
|
||||
@@ -20,7 +22,7 @@ export class AnthropicLLM implements LLM {
|
||||
if (config.baseURL) {
|
||||
clientArgs.baseURL = config.baseURL;
|
||||
}
|
||||
this.client = new Anthropic(clientArgs);
|
||||
this.clientArgs = clientArgs;
|
||||
this.model = config.model || "claude-sonnet-4-6";
|
||||
// Defaults mirror the Python provider's AnthropicConfig
|
||||
// (max_tokens=2000, temperature=0.1, top_p omitted).
|
||||
@@ -29,11 +31,22 @@ export class AnthropicLLM implements LLM {
|
||||
this.topP = config.topP;
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const sdk = await loadPeer(
|
||||
"@anthropic-ai/sdk",
|
||||
"Anthropic LLM",
|
||||
() => import("@anthropic-ai/sdk"),
|
||||
);
|
||||
this.client = new sdk.default(this.clientArgs);
|
||||
}
|
||||
|
||||
async generateResponse(
|
||||
messages: Message[],
|
||||
responseFormat?: { type: string },
|
||||
tools?: any[],
|
||||
): Promise<string | LLMResponse> {
|
||||
await this.ensureClient();
|
||||
// Extract system message if present
|
||||
const systemMessage = messages.find((msg) => msg.role === "system");
|
||||
const otherMessages = messages.filter((msg) => msg.role !== "system");
|
||||
|
||||
@@ -1,16 +1,28 @@
|
||||
import { GoogleGenAI } from "@google/genai";
|
||||
import type { GoogleGenAI } from "@google/genai";
|
||||
import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class GoogleLLM implements LLM {
|
||||
private google: GoogleGenAI;
|
||||
private google!: GoogleGenAI;
|
||||
private model: string;
|
||||
private readonly apiKey: string | undefined;
|
||||
|
||||
constructor(config: LLMConfig) {
|
||||
this.google = new GoogleGenAI({ apiKey: config.apiKey });
|
||||
this.apiKey = config.apiKey;
|
||||
this.model = config.model || "gemini-2.0-flash";
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.google) return;
|
||||
const sdk = await loadPeer(
|
||||
"@google/genai",
|
||||
"Google LLM",
|
||||
() => import("@google/genai"),
|
||||
);
|
||||
this.google = new sdk.GoogleGenAI({ apiKey: this.apiKey });
|
||||
}
|
||||
|
||||
private formatContents(messages: Message[]) {
|
||||
return messages.map((msg) => ({
|
||||
parts: [
|
||||
@@ -30,6 +42,7 @@ export class GoogleLLM implements LLM {
|
||||
responseFormat?: { type: string },
|
||||
tools?: any[],
|
||||
): Promise<string | LLMResponse> {
|
||||
await this.ensureClient();
|
||||
const contents = this.formatContents(messages);
|
||||
|
||||
// Build config with tools if provided
|
||||
@@ -72,6 +85,7 @@ export class GoogleLLM implements LLM {
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
await this.ensureClient();
|
||||
const completion = await this.google.models.generateContent({
|
||||
contents: this.formatContents(messages),
|
||||
model: this.model,
|
||||
|
||||
@@ -1,24 +1,37 @@
|
||||
import { Groq } from "groq-sdk";
|
||||
import type { Groq } from "groq-sdk";
|
||||
import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class GroqLLM implements LLM {
|
||||
private client: Groq;
|
||||
private client!: Groq;
|
||||
private model: string;
|
||||
private readonly apiKey: string;
|
||||
|
||||
constructor(config: LLMConfig) {
|
||||
const apiKey = config.apiKey || process.env.GROQ_API_KEY;
|
||||
if (!apiKey) {
|
||||
throw new Error("Groq API key is required");
|
||||
}
|
||||
this.client = new Groq({ apiKey });
|
||||
this.apiKey = apiKey;
|
||||
this.model = config.model || "llama3-70b-8192";
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const sdk = await loadPeer(
|
||||
"groq-sdk",
|
||||
"Groq LLM",
|
||||
() => import("groq-sdk"),
|
||||
);
|
||||
this.client = new sdk.Groq({ apiKey: this.apiKey });
|
||||
}
|
||||
|
||||
async generateResponse(
|
||||
messages: Message[],
|
||||
responseFormat?: { type: string },
|
||||
): Promise<string> {
|
||||
await this.ensureClient();
|
||||
const response = await this.client.chat.completions.create({
|
||||
model: this.model,
|
||||
messages: messages.map((msg) => ({
|
||||
@@ -35,6 +48,7 @@ export class GroqLLM implements LLM {
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
await this.ensureClient();
|
||||
const response = await this.client.chat.completions.create({
|
||||
model: this.model,
|
||||
messages: messages.map((msg) => ({
|
||||
|
||||
@@ -1,17 +1,16 @@
|
||||
import { BaseLanguageModel } from "@langchain/core/language_models/base";
|
||||
import {
|
||||
AIMessage,
|
||||
HumanMessage,
|
||||
SystemMessage,
|
||||
BaseMessage,
|
||||
} from "@langchain/core/messages";
|
||||
import type { BaseLanguageModel } from "@langchain/core/language_models/base";
|
||||
import type { BaseMessage } from "@langchain/core/messages";
|
||||
import { z } from "zod";
|
||||
import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types/index";
|
||||
// Import the schemas directly into LangchainLLM
|
||||
import { FactRetrievalSchema, MemoryUpdateSchema } from "../prompts";
|
||||
|
||||
const convertToLangchainMessages = (messages: Message[]): BaseMessage[] => {
|
||||
const convertToLangchainMessages = async (
|
||||
messages: Message[],
|
||||
): Promise<BaseMessage[]> => {
|
||||
const { AIMessage, HumanMessage, SystemMessage } =
|
||||
await import("@langchain/core/messages");
|
||||
return messages.map((msg) => {
|
||||
const content =
|
||||
typeof msg.content === "string"
|
||||
@@ -62,7 +61,7 @@ export class LangchainLLM implements LLM {
|
||||
response_format?: { type: string },
|
||||
tools?: any[],
|
||||
): Promise<string | LLMResponse> {
|
||||
const langchainMessages = convertToLangchainMessages(messages);
|
||||
const langchainMessages = await convertToLangchainMessages(messages);
|
||||
let runnable: any = this.llmInstance;
|
||||
const invokeOptions: Record<string, any> = {};
|
||||
let isStructuredOutput = false;
|
||||
@@ -170,7 +169,7 @@ export class LangchainLLM implements LLM {
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
const langchainMessages = convertToLangchainMessages(messages);
|
||||
const langchainMessages = await convertToLangchainMessages(messages);
|
||||
try {
|
||||
const response = await this.llmInstance.invoke(langchainMessages);
|
||||
if (response && typeof response.content === "string") {
|
||||
|
||||
@@ -1,21 +1,31 @@
|
||||
import { Mistral } from "@mistralai/mistralai";
|
||||
import type { Mistral } from "@mistralai/mistralai";
|
||||
import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class MistralLLM implements LLM {
|
||||
private client: Mistral;
|
||||
private client!: Mistral;
|
||||
private model: string;
|
||||
private readonly apiKey: string;
|
||||
|
||||
constructor(config: LLMConfig) {
|
||||
if (!config.apiKey) {
|
||||
throw new Error("Mistral API key is required");
|
||||
}
|
||||
this.client = new Mistral({
|
||||
apiKey: config.apiKey,
|
||||
});
|
||||
this.apiKey = config.apiKey;
|
||||
this.model = config.model || "mistral-tiny-latest";
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const sdk = await loadPeer(
|
||||
"@mistralai/mistralai",
|
||||
"Mistral LLM",
|
||||
() => import("@mistralai/mistralai"),
|
||||
);
|
||||
this.client = new sdk.Mistral({ apiKey: this.apiKey });
|
||||
}
|
||||
|
||||
// Helper function to convert content to string
|
||||
private contentToString(content: any): string {
|
||||
if (typeof content === "string") {
|
||||
@@ -41,6 +51,7 @@ export class MistralLLM implements LLM {
|
||||
responseFormat?: { type: string },
|
||||
tools?: any[],
|
||||
): Promise<string | LLMResponse> {
|
||||
await this.ensureClient();
|
||||
const response = await this.client.chat.complete({
|
||||
model: this.model,
|
||||
messages: messages.map((msg) => ({
|
||||
@@ -82,6 +93,7 @@ export class MistralLLM implements LLM {
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
await this.ensureClient();
|
||||
const formattedMessages = messages.map((msg) => ({
|
||||
role: msg.role as "system" | "user" | "assistant",
|
||||
content:
|
||||
|
||||
@@ -1,29 +1,36 @@
|
||||
import { Ollama } from "ollama";
|
||||
import type { Ollama } from "ollama";
|
||||
import { LLM, LLMResponse } from "./base";
|
||||
import { LLMConfig, Message } from "../types";
|
||||
import { logger } from "../utils/logger";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export class OllamaLLM implements LLM {
|
||||
private ollama: Ollama;
|
||||
private ollama!: Ollama;
|
||||
private model: string;
|
||||
private readonly host: string;
|
||||
// Using this variable to avoid calling the Ollama server multiple times
|
||||
private initialized: boolean = false;
|
||||
|
||||
constructor(config: LLMConfig) {
|
||||
this.ollama = new Ollama({
|
||||
host: config.url || config.baseURL || "http://localhost:11434",
|
||||
});
|
||||
this.host = config.url || config.baseURL || "http://localhost:11434";
|
||||
this.model = config.model || "llama3.1:8b";
|
||||
this.ensureModelExists().catch((err) => {
|
||||
logger.error(`Error ensuring model exists: ${err}`);
|
||||
});
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.ollama) return;
|
||||
const sdk = await loadPeer("ollama", "Ollama LLM", () => import("ollama"));
|
||||
this.ollama = new sdk.Ollama({ host: this.host });
|
||||
}
|
||||
|
||||
async generateResponse(
|
||||
messages: Message[],
|
||||
responseFormat?: { type: string },
|
||||
tools?: any[],
|
||||
): Promise<string | LLMResponse> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
await this.ensureModelExists();
|
||||
} catch (err) {
|
||||
@@ -63,6 +70,7 @@ export class OllamaLLM implements LLM {
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
await this.ensureModelExists();
|
||||
} catch (err) {
|
||||
@@ -93,6 +101,7 @@ export class OllamaLLM implements LLM {
|
||||
if (this.initialized) {
|
||||
return true;
|
||||
}
|
||||
await this.ensureClient();
|
||||
const local_models = await this.ollama.list();
|
||||
if (!local_models.models.find((m: any) => m.name === this.model)) {
|
||||
logger.info(`Pulling model ${this.model}...`);
|
||||
|
||||
@@ -101,6 +101,10 @@ const ENTITY_PARAMS = [
|
||||
"runId",
|
||||
];
|
||||
|
||||
// Identity keys stripped from update() metadata: ENTITY_PARAMS covers user_id/agent_id/run_id
|
||||
// in both casings (the default store promotes camelCase on read); actor_id has no camelCase alias.
|
||||
const IDENTITY_KEYS = [...ENTITY_PARAMS, "actor_id"];
|
||||
|
||||
/**
|
||||
* Validates that no top-level entity parameters are passed in config.
|
||||
* @throws Error if entity params are found at top level
|
||||
@@ -122,18 +126,19 @@ function rejectTopLevelEntityParams(
|
||||
|
||||
/**
|
||||
* Validates and normalizes an entity ID.
|
||||
* - Coerces non-string ids (e.g. numeric database keys) to string
|
||||
* - Trims leading/trailing whitespace
|
||||
* - Rejects empty or whitespace-only strings
|
||||
* - Rejects strings containing internal whitespace
|
||||
* @returns The trimmed entity ID, or undefined if input is undefined
|
||||
* @returns The trimmed entity ID, or undefined if input is undefined/null
|
||||
* @throws Error if entity ID is invalid
|
||||
*/
|
||||
function validateAndTrimEntityId(
|
||||
value: string | undefined,
|
||||
value: string | number | undefined | null,
|
||||
name: string,
|
||||
): string | undefined {
|
||||
if (value === undefined) return undefined;
|
||||
const trimmed = value.trim();
|
||||
if (value == null) return undefined;
|
||||
const trimmed = String(value).trim();
|
||||
if (trimmed === "") {
|
||||
throw new Error(
|
||||
`Invalid ${name}: cannot be empty or whitespace-only. Provide a valid identifier.`,
|
||||
@@ -1941,9 +1946,14 @@ export class Memory {
|
||||
existingEmbeddings[newData] ||
|
||||
(await this.embedder.embed(newData, "update"));
|
||||
|
||||
// Caller metadata must not overwrite or inject an identity scope (#6342 / #6367).
|
||||
const sanitizedMetadata = Object.fromEntries(
|
||||
Object.entries(metadata).filter(([k]) => !IDENTITY_KEYS.includes(k)),
|
||||
);
|
||||
|
||||
const newMetadata = {
|
||||
...existingMemory.payload,
|
||||
...metadata,
|
||||
...sanitizedMetadata,
|
||||
data: newData,
|
||||
hash: createHash("md5").update(newData).digest("hex"),
|
||||
textLemmatized: lemmatizeForBm25(newData),
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import { RerankerConfig } from "../types";
|
||||
import { Reranker, RerankResult } from "./base";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
const DEFAULT_MODEL = "rerank-v3.5";
|
||||
|
||||
@@ -41,14 +42,11 @@ export class CohereReranker implements Reranker {
|
||||
}
|
||||
|
||||
private async createClient(): Promise<any> {
|
||||
let sdk: any;
|
||||
try {
|
||||
sdk = await import("cohere-ai");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The 'cohere-ai' package is required to use the Cohere reranker. Install it with: npm install cohere-ai",
|
||||
);
|
||||
}
|
||||
const sdk = await loadPeer(
|
||||
"cohere-ai",
|
||||
"Cohere reranker",
|
||||
() => import("cohere-ai"),
|
||||
);
|
||||
return new sdk.CohereClient({ token: this.apiKey });
|
||||
}
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import { RerankerConfig } from "../types";
|
||||
import { Reranker, RerankResult } from "./base";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
const DEFAULT_MODEL = "zerank-1";
|
||||
|
||||
@@ -37,14 +38,11 @@ export class ZeroEntropyReranker implements Reranker {
|
||||
}
|
||||
|
||||
private async createClient(): Promise<any> {
|
||||
let sdk: any;
|
||||
try {
|
||||
sdk = await import("zeroentropy");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The 'zeroentropy' package is required to use the ZeroEntropy reranker. Install it with: npm install zeroentropy",
|
||||
);
|
||||
}
|
||||
const sdk = await loadPeer(
|
||||
"zeroentropy",
|
||||
"ZeroEntropy reranker",
|
||||
() => import("zeroentropy"),
|
||||
);
|
||||
return new sdk.ZeroEntropy({ apiKey: this.apiKey });
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { createClient, SupabaseClient } from "@supabase/supabase-js";
|
||||
import type { SupabaseClient } from "@supabase/supabase-js";
|
||||
import { v4 as uuidv4 } from "uuid";
|
||||
import { HistoryManager } from "./base";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface HistoryEntry {
|
||||
id: string;
|
||||
@@ -20,16 +21,32 @@ interface SupabaseHistoryConfig {
|
||||
}
|
||||
|
||||
export class SupabaseHistoryManager implements HistoryManager {
|
||||
private supabase: SupabaseClient;
|
||||
// ponytail: benign double-construct race — two concurrent first-calls may each
|
||||
// build a client; createClient opens no connection, so last-write-wins is fine.
|
||||
private supabase!: SupabaseClient;
|
||||
private readonly supabaseUrl: string;
|
||||
private readonly supabaseKey: string;
|
||||
private readonly tableName: string;
|
||||
|
||||
constructor(config: SupabaseHistoryConfig) {
|
||||
this.tableName = config.tableName || "memory_history";
|
||||
this.supabase = createClient(config.supabaseUrl, config.supabaseKey);
|
||||
this.supabaseUrl = config.supabaseUrl;
|
||||
this.supabaseKey = config.supabaseKey;
|
||||
this.initializeSupabase().catch(console.error);
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.supabase) return;
|
||||
const sdk = await loadPeer(
|
||||
"@supabase/supabase-js",
|
||||
"Supabase history manager",
|
||||
() => import("@supabase/supabase-js"),
|
||||
);
|
||||
this.supabase = sdk.createClient(this.supabaseUrl, this.supabaseKey);
|
||||
}
|
||||
|
||||
private async initializeSupabase(): Promise<void> {
|
||||
await this.ensureClient();
|
||||
// Check if table exists
|
||||
const { error } = await this.supabase
|
||||
.from(this.tableName)
|
||||
@@ -65,6 +82,7 @@ create table ${this.tableName} (
|
||||
updatedAt?: string,
|
||||
isDeleted: number = 0,
|
||||
): Promise<void> {
|
||||
await this.ensureClient();
|
||||
const historyEntry: HistoryEntry = {
|
||||
id: uuidv4(),
|
||||
memory_id: memoryId,
|
||||
@@ -87,6 +105,7 @@ create table ${this.tableName} (
|
||||
}
|
||||
|
||||
async getHistory(memoryId: string): Promise<any[]> {
|
||||
await this.ensureClient();
|
||||
const { data, error } = await this.supabase
|
||||
.from(this.tableName)
|
||||
.select("*")
|
||||
@@ -103,6 +122,7 @@ create table ${this.tableName} (
|
||||
}
|
||||
|
||||
async reset(): Promise<void> {
|
||||
await this.ensureClient();
|
||||
const { error } = await this.supabase
|
||||
.from(this.tableName)
|
||||
.delete()
|
||||
|
||||
@@ -12,8 +12,9 @@ const upsert = jest.fn().mockResolvedValue(undefined);
|
||||
const query = jest.fn().mockResolvedValue({ matches: [] });
|
||||
const fetch = jest.fn().mockResolvedValue({ records: {} });
|
||||
const deleteOne = jest.fn().mockResolvedValue(undefined);
|
||||
const deleteAll = jest.fn().mockResolvedValue(undefined);
|
||||
|
||||
const nsHandle = { upsert, query, fetch, deleteOne };
|
||||
const nsHandle = { upsert, query, fetch, deleteOne, deleteAll };
|
||||
const namespace = jest.fn().mockReturnValue(nsHandle);
|
||||
|
||||
const describeIndexStats = jest
|
||||
@@ -47,6 +48,7 @@ const __mocks__ = {
|
||||
query,
|
||||
fetch,
|
||||
deleteOne,
|
||||
deleteAll,
|
||||
namespace,
|
||||
describeIndexStats,
|
||||
index,
|
||||
@@ -94,6 +96,7 @@ beforeEach(() => {
|
||||
__mocks__.query.mockResolvedValue({ matches: [] });
|
||||
__mocks__.fetch.mockResolvedValue({ records: {} });
|
||||
__mocks__.deleteOne.mockResolvedValue(undefined);
|
||||
__mocks__.deleteAll.mockResolvedValue(undefined);
|
||||
__mocks__.describeIndexStats.mockResolvedValue({
|
||||
totalRecordCount: 0,
|
||||
namespaces: {},
|
||||
@@ -104,6 +107,7 @@ beforeEach(() => {
|
||||
query: __mocks__.query,
|
||||
fetch: __mocks__.fetch,
|
||||
deleteOne: __mocks__.deleteOne,
|
||||
deleteAll: __mocks__.deleteAll,
|
||||
};
|
||||
__mocks__.namespace.mockReturnValue(nsHandle);
|
||||
__mocks__.index.mockReturnValue({
|
||||
@@ -424,6 +428,14 @@ describe("deleteCol", () => {
|
||||
await db.initialize();
|
||||
expect(__mocks__.createIndex).toHaveBeenCalledTimes(2);
|
||||
});
|
||||
|
||||
it("clears only the namespace and never drops the shared index", async () => {
|
||||
const db = await initDb({ namespace: "tenant-a" });
|
||||
await db.deleteCol();
|
||||
expect(__mocks__.namespace).toHaveBeenCalledWith("tenant-a");
|
||||
expect(__mocks__.deleteAll).toHaveBeenCalled();
|
||||
expect(__mocks__.deleteIndex).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
||||
describe("list", () => {
|
||||
|
||||
@@ -21,6 +21,11 @@ export interface EmbeddingConfig {
|
||||
modelProperties?: Record<string, any>;
|
||||
// HuggingFace TEI / OpenAI-compatible inference endpoint base URL.
|
||||
huggingfaceBaseUrl?: string;
|
||||
// AWS Bedrock. Omit the credential fields to use the AWS default chain.
|
||||
awsRegion?: string;
|
||||
awsAccessKeyId?: string;
|
||||
awsSecretAccessKey?: string;
|
||||
awsSessionToken?: string;
|
||||
}
|
||||
|
||||
export interface VertexAIConfig extends EmbeddingConfig {
|
||||
@@ -198,6 +203,10 @@ export const MemoryConfigSchema = z.object({
|
||||
memoryAddEmbeddingType: z.string().optional(),
|
||||
memoryUpdateEmbeddingType: z.string().optional(),
|
||||
memorySearchEmbeddingType: z.string().optional(),
|
||||
awsRegion: z.string().optional(),
|
||||
awsAccessKeyId: z.string().optional(),
|
||||
awsSecretAccessKey: z.string().optional(),
|
||||
awsSessionToken: z.string().optional(),
|
||||
}),
|
||||
}),
|
||||
vectorStore: z.object({
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import { OpenAIEmbedder } from "../embeddings/openai";
|
||||
import { AWSBedrockEmbedder } from "../embeddings/aws_bedrock";
|
||||
import { OllamaEmbedder } from "../embeddings/ollama";
|
||||
import { LMStudioEmbedder } from "../embeddings/lmstudio";
|
||||
import { TogetherEmbedder } from "../embeddings/together";
|
||||
@@ -76,6 +77,8 @@ export class EmbedderFactory {
|
||||
switch (provider.toLowerCase()) {
|
||||
case "openai":
|
||||
return new OpenAIEmbedder(config);
|
||||
case "aws_bedrock":
|
||||
return new AWSBedrockEmbedder(config);
|
||||
case "ollama":
|
||||
return new OllamaEmbedder(config);
|
||||
case "lmstudio":
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
export async function loadPeer(
|
||||
pkg: string,
|
||||
label: string,
|
||||
load: () => Promise<any>,
|
||||
): Promise<any> {
|
||||
try {
|
||||
return await load();
|
||||
} catch {
|
||||
throw new Error(
|
||||
`The '${pkg}' package is required to use the ${label}. Install it with: npm install ${pkg}`,
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -1,7 +1,6 @@
|
||||
import {
|
||||
import type {
|
||||
SearchClient,
|
||||
SearchIndexClient,
|
||||
AzureKeyCredential,
|
||||
SearchIndex,
|
||||
SearchField,
|
||||
SearchFieldDataType,
|
||||
@@ -13,9 +12,9 @@ import {
|
||||
BinaryQuantizationCompression,
|
||||
VectorizedQuery,
|
||||
} from "@azure/search-documents";
|
||||
import { DefaultAzureCredential } from "@azure/identity";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
/**
|
||||
* Configuration interface for Azure AI Search vector store
|
||||
@@ -71,8 +70,8 @@ interface AzureAISearchConfig extends VectorStoreConfig {
|
||||
* Supports vector search with hybrid search, compression, and filtering
|
||||
*/
|
||||
export class AzureAISearch implements VectorStore {
|
||||
private searchClient: SearchClient<any>;
|
||||
private indexClient: SearchIndexClient;
|
||||
private searchClient!: SearchClient<any>;
|
||||
private indexClient!: SearchIndexClient;
|
||||
private readonly serviceName: string;
|
||||
private readonly indexName: string;
|
||||
private readonly embeddingModelDims: number;
|
||||
@@ -93,25 +92,42 @@ export class AzureAISearch implements VectorStore {
|
||||
this.vectorFilterMode = config.vectorFilterMode || "preFilter";
|
||||
this.apiKey = config.apiKey;
|
||||
|
||||
// Initialize the index
|
||||
this.initialize().catch(console.error);
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.searchClient) return;
|
||||
const searchSdk = await loadPeer(
|
||||
"@azure/search-documents",
|
||||
"Azure AI Search vector store",
|
||||
() => import("@azure/search-documents"),
|
||||
);
|
||||
|
||||
const serviceEndpoint = `https://${this.serviceName}.search.windows.net`;
|
||||
|
||||
// Determine authentication: API key or DefaultAzureCredential
|
||||
const credential =
|
||||
this.apiKey && this.apiKey !== "" && this.apiKey !== "your-api-key"
|
||||
? new AzureKeyCredential(this.apiKey)
|
||||
: new DefaultAzureCredential();
|
||||
let credential: any;
|
||||
if (this.apiKey && this.apiKey !== "" && this.apiKey !== "your-api-key") {
|
||||
credential = new searchSdk.AzureKeyCredential(this.apiKey);
|
||||
} else {
|
||||
const identitySdk = await loadPeer(
|
||||
"@azure/identity",
|
||||
"Azure AI Search without an apiKey",
|
||||
() => import("@azure/identity"),
|
||||
);
|
||||
credential = new identitySdk.DefaultAzureCredential();
|
||||
}
|
||||
|
||||
// Initialize clients
|
||||
this.searchClient = new SearchClient(
|
||||
this.searchClient = new searchSdk.SearchClient(
|
||||
serviceEndpoint,
|
||||
this.indexName,
|
||||
credential,
|
||||
);
|
||||
|
||||
this.indexClient = new SearchIndexClient(serviceEndpoint, credential);
|
||||
|
||||
// Initialize the index
|
||||
this.initialize().catch(console.error);
|
||||
this.indexClient = new searchSdk.SearchIndexClient(
|
||||
serviceEndpoint,
|
||||
credential,
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -125,6 +141,7 @@ export class AzureAISearch implements VectorStore {
|
||||
}
|
||||
|
||||
private async _doInitialize(): Promise<void> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
const collections = await this.listCols();
|
||||
if (!collections.includes(this.indexName)) {
|
||||
@@ -262,6 +279,7 @@ export class AzureAISearch implements VectorStore {
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
console.log(
|
||||
`Inserting ${vectors.length} vectors into index ${this.indexName}`,
|
||||
);
|
||||
@@ -334,6 +352,7 @@ export class AzureAISearch implements VectorStore {
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[] | null> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const filterExpression = filters
|
||||
? this.buildFilterExpression(filters)
|
||||
@@ -373,6 +392,7 @@ export class AzureAISearch implements VectorStore {
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
await this.initialize();
|
||||
const filterExpression = filters
|
||||
? this.buildFilterExpression(filters)
|
||||
: undefined;
|
||||
@@ -429,6 +449,7 @@ export class AzureAISearch implements VectorStore {
|
||||
* Delete a vector by ID
|
||||
*/
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
const response = await this.searchClient.deleteDocuments([
|
||||
{ id: vectorId },
|
||||
]);
|
||||
@@ -454,6 +475,7 @@ export class AzureAISearch implements VectorStore {
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
const document: Record<string, any> = { id: vectorId };
|
||||
|
||||
if (vector) {
|
||||
@@ -486,6 +508,7 @@ export class AzureAISearch implements VectorStore {
|
||||
* Retrieve a vector by ID
|
||||
*/
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const result = await this.searchClient.getDocument(vectorId);
|
||||
const payloadStr = result.payload as string;
|
||||
@@ -521,6 +544,7 @@ export class AzureAISearch implements VectorStore {
|
||||
* Delete the index
|
||||
*/
|
||||
async deleteCol(): Promise<void> {
|
||||
await this.initialize();
|
||||
await this.indexClient.deleteIndex(this.indexName);
|
||||
}
|
||||
|
||||
@@ -542,6 +566,7 @@ export class AzureAISearch implements VectorStore {
|
||||
filters?: SearchFilters,
|
||||
topK: number = 100,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
await this.initialize();
|
||||
const filterExpression = filters
|
||||
? this.buildFilterExpression(filters)
|
||||
: undefined;
|
||||
@@ -586,6 +611,7 @@ export class AzureAISearch implements VectorStore {
|
||||
* Required by VectorStore interface
|
||||
*/
|
||||
async getUserId(): Promise<string> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Check if memory_migrations index exists
|
||||
const collections = await this.listCols();
|
||||
@@ -648,6 +674,7 @@ export class AzureAISearch implements VectorStore {
|
||||
* Required by VectorStore interface
|
||||
*/
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Get existing point ID or generate new one
|
||||
const searchResults = await this.searchClient.search("*", {
|
||||
@@ -677,6 +704,7 @@ export class AzureAISearch implements VectorStore {
|
||||
* Reset the index by deleting and recreating it
|
||||
*/
|
||||
async reset(): Promise<void> {
|
||||
await this.initialize();
|
||||
console.log(`Resetting index ${this.indexName}...`);
|
||||
|
||||
try {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { Pool, RowDataPacket } from "mysql2/promise";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
const SAFE_IDENTIFIER_RE = /^[a-zA-Z_][a-zA-Z0-9_]{0,127}$/;
|
||||
|
||||
@@ -94,14 +95,11 @@ export class AzureMySQLDB implements VectorStore {
|
||||
|
||||
// Loaded dynamically: mysql2 is an optional peer dependency, so a static value import
|
||||
// would break `import { Memory } from "mem0ai/oss"` for everyone else.
|
||||
let createPool: typeof import("mysql2/promise").createPool;
|
||||
try {
|
||||
({ createPool } = await import("mysql2/promise"));
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The Azure MySQL vector store requires the 'mysql2' package. Install it with: npm install mysql2",
|
||||
);
|
||||
}
|
||||
const { createPool }: typeof import("mysql2/promise") = await loadPeer(
|
||||
"mysql2",
|
||||
"Azure MySQL vector store",
|
||||
() => import("mysql2/promise"),
|
||||
);
|
||||
|
||||
this.pool = createPool({
|
||||
host: this.config.host,
|
||||
|
||||
@@ -12,6 +12,7 @@ import type {
|
||||
} from "@mochow/mochow-sdk-node";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
type MochowSdk = typeof import("@mochow/mochow-sdk-node");
|
||||
|
||||
@@ -156,14 +157,11 @@ export class BaiduDB implements VectorStore {
|
||||
// value import would break `import { Memory } from "mem0ai/oss"` for everyone else.
|
||||
private async loadSdk(): Promise<MochowSdk> {
|
||||
if (!this.sdk) {
|
||||
let module: MochowSdk & { default?: MochowSdk };
|
||||
try {
|
||||
module = await import("@mochow/mochow-sdk-node");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The Baidu vector store requires the '@mochow/mochow-sdk-node' package. Install it with: npm install @mochow/mochow-sdk-node",
|
||||
);
|
||||
}
|
||||
const module: MochowSdk & { default?: MochowSdk } = await loadPeer(
|
||||
"@mochow/mochow-sdk-node",
|
||||
"Baidu vector store",
|
||||
() => import("@mochow/mochow-sdk-node"),
|
||||
);
|
||||
this.sdk = module.default ?? module;
|
||||
}
|
||||
return this.sdk;
|
||||
@@ -455,7 +453,13 @@ export class BaiduDB implements VectorStore {
|
||||
return (response.rows ?? []).map((result) => ({
|
||||
id: String(result.row.id),
|
||||
payload: resultPayload(result.row),
|
||||
score: result.score,
|
||||
// L2 is a distance (lower = closer); the VectorStore contract wants higher = better.
|
||||
score:
|
||||
this.metricType === "L2"
|
||||
? result.score != null
|
||||
? 1 / (1 + result.score)
|
||||
: undefined
|
||||
: result.score,
|
||||
}));
|
||||
}
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
const MIGRATION_ROW_ID = "mem0-user";
|
||||
const SAFE_IDENTIFIER_RE = /^[A-Za-z_][A-Za-z0-9_]{0,127}$/;
|
||||
@@ -377,14 +378,11 @@ export class CassandraDB implements VectorStore {
|
||||
// Loaded dynamically: cassandra-driver is an optional peer dependency, so a static
|
||||
// value import would break `import { Memory } from "mem0ai/oss"` for everyone else.
|
||||
private async loadDriver(): Promise<any> {
|
||||
let sdk: any;
|
||||
try {
|
||||
sdk = await import("cassandra-driver");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The 'cassandra-driver' package is required to use the Cassandra vector store. Install it with: npm install cassandra-driver",
|
||||
);
|
||||
}
|
||||
const sdk = await loadPeer(
|
||||
"cassandra-driver",
|
||||
"Cassandra vector store",
|
||||
() => import("cassandra-driver"),
|
||||
);
|
||||
return sdk.default ?? sdk;
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { ChromaClient, CloudClient } from "chromadb";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface ChromaConfig extends VectorStoreConfig {
|
||||
/** Pre-configured ChromaDB client instance. */
|
||||
@@ -65,14 +66,11 @@ export class ChromaDB implements VectorStore {
|
||||
return config.client;
|
||||
}
|
||||
|
||||
let sdk: any;
|
||||
try {
|
||||
sdk = await import("chromadb");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The 'chromadb' package is required to use the Chroma vector store. Install it with: npm install chromadb",
|
||||
);
|
||||
}
|
||||
const sdk = await loadPeer(
|
||||
"chromadb",
|
||||
"Chroma vector store",
|
||||
() => import("chromadb"),
|
||||
);
|
||||
|
||||
if (config.apiKey && config.tenant) {
|
||||
return new sdk.CloudClient({
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { Client } from "@elastic/elasticsearch";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface ElasticsearchConfig extends VectorStoreConfig {
|
||||
/** Pre-configured Elasticsearch client instance (typed as `any` to keep the
|
||||
@@ -100,14 +101,11 @@ export class ElasticsearchDB implements VectorStore {
|
||||
params.headers = config.headers;
|
||||
}
|
||||
|
||||
let sdk: any;
|
||||
try {
|
||||
sdk = await import("@elastic/elasticsearch");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The '@elastic/elasticsearch' package is required to use the Elasticsearch vector store. Install it with: npm install @elastic/elasticsearch",
|
||||
);
|
||||
}
|
||||
const sdk = await loadPeer(
|
||||
"@elastic/elasticsearch",
|
||||
"Elasticsearch vector store",
|
||||
() => import("@elastic/elasticsearch"),
|
||||
);
|
||||
|
||||
this.client = new sdk.Client(params);
|
||||
}
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import { VectorStore as LangchainVectorStoreInterface } from "@langchain/core/vectorstores";
|
||||
import { Document } from "@langchain/core/documents";
|
||||
import type { VectorStore as LangchainVectorStoreInterface } from "@langchain/core/vectorstores";
|
||||
import { VectorStore } from "./base"; // mem0's VectorStore interface
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
|
||||
@@ -77,6 +76,7 @@ export class LangchainVectorStore implements VectorStore {
|
||||
}
|
||||
|
||||
// Convert payloads to Langchain Document metadata format
|
||||
const { Document } = await import("@langchain/core/documents");
|
||||
const documents = payloads.map((payload, i) => {
|
||||
// Provide empty pageContent, store mem0 id and other data in metadata
|
||||
return new Document({
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { MongoClient, Collection, Db } from "mongodb";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
export interface MongoDBConfig extends VectorStoreConfig {
|
||||
url?: string;
|
||||
@@ -42,14 +43,11 @@ export class MongoDB implements VectorStore {
|
||||
if (config.client) {
|
||||
this.client = config.client;
|
||||
} else {
|
||||
let sdk: any;
|
||||
try {
|
||||
sdk = await import("mongodb");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The 'mongodb' package is required to use the MongoDB vector store. Install it with: npm install mongodb",
|
||||
);
|
||||
}
|
||||
const sdk = await loadPeer(
|
||||
"mongodb",
|
||||
"MongoDB vector store",
|
||||
() => import("mongodb"),
|
||||
);
|
||||
const url = config.url || "mongodb://localhost:27017";
|
||||
this.client = new sdk.MongoClient(url, { appName: "Mem0" });
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { Client } from "@opensearch-project/opensearch";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
type OpenSearchAuth =
|
||||
| {
|
||||
@@ -98,14 +99,11 @@ export class OpenSearchDB implements VectorStore {
|
||||
? { username: config.user, password: config.password }
|
||||
: undefined);
|
||||
|
||||
let sdk: any;
|
||||
try {
|
||||
sdk = await import("@opensearch-project/opensearch");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The '@opensearch-project/opensearch' package is required to use the OpenSearch vector store. Install it with: npm install @opensearch-project/opensearch",
|
||||
);
|
||||
}
|
||||
const sdk = await loadPeer(
|
||||
"@opensearch-project/opensearch",
|
||||
"OpenSearch vector store",
|
||||
() => import("@opensearch-project/opensearch"),
|
||||
);
|
||||
|
||||
this.client = new sdk.Client({
|
||||
node: `${useSSL ? "https" : "http"}://${host}:${port}`,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { Pinecone, Index } from "@pinecone-database/pinecone";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
const MIGRATIONS_NAMESPACE = "__mem0_migrations__";
|
||||
const MIGRATIONS_RECORD_ID = "mem0-user-id";
|
||||
@@ -82,14 +83,11 @@ export class PineconeDB implements VectorStore {
|
||||
this.client = config.client;
|
||||
} else {
|
||||
const apiKey = config.apiKey || process.env.PINECONE_API_KEY;
|
||||
let sdk: any;
|
||||
try {
|
||||
sdk = await import("@pinecone-database/pinecone");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The '@pinecone-database/pinecone' package is required to use the Pinecone vector store. Install it with: npm install @pinecone-database/pinecone",
|
||||
);
|
||||
}
|
||||
const sdk = await loadPeer(
|
||||
"@pinecone-database/pinecone",
|
||||
"Pinecone vector store",
|
||||
() => import("@pinecone-database/pinecone"),
|
||||
);
|
||||
this.client = new sdk.Pinecone({ apiKey });
|
||||
}
|
||||
}
|
||||
@@ -318,6 +316,11 @@ export class PineconeDB implements VectorStore {
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
if (this.namespace) {
|
||||
await this.initialize();
|
||||
await this.namespacedIndex().deleteAll();
|
||||
return;
|
||||
}
|
||||
if (this._initPromise) {
|
||||
await this._initPromise.catch(() => {});
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { QdrantClient } from "@qdrant/js-client-rest";
|
||||
import type { QdrantClient } from "@qdrant/js-client-rest";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
import * as fs from "fs";
|
||||
|
||||
interface QdrantConfig extends VectorStoreConfig {
|
||||
@@ -55,53 +56,64 @@ const KEY_MAP: Record<string, string> = {
|
||||
};
|
||||
|
||||
export class Qdrant implements VectorStore {
|
||||
private client: QdrantClient;
|
||||
private client!: QdrantClient;
|
||||
private readonly config: QdrantConfig;
|
||||
private readonly collectionName: string;
|
||||
private dimension: number;
|
||||
private _initPromise?: Promise<void>;
|
||||
|
||||
constructor(config: QdrantConfig) {
|
||||
if (config.client) {
|
||||
this.client = config.client;
|
||||
} else {
|
||||
const params: Record<string, any> = {};
|
||||
if (config.apiKey) {
|
||||
params.apiKey = config.apiKey;
|
||||
}
|
||||
if (config.url) {
|
||||
params.url = config.url;
|
||||
// Workaround for qdrant/qdrant-js#59: explicitly pass port to avoid "Illegal host" error
|
||||
try {
|
||||
const parsedUrl = new URL(config.url);
|
||||
params.port = parsedUrl.port ? parseInt(parsedUrl.port, 10) : 6333;
|
||||
} catch (_) {
|
||||
params.port = 6333;
|
||||
}
|
||||
}
|
||||
if (config.host && config.port) {
|
||||
params.host = config.host;
|
||||
params.port = config.port;
|
||||
}
|
||||
if (!Object.keys(params).length) {
|
||||
params.path = config.path;
|
||||
if (!config.onDisk && config.path) {
|
||||
if (
|
||||
fs.existsSync(config.path) &&
|
||||
fs.statSync(config.path).isDirectory()
|
||||
) {
|
||||
fs.rmSync(config.path, { recursive: true });
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
this.client = new QdrantClient(params);
|
||||
}
|
||||
|
||||
this.config = config;
|
||||
this.collectionName = config.collectionName;
|
||||
this.dimension = config.dimension || 1536; // Default OpenAI dimension
|
||||
this.initialize().catch(console.error);
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const config = this.config;
|
||||
if (config.client) {
|
||||
this.client = config.client;
|
||||
return;
|
||||
}
|
||||
const params: Record<string, any> = {};
|
||||
if (config.apiKey) {
|
||||
params.apiKey = config.apiKey;
|
||||
}
|
||||
if (config.url) {
|
||||
params.url = config.url;
|
||||
// Workaround for qdrant/qdrant-js#59: explicitly pass port to avoid "Illegal host" error
|
||||
try {
|
||||
const parsedUrl = new URL(config.url);
|
||||
params.port = parsedUrl.port ? parseInt(parsedUrl.port, 10) : 6333;
|
||||
} catch (_) {
|
||||
params.port = 6333;
|
||||
}
|
||||
}
|
||||
if (config.host && config.port) {
|
||||
params.host = config.host;
|
||||
params.port = config.port;
|
||||
}
|
||||
if (!Object.keys(params).length) {
|
||||
params.path = config.path;
|
||||
if (!config.onDisk && config.path) {
|
||||
if (
|
||||
fs.existsSync(config.path) &&
|
||||
fs.statSync(config.path).isDirectory()
|
||||
) {
|
||||
fs.rmSync(config.path, { recursive: true });
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const sdk = await loadPeer(
|
||||
"@qdrant/js-client-rest",
|
||||
"Qdrant vector store",
|
||||
() => import("@qdrant/js-client-rest"),
|
||||
);
|
||||
this.client = new sdk.QdrantClient(params);
|
||||
}
|
||||
|
||||
/**
|
||||
* Build a single field condition from a key-value filter pair.
|
||||
* Supports enhanced filter syntax with comparison operators.
|
||||
@@ -270,6 +282,7 @@ export class Qdrant implements VectorStore {
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
const points = vectors.map((vector, idx) => ({
|
||||
id: ids[idx],
|
||||
vector: vector,
|
||||
@@ -290,6 +303,7 @@ export class Qdrant implements VectorStore {
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
await this.initialize();
|
||||
const queryFilter = this.createFilter(filters);
|
||||
const results = await this.client.search(this.collectionName, {
|
||||
vector: query,
|
||||
@@ -305,6 +319,7 @@ export class Qdrant implements VectorStore {
|
||||
}
|
||||
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
await this.initialize();
|
||||
const results = await this.client.retrieve(this.collectionName, {
|
||||
ids: [vectorId],
|
||||
with_payload: true,
|
||||
@@ -323,6 +338,7 @@ export class Qdrant implements VectorStore {
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
const point = {
|
||||
id: vectorId,
|
||||
vector: vector,
|
||||
@@ -335,12 +351,14 @@ export class Qdrant implements VectorStore {
|
||||
}
|
||||
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
await this.client.delete(this.collectionName, {
|
||||
points: [vectorId],
|
||||
});
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
await this.initialize();
|
||||
await this.client.deleteCollection(this.collectionName);
|
||||
}
|
||||
|
||||
@@ -348,6 +366,7 @@ export class Qdrant implements VectorStore {
|
||||
filters?: SearchFilters,
|
||||
topK: number = 100,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
await this.initialize();
|
||||
const scrollRequest = {
|
||||
limit: topK,
|
||||
filter: this.createFilter(filters),
|
||||
@@ -380,6 +399,7 @@ export class Qdrant implements VectorStore {
|
||||
}
|
||||
|
||||
async getUserId(): Promise<string> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Ensure collection exists (idempotent — handles race conditions)
|
||||
await this.ensureCollection("memory_migrations", 1);
|
||||
@@ -417,6 +437,7 @@ export class Qdrant implements VectorStore {
|
||||
}
|
||||
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Get existing point ID
|
||||
const result = await this.client.scroll("memory_migrations", {
|
||||
@@ -496,6 +517,7 @@ export class Qdrant implements VectorStore {
|
||||
|
||||
private async _doInitialize(): Promise<void> {
|
||||
try {
|
||||
await this.ensureClient();
|
||||
await this.ensureCollection(this.collectionName, this.dimension);
|
||||
await this.ensureCollection("memory_migrations", 1);
|
||||
} catch (error) {
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import { createClient } from "redis";
|
||||
import type {
|
||||
RedisClientType,
|
||||
RedisDefaultModules,
|
||||
@@ -8,6 +7,7 @@ import type {
|
||||
} from "redis";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
/**
|
||||
* Escape RediSearch TAG filter special characters. Any punctuation in the
|
||||
@@ -146,15 +146,21 @@ function toCamelCase(obj: Record<string, any>): Record<string, any> {
|
||||
}
|
||||
|
||||
export class RedisDB implements VectorStore {
|
||||
private client: RedisClientType<
|
||||
private client!: RedisClientType<
|
||||
RedisDefaultModules & RedisModules & RedisFunctions & RedisScripts
|
||||
>;
|
||||
private readonly redisUrl: string;
|
||||
private readonly username?: string;
|
||||
private readonly password?: string;
|
||||
private readonly indexName: string;
|
||||
private readonly indexPrefix: string;
|
||||
private readonly schema: RedisSchema;
|
||||
private _initPromise?: Promise<void>;
|
||||
|
||||
constructor(config: RedisConfig) {
|
||||
this.redisUrl = config.redisUrl;
|
||||
this.username = config.username;
|
||||
this.password = config.password;
|
||||
this.indexName = config.collectionName;
|
||||
this.indexPrefix = `mem0:${config.collectionName}`;
|
||||
|
||||
@@ -177,12 +183,24 @@ export class RedisDB implements VectorStore {
|
||||
}),
|
||||
};
|
||||
|
||||
this.client = createClient({
|
||||
url: config.redisUrl,
|
||||
username: config.username,
|
||||
password: config.password,
|
||||
this.initialize().catch((err) => {
|
||||
console.error("Failed to initialize Redis:", err);
|
||||
});
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const sdk = await loadPeer(
|
||||
"redis",
|
||||
"Redis vector store",
|
||||
() => import("redis"),
|
||||
);
|
||||
this.client = sdk.createClient({
|
||||
url: this.redisUrl,
|
||||
username: this.username,
|
||||
password: this.password,
|
||||
socket: {
|
||||
reconnectStrategy: (retries) => {
|
||||
reconnectStrategy: (retries: number) => {
|
||||
if (retries > 10) {
|
||||
console.error("Max reconnection attempts reached");
|
||||
return new Error("Max reconnection attempts reached");
|
||||
@@ -194,10 +212,6 @@ export class RedisDB implements VectorStore {
|
||||
|
||||
this.client.on("error", (err) => console.error("Redis Client Error:", err));
|
||||
this.client.on("connect", () => console.log("Redis Client Connected"));
|
||||
|
||||
this.initialize().catch((err) => {
|
||||
console.error("Failed to initialize Redis:", err);
|
||||
});
|
||||
}
|
||||
|
||||
private async createIndex(): Promise<void> {
|
||||
@@ -260,6 +274,7 @@ export class RedisDB implements VectorStore {
|
||||
}
|
||||
|
||||
private async _doInitialize(): Promise<void> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
await this.client.connect();
|
||||
console.log("Connected to Redis");
|
||||
@@ -331,6 +346,7 @@ export class RedisDB implements VectorStore {
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
const data = vectors.map((vector, idx) => {
|
||||
const payload = toSnakeCase(payloads[idx]);
|
||||
const id = ids[idx];
|
||||
@@ -389,6 +405,7 @@ export class RedisDB implements VectorStore {
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
await this.initialize();
|
||||
const snakeFilters = filters ? toSnakeCase(filters) : undefined;
|
||||
const filterExpr = snakeFilters
|
||||
? Object.entries(snakeFilters)
|
||||
@@ -456,6 +473,7 @@ export class RedisDB implements VectorStore {
|
||||
}
|
||||
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Check if the memory exists first
|
||||
const exists = await this.client.exists(
|
||||
@@ -562,6 +580,7 @@ export class RedisDB implements VectorStore {
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
const snakePayload = toSnakeCase(payload);
|
||||
const createdAt = snakePayload.created_at
|
||||
? new Date(snakePayload.created_at).getTime()
|
||||
@@ -601,6 +620,7 @@ export class RedisDB implements VectorStore {
|
||||
}
|
||||
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Check if memory exists first
|
||||
const key = `${this.indexPrefix}:${vectorId}`;
|
||||
@@ -626,6 +646,7 @@ export class RedisDB implements VectorStore {
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
await this.initialize();
|
||||
await this.client.ft.dropIndex(this.indexName);
|
||||
}
|
||||
|
||||
@@ -633,6 +654,7 @@ export class RedisDB implements VectorStore {
|
||||
filters?: SearchFilters,
|
||||
topK: number = 100,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
await this.initialize();
|
||||
const snakeFilters = filters ? toSnakeCase(filters) : undefined;
|
||||
const filterExpr = snakeFilters
|
||||
? Object.entries(snakeFilters)
|
||||
@@ -676,10 +698,11 @@ export class RedisDB implements VectorStore {
|
||||
}
|
||||
|
||||
async close(): Promise<void> {
|
||||
await this.client.quit();
|
||||
if (this.client) await this.client.quit();
|
||||
}
|
||||
|
||||
async getUserId(): Promise<string> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Check if the user ID exists in Redis
|
||||
const userId = await this.client.get("memory_migrations:1");
|
||||
@@ -702,6 +725,7 @@ export class RedisDB implements VectorStore {
|
||||
}
|
||||
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
await this.client.set("memory_migrations:1", userId);
|
||||
} catch (error) {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { createClient, SupabaseClient } from "@supabase/supabase-js";
|
||||
import type { SupabaseClient } from "@supabase/supabase-js";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface VectorData {
|
||||
id: string;
|
||||
@@ -82,14 +83,17 @@ $$;
|
||||
*/
|
||||
|
||||
export class SupabaseDB implements VectorStore {
|
||||
private client: SupabaseClient;
|
||||
private client!: SupabaseClient;
|
||||
private readonly supabaseUrl: string;
|
||||
private readonly supabaseKey: string;
|
||||
private readonly tableName: string;
|
||||
private readonly embeddingColumnName: string;
|
||||
private readonly metadataColumnName: string;
|
||||
private _initPromise?: Promise<void>;
|
||||
|
||||
constructor(config: SupabaseConfig) {
|
||||
this.client = createClient(config.supabaseUrl, config.supabaseKey);
|
||||
this.supabaseUrl = config.supabaseUrl;
|
||||
this.supabaseKey = config.supabaseKey;
|
||||
this.tableName = config.tableName;
|
||||
this.embeddingColumnName = config.embeddingColumnName || "embedding";
|
||||
this.metadataColumnName = config.metadataColumnName || "metadata";
|
||||
@@ -99,6 +103,16 @@ export class SupabaseDB implements VectorStore {
|
||||
});
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const sdk = await loadPeer(
|
||||
"@supabase/supabase-js",
|
||||
"Supabase vector store",
|
||||
() => import("@supabase/supabase-js"),
|
||||
);
|
||||
this.client = sdk.createClient(this.supabaseUrl, this.supabaseKey);
|
||||
}
|
||||
|
||||
async initialize(): Promise<void> {
|
||||
if (!this._initPromise) {
|
||||
this._initPromise = this._doInitialize();
|
||||
@@ -107,6 +121,7 @@ export class SupabaseDB implements VectorStore {
|
||||
}
|
||||
|
||||
private async _doInitialize(): Promise<void> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
// Verify table exists and vector operations work by attempting a test insert
|
||||
const testVector = Array(1536).fill(0);
|
||||
@@ -209,6 +224,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const data = vectors.map((vector, idx) => ({
|
||||
id: ids[idx],
|
||||
@@ -237,6 +253,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const rpcQuery: VectorQueryParams = {
|
||||
query_embedding: query,
|
||||
@@ -265,6 +282,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
}
|
||||
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const { data, error } = await this.client
|
||||
.from(this.tableName)
|
||||
@@ -290,6 +308,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const { error } = await this.client
|
||||
.from(this.tableName)
|
||||
@@ -310,6 +329,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
}
|
||||
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const { error } = await this.client
|
||||
.from(this.tableName)
|
||||
@@ -324,6 +344,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const { error } = await this.client
|
||||
.from(this.tableName)
|
||||
@@ -341,6 +362,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
filters?: SearchFilters,
|
||||
topK: number = 100,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
await this.initialize();
|
||||
try {
|
||||
let query = this.client
|
||||
.from(this.tableName)
|
||||
@@ -370,6 +392,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
}
|
||||
|
||||
async getUserId(): Promise<string> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// First check if the table exists
|
||||
const { data: tableExists } = await this.client
|
||||
@@ -421,6 +444,7 @@ See the SQL migration instructions in the code comments.`,
|
||||
}
|
||||
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const { error: deleteError } = await this.client
|
||||
.from("memory_migrations")
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface TurbopufferConfig extends VectorStoreConfig {
|
||||
apiKey?: string;
|
||||
@@ -48,14 +49,11 @@ export class TurbopufferDB implements VectorStore {
|
||||
}
|
||||
|
||||
private async createClient(): Promise<any> {
|
||||
let sdk: any;
|
||||
try {
|
||||
sdk = await import("@turbopuffer/turbopuffer");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The '@turbopuffer/turbopuffer' package is required to use the Turbopuffer vector store. Install it with: npm install @turbopuffer/turbopuffer",
|
||||
);
|
||||
}
|
||||
const sdk = await loadPeer(
|
||||
"@turbopuffer/turbopuffer",
|
||||
"Turbopuffer vector store",
|
||||
() => import("@turbopuffer/turbopuffer"),
|
||||
);
|
||||
|
||||
// @turbopuffer/turbopuffer ships `Turbopuffer` as both the default export
|
||||
// and a named export pointing at the same class. Use `.default` since
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import type { Index, QueryResult, Vector } from "@upstash/vector";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface UpstashVectorConfig extends VectorStoreConfig {
|
||||
collectionName: string;
|
||||
@@ -42,14 +43,11 @@ export class UpstashVector implements VectorStore {
|
||||
if (config.client) {
|
||||
this.client = config.client;
|
||||
} else {
|
||||
let sdk: any;
|
||||
try {
|
||||
sdk = await import("@upstash/vector");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The '@upstash/vector' package is required to use the Upstash Vector store. Install it with: npm install @upstash/vector",
|
||||
);
|
||||
}
|
||||
const sdk = await loadPeer(
|
||||
"@upstash/vector",
|
||||
"Upstash Vector store",
|
||||
() => import("@upstash/vector"),
|
||||
);
|
||||
this.client = new sdk.Index({
|
||||
url: config.url,
|
||||
token: config.token,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreResult } from "../types";
|
||||
import { ValkeyConfig } from "../types/valkey";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface ValkeyClient {
|
||||
call: (...args: (string | number | Buffer)[]) => Promise<unknown>;
|
||||
@@ -142,14 +143,8 @@ function formatTimestamp(timestamp: number, timezone: string = "UTC"): string {
|
||||
return `${yyyy}-${MM}-${dd}T${HH}:${mm}:${ss}${sign}${offHH}:${offMM}`;
|
||||
}
|
||||
|
||||
async function loadIovalkey(): Promise<typeof import("iovalkey")> {
|
||||
try {
|
||||
return await import("iovalkey");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"iovalkey is required for the Valkey vector store. Install it with: npm install iovalkey",
|
||||
);
|
||||
}
|
||||
function loadIovalkey(): Promise<typeof import("iovalkey")> {
|
||||
return loadPeer("iovalkey", "Valkey vector store", () => import("iovalkey"));
|
||||
}
|
||||
|
||||
export class ValkeyDB implements VectorStore {
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
import Cloudflare from "cloudflare";
|
||||
import type Cloudflare from "cloudflare";
|
||||
import type { Vectorize, VectorizeVector } from "@cloudflare/workers-types";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface VectorizeConfig extends VectorStoreConfig {
|
||||
apiKey?: string;
|
||||
@@ -17,24 +18,36 @@ interface CloudflareVector {
|
||||
|
||||
export class VectorizeDB implements VectorStore {
|
||||
private client: Cloudflare | null = null;
|
||||
private apiKey?: string;
|
||||
private dimensions: number;
|
||||
private indexName: string;
|
||||
private accountId: string;
|
||||
private _initPromise?: Promise<void>;
|
||||
|
||||
constructor(config: VectorizeConfig) {
|
||||
this.client = new Cloudflare({ apiToken: config.apiKey });
|
||||
this.apiKey = config.apiKey;
|
||||
this.dimensions = config.dimension || 1536;
|
||||
this.indexName = config.indexName;
|
||||
this.accountId = config.accountId;
|
||||
this.initialize().catch(console.error);
|
||||
}
|
||||
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this.client) return;
|
||||
const sdk = await loadPeer(
|
||||
"cloudflare",
|
||||
"Vectorize vector store",
|
||||
() => import("cloudflare"),
|
||||
);
|
||||
this.client = new sdk.default({ apiToken: this.apiKey });
|
||||
}
|
||||
|
||||
async insert(
|
||||
vectors: number[][],
|
||||
ids: string[],
|
||||
payloads: Record<string, any>[],
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const vectorObjects: CloudflareVector[] = vectors.map(
|
||||
(vector, index) => ({
|
||||
@@ -83,6 +96,7 @@ export class VectorizeDB implements VectorStore {
|
||||
topK: number = 5,
|
||||
filters?: SearchFilters,
|
||||
): Promise<VectorStoreResult[]> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const result = await this.client?.vectorize.indexes.query(
|
||||
this.indexName,
|
||||
@@ -111,6 +125,7 @@ export class VectorizeDB implements VectorStore {
|
||||
}
|
||||
|
||||
async get(vectorId: string): Promise<VectorStoreResult | null> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const result = (await this.client?.vectorize.indexes.getByIds(
|
||||
this.indexName,
|
||||
@@ -139,6 +154,7 @@ export class VectorizeDB implements VectorStore {
|
||||
vector: number[],
|
||||
payload: Record<string, any>,
|
||||
): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const data: VectorizeVector = {
|
||||
id: vectorId,
|
||||
@@ -173,6 +189,7 @@ export class VectorizeDB implements VectorStore {
|
||||
}
|
||||
|
||||
async delete(vectorId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
await this.client?.vectorize.indexes.deleteByIds(this.indexName, {
|
||||
account_id: this.accountId,
|
||||
@@ -187,6 +204,7 @@ export class VectorizeDB implements VectorStore {
|
||||
}
|
||||
|
||||
async deleteCol(): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
await this.client?.vectorize.indexes.delete(this.indexName, {
|
||||
account_id: this.accountId,
|
||||
@@ -203,6 +221,7 @@ export class VectorizeDB implements VectorStore {
|
||||
filters?: SearchFilters,
|
||||
topK: number = 20,
|
||||
): Promise<[VectorStoreResult[], number]> {
|
||||
await this.initialize();
|
||||
try {
|
||||
const result = await this.client?.vectorize.indexes.query(
|
||||
this.indexName,
|
||||
@@ -243,6 +262,7 @@ export class VectorizeDB implements VectorStore {
|
||||
}
|
||||
|
||||
async getUserId(): Promise<string> {
|
||||
await this.initialize();
|
||||
try {
|
||||
let found = false;
|
||||
for await (const index of this.client!.vectorize.indexes.list({
|
||||
@@ -309,6 +329,7 @@ export class VectorizeDB implements VectorStore {
|
||||
}
|
||||
|
||||
async setUserId(userId: string): Promise<void> {
|
||||
await this.initialize();
|
||||
try {
|
||||
// Get existing point ID
|
||||
const result: any = await this.client?.vectorize.indexes.query(
|
||||
@@ -355,6 +376,7 @@ export class VectorizeDB implements VectorStore {
|
||||
}
|
||||
|
||||
private async _doInitialize(): Promise<void> {
|
||||
await this.ensureClient();
|
||||
try {
|
||||
// Check if the index already exists
|
||||
let indexFound = false;
|
||||
|
||||
@@ -2,6 +2,7 @@ import type { WeaviateClient } from "weaviate-client";
|
||||
import { v4 as uuidv4 } from "uuid";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
import { loadPeer } from "../utils/load_peer";
|
||||
|
||||
interface WeaviateConfig extends VectorStoreConfig {
|
||||
/** Pre-configured Weaviate client instance (typed as `any` to keep the
|
||||
@@ -50,14 +51,11 @@ export class WeaviateDB implements VectorStore {
|
||||
private async ensureClient(): Promise<void> {
|
||||
if (this._client) return;
|
||||
|
||||
let sdk: any;
|
||||
try {
|
||||
sdk = await import("weaviate-client");
|
||||
} catch {
|
||||
throw new Error(
|
||||
"The 'weaviate-client' package is required to use the Weaviate vector store. Install it with: npm install weaviate-client",
|
||||
);
|
||||
}
|
||||
const sdk = await loadPeer(
|
||||
"weaviate-client",
|
||||
"Weaviate vector store",
|
||||
() => import("weaviate-client"),
|
||||
);
|
||||
this._sdk = sdk;
|
||||
|
||||
const { client, clusterUrl, apiKey, additionalHeaders } = this._config;
|
||||
|
||||
@@ -23,11 +23,16 @@ describe("AnthropicLLM (unit)", () => {
|
||||
|
||||
// Regression #5665: a configured baseURL must reach the Anthropic client so
|
||||
// proxy/gateway users are not silently bypassed (TS parity with #5626).
|
||||
it("forwards baseURL to the Anthropic client when set", () => {
|
||||
new AnthropicLLM({
|
||||
// The client is constructed lazily on first use, so drive generateResponse.
|
||||
it("forwards baseURL to the Anthropic client when set", async () => {
|
||||
mockCreate.mockResolvedValueOnce({
|
||||
content: [{ type: "text", text: "ok" }],
|
||||
});
|
||||
const llm = new AnthropicLLM({
|
||||
apiKey: "test-key",
|
||||
baseURL: "https://proxy.example/v1",
|
||||
});
|
||||
await llm.generateResponse([{ role: "user", content: "Hi" }]);
|
||||
|
||||
expect(mockConstructor).toHaveBeenCalledTimes(1);
|
||||
const ctorArgs = mockConstructor.mock.calls[0][0];
|
||||
@@ -37,8 +42,12 @@ describe("AnthropicLLM (unit)", () => {
|
||||
|
||||
// When no baseURL is configured the client must not receive a baseURL key
|
||||
// (so the SDK default endpoint is used).
|
||||
it("does NOT set baseURL when none is configured", () => {
|
||||
new AnthropicLLM({ apiKey: "test-key" });
|
||||
it("does NOT set baseURL when none is configured", async () => {
|
||||
mockCreate.mockResolvedValueOnce({
|
||||
content: [{ type: "text", text: "ok" }],
|
||||
});
|
||||
const llm = new AnthropicLLM({ apiKey: "test-key" });
|
||||
await llm.generateResponse([{ role: "user", content: "Hi" }]);
|
||||
|
||||
expect(mockConstructor).toHaveBeenCalledTimes(1);
|
||||
const ctorArgs = mockConstructor.mock.calls[0][0];
|
||||
|
||||
@@ -0,0 +1,517 @@
|
||||
import type { InvokeModelCommand } from "@aws-sdk/client-bedrock-runtime";
|
||||
import { AWSBedrockEmbedder } from "../src/embeddings/aws_bedrock";
|
||||
import type { Embedder } from "../src/embeddings/base";
|
||||
import { EmbedderFactory } from "../src/utils/factory";
|
||||
|
||||
/**
|
||||
* Only the network boundary is faked: `BedrockRuntimeClient.send` never leaves
|
||||
* the process. `InvokeModelCommand` stays the real class from the AWS SDK, so
|
||||
* every assertion below runs against the exact payload Bedrock would receive.
|
||||
*/
|
||||
const mockSend = jest.fn();
|
||||
const mockClientConfigs: any[] = [];
|
||||
const mockClientConstructor = jest
|
||||
.fn()
|
||||
.mockImplementation((config: unknown) => {
|
||||
mockClientConfigs.push(config);
|
||||
return { send: mockSend };
|
||||
});
|
||||
|
||||
jest.mock("@aws-sdk/client-bedrock-runtime", () => {
|
||||
const actual = jest.requireActual("@aws-sdk/client-bedrock-runtime");
|
||||
return {
|
||||
...actual,
|
||||
BedrockRuntimeClient: mockClientConstructor,
|
||||
};
|
||||
});
|
||||
|
||||
const encode = (payload: unknown) => ({
|
||||
body: new TextEncoder().encode(JSON.stringify(payload)),
|
||||
});
|
||||
|
||||
const titanReply = (embedding: number[]) =>
|
||||
encode({ embedding, inputTextTokenCount: embedding.length });
|
||||
|
||||
const cohereReply = (embeddings: number[][]) =>
|
||||
encode({ embeddings, id: "req-1", response_type: "embeddings_floats" });
|
||||
|
||||
const commandAt = (index: number): InvokeModelCommand =>
|
||||
mockSend.mock.calls[index][0];
|
||||
|
||||
const requestBodyAt = (index: number) =>
|
||||
JSON.parse(new TextDecoder().decode(commandAt(index).input.body));
|
||||
|
||||
describe("AWSBedrockEmbedder", () => {
|
||||
const savedEnv = { ...process.env };
|
||||
|
||||
beforeEach(() => {
|
||||
jest.clearAllMocks();
|
||||
mockClientConfigs.length = 0;
|
||||
delete process.env.AWS_REGION;
|
||||
});
|
||||
|
||||
afterAll(() => {
|
||||
process.env = savedEnv;
|
||||
});
|
||||
|
||||
describe("Titan models", () => {
|
||||
it("sends inputText and returns the embedding vector", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([0.1, 0.2, 0.3]));
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
const embedding = await embedder.embed("hello world");
|
||||
|
||||
expect(embedding).toEqual([0.1, 0.2, 0.3]);
|
||||
expect(mockSend).toHaveBeenCalledTimes(1);
|
||||
expect(commandAt(0).input.modelId).toBe("amazon.titan-embed-text-v1");
|
||||
expect(commandAt(0).input.contentType).toBe("application/json");
|
||||
expect(commandAt(0).input.accept).toBe("application/json");
|
||||
expect(requestBodyAt(0)).toEqual({ inputText: "hello world" });
|
||||
});
|
||||
|
||||
it("forwards dimensions to Titan V2 when embeddingDims is set", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([0.1, 0.2]));
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "amazon.titan-embed-text-v2:0",
|
||||
embeddingDims: 512,
|
||||
});
|
||||
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(requestBodyAt(0)).toEqual({ inputText: "hello", dimensions: 512 });
|
||||
});
|
||||
|
||||
it("omits dimensions on Titan V1, which rejects the field", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([0.1, 0.2]));
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "amazon.titan-embed-text-v1",
|
||||
embeddingDims: 512,
|
||||
});
|
||||
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(requestBodyAt(0)).toEqual({ inputText: "hello" });
|
||||
});
|
||||
|
||||
// F6: the old guard was `model.includes("v2")`, which would also match
|
||||
// any future/other Titan model whose id merely contains "v2" somewhere
|
||||
// (e.g. an image model), wrongly sending `dimensions` to a model that may
|
||||
// reject it. Only Titan Text Embeddings V2 should get the field.
|
||||
it("does not forward dimensions to a non-Titan-V2 model whose name merely contains v2", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([0.1, 0.2]));
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "amazon.titan-embed-image-v2:0",
|
||||
embeddingDims: 512,
|
||||
});
|
||||
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(requestBodyAt(0)).toEqual({ inputText: "hello" });
|
||||
});
|
||||
|
||||
it("embedBatch issues one request per text and preserves order", async () => {
|
||||
mockSend
|
||||
.mockResolvedValueOnce(titanReply([1, 1]))
|
||||
.mockResolvedValueOnce(titanReply([2, 2]));
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
const embeddings = await embedder.embedBatch(["first", "second"]);
|
||||
|
||||
expect(embeddings).toEqual([
|
||||
[1, 1],
|
||||
[2, 2],
|
||||
]);
|
||||
expect(mockSend).toHaveBeenCalledTimes(2);
|
||||
expect(requestBodyAt(0)).toEqual({ inputText: "first" });
|
||||
expect(requestBodyAt(1)).toEqual({ inputText: "second" });
|
||||
});
|
||||
});
|
||||
|
||||
describe("Cohere models", () => {
|
||||
it("sends texts with a search_document input type", async () => {
|
||||
mockSend.mockResolvedValueOnce(cohereReply([[0.4, 0.5]]));
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "cohere.embed-english-v3",
|
||||
});
|
||||
|
||||
const embedding = await embedder.embed("hello");
|
||||
|
||||
expect(embedding).toEqual([0.4, 0.5]);
|
||||
expect(requestBodyAt(0)).toEqual({
|
||||
texts: ["hello"],
|
||||
input_type: "search_document",
|
||||
});
|
||||
});
|
||||
|
||||
it("embedBatch sends every text in a single request", async () => {
|
||||
mockSend.mockResolvedValueOnce(cohereReply([[1], [2], [3]]));
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "cohere.embed-multilingual-v3",
|
||||
});
|
||||
|
||||
const embeddings = await embedder.embedBatch(["a", "b", "c"]);
|
||||
|
||||
expect(embeddings).toEqual([[1], [2], [3]]);
|
||||
expect(mockSend).toHaveBeenCalledTimes(1);
|
||||
expect(requestBodyAt(0).texts).toEqual(["a", "b", "c"]);
|
||||
});
|
||||
|
||||
it("embedBatch splits requests at Cohere's 96 text limit", async () => {
|
||||
const texts = Array.from({ length: 100 }, (_, i) => `text-${i}`);
|
||||
mockSend
|
||||
.mockResolvedValueOnce(
|
||||
cohereReply(texts.slice(0, 96).map((_, i) => [i])),
|
||||
)
|
||||
.mockResolvedValueOnce(cohereReply(texts.slice(96).map((_, i) => [i])));
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "cohere.embed-english-v3",
|
||||
});
|
||||
|
||||
const embeddings = await embedder.embedBatch(texts);
|
||||
|
||||
expect(embeddings).toHaveLength(100);
|
||||
expect(mockSend).toHaveBeenCalledTimes(2);
|
||||
expect(requestBodyAt(0).texts).toHaveLength(96);
|
||||
expect(requestBodyAt(1).texts).toEqual([
|
||||
"text-96",
|
||||
"text-97",
|
||||
"text-98",
|
||||
"text-99",
|
||||
]);
|
||||
});
|
||||
});
|
||||
|
||||
describe("Cohere Embed v4", () => {
|
||||
// F5: only Embed v4 understands embedding_types / output_dimension, and
|
||||
// (when embedding_types is requested) replies with a nested
|
||||
// `{ embeddings: { float: [...] } }` shape instead of v3's flat array.
|
||||
it("requests embedding_types and output_dimension for v4 models", async () => {
|
||||
mockSend.mockResolvedValueOnce(
|
||||
encode({ embeddings: { float: [[0.1, 0.2, 0.3]] } }),
|
||||
);
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "cohere.embed-v4:0",
|
||||
embeddingDims: 512,
|
||||
});
|
||||
|
||||
const embedding = await embedder.embed("hello");
|
||||
|
||||
expect(embedding).toEqual([0.1, 0.2, 0.3]);
|
||||
expect(requestBodyAt(0)).toEqual({
|
||||
texts: ["hello"],
|
||||
input_type: "search_document",
|
||||
embedding_types: ["float"],
|
||||
output_dimension: 512,
|
||||
});
|
||||
});
|
||||
|
||||
it("omits output_dimension for v4 when embeddingDims is unset", async () => {
|
||||
mockSend.mockResolvedValueOnce(
|
||||
encode({ embeddings: { float: [[0.1, 0.2]] } }),
|
||||
);
|
||||
const embedder = new AWSBedrockEmbedder({ model: "cohere.embed-v4:0" });
|
||||
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(requestBodyAt(0)).toEqual({
|
||||
texts: ["hello"],
|
||||
input_type: "search_document",
|
||||
embedding_types: ["float"],
|
||||
});
|
||||
});
|
||||
|
||||
it("parses the nested embeddings.float response shape", async () => {
|
||||
mockSend.mockResolvedValueOnce(
|
||||
encode({
|
||||
embeddings: {
|
||||
float: [
|
||||
[1, 2],
|
||||
[3, 4],
|
||||
],
|
||||
},
|
||||
}),
|
||||
);
|
||||
const embedder = new AWSBedrockEmbedder({ model: "cohere.embed-v4:0" });
|
||||
|
||||
const embeddings = await embedder.embedBatch(["a", "b"]);
|
||||
|
||||
expect(embeddings).toEqual([
|
||||
[1, 2],
|
||||
[3, 4],
|
||||
]);
|
||||
});
|
||||
});
|
||||
|
||||
describe("client configuration", () => {
|
||||
it("defaults to the us-west-2 region", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([1]));
|
||||
await new AWSBedrockEmbedder({}).embed("hello");
|
||||
|
||||
expect(mockClientConfigs[0].region).toBe("us-west-2");
|
||||
});
|
||||
|
||||
it("prefers awsRegion over the AWS_REGION environment variable", async () => {
|
||||
process.env.AWS_REGION = "eu-central-1";
|
||||
mockSend.mockResolvedValueOnce(titanReply([1]));
|
||||
await new AWSBedrockEmbedder({ awsRegion: "ap-south-1" }).embed("hello");
|
||||
|
||||
expect(mockClientConfigs[0].region).toBe("ap-south-1");
|
||||
});
|
||||
|
||||
it("falls back to the AWS_REGION environment variable", async () => {
|
||||
process.env.AWS_REGION = "eu-central-1";
|
||||
mockSend.mockResolvedValueOnce(titanReply([1]));
|
||||
await new AWSBedrockEmbedder({}).embed("hello");
|
||||
|
||||
expect(mockClientConfigs[0].region).toBe("eu-central-1");
|
||||
});
|
||||
|
||||
it("passes explicitly configured credentials to the client", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([1]));
|
||||
await new AWSBedrockEmbedder({
|
||||
awsAccessKeyId: "AKIA_TEST",
|
||||
awsSecretAccessKey: "secret",
|
||||
awsSessionToken: "token",
|
||||
}).embed("hello");
|
||||
|
||||
expect(mockClientConfigs[0].credentials).toEqual({
|
||||
accessKeyId: "AKIA_TEST",
|
||||
secretAccessKey: "secret",
|
||||
sessionToken: "token",
|
||||
});
|
||||
});
|
||||
|
||||
it("leaves credentials unset so the AWS default credential chain applies", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([1]));
|
||||
await new AWSBedrockEmbedder({}).embed("hello");
|
||||
|
||||
expect(mockClientConfigs[0].credentials).toBeUndefined();
|
||||
});
|
||||
|
||||
it("rejects a half-configured credential pair", () => {
|
||||
expect(
|
||||
() => new AWSBedrockEmbedder({ awsAccessKeyId: "AKIA_TEST" }),
|
||||
).toThrow(/awsAccessKeyId and awsSecretAccessKey/);
|
||||
expect(
|
||||
() => new AWSBedrockEmbedder({ awsSecretAccessKey: "secret" }),
|
||||
).toThrow(/awsAccessKeyId and awsSecretAccessKey/);
|
||||
});
|
||||
|
||||
// Silently ignoring a lone session token would fall back to the ambient
|
||||
// credential chain, embedding under an identity the caller never chose.
|
||||
it("rejects a session token supplied without the key pair", () => {
|
||||
expect(
|
||||
() => new AWSBedrockEmbedder({ awsSessionToken: "token" }),
|
||||
).toThrow(/awsAccessKeyId and awsSecretAccessKey/);
|
||||
});
|
||||
});
|
||||
|
||||
describe("provider registration", () => {
|
||||
it("is constructed by EmbedderFactory for the aws_bedrock provider", () => {
|
||||
const embedder = EmbedderFactory.create("aws_bedrock", {
|
||||
model: "amazon.titan-embed-text-v2:0",
|
||||
});
|
||||
|
||||
expect(embedder).toBeInstanceOf(AWSBedrockEmbedder);
|
||||
});
|
||||
});
|
||||
|
||||
describe("error handling", () => {
|
||||
it("wraps Bedrock failures with the model id", async () => {
|
||||
mockSend.mockRejectedValueOnce(new Error("AccessDeniedException"));
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
"Error getting embedding from AWS Bedrock model amazon.titan-embed-text-v1: AccessDeniedException",
|
||||
);
|
||||
});
|
||||
|
||||
it("fails when the response carries no embedding", async () => {
|
||||
mockSend.mockResolvedValueOnce(encode({ inputTextTokenCount: 3 }));
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
/returned no embedding/,
|
||||
);
|
||||
});
|
||||
|
||||
// F7: `[]` is truthy, so `payload.embedding && [payload.embedding]` used to
|
||||
// turn `{"embedding": []}` into `[[]]` -- length 1, which satisfied the
|
||||
// length check for a single-input call and handed the caller an empty vector.
|
||||
it("rejects a zero-length Titan embedding as no embedding", async () => {
|
||||
mockSend.mockResolvedValueOnce(encode({ embedding: [] }));
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
/returned no embedding/,
|
||||
);
|
||||
});
|
||||
|
||||
it("returns an empty array for an empty batch without calling Bedrock", async () => {
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embedBatch([])).resolves.toEqual([]);
|
||||
expect(mockSend).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
||||
describe("dynamic SDK import", () => {
|
||||
// These tests replace the module registered for
|
||||
// @aws-sdk/client-bedrock-runtime for a single resolution. Restore the
|
||||
// working mock afterward so every other test in this file keeps getting
|
||||
// the mocked client instead of hitting module resolution for real.
|
||||
afterEach(() => {
|
||||
jest.resetModules();
|
||||
jest.doMock("@aws-sdk/client-bedrock-runtime", () => {
|
||||
const actual = jest.requireActual("@aws-sdk/client-bedrock-runtime");
|
||||
return { ...actual, BedrockRuntimeClient: mockClientConstructor };
|
||||
});
|
||||
});
|
||||
|
||||
// F1: loadSdk()'s catch used to rewrite *every* import failure into the
|
||||
// "package is required" hint, even when the package is installed but
|
||||
// failed to load for an unrelated reason. That discarded the real error.
|
||||
it("propagates a non-resolution import error unchanged", async () => {
|
||||
jest.resetModules();
|
||||
jest.doMock("@aws-sdk/client-bedrock-runtime", () => {
|
||||
const err: any = new Error("boom: unrelated crash while loading");
|
||||
err.code = "ERR_SOMETHING_ELSE";
|
||||
throw err;
|
||||
});
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
"boom: unrelated crash while loading",
|
||||
);
|
||||
});
|
||||
|
||||
// F1: a genuine resolution failure should still get the friendly install
|
||||
// hint, with the original error preserved as `cause` for debugging.
|
||||
it("gives an install hint for a genuine module-not-found error, preserving the cause", async () => {
|
||||
jest.resetModules();
|
||||
jest.doMock("@aws-sdk/client-bedrock-runtime", () => {
|
||||
const err: any = new Error(
|
||||
"Cannot find module '@aws-sdk/client-bedrock-runtime'",
|
||||
);
|
||||
err.code = "MODULE_NOT_FOUND";
|
||||
throw err;
|
||||
});
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
/npm install @aws-sdk\/client-bedrock-runtime/,
|
||||
);
|
||||
await expect(embedder.embed("hello")).rejects.toMatchObject({
|
||||
cause: expect.objectContaining({ code: "MODULE_NOT_FOUND" }),
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("client promise retry", () => {
|
||||
// F2: getClient() used to memoize the client promise before it resolved,
|
||||
// so a rejected construction (e.g. a transient credentials failure) was
|
||||
// cached forever -- every later embed() call on that instance would
|
||||
// reject immediately without ever retrying.
|
||||
it("retries client construction after a failure instead of caching the rejection", async () => {
|
||||
mockClientConstructor.mockImplementationOnce(() => {
|
||||
throw new Error("credentials not ready");
|
||||
});
|
||||
mockSend.mockResolvedValueOnce(titanReply([1, 2, 3]));
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
"credentials not ready",
|
||||
);
|
||||
await expect(embedder.embed("hello")).resolves.toEqual([1, 2, 3]);
|
||||
});
|
||||
});
|
||||
|
||||
describe("memoryAction -> Cohere input_type", () => {
|
||||
// F3: buildRequestBody() used to hardcode `input_type: "search_document"`
|
||||
// regardless of the caller's action, so `Memory.search()` (which calls
|
||||
// `embed(query, "search")`) embedded the query in document mode.
|
||||
//
|
||||
// Typed as `Embedder` (not `AWSBedrockEmbedder`) because that is how
|
||||
// memory/index.ts actually calls it: the interface already declares an
|
||||
// optional `memoryAction` second parameter, so a narrower concrete
|
||||
// `embed(text: string)` satisfies it structurally and tsc stays silent --
|
||||
// the bug is a silent behavioral one, not a compile error.
|
||||
it("sends search_query for a search action", async () => {
|
||||
mockSend.mockResolvedValueOnce(cohereReply([[0.1]]));
|
||||
const embedder: Embedder = new AWSBedrockEmbedder({
|
||||
model: "cohere.embed-english-v3",
|
||||
});
|
||||
|
||||
await embedder.embed("query text", "search");
|
||||
|
||||
expect(requestBodyAt(0)).toEqual({
|
||||
texts: ["query text"],
|
||||
input_type: "search_query",
|
||||
});
|
||||
});
|
||||
|
||||
it("sends search_document for add and update actions", async () => {
|
||||
mockSend
|
||||
.mockResolvedValueOnce(cohereReply([[0.1]]))
|
||||
.mockResolvedValueOnce(cohereReply([[0.2]]));
|
||||
const embedder: Embedder = new AWSBedrockEmbedder({
|
||||
model: "cohere.embed-english-v3",
|
||||
});
|
||||
|
||||
await embedder.embed("doc one", "add");
|
||||
await embedder.embed("doc two", "update");
|
||||
|
||||
expect(requestBodyAt(0).input_type).toBe("search_document");
|
||||
expect(requestBodyAt(1).input_type).toBe("search_document");
|
||||
});
|
||||
|
||||
// Titan has no input_type concept; buildRequestBody() must not add one
|
||||
// even when a memoryAction is explicitly passed through.
|
||||
it("Titan ignores memoryAction and never sends input_type", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([1, 2]));
|
||||
const embedder: Embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await embedder.embed("hello", "search");
|
||||
|
||||
expect(requestBodyAt(0)).toEqual({ inputText: "hello" });
|
||||
});
|
||||
});
|
||||
|
||||
describe("Titan embedBatch concurrency", () => {
|
||||
afterEach(() => {
|
||||
// This test sets a persistent mockImplementation (not a *Once), so
|
||||
// clear it explicitly -- jest.clearAllMocks() in the top beforeEach
|
||||
// clears call data but not implementations.
|
||||
mockSend.mockReset();
|
||||
});
|
||||
|
||||
// F4: embedBatch() used to Promise.all-fan-out one InvokeModel call per
|
||||
// text with no cap, so a large batch could open hundreds of concurrent
|
||||
// requests at once. TITAN_MAX_CONCURRENCY bounds this to a small pool
|
||||
// while still preserving output order.
|
||||
it("never runs more than TITAN_MAX_CONCURRENCY Titan requests at once, and preserves order", async () => {
|
||||
let active = 0;
|
||||
let peak = 0;
|
||||
mockSend.mockImplementation(async (command: InvokeModelCommand) => {
|
||||
active++;
|
||||
peak = Math.max(peak, active);
|
||||
const body = JSON.parse(
|
||||
new TextDecoder().decode(command.input.body as Uint8Array),
|
||||
);
|
||||
await new Promise((resolve) => setTimeout(resolve, 10));
|
||||
active--;
|
||||
const index = Number(body.inputText.split("-")[1]);
|
||||
return titanReply([index]);
|
||||
});
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
const texts = Array.from({ length: 10 }, (_, i) => `text-${i}`);
|
||||
|
||||
const embeddings = await embedder.embedBatch(texts);
|
||||
|
||||
expect(peak).toBeGreaterThan(1);
|
||||
expect(peak).toBeLessThanOrEqual(4);
|
||||
expect(embeddings).toEqual(texts.map((_, i) => [i]));
|
||||
expect(mockSend).toHaveBeenCalledTimes(10);
|
||||
});
|
||||
});
|
||||
});
|
||||
@@ -431,9 +431,12 @@ describe("BaiduDB reads", () => {
|
||||
],
|
||||
});
|
||||
|
||||
const results = await makeStore(client).search([1, 2, 3], 5, {
|
||||
userId: "alice",
|
||||
});
|
||||
// Non-L2 metrics already return a higher-is-better score and pass through untouched.
|
||||
const results = await makeStore(client, { metricType: "COSINE" }).search(
|
||||
[1, 2, 3],
|
||||
5,
|
||||
{ userId: "alice" },
|
||||
);
|
||||
expect(results).toEqual([{ id: "m1", payload: { data: "x" }, score: 0.8 }]);
|
||||
|
||||
const { request } = client.vectorSearch.mock.calls[0][0];
|
||||
@@ -445,6 +448,35 @@ describe("BaiduDB reads", () => {
|
||||
expect(request.config.params).toEqual({ ef: 200 });
|
||||
});
|
||||
|
||||
it("converts an L2 distance into a similarity score (higher = better)", async () => {
|
||||
// Mirrors the Python provider (#6435): 1 / (1 + distance), so closer scores higher.
|
||||
const client = fakeClient();
|
||||
client.vectorSearch.mockResolvedValue({
|
||||
...OK,
|
||||
rows: [
|
||||
{ row: { id: "near", metadata: {} }, score: 0.5 },
|
||||
{ row: { id: "far", metadata: {} }, score: 2.0 },
|
||||
],
|
||||
});
|
||||
|
||||
const results = await makeStore(client).search([1, 2, 3], 2);
|
||||
|
||||
expect(results[0].score).toBeCloseTo(1 / 1.5, 10);
|
||||
expect(results[1].score).toBeCloseTo(1 / 3.0, 10);
|
||||
});
|
||||
|
||||
it("leaves an unscored L2 row's score undefined instead of treating it as the closest match", async () => {
|
||||
const client = fakeClient();
|
||||
client.vectorSearch.mockResolvedValue({
|
||||
...OK,
|
||||
rows: [{ row: { id: "unscored", metadata: {} } }],
|
||||
});
|
||||
|
||||
const results = await makeStore(client).search([1, 2, 3], 1);
|
||||
|
||||
expect(results[0].score).toBeUndefined();
|
||||
});
|
||||
|
||||
it("omits the filter when no filters are supplied", async () => {
|
||||
const client = fakeClient();
|
||||
client.vectorSearch.mockResolvedValue({ ...OK, rows: [] });
|
||||
|
||||
@@ -207,6 +207,79 @@ describe("Memory - update()", () => {
|
||||
expect.objectContaining({ category: "hobbies", priority: "high" }),
|
||||
);
|
||||
});
|
||||
|
||||
// Regression: metadata passed to update() must never overwrite a memory's
|
||||
// identity fields (issues #6277 / #6278).
|
||||
test("metadata cannot overwrite identity fields (tenant isolation)", async () => {
|
||||
const tenantA = `tenant_a_${Date.now()}`;
|
||||
const tenantB = `tenant_b_${Date.now()}`;
|
||||
const addResult: SearchResult = await memory.add("Tenant A secret", {
|
||||
userId: tenantA,
|
||||
infer: false,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
|
||||
// Caller metadata attempts to re-scope the memory to tenant B.
|
||||
await memory.update(id, {
|
||||
text: "Updated text",
|
||||
metadata: { user_id: tenantB, agent_id: "attacker", run_id: "attacker" },
|
||||
});
|
||||
|
||||
// The memory must remain in tenant A's scope...
|
||||
const aList: SearchResult = await memory.getAll({
|
||||
filters: { user_id: tenantA },
|
||||
});
|
||||
expect(aList.results.map((m) => m.id)).toContain(id);
|
||||
|
||||
// ...and must never leak into tenant B's scope.
|
||||
const bList: SearchResult = await memory.getAll({
|
||||
filters: { user_id: tenantB },
|
||||
});
|
||||
expect(bList.results.map((m) => m.id)).not.toContain(id);
|
||||
});
|
||||
|
||||
// A memory created with only some identity fields (e.g. run_id only) must not
|
||||
// let update() metadata inject a brand-new identity key. The default vector
|
||||
// store promotes camelCase aliases to snake_case on read, so both casings must be covered.
|
||||
test("metadata cannot inject an identity field on update (snake_case or camelCase)", async () => {
|
||||
const runId = `run_only_${Date.now()}`;
|
||||
const attacker = `attacker_${Date.now()}`;
|
||||
const addResult: SearchResult = await memory.add("Run-scoped secret", {
|
||||
runId,
|
||||
infer: false,
|
||||
});
|
||||
const id = addResult.results[0].id;
|
||||
|
||||
// Memory has no user_id/agent_id/actor_id — metadata tries to inject them,
|
||||
// in both snake_case and camelCase.
|
||||
await memory.update(id, {
|
||||
text: "Updated text",
|
||||
metadata: {
|
||||
user_id: attacker,
|
||||
agent_id: attacker,
|
||||
actor_id: attacker,
|
||||
userId: attacker,
|
||||
agentId: attacker,
|
||||
},
|
||||
});
|
||||
|
||||
// Still reachable under its original run_id scope...
|
||||
const runList: SearchResult = await memory.getAll({
|
||||
filters: { run_id: runId },
|
||||
});
|
||||
expect(runList.results.map((m) => m.id)).toContain(id);
|
||||
|
||||
// ...and the injected identity scopes must not resolve to it.
|
||||
const userList: SearchResult = await memory.getAll({
|
||||
filters: { user_id: attacker },
|
||||
});
|
||||
expect(userList.results.map((m) => m.id)).not.toContain(id);
|
||||
|
||||
const agentList: SearchResult = await memory.getAll({
|
||||
filters: { agent_id: attacker },
|
||||
});
|
||||
expect(agentList.results.map((m) => m.id)).not.toContain(id);
|
||||
});
|
||||
});
|
||||
|
||||
// ─── update() options: text / data / metadata / expirationDate ───
|
||||
|
||||
@@ -253,6 +253,34 @@ describe("Memory Input Validation", () => {
|
||||
});
|
||||
});
|
||||
|
||||
describe("non-string entity ID coercion", () => {
|
||||
it("coerces an integer user_id in getAll filters to its string form", async () => {
|
||||
const listSpy = jest
|
||||
.spyOn((memory as any).vectorStore, "list")
|
||||
.mockResolvedValue([[], 0]);
|
||||
|
||||
await memory.getAll({ filters: { user_id: 42 as any } });
|
||||
|
||||
const passedFilters = listSpy.mock.calls[0][0] as Record<string, any>;
|
||||
expect(passedFilters.user_id).toBe("42");
|
||||
|
||||
listSpy.mockRestore();
|
||||
});
|
||||
|
||||
it("coerces an integer user_id in search filters to its string form", async () => {
|
||||
const searchSpy = jest
|
||||
.spyOn((memory as any).vectorStore, "search")
|
||||
.mockResolvedValue([]);
|
||||
|
||||
await memory.search("q", { filters: { user_id: 42 as any } });
|
||||
|
||||
const passedFilters = searchSpy.mock.calls[0][2] as Record<string, any>;
|
||||
expect(passedFilters.user_id).toBe("42");
|
||||
|
||||
searchSpy.mockRestore();
|
||||
});
|
||||
});
|
||||
|
||||
describe("search() filter entity ID validation", () => {
|
||||
it("should throw error when user_id in filters is whitespace-only", async () => {
|
||||
await expect(
|
||||
|
||||
@@ -22,7 +22,16 @@ function sourceFiles(dir: string, acc: string[] = []): string[] {
|
||||
for (const entry of readdirSync(dir)) {
|
||||
const full = join(dir, entry);
|
||||
if (statSync(full).isDirectory()) {
|
||||
if (entry !== "tests" && entry !== "__tests__") sourceFiles(full, acc);
|
||||
// examples/ are dev scripts and community/ is the separate @mem0/community
|
||||
// package — neither ships in files:["dist"] nor is reachable via the mem0ai/oss
|
||||
// barrel, so a static import there cannot crash a Memory() consumer.
|
||||
if (
|
||||
entry !== "tests" &&
|
||||
entry !== "__tests__" &&
|
||||
entry !== "examples" &&
|
||||
entry !== "community"
|
||||
)
|
||||
sourceFiles(full, acc);
|
||||
} else if (entry.endsWith(".ts") && !entry.endsWith(".test.ts")) {
|
||||
acc.push(full);
|
||||
}
|
||||
|
||||
@@ -40,14 +40,15 @@ beforeEach(() => {
|
||||
});
|
||||
|
||||
describe("Qdrant URL port extraction (qdrant/qdrant-js#59 workaround)", () => {
|
||||
it("extracts port from HTTPS URL with explicit port", () => {
|
||||
new Qdrant({
|
||||
it("extracts port from HTTPS URL with explicit port", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "https://my-cluster.us-west-1-0.aws.cloud.qdrant.io:6333",
|
||||
apiKey: "test-key",
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
expect(capturedParams).toBeDefined();
|
||||
expect(capturedParams!.url).toBe(
|
||||
@@ -57,47 +58,50 @@ describe("Qdrant URL port extraction (qdrant/qdrant-js#59 workaround)", () => {
|
||||
expect(capturedParams!.apiKey).toBe("test-key");
|
||||
});
|
||||
|
||||
it("extracts port from HTTP URL with explicit port", () => {
|
||||
new Qdrant({
|
||||
it("extracts port from HTTP URL with explicit port", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "http://localhost:6333",
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
expect(capturedParams).toBeDefined();
|
||||
expect(capturedParams!.url).toBe("http://localhost:6333");
|
||||
expect(capturedParams!.port).toBe(6333);
|
||||
});
|
||||
|
||||
it("defaults to port 6333 when HTTPS URL has no explicit port", () => {
|
||||
new Qdrant({
|
||||
it("defaults to port 6333 when HTTPS URL has no explicit port", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "https://my-cluster.cloud.qdrant.io",
|
||||
apiKey: "test-key",
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
expect(capturedParams).toBeDefined();
|
||||
expect(capturedParams!.url).toBe("https://my-cluster.cloud.qdrant.io");
|
||||
expect(capturedParams!.port).toBe(6333);
|
||||
});
|
||||
|
||||
it("defaults to port 6333 when HTTP URL has no explicit port", () => {
|
||||
new Qdrant({
|
||||
it("defaults to port 6333 when HTTP URL has no explicit port", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "http://localhost",
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
expect(capturedParams).toBeDefined();
|
||||
expect(capturedParams!.port).toBe(6333);
|
||||
});
|
||||
|
||||
it("host+port config overrides URL-extracted port", () => {
|
||||
new Qdrant({
|
||||
it("host+port config overrides URL-extracted port", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "https://my-cluster.cloud.qdrant.io:6333",
|
||||
host: "custom-host",
|
||||
port: 9999,
|
||||
@@ -106,28 +110,28 @@ describe("Qdrant URL port extraction (qdrant/qdrant-js#59 workaround)", () => {
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
expect(capturedParams).toBeDefined();
|
||||
expect(capturedParams!.host).toBe("custom-host");
|
||||
expect(capturedParams!.port).toBe(9999);
|
||||
});
|
||||
|
||||
it("handles invalid URL gracefully without crashing", () => {
|
||||
expect(() => {
|
||||
new Qdrant({
|
||||
url: "not-a-valid-url",
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
}).not.toThrow();
|
||||
it("handles invalid URL gracefully without crashing", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "not-a-valid-url",
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await expect(store.initialize()).resolves.not.toThrow();
|
||||
|
||||
expect(capturedParams).toBeDefined();
|
||||
expect(capturedParams!.url).toBe("not-a-valid-url");
|
||||
expect(capturedParams!.port).toBe(6333);
|
||||
});
|
||||
|
||||
it("does not pass port when using pre-configured client", () => {
|
||||
it("does not pass port when using pre-configured client", async () => {
|
||||
const mockClient: any = {
|
||||
createCollection: jest.fn().mockResolvedValue(undefined),
|
||||
getCollection: jest.fn().mockResolvedValue({
|
||||
@@ -141,12 +145,13 @@ describe("Qdrant URL port extraction (qdrant/qdrant-js#59 workaround)", () => {
|
||||
deleteCollection: jest.fn().mockResolvedValue(undefined),
|
||||
};
|
||||
|
||||
new Qdrant({
|
||||
const store = new Qdrant({
|
||||
client: mockClient,
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
// QdrantClient constructor should NOT have been called
|
||||
expect(
|
||||
@@ -154,14 +159,15 @@ describe("Qdrant URL port extraction (qdrant/qdrant-js#59 workaround)", () => {
|
||||
).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("defaults to 6333 when HTTPS URL uses default port 443", () => {
|
||||
new Qdrant({
|
||||
it("defaults to 6333 when HTTPS URL uses default port 443", async () => {
|
||||
const store = new Qdrant({
|
||||
url: "https://my-cluster.cloud.qdrant.io:443",
|
||||
apiKey: "test-key",
|
||||
collectionName: "test",
|
||||
embeddingModelDims: 768,
|
||||
dimension: 768,
|
||||
});
|
||||
await store.initialize();
|
||||
|
||||
// 443 is default for HTTPS, so URL.port returns empty string — we default to 6333
|
||||
expect(capturedParams).toBeDefined();
|
||||
|
||||
+9
-9
@@ -20,7 +20,13 @@ from mem0.client.types import (
|
||||
from mem0.client.utils import api_error_handler
|
||||
|
||||
# Exception classes are referenced in docstrings only
|
||||
from mem0.memory.setup import get_user_id, is_aliased, mark_aliased, read_anon_ids, setup_config
|
||||
from mem0.memory.setup import (
|
||||
get_user_id,
|
||||
is_aliased,
|
||||
mark_aliased,
|
||||
read_anon_ids,
|
||||
setup_config,
|
||||
)
|
||||
from mem0.memory.telemetry import capture_client_event, client_telemetry
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -725,7 +731,6 @@ class MemoryClient:
|
||||
options: Optional[ProjectUpdateOptions] = None,
|
||||
custom_instructions: Optional[str] = None,
|
||||
custom_categories: Optional[List[str]] = None,
|
||||
retrieval_criteria: Optional[List[Dict[str, Any]]] = None,
|
||||
memory_depth: Optional[str] = None,
|
||||
usecase_setting: Optional[str] = None,
|
||||
multilingual: Optional[bool] = None,
|
||||
@@ -736,7 +741,6 @@ class MemoryClient:
|
||||
options: Typed options for the update operation (ProjectUpdateOptions).
|
||||
custom_instructions: New instructions for the project.
|
||||
custom_categories: New categories for the project.
|
||||
retrieval_criteria: New retrieval criteria for the project.
|
||||
memory_depth: Memory depth for the project.
|
||||
usecase_setting: Usecase setting for the project.
|
||||
multilingual: Whether to use the input language for memory storage and retrieval.
|
||||
@@ -761,7 +765,6 @@ class MemoryClient:
|
||||
for k, v in {
|
||||
"custom_instructions": custom_instructions,
|
||||
"custom_categories": custom_categories,
|
||||
"retrieval_criteria": retrieval_criteria,
|
||||
"memory_depth": memory_depth,
|
||||
"usecase_setting": usecase_setting,
|
||||
"multilingual": multilingual,
|
||||
@@ -773,7 +776,7 @@ class MemoryClient:
|
||||
if not kwargs:
|
||||
raise ValueError(
|
||||
"Currently we only support updating custom_instructions or "
|
||||
"custom_categories or retrieval_criteria, so you must "
|
||||
"custom_categories, so you must "
|
||||
"provide at least one of them"
|
||||
)
|
||||
|
||||
@@ -1630,7 +1633,6 @@ class AsyncMemoryClient:
|
||||
options: Optional[ProjectUpdateOptions] = None,
|
||||
custom_instructions: Optional[str] = None,
|
||||
custom_categories: Optional[List[str]] = None,
|
||||
retrieval_criteria: Optional[List[Dict[str, Any]]] = None,
|
||||
memory_depth: Optional[str] = None,
|
||||
usecase_setting: Optional[str] = None,
|
||||
multilingual: Optional[bool] = None,
|
||||
@@ -1641,7 +1643,6 @@ class AsyncMemoryClient:
|
||||
options: Typed options for the update operation (ProjectUpdateOptions).
|
||||
custom_instructions: New instructions for the project.
|
||||
custom_categories: New categories for the project.
|
||||
retrieval_criteria: New retrieval criteria for the project.
|
||||
memory_depth: Memory depth for the project.
|
||||
usecase_setting: Usecase setting for the project.
|
||||
multilingual: Whether to use the input language for memory storage and retrieval.
|
||||
@@ -1666,7 +1667,6 @@ class AsyncMemoryClient:
|
||||
for k, v in {
|
||||
"custom_instructions": custom_instructions,
|
||||
"custom_categories": custom_categories,
|
||||
"retrieval_criteria": retrieval_criteria,
|
||||
"memory_depth": memory_depth,
|
||||
"usecase_setting": usecase_setting,
|
||||
"multilingual": multilingual,
|
||||
@@ -1678,7 +1678,7 @@ class AsyncMemoryClient:
|
||||
if not kwargs:
|
||||
raise ValueError(
|
||||
"Currently we only support updating custom_instructions or "
|
||||
"custom_categories or retrieval_criteria, so you must "
|
||||
"custom_categories, so you must "
|
||||
"provide at least one of them"
|
||||
)
|
||||
|
||||
|
||||
+4
-28
@@ -177,7 +177,6 @@ class BaseProject(ABC):
|
||||
self,
|
||||
custom_instructions: Optional[str] = None,
|
||||
custom_categories: Optional[List[str]] = None,
|
||||
retrieval_criteria: Optional[List[Dict[str, Any]]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Update project settings.
|
||||
@@ -185,7 +184,6 @@ class BaseProject(ABC):
|
||||
Args:
|
||||
custom_instructions: New instructions for the project
|
||||
custom_categories: New categories for the project
|
||||
retrieval_criteria: New retrieval criteria for the project
|
||||
|
||||
Returns:
|
||||
Dictionary containing the API response.
|
||||
@@ -396,7 +394,6 @@ class Project(BaseProject):
|
||||
self,
|
||||
custom_instructions: Optional[str] = None,
|
||||
custom_categories: Optional[List[str]] = None,
|
||||
retrieval_criteria: Optional[List[Dict[str, Any]]] = None,
|
||||
multilingual: Optional[bool] = None,
|
||||
decay: Optional[bool] = None,
|
||||
) -> Dict[str, Any]:
|
||||
@@ -406,7 +403,6 @@ class Project(BaseProject):
|
||||
Args:
|
||||
custom_instructions: New instructions for the project
|
||||
custom_categories: New categories for the project
|
||||
retrieval_criteria: New retrieval criteria for the project
|
||||
multilingual: Whether to use the input language for memory storage and retrieval
|
||||
decay: Toggle Memory Decay for this project. When True, search-time
|
||||
ranking boosts recently-used memories and gently dampens stale ones; when
|
||||
@@ -422,24 +418,16 @@ class Project(BaseProject):
|
||||
NetworkError: If network connectivity issues occur.
|
||||
ValueError: If org_id or project_id are not set.
|
||||
"""
|
||||
if (
|
||||
custom_instructions is None
|
||||
and custom_categories is None
|
||||
and retrieval_criteria is None
|
||||
and multilingual is None
|
||||
and decay is None
|
||||
):
|
||||
if custom_instructions is None and custom_categories is None and multilingual is None and decay is None:
|
||||
raise ValueError(
|
||||
"At least one parameter must be provided for update: "
|
||||
"custom_instructions, custom_categories, retrieval_criteria, "
|
||||
"multilingual, decay"
|
||||
"custom_instructions, custom_categories, multilingual, decay"
|
||||
)
|
||||
|
||||
payload = self._prepare_params(
|
||||
{
|
||||
"custom_instructions": custom_instructions,
|
||||
"custom_categories": custom_categories,
|
||||
"retrieval_criteria": retrieval_criteria,
|
||||
"multilingual": multilingual,
|
||||
"decay": decay,
|
||||
}
|
||||
@@ -455,7 +443,6 @@ class Project(BaseProject):
|
||||
{
|
||||
"custom_instructions": custom_instructions,
|
||||
"custom_categories": custom_categories,
|
||||
"retrieval_criteria": retrieval_criteria,
|
||||
"multilingual": multilingual,
|
||||
"decay": decay,
|
||||
"sync_type": "sync",
|
||||
@@ -720,7 +707,6 @@ class AsyncProject(BaseProject):
|
||||
self,
|
||||
custom_instructions: Optional[str] = None,
|
||||
custom_categories: Optional[List[str]] = None,
|
||||
retrieval_criteria: Optional[List[Dict[str, Any]]] = None,
|
||||
multilingual: Optional[bool] = None,
|
||||
decay: Optional[bool] = None,
|
||||
) -> Dict[str, Any]:
|
||||
@@ -730,7 +716,6 @@ class AsyncProject(BaseProject):
|
||||
Args:
|
||||
custom_instructions: New instructions for the project
|
||||
custom_categories: New categories for the project
|
||||
retrieval_criteria: New retrieval criteria for the project
|
||||
multilingual: Whether to use the input language for memory storage and retrieval
|
||||
decay: Toggle Memory Decay for this project. When True, search-time
|
||||
ranking boosts recently-used memories and gently dampens stale ones; when
|
||||
@@ -746,24 +731,16 @@ class AsyncProject(BaseProject):
|
||||
NetworkError: If network connectivity issues occur.
|
||||
ValueError: If org_id or project_id are not set.
|
||||
"""
|
||||
if (
|
||||
custom_instructions is None
|
||||
and custom_categories is None
|
||||
and retrieval_criteria is None
|
||||
and multilingual is None
|
||||
and decay is None
|
||||
):
|
||||
if custom_instructions is None and custom_categories is None and multilingual is None and decay is None:
|
||||
raise ValueError(
|
||||
"At least one parameter must be provided for update: "
|
||||
"custom_instructions, custom_categories, retrieval_criteria, "
|
||||
"multilingual, decay"
|
||||
"custom_instructions, custom_categories, multilingual, decay"
|
||||
)
|
||||
|
||||
payload = self._prepare_params(
|
||||
{
|
||||
"custom_instructions": custom_instructions,
|
||||
"custom_categories": custom_categories,
|
||||
"retrieval_criteria": retrieval_criteria,
|
||||
"multilingual": multilingual,
|
||||
"decay": decay,
|
||||
}
|
||||
@@ -779,7 +756,6 @@ class AsyncProject(BaseProject):
|
||||
{
|
||||
"custom_instructions": custom_instructions,
|
||||
"custom_categories": custom_categories,
|
||||
"retrieval_criteria": retrieval_criteria,
|
||||
"multilingual": multilingual,
|
||||
"decay": decay,
|
||||
"sync_type": "async",
|
||||
|
||||
@@ -107,4 +107,3 @@ class ProjectUpdateOptions(BaseModel):
|
||||
memory_depth: Optional[str] = Field(default=None, description="Memory depth configuration")
|
||||
usecase_setting: Optional[Any] = Field(default=None, description="Use case specific settings")
|
||||
multilingual: Optional[bool] = Field(default=None, description="Whether to enable multilingual support")
|
||||
retrieval_criteria: Optional[List[Any]] = Field(default=None, description="Criteria for memory retrieval")
|
||||
|
||||
@@ -149,4 +149,11 @@ class AnthropicLLM(LLMBase):
|
||||
elif block.type == "tool_use":
|
||||
result["tool_calls"].append({"name": block.name, "arguments": block.input})
|
||||
return result
|
||||
return response.content[0].text
|
||||
|
||||
# Thinking-enabled responses put a thinking block before the text
|
||||
# block, and a response can carry no text block at all, so find the
|
||||
# text block like the tools branch above instead of indexing content[0].
|
||||
for block in response.content:
|
||||
if block.type == "text":
|
||||
return block.text
|
||||
return ""
|
||||
|
||||
@@ -35,11 +35,6 @@ class LLMBase(ABC):
|
||||
if not hasattr(self.config, "model"):
|
||||
raise ValueError("Configuration must have a 'model' attribute")
|
||||
|
||||
if not hasattr(self.config, "api_key") and not hasattr(self.config, "api_key"):
|
||||
# Check if API key is available via environment variable
|
||||
# This will be handled by individual providers
|
||||
pass
|
||||
|
||||
def _is_reasoning_model(self, model: str) -> bool:
|
||||
"""
|
||||
Check if the model is a reasoning model or GPT-5 series that doesn't support certain parameters.
|
||||
|
||||
@@ -15,7 +15,7 @@ class OpenAIStructuredLLM(LLMBase):
|
||||
self.config.model = "gpt-5-mini"
|
||||
|
||||
api_key = self.config.api_key or os.getenv("OPENAI_API_KEY")
|
||||
base_url = self.config.openai_base_url or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1"
|
||||
base_url = self.config.openai_base_url or os.getenv("OPENAI_BASE_URL") or "https://api.openai.com/v1"
|
||||
self.client = OpenAI(api_key=api_key, base_url=base_url)
|
||||
|
||||
def generate_response(
|
||||
|
||||
+20
-12
@@ -134,6 +134,20 @@ _SENSITIVE_SUFFIXES = (
|
||||
# Entity parameters that must be passed via filters, not top-level kwargs
|
||||
ENTITY_PARAMS = frozenset({"user_id", "agent_id", "run_id"})
|
||||
|
||||
# Tenant-scoping fields that update() must never let caller-supplied metadata overwrite (issues #4490, #6277).
|
||||
_IDENTITY_KEYS = ENTITY_PARAMS | {"actor_id"}
|
||||
|
||||
|
||||
def _strip_identity_keys(metadata: Dict[str, Any], existing_payload: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Drop identity keys from caller metadata; they are immutable after creation (issues #4490, #6277)."""
|
||||
clean = {}
|
||||
for key, value in metadata.items():
|
||||
if key not in _IDENTITY_KEYS:
|
||||
clean[key] = value
|
||||
elif value != existing_payload.get(key):
|
||||
logger.warning(f"update(): ignoring metadata['{key}'] - identity fields are immutable after creation")
|
||||
return clean
|
||||
|
||||
|
||||
def _reject_top_level_entity_params(kwargs: Dict[str, Any], method_name: str) -> None:
|
||||
"""Reject top-level entity parameters - must use filters instead."""
|
||||
@@ -424,7 +438,6 @@ class _OSSProject:
|
||||
self,
|
||||
custom_instructions: Optional[str] = None,
|
||||
custom_categories: Optional[list] = None,
|
||||
retrieval_criteria: Optional[list] = None,
|
||||
multilingual: Optional[bool] = None,
|
||||
decay: Optional[bool] = None,
|
||||
):
|
||||
@@ -438,7 +451,6 @@ class _AsyncOSSProject:
|
||||
self,
|
||||
custom_instructions: Optional[str] = None,
|
||||
custom_categories: Optional[list] = None,
|
||||
retrieval_criteria: Optional[list] = None,
|
||||
multilingual: Optional[bool] = None,
|
||||
decay: Optional[bool] = None,
|
||||
):
|
||||
@@ -1785,6 +1797,8 @@ class Memory(MemoryBase):
|
||||
memory_id (str): ID of the memory to update.
|
||||
text (str, optional): New content to update the memory with.
|
||||
metadata (dict, optional): Metadata to update with the memory. Defaults to None.
|
||||
``user_id``/``agent_id``/``run_id``/``actor_id`` are ignored here - they are
|
||||
immutable after creation.
|
||||
expiration_date (Any, optional): Date in YYYY-MM-DD format, or None to clear it.
|
||||
data (str, optional): Deprecated alias for ``text``. Will be removed in the next
|
||||
major release; use ``text`` instead.
|
||||
@@ -1993,7 +2007,7 @@ class Memory(MemoryBase):
|
||||
|
||||
new_metadata = deepcopy(existing_memory.payload)
|
||||
if metadata is not None:
|
||||
new_metadata.update(metadata)
|
||||
new_metadata.update(_strip_identity_keys(metadata, existing_memory.payload))
|
||||
|
||||
new_metadata["data"] = data
|
||||
new_metadata["hash"] = hashlib.md5(data.encode()).hexdigest()
|
||||
@@ -2001,10 +2015,6 @@ class Memory(MemoryBase):
|
||||
new_metadata["created_at"] = existing_memory.payload.get("created_at")
|
||||
new_metadata["updated_at"] = datetime.now(timezone.utc).isoformat()
|
||||
|
||||
# actor_id is immutable after creation (issue #4490)
|
||||
if "actor_id" in existing_memory.payload:
|
||||
new_metadata["actor_id"] = existing_memory.payload["actor_id"]
|
||||
|
||||
if data in existing_embeddings:
|
||||
embeddings = existing_embeddings[data]
|
||||
else:
|
||||
@@ -3416,6 +3426,8 @@ class AsyncMemory(MemoryBase):
|
||||
memory_id (str): ID of the memory to update.
|
||||
text (str, optional): New content to update the memory with.
|
||||
metadata (dict, optional): Metadata to update with the memory. Defaults to None.
|
||||
``user_id``/``agent_id``/``run_id``/``actor_id`` are ignored here - they are
|
||||
immutable after creation.
|
||||
expiration_date (Any, optional): Date in YYYY-MM-DD format, or None to clear it.
|
||||
data (str, optional): Deprecated alias for ``text``. Will be removed in the next
|
||||
major release; use ``text`` instead.
|
||||
@@ -3659,7 +3671,7 @@ class AsyncMemory(MemoryBase):
|
||||
|
||||
new_metadata = deepcopy(existing_memory.payload)
|
||||
if metadata is not None:
|
||||
new_metadata.update(metadata)
|
||||
new_metadata.update(_strip_identity_keys(metadata, existing_memory.payload))
|
||||
|
||||
new_metadata["data"] = data
|
||||
new_metadata["hash"] = hashlib.md5(data.encode()).hexdigest()
|
||||
@@ -3667,10 +3679,6 @@ class AsyncMemory(MemoryBase):
|
||||
new_metadata["created_at"] = existing_memory.payload.get("created_at")
|
||||
new_metadata["updated_at"] = datetime.now(timezone.utc).isoformat()
|
||||
|
||||
# actor_id is immutable after creation (issue #4490)
|
||||
if "actor_id" in existing_memory.payload:
|
||||
new_metadata["actor_id"] = existing_memory.payload["actor_id"]
|
||||
|
||||
if data in existing_embeddings:
|
||||
embeddings = existing_embeddings[data]
|
||||
else:
|
||||
|
||||
@@ -224,8 +224,15 @@ class BaiduDB(VectorStoreBase):
|
||||
output = []
|
||||
for row in res.rows:
|
||||
row_data = row.get("row", {})
|
||||
# Mochow returns the raw L2 distance (lower = closer). Convert it to a
|
||||
# similarity score (higher = better) to satisfy the VectorStoreBase
|
||||
# contract, mirroring the milvus provider. Non-L2 metrics already
|
||||
# return a higher-is-better score.
|
||||
raw_score = row.get("score", 0.0)
|
||||
if self.metric_type in (MetricType.L2, "L2"):
|
||||
raw_score = 1.0 / (1.0 + raw_score)
|
||||
output_data = OutputData(
|
||||
id=row_data.get("id"), score=row.get("score", 0.0), payload=row_data.get("metadata", {})
|
||||
id=row_data.get("id"), score=raw_score, payload=row_data.get("metadata", {})
|
||||
)
|
||||
output.append(output_data)
|
||||
|
||||
|
||||
@@ -358,12 +358,16 @@ class PineconeDB(VectorStoreBase):
|
||||
return self.client.list_indexes()
|
||||
|
||||
def delete_col(self):
|
||||
"""Delete an index/collection."""
|
||||
"""Delete the index, or clear only the configured namespace if one is set."""
|
||||
try:
|
||||
self.client.delete_index(self.collection_name)
|
||||
logger.info(f"Index {self.collection_name} deleted successfully")
|
||||
if self.namespace is not None:
|
||||
self.index.delete(delete_all=True, namespace=self.namespace)
|
||||
logger.info(f"Namespace {self.namespace} in index {self.collection_name} cleared successfully")
|
||||
else:
|
||||
self.client.delete_index(self.collection_name)
|
||||
logger.info(f"Index {self.collection_name} deleted successfully")
|
||||
except Exception as e:
|
||||
logger.error(f"Error deleting index {self.collection_name}: {e}")
|
||||
logger.error(f"Error deleting index {self.collection_name} (namespace={self.namespace}): {e}")
|
||||
|
||||
def col_info(self) -> Dict:
|
||||
"""
|
||||
@@ -419,8 +423,7 @@ class PineconeDB(VectorStoreBase):
|
||||
int: Total number of vectors.
|
||||
"""
|
||||
stats = self.index.describe_index_stats()
|
||||
if self.namespace:
|
||||
# Safely get the namespace stats and return vector_count, defaulting to 0 if not found
|
||||
if self.namespace is not None:
|
||||
namespace_summary = (stats.namespaces or {}).get(self.namespace)
|
||||
if namespace_summary:
|
||||
return namespace_summary.vector_count or 0
|
||||
|
||||
@@ -596,3 +596,11 @@ class Qdrant(VectorStoreBase):
|
||||
logger.warning(f"Resetting index {self.collection_name}...")
|
||||
self.delete_col()
|
||||
self.create_col(self.embedding_model_dims, self.on_disk)
|
||||
if self.is_local:
|
||||
# Local delete_collection() rmtree's with ignore_errors=True and leaves its
|
||||
# sqlite handle open, so where an open file blocks unlink (Windows, NFS) the
|
||||
# recreated collection re-adopts the old storage.sqlite. Drop what survived.
|
||||
self.client.delete(
|
||||
collection_name=self.collection_name,
|
||||
points_selector=models.FilterSelector(filter=models.Filter(must=[])),
|
||||
)
|
||||
|
||||
+1
-1
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "mem0ai"
|
||||
version = "2.0.12"
|
||||
version = "2.0.13"
|
||||
description = "Long-term memory for AI Agents"
|
||||
authors = [
|
||||
{ name = "Mem0", email = "support@mem0.ai" }
|
||||
|
||||
+14
-9
@@ -3,7 +3,7 @@ import secrets
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
from db import get_db
|
||||
from db import SessionLocal
|
||||
from fastapi import Depends, HTTPException, Request
|
||||
from fastapi.security import APIKeyHeader, HTTPAuthorizationCredentials, HTTPBearer
|
||||
from jose import JWTError, jwt
|
||||
@@ -145,19 +145,24 @@ async def verify_auth(
|
||||
request: Request,
|
||||
credentials: HTTPAuthorizationCredentials | None = Depends(bearer_scheme),
|
||||
x_api_key: str | None = Depends(api_key_header),
|
||||
db: Session = Depends(get_db),
|
||||
) -> User | None:
|
||||
"""Authenticate via JWT, X-API-Key, or legacy ADMIN_API_KEY. Returns User or None."""
|
||||
"""Authenticate via JWT, X-API-Key, or legacy ADMIN_API_KEY. Returns User or None.
|
||||
|
||||
A short-lived session is opened only on the branches that query the DB, so no
|
||||
pooled connection is held for the lifetime of the (possibly long-running) request.
|
||||
"""
|
||||
if credentials is not None:
|
||||
_mark_auth_type(request, "bearer")
|
||||
return _resolve_user_from_jwt(credentials.credentials, db)
|
||||
with SessionLocal() as db:
|
||||
return _resolve_user_from_jwt(credentials.credentials, db)
|
||||
|
||||
if x_api_key is not None:
|
||||
if ADMIN_API_KEY and secrets.compare_digest(x_api_key, ADMIN_API_KEY):
|
||||
_mark_auth_type(request, "admin_api_key")
|
||||
return None
|
||||
_mark_auth_type(request, "api_key")
|
||||
return _resolve_user_from_api_key(x_api_key, db)
|
||||
with SessionLocal() as db:
|
||||
return _resolve_user_from_api_key(x_api_key, db)
|
||||
|
||||
if AUTH_DISABLED:
|
||||
_mark_auth_type(request, "disabled")
|
||||
@@ -173,12 +178,12 @@ async def verify_auth(
|
||||
async def require_auth(
|
||||
request: Request,
|
||||
user: User | None = Depends(verify_auth),
|
||||
db: Session = Depends(get_db),
|
||||
) -> User:
|
||||
"""Like verify_auth but guarantees a non-None User. Use for endpoints that require auth."""
|
||||
if user is None:
|
||||
if getattr(request.state, "auth_type", "none") in {"admin_api_key", "disabled"}:
|
||||
default_user = _get_default_user(db)
|
||||
with SessionLocal() as db:
|
||||
default_user = _get_default_user(db)
|
||||
if default_user is not None:
|
||||
return default_user
|
||||
raise HTTPException(status_code=401, detail="Authentication required.")
|
||||
@@ -193,7 +198,6 @@ _BOOTSTRAP_ADMIN = User(
|
||||
async def require_admin(
|
||||
request: Request,
|
||||
user: User | None = Depends(verify_auth),
|
||||
db: Session = Depends(get_db),
|
||||
) -> User:
|
||||
"""Like require_auth but also enforces admin role.
|
||||
|
||||
@@ -203,7 +207,8 @@ async def require_admin(
|
||||
auth_type = getattr(request.state, "auth_type", "none")
|
||||
if user is None:
|
||||
if auth_type in {"admin_api_key", "disabled"}:
|
||||
default_user = _get_default_user(db)
|
||||
with SessionLocal() as db:
|
||||
default_user = _get_default_user(db)
|
||||
if default_user is not None:
|
||||
if default_user.role != "admin":
|
||||
raise HTTPException(status_code=403, detail="Admin role required.")
|
||||
|
||||
+17
-9
@@ -175,18 +175,23 @@ def update_me(
|
||||
user: User = Depends(require_auth),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
if body.name is not None and body.name.strip():
|
||||
user.name = body.name.strip()
|
||||
# require_auth resolves the user in its own short-lived session, so `user` is
|
||||
# detached from this request's `db`. Load a session-managed copy to mutate.
|
||||
db_user = db.get(User, user.id)
|
||||
if db_user is None:
|
||||
raise HTTPException(status_code=404, detail="User not found.")
|
||||
|
||||
if body.email is not None and body.email != user.email:
|
||||
collision = db.scalar(select(User).where(User.email == body.email, User.id != user.id))
|
||||
if body.name is not None and body.name.strip():
|
||||
db_user.name = body.name.strip()
|
||||
|
||||
if body.email is not None and body.email != db_user.email:
|
||||
collision = db.scalar(select(User).where(User.email == body.email, User.id != db_user.id))
|
||||
if collision is not None:
|
||||
raise HTTPException(status_code=409, detail="Email is already in use.")
|
||||
user.email = body.email
|
||||
db_user.email = body.email
|
||||
|
||||
db.commit()
|
||||
db.refresh(user)
|
||||
return user
|
||||
return db_user
|
||||
|
||||
|
||||
@router.post("/change-password", response_model=MessageResponse)
|
||||
@@ -195,12 +200,15 @@ def change_password(
|
||||
user: User = Depends(require_auth),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
if not verify_password(body.current_password, user.password_hash):
|
||||
# require_auth resolves the user in its own short-lived session, so `user` is
|
||||
# detached from this request's `db`. Load a session-managed copy to mutate.
|
||||
db_user = db.get(User, user.id)
|
||||
if db_user is None or not verify_password(body.current_password, db_user.password_hash):
|
||||
raise HTTPException(status_code=401, detail="Current password is incorrect.")
|
||||
|
||||
_require_password_length(body.new_password)
|
||||
|
||||
user.password_hash = hash_password(body.new_password)
|
||||
db_user.password_hash = hash_password(body.new_password)
|
||||
db.commit()
|
||||
return MessageResponse(message="Password updated.")
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user