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", "-l", "-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