Files
cmroubao/backend-api/cmd/api/main_test.go
T

173 lines
4.3 KiB
Go

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
}