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 (
1717 "github.com/picosh/pico/pkg/tunkit"
1818 )
1919
20-func StartSshServer(cfg *PgsConfig, killCh chan error) {
20+func createSshServer(cfg *PgsConfig, ctx context.Context, cacheClearingQueue chan string) (*pssh.SSHServer, error) {
2121 host := shared.GetEnv("PGS_HOST", "0.0.0.0")
2222 port := shared.GetEnv("PGS_SSH_PORT", "2222")
2323 promPort := shared.GetEnv("PGS_PROM_PORT", "9222")
2424 logger := cfg.Logger
2525
26- ctx, cancel := context.WithCancel(context.Background())
27- defer cancel()
28-
29- cacheClearingQueue := make(chan string, 100)
3026 handler := NewUploadAssetHandler(
3127 cfg,
3228 cacheClearingQueue,
......@@ -68,6 +64,17 @@ func StartSshServer(cfg *PgsConfig, killCh chan error) {
6864 },
6965 )
7066
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)
7178 if err != nil {
7279 logger.Error("failed to create ssh server", "err", err.Error())
7380 os.Exit(1)
+107, -41
......@@ -1,12 +1,14 @@
11 package pgs
22
33 import (
4+ "context"
45 "crypto/ed25519"
56 "crypto/rand"
67 "encoding/pem"
78 "fmt"
89 "io"
910 "log/slog"
11+ "net"
1012 "os"
1113 "os/exec"
1214 "path/filepath"
......@@ -16,6 +18,7 @@ import (
1618
1719 pgsdb "github.com/picosh/pico/pkg/apps/pgs/db"
1820 "github.com/picosh/pico/pkg/db"
21+ "github.com/picosh/pico/pkg/pssh"
1922 "github.com/picosh/pico/pkg/shared"
2023 "github.com/picosh/pico/pkg/shared/storage"
2124 "github.com/pkg/sftp"
......@@ -23,6 +26,41 @@ import (
2326 "golang.org/x/crypto/ssh"
2427 )
2528
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+
2664 func TestSshServerSftp(t *testing.T) {
2765 opts := &slog.HandlerOptions{
2866 AddSource: true,
......@@ -43,12 +81,32 @@ func TestSshServerSftp(t *testing.T) {
4381 defer func() {
4482 _ = pubsub.Close()
4583 }()
84+
85+ // Use dynamic port for tests to avoid port conflicts
86+ _ = os.Setenv("PGS_SSH_PORT", "0")
87+
4688 cfg := NewPgsConfig(logger, dbpool, st, pubsub)
4789 done := make(chan error)
4890 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+ }
52110
53111 user := GenerateUser()
54112 // add user's pubkey to the default test account
......@@ -58,7 +116,7 @@ func TestSshServerSftp(t *testing.T) {
58116 Key: shared.KeyForKeyText(user.signer.PublicKey()),
59117 })
60118
61- client, err := user.NewClient()
119+ client, err := user.NewClientAddr(actualAddr)
62120 if err != nil {
63121 t.Error(err)
64122 return
......@@ -96,19 +154,7 @@ func TestSshServerSftp(t *testing.T) {
96154 }
97155
98156 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)
112158 }
113159
114160 func TestSshServerRsync(t *testing.T) {
......@@ -131,12 +177,32 @@ func TestSshServerRsync(t *testing.T) {
131177 defer func() {
132178 _ = pubsub.Close()
133179 }()
180+
181+ // Use dynamic port for tests to avoid port conflicts
182+ _ = os.Setenv("PGS_SSH_PORT", "0")
183+
134184 cfg := NewPgsConfig(logger, dbpool, st, pubsub)
135185 done := make(chan error)
136186 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+ }
140206
141207 user := GenerateUser()
142208 key := shared.KeyForKeyText(user.signer.PublicKey())
......@@ -147,7 +213,7 @@ func TestSshServerRsync(t *testing.T) {
147213 Key: key,
148214 })
149215
150- conn, err := user.NewClient()
216+ conn, err := user.NewClientAddr(actualAddr)
151217 if err != nil {
152218 t.Error(err)
153219 return
......@@ -215,13 +281,22 @@ func TestSshServerRsync(t *testing.T) {
215281 t.Fatal(err)
216282 }
217283
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+
218292 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,
220295 keyFile,
221296 )
222297
223298 // copy files
224- cmd := exec.Command("rsync", "-rv", "-e", eCmd, name+"/", "localhost:/test")
299+ cmd := exec.Command("rsync", "-rv", "-e", eCmd, name+"/", host+":/test")
225300 result, err := cmd.CombinedOutput()
226301 if err != nil {
227302 cfg.Logger.Error("cannot upload", "err", err, "result", string(result))
......@@ -270,7 +345,7 @@ func TestSshServerRsync(t *testing.T) {
270345 _ = os.RemoveAll(dlName)
271346 }()
272347 // download files
273- downloadCmd := exec.Command("rsync", "-rvvv", "-e", eCmd, "localhost:/test/", dlName+"/")
348+ downloadCmd := exec.Command("rsync", "-rvvv", "-e", eCmd, host+":/test/", dlName+"/")
274349 result, err = downloadCmd.CombinedOutput()
275350 if err != nil {
276351 cfg.Logger.Error("cannot download files", "err", err, "result", string(result))
......@@ -306,19 +381,7 @@ func TestSshServerRsync(t *testing.T) {
306381 }
307382
308383 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)
322385 }
323386
324387 type UserSSH struct {
......@@ -347,9 +410,7 @@ func (s UserSSH) MustCmd(client *ssh.Client, patch []byte, cmd string) string {
347410 return res
348411 }
349412
350-func (s UserSSH) NewClient() (*ssh.Client, error) {
351- host := "localhost:2222"
352-
413+func (s UserSSH) NewClientAddr(addr string) (*ssh.Client, error) {
353414 config := &ssh.ClientConfig{
354415 User: s.username,
355416 Auth: []ssh.AuthMethod{
......@@ -358,10 +419,15 @@ func (s UserSSH) NewClient() (*ssh.Client, error) {
358419 HostKeyCallback: ssh.InsecureIgnoreHostKey(),
359420 }
360421
361- client, err := ssh.Dial("tcp", host, config)
422+ client, err := ssh.Dial("tcp", addr, config)
362423 return client, err
363424 }
364425
426+func (s UserSSH) NewClient() (*ssh.Client, error) {
427+ // Default to localhost:2222 for backward compatibility
428+ return s.NewClientAddr("localhost:2222")
429+}
430+
365431 func (s UserSSH) Cmd(client *ssh.Client, patch []byte, cmd string) (string, error) {
366432 session, err := client.NewSession()
367433 if err != nil {