diff --git a/Makefile b/Makefile index 21e9b844e..81d228c87 100644 --- a/Makefile +++ b/Makefile @@ -262,7 +262,7 @@ arm-qemu-smoke-test: $(RUNTIME_BIN) load-arm-qemu simple-tests: unit-tests # Compatibility target. .PHONY: simple-tests -portforward-tests: load-basic_redis $(RUNTIME_BIN) +portforward-tests: load-basic_redis load-basic_nginx $(RUNTIME_BIN) @$(call install_runtime,$(RUNTIME),--network=sandbox) @$(call sudo,test/root:portforward_test,--runtime=$(RUNTIME) -test.v) @$(call install_runtime,$(RUNTIME),--network=host) diff --git a/images/basic/nginx/Dockerfile b/images/basic/nginx/Dockerfile index af2e62526..ad2b6eee5 100644 --- a/images/basic/nginx/Dockerfile +++ b/images/basic/nginx/Dockerfile @@ -1 +1,3 @@ FROM nginx:1.17.9 + +RUN apt-get update && apt-get install -y net-tools \ No newline at end of file diff --git a/runsc/cmd/portforward.go b/runsc/cmd/portforward.go index c2f372177..0c1c14644 100644 --- a/runsc/cmd/portforward.go +++ b/runsc/cmd/portforward.go @@ -58,8 +58,9 @@ func (*PortForward) Synopsis() string { func (*PortForward) Usage() string { return `port-forward CONTAINER_ID [LOCAL_PORT:]REMOTE_PORT - port forward to gvisor container. -Open a local port and forward connections to another port inside the specified -container. +Port forwarding has two modes. Local mode opens a local port and forwards +connections to another port inside the specified container. Stream mode +forwards a single connection on a UDS to the specified port in the container. EXAMPLES: @@ -67,11 +68,20 @@ The following will forward connections on local port 8080 to port 80 in the container named 'nginx': # runsc port-forward nginx 8080:80 + +The following will forward a single new connection on the unix domain socket at +/tmp/pipe to port 80 in the container named 'nginx': + + # runsc port-forward --stream /tmp/pipe nginx 80 + +OPTIONS: ` } // SetFlags implements subcommands.Command.SetFlags. -func (p *PortForward) SetFlags(f *flag.FlagSet) {} +func (p *PortForward) SetFlags(f *flag.FlagSet) { + f.StringVar(&p.stream, "stream", "", "Stream mode - a Unix doman socket") +} // Execute implements subcommands.Command.Execute. func (p *PortForward) Execute(ctx context.Context, f *flag.FlagSet, args ...any) subcommands.ExitStatus { @@ -90,6 +100,13 @@ func (p *PortForward) Execute(ctx context.Context, f *flag.FlagSet, args ...any) util.Fatalf("loading container: %v", err) } + if p.stream != "" { + if err := p.doStream(ctx, portStr, c); err != nil { + util.Fatalf("doStream: %v", err) + } + return subcommands.ExitSuccess + } + // Allow forwarding to a local port. ports := strings.Split(portStr, ":") if len(ports) != 2 { @@ -194,6 +211,35 @@ func localForward(ctx context.Context, c *container.Container, localPort int, co } } +// doStream does the stream version of the port-forward command. +func (p *PortForward) doStream(ctx context.Context, port string, c *container.Container) error { + var err error + p.portNum, err = strconv.Atoi(port) + if err != nil { + return fmt.Errorf("invalid port string %q: %v", port, err) + } + + if p.portNum < 0 || p.portNum > math.MaxUint16 { + return fmt.Errorf("invalid port %d: %v", p.portNum, err) + } + + f, err := openStream(p.stream) + if err != nil { + return fmt.Errorf("opening uds stream: %v", err) + } + defer f.Close() + + if err := c.PortForward(&boot.PortForwardOpts{ + Port: uint16(p.portNum), + ContainerID: c.ID, + FilePayload: urpc.FilePayload{Files: []*os.File{f}}, + }); err != nil { + return fmt.Errorf("PortForward: %v", err) + } + + return nil +} + // portCopy creates a UDS and begins copying data to and from the local // connection. func portCopy(ctx context.Context, c *container.Container, localConn net.Conn, port uint16) error { diff --git a/test/root/portforward_test.go b/test/root/portforward_test.go index fc8a6fb86..395e556e7 100644 --- a/test/root/portforward_test.go +++ b/test/root/portforward_test.go @@ -21,11 +21,15 @@ import ( "bytes" "context" "fmt" + "io" "net" "os" "os/exec" + "path" + "regexp" "strings" "testing" + "time" "github.com/syndtr/gocapability/capability" "golang.org/x/sync/errgroup" @@ -36,7 +40,7 @@ import ( "gvisor.dev/gvisor/runsc/specutils" ) -func TestPortForward(t *testing.T) { +func TestPortForwardLocalMode(t *testing.T) { ctx := context.Background() server := dockerutil.MakeContainer(ctx, t) defer server.CleanUp(ctx) @@ -61,9 +65,10 @@ func TestPortForward(t *testing.T) { var g errgroup.Group g.Go(func() error { - pf.Wait() - if pf.Error() != nil { - return fmt.Errorf("portforward command: err: %v out: %s", pf.Error(), pf.Output()) + // To end this test, we kill the portforward process, which will result in a "signal: killed" + // error. Just ignore this error. + if err := pf.Wait(); err != nil && !strings.Contains(err.Error(), "signal: killed") { + return fmt.Errorf("portforward command: err: %v process error: %v out: %s", err, pf.Error(), pf.Output()) } return nil }) @@ -94,6 +99,88 @@ func TestPortForward(t *testing.T) { } } +func TestPortForwardStreamMode(t *testing.T) { + ctx := context.Background() + sockAddrDir, err := os.MkdirTemp("", "temp-") + if err != nil { + t.Fatalf("failed to create temp dir: %v", err) + } + defer os.RemoveAll(sockAddrDir) + sockAddr := path.Join(sockAddrDir, "echo.sock") + + server := dockerutil.MakeContainer(ctx, t) + defer server.CleanUp(ctx) + + nginxPort := 80 + if err := server.Spawn(ctx, dockerutil.RunOpts{ + Image: "basic/nginx", + }); err != nil { + t.Fatalf("failed to create nginx server: %v", err) + } + + // This is a bit crude, but we need to make sure the server is up without exposing a port to the + // host. When the server container boots, the nginx process should run first. If we run nginx + // again, it will fail to bind to port 80. Run exec calls until we get that failure. + serverUpChan := make(chan struct{}, 1) + var upOut string + var upErr error + reg := regexp.MustCompile(`0\.0\.0\.0:80[\s]*0\.0\.0\.0:\*[\s]*LISTEN`) + go func() { + for { + time.Sleep(time.Millisecond * 500) + upOut, upErr = server.Exec(ctx, dockerutil.ExecOpts{}, []string{"netstat", "-l"}...) + if reg.MatchString(upOut) { + serverUpChan <- struct{}{} + return + } + } + }() + + // If the server isn't up after 10 seconds, there is probably something wrong. + select { + case <-serverUpChan: + break + case <-time.After(time.Second * 30): + t.Fatalf("could not verify server is up: err: %v out: %s", upErr, upOut) + } + + socket, err := net.Listen("unix", sockAddr) + if err != nil { + t.Fatalf("failed to listen: %v", err) + } + defer socket.Close() + + pf, err := newPortForwardStreamProcess(ctx, server, sockAddr, nginxPort) + if err != nil { + t.Fatalf("failed to create port forward process: %v", err) + } + + if err := pf.Wait(); err != nil { + t.Fatalf("failed to wait: %v out: %s", err, pf.Output()) + } + + conn, err := socket.Accept() + if err != nil { + t.Fatalf("failed to accept: %v", err) + } + defer conn.Close() + + const getMsg = "GET / HTTP/1.0\r\n\r\n" + if n, err := io.Copy(conn, bytes.NewBufferString(getMsg)); err != nil { + t.Fatalf("failed to copy: %v n: %d", err, n) + } + + buf, err := io.ReadAll(conn) + if err != nil { + t.Fatalf("failed to read: %v out: %s", err, string(buf)) + } + + const want = "Thank you for using nginx." + if !strings.Contains(string(buf), want) { + t.Fatalf("could not find %q in output: %s", want, string(buf)) + } +} + func getUnusedPort() (int, error) { l, err := net.Listen("tcp", ":0") if err != nil { @@ -114,6 +201,19 @@ func newPortForwardProcess(ctx context.Context, c *dockerutil.Container, localPo return nil, err } args := []string{"-root", rootDir, "port-forward", c.ID(), fmt.Sprintf("%d:%d", localPort, containerPort)} + return startPortForwardPorcess(ctx, args) +} + +func newPortForwardStreamProcess(ctx context.Context, c *dockerutil.Container, uds string, containerPort int) (*portForwardProcess, error) { + rootDir, err := c.RootDirectory() + if err != nil { + return nil, err + } + args := []string{"-root", rootDir, "-alsologtostderr", "port-forward", "-stream", uds, c.ID(), fmt.Sprintf("%d", containerPort)} + return startPortForwardPorcess(ctx, args) +} + +func startPortForwardPorcess(ctx context.Context, args []string) (*portForwardProcess, error) { cmd := exec.CommandContext(ctx, specutils.ExePath, args...) ret := &portForwardProcess{cmd: cmd} ret.cmd.Stdout = &ret.buf @@ -128,7 +228,7 @@ func (p *portForwardProcess) Close() error { return p.cmd.Wait() } -func (p *portForwardProcess) Wait() { p.cmd.Wait() } +func (p *portForwardProcess) Wait() error { return p.cmd.Wait() } func (p *portForwardProcess) Kill() { p.cmd.Process.Kill() }