package api_test import ( "bufio" "bytes" "encoding/json" "fmt" "net/http" "net/http/httptest" "strings" "testing" "time" "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 TestCORSPreflight(t *testing.T) { 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 != "*" { 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) } } 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) } }