diff --git a/src/renderer/modals/SessionSettings.tsx b/src/renderer/modals/SessionSettings.tsx index 29f7241c01..b26ae79efc 100644 --- a/src/renderer/modals/SessionSettings.tsx +++ b/src/renderer/modals/SessionSettings.tsx @@ -20,6 +20,7 @@ import { AssistantAvatar } from '@/components/common/Avatar' import LazyNumberInput from '@/components/common/LazyNumberInput' import MaxContextMessageCountSlider from '@/components/common/MaxContextMessageCountSlider' import { ScalableIcon } from '@/components/common/ScalableIcon' +import SegmentedControl from '@/components/common/SegmentedControl' import SliderWithInput from '@/components/common/SliderWithInput' import { handleImageInputAndSave, ImageInStorage } from '@/components/Image' import ImageStyleSelect from '@/components/ImageStyleSelect' @@ -480,6 +481,32 @@ export function ChatConfig({ /> + {settings?.provider === ModelProviderEnum.Claude && ( + + + {t('Prompt Cache')} + + + onSettingsChange({ + providerOptions: { + ...settings?.providerOptions, + claude: { + ...settings?.providerOptions?.claude, + cacheTTL: value as 'auto' | '5m' | '1h', + }, + }, + }) + } + data={[ + { label: t('Auto'), value: 'auto' }, + { label: t('5 min'), value: '5m' }, + { label: t('1 hour'), value: '1h' }, + ]} + /> + + )} ) } diff --git a/src/shared/models/anthropic-cache.ts b/src/shared/models/anthropic-cache.ts index 8738a95fec..1192ac094a 100644 --- a/src/shared/models/anthropic-cache.ts +++ b/src/shared/models/anthropic-cache.ts @@ -10,7 +10,7 @@ import type { ModelMessage } from 'ai' * Works with both direct Anthropic API and AWS Bedrock. * See: https://docs.anthropic.com/en/docs/build-with-claude/prompt-caching */ -export function addAnthropicCacheControl(messages: ModelMessage[]): ModelMessage[] { +export function addAnthropicCacheControl(messages: ModelMessage[], ttl: '5m' | '1h' = '5m'): ModelMessage[] { if (messages.length === 0) { return messages } diff --git a/src/shared/providers/definitions/claude.ts b/src/shared/providers/definitions/claude.ts index 9e45369559..33ec94dd39 100644 --- a/src/shared/providers/definitions/claude.ts +++ b/src/shared/providers/definitions/claude.ts @@ -90,6 +90,7 @@ export const claudeProvider = defineProvider({ topP: config.settings.topP, maxOutputTokens: config.settings.maxTokens, stream: config.settings.stream, + cacheTTL: config.settings.providerOptions?.claude?.cacheTTL, extraHeaders: oauthHeaders, customFetch: isOAuth && credentialManager ? createBearerOAuthFetch(config.dependencies, credentialManager) : undefined, diff --git a/src/shared/providers/definitions/models/claude.ts b/src/shared/providers/definitions/models/claude.ts index bebdb186d6..fd02afe0b3 100644 --- a/src/shared/providers/definitions/models/claude.ts +++ b/src/shared/providers/definitions/models/claude.ts @@ -17,6 +17,7 @@ interface Options { topP?: number maxOutputTokens?: number stream?: boolean + cacheTTL?: 'auto' | '5m' | '1h' extraHeaders?: Record customFetch?: typeof globalThis.fetch authToken?: string @@ -132,14 +133,30 @@ export default class Claude extends AbstractAISDKModel { } public async chat(messages: ModelMessage[], options: CallChatCompletionOptions): Promise { - return super.chat(addAnthropicCacheControl(messages), options) + const ttl = + this.options.cacheTTL === '5m' + ? '5m' + : this.options.cacheTTL === '1h' + ? '1h' + : messages.length >= 10 + ? '1h' + : '5m' + return super.chat(addAnthropicCacheControl(messages, ttl), options) } public async *chatStream( messages: ModelMessage[], options: ChatStreamOptions ): AsyncGenerator> { - yield* super.chatStream(addAnthropicCacheControl(messages), options) + const ttl = + this.options.cacheTTL === '5m' + ? '5m' + : this.options.cacheTTL === '1h' + ? '1h' + : messages.length >= 10 + ? '1h' + : '5m' + yield* super.chatStream(addAnthropicCacheControl(messages, ttl), options) } // https://docs.anthropic.com/en/docs/api/models diff --git a/src/shared/types/settings.ts b/src/shared/types/settings.ts index a6b5a6ea2e..67b83c2686 100644 --- a/src/shared/types/settings.ts +++ b/src/shared/types/settings.ts @@ -118,6 +118,8 @@ const ProviderBaseInfoSchema = z.discriminatedUnion('isCustom', [ CustomProviderBaseInfoSchema, ]) +export const ClaudeCacheTTLSchema = z.enum(['auto', '5m', '1h']) + const ClaudeParamsSchema = z.object({ thinking: z .object({ @@ -127,6 +129,7 @@ const ClaudeParamsSchema = z.object({ .optional() .catch(undefined), effort: z.enum(['low', 'medium', 'high', 'xhigh', 'max']).optional().catch(undefined), + cacheTTL: ClaudeCacheTTLSchema.optional().catch('auto'), }) const OpenAIParamsSchema = z.object({