Skip to content

Commit a75178f

Browse files
fix: synchronize CLI watchers and server shutdown
1 parent a1f8dec commit a75178f

3 files changed

Lines changed: 75 additions & 10 deletions

File tree

‎cmd/golanggraph/auto_serve_command.go‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -137,6 +137,7 @@ func runAutoServe(cmd *cobra.Command, args []string) error {
137137
// regenerating routes is intentional: endpoint generation mutates route and
138138
// agent maps and is unsafe once requests may be in flight.
139139
func runAutoServeWithContext(ctx context.Context, out io.Writer, config *server.AutoServerConfig, opts autoServeOptions) error {
140+
out = synchronizeWriter(out)
140141
var changes <-chan struct{}
141142
var watchErrors <-chan error
142143
var stopWatching func()

‎cmd/golanggraph/main.go‎

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@ import (
1919
"path/filepath"
2020
"sort"
2121
"strings"
22+
"sync"
2223
"syscall"
2324
"time"
2425

@@ -48,6 +49,28 @@ var (
4849
verbose bool
4950
)
5051

52+
// synchronizedWriter makes command output safe when a command's watcher and
53+
// serving goroutine both report progress. CLI callers commonly use bytes.Buffer
54+
// in tests, but the wrapper also protects any non-concurrent io.Writer supplied
55+
// by an embedding application.
56+
type synchronizedWriter struct {
57+
mu sync.Mutex
58+
w io.Writer
59+
}
60+
61+
func (w *synchronizedWriter) Write(p []byte) (int, error) {
62+
w.mu.Lock()
63+
defer w.mu.Unlock()
64+
return w.w.Write(p)
65+
}
66+
67+
func synchronizeWriter(w io.Writer) io.Writer {
68+
if _, ok := w.(*synchronizedWriter); ok {
69+
return w
70+
}
71+
return &synchronizedWriter{w: w}
72+
}
73+
5174
// rootCmd represents the base command when called without any subcommands
5275
var rootCmd = &cobra.Command{
5376
Use: "golanggraph",
@@ -545,6 +568,7 @@ type serverOptions struct {
545568
// printed "Server started on host:port" -- so a failure to bind was announced
546569
// as a success. And the dev command's own flags were ignored (see devCmd).
547570
func runServer(ctx context.Context, out io.Writer, opts serverOptions) error {
571+
out = synchronizeWriter(out)
548572
if opts.Port <= 0 || opts.Port > 65535 {
549573
return fmt.Errorf("invalid port %d", opts.Port)
550574
}

‎pkg/server/server.go‎

Lines changed: 50 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -70,6 +70,7 @@ type Server struct {
7070
config *ServerConfig
7171
router *mux.Router
7272
server *http.Server
73+
serverMu sync.RWMutex
7374
logger *logrus.Logger
7475
upgrader websocket.Upgrader
7576

@@ -91,6 +92,12 @@ type Server struct {
9192
// several clients can observe the same agent or graph at once.
9293
wsConnections map[string]map[*websocket.Conn]struct{}
9394
wsConnectionsMu sync.RWMutex
95+
96+
// httpConnections records connection state so Stop can close sockets that
97+
// have been accepted but have not yet reached a request. net/http otherwise
98+
// gives StateNew connections a five-second grace period during Shutdown.
99+
httpConnections map[net.Conn]http.ConnState
100+
httpConnectionsMu sync.Mutex
94101
}
95102

96103
// NewServer creates a new server
@@ -104,12 +111,13 @@ func NewServer(config *ServerConfig) *Server {
104111
}
105112

106113
server := &Server{
107-
config: config,
108-
router: mux.NewRouter(),
109-
logger: logrus.New(),
110-
graphManager: NewGraphManager(),
111-
wsConnections: make(map[string]map[*websocket.Conn]struct{}),
112-
startedAt: time.Now(),
114+
config: config,
115+
router: mux.NewRouter(),
116+
logger: logrus.New(),
117+
graphManager: NewGraphManager(),
118+
wsConnections: make(map[string]map[*websocket.Conn]struct{}),
119+
httpConnections: make(map[net.Conn]http.ConnState),
120+
startedAt: time.Now(),
113121
}
114122

115123
// Reject WebSocket upgrades from origins the API does not allow. Accepting
@@ -136,6 +144,16 @@ func NewServer(config *ServerConfig) *Server {
136144
return server
137145
}
138146

147+
func (s *Server) trackHTTPConnection(conn net.Conn, state http.ConnState) {
148+
s.httpConnectionsMu.Lock()
149+
defer s.httpConnectionsMu.Unlock()
150+
if state == http.StateClosed || state == http.StateHijacked {
151+
delete(s.httpConnections, conn)
152+
return
153+
}
154+
s.httpConnections[conn] = state
155+
}
156+
139157
// SetCheckpointer attaches the checkpointer used to serve thread history.
140158
func (s *Server) SetCheckpointer(cp persistence.Checkpointer) {
141159
s.checkpointer = cp
@@ -297,20 +315,24 @@ func (s *Server) setupRoutes() {
297315

298316
// Start starts the server
299317
func (s *Server) Start() error {
300-
s.server = &http.Server{
318+
httpServer := &http.Server{
301319
Addr: fmt.Sprintf("%s:%d", s.config.Host, s.config.Port),
302320
Handler: s.router,
303321
ReadTimeout: s.config.ReadTimeout,
304322
WriteTimeout: s.config.WriteTimeout,
305323
MaxHeaderBytes: s.config.MaxHeaderBytes,
324+
ConnState: s.trackHTTPConnection,
306325
}
326+
s.serverMu.Lock()
327+
s.server = httpServer
328+
s.serverMu.Unlock()
307329

308330
s.logger.WithFields(logrus.Fields{
309331
"host": s.config.Host,
310332
"port": s.config.Port,
311333
}).Info("Starting GoLangGraph server")
312334

313-
if err := s.server.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
335+
if err := httpServer.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
314336
return err
315337
}
316338
return nil
@@ -331,10 +353,28 @@ func (s *Server) Stop(ctx context.Context) error {
331353
}
332354
s.wsConnectionsMu.Unlock()
333355

334-
if s.server == nil {
356+
// Shutdown intentionally waits for active requests. A connection that has
357+
// not reached a request is not active work, but net/http treats StateNew as
358+
// active for five seconds; close those sockets so a stop/restart is prompt.
359+
s.httpConnectionsMu.Lock()
360+
newConnections := make([]net.Conn, 0)
361+
for conn, state := range s.httpConnections {
362+
if state == http.StateNew {
363+
newConnections = append(newConnections, conn)
364+
}
365+
}
366+
s.httpConnectionsMu.Unlock()
367+
for _, conn := range newConnections {
368+
_ = conn.Close()
369+
}
370+
371+
s.serverMu.RLock()
372+
httpServer := s.server
373+
s.serverMu.RUnlock()
374+
if httpServer == nil {
335375
return nil
336376
}
337-
return s.server.Shutdown(ctx)
377+
return httpServer.Shutdown(ctx)
338378
}
339379

340380
// Middleware

0 commit comments

Comments
 (0)