parent
ad2735a1cb
commit
ee672697c2
@ -0,0 +1,54 @@
|
|||||||
|
package redis
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strconv"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/redis/go-redis/v9"
|
||||||
|
|
||||||
|
"github.com/openimsdk/tools/errs"
|
||||||
|
)
|
||||||
|
|
||||||
|
const standaloneGatewayHashKey = "STANDALONE_GATEWAY_REGISTRY"
|
||||||
|
|
||||||
|
type StandaloneGatewayRedis struct {
|
||||||
|
rdb redis.UniversalClient
|
||||||
|
validTime time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewStandaloneGatewayRedis(rdb redis.UniversalClient, validTime time.Duration) *StandaloneGatewayRedis {
|
||||||
|
return &StandaloneGatewayRedis{rdb: rdb, validTime: validTime}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandaloneGatewayRedis) RegisterGateway(ctx context.Context, addr string) error {
|
||||||
|
pipe := s.rdb.Pipeline()
|
||||||
|
pipe.HSet(ctx, standaloneGatewayHashKey, addr, strconv.FormatInt(time.Now().UnixMilli(), 10))
|
||||||
|
pipe.Expire(ctx, standaloneGatewayHashKey, s.validTime*2)
|
||||||
|
_, err := pipe.Exec(ctx)
|
||||||
|
return errs.Wrap(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandaloneGatewayRedis) UnregisterGateway(ctx context.Context, addr string) error {
|
||||||
|
return errs.Wrap(s.rdb.HDel(ctx, standaloneGatewayHashKey, addr).Err())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *StandaloneGatewayRedis) GetGatewayAddrs(ctx context.Context) ([]string, error) {
|
||||||
|
gateways, err := s.rdb.HGetAll(ctx, standaloneGatewayHashKey).Result()
|
||||||
|
if err != nil {
|
||||||
|
return nil, errs.Wrap(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
now := time.Now()
|
||||||
|
addrs := make([]string, 0, len(gateways))
|
||||||
|
for addr, registeredAt := range gateways {
|
||||||
|
registeredAtMs, err := strconv.ParseInt(registeredAt, 10, 64)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errs.WrapMsg(err, "redis gateway register time is not int64", "addr", addr, "value", registeredAt)
|
||||||
|
}
|
||||||
|
if now.Sub(time.UnixMilli(registeredAtMs)) <= s.validTime {
|
||||||
|
addrs = append(addrs, addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return addrs, nil
|
||||||
|
}
|
||||||
@ -0,0 +1,52 @@
|
|||||||
|
package redis
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strconv"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/go-redis/redismock/v9"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestStandaloneGatewayRedisRegisterGateway(t *testing.T) {
|
||||||
|
rdb, mock := redismock.NewClientMock()
|
||||||
|
cache := NewStandaloneGatewayRedis(rdb, time.Second*10)
|
||||||
|
|
||||||
|
mock.Regexp().ExpectHSet(standaloneGatewayHashKey, "127.0.0.1:10001", `^[0-9]+$`).SetVal(1)
|
||||||
|
mock.ExpectExpire(standaloneGatewayHashKey, time.Second*20).SetVal(true)
|
||||||
|
|
||||||
|
err := cache.RegisterGateway(context.Background(), "127.0.0.1:10001")
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.NoError(t, mock.ExpectationsWereMet())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStandaloneGatewayRedisUnregisterGateway(t *testing.T) {
|
||||||
|
rdb, mock := redismock.NewClientMock()
|
||||||
|
cache := NewStandaloneGatewayRedis(rdb, time.Second)
|
||||||
|
|
||||||
|
mock.ExpectHDel(standaloneGatewayHashKey, "127.0.0.1:10001").SetVal(1)
|
||||||
|
|
||||||
|
err := cache.UnregisterGateway(context.Background(), "127.0.0.1:10001")
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.NoError(t, mock.ExpectationsWereMet())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStandaloneGatewayRedisGetGatewayAddrs(t *testing.T) {
|
||||||
|
rdb, mock := redismock.NewClientMock()
|
||||||
|
cache := NewStandaloneGatewayRedis(rdb, time.Second*10)
|
||||||
|
|
||||||
|
now := time.Now()
|
||||||
|
mock.ExpectHGetAll(standaloneGatewayHashKey).SetVal(map[string]string{
|
||||||
|
"127.0.0.1:10001": strconv.FormatInt(now.Add(-time.Second).UnixMilli(), 10),
|
||||||
|
"127.0.0.1:10002": strconv.FormatInt(now.Add(-time.Second*20).UnixMilli(), 10),
|
||||||
|
"127.0.0.1:10003": strconv.FormatInt(now.Add(time.Second).UnixMilli(), 10),
|
||||||
|
})
|
||||||
|
|
||||||
|
addrs, err := cache.GetGatewayAddrs(context.Background())
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, []string{"127.0.0.1:10001", "127.0.0.1:10003"}, addrs)
|
||||||
|
assert.NoError(t, mock.ExpectationsWereMet())
|
||||||
|
}
|
||||||
Loading…
Reference in new issue