From 4e4a71d22e46e0921dbc13217bd459cd563eb690 Mon Sep 17 00:00:00 2001 From: Sebastian Mancke Date: Sat, 29 Apr 2017 16:39:58 +0200 Subject: [PATCH] fixed caching and code exchange error handling --- login/handler_test.go | 3 +++ login/login_form.go | 1 + oauth2/manager_test.go | 39 +++++++++++++++++++++++++++++++++++++++ oauth2/oauth.go | 23 +++++++++++++++++++++-- oauth2/oauth_test.go | 41 ++++++++++++++++++++++++++++++++++++++++- 5 files changed, 104 insertions(+), 3 deletions(-) diff --git a/login/handler_test.go b/login/handler_test.go index 14bf459..0cb86ed 100644 --- a/login/handler_test.go +++ b/login/handler_test.go @@ -71,6 +71,7 @@ func TestHandler_LoginForm(t *testing.T) { assert.Contains(t, recorder.Body.String(), "form") assert.Contains(t, recorder.Body.String(), `method="POST"`) assert.Contains(t, recorder.Body.String(), `action="/context/login"`) + assert.Equal(t, "no-cache, no-store, must-revalidate", recorder.Header().Get("Cache-Control")) } func TestHandler_HEAD(t *testing.T) { @@ -135,6 +136,8 @@ func TestHandler_Logout(t *testing.T) { recorder = call(req("POST", "/context/login", "logout=true", TypeForm)) assert.Equal(t, 200, recorder.Code) assert.Contains(t, recorder.Header().Get("Set-Cookie"), "jwt_token=delete; Path=/; Expires=Thu, 01 Jan 1970 00:00:00 GMT;") + + assert.Equal(t, "no-cache, no-store, must-revalidate", recorder.Header().Get("Cache-Control")) } func TestHandler_LoginError(t *testing.T) { diff --git a/login/login_form.go b/login/login_form.go index c5513ec..f1a42cf 100644 --- a/login/login_form.go +++ b/login/login_form.go @@ -114,6 +114,7 @@ func writeLoginForm(w http.ResponseWriter, params loginFormData) { return } + w.Header().Set("Cache-Control", "no-cache, no-store, must-revalidate") w.Header().Set("Content-Type", contentTypeHtml) if params.Error { w.WriteHeader(500) diff --git a/oauth2/manager_test.go b/oauth2/manager_test.go index f551b9c..a2fbbe9 100644 --- a/oauth2/manager_test.go +++ b/oauth2/manager_test.go @@ -2,6 +2,7 @@ package oauth2 import ( "crypto/tls" + "errors" "github.com/stretchr/testify/assert" "net/http" "net/http/httptest" @@ -85,6 +86,44 @@ func Test_Manager_Positive_Flow(t *testing.T) { assert.True(t, getUserInfoCalled) } +func Test_Manager_NoAauthOnWrongCode(t *testing.T) { + var authenticateCalled, getUserInfoCalled bool + + exampleProvider := Provider{ + Name: "example", + AuthURL: "https://example.com/login/oauth/authorize", + TokenURL: "https://example.com/login/oauth/access_token", + GetUserInfo: func(token TokenInfo) (map[string]string, error) { + getUserInfoCalled = true + return map[string]string{}, nil + }, + } + RegisterProvider(exampleProvider) + defer UnRegisterProvider(exampleProvider.Name) + + m := NewManager() + m.AddConfig(exampleProvider.Name, map[string]string{ + "client_id": "foo", + "client_secret": "bar", + }) + + m.authenticate = func(cfg Config, r *http.Request) (TokenInfo, error) { + authenticateCalled = true + return TokenInfo{}, errors.New("code not valid") + } + + // callback + r, _ := http.NewRequest("GET", "http://example.com/login/"+exampleProvider.Name+"?code=xyz", nil) + + startedFlow, authenticated, userInfo, err := m.Handle(httptest.NewRecorder(), r) + assert.EqualError(t, err, "code not valid") + assert.False(t, startedFlow) + assert.False(t, authenticated) + assert.Equal(t, UserInfo{}, userInfo) + assert.True(t, authenticateCalled) + assert.False(t, getUserInfoCalled) +} + func Test_Manager_getConfig_ErrorCase(t *testing.T) { r, _ := http.NewRequest("GET", "http://example.com/login", nil) diff --git a/oauth2/oauth.go b/oauth2/oauth.go index ba8c000..72393bb 100644 --- a/oauth2/oauth.go +++ b/oauth2/oauth.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "fmt" + "io/ioutil" "math/rand" "net/http" "net/url" @@ -56,6 +57,11 @@ type TokenInfo struct { Scope string `json:"scope,omitempty"` } +// JsonError represents an oauth error response in json form. +type JsonError struct { + Error string `json:"error"` +} + const stateCookieName = "oauthState" const defaultTimeout = 5 * time.Second @@ -122,13 +128,26 @@ func getAccessToken(cfg Config, state, code string) (TokenInfo, error) { if resp.StatusCode != 200 { return TokenInfo{}, fmt.Errorf("error: expected http status 200 on token exchange, but got %v", resp.StatusCode) } + body, err := ioutil.ReadAll(resp.Body) + if err != nil { + return TokenInfo{}, fmt.Errorf("error reading token exchange response: %q", err) + } + + jsonError := JsonError{} + json.Unmarshal(body, &jsonError) + if jsonError.Error != "" { + return TokenInfo{}, fmt.Errorf("error: got %q on token exchange", jsonError.Error) + } - decoder := json.NewDecoder(resp.Body) tokenInfo := TokenInfo{} - err = decoder.Decode(&tokenInfo) + err = json.Unmarshal(body, &tokenInfo) if err != nil { return TokenInfo{}, fmt.Errorf("error on parsing oauth token: %v", err) } + + if tokenInfo.AccessToken == "" { + return TokenInfo{}, fmt.Errorf("error: no access_token on token exchange") + } return tokenInfo, nil } diff --git a/oauth2/oauth_test.go b/oauth2/oauth_test.go index 7d863b4..ecd6be6 100644 --- a/oauth2/oauth_test.go +++ b/oauth2/oauth_test.go @@ -72,6 +72,45 @@ func Test_Authenticate(t *testing.T) { assert.Equal(t, "bearer", tokenInfo.TokenType) } +func Test_Authenticate_CodeExchangeError(t *testing.T) { + var testReturnCode int + testResponseJson := `{"error":"bad_verification_code","error_description":"The code passed is incorrect or expired.","error_uri":"https://developer.github.com/v3/oauth/#bad-verification-code"}` + // mock a server for token exchange + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(testReturnCode) + w.Header().Set("Content-Type", "application/json") + w.Write([]byte(testResponseJson)) + })) + defer server.Close() + + testConfigCopy := testConfig + testConfigCopy.TokenURL = server.URL + + request, _ := http.NewRequest("GET", testConfig.RedirectURI, nil) + request.Header.Set("Cookie", "oauthState=theState") + request.URL, _ = url.Parse("http://localhost/callback?code=theCode&state=theState") + + testReturnCode = 500 + tokenInfo, err := Authenticate(testConfigCopy, request) + assert.Error(t, err) + assert.EqualError(t, err, "error: expected http status 200 on token exchange, but got 500") + assert.Equal(t, "", tokenInfo.AccessToken) + + testReturnCode = 200 + tokenInfo, err = Authenticate(testConfigCopy, request) + assert.Error(t, err) + assert.EqualError(t, err, `error: got "bad_verification_code" on token exchange`) + assert.Equal(t, "", tokenInfo.AccessToken) + + testReturnCode = 200 + testResponseJson = `{"foo": "bar"}` + tokenInfo, err = Authenticate(testConfigCopy, request) + assert.Error(t, err) + assert.EqualError(t, err, `error: no access_token on token exchange`) + assert.Equal(t, "", tokenInfo.AccessToken) + +} + func Test_Authentication_ProviderError(t *testing.T) { request, _ := http.NewRequest("GET", testConfig.RedirectURI, nil) request.URL, _ = url.Parse("http://localhost/callback?error=provider_login_error") @@ -159,5 +198,5 @@ func Test_Authentication_TokenParseError(t *testing.T) { _, err := Authenticate(testConfigCopy, request) assert.Error(t, err) - assert.Equal(t, "error on parsing oauth token: unexpected EOF", err.Error()) + assert.Equal(t, "error on parsing oauth token: unexpected end of JSON input", err.Error()) }