- hard per-traversal deadline (EXPLOREDNS_JOB_TIMEOUT, default 5m) so every job reaches a terminal state; timed-out jobs report error with any partial results instead of masquerading as complete - cap concurrent traversals (EXPLOREDNS_MAX_JOBS, default 8) returning 429 when saturated - CORS off by default (the embedded SPA is same-origin); opt in via EXPLOREDNS_CORS_ORIGIN Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
364 lines
9.1 KiB
Go
364 lines
9.1 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) {
|
|
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"])
|
|
}
|
|
}
|
|
|
|
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 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)
|
|
}
|
|
}
|
|
}
|
|
|
|
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)
|
|
}
|
|
}
|