Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
103 changes: 103 additions & 0 deletions pkg/ioutil/buffered.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
package ioutil

import (
"bytes"
"errors"
"io"
"sync"
)

var (
ErrClosed = errors.New("ioutil: closed")
ErrNotSet = errors.New("ioutil: underlying ReadWriteCloser not set")
ErrAlreadySet = errors.New("ioutil: underlying ReadWriteCloser already set")
)

// BufferedReadWriteCloser is an io.ReadWriteCloser that buffers writes until
// Set supplies the underlying one, then flushes them in order and passes
// everything through.
type BufferedReadWriteCloser interface {
io.ReadWriteCloser

// Set attaches rwc and flushes the buffered writes to it. It closes rwc
// and returns ErrClosed if Close was already called, or the write error
// if the flush fails.
Set(rwc io.ReadWriteCloser) error
}

func NewBufferedReadWriteCloser() BufferedReadWriteCloser {
return &bufferedReadWriteCloser{}
}

type bufferedReadWriteCloser struct {
// mu guards the handoff in Set: a Write that arrives while Set is flushing
// the buffer must land after the flush, on the underlying rwc.
mu sync.Mutex
rwc io.ReadWriteCloser
buf bytes.Buffer
closed bool
}

func (b *bufferedReadWriteCloser) Set(rwc io.ReadWriteCloser) error {
b.mu.Lock()
defer b.mu.Unlock()
if b.closed {
rwc.Close()
return ErrClosed
}
if b.rwc != nil {
return ErrAlreadySet
}
if b.buf.Len() > 0 {
if _, err := rwc.Write(b.buf.Bytes()); err != nil {
rwc.Close()
b.closed = true
b.buf.Reset()
return err
}
b.buf.Reset()
}
b.rwc = rwc
return nil
}

func (b *bufferedReadWriteCloser) Write(p []byte) (int, error) {
b.mu.Lock()
if b.closed {
b.mu.Unlock()
return 0, ErrClosed
}
if b.rwc == nil {
n, err := b.buf.Write(p)
b.mu.Unlock()
return n, err
}
rwc := b.rwc
b.mu.Unlock()
return rwc.Write(p)
}

func (b *bufferedReadWriteCloser) Read(p []byte) (int, error) {
b.mu.Lock()
rwc, closed := b.rwc, b.closed
b.mu.Unlock()
if closed {
return 0, ErrClosed
}
if rwc == nil {
return 0, ErrNotSet
}
return rwc.Read(p)
}

func (b *bufferedReadWriteCloser) Close() error {
b.mu.Lock()
rwc := b.rwc
b.closed = true
b.buf.Reset()
b.mu.Unlock()
if rwc != nil {
return rwc.Close()
}
return nil
}
117 changes: 117 additions & 0 deletions pkg/ioutil/buffered_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,117 @@
package ioutil

import (
"bytes"
"errors"
"io"
"testing"
)

type recordingRWC struct {
written bytes.Buffer
closed bool
writeErr error
}

func (r *recordingRWC) Write(p []byte) (int, error) {
if r.writeErr != nil {
return 0, r.writeErr
}
return r.written.Write(p)
}

func (r *recordingRWC) Read(p []byte) (int, error) { return copy(p, "read"), nil }

func (r *recordingRWC) Close() error {
r.closed = true
return nil
}

func TestWritesBeforeSetAreFlushedInOrder(t *testing.T) {
b := NewBufferedReadWriteCloser()
rwc := &recordingRWC{}

for _, chunk := range []string{"one ", "two ", "three"} {
if _, err := b.Write([]byte(chunk)); err != nil {
t.Fatalf("Write before Set: %v", err)
}
}
if _, err := b.Read(make([]byte, 4)); !errors.Is(err, ErrNotSet) {
t.Fatalf("Read before Set: got %v, want ErrNotSet", err)
}

if err := b.Set(rwc); err != nil {
t.Fatalf("Set: %v", err)
}
if got := rwc.written.String(); got != "one two three" {
t.Fatalf("flushed %q, want %q", got, "one two three")
}

if _, err := b.Write([]byte(" four")); err != nil {
t.Fatalf("Write after Set: %v", err)
}
if got := rwc.written.String(); got != "one two three four" {
t.Fatalf("after write-through got %q", got)
}

p := make([]byte, 4)
if n, err := b.Read(p); err != nil || string(p[:n]) != "read" {
t.Fatalf("Read after Set: %q, %v", p[:n], err)
}

if err := b.Set(&recordingRWC{}); !errors.Is(err, ErrAlreadySet) {
t.Fatalf("second Set: got %v, want ErrAlreadySet", err)
}
}

