diff --git a/go.mod b/go.mod index 4d05ae46..8e52b915 100644 --- a/go.mod +++ b/go.mod @@ -27,6 +27,7 @@ require ( github.com/matoous/go-nanoid/v2 v2.1.0 github.com/mozillazg/go-pinyin v0.20.0 github.com/pkoukk/tiktoken-go v0.1.7 + github.com/pquerna/otp v1.5.0 github.com/rhysd/go-github-selfupdate v1.2.3 github.com/spf13/cast v1.9.2 github.com/spf13/cobra v1.9.1 @@ -117,7 +118,6 @@ require ( github.com/pelletier/go-toml/v2 v2.2.4 // indirect github.com/pkg/errors v0.9.1 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect - github.com/pquerna/otp v1.5.0 // indirect github.com/qdrant/go-client v1.14.0 // indirect github.com/richardlehane/mscfb v1.0.4 // indirect github.com/richardlehane/msoleps v1.0.4 // indirect diff --git a/openapi/config.go b/openapi/config.go index bd28b6e5..0deb31e3 100644 --- a/openapi/config.go +++ b/openapi/config.go @@ -2,6 +2,7 @@ package openapi import ( "errors" + "fmt" "path/filepath" "strings" "time" @@ -134,6 +135,11 @@ func (config *Config) UnmarshalJSON(data []byte) error { Features: tempConfig.OAuth.Features, } + fmt.Println("----debug----") + fmt.Println("tempConfig.OAuth.IssuerURL", tempConfig.OAuth.IssuerURL) + fmt.Println("config.OAuth.IssuerURL", config.OAuth.IssuerURL) + fmt.Println("----debug----") + // Convert signing config with duration parsing config.OAuth.Signing = types.SigningConfig{ SigningCertPath: tempConfig.OAuth.Signing.SigningCertPath, @@ -379,7 +385,7 @@ func (config *Config) OAuthConfig(appConfig config.Config) (*oauth.Config, error ClientProvider: clientProvider, Cache: cacheStore, Store: dataStore, - IssuerURL: config.BaseURL, + IssuerURL: config.OAuth.IssuerURL, Signing: signingConfig, // Use the converted signing config Token: config.OAuth.Token, Security: config.OAuth.Security, diff --git a/openapi/oauth/client.go b/openapi/oauth/client.go index d42c3793..e1b3fb01 100644 --- a/openapi/oauth/client.go +++ b/openapi/oauth/client.go @@ -47,12 +47,17 @@ func (s *Service) DynamicClientRegistration(ctx context.Context, request *types. return nil, err } - // Generate client ID and secret - clientID, err := s.generateClientID() - if err != nil { - return nil, &types.ErrorResponse{ - Code: types.ErrorServerError, - ErrorDescription: "Failed to generate client ID", + // Generate client ID and secret (use the client ID from the request if provided or generate a new one) + clientID := request.ClientID + var err error + if clientID == "" { + var err error + clientID, err = s.GenerateClientID() + if err != nil { + return nil, &types.ErrorResponse{ + Code: types.ErrorServerError, + ErrorDescription: "Failed to generate client ID", + } } } @@ -64,7 +69,7 @@ func (s *Service) DynamicClientRegistration(ctx context.Context, request *types. request.TokenEndpointAuthMethod == types.TokenEndpointAuthPost || request.TokenEndpointAuthMethod == types.TokenEndpointAuthJWT { clientType = types.ClientTypeConfidential - clientSecret, err = s.generateClientSecret() + clientSecret, err = s.GenerateClientSecret() if err != nil { return nil, &types.ErrorResponse{ Code: types.ErrorServerError, @@ -144,8 +149,8 @@ func (s *Service) DynamicClientRegistration(ctx context.Context, request *types. return response, nil } -// generateClientID generates a random client ID -func (s *Service) generateClientID() (string, error) { +// GenerateClientID generates a random client ID +func (s *Service) GenerateClientID() (string, error) { length := s.config.Client.ClientIDLength if length == 0 { length = 32 @@ -160,8 +165,30 @@ func (s *Service) generateClientID() (string, error) { return strings.TrimRight(base64.URLEncoding.EncodeToString(bytes), "="), nil } -// generateClientSecret generates a random client secret -func (s *Service) generateClientSecret() (string, error) { +// ValidateClientID validates the client ID +func (s *Service) ValidateClientID(clientID string) error { + if clientID == "" { + return &types.ErrorResponse{ + Code: types.ErrorInvalidRequest, + ErrorDescription: "Client ID is required", + } + } + length := s.config.Client.ClientIDLength + if length == 0 { + length = 32 + } + + if len(clientID) != length { + return &types.ErrorResponse{ + Code: types.ErrorInvalidRequest, + ErrorDescription: fmt.Sprintf("Client ID must be %d characters long", length), + } + } + return nil +} + +// GenerateClientSecret generates a random client secret +func (s *Service) GenerateClientSecret() (string, error) { length := s.config.Client.ClientSecretLength if length == 0 { length = 64 @@ -179,7 +206,7 @@ func (s *Service) generateClientSecret() (string, error) { // validateDynamicClientRegistrationRequest validates the dynamic client registration request func (s *Service) validateDynamicClientRegistrationRequest(request *types.DynamicClientRegistrationRequest) error { // Validate redirect URIs - if len(request.RedirectURIs) == 0 { + if len(request.RedirectURIs) == 0 && (strings.Contains(request.Scope, "openid") || strings.Contains(request.Scope, "profile") || strings.Contains(request.Scope, "email")) { return &types.ErrorResponse{ Code: types.ErrorInvalidRequest, ErrorDescription: "At least one redirect URI is required", diff --git a/openapi/oauth/client_test.go b/openapi/oauth/client_test.go index 788708a9..aba8533b 100644 --- a/openapi/oauth/client_test.go +++ b/openapi/oauth/client_test.go @@ -417,7 +417,7 @@ func TestGenerateClientID(t *testing.T) { defer cleanup() t.Run("generate client ID with default length", func(t *testing.T) { - clientID, err := service.generateClientID() + clientID, err := service.GenerateClientID() assert.NoError(t, err) assert.NotEmpty(t, clientID) assert.Greater(t, len(clientID), 0) @@ -432,7 +432,7 @@ func TestGenerateClientID(t *testing.T) { clientIDs := make(map[string]bool) for i := 0; i < 100; i++ { - clientID, err := service.generateClientID() + clientID, err := service.GenerateClientID() assert.NoError(t, err) assert.NotEmpty(t, clientID) @@ -450,7 +450,7 @@ func TestGenerateClientID(t *testing.T) { service.config.Client.ClientIDLength = originalLength }() - clientID, err := service.generateClientID() + clientID, err := service.GenerateClientID() assert.NoError(t, err) assert.NotEmpty(t, clientID) @@ -465,7 +465,7 @@ func TestGenerateClientSecret(t *testing.T) { defer cleanup() t.Run("generate client secret with default length", func(t *testing.T) { - clientSecret, err := service.generateClientSecret() + clientSecret, err := service.GenerateClientSecret() assert.NoError(t, err) assert.NotEmpty(t, clientSecret) assert.Greater(t, len(clientSecret), 0) @@ -480,7 +480,7 @@ func TestGenerateClientSecret(t *testing.T) { clientSecrets := make(map[string]bool) for i := 0; i < 100; i++ { - clientSecret, err := service.generateClientSecret() + clientSecret, err := service.GenerateClientSecret() assert.NoError(t, err) assert.NotEmpty(t, clientSecret) @@ -498,7 +498,7 @@ func TestGenerateClientSecret(t *testing.T) { service.config.Client.ClientSecretLength = originalLength }() - clientSecret, err := service.generateClientSecret() + clientSecret, err := service.GenerateClientSecret() assert.NoError(t, err) assert.NotEmpty(t, clientSecret) diff --git a/openapi/oauth/providers/client/default.go b/openapi/oauth/providers/client/default.go index f1a8faf5..51f13b68 100644 --- a/openapi/oauth/providers/client/default.go +++ b/openapi/oauth/providers/client/default.go @@ -322,9 +322,9 @@ func (c *DefaultClient) ValidateClient(ctx context.Context, clientInfo *types.Cl } // Validate redirect URIs - if len(clientInfo.RedirectURIs) == 0 { + if len(clientInfo.RedirectURIs) == 0 && (strings.Contains(clientInfo.Scope, "openid") || strings.Contains(clientInfo.Scope, "profile") || strings.Contains(clientInfo.Scope, "email")) { result.Valid = false - result.Errors = append(result.Errors, "At least one redirect URI is required") + result.Errors = append(result.Errors, "At least one redirect URI is required for openid, profile, or email scope") } // Validate grant types diff --git a/openapi/oauth/providers/user/default.go b/openapi/oauth/providers/user/default.go index b43b7c21..e76264ed 100644 --- a/openapi/oauth/providers/user/default.go +++ b/openapi/oauth/providers/user/default.go @@ -168,8 +168,9 @@ type IDStrategy string // Available ID generation strategies const ( - NanoIDStrategy IDStrategy = "nanoid" // Short, URL-safe, readable (e.g., "Kx9mP2aQ7nR3") - UUIDStrategy IDStrategy = "uuid" // Traditional UUID (for compatibility) + NanoIDStrategy IDStrategy = "nanoid" // Short, URL-safe, readable (e.g., "Kx9mP2aQ7nR3") + UUIDStrategy IDStrategy = "uuid" // Traditional UUID (for compatibility) + NumericStrategy IDStrategy = "numeric" // Numeric ID (for compatibility) ) // DefaultUserOptions provides options for the DefaultUser @@ -232,7 +233,7 @@ func NewDefaultUser(options *DefaultUserOptions) *DefaultUser { // Set ID generation strategy with defaults idStrategy := options.IDStrategy if idStrategy == "" { - idStrategy = NanoIDStrategy // Default to NanoID for better UX + idStrategy = NumericStrategy // Default to Numeric for better UX } // Set ID prefix (default is empty string) diff --git a/openapi/oauth/providers/user/utils.go b/openapi/oauth/providers/user/utils.go index bbf7caa5..d8a9f5e0 100644 --- a/openapi/oauth/providers/user/utils.go +++ b/openapi/oauth/providers/user/utils.go @@ -22,8 +22,8 @@ func (u *DefaultUser) GenerateUserID(ctx context.Context, safe ...bool) (string, if len(safe) > 0 { safeMode = safe[0] // Use provided value } else { - // Default: safe for NanoID, unsafe for UUID - safeMode = u.idStrategy == NanoIDStrategy + // Default: if idStrategy is Numeric or NanoID, use safe mode. + safeMode = (u.idStrategy == NumericStrategy) || (u.idStrategy == NanoIDStrategy) } if !safeMode { @@ -66,9 +66,11 @@ func (u *DefaultUser) generateUserID() (string, error) { case UUIDStrategy: id, err = generateUUID() case NanoIDStrategy: - fallthrough - default: id, err = generateNanoID(12) // 12 characters, URL-safe, readable + case NumericStrategy: + id, err = generateNumericID(12) // 12 characters, numeric, readable (default) + default: + id, err = generateNumericID(12) // 12 characters, URL-safe, readable } if err != nil { @@ -136,6 +138,14 @@ func generateNanoID(length int) (string, error) { return gonanoid.Generate(alphabet, length) } +// generateNumericID generates a numeric ID +func generateNumericID(length int) (string, error) { + if length <= 0 || length > 16 { + return "", fmt.Errorf("length must be between 1 and 16") + } + return gonanoid.Generate("0123456789", length) +} + // generateUUID generates a traditional UUID using Google's library func generateUUID() (string, error) { return uuid.NewString(), nil diff --git a/openapi/oauth/token.go b/openapi/oauth/token.go index 7b023f20..7c336494 100644 --- a/openapi/oauth/token.go +++ b/openapi/oauth/token.go @@ -269,8 +269,8 @@ func (s *Service) Subject(clientID, userID string) (string, error) { maxRetries := 5 for i := 0; i < maxRetries; i++ { - // Generate 12-character NanoID - nanoID, err := generateNanoID(12) + // Generate 16-character NanoID + nanoID, err := generateNumericID(16) if err != nil { return "", fmt.Errorf("failed to generate NanoID: %w", err) } @@ -644,11 +644,16 @@ func (s *Service) generateToken(tokenType string, clientID string) (string, erro // User Fingerprint Methods // ============================================================================ -// generateNanoID generates a Nano ID using the library -func generateNanoID(length int) (string, error) { - // URL-safe alphabet (no ambiguous characters like 0/O, 1/l/I) - const alphabet = "23456789ABCDEFGHJKMNPQRSTUVWXYZabcdefghijkmnpqrstuvwxyz" - return gonanoid.Generate(alphabet, length) +// generateNumericID generates a deterministic numeric ID using simple hash mapping +func generateNumericID(length int) (string, error) { + if length <= 0 || length > 16 { + return "", fmt.Errorf("length must be between 1 and 16") + } + // Use only digits 0-9 for numeric ID + // This provides 10^length possible combinations + // For 16 digits, that's 10^16 = 10,000,000,000,000,000 possibilities + const numericAlphabet = "0123456789" + return gonanoid.Generate(numericAlphabet, length) } // DeleteUserFingerprint removes a fingerprint mapping diff --git a/openapi/oauth/types/types.go b/openapi/oauth/types/types.go index 21ffefd7..539b005f 100644 --- a/openapi/oauth/types/types.go +++ b/openapi/oauth/types/types.go @@ -336,6 +336,7 @@ type TokenExchangeResponse struct { // DynamicClientRegistrationRequest represents dynamic client registration request type DynamicClientRegistrationRequest struct { + ClientID string `json:"client_id,omitempty"` // Optional: Client ID to use for registration, if not provided, a new client ID will be generated RedirectURIs []string `json:"redirect_uris"` ResponseTypes []string `json:"response_types,omitempty"` GrantTypes []string `json:"grant_types,omitempty"` diff --git a/openapi/signin/login.go b/openapi/signin/login.go index 216ab00d..a0ccc820 100644 --- a/openapi/signin/login.go +++ b/openapi/signin/login.go @@ -94,17 +94,12 @@ func LoginByUserID(userid string, ip string) (*LoginResponse, error) { log.Warn("Failed to update last login: %s", err.Error()) } - var scopes []string + var scopes []string = yaoClientConfig.Scopes if v, ok := user["scopes"].([]string); ok { scopes = v } - // Get Config form app.yao config () - clientID := "1234567890" - oidcExpiresIn := 3600 - accessTokenExpiresIn := 3600 - - subject, err := oauth.OAuth.Subject(clientID, userid) + subject, err := oauth.OAuth.Subject(yaoClientConfig.ClientID, userid) if err != nil { log.Warn("Failed to store user fingerprint: %s", err.Error()) } @@ -112,27 +107,28 @@ func LoginByUserID(userid string, ip string) (*LoginResponse, error) { oidcUserInfo.Sub = subject // OIDC Token - oidcToken, err := oauth.OAuth.SignIDToken(clientID, strings.Join(scopes, " "), oidcExpiresIn, oidcUserInfo) + oidcToken, err := oauth.OAuth.SignIDToken(yaoClientConfig.ClientID, strings.Join(scopes, " "), yaoClientConfig.ExpiresIn, oidcUserInfo) if err != nil { return nil, err } // Access Token - accessToken, err := oauth.OAuth.MakeAccessToken(clientID, strings.Join(scopes, " "), subject, accessTokenExpiresIn) + accessToken, err := oauth.OAuth.MakeAccessToken(yaoClientConfig.ClientID, strings.Join(scopes, " "), subject, yaoClientConfig.ExpiresIn) if err != nil { return nil, err } // Refresh Token - refreshToken, err := oauth.OAuth.MakeRefreshToken(clientID, strings.Join(scopes, " "), subject) + refreshToken, err := oauth.OAuth.MakeRefreshToken(yaoClientConfig.ClientID, strings.Join(scopes, " "), subject) if err != nil { return nil, err } + return &LoginResponse{ AccessToken: accessToken, IDToken: oidcToken, RefreshToken: refreshToken, - ExpiresIn: accessTokenExpiresIn, + ExpiresIn: yaoClientConfig.ExpiresIn, TokenType: "Bearer", Scope: strings.Join(scopes, " "), }, nil diff --git a/openapi/signin/signin.go b/openapi/signin/signin.go index 57b2af66..8e73df7f 100644 --- a/openapi/signin/signin.go +++ b/openapi/signin/signin.go @@ -1,8 +1,8 @@ package signin import ( + "context" "fmt" - "log" "os" "path/filepath" "regexp" @@ -12,11 +12,17 @@ import ( "time" "github.com/yaoapp/gou/application" + "github.com/yaoapp/kun/log" "github.com/yaoapp/yao/config" + "github.com/yaoapp/yao/openapi/oauth" + "github.com/yaoapp/yao/openapi/oauth/types" ) // Global variables to store loaded configurations var ( + // Client config + yaoClientConfig *YaoClientConfig + // Full configurations with sensitive data (for backend use) fullConfigs = make(map[string]*Config) // Public configurations without sensitive data (for frontend use) @@ -40,21 +46,121 @@ func Load(appConfig config.Config) error { providers = make(map[string]*Provider) defaultConfig = nil - // Load providers first - err := loadProviders(appConfig.Root) - if err != nil { - return fmt.Errorf("failed to load providers: %v", err) - } - // Load signin configurations - err = loadSigninConfigs(appConfig.Root) + err := loadSigninConfigs(appConfig.Root) if err != nil { return fmt.Errorf("failed to load signin configs: %v", err) } + // Load providers first + err = loadProviders(appConfig.Root) + if err != nil { + return fmt.Errorf("failed to load providers: %v", err) + } + + // Load client config + err = loadClientConfig() + if err != nil { + return fmt.Errorf("failed to load client config: %v", err) + } + return nil } +// loadClientConfig loads the client config from the openapi/signin/client.yao file +func loadClientConfig() error { + // Check if client config exists + exists, err := application.App.Exists("openapi/signin/client.yao") + if err != nil { + return fmt.Errorf("failed to check if client config exists: %v", err) + } + + if !exists { + return fmt.Errorf("client config not found") + } + + // Read client config + clientConfigRaw, err := application.App.Read("openapi/signin/client.yao") + if err != nil { + return fmt.Errorf("failed to read client config: %v", err) + } + + var clientConfig YaoClientConfig + err = application.Parse("openapi/signin/client.yao", clientConfigRaw, &clientConfig) + if err != nil { + return fmt.Errorf("failed to parse client config: %v", err) + } + + // Process ENV variables in client config + clientConfig.ClientID = replaceENVVar(clientConfig.ClientID) + clientConfig.ClientSecret = replaceENVVar(clientConfig.ClientSecret) + + // Validate client config + err = validateClientConfig(&clientConfig) + if err != nil { + return fmt.Errorf("failed to validate client config: %v", err) + } + + yaoClientConfig = &clientConfig + return nil +} + +// validateClientConfig validates the client config +func validateClientConfig(clientConfig *YaoClientConfig) error { + + // Validate client ID + err := oauth.OAuth.ValidateClientID(clientConfig.ClientID) + if err != nil { + return err + } + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + // Validate client is registered + c := oauth.OAuth.GetClientProvider() + _, err = c.GetClientByID(ctx, clientConfig.ClientID) + if err != nil { + // If client is not registered, register it + if strings.Contains(err.Error(), "Client not found") { + yaoClientConfig, err = registerClient(clientConfig.ClientID) + if err != nil { + return fmt.Errorf("failed to register client: %v", err) + } + return nil + } + return fmt.Errorf("failed to get client: %v", err) + } + + return nil +} + +// registerClient registers the client config with the OAuth server +func registerClient(clientID string) (*YaoClientConfig, error) { + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + // Register client + response, err := oauth.OAuth.DynamicClientRegistration(ctx, &types.DynamicClientRegistrationRequest{ + ClientID: clientID, + ClientName: "Yao OpenAPI Client", + ResponseTypes: []string{"code"}, + GrantTypes: []string{"client_credentials"}, + ApplicationType: types.ApplicationTypeWeb, + }) + if err != nil { + return nil, fmt.Errorf("failed to create client: %v", err) + } + + var clientConfig *YaoClientConfig = &YaoClientConfig{} + clientConfig.ClientID = response.ClientID + clientConfig.ClientSecret = response.ClientSecret + clientConfig.ExpiresIn = 3600 * 24 // 24 hours + clientConfig.Scopes = []string{"openid", "profile", "email"} + return clientConfig, nil +} + // loadProviders loads all provider configurations from the openapi/signin/providers directory func loadProviders(rootPath string) error { // Use Walk to find all provider files in the signin/providers directory @@ -68,6 +174,11 @@ func loadProviders(rootPath string) error { return nil } + // Skip client.yao file + if filename == "client.yao" { + return nil + } + // Extract provider ID from filename (basename without extension) baseName := filepath.Base(filename) providerID := strings.TrimSuffix(baseName, ".yao") @@ -210,7 +321,7 @@ func processProviderENVVariables(provider *Provider, rootPath string) { if provider.ClientSecretGenerator.ExpiresIn != "" { normalizedDuration, err := normalizeExpiresIn(provider.ClientSecretGenerator.ExpiresIn) if err != nil { - log.Printf("Warning: Invalid expires_in format '%s' for provider '%s': %v", + log.Warn("Invalid expires_in format '%s' for provider '%s': %v", provider.ClientSecretGenerator.ExpiresIn, provider.ID, err) // Set default to 90 days provider.ClientSecretGenerator.ExpiresIn = "2160h" // 90 * 24 hours @@ -252,7 +363,7 @@ func processProviderENVVariables(provider *Provider, rootPath string) { // Log warning for missing environment variables if len(missingEnvVars) > 0 { - log.Printf("Warning: The following environment variables are not set for provider '%s': %v", provider.ID, missingEnvVars) + log.Warn("The following environment variables are not set for provider '%s': %v", provider.ID, missingEnvVars) } } @@ -298,8 +409,8 @@ func processConfigENVVariables(config *Config, rootPath string) { // Log warning for missing environment variables if len(missingEnvVars) > 0 { - log.Printf("Warning: The following environment variables are not set in signin configuration: %v", missingEnvVars) - log.Printf("Please set these environment variables to avoid exposing placeholder values in configuration") + log.Warn("The following environment variables are not set in signin configuration: %v", missingEnvVars) + log.Warn("Please set these environment variables to avoid exposing placeholder values in configuration") } } diff --git a/openapi/signin/types.go b/openapi/signin/types.go index da1404d7..0295b003 100644 --- a/openapi/signin/types.go +++ b/openapi/signin/types.go @@ -64,6 +64,14 @@ type RegisterConfig struct { Role string `json:"role,omitempty"` } +// YaoClientConfig represents the Yao OpenAPI Client config +type YaoClientConfig struct { + ClientID string `json:"client_id,omitempty"` + ClientSecret string `json:"client_secret,omitempty"` + Scopes []string `json:"scopes,omitempty"` // Default scopes if not set in the provider config + ExpiresIn int `json:"expires_in,omitempty"` // Default expires in for the access token (optional) in seconds +} + // Provider represents a third party login provider type Provider struct { ID string `json:"id,omitempty"`