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