From 40f4b3ffab7a8c01fe1bbac69b5ced830214646e Mon Sep 17 00:00:00 2001 From: youfak Date: Thu, 30 Jul 2026 16:11:04 +0800 Subject: [PATCH] feat: enforce distributed provider fetch quota --- internal/adapters/redisprovider/adapter.go | 88 ++++--- .../adapters/redisprovider/adapter_test.go | 21 ++ .../coordination_integration_test.go | 196 ++++++++++++++-- internal/adapters/redisprovider/keys.go | 33 ++- .../redisprovider/scripts/provider.lua | 139 ++++++++++- internal/controller/pool/fetch_budget.go | 38 +-- internal/controller/pool/fetch_budget_test.go | 62 ++--- internal/controller/pool/reconciler.go | 2 - internal/controller/pool/reconciler_test.go | 13 +- internal/controller/provider/coordination.go | 11 +- internal/controller/provider/reconciler.go | 84 ++++++- .../controller/provider/reconciler_test.go | 221 ++++++++++++++++-- internal/domain/upstream/fetch_capacity.go | 2 +- internal/domain/upstream/pool.go | 11 - internal/domain/upstream/pool_test.go | 11 +- 15 files changed, 730 insertions(+), 202 deletions(-) diff --git a/internal/adapters/redisprovider/adapter.go b/internal/adapters/redisprovider/adapter.go index bfafab6..4722ae6 100644 --- a/internal/adapters/redisprovider/adapter.go +++ b/internal/adapters/redisprovider/adapter.go @@ -19,6 +19,8 @@ import ( var namespacePattern = regexp.MustCompile(`^[A-Za-z0-9._-]+$`) +const maximumLuaInteger = int64(1<<53 - 1) + type Options struct { Namespace string HolderID string @@ -55,6 +57,7 @@ func (adapter *Adapter) RunLeader( ) 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 } @@ -217,9 +220,10 @@ func (session *leaderSession) Fence() controllerProvider.Fence { return controllerProvider.Fence{Generation: session.generation, Epoch: session.epoch} } -func (session *leaderSession) AcquireFetch(ctx context.Context) (controllerProvider.RequestPermit, error) { - if ctx == nil || session == nil || session.adapter == nil || session.ctx == nil { - return nil, controllerProvider.ErrInvalidCoordination +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) @@ -229,32 +233,37 @@ func (session *leaderSession) AcquireFetch(ctx context.Context) (controllerProvi }() permitToken, err := randomToken() if err != nil { - return nil, errors.Join(controllerProvider.ErrCoordinationUnavailable, err) + 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), + durationMillis(permitTTL), expected, session.limits.MaxTotal, ) if scriptErr != nil { if session.ctx.Err() != nil { - return nil, controllerProvider.ErrLeadershipLost + return nil, false, controllerProvider.ErrLeadershipLost } if operationCtx.Err() != nil { - return nil, operationCtx.Err() + return nil, false, operationCtx.Err() } if err := wait(operationCtx, session.adapter.options.RetryInterval); err != nil { - return nil, err + return nil, false, err } continue } switch reply.Status { case "ok": - return &requestPermit{adapter: session.adapter, keys: session.keys, token: permitToken}, nil + 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, controllerProvider.ErrLeadershipLost + return nil, false, controllerProvider.ErrLeadershipLost case "rate_limited", "at_capacity": delay := time.Duration(reply.WaitMS) * time.Millisecond if delay <= 0 { @@ -262,30 +271,42 @@ func (session *leaderSession) AcquireFetch(ctx context.Context) (controllerProvi } if err := wait(operationCtx, delay); err != nil { if session.ctx.Err() != nil { - return nil, controllerProvider.ErrLeadershipLost + return nil, false, controllerProvider.ErrLeadershipLost } - return nil, err + return nil, false, err } default: - return nil, controllerProvider.ErrCoordinationUnavailable + return nil, false, controllerProvider.ErrCoordinationUnavailable } } if session.ctx.Err() != nil { - return nil, controllerProvider.ErrLeadershipLost + return nil, false, controllerProvider.ErrLeadershipLost } - return nil, operationCtx.Err() + return nil, false, operationCtx.Err() } type requestPermit struct { - adapter *Adapter - keys upstreamKeys - token string - mu sync.Mutex - done bool + adapter *Adapter + keys upstreamKeys + token string + settlementTTL time.Duration + mu sync.Mutex + done bool } -func (permit *requestPermit) Release(ctx context.Context) error { - if ctx == nil || permit == nil || permit.adapter == nil || permit.token == "" { +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() @@ -293,15 +314,22 @@ func (permit *requestPermit) Release(ctx context.Context) error { if permit.done { return nil } - reply, err := runScript(ctx, permit.adapter.client, permit.keys, "release_fetch", permit.token) - if err != nil { - return err + 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 } - if reply.Status != "ok" { - return controllerProvider.ErrCoordinationUnavailable - } - permit.done = true - return nil + return ctx.Err() } func randomToken() (string, error) { diff --git a/internal/adapters/redisprovider/adapter_test.go b/internal/adapters/redisprovider/adapter_test.go index 278b52e..8ad3de9 100644 --- a/internal/adapters/redisprovider/adapter_test.go +++ b/internal/adapters/redisprovider/adapter_test.go @@ -89,6 +89,27 @@ func TestRunLeaderRejectsInvalidCalls(t *testing.T) { } } +func TestRunLeaderRejectsNegativeFetchQuota(t *testing.T) { + t.Parallel() + client := redis.NewClient(&redis.Options{Addr: "127.0.0.1:0"}) + t.Cleanup(func() { _ = client.Close() }) + adapter, err := New(client, Options{ + Namespace: "controller", HolderID: "controller-a", LeaseTTL: 3 * time.Second, + RenewEvery: time.Second, RetryInterval: 50 * time.Millisecond, + }) + if err != nil { + t.Fatalf("New(): %v", err) + } + limits := controllerProvider.CoordinationLimits{ + MaxInFlight: 1, MaxAttemptDuration: time.Second, MaxTotal: -1, + } + err = adapter.RunLeader(context.Background(), "provider-a", limits, + func(context.Context, controllerProvider.LeaderSession) error { return nil }) + if !errors.Is(err, controllerProvider.ErrInvalidCoordination) { + t.Fatalf("RunLeader() error = %v, want ErrInvalidCoordination", err) + } +} + func TestLeaderWorkResultPrefersParentCancellation(t *testing.T) { t.Parallel() ctx, cancel := context.WithCancel(context.Background()) diff --git a/internal/adapters/redisprovider/coordination_integration_test.go b/internal/adapters/redisprovider/coordination_integration_test.go index f041571..87e74bd 100644 --- a/internal/adapters/redisprovider/coordination_integration_test.go +++ b/internal/adapters/redisprovider/coordination_integration_test.go @@ -103,36 +103,201 @@ func TestRedisLeaderSessionEnforcesGlobalIntervalAndInFlightLimit(t *testing.T) }) }() session := receiveSession(t, sessions) - first, err := session.AcquireFetch(context.Background()) - if err != nil { - t.Fatalf("first AcquireFetch(): %v", err) + first, available, err := session.AcquireFetch(context.Background(), 1) + if err != nil || !available { + t.Fatalf("first AcquireFetch() = (%v, %t, %v)", first, available, err) } startedAt := time.Now() secondResult := make(chan permitResultFixture, 1) go func() { - permit, acquireErr := session.AcquireFetch(context.Background()) - secondResult <- permitResultFixture{permit: permit, err: acquireErr} + permit, permitAvailable, acquireErr := session.AcquireFetch(context.Background(), 1) + secondResult <- permitResultFixture{permit: permit, available: permitAvailable, err: acquireErr} }() select { case result := <-secondResult: t.Fatalf("second AcquireFetch() returned before release: %+v", result) case <-time.After(100 * time.Millisecond): } - if err := first.Release(context.Background()); err != nil { - t.Fatalf("first Release(): %v", err) + if err := first.Complete(context.Background(), 1); err != nil { + t.Fatalf("first Complete(): %v", err) } result := receivePermit(t, secondResult) - if result.err != nil || result.permit == nil { - t.Fatalf("second AcquireFetch() = (%v, %v)", result.permit, result.err) + if result.err != nil || !result.available || result.permit == nil { + t.Fatalf("second AcquireFetch() = (%v, %t, %v)", result.permit, result.available, result.err) } if elapsed := time.Since(startedAt); elapsed < 200*time.Millisecond { t.Fatalf("global request interval = %s, want at least 200ms", elapsed) } - if err := result.permit.Release(context.Background()); err != nil { - t.Fatalf("second Release(): %v", err) + if err := result.permit.Cancel(context.Background()); err != nil { + t.Fatalf("second Cancel(): %v", err) } - if err := result.permit.Release(context.Background()); err != nil { - t.Fatalf("idempotent second Release(): %v", err) + if err := result.permit.Cancel(context.Background()); err != nil { + t.Fatalf("idempotent second Cancel(): %v", err) + } + cancel() + waitRunner(t, done) +} + +func TestRedisLeaderSessionPreservesFetchQuotaAcrossFailover(t *testing.T) { + fixture := newRedisFixture(t) + limits := controllerProvider.CoordinationLimits{ + MaxInFlight: 1, MaxAttemptDuration: time.Second, MaxTotal: 2, + } + + firstCtx, cancelFirst := context.WithCancel(context.Background()) + firstSessions := make(chan controllerProvider.LeaderSession, 1) + firstDone := make(chan error, 1) + go func() { + firstDone <- fixture.coordinator(t, "controller-a").RunLeader( + firstCtx, "provider-a", limits, + func(workCtx context.Context, session controllerProvider.LeaderSession) error { + firstSessions <- session + <-workCtx.Done() + return nil + }, + ) + }() + firstSession := receiveSession(t, firstSessions) + firstPermit, available, err := firstSession.AcquireFetch(context.Background(), 2) + if err != nil || !available || firstPermit == nil { + t.Fatalf("first AcquireFetch() = (%v, %t, %v)", firstPermit, available, err) + } + firstFence := firstSession.Fence() + cancelFirst() + waitRunner(t, firstDone) + if err := firstPermit.Complete(context.Background(), 1); err != nil { + t.Fatalf("Complete() after leadership loss: %v", err) + } + + secondCtx, cancelSecond := context.WithCancel(context.Background()) + defer cancelSecond() + secondSessions := make(chan controllerProvider.LeaderSession, 1) + secondDone := make(chan error, 1) + go func() { + secondDone <- fixture.coordinator(t, "controller-b").RunLeader( + secondCtx, "provider-a", limits, + func(workCtx context.Context, session controllerProvider.LeaderSession) error { + secondSessions <- session + <-workCtx.Done() + return nil + }, + ) + }() + secondSession := receiveSession(t, secondSessions) + secondFence := secondSession.Fence() + if secondFence.Generation != firstFence.Generation || secondFence.Epoch <= firstFence.Epoch { + t.Fatalf("second fence = %+v, first = %+v", secondFence, firstFence) + } + + secondPermit, available, err := secondSession.AcquireFetch(context.Background(), 1) + if err != nil || !available || secondPermit == nil { + t.Fatalf("second AcquireFetch() = (%v, %t, %v)", secondPermit, available, err) + } + if err := secondPermit.Complete(context.Background(), 1); err != nil { + t.Fatalf("second Complete(): %v", err) + } + exhaustedPermit, available, err := secondSession.AcquireFetch(context.Background(), 1) + if err != nil || available || exhaustedPermit != nil { + t.Fatalf("exhausted AcquireFetch() = (%v, %t, %v), want unavailable", exhaustedPermit, available, err) + } + cancelSecond() + waitRunner(t, secondDone) +} + +func TestRedisFetchQuotaSettlementIsIdempotentAndCancellationRefundsReservation(t *testing.T) { + fixture := newRedisFixture(t) + limits := controllerProvider.CoordinationLimits{ + MaxInFlight: 2, MaxAttemptDuration: time.Second, MaxTotal: 2, + } + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + sessions := make(chan controllerProvider.LeaderSession, 1) + done := make(chan error, 1) + go func() { + done <- fixture.coordinator(t, "controller-a").RunLeader( + ctx, "provider-a", limits, + func(workCtx context.Context, session controllerProvider.LeaderSession) error { + sessions <- session + <-workCtx.Done() + return nil + }, + ) + }() + session := receiveSession(t, sessions) + + cancelled, available, err := session.AcquireFetch(context.Background(), 2) + if err != nil || !available || cancelled == nil { + t.Fatalf("cancelled AcquireFetch() = (%v, %t, %v)", cancelled, available, err) + } + if err := cancelled.Cancel(context.Background()); err != nil { + t.Fatalf("Cancel(): %v", err) + } + if err := cancelled.Cancel(context.Background()); err != nil { + t.Fatalf("idempotent Cancel(): %v", err) + } + + completed, available, err := session.AcquireFetch(context.Background(), 2) + if err != nil || !available || completed == nil { + t.Fatalf("completed AcquireFetch() = (%v, %t, %v)", completed, available, err) + } + if err := completed.Complete(context.Background(), 1); err != nil { + t.Fatalf("Complete(): %v", err) + } + if err := completed.Complete(context.Background(), 1); err != nil { + t.Fatalf("idempotent Complete(): %v", err) + } + if err := completed.Cancel(context.Background()); err != nil { + t.Fatalf("Cancel() after Complete(): %v", err) + } + + last, available, err := session.AcquireFetch(context.Background(), 1) + if err != nil || !available || last == nil { + t.Fatalf("last AcquireFetch() = (%v, %t, %v)", last, available, err) + } + if err := last.Complete(context.Background(), 1); err != nil { + t.Fatalf("last Complete(): %v", err) + } + exhausted, available, err := session.AcquireFetch(context.Background(), 1) + if err != nil || available || exhausted != nil { + t.Fatalf("exhausted AcquireFetch() = (%v, %t, %v)", exhausted, available, err) + } + cancel() + waitRunner(t, done) +} + +func TestRedisExpiredFetchReservationIsConservativelyCharged(t *testing.T) { + fixture := newRedisFixture(t) + coordinator, err := New(fixture.client, Options{ + Namespace: fixture.namespace, HolderID: "controller-a", LeaseTTL: 600 * time.Millisecond, + RenewEvery: 150 * time.Millisecond, RetryInterval: 10 * time.Millisecond, PermitGrace: 10 * time.Millisecond, + }) + if err != nil { + t.Fatalf("New(): %v", err) + } + limits := controllerProvider.CoordinationLimits{ + MaxInFlight: 1, MaxAttemptDuration: 40 * time.Millisecond, MaxTotal: 1, + } + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + sessions := make(chan controllerProvider.LeaderSession, 1) + done := make(chan error, 1) + go func() { + done <- coordinator.RunLeader(ctx, "provider-a", limits, + func(workCtx context.Context, session controllerProvider.LeaderSession) error { + sessions <- session + <-workCtx.Done() + return nil + }) + }() + session := receiveSession(t, sessions) + abandoned, available, err := session.AcquireFetch(context.Background(), 1) + if err != nil || !available || abandoned == nil { + t.Fatalf("abandoned AcquireFetch() = (%v, %t, %v)", abandoned, available, err) + } + time.Sleep(80 * time.Millisecond) + exhausted, available, err := session.AcquireFetch(context.Background(), 1) + if err != nil || available || exhausted != nil { + t.Fatalf("post-expiry AcquireFetch() = (%v, %t, %v), want charged quota", exhausted, available, err) } cancel() waitRunner(t, done) @@ -280,8 +445,9 @@ func receivePermit(t *testing.T, values <-chan permitResultFixture) permitResult } type permitResultFixture struct { - permit controllerProvider.RequestPermit - err error + permit controllerProvider.RequestPermit + available bool + err error } func waitRunner(t *testing.T, done <-chan error) { diff --git a/internal/adapters/redisprovider/keys.go b/internal/adapters/redisprovider/keys.go index e9a7da9..afbe394 100644 --- a/internal/adapters/redisprovider/keys.go +++ b/internal/adapters/redisprovider/keys.go @@ -14,11 +14,15 @@ type keyBuilder struct { } type upstreamKeys struct { - generation string - epoch string - leader string - next string - inflight string + generation string + epoch string + leader string + next string + inflight string + fetchedTotal string + pendingTotal string + permits string + permitExpiry string } func (builder keyBuilder) forUpstream(upstreamID string) (upstreamKeys, error) { @@ -28,16 +32,23 @@ func (builder keyBuilder) forUpstream(upstreamID string) (upstreamKeys, error) { digest := digestParts(upstreamID) prefix := "pp:" + builder.namespace + ":{provider:" + digest + "}" return upstreamKeys{ - generation: prefix + ":generation", - epoch: prefix + ":epoch", - leader: prefix + ":leader", - next: prefix + ":next-request", - inflight: prefix + ":inflight", + generation: prefix + ":generation", + epoch: prefix + ":epoch", + leader: prefix + ":leader", + next: prefix + ":next-request", + inflight: prefix + ":inflight", + fetchedTotal: prefix + ":fetched-total", + pendingTotal: prefix + ":pending-total", + permits: prefix + ":permits", + permitExpiry: prefix + ":permit-expiry", }, nil } func (keys upstreamKeys) all() []string { - return []string{keys.generation, keys.epoch, keys.leader, keys.next, keys.inflight} + return []string{ + keys.generation, keys.epoch, keys.leader, keys.next, keys.inflight, + keys.fetchedTotal, keys.pendingTotal, keys.permits, keys.permitExpiry, + } } func digestParts(values ...string) string { diff --git a/internal/adapters/redisprovider/scripts/provider.lua b/internal/adapters/redisprovider/scripts/provider.lua index cb6813f..331512b 100644 --- a/internal/adapters/redisprovider/scripts/provider.lua +++ b/internal/adapters/redisprovider/scripts/provider.lua @@ -1,4 +1,5 @@ local operation = ARGV[1] +local max_safe_integer = 9007199254740991 local function now_ms() local value = redis.call('TIME') @@ -33,6 +34,63 @@ local function same_leader(value, generation, holder_id, token, epoch) value.token == token and value.epoch == tostring(epoch) end +local function read_counter(key) + local encoded = redis.call('GET', key) + if not encoded then + return 0, nil + end + local value = tonumber(encoded) + if not value or value < 0 or value ~= math.floor(value) then + return nil, 'invalid' + end + return value, nil +end + +local function read_permit(token) + local encoded = redis.call('HGET', KEYS[8], token) + if not encoded then + return nil, nil + end + local ok, value = pcall(cjson.decode, encoded) + if not ok or type(value) ~= 'table' or type(value.state) ~= 'string' or + type(value.expected) ~= 'number' or value.expected <= 0 or + value.expected ~= math.floor(value.expected) then + return nil, 'invalid' + end + return value, nil +end + +local function cleanup_expired(now) + local expired = redis.call('ZRANGEBYSCORE', KEYS[9], '-inf', now, 'LIMIT', 0, 256) + if #expired == 0 then + return nil + end + local fetched, fetched_error = read_counter(KEYS[6]) + local pending, pending_error = read_counter(KEYS[7]) + if fetched_error or pending_error then + return 'invalid' + end + for _, permit_token in ipairs(expired) do + local permit, permit_error = read_permit(permit_token) + if permit_error then + return 'invalid' + end + if permit and permit.state == 'reserved' then + if pending < permit.expected or fetched > max_safe_integer - permit.expected then + return 'invalid' + end + pending = pending - permit.expected + fetched = fetched + permit.expected + end + redis.call('HDEL', KEYS[8], permit_token) + redis.call('ZREM', KEYS[5], permit_token) + redis.call('ZREM', KEYS[9], permit_token) + end + redis.call('SET', KEYS[6], fetched) + redis.call('SET', KEYS[7], pending) + return nil +end + if operation == 'acquire_leader' then local generation_candidate = ARGV[2] local holder_id = ARGV[3] @@ -41,7 +99,10 @@ if operation == 'acquire_leader' then if not lease_ttl or lease_ttl <= 0 then return reply('invalid', '', 0, 0) end - redis.call('SET', KEYS[1], generation_candidate, 'NX') + local created = redis.call('SET', KEYS[1], generation_candidate, 'NX') + if created then + redis.call('DEL', KEYS[2], KEYS[3], KEYS[4], KEYS[5], KEYS[6], KEYS[7], KEYS[8], KEYS[9]) + end local generation = redis.call('GET', KEYS[1]) local current, current_error = read_leader() if current_error then @@ -112,6 +173,13 @@ if operation == 'acquire_fetch' then local request_interval = tonumber(ARGV[7]) local max_in_flight = tonumber(ARGV[8]) local permit_ttl = tonumber(ARGV[9]) + local expected = tonumber(ARGV[10]) + local max_total = tonumber(ARGV[11]) + if not expected or expected <= 0 or expected ~= math.floor(expected) or + expected > max_safe_integer or not max_total or max_total < 0 or + max_total > max_safe_integer or max_total ~= math.floor(max_total) then + return reply('invalid', generation, epoch or 0, 0) + end local current, current_error = read_leader() if current_error then return reply('unavailable', generation, epoch or 0, 0) @@ -120,11 +188,27 @@ if operation == 'acquire_fetch' then return reply('stale', generation, epoch or 0, 0) end local now = now_ms() - redis.call('ZREMRANGEBYSCORE', KEYS[5], '-inf', now) - local existing = redis.call('ZSCORE', KEYS[5], permit_token) - if existing then + if cleanup_expired(now) then + return reply('unavailable', generation, epoch, 0) + end + local existing, existing_error = read_permit(permit_token) + if existing_error then + return reply('unavailable', generation, epoch, 0) + end + if existing and existing.state == 'reserved' then return reply('ok', generation, epoch, 0) end + if existing then + return reply('unavailable', generation, epoch, 0) + end + local fetched, fetched_error = read_counter(KEYS[6]) + local pending, pending_error = read_counter(KEYS[7]) + if fetched_error or pending_error then + return reply('unavailable', generation, epoch, 0) + end + if max_total > 0 and fetched + pending + expected > max_total then + return reply('quota_exhausted', generation, epoch, 0) + end local next_request = redis.call('GET', KEYS[4]) if next_request and not tonumber(next_request) then return reply('unavailable', generation, epoch, 0) @@ -141,7 +225,13 @@ if operation == 'acquire_fetch' then return reply('at_capacity', generation, epoch, wait_ms) end redis.call('ZADD', KEYS[5], now + permit_ttl, permit_token) - redis.call('PEXPIRE', KEYS[5], permit_ttl + 1000) + redis.call('ZADD', KEYS[9], now + permit_ttl, permit_token) + redis.call('HSET', KEYS[8], permit_token, cjson.encode({ + version = 1, + state = 'reserved', + expected = expected + })) + redis.call('SET', KEYS[7], pending + expected) if request_interval > 0 then redis.call('SET', KEYS[4], now + request_interval, 'PX', request_interval) else @@ -150,8 +240,43 @@ if operation == 'acquire_fetch' then return reply('ok', generation, epoch, 0) end -if operation == 'release_fetch' then - redis.call('ZREM', KEYS[5], ARGV[2]) +if operation == 'complete_fetch' or operation == 'cancel_fetch' then + local permit_token = ARGV[2] + local fetched_count = tonumber(ARGV[3]) + local settlement_ttl = tonumber(ARGV[4]) + if not fetched_count or fetched_count < 0 or fetched_count ~= math.floor(fetched_count) or + fetched_count > max_safe_integer or not settlement_ttl or settlement_ttl <= 0 then + return reply('invalid', '', 0, 0) + end + local now = now_ms() + if cleanup_expired(now) then + return reply('unavailable', '', 0, 0) + end + local permit, permit_error = read_permit(permit_token) + if permit_error then + return reply('unavailable', '', 0, 0) + end + if not permit or permit.state ~= 'reserved' then + return reply('ok', '', 0, 0) + end + local fetched, fetched_error = read_counter(KEYS[6]) + local pending, pending_error = read_counter(KEYS[7]) + if fetched_error or pending_error or pending < permit.expected or + fetched > max_safe_integer - fetched_count then + return reply('unavailable', '', 0, 0) + end + pending = pending - permit.expected + if operation == 'complete_fetch' then + fetched = fetched + fetched_count + permit.state = 'completed' + else + permit.state = 'cancelled' + end + redis.call('SET', KEYS[6], fetched) + redis.call('SET', KEYS[7], pending) + redis.call('HSET', KEYS[8], permit_token, cjson.encode(permit)) + redis.call('ZREM', KEYS[5], permit_token) + redis.call('ZADD', KEYS[9], now + settlement_ttl, permit_token) return reply('ok', '', 0, 0) end diff --git a/internal/controller/pool/fetch_budget.go b/internal/controller/pool/fetch_budget.go index 619c707..92e324e 100644 --- a/internal/controller/pool/fetch_budget.go +++ b/internal/controller/pool/fetch_budget.go @@ -20,48 +20,38 @@ var ( type FetchBudgetConfig struct { UpstreamID string MaxSize int - MaxTotal int64 ExpectedPerFetch int Managed int - FetchedTotal int64 } type FetchBudgetSnapshot struct { Managed int PendingExpected int - FetchedTotal int64 } -// FetchBudget owns both current-inventory and cumulative-fetch accounting for -// one upstream. Reserving the expected response before I/O closes the race -// between concurrent provider calls. +// FetchBudget owns current-inventory accounting for one upstream. Reserving +// the expected response before I/O closes the race between concurrent calls; +// distributed cumulative quota belongs to the Provider coordination permit. type FetchBudget struct { mu sync.Mutex upstreamID string maxSize int - maxTotal int64 expected int usage FetchBudgetSnapshot } func NewFetchBudget(config FetchBudgetConfig) (*FetchBudget, error) { if config.UpstreamID == "" || config.MaxSize <= 0 || config.ExpectedPerFetch <= 0 || - config.ExpectedPerFetch > config.MaxSize || config.MaxTotal < 0 || - config.Managed < 0 || config.FetchedTotal < 0 { - return nil, ErrInvalidFetchBudget - } - if config.MaxTotal > 0 && int64(config.ExpectedPerFetch) > config.MaxTotal { + config.ExpectedPerFetch > config.MaxSize || config.Managed < 0 { return nil, ErrInvalidFetchBudget } return &FetchBudget{ upstreamID: config.UpstreamID, maxSize: config.MaxSize, - maxTotal: config.MaxTotal, expected: config.ExpectedPerFetch, usage: FetchBudgetSnapshot{ - Managed: config.Managed, - FetchedTotal: config.FetchedTotal, + Managed: config.Managed, }, }, nil } @@ -95,14 +85,7 @@ func (b *FetchBudget) FetchAllowance() int { func (b *FetchBudget) canReserveLocked() bool { poolRoom := b.maxSize - b.usage.Managed - b.usage.PendingExpected - if poolRoom < b.expected { - return false - } - if b.maxTotal > 0 { - totalRoom := b.maxTotal - b.usage.FetchedTotal - int64(b.usage.PendingExpected) - return totalRoom >= int64(b.expected) - } - return true + return poolRoom >= b.expected } func (b *FetchBudget) Snapshot() FetchBudgetSnapshot { @@ -115,8 +98,8 @@ func (b *FetchBudget) Snapshot() FetchBudgetSnapshot { } // SynchronizeManaged replaces the local current-inventory count with the -// authoritative activity-store observation. Pending requests and cumulative -// fetch usage remain owned by this budget. +// authoritative activity-store observation. Pending reservations remain owned +// by this budget. func (b *FetchBudget) SynchronizeManaged(managed int) error { if b == nil || managed < 0 { return ErrInvalidManagedSynchronization @@ -158,11 +141,11 @@ func (p *fetchPermit) Expected() int { return p.expected } -func (p *fetchPermit) Complete(fetched, retained int) error { +func (p *fetchPermit) Complete(retained int) error { if p == nil || p.budget == nil { return ErrFetchPermitFinished } - if fetched < 0 || retained < 0 || retained > fetched || retained > p.expected { + if retained < 0 || retained > p.expected { return ErrInvalidFetchCompletion } p.budget.mu.Lock() @@ -173,7 +156,6 @@ func (p *fetchPermit) Complete(fetched, retained int) error { p.finished = true p.budget.usage.PendingExpected -= p.expected p.budget.usage.Managed += retained - p.budget.usage.FetchedTotal += int64(fetched) return nil } diff --git a/internal/controller/pool/fetch_budget_test.go b/internal/controller/pool/fetch_budget_test.go index 493267d..95a00a9 100644 --- a/internal/controller/pool/fetch_budget_test.go +++ b/internal/controller/pool/fetch_budget_test.go @@ -6,14 +6,12 @@ import ( "testing" ) -func TestFetchBudgetReservesExpectedCapacityAndSeparatesCounters(t *testing.T) { +func TestFetchBudgetReservesExpectedPoolCapacity(t *testing.T) { budget, err := NewFetchBudget(FetchBudgetConfig{ UpstreamID: "provider-a", MaxSize: 10, - MaxTotal: 20, ExpectedPerFetch: 4, Managed: 2, - FetchedTotal: 3, }) if err != nil { t.Fatalf("NewFetchBudget(): %v", err) @@ -31,21 +29,21 @@ func TestFetchBudgetReservesExpectedCapacityAndSeparatesCounters(t *testing.T) { t.Fatalf("third ReserveFetch() = (_, %v, %v), want no capacity", ok, err) } - if err := first.Complete(4, 3); err != nil { + if err := first.Complete(3); err != nil { t.Fatalf("first.Complete(): %v", err) } if err := second.Cancel(); err != nil { t.Fatalf("second.Cancel(): %v", err) } usage := budget.Snapshot() - if usage.Managed != 5 || usage.PendingExpected != 0 || usage.FetchedTotal != 7 { - t.Fatalf("Snapshot() = %+v, want managed=5 pending=0 fetched=7", usage) + if usage.Managed != 5 || usage.PendingExpected != 0 { + t.Fatalf("Snapshot() = %+v, want managed=5 pending=0", usage) } if err := budget.ReleaseManaged(2); err != nil { t.Fatalf("ReleaseManaged(): %v", err) } - if usage := budget.Snapshot(); usage.Managed != 3 || usage.FetchedTotal != 7 { - t.Fatalf("Snapshot() after release = %+v, want managed=3 fetched=7", usage) + if usage := budget.Snapshot(); usage.Managed != 3 { + t.Fatalf("Snapshot() after release = %+v, want managed=3", usage) } } @@ -63,7 +61,7 @@ func TestFetchBudgetRejectsManagedCounterUnderflow(t *testing.T) { func TestFetchBudgetSynchronizesAuthoritativeManagedInventory(t *testing.T) { budget, err := NewFetchBudget(FetchBudgetConfig{ - UpstreamID: "a", MaxSize: 5, ExpectedPerFetch: 2, Managed: 3, FetchedTotal: 7, + UpstreamID: "a", MaxSize: 5, ExpectedPerFetch: 2, Managed: 3, }) if err != nil { t.Fatalf("NewFetchBudget(): %v", err) @@ -77,8 +75,8 @@ func TestFetchBudgetSynchronizesAuthoritativeManagedInventory(t *testing.T) { t.Fatalf("SynchronizeManaged() error = %v, want ErrManagedSynchronizationInFlight", err) } usage := budget.Snapshot() - if usage.Managed != 3 || usage.PendingExpected != 2 || usage.FetchedTotal != 7 { - t.Fatalf("Snapshot() = %+v, want managed=3 pending=2 fetched=7", usage) + if usage.Managed != 3 || usage.PendingExpected != 2 { + t.Fatalf("Snapshot() = %+v, want managed=3 pending=2", usage) } if err := permit.Cancel(); err != nil { t.Fatalf("Cancel(): %v", err) @@ -86,8 +84,8 @@ func TestFetchBudgetSynchronizesAuthoritativeManagedInventory(t *testing.T) { if err := budget.SynchronizeManaged(1); err != nil { t.Fatalf("SynchronizeManaged() after cancel: %v", err) } - if usage := budget.Snapshot(); usage.Managed != 1 || usage.PendingExpected != 0 || usage.FetchedTotal != 7 { - t.Fatalf("Snapshot() after synchronization = %+v, want managed=1 pending=0 fetched=7", usage) + if usage := budget.Snapshot(); usage.Managed != 1 || usage.PendingExpected != 0 { + t.Fatalf("Snapshot() after synchronization = %+v, want managed=1 pending=0", usage) } } @@ -103,35 +101,15 @@ func TestFetchBudgetRejectsNegativeManagedSynchronization(t *testing.T) { } } -func TestFetchBudgetRequiresWholeExpectedBatchToFitLimits(t *testing.T) { - tests := []struct { - name string - config FetchBudgetConfig - }{ - { - name: "pool size", - config: FetchBudgetConfig{ - UpstreamID: "a", MaxSize: 10, ExpectedPerFetch: 4, Managed: 7, - }, - }, - { - name: "cumulative total", - config: FetchBudgetConfig{ - UpstreamID: "a", MaxSize: 10, MaxTotal: 5, ExpectedPerFetch: 4, FetchedTotal: 2, - }, - }, +func TestFetchBudgetRequiresWholeExpectedBatchToFitPool(t *testing.T) { + budget, err := NewFetchBudget(FetchBudgetConfig{ + UpstreamID: "a", MaxSize: 10, ExpectedPerFetch: 4, Managed: 7, + }) + if err != nil { + t.Fatalf("NewFetchBudget(): %v", err) } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - budget, err := NewFetchBudget(tt.config) - if err != nil { - t.Fatalf("NewFetchBudget(): %v", err) - } - if _, ok, err := budget.ReserveFetch("a"); err != nil || ok { - t.Fatalf("ReserveFetch() = (_, %v, %v), want no capacity", ok, err) - } - }) + if _, ok, err := budget.ReserveFetch("a"); err != nil || ok { + t.Fatalf("ReserveFetch() = (_, %v, %v), want no capacity", ok, err) } } @@ -185,7 +163,7 @@ func TestFetchPermitRejectsDoubleFinishAndInvalidCounts(t *testing.T) { if err != nil || !ok { t.Fatalf("ReserveFetch() = (_, %v, %v), want permit", ok, err) } - if err := permit.Complete(1, 2); !errors.Is(err, ErrInvalidFetchCompletion) { + if err := permit.Complete(3); !errors.Is(err, ErrInvalidFetchCompletion) { t.Fatalf("Complete() error = %v, want ErrInvalidFetchCompletion", err) } if err := permit.Cancel(); err != nil { diff --git a/internal/controller/pool/reconciler.go b/internal/controller/pool/reconciler.go index 3674f8c..d1c8730 100644 --- a/internal/controller/pool/reconciler.go +++ b/internal/controller/pool/reconciler.go @@ -26,7 +26,6 @@ type FetchNotifier interface { type ReconcileDecision struct { AvailableSlots int64 PendingExpected int - FetchedTotal int64 FetchAllowance int EffectiveSlots int64 Triggered bool @@ -82,7 +81,6 @@ func (r *Reconciler) reconcileSnapshot(inventory InventorySnapshot, synchronizeM decision := ReconcileDecision{ AvailableSlots: inventory.AvailableSlots, PendingExpected: usage.PendingExpected, - FetchedTotal: usage.FetchedTotal, FetchAllowance: r.budget.FetchAllowance(), EffectiveSlots: saturatingAdd(inventory.AvailableSlots, pendingSlots), } diff --git a/internal/controller/pool/reconciler_test.go b/internal/controller/pool/reconciler_test.go index e836e2e..94ecc2c 100644 --- a/internal/controller/pool/reconciler_test.go +++ b/internal/controller/pool/reconciler_test.go @@ -10,7 +10,7 @@ import ( func TestPoolReconcilerSignalsOnlyWhenSlotsAreLowAndWholeFetchFits(t *testing.T) { budget, err := NewFetchBudget(FetchBudgetConfig{ - UpstreamID: "provider-a", MaxSize: 10, MaxTotal: 20, ExpectedPerFetch: 2, + UpstreamID: "provider-a", MaxSize: 10, ExpectedPerFetch: 2, }) if err != nil { t.Fatalf("NewFetchBudget(): %v", err) @@ -31,8 +31,7 @@ func TestPoolReconcilerSignalsOnlyWhenSlotsAreLowAndWholeFetchFits(t *testing.T) Proxies: []upstream.ProxyCapacity{{ State: proxyDomain.StateAvailable, ExpiresAt: now.Add(time.Minute), Max: 4, Active: 3, }}, - MaxSize: 10, - MaxTotal: 20, + MaxSize: 10, } decision := reconciler.Reconcile(now, inventory) @@ -44,9 +43,9 @@ func TestPoolReconcilerSignalsOnlyWhenSlotsAreLowAndWholeFetchFits(t *testing.T) } } -func TestPoolReconcilerUsesBudgetPendingAndCumulativeCounters(t *testing.T) { +func TestPoolReconcilerUsesPendingPoolReservation(t *testing.T) { budget, err := NewFetchBudget(FetchBudgetConfig{ - UpstreamID: "provider-a", MaxSize: 2, MaxTotal: 2, ExpectedPerFetch: 2, + UpstreamID: "provider-a", MaxSize: 2, ExpectedPerFetch: 2, }) if err != nil { t.Fatalf("NewFetchBudget(): %v", err) @@ -65,7 +64,7 @@ func TestPoolReconcilerUsesBudgetPendingAndCumulativeCounters(t *testing.T) { t.Fatalf("NewReconciler(): %v", err) } - decision := reconciler.Reconcile(time.Now(), upstream.Inventory{MaxSize: 2, MaxTotal: 2}) + decision := reconciler.Reconcile(time.Now(), upstream.Inventory{MaxSize: 2}) if decision.Triggered || decision.PendingExpected != 2 || decision.FetchAllowance != 0 { t.Fatalf("Reconcile() = %+v, want pending fetch to suppress signal", decision) } @@ -133,7 +132,7 @@ func TestPoolReconcilerPendingEstimatePausesWithoutEndingRefillEpisode(t *testin if decision := reconciler.Reconcile(now, inventoryWithSlots(now, 2)); decision.Triggered || decision.EffectiveSlots != 8 { t.Fatalf("pending Reconcile() = %+v, want paused at target estimate", decision) } - if err := permit.Complete(2, 1); err != nil { + if err := permit.Complete(1); err != nil { t.Fatalf("Complete(): %v", err) } if decision := reconciler.Reconcile(now, inventoryWithSlots(now, 5)); !decision.Triggered { diff --git a/internal/controller/provider/coordination.go b/internal/controller/provider/coordination.go index a1f80b6..b40067f 100644 --- a/internal/controller/provider/coordination.go +++ b/internal/controller/provider/coordination.go @@ -17,6 +17,7 @@ type CoordinationLimits struct { RequestInterval time.Duration MaxInFlight int MaxAttemptDuration time.Duration + MaxTotal int64 } type Fence struct { @@ -32,10 +33,14 @@ type Coordinator interface { type LeaderSession interface { Fence() Fence - AcquireFetch(context.Context) (RequestPermit, error) + AcquireFetch(context.Context, int) (RequestPermit, bool, error) } type RequestPermit interface { - // Release is idempotent. A failed release expires automatically in storage. - Release(context.Context) error + // Complete atomically releases in-flight capacity and charges the actual + // successful candidate count. Repeated calls are idempotent. + Complete(context.Context, int) error + // Cancel releases a reservation that is known not to have consumed Provider + // quota. A crashed or abandoned reservation is conservatively charged. + Cancel(context.Context) error } diff --git a/internal/controller/provider/reconciler.go b/internal/controller/provider/reconciler.go index dcd2358..11b243f 100644 --- a/internal/controller/provider/reconciler.go +++ b/internal/controller/provider/reconciler.go @@ -112,7 +112,10 @@ func (r *Reconciler) Notify() { r.signal.Notify() } -func (r *Reconciler) Run(ctx context.Context) error { +func (r *Reconciler) RunLeader(ctx context.Context, session LeaderSession) error { + if ctx == nil || session == nil { + return ErrInvalidCoordination + } var workers sync.WaitGroup defer workers.Wait() for { @@ -131,14 +134,14 @@ func (r *Reconciler) Run(ctx context.Context) error { go func() { defer workers.Done() defer func() { <-r.inFlight }() - r.reconcile(ctx) + r.reconcile(ctx, session) }() } } -func (r *Reconciler) reconcile(ctx context.Context) { +func (r *Reconciler) reconcile(ctx context.Context, session LeaderSession) { for attempt := 1; attempt <= r.config.MaxAttempts; attempt++ { - response, result, retryable, ok := r.fetchAttempt(ctx, attempt) + response, result, retryable, ok := r.fetchAttempt(ctx, session, attempt) if !ok { return } @@ -157,7 +160,11 @@ func (r *Reconciler) reconcile(ctx context.Context) { } } -func (r *Reconciler) fetchAttempt(ctx context.Context, attempt int) (FetchResponse, Result, bool, bool) { +func (r *Reconciler) fetchAttempt( + ctx context.Context, + session LeaderSession, + attempt int, +) (FetchResponse, Result, bool, bool) { if err := r.waitForRequestSlot(ctx); err != nil { return FetchResponse{}, Result{}, false, false } @@ -191,6 +198,35 @@ func (r *Reconciler) fetchAttempt(ctx context.Context, attempt int) (FetchRespon _ = permit.Cancel() } }() + requestPermit, available, err := session.AcquireFetch(ctx, permit.Expected()) + if err != nil { + if ctx.Err() != nil || errors.Is(err, ErrLeadershipLost) { + return FetchResponse{}, Result{}, false, false + } + return FetchResponse{}, Result{ + UpstreamID: r.config.UpstreamID, + Class: upstream.FetchError, + Err: fmt.Errorf("acquire distributed fetch capacity: %w", err), + Attempt: attempt, + }, false, true + } + if !available { + return FetchResponse{}, Result{}, false, false + } + if requestPermit == nil { + return FetchResponse{}, Result{ + UpstreamID: r.config.UpstreamID, + Class: upstream.FetchError, + Err: fmt.Errorf("acquire distributed fetch capacity: invalid permit"), + Attempt: attempt, + }, false, true + } + requestSettlementAttempted := false + defer func() { + if !requestSettlementAttempted { + _ = r.settleRequestPermit(requestPermit, false, 0) + } + }() callCtx := ctx cancel := func() {} @@ -200,13 +236,22 @@ func (r *Reconciler) fetchAttempt(ctx context.Context, attempt int) (FetchRespon defer cancel() response, callErr := r.ports.Adapter.Fetch(callCtx) - var parseErr, candidateErr, capacityErr error + var parseErr, candidateErr, coordinationErr, capacityErr error var validCount, newCount int - if callErr == nil { + if callErr != nil { + requestSettlementAttempted = true + coordinationErr = r.settleRequestPermit(requestPermit, false, 0) + } else { candidates, err := r.ports.Parser.Parse(callCtx, response.Body) parseErr = err validCount = len(candidates) - if parseErr == nil && validCount > 0 { + charged := validCount + if parseErr != nil { + charged = permit.Expected() + } + requestSettlementAttempted = true + coordinationErr = r.settleRequestPermit(requestPermit, true, charged) + if parseErr == nil && coordinationErr == nil && validCount > 0 { retained := candidates if expected := permit.Expected(); expected < len(retained) { retained = retained[:expected] @@ -222,12 +267,17 @@ func (r *Reconciler) fetchAttempt(ctx context.Context, attempt int) (FetchRespon newCount = upserted.Inserted } } - if callErr == nil && parseErr == nil && candidateErr == nil { - capacityErr = permit.Complete(validCount, newCount) + if callErr == nil && parseErr == nil && candidateErr == nil && coordinationErr == nil { + capacityErr = permit.Complete(newCount) permitFinished = capacityErr == nil } - resultErr := errors.Join(callErr, parseErr, candidateErr, capacityErr) - class := upstream.ClassifyFetchResult(callErr, errors.Join(parseErr, candidateErr, capacityErr), validCount, newCount) + resultErr := errors.Join(callErr, parseErr, candidateErr, coordinationErr, capacityErr) + class := upstream.ClassifyFetchResult( + callErr, + errors.Join(parseErr, candidateErr, coordinationErr, capacityErr), + validCount, + newCount, + ) return response, Result{ UpstreamID: r.config.UpstreamID, Class: class, @@ -238,6 +288,16 @@ func (r *Reconciler) fetchAttempt(ctx context.Context, attempt int) (FetchRespon }, isRetryable(callErr) || parseErr != nil, true } +func (r *Reconciler) settleRequestPermit(permit RequestPermit, complete bool, fetched int) error { + timeout := min(r.config.Timeout, 5*time.Second) + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + if complete { + return permit.Complete(ctx, fetched) + } + return permit.Cancel(ctx) +} + func isRetryable(err error) bool { if err == nil { return false diff --git a/internal/controller/provider/reconciler_test.go b/internal/controller/provider/reconciler_test.go index fb29896..4c7c143 100644 --- a/internal/controller/provider/reconciler_test.go +++ b/internal/controller/provider/reconciler_test.go @@ -89,7 +89,7 @@ func TestReconcilerCoalescesConcurrentNotifications(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) done := make(chan error, 1) - go func() { done <- reconciler.Run(ctx) }() + go func() { done <- reconciler.RunLeader(ctx, unlimitedLeaderSession{}) }() select { case got := <-result: @@ -130,7 +130,7 @@ func TestReconcilerEnforcesRequestInterval(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) done := make(chan error, 1) - go func() { done <- reconciler.Run(ctx) }() + go func() { done <- reconciler.RunLeader(ctx, unlimitedLeaderSession{}) }() reconciler.Notify() <-results @@ -183,7 +183,7 @@ func TestReconcilerRetriesErrorsWithExponentialBackoffAndJitter(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) done := make(chan error, 1) - go func() { done <- reconciler.Run(ctx) }() + go func() { done <- reconciler.RunLeader(ctx, unlimitedLeaderSession{}) }() reconciler.Notify() wantClasses := []upstream.FetchClass{ @@ -291,7 +291,7 @@ func TestReconcilerHonorsRetryAfterBeforeBackoff(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) done := make(chan error, 1) - go func() { done <- reconciler.Run(ctx) }() + go func() { done <- reconciler.RunLeader(ctx, unlimitedLeaderSession{}) }() reconciler.Notify() <-results <-results @@ -332,7 +332,7 @@ func TestReconcilerCapsRetryAfter(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) done := make(chan error, 1) - go func() { done <- reconciler.Run(ctx) }() + go func() { done <- reconciler.RunLeader(ctx, unlimitedLeaderSession{}) }() reconciler.Notify() <-results <-results @@ -417,7 +417,7 @@ func TestReconcilerDropsNotificationFanoutWhileFetchIsInFlight(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) done := make(chan error, 1) - go func() { done <- reconciler.Run(ctx) }() + go func() { done <- reconciler.RunLeader(ctx, unlimitedLeaderSession{}) }() reconciler.Notify() <-started @@ -473,8 +473,8 @@ func TestReconcilerEnforcesMaxInFlightAcrossRunConsumers(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) done := make(chan error, 2) - go func() { done <- reconciler.Run(ctx) }() - go func() { done <- reconciler.Run(ctx) }() + go func() { done <- reconciler.RunLeader(ctx, unlimitedLeaderSession{}) }() + go func() { done <- reconciler.RunLeader(ctx, unlimitedLeaderSession{}) }() reconciler.Notify() <-started @@ -519,7 +519,7 @@ func TestReconcilerUsesConfiguredMaxInFlight(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) done := make(chan error, 1) - go func() { done <- reconciler.Run(ctx) }() + go func() { done <- reconciler.RunLeader(ctx, unlimitedLeaderSession{}) }() reconciler.Notify() <-started reconciler.Notify() @@ -545,6 +545,7 @@ func TestReconcilerDoesNotRefetchWhenActivitySinkFails(t *testing.T) { clock := newFakeClock(time.Date(2026, 7, 28, 12, 0, 0, 0, time.UTC)) sleeper := &fakeSleeper{clock: clock} results := make(chan Result, 3) + globalCompleted := make(chan int, 1) var calls atomic.Int64 ports := successfulPorts(func() { calls.Add(1) }, results) ports.Activity = activitySinkFunc(func(context.Context, string, activitypool.FetchedBatch) (activitypool.UpsertResult, error) { @@ -562,13 +563,20 @@ func TestReconcilerDoesNotRefetchWhenActivitySinkFails(t *testing.T) { t.Fatalf("NewReconciler(): %v", err) } - result := runSingleReconcile(t, reconciler, results) + result := runSingleReconcile(t, reconciler, results, leaderSessionFunc( + func(context.Context, int) (RequestPermit, bool, error) { + return &recordingRequestPermit{completed: globalCompleted}, true, nil + }, + )) if result.Class != upstream.FetchError { t.Fatalf("result class = %q, want %q", result.Class, upstream.FetchError) } if got := calls.Load(); got != 1 { t.Fatalf("ProviderAdapter.Fetch() calls = %d, want 1", got) } + if got := <-globalCompleted; got != 1 { + t.Fatalf("global charged count = %d, want fetched=1", got) + } if got := sleeper.Durations(); len(got) != 0 { t.Fatalf("Sleep durations = %v, want no retry backoff", got) } @@ -623,7 +631,7 @@ func TestReconcilerDoesNotCallProviderWithoutFetchCapacity(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) done := make(chan error, 1) - go func() { done <- reconciler.Run(ctx) }() + go func() { done <- reconciler.RunLeader(ctx, unlimitedLeaderSession{}) }() reconciler.Notify() time.Sleep(20 * time.Millisecond) cancel() @@ -640,9 +648,112 @@ func TestReconcilerDoesNotCallProviderWithoutFetchCapacity(t *testing.T) { } } -func TestReconcilerCompletesFetchPermitWithFetchedAndRetainedCounts(t *testing.T) { +func TestReconcilerDoesNotCallProviderWhenDistributedQuotaIsExhausted(t *testing.T) { results := make(chan Result, 1) - completed := make(chan fetchCompletion, 1) + var calls atomic.Int64 + ports := successfulPorts(func() { calls.Add(1) }, results) + ports.Capacity = fetchCapacityFunc(func(string) (upstream.FetchPermit, bool, error) { + return &recordingFetchPermit{expected: 2}, true, nil + }) + reconciler, err := NewReconciler(Config{ + UpstreamID: "provider-a", Timeout: time.Second, MaxAttempts: 1, MaxInFlight: 1, MaxSize: 100, + }, ports) + if err != nil { + t.Fatalf("NewReconciler(): %v", err) + } + session := leaderSessionFunc(func(_ context.Context, expected int) (RequestPermit, bool, error) { + if expected != 2 { + t.Errorf("distributed expected = %d, want 2", expected) + } + return nil, false, nil + }) + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { done <- reconciler.RunLeader(ctx, session) }() + reconciler.Notify() + time.Sleep(20 * time.Millisecond) + cancel() + if err := <-done; err != nil { + t.Fatalf("RunLeader(): %v", err) + } + if got := calls.Load(); got != 0 { + t.Fatalf("ProviderAdapter.Fetch() calls = %d, want 0", got) + } + select { + case result := <-results: + t.Fatalf("unexpected fetch result: %+v", result) + default: + } +} + +func TestReconcilerChargesExpectedWhenSuccessfulResponseCannotBeParsed(t *testing.T) { + results := make(chan Result, 1) + globalCompleted := make(chan int, 1) + localCancelled := make(chan struct{}, 1) + ports := successfulPorts(func() {}, results) + ports.Parser = parserFunc(func(context.Context, []byte) ([]proxyDomain.Proxy, error) { + return nil, errors.New("invalid provider payload") + }) + ports.Capacity = fetchCapacityFunc(func(string) (upstream.FetchPermit, bool, error) { + return &recordingFetchPermit{expected: 2, cancelled: localCancelled}, true, nil + }) + reconciler, err := NewReconciler(Config{ + UpstreamID: "provider-a", Timeout: time.Second, MaxAttempts: 1, MaxInFlight: 1, MaxSize: 100, + }, ports) + if err != nil { + t.Fatalf("NewReconciler(): %v", err) + } + session := leaderSessionFunc(func(context.Context, int) (RequestPermit, bool, error) { + return &recordingRequestPermit{completed: globalCompleted}, true, nil + }) + + result := runSingleReconcile(t, reconciler, results, session) + if result.Class != upstream.FetchError { + t.Fatalf("result class = %q, want %q", result.Class, upstream.FetchError) + } + if got := <-globalCompleted; got != 2 { + t.Fatalf("global charged count = %d, want expected=2", got) + } + select { + case <-localCancelled: + default: + t.Fatal("local pool reservation was not cancelled") + } +} + +func TestReconcilerCancelsDistributedPermitWhenProviderCallFails(t *testing.T) { + results := make(chan Result, 1) + globalCancelled := make(chan struct{}, 1) + ports := successfulPorts(func() {}, results) + ports.Adapter = adapterFunc(func(context.Context) (FetchResponse, error) { + return FetchResponse{}, errors.New("provider connection failed") + }) + reconciler, err := NewReconciler(Config{ + UpstreamID: "provider-a", Timeout: time.Second, MaxAttempts: 1, MaxInFlight: 1, MaxSize: 100, + }, ports) + if err != nil { + t.Fatalf("NewReconciler(): %v", err) + } + session := leaderSessionFunc(func(context.Context, int) (RequestPermit, bool, error) { + return &recordingRequestPermit{cancelled: globalCancelled}, true, nil + }) + + result := runSingleReconcile(t, reconciler, results, session) + if result.Class != upstream.FetchError { + t.Fatalf("result class = %q, want %q", result.Class, upstream.FetchError) + } + select { + case <-globalCancelled: + default: + t.Fatal("distributed request reservation was not cancelled") + } +} + +func TestReconcilerChargesFetchedGloballyAndCompletesRetainedLocally(t *testing.T) { + results := make(chan Result, 1) + localCompleted := make(chan fetchCompletion, 1) + globalCompleted := make(chan int, 1) ports := successfulPorts(func() {}, results) ports.Parser = parserFunc(func(context.Context, []byte) ([]proxyDomain.Proxy, error) { return []proxyDomain.Proxy{{ID: "p1"}, {ID: "p2"}, {ID: "p3"}}, nil @@ -654,7 +765,7 @@ func TestReconcilerCompletesFetchPermitWithFetchedAndRetainedCounts(t *testing.T return activitypool.UpsertResult{Accepted: 2, Inserted: 1}, nil }) ports.Capacity = fetchCapacityFunc(func(string) (upstream.FetchPermit, bool, error) { - return &recordingFetchPermit{expected: 2, completed: completed}, true, nil + return &recordingFetchPermit{expected: 2, completed: localCompleted}, true, nil }) reconciler, err := NewReconciler(Config{ UpstreamID: "provider-a", Timeout: time.Second, MaxAttempts: 1, MaxInFlight: 1, MaxSize: 100, @@ -663,12 +774,22 @@ func TestReconcilerCompletesFetchPermitWithFetchedAndRetainedCounts(t *testing.T t.Fatalf("NewReconciler(): %v", err) } - result := runSingleReconcile(t, reconciler, results) + result := runSingleReconcile(t, reconciler, results, leaderSessionFunc( + func(_ context.Context, expected int) (RequestPermit, bool, error) { + if expected != 2 { + t.Errorf("distributed expected = %d, want 2", expected) + } + return &recordingRequestPermit{completed: globalCompleted}, true, nil + }, + )) if result.ValidCount != 3 || result.NewCount != 1 { t.Fatalf("result = %+v, want valid=3 new=1", result) } - if got := <-completed; got.fetched != 3 || got.retained != 1 { - t.Fatalf("fetch completion = %+v, want fetched=3 retained=1", got) + if got := <-globalCompleted; got != 3 { + t.Fatalf("global fetch completion = %d, want fetched=3", got) + } + if got := <-localCompleted; got.retained != 1 { + t.Fatalf("local fetch completion = %+v, want retained=1", got) } } @@ -699,7 +820,7 @@ func TestReconcilerRetriesParserErrorsAsFetchErrors(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) done := make(chan error, 1) - go func() { done <- reconciler.Run(ctx) }() + go func() { done <- reconciler.RunLeader(ctx, unlimitedLeaderSession{}) }() reconciler.Notify() first, second := <-results, <-results cancel() @@ -744,11 +865,23 @@ func TestNewReconcilerRejectsInvalidSchedulingConfiguration(t *testing.T) { } } -func runSingleReconcile(t *testing.T, reconciler *Reconciler, results <-chan Result) Result { +func runSingleReconcile( + t *testing.T, + reconciler *Reconciler, + results <-chan Result, + sessions ...LeaderSession, +) Result { t.Helper() + session := LeaderSession(unlimitedLeaderSession{}) + if len(sessions) > 1 { + t.Fatal("runSingleReconcile accepts at most one LeaderSession") + } + if len(sessions) == 1 { + session = sessions[0] + } ctx, cancel := context.WithCancel(context.Background()) done := make(chan error, 1) - go func() { done <- reconciler.Run(ctx) }() + go func() { done <- reconciler.RunLeader(ctx, session) }() reconciler.Notify() var result Result select { @@ -867,26 +1000,66 @@ func (unlimitedFetchCapacity) ReserveFetch(string) (upstream.FetchPermit, bool, } type fetchCompletion struct { - fetched int retained int } type recordingFetchPermit struct { expected int completed chan<- fetchCompletion + cancelled chan<- struct{} } func (p *recordingFetchPermit) Expected() int { return p.expected } -func (p *recordingFetchPermit) Complete(fetched, retained int) error { +func (p *recordingFetchPermit) Complete(retained int) error { if p.completed != nil { - p.completed <- fetchCompletion{fetched: fetched, retained: retained} + p.completed <- fetchCompletion{retained: retained} } return nil } -func (*recordingFetchPermit) Cancel() error { return nil } +func (p *recordingFetchPermit) Cancel() error { + if p.cancelled != nil { + p.cancelled <- struct{}{} + } + return nil +} type resultRecorderFunc func(Result) func (f resultRecorderFunc) Record(result Result) { f(result) } + +type unlimitedLeaderSession struct{} + +func (unlimitedLeaderSession) Fence() Fence { return Fence{Generation: "test", Epoch: 1} } + +func (unlimitedLeaderSession) AcquireFetch(context.Context, int) (RequestPermit, bool, error) { + return &recordingRequestPermit{}, true, nil +} + +type leaderSessionFunc func(context.Context, int) (RequestPermit, bool, error) + +func (leaderSessionFunc) Fence() Fence { return Fence{Generation: "test", Epoch: 1} } + +func (f leaderSessionFunc) AcquireFetch(ctx context.Context, expected int) (RequestPermit, bool, error) { + return f(ctx, expected) +} + +type recordingRequestPermit struct { + completed chan<- int + cancelled chan<- struct{} +} + +func (p *recordingRequestPermit) Complete(_ context.Context, fetched int) error { + if p.completed != nil { + p.completed <- fetched + } + return nil +} + +func (p *recordingRequestPermit) Cancel(context.Context) error { + if p.cancelled != nil { + p.cancelled <- struct{}{} + } + return nil +} diff --git a/internal/domain/upstream/fetch_capacity.go b/internal/domain/upstream/fetch_capacity.go index ab32523..4b74b3c 100644 --- a/internal/domain/upstream/fetch_capacity.go +++ b/internal/domain/upstream/fetch_capacity.go @@ -8,6 +8,6 @@ type FetchCapacity interface { type FetchPermit interface { Expected() int - Complete(fetched, retained int) error + Complete(retained int) error Cancel() error } diff --git a/internal/domain/upstream/pool.go b/internal/domain/upstream/pool.go index af7aeca..7360556 100644 --- a/internal/domain/upstream/pool.go +++ b/internal/domain/upstream/pool.go @@ -17,9 +17,7 @@ type ProxyCapacity struct { type Inventory struct { Proxies []ProxyCapacity PendingExpected int - FetchedTotal int64 MaxSize int - MaxTotal int64 } func (i Inventory) AvailableSlots(now time.Time, safetyMargin time.Duration) int64 { @@ -55,15 +53,6 @@ func (i Inventory) FetchAllowance(requested int) int { return 0 } allowed := min(requested, max(i.MaxSize-i.ManagedCount(), 0)) - if i.MaxTotal > 0 { - remaining := i.MaxTotal - i.FetchedTotal - if remaining <= 0 { - return 0 - } - if int64(allowed) > remaining { - allowed = int(remaining) - } - } return allowed } diff --git a/internal/domain/upstream/pool_test.go b/internal/domain/upstream/pool_test.go index b21f52b..d5a8893 100644 --- a/internal/domain/upstream/pool_test.go +++ b/internal/domain/upstream/pool_test.go @@ -21,7 +21,7 @@ func TestInventoryAvailableSlotsUsesOnlyAllocatableCapacity(t *testing.T) { } } -func TestInventoryFetchAllowanceSeparatesPoolAndCumulativeLimits(t *testing.T) { +func TestInventoryFetchAllowanceUsesManagedPoolCapacity(t *testing.T) { inventory := Inventory{ Proxies: []ProxyCapacity{ {State: proxyDomain.StateFetched}, @@ -32,20 +32,13 @@ func TestInventoryFetchAllowanceSeparatesPoolAndCumulativeLimits(t *testing.T) { {State: proxyDomain.StateExtracted}, }, PendingExpected: 2, - FetchedTotal: 98, MaxSize: 10, - MaxTotal: 100, } if got := inventory.ManagedCount(); got != 7 { t.Fatalf("ManagedCount() = %d, want 7", got) } - if got := inventory.FetchAllowance(10); got != 2 { - t.Fatalf("FetchAllowance() = %d, want 2 from cumulative quota", got) - } - - inventory.MaxTotal = 0 if got := inventory.FetchAllowance(10); got != 3 { - t.Fatalf("FetchAllowance() with unlimited cumulative quota = %d, want 3 from pool size", got) + t.Fatalf("FetchAllowance() = %d, want 3 from pool size", got) } }