From 9f1e0b8333f5c7304a09e1dcb3f9c170dddd63f3 Mon Sep 17 00:00:00 2001 From: Dyan <0xEFEFEF@gmail.com> Date: Mon, 21 Sep 2026 14:15:15 +0800 Subject: [PATCH] feat(oauth): add OAuth logic for Cloudreve CLI (#3588) --- inventory/migration.go | 31 +++++++ inventory/migration_oauth_test.go | 44 ++++++++++ routers/controllers/oauth.go | 6 ++ routers/router.go | 5 ++ service/oauth/oauth.go | 26 ++++-- service/oauth/oauth_test.go | 138 ++++++++++++++++++++++++++++++ service/oauth/redirect.go | 47 ++++++++++ service/oauth/redirect_test.go | 44 ++++++++++ service/oauth/response.go | 3 +- 9 files changed, 335 insertions(+), 9 deletions(-) create mode 100644 inventory/migration_oauth_test.go create mode 100644 service/oauth/oauth_test.go create mode 100644 service/oauth/redirect.go create mode 100644 service/oauth/redirect_test.go diff --git a/inventory/migration.go b/inventory/migration.go index f581dbc3..9f0cae31 100644 --- a/inventory/migration.go +++ b/inventory/migration.go @@ -280,6 +280,10 @@ const ( OAuthClientDesktopSecret = "8GaQIu3lOSdqYoDHi9cR8IZ4pvuMH8ya" OAuthClientDesktopName = "application:oauth.desktop" OAuthClientDesktopRedirectURI = "/callback/desktop" + OAuthClientCLIGUID = "6326d2af-2fef-4a99-94da-1ee8ef0ca53f" + OAuthClientCLISecret = "yoFiNgbxvSCzK2Nm92T3TNaREh4qBjq4" + OAuthClientCLIName = "Cloudreve CLI" + OAuthClientCLIRedirectURI = "http://127.0.0.1/callback" OAuthClientiOSGUID = "220db97a-44a3-44f7-99b6-d767262b4daa" OAuthClientiOSSecret = "1kxOW4IyVOkPlsKCnTwzfHyP8XrbpfaF" OAuthClientiOSName = "application:setting.iOSApp" @@ -295,6 +299,10 @@ func migrateOAuthClient(l logging.Logger, client *ent.Client, ctx context.Contex return err } + if err := migrateOAuthClientCLI(l, client, ctx); err != nil { + return err + } + return nil } @@ -339,6 +347,29 @@ func migrateOAuthClientDesktop(l logging.Logger, client *ent.Client, ctx context return nil } +// migrateOAuthClientCLI preserves administrator changes to an existing built-in client. +func migrateOAuthClientCLI(l logging.Logger, client *ent.Client, ctx context.Context) error { + if _, err := client.OAuthClient.Query().Where(oauthclient.GUID(OAuthClientCLIGUID)).First(ctx); err == nil { + l.Info("Default OAuth client (GUID=%s) already exists, skip migrating.", OAuthClientCLIGUID) + return nil + } else if !ent.IsNotFound(err) { + return fmt.Errorf("failed to query default CLI OAuth client: %w", err) + } + + if _, err := client.OAuthClient.Create(). + SetGUID(OAuthClientCLIGUID). + SetSecret(OAuthClientCLISecret). + SetName(OAuthClientCLIName). + SetRedirectUris([]string{OAuthClientCLIRedirectURI}). + SetScopes([]string{"profile", "email", "openid", "offline_access", "UserInfo.Write", "Workflow.Write", "Files.Write", "Shares.Write"}). + SetProps(&types.OAuthClientProps{Icon: "/static/img/cloudreve.svg", RefreshTokenTTL: 7776000}). + SetIsEnabled(true). + Save(ctx); err != nil { + return fmt.Errorf("failed to create default CLI OAuth client: %w", err) + } + return nil +} + type ( PatchFunc func(l logging.Logger, client *ent.Client, ctx context.Context) error Patch struct { diff --git a/inventory/migration_oauth_test.go b/inventory/migration_oauth_test.go new file mode 100644 index 00000000..7e554350 --- /dev/null +++ b/inventory/migration_oauth_test.go @@ -0,0 +1,44 @@ +package inventory + +import ( + "context" + "path/filepath" + "testing" + + "github.com/cloudreve/Cloudreve/v4/ent/enttest" + "github.com/cloudreve/Cloudreve/v4/ent/oauthclient" + "github.com/cloudreve/Cloudreve/v4/pkg/logging" + "github.com/stretchr/testify/require" +) + +func TestBuiltinCLIOAuthMigration(t *testing.T) { + ctx := context.Background() + client := enttest.Open(t, "sqlite3", filepath.Join(t.TempDir(), "migration.db")) + t.Cleanup(func() { require.NoError(t, client.Close()) }) + logger := logging.NewConsoleLogger(logging.LevelError) + + require.NoError(t, migrateOAuthClient(logger, client, ctx)) + desktop, err := client.OAuthClient.Query().Where(oauthclient.GUID(OAuthClientDesktopGUID)).Only(ctx) + require.NoError(t, err) + cli, err := client.OAuthClient.Query().Where(oauthclient.GUID(OAuthClientCLIGUID)).Only(ctx) + require.NoError(t, err) + require.Equal(t, "Cloudreve CLI", cli.Name) + require.Equal(t, OAuthClientCLISecret, cli.Secret) + require.Equal(t, []string{"http://127.0.0.1/callback"}, cli.RedirectUris) + require.Equal(t, desktop.Scopes, cli.Scopes) + require.Equal(t, desktop.Props, cli.Props) + require.Equal(t, int64(7776000), cli.Props.RefreshTokenTTL) + require.True(t, cli.IsEnabled) + + _, err = client.OAuthClient.UpdateOne(cli).SetIsEnabled(false).SetName("Disabled by administrator").SetSecret("administrator-secret").Save(ctx) + require.NoError(t, err) + require.NoError(t, migrateOAuthClient(logger, client, ctx)) + preserved, err := client.OAuthClient.Get(ctx, cli.ID) + require.NoError(t, err) + require.False(t, preserved.IsEnabled) + require.Equal(t, "Disabled by administrator", preserved.Name) + require.Equal(t, "administrator-secret", preserved.Secret) + count, err := client.OAuthClient.Query().Where(oauthclient.GUID(OAuthClientCLIGUID)).Count(ctx) + require.NoError(t, err) + require.Equal(t, 1, count) +} diff --git a/routers/controllers/oauth.go b/routers/controllers/oauth.go index 9c4472bf..da2adab6 100644 --- a/routers/controllers/oauth.go +++ b/routers/controllers/oauth.go @@ -30,6 +30,12 @@ func GrantAppConsent(c *gin.Context) { c.JSON(200, serializer.Response{Data: res}) } +// DenyAppConsent validates the client redirect without granting access. +func DenyAppConsent(c *gin.Context) { + ParametersFromContext[*oauth.GrantService](c, oauth.GrantParamCtx{}).Deny = true + GrantAppConsent(c) +} + type ExchangeErrorResponse struct { Error string `json:"error"` ErrorDescription string `json:"error_description"` diff --git a/routers/router.go b/routers/router.go index 846f81af..22a613bc 100644 --- a/routers/router.go +++ b/routers/router.go @@ -329,6 +329,11 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { controllers.FromJSON[oauth.GrantService](oauth.GrantParamCtx{}), controllers.GrantAppConsent, ) + oauthRouter.POST("consent/deny", + middleware.LoginRequired(), + controllers.FromJSON[oauth.GrantService](oauth.GrantParamCtx{}), + controllers.DenyAppConsent, + ) oauthRouter.POST("token", controllers.FromForm[oauth.ExchangeTokenService](oauth.ExchangeTokenParamCtx{}), controllers.ExchangeToken, diff --git a/service/oauth/oauth.go b/service/oauth/oauth.go index a15940e1..407b0d7e 100644 --- a/service/oauth/oauth.go +++ b/service/oauth/oauth.go @@ -47,6 +47,7 @@ type ( GrantParamCtx struct{} GrantService struct { ClientID string `json:"client_id" binding:"required"` + Deny bool `json:"-"` ResponseType string `json:"response_type" binding:"required,eq=code"` RedirectURI string `json:"redirect_uri" binding:"required"` State string `json:"state" binding:"max=4096"` @@ -74,7 +75,7 @@ func (s *GrantService) Get(c *gin.Context) (*GrantResponse, error) { // 2. Validate redirect URL: must match one of the registered redirect URIs redirectValid := false for _, uri := range app.RedirectUris { - if uri == s.RedirectURI { + if redirectURIMatches(uri, s.RedirectURI) { redirectValid = true break } @@ -83,6 +84,10 @@ func (s *GrantService) Get(c *gin.Context) (*GrantResponse, error) { return nil, serializer.NewError(serializer.CodeParamErr, "Invalid redirect URI", nil) } + if s.Deny { + return &GrantResponse{Error: "access_denied", State: s.State}, nil + } + // Parse requested scopes (space-separated per OAuth 2.0 spec) requestedScopes := strings.Split(s.Scope, " ") @@ -125,11 +130,12 @@ func (s *GrantService) Get(c *gin.Context) (*GrantResponse, error) { type ( ExchangeTokenParamCtx struct{} ExchangeTokenService struct { - ClientID string `form:"client_id" binding:"required"` - ClientSecret string `form:"client_secret" binding:"required"` - GrantType string `form:"grant_type" binding:"required,eq=authorization_code"` - Code string `form:"code" binding:"required"` - CodeVerifier string `form:"code_verifier"` + ClientID string `form:"client_id" binding:"required"` + ClientSecret string `form:"client_secret" binding:"required"` + GrantType string `form:"grant_type" binding:"required,eq=authorization_code"` + Code string `form:"code" binding:"required"` + CodeVerifier string `form:"code_verifier"` + RedirectURI *string `form:"redirect_uri"` } ) @@ -160,8 +166,12 @@ func (s *ExchangeTokenService) Exchange(c *gin.Context) (*TokenResponse, error) return nil, serializer.NewError(serializer.CodeCredentialInvalid, "Client ID mismatch", nil) } - // 3. Verify PKCE: SHA256(code_verifier) should match code_challenge - if authCode.CodeChallenge != "" { + if s.RedirectURI != nil && *s.RedirectURI != authCode.RedirectURI { + return nil, serializer.NewError(serializer.CodeCredentialInvalid, "Redirect URI mismatch", nil) + } + + // 3. Verify PKCE, rejecting a verifier when the grant omitted its challenge. + if authCode.CodeChallenge != "" || s.CodeVerifier != "" { verifierHash := sha256.Sum256([]byte(s.CodeVerifier)) expectedChallenge := base64.RawURLEncoding.EncodeToString(verifierHash[:]) if expectedChallenge != authCode.CodeChallenge { diff --git a/service/oauth/oauth_test.go b/service/oauth/oauth_test.go new file mode 100644 index 00000000..29d72292 --- /dev/null +++ b/service/oauth/oauth_test.go @@ -0,0 +1,138 @@ +package oauth + +import ( + "context" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "errors" + "net/http/httptest" + "testing" + + "github.com/cloudreve/Cloudreve/v4/application/dependency" + "github.com/cloudreve/Cloudreve/v4/ent" + "github.com/cloudreve/Cloudreve/v4/inventory" + "github.com/cloudreve/Cloudreve/v4/pkg/auth" + "github.com/cloudreve/Cloudreve/v4/pkg/cache" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type oauthTestClient struct { + inventory.OAuthClientClient + app *ent.OAuthClient + grants, reads int +} + +func (c *oauthTestClient) GetByGUIDWithGrants(context.Context, string, int) (*ent.OAuthClient, error) { + if c.app == nil { + return nil, errors.New("unknown client") + } + return c.app, nil +} +func (c *oauthTestClient) UpsertGrant(context.Context, int, int, []string) error { + c.grants++ + return nil +} +func (c *oauthTestClient) GetByGUID(context.Context, string) (*ent.OAuthClient, error) { + c.reads++ + return nil, errors.New("stop after grant verification") +} + +type oauthTestCache struct { + cache.Driver + code *AuthorizationCode + writes int +} + +func (c *oauthTestCache) Get(string) (any, bool) { return c.code, c.code != nil } +func (c *oauthTestCache) Set(string, any, int) error { c.writes++; return nil } +func (c *oauthTestCache) Delete(string, ...string) error { return nil } + +type oauthTestDep struct { + dependency.Dep + client *oauthTestClient + kv *oauthTestCache +} + +func (d oauthTestDep) OAuthClientClient() inventory.OAuthClientClient { return d.client } +func (d oauthTestDep) KV() cache.Driver { return d.kv } +func (d oauthTestDep) UserClient() inventory.UserClient { return nil } +func (d oauthTestDep) TokenAuth() auth.TokenAuth { return nil } + +func oauthTestContext(client *oauthTestClient, kv *oauthTestCache) *gin.Context { + gin.SetMode(gin.TestMode) + c, engine := gin.CreateTestContext(httptest.NewRecorder()) + engine.ContextWithFallback = true + ctx := context.WithValue(context.Background(), dependency.DepCtx{}, oauthTestDep{client: client, kv: kv}) + ctx = context.WithValue(ctx, inventory.UserCtx{}, &ent.User{ID: 1}) + c.Request = httptest.NewRequest("POST", "/", nil).WithContext(ctx) + return c +} + +func TestConsentDenialValidatesRedirectWithoutIssuingGrant(t *testing.T) { + client := &oauthTestClient{app: &ent.OAuthClient{RedirectUris: []string{"http://127.0.0.1/callback"}}} + kv := &oauthTestCache{} + c := oauthTestContext(client, kv) + s := GrantService{ClientID: "client", RedirectURI: "http://127.0.0.1:49152/callback", State: "state", Deny: true} + response, err := s.Get(c) + require.NoError(t, err) + require.Equal(t, "access_denied", response.Error) + require.Equal(t, "state", response.State) + require.Empty(t, response.Code) + require.Zero(t, client.grants) + require.Zero(t, kv.writes) + + s.RedirectURI = "http://evil.example/callback" + response, err = s.Get(c) + require.ErrorContains(t, err, "Invalid redirect URI") + require.Nil(t, response) + client.app = nil + s.RedirectURI = "http://127.0.0.1:49152/callback" + response, err = s.Get(c) + require.ErrorContains(t, err, "App not found") + require.Nil(t, response) + require.Zero(t, client.grants) + require.Zero(t, kv.writes) +} + +func TestTokenExchangeBindsProvidedRedirectAndRejectsPKCEDowngrade(t *testing.T) { + uri := "http://127.0.0.1:49152/callback" + other := "http://127.0.0.1:49153/callback" + empty := "" + verifier := "verifier" + hash := sha256.Sum256([]byte(verifier)) + challenge := base64.RawURLEncoding.EncodeToString(hash[:]) + for _, test := range []struct { + name string + redirect *string + challenge, verifier string + valid bool + }{ + {"legacy omission", nil, "", "", true}, + {"exact redirect", &uri, challenge, verifier, true}, + {"other port", &other, challenge, verifier, false}, + {"empty provided redirect", &empty, "", "", false}, + {"PKCE downgrade", &uri, "", verifier, false}, + {"missing verifier", &uri, challenge, "", false}, + } { + t.Run(test.name, func(t *testing.T) { + client := &oauthTestClient{} + kv := &oauthTestCache{code: &AuthorizationCode{ClientID: "client", RedirectURI: uri, CodeChallenge: test.challenge}} + s := ExchangeTokenService{ClientID: "client", RedirectURI: test.redirect, CodeVerifier: test.verifier} + _, err := s.Exchange(oauthTestContext(client, kv)) + require.Error(t, err) + if test.valid { + require.Equal(t, 1, client.reads, "valid grant reaches client credential validation") + } else { + require.Zero(t, client.reads, "invalid grant must fail before client credential validation") + } + }) + } +} + +func TestConsentRequestCannotSetInternalDenial(t *testing.T) { + var service GrantService + require.NoError(t, json.Unmarshal([]byte(`{"deny":true}`), &service)) + require.False(t, service.Deny) +} diff --git a/service/oauth/redirect.go b/service/oauth/redirect.go new file mode 100644 index 00000000..4e27eb7d --- /dev/null +++ b/service/oauth/redirect.go @@ -0,0 +1,47 @@ +package oauth + +import ( + "net/netip" + "net/url" + "strconv" + "strings" +) + +// redirectURIMatches permits only the port variation required for native loopback clients by RFC 8252. +func redirectURIMatches(registered, requested string) bool { + registration, err := url.Parse(registered) + if err != nil { + return false + } + if registration.Scheme != "http" { + return registered == requested + } + address, err := netip.ParseAddr(registration.Hostname()) + if err != nil || !address.IsLoopback() { + return registered == requested + } + + redirect, err := url.Parse(requested) + if err != nil { + return false + } + for _, uri := range []*url.URL{registration, redirect} { + if uri.User != nil || uri.Fragment != "" || uri.RawFragment != "" || uri.Opaque != "" || strings.HasSuffix(uri.Host, ":") { + return false + } + if port := uri.Port(); port != "" { + number, err := strconv.Atoi(port) + if err != nil || number < 1 || number > 65535 { + return false + } + } + } + if strings.Contains(registered, "#") || strings.Contains(requested, "#") { + return false + } + return registration.Scheme == redirect.Scheme && + registration.Hostname() == redirect.Hostname() && + registration.EscapedPath() == redirect.EscapedPath() && + registration.RawQuery == redirect.RawQuery && + registration.ForceQuery == redirect.ForceQuery +} diff --git a/service/oauth/redirect_test.go b/service/oauth/redirect_test.go new file mode 100644 index 00000000..557f407d --- /dev/null +++ b/service/oauth/redirect_test.go @@ -0,0 +1,44 @@ +package oauth + +import "testing" + +func TestRedirectURIMatches(t *testing.T) { + for _, test := range []struct { + registered string + requested string + matches bool + }{ + {"http://127.0.0.1/callback", "http://127.0.0.1:49152/callback", true}, + {"http://[::1]/callback", "http://[::1]:49152/callback", true}, + {"http://127.0.0.1:8000/callback?a=1", "http://127.0.0.1:65535/callback?a=1", true}, + {"/callback/desktop", "/callback/desktop", true}, + {"cloudreve://mount", "cloudreve://mount", true}, + {"https://app.example/callback", "https://app.example/callback", true}, + {"http://localhost/callback", "http://localhost:49152/callback", false}, + {"https://127.0.0.1/callback", "https://127.0.0.1:49152/callback", false}, + {"http://192.168.0.1/callback", "http://192.168.0.1:49152/callback", false}, + {"http://127.0.0.1/callback", "http://127.0.0.2:49152/callback", false}, + {"http://127.0.0.1/callback", "http://127.0.0.1.evil.example:49152/callback", false}, + {"http://127.0.0.1/callback", "http://user@127.0.0.1:49152/callback", false}, + {"http://user@127.0.0.1/callback", "http://user@127.0.0.1/callback", false}, + {"http://127.0.0.1/callback", "https://127.0.0.1:49152/callback", false}, + {"http://127.0.0.1/callback", "http://127.0.0.1:0/callback", false}, + {"http://127.0.0.1:0/callback", "http://127.0.0.1:0/callback", false}, + {"http://127.0.0.1/callback", "http://127.0.0.1:65536/callback", false}, + {"http://127.0.0.1/callback", "http://127.0.0.1:abc/callback", false}, + {"http://127.0.0.1:abc/callback", "http://127.0.0.1:abc/callback", false}, + {"http://127.0.0.1/callback", "http://127.0.0.1:/callback", false}, + {"http://127.0.0.1/callback", "http://127.0.0.1:49152/callback#", false}, + {"http://127.0.0.1/callback", "http://127.0.0.1:49152/callback#fragment", false}, + {"http://127.0.0.1/callback", "http://127.0.0.1:49152/other", false}, + {"http://127.0.0.1/callback", "http://127.0.0.1:49152/%63allback", false}, + {"http://127.0.0.1/callback?a=1", "http://127.0.0.1:49152/callback?a=2", false}, + {"http://127.0.0.1/callback", "http://127.0.0.1:49152/callback?", false}, + } { + t.Run(test.registered+" -> "+test.requested, func(t *testing.T) { + if got := redirectURIMatches(test.registered, test.requested); got != test.matches { + t.Fatalf("matches = %v, want %v", got, test.matches) + } + }) + } +} diff --git a/service/oauth/response.go b/service/oauth/response.go index 49a99df1..aa8c337f 100644 --- a/service/oauth/response.go +++ b/service/oauth/response.go @@ -39,7 +39,8 @@ func BuildAppRegistration(app *ent.OAuthClient, grant *ent.OAuthGrant) *AppRegis } type GrantResponse struct { - Code string `json:"code"` + Code string `json:"code,omitempty"` + Error string `json:"error,omitempty"` State string `json:"state"` }