package main import ( "context" "errors" "net" "net/http" "net/http/httptest" "net/url" "os" "path/filepath" "testing" "time" "cmroubao/backend-api/internal/config" "cmroubao/backend-api/internal/platform/database" "cmroubao/backend-api/internal/platform/migration" ) func TestLoadConfigUsesERPEnvironmentFile(t *testing.T) { environmentFile := filepath.Join(t.TempDir(), ".env") err := os.WriteFile(environmentFile, []byte( "CMROUBAO_SHUNYUNBAO_URL=https://erp.example.test\n"+ "CMROUBAO_SHUNYUNBAO_USERNAME=dotenv-user\n"+ "CMROUBAO_SHUNYUNBAO_PASSWORD=dotenv-password\n", ), 0o600) if err != nil { t.Fatalf("os.WriteFile() error = %v", err) } cfg, err := loadConfig(func(name string) (string, bool) { if name == config.ShunyunbaoUsernameEnvironment { return "process-user", true } return "", false }, environmentFile) if err != nil { t.Fatalf("loadConfig() error = %v", err) } if cfg.ShunyunbaoURL != "https://erp.example.test" || cfg.ShunyunbaoUsername != "process-user" || cfg.ShunyunbaoPassword != "dotenv-password" { t.Fatalf( "ERP config = %#v", struct { URL string Username string Password string }{cfg.ShunyunbaoURL, cfg.ShunyunbaoUsername, cfg.ShunyunbaoPassword}, ) } } 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: 30 * time.Minute, ReadinessTTL: 2 * time.Minute, ShunyunbaoURL: "https://www.shunyunbaoerp.com", ShunyunbaoTimeout: 30 * time.Second, }, 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"), ) } for _, target := range []string{"/freight", "/freight/import", "/erp"} { request = httptest.NewRequest(http.MethodGet, target, nil) response = httptest.NewRecorder() router.ServeHTTP(response, request) if response.Code != http.StatusSeeOther || response.Header().Get("Location") != "/login?next="+url.QueryEscape(target) { t.Fatalf( "%s status/location = %d / %q", target, response.Code, response.Header().Get("Location"), ) } } request = httptest.NewRequest( http.MethodGet, "/api/v1/freight-orders", nil, ) response = httptest.NewRecorder() router.ServeHTTP(response, request) if response.Code != http.StatusUnauthorized || response.Header().Get("Cache-Control") != "no-store" { t.Fatalf( "freight API status/cache = %d / %q", response.Code, response.Header().Get("Cache-Control"), ) } } type stubMigrationStatusReader struct { statuses []migration.Status err error } func (reader stubMigrationStatusReader) Status( context.Context, ) ([]migration.Status, error) { return reader.statuses, reader.err }