From e54fc84a81b87abb3e82518b8de40a08ea4ac740 Mon Sep 17 00:00:00 2001 From: youfak Date: Thu, 30 Jul 2026 20:51:16 +0800 Subject: [PATCH] feat: distribute extraction admission limits --- .../redisactivity/status_integration_test.go | 6 +- internal/adapters/redisadmission/adapter.go | 141 +++++++++++ .../adapter_integration_test.go | 223 +++++++++++++++++ .../adapters/redisadmission/adapter_test.go | 225 ++++++++++++++++++ .../redisadmission/scripts/fixed_window.lua | 39 +++ internal/controller/bootstrap/bootstrap.go | 5 +- .../controller/bootstrap/bootstrap_test.go | 30 +++ .../controller/bootstrap/infrastructure.go | 28 +++ .../bootstrap/infrastructure_test.go | 56 +++++ internal/controller/extraction/service.go | 13 +- internal/platform/admission/admitter.go | 9 + internal/platform/admission/allow_all_test.go | 3 + internal/platform/admission/fixed_window.go | 7 +- .../platform/admission/fixed_window_test.go | 24 ++ scripts/test-redis.ps1 | 2 +- 15 files changed, 796 insertions(+), 15 deletions(-) create mode 100644 internal/adapters/redisadmission/adapter.go create mode 100644 internal/adapters/redisadmission/adapter_integration_test.go create mode 100644 internal/adapters/redisadmission/adapter_test.go create mode 100644 internal/adapters/redisadmission/scripts/fixed_window.lua create mode 100644 internal/platform/admission/admitter.go diff --git a/internal/adapters/redisactivity/status_integration_test.go b/internal/adapters/redisactivity/status_integration_test.go index f4cf057..5bf4823 100644 --- a/internal/adapters/redisactivity/status_integration_test.go +++ b/internal/adapters/redisactivity/status_integration_test.go @@ -22,7 +22,11 @@ func TestReadStateInventoryFailsClosedWhileExpiredCleanupIsBacklogged(t *testing if err != nil { t.Fatalf("New() error = %v", err) } - now := time.Date(2026, 7, 30, 12, 0, 0, 0, time.UTC) + redisTime, err := fixture.Client.Time(context.Background()).Result() + if err != nil { + t.Fatalf("Redis TIME error = %v", err) + } + now := redisTime.UTC().Add(time.Minute) _, err = bounded.UpsertFetched(context.Background(), "provider-a", activitypool.FetchedBatch{ ObservedAt: now, ConfiguredTTL: time.Second, MaxSize: 10, Proxies: []proxyDomain.Proxy{ diff --git a/internal/adapters/redisadmission/adapter.go b/internal/adapters/redisadmission/adapter.go new file mode 100644 index 0000000..4d64441 --- /dev/null +++ b/internal/adapters/redisadmission/adapter.go @@ -0,0 +1,141 @@ +package redisadmission + +import ( + "context" + "crypto/sha256" + _ "embed" + "encoding/binary" + "encoding/hex" + "errors" + "reflect" + "regexp" + "strconv" + "strings" + "time" + + "github.com/redis/go-redis/v9" + + "proxy-pool/internal/platform/admission" +) + +const MaximumLuaInteger int64 = 1<<53 - 1 + +var ( + ErrInvalidOptions = errors.New("invalid Redis admission options") + namespacePattern = regexp.MustCompile(`^[A-Za-z0-9._-]+$`) + redisAdmissionKey = "pp:{admission}:" + redisAdmissionTail = ":window" +) + +//go:embed scripts/fixed_window.lua +var fixedWindowSource string + +var fixedWindowScript = redis.NewScript(fixedWindowSource) + +type Options struct { + Namespace string + Window time.Duration + Global int64 + PerKey int64 +} + +type Adapter struct { + client redis.Scripter + key string + windowMillis int64 + global int64 + perKey int64 +} + +var _ admission.Admitter = (*Adapter)(nil) + +func New(client redis.Scripter, options Options) (*Adapter, error) { + if nilInterface(client) || options.Namespace != strings.TrimSpace(options.Namespace) || + !namespacePattern.MatchString(options.Namespace) || options.Window <= 0 || + options.Window%time.Millisecond != 0 || options.Global < 0 || options.PerKey < 0 || + (options.Global == 0 && options.PerKey == 0) || options.Global > MaximumLuaInteger || + options.PerKey > MaximumLuaInteger { + return nil, ErrInvalidOptions + } + windowMillis := options.Window.Milliseconds() + if windowMillis <= 0 || windowMillis > MaximumLuaInteger { + return nil, ErrInvalidOptions + } + return &Adapter{ + client: client, + key: redisAdmissionKey + options.Namespace + redisAdmissionTail, + windowMillis: windowMillis, + global: options.Global, + perKey: options.PerKey, + }, nil +} + +func (adapter *Adapter) Admit(ctx context.Context, identity string) error { + if ctx == nil || identity == "" { + return admission.ErrInvalidIdentity + } + if adapter == nil || nilInterface(adapter.client) || adapter.key == "" || adapter.windowMillis <= 0 || + adapter.windowMillis > MaximumLuaInteger || adapter.global < 0 || adapter.perKey < 0 || + (adapter.global == 0 && adapter.perKey == 0) || adapter.global > MaximumLuaInteger || + adapter.perKey > MaximumLuaInteger { + return admission.ErrUnavailable + } + if err := ctx.Err(); err != nil { + return err + } + + result, err := fixedWindowScript.Run(ctx, adapter.client, []string{adapter.key}, + strconv.FormatInt(adapter.windowMillis, 10), + strconv.FormatInt(adapter.global, 10), + strconv.FormatInt(adapter.perKey, 10), + hashedIdentityField(identity), + ).Result() + if err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } + return admission.ErrUnavailable + } + + var status string + switch value := result.(type) { + case string: + status = value + case []byte: + status = string(value) + default: + return admission.ErrUnavailable + } + switch status { + case "ok": + return nil + case "global": + return admission.ErrGlobalLimit + case "per_key": + return admission.ErrPerKeyLimit + default: + return admission.ErrUnavailable + } +} + +func hashedIdentityField(identity string) string { + var size [8]byte + binary.BigEndian.PutUint64(size[:], uint64(len(identity))) + digest := sha256.New() + _, _ = digest.Write(size[:]) + _, _ = digest.Write([]byte(identity)) + return "client:" + hex.EncodeToString(digest.Sum(nil)) +} + +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 + } +} diff --git a/internal/adapters/redisadmission/adapter_integration_test.go b/internal/adapters/redisadmission/adapter_integration_test.go new file mode 100644 index 0000000..d5f4fde --- /dev/null +++ b/internal/adapters/redisadmission/adapter_integration_test.go @@ -0,0 +1,223 @@ +//go:build integration + +package redisadmission + +import ( + "context" + "errors" + "fmt" + "os" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/redis/go-redis/v9" + + "proxy-pool/internal/platform/admission" +) + +var integrationNamespaceSequence atomic.Uint64 + +func TestAdaptersShareGlobalAndPerKeyLimits(t *testing.T) { + fixture := newRedisFixture(t) + first := fixture.adapter(t, fixture.namespace, time.Minute, 3, 2) + second := fixture.adapter(t, fixture.namespace, time.Minute, 3, 2) + + if err := first.Admit(t.Context(), "client-a"); err != nil { + t.Fatalf("first Admit(client-a): %v", err) + } + if err := second.Admit(t.Context(), "client-a"); err != nil { + t.Fatalf("second Admit(client-a): %v", err) + } + if err := first.Admit(t.Context(), "client-a"); !errors.Is(err, admission.ErrPerKeyLimit) { + t.Fatalf("shared per-key limit error = %v, want ErrPerKeyLimit", err) + } + if err := second.Admit(t.Context(), "client-b"); err != nil { + t.Fatalf("second Admit(client-b): %v", err) + } + if err := first.Admit(t.Context(), "client-c"); !errors.Is(err, admission.ErrGlobalLimit) { + t.Fatalf("shared global limit error = %v, want ErrGlobalLimit", err) + } +} + +func TestConcurrentAdmissionsAreExactAcrossAdapters(t *testing.T) { + fixture := newRedisFixture(t) + tests := []struct { + name string + global int64 + perKey int64 + identity func(int) string + want int64 + wantReject error + }{ + { + name: "global", global: 37, perKey: 1000, want: 37, wantReject: admission.ErrGlobalLimit, + identity: func(index int) string { return fmt.Sprintf("client-%d", index) }, + }, + { + name: "per-key", global: 1000, perKey: 23, want: 23, wantReject: admission.ErrPerKeyLimit, + identity: func(int) string { return "shared-client" }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + namespace := fixture.namespace + "-" + tt.name + first := fixture.adapter(t, namespace, time.Minute, tt.global, tt.perKey) + second := fixture.adapter(t, namespace, time.Minute, tt.global, tt.perKey) + var accepted atomic.Int64 + var wrong atomic.Int64 + var wait sync.WaitGroup + for index := range 200 { + wait.Add(1) + go func() { + defer wait.Done() + adapter := first + if index%2 == 1 { + adapter = second + } + err := adapter.Admit(t.Context(), tt.identity(index)) + switch { + case err == nil: + accepted.Add(1) + case !errors.Is(err, tt.wantReject): + wrong.Add(1) + } + }() + } + wait.Wait() + if got := accepted.Load(); got != tt.want { + t.Fatalf("accepted = %d, want %d", got, tt.want) + } + if got := wrong.Load(); got != 0 { + t.Fatalf("unexpected rejection count = %d", got) + } + }) + } +} + +func TestWindowRolloverRemovesOldClientFields(t *testing.T) { + fixture := newRedisFixture(t) + const window = 100 * time.Millisecond + adapter := fixture.adapter(t, fixture.namespace, window, 100, 10) + if err := adapter.Admit(t.Context(), "old-client"); err != nil { + t.Fatalf("Admit(old-client): %v", err) + } + if err := fixture.client.Persist(t.Context(), adapter.key).Err(); err != nil { + t.Fatalf("PERSIST admission key: %v", err) + } + storedWindow, err := fixture.client.HGet(t.Context(), adapter.key, "window").Int64() + if err != nil { + t.Fatalf("read stored window: %v", err) + } + waitForRedisWindow(t, fixture.client, window, storedWindow) + + if err := adapter.Admit(t.Context(), "new-client"); err != nil { + t.Fatalf("Admit(new-client): %v", err) + } + oldField := hashedIdentityField("old-client") + if exists, err := fixture.client.HExists(t.Context(), adapter.key, oldField).Result(); err != nil || exists { + t.Fatalf("old client field exists = %v, error = %v", exists, err) + } + newField := hashedIdentityField("new-client") + if exists, err := fixture.client.HExists(t.Context(), adapter.key, newField).Result(); err != nil || !exists { + t.Fatalf("new client field exists = %v, error = %v", exists, err) + } + if length, err := fixture.client.HLen(t.Context(), adapter.key).Result(); err != nil || length != 3 { + t.Fatalf("current window hash length = %d, error = %v, want 3", length, err) + } +} + +func TestNamespacesAreIsolated(t *testing.T) { + fixture := newRedisFixture(t) + first := fixture.adapter(t, fixture.namespace+"-a", time.Minute, 1, 1) + second := fixture.adapter(t, fixture.namespace+"-b", time.Minute, 1, 1) + + if err := first.Admit(t.Context(), "same-client"); err != nil { + t.Fatalf("first namespace Admit(): %v", err) + } + if err := second.Admit(t.Context(), "same-client"); err != nil { + t.Fatalf("second namespace Admit(): %v", err) + } + if err := first.Admit(t.Context(), "same-client"); !errors.Is(err, admission.ErrGlobalLimit) { + t.Fatalf("first namespace second Admit() error = %v", err) + } +} + +type redisFixture struct { + client *redis.Client + namespace string +} + +func newRedisFixture(t *testing.T) redisFixture { + t.Helper() + redisURL := os.Getenv("PROXY_POOL_TEST_REDIS_URL") + if redisURL == "" { + t.Skip("PROXY_POOL_TEST_REDIS_URL is not set") + } + options, err := redis.ParseURL(redisURL) + if err != nil { + t.Fatalf("parse PROXY_POOL_TEST_REDIS_URL: %v", err) + } + client := redis.NewClient(options) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + if err := client.Ping(ctx).Err(); err != nil { + _ = client.Close() + t.Fatalf("ping Redis: %v", err) + } + namespace := fmt.Sprintf("admission-it-%d-%d-%d", os.Getpid(), time.Now().UnixNano(), integrationNamespaceSequence.Add(1)) + t.Cleanup(func() { + cleanupCtx, cleanupCancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cleanupCancel() + pattern := redisAdmissionKey + namespace + "*" + var cursor uint64 + for { + keys, next, scanErr := client.Scan(cleanupCtx, cursor, pattern, 128).Result() + if scanErr != nil { + t.Errorf("scan Redis admission keys: %v", scanErr) + break + } + if len(keys) > 0 { + if unlinkErr := client.Unlink(cleanupCtx, keys...).Err(); unlinkErr != nil { + t.Errorf("remove Redis admission keys: %v", unlinkErr) + break + } + } + cursor = next + if cursor == 0 { + break + } + } + _ = client.Close() + }) + return redisFixture{client: client, namespace: namespace} +} + +func (fixture redisFixture) adapter(t *testing.T, namespace string, window time.Duration, global, perKey int64) *Adapter { + t.Helper() + return mustAdapter(t, fixture.client, Options{ + Namespace: namespace, + Window: window, + Global: global, + PerKey: perKey, + }) +} + +func waitForRedisWindow(t *testing.T, client *redis.Client, window time.Duration, previous int64) { + t.Helper() + deadline := time.Now().Add(2 * time.Second) + for time.Now().Before(deadline) { + redisTime, err := client.Time(t.Context()).Result() + if err != nil { + t.Fatalf("Redis TIME: %v", err) + } + windowID := redisTime.UnixMilli() / window.Milliseconds() + if windowID != previous { + return + } + time.Sleep(5 * time.Millisecond) + } + t.Fatal("timed out waiting for Redis window rollover") +} diff --git a/internal/adapters/redisadmission/adapter_test.go b/internal/adapters/redisadmission/adapter_test.go new file mode 100644 index 0000000..dc79331 --- /dev/null +++ b/internal/adapters/redisadmission/adapter_test.go @@ -0,0 +1,225 @@ +package redisadmission + +import ( + "context" + "crypto/sha256" + "encoding/binary" + "errors" + "fmt" + "strings" + "testing" + "time" + + "github.com/redis/go-redis/v9" + + "proxy-pool/internal/platform/admission" +) + +type scriptCall struct { + keys []string + args []any +} + +type scriptClient struct { + redis.Scripter + result any + err error + calls []scriptCall +} + +func (client *scriptClient) EvalSha(_ context.Context, _ string, keys []string, args ...any) *redis.Cmd { + client.calls = append(client.calls, scriptCall{keys: append([]string(nil), keys...), args: append([]any(nil), args...)}) + return redis.NewCmdResult(client.result, client.err) +} + +func TestNewRejectsInvalidDependenciesAndOptions(t *testing.T) { + t.Parallel() + validClient := &scriptClient{result: "ok"} + valid := Options{Namespace: "listener-a", Window: time.Minute, Global: 10, PerKey: 2} + var typedNilClient *redis.Client + + tests := []struct { + name string + client redis.Scripter + options Options + }{ + {name: "nil client", options: valid}, + {name: "typed nil client", client: typedNilClient, options: valid}, + {name: "empty namespace", client: validClient, options: withNamespace(valid, "")}, + {name: "leading namespace whitespace", client: validClient, options: withNamespace(valid, " listener-a")}, + {name: "trailing namespace whitespace", client: validClient, options: withNamespace(valid, "listener-a ")}, + {name: "cluster tag in namespace", client: validClient, options: withNamespace(valid, "listener{a}")}, + {name: "separator in namespace", client: validClient, options: withNamespace(valid, "listener:a")}, + {name: "zero window", client: validClient, options: withWindow(valid, 0)}, + {name: "sub-millisecond window", client: validClient, options: withWindow(valid, time.Microsecond)}, + {name: "fractional millisecond window", client: validClient, options: withWindow(valid, time.Millisecond+time.Microsecond)}, + {name: "negative global", client: validClient, options: withGlobal(valid, -1)}, + {name: "negative per-key", client: validClient, options: withPerKey(valid, -1)}, + {name: "zero limits", client: validClient, options: withLimits(valid, 0, 0)}, + {name: "global exceeds Lua integer", client: validClient, options: withGlobal(valid, MaximumLuaInteger+1)}, + {name: "per-key exceeds Lua integer", client: validClient, options: withPerKey(valid, MaximumLuaInteger+1)}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + if adapter, err := New(tt.client, tt.options); !errors.Is(err, ErrInvalidOptions) || adapter != nil { + t.Fatalf("New() = (%v, %v), want nil adapter and ErrInvalidOptions", adapter, err) + } + }) + } +} + +func TestAdmitMapsScriptStatusAndPreservesInputErrors(t *testing.T) { + t.Parallel() + tests := []struct { + name string + status any + want error + }{ + {name: "accepted", status: "ok"}, + {name: "global limit", status: "global", want: admission.ErrGlobalLimit}, + {name: "per-key limit", status: "per_key", want: admission.ErrPerKeyLimit}, + {name: "unexpected status", status: "unknown", want: admission.ErrUnavailable}, + {name: "unexpected reply type", status: int64(1), want: admission.ErrUnavailable}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + client := &scriptClient{result: tt.status} + adapter := mustAdapter(t, client, Options{Namespace: "listener-a", Window: time.Minute, Global: 10, PerKey: 2}) + err := adapter.Admit(context.Background(), "client-a") + if !errors.Is(err, tt.want) { + t.Fatalf("Admit() error = %v, want %v", err, tt.want) + } + }) + } + + adapter := mustAdapter(t, &scriptClient{result: "ok"}, Options{Namespace: "listener-a", Window: time.Minute, Global: 1}) + if err := adapter.Admit(nil, "client-a"); !errors.Is(err, admission.ErrInvalidIdentity) { + t.Fatalf("Admit(nil context) error = %v, want ErrInvalidIdentity", err) + } + if err := adapter.Admit(context.Background(), ""); !errors.Is(err, admission.ErrInvalidIdentity) { + t.Fatalf("Admit(empty identity) error = %v, want ErrInvalidIdentity", err) + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if err := adapter.Admit(ctx, "client-a"); !errors.Is(err, context.Canceled) { + t.Fatalf("Admit(canceled context) error = %v, want context.Canceled", err) + } + var nilAdapter *Adapter + if err := nilAdapter.Admit(context.Background(), "client-a"); !errors.Is(err, admission.ErrUnavailable) { + t.Fatalf("nil Adapter.Admit() error = %v, want ErrUnavailable", err) + } +} + +func TestAdmitUsesHashedIdentityAndExactIntegerArguments(t *testing.T) { + t.Parallel() + client := &scriptClient{result: "ok"} + adapter := mustAdapter(t, client, Options{ + Namespace: "listener-a", Window: 1500 * time.Millisecond, + Global: MaximumLuaInteger, PerKey: MaximumLuaInteger - 1, + }) + const identity = "private-client@example.test" + if err := adapter.Admit(context.Background(), identity); err != nil { + t.Fatalf("Admit(): %v", err) + } + if len(client.calls) != 1 { + t.Fatalf("script calls = %d, want 1", len(client.calls)) + } + call := client.calls[0] + if len(call.keys) != 1 || call.keys[0] != "pp:{admission}:listener-a:window" { + t.Fatalf("script keys = %q", call.keys) + } + wantArgs := []string{"1500", fmt.Sprint(MaximumLuaInteger), fmt.Sprint(MaximumLuaInteger - 1), clientField(identity)} + if got := stringify(call.args); fmt.Sprint(got) != fmt.Sprint(wantArgs) { + t.Fatalf("script args = %q, want %q", got, wantArgs) + } + serialized := fmt.Sprint(call.keys, call.args) + if strings.Contains(serialized, identity) { + t.Fatalf("Redis input leaks raw identity: %s", serialized) + } + if got := clientField(identity); len(got) != len("client:")+sha256.Size*2 || !strings.HasPrefix(got, "client:") { + t.Fatalf("client field = %q", got) + } +} + +func TestAdmitFailsClosedWithoutLeakingIdentity(t *testing.T) { + t.Parallel() + const identity = "sensitive-client" + client := &scriptClient{err: errors.New("backend failure: " + identity)} + adapter := mustAdapter(t, client, Options{Namespace: "listener-a", Window: time.Minute, Global: 1}) + err := adapter.Admit(context.Background(), identity) + if !errors.Is(err, admission.ErrUnavailable) { + t.Fatalf("Admit() error = %v, want ErrUnavailable", err) + } + if strings.Contains(err.Error(), identity) { + t.Fatalf("Admit() error leaks identity: %v", err) + } +} + +func TestAdmitFailsClosedForUninitializedAdapter(t *testing.T) { + t.Parallel() + var typedNilClient *redis.Client + tests := []*Adapter{ + {}, + {client: typedNilClient, key: "pp:{admission}:test:window", windowMillis: 60_000, global: 1}, + } + for _, adapter := range tests { + if err := adapter.Admit(context.Background(), "client-a"); !errors.Is(err, admission.ErrUnavailable) { + t.Fatalf("Admit() error = %v, want ErrUnavailable", err) + } + } +} + +func clientField(identity string) string { + var size [8]byte + binary.BigEndian.PutUint64(size[:], uint64(len(identity))) + digest := sha256.New() + _, _ = digest.Write(size[:]) + _, _ = digest.Write([]byte(identity)) + return fmt.Sprintf("client:%x", digest.Sum(nil)) +} + +func stringify(values []any) []string { + result := make([]string, len(values)) + for index, value := range values { + result[index] = fmt.Sprint(value) + } + return result +} + +func mustAdapter(t *testing.T, client redis.Scripter, options Options) *Adapter { + t.Helper() + adapter, err := New(client, options) + if err != nil { + t.Fatalf("New(): %v", err) + } + return adapter +} + +func withNamespace(options Options, namespace string) Options { + options.Namespace = namespace + return options +} + +func withWindow(options Options, window time.Duration) Options { + options.Window = window + return options +} + +func withGlobal(options Options, limit int64) Options { + options.Global = limit + return options +} + +func withPerKey(options Options, limit int64) Options { + options.PerKey = limit + return options +} + +func withLimits(options Options, global, perKey int64) Options { + options.Global = global + options.PerKey = perKey + return options +} diff --git a/internal/adapters/redisadmission/scripts/fixed_window.lua b/internal/adapters/redisadmission/scripts/fixed_window.lua new file mode 100644 index 0000000..a125449 --- /dev/null +++ b/internal/adapters/redisadmission/scripts/fixed_window.lua @@ -0,0 +1,39 @@ +local key = KEYS[1] +local window_ms = tonumber(ARGV[1]) +local global_limit = tonumber(ARGV[2]) +local per_key_limit = tonumber(ARGV[3]) +local client_field = ARGV[4] + +local redis_time = redis.call("TIME") +local now_ms = tonumber(redis_time[1]) * 1000 + math.floor(tonumber(redis_time[2]) / 1000) +local window_id = math.floor(now_ms / window_ms) +local stored_window = redis.call("HGET", key, "window") + +if not stored_window or tonumber(stored_window) ~= window_id then + redis.call("DEL", key) + redis.call("HSET", key, "window", window_id) +end + +if global_limit > 0 then + local global_used = tonumber(redis.call("HGET", key, "global") or "0") + if global_used >= global_limit then + return "global" + end +end + +if per_key_limit > 0 then + local client_used = tonumber(redis.call("HGET", key, client_field) or "0") + if client_used >= per_key_limit then + return "per_key" + end +end + +if global_limit > 0 then + redis.call("HINCRBY", key, "global", 1) +end +if per_key_limit > 0 then + redis.call("HINCRBY", key, client_field, 1) +end + +redis.call("PEXPIREAT", key, (window_id + 1) * window_ms) +return "ok" diff --git a/internal/controller/bootstrap/bootstrap.go b/internal/controller/bootstrap/bootstrap.go index 526005a..db5fd6f 100644 --- a/internal/controller/bootstrap/bootstrap.go +++ b/internal/controller/bootstrap/bootstrap.go @@ -55,6 +55,7 @@ type ports struct { activity activityStore readiness distribution.ReadinessChecker metricsReadiness platformMetrics.ReadinessChecker + admission admission.Admitter coordinator provider.Coordinator credentials credentials.Store providerResults provider.ResultRecorder @@ -145,10 +146,10 @@ func run(ctx context.Context, options Options, infrastructure infrastructure, fa dependencies := controllerRuntime.Dependencies{} if loaded.Value.Distribution.Enabled { - if nilInterface(opened.activity) || nilInterface(opened.readiness) { + if nilInterface(opened.activity) || nilInterface(opened.readiness) || nilInterface(opened.admission) { return errors.Join(ErrStartup, ErrInvalidOptions) } - service, serviceErr := extraction.NewService(opened.activity, extractionPolicy(loaded.Value), admission.AllowAll{}, options.Now) + service, serviceErr := extraction.NewService(opened.activity, extractionPolicy(loaded.Value), opened.admission, options.Now) if serviceErr != nil { return fmt.Errorf("%w: build extraction service: %w", ErrStartup, serviceErr) } diff --git a/internal/controller/bootstrap/bootstrap_test.go b/internal/controller/bootstrap/bootstrap_test.go index af3e701..fd15bf1 100644 --- a/internal/controller/bootstrap/bootstrap_test.go +++ b/internal/controller/bootstrap/bootstrap_test.go @@ -16,6 +16,7 @@ import ( "proxy-pool/internal/domain/adminstate" extractionDomain "proxy-pool/internal/domain/extraction" "proxy-pool/internal/domain/upstream" + "proxy-pool/internal/platform/admission" "proxy-pool/internal/platform/credentials" ) @@ -33,6 +34,7 @@ func TestRunLoadsOneSnapshotCommitsItAndClosesInfrastructure(t *testing.T) { closeErr := errors.New("close failed") infrastructure := &stubInfrastructure{ports: ports{ state: state, activity: activity, readiness: readyStub{}, metricsReadiness: readyStub{}, + admission: admission.AllowAll{}, coordinator: coordinatorStub{}, credentials: credentialStore, close: func() error { return closeErr }, }} @@ -113,6 +115,33 @@ func TestRunRejectsMissingAdminFingerprintKeyBeforeOpeningInfrastructure(t *test } } +func TestRunRejectsMissingDistributionAdmissionDependency(t *testing.T) { + t.Parallel() + credentialStore, err := credentials.NewMemoryStore(10) + if err != nil { + t.Fatalf("NewMemoryStore(): %v", err) + } + infrastructure := &stubInfrastructure{ports: ports{ + state: adminstate.NewMemoryStore(), activity: &stubActivityStore{}, + readiness: readyStub{}, metricsReadiness: readyStub{}, + coordinator: coordinatorStub{}, credentials: credentialStore, + close: func() error { return nil }, + }} + factory := &recordingRuntimeFactory{runner: runnerStub{err: errors.New("runtime should not start")}} + err = run(context.Background(), Options{ + ConfigPath: "controller.yaml", + Resolver: &memoryResolver{files: map[string][]byte{"controller.yaml": []byte(bootstrapTestConfig)}}, + Now: time.Now, + FingerprintKey: bootstrapTestFingerprintKey, + }, infrastructure, factory) + if !errors.Is(err, ErrStartup) || !errors.Is(err, ErrInvalidOptions) { + t.Fatalf("run() error = %v, want ErrStartup and ErrInvalidOptions", err) + } + if factory.configuration != nil { + t.Fatal("runtime factory called without admission dependency") + } +} + func TestRunSupportsProviderOnlyConfigurationWithoutHTTPRuntime(t *testing.T) { source := strings.ReplaceAll(bootstrapTestConfig, "distribution:\n enabled: true", "distribution:\n enabled: false") source = strings.ReplaceAll(source, "admin:\n enabled: true", "admin:\n enabled: false") @@ -151,6 +180,7 @@ func TestRunAdminDisableStopsActiveProviderRuntime(t *testing.T) { stopped := make(chan string, 2) infrastructure := &stubInfrastructure{ports: ports{ state: state, activity: &stubActivityStore{}, readiness: readyStub{}, metricsReadiness: readyStub{}, + admission: admission.AllowAll{}, coordinator: coordinatorFunc(func(ctx context.Context, upstreamID string) error { started <- upstreamID <-ctx.Done() diff --git a/internal/controller/bootstrap/infrastructure.go b/internal/controller/bootstrap/infrastructure.go index 8f1c4e8..8c2501f 100644 --- a/internal/controller/bootstrap/infrastructure.go +++ b/internal/controller/bootstrap/infrastructure.go @@ -13,9 +13,11 @@ import ( "proxy-pool/internal/adapters/postgresadmin" "proxy-pool/internal/adapters/redisactivity" + "proxy-pool/internal/adapters/redisadmission" "proxy-pool/internal/adapters/redisprovider" "proxy-pool/internal/config" controllerProvider "proxy-pool/internal/controller/provider" + "proxy-pool/internal/platform/admission" "proxy-pool/internal/platform/credentials" platformMetrics "proxy-pool/internal/platform/metrics" ) @@ -131,6 +133,16 @@ func (infrastructure *productionInfrastructure) Open( opened.activity = adapter opened.readiness = redisReadiness{client: redisClient} opened.credentials = credentialStore + if configuration.Distribution.Enabled { + opened.admission, err = newDistributionAdmitter( + redisClient, + namespace, + configuration.Distribution.Limits, + ) + if err != nil { + return ports{}, err + } + } if providersEnabled { stats, statsErr := controllerProvider.NewStatsRecorder(config.MaximumUpstreams) if statsErr != nil { @@ -161,6 +173,22 @@ func (infrastructure *productionInfrastructure) Open( return opened, nil } +func newDistributionAdmitter( + client redis.Scripter, + namespace string, + limits config.Limits, +) (admission.Admitter, error) { + if limits.RequestsPerMinute == 0 && limits.RequestsPerMinutePerClient == 0 { + return admission.AllowAll{}, nil + } + return redisadmission.New(client, redisadmission.Options{ + Namespace: namespace, + Window: time.Minute, + Global: int64(limits.RequestsPerMinute), + PerKey: int64(limits.RequestsPerMinutePerClient), + }) +} + func resolveRedisNamespace(configured string) (string, error) { if strings.TrimSpace(configured) != configured { return "", ErrInvalidOptions diff --git a/internal/controller/bootstrap/infrastructure_test.go b/internal/controller/bootstrap/infrastructure_test.go index bfa9d0c..f9be8dc 100644 --- a/internal/controller/bootstrap/infrastructure_test.go +++ b/internal/controller/bootstrap/infrastructure_test.go @@ -3,13 +3,34 @@ package bootstrap import ( "context" "errors" + "fmt" "strings" "sync/atomic" "testing" + "github.com/redis/go-redis/v9" + "proxy-pool/internal/config" + "proxy-pool/internal/platform/admission" ) +type recordingAdmissionScripter struct { + redis.Scripter + keys []string + args []any +} + +func (client *recordingAdmissionScripter) EvalSha( + _ context.Context, + _ string, + keys []string, + args ...any, +) *redis.Cmd { + client.keys = append([]string(nil), keys...) + client.args = append([]any(nil), args...) + return redis.NewCmdResult("ok", nil) +} + func TestProductionInfrastructureRejectsInvalidStorageWithoutLeakingURLs(t *testing.T) { t.Parallel() postgresSecret := "postgres-secret" @@ -31,6 +52,41 @@ func TestProductionInfrastructureRejectsInvalidStorageWithoutLeakingURLs(t *test } } +func TestNewDistributionAdmitterPassesConfiguredLimits(t *testing.T) { + t.Parallel() + client := &recordingAdmissionScripter{} + limiter, err := newDistributionAdmitter(client, "controller-a", config.Limits{ + RequestsPerMinute: 321, + RequestsPerMinutePerClient: 17, + }) + if err != nil { + t.Fatalf("newDistributionAdmitter(): %v", err) + } + if err := limiter.Admit(context.Background(), "client-a"); err != nil { + t.Fatalf("Admit(): %v", err) + } + if len(client.keys) != 1 || len(client.args) != 4 { + t.Fatalf("Redis admission call = keys:%v args:%v", client.keys, client.args) + } + if got := fmt.Sprintf("%s|%v|%v|%v", client.keys[0], client.args[0], client.args[1], client.args[2]); got != "pp:{admission}:controller-a:window|60000|321|17" { + t.Fatalf("Redis admission inputs = %q", got) + } +} + +func TestNewDistributionAdmitterAllowsAllWhenQuotasAreDisabled(t *testing.T) { + t.Parallel() + limiter, err := newDistributionAdmitter(nil, "controller-a", config.Limits{}) + if err != nil { + t.Fatalf("newDistributionAdmitter(): %v", err) + } + if _, ok := limiter.(admission.AllowAll); !ok { + t.Fatalf("limiter type = %T, want admission.AllowAll", limiter) + } + if err := limiter.Admit(context.Background(), "client-a"); err != nil { + t.Fatalf("AllowAll.Admit(): %v", err) + } +} + func TestSelectMetricsReadinessPreservesDistributionWhenAdminStoreFails(t *testing.T) { t.Parallel() adminCalls := &atomic.Int64{} diff --git a/internal/controller/extraction/service.go b/internal/controller/extraction/service.go index 0d3ad05..cdd1020 100644 --- a/internal/controller/extraction/service.go +++ b/internal/controller/extraction/service.go @@ -7,6 +7,7 @@ import ( "time" domain "proxy-pool/internal/domain/extraction" + "proxy-pool/internal/platform/admission" ) var ( @@ -70,16 +71,12 @@ type Response struct { type Service struct { store domain.Store policy Policy - admission Admission + admission admission.Admitter now func() time.Time } -type Admission interface { - Admit(context.Context, string) error -} - -func NewService(store domain.Store, policy Policy, admission Admission, now func() time.Time) (*Service, error) { - if admission == nil { +func NewService(store domain.Store, policy Policy, admitter admission.Admitter, now func() time.Time) (*Service, error) { + if admitter == nil { return nil, fmt.Errorf("%w: admission is required", ErrInvalidServicePolicy) } if store == nil { @@ -95,7 +92,7 @@ func NewService(store domain.Store, policy Policy, admission Admission, now func if now == nil { now = time.Now } - return &Service{store: store, policy: policy, admission: admission, now: now}, nil + return &Service{store: store, policy: policy, admission: admitter, now: now}, nil } func (s *Service) Extract(ctx context.Context, request Request) (Response, error) { diff --git a/internal/platform/admission/admitter.go b/internal/platform/admission/admitter.go new file mode 100644 index 0000000..9ac2d49 --- /dev/null +++ b/internal/platform/admission/admitter.go @@ -0,0 +1,9 @@ +package admission + +import "context" + +// Admitter is the shared admission-control contract used by controller +// services and infrastructure adapters. +type Admitter interface { + Admit(context.Context, string) error +} diff --git a/internal/platform/admission/allow_all_test.go b/internal/platform/admission/allow_all_test.go index 6ebdcb8..50ab199 100644 --- a/internal/platform/admission/allow_all_test.go +++ b/internal/platform/admission/allow_all_test.go @@ -15,6 +15,9 @@ func TestAllowAllPreservesContextAndIdentityValidation(t *testing.T) { if err := admission.Admit(context.Background(), ""); !errors.Is(err, ErrInvalidIdentity) { t.Fatalf("Admit(empty identity) error = %v", err) } + if err := admission.Admit(nil, "client-a"); !errors.Is(err, ErrInvalidIdentity) { + t.Fatalf("Admit(nil context) error = %v", err) + } ctx, cancel := context.WithCancel(context.Background()) cancel() if err := admission.Admit(ctx, "client-a"); !errors.Is(err, context.Canceled) { diff --git a/internal/platform/admission/fixed_window.go b/internal/platform/admission/fixed_window.go index 983f709..f9f913e 100644 --- a/internal/platform/admission/fixed_window.go +++ b/internal/platform/admission/fixed_window.go @@ -12,6 +12,7 @@ var ( ErrInvalidIdentity = errors.New("invalid admission identity") ErrGlobalLimit = errors.New("global admission limit exceeded") ErrPerKeyLimit = errors.New("per-key admission limit exceeded") + ErrUnavailable = errors.New("admission control unavailable") ) type FixedWindowConfig struct { @@ -54,12 +55,12 @@ func NewFixedWindow(config FixedWindowConfig) (*FixedWindow, error) { } func (l *FixedWindow) Admit(ctx context.Context, key string) error { + if ctx == nil || l == nil || key == "" { + return ErrInvalidIdentity + } if err := ctx.Err(); err != nil { return err } - if l == nil || key == "" { - return ErrInvalidIdentity - } l.mu.Lock() defer l.mu.Unlock() diff --git a/internal/platform/admission/fixed_window_test.go b/internal/platform/admission/fixed_window_test.go index 01818f1..4bc5866 100644 --- a/internal/platform/admission/fixed_window_test.go +++ b/internal/platform/admission/fixed_window_test.go @@ -74,3 +74,27 @@ func TestFixedWindowResetsAndIsConcurrencySafe(t *testing.T) { t.Fatalf("Admit(after reset): %v", err) } } + +func TestFixedWindowRejectsNilContextAndEmptyIdentity(t *testing.T) { + t.Parallel() + limiter, err := NewFixedWindow(FixedWindowConfig{ + Window: time.Minute, + Global: 1, + }) + if err != nil { + t.Fatalf("NewFixedWindow(): %v", err) + } + + if err := limiter.Admit(nil, "client-a"); !errors.Is(err, ErrInvalidIdentity) { + t.Fatalf("Admit(nil context) error = %v, want ErrInvalidIdentity", err) + } + if err := limiter.Admit(context.Background(), ""); !errors.Is(err, ErrInvalidIdentity) { + t.Fatalf("Admit(empty identity) error = %v, want ErrInvalidIdentity", err) + } +} + +func TestAdmitterImplementations(t *testing.T) { + t.Parallel() + var _ Admitter = (*FixedWindow)(nil) + var _ Admitter = AllowAll{} +} diff --git a/scripts/test-redis.ps1 b/scripts/test-redis.ps1 index a258cb4..6d4e6e2 100644 --- a/scripts/test-redis.ps1 +++ b/scripts/test-redis.ps1 @@ -14,7 +14,7 @@ try { $env:PROXY_POOL_TEST_REDIS_URL = "redis://127.0.0.1:16379/15" Push-Location $repositoryRoot try { - go test -count=1 -tags=integration -timeout 60s ./internal/adapters/redisactivity/... ./internal/adapters/redisprovider/... + go test -count=1 -tags=integration -timeout 60s ./internal/adapters/redisactivity/... ./internal/adapters/redisadmission/... ./internal/adapters/redisprovider/... if ($LASTEXITCODE -ne 0) { throw "Redis integration tests failed with exit code $LASTEXITCODE" }