From ffe5bc114874f01904397e001105c876359c1d8e Mon Sep 17 00:00:00 2001 From: Ricky Date: Tue, 11 Aug 2026 15:58:12 +0800 Subject: [PATCH] fix(oidc): refine provider integration --- inventory/migration.go | 17 ++++++++++--- inventory/setting.go | 19 ++++++++++++++- pkg/auth/oidc.go | 46 +++++------------------------------- pkg/cluster/routes/routes.go | 5 ++++ pkg/setting/provider.go | 6 +++++ service/oauth/oauth.go | 8 ++++--- service/oauth/oidc.go | 19 ++++++--------- 7 files changed, 61 insertions(+), 59 deletions(-) diff --git a/inventory/migration.go b/inventory/migration.go index f581dbc3..993e71ef 100644 --- a/inventory/migration.go +++ b/inventory/migration.go @@ -27,6 +27,14 @@ import ( // needMigration exams if required schema version is satisfied. func needMigration(client *ent.Client, ctx context.Context, requiredDbVersion string) bool { c, _ := client.Setting.Query().Where(setting.NameEQ(DBVersionPrefix + requiredDbVersion)).Count(ctx) + if c == 0 { + return true + } + + c, _ = client.Setting.Query().Where( + setting.NameEQ(OIDCSigningPrivateKeySetting), + setting.ValueNEQ(""), + ).Count(ctx) return c == 0 } @@ -66,19 +74,22 @@ func migrateDefaultSettings(l logging.Logger, client *ent.Client, ctx context.Co } // List existing settings into a map - existingSettings := make(map[string]struct{}) + existingSettings := make(map[string]*ent.Setting) settings, err := client.Setting.Query().All(ctx) if err != nil { l.Warning("Failed to query existing settings: %s", err) } for _, s := range settings { - existingSettings[s.Name] = struct{}{} + existingSettings[s.Name] = s } l.Info("Insert default settings...") for k, v := range DefaultSettings { - if _, ok := existingSettings[k]; ok { + if existing, ok := existingSettings[k]; ok { + if k == OIDCSigningPrivateKeySetting && existing.Value == "" { + client.Setting.UpdateOne(existing).SetValue(v).SaveX(ctx) + } l.Debug("Skip inserting setting %s, already exists.", k) continue } diff --git a/inventory/setting.go b/inventory/setting.go index 34ebe6f2..1b1494fa 100644 --- a/inventory/setting.go +++ b/inventory/setting.go @@ -3,8 +3,11 @@ package inventory import ( "context" "crypto/rand" + "crypto/rsa" + "crypto/x509" "encoding/base64" "encoding/json" + "encoding/pem" "fmt" "io" @@ -483,6 +486,20 @@ var mailTemplateContents = []MailTemplateContent{ }, } +const OIDCSigningPrivateKeySetting = "oidc_signing_private_key" + +func mustGenerateOIDCSigningPrivateKey() string { + key, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + panic(fmt.Errorf("failed to generate OIDC signing key: %w", err)) + } + + return string(pem.EncodeToMemory(&pem.Block{ + Type: "RSA PRIVATE KEY", + Bytes: x509.MarshalPKCS1PrivateKey(key), + })) +} + var DefaultSettings = map[string]string{ "siteURL": `http://localhost:5212`, "siteName": `Cloudreve`, @@ -524,7 +541,7 @@ var DefaultSettings = map[string]string{ "theme_options": `{"#1976d2":{"light":{"palette":{"primary":{"main":"#1976d2","light":"#42a5f5","dark":"#1565c0"},"secondary":{"main":"#9c27b0","light":"#ba68c8","dark":"#7b1fa2"}}},"dark":{"palette":{"primary":{"main":"#90caf9","light":"#e3f2fd","dark":"#42a5f5"},"secondary":{"main":"#ce93d8","light":"#f3e5f5","dark":"#ab47bc"}}}},"#3f51b5":{"light":{"palette":{"primary":{"main":"#3f51b5"},"secondary":{"main":"#f50057"}}},"dark":{"palette":{"primary":{"main":"#9fa8da"},"secondary":{"main":"#ff4081"}}}}}`, "max_parallel_transfer": `4`, "secret_key": util.RandStringRunesCrypto(256), - "oidc_signing_private_key": "", + OIDCSigningPrivateKeySetting: mustGenerateOIDCSigningPrivateKey(), "temp_path": "temp", "avatar_path": "avatar", "avatar_size": "4194304", diff --git a/pkg/auth/oidc.go b/pkg/auth/oidc.go index 987ac250..8b935253 100644 --- a/pkg/auth/oidc.go +++ b/pkg/auth/oidc.go @@ -1,8 +1,6 @@ package auth import ( - "context" - "crypto/rand" "crypto/rsa" "crypto/sha256" "crypto/x509" @@ -11,13 +9,9 @@ import ( "fmt" "math/big" - "github.com/cloudreve/Cloudreve/v4/ent" - "github.com/cloudreve/Cloudreve/v4/inventory" "github.com/golang-jwt/jwt/v5" ) -const OIDCSigningPrivateKeySetting = "oidc_signing_private_key" - type OIDCIDTokenClaims struct { jwt.RegisteredClaims Nonce string `json:"nonce,omitempty"` @@ -42,55 +36,27 @@ type JWK struct { E string `json:"e"` } -func SignOIDCIDToken(ctx context.Context, settingClient inventory.SettingClient, claims *OIDCIDTokenClaims) (string, error) { - key, kid, err := loadOrCreateOIDCSigningKey(ctx, settingClient) +func SignOIDCIDToken(privateKeyRaw string, claims *OIDCIDTokenClaims) (string, error) { + key, err := parseRSAPrivateKey(privateKeyRaw) if err != nil { return "", err } token := jwt.NewWithClaims(jwt.SigningMethodRS256, claims) - token.Header["kid"] = kid + token.Header["kid"] = oidcSigningKeyID(&key.PublicKey) return token.SignedString(key) } -func OIDCJWKSet(ctx context.Context, settingClient inventory.SettingClient) (*JWKSet, error) { - key, kid, err := loadOrCreateOIDCSigningKey(ctx, settingClient) +func OIDCJWKSet(privateKeyRaw string) (*JWKSet, error) { + key, err := parseRSAPrivateKey(privateKeyRaw) if err != nil { return nil, err } + kid := oidcSigningKeyID(&key.PublicKey) return &JWKSet{Keys: []JWK{buildJWK(&key.PublicKey, kid)}}, nil } -func loadOrCreateOIDCSigningKey(ctx context.Context, settingClient inventory.SettingClient) (*rsa.PrivateKey, string, error) { - privateKeyRaw, err := settingClient.Get(ctx, OIDCSigningPrivateKeySetting) - if err != nil && !ent.IsNotFound(err) { - return nil, "", fmt.Errorf("failed to load OIDC signing key: %w", err) - } - if privateKeyRaw != "" { - key, err := parseRSAPrivateKey(privateKeyRaw) - if err != nil { - return nil, "", err - } - return key, oidcSigningKeyID(&key.PublicKey), nil - } - - key, err := rsa.GenerateKey(rand.Reader, 2048) - if err != nil { - return nil, "", fmt.Errorf("failed to generate OIDC signing key: %w", err) - } - - privateKeyRaw = string(pem.EncodeToMemory(&pem.Block{ - Type: "RSA PRIVATE KEY", - Bytes: x509.MarshalPKCS1PrivateKey(key), - })) - if err := settingClient.Set(ctx, map[string]string{OIDCSigningPrivateKeySetting: privateKeyRaw}); err != nil { - return nil, "", fmt.Errorf("failed to persist OIDC signing key: %w", err) - } - - return key, oidcSigningKeyID(&key.PublicKey), nil -} - func parseRSAPrivateKey(privateKeyRaw string) (*rsa.PrivateKey, error) { block, _ := pem.Decode([]byte(privateKeyRaw)) if block == nil { diff --git a/pkg/cluster/routes/routes.go b/pkg/cluster/routes/routes.go index bd51d215..79b18f1b 100644 --- a/pkg/cluster/routes/routes.go +++ b/pkg/cluster/routes/routes.go @@ -43,6 +43,11 @@ func MasterPingUrl(base *url.URL) *url.URL { return base.ResolveReference(masterPing) } +func MasterOIDCEndpointUrl(base *url.URL, endpoint string) string { + route, _ := url.Parse(endpoint) + return base.ResolveReference(route).String() +} + func MasterSlaveCallbackUrl(base *url.URL, driver, id, secret string) *url.URL { apiBaseURI, _ := url.Parse(path.Join(constants.APIPrefix+"/callback", driver, id, secret)) return base.ResolveReference(apiBaseURI) diff --git a/pkg/setting/provider.go b/pkg/setting/provider.go index a95eaa24..5d427373 100644 --- a/pkg/setting/provider.go +++ b/pkg/setting/provider.go @@ -52,6 +52,8 @@ type ( SiteURL(ctx context.Context) *url.URL // SecretKey returns the secret key for general signature. SecretKey(ctx context.Context) string + // OIDCSigningPrivateKey returns the private key used to sign OIDC ID tokens. + OIDCSigningPrivateKey(ctx context.Context) string // ActivationEmailTemplate returns the email template for activation. ActivationEmailTemplate(ctx context.Context) []EmailTemplate // ResetEmailTemplate returns the email template for reset password. @@ -737,6 +739,10 @@ func (s *settingProvider) SecretKey(ctx context.Context) string { return s.getString(ctx, "secret_key", "") } +func (s *settingProvider) OIDCSigningPrivateKey(ctx context.Context) string { + return s.getString(ctx, "oidc_signing_private_key", "") +} + func (s *settingProvider) AllSiteURLs(ctx context.Context) []*url.URL { rawUrls := s.getStringList(ctx, "siteURL", []string{"http://localhost"}) if len(rawUrls) == 0 { diff --git a/service/oauth/oauth.go b/service/oauth/oauth.go index bdf71f9a..064799ad 100644 --- a/service/oauth/oauth.go +++ b/service/oauth/oauth.go @@ -132,7 +132,7 @@ type ( ClientSecret string `form:"client_secret" binding:"required"` GrantType string `form:"grant_type" binding:"required,eq=authorization_code"` Code string `form:"code" binding:"required"` - RedirectURI string `form:"redirect_uri" binding:"required"` + RedirectURI string `form:"redirect_uri"` CodeVerifier string `form:"code_verifier"` } ) @@ -163,7 +163,9 @@ func (s *ExchangeTokenService) Exchange(c *gin.Context) (*TokenResponse, error) if authCode.ClientID != s.ClientID { return nil, serializer.NewError(serializer.CodeCredentialInvalid, "Client ID mismatch", nil) } - if authCode.RedirectURI != s.RedirectURI { + if s.RedirectURI == "" { + dep.Logger().Warning("OAuth client %q did not provide redirect_uri in token request; it may become required in a future release", s.ClientID) + } else if authCode.RedirectURI != s.RedirectURI { return nil, serializer.NewError(serializer.CodeCredentialInvalid, "Redirect URI mismatch", nil) } @@ -275,7 +277,7 @@ func buildIDToken(c *gin.Context, dep dependency.Dep, user *ent.User, clientID s } } - return auth.SignOIDCIDToken(c, dep.SettingClient(), claims) + return auth.SignOIDCIDToken(dep.SettingProvider().OIDCSigningPrivateKey(c), claims) } type ( diff --git a/service/oauth/oidc.go b/service/oauth/oidc.go index 73e4a035..c37861ac 100644 --- a/service/oauth/oidc.go +++ b/service/oauth/oidc.go @@ -7,7 +7,7 @@ import ( "github.com/cloudreve/Cloudreve/v4/application/dependency" "github.com/cloudreve/Cloudreve/v4/inventory/types" "github.com/cloudreve/Cloudreve/v4/pkg/auth" - "github.com/cloudreve/Cloudreve/v4/pkg/setting" + "github.com/cloudreve/Cloudreve/v4/pkg/cluster/routes" "github.com/gin-gonic/gin" ) @@ -19,10 +19,10 @@ func (s *DiscoveryService) Get(c *gin.Context) *DiscoveryResponse { issuer := oidcIssuer(c) return &DiscoveryResponse{ Issuer: issuer.String(), - AuthorizationEndpoint: oidcEndpoint(issuer, "/session/authorize"), - TokenEndpoint: oidcEndpoint(issuer, constants.APIPrefix+"/session/oauth/token"), - UserInfoEndpoint: oidcEndpoint(issuer, constants.APIPrefix+"/session/oauth/userinfo"), - JWKSURI: oidcEndpoint(issuer, constants.APIPrefix+"/session/oauth/jwks"), + AuthorizationEndpoint: routes.MasterOIDCEndpointUrl(issuer, "/session/authorize"), + TokenEndpoint: routes.MasterOIDCEndpointUrl(issuer, constants.APIPrefix+"/session/oauth/token"), + UserInfoEndpoint: routes.MasterOIDCEndpointUrl(issuer, constants.APIPrefix+"/session/oauth/userinfo"), + JWKSURI: routes.MasterOIDCEndpointUrl(issuer, constants.APIPrefix+"/session/oauth/jwks"), ResponseTypesSupported: []string{ "code", }, @@ -60,18 +60,13 @@ func (s *DiscoveryService) Get(c *gin.Context) *DiscoveryResponse { func (s *JWKService) Get(c *gin.Context) (*auth.JWKSet, error) { dep := dependency.FromContext(c) - return auth.OIDCJWKSet(c, dep.SettingClient()) + return auth.OIDCJWKSet(dep.SettingProvider().OIDCSigningPrivateKey(c)) } func oidcIssuer(c *gin.Context) *url.URL { dep := dependency.FromContext(c) - issuer := *dep.SettingProvider().SiteURL(setting.UseFirstSiteUrl(c)) + issuer := *dep.SettingProvider().SiteURL(c) issuer.RawQuery = "" issuer.Fragment = "" return &issuer } - -func oidcEndpoint(issuer *url.URL, endpoint string) string { - route, _ := url.Parse(endpoint) - return issuer.ResolveReference(route).String() -}