Files
gochat/backend/internal/app/runtime_test.go
T
Rogeeandrogee 798ea43c2f HH-442: isolate runtime processes and harden shutdown (#90)
* HH-442: isolate runtime processes and harden shutdown

* HH-442: harden worker shutdown races

* HH-442: gate dependency shutdown on active handlers

---------

Co-authored-by: Rogee <rogee@ipao.vip>
2026-08-22 02:38:15 +08:00

238 lines
8.3 KiB
Go

package app
import (
"bufio"
"context"
"errors"
"net"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/gin-gonic/gin"
"github.com/gochat/gochat/internal/config"
"github.com/gochat/gochat/internal/model"
"github.com/gochat/gochat/internal/pubsub"
"github.com/gochat/gochat/internal/worker"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/require"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
)
type shutdownOrderPubSub struct {
db *gorm.DB
jobID uint
status string
err error
}
type shutdownRedisPubSub struct{ client *redis.Client }
func (p *shutdownRedisPubSub) Publish(ctx context.Context, topic string, _ pubsub.Event) error {
return p.client.Set(ctx, topic, "available", 0).Err()
}
func (*shutdownRedisPubSub) Subscribe(context.Context, string, pubsub.EventHandler) error {
return nil
}
func (*shutdownRedisPubSub) Unsubscribe(context.Context, string) error { return nil }
func (p *shutdownRedisPubSub) Close() error { return p.client.Close() }
func (*shutdownOrderPubSub) Publish(context.Context, string, pubsub.Event) error { return nil }
func (*shutdownOrderPubSub) Subscribe(context.Context, string, pubsub.EventHandler) error {
return nil
}
func (*shutdownOrderPubSub) Unsubscribe(context.Context, string) error { return nil }
func (p *shutdownOrderPubSub) Close() error {
var job model.BackgroundJob
p.err = p.db.First(&job, p.jobID).Error
p.status = job.Status
return nil
}
func TestWebProcessDoesNotConsumeJobs(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.BackgroundJob{}))
pool := worker.NewWorkerPool(db)
var handled atomic.Int32
pool.Register("web_must_not_consume", func(context.Context, *model.BackgroundJob) error {
handled.Add(1)
return nil
})
_, err = pool.Enqueue(context.Background(), "web_must_not_consume", nil)
require.NoError(t, err)
engine := gin.New()
engine.GET("/live", func(c *gin.Context) { c.Status(http.StatusOK) })
application := &App{
config: &config.Config{Server: config.ServerConfig{
Host: "127.0.0.1", Port: 0, ReadHeaderTimeoutS: 1, ReadTimeoutS: 1,
WriteTimeoutS: 1, IdleTimeoutS: 1, ShutdownTimeoutS: 1, MaxHeaderBytes: 1024,
}},
db: db, engine: engine, workerPool: pool,
}
ctx, cancel := context.WithCancel(context.Background())
go func() {
time.Sleep(50 * time.Millisecond)
cancel()
}()
require.NoError(t, application.runWeb(ctx))
require.Zero(t, handled.Load())
}
func TestWorkerProcessConsumesJobsWithoutHTTP(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.BackgroundJob{}))
pool := worker.NewWorkerPoolWithOptions(db, worker.WithPollInterval(5*time.Millisecond))
ctx, cancel := context.WithCancel(context.Background())
pool.Register("worker_only", func(context.Context, *model.BackgroundJob) error {
cancel()
return nil
})
_, err = pool.Enqueue(context.Background(), "worker_only", nil)
require.NoError(t, err)
occupied, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer occupied.Close()
port := occupied.Addr().(*net.TCPAddr).Port
application := &App{
config: &config.Config{Server: config.ServerConfig{Host: "127.0.0.1", Port: port, ShutdownTimeoutS: 1}},
db: db, engine: gin.New(), workerPool: pool,
}
require.NoError(t, application.runWorker(ctx))
}
func TestShutdownWaitsForCancelledJobPersistenceBeforeClosingDependencies(t *testing.T) {
db, err := gorm.Open(sqlite.Open("file:shutdown-order?mode=memory&cache=shared"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.BackgroundJob{}))
pool := worker.NewWorkerPoolWithOptions(db,
worker.WithPollInterval(5*time.Millisecond),
worker.WithBackoff(func(int) time.Duration { return 0 }),
)
started := make(chan struct{})
pool.Register("shutdown_order", func(ctx context.Context, _ *model.BackgroundJob) error {
close(started)
<-ctx.Done()
return ctx.Err()
})
job, err := pool.Enqueue(context.Background(), "shutdown_order", nil)
require.NoError(t, err)
require.NoError(t, pool.Start())
<-started
ps := &shutdownOrderPubSub{db: db, jobID: job.ID}
application := &App{db: db, pubsub: ps, workerPool: pool}
require.ErrorIs(t, application.Shutdown(20*time.Millisecond), context.DeadlineExceeded)
require.NoError(t, ps.err)
require.Equal(t, model.BackgroundJobStatusRetrying, ps.status)
}
func TestHTTPShutdownWaitsForSlowHandlerBeforeClosingDependencies(t *testing.T) {
db, err := gorm.Open(sqlite.Open("file:http-shutdown-order?mode=memory&cache=shared"), &gorm.Config{})
require.NoError(t, err)
redisServer := miniredis.RunT(t)
redisClient := redis.NewClient(&redis.Options{Addr: redisServer.Addr()})
ps := &shutdownRedisPubSub{client: redisClient}
started := make(chan struct{})
release := make(chan struct{})
handlerErr := make(chan error, 1)
engine := gin.New()
application := &App{db: db, pubsub: ps, engine: engine}
engine.GET("/slow", func(c *gin.Context) {
close(started)
<-release
var one int
dbErr := db.Raw("SELECT 1").Scan(&one).Error
redisErr := ps.Publish(c.Request.Context(), "shutdown-order", pubsub.Event{Type: "still-open"})
handlerErr <- errors.Join(dbErr, redisErr)
c.Status(http.StatusNoContent)
})
server := httptest.NewServer(application.Handler())
defer server.Close()
requestDone := make(chan error, 1)
go func() {
response, requestErr := http.Get(server.URL + "/slow")
if requestErr == nil {
_ = response.Body.Close()
}
requestDone <- requestErr
}()
<-started
shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond)
defer cancel()
require.ErrorIs(t, server.Config.Shutdown(shutdownCtx), context.DeadlineExceeded)
shutdownDone := make(chan error, 1)
go func() { shutdownDone <- application.shutdown(shutdownCtx) }()
select {
case err := <-shutdownDone:
t.Fatalf("dependencies closed before slow handler completed: %v", err)
case <-time.After(20 * time.Millisecond):
}
close(release)
require.NoError(t, <-handlerErr)
require.NoError(t, <-requestDone)
require.NoError(t, <-shutdownDone)
sqlDB, err := db.DB()
require.NoError(t, err)
require.Error(t, sqlDB.Ping())
require.Error(t, ps.Publish(context.Background(), "shutdown-order", pubsub.Event{Type: "closed"}))
}
func TestHTTPServerLimits(t *testing.T) {
engine := gin.New()
engine.GET("/", func(c *gin.Context) { c.Status(http.StatusOK) })
application := &App{config: &config.Config{Server: config.ServerConfig{
Host: "127.0.0.1", Port: 0, ReadHeaderTimeoutS: 2, ReadTimeoutS: 3,
WriteTimeoutS: 4, IdleTimeoutS: 5, MaxHeaderBytes: 1024,
}}, engine: engine}
server := application.HTTPServer()
require.Equal(t, 2*time.Second, server.ReadHeaderTimeout)
require.Equal(t, 3*time.Second, server.ReadTimeout)
require.Equal(t, 4*time.Second, server.WriteTimeout)
require.Equal(t, 5*time.Second, server.IdleTimeout)
require.Equal(t, 1024, server.MaxHeaderBytes)
server.ReadHeaderTimeout = 50 * time.Millisecond
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
serveDone := make(chan error, 1)
go func() { serveDone <- server.Serve(listener) }()
slow, err := net.Dial("tcp", listener.Addr().String())
require.NoError(t, err)
_, err = slow.Write([]byte("GET / HTTP/1.1\r\nHost: localhost\r\nX-Slow:"))
require.NoError(t, err)
require.NoError(t, slow.SetReadDeadline(time.Now().Add(500*time.Millisecond)))
started := time.Now()
line, readErr := bufio.NewReader(slow).ReadString('\n')
require.Less(t, time.Since(started), 300*time.Millisecond)
require.True(t, readErr != nil || strings.Contains(line, "400"), "expected timeout rejection, got %q (%v)", line, readErr)
require.NoError(t, slow.Close())
large, err := net.Dial("tcp", listener.Addr().String())
require.NoError(t, err)
_, err = large.Write([]byte("GET / HTTP/1.1\r\nHost: localhost\r\nX-Large: " + strings.Repeat("a", 10_000) + "\r\n\r\n"))
require.NoError(t, err)
require.NoError(t, large.SetReadDeadline(time.Now().Add(time.Second)))
status, err := bufio.NewReader(large).ReadString('\n')
require.NoError(t, err)
require.Contains(t, status, "431")
require.NoError(t, large.Close())
require.NoError(t, server.Shutdown(context.Background()))
require.ErrorIs(t, <-serveDone, http.ErrServerClosed)
}