func TestCloseBeforeSetClosesTheLateArrival(t *testing.T) {
b := NewBufferedReadWriteCloser()
_, _ = b.Write([]byte("queued"))
if err := b.Close(); err != nil {
t.Fatalf("Close: %v", err)
}

rwc := &recordingRWC{}
if err := b.Set(rwc); !errors.Is(err, ErrClosed) {
t.Fatalf("Set after Close: got %v, want ErrClosed", err)
}
if !rwc.closed {
t.Fatal("expected Set to close the ReadWriteCloser handed to a closed buffer")
}
if rwc.written.Len() != 0 {
t.Fatalf("expected nothing flushed to a closed buffer's late arrival, got %q", rwc.written.String())
}
if _, err := b.Write([]byte("x")); !errors.Is(err, ErrClosed) {
t.Fatalf("Write after Close: got %v, want ErrClosed", err)
}
}

func TestFlushFailureClosesTheUnderlying(t *testing.T) {
b := NewBufferedReadWriteCloser()
_, _ = b.Write([]byte("queued"))

rwc := &recordingRWC{writeErr: io.ErrShortWrite}
if err := b.Set(rwc); !errors.Is(err, io.ErrShortWrite) {
t.Fatalf("Set with failing flush: got %v, want ErrShortWrite", err)
}
if !rwc.closed {
t.Fatal("expected the underlying to be closed after a failed flush")
}
if _, err := b.Write([]byte("x")); !errors.Is(err, ErrClosed) {
t.Fatalf("Write after failed flush: got %v, want ErrClosed", err)
}
}

