Dual-dialect store (SQLite via modernc.org, MySQL via go-sql-driver, both pure Go) with order-tolerant start/complete upserts, filtered and paginated listing, and aggregate queries (per-day, top domains, query types, statuses, duration percentiles, top clients). Sender gains optional EXPLOREDNS_WEBHOOK_TOKEN bearer auth; a round-trip test pins receiver structs byte-compatible with the sender payloads. Note: go directive moves to 1.25.0, required by modernc.org/sqlite. CI reads the version from go.mod so GOTOOLCHAIN=local stays satisfied. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
272 lines
7.8 KiB
Go
272 lines
7.8 KiB
Go
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")
|
|
}
|
|
})
|
|
}
|
|
}
|