286 lines
8.6 KiB
Go
286 lines
8.6 KiB
Go
package service
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"io"
|
|
"os"
|
|
"os/exec"
|
|
"path/filepath"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/ThreeDotsLabs/watermill"
|
|
"github.com/ThreeDotsLabs/watermill-redisstream/pkg/redisstream"
|
|
"github.com/ThreeDotsLabs/watermill/message"
|
|
"github.com/gochat/gochat/internal/lifecycle"
|
|
"github.com/gochat/gochat/internal/pubsub"
|
|
"github.com/redis/go-redis/v9"
|
|
"github.com/stretchr/testify/require"
|
|
"gorm.io/driver/sqlite"
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
type notificationLifecycleSubscriber struct {
|
|
ready chan struct{}
|
|
messages chan *message.Message
|
|
once sync.Once
|
|
ctx context.Context
|
|
}
|
|
|
|
type notificationAckGate struct {
|
|
topic string
|
|
started chan struct{}
|
|
release chan struct{}
|
|
done chan struct{}
|
|
group string
|
|
startOnce sync.Once
|
|
doneOnce sync.Once
|
|
}
|
|
|
|
func (g *notificationAckGate) DialHook(next redis.DialHook) redis.DialHook { return next }
|
|
func (g *notificationAckGate) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook {
|
|
return next
|
|
}
|
|
func (g *notificationAckGate) ProcessHook(next redis.ProcessHook) redis.ProcessHook {
|
|
return func(ctx context.Context, cmd redis.Cmder) error {
|
|
args := cmd.Args()
|
|
if cmd.Name() != "xack" || len(args) < 3 || fmt.Sprint(args[1]) != g.topic {
|
|
return next(ctx, cmd)
|
|
}
|
|
g.group = fmt.Sprint(args[2])
|
|
g.startOnce.Do(func() { close(g.started) })
|
|
select {
|
|
case <-g.release:
|
|
err := next(ctx, cmd)
|
|
g.doneOnce.Do(func() { close(g.done) })
|
|
return err
|
|
case <-ctx.Done():
|
|
g.doneOnce.Do(func() { close(g.done) })
|
|
return ctx.Err()
|
|
}
|
|
}
|
|
}
|
|
|
|
func (s *notificationLifecycleSubscriber) Subscribe(ctx context.Context, _ string) (<-chan *message.Message, error) {
|
|
s.ctx = ctx
|
|
s.once.Do(func() { close(s.ready) })
|
|
return s.messages, nil
|
|
}
|
|
|
|
func (*notificationLifecycleSubscriber) Close() error { return nil }
|
|
|
|
func (s *notificationLifecycleSubscriber) Send(msg *message.Message) {
|
|
msg.SetContext(s.ctx)
|
|
s.messages <- msg
|
|
}
|
|
|
|
func TestNotificationCloseTimeoutWaitsForHandlerBeforeDependenciesClose(t *testing.T) {
|
|
db, err := gorm.Open(sqlite.Open("file:notification-shutdown-order?mode=memory&cache=shared"), &gorm.Config{})
|
|
require.NoError(t, err)
|
|
router, err := message.NewRouter(message.RouterConfig{CloseTimeout: 10 * time.Millisecond}, watermill.NopLogger{})
|
|
require.NoError(t, err)
|
|
subscriber := ¬ificationLifecycleSubscriber{
|
|
ready: make(chan struct{}),
|
|
messages: make(chan *message.Message),
|
|
}
|
|
handlers := lifecycle.NewHandlerGroup()
|
|
service := &NotificationDeliveryService{router: router, subscriber: subscriber, handlers: handlers}
|
|
|
|
started := make(chan struct{})
|
|
cancelled := make(chan struct{})
|
|
probeDependency := make(chan struct{})
|
|
handlerDBErr := make(chan error, 1)
|
|
router.AddNoPublisherHandler("blocked", "blocked", subscriber, service.track(func(msg *message.Message) error {
|
|
close(started)
|
|
<-msg.Context().Done()
|
|
close(cancelled)
|
|
<-probeDependency
|
|
var one int
|
|
handlerDBErr <- db.Raw("SELECT 1").Scan(&one).Error
|
|
return nil
|
|
}))
|
|
|
|
runDone := make(chan error, 1)
|
|
go func() { runDone <- service.Start(handlers.Context()) }()
|
|
<-subscriber.ready
|
|
subscriber.Send(message.NewMessage("blocked", nil))
|
|
<-started
|
|
|
|
handlers.Stop()
|
|
<-cancelled
|
|
require.ErrorContains(t, service.Close(context.Background()), "router close timeout")
|
|
|
|
waitDone := make(chan struct{})
|
|
go func() {
|
|
handlers.Wait()
|
|
close(waitDone)
|
|
}()
|
|
select {
|
|
case <-waitDone:
|
|
t.Fatal("handler gate opened before the blocked handler completed")
|
|
case <-time.After(20 * time.Millisecond):
|
|
}
|
|
require.NoError(t, db.Exec("SELECT 1").Error)
|
|
|
|
close(probeDependency)
|
|
require.NoError(t, <-handlerDBErr)
|
|
<-waitDone
|
|
sqlDB, err := db.DB()
|
|
require.NoError(t, err)
|
|
require.NoError(t, sqlDB.Close())
|
|
require.Error(t, sqlDB.Ping())
|
|
require.NoError(t, <-runDone)
|
|
}
|
|
|
|
func TestNotificationShutdownWaitsForRedisStreamXAck(t *testing.T) {
|
|
addr := startNotificationRedis(t)
|
|
topic := pubsub.TopicSystemNotification
|
|
gate := ¬ificationAckGate{
|
|
topic: topic,
|
|
started: make(chan struct{}),
|
|
release: make(chan struct{}),
|
|
done: make(chan struct{}),
|
|
}
|
|
subscriberClient := redis.NewClient(&redis.Options{Network: "unix", Addr: addr, PoolSize: 32})
|
|
subscriberClient.AddHook(gate)
|
|
publisherClient := redis.NewClient(&redis.Options{Network: "unix", Addr: addr})
|
|
t.Cleanup(func() {
|
|
_ = subscriberClient.Close()
|
|
_ = publisherClient.Close()
|
|
})
|
|
|
|
service, err := NewNotificationDeliveryService(nil, nil, nil, nil, nil, nil, nil, subscriberClient)
|
|
require.NoError(t, err)
|
|
handlers := lifecycle.NewHandlerGroup()
|
|
service.SetHandlerGroup(handlers)
|
|
|
|
releaseAck := sync.OnceFunc(func() { close(gate.release) })
|
|
t.Cleanup(func() {
|
|
releaseAck()
|
|
_ = service.Close(context.Background())
|
|
})
|
|
runDone := make(chan error, 1)
|
|
go func() { runDone <- service.Start(handlers.Context()) }()
|
|
<-service.router.Running()
|
|
|
|
publisher, err := redisstream.NewPublisher(redisstream.PublisherConfig{Client: publisherClient}, watermill.NopLogger{})
|
|
require.NoError(t, err)
|
|
t.Cleanup(func() { _ = publisher.Close() })
|
|
require.NoError(t, publisher.Publish(topic, message.NewMessage(watermill.NewUUID(), []byte("{"))))
|
|
select {
|
|
case <-gate.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("timed out waiting for XAck")
|
|
}
|
|
|
|
handlers.Wait()
|
|
pending, err := publisherClient.XPending(context.Background(), topic, gate.group).Result()
|
|
require.NoError(t, err)
|
|
require.EqualValues(t, 1, pending.Count)
|
|
|
|
handlers.Stop()
|
|
closeDone := make(chan error, 1)
|
|
go func() { closeDone <- service.Close(context.Background()) }()
|
|
require.NoError(t, publisherClient.Ping(context.Background()).Err())
|
|
|
|
releaseAck()
|
|
require.NoError(t, <-closeDone)
|
|
require.NoError(t, <-runDone)
|
|
select {
|
|
case <-gate.done:
|
|
default:
|
|
t.Fatal("router closed without completing XAck")
|
|
}
|
|
pending, err = publisherClient.XPending(context.Background(), topic, gate.group).Result()
|
|
require.NoError(t, err)
|
|
require.Zero(t, pending.Count)
|
|
require.Error(t, subscriberClient.Ping(context.Background()).Err())
|
|
}
|
|
|
|
func TestNotificationShutdownDeadlineLeavesRedisStreamMessagePending(t *testing.T) {
|
|
addr := startNotificationRedis(t)
|
|
topic := pubsub.TopicSystemNotification
|
|
gate := ¬ificationAckGate{
|
|
topic: topic,
|
|
started: make(chan struct{}),
|
|
release: make(chan struct{}),
|
|
done: make(chan struct{}),
|
|
}
|
|
subscriberClient := redis.NewClient(&redis.Options{Network: "unix", Addr: addr, PoolSize: 32})
|
|
subscriberClient.AddHook(gate)
|
|
publisherClient := redis.NewClient(&redis.Options{Network: "unix", Addr: addr})
|
|
t.Cleanup(func() {
|
|
_ = subscriberClient.Close()
|
|
_ = publisherClient.Close()
|
|
})
|
|
|
|
service, err := NewNotificationDeliveryService(nil, nil, nil, nil, nil, nil, nil, subscriberClient)
|
|
require.NoError(t, err)
|
|
handlers := lifecycle.NewHandlerGroup()
|
|
service.SetHandlerGroup(handlers)
|
|
runDone := make(chan error, 1)
|
|
go func() { runDone <- service.Start(handlers.Context()) }()
|
|
<-service.router.Running()
|
|
|
|
publisher, err := redisstream.NewPublisher(redisstream.PublisherConfig{Client: publisherClient}, watermill.NopLogger{})
|
|
require.NoError(t, err)
|
|
t.Cleanup(func() { _ = publisher.Close() })
|
|
require.NoError(t, publisher.Publish(topic, message.NewMessage(watermill.NewUUID(), []byte("{"))))
|
|
select {
|
|
case <-gate.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("timed out waiting for XAck")
|
|
}
|
|
|
|
handlers.Stop()
|
|
shutdownCtx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
|
|
defer cancel()
|
|
require.ErrorIs(t, service.Close(shutdownCtx), context.DeadlineExceeded)
|
|
require.NoError(t, <-runDone)
|
|
select {
|
|
case <-gate.done:
|
|
default:
|
|
t.Fatal("shutdown returned before the blocked XAck stopped")
|
|
}
|
|
pending, err := publisherClient.XPending(context.Background(), topic, gate.group).Result()
|
|
require.NoError(t, err)
|
|
require.EqualValues(t, 1, pending.Count)
|
|
require.NoError(t, publisherClient.Ping(context.Background()).Err())
|
|
}
|
|
|
|
func startNotificationRedis(t *testing.T) string {
|
|
t.Helper()
|
|
redisServer, err := exec.LookPath("redis-server")
|
|
if err != nil {
|
|
t.Skip("redis-server is required for this integration test")
|
|
}
|
|
dir, err := os.MkdirTemp("", "gochat-redis-")
|
|
require.NoError(t, err)
|
|
t.Cleanup(func() { _ = os.RemoveAll(dir) })
|
|
socket := filepath.Join(dir, "redis.sock")
|
|
|
|
cmd := exec.Command(redisServer,
|
|
"--port", "0",
|
|
"--unixsocket", socket,
|
|
"--unixsocketperm", "700",
|
|
"--save", "",
|
|
"--appendonly", "no",
|
|
)
|
|
cmd.Stdout = io.Discard
|
|
cmd.Stderr = io.Discard
|
|
require.NoError(t, cmd.Start())
|
|
t.Cleanup(func() {
|
|
_ = cmd.Process.Kill()
|
|
_ = cmd.Wait()
|
|
})
|
|
|
|
client := redis.NewClient(&redis.Options{Network: "unix", Addr: socket})
|
|
t.Cleanup(func() { _ = client.Close() })
|
|
require.Eventually(t, func() bool {
|
|
return client.Ping(context.Background()).Err() == nil
|
|
}, 5*time.Second, 10*time.Millisecond)
|
|
return socket
|
|
}
|