fix: don't use yamux for in-memory provisioner{,d} streams (#5136)

This commit is contained in:
Colin Adler
2022-11-22 12:19:32 -06:00
committed by GitHub
parent 2b6c229e4e
commit 1f20cab110
14 changed files with 105 additions and 57 deletions
+8 -12
View File
@@ -7,12 +7,12 @@ import (
"net"
"os"
"github.com/hashicorp/yamux"
"github.com/valyala/fasthttp/fasthttputil"
"golang.org/x/xerrors"
"storj.io/drpc/drpcmux"
"storj.io/drpc/drpcserver"
"github.com/hashicorp/yamux"
"github.com/coder/coder/provisionersdk/proto"
)
@@ -58,18 +58,14 @@ func Serve(ctx context.Context, server proto.DRPCProvisionerServer, options *Ser
// short-lived processes that can be executed concurrently.
err = srv.Serve(ctx, options.Listener)
if err != nil {
if errors.Is(err, io.EOF) {
return nil
}
if errors.Is(err, context.Canceled) {
return nil
}
if errors.Is(err, io.ErrClosedPipe) {
return nil
}
if errors.Is(err, yamux.ErrSessionShutdown) {
if errors.Is(err, io.EOF) ||
errors.Is(err, context.Canceled) ||
errors.Is(err, io.ErrClosedPipe) ||
errors.Is(err, yamux.ErrSessionShutdown) ||
errors.Is(err, fasthttputil.ErrInmemoryListenerClosed) {
return nil
}
return xerrors.Errorf("serve transport: %w", err)
}
return nil
+3 -3
View File
@@ -21,7 +21,7 @@ func TestProvisionerSDK(t *testing.T) {
t.Parallel()
t.Run("Serve", func(t *testing.T) {
t.Parallel()
client, server := provisionersdk.TransportPipe()
client, server := provisionersdk.MemTransportPipe()
defer client.Close()
defer server.Close()
@@ -34,7 +34,7 @@ func TestProvisionerSDK(t *testing.T) {
assert.NoError(t, err)
}()
api := proto.NewDRPCProvisionerClient(provisionersdk.Conn(client))
api := proto.NewDRPCProvisionerClient(client)
stream, err := api.Parse(context.Background(), &proto.Parse_Request{})
require.NoError(t, err)
_, err = stream.Recv()
@@ -43,7 +43,7 @@ func TestProvisionerSDK(t *testing.T) {
t.Run("ServeClosedPipe", func(t *testing.T) {
t.Parallel()
client, server := provisionersdk.TransportPipe()
client, server := provisionersdk.MemTransportPipe()
_ = client.Close()
_ = server.Close()
+63 -19
View File
@@ -2,10 +2,11 @@ package provisionersdk
import (
"context"
"io"
"net"
"sync"
"github.com/hashicorp/yamux"
"github.com/valyala/fasthttp/fasthttputil"
"storj.io/drpc"
"storj.io/drpc/drpcconn"
)
@@ -16,24 +17,8 @@ const (
MaxMessageSize = 4 << 20
)
// TransportPipe creates an in-memory pipe for dRPC transport.
func TransportPipe() (*yamux.Session, *yamux.Session) {
c1, c2 := net.Pipe()
yamuxConfig := yamux.DefaultConfig()
yamuxConfig.LogOutput = io.Discard
client, err := yamux.Client(c1, yamuxConfig)
if err != nil {
panic(err)
}
server, err := yamux.Server(c2, yamuxConfig)
if err != nil {
panic(err)
}
return client, server
}
// Conn returns a multiplexed dRPC connection from a yamux session.
func Conn(session *yamux.Session) drpc.Conn {
// MultiplexedConn returns a multiplexed dRPC connection from a yamux session.
func MultiplexedConn(session *yamux.Session) drpc.Conn {
return &multiplexedDRPC{session}
}
@@ -78,3 +63,62 @@ func (m *multiplexedDRPC) NewStream(ctx context.Context, rpc string, enc drpc.En
}
return stream, err
}
func MemTransportPipe() (drpc.Conn, net.Listener) {
m := &memDRPC{
closed: make(chan struct{}),
l: fasthttputil.NewInmemoryListener(),
}
return m, m.l
}
type memDRPC struct {
closeOnce sync.Once
closed chan struct{}
l *fasthttputil.InmemoryListener
}
func (m *memDRPC) Close() error {
err := m.l.Close()
m.closeOnce.Do(func() { close(m.closed) })
return err
}
func (m *memDRPC) Closed() <-chan struct{} {
return m.closed
}
func (m *memDRPC) Invoke(ctx context.Context, rpc string, enc drpc.Encoding, inMessage, outMessage drpc.Message) error {
conn, err := m.l.Dial()
if err != nil {
return err
}
dConn := drpcconn.New(conn)
defer func() {
_ = dConn.Close()
_ = conn.Close()
}()
return dConn.Invoke(ctx, rpc, enc, inMessage, outMessage)
}
func (m *memDRPC) NewStream(ctx context.Context, rpc string, enc drpc.Encoding) (drpc.Stream, error) {
conn, err := m.l.Dial()
if err != nil {
return nil, err
}
dConn := drpcconn.New(conn)
stream, err := dConn.NewStream(ctx, rpc, enc)
if err == nil {
go func() {
select {
case <-stream.Context().Done():
case <-m.closed:
}
_ = dConn.Close()
_ = conn.Close()
}()
}
return stream, err
}