diff --git a/internal/auth/auth.go b/internal/auth/auth.go index 0e80d6b..b31daa9 100644 --- a/internal/auth/auth.go +++ b/internal/auth/auth.go @@ -250,6 +250,10 @@ func openBrowser(url string) error { return cmd.Start() } +func callbackRedirectURI(addr net.Addr) string { + return fmt.Sprintf("http://%s/callback", addr.String()) +} + // Login runs the full OAuth2 Authorization Code + PKCE flow. // Returns a credential to store. func Login(server string) (*config.Credential, error) { @@ -274,8 +278,7 @@ func Login(server string) (*config.Credential, error) { if err != nil { return nil, err } - port := listener.Addr().(*net.TCPAddr).Port - redirectURI := fmt.Sprintf("http://localhost:%d/callback", port) + redirectURI := callbackRedirectURI(listener.Addr()) // Register client clientInfo, err := registerClient(regEndpoint, redirectURI) diff --git a/internal/auth/auth_test.go b/internal/auth/auth_test.go index 30f16c5..ae0e23f 100644 --- a/internal/auth/auth_test.go +++ b/internal/auth/auth_test.go @@ -3,6 +3,9 @@ package auth import ( "crypto/rand" "crypto/rsa" + "net" + "net/url" + "strconv" "testing" "time" @@ -10,6 +13,30 @@ import ( "github.com/go-jose/go-jose/v4/jwt" ) +func TestCallbackRedirectURIMatchesLoopbackListener(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + _ = listener.Close() + }) + + redirect, err := url.Parse(callbackRedirectURI(listener.Addr())) + if err != nil { + t.Fatal(err) + } + if redirect.Hostname() != "127.0.0.1" { + t.Fatalf("redirect hostname = %q, want 127.0.0.1", redirect.Hostname()) + } + if redirect.Port() != strconv.Itoa(listener.Addr().(*net.TCPAddr).Port) { + t.Fatalf("redirect port = %q, listener = %q", redirect.Port(), listener.Addr()) + } + if redirect.Path != "/callback" { + t.Fatalf("redirect path = %q, want /callback", redirect.Path) + } +} + func TestVerifiedTokenToCredentialUsesFallbacks(t *testing.T) { vt := &VerifiedToken{ AccessToken: "access",