diff --git a/web/api/handler.go b/web/api/handler.go index 6edeb15..56e3da4 100644 --- a/web/api/handler.go +++ b/web/api/handler.go @@ -86,10 +86,8 @@ func (j *TraversalJob) subscribe() <-chan ProgressEvent { return ch } -// publish sends an event to all current subscribers. -func (j *TraversalJob) publish(ev ProgressEvent) { - j.mu.Lock() - defer j.mu.Unlock() +// publishLocked sends ev to all current subscribers. Caller must hold j.mu. +func (j *TraversalJob) publishLocked(ev ProgressEvent) { for _, ch := range j.subs { select { case ch <- ev: @@ -99,6 +97,18 @@ func (j *TraversalJob) publish(ev ProgressEvent) { } } +// unsubscribe removes ch from the subscriber list. +func (j *TraversalJob) unsubscribe(ch <-chan ProgressEvent) { + j.mu.Lock() + defer j.mu.Unlock() + for i, s := range j.subs { + if s == ch { + j.subs = append(j.subs[:i], j.subs[i+1:]...) + return + } + } +} + // closeSubscribers drains and closes all subscriber channels. func (j *TraversalJob) closeSubscribers() { j.mu.Lock() @@ -153,7 +163,7 @@ type Handler struct { mux *http.ServeMux } -func newHandler() *Handler { +func newHandler(ctx context.Context) *Handler { h := &Handler{ st: newStore(), mux: http.NewServeMux(), @@ -164,12 +174,17 @@ func newHandler() *Handler { h.mux.HandleFunc("GET /api/traverse/{id}", h.getTraversal) h.mux.HandleFunc("GET /api/health", h.health) - // periodic cleanup + // periodic cleanup; exits when ctx is cancelled (e.g. on Server.Shutdown). go func() { t := time.NewTicker(10 * time.Minute) defer t.Stop() - for range t.C { - h.st.cleanup() + for { + select { + case <-ctx.Done(): + return + case <-t.C: + h.st.cleanup() + } } }() @@ -189,6 +204,7 @@ func (h *Handler) health(w http.ResponseWriter, _ *http.Request) { // startTraversal handles POST /api/traverse. func (h *Handler) startTraversal(w http.ResponseWriter, r *http.Request) { + r.Body = http.MaxBytesReader(w, r.Body, 1<<20) // 1 MB limit var req TraverseRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { writeError(w, http.StatusBadRequest, "invalid request body: "+err.Error()) @@ -285,8 +301,10 @@ func (h *Handler) streamTraversal(w http.ResponseWriter, r *http.Request) { } // Subscribe before snapshotting progress so we don't miss events between - // the two operations. + // the two operations. Unsubscribe when the client disconnects so stale + // channels don't accumulate. sub := job.subscribe() + defer job.unsubscribe(sub) // Replay events already recorded. job.mu.RLock() @@ -381,9 +399,8 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st job.mu.Lock() job.Progress = append(job.Progress, ev) + job.publishLocked(ev) // inside lock: no race with subscribe+replay job.mu.Unlock() - - job.publish(ev) }, } diff --git a/web/api/server.go b/web/api/server.go index 037270f..8d2b2e4 100644 --- a/web/api/server.go +++ b/web/api/server.go @@ -26,8 +26,9 @@ var staticFiles embed.FS // Server is the HTTP API server. type Server struct { - addr string - srv *http.Server + addr string + srv *http.Server + cancel context.CancelFunc } // NewServer creates a new Server that listens on addr (e.g. ":8080"). @@ -39,7 +40,9 @@ func NewServer(addr string) *Server { // server has accepted its first connection or the address is bound. // Call Shutdown to stop gracefully. func (s *Server) Start() error { - h := newHandler() + ctx, cancel := context.WithCancel(context.Background()) + s.cancel = cancel + h := newHandler(ctx) sub, err := fs.Sub(staticFiles, "static") if err != nil { @@ -48,11 +51,12 @@ func (s *Server) Start() error { h.registerStatic(sub) s.srv = &http.Server{ - Addr: s.addr, - Handler: corsMiddleware(h.mux), - ReadTimeout: 30 * time.Second, - WriteTimeout: 0, // SSE streams need no write timeout - IdleTimeout: 120 * time.Second, + Addr: s.addr, + Handler: corsMiddleware(h.mux), + ReadHeaderTimeout: 10 * time.Second, + ReadTimeout: 30 * time.Second, + WriteTimeout: 0, // SSE streams need no write timeout + IdleTimeout: 120 * time.Second, } ln, err := net.Listen("tcp", s.addr) @@ -78,6 +82,9 @@ func (s *Server) Addr() string { // Shutdown gracefully stops the server, waiting up to timeout for in-flight // requests to complete. func (s *Server) Shutdown(timeout time.Duration) error { + if s.cancel != nil { + s.cancel() + } if s.srv == nil { return nil }