893 lines
29 KiB
Go
893 lines
29 KiB
Go
package management
|
|
|
|
import (
|
|
"context"
|
|
"crypto/sha256"
|
|
"encoding/hex"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"net/http"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/router-for-me/CLIProxyAPI/v7/internal/auth/antigravity"
|
|
"github.com/router-for-me/CLIProxyAPI/v7/internal/auth/claude"
|
|
"github.com/router-for-me/CLIProxyAPI/v7/internal/auth/codex"
|
|
"github.com/router-for-me/CLIProxyAPI/v7/internal/auth/kimi"
|
|
xaiauth "github.com/router-for-me/CLIProxyAPI/v7/internal/auth/xai"
|
|
"github.com/router-for-me/CLIProxyAPI/v7/internal/misc"
|
|
"github.com/router-for-me/CLIProxyAPI/v7/internal/pluginhost"
|
|
"github.com/router-for-me/CLIProxyAPI/v7/internal/util"
|
|
coreauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth"
|
|
"github.com/router-for-me/CLIProxyAPI/v7/sdk/pluginapi"
|
|
log "github.com/sirupsen/logrus"
|
|
)
|
|
|
|
type codexOAuthService interface {
|
|
GenerateAuthURL(state string, pkceCodes *codex.PKCECodes) (string, error)
|
|
ExchangeCodeForTokens(ctx context.Context, code string, pkceCodes *codex.PKCECodes) (*codex.CodexAuthBundle, error)
|
|
CreateTokenStorage(bundle *codex.CodexAuthBundle) *codex.CodexTokenStorage
|
|
}
|
|
|
|
func (h *Handler) RequestAnthropicToken(c *gin.Context) {
|
|
ctx := context.Background()
|
|
ctx = PopulateAuthContext(ctx, c)
|
|
|
|
fmt.Println("Initializing Claude authentication...")
|
|
|
|
// Generate PKCE codes
|
|
pkceCodes, err := claude.GeneratePKCECodes()
|
|
if err != nil {
|
|
log.Errorf("Failed to generate PKCE codes: %v", err)
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to generate PKCE codes"})
|
|
return
|
|
}
|
|
|
|
// Generate random state parameter
|
|
state, err := misc.GenerateRandomState()
|
|
if err != nil {
|
|
log.Errorf("Failed to generate state parameter: %v", err)
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to generate state parameter"})
|
|
return
|
|
}
|
|
|
|
// Initialize Claude auth service
|
|
anthropicAuth := claude.NewClaudeAuth(h.cfg)
|
|
|
|
// Generate authorization URL (then override redirect_uri to reuse server port)
|
|
authURL, state, err := anthropicAuth.GenerateAuthURL(state, pkceCodes)
|
|
if err != nil {
|
|
log.Errorf("Failed to generate authorization URL: %v", err)
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to generate authorization url"})
|
|
return
|
|
}
|
|
|
|
RegisterOAuthSession(state, "anthropic")
|
|
|
|
isWebUI := isWebUIRequest(c)
|
|
var forwarder *callbackForwarder
|
|
if isWebUI {
|
|
targetURL, errTarget := h.managementCallbackURL("/anthropic/callback")
|
|
if errTarget != nil {
|
|
log.WithError(errTarget).Error("failed to compute anthropic callback target")
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "callback server unavailable"})
|
|
return
|
|
}
|
|
var errStart error
|
|
if forwarder, errStart = startCallbackForwarder(anthropicCallbackPort, "anthropic", targetURL); errStart != nil {
|
|
log.WithError(errStart).Error("failed to start anthropic callback forwarder")
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to start callback server"})
|
|
return
|
|
}
|
|
}
|
|
|
|
go func() {
|
|
if isWebUI {
|
|
defer stopCallbackForwarderInstance(anthropicCallbackPort, forwarder)
|
|
}
|
|
|
|
// Helper: wait for callback file
|
|
waitFile := filepath.Join(h.cfg.AuthDir, fmt.Sprintf(".oauth-anthropic-%s.oauth", state))
|
|
waitForFile := func(path string, timeout time.Duration) (map[string]string, error) {
|
|
deadline := time.Now().Add(timeout)
|
|
for {
|
|
if !IsOAuthSessionPending(state, "anthropic") {
|
|
return nil, errOAuthSessionNotPending
|
|
}
|
|
if time.Now().After(deadline) {
|
|
SetOAuthSessionError(state, "Timeout waiting for OAuth callback")
|
|
return nil, fmt.Errorf("timeout waiting for OAuth callback")
|
|
}
|
|
data, errRead := os.ReadFile(path)
|
|
if errRead == nil {
|
|
var m map[string]string
|
|
_ = json.Unmarshal(data, &m)
|
|
_ = os.Remove(path)
|
|
return m, nil
|
|
}
|
|
time.Sleep(500 * time.Millisecond)
|
|
}
|
|
}
|
|
|
|
fmt.Println("Waiting for authentication callback...")
|
|
// Wait up to 5 minutes
|
|
resultMap, errWait := waitForFile(waitFile, 5*time.Minute)
|
|
if errWait != nil {
|
|
if errors.Is(errWait, errOAuthSessionNotPending) {
|
|
return
|
|
}
|
|
authErr := claude.NewAuthenticationError(claude.ErrCallbackTimeout, errWait)
|
|
log.Error(claude.GetUserFriendlyMessage(authErr))
|
|
return
|
|
}
|
|
if errStr := resultMap["error"]; errStr != "" {
|
|
oauthErr := claude.NewOAuthError(errStr, "", http.StatusBadRequest)
|
|
log.Error(claude.GetUserFriendlyMessage(oauthErr))
|
|
SetOAuthSessionError(state, "Bad request")
|
|
return
|
|
}
|
|
if resultMap["state"] != state {
|
|
authErr := claude.NewAuthenticationError(claude.ErrInvalidState, fmt.Errorf("expected %s, got %s", state, resultMap["state"]))
|
|
log.Error(claude.GetUserFriendlyMessage(authErr))
|
|
SetOAuthSessionError(state, "State code error")
|
|
return
|
|
}
|
|
|
|
// Parse code (Claude may append state after '#')
|
|
rawCode := resultMap["code"]
|
|
code := strings.Split(rawCode, "#")[0]
|
|
|
|
// Exchange code for tokens using internal auth service
|
|
bundle, errExchange := anthropicAuth.ExchangeCodeForTokens(ctx, code, state, pkceCodes)
|
|
if errExchange != nil {
|
|
authErr := claude.NewAuthenticationError(claude.ErrCodeExchangeFailed, errExchange)
|
|
log.Errorf("Failed to exchange authorization code for tokens: %v", authErr)
|
|
SetOAuthSessionError(state, "Failed to exchange authorization code for tokens")
|
|
return
|
|
}
|
|
|
|
// Create token storage
|
|
tokenStorage := anthropicAuth.CreateTokenStorage(bundle)
|
|
metadata := map[string]any{"email": tokenStorage.Email}
|
|
if tokenStorage.AccountUUID != "" {
|
|
metadata["account_uuid"] = tokenStorage.AccountUUID
|
|
}
|
|
if tokenStorage.OrganizationUUID != "" {
|
|
metadata["organization_uuid"] = tokenStorage.OrganizationUUID
|
|
}
|
|
if tokenStorage.OrganizationName != "" {
|
|
metadata["organization_name"] = tokenStorage.OrganizationName
|
|
}
|
|
if len(tokenStorage.DeviceIDs) > 0 {
|
|
metadata[claude.ClaudeDeviceIDsMetadataKey] = append([]string(nil), tokenStorage.DeviceIDs...)
|
|
}
|
|
record := &coreauth.Auth{
|
|
ID: fmt.Sprintf("claude-%s.json", tokenStorage.Email),
|
|
Provider: "claude",
|
|
FileName: fmt.Sprintf("claude-%s.json", tokenStorage.Email),
|
|
Storage: tokenStorage,
|
|
Metadata: metadata,
|
|
}
|
|
if errGuard := guardOAuthSessionPendingForSave(state, "anthropic"); errGuard != nil {
|
|
return
|
|
}
|
|
savedPath, errSave := h.saveTokenRecord(ctx, record)
|
|
if errSave != nil {
|
|
log.Errorf("Failed to save authentication tokens: %v", errSave)
|
|
SetOAuthSessionError(state, "Failed to save authentication tokens")
|
|
return
|
|
}
|
|
|
|
fmt.Printf("Authentication successful! Token saved to %s\n", savedPath)
|
|
if bundle.APIKey != "" {
|
|
fmt.Println("API key obtained and saved")
|
|
}
|
|
fmt.Println("You can now use Claude services through this CLI")
|
|
CompleteOAuthSession(state)
|
|
}()
|
|
|
|
c.JSON(200, gin.H{"status": "ok", "url": authURL, "state": state})
|
|
}
|
|
|
|
func (h *Handler) RequestCodexToken(c *gin.Context) {
|
|
ctx := context.Background()
|
|
ctx = PopulateAuthContext(ctx, c)
|
|
|
|
fmt.Println("Initializing Codex authentication...")
|
|
|
|
// Generate PKCE codes
|
|
pkceCodes, err := codex.GeneratePKCECodes()
|
|
if err != nil {
|
|
log.Errorf("Failed to generate PKCE codes: %v", err)
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to generate PKCE codes"})
|
|
return
|
|
}
|
|
|
|
// Generate random state parameter
|
|
state, err := misc.GenerateRandomState()
|
|
if err != nil {
|
|
log.Errorf("Failed to generate state parameter: %v", err)
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to generate state parameter"})
|
|
return
|
|
}
|
|
|
|
// Initialize Codex auth service
|
|
openaiAuth := newCodexOAuthService(h.cfg)
|
|
|
|
// Generate authorization URL
|
|
authURL, err := openaiAuth.GenerateAuthURL(state, pkceCodes)
|
|
if err != nil {
|
|
log.Errorf("Failed to generate authorization URL: %v", err)
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to generate authorization url"})
|
|
return
|
|
}
|
|
|
|
RegisterOAuthSession(state, "codex")
|
|
|
|
isWebUI := isWebUIRequest(c)
|
|
var forwarder *callbackForwarder
|
|
if isWebUI {
|
|
targetURL, errTarget := h.managementCallbackURL("/codex/callback")
|
|
if errTarget != nil {
|
|
log.WithError(errTarget).Error("failed to compute codex callback target")
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "callback server unavailable"})
|
|
return
|
|
}
|
|
var errStart error
|
|
if forwarder, errStart = startCallbackForwarder(codexCallbackPort, "codex", targetURL); errStart != nil {
|
|
log.WithError(errStart).Error("failed to start codex callback forwarder")
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to start callback server"})
|
|
return
|
|
}
|
|
}
|
|
|
|
go func() {
|
|
if isWebUI {
|
|
defer stopCallbackForwarderInstance(codexCallbackPort, forwarder)
|
|
}
|
|
|
|
// Wait for callback file
|
|
waitFile := filepath.Join(h.cfg.AuthDir, fmt.Sprintf(".oauth-codex-%s.oauth", state))
|
|
deadline := time.Now().Add(5 * time.Minute)
|
|
var code string
|
|
for {
|
|
if !IsOAuthSessionPending(state, "codex") {
|
|
return
|
|
}
|
|
if time.Now().After(deadline) {
|
|
authErr := codex.NewAuthenticationError(codex.ErrCallbackTimeout, fmt.Errorf("timeout waiting for OAuth callback"))
|
|
log.Error(codex.GetUserFriendlyMessage(authErr))
|
|
SetOAuthSessionError(state, "Timeout waiting for OAuth callback")
|
|
return
|
|
}
|
|
if data, errR := os.ReadFile(waitFile); errR == nil {
|
|
var m map[string]string
|
|
_ = json.Unmarshal(data, &m)
|
|
_ = os.Remove(waitFile)
|
|
if errStr := m["error"]; errStr != "" {
|
|
oauthErr := codex.NewOAuthError(errStr, "", http.StatusBadRequest)
|
|
log.Error(codex.GetUserFriendlyMessage(oauthErr))
|
|
SetOAuthSessionError(state, "Bad Request")
|
|
return
|
|
}
|
|
if m["state"] != state {
|
|
authErr := codex.NewAuthenticationError(codex.ErrInvalidState, fmt.Errorf("expected %s, got %s", state, m["state"]))
|
|
SetOAuthSessionError(state, "State code error")
|
|
log.Error(codex.GetUserFriendlyMessage(authErr))
|
|
return
|
|
}
|
|
code = m["code"]
|
|
break
|
|
}
|
|
time.Sleep(500 * time.Millisecond)
|
|
}
|
|
|
|
log.Debug("Authorization code received, exchanging for tokens...")
|
|
// Exchange code for tokens using internal auth service
|
|
bundle, errExchange := openaiAuth.ExchangeCodeForTokens(ctx, code, pkceCodes)
|
|
if errExchange != nil {
|
|
authErr := codex.NewAuthenticationError(codex.ErrCodeExchangeFailed, errExchange)
|
|
SetOAuthSessionError(state, oauthSessionErrorWithCause("Failed to exchange authorization code for tokens", errExchange))
|
|
log.Errorf("Failed to exchange authorization code for tokens: %v", authErr)
|
|
return
|
|
}
|
|
|
|
// Extract additional info for filename generation
|
|
claims, _ := codex.ParseJWTToken(bundle.TokenData.IDToken)
|
|
planType := ""
|
|
hashAccountID := ""
|
|
if claims != nil {
|
|
planType = strings.TrimSpace(claims.CodexAuthInfo.ChatgptPlanType)
|
|
if accountID := claims.GetAccountID(); accountID != "" {
|
|
digest := sha256.Sum256([]byte(accountID))
|
|
hashAccountID = hex.EncodeToString(digest[:])[:8]
|
|
}
|
|
}
|
|
|
|
// Create token storage and persist
|
|
tokenStorage := openaiAuth.CreateTokenStorage(bundle)
|
|
fileName := codex.CredentialFileName(tokenStorage.Email, planType, hashAccountID, true)
|
|
record := &coreauth.Auth{
|
|
ID: fileName,
|
|
Provider: "codex",
|
|
FileName: fileName,
|
|
Storage: tokenStorage,
|
|
Metadata: map[string]any{
|
|
"email": tokenStorage.Email,
|
|
"account_id": tokenStorage.AccountID,
|
|
"access_token": tokenStorage.AccessToken,
|
|
"refresh_token": tokenStorage.RefreshToken,
|
|
"id_token": tokenStorage.IDToken,
|
|
"expired": tokenStorage.Expire,
|
|
"last_refresh": tokenStorage.LastRefresh,
|
|
},
|
|
}
|
|
if errGuard := guardOAuthSessionPendingForSave(state, "codex"); errGuard != nil {
|
|
return
|
|
}
|
|
savedPath, errSave := h.saveTokenRecord(ctx, record)
|
|
if errSave != nil {
|
|
SetOAuthSessionError(state, "Failed to save authentication tokens")
|
|
log.Errorf("Failed to save authentication tokens: %v", errSave)
|
|
return
|
|
}
|
|
fmt.Printf("Authentication successful! Token saved to %s\n", savedPath)
|
|
if bundle.APIKey != "" {
|
|
fmt.Println("API key obtained and saved")
|
|
}
|
|
fmt.Println("You can now use Codex services through this CLI")
|
|
CompleteOAuthSession(state)
|
|
}()
|
|
|
|
c.JSON(200, gin.H{"status": "ok", "url": authURL, "state": state})
|
|
}
|
|
|
|
func (h *Handler) RequestAntigravityToken(c *gin.Context) {
|
|
ctx := context.Background()
|
|
ctx = PopulateAuthContext(ctx, c)
|
|
|
|
fmt.Println("Initializing Antigravity authentication...")
|
|
|
|
authSvc := antigravity.NewAntigravityAuth(h.cfg, nil)
|
|
|
|
state, errState := misc.GenerateRandomState()
|
|
if errState != nil {
|
|
log.Errorf("Failed to generate state parameter: %v", errState)
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to generate state parameter"})
|
|
return
|
|
}
|
|
|
|
redirectURI := fmt.Sprintf("http://localhost:%d/oauth-callback", antigravity.CallbackPort)
|
|
authURL := authSvc.BuildAuthURL(state, redirectURI)
|
|
|
|
RegisterOAuthSession(state, "antigravity")
|
|
|
|
isWebUI := isWebUIRequest(c)
|
|
var forwarder *callbackForwarder
|
|
if isWebUI {
|
|
targetURL, errTarget := h.managementCallbackURL("/antigravity/callback")
|
|
if errTarget != nil {
|
|
log.WithError(errTarget).Error("failed to compute antigravity callback target")
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "callback server unavailable"})
|
|
return
|
|
}
|
|
var errStart error
|
|
if forwarder, errStart = startCallbackForwarder(antigravity.CallbackPort, "antigravity", targetURL); errStart != nil {
|
|
log.WithError(errStart).Error("failed to start antigravity callback forwarder")
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to start callback server"})
|
|
return
|
|
}
|
|
}
|
|
|
|
go func() {
|
|
if isWebUI {
|
|
defer stopCallbackForwarderInstance(antigravity.CallbackPort, forwarder)
|
|
}
|
|
|
|
waitFile := filepath.Join(h.cfg.AuthDir, fmt.Sprintf(".oauth-antigravity-%s.oauth", state))
|
|
deadline := time.Now().Add(5 * time.Minute)
|
|
var authCode string
|
|
for {
|
|
if !IsOAuthSessionPending(state, "antigravity") {
|
|
return
|
|
}
|
|
if time.Now().After(deadline) {
|
|
log.Error("oauth flow timed out")
|
|
SetOAuthSessionError(state, "OAuth flow timed out")
|
|
return
|
|
}
|
|
if data, errReadFile := os.ReadFile(waitFile); errReadFile == nil {
|
|
var payload map[string]string
|
|
_ = json.Unmarshal(data, &payload)
|
|
_ = os.Remove(waitFile)
|
|
if errStr := strings.TrimSpace(payload["error"]); errStr != "" {
|
|
log.Errorf("Authentication failed: %s", errStr)
|
|
SetOAuthSessionError(state, "Authentication failed")
|
|
return
|
|
}
|
|
if payloadState := strings.TrimSpace(payload["state"]); payloadState != "" && payloadState != state {
|
|
log.Errorf("Authentication failed: state mismatch")
|
|
SetOAuthSessionError(state, "Authentication failed: state mismatch")
|
|
return
|
|
}
|
|
authCode = strings.TrimSpace(payload["code"])
|
|
if authCode == "" {
|
|
log.Error("Authentication failed: code not found")
|
|
SetOAuthSessionError(state, "Authentication failed: code not found")
|
|
return
|
|
}
|
|
break
|
|
}
|
|
time.Sleep(500 * time.Millisecond)
|
|
}
|
|
|
|
tokenResp, errToken := authSvc.ExchangeCodeForTokens(ctx, authCode, redirectURI)
|
|
if errToken != nil {
|
|
log.Errorf("Failed to exchange token: %v", errToken)
|
|
SetOAuthSessionError(state, "Failed to exchange token")
|
|
return
|
|
}
|
|
|
|
accessToken := strings.TrimSpace(tokenResp.AccessToken)
|
|
if accessToken == "" {
|
|
log.Error("antigravity: token exchange returned empty access token")
|
|
SetOAuthSessionError(state, "Failed to exchange token")
|
|
return
|
|
}
|
|
|
|
email, errInfo := authSvc.FetchUserInfo(ctx, accessToken)
|
|
if errInfo != nil {
|
|
log.Errorf("Failed to fetch user info: %v", errInfo)
|
|
SetOAuthSessionError(state, "Failed to fetch user info")
|
|
return
|
|
}
|
|
email = strings.TrimSpace(email)
|
|
if email == "" {
|
|
log.Error("antigravity: user info returned empty email")
|
|
SetOAuthSessionError(state, "Failed to fetch user info")
|
|
return
|
|
}
|
|
|
|
projectID := ""
|
|
if accessToken != "" {
|
|
fetchedProjectID, errProject := authSvc.FetchProjectID(ctx, accessToken)
|
|
if errProject != nil {
|
|
log.Warnf("antigravity: failed to fetch project ID: %v", errProject)
|
|
} else {
|
|
projectID = fetchedProjectID
|
|
log.Infof("antigravity: obtained project ID %s", util.HideAPIKey(projectID))
|
|
}
|
|
}
|
|
|
|
now := time.Now()
|
|
metadata := map[string]any{
|
|
"type": "antigravity",
|
|
"access_token": tokenResp.AccessToken,
|
|
"refresh_token": tokenResp.RefreshToken,
|
|
"expires_in": tokenResp.ExpiresIn,
|
|
"timestamp": now.UnixMilli(),
|
|
"expired": now.Add(time.Duration(tokenResp.ExpiresIn) * time.Second).Format(time.RFC3339),
|
|
}
|
|
if email != "" {
|
|
metadata["email"] = email
|
|
}
|
|
if projectID != "" {
|
|
metadata["project_id"] = projectID
|
|
}
|
|
|
|
fileName := antigravity.CredentialFileName(email)
|
|
label := strings.TrimSpace(email)
|
|
if label == "" {
|
|
label = "antigravity"
|
|
}
|
|
|
|
record := &coreauth.Auth{
|
|
ID: fileName,
|
|
Provider: "antigravity",
|
|
FileName: fileName,
|
|
Label: label,
|
|
Metadata: metadata,
|
|
}
|
|
if errGuard := guardOAuthSessionPendingForSave(state, "antigravity"); errGuard != nil {
|
|
return
|
|
}
|
|
savedPath, errSave := h.saveTokenRecord(ctx, record)
|
|
if errSave != nil {
|
|
log.Errorf("Failed to save token to file: %v", errSave)
|
|
SetOAuthSessionError(state, "Failed to save token to file")
|
|
return
|
|
}
|
|
|
|
CompleteOAuthSession(state)
|
|
fmt.Printf("Authentication successful! Token saved to %s\n", savedPath)
|
|
if projectID != "" {
|
|
fmt.Printf("Using GCP project: %s\n", util.HideAPIKey(projectID))
|
|
}
|
|
fmt.Println("You can now use Antigravity services through this CLI")
|
|
}()
|
|
|
|
c.JSON(200, gin.H{"status": "ok", "url": authURL, "state": state})
|
|
}
|
|
|
|
func (h *Handler) RequestXAIToken(c *gin.Context) {
|
|
ctx := context.Background()
|
|
ctx = PopulateAuthContext(ctx, c)
|
|
|
|
fmt.Println("Initializing xAI authentication...")
|
|
|
|
state := fmt.Sprintf("xai-%d", time.Now().UnixNano())
|
|
authSvc := xaiauth.NewXAIAuth(h.cfg)
|
|
|
|
deviceFlow, errStartDeviceFlow := authSvc.StartDeviceFlow(ctx)
|
|
if errStartDeviceFlow != nil {
|
|
log.Errorf("Failed to start xAI device flow: %v", errStartDeviceFlow)
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to start device authorization flow"})
|
|
return
|
|
}
|
|
authURL := strings.TrimSpace(deviceFlow.VerificationURIComplete)
|
|
if authURL == "" {
|
|
authURL = strings.TrimSpace(deviceFlow.VerificationURI)
|
|
}
|
|
|
|
RegisterOAuthSession(state, "xai")
|
|
|
|
go func() {
|
|
pollCtx, cancelPoll := context.WithCancel(ctx)
|
|
defer cancelPoll()
|
|
go watchOAuthSessionCancel(pollCtx, cancelPoll, state, "xai")
|
|
|
|
fmt.Println("Waiting for xAI authentication...")
|
|
bundle, errWaitForAuthorization := authSvc.WaitForAuthorization(pollCtx, deviceFlow)
|
|
if errWaitForAuthorization != nil {
|
|
if !IsOAuthSessionPending(state, "xai") {
|
|
return
|
|
}
|
|
log.Errorf("xAI authentication failed: %v", errWaitForAuthorization)
|
|
SetOAuthSessionError(state, oauthSessionErrorWithCause("Authentication failed", errWaitForAuthorization))
|
|
return
|
|
}
|
|
if !IsOAuthSessionPending(state, "xai") {
|
|
return
|
|
}
|
|
|
|
tokenStorage := authSvc.CreateTokenStorage(bundle)
|
|
if tokenStorage == nil || strings.TrimSpace(tokenStorage.AccessToken) == "" {
|
|
log.Error("xAI token exchange returned empty access token")
|
|
SetOAuthSessionError(state, "Failed to exchange token")
|
|
return
|
|
}
|
|
|
|
fileName := xaiauth.CredentialFileName(tokenStorage.Email, tokenStorage.Subject)
|
|
label := strings.TrimSpace(tokenStorage.Email)
|
|
if label == "" {
|
|
label = "xAI"
|
|
}
|
|
|
|
metadata := map[string]any{
|
|
"type": "xai",
|
|
"access_token": tokenStorage.AccessToken,
|
|
"refresh_token": tokenStorage.RefreshToken,
|
|
"id_token": tokenStorage.IDToken,
|
|
"token_type": tokenStorage.TokenType,
|
|
"expires_in": tokenStorage.ExpiresIn,
|
|
"expired": tokenStorage.Expire,
|
|
"last_refresh": tokenStorage.LastRefresh,
|
|
"base_url": tokenStorage.BaseURL,
|
|
"token_endpoint": tokenStorage.TokenEndpoint,
|
|
"auth_kind": "oauth",
|
|
}
|
|
if tokenStorage.Email != "" {
|
|
metadata["email"] = tokenStorage.Email
|
|
}
|
|
if tokenStorage.Subject != "" {
|
|
metadata["sub"] = tokenStorage.Subject
|
|
}
|
|
|
|
record := &coreauth.Auth{
|
|
ID: fileName,
|
|
Provider: "xai",
|
|
FileName: fileName,
|
|
Label: label,
|
|
Storage: tokenStorage,
|
|
Metadata: metadata,
|
|
Attributes: map[string]string{
|
|
"auth_kind": "oauth",
|
|
"base_url": tokenStorage.BaseURL,
|
|
},
|
|
}
|
|
if errGuard := guardOAuthSessionPendingForSave(state, "xai"); errGuard != nil {
|
|
return
|
|
}
|
|
savedPath, errSave := h.saveTokenRecord(ctx, record)
|
|
if errSave != nil {
|
|
log.Errorf("Failed to save xAI token to file: %v", errSave)
|
|
SetOAuthSessionError(state, "Failed to save token to file")
|
|
return
|
|
}
|
|
|
|
CompleteOAuthSession(state)
|
|
fmt.Printf("Authentication successful! Token saved to %s\n", savedPath)
|
|
fmt.Println("You can now use xAI services through this CLI")
|
|
}()
|
|
|
|
response := gin.H{"status": "ok", "url": authURL, "state": state, "flow": "device"}
|
|
if userCode := strings.TrimSpace(deviceFlow.UserCode); userCode != "" {
|
|
response["user_code"] = userCode
|
|
}
|
|
if deviceFlow.ExpiresIn > 0 {
|
|
response["expires_in"] = deviceFlow.ExpiresIn
|
|
} else {
|
|
response["expires_in"] = int(xaiauth.MaxPollDuration / time.Second)
|
|
}
|
|
c.JSON(200, response)
|
|
}
|
|
|
|
func (h *Handler) RequestKimiToken(c *gin.Context) {
|
|
ctx := context.Background()
|
|
ctx = PopulateAuthContext(ctx, c)
|
|
|
|
fmt.Println("Initializing Kimi authentication...")
|
|
|
|
state := fmt.Sprintf("kmi-%d", time.Now().UnixNano())
|
|
// Initialize Kimi auth service
|
|
kimiAuth := kimi.NewKimiAuth(h.cfg)
|
|
|
|
// Generate authorization URL
|
|
deviceFlow, errStartDeviceFlow := kimiAuth.StartDeviceFlow(ctx)
|
|
if errStartDeviceFlow != nil {
|
|
log.Errorf("Failed to generate authorization URL: %v", errStartDeviceFlow)
|
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to generate authorization url"})
|
|
return
|
|
}
|
|
authURL := deviceFlow.VerificationURIComplete
|
|
if authURL == "" {
|
|
authURL = deviceFlow.VerificationURI
|
|
}
|
|
|
|
RegisterOAuthSession(state, "kimi")
|
|
|
|
go func() {
|
|
pollCtx, cancelPoll := context.WithCancel(ctx)
|
|
defer cancelPoll()
|
|
go watchOAuthSessionCancel(pollCtx, cancelPoll, state, "kimi")
|
|
|
|
fmt.Println("Waiting for authentication...")
|
|
authBundle, errWaitForAuthorization := kimiAuth.WaitForAuthorization(pollCtx, deviceFlow)
|
|
if errWaitForAuthorization != nil {
|
|
if !IsOAuthSessionPending(state, "kimi") {
|
|
return
|
|
}
|
|
SetOAuthSessionError(state, oauthSessionErrorWithCause("Authentication failed", errWaitForAuthorization))
|
|
fmt.Printf("Authentication failed: %v\n", errWaitForAuthorization)
|
|
return
|
|
}
|
|
if !IsOAuthSessionPending(state, "kimi") {
|
|
return
|
|
}
|
|
|
|
// Create token storage
|
|
tokenStorage := kimiAuth.CreateTokenStorage(authBundle)
|
|
|
|
metadata := map[string]any{
|
|
"type": "kimi",
|
|
"access_token": authBundle.TokenData.AccessToken,
|
|
"refresh_token": authBundle.TokenData.RefreshToken,
|
|
"token_type": authBundle.TokenData.TokenType,
|
|
"scope": authBundle.TokenData.Scope,
|
|
"timestamp": time.Now().UnixMilli(),
|
|
}
|
|
if authBundle.TokenData.ExpiresAt > 0 {
|
|
expired := time.Unix(authBundle.TokenData.ExpiresAt, 0).UTC().Format(time.RFC3339)
|
|
metadata["expired"] = expired
|
|
}
|
|
if strings.TrimSpace(authBundle.DeviceID) != "" {
|
|
metadata["device_id"] = strings.TrimSpace(authBundle.DeviceID)
|
|
}
|
|
|
|
fileName := fmt.Sprintf("kimi-%d.json", time.Now().UnixMilli())
|
|
record := &coreauth.Auth{
|
|
ID: fileName,
|
|
Provider: "kimi",
|
|
FileName: fileName,
|
|
Label: "Kimi User",
|
|
Storage: tokenStorage,
|
|
Metadata: metadata,
|
|
}
|
|
if errGuard := guardOAuthSessionPendingForSave(state, "kimi"); errGuard != nil {
|
|
return
|
|
}
|
|
savedPath, errSave := h.saveTokenRecord(ctx, record)
|
|
if errSave != nil {
|
|
log.Errorf("Failed to save authentication tokens: %v", errSave)
|
|
SetOAuthSessionError(state, "Failed to save authentication tokens")
|
|
return
|
|
}
|
|
|
|
fmt.Printf("Authentication successful! Token saved to %s\n", savedPath)
|
|
fmt.Println("You can now use Kimi services through this CLI")
|
|
CompleteOAuthSession(state)
|
|
}()
|
|
|
|
response := gin.H{"status": "ok", "url": authURL, "state": state, "flow": "device"}
|
|
if userCode := strings.TrimSpace(deviceFlow.UserCode); userCode != "" {
|
|
response["user_code"] = userCode
|
|
}
|
|
if deviceFlow.ExpiresIn > 0 {
|
|
response["expires_in"] = deviceFlow.ExpiresIn
|
|
}
|
|
c.JSON(200, response)
|
|
}
|
|
|
|
// watchOAuthSessionCancel cancels pollCtx once the OAuth session is no longer pending.
|
|
func watchOAuthSessionCancel(pollCtx context.Context, cancel context.CancelFunc, state, provider string) {
|
|
if cancel == nil {
|
|
return
|
|
}
|
|
ticker := time.NewTicker(2 * time.Second)
|
|
defer ticker.Stop()
|
|
for {
|
|
select {
|
|
case <-pollCtx.Done():
|
|
return
|
|
case <-ticker.C:
|
|
if !IsOAuthSessionPending(state, provider) {
|
|
cancel()
|
|
return
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
// CancelAuthSession cancels a pending OAuth session identified by state.
|
|
// Protected by management auth. Safe for both callback and device-code flows:
|
|
// waiters check IsOAuthSessionPending and exit without saving credentials.
|
|
func (h *Handler) CancelAuthSession(c *gin.Context) {
|
|
state := strings.TrimSpace(c.Query("state"))
|
|
if state == "" {
|
|
c.JSON(http.StatusBadRequest, gin.H{"status": "error", "error": "missing state"})
|
|
return
|
|
}
|
|
if err := ValidateOAuthState(state); err != nil {
|
|
c.JSON(http.StatusBadRequest, gin.H{"status": "error", "error": "invalid state"})
|
|
return
|
|
}
|
|
cancelled := CancelOAuthSession(state)
|
|
c.JSON(http.StatusOK, gin.H{"status": "ok", "cancelled": cancelled})
|
|
}
|
|
|
|
func (h *Handler) GetAuthStatus(c *gin.Context) {
|
|
state := strings.TrimSpace(c.Query("state"))
|
|
if state == "" {
|
|
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
|
return
|
|
}
|
|
if err := ValidateOAuthState(state); err != nil {
|
|
c.JSON(http.StatusBadRequest, gin.H{"status": "error", "error": "invalid state"})
|
|
return
|
|
}
|
|
|
|
provider, status, isPlugin, metadata, completed, ok := GetOAuthSessionDetails(state)
|
|
if !ok {
|
|
c.JSON(http.StatusOK, gin.H{"status": "error", "error": "unknown or expired state"})
|
|
return
|
|
}
|
|
if completed {
|
|
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
|
return
|
|
}
|
|
if status != "" {
|
|
c.JSON(http.StatusOK, gin.H{"status": "error", "error": status})
|
|
return
|
|
}
|
|
h.mu.Lock()
|
|
host := h.pluginHost
|
|
h.mu.Unlock()
|
|
if isPlugin && host != nil && host.HasAuthProvider(provider) {
|
|
ctx := PopulateAuthContext(context.Background(), c)
|
|
resp, handled, errPoll := host.PollLogin(ctx, provider, state, metadata)
|
|
if handled {
|
|
if errPoll != nil {
|
|
message := strings.TrimSpace(errPoll.Error())
|
|
if message == "" {
|
|
message = "Authentication failed"
|
|
}
|
|
SetOAuthSessionError(state, message)
|
|
c.JSON(http.StatusOK, gin.H{"status": "error", "error": message})
|
|
return
|
|
}
|
|
switch resp.Status {
|
|
case "", pluginapi.AuthLoginStatusPending:
|
|
c.JSON(http.StatusOK, gin.H{"status": "wait"})
|
|
return
|
|
case pluginapi.AuthLoginStatusError:
|
|
message := strings.TrimSpace(resp.Message)
|
|
if message == "" {
|
|
message = "Authentication failed"
|
|
}
|
|
SetOAuthSessionError(state, message)
|
|
c.JSON(http.StatusOK, gin.H{"status": "error", "error": message})
|
|
return
|
|
case pluginapi.AuthLoginStatusSuccess:
|
|
records := pluginLoginPollAuths(host, resp)
|
|
if len(records) == 0 {
|
|
SetOAuthSessionError(state, "Authentication failed")
|
|
c.JSON(http.StatusOK, gin.H{"status": "error", "error": "Authentication failed"})
|
|
return
|
|
}
|
|
if errSave := h.savePluginLoginRecords(ctx, records); errSave != nil {
|
|
log.WithError(errSave).WithField("provider", provider).Error("failed to save plugin auth tokens")
|
|
SetOAuthSessionError(state, "Failed to save authentication tokens")
|
|
c.JSON(http.StatusOK, gin.H{"status": "error", "error": "Failed to save authentication tokens"})
|
|
return
|
|
}
|
|
CompleteOAuthSession(state)
|
|
c.JSON(http.StatusOK, gin.H{"status": "ok"})
|
|
return
|
|
default:
|
|
c.JSON(http.StatusOK, gin.H{"status": "wait"})
|
|
return
|
|
}
|
|
}
|
|
}
|
|
c.JSON(http.StatusOK, gin.H{"status": "wait"})
|
|
}
|
|
|
|
func pluginLoginPollAuths(host *pluginhost.Host, resp pluginapi.AuthLoginPollResponse) []*coreauth.Auth {
|
|
if host == nil {
|
|
return nil
|
|
}
|
|
authDatas := resp.Auths
|
|
if len(authDatas) == 0 {
|
|
authDatas = []pluginapi.AuthData{resp.Auth}
|
|
}
|
|
records := make([]*coreauth.Auth, 0, len(authDatas))
|
|
for _, authData := range authDatas {
|
|
record := host.AuthDataToCoreAuth(authData, "", "")
|
|
if record == nil {
|
|
return nil
|
|
}
|
|
records = append(records, record)
|
|
}
|
|
return records
|
|
}
|
|
|
|
func (h *Handler) savePluginLoginRecords(ctx context.Context, records []*coreauth.Auth) error {
|
|
savedPaths := make([]string, 0, len(records))
|
|
for _, record := range records {
|
|
savedPath, errSave := h.saveTokenRecord(ctx, record)
|
|
if strings.TrimSpace(savedPath) != "" {
|
|
savedPaths = append(savedPaths, savedPath)
|
|
}
|
|
if errSave != nil {
|
|
h.rollbackSavedTokenRecords(ctx, savedPaths)
|
|
return errSave
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (h *Handler) rollbackSavedTokenRecords(ctx context.Context, savedPaths []string) {
|
|
for i := len(savedPaths) - 1; i >= 0; i-- {
|
|
path := strings.TrimSpace(savedPaths[i])
|
|
if path == "" {
|
|
continue
|
|
}
|
|
if errDelete := h.deleteTokenRecord(ctx, path); errDelete != nil {
|
|
log.WithError(errDelete).WithField("path", path).Warn("failed to roll back plugin auth token")
|
|
}
|
|
h.removeAuthsForPath(ctx, path, path)
|
|
}
|
|
}
|
|
|
|
// PopulateAuthContext extracts request info and adds it to the context
|
|
func PopulateAuthContext(ctx context.Context, c *gin.Context) context.Context {
|
|
info := &coreauth.RequestInfo{
|
|
Query: c.Request.URL.Query(),
|
|
Headers: c.Request.Header,
|
|
}
|
|
return coreauth.WithRequestInfo(ctx, info)
|
|
}
|