7c45a99df6decd946acc78cd8cd364c99d4425ee
- Author
- Ayman Bagabas <ayman.bagabas@gmail.com>
- Committer
- GitHub <noreply@github.com>
- Date
Message
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")