Skip to content

Commit 7a5b0f0

Browse files
fix(auth): harden Windows OAuth callback handling, preserve PKCE login flow on Windows
1 parent 00c2b27 commit 7a5b0f0

2 files changed

Lines changed: 56 additions & 9 deletions

File tree

internal/cli/app_test.go

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@ import (
66
"io"
77
"net/http"
88
"net/http/httptest"
9+
"net/url"
910
"os"
1011
"path/filepath"
1112
"regexp"
@@ -561,6 +562,36 @@ func TestWaitForOAuthCallbackMismatchAndTimeout(t *testing.T) {
561562
}
562563
}
563564

565+
func TestWaitForOAuthCallbackAcceptsIPv4WhenRedirectUsesLocalhost(t *testing.T) {
566+
server, err := waitForOAuthCallback("expected-state", time.Second)
567+
if err != nil {
568+
t.Fatal(err)
569+
}
570+
defer server.Close()
571+
if !strings.HasPrefix(server.RedirectURI, "http://localhost:") {
572+
t.Fatalf("expected localhost redirect URI for OAuth compatibility, got %s", server.RedirectURI)
573+
}
574+
if len(server.listeners) == 0 {
575+
t.Fatal("expected at least one loopback listener")
576+
}
577+
parsed, err := url.Parse(server.RedirectURI)
578+
if err != nil {
579+
t.Fatal(err)
580+
}
581+
resp, err := http.Get("http://127.0.0.1:" + parsed.Port() + "/oauth/callback?code=test-code&state=expected-state")
582+
if err != nil {
583+
t.Fatal(err)
584+
}
585+
resp.Body.Close()
586+
payload, err := server.Wait()
587+
if err != nil {
588+
t.Fatal(err)
589+
}
590+
if payload.Code != "test-code" || payload.State != "expected-state" {
591+
t.Fatalf("unexpected callback payload: %+v", payload)
592+
}
593+
}
594+
564595
func TestExchangeAuthorizationCodeFailureAndScopeArray(t *testing.T) {
565596
failServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
566597
w.WriteHeader(http.StatusBadRequest)

internal/cli/auth.go

Lines changed: 25 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -180,7 +180,7 @@ type callbackServer struct {
180180
wait chan callbackPayload
181181
errs chan error
182182
server *http.Server
183-
listener net.Listener
183+
listeners []net.Listener
184184
}
185185

186186
type callbackPayload struct {
@@ -195,19 +195,24 @@ func waitForOAuthCallback(expectedState string, timeout time.Duration) (*callbac
195195
srv := &http.Server{
196196
Handler: mux,
197197
// Set ReadHeaderTimeout to mitigate Slowloris attacks (gosec G112).
198-
// Even though this listens only on 127.0.0.1, we still bound it.
198+
// Even though this listens only on loopback interfaces, we still bound it.
199199
ReadHeaderTimeout: 10 * time.Second,
200200
}
201-
ln, err := net.Listen("tcp", "127.0.0.1:0")
201+
ln4, err := net.Listen("tcp4", "127.0.0.1:0")
202202
if err != nil {
203203
return nil, err
204204
}
205+
port := ln4.Addr().(*net.TCPAddr).Port
206+
listeners := []net.Listener{ln4}
207+
if ln6, err := net.Listen("tcp6", fmt.Sprintf("[::1]:%d", port)); err == nil {
208+
listeners = append(listeners, ln6)
209+
}
205210
cs := &callbackServer{
206-
RedirectURI: fmt.Sprintf("http://localhost:%d/oauth/callback", ln.Addr().(*net.TCPAddr).Port),
211+
RedirectURI: fmt.Sprintf("http://localhost:%d/oauth/callback", port),
207212
wait: wait,
208213
errs: errs,
209214
server: srv,
210-
listener: ln,
215+
listeners: listeners,
211216
}
212217
mux.HandleFunc("/oauth/callback", func(w http.ResponseWriter, r *http.Request) {
213218
code := r.URL.Query().Get("code")
@@ -228,9 +233,11 @@ func waitForOAuthCallback(expectedState string, timeout time.Duration) (*callbac
228233
wait <- callbackPayload{Code: code, State: state}
229234
}
230235
})
231-
go func() {
232-
_ = srv.Serve(ln)
233-
}()
236+
for _, listener := range listeners {
237+
go func(ln net.Listener) {
238+
_ = srv.Serve(ln)
239+
}(listener)
240+
}
234241
go func() {
235242
<-time.After(timeout)
236243
errs <- errors.New("Timed out waiting for the OAuth callback. Re-run with --no-browser to copy the URL manually, or check that your browser completed the login flow.")
@@ -248,7 +255,16 @@ func (c *callbackServer) Wait() (callbackPayload, error) {
248255
}
249256

250257
func (c *callbackServer) Close() error {
251-
return c.server.Close()
258+
var firstErr error
259+
if err := c.server.Close(); err != nil {
260+
firstErr = err
261+
}
262+
for _, listener := range c.listeners {
263+
if err := listener.Close(); err != nil && !errors.Is(err, net.ErrClosed) && firstErr == nil {
264+
firstErr = err
265+
}
266+
}
267+
return firstErr
252268
}
253269

254270
type tokenResponse struct {

0 commit comments

Comments
 (0)