feat(ts-sdk): add AWS Bedrock embedding provider (#6185)
This commit is contained in:
@@ -3,11 +3,27 @@ title: AWS Bedrock
|
||||
description: "Configure AWS Bedrock as an embedding provider in Mem0 with IAM credentials and boto3 authentication."
|
||||
---
|
||||
|
||||
To use AWS Bedrock embedding models, you need to have the appropriate AWS credentials and permissions. The embeddings implementation relies on the `boto3` library.
|
||||
To use AWS Bedrock embedding models, you need the appropriate AWS credentials and permissions. Python uses `boto3`, and TypeScript uses `@aws-sdk/client-bedrock-runtime`.
|
||||
|
||||
Both SDKs support the Amazon Titan and Cohere embedding model families.
|
||||
|
||||
### Setup
|
||||
- Ensure you have model access from the [AWS Bedrock Console](https://us-east-1.console.aws.amazon.com/bedrock/home?region=us-east-1#/modelaccess)
|
||||
- Authenticate the boto3 client using a method described in the [AWS documentation](https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html)
|
||||
|
||||
- Model access is automatic: Bedrock enables serverless foundation models on first invocation in AWS commercial regions, and the [Model access page has been retired](https://docs.aws.amazon.com/bedrock/latest/userguide/model-access.html). Cohere models are served from AWS Marketplace, so an account's first invocation must come from a principal with the `aws-marketplace:Subscribe` permission; after that, any user in the account can invoke them. Browse the models available to you in the [Bedrock model catalog](https://console.aws.amazon.com/bedrock/).
|
||||
- Install the AWS client for your language:
|
||||
|
||||
<CodeGroup>
|
||||
```bash Python
|
||||
pip install boto3
|
||||
```
|
||||
|
||||
```bash TypeScript
|
||||
npm install @aws-sdk/client-bedrock-runtime
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
In TypeScript this package is an optional peer dependency, so it is only required when you actually use the Bedrock embedder.
|
||||
|
||||
- Set up environment variables for authentication:
|
||||
```bash
|
||||
export AWS_REGION=us-east-1
|
||||
@@ -15,6 +31,8 @@ To use AWS Bedrock embedding models, you need to have the appropriate AWS creden
|
||||
export AWS_SECRET_ACCESS_KEY=your-secret-key
|
||||
```
|
||||
|
||||
Both SDKs fall back to the standard AWS credential chain (environment variables, shared config, SSO, or an instance role) when you do not pass credentials in the config, so you rarely need to hardcode keys. See the [boto3 credentials guide](https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html) for the Python resolution order.
|
||||
|
||||
### Usage
|
||||
|
||||
<CodeGroup>
|
||||
@@ -48,8 +66,46 @@ messages = [
|
||||
]
|
||||
m.add(messages, user_id="alice")
|
||||
```
|
||||
|
||||
```typescript TypeScript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
// Credentials are read from the AWS default chain (AWS_REGION,
|
||||
// AWS_ACCESS_KEY_ID, AWS_SECRET_ACCESS_KEY, SSO, or an instance role).
|
||||
const memory = new Memory({
|
||||
embedder: {
|
||||
provider: "aws_bedrock",
|
||||
config: {
|
||||
model: "amazon.titan-embed-text-v2:0",
|
||||
awsRegion: "us-west-2",
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
const messages = [
|
||||
{ role: "user", content: "I'm planning to watch a movie tonight. Any recommendations?" },
|
||||
{ role: "assistant", content: "How about thriller movies? They can be quite engaging." },
|
||||
{ role: "user", content: "I'm not a big fan of thriller movies but I love sci-fi movies." },
|
||||
{ role: "assistant", content: "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future." },
|
||||
];
|
||||
await memory.add(messages, { userId: "alice" });
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
### Choosing a model
|
||||
|
||||
| Model | Notes |
|
||||
| --- | --- |
|
||||
| `amazon.titan-embed-text-v1` | Default. Fixed 1536-dimension output. |
|
||||
| `amazon.titan-embed-text-v2:0` | Supports a configurable output size of 256, 512, or 1024. |
|
||||
| `cohere.embed-english-v3` | English text. Embeds up to 96 texts per request. |
|
||||
| `cohere.embed-multilingual-v3` | Multilingual text. Embeds up to 96 texts per request. |
|
||||
| `cohere.embed-v4:0` | Text. Embeds up to 96 texts per request. Supports a configurable output size of 256, 512, 1024, or 1536. TypeScript only. |
|
||||
|
||||
Custom output sizes are model specific. In Python, only Titan Text Embeddings V2 accepts one. In TypeScript, Titan Text Embeddings V2 and Cohere Embed v4 both do, and `embeddingDims` is ignored on Titan V1 and on Cohere v3, which have no such parameter. When you do set it, make sure your vector store dimension matches, otherwise inserts will fail.
|
||||
|
||||
Bedrock caps a Cohere embedding call at 96 texts. The TypeScript SDK splits larger batches into multiple requests for you, so a 200 text batch becomes 3 calls.
|
||||
|
||||
### Config
|
||||
|
||||
Here are the parameters available for configuring AWS Bedrock embedder:
|
||||
@@ -64,4 +120,16 @@ Here are the parameters available for configuring AWS Bedrock embedder:
|
||||
| `aws_secret_access_key` | AWS secret access key for authentication | `None` |
|
||||
| `aws_session_token` | AWS session token for temporary credentials | `None` |
|
||||
</Tab>
|
||||
<Tab title="TypeScript">
|
||||
| Parameter | Description | Default Value |
|
||||
| --- | --- | --- |
|
||||
| `model` | The name of the embedding model to use | `amazon.titan-embed-text-v1` |
|
||||
| `awsRegion` | AWS region for the Bedrock client. Falls back to the `AWS_REGION` environment variable | `us-west-2` |
|
||||
| `embeddingDims` | Output vector size. Titan Text Embeddings V2 (256, 512, or 1024) and Cohere Embed v4 (256, 512, 1024, or 1536) only | `undefined` |
|
||||
| `awsAccessKeyId` | AWS access key ID for authentication | `undefined` |
|
||||
| `awsSecretAccessKey` | AWS secret access key for authentication | `undefined` |
|
||||
| `awsSessionToken` | AWS session token for temporary credentials | `undefined` |
|
||||
|
||||
Omit the three credential fields to use the AWS default credential chain. If you do pass them, `awsAccessKeyId` and `awsSecretAccessKey` are both required.
|
||||
</Tab>
|
||||
</Tabs>
|
||||
|
||||
@@ -10,7 +10,7 @@ Mem0 offers support for various embedding models, allowing users to choose the o
|
||||
See the list of supported embedders below.
|
||||
|
||||
<Note>
|
||||
All embedders listed below are supported in the Python implementation. The TypeScript implementation supports: **OpenAI**, **Azure OpenAI**, **FastEmbed**, **Google AI**, **Langchain**, **LM Studio**, **Ollama**, and **Together**.
|
||||
All embedders listed below are supported in the Python implementation. The TypeScript implementation supports: **OpenAI**, **Azure OpenAI**, **AWS Bedrock**, **FastEmbed**, **Google AI**, **Hugging Face**, **Langchain**, **LM Studio**, **Ollama**, **Together**, and **Vertex AI**.
|
||||
</Note>
|
||||
|
||||
<CardGroup cols={4}>
|
||||
|
||||
@@ -0,0 +1,292 @@
|
||||
import { Embedder } from "./base";
|
||||
import { EmbeddingConfig } from "../types";
|
||||
|
||||
const DEFAULT_MODEL = "amazon.titan-embed-text-v1";
|
||||
const DEFAULT_REGION = "us-west-2";
|
||||
|
||||
// Cohere's Bedrock embed API rejects an InvokeModel call carrying more than 96
|
||||
// texts, so `embedBatch` chunks at that boundary.
|
||||
const COHERE_MAX_BATCH = 96;
|
||||
|
||||
// Titan has no server-side batch endpoint -- one InvokeModel call per text --
|
||||
// so without a cap a large embedBatch() would fan out one request per text.
|
||||
// Bounds concurrency the same way COHERE_MAX_BATCH bounds the Cohere path.
|
||||
const TITAN_MAX_CONCURRENCY = 4;
|
||||
|
||||
// Cohere wants to know whether a text is being embedded for storage or for a
|
||||
// retrieval query; embedding a search query in document mode silently
|
||||
// degrades retrieval. Titan ignores this and has no equivalent parameter.
|
||||
const COHERE_INPUT_TYPES: Record<"add" | "update" | "search", string> = {
|
||||
add: "search_document",
|
||||
update: "search_document",
|
||||
search: "search_query",
|
||||
};
|
||||
|
||||
type BedrockRuntimeModule = typeof import("@aws-sdk/client-bedrock-runtime");
|
||||
|
||||
interface BedrockCredentials {
|
||||
accessKeyId: string;
|
||||
secretAccessKey: string;
|
||||
sessionToken?: string;
|
||||
}
|
||||
|
||||
interface BedrockEmbeddingResponse {
|
||||
// Titan returns a single vector. Cohere v3 returns a flat array of vectors;
|
||||
// Cohere v4, when `embedding_types` is requested, nests it as `{ float }`.
|
||||
embedding?: number[];
|
||||
embeddings?: number[][] | { float?: number[][] };
|
||||
}
|
||||
|
||||
/**
|
||||
* Runs `fn` over `items` with at most `limit` calls in flight at once,
|
||||
* returning results in input order regardless of completion order.
|
||||
*/
|
||||
async function mapWithConcurrencyLimit<T, R>(
|
||||
items: T[],
|
||||
limit: number,
|
||||
fn: (item: T) => Promise<R>,
|
||||
): Promise<R[]> {
|
||||
const results: R[] = new Array(items.length);
|
||||
let next = 0;
|
||||
|
||||
async function worker(): Promise<void> {
|
||||
while (next < items.length) {
|
||||
const index = next++;
|
||||
results[index] = await fn(items[index]);
|
||||
}
|
||||
}
|
||||
|
||||
await Promise.all(
|
||||
Array.from({ length: Math.min(limit, items.length) }, worker),
|
||||
);
|
||||
return results;
|
||||
}
|
||||
|
||||
/**
|
||||
* AWS Bedrock embedder, mirroring `mem0/embeddings/aws_bedrock.py`.
|
||||
*
|
||||
* Supports the Amazon Titan and Cohere embedding model families. The
|
||||
* `@aws-sdk/client-bedrock-runtime` dependency is lazily imported so the
|
||||
* package stays optional: importing this module never forces the SDK to be
|
||||
* installed until a Bedrock embedder actually embeds something.
|
||||
*/
|
||||
export class AWSBedrockEmbedder implements Embedder {
|
||||
private readonly model: string;
|
||||
private readonly region: string;
|
||||
private readonly embeddingDims?: number;
|
||||
private readonly credentials?: BedrockCredentials;
|
||||
private clientPromise?: Promise<{
|
||||
sdk: BedrockRuntimeModule;
|
||||
client: { send: (command: any) => Promise<{ body?: Uint8Array }> };
|
||||
}>;
|
||||
|
||||
constructor(config: EmbeddingConfig) {
|
||||
this.model = config.model || DEFAULT_MODEL;
|
||||
this.region = config.awsRegion || process.env.AWS_REGION || DEFAULT_REGION;
|
||||
this.embeddingDims = config.embeddingDims;
|
||||
|
||||
const hasKeyPair = Boolean(
|
||||
config.awsAccessKeyId && config.awsSecretAccessKey,
|
||||
);
|
||||
const hasAnyCredential = Boolean(
|
||||
config.awsAccessKeyId ||
|
||||
config.awsSecretAccessKey ||
|
||||
config.awsSessionToken,
|
||||
);
|
||||
|
||||
// Partially configured credentials would silently fall back to the default
|
||||
// chain, embedding under an identity the caller never chose.
|
||||
if (hasAnyCredential && !hasKeyPair) {
|
||||
throw new Error(
|
||||
"AWS Bedrock requires both awsAccessKeyId and awsSecretAccessKey when any explicit credential is configured. " +
|
||||
"Omit all credential fields to use the AWS default credential chain.",
|
||||
);
|
||||
}
|
||||
|
||||
// Leaving `credentials` unset lets the AWS SDK resolve them from its
|
||||
// default chain: environment, shared config, SSO, or the instance role.
|
||||
if (hasKeyPair) {
|
||||
this.credentials = {
|
||||
accessKeyId: config.awsAccessKeyId!,
|
||||
secretAccessKey: config.awsSecretAccessKey!,
|
||||
...(config.awsSessionToken && { sessionToken: config.awsSessionToken }),
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
private async loadSdk(): Promise<BedrockRuntimeModule> {
|
||||
try {
|
||||
return await import("@aws-sdk/client-bedrock-runtime");
|
||||
} catch (error) {
|
||||
// Only a genuine module-resolution failure gets the friendly install
|
||||
// hint. Node's native ESM loader raises ERR_MODULE_NOT_FOUND; Jest's
|
||||
// and bundlers' CJS-style resolvers raise MODULE_NOT_FOUND. Anything
|
||||
// else (e.g. the package is installed but throws while loading, such
|
||||
// as on a Node version older than the SDK's own engines requirement)
|
||||
// rethrows unchanged instead of being misreported as "not installed".
|
||||
const code = (error as { code?: string } | undefined)?.code;
|
||||
if (code === "ERR_MODULE_NOT_FOUND" || code === "MODULE_NOT_FOUND") {
|
||||
throw Object.assign(
|
||||
new Error(
|
||||
"The '@aws-sdk/client-bedrock-runtime' package is required to use the AWS Bedrock embedder. " +
|
||||
"Install it with: npm install @aws-sdk/client-bedrock-runtime",
|
||||
),
|
||||
{ cause: error },
|
||||
);
|
||||
}
|
||||
throw error;
|
||||
}
|
||||
}
|
||||
|
||||
private async createClient(): Promise<{
|
||||
sdk: BedrockRuntimeModule;
|
||||
client: { send: (command: any) => Promise<{ body?: Uint8Array }> };
|
||||
}> {
|
||||
const sdk = await this.loadSdk();
|
||||
return {
|
||||
sdk,
|
||||
client: new sdk.BedrockRuntimeClient({
|
||||
region: this.region,
|
||||
...(this.credentials && { credentials: this.credentials }),
|
||||
}),
|
||||
};
|
||||
}
|
||||
|
||||
private getClient() {
|
||||
// Memoized so concurrent embed() calls share one client instead of each
|
||||
// racing to build their own. Cleared on rejection so a transient failure
|
||||
// (e.g. a network blip while resolving credentials) doesn't permanently
|
||||
// disable Bedrock for the rest of this embedder's lifetime.
|
||||
if (!this.clientPromise) {
|
||||
this.clientPromise = this.createClient().catch((err) => {
|
||||
this.clientPromise = undefined;
|
||||
throw err;
|
||||
});
|
||||
}
|
||||
return this.clientPromise;
|
||||
}
|
||||
|
||||
private isCohereModel(): boolean {
|
||||
return this.model.startsWith("cohere.");
|
||||
}
|
||||
|
||||
private isCohereV4Model(): boolean {
|
||||
return this.model.includes("embed-v4");
|
||||
}
|
||||
|
||||
private buildRequestBody(
|
||||
texts: string[],
|
||||
memoryAction?: "add" | "update" | "search",
|
||||
): Record<string, unknown> {
|
||||
if (this.isCohereModel()) {
|
||||
const body: Record<string, unknown> = {
|
||||
texts,
|
||||
input_type: memoryAction
|
||||
? COHERE_INPUT_TYPES[memoryAction]
|
||||
: "search_document",
|
||||
};
|
||||
|
||||
// Only Embed v4 understands embedding_types / output_dimension; v3
|
||||
// rejects unknown fields, so they're guarded to the v4 model family.
|
||||
if (this.isCohereV4Model()) {
|
||||
body.embedding_types = ["float"];
|
||||
if (this.embeddingDims !== undefined) {
|
||||
body.output_dimension = this.embeddingDims;
|
||||
}
|
||||
}
|
||||
return body;
|
||||
}
|
||||
|
||||
// Titan accepts one text per call. Only Titan Text Embeddings V2 supports
|
||||
// a caller-chosen output size (256/512/1024), so the field is guarded the
|
||||
// same way the Python provider guards it.
|
||||
return {
|
||||
inputText: texts[0],
|
||||
...(this.embeddingDims !== undefined &&
|
||||
this.model.includes("titan-embed-text-v2") && {
|
||||
dimensions: this.embeddingDims,
|
||||
}),
|
||||
};
|
||||
}
|
||||
|
||||
private async invoke(
|
||||
texts: string[],
|
||||
memoryAction?: "add" | "update" | "search",
|
||||
): Promise<number[][]> {
|
||||
const { sdk, client } = await this.getClient();
|
||||
|
||||
let payload: BedrockEmbeddingResponse;
|
||||
try {
|
||||
const response = await client.send(
|
||||
new sdk.InvokeModelCommand({
|
||||
modelId: this.model,
|
||||
contentType: "application/json",
|
||||
accept: "application/json",
|
||||
body: new TextEncoder().encode(
|
||||
JSON.stringify(this.buildRequestBody(texts, memoryAction)),
|
||||
),
|
||||
}),
|
||||
);
|
||||
payload = JSON.parse(new TextDecoder().decode(response.body));
|
||||
} catch (error) {
|
||||
const message = error instanceof Error ? error.message : String(error);
|
||||
throw new Error(
|
||||
`Error getting embedding from AWS Bedrock model ${this.model}: ${message}`,
|
||||
);
|
||||
}
|
||||
|
||||
// Validated outside the try so this message is not re-wrapped by the catch.
|
||||
// Cohere v3 replies with a flat `embeddings` array; v4 (when
|
||||
// embedding_types is requested) nests it under `.float`.
|
||||
const embeddings = this.isCohereModel()
|
||||
? Array.isArray(payload.embeddings)
|
||||
? payload.embeddings
|
||||
: payload.embeddings?.float
|
||||
: payload.embedding && [payload.embedding];
|
||||
|
||||
// `[]` is truthy, so a lone zero-length vector must be checked for
|
||||
// explicitly -- otherwise it passes the length check and hands the
|
||||
// caller an empty embedding instead of an error.
|
||||
if (
|
||||
!embeddings ||
|
||||
embeddings.length !== texts.length ||
|
||||
embeddings.some((embedding) => embedding.length === 0)
|
||||
) {
|
||||
throw new Error(
|
||||
`AWS Bedrock model ${this.model} returned no embedding for one or more inputs`,
|
||||
);
|
||||
}
|
||||
return embeddings;
|
||||
}
|
||||
|
||||
async embed(
|
||||
text: string,
|
||||
memoryAction?: "add" | "update" | "search",
|
||||
): Promise<number[]> {
|
||||
return (await this.invoke([text], memoryAction))[0];
|
||||
}
|
||||
|
||||
async embedBatch(
|
||||
texts: string[],
|
||||
memoryAction?: "add" | "update" | "search",
|
||||
): Promise<number[][]> {
|
||||
if (texts.length === 0) return [];
|
||||
|
||||
if (!this.isCohereModel()) {
|
||||
return mapWithConcurrencyLimit(texts, TITAN_MAX_CONCURRENCY, (text) =>
|
||||
this.embed(text, memoryAction),
|
||||
);
|
||||
}
|
||||
|
||||
const embeddings: number[][] = [];
|
||||
for (let i = 0; i < texts.length; i += COHERE_MAX_BATCH) {
|
||||
embeddings.push(
|
||||
...(await this.invoke(
|
||||
texts.slice(i, i + COHERE_MAX_BATCH),
|
||||
memoryAction,
|
||||
)),
|
||||
);
|
||||
}
|
||||
return embeddings;
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,7 @@ export * from "./memory";
|
||||
export * from "./memory/memory.types";
|
||||
export * from "./types";
|
||||
export * from "./embeddings/base";
|
||||
export * from "./embeddings/aws_bedrock";
|
||||
export * from "./embeddings/huggingface";
|
||||
export * from "./embeddings/openai";
|
||||
export * from "./embeddings/ollama";
|
||||
|
||||
@@ -21,6 +21,11 @@ export interface EmbeddingConfig {
|
||||
modelProperties?: Record<string, any>;
|
||||
// HuggingFace TEI / OpenAI-compatible inference endpoint base URL.
|
||||
huggingfaceBaseUrl?: string;
|
||||
// AWS Bedrock. Omit the credential fields to use the AWS default chain.
|
||||
awsRegion?: string;
|
||||
awsAccessKeyId?: string;
|
||||
awsSecretAccessKey?: string;
|
||||
awsSessionToken?: string;
|
||||
}
|
||||
|
||||
export interface VertexAIConfig extends EmbeddingConfig {
|
||||
@@ -198,6 +203,10 @@ export const MemoryConfigSchema = z.object({
|
||||
memoryAddEmbeddingType: z.string().optional(),
|
||||
memoryUpdateEmbeddingType: z.string().optional(),
|
||||
memorySearchEmbeddingType: z.string().optional(),
|
||||
awsRegion: z.string().optional(),
|
||||
awsAccessKeyId: z.string().optional(),
|
||||
awsSecretAccessKey: z.string().optional(),
|
||||
awsSessionToken: z.string().optional(),
|
||||
}),
|
||||
}),
|
||||
vectorStore: z.object({
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import { OpenAIEmbedder } from "../embeddings/openai";
|
||||
import { AWSBedrockEmbedder } from "../embeddings/aws_bedrock";
|
||||
import { OllamaEmbedder } from "../embeddings/ollama";
|
||||
import { LMStudioEmbedder } from "../embeddings/lmstudio";
|
||||
import { TogetherEmbedder } from "../embeddings/together";
|
||||
@@ -76,6 +77,8 @@ export class EmbedderFactory {
|
||||
switch (provider.toLowerCase()) {
|
||||
case "openai":
|
||||
return new OpenAIEmbedder(config);
|
||||
case "aws_bedrock":
|
||||
return new AWSBedrockEmbedder(config);
|
||||
case "ollama":
|
||||
return new OllamaEmbedder(config);
|
||||
case "lmstudio":
|
||||
|
||||
@@ -0,0 +1,517 @@
|
||||
import type { InvokeModelCommand } from "@aws-sdk/client-bedrock-runtime";
|
||||
import { AWSBedrockEmbedder } from "../src/embeddings/aws_bedrock";
|
||||
import type { Embedder } from "../src/embeddings/base";
|
||||
import { EmbedderFactory } from "../src/utils/factory";
|
||||
|
||||
/**
|
||||
* Only the network boundary is faked: `BedrockRuntimeClient.send` never leaves
|
||||
* the process. `InvokeModelCommand` stays the real class from the AWS SDK, so
|
||||
* every assertion below runs against the exact payload Bedrock would receive.
|
||||
*/
|
||||
const mockSend = jest.fn();
|
||||
const mockClientConfigs: any[] = [];
|
||||
const mockClientConstructor = jest
|
||||
.fn()
|
||||
.mockImplementation((config: unknown) => {
|
||||
mockClientConfigs.push(config);
|
||||
return { send: mockSend };
|
||||
});
|
||||
|
||||
jest.mock("@aws-sdk/client-bedrock-runtime", () => {
|
||||
const actual = jest.requireActual("@aws-sdk/client-bedrock-runtime");
|
||||
return {
|
||||
...actual,
|
||||
BedrockRuntimeClient: mockClientConstructor,
|
||||
};
|
||||
});
|
||||
|
||||
const encode = (payload: unknown) => ({
|
||||
body: new TextEncoder().encode(JSON.stringify(payload)),
|
||||
});
|
||||
|
||||
const titanReply = (embedding: number[]) =>
|
||||
encode({ embedding, inputTextTokenCount: embedding.length });
|
||||
|
||||
const cohereReply = (embeddings: number[][]) =>
|
||||
encode({ embeddings, id: "req-1", response_type: "embeddings_floats" });
|
||||
|
||||
const commandAt = (index: number): InvokeModelCommand =>
|
||||
mockSend.mock.calls[index][0];
|
||||
|
||||
const requestBodyAt = (index: number) =>
|
||||
JSON.parse(new TextDecoder().decode(commandAt(index).input.body));
|
||||
|
||||
describe("AWSBedrockEmbedder", () => {
|
||||
const savedEnv = { ...process.env };
|
||||
|
||||
beforeEach(() => {
|
||||
jest.clearAllMocks();
|
||||
mockClientConfigs.length = 0;
|
||||
delete process.env.AWS_REGION;
|
||||
});
|
||||
|
||||
afterAll(() => {
|
||||
process.env = savedEnv;
|
||||
});
|
||||
|
||||
describe("Titan models", () => {
|
||||
it("sends inputText and returns the embedding vector", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([0.1, 0.2, 0.3]));
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
const embedding = await embedder.embed("hello world");
|
||||
|
||||
expect(embedding).toEqual([0.1, 0.2, 0.3]);
|
||||
expect(mockSend).toHaveBeenCalledTimes(1);
|
||||
expect(commandAt(0).input.modelId).toBe("amazon.titan-embed-text-v1");
|
||||
expect(commandAt(0).input.contentType).toBe("application/json");
|
||||
expect(commandAt(0).input.accept).toBe("application/json");
|
||||
expect(requestBodyAt(0)).toEqual({ inputText: "hello world" });
|
||||
});
|
||||
|
||||
it("forwards dimensions to Titan V2 when embeddingDims is set", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([0.1, 0.2]));
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "amazon.titan-embed-text-v2:0",
|
||||
embeddingDims: 512,
|
||||
});
|
||||
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(requestBodyAt(0)).toEqual({ inputText: "hello", dimensions: 512 });
|
||||
});
|
||||
|
||||
it("omits dimensions on Titan V1, which rejects the field", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([0.1, 0.2]));
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "amazon.titan-embed-text-v1",
|
||||
embeddingDims: 512,
|
||||
});
|
||||
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(requestBodyAt(0)).toEqual({ inputText: "hello" });
|
||||
});
|
||||
|
||||
// F6: the old guard was `model.includes("v2")`, which would also match
|
||||
// any future/other Titan model whose id merely contains "v2" somewhere
|
||||
// (e.g. an image model), wrongly sending `dimensions` to a model that may
|
||||
// reject it. Only Titan Text Embeddings V2 should get the field.
|
||||
it("does not forward dimensions to a non-Titan-V2 model whose name merely contains v2", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([0.1, 0.2]));
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "amazon.titan-embed-image-v2:0",
|
||||
embeddingDims: 512,
|
||||
});
|
||||
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(requestBodyAt(0)).toEqual({ inputText: "hello" });
|
||||
});
|
||||
|
||||
it("embedBatch issues one request per text and preserves order", async () => {
|
||||
mockSend
|
||||
.mockResolvedValueOnce(titanReply([1, 1]))
|
||||
.mockResolvedValueOnce(titanReply([2, 2]));
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
const embeddings = await embedder.embedBatch(["first", "second"]);
|
||||
|
||||
expect(embeddings).toEqual([
|
||||
[1, 1],
|
||||
[2, 2],
|
||||
]);
|
||||
expect(mockSend).toHaveBeenCalledTimes(2);
|
||||
expect(requestBodyAt(0)).toEqual({ inputText: "first" });
|
||||
expect(requestBodyAt(1)).toEqual({ inputText: "second" });
|
||||
});
|
||||
});
|
||||
|
||||
describe("Cohere models", () => {
|
||||
it("sends texts with a search_document input type", async () => {
|
||||
mockSend.mockResolvedValueOnce(cohereReply([[0.4, 0.5]]));
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "cohere.embed-english-v3",
|
||||
});
|
||||
|
||||
const embedding = await embedder.embed("hello");
|
||||
|
||||
expect(embedding).toEqual([0.4, 0.5]);
|
||||
expect(requestBodyAt(0)).toEqual({
|
||||
texts: ["hello"],
|
||||
input_type: "search_document",
|
||||
});
|
||||
});
|
||||
|
||||
it("embedBatch sends every text in a single request", async () => {
|
||||
mockSend.mockResolvedValueOnce(cohereReply([[1], [2], [3]]));
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "cohere.embed-multilingual-v3",
|
||||
});
|
||||
|
||||
const embeddings = await embedder.embedBatch(["a", "b", "c"]);
|
||||
|
||||
expect(embeddings).toEqual([[1], [2], [3]]);
|
||||
expect(mockSend).toHaveBeenCalledTimes(1);
|
||||
expect(requestBodyAt(0).texts).toEqual(["a", "b", "c"]);
|
||||
});
|
||||
|
||||
it("embedBatch splits requests at Cohere's 96 text limit", async () => {
|
||||
const texts = Array.from({ length: 100 }, (_, i) => `text-${i}`);
|
||||
mockSend
|
||||
.mockResolvedValueOnce(
|
||||
cohereReply(texts.slice(0, 96).map((_, i) => [i])),
|
||||
)
|
||||
.mockResolvedValueOnce(cohereReply(texts.slice(96).map((_, i) => [i])));
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "cohere.embed-english-v3",
|
||||
});
|
||||
|
||||
const embeddings = await embedder.embedBatch(texts);
|
||||
|
||||
expect(embeddings).toHaveLength(100);
|
||||
expect(mockSend).toHaveBeenCalledTimes(2);
|
||||
expect(requestBodyAt(0).texts).toHaveLength(96);
|
||||
expect(requestBodyAt(1).texts).toEqual([
|
||||
"text-96",
|
||||
"text-97",
|
||||
"text-98",
|
||||
"text-99",
|
||||
]);
|
||||
});
|
||||
});
|
||||
|
||||
describe("Cohere Embed v4", () => {
|
||||
// F5: only Embed v4 understands embedding_types / output_dimension, and
|
||||
// (when embedding_types is requested) replies with a nested
|
||||
// `{ embeddings: { float: [...] } }` shape instead of v3's flat array.
|
||||
it("requests embedding_types and output_dimension for v4 models", async () => {
|
||||
mockSend.mockResolvedValueOnce(
|
||||
encode({ embeddings: { float: [[0.1, 0.2, 0.3]] } }),
|
||||
);
|
||||
const embedder = new AWSBedrockEmbedder({
|
||||
model: "cohere.embed-v4:0",
|
||||
embeddingDims: 512,
|
||||
});
|
||||
|
||||
const embedding = await embedder.embed("hello");
|
||||
|
||||
expect(embedding).toEqual([0.1, 0.2, 0.3]);
|
||||
expect(requestBodyAt(0)).toEqual({
|
||||
texts: ["hello"],
|
||||
input_type: "search_document",
|
||||
embedding_types: ["float"],
|
||||
output_dimension: 512,
|
||||
});
|
||||
});
|
||||
|
||||
it("omits output_dimension for v4 when embeddingDims is unset", async () => {
|
||||
mockSend.mockResolvedValueOnce(
|
||||
encode({ embeddings: { float: [[0.1, 0.2]] } }),
|
||||
);
|
||||
const embedder = new AWSBedrockEmbedder({ model: "cohere.embed-v4:0" });
|
||||
|
||||
await embedder.embed("hello");
|
||||
|
||||
expect(requestBodyAt(0)).toEqual({
|
||||
texts: ["hello"],
|
||||
input_type: "search_document",
|
||||
embedding_types: ["float"],
|
||||
});
|
||||
});
|
||||
|
||||
it("parses the nested embeddings.float response shape", async () => {
|
||||
mockSend.mockResolvedValueOnce(
|
||||
encode({
|
||||
embeddings: {
|
||||
float: [
|
||||
[1, 2],
|
||||
[3, 4],
|
||||
],
|
||||
},
|
||||
}),
|
||||
);
|
||||
const embedder = new AWSBedrockEmbedder({ model: "cohere.embed-v4:0" });
|
||||
|
||||
const embeddings = await embedder.embedBatch(["a", "b"]);
|
||||
|
||||
expect(embeddings).toEqual([
|
||||
[1, 2],
|
||||
[3, 4],
|
||||
]);
|
||||
});
|
||||
});
|
||||
|
||||
describe("client configuration", () => {
|
||||
it("defaults to the us-west-2 region", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([1]));
|
||||
await new AWSBedrockEmbedder({}).embed("hello");
|
||||
|
||||
expect(mockClientConfigs[0].region).toBe("us-west-2");
|
||||
});
|
||||
|
||||
it("prefers awsRegion over the AWS_REGION environment variable", async () => {
|
||||
process.env.AWS_REGION = "eu-central-1";
|
||||
mockSend.mockResolvedValueOnce(titanReply([1]));
|
||||
await new AWSBedrockEmbedder({ awsRegion: "ap-south-1" }).embed("hello");
|
||||
|
||||
expect(mockClientConfigs[0].region).toBe("ap-south-1");
|
||||
});
|
||||
|
||||
it("falls back to the AWS_REGION environment variable", async () => {
|
||||
process.env.AWS_REGION = "eu-central-1";
|
||||
mockSend.mockResolvedValueOnce(titanReply([1]));
|
||||
await new AWSBedrockEmbedder({}).embed("hello");
|
||||
|
||||
expect(mockClientConfigs[0].region).toBe("eu-central-1");
|
||||
});
|
||||
|
||||
it("passes explicitly configured credentials to the client", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([1]));
|
||||
await new AWSBedrockEmbedder({
|
||||
awsAccessKeyId: "AKIA_TEST",
|
||||
awsSecretAccessKey: "secret",
|
||||
awsSessionToken: "token",
|
||||
}).embed("hello");
|
||||
|
||||
expect(mockClientConfigs[0].credentials).toEqual({
|
||||
accessKeyId: "AKIA_TEST",
|
||||
secretAccessKey: "secret",
|
||||
sessionToken: "token",
|
||||
});
|
||||
});
|
||||
|
||||
it("leaves credentials unset so the AWS default credential chain applies", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([1]));
|
||||
await new AWSBedrockEmbedder({}).embed("hello");
|
||||
|
||||
expect(mockClientConfigs[0].credentials).toBeUndefined();
|
||||
});
|
||||
|
||||
it("rejects a half-configured credential pair", () => {
|
||||
expect(
|
||||
() => new AWSBedrockEmbedder({ awsAccessKeyId: "AKIA_TEST" }),
|
||||
).toThrow(/awsAccessKeyId and awsSecretAccessKey/);
|
||||
expect(
|
||||
() => new AWSBedrockEmbedder({ awsSecretAccessKey: "secret" }),
|
||||
).toThrow(/awsAccessKeyId and awsSecretAccessKey/);
|
||||
});
|
||||
|
||||
// Silently ignoring a lone session token would fall back to the ambient
|
||||
// credential chain, embedding under an identity the caller never chose.
|
||||
it("rejects a session token supplied without the key pair", () => {
|
||||
expect(
|
||||
() => new AWSBedrockEmbedder({ awsSessionToken: "token" }),
|
||||
).toThrow(/awsAccessKeyId and awsSecretAccessKey/);
|
||||
});
|
||||
});
|
||||
|
||||
describe("provider registration", () => {
|
||||
it("is constructed by EmbedderFactory for the aws_bedrock provider", () => {
|
||||
const embedder = EmbedderFactory.create("aws_bedrock", {
|
||||
model: "amazon.titan-embed-text-v2:0",
|
||||
});
|
||||
|
||||
expect(embedder).toBeInstanceOf(AWSBedrockEmbedder);
|
||||
});
|
||||
});
|
||||
|
||||
describe("error handling", () => {
|
||||
it("wraps Bedrock failures with the model id", async () => {
|
||||
mockSend.mockRejectedValueOnce(new Error("AccessDeniedException"));
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
"Error getting embedding from AWS Bedrock model amazon.titan-embed-text-v1: AccessDeniedException",
|
||||
);
|
||||
});
|
||||
|
||||
it("fails when the response carries no embedding", async () => {
|
||||
mockSend.mockResolvedValueOnce(encode({ inputTextTokenCount: 3 }));
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
/returned no embedding/,
|
||||
);
|
||||
});
|
||||
|
||||
// F7: `[]` is truthy, so `payload.embedding && [payload.embedding]` used to
|
||||
// turn `{"embedding": []}` into `[[]]` -- length 1, which satisfied the
|
||||
// length check for a single-input call and handed the caller an empty vector.
|
||||
it("rejects a zero-length Titan embedding as no embedding", async () => {
|
||||
mockSend.mockResolvedValueOnce(encode({ embedding: [] }));
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
/returned no embedding/,
|
||||
);
|
||||
});
|
||||
|
||||
it("returns an empty array for an empty batch without calling Bedrock", async () => {
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embedBatch([])).resolves.toEqual([]);
|
||||
expect(mockSend).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
||||
describe("dynamic SDK import", () => {
|
||||
// These tests replace the module registered for
|
||||
// @aws-sdk/client-bedrock-runtime for a single resolution. Restore the
|
||||
// working mock afterward so every other test in this file keeps getting
|
||||
// the mocked client instead of hitting module resolution for real.
|
||||
afterEach(() => {
|
||||
jest.resetModules();
|
||||
jest.doMock("@aws-sdk/client-bedrock-runtime", () => {
|
||||
const actual = jest.requireActual("@aws-sdk/client-bedrock-runtime");
|
||||
return { ...actual, BedrockRuntimeClient: mockClientConstructor };
|
||||
});
|
||||
});
|
||||
|
||||
// F1: loadSdk()'s catch used to rewrite *every* import failure into the
|
||||
// "package is required" hint, even when the package is installed but
|
||||
// failed to load for an unrelated reason. That discarded the real error.
|
||||
it("propagates a non-resolution import error unchanged", async () => {
|
||||
jest.resetModules();
|
||||
jest.doMock("@aws-sdk/client-bedrock-runtime", () => {
|
||||
const err: any = new Error("boom: unrelated crash while loading");
|
||||
err.code = "ERR_SOMETHING_ELSE";
|
||||
throw err;
|
||||
});
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
"boom: unrelated crash while loading",
|
||||
);
|
||||
});
|
||||
|
||||
// F1: a genuine resolution failure should still get the friendly install
|
||||
// hint, with the original error preserved as `cause` for debugging.
|
||||
it("gives an install hint for a genuine module-not-found error, preserving the cause", async () => {
|
||||
jest.resetModules();
|
||||
jest.doMock("@aws-sdk/client-bedrock-runtime", () => {
|
||||
const err: any = new Error(
|
||||
"Cannot find module '@aws-sdk/client-bedrock-runtime'",
|
||||
);
|
||||
err.code = "MODULE_NOT_FOUND";
|
||||
throw err;
|
||||
});
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
/npm install @aws-sdk\/client-bedrock-runtime/,
|
||||
);
|
||||
await expect(embedder.embed("hello")).rejects.toMatchObject({
|
||||
cause: expect.objectContaining({ code: "MODULE_NOT_FOUND" }),
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("client promise retry", () => {
|
||||
// F2: getClient() used to memoize the client promise before it resolved,
|
||||
// so a rejected construction (e.g. a transient credentials failure) was
|
||||
// cached forever -- every later embed() call on that instance would
|
||||
// reject immediately without ever retrying.
|
||||
it("retries client construction after a failure instead of caching the rejection", async () => {
|
||||
mockClientConstructor.mockImplementationOnce(() => {
|
||||
throw new Error("credentials not ready");
|
||||
});
|
||||
mockSend.mockResolvedValueOnce(titanReply([1, 2, 3]));
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await expect(embedder.embed("hello")).rejects.toThrow(
|
||||
"credentials not ready",
|
||||
);
|
||||
await expect(embedder.embed("hello")).resolves.toEqual([1, 2, 3]);
|
||||
});
|
||||
});
|
||||
|
||||
describe("memoryAction -> Cohere input_type", () => {
|
||||
// F3: buildRequestBody() used to hardcode `input_type: "search_document"`
|
||||
// regardless of the caller's action, so `Memory.search()` (which calls
|
||||
// `embed(query, "search")`) embedded the query in document mode.
|
||||
//
|
||||
// Typed as `Embedder` (not `AWSBedrockEmbedder`) because that is how
|
||||
// memory/index.ts actually calls it: the interface already declares an
|
||||
// optional `memoryAction` second parameter, so a narrower concrete
|
||||
// `embed(text: string)` satisfies it structurally and tsc stays silent --
|
||||
// the bug is a silent behavioral one, not a compile error.
|
||||
it("sends search_query for a search action", async () => {
|
||||
mockSend.mockResolvedValueOnce(cohereReply([[0.1]]));
|
||||
const embedder: Embedder = new AWSBedrockEmbedder({
|
||||
model: "cohere.embed-english-v3",
|
||||
});
|
||||
|
||||
await embedder.embed("query text", "search");
|
||||
|
||||
expect(requestBodyAt(0)).toEqual({
|
||||
texts: ["query text"],
|
||||
input_type: "search_query",
|
||||
});
|
||||
});
|
||||
|
||||
it("sends search_document for add and update actions", async () => {
|
||||
mockSend
|
||||
.mockResolvedValueOnce(cohereReply([[0.1]]))
|
||||
.mockResolvedValueOnce(cohereReply([[0.2]]));
|
||||
const embedder: Embedder = new AWSBedrockEmbedder({
|
||||
model: "cohere.embed-english-v3",
|
||||
});
|
||||
|
||||
await embedder.embed("doc one", "add");
|
||||
await embedder.embed("doc two", "update");
|
||||
|
||||
expect(requestBodyAt(0).input_type).toBe("search_document");
|
||||
expect(requestBodyAt(1).input_type).toBe("search_document");
|
||||
});
|
||||
|
||||
// Titan has no input_type concept; buildRequestBody() must not add one
|
||||
// even when a memoryAction is explicitly passed through.
|
||||
it("Titan ignores memoryAction and never sends input_type", async () => {
|
||||
mockSend.mockResolvedValueOnce(titanReply([1, 2]));
|
||||
const embedder: Embedder = new AWSBedrockEmbedder({});
|
||||
|
||||
await embedder.embed("hello", "search");
|
||||
|
||||
expect(requestBodyAt(0)).toEqual({ inputText: "hello" });
|
||||
});
|
||||
});
|
||||
|
||||
describe("Titan embedBatch concurrency", () => {
|
||||
afterEach(() => {
|
||||
// This test sets a persistent mockImplementation (not a *Once), so
|
||||
// clear it explicitly -- jest.clearAllMocks() in the top beforeEach
|
||||
// clears call data but not implementations.
|
||||
mockSend.mockReset();
|
||||
});
|
||||
|
||||
// F4: embedBatch() used to Promise.all-fan-out one InvokeModel call per
|
||||
// text with no cap, so a large batch could open hundreds of concurrent
|
||||
// requests at once. TITAN_MAX_CONCURRENCY bounds this to a small pool
|
||||
// while still preserving output order.
|
||||
it("never runs more than TITAN_MAX_CONCURRENCY Titan requests at once, and preserves order", async () => {
|
||||
let active = 0;
|
||||
let peak = 0;
|
||||
mockSend.mockImplementation(async (command: InvokeModelCommand) => {
|
||||
active++;
|
||||
peak = Math.max(peak, active);
|
||||
const body = JSON.parse(
|
||||
new TextDecoder().decode(command.input.body as Uint8Array),
|
||||
);
|
||||
await new Promise((resolve) => setTimeout(resolve, 10));
|
||||
active--;
|
||||
const index = Number(body.inputText.split("-")[1]);
|
||||
return titanReply([index]);
|
||||
});
|
||||
const embedder = new AWSBedrockEmbedder({});
|
||||
const texts = Array.from({ length: 10 }, (_, i) => `text-${i}`);
|
||||
|
||||
const embeddings = await embedder.embedBatch(texts);
|
||||
|
||||
expect(peak).toBeGreaterThan(1);
|
||||
expect(peak).toBeLessThanOrEqual(4);
|
||||
expect(embeddings).toEqual(texts.map((_, i) => [i]));
|
||||
expect(mockSend).toHaveBeenCalledTimes(10);
|
||||
});
|
||||
});
|
||||
});
|
||||
Reference in New Issue
Block a user