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 <noreply@anthropic.com>
This commit is contained in:
Gary Hansen
2026-07-07 21:50:29 +10:00
co-authored by Claude Fable 5
parent d71c7fbef2
commit 111b8bf48e
4 changed files with 178 additions and 9 deletions
+72 -1
View File
@@ -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)
@@ -202,12 +228,16 @@ func (s *store) cleanup() {
type Handler struct {
st *store
mux *http.ServeMux
jobTimeout time.Duration
maxRunning int
}
func newHandler(ctx context.Context) *Handler {
h := &Handler{
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.
+72
View File
@@ -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).
+20 -2
View File
@@ -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)
}
}
+10 -2
View File
@@ -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")