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..a6acdf85 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,100 @@ 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 + // 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 } // OSAuthProbe wires AuthProbe to the real host environment. @@ -123,7 +184,10 @@ func OSAuthProbe() AuthProbe { _, err := os.Stat(p) return err == nil }, - Home: home, + FileIdentity: hostFileIdentity, + FileMetadataIdentity: hostMetadataIdentity, + ExecutableIdentity: hostMetadataIdentity, + Home: home, } probe.APICredentials = make(map[Backend]api.ResolvedAPIKey, len(apiBackends)) for _, backend := range apiBackends { @@ -134,9 +198,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 hostMetadataIdentity(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 +446,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 +478,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 +533,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 +546,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..67249a1e 100644 --- a/pkg/ai/adapters_cache.go +++ b/pkg/ai/adapters_cache.go @@ -1,6 +1,7 @@ package ai import ( + "errors" "sync" "time" ) @@ -10,36 +11,93 @@ 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()) +// 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. +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 // 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() + if adapterCache != nil && adapterCacheFingerprint == hint.stateFingerprint && now.Sub(adapterCacheAt) < adapterCacheTTL { + adapters := ApplyDisabled(adapterCache) + adapterCacheMu.Unlock() + return adapters, nil + } + 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 && now.Sub(adapterCacheAt) < adapterCacheTTL { + if adapterCache != nil && adapterCacheFingerprint == hint.stateFingerprint && now.Sub(adapterCacheAt) < adapterCacheTTL { return ApplyDisabled(adapterCache), nil } - adapters, err := adapterProbe() - if err != nil { - return nil, err + + 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 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 = currentHint.stateFingerprint + return ApplyDisabled(adapterCache), nil + } + probe = current + } + 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 } - adapterCache = adapters - adapterCacheAt = now - return ApplyDisabled(adapters), nil + return freezeAuthProbe(probe) } diff --git a/pkg/ai/adapters_cache_test.go b/pkg/ai/adapters_cache_test.go index 3acb9b5a..35517bcb 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 { @@ -40,17 +48,84 @@ 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 - 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 +150,210 @@ 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 TestCachedAdaptersReturnsFreshestUnsettledProbeState(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(probe AuthProbe) ([]AdapterStatus, error) { + return []AdapterStatus{{ + Backend: string(BackendOpenAI), + Models: []string{probe.credentials.APIKey(BackendOpenAI).Token}, + }}, nil + } + + 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) + } +} + +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.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 7380c093..305234a0 100644 --- a/pkg/ai/availability_test.go +++ b/pkg/ai/availability_test.go @@ -56,14 +56,16 @@ var _ = Describe("runtime availability", func() { }) It("uses each adapter's live readiness in the runtime catalog", func() { - previousProbe := adapterProbe - previousCache, previousAt := adapterCache, adapterCacheAt + previousProbe, previousAuthProbe := adapterProbe, adapterAuthProbe + previousCache, previousAt, previousFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint DeferCleanup(func() { - adapterProbe = previousProbe - adapterCache, adapterCacheAt = previousCache, previousAt + adapterProbe, adapterAuthProbe = previousProbe, previousAuthProbe + adapterCache, adapterCacheAt, adapterCacheFingerprint = previousCache, previousAt, previousFingerprint }) - adapterCache, adapterCacheAt = nil, time.Time{} - adapterProbe = func() ([]AdapterStatus, error) { + 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"}, {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..a030d64f 100644 --- a/pkg/ai/catalog_disabled_ginkgo_test.go +++ b/pkg/ai/catalog_disabled_ginkgo_test.go @@ -109,14 +109,16 @@ 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 + previousProbe, previousAuthProbe := adapterProbe, adapterAuthProbe + previousCache, previousAt, previousFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint DeferCleanup(func() { - adapterProbe = previousProbe - adapterCache, adapterCacheAt = previousCache, previousAt + adapterProbe, adapterAuthProbe = previousProbe, previousAuthProbe + adapterCache, adapterCacheAt, adapterCacheFingerprint = previousCache, previousAt, previousFingerprint }) - adapterCache, adapterCacheAt = nil, time.Time{} - adapterProbe = func() ([]AdapterStatus, error) { return nil, nil } + 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) infos, err := LiveCatalogInfo(nil) diff --git a/pkg/ai/catalog_resolve.go b/pkg/ai/catalog_resolve.go index 8b0cb7a1..34549f8d 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(ctx, 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..ac074ecc 100644 --- a/pkg/ai/catalog_resolve_test.go +++ b/pkg/ai/catalog_resolve_test.go @@ -3,14 +3,20 @@ package ai import ( "context" "errors" + "fmt" + "os" "path/filepath" "strings" + "sync" "testing" + "time" + "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 +27,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 +66,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 +240,320 @@ 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 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) + 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.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 377a631d..ffac6d98 100644 --- a/pkg/ai/live_catalog_test.go +++ b/pkg/ai/live_catalog_test.go @@ -69,14 +69,16 @@ func TestMergeLiveCatalogUpsertsLiveAndPreservesStatic(t *testing.T) { } func TestLiveCatalogInfoAppliesPerProviderConfigured(t *testing.T) { - prevProbe := adapterProbe - prevCache, prevAt := adapterCache, adapterCacheAt + prevProbe, prevAuthProbe := adapterProbe, adapterAuthProbe + prevCache, prevAt, prevFingerprint := adapterCache, adapterCacheAt, adapterCacheFingerprint t.Cleanup(func() { - adapterProbe = prevProbe - adapterCache, adapterCacheAt = prevCache, prevAt + adapterProbe, adapterAuthProbe = prevProbe, prevAuthProbe + adapterCache, adapterCacheAt, adapterCacheFingerprint = prevCache, prevAt, prevFingerprint }) - adapterCache, adapterCacheAt = nil, time.Time{} - adapterProbe = func() ([]AdapterStatus, error) { + 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{ {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..85797f6c 100644 --- a/pkg/ai/model_cache.go +++ b/pkg/ai/model_cache.go @@ -1,20 +1,26 @@ package ai import ( + "context" + "crypto/rand" "encoding/json" + "errors" + "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 +31,126 @@ 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 + } + // 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 modelCacheDir() (string, error) { + root, err := modelCacheRoot() + if err != nil { return "", err } - return filepath.Join(dir, "models.json"), nil + 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 } -func loadModelCache() (*ModelCache, error) { - path, err := modelCachePath() +// 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(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 } + 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 + } + 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) + _ = 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 +168,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 +180,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 + } + 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.Rename(tmp, path) + 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/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 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)