mirror of
https://github.com/wahyd4/loginsrv.git
synced 2026-08-09 04:46:29 +10:00
fixed caching and code exchange error handling
This commit is contained in:
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
+21
-2
@@ -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
|
||||
}
|
||||
|
||||
|
||||
+40
-1
@@ -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())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user