func TestCloseAfterSetClosesTheUnderlying(t *testing.T) {
b := NewBufferedReadWriteCloser()
rwc := &recordingRWC{}
if err := b.Set(rwc); err != nil {
t.Fatalf("Set: %v", err)
}
if err := b.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
if !rwc.closed {
t.Fatal("expected Close to close the underlying")
}
}
66 changes: 47 additions & 19 deletions pkg/tunnel/tunnel.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ import (

"github.com/puzpuzpuz/xsync/v3"
bridgev1 "github.com/vercel/bridge/api/go/bridge/v1"
"github.com/vercel/bridge/pkg/ioutil"
"github.com/vercel/bridge/pkg/mitm"
"github.com/vercel/bridge/pkg/plumbing"
)
Expand Down Expand Up @@ -53,7 +54,7 @@ func New(dialer plumbing.ContextDialer, stream Stream, opts ...Option) Tunnel {
dialer: dialer,
stream: stream,
sendCh: make(chan *bridgev1.TunnelNetworkMessage, 64),
conns: xsync.NewMapOf[string, net.Conn](),
conns: xsync.NewMapOf[string, io.ReadWriteCloser](),
done: make(chan struct{}),
}
for _, o := range opts {
Expand All @@ -67,7 +68,7 @@ type tunnelImpl struct {
hijacker mitm.Hijacker
stream Stream
sendCh chan *bridgev1.TunnelNetworkMessage
conns *xsync.MapOf[string, net.Conn]
conns *xsync.MapOf[string, io.ReadWriteCloser]
done chan struct{}
ctx context.Context
cancel context.CancelFunc
Expand Down Expand Up @@ -99,8 +100,8 @@ func (t *tunnelImpl) AddConn(conn net.Conn, destOverride string, hostname string
go t.readFromConn(conn, connID, src, dst, hostname)
}

// readFromConn reads from a net.Conn and forwards data to the stream via sendCh.
func (t *tunnelImpl) readFromConn(conn net.Conn, connID string, src, dst *bridgev1.TunnelAddress, hostname string) {
// readFromConn reads from a connection and forwards data to the stream via sendCh.
func (t *tunnelImpl) readFromConn(conn io.ReadCloser, connID string, src, dst *bridgev1.TunnelAddress, hostname string) {
defer func() {
conn.Close()
t.conns.Delete(connID)
Expand Down Expand Up @@ -200,8 +201,8 @@ func (t *tunnelImpl) Start(ctx context.Context) {
continue
}

// Unknown connection ID → dial via the configured dialer.
go t.handleNewConn(msg)
// Unknown connection ID → this message opens it.
t.openConn(msg)

case err := <-recvErr:
if err != io.EOF {
Expand All @@ -218,15 +219,31 @@ func (t *tunnelImpl) Start(ctx context.Context) {
}()
}

func (t *tunnelImpl) handleNewConn(msg *bridgev1.TunnelNetworkMessage) {
dest := msg.GetDest()
if dest == nil {
slog.Info("Tunnel: ignoring message with no dest", "conn_id", msg.GetConnectionId())
// openConn registers a buffered connection for msg's ID before dialing. The
// peer has no explicit open message: the first chunk of a connection is what
// opens it, and the next chunks can arrive before the dial completes. Storing
// the buffer synchronously on the recv pump makes those chunks queue behind
// the first one instead of each dialing a second connection under the same ID
// and splitting the stream between them.
func (t *tunnelImpl) openConn(msg *bridgev1.TunnelNetworkMessage) {
connID := msg.GetConnectionId()
if msg.GetDest() == nil {
slog.Info("Tunnel: ignoring message with no dest", "conn_id", connID)
return
}

buffered := ioutil.NewBufferedReadWriteCloser()
t.conns.Store(connID, buffered)
if data := msg.GetData(); len(data) > 0 {
_, _ = buffered.Write(data)
}
go t.dial(msg, buffered)
}

func (t *tunnelImpl) dial(msg *bridgev1.TunnelNetworkMessage, buffered ioutil.BufferedReadWriteCloser) {
connID := msg.GetConnectionId()
hostname := msg.GetHostname()
dest := msg.GetDest()

// Resolve the connection: hijacker gets first shot, then fall through to the dialer.
var conn net.Conn
Expand All @@ -246,6 +263,8 @@ func (t *tunnelImpl) handleNewConn(msg *bridgev1.TunnelNetworkMessage) {

if err != nil {
slog.Info("Tunnel: connect failed", "conn_id", connID, "hostname", hostname, "error", err)
t.deleteIfSame(connID, buffered)
buffered.Close()
select {
case t.sendCh <- &bridgev1.TunnelNetworkMessage{
ConnectionId: connID,
Expand All @@ -256,20 +275,29 @@ func (t *tunnelImpl) handleNewConn(msg *bridgev1.TunnelNetworkMessage) {
return
}

t.conns.Store(connID, conn)
go t.readFromConn(conn, connID, msg.GetDest(), msg.GetSource(), hostname)
if err := buffered.Set(conn); err != nil {
slog.Debug("Tunnel: dropping dialed connection", "connection_id", connID, "error", err)
t.deleteIfSame(connID, buffered)
return
}

if data := msg.GetData(); len(data) > 0 {
if _, err := conn.Write(data); err != nil {
slog.Debug("Failed to write initial data", "connection_id", connID, "error", err)
conn.Close()
t.conns.Delete(connID)
go t.readFromConn(buffered, connID, dest, msg.GetSource(), hostname)
}

// deleteIfSame removes connID only while it still maps to conn, so a connection
// the peer has since reopened under the same ID is left alone.
func (t *tunnelImpl) deleteIfSame(connID string, conn ioutil.BufferedReadWriteCloser) {
t.conns.Compute(connID, func(cur io.ReadWriteCloser, loaded bool) (io.ReadWriteCloser, bool) {
if !loaded {
return nil, true
}
}
b, ok := cur.(ioutil.BufferedReadWriteCloser)
return cur, ok && b == conn
})
}

func (t *tunnelImpl) closeAll() {
t.conns.Range(func(key string, conn net.Conn) bool {
t.conns.Range(func(key string, conn io.ReadWriteCloser) bool {
conn.Close()
t.conns.Delete(key)
return true
Expand Down
Loading
Loading