mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
chore: separate pubsub into a new package (#8017)
* chore: rename store to dbmock for consistency * chore: remove redundant dbtype package This wasn't necessary and forked how we do DB types. * chore: separate pubsub into a new package This didn't need to be in database and was bloating it.
This commit is contained in:
@@ -0,0 +1,337 @@
|
||||
package pubsub
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/lib/pq"
|
||||
"golang.org/x/xerrors"
|
||||
)
|
||||
|
||||
// Listener represents a pubsub handler.
|
||||
type Listener func(ctx context.Context, message []byte)
|
||||
|
||||
// ListenerWithErr represents a pubsub handler that can also receive error
|
||||
// indications
|
||||
type ListenerWithErr func(ctx context.Context, message []byte, err error)
|
||||
|
||||
// ErrDroppedMessages is sent to ListenerWithErr if messages are dropped or
|
||||
// might have been dropped.
|
||||
var ErrDroppedMessages = xerrors.New("dropped messages")
|
||||
|
||||
// Pubsub is a generic interface for broadcasting and receiving messages.
|
||||
// Implementors should assume high-availability with the backing implementation.
|
||||
type Pubsub interface {
|
||||
Subscribe(event string, listener Listener) (cancel func(), err error)
|
||||
SubscribeWithErr(event string, listener ListenerWithErr) (cancel func(), err error)
|
||||
Publish(event string, message []byte) error
|
||||
Close() error
|
||||
}
|
||||
|
||||
// msgOrErr either contains a message or an error
|
||||
type msgOrErr struct {
|
||||
msg []byte
|
||||
err error
|
||||
}
|
||||
|
||||
// msgQueue implements a fixed length queue with the ability to replace elements
|
||||
// after they are queued (but before they are dequeued).
|
||||
//
|
||||
// The purpose of this data structure is to build something that works a bit
|
||||
// like a golang channel, but if the queue is full, then we can replace the
|
||||
// last element with an error so that the subscriber can get notified that some
|
||||
// messages were dropped, all without blocking.
|
||||
type msgQueue struct {
|
||||
ctx context.Context
|
||||
cond *sync.Cond
|
||||
q [BufferSize]msgOrErr
|
||||
front int
|
||||
size int
|
||||
closed bool
|
||||
l Listener
|
||||
le ListenerWithErr
|
||||
}
|
||||
|
||||
func newMsgQueue(ctx context.Context, l Listener, le ListenerWithErr) *msgQueue {
|
||||
if l == nil && le == nil {
|
||||
panic("l or le must be non-nil")
|
||||
}
|
||||
q := &msgQueue{
|
||||
ctx: ctx,
|
||||
cond: sync.NewCond(&sync.Mutex{}),
|
||||
l: l,
|
||||
le: le,
|
||||
}
|
||||
go q.run()
|
||||
return q
|
||||
}
|
||||
|
||||
func (q *msgQueue) run() {
|
||||
for {
|
||||
// wait until there is something on the queue or we are closed
|
||||
q.cond.L.Lock()
|
||||
for q.size == 0 && !q.closed {
|
||||
q.cond.Wait()
|
||||
}
|
||||
if q.closed {
|
||||
q.cond.L.Unlock()
|
||||
return
|
||||
}
|
||||
item := q.q[q.front]
|
||||
q.front = (q.front + 1) % BufferSize
|
||||
q.size--
|
||||
q.cond.L.Unlock()
|
||||
|
||||
// process item without holding lock
|
||||
if item.err == nil {
|
||||
// real message
|
||||
if q.l != nil {
|
||||
q.l(q.ctx, item.msg)
|
||||
continue
|
||||
}
|
||||
if q.le != nil {
|
||||
q.le(q.ctx, item.msg, nil)
|
||||
continue
|
||||
}
|
||||
// unhittable
|
||||
continue
|
||||
}
|
||||
// if the listener wants errors, send it.
|
||||
if q.le != nil {
|
||||
q.le(q.ctx, nil, item.err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (q *msgQueue) enqueue(msg []byte) {
|
||||
q.cond.L.Lock()
|
||||
defer q.cond.L.Unlock()
|
||||
|
||||
if q.size == BufferSize {
|
||||
// queue is full, so we're going to drop the msg we got called with.
|
||||
// We also need to record that messages are being dropped, which we
|
||||
// do at the last message in the queue. This potentially makes us
|
||||
// lose 2 messages instead of one, but it's more important at this
|
||||
// point to warn the subscriber that they're losing messages so they
|
||||
// can do something about it.
|
||||
back := (q.front + BufferSize - 1) % BufferSize
|
||||
q.q[back].msg = nil
|
||||
q.q[back].err = ErrDroppedMessages
|
||||
return
|
||||
}
|
||||
// queue is not full, insert the message
|
||||
next := (q.front + q.size) % BufferSize
|
||||
q.q[next].msg = msg
|
||||
q.q[next].err = nil
|
||||
q.size++
|
||||
q.cond.Broadcast()
|
||||
}
|
||||
|
||||
func (q *msgQueue) close() {
|
||||
q.cond.L.Lock()
|
||||
defer q.cond.L.Unlock()
|
||||
defer q.cond.Broadcast()
|
||||
q.closed = true
|
||||
}
|
||||
|
||||
// dropped records an error in the queue that messages might have been dropped
|
||||
func (q *msgQueue) dropped() {
|
||||
q.cond.L.Lock()
|
||||
defer q.cond.L.Unlock()
|
||||
|
||||
if q.size == BufferSize {
|
||||
// queue is full, but we need to record that messages are being dropped,
|
||||
// which we do at the last message in the queue. This potentially drops
|
||||
// another message, but it's more important for the subscriber to know.
|
||||
back := (q.front + BufferSize - 1) % BufferSize
|
||||
q.q[back].msg = nil
|
||||
q.q[back].err = ErrDroppedMessages
|
||||
return
|
||||
}
|
||||
// queue is not full, insert the error
|
||||
next := (q.front + q.size) % BufferSize
|
||||
q.q[next].msg = nil
|
||||
q.q[next].err = ErrDroppedMessages
|
||||
q.size++
|
||||
q.cond.Broadcast()
|
||||
}
|
||||
|
||||
// Pubsub implementation using PostgreSQL.
|
||||
type pgPubsub struct {
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
listenDone chan struct{}
|
||||
pgListener *pq.Listener
|
||||
db *sql.DB
|
||||
mut sync.Mutex
|
||||
queues map[string]map[uuid.UUID]*msgQueue
|
||||
}
|
||||
|
||||
// BufferSize is the maximum number of unhandled messages we will buffer
|
||||
// for a subscriber before dropping messages.
|
||||
const BufferSize = 2048
|
||||
|
||||
// Subscribe calls the listener when an event matching the name is received.
|
||||
func (p *pgPubsub) Subscribe(event string, listener Listener) (cancel func(), err error) {
|
||||
return p.subscribeQueue(event, newMsgQueue(p.ctx, listener, nil))
|
||||
}
|
||||
|
||||
func (p *pgPubsub) SubscribeWithErr(event string, listener ListenerWithErr) (cancel func(), err error) {
|
||||
return p.subscribeQueue(event, newMsgQueue(p.ctx, nil, listener))
|
||||
}
|
||||
|
||||
func (p *pgPubsub) subscribeQueue(event string, newQ *msgQueue) (cancel func(), err error) {
|
||||
p.mut.Lock()
|
||||
defer p.mut.Unlock()
|
||||
defer func() {
|
||||
if err != nil {
|
||||
// if we hit an error, we need to close the queue so we don't
|
||||
// leak its goroutine.
|
||||
newQ.close()
|
||||
}
|
||||
}()
|
||||
|
||||
err = p.pgListener.Listen(event)
|
||||
if errors.Is(err, pq.ErrChannelAlreadyOpen) {
|
||||
// It's ok if it's already open!
|
||||
err = nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("listen: %w", err)
|
||||
}
|
||||
|
||||
var eventQs map[uuid.UUID]*msgQueue
|
||||
var ok bool
|
||||
if eventQs, ok = p.queues[event]; !ok {
|
||||
eventQs = make(map[uuid.UUID]*msgQueue)
|
||||
p.queues[event] = eventQs
|
||||
}
|
||||
id := uuid.New()
|
||||
eventQs[id] = newQ
|
||||
return func() {
|
||||
p.mut.Lock()
|
||||
defer p.mut.Unlock()
|
||||
listeners := p.queues[event]
|
||||
q := listeners[id]
|
||||
q.close()
|
||||
delete(listeners, id)
|
||||
|
||||
if len(listeners) == 0 {
|
||||
_ = p.pgListener.Unlisten(event)
|
||||
}
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (p *pgPubsub) Publish(event string, message []byte) error {
|
||||
// This is safe because we are calling pq.QuoteLiteral. pg_notify doesn't
|
||||
// support the first parameter being a prepared statement.
|
||||
//nolint:gosec
|
||||
_, err := p.db.ExecContext(p.ctx, `select pg_notify(`+pq.QuoteLiteral(event)+`, $1)`, message)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("exec pg_notify: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close closes the pubsub instance.
|
||||
func (p *pgPubsub) Close() error {
|
||||
p.cancel()
|
||||
err := p.pgListener.Close()
|
||||
<-p.listenDone
|
||||
return err
|
||||
}
|
||||
|
||||
// listen begins receiving messages on the pq listener.
|
||||
func (p *pgPubsub) listen() {
|
||||
defer close(p.listenDone)
|
||||
defer p.pgListener.Close()
|
||||
|
||||
var (
|
||||
notif *pq.Notification
|
||||
ok bool
|
||||
)
|
||||
for {
|
||||
select {
|
||||
case <-p.ctx.Done():
|
||||
return
|
||||
case notif, ok = <-p.pgListener.Notify:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
}
|
||||
// A nil notification can be dispatched on reconnect.
|
||||
if notif == nil {
|
||||
p.recordReconnect()
|
||||
continue
|
||||
}
|
||||
p.listenReceive(notif)
|
||||
}
|
||||
}
|
||||
|
||||
func (p *pgPubsub) listenReceive(notif *pq.Notification) {
|
||||
p.mut.Lock()
|
||||
defer p.mut.Unlock()
|
||||
queues, ok := p.queues[notif.Channel]
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
extra := []byte(notif.Extra)
|
||||
for _, q := range queues {
|
||||
q.enqueue(extra)
|
||||
}
|
||||
}
|
||||
|
||||
func (p *pgPubsub) recordReconnect() {
|
||||
p.mut.Lock()
|
||||
defer p.mut.Unlock()
|
||||
for _, listeners := range p.queues {
|
||||
for _, q := range listeners {
|
||||
q.dropped()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// New creates a new Pubsub implementation using a PostgreSQL connection.
|
||||
func New(ctx context.Context, database *sql.DB, connectURL string) (Pubsub, error) {
|
||||
// Creates a new listener using pq.
|
||||
errCh := make(chan error)
|
||||
listener := pq.NewListener(connectURL, time.Second, time.Minute, func(_ pq.ListenerEventType, err error) {
|
||||
// This callback gets events whenever the connection state changes.
|
||||
// Don't send if the errChannel has already been closed.
|
||||
select {
|
||||
case <-errCh:
|
||||
return
|
||||
default:
|
||||
errCh <- err
|
||||
close(errCh)
|
||||
}
|
||||
})
|
||||
select {
|
||||
case err := <-errCh:
|
||||
if err != nil {
|
||||
_ = listener.Close()
|
||||
return nil, xerrors.Errorf("create pq listener: %w", err)
|
||||
}
|
||||
case <-ctx.Done():
|
||||
_ = listener.Close()
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
|
||||
// Start a new context that will be canceled when the pubsub is closed.
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
pgPubsub := &pgPubsub{
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
listenDone: make(chan struct{}),
|
||||
db: database,
|
||||
pgListener: listener,
|
||||
queues: make(map[string]map[uuid.UUID]*msgQueue),
|
||||
}
|
||||
go pgPubsub.listen()
|
||||
|
||||
return pgPubsub, nil
|
||||
}
|
||||
@@ -0,0 +1,140 @@
|
||||
package pubsub
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/testutil"
|
||||
)
|
||||
|
||||
func Test_msgQueue_ListenerWithError(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer cancel()
|
||||
m := make(chan string)
|
||||
e := make(chan error)
|
||||
uut := newMsgQueue(ctx, nil, func(ctx context.Context, msg []byte, err error) {
|
||||
m <- string(msg)
|
||||
e <- err
|
||||
})
|
||||
defer uut.close()
|
||||
|
||||
// We're going to enqueue 4 messages and an error in a loop -- that is, a cycle of 5.
|
||||
// PubsubBufferSize is 2048, which is a power of 2, so a pattern of 5 will not be aligned
|
||||
// when we wrap around the end of the circular buffer. This tests that we correctly handle
|
||||
// the wrapping and aren't dequeueing misaligned data.
|
||||
cycles := (BufferSize / 5) * 2 // almost twice around the ring
|
||||
for j := 0; j < cycles; j++ {
|
||||
for i := 0; i < 4; i++ {
|
||||
uut.enqueue([]byte(fmt.Sprintf("%d%d", j, i)))
|
||||
}
|
||||
uut.dropped()
|
||||
for i := 0; i < 4; i++ {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out")
|
||||
case msg := <-m:
|
||||
require.Equal(t, fmt.Sprintf("%d%d", j, i), msg)
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out")
|
||||
case err := <-e:
|
||||
require.NoError(t, err)
|
||||
}
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out")
|
||||
case msg := <-m:
|
||||
require.Equal(t, "", msg)
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out")
|
||||
case err := <-e:
|
||||
require.ErrorIs(t, err, ErrDroppedMessages)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func Test_msgQueue_Listener(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer cancel()
|
||||
m := make(chan string)
|
||||
uut := newMsgQueue(ctx, func(ctx context.Context, msg []byte) {
|
||||
m <- string(msg)
|
||||
}, nil)
|
||||
defer uut.close()
|
||||
|
||||
// We're going to enqueue 4 messages and an error in a loop -- that is, a cycle of 5.
|
||||
// PubsubBufferSize is 2048, which is a power of 2, so a pattern of 5 will not be aligned
|
||||
// when we wrap around the end of the circular buffer. This tests that we correctly handle
|
||||
// the wrapping and aren't dequeueing misaligned data.
|
||||
cycles := (BufferSize / 5) * 2 // almost twice around the ring
|
||||
for j := 0; j < cycles; j++ {
|
||||
for i := 0; i < 4; i++ {
|
||||
uut.enqueue([]byte(fmt.Sprintf("%d%d", j, i)))
|
||||
}
|
||||
uut.dropped()
|
||||
for i := 0; i < 4; i++ {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out")
|
||||
case msg := <-m:
|
||||
require.Equal(t, fmt.Sprintf("%d%d", j, i), msg)
|
||||
}
|
||||
}
|
||||
// Listener skips over errors, so we only read out the 4 real messages.
|
||||
}
|
||||
}
|
||||
|
||||
func Test_msgQueue_Full(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer cancel()
|
||||
|
||||
firstDequeue := make(chan struct{})
|
||||
allowRead := make(chan struct{})
|
||||
n := 0
|
||||
errors := make(chan error)
|
||||
uut := newMsgQueue(ctx, nil, func(ctx context.Context, msg []byte, err error) {
|
||||
if n == 0 {
|
||||
close(firstDequeue)
|
||||
}
|
||||
<-allowRead
|
||||
if err == nil {
|
||||
require.Equal(t, fmt.Sprintf("%d", n), string(msg))
|
||||
n++
|
||||
return
|
||||
}
|
||||
errors <- err
|
||||
})
|
||||
defer uut.close()
|
||||
|
||||
// we send 2 more than the capacity. One extra because the call to the ListenerFunc blocks
|
||||
// but only after we've dequeued a message, and then another extra because we want to exceed
|
||||
// the capacity, not just reach it.
|
||||
for i := 0; i < BufferSize+2; i++ {
|
||||
uut.enqueue([]byte(fmt.Sprintf("%d", i)))
|
||||
// ensure the first dequeue has happened before proceeding, so that this function isn't racing
|
||||
// against the goroutine that dequeues items.
|
||||
<-firstDequeue
|
||||
}
|
||||
close(allowRead)
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out")
|
||||
case err := <-errors:
|
||||
require.ErrorIs(t, err, ErrDroppedMessages)
|
||||
}
|
||||
// Ok, so we sent 2 more than capacity, but we only read the capacity, that's because the last
|
||||
// message we send doesn't get queued, AND, it bumps a message out of the queue to make room
|
||||
// for the error, so we read 2 less than we sent.
|
||||
require.Equal(t, BufferSize, n)
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
package pubsub
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// genericListener is either a Listener or ListenerWithErr
|
||||
type genericListener struct {
|
||||
l Listener
|
||||
le ListenerWithErr
|
||||
}
|
||||
|
||||
func (g genericListener) send(ctx context.Context, message []byte) {
|
||||
if g.l != nil {
|
||||
g.l(ctx, message)
|
||||
}
|
||||
if g.le != nil {
|
||||
g.le(ctx, message, nil)
|
||||
}
|
||||
}
|
||||
|
||||
// memoryPubsub is an in-memory Pubsub implementation.
|
||||
type memoryPubsub struct {
|
||||
mut sync.RWMutex
|
||||
listeners map[string]map[uuid.UUID]genericListener
|
||||
}
|
||||
|
||||
func (m *memoryPubsub) Subscribe(event string, listener Listener) (cancel func(), err error) {
|
||||
return m.subscribeGeneric(event, genericListener{l: listener})
|
||||
}
|
||||
|
||||
func (m *memoryPubsub) SubscribeWithErr(event string, listener ListenerWithErr) (cancel func(), err error) {
|
||||
return m.subscribeGeneric(event, genericListener{le: listener})
|
||||
}
|
||||
|
||||
func (m *memoryPubsub) subscribeGeneric(event string, listener genericListener) (cancel func(), err error) {
|
||||
m.mut.Lock()
|
||||
defer m.mut.Unlock()
|
||||
|
||||
var listeners map[uuid.UUID]genericListener
|
||||
var ok bool
|
||||
if listeners, ok = m.listeners[event]; !ok {
|
||||
listeners = map[uuid.UUID]genericListener{}
|
||||
m.listeners[event] = listeners
|
||||
}
|
||||
var id uuid.UUID
|
||||
for {
|
||||
id = uuid.New()
|
||||
if _, ok = listeners[id]; !ok {
|
||||
break
|
||||
}
|
||||
}
|
||||
listeners[id] = listener
|
||||
return func() {
|
||||
m.mut.Lock()
|
||||
defer m.mut.Unlock()
|
||||
listeners := m.listeners[event]
|
||||
delete(listeners, id)
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (m *memoryPubsub) Publish(event string, message []byte) error {
|
||||
m.mut.RLock()
|
||||
defer m.mut.RUnlock()
|
||||
listeners, ok := m.listeners[event]
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
var wg sync.WaitGroup
|
||||
for _, listener := range listeners {
|
||||
wg.Add(1)
|
||||
listener := listener
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
listener.send(context.Background(), message)
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (*memoryPubsub) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func NewInMemory() Pubsub {
|
||||
return &memoryPubsub{
|
||||
listeners: make(map[string]map[uuid.UUID]genericListener),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
package pubsub_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/coderd/database/pubsub"
|
||||
)
|
||||
|
||||
func TestPubsubMemory(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("Legacy", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
pubsub := pubsub.NewInMemory()
|
||||
event := "test"
|
||||
data := "testing"
|
||||
messageChannel := make(chan []byte)
|
||||
cancelFunc, err := pubsub.Subscribe(event, func(ctx context.Context, message []byte) {
|
||||
messageChannel <- message
|
||||
})
|
||||
require.NoError(t, err)
|
||||
defer cancelFunc()
|
||||
go func() {
|
||||
err = pubsub.Publish(event, []byte(data))
|
||||
assert.NoError(t, err)
|
||||
}()
|
||||
message := <-messageChannel
|
||||
assert.Equal(t, string(message), data)
|
||||
})
|
||||
|
||||
t.Run("WithErr", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
pubsub := pubsub.NewInMemory()
|
||||
event := "test"
|
||||
data := "testing"
|
||||
messageChannel := make(chan []byte)
|
||||
cancelFunc, err := pubsub.SubscribeWithErr(event, func(ctx context.Context, message []byte, err error) {
|
||||
assert.NoError(t, err) // memory pubsub never sends errors.
|
||||
messageChannel <- message
|
||||
})
|
||||
require.NoError(t, err)
|
||||
defer cancelFunc()
|
||||
go func() {
|
||||
err = pubsub.Publish(event, []byte(data))
|
||||
assert.NoError(t, err)
|
||||
}()
|
||||
message := <-messageChannel
|
||||
assert.Equal(t, string(message), data)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,342 @@
|
||||
//go:build linux
|
||||
|
||||
package pubsub_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"math/rand"
|
||||
"strconv"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/coderd/database/postgres"
|
||||
"github.com/coder/coder/coderd/database/pubsub"
|
||||
"github.com/coder/coder/testutil"
|
||||
)
|
||||
|
||||
// nolint:tparallel,paralleltest
|
||||
func TestPubsub(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
if testing.Short() {
|
||||
t.SkipNow()
|
||||
return
|
||||
}
|
||||
|
||||
t.Run("Postgres", func(t *testing.T) {
|
||||
ctx, cancelFunc := context.WithCancel(context.Background())
|
||||
defer cancelFunc()
|
||||
|
||||
connectionURL, closePg, err := postgres.Open()
|
||||
require.NoError(t, err)
|
||||
defer closePg()
|
||||
db, err := sql.Open("postgres", connectionURL)
|
||||
require.NoError(t, err)
|
||||
defer db.Close()
|
||||
pubsub, err := pubsub.New(ctx, db, connectionURL)
|
||||
require.NoError(t, err)
|
||||
defer pubsub.Close()
|
||||
event := "test"
|
||||
data := "testing"
|
||||
messageChannel := make(chan []byte)
|
||||
unsub, err := pubsub.Subscribe(event, func(ctx context.Context, message []byte) {
|
||||
messageChannel <- message
|
||||
})
|
||||
require.NoError(t, err)
|
||||
defer unsub()
|
||||
go func() {
|
||||
err = pubsub.Publish(event, []byte(data))
|
||||
assert.NoError(t, err)
|
||||
}()
|
||||
message := <-messageChannel
|
||||
assert.Equal(t, string(message), data)
|
||||
})
|
||||
|
||||
t.Run("PostgresCloseCancel", func(t *testing.T) {
|
||||
ctx, cancelFunc := context.WithCancel(context.Background())
|
||||
defer cancelFunc()
|
||||
connectionURL, closePg, err := postgres.Open()
|
||||
require.NoError(t, err)
|
||||
defer closePg()
|
||||
db, err := sql.Open("postgres", connectionURL)
|
||||
require.NoError(t, err)
|
||||
defer db.Close()
|
||||
pubsub, err := pubsub.New(ctx, db, connectionURL)
|
||||
require.NoError(t, err)
|
||||
defer pubsub.Close()
|
||||
cancelFunc()
|
||||
})
|
||||
|
||||
t.Run("NotClosedOnCancelContext", func(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
connectionURL, closePg, err := postgres.Open()
|
||||
require.NoError(t, err)
|
||||
defer closePg()
|
||||
db, err := sql.Open("postgres", connectionURL)
|
||||
require.NoError(t, err)
|
||||
defer db.Close()
|
||||
pubsub, err := pubsub.New(ctx, db, connectionURL)
|
||||
require.NoError(t, err)
|
||||
defer pubsub.Close()
|
||||
|
||||
// Provided context must only be active during NewPubsub, not after.
|
||||
cancel()
|
||||
|
||||
event := "test"
|
||||
data := "testing"
|
||||
messageChannel := make(chan []byte)
|
||||
unsub, err := pubsub.Subscribe(event, func(_ context.Context, message []byte) {
|
||||
messageChannel <- message
|
||||
})
|
||||
require.NoError(t, err)
|
||||
defer unsub()
|
||||
go func() {
|
||||
err = pubsub.Publish(event, []byte(data))
|
||||
assert.NoError(t, err)
|
||||
}()
|
||||
message := <-messageChannel
|
||||
assert.Equal(t, string(message), data)
|
||||
})
|
||||
|
||||
t.Run("ClosePropagatesContextCancellationToSubscription", func(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
|
||||
defer cancel()
|
||||
connectionURL, closePg, err := postgres.Open()
|
||||
require.NoError(t, err)
|
||||
defer closePg()
|
||||
db, err := sql.Open("postgres", connectionURL)
|
||||
require.NoError(t, err)
|
||||
defer db.Close()
|
||||
pubsub, err := pubsub.New(ctx, db, connectionURL)
|
||||
require.NoError(t, err)
|
||||
defer pubsub.Close()
|
||||
|
||||
event := "test"
|
||||
done := make(chan struct{})
|
||||
called := make(chan struct{})
|
||||
unsub, err := pubsub.Subscribe(event, func(subCtx context.Context, _ []byte) {
|
||||
defer close(done)
|
||||
select {
|
||||
case <-subCtx.Done():
|
||||
assert.Fail(t, "context should not be canceled")
|
||||
default:
|
||||
}
|
||||
close(called)
|
||||
select {
|
||||
case <-subCtx.Done():
|
||||
case <-ctx.Done():
|
||||
assert.Fail(t, "timeout waiting for sub context to be canceled")
|
||||
}
|
||||
})
|
||||
require.NoError(t, err)
|
||||
defer unsub()
|
||||
|
||||
go func() {
|
||||
err := pubsub.Publish(event, nil)
|
||||
assert.NoError(t, err)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-called:
|
||||
case <-ctx.Done():
|
||||
require.Fail(t, "timeout waiting for handler to be called")
|
||||
}
|
||||
err = pubsub.Close()
|
||||
require.NoError(t, err)
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-ctx.Done():
|
||||
require.Fail(t, "timeout waiting for handler to finish")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestPubsub_ordering(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancelFunc := context.WithCancel(context.Background())
|
||||
defer cancelFunc()
|
||||
|
||||
connectionURL, closePg, err := postgres.Open()
|
||||
require.NoError(t, err)
|
||||
defer closePg()
|
||||
db, err := sql.Open("postgres", connectionURL)
|
||||
require.NoError(t, err)
|
||||
defer db.Close()
|
||||
ps, err := pubsub.New(ctx, db, connectionURL)
|
||||
require.NoError(t, err)
|
||||
defer ps.Close()
|
||||
event := "test"
|
||||
messageChannel := make(chan []byte, 100)
|
||||
cancelSub, err := ps.Subscribe(event, func(ctx context.Context, message []byte) {
|
||||
// sleep a random amount of time to simulate handlers taking different amount of time
|
||||
// to process, depending on the message
|
||||
// nolint: gosec
|
||||
n := rand.Intn(100)
|
||||
time.Sleep(time.Duration(n) * time.Millisecond)
|
||||
messageChannel <- message
|
||||
})
|
||||
require.NoError(t, err)
|
||||
defer cancelSub()
|
||||
for i := 0; i < 100; i++ {
|
||||
err = ps.Publish(event, []byte(fmt.Sprintf("%d", i)))
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
for i := 0; i < 100; i++ {
|
||||
select {
|
||||
case <-time.After(testutil.WaitShort):
|
||||
t.Fatalf("timed out waiting for message %d", i)
|
||||
case message := <-messageChannel:
|
||||
assert.Equal(t, fmt.Sprintf("%d", i), string(message))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// disconnectTestPort is the hardcoded port for TestPubsub_Disconnect. In this test we need to be able to stop Postgres
|
||||
// and restart it on the same port. If we use an ephemeral port, there is a chance the OS will reallocate before we
|
||||
// start back up. The downside is that if the test crashes and leaves the container up, subsequent test runs will fail
|
||||
// until we manually kill the container.
|
||||
const disconnectTestPort = 26892
|
||||
|
||||
// nolint: paralleltest
|
||||
func TestPubsub_Disconnect(t *testing.T) {
|
||||
// we always use a Docker container for this test, even in CI, since we need to be able to kill
|
||||
// postgres and bring it back on the same port.
|
||||
connectionURL, closePg, err := postgres.OpenContainerized(disconnectTestPort)
|
||||
require.NoError(t, err)
|
||||
defer closePg()
|
||||
db, err := sql.Open("postgres", connectionURL)
|
||||
require.NoError(t, err)
|
||||
defer db.Close()
|
||||
|
||||
ctx, cancelFunc := context.WithTimeout(context.Background(), testutil.WaitSuperLong)
|
||||
defer cancelFunc()
|
||||
ps, err := pubsub.New(ctx, db, connectionURL)
|
||||
require.NoError(t, err)
|
||||
defer ps.Close()
|
||||
event := "test"
|
||||
|
||||
// buffer responses so that when the test completes, goroutines don't get blocked & leak
|
||||
errors := make(chan error, pubsub.BufferSize)
|
||||
messages := make(chan string, pubsub.BufferSize)
|
||||
readOne := func() (m string, e error) {
|
||||
t.Helper()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out")
|
||||
case m = <-messages:
|
||||
// OK
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out")
|
||||
case e = <-errors:
|
||||
// OK
|
||||
}
|
||||
return m, e
|
||||
}
|
||||
|
||||
cancelSub, err := ps.SubscribeWithErr(event, func(ctx context.Context, msg []byte, err error) {
|
||||
messages <- string(msg)
|
||||
errors <- err
|
||||
})
|
||||
require.NoError(t, err)
|
||||
defer cancelSub()
|
||||
|
||||
for i := 0; i < 100; i++ {
|
||||
err = ps.Publish(event, []byte(fmt.Sprintf("%d", i)))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
// make sure we're getting at least one message.
|
||||
m, err := readOne()
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "0", m)
|
||||
|
||||
closePg()
|
||||
// write some more messages until we hit an error
|
||||
j := 100
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out")
|
||||
default:
|
||||
// ok
|
||||
}
|
||||
err = ps.Publish(event, []byte(fmt.Sprintf("%d", j)))
|
||||
j++
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
time.Sleep(testutil.IntervalFast)
|
||||
}
|
||||
|
||||
// restart postgres on the same port --- since we only use LISTEN/NOTIFY it doesn't
|
||||
// matter that the new postgres doesn't have any persisted state from before.
|
||||
_, closeNewPg, err := postgres.OpenContainerized(disconnectTestPort)
|
||||
require.NoError(t, err)
|
||||
defer closeNewPg()
|
||||
|
||||
// now write messages until we DON'T hit an error -- pubsub is back up.
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out")
|
||||
default:
|
||||
// ok
|
||||
}
|
||||
err = ps.Publish(event, []byte(fmt.Sprintf("%d", j)))
|
||||
if err == nil {
|
||||
break
|
||||
}
|
||||
j++
|
||||
time.Sleep(testutil.IntervalFast)
|
||||
}
|
||||
// any message k or higher comes from after the restart.
|
||||
k := j
|
||||
// exceeding the buffer invalidates the test because this causes us to drop messages for reasons other than DB
|
||||
// reconnect
|
||||
require.Less(t, k, pubsub.BufferSize, "exceeded buffer")
|
||||
|
||||
// We don't know how quickly the pubsub will reconnect, so continue to send messages with increasing numbers. As
|
||||
// soon as we see k or higher we know we're getting messages after the restart.
|
||||
go func() {
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
// ok
|
||||
}
|
||||
_ = ps.Publish(event, []byte(fmt.Sprintf("%d", j)))
|
||||
j++
|
||||
time.Sleep(testutil.IntervalFast)
|
||||
}
|
||||
}()
|
||||
|
||||
gotDroppedErr := false
|
||||
for {
|
||||
m, err := readOne()
|
||||
if xerrors.Is(err, pubsub.ErrDroppedMessages) {
|
||||
gotDroppedErr = true
|
||||
continue
|
||||
}
|
||||
require.NoError(t, err, "should only get ErrDroppedMessages")
|
||||
l, err := strconv.Atoi(m)
|
||||
require.NoError(t, err)
|
||||
if l >= k {
|
||||
// exceeding the buffer invalidates the test because this causes us to drop messages for reasons other than
|
||||
// DB reconnect
|
||||
require.Less(t, l, pubsub.BufferSize, "exceeded buffer")
|
||||
break
|
||||
}
|
||||
}
|
||||
require.True(t, gotDroppedErr)
|
||||
}
|
||||
Reference in New Issue
Block a user