Files
gochat/backend/internal/service/notification_delivery_lifecycle_test.go
T
2026-08-24 12:24:34 +08:00

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 := &notificationLifecycleSubscriber{
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 := &notificationAckGate{
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 := &notificationAckGate{
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
}