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:
co-authored by
Claude Fable 5
parent
d71c7fbef2
commit
111b8bf48e
+72
-1
@@ -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.
|
||||
|
||||
@@ -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
@@ -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
@@ -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")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user