proxy-pool/internal/adapters/redisprovider/adapter.go

383 lines
11 KiB
Go

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._-]+$`)
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 > controllerProvider.MaximumCoordinationInteger ||
int64(limits.MaxInFlight) > controllerProvider.MaximumCoordinationInteger ||
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) > controllerProvider.MaximumCoordinationInteger {
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) > controllerProvider.MaximumCoordinationInteger {
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
}
}