Repository navigation
Expand file tree
/
Copy pathtransaction.go
More file actions
237 lines (215 loc) · 6.73 KB
/
Copy pathtransaction.go
File metadata and controls
237 lines (215 loc) · 6.73 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
package blockqueue
import (
"context"
"database/sql"
"errors"
"fmt"
"time"
"github.com/google/uuid"
)
// ErrInvalidTransaction reports a missing or otherwise unusable caller
// transaction.
var ErrInvalidTransaction = errors.New("invalid blockqueue transaction")
// ErrTransactionCommitUnknown means Commit returned a connection-level error
// for a caller-owned transaction. The database may have committed the
// application writes and queue changes; callers must reconcile and must not
// blindly repeat non-idempotent business operations.
var ErrTransactionCommitUnknown = errors.New("transaction commit outcome unknown")
// TransactionCommitUnknownError reports that a caller-owned transaction may
// have committed even though Commit returned an error.
type TransactionCommitUnknownError struct {
Cause error
}
func (err *TransactionCommitUnknownError) Error() string {
return fmt.Sprintf("transaction commit outcome unknown: %v", err.Cause)
}
func (err *TransactionCommitUnknownError) Unwrap() error { return err.Cause }
func (err *TransactionCommitUnknownError) Is(target error) bool {
return target == ErrTransactionCommitUnknown || errors.Is(err.Cause, target)
}
type activeTransaction struct {
topics map[uuid.UUID]struct{}
}
// WithTx runs fn in a transaction owned by the queue. It is the preferred way
// to atomically change application tables and publish or complete deliveries:
// the queue can notify local waiters only after commit succeeds. A context
// cancellation or deadline reported by Commit is outcome-unknown because the
// server may have committed before the client stopped waiting; fn is never
// retried.
func (q *Queue) WithTx(ctx context.Context, options *sql.TxOptions, fn func(*sql.Tx) error) error {
if ctx == nil {
ctx = context.Background()
}
if fn == nil {
return fmt.Errorf("%w: callback is required", ErrInvalidTransaction)
}
if err := q.reserveTransactionStart(); err != nil {
return err
}
tx, err := q.db.Database.DB().BeginTx(ctx, options)
if err != nil {
q.transactions.Done()
return err
}
q.registerTransaction(tx)
defer q.unregisterTransaction(tx)
defer func() { _ = tx.Rollback() }()
if err := fn(tx); err != nil {
return err
}
if err := tx.Commit(); err != nil {
if contextErr := ctx.Err(); contextErr != nil || isAmbiguousCommitError(err) {
cause := err
if contextErr != nil && !errors.Is(err, contextErr) {
cause = errors.Join(err, contextErr)
}
return &TransactionCommitUnknownError{Cause: cause}
}
return err
}
for _, topicID := range q.transactionTopics(tx) {
q.notify(topicID)
}
return nil
}
// PublishTx stages a canonical message and its fan-out in caller-owned tx.
// The caller owns commit or rollback; staged rows are not claimable before a
// successful commit.
func (q *Queue) PublishTx(ctx context.Context, tx *sql.Tx, topic Topic, request Message) (PublishReceipt, error) {
receipts, err := q.BatchPublishTx(ctx, tx, topic, []Message{request})
if err != nil {
return PublishReceipt{}, err
}
return receipts[0], nil
}
// BatchPublishTx validates the complete batch before writing it to tx.
func (q *Queue) BatchPublishTx(ctx context.Context, tx *sql.Tx, topic Topic, requests []Message) (PublishReceipts, error) {
if ctx == nil {
ctx = context.Background()
}
if tx == nil {
return nil, fmt.Errorf("%w: transaction is required", ErrInvalidTransaction)
}
if err := q.requireTransactionAllowed(tx); err != nil {
return nil, err
}
if len(requests) == 0 {
return PublishReceipts{}, nil
}
runtime, exists := q.getTopicRuntime(topic)
if !exists {
return nil, ErrTopicNotFound
}
q.admissionMu.RLock()
defer q.admissionMu.RUnlock()
if err := q.requireTransactionAllowed(tx); err != nil {
return nil, err
}
if current := q.registry.Load().byName[topic.Name]; current != runtime || runtime.deleted.Load() {
return nil, ErrTopicNotFound
}
now := time.Now().UTC().Truncate(time.Millisecond)
writes := make([]writeRequest, len(requests))
receipts := make(PublishReceipts, len(requests))
seenKeys := make(map[string]string)
for index, request := range requests {
write, scheduledAt, err := buildWriteRequest(runtime.id, request, now)
if err != nil {
return nil, err
}
if request.IdempotencyKey != "" {
if previous, ok := seenKeys[request.IdempotencyKey]; ok && previous != write.IdempotencyHash {
return nil, fmt.Errorf("%w: duplicate key %q has different payload", ErrInvalidPublish, request.IdempotencyKey)
}
seenKeys[request.IdempotencyKey] = write.IdempotencyHash
}
writes[index] = write
receipts[index] = PublishReceipt{
MessageID: write.MessageID,
State: PublishStateStaged,
ScheduledAt: scheduledAt,
}
}
result, err := q.db.persistWriteRequestsWithTx(ctx, tx, writes)
if err != nil {
return nil, err
}
for index := range receipts {
duplicate := result.Duplicates[index]
receipts[index].Duplicate = &duplicate
receipts[index].ScheduledAt = result.ScheduledAt[index]
}
q.markTransactionTopic(tx, runtime.id)
return receipts, nil
}
func (q *Queue) requireRunning() error {
switch q.State() {
case LifecycleRunning:
return nil
case LifecycleStopping, LifecycleStopped:
return ErrQueueStopping
default:
return ErrQueueNotRunning
}
}
func (q *Queue) requireTransactionAllowed(tx *sql.Tx) error {
if q.State() == LifecycleRunning {
return nil
}
if tx != nil {
q.transactionMu.RLock()
_, active := q.activeTx[tx]
q.transactionMu.RUnlock()
if active {
return nil
}
}
return q.requireRunning()
}
func (q *Queue) registerTransaction(tx *sql.Tx) {
q.transactionMu.Lock()
q.activeTx[tx] = &activeTransaction{topics: make(map[uuid.UUID]struct{})}
q.transactionMu.Unlock()
}
func (q *Queue) markTransactionTopic(tx *sql.Tx, topicID uuid.UUID) {
if tx == nil {
return
}
q.transactionMu.Lock()
if active := q.activeTx[tx]; active != nil {
active.topics[topicID] = struct{}{}
}
q.transactionMu.Unlock()
}
func (q *Queue) transactionTopics(tx *sql.Tx) []uuid.UUID {
q.transactionMu.RLock()
active := q.activeTx[tx]
if active == nil {
q.transactionMu.RUnlock()
return nil
}
topics := make([]uuid.UUID, 0, len(active.topics))
for topicID := range active.topics {
topics = append(topics, topicID)
}
q.transactionMu.RUnlock()
return topics
}
func (q *Queue) unregisterTransaction(tx *sql.Tx) {
q.transactionMu.Lock()
delete(q.activeTx, tx)
q.transactionMu.Unlock()
q.transactions.Done()
}
func (q *Queue) reserveTransactionStart() error {
q.admissionMu.RLock()
defer q.admissionMu.RUnlock()
if err := q.requireRunning(); err != nil {
return err
}
// Add happens behind the same lifecycle fence acquired by Shutdown before
// Wait, satisfying sync.WaitGroup's Add-before-Wait requirement without
// holding a mutex across a potentially blocking BeginTx call.
q.transactions.Add(1)
return nil
}