mirror of
https://github.com/netbirdio/gvisor.git
synced 2026-05-22 17:12:49 -07:00
Implement netstack port forward
Implement portforward support methods for sandboxes using netstack for their network stack. PiperOrigin-RevId: 503298590
This commit is contained in:
committed by
gVisor bot
parent
0e5d0cc13a
commit
373a440374
@@ -8,6 +8,8 @@ go_library(
|
||||
"portforward.go",
|
||||
"portforward_fd_rw.go",
|
||||
"portforward_hostinet.go",
|
||||
"portforward_netstack.go",
|
||||
"portforward_test_util.go",
|
||||
],
|
||||
deps = [
|
||||
"//pkg/cleanup",
|
||||
@@ -17,6 +19,11 @@ go_library(
|
||||
"//pkg/fdnotifier",
|
||||
"//pkg/log",
|
||||
"//pkg/sentry/vfs",
|
||||
"//pkg/sync",
|
||||
"//pkg/tcpip",
|
||||
"//pkg/tcpip/network/ipv4",
|
||||
"//pkg/tcpip/stack",
|
||||
"//pkg/tcpip/transport/tcp",
|
||||
"//pkg/usermem",
|
||||
"//pkg/waiter",
|
||||
"@org_golang_x_sys//unix:go_default_library",
|
||||
@@ -28,18 +35,22 @@ go_test(
|
||||
srcs = [
|
||||
"portforward_fd_rw_test.go",
|
||||
"portforward_hostinet_test.go",
|
||||
"portforward_netstack_test.go",
|
||||
],
|
||||
library = ":portforward",
|
||||
tags = [
|
||||
"manual",
|
||||
"requires-net:ipv4",
|
||||
"requires-net:loopback",
|
||||
],
|
||||
deps = [
|
||||
"//pkg/abi/linux",
|
||||
"//pkg/cleanup",
|
||||
"//pkg/context",
|
||||
"//pkg/errors/linuxerr",
|
||||
"//pkg/sentry/contexttest",
|
||||
"//pkg/sentry/vfs",
|
||||
"//pkg/tcpip",
|
||||
"//pkg/usermem",
|
||||
"//pkg/waiter",
|
||||
"@org_golang_x_sync//errgroup:go_default_library",
|
||||
|
||||
@@ -15,25 +15,17 @@
|
||||
package portforward
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"reflect"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"golang.org/x/sync/errgroup"
|
||||
"gvisor.dev/gvisor/pkg/context"
|
||||
"gvisor.dev/gvisor/pkg/errors/linuxerr"
|
||||
"gvisor.dev/gvisor/pkg/sentry/contexttest"
|
||||
"gvisor.dev/gvisor/pkg/sentry/vfs"
|
||||
"gvisor.dev/gvisor/pkg/usermem"
|
||||
"gvisor.dev/gvisor/pkg/waiter"
|
||||
)
|
||||
|
||||
func TestLocalHostSocket(t *testing.T) {
|
||||
|
||||
clientData := append(
|
||||
[]byte("do what must be done\n"),
|
||||
[]byte("do not hesitate\n")...,
|
||||
@@ -51,7 +43,6 @@ func TestLocalHostSocket(t *testing.T) {
|
||||
defer l.Close()
|
||||
|
||||
port := l.Addr().(*net.TCPAddr).Port
|
||||
|
||||
var g errgroup.Group
|
||||
|
||||
g.Go(func() error {
|
||||
@@ -121,116 +112,71 @@ func TestLocalHostSocket(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func newMockSocketPair() (*mockEndpoint, *mockEndpoint) {
|
||||
client := &mockEndpoint{}
|
||||
server := &mockEndpoint{other: client}
|
||||
client.other = server
|
||||
return client, server
|
||||
type netConnMockEndpoint struct {
|
||||
conn net.Conn
|
||||
}
|
||||
|
||||
type mockEndpoint struct {
|
||||
vfs.FileDescriptionDefaultImpl
|
||||
vfs.NoLockFD
|
||||
vfs.DentryMetadataFileDescriptionImpl
|
||||
other *mockEndpoint
|
||||
readBuf bytes.Buffer
|
||||
mu sync.Mutex
|
||||
released bool
|
||||
queue waiter.Queue
|
||||
// read implements portforwarderTestHarness.read.
|
||||
func (nc *netConnMockEndpoint) read(n int) ([]byte, error) {
|
||||
buf := make([]byte, n)
|
||||
n, err := nc.conn.Read(buf)
|
||||
return buf[:n], err
|
||||
}
|
||||
|
||||
var _ vfs.FileDescriptionImpl = (*mockEndpoint)(nil)
|
||||
|
||||
// Read implements vfs.FileDescriptionImpl.Read details for the parent mockFileDescription.
|
||||
func (s *mockEndpoint) Read(ctx context.Context, dst usermem.IOSequence, opts vfs.ReadOptions) (int64, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.released {
|
||||
return 0, io.EOF
|
||||
// write implements portforwarderTestHarness write.
|
||||
func (nc *netConnMockEndpoint) write(buf []byte) (int, error) {
|
||||
written := 0
|
||||
for {
|
||||
n, err := nc.conn.Write(buf[written:])
|
||||
if err != nil && !linuxerr.Equals(linuxerr.ErrWouldBlock, err) {
|
||||
return n, err
|
||||
}
|
||||
written += n
|
||||
if written >= len(buf) {
|
||||
return written, nil
|
||||
}
|
||||
}
|
||||
if s.readBuf.Len() == 0 {
|
||||
return 0, linuxerr.ErrWouldBlock
|
||||
}
|
||||
buf := s.readBuf.Next(s.readBuf.Len())
|
||||
n, err := dst.CopyOut(ctx, buf)
|
||||
s.queue.Notify(waiter.WritableEvents)
|
||||
return int64(n), err
|
||||
}
|
||||
|
||||
// Write implements vfs.FileDescriptionImpl.Write details for the parent mockFileDescription.
|
||||
func (s *mockEndpoint) Write(ctx context.Context, src usermem.IOSequence, opts vfs.WriteOptions) (int64, error) {
|
||||
return s.other.write(ctx, src, opts)
|
||||
}
|
||||
|
||||
func (s *mockEndpoint) write(ctx context.Context, src usermem.IOSequence, opts vfs.WriteOptions) (int64, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.released {
|
||||
return 0, io.EOF
|
||||
}
|
||||
buf := make([]byte, src.NumBytes())
|
||||
n, err := src.CopyIn(ctx, buf)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
n, err = s.readBuf.Write(buf[:n])
|
||||
s.queue.Notify(waiter.ReadableEvents)
|
||||
return int64(n), err
|
||||
}
|
||||
|
||||
func (s *mockEndpoint) IsReadable() bool {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.released {
|
||||
return false
|
||||
}
|
||||
return s.readBuf.Len() > 0
|
||||
}
|
||||
|
||||
func (s *mockEndpoint) IsWritable() bool {
|
||||
return s.other.isWritable()
|
||||
}
|
||||
|
||||
func (s *mockEndpoint) isWritable() bool {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return !s.released
|
||||
}
|
||||
|
||||
// EventRegister implements vfs.FileDescriptionImpl.EventRegister details for the parent mockFileDescription.
|
||||
func (s *mockEndpoint) EventRegister(we *waiter.Entry) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.queue.EventRegister(we)
|
||||
return nil
|
||||
}
|
||||
|
||||
// EventUnregister implements vfs.FileDescriptionImpl.Unregister details for the parent mockFileDescription.
|
||||
func (s *mockEndpoint) EventUnregister(we *waiter.Entry) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.queue.EventUnregister(we)
|
||||
}
|
||||
|
||||
// Release implements vfs.FileDescriptionImpl.Release details for the parent mockFileDescription.
|
||||
func (s *mockEndpoint) Release(context.Context) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.queue.Notify(waiter.ReadableEvents)
|
||||
s.released = true
|
||||
}
|
||||
|
||||
var responses = map[string]string{
|
||||
"PING": "PONG",
|
||||
"DING": "DONG",
|
||||
"TING": "TONG",
|
||||
}
|
||||
|
||||
func TestHostinetPortForwardConn(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
requests map[string]string
|
||||
}{
|
||||
{
|
||||
name: "single",
|
||||
requests: map[string]string{
|
||||
"PING": "PONG",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "multiple",
|
||||
requests: map[string]string{
|
||||
"PING": "PONG",
|
||||
"HELLO": "GOODBYE",
|
||||
"IMPRESSIVE": "MOST IMPRESSIVE",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "empty",
|
||||
requests: map[string]string{
|
||||
"EMPTY": "",
|
||||
"NOT": "EMPTY",
|
||||
"OTHER EMPTY": "",
|
||||
},
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
doHostinetTest(t, tc.requests)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func doHostinetTest(t *testing.T, requests map[string]string) {
|
||||
ctx := contexttest.Context(t)
|
||||
clientSock, server := newMockSocketPair()
|
||||
defer server.Release(ctx)
|
||||
client, err := newMockFileDescription(ctx, clientSock)
|
||||
appEndpoint := &mockApplicationFDImpl{}
|
||||
defer appEndpoint.Release(ctx)
|
||||
client, err := newMockFileDescription(ctx, appEndpoint)
|
||||
if err != nil {
|
||||
t.Fatalf("newMockFileDescription failed: %v", err)
|
||||
}
|
||||
@@ -252,60 +198,37 @@ func TestHostinetPortForwardConn(t *testing.T) {
|
||||
t.Fatalf("l.Accept failed: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
buf := make([]byte, 4)
|
||||
for req, resp := range responses {
|
||||
|
||||
for {
|
||||
if server.IsWritable() {
|
||||
break
|
||||
}
|
||||
harness := portforwarderTestHarness{
|
||||
app: appEndpoint,
|
||||
shim: &netConnMockEndpoint{conn},
|
||||
}
|
||||
|
||||
for req, resp := range requests {
|
||||
if _, err := harness.shimWrite([]byte(req)); err != nil {
|
||||
t.Fatalf("failed to write to shim: %v", err)
|
||||
}
|
||||
_, err := server.Write(ctx, usermem.BytesIOSequence([]byte(req)), vfs.WriteOptions{})
|
||||
|
||||
got, err := harness.appRead(len(req))
|
||||
if err != nil {
|
||||
t.Fatalf("file.Write failed: %v", err)
|
||||
t.Fatalf("failed to read from app: %v", err)
|
||||
}
|
||||
|
||||
read := 0
|
||||
for {
|
||||
n, err := conn.Read([]byte(buf)[read:])
|
||||
if err != nil && !linuxerr.Equals(linuxerr.ErrWouldBlock, err) {
|
||||
t.Fatalf("conn.Write failed: %v", err)
|
||||
}
|
||||
read += n
|
||||
if read >= len(resp) {
|
||||
break
|
||||
}
|
||||
if string(got) != req {
|
||||
t.Fatalf("app mismatch: got: %s want: %s", string(got), req)
|
||||
}
|
||||
|
||||
if string(buf) != req {
|
||||
t.Fatalf("read mismatch: got: %s want: %s", string(buf), req)
|
||||
if _, err := harness.appWrite([]byte(resp)); err != nil {
|
||||
t.Fatalf("failed to write to app: %v", err)
|
||||
}
|
||||
|
||||
written := 0
|
||||
for i := 0; i < 4; i++ {
|
||||
n, err := conn.Write([]byte(resp)[written:])
|
||||
if err != nil && !linuxerr.Equals(linuxerr.ErrWouldBlock, err) {
|
||||
t.Fatalf("conn.Write failed: %v", err)
|
||||
}
|
||||
written += n
|
||||
if written >= len(resp) {
|
||||
break
|
||||
}
|
||||
got, err = harness.shimRead(len(resp))
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read from shim: %v", err)
|
||||
}
|
||||
|
||||
for {
|
||||
if server.IsReadable() {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
_, err = server.Read(ctx, usermem.BytesIOSequence([]byte(buf[:4])), vfs.ReadOptions{})
|
||||
if err != nil && !linuxerr.Equals(linuxerr.ErrWouldBlock, err) {
|
||||
t.Fatalf("file.Read failed: %v", err)
|
||||
}
|
||||
|
||||
if string(buf) != resp {
|
||||
t.Fatalf("write mismatch: got: %s want: %s", string(buf), resp)
|
||||
if string(got) != resp {
|
||||
t.Fatalf("shim mismatch: got: %s want: %s", string(got), resp)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,232 @@
|
||||
// Copyright 2023 The gVisor Authors.
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package portforward
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
|
||||
"gvisor.dev/gvisor/pkg/cleanup"
|
||||
"gvisor.dev/gvisor/pkg/context"
|
||||
"gvisor.dev/gvisor/pkg/log"
|
||||
"gvisor.dev/gvisor/pkg/sentry/vfs"
|
||||
"gvisor.dev/gvisor/pkg/sync"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/network/ipv4"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/stack"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/transport/tcp"
|
||||
"gvisor.dev/gvisor/pkg/waiter"
|
||||
)
|
||||
|
||||
// netstackPortForwardConn is a portForwardConn implementation for netstack.
|
||||
type netstackPortForwardConn struct {
|
||||
// cid is the container id that this connection is connecting to.
|
||||
cid string
|
||||
|
||||
// ep is the tcpip.Endpoint to the application port.
|
||||
ep tcpip.Endpoint
|
||||
// wq is the endpoint waiter.Queue.
|
||||
wq *waiter.Queue
|
||||
// fd is the FileDescription for the imported host UDS fd.
|
||||
fd *vfs.FileDescription
|
||||
|
||||
// status holds the status of the connection.
|
||||
status struct {
|
||||
sync.Mutex
|
||||
// started indicates if the connection is started or not.
|
||||
started bool
|
||||
// closed indicates if the connection is closed or not.
|
||||
closed bool
|
||||
}
|
||||
|
||||
// toDone is closed when the copy to the application port is finished.
|
||||
toDone chan struct{}
|
||||
|
||||
// fromDone is closed when the copy from the application socket is finished.
|
||||
fromDone chan struct{}
|
||||
|
||||
// cu is called when the connection finishes.
|
||||
cu cleanup.Cleanup
|
||||
}
|
||||
|
||||
// newNetstackPortForward creates a new port forwarding connection to the given
|
||||
// port in netstack mode.
|
||||
func newNetstackPortForward(ctx context.Context, stack *stack.Stack, cid string, fd *vfs.FileDescription, port uint16) (portForwardConn, error) {
|
||||
var wq waiter.Queue
|
||||
ep, tcpErr := stack.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &wq)
|
||||
if tcpErr != nil {
|
||||
return nil, fmt.Errorf("creating endpoint: %v", tcpErr)
|
||||
}
|
||||
cu := cleanup.Make(func() { ep.Close() })
|
||||
defer cu.Clean()
|
||||
|
||||
waitEntry, notifyCh := waiter.NewChannelEntry(waiter.WritableEvents)
|
||||
wq.EventRegister(&waitEntry)
|
||||
defer wq.EventUnregister(&waitEntry)
|
||||
|
||||
tcpErr = ep.Connect(tcpip.FullAddress{
|
||||
Addr: "\x7f\x00\x00\x01", // 127.0.0.1
|
||||
Port: port,
|
||||
})
|
||||
if _, ok := tcpErr.(*tcpip.ErrConnectStarted); ok {
|
||||
<-notifyCh
|
||||
tcpErr = ep.LastError()
|
||||
}
|
||||
if tcpErr != nil {
|
||||
return nil, fmt.Errorf("connecting endpoint: %v", tcpErr)
|
||||
}
|
||||
|
||||
pfConn := netstackPortForwardConn{
|
||||
cid: cid,
|
||||
ep: ep,
|
||||
wq: &wq,
|
||||
fd: fd,
|
||||
toDone: make(chan struct{}),
|
||||
fromDone: make(chan struct{}),
|
||||
cu: cleanup.Cleanup{},
|
||||
}
|
||||
|
||||
cu.Release()
|
||||
return &pfConn, nil
|
||||
}
|
||||
|
||||
// start implements portForwardConn.start.
|
||||
func (c *netstackPortForwardConn) start(ctx context.Context) error {
|
||||
c.status.Lock()
|
||||
defer c.status.Unlock()
|
||||
|
||||
if c.status.closed {
|
||||
return fmt.Errorf("already closed")
|
||||
}
|
||||
if c.status.started {
|
||||
return fmt.Errorf("already started")
|
||||
}
|
||||
|
||||
log.Debugf("Start forwarding to/from container %q and localhost", c.cid)
|
||||
|
||||
go c.writeToEP(ctx)
|
||||
go c.readFromEP(ctx)
|
||||
|
||||
c.status.started = true
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// close implements portForwardConn.close.
|
||||
func (c *netstackPortForwardConn) close(ctx context.Context) error {
|
||||
c.status.Lock()
|
||||
|
||||
// This should be a no op if the connection is already closed.
|
||||
if c.status.closed {
|
||||
c.status.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
log.Debugf("Stopping forwarding to/from container %q and localhost...", c.cid)
|
||||
|
||||
// Closing the endpoint will make the other goroutine exit.
|
||||
c.ep.Close()
|
||||
c.fd.DecRef(ctx)
|
||||
|
||||
<-c.toDone
|
||||
log.Debugf("Stopped forwarding one-half of copy for %q", c.cid)
|
||||
|
||||
// Wait on the other goroutine.
|
||||
<-c.fromDone
|
||||
log.Debugf("Stopped forwarding to/from container %q and localhost", c.cid)
|
||||
|
||||
c.status.closed = true
|
||||
|
||||
c.status.Unlock()
|
||||
|
||||
// Call the cleanup object.
|
||||
c.cu.Clean()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// cleanup implements portForwardConn.cleanup.
|
||||
func (c *netstackPortForwardConn) cleanup(f func()) {
|
||||
c.cu.Add(f)
|
||||
}
|
||||
|
||||
// readFromEP reads from the tcpip.Endpoint and writes to the given Writer.
|
||||
func (c *netstackPortForwardConn) readFromEP(ctx context.Context) {
|
||||
w := &fileDescriptionReadWriter{
|
||||
file: c.fd,
|
||||
}
|
||||
|
||||
// Register for read notifications.
|
||||
waitEntry, notifyCh := waiter.NewChannelEntry(waiter.EventIn | waiter.EventHUp | waiter.EventErr)
|
||||
// Register for when the endpoint is readable or disconnected.
|
||||
c.wq.EventRegister(&waitEntry)
|
||||
|
||||
for {
|
||||
_, err := c.ep.Read(w, tcpip.ReadOptions{})
|
||||
if err != nil {
|
||||
if _, ok := err.(*tcpip.ErrWouldBlock); ok {
|
||||
<-notifyCh
|
||||
continue
|
||||
}
|
||||
log.Infof("Port forward read error; cid: %q: %v", c.cid, err)
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
// Clean up when one half of the copy is finished.
|
||||
c.wq.EventUnregister(&waitEntry)
|
||||
c.ep.Shutdown(tcpip.ShutdownRead)
|
||||
close(c.fromDone)
|
||||
c.close(ctx)
|
||||
}
|
||||
|
||||
func (c *netstackPortForwardConn) writeToEP(ctx context.Context) {
|
||||
r := &fileDescriptionReadWriter{
|
||||
file: c.fd,
|
||||
}
|
||||
|
||||
// Register for write notifications.
|
||||
waitEntry, notifyCh := waiter.NewChannelEntry(waiter.WritableEvents | waiter.EventHUp | waiter.EventErr)
|
||||
// Register for when the endpoint is writable or disconnected.
|
||||
c.wq.EventRegister(&waitEntry)
|
||||
|
||||
v := make([]byte, 16384 /* 16kb read buffer size */)
|
||||
for {
|
||||
n, err := r.Read(v)
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
var b bytes.Reader
|
||||
b.Reset(v[:n])
|
||||
for b.Len() != 0 {
|
||||
_, err := c.ep.Write(&b, tcpip.WriteOptions{Atomic: true})
|
||||
if err != nil {
|
||||
// If the channel is not ready for writing then wait until it is.
|
||||
if _, ok := err.(*tcpip.ErrWouldBlock); ok {
|
||||
<-notifyCh
|
||||
continue
|
||||
}
|
||||
log.Infof("Port forward read error; cid: %q: %v", c.cid, err)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Clean up when one half of the copy is finished.
|
||||
c.wq.EventUnregister(&waitEntry)
|
||||
c.ep.Shutdown(tcpip.ShutdownWrite)
|
||||
close(c.toDone)
|
||||
c.close(ctx)
|
||||
}
|
||||
@@ -0,0 +1,202 @@
|
||||
// Copyright 2023 The gVisor Authors.
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package portforward
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"gvisor.dev/gvisor/pkg/cleanup"
|
||||
"gvisor.dev/gvisor/pkg/sentry/contexttest"
|
||||
"gvisor.dev/gvisor/pkg/sentry/vfs"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/waiter"
|
||||
)
|
||||
|
||||
type baseTCPEndpointImpl struct {
|
||||
closed bool
|
||||
readBuf bytes.Buffer
|
||||
writeBuf bytes.Buffer
|
||||
mu sync.Mutex
|
||||
wq *waiter.Queue
|
||||
}
|
||||
|
||||
// read reads data from the buffer that "Write" writes to.
|
||||
func (b *baseTCPEndpointImpl) read(n int) ([]byte, error) {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
if b.closed {
|
||||
return nil, io.EOF
|
||||
}
|
||||
ret := b.writeBuf.Next(n)
|
||||
b.wq.Notify(waiter.WritableEvents)
|
||||
return ret, nil
|
||||
}
|
||||
|
||||
// write writes data to the read buffer that "Read" reads from.
|
||||
func (b *baseTCPEndpointImpl) write(buf []byte) (int, error) {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
if b.closed {
|
||||
return 0, io.EOF
|
||||
}
|
||||
n, err := b.readBuf.Write(buf)
|
||||
b.wq.Notify(waiter.ReadableEvents)
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (b *baseTCPEndpointImpl) Close() {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
b.closed = true
|
||||
}
|
||||
|
||||
func (b *baseTCPEndpointImpl) Read(w io.Writer, _ tcpip.ReadOptions) (tcpip.ReadResult, tcpip.Error) {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
if b.closed {
|
||||
return tcpip.ReadResult{}, &tcpip.ErrClosedForReceive{}
|
||||
}
|
||||
buf := b.readBuf.Next(b.readBuf.Len())
|
||||
n, err := w.Write(buf)
|
||||
if err != nil {
|
||||
return tcpip.ReadResult{}, &tcpip.ErrInvalidEndpointState{}
|
||||
}
|
||||
return tcpip.ReadResult{
|
||||
Count: n,
|
||||
Total: n,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *baseTCPEndpointImpl) Write(payload tcpip.Payloader, _ tcpip.WriteOptions) (int64, tcpip.Error) {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
if b.closed {
|
||||
return 0, &tcpip.ErrClosedForSend{}
|
||||
}
|
||||
buf := make([]byte, payload.Len())
|
||||
n, err := payload.Read(buf)
|
||||
if err != nil {
|
||||
return 0, &tcpip.ErrInvalidEndpointState{}
|
||||
}
|
||||
n, err = b.writeBuf.Write(buf[:n])
|
||||
if err != nil {
|
||||
return int64(n), &tcpip.ErrConnectionRefused{}
|
||||
}
|
||||
return int64(n), nil
|
||||
}
|
||||
|
||||
func (b *baseTCPEndpointImpl) Shutdown(shutdown tcpip.ShutdownFlags) tcpip.Error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func newNetstackPortForwardConnWithMock(impl mockTCPEndpointImpl, fd *vfs.FileDescription, wq *waiter.Queue) *netstackPortForwardConn {
|
||||
ep := &mockTCPEndpoint{impl}
|
||||
return &netstackPortForwardConn{
|
||||
ep: ep,
|
||||
wq: wq,
|
||||
fd: fd,
|
||||
toDone: make(chan struct{}),
|
||||
fromDone: make(chan struct{}),
|
||||
cu: cleanup.Cleanup{},
|
||||
}
|
||||
}
|
||||
|
||||
func TestNetstackPortforward(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
requests map[string]string
|
||||
}{
|
||||
{
|
||||
name: "single",
|
||||
requests: map[string]string{
|
||||
"PING": "PONG",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "multiple",
|
||||
requests: map[string]string{
|
||||
"PING": "PONG",
|
||||
"HELLO": "GOODBYE",
|
||||
"IMPRESSIVE": "MOST IMPRESSIVE",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "empty",
|
||||
requests: map[string]string{
|
||||
"EMPTY": "",
|
||||
"NOT": "EMPTY",
|
||||
"OTHER EMPTY": "",
|
||||
},
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
doNetstackTest(t, tc.requests)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func doNetstackTest(t *testing.T, responses map[string]string) {
|
||||
ctx := contexttest.Context(t)
|
||||
appEndpoint := &mockApplicationFDImpl{}
|
||||
defer appEndpoint.Release(ctx)
|
||||
fd, err := newMockFileDescription(ctx, appEndpoint)
|
||||
if err != nil {
|
||||
t.Fatalf("newMockFileDescription: %v", err)
|
||||
}
|
||||
|
||||
wq := waiter.Queue{}
|
||||
impl := &baseTCPEndpointImpl{wq: &wq}
|
||||
conn := newNetstackPortForwardConnWithMock(impl, fd, &wq)
|
||||
if err := conn.start(ctx); err != nil {
|
||||
t.Fatalf("conn.start: %v", err)
|
||||
}
|
||||
defer conn.close(ctx)
|
||||
|
||||
harness := portforwarderTestHarness{
|
||||
app: appEndpoint,
|
||||
shim: impl,
|
||||
}
|
||||
|
||||
for req, resp := range responses {
|
||||
if _, err := harness.shimWrite([]byte(req)); err != nil {
|
||||
t.Fatalf("failed to write to shim: %v", err)
|
||||
}
|
||||
|
||||
got, err := harness.appRead(len(req))
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read from app: %v", err)
|
||||
}
|
||||
|
||||
if string(got) != req {
|
||||
t.Fatalf("app mismatch: got: %s want: %s", string(got), req)
|
||||
}
|
||||
|
||||
if _, err := harness.appWrite([]byte(resp)); err != nil {
|
||||
t.Fatalf("failed to write to app: %v", err)
|
||||
}
|
||||
|
||||
got, err = harness.shimRead(len(resp))
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read from shim: %v", err)
|
||||
}
|
||||
|
||||
if string(got) != resp {
|
||||
t.Fatalf("shim mismatch: got: %s want: %s", string(got), resp)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,339 @@
|
||||
// Copyright 2023 The gVisor Authors.
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package portforward
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"sync"
|
||||
|
||||
"gvisor.dev/gvisor/pkg/context"
|
||||
"gvisor.dev/gvisor/pkg/errors/linuxerr"
|
||||
"gvisor.dev/gvisor/pkg/sentry/vfs"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/usermem"
|
||||
"gvisor.dev/gvisor/pkg/waiter"
|
||||
)
|
||||
|
||||
// mockEndpoint defines an endpoint that tests can read and write for validating portforwarders.
|
||||
type mockEndpoint interface {
|
||||
read(n int) ([]byte, error)
|
||||
write(buf []byte) (int, error)
|
||||
}
|
||||
|
||||
// portforwarderTestHarness mocks both sides of the portforwarder connection so that behavior can be
|
||||
// validated between them.
|
||||
type portforwarderTestHarness struct {
|
||||
app mockEndpoint
|
||||
shim mockEndpoint
|
||||
}
|
||||
|
||||
func (th *portforwarderTestHarness) appWrite(buf []byte) (int, error) {
|
||||
return th.app.write(buf)
|
||||
}
|
||||
|
||||
func (th *portforwarderTestHarness) appRead(n int) ([]byte, error) {
|
||||
return th.doRead(n, th.app)
|
||||
}
|
||||
|
||||
func (th *portforwarderTestHarness) shimWrite(buf []byte) (int, error) {
|
||||
return th.shim.write(buf)
|
||||
}
|
||||
|
||||
func (th *portforwarderTestHarness) shimRead(n int) ([]byte, error) {
|
||||
return th.doRead(n, th.shim)
|
||||
}
|
||||
|
||||
func (th *portforwarderTestHarness) doRead(n int, ep mockEndpoint) ([]byte, error) {
|
||||
buf := make([]byte, 0, n)
|
||||
for {
|
||||
out, err := ep.read(n - len(buf))
|
||||
if err != nil && !linuxerr.Equals(linuxerr.ErrWouldBlock, err) {
|
||||
return nil, err
|
||||
}
|
||||
buf = append(buf, out...)
|
||||
if len(buf) >= n {
|
||||
return buf, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// mockApplicationFDImpl mocks a VFS file description endpoint on which the sandboxed application
|
||||
// and the portforwarder will communicate.
|
||||
type mockApplicationFDImpl struct {
|
||||
vfs.FileDescriptionDefaultImpl
|
||||
vfs.NoLockFD
|
||||
vfs.DentryMetadataFileDescriptionImpl
|
||||
mu sync.Mutex
|
||||
readBuf bytes.Buffer
|
||||
writeBuf bytes.Buffer
|
||||
released bool
|
||||
queue waiter.Queue
|
||||
}
|
||||
|
||||
var _ vfs.FileDescriptionImpl = (*mockApplicationFDImpl)(nil)
|
||||
|
||||
// Read implements vfs.FileDescriptionImpl.Read details for the parent mockFileDescription.
|
||||
func (s *mockApplicationFDImpl) Read(ctx context.Context, dst usermem.IOSequence, opts vfs.ReadOptions) (int64, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.released {
|
||||
return 0, io.EOF
|
||||
}
|
||||
if s.readBuf.Len() == 0 {
|
||||
return 0, linuxerr.ErrWouldBlock
|
||||
}
|
||||
buf := s.readBuf.Next(s.readBuf.Len())
|
||||
n, err := dst.CopyOut(ctx, buf)
|
||||
return int64(n), err
|
||||
}
|
||||
|
||||
// Write implements vfs.FileDescriptionImpl.Write details for the parent mockFileDescription.
|
||||
func (s *mockApplicationFDImpl) Write(ctx context.Context, src usermem.IOSequence, opts vfs.WriteOptions) (int64, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.released {
|
||||
return 0, io.EOF
|
||||
}
|
||||
|
||||
buf := make([]byte, src.NumBytes())
|
||||
n, err := src.CopyIn(ctx, buf)
|
||||
if err != nil {
|
||||
return int64(n), err
|
||||
}
|
||||
res, err := s.writeBuf.Write(buf)
|
||||
return int64(res), err
|
||||
}
|
||||
|
||||
// write implements mockEndpoint.write.
|
||||
func (s *mockApplicationFDImpl) write(buf []byte) (int, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.released {
|
||||
return 0, io.EOF
|
||||
}
|
||||
ret, err := s.readBuf.Write(buf)
|
||||
s.queue.Notify(waiter.ReadableEvents)
|
||||
return ret, err
|
||||
}
|
||||
|
||||
// read implements mockEndpoint.read.
|
||||
func (s *mockApplicationFDImpl) read(n int) ([]byte, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.released {
|
||||
return nil, io.EOF
|
||||
}
|
||||
if s.writeBuf.Len() == 0 {
|
||||
return nil, linuxerr.ErrWouldBlock
|
||||
}
|
||||
ret := s.writeBuf.Next(n)
|
||||
s.queue.Notify(waiter.WritableEvents)
|
||||
return ret, nil
|
||||
}
|
||||
|
||||
func (s *mockApplicationFDImpl) IsReadable() bool {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.released {
|
||||
return false
|
||||
}
|
||||
return s.readBuf.Len() > 0
|
||||
}
|
||||
|
||||
func (s *mockApplicationFDImpl) IsWritable() bool {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return !s.released
|
||||
}
|
||||
|
||||
// EventRegister implements vfs.FileDescriptionImpl.EventRegister details for the parent mockFileDescription.
|
||||
func (s *mockApplicationFDImpl) EventRegister(we *waiter.Entry) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.queue.EventRegister(we)
|
||||
return nil
|
||||
}
|
||||
|
||||
// EventUnregister implements vfs.FileDescriptionImpl.Unregister details for the parent mockFileDescription.
|
||||
func (s *mockApplicationFDImpl) EventUnregister(we *waiter.Entry) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.queue.EventUnregister(we)
|
||||
}
|
||||
|
||||
// Release implements vfs.FileDescriptionImpl.Release details for the parent mockFileDescription.
|
||||
func (s *mockApplicationFDImpl) Release(context.Context) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.queue.Notify(waiter.ReadableEvents)
|
||||
s.released = true
|
||||
}
|
||||
|
||||
// mockTCPEndpointImpl is the subset of methods used by tests for the mockTCPEndpoint struct. This
|
||||
// is so we can quickly change implementations as needed.
|
||||
type mockTCPEndpointImpl interface {
|
||||
Close()
|
||||
Read(io.Writer, tcpip.ReadOptions) (tcpip.ReadResult, tcpip.Error)
|
||||
Write(tcpip.Payloader, tcpip.WriteOptions) (int64, tcpip.Error)
|
||||
Shutdown(tcpip.ShutdownFlags) tcpip.Error
|
||||
}
|
||||
|
||||
// mockTCPEndpoint mocks tcpip.Endpoint for tests.
|
||||
type mockTCPEndpoint struct {
|
||||
impl mockTCPEndpointImpl // impl implements the subset of methods needed for mockTCPEndpoints.
|
||||
}
|
||||
|
||||
// The below are trivial stub methods to get mockTCPEndpoint to implement tcpip.Endpoint. They
|
||||
// either panic or call the contained impl's methods.
|
||||
|
||||
// Close implements tcpip.Endpoint.Close.
|
||||
func (m *mockTCPEndpoint) Close() {
|
||||
m.impl.Close()
|
||||
}
|
||||
|
||||
// Abort implements tcpip.Endpoint.Abort.
|
||||
func (m *mockTCPEndpoint) Abort() {
|
||||
m.panicWithNotImplementedMsg()
|
||||
}
|
||||
|
||||
// Read implements tcpip.Endpoint.Read.
|
||||
func (m *mockTCPEndpoint) Read(w io.Writer, opts tcpip.ReadOptions) (tcpip.ReadResult, tcpip.Error) {
|
||||
return m.impl.Read(w, opts)
|
||||
}
|
||||
|
||||
// Write implements tcpip.Endpoint.Write.
|
||||
func (m *mockTCPEndpoint) Write(payload tcpip.Payloader, opts tcpip.WriteOptions) (int64, tcpip.Error) {
|
||||
return m.impl.Write(payload, opts)
|
||||
}
|
||||
|
||||
// Connect implements tcpip.Endpoint.Connect.
|
||||
func (m *mockTCPEndpoint) Connect(address tcpip.FullAddress) tcpip.Error {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return nil
|
||||
}
|
||||
|
||||
// Disconnect implements tcpip.Endpoint.Disconnect.
|
||||
func (m *mockTCPEndpoint) Disconnect() tcpip.Error {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return nil
|
||||
}
|
||||
|
||||
// Shutdown implements tcpip.Endpoint.Shutdown.
|
||||
func (m *mockTCPEndpoint) Shutdown(flags tcpip.ShutdownFlags) tcpip.Error {
|
||||
return m.impl.Shutdown(flags)
|
||||
}
|
||||
|
||||
// Listen implements tcpip.Endpoint.Listen.
|
||||
func (m *mockTCPEndpoint) Listen(backlog int) tcpip.Error {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return nil
|
||||
}
|
||||
|
||||
// Accept implements tcpip.Endpoint.Accept.
|
||||
func (m *mockTCPEndpoint) Accept(peerAddr *tcpip.FullAddress) (tcpip.Endpoint, *waiter.Queue, tcpip.Error) {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return nil, nil, nil
|
||||
}
|
||||
|
||||
// Bind implements tcpip.Endpoint.Bind.
|
||||
func (m *mockTCPEndpoint) Bind(address tcpip.FullAddress) tcpip.Error {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetLocalAddress implements tcpip.Endpoint.GetLocalAddress.
|
||||
func (m mockTCPEndpoint) GetLocalAddress() (tcpip.FullAddress, tcpip.Error) {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return tcpip.FullAddress{}, nil
|
||||
}
|
||||
|
||||
// GetRemoteAddress implements tcpip.Endpoint.GetRemoreAddress.
|
||||
func (m *mockTCPEndpoint) GetRemoteAddress() (tcpip.FullAddress, tcpip.Error) {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return tcpip.FullAddress{}, nil
|
||||
}
|
||||
|
||||
// Readiness implements tcpip.Endpoint.Readiness.
|
||||
func (m *mockTCPEndpoint) Readiness(mask waiter.EventMask) waiter.EventMask {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return 0
|
||||
}
|
||||
|
||||
// SetSockOpt implements tcpip.Endpoint.SetSockOpt.
|
||||
func (m *mockTCPEndpoint) SetSockOpt(opt tcpip.SettableSocketOption) tcpip.Error {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetSockOptInt implements tcpip.Endpoint.SetSockOptInt.
|
||||
func (m *mockTCPEndpoint) SetSockOptInt(opt tcpip.SockOptInt, v int) tcpip.Error {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetSockOpt implements tcpip.Endpoint.GetSockOpt.
|
||||
func (m *mockTCPEndpoint) GetSockOpt(opt tcpip.GettableSocketOption) tcpip.Error {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetSockOptInt implements tcpip.Endpoint.GetSockOpt.
|
||||
func (m *mockTCPEndpoint) GetSockOptInt(tcpip.SockOptInt) (int, tcpip.Error) {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
// State implements tcpip.Endpoint.State.
|
||||
func (m *mockTCPEndpoint) State() uint32 {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return 0
|
||||
}
|
||||
|
||||
// ModerateRecvBuf implements tcpip.Endpoint.ModerateRecvBuf
|
||||
func (m *mockTCPEndpoint) ModerateRecvBuf(copied int) {
|
||||
m.panicWithNotImplementedMsg()
|
||||
}
|
||||
|
||||
// Info implements tcpip.Endpoint.Info.
|
||||
func (m *mockTCPEndpoint) Info() tcpip.EndpointInfo {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return nil
|
||||
}
|
||||
|
||||
// Stats implements tcpip.Endpoint.Stats.
|
||||
func (m *mockTCPEndpoint) Stats() tcpip.EndpointStats {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetOwner implements tcpip.Endpoint.SetOwner.
|
||||
func (m *mockTCPEndpoint) SetOwner(owner tcpip.PacketOwner) {
|
||||
m.panicWithNotImplementedMsg()
|
||||
}
|
||||
|
||||
// LastError implements tcpip.Endpoint.LastError.
|
||||
func (m *mockTCPEndpoint) LastError() tcpip.Error {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return nil
|
||||
}
|
||||
|
||||
// SocketOptions implements tcpip.Endpoint.SocketOptions.
|
||||
func (m *mockTCPEndpoint) SocketOptions() *tcpip.SocketOptions {
|
||||
m.panicWithNotImplementedMsg()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (*mockTCPEndpoint) panicWithNotImplementedMsg() { panic("not implemented") }
|
||||
Reference in New Issue
Block a user