diff --git a/cmd/thv/app/llm.go b/cmd/thv/app/llm.go index c94ff8fdbc..8931049adb 100644 --- a/cmd/thv/app/llm.go +++ b/cmd/thv/app/llm.go @@ -93,11 +93,15 @@ func addLLMConnectionFlags(cmd *cobra.Command, opts *llm.SetOptions) { // field unchanged (nil pointer = "not provided"). Shared by "config set" and // "setup" so both commands treat these flags identically. func applyChangedLLMFlags( - cmd *cobra.Command, opts *llm.SetOptions, tlsSkipVerify, bedrockCompat, enable1M bool, models []string, + cmd *cobra.Command, opts *llm.SetOptions, + tlsSkipVerify, shortPromptCache, bedrockCompat, enable1M bool, models []string, ) { if cmd.Flags().Changed("tls-skip-verify") { opts.TLSSkipVerify = &tlsSkipVerify } + if cmd.Flags().Changed("short-prompt-cache") { + opts.ShortPromptCache = &shortPromptCache + } if cmd.Flags().Changed("bedrock-compat") { opts.BedrockCompat = &bedrockCompat } @@ -111,11 +115,12 @@ func applyChangedLLMFlags( func newConfigSetCommand() *cobra.Command { var ( - opts llm.SetOptions - tlsSkipVerify bool - bedrockCompat bool - enable1M bool - models []string + opts llm.SetOptions + tlsSkipVerify bool + shortPromptCache bool + bedrockCompat bool + enable1M bool + models []string ) cmd := &cobra.Command{ @@ -130,7 +135,7 @@ Example: --client-id my-client-id`, Args: cobra.NoArgs, RunE: func(cmd *cobra.Command, _ []string) error { - applyChangedLLMFlags(cmd, &opts, tlsSkipVerify, bedrockCompat, enable1M, models) + applyChangedLLMFlags(cmd, &opts, tlsSkipVerify, shortPromptCache, bedrockCompat, enable1M, models) return config.UpdateConfig(func(c *config.Config) error { return c.LLM.SetFields(opts) }) @@ -140,6 +145,9 @@ Example: addLLMConnectionFlags(cmd, &opts) cmd.Flags().BoolVar(&tlsSkipVerify, "tls-skip-verify", false, "Skip TLS certificate verification for the upstream gateway (local dev only; use --tls-skip-verify=false to clear)") + cmd.Flags().BoolVar(&shortPromptCache, "short-prompt-cache", false, + "Use Claude Code's five-minute prompt cache instead of ToolHive's one-hour default. "+ + "Applied by \"thv llm setup\"; use --short-prompt-cache=false to restore one-hour caching.") cmd.Flags().BoolVar(&bedrockCompat, "bedrock-compat", false, "Persist Bedrock compatibility for Claude Code (CLAUDE_CODE_DISABLE_EXPERIMENTAL_BETAS=1 + per-tier "+ "Bedrock model IDs). Applied by \"thv llm setup\". Use --bedrock-compat=false to clear.") @@ -268,6 +276,7 @@ func newLLMSetupCommand() *cobra.Command { var ( opts llm.SetOptions tlsSkipVerify bool + shortPromptCache bool bedrockCompat bool enable1M bool targetClient string @@ -323,7 +332,7 @@ Re-running is idempotent and uses the cached token (no browser prompt). Run "thv llm teardown" to revert all changes.`, Args: cobra.NoArgs, RunE: func(cmd *cobra.Command, _ []string) error { - applyChangedLLMFlags(cmd, &opts, tlsSkipVerify, bedrockCompat, enable1M, models) + applyChangedLLMFlags(cmd, &opts, tlsSkipVerify, shortPromptCache, bedrockCompat, enable1M, models) cm, err := client.NewClientManager() if err != nil { return fmt.Errorf("initializing client manager: %w", err) @@ -345,6 +354,9 @@ Run "thv llm teardown" to revert all changes.`, "For direct-mode tools (Claude Code, Gemini CLI) this sets NODE_TLS_REJECT_UNAUTHORIZED=0, "+ "disabling TLS for ALL of that tool's outbound connections. "+ "For proxy-mode tools only the proxy-to-gateway connection is affected.") + cmd.Flags().BoolVar(&shortPromptCache, "short-prompt-cache", false, + "Use Claude Code's five-minute prompt cache instead of ToolHive's one-hour default. Persisted, so a later "+ + "plain \"thv llm setup\" keeps it; use --short-prompt-cache=false to restore one-hour caching.") cmd.Flags().StringVar(&anthropicPathPrefix, "anthropic-path-prefix", "", "Path prefix appended to the gateway URL when writing ANTHROPIC_BASE_URL for direct-mode tools "+ "(e.g. /anthropic). When omitted, the gateway is probed automatically.") @@ -498,6 +510,10 @@ func (a *clientManagerAdapter) LLMGatewayModeFor(clientType string) string { return a.cm.LLMGatewayModeFor(client.ClientApp(clientType)) } +func (a *clientManagerAdapter) PromptCacheConflict(clientType string) (string, error) { + return a.cm.PromptCacheConflict(client.ClientApp(clientType)) +} + func (a *clientManagerAdapter) IsManaged(clientType string) bool { return a.cm.IsManaged(client.ClientApp(clientType)) } diff --git a/cmd/thv/app/llm_test.go b/cmd/thv/app/llm_test.go index dccecc0cc1..c51f443f52 100644 --- a/cmd/thv/app/llm_test.go +++ b/cmd/thv/app/llm_test.go @@ -7,6 +7,7 @@ import ( "bytes" "context" "errors" + "fmt" "os" "path/filepath" "runtime" @@ -47,6 +48,43 @@ func llmProvider(t *testing.T, llmCfg llm.Config) config.Provider { // Use it in tests that don't exercise the authentication path. var noopLogin llm.LoginFunc = func(context.Context, *llm.Config) error { return nil } +func TestConfigSetCommand_ShortPromptCacheFlagWiring(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + flagValue *bool + }{ + {name: "omitted"}, + {name: "enabled", flagValue: boolPtr(true)}, + {name: "explicitly disabled", flagValue: boolPtr(false)}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + cmd := newConfigSetCommand() + flag := cmd.Flags().Lookup("short-prompt-cache") + require.NotNil(t, flag) + assert.Nil(t, cmd.Flags().Lookup("extended-ttl-cache")) + if tt.flagValue != nil { + require.NoError(t, cmd.Flags().Set("short-prompt-cache", fmt.Sprintf("%t", *tt.flagValue))) + } + + parsed, err := cmd.Flags().GetBool("short-prompt-cache") + require.NoError(t, err) + var opts llm.SetOptions + applyChangedLLMFlags(cmd, &opts, false, parsed, false, false, nil) + + if tt.flagValue == nil { + assert.Nil(t, opts.ShortPromptCache) + return + } + require.NotNil(t, opts.ShortPromptCache) + assert.Equal(t, *tt.flagValue, *opts.ShortPromptCache) + }) + } +} + // errOnUpdateProvider wraps a base Provider but returns a fixed error from // UpdateConfig. Used to inject deterministic failures without relying on // filesystem permission tricks that are unreliable on Windows. diff --git a/docs/cli/thv_llm_config_set.md b/docs/cli/thv_llm_config_set.md index c732df1cbd..67c2df590e 100644 --- a/docs/cli/thv_llm_config_set.md +++ b/docs/cli/thv_llm_config_set.md @@ -40,6 +40,7 @@ thv llm config set [flags] --issuer string OIDC issuer URL --models strings Model IDs to persist and apply during "thv llm setup", comma-separated or by repeating the flag, e.g. --models=us.anthropic.claude-opus-4-8,us.anthropic.claude-sonnet-5. Credential-helper clients (Claude Desktop) write these as inferenceModels; with Bedrock compat, each ID is also mapped to a Claude Code tier by matching 'haiku', 'opus', or 'sonnet' in the ID. --proxy-port int Localhost proxy listen port (omit to keep current; default: 14000) + --short-prompt-cache Use Claude Code's five-minute prompt cache instead of ToolHive's one-hour default. Applied by "thv llm setup"; use --short-prompt-cache=false to restore one-hour caching. --tls-skip-verify Skip TLS certificate verification for the upstream gateway (local dev only; use --tls-skip-verify=false to clear) ``` diff --git a/docs/cli/thv_llm_setup.md b/docs/cli/thv_llm_setup.md index c2c0e50d45..dfdf7a02aa 100644 --- a/docs/cli/thv_llm_setup.md +++ b/docs/cli/thv_llm_setup.md @@ -77,6 +77,7 @@ thv llm setup [flags] --lazy Skip the interactive OIDC login and defer it until the first time a configured tool accesses the gateway. Tool config and persisted settings are written normally. Useful for unattended provisioning (e.g. an MDM profile). --models strings Model IDs to configure, comma-separated or by repeating the flag, e.g. --models=us.anthropic.claude-opus-4-8,us.anthropic.claude-sonnet-5. For credential-helper clients (Claude Desktop) these become inferenceModels. With --bedrock-compat, each ID is also mapped to a Claude Code tier by matching 'haiku', 'opus', or 'sonnet' in the ID (IDs matching no tier are ignored with a warning). Omit to use the built-in Bedrock defaults / gateway model discovery. --proxy-port int Localhost proxy listen port (omit to keep current; default: 14000) + --short-prompt-cache Use Claude Code's five-minute prompt cache instead of ToolHive's one-hour default. Persisted, so a later plain "thv llm setup" keeps it; use --short-prompt-cache=false to restore one-hour caching. --skip-browser Print the OIDC authorization URL instead of opening a browser, then wait for the callback. Use in headless/SSH/CI environments where no system browser is available. --tls-skip-verify Skip TLS certificate verification for the upstream gateway (local dev only). For direct-mode tools (Claude Code, Gemini CLI) this sets NODE_TLS_REJECT_UNAUTHORIZED=0, disabling TLS for ALL of that tool's outbound connections. For proxy-mode tools only the proxy-to-gateway connection is affected. ``` diff --git a/pkg/client/config.go b/pkg/client/config.go index 2505339c29..81a6b0d4c8 100644 --- a/pkg/client/config.go +++ b/pkg/client/config.go @@ -167,10 +167,11 @@ const ( // - ValueField names which ApplyConfig field to write. Valid values: // "GatewayURL", "AnthropicBaseURL", "ProxyBaseURL", "ProxyOrigin", // "TokenHelperCommand", "PlaceholderAPIKey", "ClaudeCodeHelperTTLMillis", -// "NodeTLSRejectUnauthorized", "BedrockDisableExperimentalBetas", -// "BedrockHaikuModel", "BedrockOpusModel", "BedrockSonnetModel". An -// unrecognised ValueField is a programming error and causes -// ConfigureLLMGateway to return an error. +// "NodeTLSRejectUnauthorized", "PromptCacheTTL", +// "PromptCache1HLegacy", +// "BedrockDisableExperimentalBetas", "BedrockHaikuModel", +// "BedrockOpusModel", "BedrockSonnetModel". An unrecognised ValueField is +// a programming error and causes ConfigureLLMGateway to return an error. // - Literal is written verbatim into the settings key (e.g. a fixed auth // type string). Use Literal instead of ValueField for constant values so // that typos in ValueField are caught as errors rather than silently @@ -184,8 +185,9 @@ type LLMGatewayKeySpec struct { JSONPointer string // RFC 6901 path // ValueField: "GatewayURL" | "AnthropicBaseURL" | "ProxyBaseURL" | "ProxyOrigin" | // "TokenHelperCommand" | "PlaceholderAPIKey" | "ClaudeCodeHelperTTLMillis" | - // "NodeTLSRejectUnauthorized" | "BedrockDisableExperimentalBetas" | - // "BedrockHaikuModel" | "BedrockOpusModel" | "BedrockSonnetModel" + // "NodeTLSRejectUnauthorized" | "PromptCacheTTL" | "PromptCache1HLegacy" | + // "BedrockDisableExperimentalBetas" | "BedrockHaikuModel" | + // "BedrockOpusModel" | "BedrockSonnetModel" ValueField string Literal string // constant value written verbatim; mutually exclusive with ValueField ClearWhenEmpty bool // remove the key when the resolved value is empty (ignored for Literal) @@ -549,6 +551,12 @@ var supportedClientIntegrations = []clientAppConfig{ // NODE_TLS_REJECT_UNAUTHORIZED is only written when --tls-skip-verify is set. // ClearWhenEmpty ensures it is removed when the flag is later cleared. {JSONPointer: "/env/NODE_TLS_REJECT_UNAUTHORIZED", ValueField: "NodeTLSRejectUnauthorized", ClearWhenEmpty: true}, + // Current Claude Code versions expose separate controls for the main + // conversation and auxiliary requests. ENABLE_PROMPT_CACHING_1H is the + // fallback for versions that predate those per-bucket settings. + {JSONPointer: "/promptCacheTtl", ValueField: "PromptCacheTTL", ClearWhenEmpty: true}, + {JSONPointer: "/subagentPromptCacheTtl", ValueField: "PromptCacheTTL", ClearWhenEmpty: true}, + {JSONPointer: "/env/ENABLE_PROMPT_CACHING_1H", ValueField: "PromptCache1HLegacy", ClearWhenEmpty: true}, // Bedrock-compat keys (written only with --bedrock-compat). Bedrock rejects // Claude Code's experimental anthropic-beta headers, so betas are disabled; // the per-tier model IDs pin Bedrock inference-profile IDs. All use diff --git a/pkg/client/llm_gateway.go b/pkg/client/llm_gateway.go index a5cdd093fd..8ec942ed22 100644 --- a/pkg/client/llm_gateway.go +++ b/pkg/client/llm_gateway.go @@ -9,6 +9,7 @@ import ( "log/slog" "os" "path/filepath" + "runtime" "strconv" "strings" @@ -70,9 +71,9 @@ func (cm *ClientManager) ConfigureLLMGateway(clientType ClientApp, cfg llmgatewa // Parse with hujson first so that JSONC (comments, trailing commas) is // handled correctly for all subsequent operations. - v, err := hujson.Parse(content) + v, err := parseJSONC(content, path) if err != nil { - return fmt.Errorf("parsing %s: %w", path, err) + return err } if err := applyLLMGatewayKeys(&v, appCfg.LLMGatewayKeys, cfg, path); err != nil { @@ -122,9 +123,9 @@ func applyLLMGatewayKeys(v *hujson.Value, specs []LLMGatewayKeySpec, cfg llmgate } // Standardize once for existence checks in the remove path. - standardized, err := hujson.Standardize(v.Pack()) + standardized, err := standardizeJSONC(v, filePath) if err != nil { - return fmt.Errorf("standardizing %s: %w", filePath, err) + return err } for i, spec := range specs { @@ -221,15 +222,15 @@ func revertJSONPointerGateway(appCfg *clientAppConfig, configPath string) error return nil } - v, err := hujson.Parse(content) + v, err := parseJSONC(content, configPath) if err != nil { - return fmt.Errorf("parsing %s: %w", configPath, err) + return err } // Standardize once for all existence checks below. - standardized, err := hujson.Standardize(v.Pack()) + standardized, err := standardizeJSONC(&v, configPath) if err != nil { - return fmt.Errorf("standardizing %s: %w", configPath, err) + return err } for _, spec := range appCfg.LLMGatewayKeys { @@ -260,6 +261,174 @@ func (cm *ClientManager) IsLLMGatewaySupported(clientType ClientApp) bool { return cfg != nil && cfg.LLMGatewayMode != "" } +var promptCacheEnvironmentControls = [...]string{ + "FORCE_PROMPT_CACHING_5M", + "CLAUDE_CODE_PROMPT_CACHE_TTL", + "CLAUDE_CODE_SUBAGENT_PROMPT_CACHE_TTL", +} + +// PromptCacheConflict reports a locally discoverable setting that may override +// ToolHive's one-hour Claude Code prompt-cache configuration. +func (*ClientManager) PromptCacheConflict(clientType ClientApp) (string, error) { + if clientType != ClaudeCode { + return "", nil + } + return promptCacheConflict(claudeCodeManagedSettingsDir(runtime.GOOS), os.LookupEnv) +} + +type promptCacheControl struct { + value string + source string +} + +func promptCacheConflict( + managedDir string, lookupEnv func(string) (string, bool), +) (string, error) { + controls := make(map[string]promptCacheControl, 5) + settingsPaths, err := managedPromptCacheSettingsPaths(managedDir) + if err != nil { + return "", err + } + for _, path := range settingsPaths { + if err := mergePromptCacheControls(path, controls); err != nil { + return "", err + } + } + for _, name := range promptCacheEnvironmentControls { + if value, ok := lookupEnv(name); ok { + controls[name] = promptCacheControl{strings.TrimSpace(value), "the process environment"} + } + } + return promptCacheConflictDescription(controls), nil +} + +func managedPromptCacheSettingsPaths(managedDir string) ([]string, error) { + if managedDir == "" { + return nil, nil + } + paths := []string{filepath.Join(managedDir, "managed-settings.json")} + dropInDir := filepath.Join(managedDir, "managed-settings.d") + dropIns, err := os.ReadDir(dropInDir) + if err != nil { + if os.IsNotExist(err) { + return paths, nil + } + return nil, fmt.Errorf("reading Claude Code managed settings drop-in directory: %w", err) + } + for _, entry := range dropIns { + if entry.IsDir() || strings.HasPrefix(entry.Name(), ".") || filepath.Ext(entry.Name()) != ".json" { + continue + } + paths = append(paths, filepath.Join(dropInDir, entry.Name())) + } + return paths, nil +} + +func mergePromptCacheControls(path string, controls map[string]promptCacheControl) error { + content, err := os.ReadFile(path) // #nosec G304 -- paths are registered client settings locations + if err != nil { + if os.IsNotExist(err) { + return nil + } + return fmt.Errorf("reading %s: %w", path, err) + } + if len(content) == 0 { + return nil + } + + v, err := parseJSONC(content, path) + if err != nil { + return err + } + standardized, err := standardizeJSONC(&v, path) + if err != nil { + return err + } + var settings map[string]any + if err := json.Unmarshal(standardized, &settings); err != nil { + return fmt.Errorf("decoding %s: %w", path, err) + } + env, _ := settings["env"].(map[string]any) + for _, name := range promptCacheEnvironmentControls { + value, ok := env[name].(string) + if !ok { + continue + } + controls[name] = promptCacheControl{strings.TrimSpace(value), path} + } + for _, name := range []string{"promptCacheTtl", "subagentPromptCacheTtl"} { + if value, ok := settings[name].(string); ok { + controls[name] = promptCacheControl{strings.TrimSpace(value), path} + } + } + return nil +} + +func promptCacheConflictDescription(controls map[string]promptCacheControl) string { + if force, ok := controls["FORCE_PROMPT_CACHING_5M"]; ok && force.value == "1" { + return formatPromptCacheConflict("FORCE_PROMPT_CACHING_5M", force) + } + + buckets := []struct { + environmentVariable string + setting string + }{ + {environmentVariable: "CLAUDE_CODE_PROMPT_CACHE_TTL", setting: "promptCacheTtl"}, + {environmentVariable: "CLAUDE_CODE_SUBAGENT_PROMPT_CACHE_TTL", setting: "subagentPromptCacheTtl"}, + } + for _, bucket := range buckets { + if envValue, ok := controls[bucket.environmentVariable]; ok { + switch strings.ToLower(envValue.value) { + case "5m": + return formatPromptCacheConflict(bucket.environmentVariable, envValue) + case "1h": + continue + } + } + if setting, ok := controls[bucket.setting]; ok && strings.EqualFold(setting.value, "5m") { + return formatPromptCacheConflict(bucket.setting, setting) + } + } + return "" +} + +func formatPromptCacheConflict(name string, control promptCacheControl) string { + return fmt.Sprintf("%s=%s in %s", name, control.value, control.source) +} + +func parseJSONC(content []byte, path string) (hujson.Value, error) { + v, err := hujson.Parse(content) + if err != nil { + return hujson.Value{}, fmt.Errorf("parsing %s: %w", path, err) + } + return v, nil +} + +func standardizeJSONC(v *hujson.Value, path string) ([]byte, error) { + standardized, err := hujson.Standardize(v.Pack()) + if err != nil { + return nil, fmt.Errorf("standardizing %s: %w", path, err) + } + return standardized, nil +} + +func claudeCodeManagedSettingsDir(goos string) string { + switch goos { + case "darwin": + return filepath.Join(string(filepath.Separator), "Library", "Application Support", "ClaudeCode") + case "linux": + return filepath.Join(string(filepath.Separator), "etc", "claude-code") + case "windows": + programFiles := os.Getenv("ProgramFiles") + if programFiles == "" { + programFiles = `C:\Program Files` + } + return filepath.Join(programFiles, "ClaudeCode") + default: + return "" + } +} + // IsManaged reports whether an MDM/managed-preferences profile is present for // the given client. When true, the client reads config from the managed profile // and ignores the local config "thv llm setup" writes, so setup warns the user. @@ -415,6 +584,16 @@ func resolveApplyConfigField(valueField string, cfg llmgateway.ApplyConfig) (str return "0", true } return "", true + case "PromptCacheTTL": + if !cfg.ShortPromptCache { + return "1h", true + } + return "", true + case "PromptCache1HLegacy": + if !cfg.ShortPromptCache { + return "1", true + } + return "", true default: return resolveBedrockField(valueField, cfg) } @@ -470,9 +649,9 @@ func ensureLLMAncestors(v *hujson.Value, ptr, filePath string) error { return nil // top-level key — no ancestors to create } // Standardize once for all existence checks in this call. - standardized, err := hujson.Standardize(v.Pack()) + standardized, err := standardizeJSONC(v, filePath) if err != nil { - return fmt.Errorf("standardizing JSON in %s: %w", filePath, err) + return err } ancestor := "" diff --git a/pkg/client/llm_gateway_test.go b/pkg/client/llm_gateway_test.go index 9f7b0b08f4..6b799eef7b 100644 --- a/pkg/client/llm_gateway_test.go +++ b/pkg/client/llm_gateway_test.go @@ -395,6 +395,191 @@ func TestConfigureLLMGateway_ClaudeCodeBedrock(t *testing.T) { }) } +func TestConfigureLLMGateway_ClaudeCodePromptCache(t *testing.T) { + t.Parallel() + + cachePointers := map[string]string{ + "/promptCacheTtl": "1h", + "/subagentPromptCacheTtl": "1h", + "/env/ENABLE_PROMPT_CACHING_1H": "1", + } + baseCfg := llmgateway.ApplyConfig{ + GatewayURL: "https://gw.example.com", + TokenHelperCommand: `thv llm token`, + } + + t.Run("default writes both request buckets and legacy fallback", func(t *testing.T) { + t.Parallel() + home := t.TempDir() + cm := NewTestClientManager(home, nil, supportedClientIntegrations, nil) + require.NoError(t, os.MkdirAll(filepath.Join(home, ".claude"), 0o700)) + + path, err := cm.ConfigureLLMGateway(ClaudeCode, baseCfg) + require.NoError(t, err) + + data, err := os.ReadFile(path) + require.NoError(t, err) + for ptr, want := range cachePointers { + got, ok := jsonPointerGet(data, ptr) + assert.True(t, ok, "pointer %q missing", ptr) + assert.Equal(t, want, got, "wrong value at %q", ptr) + } + }) + + t.Run("short opt-out removes keys and explicit false restores them", func(t *testing.T) { + t.Parallel() + home := t.TempDir() + cm := NewTestClientManager(home, nil, supportedClientIntegrations, nil) + require.NoError(t, os.MkdirAll(filepath.Join(home, ".claude"), 0o700)) + + path, err := cm.ConfigureLLMGateway(ClaudeCode, baseCfg) + require.NoError(t, err) + shortCfg := baseCfg + shortCfg.ShortPromptCache = true + _, err = cm.ConfigureLLMGateway(ClaudeCode, shortCfg) + require.NoError(t, err) + + data, err := os.ReadFile(path) + require.NoError(t, err) + for ptr := range cachePointers { + _, ok := jsonPointerGet(data, ptr) + assert.False(t, ok, "pointer %q should be absent with the short-cache opt-out", ptr) + } + + _, err = cm.ConfigureLLMGateway(ClaudeCode, baseCfg) + require.NoError(t, err) + data, err = os.ReadFile(path) + require.NoError(t, err) + for ptr, want := range cachePointers { + got, ok := jsonPointerGet(data, ptr) + assert.True(t, ok, "pointer %q missing after restoring the default", ptr) + assert.Equal(t, want, got, "wrong value at %q", ptr) + } + }) + + t.Run("teardown removes all prompt cache keys", func(t *testing.T) { + t.Parallel() + home := t.TempDir() + cm := NewTestClientManager(home, nil, supportedClientIntegrations, nil) + require.NoError(t, os.MkdirAll(filepath.Join(home, ".claude"), 0o700)) + + path, err := cm.ConfigureLLMGateway(ClaudeCode, baseCfg) + require.NoError(t, err) + require.NoError(t, cm.RevertLLMGateway(ClaudeCode, path)) + + data, err := os.ReadFile(path) + require.NoError(t, err) + for ptr := range cachePointers { + _, ok := jsonPointerGet(data, ptr) + assert.False(t, ok, "pointer %q should be absent after teardown", ptr) + } + }) +} + +func TestPromptCacheConflict(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + files map[string]string + env map[string]string + want string + wantErr string + }{ + {name: "no override"}, + { + name: "process global five-minute override", + env: map[string]string{"FORCE_PROMPT_CACHING_5M": "1"}, + want: "FORCE_PROMPT_CACHING_5M=1 in the process environment", + }, + { + name: "process bucket five-minute override", + env: map[string]string{"CLAUDE_CODE_PROMPT_CACHE_TTL": "5m"}, + want: "CLAUDE_CODE_PROMPT_CACHE_TTL=5m in the process environment", + }, + { + name: "managed top-level override", + files: map[string]string{"managed-settings.json": `{"promptCacheTtl":"5m"}`}, + want: "promptCacheTtl=5m", + }, + { + name: "managed drop-in environment override", + files: map[string]string{ + "managed-settings.d/10-cache.json": `{"env":{"CLAUDE_CODE_SUBAGENT_PROMPT_CACHE_TTL":"5m"}}`, + }, + want: "CLAUDE_CODE_SUBAGENT_PROMPT_CACHE_TTL=5m", + }, + { + name: "later managed drop-in restores one hour", + files: map[string]string{ + "managed-settings.json": `{"promptCacheTtl":"5m"}`, + "managed-settings.d/10-five-minutes.json": `{"promptCacheTtl":"5m"}`, + "managed-settings.d/20-one-hour.json": `{"promptCacheTtl":"1h"}`, + }, + }, + { + name: "process one-hour bucket overrides managed five minutes", + files: map[string]string{"managed-settings.json": `{"promptCacheTtl":"5m"}`}, + env: map[string]string{"CLAUDE_CODE_PROMPT_CACHE_TTL": "1h"}, + }, + { + name: "malformed managed settings", + files: map[string]string{"managed-settings.json": `{`}, + wantErr: "parsing", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + managedDir := t.TempDir() + for relativePath, contents := range tt.files { + writePromptCacheSettings(t, filepath.Join(managedDir, filepath.FromSlash(relativePath)), contents) + } + lookupEnv := func(name string) (string, bool) { + value, ok := tt.env[name] + return value, ok + } + + got, err := promptCacheConflict(managedDir, lookupEnv) + if tt.wantErr != "" { + require.ErrorContains(t, err, tt.wantErr) + return + } + require.NoError(t, err) + if tt.want == "" { + assert.Empty(t, got) + } else { + assert.Contains(t, got, tt.want) + } + }) + } +} + +func TestManagedPromptCacheSettingsPaths(t *testing.T) { + t.Parallel() + + managedDir := t.TempDir() + writePromptCacheSettings(t, filepath.Join(managedDir, "managed-settings.d", "20-later.json"), `{}`) + writePromptCacheSettings(t, filepath.Join(managedDir, "managed-settings.d", "10-earlier.json"), `{}`) + writePromptCacheSettings(t, filepath.Join(managedDir, "managed-settings.d", ".hidden.json"), `{}`) + writePromptCacheSettings(t, filepath.Join(managedDir, "managed-settings.d", "README.txt"), `{}`) + + paths, err := managedPromptCacheSettingsPaths(managedDir) + require.NoError(t, err) + assert.Equal(t, []string{ + filepath.Join(managedDir, "managed-settings.json"), + filepath.Join(managedDir, "managed-settings.d", "10-earlier.json"), + filepath.Join(managedDir, "managed-settings.d", "20-later.json"), + }, paths) +} + +func writePromptCacheSettings(t *testing.T, path, settings string) { + t.Helper() + require.NoError(t, os.MkdirAll(filepath.Dir(path), 0o700)) + require.NoError(t, os.WriteFile(path, []byte(settings), 0o600)) +} + // newLLMManager builds a ClientManager with a single direct-mode LLM entry // whose settings dir is homeDir/. func newLLMManager(t *testing.T, clientType ClientApp, mode, dir string, ptrs, vals []string) (*ClientManager, string) { diff --git a/pkg/llm/config.go b/pkg/llm/config.go index fd3631c2d6..4a34a49b5b 100644 --- a/pkg/llm/config.go +++ b/pkg/llm/config.go @@ -25,11 +25,14 @@ type OIDCConfig = pkgoidc.ClientConfig // Config holds all LLM gateway settings persisted under the llm: key in // ToolHive's config.yaml. type Config struct { - GatewayURL string `yaml:"gateway_url,omitempty" json:"gateway_url,omitempty"` - TLSSkipVerify bool `yaml:"tls_skip_verify,omitempty" json:"tls_skip_verify,omitempty"` - OIDC OIDCConfig `yaml:"oidc,omitempty" json:"oidc,omitempty"` - Proxy ProxyConfig `yaml:"proxy,omitempty" json:"proxy,omitempty"` - Bedrock BedrockConfig `yaml:"bedrock,omitempty" json:"bedrock,omitempty"` + GatewayURL string `yaml:"gateway_url,omitempty" json:"gateway_url,omitempty"` + TLSSkipVerify bool `yaml:"tls_skip_verify,omitempty" json:"tls_skip_verify,omitempty"` + // ShortPromptCache opts Claude Code out of ToolHive's default one-hour + // prompt-cache lifetime and restores Claude Code's five-minute default. + ShortPromptCache bool `yaml:"short_prompt_cache,omitempty" json:"short_prompt_cache,omitempty"` + OIDC OIDCConfig `yaml:"oidc,omitempty" json:"oidc,omitempty"` + Proxy ProxyConfig `yaml:"proxy,omitempty" json:"proxy,omitempty"` + Bedrock BedrockConfig `yaml:"bedrock,omitempty" json:"bedrock,omitempty"` // Models is the persisted, single source of truth for the model IDs applied // during setup. It feeds two consumers: credential-helper clients (Claude // Desktop) write it verbatim as inferenceModels, and — when Bedrock compat is diff --git a/pkg/llm/manage.go b/pkg/llm/manage.go index eef74295cf..71e5b0fa85 100644 --- a/pkg/llm/manage.go +++ b/pkg/llm/manage.go @@ -38,6 +38,9 @@ func (c *Config) SetFields(opts SetOptions) error { if opts.TLSSkipVerify != nil { c.TLSSkipVerify = *opts.TLSSkipVerify } + if opts.ShortPromptCache != nil { + c.ShortPromptCache = *opts.ShortPromptCache + } if opts.BedrockCompat != nil { c.Bedrock.Compat = *opts.BedrockCompat } @@ -59,13 +62,14 @@ func (c *Config) SetFields(opts SetOptions) error { // field unchanged. TLSSkipVerify uses a pointer so that false can be // distinguished from "not provided" (enabling explicit clear via config set). type SetOptions struct { - GatewayURL string - Issuer string - ClientID string - Audience string - ProxyPort int - CallbackPort int - TLSSkipVerify *bool // nil = not provided; &false = explicitly disable + GatewayURL string + Issuer string + ClientID string + Audience string + ProxyPort int + CallbackPort int + TLSSkipVerify *bool // nil = not provided; &false = explicitly disable + ShortPromptCache *bool // nil = not provided; &false = explicitly restore one-hour caching // BedrockCompat and Enable1M use pointers so false can be distinguished from // "not provided" (enabling explicit clear via config set). See BedrockConfig. BedrockCompat *bool @@ -122,6 +126,11 @@ func (c *Config) Show(w io.Writer) error { } writef("Proxy Port: %d\n", c.EffectiveProxyPort()) writef("Scopes: %v\n", c.OIDC.EffectiveScopes()) + promptCacheTTL := "1h" + if c.ShortPromptCache { + promptCacheTTL = "5m" + } + writef("Prompt Cache: %s\n", promptCacheTTL) if c.TLSSkipVerify { writef("TLS Skip Verify: true (WARNING: certificate verification disabled)\n") } diff --git a/pkg/llm/manage_test.go b/pkg/llm/manage_test.go index b253f32b18..4acc2068b7 100644 --- a/pkg/llm/manage_test.go +++ b/pkg/llm/manage_test.go @@ -6,13 +6,16 @@ package llm import ( "bytes" "context" + "encoding/json" "errors" + "fmt" "strings" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.uber.org/mock/gomock" + "gopkg.in/yaml.v3" "github.com/stacklok/toolhive/pkg/secrets" secretsmocks "github.com/stacklok/toolhive/pkg/secrets/mocks" @@ -121,6 +124,23 @@ func TestConfig_SetFields(t *testing.T) { opts: SetOptions{}, want: Config{GatewayURL: "https://gw.example.com", TLSSkipVerify: true}, }, + { + name: "ShortPromptCache pointer true sets field", + opts: SetOptions{ShortPromptCache: boolPtr(true)}, + want: Config{ShortPromptCache: true}, + }, + { + name: "ShortPromptCache pointer false clears field", + base: Config{ShortPromptCache: true}, + opts: SetOptions{ShortPromptCache: boolPtr(false)}, + want: Config{}, + }, + { + name: "nil ShortPromptCache pointer leaves existing value unchanged", + base: Config{ShortPromptCache: true}, + opts: SetOptions{}, + want: Config{ShortPromptCache: true}, + }, } for _, tt := range tests { @@ -156,6 +176,9 @@ func TestConfig_SetFields(t *testing.T) { if cfg.TLSSkipVerify != tt.want.TLSSkipVerify { t.Errorf("TLSSkipVerify = %v, want %v", cfg.TLSSkipVerify, tt.want.TLSSkipVerify) } + if cfg.ShortPromptCache != tt.want.ShortPromptCache { + t.Errorf("ShortPromptCache = %v, want %v", cfg.ShortPromptCache, tt.want.ShortPromptCache) + } }) } } @@ -365,6 +388,23 @@ func TestConfig_Show(t *testing.T) { }, absent: []string{"TLS Skip Verify"}, }, + { + name: "short prompt cache is shown as five minutes", + cfg: Config{ + GatewayURL: "https://gw.example.com", + OIDC: OIDCConfig{Issuer: "https://auth.example.com", ClientID: "client1"}, + ShortPromptCache: true, + }, + contains: []string{"Prompt Cache: 5m"}, + }, + { + name: "default prompt cache is shown as one hour", + cfg: Config{ + GatewayURL: "https://gw.example.com", + OIDC: OIDCConfig{Issuer: "https://auth.example.com", ClientID: "client1"}, + }, + contains: []string{"Prompt Cache: 1h"}, + }, } for _, tt := range tests { @@ -388,3 +428,35 @@ func TestConfig_Show(t *testing.T) { }) } } + +func TestConfig_JSONOmitsDefaultShortPromptCache(t *testing.T) { + t.Parallel() + + for _, enabled := range []bool{false, true} { + t.Run(fmt.Sprintf("enabled=%t", enabled), func(t *testing.T) { + t.Parallel() + data, err := json.Marshal(Config{ShortPromptCache: enabled}) + require.NoError(t, err) + + var got map[string]any + require.NoError(t, json.Unmarshal(data, &got)) + if enabled { + assert.Equal(t, true, got["short_prompt_cache"]) + } else { + assert.NotContains(t, got, "short_prompt_cache") + } + }) + } +} + +func TestConfig_ShortPromptCacheYAMLRoundTrip(t *testing.T) { + t.Parallel() + + data, err := yaml.Marshal(Config{ShortPromptCache: true}) + require.NoError(t, err) + assert.Contains(t, string(data), "short_prompt_cache: true") + + var got Config + require.NoError(t, yaml.Unmarshal(data, &got)) + assert.True(t, got.ShortPromptCache) +} diff --git a/pkg/llm/setup.go b/pkg/llm/setup.go index 94cc68b861..5e2b54b289 100644 --- a/pkg/llm/setup.go +++ b/pkg/llm/setup.go @@ -35,6 +35,9 @@ type GatewayManager interface { ConfigureLLMGateway(clientType string, cfg llmgateway.ApplyConfig) (string, error) // LLMGatewayModeFor returns "direct", "proxy", or "" for the given client. LLMGatewayModeFor(clientType string) string + // PromptCacheConflict names a setting that may override ToolHive's + // one-hour Claude Code prompt-cache configuration. + PromptCacheConflict(clientType string) (string, error) // IsManaged reports whether a managed-preferences profile overrides the // client's local config (so the config setup writes would be ignored). IsManaged(clientType string) bool @@ -150,7 +153,8 @@ func Setup( configured, err := configureDetectedToolsWithDiscovery( out, errOut, gm, detected, llmCfg.GatewayURL, proxyBaseURL, - tokenHelperPath, tokenHelperArgs, llmCfg.TLSSkipVerify, anthropicPrefix, llmCfg.Models, discoveredModels, llmCfg.Bedrock, + tokenHelperPath, tokenHelperArgs, llmCfg.TLSSkipVerify, anthropicPrefix, + llmCfg.Models, discoveredModels, llmCfg.ShortPromptCache, llmCfg.Bedrock, ) if err != nil { return err @@ -163,6 +167,8 @@ func Setup( // later setup that omits claude-code. llmCfg.Bedrock.Compat is the effective // (persisted + inline) compat state used for the --enable-1m check. warnBedrockNoEffect(errOut, inlineOpts, llmCfg.Bedrock.Compat, configured) + warnShortPromptCacheNoEffect(errOut, inlineOpts, configured) + warnPromptCacheBedrockCompatibility(errOut, inlineOpts, llmCfg.ShortPromptCache, configured) warnTLSSkipVerify(errOut, llmCfg.TLSSkipVerify, configured) warnCredentialHelperTools(out, errOut, gm, configured) @@ -509,6 +515,13 @@ func setupClients( const ( vsCodeClient = "vscode" vsCodeInsiderClient = "vscode-insider" + // claudeCodeClient is the canonical client identifier for Claude Code. + // Declared here as a string literal because pkg/llm does not import + // pkg/client (which owns the ClientApp constant) to avoid an import cycle. + claudeCodeClient = "claude-code" + // promptCacheVerificationCommand reports which TTL Claude Code used under + // usage.cache_creation in the command's JSON output. + promptCacheVerificationCommand = `claude -p "hello" --output-format json` ) func isVSCodeClient(clientType string) bool { @@ -631,11 +644,6 @@ func discoverGatewayModels(ctx context.Context, cfg Config) ([]string, error) { return models, nil } -// claudeCodeClient is the canonical client identifier for Claude Code. Declared -// here as a string literal because pkg/llm does not import pkg/client (which -// owns the ClientApp constant) to avoid an import cycle. -const claudeCodeClient = "claude-code" - // Default Bedrock inference-profile model IDs written for Claude Code in // bedrock-compat mode when --models does not override a tier. These track the // current generation and are expected to be bumped periodically; users override @@ -727,13 +735,34 @@ func warnBedrockNoEffect(errOut io.Writer, opts SetOptions, effectiveCompat bool } } +func warnShortPromptCacheNoEffect(errOut io.Writer, opts SetOptions, configured []ToolConfig) { + if opts.ShortPromptCache != nil && *opts.ShortPromptCache && !isTarget(configured, claudeCodeClient) { + _, _ = fmt.Fprintln(errOut, + "Warning: --short-prompt-cache was set but Claude Code was not configured; the flag had no effect on client settings.") + } +} + +func warnPromptCacheBedrockCompatibility( + errOut io.Writer, opts SetOptions, shortPromptCache bool, configured []ToolConfig, +) { + if shortPromptCache || opts.BedrockCompat == nil || !*opts.BedrockCompat || + !isTarget(configured, claudeCodeClient) { + return + } + _, _ = fmt.Fprintf(errOut, + "Warning: Claude Code's one-hour prompt-cache lifetime may not take effect with Bedrock compatibility: "+ + "the gateway may reject or strip the required beta header, and Bedrock support varies by model. "+ + "Verify the effective lifetime with: %s\n", promptCacheVerificationCommand) +} + // configureDetectedTools patches each detected tool's config file and returns // the list of successfully configured tools. An error is returned only when no // tool was configured successfully. func configureDetectedToolsWithDiscovery( out, errOut io.Writer, gm GatewayManager, detected []string, gatewayURL, proxyBaseURL, tokenHelperPath string, tokenHelperArgs []string, - tlsSkipVerify bool, anthropicPathPrefix string, models, discoveredModels []string, bedrock BedrockConfig, + tlsSkipVerify bool, anthropicPathPrefix string, + models, discoveredModels []string, shortPromptCache bool, bedrock BedrockConfig, ) ([]ToolConfig, error) { var configured []ToolConfig for _, clientType := range detected { @@ -764,6 +793,10 @@ func configureDetectedToolsWithDiscovery( DiscoveredModels: discoveredModels, } + if clientType == claudeCodeClient { + applyCfg.ShortPromptCache = shortPromptCache + } + // Bedrock-compat applies only to Claude Code: it disables the experimental // anthropic-beta headers Bedrock rejects and pins per-tier Bedrock model // IDs. Resolve defaults, tier mapping, and the optional [1m] suffix here so @@ -800,6 +833,9 @@ func configureDetectedToolsWithDiscovery( EnvFilePath: envFilePath, }) _, _ = fmt.Fprintf(out, "Configured %s (%s mode) → %s\n", clientType, mode, configPath) + if clientType == claudeCodeClient && !shortPromptCache { + warnPromptCacheConflict(errOut, gm, clientType) + } } if len(configured) == 0 { return nil, fmt.Errorf("failed to configure any detected tools") @@ -807,6 +843,22 @@ func configureDetectedToolsWithDiscovery( return configured, nil } +func warnPromptCacheConflict(errOut io.Writer, gm GatewayManager, clientType string) { + conflict, err := gm.PromptCacheConflict(clientType) + if err != nil { + _, _ = fmt.Fprintf(errOut, + "Warning: could not inspect %s for prompt-cache TTL conflicts: %v; "+ + "ToolHive wrote the one-hour settings anyway.\n", clientType, err) + return + } + if conflict != "" { + _, _ = fmt.Fprintf(errOut, + "Warning: ToolHive wrote the one-hour prompt-cache settings for %s, but %s may override them. "+ + "Remove that setting to use the one-hour lifetime. Verify the effective lifetime with: %s\n", + clientType, conflict, promptCacheVerificationCommand) + } +} + // resolveAnthropicPrefix returns the effective Anthropic path prefix. When the // caller explicitly set the flag (anthropicPathPrefixSet), the provided value is // returned as-is (including empty string, which disables the prefix). Otherwise diff --git a/pkg/llm/setup_test.go b/pkg/llm/setup_test.go index dce93485d8..59fbedd0d5 100644 --- a/pkg/llm/setup_test.go +++ b/pkg/llm/setup_test.go @@ -11,6 +11,7 @@ import ( "net" "net/http" "net/http/httptest" + "strings" "testing" "time" @@ -189,7 +190,7 @@ func TestConfigureDetectedTools_BedrockClaudeCode(t *testing.T) { []string{"claude-code"}, "https://gw.example.com", "http://localhost:14000/v1", "/usr/local/bin/thv", []string{"llm", "token", "--skip-browser"}, - false, "/anthropic", nil, nil, + false, "/anthropic", nil, nil, false, BedrockConfig{Compat: true, Enable1M: true}, ) require.NoError(t, err) @@ -213,7 +214,7 @@ func TestConfigureDetectedTools_BedrockSkippedForNonClaudeCode(t *testing.T) { []string{"cursor"}, "https://gw.example.com", "http://localhost:14000/v1", "/usr/local/bin/thv", []string{"llm", "token", "--skip-browser"}, - false, "", nil, nil, + false, "", nil, nil, false, BedrockConfig{Compat: true}, ) require.NoError(t, err) @@ -222,6 +223,88 @@ func TestConfigureDetectedTools_BedrockSkippedForNonClaudeCode(t *testing.T) { assert.Empty(t, gm.applied[0].BedrockOpusModel) } +func TestConfigureDetectedTools_PromptCache(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + clientType string + mode string + shortPromptCache bool + conflict string + conflictErr error + wantStderr []string + wantStderrAbsent []string + wantConflictCall bool + }{ + { + name: "Claude Code gets one-hour lifetime by default", + clientType: claudeCodeClient, + mode: llmgateway.ModeDirect, + wantConflictCall: true, + }, + { + name: "non-Claude client has no prompt-cache reporting", + clientType: "cursor", + mode: llmgateway.ModeProxy, + wantStderrAbsent: []string{"Warning:"}, + }, + { + name: "short-cache opt-out skips conflict inspection", + clientType: claudeCodeClient, + mode: llmgateway.ModeDirect, + shortPromptCache: true, + wantStderrAbsent: []string{"Warning:"}, + }, + { + name: "five-minute override warns but one-hour settings are still written", + clientType: claudeCodeClient, + mode: llmgateway.ModeDirect, + conflict: "CLAUDE_CODE_PROMPT_CACHE_TTL=5m in the process environment", + wantStderr: []string{"CLAUDE_CODE_PROMPT_CACHE_TTL=5m", "wrote the one-hour", promptCacheVerificationCommand}, + wantConflictCall: true, + }, + { + name: "inspection failure warns but one-hour settings are still written", + clientType: claudeCodeClient, + mode: llmgateway.ModeDirect, + conflictErr: errors.New("invalid settings"), + wantStderr: []string{"could not inspect claude-code", "wrote the one-hour settings anyway"}, + wantConflictCall: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + gm := &capturingGatewayManager{ + mode: tt.mode, + cacheConflict: tt.conflict, + cacheConflictErr: tt.conflictErr, + } + var out, errOut bytes.Buffer + + configured, err := configureDetectedToolsWithDiscovery( + &out, &errOut, gm, []string{tt.clientType}, + "https://gw.example.com", "http://localhost:14000/v1", + "/usr/local/bin/thv", []string{"llm", "token", "--skip-browser"}, + false, "", nil, nil, tt.shortPromptCache, BedrockConfig{}, + ) + require.NoError(t, err) + require.Len(t, configured, 1) + require.Len(t, gm.applied, 1) + assert.Equal(t, tt.shortPromptCache, gm.applied[0].ShortPromptCache) + assert.Equal(t, tt.wantConflictCall, gm.conflictCalls == 1) + for _, want := range tt.wantStderr { + assert.Contains(t, errOut.String(), want) + } + for _, absent := range tt.wantStderrAbsent { + assert.NotContains(t, errOut.String(), absent) + } + }) + } +} + // ── mergeToolConfigs ────────────────────────────────────────────────────────── func TestMergeToolConfigs_EmptyExisting(t *testing.T) { @@ -305,7 +388,10 @@ func (*stubGatewayManager) DetectedLLMGatewayClients() []string { return nil } func (*stubGatewayManager) ConfigureLLMGateway(_ string, _ llmgateway.ApplyConfig) (string, error) { return "", nil } -func (*stubGatewayManager) LLMGatewayModeFor(_ string) string { return "" } +func (*stubGatewayManager) LLMGatewayModeFor(_ string) string { return "" } +func (*stubGatewayManager) PromptCacheConflict(_ string) (string, error) { + return "", nil +} func (*stubGatewayManager) IsManaged(_ string) bool { return false } func (*stubGatewayManager) LLMClientDetectionHint(_ string) string { return "" } func (*stubGatewayManager) ConfigureEnvFile(_ string, _ llmgateway.ApplyConfig) (string, error) { @@ -747,7 +833,10 @@ func (g *setupGatewayManager) ConfigureLLMGateway(client string, cfg llmgateway. return "/tmp/settings.json", nil } func (g *setupGatewayManager) LLMGatewayModeFor(_ string) string { return g.mode } -func (*setupGatewayManager) IsManaged(_ string) bool { return false } +func (*setupGatewayManager) PromptCacheConflict(_ string) (string, error) { + return "", nil +} +func (*setupGatewayManager) IsManaged(_ string) bool { return false } func (g *setupGatewayManager) LLMClientDetectionHint(_ string) string { return g.hint } @@ -971,6 +1060,102 @@ func TestSetup_VSCodeDiscoveryFailureSelectionSemantics(t *testing.T) { }) } +func TestSetup_ShortPromptCachePreference(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + persistedShort bool + inline *bool + wantShort bool + }{ + {name: "default applies one-hour cache"}, + {name: "persisted opt-out is reapplied", persistedShort: true, wantShort: true}, + {name: "inline opt-out is persisted", inline: boolPtr(true), wantShort: true}, + {name: "explicit false restores one-hour cache", persistedShort: true, inline: boolPtr(false)}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + gm := &capturingGatewayManager{detected: []string{claudeCodeClient}, mode: llmgateway.ModeDirect} + provider := configuredSetupProvider() + provider.cfg.ShortPromptCache = tt.persistedShort + var stdout, stderr bytes.Buffer + err := Setup( + context.Background(), &stdout, &stderr, gm, provider, + func(context.Context, *Config) error { return nil }, + SetOptions{ShortPromptCache: tt.inline}, "", true, "", true, + ) + require.NoError(t, err) + require.Len(t, gm.applied, 1) + assert.Equal(t, tt.wantShort, gm.applied[0].ShortPromptCache) + assert.Equal(t, tt.wantShort, provider.cfg.ShortPromptCache) + }) + } +} + +func TestSetup_ShortPromptCacheWarningForNonClaudeClient(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + inline *bool + wantWarn bool + }{ + {name: "plain setup is silent"}, + {name: "explicit short-cache flag warns", inline: boolPtr(true), wantWarn: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + gm := &capturingGatewayManager{detected: []string{"cursor"}, mode: llmgateway.ModeProxy} + provider := configuredSetupProvider() + var stdout, stderr bytes.Buffer + err := Setup( + context.Background(), &stdout, &stderr, gm, provider, + func(context.Context, *Config) error { return nil }, + SetOptions{ShortPromptCache: tt.inline}, "", true, "cursor", true, + ) + require.NoError(t, err) + assert.Equal(t, tt.wantWarn, strings.Contains(stderr.String(), "--short-prompt-cache was set")) + }) + } +} + +func TestSetup_PromptCacheBedrockWarningOnlyForInlineFlag(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + persistedBedrock bool + inlineBedrock *bool + shortPromptCache bool + wantCompatibility bool + }{ + {name: "inline Bedrock compatibility warns", inlineBedrock: boolPtr(true), wantCompatibility: true}, + {name: "persisted Bedrock compatibility is silent", persistedBedrock: true}, + {name: "short cache is compatible", inlineBedrock: boolPtr(true), shortPromptCache: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + gm := &capturingGatewayManager{detected: []string{claudeCodeClient}, mode: llmgateway.ModeDirect} + provider := configuredSetupProvider() + provider.cfg.Bedrock.Compat = tt.persistedBedrock + provider.cfg.ShortPromptCache = tt.shortPromptCache + var stdout, stderr bytes.Buffer + err := Setup( + context.Background(), &stdout, &stderr, gm, provider, + func(context.Context, *Config) error { return nil }, + SetOptions{BedrockCompat: tt.inlineBedrock}, "", true, "", true, + ) + require.NoError(t, err) + assert.Equal(t, tt.wantCompatibility, + strings.Contains(stderr.String(), "may not take effect with Bedrock compatibility")) + }) + } +} + func TestFilterDetectedClients_LeftoverDirHint(t *testing.T) { t.Parallel() gm := &setupGatewayManager{ @@ -1100,17 +1285,25 @@ func TestSetup_CallbackPortInUseBeforeLogin(t *testing.T) { // capturingGatewayManager records the ApplyConfig passed to ConfigureLLMGateway. type capturingGatewayManager struct { - mode string // returned by LLMGatewayModeFor - applied []llmgateway.ApplyConfig + mode string // returned by LLMGatewayModeFor + detected []string + cacheConflict string + cacheConflictErr error + conflictCalls int + applied []llmgateway.ApplyConfig } -func (*capturingGatewayManager) DetectedLLMGatewayClients() []string { return nil } +func (g *capturingGatewayManager) DetectedLLMGatewayClients() []string { return g.detected } func (g *capturingGatewayManager) ConfigureLLMGateway(_ string, cfg llmgateway.ApplyConfig) (string, error) { g.applied = append(g.applied, cfg) return "/path/to/settings.json", nil } func (g *capturingGatewayManager) LLMGatewayModeFor(_ string) string { return g.mode } -func (*capturingGatewayManager) IsManaged(_ string) bool { return false } +func (g *capturingGatewayManager) PromptCacheConflict(_ string) (string, error) { + g.conflictCalls++ + return g.cacheConflict, g.cacheConflictErr +} +func (*capturingGatewayManager) IsManaged(_ string) bool { return false } func (*capturingGatewayManager) LLMClientDetectionHint(_ string) string { return "" } @@ -1132,7 +1325,7 @@ func TestConfigureDetectedTools_PathPrefixAppendedForDirectMode(t *testing.T) { []string{"claude-code"}, "https://gw.example.com", "http://localhost:14000/v1", "/usr/local/bin/thv", []string{"llm", "token", "--skip-browser"}, - false, "/anthropic", nil, nil, + false, "/anthropic", nil, nil, false, BedrockConfig{}, ) require.NoError(t, err) @@ -1154,7 +1347,7 @@ func TestConfigureDetectedTools_NoPrefixWhenEmpty(t *testing.T) { []string{"claude-code"}, "https://gw.example.com", "http://localhost:14000/v1", "/usr/local/bin/thv", []string{"llm", "token", "--skip-browser"}, - false, "", nil, nil, // no prefix + false, "", nil, nil, false, // no prefix BedrockConfig{}, ) require.NoError(t, err) @@ -1175,7 +1368,7 @@ func TestConfigureDetectedTools_PrefixNotAppliedForProxyMode(t *testing.T) { []string{"cursor"}, "https://gw.example.com", "http://localhost:14000/v1", "/usr/local/bin/thv", []string{"llm", "token", "--skip-browser"}, - false, "/anthropic", nil, nil, + false, "/anthropic", nil, nil, false, BedrockConfig{}, ) require.NoError(t, err) @@ -1267,6 +1460,9 @@ func (*managedGatewayManager) ConfigureLLMGateway(_ string, _ llmgateway.ApplyCo func (*managedGatewayManager) LLMGatewayModeFor(_ string) string { return llmgateway.ModeCredentialHelper } +func (*managedGatewayManager) PromptCacheConflict(_ string) (string, error) { + return "", nil +} func (g *managedGatewayManager) IsManaged(c string) bool { return g.managed[c] } func (*managedGatewayManager) LLMClientDetectionHint(_ string) string { return "" diff --git a/pkg/llmgateway/config.go b/pkg/llmgateway/config.go index c88a0ce406..6d6e8c1197 100644 --- a/pkg/llmgateway/config.go +++ b/pkg/llmgateway/config.go @@ -114,6 +114,9 @@ type ApplyConfig struct { // discovery request. It is used only by integrations that require an explicit // model catalogue, such as VS Code's customendpoint provider. DiscoveredModels []string + // ShortPromptCache opts supported clients out of ToolHive's one-hour + // prompt-cache lifetime and restores their shorter default. + ShortPromptCache bool // BedrockCompat and the per-tier Bedrock model IDs configure Claude Code for a // gateway that forwards to AWS Bedrock. When BedrockCompat is true, Claude Code // is configured with CLAUDE_CODE_DISABLE_EXPERIMENTAL_BETAS=1 (Bedrock rejects diff --git a/test/e2e/cli_llm_setup_test.go b/test/e2e/cli_llm_setup_test.go index 56d6de13c6..fb2d2c23de 100644 --- a/test/e2e/cli_llm_setup_test.go +++ b/test/e2e/cli_llm_setup_test.go @@ -472,4 +472,74 @@ var _ = Describe("thv llm setup / teardown", Label("cli", "llm", "setup", "e2e") "a fresh token should be printed to stdout after deferred login") }) }) + + Describe("thv llm setup prompt cache", func() { + It("defaults to one hour and persists the short-cache opt-out", func() { + claudeDir := filepath.Join(tempDir, ".claude") + Expect(os.MkdirAll(claudeDir, 0750)).To(Succeed()) + Expect(createFakeBinary(binDir, "claude")).To(Succeed()) + + issuerURL := fmt.Sprintf("http://localhost:%d", oidcPort) + setupArgs := []string{ + "llm", "setup", "--lazy", "--client", "claude-code", + "--anthropic-path-prefix", "", + } + + By("Applying the default one-hour cache lifetime") + stdout, stderr, err := thvCmd(append(setupArgs, + "--gateway-url", gatewayURL, + "--issuer", issuerURL, + "--client-id", clientID, + )...).RunWithTimeout(30 * time.Second) + Expect(err).ToNot(HaveOccurred(), + "setup should succeed; stdout=%q stderr=%q", stdout, stderr) + + settingsPath := filepath.Join(claudeDir, "settings.json") + expectOneHourSettings := func(present bool) { + By("Reading Claude Code settings") + data, readErr := os.ReadFile(settingsPath) + Expect(readErr).ToNot(HaveOccurred()) + var settings map[string]any + Expect(json.Unmarshal(data, &settings)).To(Succeed()) + + for pointer, expected := range map[string]string{ + "/promptCacheTtl": "1h", + "/subagentPromptCacheTtl": "1h", + "/env/ENABLE_PROMPT_CACHING_1H": "1", + } { + actual, found := jsonPointerGet(settings, pointer) + Expect(found).To(Equal(present), "unexpected presence for %s", pointer) + if present { + Expect(actual).To(Equal(expected), "unexpected value for %s", pointer) + } + } + } + expectOneHourSettings(true) + + By("Verifying the default does not persist an opt-out") + showOut, _ := thvCmd("llm", "config", "show", "--format", "json").ExpectSuccess() + var cfg llm.Config + Expect(json.Unmarshal([]byte(showOut), &cfg)).To(Succeed()) + Expect(cfg.ShortPromptCache).To(BeFalse()) + + By("Persisting the short-cache opt-out") + thvCmd(append(setupArgs, "--short-prompt-cache")...).ExpectSuccess() + expectOneHourSettings(false) + showOut, _ = thvCmd("llm", "config", "show", "--format", "json").ExpectSuccess() + Expect(json.Unmarshal([]byte(showOut), &cfg)).To(Succeed()) + Expect(cfg.ShortPromptCache).To(BeTrue()) + + By("Reapplying the persisted opt-out with a plain setup") + thvCmd(setupArgs...).ExpectSuccess() + expectOneHourSettings(false) + + By("Explicitly restoring the one-hour default") + thvCmd(append(setupArgs, "--short-prompt-cache=false")...).ExpectSuccess() + expectOneHourSettings(true) + + By("Verifying teardown removes the cache settings") + thvCmd("llm", "teardown", "claude-code").ExpectSuccess() + expectOneHourSettings(false) + }) + }) })