Files
silk/server-go/internal/middleware/state.go
T

105 lines
2.7 KiB
Go

package middleware
import (
"context"
"fmt"
"log/slog"
"time"
"github.com/redis/go-redis/v9"
)
const statePrefix = "silk:auth:"
// StateStore 跨实例认证状态存储。
type StateStore interface {
RevokeToken(ctx context.Context, id string, exp time.Time) error
IsTokenRevoked(ctx context.Context, id string) (bool, error)
CheckLoginLock(ctx context.Context, key string) (bool, time.Duration, error)
RecordLoginFailure(ctx context.Context, key string) error
RecordLoginSuccess(ctx context.Context, key string) error
}
var (
authState StateStore
appEnv string
)
// InitState 设置认证状态存储;store 为 nil 时认证相关接口保守失败。
func InitState(store StateStore, env string) {
authState = store
appEnv = env
}
// RedisState Redis 实现。
type RedisState struct {
rdb *redis.Client
prefix string
}
// NewRedisState 创建 Redis 状态存储。
func NewRedisState(rdb *redis.Client) *RedisState {
return &RedisState{rdb: rdb, prefix: statePrefix}
}
func (s *RedisState) RevokeToken(ctx context.Context, id string, exp time.Time) error {
ttl := time.Until(exp)
if ttl <= 0 {
return nil
}
return s.rdb.Set(ctx, s.prefix+"revoked:"+id, "1", ttl).Err()
}
func (s *RedisState) IsTokenRevoked(ctx context.Context, id string) (bool, error) {
count, err := s.rdb.Exists(ctx, s.prefix+"revoked:"+id).Result()
if err != nil {
return false, err
}
return count > 0, nil
}
func (s *RedisState) CheckLoginLock(ctx context.Context, key string) (bool, time.Duration, error) {
lockKey := s.prefix + "login-lock:" + key
if _, err := s.rdb.Get(ctx, lockKey).Result(); err == redis.Nil {
return false, 0, nil
} else if err != nil {
return false, 0, err
}
ttl, err := s.rdb.TTL(ctx, lockKey).Result()
if err != nil {
return false, 0, err
}
return true, ttl, nil
}
var loginFailureScript = redis.NewScript(`
local count = redis.call('INCR', KEYS[1])
redis.call('EXPIRE', KEYS[1], ARGV[1])
if tonumber(count) >= tonumber(ARGV[2]) then
redis.call('SET', KEYS[2], '1', 'PX', ARGV[3])
end
return count
`)
func (s *RedisState) RecordLoginFailure(ctx context.Context, key string) error {
return loginFailureScript.Run(ctx, s.rdb,
[]string{s.prefix + "login-failures:" + key, s.prefix + "login-lock:" + key},
int(failureWindow.Seconds()), maxFailures, int(lockDuration.Milliseconds()),
).Err()
}
func (s *RedisState) RecordLoginSuccess(ctx context.Context, key string) error {
pipe := s.rdb.Pipeline()
pipe.Del(ctx, s.prefix+"login-failures:"+key, s.prefix+"login-lock:"+key)
_, err := pipe.Exec(ctx)
return err
}
func stateUnavailable(operation string) error {
msg := "Redis 状态服务不可用,无法" + operation
if appEnv == "production" {
slog.Error(msg)
}
return fmt.Errorf("%s", msg)
}