package redisprovider import ( "context" "crypto/rand" "encoding/hex" "errors" "math" "reflect" "regexp" "strings" "sync" "time" "github.com/redis/go-redis/v9" controllerProvider "proxy-pool/internal/controller/provider" ) var namespacePattern = regexp.MustCompile(`^[A-Za-z0-9._-]+$`) const maximumLuaInteger = int64(1<<53 - 1) type Options struct { Namespace string HolderID string LeaseTTL time.Duration RenewEvery time.Duration RetryInterval time.Duration PermitGrace time.Duration } type Adapter struct { client redis.Scripter options Options keys keyBuilder } var _ controllerProvider.Coordinator = (*Adapter)(nil) func New(client redis.Scripter, options Options) (*Adapter, error) { options.Namespace = strings.TrimSpace(options.Namespace) options.HolderID = strings.TrimSpace(options.HolderID) if nilInterface(client) || !namespacePattern.MatchString(options.Namespace) || options.HolderID == "" || options.LeaseTTL <= 0 || options.RenewEvery <= 0 || options.RenewEvery > options.LeaseTTL/3 || options.RetryInterval <= 0 || options.PermitGrace < 0 { return nil, controllerProvider.ErrInvalidCoordination } return &Adapter{client: client, options: options, keys: keyBuilder{namespace: options.Namespace}}, nil } func (adapter *Adapter) RunLeader( ctx context.Context, upstreamID string, limits controllerProvider.CoordinationLimits, work func(context.Context, controllerProvider.LeaderSession) error, ) error { if ctx == nil || adapter == nil || work == nil || strings.TrimSpace(upstreamID) != upstreamID || upstreamID == "" || limits.RequestInterval < 0 || limits.MaxInFlight <= 0 || limits.MaxAttemptDuration <= 0 || limits.MaxTotal < 0 || limits.MaxTotal > maximumLuaInteger || limits.MaxAttemptDuration > time.Duration(math.MaxInt64)-adapter.options.PermitGrace { return controllerProvider.ErrInvalidCoordination } if err := ctx.Err(); err != nil { return err } keys, err := adapter.keys.forUpstream(upstreamID) if err != nil { return err } token, err := randomToken() if err != nil { return errors.Join(controllerProvider.ErrCoordinationUnavailable, err) } for ctx.Err() == nil { generationCandidate, tokenErr := randomToken() if tokenErr != nil { return errors.Join(controllerProvider.ErrCoordinationUnavailable, tokenErr) } reply, acquireErr := runScript(ctx, adapter.client, keys, "acquire_leader", generationCandidate, adapter.options.HolderID, token, durationMillis(adapter.options.LeaseTTL), ) if acquireErr != nil { if ctx.Err() != nil { return ctx.Err() } if err := wait(ctx, adapter.options.RetryInterval); err != nil { return err } continue } switch reply.Status { case "busy": delay := adapter.options.RetryInterval if reply.WaitMS > 0 && time.Duration(reply.WaitMS)*time.Millisecond < delay { delay = time.Duration(reply.WaitMS) * time.Millisecond } if err := wait(ctx, delay); err != nil { return err } continue case "ok": if reply.Generation == "" || reply.Epoch == 0 { if err := wait(ctx, adapter.options.RetryInterval); err != nil { return err } continue } default: if err := wait(ctx, adapter.options.RetryInterval); err != nil { return err } continue } session := &leaderSession{ adapter: adapter, keys: keys, upstreamID: upstreamID, limits: limits, generation: reply.Generation, holderID: adapter.options.HolderID, token: token, epoch: reply.Epoch, } lost, runErr := adapter.runLeaderTerm(ctx, session, work) if runErr != nil { return runErr } if !lost { return nil } if err := wait(ctx, adapter.options.RetryInterval); err != nil { return err } } return ctx.Err() } func (adapter *Adapter) runLeaderTerm( ctx context.Context, session *leaderSession, work func(context.Context, controllerProvider.LeaderSession) error, ) (bool, error) { leaderCtx, cancel := context.WithCancel(ctx) defer cancel() session.ctx = leaderCtx workDone := make(chan error, 1) go func() { workDone <- work(leaderCtx, session) }() ticker := time.NewTicker(adapter.options.RenewEvery) defer ticker.Stop() deadline := time.NewTimer(adapter.options.LeaseTTL - adapter.options.RenewEvery) defer deadline.Stop() for { select { case <-ctx.Done(): cancel() <-workDone adapter.releaseLeader(session) return false, ctx.Err() case workErr := <-workDone: cancel() adapter.releaseLeader(session) return false, leaderWorkResult(ctx, workErr) case <-deadline.C: cancel() <-workDone return true, nil case <-ticker.C: renewCtx, renewCancel := context.WithTimeout(leaderCtx, adapter.options.RenewEvery) reply, err := runScript(renewCtx, adapter.client, session.keys, "renew_leader", session.generation, session.holderID, session.token, session.epoch, durationMillis(adapter.options.LeaseTTL), ) renewCancel() if err != nil || reply.Status != "ok" { cancel() <-workDone return true, nil } resetTimer(deadline, adapter.options.LeaseTTL-adapter.options.RenewEvery) } } } func leaderWorkResult(ctx context.Context, workErr error) error { if err := ctx.Err(); err != nil { return err } if workErr != nil { return workErr } return controllerProvider.ErrLeaderWorkStopped } func (adapter *Adapter) releaseLeader(session *leaderSession) { ctx, cancel := context.WithTimeout(context.Background(), adapter.options.RenewEvery) defer cancel() _, _ = runScript(ctx, adapter.client, session.keys, "release_leader", session.generation, session.holderID, session.token, session.epoch, ) } type leaderSession struct { adapter *Adapter keys upstreamKeys upstreamID string limits controllerProvider.CoordinationLimits ctx context.Context generation string holderID string token string epoch uint64 } var _ controllerProvider.LeaderSession = (*leaderSession)(nil) func (session *leaderSession) Fence() controllerProvider.Fence { if session == nil { return controllerProvider.Fence{} } return controllerProvider.Fence{Generation: session.generation, Epoch: session.epoch} } func (session *leaderSession) AcquireFetch(ctx context.Context, expected int) (controllerProvider.RequestPermit, bool, error) { if ctx == nil || session == nil || session.adapter == nil || session.ctx == nil || expected <= 0 || int64(expected) > maximumLuaInteger { return nil, false, controllerProvider.ErrInvalidCoordination } operationCtx, cancel := context.WithCancel(ctx) stop := context.AfterFunc(session.ctx, cancel) defer func() { stop() cancel() }() permitToken, err := randomToken() if err != nil { return nil, false, errors.Join(controllerProvider.ErrCoordinationUnavailable, err) } permitTTL := session.limits.MaxAttemptDuration + session.adapter.options.PermitGrace for operationCtx.Err() == nil { reply, scriptErr := runScript(operationCtx, session.adapter.client, session.keys, "acquire_fetch", session.generation, session.holderID, session.token, session.epoch, permitToken, durationMillis(session.limits.RequestInterval), session.limits.MaxInFlight, durationMillis(permitTTL), expected, session.limits.MaxTotal, ) if scriptErr != nil { if session.ctx.Err() != nil { return nil, false, controllerProvider.ErrLeadershipLost } if operationCtx.Err() != nil { return nil, false, operationCtx.Err() } if err := wait(operationCtx, session.adapter.options.RetryInterval); err != nil { return nil, false, err } continue } switch reply.Status { case "ok": return &requestPermit{ adapter: session.adapter, keys: session.keys, token: permitToken, settlementTTL: permitTTL, }, true, nil case "quota_exhausted": return nil, false, nil case "stale": return nil, false, controllerProvider.ErrLeadershipLost case "rate_limited", "at_capacity": delay := time.Duration(reply.WaitMS) * time.Millisecond if delay <= 0 { delay = session.adapter.options.RetryInterval } if err := wait(operationCtx, delay); err != nil { if session.ctx.Err() != nil { return nil, false, controllerProvider.ErrLeadershipLost } return nil, false, err } default: return nil, false, controllerProvider.ErrCoordinationUnavailable } } if session.ctx.Err() != nil { return nil, false, controllerProvider.ErrLeadershipLost } return nil, false, operationCtx.Err() } type requestPermit struct { adapter *Adapter keys upstreamKeys token string settlementTTL time.Duration mu sync.Mutex done bool } func (permit *requestPermit) Complete(ctx context.Context, fetched int) error { if fetched < 0 || int64(fetched) > maximumLuaInteger { return controllerProvider.ErrInvalidCoordination } return permit.finish(ctx, "complete_fetch", fetched) } func (permit *requestPermit) Cancel(ctx context.Context) error { return permit.finish(ctx, "cancel_fetch", 0) } func (permit *requestPermit) finish(ctx context.Context, operation string, fetched int) error { if ctx == nil || permit == nil || permit.adapter == nil || permit.token == "" || permit.settlementTTL <= 0 { return controllerProvider.ErrInvalidCoordination } permit.mu.Lock() defer permit.mu.Unlock() if permit.done { return nil } for ctx.Err() == nil { reply, err := runScript(ctx, permit.adapter.client, permit.keys, operation, permit.token, fetched, durationMillis(permit.settlementTTL)) if err != nil { if waitErr := wait(ctx, permit.adapter.options.RetryInterval); waitErr != nil { return waitErr } continue } if reply.Status != "ok" { return controllerProvider.ErrCoordinationUnavailable } permit.done = true return nil } return ctx.Err() } func randomToken() (string, error) { var token [16]byte if _, err := rand.Read(token[:]); err != nil { return "", err } return hex.EncodeToString(token[:]), nil } func durationMillis(value time.Duration) int64 { milliseconds := value / time.Millisecond if value%time.Millisecond != 0 { milliseconds++ } return int64(milliseconds) } func wait(ctx context.Context, duration time.Duration) error { timer := time.NewTimer(duration) defer timer.Stop() select { case <-timer.C: return nil case <-ctx.Done(): return ctx.Err() } } func resetTimer(timer *time.Timer, duration time.Duration) { if !timer.Stop() { select { case <-timer.C: default: } } timer.Reset(duration) } func nilInterface(value any) bool { if value == nil { return true } reflected := reflect.ValueOf(value) switch reflected.Kind() { case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Pointer, reflect.Slice: return reflected.IsNil() default: return false } }