mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix: don't use yamux for in-memory provisioner{,d} streams (#5136)
This commit is contained in:
+8
-12
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user