Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
114 changes: 97 additions & 17 deletions src/llm/bedrock/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -25,17 +25,21 @@ import { ChatBedrockConverse } from '@langchain/aws';
import { AIMessageChunk } from '@langchain/core/messages';
import { ChatGenerationChunk, ChatResult } from '@langchain/core/outputs';
import {
ConverseCommand,
ConverseStreamCommand,
type ConverseStreamOutput,
type GuardrailConfiguration,
type GuardrailStreamConfiguration,
} from '@aws-sdk/client-bedrock-runtime';
import type {
ConverseStreamOutput,
GuardrailConfiguration,
GuardrailStreamConfiguration,
} from '@aws-sdk/client-bedrock-runtime';
import type { CallbackManagerForLLMRun } from '@langchain/core/callbacks/manager';
import type { BaseMessage, ResponseMetadata } from '@langchain/core/messages';
import type { ChatBedrockConverseInput } from '@langchain/aws';
import type { SmoothItem } from '@/llm/stream/smoother';
import type { ContentBlockDeltaEvent } from './types';
import {
convertConverseMessageToLangChainMessage,
convertToConverseMessages,
createConverseToolUseStopChunk,
handleConverseStreamContentBlockStart,
Expand Down Expand Up @@ -68,6 +72,51 @@ type BedrockEmittedChunk = {
callbackToken: string;
};

function extractBedrockErrorMessage(error: unknown): string | undefined {
if (typeof error === 'string') {
return error;
}
if (error == null || typeof error !== 'object') {
return undefined;
}
if ('message' in error && typeof error.message === 'string') {
return error.message;
}
if ('Message' in error && typeof error.Message === 'string') {
return error.Message;
}
if ('errors' in error && Array.isArray(error.errors)) {
const nestedErrors: unknown[] = error.errors;
const messages: string[] = [];
for (const nestedError of nestedErrors) {
if (typeof nestedError === 'string') {
messages.push(nestedError);
} else if (
nestedError != null &&
typeof nestedError === 'object' &&
'message' in nestedError &&
typeof nestedError.message === 'string'
) {
messages.push(nestedError.message);
}
}
if (messages.length > 0) {
return messages.join('; ');
}
}
return undefined;
}

function normalizeBedrockError(error: unknown): Error {
if (error instanceof Error) {
return error;
}
const message =
extractBedrockErrorMessage(error) ??
'An error occurred while calling Bedrock Converse.';
return new Error(message, { cause: error });
}

/**
* Resolves the text a delta contributes to the smoothing cadence, preferring a
* text delta over a reasoning delta and ignoring non-string payloads.
Expand Down Expand Up @@ -235,26 +284,57 @@ export class CustomChatBedrockConverse extends ChatBedrockConverse {
}

/**
* Override _generateNonStreaming to use applicationInferenceProfile as modelId.
* Uses the same model-swapping pattern as streaming for consistency.
* Prepare model-aware replay with the configured model; the inference profile
* is only the wire target and must not mutate shared model state.
*/
override async _generateNonStreaming(
messages: BaseMessage[],
options: this['ParsedCallOptions'] & CustomChatBedrockConverseCallOptions,
runManager?: CallbackManagerForLLMRun
_runManager?: CallbackManagerForLLMRun
): Promise<ChatResult> {
const originalModel = this.model;
if (
this.applicationInferenceProfile != null &&
this.applicationInferenceProfile !== ''
) {
this.model = this.applicationInferenceProfile;
}

try {
return await super._generateNonStreaming(messages, options, runManager);
} finally {
this.model = originalModel;
const { converseMessages, converseSystem } = convertToConverseMessages(
messages,
{ model: this.model }
);
const modelId = this.getModelId();
const params = this.invocationParams(options);
applyCachePointsToConversePayload({
cacheControl: options.cache_control,
system: converseSystem,
messages: converseMessages,
params,
modelId,
});
const command = new ConverseCommand({
modelId,
messages: converseMessages,
...(Array.isArray(converseSystem) && converseSystem.length > 0
? { system: converseSystem }
: {}),
requestMetadata: options.requestMetadata,
...params,
});
const { output, ...responseMetadata } = await this.client.send(command, {
abortSignal: options.signal,
});
if (!output?.message) {
throw new Error('No message found in Bedrock response.');
}
const message = convertConverseMessageToLangChainMessage(
output.message,
responseMetadata
);
return {
generations: [
{
text: typeof message.content === 'string' ? message.content : '',
message,
},
],
};
} catch (error) {
throw normalizeBedrockError(error);
}
}

Expand Down
10 changes: 3 additions & 7 deletions src/llm/bedrock/inherited.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -657,12 +657,9 @@ describe('message output usage metadata conversion', () => {
responseMetadata
);

// Fork divergence: [email protected] folds cache read+write INTO input_tokens
// (would be 20). Our fork keeps input_tokens = raw inputTokens (10) and
// surfaces cache tokens only in input_token_details (Bedrock cache is
// additive, not a subset of input_tokens). Assert OURS.
// Match the non-streaming decoder's cache-inclusive input accounting.
expect(result.usage_metadata).toEqual({
input_tokens: 10,
input_tokens: 20,
output_tokens: 5,
total_tokens: 25,
input_token_details: {
Expand Down Expand Up @@ -709,8 +706,7 @@ describe('message output usage metadata conversion', () => {
);
const message = chunk.message as AIMessageChunk;

// Same fork divergence as the non-stream case: upstream would report
// input_tokens 35 (20+9+6); our fork keeps the raw 20.
// Streaming metadata keeps cache counts separate from input_tokens.
expect(message.usage_metadata).toEqual({
input_tokens: 20,
output_tokens: 4,
Expand Down
2 changes: 1 addition & 1 deletion src/llm/bedrock/llm.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -977,7 +977,7 @@ describe('convertConverseMessageToLangChainMessage - cache token extraction', ()
);

expect(result.usage_metadata).toEqual({
input_tokens: 20,
input_tokens: 10851,
output_tokens: 5,
total_tokens: 10856,
input_token_details: {
Expand Down
Loading