From 25d2b03e909c99eb3af7cbdc05166c8dafacfc2f Mon Sep 17 00:00:00 2001 From: IanShaw027 Date: Fri, 7 Aug 2026 16:29:48 +0800 Subject: [PATCH] =?UTF-8?q?fix(grok):=20=E5=8A=A0=E5=9B=BA=20OAuth=20?= =?UTF-8?q?=E4=BC=9A=E8=AF=9D=E5=85=B1=E4=BA=AB=E4=B8=8E=E4=B8=80=E6=AC=A1?= =?UTF-8?q?=E6=80=A7=E6=B6=88=E8=B4=B9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/cmd/server/wire_gen.go | 2 +- backend/go.sum | 2 + backend/internal/pkg/redissession/store.go | 125 +++++++++++++++++ .../internal/pkg/redissession/store_test.go | 40 ++++++ backend/internal/pkg/xai/oauth.go | 127 +++++++++++++++++- .../pkg/xai/oauth_redis_fallback_test.go | 35 +++++ .../internal/service/grok_oauth_service.go | 20 ++- .../service/grok_oauth_service_test.go | 29 +++- backend/internal/service/wire.go | 4 +- 9 files changed, 371 insertions(+), 13 deletions(-) create mode 100644 backend/internal/pkg/redissession/store.go create mode 100644 backend/internal/pkg/redissession/store_test.go create mode 100644 backend/internal/pkg/xai/oauth_redis_fallback_test.go diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 91cf74dc2..4f5e675bb 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -149,7 +149,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { openAIOAuthService := service.ProvideOpenAIOAuthService(proxyRepository, openAIOAuthClient, privacyClientFactory) openAITokenProvider := service.ProvideOpenAITokenProvider(accountRepository, geminiTokenCache, openAIOAuthService, oAuthRefreshAPI) grokOAuthClient := repository.NewGrokOAuthClient() - grokOAuthService := service.ProvideGrokOAuthService(proxyRepository, grokOAuthClient, configConfig) + grokOAuthService := service.ProvideGrokOAuthService(proxyRepository, grokOAuthClient, configConfig, redisClient) grokTokenProvider := service.ProvideGrokTokenProvider(accountRepository, geminiTokenCache, grokOAuthService, oAuthRefreshAPI, tempUnschedCache) openAIGatewayService := service.NewOpenAIGatewayService(accountRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, httpUpstream, deferredService, openAITokenProvider, grokTokenProvider, modelPricingResolver, channelService, balanceNotifyService, settingService, serviceUserPlatformQuotaRepository) geminiOAuthClient := repository.NewGeminiOAuthClient(configConfig) diff --git a/backend/go.sum b/backend/go.sum index 98c57415c..cdea1a5b9 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -250,6 +250,8 @@ github.com/google/go-tpm-tools v0.3.13-0.20230620182252-4639ecce2aba/go.mod h1:E github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= +github.com/google/subcommands v1.2.0 h1:vWQspBTo2nEqTUFita5/KeEWlUL8kQObDFbub/EN9oE= +github.com/google/subcommands v1.2.0/go.mod h1:ZjhPrFU+Olkh9WazFPsl27BQ4UPiG37m3yTrtFlrHVk= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/google/wire v0.7.0 h1:JxUKI6+CVBgCO2WToKy/nQk0sS+amI9z9EjVmdaocj4= diff --git a/backend/internal/pkg/redissession/store.go b/backend/internal/pkg/redissession/store.go new file mode 100644 index 000000000..3909fceec --- /dev/null +++ b/backend/internal/pkg/redissession/store.go @@ -0,0 +1,125 @@ +// Package redissession provides a multi-instance OAuth session backend. +package redissession + +import ( + "context" + "encoding/json" + "errors" + "strings" + "time" + + "github.com/redis/go-redis/v9" +) + +var ErrNotConfigured = errors.New("redis session store not configured") + +// Store persists JSON sessions and single-use markers under one namespace. +type Store struct { + rdb *redis.Client + prefix string + ttl time.Duration +} + +func New(rdb *redis.Client, prefix string, ttl time.Duration) *Store { + if ttl <= 0 { + ttl = 30 * time.Minute + } + prefix = strings.TrimSpace(prefix) + if prefix == "" { + prefix = "oauth:session" + } + if !strings.HasSuffix(prefix, ":") { + prefix += ":" + } + return &Store{rdb: rdb, prefix: prefix, ttl: ttl} +} + +func (s *Store) dataKey(id string) string { return s.prefix + strings.TrimSpace(id) } +func (s *Store) usedKey(id string) string { return s.prefix + "used:" + strings.TrimSpace(id) } + +func (s *Store) Set(ctx context.Context, id string, value any) error { + if s == nil || s.rdb == nil { + return ErrNotConfigured + } + id = strings.TrimSpace(id) + if id == "" { + return errors.New("session id is required") + } + if ctx == nil { + ctx = context.Background() + } + raw, err := json.Marshal(value) + if err != nil { + return err + } + return s.rdb.Set(ctx, s.dataKey(id), raw, s.ttl).Err() +} + +func (s *Store) Get(ctx context.Context, id string, dest any) (bool, error) { + if s == nil || s.rdb == nil { + return false, ErrNotConfigured + } + id = strings.TrimSpace(id) + if id == "" { + return false, nil + } + if ctx == nil { + ctx = context.Background() + } + raw, err := s.rdb.Get(ctx, s.dataKey(id)).Bytes() + if errors.Is(err, redis.Nil) { + return false, nil + } + if err != nil { + return false, err + } + if err := json.Unmarshal(raw, dest); err != nil { + return false, err + } + return true, nil +} + +func (s *Store) Delete(ctx context.Context, id string) error { + if s == nil || s.rdb == nil { + return ErrNotConfigured + } + id = strings.TrimSpace(id) + if id == "" { + return nil + } + if ctx == nil { + ctx = context.Background() + } + return s.rdb.Del(ctx, s.dataKey(id), s.usedKey(id)).Err() +} + +// TryConsume returns true only for the first claim while the session exists. +func (s *Store) TryConsume(ctx context.Context, id string) (bool, error) { + if s == nil || s.rdb == nil { + return false, ErrNotConfigured + } + id = strings.TrimSpace(id) + if id == "" { + return false, nil + } + if ctx == nil { + ctx = context.Background() + } + ttl := s.ttl + if remaining, err := s.rdb.TTL(ctx, s.dataKey(id)).Result(); err == nil && remaining > 0 { + ttl = remaining + } + ok, err := s.rdb.SetNX(ctx, s.usedKey(id), "1", ttl).Result() + if err != nil || !ok { + return ok, err + } + exists, err := s.rdb.Exists(ctx, s.dataKey(id)).Result() + if err != nil { + return false, err + } + if exists == 0 { + _ = s.rdb.Del(ctx, s.usedKey(id)).Err() + return false, nil + } + return true, nil +} diff --git a/backend/internal/pkg/redissession/store_test.go b/backend/internal/pkg/redissession/store_test.go new file mode 100644 index 000000000..ad5832fb8 --- /dev/null +++ b/backend/internal/pkg/redissession/store_test.go @@ -0,0 +1,40 @@ +//go:build unit + +package redissession + +import ( + "context" + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/require" +) + +func TestStoreRoundTripAndSingleUse(t *testing.T) { + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + t.Cleanup(func() { _ = rdb.Close() }) + store := New(rdb, "oauth:test", time.Minute) + ctx := context.Background() + + require.NoError(t, store.Set(ctx, "sid", map[string]string{"state": "state"})) + var got map[string]string + ok, err := store.Get(ctx, "sid", &got) + require.NoError(t, err) + require.True(t, ok) + require.Equal(t, "state", got["state"]) + + ok, err = store.TryConsume(ctx, "sid") + require.NoError(t, err) + require.True(t, ok) + ok, err = store.TryConsume(ctx, "sid") + require.NoError(t, err) + require.False(t, ok) + + require.NoError(t, store.Delete(ctx, "sid")) + ok, err = store.Get(ctx, "sid", &got) + require.NoError(t, err) + require.False(t, ok) +} diff --git a/backend/internal/pkg/xai/oauth.go b/backend/internal/pkg/xai/oauth.go index afa44d06a..578c7278c 100644 --- a/backend/internal/pkg/xai/oauth.go +++ b/backend/internal/pkg/xai/oauth.go @@ -1,20 +1,24 @@ package xai import ( + "context" "crypto/rand" "crypto/sha256" "encoding/base64" "encoding/hex" "errors" "fmt" + "log/slog" "net/url" "os" "strings" "sync" "time" + "github.com/Wei-Shaw/sub2api/internal/pkg/redissession" "github.com/Wei-Shaw/sub2api/internal/util/logredact" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" + "github.com/redis/go-redis/v9" ) const ( @@ -56,32 +60,110 @@ type OAuthSession struct { ProxyURL string `json:"proxy_url,omitempty"` RedirectURI string `json:"redirect_uri"` CreatedAt time.Time `json:"created_at"` + + mu sync.Mutex + consumed bool } -// SessionStore manages xAI OAuth sessions in memory. +func (s *OAuthSession) TryConsume() bool { + if s == nil { + return false + } + s.mu.Lock() + defer s.mu.Unlock() + if s.consumed { + return false + } + s.consumed = true + return true +} + +// SessionStore manages xAI OAuth sessions with an optional Redis backend. type SessionStore struct { - mu sync.RWMutex - sessions map[string]*OAuthSession - stopOnce sync.Once - stopCh chan struct{} + mu sync.RWMutex + sessions map[string]*OAuthSession + localOnly map[string]struct{} + stopOnce sync.Once + stopCh chan struct{} + remote *redissession.Store +} + +type oauthSessionDTO struct { + State string `json:"state"` + CodeVerifier string `json:"code_verifier"` + CodeChallenge string `json:"code_challenge"` + ClientID string `json:"client_id,omitempty"` + Scope string `json:"scope,omitempty"` + ProxyURL string `json:"proxy_url,omitempty"` + RedirectURI string `json:"redirect_uri"` + CreatedAt time.Time `json:"created_at"` } func NewSessionStore() *SessionStore { store := &SessionStore{ - sessions: make(map[string]*OAuthSession), - stopCh: make(chan struct{}), + sessions: make(map[string]*OAuthSession), + localOnly: make(map[string]struct{}), + stopCh: make(chan struct{}), } go store.cleanup() return store } +func NewRedisSessionStore(rdb *redis.Client) *SessionStore { + store := NewSessionStore() + if rdb != nil { + store.remote = redissession.New(rdb, "oauth:session:xai", SessionTTL) + } + return store +} + func (s *SessionStore) Set(sessionID string, session *OAuthSession) { + if session == nil { + return + } + var remoteErr error + if s != nil && s.remote != nil { + remoteErr = s.remote.Set(context.Background(), sessionID, oauthSessionDTO{ + State: session.State, CodeVerifier: session.CodeVerifier, CodeChallenge: session.CodeChallenge, + ClientID: session.ClientID, Scope: session.Scope, ProxyURL: session.ProxyURL, + RedirectURI: session.RedirectURI, CreatedAt: session.CreatedAt, + }) + } s.mu.Lock() defer s.mu.Unlock() s.sessions[sessionID] = session + if remoteErr != nil { + s.localOnly[sessionID] = struct{}{} + slog.Warn("xai oauth session Redis write failed; using process-local fallback", "error", remoteErr) + } else { + delete(s.localOnly, sessionID) + } } func (s *SessionStore) Get(sessionID string) (*OAuthSession, bool) { + if s.isLocalOnly(sessionID) { + return s.getMemory(sessionID) + } + if s != nil && s.remote != nil { + var dto oauthSessionDTO + ok, err := s.remote.Get(context.Background(), sessionID, &dto) + if err != nil || !ok || time.Since(dto.CreatedAt) > SessionTTL { + return nil, false + } + session := &OAuthSession{ + State: dto.State, CodeVerifier: dto.CodeVerifier, CodeChallenge: dto.CodeChallenge, + ClientID: dto.ClientID, Scope: dto.Scope, ProxyURL: dto.ProxyURL, + RedirectURI: dto.RedirectURI, CreatedAt: dto.CreatedAt, + } + s.mu.Lock() + s.sessions[sessionID] = session + s.mu.Unlock() + return session, true + } + return s.getMemory(sessionID) +} + +func (s *SessionStore) getMemory(sessionID string) (*OAuthSession, bool) { s.mu.RLock() defer s.mu.RUnlock() session, ok := s.sessions[sessionID] @@ -95,9 +177,39 @@ func (s *SessionStore) Get(sessionID string) (*OAuthSession, bool) { } func (s *SessionStore) Delete(sessionID string) { + if s != nil && s.remote != nil { + _ = s.remote.Delete(context.Background(), sessionID) + } s.mu.Lock() defer s.mu.Unlock() delete(s.sessions, sessionID) + delete(s.localOnly, sessionID) +} + +func (s *SessionStore) TryConsumeSession(sessionID string) bool { + if s == nil { + return false + } + if s.isLocalOnly(sessionID) { + return s.tryConsumeMemory(sessionID) + } + if s.remote != nil { + ok, err := s.remote.TryConsume(context.Background(), sessionID) + return err == nil && ok + } + return s.tryConsumeMemory(sessionID) +} + +func (s *SessionStore) isLocalOnly(sessionID string) bool { + s.mu.RLock() + defer s.mu.RUnlock() + _, ok := s.localOnly[sessionID] + return ok +} + +func (s *SessionStore) tryConsumeMemory(sessionID string) bool { + session, ok := s.getMemory(sessionID) + return ok && session.TryConsume() } func (s *SessionStore) Stop() { @@ -118,6 +230,7 @@ func (s *SessionStore) cleanup() { for id, session := range s.sessions { if time.Since(session.CreatedAt) > SessionTTL { delete(s.sessions, id) + delete(s.localOnly, id) } } s.mu.Unlock() diff --git a/backend/internal/pkg/xai/oauth_redis_fallback_test.go b/backend/internal/pkg/xai/oauth_redis_fallback_test.go new file mode 100644 index 000000000..40023acd6 --- /dev/null +++ b/backend/internal/pkg/xai/oauth_redis_fallback_test.go @@ -0,0 +1,35 @@ +//go:build unit + +package xai + +import ( + "context" + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/require" +) + +func TestSessionStoreRedisFallbackIsLimitedToFailedWrites(t *testing.T) { + mr := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{Addr: mr.Addr(), MaxRetries: -1}) + t.Cleanup(func() { _ = client.Close() }) + store := NewRedisSessionStore(client) + defer store.Stop() + session := func(state string) *OAuthSession { return &OAuthSession{State: state, CreatedAt: time.Now()} } + + store.Set("remote", session("remote")) + require.NoError(t, store.remote.Delete(context.Background(), "remote")) + _, ok := store.Get("remote") + require.False(t, ok, "a remote miss must not revive the stale local copy") + + mr.Close() + store.Set("local-only", session("local")) + got, ok := store.Get("local-only") + require.True(t, ok) + require.Equal(t, "local", got.State) + require.True(t, store.TryConsumeSession("local-only")) + require.False(t, store.TryConsumeSession("local-only")) +} diff --git a/backend/internal/service/grok_oauth_service.go b/backend/internal/service/grok_oauth_service.go index 3cd9e1322..4329d90c1 100644 --- a/backend/internal/service/grok_oauth_service.go +++ b/backend/internal/service/grok_oauth_service.go @@ -11,6 +11,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/config" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/redis/go-redis/v9" ) const grokDefaultAccessTokenTTL = 6 * time.Hour @@ -34,6 +35,17 @@ func NewGrokOAuthService(proxyRepo ProxyRepository, oauthClient GrokOAuthClient, return service } +// WithRedisSessionStore enables cross-instance, single-use OAuth callbacks. +func (s *GrokOAuthService) WithRedisSessionStore(rdb *redis.Client) *GrokOAuthService { + if s != nil && rdb != nil { + if s.sessionStore != nil { + s.sessionStore.Stop() + } + s.sessionStore = xai.NewRedisSessionStore(rdb) + } + return s +} + type GrokOAuthCapabilities struct { PasswordAuthEnabled bool `json:"password_auth_enabled"` } @@ -139,7 +151,6 @@ func (s *GrokOAuthService) ExchangeCode(ctx context.Context, input *GrokExchange if !ok { return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_SESSION_NOT_FOUND", "session not found or expired") } - defer s.sessionStore.Delete(input.SessionID) parsed := xai.ParseAuthorizationInput(input.Code) code := strings.TrimSpace(parsed.Code) @@ -165,6 +176,13 @@ func (s *GrokOAuthService) ExchangeCode(ctx context.Context, input *GrokExchange return nil, err } } + if s.oauthClient == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "GROK_OAUTH_CLIENT_NOT_CONFIGURED", "oauth client is not configured") + } + if !s.sessionStore.TryConsumeSession(input.SessionID) { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_SESSION_ALREADY_USED", "oauth session has already been used") + } + defer s.sessionStore.Delete(input.SessionID) redirectURI := session.RedirectURI if strings.TrimSpace(input.RedirectURI) != "" { redirectURI = input.RedirectURI diff --git a/backend/internal/service/grok_oauth_service_test.go b/backend/internal/service/grok_oauth_service_test.go index 1b4846673..723cfe2c9 100644 --- a/backend/internal/service/grok_oauth_service_test.go +++ b/backend/internal/service/grok_oauth_service_test.go @@ -59,7 +59,7 @@ func TestGrokOAuthServiceRefreshTokenPreservesOriginalRefreshTokenWhenNotRotated require.Equal(t, "client-id", info.ClientID) } -func TestGrokOAuthServiceExchangeCodeRequiresStateForCallbackURLAndConsumesSession(t *testing.T) { +func TestGrokOAuthServiceExchangeCodeConsumesOnlyAfterValidation(t *testing.T) { client := &grokOAuthClientStub{} svc := NewGrokOAuthService(nil, client) defer svc.Stop() @@ -80,9 +80,34 @@ func TestGrokOAuthServiceExchangeCodeRequiresStateForCallbackURLAndConsumesSessi Code: "code-with-state", State: auth.State, }) + require.NoError(t, err) + require.Equal(t, 1, client.exchangeCalls) + + _, err = svc.ExchangeCode(context.Background(), &GrokExchangeCodeInput{ + SessionID: auth.SessionID, + Code: "replayed-code", + State: auth.State, + }) require.Error(t, err) require.Contains(t, err.Error(), "GROK_OAUTH_SESSION_NOT_FOUND") - require.Zero(t, client.exchangeCalls) + require.Equal(t, 1, client.exchangeCalls) +} + +func TestGrokOAuthServiceExchangeCodeRejectsMissingClientWithoutConsumingSession(t *testing.T) { + svc := NewGrokOAuthService(nil, nil) + defer svc.Stop() + auth, err := svc.GenerateAuthURL(context.Background(), nil, "") + require.NoError(t, err) + + _, err = svc.ExchangeCode(context.Background(), &GrokExchangeCodeInput{ + SessionID: auth.SessionID, + Code: "code", + State: auth.State, + }) + require.Error(t, err) + require.Contains(t, err.Error(), "GROK_OAUTH_CLIENT_NOT_CONFIGURED") + _, ok := svc.sessionStore.Get(auth.SessionID) + require.True(t, ok) } func TestGrokOAuthServiceBuildAccountCredentialsDefaultsToSubscriptionProxy(t *testing.T) { diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index 0e3cdfb99..4f6042589 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -15,8 +15,8 @@ import ( "go.uber.org/zap" ) -func ProvideGrokOAuthService(proxyRepo ProxyRepository, oauthClient GrokOAuthClient, cfg *config.Config) *GrokOAuthService { - return NewGrokOAuthService(proxyRepo, oauthClient, cfg) +func ProvideGrokOAuthService(proxyRepo ProxyRepository, oauthClient GrokOAuthClient, cfg *config.Config, redisClient *redis.Client) *GrokOAuthService { + return NewGrokOAuthService(proxyRepo, oauthClient, cfg).WithRedisSessionStore(redisClient) } // BuildInfo contains build information