mirror of
https://github.com/oauth2-proxy/oauth2-proxy.git
synced 2026-09-30 03:31:27 +02:00
Add redis lock feature (#1063)
* Add sensible logging flag to default setup for logger
* Add Redis lock
* Fix default value flag for sensitive logging
* Split RefreshSessionIfNeeded in two methods and use Redis lock
* Small adjustments to doc and code
* Remove sensible logging
* Fix method names in ticket.go
* Revert "Fix method names in ticket.go"
This reverts commit 408ba1a1a5.
* Fix methods name in ticket.go
* Remove block in Redis client get
* Increase lock time to 1 second
* Perform retries, if session store is locked
* Reverse if condition, because it should return if session does not have to be refreshed
* Update go.sum
* Update MockStore
* Return error if loading session fails
* Fix and update tests
* Change validSession to session in docs and strings
* Change validSession to session in docs and strings
* Fix docs
* Fix wrong field name
* Fix linting
* Fix imports for linting
* Revert changes except from locking functionality
* Add lock feature on session state
* Update from master
* Remove errors package, because it is not used
* Only pass context instead of request to lock
* Use lock key
* By default use NoOpLock
* Remove debug output
* Update ticket_test.go
* Map internal error to sessions error
* Add ErrLockNotObtained
* Enable lock peek for all redis clients
* Use lock key prefix consistent
* Fix imports
* Use exists method for peek lock
* Fix imports
* Fix imports
* Fix imports
* Remove own Dockerfile
* Fix imports
* Fix tests for ticket and session store
* Fix session store test
* Update pkg/apis/sessions/interfaces.go
Co-authored-by: Joel Speed <Joel.speed@hotmail.co.uk>
* Do not wrap lock method
Co-authored-by: Joel Speed <Joel.speed@hotmail.co.uk>
* Use errors package for lock constants
* Use better naming for initLock function
* Add comments
* Add session store lock test
* Fix tests
* Fix tests
* Fix tests
* Fix tests
* Add cookies after saving session
* Add mock lock
* Fix imports for mock_lock.go
* Store mock lock for key
* Apply elapsed time on mock lock
* Check if lock is initially applied
* Reuse existing lock
* Test all lock methods
* Update CHANGELOG.md
* Use redis client methods in redis.lock for release an refresh
* Use lock key suffix instead of prefix for lock key
* Add comments for Lock interface
* Update comment for Lock interface
* Update CHANGELOG.md
* Change LockSuffix to const
* Check lock on already loaded session
* Use global var for loadedSession in lock tests
* Use lock instance for refreshing and releasing of lock
* Update possible error type for Refresh
Co-authored-by: Joel Speed <Joel.speed@hotmail.co.uk>
This commit is contained in:
co-authored by
Joel Speed
parent
67bfa4b43f
commit
f648c54d87
@@ -3,6 +3,8 @@ package persistence
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/oauth2-proxy/oauth2-proxy/v7/pkg/apis/sessions"
|
||||
)
|
||||
|
||||
// Store is used for persistent session stores (IE not Cookie)
|
||||
@@ -12,4 +14,5 @@ type Store interface {
|
||||
Save(context.Context, string, []byte, time.Duration) error
|
||||
Load(context.Context, string) ([]byte, error)
|
||||
Clear(context.Context, string) error
|
||||
Lock(key string) sessions.Lock
|
||||
}
|
||||
|
||||
@@ -60,9 +60,12 @@ func (m *Manager) Load(req *http.Request) (*sessions.SessionState, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return tckt.loadSession(func(key string) ([]byte, error) {
|
||||
return m.Store.Load(req.Context(), key)
|
||||
})
|
||||
return tckt.loadSession(
|
||||
func(key string) ([]byte, error) {
|
||||
return m.Store.Load(req.Context(), key)
|
||||
},
|
||||
m.Store.Lock,
|
||||
)
|
||||
}
|
||||
|
||||
// Clear clears any saved session information for a given ticket cookie.
|
||||
|
||||
@@ -30,6 +30,10 @@ type loadFunc func(string) ([]byte, error)
|
||||
// a string key for the target of the deletion.
|
||||
type clearFunc func(string) error
|
||||
|
||||
// initLockFunc returns a lock object for a persistent store using a
|
||||
// string key
|
||||
type initLockFunc func(string) sessions.Lock
|
||||
|
||||
// ticket is a structure representing the ticket used in server based
|
||||
// session storage. It provides a unique per session decryption secret giving
|
||||
// more security than the shared CookieSecret.
|
||||
@@ -122,7 +126,8 @@ func (t *ticket) saveSession(s *sessions.SessionState, saver saveFunc) error {
|
||||
// loadSession loads a session from the disk store via the passed loadFunc
|
||||
// using the ticket.id as the key. It then decodes the SessionState using
|
||||
// ticket.secret to make the AES-GCM cipher.
|
||||
func (t *ticket) loadSession(loader loadFunc) (*sessions.SessionState, error) {
|
||||
// finally it appends a lock implementation
|
||||
func (t *ticket) loadSession(loader loadFunc, initLock initLockFunc) (*sessions.SessionState, error) {
|
||||
ciphertext, err := loader(t.id)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to load the session state with the ticket: %v", err)
|
||||
@@ -132,7 +137,13 @@ func (t *ticket) loadSession(loader loadFunc) (*sessions.SessionState, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return sessions.DecodeSessionState(ciphertext, c, false)
|
||||
sessionState, err := sessions.DecodeSessionState(ciphertext, c, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
lock := initLock(t.id)
|
||||
sessionState.Lock = lock
|
||||
return sessionState, nil
|
||||
}
|
||||
|
||||
// clearSession uses the passed clearFunc to delete a session stored with a
|
||||
|
||||
@@ -103,10 +103,17 @@ var _ = Describe("Session Ticket Tests", func() {
|
||||
c, err := t.makeCipher()
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
ss := &sessions.SessionState{User: "foobar"}
|
||||
loadedSession, err := t.loadSession(func(k string) ([]byte, error) {
|
||||
return ss.EncodeSessionState(c, false)
|
||||
})
|
||||
ss := &sessions.SessionState{
|
||||
User: "foobar",
|
||||
Lock: &sessions.NoOpLock{},
|
||||
}
|
||||
loadedSession, err := t.loadSession(
|
||||
func(k string) ([]byte, error) {
|
||||
return ss.EncodeSessionState(c, false)
|
||||
},
|
||||
func(k string) sessions.Lock {
|
||||
return &sessions.NoOpLock{}
|
||||
})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(loadedSession).To(Equal(ss))
|
||||
})
|
||||
@@ -115,9 +122,13 @@ var _ = Describe("Session Ticket Tests", func() {
|
||||
t, err := newTicket(&options.Cookie{Name: "dummy"})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
data, err := t.loadSession(func(k string) ([]byte, error) {
|
||||
return nil, errors.New("load error")
|
||||
})
|
||||
data, err := t.loadSession(
|
||||
func(k string) ([]byte, error) {
|
||||
return nil, errors.New("load error")
|
||||
},
|
||||
func(k string) sessions.Lock {
|
||||
return &sessions.NoOpLock{}
|
||||
})
|
||||
Expect(data).To(BeNil())
|
||||
Expect(err).To(MatchError(errors.New("failed to load the session state with the ticket: load error")))
|
||||
})
|
||||
|
||||
@@ -5,11 +5,13 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/go-redis/redis/v8"
|
||||
"github.com/oauth2-proxy/oauth2-proxy/v7/pkg/apis/sessions"
|
||||
)
|
||||
|
||||
// Client is wrapper interface for redis.Client and redis.ClusterClient.
|
||||
type Client interface {
|
||||
Get(ctx context.Context, key string) ([]byte, error)
|
||||
Lock(key string) sessions.Lock
|
||||
Set(ctx context.Context, key string, value []byte, expiration time.Duration) error
|
||||
Del(ctx context.Context, key string) error
|
||||
}
|
||||
@@ -21,7 +23,9 @@ type client struct {
|
||||
}
|
||||
|
||||
func newClient(c *redis.Client) Client {
|
||||
return &client{Client: c}
|
||||
return &client{
|
||||
Client: c,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *client) Get(ctx context.Context, key string) ([]byte, error) {
|
||||
@@ -36,6 +40,10 @@ func (c *client) Del(ctx context.Context, key string) error {
|
||||
return c.Client.Del(ctx, key).Err()
|
||||
}
|
||||
|
||||
func (c *client) Lock(key string) sessions.Lock {
|
||||
return NewLock(c.Client, key)
|
||||
}
|
||||
|
||||
var _ Client = (*clusterClient)(nil)
|
||||
|
||||
type clusterClient struct {
|
||||
@@ -43,7 +51,9 @@ type clusterClient struct {
|
||||
}
|
||||
|
||||
func newClusterClient(c *redis.ClusterClient) Client {
|
||||
return &clusterClient{ClusterClient: c}
|
||||
return &clusterClient{
|
||||
ClusterClient: c,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *clusterClient) Get(ctx context.Context, key string) ([]byte, error) {
|
||||
@@ -57,3 +67,7 @@ func (c *clusterClient) Set(ctx context.Context, key string, value []byte, expir
|
||||
func (c *clusterClient) Del(ctx context.Context, key string) error {
|
||||
return c.ClusterClient.Del(ctx, key).Err()
|
||||
}
|
||||
|
||||
func (c *clusterClient) Lock(key string) sessions.Lock {
|
||||
return NewLock(c.ClusterClient, key)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
package redis
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/bsm/redislock"
|
||||
"github.com/go-redis/redis/v8"
|
||||
"github.com/oauth2-proxy/oauth2-proxy/v7/pkg/apis/sessions"
|
||||
)
|
||||
|
||||
const LockSuffix = "lock"
|
||||
|
||||
type Lock struct {
|
||||
client redis.Cmdable
|
||||
locker *redislock.Client
|
||||
lock *redislock.Lock
|
||||
key string
|
||||
}
|
||||
|
||||
// NewLock instantiate a new lock instance. This will not yet apply a lock on Redis side.
|
||||
// For that you have to call Obtain(ctx context.Context, expiration time.Duration)
|
||||
func NewLock(client redis.Cmdable, key string) sessions.Lock {
|
||||
return &Lock{
|
||||
client: client,
|
||||
locker: redislock.New(client),
|
||||
key: key,
|
||||
}
|
||||
}
|
||||
|
||||
// Obtain obtains a distributed lock on Redis for the configured key.
|
||||
func (l *Lock) Obtain(ctx context.Context, expiration time.Duration) error {
|
||||
lock, err := l.locker.Obtain(ctx, l.lockKey(), expiration, nil)
|
||||
if errors.Is(err, redislock.ErrNotObtained) {
|
||||
return sessions.ErrLockNotObtained
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
l.lock = lock
|
||||
return nil
|
||||
}
|
||||
|
||||
// Refresh refreshes an already existing lock.
|
||||
func (l *Lock) Refresh(ctx context.Context, expiration time.Duration) error {
|
||||
if l.lock == nil {
|
||||
return sessions.ErrNotLocked
|
||||
}
|
||||
err := l.lock.Refresh(ctx, expiration, nil)
|
||||
if errors.Is(err, redislock.ErrNotObtained) {
|
||||
return sessions.ErrNotLocked
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// Peek returns true, if the lock is still applied.
|
||||
func (l *Lock) Peek(ctx context.Context) (bool, error) {
|
||||
v, err := l.client.Exists(ctx, l.lockKey()).Result()
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if v == 0 {
|
||||
return false, nil
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// Release releases the lock on Redis side.
|
||||
func (l *Lock) Release(ctx context.Context) error {
|
||||
if l.lock == nil {
|
||||
return sessions.ErrNotLocked
|
||||
}
|
||||
err := l.lock.Release(ctx)
|
||||
if errors.Is(err, redislock.ErrLockNotHeld) {
|
||||
return sessions.ErrNotLocked
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (l *Lock) lockKey() string {
|
||||
return fmt.Sprintf("%s.%s", l.key, LockSuffix)
|
||||
}
|
||||
@@ -35,7 +35,7 @@ func NewRedisSessionStore(opts *options.SessionOptions, cookieOpts *options.Cook
|
||||
}
|
||||
|
||||
// Save takes a sessions.SessionState and stores the information from it
|
||||
// to redies, and adds a new persistence cookie on the HTTP response writer
|
||||
// to redis, and adds a new persistence cookie on the HTTP response writer
|
||||
func (store *SessionStore) Save(ctx context.Context, key string, value []byte, exp time.Duration) error {
|
||||
err := store.Client.Set(ctx, key, value, exp)
|
||||
if err != nil {
|
||||
@@ -64,6 +64,11 @@ func (store *SessionStore) Clear(ctx context.Context, key string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Lock creates a lock object for sessions.SessionState
|
||||
func (store *SessionStore) Lock(key string) sessions.Lock {
|
||||
return store.Client.Lock(key)
|
||||
}
|
||||
|
||||
// NewRedisClient makes a redis.Client (either standalone, sentinel aware, or
|
||||
// redis cluster)
|
||||
func NewRedisClient(opts options.RedisStoreOptions) (Client, error) {
|
||||
@@ -151,7 +156,7 @@ func buildStandaloneClient(opts options.RedisStoreOptions) (Client, error) {
|
||||
}
|
||||
|
||||
// parseRedisURLs parses a list of redis urls and returns a list
|
||||
// of addresses in the form of host:port that can be used to connnect to Redis
|
||||
// of addresses in the form of host:port that can be used to connect to Redis
|
||||
func parseRedisURLs(urls []string) ([]string, error) {
|
||||
addrs := []string{}
|
||||
for _, u := range urls {
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
package tests
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/oauth2-proxy/oauth2-proxy/v7/pkg/apis/sessions"
|
||||
)
|
||||
|
||||
type MockLock struct {
|
||||
expiration time.Duration
|
||||
elapsed time.Duration
|
||||
}
|
||||
|
||||
func (l *MockLock) Obtain(ctx context.Context, expiration time.Duration) error {
|
||||
l.expiration = expiration
|
||||
return nil
|
||||
}
|
||||
|
||||
func (l *MockLock) Peek(ctx context.Context) (bool, error) {
|
||||
if l.elapsed < l.expiration {
|
||||
return true, nil
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func (l *MockLock) Refresh(ctx context.Context, expiration time.Duration) error {
|
||||
if l.expiration <= l.elapsed {
|
||||
return sessions.ErrNotLocked
|
||||
}
|
||||
l.expiration = expiration
|
||||
l.elapsed = time.Duration(0)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (l *MockLock) Release(ctx context.Context) error {
|
||||
if l.expiration <= l.elapsed {
|
||||
return sessions.ErrNotLocked
|
||||
}
|
||||
l.expiration = time.Duration(0)
|
||||
l.elapsed = time.Duration(0)
|
||||
return nil
|
||||
}
|
||||
|
||||
// FastForward simulates the flow of time to test expirations
|
||||
func (l *MockLock) FastForward(duration time.Duration) {
|
||||
l.elapsed += duration
|
||||
}
|
||||
@@ -4,6 +4,8 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/oauth2-proxy/oauth2-proxy/v7/pkg/apis/sessions"
|
||||
)
|
||||
|
||||
// entry is a MockStore cache entry with an expiration
|
||||
@@ -15,15 +17,17 @@ type entry struct {
|
||||
// MockStore is a generic in-memory implementation of persistence.Store
|
||||
// for mocking in tests
|
||||
type MockStore struct {
|
||||
cache map[string]entry
|
||||
elapsed time.Duration
|
||||
cache map[string]entry
|
||||
lockCache map[string]*MockLock
|
||||
elapsed time.Duration
|
||||
}
|
||||
|
||||
// NewMockStore creates a MockStore
|
||||
func NewMockStore() *MockStore {
|
||||
return &MockStore{
|
||||
cache: map[string]entry{},
|
||||
elapsed: 0 * time.Second,
|
||||
cache: map[string]entry{},
|
||||
lockCache: map[string]*MockLock{},
|
||||
elapsed: 0 * time.Second,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -52,7 +56,19 @@ func (s *MockStore) Clear(_ context.Context, key string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *MockStore) Lock(key string) sessions.Lock {
|
||||
if s.lockCache[key] != nil {
|
||||
return s.lockCache[key]
|
||||
}
|
||||
lock := &MockLock{}
|
||||
s.lockCache[key] = lock
|
||||
return lock
|
||||
}
|
||||
|
||||
// FastForward simulates the flow of time to test expirations
|
||||
func (s *MockStore) FastForward(duration time.Duration) {
|
||||
for _, mockLock := range s.lockCache {
|
||||
mockLock.FastForward(duration)
|
||||
}
|
||||
s.elapsed += duration
|
||||
}
|
||||
|
||||
@@ -286,6 +286,78 @@ func PersistentSessionStoreInterfaceTests(in *testInput) {
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
Context("when lock is applied", func() {
|
||||
var loadedSession *sessionsapi.SessionState
|
||||
BeforeEach(func() {
|
||||
resp := httptest.NewRecorder()
|
||||
err := in.ss().Save(resp, in.request, in.session)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
for _, cookie := range resp.Result().Cookies() {
|
||||
in.request.AddCookie(cookie)
|
||||
}
|
||||
|
||||
loadedSession, err = in.ss().Load(in.request)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
err = loadedSession.ObtainLock(in.request.Context(), 2*time.Minute)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
isLocked, err := loadedSession.PeekLock(in.request.Context())
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(isLocked).To(BeTrue())
|
||||
})
|
||||
|
||||
Context("before lock expired", func() {
|
||||
BeforeEach(func() {
|
||||
Expect(in.persistentFastForward(time.Minute)).To(Succeed())
|
||||
})
|
||||
|
||||
It("peek returns true on loaded session lock", func() {
|
||||
l := *loadedSession
|
||||
isLocked, err := l.PeekLock(in.request.Context())
|
||||
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(isLocked).To(BeTrue())
|
||||
})
|
||||
|
||||
It("lock can be released", func() {
|
||||
l := *loadedSession
|
||||
|
||||
err := l.ReleaseLock(in.request.Context())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
isLocked, err := l.PeekLock(in.request.Context())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(isLocked).To(BeFalse())
|
||||
})
|
||||
|
||||
It("lock is refreshed", func() {
|
||||
l := *loadedSession
|
||||
err := l.RefreshLock(in.request.Context(), 3*time.Minute)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
Expect(in.persistentFastForward(2 * time.Minute)).To(Succeed())
|
||||
|
||||
isLocked, err := l.PeekLock(in.request.Context())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(isLocked).To(BeTrue())
|
||||
})
|
||||
})
|
||||
|
||||
Context("after lock expired", func() {
|
||||
BeforeEach(func() {
|
||||
Expect(in.persistentFastForward(3 * time.Minute)).To(Succeed())
|
||||
})
|
||||
|
||||
It("peek returns false on loaded session lock", func() {
|
||||
l := *loadedSession
|
||||
isLocked, err := l.PeekLock(in.request.Context())
|
||||
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(isLocked).To(BeFalse())
|
||||
})
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
func SessionStoreInterfaceTests(in *testInput) {
|
||||
@@ -411,9 +483,11 @@ func LoadSessionTests(in *testInput) {
|
||||
l := *loadedSession
|
||||
l.CreatedAt = nil
|
||||
l.ExpiresOn = nil
|
||||
l.Lock = &sessionsapi.NoOpLock{}
|
||||
s := *in.session
|
||||
s.CreatedAt = nil
|
||||
s.ExpiresOn = nil
|
||||
s.Lock = &sessionsapi.NoOpLock{}
|
||||
Expect(l).To(Equal(s))
|
||||
|
||||
// Compare time.Time separately
|
||||
|
||||
Reference in New Issue
Block a user