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}