diff --git a/.gitea/workflows/release.yml b/.gitea/workflows/release.yml new file mode 100644 index 0000000..c322e0a --- /dev/null +++ b/.gitea/workflows/release.yml @@ -0,0 +1,132 @@ +name: Release + +# Builds release artifacts when a version tag (v*) is pushed: +# - cross-compiled CLI + server binaries attached to the Gitea release +# - version-tagged docker images for the CLI and web server +# The owner usually creates the Gitea release by hand with notes; this +# workflow attaches assets to it (creating a bare release only when none +# exists) and skips assets that are already attached, so re-runs are safe. + +on: + push: + tags: ["v*"] + +jobs: + binaries: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + + - uses: actions/setup-go@v5 + with: + go-version-file: 'go.mod' + cache: true + + - name: Build release binaries + env: + GOTOOLCHAIN: local + run: | + set -euo pipefail + TAG="${GITHUB_REF_NAME}" + mkdir -p dist + for target in linux/amd64 linux/arm64 darwin/amd64 darwin/arm64 windows/amd64; do + GOOS="${target%/*}" + GOARCH="${target#*/}" + EXT="" + [ "$GOOS" = "windows" ] && EXT=".exe" + OUT="dist/${GOOS}_${GOARCH}" + mkdir -p "$OUT" + CGO_ENABLED=0 GOOS="$GOOS" GOARCH="$GOARCH" \ + go build -trimpath -ldflags "-s -w -X main.version=${TAG}" \ + -o "${OUT}/exploredns${EXT}" ./cmd/exploredns + CGO_ENABLED=0 GOOS="$GOOS" GOARCH="$GOARCH" \ + go build -trimpath -ldflags "-s -w -X main.version=${TAG}" \ + -o "${OUT}/exploredns-server${EXT}" ./cmd/server + if [ "$GOOS" = "windows" ]; then + (cd "$OUT" && zip -q "../exploredns_${TAG}_${GOOS}_${GOARCH}.zip" exploredns.exe exploredns-server.exe) + else + tar -czf "dist/exploredns_${TAG}_${GOOS}_${GOARCH}.tar.gz" -C "$OUT" exploredns exploredns-server + fi + rm -rf "$OUT" + done + (cd dist && sha256sum -- * > SHA256SUMS) + ls -l dist + + - name: Attach assets to Gitea release + env: + TOKEN: ${{ secrets.GITHUB_TOKEN }} + API: ${{ github.server_url }}/api/v1 + REPO: ${{ github.repository }} + run: | + set -euo pipefail + TAG="${GITHUB_REF_NAME}" + AUTH="Authorization: token ${TOKEN}" + + # Look up the release for this tag; create a bare one only when + # none exists (the owner writes release notes by hand). + STATUS=$(curl -sS -o release.json -w '%{http_code}' -H "$AUTH" \ + "${API}/repos/${REPO}/releases/tags/${TAG}") + if [ "$STATUS" = "404" ]; then + curl -sS -f -o release.json -H "$AUTH" -H 'Content-Type: application/json' \ + -d "{\"tag_name\":\"${TAG}\",\"name\":\"${TAG}\"}" \ + "${API}/repos/${REPO}/releases" + elif [ "$STATUS" != "200" ]; then + echo "release lookup failed with HTTP ${STATUS}" >&2 + cat release.json >&2 + exit 1 + fi + RELEASE_ID=$(jq -r '.id' release.json) + echo "release id: ${RELEASE_ID}" + + # Existing asset names, so re-runs skip instead of failing. + curl -sS -f -H "$AUTH" \ + "${API}/repos/${REPO}/releases/${RELEASE_ID}/assets" \ + | jq -r '.[].name' > existing.txt + + for f in dist/*; do + NAME=$(basename "$f") + if grep -Fxq "$NAME" existing.txt; then + echo "skip ${NAME} (already attached)" + continue + fi + echo "upload ${NAME}" + curl -sS -f -o /dev/null -H "$AUTH" \ + -F "attachment=@${f}" \ + "${API}/repos/${REPO}/releases/${RELEASE_ID}/assets?name=${NAME}" + done + + docker: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + + - name: Log in to registry + uses: docker/login-action@v3 + with: + registry: gitea.hansenits.com.au + username: ${{ secrets.DOCKER_USERNAME }} + password: ${{ secrets.DOCKER_PASSWORD }} + + - name: Build and push CLI image + uses: docker/build-push-action@v6 + with: + context: . + file: Dockerfile.cli + push: true + build-args: | + VERSION=${{ github.ref_name }} + tags: | + gitea.hansenits.com.au/hits/exploredns-cli:latest + gitea.hansenits.com.au/hits/exploredns-cli:${{ github.ref_name }} + + - name: Build and push web image + uses: docker/build-push-action@v6 + with: + context: . + file: Dockerfile.web + push: true + build-args: | + VERSION=${{ github.ref_name }} + tags: | + gitea.hansenits.com.au/hits/exploredns-web:latest + gitea.hansenits.com.au/hits/exploredns-web:${{ github.ref_name }} diff --git a/Dockerfile.cli b/Dockerfile.cli index 14a05c2..fd40a4c 100644 --- a/Dockerfile.cli +++ b/Dockerfile.cli @@ -8,7 +8,8 @@ RUN go mod download COPY . . -RUN CGO_ENABLED=0 GOOS=linux go build -trimpath -ldflags="-s -w" -o /out/exploredns ./cmd/exploredns +ARG VERSION=dev +RUN CGO_ENABLED=0 GOOS=linux go build -trimpath -ldflags="-s -w -X main.version=${VERSION}" -o /out/exploredns ./cmd/exploredns # Final stage FROM alpine:3.21 diff --git a/Dockerfile.web b/Dockerfile.web index 5e6ec62..7636194 100644 --- a/Dockerfile.web +++ b/Dockerfile.web @@ -8,7 +8,8 @@ RUN go mod download COPY . . -RUN CGO_ENABLED=0 GOOS=linux go build -trimpath -ldflags="-s -w" -o /out/server ./cmd/server +ARG VERSION=dev +RUN CGO_ENABLED=0 GOOS=linux go build -trimpath -ldflags="-s -w -X main.version=${VERSION}" -o /out/server ./cmd/server # Final stage FROM alpine:3.21 diff --git a/Makefile b/Makefile index 8887fe4..d73dcb0 100644 --- a/Makefile +++ b/Makefile @@ -3,14 +3,16 @@ SERVER_BINARY_NAME=exploredns-server BUILD_DIR=bin GO=go GOFLAGS=-v +VERSION?=$(shell git describe --tags --always 2>/dev/null || echo dev) +LDFLAGS=-ldflags "-X main.version=$(VERSION)" .PHONY: build build-server build-all test lint clean cover deploy deploy-status build: - $(GO) build $(GOFLAGS) -o $(BUILD_DIR)/$(BINARY_NAME) ./cmd/exploredns + $(GO) build $(GOFLAGS) $(LDFLAGS) -o $(BUILD_DIR)/$(BINARY_NAME) ./cmd/exploredns build-server: - $(GO) build $(GOFLAGS) -o $(BUILD_DIR)/$(SERVER_BINARY_NAME) ./cmd/server + $(GO) build $(GOFLAGS) $(LDFLAGS) -o $(BUILD_DIR)/$(SERVER_BINARY_NAME) ./cmd/server build-all: build build-server diff --git a/README.md b/README.md index b52d4d2..692dcad 100644 --- a/README.md +++ b/README.md @@ -136,6 +136,9 @@ Output Options: --show-all-stats Show statistics after every node (default: false) --show-results Show the results (default: true) --show-summary-results Show the summary results (default: true) + +General Options: + --version, -V Print version ("exploredns ") and exit ``` Every `--show-X` flag can be negated with `--show-X=false` or `--no-show-X`. @@ -167,8 +170,17 @@ make build-server ``` Open `http://localhost:8080` in your browser. The SPA lets you enter a domain, -choose a record type, and watch the traversal progress in real time. When the -traversal completes the full result tree is displayed in the browser. +choose a record type, and watch the traversal progress in real time as a live +detail tree modelled on the dns.squish.net detail page: one node per referral, +indented per depth, with glue-resolution subtrees collapsed behind per-node +"show resolve" toggles (a "raw log" toggle reveals the flat event feed for +debugging). When the traversal completes the full result list is displayed, +followed by a Servers section: every nameserver queried during the traversal +is fingerprinted (`version.bind`) and shown on an OpenStreetMap/Leaflet map +plus a Country / City / Servers / Software guess table. Geolocation happens +client-side in your browser via the free [geojs.io](https://www.geojs.io/) +API (`get.geojs.io`); servers that cannot be located are still listed with a +dash location, and the table works without the map when offline. ### API endpoints @@ -177,7 +189,8 @@ traversal completes the full result tree is displayed in the browser. | `POST` | `/api/traverse` | Start an asynchronous traversal | | `GET` | `/api/traverse/{id}` | Poll traversal status and results | | `GET` | `/api/traverse/{id}/stream` | Server-Sent Events live progress stream | -| `GET` | `/api/health` | Health check — returns `{"status":"ok"}` | +| `GET` | `/api/traverse/{id}/servers` | Fingerprinted list of every server queried | +| `GET` | `/api/health` | Health check — returns `{"status":"ok","version":""}` | #### POST /api/traverse @@ -206,7 +219,9 @@ Response (`202 Accepted`): #### GET /api/traverse/{id} Returns a snapshot of the job including the full result list once complete. -`status` is one of `running`, `complete`, or `error`. +`status` is one of `running`, `complete`, or `error`. Once post-traversal +fingerprinting has finished the snapshot also carries a `servers` array (the +same list served by `GET /api/traverse/{id}/servers`; omitted before then). ```json { @@ -228,12 +243,36 @@ Returns a snapshot of the job including the full result list once complete. } ``` +#### GET /api/traverse/{id}/servers + +Every `(server name, IP)` pair queried during the traversal — including +glue-resolution subtree servers — fingerprinted via a `version.bind` CHAOS +probe once the traversal reaches a terminal state. Fingerprinting never +delays the traversal results: while it (or the traversal itself) is still in +flight the endpoint answers `202 Accepted` with `{"status":"pending"}`. +Unknown ids answer `404`. When ready: + +```json +{ + "status": "complete", + "servers": [ + {"name": "a.iana-servers.net", "ip": "199.43.135.53", "version": ""}, + {"name": "l.gtld-servers.net", "ip": "192.41.162.30", "version": "..."} + ] +} +``` + +`version` is `""` for servers that don't answer the probe. + #### GET /api/traverse/{id}/stream An [SSE](https://developer.mozilla.org/en-US/docs/Web/API/Server-sent_events) stream of `ProgressEvent` objects, one per `data:` message. Past events recorded before the client connected are replayed immediately, then live events -follow. The stream ends with `event: done`. +follow. Two synthetic stages bracket the end of a job: `{"stage":"complete"}` +when the traversal reaches its terminal status (results are fetchable) and +`{"stage":"servers"}` when the fingerprinted server list is ready. The stream +ends with `event: done`. ``` data: {"stage":"start","depth":1,"name":"www.example.com","qtype":"A","bailiwick":"com"} @@ -256,6 +295,28 @@ variables at startup: | `EXPLOREDNS_JOB_TIMEOUT` | `5m` | Hard deadline per traversal (Go duration). Timed-out jobs report `error` with any partial results. | | `EXPLOREDNS_MAX_JOBS` | `8` | Maximum concurrent traversals; further `POST /api/traverse` requests get `429`. | | `EXPLOREDNS_CORS_ORIGIN` | *(unset)* | Off by default (the SPA is same-origin). Set an origin — or `*` for development — to enable cross-origin API access. | +| `EXPLOREDNS_RATE_LIMIT` | `30/1h` | Per-client-IP token-bucket limit on `POST /api/traverse` in `N/duration` form (e.g. `10/10m`); invalid values fall back to the default. Over-limit requests get `429`. Buckets refill continuously. Direct localhost connections are exempt (dev loop, tests), but proxied requests are always limited by the real client IP from `Fly-Client-IP` / `X-Forwarded-For`. | +| `EXPLOREDNS_WEBHOOK_URL` | *(unset)* | Off by default. When set, the server POSTs a usage-reporting JSON event to this URL on every traversal start and completion (see below). | + +### Usage reporting + +When `EXPLOREDNS_WEBHOOK_URL` is set, the server sends two JSON POSTs per +traversal, each with header `X-ExploreDNS-Event` naming the event: + +- `start` — `{"event":"start","id","domain","query_type","all_roots","client_ip","started_at"}` +- `complete` — `{"event":"complete","id","domain","query_type","client_ip","started_at","done_at","duration_ms","status","error","result_count","summary"}` + where `summary` is the same grouped answers/statuses object returned by + `GET /api/traverse/{id}` and `error` is present only for failed jobs. + +`client_ip` is the requester's IP (`Fly-Client-IP`, else the first +`X-Forwarded-For` entry, else the connection address). Delivery is +fire-and-forget: a 5-second timeout, one retry after 2 seconds, and failures +are logged without ever affecting the traversal or the API response. On +Fly.io, configure it as a secret rather than in `fly.toml`: + +```sh +fly secrets set EXPLOREDNS_WEBHOOK_URL=https://example.com/hook +``` --- @@ -277,13 +338,20 @@ make deploy # flyctl deploy --remote-only the SPA at `https://.fly.dev/` with `/api/health` as the health check. Machine placement is imperative rather than part of `fly.toml`; the current -production topology is one machine in Sydney and one in Virginia: +production topology is four regions — Sydney, Virginia, Singapore, and +London — so anycast wake-up behaviour can be observed from anywhere: ```sh -flyctl scale count 2 --region syd,iad +flyctl scale count 4 --region syd,iad,sin,lhr ``` `flyctl deploy` preserves existing machines and regions on redeploys. +`GET /api/health` reports which region served the request (`region` field, +present only on Fly), making the routing easy to observe: + +```sh +curl -s https://exploredns.hansenits.com/api/health | jq -r .region +``` ### Continuous deployment @@ -294,14 +362,38 @@ dispatch). It needs a `FLY_API_TOKEN` repository secret: flyctl tokens create deploy -x 999999h ``` +### Releases + +Pushing a `v*` tag triggers the full release pipeline: + +1. `.gitea/workflows/release.yml` (`binaries` job) cross-compiles the CLI + and server for linux/amd64, linux/arm64, darwin/amd64, darwin/arm64 and + windows/amd64, packages them as + `exploredns___.tar.gz` (`.zip` on Windows) plus a + `SHA256SUMS` file, and attaches everything to the Gitea release for the + tag. Create the release with notes by hand before (or after) pushing + the tag — the workflow attaches assets to an existing release, creates + a bare one only when none exists, and skips already-attached assets so + re-runs are safe. +2. `.gitea/workflows/release.yml` (`docker` job) pushes + `gitea.hansenits.com.au/hits/exploredns-cli` and `…/exploredns-web` + images tagged `` and `latest`. +3. `.gitea/workflows/deploy.yml` deploys the web server to Fly.io. + +All binaries are stamped with the tag via +`-ldflags "-X main.version="`; check with `exploredns --version` or +`GET /api/health`. Local `make build` stamps from +`git describe --tags --always`. + ### Notes - Traversal traffic is outbound UDP/TCP port 53, which Fly machines allow; upstream root discovery uses Fly's internal resolver via `/etc/resolv.conf` and falls back to the built-in IANA root hints. -- The job timeout, job cap, and same-origin CORS defaults above are what make - unauthenticated public exposure reasonable; tighten `EXPLOREDNS_MAX_JOBS` - if the app attracts traffic. +- The job timeout, job cap, per-IP rate limit, and same-origin CORS defaults + above are what make unauthenticated public exposure reasonable; tighten + `EXPLOREDNS_MAX_JOBS` or `EXPLOREDNS_RATE_LIMIT` if the app attracts + traffic. --- diff --git a/cmd/exploredns/main.go b/cmd/exploredns/main.go index 90d60a8..c125547 100644 --- a/cmd/exploredns/main.go +++ b/cmd/exploredns/main.go @@ -13,6 +13,9 @@ import ( "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" ) +// version is stamped at build time via -ldflags "-X main.version=...". +var version = "dev" + func main() { cfg := config.DefaultConfig() @@ -42,6 +45,12 @@ func main() { flag.BoolVar(&dFlag, "debug", false, "Print debug diagnostics to stderr") flag.BoolVar(&ddFlag, "dd", false, "Like -d plus library-level debug") + // Version: long and short form share the same variable (mirrors + // dnstraverse's -V). + var showVersion bool + flag.BoolVar(&showVersion, "version", false, "Print version and exit") + flag.BoolVar(&showVersion, "V", false, "Print version and exit (shorthand)") + // Quiet: long and short form share the same variable. var quietVal bool flag.BoolVar(&quietVal, "quiet", cfg.Quiet, "Suppress the header block") @@ -69,6 +78,11 @@ func main() { flag.Parse() + if showVersion { + fmt.Printf("exploredns %s\n", version) + return + } + args := flag.Args() cfg.QueryType = *queryType diff --git a/cmd/server/main.go b/cmd/server/main.go index f6f2eb9..dc3ea8a 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -13,17 +13,21 @@ import ( "gitea.hansenits.com.au/hits/ExploreDNS/web/api" ) +// version is stamped at build time via -ldflags "-X main.version=...". +var version = "dev" + func main() { addr := flag.String("addr", ":8080", "listen address (host:port)") flag.Parse() srv := api.NewServer(*addr) + srv.SetVersion(version) if err := srv.Start(); err != nil { fmt.Fprintf(os.Stderr, "Error: %v\n", err) os.Exit(1) } - log.Printf("ExploreDNS API server listening on %s", srv.Addr()) + log.Printf("ExploreDNS API server %s listening on %s", version, srv.Addr()) quit := make(chan os.Signal, 1) signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM) diff --git a/internal/config/config.go b/internal/config/config.go index 4794f41..c12792f 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -167,6 +167,9 @@ func PrintUsage() { {"--show-results", "Show the results (default true)"}, {"--show-summary-results", "Show the summary results (default true)"}, }}, + {"General Options", [][2]string{ + {"--version, -V", "Print version and exit"}, + }}, } for _, group := range flagGroups { diff --git a/web/api/handler.go b/web/api/handler.go index 57dfd3e..07f4b87 100644 --- a/web/api/handler.go +++ b/web/api/handler.go @@ -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 { diff --git a/web/api/handler_internal_test.go b/web/api/handler_internal_test.go index c4430c6..34ac3b4 100644 --- a/web/api/handler_internal_test.go +++ b/web/api/handler_internal_test.go @@ -2,6 +2,8 @@ package api import ( "context" + "encoding/json" + "net" "net/http/httptest" "strconv" "strings" @@ -147,3 +149,171 @@ func TestSubscribeSnapshot_NoDuplicates(t *testing.T) { t.Error(msg) } } + +// fakeVersionQuerier returns canned version strings without touching the +// network. +type fakeVersionQuerier struct { + mu sync.Mutex + versions map[string]string + queried []string +} + +func (f *fakeVersionQuerier) Query(_ context.Context, ip net.IP) string { + f.mu.Lock() + defer f.mu.Unlock() + f.queried = append(f.queried, ip.String()) + return f.versions[ip.String()] +} + +// TestGetServers_PendingWhileRunning verifies the 202 pending shape while a +// job has not finished fingerprinting (running or just-completed). +func TestGetServers_PendingWhileRunning(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + h := newHandler(ctx) + + now := time.Now() + jobs := []*TraversalJob{ + {ID: "running", Status: statusRunning, StartedAt: now}, + {ID: "fingerprinting", Status: statusComplete, StartedAt: now, DoneAt: &now}, + } + for _, j := range jobs { + h.st.set(j) + } + + for _, id := range []string{"running", "fingerprinting"} { + req := httptest.NewRequest("GET", "/api/traverse/"+id+"/servers", nil) + rec := httptest.NewRecorder() + h.mux.ServeHTTP(rec, req) + + if rec.Code != 202 { + t.Fatalf("%s: want 202, got %d: %s", id, rec.Code, rec.Body.String()) + } + var body map[string]string + if err := json.NewDecoder(rec.Body).Decode(&body); err != nil { + t.Fatalf("%s: decode: %v", id, err) + } + if body["status"] != "pending" { + t.Fatalf("%s: want status=pending, got %q", id, body["status"]) + } + } +} + +// TestFingerprintServers_StoresServersAndPublishes drives the fingerprint +// step with a fake querier: (server, ip) pairs become sorted job.Servers with +// versions, pseudo/invalid entries are skipped, a {"stage":"servers"} event +// is published, and the endpoint flips from pending to the final list. +func TestFingerprintServers_StoresServersAndPublishes(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + h := newHandler(ctx) + fake := &fakeVersionQuerier{versions: map[string]string{ + "192.0.2.1": "TestDNS 1.0", + "192.0.2.2": "", + }} + h.fp = fake + + now := time.Now() + job := &TraversalJob{ID: "j", Status: statusComplete, StartedAt: now, DoneAt: &now} + h.st.set(job) + sub, _, _ := job.subscribeSnapshot() + defer job.unsubscribe(sub) + + h.fingerprintServers(job, map[string][]string{ + "b.example.net": {"192.0.2.2"}, + "a.example.net": {"192.0.2.1", "key:pseudo:entry", "not-an-ip"}, + }) + + want := []ServerInfo{ + {Name: "a.example.net", IP: "192.0.2.1", Version: "TestDNS 1.0"}, + {Name: "b.example.net", IP: "192.0.2.2", Version: ""}, + } + job.mu.RLock() + got := append([]ServerInfo(nil), job.Servers...) + done := job.serversDone + job.mu.RUnlock() + if !done { + t.Fatal("serversDone not set") + } + if len(got) != len(want) { + t.Fatalf("servers = %+v, want %+v", got, want) + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("servers[%d] = %+v, want %+v", i, got[i], want[i]) + } + } + + select { + case ev := <-sub: + if ev.Stage != "servers" { + t.Fatalf("published stage = %q, want servers", ev.Stage) + } + default: + t.Fatal("no servers event published") + } + + // Endpoint now serves the final list. + req := httptest.NewRequest("GET", "/api/traverse/j/servers", nil) + rec := httptest.NewRecorder() + h.mux.ServeHTTP(rec, req) + if rec.Code != 200 { + t.Fatalf("want 200, got %d: %s", rec.Code, rec.Body.String()) + } + var body struct { + Status string `json:"status"` + Servers []ServerInfo `json:"servers"` + } + if err := json.NewDecoder(rec.Body).Decode(&body); err != nil { + t.Fatal(err) + } + if body.Status != "complete" || len(body.Servers) != 2 { + t.Fatalf("body = %+v", body) + } + + // Snapshot includes servers too. + req = httptest.NewRequest("GET", "/api/traverse/j", nil) + rec = httptest.NewRecorder() + h.mux.ServeHTTP(rec, req) + if rec.Code != 200 { + t.Fatalf("snapshot: want 200, got %d", rec.Code) + } + var snap struct { + Servers []ServerInfo `json:"servers"` + } + if err := json.NewDecoder(rec.Body).Decode(&snap); err != nil { + t.Fatal(err) + } + if len(snap.Servers) != 2 { + t.Fatalf("snapshot servers = %+v, want 2 entries", snap.Servers) + } +} + +// TestFingerprintServers_EmptySeen still terminates the pending state and +// publishes the servers event for traversals that recorded no servers. +func TestFingerprintServers_EmptySeen(t *testing.T) { + h := &Handler{st: newStore()} + job := &TraversalJob{ID: "e", Status: statusError} + h.st.set(job) + sub, _, _ := job.subscribeSnapshot() + defer job.unsubscribe(sub) + + h.fingerprintServers(job, nil) + + job.mu.RLock() + defer job.mu.RUnlock() + if !job.serversDone { + t.Fatal("serversDone not set") + } + if len(job.Servers) != 0 { + t.Fatalf("servers = %+v, want empty", job.Servers) + } + select { + case ev := <-sub: + if ev.Stage != "servers" { + t.Fatalf("published stage = %q, want servers", ev.Stage) + } + default: + t.Fatal("no servers event published") + } +} diff --git a/web/api/handler_test.go b/web/api/handler_test.go index c30a06e..18a6361 100644 --- a/web/api/handler_test.go +++ b/web/api/handler_test.go @@ -29,6 +29,7 @@ func newTestServer(t *testing.T) *api.Server { } func TestHealth(t *testing.T) { + t.Setenv("FLY_REGION", "") // ensure region is absent regardless of host env srv := newTestServer(t) defer srv.Shutdown(5 * time.Second) //nolint:errcheck @@ -48,6 +49,55 @@ func TestHealth(t *testing.T) { if body["status"] != "ok" { t.Fatalf("want status=ok, got %q", body["status"]) } + if body["version"] != "dev" { + t.Fatalf("want version=dev, got %q", body["version"]) + } + if region, ok := body["region"]; ok { + t.Fatalf("region should be omitted outside Fly, got %q", region) + } +} + +func TestHealthReportsFlyRegion(t *testing.T) { + t.Setenv("FLY_REGION", "syd") + 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() + + var body map[string]string + if err := json.NewDecoder(resp.Body).Decode(&body); err != nil { + t.Fatal(err) + } + if body["region"] != "syd" { + t.Fatalf("want region=syd, got %q", body["region"]) + } +} + +func TestHealthReportsStampedVersion(t *testing.T) { + srv := api.NewServer("127.0.0.1:0") + srv.SetVersion("v1.2.3") + if err := srv.Start(); err != nil { + t.Fatalf("start server: %v", err) + } + 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() + + var body map[string]string + if err := json.NewDecoder(resp.Body).Decode(&body); err != nil { + t.Fatal(err) + } + if body["version"] != "v1.2.3" { + t.Fatalf("want version=v1.2.3, got %q", body["version"]) + } } func TestCORSDisabledByDefault(t *testing.T) { @@ -205,6 +255,95 @@ func TestGetTraversal_Found(t *testing.T) { } } +func TestGetServers_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/servers") + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusNotFound { + t.Fatalf("want 404, got %d", resp.StatusCode) + } +} + +// TestGetServers_AvailableAfterCompletion drives a full job through the API: +// the servers endpoint answers 202 pending while the traversal/fingerprinting +// is in flight and the fingerprinted list once everything finished. +func TestGetServers_AvailableAfterCompletion(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + 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) + } + + deadline := time.Now().Add(90 * time.Second) + for { + resp, err := http.Get("http://" + srv.Addr() + "/api/traverse/" + start.ID + "/servers") + if err != nil { + t.Fatal(err) + } + body, err := io.ReadAll(resp.Body) + resp.Body.Close() + if err != nil { + t.Fatal(err) + } + + switch resp.StatusCode { + case http.StatusAccepted: + var pending struct { + Status string `json:"status"` + } + if err := json.Unmarshal(body, &pending); err != nil { + t.Fatalf("pending body: %v (%s)", err, body) + } + if pending.Status != "pending" { + t.Fatalf("want status=pending, got %q", pending.Status) + } + case http.StatusOK: + var done struct { + Status string `json:"status"` + Servers []struct { + Name string `json:"name"` + IP string `json:"ip"` + Version string `json:"version"` + } `json:"servers"` + } + if err := json.Unmarshal(body, &done); err != nil { + t.Fatalf("servers body: %v (%s)", err, body) + } + if done.Status != "complete" { + t.Fatalf("want status=complete, got %q", done.Status) + } + if done.Servers == nil { + t.Fatalf("servers key missing or null: %s", body) + } + return + default: + t.Fatalf("unexpected status %d: %s", resp.StatusCode, body) + } + + if time.Now().After(deadline) { + t.Fatal("timed out waiting for servers to become available") + } + time.Sleep(250 * time.Millisecond) + } +} + func TestStreamTraversal_NotFound(t *testing.T) { srv := newTestServer(t) defer srv.Shutdown(5 * time.Second) //nolint:errcheck @@ -322,6 +461,71 @@ func TestStaticSPA_TypeOptions(t *testing.T) { } } +// TestStaticSPA_DetailTree asserts the SPA ships the live detail tree with +// its resolve-subtree toggle markup, plus the raw-log fallback feed so the +// old flat progress view is still reachable for debugging. +func TestStaticSPA_DetailTree(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() + + body, err := io.ReadAll(resp.Body) + if err != nil { + t.Fatal(err) + } + page := string(body) + + for _, want := range []string{ + `id="detailTree"`, // detail-tree container + `resolve-toggle`, // per-node show/hide resolve markup + `show resolve`, // toggle wording mirrors dns.squish.net + `id="progressFeed"`, // raw-log fallback feed still present + `id="rawToggle"`, // toggle that reveals it + } { + if !strings.Contains(page, want) { + t.Errorf("index.html missing %q", want) + } + } +} + +// TestStaticSPA_ServersSection asserts the SPA ships the server map/table +// section: Leaflet lazy-loaded from unpkg, geojs.io client-side geolocation, +// and the reference-style table headings. +func TestStaticSPA_ServersSection(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() + + body, err := io.ReadAll(resp.Body) + if err != nil { + t.Fatal(err) + } + page := string(body) + + for _, want := range []string{ + `unpkg.com/leaflet@1.9`, // map library CDN + `get.geojs.io`, // client-side geolocation service + `id="serversCard"`, + `id="serverMap"`, + `CountryCityServersSoftware guess`, + `/servers`, // fetches the servers endpoint + } { + if !strings.Contains(page, want) { + t.Errorf("index.html missing %q", want) + } + } +} + func TestStaticSPA_FallbackToIndex(t *testing.T) { srv := newTestServer(t) defer srv.Shutdown(5 * time.Second) //nolint:errcheck diff --git a/web/api/ratelimit.go b/web/api/ratelimit.go new file mode 100644 index 0000000..66aa25c --- /dev/null +++ b/web/api/ratelimit.go @@ -0,0 +1,120 @@ +package api + +import ( + "fmt" + "net" + "net/http" + "strconv" + "strings" + "sync" + "time" +) + +// Default per-IP rate limit for POST /api/traverse, overridable via +// EXPLOREDNS_RATE_LIMIT ("N/duration", e.g. "30/1h" or "10/10m"). +const ( + defaultRateLimitCount = 30 + defaultRateLimitWindow = time.Hour +) + +// rateLimiter is an in-memory token-bucket limiter keyed by client IP. +// Each bucket starts full with limit tokens and refills continuously at +// limit tokens per window. +type rateLimiter struct { + limit int + window time.Duration + + mu sync.Mutex + buckets map[string]*tokenBucket + now func() time.Time // overridable in tests +} + +type tokenBucket struct { + tokens float64 + last time.Time +} + +func newRateLimiter(limit int, window time.Duration) *rateLimiter { + return &rateLimiter{ + limit: limit, + window: window, + buckets: make(map[string]*tokenBucket), + now: time.Now, + } +} + +// allow reports whether one request from ip fits within the limit, +// consuming a token when it does. +func (rl *rateLimiter) allow(ip string) bool { + rl.mu.Lock() + defer rl.mu.Unlock() + + now := rl.now() + b, ok := rl.buckets[ip] + if !ok { + b = &tokenBucket{tokens: float64(rl.limit), last: now} + rl.buckets[ip] = b + } else { + refill := now.Sub(b.last).Seconds() * float64(rl.limit) / rl.window.Seconds() + b.tokens = min(b.tokens+refill, float64(rl.limit)) + b.last = now + } + if b.tokens < 1 { + return false + } + b.tokens-- + return true +} + +// sweep drops buckets idle for at least one full window; such buckets +// would be full again anyway, so dropping them loses nothing. +func (rl *rateLimiter) sweep() { + rl.mu.Lock() + defer rl.mu.Unlock() + cutoff := rl.now().Add(-rl.window) + for ip, b := range rl.buckets { + if b.last.Before(cutoff) { + delete(rl.buckets, ip) + } + } +} + +// String renders the limit for error messages, e.g. "30 requests per 1h0m0s". +func (rl *rateLimiter) String() string { + return fmt.Sprintf("%d requests per %s", rl.limit, rl.window) +} + +// rateLimitExempt reports whether r may bypass the rate limit: direct +// connections from loopback (the SPA dev loop, tests, health tooling). +// Proxied requests are never exempt — when Fly-Client-IP or +// X-Forwarded-For is present, RemoteAddr is just the proxy, so the real +// client IP must be limited even though the socket peer is local. +func rateLimitExempt(r *http.Request) bool { + if r.Header.Get("Fly-Client-IP") != "" || r.Header.Get("X-Forwarded-For") != "" { + return false + } + host, _, err := net.SplitHostPort(r.RemoteAddr) + if err != nil { + host = r.RemoteAddr + } + ip := net.ParseIP(host) + return ip != nil && ip.IsLoopback() +} + +// parseRateLimit parses "N/duration" (e.g. "30/1h"), falling back to the +// defaults when v is empty or invalid. +func parseRateLimit(v string) (int, time.Duration) { + parts := strings.SplitN(v, "/", 2) + if len(parts) != 2 { + return defaultRateLimitCount, defaultRateLimitWindow + } + n, err := strconv.Atoi(strings.TrimSpace(parts[0])) + if err != nil || n <= 0 { + return defaultRateLimitCount, defaultRateLimitWindow + } + d, err := time.ParseDuration(strings.TrimSpace(parts[1])) + if err != nil || d <= 0 { + return defaultRateLimitCount, defaultRateLimitWindow + } + return n, d +} diff --git a/web/api/ratelimit_test.go b/web/api/ratelimit_test.go new file mode 100644 index 0000000..e32aad5 --- /dev/null +++ b/web/api/ratelimit_test.go @@ -0,0 +1,169 @@ +package api + +import ( + "context" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" +) + +func TestParseRateLimit(t *testing.T) { + tests := []struct { + in string + limit int + window time.Duration + }{ + {"", defaultRateLimitCount, defaultRateLimitWindow}, + {"30/1h", 30, time.Hour}, + {"10/10m", 10, 10 * time.Minute}, + {"5 / 30s", 5, 30 * time.Second}, + {"bogus", defaultRateLimitCount, defaultRateLimitWindow}, + {"0/1h", defaultRateLimitCount, defaultRateLimitWindow}, + {"-3/1h", defaultRateLimitCount, defaultRateLimitWindow}, + {"10/-1h", defaultRateLimitCount, defaultRateLimitWindow}, + {"10/soon", defaultRateLimitCount, defaultRateLimitWindow}, + {"/1h", defaultRateLimitCount, defaultRateLimitWindow}, + } + for _, tc := range tests { + limit, window := parseRateLimit(tc.in) + if limit != tc.limit || window != tc.window { + t.Errorf("parseRateLimit(%q) = %d, %s; want %d, %s", + tc.in, limit, window, tc.limit, tc.window) + } + } +} + +func TestNewHandlerReadsRateLimitEnv(t *testing.T) { + t.Setenv("EXPLOREDNS_RATE_LIMIT", "5/10m") + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + h := newHandler(ctx) + if h.limiter.limit != 5 || h.limiter.window != 10*time.Minute { + t.Fatalf("limiter = %d/%s, want 5/10m", h.limiter.limit, h.limiter.window) + } +} + +// postTraverse sends POST /api/traverse with an empty JSON body so requests +// that pass the rate limiter fail validation (400) instead of spawning a +// real traversal. remoteAddr and headers shape the client identity. +func postTraverse(t *testing.T, h *Handler, remoteAddr string, headers map[string]string) *httptest.ResponseRecorder { + t.Helper() + req := httptest.NewRequest(http.MethodPost, "/api/traverse", strings.NewReader(`{}`)) + req.RemoteAddr = remoteAddr + for k, v := range headers { + req.Header.Set(k, v) + } + rec := httptest.NewRecorder() + h.mux.ServeHTTP(rec, req) + return rec +} + +func TestRateLimit_OverLimitReturns429(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + h := newHandler(ctx) + h.limiter = newRateLimiter(2, time.Hour) + + hdr := map[string]string{"X-Forwarded-For": "203.0.113.9"} + for i := 0; i < 2; i++ { + if rec := postTraverse(t, h, "10.0.0.1:1234", hdr); rec.Code != http.StatusBadRequest { + t.Fatalf("request %d: want 400 (under limit), got %d: %s", i, rec.Code, rec.Body.String()) + } + } + rec := postTraverse(t, h, "10.0.0.1:1234", hdr) + if rec.Code != http.StatusTooManyRequests { + t.Fatalf("want 429 over limit, got %d: %s", rec.Code, rec.Body.String()) + } + if body := rec.Body.String(); !strings.Contains(body, "2 requests per 1h0m0s") { + t.Fatalf("429 body should name the limit, got %s", body) + } + + // A distinct client IP has its own bucket and is unaffected. + other := map[string]string{"X-Forwarded-For": "203.0.113.10"} + if rec := postTraverse(t, h, "10.0.0.1:1234", other); rec.Code != http.StatusBadRequest { + t.Fatalf("distinct IP: want 400, got %d: %s", rec.Code, rec.Body.String()) + } +} + +func TestRateLimit_RefillsContinuously(t *testing.T) { + rl := newRateLimiter(2, time.Second) + now := time.Now() + rl.now = func() time.Time { return now } + + if !rl.allow("a") || !rl.allow("a") { + t.Fatal("first two requests should be allowed") + } + if rl.allow("a") { + t.Fatal("third request should be denied") + } + // Half a window refills half the bucket: one token. + now = now.Add(500 * time.Millisecond) + if !rl.allow("a") { + t.Fatal("request after refill should be allowed") + } + if rl.allow("a") { + t.Fatal("bucket should hold only the refilled token") + } +} + +func TestRateLimit_SweepDropsIdleBuckets(t *testing.T) { + rl := newRateLimiter(1, time.Minute) + now := time.Now() + rl.now = func() time.Time { return now } + + rl.allow("stale") + now = now.Add(2 * time.Minute) + rl.allow("fresh") + rl.sweep() + + rl.mu.Lock() + defer rl.mu.Unlock() + if _, ok := rl.buckets["stale"]; ok { + t.Fatal("idle bucket should have been swept") + } + if _, ok := rl.buckets["fresh"]; !ok { + t.Fatal("active bucket should survive the sweep") + } +} + +func TestRateLimit_LocalhostExempt(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + h := newHandler(ctx) + h.limiter = newRateLimiter(1, time.Hour) + + for _, addr := range []string{"127.0.0.1:5555", "[::1]:5555"} { + for i := 0; i < 3; i++ { + rec := postTraverse(t, h, addr, nil) + if rec.Code != http.StatusBadRequest { + t.Fatalf("%s request %d: localhost should be exempt, got %d: %s", + addr, i, rec.Code, rec.Body.String()) + } + } + } +} + +// TestRateLimit_ProxiedLocalhostNotExempt verifies that a request arriving +// from a local proxy (RemoteAddr loopback) is still limited by the real +// client IP carried in Fly-Client-IP / X-Forwarded-For. +func TestRateLimit_ProxiedLocalhostNotExempt(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + h := newHandler(ctx) + h.limiter = newRateLimiter(1, time.Hour) + + hdr := map[string]string{"Fly-Client-IP": "198.51.100.4"} + if rec := postTraverse(t, h, "127.0.0.1:5555", hdr); rec.Code != http.StatusBadRequest { + t.Fatalf("first proxied request: want 400, got %d", rec.Code) + } + if rec := postTraverse(t, h, "127.0.0.1:5555", hdr); rec.Code != http.StatusTooManyRequests { + t.Fatalf("second proxied request: want 429, got %d", rec.Code) + } + // The same proxy forwarding a different client is unaffected. + other := map[string]string{"Fly-Client-IP": "198.51.100.5"} + if rec := postTraverse(t, h, "127.0.0.1:5555", other); rec.Code != http.StatusBadRequest { + t.Fatalf("other client via proxy: want 400, got %d", rec.Code) + } +} diff --git a/web/api/server.go b/web/api/server.go index dec400b..9c3fb4a 100644 --- a/web/api/server.go +++ b/web/api/server.go @@ -27,14 +27,23 @@ var staticFiles embed.FS // Server is the HTTP API server. type Server struct { - addr string - srv *http.Server - cancel context.CancelFunc + addr string + version string + srv *http.Server + cancel context.CancelFunc } // NewServer creates a new Server that listens on addr (e.g. ":8080"). func NewServer(addr string) *Server { - return &Server{addr: addr} + return &Server{addr: addr, version: "dev"} +} + +// SetVersion records the build version reported by GET /api/health. +// Call before Start; empty values are ignored. +func (s *Server) SetVersion(v string) { + if v != "" { + s.version = v + } } // Start builds the HTTP handler, begins listening, and returns when the @@ -44,6 +53,7 @@ func (s *Server) Start() error { ctx, cancel := context.WithCancel(context.Background()) s.cancel = cancel h := newHandler(ctx) + h.version = s.version sub, err := fs.Sub(staticFiles, "static") if err != nil { diff --git a/web/api/static/favicon.ico b/web/api/static/favicon.ico new file mode 100644 index 0000000..eb1e2f0 Binary files /dev/null and b/web/api/static/favicon.ico differ diff --git a/web/api/static/favicon.svg b/web/api/static/favicon.svg new file mode 100644 index 0000000..f000aeb --- /dev/null +++ b/web/api/static/favicon.svg @@ -0,0 +1,13 @@ + + + + + + + + + + + + + diff --git a/web/api/static/index.html b/web/api/static/index.html index 8e101b2..f87103a 100644 --- a/web/api/static/index.html +++ b/web/api/static/index.html @@ -4,6 +4,8 @@ ExploreDNS + +