Skip to content

Commit d168d11

Browse files
committed
Address PR review feedback (#874)
- bulk insert command queue messages within the transaction bound - finish command batches after every enqueue failure - build the outstanding queue index concurrently
1 parent 06e9753 commit d168d11

12 files changed

Lines changed: 262 additions & 32 deletions

server/generated/sqlc/db.go

Lines changed: 10 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

server/generated/sqlc/querier.go

Lines changed: 1 addition & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

server/generated/sqlc/queue.sql.go

Lines changed: 37 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

server/generated/sqlc/retrying_querier.gen.go

Lines changed: 6 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.
Lines changed: 79 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,79 @@
1+
package command_test
2+
3+
import (
4+
"context"
5+
"errors"
6+
"testing"
7+
8+
"github.com/stretchr/testify/assert"
9+
"github.com/stretchr/testify/require"
10+
11+
commonpb "github.com/block/proto-fleet/server/generated/grpc/common/v1"
12+
pb "github.com/block/proto-fleet/server/generated/grpc/minercommand/v1"
13+
"github.com/block/proto-fleet/server/internal/domain/commandtype"
14+
"github.com/block/proto-fleet/server/internal/infrastructure/queue"
15+
"github.com/block/proto-fleet/server/internal/testutil"
16+
)
17+
18+
type failingEnqueueQueue struct {
19+
err error
20+
}
21+
22+
func (q failingEnqueueQueue) Enqueue(context.Context, string, commandtype.Type, []int64, interface{}) error {
23+
return q.err
24+
}
25+
26+
func (q failingEnqueueQueue) EnqueueMany(context.Context, string, commandtype.Type, []queue.EnqueueMessage) error {
27+
return q.err
28+
}
29+
30+
func (failingEnqueueQueue) Dequeue(context.Context, int32) ([]queue.Message, error) {
31+
return nil, nil
32+
}
33+
34+
func (failingEnqueueQueue) IsBatchFinished(context.Context, string) (bool, error) {
35+
return false, nil
36+
}
37+
38+
func (failingEnqueueQueue) MaxFailureRetries() int32 {
39+
return 0
40+
}
41+
42+
func TestCommandEnqueueFailureFinishesCreatedBatch(t *testing.T) {
43+
if testing.Short() {
44+
t.Skip("Skipping database integration test in short mode")
45+
}
46+
47+
// Arrange
48+
conn, dbService, user := setupRetentionTest(t)
49+
device := dbService.CreateDevice(user.OrganizationID, "proto")
50+
enqueueErr := errors.New("queue unavailable")
51+
svc := newDispatchIntegrationTestService(t, conn, failingEnqueueQueue{err: enqueueErr})
52+
ctx := testutil.MockAuthContextForTesting(t.Context(), user.DatabaseID, user.OrganizationID)
53+
54+
// Act
55+
result, err := svc.BlinkLED(ctx, &pb.DeviceSelector{
56+
SelectionType: &pb.DeviceSelector_IncludeDevices{
57+
IncludeDevices: &commonpb.DeviceIdentifierList{DeviceIdentifiers: []string{device.ID}},
58+
},
59+
})
60+
61+
// Assert
62+
require.Error(t, err)
63+
assert.Nil(t, result)
64+
assert.ErrorContains(t, err, enqueueErr.Error())
65+
var batchUUID, status string
66+
require.NoError(t, conn.QueryRowContext(t.Context(), `
67+
SELECT uuid, status
68+
FROM command_batch_log
69+
WHERE organization_id = $1
70+
ORDER BY id DESC
71+
LIMIT 1
72+
`, user.OrganizationID).Scan(&batchUUID, &status))
73+
assert.Equal(t, "FINISHED", status)
74+
var queued int
75+
require.NoError(t, conn.QueryRowContext(t.Context(), `
76+
SELECT COUNT(*) FROM queue_message WHERE command_batch_log_uuid = $1
77+
`, batchUUID).Scan(&queued))
78+
assert.Zero(t, queued)
79+
}

server/internal/domain/command/service.go

Lines changed: 8 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1061,15 +1061,16 @@ func (s *Service) processCommand(ctx context.Context, command *Command) (*Comman
10611061
return s.messageQueue.EnqueueMany(workCtx, batchLogIdentifier, command.commandType, queuePayloads)
10621062
})
10631063
if err != nil {
1064+
var enqueueErr error
10641065
switch {
10651066
case errors.Is(err, errExecutionStoppedBeforeEnqueue):
1066-
enqueueErr := fleeterror.NewInternalError("command execution service stopped before enqueue")
1067-
return nil, s.finishUnenqueuedCommandBatch(ctx, batchLogIdentifier, enqueueErr)
1067+
enqueueErr = fleeterror.NewInternalError("command execution service stopped before enqueue")
10681068
case len(queuePayloads) == 0:
1069-
return nil, fleeterror.NewInternalErrorf("error enqueuing a batch of commands: %v", err)
1069+
enqueueErr = fleeterror.NewInternalErrorf("error enqueuing a batch of commands: %v", err)
10701070
default:
1071-
return nil, fleeterror.NewInternalErrorf("error enqueuing per-device command payloads: %v", err)
1071+
enqueueErr = fleeterror.NewInternalErrorf("error enqueuing per-device command payloads: %v", err)
10721072
}
1073+
return nil, s.finishUnenqueuedCommandBatch(ctx, batchLogIdentifier, enqueueErr)
10731074
}
10741075

