FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

GitHub Viewer

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

Back | FazBrowse Home | New Git URL