7c45a99df6decd946acc78cd8cd364c99d4425ee

Author
Ayman Bagabas <ayman.bagabas@gmail.com>
Committer
GitHub <noreply@github.com>
Date

Message

fix(daemon): close listener only once (#615)

* fix(daemon): close listener only once

* refactor(daemon): rename Start to ListenAndServe and implement Serve

* fix(daemon): use atomic.Bool for server

* fix(daemon): attempt to fix idle timeout test

Diff

This diff is truncated to protect this page.

  1diff --git a/cmd/soft/serve/server.go b/cmd/soft/serve/server.go
  2index 33679091f1f177303b4651e842cdf20b01c0bc81..69d5db5a8fcf40c3c1f09786599303a61cb47c65 100644
  3--- a/cmd/soft/serve/server.go
  4+++ b/cmd/soft/serve/server.go
  5@@ -95,7 +95,7 @@ func (s *Server) Start() error {
  6 	errg, _ := errgroup.WithContext(s.ctx)
  7 	errg.Go(func() error {
  8 		s.logger.Print("Starting Git daemon", "addr", s.Config.Git.ListenAddr)
  9-		if err := s.GitDaemon.Start(); !errors.Is(err, daemon.ErrServerClosed) {
 10+		if err := s.GitDaemon.ListenAndServe(); !errors.Is(err, daemon.ErrServerClosed) {
 11 			return err
 12 		}
 13 		return nil
 14diff --git a/pkg/daemon/daemon.go b/pkg/daemon/daemon.go
 15index de65cba75d3873b5dde46ab096be431fb02b3597..e4e6340d91ea604a347b284ddd238d2d097e2f52 100644
 16--- a/pkg/daemon/daemon.go
 17+++ b/pkg/daemon/daemon.go
 18@@ -8,6 +8,7 @@ import (
 19 	"path/filepath"
 20 	"strings"
 21 	"sync"
 22+	"sync/atomic"
 23 	"time"
 24 
 25 	"github.com/charmbracelet/log"
 26@@ -43,7 +44,6 @@ var ErrServerClosed = fmt.Errorf("git: %w", net.ErrClosed)
 27 // GitDaemon represents a Git daemon.
 28 type GitDaemon struct {
 29 	ctx      context.Context
 30-	listener net.Listener
 31 	addr     string
 32 	finished chan struct{}
 33 	conns    connections
 34@@ -52,6 +52,7 @@ type GitDaemon struct {
 35 	wg       sync.WaitGroup
 36 	once     sync.Once
 37 	logger   *log.Logger
 38+	done     atomic.Bool // indicates if the server has been closed
 39 }
 40 
 41 // NewDaemon returns a new Git daemon.
 42@@ -70,26 +71,31 @@ func NewGitDaemon(ctx context.Context) (*GitDaemon, error) {
 43 	return d, nil
 44 }
 45 
 46-// Start starts the Git TCP daemon.
 47-func (d *GitDaemon) Start() error {
 48-	// listen on the socket
 49-	{
 50-		listener, err := net.Listen("tcp", d.addr)
 51-		if err != nil {
 52-			return err
 53-		}
 54-		d.listener = listener
 55+// ListenAndServe starts the Git TCP daemon.
 56+func (d *GitDaemon) ListenAndServe() error {
 57+	if d.done.Load() {
 58+		return ErrServerClosed
 59+	}
 60+	listener, err := net.Listen("tcp", d.addr)
 61+	if err != nil {
 62+		return err
 63 	}
 64+	return d.Serve(listener)
 65+}
 66 
 67-	// close eventual connections to the socket
 68-	defer d.listener.Close() // nolint: errcheck
 69+// Serve listens on the TCP network address and serves Git requests.
 70+func (d *GitDaemon) Serve(listener net.Listener) error {
 71+	if d.done.Load() {
 72+		return ErrServerClosed
 73+	}
 74 
 75 	d.wg.Add(1)
 76 	defer d.wg.Done()
 77+	defer listener.Close() //nolint:errcheck
 78 
 79 	var tempDelay time.Duration
 80 	for {
 81-		conn, err := d.listener.Accept()
 82+		conn, err := listener.Accept()
 83 		if err != nil {
 84 			select {
 85 			case <-d.finished:
 86@@ -305,21 +311,30 @@ func (d *GitDaemon) handleClient(conn net.Conn) {
 87 
 88 // Close closes the underlying listener.
 89 func (d *GitDaemon) Close() error {
 90-	d.once.Do(func() { close(d.finished) })
 91-	err := d.listener.Close()
 92+	err := d.closeListener()
 93 	d.conns.CloseAll() // nolint: errcheck
 94 	return err
 95 }
 96 
 97+// closeListener closes the listener and the finished channel.
 98+func (d *GitDaemon) closeListener() error {
 99+	if d.done.Load() {
100+		return ErrServerClosed
101+	}
102+	d.once.Do(func() {
103+		close(d.finished)
104+		d.done.Store(true)
105+	})
106+	return nil
107+}
108+
109 // Shutdown gracefully shuts down the daemon.
110 func (d *GitDaemon) Shutdown(ctx context.Context) error {
111-	// in the case when git daemon was never started
112-	if d.listener == nil {
113-		return nil
114+	if d.done.Load() {
115+		return ErrServerClosed
116 	}
117 
118diff --git a/pkg/daemon/daemon_test.go b/pkg/daemon/daemon_test.go
119index 88b4fdec7a9e2de2d4c9621b549c57c269f11ddb..b9cb77c6444707971ef47be9e09b80b2f7e689c4 100644
120--- a/pkg/daemon/daemon_test.go
121+++ b/pkg/daemon/daemon_test.go
122@@ -59,7 +59,7 @@ func TestMain(m *testing.M) {
123 	}
124 	testDaemon = d
125 	go func() {
126-		if err := d.Start(); err != ErrServerClosed {
127+		if err := d.ListenAndServe(); err != ErrServerClosed {
128 			log.Fatal(err)
129 		}
130 	}()
131@@ -75,11 +75,21 @@ func TestMain(m *testing.M) {
132 }
133 
134 func TestIdleTimeout(t *testing.T) {
135-	c, err := net.Dial("tcp", testDaemon.addr)
136-	if err != nil {
137-		t.Fatal(err)
138+	var err error
139+	var c net.Conn
140+	var tries int
141+	for {
142+		c, err = net.Dial("tcp", testDaemon.addr)
143+		if err != nil && tries >= 3 {
144+			t.Fatal(err)
145+		}
146+		tries++
147+		if testDaemon.conns.Size() != 0 {
148+			break
149+		}
150+		time.Sleep(10 * time.Millisecond)
151 	}
152-	time.Sleep(time.Second)
153+	time.Sleep(2 * time.Second)
154 	_, err = readPktline(c)
155 	if err == nil {
156 		t.Errorf("expected error, got nil")