feat(api): rate limiting, webhook telemetry, server fingerprints, region
- per-client-IP token bucket on POST /api/traverse (EXPLOREDNS_RATE_LIMIT,
default 30/1h; direct localhost exempt, proxied clients are not)
- optional usage webhooks (EXPLOREDNS_WEBHOOK_URL): start/complete JSON
events, fire-and-forget with 5s timeout + one retry so a dead receiver
never delays a job
- post-traversal version.bind fingerprinting exposed at
GET /api/traverse/{id}/servers (pending until ready) and announced via
an SSE "servers" event; never delays job completion
- /api/health reports the serving Fly region (FLY_REGION) for observing
anycast routing from a roaming client
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
94e41fe5b5
commit
d8ef805a6a
+230
-7
@@ -6,6 +6,7 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
"sort"
|
||||
@@ -16,6 +17,7 @@ import (
|
||||
|
||||
"gitea.hansenits.com.au/hits/ExploreDNS/internal/config"
|
||||
idns "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns"
|
||||
"gitea.hansenits.com.au/hits/ExploreDNS/internal/fingerprint"
|
||||
"gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse"
|
||||
)
|
||||
|
||||
@@ -38,6 +40,14 @@ const (
|
||||
defaultMaxRunningJobs = 8
|
||||
)
|
||||
|
||||
// Fingerprinting runs after a traversal reaches a terminal state: bounded
|
||||
// concurrency across the unique server IPs, with its own overall deadline so
|
||||
// a timed-out or cancelled job context never blocks the server list.
|
||||
const (
|
||||
fingerprintConcurrency = 8
|
||||
fingerprintTimeout = 15 * time.Second
|
||||
)
|
||||
|
||||
// TraverseRequest is the JSON body for POST /api/traverse.
|
||||
type TraverseRequest struct {
|
||||
Domain string `json:"domain"`
|
||||
@@ -106,6 +116,14 @@ type Summary struct {
|
||||
ByStatus []SummaryStatus `json:"by_status,omitempty"`
|
||||
}
|
||||
|
||||
// ServerInfo is one (server name, IP) pair queried during a traversal plus
|
||||
// its version.bind fingerprint ("" when the server didn't answer the probe).
|
||||
type ServerInfo struct {
|
||||
Name string `json:"name"`
|
||||
IP string `json:"ip"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
|
||||
// TraversalJob holds all state for a single asynchronous traversal.
|
||||
type TraversalJob struct {
|
||||
ID string `json:"id"`
|
||||
@@ -115,13 +133,18 @@ type TraversalJob struct {
|
||||
Results []ResultItem `json:"results,omitempty"`
|
||||
Summary *Summary `json:"summary,omitempty"`
|
||||
Progress []ProgressEvent `json:"progress,omitempty"`
|
||||
Servers []ServerInfo `json:"servers,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
StartedAt time.Time `json:"started_at"`
|
||||
DoneAt *time.Time `json:"done_at,omitempty"`
|
||||
|
||||
mu sync.RWMutex
|
||||
subs []chan ProgressEvent
|
||||
cancel context.CancelFunc
|
||||
mu sync.RWMutex
|
||||
subs []chan ProgressEvent
|
||||
cancel context.CancelFunc
|
||||
clientIP string
|
||||
// serversDone flips once the post-traversal fingerprinting step has
|
||||
// stored Servers (or was skipped); until then GET …/servers is pending.
|
||||
serversDone bool
|
||||
}
|
||||
|
||||
// subscribeSnapshot atomically registers a subscriber and snapshots the
|
||||
@@ -224,24 +247,42 @@ func (s *store) cleanup() {
|
||||
}
|
||||
}
|
||||
|
||||
// versionQuerier is the subset of fingerprint.Fingerprinter the handler
|
||||
// uses; tests substitute a fake so no probes leave the process.
|
||||
type versionQuerier interface {
|
||||
Query(ctx context.Context, ip net.IP) string
|
||||
}
|
||||
|
||||
// Handler wires together the HTTP routes and the job store.
|
||||
type Handler struct {
|
||||
st *store
|
||||
mux *http.ServeMux
|
||||
jobTimeout time.Duration
|
||||
maxRunning int
|
||||
version string
|
||||
limiter *rateLimiter
|
||||
webhook *webhookReporter
|
||||
// fp fingerprints server IPs after each traversal; shared across jobs
|
||||
// so its per-IP cache is reused.
|
||||
fp versionQuerier
|
||||
}
|
||||
|
||||
func newHandler(ctx context.Context) *Handler {
|
||||
limit, window := parseRateLimit(os.Getenv("EXPLOREDNS_RATE_LIMIT"))
|
||||
h := &Handler{
|
||||
st: newStore(),
|
||||
mux: http.NewServeMux(),
|
||||
jobTimeout: envDuration("EXPLOREDNS_JOB_TIMEOUT", defaultJobTimeout),
|
||||
maxRunning: envInt("EXPLOREDNS_MAX_JOBS", defaultMaxRunningJobs),
|
||||
version: "dev",
|
||||
limiter: newRateLimiter(limit, window),
|
||||
webhook: newWebhookReporter(os.Getenv("EXPLOREDNS_WEBHOOK_URL")),
|
||||
fp: fingerprint.New(),
|
||||
}
|
||||
|
||||
h.mux.HandleFunc("POST /api/traverse", h.startTraversal)
|
||||
h.mux.HandleFunc("GET /api/traverse/{id}/stream", h.streamTraversal)
|
||||
h.mux.HandleFunc("GET /api/traverse/{id}/servers", h.getServers)
|
||||
h.mux.HandleFunc("GET /api/traverse/{id}", h.getTraversal)
|
||||
h.mux.HandleFunc("GET /api/health", h.health)
|
||||
|
||||
@@ -255,6 +296,7 @@ func newHandler(ctx context.Context) *Handler {
|
||||
return
|
||||
case <-t.C:
|
||||
h.st.cleanup()
|
||||
h.limiter.sweep()
|
||||
}
|
||||
}
|
||||
}()
|
||||
@@ -270,11 +312,24 @@ func (h *Handler) registerStatic(sub fs.FS) {
|
||||
|
||||
// health handles GET /api/health.
|
||||
func (h *Handler) health(w http.ResponseWriter, _ *http.Request) {
|
||||
writeJSON(w, http.StatusOK, map[string]string{"status": "ok"})
|
||||
body := map[string]string{"status": "ok", "version": h.version}
|
||||
// On Fly.io this identifies which machine served the request —
|
||||
// useful for observing anycast routing and auto-start behaviour.
|
||||
if region := os.Getenv("FLY_REGION"); region != "" {
|
||||
body["region"] = region
|
||||
}
|
||||
writeJSON(w, http.StatusOK, body)
|
||||
}
|
||||
|
||||
// startTraversal handles POST /api/traverse.
|
||||
func (h *Handler) startTraversal(w http.ResponseWriter, r *http.Request) {
|
||||
ip := clientIP(r)
|
||||
if h.limiter != nil && !rateLimitExempt(r) && !h.limiter.allow(ip) {
|
||||
writeError(w, http.StatusTooManyRequests,
|
||||
fmt.Sprintf("rate limit exceeded: %s per client IP", h.limiter))
|
||||
return
|
||||
}
|
||||
|
||||
r.Body = http.MaxBytesReader(w, r.Body, 1<<20) // 1 MB limit
|
||||
var req TraverseRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
@@ -310,9 +365,20 @@ func (h *Handler) startTraversal(w http.ResponseWriter, r *http.Request) {
|
||||
QueryType: queryType,
|
||||
StartedAt: time.Now(),
|
||||
cancel: cancel,
|
||||
clientIP: ip,
|
||||
}
|
||||
h.st.set(job)
|
||||
|
||||
h.webhook.send(webhookEventStart, webhookStartEvent{
|
||||
Event: webhookEventStart,
|
||||
ID: job.ID,
|
||||
Domain: job.Domain,
|
||||
QueryType: job.QueryType,
|
||||
AllRoots: req.AllRoots,
|
||||
ClientIP: ip,
|
||||
StartedAt: job.StartedAt,
|
||||
})
|
||||
|
||||
go h.runTraversal(ctx, job, req.Domain, qtype, req.AllRoots)
|
||||
|
||||
writeJSON(w, http.StatusAccepted, TraverseStartResponse{ID: id, Status: statusRunning})
|
||||
@@ -337,6 +403,7 @@ func (h *Handler) getTraversal(w http.ResponseWriter, r *http.Request) {
|
||||
Results []ResultItem `json:"results,omitempty"`
|
||||
Summary *Summary `json:"summary,omitempty"`
|
||||
Progress []ProgressEvent `json:"progress,omitempty"`
|
||||
Servers []ServerInfo `json:"servers,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
StartedAt time.Time `json:"started_at"`
|
||||
DoneAt *time.Time `json:"done_at,omitempty"`
|
||||
@@ -348,6 +415,7 @@ func (h *Handler) getTraversal(w http.ResponseWriter, r *http.Request) {
|
||||
Results: job.Results,
|
||||
Summary: job.Summary,
|
||||
Progress: job.Progress,
|
||||
Servers: job.Servers,
|
||||
Error: job.Error,
|
||||
StartedAt: job.StartedAt,
|
||||
DoneAt: job.DoneAt,
|
||||
@@ -357,6 +425,35 @@ func (h *Handler) getTraversal(w http.ResponseWriter, r *http.Request) {
|
||||
writeJSON(w, http.StatusOK, snapshot)
|
||||
}
|
||||
|
||||
// getServers handles GET /api/traverse/{id}/servers. It answers 202 with a
|
||||
// pending body while the traversal or the post-traversal fingerprinting is
|
||||
// still in flight, then the fingerprinted server list.
|
||||
func (h *Handler) getServers(w http.ResponseWriter, r *http.Request) {
|
||||
id := r.PathValue("id")
|
||||
job, ok := h.st.get(id)
|
||||
if !ok {
|
||||
writeError(w, http.StatusNotFound, "traversal not found")
|
||||
return
|
||||
}
|
||||
|
||||
job.mu.RLock()
|
||||
done := job.serversDone
|
||||
servers := job.Servers
|
||||
job.mu.RUnlock()
|
||||
|
||||
if !done {
|
||||
writeJSON(w, http.StatusAccepted, map[string]string{"status": "pending"})
|
||||
return
|
||||
}
|
||||
if servers == nil {
|
||||
servers = []ServerInfo{}
|
||||
}
|
||||
writeJSON(w, http.StatusOK, struct {
|
||||
Status string `json:"status"`
|
||||
Servers []ServerInfo `json:"servers"`
|
||||
}{Status: "complete", Servers: servers})
|
||||
}
|
||||
|
||||
// streamTraversal handles GET /api/traverse/{id}/stream (SSE).
|
||||
func (h *Handler) streamTraversal(w http.ResponseWriter, r *http.Request) {
|
||||
id := r.PathValue("id")
|
||||
@@ -430,13 +527,21 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st
|
||||
if r := recover(); r != nil {
|
||||
now := time.Now()
|
||||
job.mu.Lock()
|
||||
job.Status = statusError
|
||||
job.Error = fmt.Sprintf("panic: %v", r)
|
||||
job.DoneAt = &now
|
||||
if job.DoneAt == nil {
|
||||
job.Status = statusError
|
||||
job.Error = fmt.Sprintf("panic: %v", r)
|
||||
job.DoneAt = &now
|
||||
}
|
||||
job.mu.Unlock()
|
||||
}
|
||||
// Never leave GET …/servers pending: fingerprinting is skipped on
|
||||
// the panic path, so flip the flag here (idempotent otherwise).
|
||||
job.mu.Lock()
|
||||
job.serversDone = true
|
||||
job.mu.Unlock()
|
||||
job.cancel()
|
||||
job.closeSubscribers()
|
||||
h.reportCompletion(job)
|
||||
}()
|
||||
|
||||
cfg := traverse.DefaultTraverserConfig()
|
||||
@@ -474,6 +579,25 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st
|
||||
tr := traverse.NewTraverser(cfg)
|
||||
root, err := tr.Run(ctx, domain)
|
||||
|
||||
h.commitResult(ctx, job, root, err)
|
||||
|
||||
// Tell streaming clients the traversal reached a terminal state so they
|
||||
// can fetch results now; the stream stays open for the servers event
|
||||
// published once fingerprinting (below) finishes.
|
||||
job.mu.Lock()
|
||||
ev := ProgressEvent{Stage: "complete", Status: job.Status}
|
||||
job.Progress = append(job.Progress, ev)
|
||||
job.publishLocked(ev)
|
||||
job.mu.Unlock()
|
||||
|
||||
// Fingerprint the servers queried during the traversal. This runs after
|
||||
// the terminal status is committed, so results never wait on versions.
|
||||
h.fingerprintServers(job, tr.ServersEncountered())
|
||||
}
|
||||
|
||||
// commitResult stores the traversal outcome and moves the job to its
|
||||
// terminal status.
|
||||
func (h *Handler) commitResult(ctx context.Context, job *TraversalJob, root *traverse.Referral, err error) {
|
||||
now := time.Now()
|
||||
job.mu.Lock()
|
||||
defer job.mu.Unlock()
|
||||
@@ -507,6 +631,105 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st
|
||||
job.Status = statusComplete
|
||||
}
|
||||
|
||||
// fingerprintServers turns the traversal's (server, ip) pairs into
|
||||
// job.Servers, probing each unique IP's version.bind with bounded
|
||||
// concurrency, then publishes a {"stage":"servers"} event so streaming
|
||||
// clients know the list is ready without polling. Pseudo "key:" entries and
|
||||
// non-address entries are skipped.
|
||||
func (h *Handler) fingerprintServers(job *TraversalJob, seen map[string][]string) {
|
||||
type pair struct{ name, ip string }
|
||||
var pairs []pair
|
||||
uniq := make(map[string]bool)
|
||||
var ips []net.IP
|
||||
for name, addrs := range seen {
|
||||
for _, addr := range addrs {
|
||||
if strings.HasPrefix(addr, "key:") {
|
||||
continue
|
||||
}
|
||||
ip := net.ParseIP(addr)
|
||||
if ip == nil {
|
||||
continue
|
||||
}
|
||||
pairs = append(pairs, pair{name: name, ip: addr})
|
||||
if !uniq[addr] {
|
||||
uniq[addr] = true
|
||||
ips = append(ips, ip)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
versions := make(map[string]string, len(ips))
|
||||
if len(ips) > 0 && h.fp != nil {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), fingerprintTimeout)
|
||||
defer cancel()
|
||||
|
||||
var (
|
||||
wg sync.WaitGroup
|
||||
mu sync.Mutex
|
||||
sem = make(chan struct{}, fingerprintConcurrency)
|
||||
)
|
||||
for _, ip := range ips {
|
||||
wg.Add(1)
|
||||
go func(ip net.IP) {
|
||||
defer wg.Done()
|
||||
sem <- struct{}{}
|
||||
defer func() { <-sem }()
|
||||
v := h.fp.Query(ctx, ip)
|
||||
mu.Lock()
|
||||
versions[ip.String()] = v
|
||||
mu.Unlock()
|
||||
}(ip)
|
||||
}
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
servers := make([]ServerInfo, 0, len(pairs))
|
||||
for _, p := range pairs {
|
||||
servers = append(servers, ServerInfo{Name: p.name, IP: p.ip, Version: versions[p.ip]})
|
||||
}
|
||||
sort.Slice(servers, func(i, j int) bool {
|
||||
if servers[i].Name != servers[j].Name {
|
||||
return servers[i].Name < servers[j].Name
|
||||
}
|
||||
return servers[i].IP < servers[j].IP
|
||||
})
|
||||
|
||||
ev := ProgressEvent{Stage: "servers"}
|
||||
job.mu.Lock()
|
||||
job.Servers = servers
|
||||
job.serversDone = true
|
||||
job.Progress = append(job.Progress, ev)
|
||||
job.publishLocked(ev)
|
||||
job.mu.Unlock()
|
||||
}
|
||||
|
||||
// reportCompletion posts the webhook "complete" event for a job that has
|
||||
// reached a terminal state. Fire-and-forget; never blocks the caller.
|
||||
func (h *Handler) reportCompletion(job *TraversalJob) {
|
||||
if h.webhook == nil {
|
||||
return
|
||||
}
|
||||
job.mu.RLock()
|
||||
ev := webhookCompleteEvent{
|
||||
Event: webhookEventComplete,
|
||||
ID: job.ID,
|
||||
Domain: job.Domain,
|
||||
QueryType: job.QueryType,
|
||||
ClientIP: job.clientIP,
|
||||
StartedAt: job.StartedAt,
|
||||
Status: job.Status,
|
||||
Error: job.Error,
|
||||
ResultCount: len(job.Results),
|
||||
Summary: job.Summary,
|
||||
}
|
||||
if job.DoneAt != nil {
|
||||
ev.DoneAt = *job.DoneAt
|
||||
ev.DurationMS = job.DoneAt.Sub(job.StartedAt).Milliseconds()
|
||||
}
|
||||
job.mu.RUnlock()
|
||||
h.webhook.send(webhookEventComplete, ev)
|
||||
}
|
||||
|
||||
// envDuration reads a Go duration from the environment, falling back to
|
||||
// def when unset or unparsable.
|
||||
func envDuration(name string, def time.Duration) time.Duration {
|
||||
|
||||
Reference in New Issue
Block a user