package main import ( "context" "errors" "net" "net/http" "net/http/httptest" "path/filepath" "testing" "time" "cmroubao/backend-api/internal/config" "cmroubao/backend-api/internal/platform/database" "cmroubao/backend-api/internal/platform/migration" ) func TestShutdownServerForceClosesAfterGracefulTimeout(t *testing.T) { handlerStarted := make(chan struct{}) releaseHandler := make(chan struct{}) server := &http.Server{ Handler: http.HandlerFunc(func(http.ResponseWriter, *http.Request) { close(handlerStarted) <-releaseHandler }), ReadHeaderTimeout: time.Second, } listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("net.Listen() error = %v", err) } serverErrors := make(chan error, 1) go func() { serverErrors <- server.Serve(listener) }() requestFinished := make(chan struct{}) go func() { defer close(requestFinished) client := http.Client{Timeout: time.Second} response, err := client.Get("http://" + listener.Addr().String()) if err == nil { _ = response.Body.Close() } }() select { case <-handlerStarted: case <-time.After(time.Second): t.Fatal("handler did not start") } err = shutdownServer(server, serverErrors, 10*time.Millisecond) if err == nil { t.Fatal("shutdownServer() error = nil") } close(releaseHandler) select { case <-requestFinished: case <-time.After(time.Second): t.Fatal("request did not finish after force close") } connection, dialErr := net.DialTimeout( "tcp", listener.Addr().String(), 100*time.Millisecond, ) if dialErr == nil { _ = connection.Close() t.Fatal("listener still accepted connections after force close") } } func TestShutdownServerAcceptsNormalServerClose(t *testing.T) { serverErrors := make(chan error, 1) serverErrors <- http.ErrServerClosed server := &http.Server{} if err := shutdownServer(server, serverErrors, time.Second); err != nil && !errors.Is(err, http.ErrServerClosed) { t.Fatalf("shutdownServer() error = %v", err) } } func TestRequireCurrentMigrationsAcceptsAppliedVersions(t *testing.T) { reader := stubMigrationStatusReader{ statuses: []migration.Status{ {Version: 1, Applied: true}, {Version: 2, Applied: true}, }, } if err := requireCurrentMigrations(context.Background(), reader); err != nil { t.Fatalf("requireCurrentMigrations() error = %v", err) } } func TestRequireCurrentMigrationsRejectsPendingVersion(t *testing.T) { reader := stubMigrationStatusReader{ statuses: []migration.Status{ {Version: 1, Applied: true}, {Version: 2, Applied: false}, }, } err := requireCurrentMigrations(context.Background(), reader) if err == nil || err.Error() != "database migrations are pending; run migrate up" { t.Fatalf("requireCurrentMigrations() error = %v", err) } } func TestRequireCurrentMigrationsHidesStatusFailure(t *testing.T) { reader := stubMigrationStatusReader{err: errors.New("database details")} err := requireCurrentMigrations(context.Background(), reader) if err == nil || err.Error() != "database migration status failed" { t.Fatalf("requireCurrentMigrations() error = %v", err) } } func TestBuildRouterRegistersProtectedLogoutRoute(t *testing.T) { ctx := context.Background() db, err := database.Open( ctx, filepath.Join(t.TempDir(), "api.db"), ) if err != nil { t.Fatalf("database.Open() error = %v", err) } defer db.Close() runner, err := migration.New(db) if err != nil { t.Fatalf("migration.New() error = %v", err) } if _, err := runner.Up(ctx); err != nil { t.Fatalf("migration.Up() error = %v", err) } router, err := buildRouter(ctx, config.Config{ AssetDirectory: filepath.Join(t.TempDir(), "assets"), ClaimLease: 10 * time.Minute, RunningLease: 90 * time.Second, ReadinessTTL: 2 * time.Minute, }, db) if err != nil { t.Fatalf("buildRouter() error = %v", err) } request := httptest.NewRequest(http.MethodPost, "/logout", nil) response := httptest.NewRecorder() router.ServeHTTP(response, request) if response.Code != http.StatusSeeOther || response.Header().Get("Location") != "/login?next=%2Ftasks" { t.Fatalf( "status/location = %d / %q", response.Code, response.Header().Get("Location"), ) } } type stubMigrationStatusReader struct { statuses []migration.Status err error } func (reader stubMigrationStatusReader) Status( context.Context, ) ([]migration.Status, error) { return reader.statuses, reader.err }