feat(oauth): built-in CLI client, loopback redirect, consent denial, PKCE downgrade fix (#231)
Port of upstream cloudreve/cloudreve#3588 adapted to fork conventions: the CLI client is a public client (empty secret, mandatory PKCE) like the desktop/iOS clients rather than carrying a hardcoded secret. Generated with [Devin](https://devin.ai) Co-authored-by: Devin <158243242+devin-ai-integration[bot]@users.noreply.github.com>pull/3589/head
parent
6da4305574
commit
2d9a3543d4
@ -0,0 +1,42 @@
|
|||||||
|
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.Empty(t, cli.Secret)
|
||||||
|
require.Equal(t, []string{"http://127.0.0.1/callback"}, cli.RedirectUris)
|
||||||
|
require.Equal(t, desktop.Scopes, cli.Scopes)
|
||||||
|
require.Equal(t, int64(7776000), cli.Props.RefreshTokenTTL)
|
||||||
|
require.True(t, cli.IsEnabled)
|
||||||
|
|
||||||
|
_, err = client.OAuthClient.UpdateOne(cli).SetIsEnabled(false).SetName("Disabled by administrator").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)
|
||||||
|
count, err := client.OAuthClient.Query().Where(oauthclient.GUID(OAuthClientCLIGUID)).Count(ctx)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, 1, count)
|
||||||
|
}
|
||||||
@ -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