mirror of
https://github.com/strands-agents/harness-sdk.git
synced 2026-10-02 02:44:48 +08:00
fix: cache unsupported models for bedrocks token counting (#999)
This commit is contained in:
@@ -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)
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user