Files
sub2api/backend/internal/service/grok_oauth_service.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

478 lines
16 KiB
Go

package service
import (
"context"
"crypto/subtle"
"errors"
"net/http"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
)
const grokDefaultAccessTokenTTL = 6 * time.Hour
type GrokOAuthService struct {
sessionStore *xai.SessionStore
proxyRepo ProxyRepository
oauthClient GrokOAuthClient
config *config.Config
}
func NewGrokOAuthService(proxyRepo ProxyRepository, oauthClient GrokOAuthClient, configs ...*config.Config) *GrokOAuthService {
service := &GrokOAuthService{
sessionStore: xai.NewSessionStore(),
proxyRepo: proxyRepo,
oauthClient: oauthClient,
}
if len(configs) > 0 {
service.config = configs[0]
}
return service
}
// WithSessionStore replaces the in-memory OAuth session store (e.g. Redis-backed
// for cross-instance single-use callbacks). Redis wiring stays in Wire providers
// so this service package does not import go-redis (depguard).
func (s *GrokOAuthService) WithSessionStore(store *xai.SessionStore) *GrokOAuthService {
if s != nil && store != nil {
if s.sessionStore != nil {
s.sessionStore.Stop()
}
s.sessionStore = store
}
return s
}
type GrokOAuthCapabilities struct {
PasswordAuthEnabled bool `json:"password_auth_enabled"`
}
func (s *GrokOAuthService) GetCapabilities() GrokOAuthCapabilities {
return GrokOAuthCapabilities{PasswordAuthEnabled: s.passwordAuthEnabled()}
}
func (s *GrokOAuthService) passwordAuthEnabled() bool {
return s.config != nil && s.config.Gateway.Grok.PasswordAuthEnabled
}
type GrokAuthURLResult struct {
AuthURL string `json:"auth_url"`
SessionID string `json:"session_id"`
State string `json:"state"`
}
func (s *GrokOAuthService) GenerateAuthURL(ctx context.Context, proxyID *int64, redirectURI string) (*GrokAuthURLResult, error) {
state, err := xai.GenerateState()
if err != nil {
return nil, infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_STATE_FAILED", "failed to generate state: %v", err)
}
nonce, err := xai.GenerateNonce()
if err != nil {
return nil, infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_NONCE_FAILED", "failed to generate nonce: %v", err)
}
codeVerifier, err := xai.GenerateCodeVerifier()
if err != nil {
return nil, infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_VERIFIER_FAILED", "failed to generate code verifier: %v", err)
}
sessionID, err := xai.GenerateSessionID()
if err != nil {
return nil, infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_SESSION_FAILED", "failed to generate session ID: %v", err)
}
proxyURL, err := s.proxyURL(ctx, proxyID)
if err != nil {
return nil, err
}
redirectURI = xai.EffectiveRedirectURI(redirectURI)
codeChallenge := xai.GenerateCodeChallenge(codeVerifier)
authURL, err := xai.BuildAuthorizationURL(state, codeChallenge, redirectURI, nonce)
if err != nil {
return nil, infraerrors.Newf(http.StatusBadRequest, "GROK_OAUTH_INVALID_AUTHORIZE_URL", "%v", err)
}
s.sessionStore.Set(sessionID, &xai.OAuthSession{
State: state,
CodeVerifier: codeVerifier,
CodeChallenge: codeChallenge,
ClientID: xai.EffectiveClientID(),
Scope: xai.EffectiveScope(),
ProxyURL: proxyURL,
RedirectURI: redirectURI,
CreatedAt: time.Now(),
})
return &GrokAuthURLResult{
AuthURL: authURL,
SessionID: sessionID,
State: state,
}, nil
}
type GrokExchangeCodeInput struct {
SessionID string
Code string
State string
RedirectURI string
ProxyID *int64
}
type GrokTokenInfo struct {
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token,omitempty"`
IDToken string `json:"id_token,omitempty"`
TokenType string `json:"token_type,omitempty"`
ExpiresIn int64 `json:"expires_in"`
ExpiresAt int64 `json:"expires_at"`
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"`
}
// GrokPasswordLoginResult is an ephemeral password-login outcome.
// SSOToken is never persisted and must only feed ConvertSSOToBuild.
type GrokPasswordLoginResult struct {
Email string `json:"email,omitempty"`
SSOToken string `json:"sso_token"`
}
func (s *GrokOAuthService) ExchangeCode(ctx context.Context, input *GrokExchangeCodeInput) (*GrokTokenInfo, error) {
if input == nil {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_INPUT", "input is required")
}
session, ok := s.sessionStore.Get(input.SessionID)
if !ok {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_SESSION_NOT_FOUND", "session not found or expired")
}
parsed := xai.ParseAuthorizationInput(input.Code)
code := strings.TrimSpace(parsed.Code)
if code == "" {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_CODE_REQUIRED", "authorization code is required")
}
state := strings.TrimSpace(input.State)
if state == "" {
state = strings.TrimSpace(parsed.State)
}
if state == "" {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_STATE_REQUIRED", "oauth state is required")
}
if subtle.ConstantTimeCompare([]byte(state), []byte(session.State)) != 1 {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_STATE", "invalid oauth state")
}
if redirectURI := strings.TrimSpace(input.RedirectURI); redirectURI != "" &&
redirectURI != strings.TrimSpace(session.RedirectURI) {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_REDIRECT_URI_MISMATCH", "redirect_uri does not match the OAuth session")
}
proxyURL := session.ProxyURL
if input.ProxyID != nil {
var err error
proxyURL, err = s.proxyURL(ctx, input.ProxyID)
if err != nil {
return nil, err
}
}
if err := s.requireOAuthClient(); err != nil {
return nil, err
}
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)
tokenResp, err := s.oauthClient.ExchangeCode(ctx, code, session.CodeVerifier, session.RedirectURI, proxyURL, session.ClientID)
if err != nil {
return nil, err
}
if err := validateGrokTokenResponse(tokenResp); err != nil {
return nil, err
}
return s.tokenInfoFromResponse(tokenResp, session.ClientID, nil), nil
}
func (s *GrokOAuthService) requireOAuthClient() error {
if s == nil || s.oauthClient == nil {
return infraerrors.New(http.StatusInternalServerError, "GROK_OAUTH_CLIENT_NOT_CONFIGURED", "oauth client is not configured")
}
return nil
}
func (s *GrokOAuthService) RefreshToken(ctx context.Context, refreshToken, proxyURL, clientID string) (*GrokTokenInfo, error) {
refreshToken = strings.TrimSpace(refreshToken)
if refreshToken == "" {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_NO_REFRESH_TOKEN", "refresh_token is required")
}
if err := s.requireOAuthClient(); err != nil {
return nil, err
}
tokenResp, err := s.oauthClient.RefreshToken(ctx, refreshToken, proxyURL, clientID)
if err != nil {
return nil, err
}
if err := validateGrokTokenResponse(tokenResp); err != nil {
return nil, err
}
tokenInfo := s.tokenInfoFromResponse(tokenResp, clientID, nil)
if tokenInfo.RefreshToken == "" {
tokenInfo.RefreshToken = refreshToken
}
return tokenInfo, nil
}
func (s *GrokOAuthService) ValidateRefreshToken(ctx context.Context, refreshToken string, proxyID *int64) (*GrokTokenInfo, error) {
proxyURL, err := s.proxyURL(ctx, proxyID)
if err != nil {
return nil, err
}
return s.RefreshToken(ctx, refreshToken, proxyURL, xai.EffectiveClientID())
}
// ValidateSSOToken converts a Web SSO cookie into Build OAuth tokens.
// The raw sso_token is never stored on GrokTokenInfo or account credentials.
func (s *GrokOAuthService) ValidateSSOToken(ctx context.Context, ssoToken string, proxyID *int64) (*GrokTokenInfo, error) {
ssoToken = strings.TrimSpace(ssoToken)
if ssoToken == "" {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_NO_SSO_TOKEN", "sso_token is required")
}
if err := s.requireOAuthClient(); err != nil {
return nil, err
}
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
}
if err := validateGrokTokenResponse(tokenResp); err != nil {
return nil, err
}
return s.tokenInfoFromResponse(tokenResp, xai.DefaultClientID, nil), nil
}
// ConvertFromSSO is the batch-import entry point; same semantics as ValidateSSOToken.
func (s *GrokOAuthService) ConvertFromSSO(ctx context.Context, ssoToken string, proxyID *int64) (*GrokTokenInfo, error) {
return s.ValidateSSOToken(ctx, ssoToken, proxyID)
}
// AuthorizePassword logs in with email/password, converts the resulting SSO cookie
// to Build OAuth, and returns OAuth tokens only. Password and raw SSO are never persisted.
func (s *GrokOAuthService) AuthorizePassword(ctx context.Context, email, password string, proxyID *int64) (*GrokTokenInfo, error) {
if !s.passwordAuthEnabled() {
return nil, infraerrors.New(http.StatusForbidden, "GROK_OAUTH_PASSWORD_AUTH_DISABLED", "Grok password authorization is disabled")
}
email = strings.TrimSpace(email)
if email == "" {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_EMAIL_REQUIRED", "email is required")
}
if strings.TrimSpace(password) == "" {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_PASSWORD_REQUIRED", "password is required")
}
if err := s.requireOAuthClient(); err != nil {
return nil, err
}
proxyURL, err := s.proxyURL(ctx, proxyID)
if err != nil {
return nil, err
}
loginResult, err := s.oauthClient.LoginWithPassword(ctx, email, password, proxyURL)
if err != nil {
return nil, err
}
if loginResult == nil || strings.TrimSpace(loginResult.SSOToken) == "" {
return nil, infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "grok password login did not return sso_token")
}
info, err := s.ValidateSSOToken(ctx, loginResult.SSOToken, proxyID)
if err != nil {
return nil, err
}
if strings.TrimSpace(info.Email) == "" {
info.Email = loginResult.Email
}
return info, nil
}
func validateGrokTokenResponse(tokenResp *xai.TokenResponse) error {
if tokenResp == nil || strings.TrimSpace(tokenResp.AccessToken) == "" {
return infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_INVALID_TOKEN_RESPONSE", "grok oauth token response missing access_token")
}
return 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")
}
if account.Type != AccountTypeOAuth {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_ACCOUNT_TYPE", "account is not an OAuth account")
}
proxyURL, err := s.proxyURL(ctx, account.ProxyID)
if err != nil {
return nil, err
}
refreshToken := account.GetCredential("refresh_token")
if strings.TrimSpace(refreshToken) == "" {
return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_NO_REFRESH_TOKEN", "no refresh token available")
}
clientID := account.GetCredential("client_id")
tokenInfo, err := s.RefreshToken(ctx, refreshToken, proxyURL, clientID)
if err != nil {
return nil, err
}
// New access-token JWT is authoritative. Keep the stored value only when
// the refreshed token has no tier claim (opaque AT / missing field).
if strings.TrimSpace(tokenInfo.SubscriptionTier) == "" {
tokenInfo.SubscriptionTier = account.GetCredential("subscription_tier")
}
if strings.TrimSpace(tokenInfo.EntitlementStatus) == "" {
tokenInfo.EntitlementStatus = account.GetCredential("entitlement_status")
}
return tokenInfo, nil
}
func (s *GrokOAuthService) BuildAccountCredentials(tokenInfo *GrokTokenInfo) map[string]any {
if tokenInfo == nil {
return nil
}
expiresAt := time.Unix(tokenInfo.ExpiresAt, 0).UTC().Format(time.RFC3339)
creds := map[string]any{
"access_token": tokenInfo.AccessToken,
"expires_at": expiresAt,
}
if tokenInfo.RefreshToken != "" {
creds["refresh_token"] = tokenInfo.RefreshToken
}
if tokenInfo.TokenType != "" {
creds["token_type"] = tokenInfo.TokenType
}
if tokenInfo.IDToken != "" {
creds["id_token"] = tokenInfo.IDToken
}
if tokenInfo.ClientID != "" {
creds["client_id"] = tokenInfo.ClientID
}
if tokenInfo.Scope != "" {
creds["scope"] = tokenInfo.Scope
}
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
}
if tokenInfo.EntitlementStatus != "" {
creds["entitlement_status"] = tokenInfo.EntitlementStatus
}
creds["base_url"] = xai.DefaultCLIBaseURL
return creds
}
func (s *GrokOAuthService) Stop() {
s.sessionStore.Stop()
}
func (s *GrokOAuthService) tokenInfoFromResponse(tokenResp *xai.TokenResponse, clientID string, existing map[string]any) *GrokTokenInfo {
now := time.Now()
expiresIn := tokenResp.ExpiresIn
if expiresIn <= 0 {
expiresIn = int64(grokDefaultAccessTokenTTL.Seconds())
}
info := &GrokTokenInfo{
AccessToken: tokenResp.AccessToken,
RefreshToken: tokenResp.RefreshToken,
IDToken: tokenResp.IDToken,
TokenType: tokenResp.TokenType,
ExpiresIn: expiresIn,
ExpiresAt: now.Add(time.Duration(expiresIn) * time.Second).Unix(),
ClientID: strings.TrimSpace(clientID),
Scope: tokenResp.Scope,
}
if info.ClientID == "" {
info.ClientID = xai.EffectiveClientID()
}
if info.TokenType == "" {
info.TokenType = "Bearer"
}
applyGrokTokenClaims(info, tokenResp.IDToken, false)
applyGrokTokenClaims(info, tokenResp.AccessToken, true)
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
}
func (s *GrokOAuthService) proxyURL(ctx context.Context, proxyID *int64) (string, error) {
if proxyID == nil {
return "", nil
}
if s.proxyRepo == nil {
return "", infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_PROXY_NOT_AVAILABLE", "proxy repository is not available")
}
proxy, err := s.proxyRepo.GetByID(ctx, *proxyID)
if err != nil {
if errors.Is(err, ErrProxyNotFound) {
return "", infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_PROXY_NOT_FOUND", "configured proxy was not found")
}
return "", infraerrors.New(http.StatusServiceUnavailable, "GROK_OAUTH_PROXY_LOOKUP_FAILED", "proxy lookup is temporarily unavailable")
}
if proxy == nil {
return "", infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_PROXY_NOT_FOUND", "configured proxy was not found")
}
return proxy.URL(), nil
}
func applyGrokTokenClaims(info *GrokTokenInfo, token string, includeTier bool) {
if info == nil || strings.TrimSpace(token) == "" {
return
}
claims := xai.DecodeJWTClaims(token)
if claims == nil {
return
}
if info.Email == "" {
info.Email = xai.JWTClaimString(claims, "email")
}
if info.Subject == "" {
info.Subject = xai.JWTClaimString(claims, "sub")
}
if info.TeamID == "" {
info.TeamID = xai.JWTClaimString(claims, "team_id")
}
if includeTier {
if tier := xai.SubscriptionTierFromJWT(token); tier != "" {
info.SubscriptionTier = tier
}
}
}