diff --git a/internal/auth/claude/identity.go b/internal/auth/claude/identity.go index c3fdd3ee..df41b99c 100644 --- a/internal/auth/claude/identity.go +++ b/internal/auth/claude/identity.go @@ -14,6 +14,12 @@ const ( claudeDeviceIDByteSize = 32 ) +// claudeDevicePoolMu guards every concurrent access to a Claude credential's +// Auth.Metadata map, not just the device pool. A single Auth is shared by all +// in-flight requests using that credential, and Go maps are not safe for +// concurrent read/write, so the account-profile and refresh paths have to take +// the same lock as the pool paths. Reaching into Auth.Metadata directly from a +// request path is a data race even when the keys differ. var claudeDevicePoolMu sync.Mutex // GenerateDeviceIDPool creates the fixed-size device pool stored with a Claude credential. @@ -105,6 +111,122 @@ func EnsureDeviceIDPool(metadata map[string]any) ([]string, bool, error) { claudeDevicePoolMu.Lock() defer claudeDevicePoolMu.Unlock() + return ensureDeviceIDPoolLocked(metadata) +} + +// EnsureDeviceIDPoolFor lazily initializes the metadata map and then ensures the +// pool, both under the device pool lock. +// +// A single *Auth is shared by every concurrent request that selects the same +// credential, so initializing the map field outside this lock races with the +// writes below and can abort the process with "concurrent map writes". Callers +// holding a shared credential must reach the pool through this package rather +// than touching the map directly. +func EnsureDeviceIDPoolFor(metadata *map[string]any) ([]string, bool, error) { + if metadata == nil { + return nil, false, fmt.Errorf("ensure Claude device pool: metadata pointer is nil") + } + claudeDevicePoolMu.Lock() + defer claudeDevicePoolMu.Unlock() + + if *metadata == nil { + *metadata = make(map[string]any) + } + return ensureDeviceIDPoolLocked(*metadata) +} + +// ReadDeviceIDPool returns the raw stored pool value, initializing the map when +// needed, under the device pool lock. +func ReadDeviceIDPool(metadata *map[string]any) any { + if metadata == nil { + return nil + } + claudeDevicePoolMu.Lock() + defer claudeDevicePoolMu.Unlock() + + if *metadata == nil { + *metadata = make(map[string]any) + return nil + } + return (*metadata)[ClaudeDeviceIDsMetadataKey] +} + +// StoreDeviceIDPool writes a defensive copy of deviceIDs under the device pool lock. +func StoreDeviceIDPool(metadata *map[string]any, deviceIDs []string) { + if metadata == nil { + return + } + claudeDevicePoolMu.Lock() + defer claudeDevicePoolMu.Unlock() + + if *metadata == nil { + *metadata = make(map[string]any) + } + (*metadata)[ClaudeDeviceIDsMetadataKey] = append([]string(nil), deviceIDs...) +} + +// ReadMetadataString reads a string-valued metadata entry under the metadata +// lock, so it cannot observe a map being concurrently written by another path. +func ReadMetadataString(metadata *map[string]any, key string) string { + if metadata == nil { + return "" + } + claudeDevicePoolMu.Lock() + defer claudeDevicePoolMu.Unlock() + + if *metadata == nil { + return "" + } + value, _ := (*metadata)[key].(string) + return value +} + +// StoreMetadataString writes a string-valued metadata entry under the metadata +// lock, initializing the map when needed. Empty values are skipped so callers can +// forward optional fields without erasing a previously resolved value. +func StoreMetadataString(metadata *map[string]any, key, value string) { + if metadata == nil || strings.TrimSpace(value) == "" { + return + } + claudeDevicePoolMu.Lock() + defer claudeDevicePoolMu.Unlock() + + if *metadata == nil { + *metadata = make(map[string]any) + } + (*metadata)[key] = value +} + +// StoreMetadataValue writes an arbitrary metadata entry under the metadata lock, +// initializing the map when needed. +func StoreMetadataValue(metadata *map[string]any, key string, value any) { + if metadata == nil { + return + } + claudeDevicePoolMu.Lock() + defer claudeDevicePoolMu.Unlock() + + if *metadata == nil { + *metadata = make(map[string]any) + } + (*metadata)[key] = value +} + +// EnsureMetadataMap initializes the metadata map under the metadata lock. +func EnsureMetadataMap(metadata *map[string]any) { + if metadata == nil { + return + } + claudeDevicePoolMu.Lock() + defer claudeDevicePoolMu.Unlock() + + if *metadata == nil { + *metadata = make(map[string]any) + } +} + +// ensureDeviceIDPoolLocked requires claudeDevicePoolMu to be held. +func ensureDeviceIDPoolLocked(metadata map[string]any) ([]string, bool, error) { if metadata == nil { return nil, false, fmt.Errorf("ensure Claude device pool: metadata is nil") } diff --git a/internal/runtime/executor/claude_executor_auth.go b/internal/runtime/executor/claude_executor_auth.go index db0ebf67..5ca5fdf2 100644 --- a/internal/runtime/executor/claude_executor_auth.go +++ b/internal/runtime/executor/claude_executor_auth.go @@ -25,20 +25,19 @@ func (e *ClaudeExecutor) ShouldPrepareRequestAuth(auth *cliproxyauth.Auth) bool if !isClaudeOAuthToken(apiKey) || auth == nil { return false } - if !claudeauth.HasCanonicalDeviceIDPool(auth.Metadata[claudeauth.ClaudeDeviceIDsMetadataKey]) { + if !claudeauth.HasCanonicalDeviceIDPool(claudeauth.ReadDeviceIDPool(&auth.Metadata)) { return true } if helps.ClaudeCredentialAccountUUID(auth) != "" { return false } - return claudeAccountProfileLookupDue(auth.Metadata, time.Now()) + return claudeAccountProfileLookupDue(claudeauth.ReadMetadataString(&auth.Metadata, claudeAccountProfileCheckedAtKey), time.Now()) } -func claudeAccountProfileLookupDue(metadata map[string]any, now time.Time) bool { - if metadata == nil { - return true - } - checkedAt, _ := metadata[claudeAccountProfileCheckedAtKey].(string) +// claudeAccountProfileLookupDue takes the already-read timestamp rather than the +// metadata map: the map belongs to a credential shared by concurrent requests and +// may only be touched under the metadata lock. +func claudeAccountProfileLookupDue(checkedAt string, now time.Time) bool { checkedAt = strings.TrimSpace(checkedAt) if checkedAt == "" { return true @@ -52,17 +51,16 @@ func (e *ClaudeExecutor) PrepareRequestAuth(ctx context.Context, auth *cliproxya return auth, nil } apiKey, _ := claudeCreds(auth) - if auth.Metadata == nil { - auth.Metadata = make(map[string]any) - } + claudeauth.EnsureMetadataMap(&auth.Metadata) if _, errDeviceIDs := helps.EnsureClaudeCredentialDevicePoolRequired(ctx, auth); errDeviceIDs != nil { return nil, errDeviceIDs } - if helps.ClaudeCredentialAccountUUID(auth) != "" || !claudeAccountProfileLookupDue(auth.Metadata, time.Now()) { + if helps.ClaudeCredentialAccountUUID(auth) != "" || + !claudeAccountProfileLookupDue(claudeauth.ReadMetadataString(&auth.Metadata, claudeAccountProfileCheckedAtKey), time.Now()) { return auth, nil } - auth.Metadata[claudeAccountProfileCheckedAtKey] = time.Now().UTC().Format(time.RFC3339) + claudeauth.StoreMetadataString(&auth.Metadata, claudeAccountProfileCheckedAtKey, time.Now().UTC().Format(time.RFC3339)) profile, errProfile := e.fetchClaudeOAuthProfile(ctx, auth, apiKey) if errProfile != nil { if errContext := ctx.Err(); errContext != nil { @@ -74,18 +72,10 @@ func (e *ClaudeExecutor) PrepareRequestAuth(ctx context.Context, auth *cliproxya if profile == nil { return auth, nil } - if accountUUID := strings.TrimSpace(profile.Account.UUID); accountUUID != "" { - auth.Metadata["account_uuid"] = accountUUID - } - if email := strings.TrimSpace(profile.Account.Email); email != "" { - auth.Metadata["email"] = email - } - if organizationUUID := strings.TrimSpace(profile.Organization.UUID); organizationUUID != "" { - auth.Metadata["organization_uuid"] = organizationUUID - } - if organizationName := strings.TrimSpace(profile.Organization.Name); organizationName != "" { - auth.Metadata["organization_name"] = organizationName - } + claudeauth.StoreMetadataString(&auth.Metadata, "account_uuid", profile.Account.UUID) + claudeauth.StoreMetadataString(&auth.Metadata, "email", profile.Account.Email) + claudeauth.StoreMetadataString(&auth.Metadata, "organization_uuid", profile.Organization.UUID) + claudeauth.StoreMetadataString(&auth.Metadata, "organization_name", profile.Organization.Name) return auth, nil } @@ -113,12 +103,7 @@ func (e *ClaudeExecutor) Refresh(ctx context.Context, auth *cliproxyauth.Auth) ( if auth == nil { return nil, fmt.Errorf("claude executor: auth is nil") } - var refreshToken string - if auth.Metadata != nil { - if v, ok := auth.Metadata["refresh_token"].(string); ok && v != "" { - refreshToken = v - } - } + refreshToken := claudeauth.ReadMetadataString(&auth.Metadata, "refresh_token") if refreshToken == "" { return auth, nil } @@ -127,26 +112,17 @@ func (e *ClaudeExecutor) Refresh(ctx context.Context, auth *cliproxyauth.Auth) ( if err != nil { return nil, err } - if auth.Metadata == nil { - auth.Metadata = make(map[string]any) - } - auth.Metadata["access_token"] = td.AccessToken - if td.RefreshToken != "" { - auth.Metadata["refresh_token"] = td.RefreshToken - } - auth.Metadata["email"] = td.Email - if td.AccountUUID != "" { - auth.Metadata["account_uuid"] = td.AccountUUID - } - if td.OrganizationUUID != "" { - auth.Metadata["organization_uuid"] = td.OrganizationUUID - } - if td.OrganizationName != "" { - auth.Metadata["organization_name"] = td.OrganizationName - } - auth.Metadata["expired"] = td.Expire - auth.Metadata["type"] = "claude" - now := time.Now().Format(time.RFC3339) - auth.Metadata["last_refresh"] = now + claudeauth.EnsureMetadataMap(&auth.Metadata) + claudeauth.StoreMetadataValue(&auth.Metadata, "access_token", td.AccessToken) + claudeauth.StoreMetadataString(&auth.Metadata, "refresh_token", td.RefreshToken) + // email is written unconditionally to preserve the previous reset-on-refresh + // behaviour; the remaining optional fields keep their prior value when absent. + claudeauth.StoreMetadataValue(&auth.Metadata, "email", td.Email) + claudeauth.StoreMetadataString(&auth.Metadata, "account_uuid", td.AccountUUID) + claudeauth.StoreMetadataString(&auth.Metadata, "organization_uuid", td.OrganizationUUID) + claudeauth.StoreMetadataString(&auth.Metadata, "organization_name", td.OrganizationName) + claudeauth.StoreMetadataValue(&auth.Metadata, "expired", td.Expire) + claudeauth.StoreMetadataValue(&auth.Metadata, "type", "claude") + claudeauth.StoreMetadataValue(&auth.Metadata, "last_refresh", time.Now().Format(time.RFC3339)) return auth, nil } diff --git a/internal/runtime/executor/claude_executor_auth_race_test.go b/internal/runtime/executor/claude_executor_auth_race_test.go new file mode 100644 index 00000000..613762a1 --- /dev/null +++ b/internal/runtime/executor/claude_executor_auth_race_test.go @@ -0,0 +1,100 @@ +package executor + +import ( + "context" + "sync" + "testing" + + claudeauth "github.com/router-for-me/CLIProxyAPI/v7/internal/auth/claude" + "github.com/router-for-me/CLIProxyAPI/v7/internal/config" + cliproxyauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" +) + +// A single Auth is shared by every in-flight request that selects the credential, +// so any request path reaching into Auth.Metadata directly races the others. An +// earlier fix locked only the device pool helpers and left the account-profile +// path unguarded, which these tests would have caught: they drive the exported +// entry points rather than the helper that was known to be broken. + +func newSharedClaudeOAuthAuth(id string) *cliproxyauth.Auth { + return &cliproxyauth.Auth{ + ID: id, + Attributes: map[string]string{"api_key": "sk-ant-oat-race-probe"}, + Metadata: map[string]any{}, + } +} + +func TestClaudeExecutorPrepareRequestAuthIsRaceFreeOnSharedCredential(t *testing.T) { + executor := NewClaudeExecutor(&config.Config{}) + executor.oauthProfileFetcher = func(context.Context, *cliproxyauth.Auth, string) (*claudeauth.OAuthProfile, error) { + profile := &claudeauth.OAuthProfile{} + profile.Account.UUID = "aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa" + profile.Account.Email = "user@example.com" + profile.Organization.UUID = "bbbbbbbb-bbbb-4bbb-8bbb-bbbbbbbbbbbb" + profile.Organization.Name = "Example Org" + return profile, nil + } + + auth := newSharedClaudeOAuthAuth("claude-race-prepare") + ctx := context.Background() + + var wg sync.WaitGroup + for i := 0; i < 32; i++ { + wg.Add(1) + go func() { + defer wg.Done() + // ShouldPrepareRequestAuth reads the same map the writers below mutate. + if executor.ShouldPrepareRequestAuth(auth) { + if _, err := executor.PrepareRequestAuth(ctx, auth); err != nil { + t.Errorf("PrepareRequestAuth() error = %v", err) + } + return + } + _ = executor.ShouldPrepareRequestAuth(auth) + }() + } + wg.Wait() + + if got := claudeauth.ReadMetadataString(&auth.Metadata, "account_uuid"); got != "aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa" { + t.Fatalf("account_uuid = %q, want the fetched profile account", got) + } + if !claudeauth.HasCanonicalDeviceIDPool(claudeauth.ReadDeviceIDPool(&auth.Metadata)) { + t.Fatal("device ID pool was not established under concurrency") + } +} + +// TestClaudeExecutorSharedCredentialMetadataMixedAccess drives the request-path +// readers against the profile writer at the same time, which is the shape that +// produced the reported data races. +func TestClaudeExecutorSharedCredentialMetadataMixedAccess(t *testing.T) { + executor := NewClaudeExecutor(&config.Config{}) + executor.oauthProfileFetcher = func(context.Context, *cliproxyauth.Auth, string) (*claudeauth.OAuthProfile, error) { + profile := &claudeauth.OAuthProfile{} + profile.Account.UUID = "cccccccc-cccc-4ccc-8ccc-cccccccccccc" + return profile, nil + } + + auth := newSharedClaudeOAuthAuth("claude-race-mixed") + ctx := context.Background() + + var wg sync.WaitGroup + for i := 0; i < 32; i++ { + wg.Add(1) + go func(i int) { + defer wg.Done() + switch i % 4 { + case 0: + if _, err := executor.PrepareRequestAuth(ctx, auth); err != nil { + t.Errorf("PrepareRequestAuth() error = %v", err) + } + case 1: + _ = executor.ShouldPrepareRequestAuth(auth) + case 2: + _ = claudeauth.ReadMetadataString(&auth.Metadata, "account_uuid") + default: + _ = claudeauth.ReadDeviceIDPool(&auth.Metadata) + } + }(i) + } + wg.Wait() +} diff --git a/internal/runtime/executor/helps/claude_credential_identity.go b/internal/runtime/executor/helps/claude_credential_identity.go index ca5fae7c..e93b0756 100644 --- a/internal/runtime/executor/helps/claude_credential_identity.go +++ b/internal/runtime/executor/helps/claude_credential_identity.go @@ -71,10 +71,7 @@ func EnsureClaudeCredentialDevicePoolRequired(ctx context.Context, auth *cliprox if auth == nil { return nil, fmt.Errorf("ensure Claude credential device pool: auth is nil") } - if auth.Metadata == nil { - auth.Metadata = make(map[string]any) - } - rawCredentialDeviceIDs := auth.Metadata[claudeauth.ClaudeDeviceIDsMetadataKey] + rawCredentialDeviceIDs := claudeauth.ReadDeviceIDPool(&auth.Metadata) if claudeauth.HasCanonicalDeviceIDPool(rawCredentialDeviceIDs) { return claudeauth.NormalizeDeviceIDPool(rawCredentialDeviceIDs), nil } @@ -82,7 +79,7 @@ func EnsureClaudeCredentialDevicePoolRequired(ctx context.Context, auth *cliprox client, homeMode, errClient := currentClaudeCredentialDevicePoolKVClient() if !homeMode { - deviceIDs, _, errEnsure := claudeauth.EnsureDeviceIDPool(auth.Metadata) + deviceIDs, _, errEnsure := claudeauth.EnsureDeviceIDPoolFor(&auth.Metadata) return deviceIDs, errEnsure } if errClient != nil { @@ -115,7 +112,7 @@ func EnsureClaudeCredentialDevicePoolRequired(ctx context.Context, auth *cliprox return nil, fmt.Errorf("ensure Claude credential device pool: canonical Home KV value was not written") } } - auth.Metadata[claudeauth.ClaudeDeviceIDsMetadataKey] = append([]string(nil), deviceIDs...) + claudeauth.StoreDeviceIDPool(&auth.Metadata, deviceIDs) return deviceIDs, nil } } @@ -151,18 +148,17 @@ func EnsureClaudeCredentialDevicePoolRequired(ctx context.Context, auth *cliprox if len(deviceIDs) != claudeauth.ClaudeDevicePoolSize { return nil, fmt.Errorf("ensure Claude credential device pool: Home KV pool has %d entries, want %d", len(deviceIDs), claudeauth.ClaudeDevicePoolSize) } - auth.Metadata[claudeauth.ClaudeDeviceIDsMetadataKey] = append([]string(nil), deviceIDs...) + claudeauth.StoreDeviceIDPool(&auth.Metadata, deviceIDs) return deviceIDs, nil } // ClaudeCredentialAccountUUID returns the selected upstream credential's account UUID. func ClaudeCredentialAccountUUID(auth *cliproxyauth.Auth) string { - if auth == nil || auth.Metadata == nil { + if auth == nil { return "" } for _, key := range []string{"account_uuid", "accountUuid"} { - value, _ := auth.Metadata[key].(string) - value = strings.TrimSpace(value) + value := strings.TrimSpace(claudeauth.ReadMetadataString(&auth.Metadata, key)) if value != "" { return value } @@ -175,10 +171,7 @@ func ApplyClaudeCredentialMetadata(payload []byte, auth *cliproxyauth.Auth, sess if auth == nil { return nil, "", fmt.Errorf("apply Claude credential metadata: auth is nil") } - if auth.Metadata == nil { - auth.Metadata = make(map[string]any) - } - deviceIDs, _, errDeviceIDs := claudeauth.EnsureDeviceIDPool(auth.Metadata) + deviceIDs, _, errDeviceIDs := claudeauth.EnsureDeviceIDPoolFor(&auth.Metadata) if errDeviceIDs != nil { return nil, "", errDeviceIDs } diff --git a/internal/runtime/executor/helps/claude_credential_identity_race_test.go b/internal/runtime/executor/helps/claude_credential_identity_race_test.go new file mode 100644 index 00000000..695d8c44 --- /dev/null +++ b/internal/runtime/executor/helps/claude_credential_identity_race_test.go @@ -0,0 +1,100 @@ +package helps + +import ( + "errors" + "sync" + "testing" + + cliproxyauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" +) + +// TestApplyClaudeCredentialMetadataConcurrentSharedAuth pins the invariant that a +// single *Auth shared by concurrent requests is safe to use. Before the device +// pool accessors were introduced these paths initialized and wrote auth.Metadata +// outside claudeDevicePoolMu, which aborts the process with "concurrent map +// writes" rather than failing a request. Run with -race. +func TestApplyClaudeCredentialMetadataConcurrentSharedAuth(t *testing.T) { + auth := &cliproxyauth.Auth{ID: "shared-credential"} + payload := []byte(`{"model":"claude-opus-4-6","messages":[{"role":"user","content":"hi"}]}`) + + const goroutines = 32 + var wg sync.WaitGroup + errs := make(chan error, goroutines) + start := make(chan struct{}) + + for i := range goroutines { + wg.Add(1) + go func(i int) { + defer wg.Done() + <-start + sessionID := "session-" + string(rune('a'+i%26)) + if _, _, err := ApplyClaudeCredentialMetadata(payload, auth, sessionID); err != nil { + errs <- err + return + } + // Concurrent readers of the same map must be safe too. + _ = ClaudeCredentialAccountUUID(auth) + }(i) + } + + close(start) + wg.Wait() + close(errs) + for err := range errs { + t.Fatalf("ApplyClaudeCredentialMetadata on shared auth: %v", err) + } + + if auth.Metadata == nil { + t.Fatal("expected metadata to be initialized") + } +} + +// TestEnsureClaudeCredentialDevicePoolConcurrentSharedAuth covers the local +// (non Home KV) branch of the pool bootstrap on a shared credential. +func TestEnsureClaudeCredentialDevicePoolConcurrentSharedAuth(t *testing.T) { + auth := &cliproxyauth.Auth{ID: "shared-credential"} + + const goroutines = 32 + var wg sync.WaitGroup + results := make(chan string, goroutines) + errs := make(chan error, goroutines) + start := make(chan struct{}) + + for range goroutines { + wg.Add(1) + go func() { + defer wg.Done() + <-start + deviceIDs, err := EnsureClaudeCredentialDevicePoolRequired(t.Context(), auth) + if err != nil { + errs <- err + return + } + if len(deviceIDs) == 0 { + errs <- errEmptyPool + return + } + results <- deviceIDs[0] + }() + } + + close(start) + wg.Wait() + close(errs) + close(results) + for err := range errs { + t.Fatalf("EnsureClaudeCredentialDevicePoolRequired on shared auth: %v", err) + } + + // Every caller must agree on the pool; a racing bootstrap would hand out + // different device IDs to different requests on the same credential. + seen := make(map[string]struct{}) + for deviceID := range results { + seen[deviceID] = struct{}{} + } + if len(seen) != 1 { + t.Fatalf("device pool bootstrap was not stable: got %d distinct device IDs, want 1", len(seen)) + } +} + +var errEmptyPool = errors.New("device pool is empty")