From f0271fe0f1fe4fc8840a7c58c1e4957d6faa0d6d Mon Sep 17 00:00:00 2001 From: imneov Date: Wed, 8 Jul 2026 02:22:08 +0000 Subject: [PATCH 1/2] feat: add local add input sources --- cmd/add.go | 162 ++++++++++++- cmd/add_from_env_test.go | 98 ++++++++ cmd/add_test.go | 163 +++++++++++++ internal/fileparse/profile_core.go | 308 ++++++++++++++++++++++++ internal/fileparse/profile_core_test.go | 100 ++++++++ 5 files changed, 829 insertions(+), 2 deletions(-) create mode 100644 cmd/add_from_env_test.go create mode 100644 internal/fileparse/profile_core.go create mode 100644 internal/fileparse/profile_core_test.go diff --git a/cmd/add.go b/cmd/add.go index 1df1640..bb460d7 100644 --- a/cmd/add.go +++ b/cmd/add.go @@ -38,6 +38,8 @@ import ( "github.com/spf13/cobra" "github.com/a2d2-dev/claudecm/internal/config" + "github.com/a2d2-dev/claudecm/internal/envextract" + "github.com/a2d2-dev/claudecm/internal/fileparse" "github.com/a2d2-dev/claudecm/internal/presets" "github.com/a2d2-dev/claudecm/internal/storage" ) @@ -84,6 +86,8 @@ var ( addSmallFastModelFlag string addSetFlag []string addPresetFlag string + addFromEnvFlag bool + addFromFileFlag string addListPresetsFlag bool addDryRunFlag bool addOverwriteFlag bool @@ -183,6 +187,8 @@ func init() { addCmd.Flags().StringVar(&addModelFlag, "model", "", "Core model name") addCmd.Flags().StringVar(&addSmallFastModelFlag, "small-fast-model", "", "Core small/fast auxiliary model name") addCmd.Flags().StringVar(&addPresetFlag, "preset", "", "Built-in provider preset name (run --list-presets to discover)") + addCmd.Flags().BoolVar(&addFromEnvFlag, "from-env", false, "Build the profile draft from Claude Code / Codex environment variables") + addCmd.Flags().StringVar(&addFromFileFlag, "from-file", "", "Build the profile draft from a dotenv, shell, JSON, YAML, or TOML file") addCmd.Flags().BoolVar(&addListPresetsFlag, "list-presets", false, "List built-in provider presets and exit") addCmd.Flags().StringArrayVar(&addSetFlag, "set", nil, "Sparse overlay entry (repeatable). Format: tools..=. "+ @@ -219,16 +225,31 @@ func runAdd(cmd *cobra.Command, args []string) error { providerFlagSet := flagWasExplicit(cmd, "provider", addProviderFlag != addProviderDefault) baseURLFlagSet := flagWasExplicit(cmd, "base-url", addBaseURLFlag != "") + apiKeyFlagSet := flagWasExplicit(cmd, "api-key", addAPIKeyFlag != "") modelFlagSet := flagWasExplicit(cmd, "model", addModelFlag != "") + smallFastModelFlagSet := flagWasExplicit(cmd, "small-fast-model", addSmallFastModelFlag != "") preset, hasPreset, err := resolveAddPreset(addPresetFlag) if err != nil { return err } + if err := validateAddInputSources(hasPreset); err != nil { + return err + } + fromInputSource := addFromEnvFlag || strings.TrimSpace(addFromFileFlag) != "" + if fromInputSource { + providerFlagSet = flagWasExplicit(cmd, "provider", false) + baseURLFlagSet = flagWasExplicit(cmd, "base-url", false) + apiKeyFlagSet = flagWasExplicit(cmd, "api-key", false) + modelFlagSet = flagWasExplicit(cmd, "model", false) + smallFastModelFlagSet = flagWasExplicit(cmd, "small-fast-model", false) + } provider := addProviderFlag baseURL := addBaseURLFlag + apiKey := addAPIKeyFlag model := addModelFlag + smallFastModel := addSmallFastModelFlag var tools map[config.ToolID]config.ToolOverlay if hasPreset { provider = preset.ProviderKey @@ -236,15 +257,48 @@ func runAdd(cmd *cobra.Command, args []string) error { model = preset.Model tools = cloneToolMap(preset.Tools) } + if addFromEnvFlag { + core, envTools, err := profileDraftFromEnv() + if err != nil { + return err + } + if core.Provider != "" { + provider = core.Provider + } + baseURL = core.BaseURL + apiKey = core.APIKey + model = core.Model + smallFastModel = core.SmallFastModel + tools = mergeToolMaps(tools, envTools) + } + if strings.TrimSpace(addFromFileFlag) != "" { + core, err := fileparse.ParseProfileCoreFile(addFromFileFlag) + if err != nil { + return err + } + if core.Provider != "" { + provider = core.Provider + } + baseURL = core.BaseURL + apiKey = core.APIKey + model = core.Model + smallFastModel = core.SmallFastModel + } if providerFlagSet { provider = addProviderFlag } if baseURLFlagSet { baseURL = addBaseURLFlag } + if apiKeyFlagSet { + apiKey = addAPIKeyFlag + } if modelFlagSet { model = addModelFlag } + if smallFastModelFlagSet { + smallFastModel = addSmallFastModelFlag + } if hasPreset { applyExplicitPresetFlagOverrides(tools, preset.Name, provider, baseURL, model, providerFlagSet, baseURLFlagSet, modelFlagSet) } @@ -255,6 +309,9 @@ func runAdd(cmd *cobra.Command, args []string) error { if hasPreset && addAPIKeyFlag == "" { return fmt.Errorf("preset %q requires --api-key in non-interactive add", preset.Name) } + if (addFromEnvFlag || strings.TrimSpace(addFromFileFlag) != "") && strings.TrimSpace(apiKey) == "" { + return fmt.Errorf("no API key found in input source") + } // Build tools overlay from --set entries. Parsing is a pure // function so an invalid entry surfaces before any I/O. @@ -274,9 +331,9 @@ func runAdd(cmd *cobra.Command, args []string) error { Core: config.CoreConfig{ Provider: provider, BaseURL: baseURL, - APIKey: addAPIKeyFlag, + APIKey: apiKey, Model: model, - SmallFastModel: addSmallFastModelFlag, + SmallFastModel: smallFastModel, }, Tools: tools, } @@ -340,6 +397,24 @@ func resolveAddPreset(raw string) (presets.Preset, bool, error) { return p, true, nil } +func validateAddInputSources(hasPreset bool) error { + fromFileSet := strings.TrimSpace(addFromFileFlag) != "" + count := 0 + if hasPreset { + count++ + } + if addFromEnvFlag { + count++ + } + if fromFileSet { + count++ + } + if count > 1 { + return fmt.Errorf("choose only one add input source: --preset, --from-env, or --from-file") + } + return nil +} + func flagWasExplicit(cmd *cobra.Command, name string, fallback bool) bool { if cmd != nil && cmd.Flags() != nil { if f := cmd.Flags().Lookup(name); f != nil && f.Changed { @@ -349,6 +424,89 @@ func flagWasExplicit(cmd *cobra.Command, name string, fallback bool) bool { return fallback } +func profileDraftFromEnv() (config.CoreConfig, map[config.ToolID]config.ToolOverlay, error) { + var core config.CoreConfig + var tools map[config.ToolID]config.ToolOverlay + + core.Provider = addProviderDefault + if v := lookupNonEmptyEnv("ANTHROPIC_BASE_URL"); v != "" { + core.BaseURL = v + } + if v := lookupNonEmptyEnv("ANTHROPIC_AUTH_TOKEN"); v != "" { + core.APIKey = v + } + if v := lookupNonEmptyEnv("ANTHROPIC_API_KEY"); v != "" { + if core.APIKey == "" { + core.APIKey = v + } else { + tools = putClaudeCodeEnv(tools, "ANTHROPIC_API_KEY", v) + } + } + if v := lookupNonEmptyEnv("ANTHROPIC_MODEL"); v != "" { + core.Model = v + } + if v := lookupNonEmptyEnv("ANTHROPIC_SMALL_FAST_MODEL"); v != "" { + core.SmallFastModel = v + } + + codexKey := lookupNonEmptyEnv("OPENAI_API_KEY") + codexBaseURL := lookupNonEmptyEnv("OPENAI_BASE_URL") + codexModel := lookupNonEmptyEnv("CODEX_MODEL") + codexProvider := normalizeCodexProvider(lookupNonEmptyEnv("CODEX_MODEL_PROVIDER")) + if core.APIKey == "" && codexKey != "" { + core.APIKey = codexKey + } + if core.BaseURL == "" && codexBaseURL != "" { + core.BaseURL = codexBaseURL + } + if core.Model == "" && codexModel != "" { + core.Model = codexModel + } + if codexProvider != "" { + core.Provider = codexProvider + } + + if strings.TrimSpace(core.APIKey) == "" { + return config.CoreConfig{}, nil, fmt.Errorf("no API key found in environment") + } + return core, tools, nil +} + +func lookupNonEmptyEnv(name string) string { + v, ok := envextract.Lookup(name) + if !ok { + return "" + } + return strings.TrimSpace(v) +} + +func normalizeCodexProvider(provider string) string { + switch strings.TrimSpace(provider) { + case "": + return "" + case "openai": + return "openai-compat" + default: + return provider + } +} + +func putClaudeCodeEnv( + tools map[config.ToolID]config.ToolOverlay, + name, value string, +) map[config.ToolID]config.ToolOverlay { + if tools == nil { + tools = map[config.ToolID]config.ToolOverlay{} + } + ov := tools[config.ToolClaudeCode] + if ov.ExtraEnv == nil { + ov.ExtraEnv = map[string]string{} + } + ov.ExtraEnv[name] = value + tools[config.ToolClaudeCode] = ov + return tools +} + func applyExplicitPresetFlagOverrides( tools map[config.ToolID]config.ToolOverlay, presetName, provider, baseURL, model string, diff --git a/cmd/add_from_env_test.go b/cmd/add_from_env_test.go new file mode 100644 index 0000000..27544a7 --- /dev/null +++ b/cmd/add_from_env_test.go @@ -0,0 +1,98 @@ +//go:build test + +package cmd + +import ( + "os" + "path/filepath" + "strings" + "testing" + + "github.com/a2d2-dev/claudecm/internal/envextract" +) + +func TestAdd_FromEnvDryRunRedactsAndDoesNotWrite(t *testing.T) { + h := newAddHarness(t) + restore := envextract.SetLookupForTest(addEnvUniverse(map[string]string{ + "ANTHROPIC_BASE_URL": "https://env.example.com", + "ANTHROPIC_AUTH_TOKEN": "sk-env-token-1234", + "ANTHROPIC_MODEL": "claude-env-model", + "ANTHROPIC_SMALL_FAST_MODEL": "claude-env-small", + })) + t.Cleanup(restore) + addFromEnvFlag = true + addDryRunFlag = true + + stdout, _, err := runAddInner(t, "envprof") + if err != nil { + t.Fatalf("runAdd --from-env: %v", err) + } + for _, want := range []string{ + "base_url: https://env.example.com", + "model: claude-env-model", + "small_fast_model: claude-env-small", + } { + if !strings.Contains(stdout, want) { + t.Fatalf("dry-run missing %q:\n%s", want, stdout) + } + } + if strings.Contains(stdout, "sk-env-token-1234") { + t.Fatalf("dry-run leaked plaintext api key:\n%s", stdout) + } + if !strings.Contains(stdout, "sk-e***1234") { + t.Fatalf("dry-run missing redacted api key:\n%s", stdout) + } + if _, statErr := os.Stat(filepath.Join(h.home, ".claudecm", "profiles", "envprof.yaml")); !os.IsNotExist(statErr) { + t.Fatalf("profile file written despite --dry-run: %v", statErr) + } +} + +func TestAdd_FromEnvNoKeyRefusesWithoutWrite(t *testing.T) { + h := newAddHarness(t) + restore := envextract.SetLookupForTest(addEnvUniverse(map[string]string{ + "ANTHROPIC_BASE_URL": "https://env.example.com", + "ANTHROPIC_MODEL": "claude-env-model", + })) + t.Cleanup(restore) + addFromEnvFlag = true + + _, _, err := runAddInner(t, "nokey") + if err == nil { + t.Fatalf("--from-env without key accepted") + } + if !strings.Contains(err.Error(), "no API key found in environment") { + t.Fatalf("error = %v", err) + } + if _, statErr := os.Stat(filepath.Join(h.home, ".claudecm", "profiles", "nokey.yaml")); !os.IsNotExist(statErr) { + t.Fatalf("profile file written despite missing env key: %v", statErr) + } +} + +func TestAdd_FromEnvExplicitModelOverridesEnv(t *testing.T) { + h := newAddHarness(t) + restore := envextract.SetLookupForTest(addEnvUniverse(map[string]string{ + "ANTHROPIC_AUTH_TOKEN": "sk-env-override-1234", + "ANTHROPIC_MODEL": "env-model", + })) + t.Cleanup(restore) + addFromEnvFlag = true + addModelFlag = "flag-model" + + if _, _, err := runAddInner(t, "envoverride"); err != nil { + t.Fatalf("runAdd --from-env override: %v", err) + } + loaded, err := h.store.LoadProfile("envoverride") + if err != nil { + t.Fatalf("LoadProfile: %v", err) + } + if loaded.Core.Model != "flag-model" { + t.Fatalf("Model = %q, want flag-model", loaded.Core.Model) + } +} + +func addEnvUniverse(m map[string]string) func(string) (string, bool) { + return func(name string) (string, bool) { + v, ok := m[name] + return v, ok + } +} diff --git a/cmd/add_test.go b/cmd/add_test.go index 67d3851..2a420b5 100644 --- a/cmd/add_test.go +++ b/cmd/add_test.go @@ -45,6 +45,8 @@ func resetAddFlags() { addSmallFastModelFlag = "" addSetFlag = nil addPresetFlag = "" + addFromEnvFlag = false + addFromFileFlag = "" addListPresetsFlag = false addDryRunFlag = false addOverwriteFlag = false @@ -94,12 +96,36 @@ func runAddInner(t *testing.T, args ...string) (stdout, stderr string, err error t.Helper() var out, errBuf bytes.Buffer cmd := &cobra.Command{Use: "add"} + bindSyntheticAddFlags(cmd) cmd.SetOut(&out) cmd.SetErr(&errBuf) err = runAdd(cmd, args) return out.String(), errBuf.String(), err } +func bindSyntheticAddFlags(cmd *cobra.Command) { + cmd.Flags().String("provider", addProviderFlag, "") + if addProviderFlag != addProviderDefault { + _ = cmd.Flags().Set("provider", addProviderFlag) + } + cmd.Flags().String("base-url", addBaseURLFlag, "") + if addBaseURLFlag != "" { + _ = cmd.Flags().Set("base-url", addBaseURLFlag) + } + cmd.Flags().String("api-key", addAPIKeyFlag, "") + if addAPIKeyFlag != "" { + _ = cmd.Flags().Set("api-key", addAPIKeyFlag) + } + cmd.Flags().String("model", addModelFlag, "") + if addModelFlag != "" { + _ = cmd.Flags().Set("model", addModelFlag) + } + cmd.Flags().String("small-fast-model", addSmallFastModelFlag, "") + if addSmallFastModelFlag != "" { + _ = cmd.Flags().Set("small-fast-model", addSmallFastModelFlag) + } +} + // --------------------------------------------------------------------------- // Happy path // --------------------------------------------------------------------------- @@ -439,6 +465,143 @@ func TestAdd_PresetUnknownAndMissingSecretRefuseWithoutWrite(t *testing.T) { } } +func TestAdd_FromFileDotenvDryRunRedactsAndDoesNotWrite(t *testing.T) { + h := newAddHarness(t) + + path := filepath.Join(h.home, "provider.env") + if err := os.WriteFile(path, []byte(strings.Join([]string{ + "ANTHROPIC_BASE_URL=https://envfile.example.com", + "ANTHROPIC_AUTH_TOKEN=sk-file-dotenv-1234", + }, "\n")), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + addFromFileFlag = path + addDryRunFlag = true + + stdout, _, err := runAddInner(t, "work") + if err != nil { + t.Fatalf("runAdd --from-file dotenv: %v", err) + } + if !strings.Contains(stdout, "base_url: https://envfile.example.com") { + t.Fatalf("dry-run missing base_url:\n%s", stdout) + } + if strings.Contains(stdout, "sk-file-dotenv-1234") { + t.Fatalf("dry-run leaked plaintext api key:\n%s", stdout) + } + if !strings.Contains(stdout, "sk-f***1234") { + t.Fatalf("dry-run missing redacted api key:\n%s", stdout) + } + if _, statErr := os.Stat(filepath.Join(h.home, ".claudecm", "profiles", "work.yaml")); !os.IsNotExist(statErr) { + t.Fatalf("profile file written despite --dry-run: %v", statErr) + } +} + +func TestAdd_FromFileJSONPopulatesProfile(t *testing.T) { + h := newAddHarness(t) + + path := filepath.Join(h.home, "provider.json") + if err := os.WriteFile(path, []byte(`{"base_url":"https://json.example.com","api_key":"sk-json-file-1234","model":"json-model"}`), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + addFromFileFlag = path + + if _, _, err := runAddInner(t, "jsonfile"); err != nil { + t.Fatalf("runAdd --from-file json: %v", err) + } + loaded, err := h.store.LoadProfile("jsonfile") + if err != nil { + t.Fatalf("LoadProfile: %v", err) + } + if loaded.Core.BaseURL != "https://json.example.com" { + t.Fatalf("BaseURL = %q", loaded.Core.BaseURL) + } + if loaded.Core.APIKey != "sk-json-file-1234" { + t.Fatalf("APIKey = %q", loaded.Core.APIKey) + } + if loaded.Core.Model != "json-model" { + t.Fatalf("Model = %q", loaded.Core.Model) + } +} + +func TestAdd_FromFileShellExportModelParsed(t *testing.T) { + h := newAddHarness(t) + + path := filepath.Join(h.home, "exports.sh") + if err := os.WriteFile(path, []byte(strings.Join([]string{ + "export ANTHROPIC_AUTH_TOKEN=sk-shell-file-1234", + "export ANTHROPIC_MODEL=claude-shell-model", + }, "\n")), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + addFromFileFlag = path + + if _, _, err := runAddInner(t, "shellfile"); err != nil { + t.Fatalf("runAdd --from-file shell: %v", err) + } + loaded, err := h.store.LoadProfile("shellfile") + if err != nil { + t.Fatalf("LoadProfile: %v", err) + } + if loaded.Core.Model != "claude-shell-model" { + t.Fatalf("Model = %q", loaded.Core.Model) + } +} + +func TestAdd_FromFileExplicitModelOverridesParsedValue(t *testing.T) { + h := newAddHarness(t) + + path := filepath.Join(h.home, "provider.json") + if err := os.WriteFile(path, []byte(`{"api_key":"sk-file-override-1234","model":"file-model"}`), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + addFromFileFlag = path + addModelFlag = "flag-model" + + if _, _, err := runAddInner(t, "overridefile"); err != nil { + t.Fatalf("runAdd --from-file override: %v", err) + } + loaded, err := h.store.LoadProfile("overridefile") + if err != nil { + t.Fatalf("LoadProfile: %v", err) + } + if loaded.Core.Model != "flag-model" { + t.Fatalf("Model = %q, want flag-model", loaded.Core.Model) + } +} + +func TestAdd_FromFileUnreadableAndGarbageRefuseWithoutWrite(t *testing.T) { + h := newAddHarness(t) + + addFromFileFlag = filepath.Join(h.home, "missing.env") + _, _, err := runAddInner(t, "missing") + if err == nil { + t.Fatalf("nonexistent --from-file accepted") + } + if !strings.Contains(err.Error(), "cannot read file") { + t.Fatalf("missing file error = %v", err) + } + if _, statErr := os.Stat(filepath.Join(h.home, ".claudecm", "profiles", "missing.yaml")); !os.IsNotExist(statErr) { + t.Fatalf("profile file written after missing file: %v", statErr) + } + + resetAddFlags() + path := filepath.Join(h.home, "garbage.bin") + if err := os.WriteFile(path, []byte{0x00, 0xff, 0x01}, 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + addFromFileFlag = path + _, _, err = runAddInner(t, "garbage") + if err == nil { + t.Fatalf("garbage --from-file accepted") + } + if !strings.Contains(err.Error(), "unrecognized config format") { + t.Fatalf("garbage file error = %v", err) + } + if _, statErr := os.Stat(filepath.Join(h.home, ".claudecm", "profiles", "garbage.yaml")); !os.IsNotExist(statErr) { + t.Fatalf("profile file written after garbage file: %v", statErr) + } +} + // --------------------------------------------------------------------------- // Edge cases // --------------------------------------------------------------------------- diff --git a/internal/fileparse/profile_core.go b/internal/fileparse/profile_core.go new file mode 100644 index 0000000..16d20b5 --- /dev/null +++ b/internal/fileparse/profile_core.go @@ -0,0 +1,308 @@ +package fileparse + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "path/filepath" + "strconv" + "strings" + "unicode/utf8" + + toml "github.com/pelletier/go-toml/v2" + "gopkg.in/yaml.v3" + + "github.com/a2d2-dev/claudecm/internal/config" +) + +const maxProfileCoreFileBytes = 1 << 20 + +var errUnrecognizedFormat = errors.New("unrecognized config format") + +// ParseProfileCoreFile reads path once and parses known local config +// formats into profile core fields. It is read-only: all writes remain +// owned by cmd/add's normal SaveProfile path. +func ParseProfileCoreFile(path string) (config.CoreConfig, error) { + clean := strings.TrimSpace(path) + if clean == "" { + return config.CoreConfig{}, fmt.Errorf("cannot read file: path is empty") + } + body, err := os.ReadFile(clean) + if err != nil { + return config.CoreConfig{}, fmt.Errorf("cannot read file %q: %w", clean, err) + } + if len(body) > maxProfileCoreFileBytes { + return config.CoreConfig{}, fmt.Errorf("cannot read file %q: file is larger than %d bytes", clean, maxProfileCoreFileBytes) + } + core, err := ParseProfileCoreBytes(filepath.Ext(clean), body) + if err != nil { + if errors.Is(err, errUnrecognizedFormat) { + return config.CoreConfig{}, err + } + return config.CoreConfig{}, fmt.Errorf("%w: %s", errUnrecognizedFormat, err) + } + if strings.TrimSpace(core.APIKey) == "" { + return config.CoreConfig{}, fmt.Errorf("no API key found in file") + } + return core, nil +} + +// ParseProfileCoreBytes auto-detects dotenv, shell export, JSON, YAML, +// and TOML content and maps recognized vocabulary into CoreConfig. +func ParseProfileCoreBytes(ext string, body []byte) (config.CoreConfig, error) { + if !utf8.Valid(body) || bytes.IndexByte(body, 0) >= 0 { + return config.CoreConfig{}, errUnrecognizedFormat + } + trimmed := strings.TrimSpace(string(body)) + if trimmed == "" { + return config.CoreConfig{}, errUnrecognizedFormat + } + + parsers := orderedParsers(ext) + var lastErr error + for _, parser := range parsers { + values, err := parser(trimmed) + if err != nil { + lastErr = err + continue + } + core := coreFromValues(values) + if !coreHasAnyField(core) { + lastErr = fmt.Errorf("no usable profile fields found") + continue + } + return core, nil + } + if lastErr != nil { + return config.CoreConfig{}, fmt.Errorf("%w: %s", errUnrecognizedFormat, lastErr) + } + return config.CoreConfig{}, errUnrecognizedFormat +} + +type valuesParser func(string) (map[string]string, error) + +func orderedParsers(ext string) []valuesParser { + switch strings.ToLower(strings.TrimSpace(ext)) { + case ".json": + return []valuesParser{parseJSONValues, parseKeyValueLines, parseYAMLValues, parseTOMLValues} + case ".yaml", ".yml": + return []valuesParser{parseYAMLValues, parseKeyValueLines, parseJSONValues, parseTOMLValues} + case ".toml": + return []valuesParser{parseTOMLValues, parseKeyValueLines, parseJSONValues, parseYAMLValues} + case ".env": + return []valuesParser{parseKeyValueLines, parseJSONValues, parseYAMLValues, parseTOMLValues} + default: + return []valuesParser{parseKeyValueLines, parseJSONValues, parseYAMLValues, parseTOMLValues} + } +} + +func parseJSONValues(text string) (map[string]string, error) { + var root any + dec := json.NewDecoder(strings.NewReader(text)) + dec.UseNumber() + if err := dec.Decode(&root); err != nil { + return nil, err + } + if err := dec.Decode(&struct{}{}); !errors.Is(err, io.EOF) { + return nil, fmt.Errorf("json has trailing content") + } + out := map[string]string{} + flattenValues(out, "", root) + if len(out) == 0 { + return nil, fmt.Errorf("json object has no scalar values") + } + return out, nil +} + +func parseYAMLValues(text string) (map[string]string, error) { + var root any + if err := yaml.Unmarshal([]byte(text), &root); err != nil { + return nil, err + } + out := map[string]string{} + flattenValues(out, "", root) + if len(out) == 0 { + return nil, fmt.Errorf("yaml document has no scalar values") + } + return out, nil +} + +func parseTOMLValues(text string) (map[string]string, error) { + var root map[string]any + if err := toml.Unmarshal([]byte(text), &root); err != nil { + return nil, err + } + out := map[string]string{} + flattenValues(out, "", root) + if len(out) == 0 { + return nil, fmt.Errorf("toml document has no scalar values") + } + return out, nil +} + +func parseKeyValueLines(text string) (map[string]string, error) { + out := map[string]string{} + seenAssignment := false + for lineNo, raw := range strings.Split(text, "\n") { + line := strings.TrimSpace(raw) + if line == "" || strings.HasPrefix(line, "#") { + continue + } + if strings.HasPrefix(line, "export ") { + line = strings.TrimSpace(strings.TrimPrefix(line, "export ")) + } + idx := strings.Index(line, "=") + if idx < 0 { + return nil, fmt.Errorf("line %d is not key=value", lineNo+1) + } + key := strings.TrimSpace(line[:idx]) + if key == "" || strings.ContainsAny(key, " \t") { + return nil, fmt.Errorf("line %d has invalid key", lineNo+1) + } + value := strings.TrimSpace(line[idx+1:]) + unquoted, err := unquoteValue(value) + if err != nil { + return nil, fmt.Errorf("line %d has invalid quoted value: %w", lineNo+1, err) + } + out[key] = unquoted + seenAssignment = true + } + if !seenAssignment { + return nil, fmt.Errorf("no key=value assignments found") + } + return out, nil +} + +func unquoteValue(value string) (string, error) { + value = strings.TrimSpace(stripInlineComment(value)) + if len(value) < 2 { + return value, nil + } + if value[0] == '\'' && value[len(value)-1] == '\'' { + return value[1 : len(value)-1], nil + } + if value[0] == '"' && value[len(value)-1] == '"' { + return strconv.Unquote(value) + } + return value, nil +} + +func stripInlineComment(value string) string { + inSingle := false + inDouble := false + escaped := false + for i, r := range value { + switch { + case escaped: + escaped = false + case r == '\\' && inDouble: + escaped = true + case r == '\'' && !inDouble: + inSingle = !inSingle + case r == '"' && !inSingle: + inDouble = !inDouble + case r == '#' && !inSingle && !inDouble: + if i == 0 || value[i-1] == ' ' || value[i-1] == '\t' { + return strings.TrimSpace(value[:i]) + } + } + } + return value +} + +func flattenValues(out map[string]string, prefix string, value any) { + switch v := value.(type) { + case map[string]any: + for k, child := range v { + flattenValues(out, joinKey(prefix, k), child) + } + case map[any]any: + for k, child := range v { + flattenValues(out, joinKey(prefix, fmt.Sprint(k)), child) + } + case []any: + return + case nil: + return + case string: + out[prefix] = v + case bool: + out[prefix] = strconv.FormatBool(v) + case int: + out[prefix] = strconv.Itoa(v) + case int64: + out[prefix] = strconv.FormatInt(v, 10) + case float64: + out[prefix] = strconv.FormatFloat(v, 'g', -1, 64) + case json.Number: + out[prefix] = v.String() + default: + out[prefix] = fmt.Sprint(v) + } +} + +func joinKey(prefix, key string) string { + if prefix == "" { + return key + } + return prefix + "." + key +} + +func coreFromValues(values map[string]string) config.CoreConfig { + var core config.CoreConfig + for key, value := range values { + assignCoreField(&core, key, value) + } + if core.Provider == "openai" { + core.Provider = "openai-compat" + } + return core +} + +func assignCoreField(core *config.CoreConfig, key, value string) { + canonical := canonicalFieldKey(key) + switch canonical { + case "provider": + core.Provider = value + case "base_url": + core.BaseURL = value + case "api_key": + core.APIKey = value + case "model": + core.Model = value + case "small_fast_model": + core.SmallFastModel = value + } +} + +func canonicalFieldKey(key string) string { + normalized := strings.ToLower(strings.TrimSpace(key)) + normalized = strings.TrimPrefix(normalized, "env.") + switch normalized { + case "provider", "core.provider", "model_provider", "codex_model_provider": + return "provider" + case "base_url", "baseurl", "core.base_url", "anthropic_base_url", "openai_base_url": + return "base_url" + case "api_key", "apikey", "auth_token", "token", "core.api_key", + "anthropic_auth_token", "anthropic_api_key", "openai_api_key": + return "api_key" + case "model", "core.model", "anthropic_model", "codex_model": + return "model" + case "small_fast_model", "smallfastmodel", "core.small_fast_model", + "anthropic_small_fast_model": + return "small_fast_model" + default: + return "" + } +} + +func coreHasAnyField(core config.CoreConfig) bool { + return core.Provider != "" || + core.BaseURL != "" || + core.APIKey != "" || + core.Model != "" || + core.SmallFastModel != "" +} diff --git a/internal/fileparse/profile_core_test.go b/internal/fileparse/profile_core_test.go new file mode 100644 index 0000000..003a787 --- /dev/null +++ b/internal/fileparse/profile_core_test.go @@ -0,0 +1,100 @@ +package fileparse + +import ( + "strings" + "testing" +) + +func TestParseProfileCoreBytes_KnownFormats(t *testing.T) { + cases := []struct { + name string + ext string + body string + baseURL string + apiKey string + model string + provider string + smallFast string + }{ + { + name: "dotenv anthropic names", + ext: ".env", + body: "ANTHROPIC_BASE_URL=https://dotenv.example.com\nANTHROPIC_AUTH_TOKEN=sk-dotenv\nANTHROPIC_SMALL_FAST_MODEL=haiku\n", + baseURL: "https://dotenv.example.com", + apiKey: "sk-dotenv", + smallFast: "haiku", + }, + { + name: "shell export", + ext: ".sh", + body: "export ANTHROPIC_MODEL='claude-shell'\nexport ANTHROPIC_AUTH_TOKEN=sk-shell\n", + apiKey: "sk-shell", + model: "claude-shell", + }, + { + name: "json core names", + ext: ".json", + body: `{"base_url":"https://json.example.com","api_key":"sk-json","model":"json-model"}`, + baseURL: "https://json.example.com", + apiKey: "sk-json", + model: "json-model", + }, + { + name: "yaml provider", + ext: ".yaml", + body: "provider: openai\napi_key: sk-yaml\nmodel: yaml-model\n", + provider: "openai-compat", + apiKey: "sk-yaml", + model: "yaml-model", + }, + { + name: "toml codex names", + ext: ".toml", + body: "OPENAI_API_KEY = \"sk-toml\"\nOPENAI_BASE_URL = \"https://toml.example.com\"\nCODEX_MODEL = \"toml-model\"\n", + baseURL: "https://toml.example.com", + apiKey: "sk-toml", + model: "toml-model", + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := ParseProfileCoreBytes(tc.ext, []byte(tc.body)) + if err != nil { + t.Fatalf("ParseProfileCoreBytes: %v", err) + } + if got.BaseURL != tc.baseURL { + t.Fatalf("BaseURL = %q, want %q", got.BaseURL, tc.baseURL) + } + if got.APIKey != tc.apiKey { + t.Fatalf("APIKey = %q, want %q", got.APIKey, tc.apiKey) + } + if got.Model != tc.model { + t.Fatalf("Model = %q, want %q", got.Model, tc.model) + } + if got.Provider != tc.provider { + t.Fatalf("Provider = %q, want %q", got.Provider, tc.provider) + } + if got.SmallFastModel != tc.smallFast { + t.Fatalf("SmallFastModel = %q, want %q", got.SmallFastModel, tc.smallFast) + } + }) + } +} + +func TestParseProfileCoreBytes_UnrecognizedAndNoFields(t *testing.T) { + for _, body := range [][]byte{ + {0x00, 0xff, 0x01}, + []byte("this is not config"), + []byte("UNRELATED=value"), + } { + _, err := ParseProfileCoreBytes("", body) + if err == nil { + t.Fatalf("ParseProfileCoreBytes(%q): want error", string(body)) + } + if !strings.Contains(err.Error(), "unrecognized") && + !strings.Contains(err.Error(), "no usable profile fields") { + t.Fatalf("unexpected error: %v", err) + } + } +} From bebdc56ec90b375fd7f35973966786ce55fcf5f5 Mon Sep 17 00:00:00 2001 From: imneov Date: Wed, 8 Jul 2026 02:37:46 +0000 Subject: [PATCH 2/2] fix: harden local add input parsing --- cmd/add.go | 3 - cmd/add_from_env_test.go | 24 ++++- cmd/add_test.go | 43 +++++++++ internal/fileparse/profile_core.go | 120 +++++++++++++++++++++--- internal/fileparse/profile_core_test.go | 48 ++++++++++ 5 files changed, 219 insertions(+), 19 deletions(-) diff --git a/cmd/add.go b/cmd/add.go index bb460d7..61bbe20 100644 --- a/cmd/add.go +++ b/cmd/add.go @@ -466,9 +466,6 @@ func profileDraftFromEnv() (config.CoreConfig, map[config.ToolID]config.ToolOver core.Provider = codexProvider } - if strings.TrimSpace(core.APIKey) == "" { - return config.CoreConfig{}, nil, fmt.Errorf("no API key found in environment") - } return core, tools, nil } diff --git a/cmd/add_from_env_test.go b/cmd/add_from_env_test.go index 27544a7..58b3bc8 100644 --- a/cmd/add_from_env_test.go +++ b/cmd/add_from_env_test.go @@ -60,7 +60,7 @@ func TestAdd_FromEnvNoKeyRefusesWithoutWrite(t *testing.T) { if err == nil { t.Fatalf("--from-env without key accepted") } - if !strings.Contains(err.Error(), "no API key found in environment") { + if !strings.Contains(err.Error(), "no API key found in input source") { t.Fatalf("error = %v", err) } if _, statErr := os.Stat(filepath.Join(h.home, ".claudecm", "profiles", "nokey.yaml")); !os.IsNotExist(statErr) { @@ -68,6 +68,28 @@ func TestAdd_FromEnvNoKeyRefusesWithoutWrite(t *testing.T) { } } +func TestAdd_FromEnvNoKeyWithExplicitAPIKeyAllowed(t *testing.T) { + h := newAddHarness(t) + restore := envextract.SetLookupForTest(addEnvUniverse(map[string]string{ + "ANTHROPIC_BASE_URL": "https://env.example.com", + "ANTHROPIC_MODEL": "claude-env-model", + })) + t.Cleanup(restore) + addFromEnvFlag = true + addAPIKeyFlag = "sk-flag-env-1234" + + if _, _, err := runAddInner(t, "envflagkey"); err != nil { + t.Fatalf("runAdd --from-env --api-key: %v", err) + } + loaded, err := h.store.LoadProfile("envflagkey") + if err != nil { + t.Fatalf("LoadProfile: %v", err) + } + if loaded.Core.APIKey != "sk-flag-env-1234" { + t.Fatalf("APIKey = %q, want flag value", loaded.Core.APIKey) + } +} + func TestAdd_FromEnvExplicitModelOverridesEnv(t *testing.T) { h := newAddHarness(t) restore := envextract.SetLookupForTest(addEnvUniverse(map[string]string{ diff --git a/cmd/add_test.go b/cmd/add_test.go index 2a420b5..b81f625 100644 --- a/cmd/add_test.go +++ b/cmd/add_test.go @@ -569,6 +569,49 @@ func TestAdd_FromFileExplicitModelOverridesParsedValue(t *testing.T) { } } +func TestAdd_FromFileKeylessWithExplicitAPIKeyAllowed(t *testing.T) { + h := newAddHarness(t) + + path := filepath.Join(h.home, "provider.json") + if err := os.WriteFile(path, []byte(`{"base_url":"https://json.example.com","model":"json-model"}`), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + addFromFileFlag = path + addAPIKeyFlag = "sk-flag-file-1234" + + if _, _, err := runAddInner(t, "fileflagkey"); err != nil { + t.Fatalf("runAdd --from-file --api-key: %v", err) + } + loaded, err := h.store.LoadProfile("fileflagkey") + if err != nil { + t.Fatalf("LoadProfile: %v", err) + } + if loaded.Core.APIKey != "sk-flag-file-1234" { + t.Fatalf("APIKey = %q, want flag value", loaded.Core.APIKey) + } +} + +func TestAdd_FromFileKeylessWithoutExplicitAPIKeyRefuses(t *testing.T) { + h := newAddHarness(t) + + path := filepath.Join(h.home, "provider.json") + if err := os.WriteFile(path, []byte(`{"base_url":"https://json.example.com","model":"json-model"}`), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + addFromFileFlag = path + + _, _, err := runAddInner(t, "filewithoutkey") + if err == nil { + t.Fatalf("keyless --from-file accepted without --api-key") + } + if !strings.Contains(err.Error(), "no API key found in input source") { + t.Fatalf("error = %v", err) + } + if _, statErr := os.Stat(filepath.Join(h.home, ".claudecm", "profiles", "filewithoutkey.yaml")); !os.IsNotExist(statErr) { + t.Fatalf("profile file written despite missing file key: %v", statErr) + } +} + func TestAdd_FromFileUnreadableAndGarbageRefuseWithoutWrite(t *testing.T) { h := newAddHarness(t) diff --git a/internal/fileparse/profile_core.go b/internal/fileparse/profile_core.go index 16d20b5..5406487 100644 --- a/internal/fileparse/profile_core.go +++ b/internal/fileparse/profile_core.go @@ -8,6 +8,7 @@ import ( "io" "os" "path/filepath" + "sort" "strconv" "strings" "unicode/utf8" @@ -20,7 +21,10 @@ import ( const maxProfileCoreFileBytes = 1 << 20 -var errUnrecognizedFormat = errors.New("unrecognized config format") +var ( + errUnrecognizedFormat = errors.New("unrecognized config format") + errInvalidQuotedValue = errors.New("invalid quoted value") +) // ParseProfileCoreFile reads path once and parses known local config // formats into profile core fields. It is read-only: all writes remain @@ -44,9 +48,6 @@ func ParseProfileCoreFile(path string) (config.CoreConfig, error) { } return config.CoreConfig{}, fmt.Errorf("%w: %s", errUnrecognizedFormat, err) } - if strings.TrimSpace(core.APIKey) == "" { - return config.CoreConfig{}, fmt.Errorf("no API key found in file") - } return core, nil } @@ -66,10 +67,16 @@ func ParseProfileCoreBytes(ext string, body []byte) (config.CoreConfig, error) { for _, parser := range parsers { values, err := parser(trimmed) if err != nil { + if errors.Is(err, errInvalidQuotedValue) { + return config.CoreConfig{}, err + } lastErr = err continue } - core := coreFromValues(values) + core, err := coreFromValues(values) + if err != nil { + return config.CoreConfig{}, err + } if !coreHasAnyField(core) { lastErr = fmt.Errorf("no usable profile fields found") continue @@ -165,7 +172,7 @@ func parseKeyValueLines(text string) (map[string]string, error) { value := strings.TrimSpace(line[idx+1:]) unquoted, err := unquoteValue(value) if err != nil { - return nil, fmt.Errorf("line %d has invalid quoted value: %w", lineNo+1, err) + return nil, fmt.Errorf("line %d has %w: %v", lineNo+1, errInvalidQuotedValue, err) } out[key] = unquoted seenAssignment = true @@ -178,13 +185,19 @@ func parseKeyValueLines(text string) (map[string]string, error) { func unquoteValue(value string) (string, error) { value = strings.TrimSpace(stripInlineComment(value)) - if len(value) < 2 { + if value == "" { return value, nil } - if value[0] == '\'' && value[len(value)-1] == '\'' { + if value[0] == '\'' { + if len(value) < 2 || value[len(value)-1] != '\'' { + return "", fmt.Errorf("unclosed single quote") + } return value[1 : len(value)-1], nil } - if value[0] == '"' && value[len(value)-1] == '"' { + if value[0] == '"' { + if len(value) < 2 || value[len(value)-1] != '"' { + return "", fmt.Errorf("unclosed double quote") + } return strconv.Unquote(value) } return value, nil @@ -251,20 +264,74 @@ func joinKey(prefix, key string) string { return prefix + "." + key } -func coreFromValues(values map[string]string) config.CoreConfig { +func coreFromValues(values map[string]string) (config.CoreConfig, error) { var core config.CoreConfig - for key, value := range values { - assignCoreField(&core, key, value) + assignments := canonicalAssignments(values) + for _, canonical := range sortedStringKeys(assignments) { + selected, err := selectCanonicalValue(canonical, assignments[canonical]) + if err != nil { + return config.CoreConfig{}, err + } + assignCoreField(&core, canonical, selected.value) } if core.Provider == "openai" { core.Provider = "openai-compat" } - return core + return core, nil +} + +type fieldAssignment struct { + key string + value string +} + +func canonicalAssignments(values map[string]string) map[string][]fieldAssignment { + out := map[string][]fieldAssignment{} + for key, value := range values { + canonical := canonicalFieldKey(key) + if canonical == "" { + continue + } + out[canonical] = append(out[canonical], fieldAssignment{key: key, value: value}) + } + for canonical := range out { + sort.Slice(out[canonical], func(i, j int) bool { + return out[canonical][i].key < out[canonical][j].key + }) + } + return out +} + +func selectCanonicalValue(canonical string, assignments []fieldAssignment) (fieldAssignment, error) { + var selected fieldAssignment + for _, assignment := range assignments { + if strings.TrimSpace(assignment.value) == "" { + continue + } + if selected.key == "" { + selected = assignment + continue + } + if assignment.value != selected.value { + return fieldAssignment{}, fmt.Errorf("conflicting values for %s: %s=%s conflicts with %s=%s", + canonical, + selected.key, + displayParsedValue(canonical, selected.value), + assignment.key, + displayParsedValue(canonical, assignment.value)) + } + } + if selected.key != "" { + return selected, nil + } + if len(assignments) == 0 { + return fieldAssignment{}, nil + } + return assignments[0], nil } func assignCoreField(core *config.CoreConfig, key, value string) { - canonical := canonicalFieldKey(key) - switch canonical { + switch key { case "provider": core.Provider = value case "base_url": @@ -278,6 +345,29 @@ func assignCoreField(core *config.CoreConfig, key, value string) { } } +func displayParsedValue(canonical, value string) string { + if canonical == "api_key" { + return redactParsedSecret(value) + } + return strconv.Quote(value) +} + +func redactParsedSecret(value string) string { + if len(value) >= 8 { + return strconv.Quote(value[:4] + "***" + value[len(value)-4:]) + } + return strconv.Quote("***") +} + +func sortedStringKeys[T any](m map[string]T) []string { + keys := make([]string, 0, len(m)) + for key := range m { + keys = append(keys, key) + } + sort.Strings(keys) + return keys +} + func canonicalFieldKey(key string) string { normalized := strings.ToLower(strings.TrimSpace(key)) normalized = strings.TrimPrefix(normalized, "env.") diff --git a/internal/fileparse/profile_core_test.go b/internal/fileparse/profile_core_test.go index 003a787..c3423bd 100644 --- a/internal/fileparse/profile_core_test.go +++ b/internal/fileparse/profile_core_test.go @@ -98,3 +98,51 @@ func TestParseProfileCoreBytes_UnrecognizedAndNoFields(t *testing.T) { } } } + +func TestParseProfileCoreBytes_ConflictingAliasesRefuseWithRedactedSecret(t *testing.T) { + _, err := ParseProfileCoreBytes(".json", []byte(`{"api_key":"sk-good-secret-1234","core":{"api_key":"sk-bad-secret-5678"}}`)) + if err == nil { + t.Fatalf("ParseProfileCoreBytes accepted conflicting api_key aliases") + } + for _, want := range []string{ + "conflicting values for api_key", + "api_key", + "core.api_key", + "sk-g***1234", + "sk-b***5678", + } { + if !strings.Contains(err.Error(), want) { + t.Fatalf("error missing %q: %v", want, err) + } + } + for _, leaked := range []string{"sk-good-secret-1234", "sk-bad-secret-5678"} { + if strings.Contains(err.Error(), leaked) { + t.Fatalf("error leaked plaintext secret %q: %v", leaked, err) + } + } +} + +func TestParseProfileCoreBytes_SameAliasValueAllowed(t *testing.T) { + got, err := ParseProfileCoreBytes(".json", []byte(`{"api_key":"sk-same-secret-1234","core":{"api_key":"sk-same-secret-1234"}}`)) + if err != nil { + t.Fatalf("ParseProfileCoreBytes: %v", err) + } + if got.APIKey != "sk-same-secret-1234" { + t.Fatalf("APIKey = %q", got.APIKey) + } +} + +func TestParseProfileCoreBytes_UnclosedQuotedValueRefuses(t *testing.T) { + _, err := ParseProfileCoreBytes(".env", []byte(`ANTHROPIC_AUTH_TOKEN="sk-unclosed-1234`)) + if err == nil { + t.Fatalf("ParseProfileCoreBytes accepted unclosed quoted value") + } + for _, want := range []string{"invalid quoted value", "unclosed double quote"} { + if !strings.Contains(err.Error(), want) { + t.Fatalf("error missing %q: %v", want, err) + } + } + if strings.Contains(err.Error(), "sk-unclosed-1234") { + t.Fatalf("error leaked plaintext secret: %v", err) + } +}