mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
feat(grok): 支持 Web SSO 批量导入并转换为 Build OAuth
新增 Grok Web SSO → xAI Device Flow → Grok Build OAuth 导入链路, 支持管理员批量粘贴 SSO key 创建 OAuth 账号。 - 后端:ConvertSSOToBuild、ConvertFromSSO、POST /admin/grok/sso-to-oauth - 批量:3 worker 并发,失败跳过并汇总 created/failed,worker panic recover - 无 refresh_token 时写入 expires_at 并强制 auto_pause_on_expired - 前端:SSO Cookie 导入入口、动态超时、中英文案、部分成功不关弹窗 - 测试:pkg/service/handler/前端超时单测;本地 Docker 真实 SSO e2e 通过
This commit is contained in:
@@ -1,16 +1,23 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/handler/dto"
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const grokSSOImportConcurrency = 3
|
||||
|
||||
type GrokOAuthHandler struct {
|
||||
grokOAuthService *service.GrokOAuthService
|
||||
adminService service.AdminService
|
||||
@@ -205,6 +212,238 @@ func (h *GrokOAuthHandler) CreateAccountFromOAuth(c *gin.Context) {
|
||||
response.Success(c, dto.AccountFromService(account))
|
||||
}
|
||||
|
||||
type GrokSSOToOAuthRequest struct {
|
||||
SSOTokens []string `json:"sso_tokens"`
|
||||
SSOToken string `json:"sso_token"`
|
||||
Name string `json:"name"`
|
||||
Notes *string `json:"notes"`
|
||||
ProxyID *int64 `json:"proxy_id"`
|
||||
GroupIDs []int64 `json:"group_ids"`
|
||||
Credentials map[string]any `json:"credentials"`
|
||||
Extra map[string]any `json:"extra"`
|
||||
Concurrency int `json:"concurrency"`
|
||||
LoadFactor *int `json:"load_factor"`
|
||||
Priority int `json:"priority"`
|
||||
RateMultiplier *float64 `json:"rate_multiplier"`
|
||||
ExpiresAt *int64 `json:"expires_at"`
|
||||
AutoPauseOnExpired *bool `json:"auto_pause_on_expired"`
|
||||
}
|
||||
|
||||
type GrokSSOToOAuthItemResult struct {
|
||||
Index int `json:"index"`
|
||||
Name string `json:"name,omitempty"`
|
||||
Email string `json:"email,omitempty"`
|
||||
Account *dto.Account `json:"account,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type GrokSSOToOAuthResponse struct {
|
||||
Created []GrokSSOToOAuthItemResult `json:"created"`
|
||||
Failed []GrokSSOToOAuthItemResult `json:"failed"`
|
||||
}
|
||||
|
||||
type grokSSOImportJob struct {
|
||||
index int
|
||||
token string
|
||||
}
|
||||
|
||||
type grokSSOImportWorkerResult struct {
|
||||
created bool
|
||||
item GrokSSOToOAuthItemResult
|
||||
}
|
||||
|
||||
func (h *GrokOAuthHandler) CreateAccountsFromSSO(c *gin.Context) {
|
||||
var req GrokSSOToOAuthRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.BadRequest(c, "Invalid request: "+err.Error())
|
||||
return
|
||||
}
|
||||
tokens := normalizeSSOImportTokens(req.SSOTokens, req.SSOToken)
|
||||
if len(tokens) == 0 {
|
||||
response.BadRequest(c, "sso_tokens is required")
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
workerCount := grokSSOImportConcurrency
|
||||
if len(tokens) < workerCount {
|
||||
workerCount = len(tokens)
|
||||
}
|
||||
jobs := make(chan grokSSOImportJob)
|
||||
items := make([]grokSSOImportWorkerResult, len(tokens))
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < workerCount; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for job := range jobs {
|
||||
items[job.index] = h.safeCreateAccountFromSSOToken(ctx, req, job.token, job.index+1, len(tokens))
|
||||
}
|
||||
}()
|
||||
}
|
||||
for i, token := range tokens {
|
||||
jobs <- grokSSOImportJob{index: i, token: token}
|
||||
}
|
||||
close(jobs)
|
||||
wg.Wait()
|
||||
|
||||
result := GrokSSOToOAuthResponse{
|
||||
Created: make([]GrokSSOToOAuthItemResult, 0, len(tokens)),
|
||||
Failed: make([]GrokSSOToOAuthItemResult, 0),
|
||||
}
|
||||
for _, item := range items {
|
||||
if item.created {
|
||||
result.Created = append(result.Created, item.item)
|
||||
} else {
|
||||
result.Failed = append(result.Failed, item.item)
|
||||
}
|
||||
}
|
||||
response.Success(c, result)
|
||||
}
|
||||
|
||||
func (h *GrokOAuthHandler) safeCreateAccountFromSSOToken(ctx context.Context, req GrokSSOToOAuthRequest, token string, index, total int) (result grokSSOImportWorkerResult) {
|
||||
defer func() {
|
||||
if recovered := recover(); recovered != nil {
|
||||
slog.Error("grok_sso_import_worker_panic", "index", index, "recover", recovered)
|
||||
result = grokSSOImportWorkerResult{
|
||||
item: GrokSSOToOAuthItemResult{
|
||||
Index: index,
|
||||
Error: fmt.Sprintf("internal worker panic: %v", recovered),
|
||||
},
|
||||
}
|
||||
}
|
||||
}()
|
||||
return h.createAccountFromSSOToken(ctx, req, token, index, total)
|
||||
}
|
||||
|
||||
func (h *GrokOAuthHandler) createAccountFromSSOToken(ctx context.Context, req GrokSSOToOAuthRequest, token string, index, total int) grokSSOImportWorkerResult {
|
||||
tokenInfo, err := h.grokOAuthService.ConvertFromSSO(ctx, token, req.ProxyID)
|
||||
if err != nil {
|
||||
return grokSSOImportWorkerResult{item: GrokSSOToOAuthItemResult{Index: index, Error: grokSSOImportErrorMessage(err)}}
|
||||
}
|
||||
|
||||
credentials := h.grokOAuthService.BuildAccountCredentials(tokenInfo)
|
||||
credentials = service.MergeCredentials(cloneGrokSSOMap(req.Credentials), credentials)
|
||||
name := grokSSOImportAccountName(req.Name, tokenInfo, index, total)
|
||||
expiresAt, autoPauseOnExpired := grokSSOImportExpiry(req.ExpiresAt, req.AutoPauseOnExpired, tokenInfo)
|
||||
account, err := h.adminService.CreateAccount(ctx, &service.CreateAccountInput{
|
||||
Name: name,
|
||||
Notes: req.Notes,
|
||||
Platform: service.PlatformGrok,
|
||||
Type: service.AccountTypeOAuth,
|
||||
Credentials: credentials,
|
||||
Extra: cloneGrokSSOMap(req.Extra),
|
||||
ProxyID: req.ProxyID,
|
||||
Concurrency: req.Concurrency,
|
||||
LoadFactor: req.LoadFactor,
|
||||
Priority: req.Priority,
|
||||
RateMultiplier: req.RateMultiplier,
|
||||
GroupIDs: append([]int64(nil), req.GroupIDs...),
|
||||
ExpiresAt: expiresAt,
|
||||
AutoPauseOnExpired: autoPauseOnExpired,
|
||||
})
|
||||
if err != nil {
|
||||
return grokSSOImportWorkerResult{item: GrokSSOToOAuthItemResult{Index: index, Name: name, Email: tokenInfo.Email, Error: grokSSOImportErrorMessage(err)}}
|
||||
}
|
||||
return grokSSOImportWorkerResult{
|
||||
created: true,
|
||||
item: GrokSSOToOAuthItemResult{
|
||||
Index: index,
|
||||
Name: name,
|
||||
Email: tokenInfo.Email,
|
||||
Account: dto.AccountFromService(account),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func grokSSOImportExpiry(requestExpiresAt *int64, requestAutoPause *bool, tokenInfo *service.GrokTokenInfo) (*int64, *bool) {
|
||||
if tokenInfo == nil || strings.TrimSpace(tokenInfo.RefreshToken) != "" || tokenInfo.ExpiresAt <= 0 {
|
||||
return requestExpiresAt, requestAutoPause
|
||||
}
|
||||
|
||||
expiresAt := tokenInfo.ExpiresAt
|
||||
if requestExpiresAt != nil && *requestExpiresAt > 0 && *requestExpiresAt < expiresAt {
|
||||
expiresAt = *requestExpiresAt
|
||||
}
|
||||
autoPause := true
|
||||
return &expiresAt, &autoPause
|
||||
}
|
||||
|
||||
func cloneGrokSSOMap(source map[string]any) map[string]any {
|
||||
if source == nil {
|
||||
return nil
|
||||
}
|
||||
clone := make(map[string]any, len(source))
|
||||
for key, value := range source {
|
||||
clone[key] = cloneGrokSSOValue(value)
|
||||
}
|
||||
return clone
|
||||
}
|
||||
|
||||
func cloneGrokSSOValue(value any) any {
|
||||
switch v := value.(type) {
|
||||
case map[string]any:
|
||||
return cloneGrokSSOMap(v)
|
||||
case []any:
|
||||
clone := make([]any, len(v))
|
||||
for i, item := range v {
|
||||
clone[i] = cloneGrokSSOValue(item)
|
||||
}
|
||||
return clone
|
||||
default:
|
||||
return value
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeSSOImportTokens(tokens []string, single string) []string {
|
||||
items := make([]string, 0, len(tokens)+1)
|
||||
if strings.TrimSpace(single) != "" {
|
||||
items = append(items, single)
|
||||
}
|
||||
items = append(items, tokens...)
|
||||
seen := make(map[string]struct{}, len(items))
|
||||
result := make([]string, 0, len(items))
|
||||
for _, item := range items {
|
||||
parts := strings.Split(strings.NewReplacer(",", "\n", "\r", "\n").Replace(item), "\n")
|
||||
for _, token := range parts {
|
||||
if token = xai.NormalizeSSOToken(token); token == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[token]; ok {
|
||||
continue
|
||||
}
|
||||
seen[token] = struct{}{}
|
||||
result = append(result, token)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func grokSSOImportAccountName(base string, tokenInfo *service.GrokTokenInfo, index, total int) string {
|
||||
base = strings.TrimSpace(base)
|
||||
if base == "" && tokenInfo != nil {
|
||||
base = strings.TrimSpace(tokenInfo.Email)
|
||||
}
|
||||
if base == "" {
|
||||
base = "Grok OAuth Account"
|
||||
}
|
||||
if total > 1 {
|
||||
return base + " #" + strconv.Itoa(index)
|
||||
}
|
||||
return base
|
||||
}
|
||||
|
||||
func grokSSOImportErrorMessage(err error) string {
|
||||
status := infraerrors.FromError(err)
|
||||
if status == nil {
|
||||
return ""
|
||||
}
|
||||
if status.Reason != "" {
|
||||
return status.Reason + ": " + status.Message
|
||||
}
|
||||
return status.Message
|
||||
}
|
||||
|
||||
func (h *GrokOAuthHandler) QueryQuota(c *gin.Context) {
|
||||
accountID, err := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
|
||||
@@ -145,3 +145,51 @@ func TestGrokOAuthHandlerRuntimeSanityDoesNotExposeSecrets(t *testing.T) {
|
||||
require.NotContains(t, rec.Body.String(), "secret")
|
||||
require.NotContains(t, rec.Body.String(), "client-secret-like-value")
|
||||
}
|
||||
|
||||
func TestGrokSSOImportExpiryUsesTokenExpiryWithoutRefreshToken(t *testing.T) {
|
||||
tokenExpiry := time.Now().Add(6 * time.Hour).Unix()
|
||||
expiresAt, autoPause := grokSSOImportExpiry(nil, nil, &service.GrokTokenInfo{
|
||||
ExpiresAt: tokenExpiry,
|
||||
})
|
||||
|
||||
require.NotNil(t, expiresAt)
|
||||
require.Equal(t, tokenExpiry, *expiresAt)
|
||||
require.NotNil(t, autoPause)
|
||||
require.True(t, *autoPause)
|
||||
}
|
||||
|
||||
func TestGrokSSOImportExpiryUsesEarlierRequestedExpiryWithoutRefreshToken(t *testing.T) {
|
||||
requestedExpiry := time.Now().Add(2 * time.Hour).Unix()
|
||||
tokenExpiry := time.Now().Add(6 * time.Hour).Unix()
|
||||
requestedAutoPause := false
|
||||
expiresAt, autoPause := grokSSOImportExpiry(&requestedExpiry, &requestedAutoPause, &service.GrokTokenInfo{
|
||||
ExpiresAt: tokenExpiry,
|
||||
})
|
||||
|
||||
require.NotNil(t, expiresAt)
|
||||
require.Equal(t, requestedExpiry, *expiresAt)
|
||||
require.NotNil(t, autoPause)
|
||||
require.True(t, *autoPause)
|
||||
}
|
||||
|
||||
func TestGrokSSOImportExpiryPreservesRequestSettingsWithRefreshToken(t *testing.T) {
|
||||
requestedExpiry := time.Now().Add(2 * time.Hour).Unix()
|
||||
requestedAutoPause := false
|
||||
expiresAt, autoPause := grokSSOImportExpiry(&requestedExpiry, &requestedAutoPause, &service.GrokTokenInfo{
|
||||
RefreshToken: "refresh-token",
|
||||
ExpiresAt: time.Now().Add(6 * time.Hour).Unix(),
|
||||
})
|
||||
|
||||
require.Same(t, &requestedExpiry, expiresAt)
|
||||
require.Same(t, &requestedAutoPause, autoPause)
|
||||
}
|
||||
|
||||
func TestGrokSSOImportWorkerRecoversPanic(t *testing.T) {
|
||||
h := &GrokOAuthHandler{}
|
||||
result := h.safeCreateAccountFromSSOToken(context.Background(), GrokSSOToOAuthRequest{}, "token", 2, 3)
|
||||
// Without a service, createAccountFromSSOToken would panic on nil service access.
|
||||
// Recovery must convert that into a failed item and keep the worker alive.
|
||||
require.False(t, result.created)
|
||||
require.Equal(t, 2, result.item.Index)
|
||||
require.Contains(t, result.item.Error, "internal worker panic")
|
||||
}
|
||||
|
||||
@@ -0,0 +1,418 @@
|
||||
package xai
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
SSOBuildScope = "openid profile email offline_access grok-cli:access api:access conversations:read conversations:write"
|
||||
SSOAccountsURL = "https://accounts.x.ai/"
|
||||
SSODeviceURL = OAuthIssuer + "/oauth2/device/code"
|
||||
SSOVerifyURL = OAuthIssuer + "/oauth2/device/verify"
|
||||
SSOApproveURL = OAuthIssuer + "/oauth2/device/approve"
|
||||
SSOTokenURL = OAuthIssuer + "/oauth2/token"
|
||||
SSOConversionTimeout = 90 * time.Second
|
||||
|
||||
ssoMaxAuthBody = 2 << 20
|
||||
ssoDefaultUA = "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36"
|
||||
ssoDefaultTokenTTL = 6 * time.Hour
|
||||
)
|
||||
|
||||
var (
|
||||
ErrSSOUnauthorized = errors.New("xai sso unauthorized")
|
||||
ErrSSOAuthorizationDenied = errors.New("xai device authorization denied")
|
||||
)
|
||||
|
||||
type SSOHTTPError struct{ Status int }
|
||||
|
||||
func (e SSOHTTPError) Error() string { return fmt.Sprintf("xAI OAuth HTTP %d", e.Status) }
|
||||
|
||||
type SSODeviceHTTPClient interface {
|
||||
Do(*http.Request) (*http.Response, error)
|
||||
}
|
||||
|
||||
type SSODeviceOptions struct {
|
||||
HTTPClient SSODeviceHTTPClient
|
||||
UserAgent string
|
||||
Sleep func(context.Context, time.Duration) error
|
||||
}
|
||||
|
||||
type ssoDeviceFlow struct {
|
||||
client SSODeviceHTTPClient
|
||||
userAgent string
|
||||
cookies map[string]string
|
||||
sleep func(context.Context, time.Duration) error
|
||||
}
|
||||
|
||||
func ConvertSSOToBuild(ctx context.Context, ssoToken string, opts *SSODeviceOptions) (*TokenResponse, error) {
|
||||
ssoToken = NormalizeSSOToken(ssoToken)
|
||||
if ssoToken == "" {
|
||||
return nil, ErrSSOUnauthorized
|
||||
}
|
||||
if opts == nil {
|
||||
opts = &SSODeviceOptions{}
|
||||
}
|
||||
client := opts.HTTPClient
|
||||
if client == nil {
|
||||
client = &http.Client{
|
||||
Timeout: SSOConversionTimeout,
|
||||
CheckRedirect: func(*http.Request, []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
},
|
||||
}
|
||||
}
|
||||
userAgent := strings.TrimSpace(opts.UserAgent)
|
||||
if userAgent == "" {
|
||||
userAgent = ssoDefaultUA
|
||||
}
|
||||
sleep := opts.Sleep
|
||||
if sleep == nil {
|
||||
sleep = sleepContext
|
||||
}
|
||||
|
||||
flow := &ssoDeviceFlow{
|
||||
client: client,
|
||||
userAgent: userAgent,
|
||||
cookies: map[string]string{"sso": ssoToken, "sso-rw": ssoToken},
|
||||
sleep: sleep,
|
||||
}
|
||||
return flow.convert(ctx)
|
||||
}
|
||||
|
||||
func (f *ssoDeviceFlow) convert(ctx context.Context) (*TokenResponse, error) {
|
||||
status, finalURL, _, err := f.do(ctx, http.MethodGet, SSOAccountsURL, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if status == http.StatusUnauthorized || strings.Contains(finalURL, "sign-in") || strings.Contains(finalURL, "sign-up") {
|
||||
return nil, ErrSSOUnauthorized
|
||||
}
|
||||
if status < 200 || status >= 400 {
|
||||
return nil, fmt.Errorf("validate Grok Web SSO: %w", SSOHTTPError{Status: status})
|
||||
}
|
||||
|
||||
status, _, body, err := f.do(ctx, http.MethodPost, SSODeviceURL, url.Values{
|
||||
"client_id": {DefaultClientID},
|
||||
"scope": {SSOBuildScope},
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if status < 200 || status >= 300 {
|
||||
return nil, fmt.Errorf("start xAI device flow: %w", SSOHTTPError{Status: status})
|
||||
}
|
||||
var device struct {
|
||||
DeviceCode string `json:"device_code"`
|
||||
UserCode string `json:"user_code"`
|
||||
VerificationURIComplete string `json:"verification_uri_complete"`
|
||||
Interval int `json:"interval"`
|
||||
ExpiresIn int `json:"expires_in"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &device); err != nil {
|
||||
return nil, fmt.Errorf("parse xAI device flow response: %w", err)
|
||||
}
|
||||
if device.DeviceCode == "" || device.UserCode == "" || !safeXAIAuthURL(device.VerificationURIComplete) {
|
||||
return nil, errors.New("xAI device flow response is incomplete")
|
||||
}
|
||||
if device.Interval <= 0 {
|
||||
device.Interval = 5
|
||||
}
|
||||
if device.ExpiresIn <= 0 {
|
||||
device.ExpiresIn = 1800
|
||||
}
|
||||
|
||||
status, _, _, err = f.do(ctx, http.MethodGet, device.VerificationURIComplete, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if status < 200 || status >= 400 {
|
||||
return nil, fmt.Errorf("open xAI device verification page: %w", SSOHTTPError{Status: status})
|
||||
}
|
||||
|
||||
status, finalURL, _, err = f.do(ctx, http.MethodPost, SSOVerifyURL, url.Values{"user_code": {device.UserCode}})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if status < 200 || status >= 400 {
|
||||
return nil, fmt.Errorf("verify xAI device code: %w", SSOHTTPError{Status: status})
|
||||
}
|
||||
if !strings.Contains(finalURL, "consent") {
|
||||
return nil, errors.New("xAI device verification did not reach consent page")
|
||||
}
|
||||
|
||||
status, finalURL, _, err = f.do(ctx, http.MethodPost, SSOApproveURL, url.Values{
|
||||
"user_code": {device.UserCode},
|
||||
"action": {"allow"},
|
||||
"principal_type": {"User"},
|
||||
"principal_id": {""},
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if status < 200 || status >= 400 {
|
||||
return nil, fmt.Errorf("approve xAI device code: %w", SSOHTTPError{Status: status})
|
||||
}
|
||||
if !strings.Contains(finalURL, "done") {
|
||||
return nil, errors.New("xAI device approval did not reach done page")
|
||||
}
|
||||
|
||||
return f.pollToken(ctx, device.DeviceCode, time.Duration(device.Interval)*time.Second, time.Duration(device.ExpiresIn)*time.Second)
|
||||
}
|
||||
|
||||
func (f *ssoDeviceFlow) pollToken(ctx context.Context, deviceCode string, interval, expiresIn time.Duration) (*TokenResponse, error) {
|
||||
if interval < time.Second {
|
||||
interval = time.Second
|
||||
}
|
||||
deadline := time.Now().Add(minDuration(expiresIn, 75*time.Second))
|
||||
for time.Now().Before(deadline) {
|
||||
if err := f.sleep(ctx, interval); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
status, _, body, err := f.do(ctx, http.MethodPost, SSOTokenURL, url.Values{
|
||||
"grant_type": {"urn:ietf:params:oauth:grant-type:device_code"},
|
||||
"client_id": {DefaultClientID},
|
||||
"device_code": {deviceCode},
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var payload struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
RefreshToken string `json:"refresh_token"`
|
||||
IDToken string `json:"id_token"`
|
||||
TokenType string `json:"token_type"`
|
||||
ExpiresIn int64 `json:"expires_in"`
|
||||
Scope string `json:"scope"`
|
||||
Error string `json:"error"`
|
||||
ErrorDescription string `json:"error_description"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &payload); err != nil {
|
||||
return nil, fmt.Errorf("parse xAI token response: %w", err)
|
||||
}
|
||||
if status >= 200 && status < 300 && payload.AccessToken != "" {
|
||||
if payload.ExpiresIn <= 0 {
|
||||
payload.ExpiresIn = int64(ssoDefaultTokenTTL.Seconds())
|
||||
}
|
||||
if payload.TokenType == "" {
|
||||
payload.TokenType = "Bearer"
|
||||
}
|
||||
return &TokenResponse{
|
||||
AccessToken: payload.AccessToken,
|
||||
RefreshToken: payload.RefreshToken,
|
||||
IDToken: payload.IDToken,
|
||||
TokenType: payload.TokenType,
|
||||
ExpiresIn: payload.ExpiresIn,
|
||||
Scope: payload.Scope,
|
||||
}, nil
|
||||
}
|
||||
switch payload.Error {
|
||||
case "authorization_pending":
|
||||
continue
|
||||
case "slow_down":
|
||||
interval += 5 * time.Second
|
||||
continue
|
||||
case "access_denied", "expired_token":
|
||||
return nil, ErrSSOAuthorizationDenied
|
||||
default:
|
||||
if status >= 400 {
|
||||
return nil, fmt.Errorf("xAI token polling failed (%s): %w", firstNonEmpty(payload.ErrorDescription, payload.Error), SSOHTTPError{Status: status})
|
||||
}
|
||||
return nil, fmt.Errorf("xAI token polling failed: %s", firstNonEmpty(payload.ErrorDescription, payload.Error, strconv.Itoa(status)))
|
||||
}
|
||||
}
|
||||
return nil, errors.New("xAI device flow token polling timed out")
|
||||
}
|
||||
|
||||
func (f *ssoDeviceFlow) do(ctx context.Context, method, endpoint string, form url.Values) (int, string, []byte, error) {
|
||||
if !safeXAIAuthURL(endpoint) {
|
||||
return 0, "", nil, errors.New("xAI OAuth URL is not trusted")
|
||||
}
|
||||
currentURL := endpoint
|
||||
currentMethod := method
|
||||
currentForm := form
|
||||
for redirects := 0; redirects <= 8; redirects++ {
|
||||
var body io.Reader
|
||||
if currentForm != nil {
|
||||
body = strings.NewReader(currentForm.Encode())
|
||||
}
|
||||
request, err := http.NewRequestWithContext(ctx, currentMethod, currentURL, body)
|
||||
if err != nil {
|
||||
return 0, currentURL, nil, err
|
||||
}
|
||||
request.Header.Set("Accept", "application/json, text/html;q=0.9, */*;q=0.8")
|
||||
request.Header.Set("Accept-Language", "zh-CN,zh;q=0.9,en;q=0.8")
|
||||
request.Header.Set("User-Agent", f.userAgent)
|
||||
if cookie := f.cookieHeader(); cookie != "" {
|
||||
request.Header.Set("Cookie", cookie)
|
||||
}
|
||||
if currentForm != nil {
|
||||
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
}
|
||||
|
||||
response, err := f.client.Do(request)
|
||||
if err != nil {
|
||||
return 0, currentURL, nil, err
|
||||
}
|
||||
f.captureCookies(response)
|
||||
data, readErr := io.ReadAll(io.LimitReader(response.Body, ssoMaxAuthBody+1))
|
||||
_ = response.Body.Close()
|
||||
if readErr != nil {
|
||||
return response.StatusCode, currentURL, nil, readErr
|
||||
}
|
||||
if len(data) > ssoMaxAuthBody {
|
||||
return response.StatusCode, currentURL, nil, errors.New("xAI OAuth response exceeds 2 MiB")
|
||||
}
|
||||
if response.StatusCode < 300 || response.StatusCode > 399 {
|
||||
return response.StatusCode, currentURL, data, nil
|
||||
}
|
||||
|
||||
location := strings.TrimSpace(response.Header.Get("Location"))
|
||||
if location == "" {
|
||||
return response.StatusCode, currentURL, data, errors.New("xAI OAuth redirect missing Location")
|
||||
}
|
||||
base, _ := url.Parse(currentURL)
|
||||
next, err := url.Parse(location)
|
||||
if err != nil {
|
||||
return response.StatusCode, currentURL, data, err
|
||||
}
|
||||
currentURL = base.ResolveReference(next).String()
|
||||
if !safeXAIAuthURL(currentURL) {
|
||||
return response.StatusCode, currentURL, data, errors.New("xAI OAuth redirected to untrusted host")
|
||||
}
|
||||
if response.StatusCode == http.StatusSeeOther || ((response.StatusCode == http.StatusMovedPermanently || response.StatusCode == http.StatusFound) && currentMethod != http.MethodGet && currentMethod != http.MethodHead) {
|
||||
currentMethod = http.MethodGet
|
||||
currentForm = nil
|
||||
}
|
||||
}
|
||||
return 0, currentURL, nil, errors.New("xAI OAuth redirected too many times")
|
||||
}
|
||||
|
||||
func (f *ssoDeviceFlow) captureCookies(response *http.Response) {
|
||||
for _, cookie := range response.Cookies() {
|
||||
name := strings.TrimSpace(cookie.Name)
|
||||
value := strings.TrimSpace(cookie.Value)
|
||||
if name == "" || len(name) > 128 || len(value) > 16384 || strings.ContainsAny(name+value, "\r\n\x00") {
|
||||
continue
|
||||
}
|
||||
if cookie.MaxAge < 0 {
|
||||
delete(f.cookies, name)
|
||||
continue
|
||||
}
|
||||
f.cookies[name] = value
|
||||
}
|
||||
}
|
||||
|
||||
func (f *ssoDeviceFlow) cookieHeader() string {
|
||||
keys := make([]string, 0, len(f.cookies))
|
||||
for key := range f.cookies {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
parts := make([]string, 0, len(keys))
|
||||
for _, key := range keys {
|
||||
parts = append(parts, key+"="+f.cookies[key])
|
||||
}
|
||||
return strings.Join(parts, "; ")
|
||||
}
|
||||
|
||||
func safeXAIAuthURL(raw string) bool {
|
||||
parsed, err := url.Parse(raw)
|
||||
if err != nil || parsed.User != nil || parsed.Hostname() == "" {
|
||||
return false
|
||||
}
|
||||
if AllowUnsafeURLOverrides() {
|
||||
return parsed.Scheme != "" && parsed.Host != ""
|
||||
}
|
||||
if parsed.Scheme != "https" {
|
||||
return false
|
||||
}
|
||||
host := strings.ToLower(parsed.Hostname())
|
||||
return host == "x.ai" || strings.HasSuffix(host, ".x.ai")
|
||||
}
|
||||
|
||||
func NormalizeSSOToken(value string) string {
|
||||
value = strings.TrimSpace(value)
|
||||
if strings.HasPrefix(strings.ToLower(value), "cookie:") {
|
||||
value = strings.TrimSpace(value[len("cookie:"):])
|
||||
}
|
||||
for _, part := range strings.Split(value, ";") {
|
||||
name, token, found := strings.Cut(strings.TrimSpace(part), "=")
|
||||
if !found {
|
||||
continue
|
||||
}
|
||||
switch strings.ToLower(strings.TrimSpace(name)) {
|
||||
case "sso", "sso-rw":
|
||||
return sanitizeSSOToken(token)
|
||||
}
|
||||
}
|
||||
if token, _, found := strings.Cut(value, ";"); found {
|
||||
value = strings.TrimSpace(token)
|
||||
}
|
||||
return sanitizeSSOToken(value)
|
||||
}
|
||||
|
||||
func sanitizeSSOToken(value string) string {
|
||||
return strings.NewReplacer("\r", "", "\n", "", "\x00", "").Replace(strings.TrimSpace(value))
|
||||
}
|
||||
|
||||
func DecodeJWTClaims(token string) map[string]any {
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) < 2 {
|
||||
return nil
|
||||
}
|
||||
payload, err := base64.RawURLEncoding.DecodeString(parts[1])
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
var claims map[string]any
|
||||
if err := json.Unmarshal(payload, &claims); err != nil {
|
||||
return nil
|
||||
}
|
||||
return claims
|
||||
}
|
||||
|
||||
func JWTClaimString(claims map[string]any, key string) string {
|
||||
value, _ := claims[key].(string)
|
||||
return strings.TrimSpace(value)
|
||||
}
|
||||
|
||||
func sleepContext(ctx context.Context, d time.Duration) error {
|
||||
timer := time.NewTimer(d)
|
||||
defer timer.Stop()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-timer.C:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func minDuration(a, b time.Duration) time.Duration {
|
||||
if a <= 0 {
|
||||
return b
|
||||
}
|
||||
if a < b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func firstNonEmpty(values ...string) string {
|
||||
for _, value := range values {
|
||||
if strings.TrimSpace(value) != "" {
|
||||
return strings.TrimSpace(value)
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,115 @@
|
||||
//go:build unit
|
||||
|
||||
package xai
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type ssoDeviceFakeClient struct {
|
||||
t *testing.T
|
||||
tokenCalls int
|
||||
cookieHeaders []string
|
||||
}
|
||||
|
||||
func (c *ssoDeviceFakeClient) Do(req *http.Request) (*http.Response, error) {
|
||||
c.cookieHeaders = append(c.cookieHeaders, req.Header.Get("Cookie"))
|
||||
switch req.URL.String() {
|
||||
case SSOAccountsURL:
|
||||
require.Equal(c.t, http.MethodGet, req.Method)
|
||||
return ssoDeviceResponse(http.StatusOK, http.Header{"Set-Cookie": {"session=web-session; Path=/"}}, `{}`), nil
|
||||
case SSODeviceURL:
|
||||
require.Equal(c.t, http.MethodPost, req.Method)
|
||||
values := readSSODeviceForm(c.t, req)
|
||||
require.Equal(c.t, DefaultClientID, values.Get("client_id"))
|
||||
require.Equal(c.t, SSOBuildScope, values.Get("scope"))
|
||||
return ssoDeviceResponse(http.StatusOK, http.Header{"Set-Cookie": {"csrf=csrf-token; Path=/"}}, `{"device_code":"device-1","user_code":"USER-1","verification_uri_complete":"https://auth.x.ai/oauth2/device/complete","interval":1,"expires_in":60}`), nil
|
||||
case "https://auth.x.ai/oauth2/device/complete":
|
||||
require.Equal(c.t, http.MethodGet, req.Method)
|
||||
return ssoDeviceResponse(http.StatusOK, nil, `<html>ok</html>`), nil
|
||||
case SSOVerifyURL:
|
||||
require.Equal(c.t, http.MethodPost, req.Method)
|
||||
values := readSSODeviceForm(c.t, req)
|
||||
require.Equal(c.t, "USER-1", values.Get("user_code"))
|
||||
return ssoDeviceResponse(http.StatusFound, http.Header{"Location": {"/oauth2/device/consent"}}, ``), nil
|
||||
case "https://auth.x.ai/oauth2/device/consent":
|
||||
require.Equal(c.t, http.MethodGet, req.Method)
|
||||
return ssoDeviceResponse(http.StatusOK, nil, `<html>consent</html>`), nil
|
||||
case SSOApproveURL:
|
||||
require.Equal(c.t, http.MethodPost, req.Method)
|
||||
values := readSSODeviceForm(c.t, req)
|
||||
require.Equal(c.t, "USER-1", values.Get("user_code"))
|
||||
require.Equal(c.t, "allow", values.Get("action"))
|
||||
require.Equal(c.t, "User", values.Get("principal_type"))
|
||||
return ssoDeviceResponse(http.StatusSeeOther, http.Header{"Location": {"/oauth2/device/done"}}, ``), nil
|
||||
case "https://auth.x.ai/oauth2/device/done":
|
||||
require.Equal(c.t, http.MethodGet, req.Method)
|
||||
return ssoDeviceResponse(http.StatusOK, nil, `<html>done</html>`), nil
|
||||
case SSOTokenURL:
|
||||
require.Equal(c.t, http.MethodPost, req.Method)
|
||||
c.tokenCalls++
|
||||
values := readSSODeviceForm(c.t, req)
|
||||
require.Equal(c.t, "urn:ietf:params:oauth:grant-type:device_code", values.Get("grant_type"))
|
||||
require.Equal(c.t, "device-1", values.Get("device_code"))
|
||||
return ssoDeviceResponse(http.StatusOK, nil, `{"access_token":"access-token","refresh_token":"refresh-token","id_token":"id-token","token_type":"Bearer","expires_in":3600,"scope":"`+SSOBuildScope+`"}`), nil
|
||||
default:
|
||||
c.t.Fatalf("unexpected request: %s %s", req.Method, req.URL.String())
|
||||
return nil, nil
|
||||
}
|
||||
}
|
||||
|
||||
func TestConvertSSOToBuildCompletesDeviceFlow(t *testing.T) {
|
||||
t.Setenv(EnvClientID, "")
|
||||
client := &ssoDeviceFakeClient{t: t}
|
||||
token, err := ConvertSSOToBuild(context.Background(), "sso=sso-token; ignored=1", &SSODeviceOptions{
|
||||
HTTPClient: client,
|
||||
Sleep: func(context.Context, time.Duration) error {
|
||||
return nil
|
||||
},
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "access-token", token.AccessToken)
|
||||
require.Equal(t, "refresh-token", token.RefreshToken)
|
||||
require.Equal(t, "id-token", token.IDToken)
|
||||
require.Equal(t, SSOBuildScope, token.Scope)
|
||||
require.Equal(t, 1, client.tokenCalls)
|
||||
require.Contains(t, client.cookieHeaders[0], "sso=sso-token")
|
||||
require.Contains(t, client.cookieHeaders[0], "sso-rw=sso-token")
|
||||
require.Contains(t, client.cookieHeaders[len(client.cookieHeaders)-1], "session=web-session")
|
||||
require.Contains(t, client.cookieHeaders[len(client.cookieHeaders)-1], "csrf=csrf-token")
|
||||
}
|
||||
|
||||
func TestNormalizeSSOTokenAcceptsCookieHeader(t *testing.T) {
|
||||
require.Equal(t, "token-1", NormalizeSSOToken("Cookie: foo=bar; sso=token-1; sso-rw=token-2"))
|
||||
require.Equal(t, "token-2", NormalizeSSOToken("sso-rw=token-2; foo=bar"))
|
||||
require.Equal(t, "raw-token", NormalizeSSOToken(" raw-token ; ignored=1"))
|
||||
}
|
||||
|
||||
func ssoDeviceResponse(status int, header http.Header, body string) *http.Response {
|
||||
if header == nil {
|
||||
header = http.Header{}
|
||||
}
|
||||
return &http.Response{
|
||||
StatusCode: status,
|
||||
Header: header,
|
||||
Body: io.NopCloser(strings.NewReader(body)),
|
||||
}
|
||||
}
|
||||
|
||||
func readSSODeviceForm(t *testing.T, req *http.Request) url.Values {
|
||||
t.Helper()
|
||||
data, err := io.ReadAll(req.Body)
|
||||
require.NoError(t, err)
|
||||
values, err := url.ParseQuery(string(data))
|
||||
require.NoError(t, err)
|
||||
return values
|
||||
}
|
||||
@@ -2,12 +2,14 @@ package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
sharedhttp "github.com/Wei-Shaw/sub2api/internal/pkg/httpclient"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/Wei-Shaw/sub2api/internal/util/logredact"
|
||||
@@ -88,6 +90,21 @@ func (c *grokOAuthClient) RefreshToken(ctx context.Context, refreshToken, proxyU
|
||||
return &tokenResp, nil
|
||||
}
|
||||
|
||||
func (c *grokOAuthClient) ConvertSSOToBuild(ctx context.Context, ssoToken, proxyURL string) (*xai.TokenResponse, error) {
|
||||
client, err := createGrokSSOHTTPClient(proxyURL)
|
||||
if err != nil {
|
||||
return nil, infraerrors.Newf(http.StatusBadGateway, "GROK_SSO_CLIENT_INIT_FAILED", "create HTTP client: %v", err)
|
||||
}
|
||||
|
||||
requestCtx, cancel := context.WithTimeout(ctx, xai.SSOConversionTimeout)
|
||||
defer cancel()
|
||||
tokenResp, err := xai.ConvertSSOToBuild(requestCtx, ssoToken, &xai.SSODeviceOptions{HTTPClient: client})
|
||||
if err != nil {
|
||||
return nil, grokSSOConversionError(err)
|
||||
}
|
||||
return tokenResp, nil
|
||||
}
|
||||
|
||||
func createGrokReqClient(proxyURL string) (*req.Client, error) {
|
||||
return getSharedReqClient(reqClientOptions{
|
||||
ProxyURL: proxyURL,
|
||||
@@ -95,6 +112,43 @@ func createGrokReqClient(proxyURL string) (*req.Client, error) {
|
||||
})
|
||||
}
|
||||
|
||||
func createGrokSSOHTTPClient(proxyURL string) (*http.Client, error) {
|
||||
client, err := sharedhttp.GetClient(sharedhttp.Options{
|
||||
ProxyURL: proxyURL,
|
||||
Timeout: xai.SSOConversionTimeout,
|
||||
ResponseHeaderTimeout: 30 * time.Second,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
clone := *client
|
||||
clone.CheckRedirect = func(*http.Request, []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
}
|
||||
return &clone, nil
|
||||
}
|
||||
|
||||
func grokSSOConversionError(err error) error {
|
||||
if errors.Is(err, xai.ErrSSOUnauthorized) {
|
||||
return infraerrors.New(http.StatusUnauthorized, "GROK_SSO_UNAUTHORIZED", "Grok Web SSO cookie is invalid or expired")
|
||||
}
|
||||
if errors.Is(err, xai.ErrSSOAuthorizationDenied) {
|
||||
return infraerrors.New(http.StatusForbidden, "GROK_SSO_AUTHORIZATION_DENIED", "xAI device authorization was denied or expired")
|
||||
}
|
||||
var statusErr xai.SSOHTTPError
|
||||
if errors.As(err, &statusErr) {
|
||||
statusCode := http.StatusBadGateway
|
||||
if statusErr.Status == http.StatusForbidden {
|
||||
statusCode = http.StatusForbidden
|
||||
}
|
||||
return infraerrors.Newf(statusCode, "GROK_SSO_UPSTREAM_FAILED", "xAI SSO conversion failed: %v", err)
|
||||
}
|
||||
if errors.Is(err, context.DeadlineExceeded) || errors.Is(err, context.Canceled) {
|
||||
return infraerrors.Newf(http.StatusGatewayTimeout, "GROK_SSO_TIMEOUT", "xAI SSO conversion timed out: %v", err)
|
||||
}
|
||||
return infraerrors.Newf(http.StatusBadGateway, "GROK_SSO_CONVERSION_FAILED", "xAI SSO conversion failed: %v", err)
|
||||
}
|
||||
|
||||
func grokOAuthStatusError(code, message string, resp *req.Response) error {
|
||||
statusCode := http.StatusBadGateway
|
||||
errorCode := code
|
||||
|
||||
@@ -399,6 +399,7 @@ func registerGrokOAuthRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
|
||||
grok.POST("/oauth/exchange-code", h.Admin.GrokOAuth.ExchangeCode)
|
||||
grok.POST("/oauth/refresh-token", h.Admin.GrokOAuth.RefreshToken)
|
||||
grok.POST("/oauth/create-from-oauth", h.Admin.GrokOAuth.CreateAccountFromOAuth)
|
||||
grok.POST("/sso-to-oauth", h.Admin.GrokOAuth.CreateAccountsFromSSO)
|
||||
grok.POST("/accounts/:id/refresh", h.Admin.GrokOAuth.RefreshAccountToken)
|
||||
grok.GET("/accounts/:id/quota", h.Admin.GrokOAuth.QueryQuota)
|
||||
grok.POST("/accounts/:id/reset-quota", h.Admin.GrokOAuth.ResetQuota)
|
||||
|
||||
@@ -3,8 +3,6 @@ package service
|
||||
import (
|
||||
"context"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -101,6 +99,8 @@ type GrokTokenInfo struct {
|
||||
ClientID string `json:"client_id,omitempty"`
|
||||
Scope string `json:"scope,omitempty"`
|
||||
Email string `json:"email,omitempty"`
|
||||
Subject string `json:"sub,omitempty"`
|
||||
TeamID string `json:"team_id,omitempty"`
|
||||
SubscriptionTier string `json:"subscription_tier,omitempty"`
|
||||
EntitlementStatus string `json:"entitlement_status,omitempty"`
|
||||
}
|
||||
@@ -175,6 +175,18 @@ func (s *GrokOAuthService) ValidateRefreshToken(ctx context.Context, refreshToke
|
||||
return s.RefreshToken(ctx, refreshToken, proxyURL, xai.EffectiveClientID())
|
||||
}
|
||||
|
||||
func (s *GrokOAuthService) ConvertFromSSO(ctx context.Context, ssoToken string, proxyID *int64) (*GrokTokenInfo, error) {
|
||||
proxyURL, err := s.proxyURL(ctx, proxyID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tokenResp, err := s.oauthClient.ConvertSSOToBuild(ctx, ssoToken, proxyURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.tokenInfoFromResponse(tokenResp, xai.DefaultClientID, nil), nil
|
||||
}
|
||||
|
||||
func (s *GrokOAuthService) RefreshAccountToken(ctx context.Context, account *Account) (*GrokTokenInfo, error) {
|
||||
if account == nil || account.Platform != PlatformGrok {
|
||||
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_ACCOUNT", "account is not a Grok account")
|
||||
@@ -229,6 +241,12 @@ func (s *GrokOAuthService) BuildAccountCredentials(tokenInfo *GrokTokenInfo) map
|
||||
if tokenInfo.Email != "" {
|
||||
creds["email"] = tokenInfo.Email
|
||||
}
|
||||
if tokenInfo.Subject != "" {
|
||||
creds["sub"] = tokenInfo.Subject
|
||||
}
|
||||
if tokenInfo.TeamID != "" {
|
||||
creds["team_id"] = tokenInfo.TeamID
|
||||
}
|
||||
if tokenInfo.SubscriptionTier != "" {
|
||||
creds["subscription_tier"] = tokenInfo.SubscriptionTier
|
||||
}
|
||||
@@ -265,12 +283,23 @@ func (s *GrokOAuthService) tokenInfoFromResponse(tokenResp *xai.TokenResponse, c
|
||||
if info.TokenType == "" {
|
||||
info.TokenType = "Bearer"
|
||||
}
|
||||
if email := parseJWTEmailClaim(tokenResp.IDToken); email != "" {
|
||||
info.Email = email
|
||||
}
|
||||
if info.Email == "" && existing != nil {
|
||||
if email, _ := existing["email"].(string); email != "" {
|
||||
info.Email = email
|
||||
applyGrokTokenClaims(info, tokenResp.IDToken)
|
||||
applyGrokTokenClaims(info, tokenResp.AccessToken)
|
||||
if existing != nil {
|
||||
if info.Email == "" {
|
||||
if email, _ := existing["email"].(string); email != "" {
|
||||
info.Email = email
|
||||
}
|
||||
}
|
||||
if info.Subject == "" {
|
||||
if subject, _ := existing["sub"].(string); subject != "" {
|
||||
info.Subject = subject
|
||||
}
|
||||
}
|
||||
if info.TeamID == "" {
|
||||
if teamID, _ := existing["team_id"].(string); teamID != "" {
|
||||
info.TeamID = teamID
|
||||
}
|
||||
}
|
||||
}
|
||||
return info
|
||||
@@ -293,20 +322,21 @@ func (s *GrokOAuthService) proxyURL(ctx context.Context, proxyID *int64) (string
|
||||
return proxy.URL(), nil
|
||||
}
|
||||
|
||||
func parseJWTEmailClaim(token string) string {
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) < 2 {
|
||||
return ""
|
||||
func applyGrokTokenClaims(info *GrokTokenInfo, token string) {
|
||||
if info == nil || strings.TrimSpace(token) == "" {
|
||||
return
|
||||
}
|
||||
payload, err := base64.RawURLEncoding.DecodeString(parts[1])
|
||||
if err != nil {
|
||||
return ""
|
||||
claims := xai.DecodeJWTClaims(token)
|
||||
if claims == nil {
|
||||
return
|
||||
}
|
||||
var claims struct {
|
||||
Email string `json:"email"`
|
||||
if info.Email == "" {
|
||||
info.Email = xai.JWTClaimString(claims, "email")
|
||||
}
|
||||
if err := json.Unmarshal(payload, &claims); err != nil {
|
||||
return ""
|
||||
if info.Subject == "" {
|
||||
info.Subject = xai.JWTClaimString(claims, "sub")
|
||||
}
|
||||
if info.TeamID == "" {
|
||||
info.TeamID = xai.JWTClaimString(claims, "team_id")
|
||||
}
|
||||
return strings.TrimSpace(claims.Email)
|
||||
}
|
||||
|
||||
@@ -4,6 +4,8 @@ package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -13,6 +15,7 @@ import (
|
||||
|
||||
type grokOAuthClientStub struct {
|
||||
refreshResponse *xai.TokenResponse
|
||||
ssoResponse *xai.TokenResponse
|
||||
exchangeCalls int
|
||||
}
|
||||
|
||||
@@ -25,6 +28,10 @@ func (s *grokOAuthClientStub) RefreshToken(context.Context, string, string, stri
|
||||
return s.refreshResponse, nil
|
||||
}
|
||||
|
||||
func (s *grokOAuthClientStub) ConvertSSOToBuild(context.Context, string, string) (*xai.TokenResponse, error) {
|
||||
return s.ssoResponse, nil
|
||||
}
|
||||
|
||||
func TestGrokOAuthServiceRefreshTokenPreservesOriginalRefreshTokenWhenNotRotated(t *testing.T) {
|
||||
svc := NewGrokOAuthService(nil, &grokOAuthClientStub{
|
||||
refreshResponse: &xai.TokenResponse{
|
||||
@@ -79,3 +86,31 @@ func TestGrokOAuthServiceBuildAccountCredentialsDefaultsToSubscriptionProxy(t *t
|
||||
|
||||
require.Equal(t, xai.DefaultCLIBaseURL, credentials["base_url"])
|
||||
}
|
||||
|
||||
func TestGrokOAuthServiceConvertFromSSOExtractsBuildClaims(t *testing.T) {
|
||||
svc := NewGrokOAuthService(nil, &grokOAuthClientStub{
|
||||
ssoResponse: &xai.TokenResponse{
|
||||
AccessToken: makeGrokOAuthJWT(map[string]any{"sub": "user-sub", "team_id": "team-1"}),
|
||||
RefreshToken: "refresh-token",
|
||||
IDToken: makeGrokOAuthJWT(map[string]any{"email": "user@example.com"}),
|
||||
ExpiresIn: 3600,
|
||||
},
|
||||
})
|
||||
defer svc.Stop()
|
||||
|
||||
info, err := svc.ConvertFromSSO(context.Background(), "sso-token", nil)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "user@example.com", info.Email)
|
||||
require.Equal(t, "user-sub", info.Subject)
|
||||
require.Equal(t, "team-1", info.TeamID)
|
||||
|
||||
credentials := svc.BuildAccountCredentials(info)
|
||||
require.Equal(t, "user@example.com", credentials["email"])
|
||||
require.Equal(t, "user-sub", credentials["sub"])
|
||||
require.Equal(t, "team-1", credentials["team_id"])
|
||||
}
|
||||
|
||||
func makeGrokOAuthJWT(claims map[string]any) string {
|
||||
payload, _ := json.Marshal(claims)
|
||||
return "header." + base64.RawURLEncoding.EncodeToString(payload) + ".signature"
|
||||
}
|
||||
|
||||
@@ -22,6 +22,7 @@ type OpenAIOAuthClient interface {
|
||||
type GrokOAuthClient interface {
|
||||
ExchangeCode(ctx context.Context, code, codeVerifier, redirectURI, proxyURL, clientID string) (*xai.TokenResponse, error)
|
||||
RefreshToken(ctx context.Context, refreshToken, proxyURL, clientID string) (*xai.TokenResponse, error)
|
||||
ConvertSSOToBuild(ctx context.Context, ssoToken, proxyURL string) (*xai.TokenResponse, error)
|
||||
}
|
||||
|
||||
// GrokOAuthTokenService is the narrow refresh port used by Grok token providers.
|
||||
|
||||
Reference in New Issue
Block a user