Skip to content
Open
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
10 changes: 7 additions & 3 deletions cmd/root.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
}
Expand Down
18 changes: 17 additions & 1 deletion internal/proxy/proxy.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ package proxy

import (
"context"
"errors"
"fmt"
"io"
"net"
Expand Down Expand Up @@ -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 {
Expand All @@ -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.
Expand Down Expand Up @@ -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
Expand Down
39 changes: 39 additions & 0 deletions internal/proxy/proxy_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -940,3 +940,42 @@
})
}
}

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")
}
}

Check failure on line 981 in internal/proxy/proxy_test.go

View workflow job for this annotation

GitHub Actions / run lint

File is not properly formatted (goimports)
Loading