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