feat(ts-sdk): add Vertex AI embedding provider (#5882)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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[][]>;
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -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";
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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({
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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" },
|
||||
]);
|
||||
});
|
||||
});
|
||||
Reference in New Issue
Block a user