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) }