feat(ts-sdk): add Vertex AI embedding provider (#5882)

Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
Fahmid Arman
2026-07-10 00:26:46 +06:00
committed by GitHub
parent 573b20cec8
commit 3b2357bfe0
14 changed files with 963 additions and 30 deletions
+99 -15
View File
@@ -4,11 +4,36 @@ description: "Configure Google Cloud Vertex AI as an embedding provider in Mem0
---
### Vertex AI
To use Google Cloud's Vertex AI for text embedding models, set the `GOOGLE_APPLICATION_CREDENTIALS` environment variable to point to the path of your service account's credentials JSON file. These credentials can be created in the [Google Cloud Console](https://console.cloud.google.com/).
Google Cloud's Vertex AI serves text embedding models such as `gemini-embedding-001`. Mem0 uses them through the provider's own SDK, which you install alongside Mem0.
### Installation
The Vertex AI client is an optional dependency, so install it yourself.
<CodeGroup>
```bash Python
pip install vertexai
```
```bash TypeScript
npm install @google-cloud/aiplatform
```
</CodeGroup>
### Authentication
Both SDKs authenticate with [Application Default Credentials](https://cloud.google.com/docs/authentication/application-default-credentials). Pick whichever fits your environment:
- **Local development:** run `gcloud auth application-default login`.
- **Service account:** create a key in the [Google Cloud Console](https://console.cloud.google.com/) and point `GOOGLE_APPLICATION_CREDENTIALS` at the JSON file, or pass its path through the embedder config.
- **Google Cloud runtimes** (Cloud Run, GKE, Compute Engine): the attached service account is picked up automatically.
The TypeScript SDK reads the project ID from `googleProjectId`, then the `GCP_PROJECT_ID`, `GOOGLE_CLOUD_PROJECT`, and `GCLOUD_PROJECT` environment variables, and finally from your credentials. Set it explicitly when your credentials cover more than one project.
### Usage
```python
<CodeGroup>
```python Python
import os
from mem0 import Memory
@@ -32,28 +57,87 @@ m = Memory.from_config(config)
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": "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="john")
```
The embedding types can be one of the following:
```typescript TypeScript
import { Memory } from "mem0ai/oss";
const config = {
embedder: {
provider: "vertexai",
config: {
model: "gemini-embedding-001",
// Optional. Falls back to GCP_PROJECT_ID / GOOGLE_CLOUD_PROJECT /
// GCLOUD_PROJECT, then to the project on your credentials.
googleProjectId: process.env.GCP_PROJECT_ID,
location: "us-central1",
// Optional. Path to a service account key file, or pass the JSON inline
// via googleServiceAccountJson.
vertexCredentialsJson: "/path/to/your/credentials.json",
embeddingDims: 256,
memoryAddEmbeddingType: "RETRIEVAL_DOCUMENT",
memoryUpdateEmbeddingType: "RETRIEVAL_DOCUMENT",
memorySearchEmbeddingType: "RETRIEVAL_QUERY",
},
},
};
const memory = new Memory(config);
await memory.add("I love sci-fi movies but not thrillers", { userId: "john" });
```
</CodeGroup>
### Embedding types
Vertex AI embeds the same text differently depending on the task you declare. The embedding types can be one of the following:
- SEMANTIC_SIMILARITY
- CLASSIFICATION
- CLUSTERING
- RETRIEVAL_DOCUMENT, RETRIEVAL_QUERY, QUESTION_ANSWERING, FACT_VERIFICATION
- CODE_RETRIEVAL_QUERY
Check out the [Vertex AI documentation](https://cloud.google.com/vertex-ai/generative-ai/docs/embeddings/task-types#supported_task_types) for more information.
- CODE_RETRIEVAL_QUERY
Check out the [Vertex AI documentation](https://cloud.google.com/vertex-ai/generative-ai/docs/embeddings/task-types#supported_task_types) for more information.
<Note>
These embedding types map to the add, update, and search memory actions in both the Python and TypeScript SDKs. Stored memories use the add or update type, and searches use the search type.
</Note>
### Choosing a model
<Warning>
`gemini-embedding-001` accepts **one input text per request**. When Mem0 embeds several texts at once, such as the memories extracted from a single conversation turn, it issues one request per text. The older `text-embedding-005` and `text-multilingual-embedding-002` models accept up to 250 texts per request, so they are faster and cheaper for large batches. See [Get text embeddings](https://cloud.google.com/vertex-ai/generative-ai/docs/embeddings/get-text-embeddings).
</Warning>
### Config
Here are the parameters available for configuring the Vertex AI embedder:
| Parameter | Description | Default Value |
| ------------------------- | ------------------------------------------------ | -------------------- |
| `model` | The name of the Vertex AI embedding model to use | `gemini-embedding-001` |
| `vertex_credentials_json` | Path to the Google Cloud credentials JSON file | `None` |
| `embedding_dims` | Dimensions of the embedding model | `256` |
| `memory_add_embedding_type` | The type of embedding to use for the add memory action | `RETRIEVAL_DOCUMENT` |
| `memory_update_embedding_type` | The type of embedding to use for the update memory action | `RETRIEVAL_DOCUMENT` |
| `memory_search_embedding_type` | The type of embedding to use for the search memory action | `RETRIEVAL_QUERY` |
<Tabs>
<Tab title="Python">
| Parameter | Description | Default Value |
| -------------------------------- | ---------------------------------------------------------- | ---------------------- |
| `model` | The name of the Vertex AI embedding model to use | `gemini-embedding-001` |
| `vertex_credentials_json` | Path to the Google Cloud credentials JSON file | `None` |
| `embedding_dims` | Dimensions of the embedding model | `256` |
| `memory_add_embedding_type` | The embedding type to use for the add memory action | `RETRIEVAL_DOCUMENT` |
| `memory_update_embedding_type` | The embedding type to use for the update memory action | `RETRIEVAL_DOCUMENT` |
| `memory_search_embedding_type` | The embedding type to use for the search memory action | `RETRIEVAL_QUERY` |
</Tab>
<Tab title="TypeScript">
| Parameter | Description | Default Value |
| ----------------------------- | -------------------------------------------------------------------------- | ---------------------- |
| `model` | The name of the Vertex AI embedding model to use | `gemini-embedding-001` |
| `googleProjectId` | Google Cloud project ID (falls back to `GCP_PROJECT_ID` env var, then to your credentials) | Resolved from credentials |
| `location` | Google Cloud region (falls back to `GCP_LOCATION` env var) | `us-central1` |
| `vertexCredentialsJson` | Path to the Google Cloud credentials JSON file | `None` |
| `googleServiceAccountJson` | Service account credentials as a JSON string or object | `None` |
| `embeddingDims` | Dimensions of the embedding model | `256` |
| `memoryAddEmbeddingType` | The embedding type to use for the add memory action | `RETRIEVAL_DOCUMENT` |
| `memoryUpdateEmbeddingType` | The embedding type to use for the update memory action | `RETRIEVAL_DOCUMENT` |
| `memorySearchEmbeddingType` | The embedding type to use for the search memory action | `RETRIEVAL_QUERY` |
</Tab>
</Tabs>
+4
View File
@@ -40,6 +40,10 @@ export class ConfigManager {
| undefined);
return {
// Spread first so provider-specific keys (e.g. the Vertex AI
// project/location/credentials) survive the merge, while the
// normalized values below still win.
...userConf,
apiKey:
userConf?.apiKey !== undefined
? userConf.apiKey
+8 -2
View File
@@ -1,4 +1,10 @@
export interface Embedder {
embed(text: string): Promise<number[]>;
embedBatch(texts: string[]): Promise<number[][]>;
embed(
text: string,
memoryAction?: "add" | "update" | "search",
): Promise<number[]>;
embedBatch(
texts: string[],
memoryAction?: "add" | "update" | "search",
): Promise<number[][]>;
}
+254
View File
@@ -0,0 +1,254 @@
import type { PredictionServiceClient } from "@google-cloud/aiplatform";
import { Embedder } from "./base";
import { VertexAIConfig } from "../types";
type AIPlatform = typeof import("@google-cloud/aiplatform");
type ClientOptions = NonNullable<
ConstructorParameters<AIPlatform["PredictionServiceClient"]>[0]
>;
interface EmbeddingResponse {
embeddings: {
values: number[];
};
}
/**
* Vertex AI caps how many input texts one `predict()` call may carry, and the
* cap depends on the model family. `gemini-embedding-*` accepts exactly one
* text per request; the older `text-embedding-*` / `text-multilingual-*`
* models accept up to 250.
* https://cloud.google.com/vertex-ai/generative-ai/docs/embeddings/get-text-embeddings
*/
function maxInstancesPerRequest(model: string): number {
return model.startsWith("gemini-embedding") ? 1 : 250;
}
function isValidEmbedding(value: unknown): value is EmbeddingResponse {
if (typeof value !== "object" || value === null) return false;
const obj = value as Record<string, unknown>;
if (typeof obj.embeddings !== "object" || obj.embeddings === null)
return false;
const embeddings = obj.embeddings as Record<string, unknown>;
const values = embeddings.values;
return (
Array.isArray(values) &&
values.every((v) => typeof v === "number" && Number.isFinite(v))
);
}
export class VertexAIEmbedder implements Embedder {
private client: PredictionServiceClient | undefined;
private helpers: AIPlatform["helpers"] | undefined;
private initPromise: Promise<void> | undefined;
private clientOptions: ClientOptions;
private model: string;
private embeddingDims: number;
private location: string;
private projectId: string;
private embeddingTypes: {
add: string;
update: string;
search: string;
};
constructor(config: VertexAIConfig) {
this.model = config.model || "gemini-embedding-001";
this.embeddingDims = config.embeddingDims || 256;
this.location =
config.location || process.env.GCP_LOCATION || "us-central1";
// Left empty when unset: initClient() resolves it from Application Default
// Credentials or the service account key file, the way the Python SDK does.
this.projectId =
config.googleProjectId ||
process.env.GCP_PROJECT_ID ||
process.env.GOOGLE_CLOUD_PROJECT ||
process.env.GCLOUD_PROJECT ||
"";
this.embeddingTypes = {
add: config.memoryAddEmbeddingType || "RETRIEVAL_DOCUMENT",
update: config.memoryUpdateEmbeddingType || "RETRIEVAL_DOCUMENT",
search: config.memorySearchEmbeddingType || "RETRIEVAL_QUERY",
};
const endpoint = `${this.location}-aiplatform.googleapis.com`;
this.clientOptions = { apiEndpoint: endpoint };
if (config.vertexCredentialsJson) {
this.clientOptions.keyFilename = config.vertexCredentialsJson;
} else if (config.googleServiceAccountJson) {
try {
this.clientOptions.credentials =
typeof config.googleServiceAccountJson === "string"
? JSON.parse(config.googleServiceAccountJson)
: config.googleServiceAccountJson;
} catch (err) {
throw new Error(
"Failed to parse googleServiceAccountJson: " + (err as Error).message,
);
}
}
}
private async initClient(): Promise<void> {
// Memoized so concurrent embed() calls share one client instead of each
// racing to build (and leak) their own gRPC channel.
if (!this.initPromise) {
this.initPromise = this.createClient().catch((err) => {
this.initPromise = undefined;
throw err;
});
}
await this.initPromise;
}
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 client = new aiplatform.PredictionServiceClient(this.clientOptions);
if (!this.projectId) {
try {
this.projectId = await client.getProjectId();
} catch (err) {
throw new Error(
"Vertex AI could not determine a Google Cloud project ID. Set googleProjectId in config, " +
"one of the GCP_PROJECT_ID / GOOGLE_CLOUD_PROJECT / GCLOUD_PROJECT env vars, or configure " +
"Application Default Credentials: " +
(err as Error).message,
);
}
}
this.client = client;
this.helpers = aiplatform.helpers;
}
private endpoint(): string {
return `projects/${this.projectId}/locations/${this.location}/publishers/google/models/${this.model}`;
}
private formatInstance(text: string, taskType: string) {
// task_type must live on the instance (snake_case), not in `parameters`.
// Vertex silently ignores an unknown `parameters.taskType`, which would
// fall back to the model's default task type. This mirrors the Python SDK's
// TextEmbeddingInput(text=..., task_type=...).
return {
content: text,
task_type: taskType,
};
}
async embed(
text: string,
memoryAction?: "add" | "update" | "search",
): Promise<number[]> {
await this.initClient();
if (!this.client || !this.helpers) {
throw new Error("Client not initialized");
}
let embeddingType = "SEMANTIC_SIMILARITY";
if (memoryAction !== undefined) {
if (!(memoryAction in this.embeddingTypes)) {
throw new Error(`Invalid memory action: ${memoryAction}`);
}
embeddingType = this.embeddingTypes[memoryAction];
}
const instance = this.formatInstance(text, embeddingType);
const parameters = {
outputDimensionality: this.embeddingDims,
};
const [response] = await this.client.predict({
endpoint: this.endpoint(),
instances: [this.helpers.toValue(instance) as any],
parameters: this.helpers.toValue(parameters) as any,
});
if (!response.predictions || response.predictions.length === 0) {
throw new Error("No predictions returned from Vertex AI");
}
const decoded = this.helpers.fromValue(response.predictions[0] as any);
if (!isValidEmbedding(decoded)) {
throw new Error("Failed to extract embedding values from response");
}
return decoded.embeddings.values;
}
async embedBatch(
texts: string[],
memoryAction: "add" | "update" | "search" = "add",
): Promise<number[][]> {
if (!texts || texts.length === 0) {
return [];
}
await this.initClient();
if (!this.client || !this.helpers) {
throw new Error("Client not initialized");
}
if (!(memoryAction in this.embeddingTypes)) {
throw new Error(`Invalid memory action: ${memoryAction}`);
}
const embeddingType = this.embeddingTypes[memoryAction];
const allEmbeddings: number[][] = [];
const batchSize = maxInstancesPerRequest(this.model);
for (let i = 0; i < texts.length; i += batchSize) {
const chunk = texts.slice(i, i + batchSize);
const instances = chunk.map(
(text) =>
this.helpers!.toValue(
this.formatInstance(text, embeddingType),
) as any,
);
const parameters = {
outputDimensionality: this.embeddingDims,
};
const [response] = await this.client.predict({
endpoint: this.endpoint(),
instances,
parameters: this.helpers.toValue(parameters) as any,
});
if (!response.predictions || response.predictions.length === 0) {
throw new Error("No predictions returned from Vertex AI batch request");
}
for (const prediction of response.predictions) {
const decoded = this.helpers.fromValue(prediction as any);
if (!isValidEmbedding(decoded)) {
throw new Error(
"Failed to extract embedding values from batch response",
);
}
allEmbeddings.push(decoded.embeddings.values);
}
}
if (allEmbeddings.length !== texts.length) {
throw new Error(
`Vertex AI embedBatch() returned ${allEmbeddings.length} embeddings for ${texts.length} texts using model '${this.model}'`,
);
}
return allEmbeddings;
}
}
+1
View File
@@ -10,6 +10,7 @@ export * from "./embeddings/together";
export * from "./embeddings/google";
export * from "./embeddings/azure";
export * from "./embeddings/langchain";
export * from "./embeddings/vertexai";
export * from "./embeddings/fastembed";
export * from "./llms/base";
export * from "./llms/openai";
+16 -12
View File
@@ -408,7 +408,7 @@ export class Memory {
}
let vec: number[];
try {
vec = await this.embedder.embed(entityText);
vec = await this.embedder.embed(entityText, "update");
} catch (e) {
console.debug(`Entity re-embed failed for '${entityText}': ${e}`);
continue;
@@ -452,7 +452,7 @@ export class Memory {
try {
let entityVec: number[];
try {
entityVec = await this.embedder.embed(entity.text);
entityVec = await this.embedder.embed(entity.text, "add");
} catch (e) {
console.debug(`Entity embed failed for '${entity.text}': ${e}`);
continue;
@@ -827,7 +827,7 @@ export class Memory {
.join("\n");
// Phase 1: Existing memory retrieval
const queryEmbedding = await this.embedder.embed(parsedMessages);
const queryEmbedding = await this.embedder.embed(parsedMessages, "search");
const existingResults = await this.vectorStore.search(
queryEmbedding,
10,
@@ -921,7 +921,7 @@ export class Memory {
.filter((t) => t.length > 0);
let embedMap: Record<string, number[]> = {};
try {
const memEmbeddingsList = await this.embedder.embedBatch(memTexts);
const memEmbeddingsList = await this.embedder.embedBatch(memTexts, "add");
for (let i = 0; i < memTexts.length; i++) {
embedMap[memTexts[i]] = memEmbeddingsList[i];
}
@@ -929,7 +929,7 @@ export class Memory {
// Fallback: embed individually
for (const text of memTexts) {
try {
embedMap[text] = await this.embedder.embed(text);
embedMap[text] = await this.embedder.embed(text, "add");
} catch (e) {
console.warn(`Failed to embed memory text: ${e}`);
}
@@ -1107,13 +1107,13 @@ export class Memory {
// 7b: Single batch embed for all unique entities
let entityEmbeddings: (number[] | null)[];
try {
entityEmbeddings = await this.embedder.embedBatch(entityTexts);
entityEmbeddings = await this.embedder.embedBatch(entityTexts, "add");
} catch {
// Fallback: embed individually
entityEmbeddings = [];
for (const t of entityTexts) {
try {
entityEmbeddings.push(await this.embedder.embed(t));
entityEmbeddings.push(await this.embedder.embed(t, "add"));
} catch {
entityEmbeddings.push(null);
}
@@ -1377,7 +1377,7 @@ export class Memory {
const queryEntities = extractEntities(query);
// Step 2: Embed query
const queryEmbedding = await this.embedder.embed(query);
const queryEmbedding = await this.embedder.embed(query, "search");
// Step 3: Semantic search (over-fetch for scoring pool)
const internalLimit = Math.max(topK * 4, 60);
@@ -1442,7 +1442,10 @@ export class Memory {
entitySearchFilters[k] = effectiveFilters[k];
}
const entityTexts = deduped.map((e) => e.text);
const embeddings = await this.embedder.embedBatch(entityTexts);
const embeddings = await this.embedder.embedBatch(
entityTexts,
"search",
);
if (embeddings.length !== entityTexts.length) {
console.warn(
@@ -1648,7 +1651,7 @@ export class Memory {
const existingEmbeddings: Record<string, number[]> = {};
if (text != null) {
existingEmbeddings[text] = await this.embedder.embed(text);
existingEmbeddings[text] = await this.embedder.embed(text, "update");
}
await this.updateMemory(memoryId, text, existingEmbeddings, updateMetadata);
@@ -1873,7 +1876,7 @@ export class Memory {
): Promise<string> {
const memoryId = uuidv4();
const embedding =
existingEmbeddings[data] || (await this.embedder.embed(data));
existingEmbeddings[data] || (await this.embedder.embed(data, "add"));
const memoryMetadata = {
...metadata,
@@ -1917,7 +1920,8 @@ export class Memory {
const textChanged = newData !== prevValue;
const embedding =
existingEmbeddings[newData] || (await this.embedder.embed(newData));
existingEmbeddings[newData] ||
(await this.embedder.embed(newData, "update"));
const newMetadata = {
...existingMemory.payload,
+18
View File
@@ -23,6 +23,15 @@ export interface EmbeddingConfig {
huggingfaceBaseUrl?: string;
}
export interface VertexAIConfig extends EmbeddingConfig {
vertexCredentialsJson?: string;
googleServiceAccountJson?: string | Record<string, any>;
googleProjectId?: string;
location?: string;
memoryAddEmbeddingType?: string;
memoryUpdateEmbeddingType?: string;
memorySearchEmbeddingType?: string;
}
export type { ValkeyConfig } from "./valkey";
export interface VectorStoreConfig {
@@ -172,6 +181,15 @@ export const MemoryConfigSchema = z.object({
baseURL: z.string().optional(),
embeddingDims: z.number().optional(),
url: z.string().optional(),
vertexCredentialsJson: z.string().optional(),
googleServiceAccountJson: z
.union([z.string(), z.record(z.string(), z.any())])
.optional(),
googleProjectId: z.string().optional(),
location: z.string().optional(),
memoryAddEmbeddingType: z.string().optional(),
memoryUpdateEmbeddingType: z.string().optional(),
memorySearchEmbeddingType: z.string().optional(),
}),
}),
vectorStore: z.object({
+3
View File
@@ -54,6 +54,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 { VertexAIEmbedder } from "../embeddings/vertexai";
import { ElasticsearchDB } from "../vector_stores/elasticsearch";
import { OpenSearchDB } from "../vector_stores/opensearch";
import { UpstashVector } from "../vector_stores/upstash_vector";
@@ -87,6 +88,8 @@ export class EmbedderFactory {
return new FastEmbedEmbedder(config);
case "langchain":
return new LangchainEmbedder(config);
case "vertexai":
return new VertexAIEmbedder(config);
case "huggingface":
return new HuggingFaceEmbedder(config);
default:
+55 -1
View File
@@ -422,6 +422,57 @@ describe("ConfigManager", () => {
expect(cfg.vectorStore.config.port).toBe(6333);
});
});
describe("mergeConfig - provider-specific embedder fields", () => {
// The embedder config used to be rebuilt from a fixed key list, which
// dropped every provider-specific field before the embedder was
// constructed. Vertex AI then authenticated against whatever ambient
// project ADC resolved to and ignored the configured task types.
it("preserves Vertex AI fields through the merge", () => {
const cfg = ConfigManager.mergeConfig({
embedder: {
provider: "vertexai",
config: {
model: "gemini-embedding-001",
googleProjectId: "my-proj",
location: "europe-west4",
vertexCredentialsJson: "/creds.json",
memoryAddEmbeddingType: "SEMANTIC_SIMILARITY",
},
},
vectorStore: { provider: "memory", config: { collectionName: "test" } },
llm: { provider: "openai", config: { apiKey: "test-key" } },
});
expect(cfg.embedder.config).toMatchObject({
model: "gemini-embedding-001",
googleProjectId: "my-proj",
location: "europe-west4",
vertexCredentialsJson: "/creds.json",
memoryAddEmbeddingType: "SEMANTIC_SIMILARITY",
});
});
it("still lets normalized values win over the raw user config", () => {
const cfg = ConfigManager.mergeConfig({
embedder: {
provider: "lmstudio",
config: {
lmstudio_base_url: "http://localhost:1234/v1",
embedding_dims: 768,
},
} as never,
vectorStore: { provider: "memory", config: { collectionName: "test" } },
llm: { provider: "openai", config: { apiKey: "test-key" } },
});
expect(cfg.embedder.config.baseURL).toBe("http://localhost:1234/v1");
expect(cfg.embedder.config.embeddingDims).toBe(768);
// Snake_case aliases are normalized, not passed through to the provider.
expect(cfg.embedder.config).not.toHaveProperty("lmstudio_base_url");
expect(cfg.embedder.config).not.toHaveProperty("embedding_dims");
});
});
});
// ─────────────────────────────────────────────────────────────────────────
@@ -629,7 +680,10 @@ describe("Memory – LM Studio end-to-end flow", () => {
filters: { user_id: "u1" },
});
expect(mockEmbedder.embed).toHaveBeenCalledWith("What does the user like?");
expect(mockEmbedder.embed).toHaveBeenCalledWith(
"What does the user like?",
"search",
);
expect(mockVStore.search).toHaveBeenCalled();
expect(result.results).toHaveLength(1);
expect(result.results[0].memory).toBe("User likes hiking");
@@ -40,6 +40,11 @@ jest.mock("../src/embeddings/lmstudio", () => ({
.fn()
.mockImplementation((config) => ({ type: "lmstudio-embedder", config })),
}));
jest.mock("../src/embeddings/vertexai", () => ({
VertexAIEmbedder: jest
.fn()
.mockImplementation((config) => ({ type: "vertexai-embedder", config })),
}));
jest.mock("../src/embeddings/together", () => ({
TogetherEmbedder: jest
.fn()
@@ -241,6 +246,7 @@ describe("EmbedderFactory", () => {
["fastembed"],
["langchain"],
["lmstudio"],
["vertexai"],
["together"],
])("creates embedder for provider '%s'", (provider) => {
expect(() =>
@@ -0,0 +1,115 @@
/// <reference types="jest" />
/**
* Verifies the memory pipeline threads the correct memory action
* ("add" | "update" | "search") into the embedder. Task-type-aware providers
* (e.g. Vertex AI) embed queries and documents differently based on this, and
* the argument is silently ignored by every other embedder, so only a
* pipeline-level test catches a dropped action.
*/
import { Memory } from "../src/memory";
const mockEmbedding = new Array(1536).fill(0.1);
// Prefixed `mock*` so jest's hoisted module factory may reference them.
const mockEmbed = jest.fn().mockResolvedValue(mockEmbedding);
const mockEmbedBatch = jest
.fn()
.mockImplementation((texts: string[]) =>
Promise.resolve(texts.map(() => mockEmbedding)),
);
const mockGenerateResponse = jest
.fn()
.mockResolvedValue(JSON.stringify({ memory: [] }));
jest.mock("../src/embeddings/google", () => ({ GoogleEmbedder: jest.fn() }));
jest.mock("../src/llms/google", () => ({ GoogleLLM: jest.fn() }));
jest.mock("../src/llms/openai", () => ({
OpenAILLM: jest.fn().mockImplementation(() => ({
generateResponse: mockGenerateResponse,
})),
}));
jest.mock("../src/embeddings/openai", () => ({
OpenAIEmbedder: jest.fn().mockImplementation(() => ({
embed: mockEmbed,
embedBatch: mockEmbedBatch,
embeddingDims: 1536,
})),
}));
function createMemory(): Memory {
return new Memory({
version: "v1.1",
embedder: {
provider: "openai",
config: { apiKey: "test-key", model: "text-embedding-3-small" },
},
vectorStore: {
provider: "memory",
config: {
collectionName: `test-action-${Date.now()}-${Math.random()}`,
dimension: 1536,
dbPath: ":memory:",
},
},
llm: {
provider: "openai",
config: { apiKey: "test-key", model: "gpt-5-mini" },
},
historyDbPath: ":memory:",
});
}
describe("embedder memory-action threading", () => {
let memory: Memory;
beforeEach(() => {
memory = createMemory();
mockEmbed.mockClear();
mockEmbedBatch.mockClear();
mockGenerateResponse.mockResolvedValue(JSON.stringify({ memory: [] }));
});
afterEach(async () => {
await memory.reset();
});
test("search() embeds the query with the 'search' action", async () => {
await memory.search("what do I like", { filters: { user_id: "u1" } });
expect(mockEmbed).toHaveBeenCalledWith("what do I like", "search");
});
test("update() embeds the new value with the 'update' action", async () => {
// Missing id: update embeds the value before it throws on the absent row.
await memory.update("missing-id", "new value").catch(() => {});
expect(mockEmbed).toHaveBeenCalledWith("new value", "update");
});
test("add() batch-embeds extracted memories and entities with the 'add' action", async () => {
mockGenerateResponse.mockResolvedValue(
JSON.stringify({
memory: [
{ id: "1", text: "John loves sci-fi movies", attributed_to: "user" },
],
}),
);
await memory.add("I love sci-fi movies", { userId: "u1" });
// Phase 1 retrieval embeds the incoming turn as a query.
expect(mockEmbed).toHaveBeenCalledWith(
expect.stringContaining("I love sci-fi movies"),
"search",
);
// Phase 3 (extracted memories) and phase 7 (linked entities) both batch
// embed as documents. Without an explicit action, a task-type-aware
// embedder falls back to its own default and silently mis-embeds.
expect(mockEmbedBatch).toHaveBeenCalledWith(
["John loves sci-fi movies"],
"add",
);
expect(mockEmbedBatch.mock.calls.length).toBeGreaterThan(0);
for (const call of mockEmbedBatch.mock.calls) {
expect(call[1]).toBe("add");
}
});
});
@@ -47,4 +47,11 @@ describe("tsup.config.ts externals", () => {
it("should have peerDependencies defined in package.json", () => {
expect(peerDeps.length).toBeGreaterThan(0);
});
it("should not list any dependency twice", () => {
const duplicates = externalDeps.filter(
(dep, i) => externalDeps.indexOf(dep) !== i,
);
expect(duplicates).toEqual([]);
});
});
@@ -0,0 +1,238 @@
/// <reference types="jest" />
const mockPredict = jest.fn();
const mockGetProjectId = jest.fn();
const mockClientConstructor = jest.fn();
jest.mock("@google-cloud/aiplatform", () => {
return {
__esModule: true,
PredictionServiceClient: jest.fn().mockImplementation((...args) => {
mockClientConstructor(...args);
return {
predict: mockPredict,
getProjectId: mockGetProjectId,
};
}),
helpers: {
toValue: jest.fn().mockImplementation((val) => val),
fromValue: jest.fn().mockImplementation((val) => val),
},
};
});
import { VertexAIEmbedder } from "../src/embeddings/vertexai";
const mockEmbedding = [0.1, 0.2, 0.3, 0.4];
/** Echoes one embedding back per instance in the request. */
function predictEchoingInstances() {
return (req: { instances: unknown[] }) =>
Promise.resolve([
{
predictions: req.instances.map(() => ({
embeddings: { values: mockEmbedding },
})),
},
]);
}
describe("VertexAIEmbedder", () => {
beforeEach(() => {
mockPredict.mockReset();
mockClientConstructor.mockReset();
mockGetProjectId.mockReset();
mockGetProjectId.mockResolvedValue("adc-project");
mockPredict.mockResolvedValue([
{
predictions: [
{
embeddings: {
values: mockEmbedding,
},
},
],
},
]);
});
describe("basic functionality", () => {
it("embed() returns the embedding vector", async () => {
const embedder = new VertexAIEmbedder({
googleProjectId: "test-project",
});
const result = await embedder.embed("hello");
expect(result).toEqual(mockEmbedding);
expect(mockPredict).toHaveBeenCalledTimes(1);
const callArgs = mockPredict.mock.calls[0][0];
expect(callArgs.endpoint).toBe(
"projects/test-project/locations/us-central1/publishers/google/models/gemini-embedding-001",
);
// task_type belongs on the instance (snake_case), parameters carries
// only outputDimensionality.
expect(callArgs.instances).toEqual([
{ content: "hello", task_type: "SEMANTIC_SIMILARITY" },
]);
expect(callArgs.parameters).toEqual({ outputDimensionality: 256 });
});
it("embed() with memory action search uses RETRIEVAL_QUERY", async () => {
const embedder = new VertexAIEmbedder({
googleProjectId: "test-project",
});
await embedder.embed("hello", "search");
const callArgs = mockPredict.mock.calls[0][0];
expect(callArgs.instances[0].task_type).toBe("RETRIEVAL_QUERY");
});
it("embed() with memory action add uses RETRIEVAL_DOCUMENT", async () => {
const embedder = new VertexAIEmbedder({
googleProjectId: "test-project",
});
await embedder.embed("hello", "add");
const callArgs = mockPredict.mock.calls[0][0];
expect(callArgs.instances[0].task_type).toBe("RETRIEVAL_DOCUMENT");
});
it("throws error when predictions are empty", async () => {
mockPredict.mockResolvedValue([{}]);
const embedder = new VertexAIEmbedder({
googleProjectId: "test-project",
});
await expect(embedder.embed("hello")).rejects.toThrow(
"No predictions returned from Vertex AI",
);
});
});
describe("embedBatch() request sizing", () => {
// gemini-embedding-001 (the default model) rejects any predict() call
// carrying more than one input text, so the batch loop must degrade to one
// request per text rather than the 250-instance chunk the older models take.
it("sends one instance per request for gemini-embedding models", async () => {
mockPredict.mockImplementation(predictEchoingInstances());
const embedder = new VertexAIEmbedder({
googleProjectId: "test-project",
});
const texts = ["a", "b", "c"];
const result = await embedder.embedBatch(texts);
expect(result).toEqual(texts.map(() => mockEmbedding));
expect(mockPredict).toHaveBeenCalledTimes(3);
for (const call of mockPredict.mock.calls) {
expect(call[0].instances.length).toBe(1);
// batch defaults to the "add" action -> RETRIEVAL_DOCUMENT
expect(call[0].instances[0].task_type).toBe("RETRIEVAL_DOCUMENT");
expect(call[0].parameters).toEqual({ outputDimensionality: 256 });
}
expect(
mockPredict.mock.calls.map((c) => c[0].instances[0].content),
).toEqual(texts);
});
it("chunks at 250 instances for text-embedding models", async () => {
mockPredict.mockImplementation(predictEchoingInstances());
const embedder = new VertexAIEmbedder({
googleProjectId: "test-project",
model: "text-embedding-005",
});
const texts = Array.from({ length: 255 }, (_, i) => `text-${i}`);
const result = await embedder.embedBatch(texts, "search");
expect(result.length).toBe(255);
expect(mockPredict).toHaveBeenCalledTimes(2);
expect(mockPredict.mock.calls[0][0].instances.length).toBe(250);
expect(mockPredict.mock.calls[0][0].instances[0].task_type).toBe(
"RETRIEVAL_QUERY",
);
expect(mockPredict.mock.calls[1][0].instances.length).toBe(5);
});
it("rejects an unknown memory action", async () => {
const embedder = new VertexAIEmbedder({
googleProjectId: "test-project",
});
await expect(
embedder.embedBatch(["a"], "delete" as unknown as "add"),
).rejects.toThrow("Invalid memory action: delete");
});
});
describe("client initialization", () => {
const PROJECT_ENV_VARS = [
"GCP_PROJECT_ID",
"GOOGLE_CLOUD_PROJECT",
"GCLOUD_PROJECT",
];
let savedEnv: Record<string, string | undefined>;
beforeEach(() => {
savedEnv = {};
for (const key of PROJECT_ENV_VARS) {
savedEnv[key] = process.env[key];
delete process.env[key];
}
});
afterEach(() => {
for (const key of PROJECT_ENV_VARS) {
if (savedEnv[key] === undefined) delete process.env[key];
else process.env[key] = savedEnv[key];
}
});
it("resolves the project ID from credentials when none is configured", async () => {
const embedder = new VertexAIEmbedder({});
await embedder.embed("hello");
expect(mockGetProjectId).toHaveBeenCalledTimes(1);
expect(mockPredict.mock.calls[0][0].endpoint).toBe(
"projects/adc-project/locations/us-central1/publishers/google/models/gemini-embedding-001",
);
});
it("surfaces a helpful error when no project ID can be resolved", async () => {
mockGetProjectId.mockRejectedValue(
new Error("Unable to detect a Project Id"),
);
const embedder = new VertexAIEmbedder({});
await expect(embedder.embed("hello")).rejects.toThrow(
"Vertex AI could not determine a Google Cloud project ID",
);
});
it("builds one client for concurrent calls", async () => {
const embedder = new VertexAIEmbedder({
googleProjectId: "test-project",
});
await Promise.all([
embedder.embed("a"),
embedder.embed("b"),
embedder.embed("c"),
]);
expect(mockClientConstructor).toHaveBeenCalledTimes(1);
});
it("retries initialization after a failure", async () => {
mockGetProjectId.mockRejectedValueOnce(new Error("transient"));
const embedder = new VertexAIEmbedder({});
await expect(embedder.embed("hello")).rejects.toThrow(
"Vertex AI could not determine a Google Cloud project ID",
);
await expect(embedder.embed("hello")).resolves.toEqual(mockEmbedding);
});
});
});
@@ -0,0 +1,139 @@
/// <reference types="jest" />
/**
* The sibling `vertexai-embedder.test.ts` stubs `helpers.toValue`/`fromValue`
* as identity functions, so it never exercises the real protobuf `Value`
* encode/decode path. Here we mock only `PredictionServiceClient` (no
* network, no credentials) and keep the REAL `helpers`, so a malformed
* instance shape or a broken decode actually fails.
*/
const mockPredict = jest.fn();
const mockGetProjectId = jest.fn();
jest.mock("@google-cloud/aiplatform", () => {
const actual = jest.requireActual("@google-cloud/aiplatform");
return {
...actual,
__esModule: true,
PredictionServiceClient: jest.fn().mockImplementation(() => ({
predict: mockPredict,
getProjectId: mockGetProjectId,
})),
};
});
import { VertexAIEmbedder } from "../src/embeddings/vertexai";
import { helpers } from "@google-cloud/aiplatform";
/** A `google.protobuf.Value` always carries `kind`+`structValue`/etc; a plain
* JS object never does. This is what "real encoding happened" looks like. */
function expectEncodedValue(value: unknown) {
expect(value).toEqual(expect.objectContaining({ kind: expect.any(String) }));
}
describe("VertexAIEmbedder protobuf boundary (real helpers.toValue/fromValue)", () => {
beforeEach(() => {
mockPredict.mockReset();
mockGetProjectId.mockReset();
mockGetProjectId.mockResolvedValue("adc-project");
});
it("encodes the instance as a genuine protobuf Value carrying {content, task_type}", async () => {
mockPredict.mockResolvedValue([
{
predictions: [helpers.toValue({ embeddings: { values: [0.1, 0.2] } })],
},
]);
const embedder = new VertexAIEmbedder({ googleProjectId: "test-project" });
await embedder.embed("hello world", "search");
const { instances } = mockPredict.mock.calls[0][0];
expect(instances).toHaveLength(1);
expectEncodedValue(instances[0]);
// Decode with the REAL fromValue -- proves the encoded instance is
// readable and matches exactly what Vertex expects on the wire.
expect(helpers.fromValue(instances[0])).toEqual({
content: "hello world",
task_type: "RETRIEVAL_QUERY",
});
});
it("encodes parameters as a genuine protobuf Value carrying {outputDimensionality}", async () => {
mockPredict.mockResolvedValue([
{
predictions: [helpers.toValue({ embeddings: { values: [0.1] } })],
},
]);
const embedder = new VertexAIEmbedder({
googleProjectId: "test-project",
embeddingDims: 768,
});
await embedder.embed("hello");
const { parameters } = mockPredict.mock.calls[0][0];
expectEncodedValue(parameters);
expect(helpers.fromValue(parameters)).toEqual({
outputDimensionality: 768,
});
});
it("decodes a real toValue()-encoded prediction back into the embedding array", async () => {
const values = [0.11, -0.22, 0.33, 0.0];
mockPredict.mockResolvedValue([
{
predictions: [helpers.toValue({ embeddings: { values } })],
},
]);
const embedder = new VertexAIEmbedder({ googleProjectId: "test-project" });
const result = await embedder.embed("hello");
expect(result).toEqual(values);
});
it("rejects a prediction that isn't a real encoded protobuf Value", async () => {
// A raw JS object (what the identity-stubbed sibling test effectively
// assumed `predict()` returns) is not a valid protobuf Value -- the real
// fromValue() throws on it instead of silently passing it through.
mockPredict.mockResolvedValue([
{ predictions: [{ embeddings: { values: [0.1, 0.2] } }] },
]);
const embedder = new VertexAIEmbedder({ googleProjectId: "test-project" });
await expect(embedder.embed("hello")).rejects.toThrow();
});
it("embedBatch() round-trips multiple real encoded predictions", async () => {
const vectors = [
[0.1, 0.2],
[0.3, 0.4],
];
mockPredict.mockImplementation(
(req: { instances: unknown[] }) =>
Promise.resolve([
{
predictions: req.instances.map((_, i) =>
helpers.toValue({ embeddings: { values: vectors[i] } }),
),
},
]) as any,
);
const embedder = new VertexAIEmbedder({
googleProjectId: "test-project",
model: "text-embedding-005",
});
const result = await embedder.embedBatch(["a", "b"]);
expect(result).toEqual(vectors);
const { instances } = mockPredict.mock.calls[0][0];
expect(instances.map((i: unknown) => helpers.fromValue(i as any))).toEqual([
{ content: "a", task_type: "RETRIEVAL_DOCUMENT" },
{ content: "b", task_type: "RETRIEVAL_DOCUMENT" },
]);
});
});