diff --git a/docs/components/vectordbs/dbs/weaviate.mdx b/docs/components/vectordbs/dbs/weaviate.mdx index 2d24e8546..4b54cd0d4 100644 --- a/docs/components/vectordbs/dbs/weaviate.mdx +++ b/docs/components/vectordbs/dbs/weaviate.mdx @@ -4,14 +4,21 @@ description: "Use Weaviate as an open-source vector search engine in Mem0 for st --- [Weaviate](https://weaviate.io/) is an open-source vector search engine. It allows efficient storage and retrieval of high-dimensional vector embeddings, enabling powerful search and retrieval capabilities. - ### Installation -```bash + + +```bash Python pip install weaviate-client ``` +```bash TypeScript +npm install weaviate-client +``` + + ### Usage + ```python Python import os from mem0 import Memory @@ -33,20 +40,73 @@ m = Memory.from_config(config) messages = [ {"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"}, {"role": "assistant", "content": "How about a thriller movie? They can be quite engaging."}, - {"role": "user", "content": "I’m not a big fan of thriller movies but I love sci-fi movies."}, + {"role": "user", "content": "I'm not a big fan of thriller movies but I love sci-fi movies."}, {"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."} ] m.add(messages, user_id="alice", metadata={"category": "movies"}) ``` +```typescript TypeScript +import { Memory } from "mem0ai/oss"; + +const config = { + vectorStore: { + provider: "weaviate", + config: { + collectionName: "test", + embeddingModelDims: 1536, + clusterUrl: "http://localhost:8080", + }, + }, +}; + +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 a thriller movie? 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", + }, +}); +``` + + +The TypeScript SDK picks the connection mode from the config you pass: + +- `clusterUrl` pointing at `localhost` connects to a local instance. +- `clusterUrl` plus `apiKey` connects to a Weaviate Cloud cluster (for example `https://my-cluster.weaviate.cloud`). +- Any other `clusterUrl` without an `apiKey` connects to a custom deployment, using the host and port from the URL. + +You can also pass a pre-configured `client` (a `WeaviateClient` instance) to reuse an existing connection. + ### Config Here are the parameters available for configuring Weaviate: -| Parameter | Description | Default Value | -| --- | --- | --- | -| `collection_name` | The name of the collection to store the vectors | `mem0` | -| `embedding_model_dims` | Dimensions of the embedding model | `1536` | -| `cluster_url` | URL for the Weaviate server | `None` | -| `auth_client_secret` | API key for Weaviate authentication | `None` | -| `additional_headers` | Additional headers to include in requests (`Dict[str, str]`) | `None` | \ No newline at end of file +| Python | TypeScript | Description | Default Value | +| --- | --- | --- | --- | +| `collection_name` | `collectionName` | The name of the collection to store the vectors | `mem0` | +| `embedding_model_dims` | `embeddingModelDims` | Dimensions of the embedding model | `1536` | +| `cluster_url` | `clusterUrl` | URL for the Weaviate server | `None` | +| `auth_client_secret` | `apiKey` | API key for Weaviate authentication | `None` | +| `additional_headers` | `additionalHeaders` | Additional headers to include in requests | `None` | diff --git a/mem0-ts/package.json b/mem0-ts/package.json index 0362feea8..48e019f73 100644 --- a/mem0-ts/package.json +++ b/mem0-ts/package.json @@ -131,6 +131,7 @@ "fastembed": "^2.1.0", "groq-sdk": "0.3.0", "mongodb": "^7.0.0", + "weaviate-client": "^3.0.0", "ollama": "^0.5.14", "pg": "8.11.3", "redis": "^4.6.13", @@ -174,6 +175,7 @@ "path-to-regexp@>=8.0.0 <8.4.0": "^8.4.0", "postcss@<8.5.10": ">=8.5.10", "uuid@<11.1.1": ">=11.1.1", + "weaviate-client>uuid": "^11.1.1", "ws@>=8.0.0 <8.20.1": ">=8.20.1", "rollup@>=4.0.0 <4.59.0": "^4.59.0", "tar-fs@>=2.0.0 <2.1.4": "^2.1.4", diff --git a/mem0-ts/pnpm-lock.yaml b/mem0-ts/pnpm-lock.yaml index f8d41837d..4a8b9311c 100644 --- a/mem0-ts/pnpm-lock.yaml +++ b/mem0-ts/pnpm-lock.yaml @@ -17,6 +17,7 @@ overrides: path-to-regexp@>=8.0.0 <8.4.0: ^8.4.0 postcss@<8.5.10: '>=8.5.10' uuid@<11.1.1: '>=11.1.1' + weaviate-client>uuid: ^11.1.1 ws@>=8.0.0 <8.20.1: '>=8.20.1' rollup@>=4.0.0 <4.59.0: ^4.59.0 tar-fs@>=2.0.0 <2.1.4: ^2.1.4 @@ -136,6 +137,9 @@ importers: uuid: specifier: ^11.1.1 version: 11.1.1 + weaviate-client: + specifier: ^3.0.0 + version: 3.13.1 zod: specifier: ^3.24.1 version: 3.25.76 @@ -647,6 +651,9 @@ packages: '@dabh/diagnostics@2.0.8': resolution: {integrity: sha512-R4MSXTVnuMzGD7bzHdW2ZhhdPC/igELENcq5IjEverBvq5hn1SXCWcsi6eSsdWP0/Ur+SItRRjAktmdoX/8R/Q==} + '@datastructures-js/deque@1.0.8': + resolution: {integrity: sha512-PSBhJ2/SmeRPRHuBv7i/fHWIdSC3JTyq56qb+Rq0wjOagi0/fdV5/B/3Md5zFZus/W6OkSPMaxMKKMNMrSmubg==} + '@elastic/elasticsearch@9.4.2': resolution: {integrity: sha512-H9myMlLUeotkZhZ4ppinoMGDFxmW3lY8/s+4TIk1vFHyCvWU1Ej4T7azX5buCzemyFApgN0ywnEuvOtpel2VZg==} engines: {node: '>=20'} @@ -829,6 +836,11 @@ packages: '@modelcontextprotocol/sdk': optional: true + '@graphql-typed-document-node/core@3.2.0': + resolution: {integrity: sha512-mB9oAsNCm9aM3/SOv4YtBMqZbYj10R7dkq8byBqxGY/ncFwhf2oQzMV+LCRlWoDSEBJ3COiR1yeDvMtsoOsuFQ==} + peerDependencies: + graphql: ^0.8.0 || ^0.9.0 || ^0.10.0 || ^0.11.0 || ^0.12.0 || ^0.13.0 || ^14.0.0 || ^15.0.0 || ^16.0.0 || ^17.0.0 + '@grpc/grpc-js@1.14.4': resolution: {integrity: sha512-k9Dj3DV/itK9D06Y8f190Qgop7/Ui+D0njFV3LHMPwPT75DpXLQohE9Wmz0QElrJnzsjB7KPWiKJbOl7IPDArQ==} engines: {node: '>=12.10.0'} @@ -1562,6 +1574,9 @@ packages: '@zilliz/milvus2-sdk-node@3.0.3': resolution: {integrity: sha512-7rC+MDmzfctc9wscIfqxkzTJCxWafSdkyFDESMIzEX1G2ndImRQzHfkrzjcPFqPBVgQHQScGnGa3W7edqcThjA==} + abort-controller-x@0.5.0: + resolution: {integrity: sha512-yTt9CI0x+nRfX6BFMenEGP8ooPvErGH6AbFz20C2IeOLIlDsrw/VHpgne3GsCEuTA410IiFiaLVFKmgM4bKEPQ==} + abort-controller@3.0.0: resolution: {integrity: sha512-h8lQ8tacZYnR3vNQTgibj+tODHI5/+l06Au2Pcriv/Gmet0eaj4TwWH41sO9wnHDiQsEj19q0drzdWdeAHtweg==} engines: {node: '>=6.5'} @@ -1967,6 +1982,9 @@ packages: create-require@1.1.1: resolution: {integrity: sha512-dcKFX3jn0MpIaXjisoRvexIJVEKzaq7z2rZKxf+MSr9TkdmHmsU4m2lcLojrj/FHl8mk5VxMmYA+ftRkP/3oKQ==} + cross-fetch@3.2.0: + resolution: {integrity: sha512-Q+xVJLoGOeIMXZmbUK4HYk+69cQH6LudR0Vu/pRm2YlU/hDV9CiS0gKUMaWY5f2NeUH9C1nV3bsTlCo0FsTV1Q==} + cross-spawn@7.0.6: resolution: {integrity: sha512-uV2QOWP2nWzsy2aMp8aRibhi9dlzF5Hgh5SHaB9OiTGEyDTiJJyx0uy51QXdyWbtAHNua4XJzUKca3OzKUd3vA==} engines: {node: '>= 8'} @@ -2371,6 +2389,15 @@ packages: resolution: {integrity: sha512-rXunEHF9M9EkMydTBux7+IryYXEZinRk6g8OBOGDBzo/qWJjhTxy86i5q7lQYpCLHN8Sqv1XX3OIOc7ka2gtvQ==} engines: {node: '>=8.0.0'} + graphql-request@6.1.0: + resolution: {integrity: sha512-p+XPfS4q7aIpKVcgmnZKhMNqhltk20hfXtkaIkTfjjmiKMJ5xrt5c743cL03y/K7y1rg3WrIC49xGiEQ4mxdNw==} + peerDependencies: + graphql: 14 - 16 + + graphql@16.14.2: + resolution: {integrity: sha512-Chq1s4CY7jmh8gO2qvLIJyfCDIN+EHLFW/9iShnp1z8FjBQMoodWP1kDC36VAMXXIvAjj4ARa7ntfAV2BrjsbA==} + engines: {node: ^12.22.0 || ^14.16.0 || ^16.0.0 || >=17.0.0} + groq-sdk@0.3.0: resolution: {integrity: sha512-Cdgjh4YoSBE2X4S9sxPGXaAy1dlN4bRtAaDZ3cnq+XsxhhN9WSBeHF64l7LWwuD5ntmw7YC5Vf4Ff1oHCg1LOg==} @@ -3031,6 +3058,15 @@ packages: neo-async@2.6.2: resolution: {integrity: sha512-Yd3UES5mWCSqR+qNT93S3UoYUkqAZ9lLg8a7g9rimsWmYGK8cVToA4/sF3RrshdyV3sAGMXVUmpMYOw+dLpOuw==} + nice-grpc-client-middleware-retry@3.1.15: + resolution: {integrity: sha512-fXfNNdtjCQzc3O/w3WsK1AOU+xdE1V7FCm4FEQWS/UUVsV1616S2oryvWqsZAF8aOvuMir8lFtNmgH3DeP+PBg==} + + nice-grpc-common@2.0.3: + resolution: {integrity: sha512-MEhnD3JMah0mgyivpb9hpRDbOBuXBxI/TVO+OK1h6rC97WM42HsPMR+zzRNQ0C5BqYJTw1nyWiQRD0DucO+pjQ==} + + nice-grpc@2.1.16: + resolution: {integrity: sha512-Cl3Pn00212Hl8/U6bpgMxmhZj5lyv3nWoJov4cd3FjWarktrMHP4DNvSjCnDwkMWYx4W1tyscEia4JX6Y4GVCQ==} + node-abi@3.92.0: resolution: {integrity: sha512-KdHvFWZjEKDf0cakgFjebl371GPsISX2oZHcuyKqM7DtogIsHrqKeLTo8wBHxaXRAQlY2PsPlZmfo+9ZCxEREQ==} engines: {node: '>=10'} @@ -3753,6 +3789,9 @@ packages: resolution: {integrity: sha512-aZbgViZrg1QNcG+LULa7nhZpJTZSLm/mXnHXnbAbjmN5aSa0y7V+wvv6+4WaBtpISJzThKy+PIPxc1Nq1EJ9mg==} engines: {node: '>= 14.0.0'} + ts-error@1.0.6: + resolution: {integrity: sha512-tLJxacIQUM82IR7JO1UUkKlYuUTmoY9HBJAmNWFzheSlDS5SPMcNIepejHJa4BpPQLAcbRhRf3GDJzyj6rbKvA==} + ts-interface-checker@0.1.13: resolution: {integrity: sha512-Y/arvbn+rrz3JCKl9C4kVNfTfSm2/mEp5FSz5EsZSANGPSlQrpRI5M4PKF+mJnE52jOO90PnPSc3Ur3bTQw0gA==} @@ -3912,6 +3951,10 @@ packages: walker@1.0.8: resolution: {integrity: sha512-ts/8E8l5b7kY0vlWLewOkDXMmPdLcVV4GmOQLyxuSswIJsweeFZtAsMF7k1Nszz+TYBQrlYRmzOnr398y1JemQ==} + weaviate-client@3.13.1: + resolution: {integrity: sha512-cimmwR8w8GnSKQsP7tyW3dT2fMxJA2Tj+WqGzaVVkTVOLH4bb4B4iY0SO3TJUM4L8bcHkL04+MY2RAWEWJpU6A==} + engines: {node: '>=22.0.0'} + web-streams-polyfill@3.3.3: resolution: {integrity: sha512-d2JWLCivmZYTSIoge9MsgFCZrt571BikcWGYkjC1khllbTeDlGqZ2D8vD8E/lJa8WGWbb7Plm8/XJYV7IJHZZw==} engines: {node: '>= 8'} @@ -4932,6 +4975,8 @@ snapshots: enabled: 2.0.0 kuler: 2.0.0 + '@datastructures-js/deque@1.0.8': {} + '@elastic/elasticsearch@9.4.2': dependencies: '@elastic/transport': 9.3.7 @@ -5047,6 +5092,10 @@ snapshots: - supports-color - utf-8-validate + '@graphql-typed-document-node/core@3.2.0(graphql@16.14.2)': + dependencies: + graphql: 16.14.2 + '@grpc/grpc-js@1.14.4': dependencies: '@grpc/proto-loader': 0.8.1 @@ -5936,6 +5985,8 @@ snapshots: - bufferutil - utf-8-validate + abort-controller-x@0.5.0: {} + abort-controller@3.0.0: dependencies: event-target-shim: 5.0.1 @@ -6337,6 +6388,12 @@ snapshots: create-require@1.1.1: {} + cross-fetch@3.2.0: + dependencies: + node-fetch: 2.7.0 + transitivePeerDependencies: + - encoding + cross-spawn@7.0.6: dependencies: path-key: 3.1.1 @@ -6781,6 +6838,16 @@ snapshots: grad-school@0.0.5: {} + graphql-request@6.1.0(graphql@16.14.2): + dependencies: + '@graphql-typed-document-node/core': 3.2.0(graphql@16.14.2) + cross-fetch: 3.2.0 + graphql: 16.14.2 + transitivePeerDependencies: + - encoding + + graphql@16.14.2: {} + groq-sdk@0.3.0: dependencies: '@types/node': 18.19.130 @@ -7614,6 +7681,21 @@ snapshots: neo-async@2.6.2: {} + nice-grpc-client-middleware-retry@3.1.15: + dependencies: + abort-controller-x: 0.5.0 + nice-grpc-common: 2.0.3 + + nice-grpc-common@2.0.3: + dependencies: + ts-error: 1.0.6 + + nice-grpc@2.1.16: + dependencies: + '@grpc/grpc-js': 1.14.4 + abort-controller-x: 0.5.0 + nice-grpc-common: 2.0.3 + node-abi@3.92.0: dependencies: semver: 7.8.4 @@ -8351,6 +8433,8 @@ snapshots: triple-beam@1.4.1: {} + ts-error@1.0.6: {} + ts-interface-checker@0.1.13: {} ts-jest@29.4.11(@babel/core@7.29.7)(@jest/transform@29.7.0)(@jest/types@29.6.3)(babel-jest@29.7.0(@babel/core@7.29.7))(esbuild@0.28.1)(jest-util@29.7.0)(jest@29.7.0(@types/node@22.19.21)(ts-node@10.9.2(@types/node@22.19.21)(typescript@5.5.4)))(typescript@5.5.4): @@ -8491,6 +8575,20 @@ snapshots: dependencies: makeerror: 1.0.12 + weaviate-client@3.13.1: + dependencies: + '@datastructures-js/deque': 1.0.8 + abort-controller-x: 0.5.0 + graphql: 16.14.2 + graphql-request: 6.1.0(graphql@16.14.2) + long: 5.3.2 + nice-grpc: 2.1.16 + nice-grpc-client-middleware-retry: 3.1.15 + nice-grpc-common: 2.0.3 + uuid: 11.1.1 + transitivePeerDependencies: + - encoding + web-streams-polyfill@3.3.3: {} web-streams-polyfill@4.0.0-beta.3: {} diff --git a/mem0-ts/src/oss/src/index.ts b/mem0-ts/src/oss/src/index.ts index 031af0fbf..8b772fbe1 100644 --- a/mem0-ts/src/oss/src/index.ts +++ b/mem0-ts/src/oss/src/index.ts @@ -44,4 +44,5 @@ export * from "./vector_stores/turbopuffer"; export * from "./vector_stores/milvus"; export * from "./vector_stores/mongodb"; export * from "./vector_stores/opensearch"; +export * from "./vector_stores/weaviate"; export * from "./utils/factory"; diff --git a/mem0-ts/src/oss/src/utils/factory.ts b/mem0-ts/src/oss/src/utils/factory.ts index 1b02978d4..50f487e23 100644 --- a/mem0-ts/src/oss/src/utils/factory.ts +++ b/mem0-ts/src/oss/src/utils/factory.ts @@ -58,6 +58,7 @@ import { S3Vectors } from "../vector_stores/s3_vectors"; import { TurbopufferDB } from "../vector_stores/turbopuffer"; import { Milvus } from "../vector_stores/milvus"; import { MongoDB } from "../vector_stores/mongodb"; +import { WeaviateDB } from "../vector_stores/weaviate"; export class EmbedderFactory { static create(provider: string, config: EmbeddingConfig): Embedder { @@ -177,6 +178,8 @@ export class VectorStoreFactory { return new Milvus(config as any); case "mongodb": return new MongoDB(config as any); + case "weaviate": + return new WeaviateDB(config as any); default: throw new Error(`Unsupported vector store provider: ${provider}`); } diff --git a/mem0-ts/src/oss/src/vector_stores/weaviate.ts b/mem0-ts/src/oss/src/vector_stores/weaviate.ts new file mode 100644 index 000000000..994a92da4 --- /dev/null +++ b/mem0-ts/src/oss/src/vector_stores/weaviate.ts @@ -0,0 +1,217 @@ +import weaviate, { Filters, type WeaviateClient } from "weaviate-client"; +import { v4 as uuidv4 } from "uuid"; +import { VectorStore } from "./base"; +import { SearchFilters, VectorStoreConfig, VectorStoreResult } from "../types"; + +interface WeaviateConfig extends VectorStoreConfig { + client?: WeaviateClient; + clusterUrl?: string; + apiKey?: string; + additionalHeaders?: Record; + collectionName: string; + embeddingModelDims: number; +} + +const RETURN_PROPERTIES = [ + "ids", + "hash", + "metadata", + "data", + "created_at", + "category", + "updated_at", + "user_id", + "agent_id", + "run_id", +]; + +export class WeaviateDB implements VectorStore { + private _config: WeaviateConfig; + private _client!: WeaviateClient; + private _col!: any; + private _userId: string; + private _initPromise?: Promise; + + constructor(config: WeaviateConfig) { + this._config = config; + this._userId = ""; + this.initialize().catch(console.error); + } + + initialize(): Promise { + return (this._initPromise ??= this._doInitialize()); + } + + private async _doInitialize(): Promise { + const { client, clusterUrl, apiKey, additionalHeaders, collectionName } = + this._config; + + if (client) { + this._client = client; + } else if (clusterUrl?.includes("localhost")) { + this._client = await weaviate.connectToLocal({ + headers: additionalHeaders, + }); + } else if (apiKey) { + this._client = await weaviate.connectToWeaviateCloud(clusterUrl!, { + authCredentials: new weaviate.ApiKey(apiKey), + headers: additionalHeaders, + }); + } else { + if (!clusterUrl) { + throw new Error( + "WeaviateDB: clusterUrl is required when client and apiKey are not provided", + ); + } + const parsed = new URL(clusterUrl); + const httpSecure = parsed.protocol === "https:"; + this._client = await weaviate.connectToCustom({ + httpHost: parsed.hostname, + httpPort: parsed.port + ? parseInt(parsed.port, 10) + : httpSecure + ? 443 + : 8080, + httpSecure, + grpcHost: parsed.hostname, + grpcPort: 50051, + grpcSecure: false, + headers: additionalHeaders, + }); + } + + const exists = await this._client.collections.exists(collectionName); + if (!exists) { + await this._client.collections.create({ + name: collectionName, + properties: RETURN_PROPERTIES.map((name) => ({ + name, + dataType: "text" as const, + })), + vectorizers: weaviate.configure.vectorizer.none(), + vectorIndex: weaviate.configure.vectorIndex.hnsw(), + } as any); + } + + this._col = this._client.collections.get(collectionName); + } + + private _buildFilters(filters?: SearchFilters) { + if (!filters) return undefined; + const conditions = (["user_id", "agent_id", "run_id"] as const) + .filter((key) => filters[key] != null) + .map((key) => this._col.filter.byProperty(key).equal(filters[key])); + return conditions.length ? Filters.and(...conditions) : undefined; + } + + async insert( + vectors: number[][], + ids: string[], + payloads: Record[], + ): Promise { + await this.initialize(); + const objects = vectors.map((vector, i) => ({ + id: ids[i], + properties: payloads[i], + vectors: vector, + })); + await this._col.data.insertMany(objects); + } + + async search( + query: number[], + topK?: number, + filters?: SearchFilters, + ): Promise { + await this.initialize(); + const result = await this._col.query.nearVector(query, { + limit: topK ?? 10, + filters: this._buildFilters(filters), + returnMetadata: ["distance"], + }); + return result.objects.map((obj: any) => ({ + id: obj.uuid, + payload: obj.properties, + score: 1 - obj.metadata.distance, + })); + } + + async keywordSearch( + query: string, + topK?: number, + filters?: SearchFilters, + ): Promise { + await this.initialize(); + const result = await this._col.query.bm25(query, { + queryProperties: ["data"], + limit: topK ?? 10, + filters: this._buildFilters(filters), + returnMetadata: ["score"], + }); + return result.objects.map((obj: any) => ({ + id: obj.uuid, + payload: obj.properties, + score: obj.metadata.score, + })); + } + + async get(vectorId: string): Promise { + await this.initialize(); + const obj = await this._col.query.fetchObjectById(vectorId, { + returnProperties: RETURN_PROPERTIES, + }); + if (!obj) return null; + return { id: obj.uuid, payload: obj.properties }; + } + + async update( + vectorId: string, + vector: number[], + payload: Record, + ): Promise { + await this.initialize(); + await this._col.data.update({ + id: vectorId, + properties: payload, + vectors: vector, + }); + } + + async delete(vectorId: string): Promise { + await this.initialize(); + await this._col.data.deleteById(vectorId); + } + + async deleteCol(): Promise { + await this.initialize(); + await this._client.collections.delete(this._config.collectionName); + } + + async list( + filters?: SearchFilters, + topK?: number, + ): Promise<[VectorStoreResult[], number]> { + await this.initialize(); + const result = await this._col.query.fetchObjects({ + limit: topK ?? 100, + filters: this._buildFilters(filters), + returnProperties: RETURN_PROPERTIES, + }); + const results = result.objects.map((obj: any) => ({ + id: obj.uuid, + payload: obj.properties, + })); + return [results, results.length]; + } + + async getUserId(): Promise { + if (!this._userId) { + this._userId = uuidv4(); + } + return this._userId; + } + + async setUserId(userId: string): Promise { + this._userId = userId; + } +} diff --git a/mem0-ts/src/oss/tests/factory.unit.test.ts b/mem0-ts/src/oss/tests/factory.unit.test.ts index c4d684106..af59c47c5 100644 --- a/mem0-ts/src/oss/tests/factory.unit.test.ts +++ b/mem0-ts/src/oss/tests/factory.unit.test.ts @@ -193,6 +193,11 @@ jest.mock("../src/vector_stores/s3_vectors", () => ({ .fn() .mockImplementation((config) => ({ type: "s3-vectors", config })), })); +jest.mock("../src/vector_stores/weaviate", () => ({ + WeaviateDB: jest + .fn() + .mockImplementation((config) => ({ type: "weaviate", config })), +})); jest.mock("../src/storage/SupabaseHistoryManager", () => ({ SupabaseHistoryManager: jest .fn() @@ -328,6 +333,7 @@ describe("VectorStoreFactory", () => { ["cassandra"], ["s3-vectors"], ["s3_vectors"], + ["weaviate"], ])("creates vector store for provider '%s'", (provider) => { expect(() => VectorStoreFactory.create(provider, dummyVSConfig), diff --git a/mem0-ts/src/oss/tests/vector-stores-compat.test.ts b/mem0-ts/src/oss/tests/vector-stores-compat.test.ts index 87d1057fa..39a259856 100644 --- a/mem0-ts/src/oss/tests/vector-stores-compat.test.ts +++ b/mem0-ts/src/oss/tests/vector-stores-compat.test.ts @@ -2281,3 +2281,152 @@ describe("Memory class – backward compat with all providers", () => { consoleSpy.mockRestore(); }); }); + +// ─────────────────────────────────────────────────────────────────────────── +// WeaviateDB — mock client, behavioral surface checks +// ─────────────────────────────────────────────────────────────────────────── +describe("WeaviateDB – backward compat with mocked client", () => { + let WeaviateDB: any; + let mockClient: any; + let mockCol: any; + + beforeEach(() => { + jest.resetModules(); + + mockCol = { + data: { + insertMany: jest.fn().mockResolvedValue({}), + deleteById: jest.fn().mockResolvedValue({}), + update: jest.fn().mockResolvedValue({}), + }, + query: { + nearVector: jest.fn().mockResolvedValue({ objects: [] }), + bm25: jest.fn().mockResolvedValue({ objects: [] }), + fetchObjectById: jest.fn().mockResolvedValue(null), + fetchObjects: jest.fn().mockResolvedValue({ objects: [] }), + }, + filter: { + byProperty: jest + .fn() + .mockReturnValue({ equal: jest.fn().mockReturnValue({}) }), + }, + }; + + mockClient = { + collections: { + exists: jest.fn().mockResolvedValue(false), + create: jest.fn().mockResolvedValue({}), + get: jest.fn().mockReturnValue(mockCol), + delete: jest.fn().mockResolvedValue({}), + }, + }; + + jest.doMock("weaviate-client", () => ({ + default: { + connectToLocal: jest.fn().mockResolvedValue(mockClient), + connectToWeaviateCloud: jest.fn().mockResolvedValue(mockClient), + connectToCustom: jest.fn().mockResolvedValue(mockClient), + ApiKey: jest.fn().mockReturnValue({}), + configure: { + vectorizer: { none: jest.fn().mockReturnValue({}) }, + vectorIndex: { hnsw: jest.fn().mockReturnValue({}) }, + }, + }, + Filters: { and: jest.fn().mockReturnValue({ __mock: "filter" }) }, + __esModule: true, + })); + + WeaviateDB = require("../src/vector_stores/weaviate").WeaviateDB; + }); + + afterEach(() => { + jest.restoreAllMocks(); + jest.resetModules(); + }); + + it("implements full VectorStore interface", () => { + const store = new WeaviateDB({ + client: mockClient, + collectionName: "test", + embeddingModelDims: 768, + }); + expect(typeof store.insert).toBe("function"); + expect(typeof store.search).toBe("function"); + expect(typeof store.keywordSearch).toBe("function"); + expect(typeof store.get).toBe("function"); + expect(typeof store.update).toBe("function"); + expect(typeof store.delete).toBe("function"); + expect(typeof store.deleteCol).toBe("function"); + expect(typeof store.list).toBe("function"); + expect(typeof store.getUserId).toBe("function"); + expect(typeof store.setUserId).toBe("function"); + expect(typeof store.initialize).toBe("function"); + }); + + it("initialize() is idempotent (same promise returned)", async () => { + const store = new WeaviateDB({ + client: mockClient, + collectionName: "test", + embeddingModelDims: 768, + }); + const p1 = store.initialize(); + const p2 = store.initialize(); + const p3 = store.initialize(); + await Promise.all([p1, p2, p3]); + expect(mockClient.collections.create).toHaveBeenCalledTimes(1); + }); + + it("insert shapes insertMany request correctly", async () => { + const store = new WeaviateDB({ + client: mockClient, + collectionName: "test", + embeddingModelDims: 3, + }); + await store.initialize(); + await store.insert([[0.1, 0.2, 0.3]], ["id-1"], [{ data: "hello" }]); + expect(mockCol.data.insertMany).toHaveBeenCalledWith( + expect.arrayContaining([ + expect.objectContaining({ + id: "id-1", + properties: { data: "hello" }, + vectors: [0.1, 0.2, 0.3], + }), + ]), + ); + }); + + it("search normalizes nearVector result to id/payload/score", async () => { + mockCol.query.nearVector.mockResolvedValue({ + objects: [ + { + uuid: "id-1", + properties: { data: "x" }, + metadata: { distance: 0.2 }, + }, + ], + }); + const store = new WeaviateDB({ + client: mockClient, + collectionName: "test", + embeddingModelDims: 3, + }); + await store.initialize(); + const results = await store.search([0.1, 0.2, 0.3], 1); + expect(results).toHaveLength(1); + expect(results[0]).toMatchObject({ + id: "id-1", + payload: { data: "x" }, + score: 0.8, + }); + }); + + it("getUserId / setUserId roundtrip", async () => { + const store = new WeaviateDB({ + client: mockClient, + collectionName: "test", + embeddingModelDims: 768, + }); + await store.setUserId("custom-user"); + expect(await store.getUserId()).toBe("custom-user"); + }); +}); diff --git a/mem0-ts/tsup.config.ts b/mem0-ts/tsup.config.ts index a7198c2fa..d79558a09 100644 --- a/mem0-ts/tsup.config.ts +++ b/mem0-ts/tsup.config.ts @@ -36,6 +36,7 @@ const external = [ "@opensearch-project/opensearch", "@elastic/elasticsearch", "chromadb", + "weaviate-client", ]; const define = {