package api import ( "context" "encoding/json" "io" "net/http" "net/http/httptest" "strings" "sync/atomic" "testing" "time" ) func TestClientIP(t *testing.T) { tests := []struct { name string remoteAddr string headers map[string]string want string }{ {"remote addr only", "192.0.2.7:4711", nil, "192.0.2.7"}, {"remote addr no port", "192.0.2.7", nil, "192.0.2.7"}, {"fly header wins", "127.0.0.1:80", map[string]string{"Fly-Client-IP": "203.0.113.1", "X-Forwarded-For": "198.51.100.1"}, "203.0.113.1"}, {"xff first entry", "127.0.0.1:80", map[string]string{"X-Forwarded-For": " 198.51.100.1 , 10.0.0.1"}, "198.51.100.1"}, {"ipv6 remote", "[::1]:9999", nil, "::1"}, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { r := httptest.NewRequest(http.MethodPost, "/api/traverse", nil) r.RemoteAddr = tc.remoteAddr for k, v := range tc.headers { r.Header.Set(k, v) } if got := clientIP(r); got != tc.want { t.Fatalf("clientIP = %q, want %q", got, tc.want) } }) } } type webhookHit struct { event string body map[string]any } // newWebhookHandler builds a Handler whose traversals fail instantly (the // per-job context is already expired) and whose webhook posts to url. func newWebhookHandler(t *testing.T, url string) *Handler { t.Helper() ctx, cancel := context.WithCancel(context.Background()) t.Cleanup(cancel) h := newHandler(ctx) h.jobTimeout = time.Nanosecond h.webhook = &webhookReporter{ url: url, client: &http.Client{}, timeout: 2 * time.Second, retryDelay: 10 * time.Millisecond, } return h } func TestWebhook_StartAndCompleteEventsDelivered(t *testing.T) { hits := make(chan webhookHit, 4) ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if ct := r.Header.Get("Content-Type"); ct != "application/json" { t.Errorf("Content-Type = %q, want application/json", ct) } raw, _ := io.ReadAll(r.Body) var body map[string]any if err := json.Unmarshal(raw, &body); err != nil { t.Errorf("unmarshal webhook body: %v (%s)", err, raw) } hits <- webhookHit{event: r.Header.Get("X-ExploreDNS-Event"), body: body} })) defer ts.Close() h := newWebhookHandler(t, ts.URL) req := httptest.NewRequest(http.MethodPost, "/api/traverse", strings.NewReader(`{"domain":"example.com","type":"A","all_roots":true}`)) req.Header.Set("X-Forwarded-For", "198.51.100.7") rec := httptest.NewRecorder() h.mux.ServeHTTP(rec, req) if rec.Code != http.StatusAccepted { t.Fatalf("want 202, got %d: %s", rec.Code, rec.Body.String()) } var start TraverseStartResponse if err := json.NewDecoder(rec.Body).Decode(&start); err != nil { t.Fatal(err) } // Both events are fired asynchronously; collect them by event name. got := map[string]map[string]any{} for len(got) < 2 { select { case hit := <-hits: got[hit.event] = hit.body case <-time.After(10 * time.Second): t.Fatalf("timed out waiting for webhook events, have %v", got) } } startEv := got["start"] if startEv == nil { t.Fatal("no start event received") } for k, want := range map[string]any{ "event": "start", "id": start.ID, "domain": "example.com", "query_type": "A", "all_roots": true, "client_ip": "198.51.100.7", } { if startEv[k] != want { t.Errorf("start event %s = %v, want %v", k, startEv[k], want) } } if s, _ := startEv["started_at"].(string); s == "" { t.Error("start event missing started_at") } compEv := got["complete"] if compEv == nil { t.Fatal("no complete event received") } for k, want := range map[string]any{ "event": "complete", "id": start.ID, "domain": "example.com", "query_type": "A", "client_ip": "198.51.100.7", "status": statusError, } { if compEv[k] != want { t.Errorf("complete event %s = %v, want %v", k, compEv[k], want) } } if msg, _ := compEv["error"].(string); !strings.Contains(msg, "timed out") { t.Errorf("complete event error = %v, want timeout message", compEv["error"]) } if _, ok := compEv["duration_ms"].(float64); !ok { t.Errorf("complete event duration_ms = %v, want a number", compEv["duration_ms"]) } if d, _ := compEv["done_at"].(string); d == "" { t.Error("complete event missing done_at") } if _, ok := compEv["result_count"].(float64); !ok { t.Errorf("complete event result_count = %v, want a number", compEv["result_count"]) } if _, ok := compEv["summary"]; !ok { t.Error("complete event missing summary key") } } // TestWebhook_SlowReceiverDoesNotDelayJob verifies that a webhook receiver // stuck for longer than the whole traversal never delays job completion. func TestWebhook_SlowReceiverDoesNotDelayJob(t *testing.T) { release := make(chan struct{}) ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { <-release // hold every delivery until the test finishes })) defer ts.Close() defer close(release) h := newWebhookHandler(t, ts.URL) startedAt := time.Now() req := httptest.NewRequest(http.MethodPost, "/api/traverse", strings.NewReader(`{"domain":"example.com"}`)) rec := httptest.NewRecorder() h.mux.ServeHTTP(rec, req) if rec.Code != http.StatusAccepted { t.Fatalf("want 202, got %d: %s", rec.Code, rec.Body.String()) } var start TraverseStartResponse if err := json.NewDecoder(rec.Body).Decode(&start); err != nil { t.Fatal(err) } job, ok := h.st.get(start.ID) if !ok { t.Fatal("job not found") } deadline := time.After(5 * time.Second) for { job.mu.RLock() done := job.DoneAt != nil job.mu.RUnlock() if done { break } select { case <-deadline: t.Fatal("job did not reach a terminal state while webhook was stalled") case <-time.After(5 * time.Millisecond): } } // The traversal fails instantly (expired context); reaching terminal // state must not have waited on the stalled webhook receiver. if elapsed := time.Since(startedAt); elapsed > 3*time.Second { t.Fatalf("job completion took %s, webhook receiver must not delay it", elapsed) } } func TestWebhook_RetriesOnceOnFailure(t *testing.T) { var calls atomic.Int32 done := make(chan struct{}) ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if calls.Add(1) == 1 { w.WriteHeader(http.StatusInternalServerError) return } close(done) })) defer ts.Close() wr := &webhookReporter{ url: ts.URL, client: &http.Client{}, timeout: 2 * time.Second, retryDelay: 10 * time.Millisecond, } wr.send("start", map[string]string{"event": "start"}) select { case <-done: case <-time.After(5 * time.Second): t.Fatal("webhook was not retried after a failed delivery") } if n := calls.Load(); n != 2 { t.Fatalf("webhook deliveries = %d, want 2 (initial + one retry)", n) } } // TestWebhook_NilReporterSafe covers the disabled (no URL) path. func TestWebhook_NilReporterSafe(t *testing.T) { var wr *webhookReporter wr.send("start", map[string]string{"event": "start"}) // must not panic if newWebhookReporter("", "token") != nil { t.Fatal("empty URL should disable the webhook reporter") } } // TestWebhook_BearerToken verifies the Authorization header is sent exactly // when a token is configured (EXPLOREDNS_WEBHOOK_TOKEN on the real path). func TestWebhook_BearerToken(t *testing.T) { for _, tc := range []struct { name, token, wantAuth string }{ {"token set", "s3cret", "Bearer s3cret"}, {"token unset", "", ""}, } { t.Run(tc.name, func(t *testing.T) { done := make(chan string, 1) ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { done <- r.Header.Get("Authorization") })) defer ts.Close() newWebhookReporter(ts.URL, tc.token).send("start", map[string]string{"event": "start"}) select { case got := <-done: if got != tc.wantAuth { t.Fatalf("Authorization = %q, want %q", got, tc.wantAuth) } case <-time.After(5 * time.Second): t.Fatal("webhook was not delivered") } }) } }