[ Web Proxy ]
URL:
Viewing: https://raw.githubusercontent.com/CeasarJackson/sshcode/master/sshcode_test.go [Back]  [Original]

package main

import (
	"context"
	"fmt"
	"io"
	"net"
	"net/http"
	"os/exec"
	"path/filepath"
	"strconv"
	"sync"
	"testing"
	"time"

	"github.com/stretchr/testify/require"
	"go.coder.com/retry"
	"golang.org/x/crypto/ssh"
)

func TestSSHCode(t *testing.T) {
	sshPort, err := randomPort()
	require.NoError(t, err)

	// start up our jank ssh server
	defer trassh(t, sshPort).Close()

	localPort := randomPortExclude(t, sshPort)
	require.NotEmpty(t, localPort)

	remotePort := randomPortExclude(t, sshPort, localPort)
	require.NotEmpty(t, remotePort)

	var wg sync.WaitGroup
	wg.Add(1)
	go func() {
		defer wg.Done()
		err := sshCode("foo@127.0.0.1", "", options{
			sshFlags:   testSSHArgs(sshPort),
			bindAddr:   net.JoinHostPort("127.0.0.1", localPort),
			remotePort: remotePort,
			noOpen:     true,
		})
		require.NoError(t, err)
	}()

	waitForSSHCode(t, localPort, time.Second*30)
	waitForSSHCode(t, remotePort, time.Second*30)

	// Typically we'd do an os.Stat call here but the os package doesn't expand '~'
	out, err := exec.Command("sh", "-c", "stat "+codeServerPath).CombinedOutput()
	require.NoError(t, err, "%s", out)

	out, err = exec.Command("pkill", filepath.Base(codeServerPath)).CombinedOutput()
	require.NoError(t, err, "%s", out)

	wg.Wait()
}

// trassh is an incomplete, local, insecure ssh server
// used for the purpose of testing the implementation without
// requiring the user to have their own remote server.
func trassh(t *testing.T, port string) io.Closer {
	private, err := ssh.ParsePrivateKey([]byte(fakeRSAKey))
	require.NoError(t, err)

	conf := &ssh.ServerConfig{
		NoClientAuth: true,
	}

	conf.AddHostKey(private)

	listener, err := net.Listen("tcp", net.JoinHostPort("127.0.0.1", port))
	require.NoError(t, err)

	go func() {
		for {
			func() {
				conn, err := listener.Accept()
				if err != nil {
					return
				}
				defer conn.Close()

				sshConn, chans, reqs, err := ssh.NewServerConn(conn, conf)
				require.NoError(t, err)

				go ssh.DiscardRequests(reqs)

				for c := range chans {
					switch c.ChannelType() {
					case "direct-tcpip":
						var req directTCPIPReq

						err := ssh.Unmarshal(c.ExtraData(), &req)
						if err != nil {
							t.Logf("failed to unmarshal tcpip data: %v", err)
							continue
						}

						ch, _, err := c.Accept()
						if err != nil {
							c.Reject(ssh.ConnectionFailed, fmt.Sprintf("unable to accept channel: %v", err))
							continue
						}

						go handleDirectTCPIP(ch, &req, t)
					case "session":
						ch, inReqs, err := c.Accept()
						if err != nil {
							c.Reject(ssh.ConnectionFailed, fmt.Sprintf("unable to accept channel: %v", err))
							continue
						}

						go handleSession(ch, inReqs, t)
					default:
						t.Logf("unsupported session type: %v\n", c.ChannelType())
						c.Reject(ssh.UnknownChannelType, "unknown channel type")
					}
				}

				sshConn.Wait()
			}()
		}
	}()
	return listener
}

func handleDirectTCPIP(ch ssh.Channel, req *directTCPIPReq, t *testing.T) {
	defer ch.Close()

	dstAddr := net.JoinHostPort(req.Host, strconv.Itoa(int(req.Port)))

	conn, err := net.Dial("tcp", dstAddr)
	if err != nil {
		return
	}
	defer conn.Close()

	var wg sync.WaitGroup

	wg.Add(1)
	go func() {
		defer wg.Done()
		defer ch.Close()

		io.Copy(ch, conn)
	}()

	wg.Add(1)
	go func() {
		defer wg.Done()
		defer conn.Close()

		io.Copy(conn, ch)
	}()
	wg.Wait()
}

// execReq describes an exec payload.
type execReq struct {
	Command string
}

// directTCPIPReq describes the extra data sent in a
// direct-tcpip request containing the host/port for the ssh server.
type directTCPIPReq struct {
	Host string
	Port uint32

	Orig     string
	OrigPort uint32
}

// exitStatus describes an 'exit-status' message
// returned after a request.
type exitStatus struct {
	Status uint32
}

func handleSession(ch ssh.Channel, in 

Web Proxy Viewer  |  New URL  |  Original Page