Merge pull request #1081 from trheyi/main
Update OAuth client configuration and ID generation methods
This commit is contained in:
commit
4987a36f26
12 changed files with 224 additions and 59 deletions
2
go.mod
2
go.mod
|
|
@ -27,6 +27,7 @@ require (
|
||||||
github.com/matoous/go-nanoid/v2 v2.1.0
|
github.com/matoous/go-nanoid/v2 v2.1.0
|
||||||
github.com/mozillazg/go-pinyin v0.20.0
|
github.com/mozillazg/go-pinyin v0.20.0
|
||||||
github.com/pkoukk/tiktoken-go v0.1.7
|
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/rhysd/go-github-selfupdate v1.2.3
|
||||||
github.com/spf13/cast v1.9.2
|
github.com/spf13/cast v1.9.2
|
||||||
github.com/spf13/cobra v1.9.1
|
github.com/spf13/cobra v1.9.1
|
||||||
|
|
@ -117,7 +118,6 @@ require (
|
||||||
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
|
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
|
||||||
github.com/pkg/errors v0.9.1 // indirect
|
github.com/pkg/errors v0.9.1 // indirect
|
||||||
github.com/pmezard/go-difflib v1.0.0 // 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/qdrant/go-client v1.14.0 // indirect
|
||||||
github.com/richardlehane/mscfb v1.0.4 // indirect
|
github.com/richardlehane/mscfb v1.0.4 // indirect
|
||||||
github.com/richardlehane/msoleps v1.0.4 // indirect
|
github.com/richardlehane/msoleps v1.0.4 // indirect
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ package openapi
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
|
"fmt"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
@ -134,6 +135,11 @@ func (config *Config) UnmarshalJSON(data []byte) error {
|
||||||
Features: tempConfig.OAuth.Features,
|
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
|
// Convert signing config with duration parsing
|
||||||
config.OAuth.Signing = types.SigningConfig{
|
config.OAuth.Signing = types.SigningConfig{
|
||||||
SigningCertPath: tempConfig.OAuth.Signing.SigningCertPath,
|
SigningCertPath: tempConfig.OAuth.Signing.SigningCertPath,
|
||||||
|
|
@ -379,7 +385,7 @@ func (config *Config) OAuthConfig(appConfig config.Config) (*oauth.Config, error
|
||||||
ClientProvider: clientProvider,
|
ClientProvider: clientProvider,
|
||||||
Cache: cacheStore,
|
Cache: cacheStore,
|
||||||
Store: dataStore,
|
Store: dataStore,
|
||||||
IssuerURL: config.BaseURL,
|
IssuerURL: config.OAuth.IssuerURL,
|
||||||
Signing: signingConfig, // Use the converted signing config
|
Signing: signingConfig, // Use the converted signing config
|
||||||
Token: config.OAuth.Token,
|
Token: config.OAuth.Token,
|
||||||
Security: config.OAuth.Security,
|
Security: config.OAuth.Security,
|
||||||
|
|
|
||||||
|
|
@ -47,12 +47,17 @@ func (s *Service) DynamicClientRegistration(ctx context.Context, request *types.
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Generate client ID and secret
|
// Generate client ID and secret (use the client ID from the request if provided or generate a new one)
|
||||||
clientID, err := s.generateClientID()
|
clientID := request.ClientID
|
||||||
if err != nil {
|
var err error
|
||||||
return nil, &types.ErrorResponse{
|
if clientID == "" {
|
||||||
Code: types.ErrorServerError,
|
var err error
|
||||||
ErrorDescription: "Failed to generate client ID",
|
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.TokenEndpointAuthPost ||
|
||||||
request.TokenEndpointAuthMethod == types.TokenEndpointAuthJWT {
|
request.TokenEndpointAuthMethod == types.TokenEndpointAuthJWT {
|
||||||
clientType = types.ClientTypeConfidential
|
clientType = types.ClientTypeConfidential
|
||||||
clientSecret, err = s.generateClientSecret()
|
clientSecret, err = s.GenerateClientSecret()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, &types.ErrorResponse{
|
return nil, &types.ErrorResponse{
|
||||||
Code: types.ErrorServerError,
|
Code: types.ErrorServerError,
|
||||||
|
|
@ -144,8 +149,8 @@ func (s *Service) DynamicClientRegistration(ctx context.Context, request *types.
|
||||||
return response, nil
|
return response, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// generateClientID generates a random client ID
|
// GenerateClientID generates a random client ID
|
||||||
func (s *Service) generateClientID() (string, error) {
|
func (s *Service) GenerateClientID() (string, error) {
|
||||||
length := s.config.Client.ClientIDLength
|
length := s.config.Client.ClientIDLength
|
||||||
if length == 0 {
|
if length == 0 {
|
||||||
length = 32
|
length = 32
|
||||||
|
|
@ -160,8 +165,30 @@ func (s *Service) generateClientID() (string, error) {
|
||||||
return strings.TrimRight(base64.URLEncoding.EncodeToString(bytes), "="), nil
|
return strings.TrimRight(base64.URLEncoding.EncodeToString(bytes), "="), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// generateClientSecret generates a random client secret
|
// ValidateClientID validates the client ID
|
||||||
func (s *Service) generateClientSecret() (string, error) {
|
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
|
length := s.config.Client.ClientSecretLength
|
||||||
if length == 0 {
|
if length == 0 {
|
||||||
length = 64
|
length = 64
|
||||||
|
|
@ -179,7 +206,7 @@ func (s *Service) generateClientSecret() (string, error) {
|
||||||
// validateDynamicClientRegistrationRequest validates the dynamic client registration request
|
// validateDynamicClientRegistrationRequest validates the dynamic client registration request
|
||||||
func (s *Service) validateDynamicClientRegistrationRequest(request *types.DynamicClientRegistrationRequest) error {
|
func (s *Service) validateDynamicClientRegistrationRequest(request *types.DynamicClientRegistrationRequest) error {
|
||||||
// Validate redirect URIs
|
// 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{
|
return &types.ErrorResponse{
|
||||||
Code: types.ErrorInvalidRequest,
|
Code: types.ErrorInvalidRequest,
|
||||||
ErrorDescription: "At least one redirect URI is required",
|
ErrorDescription: "At least one redirect URI is required",
|
||||||
|
|
|
||||||
|
|
@ -417,7 +417,7 @@ func TestGenerateClientID(t *testing.T) {
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
|
|
||||||
t.Run("generate client ID with default length", func(t *testing.T) {
|
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.NoError(t, err)
|
||||||
assert.NotEmpty(t, clientID)
|
assert.NotEmpty(t, clientID)
|
||||||
assert.Greater(t, len(clientID), 0)
|
assert.Greater(t, len(clientID), 0)
|
||||||
|
|
@ -432,7 +432,7 @@ func TestGenerateClientID(t *testing.T) {
|
||||||
clientIDs := make(map[string]bool)
|
clientIDs := make(map[string]bool)
|
||||||
|
|
||||||
for i := 0; i < 100; i++ {
|
for i := 0; i < 100; i++ {
|
||||||
clientID, err := service.generateClientID()
|
clientID, err := service.GenerateClientID()
|
||||||
assert.NoError(t, err)
|
assert.NoError(t, err)
|
||||||
assert.NotEmpty(t, clientID)
|
assert.NotEmpty(t, clientID)
|
||||||
|
|
||||||
|
|
@ -450,7 +450,7 @@ func TestGenerateClientID(t *testing.T) {
|
||||||
service.config.Client.ClientIDLength = originalLength
|
service.config.Client.ClientIDLength = originalLength
|
||||||
}()
|
}()
|
||||||
|
|
||||||
clientID, err := service.generateClientID()
|
clientID, err := service.GenerateClientID()
|
||||||
assert.NoError(t, err)
|
assert.NoError(t, err)
|
||||||
assert.NotEmpty(t, clientID)
|
assert.NotEmpty(t, clientID)
|
||||||
|
|
||||||
|
|
@ -465,7 +465,7 @@ func TestGenerateClientSecret(t *testing.T) {
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
|
|
||||||
t.Run("generate client secret with default length", func(t *testing.T) {
|
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.NoError(t, err)
|
||||||
assert.NotEmpty(t, clientSecret)
|
assert.NotEmpty(t, clientSecret)
|
||||||
assert.Greater(t, len(clientSecret), 0)
|
assert.Greater(t, len(clientSecret), 0)
|
||||||
|
|
@ -480,7 +480,7 @@ func TestGenerateClientSecret(t *testing.T) {
|
||||||
clientSecrets := make(map[string]bool)
|
clientSecrets := make(map[string]bool)
|
||||||
|
|
||||||
for i := 0; i < 100; i++ {
|
for i := 0; i < 100; i++ {
|
||||||
clientSecret, err := service.generateClientSecret()
|
clientSecret, err := service.GenerateClientSecret()
|
||||||
assert.NoError(t, err)
|
assert.NoError(t, err)
|
||||||
assert.NotEmpty(t, clientSecret)
|
assert.NotEmpty(t, clientSecret)
|
||||||
|
|
||||||
|
|
@ -498,7 +498,7 @@ func TestGenerateClientSecret(t *testing.T) {
|
||||||
service.config.Client.ClientSecretLength = originalLength
|
service.config.Client.ClientSecretLength = originalLength
|
||||||
}()
|
}()
|
||||||
|
|
||||||
clientSecret, err := service.generateClientSecret()
|
clientSecret, err := service.GenerateClientSecret()
|
||||||
assert.NoError(t, err)
|
assert.NoError(t, err)
|
||||||
assert.NotEmpty(t, clientSecret)
|
assert.NotEmpty(t, clientSecret)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -322,9 +322,9 @@ func (c *DefaultClient) ValidateClient(ctx context.Context, clientInfo *types.Cl
|
||||||
}
|
}
|
||||||
|
|
||||||
// Validate redirect URIs
|
// 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.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
|
// Validate grant types
|
||||||
|
|
|
||||||
|
|
@ -168,8 +168,9 @@ type IDStrategy string
|
||||||
|
|
||||||
// Available ID generation strategies
|
// Available ID generation strategies
|
||||||
const (
|
const (
|
||||||
NanoIDStrategy IDStrategy = "nanoid" // Short, URL-safe, readable (e.g., "Kx9mP2aQ7nR3")
|
NanoIDStrategy IDStrategy = "nanoid" // Short, URL-safe, readable (e.g., "Kx9mP2aQ7nR3")
|
||||||
UUIDStrategy IDStrategy = "uuid" // Traditional UUID (for compatibility)
|
UUIDStrategy IDStrategy = "uuid" // Traditional UUID (for compatibility)
|
||||||
|
NumericStrategy IDStrategy = "numeric" // Numeric ID (for compatibility)
|
||||||
)
|
)
|
||||||
|
|
||||||
// DefaultUserOptions provides options for the DefaultUser
|
// DefaultUserOptions provides options for the DefaultUser
|
||||||
|
|
@ -232,7 +233,7 @@ func NewDefaultUser(options *DefaultUserOptions) *DefaultUser {
|
||||||
// Set ID generation strategy with defaults
|
// Set ID generation strategy with defaults
|
||||||
idStrategy := options.IDStrategy
|
idStrategy := options.IDStrategy
|
||||||
if 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)
|
// Set ID prefix (default is empty string)
|
||||||
|
|
|
||||||
|
|
@ -22,8 +22,8 @@ func (u *DefaultUser) GenerateUserID(ctx context.Context, safe ...bool) (string,
|
||||||
if len(safe) > 0 {
|
if len(safe) > 0 {
|
||||||
safeMode = safe[0] // Use provided value
|
safeMode = safe[0] // Use provided value
|
||||||
} else {
|
} else {
|
||||||
// Default: safe for NanoID, unsafe for UUID
|
// Default: if idStrategy is Numeric or NanoID, use safe mode.
|
||||||
safeMode = u.idStrategy == NanoIDStrategy
|
safeMode = (u.idStrategy == NumericStrategy) || (u.idStrategy == NanoIDStrategy)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !safeMode {
|
if !safeMode {
|
||||||
|
|
@ -66,9 +66,11 @@ func (u *DefaultUser) generateUserID() (string, error) {
|
||||||
case UUIDStrategy:
|
case UUIDStrategy:
|
||||||
id, err = generateUUID()
|
id, err = generateUUID()
|
||||||
case NanoIDStrategy:
|
case NanoIDStrategy:
|
||||||
fallthrough
|
|
||||||
default:
|
|
||||||
id, err = generateNanoID(12) // 12 characters, URL-safe, readable
|
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 {
|
if err != nil {
|
||||||
|
|
@ -136,6 +138,14 @@ func generateNanoID(length int) (string, error) {
|
||||||
return gonanoid.Generate(alphabet, length)
|
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
|
// generateUUID generates a traditional UUID using Google's library
|
||||||
func generateUUID() (string, error) {
|
func generateUUID() (string, error) {
|
||||||
return uuid.NewString(), nil
|
return uuid.NewString(), nil
|
||||||
|
|
|
||||||
|
|
@ -269,8 +269,8 @@ func (s *Service) Subject(clientID, userID string) (string, error) {
|
||||||
|
|
||||||
maxRetries := 5
|
maxRetries := 5
|
||||||
for i := 0; i < maxRetries; i++ {
|
for i := 0; i < maxRetries; i++ {
|
||||||
// Generate 12-character NanoID
|
// Generate 16-character NanoID
|
||||||
nanoID, err := generateNanoID(12)
|
nanoID, err := generateNumericID(16)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("failed to generate NanoID: %w", err)
|
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
|
// User Fingerprint Methods
|
||||||
// ============================================================================
|
// ============================================================================
|
||||||
|
|
||||||
// generateNanoID generates a Nano ID using the library
|
// generateNumericID generates a deterministic numeric ID using simple hash mapping
|
||||||
func generateNanoID(length int) (string, error) {
|
func generateNumericID(length int) (string, error) {
|
||||||
// URL-safe alphabet (no ambiguous characters like 0/O, 1/l/I)
|
if length <= 0 || length > 16 {
|
||||||
const alphabet = "23456789ABCDEFGHJKMNPQRSTUVWXYZabcdefghijkmnpqrstuvwxyz"
|
return "", fmt.Errorf("length must be between 1 and 16")
|
||||||
return gonanoid.Generate(alphabet, length)
|
}
|
||||||
|
// 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
|
// DeleteUserFingerprint removes a fingerprint mapping
|
||||||
|
|
|
||||||
|
|
@ -336,6 +336,7 @@ type TokenExchangeResponse struct {
|
||||||
|
|
||||||
// DynamicClientRegistrationRequest represents dynamic client registration request
|
// DynamicClientRegistrationRequest represents dynamic client registration request
|
||||||
type DynamicClientRegistrationRequest struct {
|
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"`
|
RedirectURIs []string `json:"redirect_uris"`
|
||||||
ResponseTypes []string `json:"response_types,omitempty"`
|
ResponseTypes []string `json:"response_types,omitempty"`
|
||||||
GrantTypes []string `json:"grant_types,omitempty"`
|
GrantTypes []string `json:"grant_types,omitempty"`
|
||||||
|
|
|
||||||
|
|
@ -94,17 +94,12 @@ func LoginByUserID(userid string, ip string) (*LoginResponse, error) {
|
||||||
log.Warn("Failed to update last login: %s", err.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 {
|
if v, ok := user["scopes"].([]string); ok {
|
||||||
scopes = v
|
scopes = v
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get Config form app.yao config ()
|
subject, err := oauth.OAuth.Subject(yaoClientConfig.ClientID, userid)
|
||||||
clientID := "1234567890"
|
|
||||||
oidcExpiresIn := 3600
|
|
||||||
accessTokenExpiresIn := 3600
|
|
||||||
|
|
||||||
subject, err := oauth.OAuth.Subject(clientID, userid)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Warn("Failed to store user fingerprint: %s", err.Error())
|
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
|
oidcUserInfo.Sub = subject
|
||||||
|
|
||||||
// OIDC Token
|
// 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 {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Access Token
|
// 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 {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Refresh Token
|
// 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 {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
return &LoginResponse{
|
return &LoginResponse{
|
||||||
AccessToken: accessToken,
|
AccessToken: accessToken,
|
||||||
IDToken: oidcToken,
|
IDToken: oidcToken,
|
||||||
RefreshToken: refreshToken,
|
RefreshToken: refreshToken,
|
||||||
ExpiresIn: accessTokenExpiresIn,
|
ExpiresIn: yaoClientConfig.ExpiresIn,
|
||||||
TokenType: "Bearer",
|
TokenType: "Bearer",
|
||||||
Scope: strings.Join(scopes, " "),
|
Scope: strings.Join(scopes, " "),
|
||||||
}, nil
|
}, nil
|
||||||
|
|
|
||||||
|
|
@ -1,8 +1,8 @@
|
||||||
package signin
|
package signin
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"regexp"
|
"regexp"
|
||||||
|
|
@ -12,11 +12,17 @@ import (
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/yaoapp/gou/application"
|
"github.com/yaoapp/gou/application"
|
||||||
|
"github.com/yaoapp/kun/log"
|
||||||
"github.com/yaoapp/yao/config"
|
"github.com/yaoapp/yao/config"
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth"
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Global variables to store loaded configurations
|
// Global variables to store loaded configurations
|
||||||
var (
|
var (
|
||||||
|
// Client config
|
||||||
|
yaoClientConfig *YaoClientConfig
|
||||||
|
|
||||||
// Full configurations with sensitive data (for backend use)
|
// Full configurations with sensitive data (for backend use)
|
||||||
fullConfigs = make(map[string]*Config)
|
fullConfigs = make(map[string]*Config)
|
||||||
// Public configurations without sensitive data (for frontend use)
|
// Public configurations without sensitive data (for frontend use)
|
||||||
|
|
@ -40,21 +46,121 @@ func Load(appConfig config.Config) error {
|
||||||
providers = make(map[string]*Provider)
|
providers = make(map[string]*Provider)
|
||||||
defaultConfig = nil
|
defaultConfig = nil
|
||||||
|
|
||||||
// Load providers first
|
|
||||||
err := loadProviders(appConfig.Root)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to load providers: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Load signin configurations
|
// Load signin configurations
|
||||||
err = loadSigninConfigs(appConfig.Root)
|
err := loadSigninConfigs(appConfig.Root)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to load signin configs: %v", err)
|
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
|
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
|
// loadProviders loads all provider configurations from the openapi/signin/providers directory
|
||||||
func loadProviders(rootPath string) error {
|
func loadProviders(rootPath string) error {
|
||||||
// Use Walk to find all provider files in the signin/providers directory
|
// Use Walk to find all provider files in the signin/providers directory
|
||||||
|
|
@ -68,6 +174,11 @@ func loadProviders(rootPath string) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Skip client.yao file
|
||||||
|
if filename == "client.yao" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// Extract provider ID from filename (basename without extension)
|
// Extract provider ID from filename (basename without extension)
|
||||||
baseName := filepath.Base(filename)
|
baseName := filepath.Base(filename)
|
||||||
providerID := strings.TrimSuffix(baseName, ".yao")
|
providerID := strings.TrimSuffix(baseName, ".yao")
|
||||||
|
|
@ -210,7 +321,7 @@ func processProviderENVVariables(provider *Provider, rootPath string) {
|
||||||
if provider.ClientSecretGenerator.ExpiresIn != "" {
|
if provider.ClientSecretGenerator.ExpiresIn != "" {
|
||||||
normalizedDuration, err := normalizeExpiresIn(provider.ClientSecretGenerator.ExpiresIn)
|
normalizedDuration, err := normalizeExpiresIn(provider.ClientSecretGenerator.ExpiresIn)
|
||||||
if err != nil {
|
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)
|
provider.ClientSecretGenerator.ExpiresIn, provider.ID, err)
|
||||||
// Set default to 90 days
|
// Set default to 90 days
|
||||||
provider.ClientSecretGenerator.ExpiresIn = "2160h" // 90 * 24 hours
|
provider.ClientSecretGenerator.ExpiresIn = "2160h" // 90 * 24 hours
|
||||||
|
|
@ -252,7 +363,7 @@ func processProviderENVVariables(provider *Provider, rootPath string) {
|
||||||
|
|
||||||
// Log warning for missing environment variables
|
// Log warning for missing environment variables
|
||||||
if len(missingEnvVars) > 0 {
|
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
|
// Log warning for missing environment variables
|
||||||
if len(missingEnvVars) > 0 {
|
if len(missingEnvVars) > 0 {
|
||||||
log.Printf("Warning: The following environment variables are not set in signin configuration: %v", missingEnvVars)
|
log.Warn("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("Please set these environment variables to avoid exposing placeholder values in configuration")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -64,6 +64,14 @@ type RegisterConfig struct {
|
||||||
Role string `json:"role,omitempty"`
|
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
|
// Provider represents a third party login provider
|
||||||
type Provider struct {
|
type Provider struct {
|
||||||
ID string `json:"id,omitempty"`
|
ID string `json:"id,omitempty"`
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue