Parent directory

script_test.go

18471 bytes
  1package testscript
  2
  3import (
  4	"bytes"
  5	"context"
  6	"encoding/json"
  7	"flag"
  8	"fmt"
  9	"io"
 10	"math/rand"
 11	"net"
 12	"net/http"
 13	"net/url"
 14	"os"
 15	"os/exec"
 16	"path/filepath"
 17	"runtime"
 18	"strconv"
 19	"strings"
 20	"testing"
 21	"time"
 22
 23	"github.com/charmbracelet/keygen"
 24	"github.com/charmbracelet/soft-serve/pkg/config"
 25	"github.com/charmbracelet/soft-serve/pkg/db"
 26	"github.com/charmbracelet/soft-serve/pkg/test"
 27	"github.com/rogpeppe/go-internal/testscript"
 28	"github.com/spf13/cobra"
 29	"golang.org/x/crypto/ssh"
 30)
 31
 32var (
 33	update  = flag.Bool("update", false, "update script files")
 34	binPath string
 35)
 36
 37func PrepareBuildCommand(binPath string) *exec.Cmd {
 38	_, disableRaceSet := os.LookupEnv("SOFT_SERVE_DISABLE_RACE_CHECKS")
 39	if disableRaceSet {
 40		// don't add the -race flag
 41		return exec.Command("go", "build", "-cover", "-o", binPath, filepath.Join("..", "cmd", "soft")) //nolint:noctx
 42	}
 43	return exec.Command("go", "build", "-race", "-cover", "-o", binPath, filepath.Join("..", "cmd", "soft")) //nolint:noctx
 44}
 45
 46func TestMain(m *testing.M) {
 47	tmp, err := os.MkdirTemp("", "soft-serve*")
 48	if err != nil {
 49		fmt.Fprintf(os.Stderr, "failed to create temporary directory: %s", err)
 50		os.Exit(1)
 51	}
 52	defer os.RemoveAll(tmp)
 53
 54	binPath = filepath.Join(tmp, "soft")
 55	if runtime.GOOS == "windows" {
 56		binPath += ".exe"
 57	}
 58
 59	// Build the soft binary with -cover flag.
 60	cmd := PrepareBuildCommand(binPath)
 61	if err := cmd.Run(); err != nil {
 62		fmt.Fprintf(os.Stderr, "failed to build soft-serve binary: %s", err)
 63		os.Exit(1)
 64	}
 65
 66	// Run tests
 67	os.Exit(m.Run())
 68}
 69
 70func TestScript(t *testing.T) {
 71	flag.Parse()
 72
 73	mkkey := func(name string) (string, *keygen.SSHKeyPair) {
 74		path := filepath.Join(t.TempDir(), name)
 75		pair, err := keygen.New(path, keygen.WithKeyType(keygen.Ed25519), keygen.WithWrite())
 76		if err != nil {
 77			t.Fatal(err)
 78		}
 79		return path, pair
 80	}
 81
 82	admin1Key, admin1 := mkkey("admin1")
 83	_, admin2 := mkkey("admin2")
 84	user1Key, user1 := mkkey("user1")
 85	attackerKey, attacker := mkkey("attacker")
 86	attackerSigner := &maliciousSigner{
 87		publicKey: admin1.PublicKey(),
 88	}
 89
 90	testscript.Run(t, testscript.Params{
 91		Dir:                 "./testdata/",
 92		UpdateScripts:       *update,
 93		RequireExplicitExec: true,
 94		Cmds: map[string]func(ts *testscript.TestScript, neg bool, args []string){
 95			"soft":                   cmdSoft("admin", admin1.Signer()),
 96			"usoft":                  cmdSoft("user1", user1.Signer()),
 97			"attacksoft":             cmdSoft("attacker", attackerSigner, attacker.Signer()),
 98			"ksoft":                  cmdKeylessSoft,
 99			"git":                    cmdGit(admin1Key),
100			"ugit":                   cmdGit(user1Key),
101			"agit":                   cmdGit(attackerKey),
102			"curl":                   cmdCurl,
103			"mkfile":                 cmdMkfile,
104			"envfile":                cmdEnvfile,
105			"readfile":               cmdReadfile,
106			"dos2unix":               cmdDos2Unix,
107			"new-webhook":            cmdNewWebhook,
108			"ensureserverrunning":    cmdEnsureServerRunning,
109			"ensureservernotrunning": cmdEnsureServerNotRunning,
110			"stopserver":             cmdStopserver,
111			"ui":                     cmdUI(admin1.Signer()),
112			"uui":                    cmdUI(user1.Signer()),
113		},
114		Setup: func(e *testscript.Env) error {
115			// Add binPath to PATH
116			e.Setenv("PATH", fmt.Sprintf("%s%c%s", filepath.Dir(binPath), os.PathListSeparator, e.Getenv("PATH")))
117
118			data := t.TempDir()
119			sshPort := test.RandomPort()
120			sshListen := fmt.Sprintf("localhost:%d", sshPort)
121			gitPort := test.RandomPort()
122			gitListen := fmt.Sprintf("localhost:%d", gitPort)
123			httpPort := test.RandomPort()
124			httpListen := fmt.Sprintf("localhost:%d", httpPort)
125			statsPort := test.RandomPort()
126			statsListen := fmt.Sprintf("localhost:%d", statsPort)
127			serverName := "Test Soft Serve"
128
129			e.Setenv("DATA_PATH", data)
130			e.Setenv("SSH_PORT", fmt.Sprintf("%d", sshPort))
131			e.Setenv("HTTP_PORT", fmt.Sprintf("%d", httpPort))
132			e.Setenv("STATS_PORT", fmt.Sprintf("%d", statsPort))
133			e.Setenv("GIT_PORT", fmt.Sprintf("%d", gitPort))
134			e.Setenv("ADMIN1_AUTHORIZED_KEY", admin1.AuthorizedKey())
135			e.Setenv("ADMIN2_AUTHORIZED_KEY", admin2.AuthorizedKey())
136			e.Setenv("USER1_AUTHORIZED_KEY", user1.AuthorizedKey())
137			e.Setenv("ATTACKER_AUTHORIZED_KEY", attacker.AuthorizedKey())
138			e.Setenv("SSH_KNOWN_HOSTS_FILE", filepath.Join(t.TempDir(), "known_hosts"))
139			e.Setenv("SSH_KNOWN_CONFIG_FILE", filepath.Join(t.TempDir(), "config"))
140
141			// This is used to set up test specific configuration and http endpoints
142			e.Setenv("SOFT_SERVE_TESTRUN", "1")
143
144			// This will disable the default lipgloss renderer colors
145			e.Setenv("SOFT_SERVE_NO_COLOR", "1")
146
147			// Soft Serve debug environment variables
148			for _, env := range []string{
149				"SOFT_SERVE_DEBUG",
150				"SOFT_SERVE_VERBOSE",
151			} {
152				if v, ok := os.LookupEnv(env); ok {
153					e.Setenv(env, v)
154				}
155			}
156
157			// TODO: test different configs
158			cfg := config.DefaultConfig()
159			cfg.DataPath = data
160			cfg.Name = serverName
161			cfg.InitialAdminKeys = []string{admin1.AuthorizedKey()}
162			cfg.SSH.ListenAddr = sshListen
163			cfg.SSH.PublicURL = "ssh://" + sshListen
164			cfg.Git.ListenAddr = gitListen
165			cfg.HTTP.ListenAddr = httpListen
166			cfg.HTTP.PublicURL = "http://" + httpListen
167			cfg.Stats.ListenAddr = statsListen
168			cfg.LFS.Enabled = true
169
170			// Parse os SOFT_SERVE environment variables
171			if err := cfg.ParseEnv(); err != nil {
172				return err
173			}
174
175			// Override the database data source if we're using postgres
176			// so we can create a temporary database for the tests.
177			if cfg.DB.Driver == "postgres" {
178				cleanup, err := setupPostgres(e.T(), cfg)
179				if err != nil {
180					return err
181				}
182				if cleanup != nil {
183					e.Defer(cleanup)
184				}
185			}
186
187			for _, env := range cfg.Environ() {
188				parts := strings.SplitN(env, "=", 2)
189				if len(parts) != 2 {
190					e.T().Fatal("invalid environment variable", env)
191				}
192				e.Setenv(parts[0], parts[1])
193			}
194
195			return nil
196		},
197	})
198}
199
200func cmdSoft(user string, keys ...ssh.Signer) func(ts *testscript.TestScript, neg bool, args []string) {
201	return func(ts *testscript.TestScript, neg bool, args []string) {
202		cli, err := ssh.Dial(
203			"tcp",
204			net.JoinHostPort("localhost", ts.Getenv("SSH_PORT")),
205			&ssh.ClientConfig{
206				User:            user,
207				Auth:            []ssh.AuthMethod{ssh.PublicKeys(keys...)},
208				HostKeyCallback: ssh.InsecureIgnoreHostKey(),
209			},
210		)
211		ts.Check(err)
212		defer cli.Close()
213
214		sess, err := cli.NewSession()
215		ts.Check(err)
216		defer sess.Close()
217
218		sess.Stdout = ts.Stdout()
219		sess.Stderr = ts.Stderr()
220
221		check(ts, sess.Run(strings.Join(args, " ")), neg)
222	}
223}
224
225// cmdKeylessSoft is like cmdSoft, but authenticates with zero public keys,
226// forcing keyboard-interactive auth -- the actual "no key offered at all"
227// path that allow-keyless gates. This is distinct from cmdSoft/cmdUsoft
228// style helpers, which always offer a real key (registered or not) and so
229// only ever exercise the anon-access path, not allow-keyless.
230//
231// A real ssh(1) binary can't reliably be used for this in a non-interactive
232// test harness: OpenSSH's client refuses to send a keyboard-interactive
233// request at all when stdin isn't a tty ("we did not send a packet, disable
234// method"), regardless of BatchMode. golang.org/x/crypto/ssh has no such
235// restriction, so it's used directly here instead of shelling out.
236func cmdKeylessSoft(ts *testscript.TestScript, neg bool, args []string) {
237	cli, err := ssh.Dial(
238		"tcp",
239		net.JoinHostPort("localhost", ts.Getenv("SSH_PORT")),
240		&ssh.ClientConfig{
241			User: "keyless",
242			Auth: []ssh.AuthMethod{
243				ssh.KeyboardInteractive(func(_, _ string, _ []string, _ []bool) ([]string, error) {
244					return nil, nil
245				}),
246			},
247			HostKeyCallback: ssh.InsecureIgnoreHostKey(),
248		},
249	)
250	ts.Check(err)
251	defer cli.Close()
252
253	sess, err := cli.NewSession()
254	ts.Check(err)
255	defer sess.Close()
256
257	sess.Stdout = ts.Stdout()
258	sess.Stderr = ts.Stderr()
259
260	check(ts, sess.Run(strings.Join(args, " ")), neg)
261}
262
263func cmdUI(key ssh.Signer) func(ts *testscript.TestScript, neg bool, args []string) {
264	return func(ts *testscript.TestScript, neg bool, args []string) {
265		if len(args) < 1 {
266			ts.Fatalf("usage: ui <quoted string input>")
267			return
268		}
269
270		cli, err := ssh.Dial(
271			"tcp",
272			net.JoinHostPort("localhost", ts.Getenv("SSH_PORT")),
273			&ssh.ClientConfig{
274				User:            "git",
275				Auth:            []ssh.AuthMethod{ssh.PublicKeys(key)},
276				HostKeyCallback: ssh.InsecureIgnoreHostKey(),
277			},
278		)
279		check(ts, err, neg)
280		defer cli.Close()
281
282		sess, err := cli.NewSession()
283		check(ts, err, neg)
284		defer sess.Close()
285
286		// XXX: this is a hack to make the UI tests work
287		// cmp command always complains about an extra newline
288		// in the output
289		defer ts.Stdout().Write([]byte("\n"))
290
291		sess.Stdout = ts.Stdout()
292		sess.Stderr = ts.Stderr()
293
294		stdin, err := sess.StdinPipe()
295		check(ts, err, neg)
296
297		err = sess.RequestPty("dumb", 40, 80, ssh.TerminalModes{})
298		check(ts, err, neg)
299		check(ts, sess.Start(""), neg)
300
301		in, err := strconv.Unquote(args[0])
302		check(ts, err, neg)
303		reader := strings.NewReader(in)
304		go func() {
305			defer stdin.Close()
306			for {
307				r, _, err := reader.ReadRune()
308				if err == io.EOF {
309					break
310				}
311				check(ts, err, neg)
312				_, _ = io.WriteString(stdin, string(r))
313
314				// Wait for the UI to process the input
315				time.Sleep(100 * time.Millisecond)
316			}
317		}()
318
319		check(ts, sess.Wait(), neg)
320	}
321}
322
323func cmdDos2Unix(ts *testscript.TestScript, neg bool, args []string) {
324	if neg {
325		ts.Fatalf("unsupported: ! dos2unix")
326	}
327	if len(args) < 1 {
328		ts.Fatalf("usage: dos2unix paths...")
329	}
330	for _, arg := range args {
331		filename := ts.MkAbs(arg)
332		data, err := os.ReadFile(filename)
333		if err != nil {
334			ts.Fatalf("%s: %v", filename, err)
335		}
336
337		// Replace all '\r\n' with '\n'.
338		data = bytes.ReplaceAll(data, []byte{'\r', '\n'}, []byte{'\n'})
339
340		if err := os.WriteFile(filename, data, 0o644); err != nil {
341			ts.Fatalf("%s: %v", filename, err)
342		}
343	}
344}
345
346var sshConfig = `
347Host *
348  UserKnownHostsFile %q
349  StrictHostKeyChecking no
350  IdentityAgent none
351  IdentitiesOnly yes
352  ServerAliveInterval 60
353`
354
355func cmdGit(key string) func(ts *testscript.TestScript, neg bool, args []string) {
356	return func(ts *testscript.TestScript, neg bool, args []string) {
357		ts.Check(os.WriteFile(
358			ts.Getenv("SSH_KNOWN_CONFIG_FILE"),
359			[]byte(fmt.Sprintf(sshConfig, ts.Getenv("SSH_KNOWN_HOSTS_FILE"))),
360			0o600,
361		))
362		sshArgs := []string{
363			"-F", filepath.ToSlash(ts.Getenv("SSH_KNOWN_CONFIG_FILE")),
364			"-i", filepath.ToSlash(key),
365		}
366		ts.Setenv(
367			"GIT_SSH_COMMAND",
368			strings.Join(append([]string{"ssh"}, sshArgs...), " "),
369		)
370		// Disable git prompting for credentials.
371		ts.Setenv("GIT_TERMINAL_PROMPT", "0")
372		args = append([]string{
373			"-c", "user.email=john@example.com",
374			"-c", "user.name=John Doe",
375		}, args...)
376		check(ts, ts.Exec("git", args...), neg)
377	}
378}
379
380func cmdMkfile(ts *testscript.TestScript, neg bool, args []string) {
381	if len(args) < 2 {
382		ts.Fatalf("usage: mkfile path content")
383	}
384	check(ts, os.WriteFile(
385		ts.MkAbs(args[0]),
386		[]byte(strings.Join(args[1:], " ")),
387		0o644,
388	), neg)
389}
390
391func check(ts *testscript.TestScript, err error, neg bool) {
392	if neg && err == nil {
393		ts.Fatalf("expected error, got nil")
394	}
395	if !neg {
396		ts.Check(err)
397	}
398}
399
400func cmdReadfile(ts *testscript.TestScript, neg bool, args []string) {
401	ts.Stdout().Write([]byte(ts.ReadFile(args[0])))
402}
403
404func cmdEnvfile(ts *testscript.TestScript, neg bool, args []string) {
405	if len(args) < 1 {
406		ts.Fatalf("usage: envfile key=file...")
407	}
408
409	for _, arg := range args {
410		parts := strings.SplitN(arg, "=", 2)
411		if len(parts) != 2 {
412			ts.Fatalf("usage: envfile key=file...")
413		}
414		key := parts[0]
415		file := parts[1]
416		ts.Setenv(key, strings.TrimSpace(ts.ReadFile(file)))
417	}
418}
419
420func cmdNewWebhook(ts *testscript.TestScript, neg bool, args []string) {
421	type webhookSite struct {
422		UUID string `json:"uuid"`
423	}
424
425	if len(args) != 1 {
426		ts.Fatalf("usage: new-webhook <env-name>")
427	}
428
429	const whSite = "https://webhook.site"
430	req, err := http.NewRequest(http.MethodPost, whSite+"/token", nil) //nolint:noctx
431	check(ts, err, neg)
432
433	resp, err := http.DefaultClient.Do(req)
434	check(ts, err, neg)
435
436	defer resp.Body.Close()
437	var site webhookSite
438	check(ts, json.NewDecoder(resp.Body).Decode(&site), neg)
439
440	ts.Setenv(args[0], whSite+"/"+site.UUID)
441}
442
443func cmdCurl(ts *testscript.TestScript, neg bool, args []string) {
444	var verbose bool
445	var headers []string
446	var data string
447	method := http.MethodGet
448
449	cmd := &cobra.Command{
450		Use:  "curl",
451		Args: cobra.MinimumNArgs(1),
452		RunE: func(cmd *cobra.Command, args []string) error {
453			url, err := url.Parse(args[0])
454			if err != nil {
455				return err
456			}
457
458			req, err := http.NewRequest(method, url.String(), nil) //nolint:noctx
459			if err != nil {
460				return err
461			}
462
463			if data != "" {
464				req.Body = io.NopCloser(strings.NewReader(data))
465			}
466
467			if verbose {
468				fmt.Fprintf(cmd.ErrOrStderr(), "< %s %s\n", req.Method, url.String())
469			}
470
471			for _, header := range headers {
472				parts := strings.SplitN(header, ":", 2)
473				if len(parts) != 2 {
474					return fmt.Errorf("invalid header: %s", header)
475				}
476				req.Header.Add(strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1]))
477			}
478
479			if userInfo := url.User; userInfo != nil {
480				password, _ := userInfo.Password()
481				req.SetBasicAuth(userInfo.Username(), password)
482			}
483
484			if verbose {
485				for key, values := range req.Header {
486					for _, value := range values {
487						fmt.Fprintf(cmd.ErrOrStderr(), "< %s: %s\n", key, value)
488					}
489				}
490			}
491
492			resp, err := http.DefaultClient.Do(req)
493			if err != nil {
494				return err
495			}
496
497			if verbose {
498				fmt.Fprintf(ts.Stderr(), "> %s\n", resp.Status)
499				for key, values := range resp.Header {
500					for _, value := range values {
501						fmt.Fprintf(cmd.ErrOrStderr(), "> %s: %s\n", key, value)
502					}
503				}
504			}
505
506			defer resp.Body.Close()
507			buf, err := io.ReadAll(resp.Body)
508			if err != nil {
509				return err
510			}
511
512			cmd.Print(string(buf))
513
514			return nil
515		},
516	}
517
518	cmd.SetArgs(args)
519	cmd.SetOut(ts.Stdout())
520	cmd.SetErr(ts.Stderr())
521
522	cmd.Flags().BoolVarP(&verbose, "verbose", "v", verbose, "verbose")
523	cmd.Flags().StringArrayVarP(&headers, "header", "H", nil, "HTTP header")
524	cmd.Flags().StringVarP(&method, "request", "X", method, "HTTP method")
525	cmd.Flags().StringVarP(&data, "data", "d", data, "HTTP data")
526
527	check(ts, cmd.Execute(), neg)
528}
529
530func cmdEnsureServerRunning(ts *testscript.TestScript, neg bool, args []string) {
531	if len(args) < 1 {
532		ts.Fatalf("Must supply a TCP port of one of the services to connect to. " +
533			"These are set as env vars as they are randomized. " +
534			"Example usage: \"cmdensureserverrunning SSH_PORT\"\n" +
535			"Valid values for the env var: SSH_PORT|HTTP_PORT|GIT_PORT|STATS_PORT")
536	}
537
538	port := ts.Getenv(args[0])
539
540	// verify that the server is up
541	addr := net.JoinHostPort("localhost", port)
542	deadline := time.Now().Add(30 * time.Second)
543	for {
544		conn, _ := net.DialTimeout( //nolint:noctx
545			"tcp",
546			addr,
547			time.Second,
548		)
549		if conn != nil {
550			ts.Logf("Server is running on port: %s", port)
551			conn.Close()
552			break
553		}
554		if time.Now().After(deadline) {
555			ts.Fatalf("server on port %s did not start within 30 seconds", port)
556		}
557		time.Sleep(10 * time.Millisecond)
558	}
559}
560
561func cmdEnsureServerNotRunning(ts *testscript.TestScript, neg bool, args []string) {
562	if len(args) < 1 {
563		ts.Fatalf("Must supply a TCP port of one of the services to connect to. " +
564			"These are set as env vars as they are randomized. " +
565			"Example usage: \"cmdensureservernotrunning SSH_PORT\"\n" +
566			"Valid values for the env var: SSH_PORT|HTTP_PORT|GIT_PORT|STATS_PORT")
567	}
568
569	port := ts.Getenv(args[0])
570
571	// verify that the server is not up
572	addr := net.JoinHostPort("localhost", port)
573	conn, _ := net.DialTimeout( //nolint:noctx
574		"tcp",
575		addr,
576		time.Second,
577	)
578	if conn != nil {
579		ts.Fatalf("server is running on port %s while it should not be running", port)
580		conn.Close()
581	}
582}
583
584func cmdStopserver(ts *testscript.TestScript, neg bool, args []string) {
585	// stop the server
586	resp, err := http.DefaultClient.Head(fmt.Sprintf("%s/__stop", ts.Getenv("SOFT_SERVE_HTTP_PUBLIC_URL"))) //nolint:noctx
587	check(ts, err, neg)
588	resp.Body.Close()
589	time.Sleep(time.Second * 2) // Allow some time for the server to stop
590}
591
592func setupPostgres(t testscript.T, cfg *config.Config) (func(), error) {
593	// Indicates postgres
594	// Create a disposable database
595	rnd := rand.New(rand.NewSource(time.Now().UnixNano()))
596	dbName := fmt.Sprintf("softserve_test_%d", rnd.Int63())
597	dbDsn := cfg.DB.DataSource
598	if dbDsn == "" {
599		cfg.DB.DataSource = "postgres://postgres@localhost:5432/postgres?sslmode=disable"
600	}
601
602	dbUrl, err := url.Parse(cfg.DB.DataSource)
603	if err != nil {
604		return nil, err
605	}
606
607	scheme := dbUrl.Scheme
608	if scheme == "" {
609		scheme = "postgres"
610	}
611
612	host := dbUrl.Hostname()
613	if host == "" {
614		host = "localhost"
615	}
616
617	connInfo := fmt.Sprintf("host=%s sslmode=disable", host)
618	username := dbUrl.User.Username()
619	if username != "" {
620		connInfo += fmt.Sprintf(" user=%s", username)
621		password, ok := dbUrl.User.Password()
622		if ok {
623			username = fmt.Sprintf("%s:%s", username, password)
624			connInfo += fmt.Sprintf(" password=%s", password)
625		}
626		username = fmt.Sprintf("%s@", username)
627	} else {
628		connInfo += " user=postgres"
629		username = "postgres@"
630	}
631
632	port := dbUrl.Port()
633	if port != "" {
634		connInfo += fmt.Sprintf(" port=%s", port)
635		port = fmt.Sprintf(":%s", port)
636	}
637
638	cfg.DB.DataSource = fmt.Sprintf("%s://%s%s%s/%s?sslmode=disable",
639		scheme,
640		username,
641		host,
642		port,
643		dbName,
644	)
645
646	// Create the database
647	dbx, err := db.Open(context.TODO(), cfg.DB.Driver, connInfo)
648	if err != nil {
649		return nil, err
650	}
651
652	if _, err := dbx.ExecContext(context.TODO(), "CREATE DATABASE "+dbName); err != nil {
653		return nil, err
654	}
655
656	return func() {
657		dbx, err := db.Open(context.TODO(), cfg.DB.Driver, connInfo)
658		if err != nil {
659			t.Fatal("failed to open database", dbName, err)
660		}
661
662		if _, err := dbx.ExecContext(context.TODO(), "DROP DATABASE "+dbName); err != nil {
663			t.Fatal("failed to drop database", dbName, err)
664		}
665	}, nil
666}
667
668type maliciousSigner struct {
669	publicKey ssh.PublicKey
670}
671
672var _ ssh.Signer = (*maliciousSigner)(nil)
673
674// PublicKey implements ssh.Signer.
675func (m *maliciousSigner) PublicKey() ssh.PublicKey {
676	return m.publicKey
677}
678
679// Sign implements ssh.Signer.
680func (m *maliciousSigner) Sign(rand io.Reader, data []byte) (*ssh.Signature, error) {
681	// The attacker doesn't know how to sign the data without a private key.
682	return &ssh.Signature{}, nil
683}