diff --git a/doc/rfc/messagequeue-tenant-sharding.md b/doc/rfc/messagequeue-tenant-sharding.md index dea5e5b8..05873fc8 100644 --- a/doc/rfc/messagequeue-tenant-sharding.md +++ b/doc/rfc/messagequeue-tenant-sharding.md @@ -69,12 +69,12 @@ Today partition discovery runs `SELECT DISTINCT partition_key FROM queue_message The subscriber takes an explicit configured tenant list from `MQ_TENANTS`. Consumer processes reject an empty list at startup; Stovepipe also rejects ingest requests for names outside the list. Discovery becomes: ```sql -SELECT DISTINCT partition_key FROM queue_messages -WHERE tenant = ? AND topic = ? -ORDER BY partition_key +SELECT DISTINCT tenant, partition_key FROM queue_messages +WHERE tenant IN (MQ_TENANTS) AND topic = ? +ORDER BY tenant, partition_key ``` -Fair-share, orphan sweep, and idle-lease release run per `(tenant, topic)`, not across all tenants on a topic. Discovery and shutdown attempt every configured tenant and aggregate errors so one unavailable shard does not block unrelated tenants. +vtgate scatters only to shards that own those vindex values. Fair-share, orphan sweep, and idle-lease release still run per `(tenant, topic)` after grouping the result set in Go. One unavailable serving shard fails the tick for every listed tenant; the next interval retries. Poll workers stay scoped to leased `(tenant, partition_key)` rows. Discovery never uses an unscoped `WHERE topic = ?` predicate on Vitess. ## Publish diff --git a/platform/extension/messagequeue/mysql/BUILD.bazel b/platform/extension/messagequeue/mysql/BUILD.bazel index fbc1f83b..ea6bf5c3 100644 --- a/platform/extension/messagequeue/mysql/BUILD.bazel +++ b/platform/extension/messagequeue/mysql/BUILD.bazel @@ -7,6 +7,7 @@ go_library( "delivery_state_store.go", "errors.go", "identifier.go", + "inlist.go", "message_store.go", "mock_stores.go", "offset_store.go", @@ -35,6 +36,7 @@ go_test( name = "go_default_test", srcs = [ "delivery_state_store_test.go", + "inlist_test.go", "message_store_test.go", "offset_store_test.go", "partition_lease_store_test.go", diff --git a/platform/extension/messagequeue/mysql/delivery_state_store.go b/platform/extension/messagequeue/mysql/delivery_state_store.go index 8bf17fac..c3d0a63e 100644 --- a/platform/extension/messagequeue/mysql/delivery_state_store.go +++ b/platform/extension/messagequeue/mysql/delivery_state_store.go @@ -246,14 +246,13 @@ func (s *sqldeliveryStateStore) AdvanceWatermark(ctx context.Context, consumerGr } // Batch-fetch delivery state for the provided offsets. - placeholders := make([]byte, 0, len(offsets)*2-1) + placeholders, ok := inListPlaceholders(len(offsets)) + if !ok { + return currentWatermark, nil + } args := make([]interface{}, 0, 4+len(offsets)) args = append(args, tenant, consumerGroup, topic, partitionKey) - for i, offset := range offsets { - if i > 0 { - placeholders = append(placeholders, ',') - } - placeholders = append(placeholders, '?') + for _, offset := range offsets { args = append(args, offset) } @@ -261,7 +260,7 @@ func (s *sqldeliveryStateStore) AdvanceWatermark(ctx context.Context, consumerGr SELECT message_offset, acked FROM %s WHERE tenant = ? AND consumer_group = ? AND topic = ? AND partition_key = ? AND message_offset IN (%s) - `, DeliveryStateTableName, string(placeholders)), args...) + `, DeliveryStateTableName, placeholders), args...) if err != nil { return currentWatermark, fmt.Errorf("query delivery state for watermark tenant=%s topic=%s partition=%s: %w", tenant, topic, partitionKey, err) } diff --git a/platform/extension/messagequeue/mysql/inlist.go b/platform/extension/messagequeue/mysql/inlist.go new file mode 100644 index 00000000..54589ca5 --- /dev/null +++ b/platform/extension/messagequeue/mysql/inlist.go @@ -0,0 +1,32 @@ +// Copyright (c) 2026 Uber Technologies, Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package mysql + +import "strings" + +// inListPlaceholders returns "?,?,…" of length n. n < 1 is a no-op query. +func inListPlaceholders(n int) (string, bool) { + if n < 1 { + return "", false + } + return strings.Repeat(",?", n)[1:], true +} + +func appendStrings(args []any, values []string) []any { + for _, value := range values { + args = append(args, value) + } + return args +} diff --git a/platform/extension/messagequeue/mysql/inlist_test.go b/platform/extension/messagequeue/mysql/inlist_test.go new file mode 100644 index 00000000..8cc1ef95 --- /dev/null +++ b/platform/extension/messagequeue/mysql/inlist_test.go @@ -0,0 +1,39 @@ +// Copyright (c) 2026 Uber Technologies, Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package mysql + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestInListPlaceholders(t *testing.T) { + tests := []struct { + n int + want string + wantOK bool + }{ + {n: 0}, + {n: -1}, + {n: 1, want: "?", wantOK: true}, + {n: 3, want: "?,?,?", wantOK: true}, + } + for _, tt := range tests { + got, ok := inListPlaceholders(tt.n) + assert.Equal(t, tt.wantOK, ok) + assert.Equal(t, tt.want, got) + } +} diff --git a/platform/extension/messagequeue/mysql/message_store.go b/platform/extension/messagequeue/mysql/message_store.go index 981bccd6..9801dcad 100644 --- a/platform/extension/messagequeue/mysql/message_store.go +++ b/platform/extension/messagequeue/mysql/message_store.go @@ -144,22 +144,6 @@ func (s *sqlmessageStore) Insert(ctx context.Context, tenant string, topic strin return nil } -// Delete deletes a message by tenant, topic, partition key, and ID -func (s *sqlmessageStore) Delete(ctx context.Context, tenant string, topic string, partitionKey string, messageID string) (retErr error) { - op := metrics.Begin(s.scope, "delete", metrics.StorageLatencyBuckets, metrics.NewTag("topic", topic)) - defer func() { op.Complete(retErr) }() - - _, err := s.db.ExecContext(ctx, fmt.Sprintf(` - DELETE FROM %s WHERE tenant = ? AND topic = ? AND partition_key = ? AND id = ? - `, MessagesTableName), tenant, topic, partitionKey, messageID) - - if err != nil { - return fmt.Errorf("delete message tenant=%s topic=%s partition=%s message=%s: %w", tenant, topic, partitionKey, messageID, err) - } - - return nil -} - // FetchByOffset fetches messages with offset > currentOffset for a specific partition. // Messages are fetched from the immutable log; no per-message mutation occurs. func (s *sqlmessageStore) FetchByOffset(ctx context.Context, tenant string, topic string, partitionKey string, currentOffset int64, limit int) (_ []messageRow, retErr error) { diff --git a/platform/extension/messagequeue/mysql/message_store_test.go b/platform/extension/messagequeue/mysql/message_store_test.go index e726a539..3cf98275 100644 --- a/platform/extension/messagequeue/mysql/message_store_test.go +++ b/platform/extension/messagequeue/mysql/message_store_test.go @@ -131,24 +131,6 @@ func TestMessageStore_Insert(t *testing.T) { } } -func TestMessageStore_Delete(t *testing.T) { - db, mock, store := setupmessageStoreTest(t) - defer db.Close() - - ctx := context.Background() - topic := "test_topic" - partitionKey := "part1" - messageID := "msg1" - - mock.ExpectExec("DELETE FROM queue_messages"). - WithArgs(testTenant, topic, partitionKey, messageID). - WillReturnResult(sqlmock.NewResult(0, 1)) - - err := store.Delete(ctx, testTenant, topic, partitionKey, messageID) - require.NoError(t, err) - require.NoError(t, mock.ExpectationsWereMet()) -} - func TestMessageStore_FetchByOffset(t *testing.T) { db, mock, store := setupmessageStoreTest(t) defer db.Close() diff --git a/platform/extension/messagequeue/mysql/mock_stores.go b/platform/extension/messagequeue/mysql/mock_stores.go index a9be425e..b9ee7044 100644 --- a/platform/extension/messagequeue/mysql/mock_stores.go +++ b/platform/extension/messagequeue/mysql/mock_stores.go @@ -42,20 +42,6 @@ func (m *MockmessageStore) EXPECT() *MockmessageStoreMockRecorder { return m.recorder } -// Delete mocks base method. -func (m *MockmessageStore) Delete(ctx context.Context, tenant, topic, partitionKey, messageID string) error { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "Delete", ctx, tenant, topic, partitionKey, messageID) - ret0, _ := ret[0].(error) - return ret0 -} - -// Delete indicates an expected call of Delete. -func (mr *MockmessageStoreMockRecorder) Delete(ctx, tenant, topic, partitionKey, messageID any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Delete", reflect.TypeOf((*MockmessageStore)(nil).Delete), ctx, tenant, topic, partitionKey, messageID) -} - // FetchByOffset mocks base method. func (m *MockmessageStore) FetchByOffset(ctx context.Context, tenant, topic, partitionKey string, currentOffset int64, limit int) ([]messageRow, error) { m.ctrl.T.Helper() @@ -250,64 +236,63 @@ func (m *MockpartitionLeaseStore) EXPECT() *MockpartitionLeaseStoreMockRecorder return m.recorder } -// DiscoverAndAcquirePartitions mocks base method. -func (m *MockpartitionLeaseStore) DiscoverAndAcquirePartitions(ctx context.Context, tenant, topic, subscriberName, consumerGroup string, leaseDurationMs int64, maxPartitions int) (int, []string, error) { +// DiscoverPartitions mocks base method. +func (m *MockpartitionLeaseStore) DiscoverPartitions(ctx context.Context, tenants []string, topic string) (map[string][]string, error) { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "DiscoverAndAcquirePartitions", ctx, tenant, topic, subscriberName, consumerGroup, leaseDurationMs, maxPartitions) - ret0, _ := ret[0].(int) - ret1, _ := ret[1].([]string) - ret2, _ := ret[2].(error) - return ret0, ret1, ret2 + ret := m.ctrl.Call(m, "DiscoverPartitions", ctx, tenants, topic) + ret0, _ := ret[0].(map[string][]string) + ret1, _ := ret[1].(error) + return ret0, ret1 } -// DiscoverAndAcquirePartitions indicates an expected call of DiscoverAndAcquirePartitions. -func (mr *MockpartitionLeaseStoreMockRecorder) DiscoverAndAcquirePartitions(ctx, tenant, topic, subscriberName, consumerGroup, leaseDurationMs, maxPartitions any) *gomock.Call { +// DiscoverPartitions indicates an expected call of DiscoverPartitions. +func (mr *MockpartitionLeaseStoreMockRecorder) DiscoverPartitions(ctx, tenants, topic any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DiscoverAndAcquirePartitions", reflect.TypeOf((*MockpartitionLeaseStore)(nil).DiscoverAndAcquirePartitions), ctx, tenant, topic, subscriberName, consumerGroup, leaseDurationMs, maxPartitions) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DiscoverPartitions", reflect.TypeOf((*MockpartitionLeaseStore)(nil).DiscoverPartitions), ctx, tenants, topic) } -// GetAllLeases mocks base method. -func (m *MockpartitionLeaseStore) GetAllLeases(ctx context.Context, tenant, topic, consumerGroup string) ([]leaseInfo, error) { +// GetAllLeasesForTenants mocks base method. +func (m *MockpartitionLeaseStore) GetAllLeasesForTenants(ctx context.Context, tenants []string, topic, consumerGroup string) (map[string][]leaseInfo, error) { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "GetAllLeases", ctx, tenant, topic, consumerGroup) - ret0, _ := ret[0].([]leaseInfo) + ret := m.ctrl.Call(m, "GetAllLeasesForTenants", ctx, tenants, topic, consumerGroup) + ret0, _ := ret[0].(map[string][]leaseInfo) ret1, _ := ret[1].(error) return ret0, ret1 } -// GetAllLeases indicates an expected call of GetAllLeases. -func (mr *MockpartitionLeaseStoreMockRecorder) GetAllLeases(ctx, tenant, topic, consumerGroup any) *gomock.Call { +// GetAllLeasesForTenants indicates an expected call of GetAllLeasesForTenants. +func (mr *MockpartitionLeaseStoreMockRecorder) GetAllLeasesForTenants(ctx, tenants, topic, consumerGroup any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAllLeases", reflect.TypeOf((*MockpartitionLeaseStore)(nil).GetAllLeases), ctx, tenant, topic, consumerGroup) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAllLeasesForTenants", reflect.TypeOf((*MockpartitionLeaseStore)(nil).GetAllLeasesForTenants), ctx, tenants, topic, consumerGroup) } -// GetLeasedPartitions mocks base method. -func (m *MockpartitionLeaseStore) GetLeasedPartitions(ctx context.Context, tenant, topic, subscriberName, consumerGroup string) ([]string, error) { +// GetLeasedPartitionsForTenants mocks base method. +func (m *MockpartitionLeaseStore) GetLeasedPartitionsForTenants(ctx context.Context, tenants []string, topic, subscriberName, consumerGroup string) (map[string][]string, error) { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "GetLeasedPartitions", ctx, tenant, topic, subscriberName, consumerGroup) - ret0, _ := ret[0].([]string) + ret := m.ctrl.Call(m, "GetLeasedPartitionsForTenants", ctx, tenants, topic, subscriberName, consumerGroup) + ret0, _ := ret[0].(map[string][]string) ret1, _ := ret[1].(error) return ret0, ret1 } -// GetLeasedPartitions indicates an expected call of GetLeasedPartitions. -func (mr *MockpartitionLeaseStoreMockRecorder) GetLeasedPartitions(ctx, tenant, topic, subscriberName, consumerGroup any) *gomock.Call { +// GetLeasedPartitionsForTenants indicates an expected call of GetLeasedPartitionsForTenants. +func (mr *MockpartitionLeaseStoreMockRecorder) GetLeasedPartitionsForTenants(ctx, tenants, topic, subscriberName, consumerGroup any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetLeasedPartitions", reflect.TypeOf((*MockpartitionLeaseStore)(nil).GetLeasedPartitions), ctx, tenant, topic, subscriberName, consumerGroup) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetLeasedPartitionsForTenants", reflect.TypeOf((*MockpartitionLeaseStore)(nil).GetLeasedPartitionsForTenants), ctx, tenants, topic, subscriberName, consumerGroup) } -// PurgeStale mocks base method. -func (m *MockpartitionLeaseStore) PurgeStale(ctx context.Context, tenant, topic, consumerGroup string, olderThanMs int64) error { +// PurgeStaleForTenants mocks base method. +func (m *MockpartitionLeaseStore) PurgeStaleForTenants(ctx context.Context, tenants []string, topic, consumerGroup string, olderThanMs int64) error { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "PurgeStale", ctx, tenant, topic, consumerGroup, olderThanMs) + ret := m.ctrl.Call(m, "PurgeStaleForTenants", ctx, tenants, topic, consumerGroup, olderThanMs) ret0, _ := ret[0].(error) return ret0 } -// PurgeStale indicates an expected call of PurgeStale. -func (mr *MockpartitionLeaseStoreMockRecorder) PurgeStale(ctx, tenant, topic, consumerGroup, olderThanMs any) *gomock.Call { +// PurgeStaleForTenants indicates an expected call of PurgeStaleForTenants. +func (mr *MockpartitionLeaseStoreMockRecorder) PurgeStaleForTenants(ctx, tenants, topic, consumerGroup, olderThanMs any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "PurgeStale", reflect.TypeOf((*MockpartitionLeaseStore)(nil).PurgeStale), ctx, tenant, topic, consumerGroup, olderThanMs) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "PurgeStaleForTenants", reflect.TypeOf((*MockpartitionLeaseStore)(nil).PurgeStaleForTenants), ctx, tenants, topic, consumerGroup, olderThanMs) } // ReleaseLease mocks base method. @@ -324,18 +309,32 @@ func (mr *MockpartitionLeaseStoreMockRecorder) ReleaseLease(ctx, tenant, topic, return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReleaseLease", reflect.TypeOf((*MockpartitionLeaseStore)(nil).ReleaseLease), ctx, tenant, topic, partitionKey, subscriberName, consumerGroup) } -// RenewLease mocks base method. -func (m *MockpartitionLeaseStore) RenewLease(ctx context.Context, tenant, topic, partitionKey, subscriberName, consumerGroup string, leaseDurationMs int64) error { +// ReleaseOwnedLeases mocks base method. +func (m *MockpartitionLeaseStore) ReleaseOwnedLeases(ctx context.Context, tenants []string, topic, subscriberName, consumerGroup string) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ReleaseOwnedLeases", ctx, tenants, topic, subscriberName, consumerGroup) + ret0, _ := ret[0].(error) + return ret0 +} + +// ReleaseOwnedLeases indicates an expected call of ReleaseOwnedLeases. +func (mr *MockpartitionLeaseStoreMockRecorder) ReleaseOwnedLeases(ctx, tenants, topic, subscriberName, consumerGroup any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReleaseOwnedLeases", reflect.TypeOf((*MockpartitionLeaseStore)(nil).ReleaseOwnedLeases), ctx, tenants, topic, subscriberName, consumerGroup) +} + +// RenewOwnedLeases mocks base method. +func (m *MockpartitionLeaseStore) RenewOwnedLeases(ctx context.Context, tenants []string, topic, subscriberName, consumerGroup string) error { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "RenewLease", ctx, tenant, topic, partitionKey, subscriberName, consumerGroup, leaseDurationMs) + ret := m.ctrl.Call(m, "RenewOwnedLeases", ctx, tenants, topic, subscriberName, consumerGroup) ret0, _ := ret[0].(error) return ret0 } -// RenewLease indicates an expected call of RenewLease. -func (mr *MockpartitionLeaseStoreMockRecorder) RenewLease(ctx, tenant, topic, partitionKey, subscriberName, consumerGroup, leaseDurationMs any) *gomock.Call { +// RenewOwnedLeases indicates an expected call of RenewOwnedLeases. +func (mr *MockpartitionLeaseStoreMockRecorder) RenewOwnedLeases(ctx, tenants, topic, subscriberName, consumerGroup any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RenewLease", reflect.TypeOf((*MockpartitionLeaseStore)(nil).RenewLease), ctx, tenant, topic, partitionKey, subscriberName, consumerGroup, leaseDurationMs) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RenewOwnedLeases", reflect.TypeOf((*MockpartitionLeaseStore)(nil).RenewOwnedLeases), ctx, tenants, topic, subscriberName, consumerGroup) } // TryAcquireLease mocks base method. @@ -377,61 +376,61 @@ func (m *MocksubscriberHeartbeatStore) EXPECT() *MocksubscriberHeartbeatStoreMoc return m.recorder } -// ActiveSubscribers mocks base method. -func (m *MocksubscriberHeartbeatStore) ActiveSubscribers(ctx context.Context, tenant, topic, consumerGroup string, staleDurationMs int64) ([]string, error) { +// ActiveSubscribersForTenants mocks base method. +func (m *MocksubscriberHeartbeatStore) ActiveSubscribersForTenants(ctx context.Context, tenants []string, topic, consumerGroup string, staleDurationMs int64) (map[string][]string, error) { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "ActiveSubscribers", ctx, tenant, topic, consumerGroup, staleDurationMs) - ret0, _ := ret[0].([]string) + ret := m.ctrl.Call(m, "ActiveSubscribersForTenants", ctx, tenants, topic, consumerGroup, staleDurationMs) + ret0, _ := ret[0].(map[string][]string) ret1, _ := ret[1].(error) return ret0, ret1 } -// ActiveSubscribers indicates an expected call of ActiveSubscribers. -func (mr *MocksubscriberHeartbeatStoreMockRecorder) ActiveSubscribers(ctx, tenant, topic, consumerGroup, staleDurationMs any) *gomock.Call { +// ActiveSubscribersForTenants indicates an expected call of ActiveSubscribersForTenants. +func (mr *MocksubscriberHeartbeatStoreMockRecorder) ActiveSubscribersForTenants(ctx, tenants, topic, consumerGroup, staleDurationMs any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ActiveSubscribers", reflect.TypeOf((*MocksubscriberHeartbeatStore)(nil).ActiveSubscribers), ctx, tenant, topic, consumerGroup, staleDurationMs) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ActiveSubscribersForTenants", reflect.TypeOf((*MocksubscriberHeartbeatStore)(nil).ActiveSubscribersForTenants), ctx, tenants, topic, consumerGroup, staleDurationMs) } -// Deregister mocks base method. -func (m *MocksubscriberHeartbeatStore) Deregister(ctx context.Context, tenant, topic, subscriberName, consumerGroup string) error { +// DeregisterForTenants mocks base method. +func (m *MocksubscriberHeartbeatStore) DeregisterForTenants(ctx context.Context, tenants []string, topic, subscriberName, consumerGroup string) error { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "Deregister", ctx, tenant, topic, subscriberName, consumerGroup) + ret := m.ctrl.Call(m, "DeregisterForTenants", ctx, tenants, topic, subscriberName, consumerGroup) ret0, _ := ret[0].(error) return ret0 } -// Deregister indicates an expected call of Deregister. -func (mr *MocksubscriberHeartbeatStoreMockRecorder) Deregister(ctx, tenant, topic, subscriberName, consumerGroup any) *gomock.Call { +// DeregisterForTenants indicates an expected call of DeregisterForTenants. +func (mr *MocksubscriberHeartbeatStoreMockRecorder) DeregisterForTenants(ctx, tenants, topic, subscriberName, consumerGroup any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Deregister", reflect.TypeOf((*MocksubscriberHeartbeatStore)(nil).Deregister), ctx, tenant, topic, subscriberName, consumerGroup) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeregisterForTenants", reflect.TypeOf((*MocksubscriberHeartbeatStore)(nil).DeregisterForTenants), ctx, tenants, topic, subscriberName, consumerGroup) } -// Heartbeat mocks base method. -func (m *MocksubscriberHeartbeatStore) Heartbeat(ctx context.Context, tenant, topic, subscriberName, consumerGroup string) error { +// HeartbeatForTenants mocks base method. +func (m *MocksubscriberHeartbeatStore) HeartbeatForTenants(ctx context.Context, tenants []string, topic, subscriberName, consumerGroup string) error { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "Heartbeat", ctx, tenant, topic, subscriberName, consumerGroup) + ret := m.ctrl.Call(m, "HeartbeatForTenants", ctx, tenants, topic, subscriberName, consumerGroup) ret0, _ := ret[0].(error) return ret0 } -// Heartbeat indicates an expected call of Heartbeat. -func (mr *MocksubscriberHeartbeatStoreMockRecorder) Heartbeat(ctx, tenant, topic, subscriberName, consumerGroup any) *gomock.Call { +// HeartbeatForTenants indicates an expected call of HeartbeatForTenants. +func (mr *MocksubscriberHeartbeatStoreMockRecorder) HeartbeatForTenants(ctx, tenants, topic, subscriberName, consumerGroup any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Heartbeat", reflect.TypeOf((*MocksubscriberHeartbeatStore)(nil).Heartbeat), ctx, tenant, topic, subscriberName, consumerGroup) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "HeartbeatForTenants", reflect.TypeOf((*MocksubscriberHeartbeatStore)(nil).HeartbeatForTenants), ctx, tenants, topic, subscriberName, consumerGroup) } -// PurgeStale mocks base method. -func (m *MocksubscriberHeartbeatStore) PurgeStale(ctx context.Context, tenant, topic, consumerGroup string, olderThanMs int64) error { +// PurgeStaleForTenants mocks base method. +func (m *MocksubscriberHeartbeatStore) PurgeStaleForTenants(ctx context.Context, tenants []string, topic, consumerGroup string, olderThanMs int64) error { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "PurgeStale", ctx, tenant, topic, consumerGroup, olderThanMs) + ret := m.ctrl.Call(m, "PurgeStaleForTenants", ctx, tenants, topic, consumerGroup, olderThanMs) ret0, _ := ret[0].(error) return ret0 } -// PurgeStale indicates an expected call of PurgeStale. -func (mr *MocksubscriberHeartbeatStoreMockRecorder) PurgeStale(ctx, tenant, topic, consumerGroup, olderThanMs any) *gomock.Call { +// PurgeStaleForTenants indicates an expected call of PurgeStaleForTenants. +func (mr *MocksubscriberHeartbeatStoreMockRecorder) PurgeStaleForTenants(ctx, tenants, topic, consumerGroup, olderThanMs any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "PurgeStale", reflect.TypeOf((*MocksubscriberHeartbeatStore)(nil).PurgeStale), ctx, tenant, topic, consumerGroup, olderThanMs) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "PurgeStaleForTenants", reflect.TypeOf((*MocksubscriberHeartbeatStore)(nil).PurgeStaleForTenants), ctx, tenants, topic, consumerGroup, olderThanMs) } // MockdeliveryStateStore is a mock of deliveryStateStore interface. diff --git a/platform/extension/messagequeue/mysql/partition_lease_store.go b/platform/extension/messagequeue/mysql/partition_lease_store.go index 8d022c1b..1395b1ef 100644 --- a/platform/extension/messagequeue/mysql/partition_lease_store.go +++ b/platform/extension/messagequeue/mysql/partition_lease_store.go @@ -18,7 +18,6 @@ import ( "context" "database/sql" "fmt" - "sort" "time" "github.com/uber-go/tally" @@ -92,41 +91,6 @@ func (s *sqlpartitionLeaseStore) TryAcquireLease(ctx context.Context, tenant str return acquired, nil } -// RenewLease renews the lease for a partition owned by this worker -func (s *sqlpartitionLeaseStore) RenewLease(ctx context.Context, tenant string, topic string, partitionKey string, subscriberName string, consumerGroup string, leaseDurationMs int64) (retErr error) { - op := metrics.Begin(s.scope, "renew_lease", metrics.StorageLatencyBuckets, metrics.NewTag("topic", topic)) - defer func() { op.Complete(retErr) }() - - now := currentTimeMillis() - - result, err := s.db.ExecContext(ctx, fmt.Sprintf(` - UPDATE %s - SET lease_renewed_at = ? - WHERE tenant = ? AND consumer_group = ? AND topic = ? AND partition_key = ? AND leased_by = ? - `, PartitionLeasesTableName), now, tenant, consumerGroup, topic, partitionKey, subscriberName) - - if err != nil { - return fmt.Errorf("renew lease tenant=%s topic=%s partition=%s: %w", tenant, topic, partitionKey, err) - } - - rows, err := result.RowsAffected() - if err != nil { - return fmt.Errorf("check renewal result tenant=%s topic=%s partition=%s: %w", tenant, topic, partitionKey, err) - } - - if rows == 0 { - return &ErrLeaseExpired{Topic: topic, PartitionKey: partitionKey} - } - - s.logger.Debugw("renewed lease", - logTenant, tenant, - logTopic, topic, - logPartitionKey, partitionKey, - ) - - return nil -} - // ReleaseLease releases the lease for a partition owned by this worker func (s *sqlpartitionLeaseStore) ReleaseLease(ctx context.Context, tenant string, topic string, partitionKey string, subscriberName string, consumerGroup string) (retErr error) { op := metrics.Begin(s.scope, "release_lease", metrics.StorageLatencyBuckets, metrics.NewTag("topic", topic)) @@ -164,235 +128,178 @@ func (s *sqlpartitionLeaseStore) ReleaseLease(ctx context.Context, tenant string return nil } -// GetLeasedPartitions returns all partitions currently leased by this worker -func (s *sqlpartitionLeaseStore) GetLeasedPartitions(ctx context.Context, tenant string, topic string, subscriberName string, consumerGroup string) (_ []string, retErr error) { - op := metrics.Begin(s.scope, "get_leased_partitions", metrics.StorageLatencyBuckets, metrics.NewTag("topic", topic)) +func (s *sqlpartitionLeaseStore) GetLeasedPartitionsForTenants(ctx context.Context, tenants []string, topic string, subscriberName string, consumerGroup string) (_ map[string][]string, retErr error) { + op := metrics.Begin(s.scope, "get_leased_partitions_for_tenants", metrics.StorageLatencyBuckets, metrics.NewTag("topic", topic)) defer func() { op.Complete(retErr) }() - rows, err := s.db.QueryContext(ctx, fmt.Sprintf(` - SELECT partition_key FROM %s - WHERE tenant = ? AND consumer_group = ? AND topic = ? AND leased_by = ? - `, PartitionLeasesTableName), tenant, consumerGroup, topic, subscriberName) + placeholders, ok := inListPlaceholders(len(tenants)) + if !ok { + return map[string][]string{}, nil + } + args := appendStrings(nil, tenants) + args = append(args, consumerGroup, topic, subscriberName) + rows, err := s.db.QueryContext(ctx, fmt.Sprintf(` + SELECT tenant, partition_key FROM %s + WHERE tenant IN (%s) AND consumer_group = ? AND topic = ? AND leased_by = ? + `, PartitionLeasesTableName, placeholders), args...) if err != nil { - return nil, fmt.Errorf("get leased partitions tenant=%s topic=%s: %w", tenant, topic, err) + return nil, fmt.Errorf("get leased partitions topic=%s: %w", topic, err) } defer rows.Close() - var partitions []string + byTenant := make(map[string][]string) for rows.Next() { - var partition string - if err := rows.Scan(&partition); err != nil { - return nil, fmt.Errorf("scan partition tenant=%s topic=%s: %w", tenant, topic, err) + var tenant, partition string + if err := rows.Scan(&tenant, &partition); err != nil { + return nil, fmt.Errorf("scan leased partition topic=%s: %w", topic, err) } - partitions = append(partitions, partition) + byTenant[tenant] = append(byTenant[tenant], partition) } - if err := rows.Err(); err != nil { - return nil, fmt.Errorf("row iteration tenant=%s topic=%s: %w", tenant, topic, err) + return nil, fmt.Errorf("row iteration topic=%s: %w", topic, err) } - - s.logger.Debugw("retrieved leased partitions", - logTenant, tenant, - logTopic, topic, - "count", len(partitions), - ) - - return partitions, nil + return byTenant, nil } -// GetAllLeases returns the lease row for every partition currently leased -// under (tenant, topic, consumerGroup) by any subscriber. -func (s *sqlpartitionLeaseStore) GetAllLeases(ctx context.Context, tenant string, topic string, consumerGroup string) (_ []leaseInfo, retErr error) { - op := metrics.Begin(s.scope, "get_all_leases", metrics.StorageLatencyBuckets, metrics.NewTag("topic", topic)) +func (s *sqlpartitionLeaseStore) GetAllLeasesForTenants(ctx context.Context, tenants []string, topic string, consumerGroup string) (_ map[string][]leaseInfo, retErr error) { + op := metrics.Begin(s.scope, "get_all_leases_for_tenants", metrics.StorageLatencyBuckets, metrics.NewTag("topic", topic)) defer func() { op.Complete(retErr) }() - rows, err := s.db.QueryContext(ctx, fmt.Sprintf(` - SELECT partition_key, leased_by, lease_renewed_at FROM %s - WHERE tenant = ? AND consumer_group = ? AND topic = ? - `, PartitionLeasesTableName), tenant, consumerGroup, topic) + placeholders, ok := inListPlaceholders(len(tenants)) + if !ok { + return map[string][]leaseInfo{}, nil + } + args := appendStrings(nil, tenants) + args = append(args, consumerGroup, topic) + rows, err := s.db.QueryContext(ctx, fmt.Sprintf(` + SELECT tenant, partition_key, leased_by, lease_renewed_at FROM %s + WHERE tenant IN (%s) AND consumer_group = ? AND topic = ? + `, PartitionLeasesTableName, placeholders), args...) if err != nil { - return nil, fmt.Errorf("get all leases tenant=%s topic=%s: %w", tenant, topic, err) + return nil, fmt.Errorf("get all leases topic=%s: %w", topic, err) } defer rows.Close() - var leases []leaseInfo + byTenant := make(map[string][]leaseInfo) for rows.Next() { + var tenant string var lease leaseInfo - if err := rows.Scan(&lease.PartitionKey, &lease.LeasedBy, &lease.LeaseRenewedAt); err != nil { - return nil, fmt.Errorf("scan lease tenant=%s topic=%s: %w", tenant, topic, err) + if err := rows.Scan(&tenant, &lease.PartitionKey, &lease.LeasedBy, &lease.LeaseRenewedAt); err != nil { + return nil, fmt.Errorf("scan lease topic=%s: %w", topic, err) } - leases = append(leases, lease) + byTenant[tenant] = append(byTenant[tenant], lease) } - if err := rows.Err(); err != nil { - return nil, fmt.Errorf("row iteration tenant=%s topic=%s: %w", tenant, topic, err) + return nil, fmt.Errorf("row iteration topic=%s: %w", topic, err) } + return byTenant, nil +} + +func (s *sqlpartitionLeaseStore) RenewOwnedLeases(ctx context.Context, tenants []string, topic string, subscriberName string, consumerGroup string) (retErr error) { + op := metrics.Begin(s.scope, "renew_owned_leases", metrics.StorageLatencyBuckets, metrics.NewTag("topic", topic)) + defer func() { op.Complete(retErr) }() + + placeholders, ok := inListPlaceholders(len(tenants)) + if !ok { + return nil + } + now := currentTimeMillis() + args := []any{now} + args = appendStrings(args, tenants) + args = append(args, consumerGroup, topic, subscriberName) - return leases, nil + _, err := s.db.ExecContext(ctx, fmt.Sprintf(` + UPDATE %s + SET lease_renewed_at = ? + WHERE tenant IN (%s) AND consumer_group = ? AND topic = ? AND leased_by = ? + `, PartitionLeasesTableName, placeholders), args...) + if err != nil { + return fmt.Errorf("renew owned leases topic=%s: %w", topic, err) + } + return nil } -// PurgeStale deletes lease rows not renewed within olderThanMs. See the -// partitionLeaseStore interface doc. -func (s *sqlpartitionLeaseStore) PurgeStale(ctx context.Context, tenant string, topic string, consumerGroup string, olderThanMs int64) (retErr error) { - op := metrics.Begin(s.scope, "purge_stale", metrics.StorageLatencyBuckets, metrics.NewTag("topic", topic)) +func (s *sqlpartitionLeaseStore) ReleaseOwnedLeases(ctx context.Context, tenants []string, topic string, subscriberName string, consumerGroup string) (retErr error) { + op := metrics.Begin(s.scope, "release_owned_leases", metrics.StorageLatencyBuckets, metrics.NewTag("topic", topic)) defer func() { op.Complete(retErr) }() + placeholders, ok := inListPlaceholders(len(tenants)) + if !ok { + return nil + } + args := appendStrings(nil, tenants) + args = append(args, consumerGroup, topic, subscriberName) + + _, err := s.db.ExecContext(ctx, fmt.Sprintf(` + DELETE FROM %s + WHERE tenant IN (%s) AND consumer_group = ? AND topic = ? AND leased_by = ? + `, PartitionLeasesTableName, placeholders), args...) + if err != nil { + return fmt.Errorf("release owned leases topic=%s: %w", topic, err) + } + return nil +} + +func (s *sqlpartitionLeaseStore) PurgeStaleForTenants(ctx context.Context, tenants []string, topic string, consumerGroup string, olderThanMs int64) (retErr error) { + op := metrics.Begin(s.scope, "purge_stale_for_tenants", metrics.StorageLatencyBuckets, metrics.NewTag("topic", topic)) + defer func() { op.Complete(retErr) }() + + placeholders, ok := inListPlaceholders(len(tenants)) + if !ok { + return nil + } threshold := currentTimeMillis() - olderThanMs + args := appendStrings(nil, tenants) + args = append(args, consumerGroup, topic, threshold) result, err := s.db.ExecContext(ctx, fmt.Sprintf(` DELETE FROM %s - WHERE tenant = ? AND consumer_group = ? AND topic = ? AND lease_renewed_at < ? - `, PartitionLeasesTableName), tenant, consumerGroup, topic, threshold) - + WHERE tenant IN (%s) AND consumer_group = ? AND topic = ? AND lease_renewed_at < ? + `, PartitionLeasesTableName, placeholders), args...) if err != nil { - return fmt.Errorf("failed to purge stale leases tenant=%s topic=%s: %w", tenant, topic, err) + return fmt.Errorf("failed to purge stale leases topic=%s: %w", topic, err) } - - // RowsAffected error is swallowed because the DELETE itself succeeded; - // the count is for observability only. if deleted, err := result.RowsAffected(); err == nil && deleted > 0 { - metrics.NamedCounter(s.scope, "purge_stale", "rows_deleted", deleted, metrics.NewTag("topic", topic)) - s.logger.Debugw("purged stale leases", - logTenant, tenant, - logTopic, topic, - "deleted", deleted, - ) + metrics.NamedCounter(s.scope, "purge_stale_for_tenants", "rows_deleted", deleted, metrics.NewTag("topic", topic)) } - return nil } -// DiscoverAndAcquirePartitions discovers partitions from messages table and tries to acquire leases. -// Returns the number of new leases acquired and the full list of discovered partitions. -// maxPartitions limits how many total partitions this subscriber can own (0 = unlimited) -// -// Acquisition is lease-aware: one GetAllLeases read classifies every -// discovered partition, and TryAcquireLease is attempted only for partitions -// that are unleased or whose lease is stale (stealable). Partitions already -// owned by this subscriber are counted against the cap but not re-probed -// (renewal is the lease tick's job), and partitions validly held by another -// subscriber are skipped entirely — probing them is a guaranteed-futile -// write on a contended lease row. The classification is advisory (a lease -// can expire or renew between the read and the attempt); TryAcquireLease -// remains the atomic arbiter. -func (s *sqlpartitionLeaseStore) DiscoverAndAcquirePartitions(ctx context.Context, tenant string, topic string, subscriberName string, consumerGroup string, leaseDurationMs int64, maxPartitions int) (_ int, _ []string, retErr error) { - op := metrics.Begin(s.scope, "discover_and_acquire", metrics.StorageLatencyBuckets, metrics.NewTag("topic", topic)) +func (s *sqlpartitionLeaseStore) DiscoverPartitions(ctx context.Context, tenants []string, topic string) (_ map[string][]string, retErr error) { + op := metrics.Begin(s.scope, "discover_partitions", metrics.StorageLatencyBuckets, metrics.NewTag("topic", topic)) defer func() { op.Complete(retErr) }() - // Query distinct partition_keys from messages table. - // No LIMIT is applied because all partitions must be discoverable for fair - // share computation to be accurate — a LIMIT would silently exclude partitions, - // making them permanently unprocessable. The maxPartitions cap only limits how - // many leases this subscriber acquires, not how many partitions are visible. + placeholders, ok := inListPlaceholders(len(tenants)) + if !ok { + return map[string][]string{}, nil + } + args := appendStrings(nil, tenants) + args = append(args, topic) + rows, err := s.db.QueryContext(ctx, fmt.Sprintf(` - SELECT DISTINCT partition_key FROM %s WHERE tenant = ? AND topic = ? ORDER BY partition_key - `, MessagesTableName), tenant, topic) + SELECT DISTINCT tenant, partition_key FROM %s + WHERE tenant IN (%s) AND topic = ? + ORDER BY tenant, partition_key + `, MessagesTableName, placeholders), args...) if err != nil { - return 0, nil, fmt.Errorf("discover partitions tenant=%s topic=%s: %w", tenant, topic, err) + return nil, fmt.Errorf("discover partitions topic=%s: %w", topic, err) } defer rows.Close() - var partitions []string + byTenant := make(map[string][]string) for rows.Next() { - var partitionKey string - if err := rows.Scan(&partitionKey); err != nil { - return 0, nil, fmt.Errorf("scan partition key tenant=%s topic=%s: %w", tenant, topic, err) + var tenant, partitionKey string + if err := rows.Scan(&tenant, &partitionKey); err != nil { + return nil, fmt.Errorf("scan partition key topic=%s: %w", topic, err) } - partitions = append(partitions, partitionKey) + byTenant[tenant] = append(byTenant[tenant], partitionKey) } - if err := rows.Err(); err != nil { - return 0, nil, fmt.Errorf("row iteration tenant=%s topic=%s: %w", tenant, topic, err) - } - - s.logger.Debugw("discovered partitions", - logTenant, tenant, - logTopic, topic, - "count", len(partitions), - ) - - // One read of every lease row classifies the discovered partitions: - // self-owned (count toward the cap, no re-probe), validly held by - // another subscriber (skip), or unleased/stale (acquisition candidates). - allLeases, err := s.GetAllLeases(ctx, tenant, topic, consumerGroup) - if err != nil { - return 0, nil, fmt.Errorf("get all leases for acquisition tenant=%s topic=%s: %w", tenant, topic, err) - } - staleThreshold := currentTimeMillis() - leaseDurationMs - ownedCount := 0 - ownedSet := make(map[string]struct{}) - heldByOther := make(map[string]struct{}) - for _, lease := range allLeases { - switch { - case lease.LeasedBy == subscriberName: - // Self-owned, fresh or stale: a stale self-lease means our own - // renewals are lagging, not that ownership moved. - ownedSet[lease.PartitionKey] = struct{}{} - ownedCount++ - case lease.LeaseRenewedAt >= staleThreshold: - heldByOther[lease.PartitionKey] = struct{}{} - } - } - - // Sort partitions deterministically - sort.Strings(partitions) - - // Try to acquire leases for unleased or stale discovered partitions - acquiredCount := 0 - skippedCount := 0 - for _, partitionKey := range partitions { - if _, owned := ownedSet[partitionKey]; owned { - continue - } - if _, held := heldByOther[partitionKey]; held { - skippedCount++ - continue - } - - // Enforce maxPartitions cap using local count - if maxPartitions > 0 && ownedCount >= maxPartitions { - s.logger.Debugw("reached max partitions cap, stopping acquisition", - logTenant, tenant, - logTopic, topic, - "max_partitions", maxPartitions, - "owned_count", ownedCount, - ) - break - } - - acquired, err := s.TryAcquireLease(ctx, tenant, topic, partitionKey, subscriberName, consumerGroup, leaseDurationMs) - if err != nil { - // Per-partition error is swallowed because one partition's DB failure - // should not prevent acquiring leases for other partitions. The failed - // partition is retried on the next discovery cycle. - s.logger.Errorw("failed to acquire lease for partition", - logTenant, tenant, - logTopic, topic, - logPartitionKey, partitionKey, - logError, err, - ) - continue - } - if acquired { - acquiredCount++ - ownedCount++ - } + return nil, fmt.Errorf("row iteration topic=%s: %w", topic, err) } - - metrics.NamedCounter(s.scope, "discover_and_acquire", "partitions_discovered", int64(len(partitions)), metrics.NewTag("topic", topic)) - metrics.NamedCounter(s.scope, "discover_and_acquire", "partitions_acquired", int64(acquiredCount), metrics.NewTag("topic", topic)) - metrics.NamedCounter(s.scope, "discover_and_acquire", "lease_aware_skipped", int64(skippedCount), metrics.NewTag("topic", topic)) - s.logger.Debugw("completed partition discovery and acquisition", - logTenant, tenant, - logTopic, topic, - "discovered_count", len(partitions), - "acquired_count", acquiredCount, - "skipped_held_by_other", skippedCount, - ) - - return acquiredCount, partitions, nil + return byTenant, nil } // currentTimeMillis returns the current time in milliseconds since epoch. diff --git a/platform/extension/messagequeue/mysql/partition_lease_store_test.go b/platform/extension/messagequeue/mysql/partition_lease_store_test.go index 17f27804..7ea852ba 100644 --- a/platform/extension/messagequeue/mysql/partition_lease_store_test.go +++ b/platform/extension/messagequeue/mysql/partition_lease_store_test.go @@ -17,9 +17,7 @@ package mysql import ( "context" "database/sql" - "fmt" "testing" - "time" "github.com/DATA-DOG/go-sqlmock" "github.com/stretchr/testify/require" @@ -99,54 +97,6 @@ func TestPartitionLeaseStore_TryAcquireLease(t *testing.T) { } } -func TestPartitionLeaseStore_RenewLease(t *testing.T) { - tests := []struct { - name string - setup func(mock sqlmock.Sqlmock) - wantErr bool - }{ - { - name: "successfully renew lease", - setup: func(mock sqlmock.Sqlmock) { - mock.ExpectExec("UPDATE queue_partition_leases"). - WithArgs(sqlmock.AnyArg(), testTenant, testConsumerGroup, "test_topic", "part1", testSubscriberName). - WillReturnResult(sqlmock.NewResult(0, 1)) - }, - wantErr: false, - }, - { - name: "lease not owned", - setup: func(mock sqlmock.Sqlmock) { - mock.ExpectExec("UPDATE queue_partition_leases"). - WithArgs(sqlmock.AnyArg(), testTenant, testConsumerGroup, "test_topic", "part1", testSubscriberName). - WillReturnResult(sqlmock.NewResult(0, 0)) - }, - wantErr: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - db, mock, store := setuppartitionLeaseStoreTest(t) - defer db.Close() - - ctx := context.Background() - topic := "test_topic" - partitionKey := "part1" - - tt.setup(mock) - - err := store.RenewLease(ctx, testTenant, topic, partitionKey, testSubscriberName, testConsumerGroup, testLeaseDurationMs) - if tt.wantErr { - require.Error(t, err) - } else { - require.NoError(t, err) - } - require.NoError(t, mock.ExpectationsWereMet()) - }) - } -} - func TestPartitionLeaseStore_ReleaseLease(t *testing.T) { tests := []struct { name string @@ -195,272 +145,94 @@ func TestPartitionLeaseStore_ReleaseLease(t *testing.T) { } } -func TestPartitionLeaseStore_GetLeasedPartitions(t *testing.T) { +func TestPartitionLeaseStore_DiscoverPartitionsForTenants(t *testing.T) { db, mock, store := setuppartitionLeaseStoreTest(t) defer db.Close() - ctx := context.Background() - topic := "test_topic" + tenants := []string{"alpha", "beta"} + mock.ExpectQuery("SELECT DISTINCT tenant, partition_key FROM queue_messages"). + WithArgs("alpha", "beta", "test_topic"). + WillReturnRows(sqlmock.NewRows([]string{"tenant", "partition_key"}). + AddRow("alpha", "p1"). + AddRow("beta", "p2"). + AddRow("beta", "p3")) - rows := sqlmock.NewRows([]string{"partition_key"}). - AddRow("part1"). - AddRow("part2"). - AddRow("part3") - - mock.ExpectQuery("SELECT partition_key FROM queue_partition_leases"). - WithArgs(testTenant, testConsumerGroup, topic, testSubscriberName). - WillReturnRows(rows) + got, err := store.DiscoverPartitions(context.Background(), tenants, "test_topic") + require.NoError(t, err) + require.Equal(t, map[string][]string{ + "alpha": {"p1"}, + "beta": {"p2", "p3"}, + }, got) - partitions, err := store.GetLeasedPartitions(ctx, testTenant, topic, testSubscriberName, testConsumerGroup) + empty, err := store.DiscoverPartitions(context.Background(), nil, "test_topic") require.NoError(t, err) - require.Len(t, partitions, 3) - require.Equal(t, []string{"part1", "part2", "part3"}, partitions) + require.Empty(t, empty) require.NoError(t, mock.ExpectationsWereMet()) } -func TestPartitionLeaseStore_GetAllLeases(t *testing.T) { - tests := []struct { - name string - setup func(mock sqlmock.Sqlmock) - want []leaseInfo - }{ - { - name: "returns leases held by any subscriber", - setup: func(mock sqlmock.Sqlmock) { - rows := sqlmock.NewRows([]string{"partition_key", "leased_by", "lease_renewed_at"}). - AddRow("part1", testSubscriberName, int64(1000)). - AddRow("part2", "other-worker", int64(2000)) - mock.ExpectQuery("SELECT partition_key, leased_by, lease_renewed_at FROM queue_partition_leases"). - WithArgs(testTenant, testConsumerGroup, "test_topic"). - WillReturnRows(rows) - }, - want: []leaseInfo{ - {PartitionKey: "part1", LeasedBy: testSubscriberName, LeaseRenewedAt: 1000}, - {PartitionKey: "part2", LeasedBy: "other-worker", LeaseRenewedAt: 2000}, - }, - }, - { - name: "no leases returns empty", - setup: func(mock sqlmock.Sqlmock) { - mock.ExpectQuery("SELECT partition_key, leased_by, lease_renewed_at FROM queue_partition_leases"). - WithArgs(testTenant, testConsumerGroup, "test_topic"). - WillReturnRows(sqlmock.NewRows([]string{"partition_key", "leased_by", "lease_renewed_at"})) - }, - want: nil, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - db, mock, store := setuppartitionLeaseStoreTest(t) - defer db.Close() +func TestPartitionLeaseStore_GetLeasedPartitionsForTenants(t *testing.T) { + db, mock, store := setuppartitionLeaseStoreTest(t) + defer db.Close() - tt.setup(mock) + tenants := []string{"alpha", "gamma"} + mock.ExpectQuery("SELECT tenant, partition_key FROM queue_partition_leases"). + WithArgs("alpha", "gamma", testConsumerGroup, "test_topic", testSubscriberName). + WillReturnRows(sqlmock.NewRows([]string{"tenant", "partition_key"}). + AddRow("alpha", "p1")) - leases, err := store.GetAllLeases(context.Background(), testTenant, "test_topic", testConsumerGroup) - require.NoError(t, err) - require.Equal(t, tt.want, leases) - require.NoError(t, mock.ExpectationsWereMet()) - }) - } + got, err := store.GetLeasedPartitionsForTenants(context.Background(), tenants, "test_topic", testSubscriberName, testConsumerGroup) + require.NoError(t, err) + require.Equal(t, map[string][]string{"alpha": {"p1"}}, got) + require.NoError(t, mock.ExpectationsWereMet()) } -func TestPartitionLeaseStore_DiscoverAndAcquirePartitions(t *testing.T) { - leaseColumns := []string{"partition_key", "leased_by", "lease_renewed_at"} - freshMs := time.Now().UnixMilli() - staleMs := freshMs - testLeaseDurationMs - 60_000 - - // expectDiscover mocks the DISTINCT partition scan. - expectDiscover := func(mock sqlmock.Sqlmock, partitions ...string) { - rows := sqlmock.NewRows([]string{"partition_key"}) - for _, pk := range partitions { - rows.AddRow(pk) - } - mock.ExpectQuery("SELECT DISTINCT partition_key FROM queue_messages"). - WithArgs(testTenant, "test_topic"). - WillReturnRows(rows) - } +func TestPartitionLeaseStore_GetAllLeasesForTenants(t *testing.T) { + db, mock, store := setuppartitionLeaseStoreTest(t) + defer db.Close() - // expectAcquire mocks one TryAcquireLease attempt whose ownership check - // reports the given owner. - expectAcquire := func(mock sqlmock.Sqlmock, owner string) { - mock.ExpectExec("INSERT INTO queue_partition_leases"). - WillReturnResult(sqlmock.NewResult(1, 1)) - mock.ExpectQuery("SELECT leased_by FROM queue_partition_leases"). - WillReturnRows(sqlmock.NewRows([]string{"leased_by"}).AddRow(owner)) - } + tenants := []string{"alpha", "beta", "gamma"} + mock.ExpectQuery("SELECT tenant, partition_key, leased_by, lease_renewed_at FROM queue_partition_leases"). + WithArgs("alpha", "beta", "gamma", testConsumerGroup, "test_topic"). + WillReturnRows(sqlmock.NewRows([]string{"tenant", "partition_key", "leased_by", "lease_renewed_at"}). + AddRow("alpha", "p1", testSubscriberName, int64(1000)). + AddRow("beta", "p2", "other-worker", int64(2000)). + AddRow("beta", "p3", testSubscriberName, int64(3000))) - tests := []struct { - name string - maxPartitions int - setup func(mock sqlmock.Sqlmock) - wantAcquired int - }{ - { - name: "acquires unleased, skips fresh lease held by other", - maxPartitions: 0, - setup: func(mock sqlmock.Sqlmock) { - expectDiscover(mock, "part1", "part2") - mock.ExpectQuery("SELECT partition_key, leased_by, lease_renewed_at FROM queue_partition_leases"). - WithArgs(testTenant, testConsumerGroup, "test_topic"). - WillReturnRows(sqlmock.NewRows(leaseColumns). - AddRow("part2", "other-worker", freshMs)) - // Only unleased part1 is attempted; part2's fresh lease is - // never write-probed. - expectAcquire(mock, testSubscriberName) - }, - wantAcquired: 1, - }, - { - name: "stale lease held by other is stealable", - maxPartitions: 0, - setup: func(mock sqlmock.Sqlmock) { - expectDiscover(mock, "part1") - mock.ExpectQuery("SELECT partition_key, leased_by, lease_renewed_at FROM queue_partition_leases"). - WithArgs(testTenant, testConsumerGroup, "test_topic"). - WillReturnRows(sqlmock.NewRows(leaseColumns). - AddRow("part1", "other-worker", staleMs)) - expectAcquire(mock, testSubscriberName) - }, - wantAcquired: 1, - }, - { - name: "self-owned partitions are not re-probed", - maxPartitions: 0, - setup: func(mock sqlmock.Sqlmock) { - expectDiscover(mock, "part1", "part2") - mock.ExpectQuery("SELECT partition_key, leased_by, lease_renewed_at FROM queue_partition_leases"). - WithArgs(testTenant, testConsumerGroup, "test_topic"). - WillReturnRows(sqlmock.NewRows(leaseColumns). - AddRow("part1", testSubscriberName, freshMs)) - // Only part2 is attempted; renewal of part1 is the lease - // tick's job. - expectAcquire(mock, testSubscriberName) - }, - wantAcquired: 1, - }, - { - name: "stops acquiring when cap reached", - maxPartitions: 2, - setup: func(mock sqlmock.Sqlmock) { - expectDiscover(mock, "part1", "part2", "part3") - mock.ExpectQuery("SELECT partition_key, leased_by, lease_renewed_at FROM queue_partition_leases"). - WithArgs(testTenant, testConsumerGroup, "test_topic"). - WillReturnRows(sqlmock.NewRows(leaseColumns)) - // part1 and part2 acquired; part3 never attempted at the cap. - expectAcquire(mock, testSubscriberName) - expectAcquire(mock, testSubscriberName) - }, - wantAcquired: 2, - }, - { - name: "pre-owned partitions count toward cap", - maxPartitions: 3, - setup: func(mock sqlmock.Sqlmock) { - expectDiscover(mock, "part1", "part2", "part3") - mock.ExpectQuery("SELECT partition_key, leased_by, lease_renewed_at FROM queue_partition_leases"). - WithArgs(testTenant, testConsumerGroup, "test_topic"). - WillReturnRows(sqlmock.NewRows(leaseColumns). - AddRow("existing1", testSubscriberName, freshMs). - AddRow("existing2", testSubscriberName, freshMs)) - // One acquisition reaches the cap of 3; part2/part3 skipped. - expectAcquire(mock, testSubscriberName) - }, - wantAcquired: 1, - }, - { - name: "already at cap acquires nothing", - maxPartitions: 2, - setup: func(mock sqlmock.Sqlmock) { - expectDiscover(mock, "part1", "part2") - mock.ExpectQuery("SELECT partition_key, leased_by, lease_renewed_at FROM queue_partition_leases"). - WithArgs(testTenant, testConsumerGroup, "test_topic"). - WillReturnRows(sqlmock.NewRows(leaseColumns). - AddRow("existing1", testSubscriberName, freshMs). - AddRow("existing2", testSubscriberName, freshMs)) - // No acquire attempts. - }, - wantAcquired: 0, + got, err := store.GetAllLeasesForTenants(context.Background(), tenants, "test_topic", testConsumerGroup) + require.NoError(t, err) + require.Equal(t, map[string][]leaseInfo{ + "alpha": { + {PartitionKey: "p1", LeasedBy: testSubscriberName, LeaseRenewedAt: 1000}, }, - { - name: "lost race counts nothing", - maxPartitions: 0, - setup: func(mock sqlmock.Sqlmock) { - expectDiscover(mock, "part1") - mock.ExpectQuery("SELECT partition_key, leased_by, lease_renewed_at FROM queue_partition_leases"). - WithArgs(testTenant, testConsumerGroup, "test_topic"). - WillReturnRows(sqlmock.NewRows(leaseColumns)) - // Attempted while unleased, but another subscriber won the - // atomic acquire between the read and the write. - expectAcquire(mock, "other-worker") - }, - wantAcquired: 0, + "beta": { + {PartitionKey: "p2", LeasedBy: "other-worker", LeaseRenewedAt: 2000}, + {PartitionKey: "p3", LeasedBy: testSubscriberName, LeaseRenewedAt: 3000}, }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - db, mock, store := setuppartitionLeaseStoreTest(t) - defer db.Close() + }, got) - tt.setup(mock) - - acquired, discoveredPartitions, err := store.DiscoverAndAcquirePartitions(context.Background(), testTenant, "test_topic", testSubscriberName, testConsumerGroup, testLeaseDurationMs, tt.maxPartitions) - require.NoError(t, err) - require.Equal(t, tt.wantAcquired, acquired) - require.NotNil(t, discoveredPartitions) - require.NoError(t, mock.ExpectationsWereMet()) - }) - } + empty, err := store.GetAllLeasesForTenants(context.Background(), nil, "test_topic", testConsumerGroup) + require.NoError(t, err) + require.Empty(t, empty) + require.NoError(t, mock.ExpectationsWereMet()) } -func TestPartitionLeaseStore_PurgeStale(t *testing.T) { - tests := []struct { - name string - setup func(mock sqlmock.Sqlmock) - wantErr bool - }{ - { - name: "deletes rows older than threshold", - setup: func(mock sqlmock.Sqlmock) { - mock.ExpectExec("DELETE FROM queue_partition_leases"). - WithArgs(testTenant, testConsumerGroup, "test_topic", sqlmock.AnyArg()). - WillReturnResult(sqlmock.NewResult(0, 2)) - }, - }, - { - name: "no stale rows is a no-op", - setup: func(mock sqlmock.Sqlmock) { - mock.ExpectExec("DELETE FROM queue_partition_leases"). - WithArgs(testTenant, testConsumerGroup, "test_topic", sqlmock.AnyArg()). - WillReturnResult(sqlmock.NewResult(0, 0)) - }, - }, - { - name: "database error", - setup: func(mock sqlmock.Sqlmock) { - mock.ExpectExec("DELETE FROM queue_partition_leases"). - WithArgs(testTenant, testConsumerGroup, "test_topic", sqlmock.AnyArg()). - WillReturnError(fmt.Errorf("db error")) - }, - wantErr: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - db, mock, store := setuppartitionLeaseStoreTest(t) - defer db.Close() - - tt.setup(mock) +func TestPartitionLeaseStore_RenewAndReleaseOwnedLeases(t *testing.T) { + db, mock, store := setuppartitionLeaseStoreTest(t) + defer db.Close() - err := store.PurgeStale(context.Background(), testTenant, "test_topic", testConsumerGroup, 300_000) - if tt.wantErr { - require.Error(t, err) - } else { - require.NoError(t, err) - } - require.NoError(t, mock.ExpectationsWereMet()) - }) - } + tenants := []string{"alpha", "beta"} + mock.ExpectExec("UPDATE queue_partition_leases"). + WithArgs(sqlmock.AnyArg(), "alpha", "beta", testConsumerGroup, "test_topic", testSubscriberName). + WillReturnResult(sqlmock.NewResult(0, 2)) + mock.ExpectExec("DELETE FROM queue_partition_leases"). + WithArgs("alpha", "beta", testConsumerGroup, "test_topic", testSubscriberName). + WillReturnResult(sqlmock.NewResult(0, 2)) + mock.ExpectExec("DELETE FROM queue_partition_leases"). + WithArgs("alpha", "beta", testConsumerGroup, "test_topic", sqlmock.AnyArg()). + WillReturnResult(sqlmock.NewResult(0, 1)) + + require.NoError(t, store.RenewOwnedLeases(context.Background(), tenants, "test_topic", testSubscriberName, testConsumerGroup)) + require.NoError(t, store.ReleaseOwnedLeases(context.Background(), tenants, "test_topic", testSubscriberName, testConsumerGroup)) + require.NoError(t, store.PurgeStaleForTenants(context.Background(), tenants, "test_topic", testConsumerGroup, 300_000)) + require.NoError(t, mock.ExpectationsWereMet()) } diff --git a/platform/extension/messagequeue/mysql/sql.go b/platform/extension/messagequeue/mysql/sql.go index 62fa18b9..6d5c80be 100644 --- a/platform/extension/messagequeue/mysql/sql.go +++ b/platform/extension/messagequeue/mysql/sql.go @@ -91,7 +91,7 @@ func NewQueue(params Params) (extqueue.Queue, error) { messageStore := newMessageStore(params.DB, logger, params.MetricsScope) offsetStore := newOffsetStore(params.DB, params.MetricsScope) leaseStore := newPartitionLeaseStore(params.DB, logger, params.MetricsScope) - heartbeatStore := newSubscriberHeartbeatStore(params.DB, logger, params.MetricsScope, time.Now) + heartbeatStore := newSubscriberHeartbeatStore(params.DB, params.MetricsScope, time.Now) deliveryStateStore := newDeliveryStateStore(params.DB, logger, params.MetricsScope) queueMetrics := params.MetricsScope.SubScope("queue") diff --git a/platform/extension/messagequeue/mysql/stores.go b/platform/extension/messagequeue/mysql/stores.go index 2689f25a..496dc5af 100644 --- a/platform/extension/messagequeue/mysql/stores.go +++ b/platform/extension/messagequeue/mysql/stores.go @@ -67,9 +67,6 @@ type messageStore interface { // Insert inserts messages into the topic table. Insert(ctx context.Context, tenant string, topic string, messages []entityqueue.Message) error - // Delete deletes a message by tenant, topic, partition key, and ID - Delete(ctx context.Context, tenant string, topic string, partitionKey string, messageID string) error - // FetchByOffset fetches messages with offset > currentOffset for a specific partition. FetchByOffset(ctx context.Context, tenant string, topic string, partitionKey string, currentOffset int64, limit int) ([]messageRow, error) @@ -119,39 +116,45 @@ type partitionLeaseStore interface { // TryAcquireLease attempts to acquire or renew a lease for a partition TryAcquireLease(ctx context.Context, tenant string, topic string, partitionKey string, subscriberName string, consumerGroup string, leaseDurationMs int64) (bool, error) - // RenewLease renews the lease for a partition owned by this worker - RenewLease(ctx context.Context, tenant string, topic string, partitionKey string, subscriberName string, consumerGroup string, leaseDurationMs int64) error - // ReleaseLease releases the lease for a partition owned by this worker ReleaseLease(ctx context.Context, tenant string, topic string, partitionKey string, subscriberName string, consumerGroup string) error - // GetLeasedPartitions returns all partitions currently leased by this worker - GetLeasedPartitions(ctx context.Context, tenant string, topic string, subscriberName string, consumerGroup string) ([]string, error) + // GetLeasedPartitionsForTenants returns partitions leased by this worker + // across the given tenants, keyed by tenant. Tenants with no leases are absent. + GetLeasedPartitionsForTenants(ctx context.Context, tenants []string, topic string, subscriberName string, consumerGroup string) (map[string][]string, error) + + // GetAllLeasesForTenants returns every lease row under (topic, consumerGroup) + // for the given tenants, keyed by tenant. + GetAllLeasesForTenants(ctx context.Context, tenants []string, topic string, consumerGroup string) (map[string][]leaseInfo, error) + + // RenewOwnedLeases refreshes lease_renewed_at on every row this subscriber + // still holds across tenants. Stolen leases are ignored; discovery drops them. + RenewOwnedLeases(ctx context.Context, tenants []string, topic string, subscriberName string, consumerGroup string) error - // GetAllLeases returns the lease row for every partition currently leased - // under (tenant, topic, consumerGroup) by any subscriber. - GetAllLeases(ctx context.Context, tenant string, topic string, consumerGroup string) ([]leaseInfo, error) + // ReleaseOwnedLeases deletes every lease row this subscriber holds across tenants. + ReleaseOwnedLeases(ctx context.Context, tenants []string, topic string, subscriberName string, consumerGroup string) error - // PurgeStale deletes lease rows not renewed within olderThanMs. - PurgeStale(ctx context.Context, tenant string, topic string, consumerGroup string, olderThanMs int64) error + // PurgeStaleForTenants deletes lease rows not renewed within olderThanMs + // across the given tenants. + PurgeStaleForTenants(ctx context.Context, tenants []string, topic string, consumerGroup string, olderThanMs int64) error - // DiscoverAndAcquirePartitions discovers partitions from messages table and tries to acquire leases. - DiscoverAndAcquirePartitions(ctx context.Context, tenant string, topic string, subscriberName string, consumerGroup string, leaseDurationMs int64, maxPartitions int) (acquiredCount int, discoveredPartitions []string, err error) + // DiscoverPartitions returns distinct partition keys per tenant for topic. + DiscoverPartitions(ctx context.Context, tenants []string, topic string) (map[string][]string, error) } // subscriberHeartbeatStore handles subscriber heartbeat operations for fair partition leasing (internal use only) type subscriberHeartbeatStore interface { - // Heartbeat registers or renews a subscriber's heartbeat - Heartbeat(ctx context.Context, tenant string, topic string, subscriberName string, consumerGroup string) error + // HeartbeatForTenants registers or renews this subscriber's heartbeat for each tenant. + HeartbeatForTenants(ctx context.Context, tenants []string, topic string, subscriberName string, consumerGroup string) error - // ActiveSubscribers returns the names of subscribers with a recent heartbeat. - ActiveSubscribers(ctx context.Context, tenant string, topic string, consumerGroup string, staleDurationMs int64) ([]string, error) + // ActiveSubscribersForTenants returns recent subscriber names keyed by tenant. + ActiveSubscribersForTenants(ctx context.Context, tenants []string, topic string, consumerGroup string, staleDurationMs int64) (map[string][]string, error) - // Deregister removes a subscriber's heartbeat row. - Deregister(ctx context.Context, tenant string, topic string, subscriberName string, consumerGroup string) error + // DeregisterForTenants removes this subscriber's heartbeat rows across tenants. + DeregisterForTenants(ctx context.Context, tenants []string, topic string, subscriberName string, consumerGroup string) error - // PurgeStale deletes heartbeat rows whose last heartbeat is older than olderThanMs. - PurgeStale(ctx context.Context, tenant string, topic string, consumerGroup string, olderThanMs int64) error + // PurgeStaleForTenants deletes stale heartbeat rows across the given tenants. + PurgeStaleForTenants(ctx context.Context, tenants []string, topic string, consumerGroup string, olderThanMs int64) error } // DeliveryState represents the full per-message delivery tracking state. diff --git a/platform/extension/messagequeue/mysql/subscriber.go b/platform/extension/messagequeue/mysql/subscriber.go index c9d4aa8f..1dff7b76 100644 --- a/platform/extension/messagequeue/mysql/subscriber.go +++ b/platform/extension/messagequeue/mysql/subscriber.go @@ -148,8 +148,7 @@ type subscription struct { workers map[entityqueue.PartitionIdentity]*partitionWorker workersMu sync.Mutex - // lastDiscoveredPartitions is cached from the most recent - // DiscoverAndAcquirePartitions calls. + // lastDiscoveredPartitions is cached from the most recent discovery. // Used by fairShareCap during rebalance to avoid a redundant discovery query. lastDiscoveredPartitions []entityqueue.PartitionIdentity @@ -181,41 +180,6 @@ func sortPartitionIdentities(partitions []entityqueue.PartitionIdentity) { }) } -type tenantOperationResult[T any] struct { - tenant string - value T - err error -} - -func runTenantOperations[T any]( - ctx context.Context, - tenants []string, - timeout time.Duration, - operation func(context.Context, string) (T, error), -) []tenantOperationResult[T] { - results := make(chan tenantOperationResult[T], len(tenants)) - for _, tenant := range tenants { - go func() { - result := tenantOperationResult[T]{tenant: tenant} - defer func() { results <- result }() - tenantCtx, cancel := context.WithTimeout(ctx, timeout) - defer cancel() - result.value, result.err = operation(tenantCtx, tenant) - }() - } - - byTenant := make(map[string]tenantOperationResult[T], len(tenants)) - for range tenants { - result := <-results - byTenant[result.tenant] = result - } - ordered := make([]tenantOperationResult[T], 0, len(tenants)) - for _, tenant := range tenants { - ordered = append(ordered, byTenant[tenant]) - } - return ordered -} - // partitionWorker handles polling and delivering messages for a single partition. // Each worker runs in its own goroutine, polling the DB on a ticker and sending // deliveries to the shared deliveryCh. @@ -687,12 +651,11 @@ func (s *subscriber) managePartitions(ctx context.Context, sub *subscription) { // Initial heartbeat failure is non-fatal — the next leaseTicker fires within // LeaseRenewalIntervalMs and retries. tenantLeaseTimeout := time.Duration(cfg.LeaseRenewalIntervalMs) * time.Millisecond - for _, result := range runTenantOperations(ctx, s.tenants, tenantLeaseTimeout, func(tenantCtx context.Context, tenant string) (struct{}, error) { - return struct{}{}, s.sendHeartbeat(tenantCtx, sub, tenant) - }) { - if result.err != nil { - s.logger.Errorw("initial heartbeat failed", append(logFields, "tenant", result.tenant, "error", result.err)...) - } + hbCtx, hbCancel := context.WithTimeout(ctx, tenantLeaseTimeout) + err := s.sendHeartbeats(hbCtx, sub) + hbCancel() + if err != nil { + s.logger.Errorw("initial heartbeat failed", append(logFields, "error", err)...) } for { @@ -719,63 +682,7 @@ func (s *subscriber) managePartitions(ctx context.Context, sub *subscription) { return case <-leaseTicker.C: - runTenantOperations(ctx, s.tenants, tenantLeaseTimeout, func(tenantCtx context.Context, tenant string) (struct{}, error) { - tenantFields := append(logFields, "tenant", tenant) - // Fetch leased partitions once for this tenant tick — shared by - // rebalance and renewLeases to avoid redundant queries. - leasedPartitions, err := s.leaseStore.GetLeasedPartitions(tenantCtx, tenant, sub.topic, cfg.SubscriberName, cfg.ConsumerGroup) - if err != nil { - s.logger.Errorw("get leased partitions failed", append(tenantFields, "error", err)...) - // Skip rebalance+renew for this tenant; retry next tick. - if err := s.sendHeartbeat(tenantCtx, sub, tenant); err != nil { - s.logger.Errorw("heartbeat failed during lease error recovery", append(tenantFields, "error", err)...) - } - return struct{}{}, nil - } - - // Rebalance, renew, and heartbeat are independent operations. - // Each can fail without affecting the others — the next tick retries. - // Renewal covers only the partitions kept after shedding; renewing - // a just-released lease would spuriously fail with ErrLeaseExpired. - released, err := s.rebalance(tenantCtx, sub, tenant, leasedPartitions) - if err != nil { - s.logger.Errorw("rebalance failed", append(tenantFields, "error", err)...) - } - kept := leasedPartitions - if len(released) > 0 { - releasedSet := make(map[string]struct{}, len(released)) - for _, pk := range released { - releasedSet[pk] = struct{}{} - } - kept = make([]string, 0, len(leasedPartitions)) - for _, pk := range leasedPartitions { - if _, ok := releasedSet[pk]; !ok { - kept = append(kept, pk) - } - } - } - if err := s.renewLeases(tenantCtx, sub, tenant, kept); err != nil { - s.logger.Errorw("lease renewal failed", append(tenantFields, "error", err)...) - } - if err := s.sendHeartbeat(tenantCtx, sub, tenant); err != nil { - s.logger.Errorw("periodic heartbeat failed", append(tenantFields, "error", err)...) - } - // Purge heartbeat rows abandoned by subscribers that never - // deregistered (crashes) — without this the table grows - // monotonically, since every process registers under a fresh - // hostname-pid name. - if err := s.heartbeatStore.PurgeStale(tenantCtx, tenant, sub.topic, cfg.ConsumerGroup, heartbeatPurgeAfterLeaseDurations*cfg.LeaseDurationMs); err != nil { - s.logger.Errorw("stale heartbeat purge failed", append(tenantFields, "error", err)...) - } - // Purge lease rows abandoned by holders that crashed while - // owning a drained partition — acquisition only probes - // discovered partitions, so nothing else ever refreshes or - // removes a stale lease on a partition with no messages. - if err := s.leaseStore.PurgeStale(tenantCtx, tenant, sub.topic, cfg.ConsumerGroup, leasePurgeAfterLeaseDurations*cfg.LeaseDurationMs); err != nil { - s.logger.Errorw("stale lease purge failed", append(tenantFields, "error", err)...) - } - return struct{}{}, nil - }) + s.runLeaseTick(ctx, sub, tenantLeaseTimeout, logFields) s.emitSignal(SignalPartitionUpdate) case <-discoveryTicker.C: @@ -812,53 +719,52 @@ func (s *subscriber) discoverAndReconcileWorkers(ctx context.Context, sub *subsc cachedDiscovered := append([]entityqueue.PartitionIdentity(nil), sub.lastDiscoveredPartitions...) sub.workersMu.Unlock() - discoveredByTenant := make(map[string][]string, len(s.tenants)) - leasedByTenant := make(map[string][]string, len(s.tenants)) - var discoveryErrs []error - - type tenantDiscovery struct { - discovered []string - leased []string - } discoveryTimeout := max( time.Duration(cfg.PartitionDiscoveryIntervalMs)*time.Millisecond, time.Duration(cfg.LeaseRenewalIntervalMs)*time.Millisecond, ) - results := runTenantOperations(ctx, s.tenants, discoveryTimeout, func(tenantCtx context.Context, tenant string) (tenantDiscovery, error) { - leasedPartitions, err := s.leaseStore.GetLeasedPartitions(tenantCtx, tenant, sub.topic, cfg.SubscriberName, cfg.ConsumerGroup) - if err != nil { - return tenantDiscovery{}, fmt.Errorf("get leased partitions tenant=%s: %w", tenant, err) - } + tickCtx, cancel := context.WithTimeout(ctx, discoveryTimeout) + defer cancel() - cachedForTenant := partitionKeysForTenant(cachedDiscovered, tenant) - - maxPartitions := 0 - if !uncapped { - maxPartitions, err = s.fairShareCap(tenantCtx, sub, tenant, leasedPartitions, cachedForTenant) - if err != nil { - return tenantDiscovery{}, fmt.Errorf("compute fair share cap tenant=%s: %w", tenant, err) - } - } - - _, discoveredPartitions, err := s.leaseStore.DiscoverAndAcquirePartitions(tenantCtx, tenant, sub.topic, cfg.SubscriberName, cfg.ConsumerGroup, cfg.LeaseDurationMs, maxPartitions) - if err != nil { - return tenantDiscovery{}, fmt.Errorf("discover and acquire partitions tenant=%s: %w", tenant, err) - } + discoveredByTenant, err := s.leaseStore.DiscoverPartitions(tickCtx, s.tenants, sub.topic) + if err != nil { + s.keepCachedDiscovery(sub, cachedDiscovered) + s.reconcilePartitionWorkers(ctx, sub, nil) + return fmt.Errorf("discover partitions: %w", err) + } + leasedByTenant, err := s.leaseStore.GetLeasedPartitionsForTenants(tickCtx, s.tenants, sub.topic, cfg.SubscriberName, cfg.ConsumerGroup) + if err != nil { + s.keepCachedDiscovery(sub, cachedDiscovered) + s.reconcilePartitionWorkers(ctx, sub, nil) + return fmt.Errorf("get leased partitions: %w", err) + } + allLeasesByTenant, err := s.leaseStore.GetAllLeasesForTenants(tickCtx, s.tenants, sub.topic, cfg.ConsumerGroup) + var skipAcquire error + if err != nil { + skipAcquire = fmt.Errorf("get all leases: %w", err) + } - leasedPartitions, err = s.leaseStore.GetLeasedPartitions(tenantCtx, tenant, sub.topic, cfg.SubscriberName, cfg.ConsumerGroup) + var activeByTenant map[string][]string + if skipAcquire == nil && !uncapped { + activeByTenant, err = s.heartbeatStore.ActiveSubscribersForTenants(tickCtx, s.tenants, sub.topic, cfg.ConsumerGroup, cfg.LeaseDurationMs) if err != nil { - return tenantDiscovery{}, fmt.Errorf("get leased partitions after acquire tenant=%s: %w", tenant, err) + skipAcquire = fmt.Errorf("active subscribers: %w", err) } - return tenantDiscovery{discovered: discoveredPartitions, leased: leasedPartitions}, nil - }) + } - for _, result := range results { - if result.err != nil { - discoveryErrs = append(discoveryErrs, result.err) - continue + if skipAcquire == nil { + for _, tenant := range s.tenants { + discovered := discoveredByTenant[tenant] + leased := leasedByTenant[tenant] + maxPartitions := 0 + if !uncapped { + maxPartitions = s.fairShareCap(sub, leased, discovered, activeByTenant[tenant]) + } + acquired := s.acquireUnownedPartitions(tickCtx, sub, tenant, discovered, allLeasesByTenant[tenant], maxPartitions) + if len(acquired) > 0 { + leasedByTenant[tenant] = append(append([]string{}, leased...), acquired...) + } } - discoveredByTenant[result.tenant] = result.value.discovered - leasedByTenant[result.tenant] = result.value.leased } allDiscovered := make([]entityqueue.PartitionIdentity, 0) @@ -869,15 +775,7 @@ func (s *subscriber) discoverAndReconcileWorkers(ctx context.Context, sub *subsc var expired []entityqueue.PartitionIdentity for _, tenant := range s.tenants { - discoveredPartitions, succeeded := discoveredByTenant[tenant] - if !succeeded { - // Unconfirmed leases must not keep workers polling: a peer can acquire after expiry. - for _, pk := range partitionKeysForTenant(cachedDiscovered, tenant) { - allDiscovered = append(allDiscovered, entityqueue.PartitionIdentity{Tenant: tenant, PartitionKey: pk}) - } - continue - } - + discoveredPartitions := discoveredByTenant[tenant] leasedPartitions := leasedByTenant[tenant] tenantDiscovered := make([]entityqueue.PartitionIdentity, 0, len(discoveredPartitions)) for _, pk := range discoveredPartitions { @@ -959,7 +857,58 @@ func (s *subscriber) discoverAndReconcileWorkers(ctx context.Context, sub *subsc } s.reconcilePartitionWorkers(ctx, sub, allLeased) - return errors.Join(discoveryErrs...) + return skipAcquire +} + +func (s *subscriber) keepCachedDiscovery(sub *subscription, cached []entityqueue.PartitionIdentity) { + sub.workersMu.Lock() + sub.lastDiscoveredPartitions = cached + sub.workersMu.Unlock() +} + +func (s *subscriber) acquireUnownedPartitions(ctx context.Context, sub *subscription, tenant string, discovered []string, leases []leaseInfo, maxPartitions int) []string { + cfg := sub.config + staleThreshold := currentTimeMillis() - cfg.LeaseDurationMs + ownedCount := 0 + ownedSet := make(map[string]struct{}) + heldByOther := make(map[string]struct{}) + for _, lease := range leases { + switch { + case lease.LeasedBy == cfg.SubscriberName: + ownedSet[lease.PartitionKey] = struct{}{} + ownedCount++ + case lease.LeaseRenewedAt >= staleThreshold: + heldByOther[lease.PartitionKey] = struct{}{} + } + } + + var acquired []string + for _, partitionKey := range discovered { + if _, owned := ownedSet[partitionKey]; owned { + continue + } + if _, held := heldByOther[partitionKey]; held { + continue + } + if maxPartitions > 0 && ownedCount >= maxPartitions { + break + } + ok, err := s.leaseStore.TryAcquireLease(ctx, tenant, sub.topic, partitionKey, cfg.SubscriberName, cfg.ConsumerGroup, cfg.LeaseDurationMs) + if err != nil { + s.logger.Errorw("failed to acquire lease for partition", + "tenant", tenant, + "topic", sub.topic, + "partition_key", partitionKey, + "error", err, + ) + continue + } + if ok { + acquired = append(acquired, partitionKey) + ownedCount++ + } + } + return acquired } // updateDrainedTracking recomputes, for every owned partition absent from @@ -1411,72 +1360,59 @@ func (w *partitionWorker) garbageCollect(ctx context.Context) error { return nil } -// renewLeases renews leases for all partitions owned by this worker. -func (s *subscriber) renewLeases(ctx context.Context, sub *subscription, tenant string, leasedPartitions []string) error { +func (s *subscriber) runLeaseTick(ctx context.Context, sub *subscription, timeout time.Duration, logFields []any) { cfg := sub.config + tickCtx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() - for _, partitionKey := range leasedPartitions { - if err := s.leaseStore.RenewLease(ctx, tenant, sub.topic, partitionKey, cfg.SubscriberName, cfg.ConsumerGroup, cfg.LeaseDurationMs); err != nil { - return fmt.Errorf("renew lease tenant=%s partition=%s: %w", tenant, partitionKey, err) + leasedByTenant, err := s.leaseStore.GetLeasedPartitionsForTenants(tickCtx, s.tenants, sub.topic, cfg.SubscriberName, cfg.ConsumerGroup) + if err != nil { + s.logger.Errorw("get leased partitions failed", append(logFields, "error", err)...) + } else { + activeByTenant, err := s.heartbeatStore.ActiveSubscribersForTenants(tickCtx, s.tenants, sub.topic, cfg.ConsumerGroup, cfg.LeaseDurationMs) + if err != nil { + s.logger.Errorw("active subscribers failed", append(logFields, "error", err)...) + } else { + for _, tenant := range s.tenants { + if _, err := s.rebalance(tickCtx, sub, tenant, leasedByTenant[tenant], activeByTenant[tenant]); err != nil { + s.logger.Errorw("rebalance failed", append(logFields, "tenant", tenant, "error", err)...) + } + } } } - return nil + if err := s.leaseStore.RenewOwnedLeases(tickCtx, s.tenants, sub.topic, cfg.SubscriberName, cfg.ConsumerGroup); err != nil { + s.logger.Errorw("lease renewal failed", append(logFields, "error", err)...) + } + if err := s.sendHeartbeats(tickCtx, sub); err != nil { + s.logger.Errorw("periodic heartbeat failed", append(logFields, "error", err)...) + } + if err := s.heartbeatStore.PurgeStaleForTenants(tickCtx, s.tenants, sub.topic, cfg.ConsumerGroup, heartbeatPurgeAfterLeaseDurations*cfg.LeaseDurationMs); err != nil { + s.logger.Errorw("stale heartbeat purge failed", append(logFields, "error", err)...) + } + if err := s.leaseStore.PurgeStaleForTenants(tickCtx, s.tenants, sub.topic, cfg.ConsumerGroup, leasePurgeAfterLeaseDurations*cfg.LeaseDurationMs); err != nil { + s.logger.Errorw("stale lease purge failed", append(logFields, "error", err)...) + } } -// releaseAllLeases releases all leases for a topic. func (s *subscriber) releaseAllLeases(ctx context.Context, sub *subscription) error { cfg := sub.config - var releaseErrs []error - timeout := time.Duration(cfg.LeaseRenewalIntervalMs) * time.Millisecond - results := runTenantOperations(ctx, s.tenants, timeout, func(tenantCtx context.Context, tenant string) (struct{}, error) { - leasedPartitions, err := s.leaseStore.GetLeasedPartitions(tenantCtx, tenant, sub.topic, cfg.SubscriberName, cfg.ConsumerGroup) - if err != nil { - return struct{}{}, fmt.Errorf("get leased partitions for release tenant=%s: %w", tenant, err) - } - - var tenantErrs []error - for _, partitionKey := range leasedPartitions { - if err := s.leaseStore.ReleaseLease(tenantCtx, tenant, sub.topic, partitionKey, cfg.SubscriberName, cfg.ConsumerGroup); err != nil { - tenantErrs = append(tenantErrs, fmt.Errorf("release lease tenant=%s partition=%s: %w", tenant, partitionKey, err)) - } - } - return struct{}{}, errors.Join(tenantErrs...) - }) - for _, result := range results { - if result.err != nil { - releaseErrs = append(releaseErrs, result.err) - } - } - return errors.Join(releaseErrs...) + return s.leaseStore.ReleaseOwnedLeases(ctx, s.tenants, sub.topic, cfg.SubscriberName, cfg.ConsumerGroup) } -// sendHeartbeat sends a heartbeat for this subscriber. -func (s *subscriber) sendHeartbeat(ctx context.Context, sub *subscription, tenant string) error { +func (s *subscriber) sendHeartbeats(ctx context.Context, sub *subscription) error { cfg := sub.config - if err := s.heartbeatStore.Heartbeat(ctx, tenant, sub.topic, cfg.SubscriberName, cfg.ConsumerGroup); err != nil { - return fmt.Errorf("heartbeat tenant=%s: %w", tenant, err) + if err := s.heartbeatStore.HeartbeatForTenants(ctx, s.tenants, sub.topic, cfg.SubscriberName, cfg.ConsumerGroup); err != nil { + return fmt.Errorf("heartbeat: %w", err) } return nil } -// deregisterHeartbeat removes this subscriber's heartbeat entry during shutdown. func (s *subscriber) deregisterHeartbeat(ctx context.Context, sub *subscription) error { cfg := sub.config - var deregistrationErrs []error - timeout := time.Duration(cfg.LeaseRenewalIntervalMs) * time.Millisecond - results := runTenantOperations(ctx, s.tenants, timeout, func(tenantCtx context.Context, tenant string) (struct{}, error) { - err := s.heartbeatStore.Deregister(tenantCtx, tenant, sub.topic, cfg.SubscriberName, cfg.ConsumerGroup) - if err != nil { - return struct{}{}, fmt.Errorf("deregister heartbeat tenant=%s: %w", tenant, err) - } - return struct{}{}, nil - }) - for _, result := range results { - if result.err != nil { - deregistrationErrs = append(deregistrationErrs, result.err) - } + if err := s.heartbeatStore.DeregisterForTenants(ctx, s.tenants, sub.topic, cfg.SubscriberName, cfg.ConsumerGroup); err != nil { + return fmt.Errorf("deregister heartbeat: %w", err) } - return errors.Join(deregistrationErrs...) + return nil } // rebalance checks if this subscriber holds more partitions than its fair share @@ -1485,17 +1421,14 @@ func (s *subscriber) deregisterHeartbeat(ctx context.Context, sub *subscription) // renewing a just-released lease would spuriously fail with ErrLeaseExpired. // The owned slice is never mutated (the caller shares it with lease renewal). // On error, partitions released before the failure are still returned. -func (s *subscriber) rebalance(ctx context.Context, sub *subscription, tenant string, owned []string) (released []string, retErr error) { +func (s *subscriber) rebalance(ctx context.Context, sub *subscription, tenant string, owned []string, active []string) (released []string, retErr error) { cfg := sub.config sub.workersMu.Lock() discoveredPartitions := partitionKeysForTenant(sub.lastDiscoveredPartitions, tenant) sub.workersMu.Unlock() - maxPart, err := s.fairShareCap(ctx, sub, tenant, owned, discoveredPartitions) - if err != nil { - return nil, fmt.Errorf("compute fair share cap: %w", err) - } + maxPart := s.fairShareCap(sub, owned, discoveredPartitions, active) if maxPart == 0 || len(owned) <= maxPart { return nil, nil } @@ -1525,7 +1458,7 @@ func (s *subscriber) rebalance(ctx context.Context, sub *subscription, tenant st } // fairShareCap computes the max partitions this subscriber should own. -// Returns (maxPart, error). maxPart=0 means unlimited. +// maxPart=0 means unlimited. // owned is the caller-provided list of leased partitions. // discoveredPartitions is an optional pre-fetched list of all known partitions; // if nil, only owned partitions are used for fair share computation. @@ -1539,15 +1472,11 @@ func (s *subscriber) rebalance(ctx context.Context, sub *subscription, tenant st // cap implies another under its cap (rebalance sheds, the peer acquires), // and an unleased partition implies a subscriber with spare cap to claim it // — neither a starved subscriber nor a leftover partition is a stable state. -func (s *subscriber) fairShareCap(ctx context.Context, sub *subscription, tenant string, owned []string, discoveredPartitions []string) (int, error) { +func (s *subscriber) fairShareCap(sub *subscription, owned []string, discoveredPartitions []string, active []string) int { cfg := sub.config - active, err := s.heartbeatStore.ActiveSubscribers(ctx, tenant, sub.topic, cfg.ConsumerGroup, cfg.LeaseDurationMs) - if err != nil { - return 0, err - } if len(active) <= 1 { - return 0, nil + return 0 } // Count all known partitions as the union of owned + discovered. @@ -1593,7 +1522,7 @@ func (s *subscriber) fairShareCap(ctx context.Context, sub *subscription, tenant maxPart = 1 } - return maxPart, nil + return maxPart } // Close cancels every subscription and waits up to subscriptionShutdownTimeout diff --git a/platform/extension/messagequeue/mysql/subscriber_heartbeat_store.go b/platform/extension/messagequeue/mysql/subscriber_heartbeat_store.go index 06645be4..21097c4e 100644 --- a/platform/extension/messagequeue/mysql/subscriber_heartbeat_store.go +++ b/platform/extension/messagequeue/mysql/subscriber_heartbeat_store.go @@ -18,141 +18,131 @@ import ( "context" "database/sql" "fmt" + "strings" "time" "github.com/uber-go/tally" "github.com/uber/submitqueue/platform/metrics" - "go.uber.org/zap" ) // sqlSubscriberHeartbeatStore is the SQL implementation of subscriberHeartbeatStore type sqlSubscriberHeartbeatStore struct { db *sql.DB - logger *zap.SugaredLogger scope tally.Scope nowFunc func() time.Time } // newSubscriberHeartbeatStore creates a new SQL subscriber heartbeat store -func newSubscriberHeartbeatStore(db *sql.DB, logger *zap.SugaredLogger, scope tally.Scope, nowFunc func() time.Time) subscriberHeartbeatStore { +func newSubscriberHeartbeatStore(db *sql.DB, scope tally.Scope, nowFunc func() time.Time) subscriberHeartbeatStore { return &sqlSubscriberHeartbeatStore{ db: db, - logger: logger.Named("subscriber_heartbeat_store"), scope: scope.SubScope("subscriber_heartbeat_store"), nowFunc: nowFunc, } } -// Heartbeat registers or renews a subscriber's heartbeat. -func (s *sqlSubscriberHeartbeatStore) Heartbeat(ctx context.Context, tenant string, topic string, subscriberName string, consumerGroup string) (retErr error) { - op := metrics.Begin(s.scope, "heartbeat", metrics.StorageLatencyBuckets) +func (s *sqlSubscriberHeartbeatStore) HeartbeatForTenants(ctx context.Context, tenants []string, topic string, subscriberName string, consumerGroup string) (retErr error) { + op := metrics.Begin(s.scope, "heartbeat_for_tenants", metrics.StorageLatencyBuckets) defer func() { op.Complete(retErr) }() + if len(tenants) == 0 { + return nil + } now := s.nowFunc().UnixMilli() - + valueParts := make([]string, len(tenants)) + args := make([]any, 0, len(tenants)*5) + for i, tenant := range tenants { + valueParts[i] = "(?, ?, ?, ?, ?, 0)" + args = append(args, tenant, consumerGroup, topic, subscriberName, now) + } _, err := s.db.ExecContext(ctx, fmt.Sprintf(` INSERT INTO %s (tenant, consumer_group, topic, subscriber_name, heartbeat_at, deregistered_at) - VALUES (?, ?, ?, ?, ?, 0) + VALUES %s ON DUPLICATE KEY UPDATE heartbeat_at = VALUES(heartbeat_at), deregistered_at = 0 - `, SubscriberHeartbeatsTableName), tenant, consumerGroup, topic, subscriberName, now) - + `, SubscriberHeartbeatsTableName, strings.Join(valueParts, ", ")), args...) if err != nil { - return fmt.Errorf("failed to send heartbeat tenant=%s topic=%s: %w", tenant, topic, err) + return fmt.Errorf("failed to send heartbeat topic=%s: %w", topic, err) } - return nil } -// ActiveSubscribers returns the names of subscribers with a heartbeat newer than the stale threshold. -func (s *sqlSubscriberHeartbeatStore) ActiveSubscribers(ctx context.Context, tenant string, topic string, consumerGroup string, staleDurationMs int64) (_ []string, retErr error) { - op := metrics.Begin(s.scope, "active_subscribers", metrics.StorageLatencyBuckets) +func (s *sqlSubscriberHeartbeatStore) ActiveSubscribersForTenants(ctx context.Context, tenants []string, topic string, consumerGroup string, staleDurationMs int64) (_ map[string][]string, retErr error) { + op := metrics.Begin(s.scope, "active_subscribers_for_tenants", metrics.StorageLatencyBuckets) defer func() { op.Complete(retErr) }() + placeholders, ok := inListPlaceholders(len(tenants)) + if !ok { + return map[string][]string{}, nil + } staleThreshold := s.nowFunc().UnixMilli() - staleDurationMs + args := appendStrings(nil, tenants) + args = append(args, consumerGroup, topic, staleThreshold) rows, err := s.db.QueryContext(ctx, fmt.Sprintf(` - SELECT subscriber_name FROM %s - WHERE tenant = ? AND consumer_group = ? AND topic = ? AND heartbeat_at >= ? AND deregistered_at = 0 - `, SubscriberHeartbeatsTableName), tenant, consumerGroup, topic, staleThreshold) + SELECT tenant, subscriber_name FROM %s + WHERE tenant IN (%s) AND consumer_group = ? AND topic = ? AND heartbeat_at >= ? AND deregistered_at = 0 + `, SubscriberHeartbeatsTableName, placeholders), args...) if err != nil { - return nil, fmt.Errorf("failed to query active subscribers tenant=%s topic=%s: %w", tenant, topic, err) + return nil, fmt.Errorf("failed to query active subscribers topic=%s: %w", topic, err) } defer rows.Close() - var names []string + byTenant := make(map[string][]string) for rows.Next() { - var name string - if err := rows.Scan(&name); err != nil { - return nil, fmt.Errorf("failed to scan subscriber name tenant=%s topic=%s: %w", tenant, topic, err) + var tenant, name string + if err := rows.Scan(&tenant, &name); err != nil { + return nil, fmt.Errorf("failed to scan subscriber name topic=%s: %w", topic, err) } - names = append(names, name) + byTenant[tenant] = append(byTenant[tenant], name) } - if err := rows.Err(); err != nil { - return nil, fmt.Errorf("row iteration error tenant=%s topic=%s: %w", tenant, topic, err) + return nil, fmt.Errorf("row iteration error topic=%s: %w", topic, err) } - - s.logger.Debugw("found active subscribers", - logTenant, tenant, - logTopic, topic, - "count", len(names), - "subscribers", names, - ) - - return names, nil + return byTenant, nil } -// Deregister removes a subscriber's heartbeat row (hard delete — see the -// subscriberHeartbeatStore interface doc). Idempotent: no-op if already gone. -func (s *sqlSubscriberHeartbeatStore) Deregister(ctx context.Context, tenant string, topic string, subscriberName string, consumerGroup string) (retErr error) { - op := metrics.Begin(s.scope, "deregister", metrics.StorageLatencyBuckets) +func (s *sqlSubscriberHeartbeatStore) DeregisterForTenants(ctx context.Context, tenants []string, topic string, subscriberName string, consumerGroup string) (retErr error) { + op := metrics.Begin(s.scope, "deregister_for_tenants", metrics.StorageLatencyBuckets) defer func() { op.Complete(retErr) }() + placeholders, ok := inListPlaceholders(len(tenants)) + if !ok { + return nil + } + args := appendStrings(nil, tenants) + args = append(args, consumerGroup, topic, subscriberName) + _, err := s.db.ExecContext(ctx, fmt.Sprintf(` DELETE FROM %s - WHERE tenant = ? AND consumer_group = ? AND topic = ? AND subscriber_name = ? - `, SubscriberHeartbeatsTableName), tenant, consumerGroup, topic, subscriberName) - + WHERE tenant IN (%s) AND consumer_group = ? AND topic = ? AND subscriber_name = ? + `, SubscriberHeartbeatsTableName, placeholders), args...) if err != nil { - return fmt.Errorf("failed to deregister subscriber tenant=%s topic=%s: %w", tenant, topic, err) + return fmt.Errorf("failed to deregister subscriber topic=%s: %w", topic, err) } - - s.logger.Debugw("deregistered subscriber", - logTenant, tenant, - logTopic, topic, - "subscriber_name", subscriberName, - ) - return nil } -// PurgeStale deletes heartbeat rows older than olderThanMs for the topic and -// consumer group. See the subscriberHeartbeatStore interface doc. -func (s *sqlSubscriberHeartbeatStore) PurgeStale(ctx context.Context, tenant string, topic string, consumerGroup string, olderThanMs int64) (retErr error) { - op := metrics.Begin(s.scope, "purge_stale", metrics.StorageLatencyBuckets) +func (s *sqlSubscriberHeartbeatStore) PurgeStaleForTenants(ctx context.Context, tenants []string, topic string, consumerGroup string, olderThanMs int64) (retErr error) { + op := metrics.Begin(s.scope, "purge_stale_for_tenants", metrics.StorageLatencyBuckets) defer func() { op.Complete(retErr) }() + placeholders, ok := inListPlaceholders(len(tenants)) + if !ok { + return nil + } threshold := s.nowFunc().UnixMilli() - olderThanMs + args := appendStrings(nil, tenants) + args = append(args, consumerGroup, topic, threshold) result, err := s.db.ExecContext(ctx, fmt.Sprintf(` DELETE FROM %s - WHERE tenant = ? AND consumer_group = ? AND topic = ? AND heartbeat_at < ? - `, SubscriberHeartbeatsTableName), tenant, consumerGroup, topic, threshold) - + WHERE tenant IN (%s) AND consumer_group = ? AND topic = ? AND heartbeat_at < ? + `, SubscriberHeartbeatsTableName, placeholders), args...) if err != nil { - return fmt.Errorf("failed to purge stale heartbeats tenant=%s topic=%s: %w", tenant, topic, err) + return fmt.Errorf("failed to purge stale heartbeats topic=%s: %w", topic, err) } - - // RowsAffected error is swallowed because the DELETE itself succeeded; - // the count is for observability only. if deleted, err := result.RowsAffected(); err == nil && deleted > 0 { - metrics.NamedCounter(s.scope, "purge_stale", "rows_deleted", deleted, metrics.NewTag("topic", topic)) - s.logger.Debugw("purged stale heartbeats", - logTenant, tenant, - logTopic, topic, - "deleted", deleted, - ) + metrics.NamedCounter(s.scope, "purge_stale_for_tenants", "rows_deleted", deleted, metrics.NewTag("topic", topic)) } - return nil } diff --git a/platform/extension/messagequeue/mysql/subscriber_heartbeat_store_test.go b/platform/extension/messagequeue/mysql/subscriber_heartbeat_store_test.go index 37b85ad7..803b2969 100644 --- a/platform/extension/messagequeue/mysql/subscriber_heartbeat_store_test.go +++ b/platform/extension/messagequeue/mysql/subscriber_heartbeat_store_test.go @@ -24,7 +24,6 @@ import ( "github.com/DATA-DOG/go-sqlmock" "github.com/stretchr/testify/require" "github.com/uber-go/tally" - "go.uber.org/zap/zaptest" ) func setupSubscriberHeartbeatStoreTest(t *testing.T) (*sql.DB, sqlmock.Sqlmock, subscriberHeartbeatStore) { @@ -33,295 +32,138 @@ func setupSubscriberHeartbeatStoreTest(t *testing.T) (*sql.DB, sqlmock.Sqlmock, db, mock, err := sqlmock.New() require.NoError(t, err) - store := newSubscriberHeartbeatStore(db, zaptest.NewLogger(t).Sugar(), tally.NoopScope, time.Now) + store := newSubscriberHeartbeatStore(db, tally.NoopScope, time.Now) return db, mock, store } -func TestSubscriberHeartbeatStore_Heartbeat(t *testing.T) { - tests := []struct { - name string - setup func(mock sqlmock.Sqlmock) - wantErr bool - }{ - { - name: "successfully send heartbeat", - setup: func(mock sqlmock.Sqlmock) { - mock.ExpectExec("INSERT INTO queue_subscriber_heartbeats"). - WithArgs(testTenant, testConsumerGroup, "test_topic", testSubscriberName, sqlmock.AnyArg()). - WillReturnResult(sqlmock.NewResult(1, 1)) - }, - wantErr: false, - }, - { - name: "update existing heartbeat", - setup: func(mock sqlmock.Sqlmock) { - mock.ExpectExec("INSERT INTO queue_subscriber_heartbeats"). - WithArgs(testTenant, testConsumerGroup, "test_topic", testSubscriberName, sqlmock.AnyArg()). - WillReturnResult(sqlmock.NewResult(0, 2)) // ON DUPLICATE KEY UPDATE returns 2 for update - }, - wantErr: false, - }, - { - name: "database error", - setup: func(mock sqlmock.Sqlmock) { - mock.ExpectExec("INSERT INTO queue_subscriber_heartbeats"). - WithArgs(testTenant, testConsumerGroup, "test_topic", testSubscriberName, sqlmock.AnyArg()). - WillReturnError(fmt.Errorf("db error")) - }, - wantErr: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - db, mock, store := setupSubscriberHeartbeatStoreTest(t) - defer db.Close() - - ctx := context.Background() - tt.setup(mock) - - err := store.Heartbeat(ctx, testTenant, "test_topic", testSubscriberName, testConsumerGroup) - if tt.wantErr { - require.Error(t, err) - } else { - require.NoError(t, err) - } - require.NoError(t, mock.ExpectationsWereMet()) - }) - } -} - -func TestSubscriberHeartbeatStore_ActiveSubscribers(t *testing.T) { - tests := []struct { - name string - setup func(mock sqlmock.Sqlmock) - wantNames []string - wantErr bool - }{ - { - name: "multiple active subscribers", - setup: func(mock sqlmock.Sqlmock) { - rows := sqlmock.NewRows([]string{"subscriber_name"}). - AddRow("sub-1").AddRow("sub-2").AddRow("sub-3") - mock.ExpectQuery("SELECT subscriber_name"). - WithArgs(testTenant, testConsumerGroup, "test_topic", sqlmock.AnyArg()). - WillReturnRows(rows) - }, - wantNames: []string{"sub-1", "sub-2", "sub-3"}, - wantErr: false, - }, - { - name: "no active subscribers", - setup: func(mock sqlmock.Sqlmock) { - rows := sqlmock.NewRows([]string{"subscriber_name"}) - mock.ExpectQuery("SELECT subscriber_name"). - WithArgs(testTenant, testConsumerGroup, "test_topic", sqlmock.AnyArg()). - WillReturnRows(rows) - }, - wantNames: nil, - wantErr: false, - }, - { - name: "database error", - setup: func(mock sqlmock.Sqlmock) { - mock.ExpectQuery("SELECT subscriber_name"). - WithArgs(testTenant, testConsumerGroup, "test_topic", sqlmock.AnyArg()). - WillReturnError(fmt.Errorf("db error")) - }, - wantNames: nil, - wantErr: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - db, mock, store := setupSubscriberHeartbeatStoreTest(t) - defer db.Close() - - ctx := context.Background() - tt.setup(mock) - - names, err := store.ActiveSubscribers(ctx, testTenant, "test_topic", testConsumerGroup, testLeaseDurationMs) - if tt.wantErr { - require.Error(t, err) - } else { - require.NoError(t, err) - require.Equal(t, tt.wantNames, names) - } - require.NoError(t, mock.ExpectationsWereMet()) - }) - } -} - -func TestSubscriberHeartbeatStore_ActiveSubscribers_ExcludesDeregistered(t *testing.T) { +func TestSubscriberHeartbeatStore_ForTenants(t *testing.T) { db, mock, store := setupSubscriberHeartbeatStoreTest(t) defer db.Close() + tenants := []string{"alpha", "beta"} ctx := context.Background() - // Verify the query filters by deregistered_at = 0 - rows := sqlmock.NewRows([]string{"subscriber_name"}).AddRow("sub-1").AddRow("sub-2") - mock.ExpectQuery(`SELECT subscriber_name FROM queue_subscriber_heartbeats.*deregistered_at = 0`). - WithArgs(testTenant, testConsumerGroup, "test_topic", sqlmock.AnyArg()). - WillReturnRows(rows) + mock.ExpectExec("INSERT INTO queue_subscriber_heartbeats"). + WithArgs( + "alpha", testConsumerGroup, "test_topic", testSubscriberName, sqlmock.AnyArg(), + "beta", testConsumerGroup, "test_topic", testSubscriberName, sqlmock.AnyArg(), + ). + WillReturnResult(sqlmock.NewResult(2, 2)) + mock.ExpectQuery(`(?s)SELECT tenant, subscriber_name FROM queue_subscriber_heartbeats.*deregistered_at = 0`). + WithArgs("alpha", "beta", testConsumerGroup, "test_topic", sqlmock.AnyArg()). + WillReturnRows(sqlmock.NewRows([]string{"tenant", "subscriber_name"}). + AddRow("alpha", "s1"). + AddRow("alpha", "s2"). + AddRow("beta", "s1")) + mock.ExpectExec("DELETE FROM queue_subscriber_heartbeats"). + WithArgs("alpha", "beta", testConsumerGroup, "test_topic", testSubscriberName). + WillReturnResult(sqlmock.NewResult(0, 2)) + mock.ExpectExec("DELETE FROM queue_subscriber_heartbeats"). + WithArgs("alpha", "beta", testConsumerGroup, "test_topic", sqlmock.AnyArg()). + WillReturnResult(sqlmock.NewResult(0, 1)) - names, err := store.ActiveSubscribers(ctx, testTenant, "test_topic", testConsumerGroup, testLeaseDurationMs) + require.NoError(t, store.HeartbeatForTenants(ctx, tenants, "test_topic", testSubscriberName, testConsumerGroup)) + active, err := store.ActiveSubscribersForTenants(ctx, tenants, "test_topic", testConsumerGroup, testLeaseDurationMs) require.NoError(t, err) - require.Equal(t, []string{"sub-1", "sub-2"}, names) + require.Equal(t, map[string][]string{"alpha": {"s1", "s2"}, "beta": {"s1"}}, active) + require.NoError(t, store.DeregisterForTenants(ctx, tenants, "test_topic", testSubscriberName, testConsumerGroup)) + require.NoError(t, store.PurgeStaleForTenants(ctx, tenants, "test_topic", testConsumerGroup, 300_000)) require.NoError(t, mock.ExpectationsWereMet()) } -func TestSubscriberHeartbeatStore_Deregister_HardDelete(t *testing.T) { +func TestSubscriberHeartbeatStore_ForTenants_EmptyIsNoop(t *testing.T) { db, mock, store := setupSubscriberHeartbeatStoreTest(t) defer db.Close() ctx := context.Background() - - // Verify deregister deletes the row outright — subscriber names are - // unique per process, so soft-deleted rows would accumulate forever. - mock.ExpectExec(`DELETE FROM queue_subscriber_heartbeats`). - WithArgs(testTenant, testConsumerGroup, "test_topic", testSubscriberName). - WillReturnResult(sqlmock.NewResult(0, 1)) - - err := store.Deregister(ctx, testTenant, "test_topic", testSubscriberName, testConsumerGroup) + require.NoError(t, store.HeartbeatForTenants(ctx, nil, "test_topic", testSubscriberName, testConsumerGroup)) + active, err := store.ActiveSubscribersForTenants(ctx, nil, "test_topic", testConsumerGroup, testLeaseDurationMs) require.NoError(t, err) + require.Empty(t, active) + require.NoError(t, store.DeregisterForTenants(ctx, nil, "test_topic", testSubscriberName, testConsumerGroup)) + require.NoError(t, store.PurgeStaleForTenants(ctx, nil, "test_topic", testConsumerGroup, 300_000)) require.NoError(t, mock.ExpectationsWereMet()) } -func TestSubscriberHeartbeatStore_PurgeStale(t *testing.T) { - tests := []struct { - name string - setup func(mock sqlmock.Sqlmock) - wantErr bool - }{ - { - name: "deletes rows older than threshold", - setup: func(mock sqlmock.Sqlmock) { - mock.ExpectExec(`DELETE FROM queue_subscriber_heartbeats`). - WithArgs(testTenant, testConsumerGroup, "test_topic", sqlmock.AnyArg()). - WillReturnResult(sqlmock.NewResult(0, 3)) - }, - }, - { - name: "no stale rows is a no-op", - setup: func(mock sqlmock.Sqlmock) { - mock.ExpectExec(`DELETE FROM queue_subscriber_heartbeats`). - WithArgs(testTenant, testConsumerGroup, "test_topic", sqlmock.AnyArg()). - WillReturnResult(sqlmock.NewResult(0, 0)) - }, - }, - { - name: "database error", - setup: func(mock sqlmock.Sqlmock) { - mock.ExpectExec(`DELETE FROM queue_subscriber_heartbeats`). - WithArgs(testTenant, testConsumerGroup, "test_topic", sqlmock.AnyArg()). - WillReturnError(fmt.Errorf("db error")) - }, - wantErr: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - db, mock, store := setupSubscriberHeartbeatStoreTest(t) - defer db.Close() - - tt.setup(mock) - - err := store.PurgeStale(context.Background(), testTenant, "test_topic", testConsumerGroup, 300_000) - if tt.wantErr { - require.Error(t, err) - } else { - require.NoError(t, err) - } - require.NoError(t, mock.ExpectationsWereMet()) - }) - } -} - -func TestSubscriberHeartbeatStore_ReRegistration(t *testing.T) { +func TestSubscriberHeartbeatStore_ForTenants_ReregisterAfterDelete(t *testing.T) { db, mock, store := setupSubscriberHeartbeatStoreTest(t) defer db.Close() + tenants := []string{"alpha"} ctx := context.Background() - - // Step 1: Initial heartbeat registers the subscriber mock.ExpectExec("INSERT INTO queue_subscriber_heartbeats"). - WithArgs(testTenant, testConsumerGroup, "test_topic", testSubscriberName, sqlmock.AnyArg()). + WithArgs("alpha", testConsumerGroup, "test_topic", testSubscriberName, sqlmock.AnyArg()). WillReturnResult(sqlmock.NewResult(1, 1)) - - // Step 2: Deregister deletes the subscriber's row mock.ExpectExec("DELETE FROM queue_subscriber_heartbeats"). - WithArgs(testTenant, testConsumerGroup, "test_topic", testSubscriberName). + WithArgs("alpha", testConsumerGroup, "test_topic", testSubscriberName). WillReturnResult(sqlmock.NewResult(0, 1)) - - // Step 3: Heartbeat again re-registers with a fresh insert mock.ExpectExec("INSERT INTO queue_subscriber_heartbeats"). - WithArgs(testTenant, testConsumerGroup, "test_topic", testSubscriberName, sqlmock.AnyArg()). + WithArgs("alpha", testConsumerGroup, "test_topic", testSubscriberName, sqlmock.AnyArg()). WillReturnResult(sqlmock.NewResult(1, 1)) - err := store.Heartbeat(ctx, testTenant, "test_topic", testSubscriberName, testConsumerGroup) - require.NoError(t, err) - - err = store.Deregister(ctx, testTenant, "test_topic", testSubscriberName, testConsumerGroup) - require.NoError(t, err) - - err = store.Heartbeat(ctx, testTenant, "test_topic", testSubscriberName, testConsumerGroup) - require.NoError(t, err) - + require.NoError(t, store.HeartbeatForTenants(ctx, tenants, "test_topic", testSubscriberName, testConsumerGroup)) + require.NoError(t, store.DeregisterForTenants(ctx, tenants, "test_topic", testSubscriberName, testConsumerGroup)) + require.NoError(t, store.HeartbeatForTenants(ctx, tenants, "test_topic", testSubscriberName, testConsumerGroup)) require.NoError(t, mock.ExpectationsWereMet()) } -func TestSubscriberHeartbeatStore_Deregister(t *testing.T) { +func TestSubscriberHeartbeatStore_ForTenants_Errors(t *testing.T) { + tenants := []string{"alpha"} + dbErr := fmt.Errorf("db error") tests := []struct { - name string - setup func(mock sqlmock.Sqlmock) - wantErr bool + name string + setup func(sqlmock.Sqlmock) + call func(subscriberHeartbeatStore) error }{ { - name: "successfully deregister", + name: "heartbeat", setup: func(mock sqlmock.Sqlmock) { - mock.ExpectExec("DELETE FROM queue_subscriber_heartbeats"). - WithArgs(testTenant, testConsumerGroup, "test_topic", testSubscriberName). - WillReturnResult(sqlmock.NewResult(0, 1)) + mock.ExpectExec("INSERT INTO queue_subscriber_heartbeats"). + WillReturnError(dbErr) + }, + call: func(store subscriberHeartbeatStore) error { + return store.HeartbeatForTenants(context.Background(), tenants, "test_topic", testSubscriberName, testConsumerGroup) + }, + }, + { + name: "active subscribers", + setup: func(mock sqlmock.Sqlmock) { + mock.ExpectQuery(`(?s)SELECT tenant, subscriber_name FROM queue_subscriber_heartbeats.*deregistered_at = 0`). + WillReturnError(dbErr) + }, + call: func(store subscriberHeartbeatStore) error { + _, err := store.ActiveSubscribersForTenants(context.Background(), tenants, "test_topic", testConsumerGroup, testLeaseDurationMs) + return err }, - wantErr: false, }, { - name: "idempotent - already deregistered", + name: "deregister", setup: func(mock sqlmock.Sqlmock) { mock.ExpectExec("DELETE FROM queue_subscriber_heartbeats"). - WithArgs(testTenant, testConsumerGroup, "test_topic", testSubscriberName). - WillReturnResult(sqlmock.NewResult(0, 0)) + WillReturnError(dbErr) + }, + call: func(store subscriberHeartbeatStore) error { + return store.DeregisterForTenants(context.Background(), tenants, "test_topic", testSubscriberName, testConsumerGroup) }, - wantErr: false, }, { - name: "database error", + name: "purge stale", setup: func(mock sqlmock.Sqlmock) { mock.ExpectExec("DELETE FROM queue_subscriber_heartbeats"). - WithArgs(testTenant, testConsumerGroup, "test_topic", testSubscriberName). - WillReturnError(fmt.Errorf("db error")) + WillReturnError(dbErr) + }, + call: func(store subscriberHeartbeatStore) error { + return store.PurgeStaleForTenants(context.Background(), tenants, "test_topic", testConsumerGroup, 300_000) }, - wantErr: true, }, } - for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { db, mock, store := setupSubscriberHeartbeatStoreTest(t) defer db.Close() - - ctx := context.Background() tt.setup(mock) - - err := store.Deregister(ctx, testTenant, "test_topic", testSubscriberName, testConsumerGroup) - if tt.wantErr { - require.Error(t, err) - } else { - require.NoError(t, err) - } + require.Error(t, tt.call(store)) require.NoError(t, mock.ExpectationsWereMet()) }) } diff --git a/platform/extension/messagequeue/mysql/subscriber_test.go b/platform/extension/messagequeue/mysql/subscriber_test.go index 834d8176..60b6e51a 100644 --- a/platform/extension/messagequeue/mysql/subscriber_test.go +++ b/platform/extension/messagequeue/mysql/subscriber_test.go @@ -52,78 +52,26 @@ func testSubscriptionConfig() extqueue.SubscriptionConfig { return extqueue.DefaultSubscriptionConfig("test-subscriber", "test-consumer") } -func TestRunTenantOperationsConcurrentlyWithDeadlines(t *testing.T) { - slowStarted := make(chan struct{}) - releaseSlow := make(chan struct{}) - fastCompleted := make(chan struct{}) - operationCompleted := make(chan []tenantOperationResult[string]) - - go func() { - operationCompleted <- runTenantOperations( - context.Background(), - []string{"slow", "fast"}, - time.Hour, - func(ctx context.Context, tenant string) (string, error) { - _, hasDeadline := ctx.Deadline() - if !hasDeadline { - return "", errors.New("tenant operation has no deadline") - } - if tenant == "slow" { - close(slowStarted) - <-releaseSlow - } else { - close(fastCompleted) - } - return tenant, nil - }, - ) - }() - - <-slowStarted - <-fastCompleted - close(releaseSlow) - results := <-operationCompleted - assert.Equal(t, []tenantOperationResult[string]{ - {tenant: "slow", value: "slow"}, - {tenant: "fast", value: "fast"}, - }, results) -} - -func TestRunTenantOperationsPropagatesCancellation(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - started := make(chan string, 2) - operationCompleted := make(chan []tenantOperationResult[struct{}]) - - go func() { - operationCompleted <- runTenantOperations( - ctx, - []string{"tenant-a", "tenant-b"}, - time.Hour, - func(ctx context.Context, tenant string) (struct{}, error) { - started <- tenant - <-ctx.Done() - return struct{}{}, ctx.Err() - }, - ) - }() - - assert.ElementsMatch(t, []string{"tenant-a", "tenant-b"}, []string{<-started, <-started}) - cancel() - for _, result := range <-operationCompleted { - assert.ErrorIs(t, result.err, context.Canceled) - } -} - // newTestHeartbeatStore creates a mock heartbeat store that allows all calls func newTestHeartbeatStore(ctrl *gomock.Controller) *MocksubscriberHeartbeatStore { mockHB := NewMocksubscriberHeartbeatStore(ctrl) - mockHB.EXPECT().Heartbeat(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil).AnyTimes() - mockHB.EXPECT().ActiveSubscribers(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return([]string{"self"}, nil).AnyTimes() - mockHB.EXPECT().Deregister(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil).AnyTimes() - mockHB.EXPECT().PurgeStale(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil).AnyTimes() + mockHB.EXPECT().HeartbeatForTenants(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil).AnyTimes() + mockHB.EXPECT().ActiveSubscribersForTenants(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(map[string][]string{}, nil).AnyTimes() + mockHB.EXPECT().DeregisterForTenants(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil).AnyTimes() + mockHB.EXPECT().PurgeStaleForTenants(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil).AnyTimes() return mockHB } +func allowSupervisorLeaseCalls(mockLeaseStore *MockpartitionLeaseStore) { + mockLeaseStore.EXPECT().GetLeasedPartitionsForTenants(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(map[string][]string{}, nil).AnyTimes() + mockLeaseStore.EXPECT().DiscoverPartitions(gomock.Any(), gomock.Any(), gomock.Any()).Return(map[string][]string{}, nil).AnyTimes() + mockLeaseStore.EXPECT().GetAllLeasesForTenants(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(map[string][]leaseInfo{}, nil).AnyTimes() + mockLeaseStore.EXPECT().RenewOwnedLeases(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil).AnyTimes() + mockLeaseStore.EXPECT().ReleaseOwnedLeases(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil).AnyTimes() + mockLeaseStore.EXPECT().PurgeStaleForTenants(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(nil).AnyTimes() + mockLeaseStore.EXPECT().TryAcquireLease(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(false, nil).AnyTimes() +} + // newTestDeliveryStateStore creates a mock delivery state store that allows all calls func newTestDeliveryStateStore(ctrl *gomock.Controller) *MockdeliveryStateStore { mockDS := NewMockdeliveryStateStore(ctrl) @@ -181,9 +129,7 @@ func TestSubscriber_Subscribe(t *testing.T) { mockOffsetStore := NewMockoffsetStore(ctrl) mockLeaseStore := NewMockpartitionLeaseStore(ctrl) - // Reached via releaseAllLeases on the shutdown path, and by the - // discovery ticker if it fires before teardown. - mockLeaseStore.EXPECT().GetLeasedPartitions(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return([]string{}, nil).AnyTimes() + allowSupervisorLeaseCalls(mockLeaseStore) sub := setupSubscriberTest(t, mockMessageStore, mockOffsetStore, mockLeaseStore) // Close waits for managePartitions to exit; a bare cancel would only @@ -294,7 +240,7 @@ func TestSubscriber_SubscribeContextCancellation(t *testing.T) { mockMessageStore := NewMockmessageStore(ctrl) mockOffsetStore := NewMockoffsetStore(ctrl) mockLeaseStore := NewMockpartitionLeaseStore(ctrl) - mockLeaseStore.EXPECT().GetLeasedPartitions(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return([]string{}, nil).AnyTimes() + allowSupervisorLeaseCalls(mockLeaseStore) sub := setupSubscriberTest(t, mockMessageStore, mockOffsetStore, mockLeaseStore) defer func() { @@ -320,7 +266,7 @@ func TestSubscriber_SubscribeReplacesStaleSubscription(t *testing.T) { mockMessageStore := NewMockmessageStore(ctrl) mockOffsetStore := NewMockoffsetStore(ctrl) mockLeaseStore := NewMockpartitionLeaseStore(ctrl) - mockLeaseStore.EXPECT().GetLeasedPartitions(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return([]string{}, nil).AnyTimes() + allowSupervisorLeaseCalls(mockLeaseStore) sub := setupSubscriberTest(t, mockMessageStore, mockOffsetStore, mockLeaseStore) defer func() { @@ -762,7 +708,7 @@ func TestSubscriber_Close(t *testing.T) { mockLeaseStore := NewMockpartitionLeaseStore(ctrl) // Expect lease operations during cleanup - mockLeaseStore.EXPECT().GetLeasedPartitions(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return([]string{}, nil).AnyTimes() + allowSupervisorLeaseCalls(mockLeaseStore) sub := setupSubscriberTest(t, mockMessageStore, mockOffsetStore, mockLeaseStore) ctx := context.Background() @@ -992,46 +938,18 @@ func TestSubscriber_ReconcilePartitionWorkersKeepsTenantIdentity(t *testing.T) { s.stopAllWorkers(sub) } -func TestSubscriber_DiscoverAndReconcileWorkersIsolatesTenantFailures(t *testing.T) { - const ( - tenantBefore = "tenant-before" - tenantFailed = "tenant-failed" - tenantAfter = "tenant-after" - tenantFailedLast = "tenant-failed-last" - ) - +func TestSubscriber_DiscoverFailureKeepsCachedDiscovery(t *testing.T) { + const tenantFailed = "tenant-failed" ctrl := gomock.NewController(t) mockLeaseStore := NewMockpartitionLeaseStore(ctrl) discoveryErr := errors.New("tenant store unavailable") - lastDiscoveryErr := errors.New("last tenant store unavailable") cfg := testSubscriptionConfig() cfg.PollIntervalMs = int64(time.Hour / time.Millisecond) mockLeaseStore.EXPECT(). - GetLeasedPartitions(gomock.Any(), tenantBefore, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). - Return(nil, nil) - mockLeaseStore.EXPECT(). - DiscoverAndAcquirePartitions(gomock.Any(), tenantBefore, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup, cfg.LeaseDurationMs, 0). - Return(1, []string{"before-new"}, nil) - mockLeaseStore.EXPECT(). - GetLeasedPartitions(gomock.Any(), tenantBefore, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). - Return([]string{"before-new"}, nil) - mockLeaseStore.EXPECT(). - GetLeasedPartitions(gomock.Any(), tenantFailed, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + DiscoverPartitions(gomock.Any(), []string{tenantFailed}, "test-topic"). Return(nil, discoveryErr) - mockLeaseStore.EXPECT(). - GetLeasedPartitions(gomock.Any(), tenantAfter, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). - Return(nil, nil) - mockLeaseStore.EXPECT(). - DiscoverAndAcquirePartitions(gomock.Any(), tenantAfter, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup, cfg.LeaseDurationMs, 0). - Return(1, []string{"after-new"}, nil) - mockLeaseStore.EXPECT(). - GetLeasedPartitions(gomock.Any(), tenantAfter, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). - Return([]string{"after-new"}, nil) - mockLeaseStore.EXPECT(). - GetLeasedPartitions(gomock.Any(), tenantFailedLast, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). - Return(nil, lastDiscoveryErr) s := NewSubscriber( zaptest.NewLogger(t).Sugar(), @@ -1041,13 +959,13 @@ func TestSubscriber_DiscoverAndReconcileWorkersIsolatesTenantFailures(t *testing mockLeaseStore, newTestHeartbeatStore(ctrl), newTestDeliveryStateStore(ctrl), - []string{tenantBefore, tenantFailed, tenantAfter, tenantFailedLast}, + []string{tenantFailed}, ) failedWorkerDone := make(chan struct{}) close(failedWorkerDone) failedWorkerKey := entityqueue.PartitionIdentity{Tenant: tenantFailed, PartitionKey: "failed-existing"} - failedDrainSince := time.Now().Add(-time.Hour) + cached := []entityqueue.PartitionIdentity{{Tenant: tenantFailed, PartitionKey: "failed-discovered"}} sub := &subscription{ topic: "test-topic", config: cfg, @@ -1058,39 +976,14 @@ func TestSubscriber_DiscoverAndReconcileWorkersIsolatesTenantFailures(t *testing done: failedWorkerDone, }, }, - lastDiscoveredPartitions: []entityqueue.PartitionIdentity{{Tenant: tenantFailed, PartitionKey: "failed-discovered"}}, - drainedSince: map[entityqueue.PartitionIdentity]time.Time{failedWorkerKey: failedDrainSince}, + lastDiscoveredPartitions: cached, + drainedSince: map[entityqueue.PartitionIdentity]time.Time{failedWorkerKey: time.Now().Add(-time.Hour)}, } - ctx, cancel := context.WithCancel(context.Background()) - t.Cleanup(func() { - cancel() - s.stopAllWorkers(sub) - sub.workerWg.Wait() - }) - - err := s.discoverAndReconcileWorkers(ctx, sub, true) + err := s.discoverAndReconcileWorkers(context.Background(), sub, true) require.ErrorIs(t, err, discoveryErr) - require.ErrorIs(t, err, lastDiscoveryErr) - - sub.workersMu.Lock() - workerKeys := make([]entityqueue.PartitionIdentity, 0, len(sub.workers)) - for key := range sub.workers { - workerKeys = append(workerKeys, key) - } - discovered := append([]entityqueue.PartitionIdentity(nil), sub.lastDiscoveredPartitions...) - sub.workersMu.Unlock() - - assert.ElementsMatch(t, []entityqueue.PartitionIdentity{ - {Tenant: tenantBefore, PartitionKey: "before-new"}, - {Tenant: tenantAfter, PartitionKey: "after-new"}, - }, workerKeys) - assert.ElementsMatch(t, []entityqueue.PartitionIdentity{ - {Tenant: tenantBefore, PartitionKey: "before-new"}, - {Tenant: tenantFailed, PartitionKey: "failed-discovered"}, - {Tenant: tenantAfter, PartitionKey: "after-new"}, - }, discovered) - assert.NotContains(t, sub.drainedSince, failedWorkerKey) + assert.NotContains(t, sub.workers, failedWorkerKey) + assert.Equal(t, cached, sub.lastDiscoveredPartitions) } func TestSubscriber_DiscoverFailureStopsUnconfirmedWorkers(t *testing.T) { @@ -1102,7 +995,7 @@ func TestSubscriber_DiscoverFailureStopsUnconfirmedWorkers(t *testing.T) { discoveryErr := errors.New("tenant store unavailable") mockLeaseStore.EXPECT(). - GetLeasedPartitions(gomock.Any(), tenant, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + DiscoverPartitions(gomock.Any(), []string{tenant}, "test-topic"). Return(nil, discoveryErr) s := NewSubscriber( @@ -1136,6 +1029,422 @@ func TestSubscriber_DiscoverFailureStopsUnconfirmedWorkers(t *testing.T) { assert.NotContains(t, sub.workers, failedWorkerKey) } +func TestSubscriber_GetLeasedPartitionsFailureKeepsCachedDiscovery(t *testing.T) { + ctrl := gomock.NewController(t) + mockLeaseStore := NewMockpartitionLeaseStore(ctrl) + cfg := testSubscriptionConfig() + cfg.PollIntervalMs = int64(time.Hour / time.Millisecond) + leasedErr := errors.New("leased partitions unavailable") + cached := []entityqueue.PartitionIdentity{{Tenant: testTenant, PartitionKey: "cached"}} + failedWorkerKey := entityqueue.PartitionIdentity{Tenant: testTenant, PartitionKey: "existing"} + + mockLeaseStore.EXPECT(). + DiscoverPartitions(gomock.Any(), []string{testTenant}, "test-topic"). + Return(map[string][]string{testTenant: {"discovered"}}, nil) + mockLeaseStore.EXPECT(). + GetLeasedPartitionsForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + Return(nil, leasedErr) + + s := NewSubscriber( + zaptest.NewLogger(t).Sugar(), + tally.NoopScope, + NewMockmessageStore(ctrl), + NewMockoffsetStore(ctrl), + mockLeaseStore, + newTestHeartbeatStore(ctrl), + newTestDeliveryStateStore(ctrl), + []string{testTenant}, + ) + failedWorkerDone := make(chan struct{}) + close(failedWorkerDone) + sub := &subscription{ + topic: "test-topic", + config: cfg, + deliveryCh: make(chan extqueue.Delivery, 1), + workers: map[entityqueue.PartitionIdentity]*partitionWorker{ + failedWorkerKey: { + cancelFunc: func() {}, + done: failedWorkerDone, + }, + }, + lastDiscoveredPartitions: cached, + } + + err := s.discoverAndReconcileWorkers(context.Background(), sub, true) + require.ErrorIs(t, err, leasedErr) + assert.NotContains(t, sub.workers, failedWorkerKey) + assert.Equal(t, cached, sub.lastDiscoveredPartitions) +} + +func TestSubscriber_AcquireUnownedPartitions(t *testing.T) { + cfg := testSubscriptionConfig() + freshMs := time.Now().UnixMilli() + staleMs := freshMs - cfg.LeaseDurationMs - 60_000 + + expectAcquire := func(mockLeaseStore *MockpartitionLeaseStore, pk string, acquired bool) { + mockLeaseStore.EXPECT(). + TryAcquireLease(gomock.Any(), testTenant, "test-topic", pk, cfg.SubscriberName, cfg.ConsumerGroup, cfg.LeaseDurationMs). + Return(acquired, nil) + } + + tests := []struct { + name string + discovered []string + leases []leaseInfo + maxPartitions int + setup func(*MockpartitionLeaseStore) + want []string + }{ + { + name: "acquires unleased, skips fresh lease held by other", + discovered: []string{"part1", "part2"}, + leases: []leaseInfo{ + {PartitionKey: "part2", LeasedBy: "other-worker", LeaseRenewedAt: freshMs}, + }, + setup: func(mockLeaseStore *MockpartitionLeaseStore) { + expectAcquire(mockLeaseStore, "part1", true) + }, + want: []string{"part1"}, + }, + { + name: "stale lease held by other is stealable", + discovered: []string{"part1"}, + leases: []leaseInfo{ + {PartitionKey: "part1", LeasedBy: "other-worker", LeaseRenewedAt: staleMs}, + }, + setup: func(mockLeaseStore *MockpartitionLeaseStore) { + expectAcquire(mockLeaseStore, "part1", true) + }, + want: []string{"part1"}, + }, + { + name: "self-owned partitions are not re-probed", + discovered: []string{"part1", "part2"}, + leases: []leaseInfo{ + {PartitionKey: "part1", LeasedBy: cfg.SubscriberName, LeaseRenewedAt: freshMs}, + }, + setup: func(mockLeaseStore *MockpartitionLeaseStore) { + expectAcquire(mockLeaseStore, "part2", true) + }, + want: []string{"part2"}, + }, + { + name: "stops acquiring when cap reached", + discovered: []string{"part1", "part2", "part3"}, + maxPartitions: 2, + setup: func(mockLeaseStore *MockpartitionLeaseStore) { + expectAcquire(mockLeaseStore, "part1", true) + expectAcquire(mockLeaseStore, "part2", true) + }, + want: []string{"part1", "part2"}, + }, + { + name: "pre-owned partitions count toward cap", + discovered: []string{"part1", "part2", "part3"}, + maxPartitions: 3, + leases: []leaseInfo{ + {PartitionKey: "existing1", LeasedBy: cfg.SubscriberName, LeaseRenewedAt: freshMs}, + {PartitionKey: "existing2", LeasedBy: cfg.SubscriberName, LeaseRenewedAt: freshMs}, + }, + setup: func(mockLeaseStore *MockpartitionLeaseStore) { + expectAcquire(mockLeaseStore, "part1", true) + }, + want: []string{"part1"}, + }, + { + name: "already at cap acquires nothing", + discovered: []string{"part1", "part2"}, + maxPartitions: 2, + leases: []leaseInfo{ + {PartitionKey: "existing1", LeasedBy: cfg.SubscriberName, LeaseRenewedAt: freshMs}, + {PartitionKey: "existing2", LeasedBy: cfg.SubscriberName, LeaseRenewedAt: freshMs}, + }, + want: nil, + }, + { + name: "lost race counts nothing", + discovered: []string{"part1"}, + setup: func(mockLeaseStore *MockpartitionLeaseStore) { + expectAcquire(mockLeaseStore, "part1", false) + }, + want: nil, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctrl := gomock.NewController(t) + mockLeaseStore := NewMockpartitionLeaseStore(ctrl) + if tt.setup != nil { + tt.setup(mockLeaseStore) + } + s := NewSubscriber( + zaptest.NewLogger(t).Sugar(), + tally.NoopScope, + NewMockmessageStore(ctrl), + NewMockoffsetStore(ctrl), + mockLeaseStore, + NewMocksubscriberHeartbeatStore(ctrl), + NewMockdeliveryStateStore(ctrl), + []string{testTenant}, + ) + sub := &subscription{topic: "test-topic", config: cfg} + got := s.acquireUnownedPartitions(context.Background(), sub, testTenant, tt.discovered, tt.leases, tt.maxPartitions) + assert.Equal(t, tt.want, got) + }) + } +} + +func TestSubscriber_LeaseTickRenewsWhenActiveSubscribersFails(t *testing.T) { + ctrl := gomock.NewController(t) + mockLeaseStore := NewMockpartitionLeaseStore(ctrl) + mockHeartbeatStore := NewMocksubscriberHeartbeatStore(ctrl) + cfg := testSubscriptionConfig() + activeErr := errors.New("active subscribers unavailable") + + mockLeaseStore.EXPECT(). + GetLeasedPartitionsForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + Return(map[string][]string{testTenant: {"part-1"}}, nil) + mockHeartbeatStore.EXPECT(). + ActiveSubscribersForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.ConsumerGroup, cfg.LeaseDurationMs). + Return(nil, activeErr) + mockLeaseStore.EXPECT(). + RenewOwnedLeases(gomock.Any(), []string{testTenant}, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + Return(nil) + mockHeartbeatStore.EXPECT(). + HeartbeatForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + Return(nil) + mockHeartbeatStore.EXPECT(). + PurgeStaleForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.ConsumerGroup, heartbeatPurgeAfterLeaseDurations*cfg.LeaseDurationMs). + Return(nil) + mockLeaseStore.EXPECT(). + PurgeStaleForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.ConsumerGroup, leasePurgeAfterLeaseDurations*cfg.LeaseDurationMs). + Return(nil) + + s := NewSubscriber( + zaptest.NewLogger(t).Sugar(), + tally.NoopScope, + NewMockmessageStore(ctrl), + NewMockoffsetStore(ctrl), + mockLeaseStore, + mockHeartbeatStore, + newTestDeliveryStateStore(ctrl), + []string{testTenant}, + ) + sub := &subscription{ + topic: "test-topic", + config: cfg, + deliveryCh: make(chan extqueue.Delivery, 1), + workers: make(map[entityqueue.PartitionIdentity]*partitionWorker), + } + + s.runLeaseTick(context.Background(), sub, time.Second, []any{"topic", sub.topic}) +} + +func TestSubscriber_LeaseTickRenewsWhenGetLeasedPartitionsFails(t *testing.T) { + ctrl := gomock.NewController(t) + mockLeaseStore := NewMockpartitionLeaseStore(ctrl) + mockHeartbeatStore := NewMocksubscriberHeartbeatStore(ctrl) + cfg := testSubscriptionConfig() + leasedErr := errors.New("leased partitions unavailable") + + mockLeaseStore.EXPECT(). + GetLeasedPartitionsForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + Return(nil, leasedErr) + mockLeaseStore.EXPECT(). + RenewOwnedLeases(gomock.Any(), []string{testTenant}, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + Return(nil) + mockHeartbeatStore.EXPECT(). + HeartbeatForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + Return(nil) + mockHeartbeatStore.EXPECT(). + PurgeStaleForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.ConsumerGroup, heartbeatPurgeAfterLeaseDurations*cfg.LeaseDurationMs). + Return(nil) + mockLeaseStore.EXPECT(). + PurgeStaleForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.ConsumerGroup, leasePurgeAfterLeaseDurations*cfg.LeaseDurationMs). + Return(nil) + + s := NewSubscriber( + zaptest.NewLogger(t).Sugar(), + tally.NoopScope, + NewMockmessageStore(ctrl), + NewMockoffsetStore(ctrl), + mockLeaseStore, + mockHeartbeatStore, + newTestDeliveryStateStore(ctrl), + []string{testTenant}, + ) + sub := &subscription{ + topic: "test-topic", + config: cfg, + deliveryCh: make(chan extqueue.Delivery, 1), + workers: make(map[entityqueue.PartitionIdentity]*partitionWorker), + } + + s.runLeaseTick(context.Background(), sub, time.Second, []any{"topic", sub.topic}) +} + +func TestSubscriber_ActiveSubscribersFailureKeepsWorkers(t *testing.T) { + ctrl := gomock.NewController(t) + mockLeaseStore := NewMockpartitionLeaseStore(ctrl) + mockHeartbeatStore := NewMocksubscriberHeartbeatStore(ctrl) + cfg := testSubscriptionConfig() + cfg.PollIntervalMs = int64(time.Hour / time.Millisecond) + activeErr := errors.New("active subscribers unavailable") + partition := entityqueue.PartitionIdentity{Tenant: testTenant, PartitionKey: "p1"} + + mockLeaseStore.EXPECT(). + DiscoverPartitions(gomock.Any(), []string{testTenant}, "test-topic"). + Return(map[string][]string{testTenant: {partition.PartitionKey}}, nil) + mockLeaseStore.EXPECT(). + GetLeasedPartitionsForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + Return(map[string][]string{testTenant: {partition.PartitionKey}}, nil) + mockLeaseStore.EXPECT(). + GetAllLeasesForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.ConsumerGroup). + Return(map[string][]leaseInfo{}, nil) + mockHeartbeatStore.EXPECT(). + ActiveSubscribersForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.ConsumerGroup, cfg.LeaseDurationMs). + Return(nil, activeErr) + + s := NewSubscriber( + zaptest.NewLogger(t).Sugar(), + tally.NoopScope, + NewMockmessageStore(ctrl), + NewMockoffsetStore(ctrl), + mockLeaseStore, + mockHeartbeatStore, + newTestDeliveryStateStore(ctrl), + []string{testTenant}, + ) + failedWorkerDone := make(chan struct{}) + close(failedWorkerDone) + sub := &subscription{ + topic: "test-topic", + config: cfg, + deliveryCh: make(chan extqueue.Delivery, 1), + workers: map[entityqueue.PartitionIdentity]*partitionWorker{ + partition: { + cancelFunc: func() {}, + done: failedWorkerDone, + }, + }, + } + + err := s.discoverAndReconcileWorkers(context.Background(), sub, false) + require.ErrorIs(t, err, activeErr) + assert.Contains(t, sub.workers, partition) + assert.Equal(t, []entityqueue.PartitionIdentity{partition}, sub.lastDiscoveredPartitions) +} + +func TestSubscriber_GetAllLeasesFailureKeepsWorkers(t *testing.T) { + ctrl := gomock.NewController(t) + mockLeaseStore := NewMockpartitionLeaseStore(ctrl) + mockHeartbeatStore := NewMocksubscriberHeartbeatStore(ctrl) + cfg := testSubscriptionConfig() + cfg.PollIntervalMs = int64(time.Hour / time.Millisecond) + leasesErr := errors.New("all leases unavailable") + partition := entityqueue.PartitionIdentity{Tenant: testTenant, PartitionKey: "p1"} + + mockLeaseStore.EXPECT(). + DiscoverPartitions(gomock.Any(), []string{testTenant}, "test-topic"). + Return(map[string][]string{testTenant: {partition.PartitionKey}}, nil) + mockLeaseStore.EXPECT(). + GetLeasedPartitionsForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + Return(map[string][]string{testTenant: {partition.PartitionKey}}, nil) + mockLeaseStore.EXPECT(). + GetAllLeasesForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.ConsumerGroup). + Return(nil, leasesErr) + + s := NewSubscriber( + zaptest.NewLogger(t).Sugar(), + tally.NoopScope, + NewMockmessageStore(ctrl), + NewMockoffsetStore(ctrl), + mockLeaseStore, + mockHeartbeatStore, + newTestDeliveryStateStore(ctrl), + []string{testTenant}, + ) + failedWorkerDone := make(chan struct{}) + close(failedWorkerDone) + sub := &subscription{ + topic: "test-topic", + config: cfg, + deliveryCh: make(chan extqueue.Delivery, 1), + workers: map[entityqueue.PartitionIdentity]*partitionWorker{ + partition: { + cancelFunc: func() {}, + done: failedWorkerDone, + }, + }, + } + + err := s.discoverAndReconcileWorkers(context.Background(), sub, false) + require.ErrorIs(t, err, leasesErr) + assert.Contains(t, sub.workers, partition) + assert.Equal(t, []entityqueue.PartitionIdentity{partition}, sub.lastDiscoveredPartitions) +} + +func TestSubscriber_DiscoverAppliesFairSharePerTenant(t *testing.T) { + ctrl := gomock.NewController(t) + mockLeaseStore := NewMockpartitionLeaseStore(ctrl) + mockHeartbeatStore := NewMocksubscriberHeartbeatStore(ctrl) + cfg := testSubscriptionConfig() + cfg.PollIntervalMs = int64(time.Hour / time.Millisecond) + tenantA, tenantB := "tenant-a", "tenant-b" + tenants := []string{tenantA, tenantB} + discoveredA := []string{"a1", "a2", "a3", "a4"} + discoveredB := []string{"b1", "b2", "b3"} + + mockLeaseStore.EXPECT(). + DiscoverPartitions(gomock.Any(), tenants, "test-topic"). + Return(map[string][]string{tenantA: discoveredA, tenantB: discoveredB}, nil) + mockLeaseStore.EXPECT(). + GetLeasedPartitionsForTenants(gomock.Any(), tenants, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + Return(map[string][]string{}, nil) + mockLeaseStore.EXPECT(). + GetAllLeasesForTenants(gomock.Any(), tenants, "test-topic", cfg.ConsumerGroup). + Return(map[string][]leaseInfo{}, nil) + mockHeartbeatStore.EXPECT(). + ActiveSubscribersForTenants(gomock.Any(), tenants, "test-topic", cfg.ConsumerGroup, cfg.LeaseDurationMs). + Return(map[string][]string{ + tenantA: {cfg.SubscriberName, "peer"}, + tenantB: {cfg.SubscriberName}, + }, nil) + for _, pk := range discoveredA[:2] { + mockLeaseStore.EXPECT(). + TryAcquireLease(gomock.Any(), tenantA, "test-topic", pk, cfg.SubscriberName, cfg.ConsumerGroup, cfg.LeaseDurationMs). + Return(true, nil) + } + for _, pk := range discoveredB { + mockLeaseStore.EXPECT(). + TryAcquireLease(gomock.Any(), tenantB, "test-topic", pk, cfg.SubscriberName, cfg.ConsumerGroup, cfg.LeaseDurationMs). + Return(true, nil) + } + + s := NewSubscriber( + zaptest.NewLogger(t).Sugar(), + tally.NoopScope, + NewMockmessageStore(ctrl), + NewMockoffsetStore(ctrl), + mockLeaseStore, + mockHeartbeatStore, + newTestDeliveryStateStore(ctrl), + tenants, + ) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + sub := &subscription{ + topic: "test-topic", + config: cfg, + deliveryCh: make(chan extqueue.Delivery, 8), + workers: make(map[entityqueue.PartitionIdentity]*partitionWorker), + } + + require.NoError(t, s.discoverAndReconcileWorkers(ctx, sub, false)) + s.stopAllWorkers(sub) +} + func TestSubscriber_DrainedPartitionKeepsOffsetWhenLeaseReleaseFails(t *testing.T) { ctrl := gomock.NewController(t) mockLeaseStore := NewMockpartitionLeaseStore(ctrl) @@ -1143,14 +1452,14 @@ func TestSubscriber_DrainedPartitionKeepsOffsetWhenLeaseReleaseFails(t *testing. partition := entityqueue.PartitionIdentity{Tenant: testTenant, PartitionKey: "drained"} mockLeaseStore.EXPECT(). - GetLeasedPartitions(gomock.Any(), testTenant, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). - Return([]string{partition.PartitionKey}, nil) + DiscoverPartitions(gomock.Any(), []string{testTenant}, "test-topic"). + Return(map[string][]string{}, nil) mockLeaseStore.EXPECT(). - DiscoverAndAcquirePartitions(gomock.Any(), testTenant, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup, cfg.LeaseDurationMs, 0). - Return(0, nil, nil) + GetLeasedPartitionsForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + Return(map[string][]string{testTenant: {partition.PartitionKey}}, nil) mockLeaseStore.EXPECT(). - GetLeasedPartitions(gomock.Any(), testTenant, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). - Return([]string{partition.PartitionKey}, nil) + GetAllLeasesForTenants(gomock.Any(), []string{testTenant}, "test-topic", cfg.ConsumerGroup). + Return(map[string][]leaseInfo{}, nil) mockLeaseStore.EXPECT(). ReleaseLease(gomock.Any(), testTenant, "test-topic", partition.PartitionKey, cfg.SubscriberName, cfg.ConsumerGroup). Return(errors.New("release failed")) @@ -1188,71 +1497,48 @@ func TestSubscriber_ReleaseAllLeasesContinuesAfterErrors(t *testing.T) { ctrl := gomock.NewController(t) mockLeaseStore := NewMockpartitionLeaseStore(ctrl) releaseErr := errors.New("release failed") - discoveryErr := errors.New("lease lookup failed") cfg := testSubscriptionConfig() + tenants := []string{"tenant-1", "tenant-2", "tenant-3"} mockLeaseStore.EXPECT(). - GetLeasedPartitions(gomock.Any(), "tenant-1", "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). - Return([]string{"p1", "p2"}, nil) - mockLeaseStore.EXPECT(). - ReleaseLease(gomock.Any(), "tenant-1", "test-topic", "p1", cfg.SubscriberName, cfg.ConsumerGroup). + ReleaseOwnedLeases(gomock.Any(), tenants, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). Return(releaseErr) - mockLeaseStore.EXPECT(). - ReleaseLease(gomock.Any(), "tenant-1", "test-topic", "p2", cfg.SubscriberName, cfg.ConsumerGroup). - Return(nil) - mockLeaseStore.EXPECT(). - GetLeasedPartitions(gomock.Any(), "tenant-2", "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). - Return(nil, discoveryErr) - mockLeaseStore.EXPECT(). - GetLeasedPartitions(gomock.Any(), "tenant-3", "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). - Return([]string{"p3"}, nil) - mockLeaseStore.EXPECT(). - ReleaseLease(gomock.Any(), "tenant-3", "test-topic", "p3", cfg.SubscriberName, cfg.ConsumerGroup). - Return(nil) s := NewSubscriber( zaptest.NewLogger(t).Sugar(), tally.NoopScope, NewMockmessageStore(ctrl), NewMockoffsetStore(ctrl), mockLeaseStore, NewMocksubscriberHeartbeatStore(ctrl), NewMockdeliveryStateStore(ctrl), - []string{"tenant-1", "tenant-2", "tenant-3"}, + tenants, ) sub := &subscription{topic: "test-topic", config: cfg} err := s.releaseAllLeases(context.Background(), sub) require.ErrorIs(t, err, releaseErr) - require.ErrorIs(t, err, discoveryErr) } func TestSubscriber_DeregisterHeartbeatContinuesAfterErrors(t *testing.T) { ctrl := gomock.NewController(t) mockHeartbeatStore := NewMocksubscriberHeartbeatStore(ctrl) - firstErr := errors.New("first deregistration failed") - lastErr := errors.New("last deregistration failed") + deregisterErr := errors.New("deregistration failed") cfg := testSubscriptionConfig() + tenants := []string{"tenant-1", "tenant-2", "tenant-3"} mockHeartbeatStore.EXPECT(). - Deregister(gomock.Any(), "tenant-1", "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). - Return(firstErr) - mockHeartbeatStore.EXPECT(). - Deregister(gomock.Any(), "tenant-2", "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). - Return(nil) - mockHeartbeatStore.EXPECT(). - Deregister(gomock.Any(), "tenant-3", "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). - Return(lastErr) + DeregisterForTenants(gomock.Any(), tenants, "test-topic", cfg.SubscriberName, cfg.ConsumerGroup). + Return(deregisterErr) s := NewSubscriber( zaptest.NewLogger(t).Sugar(), tally.NoopScope, NewMockmessageStore(ctrl), NewMockoffsetStore(ctrl), NewMockpartitionLeaseStore(ctrl), mockHeartbeatStore, NewMockdeliveryStateStore(ctrl), - []string{"tenant-1", "tenant-2", "tenant-3"}, + tenants, ) sub := &subscription{topic: "test-topic", config: cfg} err := s.deregisterHeartbeat(context.Background(), sub) - require.ErrorIs(t, err, firstErr) - require.ErrorIs(t, err, lastErr) + require.ErrorIs(t, err, deregisterErr) } // TestSubscriber_PartitionWorkerPollAndDeliver verifies a partition worker delivers messages. @@ -1775,16 +2061,10 @@ func TestSubscriber_FairShareCap(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { ctrl := gomock.NewController(t) - mockHB := NewMocksubscriberHeartbeatStore(ctrl) - mockHB.EXPECT(). - ActiveSubscribers(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). - Return(tt.active, nil). - AnyTimes() - s := NewSubscriber( zaptest.NewLogger(t).Sugar(), tally.NoopScope, NewMockmessageStore(ctrl), NewMockoffsetStore(ctrl), - NewMockpartitionLeaseStore(ctrl), mockHB, + NewMockpartitionLeaseStore(ctrl), NewMocksubscriberHeartbeatStore(ctrl), NewMockdeliveryStateStore(ctrl), []string{testTenant}, ) @@ -1793,8 +2073,7 @@ func TestSubscriber_FairShareCap(t *testing.T) { config: extqueue.DefaultSubscriptionConfig(tt.self, "test-cg"), } - got, err := s.fairShareCap(context.Background(), sub, testTenant, tt.owned, tt.discovered) - require.NoError(t, err) + got := s.fairShareCap(sub, tt.owned, tt.discovered, tt.active) assert.Equal(t, tt.want, got) }) } @@ -1811,15 +2090,10 @@ func TestSubscriber_FairShareCap(t *testing.T) { } ctrl := gomock.NewController(t) - mockHB := NewMocksubscriberHeartbeatStore(ctrl) - mockHB.EXPECT(). - ActiveSubscribers(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). - Return(active, nil). - AnyTimes() s := NewSubscriber( zaptest.NewLogger(t).Sugar(), tally.NoopScope, NewMockmessageStore(ctrl), NewMockoffsetStore(ctrl), - NewMockpartitionLeaseStore(ctrl), mockHB, + NewMockpartitionLeaseStore(ctrl), NewMocksubscriberHeartbeatStore(ctrl), NewMockdeliveryStateStore(ctrl), []string{testTenant}, ) @@ -1830,8 +2104,7 @@ func TestSubscriber_FairShareCap(t *testing.T) { topic: "test-topic", config: extqueue.DefaultSubscriptionConfig(self, "test-cg"), } - cap, err := s.fairShareCap(context.Background(), sub, testTenant, nil, partitionKeysN(p)) - require.NoError(t, err) + cap := s.fairShareCap(sub, nil, partitionKeysN(p), active) sum += cap } require.Equal(t, p, sum, "n=%d p=%d", n, p) @@ -1861,12 +2134,6 @@ func TestSubscriber_RebalanceReleasesExcess(t *testing.T) { ctrl := gomock.NewController(t) // Two active subscribers, four partitions: self is rank 0 -> cap 2. - mockHB := NewMocksubscriberHeartbeatStore(ctrl) - mockHB.EXPECT(). - ActiveSubscribers(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). - Return([]string{"s1", "s2"}, nil) - - // The lexicographically largest partitions beyond the cap are released. mockLease := NewMockpartitionLeaseStore(ctrl) mockLease.EXPECT(). ReleaseLease(gomock.Any(), testTenant, "test-topic", "pk-c", "s1", "test-cg"). @@ -1878,7 +2145,7 @@ func TestSubscriber_RebalanceReleasesExcess(t *testing.T) { s := NewSubscriber( zaptest.NewLogger(t).Sugar(), tally.NoopScope, NewMockmessageStore(ctrl), NewMockoffsetStore(ctrl), - mockLease, mockHB, NewMockdeliveryStateStore(ctrl), + mockLease, NewMocksubscriberHeartbeatStore(ctrl), NewMockdeliveryStateStore(ctrl), []string{testTenant}, ) sub := &subscription{ @@ -1888,7 +2155,7 @@ func TestSubscriber_RebalanceReleasesExcess(t *testing.T) { } owned := []string{"pk-d", "pk-a", "pk-c", "pk-b"} - released, err := s.rebalance(context.Background(), sub, testTenant, owned) + released, err := s.rebalance(context.Background(), sub, testTenant, owned, []string{"s1", "s2"}) require.NoError(t, err) assert.Equal(t, []string{"pk-c", "pk-d"}, released) // The caller's slice is shared with lease renewal and must not be @@ -1900,16 +2167,10 @@ func TestSubscriber_RebalanceReleasesExcess(t *testing.T) { func TestSubscriber_RebalanceUnderCapReleasesNothing(t *testing.T) { ctrl := gomock.NewController(t) - mockHB := NewMocksubscriberHeartbeatStore(ctrl) - mockHB.EXPECT(). - ActiveSubscribers(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()). - Return([]string{"s1", "s2"}, nil) - - // No ReleaseLease expectations: owning exactly the cap sheds nothing. s := NewSubscriber( zaptest.NewLogger(t).Sugar(), tally.NoopScope, NewMockmessageStore(ctrl), NewMockoffsetStore(ctrl), - NewMockpartitionLeaseStore(ctrl), mockHB, NewMockdeliveryStateStore(ctrl), + NewMockpartitionLeaseStore(ctrl), NewMocksubscriberHeartbeatStore(ctrl), NewMockdeliveryStateStore(ctrl), []string{testTenant}, ) sub := &subscription{ @@ -1921,7 +2182,7 @@ func TestSubscriber_RebalanceUnderCapReleasesNothing(t *testing.T) { lastDiscoveredPartitions: tenantPartitionKeys(testTenant, []string{"pk-a", "pk-b", "pk-c", "pk-d"}), } - released, err := s.rebalance(context.Background(), sub, testTenant, []string{"pk-a", "pk-b"}) + released, err := s.rebalance(context.Background(), sub, testTenant, []string{"pk-a", "pk-b"}, []string{"s1", "s2"}) require.NoError(t, err) assert.Empty(t, released) }