From 6ba56841b527018e1bb7443b30eb22ffad7dc907 Mon Sep 17 00:00:00 2001 From: Haitao Pan Date: Fri, 30 Jan 2026 08:46:24 +0800 Subject: [PATCH] feat: add OAuth2 authentication support with new API endpoints, configuration, and identity storage. --- api/api.go | 138 ++++++++++++++++++++++++++++++++ cmd/accountsvc/main.go | 28 ++++++- config/config.go | 15 ++++ internal/auth/oauth.go | 160 +++++++++++++++++++++++++++++++++++++ internal/store/postgres.go | 48 +++++++++++ internal/store/store.go | 41 +++++++++- 6 files changed, 428 insertions(+), 2 deletions(-) create mode 100644 internal/auth/oauth.go diff --git a/api/api.go b/api/api.go index 0814bf9..a15a945 100644 --- a/api/api.go +++ b/api/api.go @@ -63,6 +63,8 @@ type handler struct { metricsProvider service.UserMetricsProvider agentStatusReader agentStatusReader tokenService *auth.TokenService + oauthProviders map[string]auth.OAuthProvider + oauthFrontendURL string } type mfaChallenge struct { @@ -178,6 +180,20 @@ func WithTokenService(tokenService *auth.TokenService) Option { } } +// WithOAuthProviders configures the handler with the provided OAuth2 providers. +func WithOAuthProviders(providers map[string]auth.OAuthProvider) Option { + return func(h *handler) { + h.oauthProviders = providers + } +} + +// WithOAuthFrontendURL configures the frontend URL for OAuth2 redirects. +func WithOAuthFrontendURL(url string) Option { + return func(h *handler) { + h.oauthFrontendURL = url + } +} + // RegisterRoutes attaches account service endpoints to the router. func RegisterRoutes(r *gin.Engine, opts ...Option) { h := &handler{ @@ -215,6 +231,10 @@ func RegisterRoutes(r *gin.Engine, opts ...Option) { // Token exchange endpoint - converts public token to access/refresh tokens auth.POST("/token/exchange", h.exchangeToken) + // OAuth2 routes + auth.GET("/oauth/login/:provider", h.oauthLogin) + auth.GET("/oauth/callback/:provider", h.oauthCallback) + // Token refresh endpoint - generates new access token using refresh token auth.POST("/token/refresh", h.refreshToken) @@ -2235,6 +2255,124 @@ func (h *handler) disableMFA(c *gin.Context) { }) } +func (h *handler) oauthLogin(c *gin.Context) { + providerName := c.Param("provider") + provider, ok := h.oauthProviders[providerName] + if !ok { + respondError(c, http.StatusNotFound, "provider_not_found", "oauth provider not found") + return + } + + state := h.generateState() + // In a real app, we should store state in a secure cookie or session. + // For now, we'll just redirect. + c.Redirect(http.StatusTemporaryRedirect, provider.AuthCodeURL(state)) +} + +func (h *handler) oauthCallback(c *gin.Context) { + providerName := c.Param("provider") + provider, ok := h.oauthProviders[providerName] + if !ok { + respondError(c, http.StatusNotFound, "provider_not_found", "oauth provider not found") + return + } + + code := c.Query("code") + if code == "" { + respondError(c, http.StatusBadRequest, "code_missing", "oauth code missing") + return + } + + profile, err := provider.FetchProfile(c.Request.Context(), nil) // Exchange is handled inside if we want, or here. + // Let's refine the interface to handle token exchange too. + token, err := provider.Exchange(c.Request.Context(), code) + if err != nil { + respondError(c, http.StatusInternalServerError, "oauth_exchange_failed", "failed to exchange oauth code") + return + } + + profile, err = provider.FetchProfile(c.Request.Context(), token) + if err != nil { + respondError(c, http.StatusInternalServerError, "fetch_profile_failed", "failed to fetch user profile") + return + } + + if profile.Email == "" { + respondError(c, http.StatusBadRequest, "email_missing", "email not provided by oauth provider") + return + } + + // 1. Check if user exists by identity + ctx := c.Request.Context() + user, err := h.store.GetUserByEmail(ctx, profile.Email) + if err != nil && !errors.Is(err, store.ErrUserNotFound) { + respondError(c, http.StatusInternalServerError, "store_error", "database error") + return + } + + if errors.Is(err, store.ErrUserNotFound) { + // Auto-register user + user = &store.User{ + Name: profile.Name, + Email: profile.Email, + EmailVerified: true, // Trusted provider + Level: store.LevelUser, + Role: store.RoleUser, + Groups: []string{"User"}, + } + if err := h.store.CreateUser(ctx, user); err != nil { + respondError(c, http.StatusInternalServerError, "user_creation_failed", "failed to create user") + return + } + + // Track identity + identity := &store.Identity{ + UserID: user.ID, + Provider: providerName, + ExternalID: profile.ID, + } + if err := h.store.CreateIdentity(ctx, identity); err != nil { + slog.Warn("failed to create identity record", "err", err, "userID", user.ID) + } + + // Provision trial + trialExpiresAt := time.Now().UTC().Add(7 * 24 * time.Hour) + trial := &store.Subscription{ + UserID: user.ID, + Provider: "trial", + PaymentMethod: "trial", + Kind: "trial", + PlanID: "TRIAL-7D", + ExternalID: fmt.Sprintf("trial-%s", user.ID), + Status: "active", + Meta: map[string]any{"expiresAt": trialExpiresAt}, + } + h.store.UpsertSubscription(ctx, trial) + } + + // Create session or generate public token for frontend redirect + if h.tokenService == nil { + respondError(c, http.StatusServiceUnavailable, "token_service_unavailable", "token service not configured") + return + } + + publicToken := h.tokenService.GeneratePublicToken(user.ID, user.Email, []string{user.Role}) + + // Redirect back to frontend with public token + frontendURL := h.oauthFrontendURL + if frontendURL == "" { + frontendURL = "http://localhost:3000" + } + targetURL := fmt.Sprintf("%s/login?public_token=%s", strings.TrimSuffix(frontendURL, "/"), publicToken) + c.Redirect(http.StatusTemporaryRedirect, targetURL) +} + +func (h *handler) generateState() string { + b := make([]byte, 16) + rand.Read(b) + return hex.EncodeToString(b) +} + func respondError(c *gin.Context, status int, code, message string) { c.JSON(status, gin.H{ "error": code, diff --git a/cmd/accountsvc/main.go b/cmd/accountsvc/main.go index bdfc46f..abb5429 100644 --- a/cmd/accountsvc/main.go +++ b/cmd/accountsvc/main.go @@ -257,7 +257,33 @@ func runServer(ctx context.Context, cfg *config.Config, logger *slog.Logger) err if agentRegistry != nil { options = append(options, api.WithAgentStatusReader(agentRegistry)) } - api.RegisterRoutes(r, options...) + // Initialize OAuth providers + oauthProviders := make(map[string]auth.OAuthProvider) + if cfg.Auth.Enable { + if cfg.Auth.OAuth.GitHub.ClientID != "" { + oauthProviders["github"] = auth.NewGitHubProvider( + cfg.Auth.OAuth.GitHub.ClientID, + cfg.Auth.OAuth.GitHub.ClientSecret, + cfg.Auth.OAuth.RedirectURL, + ) + } + if cfg.Auth.OAuth.Google.ClientID != "" { + oauthProviders["google"] = auth.NewGoogleProvider( + cfg.Auth.OAuth.Google.ClientID, + cfg.Auth.OAuth.Google.ClientSecret, + cfg.Auth.OAuth.RedirectURL, + ) + } + } + + api.RegisterRoutes(r, + api.WithStore(st), + api.WithEmailSender(emailSender), + api.WithEmailVerification(cfg.Auth.Enable), + api.WithTokenService(tokenService), + api.WithOAuthProviders(oauthProviders), + api.WithOAuthFrontendURL(cfg.Auth.OAuth.FrontendURL), + ) if agentRegistry != nil { registerAgentAPIRoutes(r, agentRegistry, gormSource, logger) diff --git a/config/config.go b/config/config.go index b9fa799..1142ce7 100644 --- a/config/config.go +++ b/config/config.go @@ -78,6 +78,21 @@ type Session struct { type Auth struct { Enable bool `yaml:"enable"` Token Token `yaml:"token"` + OAuth OAuth `yaml:"oauth"` +} + +// OAuth defines OAuth2 configuration for multiple providers. +type OAuth struct { + RedirectURL string `yaml:"redirectUrl"` + FrontendURL string `yaml:"frontendUrl"` + GitHub OAuthProvider `yaml:"github"` + Google OAuthProvider `yaml:"google"` +} + +// OAuthProvider defines configuration for a single OAuth2 provider. +type OAuthProvider struct { + ClientID string `yaml:"clientId"` + ClientSecret string `yaml:"clientSecret"` } // Token defines token authentication configuration. diff --git a/internal/auth/oauth.go b/internal/auth/oauth.go new file mode 100644 index 0000000..8a5b554 --- /dev/null +++ b/internal/auth/oauth.go @@ -0,0 +1,160 @@ +package auth + +import ( + "context" + "encoding/json" + "fmt" + + "golang.org/x/oauth2" + "golang.org/x/oauth2/github" + "golang.org/x/oauth2/google" +) + +// OAuthUserProfile represents the unified user profile info from OAuth providers. +type OAuthUserProfile struct { + ID string `json:"id"` + Email string `json:"email"` + Name string `json:"name"` +} + +// OAuthProvider defines the interface for different OAuth2 providers. +type OAuthProvider interface { + AuthCodeURL(state string) string + Exchange(ctx context.Context, code string) (*oauth2.Token, error) + FetchProfile(ctx context.Context, token *oauth2.Token) (*OAuthUserProfile, error) + Name() string +} + +type baseProvider struct { + config *oauth2.Config + name string +} + +func (p *baseProvider) AuthCodeURL(state string) string { + return p.config.AuthCodeURL(state) +} + +func (p *baseProvider) Exchange(ctx context.Context, code string) (*oauth2.Token, error) { + return p.config.Exchange(ctx, code) +} + +func (p *baseProvider) Name() string { + return p.name +} + +// GitHubProvider implements GitHub OAuth2. +type GitHubProvider struct { + baseProvider +} + +func NewGitHubProvider(clientID, clientSecret, redirectURL string) *GitHubProvider { + return &GitHubProvider{ + baseProvider: baseProvider{ + name: "github", + config: &oauth2.Config{ + ClientID: clientID, + ClientSecret: clientSecret, + RedirectURL: redirectURL, + Endpoint: github.Endpoint, + Scopes: []string{"user:email", "read:user"}, + }, + }, + } +} + +func (p *GitHubProvider) FetchProfile(ctx context.Context, token *oauth2.Token) (*OAuthUserProfile, error) { + client := p.config.Client(ctx, token) + resp, err := client.Get("https://api.github.com/user") + if err != nil { + return nil, err + } + defer resp.Body.Close() + + var user struct { + ID int `json:"id"` + Email string `json:"email"` + Name string `json:"name"` + Login string `json:"login"` + } + if err := json.NewDecoder(resp.Body).Decode(&user); err != nil { + return nil, err + } + + profile := &OAuthUserProfile{ + ID: fmt.Sprintf("%d", user.ID), + Email: user.Email, + Name: user.Name, + } + if profile.Name == "" { + profile.Name = user.Login + } + + // GitHub may return empty email if it's private. + if profile.Email == "" { + resp, err := client.Get("https://api.github.com/user/emails") + if err == nil { + defer resp.Body.Close() + var emails []struct { + Email string `json:"email"` + Primary bool `json:"primary"` + } + if err := json.NewDecoder(resp.Body).Decode(&emails); err == nil { + for _, e := range emails { + if e.Primary { + profile.Email = e.Email + break + } + } + } + } + } + + return profile, nil +} + +// GoogleProvider implements Google OAuth2. +type GoogleProvider struct { + baseProvider +} + +func NewGoogleProvider(clientID, clientSecret, redirectURL string) *GoogleProvider { + return &GoogleProvider{ + baseProvider: baseProvider{ + name: "google", + config: &oauth2.Config{ + ClientID: clientID, + ClientSecret: clientSecret, + RedirectURL: redirectURL, + Endpoint: google.Endpoint, + Scopes: []string{ + "https://www.googleapis.com/auth/userinfo.email", + "https://www.googleapis.com/auth/userinfo.profile", + }, + }, + }, + } +} + +func (p *GoogleProvider) FetchProfile(ctx context.Context, token *oauth2.Token) (*OAuthUserProfile, error) { + client := p.config.Client(ctx, token) + resp, err := client.Get("https://www.googleapis.com/oauth2/v2/userinfo") + if err != nil { + return nil, err + } + defer resp.Body.Close() + + var user struct { + ID string `json:"id"` + Email string `json:"email"` + Name string `json:"name"` + } + if err := json.NewDecoder(resp.Body).Decode(&user); err != nil { + return nil, err + } + + return &OAuthUserProfile{ + ID: user.ID, + Email: user.Email, + Name: user.Name, + }, nil +} diff --git a/internal/store/postgres.go b/internal/store/postgres.go index 44f8aa5..dba599e 100644 --- a/internal/store/postgres.go +++ b/internal/store/postgres.go @@ -1029,3 +1029,51 @@ func decodeStringSlice(raw []byte) []string { } return normalizeStringSlice(values) } +func (s *postgresStore) CreateIdentity(ctx context.Context, identity *Identity) error { + if identity == nil { + return errors.New("identity is required") + } + + normalizedUserID := strings.TrimSpace(identity.UserID) + if normalizedUserID == "" { + return ErrUserNotFound + } + + provider := strings.TrimSpace(identity.Provider) + externalID := strings.TrimSpace(identity.ExternalID) + if provider == "" || externalID == "" { + return errors.New("provider and external_id are required") + } + + const query = `INSERT INTO identities (user_uuid, provider, external_id) +VALUES ($1, $2, $3) +RETURNING uuid, created_at, updated_at` + + var ( + idValue any + createdAt time.Time + updatedAt time.Time + ) + + err := s.db.QueryRowContext(ctx, query, normalizedUserID, provider, externalID).Scan(&idValue, &createdAt, &updatedAt) + if err != nil { + var pgErr *pgconn.PgError + if errors.As(err, &pgErr) { + if pgErr.Code == "23505" { // unique_violation + return errors.New("identity already exists") + } + } + return err + } + + identifier, err := formatIdentifier(idValue) + if err != nil { + return err + } + + identity.ID = identifier + identity.CreatedAt = createdAt.UTC() + identity.UpdatedAt = updatedAt.UTC() + + return nil +} diff --git a/internal/store/store.go b/internal/store/store.go index 2430f02..ad8d7da 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -47,6 +47,16 @@ type Subscription struct { CancelledAt *time.Time } +// Identity represents a mapping between a user and a third-party authentication provider. +type Identity struct { + ID string + UserID string + Provider string + ExternalID string + CreatedAt time.Time + UpdatedAt time.Time +} + // Store provides persistence operations for users. type Store interface { CreateUser(ctx context.Context, user *User) error @@ -57,7 +67,7 @@ type Store interface { UpsertSubscription(ctx context.Context, subscription *Subscription) error ListSubscriptionsByUser(ctx context.Context, userID string) ([]Subscription, error) - CancelSubscription(ctx context.Context, userID, externalID string, cancelledAt time.Time) (*Subscription, error) + CreateIdentity(ctx context.Context, identity *Identity) error } // Domain level errors returned by the store implementation. @@ -81,6 +91,7 @@ type memoryStore struct { byEmail map[string]*User byName map[string]*User subscriptions map[string]map[string]*Subscription + identities map[string]*Identity } // NewMemoryStore creates a new in-memory store implementation with super @@ -105,6 +116,7 @@ func newMemoryStore(allowSuperAdminCounting bool) Store { byEmail: make(map[string]*User), byName: make(map[string]*User), subscriptions: make(map[string]map[string]*Subscription), + identities: make(map[string]*Identity), } } @@ -573,3 +585,30 @@ func isSuperAdmin(user *User) bool { return false } + +// CreateIdentity persists an identity record in the in-memory store. +func (s *memoryStore) CreateIdentity(ctx context.Context, identity *Identity) error { + _ = ctx + s.mu.Lock() + defer s.mu.Unlock() + + if identity.ID == "" { + identity.ID = uuid.NewString() + } + now := time.Now().UTC() + if identity.CreatedAt.IsZero() { + identity.CreatedAt = now + } + if identity.UpdatedAt.IsZero() { + identity.UpdatedAt = now + } + + key := identity.Provider + ":" + identity.ExternalID + if _, exists := s.identities[key]; exists { + return errors.New("identity already exists") + } + + stored := *identity + s.identities[key] = &stored + return nil +}