package server import ( "context" "net/http" "net/http/httptest" "path/filepath" "strings" "testing" "time" "gitea.hansenits.com.au/hits/ExploreDNS/internal/receiver/store" ) func newTestHandler(t *testing.T, token string) (http.Handler, *store.Store) { t.Helper() st, err := store.OpenSQLite(filepath.Join(t.TempDir(), "receiver.db")) if err != nil { t.Fatalf("OpenSQLite: %v", err) } t.Cleanup(func() { st.Close() }) return newHandler(st, "test-version", token, testAdminUser, testAdminPass), st } func post(h http.Handler, body string, headers map[string]string) *httptest.ResponseRecorder { req := httptest.NewRequest(http.MethodPost, "/webhook", strings.NewReader(body)) req.Header.Set("Content-Type", "application/json") for k, v := range headers { req.Header.Set(k, v) } w := httptest.NewRecorder() h.ServeHTTP(w, req) return w } const startJSON = `{"event":"start","id":"job-1","domain":"example.com","query_type":"A",` + `"all_roots":true,"client_ip":"203.0.113.9","started_at":"2026-07-01T10:00:00Z"}` const completeJSON = `{"event":"complete","id":"job-1","domain":"example.com","query_type":"A",` + `"client_ip":"203.0.113.9","started_at":"2026-07-01T10:00:00Z","done_at":"2026-07-01T10:00:03Z",` + `"duration_ms":3000,"status":"complete","result_count":2,` + `"summary":{"answers":[{"probability":1,"records":["example.com 300 IN A 192.0.2.1"]}]}}` func TestIngestAuthMatrix(t *testing.T) { tests := []struct { name string token string authHeader string want int }{ {"no token configured, no header", "", "", http.StatusNoContent}, {"no token configured, stray header", "", "Bearer whatever", http.StatusNoContent}, {"token configured, missing header", "s3cret", "", http.StatusUnauthorized}, {"token configured, wrong scheme", "s3cret", "Basic s3cret", http.StatusUnauthorized}, {"token configured, wrong token", "s3cret", "Bearer nope", http.StatusUnauthorized}, {"token configured, correct token", "s3cret", "Bearer s3cret", http.StatusNoContent}, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { h, _ := newTestHandler(t, tc.token) headers := map[string]string{} if tc.authHeader != "" { headers["Authorization"] = tc.authHeader } if got := post(h, startJSON, headers).Code; got != tc.want { t.Errorf("status = %d, want %d", got, tc.want) } }) } } func TestIngestStartEvent(t *testing.T) { h, st := newTestHandler(t, "") w := post(h, startJSON, nil) if w.Code != http.StatusNoContent { t.Fatalf("status = %d, want 204 (%s)", w.Code, w.Body) } rows, total, err := st.ListTraversals(context.Background(), store.ListFilter{}, 10, 0) if err != nil { t.Fatalf("ListTraversals: %v", err) } if total != 1 { t.Fatalf("total = %d, want 1", total) } row := rows[0] if row.ID != "job-1" || row.Domain != "example.com" || row.QueryType != "A" || !row.AllRoots || row.ClientIP != "203.0.113.9" || row.Status != store.StatusRunning { t.Errorf("row = %+v", row) } want := time.Date(2026, 7, 1, 10, 0, 0, 0, time.UTC) if !row.StartedAt.Equal(want) { t.Errorf("StartedAt = %v, want %v", row.StartedAt, want) } } func TestIngestCompleteEvent(t *testing.T) { h, st := newTestHandler(t, "") if w := post(h, startJSON, nil); w.Code != http.StatusNoContent { t.Fatalf("start status = %d", w.Code) } if w := post(h, completeJSON, nil); w.Code != http.StatusNoContent { t.Fatalf("complete status = %d", w.Code) } rows, total, err := st.ListTraversals(context.Background(), store.ListFilter{}, 10, 0) if err != nil { t.Fatalf("ListTraversals: %v", err) } if total != 1 { t.Fatalf("total = %d, want 1", total) } row := rows[0] if row.Status != store.StatusComplete { t.Errorf("Status = %q, want complete", row.Status) } if row.DurationMS == nil || *row.DurationMS != 3000 { t.Errorf("DurationMS = %v, want 3000", row.DurationMS) } if row.ResultCount == nil || *row.ResultCount != 2 { t.Errorf("ResultCount = %v, want 2", row.ResultCount) } if !strings.Contains(row.Summary, "192.0.2.1") { t.Errorf("Summary = %q, want raw summary JSON", row.Summary) } if !row.AllRoots { t.Error("AllRoots lost after complete") } } func TestIngestRejectsBadInput(t *testing.T) { tests := []struct { name string body string }{ {"malformed JSON", `{"event":`}, {"unknown event", `{"event":"pause","id":"x"}`}, {"empty event", `{"id":"x"}`}, {"wrong type for field", `{"event":"start","id":42}`}, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { h, st := newTestHandler(t, "") if got := post(h, tc.body, nil).Code; got != http.StatusBadRequest { t.Errorf("status = %d, want 400", got) } _, total, err := st.ListTraversals(context.Background(), store.ListFilter{}, 10, 0) if err != nil { t.Fatalf("ListTraversals: %v", err) } if total != 0 { t.Errorf("stored %d rows from rejected input", total) } }) } } func TestIngestBodyCap(t *testing.T) { h, _ := newTestHandler(t, "") big := `{"event":"start","id":"x","domain":"` + strings.Repeat("a", maxBodyBytes) + `"}` if got := post(h, big, nil).Code; got != http.StatusRequestEntityTooLarge { t.Errorf("status = %d, want 413", got) } } func TestHealthz(t *testing.T) { h, _ := newTestHandler(t, "s3cret") req := httptest.NewRequest(http.MethodGet, "/healthz", nil) w := httptest.NewRecorder() h.ServeHTTP(w, req) if w.Code != http.StatusOK { t.Fatalf("status = %d, want 200", w.Code) } body := w.Body.String() // healthz stays open even when an ingest token is configured. if !strings.Contains(body, `"status":"ok"`) || !strings.Contains(body, `"version":"test-version"`) { t.Errorf("body = %s", body) } } func TestServerStartShutdown(t *testing.T) { st, err := store.OpenSQLite(filepath.Join(t.TempDir(), "receiver.db")) if err != nil { t.Fatalf("OpenSQLite: %v", err) } defer st.Close() srv := New("127.0.0.1:0", st) srv.SetVersion("v-test") srv.SetAdminAuth(testAdminUser, testAdminPass) if err := srv.Start(); err != nil { t.Fatalf("Start: %v", err) } defer srv.Shutdown(2 * time.Second) resp, err := http.Get("http://" + srv.Addr() + "/healthz") if err != nil { t.Fatalf("GET /healthz: %v", err) } resp.Body.Close() if resp.StatusCode != http.StatusOK { t.Errorf("healthz status = %d, want 200", resp.StatusCode) } if err := srv.Shutdown(2 * time.Second); err != nil { t.Errorf("Shutdown: %v", err) } }