feat(vector-stores): add Databricks provider to TypeScript OSS SDK (#5824)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -6,7 +6,8 @@ description: "Use Databricks Vector Search as a serverless vector store in Mem0
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
|
||||
@@ -36,10 +37,44 @@ messages = [
|
||||
m.add(messages, user_id="alice", metadata={"category": "movies"})
|
||||
```
|
||||
|
||||
```typescript TypeScript
|
||||
// Requires the Databricks SQL driver (peer dependency): pnpm add @databricks/sql
|
||||
import { Memory } from 'mem0ai/oss';
|
||||
|
||||
const config = {
|
||||
vectorStore: {
|
||||
provider: 'databricks',
|
||||
config: {
|
||||
workspaceUrl: 'https://your-workspace.databricks.com',
|
||||
// SQL warehouse HTTP path, used for index writes (required)
|
||||
httpPath: '/sql/1.0/warehouses/your-warehouse-id',
|
||||
accessToken: 'your-access-token',
|
||||
catalog: 'your_catalog',
|
||||
schema: 'your_schema',
|
||||
tableName: 'your_table',
|
||||
collectionName: 'your_index_name',
|
||||
embeddingModelDims: 1536,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const memory = new Memory(config);
|
||||
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", metadata: { category: "movies" } });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
### Config
|
||||
|
||||
Here are the parameters available for configuring Databricks Vector Search:
|
||||
|
||||
<Tabs>
|
||||
<Tab title="Python">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `workspace_url` | The URL of your Databricks workspace | **Required** |
|
||||
@@ -60,6 +95,32 @@ Here are the parameters available for configuring Databricks Vector Search:
|
||||
| `pipeline_type` | Sync pipeline type: `TRIGGERED` or `CONTINUOUS` | `TRIGGERED` |
|
||||
| `warehouse_name` | Databricks SQL warehouse name (if using SQL warehouse) | `None` |
|
||||
| `query_type` | Query type: `ANN` or `HYBRID` | `ANN` |
|
||||
</Tab>
|
||||
<Tab title="TypeScript">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `workspaceUrl` | The URL of your Databricks workspace (or pass `host`) | **Required** |
|
||||
| `httpPath` | SQL warehouse HTTP path, used for index writes | **Required** |
|
||||
| `accessToken` | Personal Access Token for authentication | `None` |
|
||||
| `clientId` | Service principal client ID (alternative to `accessToken`) | `None` |
|
||||
| `clientSecret` | Service principal client secret (required with `clientId`) | `None` |
|
||||
| `endpointName` | Name of the Vector Search endpoint | `mem0_vector_search` |
|
||||
| `endpointType` | Type of endpoint (`STANDARD` or `STORAGE_OPTIMIZED`) | `STANDARD` |
|
||||
| `pipelineType` | Delta Sync pipeline type: `TRIGGERED` or `CONTINUOUS` | `TRIGGERED` |
|
||||
| `queryType` | Query type: `ANN` or `HYBRID` | `ANN` |
|
||||
| `catalog` | Unity Catalog catalog name | `main` |
|
||||
| `schema` | Unity Catalog schema name | `default` |
|
||||
| `collectionName` | Vector Search index name | `mem0` |
|
||||
| `tableName` | Source Delta table name | falls back to `collectionName` |
|
||||
| `embeddingModelDims` | Dimension of self-managed embeddings | `1536` |
|
||||
| `syncPollIntervalMs` | Poll interval while waiting for a `TRIGGERED` sync | `1000` |
|
||||
| `syncTimeoutMs` | Timeout while waiting for an index sync | `300000` |
|
||||
|
||||
<Note>
|
||||
The TypeScript provider uses `DELTA_SYNC` indexes with self-managed embeddings: pass vectors directly. `DIRECT_ACCESS` indexes, Databricks-computed embeddings (`embedding_model_endpoint_name`), and Azure AD auth are Python-only today. It writes to the index through a SQL warehouse, so `httpPath` is required, and `@databricks/sql` must be installed as a peer dependency.
|
||||
</Note>
|
||||
</Tab>
|
||||
</Tabs>
|
||||
|
||||
### Authentication
|
||||
|
||||
|
||||
+10
-2
@@ -109,12 +109,13 @@
|
||||
},
|
||||
"peerDependencies": {
|
||||
"@anthropic-ai/sdk": "^0.40.1",
|
||||
"@aws-sdk/client-neptune-graph": "3.966.0",
|
||||
"@aws-sdk/client-neptune-graph": ">=3.0.0 <3.968.0",
|
||||
"@aws-sdk/client-s3vectors": "3.967.0",
|
||||
"@mochow/mochow-sdk-node": "^2.1.5",
|
||||
"@azure/identity": "^4.0.0",
|
||||
"@azure/search-documents": "^12.0.0",
|
||||
"@cloudflare/workers-types": "^4.20250504.0",
|
||||
"@databricks/sql": "^1.16.0",
|
||||
"@google-cloud/aiplatform": "^6.8.0",
|
||||
"@google/genai": "^1.40.0",
|
||||
"@huggingface/transformers": "^3.0.0 || ^4.0.0",
|
||||
@@ -161,6 +162,12 @@
|
||||
},
|
||||
"@aws-sdk/client-bedrock-runtime": {
|
||||
"optional": true
|
||||
},
|
||||
"@databricks/sql": {
|
||||
"optional": true
|
||||
},
|
||||
"@aws-sdk/client-neptune-graph": {
|
||||
"optional": true
|
||||
}
|
||||
},
|
||||
"engines": {
|
||||
@@ -198,7 +205,8 @@
|
||||
"@modelcontextprotocol/sdk": "^1.25.4",
|
||||
"esbuild": ">=0.28.1",
|
||||
"undici@<6.27.0": ">=6.27.0 <8.0.0",
|
||||
"@aws-sdk/client-bedrock-runtime": "3.967.0"
|
||||
"@aws-sdk/client-bedrock-runtime": "3.967.0",
|
||||
"@aws-sdk/client-neptune-graph": "3.966.0"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Generated
+586
-7
File diff suppressed because it is too large
Load Diff
@@ -60,9 +60,14 @@ export class ConfigManager {
|
||||
})(),
|
||||
},
|
||||
vectorStore: {
|
||||
provider:
|
||||
// Every factory already matches the provider case-insensitively, so a capitalized
|
||||
// name constructs the right store -- but the `provider === "memory"` comparisons that
|
||||
// pick per-provider entity-store settings do not. Normalize once, here, so those
|
||||
// comparisons cannot silently miss.
|
||||
provider: (
|
||||
userConfig.vectorStore?.provider ||
|
||||
DEFAULT_MEMORY_CONFIG.vectorStore.provider,
|
||||
DEFAULT_MEMORY_CONFIG.vectorStore.provider
|
||||
).toLowerCase(),
|
||||
config: (() => {
|
||||
const defaultConf = DEFAULT_MEMORY_CONFIG.vectorStore.config;
|
||||
const userConf = userConfig.vectorStore?.config;
|
||||
|
||||
@@ -36,6 +36,7 @@ export * from "./vector_stores/langchain";
|
||||
export * from "./vector_stores/vectorize";
|
||||
export * from "./vector_stores/azure_ai_search";
|
||||
export * from "./vector_stores/pgvector";
|
||||
export * from "./vector_stores/databricks";
|
||||
export * from "./vector_stores/neptune_analytics";
|
||||
export * from "./vector_stores/elasticsearch";
|
||||
export * from "./vector_stores/upstash_vector";
|
||||
|
||||
@@ -7,6 +7,7 @@ import {
|
||||
Message,
|
||||
SearchFilters,
|
||||
SearchResult,
|
||||
VectorStoreConfig,
|
||||
} from "../types";
|
||||
import {
|
||||
EmbedderFactory,
|
||||
@@ -286,18 +287,24 @@ export class Memory {
|
||||
|
||||
private async getEntityStore(): Promise<VectorStore> {
|
||||
if (!this._entityStore) {
|
||||
const entityProvider = this.config.vectorStore.provider;
|
||||
const entityCollectionName = `${this.collectionName}_entities`;
|
||||
const entityConfig = {
|
||||
const entityConfig: VectorStoreConfig = {
|
||||
...this.config.vectorStore.config,
|
||||
collectionName: entityCollectionName,
|
||||
};
|
||||
// For file-based stores (memory/SQLite), always use a separate DB for entities
|
||||
if (this.config.vectorStore.provider === "memory") {
|
||||
if (entityProvider === "memory") {
|
||||
const basePath = entityConfig.dbPath || getDefaultVectorStoreDbPath();
|
||||
entityConfig.dbPath = basePath.replace(/\.db$/, "_entities.db");
|
||||
}
|
||||
if (entityProvider === "databricks") {
|
||||
entityConfig.tableName = entityConfig.tableName
|
||||
? `${entityConfig.tableName}_entities`
|
||||
: entityCollectionName;
|
||||
}
|
||||
this._entityStore = VectorStoreFactory.create(
|
||||
this.config.vectorStore.provider,
|
||||
entityProvider,
|
||||
entityConfig,
|
||||
);
|
||||
await this._entityStore.initialize();
|
||||
@@ -1747,7 +1754,7 @@ export class Memory {
|
||||
await this.db.reset();
|
||||
|
||||
// Check provider before attempting deleteCol
|
||||
if (this.config.vectorStore.provider.toLowerCase() !== "langchain") {
|
||||
if (this.config.vectorStore.provider !== "langchain") {
|
||||
try {
|
||||
await this.vectorStore.deleteCol();
|
||||
} catch (e) {
|
||||
|
||||
@@ -55,6 +55,7 @@ import { HuggingFaceEmbedder } from "../embeddings/huggingface";
|
||||
import { LangchainVectorStore } from "../vector_stores/langchain";
|
||||
import { AzureAISearch } from "../vector_stores/azure_ai_search";
|
||||
import { PGVector } from "../vector_stores/pgvector";
|
||||
import { DatabricksVectorStore } from "../vector_stores/databricks";
|
||||
import { NeptuneAnalyticsVectorStore } from "../vector_stores/neptune_analytics";
|
||||
import { VertexAIEmbedder } from "../embeddings/vertexai";
|
||||
import { ElasticsearchDB } from "../vector_stores/elasticsearch";
|
||||
@@ -173,6 +174,8 @@ export class VectorStoreFactory {
|
||||
return new VertexAIVectorSearch(config as any);
|
||||
case "pgvector":
|
||||
return new PGVector(config as any);
|
||||
case "databricks":
|
||||
return new DatabricksVectorStore(config as any);
|
||||
case "neptune":
|
||||
case "neptune-analytics":
|
||||
return new NeptuneAnalyticsVectorStore(config as any);
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,10 +1,10 @@
|
||||
import {
|
||||
ExecuteQueryCommand,
|
||||
NeptuneGraphClient,
|
||||
} from "@aws-sdk/client-neptune-graph";
|
||||
import { VectorStore } from "./base";
|
||||
import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types";
|
||||
|
||||
/**
|
||||
* The `@aws-sdk/client-neptune-graph` dependency is loaded on first use via dynamic
|
||||
* `import()` so the package stays optional (mirrors `aws_bedrock.ts`).
|
||||
*/
|
||||
interface NeptuneAnalyticsConfig extends VectorStoreConfig {
|
||||
graphIdentifier?: string;
|
||||
endpoint?: string;
|
||||
@@ -14,7 +14,14 @@ interface NeptuneAnalyticsConfig extends VectorStoreConfig {
|
||||
}
|
||||
|
||||
interface NeptuneGraphClientLike {
|
||||
send(command: ExecuteQueryCommand): Promise<NeptuneExecuteQueryOutput>;
|
||||
send(command: any): Promise<NeptuneExecuteQueryOutput>;
|
||||
}
|
||||
|
||||
interface NeptuneSDK {
|
||||
NeptuneGraphClient: new (
|
||||
config: Record<string, any>,
|
||||
) => NeptuneGraphClientLike;
|
||||
ExecuteQueryCommand: new (input: Record<string, any>) => any;
|
||||
}
|
||||
|
||||
interface NeptuneExecuteQueryOutput {
|
||||
@@ -33,7 +40,10 @@ interface WhereClauseResult {
|
||||
}
|
||||
|
||||
export class NeptuneAnalyticsVectorStore implements VectorStore {
|
||||
private readonly client: NeptuneGraphClientLike;
|
||||
private clientConfig: Record<string, any>;
|
||||
private clientOverride?: NeptuneGraphClientLike;
|
||||
private sdkPromise?: Promise<NeptuneSDK>;
|
||||
private clientPromise?: Promise<NeptuneGraphClientLike>;
|
||||
private readonly graphIdentifier: string;
|
||||
private readonly collectionName: string;
|
||||
private readonly collectionLabel: string;
|
||||
@@ -54,8 +64,8 @@ export class NeptuneAnalyticsVectorStore implements VectorStore {
|
||||
this.userLabelExpr = this.escapeLabel(this.userLabel);
|
||||
this.userNodeId = "mem0-user";
|
||||
this.dimension = config.dimension || 1536;
|
||||
this.client =
|
||||
config.client || new NeptuneGraphClient(this.buildClientConfig(config));
|
||||
this.clientConfig = this.buildClientConfig(config);
|
||||
this.clientOverride = config.client;
|
||||
|
||||
void this.initialize().catch(console.error);
|
||||
}
|
||||
@@ -171,6 +181,18 @@ export class NeptuneAnalyticsVectorStore implements VectorStore {
|
||||
const hasPayload = !!payload && Object.keys(payload).length > 0;
|
||||
const hasVector = vector.length > 0;
|
||||
|
||||
// ponytail: a combined update writes the payload before the embedding, and Neptune's vector
|
||||
// index isn't transactional -- if the upsert below fails, the new payload would otherwise be
|
||||
// left committed against the stale embedding (searches would match the old vector but return
|
||||
// the new metadata). Capture the prior node so a failed upsert can be restored; this is
|
||||
// best-effort compensation, not a rollback. Only needed when both writes happen -- a
|
||||
// payload-only or vector-only update can't desync.
|
||||
// The restore assumes a single writer per vectorId -- concurrent updates to the same node can
|
||||
// interleave and clobber each other's compensation. AWS advises against concurrent same-vertex
|
||||
// writes to the Neptune Analytics vector index for exactly this reason.
|
||||
const priorResult =
|
||||
hasPayload && hasVector ? await this.get(vectorId) : null;
|
||||
|
||||
if (hasPayload) {
|
||||
const properties = this.buildStoredPayload(payload);
|
||||
await this.executeQuery(
|
||||
@@ -187,20 +209,45 @@ export class NeptuneAnalyticsVectorStore implements VectorStore {
|
||||
}
|
||||
|
||||
if (hasVector) {
|
||||
const updateResults = await this.executeQuery(
|
||||
`
|
||||
MATCH (n:${this.collectionLabelExpr} {\`~id\`: $vectorId})
|
||||
WITH n, $embedding AS embedding
|
||||
CALL neptune.algo.vectors.upsert(n, embedding)
|
||||
YIELD success
|
||||
RETURN success
|
||||
`,
|
||||
{
|
||||
vectorId,
|
||||
embedding: vector,
|
||||
},
|
||||
);
|
||||
this.assertSuccessfulResults(updateResults, "Update");
|
||||
try {
|
||||
const updateResults = await this.executeQuery(
|
||||
`
|
||||
MATCH (n:${this.collectionLabelExpr} {\`~id\`: $vectorId})
|
||||
WITH n, $embedding AS embedding
|
||||
CALL neptune.algo.vectors.upsert(n, embedding)
|
||||
YIELD success
|
||||
RETURN success
|
||||
`,
|
||||
{
|
||||
vectorId,
|
||||
embedding: vector,
|
||||
},
|
||||
);
|
||||
this.assertSuccessfulResults(updateResults, "Update");
|
||||
} catch (error) {
|
||||
if (priorResult) {
|
||||
try {
|
||||
await this.executeQuery(
|
||||
`
|
||||
MATCH (n:${this.collectionLabelExpr} {\`~id\`: $vectorId})
|
||||
SET n = $properties
|
||||
RETURN n
|
||||
`,
|
||||
{
|
||||
vectorId,
|
||||
properties: priorResult.payload,
|
||||
},
|
||||
);
|
||||
} catch (restoreError) {
|
||||
// Do not mask the original failure with a compensation failure.
|
||||
console.error(
|
||||
"Neptune Analytics: failed to restore prior payload after a failed update upsert",
|
||||
restoreError,
|
||||
);
|
||||
}
|
||||
}
|
||||
throw error;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -985,12 +1032,58 @@ export class NeptuneAnalyticsVectorStore implements VectorStore {
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Load the optional AWS SDK on first use.
|
||||
*
|
||||
* This MUST be a dynamic `import()`, never `require()`: tsup/esbuild rewrite
|
||||
* `require()` in the published ESM bundle (`dist/oss/index.mjs`) into a
|
||||
* `__require` shim that throws `Dynamic require of "..." is not supported`,
|
||||
* so every ESM consumer would hit a dead provider even with the SDK installed.
|
||||
*/
|
||||
private async getSDK(): Promise<NeptuneSDK> {
|
||||
if (!this.sdkPromise) {
|
||||
this.sdkPromise = import("@aws-sdk/client-neptune-graph").then(
|
||||
(sdk) => sdk as unknown as NeptuneSDK,
|
||||
(err) => {
|
||||
// Let a later call retry rather than caching the rejection forever.
|
||||
this.sdkPromise = undefined;
|
||||
const detail = err instanceof Error ? err.message : String(err);
|
||||
throw new Error(
|
||||
"The '@aws-sdk/client-neptune-graph' package is required to use the Neptune Analytics vector store. " +
|
||||
`Install it with: npm install @aws-sdk/client-neptune-graph (original error: ${detail})`,
|
||||
);
|
||||
},
|
||||
);
|
||||
}
|
||||
return this.sdkPromise;
|
||||
}
|
||||
|
||||
/** Memoized Neptune client; an injected `config.client` short-circuits the SDK. */
|
||||
private async getClient(): Promise<NeptuneGraphClientLike> {
|
||||
if (this.clientOverride) return this.clientOverride;
|
||||
if (!this.clientPromise) {
|
||||
this.clientPromise = this.getSDK()
|
||||
.then(
|
||||
({ NeptuneGraphClient }) => new NeptuneGraphClient(this.clientConfig),
|
||||
)
|
||||
.catch((err) => {
|
||||
// Mirror getSDK(): drop the rejected promise so a later call retries rather than
|
||||
// replaying a cached rejection forever (a rejected promise is still truthy, so the
|
||||
// `!this.clientPromise` guard above would otherwise never re-enter).
|
||||
this.clientPromise = undefined;
|
||||
throw err;
|
||||
});
|
||||
}
|
||||
return this.clientPromise;
|
||||
}
|
||||
|
||||
private async executeQuery(
|
||||
queryString: string,
|
||||
parameters: Record<string, any> = {},
|
||||
): Promise<NeptuneQueryRecord[]> {
|
||||
const response = await this.client.send(
|
||||
new ExecuteQueryCommand({
|
||||
const [client, sdk] = await Promise.all([this.getClient(), this.getSDK()]);
|
||||
const response = await client.send(
|
||||
new sdk.ExecuteQueryCommand({
|
||||
graphIdentifier: this.graphIdentifier,
|
||||
language: "OPEN_CYPHER",
|
||||
queryString,
|
||||
|
||||
@@ -522,6 +522,34 @@ describe("ConfigManager", () => {
|
||||
expect(cfg.embedder.config).not.toHaveProperty("embedding_dims");
|
||||
});
|
||||
});
|
||||
|
||||
describe("mergeConfig - vector store provider normalization", () => {
|
||||
const baseLlm = { provider: "openai", config: { apiKey: "test-key" } };
|
||||
const baseEmbedder = { provider: "openai", config: { apiKey: "test-key" } };
|
||||
|
||||
it.each(["Memory", "DATABRICKS", "QdRaNt"])(
|
||||
"lowercases the vector store provider %p",
|
||||
(provider) => {
|
||||
const config = ConfigManager.mergeConfig({
|
||||
embedder: baseEmbedder,
|
||||
vectorStore: { provider, config: { collectionName: "test" } },
|
||||
llm: baseLlm,
|
||||
});
|
||||
|
||||
expect(config.vectorStore.provider).toBe(provider.toLowerCase());
|
||||
},
|
||||
);
|
||||
|
||||
it("still falls back to the default provider when none is given", () => {
|
||||
const config = ConfigManager.mergeConfig({
|
||||
embedder: baseEmbedder,
|
||||
vectorStore: { config: { collectionName: "test" } } as any,
|
||||
llm: baseLlm,
|
||||
});
|
||||
|
||||
expect(config.vectorStore.provider).toBe("memory");
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -183,6 +183,11 @@ jest.mock("../src/vector_stores/pgvector", () => ({
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "pgvector", config })),
|
||||
}));
|
||||
jest.mock("../src/vector_stores/databricks", () => ({
|
||||
DatabricksVectorStore: jest
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "databricks", config })),
|
||||
}));
|
||||
jest.mock("../src/vector_stores/neptune_analytics", () => ({
|
||||
NeptuneAnalyticsVectorStore: jest.fn().mockImplementation((config) => ({
|
||||
type: "neptune-analytics",
|
||||
@@ -346,6 +351,7 @@ describe("VectorStoreFactory", () => {
|
||||
["vectorize"],
|
||||
["azure-ai-search"],
|
||||
["pgvector"],
|
||||
["databricks"],
|
||||
["neptune"],
|
||||
["neptune-analytics"],
|
||||
["upstash_vector"],
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -32,6 +32,7 @@ const external = [
|
||||
"@azure/identity",
|
||||
"cloudflare",
|
||||
"@cloudflare/workers-types",
|
||||
"@databricks/sql",
|
||||
"@langchain/core",
|
||||
"fastembed",
|
||||
"compromise",
|
||||
|
||||
@@ -232,6 +232,20 @@ class NeptuneAnalyticsVector(VectorStoreBase):
|
||||
vector (Optional[List[float]]): New embedding vector.
|
||||
payload (Optional[Dict]): New metadata to replace existing payload.
|
||||
"""
|
||||
# ponytail: a combined update writes the payload before the embedding, and Neptune's
|
||||
# vector index isn't transactional -- if the upsert below fails, the new payload would
|
||||
# otherwise be left committed against the stale embedding. Capture the prior properties
|
||||
# so a failed upsert can be restored; this is best-effort compensation, not a rollback.
|
||||
# Only needed when both writes happen -- a payload-only or vector-only update can't desync.
|
||||
# The restore assumes a single writer per vector_id -- concurrent updates to the same node
|
||||
# can interleave and clobber each other's compensation. AWS advises against concurrent
|
||||
# same-vertex writes to the Neptune Analytics vector index for exactly this reason.
|
||||
prior_properties = None
|
||||
if payload and vector:
|
||||
prior = self.get(vector_id)
|
||||
if prior is not None:
|
||||
prior_properties = dict(prior.payload or {})
|
||||
prior_properties[self._FIELD_LABEL] = self.collection_name
|
||||
|
||||
if payload:
|
||||
# Replace payload
|
||||
@@ -242,9 +256,9 @@ class NeptuneAnalyticsVector(VectorStoreBase):
|
||||
"vector_id": vector_id
|
||||
}
|
||||
query_string_embedding = f"""
|
||||
MATCH (n :{self.collection_name})
|
||||
WHERE id(n) = $vector_id
|
||||
SET n = $properties
|
||||
MATCH (n :{self.collection_name})
|
||||
WHERE id(n) = $vector_id
|
||||
SET n = $properties
|
||||
"""
|
||||
self.execute_query(query_string_embedding, para_payload)
|
||||
|
||||
@@ -254,14 +268,40 @@ class NeptuneAnalyticsVector(VectorStoreBase):
|
||||
"vector_id": vector_id
|
||||
}
|
||||
query_string_embedding = f"""
|
||||
MATCH (n :{self.collection_name})
|
||||
WHERE id(n) = $vector_id
|
||||
WITH $embedding as embedding, n as n
|
||||
CALL neptune.algo.vectors.upsert(n, embedding)
|
||||
YIELD success
|
||||
RETURN success
|
||||
MATCH (n :{self.collection_name})
|
||||
WHERE id(n) = $vector_id
|
||||
WITH $embedding as embedding, n as n
|
||||
CALL neptune.algo.vectors.upsert(n, embedding)
|
||||
YIELD success
|
||||
RETURN success
|
||||
"""
|
||||
self.execute_query(query_string_embedding, para_embedding)
|
||||
try:
|
||||
result = self.execute_query(query_string_embedding, para_embedding)
|
||||
# A soft {"success": False} row desyncs the payload from the embedding just as
|
||||
# much as a thrown error, so treat it as a failure and let the rollback below fire.
|
||||
# Mirrors the TS store's assertSuccessfulResults() check (Python's
|
||||
# _process_success_message only logs, so it cannot drive the rollback).
|
||||
for row in result or []:
|
||||
if "success" in row and row["success"] is not True:
|
||||
raise RuntimeError(f"Neptune Analytics update upsert reported failure for {vector_id}")
|
||||
except Exception:
|
||||
if prior_properties is not None:
|
||||
try:
|
||||
restore_query = f"""
|
||||
MATCH (n :{self.collection_name})
|
||||
WHERE id(n) = $vector_id
|
||||
SET n = $properties
|
||||
"""
|
||||
self.execute_query(
|
||||
restore_query,
|
||||
{"properties": prior_properties, "vector_id": vector_id},
|
||||
)
|
||||
except Exception:
|
||||
logger.error(
|
||||
f"Neptune Analytics: failed to restore prior payload for {vector_id} "
|
||||
"after a failed vector upsert"
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -275,9 +275,116 @@ class TestNeptuneAnalyticsVectorInitValidation:
|
||||
def test_rejects_injection_payload_in_init(self, payload, monkeypatch):
|
||||
from mem0.vector_stores.neptune_analytics import NeptuneAnalyticsVector
|
||||
monkeypatch.setattr("mem0.vector_stores.neptune_analytics.NeptuneAnalyticsGraph", lambda *args, **kwargs: None)
|
||||
|
||||
|
||||
with pytest.raises(ValueError, match="Invalid collection_name"):
|
||||
NeptuneAnalyticsVector(
|
||||
endpoint="neptune-graph://test",
|
||||
collection_name=payload
|
||||
)
|
||||
|
||||
|
||||
class _FakeNeptuneGraph:
|
||||
"""Minimal stand-in for `NeptuneAnalyticsGraph.query()` so update()'s compensation
|
||||
path can be exercised without a real Neptune Analytics endpoint."""
|
||||
|
||||
def __init__(self):
|
||||
self.nodes = {}
|
||||
self.fail_next_upsert = False
|
||||
self.soft_fail_next_upsert = False
|
||||
self.get_call_count = 0
|
||||
|
||||
def query(self, query_string, params=None):
|
||||
params = params or {}
|
||||
|
||||
if "UNWIND $rows" in query_string:
|
||||
rows = params["rows"]
|
||||
if "CALL neptune.algo.vectors.upsert" in query_string:
|
||||
return [{"success": True} for _ in rows]
|
||||
for row in rows:
|
||||
self.nodes[row["node_id"]] = dict(row["properties"])
|
||||
return []
|
||||
|
||||
if "CALL neptune.algo.vectors.upsert" in query_string:
|
||||
if self.fail_next_upsert:
|
||||
self.fail_next_upsert = False
|
||||
raise RuntimeError("simulated Neptune upsert failure")
|
||||
if self.soft_fail_next_upsert:
|
||||
self.soft_fail_next_upsert = False
|
||||
return [{"success": False}]
|
||||
return [{"success": True}]
|
||||
|
||||
if "SET n = $properties" in query_string:
|
||||
self.nodes[params["vector_id"]] = dict(params["properties"])
|
||||
return []
|
||||
|
||||
if "RETURN n" in query_string and "node_id" in params:
|
||||
self.get_call_count += 1
|
||||
vector_id = params["node_id"]
|
||||
if vector_id not in self.nodes:
|
||||
return []
|
||||
return [{"n": {"~id": vector_id, "~properties": dict(self.nodes[vector_id])}}]
|
||||
|
||||
if "DETACH DELETE n" in query_string:
|
||||
self.nodes.pop(params.get("node_id"), None)
|
||||
return []
|
||||
|
||||
return []
|
||||
|
||||
|
||||
class TestNeptuneAnalyticsUpdateRollback:
|
||||
"""update() must not leave a payload committed against a stale embedding when the
|
||||
vector upsert step fails. See the compensation logic in `NeptuneAnalyticsVector.update()`."""
|
||||
|
||||
def _make_vec(self, monkeypatch):
|
||||
monkeypatch.setattr("mem0.vector_stores.neptune_analytics.NeptuneAnalyticsGraph", lambda *args, **kwargs: None)
|
||||
vec = NeptuneAnalyticsVector(endpoint="neptune-graph://test", collection_name="rollback")
|
||||
vec.graph = _FakeNeptuneGraph()
|
||||
return vec
|
||||
|
||||
def test_restores_prior_payload_when_upsert_fails(self, monkeypatch):
|
||||
vec = self._make_vec(monkeypatch)
|
||||
vec.insert(vectors=[[0.1, 0.2]], ids=["A"], payloads=[{"data": "alpha", "user_id": "u1"}])
|
||||
|
||||
vec.graph.fail_next_upsert = True
|
||||
with pytest.raises(RuntimeError):
|
||||
vec.update("A", vector=[0.9, 0.9], payload={"data": "beta", "user_id": "u1"})
|
||||
|
||||
restored = vec.get("A")
|
||||
assert restored.payload["data"] == "alpha"
|
||||
assert restored.payload["user_id"] == "u1"
|
||||
|
||||
def test_does_not_snapshot_prior_state_for_a_vector_only_update(self, monkeypatch):
|
||||
"""Only a combined payload+vector update can desync -- a vector-only update has
|
||||
nothing to roll back to, so it must skip the extra get() snapshot entirely."""
|
||||
vec = self._make_vec(monkeypatch)
|
||||
vec.insert(vectors=[[0.1, 0.2]], ids=["A"], payloads=[{"data": "alpha", "user_id": "u1"}])
|
||||
|
||||
vec.graph.fail_next_upsert = True
|
||||
calls_before = vec.graph.get_call_count
|
||||
with pytest.raises(RuntimeError):
|
||||
vec.update("A", vector=[0.9, 0.9])
|
||||
|
||||
assert vec.graph.get_call_count == calls_before
|
||||
|
||||
def test_succeeds_normally_when_upsert_does_not_fail(self, monkeypatch):
|
||||
vec = self._make_vec(monkeypatch)
|
||||
vec.insert(vectors=[[0.1, 0.2]], ids=["A"], payloads=[{"data": "alpha", "user_id": "u1"}])
|
||||
|
||||
vec.update("A", vector=[0.9, 0.9], payload={"data": "beta", "user_id": "u1"})
|
||||
|
||||
updated = vec.get("A")
|
||||
assert updated.payload["data"] == "beta"
|
||||
|
||||
def test_rolls_back_on_soft_upsert_failure(self, monkeypatch):
|
||||
"""A soft {"success": False} row desyncs the payload from the embedding just as much as a
|
||||
thrown error, so update() must treat it as a failure and roll the payload back too."""
|
||||
vec = self._make_vec(monkeypatch)
|
||||
vec.insert(vectors=[[0.1, 0.2]], ids=["A"], payloads=[{"data": "alpha", "user_id": "u1"}])
|
||||
|
||||
vec.graph.soft_fail_next_upsert = True
|
||||
with pytest.raises(RuntimeError):
|
||||
vec.update("A", vector=[0.9, 0.9], payload={"data": "beta", "user_id": "u1"})
|
||||
|
||||
restored = vec.get("A")
|
||||
assert restored.payload["data"] == "alpha"
|
||||
assert restored.payload["user_id"] == "u1"
|
||||
|
||||
Reference in New Issue
Block a user