| 1 | package cli |
| 2 | |
| 3 | import ( |
| 4 | "bufio" |
| 5 | "bytes" |
| 6 | "context" |
| 7 | "errors" |
| 8 | "fmt" |
| 9 | "io" |
| 10 | "os" |
| 11 | "sort" |
| 12 | "strings" |
| 13 | "time" |
| 14 | |
| 15 | "reasonix/internal/config" |
| 16 | "reasonix/internal/i18n" |
| 17 | "reasonix/internal/netclient" |
| 18 | ) |
| 19 | |
| 20 | type providerSetupSession struct { |
| 21 | cfg *config.Config |
| 22 | originalProviders map[string]config.ProviderEntry |
| 23 | originalDefault string |
| 24 | pendingCredentials map[string]string |
| 25 | // Credentials are owned by selected connections, never by their previous |
| 26 | // shared environment variable. The legacy map above is summary-only. |
| 27 | pendingConnectionCredentials map[string]string |
| 28 | removed map[string]bool |
| 29 | accessDeclared bool |
| 30 | projectScoped bool |
| 31 | declaredProviders []string |
| 32 | operations []providerSetupOperation |
| 33 | } |
| 34 | |
| 35 | const setupManagerContinue = 2 |
| 36 | |
| 37 | type providerSetupOperationKind uint8 |
| 38 | |
| 39 | const ( |
| 40 | setupOpProvider providerSetupOperationKind = iota |
| 41 | setupOpDefaultModel |
| 42 | setupOpLanguage |
| 43 | setupOpMaterializeAccess |
| 44 | setupOpAccessMembership |
| 45 | ) |
| 46 | |
| 47 | type providerSetupOperation struct { |
| 48 | kind providerSetupOperationKind |
| 49 | providerName string |
| 50 | beforeProvider *config.ProviderEntry |
| 51 | afterProvider *config.ProviderEntry |
| 52 | beforeString string |
| 53 | afterString string |
| 54 | accessName string |
| 55 | projectScoped bool |
| 56 | beforeBool bool |
| 57 | afterBool bool |
| 58 | } |
| 59 | |
| 60 | type providerSetupConflictError struct { |
| 61 | field string |
| 62 | } |
| 63 | |
| 64 | type providerSetupFileSnapshot struct { |
| 65 | exists bool |
| 66 | body []byte |
| 67 | } |
| 68 | |
| 69 | func (e *providerSetupConflictError) Error() string { |
| 70 | return e.field |
| 71 | } |
| 72 | |
| 73 | func providerSetupEntryPtr(entry config.ProviderEntry) *config.ProviderEntry { |
| 74 | copy := config.ProviderEntryConfigSnapshot(entry) |
| 75 | return © |
| 76 | } |
| 77 | |
| 78 | func readProviderSetupFileSnapshot(path string) (providerSetupFileSnapshot, error) { |
| 79 | body, err := os.ReadFile(path) |
| 80 | if err != nil { |
| 81 | if os.IsNotExist(err) { |
| 82 | return providerSetupFileSnapshot{}, nil |
| 83 | } |
| 84 | return providerSetupFileSnapshot{}, err |
| 85 | } |
| 86 | return providerSetupFileSnapshot{exists: true, body: body}, nil |
| 87 | } |
| 88 | |
| 89 | func providerSetupFileSnapshotEqual(a, b providerSetupFileSnapshot) bool { |
| 90 | return a.exists == b.exists && bytes.Equal(a.body, b.body) |
| 91 | } |
| 92 | |
| 93 | func newProviderSetupSession(cfg *config.Config) *providerSetupSession { |
| 94 | s := &providerSetupSession{ |
| 95 | cfg: cfg, |
| 96 | originalProviders: make(map[string]config.ProviderEntry, len(cfg.Providers)), |
| 97 | originalDefault: cfg.DefaultModel, |
| 98 | pendingCredentials: map[string]string{}, |
| 99 | pendingConnectionCredentials: map[string]string{}, |
| 100 | removed: map[string]bool{}, |
| 101 | } |
| 102 | for _, p := range cfg.Providers { |
| 103 | s.originalProviders[p.Name] = p |
| 104 | } |
| 105 | return s |
| 106 | } |
| 107 | |
| 108 | func newProviderSetupSessionForPath(cfg *config.Config, path string) *providerSetupSession { |
| 109 | s := newProviderSetupSession(cfg) |
| 110 | s.projectScoped = !config.IsUserConfigPath(path) |
| 111 | declarations, err := config.InspectConfigFileDeclarations(path) |
| 112 | if err != nil { |
| 113 | // LoadForEdit already reports malformed/unreadable config and falls back; |
| 114 | // keep the conservative policy here so setup never enables hidden siblings. |
| 115 | s.accessDeclared = true |
| 116 | return s |
| 117 | } |
| 118 | s.accessDeclared = declarations.DesktopProviderAccessDeclared |
| 119 | s.declaredProviders = declarations.ProviderNames |
| 120 | return s |
| 121 | } |
| 122 | |
| 123 | func (s *providerSetupSession) recordProviderMutation(name string, before, after *config.ProviderEntry) { |
| 124 | s.operations = append(s.operations, providerSetupOperation{ |
| 125 | kind: setupOpProvider, |
| 126 | providerName: name, |
| 127 | beforeProvider: before, |
| 128 | afterProvider: after, |
| 129 | }) |
| 130 | } |
| 131 | |
| 132 | func (s *providerSetupSession) setLanguage(language string) { |
| 133 | if s.cfg.Language == language { |
| 134 | return |
| 135 | } |
| 136 | s.operations = append(s.operations, providerSetupOperation{ |
| 137 | kind: setupOpLanguage, |
| 138 | beforeString: s.cfg.Language, |
| 139 | afterString: language, |
| 140 | }) |
| 141 | s.cfg.Language = language |
| 142 | } |
| 143 | |
| 144 | func (s *providerSetupSession) applyDeepSeekOfficialDefaultPricing() { |
| 145 | before := make(map[string]config.ProviderEntry, len(s.cfg.Providers)) |
| 146 | for _, provider := range s.cfg.Providers { |
| 147 | before[provider.Name] = provider |
| 148 | } |
| 149 | s.cfg.ApplyDeepSeekOfficialDefaultPricing() |
| 150 | for i := range s.cfg.Providers { |
| 151 | after := s.cfg.Providers[i] |
| 152 | previous, existed := before[after.Name] |
| 153 | if existed && config.ProviderEntriesConfigEqual(previous, after) { |
| 154 | delete(before, after.Name) |
| 155 | continue |
| 156 | } |
| 157 | var previousPtr *config.ProviderEntry |
| 158 | if existed { |
| 159 | previousPtr = providerSetupEntryPtr(previous) |
| 160 | } |
| 161 | s.recordProviderMutation(after.Name, previousPtr, providerSetupEntryPtr(after)) |
| 162 | delete(before, after.Name) |
| 163 | } |
| 164 | for name, previous := range before { |
| 165 | s.recordProviderMutation(name, providerSetupEntryPtr(previous), nil) |
| 166 | } |
| 167 | } |
| 168 | |
| 169 | func (s *providerSetupSession) resetProviderSummaryBaseline() { |
| 170 | s.originalProviders = make(map[string]config.ProviderEntry, len(s.cfg.Providers)) |
| 171 | for _, provider := range s.cfg.Providers { |
| 172 | s.originalProviders[provider.Name] = provider |
| 173 | } |
| 174 | } |
| 175 | |
| 176 | func (s *providerSetupSession) upsert(entries []config.ProviderEntry) error { |
| 177 | for _, entry := range entries { |
| 178 | var before *config.ProviderEntry |
| 179 | if current, ok := s.cfg.Provider(entry.Name); ok { |
| 180 | before = providerSetupEntryPtr(*current) |
| 181 | } |
| 182 | if err := s.cfg.UpsertProvider(entry); err != nil { |
| 183 | return err |
| 184 | } |
| 185 | current, _ := s.cfg.Provider(entry.Name) |
| 186 | if before == nil || !config.ProviderEntriesConfigEqual(*before, *current) { |
| 187 | s.recordProviderMutation(entry.Name, before, providerSetupEntryPtr(*current)) |
| 188 | } |
| 189 | delete(s.removed, entry.Name) |
| 190 | s.repairDanglingDefaultFor(*current) |
| 191 | } |
| 192 | return nil |
| 193 | } |
| 194 | |
| 195 | // repairDanglingDefaultFor re-points default_model at the provider's own default |
| 196 | // when an edit or model refresh dropped the exact model the ref named, mirroring |
| 197 | // the repair RemoveProvider performs on removal. |
| 198 | func (s *providerSetupSession) repairDanglingDefaultFor(p config.ProviderEntry) { |
| 199 | if !config.ModelRefsProvider(s.cfg.DefaultModel, p.Name) || len(p.ModelList()) == 0 { |
| 200 | return |
| 201 | } |
| 202 | if _, ok := s.cfg.ResolveModel(s.cfg.DefaultModel); ok { |
| 203 | return |
| 204 | } |
| 205 | if err := s.setDefaultModel(p.Name); err != nil { |
| 206 | fmt.Fprintln(os.Stderr, err) |
| 207 | } |
| 208 | } |
| 209 | |
| 210 | func (s *providerSetupSession) add(entries []config.ProviderEntry) error { |
| 211 | seen := make(map[string]bool, len(s.cfg.Providers)+len(entries)) |
| 212 | for _, provider := range s.cfg.Providers { |
| 213 | seen[provider.Name] = true |
| 214 | } |
| 215 | for _, entry := range entries { |
| 216 | if seen[entry.Name] { |
| 217 | return fmt.Errorf(i18n.M.SetupProviderExistsFmt, entry.Name) |
| 218 | } |
| 219 | seen[entry.Name] = true |
| 220 | } |
| 221 | return s.upsert(entries) |
| 222 | } |
| 223 | |
| 224 | func (s *providerSetupSession) remove(name string) error { |
| 225 | current, ok := s.cfg.Provider(name) |
| 226 | if !ok { |
| 227 | return fmt.Errorf("remove provider: no provider %q", name) |
| 228 | } |
| 229 | before := providerSetupEntryPtr(*current) |
| 230 | if err := s.cfg.RemoveProvider(name); err != nil { |
| 231 | return err |
| 232 | } |
| 233 | s.recordProviderMutation(name, before, nil) |
| 234 | s.removeProviderAccess(name) |
| 235 | if _, existed := s.originalProviders[name]; existed { |
| 236 | s.removed[name] = true |
| 237 | } |
| 238 | return nil |
| 239 | } |
| 240 | |
| 241 | func (s *providerSetupSession) addProviderAccess(entries []config.ProviderEntry) { |
| 242 | if len(entries) == 0 { |
| 243 | return |
| 244 | } |
| 245 | before := append([]string(nil), s.cfg.Desktop.ProviderAccess...) |
| 246 | // Preserve the legacy "undeclared means infer all configured providers" |
| 247 | // behavior before turning provider_access into an explicit list. Project |
| 248 | // setup only seeds providers declared by that project; cfg also contains |
| 249 | // built-in defaults, which must not override the user's global access policy. |
| 250 | if !s.accessDeclared && len(s.cfg.Desktop.ProviderAccess) == 0 { |
| 251 | if s.projectScoped { |
| 252 | for _, name := range s.declaredProviders { |
| 253 | provider, ok := s.cfg.Provider(name) |
| 254 | if ok && provider.Configured() && len(provider.ModelList()) > 0 { |
| 255 | s.cfg.Desktop.ProviderAccess = append(s.cfg.Desktop.ProviderAccess, name) |
| 256 | } |
| 257 | } |
| 258 | } else { |
| 259 | config.NormalizeLegacyDesktopProviderAccess(s.cfg) |
| 260 | } |
| 261 | s.accessDeclared = true |
| 262 | s.operations = append(s.operations, providerSetupOperation{ |
| 263 | kind: setupOpMaterializeAccess, |
| 264 | projectScoped: s.projectScoped, |
| 265 | }) |
| 266 | before = append([]string(nil), s.cfg.Desktop.ProviderAccess...) |
| 267 | } |
| 268 | seen := make(map[string]bool, len(s.cfg.Desktop.ProviderAccess)+len(entries)) |
| 269 | for _, name := range s.cfg.Desktop.ProviderAccess { |
| 270 | name = strings.TrimSpace(name) |
| 271 | if name != "" { |
| 272 | seen[name] = true |
| 273 | } |
| 274 | } |
| 275 | for _, entry := range entries { |
| 276 | name := strings.TrimSpace(entry.Name) |
| 277 | if name == "" || seen[name] { |
| 278 | continue |
| 279 | } |
| 280 | s.cfg.Desktop.ProviderAccess = append(s.cfg.Desktop.ProviderAccess, name) |
| 281 | seen[name] = true |
| 282 | } |
| 283 | s.accessDeclared = true |
| 284 | s.recordAccessTransition(before) |
| 285 | } |
| 286 | |
| 287 | func (s *providerSetupSession) removeProviderAccess(name string) { |
| 288 | name = strings.TrimSpace(name) |
| 289 | if name == "" || len(s.cfg.Desktop.ProviderAccess) == 0 { |
| 290 | return |
| 291 | } |
| 292 | before := append([]string(nil), s.cfg.Desktop.ProviderAccess...) |
| 293 | out := s.cfg.Desktop.ProviderAccess[:0] |
| 294 | for _, current := range s.cfg.Desktop.ProviderAccess { |
| 295 | if strings.TrimSpace(current) != name { |
| 296 | out = append(out, current) |
| 297 | } |
| 298 | } |
| 299 | s.cfg.Desktop.ProviderAccess = out |
| 300 | s.recordAccessTransition(before) |
| 301 | } |
| 302 | |
| 303 | func (s *providerSetupSession) recordAccessTransition(before []string) { |
| 304 | beforeSet := make(map[string]bool, len(before)) |
| 305 | afterSet := make(map[string]bool, len(s.cfg.Desktop.ProviderAccess)) |
| 306 | var order []string |
| 307 | seen := map[string]bool{} |
| 308 | for _, names := range [][]string{before, s.cfg.Desktop.ProviderAccess} { |
| 309 | for _, name := range names { |
| 310 | name = strings.TrimSpace(name) |
| 311 | if name == "" { |
| 312 | continue |
| 313 | } |
| 314 | if !seen[name] { |
| 315 | seen[name] = true |
| 316 | order = append(order, name) |
| 317 | } |
| 318 | } |
| 319 | } |
| 320 | for _, name := range before { |
| 321 | name = strings.TrimSpace(name) |
| 322 | if name != "" { |
| 323 | beforeSet[name] = true |
| 324 | } |
| 325 | } |
| 326 | for _, name := range s.cfg.Desktop.ProviderAccess { |
| 327 | name = strings.TrimSpace(name) |
| 328 | if name != "" { |
| 329 | afterSet[name] = true |
| 330 | } |
| 331 | } |
| 332 | for _, name := range order { |
| 333 | if beforeSet[name] == afterSet[name] { |
| 334 | continue |
| 335 | } |
| 336 | s.operations = append(s.operations, providerSetupOperation{ |
| 337 | kind: setupOpAccessMembership, |
| 338 | accessName: name, |
| 339 | beforeBool: beforeSet[name], |
| 340 | afterBool: afterSet[name], |
| 341 | }) |
| 342 | } |
| 343 | } |
| 344 | |
| 345 | func (s *providerSetupSession) setCredential(key, value string) error { |
| 346 | var names []string |
| 347 | for _, entry := range s.cfg.Providers { |
| 348 | if entry.APIKeyEnv == strings.TrimSpace(key) { |
| 349 | names = append(names, entry.Name) |
| 350 | } |
| 351 | } |
| 352 | return s.setCredentialForProviders(names, key, value) |
| 353 | } |
| 354 | |
| 355 | func (s *providerSetupSession) setCredentialForProviders(names []string, key, value string) error { |
| 356 | key = strings.TrimSpace(key) |
| 357 | if !config.IsValidCredentialKey(key) { |
| 358 | return fmt.Errorf("invalid API key variable name %q", key) |
| 359 | } |
| 360 | if strings.ContainsAny(value, "\r\n") { |
| 361 | return fmt.Errorf("API key for %s contains a newline", key) |
| 362 | } |
| 363 | s.pendingCredentials[key] = value |
| 364 | if s.pendingConnectionCredentials == nil { |
| 365 | s.pendingConnectionCredentials = map[string]string{} |
| 366 | } |
| 367 | for _, name := range names { |
| 368 | if name = strings.TrimSpace(name); name != "" { |
| 369 | entry, ok := s.cfg.Provider(name) |
| 370 | if !ok { |
| 371 | return fmt.Errorf("unknown provider %q", name) |
| 372 | } |
| 373 | // A credential-only edit still has a config precondition. Replay it |
| 374 | // at the same point as other edits so concurrent rotations conflict. |
| 375 | s.recordProviderMutation(name, providerSetupEntryPtr(*entry), providerSetupEntryPtr(*entry)) |
| 376 | s.pendingConnectionCredentials[name] = value |
| 377 | } |
| 378 | } |
| 379 | return nil |
| 380 | } |
| 381 | |
| 382 | func (s *providerSetupSession) setDefaultModel(model string) error { |
| 383 | before := s.cfg.DefaultModel |
| 384 | if err := s.cfg.SetDefaultModel(model); err != nil { |
| 385 | return err |
| 386 | } |
| 387 | if before != s.cfg.DefaultModel { |
| 388 | s.operations = append(s.operations, providerSetupOperation{ |
| 389 | kind: setupOpDefaultModel, |
| 390 | beforeString: before, |
| 391 | afterString: s.cfg.DefaultModel, |
| 392 | }) |
| 393 | } |
| 394 | return nil |
| 395 | } |
| 396 | |
| 397 | // providerUsable reports whether the provider would be selectable once this |
| 398 | // session saves: it lists models and either needs no key or has one resolvable |
| 399 | // from the credential store or staged in this session. |
| 400 | func (s *providerSetupSession) providerUsable(p *config.ProviderEntry) bool { |
| 401 | if p == nil || len(p.ModelList()) == 0 { |
| 402 | return false |
| 403 | } |
| 404 | return p.Configured() || s.pendingConnectionCredentials[p.Name] != "" |
| 405 | } |
| 406 | |
| 407 | // defaultModelUsable reports whether default_model resolves to a provider the |
| 408 | // user could actually run once this session saves. |
| 409 | func (s *providerSetupSession) defaultModelUsable() bool { |
| 410 | entry, ok := s.cfg.ResolveModel(s.cfg.DefaultModel) |
| 411 | return ok && s.providerUsable(entry) |
| 412 | } |
| 413 | |
| 414 | // promoteDefaultToNewProviders keeps the wizard's first-run contract: when the |
| 415 | // current default_model cannot run (unresolvable, or its key is neither stored |
| 416 | // nor staged), point it at the first usable provider the user just added, so a |
| 417 | // first run that only configures a custom provider boots on that provider |
| 418 | // instead of failing on the built-in default's missing key. A usable default is |
| 419 | // never hijacked. |
| 420 | func (s *providerSetupSession) promoteDefaultToNewProviders(entries []config.ProviderEntry) { |
| 421 | if s.defaultModelUsable() { |
| 422 | return |
| 423 | } |
| 424 | for _, entry := range entries { |
| 425 | current, ok := s.cfg.Provider(entry.Name) |
| 426 | if !ok || !s.providerUsable(current) { |
| 427 | continue |
| 428 | } |
| 429 | if err := s.setDefaultModel(current.Name); err == nil { |
| 430 | return |
| 431 | } |
| 432 | } |
| 433 | } |
| 434 | |
| 435 | func (s *providerSetupSession) credentialLines() []string { |
| 436 | keys := make([]string, 0, len(s.pendingCredentials)) |
| 437 | for key := range s.pendingCredentials { |
| 438 | keys = append(keys, key) |
| 439 | } |
| 440 | sort.Strings(keys) |
| 441 | lines := make([]string, 0, len(keys)) |
| 442 | for _, key := range keys { |
| 443 | lines = append(lines, key+"="+s.pendingCredentials[key]) |
| 444 | } |
| 445 | return lines |
| 446 | } |
| 447 | |
| 448 | func (s *providerSetupSession) summary() []string { |
| 449 | var added, edited []string |
| 450 | for _, p := range s.cfg.Providers { |
| 451 | old, existed := s.originalProviders[p.Name] |
| 452 | switch { |
| 453 | case !existed: |
| 454 | added = append(added, p.Name) |
| 455 | case !providerSetupEqual(old, p): |
| 456 | edited = append(edited, p.Name) |
| 457 | } |
| 458 | } |
| 459 | var out []string |
| 460 | if len(added) > 0 { |
| 461 | out = append(out, fmt.Sprintf(i18n.M.SetupSummaryAddedFmt, strings.Join(added, ", "))) |
| 462 | } |
| 463 | if len(edited) > 0 { |
| 464 | out = append(out, fmt.Sprintf(i18n.M.SetupSummaryEditedFmt, strings.Join(edited, ", "))) |
| 465 | } |
| 466 | if len(s.removed) > 0 { |
| 467 | names := make([]string, 0, len(s.removed)) |
| 468 | for name := range s.removed { |
| 469 | names = append(names, name) |
| 470 | } |
| 471 | sort.Strings(names) |
| 472 | out = append(out, fmt.Sprintf(i18n.M.SetupSummaryRemovedFmt, strings.Join(names, ", "))) |
| 473 | } |
| 474 | if s.cfg.DefaultModel != s.originalDefault { |
| 475 | out = append(out, fmt.Sprintf(i18n.M.SetupSummaryDefaultFmt, s.cfg.DefaultModel)) |
| 476 | } |
| 477 | if len(s.pendingCredentials) > 0 { |
| 478 | out = append(out, fmt.Sprintf(i18n.M.SetupSummaryKeysFmt, len(s.pendingCredentials))) |
| 479 | } |
| 480 | if len(out) == 0 { |
| 481 | out = append(out, i18n.M.SetupSummaryNoChanges) |
| 482 | } |
| 483 | return out |
| 484 | } |
| 485 | |
| 486 | func providerSetupEqual(a, b config.ProviderEntry) bool { |
| 487 | // Render-level equality is unnecessary here: the manager only changes these |
| 488 | // fields, while advanced provider fields are preserved by editing a copy. |
| 489 | return a.Name == b.Name && a.Kind == b.Kind && a.BaseURL == b.BaseURL && |
| 490 | a.Model == b.Model && strings.Join(a.Models, "\x00") == strings.Join(b.Models, "\x00") && |
| 491 | a.Default == b.Default && a.APIKeyEnv == b.APIKeyEnv |
| 492 | } |
| 493 | |
| 494 | func runProviderSetupManager(s *providerSetupSession, configPath, envPath string) int { |
| 495 | cfg := s.cfg |
| 496 | repaired, repairs := repairInvalidProviderKeyEnvs(cfg.Providers) |
| 497 | for i := range repaired { |
| 498 | if config.ProviderEntriesConfigEqual(cfg.Providers[i], repaired[i]) { |
| 499 | continue |
| 500 | } |
| 501 | before := providerSetupEntryPtr(cfg.Providers[i]) |
| 502 | cfg.Providers[i] = repaired[i] |
| 503 | s.recordProviderMutation(repaired[i].Name, before, providerSetupEntryPtr(repaired[i])) |
| 504 | } |
| 505 | for _, repair := range repairs { |
| 506 | fmt.Fprintf(os.Stderr, " %s\n", dim(fmt.Sprintf(i18n.M.RepairedAPIKeyEnvFmt, repair.provider, repair.old, repair.new))) |
| 507 | } |
| 508 | for { |
| 509 | items := providerManagerItems(s) |
| 510 | idx, err := selectOne(i18n.M.SetupManagerTitle, items) |
| 511 | if err != nil { |
| 512 | fmt.Fprintln(os.Stderr, "\n"+i18n.M.SetupCancelled) |
| 513 | return 1 |
| 514 | } |
| 515 | providerCount := len(cfg.Providers) |
| 516 | switch idx { |
| 517 | case providerCount: |
| 518 | if !addProviderToSession(s, false) { |
| 519 | continue |
| 520 | } |
| 521 | case providerCount + 1: |
| 522 | if !addProviderToSession(s, true) { |
| 523 | continue |
| 524 | } |
| 525 | case providerCount + 2: |
| 526 | rc := saveProviderSetupSession(s, configPath, envPath) |
| 527 | if rc == setupManagerContinue { |
| 528 | continue |
| 529 | } |
| 530 | return rc |
| 531 | case providerCount + 3: |
| 532 | fmt.Println(i18n.M.SetupCancelled) |
| 533 | return 1 |
| 534 | default: |
| 535 | manageProvider(s, idx) |
| 536 | } |
| 537 | } |
| 538 | } |
| 539 | |
| 540 | func providerManagerItems(s *providerSetupSession) []menuItem { |
| 541 | cfg := s.cfg |
| 542 | items := make([]menuItem, 0, len(cfg.Providers)+4) |
| 543 | for _, p := range cfg.Providers { |
| 544 | models := p.ModelList() |
| 545 | keyStatus := i18n.M.SetupKeyMissing |
| 546 | if p.APIKeyEnv == "" || config.CredentialIsSet(p.APIKeyEnv) || s.pendingConnectionCredentials[p.Name] != "" { |
| 547 | keyStatus = i18n.M.SetupKeySet |
| 548 | } |
| 549 | desc := fmt.Sprintf("%s · %d %s · %s", p.Kind, len(models), i18n.M.SetupModelsUnit, keyStatus) |
| 550 | if cfg.DefaultModel == p.Name || config.ModelRefsProvider(cfg.DefaultModel, p.Name) { |
| 551 | desc += " · " + i18n.M.SetupDefaultBadge |
| 552 | } |
| 553 | items = append(items, menuItem{name: p.Name, desc: desc}) |
| 554 | } |
| 555 | return append(items, |
| 556 | menuItem{name: i18n.M.SetupAddOpenAI, desc: i18n.M.CustomProviderDesc}, |
| 557 | menuItem{name: i18n.M.SetupAddAnthropic, desc: i18n.M.AnthropicProviderDesc}, |
| 558 | menuItem{name: i18n.M.SetupSaveExit, desc: i18n.M.SetupSaveExitDesc}, |
| 559 | menuItem{name: i18n.M.SetupCancel, desc: i18n.M.SetupCancelDesc}, |
| 560 | ) |
| 561 | } |
| 562 | |
| 563 | func addProviderToSession(s *providerSetupSession, anthropic bool) bool { |
| 564 | var result providerPromptResult |
| 565 | var err error |
| 566 | var proxy netclient.ProxySpec |
| 567 | if s != nil && s.cfg != nil { |
| 568 | proxy = s.cfg.NetworkProxySpec() |
| 569 | } |
| 570 | if anthropic { |
| 571 | result, err = promptAnthropicProvider(proxy) |
| 572 | } else { |
| 573 | result, err = promptCustomProvider(proxy) |
| 574 | } |
| 575 | if err != nil { |
| 576 | if !errors.Is(err, errCancelled) { |
| 577 | fmt.Fprintln(os.Stderr, err) |
| 578 | } |
| 579 | return false |
| 580 | } |
| 581 | for _, entry := range result.entries { |
| 582 | if !confirmSharedCredential(s.cfg, entry, "") { |
| 583 | return false |
| 584 | } |
| 585 | } |
| 586 | if err := s.add(result.entries); err != nil { |
| 587 | fmt.Fprintln(os.Stderr, err) |
| 588 | return false |
| 589 | } |
| 590 | s.addProviderAccess(result.entries) |
| 591 | for key, value := range result.credentials { |
| 592 | var names []string |
| 593 | for _, entry := range result.entries { |
| 594 | if entry.APIKeyEnv == key { |
| 595 | names = append(names, entry.Name) |
| 596 | } |
| 597 | } |
| 598 | if err := s.setCredentialForProviders(names, key, value); err != nil { |
| 599 | fmt.Fprintln(os.Stderr, err) |
| 600 | return false |
| 601 | } |
| 602 | } |
| 603 | // After the new keys are staged, so usability sees them. |
| 604 | s.promoteDefaultToNewProviders(result.entries) |
| 605 | return true |
| 606 | } |
| 607 | |
| 608 | func manageProvider(s *providerSetupSession, providerIndex int) { |
| 609 | if providerIndex < 0 || providerIndex >= len(s.cfg.Providers) { |
| 610 | return |
| 611 | } |
| 612 | p := s.cfg.Providers[providerIndex] |
| 613 | idx, err := selectOne(fmt.Sprintf(i18n.M.SetupProviderActionsFmt, p.Name), []menuItem{ |
| 614 | {name: i18n.M.SetupEditProvider}, |
| 615 | {name: i18n.M.SetupUpdateKey}, |
| 616 | {name: i18n.M.SetupTestRefresh}, |
| 617 | {name: i18n.M.SetupSetDefault}, |
| 618 | {name: i18n.M.SetupRemoveProvider}, |
| 619 | {name: i18n.M.SetupBack}, |
| 620 | }) |
| 621 | if err != nil || idx == 5 { |
| 622 | return |
| 623 | } |
| 624 | switch idx { |
| 625 | case 0: |
| 626 | editProvider(s, p) |
| 627 | case 1: |
| 628 | updateProviderKey(s, p) |
| 629 | case 2: |
| 630 | testAndRefreshProvider(s, p) |
| 631 | case 3: |
| 632 | setDefaultProvider(s, p) |
| 633 | case 4: |
| 634 | removeProviderFromSession(s, p) |
| 635 | } |
| 636 | } |
| 637 | |
| 638 | func editProvider(s *providerSetupSession, current config.ProviderEntry) { |
| 639 | in := bufio.NewScanner(os.Stdin) |
| 640 | edited := current |
| 641 | edited.BaseURL = ask(in, os.Stdout, i18n.M.CustomPromptBaseURL, current.BaseURL) |
| 642 | models := ask(in, os.Stdout, i18n.M.SetupPromptModels, strings.Join(current.ModelList(), ",")) |
| 643 | edited.Models = splitModels(models) |
| 644 | if len(edited.Models) == 1 { |
| 645 | edited.Model = edited.Models[0] |
| 646 | } else { |
| 647 | edited.Model = "" |
| 648 | } |
| 649 | if len(edited.Models) > 0 && !containsString(edited.Models, edited.Default) { |
| 650 | edited.Default = edited.Models[0] |
| 651 | } |
| 652 | edited.APIKeyEnv = promptOptionalAPIKeyEnvName(in, os.Stdout, i18n.M.CustomPromptKeyEnv, current.APIKeyEnv) |
| 653 | if !confirmSharedCredential(s.cfg, edited, current.Name) { |
| 654 | return |
| 655 | } |
| 656 | if err := s.upsert([]config.ProviderEntry{edited}); err != nil { |
| 657 | fmt.Fprintln(os.Stderr, err) |
| 658 | } |
| 659 | } |
| 660 | |
| 661 | func promptOptionalAPIKeyEnvName(in *bufio.Scanner, w io.Writer, label, def string) string { |
| 662 | for { |
| 663 | key := ask(in, w, label, def) |
| 664 | if key == "" || config.IsValidCredentialKey(key) { |
| 665 | return key |
| 666 | } |
| 667 | fmt.Fprintf(w, i18n.M.InvalidAPIKeyEnvFmt+"\n", key) |
| 668 | } |
| 669 | } |
| 670 | |
| 671 | func splitModels(raw string) []string { |
| 672 | seen := map[string]bool{} |
| 673 | var models []string |
| 674 | for model := range strings.SplitSeq(raw, ",") { |
| 675 | model = strings.TrimSpace(model) |
| 676 | if model != "" && !seen[model] { |
| 677 | seen[model] = true |
| 678 | models = append(models, model) |
| 679 | } |
| 680 | } |
| 681 | return models |
| 682 | } |
| 683 | |
| 684 | func confirmSharedCredential(cfg *config.Config, candidate config.ProviderEntry, ignoreName string) bool { |
| 685 | if candidate.APIKeyEnv == "" { |
| 686 | return true |
| 687 | } |
| 688 | for _, p := range cfg.Providers { |
| 689 | if p.Name == ignoreName || p.Name == candidate.Name || p.APIKeyEnv != candidate.APIKeyEnv || p.BaseURL == candidate.BaseURL { |
| 690 | continue |
| 691 | } |
| 692 | in := bufio.NewScanner(os.Stdin) |
| 693 | answer := ask(in, os.Stdout, fmt.Sprintf(i18n.M.SetupSharedKeyWarningFmt, candidate.APIKeyEnv, p.Name, p.BaseURL), "y/N") |
| 694 | return answer == "y" || answer == "Y" |
| 695 | } |
| 696 | return true |
| 697 | } |
| 698 | |
| 699 | func updateProviderKey(s *providerSetupSession, p config.ProviderEntry) { |
| 700 | in := bufio.NewScanner(os.Stdin) |
| 701 | keyEnvChanged := false |
| 702 | if p.APIKeyEnv == "" { |
| 703 | p.APIKeyEnv = promptAPIKeyEnvName(in, os.Stdout, i18n.M.CustomPromptKeyEnv, apiKeyEnvFromProviderName(p.Name)) |
| 704 | keyEnvChanged = true |
| 705 | } |
| 706 | if !confirmSharedCredential(s.cfg, p, p.Name) { |
| 707 | return |
| 708 | } |
| 709 | if keyEnvChanged { |
| 710 | if err := s.upsert([]config.ProviderEntry{p}); err != nil { |
| 711 | fmt.Fprintln(os.Stderr, err) |
| 712 | return |
| 713 | } |
| 714 | } |
| 715 | value := ask(in, os.Stdout, fmt.Sprintf(i18n.M.SetupPromptAPIKeyFmt, p.APIKeyEnv), "") |
| 716 | if value == "" { |
| 717 | return |
| 718 | } |
| 719 | if err := s.setCredentialForProviders([]string{p.Name}, p.APIKeyEnv, value); err != nil { |
| 720 | fmt.Fprintln(os.Stderr, err) |
| 721 | } |
| 722 | } |
| 723 | |
| 724 | func testAndRefreshProvider(s *providerSetupSession, p config.ProviderEntry) { |
| 725 | restore := temporarilySetCredential(p.APIKeyEnv, s.pendingConnectionCredentials[p.Name]) |
| 726 | defer restore() |
| 727 | p.ResolveAPIKeyFromProcessEnvForProbe() |
| 728 | ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) |
| 729 | defer cancel() |
| 730 | var proxy netclient.ProxySpec |
| 731 | if s.cfg != nil { |
| 732 | proxy = s.cfg.NetworkProxySpec() |
| 733 | } |
| 734 | models, err := p.FetchModelsWithProxy(ctx, proxy) |
| 735 | if err != nil { |
| 736 | fmt.Fprintf(os.Stderr, i18n.M.FetchModelsFailedFmt+"\n", p.Name, err) |
| 737 | return |
| 738 | } |
| 739 | if len(models) == 0 { |
| 740 | fmt.Fprintln(os.Stderr, i18n.M.CustomFetchEmpty) |
| 741 | return |
| 742 | } |
| 743 | items := make([]menuItem, len(models)) |
| 744 | for i, model := range models { |
| 745 | items[i] = menuItem{name: model} |
| 746 | } |
| 747 | idxs, err := selectMany(fmt.Sprintf(i18n.M.SelectModelsLabel, p.Name), items) |
| 748 | if err != nil || len(idxs) == 0 { |
| 749 | return |
| 750 | } |
| 751 | selected := make([]string, 0, len(idxs)) |
| 752 | for _, idx := range idxs { |
| 753 | selected = append(selected, models[idx]) |
| 754 | } |
| 755 | p.Models = selected |
| 756 | p.Model = "" |
| 757 | if !containsString(selected, p.Default) { |
| 758 | p.Default = selected[0] |
| 759 | } |
| 760 | if err := s.upsert([]config.ProviderEntry{p}); err != nil { |
| 761 | fmt.Fprintln(os.Stderr, err) |
| 762 | return |
| 763 | } |
| 764 | fmt.Printf(" %s\n", green(fmt.Sprintf(i18n.M.FetchModelsSuccessFmt, len(models), p.Name))) |
| 765 | } |
| 766 | |
| 767 | func temporarilySetCredential(key, value string) func() { |
| 768 | if key == "" || value == "" { |
| 769 | return func() {} |
| 770 | } |
| 771 | old, existed := os.LookupEnv(key) |
| 772 | _ = os.Setenv(key, value) |
| 773 | return func() { |
| 774 | if existed { |
| 775 | _ = os.Setenv(key, old) |
| 776 | } else { |
| 777 | _ = os.Unsetenv(key) |
| 778 | } |
| 779 | } |
| 780 | } |
| 781 | |
| 782 | func setDefaultProvider(s *providerSetupSession, p config.ProviderEntry) { |
| 783 | models := p.ModelList() |
| 784 | if len(models) == 0 { |
| 785 | return |
| 786 | } |
| 787 | items := make([]menuItem, len(models)) |
| 788 | for i, model := range models { |
| 789 | items[i] = menuItem{name: model} |
| 790 | } |
| 791 | idx, err := selectOne(i18n.M.SetupSelectDefaultModel, items) |
| 792 | if err != nil { |
| 793 | return |
| 794 | } |
| 795 | if err := s.setDefaultModel(p.Name + "/" + models[idx]); err != nil { |
| 796 | fmt.Fprintln(os.Stderr, err) |
| 797 | } |
| 798 | } |
| 799 | |
| 800 | func removeProviderFromSession(s *providerSetupSession, p config.ProviderEntry) { |
| 801 | in := bufio.NewScanner(os.Stdin) |
| 802 | answer := ask(in, os.Stdout, fmt.Sprintf(i18n.M.SetupConfirmRemoveFmt, p.Name), "y/N") |
| 803 | if answer != "y" && answer != "Y" { |
| 804 | return |
| 805 | } |
| 806 | if err := s.remove(p.Name); err != nil { |
| 807 | fmt.Fprintln(os.Stderr, err) |
| 808 | } |
| 809 | } |
| 810 | |
| 811 | func (s *providerSetupSession) replayOperations(cfg *config.Config, accessDeclared *bool, declaredProviders []string) error { |
| 812 | for _, operation := range s.operations { |
| 813 | switch operation.kind { |
| 814 | case setupOpProvider: |
| 815 | current, exists := cfg.Provider(operation.providerName) |
| 816 | if operation.beforeProvider == nil { |
| 817 | if exists { |
| 818 | return &providerSetupConflictError{field: fmt.Sprintf("provider %q", operation.providerName)} |
| 819 | } |
| 820 | } else if !exists || !config.ProviderEntriesConfigEqual(*current, *operation.beforeProvider) { |
| 821 | return &providerSetupConflictError{field: fmt.Sprintf("provider %q", operation.providerName)} |
| 822 | } |
| 823 | if operation.afterProvider == nil { |
| 824 | if err := cfg.RemoveProvider(operation.providerName); err != nil { |
| 825 | return fmt.Errorf("replay remove provider %q: %w", operation.providerName, err) |
| 826 | } |
| 827 | } else if err := cfg.UpsertProviderPreservingRuntime(*operation.afterProvider); err != nil { |
| 828 | return fmt.Errorf("replay provider %q: %w", operation.providerName, err) |
| 829 | } |
| 830 | case setupOpDefaultModel: |
| 831 | if cfg.DefaultModel != operation.beforeString { |
| 832 | return &providerSetupConflictError{field: "default_model"} |
| 833 | } |
| 834 | if err := cfg.SetDefaultModel(operation.afterString); err != nil { |
| 835 | return fmt.Errorf("replay default_model: %w", err) |
| 836 | } |
| 837 | case setupOpLanguage: |
| 838 | if cfg.Language != operation.beforeString { |
| 839 | return &providerSetupConflictError{field: "language"} |
| 840 | } |
| 841 | cfg.Language = operation.afterString |
| 842 | case setupOpMaterializeAccess: |
| 843 | if *accessDeclared { |
| 844 | return &providerSetupConflictError{field: "desktop.provider_access"} |
| 845 | } |
| 846 | cfg.Desktop.ProviderAccess = nil |
| 847 | if operation.projectScoped { |
| 848 | for _, name := range declaredProviders { |
| 849 | provider, ok := cfg.Provider(name) |
| 850 | if ok && provider.Configured() && len(provider.ModelList()) > 0 { |
| 851 | cfg.Desktop.ProviderAccess = append(cfg.Desktop.ProviderAccess, name) |
| 852 | } |
| 853 | } |
| 854 | } else { |
| 855 | config.NormalizeLegacyDesktopProviderAccess(cfg) |
| 856 | } |
| 857 | *accessDeclared = true |
| 858 | if cfg.Desktop.ProviderAccess == nil { |
| 859 | cfg.Desktop.ProviderAccess = []string{} |
| 860 | } |
| 861 | case setupOpAccessMembership: |
| 862 | current := providerSetupAccessContains(cfg.Desktop.ProviderAccess, operation.accessName) |
| 863 | if current != operation.beforeBool { |
| 864 | return &providerSetupConflictError{field: fmt.Sprintf("desktop.provider_access[%q]", operation.accessName)} |
| 865 | } |
| 866 | if operation.afterBool { |
| 867 | cfg.Desktop.ProviderAccess = append(cfg.Desktop.ProviderAccess, operation.accessName) |
| 868 | } else { |
| 869 | out := cfg.Desktop.ProviderAccess[:0] |
| 870 | for _, name := range cfg.Desktop.ProviderAccess { |
| 871 | if strings.TrimSpace(name) != operation.accessName { |
| 872 | out = append(out, name) |
| 873 | } |
| 874 | } |
| 875 | cfg.Desktop.ProviderAccess = out |
| 876 | } |
| 877 | default: |
| 878 | return fmt.Errorf("unknown provider setup operation %d", operation.kind) |
| 879 | } |
| 880 | } |
| 881 | return nil |
| 882 | } |
| 883 | |
| 884 | func providerSetupAccessContains(names []string, want string) bool { |
| 885 | want = strings.TrimSpace(want) |
| 886 | for _, name := range names { |
| 887 | if strings.TrimSpace(name) == want { |
| 888 | return true |
| 889 | } |
| 890 | } |
| 891 | return false |
| 892 | } |
| 893 | |
| 894 | func commitProviderSetupSession(s *providerSetupSession, configPath string) (bool, error) { |
| 895 | if len(s.operations) == 0 && len(s.pendingConnectionCredentials) == 0 { |
| 896 | return false, nil |
| 897 | } |
| 898 | unlock, err := config.LockConfigFileEdits(configPath) |
| 899 | if err != nil { |
| 900 | return false, err |
| 901 | } |
| 902 | defer unlock() |
| 903 | unlockCredentials, err := config.LockUserCredentialEdits() |
| 904 | if err != nil { |
| 905 | return false, err |
| 906 | } |
| 907 | defer unlockCredentials() |
| 908 | |
| 909 | before, err := readProviderSetupFileSnapshot(configPath) |
| 910 | if err != nil { |
| 911 | return false, err |
| 912 | } |
| 913 | declarations, err := config.InspectConfigFileDeclarations(configPath) |
| 914 | if err != nil { |
| 915 | return false, err |
| 916 | } |
| 917 | fresh, err := config.LoadForEditReadOnlyStrict(configPath) |
| 918 | if err != nil { |
| 919 | return false, err |
| 920 | } |
| 921 | if err := fresh.BeginModelCredentialCommitLocked(configPath, "cli-setup"); err != nil { |
| 922 | return false, err |
| 923 | } |
| 924 | defer fresh.CleanupStagedModelCredentialsLocked(configPath) |
| 925 | baseline := fresh.ModelSettingsBaseline() |
| 926 | accessDeclared := declarations.DesktopProviderAccessDeclared |
| 927 | if err := s.replayOperations(fresh, &accessDeclared, declarations.ProviderNames); err != nil { |
| 928 | return false, err |
| 929 | } |
| 930 | for name, value := range s.pendingConnectionCredentials { |
| 931 | slot, err := fresh.StageModelCredentialLocked(value) |
| 932 | if err != nil { |
| 933 | return false, err |
| 934 | } |
| 935 | { |
| 936 | entry, ok := fresh.Provider(name) |
| 937 | if !ok { |
| 938 | return false, &providerSetupConflictError{field: fmt.Sprintf("provider %q", name)} |
| 939 | } |
| 940 | updated := *entry |
| 941 | updated.APIKeyEnv = slot |
| 942 | if err := fresh.UpsertProvider(updated); err != nil { |
| 943 | return false, err |
| 944 | } |
| 945 | } |
| 946 | } |
| 947 | current, err := readProviderSetupFileSnapshot(configPath) |
| 948 | if err != nil { |
| 949 | return false, err |
| 950 | } |
| 951 | if !providerSetupFileSnapshotEqual(before, current) { |
| 952 | return false, &providerSetupConflictError{field: "configuration file"} |
| 953 | } |
| 954 | err = fresh.SaveModelSettingsTo(configPath, baseline) |
| 955 | if err != nil { |
| 956 | return false, err |
| 957 | } |
| 958 | if err := fresh.MarkModelCredentialConfigCommittedLocked(configPath); err != nil { |
| 959 | return false, err |
| 960 | } |
| 961 | if _, err := config.LoadForEditReadOnlyStrict(configPath); err != nil { |
| 962 | return false, err |
| 963 | } |
| 964 | if err := fresh.CompleteModelCredentialCommitLocked(); err != nil { |
| 965 | return false, err |
| 966 | } |
| 967 | return true, nil |
| 968 | } |
| 969 | |
| 970 | func saveProviderSetupSession(s *providerSetupSession, configPath, envPath string) int { |
| 971 | fmt.Println() |
| 972 | fmt.Println(i18n.M.SetupSummaryTitle) |
| 973 | for _, line := range s.summary() { |
| 974 | fmt.Println(" " + line) |
| 975 | } |
| 976 | in := bufio.NewScanner(os.Stdin) |
| 977 | answer := ask(in, os.Stdout, i18n.M.SetupConfirmSave, "Y/n") |
| 978 | if answer == "n" || answer == "N" { |
| 979 | return setupManagerContinue |
| 980 | } |
| 981 | configWritten, err := commitProviderSetupSession(s, configPath) |
| 982 | if err != nil { |
| 983 | var conflict *providerSetupConflictError |
| 984 | if errors.As(err, &conflict) { |
| 985 | fmt.Fprintf(os.Stderr, i18n.M.SetupConcurrentChangeFmt+"\n", conflict.field) |
| 986 | } else { |
| 987 | fmt.Fprintln(os.Stderr, i18n.M.WriteConfigErr, err) |
| 988 | } |
| 989 | return 1 |
| 990 | } |
| 991 | if configWritten { |
| 992 | fmt.Printf("\n%s %s\n", green("✓"), fmt.Sprintf(i18n.M.WroteFileFmt, displayPath(configPath))) |
| 993 | } |
| 994 | if len(s.pendingCredentials) > 0 { |
| 995 | target := config.CredentialsTargetDescription() |
| 996 | if target == "" { |
| 997 | target = envPath |
| 998 | } |
| 999 | fmt.Printf("%s %s\n", green("✓"), fmt.Sprintf(i18n.M.WroteFileFmt, displayPath(target))) |
| 1000 | } |
| 1001 | fmt.Printf("\n%s %s\n", accent("◆"), i18n.M.SetupComplete) |
| 1002 | return 0 |
| 1003 | } |
| 1004 |