feat(oauth): add OAuth logic for Cloudreve CLI (#3588)
parent
a8becb9f5b
commit
9f1e0b8333
@ -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)
|
||||
}
|
||||
@ -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)
|
||||
}
|
||||
@ -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
|
||||
}
|
||||
@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Loading…
Reference in new issue