Skip to content

Commit 592fa93

Browse files
committed
commandconn: don't return error if command closed successfully
--- commandconn: fix race on `Close()` During normal operation, if a `Read()` or `Write()` call results in an EOF, we call `onEOF()` to handle the terminating command, and store it's exit value. However, if a Read/Write call was blocked while `Close()` is called the in/out pipes are immediately closed which causes an EOF to be returned. Here, we shouldn't call `onEOF()`, since the reason why we got an EOF is because we're already terminating the connection. This also prevents a race between two calls to the commands `Wait()`, in the `Close()` call and `onEOF()` --- Add CLI init timeout to SSH connections --- connhelper: add 30s ssh default dialer timeout (same as non-ssh dialer) Signed-off-by: Laura Brehm <laurabrehm@hey.com>
1 parent 20923df commit 592fa93

9 files changed

Lines changed: 453 additions & 96 deletions

File tree

‎cli/command/cli.go‎

Lines changed: 2 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,6 @@ import (
88
"path/filepath"
99
"runtime"
1010
"strconv"
11-
"strings"
1211
"sync"
1312
"time"
1413

@@ -327,13 +326,8 @@ func (cli *DockerCli) getInitTimeout() time.Duration {
327326

328327
func (cli *DockerCli) initializeFromClient() {
329328
ctx := context.Background()
330-
if !strings.HasPrefix(cli.dockerEndpoint.Host, "ssh://") {
331-
// @FIXME context.WithTimeout doesn't work with connhelper / ssh connections
332-
// time="2020-04-10T10:16:26Z" level=warning msg="commandConn.CloseWrite: commandconn: failed to wait: signal: killed"
333-
var cancel func()
334-
ctx, cancel = context.WithTimeout(ctx, cli.getInitTimeout())
335-
defer cancel()
336-
}
329+
ctx, cancel := context.WithTimeout(ctx, cli.getInitTimeout())
330+
defer cancel()
337331

338332
ping, err := cli.client.Ping(ctx)
339333
if err != nil {

‎cli/connhelper/commandconn/commandconn.go‎

Lines changed: 113 additions & 88 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@ import (
2424
"runtime"
2525
"strings"
2626
"sync"
27+
"sync/atomic"
2728
"syscall"
2829
"time"
2930

@@ -64,98 +65,79 @@ func New(_ context.Context, cmd string, args ...string) (net.Conn, error) {
6465

6566
// commandConn implements net.Conn
6667
type commandConn struct {
67-
cmd *exec.Cmd
68-
cmdExited bool
69-
cmdWaitErr error
70-
cmdMutex sync.Mutex
71-
stdin io.WriteCloser
72-
stdout io.ReadCloser
73-
stderrMu sync.Mutex
74-
stderr bytes.Buffer
75-
stdioClosedMu sync.Mutex // for stdinClosed and stdoutClosed
76-
stdinClosed bool
77-
stdoutClosed bool
78-
localAddr net.Addr
79-
remoteAddr net.Addr
68+
cmdMutex sync.Mutex // for cmd, cmdWaitErr
69+
cmd *exec.Cmd
70+
cmdWaitErr error
71+
cmdExited atomic.Bool
72+
stdin io.WriteCloser
73+
stdout io.ReadCloser
74+
stderrMu sync.Mutex // for stderr
75+
stderr bytes.Buffer
76+
stdinClosed atomic.Bool
77+
stdoutClosed atomic.Bool
78+
closing atomic.Bool
79+
localAddr net.Addr
80+
remoteAddr net.Addr
8081
}
8182

82-
// killIfStdioClosed kills the cmd if both stdin and stdout are closed.
83-
func (c *commandConn) killIfStdioClosed() error {
84-
c.stdioClosedMu.Lock()
85-
stdioClosed := c.stdoutClosed && c.stdinClosed
86-
c.stdioClosedMu.Unlock()
87-
if !stdioClosed {
83+
// kill returns nil if the command terminated, regardless to the exit status.
84+
func (c *commandConn) kill() error {
85+
if c.cmdExited.Load() {
8886
return nil
8987
}
90-
return c.kill()
91-
}
92-
93-
// killAndWait tries sending SIGTERM to the process before sending SIGKILL.
94-
func killAndWait(cmd *exec.Cmd) error {
88+
c.cmdMutex.Lock()
9589
var werr error
9690
if runtime.GOOS != "windows" {
9791
werrCh := make(chan error)
98-
go func() { werrCh <- cmd.Wait() }()
99-
cmd.Process.Signal(syscall.SIGTERM)
92+
go func() { werrCh <- c.cmd.Wait() }()
93+
c.cmd.Process.Signal(syscall.SIGTERM)
10094
select {
10195
case werr = <-werrCh:
10296
case <-time.After(3 * time.Second):
103-
cmd.Process.Kill()
97+
c.cmd.Process.Kill()
10498
werr = <-werrCh
10599
}
106100
} else {
107-
cmd.Process.Kill()
108-
werr = cmd.Wait()
109-
}
110-
return werr
111-
}
112-
113-
// kill returns nil if the command terminated, regardless to the exit status.
114-
func (c *commandConn) kill() error {
115-
var werr error
116-
c.cmdMutex.Lock()
117-
if c.cmdExited {
118-
werr = c.cmdWaitErr
119-
} else {
120-
werr = killAndWait(c.cmd)
121-
c.cmdWaitErr = werr
122-
c.cmdExited = true
101+
c.cmd.Process.Kill()
102+
werr = c.cmd.Wait()
123103
}
104+
c.cmdWaitErr = werr
124105
c.cmdMutex.Unlock()
125-
if werr == nil {
126-
return nil
127-
}
128-
wExitErr, ok := werr.(*exec.ExitError)
129-
if ok {
130-
if wExitErr.ProcessState.Exited() {
131-
return nil
132-
}
133-
}
134-
return errors.Wrapf(werr, "commandconn: failed to wait")
106+
c.cmdExited.Store(true)
107+
return nil
135108
}
136109

110+
// onEOF gets called if we receive an io.EOF while reading
111+
// or writing from the undelying command pipes.
112+
//
113+
// When we've received an EOF we expect that the command will
114+
// be terminated soon. As such, we call Wait() on the command
115+
// and return EOF or the error depending on whether the command
116+
// exited with an error.
117+
//
118+
// If Wait() does not return within 10s, an error is returned
137119
func (c *commandConn) onEOF(eof error) error {
138-
// when we got EOF, the command is going to be terminated
139-
var werr error
140120
c.cmdMutex.Lock()
141-
if c.cmdExited {
121+
defer c.cmdMutex.Unlock()
122+
123+
var werr error
124+
if c.cmdExited.Load() {
142125
werr = c.cmdWaitErr
143126
} else {
144127
werrCh := make(chan error)
145128
go func() { werrCh <- c.cmd.Wait() }()
146129
select {
147130
case werr = <-werrCh:
148131
c.cmdWaitErr = werr
149-
c.cmdExited = true
132+
c.cmdExited.Store(true)
150133
case <-time.After(10 * time.Second):
151-
c.cmdMutex.Unlock()
152134
c.stderrMu.Lock()
153135
stderr := c.stderr.String()
154136
c.stderrMu.Unlock()
155137
return errors.Errorf("command %v did not exit after %v: stderr=%q", c.cmd.Args, eof, stderr)
156138
}
157139
}
158-
c.cmdMutex.Unlock()
140+
159141
if werr == nil {
160142
return eof
161143
}
@@ -178,59 +160,102 @@ func ignorableCloseError(err error) bool {
178160
return false
179161
}
180162

181-
func (c *commandConn) CloseRead() error {
182-
// NOTE: maybe already closed here
183-
if err := c.stdout.Close(); err != nil && !ignorableCloseError(err) {
184-
logrus.Warnf("commandConn.CloseRead: %v", err)
163+
func (c *commandConn) Read(p []byte) (int, error) {
164+
n, err := c.stdout.Read(p)
165+
// check after the call to Read, since
166+
// it is blocking, and while waiting on it
167+
// Close might get called
168+
if c.closing.Load() {
169+
// If we're currently closing the connection
170+
// we don't want to call onEOF, but we do want
171+
// to return an io.EOF
172+
return 0, io.EOF
185173
}
186-
c.stdioClosedMu.Lock()
187-
c.stdoutClosed = true
188-
c.stdioClosedMu.Unlock()
189-
if err := c.killIfStdioClosed(); err != nil {
190-
logrus.Warnf("commandConn.CloseRead: %v", err)
174+
175+
if err == io.EOF {
176+
err = c.onEOF(err)
191177
}
192-
return nil
178+
return n, err
193179
}
194180

195-
func (c *commandConn) Read(p []byte) (int, error) {
196-
n, err := c.stdout.Read(p)
181+
func (c *commandConn) Write(p []byte) (int, error) {
182+
n, err := c.stdin.Write(p)
183+
// check after the call to Write, since
184+
// it is blocking, and while waiting on it
185+
// Close might get called
186+
if c.closing.Load() {
187+
// If we're currently closing the connection
188+
// we don't want to call onEOF, but we do want
189+
// to return an io.EOF
190+
return 0, io.EOF
191+
}
192+
197193
if err == io.EOF {
198194
err = c.onEOF(err)
199195
}
200196
return n, err
201197
}
202198

203-
func (c *commandConn) CloseWrite() error {
199+
// CloseRead allows commandConn to implement halfCloser
200+
func (c *commandConn) CloseRead() error {
204201
// NOTE: maybe already closed here
205-
if err := c.stdin.Close(); err != nil && !ignorableCloseError(err) {
206-
logrus.Warnf("commandConn.CloseWrite: %v", err)
202+
if err := c.stdout.Close(); err != nil && !ignorableCloseError(err) {
203+
logrus.Warnf("commandConn.CloseRead: %v", err)
204+
return err
207205
}
208-
c.stdioClosedMu.Lock()
209-
c.stdinClosed = true
210-
c.stdioClosedMu.Unlock()
211-
if err := c.killIfStdioClosed(); err != nil {
212-
logrus.Warnf("commandConn.CloseWrite: %v", err)
206+
c.stdoutClosed.Store(true)
207+
208+
if c.stdinClosed.Load() {
209+
if err := c.kill(); err != nil {
210+
logrus.Warnf("commandConn.CloseRead: %v", err)
211+
return err
212+
}
213213
}
214+
214215
return nil
215216
}
216217

217-
func (c *commandConn) Write(p []byte) (int, error) {
218-
n, err := c.stdin.Write(p)
219-
if err == io.EOF {
220-
err = c.onEOF(err)
218+
// CloseWrite allows commandConn to implement halfCloser
219+
func (c *commandConn) CloseWrite() error {
220+
// NOTE: maybe already closed here
221+
if err := c.stdin.Close(); err != nil && !ignorableCloseError(err) {
222+
logrus.Warnf("commandConn.CloseWrite: %v", err)
223+
return err
221224
}
222-
return n, err
225+
c.stdinClosed.Store(true)
226+
227+
if c.stdoutClosed.Load() {
228+
if err := c.kill(); err != nil {
229+
logrus.Warnf("commandConn.CloseWrite: %v", err)
230+
return err
231+
}
232+
}
233+
234+
return nil
223235
}
224236

237+
// Close is the net.Conn func that gets called
238+
// by the transport when a dial is cancelled
239+
// due to it's context timing out. Any blocked
240+
// Read or Write calls will be unblocked and
241+
// return errors. It will block until the underlying
242+
// command has terminated.
225243
func (c *commandConn) Close() error {
226-
var err error
227-
if err = c.CloseRead(); err != nil {
244+
c.closing.Store(true)
245+
defer c.closing.Store(false)
246+
247+
err := c.CloseRead()
248+
if err != nil {
228249
logrus.Warnf("commandConn.Close: CloseRead: %v", err)
250+
return err
229251
}
230-
if err = c.CloseWrite(); err != nil {
252+
err = c.CloseWrite()
253+
if err != nil {
231254
logrus.Warnf("commandConn.Close: CloseWrite: %v", err)
255+
return err
232256
}
233-
return err
257+
258+
return nil
234259
}
235260

236261
func (c *commandConn) LocalAddr() net.Addr {

0 commit comments

Comments
 (0)