Update OAuth client configuration and ID generation methods

- Added support for numeric ID generation in the OAuth service, replacing the previous NanoID approach for better compatibility.
- Refactored client ID and secret generation methods to be public and renamed them for consistency.
- Enhanced dynamic client registration to allow optional client ID usage.
- Updated client configuration loading to include validation and registration of clients if not found.
- Improved error handling and logging for client configuration processes.
- Adjusted tests to reflect changes in ID generation and client configuration handling.
This commit is contained in:
Max 2025-08-04 11:04:56 +08:00
parent 2b7040c834
commit 4682c20903
12 changed files with 224 additions and 59 deletions

2
go.mod
View file

@ -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

View file

@ -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,

View file

@ -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",

View file

@ -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)

View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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"`

View file

@ -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

View file

@ -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")
}
}

View file

@ -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"`