Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
242 changes: 105 additions & 137 deletions authexternalbrowser.go
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package gosnowflake

import (
"bufio"
"bytes"
"context"
"encoding/base64"
Expand Down Expand Up @@ -300,162 +299,131 @@ func validExternalBrowserPreflight(request *http.Request, accountURL *url.URL) b
return true
}

func writeExternalBrowserResponse(conn net.Conn, statusCode int, headers http.Header, body string) error {
response := &http.Response{
Status: fmt.Sprintf("%d %s", statusCode, http.StatusText(statusCode)),
StatusCode: statusCode,
Proto: "HTTP/1.1",
ProtoMajor: 1,
ProtoMinor: 1,
Body: io.NopCloser(strings.NewReader(body)),
ContentLength: int64(len(body)),
Header: headers,
func writeExternalBrowserResponse(w http.ResponseWriter, statusCode int, headers http.Header, body string) error {
for key, values := range headers {
w.Header()[key] = values
}
return response.Write(conn)
}

func closeExternalBrowserConnection(conn net.Conn) {
if err := conn.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
logger.Warnf("error while closing browser connection: %v", err)
w.Header().Set("Connection", "close")
w.Header().Set("Content-Length", strconv.Itoa(len(body)))
if body != "" {
w.Header().Set("Content-Type", "text/html; charset=utf-8")
}
w.WriteHeader(statusCode)
if _, err := io.WriteString(w, body); err != nil {
return err
}
// The receiver closes the server as soon as a token is published. Flush a
// complete response first so the browser still receives the success page.
return http.NewResponseController(w).Flush()
}

func receiveExternalBrowserCallback(ctx context.Context, listener *net.TCPListener, accountURL *url.URL, application string) (string, error) {
stopListenerClose := context.AfterFunc(ctx, func() {
if err := listener.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
logger.Warnf("error while closing external browser listener: %v", err)
}
})
defer stopListenerClose()

for {
conn, err := listener.Accept()
if err != nil {
if ctx.Err() != nil {
return "", ctx.Err()
func receiveExternalBrowserCallback(ctx context.Context, listener net.Listener, accountURL *url.URL, application string) (string, error) {
tokens := make(chan string, 1)
server := &http.Server{
ReadTimeout: 10 * time.Second,
WriteTimeout: 10 * time.Second,
Handler: http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
// GET callbacks do not use the body. Do not wait for a client to send
// an unused body before writing its response.
if err := http.NewResponseController(w).EnableFullDuplex(); err != nil {
logger.Debugf("unable to enable external browser callback response: %v", err)
return
}
return "", err
}

if deadline, ok := ctx.Deadline(); ok {
if err := conn.SetDeadline(deadline); err != nil {
closeExternalBrowserConnection(conn)
if ctx.Err() != nil {
return "", ctx.Err()
respond := func(statusCode int, headers http.Header, body string) {
if err := writeExternalBrowserResponse(w, statusCode, headers, body); err != nil {
logger.Debugf("unable to write external browser callback response: %v", err)
}
logger.Debugf("unable to set external browser callback deadline: %v", err)
continue
}
}
stopConnectionClose := context.AfterFunc(ctx, func() {
closeExternalBrowserConnection(conn)
})
request, readErr := http.ReadRequest(bufio.NewReader(conn))
stopConnectionClose()
if readErr != nil {
if ctx.Err() != nil {
closeExternalBrowserConnection(conn)
return "", ctx.Err()
}
if writeErr := writeExternalBrowserResponse(conn, http.StatusBadRequest, make(http.Header), ""); writeErr != nil {
logger.Debugf("unable to write invalid external browser callback response: %v", writeErr)
if request.Method == http.MethodOptions {
headers := make(http.Header)
statusCode := http.StatusForbidden
if validExternalBrowserPreflight(request, accountURL) {
origin := request.Header.Get("Origin")
headers.Set("Access-Control-Allow-Origin", origin)
headers.Set("Access-Control-Allow-Methods", "POST, GET, OPTIONS")
headers.Set("Access-Control-Allow-Headers", httpHeaderContentType)
statusCode = http.StatusOK
}
respond(statusCode, headers, "")
return
}
closeExternalBrowserConnection(conn)
logger.Debug("ignoring invalid external browser callback request")
continue
}

if request.Method == http.MethodOptions {
headers := make(http.Header)
statusCode := http.StatusForbidden
if validExternalBrowserPreflight(request, accountURL) {
origin := request.Header.Get("Origin")
headers.Set("Access-Control-Allow-Origin", origin)
headers.Set("Access-Control-Allow-Methods", "POST, GET, OPTIONS")
headers.Set("Access-Control-Allow-Headers", httpHeaderContentType)
statusCode = http.StatusOK
origins := request.Header.Values("Origin")
originPresent := len(origins) != 0
origin := ""
if len(origins) == 1 {
origin = origins[0]
}
if writeErr := writeExternalBrowserResponse(conn, statusCode, headers, ""); writeErr != nil {
if ctx.Err() != nil {
closeExternalBrowserConnection(conn)
return "", ctx.Err()
}
logger.Debugf("unable to write external browser preflight response: %v", writeErr)
originMatchesAccount := len(origins) == 1 &&
!strings.EqualFold(strings.TrimSpace(origin), "null") &&
externalBrowserOriginMatchesAccount(origin, accountURL)
rejectOrigin := false
if request.Method == http.MethodPost {
rejectOrigin = !originMatchesAccount
} else {
rejectOrigin = originPresent && !strings.EqualFold(strings.TrimSpace(origin), "null") && !originMatchesAccount
}
closeExternalBrowserConnection(conn)
continue
}

origins := request.Header.Values("Origin")
originPresent := len(origins) != 0
origin := ""
if len(origins) == 1 {
origin = origins[0]
}
originMatchesAccount := len(origins) == 1 &&
!strings.EqualFold(strings.TrimSpace(origin), "null") &&
externalBrowserOriginMatchesAccount(origin, accountURL)
rejectOrigin := false
if request.Method == http.MethodPost {
rejectOrigin = !originMatchesAccount
} else {
rejectOrigin = originPresent && !strings.EqualFold(strings.TrimSpace(origin), "null") && !originMatchesAccount
}
if rejectOrigin {
if writeErr := writeExternalBrowserResponse(conn, http.StatusForbidden, make(http.Header), ""); writeErr != nil {
logger.Debugf("unable to write external browser callback rejection: %v", writeErr)
if rejectOrigin {
respond(http.StatusForbidden, nil, "")
return
}
closeExternalBrowserConnection(conn)
continue
}

var encodedSamlResponse string
switch request.Method {
case http.MethodPost:
encodedSamlResponse, err = getTokenFromPostRequest(request)
if err != nil || encodedSamlResponse == "" {
if writeErr := writeExternalBrowserResponse(conn, http.StatusBadRequest, make(http.Header), ""); writeErr != nil {
logger.Debugf("unable to write external browser callback response without token: %v", writeErr)
var encodedSamlResponse string
var err error
switch request.Method {
case http.MethodPost:
encodedSamlResponse, err = getTokenFromPostRequest(request)
case http.MethodGet:
if request.URL.Path != "/" || !strings.HasPrefix(request.URL.RawQuery, "token=") {
respond(http.StatusBadRequest, nil, "")
return
}
closeExternalBrowserConnection(conn)
continue
encodedSamlResponse, err = getTokenFromResponse(
request.Method + " " + request.RequestURI + " HTTP/1.1\r\n",
)
default:
respond(http.StatusMethodNotAllowed, nil, "")
return
}
case http.MethodGet:
if request.URL.Path != "/" || !strings.HasPrefix(request.URL.RawQuery, "token=") {
if writeErr := writeExternalBrowserResponse(conn, http.StatusBadRequest, make(http.Header), ""); writeErr != nil {
logger.Debugf("unable to write external browser callback response without token: %v", writeErr)
}
closeExternalBrowserConnection(conn)
continue
}
encodedSamlResponse, err = getTokenFromResponse(
request.Method + " " + request.RequestURI + " HTTP/1.1\r\n",
)
if err != nil || encodedSamlResponse == "" {
if writeErr := writeExternalBrowserResponse(conn, http.StatusBadRequest, make(http.Header), ""); writeErr != nil {
logger.Debugf("unable to write invalid external browser callback response: %v", writeErr)
}
closeExternalBrowserConnection(conn)
continue
respond(http.StatusBadRequest, nil, "")
return
}
body := fmt.Sprintf(samlSuccessHTML, application)
headers := make(http.Header)
if origin != "" && !strings.EqualFold(strings.TrimSpace(origin), "null") {
headers.Set("Access-Control-Allow-Origin", origin)
headers.Set("Vary", "Origin")
}
default:
if writeErr := writeExternalBrowserResponse(conn, http.StatusMethodNotAllowed, make(http.Header), ""); writeErr != nil {
logger.Debugf("unable to write unsupported external browser callback response: %v", writeErr)
respond(http.StatusOK, headers, body)
select {
case tokens <- encodedSamlResponse:
default:
}
closeExternalBrowserConnection(conn)
continue
}
body := fmt.Sprintf(samlSuccessHTML, application)
headers := make(http.Header)
if origin != "" && !strings.EqualFold(strings.TrimSpace(origin), "null") {
headers.Set("Access-Control-Allow-Origin", origin)
headers.Set("Vary", "Origin")
}),
}
serveErrors := make(chan error, 1)
// Each connection is handled concurrently, so idle preconnects and partial
// requests cannot block a later authentication callback.
go func() {
serveErrors <- server.Serve(listener)
}()
defer func() {
if err := server.Close(); err != nil {
logger.Warnf("error while closing external browser callback server: %v", err)
}
if err = writeExternalBrowserResponse(conn, http.StatusOK, headers, body); err != nil {
logger.Debugf("unable to write successful external browser callback response: %v", err)
}()

select {
case token := <-tokens:
return token, nil
case <-ctx.Done():
return "", ctx.Err()
case err := <-serveErrors:
if ctx.Err() != nil {
return "", ctx.Err()
}
closeExternalBrowserConnection(conn)
return encodedSamlResponse, nil
return "", err
}
}

Expand Down
111 changes: 111 additions & 0 deletions authexternalbrowser_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -591,3 +591,114 @@ func (provider *nonInteractiveSamlResponseProvider) run(url string) error {
}()
return nil
}

// Observe Accept so the real callback cannot accidentally win the race against
// the idle connection in the regression test.
type externalBrowserObservedListener struct {
net.Listener
accepted chan struct{}
}

func (l *externalBrowserObservedListener) Accept() (net.Conn, error) {
conn, err := l.Listener.Accept()
if err == nil {
l.accepted <- struct{}{}
}
return conn, err
}

func TestExternalBrowserCallbackIgnoresIdleConnection(t *testing.T) {
for _, initialRequest := range []string{
"",
"GET /?token=incomplete HTTP/1.1\r\nHost:",
"POST / HTTP/1.1\r\nHost: localhost\r\nOrigin: https://account.example.com\r\nContent-Length: 100\r\n\r\n{",
} {
t.Run(fmt.Sprintf("initialBytes=%d", len(initialRequest)), func(t *testing.T) {
accountURL, err := url.Parse("https://account.example.com:443")
assertNilF(t, err)
listener, err := createLocalTCPListener(0)
assertNilF(t, err)
defer listener.Close()
observed := &externalBrowserObservedListener{listener, make(chan struct{}, 4)}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
result := make(chan externalBrowserCallbackResult, 1)
go func() {
token, err := receiveExternalBrowserCallback(ctx, observed, accountURL, "Go")
result <- externalBrowserCallbackResult{token: token, err: err}
}()

idle, err := net.DialTimeout("tcp", listener.Addr().String(), time.Second)
assertNilF(t, err)
defer idle.Close()
_, err = io.WriteString(idle, initialRequest)
assertNilF(t, err)
select {
case <-observed.accepted:
case <-ctx.Done():
t.Fatal("idle connection was not accepted")
}

client := &http.Client{Timeout: time.Second}
baseURL := "http://" + listener.Addr().String()
// A browser must also be able to complete its CORS preflight while an
// earlier connection is idle or waiting for the rest of its body.
preflight, err := http.NewRequest(http.MethodOptions, baseURL, nil)
assertNilF(t, err)
preflight.Header.Set("Origin", "https://account.example.com")
preflight.Header.Set("Access-Control-Request-Method", "POST")
resp, err := client.Do(preflight)
assertNilF(t, err, "idle connection blocked the preflight")
resp.Body.Close()
assertEqualF(t, resp.StatusCode, http.StatusOK)
assertEqualF(t, resp.Header.Get("Access-Control-Allow-Origin"), "https://account.example.com")

expected := strings.Repeat("saml", 4096) + "+/=%2B"
resp, err = client.Get(baseURL + "/?token=" + url.QueryEscape(expected))
assertNilF(t, err, "idle connection blocked the authentication callback")
body, err := io.ReadAll(resp.Body)
resp.Body.Close()
assertNilF(t, err)
assertEqualF(t, resp.StatusCode, http.StatusOK)
assertEqualF(t, string(body), fmt.Sprintf(samlSuccessHTML, "Go"))
callback := waitExternalBrowserCallback(t, result)
assertNilF(t, callback.err)
assertEqualF(t, callback.token, url.QueryEscape(expected))
assertExternalBrowserSocketClosed(t, idle)
conn, err := net.DialTimeout("tcp", listener.Addr().String(), time.Second)
if err == nil {
conn.Close()
t.Fatal("callback listener remained open after authentication")
}
})
}
}

func assertExternalBrowserSocketClosed(t *testing.T, conn net.Conn) {
t.Helper()
assertNilF(t, conn.SetReadDeadline(time.Now().Add(time.Second)))
_, err := io.Copy(io.Discard, conn)
if timeout, ok := err.(net.Error); ok && timeout.Timeout() {
t.Fatal("callback connection remained open")
}
}

func TestExternalBrowserCallbackCancellationClosesPostBody(t *testing.T) {
accountURL, err := url.Parse("https://account.example.com")
assertNilF(t, err)
listener, err := createLocalTCPListener(0)
assertNilF(t, err)
defer listener.Close()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
result := startExternalBrowserCallback(ctx, listener, accountURL)
conn, err := net.DialTimeout("tcp", listener.Addr().String(), time.Second)
assertNilF(t, err)
defer conn.Close()
_, err = io.WriteString(conn, "POST / HTTP/1.1\r\nHost: localhost\r\nOrigin: https://account.example.com\r\nContent-Length: 100\r\n\r\n{")
assertNilF(t, err)
cancel()
callback := waitExternalBrowserCallback(t, result)
assertErrIsF(t, callback.err, context.Canceled)
assertExternalBrowserSocketClosed(t, conn)
}