internal/sshd/sshd_test.go

e2a32d5f8d59e4213571c602bd9009b6c8fa86ed
gitbay/internal/sshd/sshd_test.go history · blame · raw

228 lines · 5784 bytes

  1package sshd
  2
  3import (
  4	"bufio"
  5	"bytes"
  6	"crypto/ed25519"
  7	"crypto/rand"
  8	"errors"
  9	"net"
 10	"path/filepath"
 11	"strings"
 12	"testing"
 13	"time"
 14
 15	"golang.org/x/crypto/ssh"
 16
 17	"gitbay.org/gitbay/internal/config"
 18	"gitbay.org/gitbay/internal/store"
 19)
 20
 21// testServer is an embedded server over a fresh store holding alice
 22// with one full-scope key, and a client connected with that key.
 23type testServer struct {
 24	srv    *Server
 25	st     *store.Store
 26	client *ssh.Client
 27	uid    int64
 28	keyID  int64
 29	fp     string
 30}
 31
 32func newTestServer(t *testing.T) testServer {
 33	t.Helper()
 34	root := t.TempDir()
 35	st, err := store.Open(filepath.Join(root, "gitbay.db"))
 36	if err != nil {
 37		t.Fatal(err)
 38	}
 39	t.Cleanup(func() { st.Close() })
 40	if err := st.MigrateUp(); err != nil {
 41		t.Fatal(err)
 42	}
 43	uid, err := st.CreateUser("alice", false)
 44	if err != nil {
 45		t.Fatal(err)
 46	}
 47	_, priv, err := ed25519.GenerateKey(rand.Reader)
 48	if err != nil {
 49		t.Fatal(err)
 50	}
 51	signer, err := ssh.NewSignerFromKey(priv)
 52	if err != nil {
 53		t.Fatal(err)
 54	}
 55	pub := signer.PublicKey()
 56	fp := ssh.FingerprintSHA256(pub)
 57	if err := st.AddSSHKey(uid, fp, pub.Type(), pub.Marshal(), "full", "test"); err != nil {
 58		t.Fatal(err)
 59	}
 60	key, err := st.SSHKeyByFingerprint(fp)
 61	if err != nil {
 62		t.Fatal(err)
 63	}
 64
 65	cfg := config.Default()
 66	cfg.Server.Root = root
 67	srv, err := New(cfg, st)
 68	if err != nil {
 69		t.Fatal(err)
 70	}
 71	ln, err := net.Listen("tcp", "127.0.0.1:0")
 72	if err != nil {
 73		t.Fatal(err)
 74	}
 75	go srv.Serve(ln)
 76	t.Cleanup(func() { ln.Close() })
 77
 78	client, err := ssh.Dial("tcp", ln.Addr().String(), &ssh.ClientConfig{
 79		User:            "git",
 80		Auth:            []ssh.AuthMethod{ssh.PublicKeys(signer)},
 81		HostKeyCallback: ssh.InsecureIgnoreHostKey(),
 82		Timeout:         5 * time.Second,
 83	})
 84	if err != nil {
 85		t.Fatal(err)
 86	}
 87	t.Cleanup(func() { client.Close() })
 88	return testServer{srv: srv, st: st, client: client, uid: uid, keyID: key.ID, fp: fp}
 89}
 90
 91// withBuild gives alice the public repo alice/app and a queued build 1
 92// whose log has one line.
 93func withBuild(t *testing.T, ts testServer) {
 94	t.Helper()
 95	repoID, err := ts.st.CreateRepo("user", ts.uid, "app", "public")
 96	if err != nil {
 97		t.Fatal(err)
 98	}
 99	id, err := ts.st.CreateBuild(repoID, "unit", "abc", "main", `["true"]`, "", "", true)
100	if err != nil {
101		t.Fatal(err)
102	}
103	if err := ts.st.AppendBuildLog(id, []byte("queued\n")); err != nil {
104		t.Fatal(err)
105	}
106}
107
108// followServer starts an embedded server holding alice, her public repo
109// alice/app and a queued build 1 whose log has one line, and returns it
110// with a client connected as alice.
111func followServer(t *testing.T) (*Server, *ssh.Client) {
112	t.Helper()
113	ts := newTestServer(t)
114	withBuild(t, ts)
115	return ts.srv, ts.client
116}
117
118// startFollow runs build log --follow on a new session and returns once
119// the stored line has arrived, so the follow is past its first read.
120func startFollow(t *testing.T, client *ssh.Client, stderr *bytes.Buffer) *ssh.Session {
121	t.Helper()
122	sess, err := client.NewSession()
123	if err != nil {
124		t.Fatal(err)
125	}
126	t.Cleanup(func() { sess.Close() })
127	sess.Stderr = stderr
128	out, err := sess.StdoutPipe()
129	if err != nil {
130		t.Fatal(err)
131	}
132	if err := sess.Start("build log alice/app 1 --follow"); err != nil {
133		t.Fatal(err)
134	}
135	line := make(chan string, 1)
136	go func() {
137		l, _ := bufio.NewReader(out).ReadString('\n')
138		line <- l
139	}()
140	select {
141	case l := <-line:
142		if l != "queued\n" {
143			t.Fatalf("first line %q", l)
144		}
145	case <-time.After(5 * time.Second):
146		t.Fatal("the stored log never arrived")
147	}
148	return sess
149}
150
151func activeSessions(s *Server) int32 {
152	s.mu.Lock()
153	defer s.mu.Unlock()
154	var n int32
155	for c := range s.conns {
156		n += c.active.Load()
157	}
158	return n
159}
160
161// Closing the session channel ends a follow of a queued build while the
162// connection stays up, as a Ctrl-C does over the CLI's shared connection.
163// Nothing else would end it for ten minutes.
164func TestFollowEndsWhenChannelCloses(t *testing.T) {
165	srv, client := followServer(t)
166	var stderr bytes.Buffer
167	sess := startFollow(t, client, &stderr)
168	if n := activeSessions(srv); n != 1 {
169		t.Fatalf("%d sessions active while following, want 1", n)
170	}
171	sess.Close()
172
173	deadline := time.Now().Add(5 * time.Second)
174	for activeSessions(srv) != 0 {
175		if time.Now().After(deadline) {
176			t.Fatal("the follow outlived its channel")
177		}
178		time.Sleep(20 * time.Millisecond)
179	}
180	if _, err := client.NewSession(); err != nil {
181		t.Fatalf("the connection did not survive the channel: %v", err)
182	}
183}
184
185// A command that fails on its own while the server is stopping says
186// nothing about a restart: only a follow that Stop ended does.
187func TestStopLeavesOtherFailuresAlone(t *testing.T) {
188	srv, client := followServer(t)
189	srv.Stop()
190	sess, err := client.NewSession()
191	if err != nil {
192		t.Fatal(err)
193	}
194	defer sess.Close()
195	var stderr bytes.Buffer
196	sess.Stderr = &stderr
197	var exit *ssh.ExitError
198	if err := sess.Run("repo show nosuch/repo"); !errors.As(err, &exit) || exit.ExitStatus() != 3 {
199		t.Fatalf("repo show of a missing repository: %v, want exit 3", err)
200	}
201	if strings.Contains(stderr.String(), "restarting") {
202		t.Errorf("stderr %q", stderr.String())
203	}
204}
205
206// Stop ends a follow with exit 1 and says why, without closing the
207// connection.
208func TestStopEndsFollow(t *testing.T) {
209	srv, client := followServer(t)
210	var stderr bytes.Buffer
211	sess := startFollow(t, client, &stderr)
212	srv.Stop()
213
214	waited := make(chan error, 1)
215	go func() { waited <- sess.Wait() }()
216	select {
217	case err := <-waited:
218		var exit *ssh.ExitError
219		if !errors.As(err, &exit) || exit.ExitStatus() != 1 {
220			t.Fatalf("follow ended with %v, want exit 1", err)
221		}
222	case <-time.After(5 * time.Second):
223		t.Fatal("Stop did not end the follow")
224	}
225	if !strings.Contains(stderr.String(), "gitbay is restarting") {
226		t.Errorf("stderr %q", stderr.String())
227	}
228}