From fd4c6a3128fa00fd0d154ed7d22aa9fac881672e Mon Sep 17 00:00:00 2001 From: Aditya kumar singh <143548997+Adityakk9031@users.noreply.github.com> Date: Thu, 30 Jul 2026 16:28:27 +0530 Subject: [PATCH] fix(config): add cross-process file locking and config.Modify to prevent OAuth token clobbering (#8) --- cmd/config.go | 81 ++++++++++++++----------------- cmd/root.go | 12 +++-- cmd/update.go | 16 +++---- internal/analytics/analytics.go | 14 ++++-- internal/api/client.go | 51 +++++++++++++------- internal/auth/oauth.go | 44 +++++++---------- internal/config/config.go | 85 ++++++++++++++++++++++++++++----- internal/config/config_test.go | 54 ++++++++++++++++++++- internal/config/lock_other.go | 17 +++++++ internal/config/lock_windows.go | 30 ++++++++++++ internal/onboard/runner.go | 31 +++++++----- 11 files changed, 303 insertions(+), 132 deletions(-) create mode 100644 internal/config/lock_other.go create mode 100644 internal/config/lock_windows.go diff --git a/cmd/config.go b/cmd/config.go index d48d023..a0450e5 100644 --- a/cmd/config.go +++ b/cmd/config.go @@ -53,53 +53,46 @@ func init() { func runConfigSet(_ *cobra.Command, args []string) error { key := strings.ToLower(strings.TrimSpace(args[0])) - val := args[1] - c, err := config.Load() - if err != nil { - return err - } - // `.` routes to per-provider credentials (datadog.api_key, - // sentry.auth_token, …) — the scriptable equivalent of `codag setup`. - if provider, field, ok := strings.Cut(key, "."); ok { - if err := setup.SetProviderField(c, provider, field, val); err != nil { - return err - } - if err := config.Save(c); err != nil { - return err - } - fmt.Fprintf(os.Stderr, "saved %s -> %s\n", key, config.ConfigPath()) - return nil - } - switch key { - case "server": - if err := validateServerURL(val); err != nil { - return err - } - c.Server = val - case "telemetry": - switch strings.ToLower(val) { - case "on", "true", "1", "yes": - c.TelemetryOptOut = false - case "off", "false", "0", "no": - c.TelemetryOptOut = true - default: - return fmt.Errorf("telemetry: expected on|off, got %q", val) + val := strings.TrimSpace(args[1]) + + err := config.Modify(func(c *config.Config) error { + // `.` routes to per-provider credentials (datadog.api_key, + // sentry.auth_token, …) — the scriptable equivalent of `codag setup`. + if provider, field, ok := strings.Cut(key, "."); ok { + return setup.SetProviderField(c, provider, field, val) } - case "updates", "update-check", "update_check": - switch strings.ToLower(val) { - case "on", "true", "1", "yes": - c.UpdateCheckOptOut = false - case "off", "false", "0", "no": - c.UpdateCheckOptOut = true + switch key { + case "server": + if err := validateServerURL(val); err != nil { + return err + } + c.Server = val + case "telemetry": + switch strings.ToLower(val) { + case "on", "true", "1", "yes": + c.TelemetryOptOut = false + case "off", "false", "0", "no": + c.TelemetryOptOut = true + default: + return fmt.Errorf("telemetry: expected on|off, got %q", val) + } + case "updates", "update-check", "update_check": + switch strings.ToLower(val) { + case "on", "true", "1", "yes": + c.UpdateCheckOptOut = false + case "off", "false", "0", "no": + c.UpdateCheckOptOut = true + default: + return fmt.Errorf("updates: expected on|off, got %q", val) + } + case "api-key", "api_key": + return fmt.Errorf("api-key is no longer supported; sign in with `codag auth login`") default: - return fmt.Errorf("updates: expected on|off, got %q", val) + return fmt.Errorf("unknown key %q (valid: server, telemetry, updates, or . such as datadog.api_key)", key) } - case "api-key", "api_key": - return fmt.Errorf("api-key is no longer supported; sign in with `codag auth login`") - default: - return fmt.Errorf("unknown key %q (valid: server, telemetry, updates, or . such as datadog.api_key)", key) - } - if err := config.Save(c); err != nil { + return nil + }) + if err != nil { return err } fmt.Fprintf(os.Stderr, "saved %s -> %s\n", key, config.ConfigPath()) diff --git a/cmd/root.go b/cmd/root.go index 4814d87..8b6c52d 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -149,15 +149,19 @@ func maybeReportUpdate(cmd *cobra.Command) { return } - cfg.LastUpdateCheckAt = now.Unix() - _ = config.Save(cfg) + _ = config.Modify(func(c *config.Config) error { + c.LastUpdateCheckAt = now.Unix() + return nil + }) info, err := updatecheck.New(autoUpdateTimeout).Check(Version) if err != nil { return } - cfg.LastUpdateLatest = info.Latest - _ = config.Save(cfg) + _ = config.Modify(func(c *config.Config) error { + c.LastUpdateLatest = info.Latest + return nil + }) if !info.Available { return } diff --git a/cmd/update.go b/cmd/update.go index d3aab25..d795b31 100644 --- a/cmd/update.go +++ b/cmd/update.go @@ -61,13 +61,11 @@ func runUpdate(cmd *cobra.Command, _ []string) error { } func rememberUpdateCheck(info *updatecheck.Info) { - cfg, err := config.Load() - if err != nil || cfg == nil { - return - } - cfg.LastUpdateCheckAt = time.Now().Unix() - if info != nil { - cfg.LastUpdateLatest = info.Latest - } - _ = config.Save(cfg) + _ = config.Modify(func(cfg *config.Config) error { + cfg.LastUpdateCheckAt = time.Now().Unix() + if info != nil { + cfg.LastUpdateLatest = info.Latest + } + return nil + }) } diff --git a/internal/analytics/analytics.go b/internal/analytics/analytics.go index 5c9399c..8eccddb 100644 --- a/internal/analytics/analytics.go +++ b/internal/analytics/analytics.go @@ -50,10 +50,18 @@ func Init(cliVersion string) { return } if cfg.CLIDistinctID == "" { - cfg.CLIDistinctID = newID() - _ = config.Save(cfg) // best-effort; if we can't persist, fall through + newIDStr := newID() + _ = config.Modify(func(c *config.Config) error { + if c.CLIDistinctID == "" { + c.CLIDistinctID = newIDStr + } + return nil + }) + cfg, _ = config.Load() + } + if cfg != nil { + distinctID = cfg.CLIDistinctID } - distinctID = cfg.CLIDistinctID c, err := posthog.NewWithConfig(projectKey, posthog.Config{ Endpoint: apiHost, diff --git a/internal/api/client.go b/internal/api/client.go index a4604c2..4708143 100644 --- a/internal/api/client.go +++ b/internal/api/client.go @@ -14,6 +14,7 @@ import ( "github.com/codag-megalith/codag-cli/internal/auth" "github.com/codag-megalith/codag-cli/internal/config" + "golang.org/x/sync/singleflight" ) const maxResponseBody = 16 << 20 @@ -383,6 +384,8 @@ func (c *Client) do(ctx context.Context, method, path string, body any, out any) return nil } +var anonTokenGroup singleflight.Group + func (c *Client) ensureAnonymousToken(ctx context.Context) (string, error) { cfg, err := config.Load() if err != nil { @@ -391,30 +394,42 @@ func (c *Client) ensureAnonymousToken(ctx context.Context) (string, error) { if cfg != nil && cfg.AnonymousToken != "" { return cfg.AnonymousToken, nil } - resp, err := c.AnonymousToken(ctx) + + v, err, _ := anonTokenGroup.Do("mint", func() (any, error) { + // Double-check under lock + freshCfg, lErr := config.Load() + if lErr == nil && freshCfg != nil && freshCfg.AnonymousToken != "" { + return freshCfg.AnonymousToken, nil + } + + resp, mErr := c.AnonymousToken(ctx) + if mErr != nil { + return "", mErr + } + if resp.Token == "" { + return "", fmt.Errorf("anonymous token endpoint returned an empty token") + } + + sErr := config.Modify(func(cfg *config.Config) error { + cfg.AnonymousToken = resp.Token + return nil + }) + if sErr != nil { + return "", sErr + } + return resp.Token, nil + }) if err != nil { return "", err } - if resp.Token == "" { - return "", fmt.Errorf("anonymous token endpoint returned an empty token") - } - if cfg == nil { - cfg = &config.Config{} - } - cfg.AnonymousToken = resp.Token - if err := config.Save(cfg); err != nil { - return "", err - } - return resp.Token, nil + return v.(string), nil } func clearAnonymousToken() error { - cfg, err := config.Load() - if err != nil || cfg == nil { - return err - } - cfg.AnonymousToken = "" - return config.Save(cfg) + return config.Modify(func(cfg *config.Config) error { + cfg.AnonymousToken = "" + return nil + }) } func (c *Client) doStaticBearer(ctx context.Context, method, path string, body any, out any, bearer string) error { diff --git a/internal/auth/oauth.go b/internal/auth/oauth.go index 7d8379f..d7f3fc1 100644 --- a/internal/auth/oauth.go +++ b/internal/auth/oauth.go @@ -199,40 +199,30 @@ func (e Endpoints) Whoami(accessToken string) (map[string]any, error) { return out, nil } -// SaveTokens persists a TokenResponse into the on-disk config. Wraps -// config.Save so we don't lose other fields on disk. +// SaveTokens persists a TokenResponse into the on-disk config. Uses +// config.Modify so concurrent writes don't clobber rotated tokens. func SaveTokens(tr *TokenResponse) error { - cfg, err := config.Load() - if err != nil { - return err - } - if cfg == nil { - cfg = &config.Config{} - } - cfg.AccessToken = tr.AccessToken - if tr.RefreshToken != "" { - cfg.RefreshToken = tr.RefreshToken - } - if tr.ExpiresIn > 0 { - cfg.ExpiresAt = time.Now().Unix() + int64(tr.ExpiresIn) - } - return config.Save(cfg) + return config.Modify(func(cfg *config.Config) error { + cfg.AccessToken = tr.AccessToken + if tr.RefreshToken != "" { + cfg.RefreshToken = tr.RefreshToken + } + if tr.ExpiresIn > 0 { + cfg.ExpiresAt = time.Now().Unix() + int64(tr.ExpiresIn) + } + return nil + }) } // ClearTokens wipes only the auth-related fields, preserving server URL // and provider auth. func ClearTokens() error { - cfg, err := config.Load() - if err != nil { - return err - } - if cfg == nil { + return config.Modify(func(cfg *config.Config) error { + cfg.AccessToken = "" + cfg.RefreshToken = "" + cfg.ExpiresAt = 0 return nil - } - cfg.AccessToken = "" - cfg.RefreshToken = "" - cfg.ExpiresAt = 0 - return config.Save(cfg) + }) } // OpenBrowser tries to launch the user's default browser. Returns no diff --git a/internal/config/config.go b/internal/config/config.go index 66593ca..539c2f3 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -101,11 +101,8 @@ func path() (string, error) { return filepath.Join(dir, "codag", "config.json"), nil } -// Load reads the config file. Returns an empty Config (no error) if missing. -// Unknown JSON keys (e.g. legacy `api_key`) are silently ignored. -func Load() (*Config, error) { - fileMu.Lock() - defer fileMu.Unlock() +// loadLocked reads the config file without acquiring fileMu (caller must hold fileMu). +func loadLocked() (*Config, error) { p, err := path() if err != nil { return nil, err @@ -124,13 +121,8 @@ func Load() (*Config, error) { return &c, nil } -// Save writes the config to disk (creating the directory if needed). -// File mode 0600 because it holds the OAuth refresh token. The write is -// atomic (temp file + rename) so a concurrent reader never sees a -// truncated file. -func Save(c *Config) error { - fileMu.Lock() - defer fileMu.Unlock() +// saveLocked writes the config file without acquiring fileMu (caller must hold fileMu). +func saveLocked(c *Config) error { p, err := path() if err != nil { return err @@ -161,6 +153,75 @@ func Save(c *Config) error { return os.Rename(tmp.Name(), p) } +// withFileLock executes fn while holding a cross-process file lock on config.json.lock. +func withFileLock(fn func() error) error { + p, err := path() + if err != nil { + return err + } + lockPath := p + ".lock" + if err := os.MkdirAll(filepath.Dir(lockPath), 0o700); err != nil { + return err + } + lockFile, err := os.OpenFile(lockPath, os.O_CREATE|os.O_RDWR, 0o600) + if err != nil { + return err + } + defer lockFile.Close() + + if err := lockFileFLock(lockFile); err != nil { + return err + } + defer unlockFileFLock(lockFile) + + return fn() +} + +// Load reads the config file under a cross-process file lock. Returns an empty Config (no error) if missing. +func Load() (*Config, error) { + fileMu.Lock() + defer fileMu.Unlock() + var c *Config + err := withFileLock(func() error { + var lErr error + c, lErr = loadLocked() + return lErr + }) + if err != nil { + return nil, err + } + return c, nil +} + +// Save writes the config to disk under a cross-process file lock. +func Save(c *Config) error { + fileMu.Lock() + defer fileMu.Unlock() + return withFileLock(func() error { + return saveLocked(c) + }) +} + +// Modify atomically reloads the latest config from disk under a cross-process lock, +// applies the mutate function, and saves it back to disk. +func Modify(mutate func(cfg *Config) error) error { + fileMu.Lock() + defer fileMu.Unlock() + return withFileLock(func() error { + cfg, err := loadLocked() + if err != nil { + return err + } + if cfg == nil { + cfg = &Config{} + } + if err := mutate(cfg); err != nil { + return err + } + return saveLocked(cfg) + }) +} + // ResolveServer: flag > env > config > default. func ResolveServer(flagVal string) string { if flagVal != "" { diff --git a/internal/config/config_test.go b/internal/config/config_test.go index ff8ccae..7b0abad 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -41,13 +41,13 @@ func TestSaveIsAtomicAndKeepsMode0600(t *testing.T) { if runtime.GOOS != "windows" && info.Mode().Perm() != 0o600 { t.Fatalf("config mode = %v, want 0600", info.Mode().Perm()) } - // No leftover temp files from the atomic write. + // No leftover temp files from the atomic write. Only config.json and config.json.lock are expected. entries, err := os.ReadDir(filepath.Dir(ConfigPath())) if err != nil { t.Fatalf("readdir: %v", err) } for _, e := range entries { - if e.Name() != "config.json" { + if e.Name() != "config.json" && e.Name() != "config.json.lock" { t.Fatalf("unexpected leftover file %q next to config", e.Name()) } } @@ -160,3 +160,53 @@ func TestOptOutEnvParsing(t *testing.T) { t.Error("config opt-outs not honored") } } + +func TestModifyConcurrentDeltas(t *testing.T) { + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + + // Initialize config with base tokens + if err := Save(&Config{ + AccessToken: "AT-original", + RefreshToken: "RT-1", + }); err != nil { + t.Fatalf("initial Save: %v", err) + } + + var wg sync.WaitGroup + + // Task 1: Simulates OAuth refresh (rotates AccessToken & RefreshToken) + wg.Add(1) + go func() { + defer wg.Done() + _ = Modify(func(c *Config) error { + c.AccessToken = "AT-rotated" + c.RefreshToken = "RT-2" + return nil + }) + }() + + // Task 2: Simulates update checker / telemetry writing distinct fields + wg.Add(1) + go func() { + defer wg.Done() + _ = Modify(func(c *Config) error { + c.LastUpdateCheckAt = 1234567890 + c.CLIDistinctID = "distinct-id-xyz" + return nil + }) + }() + + wg.Wait() + + finalCfg, err := Load() + if err != nil { + t.Fatalf("Load after Modify: %v", err) + } + + if finalCfg.AccessToken != "AT-rotated" || finalCfg.RefreshToken != "RT-2" { + t.Errorf("rotated tokens were lost or overwritten: AT=%q RT=%q", finalCfg.AccessToken, finalCfg.RefreshToken) + } + if finalCfg.LastUpdateCheckAt != 1234567890 || finalCfg.CLIDistinctID != "distinct-id-xyz" { + t.Errorf("concurrent deltas were lost: LastUpdateCheckAt=%d CLIDistinctID=%q", finalCfg.LastUpdateCheckAt, finalCfg.CLIDistinctID) + } +} diff --git a/internal/config/lock_other.go b/internal/config/lock_other.go new file mode 100644 index 0000000..deb872e --- /dev/null +++ b/internal/config/lock_other.go @@ -0,0 +1,17 @@ +//go:build !windows + +package config + +import ( + "os" + + "golang.org/x/sys/unix" +) + +func lockFileFLock(f *os.File) error { + return unix.Flock(int(f.Fd()), unix.LOCK_EX) +} + +func unlockFileFLock(f *os.File) error { + return unix.Flock(int(f.Fd()), unix.LOCK_UN) +} diff --git a/internal/config/lock_windows.go b/internal/config/lock_windows.go new file mode 100644 index 0000000..0fcf432 --- /dev/null +++ b/internal/config/lock_windows.go @@ -0,0 +1,30 @@ +package config + +import ( + "os" + + "golang.org/x/sys/windows" +) + +func lockFileFLock(f *os.File) error { + var lockfileOverlapped windows.Overlapped + return windows.LockFileEx( + windows.Handle(f.Fd()), + windows.LOCKFILE_EXCLUSIVE_LOCK, + 0, + 1, + 0, + &lockfileOverlapped, + ) +} + +func unlockFileFLock(f *os.File) error { + var lockfileOverlapped windows.Overlapped + return windows.UnlockFileEx( + windows.Handle(f.Fd()), + 0, + 1, + 0, + &lockfileOverlapped, + ) +} diff --git a/internal/onboard/runner.go b/internal/onboard/runner.go index 4559e15..bf49f0c 100644 --- a/internal/onboard/runner.go +++ b/internal/onboard/runner.go @@ -82,24 +82,29 @@ func (r *Runner) Run(ctx context.Context, cfg *config.Config, sources []Source, // --since to the delta. Only update for sources that processed // something AND finished without a fatal error. now := time.Now().Unix() - dirty := false + var updated []string for name, ss := range stats.bySource { if ss.LinesProcessed == 0 || ss.Err != nil { continue } - pa, ok := cfg.Providers[name] - if !ok { - pa = config.ProviderAuth{} - } - pa.LastOnboardedAt = now - if cfg.Providers == nil { - cfg.Providers = map[string]config.ProviderAuth{} - } - cfg.Providers[name] = pa - dirty = true + updated = append(updated, name) } - if dirty { - _ = config.Save(cfg) + + if len(updated) > 0 { + _ = config.Modify(func(cfg *config.Config) error { + if cfg.Providers == nil { + cfg.Providers = map[string]config.ProviderAuth{} + } + for _, name := range updated { + pa, ok := cfg.Providers[name] + if !ok { + pa = config.ProviderAuth{} + } + pa.LastOnboardedAt = now + cfg.Providers[name] = pa + } + return nil + }) } return stats, firstErr }