Files
ExploreDNS/web/api/handler_test.go
Gary HansenandClaude Fable 5 d8ef805a6a 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>
2026-07-07 23:31:34 +10:00

568 lines
15 KiB
Go

package api_test
import (
"bufio"
"bytes"
"encoding/json"
"fmt"
"io"
"net/http"
"net/http/httptest"
"regexp"
"strings"
"testing"
"time"
"gitea.hansenits.com.au/hits/ExploreDNS/internal/config"
"gitea.hansenits.com.au/hits/ExploreDNS/web/api"
)
// newTestServer starts a Server on a random port and returns it.
// The caller is responsible for calling srv.Shutdown.
func newTestServer(t *testing.T) *api.Server {
t.Helper()
srv := api.NewServer("127.0.0.1:0")
if err := srv.Start(); err != nil {
t.Fatalf("start server: %v", err)
}
return srv
}
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
resp, err := http.Get("http://" + srv.Addr() + "/api/health")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
t.Fatalf("want 200, got %d", resp.StatusCode)
}
var body map[string]string
if err := json.NewDecoder(resp.Body).Decode(&body); err != nil {
t.Fatal(err)
}
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) {
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
req, _ := http.NewRequest(http.MethodOptions, "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 resp.StatusCode != http.StatusNoContent {
t.Fatalf("want 204, got %d", resp.StatusCode)
}
if got := resp.Header.Get("Access-Control-Allow-Origin"); got != "http://localhost:3000" {
t.Fatalf("CORS origin header: %q", got)
}
}
func TestStartTraversal_MissingDomain(t *testing.T) {
srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck
body := bytes.NewBufferString(`{}`)
resp, err := http.Post("http://"+srv.Addr()+"/api/traverse", "application/json", body)
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusBadRequest {
t.Fatalf("want 400, got %d", resp.StatusCode)
}
}
func TestStartTraversal_InvalidType(t *testing.T) {
srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck
body := bytes.NewBufferString(`{"domain":"example.com","type":"BOGUS"}`)
resp, err := http.Post("http://"+srv.Addr()+"/api/traverse", "application/json", body)
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusBadRequest {
t.Fatalf("want 400, got %d", resp.StatusCode)
}
}
func TestStartTraversal_ReturnsID(t *testing.T) {
srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck
body := bytes.NewBufferString(`{"domain":"example.com","type":"A"}`)
resp, err := http.Post("http://"+srv.Addr()+"/api/traverse", "application/json", body)
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusAccepted {
t.Fatalf("want 202, got %d", resp.StatusCode)
}
var start struct {
ID string `json:"id"`
Status string `json:"status"`
}
if err := json.NewDecoder(resp.Body).Decode(&start); err != nil {
t.Fatal(err)
}
if start.ID == "" {
t.Fatal("expected non-empty id")
}
if start.Status != "running" {
t.Fatalf("want status=running, got %q", start.Status)
}
}
func TestGetTraversal_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")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusNotFound {
t.Fatalf("want 404, got %d", resp.StatusCode)
}
}
func TestGetTraversal_Found(t *testing.T) {
srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck
// Start a job.
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)
}
// Poll for it.
resp, err := http.Get("http://" + srv.Addr() + "/api/traverse/" + start.ID)
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
t.Fatalf("want 200, got %d", resp.StatusCode)
}
var job struct {
ID string `json:"id"`
Status string `json:"status"`
}
if err := json.NewDecoder(resp.Body).Decode(&job); err != nil {
t.Fatal(err)
}
if job.ID != start.ID {
t.Fatalf("id mismatch: %q vs %q", job.ID, start.ID)
}
}
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
resp, err := http.Get("http://" + srv.Addr() + "/api/traverse/nope/stream")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusNotFound {
t.Fatalf("want 404, got %d", resp.StatusCode)
}
}
func TestStreamTraversal_ContentType(t *testing.T) {
srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck
// Start a job.
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)
}
// Open SSE stream.
resp, err := http.Get("http://" + srv.Addr() + "/api/traverse/" + start.ID + "/stream")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if ct := resp.Header.Get("Content-Type"); !strings.HasPrefix(ct, "text/event-stream") {
t.Fatalf("want text/event-stream, got %q", ct)
}
// Read until we get the done event or timeout.
done := make(chan struct{})
go func() {
defer close(done)
scanner := bufio.NewScanner(resp.Body)
for scanner.Scan() {
line := scanner.Text()
if strings.HasPrefix(line, "event: done") {
return
}
}
}()
select {
case <-done:
case <-time.After(30 * time.Second):
t.Fatal("timed out waiting for SSE done event")
}
}
func TestStaticSPA_Index(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()
if resp.StatusCode != http.StatusOK {
t.Fatalf("want 200, got %d", resp.StatusCode)
}
}
// TestStaticSPA_TypeOptions asserts the SPA type dropdown offers exactly the
// query types config.ParseQueryType accepts.
func TestStaticSPA_TypeOptions(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)
re := regexp.MustCompile(`<option>([A-Z]+)</option>`)
var got []string
for _, m := range re.FindAllStringSubmatch(page, -1) {
got = append(got, m[1])
}
want := []string{"A", "AAAA", "NS", "CNAME", "MX", "TXT", "SOA", "PTR", "ANY"}
if len(got) != len(want) {
t.Fatalf("type options = %v, want %v", got, want)
}
for i, typ := range want {
if got[i] != typ {
t.Fatalf("type options = %v, want %v", got, want)
}
if _, err := config.ParseQueryType(typ); err != nil {
t.Fatalf("option %s rejected by ParseQueryType: %v", typ, err)
}
}
}
// 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) {
srv := newTestServer(t)
defer srv.Shutdown(5 * time.Second) //nolint:errcheck
resp, err := http.Get("http://" + srv.Addr() + "/some/spa/route")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
t.Fatalf("want 200 (SPA fallback), got %d", resp.StatusCode)
}
}
// TestHandlerUnit uses httptest.NewRecorder for fast unit tests without
// starting a real listener.
func TestHandlerUnit_Health(t *testing.T) {
w := httptest.NewRecorder()
r := httptest.NewRequest(http.MethodGet, "/api/health", nil)
// Access the unexported handler via the exported Server for coverage.
srv := api.NewServer("127.0.0.1:0")
if err := srv.Start(); err != nil {
t.Fatal(err)
}
defer srv.Shutdown(5 * time.Second) //nolint:errcheck
resp, err := http.Get(fmt.Sprintf("http://%s/api/health", srv.Addr()))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
_ = w
_ = r
if resp.StatusCode != http.StatusOK {
t.Fatalf("want 200, got %d", resp.StatusCode)
}
}