From 01e5a009e056793db41d77348faf0ef2d9e92f4f Mon Sep 17 00:00:00 2001 From: Boshi Lian Date: Mon, 28 Nov 2022 05:56:12 -0800 Subject: [PATCH] add banner opt back (#106) * add banner opt back * happy lint --- cmd/sshpiperd/daemon.go | 17 ++++++ cmd/sshpiperd/main.go | 12 ++++ e2e/banner_test.go | 118 ++++++++++++++++++++++++++++++++++++++++ e2e/main_test.go | 16 ++++-- 4 files changed, 158 insertions(+), 5 deletions(-) create mode 100644 e2e/banner_test.go diff --git a/cmd/sshpiperd/daemon.go b/cmd/sshpiperd/daemon.go index 2e8ec486..93480228 100644 --- a/cmd/sshpiperd/daemon.go +++ b/cmd/sshpiperd/daemon.go @@ -56,6 +56,23 @@ func newDaemon(ctx *cli.Context) (*daemon, error) { return nil, fmt.Errorf("failed to listen for connection: %v", err) } + bannertext := ctx.String("banner-text") + bannerfile := ctx.String("banner-file") + + if bannertext != "" || bannerfile != "" { + config.BannerCallback = func(_ ssh.ConnMetadata, _ ssh.ChallengeContext) string { + if bannerfile != "" { + text, err := os.ReadFile(bannerfile) + if err != nil { + log.Warnf("cannot read banner file %v: %v", bannerfile, err) + } else { + return string(text) + } + } + return bannertext + } + } + return &daemon{ config: config, lis: lis, diff --git a/cmd/sshpiperd/main.go b/cmd/sshpiperd/main.go index 747b4fe4..74ecc343 100644 --- a/cmd/sshpiperd/main.go +++ b/cmd/sshpiperd/main.go @@ -115,6 +115,18 @@ func main() { Usage: "create typescript format screen recording and save into the directory see https://linux.die.net/man/1/script", EnvVars: []string{"SSHPIPERD_TYPESCRIPT_LOG_DIR"}, }, + &cli.StringFlag{ + Name: "banner-text", + Value: "", + Usage: "display a banner before authentication, would be ignored if banner file was set", + EnvVars: []string{"SSHPIPERD_BANNERTEXT"}, + }, + &cli.StringFlag{ + Name: "banner-file", + Value: "", + Usage: "display a banner from file before authentication", + EnvVars: []string{"SSHPIPERD_BANNERFILE"}, + }, }, Action: func(ctx *cli.Context) error { level, err := log.ParseLevel(ctx.String("log-level")) diff --git a/e2e/banner_test.go b/e2e/banner_test.go new file mode 100644 index 00000000..63263299 --- /dev/null +++ b/e2e/banner_test.go @@ -0,0 +1,118 @@ +package e2e_test + +import ( + "os" + "testing" + + "github.com/google/uuid" +) + +func TestBanner(t *testing.T) { + + t.Run("args", func(t *testing.T) { + piperaddr, piperport := nextAvailablePiperAddress() + randtext := uuid.New().String() + + piper, _, _, err := runCmd("/sshpiperd/sshpiperd", + "--banner-text", + randtext, + "-p", + piperport, + "/sshpiperd/plugins/fixed", + "--target", + "host-password:2222", + ) + + if err != nil { + t.Errorf("failed to run sshpiperd: %v", err) + } + + defer killCmd(piper) + + waitForEndpointReady(piperaddr) + + c, _, stdout, err := runCmd( + "ssh", + "-v", + "-o", + "StrictHostKeyChecking=no", + "-o", + "UserKnownHostsFile=/dev/null", + "-p", + piperport, + "-l", + "user", + "127.0.0.1", + ) + + if err != nil { + t.Errorf("failed to ssh to piper, %v", err) + } + + defer killCmd(c) + + waitForStdoutContains(stdout, randtext, func(_ string) { + }) + }) + + t.Run("file", func(t *testing.T) { + + piperaddr, piperport := nextAvailablePiperAddress() + randtext := uuid.New().String() + + bannerfile, err := os.CreateTemp("", "banner") + if err != nil { + t.Errorf("failed to create temp file: %v", err) + } + defer os.Remove(bannerfile.Name()) + + if _, err := bannerfile.WriteString(randtext); err != nil { + t.Errorf("failed to write to temp file: %v", err) + } + + if err := bannerfile.Close(); err != nil { + t.Errorf("failed to close temp file: %v", err) + } + + piper, _, _, err := runCmd("/sshpiperd/sshpiperd", + "--banner-file", + bannerfile.Name(), + "-p", + piperport, + "/sshpiperd/plugins/fixed", + "--target", + "host-password:2222", + ) + + if err != nil { + t.Errorf("failed to run sshpiperd: %v", err) + } + + defer killCmd(piper) + + waitForEndpointReady(piperaddr) + + c, _, stdout, err := runCmd( + "ssh", + "-v", + "-o", + "StrictHostKeyChecking=no", + "-o", + "UserKnownHostsFile=/dev/null", + "-p", + piperport, + "-l", + "user", + "127.0.0.1", + ) + + if err != nil { + t.Errorf("failed to ssh to piper, %v", err) + } + + defer killCmd(c) + + waitForStdoutContains(stdout, randtext, func(_ string) { + }) + }) +} diff --git a/e2e/main_test.go b/e2e/main_test.go index f161236d..d185bb73 100644 --- a/e2e/main_test.go +++ b/e2e/main_test.go @@ -113,21 +113,20 @@ func runCmdAndWait(cmd string, args ...string) error { return c.Wait() } -func enterPassword(stdin io.Writer, stdout io.Reader, password string) { +func waitForStdoutContains(stdout io.Reader, text string, cb func(string)) { st := time.Now() for { scanner := bufio.NewScanner(stdout) for scanner.Scan() { line := scanner.Text() - if strings.Contains(line, "'s password") { - _, _ = stdin.Write([]byte(fmt.Sprintf("%v\n", password))) - log.Printf("got password prompt, sending password") + if strings.Contains(line, text) { + cb(line) return } } if time.Since(st) > waitTimeout { - log.Panic("timeout waiting for password prompt") + log.Panicf("timeout waiting for [%s] from prompt", text) return } @@ -135,6 +134,13 @@ func enterPassword(stdin io.Writer, stdout io.Reader, password string) { } } +func enterPassword(stdin io.Writer, stdout io.Reader, password string) { + waitForStdoutContains(stdout, "'s password", func(_ string) { + _, _ = stdin.Write([]byte(fmt.Sprintf("%v\n", password))) + log.Printf("got password prompt, sending password") + }) +} + func checkSharedFileContent(t *testing.T, targetfie string, expected string) { f, err := os.Open(fmt.Sprintf("/shared/%v", targetfie)) if err != nil {