add jwt refresh
This commit is contained in:
+6
-2
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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) {
|
||||
|
||||
+35
-17
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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(),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -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()}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user