diff --git a/changes/issue-1969-redis-config b/changes/issue-1969-redis-config new file mode 100644 index 0000000000..dbbbcf214e --- /dev/null +++ b/changes/issue-1969-redis-config @@ -0,0 +1,2 @@ +* Add redis configuration option to retry failed connections. +* Add redis configuration option to follow cluster redirections. diff --git a/cmd/fleet/serve.go b/cmd/fleet/serve.go index 13dbf29b6c..d95beda2f3 100644 --- a/cmd/fleet/serve.go +++ b/cmd/fleet/serve.go @@ -200,21 +200,21 @@ the way that the Fleet server works. } } - redisPool, err := redis.NewRedisPool( - config.Redis.Address, - config.Redis.Password, - config.Redis.Database, - config.Redis.UseTLS, - config.Redis.ConnectTimeout, - config.Redis.KeepAlive, - ) + redisPool, err := redis.NewRedisPool(redis.PoolConfig{ + Server: config.Redis.Address, + Password: config.Redis.Password, + Database: config.Redis.Database, + UseTLS: config.Redis.UseTLS, + ConnTimeout: config.Redis.ConnectTimeout, + KeepAlive: config.Redis.KeepAlive, + ConnectRetryAttempts: config.Redis.ConnectRetryAttempts, + ClusterFollowRedirections: config.Redis.ClusterFollowRedirections, + }) if err != nil { initFatal(err, "initialize Redis") } resultStore := pubsub.NewRedisQueryResults(redisPool, config.Redis.DuplicateResults) liveQueryStore := live_query.NewRedisLiveQuery(redisPool) - // TODO: should that only be done when a certain "migrate" flag is set, - // to prevent affecting every startup? if err := liveQueryStore.MigrateKeys(); err != nil { level.Info(logger).Log( "err", err, diff --git a/docs/2-Deploying/2-Configuration.md b/docs/2-Deploying/2-Configuration.md index 1a41211d35..b7da00f7f6 100644 --- a/docs/2-Deploying/2-Configuration.md +++ b/docs/2-Deploying/2-Configuration.md @@ -289,7 +289,7 @@ Maximum idle connections to database. This value should be equal to or less than max_idle_conns: 50 ``` -###### conn_max_lifetime +###### mysql_conn_max_lifetime Maximum amount of time, in seconds, a connection may be reused. @@ -358,7 +358,7 @@ Whether or not to duplicate Live Query results to another Redis channel named `L ###### redis_connect_timeout -Timeout for redis connection. +Timeout for redis connection. - Default value: 5s - Environment variable: `FLEET_REDIS_CONNECT_TIMEOUT` @@ -382,6 +382,38 @@ Interval between keep alive probes. keep_alive: 30s ``` +###### redis_connect_retry_attempts + +Maximum number of attempts to retry a failed connection to a redis node. Only +certain type of errors are retried, such as connection timeouts. + +- Default value: 0 (no retry) +- Environment variable: `FLEET_REDIS_CONNECT_RETRY_ATTEMPTS` +- Config file format: + + ``` + redis: + connect_retry_attempts: 2 + ``` + +###### redis_cluster_follow_redirections + +Whether or not to automatically follow redirection errors received from the +Redis server. Applies only to Redis Cluster setups, ignored in standalone +Redis. In Redis Cluster, keys can be moved around to different nodes when the +cluster is unstable and reorganizing the data. With this configuration option +set to true, those (typically short and transient) redirection errors can be +handled transparently instead of ending in an error. + +- Default value: false +- Environment variable: `FLEET_REDIS_CLUSTER_FOLLOW_REDIRECTIONS` +- Config file format: + + ``` + redis: + cluster_follow_redirections: true + ``` + ##### Server ###### server_address diff --git a/server/config/config.go b/server/config/config.go index e687581014..1d7f977c84 100644 --- a/server/config/config.go +++ b/server/config/config.go @@ -37,13 +37,15 @@ type MysqlConfig struct { // RedisConfig defines configs related to Redis type RedisConfig struct { - Address string - Password string - Database int - UseTLS bool `yaml:"use_tls"` - DuplicateResults bool `yaml:"duplicate_results"` - ConnectTimeout time.Duration `yaml:"connect_timeout"` - KeepAlive time.Duration `yaml:"keep_alive"` + Address string + Password string + Database int + UseTLS bool `yaml:"use_tls"` + DuplicateResults bool `yaml:"duplicate_results"` + ConnectTimeout time.Duration `yaml:"connect_timeout"` + KeepAlive time.Duration `yaml:"keep_alive"` + ConnectRetryAttempts int `yaml:"connect_retry_attempts"` + ClusterFollowRedirections bool `yaml:"cluster_follow_redirections"` } const ( @@ -243,6 +245,8 @@ func (man Manager) addConfigs() { man.addConfigBool("redis.duplicate_results", false, "Duplicate Live Query results to another Redis channel") man.addConfigDuration("redis.connect_timeout", 5*time.Second, "Timeout at connection time") man.addConfigDuration("redis.keep_alive", 10*time.Second, "Interval between keep alive probes") + man.addConfigInt("redis.connect_retry_attempts", 0, "Number of attempts to retry a failed connection") + man.addConfigBool("redis.cluster_follow_redirections", false, "Automatically follow Redis Cluster redirections") // Server man.addConfigString("server.address", "0.0.0.0:8080", @@ -417,13 +421,15 @@ func (man Manager) LoadConfig() FleetConfig { Mysql: loadMysqlConfig("mysql"), MysqlReadReplica: loadMysqlConfig("mysql_read_replica"), Redis: RedisConfig{ - Address: man.getConfigString("redis.address"), - Password: man.getConfigString("redis.password"), - Database: man.getConfigInt("redis.database"), - UseTLS: man.getConfigBool("redis.use_tls"), - DuplicateResults: man.getConfigBool("redis.duplicate_results"), - ConnectTimeout: man.getConfigDuration("redis.connect_timeout"), - KeepAlive: man.getConfigDuration("redis.keep_alive"), + Address: man.getConfigString("redis.address"), + Password: man.getConfigString("redis.password"), + Database: man.getConfigInt("redis.database"), + UseTLS: man.getConfigBool("redis.use_tls"), + DuplicateResults: man.getConfigBool("redis.duplicate_results"), + ConnectTimeout: man.getConfigDuration("redis.connect_timeout"), + KeepAlive: man.getConfigDuration("redis.keep_alive"), + ConnectRetryAttempts: man.getConfigInt("redis.connect_retry_attempts"), + ClusterFollowRedirections: man.getConfigBool("redis.cluster_follow_redirections"), }, Server: ServerConfig{ Address: man.getConfigString("server.address"), diff --git a/server/datastore/redis/redis.go b/server/datastore/redis/redis.go index d9046af6e8..3679b7a61a 100644 --- a/server/datastore/redis/redis.go +++ b/server/datastore/redis/redis.go @@ -1,9 +1,11 @@ package redis import ( + "net" "strings" "time" + "github.com/cenkalti/backoff/v4" "github.com/fleetdm/fleet/v4/server/fleet" "github.com/gomodule/redigo/redis" "github.com/mna/redisc" @@ -17,29 +19,64 @@ type standalonePool struct { addr string } +func (p *standalonePool) ConfigureDoer(conn redis.Conn) redis.Conn { + return conn +} + func (p *standalonePool) Stats() map[string]redis.PoolStats { return map[string]redis.PoolStats{ p.addr: p.Pool.Stats(), } } +type clusterPool struct { + *redisc.Cluster + followRedirs bool +} + +// ConfigureDoer configures conn to follow redirections if the redis +// configuration requested it. If the conn is already in error, or +// if it is not a redisc cluster connection, it is returned unaltered. +func (p *clusterPool) ConfigureDoer(conn redis.Conn) redis.Conn { + if err := conn.Err(); err == nil && p.followRedirs { + rc, err := redisc.RetryConn(conn, 3, 300*time.Millisecond) + if err == nil { + return rc + } + } + return conn +} + +// PoolConfig holds the redis pool configuration options. +type PoolConfig struct { + Server string + Password string + Database int + UseTLS bool + ConnTimeout time.Duration + KeepAlive time.Duration + ConnectRetryAttempts int + ClusterFollowRedirections bool + + // allows for testing dial retries and other dial-related scenarios + testRedisDialFunc func(net, addr string, opts ...redis.DialOption) (redis.Conn, error) +} + // NewRedisPool creates a Redis connection pool using the provided server // address, password and database. -func NewRedisPool( - server, password string, database int, useTLS bool, connTimeout, keepAlive time.Duration, -) (fleet.RedisPool, error) { - cluster := newCluster(server, password, database, useTLS, connTimeout, keepAlive) +func NewRedisPool(config PoolConfig) (fleet.RedisPool, error) { + cluster := newCluster(config) if err := cluster.Refresh(); err != nil { if isClusterDisabled(err) || isClusterCommandUnknown(err) { // not a Redis Cluster setup, use a standalone Redis pool - pool, _ := cluster.CreatePool(server) + pool, _ := cluster.CreatePool(config.Server) cluster.Close() - return &standalonePool{pool, server}, nil + return &standalonePool{pool, config.Server}, nil } return nil, errors.Wrap(err, "refresh cluster") } - return cluster, nil + return &clusterPool{cluster, config.ClusterFollowRedirections}, nil } // SplitRedisKeysBySlot takes a list of redis keys and groups them by hash slot @@ -49,7 +86,7 @@ func NewRedisPool( // simply returns all keys in the same group (i.e. the top-level slice has a // length of 1). func SplitRedisKeysBySlot(pool fleet.RedisPool, keys ...string) [][]string { - if _, isCluster := pool.(*redisc.Cluster); isCluster { + if _, isCluster := pool.(*clusterPool); isCluster { return redisc.SplitBySlot(keys...) } return [][]string{keys} @@ -61,7 +98,7 @@ func SplitRedisKeysBySlot(pool fleet.RedisPool, keys ...string) [][]string { // of nodes stops and EachRedisNode returns that error. For standalone redis, // fn is called only once. func EachRedisNode(pool fleet.RedisPool, fn func(conn redis.Conn) error) error { - if cluster, isCluster := pool.(*redisc.Cluster); isCluster { + if cluster, isCluster := pool.(*clusterPool); isCluster { return cluster.EachNode(false, func(_ string, conn redis.Conn) error { return fn(conn) }) @@ -72,35 +109,64 @@ func EachRedisNode(pool fleet.RedisPool, fn func(conn redis.Conn) error) error { return fn(conn) } -func newCluster(server, password string, database int, useTLS bool, connTimeout, keepAlive time.Duration) *redisc.Cluster { +func newCluster(config PoolConfig) *redisc.Cluster { + opts := []redis.DialOption{ + redis.DialDatabase(config.Database), + redis.DialUseTLS(config.UseTLS), + redis.DialConnectTimeout(config.ConnTimeout), + redis.DialKeepAlive(config.KeepAlive), + // Read/Write timeouts not set here because we may see results + // only rarely on the pub/sub channel. + } + if config.Password != "" { + opts = append(opts, redis.DialPassword(config.Password)) + } + + dialFn := redis.Dial + if config.testRedisDialFunc != nil { + dialFn = config.testRedisDialFunc + } + return &redisc.Cluster{ - StartupNodes: []string{server}, - CreatePool: func(server string, opts ...redis.DialOption) (*redis.Pool, error) { + StartupNodes: []string{config.Server}, + CreatePool: func(server string, _ ...redis.DialOption) (*redis.Pool, error) { return &redis.Pool{ MaxIdle: 3, IdleTimeout: 240 * time.Second, + Dial: func() (redis.Conn, error) { - c, err := redis.Dial( - "tcp", - server, - redis.DialDatabase(database), - redis.DialUseTLS(useTLS), - redis.DialConnectTimeout(connTimeout), - redis.DialKeepAlive(keepAlive), - // Read/Write timeouts not set here because we may see results - // only rarely on the pub/sub channel. - ) - if err != nil { - return nil, err + var conn redis.Conn + op := func() error { + c, err := dialFn("tcp", server, opts...) + + var netErr net.Error + if errors.As(err, &netErr) { + if netErr.Temporary() || netErr.Timeout() { + // retryable error + return err + } + } + if err != nil { + // at this point, this is a non-retryable error + return backoff.Permanent(err) + } + + // success, store the connection to use + conn = c + return nil } - if password != "" { - if _, err := c.Do("AUTH", password); err != nil { - c.Close() + + if config.ConnectRetryAttempts > 0 { + boff := backoff.WithMaxRetries(backoff.NewExponentialBackOff(), uint64(config.ConnectRetryAttempts)) + if err := backoff.Retry(op, boff); err != nil { return nil, err } + } else if err := op(); err != nil { + return nil, err } - return c, err + return conn, nil }, + TestOnBorrow: func(c redis.Conn, t time.Time) error { if time.Since(t) < time.Minute { return nil diff --git a/server/datastore/redis/redis_test.go b/server/datastore/redis/redis_test.go index ad3da7d6b5..30b55a3c16 100644 --- a/server/datastore/redis/redis_test.go +++ b/server/datastore/redis/redis_test.go @@ -2,15 +2,171 @@ package redis import ( "fmt" + "io" + "runtime" "testing" "time" "github.com/fleetdm/fleet/v4/server/fleet" "github.com/gomodule/redigo/redis" "github.com/mna/redisc" + "github.com/pkg/errors" "github.com/stretchr/testify/require" ) +type netError struct { + error + timeout bool + temporary bool + allowedCalls int // once this reaches 0, mockDial does not return an error + countCalls int +} + +func (t *netError) Timeout() bool { return t.timeout } +func (t *netError) Temporary() bool { return t.temporary } + +var errFromConn = errors.New("SUCCESS") + +type redisConn struct{} + +func (redisConn) Close() error { return errFromConn } +func (redisConn) Err() error { return errFromConn } +func (redisConn) Do(_ string, _ ...interface{}) (interface{}, error) { return nil, errFromConn } +func (redisConn) Send(_ string, _ ...interface{}) error { return errFromConn } +func (redisConn) Flush() error { return errFromConn } +func (redisConn) Receive() (interface{}, error) { return nil, errFromConn } + +func TestConnectRetry(t *testing.T) { + mockDial := func(err error) func(net, addr string, opts ...redis.DialOption) (redis.Conn, error) { + return func(net, addr string, opts ...redis.DialOption) (redis.Conn, error) { + var ne *netError + if errors.As(err, &ne) { + ne.countCalls++ + if ne.allowedCalls <= 0 { + return redisConn{}, nil + } + ne.allowedCalls-- + } + return nil, err + } + } + + cases := []struct { + err error + retries int + wantCalls int + min, max time.Duration + }{ + // the min-max time intervals are based on the backoff default configuration as + // used in the Dial func of the redis pool. It starts with 500ms interval, + // multiplies by 1.5 on each attempt, and has a randomization of 0.5 that must + // be accounted for. Example ranges of intervals are given at + // https://github.com/fleetdm/fleet/pull/1962#issue-729635664 + // and were used to calculate the (approximate) expected range. + { + io.EOF, 0, 1, 0, 100 * time.Millisecond, + }, // non-retryable, no retry configured + { + &netError{error: io.EOF, timeout: true, allowedCalls: 10}, 0, 1, 0, 100 * time.Millisecond, + }, // retryable, but no retry configured + { + io.EOF, 3, 1, 0, 100 * time.Millisecond, + }, // non-retryable, retry configured + { + &netError{error: io.EOF, timeout: true, allowedCalls: 10}, 2, 3, 625 * time.Millisecond, 3500 * time.Millisecond, + }, // retryable, retry configured + { + &netError{error: io.EOF, temporary: true, allowedCalls: 10}, 2, 3, 625 * time.Millisecond, 3500 * time.Millisecond, + }, // retryable, retry configured + { + &netError{error: io.EOF, allowedCalls: 10}, 2, 1, 0, 100 * time.Millisecond, + }, // net error, but non-retryable + { + &netError{error: io.EOF, timeout: true, allowedCalls: 1}, 10, 2, 250 * time.Millisecond, 750 * time.Millisecond, + }, // retryable, but succeeded after one retry + } + for _, c := range cases { + t.Run(c.err.Error(), func(t *testing.T) { + start := time.Now() + _, err := NewRedisPool(PoolConfig{ + Server: "127.0.0.1:12345", + ConnectRetryAttempts: c.retries, + testRedisDialFunc: mockDial(c.err), + }) + diff := time.Since(start) + require.GreaterOrEqual(t, diff, c.min) + require.LessOrEqual(t, diff, c.max) + require.Error(t, err) + + wantErr := io.EOF + var ne *netError + if errors.As(c.err, &ne) { + require.Equal(t, c.wantCalls, ne.countCalls) + if ne.allowedCalls == 0 { + wantErr = errFromConn + } + } else { + require.Equal(t, c.wantCalls, 1) + } + + // the error is returned as part of the cluster.Refresh error, hence the + // check with Contains. + require.Contains(t, err.Error(), wantErr.Error()) + }) + } +} + +func TestRedisPoolConfigureDoer(t *testing.T) { + const prefix = "TestRedisPoolConfigureDoer:" + + t.Run("standalone", func(t *testing.T) { + pool, teardown := setupRedisForTest(t, false, false) + defer teardown() + + c1 := pool.Get() + defer c1.Close() + c2 := pool.ConfigureDoer(pool.Get()) + defer c2.Close() + + // both conns work equally well, get nil because keys do not exist, + // but no redirection error (this is standalone redis). + _, err := redis.String(c1.Do("GET", prefix+"{a}")) + require.Equal(t, redis.ErrNil, err) + _, err = redis.String(c1.Do("GET", prefix+"{b}")) + require.Equal(t, redis.ErrNil, err) + + _, err = redis.String(c2.Do("GET", prefix+"{a}")) + require.Equal(t, redis.ErrNil, err) + _, err = redis.String(c2.Do("GET", prefix+"{b}")) + require.Equal(t, redis.ErrNil, err) + }) + + t.Run("cluster", func(t *testing.T) { + pool, teardown := setupRedisForTest(t, true, true) + defer teardown() + + c1 := pool.Get() + defer c1.Close() + c2 := pool.ConfigureDoer(pool.Get()) + defer c2.Close() + + // unconfigured conn gets MOVED error on the second key + // (it is bound to {a}, {b} is on a different node) + _, err := redis.String(c1.Do("GET", prefix+"{a}")) + require.Equal(t, redis.ErrNil, err) + _, err = redis.String(c1.Do("GET", prefix+"{b}")) + rerr := redisc.ParseRedir(err) + require.Error(t, rerr) + require.Equal(t, "MOVED", rerr.Type) + + // configured conn gets the nil value, it redirected automatically + _, err = redis.String(c2.Do("GET", prefix+"{a}")) + require.Equal(t, redis.ErrNil, err) + _, err = redis.String(c2.Do("GET", prefix+"{b}")) + require.Equal(t, redis.ErrNil, err) + }) +} + func TestEachRedisNode(t *testing.T) { const prefix = "TestEachRedisNode:" @@ -49,19 +205,23 @@ func TestEachRedisNode(t *testing.T) { } t.Run("standalone", func(t *testing.T) { - pool, teardown := setupRedisForTest(t, false) + pool, teardown := setupRedisForTest(t, false, false) defer teardown() runTest(t, pool) }) t.Run("cluster", func(t *testing.T) { - pool, teardown := setupRedisForTest(t, true) + pool, teardown := setupRedisForTest(t, true, false) defer teardown() runTest(t, pool) }) } -func setupRedisForTest(t *testing.T, cluster bool) (pool fleet.RedisPool, teardown func()) { +func setupRedisForTest(t *testing.T, cluster, redir bool) (pool fleet.RedisPool, teardown func()) { + if cluster && (runtime.GOOS == "darwin" || runtime.GOOS == "windows") { + t.Skipf("docker networking limitations prevent running redis cluster tests on %s", runtime.GOOS) + } + var ( addr = "127.0.0.1:" password = "" @@ -74,7 +234,15 @@ func setupRedisForTest(t *testing.T, cluster bool) (pool fleet.RedisPool, teardo } addr += port - pool, err := NewRedisPool(addr, password, database, useTLS, 5*time.Second, 10*time.Second) + pool, err := NewRedisPool(PoolConfig{ + Server: addr, + Password: password, + Database: database, + UseTLS: useTLS, + ConnTimeout: 5 * time.Second, + KeepAlive: 10 * time.Second, + ClusterFollowRedirections: redir, + }) require.NoError(t, err) conn := pool.Get() diff --git a/server/fleet/redis_pool.go b/server/fleet/redis_pool.go index 18fa088b77..e1d86daf2c 100644 --- a/server/fleet/redis_pool.go +++ b/server/fleet/redis_pool.go @@ -5,7 +5,17 @@ import "github.com/gomodule/redigo/redis" // RedisPool is the common interface for redigo's Pool for standalone Redis // and redisc's Cluster for Redis Cluster. type RedisPool interface { + // Get returns a redis connection. It must always be closed after use. Get() redis.Conn + + // Close closes the redis connection. Close() error + + // Stats returns a map of redis pool statistics for each server address. Stats() map[string]redis.PoolStats + + // ConfigureDoer returns a redis connection that is properly configured + // to execute Do commands. This should only be called when the actions + // to execute are all done with conn.Do. + ConfigureDoer(redis.Conn) redis.Conn } diff --git a/server/live_query/redis_live_query.go b/server/live_query/redis_live_query.go index b81b88b3fb..a54e480b37 100644 --- a/server/live_query/redis_live_query.go +++ b/server/live_query/redis_live_query.go @@ -194,7 +194,7 @@ func (r *redisLiveQuery) RunQuery(name, sql string, hostIDs []uint) error { } func (r *redisLiveQuery) StopQuery(name string) error { - conn := r.pool.Get() + conn := r.pool.ConfigureDoer(r.pool.Get()) defer conn.Close() targetKey, sqlKey := generateKeys(name) @@ -279,7 +279,7 @@ func (r *redisLiveQuery) collectBatchQueriesForHost(hostID uint, queryKeys []str } func (r *redisLiveQuery) QueryCompletedByHost(name string, hostID uint) error { - conn := r.pool.Get() + conn := r.pool.ConfigureDoer(r.pool.Get()) defer conn.Close() targetKey, _ := generateKeys(name) diff --git a/server/live_query/redis_live_query_test.go b/server/live_query/redis_live_query_test.go index bdabe609f2..5a13411b32 100644 --- a/server/live_query/redis_live_query_test.go +++ b/server/live_query/redis_live_query_test.go @@ -1,6 +1,7 @@ package live_query import ( + "runtime" "testing" "time" @@ -99,6 +100,10 @@ func TestMigrateKeys(t *testing.T) { } func setupRedisLiveQuery(t *testing.T, cluster bool) (store *redisLiveQuery, teardown func()) { + if cluster && (runtime.GOOS == "darwin" || runtime.GOOS == "windows") { + t.Skipf("docker networking limitations prevent running redis cluster tests on %s", runtime.GOOS) + } + var ( addr = "127.0.0.1:" password = "" @@ -111,7 +116,14 @@ func setupRedisLiveQuery(t *testing.T, cluster bool) (store *redisLiveQuery, tea } addr += port - pool, err := redis.NewRedisPool(addr, password, database, useTLS, 5*time.Second, 10*time.Second) + pool, err := redis.NewRedisPool(redis.PoolConfig{ + Server: addr, + Password: password, + Database: database, + UseTLS: useTLS, + ConnTimeout: 5 * time.Second, + KeepAlive: 10 * time.Second, + }) require.NoError(t, err) store = NewRedisLiveQuery(pool) diff --git a/server/pubsub/testing_utils.go b/server/pubsub/testing_utils.go index 2b8a02cc41..9ff5b39c1d 100644 --- a/server/pubsub/testing_utils.go +++ b/server/pubsub/testing_utils.go @@ -1,6 +1,7 @@ package pubsub import ( + "runtime" "testing" "time" @@ -10,6 +11,10 @@ import ( ) func SetupRedisForTest(t *testing.T, cluster bool) (store *redisQueryResults, teardown func()) { + if cluster && (runtime.GOOS == "darwin" || runtime.GOOS == "windows") { + t.Skipf("docker networking limitations prevent running redis cluster tests on %s", runtime.GOOS) + } + var ( addr = "127.0.0.1:" password = "" @@ -23,7 +28,14 @@ func SetupRedisForTest(t *testing.T, cluster bool) (store *redisQueryResults, te } addr += port - pool, err := redis.NewRedisPool(addr, password, database, useTLS, 5*time.Second, 10*time.Second) + pool, err := redis.NewRedisPool(redis.PoolConfig{ + Server: addr, + Password: password, + Database: database, + UseTLS: useTLS, + ConnTimeout: 5 * time.Second, + KeepAlive: 10 * time.Second, + }) require.NoError(t, err) store = NewRedisQueryResults(pool, dupResults) diff --git a/server/sso/session_store.go b/server/sso/session_store.go index b2277c5619..7b0bb36fdc 100644 --- a/server/sso/session_store.go +++ b/server/sso/session_store.go @@ -46,7 +46,7 @@ func (s *store) create(requestID, originalURL, metadata string, lifetimeSecs uin if len(requestID) < 8 { return errors.New("request id must be 8 or more characters in length") } - conn := s.pool.Get() + conn := s.pool.ConfigureDoer(s.pool.Get()) defer conn.Close() sess := Session{OriginalURL: originalURL, Metadata: metadata} var writer bytes.Buffer @@ -59,7 +59,7 @@ func (s *store) create(requestID, originalURL, metadata string, lifetimeSecs uin } func (s *store) Get(requestID string) (*Session, error) { - conn := s.pool.Get() + conn := s.pool.ConfigureDoer(s.pool.Get()) defer conn.Close() val, err := redis.String(conn.Do("GET", requestID)) if err != nil { @@ -81,7 +81,7 @@ func (s *store) Get(requestID string) (*Session, error) { var ErrSessionNotFound = errors.New("session not found") func (s *store) Expire(requestID string) error { - conn := s.pool.Get() + conn := s.pool.ConfigureDoer(s.pool.Get()) defer conn.Close() _, err := conn.Do("DEL", requestID) return err diff --git a/server/sso/session_store_test.go b/server/sso/session_store_test.go index 45b46af866..318ba12acd 100644 --- a/server/sso/session_store_test.go +++ b/server/sso/session_store_test.go @@ -2,6 +2,7 @@ package sso import ( "os" + "runtime" "testing" "time" @@ -12,6 +13,10 @@ import ( ) func newPool(t *testing.T, cluster bool) fleet.RedisPool { + if cluster && (runtime.GOOS == "darwin" || runtime.GOOS == "windows") { + t.Skipf("docker networking limitations prevent running redis cluster tests on %s", runtime.GOOS) + } + if _, ok := os.LookupEnv("REDIS_TEST"); ok { var ( addr = "127.0.0.1:" @@ -25,7 +30,14 @@ func newPool(t *testing.T, cluster bool) fleet.RedisPool { } addr += port - pool, err := redis.NewRedisPool(addr, password, database, useTLS, 5*time.Second, 10*time.Second) + pool, err := redis.NewRedisPool(redis.PoolConfig{ + Server: addr, + Password: password, + Database: database, + UseTLS: useTLS, + ConnTimeout: 5 * time.Second, + KeepAlive: 10 * time.Second, + }) require.NoError(t, err) conn := pool.Get() defer conn.Close()