feat(api): rate limiting, webhook telemetry, server fingerprints, region

- per-client-IP token bucket on POST /api/traverse (EXPLOREDNS_RATE_LIMIT,
  default 30/1h; direct localhost exempt, proxied clients are not)
- optional usage webhooks (EXPLOREDNS_WEBHOOK_URL): start/complete JSON
  events, fire-and-forget with 5s timeout + one retry so a dead receiver
  never delays a job
- post-traversal version.bind fingerprinting exposed at
  GET /api/traverse/{id}/servers (pending until ready) and announced via
  an SSE "servers" event; never delays job completion
- /api/health reports the serving Fly region (FLY_REGION) for observing
  anycast routing from a roaming client

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Gary Hansen
2026-07-07 23:31:34 +10:00
co-authored by Claude Fable 5
parent 94e41fe5b5
commit d8ef805a6a
7 changed files with 1271 additions and 7 deletions
+230 -7
View File
@@ -6,6 +6,7 @@ import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"io/fs" "io/fs"
"net"
"net/http" "net/http"
"os" "os"
"sort" "sort"
@@ -16,6 +17,7 @@ import (
"gitea.hansenits.com.au/hits/ExploreDNS/internal/config" "gitea.hansenits.com.au/hits/ExploreDNS/internal/config"
idns "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" idns "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns"
"gitea.hansenits.com.au/hits/ExploreDNS/internal/fingerprint"
"gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse"
) )
@@ -38,6 +40,14 @@ const (
defaultMaxRunningJobs = 8 defaultMaxRunningJobs = 8
) )
// Fingerprinting runs after a traversal reaches a terminal state: bounded
// concurrency across the unique server IPs, with its own overall deadline so
// a timed-out or cancelled job context never blocks the server list.
const (
fingerprintConcurrency = 8
fingerprintTimeout = 15 * time.Second
)
// TraverseRequest is the JSON body for POST /api/traverse. // TraverseRequest is the JSON body for POST /api/traverse.
type TraverseRequest struct { type TraverseRequest struct {
Domain string `json:"domain"` Domain string `json:"domain"`
@@ -106,6 +116,14 @@ type Summary struct {
ByStatus []SummaryStatus `json:"by_status,omitempty"` ByStatus []SummaryStatus `json:"by_status,omitempty"`
} }
// ServerInfo is one (server name, IP) pair queried during a traversal plus
// its version.bind fingerprint ("" when the server didn't answer the probe).
type ServerInfo struct {
Name string `json:"name"`
IP string `json:"ip"`
Version string `json:"version"`
}
// TraversalJob holds all state for a single asynchronous traversal. // TraversalJob holds all state for a single asynchronous traversal.
type TraversalJob struct { type TraversalJob struct {
ID string `json:"id"` ID string `json:"id"`
@@ -115,13 +133,18 @@ type TraversalJob struct {
Results []ResultItem `json:"results,omitempty"` Results []ResultItem `json:"results,omitempty"`
Summary *Summary `json:"summary,omitempty"` Summary *Summary `json:"summary,omitempty"`
Progress []ProgressEvent `json:"progress,omitempty"` Progress []ProgressEvent `json:"progress,omitempty"`
Servers []ServerInfo `json:"servers,omitempty"`
Error string `json:"error,omitempty"` Error string `json:"error,omitempty"`
StartedAt time.Time `json:"started_at"` StartedAt time.Time `json:"started_at"`
DoneAt *time.Time `json:"done_at,omitempty"` DoneAt *time.Time `json:"done_at,omitempty"`
mu sync.RWMutex mu sync.RWMutex
subs []chan ProgressEvent subs []chan ProgressEvent
cancel context.CancelFunc cancel context.CancelFunc
clientIP string
// serversDone flips once the post-traversal fingerprinting step has
// stored Servers (or was skipped); until then GET …/servers is pending.
serversDone bool
} }
// subscribeSnapshot atomically registers a subscriber and snapshots the // subscribeSnapshot atomically registers a subscriber and snapshots the
@@ -224,24 +247,42 @@ func (s *store) cleanup() {
} }
} }
// versionQuerier is the subset of fingerprint.Fingerprinter the handler
// uses; tests substitute a fake so no probes leave the process.
type versionQuerier interface {
Query(ctx context.Context, ip net.IP) string
}
// Handler wires together the HTTP routes and the job store. // Handler wires together the HTTP routes and the job store.
type Handler struct { type Handler struct {
st *store st *store
mux *http.ServeMux mux *http.ServeMux
jobTimeout time.Duration jobTimeout time.Duration
maxRunning int maxRunning int
version string
limiter *rateLimiter
webhook *webhookReporter
// fp fingerprints server IPs after each traversal; shared across jobs
// so its per-IP cache is reused.
fp versionQuerier
} }
func newHandler(ctx context.Context) *Handler { func newHandler(ctx context.Context) *Handler {
limit, window := parseRateLimit(os.Getenv("EXPLOREDNS_RATE_LIMIT"))
h := &Handler{ h := &Handler{
st: newStore(), st: newStore(),
mux: http.NewServeMux(), mux: http.NewServeMux(),
jobTimeout: envDuration("EXPLOREDNS_JOB_TIMEOUT", defaultJobTimeout), jobTimeout: envDuration("EXPLOREDNS_JOB_TIMEOUT", defaultJobTimeout),
maxRunning: envInt("EXPLOREDNS_MAX_JOBS", defaultMaxRunningJobs), maxRunning: envInt("EXPLOREDNS_MAX_JOBS", defaultMaxRunningJobs),
version: "dev",
limiter: newRateLimiter(limit, window),
webhook: newWebhookReporter(os.Getenv("EXPLOREDNS_WEBHOOK_URL")),
fp: fingerprint.New(),
} }
h.mux.HandleFunc("POST /api/traverse", h.startTraversal) h.mux.HandleFunc("POST /api/traverse", h.startTraversal)
h.mux.HandleFunc("GET /api/traverse/{id}/stream", h.streamTraversal) h.mux.HandleFunc("GET /api/traverse/{id}/stream", h.streamTraversal)
h.mux.HandleFunc("GET /api/traverse/{id}/servers", h.getServers)
h.mux.HandleFunc("GET /api/traverse/{id}", h.getTraversal) h.mux.HandleFunc("GET /api/traverse/{id}", h.getTraversal)
h.mux.HandleFunc("GET /api/health", h.health) h.mux.HandleFunc("GET /api/health", h.health)
@@ -255,6 +296,7 @@ func newHandler(ctx context.Context) *Handler {
return return
case <-t.C: case <-t.C:
h.st.cleanup() h.st.cleanup()
h.limiter.sweep()
} }
} }
}() }()
@@ -270,11 +312,24 @@ func (h *Handler) registerStatic(sub fs.FS) {
// health handles GET /api/health. // health handles GET /api/health.
func (h *Handler) health(w http.ResponseWriter, _ *http.Request) { func (h *Handler) health(w http.ResponseWriter, _ *http.Request) {
writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) body := map[string]string{"status": "ok", "version": h.version}
// On Fly.io this identifies which machine served the request —
// useful for observing anycast routing and auto-start behaviour.
if region := os.Getenv("FLY_REGION"); region != "" {
body["region"] = region
}
writeJSON(w, http.StatusOK, body)
} }
// startTraversal handles POST /api/traverse. // startTraversal handles POST /api/traverse.
func (h *Handler) startTraversal(w http.ResponseWriter, r *http.Request) { func (h *Handler) startTraversal(w http.ResponseWriter, r *http.Request) {
ip := clientIP(r)
if h.limiter != nil && !rateLimitExempt(r) && !h.limiter.allow(ip) {
writeError(w, http.StatusTooManyRequests,
fmt.Sprintf("rate limit exceeded: %s per client IP", h.limiter))
return
}
r.Body = http.MaxBytesReader(w, r.Body, 1<<20) // 1 MB limit r.Body = http.MaxBytesReader(w, r.Body, 1<<20) // 1 MB limit
var req TraverseRequest var req TraverseRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil { if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
@@ -310,9 +365,20 @@ func (h *Handler) startTraversal(w http.ResponseWriter, r *http.Request) {
QueryType: queryType, QueryType: queryType,
StartedAt: time.Now(), StartedAt: time.Now(),
cancel: cancel, cancel: cancel,
clientIP: ip,
} }
h.st.set(job) h.st.set(job)
h.webhook.send(webhookEventStart, webhookStartEvent{
Event: webhookEventStart,
ID: job.ID,
Domain: job.Domain,
QueryType: job.QueryType,
AllRoots: req.AllRoots,
ClientIP: ip,
StartedAt: job.StartedAt,
})
go h.runTraversal(ctx, job, req.Domain, qtype, req.AllRoots) go h.runTraversal(ctx, job, req.Domain, qtype, req.AllRoots)
writeJSON(w, http.StatusAccepted, TraverseStartResponse{ID: id, Status: statusRunning}) writeJSON(w, http.StatusAccepted, TraverseStartResponse{ID: id, Status: statusRunning})
@@ -337,6 +403,7 @@ func (h *Handler) getTraversal(w http.ResponseWriter, r *http.Request) {
Results []ResultItem `json:"results,omitempty"` Results []ResultItem `json:"results,omitempty"`
Summary *Summary `json:"summary,omitempty"` Summary *Summary `json:"summary,omitempty"`
Progress []ProgressEvent `json:"progress,omitempty"` Progress []ProgressEvent `json:"progress,omitempty"`
Servers []ServerInfo `json:"servers,omitempty"`
Error string `json:"error,omitempty"` Error string `json:"error,omitempty"`
StartedAt time.Time `json:"started_at"` StartedAt time.Time `json:"started_at"`
DoneAt *time.Time `json:"done_at,omitempty"` DoneAt *time.Time `json:"done_at,omitempty"`
@@ -348,6 +415,7 @@ func (h *Handler) getTraversal(w http.ResponseWriter, r *http.Request) {
Results: job.Results, Results: job.Results,
Summary: job.Summary, Summary: job.Summary,
Progress: job.Progress, Progress: job.Progress,
Servers: job.Servers,
Error: job.Error, Error: job.Error,
StartedAt: job.StartedAt, StartedAt: job.StartedAt,
DoneAt: job.DoneAt, DoneAt: job.DoneAt,
@@ -357,6 +425,35 @@ func (h *Handler) getTraversal(w http.ResponseWriter, r *http.Request) {
writeJSON(w, http.StatusOK, snapshot) writeJSON(w, http.StatusOK, snapshot)
} }
// getServers handles GET /api/traverse/{id}/servers. It answers 202 with a
// pending body while the traversal or the post-traversal fingerprinting is
// still in flight, then the fingerprinted server list.
func (h *Handler) getServers(w http.ResponseWriter, r *http.Request) {
id := r.PathValue("id")
job, ok := h.st.get(id)
if !ok {
writeError(w, http.StatusNotFound, "traversal not found")
return
}
job.mu.RLock()
done := job.serversDone
servers := job.Servers
job.mu.RUnlock()
if !done {
writeJSON(w, http.StatusAccepted, map[string]string{"status": "pending"})
return
}
if servers == nil {
servers = []ServerInfo{}
}
writeJSON(w, http.StatusOK, struct {
Status string `json:"status"`
Servers []ServerInfo `json:"servers"`
}{Status: "complete", Servers: servers})
}
// streamTraversal handles GET /api/traverse/{id}/stream (SSE). // streamTraversal handles GET /api/traverse/{id}/stream (SSE).
func (h *Handler) streamTraversal(w http.ResponseWriter, r *http.Request) { func (h *Handler) streamTraversal(w http.ResponseWriter, r *http.Request) {
id := r.PathValue("id") id := r.PathValue("id")
@@ -430,13 +527,21 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st
if r := recover(); r != nil { if r := recover(); r != nil {
now := time.Now() now := time.Now()
job.mu.Lock() job.mu.Lock()
job.Status = statusError if job.DoneAt == nil {
job.Error = fmt.Sprintf("panic: %v", r) job.Status = statusError
job.DoneAt = &now job.Error = fmt.Sprintf("panic: %v", r)
job.DoneAt = &now
}
job.mu.Unlock() job.mu.Unlock()
} }
// Never leave GET …/servers pending: fingerprinting is skipped on
// the panic path, so flip the flag here (idempotent otherwise).
job.mu.Lock()
job.serversDone = true
job.mu.Unlock()
job.cancel() job.cancel()
job.closeSubscribers() job.closeSubscribers()
h.reportCompletion(job)
}() }()
cfg := traverse.DefaultTraverserConfig() cfg := traverse.DefaultTraverserConfig()
@@ -474,6 +579,25 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st
tr := traverse.NewTraverser(cfg) tr := traverse.NewTraverser(cfg)
root, err := tr.Run(ctx, domain) root, err := tr.Run(ctx, domain)
h.commitResult(ctx, job, root, err)
// Tell streaming clients the traversal reached a terminal state so they
// can fetch results now; the stream stays open for the servers event
// published once fingerprinting (below) finishes.
job.mu.Lock()
ev := ProgressEvent{Stage: "complete", Status: job.Status}
job.Progress = append(job.Progress, ev)
job.publishLocked(ev)
job.mu.Unlock()
// Fingerprint the servers queried during the traversal. This runs after
// the terminal status is committed, so results never wait on versions.
h.fingerprintServers(job, tr.ServersEncountered())
}
// commitResult stores the traversal outcome and moves the job to its
// terminal status.
func (h *Handler) commitResult(ctx context.Context, job *TraversalJob, root *traverse.Referral, err error) {
now := time.Now() now := time.Now()
job.mu.Lock() job.mu.Lock()
defer job.mu.Unlock() defer job.mu.Unlock()
@@ -507,6 +631,105 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st
job.Status = statusComplete job.Status = statusComplete
} }
// fingerprintServers turns the traversal's (server, ip) pairs into
// job.Servers, probing each unique IP's version.bind with bounded
// concurrency, then publishes a {"stage":"servers"} event so streaming
// clients know the list is ready without polling. Pseudo "key:" entries and
// non-address entries are skipped.
func (h *Handler) fingerprintServers(job *TraversalJob, seen map[string][]string) {
type pair struct{ name, ip string }
var pairs []pair
uniq := make(map[string]bool)
var ips []net.IP
for name, addrs := range seen {
for _, addr := range addrs {
if strings.HasPrefix(addr, "key:") {
continue
}
ip := net.ParseIP(addr)
if ip == nil {
continue
}
pairs = append(pairs, pair{name: name, ip: addr})
if !uniq[addr] {
uniq[addr] = true
ips = append(ips, ip)
}
}
}
versions := make(map[string]string, len(ips))
if len(ips) > 0 && h.fp != nil {
ctx, cancel := context.WithTimeout(context.Background(), fingerprintTimeout)
defer cancel()
var (
wg sync.WaitGroup
mu sync.Mutex
sem = make(chan struct{}, fingerprintConcurrency)
)
for _, ip := range ips {
wg.Add(1)
go func(ip net.IP) {
defer wg.Done()
sem <- struct{}{}
defer func() { <-sem }()
v := h.fp.Query(ctx, ip)
mu.Lock()
versions[ip.String()] = v
mu.Unlock()
}(ip)
}
wg.Wait()
}
servers := make([]ServerInfo, 0, len(pairs))
for _, p := range pairs {
servers = append(servers, ServerInfo{Name: p.name, IP: p.ip, Version: versions[p.ip]})
}
sort.Slice(servers, func(i, j int) bool {
if servers[i].Name != servers[j].Name {
return servers[i].Name < servers[j].Name
}
return servers[i].IP < servers[j].IP
})
ev := ProgressEvent{Stage: "servers"}
job.mu.Lock()
job.Servers = servers
job.serversDone = true
job.Progress = append(job.Progress, ev)
job.publishLocked(ev)
job.mu.Unlock()
}
// reportCompletion posts the webhook "complete" event for a job that has
// reached a terminal state. Fire-and-forget; never blocks the caller.
func (h *Handler) reportCompletion(job *TraversalJob) {
if h.webhook == nil {
return
}
job.mu.RLock()
ev := webhookCompleteEvent{
Event: webhookEventComplete,
ID: job.ID,
Domain: job.Domain,
QueryType: job.QueryType,
ClientIP: job.clientIP,
StartedAt: job.StartedAt,
Status: job.Status,
Error: job.Error,
ResultCount: len(job.Results),
Summary: job.Summary,
}
if job.DoneAt != nil {
ev.DoneAt = *job.DoneAt
ev.DurationMS = job.DoneAt.Sub(job.StartedAt).Milliseconds()
}
job.mu.RUnlock()
h.webhook.send(webhookEventComplete, ev)
}
// envDuration reads a Go duration from the environment, falling back to // envDuration reads a Go duration from the environment, falling back to
// def when unset or unparsable. // def when unset or unparsable.
func envDuration(name string, def time.Duration) time.Duration { func envDuration(name string, def time.Duration) time.Duration {
+170
View File
@@ -2,6 +2,8 @@ package api
import ( import (
"context" "context"
"encoding/json"
"net"
"net/http/httptest" "net/http/httptest"
"strconv" "strconv"
"strings" "strings"
@@ -147,3 +149,171 @@ func TestSubscribeSnapshot_NoDuplicates(t *testing.T) {
t.Error(msg) t.Error(msg)
} }
} }
// fakeVersionQuerier returns canned version strings without touching the
// network.
type fakeVersionQuerier struct {
mu sync.Mutex
versions map[string]string
queried []string
}
func (f *fakeVersionQuerier) Query(_ context.Context, ip net.IP) string {
f.mu.Lock()
defer f.mu.Unlock()
f.queried = append(f.queried, ip.String())
return f.versions[ip.String()]
}
// TestGetServers_PendingWhileRunning verifies the 202 pending shape while a
// job has not finished fingerprinting (running or just-completed).
func TestGetServers_PendingWhileRunning(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
h := newHandler(ctx)
now := time.Now()
jobs := []*TraversalJob{
{ID: "running", Status: statusRunning, StartedAt: now},
{ID: "fingerprinting", Status: statusComplete, StartedAt: now, DoneAt: &now},
}
for _, j := range jobs {
h.st.set(j)
}
for _, id := range []string{"running", "fingerprinting"} {
req := httptest.NewRequest("GET", "/api/traverse/"+id+"/servers", nil)
rec := httptest.NewRecorder()
h.mux.ServeHTTP(rec, req)
if rec.Code != 202 {
t.Fatalf("%s: want 202, got %d: %s", id, rec.Code, rec.Body.String())
}
var body map[string]string
if err := json.NewDecoder(rec.Body).Decode(&body); err != nil {
t.Fatalf("%s: decode: %v", id, err)
}
if body["status"] != "pending" {
t.Fatalf("%s: want status=pending, got %q", id, body["status"])
}
}
}
// TestFingerprintServers_StoresServersAndPublishes drives the fingerprint
// step with a fake querier: (server, ip) pairs become sorted job.Servers with
// versions, pseudo/invalid entries are skipped, a {"stage":"servers"} event
// is published, and the endpoint flips from pending to the final list.
func TestFingerprintServers_StoresServersAndPublishes(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
h := newHandler(ctx)
fake := &fakeVersionQuerier{versions: map[string]string{
"192.0.2.1": "TestDNS 1.0",
"192.0.2.2": "",
}}
h.fp = fake
now := time.Now()
job := &TraversalJob{ID: "j", Status: statusComplete, StartedAt: now, DoneAt: &now}
h.st.set(job)
sub, _, _ := job.subscribeSnapshot()
defer job.unsubscribe(sub)
h.fingerprintServers(job, map[string][]string{
"b.example.net": {"192.0.2.2"},
"a.example.net": {"192.0.2.1", "key:pseudo:entry", "not-an-ip"},
})
want := []ServerInfo{
{Name: "a.example.net", IP: "192.0.2.1", Version: "TestDNS 1.0"},
{Name: "b.example.net", IP: "192.0.2.2", Version: ""},
}
job.mu.RLock()
got := append([]ServerInfo(nil), job.Servers...)
done := job.serversDone
job.mu.RUnlock()
if !done {
t.Fatal("serversDone not set")
}
if len(got) != len(want) {
t.Fatalf("servers = %+v, want %+v", got, want)
}
for i := range want {
if got[i] != want[i] {
t.Fatalf("servers[%d] = %+v, want %+v", i, got[i], want[i])
}
}
select {
case ev := <-sub:
if ev.Stage != "servers" {
t.Fatalf("published stage = %q, want servers", ev.Stage)
}
default:
t.Fatal("no servers event published")
}
// Endpoint now serves the final list.
req := httptest.NewRequest("GET", "/api/traverse/j/servers", nil)
rec := httptest.NewRecorder()
h.mux.ServeHTTP(rec, req)
if rec.Code != 200 {
t.Fatalf("want 200, got %d: %s", rec.Code, rec.Body.String())
}
var body struct {
Status string `json:"status"`
Servers []ServerInfo `json:"servers"`
}
if err := json.NewDecoder(rec.Body).Decode(&body); err != nil {
t.Fatal(err)
}
if body.Status != "complete" || len(body.Servers) != 2 {
t.Fatalf("body = %+v", body)
}
// Snapshot includes servers too.
req = httptest.NewRequest("GET", "/api/traverse/j", nil)
rec = httptest.NewRecorder()
h.mux.ServeHTTP(rec, req)
if rec.Code != 200 {
t.Fatalf("snapshot: want 200, got %d", rec.Code)
}
var snap struct {
Servers []ServerInfo `json:"servers"`
}
if err := json.NewDecoder(rec.Body).Decode(&snap); err != nil {
t.Fatal(err)
}
if len(snap.Servers) != 2 {
t.Fatalf("snapshot servers = %+v, want 2 entries", snap.Servers)
}
}
// TestFingerprintServers_EmptySeen still terminates the pending state and
// publishes the servers event for traversals that recorded no servers.
func TestFingerprintServers_EmptySeen(t *testing.T) {
h := &Handler{st: newStore()}
job := &TraversalJob{ID: "e", Status: statusError}
h.st.set(job)
sub, _, _ := job.subscribeSnapshot()
defer job.unsubscribe(sub)
h.fingerprintServers(job, nil)
job.mu.RLock()
defer job.mu.RUnlock()
if !job.serversDone {
t.Fatal("serversDone not set")
}
if len(job.Servers) != 0 {
t.Fatalf("servers = %+v, want empty", job.Servers)
}
select {
case ev := <-sub:
if ev.Stage != "servers" {
t.Fatalf("published stage = %q, want servers", ev.Stage)
}
default:
t.Fatal("no servers event published")
}
}
+204
View File
@@ -29,6 +29,7 @@ func newTestServer(t *testing.T) *api.Server {
} }
func TestHealth(t *testing.T) { func TestHealth(t *testing.T) {
t.Setenv("FLY_REGION", "") // ensure region is absent regardless of host env
srv := newTestServer(t) srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck defer srv.Shutdown(5 * time.Second) //nolint:errcheck
@@ -48,6 +49,55 @@ func TestHealth(t *testing.T) {
if body["status"] != "ok" { if body["status"] != "ok" {
t.Fatalf("want status=ok, got %q", body["status"]) t.Fatalf("want status=ok, got %q", body["status"])
} }
if body["version"] != "dev" {
t.Fatalf("want version=dev, got %q", body["version"])
}
if region, ok := body["region"]; ok {
t.Fatalf("region should be omitted outside Fly, got %q", region)
}
}
func TestHealthReportsFlyRegion(t *testing.T) {
t.Setenv("FLY_REGION", "syd")
srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck
resp, err := http.Get("http://" + srv.Addr() + "/api/health")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
var body map[string]string
if err := json.NewDecoder(resp.Body).Decode(&body); err != nil {
t.Fatal(err)
}
if body["region"] != "syd" {
t.Fatalf("want region=syd, got %q", body["region"])
}
}
func TestHealthReportsStampedVersion(t *testing.T) {
srv := api.NewServer("127.0.0.1:0")
srv.SetVersion("v1.2.3")
if err := srv.Start(); err != nil {
t.Fatalf("start server: %v", err)
}
defer srv.Shutdown(5 * time.Second) //nolint:errcheck
resp, err := http.Get("http://" + srv.Addr() + "/api/health")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
var body map[string]string
if err := json.NewDecoder(resp.Body).Decode(&body); err != nil {
t.Fatal(err)
}
if body["version"] != "v1.2.3" {
t.Fatalf("want version=v1.2.3, got %q", body["version"])
}
} }
func TestCORSDisabledByDefault(t *testing.T) { func TestCORSDisabledByDefault(t *testing.T) {
@@ -205,6 +255,95 @@ func TestGetTraversal_Found(t *testing.T) {
} }
} }
func TestGetServers_NotFound(t *testing.T) {
srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck
resp, err := http.Get("http://" + srv.Addr() + "/api/traverse/does-not-exist/servers")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusNotFound {
t.Fatalf("want 404, got %d", resp.StatusCode)
}
}
// TestGetServers_AvailableAfterCompletion drives a full job through the API:
// the servers endpoint answers 202 pending while the traversal/fingerprinting
// is in flight and the fingerprinted list once everything finished.
func TestGetServers_AvailableAfterCompletion(t *testing.T) {
srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck
startBody := bytes.NewBufferString(`{"domain":"example.com","type":"A"}`)
startResp, err := http.Post("http://"+srv.Addr()+"/api/traverse", "application/json", startBody)
if err != nil {
t.Fatal(err)
}
defer startResp.Body.Close()
var start struct {
ID string `json:"id"`
}
if err := json.NewDecoder(startResp.Body).Decode(&start); err != nil {
t.Fatal(err)
}
deadline := time.Now().Add(90 * time.Second)
for {
resp, err := http.Get("http://" + srv.Addr() + "/api/traverse/" + start.ID + "/servers")
if err != nil {
t.Fatal(err)
}
body, err := io.ReadAll(resp.Body)
resp.Body.Close()
if err != nil {
t.Fatal(err)
}
switch resp.StatusCode {
case http.StatusAccepted:
var pending struct {
Status string `json:"status"`
}
if err := json.Unmarshal(body, &pending); err != nil {
t.Fatalf("pending body: %v (%s)", err, body)
}
if pending.Status != "pending" {
t.Fatalf("want status=pending, got %q", pending.Status)
}
case http.StatusOK:
var done struct {
Status string `json:"status"`
Servers []struct {
Name string `json:"name"`
IP string `json:"ip"`
Version string `json:"version"`
} `json:"servers"`
}
if err := json.Unmarshal(body, &done); err != nil {
t.Fatalf("servers body: %v (%s)", err, body)
}
if done.Status != "complete" {
t.Fatalf("want status=complete, got %q", done.Status)
}
if done.Servers == nil {
t.Fatalf("servers key missing or null: %s", body)
}
return
default:
t.Fatalf("unexpected status %d: %s", resp.StatusCode, body)
}
if time.Now().After(deadline) {
t.Fatal("timed out waiting for servers to become available")
}
time.Sleep(250 * time.Millisecond)
}
}
func TestStreamTraversal_NotFound(t *testing.T) { func TestStreamTraversal_NotFound(t *testing.T) {
srv := newTestServer(t) srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck defer srv.Shutdown(5 * time.Second) //nolint:errcheck
@@ -322,6 +461,71 @@ func TestStaticSPA_TypeOptions(t *testing.T) {
} }
} }
// TestStaticSPA_DetailTree asserts the SPA ships the live detail tree with
// its resolve-subtree toggle markup, plus the raw-log fallback feed so the
// old flat progress view is still reachable for debugging.
func TestStaticSPA_DetailTree(t *testing.T) {
srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck
resp, err := http.Get("http://" + srv.Addr() + "/")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
t.Fatal(err)
}
page := string(body)
for _, want := range []string{
`id="detailTree"`, // detail-tree container
`resolve-toggle`, // per-node show/hide resolve markup
`show resolve`, // toggle wording mirrors dns.squish.net
`id="progressFeed"`, // raw-log fallback feed still present
`id="rawToggle"`, // toggle that reveals it
} {
if !strings.Contains(page, want) {
t.Errorf("index.html missing %q", want)
}
}
}
// TestStaticSPA_ServersSection asserts the SPA ships the server map/table
// section: Leaflet lazy-loaded from unpkg, geojs.io client-side geolocation,
// and the reference-style table headings.
func TestStaticSPA_ServersSection(t *testing.T) {
srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck
resp, err := http.Get("http://" + srv.Addr() + "/")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
t.Fatal(err)
}
page := string(body)
for _, want := range []string{
`unpkg.com/leaflet@1.9`, // map library CDN
`get.geojs.io`, // client-side geolocation service
`id="serversCard"`,
`id="serverMap"`,
`<th>Country</th><th>City</th><th>Servers</th><th>Software guess</th>`,
`/servers`, // fetches the servers endpoint
} {
if !strings.Contains(page, want) {
t.Errorf("index.html missing %q", want)
}
}
}
func TestStaticSPA_FallbackToIndex(t *testing.T) { func TestStaticSPA_FallbackToIndex(t *testing.T) {
srv := newTestServer(t) srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck defer srv.Shutdown(5 * time.Second) //nolint:errcheck
+120
View File
@@ -0,0 +1,120 @@
package api
import (
"fmt"
"net"
"net/http"
"strconv"
"strings"
"sync"
"time"
)
// Default per-IP rate limit for POST /api/traverse, overridable via
// EXPLOREDNS_RATE_LIMIT ("N/duration", e.g. "30/1h" or "10/10m").
const (
defaultRateLimitCount = 30
defaultRateLimitWindow = time.Hour
)
// rateLimiter is an in-memory token-bucket limiter keyed by client IP.
// Each bucket starts full with limit tokens and refills continuously at
// limit tokens per window.
type rateLimiter struct {
limit int
window time.Duration
mu sync.Mutex
buckets map[string]*tokenBucket
now func() time.Time // overridable in tests
}
type tokenBucket struct {
tokens float64
last time.Time
}
func newRateLimiter(limit int, window time.Duration) *rateLimiter {
return &rateLimiter{
limit: limit,
window: window,
buckets: make(map[string]*tokenBucket),
now: time.Now,
}
}
// allow reports whether one request from ip fits within the limit,
// consuming a token when it does.
func (rl *rateLimiter) allow(ip string) bool {
rl.mu.Lock()
defer rl.mu.Unlock()
now := rl.now()
b, ok := rl.buckets[ip]
if !ok {
b = &tokenBucket{tokens: float64(rl.limit), last: now}
rl.buckets[ip] = b
} else {
refill := now.Sub(b.last).Seconds() * float64(rl.limit) / rl.window.Seconds()
b.tokens = min(b.tokens+refill, float64(rl.limit))
b.last = now
}
if b.tokens < 1 {
return false
}
b.tokens--
return true
}
// sweep drops buckets idle for at least one full window; such buckets
// would be full again anyway, so dropping them loses nothing.
func (rl *rateLimiter) sweep() {
rl.mu.Lock()
defer rl.mu.Unlock()
cutoff := rl.now().Add(-rl.window)
for ip, b := range rl.buckets {
if b.last.Before(cutoff) {
delete(rl.buckets, ip)
}
}
}
// String renders the limit for error messages, e.g. "30 requests per 1h0m0s".
func (rl *rateLimiter) String() string {
return fmt.Sprintf("%d requests per %s", rl.limit, rl.window)
}
// rateLimitExempt reports whether r may bypass the rate limit: direct
// connections from loopback (the SPA dev loop, tests, health tooling).
// Proxied requests are never exempt — when Fly-Client-IP or
// X-Forwarded-For is present, RemoteAddr is just the proxy, so the real
// client IP must be limited even though the socket peer is local.
func rateLimitExempt(r *http.Request) bool {
if r.Header.Get("Fly-Client-IP") != "" || r.Header.Get("X-Forwarded-For") != "" {
return false
}
host, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
host = r.RemoteAddr
}
ip := net.ParseIP(host)
return ip != nil && ip.IsLoopback()
}
// parseRateLimit parses "N/duration" (e.g. "30/1h"), falling back to the
// defaults when v is empty or invalid.
func parseRateLimit(v string) (int, time.Duration) {
parts := strings.SplitN(v, "/", 2)
if len(parts) != 2 {
return defaultRateLimitCount, defaultRateLimitWindow
}
n, err := strconv.Atoi(strings.TrimSpace(parts[0]))
if err != nil || n <= 0 {
return defaultRateLimitCount, defaultRateLimitWindow
}
d, err := time.ParseDuration(strings.TrimSpace(parts[1]))
if err != nil || d <= 0 {
return defaultRateLimitCount, defaultRateLimitWindow
}
return n, d
}
+169
View File
@@ -0,0 +1,169 @@
package api
import (
"context"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
)
func TestParseRateLimit(t *testing.T) {
tests := []struct {
in string
limit int
window time.Duration
}{
{"", defaultRateLimitCount, defaultRateLimitWindow},
{"30/1h", 30, time.Hour},
{"10/10m", 10, 10 * time.Minute},
{"5 / 30s", 5, 30 * time.Second},
{"bogus", defaultRateLimitCount, defaultRateLimitWindow},
{"0/1h", defaultRateLimitCount, defaultRateLimitWindow},
{"-3/1h", defaultRateLimitCount, defaultRateLimitWindow},
{"10/-1h", defaultRateLimitCount, defaultRateLimitWindow},
{"10/soon", defaultRateLimitCount, defaultRateLimitWindow},
{"/1h", defaultRateLimitCount, defaultRateLimitWindow},
}
for _, tc := range tests {
limit, window := parseRateLimit(tc.in)
if limit != tc.limit || window != tc.window {
t.Errorf("parseRateLimit(%q) = %d, %s; want %d, %s",
tc.in, limit, window, tc.limit, tc.window)
}
}
}
func TestNewHandlerReadsRateLimitEnv(t *testing.T) {
t.Setenv("EXPLOREDNS_RATE_LIMIT", "5/10m")
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
h := newHandler(ctx)
if h.limiter.limit != 5 || h.limiter.window != 10*time.Minute {
t.Fatalf("limiter = %d/%s, want 5/10m", h.limiter.limit, h.limiter.window)
}
}
// postTraverse sends POST /api/traverse with an empty JSON body so requests
// that pass the rate limiter fail validation (400) instead of spawning a
// real traversal. remoteAddr and headers shape the client identity.
func postTraverse(t *testing.T, h *Handler, remoteAddr string, headers map[string]string) *httptest.ResponseRecorder {
t.Helper()
req := httptest.NewRequest(http.MethodPost, "/api/traverse", strings.NewReader(`{}`))
req.RemoteAddr = remoteAddr
for k, v := range headers {
req.Header.Set(k, v)
}
rec := httptest.NewRecorder()
h.mux.ServeHTTP(rec, req)
return rec
}
func TestRateLimit_OverLimitReturns429(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
h := newHandler(ctx)
h.limiter = newRateLimiter(2, time.Hour)
hdr := map[string]string{"X-Forwarded-For": "203.0.113.9"}
for i := 0; i < 2; i++ {
if rec := postTraverse(t, h, "10.0.0.1:1234", hdr); rec.Code != http.StatusBadRequest {
t.Fatalf("request %d: want 400 (under limit), got %d: %s", i, rec.Code, rec.Body.String())
}
}
rec := postTraverse(t, h, "10.0.0.1:1234", hdr)
if rec.Code != http.StatusTooManyRequests {
t.Fatalf("want 429 over limit, got %d: %s", rec.Code, rec.Body.String())
}
if body := rec.Body.String(); !strings.Contains(body, "2 requests per 1h0m0s") {
t.Fatalf("429 body should name the limit, got %s", body)
}
// A distinct client IP has its own bucket and is unaffected.
other := map[string]string{"X-Forwarded-For": "203.0.113.10"}
if rec := postTraverse(t, h, "10.0.0.1:1234", other); rec.Code != http.StatusBadRequest {
t.Fatalf("distinct IP: want 400, got %d: %s", rec.Code, rec.Body.String())
}
}
func TestRateLimit_RefillsContinuously(t *testing.T) {
rl := newRateLimiter(2, time.Second)
now := time.Now()
rl.now = func() time.Time { return now }
if !rl.allow("a") || !rl.allow("a") {
t.Fatal("first two requests should be allowed")
}
if rl.allow("a") {
t.Fatal("third request should be denied")
}
// Half a window refills half the bucket: one token.
now = now.Add(500 * time.Millisecond)
if !rl.allow("a") {
t.Fatal("request after refill should be allowed")
}
if rl.allow("a") {
t.Fatal("bucket should hold only the refilled token")
}
}
func TestRateLimit_SweepDropsIdleBuckets(t *testing.T) {
rl := newRateLimiter(1, time.Minute)
now := time.Now()
rl.now = func() time.Time { return now }
rl.allow("stale")
now = now.Add(2 * time.Minute)
rl.allow("fresh")
rl.sweep()
rl.mu.Lock()
defer rl.mu.Unlock()
if _, ok := rl.buckets["stale"]; ok {
t.Fatal("idle bucket should have been swept")
}
if _, ok := rl.buckets["fresh"]; !ok {
t.Fatal("active bucket should survive the sweep")
}
}
func TestRateLimit_LocalhostExempt(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
h := newHandler(ctx)
h.limiter = newRateLimiter(1, time.Hour)
for _, addr := range []string{"127.0.0.1:5555", "[::1]:5555"} {
for i := 0; i < 3; i++ {
rec := postTraverse(t, h, addr, nil)
if rec.Code != http.StatusBadRequest {
t.Fatalf("%s request %d: localhost should be exempt, got %d: %s",
addr, i, rec.Code, rec.Body.String())
}
}
}
}
// TestRateLimit_ProxiedLocalhostNotExempt verifies that a request arriving
// from a local proxy (RemoteAddr loopback) is still limited by the real
// client IP carried in Fly-Client-IP / X-Forwarded-For.
func TestRateLimit_ProxiedLocalhostNotExempt(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
h := newHandler(ctx)
h.limiter = newRateLimiter(1, time.Hour)
hdr := map[string]string{"Fly-Client-IP": "198.51.100.4"}
if rec := postTraverse(t, h, "127.0.0.1:5555", hdr); rec.Code != http.StatusBadRequest {
t.Fatalf("first proxied request: want 400, got %d", rec.Code)
}
if rec := postTraverse(t, h, "127.0.0.1:5555", hdr); rec.Code != http.StatusTooManyRequests {
t.Fatalf("second proxied request: want 429, got %d", rec.Code)
}
// The same proxy forwarding a different client is unaffected.
other := map[string]string{"Fly-Client-IP": "198.51.100.5"}
if rec := postTraverse(t, h, "127.0.0.1:5555", other); rec.Code != http.StatusBadRequest {
t.Fatalf("other client via proxy: want 400, got %d", rec.Code)
}
}
+137
View File
@@ -0,0 +1,137 @@
package api
import (
"bytes"
"context"
"encoding/json"
"fmt"
"log"
"net"
"net/http"
"strings"
"time"
)
// Webhook event names, sent both in the JSON body and in the
// X-ExploreDNS-Event request header.
const (
webhookEventStart = "start"
webhookEventComplete = "complete"
)
// webhookStartEvent is posted when a traversal job is accepted.
type webhookStartEvent struct {
Event string `json:"event"`
ID string `json:"id"`
Domain string `json:"domain"`
QueryType string `json:"query_type"`
AllRoots bool `json:"all_roots"`
ClientIP string `json:"client_ip"`
StartedAt time.Time `json:"started_at"`
}
// webhookCompleteEvent is posted when a traversal job reaches a terminal
// state. Summary reuses the API Summary shape.
type webhookCompleteEvent struct {
Event string `json:"event"`
ID string `json:"id"`
Domain string `json:"domain"`
QueryType string `json:"query_type"`
ClientIP string `json:"client_ip"`
StartedAt time.Time `json:"started_at"`
DoneAt time.Time `json:"done_at"`
DurationMS int64 `json:"duration_ms"`
Status string `json:"status"`
Error string `json:"error,omitempty"`
ResultCount int `json:"result_count"`
Summary *Summary `json:"summary"`
}
// webhookReporter posts usage events to a configured URL. Sends are
// fire-and-forget: each runs in its own goroutine with a timeout and a
// single retry, and failures are logged but never surface to callers.
type webhookReporter struct {
url string
client *http.Client
timeout time.Duration
retryDelay time.Duration
}
// newWebhookReporter returns a reporter for url, or nil when url is empty
// (webhook reporting disabled). A nil reporter is safe to call.
func newWebhookReporter(url string) *webhookReporter {
if url == "" {
return nil
}
return &webhookReporter{
url: url,
client: &http.Client{},
timeout: 5 * time.Second,
retryDelay: 2 * time.Second,
}
}
// send marshals payload and posts it asynchronously with one retry on
// failure. Errors are logged and never affect the caller.
func (wr *webhookReporter) send(event string, payload any) {
if wr == nil {
return
}
body, err := json.Marshal(payload)
if err != nil {
log.Printf("webhook: marshal %s event: %v", event, err)
return
}
go func() {
err := wr.post(event, body)
if err == nil {
return
}
log.Printf("webhook: %s event failed, retrying in %s: %v", event, wr.retryDelay, err)
time.Sleep(wr.retryDelay)
if err := wr.post(event, body); err != nil {
log.Printf("webhook: %s event failed after retry: %v", event, err)
}
}()
}
// post performs one synchronous webhook delivery attempt.
func (wr *webhookReporter) post(event string, body []byte) error {
ctx, cancel := context.WithTimeout(context.Background(), wr.timeout)
defer cancel()
req, err := http.NewRequestWithContext(ctx, http.MethodPost, wr.url, bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-ExploreDNS-Event", event)
resp, err := wr.client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode > 299 {
return fmt.Errorf("webhook returned status %d", resp.StatusCode)
}
return nil
}
// clientIP resolves the requesting client's IP: the Fly-Client-IP header if
// present, else the first entry of X-Forwarded-For, else the host part of
// RemoteAddr.
func clientIP(r *http.Request) string {
if ip := strings.TrimSpace(r.Header.Get("Fly-Client-IP")); ip != "" {
return ip
}
if xff := r.Header.Get("X-Forwarded-For"); xff != "" {
if first := strings.TrimSpace(strings.Split(xff, ",")[0]); first != "" {
return first
}
}
if host, _, err := net.SplitHostPort(r.RemoteAddr); err == nil {
return host
}
return r.RemoteAddr
}
+241
View File
@@ -0,0 +1,241 @@
package api
import (
"context"
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
)
func TestClientIP(t *testing.T) {
tests := []struct {
name string
remoteAddr string
headers map[string]string
want string
}{
{"remote addr only", "192.0.2.7:4711", nil, "192.0.2.7"},
{"remote addr no port", "192.0.2.7", nil, "192.0.2.7"},
{"fly header wins", "127.0.0.1:80",
map[string]string{"Fly-Client-IP": "203.0.113.1", "X-Forwarded-For": "198.51.100.1"},
"203.0.113.1"},
{"xff first entry", "127.0.0.1:80",
map[string]string{"X-Forwarded-For": " 198.51.100.1 , 10.0.0.1"},
"198.51.100.1"},
{"ipv6 remote", "[::1]:9999", nil, "::1"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
r := httptest.NewRequest(http.MethodPost, "/api/traverse", nil)
r.RemoteAddr = tc.remoteAddr
for k, v := range tc.headers {
r.Header.Set(k, v)
}
if got := clientIP(r); got != tc.want {
t.Fatalf("clientIP = %q, want %q", got, tc.want)
}
})
}
}
type webhookHit struct {
event string
body map[string]any
}
// newWebhookHandler builds a Handler whose traversals fail instantly (the
// per-job context is already expired) and whose webhook posts to url.
func newWebhookHandler(t *testing.T, url string) *Handler {
t.Helper()
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
h := newHandler(ctx)
h.jobTimeout = time.Nanosecond
h.webhook = &webhookReporter{
url: url,
client: &http.Client{},
timeout: 2 * time.Second,
retryDelay: 10 * time.Millisecond,
}
return h
}
func TestWebhook_StartAndCompleteEventsDelivered(t *testing.T) {
hits := make(chan webhookHit, 4)
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if ct := r.Header.Get("Content-Type"); ct != "application/json" {
t.Errorf("Content-Type = %q, want application/json", ct)
}
raw, _ := io.ReadAll(r.Body)
var body map[string]any
if err := json.Unmarshal(raw, &body); err != nil {
t.Errorf("unmarshal webhook body: %v (%s)", err, raw)
}
hits <- webhookHit{event: r.Header.Get("X-ExploreDNS-Event"), body: body}
}))
defer ts.Close()
h := newWebhookHandler(t, ts.URL)
req := httptest.NewRequest(http.MethodPost, "/api/traverse",
strings.NewReader(`{"domain":"example.com","type":"A","all_roots":true}`))
req.Header.Set("X-Forwarded-For", "198.51.100.7")
rec := httptest.NewRecorder()
h.mux.ServeHTTP(rec, req)
if rec.Code != http.StatusAccepted {
t.Fatalf("want 202, got %d: %s", rec.Code, rec.Body.String())
}
var start TraverseStartResponse
if err := json.NewDecoder(rec.Body).Decode(&start); err != nil {
t.Fatal(err)
}
// Both events are fired asynchronously; collect them by event name.
got := map[string]map[string]any{}
for len(got) < 2 {
select {
case hit := <-hits:
got[hit.event] = hit.body
case <-time.After(10 * time.Second):
t.Fatalf("timed out waiting for webhook events, have %v", got)
}
}
startEv := got["start"]
if startEv == nil {
t.Fatal("no start event received")
}
for k, want := range map[string]any{
"event": "start", "id": start.ID, "domain": "example.com",
"query_type": "A", "all_roots": true, "client_ip": "198.51.100.7",
} {
if startEv[k] != want {
t.Errorf("start event %s = %v, want %v", k, startEv[k], want)
}
}
if s, _ := startEv["started_at"].(string); s == "" {
t.Error("start event missing started_at")
}
compEv := got["complete"]
if compEv == nil {
t.Fatal("no complete event received")
}
for k, want := range map[string]any{
"event": "complete", "id": start.ID, "domain": "example.com",
"query_type": "A", "client_ip": "198.51.100.7", "status": statusError,
} {
if compEv[k] != want {
t.Errorf("complete event %s = %v, want %v", k, compEv[k], want)
}
}
if msg, _ := compEv["error"].(string); !strings.Contains(msg, "timed out") {
t.Errorf("complete event error = %v, want timeout message", compEv["error"])
}
if _, ok := compEv["duration_ms"].(float64); !ok {
t.Errorf("complete event duration_ms = %v, want a number", compEv["duration_ms"])
}
if d, _ := compEv["done_at"].(string); d == "" {
t.Error("complete event missing done_at")
}
if _, ok := compEv["result_count"].(float64); !ok {
t.Errorf("complete event result_count = %v, want a number", compEv["result_count"])
}
if _, ok := compEv["summary"]; !ok {
t.Error("complete event missing summary key")
}
}
// TestWebhook_SlowReceiverDoesNotDelayJob verifies that a webhook receiver
// stuck for longer than the whole traversal never delays job completion.
func TestWebhook_SlowReceiverDoesNotDelayJob(t *testing.T) {
release := make(chan struct{})
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
<-release // hold every delivery until the test finishes
}))
defer ts.Close()
defer close(release)
h := newWebhookHandler(t, ts.URL)
startedAt := time.Now()
req := httptest.NewRequest(http.MethodPost, "/api/traverse",
strings.NewReader(`{"domain":"example.com"}`))
rec := httptest.NewRecorder()
h.mux.ServeHTTP(rec, req)
if rec.Code != http.StatusAccepted {
t.Fatalf("want 202, got %d: %s", rec.Code, rec.Body.String())
}
var start TraverseStartResponse
if err := json.NewDecoder(rec.Body).Decode(&start); err != nil {
t.Fatal(err)
}
job, ok := h.st.get(start.ID)
if !ok {
t.Fatal("job not found")
}
deadline := time.After(5 * time.Second)
for {
job.mu.RLock()
done := job.DoneAt != nil
job.mu.RUnlock()
if done {
break
}
select {
case <-deadline:
t.Fatal("job did not reach a terminal state while webhook was stalled")
case <-time.After(5 * time.Millisecond):
}
}
// The traversal fails instantly (expired context); reaching terminal
// state must not have waited on the stalled webhook receiver.
if elapsed := time.Since(startedAt); elapsed > 3*time.Second {
t.Fatalf("job completion took %s, webhook receiver must not delay it", elapsed)
}
}
func TestWebhook_RetriesOnceOnFailure(t *testing.T) {
var calls atomic.Int32
done := make(chan struct{})
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if calls.Add(1) == 1 {
w.WriteHeader(http.StatusInternalServerError)
return
}
close(done)
}))
defer ts.Close()
wr := &webhookReporter{
url: ts.URL,
client: &http.Client{},
timeout: 2 * time.Second,
retryDelay: 10 * time.Millisecond,
}
wr.send("start", map[string]string{"event": "start"})
select {
case <-done:
case <-time.After(5 * time.Second):
t.Fatal("webhook was not retried after a failed delivery")
}
if n := calls.Load(); n != 2 {
t.Fatalf("webhook deliveries = %d, want 2 (initial + one retry)", n)
}
}
// TestWebhook_NilReporterSafe covers the disabled (no URL) path.
func TestWebhook_NilReporterSafe(t *testing.T) {
var wr *webhookReporter
wr.send("start", map[string]string{"event": "start"}) // must not panic
if newWebhookReporter("") != nil {
t.Fatal("empty URL should disable the webhook reporter")
}
}