fix(grok): 加固 OAuth 会话共享与一次性消费

This commit is contained in:
IanShaw027
2026-08-07 16:29:48 +08:00
parent a061a758f4
commit 25d2b03e90
9 changed files with 371 additions and 13 deletions
+1 -1
View File
@@ -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)
+2
View File
@@ -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=
+125
View File
@@ -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
}
@@ -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)
}
+120 -7
View File
@@ -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()
@@ -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"))
}
+19 -1
View File
@@ -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
@@ -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) {
+2 -2
View File
@@ -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