Commit da9a246
Eric Bower
·
2026-02-26 09:04:53 -0500 EST
parent 391c4f9
refactor(pgs): use dynamic port for ssh tests
2 files changed,
+119,
-46
+12,
-5
| ... | ... | @@ -17,16 +17,12 @@ import ( | |
| 17 | 17 | "github.com/picosh/pico/pkg/tunkit" | |
| 18 | 18 | ) | |
| 19 | 19 | ||
| 20 | - | func StartSshServer(cfg *PgsConfig, killCh chan error) { | |
| 20 | + | func createSshServer(cfg *PgsConfig, ctx context.Context, cacheClearingQueue chan string) (*pssh.SSHServer, error) { | |
| 21 | 21 | host := shared.GetEnv("PGS_HOST", "0.0.0.0") | |
| 22 | 22 | port := shared.GetEnv("PGS_SSH_PORT", "2222") | |
| 23 | 23 | promPort := shared.GetEnv("PGS_PROM_PORT", "9222") | |
| 24 | 24 | logger := cfg.Logger | |
| 25 | 25 | ||
| 26 | - | ctx, cancel := context.WithCancel(context.Background()) | |
| 27 | - | defer cancel() | |
| 28 | - | ||
| 29 | - | cacheClearingQueue := make(chan string, 100) | |
| 30 | 26 | handler := NewUploadAssetHandler( | |
| 31 | 27 | cfg, | |
| 32 | 28 | cacheClearingQueue, |
| ... | ... | @@ -68,6 +64,17 @@ func StartSshServer(cfg *PgsConfig, killCh chan error) { | |
| 68 | 64 | }, | |
| 69 | 65 | ) | |
| 70 | 66 | ||
| 67 | + | return server, err | |
| 68 | + | } | |
| 69 | + | ||
| 70 | + | func StartSshServer(cfg *PgsConfig, killCh chan error) { | |
| 71 | + | ctx, cancel := context.WithCancel(context.Background()) | |
| 72 | + | defer cancel() | |
| 73 | + | ||
| 74 | + | cacheClearingQueue := make(chan string, 100) | |
| 75 | + | logger := cfg.Logger | |
| 76 | + | ||
| 77 | + | server, err := createSshServer(cfg, ctx, cacheClearingQueue) | |
| 71 | 78 | if err != nil { | |
| 72 | 79 | logger.Error("failed to create ssh server", "err", err.Error()) | |
| 73 | 80 | os.Exit(1) |
+107,
-41
| ... | ... | @@ -1,12 +1,14 @@ | |
| 1 | 1 | package pgs | |
| 2 | 2 | ||
| 3 | 3 | import ( | |
| 4 | + | "context" | |
| 4 | 5 | "crypto/ed25519" | |
| 5 | 6 | "crypto/rand" | |
| 6 | 7 | "encoding/pem" | |
| 7 | 8 | "fmt" | |
| 8 | 9 | "io" | |
| 9 | 10 | "log/slog" | |
| 11 | + | "net" | |
| 10 | 12 | "os" | |
| 11 | 13 | "os/exec" | |
| 12 | 14 | "path/filepath" |
| ... | ... | @@ -16,6 +18,7 @@ import ( | |
| 16 | 18 | ||
| 17 | 19 | pgsdb "github.com/picosh/pico/pkg/apps/pgs/db" | |
| 18 | 20 | "github.com/picosh/pico/pkg/db" | |
| 21 | + | "github.com/picosh/pico/pkg/pssh" | |
| 19 | 22 | "github.com/picosh/pico/pkg/shared" | |
| 20 | 23 | "github.com/picosh/pico/pkg/shared/storage" | |
| 21 | 24 | "github.com/pkg/sftp" |
| ... | ... | @@ -23,6 +26,41 @@ import ( | |
| 23 | 26 | "golang.org/x/crypto/ssh" | |
| 24 | 27 | ) | |
| 25 | 28 | ||
| 29 | + | func StartSshServerForTesting(cfg *PgsConfig, killCh chan error) *pssh.SSHServer { | |
| 30 | + | ctx, cancel := context.WithCancel(context.Background()) | |
| 31 | + | defer func() { | |
| 32 | + | // Cancel is deferred to avoid being called in error path | |
| 33 | + | // It will be called from the cleanup goroutine | |
| 34 | + | _ = cancel | |
| 35 | + | }() | |
| 36 | + | ||
| 37 | + | cacheClearingQueue := make(chan string, 100) | |
| 38 | + | logger := cfg.Logger | |
| 39 | + | ||
| 40 | + | server, err := createSshServer(cfg, ctx, cacheClearingQueue) | |
| 41 | + | if err != nil { | |
| 42 | + | logger.Error("failed to create ssh server", "err", err.Error()) | |
| 43 | + | cancel() // Clean up if server creation fails | |
| 44 | + | return nil | |
| 45 | + | } | |
| 46 | + | ||
| 47 | + | logger.Info("Starting SSH server", "addr", server.Config.ListenAddr) | |
| 48 | + | go func() { | |
| 49 | + | if err = server.ListenAndServe(); err != nil { | |
| 50 | + | logger.Error("serve", "err", err.Error()) | |
| 51 | + | } | |
| 52 | + | }() | |
| 53 | + | ||
| 54 | + | go func() { | |
| 55 | + | // Wait for kill signal and clean up | |
| 56 | + | <-killCh | |
| 57 | + | logger.Info("stopping ssh server") | |
| 58 | + | cancel() | |
| 59 | + | }() | |
| 60 | + | ||
| 61 | + | return server | |
| 62 | + | } | |
| 63 | + | ||
| 26 | 64 | func TestSshServerSftp(t *testing.T) { | |
| 27 | 65 | opts := &slog.HandlerOptions{ | |
| 28 | 66 | AddSource: true, |
| ... | ... | @@ -43,12 +81,32 @@ func TestSshServerSftp(t *testing.T) { | |
| 43 | 81 | defer func() { | |
| 44 | 82 | _ = pubsub.Close() | |
| 45 | 83 | }() | |
| 84 | + | ||
| 85 | + | // Use dynamic port for tests to avoid port conflicts | |
| 86 | + | _ = os.Setenv("PGS_SSH_PORT", "0") | |
| 87 | + | ||
| 46 | 88 | cfg := NewPgsConfig(logger, dbpool, st, pubsub) | |
| 47 | 89 | done := make(chan error) | |
| 48 | 90 | prometheus.DefaultRegisterer = prometheus.NewRegistry() | |
| 49 | - | go StartSshServer(cfg, done) | |
| 50 | - | // Hack to wait for startup | |
| 51 | - | time.Sleep(time.Millisecond * 100) | |
| 91 | + | ||
| 92 | + | var server *pssh.SSHServer | |
| 93 | + | go func() { | |
| 94 | + | server = StartSshServerForTesting(cfg, done) | |
| 95 | + | }() | |
| 96 | + | ||
| 97 | + | // Wait for server to be ready and get the actual listening address | |
| 98 | + | var actualAddr string | |
| 99 | + | for i := 0; i < 100; i++ { | |
| 100 | + | if server != nil && server.Listener != nil { | |
| 101 | + | actualAddr = server.Listener.Addr().String() | |
| 102 | + | break | |
| 103 | + | } | |
| 104 | + | time.Sleep(20 * time.Millisecond) | |
| 105 | + | } | |
| 106 | + | ||
| 107 | + | if actualAddr == "" { | |
| 108 | + | t.Fatal("server listener not ready") | |
| 109 | + | } | |
| 52 | 110 | ||
| 53 | 111 | user := GenerateUser() | |
| 54 | 112 | // add user's pubkey to the default test account |
| ... | ... | @@ -58,7 +116,7 @@ func TestSshServerSftp(t *testing.T) { | |
| 58 | 116 | Key: shared.KeyForKeyText(user.signer.PublicKey()), | |
| 59 | 117 | }) | |
| 60 | 118 | ||
| 61 | - | client, err := user.NewClient() | |
| 119 | + | client, err := user.NewClientAddr(actualAddr) | |
| 62 | 120 | if err != nil { | |
| 63 | 121 | t.Error(err) | |
| 64 | 122 | return |
| ... | ... | @@ -96,19 +154,7 @@ func TestSshServerSftp(t *testing.T) { | |
| 96 | 154 | } | |
| 97 | 155 | ||
| 98 | 156 | close(done) | |
| 99 | - | ||
| 100 | - | p, err := os.FindProcess(os.Getpid()) | |
| 101 | - | if err != nil { | |
| 102 | - | t.Fatal(err) | |
| 103 | - | return | |
| 104 | - | } | |
| 105 | - | ||
| 106 | - | err = p.Signal(os.Interrupt) | |
| 107 | - | if err != nil { | |
| 108 | - | t.Fatal(err) | |
| 109 | - | return | |
| 110 | - | } | |
| 111 | - | <-time.After(10 * time.Millisecond) | |
| 157 | + | time.Sleep(100 * time.Millisecond) | |
| 112 | 158 | } | |
| 113 | 159 | ||
| 114 | 160 | func TestSshServerRsync(t *testing.T) { |
| ... | ... | @@ -131,12 +177,32 @@ func TestSshServerRsync(t *testing.T) { | |
| 131 | 177 | defer func() { | |
| 132 | 178 | _ = pubsub.Close() | |
| 133 | 179 | }() | |
| 180 | + | ||
| 181 | + | // Use dynamic port for tests to avoid port conflicts | |
| 182 | + | _ = os.Setenv("PGS_SSH_PORT", "0") | |
| 183 | + | ||
| 134 | 184 | cfg := NewPgsConfig(logger, dbpool, st, pubsub) | |
| 135 | 185 | done := make(chan error) | |
| 136 | 186 | prometheus.DefaultRegisterer = prometheus.NewRegistry() | |
| 137 | - | go StartSshServer(cfg, done) | |
| 138 | - | // Hack to wait for startup | |
| 139 | - | time.Sleep(time.Millisecond * 100) | |
| 187 | + | ||
| 188 | + | var server *pssh.SSHServer | |
| 189 | + | go func() { | |
| 190 | + | server = StartSshServerForTesting(cfg, done) | |
| 191 | + | }() | |
| 192 | + | ||
| 193 | + | // Wait for server to be ready and get the actual listening address | |
| 194 | + | var actualAddr string | |
| 195 | + | for i := 0; i < 100; i++ { | |
| 196 | + | if server != nil && server.Listener != nil { | |
| 197 | + | actualAddr = server.Listener.Addr().String() | |
| 198 | + | break | |
| 199 | + | } | |
| 200 | + | time.Sleep(20 * time.Millisecond) | |
| 201 | + | } | |
| 202 | + | ||
| 203 | + | if actualAddr == "" { | |
| 204 | + | t.Fatal("server listener not ready") | |
| 205 | + | } | |
| 140 | 206 | ||
| 141 | 207 | user := GenerateUser() | |
| 142 | 208 | key := shared.KeyForKeyText(user.signer.PublicKey()) |
| ... | ... | @@ -147,7 +213,7 @@ func TestSshServerRsync(t *testing.T) { | |
| 147 | 213 | Key: key, | |
| 148 | 214 | }) | |
| 149 | 215 | ||
| 150 | - | conn, err := user.NewClient() | |
| 216 | + | conn, err := user.NewClientAddr(actualAddr) | |
| 151 | 217 | if err != nil { | |
| 152 | 218 | t.Error(err) | |
| 153 | 219 | return |
| ... | ... | @@ -215,13 +281,22 @@ func TestSshServerRsync(t *testing.T) { | |
| 215 | 281 | t.Fatal(err) | |
| 216 | 282 | } | |
| 217 | 283 | ||
| 284 | + | // Extract port from actualAddr (format: "0.0.0.0:XXXXX") | |
| 285 | + | _, port, err := net.SplitHostPort(actualAddr) | |
| 286 | + | if err != nil { | |
| 287 | + | t.Fatalf("failed to parse server address: %v", err) | |
| 288 | + | } | |
| 289 | + | // Use localhost for rsync (works regardless of IPv4/IPv6 binding) | |
| 290 | + | host := "localhost" | |
| 291 | + | ||
| 218 | 292 | eCmd := fmt.Sprintf( | |
| 219 | - | "ssh -p 2222 -o IdentitiesOnly=yes -i %s -o StrictHostKeyChecking=no", | |
| 293 | + | "ssh -p %s -o IdentitiesOnly=yes -i %s -o StrictHostKeyChecking=no", | |
| 294 | + | port, | |
| 220 | 295 | keyFile, | |
| 221 | 296 | ) | |
| 222 | 297 | ||
| 223 | 298 | // copy files | |
| 224 | - | cmd := exec.Command("rsync", "-rv", "-e", eCmd, name+"/", "localhost:/test") | |
| 299 | + | cmd := exec.Command("rsync", "-rv", "-e", eCmd, name+"/", host+":/test") | |
| 225 | 300 | result, err := cmd.CombinedOutput() | |
| 226 | 301 | if err != nil { | |
| 227 | 302 | cfg.Logger.Error("cannot upload", "err", err, "result", string(result)) |
| ... | ... | @@ -270,7 +345,7 @@ func TestSshServerRsync(t *testing.T) { | |
| 270 | 345 | _ = os.RemoveAll(dlName) | |
| 271 | 346 | }() | |
| 272 | 347 | // download files | |
| 273 | - | downloadCmd := exec.Command("rsync", "-rvvv", "-e", eCmd, "localhost:/test/", dlName+"/") | |
| 348 | + | downloadCmd := exec.Command("rsync", "-rvvv", "-e", eCmd, host+":/test/", dlName+"/") | |
| 274 | 349 | result, err = downloadCmd.CombinedOutput() | |
| 275 | 350 | if err != nil { | |
| 276 | 351 | cfg.Logger.Error("cannot download files", "err", err, "result", string(result)) |
| ... | ... | @@ -306,19 +381,7 @@ func TestSshServerRsync(t *testing.T) { | |
| 306 | 381 | } | |
| 307 | 382 | ||
| 308 | 383 | close(done) | |
| 309 | - | ||
| 310 | - | p, err := os.FindProcess(os.Getpid()) | |
| 311 | - | if err != nil { | |
| 312 | - | t.Fatal(err) | |
| 313 | - | return | |
| 314 | - | } | |
| 315 | - | ||
| 316 | - | err = p.Signal(os.Interrupt) | |
| 317 | - | if err != nil { | |
| 318 | - | t.Fatal(err) | |
| 319 | - | return | |
| 320 | - | } | |
| 321 | - | <-time.After(10 * time.Millisecond) | |
| 384 | + | time.Sleep(100 * time.Millisecond) | |
| 322 | 385 | } | |
| 323 | 386 | ||
| 324 | 387 | type UserSSH struct { |
| ... | ... | @@ -347,9 +410,7 @@ func (s UserSSH) MustCmd(client *ssh.Client, patch []byte, cmd string) string { | |
| 347 | 410 | return res | |
| 348 | 411 | } | |
| 349 | 412 | ||
| 350 | - | func (s UserSSH) NewClient() (*ssh.Client, error) { | |
| 351 | - | host := "localhost:2222" | |
| 352 | - | ||
| 413 | + | func (s UserSSH) NewClientAddr(addr string) (*ssh.Client, error) { | |
| 353 | 414 | config := &ssh.ClientConfig{ | |
| 354 | 415 | User: s.username, | |
| 355 | 416 | Auth: []ssh.AuthMethod{ |
| ... | ... | @@ -358,10 +419,15 @@ func (s UserSSH) NewClient() (*ssh.Client, error) { | |
| 358 | 419 | HostKeyCallback: ssh.InsecureIgnoreHostKey(), | |
| 359 | 420 | } | |
| 360 | 421 | ||
| 361 | - | client, err := ssh.Dial("tcp", host, config) | |
| 422 | + | client, err := ssh.Dial("tcp", addr, config) | |
| 362 | 423 | return client, err | |
| 363 | 424 | } | |
| 364 | 425 | ||
| 426 | + | func (s UserSSH) NewClient() (*ssh.Client, error) { | |
| 427 | + | // Default to localhost:2222 for backward compatibility | |
| 428 | + | return s.NewClientAddr("localhost:2222") | |
| 429 | + | } | |
| 430 | + | ||
| 365 | 431 | func (s UserSSH) Cmd(client *ssh.Client, patch []byte, cmd string) (string, error) { | |
| 366 | 432 | session, err := client.NewSession() | |
| 367 | 433 | if err != nil { |