@@ -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
6667type 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
137119func (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.
225243func (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
236261func (c * commandConn ) LocalAddr () net.Addr {
0 commit comments