diff --git a/file.go b/file.go index 9700b61..0353160 100644 --- a/file.go +++ b/file.go @@ -71,7 +71,7 @@ func makeWin32File(h syscall.Handle) (*win32File, error) { } func MakeOpenFile(h syscall.Handle) (io.ReadWriteCloser, error) { - return makeWin32File(h) + return makeWin32File(h) } // closeHandle closes the resources associated with a Win32 handle diff --git a/pipe.go b/pipe.go index 32523c0..058a2e1 100644 --- a/pipe.go +++ b/pipe.go @@ -114,8 +114,9 @@ type win32PipeListener struct { firstHandle syscall.Handle path string securityDescriptor string - closeCh chan (chan int) acceptCh chan (chan acceptResponse) + closeCh chan int + doneCh chan int } func makeServerPipeHandle(path, securityDescriptor string, first bool) (syscall.Handle, error) { @@ -157,40 +158,43 @@ func (l *win32PipeListener) makeServerPipe() (*win32Pipe, error) { } func (l *win32PipeListener) listenerRoutine() { - var closeResponseCh (chan int) -Loop: - for { - var responseCh (chan acceptResponse) + closed := false + for !closed { select { - case closeResponseCh = <-l.closeCh: - break Loop - case responseCh = <-l.acceptCh: - } - p, err := l.makeServerPipe() - if err != nil { - responseCh <- acceptResponse{nil, err} - } else { - ch := make(chan error) - go func() { - ch <- connectPipe(p) - }() - select { - case closeResponseCh = <-l.closeCh: - p.Close() - <-ch - break Loop - case err = <-ch: - if err != nil { + case <-l.closeCh: + closed = true + case responseCh := <-l.acceptCh: + p, err := l.makeServerPipe() + if err == nil { + // Wait for the client to connect. + ch := make(chan error) + go func() { + ch <- connectPipe(p) + }() + select { + case err = <-ch: + if err != nil { + p.Close() + p = nil + } + case <-l.closeCh: + // Abort the connect request by closing the handle. p.Close() p = nil + err = <-ch + if err == nil { + err = FileClosed + } + closed = true } - responseCh <- acceptResponse{p, err} } + responseCh <- acceptResponse{p, err} } } syscall.CloseHandle(l.firstHandle) l.firstHandle = syscall.Handle(0) - closeResponseCh <- 1 + // Notify Close() and Accept() callers that the handle has been closed. + close(l.doneCh) } func ListenPipe(path, securityDescriptor string) (net.Listener, error) { @@ -210,8 +214,9 @@ func ListenPipe(path, securityDescriptor string) (net.Listener, error) { firstHandle: h, path: path, securityDescriptor: securityDescriptor, - closeCh: make(chan (chan int)), acceptCh: make(chan (chan acceptResponse)), + closeCh: make(chan int), + doneCh: make(chan int), } go l.listenerRoutine() return l, nil @@ -232,15 +237,21 @@ func connectPipe(p *win32Pipe) error { func (l *win32PipeListener) Accept() (net.Conn, error) { ch := make(chan acceptResponse) - l.acceptCh <- ch - response := <-ch - return response.p, response.err + select { + case l.acceptCh <- ch: + response := <-ch + return response.p, response.err + case <-l.doneCh: + return nil, FileClosed + } } func (l *win32PipeListener) Close() error { - ch := make(chan int) - l.closeCh <- ch - <-ch + select { + case l.closeCh <- 1: + <-l.doneCh + case <-l.doneCh: + } return nil } diff --git a/pipe_test.go b/pipe_test.go index 4e69fcc..2e4df70 100644 --- a/pipe_test.go +++ b/pipe_test.go @@ -93,12 +93,11 @@ func TestReadTimeout(t *testing.T) { } } -func server(l net.Listener) { +func server(l net.Listener, ch chan int) { c, err := l.Accept() if err != nil { panic(err) } - defer c.Close() rw := bufio.NewReadWriter(bufio.NewReader(c), bufio.NewWriter(c)) s, err := rw.ReadString('\n') if err != nil { @@ -112,6 +111,8 @@ func server(l net.Listener) { if err != nil { panic(err) } + c.Close() + ch <- 1 } func TestFullListenDialReadWrite(t *testing.T) { @@ -119,13 +120,16 @@ func TestFullListenDialReadWrite(t *testing.T) { if err != nil { t.Fatal(err) } + defer l.Close() - go server(l) + ch := make(chan int) + go server(l, ch) c, err := DialPipe(testPipeName, nil) if err != nil { t.Fatal(err) } + defer c.Close() rw := bufio.NewReadWriter(bufio.NewReader(c), bufio.NewWriter(c)) _, err = rw.WriteString("hello world\n") @@ -145,4 +149,39 @@ func TestFullListenDialReadWrite(t *testing.T) { if s != ms { t.Errorf("expected '%s', got '%s'", ms, s) } + + <-ch +} + +func TestCloseAbortsListen(t *testing.T) { + l, err := ListenPipe(testPipeName, "") + if err != nil { + t.Fatal(err) + } + + ch := make(chan int) + go func() { + _, err := l.Accept() + if err != syscall.ERROR_OPERATION_ABORTED { + t.Fatalf("expected ERROR_OPERATION_ABORTED, got %v", err) + } + ch <- 1 + }() + + time.Sleep(30 * time.Millisecond) + l.Close() + + <-ch +} + +func TestAcceptAfterCloseFails(t *testing.T) { + l, err := ListenPipe(testPipeName, "") + if err != nil { + t.Fatal(err) + } + l.Close() + _, err = l.Accept() + if err != FileClosed { + t.Fatalf("expected FileClosed, got %v", err) + } }