diff --git a/web/api/handler.go b/web/api/handler.go index 57dfd3e..07f4b87 100644 --- a/web/api/handler.go +++ b/web/api/handler.go @@ -6,6 +6,7 @@ import ( "encoding/json" "fmt" "io/fs" + "net" "net/http" "os" "sort" @@ -16,6 +17,7 @@ import ( "gitea.hansenits.com.au/hits/ExploreDNS/internal/config" 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" ) @@ -38,6 +40,14 @@ const ( 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. type TraverseRequest struct { Domain string `json:"domain"` @@ -106,6 +116,14 @@ type Summary struct { 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. type TraversalJob struct { ID string `json:"id"` @@ -115,13 +133,18 @@ type TraversalJob struct { Results []ResultItem `json:"results,omitempty"` Summary *Summary `json:"summary,omitempty"` Progress []ProgressEvent `json:"progress,omitempty"` + Servers []ServerInfo `json:"servers,omitempty"` Error string `json:"error,omitempty"` StartedAt time.Time `json:"started_at"` DoneAt *time.Time `json:"done_at,omitempty"` - mu sync.RWMutex - subs []chan ProgressEvent - cancel context.CancelFunc + mu sync.RWMutex + subs []chan ProgressEvent + 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 @@ -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. type Handler struct { st *store mux *http.ServeMux jobTimeout time.Duration 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 { + limit, window := parseRateLimit(os.Getenv("EXPLOREDNS_RATE_LIMIT")) h := &Handler{ st: newStore(), mux: http.NewServeMux(), jobTimeout: envDuration("EXPLOREDNS_JOB_TIMEOUT", defaultJobTimeout), 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("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/health", h.health) @@ -255,6 +296,7 @@ func newHandler(ctx context.Context) *Handler { return case <-t.C: h.st.cleanup() + h.limiter.sweep() } } }() @@ -270,11 +312,24 @@ func (h *Handler) registerStatic(sub fs.FS) { // health handles GET /api/health. 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. 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 var req TraverseRequest 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, StartedAt: time.Now(), cancel: cancel, + clientIP: ip, } 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) 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"` Summary *Summary `json:"summary,omitempty"` Progress []ProgressEvent `json:"progress,omitempty"` + Servers []ServerInfo `json:"servers,omitempty"` Error string `json:"error,omitempty"` StartedAt time.Time `json:"started_at"` DoneAt *time.Time `json:"done_at,omitempty"` @@ -348,6 +415,7 @@ func (h *Handler) getTraversal(w http.ResponseWriter, r *http.Request) { Results: job.Results, Summary: job.Summary, Progress: job.Progress, + Servers: job.Servers, Error: job.Error, StartedAt: job.StartedAt, DoneAt: job.DoneAt, @@ -357,6 +425,35 @@ func (h *Handler) getTraversal(w http.ResponseWriter, r *http.Request) { 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). func (h *Handler) streamTraversal(w http.ResponseWriter, r *http.Request) { id := r.PathValue("id") @@ -430,13 +527,21 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st if r := recover(); r != nil { now := time.Now() job.mu.Lock() - job.Status = statusError - job.Error = fmt.Sprintf("panic: %v", r) - job.DoneAt = &now + if job.DoneAt == nil { + job.Status = statusError + job.Error = fmt.Sprintf("panic: %v", r) + job.DoneAt = &now + } 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.closeSubscribers() + h.reportCompletion(job) }() cfg := traverse.DefaultTraverserConfig() @@ -474,6 +579,25 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st tr := traverse.NewTraverser(cfg) 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() job.mu.Lock() defer job.mu.Unlock() @@ -507,6 +631,105 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st 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 // def when unset or unparsable. func envDuration(name string, def time.Duration) time.Duration { diff --git a/web/api/handler_internal_test.go b/web/api/handler_internal_test.go index c4430c6..34ac3b4 100644 --- a/web/api/handler_internal_test.go +++ b/web/api/handler_internal_test.go @@ -2,6 +2,8 @@ package api import ( "context" + "encoding/json" + "net" "net/http/httptest" "strconv" "strings" @@ -147,3 +149,171 @@ func TestSubscribeSnapshot_NoDuplicates(t *testing.T) { 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") + } +} diff --git a/web/api/handler_test.go b/web/api/handler_test.go index c30a06e..18a6361 100644 --- a/web/api/handler_test.go +++ b/web/api/handler_test.go @@ -29,6 +29,7 @@ func newTestServer(t *testing.T) *api.Server { } func TestHealth(t *testing.T) { + t.Setenv("FLY_REGION", "") // ensure region is absent regardless of host env srv := newTestServer(t) defer srv.Shutdown(5 * time.Second) //nolint:errcheck @@ -48,6 +49,55 @@ func TestHealth(t *testing.T) { if body["status"] != "ok" { 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) { @@ -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) { srv := newTestServer(t) 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"`, + `CountryCityServersSoftware guess`, + `/servers`, // fetches the servers endpoint + } { + if !strings.Contains(page, want) { + t.Errorf("index.html missing %q", want) + } + } +} + func TestStaticSPA_FallbackToIndex(t *testing.T) { srv := newTestServer(t) defer srv.Shutdown(5 * time.Second) //nolint:errcheck diff --git a/web/api/ratelimit.go b/web/api/ratelimit.go new file mode 100644 index 0000000..66aa25c --- /dev/null +++ b/web/api/ratelimit.go @@ -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 +} diff --git a/web/api/ratelimit_test.go b/web/api/ratelimit_test.go new file mode 100644 index 0000000..e32aad5 --- /dev/null +++ b/web/api/ratelimit_test.go @@ -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) + } +} diff --git a/web/api/webhook.go b/web/api/webhook.go new file mode 100644 index 0000000..3c8be5f --- /dev/null +++ b/web/api/webhook.go @@ -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 +} diff --git a/web/api/webhook_test.go b/web/api/webhook_test.go new file mode 100644 index 0000000..a01998d --- /dev/null +++ b/web/api/webhook_test.go @@ -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") + } +}