fix: cache unsupported models for bedrocks token counting (#999)

This commit is contained in:
opieter-aws
2026-05-06 16:11:46 +00:00
committed by GitHub
parent 2614968b97
commit b21144c032
2 changed files with 68 additions and 1 deletions
@@ -4168,6 +4168,7 @@ describe('BedrockModel', () => {
beforeEach(() => {
vi.clearAllMocks()
BedrockModel.clearCountTokensCache()
})
it('should return native token count on success', async () => {
@@ -4274,5 +4275,38 @@ describe('BedrockModel', () => {
expect(typeof result).toBe('number')
expect(result).toBeGreaterThanOrEqual(0)
})
it('should cache model ID and skip API call when model does not support counting tokens', async () => {
const unsupportedError = new Error("The provided model doesn't support counting tokens")
unsupportedError.name = 'ValidationException'
const mockSend = vi.fn(async () => {
throw unsupportedError
})
mockBedrockClientImplementation({ send: mockSend })
const model = new BedrockModel()
// First call: hits API, gets error, caches
await model.countTokens(messages)
expect(mockSend).toHaveBeenCalledOnce()
// Second call: skips API entirely
await model.countTokens(messages)
expect(mockSend).toHaveBeenCalledOnce()
})
it('should not cache model ID for other errors', async () => {
const mockSend = vi.fn(async () => {
throw new Error('Transient network error')
})
mockBedrockClientImplementation({ send: mockSend })
const model = new BedrockModel()
await model.countTokens(messages)
expect(mockSend).toHaveBeenCalledTimes(1)
// Second call should still attempt the API
await model.countTokens(messages)
expect(mockSend).toHaveBeenCalledTimes(2)
})
})
})
+34 -1
View File
@@ -96,6 +96,12 @@ const BEDROCK_CONTEXT_WINDOW_OVERFLOW_MESSAGES = [
'prompt is too long',
]
/**
* Cache of model IDs that do not support the CountTokens API.
* Prevents repeated failing API calls for models that will never support token counting.
*/
const UNSUPPORTED_COUNT_TOKENS_MODELS = new Set<string>()
/**
* Mapping of Bedrock stop reasons to SDK stop reasons.
*/
@@ -333,6 +339,16 @@ export class BedrockModel extends Model<BedrockModelConfig> {
private _config: BedrockModelConfig
private _client: BedrockRuntimeClient
/**
* Clears the cache of model IDs that do not support the CountTokens API.
* After calling this, the next countTokens invocation will attempt the API again.
*
* @internal
*/
static clearCountTokensCache(): void {
UNSUPPORTED_COUNT_TOKENS_MODELS.clear()
}
/**
* Creates a new BedrockModel instance.
*
@@ -485,6 +501,12 @@ export class BedrockModel extends Model<BedrockModelConfig> {
* @returns Total input token count
*/
override async countTokens(messages: Message[], options?: CountTokensOptions): Promise<number> {
const modelId = this._config.modelId ?? MODEL_DEFAULTS.bedrock.modelId
if (UNSUPPORTED_COUNT_TOKENS_MODELS.has(modelId)) {
return super.countTokens(messages, options)
}
try {
const request = this._formatRequest(messages, options)
const converseInput: Record<string, unknown> = {}
@@ -506,7 +528,18 @@ export class BedrockModel extends Model<BedrockModelConfig> {
logger.debug(`total_tokens=<${response.inputTokens}> | native token count`)
return response.inputTokens
} catch (error) {
logger.debug(`error=<${error}> | native token counting failed, falling back to estimation`)
if (
error instanceof Error &&
error.name === 'ValidationException' &&
error.message.includes("doesn't support counting tokens")
) {
logger.debug(
`model_id=<${modelId}> | model does not support CountTokens, caching for future calls, falling back to estimation`
)
UNSUPPORTED_COUNT_TOKENS_MODELS.add(modelId)
} else {
logger.debug(`error=<${error}> | native token counting failed, falling back to estimation`)
}
return super.countTokens(messages, options)
}
}