package admin import ( "context" "fmt" "log/slog" "strconv" "strings" "sync" "time" "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 quotaService *service.GrokQuotaService importProber grokImportProber reconciler service.GrokOAuthReconciler } func NewGrokOAuthHandler( grokOAuthService *service.GrokOAuthService, adminService service.AdminService, quotaService *service.GrokQuotaService, reconciler service.GrokOAuthReconciler, ) *GrokOAuthHandler { return &GrokOAuthHandler{ grokOAuthService: grokOAuthService, adminService: adminService, quotaService: quotaService, importProber: quotaService, reconciler: reconciler, } } type GrokGenerateAuthURLRequest struct { ProxyID *int64 `json:"proxy_id"` RedirectURI string `json:"redirect_uri"` } func (h *GrokOAuthHandler) GetCapabilities(c *gin.Context) { response.Success(c, h.grokOAuthService.GetCapabilities()) } func (h *GrokOAuthHandler) GenerateAuthURL(c *gin.Context) { var req GrokGenerateAuthURLRequest if err := c.ShouldBindJSON(&req); err != nil { req = GrokGenerateAuthURLRequest{} } result, err := h.grokOAuthService.GenerateAuthURL(c.Request.Context(), req.ProxyID, req.RedirectURI) if err != nil { response.ErrorFrom(c, err) return } response.Success(c, result) } type GrokExchangeCodeRequest struct { SessionID string `json:"session_id" binding:"required"` Code string `json:"code" binding:"required"` State string `json:"state"` RedirectURI string `json:"redirect_uri"` ProxyID *int64 `json:"proxy_id"` } func (h *GrokOAuthHandler) ExchangeCode(c *gin.Context) { var req GrokExchangeCodeRequest if err := c.ShouldBindJSON(&req); err != nil { response.BadRequest(c, "Invalid request: "+err.Error()) return } tokenInfo, err := h.grokOAuthService.ExchangeCode(c.Request.Context(), &service.GrokExchangeCodeInput{ SessionID: req.SessionID, Code: req.Code, State: req.State, RedirectURI: req.RedirectURI, ProxyID: req.ProxyID, }) if err != nil { response.ErrorFrom(c, err) return } response.Success(c, tokenInfo) } type GrokRefreshTokenRequest struct { RefreshToken string `json:"refresh_token"` RT string `json:"rt"` ClientID string `json:"client_id"` ProxyID *int64 `json:"proxy_id"` } type GrokSSOTokenRequest struct { SSOToken string `json:"sso_token"` ProxyID *int64 `json:"proxy_id"` } type GrokPasswordAuthorizeRequest struct { Email string `json:"email"` Password string `json:"password"` ProxyID *int64 `json:"proxy_id"` } func (h *GrokOAuthHandler) RefreshToken(c *gin.Context) { var req GrokRefreshTokenRequest if err := c.ShouldBindJSON(&req); err != nil { response.BadRequest(c, "Invalid request: "+err.Error()) return } refreshToken := strings.TrimSpace(req.RefreshToken) if refreshToken == "" { refreshToken = strings.TrimSpace(req.RT) } if refreshToken == "" { response.BadRequest(c, "refresh_token is required") return } var proxyURL string if req.ProxyID != nil { proxy, err := h.adminService.GetProxy(c.Request.Context(), *req.ProxyID) if err != nil { response.ErrorFrom(c, err) return } if proxy == nil { response.BadRequest(c, "GROK_OAUTH_PROXY_NOT_FOUND: proxy not found") return } proxyURL = proxy.URL() } tokenInfo, err := h.grokOAuthService.RefreshToken(c.Request.Context(), refreshToken, proxyURL, req.ClientID) if err != nil { response.ErrorFrom(c, err) return } response.Success(c, tokenInfo) } // ValidateSSOToken converts a Web SSO cookie into Build OAuth tokens. // Response contains OAuth token info only — never echoes sso_token. func (h *GrokOAuthHandler) ValidateSSOToken(c *gin.Context) { var req GrokSSOTokenRequest if err := c.ShouldBindJSON(&req); err != nil { response.BadRequest(c, "Invalid request: "+err.Error()) return } tokenInfo, err := h.grokOAuthService.ValidateSSOToken(c.Request.Context(), req.SSOToken, req.ProxyID) if err != nil { response.ErrorFrom(c, err) return } response.Success(c, tokenInfo) } // AuthorizePassword exchanges email/password for Build OAuth tokens via SSO conversion. // Response never includes password or raw sso_token. func (h *GrokOAuthHandler) AuthorizePassword(c *gin.Context) { var req GrokPasswordAuthorizeRequest if err := c.ShouldBindJSON(&req); err != nil { response.BadRequest(c, "Invalid request: "+err.Error()) return } tokenInfo, err := h.grokOAuthService.AuthorizePassword(c.Request.Context(), req.Email, req.Password, req.ProxyID) if err != nil { response.ErrorFrom(c, err) return } response.Success(c, tokenInfo) } func (h *GrokOAuthHandler) RefreshAccountToken(c *gin.Context) { accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) if err != nil { response.BadRequest(c, "Invalid account ID") return } account, err := h.adminService.GetAccount(c.Request.Context(), accountID) if err != nil { response.ErrorFrom(c, err) return } if account.Platform != service.PlatformGrok { response.BadRequest(c, "Account platform does not match Grok OAuth endpoint") return } if !account.IsOAuth() { response.BadRequest(c, "Cannot refresh non-OAuth account credentials") return } tokenInfo, err := h.grokOAuthService.RefreshAccountToken(c.Request.Context(), account) if err != nil { response.ErrorFrom(c, err) return } newCredentials := h.grokOAuthService.BuildAccountCredentials(tokenInfo) newCredentials = service.MergeCredentials(account.Credentials, newCredentials) if baseURL := strings.TrimSpace(account.GetCredential("base_url")); baseURL != "" { newCredentials["base_url"] = baseURL } updatedAccount, err := h.adminService.UpdateAccount(c.Request.Context(), accountID, &service.UpdateAccountInput{ Credentials: newCredentials, }) if err != nil { response.ErrorFrom(c, err) return } response.Success(c, dto.AccountFromService(updatedAccount)) } type GrokOAuthReconcileRequest struct { DryRun *bool `json:"dry_run"` Apply bool `json:"apply"` AfterID int64 `json:"after_id"` Limit int `json:"limit"` RefreshWindowSeconds int64 `json:"refresh_window_seconds"` } func (h *GrokOAuthHandler) ReconcileOAuthAccounts(c *gin.Context) { var req GrokOAuthReconcileRequest if err := c.ShouldBindJSON(&req); err != nil { response.BadRequest(c, "Invalid request") return } dryRun := true if req.DryRun != nil { dryRun = *req.DryRun } if req.Apply == dryRun { response.ErrorFrom(c, service.ErrGrokOAuthReconcileMode) return } if req.RefreshWindowSeconds < 0 || req.RefreshWindowSeconds > int64((24*time.Hour)/time.Second) { response.ErrorFrom(c, service.ErrGrokOAuthReconcileWindow) return } if h.reconciler == nil { response.InternalError(c, "Grok OAuth reconciliation service is unavailable") return } result, err := h.reconciler.ReconcileGrokOAuth(c.Request.Context(), service.GrokOAuthReconcileInput{ DryRun: dryRun, Apply: req.Apply, AfterID: req.AfterID, Limit: req.Limit, RefreshWindow: time.Duration(req.RefreshWindowSeconds) * time.Second, }) if err != nil { response.ErrorFrom(c, err) return } response.Success(c, result) } func (h *GrokOAuthHandler) CreateAccountFromOAuth(c *gin.Context) { var req struct { SessionID string `json:"session_id" binding:"required"` Code string `json:"code" binding:"required"` State string `json:"state"` RedirectURI string `json:"redirect_uri"` ProxyID *int64 `json:"proxy_id"` Name string `json:"name"` Concurrency int `json:"concurrency"` Priority int `json:"priority"` GroupIDs []int64 `json:"group_ids"` } if err := c.ShouldBindJSON(&req); err != nil { response.BadRequest(c, "Invalid request: "+err.Error()) return } tokenInfo, err := h.grokOAuthService.ExchangeCode(c.Request.Context(), &service.GrokExchangeCodeInput{ SessionID: req.SessionID, Code: req.Code, State: req.State, RedirectURI: req.RedirectURI, ProxyID: req.ProxyID, }) if err != nil { response.ErrorFrom(c, err) return } credentials := h.grokOAuthService.BuildAccountCredentials(tokenInfo) name := strings.TrimSpace(req.Name) if name == "" && tokenInfo.Email != "" { name = tokenInfo.Email } if name == "" { name = "Grok OAuth Account" } account, err := h.adminService.CreateAccount(c.Request.Context(), &service.CreateAccountInput{ Name: name, Platform: service.PlatformGrok, Type: service.AccountTypeOAuth, Credentials: credentials, ProxyID: req.ProxyID, Concurrency: req.Concurrency, Priority: req.Priority, GroupIDs: req.GroupIDs, }) if err != nil { response.ErrorFrom(c, err) return } h.scheduleGrokImportProbe(account) 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 := grokSSOImportCredentials(h.grokOAuthService.BuildAccountCredentials(tokenInfo), req.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)}} } h.scheduleGrokImportProbe(account) return grokSSOImportWorkerResult{ created: true, item: GrokSSOToOAuthItemResult{ Index: index, Name: name, Email: tokenInfo.Email, Account: dto.AccountFromService(account), }, } } // grokSSOImportCredentials 合并 SSO 兑换出的凭据与导入请求携带的运营侧配置。 // token 字段以 BuildAccountCredentials 为准(请求不可覆盖);但 base_url 是运营侧 // 配置且 Build 恒写官方地址,会吞掉导入时指定的自定义转发地址——与 // RefreshAccountToken 的保留逻辑对齐,请求显式提供时以请求为准。 func grokSSOImportCredentials(built map[string]any, reqCredentials map[string]any) map[string]any { // Only merge operator config from the request — never free-form secrets // (password / sso_token / cookie / etc.) into stored credentials. allowedReqKeys := map[string]struct{}{ "base_url": {}, "model_mapping": {}, "header_override": {}, "header_overrides": {}, "header_override_enabled": {}, "custom_headers": {}, } ops := map[string]any{} for k, v := range reqCredentials { if _, ok := allowedReqKeys[k]; !ok { continue } if service.IsSensitiveCredentialKey(k) { continue } ops[k] = v } credentials := service.MergeCredentials(ops, built) // Strip any sensitive keys that might have slipped in via older callers. for k := range credentials { if service.IsSensitiveCredentialKey(k) { // Keep only keys produced by BuildAccountCredentials (tokens). if k == "access_token" || k == "refresh_token" || k == "id_token" { continue } delete(credentials, k) } } if reqBaseURL, ok := reqCredentials["base_url"].(string); ok && strings.TrimSpace(reqBaseURL) != "" { credentials["base_url"] = strings.TrimSpace(reqBaseURL) } return service.SanitizeStoredCredentials(service.PlatformGrok, credentials) } 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 { response.BadRequest(c, "Invalid account ID") return } if h.quotaService == nil { response.BadRequest(c, "grok quota service is not enabled") return } result, err := h.quotaService.QueryQuota(c.Request.Context(), accountID) if err != nil { response.ErrorFrom(c, err) return } response.Success(c, result) } func (h *GrokOAuthHandler) ResetQuota(c *gin.Context) { accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) if err != nil { response.BadRequest(c, "Invalid account ID") return } if h.quotaService == nil { response.BadRequest(c, "grok quota service is not enabled") return } result, err := h.quotaService.ResetQuota(c.Request.Context(), accountID) if err != nil { response.ErrorFrom(c, err) return } response.Success(c, result) } func (h *GrokOAuthHandler) RuntimeSanity(c *gin.Context) { response.Success(c, xai.RuntimeSanity()) }