diff --git a/cmd/root.go b/cmd/root.go index 76eeb8913..132261ca7 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -1243,8 +1243,10 @@ func runSignalWrapper(cmd *Command) (err error) { go func() { err := p.Serve(ctx, notifyStarted) - cmd.logger.Debugf("proxy server error: %v", err) - shutdownCh <- err + if err != nil { + cmd.logger.Debugf("proxy server error: %v", err) + shutdownCh <- err + } }() err = <-shutdownCh @@ -1262,7 +1264,9 @@ func runSignalWrapper(cmd *Command) (err error) { cmd.logger.Infof("/quitquitquit received request. Shutting down...") time.Sleep(cmd.conf.WaitBeforeClose) default: - cmd.logger.Errorf("The proxy has encountered a terminal error: %v", err) + if err != nil { + cmd.logger.Errorf("The proxy has encountered a terminal error: %v", err) + } } return err } diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index 5c8732614..07f3d35f0 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -16,6 +16,7 @@ package proxy import ( "context" + "errors" "fmt" "io" "net" @@ -693,8 +694,11 @@ func (c *Client) Serve(ctx context.Context, notify func()) error { } exitCh := make(chan error) + var wg sync.WaitGroup for _, m := range c.mnts { + wg.Add(1) go func(mnt *socketMount) { + defer wg.Done() err := c.serveSocketMount(ctx, mnt) if err != nil { select { @@ -710,8 +714,17 @@ func (c *Client) Serve(ctx context.Context, notify func()) error { } }(m) } + go func() { + wg.Wait() + close(exitCh) + }() notify() - return <-exitCh + select { + case <-ctx.Done(): + return nil + case err := <-exitCh: + return err + } } // MultiErr is a group of errors wrapped into one. @@ -799,6 +812,9 @@ func (c *Client) serveSocketMount(ctx context.Context, s *socketMount) error { } cConn, err := s.Accept() if err != nil { + if errors.Is(err, net.ErrClosed) || strings.Contains(err.Error(), "use of closed network connection") { + return nil + } if nerr, ok := err.(net.Error); ok && nerr.Timeout() { c.logger.Errorf("[%s] Error accepting connection: %v", s.inst, err) // For transient errors, wait a small amount of time to see if it resolves itself diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go index c529d6a84..51349ef99 100644 --- a/internal/proxy/proxy_test.go +++ b/internal/proxy/proxy_test.go @@ -940,3 +940,42 @@ func TestProxyMultiInstances(t *testing.T) { }) } } + +func TestServeExitsCleanlyOnClose(t *testing.T) { + in := &proxy.Config{ + Addr: "127.0.0.1", + Port: 24018, + Instances: []proxy.InstanceConnConfig{ + {Name: "proj:region:pg"}, + }, + } + d := &fakeDialer{} + c, err := proxy.NewClient(context.Background(), d, testLogger, in, nil) + if err != nil { + t.Fatalf("proxy.NewClient error: %v", err) + } + + serveErrCh := make(chan error, 1) + started := make(chan struct{}) + go func() { + serveErrCh <- c.Serve(context.Background(), func() { + close(started) + }) + }() + + <-started + + if err := c.Close(); err != nil { + t.Fatalf("c.Close() error: %v", err) + } + + select { + case err := <-serveErrCh: + if err != nil { + t.Fatalf("Serve returned non-nil error on Close: %v", err) + } + case <-time.After(5 * time.Second): + t.Fatal("Serve did not exit after Close") + } +} +