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(``) 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"`, `