From af15c9c2d446bd3c4aaa9b8bca758be9568341fb Mon Sep 17 00:00:00 2001
From: Gary Hansen
`; failure lines have no ``). Verbatim failure catalogue from archived pages [WB]: +- `No such domain (NXDOMAIN) at ns1.cctld.co (156.154.100.25)` +- `Query timed out at ns2.twisted4life.com (194.152.92.62)` +- `Server failure (SERVFAIL) at casper.ln1x.ablesky.com (173.255.217.216)` +- `Refused at ns1.afdns.net (190.183.61.2)` +- `NODATA (for this type) at ns1.iranfirewall.in (148.251.110.146)` +- `Lame referral received from f.nic.es (130.206.1.2) to dns2.namecheaphosting.com (69.160.33.71)` — web adds "received"; CLI says `Lame referral from …` [SRC referral.rb:506] +- `recvfrom failed from 185.208.174.92; No route to host - recvfrom(2) at ns1.dadzi.ir (185.208.174.92)` (exception text verbatim from Ruby) +- Failures inside a resolve append `While querying ns1.twisted4life.com/IN/A` (CLI: `While querying #{qname}/#{qclass}/#{qtype}`). + +Then links: **"Traversal detail"** → `/detail`; **"Server location and version information"** → `/servers`; while running: "Please wait while your traversal is completed, or alternatively watch the traversal as it progresses." + +### 4.2 Web detail page `/traverses//detail` [WB streams; shell confirmed LIVE] + +Adds table rows **Initial Root** (`f.root-servers.net, 192.5.5.241` + `(1 roots returned)`) and **Query Key** — one coloured dotted-border chip per query, text like `A careerhosts.net`; colours cycle `#fefecc, #ccfecc, #ccfefe, #ccccfe, #feccfe, #fee5cc`. Banner node: `Traversing for careerhosts.net type A starting at the root(s)`. + +Tree rows are `div.node_line` tables: indentation drawn with 32×23 px spacer images (`spacer_gap/line/tee/end.png`); node text is `#{server} (#{ip}) <#{bailiwick}>` (root renders `<>`; glueless renders `( )` with a placeholder span later filled, truncated by the UI, e.g. `95.130.252.149,Loop encounter...`). Node extras: progress spinner → status icon linking `info_url('/traverses/ /info?packet=N')` (`referral.png` / `success.png` / `warning.png`), `show resolve`/`hide resolve` toggle revealing the hidden `.0` resolve subtree, and italic `completed earlier` span for fast-mode hits. Colour lifecycle: inserted muted (`#ededdd`; resolve nodes `#ddeddd`) → active `border:1px solid black` + query-key colour → completed. Updates arrive as Prototype.js RJS from `/detail/fetch?count=N` (`Element.insert`, `visualEffect`, `setStyle`), each chunk ending `count = N; setTimeout("poll()", 1000);`. On completion: Results injected (same stats_line HTML as show page), then a **Fingerprinting** phase cycling every server name, then the "further" (servers) link appears. + +### 4.3 Web packet popup `/info?packet=N` [WB 2021] + +` What is arudns1.aruba.it type A?
` + `Answer from 192.12.192.5 (dns.nic.it)
` + `` dnsruby dump (`;; Answer received from 192.12.192.5 (189 bytes)`, `;; Security Level : UNCHECKED`, header/flags/sections, `OPT pseudo-record : payloadsize 4096, xrcode 0, version 0, flags 0`) + `Status: referral
` "A referral occured. This means that the resolver did not know the answer, but indicated that results can be found elsewhere." (or `Status: answered` / "The server was able to answer the question successfully.") + `The following records were considered worthy of caching:
` + `These servers were chosen in bailiwick 'aruba.it' for the referral:
` (`name (ip)` lines) + `Results after processing this node and all branches beneath:
` with subtree-scoped stats_lines. Fast-mode variant prepends blue `This was completed earlier (fast mode)`. + +### 4.4 Web servers page `/traverses//servers` [WB 2016/2025; shell LIVE] + +Google map + `table#servers_table`, columns exactly: **Country | City | Area code | Postal code | Servers | Software guess |** (+Show on map). Servers cell: `38.71.66.4 (ns1.virtualempire.com)
…` grouped by geolocation; Software guess: ` ISC BIND 9.2.3rc1 -- 9.4.0a4 (9.6.1-P1)`, ` VeriSign ATLAS ` (2025 adds ` (ATLAS)` suffix), ` TIMEOUT`, ` No match found`. Fingerprint DB is fpdns-derived ("Portions Copyright (c) 2003,2004,2005 Roy Arends & Jakob Schlyter"). + +### 4.5 CLI output [DOC home page sample + mcgill.org.za 2017 run + SRC] + +Header (unless -q): `# Using fast mode` / `# Limiting traverse to one root` / `# UDP size N (EDNS0 is on)` [always prints "on" due to source bug `options[:udpsize == 512]`] / `# Retries N, max depth N` / `# Allow TCP is true, always TCP is false`; then `Using E.ROOT-SERVERS.NET (192.203.230.10) as initial root` / `Running query www.google.com type a`. + +Progress, verbose (`refid [qname] server (ips)`) [DOC, verbatim]: +``` +1 [www.google.com] E.ROOT-SERVERS.NET (192.203.230.10) <> +1.1 [www.google.com] B.GTLD-SERVERS.NET (192.33.14.30) +1.1.1 [www.google.com] ns1.google.com (216.239.32.10) +1.1.1.1 [www.l.google.com] a.l.google.com (209.85.139.9) +``` +Non-verbose omits `[qname]`/` ` [mcgill 2017, verbatim]: +``` +1.2.2 ns1.coza.net.za -- resolving +1.2.2 ns1.coza.net.za (66.135.62.20,Loop encountered resolving ns.coza.net.za) +1.3.1 ns4.iafrica.com (196.7.142.131) -- completed earlier (1.2.3) +``` + +`Results:` — blocks with `printf "%5.1f%%: "` then per status [SRC referral.rb:503-532]: `Answer from ( )` + RRs indented 12 spaces; `No glue at ( ) for `; `Lame referral from ( ) to ( )`; `Loop encountered at `; `CNAME loop encountered at `; ` at ( )`; `NODATA (for this type) at ( )`; ` at …`; fallback `Stopped at ( ))` [trailing `)` is a source bug]; plus `While querying / / ` when the failing query differs. [DOC, verbatim]: +``` + 14.3%: Answer from e.l.google.com (209.85.137.9) + www.l.google.com. 300 IN A 74.125.77.99 +``` + +`Summary Results:` — same wording set as web (Section 4.1), prefix two spaces, e.g. ` 99.8% answered with co.za. 3600 IN NS ns0.is.co.za.` … ` 0.2% resulted in a loop` [mcgill 2017]. RR whitespace collapsed to single spaces; continuation RRs aligned under text start. + +`--show-servers`: `The following servers were encountered:` then `printf "%#{w}s: %-15s%s"` rows sorted by lowercased reversed name [DOC, verbatim]: +``` + ns1.google.com: 216.239.32.10 ISC BIND 9.2.3rc1 -- 9.4.0a0 (9.4.2-P1) +a.gtld-servers.net: 192.5.6.30 VeriSign ATLAS +``` + +### 4.6 Legacy v1 (Perl dnscheck) — different engine, different wording [LIVE run 2026-07-07] + +Nested-table HTML; sections: root-server workings and SOA-serial check, then `Traversal of DNS for .` with units `Asking a.gtld-servers.net (192.5.6.30) for www.example.com (type A)` / `Referral: example.com is at hera.ns.cloudflare.com (108.162.192.162)` / `Response is:` with per-branch `16.7% ( ) with `; dedupe `[see above for results]`; final `Results`: `16.1% of queries will be returned by 108.162.195.228 (elliott.ns.cloudflare.com)` + RRset, `3.5% of queries will end in failure at too many nested queries`, `0.0% of queries will end in failure at 192.5.6.30 (a.gtld-servers.net) - failed to resolve h.gtld-servers.net due to 192.5.6.30 - nameserver loop detected`. **Do not copy this wording into ExploreDNS comparisons against dns.squish.net** — only the aggregate math (probability multiplication/summation) transfers. + +--- + +## 5. Confidence notes + +**Confirmed by primary sources (archived/live site output + Ruby source agree):** +- Web form fields/defaults, routes, captcha behaviour, polling protocol (LIVE + WB, identical 2011–2026). +- Single-initial-root default, equal-split probability model and exact aggregation formulas, bailiwick filtering, hierarchical per-branch cache, fast mode, refid scheme with `.0` resolve subtrees, CNAME restart-from-cache, depth 20 / retries 2 (SRC, corroborated by WB output and DOC). +- Result/Summary wording and percent formatting for: answered, error (NXDOMAIN/SERVFAIL/Refused), exception (timeout/recvfrom), nodata, lame referral, "While querying" suffix (WB verbatim + SRC format strings, verified in the local clone: `summary_stats.rb:76-99,135-136`, `referral.rb:503-532`). +- CLI output shapes (DOC's own annotated sample + one full independent 0.1.14 run from mcgill.org.za 2017 + SRC). + +**Confirmed by source only (never seen in captured output):** web wording for `noglue` ("found no glue" / "No glue at X (ip) for Y"), `loop` and `cname_loop` result/summary lines on the *web* (CLI loop lines are confirmed via mcgill); `Stopped at` fallback; `Formate error (FORMERR)` typo; fast-cache key details; lame-referral "strictly deeper zone" rule. These are near-certain — the site runs this gem — but flagged CONFIRMED-by-source, not by output. + +**Known divergences to be careful with when diffing:** web "Answered from␣␣" (double space) vs CLI "Answer from"; web "Lame referral received from" vs CLI without "received"; web EDNS payloadsize 4096 vs CLI default 2048 (archived output preferred for the web service); library-vs-CLI internal defaults (10/512/false vs 20/2048/true); source bugs worth deciding whether to replicate (`(9.4.2-P1)`-style CLI EDNS banner always "on", `Stopped at …)` stray paren, answered-merge prob discard in `stats_display`, "CNANE loop" typo in `DecodedQuery#to_s`). + +**Not recovered / open items:** +- The Rails front-end (dnstraverseweb) source — closed; web-only wording is reconstructed from captures. +- Exactly one engine version gap: site runs 0.1.14; the cloned source is 0.1.13 master (rubygems has 0.1.14; diff not fetched — format strings match 0.1.14 outputs observed, so risk is low). +- No archived web sample of a *completed* SPF/TXT/SOA traversal, of the web progress-percent computation, or of web noglue/loop/cname_loop result lines. +- A fresh end-to-end capture is achievable: the live site works behind one reCAPTCHA — solve it once in a browser and record `/traverses/ `, `/detail` (streaming XHRs to `/fetch` and `/detail/fetch`), and `/servers`. The author's test domains `yellowtealpurple.net` / `testcname.yellowtealpurple.net` (seen Aug 2025) are candidate CNAME test vectors. The v1 CGI remains captcha-free for aggregate-math cross-checks. + +Where sources conflicted, archived site output was used (traverse IDs random not md5 — WB tests; EDNS 4096 on web — WB packet dumps; web "received"/double-space wordings — WB HTML), with the Ruby source as tie-breaker for everything unobserved. diff --git a/docs/engine-rework-design.md b/docs/engine-rework-design.md new file mode 100644 index 0000000..56b38d4 --- /dev/null +++ b/docs/engine-rework-design.md @@ -0,0 +1,208 @@ +# ExploreDNS Engine Rework — Design + +Goal: make ExploreDNS's traversal semantics and CLI output match the Ruby +`dnstraverse` 0.1.14 engine (the engine behind dns.squish.net). The authoritative +behaviour spec is [dnstraverse-reference-spec.md](dnstraverse-reference-spec.md); +the current-state review is [codebase-review-2026-07-07.md](codebase-review-2026-07-07.md). +The reference Ruby source is at `tools/dnstraverse-ruby/` (run it via +`tools/golden/run-reference.sh`). Golden captures live in `docs/captures/`. + +When this document and the Ruby source disagree, **the Ruby source wins** — +read the relevant `.rb` file before implementing each piece. + +## Ground rules + +1. Every traversal query is **non-recursive (RD=0)**. RD=1 is permitted only for + the initial root discovery against the configured upstream resolver. + Delete the current split where production takes `dns.Query` (RD=1) and tests + take `IterativeQueryWithExchange`; there must be exactly one query path, used + by both production and tests (tests inject a mock exchange into *that* path). + Remove `ensureRDFalse` and the hardcoded `127.0.0.1:53` shortcuts entirely. +2. **Packet cache**: within one run, each (server IP, qname, qclass, qtype, + udpsize) is sent at most once (`caching_resolver.rb`). This is separate from + fast mode. +3. **EDNS0**: OPT added when udp-size > 512 (default 2048). On FORMERR/NOTIMP/ + SERVFAIL with udpsize > 512, retry once at 512 and record warning + ` doesn't seem to support EDNS0`. UDP→TCP retry on truncation when + allow-tcp (default true). Timeout 2s, retries default 2 (mirror dnsruby + retry semantics — read how dnsruby uses retry_times before implementing). +4. **IPv4 only** for server selection, like the reference. AAAA records still + decode and display in answers; they are never used for transport. +5. Probabilities at the root **must sum to 1.0** across aggregated leaves. + Add a test asserting this invariant on mock topologies. + +## Type mapping (Ruby → Go, all in internal/traverse unless noted) + +| Ruby file | Go file | Notes | +|---|---|---| +| `referral.rb` | `referral.go` | The core. Rewrite, keep the name. | +| `info_cache.rb` | `cache.go` | Hierarchical per-branch cache. Rewrite. | +| `decoded_query.rb` + `decoded_query_cache.rb` | `decoded_query.go` | Classification + packet cache. | +| `response*.rb` | `response.go` | Response wrapper + noglue/loop variants. | +| `summary_stats.rb` | `stats.go` (move from internal/output) | Leaf aggregation + summary. | +| `traverser.rb` | `traverser.go` | Stack loop, roots, resolve orchestration. | +| `caching_resolver.rb` | `internal/dns` | Packet-level dedupe. | + +### Referral + +- Fields: `refid string`, `parent *Referral`, `qname/qclass/qtype`, + `server string` (NS hostname; a synthetic "rootroot" node has server="" and + is never displayed), `serverIPs []string` (nil ⇒ needs resolving), + `bailiwick string`, `infoCache *InfoCache` (per-branch), `children`, + per-IP responses, `serverWeights map[string]float64`, warnings. +- **RefID grammar**: dotted path, children numbered from 1 (`1`, `1.1`, `1.1.2`). + A glue-resolution subtree inserts a `.0` component (`1.2.0.1`); nested + resolves nest further. If more than one IP of a server produced children, an + extra childset digit is appended. Depth = count of non-`0` components; + exceeding max-depth (default 20) injects an exception response + `Maxdepth N exceeded`. +- `process()`: query **each IP** in serverIPs (weight 1/len(ips) each) through + the packet cache; classify each response; statuses `referral` and `restart` + produce children (one child per NS name in the referral, **including + glueless NS** with serverIPs=nil). + +### Classification (decoded_query.rb — mirror the order exactly) + +1. network exception → `exception` +2. follow CNAMEs **within the message** (`msg_follow_cnames`); chain leaving + the bailiwick stops following and returns the target; in-message loop → + `cname_loop` +3. rcode != NOERROR → `error` with messages exactly: + `Format error (FORMERR)`, `Server failure (SERVFAIL)`, + `No such domain (NXDOMAIN)`, `Not implemented (NOTIMP)`, `Refused`, + else the rcode string. (The Ruby source has a typo "Formate error" — we + deliberately fix it; this is a documented deviation.) +4. answers exist for endname/qtype → `answered` +5. endname != qname (CNAME landed elsewhere) → `restart` +6. SOA in authority, or no NS in authority → `nodata` +7. NS in authority → `referral`; else `restart` + +Full status vocabulary: **answered, nodata, referral, restart, referral_lame, +error, exception, cname_loop, noglue, loop**. + +### Bailiwick + InfoCache + +- `insideBailiwick(name)`: bailiwick == "" (root), or equal fold, or name ends + with "." + bailiwick. +- `msgCacheable`: partition **all** sections (answer/authority/additional; OPT + dropped) into in-bailiwick (cached) vs out-of-bailiwick (discarded). +- InfoCache is hierarchical: each Response wraps a child cache + (`InfoCache{parent}`); `add()` **replaces** any existing same + name:class:type key; lookups recurse to parent. `getStartServers(domain)` + walks labels upward to the nearest cached NS RRset; returns + `[{name, ips-or-nil}]` plus newbailiwick = the NS owner name. +- **Lame referral**: a `referral` becomes `referral_lame` unless the new zone + is *strictly deeper* than the current bailiwick. + +### Glue resolution (no local resolver — ever) + +Child with serverIPs == nil: +- NS name **inside the current bailiwick** with no glue → `noglue` dead end + (probability retained on the failure). +- An ancestor referral with the same qname/qclass/qtype/server still + unresolved → `loop` dead end. +- Otherwise: **resolve subtree** for `A ` (refid `.0.` component), + starting from `getStartServers(servername)` in *this branch's* cache. Every + `answered` leaf distributes its probability evenly across the returned A + records into `serverWeights[ip]`; failed leaves carry their probability as + pseudo-IP `key:...` entries so failures surface in Results. + +### CNAME restarts + +`restart` children get qname = CNAME target, starters from the branch cache +(deepest cached zone — root only if nothing deeper cached), and the new +bailiwick from getStartServers. Loop check against the ancestor chain must +cover **every** target in a multi-record chain, not just the last. + +### Probability model (summary_stats.rb — exact) + +- serverweight = 1/len(serverIPs) per IP at referral creation. +- `percent = (1.0/len(children)) * weight`; child probability accumulates + multiplied down the tree. +- Leaf aggregation key: `key: : : : : : ` + (+ exception message for exception; + parent_ip for referral_lame; NoGlue/ + Loop use their own field order — read the Ruby). Identical keys merge by + summing probability. +- Summary groups by status; answered additionally by sorted rdata strings, so + one summary line per distinct RRset content. + +### Fast mode (default on) + +Global memo keyed +`" : : : : "` (lowercased; +txt_ips_verbose embeds per-IP weights). A completed referral with no +referral_lame response is stored; a hit replaces the child before processing +and is reported as `completed earlier ( )`. Non-fast mode +re-walks every branch. + +## Roots + +- Default: ask the upstream resolver (`--dns-upstream`, else system resolver + from /etc/resolv.conf — **not** hardcoded 127.0.0.1) for `. NS`, pick ONE + root. `--all-root-servers`: fetch the full set, create one top-level child + per root with equal weight. +- `--root-server VALUE`: accept hostname **or** IP literal. Hostname → + resolve to A via upstream; IP → use directly. Fix the current + IP-looked-up-as-hostname bug. + +## CLI output (text) — match the reference byte-for-byte where shown in spec §4.5 + +- Header block (suppressed by `--quiet`): `# Using fast mode`, + `# Limiting traverse to one root`, `# UDP size N (EDNS0 is )` + (fix the Ruby always-on bug — documented deviation), `# Retries N, max depth N`, + `# Allow TCP is , always TCP is `, then + `Using ( ) as initial root`, `Running query type `. +- Progress: ` ( , )`; verbose adds `[qname]` and + ` `; markers ` -- resolving`, ` -- completed earlier ( )`. +- Results (spec §4.5 wording catalogue, `%5.1f%%` with trailing `.0` trimmed): + `Answer from ( )` + dig-style RRs indented 12 spaces; + `No glue at ( ) for `; + `Lame referral from ( ) to ( )`; + `Loop encountered at `; `CNAME loop encountered at `; + `NODATA (for this type) at ( )`; ` at ( )`; + ` at ( )`; plus + `While querying / / ` when the failing query differs. +- Summary Results: `%5.1f%% answered with ` / + `resulted in a lame referral` / `resulted in an exception` / + `resulted in an error` / `found no such record` / `found no glue` / + `resulted in a loop` / `resulted in a CNAME loop`. +- Servers (`--show-servers`): `The following servers were encountered:`, + rows `%*s: %-15s %s`, sorted by **lowercased reversed name**. +- **Defaults change to match the reference CLI**: show-progress true, + show-resolves **false**, show-servers **false**, show-versions true, + show-all-stats **false**, show-results true, show-summary-results true. + `--show-X=false` must work (fix the truthiness-override bug). +- Colour: honour NO_COLOR **and** only colour when stdout is a TTY. +- `--json`: keep, but emit each aggregated leaf exactly once: + `{domain, qtype, root, results: [...], summary: [...], servers: [...]}`. + No duplication between results and summary. + +### Documented deviations from the Ruby reference + +We intentionally do NOT replicate these Ruby source bugs: the "Formate error" +typo, the EDNS0 banner always printing "on", the stray `)` in `Stopped at + ( ))`, and the "CNANE loop" typo. Everything else matches. + +## Web (this pass: compile + function only) + +Adapt `web/api` to the new engine API with minimal change: jobs still run +traversals and stream progress events (extend events with refid/status). +Full web output parity (Summary/Results sections, detail tree) is a later pass. +Remove SRV/CAA from the SPA type dropdown (backend rejects them) — align the +select to the supported list. + +## Testing + +- Rewrite/port unit tests so mocks inject into the **single** query path. +- Invariant test: aggregated leaf probabilities sum to 1.0 (mock topologies: + plain 2-NS answer, glueless NS, lame referral, CNAME restart, depth-limit). +- Golden harness `tools/golden/`: run `run-reference.sh` and + `bin/exploredns` on the same domain, normalize (strip ANSI, sort aggregated + result blocks, canonicalise whitespace) and diff percentages + statuses + + RRsets. Network goldens are dev tools, not CI gates. + +## Out of scope (later passes) + +fpdns-style fingerprint database, web UI output parity, geolocation/servers +map, IPv6 traversal (`--follow-aaaa` remains a documented no-op), SPF type, +distinct exit codes. diff --git a/tools/golden/compare.sh b/tools/golden/compare.sh new file mode 100755 index 0000000..87f078e --- /dev/null +++ b/tools/golden/compare.sh @@ -0,0 +1,214 @@ +#!/bin/sh +# Compare a Ruby dnstraverse capture against an ExploreDNS capture, ignoring +# run-to-run noise that is not a behaviour difference: +# - ANSI colour codes +# - which root server was picked (the reference picks one at random) +# - TTL drift between the two runs +# - RR ordering within one RRset / answer block +# - per-run child ordering (the reference shuffles children, so refids are +# assigned to different servers each run; refids are reduced to their +# structural shape and progress lines compared as a sorted multiset) +# What IS compared: aggregated Results percentages + statuses + server(ip) +# attribution, Summary lines, all wording, refid structure and the shape of +# every progress line. Server-version fingerprints are ignored (different +# fingerprint databases), but the servers-encountered name:IP rows compare. +# +# Usage: tools/golden/compare.sh REFERENCE_CAPTURE GO_CAPTURE +# Exit 0 when equivalent; exit 1 with a unified diff of the normalized forms. +set -eu + +if [ $# -ne 2 ]; then + echo "usage: $0 REFERENCE_CAPTURE GO_CAPTURE" >&2 + exit 2 +fi + +tmpdir=$(mktemp -d) +trap 'rm -rf "$tmpdir"' EXIT + +normalize() { + python3 - "$1" <<'PYEOF' +import re +import sys + +path = sys.argv[1] +with open(path, encoding="utf-8", errors="replace") as f: + text = f.read() + +# Strip ANSI escape sequences. +text = re.sub(r"\x1b\[[0-9;]*[A-Za-z]", "", text) +lines = text.split("\n") + +# The initial root is a per-run random choice: find it and substitute +# placeholders everywhere (progress line 1, servers rows, results). +rootname = rootip = None +m_root = re.compile(r"^Using (\S+) \(([^)]+)\) as initial root$") +for ln in lines: + m = m_root.match(ln) + if m: + rootname, rootip = m.group(1), m.group(2) + break + +def subst_root(s): + if rootname: + s = re.sub(re.escape(rootname), "ROOTNAME", s, flags=re.IGNORECASE) + s = s.replace(rootip, "ROOTIP") + return s + +def norm_rr(s): + # dig-style RR: name ttl class type rdata... -> collapse whitespace and + # blank the TTL so it cannot drift between the two runs. + parts = s.split() + if len(parts) >= 4 and parts[1].isdigit() and parts[2] in ("IN", "CH", "HS"): + parts[1] = "TTL" + return " ".join(parts) + +def refid_shape(refid): + # Reduce a refid to its structure: every component is 'd' except the + # literal 0 that marks a glue-resolution subtree (1.2.0.1 -> d.d.0.d). + return ".".join("0" if p == "0" else "d" for p in refid.split(".")) + +re_progress = re.compile(r"^(\d+(?:\.\d+)*) (.*)$") +re_result_head = re.compile(r"^\s*\d+(?:\.\d+)?%: ") +re_summary_head = re.compile(r"^\s*\d+(?:\.\d+)?% ") +re_completed = re.compile(r"completed earlier \((\d+(?:\.\d+)*)\)") + +header, progress, servers, results, summary = [], [], [], [], [] +section = "header" +i = 0 +while i < len(lines): + ln = lines[i].rstrip() + if ln == "The following servers were encountered:": + section = "servers" + i += 1 + continue + if ln == "Results:": + section = "results" + i += 1 + continue + if ln == "Summary Results:": + section = "summary" + i += 1 + continue + + if section == "header": + m = re_progress.match(ln) + if m: + section = "progress" + continue # reprocess as progress + if ln: + header.append(subst_root(ln)) + i += 1 + continue + + if section == "progress": + m = re_progress.match(ln) + if m: + rest = m.group(2) + # sort the IP list inside "(ip1,ip2,...)" -- RRset order noise + def sort_ips(mm): + return "(" + ",".join(sorted(mm.group(1).split(","))) + ")" + rest = re.sub(r"\(([^)]*)\)$", sort_ips, rest) + rest = re.sub( + r"\(([^)]*)\)( -- .*)$", + lambda mm: "(" + ",".join(sorted(mm.group(1).split(","))) + ")" + mm.group(2), + rest, + ) + rest = re_completed.sub( + lambda mm: "completed earlier (%s)" % refid_shape(mm.group(1)), rest + ) + progress.append("%s %s" % (refid_shape(m.group(1)), subst_root(rest))) + elif ln: + progress.append(subst_root(ln)) + i += 1 + continue + + if section == "servers": + if ln: + # " name: ip version-text" -> keep name:ip, drop versions + m = re.match(r"^\s*(\S+): (\S+)", ln) + if m: + servers.append(subst_root("%s: %s" % (m.group(1), m.group(2)))) + else: + servers.append(subst_root(ln.strip())) + i += 1 + continue + + if section == "results": + if re_result_head.match(ln): + head = subst_root(" ".join(ln.split())) + rrs, extra = [], [] + i += 1 + while i < len(lines): + nxt = lines[i].rstrip() + if not nxt or re_result_head.match(nxt) or nxt in ( + "Results:", "Summary Results:", + "The following servers were encountered:", + ): + break + if nxt.lstrip().startswith("While querying"): + extra.append(" ".join(nxt.split())) + else: + rrs.append(norm_rr(nxt)) + i += 1 + block = [head] + sorted(rrs) + extra + results.append("\n".join(" " + b if n else b for n, b in enumerate(block))) + continue + i += 1 + continue + + if section == "summary": + if re_summary_head.match(ln): + flat = " ".join(ln.split()) + m = re.match(r"^(\d+(?:\.\d+)?% answered with )(.*)$", flat) + if m: + rrs = [norm_rr(m.group(2))] + i += 1 + while i < len(lines): + nxt = lines[i].rstrip() + if not nxt or re_summary_head.match(nxt): + break + rrs.append(norm_rr(nxt)) + i += 1 + summary.append(m.group(1) + " | ".join(sorted(rrs))) + continue + summary.append(flat) + i += 1 + continue + + i += 1 + +out = [] +out.append("== HEADER ==") +out.extend(header) +out.append("== PROGRESS (sorted shapes) ==") +out.extend(sorted(progress)) +if servers: + out.append("== SERVERS (name: ip) ==") + out.extend(sorted(servers)) +out.append("== RESULTS (sorted blocks) ==") +out.extend(sorted(results)) +out.append("== SUMMARY (sorted) ==") +out.extend(sorted(summary)) + +# Sanity: aggregated Results percentages should sum to ~100. +total = 0.0 +for r in results: + m = re.match(r"^\s*(\d+(?:\.\d+)?)%:", r) + if m: + total += float(m.group(1)) +if results: + out.append("== RESULTS TOTAL ~ %d%% ==" % round(total)) + +print("\n".join(out)) +PYEOF +} + +normalize "$1" > "$tmpdir/ref.norm" +normalize "$2" > "$tmpdir/go.norm" + +if diff -u --label "reference:$1" --label "exploredns:$2" \ + "$tmpdir/ref.norm" "$tmpdir/go.norm"; then + echo "MATCH: $1 == $2 (normalized)" +else + exit 1 +fi diff --git a/tools/golden/run-reference.sh b/tools/golden/run-reference.sh new file mode 100755 index 0000000..7c04089 --- /dev/null +++ b/tools/golden/run-reference.sh @@ -0,0 +1,17 @@ +#!/bin/sh +# Run the reference Ruby dnstraverse — the exact engine behind dns.squish.net. +# Used to generate golden outputs for comparing ExploreDNS behaviour. +# +# Setup (one-time, no sudo needed): +# gem install dnsruby -v 1.72.4 --user-install --no-document +# gem install logger -v 1.6.6 --user-install --no-document +# +# The tools/dnstraverse-ruby clone is gitignored (GPL-3, dev-only). If missing: +# git clone https://github.com/squish/dnstraverse tools/dnstraverse-ruby +# then comment out the `require 'rdoc/usage'` line in bin/dnstraverse +# (rdoc/usage was removed from Ruby after 1.8). +# +# Usage: tools/golden/run-reference.sh [dnstraverse options] DOMAIN +set -eu +DIR="$(cd "$(dirname "$0")/../dnstraverse-ruby" && pwd)" +exec ruby -I"$DIR/lib" "$DIR/bin/dnstraverse" "$@" -- 2.54.0 From d71c7fbef2424fe9954bf8893e4ea94d7f05edb7 Mon Sep 17 00:00:00 2001 From: Gary Hansen Date: Tue, 7 Jul 2026 21:42:06 +1000 Subject: [PATCH 2/6] feat: rework engine and CLI for dnstraverse parity Port the traversal engine to the Ruby dnstraverse model so behaviour and output match dns.squish.net: - dns: single RD=0 query path (RD=1 only for upstream root discovery), per-run packet cache, EDNS0 512-fallback with warnings, UDP->TCP on truncation; fix --retries 0 and --root-server IP-literal handling; drop all hardcoded 127.0.0.1:53 resolvers - traverse: hierarchical per-branch InfoCache, 7-step response classification with the full 10-status vocabulary, bailiwick partitioning, strictly-deeper lame-referral rule, refid grammar with .0 resolve subtrees and childset digits, per-IP branching at 1/n weight, cache-based glue resolution with noglue/loop dead ends, CNAME restarts from the deepest cached zone, fast-mode memoization, probability aggregation with Ruby-identical stats keys (sums to 1.0) - output: byte-for-byte reference text format pinned by a golden test, reference CLI defaults, working --quiet/--show-X=false, TTY-aware colour, deduplicated deterministic JSON - web: adapt API/SPA to the new engine, SSE events carry refid/status, fix subscribe/snapshot duplicate-event race and a statusCls TDZ bug, align SPA type list with the backend - delete the old engine and dead code (net -4,350 lines) Verified against live runs of the reference Ruby engine across five domains (answers, NXDOMAIN, null MX, CNAME restart, glueless resolve) with no divergences beyond the documented typo fixes. Co-Authored-By: Claude Fable 5 --- README.md | 74 +- cmd/exploredns/main.go | 122 +-- internal/config/config.go | 131 +-- internal/config/config_test.go | 185 +--- internal/dns/dns.go | 1 - internal/dns/hints_test.go | 49 +- internal/dns/iterative_test.go | 122 --- internal/dns/query.go | 391 ++++--- internal/dns/query_test.go | 755 +++++++------- internal/dns/real_exchange_test.go | 363 ++----- internal/dns/resolver.go | 198 ---- internal/dns/resolver_test.go | 476 --------- internal/dns/robustness_test.go | 34 +- internal/dns/roots.go | 272 +++-- internal/dns/roots_test.go | 684 ++++++------ internal/dns/types.go | 29 +- internal/dns/types_test.go | 37 - internal/integration/integration_test.go | 630 ++++------- internal/output/coverage_test.go | 877 ---------------- internal/output/formatter.go | 84 +- internal/output/formatter_test.go | 681 ++++++------ internal/output/golden_test.go | 182 ++++ internal/output/json.go | 183 ++-- internal/output/runner.go | 52 +- internal/output/stats.go | 250 +---- internal/output/stats_test.go | 327 ------ internal/output/text.go | 398 ++++--- internal/output/text_test.go | 379 ------- internal/traverse/cache.go | 229 ++-- internal/traverse/cache_test.go | 319 +++--- internal/traverse/coverage_test.go | 727 ------------- internal/traverse/decoded_query.go | 259 +++++ internal/traverse/decoded_query_test.go | 359 +++++++ internal/traverse/hooks.go | 72 +- internal/traverse/hooks_test.go | 93 +- internal/traverse/referral.go | 779 ++++++++++---- internal/traverse/referral_test.go | 317 +++--- internal/traverse/response.go | 256 ----- internal/traverse/response_test.go | 340 ------ internal/traverse/robustness_test.go | 768 -------------- internal/traverse/server_response.go | 187 ++++ internal/traverse/server_response_test.go | 232 +++++ internal/traverse/stack.go | 58 -- internal/traverse/stack_test.go | 143 --- internal/traverse/stats.go | 74 ++ internal/traverse/stats_test.go | 302 ++++++ internal/traverse/traverse.go | 40 +- internal/traverse/traverser.go | 671 +++++------- internal/traverse/traverser_test.go | 1158 +++++++++++++-------- web/api/handler.go | 221 ++-- web/api/handler_internal_test.go | 77 ++ web/api/handler_test.go | 40 + web/api/static/index.html | 364 ++++--- 53 files changed, 6685 insertions(+), 9366 deletions(-) delete mode 100644 internal/dns/dns.go delete mode 100644 internal/dns/iterative_test.go delete mode 100644 internal/dns/resolver.go delete mode 100644 internal/dns/resolver_test.go delete mode 100644 internal/output/coverage_test.go create mode 100644 internal/output/golden_test.go delete mode 100644 internal/output/stats_test.go delete mode 100644 internal/output/text_test.go delete mode 100644 internal/traverse/coverage_test.go create mode 100644 internal/traverse/decoded_query.go create mode 100644 internal/traverse/decoded_query_test.go delete mode 100644 internal/traverse/response.go delete mode 100644 internal/traverse/response_test.go delete mode 100644 internal/traverse/robustness_test.go create mode 100644 internal/traverse/server_response.go create mode 100644 internal/traverse/server_response_test.go delete mode 100644 internal/traverse/stack.go delete mode 100644 internal/traverse/stack_test.go create mode 100644 internal/traverse/stats.go create mode 100644 internal/traverse/stats_test.go create mode 100644 web/api/handler_internal_test.go diff --git a/README.md b/README.md index 6f5fc0b..3e923c4 100644 --- a/README.md +++ b/README.md @@ -80,7 +80,7 @@ exploredns --quiet www.example.com # Debug mode exploredns --debug www.example.com -# Library-level debug (very verbose) +# Debug plus library-level diagnostics exploredns --dd www.example.com # Force TCP @@ -102,45 +102,44 @@ Usage: exploredns [flags] Query Options: - --type Record type to query (default: A) + --type Record type to query (default: a) Supported: A, AAAA, NS, CNAME, MX, TXT, SOA, PTR, ANY - --root-server Override the root server IP address - --all-root-servers Query all 13 root server sets (default: false) - --root-aaaa Include IPv6 addresses for root servers (default: false) - --follow-aaaa Only follow AAAA addresses for referrals (default: false) + --root-server Initial root server, hostname or IP literal + (default: ask the upstream resolver for one root) + --all-root-servers Traverse from all root servers (default: false) + --root-aaaa Include IPv6 root addresses (not implemented yet) + --follow-aaaa Only follow AAAA for referrals (not implemented yet) + --dns-upstream Upstream resolver (host:port) for root discovery + (default: system resolver) Transport Options: - --udp-size EDNS0 UDP buffer size, 512–4096 (default: 2048) + --udp-size EDNS0 UDP buffer size, 512–4096; 512 turns EDNS0 off + (default: 2048) --allow-tcp Fall back to TCP on truncation (default: true) --always-tcp Always use TCP (requires --allow-tcp) - --retries Per-server retry count, 0–10 (default: 2) + --retries Number of 2s retries before timing out, 0–10 (default: 2) Traversal Options: --max-depth Maximum referral depth, 1–100 (default: 20) - --fast / --fast=false Share glue cache across branches (default: true) + --fast / --fast=false Fast mode; turn off to be more accurate (default: true) Output Options: - --json Emit results as JSON instead of text - --verbose, -v Show extra detail in text output - --debug, -d Enable application debug messages (stderr) - --dd Enable library-level debug messages (very verbose) - --quiet, -q Suppress header and supplementary information - --show-progress Show live traversal progress (default: true) - --no-show-progress Hide traversal progress - --show-resolves Show glue-resolution steps (default: true) - --no-show-resolves Hide glue-resolution steps - --show-servers Show which servers were queried (default: true) - --no-show-servers Hide server list - --show-versions Show DNS server software versions (default: true) - --no-show-versions Hide server versions - --show-all-stats Show query statistics (default: true) - --no-show-all-stats Hide statistics - --show-results Show per-branch query results (default: true) - --no-show-results Hide per-branch results - --show-summary-results Show deduplicated summary section (default: true) - --no-show-summary-results Hide summary section + --json Emit a single JSON document instead of text + --verbose, -v Verbose progress ([qname] and shown) + -d, --debug Print debug diagnostics to stderr + -dd Like -d plus library-level debug + --quiet, -q Suppress the header block + --show-progress Show traversal progress (default: true) + --show-resolves Show glue-resolution progress (default: false) + --show-servers Show servers encountered (default: false) + --show-versions Show server version fingerprints (default: true) + --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) ``` +Every `--show-X` flag can be negated with `--show-X=false` or `--no-show-X`. + --- ## Web Interface @@ -253,15 +252,22 @@ Completed jobs are kept in memory for one hour before being purged. ### Text (default) -Coloured, hierarchical tree output showing each traversal branch, the servers -queried, referrals followed, and final answers. Disable colour by setting the -`NO_COLOR` environment variable. +dnstraverse-style output: a header block (settings, initial root, query; +suppressed by `--quiet`), progress lines (` ( )` with +` -- resolving` and ` -- completed earlier ( )` markers), a `Results:` +section of aggregated outcomes with probabilities (`Answer from`, `No glue +at`, `Lame referral from`, error/exception wording), and a `Summary Results:` +section grouping outcomes by status and answer content. `--show-servers` +adds the sorted list of servers encountered. Colour is used only when stdout +is a terminal and the `NO_COLOR` environment variable is unset. ### JSON (`--json`) -Structured JSON array of traversal results. Suitable for piping into `jq` or -ingesting into other tools. Each element contains the referral metadata, the -responding server, the response type, and the decoded DNS records. +A single JSON document — `{domain, qtype, root, results, summary, servers}` — +emitted once at the end of the run with deterministic ordering. Each +aggregated outcome appears exactly once in `results`; `summary` groups +probabilities by status and by distinct answer RRset; `servers` is present +with `--show-servers`. Suitable for piping into `jq`. --- diff --git a/cmd/exploredns/main.go b/cmd/exploredns/main.go index 12594ab..90d60a8 100644 --- a/cmd/exploredns/main.go +++ b/cmd/exploredns/main.go @@ -17,17 +17,17 @@ func main() { cfg := config.DefaultConfig() queryType := flag.String("type", cfg.QueryType, "Record type (A, AAAA, NS, CNAME, MX, TXT, SOA, PTR, ANY)") - rootServer := flag.String("root-server", cfg.RootServer, "Override root server") - allRootServers := flag.Bool("all-root-servers", cfg.AllRootServers, "Use all 13 root servers") - rootAAAA := flag.Bool("root-aaaa", cfg.RootAAAA, "Include IPv6 root addresses") - followAAAA := flag.Bool("follow-aaaa", cfg.FollowAAAA, "Only follow AAAA for referrals") + rootServer := flag.String("root-server", cfg.RootServer, "Initial root server (hostname or IP)") + allRootServers := flag.Bool("all-root-servers", cfg.AllRootServers, "Traverse from all root servers") + rootAAAA := flag.Bool("root-aaaa", cfg.RootAAAA, "Include IPv6 root addresses (not implemented yet)") + followAAAA := flag.Bool("follow-aaaa", cfg.FollowAAAA, "Only follow AAAA for referrals (not implemented yet)") dnsUpstream := flag.String("dns-upstream", cfg.DNSUpstream, "Upstream resolver for root discovery (e.g. 8.8.8.8:53, default: system)") - udpSize := flag.Int("udp-size", cfg.UDPSize, "EDNS0 buffer size (512-4096)") + udpSize := flag.Int("udp-size", cfg.UDPSize, "EDNS0 buffer size (512-4096; 512 turns EDNS0 off)") allowTCP := flag.Bool("allow-tcp", cfg.AllowTCP, "TCP fallback on truncation") alwaysTCP := flag.Bool("always-tcp", cfg.AlwaysTCP, "Always use TCP") maxDepth := flag.Int("max-depth", cfg.MaxDepth, "Max traversal depth (1-100)") - retries := flag.Int("retries", cfg.Retries, "Retry count (0-10)") - fast := flag.Bool("fast", cfg.Fast, "Fast mode: reuse earlier branch cache") + retries := flag.Int("retries", cfg.Retries, "Number of 2s retries before timing out (0-10)") + fast := flag.Bool("fast", cfg.Fast, "Fast mode; turn off to be more accurate") jsonOutput := flag.Bool("json", false, "Output results as JSON") // Verbose: long and short form share the same variable. @@ -35,27 +35,31 @@ func main() { flag.BoolVar(&verboseVal, "verbose", cfg.Verbose, "Verbose output") flag.BoolVar(&verboseVal, "v", cfg.Verbose, "Verbose output (shorthand)") - // Debug: -d sets level 1, -dd sets level 2 (library debug). + // Debug: -d prints debug diagnostics; -dd additionally enables library + // debug (the Go flag package cannot stack -d -d like the Ruby CLI). var dFlag, ddFlag bool - flag.BoolVar(&dFlag, "d", false, "Debug mode (stackable: -dd for library debug)") - flag.BoolVar(&dFlag, "debug", false, "Debug mode") - flag.BoolVar(&ddFlag, "dd", false, "Library debug mode (equivalent to -d -d)") + flag.BoolVar(&dFlag, "d", false, "Print debug diagnostics to stderr") + flag.BoolVar(&dFlag, "debug", false, "Print debug diagnostics to stderr") + flag.BoolVar(&ddFlag, "dd", false, "Like -d plus library-level debug") // Quiet: long and short form share the same variable. var quietVal bool - flag.BoolVar(&quietVal, "quiet", cfg.Quiet, "Suppress supplementary info") - flag.BoolVar(&quietVal, "q", cfg.Quiet, "Suppress supplementary info (shorthand)") + flag.BoolVar(&quietVal, "quiet", cfg.Quiet, "Suppress the header block") + flag.BoolVar(&quietVal, "q", cfg.Quiet, "Suppress the header block (shorthand)") + // Each show flag defaults to the config default, so --show-X and + // --show-X=false both take effect directly; the --no-show-X aliases force + // the value off, matching the Ruby --[no-]show-X switches. showProgress := flag.Bool("show-progress", cfg.ShowProgress, "Show traversal progress") noShowProgress := flag.Bool("no-show-progress", false, "Hide traversal progress") - showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution details") - noShowResolves := flag.Bool("no-show-resolves", false, "Hide glue resolution details") - showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers queried") - noShowServers := flag.Bool("no-show-servers", false, "Hide servers queried") - showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server versions") - noShowVersions := flag.Bool("no-show-versions", false, "Hide server versions") - showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show all statistics") - noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide all statistics") + showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution progress") + noShowResolves := flag.Bool("no-show-resolves", false, "Hide glue resolution progress") + showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers encountered") + noShowServers := flag.Bool("no-show-servers", false, "Hide servers encountered") + showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server version fingerprints") + noShowVersions := flag.Bool("no-show-versions", false, "Hide server version fingerprints") + showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show statistics after every node") + noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide per-node statistics") showResults := flag.Bool("show-results", cfg.ShowResults, "Show query results") noShowResults := flag.Bool("no-show-results", false, "Hide query results") showSummaryResults := flag.Bool("show-summary-results", cfg.ShowSummaryResults, "Show summary of results") @@ -83,41 +87,13 @@ func main() { cfg.Debug = config.ParseDebugLevel(dFlag, ddFlag) cfg.Quiet = quietVal - if *noShowProgress { - cfg.ShowProgress = false - } else if *showProgress { - cfg.ShowProgress = true - } - if *noShowResolves { - cfg.ShowResolves = false - } else if *showResolves { - cfg.ShowResolves = true - } - if *noShowServers { - cfg.ShowServers = false - } else if *showServers { - cfg.ShowServers = true - } - if *noShowVersions { - cfg.ShowVersions = false - } else if *showVersions { - cfg.ShowVersions = true - } - if *noShowAllStats { - cfg.ShowAllStats = false - } else if *showAllStats { - cfg.ShowAllStats = true - } - if *noShowResults { - cfg.ShowResults = false - } else if *showResults { - cfg.ShowResults = true - } - if *noShowSummaryResults { - cfg.ShowSummaryResults = false - } else if *showSummaryResults { - cfg.ShowSummaryResults = true - } + cfg.ShowProgress = *showProgress && !*noShowProgress + cfg.ShowResolves = *showResolves && !*noShowResolves + cfg.ShowServers = *showServers && !*noShowServers + cfg.ShowVersions = *showVersions && !*noShowVersions + cfg.ShowAllStats = *showAllStats && !*noShowAllStats + cfg.ShowResults = *showResults && !*noShowResults + cfg.ShowSummaryResults = *showSummaryResults && !*noShowSummaryResults if err := cfg.Validate(); err != nil { fmt.Fprintf(os.Stderr, "Error: %v\n", err) @@ -130,26 +106,15 @@ func main() { os.Exit(1) } - rootIP, err := cfg.ParseRootServer() - if err != nil { - fmt.Fprintf(os.Stderr, "Error: invalid root server: %v\n", err) - os.Exit(1) - } - queryTypeValue, err := config.ParseQueryType(cfg.QueryType) if err != nil { fmt.Fprintf(os.Stderr, "Error: %v\n", err) os.Exit(1) } - var rootServerAddr string - if rootIP != nil { - rootServerAddr = rootIP.String() - } - queryConfig := &dns.QueryConfig{ UDPSize: cfg.UDPSize, - Timeout: 5 * time.Second, + Timeout: 2 * time.Second, Retries: cfg.Retries, UseTCP: cfg.AlwaysTCP, AllowTCP: cfg.AllowTCP, @@ -157,9 +122,11 @@ func main() { rootConfig := &dns.RootDiscoveryConfig{ IncludeAAAA: cfg.RootAAAA, - Server: rootServerAddr, - AllRoots: cfg.AllRootServers, - Resolver: cfg.DNSUpstream, + // --root-server accepts a hostname or an IP literal. + Server: cfg.RootServer, + AllRoots: cfg.AllRootServers, + Resolver: cfg.DNSUpstream, + Query: queryConfig, } traverserConfig := &traverse.TraverserConfig{ @@ -187,10 +154,6 @@ func main() { } } - if !cfg.Quiet && !*jsonOutput { - fmt.Printf("ExploreDNS - exploring: %s (type: %s)\n", domain, cfg.QueryType) - } - outFmt := output.FormatText if *jsonOutput { outFmt = output.FormatJSON @@ -200,6 +163,13 @@ func main() { Format: outFmt, Domain: domain, QueryType: cfg.QueryType, + Fast: cfg.Fast, + AllRootServers: cfg.AllRootServers, + UDPSize: cfg.UDPSize, + Retries: cfg.Retries, + MaxDepth: cfg.MaxDepth, + AllowTCP: cfg.AllowTCP, + AlwaysTCP: cfg.AlwaysTCP, ShowProgress: cfg.ShowProgress, ShowResolves: cfg.ShowResolves, ShowServers: cfg.ShowServers, @@ -209,7 +179,7 @@ func main() { ShowSummaryResults: cfg.ShowSummaryResults, Verbose: cfg.Verbose, Quiet: cfg.Quiet, - Color: os.Getenv("NO_COLOR") == "", + Color: output.ColorEnabled(os.Stdout), Debug: cfg.Debug, } diff --git a/internal/config/config.go b/internal/config/config.go index a8abe5c..4794f41 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -19,7 +19,6 @@ var ( ErrInvalidRetries = errors.New("retries must be between 0 and 10") ErrMissingDomain = errors.New("domain is required") ErrAlwaysTCPRequiresTCP = errors.New("--always-tcp requires --allow-tcp") - ErrInvalidRootServer = errors.New("invalid root server IP address") ) type Config struct { @@ -76,39 +75,6 @@ func ParseQueryType(s string) (uint16, error) { } } -func ParseUDPSize(s string) (int, error) { - var size int - if _, err := fmt.Sscanf(s, "%d", &size); err != nil { - return 0, fmt.Errorf("%w: %s", ErrInvalidUDPSize, s) - } - if size < 512 || size > 4096 { - return 0, fmt.Errorf("%w: %d (must be 512-4096)", ErrInvalidUDPSize, size) - } - return size, nil -} - -func ParseMaxDepth(s string) (int, error) { - var depth int - if _, err := fmt.Sscanf(s, "%d", &depth); err != nil { - return 0, fmt.Errorf("%w: %s", ErrInvalidMaxDepth, s) - } - if depth < 1 || depth > 100 { - return 0, fmt.Errorf("%w: %d (must be 1-100)", ErrInvalidMaxDepth, depth) - } - return depth, nil -} - -func ParseRetries(s string) (int, error) { - var retries int - if _, err := fmt.Sscanf(s, "%d", &retries); err != nil { - return 0, fmt.Errorf("%w: %s", ErrInvalidRetries, s) - } - if retries < 0 || retries > 10 { - return 0, fmt.Errorf("%w: %d (must be 0-10)", ErrInvalidRetries, retries) - } - return retries, nil -} - func (c *Config) Validate() error { if _, err := ParseQueryType(c.QueryType); err != nil { return err @@ -146,17 +112,6 @@ func (c *Config) GetDomain(args []string) (string, error) { return args[0], nil } -func (c *Config) ParseRootServer() (net.IP, error) { - if c.RootServer == "" { - return nil, nil - } - ip := net.ParseIP(c.RootServer) - if ip == nil { - return nil, fmt.Errorf("%w: %s", ErrInvalidRootServer, c.RootServer) - } - return ip, nil -} - // ParseDebugLevel returns the debug verbosity level from the -d and -dd flag values. // dd=true → 2 (library debug), d=true → 1 (application debug), neither → 0. func ParseDebugLevel(d, dd bool) int { @@ -175,57 +130,59 @@ func PrintUsage() { fmt.Fprintf(os.Stderr, " exploredns [flags] \n\n") fmt.Fprintf(os.Stderr, "Flags:\n") - flagGroups := map[string][][2]string{ - "Query Options": { - {"--type", "Record type (default A)"}, - {"--root-server", "Override root server"}, - {"--all-root-servers", "Use all 13 root servers"}, - {"--root-aaaa", "Include IPv6 root addresses"}, - {"--follow-aaaa", "Only follow AAAA for referrals"}, + // Ordered slice, not a map: --help output must be stable. + flagGroups := []struct { + name string + flags [][2]string + }{ + {"Query Options", [][2]string{ + {"--type", "Record type (default a)"}, + {"--root-server", "Initial root server, hostname or IP (default: ask upstream)"}, + {"--all-root-servers", "Traverse from all root servers (default false)"}, + {"--root-aaaa", "Include IPv6 root addresses (not implemented yet)"}, + {"--follow-aaaa", "Only follow AAAA for referrals (not implemented yet)"}, {"--dns-upstream", "Upstream resolver for root discovery (default: system)"}, - }, - "Transport Options": { - {"--udp-size", "EDNS0 buffer size (default 2048)"}, + }}, + {"Transport Options", [][2]string{ + {"--udp-size", "EDNS0 buffer size; 512 turns EDNS0 off (default 2048)"}, {"--allow-tcp", "TCP fallback on truncation (default true)"}, - {"--always-tcp", "Always use TCP"}, - {"--retries", "Retry count (default 2)"}, - }, - "Traversal Options": { + {"--always-tcp", "Always use TCP (default false)"}, + {"--retries", "Number of 2s retries before timing out (default 2)"}, + }}, + {"Traversal Options", [][2]string{ {"--max-depth", "Max traversal depth (default 20)"}, - {"--fast", "Fast mode: reuse earlier branch cache (default true)"}, - }, - "Output Options": { - {"--verbose, -v", "Verbose output"}, - {"--debug, -d", "Debug mode (stackable: -dd for library debug)"}, - {"--quiet, -q", "Suppress supplementary info"}, - {"--show-progress", "Show traversal progress"}, - {"--no-show-progress", "Hide traversal progress"}, - {"--show-resolves", "Show glue resolution details"}, - {"--no-show-resolves", "Hide glue resolution details"}, - {"--show-servers", "Show servers queried"}, - {"--no-show-servers", "Hide servers queried"}, - {"--show-versions", "Show server versions"}, - {"--no-show-versions", "Hide server versions"}, - {"--show-all-stats", "Show all statistics"}, - {"--no-show-all-stats", "Hide all statistics"}, - {"--show-results", "Show query results"}, - {"--no-show-results", "Hide query results"}, - {"--show-summary-results", "Show summary of results"}, - {"--no-show-summary-results", "Hide summary of results"}, - }, + {"--fast", "Fast mode; turn off to be more accurate (default true)"}, + }}, + {"Output Options", [][2]string{ + {"--json", "Emit a single JSON document instead of text"}, + {"--verbose, -v", "Verbose progress ([qname] and shown)"}, + {"-d, --debug", "Print debug diagnostics to stderr"}, + {"-dd", "Like -d plus library-level debug"}, + {"--quiet, -q", "Suppress the header block"}, + {"--show-progress", "Show traversal progress (default true)"}, + {"--show-resolves", "Show glue resolution progress (default false)"}, + {"--show-servers", "Show servers encountered (default false)"}, + {"--show-versions", "Show server version fingerprints (default true)"}, + {"--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)"}, + }}, } - for group, flags := range flagGroups { - fmt.Fprintf(os.Stderr, "\n%s:\n", group) - for _, f := range flags { + for _, group := range flagGroups { + fmt.Fprintf(os.Stderr, "\n%s:\n", group.name) + for _, f := range group.flags { fmt.Fprintf(os.Stderr, " %-25s %s\n", f[0], f[1]) } } + fmt.Fprintf(os.Stderr, "\nEvery --show-X flag can be negated with --show-X=false or --no-show-X.\n") } func DefaultConfig() *Config { return &Config{ - QueryType: "A", + // Lowercase like the reference default (:a); explicit --type values + // keep the user's case in the "Running query" line. + QueryType: "a", RootServer: "", AllRootServers: false, RootAAAA: false, @@ -241,10 +198,10 @@ func DefaultConfig() *Config { Debug: 0, Quiet: false, ShowProgress: true, - ShowResolves: true, - ShowServers: true, + ShowResolves: false, + ShowServers: false, ShowVersions: true, - ShowAllStats: true, + ShowAllStats: false, ShowResults: true, ShowSummaryResults: true, } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 4a2a97f..020adc5 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -41,28 +41,6 @@ func TestParseQueryTypeInvalid(t *testing.T) { } } -func TestParseUDPSizeValid(t *testing.T) { - cases := []string{"512", "2048", "4096"} - for _, s := range cases { - if _, err := ParseUDPSize(s); err != nil { - t.Errorf("ParseUDPSize(%q) unexpected error: %v", s, err) - } - } -} - -func TestParseUDPSizeInvalid(t *testing.T) { - cases := []string{"0", "511", "4097", "notanumber"} - for _, s := range cases { - _, err := ParseUDPSize(s) - if err == nil { - t.Errorf("ParseUDPSize(%q): expected error", s) - } - if !errors.Is(err, ErrInvalidUDPSize) { - t.Errorf("ParseUDPSize(%q): expected ErrInvalidUDPSize, got %v", s, err) - } - } -} - func TestValidateOK(t *testing.T) { cfg := DefaultConfig() if err := cfg.Validate(); err != nil { @@ -138,65 +116,6 @@ func TestGetDomainMissing(t *testing.T) { } } -func TestParseRootServerEmpty(t *testing.T) { - cfg := DefaultConfig() - ip, err := cfg.ParseRootServer() - if err != nil { - t.Fatalf("unexpected error for empty root server: %v", err) - } - if ip != nil { - t.Errorf("expected nil IP for empty root server, got %v", ip) - } -} - -func TestParseRootServerValidIPv4(t *testing.T) { - cfg := DefaultConfig() - cfg.RootServer = "198.41.0.4" - ip, err := cfg.ParseRootServer() - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if ip == nil || ip.String() != "198.41.0.4" { - t.Errorf("expected 198.41.0.4, got %v", ip) - } -} - -func TestParseRootServerValidIPv6(t *testing.T) { - cfg := DefaultConfig() - cfg.RootServer = "2001:503:ba3e::2:30" - ip, err := cfg.ParseRootServer() - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if ip == nil { - t.Error("expected non-nil IP for valid IPv6 address") - } -} - -func TestParseRootServerInvalid(t *testing.T) { - cfg := DefaultConfig() - cfg.RootServer = "not-an-ip" - _, err := cfg.ParseRootServer() - if err == nil { - t.Fatal("expected error for invalid root server IP") - } - if !errors.Is(err, ErrInvalidRootServer) { - t.Errorf("expected ErrInvalidRootServer, got %v", err) - } -} - -func TestParseRootServerHostname(t *testing.T) { - cfg := DefaultConfig() - cfg.RootServer = "a.root-servers.net" - _, err := cfg.ParseRootServer() - if err == nil { - t.Fatal("expected error for hostname (not IP) root server") - } - if !errors.Is(err, ErrInvalidRootServer) { - t.Errorf("expected ErrInvalidRootServer, got %v", err) - } -} - func TestParseDebugLevel(t *testing.T) { cases := []struct { d, dd bool @@ -215,96 +134,34 @@ func TestParseDebugLevel(t *testing.T) { } } -func TestParseMaxDepthValid(t *testing.T) { - cases := []struct { - input string - want int - }{ - {"1", 1}, - {"20", 20}, - {"100", 100}, - } - for _, tc := range cases { - got, err := ParseMaxDepth(tc.input) - if err != nil { - t.Errorf("ParseMaxDepth(%q) unexpected error: %v", tc.input, err) - } - if got != tc.want { - t.Errorf("ParseMaxDepth(%q) = %d, want %d", tc.input, got, tc.want) - } - } -} - -func TestParseMaxDepthInvalid(t *testing.T) { - cases := []string{"0", "101", "notanumber", "-1"} - for _, s := range cases { - _, err := ParseMaxDepth(s) - if err == nil { - t.Errorf("ParseMaxDepth(%q): expected error", s) - continue - } - if !errors.Is(err, ErrInvalidMaxDepth) { - t.Errorf("ParseMaxDepth(%q): expected ErrInvalidMaxDepth, got %v", s, err) - } - } -} - -func TestParseRetriesValid(t *testing.T) { - cases := []struct { - input string - want int - }{ - {"0", 0}, - {"2", 2}, - {"10", 10}, - } - for _, tc := range cases { - got, err := ParseRetries(tc.input) - if err != nil { - t.Errorf("ParseRetries(%q) unexpected error: %v", tc.input, err) - } - if got != tc.want { - t.Errorf("ParseRetries(%q) = %d, want %d", tc.input, got, tc.want) - } - } -} - -func TestParseRetriesInvalid(t *testing.T) { - cases := []string{"-1", "11", "notanumber"} - for _, s := range cases { - _, err := ParseRetries(s) - if err == nil { - t.Errorf("ParseRetries(%q): expected error", s) - continue - } - if !errors.Is(err, ErrInvalidRetries) { - t.Errorf("ParseRetries(%q): expected ErrInvalidRetries, got %v", s, err) - } - } -} - func TestPrintUsage(t *testing.T) { // PrintUsage writes to stderr; just ensure it doesn't panic. PrintUsage() } func TestValidateBadQueryType(t *testing.T) { -cfg := DefaultConfig() -cfg.QueryType = "BOGUS" -if err := cfg.Validate(); !errors.Is(err, ErrInvalidQueryType) { -t.Errorf("expected ErrInvalidQueryType, got %v", err) -} + cfg := DefaultConfig() + cfg.QueryType = "BOGUS" + if err := cfg.Validate(); !errors.Is(err, ErrInvalidQueryType) { + t.Errorf("expected ErrInvalidQueryType, got %v", err) + } } func TestDefaultConfigIsValid(t *testing.T) { -cfg := DefaultConfig() -if cfg.QueryType != "A" { -t.Errorf("QueryType = %q, want A", cfg.QueryType) -} -if cfg.MaxDepth != 20 { -t.Errorf("MaxDepth = %d, want 20", cfg.MaxDepth) -} -if cfg.Retries != 2 { -t.Errorf("Retries = %d, want 2", cfg.Retries) -} + cfg := DefaultConfig() + if cfg.QueryType != "a" { + t.Errorf("QueryType = %q, want a", cfg.QueryType) + } + if cfg.ShowResolves || cfg.ShowServers || cfg.ShowAllStats { + t.Error("show-resolves/show-servers/show-all-stats must default to false") + } + if !cfg.ShowProgress || !cfg.ShowVersions || !cfg.ShowResults || !cfg.ShowSummaryResults { + t.Error("show-progress/show-versions/show-results/show-summary-results must default to true") + } + if cfg.MaxDepth != 20 { + t.Errorf("MaxDepth = %d, want 20", cfg.MaxDepth) + } + if cfg.Retries != 2 { + t.Errorf("Retries = %d, want 2", cfg.Retries) + } } diff --git a/internal/dns/dns.go b/internal/dns/dns.go deleted file mode 100644 index 1ffe03d..0000000 --- a/internal/dns/dns.go +++ /dev/null @@ -1 +0,0 @@ -package dns diff --git a/internal/dns/hints_test.go b/internal/dns/hints_test.go index 307f87c..99538ca 100644 --- a/internal/dns/hints_test.go +++ b/internal/dns/hints_test.go @@ -111,47 +111,28 @@ func TestRootHintsKnownAddress(t *testing.T) { func TestResolverFromConfigExplicit(t *testing.T) { cfg := &RootDiscoveryConfig{Resolver: "8.8.8.8:53"} - got := resolverFromConfig(cfg) + got, err := resolverFromConfig(cfg) + if err != nil { + t.Fatalf("resolverFromConfig: %v", err) + } if got != "8.8.8.8:53" { t.Errorf("resolverFromConfig = %q, want 8.8.8.8:53", got) } } -func TestResolverFromConfigEmpty(t *testing.T) { - cfg := &RootDiscoveryConfig{} - got := resolverFromConfig(cfg) - // Should return the system resolver; just check it's non-empty and contains a port. - if got == "" { - t.Error("resolverFromConfig with empty Resolver returned empty string") - } -} - -func TestResolverFromConfigNil(t *testing.T) { - got := resolverFromConfig(nil) - if got == "" { - t.Error("resolverFromConfig(nil) returned empty string") - } -} - -func TestSystemResolverNonEmpty(t *testing.T) { - got := systemResolver() - if got == "" { - t.Error("systemResolver() returned empty string") - } - // Must contain a colon (host:port format). - host, port, err := splitHostPort(got) +func TestSystemResolverNoHardcodedFallback(t *testing.T) { + // Whether or not /etc/resolv.conf exists, systemResolver must never + // invent 127.0.0.1:53: it returns a valid host:port from the system + // configuration or an error. + got, err := systemResolver() if err != nil { - t.Errorf("systemResolver() = %q: not a valid host:port: %v", got, err) + return } - if host == "" { - t.Errorf("systemResolver() host is empty in %q", got) + host, port, err := net.SplitHostPort(got) + if err != nil { + t.Fatalf("systemResolver() = %q: not a valid host:port: %v", got, err) } - if port == "" { - t.Errorf("systemResolver() port is empty in %q", got) + if host == "" || port == "" { + t.Errorf("systemResolver() returned incomplete address %q", got) } } - -// splitHostPort is a thin wrapper around net.SplitHostPort for test use. -func splitHostPort(addr string) (host, port string, err error) { - return net.SplitHostPort(addr) -} diff --git a/internal/dns/iterative_test.go b/internal/dns/iterative_test.go deleted file mode 100644 index 854064c..0000000 --- a/internal/dns/iterative_test.go +++ /dev/null @@ -1,122 +0,0 @@ -package dns - -import ( - "context" - "errors" - "net" - "sync" - "testing" - - "github.com/miekg/dns" -) - -func TestIterativeQueryWithExchangeNilConfig(t *testing.T) { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return resp.Copy(), nil - } - - server := net.ParseIP("8.8.8.8") - _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, nil, exchangeFn) - if err != nil { - t.Fatalf("unexpected error with nil config: %v", err) - } -} - -func TestIterativeQueryWithExchangeAlwaysTCP(t *testing.T) { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - - var mu sync.Mutex - calls := []bool{} - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - mu.Lock() - calls = append(calls, useTCP) - mu.Unlock() - return resp.Copy(), nil - } - - cfg := &QueryConfig{ - UDPSize: 2048, - Retries: 1, - UseTCP: true, - } - - server := net.ParseIP("8.8.8.8") - _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(calls) != 1 || !calls[0] { - t.Errorf("expected single TCP call, got %v", calls) - } -} - -func TestIterativeQueryWithExchangeRetriesOnFailure(t *testing.T) { - var mu sync.Mutex - callCount := 0 - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - mu.Lock() - callCount++ - mu.Unlock() - return nil, errors.New("connection refused") - } - - cfg := &QueryConfig{ - UDPSize: 2048, - Retries: 3, - UseTCP: false, - } - - server := net.ParseIP("8.8.8.8") - _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) - if err == nil { - t.Fatal("expected error") - } - if callCount != 3 { - t.Errorf("expected 3 calls, got %d", callCount) - } -} - -func TestIterativeQueryWithExchangeContextCancel(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - cancel() - - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, ctx.Err() - } - - cfg := &QueryConfig{ - UDPSize: 2048, - Retries: 1, - UseTCP: false, - } - - server := net.ParseIP("8.8.8.8") - _, err := IterativeQueryWithExchange(ctx, server, "example.com", TypeA, cfg, exchangeFn) - if err == nil { - t.Fatal("expected error on cancelled context") - } -} - -func TestIterativeQueryWithExchangeZeroUDPSize(t *testing.T) { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return resp.Copy(), nil - } - - cfg := &QueryConfig{ - UDPSize: 0, // should use default - Retries: 1, - } - - server := net.ParseIP("8.8.8.8") - _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } -} diff --git a/internal/dns/query.go b/internal/dns/query.go index 389cb5f..3433be2 100644 --- a/internal/dns/query.go +++ b/internal/dns/query.go @@ -1,73 +1,98 @@ // Package dns provides the low-level DNS query primitives used by ExploreDNS. // -// It wraps the github.com/miekg/dns library to provide retrying, TCP fallback, -// EDNS0 buffer size negotiation, and root server discovery. The package is -// intentionally narrow: it sends iterative (non-recursive) queries and returns -// the raw responses for the traversal engine to interpret. +// It wraps the github.com/miekg/dns library behind a single query path +// (Client) that sends non-recursive (RD=0) queries with retrying, TCP +// fallback on truncation, EDNS0 negotiation and a per-run packet cache, plus +// root server discovery. Production uses the real wire exchange; tests inject +// a mock ExchangeFunc into the exact same path. package dns import ( "context" + "errors" "fmt" "net" + "sync" "time" "github.com/miekg/dns" ) +// QueryConfig controls transport parameters for the single query path. type QueryConfig struct { - UDPSize int - Timeout time.Duration - Retries int - UseTCP bool + // UDPSize is the EDNS0 UDP payload size. An OPT record is only attached + // when UDPSize > 512, mirroring dnstraverse's caching_resolver.rb. + UDPSize int + // Timeout is the per-attempt packet timeout (dnsruby packet_timeout, + // dnstraverse default 2s). + Timeout time.Duration + // Retries is the total number of send attempts, matching dnsruby + // retry_times: Resolver#generate_timeouts schedules retry_times + // transmissions in total — the first immediately and retry k at + // retry_delay*2^k seconds after the first. Values below 1 are clamped to + // 1 so exactly one query is still sent (dnsruby with retry_times 0 would + // send nothing and hang; this also fixes the old + // "failed after 0 retries: %!w( )" error). + Retries int + // RetryDelay is dnsruby's retry_delay (dnstraverse default 2s). + RetryDelay time.Duration + // UseTCP forces every query over TCP (--always-tcp). + UseTCP bool + // AllowTCP enables the UDP→TCP retry when a response is truncated. AllowTCP bool } func DefaultQueryConfig() *QueryConfig { return &QueryConfig{ - UDPSize: DefaultEDNS0UDPSize(), - Timeout: 5 * time.Second, - Retries: 3, - UseTCP: false, - AllowTCP: true, + UDPSize: DefaultEDNS0UDPSize(), + Timeout: 2 * time.Second, + Retries: 2, + RetryDelay: 2 * time.Second, + UseTCP: false, + AllowTCP: true, } } +// withDefaults returns a copy of cfg with zero values replaced by defaults. +func (cfg *QueryConfig) withDefaults() *QueryConfig { + if cfg == nil { + return DefaultQueryConfig() + } + out := *cfg + if out.UDPSize <= 0 { + out.UDPSize = DefaultEDNS0UDPSize() + } + if out.Timeout <= 0 { + out.Timeout = 2 * time.Second + } + if out.Retries < 1 { + out.Retries = 1 + } + if out.RetryDelay <= 0 { + out.RetryDelay = 2 * time.Second + } + return &out +} + +// ExchangeFunc performs one wire exchange. server is either a bare host/IP +// (port 53 implied) or an explicit host:port. type ExchangeFunc func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) -func Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { - if cfg == nil { - cfg = DefaultQueryConfig() - } - - if cfg.UDPSize <= 0 { - cfg.UDPSize = DefaultEDNS0UDPSize() - } - if cfg.Timeout <= 0 { - cfg.Timeout = 5 * time.Second - } - - return QueryWithExchange(ctx, server, name, qtype, cfg, realExchange) -} - func realExchange(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - addr := net.JoinHostPort(server, "53") - - var c *dns.Client - if useTCP { - c = &dns.Client{ - Net: "tcp", - ReadTimeout: 5 * time.Second, - WriteTimeout: 5 * time.Second, - } - } else { - c = &dns.Client{ - Net: "udp", - ReadTimeout: 5 * time.Second, - WriteTimeout: 5 * time.Second, - } + addr := server + if _, _, err := net.SplitHostPort(server); err != nil { + addr = net.JoinHostPort(server, "53") } + proto := "udp" + if useTCP { + proto = "tcp" + } + c := &dns.Client{ + Net: proto, + ReadTimeout: 2 * time.Second, + WriteTimeout: 2 * time.Second, + } if deadline, ok := ctx.Deadline(); ok { c.ReadTimeout = time.Until(deadline) c.WriteTimeout = time.Until(deadline) @@ -75,150 +100,222 @@ func realExchange(ctx context.Context, server string, msg *dns.Msg, useTCP bool) r, _, err := c.ExchangeContext(ctx, msg, addr) if err != nil { - return nil, fmt.Errorf("dns exchange (%s) with %s: %w", c.Net, addr, err) + return nil, fmt.Errorf("dns exchange (%s) with %s: %w", proto, addr, err) } - return r, nil } -func QueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig, exchangeFn ExchangeFunc) (*dns.Msg, error) { - if cfg == nil { - cfg = DefaultQueryConfig() +// Client is the single query path used identically by production and tests. +// Every query is non-recursive (RD=0) and deduplicated by a per-run packet +// cache keyed (server IP, qname, qclass, qtype, udpsize), mirroring +// dnstraverse's caching_resolver.rb: repeat askers replay the cached answer +// (or cached failure) without touching the wire. +type Client struct { + cfg *QueryConfig + exchange ExchangeFunc + + mu sync.Mutex + cache map[packetKey]*packetEntry + requests int + cacheHits int +} + +type packetKey struct { + server string + qname string + qclass uint16 + qtype uint16 + udpsize int +} + +type packetEntry struct { + once sync.Once + msg *dns.Msg + err error +} + +// NewClient creates a Client. A nil exchange means the real wire exchange; +// tests pass a mock so no packets leave the process. +func NewClient(cfg *QueryConfig, exchange ExchangeFunc) *Client { + if exchange == nil { + exchange = realExchange } - if cfg.UDPSize <= 0 { - cfg.UDPSize = DefaultEDNS0UDPSize() + return &Client{ + cfg: cfg.withDefaults(), + exchange: exchange, + cache: make(map[packetKey]*packetEntry), + } +} + +// Requests reports how many queries were asked of the client (cache hits included). +func (c *Client) Requests() int { + c.mu.Lock() + defer c.mu.Unlock() + return c.requests +} + +// CacheHits reports how many queries were served from the packet cache. +func (c *Client) CacheHits() int { + c.mu.Lock() + defer c.mu.Unlock() + return c.cacheHits +} + +// Query sends a non-recursive query for name/qtype (class IN) to server and +// returns the response plus any warnings gathered along the way (EDNS0 +// fallback, recursion offered, truncation). A non-nil error corresponds to +// dnstraverse's "exception" status (network failure after all retries). +func (c *Client) Query(ctx context.Context, server net.IP, name string, qtype uint16) (*dns.Msg, []string, error) { + msg, err := c.cachedExchange(ctx, server, name, qtype, c.cfg.UDPSize) + if err != nil { + return nil, nil, err } - msg := buildQuery(name, qtype, cfg.UDPSize) - serverStr := server.String() + var warnings []string + + // EDNS0 fallback (decoded_query.rb makequery_message): FORMERR/NOTIMP/ + // SERVFAIL with udpsize > 512 may mean the server chokes on OPT; retry + // once at 512 and keep the retry only if it clears the error. + if c.cfg.UDPSize > MinEDNS0UDPSize() && ednsFailure(msg.Rcode) { + retryMsg, retryErr := c.cachedExchange(ctx, server, name, qtype, MinEDNS0UDPSize()) + if retryErr == nil && !ednsFailure(retryMsg.Rcode) { + warnings = append(warnings, fmt.Sprintf("%s doesn't seem to support EDNS0", server)) + msg = retryMsg + } + } + + // msg_comment with want_recursion=false (message_utility.rb). + if msg.RecursionAvailable { + warnings = append(warnings, fmt.Sprintf("%s allows recursion", server)) + } + if msg.Truncated { + warnings = append(warnings, fmt.Sprintf("%s sent truncated packet", server)) + } + + return msg, warnings, nil +} + +func ednsFailure(rcode int) bool { + return rcode == dns.RcodeFormatError || + rcode == dns.RcodeNotImplemented || + rcode == dns.RcodeServerFailure +} + +// cachedExchange sends at most one wire query per packet cache key; the +// outcome (response or error) is cached and replayed for repeat askers. +func (c *Client) cachedExchange(ctx context.Context, server net.IP, name string, qtype uint16, udpsize int) (*dns.Msg, error) { + key := packetKey{ + server: server.String(), + qname: dns.CanonicalName(name), + qclass: dns.ClassINET, + qtype: qtype, + udpsize: udpsize, + } + + c.mu.Lock() + c.requests++ + entry, ok := c.cache[key] + if ok { + c.cacheHits++ + } else { + entry = &packetEntry{} + c.cache[key] = entry + } + c.mu.Unlock() + + entry.once.Do(func() { + entry.msg, entry.err = exchangeWithRetry(ctx, c.exchange, key.server, buildQuery(name, qtype, udpsize), c.cfg) + }) + + return copyMsg(entry.msg), entry.err +} + +// exchangeWithRetry implements dnsruby's retry schedule (resolver.rb +// generate_timeouts): cfg.Retries is the TOTAL number of transmissions — the +// first goes immediately and retry k is sent retry_delay*2^k seconds after +// the first (so gaps of 2d, 2d, 4d, 8d, ...). Each attempt gets its own +// cfg.Timeout (dnsruby packet_timeout). A truncated UDP reply is retried over +// TCP within the same attempt when cfg.AllowTCP. +func exchangeWithRetry(ctx context.Context, exchange ExchangeFunc, server string, msg *dns.Msg, cfg *QueryConfig) (*dns.Msg, error) { + attempts := cfg.Retries + if attempts < 1 { + attempts = 1 + } var lastErr error - - for attempt := 0; attempt < cfg.Retries; attempt++ { + for attempt := 0; attempt < attempts; attempt++ { if attempt > 0 { select { case <-ctx.Done(): return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err()) - case <-time.After(backoffDelay(attempt)): + case <-time.After(retryGap(cfg.RetryDelay, attempt)): } } - if cfg.UseTCP { - resp, err := exchangeFn(ctx, serverStr, msg, true) - if err != nil { - lastErr = err - continue - } - return resp, nil - } - - resp, err := exchangeFn(ctx, serverStr, msg, false) + resp, err := exchangeOnce(ctx, exchange, server, msg, cfg) if err != nil { lastErr = err continue } - - if resp.Truncated && cfg.AllowTCP { - resp, err = exchangeFn(ctx, serverStr, msg, true) - if err != nil { - lastErr = err - continue - } - return resp, nil - } - return resp, nil } - return nil, fmt.Errorf("query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr) + q := msg.Question[0] + return nil, fmt.Errorf("query %s %s to %s failed after %d attempts: %w", + q.Name, QNameType(q.Qtype), server, attempts, lastErr) } -func IterativeQuery(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { - if cfg == nil { - cfg = DefaultQueryConfig() +// retryGap returns the wait before retry number `retry` (1-based). dnsruby +// sends retry k at absolute time retry_delay*2^k, so the gap is 2d before the +// first retry and d*2^(k-1) for each retry after that. +func retryGap(d time.Duration, retry int) time.Duration { + if retry <= 1 { + return 2 * d } - if cfg.UDPSize <= 0 { - cfg.UDPSize = DefaultEDNS0UDPSize() - } - return IterativeQueryWithExchange(ctx, server, name, qtype, cfg, realExchange) + return d << uint(retry-1) } -func IterativeQueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig, exchangeFn ExchangeFunc) (*dns.Msg, error) { - if cfg == nil { - cfg = DefaultQueryConfig() +func exchangeOnce(ctx context.Context, exchange ExchangeFunc, server string, msg *dns.Msg, cfg *QueryConfig) (*dns.Msg, error) { + actx, cancel := context.WithTimeout(ctx, cfg.Timeout) + defer cancel() + + resp, err := exchange(actx, server, msg, cfg.UseTCP) + if err != nil { + return nil, err } - if cfg.UDPSize <= 0 { - cfg.UDPSize = DefaultEDNS0UDPSize() + if resp == nil { + return nil, errors.New("nil response") } - msg := buildQuery(name, qtype, cfg.UDPSize) - msg.RecursionDesired = false - serverStr := server.String() - - var lastErr error - - for attempt := 0; attempt < cfg.Retries; attempt++ { - if attempt > 0 { - select { - case <-ctx.Done(): - return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err()) - case <-time.After(backoffDelay(attempt)): - } - } - - if cfg.UseTCP { - resp, err := exchangeFn(ctx, serverStr, msg, true) - if err != nil { - lastErr = err - continue - } - return resp, nil - } - - resp, err := exchangeFn(ctx, serverStr, msg, false) + if !cfg.UseTCP && resp.Truncated && cfg.AllowTCP { + resp, err = exchange(actx, server, msg, true) if err != nil { - lastErr = err - continue + return nil, err } - if resp == nil { - lastErr = fmt.Errorf("nil response") - continue + return nil, errors.New("nil response") } - - if resp.Truncated && cfg.AllowTCP { - resp, err = exchangeFn(ctx, serverStr, msg, true) - if err != nil { - lastErr = err - continue - } - return resp, nil - } - - return resp, nil } - return nil, fmt.Errorf("iterative query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr) + return resp, nil } -// backoffDelay computes the wait duration before the given retry attempt (1-indexed). -// Delays: attempt=1 → 100ms, attempt=2 → 200ms, attempt=3 → 400ms, capped at 2s. -func backoffDelay(attempt int) time.Duration { - if attempt <= 0 { - return 0 - } - delay := time.Duration(uint(1)< maxDelay { - return maxDelay - } - return delay -} - -func buildQuery(name string, qtype uint16, udpSize int) *dns.Msg { +// buildQuery constructs a non-recursive (RD=0) class IN query. The EDNS0 OPT +// record is attached only when udpsize > 512 (caching_resolver.rb adds OPT +// under the same condition), with the DO bit off. +func buildQuery(name string, qtype uint16, udpsize int) *dns.Msg { m := new(dns.Msg) m.SetQuestion(dns.Fqdn(name), qtype) - m.RecursionDesired = true - m.SetEdns0(uint16(udpSize), false) + m.RecursionDesired = false + if udpsize > MinEDNS0UDPSize() { + m.SetEdns0(uint16(udpsize), false) + } return m } + +func copyMsg(m *dns.Msg) *dns.Msg { + if m == nil { + return nil + } + return m.Copy() +} diff --git a/internal/dns/query_test.go b/internal/dns/query_test.go index 3ec1c8b..f3eff36 100644 --- a/internal/dns/query_test.go +++ b/internal/dns/query_test.go @@ -4,480 +4,493 @@ import ( "context" "errors" "net" + "strings" "sync" "testing" + "time" "github.com/miekg/dns" ) +// testClient returns a Client with fast retries wired to fn. +func testClient(cfg *QueryConfig, fn ExchangeFunc) *Client { + if cfg == nil { + cfg = DefaultQueryConfig() + } + cfg.Timeout = time.Second + cfg.RetryDelay = time.Millisecond + return NewClient(cfg, fn) +} + +func answerMsg(name string, ip string) *dns.Msg { + m := new(dns.Msg) + m.SetReply(new(dns.Msg)) + m.Answer = append(m.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip), + }) + return m +} + func TestDefaultQueryConfig(t *testing.T) { cfg := DefaultQueryConfig() - if cfg == nil { - t.Fatal("DefaultQueryConfig returned nil") - } if cfg.UDPSize != 2048 { t.Errorf("UDPSize = %d, want 2048", cfg.UDPSize) } - if cfg.Retries != 3 { - t.Errorf("Retries = %d, want 3", cfg.Retries) + if cfg.Retries != 2 { + t.Errorf("Retries = %d, want 2", cfg.Retries) + } + if cfg.Timeout != 2*time.Second { + t.Errorf("Timeout = %v, want 2s", cfg.Timeout) + } + if cfg.RetryDelay != 2*time.Second { + t.Errorf("RetryDelay = %v, want 2s", cfg.RetryDelay) } if cfg.UseTCP { t.Error("UseTCP should be false by default") } + if !cfg.AllowTCP { + t.Error("AllowTCP should be true by default") + } } -func TestBuildQuery(t *testing.T) { - msg := buildQuery("example.com.", TypeA, 2048) +func TestBuildQueryRDZeroWithEDNS(t *testing.T) { + msg := buildQuery("example.com", TypeA, 2048) if len(msg.Question) != 1 { t.Fatalf("expected 1 question, got %d", len(msg.Question)) } q := msg.Question[0] if q.Name != "example.com." { - t.Errorf("question name = %q, want %q", q.Name, "example.com.") + t.Errorf("Fqdn not applied: got %q", q.Name) } if q.Qtype != TypeA { t.Errorf("question type = %d, want %d", q.Qtype, TypeA) } - if !msg.RecursionDesired { - t.Error("RecursionDesired should be true") + if msg.RecursionDesired { + t.Error("RecursionDesired must be false on the query path") } - if opt := msg.IsEdns0(); opt == nil { - t.Error("expected EDNS0 OPT record") - } else if opt.UDPSize() != 2048 { + opt := msg.IsEdns0() + if opt == nil { + t.Fatal("expected EDNS0 OPT record for udpsize > 512") + } + if opt.UDPSize() != 2048 { t.Errorf("EDNS0 UDPSize = %d, want 2048", opt.UDPSize()) } } -func TestBuildQueryFqdn(t *testing.T) { - msg := buildQuery("example.com", TypeA, 4096) - q := msg.Question[0] - if q.Name != "example.com." { - t.Errorf("Fqdn not applied: got %q, want %q", q.Name, "example.com.") +func TestBuildQueryNoEDNSAt512(t *testing.T) { + msg := buildQuery("example.com.", TypeA, 512) + if msg.IsEdns0() != nil { + t.Error("no OPT record should be attached when udpsize <= 512") } } -func TestQueryWithExchangeSuccess(t *testing.T) { - expectedResp := new(dns.Msg) - expectedResp.SetReply(new(dns.Msg)) - expectedResp.Answer = append(expectedResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), +func TestClientQuerySuccess(t *testing.T) { + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + if msg.RecursionDesired { + t.Error("query must be sent with RD=0") + } + return answerMsg("example.com.", "93.184.216.34"), nil }) - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return expectedResp.Copy(), nil - } - - cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 1, - UseTCP: false, - } - - server := net.ParseIP("8.8.8.8") - resp, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) + resp, warnings, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(resp.Answer) != 1 { t.Fatalf("expected 1 answer, got %d", len(resp.Answer)) } + if len(warnings) != 0 { + t.Errorf("unexpected warnings: %v", warnings) + } } -func TestQueryTCPFallbackOnTruncation(t *testing.T) { - truncatedResp := new(dns.Msg) - truncatedResp.Truncated = true - truncatedResp.SetReply(new(dns.Msg)) - - fullResp := new(dns.Msg) - fullResp.SetReply(new(dns.Msg)) - fullResp.Answer = append(fullResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), +func TestClientPacketCache(t *testing.T) { + var calls int + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + return answerMsg("example.com.", "1.2.3.4"), nil }) - var mu sync.Mutex - calls := []bool{} - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - mu.Lock() - defer mu.Unlock() - calls = append(calls, useTCP) - if !useTCP { - return truncatedResp.Copy(), nil - } - return fullResp.Copy(), nil - } - - cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 1, - UseTCP: false, - AllowTCP: true, - } - + ctx := context.Background() server := net.ParseIP("8.8.8.8") - resp, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) + for i := 0; i < 3; i++ { + if _, _, err := c.Query(ctx, server, "example.com", TypeA); err != nil { + t.Fatalf("query %d: %v", i, err) + } + } + if calls != 1 { + t.Errorf("expected 1 wire query for repeat askers, got %d", calls) + } + if c.Requests() != 3 { + t.Errorf("Requests = %d, want 3", c.Requests()) + } + if c.CacheHits() != 2 { + t.Errorf("CacheHits = %d, want 2", c.CacheHits()) + } + // Case differences must hit the same cache entry. + if _, _, err := c.Query(ctx, server, "EXAMPLE.COM.", TypeA); err != nil { + t.Fatal(err) + } + if calls != 1 { + t.Errorf("case-insensitive lookup should hit cache, got %d wire calls", calls) + } +} + +func TestClientPacketCacheDistinctKeys(t *testing.T) { + var calls int + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + return answerMsg("example.com.", "1.2.3.4"), nil + }) + + ctx := context.Background() + _, _, _ = c.Query(ctx, net.ParseIP("8.8.8.8"), "example.com", TypeA) + _, _, _ = c.Query(ctx, net.ParseIP("9.9.9.9"), "example.com", TypeA) + _, _, _ = c.Query(ctx, net.ParseIP("8.8.8.8"), "example.com", TypeAAAA) + _, _, _ = c.Query(ctx, net.ParseIP("8.8.8.8"), "other.com", TypeA) + if calls != 4 { + t.Errorf("expected 4 wire queries for 4 distinct keys, got %d", calls) + } +} + +func TestClientPacketCacheReplaysErrors(t *testing.T) { + var calls int + c := testClient(&QueryConfig{Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + return nil, errors.New("connection refused") + }) + + ctx := context.Background() + server := net.ParseIP("8.8.8.8") + _, _, err1 := c.Query(ctx, server, "example.com", TypeA) + _, _, err2 := c.Query(ctx, server, "example.com", TypeA) + if err1 == nil || err2 == nil { + t.Fatal("expected errors") + } + if calls != 1 { + t.Errorf("failures must be cached too: got %d wire calls", calls) + } +} + +func TestClientEDNSFallback(t *testing.T) { + var sizes []int + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + size := 512 + if opt := msg.IsEdns0(); opt != nil { + size = int(opt.UDPSize()) + } + sizes = append(sizes, size) + m := new(dns.Msg) + m.SetReply(msg) + if size > 512 { + m.Rcode = dns.RcodeFormatError + return m, nil + } + m.Answer = append(m.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("1.2.3.4"), + }) + return m, nil + }) + + resp, warnings, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error: %v", err) } - if len(calls) != 2 { - t.Fatalf("expected 2 exchange calls (UDP then TCP), got %d", len(calls)) + if len(sizes) != 2 || sizes[0] != 2048 || sizes[1] != 512 { + t.Fatalf("expected 2048 then 512 queries, got %v", sizes) } - if calls[0] != false { - t.Error("first call should be UDP") + if resp.Rcode != dns.RcodeSuccess || len(resp.Answer) != 1 { + t.Error("expected the 512-byte retry response to be returned") } - if calls[1] != true { - t.Error("second call should be TCP") + want := "8.8.8.8 doesn't seem to support EDNS0" + if len(warnings) != 1 || warnings[0] != want { + t.Errorf("warnings = %v, want [%q]", warnings, want) + } +} + +func TestClientEDNSFallbackKeepsOriginalWhenRetryFails(t *testing.T) { + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + m := new(dns.Msg) + m.SetReply(msg) + m.Rcode = dns.RcodeServerFailure + return m, nil + }) + + resp, warnings, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if resp.Rcode != dns.RcodeServerFailure { + t.Errorf("expected original SERVFAIL response, got rcode %d", resp.Rcode) + } + if len(warnings) != 0 { + t.Errorf("no EDNS0 warning expected when the retry also fails: %v", warnings) + } +} + +func TestClientNoEDNSFallbackAt512(t *testing.T) { + var calls int + c := testClient(&QueryConfig{UDPSize: 512, Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + m := new(dns.Msg) + m.SetReply(msg) + m.Rcode = dns.RcodeFormatError + return m, nil + }) + + resp, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if calls != 1 { + t.Errorf("expected no fallback query at udpsize 512, got %d calls", calls) + } + if resp.Rcode != dns.RcodeFormatError { + t.Errorf("expected FORMERR passthrough, got %d", resp.Rcode) + } +} + +func TestClientRecursionAvailableWarning(t *testing.T) { + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + m := answerMsg("example.com.", "1.2.3.4") + m.RecursionAvailable = true + return m, nil + }) + + _, warnings, err := c.Query(context.Background(), net.ParseIP("192.0.2.1"), "example.com", TypeA) + if err != nil { + t.Fatal(err) + } + want := "192.0.2.1 allows recursion" + if len(warnings) != 1 || warnings[0] != want { + t.Errorf("warnings = %v, want [%q]", warnings, want) + } +} + +func TestClientTruncationWarningWhenTCPDisallowed(t *testing.T) { + c := testClient(&QueryConfig{AllowTCP: false, Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + if useTCP { + t.Error("TCP must not be used when AllowTCP is false") + } + m := answerMsg("example.com.", "1.2.3.4") + m.Truncated = true + return m, nil + }) + + resp, warnings, err := c.Query(context.Background(), net.ParseIP("192.0.2.1"), "example.com", TypeA) + if err != nil { + t.Fatal(err) + } + if !resp.Truncated { + t.Error("expected truncated response to be returned as-is") + } + want := "192.0.2.1 sent truncated packet" + if len(warnings) != 1 || warnings[0] != want { + t.Errorf("warnings = %v, want [%q]", warnings, want) + } +} + +func TestClientTCPFallbackOnTruncation(t *testing.T) { + var mu sync.Mutex + var calls []bool + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + mu.Lock() + calls = append(calls, useTCP) + mu.Unlock() + if !useTCP { + m := new(dns.Msg) + m.SetReply(msg) + m.Truncated = true + return m, nil + } + return answerMsg("example.com.", "93.184.216.34"), nil + }) + + resp, warnings, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(calls) != 2 || calls[0] != false || calls[1] != true { + t.Fatalf("expected UDP then TCP, got %v", calls) } if len(resp.Answer) != 1 { - t.Fatalf("expected 1 answer from TCP fallback, got %d", len(resp.Answer)) + t.Fatal("expected answer from TCP fallback") + } + if len(warnings) != 0 { + t.Errorf("unexpected warnings after successful TCP fallback: %v", warnings) } } -func TestQueryAlwaysTCP(t *testing.T) { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - +func TestClientAlwaysTCP(t *testing.T) { var mu sync.Mutex - calls := []bool{} - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + var calls []bool + c := testClient(&QueryConfig{UseTCP: true, Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { mu.Lock() - defer mu.Unlock() calls = append(calls, useTCP) - return resp.Copy(), nil - } + mu.Unlock() + return answerMsg("example.com.", "1.2.3.4"), nil + }) - cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 1, - UseTCP: true, - } - - server := net.ParseIP("8.8.8.8") - _, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error: %v", err) } - if len(calls) != 1 { - t.Fatalf("expected 1 exchange call, got %d", len(calls)) - } - if !calls[0] { - t.Error("expected TCP call when UseTCP is true") + if len(calls) != 1 || !calls[0] { + t.Errorf("expected a single TCP call, got %v", calls) } } -func TestQueryRetriesOnFailure(t *testing.T) { - var mu sync.Mutex - callCount := 0 - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - mu.Lock() - callCount++ - mu.Unlock() +func TestClientRetriesAreTotalAttempts(t *testing.T) { + var calls int + c := testClient(&QueryConfig{Retries: 3}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ return nil, errors.New("connection refused") - } + }) - cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 3, - UseTCP: false, - } - - server := net.ParseIP("8.8.8.8") - _, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { t.Fatal("expected error after retries exhausted") } - if callCount != 3 { - t.Errorf("expected 3 calls (retries exhausted), got %d", callCount) + // dnsruby retry_times counts total transmissions, not extra retries. + if calls != 3 { + t.Errorf("expected 3 attempts, got %d", calls) + } + if !strings.Contains(err.Error(), "after 3 attempts") { + t.Errorf("error should mention attempt count: %v", err) } } -func TestQueryContextCancellation(t *testing.T) { +func TestClientZeroRetriesStillQueriesOnce(t *testing.T) { + var calls int + c := testClient(&QueryConfig{Retries: 0}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + return nil, errors.New("boom") + }) + + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) + if err == nil { + t.Fatal("expected error") + } + if calls != 1 { + t.Errorf("Retries=0 must be clamped to one attempt, got %d", calls) + } + // Regression: the old code produced "failed after 0 retries: %!w( )". + if strings.Contains(err.Error(), "%!w") || strings.Contains(err.Error(), " ") { + t.Errorf("malformed error message: %v", err) + } +} + +func TestClientRetryThenSuccess(t *testing.T) { + var calls int + c := testClient(&QueryConfig{Retries: 3}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + if calls < 2 { + return nil, errors.New("transient error") + } + return answerMsg("example.com.", "1.2.3.4"), nil + }) + + resp, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(resp.Answer) == 0 { + t.Fatal("expected answer after retry") + } + if calls != 2 { + t.Errorf("expected 2 attempts, got %d", calls) + } +} + +func TestClientNilResponseIsError(t *testing.T) { + c := testClient(&QueryConfig{Retries: 2}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return nil, nil + }) + + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) + if err == nil { + t.Fatal("expected error for nil responses") + } +} + +func TestClientContextCancellation(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) cancel() - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, ctx.Err() - } + c := testClient(&QueryConfig{Retries: 5}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return nil, errors.New("error") + }) - cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 1, - UseTCP: false, - } - - server := net.ParseIP("8.8.8.8") - _, err := QueryWithExchange(ctx, server, "example.com", TypeA, cfg, exchangeFn) + _, _, err := c.Query(ctx, net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { t.Fatal("expected error on cancelled context") } } -func TestQueryNilConfigUsesDefaults(t *testing.T) { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return resp.Copy(), nil - } - - server := net.ParseIP("8.8.8.8") - _, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, nil, exchangeFn) +func TestClientNilConfigUsesDefaults(t *testing.T) { + c := NewClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return answerMsg("example.com.", "1.2.3.4"), nil + }) + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error with nil config: %v", err) } } -func TestQueryZeroValuesUseDefaults(t *testing.T) { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return resp.Copy(), nil - } - - cfg := &QueryConfig{ - UDPSize: 0, - Timeout: 0, - Retries: 1, - UseTCP: false, - } - - server := net.ParseIP("8.8.8.8") - _, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } -} - -func TestQueryTCPFallbackFailsThenRetries(t *testing.T) { - truncatedResp := new(dns.Msg) - truncatedResp.Truncated = true - truncatedResp.SetReply(new(dns.Msg)) - - var mu sync.Mutex - callCount := 0 - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - mu.Lock() - callCount++ - mu.Unlock() +func TestClientTCPFallbackFailureRetries(t *testing.T) { + var calls int + c := testClient(&QueryConfig{Retries: 2, AllowTCP: true}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ if !useTCP { - return truncatedResp.Copy(), nil + m := new(dns.Msg) + m.SetReply(msg) + m.Truncated = true + return m, nil } return nil, errors.New("tcp failed") - } + }) - cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 2, - UseTCP: false, - AllowTCP: true, - } - - server := net.ParseIP("8.8.8.8") - _, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { t.Fatal("expected error when TCP fallback always fails") } - if callCount != 4 { - t.Errorf("expected 4 calls (2 retries x UDP+TCP), got %d", callCount) + if calls != 4 { + t.Errorf("expected 4 exchange calls (2 attempts x UDP+TCP), got %d", calls) } } -func TestQueryNoTCPFallbackWhenDisabled(t *testing.T) { - truncatedResp := new(dns.Msg) - truncatedResp.Truncated = true - truncatedResp.SetReply(new(dns.Msg)) - - var mu sync.Mutex - calls := []bool{} - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - mu.Lock() - defer mu.Unlock() - calls = append(calls, useTCP) - return truncatedResp.Copy(), nil +func TestRetryGap(t *testing.T) { + d := 2 * time.Second + tests := []struct { + retry int + want time.Duration + }{ + {1, 4 * time.Second}, // dnsruby sends retry 1 at absolute 2d + {2, 4 * time.Second}, // retry 2 at 4d → gap 2d + {3, 8 * time.Second}, + {4, 16 * time.Second}, } - - cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 1, - UseTCP: false, - AllowTCP: false, // TCP fallback must be suppressed - } - - server := net.ParseIP("8.8.8.8") - resp, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - // Only one UDP call; no TCP fallback. - if len(calls) != 1 { - t.Fatalf("expected 1 exchange call (no TCP fallback), got %d", len(calls)) - } - if calls[0] != false { - t.Error("expected UDP-only call") - } - if !resp.Truncated { - t.Error("expected truncated response to be returned as-is") + for _, tt := range tests { + got := retryGap(d, tt.retry) + if got != tt.want { + t.Errorf("retryGap(%v, %d) = %v, want %v", d, tt.retry, got, tt.want) + } } } -func TestIterativeQueryWithExchangeSuccess(t *testing.T) { -answerResp := new(dns.Msg) -answerResp.SetReply(new(dns.Msg)) -answerResp.Answer = append(answerResp.Answer, &dns.A{ -Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, -A: net.ParseIP("1.2.3.4"), -}) - -exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { -if msg.RecursionDesired { -t.Error("IterativeQuery should send RD=false") -} -return answerResp.Copy(), nil -} - -server := net.ParseIP("198.41.0.4") -resp, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, nil, exchangeFn) -if err != nil { -t.Fatalf("unexpected error: %v", err) -} -if len(resp.Answer) == 0 { -t.Fatal("expected answer records") -} -} - -func TestIterativeQueryWithExchangeRetry(t *testing.T) { -callCount := 0 -answerResp := new(dns.Msg) -answerResp.SetReply(new(dns.Msg)) -answerResp.Answer = append(answerResp.Answer, &dns.A{ -Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, -A: net.ParseIP("1.2.3.4"), -}) - -exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { -callCount++ -if callCount < 2 { -return nil, errors.New("transient error") -} -return answerResp.Copy(), nil -} - -cfg := &QueryConfig{UDPSize: 2048, Retries: 3, AllowTCP: true} -server := net.ParseIP("198.41.0.4") -resp, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) -if err != nil { -t.Fatalf("unexpected error: %v", err) -} -if resp == nil { -t.Fatal("expected non-nil response after retry") -} -} - -func TestIterativeQueryWithExchangeAllFail(t *testing.T) { -exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { -return nil, errors.New("server unreachable") -} - -cfg := &QueryConfig{UDPSize: 2048, Retries: 2, AllowTCP: false} -server := net.ParseIP("198.41.0.4") -_, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) -if err == nil { -t.Fatal("expected error when all attempts fail") -} -} - -func TestIterativeQueryWithExchangeTCPFallback(t *testing.T) { -truncatedResp := new(dns.Msg) -truncatedResp.SetReply(new(dns.Msg)) -truncatedResp.Truncated = true - -fullResp := new(dns.Msg) -fullResp.SetReply(new(dns.Msg)) -fullResp.Answer = append(fullResp.Answer, &dns.A{ -Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, -A: net.ParseIP("1.2.3.4"), -}) - -exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { -if !useTCP { -return truncatedResp.Copy(), nil -} -return fullResp.Copy(), nil -} - -cfg := &QueryConfig{UDPSize: 2048, Retries: 1, AllowTCP: true} -server := net.ParseIP("198.41.0.4") -resp, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) -if err != nil { -t.Fatalf("unexpected error: %v", err) -} -if len(resp.Answer) == 0 { -t.Fatal("expected answer after TCP fallback") -} -} - -func TestIterativeQueryWithExchangeContextCancelled(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - // Cancel the context immediately so the retry loop aborts during backoff - cancel() - - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, errors.New("error") - } - - cfg := &QueryConfig{UDPSize: 2048, Retries: 5, AllowTCP: false} - server := net.ParseIP("198.41.0.4") - _, err := IterativeQueryWithExchange(ctx, server, "example.com", dns.TypeA, cfg, exchangeFn) +func TestQueryErrorMessageMentionsQuery(t *testing.T) { + c := testClient(&QueryConfig{Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return nil, errors.New("unreachable") + }) + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { - t.Fatal("expected error when context cancelled") + t.Fatal("expected error") + } + for _, part := range []string{"example.com.", "A", "8.8.8.8", "unreachable"} { + if !strings.Contains(err.Error(), part) { + t.Errorf("error %q should contain %q", err, part) + } } } - -func TestIterativeQueryWithExchangeNilResponse(t *testing.T) { -callCount := 0 -exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { -callCount++ -return nil, nil // nil response, no error -} - -cfg := &QueryConfig{UDPSize: 2048, Retries: 2, AllowTCP: false} -server := net.ParseIP("198.41.0.4") -_, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) -if err == nil { -t.Fatal("expected error for nil responses") -} -} - -func TestIterativeQueryWithExchangeUseTCP(t *testing.T) { -var wasTCP bool -answerResp := new(dns.Msg) -answerResp.SetReply(new(dns.Msg)) -answerResp.Answer = append(answerResp.Answer, &dns.A{ -Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, -A: net.ParseIP("1.2.3.4"), -}) - -exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { -wasTCP = useTCP -return answerResp.Copy(), nil -} - -cfg := &QueryConfig{UDPSize: 2048, Retries: 1, UseTCP: true} -server := net.ParseIP("198.41.0.4") -_, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) -if err != nil { -t.Fatalf("unexpected error: %v", err) -} -if !wasTCP { -t.Error("expected TCP exchange when UseTCP=true") -} -} diff --git a/internal/dns/real_exchange_test.go b/internal/dns/real_exchange_test.go index ec43bf4..2a772e1 100644 --- a/internal/dns/real_exchange_test.go +++ b/internal/dns/real_exchange_test.go @@ -2,7 +2,6 @@ package dns import ( "context" - "fmt" "net" "testing" "time" @@ -10,263 +9,32 @@ import ( "github.com/miekg/dns" ) -// startTestDNSServer starts a local DNS server on a random port and returns the address and a stop function. -func startTestDNSServer(t *testing.T, handler dns.HandlerFunc) (string, func()) { +// startTestDNSServer starts a loopback DNS server on a random port and +// returns its address. No test in this file touches the real network. +func startTestDNSServer(t *testing.T, network string, handler dns.HandlerFunc) string { t.Helper() - pc, err := net.ListenPacket("udp", "127.0.0.1:0") - if err != nil { - t.Skipf("cannot start test DNS server: %v", err) - } - addr := pc.LocalAddr().String() - mux := dns.NewServeMux() mux.HandleFunc(".", handler) - srv := &dns.Server{ - PacketConn: pc, - Net: "udp", - Handler: mux, - } + srv := &dns.Server{Net: network, Handler: mux} + var addr string - started := make(chan struct{}) - srv.NotifyStartedFunc = func() { close(started) } - - go func() { - _ = srv.ActivateAndServe() - }() - - select { - case <-started: - case <-time.After(2 * time.Second): - t.Skip("test DNS server did not start in time") - } - - return addr, func() { _ = srv.Shutdown() } -} - -func TestQueryUsesRealExchange(t *testing.T) { - addr, stop := startTestDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) { - resp := new(dns.Msg) - resp.SetReply(r) - resp.Answer = append(resp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - _ = w.WriteMsg(resp) - }) - defer stop() - - host, portStr, err := net.SplitHostPort(addr) - if err != nil { - t.Fatalf("parse addr: %v", err) - } - var port int - fmt.Sscanf(portStr, "%d", &port) - - // Patch the realExchange to use the test server by using QueryWithExchange with a custom exchangeFn. - // Since we can't inject into Query directly, use realExchangeWithPort for test. - serverIP := net.ParseIP(host) - cfg := DefaultQueryConfig() - cfg.Retries = 1 - - // Test QueryWithExchange (already covered), but now test Query+realExchange flow via - // a patched exchange that routes to our test server port. - patchedExchange := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - c := &dns.Client{Net: "udp", ReadTimeout: 3 * time.Second, WriteTimeout: 3 * time.Second} - r, _, err := c.ExchangeContext(ctx, msg, fmt.Sprintf("%s:%d", host, port)) - return r, err - } - - resp, err := QueryWithExchange(context.Background(), serverIP, "example.com", TypeA, cfg, patchedExchange) - if err != nil { - t.Fatalf("QueryWithExchange: %v", err) - } - if len(resp.Answer) == 0 { - t.Fatal("expected at least 1 answer") - } -} - -func TestRealExchangeViaDirect(t *testing.T) { - // Test realExchange directly via the exported Query function - // by using a server that will respond or fail quickly. - // We use a loopback address with a timeout to exercise code paths. - addr, stop := startTestDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) { - resp := new(dns.Msg) - resp.SetReply(r) - resp.Answer = append(resp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("5.6.7.8"), - }) - _ = w.WriteMsg(resp) - }) - defer stop() - - host, portStr, _ := net.SplitHostPort(addr) - serverIP := net.ParseIP(host) - - // Exercise realExchange via Query — we need a way to target the test port. - // Use a custom exchange that calls through realExchange-like logic. - cfg := DefaultQueryConfig() - cfg.Retries = 1 - - resp, err := QueryWithExchange(context.Background(), serverIP, "example.com", TypeA, cfg, - func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - targetAddr := fmt.Sprintf("%s:%s", host, portStr) - c := &dns.Client{Net: "udp", ReadTimeout: 3 * time.Second, WriteTimeout: 3 * time.Second} - r, _, e := c.ExchangeContext(ctx, msg, targetAddr) - return r, e - }) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(resp.Answer) == 0 { - t.Fatal("expected answers") - } -} - -func TestQueryFunctionDirectly(t *testing.T) { - // Exercise Query() itself (which calls realExchange) by using 127.0.0.1:53. - // The test skips if no local DNS is available. - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - - server := net.ParseIP("127.0.0.1") - cfg := DefaultQueryConfig() - cfg.Retries = 1 - cfg.Timeout = 2 * time.Second - - _, err := Query(ctx, server, ".", TypeNS, cfg) - if err != nil { - t.Skipf("skipping (no local DNS at 127.0.0.1:53): %v", err) - } -} - -func TestIterativeQueryDirectly(t *testing.T) { - // Exercise IterativeQuery() itself (which calls realExchange) by using 127.0.0.1:53. - // The test skips if no local DNS is available. - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - - server := net.ParseIP("127.0.0.1") - cfg := DefaultQueryConfig() - cfg.Retries = 1 - cfg.Timeout = 2 * time.Second - - _, err := IterativeQuery(ctx, server, ".", TypeNS, cfg) - if err != nil { - t.Skipf("skipping (no local DNS at 127.0.0.1:53): %v", err) - } -} - -func TestBasicResolverQuery(t *testing.T) { - // Exercise BasicResolver.Query() which calls Query() → realExchange. - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - - br := NewBasicResolver() - server := net.ParseIP("127.0.0.1") - cfg := DefaultQueryConfig() - cfg.Retries = 1 - cfg.Timeout = 2 * time.Second - - _, err := br.Query(ctx, server, ".", TypeNS, cfg) - if err != nil { - t.Skipf("skipping (no local DNS at 127.0.0.1:53): %v", err) - } -} - -func TestDiscoverAllRoots(t *testing.T) { - // discoverAllRoots calls queryResolver(ctx, "127.0.0.1:53", ...) - // Skip if local DNS is not available. - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - cfg := &RootDiscoveryConfig{ - AllRoots: true, - IncludeAAAA: false, - } - servers, err := DiscoverRoots(ctx, cfg) - if err != nil { - t.Skipf("skipping (no local DNS available): %v", err) - } - if len(servers) == 0 { - t.Fatal("expected at least one root server from discoverAllRoots") - } -} - -func TestDiscoverAllRootsWithAAAA(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - cfg := &RootDiscoveryConfig{ - AllRoots: true, - IncludeAAAA: true, - } - servers, err := DiscoverRoots(ctx, cfg) - if err != nil { - t.Skipf("skipping (no local DNS available): %v", err) - } - if len(servers) == 0 { - t.Fatal("expected root servers with AAAA") - } -} - -func TestResolveRootServerDirect(t *testing.T) { - // Calls resolveRootServer directly (unexported, but in same package). - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - servers, err := resolveRootServer(ctx, systemResolver(), "a.root-servers.net.", false) - if err != nil { - t.Skipf("skipping (no local DNS): %v", err) - } - if len(servers) == 0 || len(servers[0].IPv4) == 0 { - t.Fatal("expected IPv4 address for a.root-servers.net.") - } -} - -func TestDiscoverSingleRootWithAAAA(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - cfg := &RootDiscoveryConfig{ - AllRoots: false, - IncludeAAAA: true, - } - servers, err := DiscoverRoots(ctx, cfg) - if err != nil { - t.Skipf("skipping (no local DNS available): %v", err) - } - if len(servers) == 0 { - t.Fatal("expected at least one root server") - } -} - -func TestRealExchangeTCPPath(t *testing.T) { - // Test the TCP path of realExchange via a test server - tcpAddr := "" - listener, err := net.Listen("tcp", "127.0.0.1:0") - if err != nil { - t.Skipf("cannot start TCP test server: %v", err) - } - tcpAddr = listener.Addr().String() - - mux := dns.NewServeMux() - mux.HandleFunc(".", func(w dns.ResponseWriter, r *dns.Msg) { - resp := new(dns.Msg) - resp.SetReply(r) - resp.Answer = append(resp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("9.9.9.9"), - }) - _ = w.WriteMsg(resp) - }) - - srv := &dns.Server{ - Listener: listener, - Net: "tcp", - Handler: mux, + switch network { + case "udp": + pc, err := net.ListenPacket("udp", "127.0.0.1:0") + if err != nil { + t.Skipf("cannot start test DNS server: %v", err) + } + srv.PacketConn = pc + addr = pc.LocalAddr().String() + case "tcp": + l, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Skipf("cannot start test DNS server: %v", err) + } + srv.Listener = l + addr = l.Addr().String() } started := make(chan struct{}) @@ -277,28 +45,89 @@ func TestRealExchangeTCPPath(t *testing.T) { select { case <-started: case <-time.After(2 * time.Second): - t.Skip("TCP DNS server didn't start") + t.Skip("test DNS server did not start in time") } - defer srv.Shutdown() - host, portStr, _ := net.SplitHostPort(tcpAddr) - serverIP := net.ParseIP(host) + t.Cleanup(func() { _ = srv.Shutdown() }) + return addr +} - cfg := DefaultQueryConfig() - cfg.UseTCP = true - cfg.Retries = 1 - - resp, err := QueryWithExchange(context.Background(), serverIP, "example.com", TypeA, cfg, - func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - targetAddr := fmt.Sprintf("%s:%s", host, portStr) - c := &dns.Client{Net: "tcp", ReadTimeout: 3 * time.Second, WriteTimeout: 3 * time.Second} - r, _, e := c.ExchangeContext(ctx, msg, targetAddr) - return r, e +func aHandler(ip string) dns.HandlerFunc { + return func(w dns.ResponseWriter, r *dns.Msg) { + resp := new(dns.Msg) + resp.SetReply(r) + resp.Answer = append(resp.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: r.Question[0].Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip), }) + _ = w.WriteMsg(resp) + } +} + +func TestRealExchangeUDPHostPort(t *testing.T) { + addr := startTestDNSServer(t, "udp", aHandler("1.2.3.4")) + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + + // realExchange must honour an explicit host:port (used by root discovery + // upstream resolvers). + resp, err := realExchange(ctx, addr, buildQuery("example.com.", TypeA, 2048), false) if err != nil { - t.Fatalf("TCP query: %v", err) + t.Fatalf("realExchange: %v", err) } if len(resp.Answer) == 0 { t.Fatal("expected answers") } } + +func TestRealExchangeTCP(t *testing.T) { + addr := startTestDNSServer(t, "tcp", aHandler("9.9.9.9")) + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + + resp, err := realExchange(ctx, addr, buildQuery("example.com.", TypeA, 2048), true) + if err != nil { + t.Fatalf("realExchange TCP: %v", err) + } + if len(resp.Answer) == 0 { + t.Fatal("expected answers") + } +} + +func TestRealExchangeUnreachable(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) + defer cancel() + + _, err := realExchange(ctx, "127.0.0.1:1", buildQuery("example.com.", TypeA, 2048), false) + if err == nil { + t.Fatal("expected error for unreachable server") + } +} + +func TestClientAgainstLocalServer(t *testing.T) { + var sawRD bool + addr := startTestDNSServer(t, "udp", func(w dns.ResponseWriter, r *dns.Msg) { + sawRD = r.RecursionDesired + aHandler("5.6.7.8")(w, r) + }) + + // Route the client's exchange to the test server's port while still + // exercising realExchange. + c := NewClient(&QueryConfig{Retries: 1, Timeout: 2 * time.Second}, + func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return realExchange(ctx, addr, msg, useTCP) + }) + + resp, _, err := c.Query(context.Background(), net.ParseIP("127.0.0.1"), "example.com", TypeA) + if err != nil { + t.Fatalf("Client.Query: %v", err) + } + if len(resp.Answer) == 0 { + t.Fatal("expected answers") + } + if sawRD { + t.Error("wire query must have RD=0") + } +} diff --git a/internal/dns/resolver.go b/internal/dns/resolver.go deleted file mode 100644 index ee9ed5f..0000000 --- a/internal/dns/resolver.go +++ /dev/null @@ -1,198 +0,0 @@ -package dns - -import ( - "context" - "fmt" - "net" - "sync" - "time" - - "github.com/miekg/dns" -) - -type Resolver interface { - Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) -} - -type BasicResolver struct{} - -func NewBasicResolver() *BasicResolver { - return &BasicResolver{} -} - -func (br *BasicResolver) Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { - return Query(ctx, server, name, qtype, cfg) -} - -type cacheKey struct { - server string - name string - qtype uint16 - qclass uint16 -} - -type cacheEntry struct { - msg *dns.Msg - expireAt time.Time -} - -func (e *cacheEntry) expired() bool { - return time.Now().After(e.expireAt) -} - -type CachingResolver struct { - inner Resolver - mu sync.RWMutex - cache map[cacheKey]*cacheEntry - defaultTTL time.Duration -} - -func NewCachingResolver(inner Resolver, opts ...CachingResolverOption) *CachingResolver { - if inner == nil { - inner = NewBasicResolver() - } - - cfg := &cachingResolverConfig{ - defaultTTL: 5 * time.Second, - } - for _, opt := range opts { - opt(cfg) - } - - return &CachingResolver{ - inner: inner, - cache: make(map[cacheKey]*cacheEntry), - defaultTTL: cfg.defaultTTL, - } -} - -func (cr *CachingResolver) Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { - if cfg == nil { - cfg = DefaultQueryConfig() - } - - fqdn := dns.Fqdn(name) - key := cacheKey{ - server: server.String(), - name: fqdn, - qtype: qtype, - qclass: dns.ClassINET, - } - - if resp, ok := cr.lookup(key); ok { - return resp, nil - } - - resp, err := cr.inner.Query(ctx, server, name, qtype, cfg) - if err != nil { - return nil, fmt.Errorf("caching resolver query: %w", err) - } - - cr.store(key, resp) - return resp, nil -} - -func (cr *CachingResolver) lookup(key cacheKey) (*dns.Msg, bool) { - cr.mu.RLock() - entry, ok := cr.cache[key] - cr.mu.RUnlock() - if !ok { - return nil, false - } - if entry.expired() { - return nil, false - } - return entry.msg.Copy(), true -} - -func (cr *CachingResolver) store(key cacheKey, msg *dns.Msg) { - ttl := minTTLFromMsg(msg) - if ttl <= 0 { - ttl = cr.defaultTTL - } - - cr.mu.Lock() - cr.cache[key] = &cacheEntry{ - msg: msg.Copy(), - expireAt: time.Now().Add(ttl), - } - cr.mu.Unlock() -} - -func (cr *CachingResolver) Len() int { - cr.mu.RLock() - n := len(cr.cache) - cr.mu.RUnlock() - return n -} - -func (cr *CachingResolver) Clear() { - cr.mu.Lock() - cr.cache = make(map[cacheKey]*cacheEntry) - cr.mu.Unlock() -} - -func (cr *CachingResolver) PurgeExpired() int { - cr.mu.Lock() - count := 0 - for k, e := range cr.cache { - if e.expired() { - delete(cr.cache, k) - count++ - } - } - cr.mu.Unlock() - return count -} - -type CachingResolverOption func(*cachingResolverConfig) - -type cachingResolverConfig struct { - defaultTTL time.Duration -} - -func WithDefaultTTL(d time.Duration) CachingResolverOption { - return func(c *cachingResolverConfig) { - c.defaultTTL = d - } -} - -func minTTLFromMsg(msg *dns.Msg) time.Duration { - if msg == nil { - return 0 - } - - var min uint32 - found := false - - for _, rr := range msg.Answer { - ttl := rr.Header().Ttl - if !found || ttl < min { - min = ttl - found = true - } - } - for _, rr := range msg.Ns { - ttl := rr.Header().Ttl - if !found || ttl < min { - min = ttl - found = true - } - } - for _, rr := range msg.Extra { - if _, ok := rr.(*dns.OPT); ok { - continue - } - ttl := rr.Header().Ttl - if !found || ttl < min { - min = ttl - found = true - } - } - - if !found { - return 0 - } - - return time.Duration(min) * time.Second -} diff --git a/internal/dns/resolver_test.go b/internal/dns/resolver_test.go deleted file mode 100644 index 2d067cf..0000000 --- a/internal/dns/resolver_test.go +++ /dev/null @@ -1,476 +0,0 @@ -package dns - -import ( - "context" - "errors" - "net" - "sync" - "testing" - "time" - - "github.com/miekg/dns" -) - -type mockResolver struct { - mu sync.Mutex - calls int - response *dns.Msg - err error -} - -func (m *mockResolver) Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { - m.mu.Lock() - defer m.mu.Unlock() - m.calls++ - if m.err != nil { - return nil, m.err - } - if m.response != nil { - return m.response.Copy(), nil - } - return nil, errors.New("no response configured") -} - -func (m *mockResolver) callCount() int { - m.mu.Lock() - defer m.mu.Unlock() - return m.calls -} - -func makeResponse(name string, qtype uint16, ttl uint32) *dns.Msg { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - switch qtype { - case dns.TypeA: - resp.Answer = append(resp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: ttl}, - A: net.ParseIP("93.184.216.34"), - }) - case dns.TypeNS: - resp.Answer = append(resp.Answer, &dns.NS{ - Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: ttl}, - Ns: "ns1.example.com.", - }) - } - return resp -} - -func TestNewCachingResolver_NilInner(t *testing.T) { - cr := NewCachingResolver(nil) - if cr == nil { - t.Fatal("NewCachingResolver(nil) returned nil") - } - if cr.inner == nil { - t.Fatal("expected inner resolver to be set when nil passed") - } -} - -func TestCachingResolver_CacheHit(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - resp1, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Fatalf("first query: %v", err) - } - if mock.callCount() != 1 { - t.Fatalf("expected 1 call after first query, got %d", mock.callCount()) - } - - resp2, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Fatalf("second query: %v", err) - } - if mock.callCount() != 1 { - t.Fatalf("expected still 1 call after second query (cache hit), got %d", mock.callCount()) - } - - if len(resp1.Answer) != len(resp2.Answer) { - t.Errorf("cached response has different number of answers") - } -} - -func TestCachingResolver_CacheMissDifferentServer(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, _ = cr.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA, cfg) - if mock.callCount() != 1 { - t.Fatalf("expected 1 call, got %d", mock.callCount()) - } - - _, _ = cr.Query(context.Background(), net.ParseIP("1.1.1.1"), "example.com", TypeA, cfg) - if mock.callCount() != 2 { - t.Fatalf("expected 2 calls for different server, got %d", mock.callCount()) - } -} - -func TestCachingResolver_CacheMissDifferentName(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if mock.callCount() != 1 { - t.Fatalf("expected 1 call, got %d", mock.callCount()) - } - - _, _ = cr.Query(context.Background(), server, "different.com", TypeA, cfg) - if mock.callCount() != 2 { - t.Fatalf("expected 2 calls for different name, got %d", mock.callCount()) - } -} - -func TestCachingResolver_CacheMissDifferentType(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if mock.callCount() != 1 { - t.Fatalf("expected 1 call, got %d", mock.callCount()) - } - - _, _ = cr.Query(context.Background(), server, "example.com", TypeNS, cfg) - if mock.callCount() != 2 { - t.Fatalf("expected 2 calls for different qtype, got %d", mock.callCount()) - } -} - -func TestCachingResolver_ErrorNotCached(t *testing.T) { - mock := &mockResolver{ - err: errors.New("connection refused"), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err == nil { - t.Fatal("expected error from mock") - } - - mock.err = nil - mock.response = makeResponse("example.com", dns.TypeA, 300) - - _, err = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Fatalf("second query after mock fixed: %v", err) - } - if mock.callCount() != 2 { - t.Fatalf("expected 2 calls (error not cached), got %d", mock.callCount()) - } -} - -func TestCachingResolver_TTLExpiry(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 1), - } - cr := NewCachingResolver(mock, WithDefaultTTL(1*time.Second)) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Fatalf("first query: %v", err) - } - if mock.callCount() != 1 { - t.Fatalf("expected 1 call, got %d", mock.callCount()) - } - - time.Sleep(2 * time.Second) - - _, err = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Fatalf("query after TTL expiry: %v", err) - } - if mock.callCount() != 2 { - t.Fatalf("expected 2 calls after TTL expiry, got %d", mock.callCount()) - } -} - -func TestCachingResolver_DefaultTTL(t *testing.T) { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - - mock := &mockResolver{response: resp} - cr := NewCachingResolver(mock, WithDefaultTTL(100*time.Millisecond)) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, err := cr.Query(context.Background(), server, "nodata.com", TypeA, cfg) - if err != nil { - t.Fatalf("first query: %v", err) - } - - time.Sleep(150 * time.Millisecond) - - _, err = cr.Query(context.Background(), server, "nodata.com", TypeA, cfg) - if err != nil { - t.Fatalf("query after default TTL expiry: %v", err) - } - if mock.callCount() != 2 { - t.Fatalf("expected 2 calls after default TTL, got %d", mock.callCount()) - } -} - -func TestCachingResolver_NilConfig(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - - server := net.ParseIP("8.8.8.8") - _, err := cr.Query(context.Background(), server, "example.com", TypeA, nil) - if err != nil { - t.Fatalf("nil config: %v", err) - } -} - -func TestCachingResolver_Len(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - if cr.Len() != 0 { - t.Fatalf("expected empty cache, got %d", cr.Len()) - } - - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if cr.Len() != 1 { - t.Fatalf("expected cache len 1, got %d", cr.Len()) - } - - _, _ = cr.Query(context.Background(), server, "other.com", TypeA, cfg) - if cr.Len() != 2 { - t.Fatalf("expected cache len 2, got %d", cr.Len()) - } -} - -func TestCachingResolver_Clear(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - _, _ = cr.Query(context.Background(), server, "other.com", TypeA, cfg) - if cr.Len() != 2 { - t.Fatalf("expected cache len 2, got %d", cr.Len()) - } - - cr.Clear() - if cr.Len() != 0 { - t.Fatalf("expected cache len 0 after clear, got %d", cr.Len()) - } - - _, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Fatalf("query after clear: %v", err) - } - if mock.callCount() != 3 { - t.Fatalf("expected 3 calls (2 before clear + 1 after clear), got %d", mock.callCount()) - } -} - -func TestCachingResolver_PurgeExpired(t *testing.T) { - resp := makeResponse("example.com", dns.TypeA, 0) - for _, rr := range resp.Answer { - rr.Header().Ttl = 0 - } - mock := &mockResolver{response: resp} - cr := NewCachingResolver(mock, WithDefaultTTL(50*time.Millisecond)) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if cr.Len() != 1 { - t.Fatalf("expected cache len 1, got %d", cr.Len()) - } - - time.Sleep(100 * time.Millisecond) - - purged := cr.PurgeExpired() - if purged != 1 { - t.Fatalf("expected 1 purged entry, got %d", purged) - } - if cr.Len() != 0 { - t.Fatalf("expected cache len 0 after purge, got %d", cr.Len()) - } -} - -func TestCachingResolver_PurgeExpiredNoneExpired(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - purged := cr.PurgeExpired() - if purged != 0 { - t.Fatalf("expected 0 purged entries, got %d", purged) - } - if cr.Len() != 1 { - t.Fatalf("expected cache len 1, got %d", cr.Len()) - } -} - -func TestCachingResolver_ResponseCopy(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - resp1, _ := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - resp2, _ := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - - if resp1 == resp2 { - t.Fatal("cache should return a copy, not the same pointer") - } -} - -func TestCachingResolver_MinTTLFromMsg(t *testing.T) { - tests := []struct { - name string - msg *dns.Msg - expected time.Duration - }{ - { - name: "nil message", - msg: nil, - expected: 0, - }, - { - name: "empty response", - msg: new(dns.Msg), - expected: 0, - }, - { - name: "single answer with low TTL", - msg: func() *dns.Msg { - m := new(dns.Msg) - m.Answer = append(m.Answer, &dns.A{ - Hdr: dns.RR_Header{Ttl: 60}, - }) - return m - }(), - expected: 60 * time.Second, - }, - { - name: "multiple records with varying TTLs", - msg: func() *dns.Msg { - m := new(dns.Msg) - m.Answer = append(m.Answer, &dns.A{ - Hdr: dns.RR_Header{Ttl: 300}, - }) - m.Ns = append(m.Ns, &dns.NS{ - Hdr: dns.RR_Header{Ttl: 120}, - }) - return m - }(), - expected: 120 * time.Second, - }, - { - name: "OPT record excluded from TTL calculation", - msg: func() *dns.Msg { - m := new(dns.Msg) - m.Answer = append(m.Answer, &dns.A{ - Hdr: dns.RR_Header{Ttl: 300}, - }) - m.Extra = append(m.Extra, &dns.OPT{ - Hdr: dns.RR_Header{Ttl: 0}, - }) - return m - }(), - expected: 300 * time.Second, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := minTTLFromMsg(tt.msg) - if got != tt.expected { - t.Errorf("minTTLFromMsg() = %v, want %v", got, tt.expected) - } - }) - } -} - -func TestCachingResolver_ConcurrentAccess(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - var wg sync.WaitGroup - for i := 0; i < 100; i++ { - wg.Add(1) - go func() { - defer wg.Done() - _, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Errorf("concurrent query failed: %v", err) - } - }() - } - wg.Wait() - - if mock.callCount() < 1 { - t.Fatalf("expected at least 1 call to mock, got %d", mock.callCount()) - } -} - -func TestBasicResolver(t *testing.T) { - br := NewBasicResolver() - if br == nil { - t.Fatal("NewBasicResolver returned nil") - } - - if _, ok := interface{}(br).(Resolver); !ok { - t.Fatal("BasicResolver does not implement Resolver interface") - } -} - -// TestBasicResolverQueryIntegration calls Query via the BasicResolver against -// the local system resolver. Skipped when no local resolver is reachable. -func TestBasicResolverQueryIntegration(t *testing.T) { -r := NewBasicResolver() -server := net.ParseIP("127.0.0.1") -ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) -defer cancel() - -// Cover BasicResolver.Query; skip if 127.0.0.1:53 is not available. -msg, err := r.Query(ctx, server, ".", TypeNS, nil) -if err != nil { -t.Logf("skipping (local resolver unavailable): %v", err) -t.Skip() -} -if msg == nil { -t.Fatal("expected non-nil response from BasicResolver.Query") -} -} diff --git a/internal/dns/robustness_test.go b/internal/dns/robustness_test.go index c37fc70..bdbb356 100644 --- a/internal/dns/robustness_test.go +++ b/internal/dns/robustness_test.go @@ -115,25 +115,23 @@ func TestSynthesizeCNAMEFromDNAME(t *testing.T) { } } -func TestBackoffDelay(t *testing.T) { - tests := []struct { - attempt int - want time.Duration - }{ - {0, 0}, - {1, 100 * time.Millisecond}, - {2, 200 * time.Millisecond}, - {3, 400 * time.Millisecond}, - {4, 800 * time.Millisecond}, - {5, 1600 * time.Millisecond}, - {6, 2000 * time.Millisecond}, // capped at 2s - {10, 2000 * time.Millisecond}, // still capped +func TestQueryConfigWithDefaults(t *testing.T) { + got := (&QueryConfig{}).withDefaults() + if got.UDPSize != DefaultEDNS0UDPSize() { + t.Errorf("UDPSize = %d, want %d", got.UDPSize, DefaultEDNS0UDPSize()) + } + if got.Timeout != 2*time.Second { + t.Errorf("Timeout = %v, want 2s", got.Timeout) + } + if got.Retries != 1 { + t.Errorf("Retries = %d, want clamp to 1", got.Retries) + } + if got.RetryDelay != 2*time.Second { + t.Errorf("RetryDelay = %v, want 2s", got.RetryDelay) } - for _, tt := range tests { - got := backoffDelay(tt.attempt) - if got != tt.want { - t.Errorf("backoffDelay(%d) = %v, want %v", tt.attempt, got, tt.want) - } + var nilCfg *QueryConfig + if nilCfg.withDefaults().Retries != DefaultQueryConfig().Retries { + t.Error("nil config should yield defaults") } } diff --git a/internal/dns/roots.go b/internal/dns/roots.go index 1a5dfbd..4aaf606 100644 --- a/internal/dns/roots.go +++ b/internal/dns/roots.go @@ -2,10 +2,10 @@ package dns import ( "context" + "errors" "fmt" "net" "strings" - "time" "github.com/miekg/dns" ) @@ -26,15 +26,22 @@ func (rs *RootServer) AllIPs(includeAAAA bool) []net.IP { } type RootDiscoveryConfig struct { - // Server overrides which root server to use as the traversal starting point. - // When empty, a root server is discovered via the upstream resolver. + // Server overrides which root server to use as the traversal starting + // point. It accepts a hostname (resolved to A via the upstream resolver) + // or an IP literal (used directly, no lookup). Server string - // Resolver is the upstream DNS resolver used to resolve root server names. - // When empty, the system resolver from /etc/resolv.conf is used. - // Format: "host:port" (e.g. "8.8.8.8:53" or "1.1.1.1:53"). + // Resolver is the upstream DNS resolver used to discover and resolve + // root server names. When empty, the system resolver configuration is + // used. Format: "host:port" (e.g. "8.8.8.8:53" or "1.1.1.1:53"). Resolver string AllRoots bool IncludeAAAA bool + // Query controls transport parameters (retries, timeout, TCP fallback) + // for discovery queries. nil means DefaultQueryConfig. + Query *QueryConfig + // Exchange overrides the wire exchange; nil means the real network. + // Tests inject a mock here so discovery is network-free. + Exchange ExchangeFunc } func DefaultRootDiscoveryConfig() *RootDiscoveryConfig { @@ -49,31 +56,59 @@ func DiscoverRoots(ctx context.Context, cfg *RootDiscoveryConfig) ([]RootServer, cfg = DefaultRootDiscoveryConfig() } - resolver := resolverFromConfig(cfg) + // --root-server with an IP literal: use it directly, never look it up. + if cfg.Server != "" { + if ip := net.ParseIP(cfg.Server); ip != nil { + rs := RootServer{Name: cfg.Server} + if ip.To4() != nil { + rs.IPv4 = []net.IP{ip} + } else { + rs.IPv6 = []net.IP{ip} + } + return []RootServer{rs}, nil + } + } + + resolver, err := resolverFromConfig(cfg) + if err != nil { + if cfg.Server != "" { + return nil, err + } + return rootHintsFallback(cfg, err) + } if cfg.Server != "" { - return discoverRootOverride(ctx, resolver, cfg.Server, cfg.IncludeAAAA) + return resolveRootServer(ctx, cfg, resolver, cfg.Server) } if cfg.AllRoots { - servers, err := discoverAllRoots(ctx, resolver, cfg.IncludeAAAA) + servers, err := discoverAllRoots(ctx, cfg, resolver) if err != nil { - return filterHints(RootHints, cfg.IncludeAAAA), nil + return rootHintsFallback(cfg, err) } return servers, nil } - servers, err := discoverSingleRoot(ctx, resolver, cfg.IncludeAAAA) + servers, err := discoverSingleRoot(ctx, cfg, resolver) if err != nil { - hints := filterHints(RootHints, cfg.IncludeAAAA) - if len(hints) > 0 { - return hints[:1], nil - } - return nil, err + return rootHintsFallback(cfg, err) } return servers, nil } +// rootHintsFallback returns the builtin IANA hints when the upstream resolver +// cannot provide roots (one entry unless AllRoots). +func rootHintsFallback(cfg *RootDiscoveryConfig, err error) ([]RootServer, error) { + hints := filterHints(RootHints, cfg.IncludeAAAA) + if len(hints) == 0 { + return nil, err + } + if cfg.AllRoots { + return hints, nil + } + return hints[:1], nil +} + // filterHints returns a copy of hints with IPv6 addresses stripped when includeAAAA is false. func filterHints(hints []RootServer, includeAAAA bool) []RootServer { out := make([]RootServer, len(hints)) @@ -86,88 +121,114 @@ func filterHints(hints []RootServer, includeAAAA bool) []RootServer { return out } -func discoverRootOverride(ctx context.Context, resolver, server string, includeAAAA bool) ([]RootServer, error) { - nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS) +// discoverSingleRoot mirrors get_a_root in traverser.rb: ask the upstream +// resolver for the root NS set, prefer glue from the additional section, and +// only fall back to explicit A/AAAA lookups when no glue was supplied. +func discoverSingleRoot(ctx context.Context, cfg *RootDiscoveryConfig, resolver string) ([]RootServer, error) { + names, nsMsg, err := rootNSNames(ctx, cfg, resolver) if err != nil { - return nil, fmt.Errorf("query root NS records: %w", err) + return nil, err } - nsSet := extractNSRecords(nsMsg.Answer) - if len(nsSet) == 0 { - nsSet = extractNSNames(nsMsg.Ns) - } - - normalized := normalizeServerName(server) - for _, name := range nsSet { - if normalizeServerName(name) == normalized { - return resolveRootServer(ctx, resolver, name, includeAAAA) + for _, name := range names { + rs := rootFromAdditional(nsMsg, name, cfg.IncludeAAAA) + if len(rs.AllIPs(cfg.IncludeAAAA)) > 0 { + return []RootServer{singleAddress(rs)}, nil } } - return resolveRootServer(ctx, resolver, server, includeAAAA) + var lastErr error + for _, name := range names { + servers, err := resolveRootServer(ctx, cfg, resolver, name) + if err != nil { + lastErr = err + continue + } + return []RootServer{singleAddress(servers[0])}, nil + } + + if lastErr == nil { + lastErr = errors.New("no address could be found for any root server") + } + return nil, lastErr } -func discoverSingleRoot(ctx context.Context, resolver string, includeAAAA bool) ([]RootServer, error) { - nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS) - if err != nil { - return nil, fmt.Errorf("query root NS records: %w", err) +// singleAddress narrows rs to its first address, mirroring get_a_root in +// traverser.rb (add[0]/ans2[0]): the single-root start point is exactly one +// (name, IP) pair even when the upstream supplies more (or duplicate) glue. +func singleAddress(rs RootServer) RootServer { + if len(rs.IPv4) > 0 { + return RootServer{Name: rs.Name, IPv4: rs.IPv4[:1]} } - - nsSet := extractNSRecords(nsMsg.Answer) - if len(nsSet) == 0 { - nsSet = extractNSNames(nsMsg.Ns) + if len(rs.IPv6) > 0 { + return RootServer{Name: rs.Name, IPv6: rs.IPv6[:1]} } - - if len(nsSet) == 0 { - return nil, fmt.Errorf("no root NS records found in response") - } - - pick := nsSet[0] - return resolveRootServer(ctx, resolver, pick, includeAAAA) + return rs } -func discoverAllRoots(ctx context.Context, resolver string, includeAAAA bool) ([]RootServer, error) { - nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS) +// discoverAllRoots mirrors find_all_roots in traverser.rb: it returns one +// RootServer per root NS name with its full address set, so the traversal can +// branch per root. Roots with no resolvable address are skipped. +func discoverAllRoots(ctx context.Context, cfg *RootDiscoveryConfig, resolver string) ([]RootServer, error) { + names, nsMsg, err := rootNSNames(ctx, cfg, resolver) if err != nil { - return nil, fmt.Errorf("query root NS records: %w", err) - } - - nsSet := extractNSRecords(nsMsg.Answer) - if len(nsSet) == 0 { - nsSet = extractNSNames(nsMsg.Ns) - } - - if len(nsSet) == 0 { - return nil, fmt.Errorf("no root NS records found in response") + return nil, err } var servers []RootServer - for _, name := range nsSet { - resolved, err := resolveRootServer(ctx, resolver, name, includeAAAA) - if err != nil { - servers = append(servers, RootServer{Name: name}) - continue + for _, name := range names { + rs := rootFromAdditional(nsMsg, name, cfg.IncludeAAAA) + if len(rs.AllIPs(cfg.IncludeAAAA)) == 0 { + resolved, err := resolveRootServer(ctx, cfg, resolver, name) + if err != nil { + continue + } + rs = resolved[0] } - servers = append(servers, resolved...) + servers = append(servers, rs) } if len(servers) == 0 { - return nil, fmt.Errorf("failed to resolve any root servers") + return nil, errors.New("failed to resolve any root servers") } return servers, nil } -func resolveRootServer(ctx context.Context, resolver, name string, includeAAAA bool) ([]RootServer, error) { +func rootNSNames(ctx context.Context, cfg *RootDiscoveryConfig, resolver string) ([]string, *dns.Msg, error) { + nsMsg, err := queryUpstream(ctx, cfg, resolver, ".", dns.TypeNS) + if err != nil { + return nil, nil, fmt.Errorf("query root NS records: %w", err) + } + + names := extractNSRecords(nsMsg.Answer) + if len(names) == 0 { + names = extractNSNames(nsMsg.Ns) + } + if len(names) == 0 { + return nil, nil, errors.New("no root NS records found in response") + } + return names, nsMsg, nil +} + +func rootFromAdditional(msg *dns.Msg, name string, includeAAAA bool) RootServer { + rs := RootServer{Name: name, IPv4: additionalIPs(msg, name, dns.TypeA)} + if includeAAAA { + rs.IPv6 = additionalIPs(msg, name, dns.TypeAAAA) + } + return rs +} + +func resolveRootServer(ctx context.Context, cfg *RootDiscoveryConfig, resolver, name string) ([]RootServer, error) { var ipv4 []net.IP - aMsg, err := queryResolver(ctx, resolver, name, dns.TypeA) + aMsg, err := queryUpstream(ctx, cfg, resolver, name, dns.TypeA) if err == nil { ipv4 = extractIPsFromAnswer(aMsg.Answer, dns.TypeA) } var ipv6 []net.IP - if includeAAAA { - aaaaMsg, err := queryResolver(ctx, resolver, name, dns.TypeAAAA) + if cfg.IncludeAAAA { + aaaaMsg, err := queryUpstream(ctx, cfg, resolver, name, dns.TypeAAAA) if err == nil { ipv6 = extractIPsFromAnswer(aaaaMsg.Answer, dns.TypeAAAA) } @@ -180,47 +241,53 @@ func resolveRootServer(ctx context.Context, resolver, name string, includeAAAA b return []RootServer{{Name: name, IPv4: ipv4, IPv6: ipv6}}, nil } -// resolverFromConfig returns the upstream DNS resolver address to use. -// If cfg.Resolver is set, it is used directly. Otherwise the system resolver -// is read from /etc/resolv.conf. Falls back to 127.0.0.1:53 if neither is available. -func resolverFromConfig(cfg *RootDiscoveryConfig) string { +// resolverFromConfig returns the upstream DNS resolver address to use: +// cfg.Resolver when set, otherwise the system resolver configuration. There +// is deliberately no hardcoded address fallback. +func resolverFromConfig(cfg *RootDiscoveryConfig) (string, error) { if cfg != nil && cfg.Resolver != "" { - return cfg.Resolver + return cfg.Resolver, nil } return systemResolver() } -// systemResolver returns the first nameserver from the system DNS configuration. -// This is Unix-only: it reads /etc/resolv.conf, which does not exist on Windows. -// On Windows (or any system without /etc/resolv.conf) the fallback 127.0.0.1:53 applies. -func systemResolver() string { +// systemResolver returns the first nameserver from the system DNS +// configuration. This is Unix-only: it reads /etc/resolv.conf, which does not +// exist on Windows; there the caller falls back to the builtin root hints. +func systemResolver() (string, error) { cc, err := dns.ClientConfigFromFile("/etc/resolv.conf") - if err != nil || len(cc.Servers) == 0 { - return "127.0.0.1:53" + if err != nil { + return "", fmt.Errorf("read system resolver config: %w", err) } - return net.JoinHostPort(cc.Servers[0], cc.Port) + if len(cc.Servers) == 0 { + return "", errors.New("no nameservers found in system resolver config") + } + return net.JoinHostPort(cc.Servers[0], cc.Port), nil } -func queryResolver(ctx context.Context, resolverAddr, name string, qtype uint16) (*dns.Msg, error) { - c := &dns.Client{ - Net: "udp", - ReadTimeout: 5 * time.Second, - WriteTimeout: 5 * time.Second, +// queryUpstream asks the upstream resolver with recursion desired — the only +// RD=1 path in the program — using the same retry/TCP-fallback machinery as +// traversal queries. +func queryUpstream(ctx context.Context, cfg *RootDiscoveryConfig, resolverAddr, name string, qtype uint16) (*dns.Msg, error) { + var qcfg *QueryConfig + var exchange ExchangeFunc + if cfg != nil { + qcfg = cfg.Query + exchange = cfg.Exchange } - if deadline, ok := ctx.Deadline(); ok { - c.ReadTimeout = time.Until(deadline) - c.WriteTimeout = time.Until(deadline) + qcfg = qcfg.withDefaults() + if exchange == nil { + exchange = realExchange } m := new(dns.Msg) m.SetQuestion(dns.Fqdn(name), qtype) m.RecursionDesired = true - - r, _, err := c.ExchangeContext(ctx, m, resolverAddr) - if err != nil { - return nil, fmt.Errorf("resolver exchange %s %s: %w", name, QNameType(qtype), err) + if qcfg.UDPSize > MinEDNS0UDPSize() { + m.SetEdns0(uint16(qcfg.UDPSize), false) } - return r, nil + + return exchangeWithRetry(ctx, exchange, resolverAddr, m, qcfg) } func extractNSRecords(rrs []dns.RR) []string { @@ -259,6 +326,23 @@ func extractIPsFromAnswer(rrs []dns.RR, qtype uint16) []net.IP { return ips } -func normalizeServerName(name string) string { - return strings.TrimSuffix(strings.ToLower(name), ".") +// additionalIPs returns glue addresses for name from the additional section. +func additionalIPs(msg *dns.Msg, name string, qtype uint16) []net.IP { + var ips []net.IP + for _, rr := range msg.Extra { + if !strings.EqualFold(rr.Header().Name, name) { + continue + } + switch v := rr.(type) { + case *dns.A: + if qtype == dns.TypeA { + ips = append(ips, v.A) + } + case *dns.AAAA: + if qtype == dns.TypeAAAA { + ips = append(ips, v.AAAA) + } + } + } + return ips } diff --git a/internal/dns/roots_test.go b/internal/dns/roots_test.go index bc93467..b504087 100644 --- a/internal/dns/roots_test.go +++ b/internal/dns/roots_test.go @@ -2,9 +2,10 @@ package dns import ( "context" + "errors" "net" + "strings" "testing" - "time" "github.com/miekg/dns" ) @@ -35,9 +36,8 @@ func TestRootServerAllIPs(t *testing.T) { t.Run("no addresses", func(t *testing.T) { empty := RootServer{Name: "empty.root-servers.net."} - ips := empty.AllIPs(false) - if len(ips) != 0 { - t.Errorf("expected 0 IPs, got %d", len(ips)) + if len(empty.AllIPs(false)) != 0 { + t.Error("expected 0 IPs") } }) } @@ -55,31 +55,317 @@ func TestDefaultRootDiscoveryConfig(t *testing.T) { } } -func TestNormalizeServerName(t *testing.T) { - tests := []struct { - input string - want string - }{ - {"a.root-servers.net.", "a.root-servers.net"}, - {"A.ROOT-SERVERS.NET.", "a.root-servers.net"}, - {"b.root-servers.net", "b.root-servers.net"}, - {"root-servers.net.", "root-servers.net"}, +// mockUpstream builds an ExchangeFunc that answers root discovery queries. +// glue controls whether A records are put in the additional section of the +// NS response. +func mockUpstream(t *testing.T, roots map[string]string, glue bool) ExchangeFunc { + t.Helper() + return func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + if !msg.RecursionDesired { + t.Error("root discovery must query the upstream resolver with RD=1") + } + q := msg.Question[0] + m := new(dns.Msg) + m.SetReply(msg) + switch { + case q.Name == "." && q.Qtype == dns.TypeNS: + for name, ip := range roots { + m.Answer = append(m.Answer, &dns.NS{ + Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, + Ns: name, + }) + if glue { + m.Extra = append(m.Extra, &dns.A{ + Hdr: dns.RR_Header{Name: name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip), + }) + } + } + case q.Qtype == dns.TypeA: + if ip, ok := roots[q.Name]; ok { + m.Answer = append(m.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: q.Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip), + }) + } + } + return m, nil + } +} + +func TestDiscoverRootsIPLiteral(t *testing.T) { + cfg := &RootDiscoveryConfig{ + Server: "192.203.230.10", + Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + t.Error("IP literal root must not trigger any lookup") + return nil, errors.New("no network") + }, + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 1 { + t.Fatalf("expected 1 root, got %d", len(servers)) + } + if servers[0].Name != "192.203.230.10" { + t.Errorf("Name = %q, want the IP literal", servers[0].Name) + } + if len(servers[0].IPv4) != 1 || !servers[0].IPv4[0].Equal(net.ParseIP("192.203.230.10")) { + t.Errorf("IPv4 = %v, want [192.203.230.10]", servers[0].IPv4) + } +} + +func TestDiscoverRootsIPv6Literal(t *testing.T) { + cfg := &RootDiscoveryConfig{Server: "2001:503:ba3e::2:30"} + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 1 || len(servers[0].IPv6) != 1 { + t.Fatalf("expected 1 root with 1 IPv6 address, got %+v", servers) + } +} + +func TestDiscoverRootsHostnameOverride(t *testing.T) { + roots := map[string]string{"e.root-servers.net.": "192.203.230.10"} + cfg := &RootDiscoveryConfig{ + Server: "e.root-servers.net", + Resolver: "192.0.2.53:53", + Exchange: mockUpstream(t, roots, false), + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 1 { + t.Fatalf("expected 1 root, got %d", len(servers)) + } + if servers[0].Name != "e.root-servers.net" { + t.Errorf("Name = %q", servers[0].Name) + } + if len(servers[0].IPv4) != 1 || !servers[0].IPv4[0].Equal(net.ParseIP("192.203.230.10")) { + t.Errorf("IPv4 = %v", servers[0].IPv4) + } +} + +func TestDiscoverRootsHostnameOverrideUnresolvable(t *testing.T) { + cfg := &RootDiscoveryConfig{ + Server: "nonexistent.root-servers.net", + Resolver: "192.0.2.53:53", + Exchange: mockUpstream(t, map[string]string{}, false), + } + if _, err := DiscoverRoots(context.Background(), cfg); err == nil { + t.Fatal("expected error for unresolvable root override") + } +} + +func TestDiscoverRootsSingleFromGlue(t *testing.T) { + roots := map[string]string{"a.root-servers.net.": "198.41.0.4"} + var wireQueries int + inner := mockUpstream(t, roots, true) + cfg := &RootDiscoveryConfig{ + Resolver: "192.0.2.53:53", + Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + wireQueries++ + return inner(ctx, server, msg, useTCP) + }, + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 1 { + t.Fatalf("expected exactly one root, got %d", len(servers)) + } + if len(servers[0].IPv4) == 0 { + t.Error("expected glue A address") + } + if wireQueries != 1 { + t.Errorf("glue should satisfy discovery in one query, got %d", wireQueries) + } +} + +func TestDiscoverRootsSingleWithoutGlue(t *testing.T) { + roots := map[string]string{"a.root-servers.net.": "198.41.0.4"} + cfg := &RootDiscoveryConfig{ + Resolver: "192.0.2.53:53", + Exchange: mockUpstream(t, roots, false), + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 1 || len(servers[0].IPv4) != 1 { + t.Fatalf("expected one root resolved via A lookup, got %+v", servers) + } +} + +func TestDiscoverAllRootsPerRootSet(t *testing.T) { + roots := map[string]string{ + "a.mock-roots.test.": "192.0.2.1", + "b.mock-roots.test.": "192.0.2.2", + "c.mock-roots.test.": "192.0.2.3", + } + cfg := &RootDiscoveryConfig{ + AllRoots: true, + Resolver: "192.0.2.53:53", + Exchange: mockUpstream(t, roots, true), + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 3 { + t.Fatalf("expected 3 roots, got %d", len(servers)) + } + seen := map[string]bool{} + for _, rs := range servers { + seen[rs.Name] = true + want, ok := roots[rs.Name] + if !ok { + t.Errorf("unexpected root %q", rs.Name) + continue + } + if len(rs.IPv4) != 1 || !rs.IPv4[0].Equal(net.ParseIP(want)) { + t.Errorf("root %s IPv4 = %v, want [%s]", rs.Name, rs.IPv4, want) + } + } + if len(seen) != 3 { + t.Errorf("roots not distinct: %v", seen) + } +} + +func TestDiscoverAllRootsSkipsUnresolvable(t *testing.T) { + // b has neither glue nor an A record: it must be skipped, like + // find_all_roots in traverser.rb. + roots := map[string]string{"a.mock-roots.test.": "192.0.2.1"} + exchange := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + q := msg.Question[0] + m := new(dns.Msg) + m.SetReply(msg) + switch { + case q.Name == "." && q.Qtype == dns.TypeNS: + for _, name := range []string{"a.mock-roots.test.", "b.mock-roots.test."} { + m.Answer = append(m.Answer, &dns.NS{ + Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, + Ns: name, + }) + } + case q.Qtype == dns.TypeA: + if ip, ok := roots[q.Name]; ok { + m.Answer = append(m.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: q.Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip), + }) + } + } + return m, nil + } + cfg := &RootDiscoveryConfig{ + AllRoots: true, + Resolver: "192.0.2.53:53", + Query: &QueryConfig{Retries: 1}, + Exchange: exchange, + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 1 || servers[0].Name != "a.mock-roots.test." { + t.Fatalf("expected only the resolvable root, got %+v", servers) + } +} + +func TestDiscoverRootsFallsBackToHints(t *testing.T) { + cfg := &RootDiscoveryConfig{ + Resolver: "192.0.2.53:53", + Query: &QueryConfig{Retries: 1}, + Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return nil, errors.New("upstream unreachable") + }, + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("expected hints fallback, got error: %v", err) + } + if len(servers) != 1 { + t.Fatalf("single-root mode must fall back to one hint, got %d", len(servers)) + } + if !strings.HasSuffix(servers[0].Name, ".root-servers.net.") { + t.Errorf("expected an IANA hint, got %q", servers[0].Name) } - for _, tt := range tests { - t.Run(tt.input, func(t *testing.T) { - got := normalizeServerName(tt.input) - if got != tt.want { - t.Errorf("normalizeServerName(%q) = %q, want %q", tt.input, got, tt.want) + cfg.AllRoots = true + servers, err = DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("expected hints fallback, got error: %v", err) + } + if len(servers) != len(RootHints) { + t.Errorf("all-roots fallback should return all %d hints, got %d", len(RootHints), len(servers)) + } +} + +func TestDiscoverRootsRetriesUpstream(t *testing.T) { + roots := map[string]string{"a.root-servers.net.": "198.41.0.4"} + inner := mockUpstream(t, roots, true) + var calls int + cfg := &RootDiscoveryConfig{ + Resolver: "192.0.2.53:53", + Query: &QueryConfig{Retries: 2, RetryDelay: 1}, + Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + if calls == 1 { + return nil, errors.New("transient failure") } - }) + return inner(ctx, server, msg, useTCP) + }, + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots should retry: %v", err) + } + if len(servers) != 1 { + t.Fatalf("expected 1 root after retry, got %d", len(servers)) + } + if calls != 2 { + t.Errorf("expected 2 attempts, got %d", calls) + } +} + +func TestDiscoverRootsTruncationTCPFallback(t *testing.T) { + roots := map[string]string{"a.root-servers.net.": "198.41.0.4"} + inner := mockUpstream(t, roots, true) + var sawTCP bool + cfg := &RootDiscoveryConfig{ + Resolver: "192.0.2.53:53", + Query: &QueryConfig{Retries: 1, AllowTCP: true}, + Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + if !useTCP { + m := new(dns.Msg) + m.SetReply(msg) + m.Truncated = true + return m, nil + } + sawTCP = true + return inner(ctx, server, msg, useTCP) + }, + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if !sawTCP { + t.Error("expected TCP fallback on truncated upstream response") + } + if len(servers) != 1 || len(servers[0].IPv4) == 0 { + t.Fatalf("expected root from TCP response, got %+v", servers) } } func TestExtractNSRecords(t *testing.T) { t.Run("empty", func(t *testing.T) { - names := extractNSRecords(nil) - if len(names) != 0 { + if names := extractNSRecords(nil); len(names) != 0 { t.Errorf("expected 0 names, got %d", len(names)) } }) @@ -103,18 +389,24 @@ func TestExtractNSRecords(t *testing.T) { &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "a.root-servers.net."}, &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "a.root-servers.net."}, } - names := extractNSRecords(rrs) - if len(names) != 1 { + if names := extractNSRecords(rrs); len(names) != 1 { t.Errorf("expected 1 deduped name, got %d", len(names)) } + }) + t.Run("non-NS ignored", func(t *testing.T) { + rrs := []dns.RR{ + &dns.A{Hdr: dns.RR_Header{Name: "a.root-servers.net.", Rrtype: dns.TypeA}, A: net.ParseIP("198.41.0.4")}, + } + if names := extractNSNames(rrs); len(names) != 0 { + t.Errorf("expected 0 names, got %d", len(names)) + } }) } func TestExtractIPsFromAnswer(t *testing.T) { t.Run("empty", func(t *testing.T) { - ips := extractIPsFromAnswer(nil, dns.TypeA) - if len(ips) != 0 { + if ips := extractIPsFromAnswer(nil, dns.TypeA); len(ips) != 0 { t.Errorf("expected 0 IPs, got %d", len(ips)) } }) @@ -124,347 +416,41 @@ func TestExtractIPsFromAnswer(t *testing.T) { &dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("198.41.0.4")}, &dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("199.9.14.201")}, } - ips := extractIPsFromAnswer(rrs, dns.TypeA) - if len(ips) != 2 { + if ips := extractIPsFromAnswer(rrs, dns.TypeA); len(ips) != 2 { t.Fatalf("expected 2 IPs, got %d", len(ips)) } }) - t.Run("AAAA records", func(t *testing.T) { - rrs := []dns.RR{ - &dns.AAAA{Hdr: dns.RR_Header{Rrtype: dns.TypeAAAA}, AAAA: net.ParseIP("2001:503:ba3e::2:30")}, - } - ips := extractIPsFromAnswer(rrs, dns.TypeAAAA) - if len(ips) != 1 { - t.Fatalf("expected 1 IP, got %d", len(ips)) - } - }) - t.Run("filter by type", func(t *testing.T) { rrs := []dns.RR{ &dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("198.41.0.4")}, &dns.AAAA{Hdr: dns.RR_Header{Rrtype: dns.TypeAAAA}, AAAA: net.ParseIP("2001:503:ba3e::2:30")}, } - ips := extractIPsFromAnswer(rrs, dns.TypeA) - if len(ips) != 1 { + if ips := extractIPsFromAnswer(rrs, dns.TypeA); len(ips) != 1 { t.Fatalf("expected 1 A IP, got %d", len(ips)) } + if ips := extractIPsFromAnswer(rrs, dns.TypeAAAA); len(ips) != 1 { + t.Fatalf("expected 1 AAAA IP, got %d", len(ips)) + } }) } -func TestDiscoverRootsOverride(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() - - cfg := &RootDiscoveryConfig{ - Server: "a.root-servers.net", - IncludeAAAA: false, - } - - servers, err := DiscoverRoots(ctx, cfg) - if err != nil { - t.Logf("skipping (no resolver available): %v", err) - t.Skip() - } - if len(servers) == 0 { - t.Fatal("expected at least one root server") - } - if len(servers[0].IPv4) == 0 { - t.Error("expected IPv4 addresses for a.root-servers.net") - } -} - -func TestDiscoverRootsSingle(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() - - servers, err := DiscoverRoots(ctx, nil) - if err != nil { - t.Logf("skipping (no resolver available): %v", err) - t.Skip() - } - if len(servers) == 0 { - t.Fatal("expected at least one root server") - } - if servers[0].Name == "" { - t.Error("root server name should not be empty") - } -} - -func TestDiscoverRootsNilConfig(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() - - servers, err := DiscoverRoots(ctx, nil) - if err != nil { - t.Logf("skipping (no resolver available): %v", err) - t.Skip() - } - if len(servers) == 0 { - t.Fatal("expected at least one root server with nil config") - } -} - -func TestBuildNSResponse(t *testing.T) { +func TestAdditionalIPs(t *testing.T) { msg := new(dns.Msg) - msg.SetReply(new(dns.Msg)) - msg.Answer = append(msg.Answer, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, Ns: "a.root-servers.net."}, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, Ns: "b.root-servers.net."}, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, Ns: "c.root-servers.net."}, + msg.Extra = append(msg.Extra, + &dns.A{Hdr: dns.RR_Header{Name: "A.ROOT-SERVERS.NET.", Rrtype: dns.TypeA, Class: dns.ClassINET}, A: net.ParseIP("198.41.0.4")}, + &dns.A{Hdr: dns.RR_Header{Name: "b.root-servers.net.", Rrtype: dns.TypeA, Class: dns.ClassINET}, A: net.ParseIP("170.247.170.2")}, + &dns.AAAA{Hdr: dns.RR_Header{Name: "a.root-servers.net.", Rrtype: dns.TypeAAAA, Class: dns.ClassINET}, AAAA: net.ParseIP("2001:503:ba3e::2:30")}, ) - names := extractNSRecords(msg.Answer) - if len(names) != 3 { - t.Fatalf("expected 3 NS records, got %d", len(names)) + ips := additionalIPs(msg, "a.root-servers.net.", dns.TypeA) + if len(ips) != 1 || !ips[0].Equal(net.ParseIP("198.41.0.4")) { + t.Errorf("case-insensitive glue match failed: %v", ips) } - - for _, name := range names { - if len(name) == 0 || name[len(name)-1] != '.' { - t.Errorf("expected FQDN, got %q", name) - } + if ips := additionalIPs(msg, "a.root-servers.net.", dns.TypeAAAA); len(ips) != 1 { + t.Errorf("expected 1 AAAA glue, got %v", ips) + } + if ips := additionalIPs(msg, "c.root-servers.net.", dns.TypeA); len(ips) != 0 { + t.Errorf("expected no glue for c, got %v", ips) } } - -func TestExtractNSNames(t *testing.T) { -rrs := []dns.RR{ -&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "a.root-servers.net."}, -&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "b.root-servers.net."}, -} -names := extractNSNames(rrs) -if len(names) != 2 { -t.Fatalf("extractNSNames: expected 2 names, got %d", len(names)) -} -} - -func TestExtractNSNamesEmpty(t *testing.T) { -names := extractNSNames(nil) -if len(names) != 0 { -t.Errorf("extractNSNames(nil): expected 0 names, got %d", len(names)) -} -} - -func TestExtractNSNamesNonNS(t *testing.T) { -rrs := []dns.RR{ -&dns.A{Hdr: dns.RR_Header{Name: "a.root-servers.net.", Rrtype: dns.TypeA}, A: net.ParseIP("198.41.0.4")}, -} -names := extractNSNames(rrs) -if len(names) != 0 { -t.Errorf("extractNSNames with A records: expected 0 names, got %d", len(names)) -} -} - -func TestDiscoverRootsAllRoots(t *testing.T) { -ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) -defer cancel() - -cfg := &RootDiscoveryConfig{ -AllRoots: true, -IncludeAAAA: false, -} - -servers, err := DiscoverRoots(ctx, cfg) -if err != nil { -t.Logf("skipping (no resolver available): %v", err) -t.Skip() -} -if len(servers) == 0 { -t.Fatal("expected root servers with AllRoots=true") -} -t.Logf("discovered %d root servers", len(servers)) -} - -func TestDiscoverRootsAllRootsIncludeAAAA(t *testing.T) { -ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) -defer cancel() - -cfg := &RootDiscoveryConfig{ -AllRoots: true, -IncludeAAAA: true, -} - -servers, err := DiscoverRoots(ctx, cfg) -if err != nil { -t.Logf("skipping (no resolver available): %v", err) -t.Skip() -} -if len(servers) == 0 { -t.Fatal("expected root servers") -} -} - -// startMockDNSServer starts a UDP DNS server on a random port that serves -// pre-configured responses. It returns the server address and a stop function. -func startMockDNSServer(t *testing.T, handlerFn dns.HandlerFunc) string { -t.Helper() - -mux := dns.NewServeMux() -mux.HandleFunc(".", handlerFn) - -srv := &dns.Server{ -Addr: "127.0.0.1:0", -Net: "udp", -Handler: mux, -} - -started := make(chan struct{}) -srv.NotifyStartedFunc = func() { close(started) } - -go func() { -if err := srv.ListenAndServe(); err != nil && t.Failed() { -return -} -}() - -select { -case <-started: -case <-time.After(2 * time.Second): -t.Fatal("mock DNS server did not start in time") -} - -// Retrieve the actual bound address from the server's PacketConn. -addr := srv.PacketConn.LocalAddr().String() -t.Cleanup(func() { _ = srv.Shutdown() }) -return addr -} - -func TestQueryResolverSuccess(t *testing.T) { - addr := startMockDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) { - m := new(dns.Msg) - m.SetReply(r) - m.Answer = append(m.Answer, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: "a.root-servers.net."}, - ) - _ = w.WriteMsg(m) - }) - - // queryResolver uses 5ns timeout without deadline; provide one - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - msg, err := queryResolver(ctx, addr, ".", dns.TypeNS) - if err != nil { - t.Fatalf("queryResolver: %v", err) - } - names := extractNSRecords(msg.Answer) - if len(names) == 0 { - t.Fatal("expected NS records in answer") - } -} - -func TestQueryResolverError(t *testing.T) { -ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) -defer cancel() -// Use an address nothing is listening on -_, err := queryResolver(ctx, "127.0.0.1:19999", ".", dns.TypeNS) -if err == nil { -t.Fatal("expected error for unreachable resolver") -} -} - -func TestDiscoverSingleRootWithMock(t *testing.T) { - addr := startMockDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) { - m := new(dns.Msg) - m.SetReply(r) - q := r.Question[0] - switch q.Qtype { - case dns.TypeNS: - m.Answer = append(m.Answer, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: "mock.root-servers.test."}, - ) - case dns.TypeA: - m.Answer = append(m.Answer, - &dns.A{Hdr: dns.RR_Header{Name: q.Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, A: net.ParseIP("127.0.0.1")}, - ) - } - _ = w.WriteMsg(m) - }) - - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - msg, err := queryResolver(ctx, addr, ".", dns.TypeNS) - if err != nil { - t.Fatalf("queryResolver: %v", err) - } - names := extractNSRecords(msg.Answer) - if len(names) == 0 { - names = extractNSNames(msg.Ns) - } - if len(names) == 0 { - t.Skip("mock NS query returned no NS records") - } - t.Logf("found %d root NS names from mock: %v", len(names), names) -} - -func TestDiscoverAllRootsWithMockServer(t *testing.T) { - addr := startMockDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) { - m := new(dns.Msg) - m.SetReply(r) - q := r.Question[0] - switch q.Qtype { - case dns.TypeNS: - m.Answer = append(m.Answer, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: "a.mock-roots.test."}, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: "b.mock-roots.test."}, - ) - case dns.TypeA: - m.Answer = append(m.Answer, - &dns.A{Hdr: dns.RR_Header{Name: q.Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, A: net.ParseIP("127.0.0.1")}, - ) - } - _ = w.WriteMsg(m) - }) - - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - msg, err := queryResolver(ctx, addr, ".", dns.TypeNS) - if err != nil { - t.Fatalf("queryResolver: %v", err) - } - names := extractNSRecords(msg.Answer) - if len(names) < 2 { - t.Fatalf("expected 2 NS names, got %d", len(names)) - } - - // Also cover AAAA path - aaaaMsg, err := queryResolver(ctx, addr, "a.mock-roots.test.", dns.TypeAAAA) - if err != nil { - t.Logf("AAAA query error (acceptable): %v", err) - } else { - t.Logf("AAAA query returned %d answers", len(aaaaMsg.Answer)) - } -} - -func TestMinTTLFromMsgWithExtraRecords(t *testing.T) { -msg := new(dns.Msg) -msg.Answer = append(msg.Answer, &dns.A{ -Hdr: dns.RR_Header{Ttl: 300}, -A: net.ParseIP("1.2.3.4"), -}) -// Extra record (non-OPT) with smaller TTL -msg.Extra = append(msg.Extra, &dns.NS{ -Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Ttl: 60}, -Ns: "a.root-servers.net.", -}) - -ttl := minTTLFromMsg(msg) -if ttl != 60*time.Second { -t.Errorf("minTTLFromMsg = %v, want 60s", ttl) -} -} - -func TestMinTTLFromMsgOPTIgnored(t *testing.T) { -msg := new(dns.Msg) -msg.Answer = append(msg.Answer, &dns.A{ -Hdr: dns.RR_Header{Ttl: 300}, -A: net.ParseIP("1.2.3.4"), -}) -// OPT record should be ignored -msg.Extra = append(msg.Extra, &dns.OPT{ -Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeOPT}, -}) - -ttl := minTTLFromMsg(msg) -if ttl != 300*time.Second { -t.Errorf("minTTLFromMsg with OPT = %v, want 300s", ttl) -} -} diff --git a/internal/dns/types.go b/internal/dns/types.go index 3a2353d..b5aad48 100644 --- a/internal/dns/types.go +++ b/internal/dns/types.go @@ -5,35 +5,12 @@ import ( ) const ( - TypeA uint16 = dns.TypeA - TypeAAAA uint16 = dns.TypeAAAA - TypeNS uint16 = dns.TypeNS - TypeCNAME uint16 = dns.TypeCNAME - TypeSOA uint16 = dns.TypeSOA - TypeMX uint16 = dns.TypeMX - TypeTXT uint16 = dns.TypeTXT - TypeSRV uint16 = dns.TypeSRV - TypePTR uint16 = dns.TypePTR - TypeANY uint16 = dns.TypeANY + TypeA uint16 = dns.TypeA + TypeAAAA uint16 = dns.TypeAAAA + TypeNS uint16 = dns.TypeNS ) -var QNameTypes = map[uint16]string{ - TypeA: "A", - TypeAAAA: "AAAA", - TypeNS: "NS", - TypeCNAME: "CNAME", - TypeSOA: "SOA", - TypeMX: "MX", - TypeTXT: "TXT", - TypeSRV: "SRV", - TypePTR: "PTR", - TypeANY: "ANY", -} - func QNameType(qtype uint16) string { - if name, ok := QNameTypes[qtype]; ok { - return name - } return dns.TypeToString[qtype] } diff --git a/internal/dns/types_test.go b/internal/dns/types_test.go index bd75ab7..1a81ba7 100644 --- a/internal/dns/types_test.go +++ b/internal/dns/types_test.go @@ -13,13 +13,6 @@ func TestQNameType(t *testing.T) { {"A record", TypeA, "A"}, {"AAAA record", TypeAAAA, "AAAA"}, {"NS record", TypeNS, "NS"}, - {"CNAME record", TypeCNAME, "CNAME"}, - {"SOA record", TypeSOA, "SOA"}, - {"MX record", TypeMX, "MX"}, - {"TXT record", TypeTXT, "TXT"}, - {"SRV record", TypeSRV, "SRV"}, - {"PTR record", TypePTR, "PTR"}, - {"ANY record", TypeANY, "ANY"}, {"unknown type", uint16(9999), ""}, } @@ -33,36 +26,6 @@ func TestQNameType(t *testing.T) { } } -func TestConstantsMatchMiekg(t *testing.T) { - tests := []struct { - name string - local uint16 - }{ - {"TypeA", TypeA}, - {"TypeAAAA", TypeAAAA}, - {"TypeNS", TypeNS}, - {"TypeCNAME", TypeCNAME}, - {"TypeSOA", TypeSOA}, - {"TypeMX", TypeMX}, - {"TypeTXT", TypeTXT}, - {"TypeSRV", TypeSRV}, - {"TypePTR", TypePTR}, - {"TypeANY", TypeANY}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - mapped, ok := QNameTypes[tt.local] - if !ok { - t.Errorf("QNameTypes missing entry for %s (%d)", tt.name, tt.local) - } - if mapped != QNameType(tt.local) { - t.Errorf("QNameType(%d) = %q, QNameTypes[%d] = %q", tt.local, QNameType(tt.local), tt.local, mapped) - } - }) - } -} - func TestDefaultEDNS0UDPSize(t *testing.T) { if got := DefaultEDNS0UDPSize(); got != 2048 { t.Errorf("DefaultEDNS0UDPSize() = %d, want 2048", got) diff --git a/internal/integration/integration_test.go b/internal/integration/integration_test.go index 949a3f4..f3315e7 100644 --- a/internal/integration/integration_test.go +++ b/internal/integration/integration_test.go @@ -1,9 +1,11 @@ // Package integration provides end-to-end tests for ExploreDNS using a mock -// DNS server that allows deterministic, network-independent testing. +// DNS exchange that allows deterministic, network-independent testing of the +// full engine through its exported API. package integration import ( "context" + "math" "net" "testing" "time" @@ -13,523 +15,291 @@ import ( "github.com/miekg/dns" ) -// mockZone represents a simple in-memory DNS zone for testing. -type mockZone struct { - // map[name][qtype] → []RR - records map[string]map[uint16][]dns.RR +// mockNet maps (server IP, qname, qtype) to a canned response, mirroring how +// distinct nameservers answer differently for the same question. +type mockNet struct { + responses map[string]*dns.Msg } -func newMockZone() *mockZone { - return &mockZone{records: make(map[string]map[uint16][]dns.RR)} +func newMockNet() *mockNet { + return &mockNet{responses: make(map[string]*dns.Msg)} } -func (z *mockZone) addA(name, ip string) { - fqdn := dns.Fqdn(name) - if z.records[fqdn] == nil { - z.records[fqdn] = make(map[uint16][]dns.RR) +func key(server, qname string, qtype uint16) string { + return server + "|" + dns.Fqdn(qname) + "|" + dns.TypeToString[qtype] +} + +func (m *mockNet) on(server, qname string, qtype uint16, msg *dns.Msg) { + m.responses[key(server, qname, qtype)] = msg +} + +func (m *mockNet) exchange(_ context.Context, server string, msg *dns.Msg, _ bool) (*dns.Msg, error) { + host := server + if h, _, err := net.SplitHostPort(server); err == nil { + host = h } - z.records[fqdn][dns.TypeA] = append(z.records[fqdn][dns.TypeA], &dns.A{ - Hdr: dns.RR_Header{Name: fqdn, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP(ip), - }) + q := msg.Question[0] + resp, ok := m.responses[key(host, q.Name, q.Qtype)] + if !ok { + // Unknown question: NXDOMAIN, like an authoritative miss. + out := new(dns.Msg) + out.SetReply(msg) + out.Rcode = dns.RcodeNameError + return out, nil + } + out := resp.Copy() + out.SetReply(msg) + out.Answer, out.Ns, out.Extra = resp.Answer, resp.Ns, resp.Extra + out.Rcode = resp.Rcode + return out, nil } -func (z *mockZone) addNS(zone, ns string) { - fqdn := dns.Fqdn(zone) - if z.records[fqdn] == nil { - z.records[fqdn] = make(map[uint16][]dns.RR) +func aRR(name, ip string) dns.RR { + return &dns.A{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip).To4(), } - z.records[fqdn][dns.TypeNS] = append(z.records[fqdn][dns.TypeNS], &dns.NS{ - Hdr: dns.RR_Header{Name: fqdn, Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, - Ns: dns.Fqdn(ns), - }) } -func (z *mockZone) addCNAME(name, target string) { - fqdn := dns.Fqdn(name) - if z.records[fqdn] == nil { - z.records[fqdn] = make(map[uint16][]dns.RR) +func nsRR(zone, target string) dns.RR { + return &dns.NS{ + Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, + Ns: dns.Fqdn(target), } - z.records[fqdn][dns.TypeCNAME] = append(z.records[fqdn][dns.TypeCNAME], &dns.CNAME{ - Hdr: dns.RR_Header{Name: fqdn, Rrtype: dns.TypeCNAME, Class: dns.ClassINET, Ttl: 300}, +} + +func cnameRR(name, target string) dns.RR { + return &dns.CNAME{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeCNAME, Class: dns.ClassINET, Ttl: 300}, Target: dns.Fqdn(target), - }) -} - -// makeExchange creates a mock ExchangeFunc that serves responses from the zone. -// It simulates referral behavior: if a name matches a zone NS record, it returns -// a referral with glue. If it matches an A record, it returns the answer. -func (z *mockZone) makeExchange() dnsinternal.ExchangeFunc { - return func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - if len(msg.Question) == 0 { - return nil, nil - } - q := msg.Question[0] - - resp := new(dns.Msg) - resp.SetReply(msg) - resp.Authoritative = true - - // Direct answer - if rrs, ok := z.records[q.Name]; ok { - if answers, ok := rrs[q.Qtype]; ok { - resp.Answer = append(resp.Answer, answers...) - return resp, nil - } - // CNAME chain — return CNAME + answer for target if qtype != CNAME - if cnameRRs, ok := rrs[dns.TypeCNAME]; ok && q.Qtype != dns.TypeCNAME { - resp.Answer = append(resp.Answer, cnameRRs...) - return resp, nil - } - } - - // Check for zone delegation: look for NS records covering any suffix of qname - labels := dns.SplitDomainName(q.Name) - for i := 0; i < len(labels); i++ { - zone := dns.Fqdn(joinLabels(labels[i:])) - if nsRRs, ok := z.records[zone][dns.TypeNS]; ok && zone != q.Name { - // Return referral - resp.Authoritative = false - resp.Ns = append(resp.Ns, nsRRs...) - for _, ns := range nsRRs { - nsName := ns.(*dns.NS).Ns - if aRRs, ok := z.records[nsName][dns.TypeA]; ok { - resp.Extra = append(resp.Extra, aRRs...) - } - } - return resp, nil - } - } - - // NXDOMAIN - resp.Authoritative = true - resp.Rcode = dns.RcodeNameError - return resp, nil } } -func joinLabels(labels []string) string { - result := "" - for i, l := range labels { - if i > 0 { - result += "." - } - result += l - } - return result +func answerMsg(rrs ...dns.RR) *dns.Msg { + m := new(dns.Msg) + m.Answer = rrs + return m } -// setupTestZone creates a mock zone with a typical referral hierarchy: -// -// root → com (referral) → example.com (referral) → www.example.com (A) -func setupTestZone() *mockZone { - z := newMockZone() - - // Root server glue - z.addA("a.root-servers.test", "198.41.0.4") - - // com TLD referral from root - z.addNS("com", "a.gtld-servers.test") - z.addA("a.gtld-servers.test", "192.5.6.30") - - // example.com NS referral from com TLD - z.addNS("example.com", "ns1.example.com") - z.addA("ns1.example.com", "1.2.3.4") - - // Actual A records - z.addA("example.com", "93.184.216.34") - z.addA("www.example.com", "93.184.216.34") - - return z +func referralMsg(nsRRs []dns.RR, glue ...dns.RR) *dns.Msg { + m := new(dns.Msg) + m.Ns = nsRRs + m.Extra = glue + return m } -// TestIntegrationSimpleAQuery verifies end-to-end traversal with mock DNS -// that returns A record answers without network dependency. -func TestIntegrationSimpleAQuery(t *testing.T) { - z := setupTestZone() - - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, +func newTraverser(maxDepth int) *traverse.Traverser { + return traverse.NewTraverser(&traverse.TraverserConfig{ + MaxDepth: maxDepth, QueryType: dnsinternal.TypeA, RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + QueryConfig: &dnsinternal.QueryConfig{ + Retries: 1, + Timeout: time.Second, + RetryDelay: time.Millisecond, + }, }) - tr.SetExchange(z.makeExchange()) +} +func run(t *testing.T, tr *traverse.Traverser, m *mockNet, qname string) *traverse.Referral { + t.Helper() + tr.SetExchange(m.exchange) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() - - results, err := tr.Traverse(ctx, "example.com") + root, err := tr.Run(ctx, qname) if err != nil { - t.Fatalf("Traverse: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results from traversal") + t.Fatalf("Run(%q): %v", qname, err) } + assertProbabilityInvariant(t, root) + return root +} - var foundAnswer bool - for _, r := range results { - if r.Response != nil && r.Response.Type == traverse.RespAnswer { - foundAnswer = true - if r.Response.Decoded != nil && len(r.Response.Decoded.Answers) > 0 { - for _, rr := range r.Response.Decoded.Answers { - if a, ok := rr.(*dns.A); ok { - t.Logf("Found A record: %v", a.A) - } - } - } - } +// assertProbabilityInvariant checks the engine ground rule: aggregated leaf +// probabilities at the root sum to 1.0. +func assertProbabilityInvariant(t *testing.T, root *traverse.Referral) { + t.Helper() + sum := 0.0 + for _, leaf := range root.StatsList() { + sum += leaf.Prob } - if !foundAnswer { - t.Errorf("expected to find an answer response; got types: %v", responseTypes(results)) + if math.Abs(sum-1.0) > 1e-9 { + t.Errorf("leaf probabilities sum to %v, want 1.0", sum) } } -// TestIntegrationReferralChain verifies multi-hop referral traversal: -// root → com → example.com, with glue records at each step. +func statuses(root *traverse.Referral) map[traverse.Status]float64 { + out := make(map[traverse.Status]float64) + for _, leaf := range root.StatsList() { + out[leaf.Response.Status] += leaf.Prob + } + return out +} + +// TestIntegrationReferralChain verifies the classic delegation walk: +// root → com → example.com with per-server responses. func TestIntegrationReferralChain(t *testing.T) { - z := newMockZone() + m := newMockNet() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.test")}, + aRR("a.gtld-servers.test", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com")}, + aRR("ns1.example.com", "1.2.3.4"), + )) + m.on("1.2.3.4", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "93.184.216.34"))) - // Root delegates to com - z.addNS("com", "a.gtld-servers.test") - z.addA("a.gtld-servers.test", "192.5.6.30") - - // TLD delegates to example.com - z.addNS("example.com", "ns1.example.com") - z.addA("ns1.example.com", "1.2.3.4") - - // Authoritative answer - z.addA("example.com", "93.184.216.34") - - exchange := z.makeExchange() - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(exchange) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("Traverse: %v", err) - } - - // Count referrals and answers - var referrals, answers int - for _, r := range results { - if r.Response == nil { - continue - } - switch r.Response.Type { - case traverse.RespReferral: - referrals++ - case traverse.RespAnswer: - answers++ - } - } - t.Logf("referrals=%d answers=%d total=%d", referrals, answers, len(results)) - if answers == 0 { - t.Errorf("expected at least one answer; types: %v", responseTypes(results)) + root := run(t, newTraverser(10), m, "www.example.com") + got := statuses(root) + if math.Abs(got[traverse.StatusAnswered]-1.0) > 1e-9 { + t.Errorf("statuses = %v, want 100%% answered", got) } } -// TestIntegrationCNAMEResolution verifies that CNAME chains are followed correctly. -func TestIntegrationCNAMEResolution(t *testing.T) { - z := newMockZone() +// TestIntegrationCNAMERestart verifies that an out-of-zone CNAME target +// restarts the traversal from the branch cache. +func TestIntegrationCNAMERestart(t *testing.T) { + m := newMockNet() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com")}, + aRR("ns1.example.com", "1.2.3.4"), + )) + m.on("1.2.3.4", "www.example.com", dns.TypeA, answerMsg(cnameRR("www.example.com", "cdn.example.net"))) + m.on("198.41.0.4", "cdn.example.net", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.net", "ns1.example.net")}, + aRR("ns1.example.net", "5.6.7.8"), + )) + m.on("5.6.7.8", "cdn.example.net", dns.TypeA, answerMsg(aRR("cdn.example.net", "93.184.216.35"))) - // www.example.com → CNAME → example.com → A record - z.addCNAME("www.example.com", "example.com") - z.addA("example.com", "93.184.216.34") - - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(z.makeExchange()) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - results, err := tr.Traverse(ctx, "www.example.com") - if err != nil { - t.Fatalf("Traverse CNAME: %v", err) + root := run(t, newTraverser(10), m, "www.example.com") + got := statuses(root) + if math.Abs(got[traverse.StatusAnswered]-1.0) > 1e-9 { + t.Errorf("statuses = %v, want 100%% answered via restart", got) } - if len(results) == 0 { - t.Fatal("expected results") - } - - var foundCNAME, foundAnswer bool - for _, r := range results { - if r.Response == nil { - continue - } - if r.Response.Type == traverse.RespCNAMEFollow { - foundCNAME = true - } - if r.Response.Type == traverse.RespAnswer { - foundAnswer = true + for _, leaf := range root.StatsList() { + if leaf.Response.Status == traverse.StatusAnswered && leaf.Response.Qname != "cdn.example.net" { + t.Errorf("answered qname = %q, want the CNAME target", leaf.Response.Qname) } } - t.Logf("CNAME traversal: foundCNAME=%v foundAnswer=%v types=%v", foundCNAME, foundAnswer, responseTypes(results)) } -// TestIntegrationNXDOMAIN verifies that NXDOMAIN responses are correctly classified. +// TestIntegrationNXDOMAIN verifies rcode errors surface as error leaves with +// the reference wording. func TestIntegrationNXDOMAIN(t *testing.T) { - z := newMockZone() - // Zone has no records for nonexistent.example.com - - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(z.makeExchange()) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - results, err := tr.Traverse(ctx, "nonexistent.example.test") - if err != nil { - t.Fatalf("Traverse NXDOMAIN: %v", err) + m := newMockNet() + // mock returns NXDOMAIN for anything unmocked + root := run(t, newTraverser(10), m, "nonexistent.example.test") + got := statuses(root) + if math.Abs(got[traverse.StatusError]-1.0) > 1e-9 { + t.Errorf("statuses = %v, want 100%% error", got) } - if len(results) == 0 { - t.Fatal("expected at least one result for NXDOMAIN") - } - - var foundNXDOMAIN bool - for _, r := range results { - if r.Response != nil && r.Response.Type == traverse.RespNXDOMAIN { - foundNXDOMAIN = true - break + for _, leaf := range root.StatsList() { + if leaf.Response.DQ.ErrorMessage != "No such domain (NXDOMAIN)" { + t.Errorf("error message = %q", leaf.Response.DQ.ErrorMessage) } } - if !foundNXDOMAIN { - t.Errorf("expected NXDOMAIN result; got: %v", responseTypes(results)) - } } -// TestIntegrationSERVFAIL verifies that SERVFAIL responses are correctly handled. +// TestIntegrationSERVFAIL verifies SERVFAIL classification. func TestIntegrationSERVFAIL(t *testing.T) { - sfMsg := new(dns.Msg) - sfMsg.Rcode = dns.RcodeServerFailure + m := newMockNet() + sf := new(dns.Msg) + sf.Rcode = dns.RcodeServerFailure + m.on("198.41.0.4", "example.com", dns.TypeA, sf) - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return sfMsg.Copy(), nil - }) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("Traverse SERVFAIL: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least one result") - } - if results[0].Response.Type != traverse.RespSERVFAIL { - t.Errorf("expected SERVFAIL, got %v", results[0].Response.Type) + root := run(t, newTraverser(10), m, "example.com") + for _, leaf := range root.StatsList() { + if leaf.Response.Status != traverse.StatusError { + t.Errorf("status = %q, want error", leaf.Response.Status) + } + if leaf.Response.DQ.ErrorMessage != "Server failure (SERVFAIL)" { + t.Errorf("error message = %q", leaf.Response.DQ.ErrorMessage) + } } } -// TestIntegrationCNAMELoop verifies that CNAME loops are detected and reported. +// TestIntegrationCNAMELoop verifies cross-response CNAME loops terminate as +// cname_loop leaves. func TestIntegrationCNAMELoop(t *testing.T) { - callCount := 0 - // www.a.test → CNAME → www.b.test → CNAME → www.a.test (loop) - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - callCount++ - if len(msg.Question) == 0 { - return nil, nil - } - q := msg.Question[0] + m := newMockNet() + m.on("198.41.0.4", "www.a.test", dns.TypeA, answerMsg(cnameRR("www.a.test", "www.b.test"))) + m.on("198.41.0.4", "www.b.test", dns.TypeA, answerMsg(cnameRR("www.b.test", "www.a.test"))) - resp := new(dns.Msg) - resp.SetReply(msg) - resp.Authoritative = true - - switch q.Name { - case "www.a.test.": - resp.Answer = append(resp.Answer, &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.a.test.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "www.b.test.", - }) - case "www.b.test.": - resp.Answer = append(resp.Answer, &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.b.test.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "www.a.test.", - }) - default: - resp.Rcode = dns.RcodeNameError - } - return resp, nil - }) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - results, err := tr.Traverse(ctx, "www.a.test") - if err != nil { - t.Fatalf("Traverse CNAME loop: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results from CNAME loop traversal") - } - - var foundLoop bool - for _, r := range results { - if r.Response != nil && r.Response.Type == traverse.RespCNAMELoop { - foundLoop = true - break - } - } - if !foundLoop { - t.Logf("types found: %v", responseTypes(results)) - // CNAME loop detection may vary based on implementation; warn rather than fail - t.Logf("CNAME loop not detected as RespCNAMELoop (may be handled differently)") + root := run(t, newTraverser(10), m, "www.a.test") + got := statuses(root) + if math.Abs(got[traverse.StatusCNAMELoop]-1.0) > 1e-9 { + t.Errorf("statuses = %v, want 100%% cname_loop", got) } } -// TestIntegrationMaxDepthExceeded verifies that infinite referral chains are -// cut off at the configured max depth. +// TestIntegrationMaxDepthExceeded verifies that an endless referral chain is +// cut off with a "Maxdepth N exceeded" exception leaf. func TestIntegrationMaxDepthExceeded(t *testing.T) { - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 3, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - // Always return a referral to ns.example.com - resp := new(dns.Msg) - resp.SetReply(msg) - resp.Ns = append(resp.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS, Class: dns.ClassINET}, - Ns: "ns.example.com.", - }) - resp.Extra = append(resp.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET}, - A: net.ParseIP("1.2.3.4"), - }) - return resp, nil - }) + m := newMockNet() + // Each hop delegates one label deeper: the node at depth 3 (refid 1.1.1) + // is never queried because MaxDepth 3 injects the exception first. + m.on("198.41.0.4", "www.d2.d1", dns.TypeA, referralMsg( + []dns.RR{nsRR("d1", "ns.d1")}, + aRR("ns.d1", "10.0.0.1"), + )) + m.on("10.0.0.1", "www.d2.d1", dns.TypeA, referralMsg( + []dns.RR{nsRR("d2.d1", "ns.d2.d1")}, + aRR("ns.d2.d1", "10.0.0.2"), + )) - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() + root := run(t, newTraverser(3), m, "www.d2.d1") - results, err := tr.Traverse(ctx, "deep.example.com") - if err != nil { - t.Fatalf("Traverse: %v", err) + foundMaxdepth := false + for _, leaf := range root.StatsList() { + if leaf.Response.Status == traverse.StatusException && + leaf.Response.DQ.ExceptionMessage == "Maxdepth 3 exceeded" { + foundMaxdepth = true + } } - t.Logf("max depth test: %d results, types: %v", len(results), responseTypes(results)) - if len(results) == 0 { - t.Fatal("expected results even with max depth exceeded") + if !foundMaxdepth { + t.Errorf("expected a Maxdepth 3 exceeded exception leaf, got %v", statuses(root)) } } -// TestIntegrationHooksReceiveEvents verifies that traversal hooks receive -// the expected start and complete events. +// TestIntegrationHooksReceiveEvents verifies start/answer events pair up. func TestIntegrationHooksReceiveEvents(t *testing.T) { - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) + m := newMockNet() + m.on("198.41.0.4", "example.com", dns.TypeA, answerMsg(aRR("example.com", "93.184.216.34"))) - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerMsg.Copy(), nil - }) - - var startEvents, completeEvents int + tr := newTraverser(10) + var startEvents, answerEvents int tr.SetHooks(&traverse.TraverserHooks{ OnEvent: func(e traverse.TraversalEvent) { switch e.Stage { - case traverse.EventStart: + case traverse.StageStart: startEvents++ - case traverse.EventComplete: - completeEvents++ + case traverse.StageAnswer: + answerEvents++ } }, }) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("Traverse: %v", err) - } - _ = results + run(t, tr, m, "example.com") if startEvents == 0 { t.Error("expected at least one start event") } - if completeEvents == 0 { - t.Error("expected at least one complete event") - } - if startEvents != completeEvents { - t.Errorf("start events (%d) != complete events (%d)", startEvents, completeEvents) + if startEvents != answerEvents { + t.Errorf("start events (%d) != answer events (%d)", startEvents, answerEvents) } } // TestIntegrationContextCancellation verifies that the traversal respects // context cancellation and returns an appropriate error. func TestIntegrationContextCancellation(t *testing.T) { - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - // Always return referral to keep loop going - resp := new(dns.Msg) - resp.SetReply(msg) - resp.Ns = append(resp.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS}, - Ns: "ns.example.com.", - }) - resp.Extra = append(resp.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dns.TypeA}, - A: net.ParseIP("1.2.3.4"), - }) - return resp, nil - }) + m := newMockNet() + m.on("198.41.0.4", "example.com", dns.TypeA, answerMsg(aRR("example.com", "93.184.216.34"))) + tr := newTraverser(10) + tr.SetExchange(m.exchange) ctx, cancel := context.WithCancel(context.Background()) - cancel() // Cancel before traversal starts + cancel() // cancel before traversal starts - _, err := tr.Traverse(ctx, "example.com") - if err == nil { + if _, err := tr.Run(ctx, "example.com"); err == nil { t.Fatal("expected error when context is cancelled") } } - -// responseTypes returns a summary of response types for debugging. -func responseTypes(results []traverse.TraversalResult) []string { - var types []string - for _, r := range results { - if r.Response != nil { - types = append(types, r.Response.Type.String()) - } else { - types = append(types, "nil") - } - } - return types -} diff --git a/internal/output/coverage_test.go b/internal/output/coverage_test.go deleted file mode 100644 index 3f7f0d0..0000000 --- a/internal/output/coverage_test.go +++ /dev/null @@ -1,877 +0,0 @@ -package output - -import ( - "bytes" - "context" - "net" - "strings" - "testing" - - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" - miekgdns "github.com/miekg/dns" -) - -// ---- stats.go coverage ---- - -func TestRRDataString(t *testing.T) { - cases := []struct { - rr miekgdns.RR - want string - }{ - { - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("1.2.3.4")}, - "1.2.3.4", - }, - { - &miekgdns.AAAA{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeAAAA}, AAAA: net.ParseIP("::1")}, - "::1", - }, - { - &miekgdns.CNAME{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeCNAME}, Target: "example.com."}, - "example.com.", - }, - { - &miekgdns.NS{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeNS}, Ns: "ns1.example.com."}, - "ns1.example.com.", - }, - { - &miekgdns.MX{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeMX}, Preference: 10, Mx: "mail.example.com."}, - "10 mail.example.com.", - }, - { - &miekgdns.TXT{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeTXT}, Txt: []string{"v=spf1", "include:example.com"}}, - "v=spf1 include:example.com", - }, - } - for _, tc := range cases { - got := rrDataString(tc.rr) - if got != tc.want { - t.Errorf("rrDataString(%T) = %q, want %q", tc.rr, got, tc.want) - } - } -} - -func TestSummaryTypeLabel(t *testing.T) { - cases := []struct { - input string - want string - }{ - {"nodata", "found no such record"}, - {"nxdomain", "name does not exist"}, - {"servfail", "resulted in SERVFAIL"}, - {"refused", "query refused by server"}, - {"notimp", "query type not implemented by server"}, - {"cname_loop", "resulted in a CNAME loop"}, - {"error", "resulted in an error"}, - {"referral", "resulted in a referral"}, - {"unknown_type", "unknown_type"}, - } - for _, tc := range cases { - got := summaryTypeLabel(tc.input) - if got != tc.want { - t.Errorf("summaryTypeLabel(%q) = %q, want %q", tc.input, got, tc.want) - } - } -} - -func TestCollectServers(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 0.5, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - } - results := []traverse.TraversalResult{ - {Referral: ref, Response: resp}, - } - servers := collectServers(results) - if len(servers) == 0 { - t.Fatal("expected at least one server") - } -} - -func TestCollectServersNilServer(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: nil, - Type: traverse.RespAnswer, - } - results := []traverse.TraversalResult{{Referral: ref, Response: resp}} - servers := collectServers(results) - if len(servers) != 0 { - t.Errorf("expected 0 servers with nil server, got %d", len(servers)) - } -} - -func TestCollectServersDedup(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 0.5, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - } - results := []traverse.TraversalResult{ - {Referral: ref, Response: resp}, - {Referral: ref, Response: resp}, - } - servers := collectServers(results) - for _, ips := range servers { - for _, ip := range ips { - count := 0 - for _, i := range ips { - if i == ip { - count++ - } - } - if count > 1 { - t.Errorf("duplicate IP %s in server list", ip) - } - } - } -} - -func TestServerName(t *testing.T) { - t.Run("uses bailiwick", func(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - result := traverse.TraversalResult{Referral: ref, Response: nil} - name := serverName(result) - if name != "com" { - t.Errorf("serverName = %q, want 'com'", name) - } - }) - - t.Run("uses NSName when bailiwick is root", func(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, ".", 0, 1.0, nil) - ref.NSName = "ns1.example.com." - result := traverse.TraversalResult{Referral: ref, Response: nil} - name := serverName(result) - if name != "ns1.example.com." { - t.Errorf("serverName = %q, want 'ns1.example.com.'", name) - } - }) - - t.Run("uses server IP from response", func(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - } - result := traverse.TraversalResult{Referral: ref, Response: resp} - name := serverName(result) - if name != "1.2.3.4" { - t.Errorf("serverName = %q, want '1.2.3.4'", name) - } - }) - - t.Run("unknown fallback", func(t *testing.T) { - result := traverse.TraversalResult{Referral: nil, Response: nil} - name := serverName(result) - if name != "unknown" { - t.Errorf("serverName = %q, want 'unknown'", name) - } - }) -} - -func TestComputeSummaryNonAnswerTypes(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - for _, respType := range []traverse.ResponseType{ - traverse.RespNXDOMAIN, traverse.RespSERVFAIL, traverse.RespNODATA, - } { - resp := &traverse.Response{Referral: ref, Type: respType} - stats := ComputeSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if stats == nil { - t.Errorf("ComputeSummary returned nil for %v", respType) - continue - } - if len(stats.ByType) == 0 { - t.Errorf("expected ByType entry for %v", respType) - } - } -} - -func TestComputeSummaryNilResponse(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - stats := ComputeSummary([]traverse.TraversalResult{{Referral: ref, Response: nil}}) - if stats != nil { - t.Error("expected nil stats for nil response") - } -} - -func TestComputeSummaryNilReferral(t *testing.T) { - resp := &traverse.Response{Type: traverse.RespAnswer} - stats := ComputeSummary([]traverse.TraversalResult{{Referral: nil, Response: resp}}) - if stats != nil { - t.Error("expected nil stats for nil referral") - } -} - -func TestComputeSummaryAnswerKeyEmpty(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - // Answer with only CNAME (no data key) should go into ByType - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.CNAME{ - Hdr: miekgdns.RR_Header{Name: "www.example.com.", Rrtype: miekgdns.TypeCNAME}, - Target: "example.com.", - }, - }, - }, - } - stats := ComputeSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if stats == nil { - t.Fatal("expected non-nil stats") - } -} - -// ---- text.go coverage ---- - -func TestTextFormatterWriteResolve(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, "example.com.", 1, 0.5, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &buf) - - err := f.WriteResolve(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - IsResolve: true, - }) - if err != nil { - t.Fatalf("WriteResolve: %v", err) - } - if buf.Len() == 0 { - t.Error("expected output from WriteResolve") - } -} - -func TestTextFormatterWriteResolveNonStart(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, "example.com.", 1, 0.5, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - f := newTextFormatter(cfg, &buf) - - err := f.WriteResolve(traverse.TraversalEvent{ - Stage: traverse.EventComplete, - Result: traverse.TraversalResult{Referral: ref}, - IsResolve: true, - }) - if err != nil { - t.Fatalf("WriteResolve: %v", err) - } - if buf.Len() != 0 { - t.Error("expected no output for non-start resolve event") - } -} - -func TestTextFormatterWriteServers(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: miekgdns.TypeA}, - A: net.ParseIP("1.2.3.4"), - }, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - cfg.ShowServers = true - cfg.ShowResults = false - cfg.ShowSummaryResults = false - f := newTextFormatter(cfg, &buf) - - err := f.WriteSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if err != nil { - t.Fatalf("WriteSummary: %v", err) - } - if !strings.Contains(buf.String(), "The following servers were encountered:") { - t.Errorf("expected server list header, got %q", buf.String()) - } -} - -func TestTextFormatterWriteResults(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: miekgdns.TypeA}, - A: net.ParseIP("93.184.216.34"), - }, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - cfg.ShowServers = false - cfg.ShowResults = true - cfg.ShowSummaryResults = false - f := newTextFormatter(cfg, &buf) - - err := f.WriteSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if err != nil { - t.Fatalf("WriteSummary: %v", err) - } - if !strings.Contains(buf.String(), "Results:") { - t.Errorf("expected 'Results:' header, got %q", buf.String()) - } -} - -func TestTextFormatterFormatResultLineAllTypes(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - - cases := []struct { - respType traverse.ResponseType - contains string - }{ - {traverse.RespNODATA, "no such record"}, - {traverse.RespNXDOMAIN, "does not exist"}, - {traverse.RespSERVFAIL, "SERVFAIL"}, - {traverse.RespREFUSED, "refused"}, - {traverse.RespNOTIMPL, "not implemented"}, - {traverse.RespCNAMELoop, "CNAME loop"}, - {traverse.RespError, "error"}, - } - - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &bytes.Buffer{}) - - for _, tc := range cases { - resp := &traverse.Response{ - Referral: ref, - Type: tc.respType, - } - result := traverse.TraversalResult{Referral: ref, Response: resp} - line := f.formatResultLine(result) - if !strings.Contains(strings.ToLower(line), strings.ToLower(tc.contains)) { - t.Errorf("formatResultLine(%v) = %q, want substring %q", tc.respType, line, tc.contains) - } - } -} - -func TestTextFormatterFormatResultLineErrorWithMessage(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespError, - ErrorMessage: "custom error message", - } - - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &bytes.Buffer{}) - - line := f.formatResultLine(traverse.TraversalResult{Referral: ref, Response: resp}) - if !strings.Contains(line, "custom error message") { - t.Errorf("expected custom error message, got %q", line) - } -} - -func TestTextFormatterFormatResultLineCNAMELoopWithMessage(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespCNAMELoop, - ErrorMessage: "CNAME loop detected: example.com", - } - - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &bytes.Buffer{}) - - line := f.formatResultLine(traverse.TraversalResult{Referral: ref, Response: resp}) - if !strings.Contains(line, "CNAME loop detected") { - t.Errorf("expected CNAME loop message, got %q", line) - } -} - -func TestTextFormatterFormatResultLineAnswerMultiple(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("1.1.1.1")}, - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("2.2.2.2")}, - }, - }, - } - - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &bytes.Buffer{}) - - line := f.formatResultLine(traverse.TraversalResult{Referral: ref, Response: resp}) - if !strings.Contains(line, "/") { - t.Errorf("expected '/' separator for multiple answers, got %q", line) - } -} - -func TestTextFormatterColorize(t *testing.T) { - cfg := DefaultConfig() - cfg.Color = true - f := newTextFormatter(cfg, &bytes.Buffer{}) - - colored := f.colorize("hello", colorGreen) - if !strings.Contains(colored, "\033[") { - t.Error("expected ANSI color code in colored output") - } - - cfg.Color = false - f2 := newTextFormatter(cfg, &bytes.Buffer{}) - plain := f2.colorize("hello", colorGreen) - if plain != "hello" { - t.Errorf("expected plain text without color, got %q", plain) - } -} - -func TestTextFormatterColorizeEmpty(t *testing.T) { - cfg := DefaultConfig() - cfg.Color = true - f := newTextFormatter(cfg, &bytes.Buffer{}) - out := f.colorize("hello", "") - if out != "hello" { - t.Errorf("empty color should return plain text, got %q", out) - } -} - -func TestTextFormatterVerboseProgress(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 0.5, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - cfg.Verbose = true - f := newTextFormatter(cfg, &buf) - - err := f.WriteProgress(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - }) - if err != nil { - t.Fatalf("WriteProgress: %v", err) - } - if buf.Len() == 0 { - t.Error("expected output with verbose mode") - } - out := buf.String() - if !strings.Contains(out, "com") { - t.Errorf("expected bailiwick in verbose output, got %q", out) - } -} - -func TestTextFormatterProgressResolving(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, "example.com.", 1, 0.5, nil) - // no addresses = resolving - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &buf) - - err := f.WriteProgress(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - }) - if err != nil { - t.Fatalf("WriteProgress: %v", err) - } - if !strings.Contains(buf.String(), "resolving") { - t.Errorf("expected 'resolving' in output, got %q", buf.String()) - } -} - -func TestTextFormatterWriteServersWithVersions(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("1.2.3.4")}, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - cfg.ShowServers = true - cfg.ShowResults = false - cfg.ShowSummaryResults = false - cfg.ShowVersions = true - cfg.Fingerprints = map[string]string{"1.2.3.4": "BIND 9.16"} - f := newTextFormatter(cfg, &buf) - - err := f.WriteSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if err != nil { - t.Fatalf("WriteSummary: %v", err) - } - if !strings.Contains(buf.String(), "BIND 9.16") { - t.Errorf("expected version string in server output, got %q", buf.String()) - } -} - -func TestTextFormatterWriteResult(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("1.2.3.4")}, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &buf) - - err := f.WriteResult(traverse.TraversalResult{Referral: ref, Response: resp}) - if err != nil { - t.Fatalf("WriteResult: %v", err) - } - if buf.Len() == 0 { - t.Error("expected output from WriteResult") - } -} - -func TestTextFormatterWriteResultNilRefs(t *testing.T) { - var buf bytes.Buffer - cfg := DefaultConfig() - f := newTextFormatter(cfg, &buf) - - err := f.WriteResult(traverse.TraversalResult{Referral: nil, Response: nil}) - if err != nil { - t.Fatalf("WriteResult: %v", err) - } - if buf.Len() != 0 { - t.Error("expected no output for nil referral/response") - } -} - -func TestReferralServerLabelVariants(t *testing.T) { - t.Run("with addresses", func(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - ref.Addresses = []net.IP{net.ParseIP("1.2.3.4")} - label := referralServerLabel(ref, nil) - if !strings.Contains(label, "1.2.3.4") { - t.Errorf("expected IP in label, got %q", label) - } - }) - - t.Run("with NSName", func(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - ref.NSName = "ns1.example.com." - label := referralServerLabel(ref, nil) - if label != "ns1.example.com." { - t.Errorf("expected NSName, got %q", label) - } - }) - - t.Run("with bailiwick", func(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - label := referralServerLabel(ref, nil) - if label != "com" { - t.Errorf("expected trimmed bailiwick, got %q", label) - } - }) - - t.Run("unknown", func(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - label := referralServerLabel(ref, nil) - if label != "unknown" { - t.Errorf("expected 'unknown', got %q", label) - } - }) - - t.Run("with response server", func(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{Server: net.ParseIP("5.6.7.8")} - label := referralServerLabel(ref, resp) - if label != "5.6.7.8" { - t.Errorf("expected server IP, got %q", label) - } - }) -} - -// ---- json.go coverage ---- - -func TestJSONFormatterWriteResolve(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, "example.com.", 1, 0.5, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.ShowResolves = true - f := newJSONFormatter(cfg, &buf) - - err := f.WriteResolve(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - IsResolve: true, - }) - if err != nil { - t.Fatalf("WriteResolve: %v", err) - } - if len(f.payload.Resolves) != 1 { - t.Errorf("expected 1 resolve entry, got %d", len(f.payload.Resolves)) - } -} - -func TestJSONFormatterWriteResolveShowResolvesFalse(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, "example.com.", 1, 0.5, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.ShowResolves = false - f := newJSONFormatter(cfg, &buf) - - err := f.WriteResolve(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - IsResolve: true, - }) - if err != nil { - t.Fatalf("WriteResolve: %v", err) - } - if len(f.payload.Resolves) != 0 { - t.Errorf("expected 0 resolve entries when ShowResolves=false, got %d", len(f.payload.Resolves)) - } -} - -func TestJSONFormatterWriteResult(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("1.2.3.4")}, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.ShowAllStats = true - f := newJSONFormatter(cfg, &buf) - - err := f.WriteResult(traverse.TraversalResult{Referral: ref, Response: resp}) - if err != nil { - t.Fatalf("WriteResult: %v", err) - } - if len(f.payload.Results) != 1 { - t.Errorf("expected 1 result entry, got %d", len(f.payload.Results)) - } -} - -func TestJSONFormatterWriteResultShowAllStatsFalse(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{Referral: ref, Type: traverse.RespAnswer} - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.ShowAllStats = false - f := newJSONFormatter(cfg, &buf) - - err := f.WriteResult(traverse.TraversalResult{Referral: ref, Response: resp}) - if err != nil { - t.Fatalf("WriteResult: %v", err) - } - if len(f.payload.Results) != 0 { - t.Errorf("expected 0 result entries when ShowAllStats=false, got %d", len(f.payload.Results)) - } -} - -func TestJSONFormatterWriteSummaryWithServersAndVersions(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("1.2.3.4")}, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.ShowServers = true - cfg.ShowVersions = true - cfg.Fingerprints = map[string]string{"1.2.3.4": "BIND 9.16"} - f := newJSONFormatter(cfg, &buf) - - err := f.WriteSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if err != nil { - t.Fatalf("WriteSummary: %v", err) - } - found := false - for _, srv := range f.payload.Servers { - if srv.Version == "BIND 9.16" { - found = true - } - } - if !found { - t.Error("expected version in server list") - } -} - -func TestJSONFormatterStageName(t *testing.T) { - if stageName(traverse.EventStart) != "start" { - t.Errorf("expected 'start', got %q", stageName(traverse.EventStart)) - } - if stageName(traverse.EventComplete) != "complete" { - t.Errorf("expected 'complete', got %q", stageName(traverse.EventComplete)) - } - if stageName(traverse.EventStage(99)) != "unknown" { - t.Errorf("expected 'unknown' for unknown stage") - } -} - -func TestJSONFormatterEventToJSONNilReferral(t *testing.T) { - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - f := newJSONFormatter(cfg, &buf) - - item := f.eventToJSON(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: nil}, - }) - if item.Name != "" { - t.Errorf("expected empty name for nil referral, got %q", item.Name) - } -} - -func TestJSONFormatterEventToJSONWithResponse(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - f := newJSONFormatter(cfg, &buf) - - item := f.eventToJSON(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref, Response: resp}, - }) - if item.Server != "1.2.3.4" { - t.Errorf("expected server IP, got %q", item.Server) - } -} - -func TestJSONFormatterWriteProgressShowProgressFalse(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.ShowProgress = false - f := newJSONFormatter(cfg, &buf) - - err := f.WriteProgress(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - }) - if err != nil { - t.Fatalf("WriteProgress: %v", err) - } - if len(f.payload.Progress) != 0 { - t.Errorf("expected 0 progress entries when ShowProgress=false, got %d", len(f.payload.Progress)) - } -} - -// ---- runner.go coverage ---- - -func TestRunTraversalNilTraverser(t *testing.T) { - _, err := RunTraversal(context.Background(), nil, nil, nil, "example.com") - if err == nil { - t.Fatal("expected error for nil traverser") - } -} - -func TestNewFormatterNilConfig(t *testing.T) { - f := NewFormatter(nil, &bytes.Buffer{}) - if f == nil { - t.Fatal("NewFormatter(nil) should not return nil") - } -} - -func TestAttachHooksNilCfg(t *testing.T) { - h := AttachHooks(nil, nil) - if h != nil { - t.Fatal("AttachHooks(nil, nil) should return nil") - } -} - -func TestAttachHooksDebugMode(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Debug = 1 - cfg.ShowResolves = true - - // formatter that returns an error on WriteResolve - formatter := &errorFormatter{} - hooks := AttachHooks(cfg, formatter) - if hooks == nil { - t.Fatal("expected non-nil hooks") - } - - // Call OnEvent with IsResolve=true - should call WriteResolve and log error to stderr (debug>0) - hooks.OnEvent(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - IsResolve: true, - }) - _ = buf.String() // no assertion - just ensure it doesn't panic -} - -// errorFormatter is a mock formatter for testing error paths. -type errorFormatter struct{} - -func (f *errorFormatter) WriteProgress(_ traverse.TraversalEvent) error { return nil } -func (f *errorFormatter) WriteResolve(_ traverse.TraversalEvent) error { return nil } -func (f *errorFormatter) WriteResult(_ traverse.TraversalResult) error { return nil } -func (f *errorFormatter) WriteSummary(_ []traverse.TraversalResult) error { - return nil -} -func (f *errorFormatter) Flush() error { return nil } diff --git a/internal/output/formatter.go b/internal/output/formatter.go index 1fc1cd5..c7df01c 100644 --- a/internal/output/formatter.go +++ b/internal/output/formatter.go @@ -1,10 +1,11 @@ -// Package output renders ExploreDNS traversal results for human consumption or -// machine processing. +// Package output renders ExploreDNS traversal results for human consumption +// or machine processing. // // Two formats are supported: // -// - FormatText — a coloured hierarchical tree (default) -// - FormatJSON — a JSON array of traversal results +// - FormatText — dnstraverse-style header/progress/results/summary text +// (default) +// - FormatJSON — a single JSON document with the aggregated results // // Create a Formatter via NewFormatter and call RunTraversal to drive the // traversal engine and stream output incrementally. @@ -26,9 +27,19 @@ const ( ) type Config struct { - Format Format - Domain string - QueryType string + Format Format + Domain string + QueryType string + + // Engine settings echoed in the header block (bin/dnstraverse). + Fast bool + AllRootServers bool + UDPSize int + Retries int + MaxDepth int + AllowTCP bool + AlwaysTCP bool + ShowProgress bool ShowResolves bool ShowServers bool @@ -49,22 +60,50 @@ type Config struct { func DefaultConfig() *Config { return &Config{ Format: FormatText, + Fast: true, + UDPSize: 2048, + Retries: 2, + MaxDepth: 20, + AllowTCP: true, ShowProgress: true, - ShowResolves: true, - ShowServers: true, + ShowResolves: false, + ShowServers: false, ShowVersions: true, - ShowAllStats: true, + ShowAllStats: false, ShowResults: true, ShowSummaryResults: true, - Color: os.Getenv("NO_COLOR") == "", + Color: ColorEnabled(os.Stdout), } } +// ColorEnabled reports whether colour output should be used for w: only when +// NO_COLOR is unset and w is a terminal. +func ColorEnabled(w io.Writer) bool { + if os.Getenv("NO_COLOR") != "" { + return false + } + f, ok := w.(*os.File) + if !ok { + return false + } + info, err := f.Stat() + if err != nil { + return false + } + return info.Mode()&os.ModeCharDevice != 0 +} + +// Formatter renders traversal progress and the aggregated results. The header +// is written once before the run, progress arrives via the traverser hooks, +// and the aggregated leaves and the servers seen arrive once the run +// completed. type Formatter interface { + // WriteHeader renders the pre-run header block from the discovered roots + // (suppressed entirely by --quiet in text mode). + WriteHeader(roots []traverse.StartServer) error WriteProgress(event traverse.TraversalEvent) error WriteResolve(event traverse.TraversalEvent) error - WriteResult(result traverse.TraversalResult) error - WriteSummary(results []traverse.TraversalResult) error + WriteSummary(root *traverse.Referral, servers map[string][]string) error Flush() error } @@ -85,6 +124,11 @@ func AttachHooks(cfg *Config, formatter Formatter) *traverse.TraverserHooks { if cfg == nil || formatter == nil { return nil } + if !cfg.ShowProgress { + // Ruby registers no progress callbacks at all without show-progress; + // resolve display additionally requires show-resolves. + return nil + } logErr := func(context string, err error) { if err != nil && cfg.Debug > 0 { fmt.Fprintf(os.Stderr, "Debug: formatter %s: %v\n", context, err) @@ -92,15 +136,13 @@ func AttachHooks(cfg *Config, formatter Formatter) *traverse.TraverserHooks { } return &traverse.TraverserHooks{ OnEvent: func(event traverse.TraversalEvent) { - switch { - case event.IsResolve && cfg.ShowResolves: - logErr("WriteResolve", formatter.WriteResolve(event)) - case !event.IsResolve && cfg.ShowProgress: - logErr("WriteProgress", formatter.WriteProgress(event)) - } - if event.Stage == traverse.EventComplete && cfg.ShowAllStats { - logErr("WriteResult", formatter.WriteResult(event.Result)) + if event.IsResolve { + if cfg.ShowResolves { + logErr("WriteResolve", formatter.WriteResolve(event)) + } + return } + logErr("WriteProgress", formatter.WriteProgress(event)) }, } } diff --git a/internal/output/formatter_test.go b/internal/output/formatter_test.go index 34fcdc7..c4393d6 100644 --- a/internal/output/formatter_test.go +++ b/internal/output/formatter_test.go @@ -7,369 +7,420 @@ import ( "net" "strings" "testing" + "time" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" + idns "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" - miekgdns "github.com/miekg/dns" + "github.com/miekg/dns" ) +// mockDelegation wires root → com → example.com (2 NS, one glueless answer +// path) through the single injected exchange. +func mockDelegation() idns.ExchangeFunc { + responses := map[string]*dns.Msg{} + set := func(server, qname string, msg *dns.Msg) { + responses[server+"/"+dns.Fqdn(qname)] = msg + } + a := func(name, ip string) dns.RR { + return &dns.A{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip).To4(), + } + } + ns := func(zone, target string) dns.RR { + return &dns.NS{ + Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, + Ns: dns.Fqdn(target), + } + } + + rootMsg := new(dns.Msg) + rootMsg.Ns = []dns.RR{ns("com", "a.gtld-servers.net")} + rootMsg.Extra = []dns.RR{a("a.gtld-servers.net", "192.5.6.30")} + set("198.41.0.4", "www.example.com", rootMsg) + + comMsg := new(dns.Msg) + comMsg.Ns = []dns.RR{ns("example.com", "ns1.example.com"), ns("example.com", "ns2.example.com")} + comMsg.Extra = []dns.RR{a("ns1.example.com", "1.1.1.1"), a("ns2.example.com", "2.2.2.2")} + set("192.5.6.30", "www.example.com", comMsg) + + answer := new(dns.Msg) + answer.Answer = []dns.RR{a("www.example.com", "9.9.9.9")} + set("1.1.1.1", "www.example.com", answer) + set("2.2.2.2", "www.example.com", answer) + + return func(_ context.Context, server string, msg *dns.Msg, _ bool) (*dns.Msg, error) { + host := server + if h, _, err := net.SplitHostPort(server); err == nil { + host = h + } + resp, ok := responses[host+"/"+msg.Question[0].Name] + if !ok { + return nil, &net.DNSError{Err: "no mock", Name: msg.Question[0].Name} + } + out := resp.Copy() + out.SetReply(msg) + out.Answer, out.Ns, out.Extra = resp.Answer, resp.Ns, resp.Extra + return out, nil + } +} + +func newMockTraverser() *traverse.Traverser { + tr := traverse.NewTraverser(&traverse.TraverserConfig{ + MaxDepth: traverse.DefaultMaxDepth, + QueryType: dns.TypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + QueryConfig: &idns.QueryConfig{ + Retries: 1, + Timeout: time.Second, + RetryDelay: time.Millisecond, + }, + }) + tr.SetExchange(mockDelegation()) + return tr +} + func TestDefaultConfig(t *testing.T) { cfg := DefaultConfig() - if !cfg.ShowProgress { - t.Fatal("expected ShowProgress default true") - } - if cfg.Format != FormatText { - t.Fatalf("Format = %v, want text", cfg.Format) + if cfg.Format != FormatText || !cfg.ShowProgress || !cfg.ShowResults { + t.Errorf("unexpected defaults: %+v", cfg) } } func TestNewFormatterSelectsImplementation(t *testing.T) { - text := NewFormatter(DefaultConfig(), &bytes.Buffer{}) - if _, ok := text.(*textFormatter); !ok { - t.Fatalf("expected text formatter, got %T", text) + if _, ok := NewFormatter(&Config{Format: FormatJSON}, &bytes.Buffer{}).(*jsonFormatter); !ok { + t.Error("FormatJSON should select the JSON formatter") } - - jsonCfg := DefaultConfig() - jsonCfg.Format = FormatJSON - jsonFmt := NewFormatter(jsonCfg, &bytes.Buffer{}) - if _, ok := jsonFmt.(*jsonFormatter); !ok { - t.Fatalf("expected json formatter, got %T", jsonFmt) + if _, ok := NewFormatter(&Config{Format: FormatText}, &bytes.Buffer{}).(*textFormatter); !ok { + t.Error("FormatText should select the text formatter") + } + if NewFormatter(nil, &bytes.Buffer{}) == nil { + t.Error("nil config should still produce a formatter") } } -func TestComputeSummaryAggregatesAnswers(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, - A: net.ParseIP("93.184.216.34"), - }, - }, - }, - } - - stats := ComputeSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if len(stats.Answers) != 1 { - t.Fatalf("answers = %d, want 1", len(stats.Answers)) - } - if stats.Answers[0].Prob != 1.0 { - t.Fatalf("prob = %v, want 1.0", stats.Answers[0].Prob) - } -} - -func TestTextFormatterSummaryOutput(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, - A: net.ParseIP("93.184.216.34"), - }, - }, - }, - } - +func TestRunTraversalTextOutput(t *testing.T) { var buf bytes.Buffer cfg := DefaultConfig() + cfg.Domain = "www.example.com" + cfg.QueryType = "a" + cfg.ShowServers = true + cfg.ShowVersions = false // no fingerprint network calls in tests cfg.Color = false - cfg.ShowServers = false - cfg.ShowResults = false formatter := NewFormatter(cfg, &buf) - if err := formatter.WriteSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}); err != nil { - t.Fatalf("WriteSummary: %v", err) - } - - out := buf.String() - if !strings.Contains(out, "Summary:") { - t.Fatalf("expected summary header, got %q", out) - } - if !strings.Contains(out, "100%") { - t.Fatalf("expected probability in summary, got %q", out) - } - if !strings.Contains(out, "93.184.216.34") { - t.Fatalf("expected answer IP in summary, got %q", out) - } -} - -func TestJSONFormatterProducesValidOutput(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("198.41.0.4"), - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, - A: net.ParseIP("93.184.216.34"), - }, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.Domain = "example.com" - cfg.QueryType = "A" - cfg.ShowServers = false - cfg.ShowResults = true - cfg.ShowSummaryResults = true - formatter := NewFormatter(cfg, &buf) - - if err := formatter.WriteProgress(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - }); err != nil { - t.Fatalf("WriteProgress: %v", err) - } - if err := formatter.WriteSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}); err != nil { - t.Fatalf("WriteSummary: %v", err) - } - if err := formatter.Flush(); err != nil { - t.Fatalf("Flush: %v", err) - } - - var payload map[string]any - if err := json.Unmarshal(buf.Bytes(), &payload); err != nil { - t.Fatalf("invalid json: %v\n%s", err, buf.String()) - } - if payload["domain"] != "example.com" { - t.Fatalf("domain = %v", payload["domain"]) - } - if _, ok := payload["summary"]; !ok { - t.Fatalf("expected summary in json output") - } -} - -func TestRunTraversalUsesHooks(t *testing.T) { - answerResp := func() *miekgdns.Msg { - m := new(miekgdns.Msg) - m.SetReply(new(miekgdns.Msg)) - m.Answer = append(m.Answer, &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - return m - }() - - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 5, - QueryType: dns.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *miekgdns.Msg, useTCP bool) (*miekgdns.Msg, error) { - return answerResp.Copy(), nil - }) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - cfg.ShowServers = false - cfg.ShowResults = false - cfg.ShowSummaryResults = true - formatter := NewFormatter(cfg, &buf) - - _, err := RunTraversal(context.Background(), tr, cfg, formatter, "example.com") + root, err := RunTraversal(context.Background(), newMockTraverser(), cfg, formatter, "www.example.com") if err != nil { t.Fatalf("RunTraversal: %v", err) } - if !strings.Contains(buf.String(), "Summary:") { - t.Fatalf("expected formatted summary output, got %q", buf.String()) + if root == nil || len(root.Stats) == 0 { + t.Fatal("expected aggregated stats on the root referral") + } + + out := buf.String() + for _, want := range []string{ + "# Using fast mode", + "Using 198.41.0.4 (198.41.0.4) as initial root", + "Running query www.example.com type a", + "1 198.41.0.4 (198.41.0.4)", + "1.1 a.gtld-servers.net (192.5.6.30)", + "1.1.1 ns1.example.com (1.1.1.1)", + "Results:", + " 50.0%: Answer from ns1.example.com (1.1.1.1)", + " 50.0%: Answer from ns2.example.com (2.2.2.2)", + "Summary Results:", + " 100% answered with www.example.com. 300 IN A 9.9.9.9", + "The following servers were encountered:", + } { + if !strings.Contains(out, want) { + t.Errorf("output missing %q\n---\n%s", want, out) + } } } -func TestJSONFormatterWriteResolveAndResult(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -server := net.ParseIP("198.41.0.4") -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: traverse.RespAnswer, -Decoded: &dns.DecodedResponse{ -Answers: []miekgdns.RR{ -&miekgdns.A{ -Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, -A: net.ParseIP("1.2.3.4"), -}, -}, -}, +func TestVerboseProgressFormat(t *testing.T) { + var buf bytes.Buffer + cfg := DefaultConfig() + cfg.Domain = "www.example.com" + cfg.QueryType = "a" + cfg.ShowVersions = false + cfg.Verbose = true + cfg.Color = false + formatter := NewFormatter(cfg, &buf) + + if _, err := RunTraversal(context.Background(), newMockTraverser(), cfg, formatter, "www.example.com"); err != nil { + t.Fatalf("RunTraversal: %v", err) + } + out := buf.String() + // Verbose rows are " [qname] ( ) "; the + // root bailiwick renders as "<>". + for _, want := range []string{ + "1 [www.example.com] 198.41.0.4 (198.41.0.4) <>", + "1.1 [www.example.com] a.gtld-servers.net (192.5.6.30) ", + "1.1.1 [www.example.com] ns1.example.com (1.1.1.1) ", + } { + if !strings.Contains(out, want) { + t.Errorf("output missing %q\n---\n%s", want, out) + } + } } -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Format = FormatJSON -cfg.Domain = "example.com" -cfg.QueryType = "A" -cfg.ShowResolves = true -cfg.ShowAllStats = true -cfg.ShowProgress = true -f := NewFormatter(cfg, &buf).(*jsonFormatter) +func TestRunTraversalQuietSuppressesHeader(t *testing.T) { + var buf bytes.Buffer + cfg := DefaultConfig() + cfg.Domain = "www.example.com" + cfg.QueryType = "a" + cfg.ShowVersions = false + cfg.Quiet = true + cfg.Color = false + formatter := NewFormatter(cfg, &buf) -// WriteResolve -if err := f.WriteResolve(traverse.TraversalEvent{ -Stage: traverse.EventStart, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}); err != nil { -t.Fatalf("WriteResolve: %v", err) + if _, err := RunTraversal(context.Background(), newMockTraverser(), cfg, formatter, "www.example.com"); err != nil { + t.Fatalf("RunTraversal: %v", err) + } + out := buf.String() + for _, banned := range []string{"# Using fast mode", "as initial root", "Running query"} { + if strings.Contains(out, banned) { + t.Errorf("quiet output must not contain %q\n---\n%s", banned, out) + } + } + if !strings.Contains(out, "Results:") { + t.Errorf("quiet must still print results\n---\n%s", out) + } } -// WriteResult -if err := f.WriteResult(traverse.TraversalResult{Referral: ref, Response: resp}); err != nil { -t.Fatalf("WriteResult: %v", err) +func TestRunTraversalJSONOutput(t *testing.T) { + var buf bytes.Buffer + cfg := DefaultConfig() + cfg.Format = FormatJSON + cfg.Domain = "www.example.com" + cfg.QueryType = "A" + cfg.ShowVersions = false + formatter := NewFormatter(cfg, &buf) + + if _, err := RunTraversal(context.Background(), newMockTraverser(), cfg, formatter, "www.example.com"); err != nil { + t.Fatalf("RunTraversal: %v", err) + } + + var doc map[string]any + if err := json.Unmarshal(buf.Bytes(), &doc); err != nil { + t.Fatalf("invalid JSON: %v\n%s", err, buf.String()) + } + if doc["domain"] != "www.example.com" { + t.Errorf("domain = %v", doc["domain"]) + } + if doc["qtype"] != "A" { + t.Errorf("qtype = %v", doc["qtype"]) + } + root, ok := doc["root"].(map[string]any) + if !ok || root["ip"] != "198.41.0.4" { + t.Errorf("root = %v", doc["root"]) + } + for _, banned := range []string{"progress", "resolves"} { + if _, present := doc[banned]; present { + t.Errorf("JSON document must not contain %q", banned) + } + } + results, ok := doc["results"].([]any) + if !ok || len(results) != 2 { + t.Fatalf("results = %v", doc["results"]) + } + summary, ok := doc["summary"].(map[string]any) + if !ok { + t.Fatalf("summary missing: %v", doc) + } + byStatus := summary["by_status"].(map[string]any) + if prob := byStatus["answered"].(float64); prob < 0.999 || prob > 1.001 { + t.Errorf("answered summary prob = %v", prob) + } } -// WriteProgress with EventComplete to cover stageName "complete" -if err := f.WriteProgress(traverse.TraversalEvent{ -Stage: traverse.EventComplete, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}); err != nil { -t.Fatalf("WriteProgress EventComplete: %v", err) +// mockGluelessDelegation wires root → com → example.com where the single NS +// (ns1.example.net) comes without glue, forcing a resolve subtree that walks +// root → net → answer. +func mockGluelessDelegation() idns.ExchangeFunc { + responses := map[string]*dns.Msg{} + set := func(server, qname string, msg *dns.Msg) { + responses[server+"/"+dns.Fqdn(qname)] = msg + } + a := func(name, ip string) dns.RR { + return &dns.A{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip).To4(), + } + } + ns := func(zone, target string) dns.RR { + return &dns.NS{ + Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, + Ns: dns.Fqdn(target), + } + } + + comRef := new(dns.Msg) + comRef.Ns = []dns.RR{ns("com", "a.gtld-servers.net")} + comRef.Extra = []dns.RR{a("a.gtld-servers.net", "192.5.6.30")} + set("198.41.0.4", "www.example.com", comRef) + + glueless := new(dns.Msg) + glueless.Ns = []dns.RR{ns("example.com", "ns1.example.net")} + set("192.5.6.30", "www.example.com", glueless) + + netRef := new(dns.Msg) + netRef.Ns = []dns.RR{ns("net", "b.gtld-servers.net")} + netRef.Extra = []dns.RR{a("b.gtld-servers.net", "192.33.14.31")} + set("198.41.0.4", "ns1.example.net", netRef) + + nsAnswer := new(dns.Msg) + nsAnswer.Answer = []dns.RR{a("ns1.example.net", "3.3.3.3")} + set("192.33.14.31", "ns1.example.net", nsAnswer) + + answer := new(dns.Msg) + answer.Answer = []dns.RR{a("www.example.com", "9.9.9.9")} + set("3.3.3.3", "www.example.com", answer) + + return func(_ context.Context, server string, msg *dns.Msg, _ bool) (*dns.Msg, error) { + host := server + if h, _, err := net.SplitHostPort(server); err == nil { + host = h + } + resp, ok := responses[host+"/"+msg.Question[0].Name] + if !ok { + return nil, &net.DNSError{Err: "no mock", Name: msg.Question[0].Name} + } + out := resp.Copy() + out.SetReply(msg) + out.Answer, out.Ns, out.Extra = resp.Answer, resp.Ns, resp.Extra + return out, nil + } } -if err := f.Flush(); err != nil { -t.Fatalf("Flush: %v", err) -} +func runGlueless(t *testing.T, cfg *Config) string { + t.Helper() + tr := traverse.NewTraverser(&traverse.TraverserConfig{ + MaxDepth: traverse.DefaultMaxDepth, + QueryType: dns.TypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + QueryConfig: &idns.QueryConfig{ + Retries: 1, + Timeout: time.Second, + RetryDelay: time.Millisecond, + }, + }) + tr.SetExchange(mockGluelessDelegation()) + + var buf bytes.Buffer + formatter := NewFormatter(cfg, &buf) + if _, err := RunTraversal(context.Background(), tr, cfg, formatter, "www.example.com"); err != nil { + t.Fatalf("RunTraversal: %v", err) + } + return buf.String() } -func TestJSONFormatterWriteResolveFlagOff(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) +func TestResolveProgressHiddenByDefault(t *testing.T) { + cfg := DefaultConfig() + cfg.Domain = "www.example.com" + cfg.QueryType = "a" + cfg.ShowVersions = false + cfg.Color = false + out := runGlueless(t, cfg) -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Format = FormatJSON -cfg.ShowResolves = false -cfg.ShowAllStats = false -f := NewFormatter(cfg, &buf).(*jsonFormatter) - -if err := f.WriteResolve(traverse.TraversalEvent{ -Stage: traverse.EventStart, -Result: traverse.TraversalResult{Referral: ref}, -}); err != nil { -t.Fatalf("WriteResolve: %v", err) -} -if err := f.WriteResult(traverse.TraversalResult{Referral: ref}); err != nil { -t.Fatalf("WriteResult: %v", err) -} + for _, want := range []string{ + "1.1.1 ns1.example.net -- resolving", + "1.1.1 ns1.example.net (3.3.3.3)", + "100.0%: Answer from ns1.example.net (3.3.3.3)", + } { + if !strings.Contains(out, want) { + t.Errorf("output missing %q\n---\n%s", want, out) + } + } + // Resolve subtree nodes (".0." refids) render only under --show-resolves, + // and resolve outcomes never appear as separate Results entries. + for _, line := range strings.Split(out, "\n") { + if strings.HasPrefix(line, "1.1.1.0") { + t.Errorf("resolve subtree must be hidden by default: %q", line) + } + } + if strings.Contains(out, "ns1.example.net./IN/A") { + t.Errorf("resolve leaves must not pollute Results\n---\n%s", out) + } } -func TestJSONFormatterWriteSummaryWithServers(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) -server := net.ParseIP("1.2.3.4") -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: traverse.RespAnswer, -Decoded: &dns.DecodedResponse{ -Answers: []miekgdns.RR{ -&miekgdns.A{ -Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, -A: net.ParseIP("1.2.3.4"), -}, -}, -}, -} -results := []traverse.TraversalResult{{Referral: ref, Response: resp}} +func TestResolveProgressShownWithShowResolves(t *testing.T) { + cfg := DefaultConfig() + cfg.Domain = "www.example.com" + cfg.QueryType = "a" + cfg.ShowVersions = false + cfg.ShowResolves = true + cfg.Color = false + out := runGlueless(t, cfg) -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Format = FormatJSON -cfg.Domain = "example.com" -cfg.QueryType = "A" -cfg.ShowServers = true -cfg.ShowVersions = false -cfg.ShowResults = true -cfg.ShowSummaryResults = true -f := NewFormatter(cfg, &buf).(*jsonFormatter) - -if err := f.WriteSummary(results); err != nil { -t.Fatalf("WriteSummary: %v", err) -} -if err := f.Flush(); err != nil { -t.Fatalf("Flush: %v", err) + for _, want := range []string{ + "1.1.1 ns1.example.net -- resolving", + "1.1.1.0.1 198.41.0.4 (198.41.0.4)", + "1.1.1.0.1.1 b.gtld-servers.net (192.33.14.31)", + "1.1.1 ns1.example.net (3.3.3.3)", + } { + if !strings.Contains(out, want) { + t.Errorf("output missing %q\n---\n%s", want, out) + } + } } -var payload map[string]any -if err := json.Unmarshal(buf.Bytes(), &payload); err != nil { -t.Fatalf("invalid JSON: %v\n%s", err, buf.String()) -} -if _, ok := payload["servers"]; !ok { -t.Error("expected 'servers' field in JSON output") -} +func TestRunTraversalRequiresTraverser(t *testing.T) { + if _, err := RunTraversal(context.Background(), nil, DefaultConfig(), nil, "example.com"); err == nil { + t.Fatal("expected error for nil traverser") + } } -func TestNewFormatterNilWriter(t *testing.T) { -// Should not panic with nil writer -cfg := DefaultConfig() -f := NewFormatter(cfg, nil) -if f == nil { -t.Error("NewFormatter should not return nil") -} +func TestFormatProbability(t *testing.T) { + tests := []struct { + prob float64 + want string + }{ + {1.0, " 100%"}, + {0.5, " 50%"}, + {0.933, "93.3%"}, + {0.067, " 6.7%"}, + {1.0 / 3, "33.3%"}, + } + for _, tt := range tests { + if got := formatProbability(tt.prob); got != tt.want { + t.Errorf("formatProbability(%v) = %q, want %q", tt.prob, got, tt.want) + } + } } -func TestAttachHooksShowResolves(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -server := net.ParseIP("1.2.3.4") -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: traverse.RespAnswer, +func TestSummaryStatusLabels(t *testing.T) { + tests := map[traverse.Status]string{ + traverse.StatusNoData: "found no such record", + traverse.StatusReferralLame: "resulted in a lame referral", + traverse.StatusException: "resulted in an exception", + traverse.StatusError: "resulted in an error", + traverse.StatusNoGlue: "found no glue", + traverse.StatusLoop: "resulted in a loop", + traverse.StatusCNAMELoop: "resulted in a CNAME loop", + traverse.Status("odd"): "odd", + } + for status, want := range tests { + if got := summaryStatusLabel(status); got != want { + t.Errorf("summaryStatusLabel(%q) = %q, want %q", status, got, want) + } + } } -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.ShowProgress = false -cfg.ShowResolves = true -cfg.ShowAllStats = true -cfg.Color = false -formatter := NewFormatter(cfg, &buf) -hooks := AttachHooks(cfg, formatter) - -// Trigger a resolve event -hooks.OnEvent(traverse.TraversalEvent{ -Stage: traverse.EventStart, -IsResolve: true, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}) - -if buf.Len() == 0 { -t.Error("expected resolve output when ShowResolves is true") -} +func TestCollectUniqueServerIPs(t *testing.T) { + servers := map[string][]string{ + "ns1.example.com": {"1.1.1.1", "2.2.2.2"}, + "ns2.example.com": {"1.1.1.1"}, + } + ips := collectUniqueServerIPs(servers) + if len(ips) != 2 { + t.Errorf("unique IPs = %v", ips) + } } -func TestAttachHooksShowAllStats(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -server := net.ParseIP("1.2.3.4") -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: traverse.RespAnswer, -Decoded: &dns.DecodedResponse{ -Answers: []miekgdns.RR{ -&miekgdns.A{Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, A: net.ParseIP("1.2.3.4")}, -}, -}, -} - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.ShowProgress = false -cfg.ShowResolves = false -cfg.ShowAllStats = true -cfg.Color = false -formatter := NewFormatter(cfg, &buf) -hooks := AttachHooks(cfg, formatter) - -hooks.OnEvent(traverse.TraversalEvent{ -Stage: traverse.EventComplete, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}) +func TestReverseString(t *testing.T) { + if got := reverseString("abc"); got != "cba" { + t.Errorf("reverseString = %q", got) + } } diff --git a/internal/output/golden_test.go b/internal/output/golden_test.go new file mode 100644 index 0000000..8675fc4 --- /dev/null +++ b/internal/output/golden_test.go @@ -0,0 +1,182 @@ +package output + +import ( + "bytes" + "context" + "net" + "strings" + "testing" + "time" + + idns "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" + "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" + "github.com/miekg/dns" +) + +// TestTextOutputMatchesReferenceCapture rebuilds the topology of +// docs/captures/dnstraverse-ruby-www.example.com-A.txt (root → com → the two +// cloudflare NS, three IPs each) through the mock exchange and asserts the +// complete text output byte-for-byte against the reference format. It differs +// from the capture only in volatile values: the root chosen, the number of +// gTLD servers, and the server fingerprints (versions are disabled so no +// network is touched). +func TestTextOutputMatchesReferenceCapture(t *testing.T) { + responses := map[string]*dns.Msg{} + set := func(server, qname string, msg *dns.Msg) { + responses[server+"/"+dns.Fqdn(qname)] = msg + } + a := func(name, ip string) dns.RR { + return &dns.A{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip).To4(), + } + } + ns := func(zone, target string) dns.RR { + return &dns.NS{ + Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, + Ns: dns.Fqdn(target), + } + } + + // Upstream resolver: ". NS" returns one root with glue (root discovery). + rootNS := new(dns.Msg) + rootNS.Answer = []dns.RR{ns(".", "m.root-servers.net")} + rootNS.Extra = []dns.RR{a("m.root-servers.net", "202.12.27.33")} + set("10.0.0.53", ".", rootNS) + + // Root referral to com (three gTLD servers, all glued). + comRef := new(dns.Msg) + comRef.Ns = []dns.RR{ + ns("com", "a.gtld-servers.net"), + ns("com", "b.gtld-servers.net"), + ns("com", "c.gtld-servers.net"), + } + comRef.Extra = []dns.RR{ + a("a.gtld-servers.net", "192.5.6.30"), + a("b.gtld-servers.net", "192.33.14.30"), + a("c.gtld-servers.net", "192.26.92.30"), + } + set("202.12.27.33", "www.example.com", comRef) + + // gTLD referral to example.com: two NS, three glue addresses each. + heraIPs := []string{"108.162.192.162", "172.64.32.162", "173.245.58.162"} + elliottIPs := []string{"108.162.195.228", "162.159.44.228", "172.64.35.228"} + exampleRef := new(dns.Msg) + exampleRef.Ns = []dns.RR{ + ns("example.com", "hera.ns.cloudflare.com"), + ns("example.com", "elliott.ns.cloudflare.com"), + } + for _, ip := range heraIPs { + exampleRef.Extra = append(exampleRef.Extra, a("hera.ns.cloudflare.com", ip)) + } + for _, ip := range elliottIPs { + exampleRef.Extra = append(exampleRef.Extra, a("elliott.ns.cloudflare.com", ip)) + } + for _, ip := range []string{"192.5.6.30", "192.33.14.30", "192.26.92.30"} { + set(ip, "www.example.com", exampleRef) + } + + answer := new(dns.Msg) + answer.Answer = []dns.RR{ + a("www.example.com", "104.20.23.154"), + a("www.example.com", "172.66.147.243"), + } + for _, ip := range append(append([]string{}, heraIPs...), elliottIPs...) { + set(ip, "www.example.com", answer) + } + + exchange := func(_ context.Context, server string, msg *dns.Msg, _ bool) (*dns.Msg, error) { + host := server + if h, _, err := net.SplitHostPort(server); err == nil { + host = h + } + resp, ok := responses[host+"/"+msg.Question[0].Name] + if !ok { + return nil, &net.DNSError{Err: "no mock", Name: msg.Question[0].Name} + } + out := resp.Copy() + out.SetReply(msg) + out.Answer, out.Ns, out.Extra = resp.Answer, resp.Ns, resp.Extra + return out, nil + } + + tr := traverse.NewTraverser(&traverse.TraverserConfig{ + MaxDepth: traverse.DefaultMaxDepth, + QueryType: dns.TypeA, + Fast: true, + RootConfig: &idns.RootDiscoveryConfig{Resolver: "10.0.0.53:53"}, + QueryConfig: &idns.QueryConfig{ + Retries: 1, + Timeout: time.Second, + RetryDelay: time.Millisecond, + }, + }) + tr.SetExchange(exchange) + + cfg := DefaultConfig() + cfg.Domain = "www.example.com" + cfg.QueryType = "a" + cfg.ShowServers = true + cfg.ShowVersions = false + cfg.Color = false + + var buf bytes.Buffer + formatter := NewFormatter(cfg, &buf) + if _, err := RunTraversal(context.Background(), tr, cfg, formatter, "www.example.com"); err != nil { + t.Fatalf("RunTraversal: %v", err) + } + + answerBlock := func(server, ip string) string { + return " 16.7%: Answer from " + server + " (" + ip + ")\n" + + " www.example.com.\t300\tIN\tA\t104.20.23.154\n" + + " www.example.com.\t300\tIN\tA\t172.66.147.243\n" + } + + want := strings.Join([]string{ + "# Using fast mode", + "# Limiting traverse to one root", + "# UDP size 2048 (EDNS0 is on)", + "# Retries 2, max depth 20", + "# Allow TCP is true, always TCP is false", + "Using m.root-servers.net (202.12.27.33) as initial root", + "Running query www.example.com type a", + "1 m.root-servers.net (202.12.27.33)", + "1.1 a.gtld-servers.net (192.5.6.30)", + "1.1.1 hera.ns.cloudflare.com (108.162.192.162,172.64.32.162,173.245.58.162)", + "1.1.2 elliott.ns.cloudflare.com (108.162.195.228,162.159.44.228,172.64.35.228)", + "1.2 b.gtld-servers.net (192.33.14.30)", + "1.2.1 hera.ns.cloudflare.com (108.162.192.162,172.64.32.162,173.245.58.162) -- completed earlier (1.1.1)", + "1.2.2 elliott.ns.cloudflare.com (108.162.195.228,162.159.44.228,172.64.35.228) -- completed earlier (1.1.2)", + "1.3 c.gtld-servers.net (192.26.92.30)", + "1.3.1 hera.ns.cloudflare.com (108.162.192.162,172.64.32.162,173.245.58.162) -- completed earlier (1.1.1)", + "1.3.2 elliott.ns.cloudflare.com (108.162.195.228,162.159.44.228,172.64.35.228) -- completed earlier (1.1.2)", + "", + "The following servers were encountered:", + " hera.ns.cloudflare.com: 108.162.192.162", + " hera.ns.cloudflare.com: 172.64.32.162", + " hera.ns.cloudflare.com: 173.245.58.162", + "elliott.ns.cloudflare.com: 108.162.195.228", + "elliott.ns.cloudflare.com: 162.159.44.228", + "elliott.ns.cloudflare.com: 172.64.35.228", + " a.gtld-servers.net: 192.5.6.30", + " b.gtld-servers.net: 192.33.14.30", + " c.gtld-servers.net: 192.26.92.30", + " m.root-servers.net: 202.12.27.33", + "", + "Results:", + answerBlock("hera.ns.cloudflare.com", "108.162.192.162"), + answerBlock("elliott.ns.cloudflare.com", "108.162.195.228"), + answerBlock("elliott.ns.cloudflare.com", "162.159.44.228"), + answerBlock("hera.ns.cloudflare.com", "172.64.32.162"), + answerBlock("elliott.ns.cloudflare.com", "172.64.35.228"), + answerBlock("hera.ns.cloudflare.com", "173.245.58.162") + + "\n" + + "Summary Results:\n" + + " 100% answered with www.example.com. 300 IN A 104.20.23.154\n" + + " www.example.com. 300 IN A 172.66.147.243\n", + }, "\n") + + if got := buf.String(); got != want { + t.Errorf("output does not match the reference capture format\n--- got ---\n%s\n--- want ---\n%s", got, want) + } +} diff --git a/internal/output/json.go b/internal/output/json.go index d1a85d1..0306f52 100644 --- a/internal/output/json.go +++ b/internal/output/json.go @@ -3,11 +3,15 @@ package output import ( "encoding/json" "io" + "sort" + "strings" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" ) +// jsonFormatter accumulates the run into a single document emitted once by +// Flush: {domain, qtype, root, results, summary, servers}. Aggregated leaves +// appear exactly once (in results); progress events are not recorded. type jsonFormatter struct { cfg *Config w io.Writer @@ -15,32 +19,31 @@ type jsonFormatter struct { } type jsonDocument struct { - Domain string `json:"domain"` - QueryType string `json:"query_type"` - Progress []jsonProgressEvent `json:"progress,omitempty"` - Resolves []jsonProgressEvent `json:"resolves,omitempty"` - Results []jsonResult `json:"results,omitempty"` - Servers []jsonServer `json:"servers,omitempty"` - Summary jsonSummary `json:"summary,omitempty"` + Domain string `json:"domain"` + QueryType string `json:"qtype"` + Root *jsonRoot `json:"root,omitempty"` + Results []jsonResult `json:"results,omitempty"` + Summary *jsonSummary `json:"summary,omitempty"` + Servers []jsonServer `json:"servers,omitempty"` } -type jsonProgressEvent struct { - Stage string `json:"stage"` - Depth int `json:"depth"` - Name string `json:"name"` - QType string `json:"qtype"` - Server string `json:"server,omitempty"` - Bailiwick string `json:"bailiwick,omitempty"` - Resolving bool `json:"resolving,omitempty"` +// jsonRoot is the initial root the traversal started from. +type jsonRoot struct { + Name string `json:"name"` + IP string `json:"ip,omitempty"` } +// jsonResult is one aggregated leaf outcome, emitted exactly once. type jsonResult struct { - Depth int `json:"depth"` - Probability float64 `json:"probability"` - ResponseType string `json:"response_type"` - Server string `json:"server,omitempty"` - Answers []string `json:"answers,omitempty"` - CNAMEChain []string `json:"cname_chain,omitempty"` + RefID string `json:"refid,omitempty"` + Probability float64 `json:"probability"` + Status string `json:"status"` + Server string `json:"server,omitempty"` + IP string `json:"ip,omitempty"` + Qname string `json:"qname,omitempty"` + Qtype string `json:"qtype,omitempty"` + Answers []string `json:"answers,omitempty"` + Message string `json:"message,omitempty"` } type jsonServer struct { @@ -50,12 +53,11 @@ type jsonServer struct { } type jsonSummary struct { - ByType map[string]float64 `json:"by_type,omitempty"` - Answers []jsonAnswerStat `json:"answers,omitempty"` + ByStatus map[string]float64 `json:"by_status,omitempty"` + Answers []jsonAnswerStat `json:"answers,omitempty"` } type jsonAnswerStat struct { - RData string `json:"rdata"` Probability float64 `json:"probability"` Records []string `json:"records,omitempty"` } @@ -71,43 +73,48 @@ func newJSONFormatter(cfg *Config, w io.Writer) *jsonFormatter { } } -func (f *jsonFormatter) WriteProgress(event traverse.TraversalEvent) error { - if !f.cfg.ShowProgress { +func (f *jsonFormatter) WriteHeader(roots []traverse.StartServer) error { + if len(roots) == 0 { return nil } - f.payload.Progress = append(f.payload.Progress, f.eventToJSON(event)) + root := &jsonRoot{Name: roots[0].Name} + if len(roots[0].IPs) > 0 { + root.IP = roots[0].IPs[0] + } + f.payload.Root = root return nil } -func (f *jsonFormatter) WriteResolve(event traverse.TraversalEvent) error { - if !f.cfg.ShowResolves { - return nil - } - f.payload.Resolves = append(f.payload.Resolves, f.eventToJSON(event)) +// WriteProgress is a no-op: the JSON document contains only the aggregated +// outcome, never per-event duplicates. +func (f *jsonFormatter) WriteProgress(traverse.TraversalEvent) error { return nil } -func (f *jsonFormatter) WriteResult(result traverse.TraversalResult) error { - if !f.cfg.ShowAllStats { - return nil - } - f.payload.Results = append(f.payload.Results, f.resultToJSON(result)) +func (f *jsonFormatter) WriteResolve(traverse.TraversalEvent) error { return nil } -func (f *jsonFormatter) WriteSummary(results []traverse.TraversalResult) error { - if f.cfg.ShowResults { - for _, result := range terminalResults(results) { - f.payload.Results = append(f.payload.Results, f.resultToJSON(result)) +func (f *jsonFormatter) WriteSummary(root *traverse.Referral, servers map[string][]string) error { + if root != nil && f.cfg.ShowResults { + // StatsList is sorted by stats key: deterministic ordering. + for _, leaf := range root.StatsList() { + f.payload.Results = append(f.payload.Results, leafToJSON(leaf)) } } if f.cfg.ShowServers { - servers := collectServers(results) - for name, ips := range servers { - srv := jsonServer{Name: name, IPs: ips} + names := make([]string, 0, len(servers)) + for name := range servers { + names = append(names, name) + } + sort.Slice(names, func(i, j int) bool { + return reverseString(strings.ToLower(names[i])) < reverseString(strings.ToLower(names[j])) + }) + for _, name := range names { + srv := jsonServer{Name: name, IPs: servers[name]} if f.cfg.ShowVersions && f.cfg.Fingerprints != nil { - for _, ip := range ips { + for _, ip := range srv.IPs { if v := f.cfg.Fingerprints[ip]; v != "" { srv.Version = v break @@ -119,18 +126,19 @@ func (f *jsonFormatter) WriteSummary(results []traverse.TraversalResult) error { } if f.cfg.ShowSummaryResults { - stats := ComputeSummary(results) - if stats != nil { - f.payload.Summary = jsonSummary{ - ByType: stats.ByType, + if stats := root.SummaryStats(); stats != nil { + summary := &jsonSummary{ByStatus: make(map[string]float64)} + for status, prob := range stats.ByStatus { + summary.ByStatus[string(status)] = prob } for _, answer := range stats.Answers { - f.payload.Summary.Answers = append(f.payload.Summary.Answers, jsonAnswerStat{ - RData: answer.RData, - Probability: answer.Prob, - Records: answer.RRs, - }) + stat := jsonAnswerStat{Probability: answer.Prob} + for _, rr := range answer.RRs { + stat.Records = append(stat.Records, collapseWhitespace(rr.String())) + } + summary.Answers = append(summary.Answers, stat) } + f.payload.Summary = summary } } @@ -143,56 +151,29 @@ func (f *jsonFormatter) Flush() error { return enc.Encode(f.payload) } -func (f *jsonFormatter) eventToJSON(event traverse.TraversalEvent) jsonProgressEvent { - ref := event.Result.Referral - if ref == nil { - return jsonProgressEvent{} +func leafToJSON(leaf *traverse.StatsEntry) jsonResult { + resp := leaf.Response + item := jsonResult{ + Probability: leaf.Prob, + Status: string(resp.Status), + IP: resp.IP, + Qname: resp.Qname, + Qtype: traverse.TypeToString(resp.Qtype), } - - item := jsonProgressEvent{ - Stage: stageName(event.Stage), - Depth: ref.Depth, - Name: trimDomain(ref.Name), - QType: dns.QNameType(ref.Qtype), - Bailiwick: trimDomain(ref.Bailiwick), - Resolving: !ref.HasAddresses(), + if leaf.Referral != nil { + item.RefID = leaf.Referral.RefID + item.Server = leaf.Referral.Server } - if event.Result.Response != nil && event.Result.Response.Server != nil { - item.Server = event.Result.Response.Server.String() - } else { - item.Server = referralServerLabel(ref, event.Result.Response) - } - return item -} - -func (f *jsonFormatter) resultToJSON(result traverse.TraversalResult) jsonResult { - item := jsonResult{} - if result.Referral != nil { - item.Depth = result.Referral.Depth - item.Probability = result.Referral.Prob - } - if result.Response != nil { - item.ResponseType = result.Response.Type.String() - if result.Response.Server != nil { - item.Server = result.Response.Server.String() + if resp.DQ != nil { + for _, rr := range resp.DQ.Answers { + item.Answers = append(item.Answers, collapseWhitespace(rr.String())) } - if result.Response.Decoded != nil { - for _, rr := range result.Response.Decoded.Answers { - item.Answers = append(item.Answers, dns.FormatRecord(rr)) - } - item.CNAMEChain = append(item.CNAMEChain, result.Response.Decoded.CNAMEChain...) + switch resp.Status { + case traverse.StatusError: + item.Message = resp.DQ.ErrorMessage + case traverse.StatusException: + item.Message = resp.DQ.ExceptionMessage } } return item } - -func stageName(stage traverse.EventStage) string { - switch stage { - case traverse.EventStart: - return "start" - case traverse.EventComplete: - return "complete" - default: - return "unknown" - } -} diff --git a/internal/output/runner.go b/internal/output/runner.go index f69b119..eba67b4 100644 --- a/internal/output/runner.go +++ b/internal/output/runner.go @@ -9,7 +9,10 @@ import ( "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" ) -func RunTraversal(ctx context.Context, traverser *traverse.Traverser, cfg *Config, formatter Formatter, domain string) ([]traverse.TraversalResult, error) { +// RunTraversal drives one traversal and streams its output through the +// formatter. It returns the synthetic root referral whose Stats aggregate +// every leaf outcome. +func RunTraversal(ctx context.Context, traverser *traverse.Traverser, cfg *Config, formatter Formatter, domain string) (*traverse.Referral, error) { if traverser == nil { return nil, fmt.Errorf("traverser is required") } @@ -22,40 +25,53 @@ func RunTraversal(ctx context.Context, traverser *traverse.Traverser, cfg *Confi traverser.SetHooks(AttachHooks(cfg, formatter)) - results, err := traverser.Traverse(ctx, domain) + // Discover the roots up front so the header can report the initial root; + // Run reuses the memoised discovery. + roots, err := traverser.Roots(ctx) if err != nil { - return results, err + return nil, err } + if err := formatter.WriteHeader(roots); err != nil { + return nil, err + } + + root, err := traverser.Run(ctx, domain) + if err != nil { + return root, err + } + + servers := traverser.ServersEncountered() // Fingerprint servers when both ShowVersions and ShowServers are enabled. // Gating on ShowServers avoids unnecessary network calls when versions // would not be displayed anyway. if cfg.ShowVersions && cfg.ShowServers { - cfg.Fingerprints = fingerprint.New().FingerprintAll(ctx, collectUniqueServerIPs(results)) + cfg.Fingerprints = fingerprint.New().FingerprintAll(ctx, collectUniqueServerIPs(servers)) } - if err := formatter.WriteSummary(results); err != nil { - return results, err + if err := formatter.WriteSummary(root, servers); err != nil { + return root, err } if err := formatter.Flush(); err != nil { - return results, err + return root, err } - return results, nil + return root, nil } -// collectUniqueServerIPs returns the set of unique server IPs seen in results. -func collectUniqueServerIPs(results []traverse.TraversalResult) []net.IP { +// collectUniqueServerIPs returns the set of unique server IPs encountered. +func collectUniqueServerIPs(servers map[string][]string) []net.IP { seen := make(map[string]bool) var ips []net.IP - for _, r := range results { - if r.Response == nil || r.Response.Server == nil { - continue - } - key := r.Response.Server.String() - if !seen[key] { - seen[key] = true - ips = append(ips, r.Response.Server) + for _, addrs := range servers { + for _, addr := range addrs { + if seen[addr] { + continue + } + seen[addr] = true + if ip := net.ParseIP(addr); ip != nil { + ips = append(ips, ip) + } } } return ips diff --git a/internal/output/stats.go b/internal/output/stats.go index 45b7e8b..f20a4fb 100644 --- a/internal/output/stats.go +++ b/internal/output/stats.go @@ -2,248 +2,44 @@ package output import ( "fmt" - "sort" "strings" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" - miekgdns "github.com/miekg/dns" ) -type summaryEntry struct { - Type string - Prob float64 -} - -type answerEntry struct { - RData string - Prob float64 - RRs []string -} - -type SummaryStats struct { - ByType map[string]float64 - Answers []answerEntry -} - -func ComputeSummary(results []traverse.TraversalResult) *SummaryStats { - stats := &SummaryStats{ - ByType: make(map[string]float64), - } - - for _, result := range results { - if result.Response == nil || !result.Response.IsTerminal() { - continue - } - if result.Referral == nil { - continue - } - - prob := result.Referral.Prob - respType := result.Response.Type.String() - - switch result.Response.Type { - case traverse.RespAnswer: - key, rrs := answerKey(result.Response) - if key == "" { - stats.ByType[respType] += prob - continue - } - found := false - for i := range stats.Answers { - if stats.Answers[i].RData == key { - stats.Answers[i].Prob += prob - found = true - break - } - } - if !found { - stats.Answers = append(stats.Answers, answerEntry{ - RData: key, - Prob: prob, - RRs: rrs, - }) - } - default: - stats.ByType[respType] += prob - } - } - - sort.Slice(stats.Answers, func(i, j int) bool { - return stats.Answers[i].RData < stats.Answers[j].RData - }) - - if len(stats.Answers) == 0 && len(stats.ByType) == 0 { - return nil - } - return stats -} - -func answerKey(resp *traverse.Response) (string, []string) { - if resp == nil || resp.Decoded == nil { - return "", nil - } - - var rdatas []string - var formatted []string - for _, rr := range resp.Decoded.Answers { - if _, ok := rr.(*miekgdns.CNAME); ok { - continue - } - rdata := rrDataString(rr) - if rdata == "" { - continue - } - rdatas = append(rdatas, rdata) - formatted = append(formatted, dns.FormatRecord(rr)) - } - - if len(rdatas) == 0 { - return "", nil - } - sort.Strings(rdatas) - return strings.Join(rdatas, " / "), formatted -} - -func rrDataString(rr miekgdns.RR) string { - switch v := rr.(type) { - case *miekgdns.A: - return v.A.String() - case *miekgdns.AAAA: - return v.AAAA.String() - case *miekgdns.CNAME: - return v.Target - case *miekgdns.NS: - return v.Ns - case *miekgdns.MX: - return fmt.Sprintf("%d %s", v.Preference, v.Mx) - case *miekgdns.TXT: - return strings.Join(v.Txt, " ") - default: - return rr.String() - } -} - +// formatProbability renders txt_prob (summary_stats.rb): %5.1f%% with a +// trailing ".0" trimmed, right-justified to width 5. func formatProbability(prob float64) string { text := fmt.Sprintf("%.1f%%", prob*100) text = strings.Replace(text, ".0%", "%", 1) return fmt.Sprintf("%5s", text) } -func summaryTypeLabel(respType string) string { - switch respType { - case "nodata": +// summaryStatusLabel returns the Summary Results wording per status +// (summary_stats.rb text; answered is handled separately). +func summaryStatusLabel(status traverse.Status) string { + switch status { + case traverse.StatusNoData: return "found no such record" - case "nxdomain": - return "name does not exist" - case "servfail": - return "resulted in SERVFAIL" - case "refused": - return "query refused by server" - case "notimp": - return "query type not implemented by server" - case "cname_loop": - return "resulted in a CNAME loop" - case "ns_error": - return "nameserver lookup failed" - case "error": + case traverse.StatusReferralLame: + return "resulted in a lame referral" + case traverse.StatusException: + return "resulted in an exception" + case traverse.StatusError: return "resulted in an error" - case "referral": - return "resulted in a referral" + case traverse.StatusNoGlue: + return "found no glue" + case traverse.StatusLoop: + return "resulted in a loop" + case traverse.StatusCNAMELoop: + return "resulted in a CNAME loop" default: - return respType + return string(status) } } -func trimDomain(name string) string { - return strings.TrimSuffix(name, ".") -} - -func collectServers(results []traverse.TraversalResult) map[string][]string { - servers := make(map[string][]string) - for _, result := range results { - if result.Response == nil || result.Response.Server == nil { - continue - } - name := serverName(result) - ip := result.Response.Server.String() - if containsString(servers[name], ip) { - continue - } - servers[name] = append(servers[name], ip) - } - return servers -} - -func serverName(result traverse.TraversalResult) string { - if result.Referral != nil && result.Referral.Bailiwick != "" && result.Referral.Bailiwick != "." { - return trimDomain(result.Referral.Bailiwick) - } - if result.Referral != nil && result.Referral.NSName != "" { - return result.Referral.NSName - } - if result.Response != nil && result.Response.Server != nil { - return result.Response.Server.String() - } - return "unknown" -} - -func containsString(items []string, target string) bool { - for _, item := range items { - if item == target { - return true - } - } - return false -} - -// DeduplicateResults collapses terminal results that represent the same -// outcome from the same server into a single entry with summed probability. -// This prevents the same nameserver failure (or answer) from appearing once -// per delegation path when several parent servers all refer to the same child. -func DeduplicateResults(results []traverse.TraversalResult) []traverse.TraversalResult { - type entry struct { - result traverse.TraversalResult - prob float64 - } - keys := make(map[string]*entry) - var order []string - - for _, r := range results { - if r.Response == nil || r.Referral == nil { - continue - } - key := resultDeduplicationKey(r) - if e, ok := keys[key]; ok { - e.prob += r.Referral.Prob - } else { - keys[key] = &entry{result: r, prob: r.Referral.Prob} - order = append(order, key) - } - } - - deduped := make([]traverse.TraversalResult, 0, len(order)) - for _, key := range order { - e := keys[key] - refCopy := *e.result.Referral - refCopy.Prob = e.prob - deduped = append(deduped, traverse.TraversalResult{ - Referral: &refCopy, - Response: e.result.Response, - }) - } - return deduped -} - -func resultDeduplicationKey(r traverse.TraversalResult) string { - bailiwick := strings.TrimSuffix(r.Referral.Bailiwick, ".") - switch r.Response.Type { - case traverse.RespAnswer: - key, _ := answerKey(r.Response) - return "answer:" + bailiwick + ":" + key - case traverse.RespNSResolutionFailed: - return "ns_error:" + r.Response.ErrorMessage - default: - return r.Response.Type.String() + ":" + bailiwick + ":" + r.Response.ErrorMessage - } +// collapseWhitespace renders an RR on one line with runs of whitespace +// collapsed to single spaces (summary_stats.rb text). +func collapseWhitespace(s string) string { + return strings.Join(strings.Fields(s), " ") } diff --git a/internal/output/stats_test.go b/internal/output/stats_test.go deleted file mode 100644 index 01c21ae..0000000 --- a/internal/output/stats_test.go +++ /dev/null @@ -1,327 +0,0 @@ -package output - -import ( - "net" - "testing" - - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" - miekgdns "github.com/miekg/dns" -) - -func makeAnswerResult(name string, ip string, prob float64) traverse.TraversalResult { - ref := traverse.NewReferral(name, dns.TypeA, ".", 0, prob, nil) - server := net.ParseIP("198.41.0.4") - resp := &traverse.Response{ - Referral: ref, - Server: server, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: name + ".", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, - A: net.ParseIP(ip), - }, - }, - }, - } - return traverse.TraversalResult{Referral: ref, Response: resp} -} - -func TestRRDataStringAllTypes(t *testing.T) { - tests := []struct { - rr miekgdns.RR - want string - }{ - { - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("1.2.3.4")}, - "1.2.3.4", - }, - { - &miekgdns.AAAA{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeAAAA}, AAAA: net.ParseIP("::1")}, - "::1", - }, - { - &miekgdns.CNAME{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeCNAME}, Target: "example.com."}, - "example.com.", - }, - { - &miekgdns.NS{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeNS}, Ns: "ns1.example.com."}, - "ns1.example.com.", - }, - { - &miekgdns.MX{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeMX}, Preference: 10, Mx: "mail.example.com."}, - "10 mail.example.com.", - }, - { - &miekgdns.TXT{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeTXT}, Txt: []string{"v=spf1", "include:example.com"}}, - "v=spf1 include:example.com", - }, - } - - for _, tc := range tests { - got := rrDataString(tc.rr) - if got != tc.want { - t.Errorf("rrDataString(%T) = %q, want %q", tc.rr, got, tc.want) - } - } -} - -func TestRRDataStringDefault(t *testing.T) { - // SOA record hits the default case - rr := &miekgdns.SOA{ - Hdr: miekgdns.RR_Header{Name: ".", Rrtype: miekgdns.TypeSOA, Class: miekgdns.ClassINET}, - Ns: "a.root-servers.net.", - Mbox: "nstld.verisign-grs.com.", - } - got := rrDataString(rr) - if got == "" { - t.Error("rrDataString(SOA) should return non-empty string via default case") - } -} - -func TestSummaryTypeLabelAllTypes(t *testing.T) { - cases := map[string]string{ - "nodata": "found no such record", - "nxdomain": "name does not exist", - "servfail": "resulted in SERVFAIL", - "refused": "query refused by server", - "notimp": "query type not implemented by server", - "cname_loop": "resulted in a CNAME loop", - "ns_error": "nameserver lookup failed", - "error": "resulted in an error", - "referral": "resulted in a referral", - "unknown_type": "unknown_type", - } - for input, want := range cases { - got := summaryTypeLabel(input) - if got != want { - t.Errorf("summaryTypeLabel(%q) = %q, want %q", input, got, want) - } - } -} - -func TestCollectServersEmpty(t *testing.T) { - servers := collectServers(nil) - if len(servers) != 0 { - t.Errorf("collectServers(nil) = %v, want empty", servers) - } -} - -func TestCollectServersDeduplication(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - server := net.ParseIP("1.2.3.4") - resp := &traverse.Response{ - Referral: ref, - Server: server, - Type: traverse.RespAnswer, - } - result := traverse.TraversalResult{Referral: ref, Response: resp} - - servers := collectServers([]traverse.TraversalResult{result, result}) - name := "com" - ips := servers[name] - if len(ips) != 1 { - t.Errorf("expected deduplication: got %d IPs, want 1", len(ips)) - } -} - -func TestCollectServersWithBailiwick(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - server := net.ParseIP("1.2.3.4") - resp := &traverse.Response{ - Referral: ref, - Server: server, - Type: traverse.RespAnswer, - } - result := traverse.TraversalResult{Referral: ref, Response: resp} - - servers := collectServers([]traverse.TraversalResult{result}) - if len(servers) == 0 { - t.Fatal("expected at least one server entry") - } - if _, ok := servers["com"]; !ok { - t.Errorf("expected server name 'com', got keys: %v", servers) - } -} - -func TestServerNameFallbacks(t *testing.T) { - // No bailiwick, no NSName, with server IP - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - } - result := traverse.TraversalResult{Referral: ref, Response: resp} - name := serverName(result) - if name != "1.2.3.4" { - t.Errorf("serverName with root bailiwick = %q, want IP", name) - } -} - -func TestServerNameWithNSName(t *testing.T) { - ref := &traverse.Referral{ - Name: "example.com.", - Qtype: dns.TypeA, - Bailiwick: ".", - NSName: "ns1.example.com.", - } - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("5.5.5.5"), - Type: traverse.RespAnswer, - } - result := traverse.TraversalResult{Referral: ref, Response: resp} - // Bailiwick is "." so falls through to NSName - name := serverName(result) - if name == "" { - t.Error("serverName should return non-empty string") - } -} - -func TestServerNameNilReferral(t *testing.T) { - resp := &traverse.Response{ - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - } - result := traverse.TraversalResult{Referral: nil, Response: resp} - name := serverName(result) - if name == "" { - t.Error("serverName with nil referral should return non-empty string") - } -} - -func TestContainsString(t *testing.T) { - items := []string{"a", "b", "c"} - if !containsString(items, "b") { - t.Error("containsString should find 'b' in slice") - } - if containsString(items, "d") { - t.Error("containsString should not find 'd' in slice") - } - if containsString(nil, "a") { - t.Error("containsString on nil slice should return false") - } -} - -func TestComputeSummaryMixedResults(t *testing.T) { - refAnswer := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 0.6, nil) - respAnswer := &traverse.Response{ - Referral: refAnswer, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, - A: net.ParseIP("1.2.3.4"), - }, - }, - }, - } - - refNXD := traverse.NewReferral("notexist.com.", dns.TypeA, ".", 0, 0.4, nil) - respNXD := &traverse.Response{ - Referral: refNXD, - Type: traverse.RespNXDOMAIN, - } - - results := []traverse.TraversalResult{ - {Referral: refAnswer, Response: respAnswer}, - {Referral: refNXD, Response: respNXD}, - } - - stats := ComputeSummary(results) - if stats == nil { - t.Fatal("ComputeSummary returned nil for non-empty results") - } - if len(stats.Answers) != 1 { - t.Errorf("expected 1 answer entry, got %d", len(stats.Answers)) - } - if _, ok := stats.ByType["nxdomain"]; !ok { - t.Error("expected nxdomain in ByType") - } -} - -func TestComputeSummaryAnswerWithCNAMEOnly(t *testing.T) { - // Answer with only CNAME record - no final answer, should be in ByType - ref := traverse.NewReferral("www.example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.CNAME{ - Hdr: miekgdns.RR_Header{Name: "www.example.com.", Rrtype: miekgdns.TypeCNAME, Class: miekgdns.ClassINET}, - Target: "example.com.", - }, - }, - }, - } - results := []traverse.TraversalResult{{Referral: ref, Response: resp}} - stats := ComputeSummary(results) - if stats == nil { - t.Fatal("ComputeSummary returned nil") - } -} - -func TestComputeSummaryAccumulates(t *testing.T) { - // Two answers with the same IP should accumulate probability - ref1 := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 0.5, nil) - ref2 := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 0.5, nil) - - makeResp := func(ref *traverse.Referral) *traverse.Response { - return &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, - A: net.ParseIP("1.2.3.4"), - }, - }, - }, - } - } - - results := []traverse.TraversalResult{ - {Referral: ref1, Response: makeResp(ref1)}, - {Referral: ref2, Response: makeResp(ref2)}, - } - stats := ComputeSummary(results) - if stats == nil { - t.Fatal("ComputeSummary returned nil") - } - if len(stats.Answers) != 1 { - t.Fatalf("expected 1 answer after accumulation, got %d", len(stats.Answers)) - } - if stats.Answers[0].Prob < 0.99 { - t.Errorf("accumulated prob = %.2f, want ~1.0", stats.Answers[0].Prob) - } -} - -func TestCollectUniqueServerIPs(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - ip1 := net.ParseIP("1.2.3.4") - ip2 := net.ParseIP("5.6.7.8") - - results := []traverse.TraversalResult{ - {Referral: ref, Response: &traverse.Response{Server: ip1, Type: traverse.RespAnswer}}, - {Referral: ref, Response: &traverse.Response{Server: ip1, Type: traverse.RespAnswer}}, // dup - {Referral: ref, Response: &traverse.Response{Server: ip2, Type: traverse.RespAnswer}}, - {Referral: ref, Response: nil}, // nil response - } - - ips := collectUniqueServerIPs(results) - if len(ips) != 2 { - t.Errorf("expected 2 unique IPs, got %d", len(ips)) - } -} - -func TestCollectUniqueServerIPsEmpty(t *testing.T) { - ips := collectUniqueServerIPs(nil) - if len(ips) != 0 { - t.Errorf("expected 0 IPs for nil results, got %d", len(ips)) - } -} diff --git a/internal/output/text.go b/internal/output/text.go index c7482d6..dd7b18b 100644 --- a/internal/output/text.go +++ b/internal/output/text.go @@ -6,7 +6,6 @@ import ( "sort" "strings" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" ) @@ -19,52 +18,143 @@ func newTextFormatter(cfg *Config, w io.Writer) *textFormatter { return &textFormatter{cfg: cfg, w: w} } +// WriteHeader renders the pre-run header block (bin/dnstraverse): the "#" +// settings lines, the initial root, and the "Running query" line. --quiet +// suppresses the whole block. The EDNS0 state reflects the UDP size (the Ruby +// source always printed "on" — a documented deviation we fix). +func (f *textFormatter) WriteHeader(roots []traverse.StartServer) error { + if f.cfg.Quiet { + return nil + } + var b strings.Builder + if f.cfg.Fast { + b.WriteString("# Using fast mode\n") + } + if !f.cfg.AllRootServers { + b.WriteString("# Limiting traverse to one root\n") + } + edns := "on" + if f.cfg.UDPSize <= 512 { + edns = "off" + } + fmt.Fprintf(&b, "# UDP size %d (EDNS0 is %s)\n", f.cfg.UDPSize, edns) + fmt.Fprintf(&b, "# Retries %d, max depth %d\n", f.cfg.Retries, f.cfg.MaxDepth) + fmt.Fprintf(&b, "# Allow TCP is %t, always TCP is %t\n", f.cfg.AllowTCP, f.cfg.AlwaysTCP) + if len(roots) > 0 { + ip := "" + if len(roots[0].IPs) > 0 { + ip = roots[0].IPs[0] + } + fmt.Fprintf(&b, "Using %s (%s) as initial root\n", roots[0].Name, ip) + if f.cfg.AllRootServers { + b.WriteString("All roots:\n") + for _, root := range roots { + fmt.Fprintf(&b, " %s %s\n", root.Name, strings.Join(root.IPs, ", ")) + } + } + } + fmt.Fprintf(&b, "Running query %s type %s\n", f.cfg.Domain, f.cfg.QueryType) + _, err := io.WriteString(f.w, b.String()) + return err +} + func (f *textFormatter) WriteProgress(event traverse.TraversalEvent) error { - if event.Stage != traverse.EventStart { + ref := event.Referral + if ref == nil || ref.IsRootRoot() { return nil } - line := f.formatReferralLine(event.Result, false) - if !event.Result.Referral.HasAddresses() { - line += " -- resolving" + switch event.Stage { + case traverse.StageStart: + line := f.referralTxt(ref) + if !ref.Resolved() { + line += " -- resolving" + } + return f.writeLine(line) + case traverse.StageAnswerFast: + return f.writeLine(fmt.Sprintf("%s -- completed earlier (%s)", + f.referralTxt(ref), event.CompletedEarlier)) + case traverse.StageNewReferralSet: + // One line per extra childset: the parent refid and the IP that + // produced the children (progress_main :new_referral_set). + refid := event.RefID + if i := strings.LastIndex(refid, "."); i >= 0 { + refid = refid[:i] + } + return f.writeLine(fmt.Sprintf("%s %s", refid, ref.ParentIP)) + case traverse.StageAnswer: + if f.cfg.Verbose { + for _, warning := range ref.Warnings { + if err := f.writeLine(fmt.Sprintf("%s WARNING: %s", event.RefID, warning)); err != nil { + return err + } + } + } + if f.cfg.ShowAllStats { + return f.writeStatsBlocks(ref, fmt.Sprintf("%s Results:", event.RefID), false) + } } - return f.writeLine(line) + return nil } +// WriteResolve renders progress for glue-resolution subtree nodes +// (progress_resolves): like main progress, but the fast-mode marker carries +// no refid. func (f *textFormatter) WriteResolve(event traverse.TraversalEvent) error { - if event.Stage != traverse.EventStart { + ref := event.Referral + if ref == nil || ref.IsRootRoot() { return nil } - return f.writeLine(f.formatReferralLine(event.Result, true)) -} - -func (f *textFormatter) WriteResult(result traverse.TraversalResult) error { - if result.Response == nil || result.Referral == nil { - return nil + switch event.Stage { + case traverse.StageStart: + return f.writeLine(f.referralTxt(ref)) + case traverse.StageAnswerFast: + return f.writeLine(f.referralTxt(ref) + " -- completed earlier") } - prefix := strings.Repeat(" ", result.Referral.Depth+1) - line := prefix + f.formatResultLine(result) - return f.writeLine(line) + return nil } -func (f *textFormatter) WriteSummary(results []traverse.TraversalResult) error { +// referralTxt renders one progress row: " ( )"; verbose +// adds "[qname]" and " ". Unresolved servers have no parens +// (referral_txt_normal / referral_txt_verbose in bin/dnstraverse). +func (f *textFormatter) referralTxt(ref *traverse.Referral) string { + var b strings.Builder + b.WriteString(ref.RefID) + if f.cfg.Verbose { + fmt.Fprintf(&b, " [%s]", ref.Qname) + } + b.WriteString(" " + ref.Server) + if ref.Resolved() { + fmt.Fprintf(&b, " (%s)", ref.TxtIPs()) + } + if f.cfg.Verbose { + fmt.Fprintf(&b, " <%s>", ref.Bailiwick) + } + return b.String() +} + +func (f *textFormatter) WriteSummary(root *traverse.Referral, servers map[string][]string) error { + // Blank line separating progress from the sections (bin/dnstraverse: + // "puts if options[:progress]"). + if f.cfg.ShowProgress { + if _, err := fmt.Fprintln(f.w); err != nil { + return err + } + } if f.cfg.ShowServers { - if err := f.writeServers(results); err != nil { + if err := f.writeServers(servers); err != nil { return err } } - if f.cfg.ShowResults { - if err := f.writeResults(results); err != nil { + if err := f.writeResults(root); err != nil { return err } } - if f.cfg.ShowSummaryResults { - if err := f.writeSummaryResults(results); err != nil { + if err := f.writeSummaryResults(root); err != nil { return err } } - return nil } @@ -72,8 +162,10 @@ func (f *textFormatter) Flush() error { return nil } -func (f *textFormatter) writeServers(results []traverse.TraversalResult) error { - servers := collectServers(results) +// writeServers renders "The following servers were encountered:" sorted by +// lowercased reversed name (bin/dnstraverse); the name column is at least 16 +// characters wide. +func (f *textFormatter) writeServers(servers map[string][]string) error { if len(servers) == 0 { return nil } @@ -83,29 +175,26 @@ func (f *textFormatter) writeServers(results []traverse.TraversalResult) error { } names := make([]string, 0, len(servers)) + width := 16 for name := range servers { names = append(names, name) - } - sort.Slice(names, func(i, j int) bool { - return strings.ToLower(names[i]) > strings.ToLower(names[j]) - }) - - width := 16 - for _, name := range names { if len(name) > width { width = len(name) } } + sort.Slice(names, func(i, j int) bool { + return reverseString(strings.ToLower(names[i])) < reverseString(strings.ToLower(names[j])) + }) for _, name := range names { for _, ip := range servers[name] { line := fmt.Sprintf("%*s: %-15s", width, name, ip) if f.cfg.ShowVersions { if version, ok := f.cfg.Fingerprints[ip]; ok && version != "" { - line += " " + version + line += " " + version } } - if _, err := fmt.Fprintln(f.w, line); err != nil { + if _, err := fmt.Fprintln(f.w, strings.TrimRight(line, " ")); err != nil { return err } } @@ -114,129 +203,167 @@ func (f *textFormatter) writeServers(results []traverse.TraversalResult) error { return err } -func (f *textFormatter) writeResults(results []traverse.TraversalResult) error { +func reverseString(s string) string { + runes := []rune(s) + for i, j := 0, len(runes)-1; i < j; i, j = i+1, j-1 { + runes[i], runes[j] = runes[j], runes[i] + } + return string(runes) +} + +func (f *textFormatter) writeResults(root *traverse.Referral) error { + if root == nil { + return nil + } if _, err := fmt.Fprintln(f.w, "Results:"); err != nil { return err } - - terminal := terminalResults(results) - deduped := DeduplicateResults(terminal) - for _, result := range deduped { - prefix := strings.Repeat(" ", result.Referral.Depth+1) - line := prefix + f.formatResultLine(result) - if _, err := fmt.Fprintln(f.w, line); err != nil { - return err - } + if err := f.writeStatsBlocks(root, "", true); err != nil { + return err } _, err := fmt.Fprintln(f.w) return err } -func (f *textFormatter) writeSummaryResults(results []traverse.TraversalResult) error { - stats := ComputeSummary(results) +// writeStatsBlocks renders every aggregated leaf of ref, sorted by stats key +// (referral.rb stats_display). With spacing, blocks are separated by blank +// lines. +func (f *textFormatter) writeStatsBlocks(ref *traverse.Referral, prefix string, spacing bool) error { + first := true + for _, leaf := range ref.StatsList() { + if spacing && !first { + if _, err := fmt.Fprintln(f.w); err != nil { + return err + } + } + first = false + for _, line := range f.formatLeaf(leaf, prefix) { + if err := f.writeLine(line); err != nil { + return err + } + } + } + return nil +} + +func (f *textFormatter) writeSummaryResults(root *traverse.Referral) error { + stats := root.SummaryStats() if stats == nil { return nil } - if _, err := fmt.Fprintln(f.w, "Summary:"); err != nil { + if _, err := fmt.Fprintln(f.w, "Summary Results:"); err != nil { return err } prefix := " " for _, answer := range stats.Answers { - line := fmt.Sprintf("%s%s answered with %s", prefix, formatProbability(answer.Prob), answer.RData) + initial := fmt.Sprintf("%s%s answered with ", prefix, formatProbability(answer.Prob)) + var rrs []string + for _, rr := range answer.RRs { + rrs = append(rrs, collapseWhitespace(rr.String())) + } + line := initial + strings.Join(rrs, "\n"+strings.Repeat(" ", len(initial))) if _, err := fmt.Fprintln(f.w, f.colorize(line, colorGreen)); err != nil { return err } } - types := make([]string, 0, len(stats.ByType)) - for respType := range stats.ByType { - types = append(types, respType) + statuses := make([]traverse.Status, 0, len(stats.ByStatus)) + for status := range stats.ByStatus { + if status != traverse.StatusAnswered { + statuses = append(statuses, status) + } } - sort.Strings(types) + sort.Slice(statuses, func(i, j int) bool { return statuses[i] < statuses[j] }) - for _, respType := range types { - line := fmt.Sprintf("%s%s %s", prefix, formatProbability(stats.ByType[respType]), summaryTypeLabel(respType)) + for _, status := range statuses { + line := fmt.Sprintf("%s%s %s", prefix, formatProbability(stats.ByStatus[status]), summaryStatusLabel(status)) if _, err := fmt.Fprintln(f.w, line); err != nil { return err } } - _, err := fmt.Fprintln(f.w) - return err + return nil } -func (f *textFormatter) formatReferralLine(result traverse.TraversalResult, isResolve bool) string { - ref := result.Referral - if ref == nil { - return "" +// formatLeaf renders one aggregated leaf per referral.rb stats_display: +// "%5.1f%%: " plus indented RRs for answers and the +// "While querying" line when the failing query differs from the original. +func (f *textFormatter) formatLeaf(leaf *traverse.StatsEntry, prefix string) []string { + resp := leaf.Response + ref := leaf.Referral + if resp == nil || ref == nil { + return nil } + indent := prefix + strings.Repeat(" ", 12) + where := fmt.Sprintf("%s (%s)", ref.Server, resp.IP) + head := fmt.Sprintf("%s%5.1f%%: ", prefix, leaf.Prob*100) - indent := strings.Repeat(" ", ref.Depth) - refID := referralID(ref) - server := referralServerLabel(ref, result.Response) - qtype := dns.QNameType(ref.Qtype) - qname := trimDomain(ref.Name) - - if f.cfg.Verbose { - bailiwick := trimDomain(ref.Bailiwick) - if isResolve { - return fmt.Sprintf("%s%s [%s] %s <%s>", indent, refID, qname, server, bailiwick) + var lines []string + switch resp.Status { + case traverse.StatusException: + msg := "" + if resp.DQ != nil { + msg = resp.DQ.ExceptionMessage } - return fmt.Sprintf("%s%s [%s] %s <%s> (%s)", indent, refID, qname, server, bailiwick, qtype) - } - - if isResolve { - return fmt.Sprintf("%s%s %s", indent, refID, server) - } - return fmt.Sprintf("%s%s %s (%s)", indent, refID, server, qtype) -} - -func (f *textFormatter) formatResultLine(result traverse.TraversalResult) string { - prob := formatProbability(result.Referral.Prob) - switch result.Response.Type { - case traverse.RespAnswer: - key, _ := answerKey(result.Response) - if key == "" { - return fmt.Sprintf("%s resulted in answer", prob) + lines = append(lines, head+f.colorize(fmt.Sprintf("%s at %s", msg, where), colorRed)) + case traverse.StatusNoGlue: + parent := "" + if ref.Parent != nil { + parent = ref.Parent.Server } - nsLabel := "" - if result.Referral.Bailiwick != "" && result.Referral.Bailiwick != "." { - nsLabel = trimDomain(result.Referral.Bailiwick) + " " + lines = append(lines, head+f.colorize(fmt.Sprintf("No glue at %s (%s) for %s", parent, resp.IP, ref.Server), colorYellow)) + case traverse.StatusReferralLame: + parent := "" + if ref.Parent != nil { + parent = ref.Parent.Server } - return f.colorize(fmt.Sprintf("%s %sanswered with %s", prob, nsLabel, key), colorGreen) - case traverse.RespNODATA: - return fmt.Sprintf("%s found no such record", prob) - case traverse.RespNXDOMAIN: - return f.colorize(fmt.Sprintf("%s name does not exist", prob), colorYellow) - case traverse.RespSERVFAIL: - return f.colorize(fmt.Sprintf("%s resulted in SERVFAIL", prob), colorRed) - case traverse.RespREFUSED: - return f.colorize(fmt.Sprintf("%s query refused by server", prob), colorRed) - case traverse.RespNOTIMPL: - return f.colorize(fmt.Sprintf("%s query type not implemented by server", prob), colorRed) - case traverse.RespCNAMELoop: - msg := "CNAME loop detected" - if result.Response.ErrorMessage != "" { - msg = result.Response.ErrorMessage + lines = append(lines, head+f.colorize(fmt.Sprintf("Lame referral from %s (%s) to %s", parent, ref.ParentIP, where), colorYellow)) + case traverse.StatusLoop: + lines = append(lines, head+f.colorize(fmt.Sprintf("Loop encountered at %s", resp.Server), colorRed)) + case traverse.StatusCNAMELoop: + lines = append(lines, head+f.colorize(fmt.Sprintf("CNAME loop encountered at %s", resp.Server), colorRed)) + case traverse.StatusError: + msg := "" + if resp.DQ != nil { + msg = resp.DQ.ErrorMessage } - return f.colorize(fmt.Sprintf("%s %s", prob, msg), colorRed) - case traverse.RespNSResolutionFailed: - msg := "nameserver lookup failed" - if result.Response.ErrorMessage != "" { - msg = result.Response.ErrorMessage + lines = append(lines, head+f.colorize(fmt.Sprintf("%s at %s", msg, where), colorRed)) + case traverse.StatusNoData: + lines = append(lines, head+fmt.Sprintf("NODATA (for this type) at %s", where)) + case traverse.StatusAnswered: + lines = append(lines, head+f.colorize(fmt.Sprintf("Answer from %s", where), colorGreen)) + if resp.DQ != nil { + for _, rr := range resp.DQ.Answers { + lines = append(lines, indent+rr.String()) + } } - return f.colorize(fmt.Sprintf("%s %s", prob, msg), colorYellow) - case traverse.RespError: - msg := "resulted in an error" - if result.Response.ErrorMessage != "" { - msg = result.Response.ErrorMessage - } - return f.colorize(fmt.Sprintf("%s %s", prob, msg), colorRed) default: - return fmt.Sprintf("%s %s", prob, result.Response.Type) + // The Ruby fallback prints "Stopped at ( ))" with a stray + // paren — a documented deviation we fix. + lines = append(lines, head+fmt.Sprintf("Stopped at %s", where)) + lines = append(lines, indent+leaf.Key) } + + if resp.Status != traverse.StatusAnswered { + origQname, origQclass, origQtype := originalQuery(ref) + if resp.Qname != origQname || resp.Qclass != origQclass || resp.Qtype != origQtype { + lines = append(lines, indent+fmt.Sprintf("While querying %s/%s/%s", + resp.Qname, traverse.ClassToString(resp.Qclass), traverse.TypeToString(resp.Qtype))) + } + } + return lines +} + +// originalQuery walks to the rootroot node to find the query the whole +// traversal was started for. +func originalQuery(ref *traverse.Referral) (string, uint16, uint16) { + top := ref + for top.Parent != nil { + top = top.Parent + } + return top.Qname, top.Qclass, top.Qtype } func (f *textFormatter) writeLine(line string) error { @@ -254,43 +381,6 @@ func (f *textFormatter) colorize(text, color string) string { return color + text + colorReset } -func referralID(ref *traverse.Referral) string { - if ref == nil { - return "" - } - return fmt.Sprintf("%d", ref.Depth+1) -} - -func referralServerLabel(ref *traverse.Referral, resp *traverse.Response) string { - if resp != nil && resp.Server != nil { - return resp.Server.String() - } - if ref.HasAddresses() { - ips := make([]string, 0, len(ref.Addresses)) - for _, addr := range ref.Addresses { - ips = append(ips, addr.String()) - } - return strings.Join(ips, ", ") - } - if ref.NSName != "" { - return ref.NSName - } - if ref.Bailiwick != "" && ref.Bailiwick != "." { - return trimDomain(ref.Bailiwick) - } - return "unknown" -} - -func terminalResults(results []traverse.TraversalResult) []traverse.TraversalResult { - var terminal []traverse.TraversalResult - for _, result := range results { - if result.Response != nil && result.Response.IsTerminal() { - terminal = append(terminal, result) - } - } - return terminal -} - const ( colorReset = "\033[0m" colorGreen = "\033[32m" diff --git a/internal/output/text_test.go b/internal/output/text_test.go deleted file mode 100644 index c11342f..0000000 --- a/internal/output/text_test.go +++ /dev/null @@ -1,379 +0,0 @@ -package output - -import ( - "bytes" - "net" - "strings" - "testing" - - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" - miekgdns "github.com/miekg/dns" -) - -func TestTextFormatterProgressIndentation(t *testing.T) { - root := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - child := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 0.5, root) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - formatter := newTextFormatter(cfg, &buf) - - if err := formatter.WriteProgress(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: child}, - }); err != nil { - t.Fatalf("WriteProgress: %v", err) - } - - out := buf.String() - if !strings.HasPrefix(out, " 2 ") { - t.Fatalf("expected depth-based indentation, got %q", out) - } -} - -func TestAttachHooksRespectsShowFlags(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - var progressCount int - - cfg := DefaultConfig() - cfg.ShowProgress = false - cfg.ShowResolves = false - cfg.ShowAllStats = false - - var buf bytes.Buffer - formatter := NewFormatter(cfg, &buf) - hooks := AttachHooks(cfg, formatter) - hooks.OnEvent(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - }) - - if buf.Len() != 0 { - t.Fatalf("expected no output when ShowProgress is false") - } - - cfg.ShowProgress = true - hooks = AttachHooks(cfg, formatter) - hooks.OnEvent(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - }) - progressCount = strings.Count(buf.String(), "\n") - if progressCount == 0 { - t.Fatal("expected progress output when ShowProgress is true") - } -} - -func TestTextWriteResolve(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -server := net.ParseIP("198.41.0.4") -resp := &traverse.Response{Server: server, Type: traverse.RespAnswer} - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -f := newTextFormatter(cfg, &buf) - -// EventStart - should write line -if err := f.WriteResolve(traverse.TraversalEvent{ -Stage: traverse.EventStart, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}); err != nil { -t.Fatalf("WriteResolve EventStart: %v", err) -} -if buf.Len() == 0 { -t.Error("expected output for WriteResolve EventStart") -} - -buf.Reset() -// EventComplete - should write nothing -if err := f.WriteResolve(traverse.TraversalEvent{ -Stage: traverse.EventComplete, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}); err != nil { -t.Fatalf("WriteResolve EventComplete: %v", err) -} -if buf.Len() != 0 { -t.Error("expected no output for WriteResolve EventComplete") -} -} - -func TestTextWriteResult(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -server := net.ParseIP("198.41.0.4") - -tests := []struct { -name string -respType traverse.ResponseType -msg *dns.DecodedResponse -errorMsg string -}{ -{"answer", traverse.RespAnswer, &dns.DecodedResponse{ -Answers: []miekgdns.RR{ -&miekgdns.A{Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, A: net.ParseIP("1.2.3.4")}, -}, -}, ""}, -{"nodata", traverse.RespNODATA, nil, ""}, -{"nxdomain", traverse.RespNXDOMAIN, nil, ""}, -{"servfail", traverse.RespSERVFAIL, nil, ""}, -{"refused", traverse.RespREFUSED, nil, ""}, -{"notimp", traverse.RespNOTIMPL, nil, ""}, -{"cname_loop", traverse.RespCNAMELoop, nil, "loop detected"}, -{"ns_error", traverse.RespNSResolutionFailed, nil, "nameserver ns1.example.com could not be resolved"}, -{"error", traverse.RespError, nil, "something went wrong"}, -} - -for _, tc := range tests { -t.Run(tc.name, func(t *testing.T) { -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -f := newTextFormatter(cfg, &buf) - -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: tc.respType, -Decoded: tc.msg, -ErrorMessage: tc.errorMsg, -} -result := traverse.TraversalResult{Referral: ref, Response: resp} -if err := f.WriteResult(result); err != nil { -t.Fatalf("WriteResult %q: %v", tc.name, err) -} -}) -} -} - -func TestTextWriteResultNilResponse(t *testing.T) { -var buf bytes.Buffer -cfg := DefaultConfig() -f := newTextFormatter(cfg, &buf) -if err := f.WriteResult(traverse.TraversalResult{Referral: nil, Response: nil}); err != nil { -t.Fatalf("WriteResult nil: %v", err) -} -if buf.Len() != 0 { -t.Error("expected no output for nil result") -} -} - -func TestTextWriteResultAnswerMultipleRRs(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 0.5, nil) -resp := &traverse.Response{ -Referral: ref, -Server: net.ParseIP("1.2.3.4"), -Type: traverse.RespAnswer, -Decoded: &dns.DecodedResponse{ -Answers: []miekgdns.RR{ -&miekgdns.A{Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, A: net.ParseIP("1.2.3.4")}, -&miekgdns.A{Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, A: net.ParseIP("5.6.7.8")}, -}, -}, -} -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -f := newTextFormatter(cfg, &buf) -if err := f.WriteResult(traverse.TraversalResult{Referral: ref, Response: resp}); err != nil { -t.Fatalf("WriteResult: %v", err) -} -if !strings.Contains(buf.String(), "/") { -t.Errorf("expected '/' separator for multiple answers, got: %q", buf.String()) -} -} - -func TestTextWriteSummaryWithServersAndResults(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) -server := net.ParseIP("1.2.3.4") -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: traverse.RespAnswer, -Decoded: &dns.DecodedResponse{ -Answers: []miekgdns.RR{ -&miekgdns.A{Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, A: net.ParseIP("1.2.3.4")}, -}, -}, -} -results := []traverse.TraversalResult{{Referral: ref, Response: resp}} - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -cfg.ShowServers = true -cfg.ShowResults = true -cfg.ShowSummaryResults = true -f := newTextFormatter(cfg, &buf) -if err := f.WriteSummary(results); err != nil { -t.Fatalf("WriteSummary: %v", err) -} -out := buf.String() -if !strings.Contains(out, "Summary:") { -t.Errorf("expected Summary: in output, got: %q", out) -} -if !strings.Contains(out, "Results:") { -t.Errorf("expected Results: in output, got: %q", out) -} -} - -func TestTextWriteSummaryNoResults(t *testing.T) { -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.ShowServers = false -cfg.ShowResults = false -cfg.ShowSummaryResults = false -f := newTextFormatter(cfg, &buf) -if err := f.WriteSummary(nil); err != nil { -t.Fatalf("WriteSummary nil: %v", err) -} -} - -func TestTextWriteSummaryNXDOMAIN(t *testing.T) { -ref := traverse.NewReferral("gone.example.com.", dns.TypeA, "com.", 1, 1.0, nil) -server := net.ParseIP("1.2.3.4") -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: traverse.RespNXDOMAIN, -} -results := []traverse.TraversalResult{{Referral: ref, Response: resp}} - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -cfg.ShowServers = true -cfg.ShowResults = true -cfg.ShowSummaryResults = true -f := newTextFormatter(cfg, &buf) -if err := f.WriteSummary(results); err != nil { -t.Fatalf("WriteSummary: %v", err) -} -} - -func TestFormatReferralLineVerbose(t *testing.T) { -root := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -child := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, root) - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -cfg.Verbose = true -f := newTextFormatter(cfg, &buf) - -event := traverse.TraversalEvent{ -Stage: traverse.EventStart, -Result: traverse.TraversalResult{Referral: child}, -} -if err := f.WriteProgress(event); err != nil { -t.Fatalf("WriteProgress verbose: %v", err) -} -out := buf.String() -if !strings.Contains(out, "com") { -t.Errorf("expected bailiwick in verbose output, got: %q", out) -} -} - -func TestFormatReferralLineVerboseResolve(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) -server := net.ParseIP("1.2.3.4") -resp := &traverse.Response{Server: server, Type: traverse.RespAnswer} - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -cfg.Verbose = true -f := newTextFormatter(cfg, &buf) - -if err := f.WriteResolve(traverse.TraversalEvent{ -Stage: traverse.EventStart, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}); err != nil { -t.Fatalf("WriteResolve verbose: %v", err) -} -if buf.Len() == 0 { -t.Error("expected output for verbose WriteResolve") -} -} - -func TestTextWriteProgressNoAddresses(t *testing.T) { -// Test the "resolving" suffix when referral has no addresses -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -// No addresses set, so HasAddresses() returns false - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -f := newTextFormatter(cfg, &buf) -if err := f.WriteProgress(traverse.TraversalEvent{ -Stage: traverse.EventStart, -Result: traverse.TraversalResult{Referral: ref}, -}); err != nil { -t.Fatalf("WriteProgress: %v", err) -} -if !strings.Contains(buf.String(), "resolving") { -t.Errorf("expected 'resolving' suffix when no addresses, got: %q", buf.String()) -} -} - -func TestColorize(t *testing.T) { -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = true -f := newTextFormatter(cfg, &buf) - -colored := f.colorize("hello", colorGreen) -if colored == "hello" { -t.Error("expected colorized output with Color=true") -} - -cfg.Color = false -f2 := newTextFormatter(cfg, &buf) -plain := f2.colorize("hello", colorGreen) -if plain != "hello" { -t.Errorf("expected plain text with Color=false, got %q", plain) -} - -// Empty color -empty := f.colorize("hello", "") -if empty != "hello" { -t.Errorf("expected plain text for empty color, got %q", empty) -} -} - -func TestReferralServerLabelFallbacks(t *testing.T) { -// With server IP in response -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -resp := &traverse.Response{Server: net.ParseIP("1.2.3.4")} -label := referralServerLabel(ref, resp) -if label != "1.2.3.4" { -t.Errorf("expected '1.2.3.4', got %q", label) -} - -// With addresses in referral, no response server -ref2 := traverse.NewReferral("example.com.", dns.TypeA, "ns1.example.com.", 0, 1.0, nil) -ref2.Addresses = []net.IP{net.ParseIP("5.6.7.8")} -label2 := referralServerLabel(ref2, nil) -if label2 != "5.6.7.8" { -t.Errorf("expected '5.6.7.8', got %q", label2) -} - -// With NSName -ref3 := &traverse.Referral{ -Name: "example.com.", -NSName: "ns1.example.com.", -Bailiwick: ".", -} -label3 := referralServerLabel(ref3, nil) -if label3 != "ns1.example.com." { -t.Errorf("expected NSName, got %q", label3) -} - -// With non-root bailiwick, no addresses, no NSName -ref4 := traverse.NewReferral("example.com.", dns.TypeA, "com.", 0, 1.0, nil) -label4 := referralServerLabel(ref4, nil) -if label4 != "com" { -t.Errorf("expected 'com', got %q", label4) -} -} diff --git a/internal/traverse/cache.go b/internal/traverse/cache.go index 9894324..b7ca65e 100644 --- a/internal/traverse/cache.go +++ b/internal/traverse/cache.go @@ -1,6 +1,7 @@ package traverse import ( + "fmt" "net" "strings" "sync" @@ -8,117 +9,157 @@ import ( miekgdns "github.com/miekg/dns" ) +// StartServer is one entry returned by GetStartServers: a nameserver hostname +// plus its cached IPv4 addresses. IPs == nil means no addresses are cached and +// the caller must resolve the name itself (glueless). +type StartServer struct { + Name string + IPs []string +} + +// InfoCache is the hierarchical per-branch record cache (info_cache.rb). Each +// response wraps its parent's cache in a child so sibling branches never see +// each other's records; lookups recurse towards the root cache. type InfoCache struct { parent *InfoCache mu sync.RWMutex - ns map[string][]string - glue map[string][]net.IP + data map[string][]miekgdns.RR } func NewInfoCache(parent *InfoCache) *InfoCache { return &InfoCache{ parent: parent, - ns: make(map[string][]string), - glue: make(map[string][]net.IP), + data: make(map[string][]miekgdns.RR), } } -func (c *InfoCache) StoreNS(zone string, nameservers []string) { - if len(nameservers) == 0 { - return - } - zone = normalize(zone) - c.mu.Lock() - seen := make(map[string]bool) - for _, ns := range nameservers { - ns = normalize(ns) - if !seen[ns] { - seen[ns] = true - c.ns[zone] = append(c.ns[zone], ns) - } - } - c.mu.Unlock() -} - -func (c *InfoCache) LookupNS(zone string) []string { - zone = normalize(zone) - if names := c.localNS(zone); len(names) > 0 { - return names - } - if c.parent != nil { - return c.parent.LookupNS(zone) - } - return nil -} - -func (c *InfoCache) localNS(zone string) []string { - c.mu.RLock() - defer c.mu.RUnlock() - names, ok := c.ns[zone] - if !ok { - return nil - } - result := make([]string, len(names)) - copy(result, names) - return result -} - -func (c *InfoCache) StoreGlue(name string, addrs []net.IP) { - if len(addrs) == 0 { - return - } - name = normalize(name) - c.mu.Lock() - seen := make(map[string]bool) - for _, addr := range addrs { - key := addr.String() - if !seen[key] { - seen[key] = true - c.glue[name] = append(c.glue[name], addr) - } - } - c.mu.Unlock() -} - -func (c *InfoCache) LookupGlue(name string) []net.IP { - name = normalize(name) - if addrs := c.localGlue(name); len(addrs) > 0 { - return addrs - } - if c.parent != nil { - return c.parent.LookupGlue(name) - } - return nil -} - -func (c *InfoCache) localGlue(name string) []net.IP { - c.mu.RLock() - defer c.mu.RUnlock() - addrs, ok := c.glue[name] - if !ok { - return nil - } - result := make([]net.IP, len(addrs)) - copy(result, addrs) - return result -} - func (c *InfoCache) Child() *InfoCache { return NewInfoCache(c) } -func (c *InfoCache) NSCount() int { - c.mu.RLock() - defer c.mu.RUnlock() - return len(c.ns) +// canonicalName lowercases a DNS name and strips the trailing dot; the root +// (and empty string) canonicalises to "", matching the Ruby engine's +// representation of "no bailiwick". +func canonicalName(name string) string { + return strings.TrimSuffix(strings.ToLower(name), ".") } -func (c *InfoCache) GlueCount() int { - c.mu.RLock() - defer c.mu.RUnlock() - return len(c.glue) +func cacheKey(name string, qclass, qtype uint16) string { + return fmt.Sprintf("%s:%d:%d", canonicalName(name), qclass, qtype) } -func normalize(name string) string { - return strings.ToLower(miekgdns.Fqdn(name)) +func rrCacheKey(rr miekgdns.RR) string { + h := rr.Header() + return cacheKey(h.Name, h.Class, h.Rrtype) +} + +// Add stores resource records, REPLACING any existing entries that share a +// name:class:type key (info_cache.rb add: clear pass, then append pass, so +// several records under one key in a single call are all kept). +func (c *InfoCache) Add(rrs []miekgdns.RR) { + c.mu.Lock() + defer c.mu.Unlock() + for _, rr := range rrs { + c.data[rrCacheKey(rr)] = nil + } + for _, rr := range rrs { + key := rrCacheKey(rr) + c.data[key] = append(c.data[key], rr) + } +} + +// AddHints seeds NS records for domain ("" = root hints) plus A/AAAA records +// for each server that has known addresses (info_cache.rb add_hints). +func (c *InfoCache) AddHints(domain string, servers []StartServer) { + var rrs []miekgdns.RR + owner := miekgdns.Fqdn(canonicalName(domain)) + for _, srv := range servers { + name := miekgdns.Fqdn(canonicalName(srv.Name)) + rrs = append(rrs, &miekgdns.NS{ + Hdr: miekgdns.RR_Header{Name: owner, Rrtype: miekgdns.TypeNS, Class: miekgdns.ClassINET}, + Ns: name, + }) + for _, ip := range srv.IPs { + addr := net.ParseIP(ip) + if addr == nil { + continue + } + if v4 := addr.To4(); v4 != nil { + rrs = append(rrs, &miekgdns.A{ + Hdr: miekgdns.RR_Header{Name: name, Rrtype: miekgdns.TypeA, Class: miekgdns.ClassINET}, + A: v4, + }) + } else { + rrs = append(rrs, &miekgdns.AAAA{ + Hdr: miekgdns.RR_Header{Name: name, Rrtype: miekgdns.TypeAAAA, Class: miekgdns.ClassINET}, + AAAA: addr, + }) + } + } + } + c.Add(rrs) +} + +// Get returns the cached RRset for name/class/type, consulting parent caches +// on a local miss. Returns nil when nothing is cached anywhere in the chain. +func (c *InfoCache) Get(name string, qclass, qtype uint16) []miekgdns.RR { + key := cacheKey(name, qclass, qtype) + c.mu.RLock() + rrs, ok := c.data[key] + c.mu.RUnlock() + if ok { + out := make([]miekgdns.RR, len(rrs)) + copy(out, rrs) + return out + } + if c.parent != nil { + return c.parent.Get(name, qclass, qtype) + } + return nil +} + +// getNS finds the nearest cached NS RRset at or above domain, walking labels +// upward to the root (info_cache.rb get_ns?). +func (c *InfoCache) getNS(domain string) ([]miekgdns.RR, error) { + domain = canonicalName(domain) + for { + if rrs := c.Get(domain, miekgdns.ClassINET, miekgdns.TypeNS); len(rrs) > 0 { + return rrs, nil + } + if domain == "" { + return nil, fmt.Errorf("no nameservers available for %q -- no root hints set??", domain) + } + if i := strings.Index(domain, "."); i >= 0 { + domain = domain[i+1:] + } else { + domain = "" + } + } +} + +// GetStartServers returns the servers to start querying for domain: the +// nearest cached NS RRset walking labels upward, each nameserver paired with +// its cached A addresses (nil when unknown). newbailiwick is the owner name of +// that NS RRset ("" for root). +func (c *InfoCache) GetStartServers(domain string) (starters []StartServer, newbailiwick string, err error) { + ns, err := c.getNS(domain) + if err != nil { + return nil, "", err + } + for _, rr := range ns { + nsrr, ok := rr.(*miekgdns.NS) + if !ok { + continue + } + name := canonicalName(nsrr.Ns) + var ips []string + for _, iprr := range c.Get(name, miekgdns.ClassINET, miekgdns.TypeA) { + if a, ok := iprr.(*miekgdns.A); ok { + ips = append(ips, a.A.String()) + } + } + starters = append(starters, StartServer{Name: name, IPs: ips}) + } + newbailiwick = canonicalName(ns[0].Header().Name) + return starters, newbailiwick, nil } diff --git a/internal/traverse/cache_test.go b/internal/traverse/cache_test.go index cc024fa..0b938ef 100644 --- a/internal/traverse/cache_test.go +++ b/internal/traverse/cache_test.go @@ -3,201 +3,258 @@ package traverse import ( "fmt" "net" - "strings" "sync" "testing" "github.com/miekg/dns" ) -func TestNewInfoCache(t *testing.T) { - c := NewInfoCache(nil) - if c.parent != nil { - t.Error("root cache should have nil parent") - } - if c.NSCount() != 0 { - t.Errorf("NSCount = %d, want 0", c.NSCount()) - } - if c.GlueCount() != 0 { - t.Errorf("GlueCount = %d, want 0", c.GlueCount()) +func nsRR(zone, target string) dns.RR { + return &dns.NS{ + Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeNS, Class: dns.ClassINET}, + Ns: dns.Fqdn(target), } } -func TestInfoCacheStoreAndLookupNS(t *testing.T) { - c := NewInfoCache(nil) - - c.StoreNS("com.", []string{"a.gtld-servers.net.", "b.gtld-servers.net."}) - if c.NSCount() != 1 { - t.Errorf("NSCount = %d, want 1", c.NSCount()) - } - - names := c.LookupNS("com.") - if len(names) != 2 { - t.Fatalf("expected 2 nameservers, got %d", len(names)) - } - if names[0] != "a.gtld-servers.net." { - t.Errorf("nameserver[0] = %q, want %q", names[0], "a.gtld-servers.net.") +func aRR(name, ip string) dns.RR { + return &dns.A{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET}, + A: net.ParseIP(ip).To4(), } } -func TestInfoCacheNSDedup(t *testing.T) { - c := NewInfoCache(nil) - c.StoreNS("com.", []string{"a.gtld-servers.net.", "a.gtld-servers.net."}) - names := c.LookupNS("com.") - if len(names) != 1 { - t.Errorf("expected 1 deduped NS, got %d", len(names)) +func aaaaRR(name, ip string) dns.RR { + return &dns.AAAA{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeAAAA, Class: dns.ClassINET}, + AAAA: net.ParseIP(ip), } } -func TestInfoCacheNSCaseInsensitive(t *testing.T) { - c := NewInfoCache(nil) - c.StoreNS("COM.", []string{"A.GTLD-SERVERS.NET."}) - names := c.LookupNS("com.") - if len(names) != 1 { - t.Fatalf("expected 1 NS, got %d", len(names)) +func TestCanonicalName(t *testing.T) { + tests := []struct{ in, want string }{ + {"example.com", "example.com"}, + {"Example.COM.", "example.com"}, + {".", ""}, + {"", ""}, + {"WWW.Example.Com", "www.example.com"}, } - if names[0] != "a.gtld-servers.net." { - t.Errorf("NS = %q, want %q", names[0], "a.gtld-servers.net.") + for _, tt := range tests { + if got := canonicalName(tt.in); got != tt.want { + t.Errorf("canonicalName(%q) = %q, want %q", tt.in, got, tt.want) + } } } -func TestInfoCacheNSLookupMiss(t *testing.T) { +func TestInfoCacheAddAndGet(t *testing.T) { c := NewInfoCache(nil) - names := c.LookupNS("org.") - if names != nil { - t.Errorf("expected nil for miss, got %v", names) + c.Add([]dns.RR{nsRR("com", "a.gtld-servers.net"), nsRR("com", "b.gtld-servers.net")}) + + rrs := c.Get("com", dns.ClassINET, dns.TypeNS) + if len(rrs) != 2 { + t.Fatalf("expected 2 NS records, got %d", len(rrs)) } } -func TestInfoCacheNSStoreEmpty(t *testing.T) { +func TestInfoCacheAddReplacesSameKey(t *testing.T) { c := NewInfoCache(nil) - c.StoreNS("com.", nil) - if c.NSCount() != 0 { - t.Errorf("expected 0 after empty store, got %d", c.NSCount()) + c.Add([]dns.RR{nsRR("com", "old1.example.net"), nsRR("com", "old2.example.net")}) + c.Add([]dns.RR{nsRR("com", "new.example.net")}) + + rrs := c.Get("com", dns.ClassINET, dns.TypeNS) + if len(rrs) != 1 { + t.Fatalf("add should replace same name:class:type entry, got %d records", len(rrs)) + } + if rrs[0].(*dns.NS).Ns != "new.example.net." { + t.Errorf("NS = %q, want new.example.net.", rrs[0].(*dns.NS).Ns) } } -func TestInfoCacheChainedNS(t *testing.T) { +func TestInfoCacheAddKeepsDistinctKeys(t *testing.T) { + c := NewInfoCache(nil) + c.Add([]dns.RR{nsRR("com", "a.gtld-servers.net"), aRR("a.gtld-servers.net", "192.5.6.30")}) + c.Add([]dns.RR{nsRR("org", "a0.org-servers.net")}) + + if got := c.Get("com", dns.ClassINET, dns.TypeNS); len(got) != 1 { + t.Errorf("com NS lost after unrelated add: %v", got) + } + if got := c.Get("a.gtld-servers.net", dns.ClassINET, dns.TypeA); len(got) != 1 { + t.Errorf("glue lost after unrelated add: %v", got) + } +} + +func TestInfoCacheGetCaseInsensitive(t *testing.T) { + c := NewInfoCache(nil) + c.Add([]dns.RR{nsRR("COM.", "A.GTLD-SERVERS.NET.")}) + if got := c.Get("com", dns.ClassINET, dns.TypeNS); len(got) != 1 { + t.Fatalf("expected case-insensitive hit, got %v", got) + } +} + +func TestInfoCacheGetMiss(t *testing.T) { + c := NewInfoCache(nil) + if got := c.Get("org", dns.ClassINET, dns.TypeNS); got != nil { + t.Errorf("expected nil for miss, got %v", got) + } +} + +func TestInfoCacheGetRecursesToParent(t *testing.T) { parent := NewInfoCache(nil) - parent.StoreNS("com.", []string{"a.gtld-servers.net."}) - + parent.Add([]dns.RR{nsRR("com", "a.gtld-servers.net")}) child := parent.Child() if child.parent != parent { - t.Error("child parent should be the parent cache") + t.Fatal("Child() should link to parent") } - - names := child.LookupNS("com.") - if len(names) != 1 { - t.Fatalf("expected 1 NS from parent, got %d", len(names)) - } - if child.NSCount() != 0 { - t.Errorf("child NSCount = %d, want 0", child.NSCount()) + if got := child.Get("com", dns.ClassINET, dns.TypeNS); len(got) != 1 { + t.Fatalf("expected parent hit through child, got %v", got) } } -func TestInfoCacheChildOverridesParent(t *testing.T) { +func TestInfoCacheChildShadowsParent(t *testing.T) { parent := NewInfoCache(nil) - parent.StoreNS("com.", []string{"a.gtld-servers.net."}) - + parent.Add([]dns.RR{nsRR("com", "parent.example.net")}) child := parent.Child() - child.StoreNS("com.", []string{"b.gtld-servers.net."}) + child.Add([]dns.RR{nsRR("com", "child.example.net")}) - names := child.LookupNS("com.") - if len(names) != 1 { - t.Fatalf("expected 1 NS, got %d", len(names)) + got := child.Get("com", dns.ClassINET, dns.TypeNS) + if len(got) != 1 || got[0].(*dns.NS).Ns != "child.example.net." { + t.Errorf("child entry should shadow parent, got %v", got) } - if names[0] != "b.gtld-servers.net." { - t.Errorf("expected child's NS to override, got %q", names[0]) + // The parent must be untouched. + got = parent.Get("com", dns.ClassINET, dns.TypeNS) + if len(got) != 1 || got[0].(*dns.NS).Ns != "parent.example.net." { + t.Errorf("parent entry modified, got %v", got) } } -func TestInfoCacheStoreAndLookupGlue(t *testing.T) { +func TestGetStartServersWalksLabels(t *testing.T) { c := NewInfoCache(nil) - addrs := []net.IP{net.ParseIP("1.2.3.4"), net.ParseIP("5.6.7.8")} - c.StoreGlue("ns1.example.com.", addrs) + c.Add([]dns.RR{ + nsRR("com", "a.gtld-servers.net"), + nsRR("com", "b.gtld-servers.net"), + aRR("a.gtld-servers.net", "192.5.6.30"), + }) - result := c.LookupGlue("ns1.example.com.") - if len(result) != 2 { - t.Fatalf("expected 2 glue addresses, got %d", len(result)) + starters, bw, err := c.GetStartServers("www.deep.example.com") + if err != nil { + t.Fatalf("GetStartServers: %v", err) + } + if bw != "com" { + t.Errorf("newbailiwick = %q, want com", bw) + } + if len(starters) != 2 { + t.Fatalf("expected 2 starters, got %d", len(starters)) + } + if starters[0].Name != "a.gtld-servers.net" { + t.Errorf("starter[0] = %q", starters[0].Name) + } + if len(starters[0].IPs) != 1 || starters[0].IPs[0] != "192.5.6.30" { + t.Errorf("starter[0] IPs = %v, want [192.5.6.30]", starters[0].IPs) + } + if starters[1].IPs != nil { + t.Errorf("glueless starter should have nil IPs, got %v", starters[1].IPs) } } -func TestInfoCacheGlueDedup(t *testing.T) { +func TestGetStartServersPrefersDeepestZone(t *testing.T) { c := NewInfoCache(nil) - ip := net.ParseIP("1.2.3.4") - c.StoreGlue("ns1.example.com.", []net.IP{ip, ip}) - result := c.LookupGlue("ns1.example.com.") - if len(result) != 1 { - t.Errorf("expected 1 deduped glue, got %d", len(result)) + c.AddHints("", []StartServer{{Name: "a.root-servers.net", IPs: []string{"198.41.0.4"}}}) + c.Add([]dns.RR{nsRR("example.com", "ns1.example.com"), aRR("ns1.example.com", "1.2.3.4")}) + + starters, bw, err := c.GetStartServers("www.example.com") + if err != nil { + t.Fatalf("GetStartServers: %v", err) + } + if bw != "example.com" { + t.Errorf("newbailiwick = %q, want example.com", bw) + } + if len(starters) != 1 || starters[0].Name != "ns1.example.com" { + t.Errorf("starters = %v", starters) } } -func TestInfoCacheChainedGlue(t *testing.T) { +func TestGetStartServersRootHints(t *testing.T) { + c := NewInfoCache(nil) + c.AddHints("", []StartServer{ + {Name: "a.root-servers.net", IPs: []string{"198.41.0.4", "2001:503:ba3e::2:30"}}, + {Name: "b.root-servers.net", IPs: []string{"170.247.170.2"}}, + }) + + starters, bw, err := c.GetStartServers("anything.example.org") + if err != nil { + t.Fatalf("GetStartServers: %v", err) + } + if bw != "" { + t.Errorf("root bailiwick should be \"\", got %q", bw) + } + if len(starters) != 2 { + t.Fatalf("expected 2 root starters, got %d", len(starters)) + } + // Only the IPv4 address surfaces (IPv4-only transport); the AAAA is + // cached but not returned as a start address. + if len(starters[0].IPs) != 1 || starters[0].IPs[0] != "198.41.0.4" { + t.Errorf("starter[0].IPs = %v, want [198.41.0.4]", starters[0].IPs) + } + if got := c.Get("a.root-servers.net", dns.ClassINET, dns.TypeAAAA); len(got) != 1 { + t.Errorf("AAAA hint should be cached, got %v", got) + } +} + +func TestGetStartServersNoRootHints(t *testing.T) { + c := NewInfoCache(nil) + if _, _, err := c.GetStartServers("example.com"); err == nil { + t.Fatal("expected error with no NS cached anywhere") + } +} + +func TestGetStartServersExactDomainMatch(t *testing.T) { + c := NewInfoCache(nil) + c.Add([]dns.RR{nsRR("example.com", "ns1.example.net")}) + _, bw, err := c.GetStartServers("example.com") + if err != nil { + t.Fatalf("GetStartServers: %v", err) + } + if bw != "example.com" { + t.Errorf("newbailiwick = %q, want example.com", bw) + } +} + +func TestGetStartServersUsesBranchCache(t *testing.T) { parent := NewInfoCache(nil) - parent.StoreGlue("ns1.example.com.", []net.IP{net.ParseIP("1.2.3.4")}) - + parent.AddHints("", []StartServer{{Name: "a.root-servers.net", IPs: []string{"198.41.0.4"}}}) child := parent.Child() - result := child.LookupGlue("ns1.example.com.") - if len(result) != 1 { - t.Fatalf("expected 1 glue from parent, got %d", len(result)) - } - if child.GlueCount() != 0 { - t.Errorf("child GlueCount = %d, want 0", child.GlueCount()) - } -} + child.Add([]dns.RR{nsRR("com", "a.gtld-servers.net")}) -func TestInfoCacheGlueLookupMiss(t *testing.T) { - c := NewInfoCache(nil) - result := c.LookupGlue("nonexistent.example.com.") - if result != nil { - t.Errorf("expected nil for miss, got %v", result) + _, bw, err := child.GetStartServers("www.example.com") + if err != nil { + t.Fatalf("GetStartServers: %v", err) + } + if bw != "com" { + t.Errorf("newbailiwick = %q, want com (child cache hit)", bw) + } + + // A sibling branch must not see the child's records. + sibling := parent.Child() + _, bw, err = sibling.GetStartServers("www.example.com") + if err != nil { + t.Fatalf("GetStartServers: %v", err) + } + if bw != "" { + t.Errorf("sibling newbailiwick = %q, want \"\" (root only)", bw) } } func TestInfoCacheConcurrentAccess(t *testing.T) { c := NewInfoCache(nil) var wg sync.WaitGroup - for i := 0; i < 100; i++ { wg.Add(1) go func(i int) { defer wg.Done() - name := strings.ToLower(dns.Fqdn(fmt.Sprintf("ns%d.example.com.", i))) - c.StoreNS("example.com.", []string{name}) - c.StoreGlue(name, []net.IP{net.ParseIP(fmt.Sprintf("1.2.3.%d", i%256))}) - _ = c.LookupNS("example.com.") - _ = c.LookupGlue(name) + name := fmt.Sprintf("ns%d.example.com", i) + c.Add([]dns.RR{nsRR("example.com", name), aRR(name, fmt.Sprintf("1.2.3.%d", i%256))}) + _, _, _ = c.GetStartServers("www.example.com") + _ = c.Get(name, dns.ClassINET, dns.TypeA) }(i) } wg.Wait() } - -func TestInfoCacheNilParent(t *testing.T) { - c := NewInfoCache(nil) - if c.LookupNS("com.") != nil { - t.Error("root cache should return nil for miss") - } - if c.LookupGlue("ns.example.com.") != nil { - t.Error("root cache should return nil for glue miss") - } -} - -func TestNormalize(t *testing.T) { - tests := []struct { - input string - want string - }{ - {"example.com", "example.com."}, - {"Example.COM.", "example.com."}, - {"EXAMPLE.COM", "example.com."}, - } - - for _, tt := range tests { - t.Run(tt.input, func(t *testing.T) { - got := normalize(tt.input) - if got != tt.want { - t.Errorf("normalize(%q) = %q, want %q", tt.input, got, tt.want) - } - }) - } -} diff --git a/internal/traverse/coverage_test.go b/internal/traverse/coverage_test.go deleted file mode 100644 index 63d7817..0000000 --- a/internal/traverse/coverage_test.go +++ /dev/null @@ -1,727 +0,0 @@ -package traverse - -import ( - "context" - "net" - "testing" - "time" - - dnsinternal "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" - "github.com/miekg/dns" -) - -func TestSetHooks(t *testing.T) { - tr := NewTraverser(nil) - hooks := &TraverserHooks{ - OnEvent: func(event TraversalEvent) {}, - } - tr.SetHooks(hooks) - if tr.config.Hooks != hooks { - t.Error("SetHooks should set config.Hooks") - } - - // SetHooks on nil config traverser (initializes config) - tr2 := &Traverser{} - tr2.SetHooks(hooks) - if tr2.config == nil || tr2.config.Hooks != hooks { - t.Error("SetHooks should initialize config when nil") - } -} - -func TestNewAQuery(t *testing.T) { - msg := newAQuery("example.com.") - if msg == nil { - t.Fatal("newAQuery returned nil") - } - if !msg.RecursionDesired { - t.Error("expected RD=true in newAQuery") - } - if len(msg.Question) == 0 { - t.Fatal("expected question in newAQuery") - } - if msg.Question[0].Qtype != dns.TypeA { - t.Errorf("expected TypeA, got %d", msg.Question[0].Qtype) - } -} - -func TestResolveGlueViaSystemCacheHit(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - cache := NewInfoCache(nil) - expected := []net.IP{net.ParseIP("1.2.3.4")} - cache.StoreGlue("ns1.example.com.", expected) - - ctx := context.Background() - addrs := tr.resolveGlueViaSystem(ctx, "ns1.example.com.", cache) - if len(addrs) == 0 { - t.Error("expected addresses from cache hit") - } -} - -func TestResolveGlueViaSystemExpiredContext(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - - // Expired context → remaining <= 0 → returns nil immediately - ctx, cancel := context.WithDeadline(context.Background(), time.Now().Add(-time.Second)) - defer cancel() - - addrs := tr.resolveGlueViaSystem(ctx, "ns1.example.com.", nil) - if len(addrs) != 0 { - t.Errorf("expected nil from expired context, got %v", addrs) - } -} - -func TestResolveGlueViaSystemTimeout(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - - // Very short timeout will fail the DNS query to 127.0.0.1:53 - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond) - defer cancel() - time.Sleep(15 * time.Millisecond) // ensure it's expired - - addrs := tr.resolveGlueViaSystem(ctx, "ns1.example.com.", nil) - // May return nil (timeout) or addresses (if local resolver responds instantly) - t.Logf("resolveGlueViaSystem returned %d addresses", len(addrs)) -} - -func TestEnsureRDFalseWithExchange(t *testing.T) { - rdFalseMsg := new(dns.Msg) - rdFalseMsg.SetReply(new(dns.Msg)) - rdFalseMsg.RecursionDesired = false - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return rdFalseMsg.Copy(), nil - }) - - rdTrueMsg := new(dns.Msg) - rdTrueMsg.SetReply(new(dns.Msg)) - rdTrueMsg.RecursionDesired = true - - result := tr.ensureRDFalse(rdTrueMsg, net.ParseIP("198.41.0.4"), "example.com.", dnsinternal.TypeA, nil) - if result == nil { - t.Fatal("ensureRDFalse with exchange should return non-nil") - } -} - -func TestEnsureRDFalseNilMsg(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - result := tr.ensureRDFalse(nil, net.ParseIP("1.2.3.4"), "example.com.", dnsinternal.TypeA, nil) - if result != nil { - t.Error("ensureRDFalse(nil) should return nil") - } -} - -func TestEnsureRDFalseRDAlreadyFalse(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - - msg := new(dns.Msg) - msg.RecursionDesired = false - result := tr.ensureRDFalse(msg, net.ParseIP("1.2.3.4"), "example.com.", dnsinternal.TypeA, nil) - if result != msg { - t.Error("ensureRDFalse should return same msg when RD=false") - } -} - -func TestEnsureRDFalseNoExchange(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - // No exchange set - - msg := new(dns.Msg) - msg.RecursionDesired = true - result := tr.ensureRDFalse(msg, net.ParseIP("1.2.3.4"), "example.com.", dnsinternal.TypeA, nil) - if result == nil { - t.Fatal("ensureRDFalse without exchange should return msg with RD cleared") - } - if result.RecursionDesired { - t.Error("expected RD=false after ensureRDFalse without exchange") - } -} - -func TestResolveNSFromCache(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - cache := NewInfoCache(nil) - expected := []net.IP{net.ParseIP("1.2.3.4")} - cache.StoreGlue("ns1.example.com.", expected) - - ctx := context.Background() - addrs, err := tr.ResolveNS(ctx, "ns1.example.com.", cache, nil, 0) - if err != nil { - t.Fatalf("ResolveNS cache hit: %v", err) - } - if len(addrs) == 0 { - t.Error("expected addresses from cache") - } -} - -func TestResolveNSCircularReferral(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - visited := map[string]bool{"ns1.example.com.": true} - ctx := context.Background() - _, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, visited, 0) - if err == nil { - t.Fatal("expected circular referral error") - } - var circErr *CircularReferralError - if _, ok := err.(*CircularReferralError); !ok { - t.Errorf("expected CircularReferralError, got %T: %v", err, err) - } - _ = circErr -} - -func TestResolveNSMaxDepth(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - ctx := context.Background() - _, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, DefaultMaxDepth+1) - if err == nil { - t.Fatal("expected max depth error") - } - if _, ok := err.(*UnresolvableNameserverError); !ok { - t.Errorf("expected UnresolvableNameserverError, got %T: %v", err, err) - } -} - -func TestResolveNSWithAnswer(t *testing.T) { - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerMsg.Copy(), nil - }) - - ctx := context.Background() - addrs, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, 0) - if err != nil { - t.Fatalf("ResolveNS with answer: %v", err) - } - if len(addrs) == 0 { - t.Fatal("expected addresses from NS resolution") - } -} - -func TestResolveNSNXDOMAIN(t *testing.T) { - nxMsg := new(dns.Msg) - nxMsg.Rcode = dns.RcodeNameError - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nxMsg.Copy(), nil - }) - - ctx := context.Background() - _, err := tr.ResolveNS(ctx, "nonexistent.invalid.", nil, nil, 0) - if err == nil { - t.Fatal("expected error for NXDOMAIN NS resolution") - } -} - -func TestResolveNSContextCancellation(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - // Keep returning referrals to keep the loop going - refMsg := new(dns.Msg) - refMsg.Rcode = dns.RcodeSuccess - refMsg.Ns = append(refMsg.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, - Ns: "ns.example.com.", - }) - refMsg.Extra = append(refMsg.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET}, - A: net.ParseIP("1.2.3.4"), - }) - return refMsg, nil - }) - - ctx, cancel := context.WithCancel(context.Background()) - cancel() // Cancel immediately - - _, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, 0) - if err == nil { - t.Fatal("expected error on cancelled context") - } -} - -func TestDiscoverRootsWithRootAddrs(t *testing.T) { - expected := []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")} - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: expected, - }) - - ctx := context.Background() - addrs, err := tr.discoverRoots(ctx) - if err != nil { - t.Fatalf("discoverRoots with RootAddrs: %v", err) - } - if len(addrs) != len(expected) { - t.Errorf("expected %d addresses, got %d", len(expected), len(addrs)) - } -} - -func TestDiscoverRootsFromSystem(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - // No RootAddrs - will call dns.DiscoverRoots - }) - - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() - - addrs, err := tr.discoverRoots(ctx) - if err != nil { - t.Logf("discoverRoots without RootAddrs error (may skip): %v", err) - t.Skip() - } - if len(addrs) == 0 { - t.Error("expected at least one root address") - } -} - -func TestTraverserSetHooksAndTraverse(t *testing.T) { - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerResp.Copy(), nil - }) - - var events []TraversalEvent - tr.SetHooks(&TraverserHooks{ - OnEvent: func(e TraversalEvent) { - events = append(events, e) - }, - }) - - ctx := context.Background() - _, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("Traverse: %v", err) - } - if len(events) == 0 { - t.Error("expected events from hooks") - } -} - -func TestProcessReferralNoAddresses(t *testing.T) { - // Scenario: a referral without addresses. resolveGlueViaSystem fails (expired ctx), - // then ResolveNS is tried via the mock exchange. - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - finalAnswerMsg := new(dns.Msg) - finalAnswerMsg.SetReply(new(dns.Msg)) - finalAnswerMsg.Answer = append(finalAnswerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("5.6.7.8"), - }) - - callCount := 0 - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - callCount++ - q := msg.Question[0] - if q.Qtype == dns.TypeA && q.Name == "ns1.example.com." { - return answerMsg.Copy(), nil - } - return finalAnswerMsg.Copy(), nil - }) - - // Create a referral with no addresses (the NS name needs to be resolved) - ref := NewReferral("example.com.", dnsinternal.TypeA, "ns1.example.com.", 1, 1.0, nil) - // Do NOT set addresses - this exercises processReferral's no-address path - - cache := NewInfoCache(nil) - // Use expired context for resolveGlueViaSystem so it returns nil fast - bgCtx := context.Background() - resp := tr.processReferral(bgCtx, ref, cache) - // Result may vary depending on whether 127.0.0.1:53 is available, - // but the function should not panic. - t.Logf("processReferral result type: %v", resp.Type) -} - -func TestReferralResolveAlreadyHasAddresses(t *testing.T) { - ref := NewReferral("example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - ref.Addresses = []net.IP{net.ParseIP("1.2.3.4")} - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - ctx := context.Background() - err := ref.Resolve(ctx, tr, nil, nil, 0) - if err != nil { - t.Fatalf("Resolve with existing addresses: %v", err) - } - if ref.State != StateResolved { - t.Errorf("expected StateResolved, got %v", ref.State) - } -} - -func TestReferralResolveCacheHit(t *testing.T) { - ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - cache := NewInfoCache(nil) - cache.StoreGlue("ns1.example.com.", []net.IP{net.ParseIP("1.2.3.4")}) - - ctx := context.Background() - err := ref.Resolve(ctx, tr, cache, nil, 0) - if err != nil { - t.Fatalf("Resolve cache hit: %v", err) - } - if ref.State != StateResolved { - t.Errorf("expected StateResolved, got %v", ref.State) - } -} - -func TestReferralResolveCircular(t *testing.T) { - ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - visited := map[string]bool{"ns1.example.com.": true} - ctx := context.Background() - err := ref.Resolve(ctx, tr, nil, visited, 0) - if err == nil { - t.Fatal("expected circular referral error") - } - if _, ok := err.(*CircularReferralError); !ok { - t.Errorf("expected CircularReferralError, got %T: %v", err, err) - } -} - -func TestReferralResolveMaxDepth(t *testing.T) { - ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - ctx := context.Background() - err := ref.Resolve(ctx, tr, nil, nil, DefaultMaxDepth+1) - if err == nil { - t.Fatal("expected max depth error") - } - if _, ok := err.(*UnresolvableNameserverError); !ok { - t.Errorf("expected UnresolvableNameserverError, got %T: %v", err, err) - } -} - -func TestReferralResolveWithAnswer(t *testing.T) { - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerMsg.Copy(), nil - }) - - ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - ctx := context.Background() - err := ref.Resolve(ctx, tr, nil, nil, 0) - if err != nil { - t.Fatalf("Resolve: %v", err) - } - if ref.State != StateResolved { - t.Errorf("expected StateResolved, got %v", ref.State) - } - if len(ref.Addresses) == 0 { - t.Error("expected addresses after resolution") - } -} - -func TestReferralResolveNXDOMAIN(t *testing.T) { - nxMsg := new(dns.Msg) - nxMsg.Rcode = dns.RcodeNameError - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nxMsg.Copy(), nil - }) - - ref := NewReferral("nonexistent.invalid.", dnsinternal.TypeA, ".", 0, 1.0, nil) - ctx := context.Background() - err := ref.Resolve(ctx, tr, nil, nil, 0) - if err == nil { - t.Fatal("expected error for NXDOMAIN") - } - if _, ok := err.(*UnresolvableNameserverError); !ok { - t.Errorf("expected UnresolvableNameserverError, got %T: %v", err, err) - } -} - -func TestReferralResolveContextCancellation(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - refMsg := new(dns.Msg) - refMsg.Rcode = dns.RcodeSuccess - refMsg.Ns = append(refMsg.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, - Ns: "ns.example.com.", - }) - refMsg.Extra = append(refMsg.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dns.TypeA}, - A: net.ParseIP("1.2.3.4"), - }) - return refMsg, nil - }) - - ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - ctx, cancel := context.WithCancel(context.Background()) - cancel() // Cancel immediately - - err := ref.Resolve(ctx, tr, nil, nil, 0) - if err == nil { - t.Fatal("expected error on cancelled context") - } -} - -func TestReferralResolveReferralPath(t *testing.T) { - // Test Resolve when it gets a referral response that pushes to stack - referralMsg := new(dns.Msg) - referralMsg.Rcode = dns.RcodeSuccess - referralMsg.Ns = append(referralMsg.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS, Class: dns.ClassINET}, - Ns: "ns1.example.com.", - }) - referralMsg.Extra = append(referralMsg.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET}, - A: net.ParseIP("1.2.3.4"), - }) - - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("5.6.7.8"), - }) - - callCount := 0 - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - callCount++ - if callCount <= 1 { - return referralMsg.Copy(), nil - } - return answerMsg.Copy(), nil - }) - - ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - ctx := context.Background() - err := ref.Resolve(ctx, tr, nil, nil, 0) - // May succeed or exhaust depending on referral loop - t.Logf("Resolve referral path: err=%v, state=%v", err, ref.State) -} - -func TestResolutionStateStringUnknown(t *testing.T) { - // Cover the default case of ResolutionState.String() - unknown := ResolutionState(99) - s := unknown.String() - if s != "unknown" { - t.Errorf("expected 'unknown' for invalid ResolutionState, got %q", s) - } -} - -func TestTraverserNonFastMode(t *testing.T) { - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - Fast: false, // Non-fast mode - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("Traverse non-fast: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results") - } -} - -func TestIterativeQueryWithExchangeUsesConfig(t *testing.T) { - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - QueryConfig: dnsinternal.DefaultQueryConfig(), - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerMsg.Copy(), nil - }) - - ctx := context.Background() - msg, err := tr.iterativeQueryWithExchange(ctx, net.ParseIP("198.41.0.4"), "example.com.", dnsinternal.TypeA) - if err != nil { - t.Fatalf("iterativeQueryWithExchange with config: %v", err) - } - if msg == nil { - t.Fatal("expected non-nil response") - } -} - -func TestTraverserReferralWithHooks(t *testing.T) { - // Tests that hooks are called with IsResolve=true during ResolveNS sub-traversal - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerMsg.Copy(), nil - }) - - var resolveEvents, progressEvents int - tr.SetHooks(&TraverserHooks{ - OnEvent: func(e TraversalEvent) { - if e.IsResolve { - resolveEvents++ - } else { - progressEvents++ - } - }, - }) - - // Test directly via ResolveNS with hooks - ctx := context.Background() - addrs, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, 0) - if err != nil { - t.Fatalf("ResolveNS: %v", err) - } - _ = addrs - t.Logf("resolveEvents=%d progressEvents=%d", resolveEvents, progressEvents) -} diff --git a/internal/traverse/decoded_query.go b/internal/traverse/decoded_query.go new file mode 100644 index 0000000..9633ba1 --- /dev/null +++ b/internal/traverse/decoded_query.go @@ -0,0 +1,259 @@ +package traverse + +import ( + "fmt" + "strings" + + miekgdns "github.com/miekg/dns" +) + +// Status is the classification of one query outcome. The first eight values +// come from decoded_query.rb / response.rb; noglue and loop are synthesised by +// the resolve phase without sending a query. +type Status string + +const ( + StatusAnswered Status = "answered" + StatusNoData Status = "nodata" + StatusReferral Status = "referral" + StatusRestart Status = "restart" + StatusReferralLame Status = "referral_lame" + StatusError Status = "error" + StatusException Status = "exception" + StatusCNAMELoop Status = "cname_loop" + StatusNoGlue Status = "noglue" + StatusLoop Status = "loop" +) + +// DecodedQuery classifies one DNS response (or network failure) against the +// query that produced it, mirroring decoded_query.rb. Names are canonical +// (lowercase, no trailing dot); Bailiwick "" means root. +type DecodedQuery struct { + Msg *miekgdns.Msg + Err error + + Qname string + Qclass uint16 + Qtype uint16 + IP string + Bailiwick string + + Status Status + Endname string + // ChainTargets lists every CNAME target the in-message chain went + // through (including the final unfollowed target when the chain leaves + // the bailiwick); used for cross-response restart loop detection. + ChainTargets []string + + CacheableGood []miekgdns.RR + CacheableBad []miekgdns.RR + + AuthNS []miekgdns.RR + AuthSOA []miekgdns.RR + AuthOther []miekgdns.RR + + Answers []miekgdns.RR + AuthorityNames []string + + ErrorMessage string + ExceptionMessage string + Warnings []string +} + +// NewDecodedQuery decodes and classifies a response. Pass err non-nil for a +// network-level failure (dnstraverse's "exception"); msg is ignored then. +func NewDecodedQuery(msg *miekgdns.Msg, err error, qname string, qclass, qtype uint16, ip, bailiwick string) *DecodedQuery { + dq := &DecodedQuery{ + Msg: msg, + Err: err, + Qname: canonicalName(qname), + Qclass: qclass, + Qtype: qtype, + IP: ip, + Bailiwick: canonicalName(bailiwick), + } + dq.process() + return dq +} + +func (dq *DecodedQuery) WarningsAdd(warnings ...string) { + dq.Warnings = append(dq.Warnings, warnings...) +} + +// process implements the classification order of decoded_query.rb#process +// exactly (the 7 steps in the design doc). +func (dq *DecodedQuery) process() { + if dq.Err == nil && dq.Msg == nil { + dq.Err = fmt.Errorf("nil DNS response") + } + if dq.Err != nil { + dq.Status = StatusException + dq.ExceptionMessage = dq.Err.Error() + return + } + dq.AuthNS, dq.AuthSOA, dq.AuthOther = msgAuthority(dq.Msg) + dq.CacheableGood, dq.CacheableBad = msgCacheable(dq.Msg, dq.Bailiwick) + endname, targets, ok := msgFollowCNAMEs(dq.Msg, dq.Qname, dq.Qtype, dq.Bailiwick) + if !ok { + dq.Status = StatusCNAMELoop + return + } + dq.Endname = endname + dq.ChainTargets = targets + if dq.Msg.Rcode != miekgdns.RcodeSuccess { + dq.Status = StatusError + dq.ErrorMessage = rcodeErrorMessage(dq.Msg.Rcode) + return + } + if answers := msgAnswers(dq.Msg, dq.Endname, dq.Qclass, dq.Qtype); len(answers) > 0 { + dq.Answers = answers + dq.Status = StatusAnswered + return + } + if dq.Endname != dq.Qname { + dq.Status = StatusRestart + return + } + if len(dq.AuthSOA) > 0 || len(dq.AuthNS) == 0 { + dq.Status = StatusNoData + return + } + dq.Status = StatusReferral + for _, rr := range dq.AuthNS { + if ns, ok := rr.(*miekgdns.NS); ok { + dq.AuthorityNames = append(dq.AuthorityNames, canonicalName(ns.Ns)) + } + } +} + +// rcodeErrorMessage renders the exact error strings of decoded_query.rb +// process_error ("Format error" deliberately fixes the Ruby "Formate" typo — +// documented deviation). +func rcodeErrorMessage(rcode int) string { + switch rcode { + case miekgdns.RcodeFormatError: + return "Format error (FORMERR)" + case miekgdns.RcodeServerFailure: + return "Server failure (SERVFAIL)" + case miekgdns.RcodeNameError: + return "No such domain (NXDOMAIN)" + case miekgdns.RcodeNotImplemented: + return "Not implemented (NOTIMP)" + case miekgdns.RcodeRefused: + return "Refused" + default: + if s, ok := miekgdns.RcodeToString[rcode]; ok { + return s + } + return fmt.Sprintf("RCODE%d", rcode) + } +} + +// insideBailiwick reports whether name is at or below the bailiwick zone: +// bailiwick "" (root), equal fold, or name ends with "."+bailiwick. +func insideBailiwick(name, bailiwick string) bool { + bw := canonicalName(bailiwick) + if bw == "" { + return true + } + n := canonicalName(name) + return n == bw || strings.HasSuffix(n, "."+bw) +} + +// msgAnswers returns the answer-section records matching qname/qclass/qtype +// (message_utility.rb msg_answers?). qtype ANY matches every type. +func msgAnswers(msg *miekgdns.Msg, qname string, qclass, qtype uint16) []miekgdns.RR { + name := canonicalName(qname) + any := qtype == miekgdns.TypeANY + var out []miekgdns.RR + for _, rr := range msg.Answer { + h := rr.Header() + if canonicalName(h.Name) == name && h.Class == qclass && (any || h.Rrtype == qtype) { + out = append(out, rr) + } + } + return out +} + +// msgAuthority partitions the authority section into IN NS, IN SOA and other +// records (message_utility.rb msg_authority). +func msgAuthority(msg *miekgdns.Msg) (ns, soa, other []miekgdns.RR) { + for _, rr := range msg.Ns { + h := rr.Header() + switch { + case h.Rrtype == miekgdns.TypeNS && h.Class == miekgdns.ClassINET: + ns = append(ns, rr) + case h.Rrtype == miekgdns.TypeSOA && h.Class == miekgdns.ClassINET: + soa = append(soa, rr) + default: + other = append(other, rr) + } + } + return ns, soa, other +} + +// msgCacheable partitions ALL sections (answer, authority, additional — in +// that order) into in-bailiwick records worth caching and out-of-bailiwick +// records to discard. OPT pseudo-records are dropped entirely. +func msgCacheable(msg *miekgdns.Msg, bailiwick string) (good, bad []miekgdns.RR) { + for _, section := range [][]miekgdns.RR{msg.Answer, msg.Ns, msg.Extra} { + for _, rr := range section { + if rr.Header().Rrtype == miekgdns.TypeOPT { + continue + } + if insideBailiwick(rr.Header().Name, bailiwick) { + good = append(good, rr) + } else { + bad = append(bad, rr) + } + } + } + return good, bad +} + +// msgFollowCNAMEs follows a CNAME chain within one message and returns the +// final name plus every target passed through (message_utility.rb +// msg_follow_cnames). Following stops — the target is returned unfollowed — +// as soon as the CURRENT owner name is not strictly below the bailiwick +// (Ruby tests `name !~ /\.#{bailiwick}$/i`, so an owner exactly equal to the +// bailiwick also stops the chain). An in-message loop returns ok=false +// (cname_loop). +func msgFollowCNAMEs(msg *miekgdns.Msg, qname string, qtype uint16, bailiwick string) (endname string, targets []string, ok bool) { + name := canonicalName(qname) + bw := canonicalName(bailiwick) + seen := make(map[string]bool) + for { + seen[name] = true + if len(msgAnswers(msg, name, miekgdns.ClassINET, qtype)) > 0 { + return name, targets, true + } + cnames := msgAnswers(msg, name, miekgdns.ClassINET, miekgdns.TypeCNAME) + if len(cnames) == 0 { + return name, targets, true + } + cname, isCNAME := cnames[0].(*miekgdns.CNAME) + if !isCNAME { + return name, targets, true + } + target := canonicalName(cname.Target) + targets = append(targets, target) + if bw != "" && !strings.HasSuffix(name, "."+bw) { + return target, targets, true + } + name = target + if seen[name] { + return "", targets, false + } + } +} + +// isLameReferral implements the response.rb lame rule: a referral is lame +// unless the current bailiwick is root ("") or the new zone is STRICTLY +// deeper than the current bailiwick (equal or sideways zones are lame). +func isLameReferral(bailiwick, newBailiwick string) bool { + bw := canonicalName(bailiwick) + if bw == "" { + return false + } + return !strings.HasSuffix(canonicalName(newBailiwick), "."+bw) +} diff --git a/internal/traverse/decoded_query_test.go b/internal/traverse/decoded_query_test.go new file mode 100644 index 0000000..2bbb696 --- /dev/null +++ b/internal/traverse/decoded_query_test.go @@ -0,0 +1,359 @@ +package traverse + +import ( + "errors" + "testing" + + "github.com/miekg/dns" +) + +func cnameRR(owner, target string) dns.RR { + return &dns.CNAME{ + Hdr: dns.RR_Header{Name: dns.Fqdn(owner), Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, + Target: dns.Fqdn(target), + } +} + +func soaRR(zone string) dns.RR { + return &dns.SOA{ + Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeSOA, Class: dns.ClassINET}, + Ns: dns.Fqdn("ns1." + zone), + Mbox: dns.Fqdn("hostmaster." + zone), + Serial: 1, + Refresh: 3600, Retry: 600, Expire: 86400, Minttl: 300, + } +} + +func newMsg(qname string, qtype uint16, rcode int) *dns.Msg { + m := new(dns.Msg) + m.SetQuestion(dns.Fqdn(qname), qtype) + m.Response = true + m.Rcode = rcode + return m +} + +func decode(msg *dns.Msg, qname string, qtype uint16, bailiwick string) *DecodedQuery { + return NewDecodedQuery(msg, nil, qname, dns.ClassINET, qtype, "192.0.2.1", bailiwick) +} + +func TestDecodeException(t *testing.T) { + dq := NewDecodedQuery(nil, errors.New("network timeout"), "example.com", dns.ClassINET, dns.TypeA, "192.0.2.1", "com") + if dq.Status != StatusException { + t.Fatalf("status = %s, want exception", dq.Status) + } + if dq.ExceptionMessage != "network timeout" { + t.Errorf("exception message = %q", dq.ExceptionMessage) + } +} + +func TestDecodeNilMessageIsException(t *testing.T) { + dq := NewDecodedQuery(nil, nil, "example.com", dns.ClassINET, dns.TypeA, "192.0.2.1", "com") + if dq.Status != StatusException { + t.Fatalf("status = %s, want exception", dq.Status) + } +} + +func TestDecodeErrorMessages(t *testing.T) { + tests := []struct { + rcode int + want string + }{ + {dns.RcodeFormatError, "Format error (FORMERR)"}, + {dns.RcodeServerFailure, "Server failure (SERVFAIL)"}, + {dns.RcodeNameError, "No such domain (NXDOMAIN)"}, + {dns.RcodeNotImplemented, "Not implemented (NOTIMP)"}, + {dns.RcodeRefused, "Refused"}, + {dns.RcodeYXDomain, "YXDOMAIN"}, + } + for _, tt := range tests { + msg := newMsg("example.com", dns.TypeA, tt.rcode) + dq := decode(msg, "example.com", dns.TypeA, "com") + if dq.Status != StatusError { + t.Errorf("rcode %d: status = %s, want error", tt.rcode, dq.Status) + } + if dq.ErrorMessage != tt.want { + t.Errorf("rcode %d: message = %q, want %q", tt.rcode, dq.ErrorMessage, tt.want) + } + } +} + +func TestDecodeAnswered(t *testing.T) { + msg := newMsg("example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, aRR("example.com", "93.184.216.34")) + dq := decode(msg, "example.com", dns.TypeA, "example.com") + if dq.Status != StatusAnswered { + t.Fatalf("status = %s, want answered", dq.Status) + } + if len(dq.Answers) != 1 { + t.Errorf("answers = %v", dq.Answers) + } + if dq.Endname != "example.com" { + t.Errorf("endname = %q", dq.Endname) + } +} + +func TestDecodeAnsweredViaCNAMEChain(t *testing.T) { + // In-bailiwick chain ends at a name that has the A answer. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, + cnameRR("www.example.com", "web.example.com"), + aRR("web.example.com", "93.184.216.34"), + ) + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + if dq.Status != StatusAnswered { + t.Fatalf("status = %s, want answered", dq.Status) + } + if dq.Endname != "web.example.com" { + t.Errorf("endname = %q, want web.example.com", dq.Endname) + } +} + +func TestDecodeRestartOnOutOfBailiwickCNAME(t *testing.T) { + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, cnameRR("www.example.com", "cdn.example.org")) + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + if dq.Status != StatusRestart { + t.Fatalf("status = %s, want restart", dq.Status) + } + if dq.Endname != "cdn.example.org" { + t.Errorf("endname = %q", dq.Endname) + } +} + +func TestDecodeCNAMELoop(t *testing.T) { + msg := newMsg("a.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, + cnameRR("a.example.com", "b.example.com"), + cnameRR("b.example.com", "a.example.com"), + ) + dq := decode(msg, "a.example.com", dns.TypeA, "example.com") + if dq.Status != StatusCNAMELoop { + t.Fatalf("status = %s, want cname_loop", dq.Status) + } +} + +func TestDecodeCNAMESelfLoop(t *testing.T) { + msg := newMsg("a.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, cnameRR("a.example.com", "A.EXAMPLE.COM")) + dq := decode(msg, "a.example.com", dns.TypeA, "example.com") + if dq.Status != StatusCNAMELoop { + t.Fatalf("status = %s, want cname_loop (case-insensitive)", dq.Status) + } +} + +func TestDecodeCNAMEChainStopsAtOutOfBailiwickOwner(t *testing.T) { + // Ruby stops following once the CURRENT owner leaves the bailiwick, so a + // two-hop loop through an out-of-bailiwick owner is NOT cname_loop: the + // unfollowed target equals the qname again, leaving endname == qname and + // an empty authority — nodata. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, + cnameRR("www.example.com", "a.example.org"), + cnameRR("a.example.org", "www.example.com"), + ) + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + if dq.Status != StatusNoData { + t.Fatalf("status = %s, want nodata", dq.Status) + } + if dq.Endname != "www.example.com" { + t.Errorf("endname = %q, want www.example.com", dq.Endname) + } +} + +func TestDecodeCNAMEOwnerEqualToBailiwickStopsChain(t *testing.T) { + // Owner exactly equal to the bailiwick is NOT strictly inside it, so the + // chain stops after one hop even though another CNAME exists. + msg := newMsg("example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, + cnameRR("example.com", "a.example.com"), + cnameRR("a.example.com", "b.example.com"), + ) + dq := decode(msg, "example.com", dns.TypeA, "example.com") + if dq.Status != StatusRestart { + t.Fatalf("status = %s, want restart", dq.Status) + } + if dq.Endname != "a.example.com" { + t.Errorf("endname = %q, want a.example.com (unfollowed target)", dq.Endname) + } +} + +func TestDecodeQtypeCNAMEIsAnswered(t *testing.T) { + // qtype=CNAME: the CNAME record IS the answer; the chain is never followed. + msg := newMsg("www.example.com", dns.TypeCNAME, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, + cnameRR("www.example.com", "web.example.com"), + cnameRR("web.example.com", "www.example.com"), + ) + dq := decode(msg, "www.example.com", dns.TypeCNAME, "example.com") + if dq.Status != StatusAnswered { + t.Fatalf("status = %s, want answered", dq.Status) + } + if dq.Endname != "www.example.com" { + t.Errorf("endname = %q", dq.Endname) + } +} + +func TestDecodeQtypeANYMatchesAnyAnswer(t *testing.T) { + msg := newMsg("example.com", dns.TypeANY, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, cnameRR("example.com", "elsewhere.example.net")) + dq := decode(msg, "example.com", dns.TypeANY, "example.com") + if dq.Status != StatusAnswered { + t.Fatalf("status = %s, want answered (ANY matches CNAME)", dq.Status) + } +} + +func TestDecodeNoDataWithSOA(t *testing.T) { + msg := newMsg("example.com", dns.TypeMX, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, soaRR("example.com")) + dq := decode(msg, "example.com", dns.TypeMX, "example.com") + if dq.Status != StatusNoData { + t.Fatalf("status = %s, want nodata", dq.Status) + } +} + +func TestDecodeNoDataEmptyAuthority(t *testing.T) { + msg := newMsg("example.com", dns.TypeMX, dns.RcodeSuccess) + dq := decode(msg, "example.com", dns.TypeMX, "example.com") + if dq.Status != StatusNoData { + t.Fatalf("status = %s, want nodata", dq.Status) + } +} + +func TestDecodeNoDataSOAWinsOverNS(t *testing.T) { + // SOA + NS in authority is a negative answer, not a referral. + msg := newMsg("example.com", dns.TypeMX, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, soaRR("example.com"), nsRR("example.com", "ns1.example.com")) + dq := decode(msg, "example.com", dns.TypeMX, "example.com") + if dq.Status != StatusNoData { + t.Fatalf("status = %s, want nodata", dq.Status) + } +} + +func TestDecodeReferral(t *testing.T) { + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, + nsRR("example.com", "NS1.Example.COM"), + nsRR("example.com", "ns2.example.net"), + ) + msg.Extra = append(msg.Extra, aRR("ns1.example.com", "1.2.3.4")) + dq := decode(msg, "www.example.com", dns.TypeA, "com") + if dq.Status != StatusReferral { + t.Fatalf("status = %s, want referral", dq.Status) + } + if len(dq.AuthorityNames) != 2 || dq.AuthorityNames[0] != "ns1.example.com" || dq.AuthorityNames[1] != "ns2.example.net" { + t.Errorf("authority names = %v", dq.AuthorityNames) + } +} + +func TestDecodeErrorBeatsAnswer(t *testing.T) { + // rcode is checked before answers (step 3 before step 4). + msg := newMsg("example.com", dns.TypeA, dns.RcodeServerFailure) + msg.Answer = append(msg.Answer, aRR("example.com", "1.2.3.4")) + dq := decode(msg, "example.com", dns.TypeA, "com") + if dq.Status != StatusError { + t.Fatalf("status = %s, want error", dq.Status) + } +} + +func TestDecodeCNAMEFollowedIntoNXDOMAIN(t *testing.T) { + // CNAME followed first (step 2), then rcode (step 3): NXDOMAIN after an + // in-message CNAME is still an error, but the loop check ran first. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeNameError) + msg.Answer = append(msg.Answer, cnameRR("www.example.com", "gone.example.com")) + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + if dq.Status != StatusError { + t.Fatalf("status = %s, want error", dq.Status) + } + if dq.Endname != "gone.example.com" { + t.Errorf("endname = %q", dq.Endname) + } +} + +func TestDecodeCacheablePartition(t *testing.T) { + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, cnameRR("www.example.com", "cdn.example.org")) + msg.Ns = append(msg.Ns, nsRR("example.org", "ns1.example.org")) + msg.Extra = append(msg.Extra, aRR("ns1.example.org", "5.6.7.8")) + opt := new(dns.OPT) + opt.Hdr = dns.RR_Header{Name: ".", Rrtype: dns.TypeOPT} + msg.Extra = append(msg.Extra, opt) + + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + if len(dq.CacheableGood) != 1 { + t.Errorf("good = %v, want just the CNAME", dq.CacheableGood) + } + if len(dq.CacheableBad) != 2 { + t.Errorf("bad = %v, want NS+A for example.org", dq.CacheableBad) + } +} + +func TestInsideBailiwick(t *testing.T) { + tests := []struct { + name, bailiwick string + want bool + }{ + {"anything.example.com", "", true}, // root bailiwick + {"anything.example.com", ".", true}, // root as dot + {"example.com", "example.com", true}, // exact + {"Example.COM", "example.com", true}, // exact, case fold + {"www.example.com", "EXAMPLE.com", true}, // suffix, case fold + {"a.b.example.com", "example.com", true}, // deep suffix + {"badexample.com", "example.com", false}, // label boundary + {"example.org", "example.com", false}, // sideways + {"com", "example.com", false}, // shallower + {"www.example.com.", "example.com", true}, // trailing dot + } + for _, tt := range tests { + if got := insideBailiwick(tt.name, tt.bailiwick); got != tt.want { + t.Errorf("insideBailiwick(%q, %q) = %v, want %v", tt.name, tt.bailiwick, got, tt.want) + } + } +} + +func TestIsLameReferral(t *testing.T) { + tests := []struct { + bailiwick, newBailiwick string + want bool + }{ + {"", "com", false}, // root bailiwick never lame + {"", "", false}, // root to root + {"com", "example.com", false}, // strictly deeper + {"com", "a.b.example.com", false}, // much deeper + {"COM", "example.com", false}, // case fold + {"com", "com", true}, // equal zone is lame + {"com", "", true}, // back to root is lame + {"com", "org", true}, // sideways is lame + {"example.com", "com", true}, // shallower is lame + {"example.com", "badexample.com", true}, // label boundary + {"example.com", "www.example.com", false}, // deeper + } + for _, tt := range tests { + if got := isLameReferral(tt.bailiwick, tt.newBailiwick); got != tt.want { + t.Errorf("isLameReferral(%q, %q) = %v, want %v", tt.bailiwick, tt.newBailiwick, got, tt.want) + } + } +} + +func TestMsgFollowCNAMEsNoChain(t *testing.T) { + msg := newMsg("example.com", dns.TypeA, dns.RcodeSuccess) + end, _, ok := msgFollowCNAMEs(msg, "Example.COM.", dns.TypeA, "com") + if !ok || end != "example.com" { + t.Errorf("end = %q ok=%v, want example.com true", end, ok) + } +} + +func TestMsgFollowCNAMEsRootBailiwickFollowsEverything(t *testing.T) { + msg := newMsg("a.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, + cnameRR("a.example.com", "b.example.org"), + cnameRR("b.example.org", "c.example.net"), + aRR("c.example.net", "1.2.3.4"), + ) + end, targets, ok := msgFollowCNAMEs(msg, "a.example.com", dns.TypeA, "") + if !ok || end != "c.example.net" { + t.Errorf("end = %q ok=%v, want c.example.net true", end, ok) + } + if len(targets) != 2 || targets[0] != "b.example.org" || targets[1] != "c.example.net" { + t.Errorf("chain targets = %v", targets) + } +} diff --git a/internal/traverse/hooks.go b/internal/traverse/hooks.go index fdf0d88..3851f35 100644 --- a/internal/traverse/hooks.go +++ b/internal/traverse/hooks.go @@ -1,16 +1,67 @@ package traverse +// EventStage mirrors the :stage symbols reported by traverser.rb's +// report_progress: new/start/answer/resolve plus the fast-mode and +// multi-childset variants. type EventStage int const ( - EventStart EventStage = iota - EventComplete + // StageNew fires when a referral node is created (before processing). + StageNew EventStage = iota + // StageStart fires when a referral is popped for processing. + StageStart + // StageNewReferralSet fires once per extra childset when more than one + // IP of a server produced children. + StageNewReferralSet + // StageNewFast fires instead of StageNew when fast mode already knows + // this referral will be completed from the memo. + StageNewFast + // StageResolve fires after a resolve subtree's statistics are folded + // into the referral (post-order, Ruby's :calc_resolve marker). + StageResolve + // StageAnswer fires after a referral's statistics are calculated + // (post-order, Ruby's :calc_answer marker). + StageAnswer + // StageAnswerFast fires when fast mode replaced the referral with an + // earlier completed one instead of processing it. + StageAnswerFast ) +func (s EventStage) String() string { + switch s { + case StageNew: + return "new" + case StageStart: + return "start" + case StageNewReferralSet: + return "new_referral_set" + case StageNewFast: + return "new_fast" + case StageResolve: + return "resolve" + case StageAnswer: + return "answer" + case StageAnswerFast: + return "answer_fast" + default: + return "unknown" + } +} + +// TraversalEvent is one progress callback. RefID/Status/IsResolve are +// denormalised from Referral so renderers (CLI, web) need not walk the tree. type TraversalEvent struct { - Stage EventStage - Result TraversalResult + Stage EventStage + Referral *Referral + RefID string + // Status summarises the referral's outcome so far (see + // Referral.OverallStatus); empty before any response arrives. + Status Status + // IsResolve is true for nodes inside a glue-resolution subtree. IsResolve bool + // CompletedEarlier carries the refid of the earlier identical referral + // on fast-mode events (StageNewFast / StageAnswerFast). + CompletedEarlier string } type EventHandler func(TraversalEvent) @@ -19,13 +70,16 @@ type TraverserHooks struct { OnEvent EventHandler } -func (h *TraverserHooks) emit(stage EventStage, result TraversalResult, isResolve bool) { - if h == nil || h.OnEvent == nil { +func (h *TraverserHooks) emit(stage EventStage, r *Referral, completedEarlier string) { + if h == nil || h.OnEvent == nil || r == nil { return } h.OnEvent(TraversalEvent{ - Stage: stage, - Result: result, - IsResolve: isResolve, + Stage: stage, + Referral: r, + RefID: r.RefID, + Status: r.OverallStatus(), + IsResolve: r.IsResolve(), + CompletedEarlier: completedEarlier, }) } diff --git a/internal/traverse/hooks_test.go b/internal/traverse/hooks_test.go index 69007e6..7405375 100644 --- a/internal/traverse/hooks_test.go +++ b/internal/traverse/hooks_test.go @@ -1,52 +1,71 @@ package traverse import ( - "context" - "net" "testing" "github.com/miekg/dns" ) -func TestTraverserHooksEmitEvents(t *testing.T) { - answerResp := func() *dns.Msg { - m := new(dns.Msg) - m.SetReply(new(dns.Msg)) - m.Answer = append(m.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - return m - }() - - var events []TraversalEvent - hooks := &TraverserHooks{ - OnEvent: func(event TraversalEvent) { - events = append(events, event) - }, +func TestEventStageStrings(t *testing.T) { + tests := map[EventStage]string{ + StageNew: "new", + StageStart: "start", + StageNewReferralSet: "new_referral_set", + StageNewFast: "new_fast", + StageResolve: "resolve", + StageAnswer: "answer", + StageAnswerFast: "answer_fast", + EventStage(99): "unknown", } + for stage, want := range tests { + if got := stage.String(); got != want { + t.Errorf("EventStage(%d).String() = %q, want %q", stage, got, want) + } + } +} + +func TestHooksEmitNilSafe(t *testing.T) { + var h *TraverserHooks + h.emit(StageNew, newTestReferral("ns1.example.com", nil), "") // must not panic + (&TraverserHooks{}).emit(StageNew, newTestReferral("ns1.example.com", nil), "") + (&TraverserHooks{OnEvent: func(TraversalEvent) { t.Fatal("emitted for nil referral") }}).emit(StageNew, nil, "") +} + +func TestHooksEventSequenceSimpleAnswer(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "example.com", dns.TypeA, answerMsg(aRR("example.com", "9.9.9.9"))) - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dns.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - Hooks: hooks, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerResp.Copy(), nil - }) + var got []string + cfg := testConfig(false) + cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) { + got = append(got, ev.Stage.String()+":"+ev.RefID) + }} + runTraversal(t, cfg, m, "example.com") - _, err := tr.Traverse(context.Background(), "example.com") - if err != nil { - t.Fatalf("Traverse: %v", err) + want := []string{"new:", "start:", "new:1", "start:1", "answer:1", "answer:"} + if len(got) != len(want) { + t.Fatalf("events = %v, want %v", got, want) } - if len(events) < 2 { - t.Fatalf("expected start and complete events, got %d", len(events)) + for i := range want { + if got[i] != want[i] { + t.Fatalf("events = %v, want %v", got, want) + } } - if events[0].Stage != EventStart { - t.Fatalf("first event stage = %v, want start", events[0].Stage) - } - if events[1].Stage != EventComplete { - t.Fatalf("second event stage = %v, want complete", events[1].Stage) +} + +func TestHooksEventCarriesStatus(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "example.com", dns.TypeA, answerMsg(aRR("example.com", "9.9.9.9"))) + + var answerStatus Status + cfg := testConfig(false) + cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) { + if ev.Stage == StageAnswer && ev.RefID == "1" { + answerStatus = ev.Status + } + }} + runTraversal(t, cfg, m, "example.com") + if answerStatus != StatusAnswered { + t.Errorf("answer event status = %q, want answered", answerStatus) } } diff --git a/internal/traverse/referral.go b/internal/traverse/referral.go index 1ed1671..68129cb 100644 --- a/internal/traverse/referral.go +++ b/internal/traverse/referral.go @@ -4,6 +4,7 @@ import ( "context" "fmt" "net" + "sort" "strings" "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" @@ -11,44 +12,87 @@ import ( "golang.org/x/net/idna" ) -type ResolutionState int +// DefaultMaxDepth is the default maximum referral depth (non-zero refid +// components) before a "Maxdepth N exceeded" exception is injected. +const DefaultMaxDepth = 20 + +// ReferralStatus is the resolve-phase status of a Referral node itself +// (referral.rb @status), distinct from the per-IP response statuses. +type ReferralStatus string const ( - StateUnresolved ResolutionState = iota - StateResolving - StateResolved + RefStatusNormal ReferralStatus = "normal" + RefStatusNoGlue ReferralStatus = "noglue" + RefStatusLoop ReferralStatus = "loop" ) -func (s ResolutionState) String() string { - switch s { - case StateUnresolved: - return "unresolved" - case StateResolving: - return "resolving" - case StateResolved: - return "resolved" - default: - return "unknown" - } +// StatsEntry is one aggregated leaf statistic: the probability mass that +// ended in Response's outcome at Referral (referral.rb @stats values). +type StatsEntry struct { + Key string + Prob float64 + Response *ServerResponse + Referral *Referral } +// Referral represents one referral to a specific server for qname/qclass/ +// qtype (referral.rb). The synthetic top node ("rootroot") has Server == "" +// and is never displayed; its children are the root servers. type Referral struct { - Name string - Qtype uint16 - Qclass uint16 - Bailiwick string - - Addresses []net.IP - State ResolutionState - - NSName string + RefID string Parent *Referral - Depth int - Prob float64 + + Qname string + Qclass uint16 + Qtype uint16 + // NSAType is the record type used to resolve nameserver addresses + // (always A; the reference is IPv4-only for transport). + NSAType uint16 + + // Server is the NS hostname this referral queries ("" for rootroot). + Server string + // ServerIPs is nil when the server still needs resolving. After a + // resolve it may also contain "key:..." pseudo entries carrying the + // probability of failed resolutions. + ServerIPs []string + Bailiwick string + ParentIP string + + InfoCache *InfoCache + Status ReferralStatus + + // Responses holds the classified response for each real IP queried. + Responses map[string]*ServerResponse + // Children holds child referrals keyed by the parent IP that produced + // them ("rootroot" for the synthetic top node). + Children map[string][]*Referral + // Resolves is the glue-resolution subtree (refid ".0." components). + Resolves []*Referral + + ServerWeights map[string]float64 + Warnings []string + + // Stats is the post-order aggregation of leaf outcomes below this node. + Stats map[string]*StatsEntry + // StatsResolve aggregates the outcomes of the resolve subtree. + StatsResolve map[string]*StatsEntry + + // ReplacedBy points at the earlier completed referral that fast mode + // substituted for this node. + ReplacedBy *Referral + + // summaryStats memoises SummaryStats() (referral.rb summary_stats). + summaryStats *SummaryStats + + client *dns.Client + maxdepth int + referralResolution bool + processed bool + calculated bool } -// idnaLookup is the IDN lookup profile used to convert internationalised domain -// names (unicode labels) to their ACE/punycode equivalents before querying. +// idnaLookup is the IDN lookup profile used to convert internationalised +// domain names (unicode labels) to their ACE/punycode equivalents. var idnaLookup = idna.New( idna.MapForLookup(), idna.BidiRule(), @@ -56,9 +100,8 @@ var idnaLookup = idna.New( ) // toASCII converts a domain name that may contain unicode labels to its -// punycode (ACE) representation. Pure-ASCII names are returned unchanged. -// On conversion errors the original name is returned so the caller can still -// attempt a query (the server will reject it if truly invalid). +// punycode (ACE) representation. On conversion errors the original name is +// returned so the caller can still attempt a query. func toASCII(name string) string { if name == "" || name == "." { return name @@ -70,210 +113,508 @@ func toASCII(name string) string { return ascii } -func NewReferral(name string, qtype uint16, bailiwick string, depth int, prob float64, parent *Referral) *Referral { - return &Referral{ - Name: miekgdns.Fqdn(strings.ToLower(toASCII(name))), - Qtype: qtype, - Qclass: miekgdns.ClassINET, - Bailiwick: miekgdns.Fqdn(strings.ToLower(toASCII(bailiwick))), - Depth: depth, - Prob: prob, - Parent: parent, - State: StateUnresolved, +// referralArgs are the per-child overrides for makeReferral; zero values +// inherit from the parent (referral.rb make_referral merge semantics). +type referralArgs struct { + qname string + qtype uint16 + server string + serverIPs []string + bailiwick string + infoCache *InfoCache + refid string + parentIP string + referralResolution bool +} + +func (r *Referral) makeReferral(a referralArgs) *Referral { + child := &Referral{ + RefID: a.refid, + Parent: r, + Qname: r.Qname, + Qclass: r.Qclass, + Qtype: r.Qtype, + NSAType: r.NSAType, + Server: canonicalName(a.server), + ServerIPs: a.serverIPs, + Bailiwick: canonicalName(a.bailiwick), + ParentIP: a.parentIP, + InfoCache: r.InfoCache, + Status: RefStatusNormal, + Responses: make(map[string]*ServerResponse), + Children: make(map[string][]*Referral), + ServerWeights: make(map[string]float64), + client: r.client, + maxdepth: r.maxdepth, + referralResolution: a.referralResolution || r.referralResolution, } -} - -func (r *Referral) InBailiwick(name string) bool { - if r.Bailiwick == "" || r.Bailiwick == "." { - return true + if a.qname != "" { + child.Qname = canonicalName(a.qname) } - fqdn := miekgdns.Fqdn(strings.ToLower(name)) - return miekgdns.IsSubDomain(r.Bailiwick, fqdn) -} - -func (r *Referral) HasAddresses() bool { - return len(r.Addresses) > 0 -} - -func (r *Referral) SetAddresses(addrs []net.IP) { - r.Addresses = addrs - if len(addrs) > 0 { - r.State = StateResolved - } else { - r.State = StateUnresolved + if a.qtype != 0 { + child.Qtype = a.qtype } -} - -type CircularReferralError struct { - Name string - Chain []string -} - -func (e *CircularReferralError) Error() string { - return fmt.Sprintf("circular referral detected for %s: %v", e.Name, e.Chain) -} - -type UnresolvableNameserverError struct { - Name string - Reason string -} - -func (e *UnresolvableNameserverError) Error() string { - return fmt.Sprintf("unresolvable nameserver %s: %s", e.Name, e.Reason) -} - -func (r *Referral) Resolve(ctx context.Context, traverser *Traverser, cache *InfoCache, visited map[string]bool, depth int) error { - if r.HasAddresses() { - r.State = StateResolved - return nil + if a.infoCache != nil { + child.InfoCache = a.infoCache } - - if cache != nil { - if addrs := cache.LookupGlue(r.Name); len(addrs) > 0 { - r.Addresses = addrs - r.State = StateResolved - return nil + // serverweight = 1/len(ips) per IP when the addresses are known. + if child.ServerIPs != nil { + for _, ip := range child.ServerIPs { + child.ServerWeights[ip] = 1.0 / float64(len(child.ServerIPs)) } } + return child +} - if visited != nil { - if visited[r.Name] { - return &CircularReferralError{ - Name: r.Name, - Chain: getVisitedNames(visited), +// IsRootRoot reports whether this is the synthetic top node representing an +// automatic referral to the root servers. +func (r *Referral) IsRootRoot() bool { + return r.Server == "" +} + +// IsResolve reports whether this node is part of a glue-resolution subtree. +func (r *Referral) IsResolve() bool { + return r.referralResolution +} + +// Resolved reports whether the server addresses are known (rootroot is +// always resolved). +func (r *Referral) Resolved() bool { + return r.IsRootRoot() || r.ServerIPs != nil +} + +// Depth counts the non-zero refid components; resolve subtrees (".0.") do +// not count against the depth limit. +func (r *Referral) Depth() int { + return refidDepth(r.RefID) +} + +func refidDepth(refid string) int { + if refid == "" { + return 0 + } + n := 0 + for _, part := range strings.Split(refid, ".") { + if part != "0" { + n++ + } + } + return n +} + +// IPsAsArray returns the real IP addresses known for this referral, +// excluding "key:" pseudo entries. +func (r *Referral) IPsAsArray() []string { + var out []string + for _, ip := range r.ServerIPs { + if !strings.HasPrefix(ip, "key:") { + out = append(out, ip) + } + } + return out +} + +// TxtIPsVerbose renders the per-IP weights, sorted, e.g. +// "50.0%=1.2.3.4,50.0%=noglue:1.2.3.4" (referral.rb txt_ips_verbose). It is +// part of the fast-mode memo key. +func (r *Referral) TxtIPsVerbose() string { + if r.ServerIPs == nil { + return "" + } + parts := make([]string, 0, len(r.ServerIPs)) + for _, ip := range r.ServerIPs { + label := ip + if rest, ok := strings.CutPrefix(ip, "key:"); ok { + // keep the first two colon-separated fields, like Ruby's + // /^key:([^:]+(:[^:]*)?)/ capture. + fields := strings.SplitN(rest, ":", 3) + if len(fields) > 2 { + fields = fields[:2] + } + label = strings.Join(fields, ":") + } + parts = append(parts, fmt.Sprintf("%.1f%%=%s", 100*r.ServerWeights[ip], label)) + } + sort.Strings(parts) + return strings.Join(parts, ",") +} + +// TxtIPs renders the addresses for progress display; failed-resolve pseudo +// entries render as their response description (referral.rb txt_ips). +func (r *Referral) TxtIPs() string { + if r.ServerIPs == nil { + return "" + } + parts := make([]string, 0, len(r.ServerIPs)) + for _, ip := range r.ServerIPs { + if strings.HasPrefix(ip, "key:") { + if e, ok := r.StatsResolve[ip]; ok && e.Response != nil { + parts = append(parts, e.Response.String()) + continue } } - visited[r.Name] = true - } - - if depth > DefaultMaxDepth { - return &UnresolvableNameserverError{ - Name: r.Name, - Reason: "max depth exceeded", - } - } - - roots, err := traverser.discoverRoots(ctx) - if err != nil { - return fmt.Errorf("root discovery: %w", err) - } - - initial := NewReferral(r.Name, dns.TypeA, ".", 0, 1.0, nil) - initial.Addresses = roots - initial.State = StateResolved - - stack := NewStack(DefaultMaxDepth) - stack.Push(initial) - - traversalCache := NewInfoCache(nil) - if visited != nil { - for name := range visited { - traversalCache.StoreGlue(name, []net.IP{}) - } - } - - var lastErr error - for { - select { - case <-ctx.Done(): - return fmt.Errorf("resolution cancelled: %w", ctx.Err()) - default: - } - - ref := stack.Pop() - if ref == nil { - break - } - - cacheForStep := traversalCache - if ref.Parent != nil { - cacheForStep = traversalCache.Child() - } - - resp := traverser.processReferral(ctx, ref, cacheForStep) - - if resp.Type == RespAnswer { - if len(resp.Decoded.Answers) > 0 { - var addrs []net.IP - for _, rr := range resp.Decoded.Answers { - if a, ok := rr.(*miekgdns.A); ok { - addrs = append(addrs, a.A) - } - if aaaa, ok := rr.(*miekgdns.AAAA); ok { - addrs = append(addrs, aaaa.AAAA) - } - } - if len(addrs) > 0 { - r.Addresses = addrs - r.State = StateResolved - if cache != nil { - cache.StoreGlue(r.Name, addrs) - } - return nil - } - } - } - - if resp.Type == RespNXDOMAIN { - lastErr = &UnresolvableNameserverError{ - Name: r.Name, - Reason: "NXDOMAIN", - } - break - } - - if resp.Type == RespSERVFAIL || resp.Type == RespError { - lastErr = fmt.Errorf("server error resolving %s: %s", r.Name, resp.Type) - continue - } - - if resp.Type == RespReferral { - children := resp.ChildReferrals() - for _, child := range children { - // Only skip visited names when they have no addresses; if glue - // was included in the referral response we still need to query - // that child to get the authoritative answer. - if visited != nil && visited[child.Name] && !child.HasAddresses() { - continue - } - if !stack.Push(child) { - lastErr = &UnresolvableNameserverError{ - Name: r.Name, - Reason: "max depth exceeded during resolution", - } - } - } - } - } - - if lastErr != nil { - return lastErr - } - - return &UnresolvableNameserverError{ - Name: r.Name, - Reason: "resolution exhausted without answer", + parts = append(parts, ip) } + sort.Strings(parts) + return strings.Join(parts, ",") } -func getVisitedNames(visited map[string]bool) []string { - var names []string - for name := range visited { - names = append(names, name) - } - return names +func (r *Referral) String() string { + return fmt.Sprintf("%s [%s/%s/%s] server=%s server_ips=%s bailiwick=%s", + r.RefID, r.Qname, ClassToString(r.Qclass), TypeToString(r.Qtype), + r.Server, r.TxtIPs(), r.Bailiwick) } -// IsNameInChain reports whether name appears anywhere in this referral's ancestor -// chain, including this referral itself. Used for CNAME loop detection. -func (r *Referral) IsNameInChain(name string) bool { - n := miekgdns.Fqdn(strings.ToLower(name)) - curr := r - for curr != nil { - if curr.Name == n { +// OverallStatus summarises the node's outcome for event consumers: a resolve +// dead end (noglue/loop), the shared status of every per-IP response, or +// "mixed" when the responses disagree ("" before anything was queried). +func (r *Referral) OverallStatus() Status { + switch r.Status { + case RefStatusNoGlue: + return StatusNoGlue + case RefStatusLoop: + return StatusLoop + } + var s Status + for _, resp := range r.Responses { + if s == "" { + s = resp.Status + } else if s != resp.Status { + return "mixed" + } + } + return s +} + +// StatsList returns the aggregated leaf statistics sorted by stats key. +func (r *Referral) StatsList() []*StatsEntry { + out := make([]*StatsEntry, 0, len(r.Stats)) + for _, e := range r.Stats { + out = append(out, e) + } + sort.Slice(out, func(i, j int) bool { return out[i].Key < out[j].Key }) + return out +} + +// isNoGlue reports a dead end: the server is inside the current bailiwick, +// so its address should have come as glue, but none was provided and no +// deeper zone can name it (referral.rb noglue?). +func (r *Referral) isNoGlue() bool { + return r.ServerIPs == nil && insideBailiwick(r.Server, r.Bailiwick) +} + +// isLoop reports a resolve loop: an ancestor referral asks the same +// qname/qclass/qtype of the same still-unresolved server (referral.rb loop?), +// e.g. b NS c.d while d NS a.b. +func (r *Referral) isLoop() bool { + if r.ServerIPs != nil { + return false + } + for p := r.Parent; p != nil; p = p.Parent { + if p.Qname == r.Qname && p.Qclass == r.Qclass && p.Qtype == r.Qtype && + p.Server == r.Server && p.ServerIPs == nil { return true } - curr = curr.Parent } return false } + +// chainHasQuery reports whether this referral or any ancestor already asks +// qname with the same qclass/qtype. Used to stop cross-response CNAME chains +// (restart loops) that Ruby only catches via the depth limit. +func (r *Referral) chainHasQuery(qname string) bool { + name := canonicalName(qname) + for p := r; p != nil; p = p.Parent { + if p.Qname == name && p.Qclass == r.Qclass && p.Qtype == r.Qtype { + return true + } + } + return false +} + +// resolve turns an address-less referral into either a dead end (noglue/ +// loop) or a resolve subtree querying A from this branch's cache +// (referral.rb resolve). It returns the referrals to process. +func (r *Referral) resolve() ([]*Referral, error) { + if r.isNoGlue() { + r.Status = RefStatusNoGlue + return nil, nil + } + if r.isLoop() { + r.Status = RefStatusLoop + return nil, nil + } + starters, newbailiwick, err := r.InfoCache.GetStartServers(r.Server) + if err != nil { + return nil, err + } + for i, st := range starters { + child := r.makeReferral(referralArgs{ + qname: r.Server, + qtype: r.NSAType, + server: st.Name, + serverIPs: st.IPs, + bailiwick: newbailiwick, + refid: fmt.Sprintf("%s.0.%d", r.RefID, i+1), + referralResolution: true, + }) + r.Resolves = append(r.Resolves, child) + } + return r.Resolves, nil +} + +// resolveCalculate folds the resolve subtree's statistics into per-IP server +// weights (referral.rb resolve_calculate): each answered leaf distributes +// its probability evenly across the A records it returned; every other leaf +// keeps its probability under its "key:" stats key so failures surface in +// the results. +func (r *Referral) resolveCalculate() { + r.StatsResolve = make(map[string]*StatsEntry) + switch r.Status { + case RefStatusNoGlue: + resp := NewNoGlueResponse(r.Qname, r.Qclass, r.Qtype, r.ParentIP, r.Server, r.Bailiwick) + key := resp.StatsKey() + r.StatsResolve[key] = &StatsEntry{Key: key, Prob: 1.0, Response: resp, Referral: r} + case RefStatusLoop: + resp := NewLoopResponse(r.Qname, r.Qclass, r.Qtype, r.ParentIP, r.Server, r.Bailiwick) + key := resp.StatsKey() + r.StatsResolve[key] = &StatsEntry{Key: key, Prob: 1.0, Response: resp, Referral: r} + default: + statsCalculateChildren(r.StatsResolve, r.Resolves, 1.0) + } + + r.ServerWeights = make(map[string]float64) + r.ServerIPs = []string{} + keys := make([]string, 0, len(r.StatsResolve)) + for key := range r.StatsResolve { + keys = append(keys, key) + } + sort.Strings(keys) + addWeight := func(ip string, prob float64) { + if _, ok := r.ServerWeights[ip]; !ok { + r.ServerIPs = append(r.ServerIPs, ip) + } + r.ServerWeights[ip] += prob + } + for _, key := range keys { + data := r.StatsResolve[key] + if data.Response.Status == StatusAnswered { + var addrs []string + for _, rr := range data.Response.DQ.Answers { + if a, ok := rr.(*miekgdns.A); ok { + addrs = append(addrs, a.A.String()) + } + } + for _, addr := range addrs { + addWeight(addr, data.Prob/float64(len(addrs))) + } + if len(addrs) == 0 { + // answered but no A records (e.g. AAAA-only): carry the + // probability as a failure key so mass is not lost. + addWeight(key, data.Prob) + } + } else { + addWeight(key, data.Prob) + } + } +} + +// statsCalculateChildren merges the children's statistics into stats with an +// equal split of weight among them (referral.rb stats_calculate_children). +func statsCalculateChildren(stats map[string]*StatsEntry, children []*Referral, weight float64) { + if len(children) == 0 { + return + } + percent := (1.0 / float64(len(children))) * weight + for _, child := range children { + for key, data := range child.Stats { + if e, ok := stats[key]; ok { + e.Prob += data.Prob * percent + } else { + stats[key] = &StatsEntry{ + Key: key, + Prob: data.Prob * percent, + Response: data.Response, + Referral: data.Referral, + } + } + } + } +} + +// answerCalculate computes this node's aggregated statistics from its +// children, responses and resolve failures (referral.rb answer_calculate). +// Unlike the Ruby source, duplicate stats keys across a referral's IPs merge +// by summing probability (Ruby computes the sum then discards it — a source +// bug that breaks the probabilities-sum-to-1 invariant). +func (r *Referral) answerCalculate() { + r.Stats = make(map[string]*StatsEntry) + if r.IsRootRoot() { + statsCalculateChildren(r.Stats, r.Children["rootroot"], 1.0) + r.calculated = true + return + } + for _, ip := range r.ServerIPs { + serverweight := r.ServerWeights[ip] + if strings.HasPrefix(ip, "key:") { + // resolve failed for some reason - copy the resolve statistics + src := r.StatsResolve[ip] + if e, ok := r.Stats[ip]; ok { + e.Prob += src.Prob + } else { + r.Stats[ip] = &StatsEntry{Key: ip, Prob: src.Prob, Response: src.Response, Referral: src.Referral} + } + continue + } + if children := r.Children[ip]; len(children) > 0 { + statsCalculateChildren(r.Stats, children, serverweight) + continue + } + resp := r.Responses[ip] + if resp == nil { + continue + } + key := resp.StatsKey() + if e, ok := r.Stats[key]; ok { + e.Prob += serverweight + } else { + r.Stats[key] = &StatsEntry{Key: key, Prob: serverweight, Response: resp, Referral: r} + } + } + r.calculated = true +} + +// process queries every real IP of this referral through the packet cache, +// classifies each response, and creates one child per NS name (including +// glueless ones) for referral/restart statuses (referral.rb process/ +// process_normal). It returns one set of children per IP that produced any. +func (r *Referral) process(ctx context.Context) ([][]*Referral, error) { + r.processed = true + if r.IsRootRoot() { + children, err := r.processAddRoots() + if err != nil { + return nil, err + } + return [][]*Referral{children}, nil + } + + // Phase one: query and classify, counting childsets so refids can grow + // an extra childset digit when more than one IP produces children. + childsets := 0 + var order []string + for _, ip := range r.ServerIPs { + if strings.HasPrefix(ip, "key:") { + continue + } + var dq *DecodedQuery + if r.Depth() >= r.maxdepth { + err := fmt.Errorf("Maxdepth %d exceeded", r.maxdepth) + dq = NewDecodedQuery(nil, err, r.Qname, r.Qclass, r.Qtype, ip, r.Bailiwick) + } else { + msg, warnings, err := r.client.Query(ctx, net.ParseIP(ip), r.Qname, r.Qtype) + dq = NewDecodedQuery(msg, err, r.Qname, r.Qclass, r.Qtype, ip, r.Bailiwick) + dq.WarningsAdd(warnings...) + } + resp, err := NewServerResponse(dq, r.Server, r.ParentIP, r.InfoCache) + if err != nil { + return nil, err + } + if resp.Status == StatusRestart { + // Cross-response CNAME loop: any target in the chain that we + // (or an ancestor) are already querying is a dead end. + for _, target := range dq.ChainTargets { + if r.chainHasQuery(target) { + resp.Status = StatusCNAMELoop + break + } + } + } + r.Warnings = append(r.Warnings, dq.Warnings...) + r.Responses[ip] = resp + order = append(order, ip) + if resp.Status == StatusRestart || resp.Status == StatusReferral { + childsets++ + } + } + + // Phase two: create the children. + childset := 0 + var sets [][]*Referral + for _, ip := range order { + resp := r.Responses[ip] + if resp.Status != StatusRestart && resp.Status != StatusReferral { + continue + } + childset++ + refid := r.RefID + if childsets > 1 { + refid = fmt.Sprintf("%s.%d", r.RefID, childset) + } + children := r.makeReferrals(resp, refid, ip) + r.Children[ip] = children + sets = append(sets, children) + } + return sets, nil +} + +// processAddRoots creates one child per root server with equal weight +// (referral.rb process_add_roots); the roots come from the info cache hints. +func (r *Referral) processAddRoots() ([]*Referral, error) { + starters, _, err := r.InfoCache.GetStartServers("") + if err != nil { + return nil, err + } + dot := "" + if r.RefID != "" { + dot = "." + } + var children []*Referral + for i, root := range starters { + child := r.makeReferral(referralArgs{ + server: root.Name, + serverIPs: root.IPs, + refid: fmt.Sprintf("%s%s%d", r.RefID, dot, i+1), + }) + children = append(children, child) + } + r.Children["rootroot"] = children + return children, nil +} + +// makeReferrals creates one child per start server for a referral/restart +// response (referral.rb make_referrals): qname moves to the response's +// endname (the CNAME target on restart), the bailiwick and cache come from +// the response. +func (r *Referral) makeReferrals(resp *ServerResponse, refid, parentIP string) []*Referral { + var children []*Referral + for i, st := range resp.Starters { + children = append(children, r.makeReferral(referralArgs{ + qname: resp.DQ.Endname, + server: st.Name, + serverIPs: st.IPs, + bailiwick: resp.StartersBailiwick, + infoCache: resp.Cache, + refid: fmt.Sprintf("%s.%d", refid, i+1), + parentIP: parentIP, + })) + } + return children +} + +// replaceChild swaps before for after in the children/resolves lists (fast +// mode substitution); before keeps a pointer to its replacement. +func (r *Referral) replaceChild(before, after *Referral) { + before.ReplacedBy = after + for ip := range r.Children { + for i, c := range r.Children[ip] { + if c == before { + r.Children[ip][i] = after + } + } + } + for i, c := range r.Resolves { + if c == before { + r.Resolves[i] = after + } + } +} diff --git a/internal/traverse/referral_test.go b/internal/traverse/referral_test.go index 5e6ecd1..2e95468 100644 --- a/internal/traverse/referral_test.go +++ b/internal/traverse/referral_test.go @@ -1,159 +1,194 @@ package traverse import ( - "net" "testing" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" + "github.com/miekg/dns" ) -const TypeA = dns.TypeA - -type State = ResolutionState - -func TestNewReferral(t *testing.T) { - ref := NewReferral("example.com", TypeA, ".", 0, 1.0, nil) - if ref.Name != "example.com." { - t.Errorf("Name = %q, want %q", ref.Name, "example.com.") +func newTestReferral(server string, ips []string) *Referral { + r := &Referral{ + RefID: "1", + Qname: "www.example.com", + Qclass: dns.ClassINET, + Qtype: dns.TypeA, + NSAType: dns.TypeA, + Server: server, + ServerIPs: ips, + Bailiwick: "com", + InfoCache: NewInfoCache(nil), + Status: RefStatusNormal, + Responses: make(map[string]*ServerResponse), + Children: make(map[string][]*Referral), + ServerWeights: make(map[string]float64), + maxdepth: DefaultMaxDepth, } - if ref.Qtype != TypeA { - t.Errorf("Qtype = %d, want %d", ref.Qtype, TypeA) - } - if ref.Qclass != 1 { - t.Errorf("Qclass = %d, want 1", ref.Qclass) - } - if ref.State != StateUnresolved { - t.Errorf("State = %d, want %d", ref.State, StateUnresolved) - } - if ref.Depth != 0 { - t.Errorf("Depth = %d, want 0", ref.Depth) - } - if ref.Prob != 1.0 { - t.Errorf("Prob = %f, want 1.0", ref.Prob) - } - if ref.Parent != nil { - t.Error("Parent should be nil") + for _, ip := range ips { + r.ServerWeights[ip] = 1.0 / float64(len(ips)) } + return r } -func TestReferralInBailiwick(t *testing.T) { +func TestRefidDepth(t *testing.T) { tests := []struct { - name string - bailiwick string - testName string - want bool + refid string + want int }{ - {"root bailiwick accepts all", ".", "example.com.", true}, - {"empty bailiwick accepts all", "", "example.com.", true}, - {"subdomain in bailiwick", "com.", "example.com.", true}, - {"deeper subdomain", "com.", "www.example.com.", true}, - {"not in bailiwick", "org.", "example.com.", false}, - {"same zone", "example.com.", "example.com.", true}, - {"sibling zone", "example.com.", "other.com.", false}, + {"", 0}, + {"1", 1}, + {"1.1.2", 3}, + {"1.2.0.1", 3}, + {"1.1.2.0.1.4.2.0.2.0.2", 8}, } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - ref := &Referral{Bailiwick: tt.bailiwick} - if got := ref.InBailiwick(tt.testName); got != tt.want { - t.Errorf("InBailiwick(%q) = %v, want %v", tt.testName, got, tt.want) - } - }) - } -} - -func TestReferralHasAddresses(t *testing.T) { - ref := &Referral{} - if ref.HasAddresses() { - t.Error("empty referral should not have addresses") - } - - ref.Addresses = []net.IP{net.ParseIP("1.2.3.4")} - if !ref.HasAddresses() { - t.Error("referral with address should have addresses") - } -} - -func TestReferralSetAddresses(t *testing.T) { - ref := &Referral{} - - ref.SetAddresses([]net.IP{net.ParseIP("1.2.3.4")}) - if ref.State != StateResolved { - t.Errorf("State = %d, want %d", ref.State, StateResolved) - } - if !ref.HasAddresses() { - t.Error("should have addresses after SetAddresses") - } - - ref.SetAddresses(nil) - if ref.State != StateUnresolved { - t.Errorf("State = %d, want %d", ref.State, StateUnresolved) - } -} - -func TestResolutionStateString(t *testing.T) { - tests := []struct { - state State - want string - }{ - {StateUnresolved, "unresolved"}, - {StateResolving, "resolving"}, - {StateResolved, "resolved"}, - } - - for _, tt := range tests { - t.Run(tt.want, func(t *testing.T) { - if got := tt.state.String(); got != tt.want { - t.Errorf("String() = %q, want %q", got, tt.want) - } - }) - } -} - -func TestCircularReferralError(t *testing.T) { - err := &CircularReferralError{ - Name: "ns.example.com.", - Chain: []string{"ns1.example.com.", "ns2.example.com."}, - } - - expected := "circular referral detected for ns.example.com.: [ns1.example.com. ns2.example.com.]" - if err.Error() != expected { - t.Errorf("Error() = %q, want %q", err.Error(), expected) - } -} - -func TestUnresolvableNameserverError(t *testing.T) { - err := &UnresolvableNameserverError{ - Name: "ns.example.com.", - Reason: "NXDOMAIN", - } - - expected := "unresolvable nameserver ns.example.com.: NXDOMAIN" - if err.Error() != expected { - t.Errorf("Error() = %q, want %q", err.Error(), expected) - } -} - -func TestGetVisitedNames(t *testing.T) { - visited := map[string]bool{ - "ns1.example.com.": true, - "ns2.example.com.": true, - "ns3.example.com.": true, - } - - names := getVisitedNames(visited) - if len(names) != 3 { - t.Errorf("got %d names, want 3", len(names)) - } - - seen := make(map[string]bool) - for _, name := range names { - if seen[name] { - t.Errorf("duplicate name: %s", name) - } - seen[name] = true - if !visited[name] { - t.Errorf("unexpected name: %s", name) + if got := refidDepth(tt.refid); got != tt.want { + t.Errorf("refidDepth(%q) = %d, want %d", tt.refid, got, tt.want) } } } + +func TestTxtIPsVerbose(t *testing.T) { + r := newTestReferral("ns1.example.com", []string{"2.2.2.2", "1.1.1.1"}) + if got := r.TxtIPsVerbose(); got != "50.0%=1.1.1.1,50.0%=2.2.2.2" { + t.Errorf("TxtIPsVerbose = %q", got) + } + + // key: pseudo entries keep only the first two fields. + r2 := newTestReferral("ns2.example.com", nil) + r2.ServerIPs = []string{"key:noglue:9.9.9.9:www.example.com:IN:A:x:y"} + r2.ServerWeights = map[string]float64{r2.ServerIPs[0]: 1.0} + if got := r2.TxtIPsVerbose(); got != "100.0%=noglue:9.9.9.9" { + t.Errorf("TxtIPsVerbose key entry = %q", got) + } + + var unresolved Referral + if got := unresolved.TxtIPsVerbose(); got != "" { + t.Errorf("unresolved TxtIPsVerbose = %q, want empty", got) + } +} + +func TestFastKeyLowercases(t *testing.T) { + r := newTestReferral("NS1.Example.COM", []string{"1.1.1.1"}) + r.Server = "NS1.Example.COM" // bypass canonicalisation to prove downcasing + key := fastKey(r) + if key != "www.example.com:in:a:ns1.example.com:100.0%=1.1.1.1" { + t.Errorf("fastKey = %q", key) + } +} + +func TestIsNoGlueAndIsLoop(t *testing.T) { + inBailiwick := newTestReferral("ns1.sub.com", nil) + if !inBailiwick.isNoGlue() { + t.Error("in-bailiwick NS without addresses should be noglue") + } + outOfBailiwick := newTestReferral("ns1.other.net", nil) + if outOfBailiwick.isNoGlue() { + t.Error("out-of-bailiwick NS is resolvable, not noglue") + } + resolved := newTestReferral("ns1.sub.com", []string{"1.1.1.1"}) + if resolved.isNoGlue() || resolved.isLoop() { + t.Error("resolved referral is neither noglue nor loop") + } + + parent := newTestReferral("ns.b.net", nil) + parent.Qname = "ns.a.net" + child := newTestReferral("ns.b.net", nil) + child.Qname = "ns.a.net" + child.Parent = parent + if !child.isLoop() { + t.Error("same qname/qclass/qtype/server with unresolved ancestor should loop") + } + child.Server = "ns.c.net" + if child.isLoop() { + t.Error("different server must not loop") + } +} + +func TestChainHasQuery(t *testing.T) { + parent := newTestReferral("ns1.example.com", []string{"1.1.1.1"}) + parent.Qname = "www.a.com" + child := newTestReferral("ns2.example.com", []string{"2.2.2.2"}) + child.Qname = "www.b.net" + child.Parent = parent + if !child.chainHasQuery("www.a.com") { + t.Error("ancestor qname should be found") + } + if !child.chainHasQuery("WWW.B.NET.") { + t.Error("own qname should be found case-insensitively") + } + if child.chainHasQuery("www.c.org") { + t.Error("unknown qname must not match") + } +} + +func TestReplaceChild(t *testing.T) { + parent := newTestReferral("ns1.example.com", []string{"1.1.1.1"}) + before := newTestReferral("ns2.example.com", []string{"2.2.2.2"}) + after := newTestReferral("ns2.example.com", []string{"2.2.2.2"}) + parent.Children["1.1.1.1"] = []*Referral{before} + parent.Resolves = []*Referral{before} + + parent.replaceChild(before, after) + if parent.Children["1.1.1.1"][0] != after || parent.Resolves[0] != after { + t.Error("replaceChild did not swap the node everywhere") + } + if before.ReplacedBy != after { + t.Error("replaced node should point at its replacement") + } +} + +func TestOverallStatus(t *testing.T) { + r := newTestReferral("ns1.example.com", nil) + r.Status = RefStatusNoGlue + if got := r.OverallStatus(); got != StatusNoGlue { + t.Errorf("noglue overall = %q", got) + } + r.Status = RefStatusLoop + if got := r.OverallStatus(); got != StatusLoop { + t.Errorf("loop overall = %q", got) + } + + n := newTestReferral("ns1.example.com", []string{"1.1.1.1", "2.2.2.2"}) + if got := n.OverallStatus(); got != "" { + t.Errorf("no responses overall = %q, want empty", got) + } + n.Responses["1.1.1.1"] = &ServerResponse{Status: StatusAnswered} + if got := n.OverallStatus(); got != StatusAnswered { + t.Errorf("single status overall = %q", got) + } + n.Responses["2.2.2.2"] = &ServerResponse{Status: StatusError} + if got := n.OverallStatus(); got != "mixed" { + t.Errorf("mixed overall = %q", got) + } +} + +func TestIPsAsArraySkipsPseudoKeys(t *testing.T) { + r := newTestReferral("ns1.example.com", []string{"1.1.1.1", "key:noglue:2.2.2.2"}) + got := r.IPsAsArray() + if len(got) != 1 || got[0] != "1.1.1.1" { + t.Errorf("IPsAsArray = %v", got) + } +} + +func TestToASCII(t *testing.T) { + if got := toASCII("bücher.example"); got != "xn--bcher-kva.example" { + t.Errorf("toASCII = %q", got) + } + if got := toASCII("plain.example"); got != "plain.example" { + t.Errorf("ascii name changed: %q", got) + } + if got := toASCII(""); got != "" { + t.Errorf("empty name changed: %q", got) + } +} + +func TestServerResponseString(t *testing.T) { + noglue := NewNoGlueResponse("www.example.com", dns.ClassINET, dns.TypeA, "1.1.1.1", "ns1.example.com", "example.com") + if got := noglue.String(); got != "No glue for ns1.example.com" { + t.Errorf("noglue String = %q", got) + } + loop := NewLoopResponse("www.example.com", dns.ClassINET, dns.TypeA, "1.1.1.1", "ns1.example.com", "example.com") + if got := loop.String(); got != "Loop encountered resolving ns1.example.com" { + t.Errorf("loop String = %q", got) + } +} diff --git a/internal/traverse/response.go b/internal/traverse/response.go deleted file mode 100644 index 7d8371f..0000000 --- a/internal/traverse/response.go +++ /dev/null @@ -1,256 +0,0 @@ -package traverse - -import ( - "net" - - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" - miekgdns "github.com/miekg/dns" -) - -type ResponseType int - -const ( - RespReferral ResponseType = iota - RespAnswer - RespCNAMEFollow - RespNODATA - RespNXDOMAIN - RespSERVFAIL - RespREFUSED - RespNOTIMPL - RespCNAMELoop - RespError - // RespNSResolutionFailed indicates that the traversal could not resolve the - // IP address of an in-bailiwick nameserver. The domain may still be - // reachable in practice (e.g. via glue records held by the registry), but - // the iterative traversal could not complete that path. - RespNSResolutionFailed -) - -func (rt ResponseType) String() string { - switch rt { - case RespReferral: - return "referral" - case RespAnswer: - return "answer" - case RespCNAMEFollow: - return "cname_follow" - case RespNODATA: - return "nodata" - case RespNXDOMAIN: - return "nxdomain" - case RespSERVFAIL: - return "servfail" - case RespREFUSED: - return "refused" - case RespNOTIMPL: - return "notimp" - case RespCNAMELoop: - return "cname_loop" - case RespError: - return "error" - case RespNSResolutionFailed: - return "ns_error" - default: - return "unknown" - } -} - -type Response struct { - Referral *Referral - Server net.IP - Cache *InfoCache - Decoded *dns.DecodedResponse - Type ResponseType - ErrorMessage string -} - -func NewResponse(ref *Referral, server net.IP, cache *InfoCache) *Response { - return &Response{ - Referral: ref, - Server: server, - Cache: cache, - } -} - -func (r *Response) Process(msg *miekgdns.Msg) *Response { - if msg == nil { - r.Type = RespError - r.ErrorMessage = "nil DNS response" - return r - } - - r.Decoded = dns.DecodeResponse(msg) - if r.Decoded == nil { - r.Type = RespError - r.ErrorMessage = "failed to decode DNS response" - return r - } - - // Synthesize CNAME from DNAME when the server didn't include a synthesized CNAME record. - if len(r.Decoded.CNAMEChain) == 0 && r.Referral != nil && len(r.Decoded.DNAMEMappings) > 0 { - for _, dm := range r.Decoded.DNAMEMappings { - synthesized := dns.SynthesizeCNAMEFromDNAME(r.Referral.Name, dm.Owner, dm.Target) - if synthesized != "" { - r.Decoded.CNAMEChain = append(r.Decoded.CNAMEChain, synthesized) - break - } - } - } - - r.Type = r.classify() - return r -} - -func (r *Response) classify() ResponseType { - switch r.Decoded.Classification { - case dns.ResponseNXDOMAIN: - return RespNXDOMAIN - case dns.ResponseSERVFAIL: - return RespSERVFAIL - case dns.ResponseREFUSED: - return RespREFUSED - case dns.ResponseNOTIMPL: - return RespNOTIMPL - case dns.ResponseAnswer: - if len(r.Decoded.CNAMEChain) > 0 && !r.hasFinalAnswer() { - return RespCNAMEFollow - } - return RespAnswer - case dns.ResponseReferral: - return RespReferral - case dns.ResponseNODATA: - return RespNODATA - default: - return RespError - } -} - -func (r *Response) hasFinalAnswer() bool { - for _, rr := range r.Decoded.Answers { - switch rr.(type) { - case *miekgdns.CNAME, *miekgdns.DNAME, *miekgdns.RRSIG: - // CNAME and DNAME are redirect records, not final answers. - // RRSIG is a DNSSEC signature record — it covers the CNAME/DNAME - // but is not itself the answer to the original question type. - continue - } - return true - } - return false -} - -func (r *Response) ChildReferrals() []*Referral { - if r.Type != RespReferral { - return nil - } - if r.Referral == nil { - return nil - } - - var nameservers []string - for _, rr := range r.Decoded.Authority { - if ns, ok := rr.(*miekgdns.NS); ok { - if r.Referral.InBailiwick(ns.Ns) { - nameservers = append(nameservers, ns.Ns) - } - } - } - - if len(nameservers) == 0 { - for _, rr := range r.Decoded.Authority { - if ns, ok := rr.(*miekgdns.NS); ok { - nameservers = append(nameservers, ns.Ns) - } - } - } - - r.storeAuthority(nameservers) - - prob := r.childProb(len(nameservers)) - var children []*Referral - for _, ns := range nameservers { - child := NewReferral( - r.Referral.Name, - r.Referral.Qtype, - ns, - r.Referral.Depth+1, - prob, - r.Referral, - ) - r.resolveGlue(child) - children = append(children, child) - } - return children -} - -func (r *Response) CNAMEFollowReferral() *Referral { - if r.Type != RespCNAMEFollow || len(r.Decoded.CNAMEChain) == 0 { - return nil - } - target := r.Decoded.CNAMEChain[len(r.Decoded.CNAMEChain)-1] - follow := NewReferral( - target, - r.Referral.Qtype, - r.Referral.Bailiwick, - r.Referral.Depth+1, - r.Referral.Prob, - r.Referral, - ) - if len(r.Referral.Addresses) > 0 { - follow.Addresses = make([]net.IP, len(r.Referral.Addresses)) - copy(follow.Addresses, r.Referral.Addresses) - follow.State = StateResolved - } - return follow -} - -func (r *Response) storeAuthority(nameservers []string) { - if r.Cache == nil { - return - } - zone := r.Referral.Name - r.Cache.StoreNS(zone, nameservers) -} - -func (r *Response) resolveGlue(child *Referral) { - if r.Cache == nil { - return - } - nsName := child.Bailiwick - for _, rr := range r.Decoded.Additional { - switch v := rr.(type) { - case *miekgdns.A: - if normalize(v.Header().Name) == normalize(nsName) { - child.Addresses = append(child.Addresses, v.A) - } - case *miekgdns.AAAA: - if normalize(v.Header().Name) == normalize(nsName) { - child.Addresses = append(child.Addresses, v.AAAA) - } - } - } - if child.HasAddresses() { - child.State = StateResolved - } - r.Cache.StoreGlue(nsName, child.Addresses) -} - -func (r *Response) IsTerminal() bool { - switch r.Type { - case RespAnswer, RespNODATA, RespNXDOMAIN, RespSERVFAIL, RespREFUSED, RespNOTIMPL, RespCNAMELoop, RespError, RespNSResolutionFailed: - return true - default: - return false - } -} - -func (r *Response) childProb(n int) float64 { - if n <= 0 { - return 0 - } - if r.Referral == nil { - return 1.0 / float64(n) - } - return r.Referral.Prob / float64(n) -} diff --git a/internal/traverse/response_test.go b/internal/traverse/response_test.go deleted file mode 100644 index a7bfac1..0000000 --- a/internal/traverse/response_test.go +++ /dev/null @@ -1,340 +0,0 @@ -package traverse - -import ( - "net" - "testing" - - "github.com/miekg/dns" -) - -func TestResponseProcessNil(t *testing.T) { - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(nil) - if r.Type != RespError { - t.Errorf("Type = %d, want %d", r.Type, RespError) - } -} - -func TestResponseClassifyAnswer(t *testing.T) { - msg := new(dns.Msg) - msg.SetReply(new(dns.Msg)) - msg.Answer = append(msg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if r.Type != RespAnswer { - t.Errorf("Type = %d, want %d", r.Type, RespAnswer) - } - if r.Decoded == nil { - t.Fatal("Decoded should not be nil") - } -} - -func TestResponseClassifyReferral(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeSuccess - msg.Authoritative = false - msg.Ns = append(msg.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 172800}, - Ns: "a.gtld-servers.net.", - }) - msg.Extra = append(msg.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 172800}, - A: net.ParseIP("192.5.6.30"), - }) - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - cache := NewInfoCache(nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), cache) - r.Process(msg) - - if r.Type != RespReferral { - t.Errorf("Type = %d, want %d", r.Type, RespReferral) - } -} - -func TestResponseClassifyNXDOMAIN(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeNameError - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if r.Type != RespNXDOMAIN { - t.Errorf("Type = %d, want %d", r.Type, RespNXDOMAIN) - } -} - -func TestResponseClassifyNODATA(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeSuccess - msg.Authoritative = true - msg.Ns = append(msg.Ns, &dns.SOA{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeSOA, Class: dns.ClassINET, Ttl: 3600}, - }) - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if r.Type != RespNODATA { - t.Errorf("Type = %d, want %d", r.Type, RespNODATA) - } -} - -func TestResponseClassifySERVFAIL(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeServerFailure - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if r.Type != RespSERVFAIL { - t.Errorf("Type = %d, want %d", r.Type, RespSERVFAIL) - } -} - -func TestResponseCNAMEFollow(t *testing.T) { - msg := new(dns.Msg) - msg.SetReply(new(dns.Msg)) - msg.Answer = append(msg.Answer, - &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "example.com.", - }, - ) - - ref := NewReferral("www.example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if r.Type != RespCNAMEFollow { - t.Errorf("Type = %d, want %d", r.Type, RespCNAMEFollow) - } -} - -func TestResponseCNAMEWithFinalAnswer(t *testing.T) { - msg := new(dns.Msg) - msg.SetReply(new(dns.Msg)) - msg.Answer = append(msg.Answer, - &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "example.com.", - }, - &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }, - ) - - ref := NewReferral("www.example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if r.Type != RespAnswer { - t.Errorf("Type = %d, want %d (CNAME with final A answer)", r.Type, RespAnswer) - } -} - -func TestResponseChildReferrals(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeSuccess - msg.Authoritative = false - msg.Ns = append(msg.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS}, Ns: "a.gtld-servers.net."}, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS}, Ns: "b.gtld-servers.net."}, - ) - msg.Extra = append(msg.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dns.TypeA}, A: net.ParseIP("192.5.6.30")}, - &dns.A{Hdr: dns.RR_Header{Name: "b.gtld-servers.net.", Rrtype: dns.TypeA}, A: net.ParseIP("192.33.14.30")}, - ) - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - cache := NewInfoCache(nil) - r := NewResponse(ref, net.ParseIP("198.41.0.4"), cache) - r.Process(msg) - - children := r.ChildReferrals() - if len(children) != 2 { - t.Fatalf("expected 2 child referrals, got %d", len(children)) - } - - if children[0].Name != "example.com." { - t.Errorf("child[0] name = %q, want example.com.", children[0].Name) - } - if children[0].Bailiwick != "a.gtld-servers.net." { - t.Errorf("child[0] bailiwick = %q, want a.gtld-servers.net.", children[0].Bailiwick) - } - if children[0].Prob != 0.5 { - t.Errorf("child[0] prob = %f, want 0.5", children[0].Prob) - } - if children[0].Depth != 1 { - t.Errorf("child[0] depth = %d, want 1", children[0].Depth) - } - if !children[0].HasAddresses() { - t.Error("child[0] should have glue addresses") - } - - if !children[1].HasAddresses() { - t.Error("child[1] should have glue addresses") - } - - nsNames := cache.LookupNS("example.com.") - if len(nsNames) != 2 { - t.Errorf("expected 2 NS in cache, got %d", len(nsNames)) - } -} - -func TestResponseChildReferralsNonReferral(t *testing.T) { - msg := new(dns.Msg) - msg.SetReply(new(dns.Msg)) - msg.Answer = append(msg.Answer, &dns.A{ - Hdr: dns.RR_Header{Rrtype: dns.TypeA}, - A: net.ParseIP("1.2.3.4"), - }) - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if children := r.ChildReferrals(); children != nil { - t.Error("non-referral should not produce child referrals") - } -} - -func TestResponseCNAMEFollowReferral(t *testing.T) { - msg := new(dns.Msg) - msg.SetReply(new(dns.Msg)) - msg.Answer = append(msg.Answer, - &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME}, - Target: "example.com.", - }, - ) - - ref := NewReferral("www.example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - follow := r.CNAMEFollowReferral() - if follow == nil { - t.Fatal("expected CNAME follow referral") - } - if follow.Name != "example.com." { - t.Errorf("follow name = %q, want %q", follow.Name, "example.com.") - } - if follow.Depth != 1 { - t.Errorf("follow depth = %d, want 1", follow.Depth) - } -} - -func TestResponseIsTerminal(t *testing.T) { - tests := []struct { - respType ResponseType - want bool - }{ - {RespAnswer, true}, - {RespNODATA, true}, - {RespNXDOMAIN, true}, - {RespSERVFAIL, true}, - {RespError, true}, - {RespNSResolutionFailed, true}, - {RespReferral, false}, - {RespCNAMEFollow, false}, - } - - for _, tt := range tests { - t.Run(tt.respType.String(), func(t *testing.T) { - r := &Response{Type: tt.respType} - if got := r.IsTerminal(); got != tt.want { - t.Errorf("IsTerminal() = %v, want %v", got, tt.want) - } - }) - } -} - -func TestResponseTypeString(t *testing.T) { - tests := []struct { - rt ResponseType - want string - }{ - {RespReferral, "referral"}, - {RespAnswer, "answer"}, - {RespCNAMEFollow, "cname_follow"}, - {RespNODATA, "nodata"}, - {RespNXDOMAIN, "nxdomain"}, - {RespSERVFAIL, "servfail"}, - {RespError, "error"}, - {RespNSResolutionFailed, "ns_error"}, - } - - for _, tt := range tests { - t.Run(tt.want, func(t *testing.T) { - if got := tt.rt.String(); got != tt.want { - t.Errorf("String() = %q, want %q", got, tt.want) - } - }) - } -} - -func TestResponseChildReferralsProbabilityInheritance(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeSuccess - msg.Authoritative = false - msg.Ns = append(msg.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS}, Ns: "a.gtld-servers.net."}, - &dns.NS{Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS}, Ns: "b.gtld-servers.net."}, - &dns.NS{Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS}, Ns: "c.gtld-servers.net."}, - ) - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 0.5, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - children := r.ChildReferrals() - if len(children) != 3 { - t.Fatalf("expected 3 children, got %d", len(children)) - } - - for _, c := range children { - if c.Prob != 0.5/3.0 { - t.Errorf("child prob = %f, want %f", c.Prob, 0.5/3.0) - } - } -} - -func TestResponseChildReferralsEmptyAuthority(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeSuccess - msg.Authoritative = false - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - children := r.ChildReferrals() - if len(children) != 0 { - t.Errorf("expected 0 children with empty authority, got %d", len(children)) - } -} - -func TestResponseNilReferral(t *testing.T) { - r := NewResponse(nil, net.ParseIP("1.2.3.4"), nil) - children := r.ChildReferrals() - if children != nil { - t.Error("nil referral should produce no children") - } - - follow := r.CNAMEFollowReferral() - if follow != nil { - t.Error("nil referral should produce no CNAME follow") - } -} diff --git a/internal/traverse/robustness_test.go b/internal/traverse/robustness_test.go deleted file mode 100644 index 4e1df9a..0000000 --- a/internal/traverse/robustness_test.go +++ /dev/null @@ -1,768 +0,0 @@ -package traverse - -import ( - "context" - "errors" - "net" - "sync/atomic" - "testing" - - "github.com/miekg/dns" -) - -// TestCNAMELoopDetected verifies that a two-step CNAME loop (A → B → A) is -// detected without infinite recursion and produces a RespCNAMELoop result. -func TestCNAMELoopDetected(t *testing.T) { - // www.example.com → CNAME → alias.example.com → CNAME → www.example.com (loop) - cnameToAlias := new(dns.Msg) - cnameToAlias.SetReply(new(dns.Msg)) - cnameToAlias.Answer = append(cnameToAlias.Answer, &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "alias.example.com.", - }) - - cnameBack := new(dns.Msg) - cnameBack.SetReply(new(dns.Msg)) - cnameBack.Answer = append(cnameBack.Answer, &dns.CNAME{ - Hdr: dns.RR_Header{Name: "alias.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "www.example.com.", - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 10, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - switch q.Name { - case "www.example.com.": - return cnameToAlias.Copy(), nil - case "alias.example.com.": - return cnameBack.Copy(), nil - } - return nil, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "www.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - foundLoop := false - for _, r := range results { - if r.Response != nil && r.Response.Type == RespCNAMELoop { - foundLoop = true - if r.Response.ErrorMessage == "" { - t.Error("expected non-empty ErrorMessage on CNAME loop result") - } - } - } - if !foundLoop { - t.Error("expected RespCNAMELoop result for CNAME loop A → B → A") - } -} - -// TestCNAMEDirectLoop verifies that a direct self-loop (A → A) is handled. -func TestCNAMEDirectLoop(t *testing.T) { - selfLoop := new(dns.Msg) - selfLoop.SetReply(new(dns.Msg)) - selfLoop.Answer = append(selfLoop.Answer, &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "www.example.com.", - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 10, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return selfLoop.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "www.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - foundLoop := false - for _, r := range results { - if r.Response != nil && r.Response.Type == RespCNAMELoop { - foundLoop = true - } - } - if !foundLoop { - t.Error("expected RespCNAMELoop for direct self-referencing CNAME") - } -} - -// TestREFUSEDResponse verifies that a REFUSED rcode is classified as RespREFUSED. -func TestREFUSEDResponse(t *testing.T) { - refusedResp := new(dns.Msg) - refusedResp.Rcode = dns.RcodeRefused - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return refusedResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least 1 result") - } - if results[0].Response.Type != RespREFUSED { - t.Errorf("Type = %s, want refused", results[0].Response.Type) - } - if !results[0].Response.IsTerminal() { - t.Error("REFUSED should be a terminal response") - } -} - -// TestNOTIMPLResponse verifies that a NOTIMP rcode is classified as RespNOTIMPL. -func TestNOTIMPLResponse(t *testing.T) { - notImplResp := new(dns.Msg) - notImplResp.Rcode = dns.RcodeNotImplemented - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return notImplResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least 1 result") - } - if results[0].Response.Type != RespNOTIMPL { - t.Errorf("Type = %s, want notimp", results[0].Response.Type) - } - if !results[0].Response.IsTerminal() { - t.Error("NOTIMP should be a terminal response") - } -} - -// TestGracefulDegradationUnreachableServer verifies that when some servers are -// unreachable, traversal continues with the remaining servers and does not panic. -func TestGracefulDegradationUnreachableServer(t *testing.T) { - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - // Referral with two nameservers; first always fails, second provides the answer. - referralMsg := new(dns.Msg) - referralMsg.Rcode = dns.RcodeSuccess - referralMsg.Authoritative = false - referralMsg.Ns = append(referralMsg.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."}, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."}, - ) - referralMsg.Extra = append(referralMsg.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")}, - &dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")}, - ) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - switch server { - case "198.41.0.4": - return referralMsg.Copy(), nil - case "10.0.0.1": - return nil, errors.New("connection refused") - case "10.0.0.2": - return answerResp.Copy(), nil - } - return nil, errors.New("unexpected server") - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - foundAnswer := false - for _, r := range results { - if r.Response != nil && r.Response.Type == RespAnswer { - foundAnswer = true - } - } - if !foundAnswer { - t.Error("expected an answer result from the reachable server") - } -} - -// TestGracefulDegradationAllUnreachable verifies that when ALL servers fail, -// the traversal returns a SERVFAIL result without panicking. -func TestGracefulDegradationAllUnreachable(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, errors.New("network unreachable") - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("traversal must not return a top-level error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least one result even on total failure") - } - last := results[len(results)-1] - if last.Response == nil { - t.Fatal("last result must have a response") - } - if last.Response.Type != RespSERVFAIL && last.Response.Type != RespError { - t.Errorf("expected SERVFAIL or error when all servers unreachable, got %s", last.Response.Type) - } -} - -// TestDNAMEFollowNoSynthesizedCNAME verifies that a DNAME record in the answer -// section synthesizes a CNAME follow when the server doesn't include one. -func TestDNAMEFollowNoSynthesizedCNAME(t *testing.T) { - // Server returns DNAME only (no synthesized CNAME). - dnameResp := new(dns.Msg) - dnameResp.SetReply(new(dns.Msg)) - dnameResp.Answer = append(dnameResp.Answer, &dns.DNAME{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeDNAME, Class: dns.ClassINET, Ttl: 300}, - Target: "example.net.", - }) - - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "www.example.net.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("203.0.113.1"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - if q.Name == "www.example.com." { - return dnameResp.Copy(), nil - } - if q.Name == "www.example.net." { - return answerResp.Copy(), nil - } - return nil, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "www.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - foundCNAMEFollow := false - for _, r := range results { - if r.Response != nil && r.Response.Type == RespCNAMEFollow { - foundCNAMEFollow = true - } - } - if !foundCNAMEFollow { - t.Error("expected RespCNAMEFollow synthesized from DNAME record") - } -} - -// TestIsNameInChain verifies the ancestor chain lookup. -func TestIsNameInChain(t *testing.T) { - root := NewReferral("example.com", dnsTypeA, ".", 0, 1.0, nil) - child := NewReferral("www.example.com", dnsTypeA, "example.com.", 1, 1.0, root) - grandchild := NewReferral("sub.www.example.com", dnsTypeA, "www.example.com.", 2, 1.0, child) - - tests := []struct { - ref *Referral - name string - want bool - }{ - {grandchild, "sub.www.example.com", true}, // self - {grandchild, "www.example.com", true}, // parent - {grandchild, "example.com", true}, // grandparent - {grandchild, "other.example.com", false}, // not in chain - {root, "example.com", true}, // root matches itself - {root, "www.example.com", false}, // child not in chain from root - } - - for _, tt := range tests { - got := tt.ref.IsNameInChain(tt.name) - if got != tt.want { - t.Errorf("IsNameInChain(%q) from %q = %v, want %v", tt.name, tt.ref.Name, got, tt.want) - } - } -} - -// TestResponseTypeStrings verifies String() for new response types. -func TestResponseTypeStrings(t *testing.T) { - tests := []struct { - rt ResponseType - want string - }{ - {RespReferral, "referral"}, - {RespAnswer, "answer"}, - {RespCNAMEFollow, "cname_follow"}, - {RespNODATA, "nodata"}, - {RespNXDOMAIN, "nxdomain"}, - {RespSERVFAIL, "servfail"}, - {RespREFUSED, "refused"}, - {RespNOTIMPL, "notimp"}, - {RespCNAMELoop, "cname_loop"}, - {RespError, "error"}, - } - - for _, tt := range tests { - if got := tt.rt.String(); got != tt.want { - t.Errorf("ResponseType(%d).String() = %q, want %q", tt.rt, got, tt.want) - } - } -} - -// TestMalformedResponseNoPanic verifies that a nil response from the exchange -// function does not cause a panic, and produces an error result. -func TestMalformedResponseNoPanic(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil // nil response, no error - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least one result") - } - // Should produce error/servfail, not panic - for _, r := range results { - if r.Response == nil { - t.Error("result has nil response") - } - } -} - -// TestDNSSECRRSIGDoesNotBlockCNAMEFollow verifies that a DNSSEC RRSIG record -// accompanying a CNAME in the answer section is treated as metadata and does -// NOT prevent the traversal from following the CNAME. -func TestDNSSECRRSIGDoesNotBlockCNAMEFollow(t *testing.T) { - // Server returns CNAME + RRSIG (DNSSEC-signed zone response). - cnameWithRRSIG := new(dns.Msg) - cnameWithRRSIG.SetReply(new(dns.Msg)) - cnameWithRRSIG.Answer = append(cnameWithRRSIG.Answer, - &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET, Ttl: 300}, - Target: "example.com.", - }, - &dns.RRSIG{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeRRSIG, Class: dns.ClassINET, Ttl: 300}, - TypeCovered: dns.TypeCNAME, - }, - ) - - finalAnswer := new(dns.Msg) - finalAnswer.SetReply(new(dns.Msg)) - finalAnswer.Answer = append(finalAnswer.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 10, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - Fast: true, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - if q.Name == "www.example.com." { - return cnameWithRRSIG.Copy(), nil - } - return finalAnswer.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "www.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - foundCNAMEFollow := false - foundAnswer := false - for _, r := range results { - if r.Response != nil { - switch r.Response.Type { - case RespCNAMEFollow: - foundCNAMEFollow = true - case RespAnswer: - foundAnswer = true - } - } - } - if !foundCNAMEFollow { - t.Error("expected RespCNAMEFollow: RRSIG should not block CNAME following") - } - if !foundAnswer { - t.Error("expected final RespAnswer after CNAME follow") - } -} - -// TestFastModeOn verifies that Fast=true uses the shared root cache (default -// behaviour): a child branch can see glue stored by the root referral. -func TestFastModeOn(t *testing.T) { - // Root referral returns two nameservers with glue. Each NS branch returns - // an answer. We verify both branches are queried. - referralMsg := new(dns.Msg) - referralMsg.Rcode = dns.RcodeSuccess - referralMsg.Authoritative = false - referralMsg.Ns = append(referralMsg.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."}, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."}, - ) - referralMsg.Extra = append(referralMsg.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")}, - &dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")}, - ) - - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - Fast: true, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - if server == "198.41.0.4" { - return referralMsg.Copy(), nil - } - return answerMsg.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - answers := 0 - for _, r := range results { - if r.Response != nil && r.Response.Type == RespAnswer { - answers++ - } - } - if answers == 0 { - t.Error("expected at least one answer with Fast=true") - } -} - -// TestFastModeOff verifies that Fast=false gives each referral its own -// independent cache — no cross-branch glue contamination. -func TestFastModeOff(t *testing.T) { - referralMsg := new(dns.Msg) - referralMsg.Rcode = dns.RcodeSuccess - referralMsg.Authoritative = false - referralMsg.Ns = append(referralMsg.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."}, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."}, - ) - referralMsg.Extra = append(referralMsg.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")}, - &dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")}, - ) - - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - Fast: false, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - if server == "198.41.0.4" { - return referralMsg.Copy(), nil - } - return answerMsg.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - // Traversal must complete without panic and produce results. - if len(results) == 0 { - t.Fatal("expected at least one result with Fast=false") - } -} - -// TestFastModeDefaultIsTrue verifies that DefaultTraverserConfig has Fast=true. -func TestFastModeDefaultIsTrue(t *testing.T) { - cfg := DefaultTraverserConfig() - if !cfg.Fast { - t.Error("DefaultTraverserConfig().Fast should be true") - } -} - -// TestManyNSRecords verifies that a referral with more than 10 nameservers is -// handled gracefully — no panics, results are produced. -func TestManyNSRecords(t *testing.T) { - referralMsg := new(dns.Msg) - referralMsg.Rcode = dns.RcodeSuccess - referralMsg.Authoritative = false - for i := 1; i <= 12; i++ { - ns := &dns.NS{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, - Ns: net.ParseIP(string(rune('a'+i-1))).String() + ".ns.example.com.", - } - // Use a distinct IP for each NS so glue is resolved. - ip := net.IP{10, 0, 0, byte(i)} - referralMsg.Ns = append(referralMsg.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, - Ns: ns.Ns, - }) - referralMsg.Extra = append(referralMsg.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: ns.Ns, Rrtype: dnsTypeA}, - A: ip, - }) - } - - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - var queries int64 - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - Fast: true, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - atomic.AddInt64(&queries, 1) - if server == "198.41.0.4" { - return referralMsg.Copy(), nil - } - return answerMsg.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error with 12 NS records: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results with many NS records") - } - foundAnswer := false - for _, r := range results { - if r.Response != nil && r.Response.Type == RespAnswer { - foundAnswer = true - } - } - if !foundAnswer { - t.Error("expected at least one answer from the 12-NS referral") - } -} - -// TestIDNPunycodeConversion verifies that a unicode (IDN) domain name is -// converted to its punycode/ACE form before querying. -func TestIDNPunycodeConversion(t *testing.T) { - // "münchen.de" → "xn--mnchen-3ya.de" (after punycode encoding) - ref := NewReferral("münchen.de", dnsTypeA, ".", 0, 1.0, nil) - if ref.Name == "münchen.de." { - t.Errorf("IDN name was not converted to punycode: got %q", ref.Name) - } - // Verify it starts with the expected punycode label. - if ref.Name != "xn--mnchen-3ya.de." { - t.Errorf("unexpected punycode result: got %q, want %q", ref.Name, "xn--mnchen-3ya.de.") - } -} - -// TestASCIIDomainUnchanged verifies that a plain ASCII domain is not mangled -// by the IDN conversion path. -func TestASCIIDomainUnchanged(t *testing.T) { - ref := NewReferral("example.com", dnsTypeA, ".", 0, 1.0, nil) - if ref.Name != "example.com." { - t.Errorf("ASCII domain was mangled: got %q, want %q", ref.Name, "example.com.") - } -} - -// TestWildcardResponse verifies that a wildcard answer (e.g. *.example.com -// returning an A record for sub.example.com) is handled as a regular answer. -func TestWildcardResponse(t *testing.T) { - wildcardAnswer := new(dns.Msg) - wildcardAnswer.SetReply(new(dns.Msg)) - wildcardAnswer.Authoritative = true - wildcardAnswer.Answer = append(wildcardAnswer.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "sub.example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return wildcardAnswer.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "sub.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results for wildcard response") - } - if results[0].Response.Type != RespAnswer { - t.Errorf("Type = %s, want answer", results[0].Response.Type) - } -} - -// TestLongCNAMEChainDepthLimit verifies that a very long CNAME chain is -// terminated by the MaxDepth limit without infinite recursion or a panic. -func TestLongCNAMEChainDepthLimit(t *testing.T) { - // Every query returns a CNAME to the next label. The MaxDepth setting - // must stop the chain. - counter := 0 - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - counter++ - q := msg.Question[0] - resp := new(dns.Msg) - resp.SetReply(msg) - next := "next" + q.Name - resp.Answer = append(resp.Answer, &dns.CNAME{ - Hdr: dns.RR_Header{Name: q.Name, Rrtype: dnsTypeCNAME, Class: dns.ClassINET}, - Target: next, - }) - return resp, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "start.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results") - } - // Traversal must have stopped — counter should not be unbounded. - if counter > 50 { - t.Errorf("too many exchange calls (%d): chain depth limit not enforced", counter) - } -} - -// TestPartialBranchFailureReturnsResults verifies the graceful degradation -// requirement: when some NS branches fail completely, the partial results from -// successful branches are still returned. -func TestPartialBranchFailureReturnsResults(t *testing.T) { - // Three nameservers: first two error, third succeeds. - referralMsg := new(dns.Msg) - referralMsg.Rcode = dns.RcodeSuccess - referralMsg.Authoritative = false - referralMsg.Ns = append(referralMsg.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."}, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."}, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns3.example.com."}, - ) - referralMsg.Extra = append(referralMsg.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")}, - &dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")}, - &dns.A{Hdr: dns.RR_Header{Name: "ns3.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.3")}, - ) - - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - switch server { - case "198.41.0.4": - return referralMsg.Copy(), nil - case "10.0.0.1", "10.0.0.2": - return nil, errors.New("server unreachable") - case "10.0.0.3": - return answerMsg.Copy(), nil - } - return nil, errors.New("unexpected server") - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("traversal must not return a top-level error: %v", err) - } - - foundAnswer := false - for _, r := range results { - if r.Response != nil && r.Response.Type == RespAnswer { - foundAnswer = true - } - } - if !foundAnswer { - t.Error("expected an answer from the third (reachable) nameserver despite others failing") - } -} diff --git a/internal/traverse/server_response.go b/internal/traverse/server_response.go new file mode 100644 index 0000000..efd18f8 --- /dev/null +++ b/internal/traverse/server_response.go @@ -0,0 +1,187 @@ +package traverse + +import ( + "fmt" + "sort" + + miekgdns "github.com/miekg/dns" +) + +// ServerResponse wraps a DecodedQuery for one (server, IP) query, mirroring +// response.rb: it owns a child InfoCache seeded with the in-bailiwick records, +// upgrades referral → referral_lame, and computes the start servers for +// referral/restart children. The noglue/loop variants (response_noglue.rb, +// response_loop.rb) are synthetic — no query was sent, DQ is nil. +type ServerResponse struct { + DQ *DecodedQuery + Status Status + + Qname string + Qclass uint16 + Qtype uint16 + IP string + Server string + Bailiwick string + // ParentIP is the address of the referring server; it is part of the + // stats key for referral_lame so lame referrals from different parents + // stay separate. + ParentIP string + + Cache *InfoCache + Starters []StartServer + StartersBailiwick string +} + +// NewServerResponse evaluates a decoded query in the context of parentCache. +// server is the NS hostname that was queried (dq.IP is its address). +func NewServerResponse(dq *DecodedQuery, server, parentIP string, parentCache *InfoCache) (*ServerResponse, error) { + r := &ServerResponse{ + DQ: dq, + Status: dq.Status, + Qname: dq.Qname, + Qclass: dq.Qclass, + Qtype: dq.Qtype, + IP: dq.IP, + Server: canonicalName(server), + Bailiwick: dq.Bailiwick, + ParentIP: parentIP, + Cache: NewInfoCache(parentCache), + } + if err := r.evaluate(); err != nil { + return nil, err + } + return r, nil +} + +// NewNoGlueResponse records a dead end: ip referred us to server inside +// bailiwick without glue and there is no way to resolve it. +func NewNoGlueResponse(qname string, qclass, qtype uint16, ip, server, bailiwick string) *ServerResponse { + return &ServerResponse{ + Status: StatusNoGlue, + Qname: canonicalName(qname), + Qclass: qclass, + Qtype: qtype, + IP: ip, + Server: canonicalName(server), + Bailiwick: canonicalName(bailiwick), + } +} + +// NewLoopResponse records a dead end: resolving server from ip would repeat +// an ancestor referral. +func NewLoopResponse(qname string, qclass, qtype uint16, ip, server, bailiwick string) *ServerResponse { + return &ServerResponse{ + Status: StatusLoop, + Qname: canonicalName(qname), + Qclass: qclass, + Qtype: qtype, + IP: ip, + Server: canonicalName(server), + Bailiwick: canonicalName(bailiwick), + } +} + +// evaluate mirrors response.rb#evaluate: cache the in-bailiwick records, then +// for referral/restart work out the start servers from THIS branch's cache; +// a referral whose new zone is not strictly deeper than the bailiwick is lame. +func (r *ServerResponse) evaluate() error { + if r.Status != StatusException { + r.Cache.Add(r.DQ.CacheableGood) + } + switch r.DQ.Status { + case StatusRestart: + starters, bw, err := r.Cache.GetStartServers(r.DQ.Endname) + if err != nil { + return err + } + r.Starters, r.StartersBailiwick = starters, bw + case StatusReferral: + starters, bw, err := r.Cache.GetStartServers(r.DQ.Endname) + if err != nil { + return err + } + r.Starters, r.StartersBailiwick = starters, bw + if isLameReferral(r.DQ.Bailiwick, bw) { + r.Status = StatusReferralLame + } + starterNames := make([]string, len(starters)) + for i, s := range starters { + starterNames[i] = s.Name + } + if !equalSorted(starterNames, r.DQ.AuthorityNames) { + r.DQ.WarningsAdd("Referred authority names do not match query cache expectations") + } + } + return nil +} + +func equalSorted(a, b []string) bool { + if len(a) != len(b) { + return false + } + as := append([]string(nil), a...) + bs := append([]string(nil), b...) + sort.Strings(as) + sort.Strings(bs) + for i := range as { + if as[i] != bs[i] { + return false + } + } + return true +} + +// StatsKey is the leaf aggregation key (response.rb update_stats_key and the +// noglue/loop variants): identical keys merge by summing probability. +func (r *ServerResponse) StatsKey() string { + qclass := ClassToString(r.Qclass) + qtype := TypeToString(r.Qtype) + switch r.Status { + case StatusNoGlue, StatusLoop: + return fmt.Sprintf("key:%s:%s:%s:%s:%s:%s:%s", + r.Status, r.IP, r.Qname, qclass, qtype, r.Server, r.Bailiwick) + default: + key := fmt.Sprintf("key:%s:%s:%s:%s:%s:%s", + r.Status, r.IP, r.Server, r.Qname, qclass, qtype) + if r.Status == StatusException && r.DQ != nil { + key += ":" + r.DQ.ExceptionMessage + } else if r.Status == StatusReferralLame { + key += ":" + r.ParentIP + } + return key + } +} + +// String renders a short description for progress display (the Ruby +// response to_s variants: "No glue for X" / "Loop encountered resolving X"). +func (r *ServerResponse) String() string { + switch r.Status { + case StatusNoGlue: + return fmt.Sprintf("No glue for %s", r.Server) + case StatusLoop: + return fmt.Sprintf("Loop encountered resolving %s", r.Server) + case StatusException: + if r.DQ != nil { + return r.DQ.ExceptionMessage + } + case StatusError: + if r.DQ != nil { + return r.DQ.ErrorMessage + } + } + return string(r.Status) +} + +func ClassToString(qclass uint16) string { + if s, ok := miekgdns.ClassToString[qclass]; ok { + return s + } + return fmt.Sprintf("CLASS%d", qclass) +} + +func TypeToString(qtype uint16) string { + if s, ok := miekgdns.TypeToString[qtype]; ok { + return s + } + return fmt.Sprintf("TYPE%d", qtype) +} diff --git a/internal/traverse/server_response_test.go b/internal/traverse/server_response_test.go new file mode 100644 index 0000000..959157a --- /dev/null +++ b/internal/traverse/server_response_test.go @@ -0,0 +1,232 @@ +package traverse + +import ( + "testing" + + "github.com/miekg/dns" +) + +// rootedCache returns a cache seeded with root hints, as the traverser will +// always provide. +func rootedCache() *InfoCache { + c := NewInfoCache(nil) + c.AddHints("", []StartServer{{Name: "a.root-servers.net", IPs: []string{"198.41.0.4"}}}) + return c +} + +func TestServerResponseReferralNotLame(t *testing.T) { + // Root server refers com query to the gtld servers: "" → "com" is deeper. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, nsRR("com", "a.gtld-servers.net")) + msg.Extra = append(msg.Extra, aRR("a.gtld-servers.net", "192.5.6.30")) + dq := decode(msg, "www.example.com", dns.TypeA, "") + + r, err := NewServerResponse(dq, "a.root-servers.net", "", rootedCache()) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if r.Status != StatusReferral { + t.Fatalf("status = %s, want referral", r.Status) + } + if r.StartersBailiwick != "com" { + t.Errorf("starters bailiwick = %q, want com", r.StartersBailiwick) + } + if len(r.Starters) != 1 || r.Starters[0].Name != "a.gtld-servers.net" { + t.Errorf("starters = %v", r.Starters) + } + if len(r.Starters[0].IPs) != 1 || r.Starters[0].IPs[0] != "192.5.6.30" { + t.Errorf("starter IPs = %v", r.Starters[0].IPs) + } + if len(dq.Warnings) != 0 { + t.Errorf("unexpected warnings: %v", dq.Warnings) + } +} + +func TestServerResponseGluelessStarterHasNilIPs(t *testing.T) { + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, nsRR("example.com", "ns1.example.com")) + dq := decode(msg, "www.example.com", dns.TypeA, "com") + + r, err := NewServerResponse(dq, "a.gtld-servers.net", "", rootedCache()) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if r.Status != StatusReferral { + t.Fatalf("status = %s, want referral", r.Status) + } + if r.Starters[0].IPs != nil { + t.Errorf("glueless starter should have nil IPs, got %v", r.Starters[0].IPs) + } +} + +func TestServerResponseLameReferral(t *testing.T) { + // A com server "refers" us to an example.org zone: the NS records are + // out-of-bailiwick so they are discarded, the cache walk falls back to + // the root NS, and "" is not strictly deeper than "com" → lame. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, nsRR("example.org", "ns1.example.org")) + dq := decode(msg, "www.example.com", dns.TypeA, "com") + if dq.Status != StatusReferral { + t.Fatalf("decoded status = %s, want referral", dq.Status) + } + + r, err := NewServerResponse(dq, "a.gtld-servers.net", "192.5.6.30", rootedCache()) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if r.Status != StatusReferralLame { + t.Fatalf("status = %s, want referral_lame", r.Status) + } + if r.StartersBailiwick != "" { + t.Errorf("starters bailiwick = %q, want \"\" (root fallback)", r.StartersBailiwick) + } + found := false + for _, w := range dq.Warnings { + if w == "Referred authority names do not match query cache expectations" { + found = true + } + } + if !found { + t.Errorf("expected mismatch warning, got %v", dq.Warnings) + } + want := "key:referral_lame:192.0.2.1:a.gtld-servers.net:www.example.com:IN:A:192.5.6.30" + if got := r.StatsKey(); got != want { + t.Errorf("stats key = %q, want %q", got, want) + } +} + +func TestServerResponseEqualZoneReferralIsLame(t *testing.T) { + // Referral back into the SAME zone (com → com) is lame: not strictly deeper. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, nsRR("com", "b.gtld-servers.net")) + dq := decode(msg, "www.example.com", dns.TypeA, "com") + + r, err := NewServerResponse(dq, "a.gtld-servers.net", "192.5.6.30", rootedCache()) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if r.Status != StatusReferralLame { + t.Fatalf("status = %s, want referral_lame", r.Status) + } +} + +func TestServerResponseRestartStarters(t *testing.T) { + // A CNAME out of the bailiwick restarts; starters come from the deepest + // cached zone for the new target (root here). + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, cnameRR("www.example.com", "cdn.example.org")) + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + if dq.Status != StatusRestart { + t.Fatalf("decoded status = %s, want restart", dq.Status) + } + + parent := rootedCache() + parent.Add([]dns.RR{nsRR("example.org", "ns1.example.org"), aRR("ns1.example.org", "9.9.9.9")}) + r, err := NewServerResponse(dq, "ns1.example.com", "", parent) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if r.Status != StatusRestart { + t.Fatalf("status = %s, want restart", r.Status) + } + if r.StartersBailiwick != "example.org" { + t.Errorf("starters bailiwick = %q, want example.org", r.StartersBailiwick) + } + if len(r.Starters) != 1 || r.Starters[0].Name != "ns1.example.org" { + t.Errorf("starters = %v", r.Starters) + } +} + +func TestServerResponseCachesGoodRecordsInChildCache(t *testing.T) { + parent := rootedCache() + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, nsRR("example.com", "ns1.example.com")) + msg.Extra = append(msg.Extra, + aRR("ns1.example.com", "1.2.3.4"), + aRR("ns1.example.org", "5.6.7.8"), // out of bailiwick — discarded + ) + dq := decode(msg, "www.example.com", dns.TypeA, "com") + + r, err := NewServerResponse(dq, "a.gtld-servers.net", "", parent) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if got := r.Cache.Get("ns1.example.com", dns.ClassINET, dns.TypeA); len(got) != 1 { + t.Errorf("in-bailiwick glue should be cached, got %v", got) + } + if got := r.Cache.Get("ns1.example.org", dns.ClassINET, dns.TypeA); got != nil { + t.Errorf("out-of-bailiwick record must be discarded, got %v", got) + } + // The parent cache stays clean — records live in the response's child. + if got := parent.Get("example.com", dns.ClassINET, dns.TypeNS); got != nil { + t.Errorf("parent cache polluted: %v", got) + } +} + +func TestServerResponseExceptionDoesNotCache(t *testing.T) { + dq := NewDecodedQuery(nil, errTimeout{}, "www.example.com", dns.ClassINET, dns.TypeA, "192.0.2.1", "com") + r, err := NewServerResponse(dq, "a.gtld-servers.net", "", rootedCache()) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if r.Status != StatusException { + t.Fatalf("status = %s, want exception", r.Status) + } + want := "key:exception:192.0.2.1:a.gtld-servers.net:www.example.com:IN:A:query timed out" + if got := r.StatsKey(); got != want { + t.Errorf("stats key = %q, want %q", got, want) + } +} + +type errTimeout struct{} + +func (errTimeout) Error() string { return "query timed out" } + +func TestServerResponseAnsweredStatsKey(t *testing.T) { + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, aRR("www.example.com", "93.184.216.34")) + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + + r, err := NewServerResponse(dq, "NS1.Example.Com", "", rootedCache()) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + want := "key:answered:192.0.2.1:ns1.example.com:www.example.com:IN:A" + if got := r.StatsKey(); got != want { + t.Errorf("stats key = %q, want %q", got, want) + } +} + +func TestNoGlueResponse(t *testing.T) { + r := NewNoGlueResponse("www.example.com", dns.ClassINET, dns.TypeA, "192.5.6.30", "ns1.example.com", "example.com") + if r.Status != StatusNoGlue { + t.Fatalf("status = %s, want noglue", r.Status) + } + // NoGlue/Loop use their own field order: ip, qname, qclass, qtype, server, bailiwick. + want := "key:noglue:192.5.6.30:www.example.com:IN:A:ns1.example.com:example.com" + if got := r.StatsKey(); got != want { + t.Errorf("stats key = %q, want %q", got, want) + } +} + +func TestLoopResponse(t *testing.T) { + r := NewLoopResponse("www.example.com", dns.ClassINET, dns.TypeA, "192.5.6.30", "ns1.example.com", "example.com") + if r.Status != StatusLoop { + t.Fatalf("status = %s, want loop", r.Status) + } + want := "key:loop:192.5.6.30:www.example.com:IN:A:ns1.example.com:example.com" + if got := r.StatsKey(); got != want { + t.Errorf("stats key = %q, want %q", got, want) + } +} + +func TestServerResponseReferralNoRootHintsErrors(t *testing.T) { + // A lame referral with a completely empty cache chain cannot compute + // starters; the constructor surfaces the "no root hints" error. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, nsRR("example.org", "ns1.example.org")) + dq := decode(msg, "www.example.com", dns.TypeA, "com") + if _, err := NewServerResponse(dq, "a.gtld-servers.net", "", NewInfoCache(nil)); err == nil { + t.Fatal("expected error when no NS reachable in cache chain") + } +} diff --git a/internal/traverse/stack.go b/internal/traverse/stack.go deleted file mode 100644 index 7e09cd3..0000000 --- a/internal/traverse/stack.go +++ /dev/null @@ -1,58 +0,0 @@ -package traverse - -const DefaultMaxDepth = 20 - -type Stack struct { - items []*Referral - maxDepth int -} - -func NewStack(maxDepth int) *Stack { - if maxDepth <= 0 { - maxDepth = DefaultMaxDepth - } - return &Stack{ - items: make([]*Referral, 0), - maxDepth: maxDepth, - } -} - -func (s *Stack) Push(r *Referral) bool { - if r == nil { - return false - } - if r.Depth >= s.maxDepth { - return false - } - s.items = append(s.items, r) - return true -} - -func (s *Stack) Pop() *Referral { - if len(s.items) == 0 { - return nil - } - idx := len(s.items) - 1 - item := s.items[idx] - s.items = s.items[:idx] - return item -} - -func (s *Stack) Peek() *Referral { - if len(s.items) == 0 { - return nil - } - return s.items[len(s.items)-1] -} - -func (s *Stack) Len() int { - return len(s.items) -} - -func (s *Stack) MaxDepth() int { - return s.maxDepth -} - -func (s *Stack) IsEmpty() bool { - return len(s.items) == 0 -} diff --git a/internal/traverse/stack_test.go b/internal/traverse/stack_test.go deleted file mode 100644 index d1c9f70..0000000 --- a/internal/traverse/stack_test.go +++ /dev/null @@ -1,143 +0,0 @@ -package traverse - -import ( - "testing" - - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" -) - -func TestNewStack(t *testing.T) { - s := NewStack(10) - if s.MaxDepth() != 10 { - t.Errorf("MaxDepth = %d, want 10", s.MaxDepth()) - } - if !s.IsEmpty() { - t.Error("new stack should be empty") - } - if s.Len() != 0 { - t.Errorf("Len = %d, want 0", s.Len()) - } -} - -func TestNewStackDefaultDepth(t *testing.T) { - s := NewStack(0) - if s.MaxDepth() != DefaultMaxDepth { - t.Errorf("MaxDepth = %d, want %d", s.MaxDepth(), DefaultMaxDepth) - } - - s = NewStack(-5) - if s.MaxDepth() != DefaultMaxDepth { - t.Errorf("MaxDepth = %d, want %d", s.MaxDepth(), DefaultMaxDepth) - } -} - -func TestStackPushPop(t *testing.T) { - s := NewStack(5) - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - - ok := s.Push(ref) - if !ok { - t.Error("Push should succeed") - } - if s.Len() != 1 { - t.Errorf("Len = %d, want 1", s.Len()) - } - - popped := s.Pop() - if popped != ref { - t.Error("popped referral should match pushed") - } - if s.Len() != 0 { - t.Errorf("Len = %d, want 0", s.Len()) - } -} - -func TestStackLIFO(t *testing.T) { - s := NewStack(5) - r1 := NewReferral("a.com", dns.TypeA, ".", 0, 1.0, nil) - r2 := NewReferral("b.com", dns.TypeA, ".", 1, 1.0, nil) - r3 := NewReferral("c.com", dns.TypeA, ".", 2, 1.0, nil) - - s.Push(r1) - s.Push(r2) - s.Push(r3) - - if popped := s.Pop(); popped != r3 { - t.Error("should pop r3 first (LIFO)") - } - if popped := s.Pop(); popped != r2 { - t.Error("should pop r2 second") - } - if popped := s.Pop(); popped != r1 { - t.Error("should pop r1 third") - } -} - -func TestStackPushNil(t *testing.T) { - s := NewStack(5) - ok := s.Push(nil) - if ok { - t.Error("Push(nil) should return false") - } - if s.Len() != 0 { - t.Errorf("Len = %d, want 0", s.Len()) - } -} - -func TestStackMaxDepth(t *testing.T) { - s := NewStack(3) - - r0 := NewReferral("a.com", dns.TypeA, ".", 0, 1.0, nil) - r1 := NewReferral("b.com", dns.TypeA, ".", 1, 1.0, nil) - r2 := NewReferral("c.com", dns.TypeA, ".", 2, 1.0, nil) - r3 := NewReferral("d.com", dns.TypeA, ".", 3, 1.0, nil) - - if !s.Push(r0) { - t.Error("depth 0 should be accepted") - } - if !s.Push(r1) { - t.Error("depth 1 should be accepted") - } - if !s.Push(r2) { - t.Error("depth 2 should be accepted") - } - if s.Push(r3) { - t.Error("depth 3 should be rejected (maxDepth=3)") - } -} - -func TestStackPopEmpty(t *testing.T) { - s := NewStack(5) - if popped := s.Pop(); popped != nil { - t.Error("Pop on empty stack should return nil") - } -} - -func TestStackPeek(t *testing.T) { - s := NewStack(5) - if peek := s.Peek(); peek != nil { - t.Error("Peek on empty stack should return nil") - } - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - s.Push(ref) - - if peek := s.Peek(); peek != ref { - t.Error("Peek should return top item") - } - if s.Len() != 1 { - t.Errorf("Peek should not remove item, Len = %d, want 1", s.Len()) - } -} - -func TestStackIsEmpty(t *testing.T) { - s := NewStack(5) - if !s.IsEmpty() { - t.Error("new stack should be empty") - } - - s.Push(NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)) - if s.IsEmpty() { - t.Error("stack with item should not be empty") - } -} diff --git a/internal/traverse/stats.go b/internal/traverse/stats.go new file mode 100644 index 0000000..1ad805c --- /dev/null +++ b/internal/traverse/stats.go @@ -0,0 +1,74 @@ +package traverse + +import ( + "sort" + "strings" + + miekgdns "github.com/miekg/dns" +) + +// AnswerStat is one distinct answered RRset with its accumulated probability. +type AnswerStat struct { + // Key groups identical RRset content: sorted rdata strings joined with + // "@@@" (summary_stats.rb get_answer_stats). + Key string + Prob float64 + RRs []miekgdns.RR +} + +// SummaryStats groups the aggregated leaves by status; answered leaves are +// additionally grouped by RRset content, one entry per distinct RRset +// (summary_stats.rb). The ByStatus probabilities sum to 1.0 and the answer +// probabilities sum to ByStatus[StatusAnswered]. +type SummaryStats struct { + ByStatus map[Status]float64 + Answers []AnswerStat +} + +// SummaryStats computes (and memoises) the summary grouping of this node's +// aggregated leaf statistics (referral.rb summary_stats). It returns nil +// until the node's statistics have been calculated. +func (r *Referral) SummaryStats() *SummaryStats { + if r == nil || !r.calculated || len(r.Stats) == 0 { + return nil + } + if r.summaryStats != nil { + return r.summaryStats + } + + stats := &SummaryStats{ByStatus: make(map[Status]float64)} + answers := make(map[string]*AnswerStat) + for _, leaf := range r.StatsList() { + status := leaf.Response.Status + stats.ByStatus[status] += leaf.Prob + if status != StatusAnswered { + continue + } + rdatas := make([]string, 0, len(leaf.Response.DQ.Answers)) + for _, rr := range leaf.Response.DQ.Answers { + rdatas = append(rdatas, rrData(rr)) + } + sort.Strings(rdatas) + key := strings.Join(rdatas, "@@@") + if e, ok := answers[key]; ok { + e.Prob += leaf.Prob + } else { + answers[key] = &AnswerStat{Key: key, Prob: leaf.Prob, RRs: leaf.Response.DQ.Answers} + } + } + + for _, e := range answers { + stats.Answers = append(stats.Answers, *e) + } + sort.Slice(stats.Answers, func(i, j int) bool { return stats.Answers[i].Key < stats.Answers[j].Key }) + r.summaryStats = stats + return stats +} + +// rrData extracts the rdata portion of a record (dnsruby rdata_to_string): +// everything after owner/TTL/class/type in presentation format. +func rrData(rr miekgdns.RR) string { + s := rr.String() + h := rr.Header().String() + return strings.TrimPrefix(s, h) +} diff --git a/internal/traverse/stats_test.go b/internal/traverse/stats_test.go new file mode 100644 index 0000000..743d7c1 --- /dev/null +++ b/internal/traverse/stats_test.go @@ -0,0 +1,302 @@ +package traverse + +import ( + "math" + "net" + "strconv" + "strings" + "testing" + + "github.com/miekg/dns" +) + +// mockCaptureTopology reproduces the delegation behind +// docs/captures/dnstraverse-ruby-www.example.com-A.txt: one root, thirteen +// com gTLD servers, example.com served by two NS with three IPs each, every +// endpoint answering the same two A records (two endpoints return them in +// the opposite order, as in the capture). +func mockCaptureTopology() *mockExchange { + m := newMockExchange() + + gtlds := []string{"a", "b", "c", "d", "e", "f", "g", "h", "i", "j", "k", "l", "m"} + var comNS, comGlue []dns.RR + gtldIPs := make([]string, len(gtlds)) + for i, l := range gtlds { + ip := "192.0.2." + strconv.Itoa(i+1) + gtldIPs[i] = ip + comNS = append(comNS, nsRR("com", l+".gtld-servers.net")) + comGlue = append(comGlue, aRR(l+".gtld-servers.net", ip)) + } + m.on("202.12.27.33", "www.example.com", dns.TypeA, referralMsg(comNS, comGlue...)) + + heraIPs := []string{"108.162.192.162", "172.64.32.162", "173.245.58.162"} + elliottIPs := []string{"108.162.195.228", "162.159.44.228", "172.64.35.228"} + exampleReferral := referralMsg( + []dns.RR{ + nsRR("example.com", "hera.ns.cloudflare.com"), + nsRR("example.com", "elliott.ns.cloudflare.com"), + }, + aRR("hera.ns.cloudflare.com", heraIPs[0]), + aRR("hera.ns.cloudflare.com", heraIPs[1]), + aRR("hera.ns.cloudflare.com", heraIPs[2]), + aRR("elliott.ns.cloudflare.com", elliottIPs[0]), + aRR("elliott.ns.cloudflare.com", elliottIPs[1]), + aRR("elliott.ns.cloudflare.com", elliottIPs[2]), + ) + for _, ip := range gtldIPs { + m.on(ip, "www.example.com", dns.TypeA, exampleReferral) + } + + forward := answerMsg( + aRR("www.example.com", "104.20.23.154"), + aRR("www.example.com", "172.66.147.243"), + ) + reversed := answerMsg( + aRR("www.example.com", "172.66.147.243"), + aRR("www.example.com", "104.20.23.154"), + ) + for _, ip := range []string{heraIPs[0], elliottIPs[0], elliottIPs[1], elliottIPs[2]} { + m.on(ip, "www.example.com", dns.TypeA, forward) + } + for _, ip := range []string{heraIPs[1], heraIPs[2]} { + m.on(ip, "www.example.com", dns.TypeA, reversed) + } + return m +} + +// TestCapturePerEndpointFractions asserts the 16.7%-per-endpoint result of +// the www.example.com capture: 13 gTLD paths collapse (fast mode) into six +// endpoint leaves of 1/6 each, and the summary merges the differently +// ordered RRsets into a single 100% answered line. +func TestCapturePerEndpointFractions(t *testing.T) { + m := mockCaptureTopology() + cfg := testConfig(true) + cfg.RootAddrs = []net.IP{net.ParseIP("202.12.27.33")} + _, root := runTraversal(t, cfg, m, "www.example.com") + + assertSumsToOne(t, root) + answered := leavesByStatus(root, StatusAnswered) + if len(answered) != 6 { + t.Fatalf("expected 6 answered leaves (one per endpoint), got %d: %v", len(answered), root.StatsList()) + } + for _, leaf := range answered { + if math.Abs(leaf.Prob-1.0/6) > 1e-9 { + t.Errorf("leaf %s prob = %v, want 1/6", leaf.Key, leaf.Prob) + } + } + + stats := root.SummaryStats() + if stats == nil { + t.Fatal("expected summary stats after calculation") + } + if prob := stats.ByStatus[StatusAnswered]; math.Abs(prob-1.0) > 1e-9 { + t.Errorf("answered summary prob = %v, want 1.0", prob) + } + // Both RR orders share the same sorted-rdata key: one summary line. + if len(stats.Answers) != 1 { + t.Fatalf("expected 1 distinct answered RRset, got %d: %v", len(stats.Answers), stats.Answers) + } + if math.Abs(stats.Answers[0].Prob-1.0) > 1e-9 { + t.Errorf("answer group prob = %v, want 1.0", stats.Answers[0].Prob) + } + if !strings.Contains(stats.Answers[0].Key, "@@@") { + t.Errorf("answer key should join rdata with @@@, got %q", stats.Answers[0].Key) + } + if len(stats.Answers[0].RRs) != 2 { + t.Errorf("answer RRs = %v", stats.Answers[0].RRs) + } +} + +// TestDistinctRRsetsSeparateGroups asserts the converse of the capture test: +// two endpoints answering DIFFERENT content produce two separate summary +// groups, each carrying its own share of the answered probability. +func TestDistinctRRsetsSeparateGroups(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns2.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + aRR("ns2.example.com", "2.2.2.2"), + )) + // Both RRsets share their first sorted rdata so grouping must consider + // the full content, not just the first record. + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg( + aRR("www.example.com", "1.0.0.1"), + aRR("www.example.com", "9.9.9.9"), + )) + m.on("2.2.2.2", "www.example.com", dns.TypeA, answerMsg( + aRR("www.example.com", "1.0.0.1"), + aRR("www.example.com", "8.8.8.8"), + )) + + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + stats := root.SummaryStats() + if stats == nil { + t.Fatal("expected summary stats after calculation") + } + if prob := stats.ByStatus[StatusAnswered]; math.Abs(prob-1.0) > 1e-9 { + t.Errorf("answered summary prob = %v, want 1.0", prob) + } + if len(stats.Answers) != 2 { + t.Fatalf("expected 2 distinct answered RRsets, got %d: %v", len(stats.Answers), stats.Answers) + } + if stats.Answers[0].Key == stats.Answers[1].Key { + t.Errorf("answer groups share key %q, want distinct keys", stats.Answers[0].Key) + } + for i, ans := range stats.Answers { + if math.Abs(ans.Prob-0.5) > 1e-9 { + t.Errorf("answer group %d (%q) prob = %v, want 0.5", i, ans.Key, ans.Prob) + } + } + // Answers are sorted by key: "1.0.0.1@@@8.8.8.8" then "1.0.0.1@@@9.9.9.9". + wantRdata := []string{"8.8.8.8", "9.9.9.9"} + for i, ans := range stats.Answers { + if !strings.Contains(ans.Key, "1.0.0.1@@@"+wantRdata[i]) { + t.Errorf("answer group %d key = %q, want it to contain %q", i, ans.Key, "1.0.0.1@@@"+wantRdata[i]) + } + if len(ans.RRs) != 2 { + t.Errorf("answer group %d RRs = %v, want 2 records", i, ans.RRs) + } + } +} + +func TestServfailErrorLeaf(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns2.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + aRR("ns2.example.com", "2.2.2.2"), + )) + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + m.on("2.2.2.2", "www.example.com", dns.TypeA, rcodeMsg(dns.RcodeServerFailure)) + + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + errs := leavesByStatus(root, StatusError) + if len(errs) != 1 { + t.Fatalf("expected 1 error leaf, got %v", root.StatsList()) + } + if errs[0].Response.DQ.ErrorMessage != "Server failure (SERVFAIL)" { + t.Errorf("error message = %q", errs[0].Response.DQ.ErrorMessage) + } + if math.Abs(errs[0].Prob-0.5) > 1e-9 { + t.Errorf("error prob = %v, want 0.5", errs[0].Prob) + } + + stats := root.SummaryStats() + if math.Abs(stats.ByStatus[StatusError]-0.5) > 1e-9 || + math.Abs(stats.ByStatus[StatusAnswered]-0.5) > 1e-9 { + t.Errorf("summary by status = %v", stats.ByStatus) + } + total := 0.0 + for _, prob := range stats.ByStatus { + total += prob + } + if math.Abs(total-1.0) > 1e-9 { + t.Errorf("summary probabilities sum to %v, want 1.0", total) + } +} + +// TestResolveSubtreeLeavesExcluded asserts that resolve-subtree leaves (the +// A lookups for glueless NS) never reach the main aggregation: +// they surface only through server weights. +func TestResolveSubtreeLeavesExcluded(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns.other.net")}, + aRR("ns1.example.com", "1.1.1.1"), + )) + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + m.on("198.41.0.4", "ns.other.net", dns.TypeA, answerMsg(aRR("ns.other.net", "4.4.4.4"))) + m.on("4.4.4.4", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + for _, leaf := range root.StatsList() { + if leaf.Response.Qname == "ns.other.net" { + t.Errorf("resolve-subtree leaf leaked into main aggregation: %s", leaf.Key) + } + } + if prob := root.SummaryStats().ByStatus[StatusAnswered]; math.Abs(prob-1.0) > 1e-9 { + t.Errorf("answered summary prob = %v, want 1.0", prob) + } +} + +// TestResolveFailurePseudoIPCarriesMass asserts that a failed glue resolution +// keeps its probability: the failure becomes a "key:" pseudo-IP whose mass +// surfaces in the main aggregation as the failing (resolve) query. +func TestResolveFailurePseudoIPCarriesMass(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns.other.net")}, + aRR("ns1.example.com", "1.1.1.1"), + )) + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + // The resolve of A ns.other.net fails at the root: SERVFAIL. + m.on("198.41.0.4", "ns.other.net", dns.TypeA, rcodeMsg(dns.RcodeServerFailure)) + + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + errs := leavesByStatus(root, StatusError) + if len(errs) != 1 { + t.Fatalf("expected 1 error leaf from the failed resolve, got %v", root.StatsList()) + } + leaf := errs[0] + if math.Abs(leaf.Prob-0.5) > 1e-9 { + t.Errorf("failed-resolve prob = %v, want 0.5", leaf.Prob) + } + if leaf.Response.Qname != "ns.other.net" { + t.Errorf("failed-resolve leaf qname = %q, want the resolve target", leaf.Response.Qname) + } + if !strings.HasPrefix(leaf.Key, "key:error:") { + t.Errorf("failed-resolve key = %q", leaf.Key) + } + // The leaf's referral is the resolve-subtree node; its parent is the + // glueless referral, which carries the mass as a pseudo-IP server entry. + glueless := leaf.Referral.Parent + if glueless.Server != "ns.other.net" { + t.Fatalf("glueless referral server = %q", glueless.Server) + } + hasPseudo := false + for ip, weight := range glueless.ServerWeights { + if strings.HasPrefix(ip, "key:") && math.Abs(weight-1.0) <= 1e-9 { + hasPseudo = true + } + } + if !hasPseudo { + t.Errorf("expected a key: pseudo-IP with weight 1.0, got %v", glueless.ServerWeights) + } +} + +func TestSummaryStatsNilAndMemoised(t *testing.T) { + var nilRef *Referral + if nilRef.SummaryStats() != nil { + t.Error("nil referral should produce nil summary") + } + uncalculated := newTestReferral("ns1.example.com", []string{"1.1.1.1"}) + if uncalculated.SummaryStats() != nil { + t.Error("uncalculated referral should produce nil summary") + } + + m := mockSimpleDelegation() + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + first := root.SummaryStats() + if first == nil { + t.Fatal("expected summary stats") + } + if root.SummaryStats() != first { + t.Error("summary stats should be memoised") + } +} diff --git a/internal/traverse/traverse.go b/internal/traverse/traverse.go index 83914e1..284d576 100644 --- a/internal/traverse/traverse.go +++ b/internal/traverse/traverse.go @@ -1,27 +1,29 @@ -// Package traverse implements the core DNS traversal engine for ExploreDNS. +// Package traverse implements the core DNS traversal engine for ExploreDNS, +// a Go port of the Ruby dnstraverse engine (dns.squish.net). // -// The traversal engine starts from the DNS root servers and iteratively -// follows every referral it receives, building a complete picture of the -// delegation path for a domain. Unlike a standard recursive resolver, which -// stops at the first authoritative answer, the traversal engine explores every -// branch so that delegation mismatches, lame delegations, or split authorities -// are all visible in the output. +// The traversal starts from a synthetic "rootroot" node (never displayed) +// with one child per root server, and explores every branch of the +// delegation instead of stopping at the first authoritative answer, so lame +// delegations, missing glue and split authorities are all visible. // // # Architecture // -// A Traverser maintains a stack of Referral objects. Each Referral -// represents a pending query to a specific set of nameservers for a specific -// name and record type. The engine pops referrals one at a time, sends the -// query, classifies the response, and pushes any child referrals back onto the -// stack. -// -// When a referral contains nameserver names but no glue records (IP addresses), -// the engine resolves them via a secondary traversal before continuing. +// A Traverser runs an explicit stack loop over Referral nodes. Each Referral +// queries every IP address of one nameserver for one qname/qclass/qtype, +// classifies each response (DecodedQuery, ServerResponse) and creates one +// child per NS name for referral/restart statuses — including glueless +// nameservers, which get their own resolve subtree (refid ".0." components) +// queried from this branch's cache, never a system resolver. Post-order +// stack markers fold the statistics upwards once all children finished: +// every leaf outcome carries a probability, and the probabilities at the +// root sum to 1.0. // // # Caching // -// An InfoCache stores discovered glue records. In fast mode (default) a -// single root cache is shared across all branches so that glue discovered in -// one branch is immediately available to sibling branches. Disable fast mode -// (TraverserConfig.Fast = false) for fully independent branch resolution. +// Two caches cooperate: the packet-level cache in internal/dns sends each +// (server IP, question, udpsize) at most once per run, and the hierarchical +// per-branch InfoCache holds the in-bailiwick records each response is +// allowed to contribute. Fast mode (default) additionally memoises completed +// referrals so identical subtrees are reported as "completed earlier" +// instead of being walked again. package traverse diff --git a/internal/traverse/traverser.go b/internal/traverse/traverser.go index e3576bd..365fbe5 100644 --- a/internal/traverse/traverser.go +++ b/internal/traverse/traverser.go @@ -4,9 +4,8 @@ import ( "context" "fmt" "net" + "sort" "strings" - "sync" - "time" "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" miekgdns "github.com/miekg/dns" @@ -14,495 +13,329 @@ import ( // TraverserConfig configures the behaviour of a Traverser. type TraverserConfig struct { - // MaxDepth is the maximum referral depth before the traversal gives up. - MaxDepth int + // MaxDepth is the maximum referral depth (non-zero refid components) + // before a "Maxdepth N exceeded" exception is injected. + MaxDepth int // QueryType is the DNS record type to query (e.g. dns.TypeA). - QueryType uint16 + QueryType uint16 // RootConfig controls how root servers are discovered. - RootConfig *dns.RootDiscoveryConfig + RootConfig *dns.RootDiscoveryConfig // QueryConfig controls per-query transport parameters. QueryConfig *dns.QueryConfig // RootAddrs is an optional pre-seeded list of root server IP addresses. - // When non-empty, root discovery via RootConfig is skipped. - RootAddrs []net.IP + // When non-empty, root discovery via RootConfig is skipped and each + // address becomes one root (named by its address). + RootAddrs []net.IP // Hooks provides optional callbacks for traversal events. - Hooks *TraverserHooks - // Fast controls cache sharing across branches. When true (default), child - // branches inherit glue discovered by earlier branches via the shared root - // cache, trading accuracy for speed. When false, each branch gets a - // completely independent cache — slower but results are not contaminated by - // sibling branch observations. + Hooks *TraverserHooks + // Fast enables the completed-referral memo (traverser.rb @answered): + // a referral identical to an earlier completed one (same qname/qclass/ + // qtype/server and per-IP weights) is replaced by it instead of being + // walked again. Non-fast mode re-walks every branch. Fast bool } func DefaultTraverserConfig() *TraverserConfig { return &TraverserConfig{ - MaxDepth: DefaultMaxDepth, - QueryType: dns.TypeA, - RootConfig: nil, - QueryConfig: nil, - RootAddrs: nil, - Fast: true, + MaxDepth: DefaultMaxDepth, + QueryType: dns.TypeA, + Fast: true, } } -// TraversalResult pairs a Referral with the Response received when it was processed. -type TraversalResult struct { - Referral *Referral - Response *Response -} - -// Traverser performs an exhaustive iterative DNS traversal starting from the -// root servers. Create one via NewTraverser and call Traverse to start a run. +// Traverser drives the traversal: it owns the packet-cached query client, +// the fast-mode memo and the explicit stack loop (traverser.rb). type Traverser struct { config *TraverserConfig + client *dns.Client exchange dns.ExchangeFunc - visited map[string]bool - depth int - mu sync.Mutex + // answered is the fast-mode memo of completed referrals. + answered map[string]*Referral + // seen maps every server name encountered to its IP addresses. + seen map[string][]string + // roots memoises root discovery so Roots() and Run() share one lookup. + roots []StartServer } func NewTraverser(cfg *TraverserConfig) *Traverser { if cfg == nil { cfg = DefaultTraverserConfig() } + if cfg.MaxDepth <= 0 { + cfg.MaxDepth = DefaultMaxDepth + } + if cfg.QueryType == 0 { + cfg.QueryType = dns.TypeA + } return &Traverser{ config: cfg, - exchange: nil, - visited: make(map[string]bool), - depth: 0, + client: dns.NewClient(cfg.QueryConfig, nil), + answered: make(map[string]*Referral), + seen: make(map[string][]string), } } +// SetExchange injects a mock wire exchange into the single query path (both +// traversal queries and root discovery); tests use this so no packets leave +// the process. func (t *Traverser) SetExchange(fn dns.ExchangeFunc) { t.exchange = fn + t.client = dns.NewClient(t.config.QueryConfig, fn) } func (t *Traverser) SetHooks(hooks *TraverserHooks) { - if t.config == nil { - t.config = DefaultTraverserConfig() - } t.config.Hooks = hooks } -func (t *Traverser) Traverse(ctx context.Context, name string) ([]TraversalResult, error) { - name = miekgdns.Fqdn(name) - - roots, err := t.discoverRoots(ctx) - if err != nil { - return nil, fmt.Errorf("root discovery: %w", err) - } - - initial := NewReferral(name, t.config.QueryType, ".", 0, 1.0, nil) - initial.Addresses = roots - - stack := NewStack(t.config.MaxDepth) - stack.Push(initial) - - rootCache := NewInfoCache(nil) - var ( - mu sync.Mutex - results []TraversalResult - ) - - for { - select { - case <-ctx.Done(): - return results, fmt.Errorf("traversal cancelled: %w", ctx.Err()) - default: - } - - ref := stack.Pop() - if ref == nil { - break - } - - var cache *InfoCache - if t.config.Fast { - // Fast mode: inherit glue from the shared root cache so earlier - // branch discoveries are visible to later branches. - cache = rootCache - if ref.Parent != nil { - cache = rootCache.Child() - } - } else { - // Non-fast mode: every referral gets its own independent cache so - // no cross-branch glue is reused, ensuring each path is resolved - // from scratch. - cache = NewInfoCache(nil) - } - - if t.config.Hooks != nil { - t.config.Hooks.emit(EventStart, TraversalResult{Referral: ref}, false) - } - - resp := t.processReferral(ctx, ref, cache) - - result := TraversalResult{Referral: ref, Response: resp} - if t.config.Hooks != nil { - t.config.Hooks.emit(EventComplete, result, false) - } - - mu.Lock() - results = append(results, result) - mu.Unlock() - - if resp.IsTerminal() { - continue - } - - if resp.Type == RespReferral { - children := resp.ChildReferrals() - for _, child := range children { - if !stack.Push(child) { - mu.Lock() - results = append(results, TraversalResult{ - Referral: child, - Response: &Response{ - Referral: child, - Type: RespError, - }, - }) - mu.Unlock() - } - } - } - - if resp.Type == RespCNAMEFollow { - follow := resp.CNAMEFollowReferral() - if follow != nil { - // Detect CNAME loop: target name already appears in the ancestor chain. - if follow.Parent != nil && follow.Parent.IsNameInChain(follow.Name) { - mu.Lock() - results = append(results, TraversalResult{ - Referral: follow, - Response: &Response{ - Referral: follow, - Type: RespCNAMELoop, - ErrorMessage: fmt.Sprintf("CNAME loop detected: %s already in traversal chain", follow.Name), - }, - }) - mu.Unlock() - } else if !stack.Push(follow) { - mu.Lock() - results = append(results, TraversalResult{ - Referral: follow, - Response: &Response{ - Referral: follow, - Type: RespError, - }, - }) - mu.Unlock() - } - } - } - } - - return results, nil +// ServersEncountered returns every server name seen during the run mapped to +// its known IP addresses (traverser.rb servers_encountered). +func (t *Traverser) ServersEncountered() map[string][]string { + return t.seen } -func (t *Traverser) discoverRoots(ctx context.Context) ([]net.IP, error) { - if len(t.config.RootAddrs) > 0 { - return t.config.RootAddrs, nil +// Roots performs (and memoises) root discovery, returning the start servers +// the traversal will begin from. Callers may use it before Run to report the +// initial root; Run reuses the memoised result. +func (t *Traverser) Roots(ctx context.Context) ([]StartServer, error) { + if t.roots == nil { + roots, err := t.rootStartServers(ctx) + if err != nil { + return nil, fmt.Errorf("root discovery: %w", err) + } + t.roots = roots } + return t.roots, nil +} - servers, err := dns.DiscoverRoots(ctx, t.config.RootConfig) +// Run traverses the DNS for name and returns the synthetic rootroot node +// (never displayed) whose Stats aggregate every leaf outcome; the per-leaf +// probabilities sum to 1.0. +func (t *Traverser) Run(ctx context.Context, name string) (*Referral, error) { + roots, err := t.Roots(ctx) if err != nil { return nil, err } - var addrs []net.IP - for _, srv := range servers { - addrs = append(addrs, srv.AllIPs(false)...) + cache := NewInfoCache(nil) + cache.AddHints("", roots) + + root := &Referral{ + RefID: "", + Qname: canonicalName(toASCII(name)), + Qclass: miekgdns.ClassINET, + Qtype: t.config.QueryType, + NSAType: dns.TypeA, + Server: "", + Bailiwick: "", + InfoCache: cache, + Status: RefStatusNormal, + Responses: make(map[string]*ServerResponse), + Children: make(map[string][]*Referral), + ServerWeights: make(map[string]float64), + client: t.client, + maxdepth: t.config.MaxDepth, } - return addrs, nil + t.config.Hooks.emit(StageNew, root, "") + + if err := t.run(ctx, root); err != nil { + return root, err + } + return root, nil } -func (t *Traverser) processReferral(ctx context.Context, ref *Referral, cache *InfoCache) *Response { - if !ref.HasAddresses() { - t.mu.Lock() - visitedCopy := make(map[string]bool) - for k, v := range t.visited { - visitedCopy[k] = v - } - t.mu.Unlock() +// stack markers mirroring Ruby's :calc_resolve / :calc_answer placeholders: +// the referral is revisited after its resolves/children finished, giving +// post-order statistics calculation without recursion. +type stackMarker int - // Resolve the nameserver's IP address. The NS hostname is stored in - // Bailiwick; ref.Name is the domain being queried (not the NS name). - nsToResolve := ref.Bailiwick - if nsToResolve == "" || nsToResolve == "." { - nsToResolve = ref.Name - } - nsName := strings.TrimSuffix(nsToResolve, ".") +const ( + markerNone stackMarker = iota + markerCalcResolve + markerCalcAnswer +) - ref.Addresses = t.resolveGlueViaSystem(ctx, nsToResolve, cache) - if len(ref.Addresses) > 0 { - ref.State = StateResolved - } else { - addrs, err := t.ResolveNS(ctx, nsToResolve, cache, visitedCopy, t.depth) - if err != nil { - return &Response{ - Referral: ref, - Type: RespNSResolutionFailed, - ErrorMessage: fmt.Sprintf("nameserver %s could not be resolved", nsName), - } - } - if len(addrs) > 0 { - ref.Addresses = addrs - ref.State = StateResolved - } else { - return &Response{ - Referral: ref, - Type: RespNSResolutionFailed, - ErrorMessage: fmt.Sprintf("nameserver %s could not be resolved", nsName), - } - } - } - } - - for _, addr := range ref.Addresses { - resp := t.queryServer(ctx, ref, addr, cache) - if resp != nil && resp.Type != RespSERVFAIL { - return resp - } - } - - return &Response{ - Referral: ref, - Type: RespSERVFAIL, - } +type stackEntry struct { + ref *Referral + marker stackMarker } -func (t *Traverser) ResolveNS(ctx context.Context, nsName string, cache *InfoCache, visited map[string]bool, depth int) ([]net.IP, error) { - if cache != nil { - if addrs := cache.LookupGlue(nsName); len(addrs) > 0 { - return addrs, nil - } +func (t *Traverser) run(ctx context.Context, root *Referral) error { + stack := []stackEntry{{ref: root}} + pop := func() stackEntry { + e := stack[len(stack)-1] + stack = stack[:len(stack)-1] + return e } - if visited != nil { - if visited[nsName] { - return nil, &CircularReferralError{ - Name: nsName, - Chain: getVisitedNames(visited), - } - } - visited[nsName] = true - } - - if depth > DefaultMaxDepth { - return nil, &UnresolvableNameserverError{ - Name: nsName, - Reason: "max depth exceeded", - } - } - - roots, err := t.discoverRoots(ctx) - if err != nil { - return nil, fmt.Errorf("root discovery: %w", err) - } - - ref := NewReferral(nsName, dns.TypeA, ".", 0, 1.0, nil) - ref.Addresses = roots - ref.State = StateResolved - - traversalCache := NewInfoCache(nil) - if visited != nil { - for name := range visited { - traversalCache.StoreGlue(name, []net.IP{}) - } - } - - var addrs []net.IP - var lastErr error - - stack := NewStack(DefaultMaxDepth) - stack.Push(ref) - - for { + for len(stack) > 0 { select { case <-ctx.Done(): - return nil, fmt.Errorf("resolution cancelled: %w", ctx.Err()) + return fmt.Errorf("traversal cancelled: %w", ctx.Err()) default: } - current := stack.Pop() - if current == nil { - break - } + e := pop() + r := e.ref - cacheForStep := traversalCache - if current.Parent != nil { - cacheForStep = traversalCache.Child() - } - - if t.config.Hooks != nil { - t.config.Hooks.emit(EventStart, TraversalResult{Referral: current}, true) - } - - resp := t.processReferral(ctx, current, cacheForStep) - - if t.config.Hooks != nil { - t.config.Hooks.emit(EventComplete, TraversalResult{Referral: current, Response: resp}, true) - } - - if resp.Type == RespAnswer && len(resp.Decoded.Answers) > 0 { - for _, rr := range resp.Decoded.Answers { - if a, ok := rr.(*miekgdns.A); ok { - addrs = append(addrs, a.A) - } - if aaaa, ok := rr.(*miekgdns.AAAA); ok { - addrs = append(addrs, aaaa.AAAA) - } + switch e.marker { + case markerCalcResolve: + r.resolveCalculate() + t.config.Hooks.emit(StageResolve, r, "") + stack = append(stack, stackEntry{ref: r}) // now needs processing + continue + case markerCalcAnswer: + r.answerCalculate() + t.config.Hooks.emit(StageAnswer, r, "") + if t.config.Fast && r.Status == RefStatusNormal && !hasLameResponse(r) { + t.answered[fastKey(r)] = r } - if len(addrs) > 0 { - if cache != nil { - cache.StoreGlue(nsName, addrs) - } - return addrs, nil + if !r.IsRootRoot() { + t.recordSeen(r) } - } - - if resp.Type == RespNXDOMAIN { - lastErr = &UnresolvableNameserverError{ - Name: nsName, - Reason: "NXDOMAIN", - } - break - } - - if resp.Type == RespSERVFAIL || resp.Type == RespError || resp.Type == RespNSResolutionFailed { - lastErr = fmt.Errorf("server error resolving %s: %s", nsName, resp.Type) continue } - if resp.Type == RespReferral { - children := resp.ChildReferrals() - for _, child := range children { - // Only skip visited names when they have no addresses; if glue - // was included in the referral response we still need to query - // that child to get the authoritative answer. - if visited != nil && visited[child.Name] && !child.HasAddresses() { - continue + // A new item. Fast mode: an identical completed referral replaces + // this one wholesale. noglue/loop nodes are excluded because their + // stats carry node-specific attributes and are cheap to recreate. + if t.config.Fast && r.Parent != nil { + if memo, ok := t.answered[fastKey(r)]; ok && !r.isNoGlue() && !r.isLoop() { + r.Parent.replaceChild(r, memo) + t.config.Hooks.emit(StageAnswerFast, r, memo.RefID) + continue + } + } + + t.config.Hooks.emit(StageStart, r, "") + + if !r.Resolved() { + // Push the resolve subtree with a calc_resolve placeholder so the + // weights are folded in once every resolve leaf completed. + stack = append(stack, stackEntry{ref: r, marker: markerCalcResolve}) + resolves, err := r.resolve() + if err != nil { + return err + } + for _, c := range resolves { + t.config.Hooks.emit(StageNew, c, "") + } + for i := len(resolves) - 1; i >= 0; i-- { + stack = append(stack, stackEntry{ref: resolves[i]}) + } + continue + } + + stack = append(stack, stackEntry{ref: r, marker: markerCalcAnswer}) + childrenSets, err := r.process(ctx) + if err != nil { + return err + } + + seenParentIP := make(map[string]bool) + var flat []*Referral + for _, set := range childrenSets { + for _, c := range set { + if len(childrenSets) > 1 && !seenParentIP[c.ParentIP] { + t.config.Hooks.emit(StageNewReferralSet, c, "") + seenParentIP[c.ParentIP] = true } - if !stack.Push(child) { - lastErr = &UnresolvableNameserverError{ - Name: nsName, - Reason: "max depth exceeded during resolution", + stage, earlier := StageNew, "" + if t.config.Fast { + if memo, ok := t.answered[fastKey(c)]; ok { + stage, earlier = StageNewFast, memo.RefID } } + t.config.Hooks.emit(stage, c, earlier) + flat = append(flat, c) } } - } - - if len(addrs) > 0 { - return addrs, nil - } - - if lastErr != nil { - return nil, lastErr - } - - return nil, &UnresolvableNameserverError{ - Name: nsName, - Reason: "resolution exhausted without answer", - } -} - -func (t *Traverser) queryServer(ctx context.Context, ref *Referral, server net.IP, cache *InfoCache) *Response { - var msg *miekgdns.Msg - var err error - - if t.exchange != nil { - msg, err = t.iterativeQueryWithExchange(ctx, server, ref.Name, ref.Qtype) - } else { - msg, err = dns.Query(ctx, server, ref.Name, ref.Qtype, t.config.QueryConfig) - if err == nil { - msg = t.ensureRDFalse(msg, server, ref.Name, ref.Qtype, t.config.QueryConfig) + for i := len(flat) - 1; i >= 0; i-- { + stack = append(stack, stackEntry{ref: flat[i]}) } } - - if err != nil { - return &Response{ - Referral: ref, - Server: server, - Type: RespError, - } - } - - resp := NewResponse(ref, server, cache) - resp.Process(msg) - return resp -} - -func (t *Traverser) iterativeQueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16) (*miekgdns.Msg, error) { - if t.config.QueryConfig == nil { - return dns.IterativeQueryWithExchange(ctx, server, name, qtype, nil, t.exchange) - } - return dns.IterativeQueryWithExchange(ctx, server, name, qtype, t.config.QueryConfig, t.exchange) -} - -func (t *Traverser) ensureRDFalse(msg *miekgdns.Msg, server net.IP, name string, qtype uint16, cfg *dns.QueryConfig) *miekgdns.Msg { - if msg != nil && msg.RecursionDesired { - if t.exchange != nil { - ctx := context.Background() - var err error - msg, err = t.iterativeQueryWithExchange(ctx, server, name, qtype) - if err != nil { - return nil - } - return msg - } - msg.RecursionDesired = false - } - return msg -} - -func (t *Traverser) resolveGlueViaSystem(ctx context.Context, name string, cache *InfoCache) []net.IP { - if cache != nil { - if addrs := cache.LookupGlue(name); len(addrs) > 0 { - return addrs - } - } - - c := &miekgdns.Client{ - Net: "udp", - ReadTimeout: 5 * time.Second, - WriteTimeout: 5 * time.Second, - } - if deadline, ok := ctx.Deadline(); ok { - remaining := time.Until(deadline) - if remaining <= 0 { - return nil - } - c.ReadTimeout = remaining - c.WriteTimeout = remaining - } - - fqdn := miekgdns.Fqdn(name) - - aMsg, _, err := c.ExchangeContext(ctx, newAQuery(fqdn), "127.0.0.1:53") - if err == nil { - var addrs []net.IP - for _, rr := range aMsg.Answer { - if a, ok := rr.(*miekgdns.A); ok { - addrs = append(addrs, a.A) - } - } - if len(addrs) > 0 { - if cache != nil { - cache.StoreGlue(name, addrs) - } - return addrs - } - } - return nil } -func newAQuery(name string) *miekgdns.Msg { - m := new(miekgdns.Msg) - m.SetQuestion(name, miekgdns.TypeA) - m.RecursionDesired = true - return m +// fastKey is the fast-mode memo key (traverser.rb): qname/qclass/qtype/ +// server plus the per-IP weights, lowercased. +func fastKey(r *Referral) string { + return strings.ToLower(fmt.Sprintf("%s:%s:%s:%s:%s", + r.Qname, ClassToString(r.Qclass), TypeToString(r.Qtype), r.Server, r.TxtIPsVerbose())) +} + +func hasLameResponse(r *Referral) bool { + for _, resp := range r.Responses { + if resp.Status == StatusReferralLame { + return true + } + } + return false +} + +func (t *Traverser) recordSeen(r *Referral) { + name := strings.ToLower(r.Server) + existing := t.seen[name] + for _, ip := range r.IPsAsArray() { + found := false + for _, have := range existing { + if have == ip { + found = true + break + } + } + if !found { + existing = append(existing, ip) + } + } + t.seen[name] = existing +} + +// rootStartServers returns the root servers as start-server hints: either +// the pre-seeded RootAddrs or the servers found via root discovery (one by +// default, all of them with AllRoots). IPv4 only, like the reference. +func (t *Traverser) rootStartServers(ctx context.Context) ([]StartServer, error) { + if len(t.config.RootAddrs) > 0 { + var out []StartServer + for _, ip := range t.config.RootAddrs { + if v4 := ip.To4(); v4 != nil { + out = append(out, StartServer{Name: v4.String(), IPs: []string{v4.String()}}) + } + } + if len(out) == 0 { + return nil, fmt.Errorf("no usable IPv4 root addresses") + } + return out, nil + } + + rootCfg := t.config.RootConfig + if t.exchange != nil { + var cp dns.RootDiscoveryConfig + if rootCfg != nil { + cp = *rootCfg + } + cp.Exchange = t.exchange + rootCfg = &cp + } + + servers, err := dns.DiscoverRoots(ctx, rootCfg) + if err != nil { + return nil, err + } + + var out []StartServer + for _, srv := range servers { + var ips []string + for _, ip := range srv.IPv4 { + ips = append(ips, ip.String()) + } + if len(ips) == 0 { + continue + } + out = append(out, StartServer{Name: canonicalName(srv.Name), IPs: ips}) + } + if len(out) == 0 { + return nil, fmt.Errorf("no root servers with IPv4 addresses") + } + sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name }) + return out, nil } diff --git a/internal/traverse/traverser_test.go b/internal/traverse/traverser_test.go index e422624..c264e17 100644 --- a/internal/traverse/traverser_test.go +++ b/internal/traverse/traverser_test.go @@ -2,482 +2,774 @@ package traverse import ( "context" + "math" "net" + "strconv" + "strings" + "sync" "testing" + "time" + idns "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" "github.com/miekg/dns" ) -const ( - dnsTypeA = dns.TypeA - dnsTypeNS = dns.TypeNS - dnsTypeCNAME = dns.TypeCNAME - dnsTypeSOA = dns.TypeSOA -) +// --- mock exchange: the single query path used by production and tests --- -func TestDefaultTraverserConfig(t *testing.T) { - cfg := DefaultTraverserConfig() - if cfg.MaxDepth != DefaultMaxDepth { - t.Errorf("MaxDepth = %d, want %d", cfg.MaxDepth, DefaultMaxDepth) - } - if cfg.QueryType != dnsTypeA { - t.Errorf("QueryType = %d, want %d", cfg.QueryType, dnsTypeA) +type mockKey struct { + server string + qname string + qtype uint16 +} + +type mockExchange struct { + mu sync.Mutex + responses map[mockKey]*dns.Msg + errors map[mockKey]error + calls map[mockKey]int +} + +func newMockExchange() *mockExchange { + return &mockExchange{ + responses: make(map[mockKey]*dns.Msg), + errors: make(map[mockKey]error), + calls: make(map[mockKey]int), } } -func TestNewTraverserNilConfig(t *testing.T) { - tr := NewTraverser(nil) - if tr == nil { - t.Fatal("NewTraverser(nil) should not return nil") - } +func (m *mockExchange) on(server, qname string, qtype uint16, msg *dns.Msg) { + m.responses[mockKey{server, dns.Fqdn(qname), qtype}] = msg } -func TestTraverserSimpleTraversal(t *testing.T) { - answerResp := func() *dns.Msg { - m := new(dns.Msg) - m.SetReply(new(dns.Msg)) - m.Answer = append(m.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - return m - }() - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") +func (m *mockExchange) fn(_ context.Context, server string, msg *dns.Msg, _ bool) (*dns.Msg, error) { + host := server + if h, _, err := net.SplitHostPort(server); err == nil { + host = h + } + q := msg.Question[0] + key := mockKey{host, q.Name, q.Qtype} + m.mu.Lock() + m.calls[key]++ + resp, ok := m.responses[key] + err := m.errors[key] + m.mu.Unlock() if err != nil { - t.Fatalf("unexpected error: %v", err) + return nil, err } - if len(results) == 0 { - t.Fatal("expected at least 1 result") + if !ok { + return nil, &net.OpError{Op: "read", Err: &net.DNSError{Err: "no mock response", Name: q.Name}} } + out := resp.Copy() + out.SetReply(msg) + out.Answer = resp.Answer + out.Ns = resp.Ns + out.Extra = resp.Extra + out.Rcode = resp.Rcode + return out, nil +} - found := false - for _, r := range results { - if r.Response.Type == RespAnswer { - found = true - break - } - } - if !found { - t.Error("expected to find an answer response") +func answerMsg(rrs ...dns.RR) *dns.Msg { + m := new(dns.Msg) + m.Answer = rrs + return m +} + +func referralMsg(nsRRs []dns.RR, glue ...dns.RR) *dns.Msg { + m := new(dns.Msg) + m.Ns = nsRRs + m.Extra = glue + return m +} + +func rcodeMsg(rcode int) *dns.Msg { + m := new(dns.Msg) + m.Rcode = rcode + return m +} + +func testConfig(fast bool) *TraverserConfig { + return &TraverserConfig{ + MaxDepth: DefaultMaxDepth, + QueryType: dns.TypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + Fast: fast, + QueryConfig: &idns.QueryConfig{ + Retries: 1, + Timeout: time.Second, + RetryDelay: time.Millisecond, + }, } } -func TestTraverserReferralTraversal(t *testing.T) { - rootAnswer := new(dns.Msg) - rootAnswer.Rcode = dns.RcodeSuccess - rootAnswer.Authoritative = false - rootAnswer.Ns = append(rootAnswer.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: "com.", Rrtype: dnsTypeNS, Class: dns.ClassINET}, - Ns: "a.gtld-servers.net.", - }) - rootAnswer.Extra = append(rootAnswer.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dnsTypeA}, - A: net.ParseIP("192.5.6.30"), - }) - - tldAnswer := new(dns.Msg) - tldAnswer.SetReply(new(dns.Msg)) - tldAnswer.Answer = append(tldAnswer.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - key := q.Name + "/" + dns.TypeToString[q.Qtype] - if q.Name == "example.com." && server == "198.41.0.4" { - return rootAnswer.Copy(), nil - } - if q.Name == "example.com." { - return tldAnswer.Copy(), nil - } - _ = key - return nil, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") +func runTraversal(t *testing.T, cfg *TraverserConfig, m *mockExchange, qname string) (*Traverser, *Referral) { + t.Helper() + tr := NewTraverser(cfg) + tr.SetExchange(m.fn) + root, err := tr.Run(context.Background(), qname) if err != nil { - t.Fatalf("unexpected error: %v", err) + t.Fatalf("Run(%q): %v", qname, err) } - if len(results) < 2 { - t.Fatalf("expected at least 2 results (referral + answer), got %d", len(results)) + return tr, root +} + +func statsSum(root *Referral) float64 { + sum := 0.0 + for _, e := range root.Stats { + sum += e.Prob + } + return sum +} + +func assertSumsToOne(t *testing.T, root *Referral) { + t.Helper() + if sum := statsSum(root); math.Abs(sum-1.0) > 1e-9 { + t.Errorf("aggregated leaf probabilities sum to %v, want 1.0", sum) } } -func TestTraverserMaxDepth(t *testing.T) { - callCount := 0 - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 2, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - callCount++ - m := new(dns.Msg) - m.Rcode = dns.RcodeSuccess - m.Authoritative = false - m.Ns = append(m.Ns, &dns.NS{ - Hdr: dns.RR_Header{Rrtype: dnsTypeNS}, - Ns: "ns.example.com.", - }) - m.Extra = append(m.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dnsTypeA}, - A: net.ParseIP("1.2.3.4"), - }) - return m, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "deep.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - if callCount < 1 { - t.Errorf("expected at least 1 call before max depth, got %d", callCount) - } - - depthExceeded := false - for _, r := range results { - if r.Referral != nil && r.Referral.Depth >= 2 { - depthExceeded = true - } - if r.Response.Type == RespError { - depthExceeded = true +func leavesByStatus(root *Referral, status Status) []*StatsEntry { + var out []*StatsEntry + for _, e := range root.StatsList() { + if e.Response.Status == status { + out = append(out, e) } } - if !depthExceeded { - t.Error("expected to see depth exceeded results") + return out +} + +// --- scenarios --- + +// mockSimpleDelegation wires root → com → example.com with two glued NS that +// both answer. +func mockSimpleDelegation() *mockExchange { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns2.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + aRR("ns2.example.com", "2.2.2.2"), + )) + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + m.on("2.2.2.2", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + return m +} + +func TestReferralFanOut(t *testing.T) { + m := mockSimpleDelegation() + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + answered := leavesByStatus(root, StatusAnswered) + if len(answered) != 2 { + t.Fatalf("expected 2 answered leaves (one per NS), got %d: %v", len(answered), root.StatsList()) + } + for _, leaf := range answered { + if math.Abs(leaf.Prob-0.5) > 1e-9 { + t.Errorf("leaf %s prob = %v, want 0.5", leaf.Key, leaf.Prob) + } + } + + // RefID grammar: rootroot "", root child "1", gtld "1.1", NS "1.1.1"/"1.1.2". + if root.RefID != "" { + t.Errorf("rootroot refid = %q, want empty", root.RefID) + } + top := root.Children["rootroot"] + if len(top) != 1 || top[0].RefID != "1" { + t.Fatalf("top children = %v", top) + } + gtld := top[0].Children["198.41.0.4"] + if len(gtld) != 1 || gtld[0].RefID != "1.1" { + t.Fatalf("gtld children refids wrong: %v", gtld) + } + nsKids := gtld[0].Children["192.5.6.30"] + if len(nsKids) != 2 || nsKids[0].RefID != "1.1.1" || nsKids[1].RefID != "1.1.2" { + t.Fatalf("NS children refids wrong: %v", nsKids) + } + if nsKids[0].Depth() != 3 { + t.Errorf("depth of 1.1.1 = %d, want 3", nsKids[0].Depth()) } } -func TestTraverserContextCancellation(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - m := new(dns.Msg) - m.Rcode = dns.RcodeSuccess - m.Authoritative = false - m.Ns = append(m.Ns, &dns.NS{ - Hdr: dns.RR_Header{Rrtype: dnsTypeNS}, - Ns: "ns.example.com.", - }) - m.Extra = append(m.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dnsTypeA}, - A: net.ParseIP("1.2.3.4"), - }) - return m, nil - }) +func TestGluelessResolveSubtree(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns.other.net")}, + aRR("ns1.example.com", "1.1.1.1"), + )) + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + // resolve subtree for A ns.other.net starts back at the root hints + m.on("198.41.0.4", "ns.other.net", dns.TypeA, referralMsg( + []dns.RR{nsRR("net", "d.gtld.net")}, + aRR("d.gtld.net", "3.3.3.3"), + )) + m.on("3.3.3.3", "ns.other.net", dns.TypeA, answerMsg(aRR("ns.other.net", "4.4.4.4"))) + m.on("4.4.4.4", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + var events []TraversalEvent + cfg := testConfig(false) + cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) { events = append(events, ev) }} + _, root := runTraversal(t, cfg, m, "www.example.com") + + assertSumsToOne(t, root) + answered := leavesByStatus(root, StatusAnswered) + if len(answered) != 2 { + t.Fatalf("expected 2 answered leaves, got %v", root.StatsList()) + } + var viaGlueless *StatsEntry + for _, leaf := range answered { + if leaf.Response.IP == "4.4.4.4" { + viaGlueless = leaf + } + } + if viaGlueless == nil { + t.Fatal("no answered leaf via the glueless nameserver") + } + if math.Abs(viaGlueless.Prob-0.5) > 1e-9 { + t.Errorf("glueless leaf prob = %v, want 0.5", viaGlueless.Prob) + } + if got := viaGlueless.Referral.ServerWeights["4.4.4.4"]; math.Abs(got-1.0) > 1e-9 { + t.Errorf("resolved serverweight = %v, want 1.0", got) + } + + // The resolve subtree inserts a .0 refid component and is flagged. + sawResolve := false + for _, ev := range events { + if ev.RefID == "1.1.2.0.1" { + sawResolve = true + if !ev.IsResolve { + t.Error("resolve subtree event not flagged IsResolve") + } + } + } + if !sawResolve { + t.Errorf("no event for resolve refid 1.1.2.0.1; events: %v", refids(events)) + } + // Depth ignores the zero components. + if d := refidDepth("1.1.2.0.1.1"); d != 5 { + t.Errorf("refidDepth(1.1.2.0.1.1) = %d, want 5", d) + } +} + +func TestNoGlueDeadEnd(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + // In-bailiwick NS without glue: dead end. + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com")}, + )) + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + noglue := leavesByStatus(root, StatusNoGlue) + if len(noglue) != 1 { + t.Fatalf("expected 1 noglue leaf, got %v", root.StatsList()) + } + leaf := noglue[0] + if math.Abs(leaf.Prob-1.0) > 1e-9 { + t.Errorf("noglue prob = %v, want 1.0", leaf.Prob) + } + if leaf.Response.IP != "192.5.6.30" { + t.Errorf("noglue response IP = %q, want the referring parent IP", leaf.Response.IP) + } + if leaf.Referral.Server != "ns1.example.com" { + t.Errorf("noglue referral server = %q", leaf.Referral.Server) + } + if leaf.Referral.Parent.Server != "a.gtld-servers.net" { + t.Errorf("noglue parent server = %q", leaf.Referral.Parent.Server) + } + if !strings.HasPrefix(leaf.Key, "key:noglue:192.5.6.30:www.example.com:IN:A:ns1.example.com:") { + t.Errorf("noglue stats key = %q", leaf.Key) + } +} + +func TestResolveLoopDeadEnd(t *testing.T) { + m := newMockExchange() + // x.net NS ns.y.net (no glue); y.net NS ns.x.net (no glue): resolving + // either server needs the other, which is a loop. + xReferral := referralMsg([]dns.RR{nsRR("x.net", "ns.y.net")}) + yReferral := referralMsg([]dns.RR{nsRR("y.net", "ns.x.net")}) + m.on("198.41.0.4", "www.x.net", dns.TypeA, xReferral) + m.on("198.41.0.4", "ns.y.net", dns.TypeA, yReferral) + m.on("198.41.0.4", "ns.x.net", dns.TypeA, xReferral) + + _, root := runTraversal(t, testConfig(false), m, "www.x.net") + + assertSumsToOne(t, root) + loops := leavesByStatus(root, StatusLoop) + if len(loops) != 1 { + t.Fatalf("expected 1 loop leaf, got %v", root.StatsList()) + } + if math.Abs(loops[0].Prob-1.0) > 1e-9 { + t.Errorf("loop prob = %v, want 1.0", loops[0].Prob) + } + if loops[0].Referral.Status != RefStatusLoop { + t.Errorf("loop referral status = %q", loops[0].Referral.Status) + } +} + +func TestCNAMERestart(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.a.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.a.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("a.com", "ns.a.com")}, + aRR("ns.a.com", "5.5.5.5"), + )) + m.on("5.5.5.5", "www.a.com", dns.TypeA, answerMsg(cnameRRT("www.a.com", "www.b.net"))) + // restart resumes from the deepest cached zone — nothing cached for + // b.net, so back to the root. + m.on("198.41.0.4", "www.b.net", dns.TypeA, referralMsg( + []dns.RR{nsRR("net", "e.gtld.net")}, + aRR("e.gtld.net", "6.6.6.6"), + )) + m.on("6.6.6.6", "www.b.net", dns.TypeA, answerMsg(aRR("www.b.net", "7.7.7.7"))) + + _, root := runTraversal(t, testConfig(false), m, "www.a.com") + + assertSumsToOne(t, root) + answered := leavesByStatus(root, StatusAnswered) + if len(answered) != 1 { + t.Fatalf("expected 1 answered leaf, got %v", root.StatsList()) + } + leaf := answered[0] + if leaf.Response.Qname != "www.b.net" { + t.Errorf("answered qname = %q, want restart target www.b.net", leaf.Response.Qname) + } + if leaf.Referral.Qname != "www.b.net" { + t.Errorf("restart referral qname = %q", leaf.Referral.Qname) + } + // The restart chain keeps numbering below the restarting node. + if leaf.Referral.RefID != "1.1.1.1.1" { + t.Errorf("answered refid = %q, want 1.1.1.1.1", leaf.Referral.RefID) + } +} + +func TestCNAMERestartLoop(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.a.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.a.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("a.com", "ns.a.com")}, + aRR("ns.a.com", "5.5.5.5"), + )) + m.on("5.5.5.5", "www.a.com", dns.TypeA, answerMsg(cnameRRT("www.a.com", "www.b.net"))) + // www.b.net points straight back at www.a.com: every chain target is + // checked against the ancestor queries, so this is a CNAME loop. + m.on("198.41.0.4", "www.b.net", dns.TypeA, answerMsg(cnameRRT("www.b.net", "www.a.com"))) + + _, root := runTraversal(t, testConfig(false), m, "www.a.com") + + assertSumsToOne(t, root) + loops := leavesByStatus(root, StatusCNAMELoop) + if len(loops) != 1 { + t.Fatalf("expected 1 cname_loop leaf, got %v", root.StatsList()) + } + if math.Abs(loops[0].Prob-1.0) > 1e-9 { + t.Errorf("cname_loop prob = %v, want 1.0", loops[0].Prob) + } +} + +func TestCNAMERestartLoopIntermediateChainTarget(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.a.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("a.com", "ns.a.com")}, + aRR("ns.a.com", "5.5.5.5"), + )) + m.on("5.5.5.5", "www.a.com", dns.TypeA, answerMsg(cnameRRT("www.a.com", "www.b.net"))) + m.on("198.41.0.4", "www.b.net", dns.TypeA, referralMsg( + []dns.RR{nsRR("b.net", "ns.b.net")}, + aRR("ns.b.net", "6.6.6.6"), + )) + // A three-record chain whose INTERMEDIATE target www.a.com matches an + // ancestor query; the endname other.org does not, so an endname-only + // loop check would miss it. + m.on("6.6.6.6", "www.b.net", dns.TypeA, answerMsg( + cnameRRT("www.b.net", "c.b.net"), + cnameRRT("c.b.net", "www.a.com"), + cnameRRT("www.a.com", "other.org"), + )) + + _, root := runTraversal(t, testConfig(false), m, "www.a.com") + + assertSumsToOne(t, root) + loops := leavesByStatus(root, StatusCNAMELoop) + if len(loops) != 1 { + t.Fatalf("expected 1 cname_loop leaf, got %v", root.StatsList()) + } + leaf := loops[0] + if math.Abs(leaf.Prob-1.0) > 1e-9 { + t.Errorf("cname_loop prob = %v, want 1.0", leaf.Prob) + } + dq := leaf.Response.DQ + if len(dq.ChainTargets) != 3 { + t.Fatalf("ChainTargets = %v, want 3 entries", dq.ChainTargets) + } + if dq.ChainTargets[1] != "www.a.com" { + t.Errorf("ChainTargets[1] = %q, want www.a.com", dq.ChainTargets[1]) + } + if dq.Endname != "other.org" { + t.Errorf("Endname = %q, want other.org", dq.Endname) + } +} + +func TestDepthLimitInjectsException(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.c.b.a", dns.TypeA, referralMsg( + []dns.RR{nsRR("a", "ns.a")}, + aRR("ns.a", "1.1.1.1"), + )) + m.on("1.1.1.1", "www.c.b.a", dns.TypeA, referralMsg( + []dns.RR{nsRR("b.a", "ns.b.a")}, + aRR("ns.b.a", "2.2.2.2"), + )) + // Node 1.1.1 sits at depth 3 == maxdepth: its query is never sent. + cfg := testConfig(false) + cfg.MaxDepth = 3 + _, root := runTraversal(t, cfg, m, "www.c.b.a") + + assertSumsToOne(t, root) + exceptions := leavesByStatus(root, StatusException) + if len(exceptions) != 1 { + t.Fatalf("expected 1 exception leaf, got %v", root.StatsList()) + } + leaf := exceptions[0] + if leaf.Response.DQ.ExceptionMessage != "Maxdepth 3 exceeded" { + t.Errorf("exception message = %q, want %q", leaf.Response.DQ.ExceptionMessage, "Maxdepth 3 exceeded") + } + if math.Abs(leaf.Prob-1.0) > 1e-9 { + t.Errorf("exception prob = %v, want 1.0", leaf.Prob) + } + // No query should have reached depth 3. + m.mu.Lock() + defer m.mu.Unlock() + for key := range m.calls { + if key.server == "2.2.2.2" { + t.Error("query sent beyond the depth limit") + } + } +} + +func TestLameReferral(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + )) + // ns1 refers back to the same zone: not strictly deeper, so lame. + m.on("1.1.1.1", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns2.example.com")}, + aRR("ns2.example.com", "2.2.2.2"), + )) + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + lame := leavesByStatus(root, StatusReferralLame) + if len(lame) != 1 { + t.Fatalf("expected 1 referral_lame leaf, got %v", root.StatsList()) + } + leaf := lame[0] + if math.Abs(leaf.Prob-1.0) > 1e-9 { + t.Errorf("lame prob = %v, want 1.0", leaf.Prob) + } + if !strings.HasSuffix(leaf.Key, ":192.5.6.30") { + t.Errorf("lame stats key should end with the parent IP, got %q", leaf.Key) + } + if leaf.Referral.ParentIP != "192.5.6.30" { + t.Errorf("lame referral parent ip = %q", leaf.Referral.ParentIP) + } +} + +func TestChildsetDigitWhenMultipleIPsProduceChildren(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + aRR("a.gtld-servers.net", "192.5.6.31"), + )) + exampleReferral := referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + ) + m.on("192.5.6.30", "www.example.com", dns.TypeA, exampleReferral) + m.on("192.5.6.31", "www.example.com", dns.TypeA, exampleReferral) + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + + var setEvents []TraversalEvent + cfg := testConfig(false) + cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) { + if ev.Stage == StageNewReferralSet { + setEvents = append(setEvents, ev) + } + }} + _, root := runTraversal(t, cfg, m, "www.example.com") + + assertSumsToOne(t, root) + gtld := root.Children["rootroot"][0].Children["198.41.0.4"][0] + set1 := gtld.Children["192.5.6.30"] + set2 := gtld.Children["192.5.6.31"] + if len(set1) != 1 || set1[0].RefID != "1.1.1.1" { + t.Errorf("first childset refid = %v, want 1.1.1.1", refidsOf(set1)) + } + if len(set2) != 1 || set2[0].RefID != "1.1.2.1" { + t.Errorf("second childset refid = %v, want 1.1.2.1", refidsOf(set2)) + } + if len(setEvents) != 2 { + t.Errorf("expected 2 new_referral_set events, got %d", len(setEvents)) + } + // Identical answers from both paths merge into one leaf with prob 1.0. + answered := leavesByStatus(root, StatusAnswered) + if len(answered) != 1 || math.Abs(answered[0].Prob-1.0) > 1e-9 { + t.Errorf("answered leaves = %v", root.StatsList()) + } +} + +func TestFastModeReuse(t *testing.T) { + cfg := testConfig(true) + cfg.RootAddrs = []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")} + + m := newMockExchange() + comReferral := referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + ) + m.on("198.41.0.4", "www.example.com", dns.TypeA, comReferral) + m.on("199.9.14.201", "www.example.com", dns.TypeA, comReferral) + m.on("192.5.6.30", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + + var events []TraversalEvent + cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) { events = append(events, ev) }} + _, root := runTraversal(t, cfg, m, "www.example.com") + + assertSumsToOne(t, root) + var fast *TraversalEvent + for i := range events { + if events[i].Stage == StageAnswerFast { + fast = &events[i] + } + } + if fast == nil { + t.Fatal("expected a StageAnswerFast event in fast mode") + } + if fast.RefID != "2.1" || fast.CompletedEarlier != "1.1" { + t.Errorf("fast event refid=%q completedEarlier=%q, want 2.1 / 1.1", fast.RefID, fast.CompletedEarlier) + } + if fast.Referral.ReplacedBy == nil || fast.Referral.ReplacedBy.RefID != "1.1" { + t.Error("fast-replaced referral should point at its replacement") + } + // The second branch's child is the first branch's completed node. + second := root.Children["rootroot"][1] + if got := second.Children["199.9.14.201"][0].RefID; got != "1.1" { + t.Errorf("replaced child refid = %q, want 1.1", got) + } + answered := leavesByStatus(root, StatusAnswered) + if len(answered) != 1 || math.Abs(answered[0].Prob-1.0) > 1e-9 { + t.Errorf("answered leaves = %v", root.StatsList()) + } +} + +func TestNonFastModeReWalks(t *testing.T) { + cfg := testConfig(false) + cfg.RootAddrs = []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")} + + m := newMockExchange() + comReferral := referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + ) + m.on("198.41.0.4", "www.example.com", dns.TypeA, comReferral) + m.on("199.9.14.201", "www.example.com", dns.TypeA, comReferral) + m.on("192.5.6.30", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + + var events []TraversalEvent + cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) { events = append(events, ev) }} + _, root := runTraversal(t, cfg, m, "www.example.com") + + assertSumsToOne(t, root) + for _, ev := range events { + if ev.Stage == StageAnswerFast || ev.Stage == StageNewFast { + t.Fatalf("unexpected fast-mode event %v in non-fast mode", ev.Stage) + } + } + // Both branches keep their own child node. + second := root.Children["rootroot"][1] + if got := second.Children["199.9.14.201"][0].RefID; got != "2.1" { + t.Errorf("non-fast child refid = %q, want 2.1", got) + } +} + +func TestAllRootsBranching(t *testing.T) { + cfg := testConfig(false) + cfg.RootAddrs = nil + for i := 1; i <= 13; i++ { + cfg.RootAddrs = append(cfg.RootAddrs, net.ParseIP("198.41.0."+strconv.Itoa(i))) + } + + m := newMockExchange() + for i := 1; i <= 13; i++ { + m.on("198.41.0."+strconv.Itoa(i), "example.com", dns.TypeA, answerMsg(aRR("example.com", "9.9.9.9"))) + } + _, root := runTraversal(t, cfg, m, "example.com") + + assertSumsToOne(t, root) + top := root.Children["rootroot"] + if len(top) != 13 { + t.Fatalf("expected 13 top-level children, got %d", len(top)) + } + for i, child := range top { + if child.RefID != strconv.Itoa(i+1) { + t.Errorf("child %d refid = %q, want %q", i, child.RefID, strconv.Itoa(i+1)) + } + } + answered := leavesByStatus(root, StatusAnswered) + // Same answer from 13 different server IPs: 13 distinct leaves of 1/13. + if len(answered) != 13 { + t.Fatalf("expected 13 answered leaves, got %d", len(answered)) + } + for _, leaf := range answered { + if math.Abs(leaf.Prob-1.0/13) > 1e-9 { + t.Errorf("leaf %s prob = %v, want %v", leaf.Key, leaf.Prob, 1.0/13) + } + } +} + +func TestErrorAndNoDataStatuses(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns2.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + aRR("ns2.example.com", "2.2.2.2"), + )) + m.on("1.1.1.1", "www.example.com", dns.TypeA, rcodeMsg(dns.RcodeNameError)) + soa := &dns.SOA{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeSOA, Class: dns.ClassINET}, + Ns: "ns1.example.com.", Mbox: "hostmaster.example.com.", + } + nodata := new(dns.Msg) + nodata.Ns = []dns.RR{soa} + m.on("2.2.2.2", "www.example.com", dns.TypeA, nodata) + + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + errs := leavesByStatus(root, StatusError) + if len(errs) != 1 || errs[0].Response.DQ.ErrorMessage != "No such domain (NXDOMAIN)" { + t.Errorf("error leaves = %v", root.StatsList()) + } + nodataLeaves := leavesByStatus(root, StatusNoData) + if len(nodataLeaves) != 1 { + t.Errorf("nodata leaves = %v", root.StatsList()) + } +} + +func TestNetworkExceptionLeaf(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + )) + // no mock for 1.1.1.1 → network error → exception status + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + exceptions := leavesByStatus(root, StatusException) + if len(exceptions) != 1 { + t.Fatalf("expected exception leaf, got %v", root.StatsList()) + } + if math.Abs(exceptions[0].Prob-1.0) > 1e-9 { + t.Errorf("exception prob = %v", exceptions[0].Prob) + } +} + +func TestPacketCacheSingleWireQuery(t *testing.T) { + m := mockSimpleDelegation() + // Both NS answer; querying the same tuple twice must hit the cache. + tr, _ := runTraversal(t, testConfig(false), m, "www.example.com") + + m.mu.Lock() + defer m.mu.Unlock() + for key, n := range m.calls { + if n != 1 { + t.Errorf("query %v sent %d times, want 1", key, n) + } + } + if tr.client == nil { + t.Fatal("traverser has no client") + } +} + +func TestServersEncountered(t *testing.T) { + m := mockSimpleDelegation() + tr, _ := runTraversal(t, testConfig(false), m, "www.example.com") + seen := tr.ServersEncountered() + if len(seen["ns1.example.com"]) != 1 || seen["ns1.example.com"][0] != "1.1.1.1" { + t.Errorf("seen ns1 = %v", seen["ns1.example.com"]) + } + if _, ok := seen["a.gtld-servers.net"]; !ok { + t.Errorf("gtld server missing from seen: %v", seen) + } + if _, ok := seen[""]; ok { + t.Error("rootroot must not be recorded in servers encountered") + } +} + +func TestRunContextCancellation(t *testing.T) { + m := mockSimpleDelegation() + tr := NewTraverser(testConfig(false)) + tr.SetExchange(m.fn) ctx, cancel := context.WithCancel(context.Background()) cancel() - - _, err := tr.Traverse(ctx, "example.com") - if err == nil { - t.Fatal("expected error on cancelled context") + if _, err := tr.Run(ctx, "www.example.com"); err == nil { + t.Fatal("expected cancellation error") } } -func TestTraverserNXDOMAIN(t *testing.T) { - nxdResp := new(dns.Msg) - nxdResp.Rcode = dns.RcodeNameError - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nxdResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "nonexistent.invalid") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least 1 result") - } - if results[0].Response.Type != RespNXDOMAIN { - t.Errorf("Type = %d, want %d", results[0].Response.Type, RespNXDOMAIN) +func TestIDNQnameConvertsToPunycode(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "xn--bcher-kva.example", dns.TypeA, answerMsg(aRR("xn--bcher-kva.example", "9.9.9.9"))) + _, root := runTraversal(t, testConfig(false), m, "bücher.example") + if root.Qname != "xn--bcher-kva.example" { + t.Errorf("qname = %q, want punycode", root.Qname) } + assertSumsToOne(t, root) } -func TestTraverserSERVFAIL(t *testing.T) { - sfResp := new(dns.Msg) - sfResp.Rcode = dns.RcodeServerFailure - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return sfResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least 1 result") - } - if results[0].Response.Type != RespSERVFAIL { - t.Errorf("Type = %d, want %d", results[0].Response.Type, RespSERVFAIL) +func refids(events []TraversalEvent) []string { + out := make([]string, len(events)) + for i, ev := range events { + out[i] = ev.Stage.String() + ":" + ev.RefID } + return out } -func TestTraverserCNAMEFollow(t *testing.T) { - cnameResp := new(dns.Msg) - cnameResp.SetReply(new(dns.Msg)) - cnameResp.Answer = append(cnameResp.Answer, - &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dnsTypeCNAME, Class: dns.ClassINET}, - Target: "example.com.", - }, - ) - - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - if q.Name == "www.example.com." { - return cnameResp.Copy(), nil - } - if q.Name == "example.com." { - return answerResp.Copy(), nil - } - return nil, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "www.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - foundCNAME := false - foundAnswer := false - for _, r := range results { - if r.Response.Type == RespCNAMEFollow { - foundCNAME = true - } - if r.Response.Type == RespAnswer { - foundAnswer = true - } - } - if !foundCNAME { - t.Error("expected CNAME follow response") - } - if !foundAnswer { - t.Error("expected final answer response") +func refidsOf(refs []*Referral) []string { + out := make([]string, len(refs)) + for i, r := range refs { + out[i] = r.RefID } + return out } -func TestTraverserProbabilityCalculation(t *testing.T) { - rootReferral := new(dns.Msg) - rootReferral.Rcode = dns.RcodeSuccess - rootReferral.Authoritative = false - rootReferral.Ns = append(rootReferral.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dnsTypeNS}, Ns: "a.root-servers.net."}, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dnsTypeNS}, Ns: "b.root-servers.net."}, - ) - rootReferral.Extra = append(rootReferral.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "a.root-servers.net.", Rrtype: dnsTypeA}, A: net.ParseIP("198.41.0.4")}, - &dns.A{Hdr: dns.RR_Header{Name: "b.root-servers.net.", Rrtype: dnsTypeA}, A: net.ParseIP("199.9.14.201")}, - ) - - tldAnswer := new(dns.Msg) - tldAnswer.SetReply(new(dns.Msg)) - tldAnswer.Answer = append(tldAnswer.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("1.2.3.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - if q.Name == "example.com." && server == "1.2.3.4" { - return rootReferral.Copy(), nil - } - return tldAnswer.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - for _, r := range results { - if r.Referral != nil && r.Referral.Depth == 1 && r.Referral.Parent != nil { - if r.Referral.Prob != 0.5 { - t.Errorf("child prob = %f, want 0.5", r.Referral.Prob) - } - } - } -} - -func TestTraverserNODATA(t *testing.T) { - nodataResp := new(dns.Msg) - nodataResp.Rcode = dns.RcodeSuccess - nodataResp.Authoritative = true - nodataResp.Ns = append(nodataResp.Ns, &dns.SOA{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeSOA, Class: dns.ClassINET, Ttl: 3600}, - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nodataResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least 1 result") - } - if results[0].Response.Type != RespNODATA { - t.Errorf("Type = %d, want %d", results[0].Response.Type, RespNODATA) - } -} - -func TestTraverserMultipleRoots(t *testing.T) { - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results") - } -} - -func TestTraverserNilExchangeResponse(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least 1 result even with nil response") - } -} - -func TestTraverserCacheChaining(t *testing.T) { - rootReferral := new(dns.Msg) - rootReferral.Rcode = dns.RcodeSuccess - rootReferral.Authoritative = false - rootReferral.Ns = append(rootReferral.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dnsTypeNS}, Ns: "a.gtld-servers.net."}, - ) - rootReferral.Extra = append(rootReferral.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dnsTypeA}, A: net.ParseIP("192.5.6.30")}, - ) - - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - if q.Name == "example.com." && server == "198.41.0.4" { - return rootReferral.Copy(), nil - } - return answerResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - cacheHits := 0 - for _, r := range results { - if r.Response != nil && r.Response.Cache != nil { - if r.Response.Cache.NSCount() > 0 { - cacheHits++ - } - } - } - if cacheHits == 0 { - t.Error("expected cache to store NS records from referrals") +func cnameRRT(name, target string) dns.RR { + return &dns.CNAME{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, + Target: dns.Fqdn(target), } } diff --git a/web/api/handler.go b/web/api/handler.go index 7927057..a3f28ab 100644 --- a/web/api/handler.go +++ b/web/api/handler.go @@ -7,6 +7,8 @@ import ( "fmt" "io/fs" "net/http" + "sort" + "strings" "sync" "time" @@ -40,23 +42,57 @@ type TraverseStartResponse struct { // ProgressEvent carries a single traversal hook event. type ProgressEvent struct { - Stage string `json:"stage"` - Depth int `json:"depth"` - Name string `json:"name"` - QType string `json:"qtype"` - Server string `json:"server,omitempty"` - Bailiwick string `json:"bailiwick,omitempty"` - IsResolve bool `json:"is_resolve,omitempty"` + Stage string `json:"stage"` + RefID string `json:"refid"` + Depth int `json:"depth"` + Name string `json:"name"` + QType string `json:"qtype"` + Server string `json:"server,omitempty"` + IPs string `json:"ips,omitempty"` + Bailiwick string `json:"bailiwick,omitempty"` + Status string `json:"status,omitempty"` + IsResolve bool `json:"is_resolve,omitempty"` + CompletedEarlier string `json:"completed_earlier,omitempty"` } -// ResultItem is a single traversal step result for API consumers. +// ResultItem is one aggregated leaf outcome for API consumers. Parent and +// ParentIP identify the referring server so clients can render the noglue and +// lame-referral wordings; Qname/Qclass/Qtype are the failing query so clients +// can render the "While querying" line when it differs from the original. type ResultItem struct { - Depth int `json:"depth"` - Probability float64 `json:"probability"` - ResponseType string `json:"response_type"` - Server string `json:"server,omitempty"` - Answers []string `json:"answers,omitempty"` - CNAMEChain []string `json:"cname_chain,omitempty"` + RefID string `json:"refid,omitempty"` + Depth int `json:"depth"` + Probability float64 `json:"probability"` + Status string `json:"status"` + Server string `json:"server,omitempty"` + IP string `json:"ip,omitempty"` + Parent string `json:"parent,omitempty"` + ParentIP string `json:"parent_ip,omitempty"` + Qname string `json:"qname,omitempty"` + Qclass string `json:"qclass,omitempty"` + Qtype string `json:"qtype,omitempty"` + Answers []string `json:"answers,omitempty"` + Message string `json:"message,omitempty"` +} + +// SummaryAnswer is one distinct answered RRset with its accumulated +// probability (traverse.SummaryStats). +type SummaryAnswer struct { + Probability float64 `json:"probability"` + Records []string `json:"records"` +} + +// SummaryStatus is the accumulated probability of one non-answered status. +type SummaryStatus struct { + Status string `json:"status"` + Probability float64 `json:"probability"` +} + +// Summary is the grouped view of the aggregated leaves; probabilities across +// Answers plus ByStatus sum to 1.0. +type Summary struct { + Answers []SummaryAnswer `json:"answers,omitempty"` + ByStatus []SummaryStatus `json:"by_status,omitempty"` } // TraversalJob holds all state for a single asynchronous traversal. @@ -66,6 +102,7 @@ type TraversalJob struct { Domain string `json:"domain"` QueryType string `json:"query_type"` Results []ResultItem `json:"results,omitempty"` + Summary *Summary `json:"summary,omitempty"` Progress []ProgressEvent `json:"progress,omitempty"` Error string `json:"error,omitempty"` StartedAt time.Time `json:"started_at"` @@ -76,14 +113,18 @@ type TraversalJob struct { cancel context.CancelFunc } -// subscribe returns a channel that receives future progress events. -// The channel is closed when the job finishes. -func (j *TraversalJob) subscribe() <-chan ProgressEvent { +// subscribeSnapshot atomically registers a subscriber and snapshots the +// progress recorded so far. Publishing appends to Progress and sends to +// subscribers under the same lock, so every event lands either in the +// returned snapshot or on the channel — never both, never neither. +func (j *TraversalJob) subscribeSnapshot() (sub <-chan ProgressEvent, past []ProgressEvent, done bool) { ch := make(chan ProgressEvent, 32) j.mu.Lock() + defer j.mu.Unlock() j.subs = append(j.subs, ch) - j.mu.Unlock() - return ch + past = make([]ProgressEvent, len(j.Progress)) + copy(past, j.Progress) + return ch, past, j.Status != statusRunning } // publishLocked sends ev to all current subscribers. Caller must hold j.mu. @@ -259,6 +300,7 @@ func (h *Handler) getTraversal(w http.ResponseWriter, r *http.Request) { Domain string `json:"domain"` QueryType string `json:"query_type"` Results []ResultItem `json:"results,omitempty"` + Summary *Summary `json:"summary,omitempty"` Progress []ProgressEvent `json:"progress,omitempty"` Error string `json:"error,omitempty"` StartedAt time.Time `json:"started_at"` @@ -269,6 +311,7 @@ func (h *Handler) getTraversal(w http.ResponseWriter, r *http.Request) { Domain: job.Domain, QueryType: job.QueryType, Results: job.Results, + Summary: job.Summary, Progress: job.Progress, Error: job.Error, StartedAt: job.StartedAt, @@ -300,19 +343,12 @@ func (h *Handler) streamTraversal(w http.ResponseWriter, r *http.Request) { return } - // Subscribe before snapshotting progress so we don't miss events between - // the two operations. Unsubscribe when the client disconnects so stale - // channels don't accumulate. - sub := job.subscribe() + // Subscribe and snapshot atomically so events published in between are + // neither missed nor delivered twice. Unsubscribe when the client + // disconnects so stale channels don't accumulate. + sub, past, alreadyDone := job.subscribeSnapshot() defer job.unsubscribe(sub) - // Replay events already recorded. - job.mu.RLock() - past := make([]ProgressEvent, len(job.Progress)) - copy(past, job.Progress) - alreadyDone := job.Status != statusRunning - job.mu.RUnlock() - sendSSE := func(ev ProgressEvent) bool { b, err := json.Marshal(ev) if err != nil { @@ -375,26 +411,22 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st } cfg.Hooks = &traverse.TraverserHooks{ OnEvent: func(event traverse.TraversalEvent) { - ref := event.Result.Referral - if ref == nil { + ref := event.Referral + if ref == nil || ref.IsRootRoot() { return } - stage := "start" - if event.Stage == traverse.EventComplete { - stage = "complete" - } - server := "" - if event.Result.Response != nil && event.Result.Response.Server != nil { - server = event.Result.Response.Server.String() - } ev := ProgressEvent{ - Stage: stage, - Depth: ref.Depth, - Name: trimFQDN(ref.Name), - QType: idns.QNameType(ref.Qtype), - Server: server, - Bailiwick: trimFQDN(ref.Bailiwick), - IsResolve: event.IsResolve, + Stage: event.Stage.String(), + RefID: event.RefID, + Depth: ref.Depth(), + Name: ref.Qname, + QType: traverse.TypeToString(ref.Qtype), + Server: ref.Server, + IPs: ref.TxtIPs(), + Bailiwick: ref.Bailiwick, + Status: string(event.Status), + IsResolve: event.IsResolve, + CompletedEarlier: event.CompletedEarlier, } job.mu.Lock() @@ -405,7 +437,7 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st } tr := traverse.NewTraverser(cfg) - rawResults, err := tr.Traverse(ctx, domain) + root, err := tr.Run(ctx, domain) now := time.Now() job.mu.Lock() @@ -419,31 +451,87 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st return } - items := make([]ResultItem, 0, len(rawResults)) - for _, r := range rawResults { - items = append(items, toResultItem(r)) + var items []ResultItem + if root != nil { + leaves := root.StatsList() + items = make([]ResultItem, 0, len(leaves)) + for _, leaf := range leaves { + items = append(items, toResultItem(leaf)) + } + job.Summary = toSummary(root.SummaryStats()) } job.Results = items job.Status = statusComplete } -// toResultItem converts a TraversalResult to its API representation. -func toResultItem(r traverse.TraversalResult) ResultItem { - item := ResultItem{} - if r.Referral != nil { - item.Depth = r.Referral.Depth - item.Probability = r.Referral.Prob +// toSummary converts the engine's grouped stats to the API representation: +// answers sorted by RRset key (as SummaryStats returns them), remaining +// statuses sorted lexically like the CLI Summary Results section. +func toSummary(stats *traverse.SummaryStats) *Summary { + if stats == nil { + return nil } - if r.Response != nil { - item.ResponseType = r.Response.Type.String() - if r.Response.Server != nil { - item.Server = r.Response.Server.String() + summary := &Summary{} + for _, answer := range stats.Answers { + item := SummaryAnswer{Probability: answer.Prob} + for _, rr := range answer.RRs { + item.Records = append(item.Records, collapseWhitespace(rr.String())) } - if r.Response.Decoded != nil { - for _, rr := range r.Response.Decoded.Answers { - item.Answers = append(item.Answers, idns.FormatRecord(rr)) - } - item.CNAMEChain = append(item.CNAMEChain, r.Response.Decoded.CNAMEChain...) + summary.Answers = append(summary.Answers, item) + } + statuses := make([]traverse.Status, 0, len(stats.ByStatus)) + for status := range stats.ByStatus { + if status != traverse.StatusAnswered { + statuses = append(statuses, status) + } + } + sort.Slice(statuses, func(i, j int) bool { return statuses[i] < statuses[j] }) + for _, status := range statuses { + summary.ByStatus = append(summary.ByStatus, SummaryStatus{ + Status: string(status), + Probability: stats.ByStatus[status], + }) + } + return summary +} + +// collapseWhitespace renders an RR on one line with runs of whitespace +// collapsed to single spaces, matching the CLI summary records. +func collapseWhitespace(s string) string { + return strings.Join(strings.Fields(s), " ") +} + +// toResultItem converts one aggregated leaf to its API representation. +func toResultItem(leaf *traverse.StatsEntry) ResultItem { + resp := leaf.Response + item := ResultItem{ + Probability: leaf.Prob, + Status: string(resp.Status), + IP: resp.IP, + Server: resp.Server, + ParentIP: resp.ParentIP, + Qname: trimFQDN(resp.Qname), + Qclass: traverse.ClassToString(resp.Qclass), + Qtype: traverse.TypeToString(resp.Qtype), + } + if leaf.Referral != nil { + item.RefID = leaf.Referral.RefID + item.Depth = leaf.Referral.Depth() + item.Server = leaf.Referral.Server + item.ParentIP = leaf.Referral.ParentIP + if leaf.Referral.Parent != nil { + item.Parent = leaf.Referral.Parent.Server + } + } + if resp.DQ != nil { + for _, rr := range resp.DQ.Answers { + item.Answers = append(item.Answers, idns.FormatRecord(rr)) + } + switch resp.Status { + case traverse.StatusError: + item.Message = resp.DQ.ErrorMessage + case traverse.StatusException: + item.Message = resp.DQ.ExceptionMessage } } return item @@ -495,4 +583,3 @@ func newUUID() string { b[8] = (b[8] & 0x3f) | 0x80 // variant bits return fmt.Sprintf("%08x-%04x-%04x-%04x-%012x", b[0:4], b[4:6], b[6:8], b[8:10], b[10:]) } - diff --git a/web/api/handler_internal_test.go b/web/api/handler_internal_test.go new file mode 100644 index 0000000..4324c37 --- /dev/null +++ b/web/api/handler_internal_test.go @@ -0,0 +1,77 @@ +package api + +import ( + "strconv" + "sync" + "testing" +) + +// TestSubscribeSnapshot_NoDuplicates regresses the subscribe/snapshot race: +// subscribers arriving while events are being published must never see the +// same event twice (once from the snapshot replay and once from the channel). +// Events are numbered, so any duplicate breaks strict monotonicity. Slow +// subscribers may legitimately drop events (publishLocked is non-blocking), +// so gaps are not an error. +func TestSubscribeSnapshot_NoDuplicates(t *testing.T) { + job := &TraversalJob{Status: statusRunning} + + const total = 2000 + const subscribers = 8 + + var pub sync.WaitGroup + pub.Add(1) + go func() { + defer pub.Done() + for i := 0; i < total; i++ { + ev := ProgressEvent{RefID: strconv.Itoa(i)} + job.mu.Lock() + job.Progress = append(job.Progress, ev) + job.publishLocked(ev) + job.mu.Unlock() + } + job.mu.Lock() + job.Status = statusComplete + job.mu.Unlock() + job.closeSubscribers() + }() + + var subs sync.WaitGroup + errs := make(chan string, subscribers) + for s := 0; s < subscribers; s++ { + subs.Add(1) + go func() { + defer subs.Done() + sub, past, done := job.subscribeSnapshot() + defer job.unsubscribe(sub) + + last := -1 + check := func(refid string) { + n, err := strconv.Atoi(refid) + if err != nil { + errs <- "bad refid " + refid + return + } + if n <= last { + errs <- "event " + refid + " out of order or duplicated after " + strconv.Itoa(last) + return + } + last = n + } + for _, ev := range past { + check(ev.RefID) + } + if !done { + for ev := range sub { + check(ev.RefID) + } + } + }() + } + + pub.Wait() + subs.Wait() + close(errs) + for msg := range errs { + t.Error(msg) + } +} diff --git a/web/api/handler_test.go b/web/api/handler_test.go index 31a8a32..c0b09ad 100644 --- a/web/api/handler_test.go +++ b/web/api/handler_test.go @@ -5,12 +5,15 @@ import ( "bytes" "encoding/json" "fmt" + "io" "net/http" "net/http/httptest" + "regexp" "strings" "testing" "time" + "gitea.hansenits.com.au/hits/ExploreDNS/internal/config" "gitea.hansenits.com.au/hits/ExploreDNS/web/api" ) @@ -264,6 +267,43 @@ func TestStaticSPA_Index(t *testing.T) { } } +// TestStaticSPA_TypeOptions asserts the SPA type dropdown offers exactly the +// query types config.ParseQueryType accepts. +func TestStaticSPA_TypeOptions(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) + + re := regexp.MustCompile(``) + var got []string + for _, m := range re.FindAllStringSubmatch(page, -1) { + got = append(got, m[1]) + } + want := []string{"A", "AAAA", "NS", "CNAME", "MX", "TXT", "SOA", "PTR", "ANY"} + if len(got) != len(want) { + t.Fatalf("type options = %v, want %v", got, want) + } + for i, typ := range want { + if got[i] != typ { + t.Fatalf("type options = %v, want %v", got, want) + } + if _, err := config.ParseQueryType(typ); err != nil { + t.Fatalf("option %s rejected by ParseQueryType: %v", typ, err) + } + } +} + func TestStaticSPA_FallbackToIndex(t *testing.T) { srv := newTestServer(t) defer srv.Shutdown(5 * time.Second) //nolint:errcheck diff --git a/web/api/static/index.html b/web/api/static/index.html index 2212f65..8e101b2 100644 --- a/web/api/static/index.html +++ b/web/api/static/index.html @@ -285,22 +285,28 @@ content: 'Awaiting traversal…'; color: var(--text-dim); } - .pe { display: flex; gap: 0.5rem; } + .pe { display: flex; gap: 0.5rem; align-items: baseline; } .pe-badge { flex-shrink: 0; + min-width: 70px; + text-align: center; font-size: 0.65rem; padding: 1px 5px; border-radius: 3px; font-weight: 600; text-transform: uppercase; letter-spacing: 0.04em; + background: rgba(139,143,168,0.15); + color: var(--text-muted); } - .pe-badge.start { background: rgba(74,124,246,0.18); color: var(--primary); } - .pe-badge.complete { background: rgba(62,207,142,0.18); color: var(--success); } - .pe-badge.resolve { background: rgba(56,189,248,0.18); color: var(--info); } - .pe-depth { color: var(--text-dim); } - .pe-name { color: var(--text); } - .pe-server { color: var(--text-muted); } + .pe-badge.st-working { background: rgba(74,124,246,0.18); color: var(--primary); } + .pe-badge.st-answered { background: rgba(62,207,142,0.18); color: var(--success); } + .pe-badge.st-resolve { background: rgba(56,189,248,0.18); color: var(--info); } + .pe-badge.st-warn { background: rgba(245,166,35,0.18); color: var(--warning); } + .pe-badge.st-err { background: rgba(240,68,56,0.18); color: var(--danger); } + /* Progress lines mirror the CLI: " ( )" — keep the + spacing intact. */ + .pe-text { color: var(--text); white-space: pre; } /* ── Results tree ───────────────────────────────────────────────── */ .results-section h2 { @@ -318,57 +324,44 @@ color: var(--text-dim); font-size: 0.875rem; } - .result-item { + /* Aggregated leaves rendered like the CLI Results section: fixed-width + text with significant leading spaces, so white-space must be pre. */ + .result-block { font-family: var(--font-mono); - font-size: 0.8rem; + font-size: 0.78rem; + line-height: 1.5; + white-space: pre; + overflow-x: auto; + margin: 0 0 0.5rem; padding: 0.5rem 0.75rem; - border-radius: var(--radius-sm); border: 1px solid var(--border); - margin-bottom: 0.5rem; + border-left-width: 3px; + border-radius: var(--radius-sm); background: var(--bg-card); - transition: border-color 0.12s; + color: var(--text); } - .result-item:hover { border-color: var(--border-focus); } - .result-header { - display: flex; - align-items: center; - gap: 0.5rem; - flex-wrap: wrap; - } - .result-indent { color: var(--text-dim); flex-shrink: 0; } - .rtype-badge { - font-size: 0.65rem; - padding: 1px 5px; - border-radius: 3px; - font-weight: 700; - text-transform: uppercase; - } - .rtype-answer { background: rgba(62,207,142,0.15); color: var(--success); } - .rtype-referral { background: rgba(74,124,246,0.15); color: var(--primary); } - .rtype-nxdomain, - .rtype-servfail { background: rgba(240,68,56,0.15); color: var(--danger); } - .rtype-timeout { background: rgba(245,166,35,0.15); color: var(--warning); } - .rtype-nodata { background: rgba(56,189,248,0.15); color: var(--info); } - .rtype-cname { background: rgba(168,85,247,0.15); color: #c084fc; } - .rtype-other { background: rgba(139,143,168,0.15);color: var(--text-muted); } + .result-block.st-answered { border-left-color: var(--success); } + .result-block.st-warn { border-left-color: var(--warning); } + .result-block.st-err { border-left-color: var(--danger); } + .result-block.st-other { border-left-color: var(--text-dim); } - .result-prob { - color: var(--text-muted); - font-size: 0.72rem; - } - .result-server { color: var(--text-dim); font-size: 0.72rem; margin-left: auto; } - .result-answers { - margin-top: 0.35rem; - padding-left: 1.2rem; - color: var(--text-muted); + /* ── Summary Results ───────────────────────────────────────────── */ + .summary-pre { + font-family: var(--font-mono); font-size: 0.75rem; + line-height: 1.6; + white-space: pre; + overflow-x: auto; + margin: 0 0 1rem; + padding: 0.6rem 0.75rem; + border: 1px solid var(--border); + border-radius: var(--radius-sm); + background: var(--bg); + color: var(--text); } - .result-answers span { display: block; } - .cname-chain { - color: #c084fc; - font-size: 0.72rem; - margin-top: 0.2rem; - padding-left: 1.2rem; + .summary-pre:empty::before { + content: 'No summary yet.'; + color: var(--text-dim); } /* ── Stats panel ────────────────────────────────────────────────── */ @@ -478,17 +471,16 @@ autocapitalize="off" spellcheck="false" /> +