Commit 47413a9

Eric Bower  ·  2025-12-25 23:32:29 -0500 EST
parent 1d500f9
chore(pipe): added tests
1 files changed,  +661, -0
+661, -0
......@@ -0,0 +1,661 @@
1+package pipe
2+
3+import (
4+ "context"
5+ "crypto/ed25519"
6+ "crypto/rand"
7+ "fmt"
8+ "io"
9+ "log/slog"
10+ "os"
11+ "strings"
12+ "testing"
13+ "time"
14+
15+ "github.com/antoniomika/syncmap"
16+ "github.com/picosh/pico/pkg/db"
17+ "github.com/picosh/pico/pkg/db/stub"
18+ "github.com/picosh/pico/pkg/pssh"
19+ "github.com/picosh/pico/pkg/shared"
20+ psub "github.com/picosh/pubsub"
21+ "github.com/picosh/utils"
22+ "github.com/prometheus/client_golang/prometheus"
23+ "golang.org/x/crypto/ssh"
24+)
25+
26+type TestDB struct {
27+ *stub.StubDB
28+ Users []*db.User
29+ Pubkeys []*db.PublicKey
30+ Features []*db.FeatureFlag
31+}
32+
33+func NewTestDB(logger *slog.Logger) *TestDB {
34+ return &TestDB{
35+ StubDB: stub.NewStubDB(logger),
36+ }
37+}
38+
39+func (t *TestDB) FindUserByPubkey(key string) (*db.User, error) {
40+ for _, pk := range t.Pubkeys {
41+ if pk.Key == key {
42+ return t.FindUser(pk.UserID)
43+ }
44+ }
45+ return nil, fmt.Errorf("user not found for pubkey")
46+}
47+
48+func (t *TestDB) FindUser(userID string) (*db.User, error) {
49+ for _, user := range t.Users {
50+ if user.ID == userID {
51+ return user, nil
52+ }
53+ }
54+ return nil, fmt.Errorf("user not found")
55+}
56+
57+func (t *TestDB) FindUserByName(name string) (*db.User, error) {
58+ for _, user := range t.Users {
59+ if user.Name == name {
60+ return user, nil
61+ }
62+ }
63+ return nil, fmt.Errorf("user not found")
64+}
65+
66+func (t *TestDB) FindFeature(userID, name string) (*db.FeatureFlag, error) {
67+ for _, ff := range t.Features {
68+ if ff.UserID == userID && ff.Name == name {
69+ return ff, nil
70+ }
71+ }
72+ return nil, fmt.Errorf("feature not found")
73+}
74+
75+func (t *TestDB) HasFeatureByUser(userID string, feature string) bool {
76+ ff, err := t.FindFeature(userID, feature)
77+ if err != nil {
78+ return false
79+ }
80+ return ff.IsValid()
81+}
82+
83+func (t *TestDB) InsertAccessLog(_ *db.AccessLog) error {
84+ return nil
85+}
86+
87+func (t *TestDB) Close() error {
88+ return nil
89+}
90+
91+func (t *TestDB) AddUser(user *db.User) {
92+ t.Users = append(t.Users, user)
93+}
94+
95+func (t *TestDB) AddPubkey(pubkey *db.PublicKey) {
96+ t.Pubkeys = append(t.Pubkeys, pubkey)
97+}
98+
99+type TestSSHServer struct {
100+ Cfg *shared.ConfigSite
101+ DBPool *TestDB
102+ Cancel context.CancelFunc
103+}
104+
105+func NewTestSSHServer(t *testing.T) *TestSSHServer {
106+ t.Helper()
107+
108+ opts := &slog.HandlerOptions{
109+ AddSource: true,
110+ Level: slog.LevelDebug,
111+ }
112+ logger := slog.New(slog.NewTextHandler(os.Stdout, opts))
113+
114+ dbpool := NewTestDB(logger)
115+
116+ cfg := &shared.ConfigSite{
117+ Domain: "pipe.test",
118+ Port: "2222",
119+ PortOverride: "2222",
120+ Protocol: "ssh",
121+ Logger: logger,
122+ Space: "pipe",
123+ }
124+
125+ ctx, cancel := context.WithCancel(context.Background())
126+
127+ pubsub := psub.NewMulticast(logger)
128+ handler := &CliHandler{
129+ Logger: logger,
130+ DBPool: dbpool,
131+ PubSub: pubsub,
132+ Cfg: cfg,
133+ Waiters: syncmap.New[string, []string](),
134+ Access: syncmap.New[string, []string](),
135+ }
136+
137+ sshAuth := shared.NewSshAuthHandler(dbpool, logger, "pipe")
138+
139+ prometheus.DefaultRegisterer = prometheus.NewRegistry()
140+
141+ server, err := pssh.NewSSHServerWithConfig(
142+ ctx,
143+ logger,
144+ "pipe-ssh-test",
145+ "localhost",
146+ "2222",
147+ "9222",
148+ "../../ssh_data/term_info_ed25519",
149+ func(conn ssh.ConnMetadata, key ssh.PublicKey) (*ssh.Permissions, error) {
150+ perms, _ := sshAuth.PubkeyAuthHandler(conn, key)
151+ if perms == nil {
152+ perms = &ssh.Permissions{
153+ Extensions: map[string]string{
154+ "pubkey": utils.KeyForKeyText(key),
155+ },
156+ }
157+ }
158+ return perms, nil
159+ },
160+ []pssh.SSHServerMiddleware{
161+ Middleware(handler),
162+ pssh.LogMiddleware(handler, dbpool),
163+ },
164+ nil,
165+ nil,
166+ )
167+
168+ if err != nil {
169+ t.Fatalf("failed to create ssh server: %v", err)
170+ }
171+
172+ go func() {
173+ if err := server.ListenAndServe(); err != nil {
174+ logger.Error("serve", "err", err.Error())
175+ }
176+ }()
177+
178+ time.Sleep(100 * time.Millisecond)
179+
180+ return &TestSSHServer{
181+ Cfg: cfg,
182+ DBPool: dbpool,
183+ Cancel: cancel,
184+ }
185+}
186+
187+func (s *TestSSHServer) Shutdown() {
188+ s.Cancel()
189+ time.Sleep(10 * time.Millisecond)
190+}
191+
192+type UserSSH struct {
193+ username string
194+ signer ssh.Signer
195+ privateKey []byte
196+}
197+
198+func GenerateUser(username string) UserSSH {
199+ _, userKey, err := ed25519.GenerateKey(rand.Reader)
200+ if err != nil {
201+ panic(err)
202+ }
203+
204+ b, err := ssh.MarshalPrivateKey(userKey, "")
205+ if err != nil {
206+ panic(err)
207+ }
208+
209+ userSigner, err := ssh.NewSignerFromKey(userKey)
210+ if err != nil {
211+ panic(err)
212+ }
213+
214+ return UserSSH{
215+ username: username,
216+ signer: userSigner,
217+ privateKey: b.Bytes,
218+ }
219+}
220+
221+func (u UserSSH) PublicKey() string {
222+ return utils.KeyForKeyText(u.signer.PublicKey())
223+}
224+
225+func (u UserSSH) NewClient() (*ssh.Client, error) {
226+ config := &ssh.ClientConfig{
227+ User: u.username,
228+ Auth: []ssh.AuthMethod{
229+ ssh.PublicKeys(u.signer),
230+ },
231+ HostKeyCallback: ssh.InsecureIgnoreHostKey(),
232+ }
233+
234+ return ssh.Dial("tcp", "localhost:2222", config)
235+}
236+
237+func (u UserSSH) RunCommand(client *ssh.Client, cmd string) (string, error) {
238+ session, err := client.NewSession()
239+ if err != nil {
240+ return "", err
241+ }
242+ defer func() { _ = session.Close() }()
243+
244+ stdoutPipe, err := session.StdoutPipe()
245+ if err != nil {
246+ return "", err
247+ }
248+
249+ stderrPipe, err := session.StderrPipe()
250+ if err != nil {
251+ return "", err
252+ }
253+
254+ if err := session.Start(cmd); err != nil {
255+ return "", err
256+ }
257+
258+ stdout := new(strings.Builder)
259+ stderr := new(strings.Builder)
260+ _, _ = io.Copy(stdout, stdoutPipe)
261+ _, _ = io.Copy(stderr, stderrPipe)
262+
263+ _ = session.Wait()
264+ return stdout.String() + stderr.String(), nil
265+}
266+
267+func (u UserSSH) RunCommandWithStdin(client *ssh.Client, cmd string, stdin string) (string, error) {
268+ session, err := client.NewSession()
269+ if err != nil {
270+ return "", err
271+ }
272+ defer func() { _ = session.Close() }()
273+
274+ stdinPipe, err := session.StdinPipe()
275+ if err != nil {
276+ return "", err
277+ }
278+
279+ stdoutPipe, err := session.StdoutPipe()
280+ if err != nil {
281+ return "", err
282+ }
283+
284+ if err := session.Start(cmd); err != nil {
285+ return "", err
286+ }
287+
288+ _, err = stdinPipe.Write([]byte(stdin))
289+ if err != nil {
290+ return "", err
291+ }
292+ _ = stdinPipe.Close()
293+
294+ buf := new(strings.Builder)
295+ _, err = io.Copy(buf, stdoutPipe)
296+ if err != nil {
297+ return "", err
298+ }
299+
300+ _ = session.Wait()
301+ return buf.String(), nil
302+}
303+
304+func RegisterUserWithServer(server *TestSSHServer, user UserSSH) {
305+ dbUser := &db.User{
306+ ID: user.username + "-id",
307+ Name: user.username,
308+ }
309+ server.DBPool.AddUser(dbUser)
310+ server.DBPool.AddPubkey(&db.PublicKey{
311+ ID: user.username + "-pubkey-id",
312+ UserID: dbUser.ID,
313+ Key: user.PublicKey(),
314+ })
315+}
316+
317+func TestLs_UnauthenticatedUserDenied(t *testing.T) {
318+ server := NewTestSSHServer(t)
319+ defer server.Shutdown()
320+
321+ user := GenerateUser("anonymous")
322+
323+ client, err := user.NewClient()
324+ if err != nil {
325+ t.Fatalf("failed to connect: %v", err)
326+ }
327+ defer func() { _ = client.Close() }()
328+
329+ output, err := user.RunCommand(client, "ls")
330+ if err != nil {
331+ t.Logf("command error (expected): %v", err)
332+ }
333+
334+ if !strings.Contains(output, "access denied") {
335+ t.Errorf("expected 'access denied', got: %s", output)
336+ }
337+}
338+
339+func TestLs_AuthenticatedUser(t *testing.T) {
340+ server := NewTestSSHServer(t)
341+ defer server.Shutdown()
342+
343+ user := GenerateUser("alice")
344+ RegisterUserWithServer(server, user)
345+
346+ client, err := user.NewClient()
347+ if err != nil {
348+ t.Fatalf("failed to connect: %v", err)
349+ }
350+ defer func() { _ = client.Close() }()
351+
352+ output, err := user.RunCommand(client, "ls")
353+ if err != nil {
354+ t.Logf("command completed with: %v", err)
355+ }
356+
357+ if strings.Contains(output, "access denied") {
358+ t.Errorf("authenticated user should not get access denied, got: %s", output)
359+ }
360+
361+ if !strings.Contains(output, "no pubsub channels found") {
362+ t.Errorf("expected 'no pubsub channels found' for empty state, got: %s", output)
363+ }
364+}
365+
366+func TestPubSub_BasicFlow(t *testing.T) {
367+ server := NewTestSSHServer(t)
368+ defer server.Shutdown()
369+
370+ user := GenerateUser("alice")
371+ RegisterUserWithServer(server, user)
372+
373+ subClient, err := user.NewClient()
374+ if err != nil {
375+ t.Fatalf("failed to connect subscriber: %v", err)
376+ }
377+ defer func() { _ = subClient.Close() }()
378+
379+ pubClient, err := user.NewClient()
380+ if err != nil {
381+ t.Fatalf("failed to connect publisher: %v", err)
382+ }
383+ defer func() { _ = pubClient.Close() }()
384+
385+ subSession, err := subClient.NewSession()
386+ if err != nil {
387+ t.Fatalf("failed to create sub session: %v", err)
388+ }
389+ defer func() { _ = subSession.Close() }()
390+
391+ subStdout, err := subSession.StdoutPipe()
392+ if err != nil {
393+ t.Fatalf("failed to get sub stdout: %v", err)
394+ }
395+
396+ if err := subSession.Start("sub testtopic -c"); err != nil {
397+ t.Fatalf("failed to start sub: %v", err)
398+ }
399+
400+ time.Sleep(100 * time.Millisecond)
401+
402+ testMessage := "hello from pub"
403+ _, err = user.RunCommandWithStdin(pubClient, "pub testtopic -c", testMessage)
404+ if err != nil {
405+ t.Logf("pub command completed: %v", err)
406+ }
407+
408+ received := make([]byte, len(testMessage)+10)
409+ n, err := subStdout.Read(received)
410+ if err != nil && err != io.EOF {
411+ t.Logf("read error: %v", err)
412+ }
413+
414+ receivedStr := string(received[:n])
415+ if !strings.Contains(receivedStr, testMessage) {
416+ t.Errorf("subscriber did not receive message, got: %q, want: %q", receivedStr, testMessage)
417+ }
418+}
419+
420+func TestPubSub_PublicTopic(t *testing.T) {
421+ server := NewTestSSHServer(t)
422+ defer server.Shutdown()
423+
424+ alice := GenerateUser("alice")
425+ bob := GenerateUser("bob")
426+ RegisterUserWithServer(server, alice)
427+ RegisterUserWithServer(server, bob)
428+
429+ subClient, err := bob.NewClient()
430+ if err != nil {
431+ t.Fatalf("failed to connect subscriber: %v", err)
432+ }
433+ defer func() { _ = subClient.Close() }()
434+
435+ pubClient, err := alice.NewClient()
436+ if err != nil {
437+ t.Fatalf("failed to connect publisher: %v", err)
438+ }
439+ defer func() { _ = pubClient.Close() }()
440+
441+ subSession, err := subClient.NewSession()
442+ if err != nil {
443+ t.Fatalf("failed to create sub session: %v", err)
444+ }
445+ defer func() { _ = subSession.Close() }()
446+
447+ subStdout, err := subSession.StdoutPipe()
448+ if err != nil {
449+ t.Fatalf("failed to get sub stdout: %v", err)
450+ }
451+
452+ if err := subSession.Start("sub publictopic -p -c"); err != nil {
453+ t.Fatalf("failed to start sub: %v", err)
454+ }
455+
456+ time.Sleep(100 * time.Millisecond)
457+
458+ testMessage := "public message"
459+ _, err = alice.RunCommandWithStdin(pubClient, "pub publictopic -p -c", testMessage)
460+ if err != nil {
461+ t.Logf("pub command completed: %v", err)
462+ }
463+
464+ received := make([]byte, len(testMessage)+10)
465+ n, err := subStdout.Read(received)
466+ if err != nil && err != io.EOF {
467+ t.Logf("read error: %v", err)
468+ }
469+
470+ receivedStr := string(received[:n])
471+ if !strings.Contains(receivedStr, testMessage) {
472+ t.Errorf("subscriber did not receive public message, got: %q, want: %q", receivedStr, testMessage)
473+ }
474+}
475+
476+func TestPipe_Bidirectional(t *testing.T) {
477+ server := NewTestSSHServer(t)
478+ defer server.Shutdown()
479+
480+ alice := GenerateUser("alice")
481+ bob := GenerateUser("bob")
482+ RegisterUserWithServer(server, alice)
483+ RegisterUserWithServer(server, bob)
484+
485+ aliceClient, err := alice.NewClient()
486+ if err != nil {
487+ t.Fatalf("failed to connect alice: %v", err)
488+ }
489+ defer func() { _ = aliceClient.Close() }()
490+
491+ bobClient, err := bob.NewClient()
492+ if err != nil {
493+ t.Fatalf("failed to connect bob: %v", err)
494+ }
495+ defer func() { _ = bobClient.Close() }()
496+
497+ aliceSession, err := aliceClient.NewSession()
498+ if err != nil {
499+ t.Fatalf("failed to create alice session: %v", err)
500+ }
501+ defer func() { _ = aliceSession.Close() }()
502+
503+ aliceStdin, err := aliceSession.StdinPipe()
504+ if err != nil {
505+ t.Fatalf("failed to get alice stdin: %v", err)
506+ }
507+
508+ aliceStdout, err := aliceSession.StdoutPipe()
509+ if err != nil {
510+ t.Fatalf("failed to get alice stdout: %v", err)
511+ }
512+
513+ if err := aliceSession.Start("pipe pipetopic -p -c"); err != nil {
514+ t.Fatalf("failed to start alice pipe: %v", err)
515+ }
516+
517+ time.Sleep(100 * time.Millisecond)
518+
519+ bobSession, err := bobClient.NewSession()
520+ if err != nil {
521+ t.Fatalf("failed to create bob session: %v", err)
522+ }
523+ defer func() { _ = bobSession.Close() }()
524+
525+ bobStdin, err := bobSession.StdinPipe()
526+ if err != nil {
527+ t.Fatalf("failed to get bob stdin: %v", err)
528+ }
529+
530+ bobStdout, err := bobSession.StdoutPipe()
531+ if err != nil {
532+ t.Fatalf("failed to get bob stdout: %v", err)
533+ }
534+
535+ if err := bobSession.Start("pipe pipetopic -p -c"); err != nil {
536+ t.Fatalf("failed to start bob pipe: %v", err)
537+ }
538+
539+ time.Sleep(100 * time.Millisecond)
540+
541+ aliceMsg := "hello from alice\n"
542+ _, err = aliceStdin.Write([]byte(aliceMsg))
543+ if err != nil {
544+ t.Fatalf("alice failed to write: %v", err)
545+ }
546+
547+ bobReceived := make([]byte, 100)
548+ n, err := bobStdout.Read(bobReceived)
549+ if err != nil && err != io.EOF {
550+ t.Logf("bob read error: %v", err)
551+ }
552+ if !strings.Contains(string(bobReceived[:n]), "hello from alice") {
553+ t.Errorf("bob did not receive alice's message, got: %q", string(bobReceived[:n]))
554+ }
555+
556+ bobMsg := "hello from bob\n"
557+ _, err = bobStdin.Write([]byte(bobMsg))
558+ if err != nil {
559+ t.Fatalf("bob failed to write: %v", err)
560+ }
561+
562+ aliceReceived := make([]byte, 100)
563+ n, err = aliceStdout.Read(aliceReceived)
564+ if err != nil && err != io.EOF {
565+ t.Logf("alice read error: %v", err)
566+ }
567+ if !strings.Contains(string(aliceReceived[:n]), "hello from bob") {
568+ t.Errorf("alice did not receive bob's message, got: %q", string(aliceReceived[:n]))
569+ }
570+}
571+
572+func TestPipe_AutoGeneratedTopic(t *testing.T) {
573+ server := NewTestSSHServer(t)
574+ defer server.Shutdown()
575+
576+ user := GenerateUser("alice")
577+ RegisterUserWithServer(server, user)
578+
579+ client, err := user.NewClient()
580+ if err != nil {
581+ t.Fatalf("failed to connect: %v", err)
582+ }
583+ defer func() { _ = client.Close() }()
584+
585+ session, err := client.NewSession()
586+ if err != nil {
587+ t.Fatalf("failed to create session: %v", err)
588+ }
589+ defer func() { _ = session.Close() }()
590+
591+ stdout, err := session.StdoutPipe()
592+ if err != nil {
593+ t.Fatalf("failed to get stdout: %v", err)
594+ }
595+
596+ if err := session.Start("pipe"); err != nil {
597+ t.Fatalf("failed to start pipe: %v", err)
598+ }
599+
600+ received := make([]byte, 200)
601+ n, err := stdout.Read(received)
602+ if err != nil && err != io.EOF {
603+ t.Logf("read error: %v", err)
604+ }
605+
606+ output := string(received[:n])
607+ if !strings.Contains(output, "subscribe to this topic") {
608+ t.Errorf("expected topic subscription instructions, got: %q", output)
609+ }
610+}
611+
612+func TestAccessControl_AllowedUserViaFullPath(t *testing.T) {
613+ server := NewTestSSHServer(t)
614+ defer server.Shutdown()
615+
616+ alice := GenerateUser("alice")
617+ bob := GenerateUser("bob")
618+ RegisterUserWithServer(server, alice)
619+ RegisterUserWithServer(server, bob)
620+
621+ aliceClient, err := alice.NewClient()
622+ if err != nil {
623+ t.Fatalf("failed to connect alice: %v", err)
624+ }
625+ defer func() { _ = aliceClient.Close() }()
626+
627+ aliceSession, err := aliceClient.NewSession()
628+ if err != nil {
629+ t.Fatalf("failed to create alice session: %v", err)
630+ }
631+ defer func() { _ = aliceSession.Close() }()
632+
633+ aliceStdout, err := aliceSession.StdoutPipe()
634+ if err != nil {
635+ t.Fatalf("failed to get alice stdout: %v", err)
636+ }
637+
638+ if err := aliceSession.Start("sub sharedtopic -a alice,bob -c"); err != nil {
639+ t.Fatalf("failed to start alice sub: %v", err)
640+ }
641+
642+ time.Sleep(100 * time.Millisecond)
643+
644+ bobClient, err := bob.NewClient()
645+ if err != nil {
646+ t.Fatalf("failed to connect bob: %v", err)
647+ }
648+ defer func() { _ = bobClient.Close() }()
649+
650+ _, err = bob.RunCommandWithStdin(bobClient, "pub alice/sharedtopic -c", "bob allowed")
651+ if err != nil {
652+ t.Logf("bob pub completed: %v", err)
653+ }
654+
655+ aliceReceived := make([]byte, 100)
656+ n, _ := aliceStdout.Read(aliceReceived)
657+
658+ if !strings.Contains(string(aliceReceived[:n]), "bob allowed") {
659+ t.Errorf("alice should receive bob's message on shared topic, got: %q", string(aliceReceived[:n]))
660+ }
661+}