From 27ab7455212c3056e02f8151c0ec610f302215f0 Mon Sep 17 00:00:00 2001 From: Umputun Date: Mon, 21 May 2018 22:59:01 -0500 Subject: [PATCH] add jwt refresh --- app/main.go | 8 ++++-- app/rest/auth/auth.go | 5 +++- app/rest/auth/auth_test.go | 7 ++--- app/rest/auth/jwt.go | 52 +++++++++++++++++++++++----------- app/rest/auth/jwt_test.go | 51 ++++++++++++++++++++++++++++----- app/rest/auth/provider.go | 14 ++++----- app/rest/auth/provider_test.go | 3 +- 7 files changed, 98 insertions(+), 42 deletions(-) diff --git a/app/main.go b/app/main.go index 784426df..df459943 100644 --- a/app/main.go +++ b/app/main.go @@ -106,6 +106,8 @@ func main() { RemarkURL: strings.TrimSuffix(opts.RemarkURL, "/"), } + jwtService := auth.NewJWT(opts.SecretKey, strings.HasPrefix(opts.RemarkURL, "https://"), time.Duration(7*24*time.Hour)) + srv := api.Rest{ Version: revision, DataService: dataService, @@ -113,8 +115,9 @@ func main() { WebRoot: opts.WebRoot, ImageProxy: proxy.Image{Enabled: opts.ImageProxy, RoutePath: "/api/v1/img", RemarkURL: opts.RemarkURL}, Authenticator: auth.Authenticator{ + JWTService: jwtService, Admins: opts.Admins, - Providers: makeAuthProviders(avatarProxy), + Providers: makeAuthProviders(jwtService, avatarProxy), AvatarProxy: avatarProxy, DevPasswd: opts.DevPasswd, }, @@ -180,10 +183,11 @@ func makeDirs(dirs ...string) error { return nil } -func makeAuthProviders(avatarProxy *proxy.Avatar) (providers []auth.Provider) { +func makeAuthProviders(jwtService *auth.JWT, avatarProxy *proxy.Avatar) (providers []auth.Provider) { makeParams := func(cid, secret string) auth.Params { return auth.Params{ + JwtService: jwtService, AvatarProxy: avatarProxy, RemarkURL: opts.RemarkURL, Cid: cid, diff --git a/app/rest/auth/auth.go b/app/rest/auth/auth.go index e86ec5b9..aac4594a 100644 --- a/app/rest/auth/auth.go +++ b/app/rest/auth/auth.go @@ -14,11 +14,11 @@ import ( // Authenticator is top level auth object providing middlewares type Authenticator struct { + JWTService *JWT AvatarProxy *proxy.Avatar Admins []string Providers []Provider DevPasswd string - JWTService JWT } var devUser = store.User{ @@ -65,6 +65,9 @@ func (a *Authenticator) Auth(reqAuth bool) func(http.Handler) http.Handler { user.Admin = true break } + if _, err := a.JWTService.Refresh(w, r); err != nil { + log.Printf("[WARN] can't refresh jwt, %s", err) + } } r = rest.SetUserInfo(r, user) diff --git a/app/rest/auth/auth_test.go b/app/rest/auth/auth_test.go index a07df6b5..b6df5ac2 100644 --- a/app/rest/auth/auth_test.go +++ b/app/rest/auth/auth_test.go @@ -9,13 +9,12 @@ import ( "time" "github.com/go-chi/chi" - "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) func TestAuthJWTCookie(t *testing.T) { - a := Authenticator{DevPasswd: "123456", JWTService: JWT{secret: "xyz 12345", secureCookies: false}} + a := Authenticator{DevPasswd: "123456", JWTService: NewJWT("xyz 12345", false, time.Hour)} router := chi.NewRouter() router.With(a.Auth(true)).Get("/auth", func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(201) @@ -52,8 +51,7 @@ func TestAuthJWTCookie(t *testing.T) { } func TestAuthJWTHeader(t *testing.T) { - a := Authenticator{DevPasswd: "123456", JWTService: JWT{secret: "xyz 12345", secureCookies: false}} - + a := Authenticator{DevPasswd: "123456", JWTService: NewJWT("xyz 12345", false, time.Hour)} router := chi.NewRouter() router.With(a.Auth(true)).Get("/auth", func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(201) @@ -79,7 +77,6 @@ func TestAuthJWTHeader(t *testing.T) { assert.Equal(t, 401, resp.StatusCode, "invalid auth token") } func TestAuthRequired(t *testing.T) { - a := Authenticator{DevPasswd: "123456"} router := chi.NewRouter() router.With(a.Auth(true)).Get("/auth", func(w http.ResponseWriter, r *http.Request) { diff --git a/app/rest/auth/jwt.go b/app/rest/auth/jwt.go index 97ea05f4..73525435 100644 --- a/app/rest/auth/jwt.go +++ b/app/rest/auth/jwt.go @@ -16,15 +16,15 @@ import ( type JWT struct { secret string secureCookies bool + exp time.Duration } // CustomClaims stores user info for auth and state & from from login type CustomClaims struct { jwt.StandardClaims - User *store.User `json:"user,omitempty"` - - State string `json:"state,omitempty"` - From string `json:"from,omitempty"` + User *store.User `json:"user,omitempty"` + State string `json:"state,omitempty"` + From string `json:"from,omitempty"` } const jwtCookieName = "JWT" @@ -32,15 +32,29 @@ const jwtHeaderKey = "X-JWT" const xsrfCookieName = "XSRF-TOKEN" const xsrfHeaderKey = "X-XSRF-TOKEN" +// NewJWT makes JWT service +func NewJWT(secret string, secureCookies bool, exp time.Duration) *JWT { + res := JWT{ + secret: secret, + secureCookies: secureCookies, + exp: exp, + } + return &res +} + // Set creates jwt cookie with xsrf cookie and put it to ResponseWriter +// accepts claims and sets expiration func (j *JWT) Set(w http.ResponseWriter, claims *CustomClaims) error { + if claims.ExpiresAt == 0 { + claims.ExpiresAt = time.Now().Add(j.exp).Unix() + } token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims) tokenString, err := token.SignedString([]byte(j.secret)) if err != nil { return errors.Wrap(err, "can't sign jwt token") } - cookieExpiration := 365 * 24 * 3600 // 1year + cookieExpiration := 365 * 24 * 3600 // 1 year jwtCookie := http.Cookie{Name: jwtCookieName, Value: tokenString, HttpOnly: true, Path: "/", MaxAge: cookieExpiration, Secure: j.secureCookies} @@ -100,6 +114,22 @@ func (j *JWT) Get(r *http.Request) (*CustomClaims, error) { return claims, nil } +// Refresh gets jwt from request, checks if it will be expiring soon and create new onw +func (j *JWT) Refresh(w http.ResponseWriter, r *http.Request) (*CustomClaims, error) { + claims, err := j.Get(r) + if err != nil { + return nil, err + } + untilExp := time.Unix(claims.ExpiresAt, 0).Sub(time.Now()).Seconds() + log.Print(untilExp) + if untilExp < j.exp.Seconds()/2 { + claims.ExpiresAt = time.Now().Add(j.exp).Unix() + e := j.Set(w, claims) + return claims, e + } + return claims, nil +} + // Reset token's cookies func (j *JWT) Reset(w http.ResponseWriter) { jwtCookie := http.Cookie{Name: jwtCookieName, Value: "", HttpOnly: false, Path: "/", @@ -110,15 +140,3 @@ func (j *JWT) Reset(w http.ResponseWriter) { MaxAge: -1, Expires: time.Unix(0, 0), Secure: true} http.SetCookie(w, &xsrfCookie) } - -func (j *JWT) verify(claims CustomClaims) error { - - if time.Now().Unix() > claims.ExpiresAt { - return errors.Errorf("token exp failed %d:%d", claims.ExpiresAt, time.Now().Unix()) - } - - if time.Now().Unix() < claims.NotBefore { - return errors.Errorf("token nbf failed %d:%d", claims.NotBefore, time.Now().Unix()) - } - return nil -} diff --git a/app/rest/auth/jwt_test.go b/app/rest/auth/jwt_test.go index ca62b4ce..3cdf00ae 100644 --- a/app/rest/auth/jwt_test.go +++ b/app/rest/auth/jwt_test.go @@ -7,12 +7,11 @@ import ( "testing" "time" + jwt "github.com/dgrijalva/jwt-go" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "github.com/stretchr/testify/assert" "github.com/umputun/remark/app/store" - - jwt "github.com/dgrijalva/jwt-go" ) var testJwtValid = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJleHAiOjI3ODkxOTE4MjIsImp0aSI6InJhbmRvbSBpZCI" + @@ -23,7 +22,7 @@ var testJwtExpired = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJleHAiOjE1MjY4ODc4M "sImFkbWluIjpmYWxzZX0sInN0YXRlIjoiMTIzNDU2IiwiZnJvbSI6ImZyb20ifQ.4_dCrY9ihyfZIedz-kZwBTxmxU1a52V7IqeJrOqTzE4" func TestJWT_Set(t *testing.T) { - j := JWT{secret: "xyz 12345"} + j := NewJWT("xyz 12345", false, time.Hour) claims := &CustomClaims{ State: "123456", @@ -54,7 +53,7 @@ func TestJWT_Set(t *testing.T) { } func TestJWT_GetFromHeader(t *testing.T) { - j := JWT{secret: "xyz 12345"} + j := NewJWT("xyz 12345", false, time.Hour) req := httptest.NewRequest("GET", "/", nil) req.Header.Add(jwtHeaderKey, testJwtValid) @@ -78,7 +77,7 @@ func TestJWT_GetFromHeader(t *testing.T) { } func TestJWT_SetAndGetWithCookies(t *testing.T) { - j := JWT{secret: "xyz 12345"} + j := NewJWT("xyz 12345", false, time.Hour) claims := &CustomClaims{ State: "123456", @@ -117,7 +116,7 @@ func TestJWT_SetAndGetWithCookies(t *testing.T) { } func TestJWT_SetAndGetWithCookiesExpired(t *testing.T) { - j := JWT{secret: "xyz 12345"} + j := NewJWT("xyz 12345", false, time.Hour) claims := &CustomClaims{ State: "123456", @@ -153,3 +152,41 @@ func TestJWT_SetAndGetWithCookiesExpired(t *testing.T) { assert.NotNil(t, err) assert.True(t, strings.HasPrefix(err.Error(), "can't parse jwt: token is expired by"), err.Error()) } + +func TestJWT_Refresh(t *testing.T) { + j := NewJWT("xyz 12345", false, 2*time.Second) + + claims := &CustomClaims{ + State: "123456", + From: "from", + User: &store.User{ + ID: "id1", + Name: "name1", + }, + StandardClaims: jwt.StandardClaims{ + Id: "random id", + Issuer: "remark42", + }, + } + // set token + rr := httptest.NewRecorder() + err := j.Set(rr, claims) + assert.Nil(t, err) + cookies := rr.Result().Cookies() + require.Equal(t, 2, len(cookies)) + + req, err := http.NewRequest("GET", "http://example.com/blah", nil) + require.Nil(t, err) + req.AddCookie(cookies[0]) + req.Header.Add(xsrfHeaderKey, "random id") + + claims2, err := j.Refresh(rr, req) + require.Nil(t, err) + assert.Equal(t, claims.ExpiresAt, claims2.ExpiresAt, "no refresh yet") + + time.Sleep(1 * time.Second) + claims2, err = j.Refresh(rr, req) + assert.Nil(t, err) + assert.True(t, claims.ExpiresAt < claims2.ExpiresAt, "refreshed") + t.Log(claims.ExpiresAt, claims2.ExpiresAt) +} diff --git a/app/rest/auth/provider.go b/app/rest/auth/provider.go index 9f20cd10..b620a8f7 100644 --- a/app/rest/auth/provider.go +++ b/app/rest/auth/provider.go @@ -10,7 +10,6 @@ import ( "io/ioutil" "log" "net/http" - "strings" "time" jwt "github.com/dgrijalva/jwt-go" @@ -44,6 +43,7 @@ type Params struct { Csecret string RemarkURL string AvatarProxy *proxy.Avatar + JwtService *JWT } type userData map[string]interface{} @@ -69,8 +69,7 @@ func initProvider(p Params, provider Provider) Provider { provider.conf = &conf provider.avatarProxy = p.AvatarProxy - provider.jwtService = &JWT{secret: provider.Secret, secureCookies: strings.HasPrefix(p.RemarkURL, "https://")} - + provider.jwtService = p.JwtService return provider } @@ -89,7 +88,7 @@ func (p Provider) loginHandler(w http.ResponseWriter, r *http.Request) { // make state (random) and store in session state := p.randToken() - claims := &CustomClaims{ + claims := CustomClaims{ State: state, From: r.URL.Query().Get("from"), StandardClaims: jwt.StandardClaims{ @@ -100,7 +99,7 @@ func (p Provider) loginHandler(w http.ResponseWriter, r *http.Request) { }, } - if err := p.jwtService.Set(w, claims); err != nil { + if err := p.jwtService.Set(w, &claims); err != nil { rest.SendErrorJSON(w, r, http.StatusInternalServerError, err, "failed to set jwt") return } @@ -173,9 +172,8 @@ func (p Provider) authHandler(w http.ResponseWriter, r *http.Request) { authClaims := &CustomClaims{ User: &u, StandardClaims: jwt.StandardClaims{ - Issuer: "remark42", - Id: p.randToken(), - ExpiresAt: time.Now().Add(7 * 24 * time.Hour).Unix(), + Issuer: "remark42", + Id: p.randToken(), }, } diff --git a/app/rest/auth/provider_test.go b/app/rest/auth/provider_test.go index 31e5665c..e6f6dcd6 100644 --- a/app/rest/auth/provider_test.go +++ b/app/rest/auth/provider_test.go @@ -95,10 +95,9 @@ func mockProvider(t *testing.T, loginPort, authPort int) (provider Provider, ts } return userInfo }, - jwtService: &JWT{secret: "12345", secureCookies: false}, } - provider = initProvider(Params{Cid: "cid", Csecret: "csecret"}, provider) + provider = initProvider(Params{Cid: "cid", Csecret: "csecret", JwtService: NewJWT("12345", false, time.Hour)}, provider) ts = &http.Server{Addr: fmt.Sprintf(":%d", loginPort), Handler: provider.Routes()}