diff --git a/login/handler.go b/login/handler.go index 400b895e..93d1704b 100644 --- a/login/handler.go +++ b/login/handler.go @@ -20,7 +20,8 @@ const contentTypeJWT = "application/jwt" const contentTypeJSON = "application/json" const contentTypePlain = "text/plain" -type userClaimsFunc func(userInfo model.UserInfo) (jwt.Claims, error) +// UserClaimsFunc returns jwt.Claims +type UserClaimsFunc func(userInfo model.UserInfo) (jwt.Claims, error) // Handler is the mail login handler. // It serves the login ressource and does the authentication against the backends or oauth provider. @@ -31,7 +32,7 @@ type Handler struct { signingMethod jwt.SigningMethod signingKey interface{} signingVerifyKey interface{} - userClaims userClaimsFunc + UserClaims UserClaimsFunc } // NewHandler creates a login handler based on the supplied configuration. @@ -70,7 +71,7 @@ func NewHandler(config *Config) (*Handler, error) { backends: backends, config: config, oauth: oauth, - userClaims: userClaims.Claims, + UserClaims: userClaims.Claims, }, nil } @@ -146,11 +147,12 @@ func (h *Handler) handleLogin(w http.ResponseWriter, r *http.Request) { if r.Method == "GET" { userInfo, valid := h.GetToken(r) + claims := h.getUserClaims(userInfo) if wantJSON(r) { if valid { w.Header().Set("Content-Type", contentTypeJSON) enc := json.NewEncoder(w) - enc.Encode(userInfo) // ignore error of encoding + enc.Encode(claims) // ignore error of encoding } else { h.respondAuthFailure(w, r) } @@ -276,9 +278,9 @@ func (h *Handler) respondAuthenticatedHTML(w http.ResponseWriter, r *http.Reques func (h *Handler) createToken(userInfo model.UserInfo) (string, error) { var claims jwt.Claims = userInfo - if h.userClaims != nil { + if h.UserClaims != nil { var err error - claims, err = h.userClaims(userInfo) + claims, err = h.UserClaims(userInfo) if err != nil { return "", err } @@ -292,6 +294,17 @@ func (h *Handler) createToken(userInfo model.UserInfo) (string, error) { return token.SignedString(key) } +func (h *Handler) getUserClaims(userInfo model.UserInfo) jwt.Claims { + var claims jwt.Claims = userInfo + if h.UserClaims != nil { + uc, err := h.UserClaims(userInfo) + if err == nil { + claims = uc + } + } + return claims +} + func (h *Handler) GetToken(r *http.Request) (userInfo model.UserInfo, valid bool) { c, err := r.Cookie(h.config.CookieName) if err != nil { diff --git a/login/handler_test.go b/login/handler_test.go index 3d314667..b91b53c7 100644 --- a/login/handler_test.go +++ b/login/handler_test.go @@ -485,6 +485,45 @@ func TestHandler_ReturnUserInfoJSON(t *testing.T) { Equal(t, input, output) } +func callWithHandler(h *Handler, req *http.Request) *httptest.ResponseRecorder { + recorder := httptest.NewRecorder() + h.ServeHTTP(recorder, req) + return recorder +} + +func TestHandler_ReturnUserInfoJSON_CustomClaims(t *testing.T) { + h := testHandler() + input := model.UserInfo{Sub: "marvin", Expiry: time.Now().Add(time.Second).Unix()} + claims := customClaims{"sub": "marvin", "exp": json.Number(strconv.FormatInt(input.Expiry, 10)), "foo": "bar"} + h.UserClaims = func(userInfo model.UserInfo) (jwt.Claims, error) { + return claims, nil + } + token, err := h.createToken(input) + NoError(t, err) + url, _ := url.Parse("/context/login") + r := &http.Request{ + Method: "GET", + URL: url, + Header: http.Header{ + "Cookie": {h.config.CookieName + "=" + token + ";"}, + "Accept": {"application/json"}, + }, + } + + recorder := callWithHandler(h, r) + Equal(t, 200, recorder.Code) + Equal(t, "application/json", recorder.Header().Get("Content-Type")) + + d := json.NewDecoder(strings.NewReader(recorder.Body.String())) + d.UseNumber() + var output customClaims + if err := d.Decode(&output); err != nil { + t.Error(err) + } + + Equal(t, claims, output) +} + func TestHandler_ReturnUserInfoJSON_InvalidToken(t *testing.T) { h := testHandler() url, _ := url.Parse("/context/login") @@ -628,7 +667,7 @@ func TestHandler_getToken_InvalidNoToken(t *testing.T) { func TestHandler_getToken_WithUserClaims(t *testing.T) { h := testHandler() input := model.UserInfo{Sub: "marvin", Expiry: time.Now().Add(time.Second).Unix()} - h.userClaims = func(userInfo model.UserInfo) (jwt.Claims, error) { + h.UserClaims = func(userInfo model.UserInfo) (jwt.Claims, error) { return customClaims{"sub": "Zappod", "origin": "fake", "exp": userInfo.Expiry}, nil } token, err := h.createToken(input)