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
+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" },
]);
});
});