From 778868f1c1f7f9de330474baf77b92389112a4de Mon Sep 17 00:00:00 2001 From: Aditya Thebe Date: Tue, 11 Aug 2026 17:02:51 +0000 Subject: [PATCH 1/2] fix(ai): make model discovery auth-consistent --- README.md | 2 +- pkg/ai/adapter_models.go | 17 +- pkg/ai/adapters.go | 296 ++++++++++++++-- pkg/ai/adapters_cache.go | 53 ++- pkg/ai/adapters_cache_test.go | 225 +++++++++++- pkg/ai/adapters_test.go | 51 +++ pkg/ai/availability_test.go | 8 +- pkg/ai/catalog_disabled_ginkgo_test.go | 8 +- pkg/ai/catalog_resolve.go | 185 +++++++--- pkg/ai/catalog_resolve_test.go | 320 +++++++++++++++++- pkg/ai/live_catalog_test.go | 8 +- pkg/ai/model_cache.go | 206 ++++++++++- pkg/ai/models_remote.go | 187 ++++++---- pkg/ai/models_remote_test.go | 77 +++++ pkg/database/session_chat_store.go | 3 +- .../session_last_activity_integration_test.go | 45 --- 16 files changed, 1465 insertions(+), 226 deletions(-) diff --git a/README.md b/README.md index e8fb8b91..5e9de221 100644 --- a/README.md +++ b/README.md @@ -496,7 +496,7 @@ captain whoami --no-cache Lists every AI adapter (API providers and CLI agents: `anthropic`, `openai`, `gemini`, `claude-cli`, `claude-agent`, `codex-cli`, `gemini-cli`), their authentication method, binary availability, and a live model listing. -Provider model listings are resolved through the persisted cache in `~/.config/captain/models.json` (24h TTL, invalidated when the set of configured API keys changes), and priced from the OpenRouter snapshot in `~/Library/Caches/flanksource/openrouter-pricing.json` (24h TTL). `--no-cache` skips both, re-queries every provider's model endpoint plus OpenRouter pricing, and rewrites each cache with the fresh result. Add `-v` to see an access line per request, `-vv` for headers and query params, `-vvv` for request bodies, `-vvvv` for response bodies; credentials are redacted at every rung. Failed requests (status >= 400 or a transport error) are logged at the default verbosity. `-Plog.level.http=` raises HTTP logging alone, and `-Phttp.log.base-level=` (or `HTTP_LOG_BASE_LEVEL`) shifts the whole ladder. +Provider model listings are resolved through credential- and endpoint-scoped entries under `~/.config/captain/models/` (24h TTL), and priced from the OpenRouter snapshot in `~/Library/Caches/flanksource/openrouter-pricing.json` (24h TTL). The model cache uses a machine-local keyed namespace, so different provider accounts cannot reuse one another's availability. `--no-cache` skips both caches, re-queries every provider's model endpoint plus OpenRouter pricing, and rewrites only the affected model-cache entries with the fresh result. Add `-v` to see an access line per request, `-vv` for headers and query params, `-vvv` for request bodies, `-vvvv` for response bodies; credentials are redacted at every rung. Failed requests (status >= 400 or a transport error) are logged at the default verbosity. `-Plog.level.http=` raises HTTP logging alone, and `-Phttp.log.base-level=` (or `HTTP_LOG_BASE_LEVEL`) shifts the whole ladder. To keep the traffic instead of watching it scroll past, `-Phttp.har=` writes every request/response pair — including redirect hops and retries — to a HAR 1.2 archive you can open in browser DevTools. It applies to any command, not just `whoami`, and the file is written even when the command fails. `-Phttp.har.level=metadata` records headers, query strings and timings without bodies; `-Phttp.har.maxBodySize=` changes the 64 KB per-body cap (`0` for none). Credentials are masked the same way as in the wire log, which means the archive is safe to attach to a bug report but cannot be replayed — `-Phttp.har.sensitive=true` keeps them verbatim and writes the file `0600`. Use `-Phttp.captain.har=` to capture captain's own traffic when a shared `http.har` is already set. diff --git a/pkg/ai/adapter_models.go b/pkg/ai/adapter_models.go index 6e662c21..c94ee1ce 100644 --- a/pkg/ai/adapter_models.go +++ b/pkg/ai/adapter_models.go @@ -25,7 +25,7 @@ var resolveModelRows = ResolveModels // The resolver is Captain's cached model path, so repeated probes reuse a fresh // cache instead of hitting providers every time; refresh bypasses that cache and // re-queries every provider listing. -func fetchAPIModels(backends []Backend, probe AuthProbe, refresh bool) map[Backend]modelFetch { +func fetchAPIModels(backends []Backend, credentials CredentialSnapshot, apiURLs map[Backend]string, refresh bool) map[Backend]modelFetch { apis := map[Backend]bool{} for _, b := range backends { if b.Kind() != "api" { @@ -35,7 +35,7 @@ func fetchAPIModels(backends []Backend, probe AuthProbe, refresh bool) map[Backe if source == "" { continue } - if effectiveAPIKey(source, probe) != "" { + if strings.TrimSpace(credentials.APIKey(source).Token) != "" { apis[source] = true } } @@ -47,7 +47,13 @@ func fetchAPIModels(backends []Backend, probe AuthProbe, refresh bool) map[Backe wg.Add(1) go func(backend Backend) { defer wg.Done() - rows, err := resolveModelRows(context.Background(), ResolveOptions{Backend: backend, UseTokens: true, Refresh: refresh}) + rows, err := resolveModelRows(context.Background(), ResolveOptions{ + Backend: backend, + UseTokens: true, + Refresh: refresh, + Credentials: credentials, + APIURL: apiURLs[backend], + }) m := liveModelDefs(rows, backend) mu.Lock() out[backend] = modelFetch{models: m, err: err} @@ -140,7 +146,7 @@ func applyModels(st *AdapterStatus, b Backend, cache map[Backend]modelFetch, cod } envVars := AuthEnvVars(source) - if effectiveAPIKey(source, probe) == "" { + if strings.TrimSpace(effectiveAPIKey(source, probe)) == "" { st.ModelError = "configure a Captain vault token or set " + strings.Join(envVars, " or ") + " to list models" return } @@ -157,6 +163,9 @@ func applyModels(st *AdapterStatus, b Backend, cache map[Backend]modelFetch, cod } func effectiveAPIKey(backend Backend, probe AuthProbe) string { + if probe.credentials.supplied { + return probe.credentials.APIKey(backend).Token + } if probe.APICredentials != nil { return probe.APICredentials[backend].Token } diff --git a/pkg/ai/adapters.go b/pkg/ai/adapters.go index d2c66384..0c0c222f 100644 --- a/pkg/ai/adapters.go +++ b/pkg/ai/adapters.go @@ -2,6 +2,8 @@ package ai import ( "context" + "crypto/sha256" + "encoding/json" "fmt" "os" "os/exec" @@ -75,41 +77,97 @@ func (a AdapterStatus) Ready() bool { // CachedAdapters so the cache keeps raw probe data and a toggle takes effect on // the next read instead of after the TTL expires. func ApplyDisabled(adapters []AdapterStatus) []AdapterStatus { + out := cloneAdapterStatuses(adapters) disabled := Disabled() if disabled.Empty() { - return adapters + return out } - out := make([]AdapterStatus, len(adapters)) - for i, a := range adapters { + for i, a := range out { backend := Backend(a.Backend) a.Disabled = disabled.Backend(backend) a.DisabledReason = disabled.Reason(backend) - details := make([]ModelDef, len(a.ModelDetails)) for j, md := range a.ModelDetails { md.Disabled = a.Disabled || disabled.Model(backend, md.ID) md.SupportedEfforts = disabled.Efforts(md.SupportedEfforts) if disabled.Effort(md.DefaultEffort) { md.DefaultEffort = api.EffortNone } - details[j] = md + a.ModelDetails[j] = md } - a.ModelDetails = details out[i] = a } return out } +func cloneAdapterStatuses(adapters []AdapterStatus) []AdapterStatus { + out := make([]AdapterStatus, len(adapters)) + for i, adapter := range adapters { + adapter.Models = append([]string(nil), adapter.Models...) + adapter.ModelDetails = cloneModelDefs(adapter.ModelDetails) + out[i] = adapter + } + return out +} + +func cloneModelDefs(models []ModelDef) []ModelDef { + out := make([]ModelDef, len(models)) + for i, model := range models { + model.InputMediaTypes = append([]string(nil), model.InputMediaTypes...) + model.SupportedEfforts = append([]api.Effort(nil), model.SupportedEfforts...) + out[i] = model + } + return out +} + +// CredentialSnapshot is an immutable set of already-resolved API credentials. +// NewCredentialSnapshot clones its input, including an empty map, so callers can +// explicitly suppress process-global credential lookup. The zero value means no +// snapshot was supplied and lets ResolveModels resolve the relevant credentials +// once at operation start for backwards compatibility. +type CredentialSnapshot struct { + apiKeys map[Backend]api.ResolvedAPIKey + supplied bool +} + +func NewCredentialSnapshot(apiKeys map[Backend]api.ResolvedAPIKey) CredentialSnapshot { + cloned := make(map[Backend]api.ResolvedAPIKey, len(apiKeys)) + for backend, resolved := range apiKeys { + cloned[backend] = resolved + } + return CredentialSnapshot{apiKeys: cloned, supplied: true} +} + +// APIKey returns the resolved credential for backend, or an empty value when +// that backend was not present in the snapshot. +func (s CredentialSnapshot) APIKey(backend Backend) api.ResolvedAPIKey { + return s.apiKeys[backend] +} + +func (s CredentialSnapshot) clone() CredentialSnapshot { + if !s.supplied { + return CredentialSnapshot{} + } + return NewCredentialSnapshot(s.apiKeys) +} + // AuthProbe abstracts the host environment (env vars, PATH, credential files) // so resolveAdapter stays pure and testable. Fields are exported so callers in // other packages (and their tests) can construct a hermetic probe. type AuthProbe struct { - Getenv func(string) string - LookPath func(string) (string, error) - FileExists func(string) bool - CodexModels func(context.Context, string) ([]ModelDef, error) - APICredentials map[Backend]api.ResolvedAPIKey - ProbeError error - Home string + Getenv func(string) string + LookPath func(string) (string, error) + FileExists func(string) bool + FileIdentity func(string) string + ExecutableIdentity func(string) string + CodexModels func(context.Context, string) ([]ModelDef, error) + APICredentials map[Backend]api.ResolvedAPIKey + APIURLs map[Backend]string + RuntimeStatuses map[Backend]RuntimeStatus + ProbeError error + Home string + + credentials CredentialSnapshot + stateFingerprint string } // OSAuthProbe wires AuthProbe to the real host environment. @@ -123,7 +181,9 @@ func OSAuthProbe() AuthProbe { _, err := os.Stat(p) return err == nil }, - Home: home, + FileIdentity: hostFileIdentity, + ExecutableIdentity: hostExecutableIdentity, + Home: home, } probe.APICredentials = make(map[Backend]api.ResolvedAPIKey, len(apiBackends)) for _, backend := range apiBackends { @@ -134,9 +194,201 @@ func OSAuthProbe() AuthProbe { } probe.APICredentials[backend] = resolved } + probe.APIURLs = make(map[Backend]string, len(apiBackends)) + for _, backend := range apiBackends { + if apiURL := firstEnv(modelAPIURLEnvVars(backend), os.Getenv); apiURL != "" { + probe.APIURLs[backend] = apiURL + } + } + return probe +} + +type frozenPathState struct { + Path string `json:"path,omitempty"` + Identity string `json:"identity,omitempty"` +} + +type frozenFileState struct { + Exists bool `json:"exists"` + Identity string `json:"identity,omitempty"` +} + +type frozenCredentialState struct { + Token string `json:"token,omitempty"` + Source string `json:"source,omitempty"` + Detail string `json:"detail,omitempty"` +} + +type frozenProbeState struct { + Home string `json:"home"` + Credentials map[Backend]frozenCredentialState `json:"credentials"` + Environment map[string]string `json:"environment"` + APIURLs map[Backend]string `json:"apiURLs"` + Paths map[string]frozenPathState `json:"paths"` + Files map[string]frozenFileState `json:"files"` + Runtimes map[Backend]RuntimeStatus `json:"runtimes"` +} + +// freezeAuthProbe eagerly captures every identity-bearing host observation used +// by adapter probing. Downstream auth reporting, cache identity, binary checks, +// and model discovery then consume values from one point-in-time snapshot rather +// than invoking live OS callbacks independently. +func freezeAuthProbe(probe AuthProbe) AuthProbe { + if probe.stateFingerprint != "" { + return probe + } + getenv := probe.Getenv + if getenv == nil { + getenv = func(string) string { return "" } + } + environment := map[string]string{} + for _, backend := range AllBackends() { + if backend.Kind() == "cli" { + for _, name := range AuthEnvVars(backend) { + environment[name] = getenv(name) + } + } + for _, name := range modelAPIURLEnvVars(backend) { + environment[name] = getenv(name) + } + } + probe.Getenv = func(name string) string { return environment[name] } + + if probe.credentials.supplied { + probe.credentials = probe.credentials.clone() + } else if probe.APICredentials != nil { + probe.credentials = NewCredentialSnapshot(probe.APICredentials) + } else { + resolved := make(map[Backend]api.ResolvedAPIKey, len(apiBackends)) + for _, backend := range apiBackends { + for _, name := range AuthEnvVars(backend) { + if token := getenv(name); strings.TrimSpace(token) != "" { + resolved[backend] = api.ResolvedAPIKey{Token: token, Source: credentials.SourceEnvironment, Detail: name} + break + } + } + } + probe.credentials = NewCredentialSnapshot(resolved) + } + probe.APICredentials = nil + + apiURLsSupplied := probe.APIURLs != nil + apiURLs := make(map[Backend]string, len(probe.APIURLs)) + for backend, apiURL := range probe.APIURLs { + apiURLs[backend] = apiURL + } + if !apiURLsSupplied { + for _, backend := range apiBackends { + if apiURL := firstEnv(modelAPIURLEnvVars(backend), getenv); apiURL != "" { + apiURLs[backend] = apiURL + } + } + } + probe.APIURLs = apiURLs + + lookPath := probe.LookPath + if lookPath == nil { + lookPath = func(string) (string, error) { return "", os.ErrNotExist } + } + paths := map[string]frozenPathState{} + executableIdentities := map[string]string{} + for _, adapter := range cliAdapters() { + if _, exists := paths[adapter.binary]; exists { + continue + } + path, err := lookPath(adapter.binary) + if err != nil || strings.TrimSpace(path) == "" { + paths[adapter.binary] = frozenPathState{} + continue + } + identity := "" + if probe.ExecutableIdentity != nil { + identity = probe.ExecutableIdentity(path) + } + executableIdentities[path] = identity + paths[adapter.binary] = frozenPathState{Path: path, Identity: identity} + } + probe.LookPath = func(binary string) (string, error) { + if state, ok := paths[binary]; ok && state.Path != "" { + return state.Path, nil + } + return "", os.ErrNotExist + } + probe.ExecutableIdentity = func(path string) string { return executableIdentities[path] } + + fileExists := probe.FileExists + if fileExists == nil { + fileExists = func(string) bool { return false } + } + files := map[string]frozenFileState{} + for _, adapter := range cliAdapters() { + for _, login := range adapter.logins { + path := filepath.Join(probe.Home, login.rel) + if _, captured := files[path]; captured { + continue + } + exists := fileExists(path) + identity := "" + if exists && probe.FileIdentity != nil { + identity = probe.FileIdentity(path) + } + files[path] = frozenFileState{Exists: exists, Identity: identity} + } + } + probe.FileExists = func(path string) bool { return files[path].Exists } + probe.FileIdentity = func(path string) string { return files[path].Identity } + + runtimesSupplied := probe.RuntimeStatuses != nil + runtimes := make(map[Backend]RuntimeStatus, len(probe.RuntimeStatuses)) + for backend, status := range probe.RuntimeStatuses { + runtimes[backend] = status + } + if !runtimesSupplied { + for backend := range cliAdapters() { + if status, custom := probeRuntime(backend); custom { + runtimes[backend] = status + } + } + } + probe.RuntimeStatuses = runtimes + + credentialState := make(map[Backend]frozenCredentialState, len(apiBackends)) + for _, backend := range apiBackends { + resolved := probe.credentials.APIKey(backend) + credentialState[backend] = frozenCredentialState{Token: resolved.Token, Source: resolved.Source, Detail: resolved.Detail} + } + state := frozenProbeState{ + Home: probe.Home, + Credentials: credentialState, + Environment: environment, + APIURLs: apiURLs, + Paths: paths, + Files: files, + Runtimes: runtimes, + } + encoded, _ := json.Marshal(state) + fingerprint := sha256.Sum256(encoded) + probe.stateFingerprint = fmt.Sprintf("%x", fingerprint) return probe } +func hostFileIdentity(path string) string { + data, err := os.ReadFile(path) + if err != nil { + return "unreadable" + } + digest := sha256.Sum256(data) + return fmt.Sprintf("sha256:%x", digest) +} + +func hostExecutableIdentity(path string) string { + info, err := os.Stat(path) + if err != nil { + return "unreadable" + } + return fmt.Sprintf("size=%d|mode=%d|mtime=%d", info.Size(), info.Mode(), info.ModTime().UnixNano()) +} + // loginFile is a credential file whose presence indicates a CLI has been logged // in out-of-band (subscription/OAuth) rather than via an API-key env var. type loginFile struct { @@ -190,13 +442,18 @@ func cliAdapters() map[Backend]cliAdapter { // wins over a CLI login file because that is the path NewProvider/ListModels // actually take. func resolveAdapter(backend Backend, p AuthProbe) AdapterStatus { + p = freezeAuthProbe(p) + return resolveAdapterFrozen(backend, p) +} + +func resolveAdapterFrozen(backend Backend, p AuthProbe) AdapterStatus { st := AdapterStatus{ Backend: string(backend), Type: backend.Kind(), Provider: string(backend.Provider()), Mode: string(backend.Mode()), } - if backend.Kind() == "api" && p.APICredentials != nil { - if resolved := p.APICredentials[backend]; resolved.Token != "" { + if backend.Kind() == "api" { + if resolved := p.credentials.APIKey(backend); strings.TrimSpace(resolved.Token) != "" { st.Authenticated = true st.AuthDetail = MaskKey(resolved.Token) if resolved.Source == credentials.SourceVault { @@ -217,7 +474,7 @@ func resolveAdapter(backend Backend, p AuthProbe) AdapterStatus { } if cli, ok := cliAdapters()[backend]; ok { - if runtime, custom := probeRuntime(backend); custom { + if runtime, custom := p.RuntimeStatuses[backend]; custom { st.Binary = runtime.Binary st.BinaryMissing = runtime.BinaryMissing st.DependencyMissing = runtime.DependencyMissing @@ -272,6 +529,7 @@ func ProbeAdapters(opts WhoamiOptions, probe AuthProbe) ([]AdapterStatus, error) if probe.ProbeError != nil { return nil, probe.ProbeError } + probe = freezeAuthProbe(probe) backends := AllBackends() if opts.Backend != "" { b := Backend(opts.Backend) @@ -284,13 +542,13 @@ func ProbeAdapters(opts WhoamiOptions, probe AuthProbe) ([]AdapterStatus, error) var models map[Backend]modelFetch var codexModels modelFetch if opts.Models { - models = fetchAPIModels(backends, probe, opts.NoCache) + models = fetchAPIModels(backends, probe.credentials, probe.APIURLs, opts.NoCache) codexModels = fetchCodexModels(backends, probe) } adapters := make([]AdapterStatus, 0, len(backends)) for _, b := range backends { - st := resolveAdapter(b, probe) + st := resolveAdapterFrozen(b, probe) if opts.Models { applyModels(&st, b, models, codexModels, probe) } diff --git a/pkg/ai/adapters_cache.go b/pkg/ai/adapters_cache.go index d1a08150..d6e83030 100644 --- a/pkg/ai/adapters_cache.go +++ b/pkg/ai/adapters_cache.go @@ -1,6 +1,7 @@ package ai import ( + "errors" "sync" "time" ) @@ -10,16 +11,19 @@ import ( // key/login/model changes surface without a probe per request. const adapterCacheTTL = 60 * time.Second -// adapterProbe is the live probe sourcing the cache. It is a package var so -// tests can substitute a deterministic, network-free stub. -var adapterProbe = func() ([]AdapterStatus, error) { - return ProbeAdapters(WhoamiOptions{Models: true}, OSAuthProbe()) +// adapterAuthProbe captures the current host identity before a cache lookup; +// adapterProbe resolves adapters from that same frozen snapshot. Both are +// package vars so tests can substitute deterministic, network-free stubs. +var adapterAuthProbe = OSAuthProbe +var adapterProbe = func(probe AuthProbe) ([]AdapterStatus, error) { + return ProbeAdapters(WhoamiOptions{Models: true}, probe) } var ( - adapterCacheMu sync.Mutex - adapterCache []AdapterStatus - adapterCacheAt time.Time + adapterCacheMu sync.Mutex + adapterCache []AdapterStatus + adapterCacheAt time.Time + adapterCacheFingerprint string ) // CachedAdapters returns the probed adapters, reusing a cached probe within the @@ -32,14 +36,35 @@ var ( func CachedAdapters(now time.Time) ([]AdapterStatus, error) { adapterCacheMu.Lock() defer adapterCacheMu.Unlock() - if adapterCache != nil && now.Sub(adapterCacheAt) < adapterCacheTTL { + probe := freezeAuthProbe(adapterAuthProbe()) + if probe.ProbeError != nil { + return nil, probe.ProbeError + } + if adapterCache != nil && adapterCacheFingerprint == probe.stateFingerprint && now.Sub(adapterCacheAt) < adapterCacheTTL { return ApplyDisabled(adapterCache), nil } - adapters, err := adapterProbe() - if err != nil { - return nil, err + for attempt := 0; attempt < 2; attempt++ { + adapters, err := adapterProbe(probe) + if err != nil { + return nil, err + } + // Codex model discovery runs an external process whose account cannot be + // injected. Re-capture cheap host state before publishing its result; one + // retry closes a login/binary change that happened during the probe. + current := freezeAuthProbe(adapterAuthProbe()) + if current.ProbeError == nil && current.stateFingerprint == probe.stateFingerprint { + adapterCache = cloneAdapterStatuses(adapters) + adapterCacheAt = now + adapterCacheFingerprint = probe.stateFingerprint + return ApplyDisabled(adapterCache), nil + } + if current.ProbeError != nil { + return nil, current.ProbeError + } + if attempt == 1 { + break + } + probe = current } - adapterCache = adapters - adapterCacheAt = now - return ApplyDisabled(adapters), nil + return nil, errors.New("adapter probe did not settle") } diff --git a/pkg/ai/adapters_cache_test.go b/pkg/ai/adapters_cache_test.go index 3acb9b5a..8d5c9764 100644 --- a/pkg/ai/adapters_cache_test.go +++ b/pkg/ai/adapters_cache_test.go @@ -2,25 +2,33 @@ package ai import ( "errors" + "path/filepath" "testing" "time" + + "github.com/flanksource/captain/pkg/api" + "github.com/flanksource/captain/pkg/credentials" ) func TestCachedAdaptersReusesWithinTTL(t *testing.T) { prevProbe := adapterProbe - prevCache, prevAt := adapterCache, adapterCacheAt + prevAuthProbe := adapterAuthProbe + prevCache, prevAt, prevFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint t.Cleanup(func() { adapterProbe = prevProbe - adapterCache, adapterCacheAt = prevCache, prevAt + adapterAuthProbe = prevAuthProbe + adapterCache, adapterCacheAt, adapterCacheFingerprint = prevCache, prevAt, prevFingerprint }) + probe := fakeProbe(nil, nil, nil, t.TempDir()) + adapterAuthProbe = func() AuthProbe { return probe } stub := []AdapterStatus{{Backend: string(BackendAnthropic), Type: "api"}} calls := 0 - adapterProbe = func() ([]AdapterStatus, error) { + adapterProbe = func(AuthProbe) ([]AdapterStatus, error) { calls++ return stub, nil } - adapterCache, adapterCacheAt = nil, time.Time{} + adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" base := time.Unix(1_000_000, 0) if _, err := CachedAdapters(base); err != nil { @@ -42,15 +50,19 @@ func TestCachedAdaptersReusesWithinTTL(t *testing.T) { func TestCachedAdaptersDoesNotCacheErrors(t *testing.T) { prevProbe := adapterProbe - prevCache, prevAt := adapterCache, adapterCacheAt + prevAuthProbe := adapterAuthProbe + prevCache, prevAt, prevFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint t.Cleanup(func() { adapterProbe = prevProbe - adapterCache, adapterCacheAt = prevCache, prevAt + adapterAuthProbe = prevAuthProbe + adapterCache, adapterCacheAt, adapterCacheFingerprint = prevCache, prevAt, prevFingerprint }) - adapterCache, adapterCacheAt = nil, time.Time{} + probe := fakeProbe(nil, nil, nil, t.TempDir()) + adapterAuthProbe = func() AuthProbe { return probe } + adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" calls := 0 - adapterProbe = func() ([]AdapterStatus, error) { + adapterProbe = func(AuthProbe) ([]AdapterStatus, error) { calls++ if calls == 1 { return nil, errors.New("transient probe failure") @@ -75,3 +87,200 @@ func TestCachedAdaptersDoesNotCacheErrors(t *testing.T) { t.Errorf("probe called %d times, want 2 (error not cached)", calls) } } + +func TestCachedAdaptersInvalidatesCredentialSourceTokenAndEndpointImmediately(t *testing.T) { + prevProbe, prevAuthProbe := adapterProbe, adapterAuthProbe + prevCache, prevAt, prevFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint + t.Cleanup(func() { + adapterProbe, adapterAuthProbe = prevProbe, prevAuthProbe + adapterCache, adapterCacheAt, adapterCacheFingerprint = prevCache, prevAt, prevFingerprint + }) + adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" + + token := "same-effective-token" + source := credentials.SourceVault + detail := "vault" + apiURL := "https://one.example/v1" + home := t.TempDir() + adapterAuthProbe = func() AuthProbe { + probe := fakeProbe(nil, nil, nil, home) + probe.APICredentials = map[Backend]api.ResolvedAPIKey{ + BackendOpenAI: {Token: token, Source: source, Detail: detail}, + } + probe.APIURLs = map[Backend]string{BackendOpenAI: apiURL} + return probe + } + calls := 0 + adapterProbe = func(probe AuthProbe) ([]AdapterStatus, error) { + calls++ + return []AdapterStatus{resolveAdapterFrozen(BackendOpenAI, probe)}, nil + } + + base := time.Unix(3_000_000, 0) + got, err := CachedAdapters(base) + if err != nil || len(got) != 1 || got[0].AuthMethod != "Captain vault" { + t.Fatalf("initial adapters = %+v err=%v", got, err) + } + if _, err := CachedAdapters(base.Add(time.Second)); err != nil || calls != 1 { + t.Fatalf("unchanged cache: calls=%d err=%v", calls, err) + } + + // The effective token is unchanged, but auth reporting must immediately move + // from vault to environment. + source, detail = credentials.SourceEnvironment, "OPENAI_API_KEY" + got, err = CachedAdapters(base.Add(2 * time.Second)) + if err != nil || got[0].AuthMethod != "OPENAI_API_KEY (env)" || calls != 2 { + t.Fatalf("source change adapters = %+v calls=%d err=%v", got, calls, err) + } + apiURL = "https://two.example/v1" + if _, err := CachedAdapters(base.Add(3 * time.Second)); err != nil || calls != 3 { + t.Fatalf("endpoint change: calls=%d err=%v", calls, err) + } + token = "rotated-effective-token" + if _, err := CachedAdapters(base.Add(4 * time.Second)); err != nil || calls != 4 { + t.Fatalf("token rotation: calls=%d err=%v", calls, err) + } + token = "" + got, err = CachedAdapters(base.Add(5 * time.Second)) + if err != nil || got[0].Authenticated || calls != 5 { + t.Fatalf("token removal adapters = %+v calls=%d err=%v", got, calls, err) + } +} + +func TestCachedAdaptersInvalidatesLocalLoginIdentityImmediately(t *testing.T) { + prevProbe, prevAuthProbe := adapterProbe, adapterAuthProbe + prevCache, prevAt, prevFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint + t.Cleanup(func() { + adapterProbe, adapterAuthProbe = prevProbe, prevAuthProbe + adapterCache, adapterCacheAt, adapterCacheFingerprint = prevCache, prevAt, prevFingerprint + }) + adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" + + home := t.TempDir() + authFile := filepath.Join(home, ".codex", "auth.json") + accountIdentity := "account-a" + adapterAuthProbe = func() AuthProbe { + probe := fakeProbe(nil, map[string]string{"codex": "/bin/codex"}, map[string]bool{authFile: true}, home) + probe.FileIdentity = func(path string) string { + if path == authFile { + return accountIdentity + } + return "" + } + return probe + } + calls := 0 + adapterProbe = func(probe AuthProbe) ([]AdapterStatus, error) { + calls++ + status := resolveAdapterFrozen(BackendCodexCLI, probe) + status.Models = []string{probe.FileIdentity(authFile)} + return []AdapterStatus{status}, nil + } + + base := time.Unix(4_000_000, 0) + got, err := CachedAdapters(base) + if err != nil || got[0].Models[0] != "account-a" || calls != 1 { + t.Fatalf("initial local account = %+v calls=%d err=%v", got, calls, err) + } + accountIdentity = "account-b" + got, err = CachedAdapters(base.Add(time.Second)) + if err != nil || got[0].Models[0] != "account-b" || calls != 2 { + t.Fatalf("changed local account = %+v calls=%d err=%v", got, calls, err) + } +} + +func TestCachedAdaptersReturnsDeepCopies(t *testing.T) { + prevProbe, prevAuthProbe := adapterProbe, adapterAuthProbe + prevCache, prevAt, prevFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint + t.Cleanup(func() { + adapterProbe, adapterAuthProbe = prevProbe, prevAuthProbe + adapterCache, adapterCacheAt, adapterCacheFingerprint = prevCache, prevAt, prevFingerprint + }) + adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" + probe := fakeProbe(nil, nil, nil, t.TempDir()) + adapterAuthProbe = func() AuthProbe { return probe } + adapterProbe = func(AuthProbe) ([]AdapterStatus, error) { + return []AdapterStatus{{ + Backend: string(BackendOpenAI), Models: []string{"model-a"}, + ModelDetails: []ModelDef{{ + ID: "model-a", InputMediaTypes: []string{"image/png"}, SupportedEfforts: []api.Effort{api.EffortLow}, + }}, + }}, nil + } + + base := time.Unix(5_000_000, 0) + first, err := CachedAdapters(base) + if err != nil { + t.Fatal(err) + } + first[0].Backend = "poisoned" + first[0].Models[0] = "poisoned" + first[0].ModelDetails[0].ID = "poisoned" + first[0].ModelDetails[0].InputMediaTypes[0] = "poisoned" + first[0].ModelDetails[0].SupportedEfforts[0] = api.EffortHigh + + second, err := CachedAdapters(base.Add(time.Second)) + if err != nil { + t.Fatal(err) + } + if second[0].Backend != string(BackendOpenAI) || second[0].Models[0] != "model-a" || second[0].ModelDetails[0].ID != "model-a" || second[0].ModelDetails[0].InputMediaTypes[0] != "image/png" || second[0].ModelDetails[0].SupportedEfforts[0] != api.EffortLow { + t.Fatalf("cached adapters were mutated through a returned value: %+v", second) + } +} + +func TestCachedAdaptersRejectsUnsettledProbeState(t *testing.T) { + prevProbe, prevAuthProbe := adapterProbe, adapterAuthProbe + prevCache, prevAt, prevFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint + t.Cleanup(func() { + adapterProbe, adapterAuthProbe = prevProbe, prevAuthProbe + adapterCache, adapterCacheAt, adapterCacheFingerprint = prevCache, prevAt, prevFingerprint + }) + adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" + + home := t.TempDir() + captures := 0 + adapterAuthProbe = func() AuthProbe { + captures++ + probe := fakeProbe(nil, nil, nil, home) + probe.APICredentials = map[Backend]api.ResolvedAPIKey{ + BackendOpenAI: {Token: string(rune('a' + captures))}, + } + return probe + } + adapterProbe = func(AuthProbe) ([]AdapterStatus, error) { + return []AdapterStatus{{Backend: string(BackendOpenAI)}}, nil + } + + if got, err := CachedAdapters(time.Unix(6_000_000, 0)); err == nil || got != nil { + t.Fatalf("unsettled probe returned adapters=%+v err=%v", got, err) + } +} + +func TestCachedAdaptersRejectsRecaptureError(t *testing.T) { + prevProbe, prevAuthProbe := adapterProbe, adapterAuthProbe + prevCache, prevAt, prevFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint + t.Cleanup(func() { + adapterProbe, adapterAuthProbe = prevProbe, prevAuthProbe + adapterCache, adapterCacheAt, adapterCacheFingerprint = prevCache, prevAt, prevFingerprint + }) + adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" + + wantErr := errors.New("credential vault became unreadable") + captures := 0 + adapterAuthProbe = func() AuthProbe { + captures++ + probe := fakeProbe(nil, nil, nil, t.TempDir()) + if captures > 1 { + probe.ProbeError = wantErr + } + return probe + } + adapterProbe = func(AuthProbe) ([]AdapterStatus, error) { + return []AdapterStatus{{Backend: string(BackendOpenAI)}}, nil + } + + got, err := CachedAdapters(time.Unix(7_000_000, 0)) + if !errors.Is(err, wantErr) || got != nil { + t.Fatalf("recapture error returned adapters=%+v err=%v", got, err) + } +} diff --git a/pkg/ai/adapters_test.go b/pkg/ai/adapters_test.go index f0dfb826..acfec39b 100644 --- a/pkg/ai/adapters_test.go +++ b/pkg/ai/adapters_test.go @@ -310,6 +310,57 @@ func TestProbeAdaptersNoCacheBypassesPersistedModelCache(t *testing.T) { } } +func TestProbeAdaptersUsesExactCallerCredentialSnapshot(t *testing.T) { + credentials := map[Backend]api.ResolvedAPIKey{ + BackendOpenAI: {Token: "caller-token-b", Source: "test", Detail: "caller"}, + } + probe := fakeProbe(map[string]string{"OPENAI_API_KEY": "process-token-a"}, nil, nil, "/home/u") + probe.APICredentials = credentials + + prev := resolveModelRows + resolveModelRows = func(_ context.Context, opts ResolveOptions) ([]ResolvedModel, error) { + // Mutating the caller-owned map after ProbeAdapters starts must not alter + // the operation snapshot passed to model resolution. + credentials[BackendOpenAI] = api.ResolvedAPIKey{Token: "mutated-token-c"} + if got := opts.Credentials.APIKey(BackendOpenAI).Token; got != "caller-token-b" { + t.Fatalf("model resolver credential = %q, want caller-token-b", got) + } + return []ResolvedModel{{ + Model: Model{ID: "openai/private-model-b", Backend: BackendOpenAI}, + Live: true, + }}, nil + } + t.Cleanup(func() { resolveModelRows = prev }) + + adapters, err := ProbeAdapters(WhoamiOptions{Backend: string(BackendOpenAI), Models: true}, probe) + if err != nil { + t.Fatalf("ProbeAdapters: %v", err) + } + if len(adapters) != 1 || adapters[0].AuthDetail != "call…en-b" || !stringSliceContains(adapters[0].Models, "private-model-b") { + t.Fatalf("adapter = %+v, want auth and availability from caller-token-b", adapters) + } +} + +func TestProbeAdaptersExplicitEmptyCredentialsDoNotFallBackToEnvironment(t *testing.T) { + probe := fakeProbe(map[string]string{"OPENAI_API_KEY": "process-token-a"}, nil, nil, "/home/u") + probe.APICredentials = map[Backend]api.ResolvedAPIKey{} + + prev := resolveModelRows + resolveModelRows = func(context.Context, ResolveOptions) ([]ResolvedModel, error) { + t.Fatal("model resolver must not run for an explicitly empty credential snapshot") + return nil, nil + } + t.Cleanup(func() { resolveModelRows = prev }) + + adapters, err := ProbeAdapters(WhoamiOptions{Backend: string(BackendOpenAI), Models: true}, probe) + if err != nil { + t.Fatalf("ProbeAdapters: %v", err) + } + if len(adapters) != 1 || adapters[0].Authenticated || adapters[0].ModelError == "" { + t.Fatalf("adapter = %+v, want no auth or live availability", adapters) + } +} + func TestProbeAdaptersUsesCodexDebugModelsOnceRegardlessOfAPIKey(t *testing.T) { probe := fakeProbe(map[string]string{"OPENAI_API_KEY": "sk-test"}, map[string]string{"codex": "/usr/local/bin/codex"}, nil, "/home/u") calls := 0 diff --git a/pkg/ai/availability_test.go b/pkg/ai/availability_test.go index 7380c093..f78a9f53 100644 --- a/pkg/ai/availability_test.go +++ b/pkg/ai/availability_test.go @@ -57,13 +57,13 @@ var _ = Describe("runtime availability", func() { It("uses each adapter's live readiness in the runtime catalog", func() { previousProbe := adapterProbe - previousCache, previousAt := adapterCache, adapterCacheAt + previousCache, previousAt, previousFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint DeferCleanup(func() { adapterProbe = previousProbe - adapterCache, adapterCacheAt = previousCache, previousAt + adapterCache, adapterCacheAt, adapterCacheFingerprint = previousCache, previousAt, previousFingerprint }) - adapterCache, adapterCacheAt = nil, time.Time{} - adapterProbe = func() ([]AdapterStatus, error) { + adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" + adapterProbe = func(AuthProbe) ([]AdapterStatus, error) { return []AdapterStatus{ {Backend: string(BackendOpenAI), Type: "api"}, {Backend: string(BackendCodexCLI), Type: "cli", Binary: "/bin/codex"}, diff --git a/pkg/ai/catalog_disabled_ginkgo_test.go b/pkg/ai/catalog_disabled_ginkgo_test.go index 3adc1344..23f9c298 100644 --- a/pkg/ai/catalog_disabled_ginkgo_test.go +++ b/pkg/ai/catalog_disabled_ginkgo_test.go @@ -110,13 +110,13 @@ var _ = Describe("catalog opt-out filtering", func() { It("retains disabled models with remediation in the live descriptive menu", func() { previousProbe := adapterProbe - previousCache, previousAt := adapterCache, adapterCacheAt + previousCache, previousAt, previousFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint DeferCleanup(func() { adapterProbe = previousProbe - adapterCache, adapterCacheAt = previousCache, previousAt + adapterCache, adapterCacheAt, adapterCacheFingerprint = previousCache, previousAt, previousFingerprint }) - adapterCache, adapterCacheAt = nil, time.Time{} - adapterProbe = func() ([]AdapterStatus, error) { return nil, nil } + adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" + adapterProbe = func(AuthProbe) ([]AdapterStatus, error) { return nil, nil } disable(nil, []string{"deepseek"}, nil, nil, nil) infos, err := LiveCatalogInfo(nil) diff --git a/pkg/ai/catalog_resolve.go b/pkg/ai/catalog_resolve.go index 8b0cb7a1..1c602cbc 100644 --- a/pkg/ai/catalog_resolve.go +++ b/pkg/ai/catalog_resolve.go @@ -2,9 +2,9 @@ package ai import ( "context" - "crypto/pbkdf2" - "crypto/sha512" - "encoding/hex" + "crypto/hmac" + "crypto/sha256" + "encoding/json" "fmt" "sort" "strings" @@ -50,51 +50,104 @@ func (r ResolvedModel) Context() int { // ResolveOptions controls a ResolveModels query. type ResolveOptions struct { - Backend Backend // empty = all backends - Filter string // substring filter on id/label; non-empty also reveals legacy ids - UseTokens bool // when true, augment API backends (that have a key) with live /v1/models - Refresh bool // bypass the persisted cache and re-resolve + Backend Backend // empty = all backends + Filter string // substring filter on id/label; non-empty also reveals legacy ids + UseTokens bool // when true, augment API backends (that have a key) with live /v1/models + Refresh bool // bypass the persisted cache and re-resolve + Credentials CredentialSnapshot + APIURL string // optional provider base URL; valid only for a selected API backend } // liveModelFetcher fetches a backend's live model list. It is a package var so // tests can stub it without hitting the network. -var liveModelFetcher = func(ctx context.Context, b Backend) ([]ModelDef, error) { - return ListModels(ctx, b) +var liveModelFetcher = func(ctx context.Context, b Backend, token, endpoint string) ([]ModelDef, error) { + return listModelsWithAPIKeyAtEndpoint(ctx, b, token, endpoint) } // apiBackends are the direct-API backends the resolver can list live. var apiBackends = []Backend{BackendAnthropic, BackendOpenAI, BackendGemini, BackendDeepSeek} // ResolveModels returns the merged catalog ∪ live-API view for opts, joined to -// merged OpenRouter/static pricing and persisted to ~/.config/captain/models.json. +// merged OpenRouter/static pricing and persisted in an auth-scoped cache entry. // A fresh, fingerprint-matching cache is reused (unless opts.Refresh). func ResolveModels(ctx context.Context, opts ResolveOptions) ([]ResolvedModel, error) { - fp, err := resolveFingerprint(opts) + credentials, err := credentialSnapshotForOptions(opts) + if err != nil { + return nil, err + } + fp, cacheable, err := resolveFingerprint(opts, credentials) if err != nil { return nil, err } - rows, ok := cachedRows(opts, fp) - if !ok { - var err error - rows, err = resolveFresh(ctx, opts) - if err != nil { - return nil, err + if !cacheable { + rows, err := resolveFresh(ctx, opts, credentials) + return filterResolved(rows, opts.Filter), err + } + + unlock, err := lockModelCache(fp) + if err != nil { + catalogLog.Warnf("model disk cache disabled: failed to lock cache entry: %v", err) + rows, resolveErr := resolveFresh(ctx, opts, credentials) + return filterResolved(rows, opts.Filter), resolveErr + } + defer unlock() + + if rows, ok := cachedRows(opts, fp); ok { + return filterResolved(rows, opts.Filter), nil + } + rows, err := resolveFresh(ctx, opts, credentials) + if err != nil { + return nil, err + } + if err := saveModelCache(fp, rows); err != nil { + catalogLog.Warnf("failed to persist model cache: %v", err) + } + return filterResolved(rows, opts.Filter), nil +} + +func credentialSnapshotForOptions(opts ResolveOptions) (CredentialSnapshot, error) { + if strings.TrimSpace(opts.APIURL) != "" && (opts.Backend == "" || opts.Backend.Kind() != "api") { + return CredentialSnapshot{}, fmt.Errorf("APIURL requires one selected API backend") + } + if opts.Credentials.supplied { + return opts.Credentials.clone(), nil + } + resolved := make(map[Backend]api.ResolvedAPIKey) + if opts.UseTokens { + for _, backend := range selectedAPIBackends(opts.Backend) { + credential, err := ResolveAPIKey(backend) + if err != nil { + return CredentialSnapshot{}, err + } + resolved[backend] = credential } - if err := saveModelCache(fp, rows); err != nil { - catalogLog.Warnf("failed to persist model cache: %v", err) + } + return NewCredentialSnapshot(resolved), nil +} + +func selectedAPIBackends(backend Backend) []Backend { + if backend == "" { + return append([]Backend(nil), apiBackends...) + } + if backend.Kind() != "api" { + return nil + } + for _, candidate := range apiBackends { + if candidate == backend { + return []Backend{backend} } } - return filterResolved(rows, opts.Filter), nil + return nil } // cachedRows returns the persisted rows when they are fresh and match the -// current fingerprint. +// current fingerprint. Callers hold the entry lock while reading. func cachedRows(opts ResolveOptions, fp string) ([]ResolvedModel, bool) { if opts.Refresh { return nil, false } - c, err := loadModelCache() + c, err := loadModelCache(fp) if err != nil || c == nil || c.expired() || c.KeyFingerprint != fp { return nil, false } @@ -105,11 +158,11 @@ func cachedRows(opts ResolveOptions, fp string) ([]ResolvedModel, bool) { // when tokens are present, and joins each row to pricing. opts.Refresh reaches // the pricing snapshot too: bypassing the model cache while still pricing from a // day-old OpenRouter snapshot would only half-honour --no-cache. -func resolveFresh(ctx context.Context, opts ResolveOptions) ([]ResolvedModel, error) { +func resolveFresh(ctx context.Context, opts ResolveOptions, credentials CredentialSnapshot) ([]ResolvedModel, error) { rows, index := seedCatalog(opts.Backend) if opts.UseTokens { - if err := unionLive(ctx, opts.Backend, &rows, index); err != nil { + if err := unionLive(ctx, opts, credentials, &rows, index); err != nil { return nil, err } } @@ -159,19 +212,17 @@ func catalogBackendMatch(want, modelBackend Backend) bool { // unionLive fetches live /v1/models for every API backend (matching the filter) // that has a key set, and merges them into rows. A fetch error fails loud. -func unionLive(ctx context.Context, backend Backend, rows *[]ResolvedModel, index map[modelKey]int) error { - for _, b := range apiBackends { - if backend != "" && b != backend { +func unionLive(ctx context.Context, opts ResolveOptions, credentials CredentialSnapshot, rows *[]ResolvedModel, index map[modelKey]int) error { + for _, b := range selectedAPIBackends(opts.Backend) { + resolved := credentials.APIKey(b) + if strings.TrimSpace(resolved.Token) == "" { continue } - resolved, err := ResolveAPIKey(b) + endpoint, err := modelListEndpoint(b, opts.APIURL) if err != nil { return err } - if resolved.Token == "" { - continue - } - live, err := liveModelFetcher(ctx, b) + live, err := liveModelFetcher(ctx, b, resolved.Token, endpoint) if err != nil { return fmt.Errorf("%s: %w", b, err) } @@ -218,26 +269,70 @@ func filterResolved(rows []ResolvedModel, filter string) []ResolvedModel { // - v2: model identity unified on the provider descriptors; catalog pricing now // resolves through the same prefixed-first path as billing, so cached Claude // prices from the static fallback table are stale. -const resolveSchemaVersion = "v2" -const resolveFingerprintSalt = "captain/model-cache/api-key" +// - v3: live sources are namespaced by exact credential and model endpoint, +// using a machine-local keyed HMAC instead of a comparable fixed-salt hash. +const resolveSchemaVersion = "v3" -func resolveFingerprint(opts ResolveOptions) (string, error) { - var present []string - for _, b := range apiBackends { - resolved, err := ResolveAPIKey(b) +type resolveCacheSource struct { + Backend Backend `json:"backend"` + EndpointHash string `json:"endpointHash"` + TokenHMAC string `json:"tokenHMAC"` +} + +type resolveCacheDescriptor struct { + Schema string `json:"schema"` + Backend Backend `json:"backend"` + UseTokens bool `json:"useTokens"` + LiveSources []resolveCacheSource `json:"liveSources,omitempty"` +} + +// resolveFingerprint returns a canonical, non-secret cache identity. A secure +// HMAC-key failure disables token-bearing disk caching rather than failing live +// model discovery or falling back to a comparable unkeyed token hash. +func resolveFingerprint(opts ResolveOptions, credentials CredentialSnapshot) (string, bool, error) { + descriptor := resolveCacheDescriptor{ + Schema: resolveSchemaVersion, + Backend: opts.Backend, + UseTokens: opts.UseTokens, + } + var hmacKey []byte + for _, backend := range selectedAPIBackends(opts.Backend) { + if !opts.UseTokens { + break + } + resolved := credentials.APIKey(backend) + if strings.TrimSpace(resolved.Token) == "" { + continue + } + endpoint, err := modelListEndpoint(backend, opts.APIURL) if err != nil { - return "", err + return "", false, err } - if resolved.Token != "" { - fingerprint, err := pbkdf2.Key(sha512.New, resolved.Token, []byte(resolveFingerprintSalt), 4096, 8) + if hmacKey == nil { + hmacKey, err = modelCacheHMACKey() if err != nil { - return "", fmt.Errorf("fingerprint %s API key: %w", b, err) + catalogLog.Warnf("model disk cache disabled: secure identity key unavailable") + return "", false, nil } - present = append(present, string(b)+":"+hex.EncodeToString(fingerprint)) } + endpointHash := sha256.Sum256([]byte(endpoint)) + mac := hmac.New(sha256.New, hmacKey) + _, _ = mac.Write([]byte(resolved.Token)) + descriptor.LiveSources = append(descriptor.LiveSources, resolveCacheSource{ + Backend: backend, + EndpointHash: fmt.Sprintf("%x", endpointHash), + TokenHMAC: fmt.Sprintf("%x", mac.Sum(nil)), + }) + } + sort.Slice(descriptor.LiveSources, func(i, j int) bool { + return descriptor.LiveSources[i].Backend < descriptor.LiveSources[j].Backend + }) + encoded, err := json.Marshal(descriptor) + if err != nil { + return "", false, err } - sort.Strings(present) - return fmt.Sprintf("v=%s|b=%s|tok=%v|keys=%s", resolveSchemaVersion, opts.Backend, opts.UseTokens, strings.Join(present, ",")), nil + fingerprint := sha256.Sum256(encoded) + return fmt.Sprintf("%x", fingerprint), true, nil } func bareModelID(id string) string { diff --git a/pkg/ai/catalog_resolve_test.go b/pkg/ai/catalog_resolve_test.go index 95297d42..24a95e50 100644 --- a/pkg/ai/catalog_resolve_test.go +++ b/pkg/ai/catalog_resolve_test.go @@ -3,14 +3,19 @@ package ai import ( "context" "errors" + "fmt" + "os" "path/filepath" "strings" + "sync" "testing" + "github.com/flanksource/captain/pkg/api" "github.com/flanksource/captain/pkg/credentials" ) func TestResolveFingerprintChangesWhenVaultTokenRotates(t *testing.T) { + t.Setenv("HOME", t.TempDir()) credentials.SetPathForTesting(filepath.Join(t.TempDir(), "vault")) t.Cleanup(func() { credentials.SetPathForTesting("") }) t.Setenv("OPENAI_API_KEY", "") @@ -21,16 +26,24 @@ func TestResolveFingerprintChangesWhenVaultTokenRotates(t *testing.T) { if err := vault.Set("openai", "first-secret-token"); err != nil { t.Fatalf("Set first token: %v", err) } - first, err := resolveFingerprint(ResolveOptions{Backend: BackendOpenAI, UseTokens: true}) + firstCredentials, err := credentialSnapshotForOptions(ResolveOptions{Backend: BackendOpenAI, UseTokens: true}) if err != nil { - t.Fatalf("first fingerprint: %v", err) + t.Fatalf("first credential snapshot: %v", err) + } + first, cacheable, err := resolveFingerprint(ResolveOptions{Backend: BackendOpenAI, UseTokens: true}, firstCredentials) + if err != nil || !cacheable { + t.Fatalf("first fingerprint: cacheable=%v err=%v", cacheable, err) } if err := vault.Set("openai", "second-secret-token"); err != nil { t.Fatalf("Set second token: %v", err) } - second, err := resolveFingerprint(ResolveOptions{Backend: BackendOpenAI, UseTokens: true}) + secondCredentials, err := credentialSnapshotForOptions(ResolveOptions{Backend: BackendOpenAI, UseTokens: true}) if err != nil { - t.Fatalf("second fingerprint: %v", err) + t.Fatalf("second credential snapshot: %v", err) + } + second, cacheable, err := resolveFingerprint(ResolveOptions{Backend: BackendOpenAI, UseTokens: true}, secondCredentials) + if err != nil || !cacheable { + t.Fatalf("second fingerprint: cacheable=%v err=%v", cacheable, err) } if first == second { t.Fatal("token rotation must invalidate the model cache fingerprint") @@ -52,7 +65,7 @@ func stubLiveFetcher(t *testing.T, fn func(b Backend) ([]ModelDef, error)) { t.Setenv("GOOGLE_API_KEY", "") prev := liveModelFetcher - liveModelFetcher = func(_ context.Context, b Backend) ([]ModelDef, error) { return fn(b) } + liveModelFetcher = func(_ context.Context, b Backend, _, _ string) ([]ModelDef, error) { return fn(b) } t.Cleanup(func() { liveModelFetcher = prev }) } @@ -226,3 +239,300 @@ func TestResolveModels_CacheWrittenAndReused(t *testing.T) { t.Fatalf("refresh should re-fetch, got %d calls", calls) } } + +func TestResolveModelsCredentialCacheIsolationAndSourceConsistency(t *testing.T) { + t.Setenv("HOME", t.TempDir()) + prev := liveModelFetcher + calls := map[string]int{} + liveModelFetcher = func(_ context.Context, backend Backend, token, _ string) ([]ModelDef, error) { + if backend != BackendOpenAI { + t.Fatalf("unexpected live backend %s", backend) + } + calls[token]++ + modelID := map[string]string{"token-a": "private-model-a", "token-b": "private-model-b"}[token] + return []ModelDef{{ID: modelID, Backend: backend}}, nil + } + t.Cleanup(func() { liveModelFetcher = prev }) + + resolve := func(openAIToken, source, unrelatedToken string) []ResolvedModel { + t.Helper() + rows, err := ResolveModels(context.Background(), ResolveOptions{ + Backend: BackendOpenAI, + UseTokens: true, + Credentials: NewCredentialSnapshot(map[Backend]api.ResolvedAPIKey{ + BackendOpenAI: {Token: openAIToken, Source: source, Detail: source}, + BackendAnthropic: {Token: unrelatedToken, Source: source, Detail: source}, + }), + }) + if err != nil { + t.Fatalf("ResolveModels(%s): %v", openAIToken, err) + } + return rows + } + + if _, ok := hasModelID(resolve("token-b", credentials.SourceEnvironment, "unrelated-a"), "private-model-b"); !ok { + t.Fatal("token B availability missing") + } + // Source and unrelated-provider changes must reuse the same OpenAI entry. + if _, ok := hasModelID(resolve("token-b", credentials.SourceVault, "unrelated-b"), "private-model-b"); !ok { + t.Fatal("same effective token should reuse availability across sources") + } + if _, ok := hasModelID(resolve("token-a", credentials.SourceEnvironment, "unrelated-b"), "private-model-a"); !ok { + t.Fatal("token A availability missing") + } + if _, ok := hasModelID(resolve("token-b", credentials.SourceEnvironment, "unrelated-c"), "private-model-b"); !ok { + t.Fatal("token B cache should survive resolving token A") + } + if calls["token-a"] != 1 || calls["token-b"] != 1 { + t.Fatalf("live calls = %v, want one isolated fetch per effective OpenAI token", calls) + } +} + +func TestResolveModelsEndpointIsolation(t *testing.T) { + t.Setenv("HOME", t.TempDir()) + prev := liveModelFetcher + var endpoints []string + liveModelFetcher = func(_ context.Context, backend Backend, _ string, endpoint string) ([]ModelDef, error) { + endpoints = append(endpoints, endpoint) + return []ModelDef{{ID: fmt.Sprintf("endpoint-%d", len(endpoints)), Backend: backend}}, nil + } + t.Cleanup(func() { liveModelFetcher = prev }) + credentials := NewCredentialSnapshot(map[Backend]api.ResolvedAPIKey{ + BackendOpenAI: {Token: "endpoint-token"}, + }) + + for _, apiURL := range []string{ + "https://tenant-one.example/v1?account=private-one", + "https://tenant-one.example/v1?account=private-one", + "https://tenant-two.example/v1?account=private-two", + } { + if _, err := ResolveModels(context.Background(), ResolveOptions{ + Backend: BackendOpenAI, UseTokens: true, Credentials: credentials, APIURL: apiURL, + }); err != nil { + t.Fatalf("ResolveModels(%s): %v", apiURL, err) + } + } + if len(endpoints) != 2 { + t.Fatalf("live endpoints = %v, want one fetch per exact endpoint", endpoints) + } + if endpoints[0] != "https://tenant-one.example/v1/models?account=private-one" || endpoints[1] != "https://tenant-two.example/v1/models?account=private-two" { + t.Fatalf("resolved endpoints = %v", endpoints) + } + entries, err := os.ReadDir(filepath.Join(os.Getenv("HOME"), ".config", "captain", "models")) + if err != nil { + t.Fatalf("ReadDir model cache: %v", err) + } + for _, entry := range entries { + data, err := os.ReadFile(filepath.Join(os.Getenv("HOME"), ".config", "captain", "models", entry.Name())) + if err != nil { + t.Fatal(err) + } + if strings.Contains(string(data), "private-one") || strings.Contains(string(data), "private-two") { + t.Fatalf("cache %s persisted a plaintext endpoint identifier", entry.Name()) + } + } +} + +func TestResolveModelsConcurrentBackendsRetainSecureIndependentEntries(t *testing.T) { + home := t.TempDir() + t.Setenv("HOME", home) + root := filepath.Join(home, ".config", "captain") + if err := os.MkdirAll(root, 0o755); err != nil { + t.Fatal(err) + } + legacy := filepath.Join(root, "models.json") + if err := os.WriteFile(legacy, []byte(`{"models":[{"id":"legacy-private"}]}`), 0o644); err != nil { + t.Fatal(err) + } + + prev := liveModelFetcher + var mu sync.Mutex + calls := map[Backend]int{} + liveModelFetcher = func(_ context.Context, backend Backend, _, _ string) ([]ModelDef, error) { + mu.Lock() + calls[backend]++ + mu.Unlock() + return []ModelDef{{ID: string(backend) + "-private", Backend: backend}}, nil + } + t.Cleanup(func() { liveModelFetcher = prev }) + + options := []ResolveOptions{ + {Backend: BackendOpenAI, UseTokens: true, Credentials: NewCredentialSnapshot(map[Backend]api.ResolvedAPIKey{BackendOpenAI: {Token: "openai-secret-token"}})}, + {Backend: BackendAnthropic, UseTokens: true, Credentials: NewCredentialSnapshot(map[Backend]api.ResolvedAPIKey{BackendAnthropic: {Token: "anthropic-secret-token"}})}, + } + runConcurrent := func(options []ResolveOptions) { + t.Helper() + var wg sync.WaitGroup + errs := make(chan error, len(options)) + for _, opts := range options { + wg.Add(1) + go func(opts ResolveOptions) { + defer wg.Done() + _, err := ResolveModels(context.Background(), opts) + errs <- err + }(opts) + } + wg.Wait() + close(errs) + for err := range errs { + if err != nil { + t.Fatalf("concurrent ResolveModels: %v", err) + } + } + } + runConcurrent(options) + runConcurrent(options) + if calls[BackendOpenAI] != 1 || calls[BackendAnthropic] != 1 { + t.Fatalf("warm-cache calls = %v, want one fetch per backend", calls) + } + refresh := options[0] + refresh.Refresh = true + if _, err := ResolveModels(context.Background(), refresh); err != nil { + t.Fatalf("refresh OpenAI: %v", err) + } + if _, err := ResolveModels(context.Background(), options[1]); err != nil { + t.Fatalf("warm Anthropic after OpenAI refresh: %v", err) + } + if calls[BackendOpenAI] != 2 || calls[BackendAnthropic] != 1 { + t.Fatalf("refresh calls = %v, want refresh isolated to OpenAI", calls) + } + + assertMode := func(path string, want os.FileMode) { + t.Helper() + info, err := os.Stat(path) + if err != nil { + t.Fatalf("stat %s: %v", path, err) + } + if got := info.Mode().Perm(); got != want { + t.Fatalf("mode %s = %04o, want %04o", path, got, want) + } + } + assertMode(root, 0o700) + assertMode(filepath.Join(root, "models"), 0o700) + assertMode(filepath.Join(root, "model-cache.key"), 0o600) + assertMode(legacy, 0o600) + entries, err := os.ReadDir(filepath.Join(root, "models")) + if err != nil { + t.Fatal(err) + } + jsonEntries := 0 + for _, entry := range entries { + if strings.Contains(entry.Name(), ".tmp") { + t.Fatalf("temporary cache file remains: %s", entry.Name()) + } + path := filepath.Join(root, "models", entry.Name()) + assertMode(path, 0o600) + if filepath.Ext(entry.Name()) != ".json" { + continue + } + jsonEntries++ + data, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + if strings.Contains(string(data), "openai-secret-token") || strings.Contains(string(data), "anthropic-secret-token") { + t.Fatalf("cache %s persisted raw credential material", entry.Name()) + } + } + if jsonEntries != 2 { + t.Fatalf("cache entries = %v, want one JSON entry per backend", entries) + } +} + +func TestResolveModelsSerializesSameCredentialCacheEntry(t *testing.T) { + t.Setenv("HOME", t.TempDir()) + previous := liveModelFetcher + started := make(chan struct{}) + release := make(chan struct{}) + calls := 0 + var mu sync.Mutex + liveModelFetcher = func(_ context.Context, backend Backend, _, _ string) ([]ModelDef, error) { + mu.Lock() + calls++ + call := calls + mu.Unlock() + if call == 1 { + close(started) + <-release + } + return []ModelDef{{ID: fmt.Sprintf("private-model-%d", call), Backend: backend}}, nil + } + t.Cleanup(func() { liveModelFetcher = previous }) + opts := ResolveOptions{ + Backend: BackendOpenAI, UseTokens: true, + Credentials: NewCredentialSnapshot(map[Backend]api.ResolvedAPIKey{BackendOpenAI: {Token: "shared-token"}}), + } + + firstDone := make(chan error, 1) + go func() { + _, err := ResolveModels(context.Background(), opts) + firstDone <- err + }() + <-started + secondDone := make(chan []ResolvedModel, 1) + go func() { + rows, _ := ResolveModels(context.Background(), opts) + secondDone <- rows + }() + close(release) + if err := <-firstDone; err != nil { + t.Fatalf("first resolve: %v", err) + } + rows := <-secondDone + if _, ok := hasModelID(rows, "private-model-1"); !ok { + t.Fatalf("second resolve did not reuse serialized cache entry: %+v", rows) + } + mu.Lock() + defer mu.Unlock() + if calls != 1 { + t.Fatalf("live fetches = %d, want one for the shared cache identity", calls) + } +} + +func TestResolveModelsCatalogOnlyCreatesNoCredentialVerifier(t *testing.T) { + home := t.TempDir() + t.Setenv("HOME", home) + prev := liveModelFetcher + liveModelFetcher = func(_ context.Context, backend Backend, _, _ string) ([]ModelDef, error) { + t.Fatalf("catalog-only resolution fetched %s", backend) + return nil, nil + } + t.Cleanup(func() { liveModelFetcher = prev }) + if _, err := ResolveModels(context.Background(), ResolveOptions{ + Backend: BackendAnthropic, UseTokens: false, + Credentials: NewCredentialSnapshot(map[Backend]api.ResolvedAPIKey{BackendAnthropic: {Token: "must-not-be-fingerprinted"}}), + }); err != nil { + t.Fatalf("ResolveModels: %v", err) + } + if _, err := os.Stat(filepath.Join(home, ".config", "captain", "model-cache.key")); !os.IsNotExist(err) { + t.Fatalf("catalog-only resolution created a credential verifier: %v", err) + } +} + +func TestResolveModelsSecureKeyFailureFallsBackToUncachedLiveResolution(t *testing.T) { + t.Setenv("HOME", t.TempDir()) + prevFetcher := liveModelFetcher + prevKey := modelCacheHMACKey + calls := 0 + liveModelFetcher = func(_ context.Context, backend Backend, _, _ string) ([]ModelDef, error) { + calls++ + return []ModelDef{{ID: "uncached-private", Backend: backend}}, nil + } + modelCacheHMACKey = func() ([]byte, error) { return nil, errors.New("unavailable") } + t.Cleanup(func() { + liveModelFetcher = prevFetcher + modelCacheHMACKey = prevKey + }) + opts := ResolveOptions{ + Backend: BackendOpenAI, UseTokens: true, + Credentials: NewCredentialSnapshot(map[Backend]api.ResolvedAPIKey{BackendOpenAI: {Token: "still-usable-live"}}), + } + for i := 0; i < 2; i++ { + if _, err := ResolveModels(context.Background(), opts); err != nil { + t.Fatalf("ResolveModels attempt %d: %v", i+1, err) + } + } + if calls != 2 { + t.Fatalf("live calls = %d, want uncached fetch after each secure-key failure", calls) + } +} diff --git a/pkg/ai/live_catalog_test.go b/pkg/ai/live_catalog_test.go index 377a631d..13fcd3db 100644 --- a/pkg/ai/live_catalog_test.go +++ b/pkg/ai/live_catalog_test.go @@ -70,13 +70,13 @@ func TestMergeLiveCatalogUpsertsLiveAndPreservesStatic(t *testing.T) { func TestLiveCatalogInfoAppliesPerProviderConfigured(t *testing.T) { prevProbe := adapterProbe - prevCache, prevAt := adapterCache, adapterCacheAt + prevCache, prevAt, prevFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint t.Cleanup(func() { adapterProbe = prevProbe - adapterCache, adapterCacheAt = prevCache, prevAt + adapterCache, adapterCacheAt, adapterCacheFingerprint = prevCache, prevAt, prevFingerprint }) - adapterCache, adapterCacheAt = nil, time.Time{} - adapterProbe = func() ([]AdapterStatus, error) { + adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" + adapterProbe = func(AuthProbe) ([]AdapterStatus, error) { return []AdapterStatus{ {Backend: string(BackendAnthropic), Type: "api", ModelDetails: []ModelDef{ {ID: "claude-sonnet-5", Name: "Claude Sonnet 5", Reasoning: true}, diff --git a/pkg/ai/model_cache.go b/pkg/ai/model_cache.go index 45546c6e..e76ecd61 100644 --- a/pkg/ai/model_cache.go +++ b/pkg/ai/model_cache.go @@ -1,20 +1,24 @@ package ai import ( + "crypto/rand" "encoding/json" + "fmt" "os" "path/filepath" + "sync" "time" + + "golang.org/x/sys/unix" ) // modelCacheTTL bounds how long a persisted resolve is reused before the // resolver re-queries live model lists. const modelCacheTTL = 24 * time.Hour -// ModelCache is the persisted merged model view written to -// ~/.config/captain/models.json. KeyFingerprint records which resolve produced -// it (backend filter + token use + which API keys were present) so a changed -// environment re-resolves instead of serving a stale live view. +// ModelCache is one persisted merged model view. KeyFingerprint is a canonical, +// non-secret identity covering the resolution schema, backend, effective model +// endpoint, and a machine-keyed HMAC of each exact credential used. type ModelCache struct { Timestamp time.Time `json:"timestamp"` KeyFingerprint string `json:"keyFingerprint"` @@ -25,24 +29,102 @@ func (c *ModelCache) expired() bool { return time.Since(c.Timestamp) >= modelCacheTTL } -// modelCachePath returns ~/.config/captain/models.json, creating the directory. -func modelCachePath() (string, error) { +const modelCacheHMACKeySize = 32 + +var ( + modelCacheKeyMu sync.Mutex + modelCacheHMACKey = loadOrCreateModelCacheHMACKey +) + +func modelCacheRoot() (string, error) { home, err := os.UserHomeDir() if err != nil { return "", err } dir := filepath.Join(home, ".config", "captain") - if err := os.MkdirAll(dir, 0o755); err != nil { + if err := os.MkdirAll(dir, 0o700); err != nil { + return "", err + } + if err := os.Chmod(dir, 0o700); err != nil { return "", err } - return filepath.Join(dir, "models.json"), nil + // Never reuse the legacy single-slot cache because it cannot be associated + // with one credential. Tighten it in place so private model IDs left by an + // older Captain are no longer world-readable. + legacy := filepath.Join(dir, "models.json") + if info, statErr := os.Lstat(legacy); statErr == nil && info.Mode().IsRegular() { + _ = os.Chmod(legacy, 0o600) + } + return dir, nil } -func loadModelCache() (*ModelCache, error) { - path, err := modelCachePath() +func modelCacheDir() (string, error) { + root, err := modelCacheRoot() + if err != nil { + return "", err + } + dir := filepath.Join(root, "models") + if err := os.MkdirAll(dir, 0o700); err != nil { + return "", err + } + if err := os.Chmod(dir, 0o700); err != nil { + return "", err + } + return dir, nil +} + +func modelCachePath(fingerprint string) (string, error) { + dir, err := modelCacheDir() + if err != nil { + return "", err + } + return filepath.Join(dir, fingerprint+".json"), nil +} + +// lockModelCache serializes one cache identity across goroutines and Captain +// processes. The lock spans cache re-check, live fetch, and publish so a slow +// stale fetch cannot overwrite a newer refresh with a fresh timestamp. +func lockModelCache(fingerprint string) (func(), error) { + dir, err := modelCacheDir() if err != nil { return nil, err } + lock, err := os.OpenFile(filepath.Join(dir, fingerprint+".lock"), os.O_CREATE|os.O_RDWR, 0o600) + if err != nil { + return nil, err + } + if err := lock.Chmod(0o600); err != nil { + _ = lock.Close() + return nil, err + } + if err := unix.Flock(int(lock.Fd()), unix.LOCK_EX); err != nil { + _ = lock.Close() + return nil, err + } + return func() { + _ = unix.Flock(int(lock.Fd()), unix.LOCK_UN) + _ = lock.Close() + }, nil +} + +func loadModelCache(fingerprint string) (*ModelCache, error) { + path, err := modelCachePath(fingerprint) + if err != nil { + return nil, err + } + info, err := os.Lstat(path) + if err != nil { + if os.IsNotExist(err) { + return nil, nil + } + return nil, err + } + if !info.Mode().IsRegular() { + return nil, fmt.Errorf("model cache entry is not a regular file") + } + if err := os.Chmod(path, 0o600); err != nil { + return nil, err + } data, err := os.ReadFile(path) if err != nil { if os.IsNotExist(err) { @@ -60,7 +142,7 @@ func loadModelCache() (*ModelCache, error) { // saveModelCache atomically writes the merged view (tmp + rename). func saveModelCache(fingerprint string, models []ResolvedModel) error { - path, err := modelCachePath() + path, err := modelCachePath(fingerprint) if err != nil { return err } @@ -72,9 +154,105 @@ func saveModelCache(fingerprint string, models []ResolvedModel) error { if err != nil { return err } - tmp := path + ".tmp" - if err := os.WriteFile(tmp, data, 0o644); err != nil { + return atomicWriteModelCache(path, data) +} + +func atomicWriteModelCache(path string, data []byte) error { + tmp, err := os.CreateTemp(filepath.Dir(path), ".models-*.tmp") + if err != nil { return err } - return os.Rename(tmp, path) + tmpPath := tmp.Name() + defer func() { _ = os.Remove(tmpPath) }() + if err := tmp.Chmod(0o600); err != nil { + _ = tmp.Close() + return err + } + if _, err := tmp.Write(data); err != nil { + _ = tmp.Close() + return err + } + if err := tmp.Sync(); err != nil { + _ = tmp.Close() + return err + } + if err := tmp.Close(); err != nil { + return err + } + if err := os.Rename(tmpPath, path); err != nil { + return err + } + return os.Chmod(path, 0o600) +} + +func loadOrCreateModelCacheHMACKey() ([]byte, error) { + modelCacheKeyMu.Lock() + defer modelCacheKeyMu.Unlock() + root, err := modelCacheRoot() + if err != nil { + return nil, err + } + path := filepath.Join(root, "model-cache.key") + if key, err := readModelCacheHMACKey(path); err == nil { + return key, nil + } else if !os.IsNotExist(err) { + return nil, err + } + + key := make([]byte, modelCacheHMACKeySize) + if _, err := rand.Read(key); err != nil { + return nil, fmt.Errorf("generate model cache identity key: %w", err) + } + tmp, err := os.CreateTemp(root, ".model-cache-key-*.tmp") + if err != nil { + return nil, err + } + tmpPath := tmp.Name() + defer func() { _ = os.Remove(tmpPath) }() + if err := tmp.Chmod(0o600); err != nil { + _ = tmp.Close() + return nil, err + } + if _, err := tmp.Write(key); err != nil { + _ = tmp.Close() + return nil, err + } + if err := tmp.Sync(); err != nil { + _ = tmp.Close() + return nil, err + } + if err := tmp.Close(); err != nil { + return nil, err + } + // A hard link publishes a fully written key without replacing a key another + // Captain process may have created concurrently. If the filesystem cannot + // provide this guarantee, callers disable token-bearing disk caching. + if err := os.Link(tmpPath, path); err != nil { + if key, readErr := readModelCacheHMACKey(path); readErr == nil { + return key, nil + } + return nil, fmt.Errorf("publish model cache identity key: %w", err) + } + return readModelCacheHMACKey(path) +} + +func readModelCacheHMACKey(path string) ([]byte, error) { + info, err := os.Lstat(path) + if err != nil { + return nil, err + } + if !info.Mode().IsRegular() { + return nil, fmt.Errorf("model cache identity key is not a regular file") + } + if err := os.Chmod(path, 0o600); err != nil { + return nil, err + } + key, err := os.ReadFile(path) + if err != nil { + return nil, err + } + if len(key) != modelCacheHMACKeySize { + return nil, fmt.Errorf("model cache identity key has invalid length") + } + return key, nil } diff --git a/pkg/ai/models_remote.go b/pkg/ai/models_remote.go index 2df6bfcf..74ca789b 100644 --- a/pkg/ai/models_remote.go +++ b/pkg/ai/models_remote.go @@ -3,8 +3,10 @@ package ai import ( "context" "encoding/json" + "errors" "fmt" "net/http" + "net/url" "sort" "strings" "time" @@ -41,6 +43,13 @@ type ModelDef struct { // caller surfaces an error to the user instead of blocking the form. const remoteModelsTimeout = 5 * time.Second +var defaultModelListEndpoints = map[Backend]string{ + BackendAnthropic: "https://api.anthropic.com/v1/models", + BackendOpenAI: "https://api.openai.com/v1/models", + BackendGemini: "https://generativelanguage.googleapis.com/v1beta/models", + BackendDeepSeek: "https://api.deepseek.com/models", +} + // modelsListResponse covers the wire shape used by OpenAI, Anthropic, and // Google's Generative Language API (with the field aliases each provider // returns). Decoding into this permissive shape lets a single helper handle @@ -67,20 +76,25 @@ func (e ModelHTTPError) Error() string { return fmt.Sprintf("%s models: HTTP %d", e.Backend, e.StatusCode) } +// ModelTransportError keeps custom endpoint details out of user-visible errors +// while preserving the underlying cause for errors.Is/errors.As callers. +type ModelTransportError struct { + Backend Backend + Err error +} + +func (e ModelTransportError) Error() string { + return fmt.Sprintf("%s models: transport request failed", e.Backend) +} + +func (e ModelTransportError) Unwrap() error { return e.Err } + // FetchOpenAIModels calls https://api.openai.com/v1/models and returns the // available model IDs as ModelDefs scoped to BackendOpenAI. apiKey is sent // as a Bearer token. An empty apiKey returns an error without making a // request. func FetchOpenAIModels(ctx context.Context, apiKey string) ([]ModelDef, error) { - if strings.TrimSpace(apiKey) == "" { - return nil, fmt.Errorf("OPENAI_API_KEY is not set") - } - req, err := http.NewRequestWithContext(ctx, http.MethodGet, "https://api.openai.com/v1/models", nil) - if err != nil { - return nil, err - } - req.Header.Set("Authorization", "Bearer "+apiKey) - return doModelsRequest(req, BackendOpenAI) + return fetchModelsAtEndpoint(ctx, BackendOpenAI, apiKey, defaultModelListEndpoints[BackendOpenAI]) } // FetchAnthropicModels calls https://api.anthropic.com/v1/models and returns @@ -88,16 +102,7 @@ func FetchOpenAIModels(ctx context.Context, apiKey string) ([]ModelDef, error) { // sent via the x-api-key header along with the required anthropic-version // header. An empty apiKey returns an error without making a request. func FetchAnthropicModels(ctx context.Context, apiKey string) ([]ModelDef, error) { - if strings.TrimSpace(apiKey) == "" { - return nil, fmt.Errorf("ANTHROPIC_API_KEY is not set") - } - req, err := http.NewRequestWithContext(ctx, http.MethodGet, "https://api.anthropic.com/v1/models", nil) - if err != nil { - return nil, err - } - req.Header.Set("x-api-key", apiKey) - req.Header.Set("anthropic-version", "2023-06-01") - return doModelsRequest(req, BackendAnthropic) + return fetchModelsAtEndpoint(ctx, BackendAnthropic, apiKey, defaultModelListEndpoints[BackendAnthropic]) } // FetchGeminiModels calls Google's Generative Language ListModels endpoint @@ -105,15 +110,7 @@ func FetchAnthropicModels(ctx context.Context, apiKey string) ([]ModelDef, error // the x-goog-api-key header. The returned `name` field is shaped // "models/gemini-2.5-flash"; we strip the prefix so callers see the bare id. func FetchGeminiModels(ctx context.Context, apiKey string) ([]ModelDef, error) { - if strings.TrimSpace(apiKey) == "" { - return nil, fmt.Errorf("GEMINI_API_KEY is not set") - } - req, err := http.NewRequestWithContext(ctx, http.MethodGet, "https://generativelanguage.googleapis.com/v1beta/models", nil) - if err != nil { - return nil, err - } - req.Header.Set("x-goog-api-key", apiKey) - return doModelsRequest(req, BackendGemini) + return fetchModelsAtEndpoint(ctx, BackendGemini, apiKey, defaultModelListEndpoints[BackendGemini]) } // FetchDeepSeekModels calls https://api.deepseek.com/models and returns the @@ -121,26 +118,108 @@ func FetchGeminiModels(ctx context.Context, apiKey string) ([]ModelDef, error) { // OpenAI-compatible, so the endpoint is a Bearer-authenticated, OpenAI-shaped // listing. An empty apiKey returns an error without making a request. func FetchDeepSeekModels(ctx context.Context, apiKey string) ([]ModelDef, error) { + return fetchModelsAtEndpoint(ctx, BackendDeepSeek, apiKey, defaultModelListEndpoints[BackendDeepSeek]) +} + +// modelListEndpoint resolves a provider base URL to the exact URL used by the +// model-list request. APIURL follows Config.APIURL's base-URL convention; a URL +// already ending in /models is also accepted. Gemini deliberately retains its +// existing no-custom-endpoint contract. +func modelListEndpoint(backend Backend, apiURL string) (string, error) { + apiURL = strings.TrimSpace(apiURL) + if apiURL == "" { + endpoint, ok := defaultModelListEndpoints[backend] + if !ok { + return "", fmt.Errorf("backend %s has no live model listing", backend) + } + return endpoint, nil + } + if backend == BackendGemini { + return "", fmt.Errorf("backend %s does not support a custom model-list endpoint", backend) + } + u, err := url.Parse(apiURL) + if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" { + return "", fmt.Errorf("invalid model-list endpoint for %s", backend) + } + u.Fragment = "" + path := strings.TrimRight(u.Path, "/") + if !strings.HasSuffix(path, "/models") { + switch backend { + case BackendAnthropic: + if strings.HasSuffix(path, "/v1") { + path += "/models" + } else { + path += "/v1/models" + } + case BackendOpenAI, BackendDeepSeek: + path += "/models" + default: + return "", fmt.Errorf("backend %s has no live model listing", backend) + } + } + u.Path = path + return u.String(), nil +} + +func modelAPIURLEnvVars(backend Backend) []string { + switch backend { + case BackendAnthropic: + return []string{"ANTHROPIC_BASE_URL"} + case BackendOpenAI: + return []string{"OPENAI_BASE_URL"} + case BackendDeepSeek: + return []string{"DEEPSEEK_BASE_URL"} + default: + return nil + } +} + +func fetchModelsAtEndpoint(ctx context.Context, backend Backend, apiKey, endpoint string) ([]ModelDef, error) { if strings.TrimSpace(apiKey) == "" { - return nil, fmt.Errorf("DEEPSEEK_API_KEY is not set") + envVars := AuthEnvVars(backend) + if len(envVars) == 0 { + return nil, fmt.Errorf("backend %s has no live model listing", backend) + } + return nil, fmt.Errorf("%s is not set", envVars[0]) } - req, err := http.NewRequestWithContext(ctx, http.MethodGet, "https://api.deepseek.com/models", nil) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil) if err != nil { - return nil, err + return nil, fmt.Errorf("build %s models request: %w", backend, err) } - req.Header.Set("Authorization", "Bearer "+apiKey) - return doModelsRequest(req, BackendDeepSeek) + switch backend { + case BackendAnthropic: + req.Header.Set("x-api-key", apiKey) + req.Header.Set("anthropic-version", "2023-06-01") + case BackendGemini: + req.Header.Set("x-goog-api-key", apiKey) + case BackendOpenAI, BackendDeepSeek: + req.Header.Set("Authorization", "Bearer "+apiKey) + default: + return nil, fmt.Errorf("backend %s has no live model listing", backend) + } + return doModelsRequest(req, backend) } // doModelsRequest issues req with the default client and decodes the // permissive listing shape, projecting each entry into a ModelDef tagged -// with the supplied backend. Centralising this keeps the three fetchers +// with the supplied backend. Centralising this keeps the provider fetchers // behaviourally identical (same timeouts, same error messages, same name // normalisation). func doModelsRequest(req *http.Request, backend Backend) ([]ModelDef, error) { - resp, err := http.DefaultClient.Do(req) + client := *http.DefaultClient + client.CheckRedirect = func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + } + resp, err := client.Do(req) if err != nil { - return nil, fmt.Errorf("%s models: %w", backend, err) + // url.Error includes the full request URL, which may carry tenant or + // credential-bearing query data on custom endpoints. Preserve the cause + // for programmatic inspection but never render it to the user. + var urlErr *url.Error + if errors.As(err, &urlErr) { + err = urlErr.Err + } + return nil, ModelTransportError{Backend: backend, Err: err} } defer func() { _ = resp.Body.Close() }() if resp.StatusCode != http.StatusOK { @@ -212,15 +291,25 @@ func ListModels(ctx context.Context, backend Backend) ([]ModelDef, error) { // ListModelsWithAPIKey validates a candidate credential directly against the // provider model endpoint without reading or writing Captain's credential vault. func ListModelsWithAPIKey(ctx context.Context, backend Backend, apiKey string) ([]ModelDef, error) { - fetch := remoteFetcherFor(backend) - if fetch == nil { - return nil, fmt.Errorf("backend %s has no live model listing", backend) + return ListModelsWithAPIKeyAndURL(ctx, backend, apiKey, "") +} + +// ListModelsWithAPIKeyAndURL is ListModelsWithAPIKey with an optional provider +// base-URL override. The exact resolved request URL is shared with the persisted +// model-cache identity so endpoint-specific availability cannot cross caches. +func ListModelsWithAPIKeyAndURL(ctx context.Context, backend Backend, apiKey, apiURL string) ([]ModelDef, error) { + endpoint, err := modelListEndpoint(backend, apiURL) + if err != nil { + return nil, err } + return listModelsWithAPIKeyAtEndpoint(ctx, backend, apiKey, endpoint) +} +func listModelsWithAPIKeyAtEndpoint(ctx context.Context, backend Backend, apiKey, endpoint string) ([]ModelDef, error) { fetchCtx, cancel := context.WithTimeout(ctx, remoteModelsTimeout) defer cancel() - models, err := fetch(fetchCtx, apiKey) + models, err := fetchModelsAtEndpoint(fetchCtx, backend, apiKey, endpoint) if err != nil { return nil, err } @@ -228,21 +317,3 @@ func ListModelsWithAPIKey(ctx context.Context, backend Backend, apiKey string) ( sort.SliceStable(models, func(i, j int) bool { return models[i].ID < models[j].ID }) return models, nil } - -// remoteFetcherFor returns the live-list function for an API backend. It -// returns nil for any backend without a live listing endpoint -// (every CLI/agent backend, which lists from the static catalog instead). -func remoteFetcherFor(backend Backend) func(context.Context, string) ([]ModelDef, error) { - switch backend { - case BackendOpenAI: - return FetchOpenAIModels - case BackendAnthropic: - return FetchAnthropicModels - case BackendGemini: - return FetchGeminiModels - case BackendDeepSeek: - return FetchDeepSeekModels - default: - return nil - } -} diff --git a/pkg/ai/models_remote_test.go b/pkg/ai/models_remote_test.go index 44d97169..682f4fa4 100644 --- a/pkg/ai/models_remote_test.go +++ b/pkg/ai/models_remote_test.go @@ -3,6 +3,7 @@ package ai import ( "context" "encoding/json" + "errors" "io" "net/http" "net/http/httptest" @@ -29,6 +30,10 @@ type rewriteTransport struct { inner http.RoundTripper } +type roundTripFunc func(*http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { return f(req) } + func (r rewriteTransport) RoundTrip(req *http.Request) (*http.Response, error) { // Strip scheme+host while preserving request metadata for provider checks. target := r.base + req.URL.Path @@ -257,6 +262,78 @@ func TestListModels_SortsAlphabetically(t *testing.T) { } } +func TestListModelsWithAPIKeyAndURLUsesExactResolvedEndpointAndCredential(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if got := r.URL.RequestURI(); got != "/tenant/v1/models?account=private" { + t.Errorf("request URI = %q", got) + } + if got := r.Header.Get("Authorization"); got != "Bearer caller-token-b" { + t.Errorf("Authorization = %q", got) + } + _ = json.NewEncoder(w).Encode(map[string]any{ + "data": []map[string]any{{"id": "private-model-b"}}, + }) + })) + defer srv.Close() + + models, err := ListModelsWithAPIKeyAndURL( + context.Background(), BackendOpenAI, "caller-token-b", srv.URL+"/tenant/v1?account=private", + ) + if err != nil { + t.Fatalf("ListModelsWithAPIKeyAndURL: %v", err) + } + if len(models) != 1 || models[0].ID != "private-model-b" { + t.Fatalf("models = %+v", models) + } +} + +func TestListModelsWithAPIKeyAndURLRedactsEndpointFromTransportErrors(t *testing.T) { + original := http.DefaultClient.Transport + transportErr := errors.New("https://tenant.example/v1/models?access_token=query-secret: transport unavailable") + http.DefaultClient.Transport = roundTripFunc(func(*http.Request) (*http.Response, error) { + return nil, transportErr + }) + t.Cleanup(func() { http.DefaultClient.Transport = original }) + + _, err := ListModelsWithAPIKeyAndURL( + context.Background(), BackendOpenAI, "caller-token", "https://tenant.example/v1?access_token=query-secret", + ) + if err == nil { + t.Fatal("expected transport error") + } + if strings.Contains(err.Error(), "query-secret") || strings.Contains(err.Error(), "tenant.example") { + t.Fatalf("transport error exposed the custom endpoint: %v", err) + } + if !errors.Is(err, transportErr) { + t.Fatalf("transport error does not preserve its cause: %v", err) + } +} + +func TestListModelsWithAPIKeyAndURLDoesNotFollowRedirects(t *testing.T) { + redirected := make(chan *http.Request, 1) + target := httptest.NewServer(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { + redirected <- r.Clone(r.Context()) + })) + defer target.Close() + source := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, target.URL+"/capture", http.StatusTemporaryRedirect) + })) + defer source.Close() + + _, err := ListModelsWithAPIKeyAndURL( + context.Background(), BackendAnthropic, "caller-token", source.URL+"?access_token=query-secret", + ) + var httpErr ModelHTTPError + if !errors.As(err, &httpErr) || httpErr.StatusCode != http.StatusTemporaryRedirect { + t.Fatalf("redirect error = %v, want HTTP %d", err, http.StatusTemporaryRedirect) + } + select { + case request := <-redirected: + t.Fatalf("redirect target received credential-bearing request: %+v", request) + case <-time.After(50 * time.Millisecond): + } +} + func TestFetchGeminiModels_StripsModelsPrefix(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if got := r.Header.Get("x-goog-api-key"); got != "g-test" { diff --git a/pkg/database/session_chat_store.go b/pkg/database/session_chat_store.go index 0969032f..4ceba0c1 100644 --- a/pkg/database/session_chat_store.go +++ b/pkg/database/session_chat_store.go @@ -354,7 +354,8 @@ func (db *DB) DeleteChatSession(ctx context.Context, id uuid.UUID) error { func (db *DB) touchChatSession(ctx context.Context, id uuid.UUID) error { result := db.gorm.WithContext(ctx).Model(&sessionRecord{}).Where("id = ?", id). Updates(map[string]any{ - "last_activity_at": gorm.Expr("GREATEST(last_activity_at, ?)", time.Now().UTC()), + "state_version": gorm.Expr("state_version + 1"), + "last_activity_at": time.Now().UTC(), }) if result.Error != nil { return fmt.Errorf("touch Captain chat session: %w", result.Error) diff --git a/pkg/database/session_last_activity_integration_test.go b/pkg/database/session_last_activity_integration_test.go index b1a1be74..3800cfc4 100644 --- a/pkg/database/session_last_activity_integration_test.go +++ b/pkg/database/session_last_activity_integration_test.go @@ -2,7 +2,6 @@ package database import ( "database/sql" - "encoding/json" "testing" "time" @@ -171,50 +170,6 @@ func TestLastActivityAt_IngestIsMonotonic(t *testing.T) { }) } -func TestLastActivityAt_ChatMessageIsMonotonic(t *testing.T) { - db := openLastActivityTestDB(t) - session, err := db.CreateOrGetSession(t.Context(), CreateSessionInput{ - ID: uuid.New(), Source: "aichat", Provider: "captain", HostID: "local", - }) - require.NoError(t, err) - turn, created, err := db.CreateChatTurn(t.Context(), CreateChatTurnInput{ - SessionID: session.ID, ProviderTurnID: "turn-1", - }) - require.NoError(t, err) - require.True(t, created) - - newer := time.Now().UTC().Truncate(time.Second).Add(time.Hour) - pinLastActivity(t, db, session.ID, newer) - require.NoError(t, db.PutChatMessage(t.Context(), PutChatMessageInput{ - SessionID: session.ID, TurnID: turn.ID, ProviderMessageID: "message-1", - Role: "user", Parts: json.RawMessage(`[{"type":"text","text":"hello"}]`), - })) - - got := lastActivityAt(t, db, session.ID) - require.NotNil(t, got) - assert.True(t, got.Equal(newer), "chat message moved last_activity_at backwards: got %v, want %v", got, newer) -} - -func TestTouchChatSessionDoesNotAdvanceStateVersion(t *testing.T) { - db := openLastActivityTestDB(t) - session, err := db.CreateOrGetSession(t.Context(), CreateSessionInput{ - ID: uuid.New(), Source: "aichat", Provider: "captain", HostID: "local", - }) - require.NoError(t, err) - - require.NoError(t, db.touchChatSession(t.Context(), session.ID)) - touched, err := db.GetSession(t.Context(), session.ID) - require.NoError(t, err) - assert.Equal(t, session.StateVersion, touched.StateVersion) - - running := SessionLifecycleRunning - updated, err := db.UpdateSessionState(t.Context(), UpdateSessionStateInput{ - ID: session.ID, ExpectedVersion: session.StateVersion, LifecycleStatus: &running, - }) - require.NoError(t, err) - assert.Equal(t, session.StateVersion+1, updated.StateVersion) -} - func TestLastActivityAt_ChildActivityOnlyRewritesSessionWhenAdvancing(t *testing.T) { db := openLastActivityTestDB(t) now := time.Now().UTC().Truncate(time.Second) From 03b31a9b67dbcaa5f2b6cca53fccb3311233b76a Mon Sep 17 00:00:00 2001 From: Aditya Thebe Date: Tue, 11 Aug 2026 23:29:53 +0545 Subject: [PATCH 2/2] fix(ai): harden adapter and model cache refresh Adapter cache hits rehashed OAuth files while holding the global cache lock, and repeated credential rewrites discarded otherwise usable probe results. Validate hits with cheap file metadata, return unsettled snapshots uncached through a sentinel, and isolate host auth in catalog tests. Make model-cache lock acquisition context-aware so a contended entry cannot outlive the caller deadline. --- pkg/ai/adapters.go | 34 ++++++----- pkg/ai/adapters_cache.go | 71 ++++++++++++++++------ pkg/ai/adapters_cache_test.go | 83 ++++++++++++++++++++++++-- pkg/ai/availability.go | 3 +- pkg/ai/availability_test.go | 6 +- pkg/ai/catalog_disabled_ginkgo_test.go | 6 +- pkg/ai/catalog_resolve.go | 2 +- pkg/ai/catalog_resolve_test.go | 21 +++++++ pkg/ai/live_catalog.go | 5 +- pkg/ai/live_catalog_test.go | 6 +- pkg/ai/model_cache.go | 34 +++++++++-- pkg/cli/prompt_schema.go | 7 ++- 12 files changed, 224 insertions(+), 54 deletions(-) diff --git a/pkg/ai/adapters.go b/pkg/ai/adapters.go index 0c0c222f..a6acdf85 100644 --- a/pkg/ai/adapters.go +++ b/pkg/ai/adapters.go @@ -154,17 +154,20 @@ func (s CredentialSnapshot) clone() CredentialSnapshot { // so resolveAdapter stays pure and testable. Fields are exported so callers in // other packages (and their tests) can construct a hermetic probe. type AuthProbe struct { - Getenv func(string) string - LookPath func(string) (string, error) - FileExists func(string) bool - FileIdentity func(string) string - ExecutableIdentity func(string) string - CodexModels func(context.Context, string) ([]ModelDef, error) - APICredentials map[Backend]api.ResolvedAPIKey - APIURLs map[Backend]string - RuntimeStatuses map[Backend]RuntimeStatus - ProbeError error - Home string + Getenv func(string) string + LookPath func(string) (string, error) + FileExists func(string) bool + FileIdentity func(string) string + // FileMetadataIdentity is a cheap change detector for FileIdentity. When it + // is absent, cache validation falls back to FileIdentity. + FileMetadataIdentity func(string) string + ExecutableIdentity func(string) string + CodexModels func(context.Context, string) ([]ModelDef, error) + APICredentials map[Backend]api.ResolvedAPIKey + APIURLs map[Backend]string + RuntimeStatuses map[Backend]RuntimeStatus + ProbeError error + Home string credentials CredentialSnapshot stateFingerprint string @@ -181,9 +184,10 @@ func OSAuthProbe() AuthProbe { _, err := os.Stat(p) return err == nil }, - FileIdentity: hostFileIdentity, - ExecutableIdentity: hostExecutableIdentity, - Home: home, + FileIdentity: hostFileIdentity, + FileMetadataIdentity: hostMetadataIdentity, + ExecutableIdentity: hostMetadataIdentity, + Home: home, } probe.APICredentials = make(map[Backend]api.ResolvedAPIKey, len(apiBackends)) for _, backend := range apiBackends { @@ -381,7 +385,7 @@ func hostFileIdentity(path string) string { return fmt.Sprintf("sha256:%x", digest) } -func hostExecutableIdentity(path string) string { +func hostMetadataIdentity(path string) string { info, err := os.Stat(path) if err != nil { return "unreadable" diff --git a/pkg/ai/adapters_cache.go b/pkg/ai/adapters_cache.go index d6e83030..67249a1e 100644 --- a/pkg/ai/adapters_cache.go +++ b/pkg/ai/adapters_cache.go @@ -11,6 +11,11 @@ import ( // key/login/model changes surface without a probe per request. const adapterCacheTTL = 60 * time.Second +// ErrAdapterProbeUnsettled reports that authentication state changed throughout +// adapter discovery. The returned adapter snapshot remains usable but is not +// cached. +var ErrAdapterProbeUnsettled = errors.New("adapter probe did not settle") + // adapterAuthProbe captures the current host identity before a cache lookup; // adapterProbe resolves adapters from that same frozen snapshot. Both are // package vars so tests can substitute deterministic, network-free stubs. @@ -28,43 +33,71 @@ var ( // CachedAdapters returns the probed adapters, reusing a cached probe within the // TTL. A probe error is never cached: the next call retries so a transient -// failure does not permanently empty the catalog. `now` is a parameter so tests -// can advance time deterministically. +// failure does not permanently empty the catalog. If host state never settles, +// the freshest result is returned with ErrAdapterProbeUnsettled and is not +// cached. `now` is a parameter so tests can advance time deterministically. // // The cache stores the raw probe; the user's opt-out set is applied on the way // out. Baking it in would make a toggle wait out the TTL. func CachedAdapters(now time.Time) ([]AdapterStatus, error) { + rawProbe := adapterAuthProbe() + if rawProbe.ProbeError != nil { + return nil, rawProbe.ProbeError + } + hint := freezeAuthProbeHint(rawProbe) + adapterCacheMu.Lock() - defer adapterCacheMu.Unlock() - probe := freezeAuthProbe(adapterAuthProbe()) - if probe.ProbeError != nil { - return nil, probe.ProbeError + if adapterCache != nil && adapterCacheFingerprint == hint.stateFingerprint && now.Sub(adapterCacheAt) < adapterCacheTTL { + adapters := ApplyDisabled(adapterCache) + adapterCacheMu.Unlock() + return adapters, nil } - if adapterCache != nil && adapterCacheFingerprint == probe.stateFingerprint && now.Sub(adapterCacheAt) < adapterCacheTTL { + adapterCacheMu.Unlock() + + // Full credential-file hashing is only needed after the metadata hint says + // the cache may be stale. Do it outside the cache lock so unrelated readers + // do not queue behind disk I/O. + probe := freezeAuthProbe(rawProbe) + + adapterCacheMu.Lock() + defer adapterCacheMu.Unlock() + if adapterCache != nil && adapterCacheFingerprint == hint.stateFingerprint && now.Sub(adapterCacheAt) < adapterCacheTTL { return ApplyDisabled(adapterCache), nil } + + var latest []AdapterStatus for attempt := 0; attempt < 2; attempt++ { adapters, err := adapterProbe(probe) if err != nil { return nil, err } + latest = adapters + // Codex model discovery runs an external process whose account cannot be - // injected. Re-capture cheap host state before publishing its result; one - // retry closes a login/binary change that happened during the probe. - current := freezeAuthProbe(adapterAuthProbe()) - if current.ProbeError == nil && current.stateFingerprint == probe.stateFingerprint { + // injected. Re-capture host state before publishing its result; one retry + // closes a login or binary change that happened during the probe. + currentRaw := adapterAuthProbe() + if currentRaw.ProbeError != nil { + return nil, currentRaw.ProbeError + } + currentHint := freezeAuthProbeHint(currentRaw) + current := freezeAuthProbe(currentRaw) + if current.stateFingerprint == probe.stateFingerprint { adapterCache = cloneAdapterStatuses(adapters) adapterCacheAt = now - adapterCacheFingerprint = probe.stateFingerprint + adapterCacheFingerprint = currentHint.stateFingerprint return ApplyDisabled(adapterCache), nil } - if current.ProbeError != nil { - return nil, current.ProbeError - } - if attempt == 1 { - break - } probe = current } - return nil, errors.New("adapter probe did not settle") + return ApplyDisabled(latest), ErrAdapterProbeUnsettled +} + +// freezeAuthProbeHint captures the same state dimensions as freezeAuthProbe but +// uses file metadata in place of file contents when the host probe supports it. +func freezeAuthProbeHint(probe AuthProbe) AuthProbe { + if probe.FileMetadataIdentity != nil { + probe.FileIdentity = probe.FileMetadataIdentity + } + return freezeAuthProbe(probe) } diff --git a/pkg/ai/adapters_cache_test.go b/pkg/ai/adapters_cache_test.go index 8d5c9764..35517bcb 100644 --- a/pkg/ai/adapters_cache_test.go +++ b/pkg/ai/adapters_cache_test.go @@ -48,6 +48,69 @@ func TestCachedAdaptersReusesWithinTTL(t *testing.T) { } } +func TestCachedAdaptersUsesMetadataForCacheHits(t *testing.T) { + prevProbe, prevAuthProbe := adapterProbe, adapterAuthProbe + prevCache, prevAt, prevFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint + t.Cleanup(func() { + adapterProbe, adapterAuthProbe = prevProbe, prevAuthProbe + adapterCache, adapterCacheAt, adapterCacheFingerprint = prevCache, prevAt, prevFingerprint + }) + adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" + + home := t.TempDir() + authFile := filepath.Join(home, ".codex", "auth.json") + metadataIdentity := "mtime-a" + contentIdentity := "content-a" + hashReads := 0 + adapterAuthProbe = func() AuthProbe { + probe := fakeProbe(nil, nil, map[string]bool{authFile: true}, home) + probe.FileMetadataIdentity = func(path string) string { + if path == authFile { + return metadataIdentity + } + return "" + } + probe.FileIdentity = func(path string) string { + if path == authFile { + hashReads++ + return contentIdentity + } + return "" + } + return probe + } + calls := 0 + adapterProbe = func(probe AuthProbe) ([]AdapterStatus, error) { + calls++ + return []AdapterStatus{{Backend: string(BackendCodexCLI), Models: []string{probe.FileIdentity(authFile)}}}, nil + } + + base := time.Unix(1_500_000, 0) + got, err := CachedAdapters(base) + if err != nil || got[0].Models[0] != "content-a" { + t.Fatalf("initial adapters = %+v err=%v", got, err) + } + if hashReads != 2 { + t.Fatalf("initial credential hashes = %d, want 2", hashReads) + } + if _, err := CachedAdapters(base.Add(time.Second)); err != nil { + t.Fatal(err) + } + if hashReads != 2 || calls != 1 { + t.Fatalf("cache hit: credential hashes=%d adapter probes=%d, want 2 and 1", hashReads, calls) + } + + metadataIdentity = "mtime-b" + contentIdentity = "content-b" + got, err = CachedAdapters(base.Add(2 * time.Second)) + if err != nil || got[0].Models[0] != "content-b" { + t.Fatalf("changed credentials adapters = %+v err=%v", got, err) + } + if hashReads != 4 || calls != 2 { + t.Fatalf("invalidated cache: credential hashes=%d adapter probes=%d, want 4 and 2", hashReads, calls) + } +} + func TestCachedAdaptersDoesNotCacheErrors(t *testing.T) { prevProbe := adapterProbe prevAuthProbe := adapterAuthProbe @@ -228,7 +291,7 @@ func TestCachedAdaptersReturnsDeepCopies(t *testing.T) { } } -func TestCachedAdaptersRejectsUnsettledProbeState(t *testing.T) { +func TestCachedAdaptersReturnsFreshestUnsettledProbeState(t *testing.T) { prevProbe, prevAuthProbe := adapterProbe, adapterAuthProbe prevCache, prevAt, prevFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint t.Cleanup(func() { @@ -247,12 +310,22 @@ func TestCachedAdaptersRejectsUnsettledProbeState(t *testing.T) { } return probe } - adapterProbe = func(AuthProbe) ([]AdapterStatus, error) { - return []AdapterStatus{{Backend: string(BackendOpenAI)}}, nil + adapterProbe = func(probe AuthProbe) ([]AdapterStatus, error) { + return []AdapterStatus{{ + Backend: string(BackendOpenAI), + Models: []string{probe.credentials.APIKey(BackendOpenAI).Token}, + }}, nil } - if got, err := CachedAdapters(time.Unix(6_000_000, 0)); err == nil || got != nil { - t.Fatalf("unsettled probe returned adapters=%+v err=%v", got, err) + got, err := CachedAdapters(time.Unix(6_000_000, 0)) + if !errors.Is(err, ErrAdapterProbeUnsettled) { + t.Fatalf("unsettled probe error = %v, want ErrAdapterProbeUnsettled", err) + } + if len(got) != 1 || len(got[0].Models) != 1 || got[0].Models[0] != "c" { + t.Fatalf("unsettled probe adapters = %+v, want freshest observation", got) + } + if adapterCache != nil { + t.Fatalf("unsettled probe published cache: %+v", adapterCache) } } diff --git a/pkg/ai/availability.go b/pkg/ai/availability.go index 0f8524ed..119d80ac 100644 --- a/pkg/ai/availability.go +++ b/pkg/ai/availability.go @@ -1,6 +1,7 @@ package ai import ( + "errors" "fmt" "strings" "time" @@ -50,7 +51,7 @@ func AvailabilityForAdapter(status AdapterStatus) api.Availability { // LiveRuntimeCatalog annotates the registry runtime catalog with host readiness. func LiveRuntimeCatalog() ([]api.RuntimeFamily, error) { adapters, err := CachedAdapters(time.Now()) - if err != nil { + if err != nil && !errors.Is(err, ErrAdapterProbeUnsettled) { return nil, err } byBackend := make(map[string]AdapterStatus, len(adapters)) diff --git a/pkg/ai/availability_test.go b/pkg/ai/availability_test.go index f78a9f53..305234a0 100644 --- a/pkg/ai/availability_test.go +++ b/pkg/ai/availability_test.go @@ -56,13 +56,15 @@ var _ = Describe("runtime availability", func() { }) It("uses each adapter's live readiness in the runtime catalog", func() { - previousProbe := adapterProbe + previousProbe, previousAuthProbe := adapterProbe, adapterAuthProbe previousCache, previousAt, previousFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint DeferCleanup(func() { - adapterProbe = previousProbe + adapterProbe, adapterAuthProbe = previousProbe, previousAuthProbe adapterCache, adapterCacheAt, adapterCacheFingerprint = previousCache, previousAt, previousFingerprint }) adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" + home := GinkgoT().TempDir() + adapterAuthProbe = func() AuthProbe { return fakeProbe(nil, nil, nil, home) } adapterProbe = func(AuthProbe) ([]AdapterStatus, error) { return []AdapterStatus{ {Backend: string(BackendOpenAI), Type: "api"}, diff --git a/pkg/ai/catalog_disabled_ginkgo_test.go b/pkg/ai/catalog_disabled_ginkgo_test.go index 23f9c298..a030d64f 100644 --- a/pkg/ai/catalog_disabled_ginkgo_test.go +++ b/pkg/ai/catalog_disabled_ginkgo_test.go @@ -109,13 +109,15 @@ var _ = Describe("catalog opt-out filtering", func() { }) It("retains disabled models with remediation in the live descriptive menu", func() { - previousProbe := adapterProbe + previousProbe, previousAuthProbe := adapterProbe, adapterAuthProbe previousCache, previousAt, previousFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint DeferCleanup(func() { - adapterProbe = previousProbe + adapterProbe, adapterAuthProbe = previousProbe, previousAuthProbe adapterCache, adapterCacheAt, adapterCacheFingerprint = previousCache, previousAt, previousFingerprint }) adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" + home := GinkgoT().TempDir() + adapterAuthProbe = func() AuthProbe { return fakeProbe(nil, nil, nil, home) } adapterProbe = func(AuthProbe) ([]AdapterStatus, error) { return nil, nil } disable(nil, []string{"deepseek"}, nil, nil, nil) diff --git a/pkg/ai/catalog_resolve.go b/pkg/ai/catalog_resolve.go index 1c602cbc..34549f8d 100644 --- a/pkg/ai/catalog_resolve.go +++ b/pkg/ai/catalog_resolve.go @@ -85,7 +85,7 @@ func ResolveModels(ctx context.Context, opts ResolveOptions) ([]ResolvedModel, e return filterResolved(rows, opts.Filter), err } - unlock, err := lockModelCache(fp) + unlock, err := lockModelCache(ctx, fp) if err != nil { catalogLog.Warnf("model disk cache disabled: failed to lock cache entry: %v", err) rows, resolveErr := resolveFresh(ctx, opts, credentials) diff --git a/pkg/ai/catalog_resolve_test.go b/pkg/ai/catalog_resolve_test.go index 24a95e50..ac074ecc 100644 --- a/pkg/ai/catalog_resolve_test.go +++ b/pkg/ai/catalog_resolve_test.go @@ -9,6 +9,7 @@ import ( "strings" "sync" "testing" + "time" "github.com/flanksource/captain/pkg/api" "github.com/flanksource/captain/pkg/credentials" @@ -489,6 +490,26 @@ func TestResolveModelsSerializesSameCredentialCacheEntry(t *testing.T) { } } +func TestLockModelCacheHonorsContextWhileContended(t *testing.T) { + t.Setenv("HOME", t.TempDir()) + unlock, err := lockModelCache(context.Background(), "contended") + if err != nil { + t.Fatalf("first lock: %v", err) + } + defer unlock() + + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + unexpectedUnlock, err := lockModelCache(ctx, "contended") + if unexpectedUnlock != nil { + unexpectedUnlock() + t.Fatal("contended lock was acquired") + } + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("contended lock error = %v, want context deadline", err) + } +} + func TestResolveModelsCatalogOnlyCreatesNoCredentialVerifier(t *testing.T) { home := t.TempDir() t.Setenv("HOME", home) diff --git a/pkg/ai/live_catalog.go b/pkg/ai/live_catalog.go index 29a31b6b..436a9412 100644 --- a/pkg/ai/live_catalog.go +++ b/pkg/ai/live_catalog.go @@ -1,6 +1,7 @@ package ai import ( + "errors" "time" "github.com/flanksource/captain/pkg/api" @@ -15,7 +16,7 @@ import ( // overrides its static counterpart. func LiveCatalog() ([]Model, error) { adapters, err := CachedAdapters(time.Now()) - if err != nil { + if err != nil && !errors.Is(err, ErrAdapterProbeUnsettled) { return nil, err } return mergeLiveCatalog(Catalog(), adapters, liveCatalogOptions{}), nil @@ -27,7 +28,7 @@ func LiveCatalog() ([]Model, error) { // local backend binary is installed. func LiveCatalogInfo(configuredProviders []string) ([]ModelInfo, error) { adapters, err := CachedAdapters(time.Now()) - if err != nil { + if err != nil && !errors.Is(err, ErrAdapterProbeUnsettled) { return nil, err } models := mergeLiveCatalog(catalogSnapshot(), adapters, liveCatalogOptions{IncludeDisabled: true}) diff --git a/pkg/ai/live_catalog_test.go b/pkg/ai/live_catalog_test.go index 13fcd3db..ffac6d98 100644 --- a/pkg/ai/live_catalog_test.go +++ b/pkg/ai/live_catalog_test.go @@ -69,13 +69,15 @@ func TestMergeLiveCatalogUpsertsLiveAndPreservesStatic(t *testing.T) { } func TestLiveCatalogInfoAppliesPerProviderConfigured(t *testing.T) { - prevProbe := adapterProbe + prevProbe, prevAuthProbe := adapterProbe, adapterAuthProbe prevCache, prevAt, prevFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint t.Cleanup(func() { - adapterProbe = prevProbe + adapterProbe, adapterAuthProbe = prevProbe, prevAuthProbe adapterCache, adapterCacheAt, adapterCacheFingerprint = prevCache, prevAt, prevFingerprint }) adapterCache, adapterCacheAt, adapterCacheFingerprint = nil, time.Time{}, "" + home := t.TempDir() + adapterAuthProbe = func() AuthProbe { return fakeProbe(nil, nil, nil, home) } adapterProbe = func(AuthProbe) ([]AdapterStatus, error) { return []AdapterStatus{ {Backend: string(BackendAnthropic), Type: "api", ModelDetails: []ModelDef{ diff --git a/pkg/ai/model_cache.go b/pkg/ai/model_cache.go index e76ecd61..85797f6c 100644 --- a/pkg/ai/model_cache.go +++ b/pkg/ai/model_cache.go @@ -1,8 +1,10 @@ package ai import ( + "context" "crypto/rand" "encoding/json" + "errors" "fmt" "os" "path/filepath" @@ -84,7 +86,10 @@ func modelCachePath(fingerprint string) (string, error) { // lockModelCache serializes one cache identity across goroutines and Captain // processes. The lock spans cache re-check, live fetch, and publish so a slow // stale fetch cannot overwrite a newer refresh with a fresh timestamp. -func lockModelCache(fingerprint string) (func(), error) { +func lockModelCache(ctx context.Context, fingerprint string) (func(), error) { + if err := ctx.Err(); err != nil { + return nil, err + } dir, err := modelCacheDir() if err != nil { return nil, err @@ -97,9 +102,30 @@ func lockModelCache(fingerprint string) (func(), error) { _ = lock.Close() return nil, err } - if err := unix.Flock(int(lock.Fd()), unix.LOCK_EX); err != nil { - _ = lock.Close() - return nil, err + for { + if err := ctx.Err(); err != nil { + _ = lock.Close() + return nil, err + } + err := unix.Flock(int(lock.Fd()), unix.LOCK_EX|unix.LOCK_NB) + if err == nil { + if err := ctx.Err(); err != nil { + _ = unix.Flock(int(lock.Fd()), unix.LOCK_UN) + _ = lock.Close() + return nil, err + } + break + } + if !errors.Is(err, unix.EWOULDBLOCK) { + _ = lock.Close() + return nil, err + } + select { + case <-ctx.Done(): + _ = lock.Close() + return nil, ctx.Err() + case <-time.After(25 * time.Millisecond): + } } return func() { _ = unix.Flock(int(lock.Fd()), unix.LOCK_UN) diff --git a/pkg/cli/prompt_schema.go b/pkg/cli/prompt_schema.go index 57a5572e..f8772361 100644 --- a/pkg/cli/prompt_schema.go +++ b/pkg/cli/prompt_schema.go @@ -2,6 +2,7 @@ package cli import ( "encoding/json" + "errors" "fmt" "io" "sync" @@ -80,7 +81,11 @@ func PromptSchemaDocument() (map[string]any, error) { // schemaAdapters sources the probed adapters through pkg/ai's cache. It is a // package var so tests can substitute a deterministic, network-free stub. var schemaAdapters = func() ([]AdapterStatus, error) { - return ai.CachedAdapters(time.Now()) + adapters, err := ai.CachedAdapters(time.Now()) + if errors.Is(err, ai.ErrAdapterProbeUnsettled) { + return adapters, nil + } + return adapters, err } // reflectedSchemaBytes holds the JSON of the reflection-derived schemas. They are