Skip to content

Commit f47133b

Browse files
PierrunoYTampagent
andcommitted
fix(provider): fail closed on ambiguous identities
Amp-Thread-ID: https://ampcode.com/threads/T-01a0246b-e61a-70e8-aa71-24f1ea7804c8 Co-authored-by: Amp <amp@ampcode.com>
1 parent 54b8cea commit f47133b

4 files changed

Lines changed: 31 additions & 36 deletions

File tree

‎internal/cli/auth_test.go‎

Lines changed: 9 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -851,12 +851,15 @@ func TestRunAuthRefreshRejectsEmptyCredentialCandidates(t *testing.T) {
851851
if err := os.WriteFile(configPath, []byte(`{"providers":[]}`), 0o600); err != nil {
852852
t.Fatal(err)
853853
}
854-
var stdout, stderr bytes.Buffer
855-
code := runWithDeps([]string{"auth", "refresh", " "}, &stdout, &stderr, appDeps{
856-
userConfigPath: func() (string, error) { return configPath, nil },
857-
})
858-
if code != exitCrash || !strings.Contains(stderr.String(), "no credential candidates") {
859-
t.Fatalf("exit = %d stderr = %q, want empty-candidate error", code, stderr.String())
854+
for _, extra := range [][]string{nil, {"--watch"}} {
855+
args := append([]string{"auth", "refresh", " "}, extra...)
856+
var stdout, stderr bytes.Buffer
857+
code := runWithDeps(args, &stdout, &stderr, appDeps{
858+
userConfigPath: func() (string, error) { return configPath, nil },
859+
})
860+
if code != exitCrash || stdout.Len() != 0 || !strings.Contains(stderr.String(), "no credential candidates") {
861+
t.Fatalf("args = %q exit = %d stdout = %q stderr = %q, want explicit empty-candidate app error", args, code, stdout.String(), stderr.String())
862+
}
860863
}
861864
}
862865

‎internal/cli/provider_onboarding_test.go‎

Lines changed: 9 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -105,6 +105,10 @@ func TestProviderRemoveRejectsAmbiguousFoldedName(t *testing.T) {
105105
{Name: "WORK", ProviderKind: config.ProviderKindOpenAICompatible, BaseURL: "https://upper.example/v1", Model: "m2", APIKeyStored: true},
106106
},
107107
})
108+
configBefore, err := os.ReadFile(configPath)
109+
if err != nil {
110+
t.Fatal(err)
111+
}
108112
store, err := config.ProviderKeyStoreAt(filepath.Dir(configPath))
109113
if err != nil {
110114
t.Fatal(err)
@@ -121,13 +125,13 @@ func TestProviderRemoveRejectsAmbiguousFoldedName(t *testing.T) {
121125
t.Fatalf("stderr = %q, want an ambiguous-provider-name error", stderr.String())
122126
}
123127

124-
cfg := readFileConfig(t, configPath)
125-
if len(cfg.Providers) != 2 || cfg.Providers[0].Name != "work" || cfg.Providers[1].Name != "WORK" {
126-
t.Fatalf("providers = %+v, want both rows untouched", cfg.Providers)
128+
configAfter, err := os.ReadFile(configPath)
129+
if err != nil || !bytes.Equal(configAfter, configBefore) {
130+
t.Fatalf("config changed after rejected removal: err=%v\nbefore=%s\nafter=%s", err, configBefore, configAfter)
127131
}
128132
key, ok, err := store.Get("work")
129133
if err != nil || !ok || key != "sk-lower" {
130-
t.Fatalf("stored key = %q ok=%v err=%v, want the credential preserved", key, ok, err)
134+
t.Fatalf("stored key = %q ok=%v err=%v, want the credential state preserved", key, ok, err)
131135
}
132136

133137
// The exact spelling still works, so a legacy config remains repairable.
@@ -136,7 +140,7 @@ func TestProviderRemoveRejectsAmbiguousFoldedName(t *testing.T) {
136140
if code := runWithDeps([]string{"providers", "remove", "WORK"}, &stdout, &stderr, providerSetupDeps(configPath)); code != exitSuccess {
137141
t.Fatalf("exact-name removal exit = %d stderr = %q", code, stderr.String())
138142
}
139-
cfg = readFileConfig(t, configPath)
143+
cfg := readFileConfig(t, configPath)
140144
if len(cfg.Providers) != 1 || cfg.Providers[0].Name != "work" {
141145
t.Fatalf("providers = %+v, want only the exact row removed", cfg.Providers)
142146
}

‎internal/config/writer.go‎

Lines changed: 9 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,6 @@ package config
22

33
import (
44
"encoding/json"
5-
"errors"
65
"fmt"
76
"os"
87
"path/filepath"
@@ -220,20 +219,19 @@ func PersistedProviderIdentity(path, identity string) (string, bool, error) {
220219
type PersistedIdentityMatch uint8
221220

222221
const (
223-
// PersistedIdentityNone means no row owns the identity, or a catalog id was
224-
// given that more than one row claims (an ambiguous request this package
225-
// refuses to guess at).
222+
// PersistedIdentityNone means no row owns the identity.
226223
PersistedIdentityNone PersistedIdentityMatch = iota
227224
// PersistedIdentityName means a row's own name matched, exactly or as a
228225
// case variant.
229226
PersistedIdentityName
230227
// PersistedIdentityCatalog means the identity matched only the catalog id of
231228
// exactly one row.
232229
PersistedIdentityCatalog
230+
// PersistedIdentityAmbiguous means multiple rows own the folded name or
231+
// catalog id. No caller may select a row or credential from this result.
232+
PersistedIdentityAmbiguous
233233
)
234234

235-
var errAmbiguousPersistedProviderIdentity = errors.New("ambiguous persisted provider identity")
236-
237235
// ResolvePersistedProviderIdentity finds the persisted user-config row that
238236
// owns identity and reports how it was addressed.
239237
//
@@ -292,11 +290,14 @@ func ResolvePersistedProviderIdentity(path, identity string) (ProviderProfile, P
292290
return *foldedName, PersistedIdentityName, nil
293291
}
294292
if foldedMatches > 1 {
295-
return ProviderProfile{}, PersistedIdentityNone, fmt.Errorf("%w: ambiguous provider name %q matches multiple persisted rows that differ only by case; use the exact spelling from config.json", errAmbiguousPersistedProviderIdentity, identity)
293+
return ProviderProfile{}, PersistedIdentityAmbiguous, fmt.Errorf("ambiguous provider name %q matches multiple persisted rows that differ only by case; use the exact spelling from config.json", identity)
296294
}
297295
if catalogMatches == 1 {
298296
return *catalogRow, PersistedIdentityCatalog, nil
299297
}
298+
if catalogMatches > 1 {
299+
return ProviderProfile{}, PersistedIdentityAmbiguous, fmt.Errorf("provider identity %q is ambiguous: %d saved profiles use it as a catalog id", identity, catalogMatches)
300+
}
300301
return ProviderProfile{}, PersistedIdentityNone, nil
301302
}
302303

@@ -353,25 +354,12 @@ func ProviderCredentialCandidates(path, addressedName string) (candidates []stri
353354
add(canonicalName)
354355
row, match, err := ResolvePersistedProviderIdentity(path, addressedName)
355356
if err != nil {
356-
if errors.Is(err, errAmbiguousPersistedProviderIdentity) {
357+
if match == PersistedIdentityAmbiguous {
357358
return nil, canonicalName, err
358359
}
359360
return candidates, canonicalName, err
360361
}
361362
if match == PersistedIdentityNone {
362-
providers, err := persistedProviders(path)
363-
if err != nil {
364-
return candidates, canonicalName, err
365-
}
366-
catalogMatches := 0
367-
for _, provider := range providers {
368-
if sameProviderIdentity(provider.CatalogID, addressedName) {
369-
catalogMatches++
370-
}
371-
}
372-
if catalogMatches > 1 {
373-
return nil, canonicalName, fmt.Errorf("provider identity %q is ambiguous: %d saved profiles use it as a catalog id", addressedName, catalogMatches)
374-
}
375363
return candidates, canonicalName, nil
376364
}
377365
canonicalName = strings.TrimSpace(row.Name)

‎internal/config/writer_test.go‎

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1414,8 +1414,8 @@ func TestResolvePersistedProviderIdentityPrefersNames(t *testing.T) {
14141414
writeConfigFixture(t, variants, FileConfig{
14151415
Providers: []ProviderProfile{{Name: "work"}, {Name: "WORK"}},
14161416
}, 0o600)
1417-
if _, _, err := ResolvePersistedProviderIdentity(variants, "wOrK"); err == nil {
1418-
t.Fatal("resolve succeeded for a name matching two case-variant rows, want an ambiguity error")
1417+
if _, match, err := ResolvePersistedProviderIdentity(variants, "wOrK"); err == nil || match != PersistedIdentityAmbiguous {
1418+
t.Fatalf("match = %v err = %v, want a distinct ambiguity result", match, err)
14191419
}
14201420
// An exact spelling still addresses one row, so a legacy config with such
14211421
// a pair stays repairable.
@@ -1450,8 +1450,8 @@ func TestResolvePersistedProviderIdentityPrefersNames(t *testing.T) {
14501450
{Name: "personal-xai", CatalogID: "xai"},
14511451
},
14521452
}, 0o600)
1453-
if _, match, err := ResolvePersistedProviderIdentity(shared, "xai"); err != nil || match != PersistedIdentityNone {
1454-
t.Fatalf("match = %v err = %v, want no guess at a shared catalog id", match, err)
1453+
if _, match, err := ResolvePersistedProviderIdentity(shared, "xai"); err == nil || match != PersistedIdentityAmbiguous {
1454+
t.Fatalf("match = %v err = %v, want distinct ambiguity for a shared catalog id", match, err)
14551455
}
14561456
})
14571457
}

0 commit comments

Comments
 (0)