From 111b8bf48eff748601506852ab63c8c82163ee73 Mon Sep 17 00:00:00 2001 From: Gary Hansen Date: Tue, 7 Jul 2026 21:50:29 +1000 Subject: [PATCH] feat(web): harden server for public exposure - hard per-traversal deadline (EXPLOREDNS_JOB_TIMEOUT, default 5m) so every job reaches a terminal state; timed-out jobs report error with any partial results instead of masquerading as complete - cap concurrent traversals (EXPLOREDNS_MAX_JOBS, default 8) returning 429 when saturated - CORS off by default (the embedded SPA is same-origin); opt in via EXPLOREDNS_CORS_ORIGIN Co-Authored-By: Claude Fable 5 --- web/api/handler.go | 81 ++++++++++++++++++++++++++++++-- web/api/handler_internal_test.go | 72 ++++++++++++++++++++++++++++ web/api/handler_test.go | 22 ++++++++- web/api/server.go | 12 ++++- 4 files changed, 178 insertions(+), 9 deletions(-) diff --git a/web/api/handler.go b/web/api/handler.go index a3f28ab..57dfd3e 100644 --- a/web/api/handler.go +++ b/web/api/handler.go @@ -7,7 +7,9 @@ import ( "fmt" "io/fs" "net/http" + "os" "sort" + "strconv" "strings" "sync" "time" @@ -27,6 +29,15 @@ const ( // jobTTL is how long completed jobs are retained in memory. const jobTTL = time.Hour +// Public-exposure guards. Every traversal gets a hard deadline so jobs +// always reach a terminal state (the TTL cleanup only purges finished +// jobs), and the number of in-flight traversals is capped. Overridable +// via EXPLOREDNS_JOB_TIMEOUT (Go duration) and EXPLOREDNS_MAX_JOBS. +const ( + defaultJobTimeout = 5 * time.Minute + defaultMaxRunningJobs = 8 +) + // TraverseRequest is the JSON body for POST /api/traverse. type TraverseRequest struct { Domain string `json:"domain"` @@ -183,6 +194,21 @@ func (s *store) set(j *TraversalJob) { s.jobs[j.ID] = j } +// runningCount reports how many jobs have not yet reached a terminal state. +func (s *store) runningCount() int { + s.mu.RLock() + defer s.mu.RUnlock() + n := 0 + for _, j := range s.jobs { + j.mu.RLock() + if j.DoneAt == nil { + n++ + } + j.mu.RUnlock() + } + return n +} + // cleanup removes completed jobs older than jobTTL. func (s *store) cleanup() { cutoff := time.Now().Add(-jobTTL) @@ -200,14 +226,18 @@ func (s *store) cleanup() { // Handler wires together the HTTP routes and the job store. type Handler struct { - st *store - mux *http.ServeMux + st *store + mux *http.ServeMux + jobTimeout time.Duration + maxRunning int } func newHandler(ctx context.Context) *Handler { h := &Handler{ - st: newStore(), - mux: http.NewServeMux(), + st: newStore(), + mux: http.NewServeMux(), + jobTimeout: envDuration("EXPLOREDNS_JOB_TIMEOUT", defaultJobTimeout), + maxRunning: envInt("EXPLOREDNS_MAX_JOBS", defaultMaxRunningJobs), } h.mux.HandleFunc("POST /api/traverse", h.startTraversal) @@ -265,8 +295,13 @@ func (h *Handler) startTraversal(w http.ResponseWriter, r *http.Request) { return } + if h.st.runningCount() >= h.maxRunning { + writeError(w, http.StatusTooManyRequests, "too many concurrent traversals, try again shortly") + return + } + id := newUUID() - ctx, cancel := context.WithCancel(context.Background()) + ctx, cancel := context.WithTimeout(context.Background(), h.jobTimeout) job := &TraversalJob{ ID: id, @@ -445,6 +480,7 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st job.DoneAt = &now + timedOut := ctx.Err() == context.DeadlineExceeded if err != nil && ctx.Err() == nil { job.Status = statusError job.Error = err.Error() @@ -461,9 +497,44 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st job.Summary = toSummary(root.SummaryStats()) } job.Results = items + if timedOut { + // Keep any partial results, but do not report a truncated + // traversal as a successful one. + job.Status = statusError + job.Error = fmt.Sprintf("traversal timed out after %s", h.jobTimeout) + return + } job.Status = statusComplete } +// envDuration reads a Go duration from the environment, falling back to +// def when unset or unparsable. +func envDuration(name string, def time.Duration) time.Duration { + v := os.Getenv(name) + if v == "" { + return def + } + d, err := time.ParseDuration(v) + if err != nil || d <= 0 { + return def + } + return d +} + +// envInt reads a positive integer from the environment, falling back to +// def when unset or unparsable. +func envInt(name string, def int) int { + v := os.Getenv(name) + if v == "" { + return def + } + n, err := strconv.Atoi(v) + if err != nil || n <= 0 { + return def + } + return n +} + // toSummary converts the engine's grouped stats to the API representation: // answers sorted by RRset key (as SummaryStats returns them), remaining // statuses sorted lexically like the CLI Summary Results section. diff --git a/web/api/handler_internal_test.go b/web/api/handler_internal_test.go index 4324c37..c4430c6 100644 --- a/web/api/handler_internal_test.go +++ b/web/api/handler_internal_test.go @@ -1,11 +1,83 @@ package api import ( + "context" + "net/http/httptest" "strconv" + "strings" "sync" "testing" + "time" ) +// TestStartTraversalJobCap verifies that new traversals are rejected with +// 429 once maxRunning jobs are in flight. The store is pre-filled with +// running jobs so no real traversal is spawned. +func TestStartTraversalJobCap(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + h := newHandler(ctx) + h.maxRunning = 2 + + for i := 0; i < 2; i++ { + h.st.set(&TraversalJob{ID: strconv.Itoa(i), Status: statusRunning, StartedAt: time.Now()}) + } + + req := httptest.NewRequest("POST", "/api/traverse", + strings.NewReader(`{"domain":"example.com"}`)) + rec := httptest.NewRecorder() + h.mux.ServeHTTP(rec, req) + + if rec.Code != 429 { + t.Fatalf("want 429 with %d running jobs, got %d: %s", h.st.runningCount(), rec.Code, rec.Body.String()) + } + + // A finished job frees a slot: runningCount must drop below the cap. + done := time.Now() + if j, ok := h.st.get("0"); ok { + j.mu.Lock() + j.Status = statusComplete + j.DoneAt = &done + j.mu.Unlock() + } + if got := h.st.runningCount(); got != 1 { + t.Fatalf("runningCount after completion = %d, want 1", got) + } +} + +// TestJobTimeoutReachesTerminalState verifies that a traversal launched +// with an already-expired deadline still drives the job to a terminal +// error state (the TTL cleanup only ever purges finished jobs). +func TestJobTimeoutReachesTerminalState(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), time.Nanosecond) + defer cancel() + <-ctx.Done() // deadline already exceeded + + h := &Handler{st: newStore(), jobTimeout: time.Nanosecond} + job := &TraversalJob{ID: "t", Status: statusRunning, StartedAt: time.Now(), cancel: cancel} + h.st.set(job) + + doneCh := make(chan struct{}) + go func() { + h.runTraversal(ctx, job, "example.com", 1, false) + close(doneCh) + }() + select { + case <-doneCh: + case <-time.After(30 * time.Second): + t.Fatal("runTraversal did not return with an expired context") + } + + job.mu.RLock() + defer job.mu.RUnlock() + if job.DoneAt == nil { + t.Fatal("job never reached a terminal state") + } + if job.Status != statusError || !strings.Contains(job.Error, "timed out") { + t.Fatalf("want error status mentioning timeout, got status=%q error=%q", job.Status, job.Error) + } +} + // TestSubscribeSnapshot_NoDuplicates regresses the subscribe/snapshot race: // subscribers arriving while events are being published must never see the // same event twice (once from the snapshot replay and once from the channel). diff --git a/web/api/handler_test.go b/web/api/handler_test.go index c0b09ad..c30a06e 100644 --- a/web/api/handler_test.go +++ b/web/api/handler_test.go @@ -50,7 +50,25 @@ func TestHealth(t *testing.T) { } } -func TestCORSPreflight(t *testing.T) { +func TestCORSDisabledByDefault(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + req, _ := http.NewRequest(http.MethodGet, "http://"+srv.Addr()+"/api/health", nil) + req.Header.Set("Origin", "http://localhost:3000") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if got := resp.Header.Get("Access-Control-Allow-Origin"); got != "" { + t.Fatalf("CORS should be off by default, got origin header %q", got) + } +} + +func TestCORSPreflightWithConfiguredOrigin(t *testing.T) { + t.Setenv("EXPLOREDNS_CORS_ORIGIN", "http://localhost:3000") srv := newTestServer(t) defer srv.Shutdown(5 * time.Second) //nolint:errcheck @@ -65,7 +83,7 @@ func TestCORSPreflight(t *testing.T) { if resp.StatusCode != http.StatusNoContent { t.Fatalf("want 204, got %d", resp.StatusCode) } - if got := resp.Header.Get("Access-Control-Allow-Origin"); got != "*" { + if got := resp.Header.Get("Access-Control-Allow-Origin"); got != "http://localhost:3000" { t.Fatalf("CORS origin header: %q", got) } } diff --git a/web/api/server.go b/web/api/server.go index 8d2b2e4..dec400b 100644 --- a/web/api/server.go +++ b/web/api/server.go @@ -18,6 +18,7 @@ import ( "log" "net" "net/http" + "os" "time" ) @@ -93,10 +94,17 @@ func (s *Server) Shutdown(timeout time.Duration) error { return s.srv.Shutdown(ctx) } -// corsMiddleware adds CORS headers for cross-origin SPA access. +// corsMiddleware adds CORS headers for cross-origin API access. The +// embedded SPA is served same-origin and needs none, so CORS is off +// unless EXPLOREDNS_CORS_ORIGIN names an allowed origin (use "*" to +// restore the old allow-all behaviour for development). func corsMiddleware(next http.Handler) http.Handler { + origin := os.Getenv("EXPLOREDNS_CORS_ORIGIN") + if origin == "" { + return next + } return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Access-Control-Allow-Origin", "*") + w.Header().Set("Access-Control-Allow-Origin", origin) w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") w.Header().Set("Access-Control-Allow-Headers", "Content-Type")