10751076
return &CommandResult{
@@ -1548,14 +1549,10 @@ func (s *Service) ReapplyCurrentPoolsWithWorkerNames(
15481549
return s.enqueueWorkerNameReapplyMessages(workCtx, commandBatchLogUUID, deviceIdentifiers, deviceIDsByIdentifier, desiredWorkerNamesByDeviceIdentifier)
15491550
})
15501551
if err != nil {
1551-
enqueueErr := err
1552-
switch {
1553-
case errors.Is(err, errExecutionStoppedBeforeEnqueue):
1554-
enqueueErr = fleeterror.NewInternalError("command execution service stopped before enqueue")
1555-
return "", s.finishUnenqueuedCommandBatch(ctx, commandBatchLogUUID, enqueueErr)
1556-
default:
1557-
return "", enqueueErr
1552+
if errors.Is(err, errExecutionStoppedBeforeEnqueue) {
1553+
err = fleeterror.NewInternalError("command execution service stopped before enqueue")
15581554
}
1555+
return "", s.finishUnenqueuedCommandBatch(ctx, commandBatchLogUUID, err)
15591556
}
15601557

15611558
s.initializeStatusUpdateRoutine(commandBatchLogUUID, nil)

server/internal/domain/command/zero_target_integration_test.go

Lines changed: 8 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,14 @@ import (
2323

2424
func newZeroTargetDispatchTestService(t *testing.T, conn *sql.DB) *command.Service {
2525
t.Helper()
26+
return newDispatchIntegrationTestService(t, conn, queue.NewDatabaseMessageQueue(&queue.Config{
27+
DequeLimit: 10,
28+
MaxFailureRetries: 1,
29+
}, conn))
30+
}
31+
32+
func newDispatchIntegrationTestService(t *testing.T, conn *sql.DB, messageQueue queue.MessageQueue) *command.Service {
33+
t.Helper()
2634

2735
commandConfig := &command.Config{
2836
MaxWorkers: 1,
@@ -33,11 +41,6 @@ func newZeroTargetDispatchTestService(t *testing.T, conn *sql.DB) *command.Servi
3341
StuckMessageTimeout: time.Hour,
3442
ReaperInterval: time.Hour,
3543
}
36-
queueConfig := &queue.Config{
37-
DequeLimit: 10,
38-
MaxFailureRetries: 1,
39-
}
40-
messageQueue := queue.NewDatabaseMessageQueue(queueConfig, conn)
4144
executionCtx, cancel := context.WithCancel(context.Background())
4245
t.Cleanup(cancel)
4346

server/internal/infrastructure/queue/service.go

Lines changed: 13 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -5,8 +5,6 @@ import (
55
"database/sql"
66
"encoding/json"
77

8-
"github.com/sqlc-dev/pqtype"
9-
108
"github.com/block/proto-fleet/server/generated/sqlc"
119
"github.com/block/proto-fleet/server/internal/domain/commandtype"
1210
"github.com/block/proto-fleet/server/internal/domain/fleeterror"
@@ -58,19 +56,20 @@ func (d DatabaseMessageQueue) EnqueueMany(ctx context.Context, commandBatchLogUU
5856
}
5957

6058
func (d DatabaseMessageQueue) enqueueEncoded(ctx context.Context, commandBatchLogUUID string, commandType commandtype.Type, messages []encodedMessage) error {
59+
deviceIDs := make([]int64, len(messages))
60+
payloads := make([]string, len(messages))
61+
for i, message := range messages {
62+
deviceIDs[i] = message.deviceID
63+
payloads[i] = string(message.payload)
64+
}
6165
return db.WithTransactionTimeoutNoResult(ctx, d.conn, runtimepolicy.CommandTransactionBound, func(q sqlc.Querier) error {
62-
for _, message := range messages {
63-
err := q.CreateQueueMessage(ctx, sqlc.CreateQueueMessageParams{
64-
CommandBatchLogUuid: commandBatchLogUUID,
65-
CommandType: commandType.String(),
66-
DeviceID: message.deviceID,
67-
Status: sqlc.QueueStatusEnumPENDING,
68-
RetryCount: 0,
69-
Payload: pqtype.NullRawMessage{RawMessage: message.payload, Valid: true},
70-
})
71-
if err != nil {
72-
return fleeterror.NewInternalErrorf("failed to enqueue message: %v", err)
73-
}
66+
if err := q.CreateQueueMessages(ctx, sqlc.CreateQueueMessagesParams{
67+
CommandBatchLogUuid: commandBatchLogUUID,
68+
CommandType: commandType.String(),
69+
DeviceIds: deviceIDs,
70+
Payloads: payloads,
71+
}); err != nil {
72+
return fleeterror.NewInternalErrorf("failed to enqueue messages: %v", err)
7473
}
7574
return nil
7675
})
Lines changed: 79 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,79 @@
1+
package queue_test
2+
3+
import (
4+
"database/sql"
5+
"encoding/json"
6+
"testing"
7+
"time"
8+
9+
"github.com/sqlc-dev/pqtype"
10+
"github.com/stretchr/testify/assert"
11+
"github.com/stretchr/testify/require"
12+
13+
"github.com/block/proto-fleet/server/generated/sqlc"
14+
"github.com/block/proto-fleet/server/internal/domain/commandtype"
15+
"github.com/block/proto-fleet/server/internal/infrastructure/id"
16+
"github.com/block/proto-fleet/server/internal/infrastructure/queue"
17+
"github.com/block/proto-fleet/server/internal/testutil"
18+
)
19+
20+
func TestDatabaseMessageQueueEnqueueManyInsertsPerDevicePayloads(t *testing.T) {
21+
if testing.Short() {
22+
t.Skip("Skipping database integration test in short mode")
23+
}
24+
25+
// Arrange
26+
cfg, err := testutil.GetTestConfig()
27+
require.NoError(t, err)
28+
dbService := testutil.NewDatabaseService(t, cfg)
29+
user := dbService.CreateSuperAdminUser()
30+
firstDevice := dbService.CreateDevice(user.OrganizationID, "proto")
31+
secondDevice := dbService.CreateDevice(user.OrganizationID, "proto")
32+
batchUUID := id.GenerateID()
33+
commandType := commandtype.UpdateMiningPools
34+
_, err = sqlc.New(dbService.DB).CreateCommandBatchLog(t.Context(), sqlc.CreateCommandBatchLogParams{
35+
Uuid: batchUUID,
36+
Type: commandType.String(),
37+
CreatedBy: user.DatabaseID,
38+
CreatedAt: time.Now(),
39+
Status: sqlc.BatchStatusEnumPENDING,
40+
DevicesCount: 2,
41+
Payload: pqtype.NullRawMessage{},
42+
OrganizationID: sql.NullInt64{Int64: user.OrganizationID, Valid: true},
43+
})
44+
require.NoError(t, err)
45+
messageQueue := queue.NewDatabaseMessageQueue(&queue.Config{}, dbService.DB)
46+
messages := []queue.EnqueueMessage{
47+
{DeviceID: firstDevice.DatabaseID, Payload: map[string]string{"worker_name": "first"}},
48+
{DeviceID: secondDevice.DatabaseID, Payload: map[string]string{"worker_name": "second"}},
49+
}
50+
51+
// Act
52+
err = messageQueue.EnqueueMany(t.Context(), batchUUID, commandType, messages)
53+
54+
// Assert
55+
require.NoError(t, err)
56+
rows, err := dbService.DB.QueryContext(t.Context(), `
57+
SELECT device_id, payload, status
58+
FROM queue_message
59+
WHERE command_batch_log_uuid = $1
60+
`, batchUUID)
61+
require.NoError(t, err)
62+
defer rows.Close()
63+
gotPayloads := make(map[int64]map[string]string)
64+
for rows.Next() {
65+
var deviceID int64
66+
var payload []byte
67+
var status string
68+
require.NoError(t, rows.Scan(&deviceID, &payload, &status))
69+
var decoded map[string]string
70+
require.NoError(t, json.Unmarshal(payload, &decoded))
71+
gotPayloads[deviceID] = decoded
72+
assert.Equal(t, "PENDING", status)
73+
}
74+
require.NoError(t, rows.Err())
75+
assert.Equal(t, map[int64]map[string]string{
76+
firstDevice.DatabaseID: {"worker_name": "first"},
77+
secondDevice.DatabaseID: {"worker_name": "second"},
78+
}, gotPayloads)
79+
}
Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1 @@
1-
DROP INDEX IF EXISTS idx_queue_message_reaper_outstanding;
1+
DROP INDEX CONCURRENTLY IF EXISTS idx_queue_message_reaper_outstanding;

0 commit comments

Comments
 (0)