mirror of
https://github.com/Control-D-Inc/ctrld.git
synced 2026-09-04 13:36:35 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3fa03d92d6 | ||
|
|
8013bb72d2 | ||
|
|
d2aa155ff3 | ||
|
|
f8f66609da | ||
|
|
d78e9bcf5b | ||
|
|
30acb846ca | ||
|
|
6615e431dc | ||
|
|
1f001a559a | ||
|
|
753d245029 | ||
|
|
8ce3b7ca6c | ||
|
|
084c785ed5 | ||
|
|
5c9d3dec4e | ||
|
|
2400f27962 | ||
|
|
4f730167d4 | ||
|
|
b74937fcf3 | ||
|
|
dfaad4a20d | ||
|
|
d7c30b18ed | ||
|
|
a828c8853a | ||
|
|
4d026d836c | ||
|
|
779fe015f0 | ||
|
|
246c1b9691 | ||
|
|
959f49dae3 | ||
|
|
8ebe911b1a | ||
|
|
d38538f593 | ||
|
|
737fc79b58 | ||
|
|
fa074f1f5e | ||
|
|
c596ef586b | ||
|
|
a4cfd4e479 | ||
|
|
b79098658a | ||
|
|
d29e7d131e | ||
|
|
836c9ccf12 | ||
|
|
53d3d3d44a | ||
|
|
d7f43ea4bf | ||
|
|
41ca69849a | ||
|
|
0d8df38dc1 | ||
|
|
3ef17bc5b9 | ||
|
|
5bf26da585 | ||
|
|
a5d536ab79 | ||
|
|
735590d244 | ||
|
|
18f01baa01 | ||
|
|
723c7827ba | ||
|
|
1e1c998c89 | ||
|
|
da454db8ef | ||
|
|
3fe9b27fb4 | ||
|
|
35455eb0b9 | ||
|
|
f1309121ae | ||
|
|
06668a2b6c | ||
|
|
97e5e99b8d | ||
|
|
c54ff701bd | ||
|
|
33682e2312 | ||
|
|
d629ecda33 | ||
|
|
87ddf03b90 | ||
|
|
d49a4c67c9 | ||
|
|
2c38ff74c3 | ||
|
|
75e8447c75 | ||
|
|
4395efcb22 |
@@ -9,7 +9,7 @@ jobs:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
os: ["windows-latest", "ubuntu-latest", "macOS-latest"]
|
||||
go: ["1.24.x"]
|
||||
go: ["1.26.x"]
|
||||
runs-on: ${{ matrix.os }}
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
@@ -19,8 +19,8 @@ jobs:
|
||||
with:
|
||||
go-version: ${{ matrix.go }}
|
||||
- run: "go test -race ./..."
|
||||
- uses: dominikh/staticcheck-action@v1.4.0
|
||||
- uses: dominikh/staticcheck-action@v1.4.1
|
||||
with:
|
||||
version: "2025.1.1"
|
||||
version: "2026.2"
|
||||
install-go: false
|
||||
cache-key: ${{ matrix.go }}
|
||||
|
||||
@@ -0,0 +1,206 @@
|
||||
# SPEC: Stable customer-visible provisioning failure codes
|
||||
|
||||
Issue: [#586](https://gitlab.int.windscribe.com/controld/clients/ctrld/-/issues/586)
|
||||
Requested by: Catt Garrod (@catt). Scope expanded by: Anthony Wong (@anthony).
|
||||
|
||||
## 1. Objective
|
||||
|
||||
Terminal provisioning failures in ctrld — bootstrap/API setup, listener
|
||||
binding, and service installation/startup — must produce a stable,
|
||||
support-facing failure identifier that survives process exit and reaches
|
||||
both manual CLI users and MDM-driven installs. A customer or admin reports
|
||||
one code; Support maps it to a scenario and a next action without asking
|
||||
for reruns or verbose logs.
|
||||
|
||||
Motivating incident (v1.5.5, macOS): provisioning reached the Control D
|
||||
API, then died with only `FTL listener.0 could not find available listen
|
||||
ip and port`. The per-address UDP/TCP bind errors existed only at Info
|
||||
level in an in-memory logger and vanished on exit. The macOS pkg
|
||||
`postinstall` discards ctrld's stdout/stderr entirely and judges success
|
||||
by plist existence, so nothing useful reached the MDM log.
|
||||
|
||||
**Users:** end customers and IT admins reporting failures; Support agents
|
||||
triaging them; MDM/RMM operators reading installer logs.
|
||||
|
||||
### Failure contract (agreed design)
|
||||
|
||||
Three surfaces, all carrying the same identifier:
|
||||
|
||||
1. **Result file** — on terminal provisioning failure, ctrld writes a
|
||||
small redacted JSON file (atomic write: temp + rename) in the ctrld
|
||||
home directory (same base dir as the internal `ctrld.log`,
|
||||
via `absHomeDir`). Removed/overwritten on later successful
|
||||
provisioning so stale failures don't mislead. Schema:
|
||||
|
||||
```json
|
||||
{
|
||||
"version": 1,
|
||||
"timestamp": "2026-08-18T12:00:00Z",
|
||||
"stage": "listener",
|
||||
"code": "LISTENER_BIND_FAILED",
|
||||
"exit_code": 41,
|
||||
"message": "could not find available listen ip and port",
|
||||
"detail": {
|
||||
"attempts": [
|
||||
{"addr": "127.0.0.1:53", "proto": "udp", "os_error": "address already in use"}
|
||||
]
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
`detail` is bounded (cap recorded bind attempts; cap string lengths)
|
||||
and redacted by construction: no provisioning tokens, resolver IDs,
|
||||
config contents, or unrelated host data.
|
||||
|
||||
2. **Exit code + final stderr line** — the installer-facing command
|
||||
(`ctrld start`, and `ctrld run` when run manually in the foreground)
|
||||
exits with a stage-scoped code and prints one final line containing
|
||||
the string code and stage, e.g.
|
||||
`provisioning failed: stage=listener code=LISTENER_BIND_FAILED (exit 41)`.
|
||||
|
||||
3. **Installer log (MDM path)** — `scripts/pkg/postinstall` stops
|
||||
discarding the signal: it captures `ctrld start`'s output to a
|
||||
private temp file, extracts only the fixed-charset identifier line
|
||||
(`stage=[a-z]* code=[A-Z_]* (exit [0-9]*)` — structurally unable to
|
||||
carry the token), and echoes it with the exit code into the
|
||||
installer log. The result file's `message`/`detail` fields are
|
||||
deliberately never surfaced there. The plist-existence check remains
|
||||
the final success gate.
|
||||
|
||||
### Identifier format
|
||||
|
||||
- **Primary identifier: stable string codes.** Initial set —
|
||||
bootstrap: `API_UNREACHABLE`, `API_REJECTED`, `API_DEVICE_INVALID`;
|
||||
listener: `LISTENER_BIND_FAILED`, `LISTENER_CONFIGURED_ADDR_UNAVAILABLE`;
|
||||
service: `SERVICE_INSTALL_FAILED`, `SERVICE_START_FAILED`,
|
||||
`SERVICE_SELFCHECK_FAILED`. Codes are append-only; renames are new
|
||||
codes plus a deprecation note in the mapping doc.
|
||||
- **Secondary: stage-scoped process exit codes** as a coarse machine
|
||||
signal: bootstrap 30–39, listener 40–49, service install/start 50–59.
|
||||
Each string code owns one exit code. Existing contracts are untouched:
|
||||
`ctrld status` 0–3, deactivation-pin 126, success 0.
|
||||
- One underlying failure maps to one code on every path (manual CLI and
|
||||
MDM), on both branches.
|
||||
|
||||
### Propagation (daemon → installer)
|
||||
|
||||
The listener/bootstrap fatals fire inside the daemon process
|
||||
(`ctrld run` under launchd/systemd/SCM), not in `ctrld start`. The
|
||||
daemon writes the result file before exiting; the existing log-socket
|
||||
exit notification (`notifyExitToLogServer`) already unblocks `ctrld
|
||||
start`'s self-check. `ctrld start` then reads the result file, prints
|
||||
the identifier, and exits with the mapped stage exit code. The daemon's
|
||||
own exit-status semantics toward service managers are preserved —
|
||||
in particular the deliberate exit-0 on permanent API rejection that
|
||||
protects the restart-policy budget; the result file carries the failure
|
||||
identity in that case.
|
||||
|
||||
### Support mapping
|
||||
|
||||
`docs/provisioning-failure-codes.md` in this repo: one row per code —
|
||||
code, stage, exit code, failure scenario, next safe troubleshooting
|
||||
action or evidence request. Updated in the same MR whenever a code is
|
||||
added or changed.
|
||||
|
||||
### Branch scope
|
||||
|
||||
Full implementation on **both** `v1.0` (release line for v1.5.5) and
|
||||
`master`. The branches diverge heavily (`v1.0`: zerolog fork,
|
||||
`commands.go`, `service_status.go`, macOS pkg scripts; `master`: zap,
|
||||
inline commands, no pkg scripts), so this is one shared contract
|
||||
(codes, exit-code ranges, file schema, doc) implemented twice, as two
|
||||
MRs referencing #586.
|
||||
|
||||
## 2. Commands
|
||||
|
||||
- Build: `go build ./...`
|
||||
- Test: `go test ./cmd/cli/...` (full: `go test ./...`)
|
||||
- Vet: `go vet ./...`
|
||||
- Branch workflow: feature branch off `v1.0` for the v1.0 MR; separate
|
||||
feature branch off `master` for the port MR. Rebase, never merge the
|
||||
base branch in.
|
||||
|
||||
## 3. Project structure
|
||||
|
||||
New and touched files on `v1.0` (master port mirrors the same contract
|
||||
at its equivalent emission points in its `cli.go`):
|
||||
|
||||
- `cmd/cli/provision_result.go` (new) — stage + code enums, exit-code
|
||||
mapping, result-file schema, atomic write/read/clear helpers,
|
||||
bounded/redacted detail builders. Pattern follows `service_status.go`
|
||||
(small file: named constants + classifier + dedicated tests).
|
||||
- `cmd/cli/provision_result_test.go` (new).
|
||||
- `cmd/cli/cli.go` — emission points: `run()` bootstrap failure branches
|
||||
(permanent rejection, invalid-device, fatal fetch), and
|
||||
`tryUpdateListenerConfig` / `tryUpdateListenerConfigIntercept` fatals,
|
||||
which now record per-attempt `{addr, proto, os_error}` bind detail.
|
||||
- `cmd/cli/commands.go` — `initStartCmd`: doTasks install/start failures
|
||||
and the self-check failure branch read the result file, print the
|
||||
identifier, and exit with the stage code (replacing bare `os.Exit(1)`
|
||||
on those paths).
|
||||
- `scripts/pkg/postinstall` — propagate exit code + result-file contents
|
||||
into the installer log (v1.0 only; master has no pkg scripts).
|
||||
- `docs/provisioning-failure-codes.md` (new) — support mapping.
|
||||
|
||||
## 4. Code style
|
||||
|
||||
- Per repo conventions and global rules: guard clauses, small functions,
|
||||
descriptive names, explicit error handling — never weaken existing
|
||||
handling (e.g. keep the permanent-rejection exit-0 rationale intact).
|
||||
- Comments only for non-obvious constraints (e.g. why the daemon must
|
||||
still exit 0 on permanent rejection), simple-english, self-contained —
|
||||
no issue/MR references in code.
|
||||
- Match each branch's logging idiom: zerolog fork on `v1.0`, zap on
|
||||
`master`. No new dependencies.
|
||||
- Conventional Commits; MR titles in simple-english; both MRs reference
|
||||
#586 (release-line MR carries `Closes #586`).
|
||||
|
||||
## 5. Testing strategy
|
||||
|
||||
Test-first where the harness allows. Coverage required by the issue:
|
||||
|
||||
- **Code/mapping unit tests** — every string code maps to exactly one
|
||||
stage and one in-range exit code; ranges don't collide with existing
|
||||
contracts (0–3 status, 126 pin).
|
||||
- **Result file round-trip** — write/read/clear; atomic write; stale
|
||||
file removed on success.
|
||||
- **Redaction** — serialize a result built from inputs containing a
|
||||
provision token, resolver ID, and config content; assert none appear.
|
||||
- **Listener bind failure (regression test for the incident)** — occupy
|
||||
a port, drive the listener-config path to exhaustion, assert the
|
||||
result records `LISTENER_BIND_FAILED` with attempted address, UDP/TCP
|
||||
operation, and OS error (`address already in use`-class).
|
||||
- **Bootstrap failures** — mock API: permanent 4xx → `API_REJECTED`;
|
||||
invalid-device 40402 → `API_DEVICE_INVALID`; unreachable →
|
||||
`API_UNREACHABLE`.
|
||||
- **Service install/start/self-check failures** — injected task
|
||||
failures assert code selection and `ctrld start` exit code.
|
||||
- **MDM surface** — shell-level check of `postinstall` failure branch
|
||||
(result file present → correct log line and exit), aligned with the
|
||||
existing `test-scripts/` approach; manual pkg verification steps
|
||||
documented in the MR.
|
||||
- Both branches: the shared contract tests exist on both; branch-specific
|
||||
emission tests match each branch's structure.
|
||||
|
||||
## 6. Boundaries
|
||||
|
||||
**Always:**
|
||||
- Redact tokens, resolver IDs, config contents, host data from every
|
||||
customer-visible surface (result file, stderr line, installer log).
|
||||
- Preserve existing exit-code contracts (`ctrld status` 0–3, pin 126)
|
||||
and the daemon's service-manager-facing exit semantics.
|
||||
- Bound all recorded detail (attempt counts, string lengths).
|
||||
- Keep codes append-only once merged.
|
||||
|
||||
**Ask first:**
|
||||
- Changing the daemon's (`ctrld run` under a service manager) exit codes
|
||||
or restart-relevant behavior beyond writing the result file.
|
||||
- Adding any persisted file outside the ctrld home directory.
|
||||
- Expanding scope to runtime (post-provisioning) failures — this ticket
|
||||
owns terminal provisioning failures only.
|
||||
|
||||
**Never:**
|
||||
- Print or persist the provisioning token (the reason postinstall
|
||||
discards output today — the replacement surface must stay token-free).
|
||||
- Auto-detect or kill conflicting processes (explicitly out of scope).
|
||||
- Break `ctrld status`'s documented exit-code contract.
|
||||
+432
-84
@@ -147,6 +147,25 @@ func isMobile() bool {
|
||||
return runtime.GOOS == "android" || runtime.GOOS == "ios"
|
||||
}
|
||||
|
||||
func updateConfigInterceptMode(cfg *ctrld.Config, mode string) bool {
|
||||
desired := ""
|
||||
switch mode {
|
||||
case "dns", "hard":
|
||||
desired = mode
|
||||
case "off":
|
||||
desired = ""
|
||||
case "":
|
||||
return false
|
||||
default:
|
||||
return false
|
||||
}
|
||||
if cfg.Service.InterceptMode == desired {
|
||||
return false
|
||||
}
|
||||
cfg.Service.InterceptMode = desired
|
||||
return true
|
||||
}
|
||||
|
||||
// isAndroid reports whether the current OS is Android.
|
||||
func isAndroid() bool {
|
||||
return runtime.GOOS == "android"
|
||||
@@ -318,41 +337,55 @@ func run(appCallback *AppCallback, stopCh chan struct{}) {
|
||||
}
|
||||
if cdUID != "" {
|
||||
validateCdUpstreamProtocol()
|
||||
if rc, err := processCDFlags(&cfg); err != nil {
|
||||
// Bound API preflight by the service lifetime. Without this, a stop request
|
||||
// arriving while the API is unreachable leaves this retry/backoff loop running
|
||||
// after "service stopped" was logged, so the process keeps working on behalf of
|
||||
// a service the OS considers stopped.
|
||||
pf := runAPIPreflight(p.stopCh, &cfg)
|
||||
switch {
|
||||
case pf.stopRequested:
|
||||
// Stop requested during preflight, whether or not the fetch itself
|
||||
// succeeded. A successful fetch does not entitle startup to continue: the
|
||||
// operator asked for a stop, and carrying on would set up listeners and
|
||||
// interception for a service the OS already considers stopping.
|
||||
//
|
||||
// Exit the way a normal stop does: no Fatal, so the OS service manager does
|
||||
// not see a failed start and apply its restart policy to a service the
|
||||
// operator just asked to stop.
|
||||
mainLog.Load().Notice().Msg("stop requested while fetching resolver config, shutting down")
|
||||
notifyExitToLogServer()
|
||||
return
|
||||
case pf.err != nil:
|
||||
if isMobile() {
|
||||
appCallback.Exit(err.Error())
|
||||
appCallback.Exit(pf.err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
cdLogger := mainLog.Load().With().Str("mode", "cd").Logger()
|
||||
// Performs self-uninstallation if the ControlD device does not exist.
|
||||
var uer *controld.ErrorResponse
|
||||
if errors.As(err, &uer) && uer.ErrorField.Code == controld.InvalidConfigCode {
|
||||
_ = uninstallInvalidCdUID(p, cdLogger, false)
|
||||
}
|
||||
notifyExitToLogServer()
|
||||
cdLogger.Fatal().Err(err).Msg("failed to fetch resolver config")
|
||||
} else {
|
||||
handleAPIPreflightFailure(p, pf.err, notifyExitToLogServer)
|
||||
return
|
||||
default:
|
||||
p.mu.Lock()
|
||||
p.rc = rc
|
||||
p.rc = pf.rc
|
||||
p.mu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
updated := updateListenerConfig(&cfg, notifyExitToLogServer)
|
||||
|
||||
// Bootstrap and listener binding both succeeded, so an earlier run's
|
||||
// recorded failure no longer describes this install.
|
||||
clearProvisionResult()
|
||||
|
||||
if cdUID != "" {
|
||||
processLogAndCacheFlags(v, &cfg)
|
||||
}
|
||||
|
||||
// Persist intercept_mode to config when provided via CLI flag on full install.
|
||||
// This ensures the config file reflects the actual running mode for RMM/MDM visibility.
|
||||
if interceptMode == "dns" || interceptMode == "hard" {
|
||||
if cfg.Service.InterceptMode != interceptMode {
|
||||
cfg.Service.InterceptMode = interceptMode
|
||||
updated = true
|
||||
mainLog.Load().Info().Msgf("writing intercept_mode = %q to config", interceptMode)
|
||||
}
|
||||
// Keep config and the explicit CLI/service mode in sync. In particular, "off"
|
||||
// must clear a previously persisted dns/hard value or the next service start
|
||||
// would silently re-enable interception from config.
|
||||
if updateConfigInterceptMode(&cfg, interceptMode) {
|
||||
updated = true
|
||||
mainLog.Load().Info().Msgf("writing intercept_mode = %q to config (requested %q)", cfg.Service.InterceptMode, interceptMode)
|
||||
}
|
||||
|
||||
if updated {
|
||||
@@ -451,18 +484,7 @@ func run(appCallback *AppCallback, stopCh chan struct{}) {
|
||||
p.onStopped = append(p.onStopped, func() {
|
||||
// restore static DNS settings or DHCP
|
||||
p.resetDNS(false, true)
|
||||
// Iterate over all physical interfaces and restore static DNS if a saved static config exists.
|
||||
withEachPhysicalInterfaces("", "restore static DNS", func(i *net.Interface) error {
|
||||
file := savedStaticDnsSettingsFilePath(i)
|
||||
if _, err := os.Stat(file); err == nil {
|
||||
if err := restoreDNS(i); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msgf("Could not restore static DNS on interface %s", i.Name)
|
||||
} else {
|
||||
mainLog.Load().Debug().Msgf("Restored static DNS on interface %s successfully", i.Name)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
restoreSavedStaticDNS("", false)
|
||||
})
|
||||
|
||||
close(waitCh)
|
||||
@@ -638,6 +660,19 @@ const defaultDeactivationPin = -1
|
||||
// cdDeactivationPin is used in cd mode to decide whether stop and uninstall commands can be run.
|
||||
var cdDeactivationPin atomic.Int64
|
||||
|
||||
// Brute-force protection for the deactivation PIN endpoint on the control socket.
|
||||
// After deactivationMaxFailedAttempts consecutive wrong PINs, further attempts are
|
||||
// rejected for deactivationLockoutSeconds. Counter resets on a correct PIN.
|
||||
const (
|
||||
deactivationMaxFailedAttempts = 5
|
||||
deactivationLockoutSeconds = 60
|
||||
)
|
||||
|
||||
var (
|
||||
deactivationFailedAttempts atomic.Int64
|
||||
deactivationLockedUntil atomic.Int64
|
||||
)
|
||||
|
||||
func init() {
|
||||
cdDeactivationPin.Store(defaultDeactivationPin)
|
||||
}
|
||||
@@ -647,24 +682,218 @@ func deactivationPinSet() bool {
|
||||
return cdDeactivationPin.Load() != defaultDeactivationPin
|
||||
}
|
||||
|
||||
func processCDFlags(cfg *ctrld.Config) (*controld.ResolverConfig, error) {
|
||||
// fetchResolverConfig is a test seam for the ControlD resolver-config API call.
|
||||
var fetchResolverConfig = controld.FetchResolverConfig
|
||||
|
||||
// apiPreflight is the outcome of the API preflight fetch: the resolver config, the
|
||||
// error if any, and whether the service was asked to stop while it ran.
|
||||
type apiPreflight struct {
|
||||
rc *controld.ResolverConfig
|
||||
err error
|
||||
stopRequested bool
|
||||
}
|
||||
|
||||
// runAPIPreflight fetches the ControlD resolver config bounded by the service
|
||||
// lifetime, and reports whether a stop was requested while it ran.
|
||||
//
|
||||
// The distinction matters because the caller does very different things with it: a stop
|
||||
// exits quietly, while a failure self-uninstalls a deleted device, surfaces the error to
|
||||
// a mobile app, and reports a failed start to the service manager.
|
||||
//
|
||||
// stopRequested must not be derived from the context once it has been cancelled.
|
||||
// context.CancelFunc sets ctx.Err() unconditionally, so reading it after the cancel
|
||||
// classifies *every* failure - a deleted device, an exhausted retry, a mobile caller
|
||||
// with no stop channel - as an operator stop. Reading the stop channel directly is also
|
||||
// independent of whether the context's watcher goroutine has been scheduled yet.
|
||||
func runAPIPreflight(stopCh <-chan struct{}, cfg *ctrld.Config) apiPreflight {
|
||||
rc, err := fetchCDConfigBoundedBy(stopCh, cfg)
|
||||
return apiPreflight{rc: rc, err: err, stopRequested: stopRequested(stopCh)}
|
||||
}
|
||||
|
||||
// permanentAPIRejection reports whether err is the API refusing this request in a way
|
||||
// that a restart cannot change, and returns the rejection when it is.
|
||||
//
|
||||
// The type alone does not answer this. controld builds an *ErrorResponse for *any*
|
||||
// non-200 whose body decodes, so a 502 from a load balancer and a 404 for a deleted
|
||||
// device arrive as the same Go type. Treating both as permanent would let a few minutes
|
||||
// of API trouble stop ctrld on every host with no service-manager retry behind it, which
|
||||
// is strictly worse than the abnormal exit it replaced.
|
||||
//
|
||||
// So the HTTP status decides, and only a client-error status counts:
|
||||
//
|
||||
// - 4xx: the API examined this request and refused it - a deleted device, a revoked
|
||||
// token, a malformed UID. The same request will be refused again.
|
||||
// - 408 and 429 are the exceptions: they are the API asking for another attempt later.
|
||||
// - 5xx, or no recorded status, says nothing about this configuration. Retry.
|
||||
func permanentAPIRejection(err error) (*controld.ErrorResponse, bool) {
|
||||
var uer *controld.ErrorResponse
|
||||
if !errors.As(err, &uer) {
|
||||
return nil, false
|
||||
}
|
||||
switch uer.StatusCode {
|
||||
case http.StatusRequestTimeout, http.StatusTooManyRequests:
|
||||
return nil, false
|
||||
}
|
||||
if uer.StatusCode < 400 || uer.StatusCode >= 500 {
|
||||
return nil, false
|
||||
}
|
||||
return uer, true
|
||||
}
|
||||
|
||||
// apiFailureCode maps a bootstrap preflight error to its provisioning code.
|
||||
// A deleted device gets its own code because it triggers self-uninstall;
|
||||
// other permanent rejections are generic; anything else counts as
|
||||
// reachability trouble worth retrying.
|
||||
func apiFailureCode(err error) (provisionFailureCode, bool) {
|
||||
if err == nil {
|
||||
return "", false
|
||||
}
|
||||
var uer *controld.ErrorResponse
|
||||
if errors.As(err, &uer) && uer.ErrorField.Code == controld.InvalidConfigCode {
|
||||
return provisionCodeAPIDeviceInvalid, true
|
||||
}
|
||||
if _, ok := permanentAPIRejection(err); ok {
|
||||
return provisionCodeAPIRejected, true
|
||||
}
|
||||
return provisionCodeAPIUnreachable, true
|
||||
}
|
||||
|
||||
// apiRejectionSummary reports the HTTP status only. The API's raw error body
|
||||
// can echo back the value the caller sent, so it stays out of the artifact.
|
||||
func apiRejectionSummary(statusCode int) string {
|
||||
return fmt.Sprintf("ControlD API rejected this configuration (HTTP status %d)", statusCode)
|
||||
}
|
||||
|
||||
// provisionSecrets lists every secret-bearing value to strip from provisioning
|
||||
// artifacts, including both parts of a composite "<uid>/<clientID>" --cd
|
||||
// value, which the API may echo back separately.
|
||||
func provisionSecrets() []string {
|
||||
uid, clientID := controld.ParseRawUID(cdUID)
|
||||
return []string{cdUID, cdOrg, uid, clientID}
|
||||
}
|
||||
|
||||
// uninstallInvalidCdUIDFn is a var so tests can observe the self-uninstall
|
||||
// without driving the OS service manager.
|
||||
var uninstallInvalidCdUIDFn = uninstallInvalidCdUID
|
||||
|
||||
// handleAPIPreflightFailure reports a failed resolver-config fetch. A deleted
|
||||
// device self-uninstalls; it and any other permanent rejection return cleanly
|
||||
// so a config problem cannot burn the service manager's restart budget (on
|
||||
// Windows those restarts are what bring enforcement back after a real crash).
|
||||
// Anything else exits nonzero through failProvision so the manager retries.
|
||||
func handleAPIPreflightFailure(p *prog, err error, notify func()) {
|
||||
cdLogger := mainLog.Load().With().Str("mode", "cd").Logger()
|
||||
code, _ := apiFailureCode(err)
|
||||
var uer *controld.ErrorResponse
|
||||
if errors.As(err, &uer) && uer.ErrorField.Code == controld.InvalidConfigCode {
|
||||
r := newProvisionResult(code, apiRejectionSummary(uer.StatusCode), nil, provisionSecrets()...)
|
||||
if werr := writeProvisionResult(r); werr != nil {
|
||||
cdLogger.Warn().Err(werr).Msg("could not persist provision result")
|
||||
}
|
||||
_ = uninstallInvalidCdUIDFn(p, cdLogger, false)
|
||||
cdLogger.Error().Err(err).Int("status", uer.StatusCode).Msg("failed to fetch resolver config, the device no longer exists")
|
||||
cdLogger.Error().Msg(r.failureLine())
|
||||
notify()
|
||||
return
|
||||
}
|
||||
if rejection, ok := permanentAPIRejection(err); ok {
|
||||
r := newProvisionResult(code, apiRejectionSummary(rejection.StatusCode), nil, provisionSecrets()...)
|
||||
if werr := writeProvisionResult(r); werr != nil {
|
||||
cdLogger.Warn().Err(werr).Msg("could not persist provision result")
|
||||
}
|
||||
cdLogger.Error().Err(err).Int("status", rejection.StatusCode).Msg("failed to fetch resolver config, the API rejected this configuration")
|
||||
cdLogger.Error().Msg(r.failureLine())
|
||||
notify()
|
||||
return
|
||||
}
|
||||
cdLogger.Error().Err(err).Msg("failed to fetch resolver config")
|
||||
failProvision(newProvisionResult(code, fmt.Sprintf("failed to fetch resolver config: %v", err), nil, provisionSecrets()...), notify)
|
||||
}
|
||||
|
||||
// processCDFlagsFn is the API fetch, indirected so the lifetime binding around it can be
|
||||
// tested without reaching the network.
|
||||
var processCDFlagsFn = processCDFlags
|
||||
|
||||
// fetchCDConfigBoundedBy runs the API fetch bounded by stopCh, so a fetch that cannot
|
||||
// reach the API stops when the service is asked to stop instead of working on behalf of a
|
||||
// service the OS already considers stopped. The derived context is always cancelled, which
|
||||
// releases the goroutine watching stopCh.
|
||||
func fetchCDConfigBoundedBy(stopCh <-chan struct{}, cfg *ctrld.Config) (*controld.ResolverConfig, error) {
|
||||
ctx, cancel := contextFromStopCh(stopCh)
|
||||
defer cancel()
|
||||
return processCDFlagsFn(ctx, cfg)
|
||||
}
|
||||
|
||||
// fetchCDConfigBoundedByLifetime is the reload path's fetch. Reload binds the same stop
|
||||
// primitives as startup - it used to wire them up itself, where a dropped cancel or the
|
||||
// wrong channel would have failed nothing.
|
||||
func (p *prog) fetchCDConfigBoundedByLifetime(cfg *ctrld.Config) (*controld.ResolverConfig, error) {
|
||||
return fetchCDConfigBoundedBy(p.stopCh, cfg)
|
||||
}
|
||||
|
||||
// stopRequested reports whether stopCh has been closed. A nil channel - mobile passes
|
||||
// none - blocks forever, so the default case is taken and it reads as "no stop".
|
||||
func stopRequested(stopCh <-chan struct{}) bool {
|
||||
select {
|
||||
case <-stopCh:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// contextFromStopCh returns a context that is cancelled when stopCh closes, so
|
||||
// long-running startup work stops as soon as the service is asked to stop. The
|
||||
// returned cancel func must be called to release the watcher goroutine.
|
||||
func contextFromStopCh(stopCh <-chan struct{}) (context.Context, context.CancelFunc) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
if stopCh == nil {
|
||||
return ctx, cancel
|
||||
}
|
||||
go func() {
|
||||
select {
|
||||
case <-stopCh:
|
||||
cancel()
|
||||
case <-ctx.Done():
|
||||
}
|
||||
}()
|
||||
return ctx, cancel
|
||||
}
|
||||
|
||||
// processCDFlags fetches the ControlD configuration for cdUID and applies it to cfg.
|
||||
//
|
||||
// ctx bounds the bootstrap-DNS retry loop below. That loop retries indefinitely by
|
||||
// design (a device with no network yet must eventually come up), so it must be
|
||||
// cancellable: otherwise a stop request during preflight is ignored and the process
|
||||
// keeps retrying after the service reports itself stopped.
|
||||
func processCDFlags(ctx context.Context, cfg *ctrld.Config) (*controld.ResolverConfig, error) {
|
||||
logger := mainLog.Load().With().Str("mode", "cd").Logger()
|
||||
logger.Info().Msgf("fetching Controld D configuration from API: %s", cdUID)
|
||||
bo := backoff.NewBackoff("processCDFlags", logf, 30*time.Second)
|
||||
bo.LogLongerThan = 30 * time.Second
|
||||
|
||||
ctx := context.Background()
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
req := &controld.ResolverConfigRequest{
|
||||
RawUID: cdUID,
|
||||
Version: rootCmd.Version,
|
||||
Metadata: ctrld.SystemMetadataRuntime(context.Background()),
|
||||
Metadata: ctrld.SystemMetadataRuntime(ctx),
|
||||
}
|
||||
resolverConfig, err := controld.FetchResolverConfig(req, cdDev)
|
||||
resolverConfig, err := fetchResolverConfig(ctx, req, cdDev)
|
||||
for {
|
||||
if ctxErr := ctx.Err(); ctxErr != nil {
|
||||
logger.Debug().Msg("resolver config fetch cancelled")
|
||||
return nil, ctxErr
|
||||
}
|
||||
if errUrlNetworkError(err) {
|
||||
bo.BackOff(ctx, err)
|
||||
if ctxErr := ctx.Err(); ctxErr != nil {
|
||||
logger.Debug().Msg("resolver config fetch cancelled during backoff")
|
||||
return nil, ctxErr
|
||||
}
|
||||
logger.Warn().Msg("could not fetch resolver using bootstrap DNS, retrying...")
|
||||
resolverConfig, err = controld.FetchResolverConfig(req, cdDev)
|
||||
resolverConfig, err = fetchResolverConfig(ctx, req, cdDev)
|
||||
continue
|
||||
}
|
||||
break
|
||||
@@ -696,7 +925,10 @@ func processCDFlags(cfg *ctrld.Config) (*controld.ResolverConfig, error) {
|
||||
return resolverConfig, nil
|
||||
}
|
||||
}
|
||||
mainLog.Load().Warn().Err(err).Msg("disregarding invalid custom config")
|
||||
// cfgErr, not err: err is the resolver-config fetch error from above, which is
|
||||
// nil on every path that reaches here, so logging it said nothing about why the
|
||||
// custom config was rejected.
|
||||
mainLog.Load().Warn().Err(cfgErr).Msg("disregarding invalid custom config")
|
||||
}
|
||||
|
||||
bootstrapIP := func(endpoint string) string {
|
||||
@@ -813,7 +1045,7 @@ func processLogAndCacheFlags(v *viper.Viper, cfg *ctrld.Config) {
|
||||
}
|
||||
|
||||
func netInterface(ifaceName string) (*net.Interface, error) {
|
||||
if ifaceName == "auto" {
|
||||
if ifaceName == autoIface {
|
||||
ifaceName = defaultIfaceName()
|
||||
}
|
||||
var iface *net.Interface
|
||||
@@ -1099,23 +1331,7 @@ func uninstall(p *prog, s service.Service) {
|
||||
}
|
||||
// restore static DNS settings or DHCP
|
||||
p.resetDNS(false, true)
|
||||
|
||||
// Iterate over all physical interfaces and restore DNS if a saved static config exists.
|
||||
withEachPhysicalInterfaces(p.runningIface, "restore static DNS", func(i *net.Interface) error {
|
||||
file := savedStaticDnsSettingsFilePath(i)
|
||||
if _, err := os.Stat(file); err == nil {
|
||||
if err := restoreDNS(i); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msgf("Could not restore static DNS on interface %s", i.Name)
|
||||
} else {
|
||||
mainLog.Load().Debug().Msgf("Restored static DNS on interface %s successfully", i.Name)
|
||||
err = os.Remove(file)
|
||||
if err != nil {
|
||||
mainLog.Load().Debug().Err(err).Msgf("Could not remove saved static DNS file for interface %s", i.Name)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
restoreSavedStaticDNS(p.runningIface, true)
|
||||
|
||||
if router.Name() != "" {
|
||||
mainLog.Load().Debug().Msg("Router cleanup")
|
||||
@@ -1128,6 +1344,26 @@ func uninstall(p *prog, s service.Service) {
|
||||
}
|
||||
}
|
||||
|
||||
// restoreSavedStaticDNS restores DNS from saved static config files on physical interfaces.
|
||||
func restoreSavedStaticDNS(excludeIfaceName string, removeSaved bool) {
|
||||
withEachPhysicalInterfaces(excludeIfaceName, "restore static DNS", func(i *net.Interface) error {
|
||||
file := savedStaticDnsSettingsFilePath(i)
|
||||
if _, err := os.Stat(file); err == nil {
|
||||
if err := restoreDNS(i); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msgf("Could not restore static DNS on interface %s", i.Name)
|
||||
} else {
|
||||
mainLog.Load().Debug().Msgf("Restored static DNS on interface %s successfully", i.Name)
|
||||
if removeSaved {
|
||||
if err := os.Remove(file); err != nil {
|
||||
mainLog.Load().Debug().Err(err).Msgf("Could not remove saved static DNS file for interface %s", i.Name)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func validateConfig(cfg *ctrld.Config) error {
|
||||
if err := ctrld.ValidateConfig(validator.New(), cfg); err != nil {
|
||||
var ve validator.ValidationErrors
|
||||
@@ -1245,7 +1481,7 @@ func tryUpdateListenerConfigIntercept(cfg *ctrld.Config, notifyFunc func(), fata
|
||||
return false, true
|
||||
}
|
||||
|
||||
hasExplicitConfig := lc.IP != "" && lc.IP != "0.0.0.0" && lc.Port != 0
|
||||
hasExplicitConfig := isExplicitInterceptListener(lc.IP, lc.Port)
|
||||
if !hasExplicitConfig {
|
||||
// Set defaults for intercept mode
|
||||
if lc.IP == "" || lc.IP == "0.0.0.0" {
|
||||
@@ -1258,16 +1494,27 @@ func tryUpdateListenerConfigIntercept(cfg *ctrld.Config, notifyFunc func(), fata
|
||||
}
|
||||
}
|
||||
|
||||
// bindAttempts feeds the provisioning result detail. newProvisionResult
|
||||
// caps it, so it grows freely here.
|
||||
var bindAttempts []provisionBindAttempt
|
||||
recordBindAttempt := func(addr, proto string, err error) {
|
||||
if err != nil {
|
||||
bindAttempts = append(bindAttempts, provisionBindAttempt{Addr: addr, Proto: proto, OSError: err.Error()})
|
||||
}
|
||||
}
|
||||
|
||||
tryListen := func(ip string, port int) bool {
|
||||
addr := net.JoinHostPort(ip, strconv.Itoa(port))
|
||||
udpLn, udpErr := net.ListenPacket("udp", addr)
|
||||
if udpLn != nil {
|
||||
udpLn.Close()
|
||||
}
|
||||
recordBindAttempt(addr, "udp", udpErr)
|
||||
tcpLn, tcpErr := net.Listen("tcp", addr)
|
||||
if tcpLn != nil {
|
||||
tcpLn.Close()
|
||||
}
|
||||
recordBindAttempt(addr, "tcp", tcpErr)
|
||||
return udpErr == nil && tcpErr == nil
|
||||
}
|
||||
|
||||
@@ -1282,8 +1529,10 @@ func tryUpdateListenerConfigIntercept(cfg *ctrld.Config, notifyFunc func(), fata
|
||||
if hasExplicitConfig {
|
||||
// User specified explicit address — don't guess, just fail
|
||||
if fatal {
|
||||
notifyFunc()
|
||||
mainLog.Load().Fatal().Msgf("DNS intercept: cannot listen on configured address %s", addr)
|
||||
msg := fmt.Sprintf("DNS intercept: cannot listen on configured address %s", addr)
|
||||
mainLog.Load().Error().Msg(msg)
|
||||
failProvision(newProvisionResult(provisionCodeListenerAddrUnavail, msg, bindAttempts, provisionSecrets()...), notifyFunc)
|
||||
return updated, false
|
||||
}
|
||||
return updated, false
|
||||
}
|
||||
@@ -1297,12 +1546,35 @@ func tryUpdateListenerConfigIntercept(cfg *ctrld.Config, notifyFunc func(), fata
|
||||
}
|
||||
|
||||
if fatal {
|
||||
notifyFunc()
|
||||
mainLog.Load().Fatal().Msg("DNS intercept: cannot bind 127.0.0.1:53 or 127.0.0.1:5354")
|
||||
const msg = "DNS intercept: cannot bind 127.0.0.1:53 or 127.0.0.1:5354"
|
||||
mainLog.Load().Error().Msg(msg)
|
||||
failProvision(newProvisionResult(provisionCodeListenerBindFailed, msg, bindAttempts, provisionSecrets()...), notifyFunc)
|
||||
return updated, false
|
||||
}
|
||||
return updated, false
|
||||
}
|
||||
|
||||
func isExplicitInterceptListener(ip string, port int) bool {
|
||||
if ip == "" || ip == "0.0.0.0" || port == 0 {
|
||||
return false
|
||||
}
|
||||
// 127.0.0.1:53 is the default macOS DNS-intercept listener. It can appear
|
||||
// in generated/custom Control D configs, but it should still be allowed to
|
||||
// fall back to 127.0.0.1:5354 when mDNSResponder already owns port 53.
|
||||
return !(ip == "127.0.0.1" && port == 53)
|
||||
}
|
||||
|
||||
// listenerInterceptMode resolves the mode that selects the listener binding
|
||||
// strategy. An explicit "off" is final here, the same as in setDNS. A fallback
|
||||
// to the config value would select the intercept strategy from a stale
|
||||
// persisted mode on the first start after a revert to standard mode.
|
||||
func listenerInterceptMode(cfg *ctrld.Config) string {
|
||||
if interceptMode == "" {
|
||||
return cfg.Service.InterceptMode
|
||||
}
|
||||
return interceptMode
|
||||
}
|
||||
|
||||
// tryUpdateListenerConfig tries updating listener config with a working one.
|
||||
// If fatal is true, and there's listen address conflicted, the function do
|
||||
// fatal error.
|
||||
@@ -1312,13 +1584,9 @@ func tryUpdateListenerConfig(cfg *ctrld.Config, infoLogger *zerolog.Logger, noti
|
||||
// 1. If config has explicit non-default IP:port, use exactly that
|
||||
// 2. Otherwise: try 127.0.0.1:53, then 127.0.0.1:5354, then fatal
|
||||
// This bypasses the full cd-mode listener probing loop entirely.
|
||||
// Check interceptMode (CLI flag) first, then fall back to config value.
|
||||
// dnsIntercept bool is derived later in prog.run(), but we need to know
|
||||
// the intercept mode here to select the right listener probing strategy.
|
||||
im := interceptMode
|
||||
if im == "" || im == "off" {
|
||||
im = cfg.Service.InterceptMode
|
||||
}
|
||||
im := listenerInterceptMode(cfg)
|
||||
if (im == "dns" || im == "hard") && runtime.GOOS == "darwin" {
|
||||
return tryUpdateListenerConfigIntercept(cfg, notifyFunc, fatal)
|
||||
}
|
||||
@@ -1390,6 +1658,15 @@ func tryUpdateListenerConfig(cfg *ctrld.Config, infoLogger *zerolog.Logger, noti
|
||||
_ = closer.Close()
|
||||
}
|
||||
}()
|
||||
// bindAttempts feeds the provisioning result detail. newProvisionResult
|
||||
// caps it, so it grows freely here.
|
||||
var bindAttempts []provisionBindAttempt
|
||||
recordBindAttempt := func(addr, proto string, err error) {
|
||||
if err != nil {
|
||||
bindAttempts = append(bindAttempts, provisionBindAttempt{Addr: addr, Proto: proto, OSError: err.Error()})
|
||||
}
|
||||
}
|
||||
|
||||
// tryListen attempts to listen on given udp and tcp address.
|
||||
// Created listeners will be kept in listeners slice above, and close
|
||||
// before function finished.
|
||||
@@ -1398,16 +1675,21 @@ func tryUpdateListenerConfig(cfg *ctrld.Config, infoLogger *zerolog.Logger, noti
|
||||
if udpLn != nil {
|
||||
closers = append(closers, udpLn)
|
||||
}
|
||||
recordBindAttempt(addr, "udp", udpErr)
|
||||
tcpLn, tcpErr := net.Listen("tcp", addr)
|
||||
if tcpLn != nil {
|
||||
closers = append(closers, tcpLn)
|
||||
}
|
||||
recordBindAttempt(addr, "tcp", tcpErr)
|
||||
return errors.Join(udpErr, tcpErr)
|
||||
}
|
||||
|
||||
listenerMsg := func(listenerNum int, format string, v ...any) string {
|
||||
return fmt.Sprintf("listener.%d %s", listenerNum, fmt.Sprintf(format, v...))
|
||||
}
|
||||
logMsg := func(e *zerolog.Event, listenerNum int, format string, v ...any) {
|
||||
e.MsgFunc(func() string {
|
||||
return fmt.Sprintf("listener.%d %s", listenerNum, fmt.Sprintf(format, v...))
|
||||
return listenerMsg(listenerNum, format, v...)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1459,8 +1741,10 @@ func tryUpdateListenerConfig(cfg *ctrld.Config, infoLogger *zerolog.Logger, noti
|
||||
maxAttempts := 10
|
||||
for {
|
||||
if attempts == maxAttempts {
|
||||
notifyFunc()
|
||||
logMsg(mainLog.Load().Fatal(), n, "could not find available listen ip and port")
|
||||
logMsg(mainLog.Load().Error(), n, "could not find available listen ip and port")
|
||||
msg := listenerMsg(n, "could not find available listen ip and port")
|
||||
failProvision(newProvisionResult(provisionCodeListenerBindFailed, msg, bindAttempts, provisionSecrets()...), notifyFunc)
|
||||
return updated, false
|
||||
}
|
||||
addr := net.JoinHostPort(listener.IP, strconv.Itoa(listener.Port))
|
||||
err := tryListen(addr)
|
||||
@@ -1472,8 +1756,10 @@ func tryUpdateListenerConfig(cfg *ctrld.Config, infoLogger *zerolog.Logger, noti
|
||||
|
||||
if !check.IP && !check.Port {
|
||||
if fatal {
|
||||
notifyFunc()
|
||||
logMsg(mainLog.Load().Fatal(), n, "failed to listen: %v", err)
|
||||
logMsg(mainLog.Load().Error(), n, "failed to listen: %v", err)
|
||||
msg := listenerMsg(n, "failed to listen: %v", err)
|
||||
failProvision(newProvisionResult(provisionCodeListenerAddrUnavail, msg, bindAttempts, provisionSecrets()...), notifyFunc)
|
||||
return updated, false
|
||||
}
|
||||
ok = false
|
||||
break
|
||||
@@ -1540,8 +1826,11 @@ func tryUpdateListenerConfig(cfg *ctrld.Config, infoLogger *zerolog.Logger, noti
|
||||
}
|
||||
if listener.IP == oldIP && listener.Port == oldPort {
|
||||
if fatal {
|
||||
notifyFunc()
|
||||
logMsg(mainLog.Load().Fatal(), n, "could not listen on %s: %v", net.JoinHostPort(listener.IP, strconv.Itoa(listener.Port)), err)
|
||||
triedAddr := net.JoinHostPort(listener.IP, strconv.Itoa(listener.Port))
|
||||
logMsg(mainLog.Load().Error(), n, "could not listen on %s: %v", triedAddr, err)
|
||||
msg := listenerMsg(n, "could not listen on %s: %v", triedAddr, err)
|
||||
failProvision(newProvisionResult(provisionCodeListenerBindFailed, msg, bindAttempts, provisionSecrets()...), notifyFunc)
|
||||
return updated, false
|
||||
}
|
||||
ok = false
|
||||
break
|
||||
@@ -1579,8 +1868,10 @@ func tryUpdateListenerConfig(cfg *ctrld.Config, infoLogger *zerolog.Logger, noti
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
notifyFunc()
|
||||
logMsg(mainLog.Load().Fatal(), n, "could not use %q as DNS nameserver with systemd resolved", listener.IP)
|
||||
logMsg(mainLog.Load().Error(), n, "could not use %q as DNS nameserver with systemd resolved", listener.IP)
|
||||
msg := listenerMsg(n, "could not use %q as DNS nameserver with systemd resolved", listener.IP)
|
||||
failProvision(newProvisionResult(provisionCodeListenerAddrUnavail, msg, bindAttempts, provisionSecrets()...), notifyFunc)
|
||||
return updated, false
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1628,13 +1919,23 @@ func cdUIDFromProvToken() string {
|
||||
Metadata: ctrld.SystemMetadata(context.Background()),
|
||||
}
|
||||
// Process provision token if provided.
|
||||
resolverConfig, err := controld.FetchResolverUID(req, rootCmd.Version, cdDev)
|
||||
resolverConfig, err := fetchResolverUIDFn(context.Background(), req, rootCmd.Version, cdDev)
|
||||
if err != nil {
|
||||
mainLog.Load().Fatal().Err(err).Msgf("failed to fetch resolver uid with provision token: %s", cdOrg)
|
||||
// The token exchange is the first API call of an org/MDM install, so
|
||||
// its failure must carry a code like every other bootstrap failure.
|
||||
code, _ := apiFailureCode(err)
|
||||
mainLog.Load().Error().Msgf("failed to fetch resolver uid with provision token: %s: %s",
|
||||
redactToken(cdOrg), redactSecrets(err.Error(), provisionSecrets()...))
|
||||
failProvision(newProvisionResult(code, fmt.Sprintf("provision token exchange failed: %v", err), nil, provisionSecrets()...), nil)
|
||||
return ""
|
||||
}
|
||||
return resolverConfig.UID
|
||||
}
|
||||
|
||||
// fetchResolverUIDFn is a var so tests can drive token-exchange failures
|
||||
// without reaching the network.
|
||||
var fetchResolverUIDFn = controld.FetchResolverUID
|
||||
|
||||
// removeOrgFlagsFromArgs removes organization flags from command line arguments.
|
||||
// The flags are:
|
||||
//
|
||||
@@ -1809,6 +2110,9 @@ var errInvalidDeactivationPin = errors.New("deactivation pin is invalid")
|
||||
// errRequiredDeactivationPin indicates that the deactivation pin is required but not provided by users.
|
||||
var errRequiredDeactivationPin = errors.New("deactivation pin is required to stop or uninstall the service")
|
||||
|
||||
// errTooManyDeactivationPin represents an error indicating excessive deactivation PIN request attempts.
|
||||
var errTooManyDeactivationPin = errors.New("too many request attempts")
|
||||
|
||||
// checkDeactivationPin validates if the deactivation pin matches one in ControlD config.
|
||||
func checkDeactivationPin(s service.Service, stopCh chan struct{}) error {
|
||||
mainLog.Load().Debug().Msg("Checking deactivation pin")
|
||||
@@ -1837,6 +2141,9 @@ func checkDeactivationPin(s service.Service, stopCh chan struct{}) error {
|
||||
case http.StatusBadRequest:
|
||||
mainLog.Load().Error().Msg(errRequiredDeactivationPin.Error())
|
||||
return errRequiredDeactivationPin // pin is required
|
||||
case http.StatusTooManyRequests:
|
||||
mainLog.Load().Error().Msg(errTooManyDeactivationPin.Error())
|
||||
return errTooManyDeactivationPin
|
||||
case http.StatusOK:
|
||||
return nil // valid pin
|
||||
case http.StatusNotFound:
|
||||
@@ -1849,7 +2156,9 @@ func checkDeactivationPin(s service.Service, stopCh chan struct{}) error {
|
||||
|
||||
// isCheckDeactivationPinErr reports whether there is an error during check deactivation pin process.
|
||||
func isCheckDeactivationPinErr(err error) bool {
|
||||
return errors.Is(err, errInvalidDeactivationPin) || errors.Is(err, errRequiredDeactivationPin)
|
||||
return errors.Is(err, errInvalidDeactivationPin) ||
|
||||
errors.Is(err, errRequiredDeactivationPin) ||
|
||||
errors.Is(err, errTooManyDeactivationPin)
|
||||
}
|
||||
|
||||
// ensureUninstall ensures that s.Uninstall will remove ctrld service from system completely.
|
||||
@@ -1974,7 +2283,7 @@ func doValidateCdRemoteConfig(cdUID string, fatal bool) error {
|
||||
Version: rootCmd.Version,
|
||||
Metadata: ctrld.SystemMetadataRuntime(context.Background()),
|
||||
}
|
||||
rc, err := controld.FetchResolverConfig(req, cdDev)
|
||||
rc, err := controld.FetchResolverConfig(context.Background(), req, cdDev)
|
||||
if err != nil {
|
||||
logger := mainLog.Load().Fatal()
|
||||
if !fatal {
|
||||
@@ -2002,17 +2311,25 @@ func doValidateCdRemoteConfig(cdUID string, fatal bool) error {
|
||||
} else {
|
||||
if errors.As(cfgErr, &viper.ConfigParseError{}) {
|
||||
if configStr, _ := base64.StdEncoding.DecodeString(rc.Ctrld.CustomConfig); len(configStr) > 0 {
|
||||
tmpDir := os.TempDir()
|
||||
tmpConfFile := filepath.Join(tmpDir, "ctrld.toml")
|
||||
errorLogged := false
|
||||
// Write remote config to a temporary file to get details error.
|
||||
if we := os.WriteFile(tmpConfFile, configStr, 0600); we == nil {
|
||||
// Write remote config to a uniquely named temporary file to get detailed error.
|
||||
if tmpFile, tmpErr := os.CreateTemp("", "ctrld-*.toml"); tmpErr == nil {
|
||||
tmpConfFile := tmpFile.Name()
|
||||
if _, err := tmpFile.Write(configStr); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to write temporary config file")
|
||||
}
|
||||
if err := tmpFile.Close(); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to save temporary config file")
|
||||
|
||||
}
|
||||
if de := decoderErrorFromTomlFile(tmpConfFile); de != nil {
|
||||
row, col := de.Position()
|
||||
mainLog.Load().Error().Msgf("failed to parse custom config at line: %d, column: %d, error: %s", row, col, de.Error())
|
||||
errorLogged = true
|
||||
}
|
||||
_ = os.Remove(tmpConfFile)
|
||||
if err := os.Remove(tmpConfFile); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to remove temporary config file")
|
||||
}
|
||||
}
|
||||
// If we could not log details error, emit what we have already got.
|
||||
if !errorLogged {
|
||||
@@ -2030,6 +2347,23 @@ func doValidateCdRemoteConfig(cdUID string, fatal bool) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ensureRunningIfaceForInvalidUninstall populates p.runningIface before the
|
||||
// invalid-device self-uninstall resets DNS. This path can run early during
|
||||
// service startup (e.g. right after a reboot) before the running interface is
|
||||
// otherwise known. resetDNS, via resetDNSForRunningIface, silently skips DNS
|
||||
// restoration when p.runningIface is empty, which would leave the OS pointed at
|
||||
// ctrld's local listener after the service is removed. See issue-556.
|
||||
func ensureRunningIfaceForInvalidUninstall(p *prog, s service.Service) {
|
||||
if iface == "" {
|
||||
iface = autoIface
|
||||
}
|
||||
p.preRun()
|
||||
if ir := runningIface(s); ir != nil {
|
||||
p.runningIface = ir.Name
|
||||
p.requiredMultiNICsConfig = ir.All
|
||||
}
|
||||
}
|
||||
|
||||
// uninstallInvalidCdUID performs self-uninstallation because the ControlD device does not exist.
|
||||
func uninstallInvalidCdUID(p *prog, logger zerolog.Logger, doStop bool) bool {
|
||||
s, err := newService(p, svcConfig)
|
||||
@@ -2037,8 +2371,13 @@ func uninstallInvalidCdUID(p *prog, logger zerolog.Logger, doStop bool) bool {
|
||||
logger.Warn().Err(err).Msg("failed to create new service")
|
||||
return false
|
||||
}
|
||||
ensureRunningIfaceForInvalidUninstall(p, s)
|
||||
// restore static DNS settings or DHCP
|
||||
p.resetDNS(false, true)
|
||||
// The invalid-device path may run early during service startup before runningIface
|
||||
// is known. Restore every saved static DNS file so uninstalling does not leave the
|
||||
// OS pointed at ctrld's local listener after the service is removed.
|
||||
restoreSavedStaticDNS("", true)
|
||||
|
||||
tasks := []task{{s.Uninstall, true, "Uninstall"}}
|
||||
if doTasks(tasks) {
|
||||
@@ -2050,3 +2389,12 @@ func uninstallInvalidCdUID(p *prog, logger zerolog.Logger, doStop bool) bool {
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// redactToken returns the first 4 characters of a token followed by ***,
|
||||
// or just *** if the token is 4 characters or shorter.
|
||||
func redactToken(s string) string {
|
||||
if len(s) <= 4 {
|
||||
return "***"
|
||||
}
|
||||
return s[:4] + "***"
|
||||
}
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
func TestIsExplicitInterceptListener(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
ip string
|
||||
port int
|
||||
want bool
|
||||
}{
|
||||
{name: "empty", ip: "", port: 0, want: false},
|
||||
{name: "wildcard", ip: "0.0.0.0", port: 53, want: false},
|
||||
{name: "zero port", ip: "127.0.0.1", port: 0, want: false},
|
||||
{name: "default intercept listener", ip: "127.0.0.1", port: 53, want: false},
|
||||
{name: "fallback port explicit", ip: "127.0.0.1", port: 5354, want: true},
|
||||
{name: "custom loopback explicit", ip: "127.0.0.2", port: 53, want: true},
|
||||
{name: "custom address explicit", ip: "192.0.2.10", port: 53, want: true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := isExplicitInterceptListener(tt.ip, tt.port); got != tt.want {
|
||||
t.Fatalf("isExplicitInterceptListener(%q, %d) = %v, want %v", tt.ip, tt.port, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestPreserveBoundListeners is a regression test for #551: on reload, the on-disk
|
||||
// generated config still declares 127.0.0.1:53, but the running listener has fallen back
|
||||
// to 127.0.0.1:5354. preserveBoundListeners must keep the in-memory config on the actual
|
||||
// bound port so pf rdr rules and probes do not target the dead default port.
|
||||
func TestPreserveBoundListeners(t *testing.T) {
|
||||
// cur = actual running listener (fell back to 5354); newCfg = freshly read from disk (53).
|
||||
cur := map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.1", Port: 5354}}
|
||||
newListeners := map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.1", Port: 53}}
|
||||
|
||||
preserveBoundListeners(newListeners, cur)
|
||||
|
||||
if got := newListeners["0"].Port; got != 5354 {
|
||||
t.Errorf("listener port after reload = %d, want 5354 (actual bound port)", got)
|
||||
}
|
||||
if got := newListeners["0"].IP; got != "127.0.0.1" {
|
||||
t.Errorf("listener IP after reload = %q, want 127.0.0.1", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPreserveBoundListeners_NoChange verifies that when the on-disk config matches the
|
||||
// running listener, the config is left untouched (a legitimate reload with the same port).
|
||||
func TestPreserveBoundListeners_NoChange(t *testing.T) {
|
||||
cur := map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.1", Port: 5354}}
|
||||
newListeners := map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.1", Port: 5354}}
|
||||
|
||||
preserveBoundListeners(newListeners, cur)
|
||||
|
||||
if got := newListeners["0"].Port; got != 5354 {
|
||||
t.Errorf("listener port = %d, want 5354", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPreserveBoundListeners_MissingCurrent verifies that a listener present on disk but not
|
||||
// in the current running set (e.g. newly added) is left as configured.
|
||||
func TestPreserveBoundListeners_MissingCurrent(t *testing.T) {
|
||||
cur := map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.1", Port: 5354}}
|
||||
newListeners := map[string]*ctrld.ListenerConfig{
|
||||
"0": {IP: "127.0.0.1", Port: 53},
|
||||
"1": {IP: "127.0.0.1", Port: 5355},
|
||||
}
|
||||
|
||||
preserveBoundListeners(newListeners, cur)
|
||||
|
||||
if got := newListeners["0"].Port; got != 5354 {
|
||||
t.Errorf("listener 0 port = %d, want 5354 (preserved)", got)
|
||||
}
|
||||
if got := newListeners["1"].Port; got != 5355 {
|
||||
t.Errorf("listener 1 port = %d, want 5355 (unchanged, no current binding)", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPreserveBoundListeners_ExplicitChangeNotMasked verifies that an explicit, non-default
|
||||
// listener in the reloaded config is applied rather than reverted to the old bound listener.
|
||||
// Reverting an explicit change would make the control-server reload comparison return 200
|
||||
// instead of 201, silently dropping the new listener. Regression guard for #551 review.
|
||||
func TestPreserveBoundListeners_ExplicitChangeNotMasked(t *testing.T) {
|
||||
// Running listener fell back to 5354; user reloads with an explicit new listener.
|
||||
cur := map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.1", Port: 5354}}
|
||||
newListeners := map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.2", Port: 5399}}
|
||||
|
||||
preserveBoundListeners(newListeners, cur)
|
||||
|
||||
if got := newListeners["0"].IP; got != "127.0.0.2" {
|
||||
t.Errorf("explicit listener IP = %q, want 127.0.0.2 (not reverted)", got)
|
||||
}
|
||||
if got := newListeners["0"].Port; got != 5399 {
|
||||
t.Errorf("explicit listener port = %d, want 5399 (not reverted)", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPreserveBoundListeners_ExplicitDefaultPreserved verifies that the default
|
||||
// 127.0.0.1:53 listener remains fallback-eligible: when it diverges from the running
|
||||
// fallback port it is still preserved (isExplicitInterceptListener treats :53 as non-explicit).
|
||||
func TestPreserveBoundListeners_ExplicitDefaultPreserved(t *testing.T) {
|
||||
cur := map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.1", Port: 5354}}
|
||||
newListeners := map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.1", Port: 53}}
|
||||
|
||||
preserveBoundListeners(newListeners, cur)
|
||||
|
||||
if got := newListeners["0"].Port; got != 5354 {
|
||||
t.Errorf("default listener port = %d, want 5354 (preserved fallback)", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,403 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sync/atomic"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
"github.com/Control-D-Inc/ctrld/internal/controld"
|
||||
)
|
||||
|
||||
func TestContextFromStopCh(t *testing.T) {
|
||||
t.Run("cancels when stopCh closes", func(t *testing.T) {
|
||||
stopCh := make(chan struct{})
|
||||
ctx, cancel := contextFromStopCh(stopCh)
|
||||
defer cancel()
|
||||
|
||||
if ctx.Err() != nil {
|
||||
t.Fatalf("context cancelled before the stop request: %v", ctx.Err())
|
||||
}
|
||||
close(stopCh)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("context was not cancelled after stopCh closed")
|
||||
}
|
||||
if !errors.Is(ctx.Err(), context.Canceled) {
|
||||
t.Errorf("ctx.Err() = %v, want %v", ctx.Err(), context.Canceled)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("cancel releases the watcher", func(t *testing.T) {
|
||||
// stopCh is never closed: cancel() must still end the goroutine watching it.
|
||||
ctx, cancel := contextFromStopCh(make(chan struct{}))
|
||||
cancel()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("context was not cancelled by cancel()")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("nil stopCh is usable", func(t *testing.T) {
|
||||
// Mobile callers have no stop channel; preflight must still run.
|
||||
ctx, cancel := contextFromStopCh(nil)
|
||||
defer cancel()
|
||||
if ctx.Err() != nil {
|
||||
t.Fatalf("context cancelled immediately: %v", ctx.Err())
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// retryableNetworkErr is the shape processCDFlags treats as "retry with bootstrap
|
||||
// DNS": a url.Error wrapping a network failure.
|
||||
func retryableNetworkErr() error {
|
||||
return &url.Error{
|
||||
Op: "Post",
|
||||
URL: "https://api.controld.com/utility",
|
||||
Err: &net.OpError{Op: "dial", Net: "tcp", Err: syscall.ECONNREFUSED},
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessCDFlagsStopsWhenCancelled(t *testing.T) {
|
||||
oldFetch := fetchResolverConfig
|
||||
oldUID := cdUID
|
||||
t.Cleanup(func() {
|
||||
fetchResolverConfig = oldFetch
|
||||
cdUID = oldUID
|
||||
})
|
||||
cdUID = "testuid"
|
||||
|
||||
var calls atomic.Int64
|
||||
fetchResolverConfig = func(ctx context.Context, req *controld.ResolverConfigRequest, dev bool) (*controld.ResolverConfig, error) {
|
||||
calls.Add(1)
|
||||
return nil, retryableNetworkErr()
|
||||
}
|
||||
|
||||
// A stop request arriving while the API is unreachable. Before this was
|
||||
// cancellable, the retry loop kept running after the service reported itself
|
||||
// stopped, which is what kept the incident's process alive and enforcing.
|
||||
stopCh := make(chan struct{})
|
||||
ctx, cancel := contextFromStopCh(stopCh)
|
||||
defer cancel()
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
cfg := ctrld.Config{}
|
||||
_, err := processCDFlags(ctx, &cfg)
|
||||
done <- err
|
||||
}()
|
||||
|
||||
// Let it fail at least once and settle into backoff before stopping.
|
||||
deadline := time.After(10 * time.Second)
|
||||
for calls.Load() == 0 {
|
||||
select {
|
||||
case <-deadline:
|
||||
t.Fatal("resolver config was never fetched")
|
||||
case err := <-done:
|
||||
t.Fatalf("processCDFlags returned before any fetch: %v", err)
|
||||
default:
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
}
|
||||
close(stopCh)
|
||||
|
||||
select {
|
||||
case err := <-done:
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Errorf("processCDFlags err = %v, want it to report %v", err, context.Canceled)
|
||||
}
|
||||
case <-time.After(30 * time.Second):
|
||||
t.Fatal("processCDFlags did not return after the stop request")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessCDFlagsReturnsImmediatelyWhenAlreadyCancelled(t *testing.T) {
|
||||
oldFetch := fetchResolverConfig
|
||||
oldUID := cdUID
|
||||
t.Cleanup(func() {
|
||||
fetchResolverConfig = oldFetch
|
||||
cdUID = oldUID
|
||||
})
|
||||
cdUID = "testuid"
|
||||
|
||||
var calls atomic.Int64
|
||||
fetchResolverConfig = func(ctx context.Context, req *controld.ResolverConfigRequest, dev bool) (*controld.ResolverConfig, error) {
|
||||
calls.Add(1)
|
||||
return nil, retryableNetworkErr()
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
cfg := ctrld.Config{}
|
||||
_, err := processCDFlags(ctx, &cfg)
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Errorf("processCDFlags err = %v, want %v", err, context.Canceled)
|
||||
}
|
||||
// One attempt is made before the loop notices; it must not retry past that.
|
||||
if got := calls.Load(); got > 1 {
|
||||
t.Errorf("fetched %d times with a cancelled context, want at most 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestRunAPIPreflightClassification is the regression guard for classifying a preflight
|
||||
// failure as an operator stop.
|
||||
//
|
||||
// runAPIPreflight cancels the context it derived from stopCh. Sampling the stop state
|
||||
// from that context afterwards reports "stopped" unconditionally, because
|
||||
// context.CancelFunc sets ctx.Err() whether or not anyone asked to stop. run() then
|
||||
// takes the stop branch for every failure, which skips self-uninstalling a deleted
|
||||
// device, skips the mobile exit callback, and tells the service manager a failed start
|
||||
// was a clean exit.
|
||||
func TestRunAPIPreflightClassification(t *testing.T) {
|
||||
oldFetch := fetchResolverConfig
|
||||
oldUID := cdUID
|
||||
t.Cleanup(func() {
|
||||
fetchResolverConfig = oldFetch
|
||||
cdUID = oldUID
|
||||
})
|
||||
cdUID = "testuid"
|
||||
|
||||
// A deleted ControlD device: non-retryable, so preflight returns promptly.
|
||||
deletedDevice := func() error {
|
||||
e := &controld.ErrorResponse{}
|
||||
e.ErrorField.Code = controld.InvalidConfigCode
|
||||
e.ErrorField.Message = "device does not exist"
|
||||
return e
|
||||
}
|
||||
|
||||
openCh := make(chan struct{})
|
||||
closedCh := make(chan struct{})
|
||||
close(closedCh)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
stopCh <-chan struct{}
|
||||
fetchErr func() error
|
||||
wantStop bool
|
||||
}{
|
||||
{
|
||||
// The P1: no stop was requested, so this must reach the failure branch.
|
||||
name: "api error with no stop request",
|
||||
stopCh: openCh,
|
||||
fetchErr: deletedDevice,
|
||||
},
|
||||
{
|
||||
// Mobile passes no stop channel at all, so it could never have stopped.
|
||||
name: "api error with a nil stop channel",
|
||||
stopCh: nil,
|
||||
fetchErr: deletedDevice,
|
||||
},
|
||||
{
|
||||
name: "stop requested during preflight",
|
||||
stopCh: closedCh,
|
||||
fetchErr: func() error { return retryableNetworkErr() },
|
||||
wantStop: true,
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
fetchResolverConfig = func(context.Context, *controld.ResolverConfigRequest, bool) (*controld.ResolverConfig, error) {
|
||||
return nil, tc.fetchErr()
|
||||
}
|
||||
cfg := ctrld.Config{}
|
||||
pf := runAPIPreflight(tc.stopCh, &cfg)
|
||||
|
||||
if pf.err == nil {
|
||||
t.Fatal("expected preflight to fail")
|
||||
}
|
||||
if pf.stopRequested != tc.wantStop {
|
||||
t.Errorf("stopRequested = %v, want %v", pf.stopRequested, tc.wantStop)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestRunAPIPreflightPreservesAPIError verifies the error reaches the caller in a form
|
||||
// the failure branch can still act on: self-uninstall keys off an *ErrorResponse with
|
||||
// InvalidConfigCode, and it only runs if that error is both classified as a failure and
|
||||
// still unwrappable.
|
||||
func TestRunAPIPreflightPreservesAPIError(t *testing.T) {
|
||||
oldFetch := fetchResolverConfig
|
||||
oldUID := cdUID
|
||||
t.Cleanup(func() {
|
||||
fetchResolverConfig = oldFetch
|
||||
cdUID = oldUID
|
||||
})
|
||||
cdUID = "testuid"
|
||||
|
||||
want := &controld.ErrorResponse{}
|
||||
want.ErrorField.Code = controld.InvalidConfigCode
|
||||
fetchResolverConfig = func(context.Context, *controld.ResolverConfigRequest, bool) (*controld.ResolverConfig, error) {
|
||||
return nil, want
|
||||
}
|
||||
|
||||
cfg := ctrld.Config{}
|
||||
pf := runAPIPreflight(make(chan struct{}), &cfg)
|
||||
|
||||
if pf.stopRequested {
|
||||
t.Error("a device-deleted failure must not be reported as an operator stop")
|
||||
}
|
||||
var got *controld.ErrorResponse
|
||||
if !errors.As(pf.err, &got) {
|
||||
t.Fatalf("error no longer unwraps to *controld.ErrorResponse: %v", pf.err)
|
||||
}
|
||||
if got.ErrorField.Code != controld.InvalidConfigCode {
|
||||
t.Errorf("code = %d, want %d (self-uninstall would not trigger)", got.ErrorField.Code, controld.InvalidConfigCode)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPermanentAPIRejectionNarrowsToClientErrors is the regression guard for the clean
|
||||
// exit added above.
|
||||
//
|
||||
// controld builds an *ErrorResponse for any non-200 whose body decodes, so the Go type
|
||||
// says nothing about whether the API's answer will change on a retry. Keying the clean
|
||||
// exit off the type alone meant a 502 from a load balancer, or an API having a bad ten
|
||||
// minutes, stopped ctrld on every affected host with no service-manager retry behind it -
|
||||
// worse than the abnormal exit it replaced, because a Fatal at least gets restarted.
|
||||
//
|
||||
// Only a client-error status may take that path.
|
||||
func TestPermanentAPIRejectionNarrowsToClientErrors(t *testing.T) {
|
||||
rejection := func(status, code int) error {
|
||||
e := &controld.ErrorResponse{StatusCode: status}
|
||||
e.ErrorField.Code = code
|
||||
e.ErrorField.Message = "api said no"
|
||||
return e
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
err error
|
||||
wantPermanent bool
|
||||
}{
|
||||
{
|
||||
// The case the clean exit exists for: the device is gone, and every restart
|
||||
// will be told the same thing.
|
||||
name: "deleted device",
|
||||
err: rejection(http.StatusNotFound, controld.InvalidConfigCode),
|
||||
wantPermanent: true,
|
||||
},
|
||||
{"revoked credentials", rejection(http.StatusUnauthorized, 0), true},
|
||||
{"forbidden", rejection(http.StatusForbidden, 0), true},
|
||||
{"malformed request", rejection(http.StatusBadRequest, 0), true},
|
||||
|
||||
// Server-side trouble. These must keep the abnormal exit so the service
|
||||
// manager's recovery policy retries.
|
||||
{"bad gateway", rejection(http.StatusBadGateway, 0), false},
|
||||
{"internal error", rejection(http.StatusInternalServerError, 0), false},
|
||||
{"service unavailable", rejection(http.StatusServiceUnavailable, 0), false},
|
||||
|
||||
// 4xx, but both are the API asking for a later attempt rather than refusing
|
||||
// this configuration.
|
||||
{"request timeout", rejection(http.StatusRequestTimeout, 0), false},
|
||||
{"rate limited", rejection(http.StatusTooManyRequests, 0), false},
|
||||
|
||||
// An *ErrorResponse built without a recorded status carries no verdict. A
|
||||
// hand-constructed one, or a decode path that forgets to record the status,
|
||||
// must not silently gain the clean exit.
|
||||
{"no recorded status", rejection(0, controld.InvalidConfigCode), false},
|
||||
|
||||
// Not an API answer at all: the incident's denied socket reaches Fatal.
|
||||
{"network failure", retryableNetworkErr(), false},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, ok := permanentAPIRejection(tc.err)
|
||||
if ok != tc.wantPermanent {
|
||||
t.Errorf("permanentAPIRejection() = %v, want %v", ok, tc.wantPermanent)
|
||||
}
|
||||
if ok && got == nil {
|
||||
t.Error("a permanent rejection must return the rejection for reporting")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// The wrapped form matters too: preflight composes the fetch error, and errors.As has
|
||||
// to reach through that for either branch to be chosen correctly.
|
||||
wrapped := fmt.Errorf("processCDFlags: %w", rejection(http.StatusNotFound, controld.InvalidConfigCode))
|
||||
if _, ok := permanentAPIRejection(wrapped); !ok {
|
||||
t.Error("a wrapped API rejection must still be recognised")
|
||||
}
|
||||
wrappedTransient := fmt.Errorf("processCDFlags: %w", rejection(http.StatusBadGateway, 0))
|
||||
if _, ok := permanentAPIRejection(wrappedTransient); ok {
|
||||
t.Error("a wrapped 502 must not be treated as a permanent rejection")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStopRequested(t *testing.T) {
|
||||
closedCh := make(chan struct{})
|
||||
close(closedCh)
|
||||
|
||||
if stopRequested(nil) {
|
||||
t.Error("a nil stop channel must read as no stop (mobile passes none)")
|
||||
}
|
||||
if stopRequested(make(chan struct{})) {
|
||||
t.Error("an open stop channel must read as no stop")
|
||||
}
|
||||
if !stopRequested(closedCh) {
|
||||
t.Error("a closed stop channel must read as a stop")
|
||||
}
|
||||
}
|
||||
|
||||
// TestReloadFetchIsBoundedByServiceLifetime covers the reload path's stop wiring.
|
||||
//
|
||||
// Reload fetches the ControlD config too, and it used to build the bounded context
|
||||
// itself. Nothing tested that: the wrong channel, or a dropped cancel, would have left a
|
||||
// reload retrying against an unreachable API after "service stopped" was logged, and no
|
||||
// test would have failed. Both paths now go through one bounded fetch, so this pins it.
|
||||
func TestReloadFetchIsBoundedByServiceLifetime(t *testing.T) {
|
||||
original := processCDFlagsFn
|
||||
t.Cleanup(func() { processCDFlagsFn = original })
|
||||
|
||||
t.Run("a stop request cancels the reload fetch", func(t *testing.T) {
|
||||
stopCh := make(chan struct{})
|
||||
close(stopCh)
|
||||
|
||||
var sawCancelled bool
|
||||
processCDFlagsFn = func(ctx context.Context, _ *ctrld.Config) (*controld.ResolverConfig, error) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
sawCancelled = true
|
||||
case <-time.After(2 * time.Second):
|
||||
}
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
|
||||
p := &prog{stopCh: stopCh}
|
||||
if _, err := p.fetchCDConfigBoundedByLifetime(&ctrld.Config{}); !errors.Is(err, context.Canceled) {
|
||||
t.Errorf("reload fetch err = %v, want %v", err, context.Canceled)
|
||||
}
|
||||
if !sawCancelled {
|
||||
t.Error("the reload fetch did not observe the stop request: it is not bound to the service lifetime")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("the derived context is always released", func(t *testing.T) {
|
||||
// stopCh stays open: the fetch's own cancel is what must end the watcher, or
|
||||
// every reload leaks a goroutine.
|
||||
var captured context.Context
|
||||
processCDFlagsFn = func(ctx context.Context, _ *ctrld.Config) (*controld.ResolverConfig, error) {
|
||||
captured = ctx
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
p := &prog{stopCh: make(chan struct{})}
|
||||
if _, err := p.fetchCDConfigBoundedByLifetime(&ctrld.Config{}); err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
select {
|
||||
case <-captured.Done():
|
||||
case <-time.After(time.Second):
|
||||
t.Error("the reload fetch left its context uncancelled")
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,329 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/rs/zerolog"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
"github.com/Control-D-Inc/ctrld/internal/controld"
|
||||
)
|
||||
|
||||
// TestApiFailureCode covers the preflight-error mapping: a deleted device
|
||||
// gets its own code (it drives self-uninstall), other permanent rejections
|
||||
// are generic, anything else is retryable reachability trouble.
|
||||
func TestApiFailureCode(t *testing.T) {
|
||||
rejection := func(status, code int) error {
|
||||
e := &controld.ErrorResponse{StatusCode: status}
|
||||
e.ErrorField.Code = code
|
||||
e.ErrorField.Message = "api said no"
|
||||
return e
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
err error
|
||||
wantCode provisionFailureCode
|
||||
wantOk bool
|
||||
}{
|
||||
{name: "nil error", err: nil, wantCode: "", wantOk: false},
|
||||
{
|
||||
name: "deleted device maps to device invalid",
|
||||
err: rejection(http.StatusNotFound, controld.InvalidConfigCode),
|
||||
wantCode: provisionCodeAPIDeviceInvalid,
|
||||
wantOk: true,
|
||||
},
|
||||
{
|
||||
name: "revoked credentials map to rejected",
|
||||
err: rejection(http.StatusUnauthorized, 0),
|
||||
wantCode: provisionCodeAPIRejected,
|
||||
wantOk: true,
|
||||
},
|
||||
{
|
||||
name: "server error maps to unreachable",
|
||||
err: rejection(http.StatusBadGateway, 0),
|
||||
wantCode: provisionCodeAPIUnreachable,
|
||||
wantOk: true,
|
||||
},
|
||||
{
|
||||
name: "network failure maps to unreachable",
|
||||
err: retryableNetworkErr(),
|
||||
wantCode: provisionCodeAPIUnreachable,
|
||||
wantOk: true,
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
code, ok := apiFailureCode(tc.err)
|
||||
if ok != tc.wantOk {
|
||||
t.Fatalf("apiFailureCode() ok = %v, want %v", ok, tc.wantOk)
|
||||
}
|
||||
if code != tc.wantCode {
|
||||
t.Errorf("apiFailureCode() code = %s, want %s", code, tc.wantCode)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func stubProvisionGlobals(t *testing.T) (exitCode *int, notified *bool) {
|
||||
t.Helper()
|
||||
oldCdUID, oldCdOrg := cdUID, cdOrg
|
||||
oldExit, oldUninstall := provisionExit, uninstallInvalidCdUIDFn
|
||||
t.Cleanup(func() {
|
||||
cdUID, cdOrg = oldCdUID, oldCdOrg
|
||||
provisionExit, uninstallInvalidCdUIDFn = oldExit, oldUninstall
|
||||
})
|
||||
overrideProvisionResultPath(t)
|
||||
code := -1
|
||||
provisionExit = func(c int) { code = c }
|
||||
n := false
|
||||
return &code, &n
|
||||
}
|
||||
|
||||
func TestHandleAPIPreflightFailure(t *testing.T) {
|
||||
deviceInvalid := func() error {
|
||||
e := &controld.ErrorResponse{StatusCode: http.StatusNotFound}
|
||||
e.ErrorField.Code = controld.InvalidConfigCode
|
||||
e.ErrorField.Message = "device does not exist"
|
||||
return e
|
||||
}
|
||||
rejected := func() error {
|
||||
e := &controld.ErrorResponse{StatusCode: http.StatusUnauthorized}
|
||||
e.ErrorField.Message = "bad token"
|
||||
return e
|
||||
}
|
||||
|
||||
t.Run("permanent rejection returns cleanly", func(t *testing.T) {
|
||||
exitCode, notified := stubProvisionGlobals(t)
|
||||
handleAPIPreflightFailure(&prog{}, rejected(), func() { *notified = true })
|
||||
if *exitCode != -1 {
|
||||
t.Errorf("provisionExit called with %d, want a clean return", *exitCode)
|
||||
}
|
||||
if !*notified {
|
||||
t.Error("notify not called")
|
||||
}
|
||||
r, err := readProvisionResult()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if r.Code != string(provisionCodeAPIRejected) {
|
||||
t.Errorf("code = %q, want API_REJECTED", r.Code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("deleted device self-uninstalls and returns cleanly", func(t *testing.T) {
|
||||
exitCode, notified := stubProvisionGlobals(t)
|
||||
uninstalled := false
|
||||
uninstallInvalidCdUIDFn = func(_ *prog, _ zerolog.Logger, _ bool) bool {
|
||||
uninstalled = true
|
||||
return true
|
||||
}
|
||||
handleAPIPreflightFailure(&prog{}, deviceInvalid(), func() { *notified = true })
|
||||
if *exitCode != -1 {
|
||||
t.Errorf("provisionExit called with %d, want a clean return", *exitCode)
|
||||
}
|
||||
if !uninstalled {
|
||||
t.Error("self-uninstall not attempted")
|
||||
}
|
||||
if !*notified {
|
||||
t.Error("notify not called")
|
||||
}
|
||||
r, err := readProvisionResult()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if r.Code != string(provisionCodeAPIDeviceInvalid) {
|
||||
t.Errorf("code = %q, want API_DEVICE_INVALID", r.Code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unreachable exits nonzero", func(t *testing.T) {
|
||||
exitCode, notified := stubProvisionGlobals(t)
|
||||
handleAPIPreflightFailure(&prog{}, retryableNetworkErr(), func() { *notified = true })
|
||||
if *exitCode != provisionExitCodeForCode[provisionCodeAPIUnreachable] {
|
||||
t.Errorf("exit = %d, want %d", *exitCode, provisionExitCodeForCode[provisionCodeAPIUnreachable])
|
||||
}
|
||||
if !*notified {
|
||||
t.Error("notify not called")
|
||||
}
|
||||
r, err := readProvisionResult()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if r.Code != string(provisionCodeAPIUnreachable) {
|
||||
t.Errorf("code = %q, want API_UNREACHABLE", r.Code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("bare uid from a composite --cd value is redacted", func(t *testing.T) {
|
||||
_, _ = stubProvisionGlobals(t)
|
||||
cdUID = "deviceabc/clientxyz"
|
||||
cdOrg = ""
|
||||
err := fmt.Errorf("failed: api says deviceabc is unknown")
|
||||
handleAPIPreflightFailure(&prog{}, err, func() {})
|
||||
r, rerr := readProvisionResult()
|
||||
if rerr != nil {
|
||||
t.Fatal(rerr)
|
||||
}
|
||||
if strings.Contains(r.Message, "deviceabc") {
|
||||
t.Errorf("bare uid leaked into message: %q", r.Message)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestCdUIDFromProvTokenFailureEmitsCode(t *testing.T) {
|
||||
exitCode, _ := stubProvisionGlobals(t)
|
||||
oldFetch, oldHostname := fetchResolverUIDFn, customHostname
|
||||
t.Cleanup(func() { fetchResolverUIDFn, customHostname = oldFetch, oldHostname })
|
||||
cdUID = ""
|
||||
cdOrg = "org-secret-token-123"
|
||||
customHostname = ""
|
||||
|
||||
rejected := &controld.ErrorResponse{StatusCode: http.StatusUnauthorized}
|
||||
rejected.ErrorField.Message = "bad provision token org-secret-token-123"
|
||||
fetchResolverUIDFn = func(context.Context, *controld.UtilityOrgRequest, string, bool) (*controld.ResolverConfig, error) {
|
||||
return nil, rejected
|
||||
}
|
||||
|
||||
if got := cdUIDFromProvToken(); got != "" {
|
||||
t.Errorf("cdUIDFromProvToken() = %q, want empty on failure", got)
|
||||
}
|
||||
if *exitCode != provisionExitCodeForCode[provisionCodeAPIRejected] {
|
||||
t.Errorf("exit = %d, want API_REJECTED exit %d", *exitCode, provisionExitCodeForCode[provisionCodeAPIRejected])
|
||||
}
|
||||
r, err := readProvisionResult()
|
||||
if err != nil {
|
||||
t.Fatalf("no provision result written: %v", err)
|
||||
}
|
||||
if r.Code != string(provisionCodeAPIRejected) {
|
||||
t.Errorf("code = %q, want API_REJECTED", r.Code)
|
||||
}
|
||||
if strings.Contains(r.Message, cdOrg) {
|
||||
t.Errorf("token leaked into result message: %q", r.Message)
|
||||
}
|
||||
}
|
||||
|
||||
// Regression test: an explicit ip:port that fails to bind used to die with a
|
||||
// bare fatal log automation could not tell apart from any other crash. It
|
||||
// must report a stable code through the provisioning result instead.
|
||||
func TestTryUpdateListenerConfigConfiguredAddrUnavailable(t *testing.T) {
|
||||
// Occupy one localhost port on both udp and tcp, and hold both for the
|
||||
// whole test so ctrld's own bind attempt is guaranteed to fail.
|
||||
udpConn, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("could not reserve a udp port: %v", err)
|
||||
}
|
||||
defer udpConn.Close()
|
||||
|
||||
host, portStr, err := net.SplitHostPort(udpConn.LocalAddr().String())
|
||||
if err != nil {
|
||||
t.Fatalf("could not parse reserved address: %v", err)
|
||||
}
|
||||
port, err := strconv.Atoi(portStr)
|
||||
if err != nil {
|
||||
t.Fatalf("could not parse reserved port: %v", err)
|
||||
}
|
||||
|
||||
tcpLn, err := net.Listen("tcp", net.JoinHostPort(host, portStr))
|
||||
if err != nil {
|
||||
t.Fatalf("could not reserve the same port on tcp: %v", err)
|
||||
}
|
||||
defer tcpLn.Close()
|
||||
|
||||
oldCdUID, oldCdOrg, oldNextdns, oldIntercept := cdUID, cdOrg, nextdns, interceptMode
|
||||
oldPath, oldExit := provisionResultPath, provisionExit
|
||||
t.Cleanup(func() {
|
||||
cdUID, cdOrg, nextdns, interceptMode = oldCdUID, oldCdOrg, oldNextdns, oldIntercept
|
||||
provisionResultPath, provisionExit = oldPath, oldExit
|
||||
})
|
||||
// Non-cd, non-nextdns mode with an explicit ip:port: no fallback checks,
|
||||
// the path that used to reach the fatal exit directly.
|
||||
cdUID = ""
|
||||
cdOrg = ""
|
||||
nextdns = ""
|
||||
interceptMode = ""
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
provisionResultPath = func() string { return filepath.Join(tmpDir, "provision_result.json") }
|
||||
|
||||
var exitCode int
|
||||
var exited bool
|
||||
provisionExit = func(code int) { exitCode = code; exited = true }
|
||||
|
||||
cfg := &ctrld.Config{
|
||||
Listener: map[string]*ctrld.ListenerConfig{
|
||||
"0": {IP: host, Port: port},
|
||||
},
|
||||
}
|
||||
|
||||
notified := false
|
||||
_, ok := tryUpdateListenerConfig(cfg, nil, func() { notified = true }, true)
|
||||
|
||||
if ok {
|
||||
t.Error("tryUpdateListenerConfig ok = true, want false")
|
||||
}
|
||||
if !notified {
|
||||
t.Error("expected notifyFunc to run before the recorded exit")
|
||||
}
|
||||
if !exited {
|
||||
t.Fatal("expected provisionExit to be called")
|
||||
}
|
||||
if exitCode != 42 {
|
||||
t.Errorf("exit code = %d, want 42 (LISTENER_CONFIGURED_ADDR_UNAVAILABLE)", exitCode)
|
||||
}
|
||||
|
||||
result, err := readProvisionResult()
|
||||
if err != nil {
|
||||
t.Fatalf("could not read provision result: %v", err)
|
||||
}
|
||||
if result.Code != string(provisionCodeListenerAddrUnavail) {
|
||||
t.Errorf("result code = %s, want %s", result.Code, provisionCodeListenerAddrUnavail)
|
||||
}
|
||||
if result.Stage != string(provisionStageListener) {
|
||||
t.Errorf("result stage = %s, want %s", result.Stage, provisionStageListener)
|
||||
}
|
||||
if result.ExitCode != 42 {
|
||||
t.Errorf("result exit code = %d, want 42", result.ExitCode)
|
||||
}
|
||||
if result.Detail == nil || len(result.Detail.Attempts) == 0 {
|
||||
t.Fatal("expected the occupied address to appear as a recorded bind attempt")
|
||||
}
|
||||
|
||||
occupiedAddr := net.JoinHostPort(host, portStr)
|
||||
// Windows words WSAEADDRINUSE differently, so only require the canonical
|
||||
// message on platforms that produce it.
|
||||
requireInUseText := runtime.GOOS != "windows"
|
||||
var sawUDP, sawTCP bool
|
||||
for _, a := range result.Detail.Attempts {
|
||||
if a.Addr != occupiedAddr || a.OSError == "" {
|
||||
continue
|
||||
}
|
||||
if requireInUseText && !strings.Contains(strings.ToLower(a.OSError), "address already in use") {
|
||||
continue
|
||||
}
|
||||
switch a.Proto {
|
||||
case "udp":
|
||||
sawUDP = true
|
||||
case "tcp":
|
||||
sawTCP = true
|
||||
}
|
||||
}
|
||||
if !sawUDP {
|
||||
t.Error("expected a udp attempt on the occupied address with a bind error")
|
||||
}
|
||||
if !sawTCP {
|
||||
t.Error("expected a tcp attempt on the occupied address with a bind error")
|
||||
}
|
||||
}
|
||||
|
||||
// The exhaustion path (exit 41) is not covered: forcing every fallback,
|
||||
// including a freshly randomized ip/port, to fail has no deterministic seam,
|
||||
// so a test would race whatever ports are free on the host.
|
||||
+200
-97
@@ -275,7 +275,7 @@ func initRunCmd() *cobra.Command {
|
||||
_ = runCmd.Flags().MarkHidden("iface")
|
||||
runCmd.Flags().StringVarP(&cdUpstreamProto, "proto", "", ctrld.ResolverTypeDOH, `Control D upstream type, either "doh" or "doh3"`)
|
||||
runCmd.Flags().BoolVarP(&rfc1918, "rfc1918", "", false, "Listen on RFC1918 addresses when 127.0.0.1 is the only listener")
|
||||
runCmd.Flags().StringVarP(&interceptMode, "intercept-mode", "", "", "OS-level DNS interception mode: 'dns' (with VPN split routing) or 'hard' (all DNS through ctrld, no VPN split routing)")
|
||||
runCmd.Flags().StringVarP(&interceptMode, "intercept-mode", "", "", "OS-level DNS interception mode: 'off' (disable interception and clear a persisted intercept_mode), 'dns' (with VPN split routing), or 'hard' (all DNS through ctrld, no VPN split routing)")
|
||||
|
||||
runCmd.FParseErrWhitelist = cobra.FParseErrWhitelist{UnknownFlags: true}
|
||||
rootCmd.AddCommand(runCmd)
|
||||
@@ -283,6 +283,46 @@ func initRunCmd() *cobra.Command {
|
||||
return runCmd
|
||||
}
|
||||
|
||||
// serviceStageFailureCode maps an aborted service-manager task to its
|
||||
// provisioning code. Other abortOnError tasks (like config validation) keep
|
||||
// their own error paths.
|
||||
func serviceStageFailureCode(taskName string) (provisionFailureCode, bool) {
|
||||
switch taskName {
|
||||
case "Install":
|
||||
return provisionCodeServiceInstall, true
|
||||
case "Start":
|
||||
return provisionCodeServiceStartFailed, true
|
||||
default:
|
||||
return "", false
|
||||
}
|
||||
}
|
||||
|
||||
// serviceTaskErrorSummary describes which service-manager task failed and why,
|
||||
// for use as a provisioning result message.
|
||||
func serviceTaskErrorSummary(taskName string, err error) string {
|
||||
return fmt.Sprintf("%s failed: %v", taskName, err)
|
||||
}
|
||||
|
||||
// resultStalenessTolerance absorbs clock granularity between "ctrld start"
|
||||
// recording its start time and the daemon writing its result file.
|
||||
const resultStalenessTolerance = 2 * time.Second
|
||||
|
||||
// reportStartFailure reports why "ctrld start" failed after install/start
|
||||
// looked fine. A result file the daemon wrote during this attempt names the
|
||||
// failure better than a generic self-check code, so it wins.
|
||||
func reportStartFailure(startedAt time.Time, fallbackMsg string) {
|
||||
if r, err := readProvisionResult(); err == nil && provisionResultTrusted(r) {
|
||||
if ts, err := time.Parse(time.RFC3339, r.Timestamp); err == nil {
|
||||
if !ts.Before(startedAt.Add(-resultStalenessTolerance)) {
|
||||
mainLog.Load().Error().Msg(r.failureLine())
|
||||
provisionExit(r.ExitCode)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
failProvision(newProvisionResult(provisionCodeServiceSelfCheck, fallbackMsg, nil, provisionSecrets()...), nil)
|
||||
}
|
||||
|
||||
func initStartCmd() *cobra.Command {
|
||||
startCmd := &cobra.Command{
|
||||
PreRun: func(cmd *cobra.Command, args []string) {
|
||||
@@ -354,21 +394,22 @@ NOTE: running "ctrld start" without any arguments will start already installed c
|
||||
svcExists := serviceConfigFileExists()
|
||||
mainLog.Load().Debug().Msgf("intercept upgrade check: args=%v interceptOnly=%v svcConfigExists=%v interceptMode=%q", osArgsEarly, interceptOnly, svcExists, interceptMode)
|
||||
if interceptOnly && svcExists {
|
||||
// Remove any existing intercept flags before applying the new value.
|
||||
_ = removeServiceFlag("--intercept-mode")
|
||||
// An explicit "off" argument must override a previously persisted config
|
||||
// value while the service clears that value on startup.
|
||||
if err := removeServiceFlag("--intercept-mode"); err != nil {
|
||||
mainLog.Load().Fatal().Err(err).Msg("failed to remove existing intercept mode from service arguments")
|
||||
}
|
||||
|
||||
if interceptMode == "off" {
|
||||
// "off" = remove intercept mode entirely (just the removal above).
|
||||
mainLog.Load().Notice().Msg("Existing service detected — removing --intercept-mode from service arguments")
|
||||
mainLog.Load().Notice().Msg("Existing service detected — disabling intercept mode")
|
||||
} else {
|
||||
// Add the new mode value.
|
||||
mainLog.Load().Notice().Msgf("Existing service detected — appending --intercept-mode %s to service arguments", interceptMode)
|
||||
if err := appendServiceFlag("--intercept-mode"); err != nil {
|
||||
mainLog.Load().Fatal().Err(err).Msg("failed to append intercept flag to service arguments")
|
||||
}
|
||||
if err := appendServiceFlag(interceptMode); err != nil {
|
||||
mainLog.Load().Fatal().Err(err).Msg("failed to append intercept mode value to service arguments")
|
||||
}
|
||||
}
|
||||
if err := appendServiceFlag("--intercept-mode"); err != nil {
|
||||
mainLog.Load().Fatal().Err(err).Msg("failed to append intercept flag to service arguments")
|
||||
}
|
||||
if err := appendServiceFlag(interceptMode); err != nil {
|
||||
mainLog.Load().Fatal().Err(err).Msg("failed to append intercept mode value to service arguments")
|
||||
}
|
||||
|
||||
// Stop the service if running (bypasses ctrld pin — this is an
|
||||
@@ -399,7 +440,7 @@ NOTE: running "ctrld start" without any arguments will start already installed c
|
||||
reportSetDnsOk := func(sockDir string) {
|
||||
if cc := newSocketControlClient(ctx, s, sockDir); cc != nil {
|
||||
if resp, _ := cc.post(ifacePath, nil); resp != nil && resp.StatusCode == http.StatusOK {
|
||||
if iface == "auto" {
|
||||
if iface == autoIface {
|
||||
iface = defaultIfaceName()
|
||||
}
|
||||
res := &ifaceResponse{}
|
||||
@@ -524,23 +565,50 @@ NOTE: running "ctrld start" without any arguments will start already installed c
|
||||
{s.Start, true, "Start"},
|
||||
{noticeWritingControlDConfig, false, "Notice writing ControlD config"},
|
||||
}
|
||||
// Any result found later must come from this attempt, not a stale run.
|
||||
clearProvisionResult()
|
||||
startAttemptAt := time.Now()
|
||||
mainLog.Load().Notice().Msg("Starting existing ctrld service")
|
||||
if doTasks(tasks) {
|
||||
mainLog.Load().Notice().Msg("Service started")
|
||||
sockDir, err := socketDir()
|
||||
if err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("Failed to get socket directory")
|
||||
os.Exit(1)
|
||||
failedTask, taskErr := doTasksE(tasks)
|
||||
if taskErr != nil {
|
||||
if code, ok := serviceStageFailureCode(failedTask); ok {
|
||||
failProvision(newProvisionResult(code, serviceTaskErrorSummary(failedTask, taskErr), nil, provisionSecrets()...), nil)
|
||||
return
|
||||
}
|
||||
reportSetDnsOk(sockDir)
|
||||
// Verify service registration after successful start.
|
||||
if err := verifyServiceRegistration(); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("Service registry verification failed")
|
||||
}
|
||||
} else {
|
||||
mainLog.Load().Error().Err(err).Msg("Failed to start existing ctrld service")
|
||||
os.Exit(1)
|
||||
}
|
||||
sockDir, err := socketDir()
|
||||
if err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("Failed to get socket directory")
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
// The daemon can start and still fail provisioning (for example a
|
||||
// listener bind conflict). Self-check like a fresh install so this
|
||||
// path reports the daemon's failure code instead of a false
|
||||
// "Service started" — but never uninstall an existing service.
|
||||
time.Sleep(1 * time.Second)
|
||||
ok, status, err := selfCheckStatus(ctx, s, sockDir)
|
||||
if !ok || status != service.StatusRunning {
|
||||
fallbackMsg := "ctrld service did not pass its post-start self-check"
|
||||
if err != nil {
|
||||
fallbackMsg = fmt.Sprintf("An error occurred while performing test query: %s", err)
|
||||
mainLog.Load().Error().Msg(fallbackMsg)
|
||||
}
|
||||
if status == service.StatusRunning && err == nil {
|
||||
fallbackMsg = "ctrld service was running, but a DNS query could not be sent to its listener; check firewall rules blocking/intercepting/redirecting DNS queries"
|
||||
mainLog.Load().Error().Msg(fallbackMsg)
|
||||
}
|
||||
reportStartFailure(startAttemptAt, fallbackMsg)
|
||||
return
|
||||
}
|
||||
mainLog.Load().Notice().Msg("Service started")
|
||||
clearProvisionResult()
|
||||
reportSetDnsOk(sockDir)
|
||||
// Verify service registration after successful start.
|
||||
if err := verifyServiceRegistration(); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("Service registry verification failed")
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
@@ -605,7 +673,7 @@ NOTE: running "ctrld start" without any arguments will start already installed c
|
||||
})
|
||||
return nil
|
||||
}, false, "Save current DNS"},
|
||||
{s.Install, false, "Install"},
|
||||
{s.Install, true, "Install"},
|
||||
{func() error {
|
||||
return ConfigureWindowsServiceFailureActions(ctrldServiceName)
|
||||
}, false, "Configure Windows service failure actions"},
|
||||
@@ -614,59 +682,77 @@ NOTE: running "ctrld start" without any arguments will start already installed c
|
||||
// generated after s.Start, so we notice users here for consistent with nextdns mode.
|
||||
{noticeWritingControlDConfig, false, "Notice writing ControlD config"},
|
||||
}
|
||||
// Any result found later must come from this attempt, not a stale run.
|
||||
clearProvisionResult()
|
||||
startAttemptAt := time.Now()
|
||||
mainLog.Load().Notice().Msg("Starting service")
|
||||
if doTasks(tasks) {
|
||||
if err := p.router.Install(sc); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("post installation failed, please check system/service log for details error")
|
||||
failedTask, taskErr := doTasksE(tasks)
|
||||
if taskErr != nil {
|
||||
if code, ok := serviceStageFailureCode(failedTask); ok {
|
||||
failProvision(newProvisionResult(code, serviceTaskErrorSummary(failedTask, taskErr), nil, provisionSecrets()...), nil)
|
||||
return
|
||||
}
|
||||
// Not a service-stage task. doTasksE already logged the cause; exit
|
||||
// non-zero instead of the old silent fall-through that exited 0.
|
||||
os.Exit(1)
|
||||
return
|
||||
}
|
||||
|
||||
// add a small delay to ensure the service is started and did not crash
|
||||
time.Sleep(1 * time.Second)
|
||||
if err := p.router.Install(sc); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("post installation failed, please check system/service log for details error")
|
||||
return
|
||||
}
|
||||
|
||||
ok, status, err := selfCheckStatus(ctx, s, sockDir)
|
||||
switch {
|
||||
case ok && status == service.StatusRunning:
|
||||
mainLog.Load().Notice().Msg("Service started")
|
||||
default:
|
||||
marker := bytes.Repeat([]byte("="), 32)
|
||||
// If ctrld service is not running, emitting log obtained from ctrld process.
|
||||
if status != service.StatusRunning || ctx.Err() != nil {
|
||||
mainLog.Load().Error().Msg("ctrld service may not have started due to an error or misconfiguration, service log:")
|
||||
_, _ = mainLog.Load().Write(marker)
|
||||
haveLog := false
|
||||
for msg := range runCmdLogCh {
|
||||
_, _ = mainLog.Load().Write([]byte(strings.ReplaceAll(msg, msgExit, "")))
|
||||
haveLog = true
|
||||
}
|
||||
// If we're unable to get log from "ctrld run", notice users about it.
|
||||
if !haveLog {
|
||||
mainLog.Load().Write([]byte(`<no log output is obtained from ctrld process>"`))
|
||||
}
|
||||
}
|
||||
// Report any error if occurred.
|
||||
if err != nil {
|
||||
_, _ = mainLog.Load().Write(marker)
|
||||
msg := fmt.Sprintf("An error occurred while performing test query: %s", err)
|
||||
mainLog.Load().Write([]byte(msg))
|
||||
}
|
||||
// If ctrld service is running but selfCheckStatus failed, it could be related
|
||||
// to user's system firewall configuration, notice users about it.
|
||||
if status == service.StatusRunning && err == nil {
|
||||
_, _ = mainLog.Load().Write(marker)
|
||||
mainLog.Load().Write([]byte(`ctrld service was running, but a DNS query could not be sent to its listener`))
|
||||
mainLog.Load().Write([]byte(`Please check your system firewall if it is configured to block/intercept/redirect DNS queries`))
|
||||
}
|
||||
// add a small delay to ensure the service is started and did not crash
|
||||
time.Sleep(1 * time.Second)
|
||||
|
||||
ok, status, err := selfCheckStatus(ctx, s, sockDir)
|
||||
switch {
|
||||
case ok && status == service.StatusRunning:
|
||||
mainLog.Load().Notice().Msg("Service started")
|
||||
clearProvisionResult()
|
||||
default:
|
||||
marker := bytes.Repeat([]byte("="), 32)
|
||||
fallbackMsg := "ctrld service did not pass its post-start self-check"
|
||||
// If ctrld service is not running, emitting log obtained from ctrld process.
|
||||
if status != service.StatusRunning || ctx.Err() != nil {
|
||||
mainLog.Load().Error().Msg("ctrld service may not have started due to an error or misconfiguration, service log:")
|
||||
_, _ = mainLog.Load().Write(marker)
|
||||
uninstall(p, s)
|
||||
os.Exit(1)
|
||||
haveLog := false
|
||||
for msg := range runCmdLogCh {
|
||||
_, _ = mainLog.Load().Write([]byte(strings.ReplaceAll(msg, msgExit, "")))
|
||||
haveLog = true
|
||||
}
|
||||
// If we're unable to get log from "ctrld run", notice users about it.
|
||||
if !haveLog {
|
||||
mainLog.Load().Write([]byte(`<no log output is obtained from ctrld process>"`))
|
||||
}
|
||||
}
|
||||
reportSetDnsOk(sockDir)
|
||||
// Verify service registration after successful start.
|
||||
if err := verifyServiceRegistration(); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("Service registry verification failed")
|
||||
// Report any error if occurred.
|
||||
if err != nil {
|
||||
_, _ = mainLog.Load().Write(marker)
|
||||
msg := fmt.Sprintf("An error occurred while performing test query: %s", err)
|
||||
mainLog.Load().Write([]byte(msg))
|
||||
fallbackMsg = msg
|
||||
}
|
||||
// If ctrld service is running but selfCheckStatus failed, it could be related
|
||||
// to user's system firewall configuration, notice users about it.
|
||||
if status == service.StatusRunning && err == nil {
|
||||
_, _ = mainLog.Load().Write(marker)
|
||||
mainLog.Load().Write([]byte(`ctrld service was running, but a DNS query could not be sent to its listener`))
|
||||
mainLog.Load().Write([]byte(`Please check your system firewall if it is configured to block/intercept/redirect DNS queries`))
|
||||
fallbackMsg = "ctrld service was running, but a DNS query could not be sent to its listener; check firewall rules blocking/intercepting/redirecting DNS queries"
|
||||
}
|
||||
|
||||
_, _ = mainLog.Load().Write(marker)
|
||||
uninstall(p, s)
|
||||
reportStartFailure(startAttemptAt, fallbackMsg)
|
||||
return
|
||||
}
|
||||
reportSetDnsOk(sockDir)
|
||||
// Verify service registration after successful start.
|
||||
if err := verifyServiceRegistration(); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("Service registry verification failed")
|
||||
}
|
||||
},
|
||||
}
|
||||
@@ -691,7 +777,7 @@ NOTE: running "ctrld start" without any arguments will start already installed c
|
||||
startCmd.Flags().BoolVarP(&startOnly, "start_only", "", false, "Do not install new service")
|
||||
_ = startCmd.Flags().MarkHidden("start_only")
|
||||
startCmd.Flags().BoolVarP(&rfc1918, "rfc1918", "", false, "Listen on RFC1918 addresses when 127.0.0.1 is the only listener")
|
||||
startCmd.Flags().StringVarP(&interceptMode, "intercept-mode", "", "", "OS-level DNS interception mode: 'dns' (with VPN split routing) or 'hard' (all DNS through ctrld, no VPN split routing)")
|
||||
startCmd.Flags().StringVarP(&interceptMode, "intercept-mode", "", "", "OS-level DNS interception mode: 'off' (disable interception and clear a persisted intercept_mode), 'dns' (with VPN split routing), or 'hard' (all DNS through ctrld, no VPN split routing)")
|
||||
|
||||
routerCmd := &cobra.Command{
|
||||
Use: "setup",
|
||||
@@ -748,7 +834,7 @@ NOTE: running "ctrld start" without any arguments will start already installed c
|
||||
startCmd.Run(cmd, args)
|
||||
},
|
||||
}
|
||||
startCmdAlias.Flags().StringVarP(&ifaceStartStop, "iface", "", "auto", `Update DNS setting for iface, "auto" means the default interface gateway`)
|
||||
startCmdAlias.Flags().StringVarP(&ifaceStartStop, "iface", "", autoIface, `Update DNS setting for iface, "auto" means the default interface gateway`)
|
||||
startCmdAlias.Flags().AddFlagSet(startCmd.Flags())
|
||||
rootCmd.AddCommand(startCmdAlias)
|
||||
|
||||
@@ -833,7 +919,7 @@ func initStopCmd() *cobra.Command {
|
||||
stopCmd.Run(cmd, args)
|
||||
},
|
||||
}
|
||||
stopCmdAlias.Flags().StringVarP(&ifaceStartStop, "iface", "", "auto", `Reset DNS setting for iface, "auto" means the default interface gateway`)
|
||||
stopCmdAlias.Flags().StringVarP(&ifaceStartStop, "iface", "", autoIface, `Reset DNS setting for iface, "auto" means the default interface gateway`)
|
||||
stopCmdAlias.Flags().AddFlagSet(stopCmd.Flags())
|
||||
rootCmd.AddCommand(stopCmdAlias)
|
||||
|
||||
@@ -865,7 +951,7 @@ func initRestartCmd() *cobra.Command {
|
||||
return
|
||||
}
|
||||
if iface == "" {
|
||||
iface = "auto"
|
||||
iface = autoIface
|
||||
}
|
||||
p.preRun()
|
||||
if ir := runningIface(s); ir != nil {
|
||||
@@ -1047,6 +1133,7 @@ func initStatusCmd() *cobra.Command {
|
||||
statusCmd := &cobra.Command{
|
||||
Use: "status",
|
||||
Short: "Show status of the ctrld service",
|
||||
Long: statusCmdLong,
|
||||
Args: cobra.NoArgs,
|
||||
Run: func(cmd *cobra.Command, args []string) {
|
||||
s, err := newService(&prog{}, svcConfig)
|
||||
@@ -1062,13 +1149,25 @@ func initStatusCmd() *cobra.Command {
|
||||
switch status {
|
||||
case service.StatusUnknown:
|
||||
mainLog.Load().Notice().Msg("Unknown status")
|
||||
os.Exit(2)
|
||||
os.Exit(statusExitUnknown)
|
||||
case service.StatusRunning:
|
||||
mainLog.Load().Notice().Msg("Service is running")
|
||||
os.Exit(0)
|
||||
// The service manager only knows a process was created. It reports a
|
||||
// service as running even when the process is still in startup, with
|
||||
// no control socket, no DNS listener and no policy applied - so
|
||||
// "Service is running" can describe a host with no working DNS.
|
||||
// Probe readiness before claiming it.
|
||||
ready, probeErr := serviceReady()
|
||||
if probeErr != nil {
|
||||
mainLog.Load().Debug().Err(probeErr).Msg("Readiness probe did not confirm startup")
|
||||
}
|
||||
r := classifyReadiness(ready, probeErr, readinessVerifiable())
|
||||
for _, msg := range r.messages {
|
||||
mainLog.Load().Notice().Msg(msg)
|
||||
}
|
||||
os.Exit(r.exitCode)
|
||||
case service.StatusStopped:
|
||||
mainLog.Load().Notice().Msg("Service is stopped")
|
||||
os.Exit(1)
|
||||
os.Exit(statusExitStopped)
|
||||
}
|
||||
},
|
||||
}
|
||||
@@ -1082,6 +1181,7 @@ func initStatusCmd() *cobra.Command {
|
||||
statusCmdAlias := &cobra.Command{
|
||||
Use: "status",
|
||||
Short: "Show status of the ctrld service",
|
||||
Long: statusCmdLong,
|
||||
Args: cobra.NoArgs,
|
||||
Run: statusCmd.Run,
|
||||
}
|
||||
@@ -1111,7 +1211,7 @@ NOTE: Uninstalling will set DNS to values provided by DHCP.`,
|
||||
return
|
||||
}
|
||||
if iface == "" {
|
||||
iface = "auto"
|
||||
iface = autoIface
|
||||
}
|
||||
p.preRun()
|
||||
if ir := runningIface(s); ir != nil {
|
||||
@@ -1207,7 +1307,7 @@ NOTE: Uninstalling will set DNS to values provided by DHCP.`,
|
||||
uninstallCmd.Run(cmd, args)
|
||||
},
|
||||
}
|
||||
uninstallCmdAlias.Flags().StringVarP(&ifaceStartStop, "iface", "", "auto", `Reset DNS setting for iface, "auto" means the default interface gateway`)
|
||||
uninstallCmdAlias.Flags().StringVarP(&ifaceStartStop, "iface", "", autoIface, `Reset DNS setting for iface, "auto" means the default interface gateway`)
|
||||
uninstallCmdAlias.Flags().AddFlagSet(uninstallCmd.Flags())
|
||||
rootCmd.AddCommand(uninstallCmdAlias)
|
||||
|
||||
@@ -1400,7 +1500,7 @@ func initUpgradeCmd() *cobra.Command {
|
||||
return
|
||||
}
|
||||
if iface == "" {
|
||||
iface = "auto"
|
||||
iface = autoIface
|
||||
}
|
||||
p.preRun()
|
||||
if ir := runningIface(s); ir != nil {
|
||||
@@ -1501,28 +1601,31 @@ func initUpgradeCmd() *cobra.Command {
|
||||
if doRestart() {
|
||||
_ = os.Remove(oldBin)
|
||||
_ = os.Chmod(bin, 0755)
|
||||
ver := "unknown version"
|
||||
out, err := exec.Command(bin, "--version").CombinedOutput()
|
||||
ver, err := binaryVersion(bin)
|
||||
if err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("Failed to get new binary version")
|
||||
}
|
||||
if after, found := strings.CutPrefix(string(out), "ctrld version "); found {
|
||||
ver = after
|
||||
ver = "unknown version"
|
||||
}
|
||||
mainLog.Load().Notice().Msgf("Upgrade successful - %s", ver)
|
||||
return
|
||||
}
|
||||
|
||||
mainLog.Load().Warn().Msgf("Upgrade failed, restoring previous binary: %s", oldBin)
|
||||
if err := os.Remove(bin); err != nil {
|
||||
mainLog.Load().Fatal().Err(err).Msg("failed to remove new binary")
|
||||
mainLog.Load().Warn().Msg("Upgrade failed: the new binary did not become ready")
|
||||
stop := func() error {
|
||||
if !svcInstalled {
|
||||
return nil
|
||||
}
|
||||
if err := stopServiceAndWait(s, upgradeStopTimeout); err != nil {
|
||||
return err
|
||||
}
|
||||
// Mirror the Cleanup task in doRestart: leave DNS settings as the OS
|
||||
// had them, not as a half-started ctrld left them.
|
||||
p.router.Cleanup()
|
||||
p.resetDNS(false, true)
|
||||
return nil
|
||||
}
|
||||
if err := os.Rename(oldBin, bin); err != nil {
|
||||
mainLog.Load().Fatal().Err(err).Msg("failed to restore old binary")
|
||||
}
|
||||
if doRestart() {
|
||||
mainLog.Load().Notice().Msg("Restored previous binary successfully")
|
||||
return
|
||||
if err := rollbackToPreviousBinary(bin, oldBin, stop, doRestart); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("Rollback did not complete")
|
||||
}
|
||||
},
|
||||
}
|
||||
@@ -1595,7 +1698,7 @@ func onlyInterceptFlags(args []string) bool {
|
||||
} else {
|
||||
return false
|
||||
}
|
||||
case arg == "--iface=auto" || arg == "--iface" || arg == "auto":
|
||||
case arg == "--iface="+autoIface || arg == "--iface" || arg == autoIface:
|
||||
// Auto-added by startCmdAlias or its value; safe to ignore.
|
||||
continue
|
||||
default:
|
||||
|
||||
@@ -0,0 +1,122 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestServiceStageFailureCode(t *testing.T) {
|
||||
tests := []struct {
|
||||
taskName string
|
||||
wantCode provisionFailureCode
|
||||
wantOK bool
|
||||
}{
|
||||
{"Install", provisionCodeServiceInstall, true},
|
||||
{"Start", provisionCodeServiceStartFailed, true},
|
||||
{"Checking config", "", false},
|
||||
{"", "", false},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
code, ok := serviceStageFailureCode(tc.taskName)
|
||||
if code != tc.wantCode || ok != tc.wantOK {
|
||||
t.Errorf("serviceStageFailureCode(%q) = (%q, %v), want (%q, %v)", tc.taskName, code, ok, tc.wantCode, tc.wantOK)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func stubProvisionExit(t *testing.T) *int {
|
||||
t.Helper()
|
||||
exitCode := -1
|
||||
old := provisionExit
|
||||
provisionExit = func(code int) { exitCode = code }
|
||||
t.Cleanup(func() { provisionExit = old })
|
||||
return &exitCode
|
||||
}
|
||||
|
||||
func TestReportStartFailureUsesFreshDaemonResult(t *testing.T) {
|
||||
overrideProvisionResultPath(t)
|
||||
exitCode := stubProvisionExit(t)
|
||||
|
||||
startedAt := time.Now()
|
||||
daemonResult := newProvisionResult(provisionCodeAPIUnreachable, "daemon could not reach the API", nil)
|
||||
if err := writeProvisionResult(daemonResult); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
reportStartFailure(startedAt, "generic self-check failure")
|
||||
|
||||
if *exitCode != provisionExitCodeForCode[provisionCodeAPIUnreachable] {
|
||||
t.Errorf("exit code = %d, want the daemon's own exit code %d", *exitCode, provisionExitCodeForCode[provisionCodeAPIUnreachable])
|
||||
}
|
||||
out, err := readProvisionResult()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if out.Code != string(provisionCodeAPIUnreachable) {
|
||||
t.Errorf("persisted code = %q, want the daemon's own code untouched", out.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReportStartFailureFallsBackOnStaleDaemonResult(t *testing.T) {
|
||||
overrideProvisionResultPath(t)
|
||||
exitCode := stubProvisionExit(t)
|
||||
|
||||
stale := newProvisionResult(provisionCodeAPIUnreachable, "an old failure", nil)
|
||||
stale.Timestamp = time.Now().Add(-1 * time.Hour).UTC().Format(time.RFC3339)
|
||||
if err := writeProvisionResult(stale); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
startedAt := time.Now()
|
||||
reportStartFailure(startedAt, "test query failed: timeout")
|
||||
|
||||
if *exitCode != provisionExitCodeForCode[provisionCodeServiceSelfCheck] {
|
||||
t.Errorf("exit code = %d, want SERVICE_SELFCHECK_FAILED exit %d", *exitCode, provisionExitCodeForCode[provisionCodeServiceSelfCheck])
|
||||
}
|
||||
out, err := readProvisionResult()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if out.Code != string(provisionCodeServiceSelfCheck) {
|
||||
t.Errorf("persisted code = %q, want %q", out.Code, provisionCodeServiceSelfCheck)
|
||||
}
|
||||
if out.Message != "test query failed: timeout" {
|
||||
t.Errorf("persisted message = %q, want the fallback message", out.Message)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReportStartFailureRejectsUntrustedFile(t *testing.T) {
|
||||
overrideProvisionResultPath(t)
|
||||
exitCode := stubProvisionExit(t)
|
||||
|
||||
planted := newProvisionResult(provisionCodeAPIUnreachable, "planted", nil)
|
||||
planted.Code = "FAKE_CODE"
|
||||
planted.ExitCode = 99
|
||||
if err := writeProvisionResult(planted); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
reportStartFailure(time.Now().Add(-time.Minute), "self-check failed")
|
||||
|
||||
if *exitCode != provisionExitCodeForCode[provisionCodeServiceSelfCheck] {
|
||||
t.Errorf("exit = %d, want the fallback %d, never the planted 99", *exitCode, provisionExitCodeForCode[provisionCodeServiceSelfCheck])
|
||||
}
|
||||
}
|
||||
|
||||
func TestReportStartFailureFallsBackWhenResultFileMissing(t *testing.T) {
|
||||
overrideProvisionResultPath(t)
|
||||
exitCode := stubProvisionExit(t)
|
||||
|
||||
reportStartFailure(time.Now(), "firewall hint")
|
||||
|
||||
if *exitCode != provisionExitCodeForCode[provisionCodeServiceSelfCheck] {
|
||||
t.Errorf("exit code = %d, want SERVICE_SELFCHECK_FAILED exit %d", *exitCode, provisionExitCodeForCode[provisionCodeServiceSelfCheck])
|
||||
}
|
||||
out, err := readProvisionResult()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if out.Message != "firewall hint" {
|
||||
t.Errorf("persisted message = %q, want the fallback message", out.Message)
|
||||
}
|
||||
}
|
||||
@@ -59,12 +59,18 @@ func newControlServer(addr string) (*controlServer, error) {
|
||||
func (s *controlServer) start() error {
|
||||
_ = os.Remove(s.addr)
|
||||
unixListener, err := net.Listen("unix", s.addr)
|
||||
if l, ok := unixListener.(*net.UnixListener); ok {
|
||||
l.SetUnlinkOnClose(true)
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// Restrict socket permissions to owner-only (0600) so that only the
|
||||
// process owner (typically root) can connect. Defense-in-depth since
|
||||
// the control server endpoints carry no authentication of their own.
|
||||
if err := os.Chmod(s.addr, 0600); err != nil {
|
||||
return err
|
||||
}
|
||||
if l, ok := unixListener.(*net.UnixListener); ok {
|
||||
l.SetUnlinkOnClose(true)
|
||||
}
|
||||
go s.server.Serve(unixListener)
|
||||
return nil
|
||||
}
|
||||
@@ -219,13 +225,19 @@ func (p *prog) registerControlServerHandler() {
|
||||
return
|
||||
}
|
||||
|
||||
// Reject further attempts while locked out due to repeated wrong PINs.
|
||||
if now := time.Now().Unix(); now < deactivationLockedUntil.Load() {
|
||||
w.WriteHeader(http.StatusTooManyRequests)
|
||||
return
|
||||
}
|
||||
|
||||
// Re-fetch pin code from API.
|
||||
rcReq := &controld.ResolverConfigRequest{
|
||||
RawUID: cdUID,
|
||||
Version: rootCmd.Version,
|
||||
Metadata: ctrld.SystemMetadataRuntime(context.Background()),
|
||||
}
|
||||
if rc, err := controld.FetchResolverConfig(rcReq, cdDev); rc != nil {
|
||||
if rc, err := controld.FetchResolverConfig(context.Background(), rcReq, cdDev); rc != nil {
|
||||
if rc.DeactivationPin != nil {
|
||||
cdDeactivationPin.Store(*rc.DeactivationPin)
|
||||
} else {
|
||||
@@ -252,6 +264,7 @@ func (p *prog) registerControlServerHandler() {
|
||||
switch req.Pin {
|
||||
case cdDeactivationPin.Load():
|
||||
code = http.StatusOK
|
||||
deactivationFailedAttempts.Store(0)
|
||||
select {
|
||||
case p.pinCodeValidCh <- struct{}{}:
|
||||
default:
|
||||
@@ -259,6 +272,11 @@ func (p *prog) registerControlServerHandler() {
|
||||
case defaultDeactivationPin:
|
||||
// If the pin code was set, but users do not provide --pin, return proper code to client.
|
||||
code = http.StatusBadRequest
|
||||
default:
|
||||
if deactivationFailedAttempts.Add(1) >= deactivationMaxFailedAttempts {
|
||||
deactivationLockedUntil.Store(time.Now().Unix() + deactivationLockoutSeconds)
|
||||
deactivationFailedAttempts.Store(0)
|
||||
}
|
||||
}
|
||||
w.WriteHeader(code)
|
||||
}))
|
||||
@@ -333,7 +351,7 @@ func (p *prog) registerControlServerHandler() {
|
||||
}
|
||||
mainLog.Load().Debug().Msg("sending log file to ControlD server")
|
||||
resp := logSentResponse{Size: r.size}
|
||||
if err := controld.SendLogs(req, cdDev); err != nil {
|
||||
if err := controld.SendLogs(context.Background(), req, cdDev); err != nil {
|
||||
mainLog.Load().Error().Msgf("could not send log file to ControlD server: %v", err)
|
||||
resp.Error = err.Error()
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
|
||||
+718
-351
File diff suppressed because it is too large
Load Diff
@@ -3,8 +3,16 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"tailscale.com/net/netmon"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
@@ -122,6 +130,35 @@ func TestPFBuildAnchorRules_Ordering(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestPFBuildAnchorRules_FallbackPort verifies that when the listener falls back
|
||||
// to an alternate local port (e.g. 5354 because mDNSResponder owns *:53), the pf
|
||||
// rdr rules redirect DNS to the ACTUAL bound port, not the configured default 53.
|
||||
// Regression test for #551: pf redirected to a dead port after listener fallback.
|
||||
func TestPFBuildAnchorRules_FallbackPort(t *testing.T) {
|
||||
// Configured/generated listener is 127.0.0.1:53, but the runtime bound port is 5354.
|
||||
p := &prog{cfg: &ctrld.Config{Listener: map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.1", Port: 5354}}}}
|
||||
rules := p.buildPFAnchorRules(nil)
|
||||
|
||||
// rdr must redirect to the actual bound port 5354.
|
||||
if !strings.Contains(rules, "rdr on lo0 inet proto udp from any to ! 127.0.0.1 port 53 -> 127.0.0.1 port 5354") {
|
||||
t.Errorf("UDP rdr must redirect to bound port 5354, got:\n%s", rules)
|
||||
}
|
||||
if !strings.Contains(rules, "rdr on lo0 inet proto tcp from any to ! 127.0.0.1 port 53 -> 127.0.0.1 port 5354") {
|
||||
t.Errorf("TCP rdr must redirect to bound port 5354, got:\n%s", rules)
|
||||
}
|
||||
|
||||
// The rdr redirect target must NOT point at the dead default port 53.
|
||||
// Match the exact port at line end so "port 5354" is not a false positive.
|
||||
if strings.Contains(rules, "-> 127.0.0.1 port 53\n") {
|
||||
t.Errorf("rdr must not redirect to dead port 53 after fallback, got:\n%s", rules)
|
||||
}
|
||||
|
||||
// The inbound accept rule must also target the actual bound port.
|
||||
if !strings.Contains(rules, "127.0.0.1 port 5354") {
|
||||
t.Errorf("pass in rule must reference bound port 5354, got:\n%s", rules)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPFAddressFamily tests the pfAddressFamily helper.
|
||||
func TestPFAddressFamily(t *testing.T) {
|
||||
tests := []struct {
|
||||
@@ -141,3 +178,833 @@ func TestPFAddressFamily(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsResourceExhaustion(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
err error
|
||||
output []byte
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "exec start failure",
|
||||
err: errors.New("fork/exec /sbin/pfctl: resource temporarily unavailable"),
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "fd exhaustion from stderr output",
|
||||
err: errors.New("exit status 1"),
|
||||
output: []byte("pfctl: Pipe: Too many open files"),
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "process exhaustion from wrapped restore error",
|
||||
err: errors.New("failed to dump running filter rules: exit status 1 (output: too many processes)"),
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "ordinary pf syntax failure",
|
||||
err: errors.New("exit status 1"),
|
||||
output: []byte("pfctl: syntax error"),
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "nil error and empty output",
|
||||
want: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := isResourceExhaustion(tt.err, tt.output); got != tt.want {
|
||||
t.Fatalf("isResourceExhaustion() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func stubPFAnchorCheckCommand(t *testing.T, outputs map[string]string) {
|
||||
t.Helper()
|
||||
original := runPFAnchorCheckCommand
|
||||
runPFAnchorCheckCommand = func(args ...string) ([]byte, error) {
|
||||
key := strings.Join(args, " ")
|
||||
output, ok := outputs[key]
|
||||
if !ok {
|
||||
t.Fatalf("unexpected pf anchor check command: pfctl %s", key)
|
||||
}
|
||||
return []byte(output), nil
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
runPFAnchorCheckCommand = original
|
||||
})
|
||||
}
|
||||
|
||||
func TestEnsurePFAnchorActiveRecentRestoreWithIntactRulesDoesNotStabilize(t *testing.T) {
|
||||
stubPFAnchorCheckCommand(t, map[string]string{
|
||||
"-sn": `rdr-anchor "com.controld.ctrld"`,
|
||||
"-sr": `anchor "com.controld.ctrld"`,
|
||||
"-a com.controld.ctrld -sr": "pass in quick on lo0",
|
||||
"-a com.controld.ctrld -sn": "rdr on lo0",
|
||||
})
|
||||
|
||||
p := &prog{
|
||||
dnsInterceptState: &pfState{},
|
||||
stopCh: make(chan struct{}),
|
||||
}
|
||||
restoredAt := time.Now().Add(-time.Second).UnixMilli()
|
||||
p.pfLastRestoreTime.Store(restoredAt)
|
||||
|
||||
if result := p.ensurePFAnchorActive(); result != pfAnchorCheckIntact {
|
||||
t.Fatalf("intact rules result = %v, want intact", result)
|
||||
}
|
||||
if p.pfBackoffMultiplier.Load() != 0 {
|
||||
t.Fatalf("intact rules incremented backoff to %d", p.pfBackoffMultiplier.Load())
|
||||
}
|
||||
if p.pfStabilizing.Load() {
|
||||
t.Fatal("intact rules must not enter stabilization")
|
||||
}
|
||||
if got := p.pfLastRestoreTime.Load(); got != restoredAt {
|
||||
t.Fatalf("intact check changed restore timestamp: got %d, want %d", got, restoredAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsurePFAnchorActiveCheckFailureIsNotIntact(t *testing.T) {
|
||||
original := runPFAnchorCheckCommand
|
||||
runPFAnchorCheckCommand = func(...string) ([]byte, error) {
|
||||
return nil, errors.New("pfctl unavailable")
|
||||
}
|
||||
t.Cleanup(func() { runPFAnchorCheckCommand = original })
|
||||
|
||||
p := &prog{dnsInterceptState: &pfState{}}
|
||||
if result := p.ensurePFAnchorActive(); result != pfAnchorCheckFailed {
|
||||
t.Fatalf("failed PF inspection result = %v, want failed", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsurePFAnchorActiveRecentActualWipeStartsStabilization(t *testing.T) {
|
||||
stubPFAnchorCheckCommand(t, map[string]string{
|
||||
"-sn": "",
|
||||
})
|
||||
|
||||
stopCh := make(chan struct{})
|
||||
close(stopCh)
|
||||
p := &prog{
|
||||
dnsInterceptState: &pfState{},
|
||||
stopCh: stopCh,
|
||||
}
|
||||
restoredAt := time.Now().Add(-time.Second).UnixMilli()
|
||||
p.pfLastRestoreTime.Store(restoredAt)
|
||||
|
||||
if result := p.ensurePFAnchorActive(); result != pfAnchorCheckDeferred {
|
||||
t.Fatalf("recent repeated wipe result = %v, want deferred", result)
|
||||
}
|
||||
if got := p.pfBackoffMultiplier.Load(); got != 1 {
|
||||
t.Fatalf("recent repeated wipe backoff = %d, want 1", got)
|
||||
}
|
||||
if got := p.pfLastRestoreTime.Load(); got != restoredAt {
|
||||
t.Fatalf("deferred restore changed restore timestamp: got %d, want %d", got, restoredAt)
|
||||
}
|
||||
deadline := time.Now().Add(time.Second)
|
||||
for p.pfStabilizing.Load() && time.Now().Before(deadline) {
|
||||
time.Sleep(time.Millisecond)
|
||||
}
|
||||
if p.pfStabilizing.Load() {
|
||||
t.Fatal("stabilization goroutine did not observe closed stop channel")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDNSInterceptIgnoredChangeReconcileDue(t *testing.T) {
|
||||
p := &prog{}
|
||||
start := time.Unix(1_000_000, 0)
|
||||
|
||||
if !p.dnsInterceptIgnoredChangeReconcileDue(start) {
|
||||
t.Fatal("first ignored change must reconcile immediately")
|
||||
}
|
||||
if p.dnsInterceptIgnoredChangeReconcileDue(start.Add(pfIgnoredChangeReconcileInterval - time.Millisecond)) {
|
||||
t.Fatal("ignored changes inside the interval must be coalesced")
|
||||
}
|
||||
if !p.dnsInterceptIgnoredChangeReconcileDue(start.Add(pfIgnoredChangeReconcileInterval)) {
|
||||
t.Fatal("continuous ignored changes must reconcile again at the interval boundary")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIgnoredNetworkChangeCallbackBoundsWorkWithoutBurningStabilizedSlot(t *testing.T) {
|
||||
outputs := map[string]string{
|
||||
"-sn": `rdr-anchor "com.controld.ctrld"`,
|
||||
"-sr": `anchor "com.controld.ctrld"`,
|
||||
"-a com.controld.ctrld -sr": "pass in quick on lo0",
|
||||
"-a com.controld.ctrld -sn": "rdr on lo0",
|
||||
}
|
||||
originalCheck := runPFAnchorCheckCommand
|
||||
pfChecks := 0
|
||||
runPFAnchorCheckCommand = func(args ...string) ([]byte, error) {
|
||||
key := strings.Join(args, " ")
|
||||
output, ok := outputs[key]
|
||||
if !ok {
|
||||
t.Fatalf("unexpected pf anchor check command: pfctl %s", key)
|
||||
}
|
||||
if key == "-sn" {
|
||||
pfChecks++
|
||||
}
|
||||
return []byte(output), nil
|
||||
}
|
||||
originalDiscover := discoverTunnelInterfacesForReconcile
|
||||
discoverTunnelInterfacesForReconcile = func() []string { return nil }
|
||||
t.Cleanup(func() {
|
||||
runPFAnchorCheckCommand = originalCheck
|
||||
discoverTunnelInterfacesForReconcile = originalDiscover
|
||||
})
|
||||
|
||||
refreshes := 0
|
||||
vpnDNS := newVPNDNSManager(nil)
|
||||
vpnDNS.discoverVPNDNS = func(context.Context) []ctrld.VPNDNSConfig {
|
||||
refreshes++
|
||||
return nil
|
||||
}
|
||||
p := &prog{dnsInterceptState: &pfState{}, vpnDNS: vpnDNS}
|
||||
t.Cleanup(func() {
|
||||
p.pfDelayedRecheckMu.Lock()
|
||||
defer p.pfDelayedRecheckMu.Unlock()
|
||||
for _, timer := range p.pfDelayedRecheckTimers {
|
||||
if timer != nil {
|
||||
timer.Stop()
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
delta := &netmon.ChangeDelta{
|
||||
Old: &netmon.State{Interface: map[string]netmon.Interface{}},
|
||||
New: &netmon.State{Interface: map[string]netmon.Interface{}},
|
||||
}
|
||||
start := time.Unix(1_000_000, 0)
|
||||
p.handleDNSInterceptIgnoredNetworkChange(delta, start)
|
||||
if pfChecks != 1 || refreshes != 1 {
|
||||
t.Fatalf("first ignored delta work: pf checks=%d refreshes=%d, want 1 each", pfChecks, refreshes)
|
||||
}
|
||||
|
||||
p.pfStabilizing.Store(true)
|
||||
p.handleDNSInterceptIgnoredNetworkChange(delta, start.Add(pfIgnoredChangeReconcileInterval))
|
||||
if pfChecks != 1 || refreshes != 1 {
|
||||
t.Fatalf("stabilized delta ran leading reconciliation: pf checks=%d refreshes=%d", pfChecks, refreshes)
|
||||
}
|
||||
|
||||
p.pfStabilizing.Store(false)
|
||||
resumeAt := start.Add(pfIgnoredChangeReconcileInterval + time.Millisecond)
|
||||
p.handleDNSInterceptIgnoredNetworkChange(delta, resumeAt)
|
||||
if pfChecks != 2 || refreshes != 2 {
|
||||
t.Fatalf("first post-stabilization delta did not reconcile immediately: pf checks=%d refreshes=%d", pfChecks, refreshes)
|
||||
}
|
||||
|
||||
for i := 1; i <= 8; i++ {
|
||||
p.handleDNSInterceptIgnoredNetworkChange(delta, resumeAt.Add(time.Duration(i)*100*time.Millisecond))
|
||||
}
|
||||
if pfChecks != 2 || refreshes != 2 {
|
||||
t.Fatalf("ignored delta burst was not coalesced: pf checks=%d refreshes=%d", pfChecks, refreshes)
|
||||
}
|
||||
|
||||
p.handleDNSInterceptIgnoredNetworkChange(delta, resumeAt.Add(pfIgnoredChangeReconcileInterval))
|
||||
if pfChecks != 3 || refreshes != 3 {
|
||||
t.Fatalf("interval boundary did not reconcile: pf checks=%d refreshes=%d, want 3 each", pfChecks, refreshes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRestorePFAnchorFailureIsNotReportedOrTimestamped(t *testing.T) {
|
||||
originalReference := ensurePFAnchorReferenceForRestore
|
||||
originalRebuild := rebuildPFAnchorRulesForReconcile
|
||||
ensurePFAnchorReferenceForRestore = func(*prog) error { return nil }
|
||||
rebuildPFAnchorRulesForReconcile = func(*prog, []vpnDNSExemption) ([]string, error) {
|
||||
return nil, errors.New("pf load failed")
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
ensurePFAnchorReferenceForRestore = originalReference
|
||||
rebuildPFAnchorRulesForReconcile = originalRebuild
|
||||
})
|
||||
|
||||
p := &prog{dnsInterceptState: &pfState{}}
|
||||
if result := p.restorePFAnchor("test"); result != pfAnchorCheckFailed {
|
||||
t.Fatalf("failed restore result = %v, want failed", result)
|
||||
}
|
||||
if got := p.pfLastRestoreTime.Load(); got != 0 {
|
||||
t.Fatalf("failed restore changed timestamp to %d", got)
|
||||
}
|
||||
if len(p.lastTunnelIfaces) != 0 {
|
||||
t.Fatalf("failed restore committed tunnel state: %v", p.lastTunnelIfaces)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPFStabilizationTimeoutReturnsOwnershipToDelayedRecovery(t *testing.T) {
|
||||
p := &prog{dnsInterceptState: &pfState{}}
|
||||
p.pfStabilizing.Store(true)
|
||||
p.pfStabilizationLoopWithMaxWait(t.Context(), time.Hour, 25*time.Millisecond)
|
||||
|
||||
if p.pfStabilizing.Load() {
|
||||
t.Fatal("stabilization retained ownership after the maximum wait")
|
||||
}
|
||||
p.pfDelayedRecheckMu.Lock()
|
||||
timers := append([]*time.Timer(nil), p.pfDelayedRecheckTimers...)
|
||||
p.pfDelayedRecheckTimers = nil
|
||||
p.pfDelayedRecheckMu.Unlock()
|
||||
if len(timers) != 2 {
|
||||
t.Fatalf("expected bounded timeout to schedule delayed recovery, got %d timers", len(timers))
|
||||
}
|
||||
for _, timer := range timers {
|
||||
timer.Stop()
|
||||
}
|
||||
}
|
||||
|
||||
func TestStopDNSInterceptWaitsForInFlightPFMutation(t *testing.T) {
|
||||
binDir := t.TempDir()
|
||||
pfctlPath := filepath.Join(binDir, "pfctl")
|
||||
if err := os.WriteFile(pfctlPath, []byte("#!/bin/sh\nexit 0\n"), 0755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Setenv("PATH", binDir+":"+os.Getenv("PATH"))
|
||||
|
||||
anchorFile := filepath.Join(t.TempDir(), "anchor")
|
||||
if err := os.WriteFile(anchorFile, []byte("rules"), 0600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
p := &prog{dnsInterceptState: &pfState{anchorName: pfAnchorName, anchorFile: anchorFile}}
|
||||
p.pfEnsureRunning.Store(true)
|
||||
|
||||
revoked := make(chan struct{})
|
||||
originalRevokedHook := pfShutdownStateRevokedForTest
|
||||
pfShutdownStateRevokedForTest = func() { close(revoked) }
|
||||
t.Cleanup(func() { pfShutdownStateRevokedForTest = originalRevokedHook })
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() { done <- p.stopDNSIntercept() }()
|
||||
|
||||
select {
|
||||
case <-revoked:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("shutdown did not revoke PF lifecycle state before waiting")
|
||||
}
|
||||
select {
|
||||
case err := <-done:
|
||||
t.Fatalf("shutdown completed before in-flight PF owner released: %v", err)
|
||||
case <-time.After(25 * time.Millisecond):
|
||||
}
|
||||
|
||||
p.pfEnsureRunning.Store(false)
|
||||
if err := <-done; err != nil {
|
||||
t.Fatalf("stopDNSIntercept() error: %v", err)
|
||||
}
|
||||
if _, err := os.Stat(anchorFile); !os.IsNotExist(err) {
|
||||
t.Fatalf("anchor file remained after serialized shutdown: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostStabilizationReconcileRetainsOwnershipAndForcesRebuild(t *testing.T) {
|
||||
stubPFAnchorCheckCommand(t, map[string]string{
|
||||
"-sn": `rdr-anchor "com.controld.ctrld"`,
|
||||
"-sr": `anchor "com.controld.ctrld"`,
|
||||
"-a com.controld.ctrld -sr": "pass in quick on lo0",
|
||||
"-a com.controld.ctrld -sn": "rdr on lo0",
|
||||
})
|
||||
originalRestore := restorePFAnchorForReconcile
|
||||
calls := 0
|
||||
restorePFAnchorForReconcile = func(*prog, string) pfAnchorCheckResult {
|
||||
calls++
|
||||
return pfAnchorCheckRestored
|
||||
}
|
||||
t.Cleanup(func() { restorePFAnchorForReconcile = originalRestore })
|
||||
|
||||
p := &prog{
|
||||
dnsInterceptState: &pfState{},
|
||||
pendingTunnelIfaces: []string{"utun9"},
|
||||
hasPendingTunnelIfaces: true,
|
||||
}
|
||||
p.pfStabilizing.Store(true)
|
||||
if result := p.reconcilePFAnchorAfterStabilization(); result != pfAnchorCheckRestored {
|
||||
t.Fatalf("post-stabilization result = %v, want restored", result)
|
||||
}
|
||||
if calls != 1 {
|
||||
t.Fatalf("post-stabilization restore calls = %d, want 1", calls)
|
||||
}
|
||||
if !p.pfStabilizing.Load() {
|
||||
t.Fatal("post-stabilization reconcile released loop ownership")
|
||||
}
|
||||
if p.pfBackoffMultiplier.Load() != 0 {
|
||||
t.Fatalf("post-stabilization reconcile changed backoff to %d", p.pfBackoffMultiplier.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostStabilizationIntactWithoutPendingAvoidsRebuild(t *testing.T) {
|
||||
stubPFAnchorCheckCommand(t, map[string]string{
|
||||
"-sn": `rdr-anchor "com.controld.ctrld"`,
|
||||
"-sr": `anchor "com.controld.ctrld"`,
|
||||
"-a com.controld.ctrld -sr": "pass in quick on lo0",
|
||||
"-a com.controld.ctrld -sn": "rdr on lo0",
|
||||
})
|
||||
originalRestore := restorePFAnchorForReconcile
|
||||
calls := 0
|
||||
restorePFAnchorForReconcile = func(*prog, string) pfAnchorCheckResult {
|
||||
calls++
|
||||
return pfAnchorCheckRestored
|
||||
}
|
||||
t.Cleanup(func() { restorePFAnchorForReconcile = originalRestore })
|
||||
|
||||
p := &prog{dnsInterceptState: &pfState{}}
|
||||
p.pfStabilizing.Store(true)
|
||||
if result := p.reconcilePFAnchorAfterStabilization(); result != pfAnchorCheckIntact {
|
||||
t.Fatalf("post-stabilization result = %v, want intact", result)
|
||||
}
|
||||
if calls != 0 {
|
||||
t.Fatalf("intact post-stabilization anchor rebuilt %d times", calls)
|
||||
}
|
||||
if !p.pfStabilizing.Load() {
|
||||
t.Fatal("intact post-stabilization reconcile released loop ownership")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelRemovalFailureRetriesBeforeCommittingBaseline(t *testing.T) {
|
||||
originalDiscover := discoverTunnelInterfacesForReconcile
|
||||
originalRestore := restorePFAnchorForReconcile
|
||||
current := []string{}
|
||||
discoverTunnelInterfacesForReconcile = func() []string {
|
||||
return append([]string(nil), current...)
|
||||
}
|
||||
calls := 0
|
||||
restorePFAnchorForReconcile = func(p *prog, _ string) pfAnchorCheckResult {
|
||||
calls++
|
||||
if calls == 1 {
|
||||
return pfAnchorCheckFailed
|
||||
}
|
||||
p.commitPFReconcileState(current)
|
||||
return pfAnchorCheckRestored
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
discoverTunnelInterfacesForReconcile = originalDiscover
|
||||
restorePFAnchorForReconcile = originalRestore
|
||||
})
|
||||
|
||||
p := &prog{
|
||||
dnsInterceptState: &pfState{},
|
||||
lastTunnelIfaces: []string{"utun7"},
|
||||
}
|
||||
if !p.checkTunnelInterfaceChanges() {
|
||||
t.Fatal("first tunnel removal was not detected")
|
||||
}
|
||||
if !stringSlicesEqual(p.lastTunnelIfaces, []string{"utun7"}) {
|
||||
t.Fatalf("failed removal committed baseline: %v", p.lastTunnelIfaces)
|
||||
}
|
||||
if !p.hasPendingTunnelReconcile() {
|
||||
t.Fatal("failed removal did not retain desired tunnel state for retry")
|
||||
}
|
||||
if !p.checkTunnelInterfaceChanges() {
|
||||
t.Fatal("failed tunnel removal was not retried")
|
||||
}
|
||||
if len(p.lastTunnelIfaces) != 0 {
|
||||
t.Fatalf("successful retry did not commit empty tunnel baseline: %v", p.lastTunnelIfaces)
|
||||
}
|
||||
if calls != 2 {
|
||||
t.Fatalf("restore calls = %d, want 2", calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPendingTunnelStateRetriesAfterStabilization(t *testing.T) {
|
||||
originalDiscover := discoverTunnelInterfacesForReconcile
|
||||
originalRestore := restorePFAnchorForReconcile
|
||||
current := []string{}
|
||||
discoverTunnelInterfacesForReconcile = func() []string { return nil }
|
||||
calls := 0
|
||||
restorePFAnchorForReconcile = func(p *prog, _ string) pfAnchorCheckResult {
|
||||
calls++
|
||||
p.commitPFReconcileState(current)
|
||||
return pfAnchorCheckRestored
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
discoverTunnelInterfacesForReconcile = originalDiscover
|
||||
restorePFAnchorForReconcile = originalRestore
|
||||
})
|
||||
|
||||
p := &prog{
|
||||
dnsInterceptState: &pfState{},
|
||||
lastTunnelIfaces: []string{"utun7"},
|
||||
pendingTunnelIfaces: current,
|
||||
hasPendingTunnelIfaces: true,
|
||||
}
|
||||
if !p.checkTunnelInterfaceChanges() {
|
||||
t.Fatal("pending tunnel removal was not retried after stabilization")
|
||||
}
|
||||
if calls != 1 || len(p.lastTunnelIfaces) != 0 || p.hasPendingTunnelReconcile() {
|
||||
t.Fatalf("pending retry result: calls=%d baseline=%v pending=%v", calls, p.lastTunnelIfaces, p.hasPendingTunnelReconcile())
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelReconcileHonorsPFExecBackoff(t *testing.T) {
|
||||
originalDiscover := discoverTunnelInterfacesForReconcile
|
||||
originalRestore := restorePFAnchorForReconcile
|
||||
current := []string{}
|
||||
discoverTunnelInterfacesForReconcile = func() []string { return nil }
|
||||
calls := 0
|
||||
restorePFAnchorForReconcile = func(p *prog, _ string) pfAnchorCheckResult {
|
||||
calls++
|
||||
p.commitPFReconcileState(current)
|
||||
return pfAnchorCheckRestored
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
discoverTunnelInterfacesForReconcile = originalDiscover
|
||||
restorePFAnchorForReconcile = originalRestore
|
||||
})
|
||||
|
||||
p := &prog{dnsInterceptState: &pfState{}, lastTunnelIfaces: []string{"utun7"}}
|
||||
p.pfExecBackoffUntil.Store(time.Now().Add(time.Minute).UnixMilli())
|
||||
if !p.checkTunnelInterfaceChanges() {
|
||||
t.Fatal("tunnel removal was not detected during PF exec backoff")
|
||||
}
|
||||
if calls != 0 || !stringSlicesEqual(p.lastTunnelIfaces, []string{"utun7"}) {
|
||||
t.Fatalf("PF restore ran during exec backoff: calls=%d baseline=%v", calls, p.lastTunnelIfaces)
|
||||
}
|
||||
if p.checkTunnelInterfaceChanges() {
|
||||
t.Fatal("identical deferred tunnel retry bypassed the ignored-event limiter")
|
||||
}
|
||||
p.pfExecBackoffUntil.Store(0)
|
||||
if !p.checkTunnelInterfaceChanges() || calls != 1 || len(p.lastTunnelIfaces) != 0 {
|
||||
t.Fatalf("tunnel removal did not retry after backoff: calls=%d baseline=%v", calls, p.lastTunnelIfaces)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelRapidReversalClearsUnappliedPendingState(t *testing.T) {
|
||||
originalDiscover := discoverTunnelInterfacesForReconcile
|
||||
discoverTunnelInterfacesForReconcile = func() []string { return nil }
|
||||
t.Cleanup(func() { discoverTunnelInterfacesForReconcile = originalDiscover })
|
||||
|
||||
p := &prog{
|
||||
dnsInterceptState: &pfState{},
|
||||
pendingTunnelIfaces: []string{"utun9"},
|
||||
hasPendingTunnelIfaces: true,
|
||||
}
|
||||
p.pfStabilizing.Store(true)
|
||||
if !p.checkTunnelInterfaceChanges() {
|
||||
t.Fatal("rapid tunnel reversal was not observed")
|
||||
}
|
||||
if p.hasPendingTunnelReconcile() || len(p.lastTunnelIfaces) != 0 {
|
||||
t.Fatalf("rapid reversal left unapplied tunnel state: baseline=%v pending=%v", p.lastTunnelIfaces, p.hasPendingTunnelReconcile())
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelAdditionIsCoalescedUntilSuccessfulRebuild(t *testing.T) {
|
||||
originalDiscover := discoverTunnelInterfacesForReconcile
|
||||
current := []string{"utun9"}
|
||||
discoverTunnelInterfacesForReconcile = func() []string {
|
||||
return append([]string(nil), current...)
|
||||
}
|
||||
t.Cleanup(func() { discoverTunnelInterfacesForReconcile = originalDiscover })
|
||||
|
||||
p := &prog{dnsInterceptState: &pfState{}}
|
||||
p.pfStabilizing.Store(true)
|
||||
if !p.checkTunnelInterfaceChanges() {
|
||||
t.Fatal("new tunnel was not detected")
|
||||
}
|
||||
if p.checkTunnelInterfaceChanges() {
|
||||
t.Fatal("identical pending tunnel state was not coalesced")
|
||||
}
|
||||
if len(p.lastTunnelIfaces) != 0 {
|
||||
t.Fatalf("pending tunnel was committed before PF rebuild: %v", p.lastTunnelIfaces)
|
||||
}
|
||||
if !p.hasPendingTunnelReconcile() {
|
||||
t.Fatal("new tunnel was not retained as pending")
|
||||
}
|
||||
|
||||
p.commitPFReconcileState(current)
|
||||
if !stringSlicesEqual(p.lastTunnelIfaces, current) {
|
||||
t.Fatalf("successful rebuild baseline = %v, want %v", p.lastTunnelIfaces, current)
|
||||
}
|
||||
if p.hasPendingTunnelReconcile() {
|
||||
t.Fatal("successful rebuild did not clear pending tunnel state")
|
||||
}
|
||||
}
|
||||
|
||||
// TestVPNDNSRefreshDeferredWhileStabilizing covers the ignored network-change path,
|
||||
// which can trigger a VPN DNS refresh from outside stabilization.
|
||||
//
|
||||
// A refresh rebuilds and reloads the pf anchor. Stabilization owns pf while a VPN's
|
||||
// ruleset is still settling, so refreshing then is the mutual-overwrite collision
|
||||
// stabilization exists to prevent - and these deltas arrive exactly when a VPN is
|
||||
// coming up. Deferring is safe: checkTunnelInterfaceChanges keeps the observation
|
||||
// pending, so the transition is retried afterwards.
|
||||
//
|
||||
// The watchdog tick carries the same guard for the same reason; it is not driven here
|
||||
// because that would mean running its 30s loop.
|
||||
func TestVPNDNSRefreshDeferredWhileStabilizing(t *testing.T) {
|
||||
newProg := func(t *testing.T, refreshes *int, tunnels []string) *prog {
|
||||
t.Helper()
|
||||
outputs := map[string]string{
|
||||
"-sn": `rdr-anchor "com.controld.ctrld"`,
|
||||
"-sr": `anchor "com.controld.ctrld"`,
|
||||
"-a com.controld.ctrld -sr": "pass in quick on lo0",
|
||||
"-a com.controld.ctrld -sn": "rdr on lo0",
|
||||
}
|
||||
originalCheck := runPFAnchorCheckCommand
|
||||
runPFAnchorCheckCommand = func(args ...string) ([]byte, error) {
|
||||
output, ok := outputs[strings.Join(args, " ")]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("unexpected pf anchor check command")
|
||||
}
|
||||
return []byte(output), nil
|
||||
}
|
||||
// Discovery reports no tunnels. With a seeded baseline that is a removal, which
|
||||
// checkTunnelInterfaceChanges reports as a change without touching pf while
|
||||
// stabilizing - so this fixture never reaches a real pfctl write.
|
||||
originalDiscover := discoverTunnelInterfacesForReconcile
|
||||
discoverTunnelInterfacesForReconcile = func() []string { return nil }
|
||||
t.Cleanup(func() {
|
||||
runPFAnchorCheckCommand = originalCheck
|
||||
discoverTunnelInterfacesForReconcile = originalDiscover
|
||||
})
|
||||
|
||||
vpnDNS := newVPNDNSManager(nil)
|
||||
vpnDNS.discoverVPNDNS = func(context.Context) []ctrld.VPNDNSConfig {
|
||||
*refreshes++
|
||||
return nil
|
||||
}
|
||||
p := &prog{dnsInterceptState: &pfState{}, vpnDNS: vpnDNS, lastTunnelIfaces: tunnels}
|
||||
t.Cleanup(func() {
|
||||
p.pfDelayedRecheckMu.Lock()
|
||||
defer p.pfDelayedRecheckMu.Unlock()
|
||||
for _, timer := range p.pfDelayedRecheckTimers {
|
||||
if timer != nil {
|
||||
timer.Stop()
|
||||
}
|
||||
}
|
||||
})
|
||||
return p
|
||||
}
|
||||
delta := func() *netmon.ChangeDelta {
|
||||
return &netmon.ChangeDelta{
|
||||
Old: &netmon.State{Interface: map[string]netmon.Interface{}},
|
||||
New: &netmon.State{Interface: map[string]netmon.Interface{}},
|
||||
}
|
||||
}
|
||||
|
||||
t.Run("tunnel change during stabilization does not refresh", func(t *testing.T) {
|
||||
refreshes := 0
|
||||
// Seeded baseline plus empty discovery = a tunnel transition to report, so the
|
||||
// refresh is eligible on everything except the stabilization guard.
|
||||
p := newProg(t, &refreshes, []string{"utun9"})
|
||||
p.pfStabilizing.Store(true)
|
||||
|
||||
p.handleDNSInterceptIgnoredNetworkChange(delta(), time.Unix(1_000_000, 0))
|
||||
|
||||
if refreshes != 0 {
|
||||
t.Errorf("refreshed %d time(s) while stabilizing — that rebuilds the anchor under a settling VPN ruleset", refreshes)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("refresh still happens outside stabilization", func(t *testing.T) {
|
||||
refreshes := 0
|
||||
p := newProg(t, &refreshes, nil)
|
||||
|
||||
p.handleDNSInterceptIgnoredNetworkChange(delta(), time.Unix(1_000_000, 0))
|
||||
|
||||
if refreshes == 0 {
|
||||
t.Error("no refresh outside stabilization — the guard must defer, not disable")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestExemptVPNDNSServersDeferredWhileStabilizing checks the mutation point itself,
|
||||
// not just the call sites: any future caller reaching it during stabilization is
|
||||
// refused before the anchor is rewritten.
|
||||
//
|
||||
// It returns before pfEnsureRunning is taken and before any pfctl work, so this drives
|
||||
// the real function without touching the host's pf state.
|
||||
func TestExemptVPNDNSServersDeferredWhileStabilizing(t *testing.T) {
|
||||
p := &prog{dnsInterceptState: &pfState{}}
|
||||
p.pfStabilizing.Store(true)
|
||||
|
||||
err := p.exemptVPNDNSServers([]vpnDNSExemption{{Server: "192.168.1.1"}})
|
||||
if err == nil {
|
||||
t.Fatal("exemption applied while stabilizing — that rewrites the anchor under a settling VPN ruleset")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "stabilization") {
|
||||
t.Errorf("error does not name the reason: %v", err)
|
||||
}
|
||||
// The refusal must happen before the reconcile latch is claimed, or a deferral
|
||||
// would lock out the reconcile that runs once stabilization ends.
|
||||
if p.pfEnsureRunning.Load() {
|
||||
t.Error("pfEnsureRunning was left held by a deferred exemption")
|
||||
}
|
||||
}
|
||||
|
||||
// stubStabilizationProbe replaces the post-stabilization verification seams and returns
|
||||
// counters for probe and forced-reload calls.
|
||||
func stubStabilizationProbe(t *testing.T, probeResults []bool, reloadOK bool) (probes, reloads *int) {
|
||||
t.Helper()
|
||||
originalProbe, originalReload := probePFInterceptFn, forceReloadPFInterceptFn
|
||||
t.Cleanup(func() {
|
||||
probePFInterceptFn, forceReloadPFInterceptFn = originalProbe, originalReload
|
||||
})
|
||||
probeCalls, reloadCalls := 0, 0
|
||||
probePFInterceptFn = func(*prog) bool {
|
||||
result := false
|
||||
if probeCalls < len(probeResults) {
|
||||
result = probeResults[probeCalls]
|
||||
}
|
||||
probeCalls++
|
||||
return result
|
||||
}
|
||||
forceReloadPFInterceptFn = func(*prog) bool {
|
||||
reloadCalls++
|
||||
return reloadOK
|
||||
}
|
||||
return &probeCalls, &reloadCalls
|
||||
}
|
||||
|
||||
// TestPostStabilizationVerifiesInterceptionFunctionally is the post-wake continuity
|
||||
// boundary: the reconcile above it only proves rule text, and QA saw rules intact,
|
||||
// references intact and post-load verification passed while every query through the system
|
||||
// resolver timed out. Nothing else probes until the periodic watchdog, because the probe
|
||||
// monitor stands down while stabilization owns pf, so recovery waited for that tick.
|
||||
func TestPostStabilizationVerifiesInterceptionFunctionally(t *testing.T) {
|
||||
// Probe fails once, then passes after the reload.
|
||||
probes, reloads := stubStabilizationProbe(t, []bool{false, true}, true)
|
||||
|
||||
p := &prog{dnsInterceptState: &pfState{}}
|
||||
p.verifyInterceptAfterStabilization()
|
||||
|
||||
if *probes != 2 {
|
||||
t.Errorf("probe calls = %d, want 2: one to detect and one to confirm the repair", *probes)
|
||||
}
|
||||
if *reloads != 1 {
|
||||
t.Errorf("forced reloads = %d, want exactly 1 bounded repair", *reloads)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPostStabilizationProbePassSkipsReload keeps the healthy path free of a pf reload,
|
||||
// which would flush states and kill in-flight DoH connections for nothing.
|
||||
func TestPostStabilizationProbePassSkipsReload(t *testing.T) {
|
||||
probes, reloads := stubStabilizationProbe(t, []bool{true}, true)
|
||||
|
||||
p := &prog{dnsInterceptState: &pfState{}}
|
||||
p.verifyInterceptAfterStabilization()
|
||||
|
||||
if *probes != 1 || *reloads != 0 {
|
||||
t.Errorf("probe calls = %d, forced reloads = %d, want 1/0", *probes, *reloads)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPostStabilizationRepairIsBounded pins the "one bounded recovery" contract: a probe
|
||||
// that never passes must not turn into a reload loop here - the watchdog owns retries.
|
||||
func TestPostStabilizationRepairIsBounded(t *testing.T) {
|
||||
probes, reloads := stubStabilizationProbe(t, []bool{false, false, false}, true)
|
||||
|
||||
p := &prog{dnsInterceptState: &pfState{}}
|
||||
p.verifyInterceptAfterStabilization()
|
||||
|
||||
if *reloads != 1 {
|
||||
t.Errorf("forced reloads = %d, want 1: the repair must not loop", *reloads)
|
||||
}
|
||||
if *probes != 2 {
|
||||
t.Errorf("probe calls = %d, want 2", *probes)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPostStabilizationWaitsForAProberThatStandsDown is the interleaving that made
|
||||
// "skip when the flag is set" wrong. A probe monitor started by an ignored network change
|
||||
// claims functional-probe ownership and then aborts, because stabilization still owns pf.
|
||||
// If the verifier treats the claimed flag as "somebody is probing", neither path probes and
|
||||
// the outage lasts until the next watchdog tick - the exact window this is meant to close.
|
||||
func TestPostStabilizationWaitsForAProberThatStandsDown(t *testing.T) {
|
||||
probes, reloads := stubStabilizationProbe(t, []bool{false, true}, true)
|
||||
|
||||
p := &prog{dnsInterceptState: &pfState{}}
|
||||
// Model the monitor's claim-then-abort: ownership is held, then released.
|
||||
p.pfMonitorRunning.Store(true)
|
||||
released := make(chan struct{})
|
||||
go func() {
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
p.pfMonitorRunning.Store(false)
|
||||
close(released)
|
||||
}()
|
||||
|
||||
p.verifyInterceptAfterStabilization()
|
||||
<-released
|
||||
|
||||
if *probes != 2 {
|
||||
t.Errorf("probe calls = %d, want 2: the verifier must wait out a prober that stands down", *probes)
|
||||
}
|
||||
if *reloads != 1 {
|
||||
t.Errorf("forced reloads = %d, want 1", *reloads)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPostStabilizationYieldsToAProberThatKeepsProbing is the other half of the handoff:
|
||||
// when the holder is genuinely working through its probe sequence, the verifier must step
|
||||
// aside rather than run a second prober against the same pf state.
|
||||
func TestPostStabilizationYieldsToAProberThatKeepsProbing(t *testing.T) {
|
||||
originalWait := pfFunctionalProbeOwnerWait
|
||||
pfFunctionalProbeOwnerWait = 30 * time.Millisecond
|
||||
t.Cleanup(func() { pfFunctionalProbeOwnerWait = originalWait })
|
||||
|
||||
probes, reloads := stubStabilizationProbe(t, []bool{false}, true)
|
||||
|
||||
p := &prog{dnsInterceptState: &pfState{}}
|
||||
p.pfMonitorRunning.Store(true) // held for the whole wait
|
||||
p.verifyInterceptAfterStabilization()
|
||||
|
||||
if *probes != 0 || *reloads != 0 {
|
||||
t.Errorf("probe calls = %d, forced reloads = %d, want 0/0 while another prober is working", *probes, *reloads)
|
||||
}
|
||||
}
|
||||
|
||||
// TestInterceptMonitorDoesNotClaimOwnershipWhileStabilizing pins the source of that race:
|
||||
// a monitor which cannot do useful work must not take functional-probe ownership on its way
|
||||
// out, or it starves the post-stabilization verifier.
|
||||
func TestInterceptMonitorDoesNotClaimOwnershipWhileStabilizing(t *testing.T) {
|
||||
probes, reloads := stubStabilizationProbe(t, []bool{false}, true)
|
||||
|
||||
p := &prog{dnsInterceptState: &pfState{}}
|
||||
p.pfStabilizing.Store(true)
|
||||
if p.interceptProbeMonitorAllowed() {
|
||||
t.Fatal("the probe monitor considers itself eligible while stabilization owns pf")
|
||||
}
|
||||
|
||||
p.pfInterceptMonitor()
|
||||
|
||||
if *probes != 0 || *reloads != 0 {
|
||||
t.Errorf("probe calls = %d, forced reloads = %d, want 0/0 from a monitor that cannot run", *probes, *reloads)
|
||||
}
|
||||
if !p.claimFunctionalProbeOwner(0) {
|
||||
t.Error("the aborted monitor left functional-probe ownership taken; the verifier would skip")
|
||||
}
|
||||
p.pfMonitorRunning.Store(false)
|
||||
}
|
||||
|
||||
// TestFinishPFStabilizationRunsFunctionalVerification wires the fix to the production
|
||||
// completion path: deleting the verification call, or reordering it before the reconcile,
|
||||
// makes this fail.
|
||||
func TestFinishPFStabilizationRunsFunctionalVerification(t *testing.T) {
|
||||
stubPFAnchorCheckCommand(t, map[string]string{
|
||||
"-sn": `rdr-anchor "com.controld.ctrld"`,
|
||||
"-sr": `anchor "com.controld.ctrld"`,
|
||||
"-a com.controld.ctrld -sr": "pass in quick on lo0",
|
||||
"-a com.controld.ctrld -sn": "rdr on lo0",
|
||||
})
|
||||
originalResolver := initializeOsResolver
|
||||
initializeOsResolver = func(bool) []string { return nil }
|
||||
t.Cleanup(func() { initializeOsResolver = originalResolver })
|
||||
|
||||
probes, reloads := stubStabilizationProbe(t, []bool{false, true}, true)
|
||||
|
||||
p := &prog{dnsInterceptState: &pfState{}}
|
||||
p.pfStabilizing.Store(true)
|
||||
p.finishPFStabilization(time.Millisecond)
|
||||
|
||||
if *probes == 0 {
|
||||
t.Fatal("stabilization completed without probing functional interception; recovery would wait for the watchdog")
|
||||
}
|
||||
if *reloads != 1 {
|
||||
t.Errorf("forced reloads = %d, want 1", *reloads)
|
||||
}
|
||||
|
||||
p.pfDelayedRecheckMu.Lock()
|
||||
timers := append([]*time.Timer(nil), p.pfDelayedRecheckTimers...)
|
||||
p.pfDelayedRecheckTimers = nil
|
||||
p.pfDelayedRecheckMu.Unlock()
|
||||
for _, timer := range timers {
|
||||
timer.Stop()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
//go:build windows
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestDNSInterceptIgnoredChangeReconcileDueWindowsPreservesImmediateBehavior(t *testing.T) {
|
||||
p := &prog{}
|
||||
now := time.Now()
|
||||
|
||||
if !p.dnsInterceptIgnoredChangeReconcileDue(now) {
|
||||
t.Fatal("first ignored Windows change must reconcile immediately")
|
||||
}
|
||||
if !p.dnsInterceptIgnoredChangeReconcileDue(now) {
|
||||
t.Fatal("Windows ignored changes must not inherit the macOS pf rate limit")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,319 @@
|
||||
//go:build windows
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"runtime"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// newInterceptTestProg returns a prog with a published intercept state, fake NRPT
|
||||
// operations already installed, and no WFP engine (engineHandle 0).
|
||||
//
|
||||
// The fake is installed here, before anything can inspect registry state, and it is the
|
||||
// safety boundary - not the empty wfpState. A zero-valued state has owner None, and
|
||||
// shutdown's None branch sweeps orphaned ctrld rules, so an unfaked stopDNSIntercept would
|
||||
// reach the production nrptCatchAllRuleExists / removeNRPTCatchAllRule / signalNRPTChange.
|
||||
// On a host that has ctrld's deterministic key - a developer box, or a CI runner where
|
||||
// ctrld is installed - that deletes live policy and forces a Group Policy refresh, a
|
||||
// Dnscache paramchange and a cache flush. A green run on a clean runner proves nothing
|
||||
// about that.
|
||||
func newInterceptTestProg(t *testing.T) (*prog, *wfpState, *fakeNRPTOps) {
|
||||
t.Helper()
|
||||
f := fakeNRPTOpsForTest(t)
|
||||
// Prove the fake is in effect before anything can inspect registry state. Asserting
|
||||
// zero side effects afterwards cannot do that: an uninstalled fake reports zero
|
||||
// whether it was consulted or bypassed.
|
||||
requireFakeNRPTOpsInstalled(t, f)
|
||||
state := &wfpState{stopCh: make(chan struct{}), listenerIP: "127.0.0.1"}
|
||||
p := &prog{}
|
||||
p.dnsInterceptState = state
|
||||
return p, state, f
|
||||
}
|
||||
|
||||
// assertNoNRPTSideEffects fails when a lifecycle path wrote NRPT policy or signalled the
|
||||
// DNS Client. Every test in this file exercises a guard that is supposed to stand down, so
|
||||
// any registry write or signal here means the guard did not hold - and, without the fake,
|
||||
// would have hit the host's real policy.
|
||||
func assertNoNRPTSideEffects(t *testing.T, f *fakeNRPTOps) {
|
||||
t.Helper()
|
||||
add, remove, signal, _ := f.counts()
|
||||
if add != 0 || remove != 0 || signal != 0 {
|
||||
t.Errorf("addRule = %d, removeRule = %d, signal = %d, want 0/0/0: this path must not write NRPT policy",
|
||||
add, remove, signal)
|
||||
}
|
||||
if flush := f.flushCount(); flush != 0 {
|
||||
t.Errorf("flush calls = %d, want 0: this path must not flush the resolver cache", flush)
|
||||
}
|
||||
}
|
||||
|
||||
// TestStopDNSInterceptRevokesBeforeTeardown pins the ordering the shutdown/monitor race
|
||||
// depends on. Teardown deletes our WFP sublayer, and a missing sublayer is precisely what
|
||||
// the health monitor treats as "our filters were wiped, rebuild everything". Were the
|
||||
// state revoked only after teardown, a monitor tick inside that window would rebuild the
|
||||
// intercept during shutdown.
|
||||
func TestStopDNSInterceptRevokesBeforeTeardown(t *testing.T) {
|
||||
p, state, f := newInterceptTestProg(t)
|
||||
defer assertNoNRPTSideEffects(t, f)
|
||||
|
||||
if p.interceptStateRevoked(state) {
|
||||
t.Fatal("a freshly published intercept state must not read as retired")
|
||||
}
|
||||
if err := p.stopDNSIntercept(); err != nil {
|
||||
t.Fatalf("stopDNSIntercept() = %v", err)
|
||||
}
|
||||
if !p.interceptStateRevoked(state) {
|
||||
t.Error("state still reads live after shutdown: the monitor and heal flows would keep writing host DNS state")
|
||||
}
|
||||
if p.dnsInterceptState != nil {
|
||||
t.Error("dnsInterceptState survived shutdown")
|
||||
}
|
||||
if p.dnsInterceptStopRequested.Load() {
|
||||
t.Error("stop-requested flag was left set; a later start would see a phantom shutdown")
|
||||
}
|
||||
}
|
||||
|
||||
// TestRebuildDNSInterceptRefusedAfterShutdown is the regression test for the reported
|
||||
// race: SCM stop runs resetDNS -> stopDNSIntercept while the health monitor is mid-tick,
|
||||
// and the monitor then reaches the rebuild path before the process exits. The rebuild
|
||||
// must refuse - completing it would re-add the NRPT catch-all and the WFP filters moments
|
||||
// before ctrld disappears, leaving Windows resolving through a listener that is gone.
|
||||
//
|
||||
// That refusal is also what keeps this test safe on a real Windows host: a rebuild that
|
||||
// did not refuse would run startDNSIntercept and write NRPT policy to the machine
|
||||
// running the tests.
|
||||
func TestRebuildDNSInterceptRefusedAfterShutdown(t *testing.T) {
|
||||
p, state, f := newInterceptTestProg(t)
|
||||
defer assertNoNRPTSideEffects(t, f)
|
||||
if err := p.stopDNSIntercept(); err != nil {
|
||||
t.Fatalf("stopDNSIntercept() = %v", err)
|
||||
}
|
||||
|
||||
if got := p.rebuildDNSIntercept(state, "WFP sublayer missing during health check"); got != interceptRebuildRetired {
|
||||
t.Fatalf("rebuildDNSIntercept() = %v, want interceptRebuildRetired - a post-shutdown rebuild resurrects DNS interception", got)
|
||||
}
|
||||
if p.dnsInterceptState != nil {
|
||||
t.Error("rebuild published new intercept state after shutdown")
|
||||
}
|
||||
}
|
||||
|
||||
// TestRebuildDNSInterceptRefusedForReplacedState covers the other stale-owner case: an
|
||||
// earlier rebuild already replaced the state, so a goroutine still holding the old one
|
||||
// must not tear down its successor.
|
||||
func TestRebuildDNSInterceptRefusedForReplacedState(t *testing.T) {
|
||||
p, old, f := newInterceptTestProg(t)
|
||||
defer assertNoNRPTSideEffects(t, f)
|
||||
current := &wfpState{stopCh: make(chan struct{}), listenerIP: "127.0.0.1"}
|
||||
p.dnsInterceptState = current
|
||||
|
||||
if got := p.rebuildDNSIntercept(old, "WFP sublayer missing during health check"); got != interceptRebuildRetired {
|
||||
t.Fatalf("rebuildDNSIntercept() = %v, want interceptRebuildRetired for a superseded state", got)
|
||||
}
|
||||
if p.dnsInterceptState != any(current) {
|
||||
t.Error("a superseded state's rebuild replaced the live intercept")
|
||||
}
|
||||
if p.interceptStateRevoked(current) {
|
||||
t.Error("the live state was revoked by a superseded rebuild")
|
||||
}
|
||||
}
|
||||
|
||||
// TestRepairMissingWFPStandsDownAfterShutdown checks the monitor's entry point. It must
|
||||
// not even query WFP for a retired state - the sublayer it looks for is what teardown
|
||||
// just deleted - and it must tell the monitor goroutine to exit.
|
||||
func TestRepairMissingWFPStandsDownAfterShutdown(t *testing.T) {
|
||||
p, state, f := newInterceptTestProg(t)
|
||||
defer assertNoNRPTSideEffects(t, f)
|
||||
if err := p.stopDNSIntercept(); err != nil {
|
||||
t.Fatalf("stopDNSIntercept() = %v", err)
|
||||
}
|
||||
// Set the handle only after teardown. A fake handle proves the revocation check
|
||||
// comes first, but must never reach the real WFP calls in cleanupWFPFilters.
|
||||
state.engineHandle = 1
|
||||
|
||||
if !p.repairMissingWFP(state) {
|
||||
t.Error("repairMissingWFP() = false after shutdown; the health monitor would keep running for a dead intercept")
|
||||
}
|
||||
if p.dnsInterceptState != nil {
|
||||
t.Error("repairMissingWFP rebuilt the intercept after shutdown")
|
||||
}
|
||||
}
|
||||
|
||||
// TestPendingStopSignalsRevocation covers how a stop avoids waiting: while it is blocked
|
||||
// on the lifecycle lock it must already read as revoked, so an in-flight NRPT heal
|
||||
// abandons its probe backoff instead of making the service stop wait it out. A stop that
|
||||
// waits too long is killed by the Service Control Manager, which cleans up nothing.
|
||||
func TestPendingStopSignalsRevocation(t *testing.T) {
|
||||
p, state, f := newInterceptTestProg(t)
|
||||
defer assertNoNRPTSideEffects(t, f)
|
||||
|
||||
p.dnsInterceptMu.Lock()
|
||||
stopped := make(chan struct{})
|
||||
go func() {
|
||||
defer close(stopped)
|
||||
_ = p.stopDNSIntercept()
|
||||
}()
|
||||
|
||||
// Wait for the stop to announce itself while it is blocked on the lock.
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for !p.dnsInterceptStopRequested.Load() {
|
||||
if time.Now().After(deadline) {
|
||||
p.dnsInterceptMu.Unlock()
|
||||
<-stopped
|
||||
t.Fatal("stop never announced itself before waiting for the lifecycle lock")
|
||||
}
|
||||
runtime.Gosched()
|
||||
}
|
||||
if !p.interceptStateRevoked(state) {
|
||||
t.Error("a pending stop does not read as revoked; the heal flows would keep it waiting")
|
||||
}
|
||||
p.dnsInterceptMu.Unlock()
|
||||
<-stopped
|
||||
|
||||
if p.dnsInterceptState != nil {
|
||||
t.Error("the pending stop did not tear down the intercept once it acquired the lock")
|
||||
}
|
||||
}
|
||||
|
||||
// TestInterceptWaitAbandonsPromptlyOnPendingStop is the bound on how long a stop can be
|
||||
// delayed by a recovery flow: the heal sequence's waits add up to tens of seconds, and
|
||||
// each one must end as soon as a stop is pending.
|
||||
func TestInterceptWaitAbandonsPromptlyOnPendingStop(t *testing.T) {
|
||||
p, state, f := newInterceptTestProg(t)
|
||||
defer assertNoNRPTSideEffects(t, f)
|
||||
p.dnsInterceptStopRequested.Store(true)
|
||||
|
||||
start := time.Now()
|
||||
if p.interceptWait(state, 30*time.Second) {
|
||||
t.Fatal("interceptWait() = true with a stop pending; the caller would carry on writing host DNS state")
|
||||
}
|
||||
if elapsed := time.Since(start); elapsed > 2*time.Second {
|
||||
t.Errorf("interceptWait took %v to notice a pending stop; shutdown would inherit that delay", elapsed)
|
||||
}
|
||||
}
|
||||
|
||||
// TestInterceptWaitRunsToCompletionWhileLive guards the other direction: the cancellable
|
||||
// wait must still actually wait, or the recovery flows lose their backoff.
|
||||
func TestInterceptWaitRunsToCompletionWhileLive(t *testing.T) {
|
||||
p, state, f := newInterceptTestProg(t)
|
||||
defer assertNoNRPTSideEffects(t, f)
|
||||
|
||||
start := time.Now()
|
||||
if !p.interceptWait(state, 250*time.Millisecond) {
|
||||
t.Fatal("interceptWait() = false for a live intercept")
|
||||
}
|
||||
if elapsed := time.Since(start); elapsed < 250*time.Millisecond {
|
||||
t.Errorf("interceptWait returned after %v, want at least 250ms", elapsed)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNRPTNeedsCtrldActivation covers the recovery gap that left a machine unfiltered
|
||||
// until restart: a failed NRPT write clears ownership, and an owner-None tick used to do
|
||||
// nothing at all, so nothing ever retried the write.
|
||||
func TestNRPTNeedsCtrldActivation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
owner nrptRuleOwner
|
||||
ruleExists bool
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
// The reported hole: activation failed, ownership was cleared, and no
|
||||
// other path re-arms it. In hard mode WFP keeps blocking DNS meanwhile.
|
||||
name: "no owner retries the failed write",
|
||||
owner: nrptRuleOwnerNone,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "no owner retries even if a rule is somehow present",
|
||||
owner: nrptRuleOwnerNone,
|
||||
ruleExists: true,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "ctrld-owned rule removed externally is re-added",
|
||||
owner: nrptRuleOwnerCtrld,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "healthy ctrld-owned rule is left alone",
|
||||
owner: nrptRuleOwnerCtrld,
|
||||
ruleExists: true,
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
// Writing beside external policy would be ambiguous policy, not recovery.
|
||||
name: "external policy is never overwritten",
|
||||
owner: nrptRuleOwnerGroupPolicy,
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "external policy is never overwritten even with a ctrld rule present",
|
||||
owner: nrptRuleOwnerGroupPolicy,
|
||||
ruleExists: true,
|
||||
want: false,
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := nrptNeedsCtrldActivation(tc.owner, tc.ruleExists); got != tc.want {
|
||||
t.Errorf("nrptNeedsCtrldActivation(%v, %v) = %v, want %v", tc.owner, tc.ruleExists, got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestActivateCtrldNRPTFallbackRefusedAfterShutdown guards the worst leftover. A
|
||||
// catch-all re-added after shutdown points every DNS query on the machine at a listener
|
||||
// that no longer exists, so nothing resolves at all. Refusing early also keeps this test
|
||||
// from writing NRPT policy on the machine running it.
|
||||
func TestActivateCtrldNRPTFallbackRefusedAfterShutdown(t *testing.T) {
|
||||
p, state, f := newInterceptTestProg(t)
|
||||
defer assertNoNRPTSideEffects(t, f)
|
||||
if err := p.stopDNSIntercept(); err != nil {
|
||||
t.Fatalf("stopDNSIntercept() = %v", err)
|
||||
}
|
||||
|
||||
if p.activateCtrldNRPTFallback(state, "ctrld-owned rule missing during health check") {
|
||||
t.Error("activateCtrldNRPTFallback() = true after shutdown: the catch-all would outlive ctrld")
|
||||
}
|
||||
if owner, _ := state.nrptPolicyOwner(); owner != nrptRuleOwnerNone {
|
||||
t.Errorf("NRPT owner = %v after a refused fallback, want nrptRuleOwnerNone", owner)
|
||||
}
|
||||
}
|
||||
|
||||
// TestAllowHandbackAttemptRateLimits covers the throttle on testing an external
|
||||
// catch-all. Each attempt takes ctrld's rule out of the way for a probe, so a rule that
|
||||
// never routes would cost a brief DNS outage on every 30s health tick without this - in
|
||||
// hard mode a window where WFP blocks DNS and nothing redirects it.
|
||||
func TestHandbackThrottleIsPerRule(t *testing.T) {
|
||||
state := &wfpState{stopCh: make(chan struct{})}
|
||||
now := time.Now()
|
||||
|
||||
if !state.handbackAllowed(now, "{GP-RULE}", nrptHandbackRetryInterval) {
|
||||
t.Fatal("first handback attempt must be allowed")
|
||||
}
|
||||
// Checking alone must not spend the budget: a pre-probe can still abort the attempt
|
||||
// without disturbing NRPT, and that must not cost the rule its next window.
|
||||
if !state.handbackAllowed(now, "{GP-RULE}", nrptHandbackRetryInterval) {
|
||||
t.Error("handbackAllowed must not consume the budget by itself")
|
||||
}
|
||||
|
||||
state.recordHandbackAttempt(now, "{GP-RULE}", nrptHandbackRetryInterval)
|
||||
if state.handbackAllowed(now.Add(nrptHandbackRetryInterval-time.Second), "{GP-RULE}", nrptHandbackRetryInterval) {
|
||||
t.Error("re-testing the same rule inside the interval must be suppressed")
|
||||
}
|
||||
// Group Policy alternating between two names must not erase either one's memory:
|
||||
// with a single slot every swap costs another removal of the live rule.
|
||||
if !state.handbackAllowed(now.Add(time.Second), "{OTHER-RULE}", nrptHandbackRetryInterval) {
|
||||
t.Error("a different rule name means the administrator changed policy: test it now")
|
||||
}
|
||||
state.recordHandbackAttempt(now.Add(time.Second), "{OTHER-RULE}", nrptHandbackRetryInterval)
|
||||
if state.handbackAllowed(now.Add(2*time.Second), "{GP-RULE}", nrptHandbackRetryInterval) {
|
||||
t.Error("testing another rule must not clear the first rule's throttle")
|
||||
}
|
||||
|
||||
if !state.handbackAllowed(now.Add(2*nrptHandbackRetryInterval), "{GP-RULE}", nrptHandbackRetryInterval) {
|
||||
t.Error("the same rule must be testable again after the interval")
|
||||
}
|
||||
}
|
||||
@@ -4,6 +4,7 @@ package cli
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// startDNSIntercept is not supported on this platform.
|
||||
@@ -17,14 +18,17 @@ func (p *prog) stopDNSIntercept() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// skipInitialDNSReset is Windows-only; other platforms keep the normal reset.
|
||||
func (p *prog) skipInitialDNSReset() bool { return false }
|
||||
|
||||
// exemptVPNDNSServers is a no-op on unsupported platforms.
|
||||
func (p *prog) exemptVPNDNSServers(exemptions []vpnDNSExemption) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ensurePFAnchorActive is a no-op on unsupported platforms.
|
||||
func (p *prog) ensurePFAnchorActive() bool {
|
||||
return false
|
||||
func (p *prog) ensurePFAnchorActive() pfAnchorCheckResult {
|
||||
return pfAnchorCheckSkipped
|
||||
}
|
||||
|
||||
// checkTunnelInterfaceChanges is a no-op on unsupported platforms.
|
||||
@@ -32,6 +36,10 @@ func (p *prog) checkTunnelInterfaceChanges() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (p *prog) dnsInterceptIgnoredChangeReconcileDue(time.Time) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// scheduleDelayedRechecks is a no-op on unsupported platforms.
|
||||
func (p *prog) scheduleDelayedRechecks() {}
|
||||
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
package cli
|
||||
|
||||
import "github.com/Control-D-Inc/ctrld"
|
||||
|
||||
var initializeOsResolver = ctrld.InitializeOsResolver
|
||||
|
||||
func (p *prog) refreshDNSAfterVPNSettle(reason string) (routes, domainlessServers, exemptions int) {
|
||||
mainLog.Load().Info().Msgf("DNS intercept: refreshing OS/VPN DNS route state after VPN settle (%s)", reason)
|
||||
ns := initializeOsResolver(true)
|
||||
mainLog.Load().Debug().Msgf("DNS intercept: post-settle OS resolver nameservers: %v", ns)
|
||||
|
||||
if p.vpnDNS == nil {
|
||||
mainLog.Load().Debug().Msg("DNS intercept: post-settle VPN DNS route refresh skipped — manager unavailable")
|
||||
return 0, 0, 0
|
||||
}
|
||||
|
||||
routes, domainlessServers, exemptions = p.vpnDNS.RefreshRoutesOnly()
|
||||
mainLog.Load().Info().Msgf("DNS intercept: post-settle VPN DNS route refresh completed — %d routes, %d domainless servers, %d exemptions",
|
||||
routes, domainlessServers, exemptions)
|
||||
return routes, domainlessServers, exemptions
|
||||
}
|
||||
|
||||
func vpnDNSExemptionsEqual(a, b []vpnDNSExemption) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
seen := make(map[vpnDNSExemption]int, len(a))
|
||||
for _, ex := range a {
|
||||
seen[ex]++
|
||||
}
|
||||
for _, ex := range b {
|
||||
if seen[ex] == 0 {
|
||||
return false
|
||||
}
|
||||
seen[ex]--
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
func TestRefreshDNSAfterVPNSettleRefreshesOSResolverAndVPNRoutes(t *testing.T) {
|
||||
oldInitialize := initializeOsResolver
|
||||
defer func() { initializeOsResolver = oldInitialize }()
|
||||
|
||||
var initialized []bool
|
||||
initializeOsResolver = func(force bool) []string {
|
||||
initialized = append(initialized, force)
|
||||
return []string{"10.102.26.10:53"}
|
||||
}
|
||||
|
||||
var exemptionUpdates [][]vpnDNSExemption
|
||||
p := &prog{}
|
||||
p.vpnDNS = newVPNDNSManager(func(exemptions []vpnDNSExemption) error {
|
||||
exemptionUpdates = append(exemptionUpdates, append([]vpnDNSExemption{}, exemptions...))
|
||||
return nil
|
||||
})
|
||||
p.vpnDNS.discoverVPNDNS = func(context.Context) []ctrld.VPNDNSConfig {
|
||||
return []ctrld.VPNDNSConfig{{
|
||||
InterfaceName: "utun4",
|
||||
Servers: []string{"10.102.26.10"},
|
||||
Domains: []string{"bmwgroup.net"},
|
||||
}}
|
||||
}
|
||||
|
||||
routes, domainlessServers, exemptions := p.refreshDNSAfterVPNSettle("test")
|
||||
|
||||
if routes != 1 || domainlessServers != 0 || exemptions != 1 {
|
||||
t.Fatalf("expected 1 route, 0 domainless servers, 1 exemption, got routes=%d domainless=%d exemptions=%d",
|
||||
routes, domainlessServers, exemptions)
|
||||
}
|
||||
if len(initialized) != 1 || !initialized[0] {
|
||||
t.Fatalf("expected forced OS resolver refresh once, got %v", initialized)
|
||||
}
|
||||
if got := p.vpnDNS.UpstreamForDomain("jira.cc.bmwgroup.net."); len(got) != 1 || got[0] != "10.102.26.10" {
|
||||
t.Fatalf("expected refreshed VPN DNS route, got %v", got)
|
||||
}
|
||||
if len(exemptionUpdates) != 1 || len(exemptionUpdates[0]) != 1 || exemptionUpdates[0][0].Server != "10.102.26.10" {
|
||||
t.Fatalf("expected one serialized pf exemption update for the late VPN DNS server, got %+v", exemptionUpdates)
|
||||
}
|
||||
|
||||
p.refreshDNSAfterVPNSettle("test-repeat")
|
||||
if len(exemptionUpdates) != 1 {
|
||||
t.Fatalf("unchanged post-settle VPN DNS state rewrote pf: %+v", exemptionUpdates)
|
||||
}
|
||||
}
|
||||
+1638
-228
File diff suppressed because it is too large
Load Diff
+261
-121
@@ -130,13 +130,7 @@ func (p *prog) serveDNS(listenerNum string) error {
|
||||
// signal the prober and respond NXDOMAIN. Used by both macOS pf probes
|
||||
// (_pf-probe-*) and Windows NRPT probes (_nrpt-probe-*) to verify that
|
||||
// DNS interception is actually routing queries to ctrld's listener.
|
||||
if probeID, ok := p.pfProbeExpected.Load().(string); ok && probeID != "" && domain == probeID {
|
||||
if chPtr, ok := p.pfProbeCh.Load().(*chan struct{}); ok && chPtr != nil {
|
||||
select {
|
||||
case *chPtr <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
if p.signalInterceptProbe(domain) {
|
||||
answer := new(dns.Msg)
|
||||
answer.SetRcode(m, dns.RcodeNameError) // NXDOMAIN
|
||||
_ = w.WriteMsg(answer)
|
||||
@@ -548,6 +542,7 @@ func (p *prog) proxy(ctx context.Context, req *proxyRequest) *proxyResponse {
|
||||
if vpnServers := p.vpnDNS.UpstreamForDomain(domain); len(vpnServers) > 0 {
|
||||
ctrld.Log(ctx, mainLog.Load().Debug(), "VPN DNS route matched for domain %s, using servers: %v", domain, vpnServers)
|
||||
|
||||
var gotTransportFailure bool
|
||||
for _, server := range vpnServers {
|
||||
upstreamConfig := p.vpnDNS.upstreamConfigFor(server)
|
||||
ctrld.Log(ctx, mainLog.Load().Debug(), "Querying VPN DNS server: %s", server)
|
||||
@@ -561,6 +556,7 @@ func (p *prog) proxy(ctx context.Context, req *proxyRequest) *proxyResponse {
|
||||
answer, err := dnsResolver.Resolve(resolveCtx, req.msg)
|
||||
cancel()
|
||||
if answer != nil {
|
||||
p.vpnDNS.VPNDNSReachable()
|
||||
ctrld.Log(ctx, mainLog.Load().Debug(), "VPN DNS query successful")
|
||||
if p.cache != nil {
|
||||
ttl := 60 * time.Second
|
||||
@@ -573,9 +569,22 @@ func (p *prog) proxy(ctx context.Context, req *proxyRequest) *proxyResponse {
|
||||
}
|
||||
return &proxyResponse{answer: answer}
|
||||
}
|
||||
gotTransportFailure = true
|
||||
ctrld.Log(ctx, mainLog.Load().Debug().Err(err), "VPN DNS server %s failed", server)
|
||||
}
|
||||
|
||||
// Explicit VPN DNS routes are authoritative for their suffix. If all
|
||||
// routed servers fail at the transport layer while Windows is serving
|
||||
// retained VPN DNS state, fail closed instead of leaking VPN/internal
|
||||
// names to normal upstreams.
|
||||
if gotTransportFailure && p.vpnDNS.ShouldFailClosedAfterVPNDNSTransportFailure(domain, vpnServers) {
|
||||
ctrld.Log(ctx, mainLog.Load().Debug(),
|
||||
"All VPN DNS servers had transport failures for %s; returning SERVFAIL while retained VPN DNS state is active", domain)
|
||||
answer := new(dns.Msg)
|
||||
answer.SetRcode(req.msg, dns.RcodeServerFailure)
|
||||
return &proxyResponse{answer: answer}
|
||||
}
|
||||
|
||||
ctrld.Log(ctx, mainLog.Load().Debug(), "All VPN DNS servers failed, falling back to normal upstreams")
|
||||
}
|
||||
}
|
||||
@@ -589,13 +598,15 @@ func (p *prog) proxy(ctx context.Context, req *proxyRequest) *proxyResponse {
|
||||
// polluting captive portal / DHCP flows.
|
||||
if dnsIntercept && p.vpnDNS != nil && req.ufr.matched &&
|
||||
len(upstreams) > 0 && upstreams[0] == upstreamOS &&
|
||||
len(req.msg.Question) > 0 && !p.isAdDomainQuery(req.msg) {
|
||||
len(req.msg.Question) > 0 {
|
||||
if dlServers := p.vpnDNS.DomainlessServers(); len(dlServers) > 0 {
|
||||
domain := req.msg.Question[0].Name
|
||||
ctrld.Log(ctx, mainLog.Load().Debug(),
|
||||
"Split-rule query %s going to upstream.os, trying %d domain-less VPN DNS servers first: %v",
|
||||
domain, len(dlServers), dlServers)
|
||||
|
||||
var gotDNSAnswer bool
|
||||
var gotTransportFailure bool
|
||||
for _, server := range dlServers {
|
||||
upstreamCfg := p.vpnDNS.upstreamConfigFor(server)
|
||||
ctrld.Log(ctx, mainLog.Load().Debug(), "Querying domain-less VPN DNS server: %s", server)
|
||||
@@ -608,6 +619,10 @@ func (p *prog) proxy(ctx context.Context, req *proxyRequest) *proxyResponse {
|
||||
resolveCtx, cancel := upstreamCfg.Context(ctx)
|
||||
answer, err := dnsResolver.Resolve(resolveCtx, req.msg)
|
||||
cancel()
|
||||
if answer != nil {
|
||||
gotDNSAnswer = true
|
||||
p.vpnDNS.VPNDNSReachable()
|
||||
}
|
||||
if answer != nil && answer.Rcode == dns.RcodeSuccess {
|
||||
ctrld.Log(ctx, mainLog.Load().Debug(),
|
||||
"Domain-less VPN DNS server %s answered %s successfully", server, domain)
|
||||
@@ -618,10 +633,25 @@ func (p *prog) proxy(ctx context.Context, req *proxyRequest) *proxyResponse {
|
||||
"Domain-less VPN DNS server %s returned %s for %s, trying next",
|
||||
server, dns.RcodeToString[answer.Rcode], domain)
|
||||
} else {
|
||||
gotTransportFailure = true
|
||||
ctrld.Log(ctx, mainLog.Load().Debug().Err(err),
|
||||
"Domain-less VPN DNS server %s failed for %s", server, domain)
|
||||
}
|
||||
}
|
||||
|
||||
// If every domainless VPN DNS attempt failed before receiving a DNS
|
||||
// packet while Windows is serving retained VPN DNS state, fail closed
|
||||
// instead of asking LAN/public DNS about internal split-rule names and
|
||||
// caching false negatives. Reachable negative DNS responses still fall
|
||||
// through to the old OS fallback behavior below.
|
||||
if !gotDNSAnswer && gotTransportFailure && p.vpnDNS.ShouldFailClosedAfterVPNDNSTransportFailure(domain, dlServers) {
|
||||
ctrld.Log(ctx, mainLog.Load().Debug(),
|
||||
"All domain-less VPN DNS servers had transport failures for %s; returning SERVFAIL while retained VPN DNS state is active", domain)
|
||||
answer := new(dns.Msg)
|
||||
answer.SetRcode(req.msg, dns.RcodeServerFailure)
|
||||
return &proxyResponse{answer: answer}
|
||||
}
|
||||
|
||||
ctrld.Log(ctx, mainLog.Load().Debug(),
|
||||
"All domain-less VPN DNS servers failed for %s, falling back to OS resolver", domain)
|
||||
}
|
||||
@@ -700,6 +730,17 @@ func (p *prog) proxy(ctx context.Context, req *proxyRequest) *proxyResponse {
|
||||
}
|
||||
continue
|
||||
}
|
||||
// Reject an answer whose question does not match the request before it
|
||||
// can be served or cached. A mismatched question means the upstream
|
||||
// answered a different name/type than asked; caching it would poison
|
||||
// the shared cache with wrong-domain records for the requested name.
|
||||
// See github.com/Control-D-Inc/ctrld/issues/322.
|
||||
if !sameQuestion(req.msg, answer) {
|
||||
ctrld.Log(ctx, mainLog.Load().Debug(),
|
||||
"discarding answer from %s: question mismatch (asked %q, got %q)",
|
||||
upstreams[n], questionString(req.msg), questionString(answer))
|
||||
continue
|
||||
}
|
||||
// We are doing LAN/PTR lookup using private resolver, so always process next one.
|
||||
// Except for the last, we want to send response instead of saying all upstream failed.
|
||||
if answer.Rcode != dns.RcodeSuccess && isLanOrPtrQuery && n != len(upstreamConfigs)-1 {
|
||||
@@ -863,6 +904,33 @@ func containRcode(rcodes []int, rcode int) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// sameQuestion reports whether the upstream answer echoes the request's
|
||||
// question. A well-behaved resolver always copies the question section from
|
||||
// the query (RFC 1035 section 4.1.2); names are compared case-insensitively
|
||||
// because DNS names are case-insensitive. A mismatch means the upstream
|
||||
// answered a different name/type than asked - malformed or malicious - and the
|
||||
// answer must not be served or cached, or it would poison the shared cache with
|
||||
// wrong-domain records. See github.com/Control-D-Inc/ctrld/issues/322.
|
||||
func sameQuestion(req, answer *dns.Msg) bool {
|
||||
if req == nil || answer == nil {
|
||||
return false
|
||||
}
|
||||
if len(req.Question) == 0 || len(answer.Question) == 0 {
|
||||
return false
|
||||
}
|
||||
rq, aq := req.Question[0], answer.Question[0]
|
||||
return rq.Qtype == aq.Qtype && rq.Qclass == aq.Qclass && strings.EqualFold(rq.Name, aq.Name)
|
||||
}
|
||||
|
||||
// questionString renders a message's first question as "name/type" for logging.
|
||||
func questionString(msg *dns.Msg) string {
|
||||
if msg == nil || len(msg.Question) == 0 {
|
||||
return "<none>"
|
||||
}
|
||||
q := msg.Question[0]
|
||||
return q.Name + "/" + dns.TypeToString[q.Qtype]
|
||||
}
|
||||
|
||||
func setCachedAnswerTTL(answer *dns.Msg, now, expiredTime time.Time) {
|
||||
ttlSecs := expiredTime.Sub(now).Seconds()
|
||||
if ttlSecs < 0 {
|
||||
@@ -1122,7 +1190,7 @@ func (p *prog) doSelfUninstall(answer *dns.Msg) {
|
||||
Version: rootCmd.Version,
|
||||
Metadata: ctrld.SystemMetadataRuntime(context.Background()),
|
||||
}
|
||||
_, err := controld.FetchResolverConfig(req, cdDev)
|
||||
_, err := controld.FetchResolverConfig(context.Background(), req, cdDev)
|
||||
logger.Debug().Msg("maximum number of refused queries reached, checking device status")
|
||||
selfUninstallCheck(err, p, logger)
|
||||
|
||||
@@ -1263,7 +1331,8 @@ func isPrivatePtrLookup(m *dns.Msg) bool {
|
||||
return addr.IsPrivate() ||
|
||||
addr.IsLoopback() ||
|
||||
addr.IsLinkLocalUnicast() ||
|
||||
tsaddr.CGNATRange().Contains(addr)
|
||||
tsaddr.CGNATRange().Contains(addr) ||
|
||||
isServiceContinuityAddr(addr)
|
||||
}
|
||||
}
|
||||
return false
|
||||
@@ -1301,6 +1370,20 @@ func isLanHostname(name string) bool {
|
||||
strings.HasSuffix(name, ".local")
|
||||
}
|
||||
|
||||
// ipv4ServiceContinuityPrefix is the RFC 7335 IPv4 Service Continuity Prefix
|
||||
// (192.0.0.0/29), used by the CLAT in 464XLAT/DS-Lite transition setups. On such
|
||||
// networks (common on IPv6-only cellular carriers and iPhone hotspots) the local
|
||||
// machine's DNS queries reach ctrld with a source in this range (e.g. 192.0.0.2),
|
||||
// so they must be treated as local, not WAN. Go's netip.IsPrivate does not cover
|
||||
// this range — the same reason the CGNAT range is special-cased below. See #552.
|
||||
var ipv4ServiceContinuityPrefix = netip.MustParsePrefix("192.0.0.0/29")
|
||||
|
||||
// isServiceContinuityAddr reports whether ip is in the RFC 7335 IPv4 Service
|
||||
// Continuity Prefix (464XLAT/DS-Lite CLAT).
|
||||
func isServiceContinuityAddr(ip netip.Addr) bool {
|
||||
return ipv4ServiceContinuityPrefix.Contains(ip)
|
||||
}
|
||||
|
||||
// isWanClient reports whether the input is a WAN address.
|
||||
func isWanClient(na net.Addr) bool {
|
||||
var ip netip.Addr
|
||||
@@ -1311,7 +1394,8 @@ func isWanClient(na net.Addr) bool {
|
||||
!ip.IsPrivate() &&
|
||||
!ip.IsLinkLocalUnicast() &&
|
||||
!ip.IsLinkLocalMulticast() &&
|
||||
!tsaddr.CGNATRange().Contains(ip)
|
||||
!tsaddr.CGNATRange().Contains(ip) &&
|
||||
!isServiceContinuityAddr(ip)
|
||||
}
|
||||
|
||||
// isIPv6LoopbackListener reports whether the listener address is [::1].
|
||||
@@ -1451,64 +1535,13 @@ func (p *prog) monitorNetworkChanges() error {
|
||||
mainLog.Load().Debug().Msg("Ignoring interface change - no valid interfaces affected")
|
||||
// check if the default IPs are still on an interface that is up
|
||||
ValidateDefaultLocalIPsFromDelta(delta.New)
|
||||
// Even minor interface changes can trigger macOS pf reloads — verify anchor.
|
||||
// We check immediately AND schedule delayed re-checks (2s + 4s) to catch
|
||||
// programs like Windscribe that modify pf rules and DNS settings
|
||||
// asynchronously after the network change event fires.
|
||||
// Minor interface changes can still accompany pf/WFP or VPN DNS changes.
|
||||
// On macOS, bound the immediate full reconciliation so link-local-only
|
||||
// notification storms do not run pfctl/scutil work for every event.
|
||||
// Windows keeps the existing immediate behavior. Tunnel changes always
|
||||
// bypass the macOS limit, and delayed checks provide a trailing refresh.
|
||||
if dnsIntercept && p.dnsInterceptState != nil {
|
||||
if !p.pfStabilizing.Load() {
|
||||
p.ensurePFAnchorActive()
|
||||
}
|
||||
// Check tunnel interfaces unconditionally — it decides internally
|
||||
// whether to enter stabilization or rebuild immediately.
|
||||
p.checkTunnelInterfaceChanges()
|
||||
// Schedule delayed re-checks to catch async VPN teardown changes.
|
||||
// These also refresh the OS resolver and VPN DNS routes.
|
||||
p.scheduleDelayedRechecks()
|
||||
|
||||
// Detect interface appearance/disappearance — hypervisors (Parallels,
|
||||
// VMware, VirtualBox) reload pf when creating/destroying virtual network
|
||||
// interfaces, which can corrupt pf's internal translation state. The rdr
|
||||
// rules survive in text form (watchdog says "intact") but stop evaluating.
|
||||
// Spawn an async monitor that probes pf interception with backoff and
|
||||
// forces a full pf reload if broken.
|
||||
if delta.Old != nil {
|
||||
interfaceChanged := false
|
||||
var changedIface string
|
||||
for ifaceName := range delta.Old.Interface {
|
||||
if ifaceName == "lo0" {
|
||||
continue
|
||||
}
|
||||
if _, exists := delta.New.Interface[ifaceName]; !exists {
|
||||
interfaceChanged = true
|
||||
changedIface = ifaceName
|
||||
break
|
||||
}
|
||||
}
|
||||
if !interfaceChanged {
|
||||
for ifaceName := range delta.New.Interface {
|
||||
if ifaceName == "lo0" {
|
||||
continue
|
||||
}
|
||||
if _, exists := delta.Old.Interface[ifaceName]; !exists {
|
||||
interfaceChanged = true
|
||||
changedIface = ifaceName
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
if interfaceChanged {
|
||||
mainLog.Load().Info().Str("interface", changedIface).
|
||||
Msg("DNS intercept: interface appeared/disappeared — starting interception probe monitor")
|
||||
go p.pfInterceptMonitor()
|
||||
}
|
||||
}
|
||||
}
|
||||
// Refresh VPN DNS on tunnel interface changes (e.g., Tailscale connect/disconnect)
|
||||
// even though the physical interface didn't change. Runs after tunnel checks
|
||||
// so the pf anchor rebuild includes current VPN DNS exemptions.
|
||||
if dnsIntercept && p.vpnDNS != nil {
|
||||
p.vpnDNS.Refresh(true)
|
||||
p.handleDNSInterceptIgnoredNetworkChange(delta, time.Now())
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -1610,6 +1643,76 @@ func (p *prog) monitorNetworkChanges() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// handleDNSInterceptIgnoredNetworkChange runs the DNS-intercept work for a
|
||||
// network delta that did not affect a usable interface. Keeping this path in a
|
||||
// method lets tests exercise the callback wiring with synthetic deltas.
|
||||
func (p *prog) handleDNSInterceptIgnoredNetworkChange(delta *netmon.ChangeDelta, now time.Time) {
|
||||
reconcileNow := false
|
||||
// Stabilization owns PF repair. Do not consume the next leading-edge slot
|
||||
// until an ignored delta can actually perform the corresponding PF check.
|
||||
if !p.pfStabilizing.Load() {
|
||||
reconcileNow = p.dnsInterceptIgnoredChangeReconcileDue(now)
|
||||
if reconcileNow {
|
||||
p.ensurePFAnchorActive()
|
||||
}
|
||||
}
|
||||
|
||||
// Check tunnel interfaces unconditionally — it decides internally whether
|
||||
// to enter stabilization or rebuild immediately.
|
||||
tunnelChanged := p.checkTunnelInterfaceChanges()
|
||||
// Schedule delayed re-checks to catch async VPN teardown changes. These also
|
||||
// refresh the OS resolver and VPN DNS routes.
|
||||
p.scheduleDelayedRechecks()
|
||||
|
||||
// Detect interface appearance/disappearance — hypervisors (Parallels,
|
||||
// VMware, VirtualBox) reload pf when creating/destroying virtual network
|
||||
// interfaces, which can corrupt pf's internal translation state. The rdr
|
||||
// rules survive in text form (watchdog says "intact") but stop evaluating.
|
||||
// Spawn an async monitor that probes pf interception with backoff and forces
|
||||
// a full pf reload if broken.
|
||||
if delta.Old != nil {
|
||||
interfaceChanged := false
|
||||
var changedIface string
|
||||
for ifaceName := range delta.Old.Interface {
|
||||
if ifaceName == "lo0" {
|
||||
continue
|
||||
}
|
||||
if _, exists := delta.New.Interface[ifaceName]; !exists {
|
||||
interfaceChanged = true
|
||||
changedIface = ifaceName
|
||||
break
|
||||
}
|
||||
}
|
||||
if !interfaceChanged {
|
||||
for ifaceName := range delta.New.Interface {
|
||||
if ifaceName == "lo0" {
|
||||
continue
|
||||
}
|
||||
if _, exists := delta.Old.Interface[ifaceName]; !exists {
|
||||
interfaceChanged = true
|
||||
changedIface = ifaceName
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
if interfaceChanged {
|
||||
mainLog.Load().Info().Str("interface", changedIface).
|
||||
Msg("DNS intercept: interface appeared/disappeared — starting interception probe monitor")
|
||||
go p.pfInterceptMonitor()
|
||||
}
|
||||
}
|
||||
|
||||
// Refresh VPN DNS immediately for real tunnel changes even when the periodic
|
||||
// ignored-change reconciliation is currently rate-limited - but not while
|
||||
// stabilization owns pf. A refresh rebuilds the anchor, and these deltas arrive
|
||||
// exactly when a VPN is bringing its own ruleset up, which is the collision
|
||||
// stabilization is there to prevent. checkTunnelInterfaceChanges keeps the
|
||||
// observation pending, so the transition is retried rather than dropped.
|
||||
if p.vpnDNS != nil && (reconcileNow || tunnelChanged) && !p.pfStabilizing.Load() {
|
||||
p.vpnDNS.Refresh(true)
|
||||
}
|
||||
}
|
||||
|
||||
// interfaceStatesEqual compares two interface states
|
||||
func interfaceStatesEqual(a, b *netmon.Interface) bool {
|
||||
if a == nil || b == nil {
|
||||
@@ -1697,6 +1800,15 @@ func (p *prog) checkUpstreamOnce(upstream string, uc *ctrld.UpstreamConfig) erro
|
||||
mainLog.Load().Debug().Err(err).Msgf("Upstream %s check failed after %v (WFP loopback protect active)", upstream, duration)
|
||||
return errOsHealthcheckSuppressed
|
||||
}
|
||||
// A no-route/network-unreachable failure means the endpoint's address
|
||||
// family is available locally but unroutable (e.g. an IPv6 DoH endpoint
|
||||
// while IPv6 is up but has no route). These repeat until the route
|
||||
// returns and are handled by bounded backoff in the recovery loop, so
|
||||
// keep them at debug to avoid sustained error-log spam.
|
||||
if ctrldnet.IsUnreachable(err) {
|
||||
mainLog.Load().Debug().Err(err).Msgf("Upstream %s check failed after %v (network unreachable)", upstream, duration)
|
||||
return err
|
||||
}
|
||||
mainLog.Load().Error().Err(err).Msgf("Upstream %s check failed after %v", upstream, duration)
|
||||
return err
|
||||
}
|
||||
@@ -1740,24 +1852,13 @@ func (p *prog) debounceRecovery() {
|
||||
func (p *prog) handleRecovery(reason RecoveryReason) {
|
||||
mainLog.Load().Debug().Msg("Starting recovery process: removing DNS settings")
|
||||
|
||||
// For network changes, cancel any existing recovery check because the network state has changed.
|
||||
recoveryCtx, gen, interceptRecovery, ok := p.beginRecovery(reason)
|
||||
if !ok {
|
||||
mainLog.Load().Debug().Msg("Upstream recovery already in progress; skipping duplicate trigger")
|
||||
return
|
||||
}
|
||||
if reason == RecoveryReasonNetworkChange {
|
||||
p.recoveryCancelMu.Lock()
|
||||
if p.recoveryCancel != nil {
|
||||
mainLog.Load().Debug().Msg("Cancelling existing recovery check (network change)")
|
||||
p.recoveryCancel()
|
||||
p.recoveryCancel = nil
|
||||
}
|
||||
p.recoveryCancelMu.Unlock()
|
||||
} else {
|
||||
// For upstream failures, if a recovery is already in progress, do nothing new.
|
||||
p.recoveryCancelMu.Lock()
|
||||
if p.recoveryCancel != nil {
|
||||
mainLog.Load().Debug().Msg("Upstream recovery already in progress; skipping duplicate trigger")
|
||||
p.recoveryCancelMu.Unlock()
|
||||
return
|
||||
}
|
||||
p.recoveryCancelMu.Unlock()
|
||||
mainLog.Load().Debug().Msg("Network change recovery now owns shared recovery state")
|
||||
}
|
||||
|
||||
// For network changes, force-reset all upstream transports synchronously.
|
||||
@@ -1775,32 +1876,28 @@ func (p *prog) handleRecovery(reason RecoveryReason) {
|
||||
mainLog.Load().Info().Msg("Force-reset upstream transports for network change recovery")
|
||||
}
|
||||
|
||||
// Create a new recovery context without a fixed timeout.
|
||||
p.recoveryCancelMu.Lock()
|
||||
recoveryCtx, cancel := context.WithCancel(context.Background())
|
||||
p.recoveryCancel = cancel
|
||||
p.recoveryCancelMu.Unlock()
|
||||
|
||||
// set recoveryRunning to true to prevent watchdogs from putting the listener back on the interface
|
||||
p.recoveryRunning.Store(true)
|
||||
|
||||
// In DNS intercept mode, don't tear down WFP/pf filters.
|
||||
// Instead, enable recovery bypass so proxy() forwards queries to
|
||||
// the OS/DHCP resolver. This handles captive portal authentication
|
||||
// without the overhead of filter teardown/rebuild.
|
||||
if dnsIntercept && p.dnsInterceptState != nil {
|
||||
p.recoveryBypass.Store(true)
|
||||
if interceptRecovery {
|
||||
mainLog.Load().Info().Msg("DNS intercept recovery: enabling DHCP bypass (filters stay active)")
|
||||
|
||||
// Reinitialize OS resolver to discover DHCP servers on the new network.
|
||||
mainLog.Load().Debug().Msg("DNS intercept recovery: discovering DHCP nameservers")
|
||||
dhcpServers := ctrld.InitializeOsResolver(true)
|
||||
dhcpServers, systemNameservers := ctrld.InitializeOsResolverWithSystemNameservers(true)
|
||||
if len(dhcpServers) == 0 {
|
||||
mainLog.Load().Warn().Msg("DNS intercept recovery: no DHCP nameservers found")
|
||||
} else {
|
||||
mainLog.Load().Info().Msgf("DNS intercept recovery: found DHCP nameservers: %v", dhcpServers)
|
||||
}
|
||||
|
||||
// If the new network provides no usable IPv4 DNS (e.g. IPv6-only
|
||||
// tethering with 464XLAT), macOS cannot emit DNS queries at all and
|
||||
// pf has nothing to intercept. Ensure a loopback DNS target exists
|
||||
// so the OS keeps sending queries to ctrld's listener (issue #533).
|
||||
ensureInterceptDNSTargetFn(p, systemNameservers)
|
||||
|
||||
// Exempt DHCP nameservers from intercept filters so the OS resolver
|
||||
// can actually reach them on port 53.
|
||||
if len(dhcpServers) > 0 {
|
||||
@@ -1842,21 +1939,20 @@ func (p *prog) handleRecovery(reason RecoveryReason) {
|
||||
recovered, err := p.waitForUpstreamRecovery(recoveryCtx, upstreams)
|
||||
if err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("Recovery canceled; DNS settings remain removed")
|
||||
p.recoveryCancelMu.Lock()
|
||||
p.recoveryCancel = nil
|
||||
p.recoveryCancelMu.Unlock()
|
||||
p.recoveryCanceledCleanup(gen)
|
||||
return
|
||||
}
|
||||
if !p.recoveryOwnsState(gen) {
|
||||
mainLog.Load().Debug().Msgf("Recovery generation %d was superseded after upstream success; skipping stale completion", gen)
|
||||
return
|
||||
}
|
||||
mainLog.Load().Info().Msgf("Upstream %q recovered; re-applying DNS settings", recovered)
|
||||
|
||||
// reset the upstream failure count and down state
|
||||
// Reset the upstream failure count and down state while this generation
|
||||
// still owns recovery completion.
|
||||
p.um.reset(recovered)
|
||||
|
||||
// In DNS intercept mode, just disable the bypass — filters are still active.
|
||||
if dnsIntercept && p.dnsInterceptState != nil {
|
||||
p.recoveryBypass.Store(false)
|
||||
mainLog.Load().Info().Msg("DNS intercept recovery complete: disabling DHCP bypass, resuming normal flow")
|
||||
|
||||
if interceptRecovery {
|
||||
// Refresh VPN DNS routes in case VPN state changed during recovery.
|
||||
if p.vpnDNS != nil {
|
||||
p.vpnDNS.Refresh(true)
|
||||
@@ -1871,11 +1967,15 @@ func (p *prog) handleRecovery(reason RecoveryReason) {
|
||||
mainLog.Load().Info().Msgf("Reinitialized OS resolver with nameservers: %v", ns)
|
||||
}
|
||||
}
|
||||
|
||||
p.recoveryRunning.Store(false)
|
||||
} else {
|
||||
// For network changes we also reinitialize the OS resolver.
|
||||
if reason == RecoveryReasonNetworkChange {
|
||||
var systemNameservers []string
|
||||
if dnsIntercept {
|
||||
// Intercept was requested but no interceptor was active when recovery
|
||||
// began. Rediscover on every recovery reason before retrying setDNS;
|
||||
// passing nil could make a successful retry install a loopback target
|
||||
// on a healthy DHCP network.
|
||||
systemNameservers = systemNameserversForInterceptRetry()
|
||||
} else if reason == RecoveryReasonNetworkChange {
|
||||
ns := ctrld.InitializeOsResolver(true)
|
||||
if len(ns) == 0 {
|
||||
mainLog.Load().Warn().Msg("No nameservers found for OS resolver during network-change recovery; using existing values")
|
||||
@@ -1885,22 +1985,25 @@ func (p *prog) handleRecovery(reason RecoveryReason) {
|
||||
}
|
||||
|
||||
// Apply our DNS settings back and log the interface state.
|
||||
p.setDNS()
|
||||
p.setDNS(systemNameservers)
|
||||
p.logInterfacesState()
|
||||
|
||||
// allow watchdogs to put the listener back on the interface if its changed for any reason
|
||||
p.recoveryRunning.Store(false)
|
||||
}
|
||||
|
||||
// Clear the recovery cancellation for a clean slate.
|
||||
p.recoveryCancelMu.Lock()
|
||||
p.recoveryCancel = nil
|
||||
p.recoveryCancelMu.Unlock()
|
||||
if !p.completeRecovery(gen) {
|
||||
mainLog.Load().Debug().Msgf("Recovery generation %d was superseded during completion; preserving successor state", gen)
|
||||
return
|
||||
}
|
||||
if interceptRecovery {
|
||||
mainLog.Load().Info().Msg("DNS intercept recovery complete: disabling DHCP bypass, resuming normal flow")
|
||||
}
|
||||
}
|
||||
|
||||
// waitForUpstreamRecovery checks the provided upstreams concurrently until one recovers.
|
||||
// It returns the name of the recovered upstream or an error if the check times out.
|
||||
func (p *prog) waitForUpstreamRecovery(ctx context.Context, upstreams map[string]*ctrld.UpstreamConfig) (string, error) {
|
||||
recoveryCtx, cancel := context.WithCancel(ctx)
|
||||
defer cancel()
|
||||
|
||||
recoveredCh := make(chan string, 1)
|
||||
var wg sync.WaitGroup
|
||||
|
||||
@@ -1912,9 +2015,10 @@ func (p *prog) waitForUpstreamRecovery(ctx context.Context, upstreams map[string
|
||||
defer wg.Done()
|
||||
mainLog.Load().Debug().Msgf("Starting recovery check loop for upstream: %s", name)
|
||||
attempts := 0
|
||||
unreachableStreak := 0
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
case <-recoveryCtx.Done():
|
||||
mainLog.Load().Debug().Msgf("Context canceled for upstream %s", name)
|
||||
return
|
||||
default:
|
||||
@@ -1926,13 +2030,30 @@ func (p *prog) waitForUpstreamRecovery(ctx context.Context, upstreams map[string
|
||||
select {
|
||||
case recoveredCh <- name:
|
||||
mainLog.Load().Debug().Msgf("Sent recovery notification for upstream %s", name)
|
||||
cancel()
|
||||
default:
|
||||
mainLog.Load().Debug().Msg("Recovery channel full, another upstream already recovered")
|
||||
}
|
||||
return
|
||||
}
|
||||
mainLog.Load().Debug().Msgf("Upstream %s check failed, sleeping before retry", name)
|
||||
time.Sleep(checkUpstreamBackoffSleep)
|
||||
// Back off the retry cadence for an unroutable endpoint so a
|
||||
// host with IPv6 up but no route to the IPv6 DoH endpoint does
|
||||
// not re-bootstrap/re-check every checkUpstreamBackoffSleep and
|
||||
// spam the log. The backoff is bounded (checkUpstreamUnreachableBackoffMax)
|
||||
// so the endpoint is still re-probed and recovers when the route
|
||||
// returns; any other failure resets to the base cadence.
|
||||
sleep := checkUpstreamBackoffSleep
|
||||
if ctrldnet.IsUnreachable(err) {
|
||||
unreachableStreak++
|
||||
sleep = unreachableRecoveryBackoff(unreachableStreak)
|
||||
mainLog.Load().Debug().Msgf("Upstream %s unreachable (streak %d), backing off %s before retry", name, unreachableStreak, sleep)
|
||||
} else {
|
||||
unreachableStreak = 0
|
||||
mainLog.Load().Debug().Msgf("Upstream %s check failed, sleeping before retry", name)
|
||||
}
|
||||
if !sleepWithContext(recoveryCtx, sleep) {
|
||||
return
|
||||
}
|
||||
|
||||
// if this is the upstreamOS and it's the 3rd attempt (or multiple of 3),
|
||||
// we should try to reinit the OS resolver to ensure we can recover
|
||||
@@ -1952,7 +2073,15 @@ func (p *prog) waitForUpstreamRecovery(ctx context.Context, upstreams map[string
|
||||
|
||||
var recovered string
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return "", ctx.Err()
|
||||
default:
|
||||
}
|
||||
select {
|
||||
case recovered = <-recoveredCh:
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
case <-ctx.Done():
|
||||
return "", ctx.Err()
|
||||
}
|
||||
@@ -1960,6 +2089,17 @@ func (p *prog) waitForUpstreamRecovery(ctx context.Context, upstreams map[string
|
||||
return recovered, nil
|
||||
}
|
||||
|
||||
func sleepWithContext(ctx context.Context, d time.Duration) bool {
|
||||
timer := time.NewTimer(d)
|
||||
defer timer.Stop()
|
||||
select {
|
||||
case <-timer.C:
|
||||
return true
|
||||
case <-ctx.Done():
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// buildRecoveryUpstreams constructs the map of upstream configurations to test.
|
||||
// For OS failures we supply the manual OS resolver upstream configuration.
|
||||
// For network change or regular failure we use the upstreams defined in p.cfg (ignoring OS).
|
||||
|
||||
@@ -405,6 +405,8 @@ func Test_isPrivatePtrLookup(t *testing.T) {
|
||||
{"CGNAT", newDnsMsgPtr("100.66.27.28", t), true},
|
||||
{"Loopback", newDnsMsgPtr("127.0.0.1", t), true},
|
||||
{"Link Local Unicast", newDnsMsgPtr("fe80::69f6:e16e:8bdb:433f", t), true},
|
||||
// RFC 7335 IPv4 Service Continuity Prefix (464XLAT/DS-Lite CLAT), see #552.
|
||||
{"464XLAT CLAT host", newDnsMsgPtr("192.0.0.2", t), true},
|
||||
{"Public IP", newDnsMsgPtr("8.8.8.8", t), false},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
@@ -452,6 +454,11 @@ func Test_isWanClient(t *testing.T) {
|
||||
{"CGNAT", &net.UDPAddr{IP: net.ParseIP("100.66.27.28")}, false},
|
||||
{"Loopback", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")}, false},
|
||||
{"Link Local Unicast", &net.UDPAddr{IP: net.ParseIP("fe80::69f6:e16e:8bdb:433f")}, false},
|
||||
// RFC 7335 IPv4 Service Continuity Prefix (464XLAT/DS-Lite CLAT), see #552.
|
||||
{"464XLAT PLAT side", &net.UDPAddr{IP: net.ParseIP("192.0.0.1")}, false},
|
||||
{"464XLAT CLAT host", &net.UDPAddr{IP: net.ParseIP("192.0.0.2")}, false},
|
||||
// Outside the /29 but inside 192.0.0.0/24: still WAN (fix is scoped to /29).
|
||||
{"192.0.0.0/24 outside /29", &net.UDPAddr{IP: net.ParseIP("192.0.0.100")}, true},
|
||||
{"Public", &net.UDPAddr{IP: net.ParseIP("8.8.8.8")}, true},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
@@ -474,3 +481,33 @@ func Test_prog_queryFromSelf(t *testing.T) {
|
||||
p.queryFromSelf("foo")
|
||||
})
|
||||
}
|
||||
|
||||
func Test_sameQuestion(t *testing.T) {
|
||||
mk := func(name string, qtype uint16) *dns.Msg {
|
||||
m := new(dns.Msg)
|
||||
m.SetQuestion(name, qtype)
|
||||
return m
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
req *dns.Msg
|
||||
answer *dns.Msg
|
||||
want bool
|
||||
}{
|
||||
{"identical", mk("example.com.", dns.TypeA), mk("example.com.", dns.TypeA), true},
|
||||
{"case insensitive", mk("Example.COM.", dns.TypeA), mk("example.com.", dns.TypeA), true},
|
||||
{"different name", mk("victim.example.", dns.TypeA), mk("attacker.example.", dns.TypeA), false},
|
||||
{"different type", mk("example.com.", dns.TypeA), mk("example.com.", dns.TypeAAAA), false},
|
||||
{"nil req", nil, mk("example.com.", dns.TypeA), false},
|
||||
{"nil answer", mk("example.com.", dns.TypeA), nil, false},
|
||||
{"empty answer question", mk("example.com.", dns.TypeA), new(dns.Msg), false},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := sameQuestion(tc.req, tc.answer); got != tc.want {
|
||||
t.Errorf("sameQuestion() = %v, want %v", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
package cli
|
||||
|
||||
import "net"
|
||||
|
||||
// interceptDNSRdrTarget is the loopback address used as the macOS service
|
||||
// DNS value when ctrld's listener is NOT reachable at <listener IP>:53
|
||||
// directly (non-53 port, e.g. 127.0.0.1:5354 when mDNSResponder holds *:53).
|
||||
//
|
||||
// macOS resolvers always send DNS to port 53, so a direct-hit value is
|
||||
// impossible in that case; delivery must go through the pf rdr rule
|
||||
// ("rdr on lo0 ... to ! <listenerIP> port 53 -> <listenerIP> port <port>").
|
||||
// The value therefore must be a loopback address DIFFERENT from the listener
|
||||
// IP so the rdr's "! <listenerIP>" matches. Any 127/8 address routes via lo0
|
||||
// on macOS.
|
||||
const interceptDNSRdrTarget = "127.0.0.53"
|
||||
|
||||
// interceptDNSTargetValue returns the nameserver value to set on a DNS-less
|
||||
// macOS service so the OS emits DNS queries that reach ctrld, respecting the
|
||||
// configured listener. The listener IP/port derivation mirrors
|
||||
// buildPFAnchorRulesForTunnels so the value and the pf rules always agree.
|
||||
//
|
||||
// - listener on port 53: return the effective listener IP — queries hit the
|
||||
// listener directly, no pf dependency for this leg.
|
||||
// - listener on another port: return interceptDNSRdrTarget so the lo0 rdr
|
||||
// rule fires and rewrites to the real listener address.
|
||||
func (p *prog) interceptDNSTargetValue() string {
|
||||
listenerIP := "127.0.0.1"
|
||||
listenerPort := 53
|
||||
// FirstListener panics when no listener is configured; guard like the
|
||||
// startup paths do.
|
||||
if p.cfg != nil && len(p.cfg.Listener) > 0 {
|
||||
if lc := p.cfg.FirstListener(); lc != nil {
|
||||
if lc.IP != "" && lc.IP != "0.0.0.0" && lc.IP != "::" {
|
||||
listenerIP = lc.IP
|
||||
}
|
||||
if lc.Port != 0 {
|
||||
listenerPort = lc.Port
|
||||
}
|
||||
}
|
||||
}
|
||||
if listenerPort == 53 {
|
||||
return listenerIP
|
||||
}
|
||||
if listenerIP == interceptDNSRdrTarget {
|
||||
// Pathological config: the listener itself sits on the rdr target
|
||||
// address (with a non-53 port). Pick a different loopback so the
|
||||
// rdr's "! <listenerIP>" still matches.
|
||||
return "127.0.0.54"
|
||||
}
|
||||
return interceptDNSRdrTarget
|
||||
}
|
||||
|
||||
// hasIPv4DNS reports whether any of the given nameserver strings (bare IPs or
|
||||
// host:port) is an IPv4 address. Loopback counts: an existing local resolver
|
||||
// is treated conservatively as an intentional emittable DNS target; ctrld does
|
||||
// not probe or replace another resolver's ownership.
|
||||
func hasIPv4DNS(nameservers []string) bool {
|
||||
for _, s := range nameservers {
|
||||
host := s
|
||||
if h, _, err := net.SplitHostPort(s); err == nil {
|
||||
host = h
|
||||
}
|
||||
ip := net.ParseIP(host)
|
||||
if ip == nil {
|
||||
continue
|
||||
}
|
||||
if ip.To4() != nil {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// needsInterceptDNSTarget reports whether the OS is left without any usable
|
||||
// IPv4 DNS target: neither the default-route service's static DNS nor the
|
||||
// discovered (DHCP/scutil) nameservers contain an IPv4 address.
|
||||
//
|
||||
// IPv6-only DNS is not usable under DNS intercept mode on macOS: the pf
|
||||
// ruleset blocks all outbound IPv6 port-53 traffic (IPv6 interception is not
|
||||
// supported, see issues #507/#533), and with no IPv4 DNS configured
|
||||
// mDNSResponder emits no DNS packets at all — leaving pf nothing to
|
||||
// intercept despite a healthy upstream. Observed in production on IPv6-only
|
||||
// iPhone tethering with 464XLAT (issue #533).
|
||||
func needsInterceptDNSTarget(staticDNS, discovered []string) bool {
|
||||
return !hasIPv4DNS(staticDNS) && !hasIPv4DNS(discovered)
|
||||
}
|
||||
|
||||
// isInterceptDNSTargetOnly reports whether the given static DNS list is
|
||||
// exactly the entry ctrld set via ensureInterceptDNSTarget (recorded in
|
||||
// target), meaning it is safe for ctrld to remove.
|
||||
func isInterceptDNSTargetOnly(nameservers []string, target string) bool {
|
||||
return target != "" && len(nameservers) == 1 && nameservers[0] == target
|
||||
}
|
||||
|
||||
// filterOwnTarget returns nameservers with ctrld's own recorded target
|
||||
// removed. A previously-set target must never be mistaken for user/network
|
||||
// IPv4 DNS when judging whether the network still needs one — otherwise the
|
||||
// second recovery on the same DNS-less network would see "IPv4 DNS present"
|
||||
// and remove the entry, and the third would re-add it, oscillating on every
|
||||
// recovery.
|
||||
func filterOwnTarget(nameservers []string, target string) []string {
|
||||
if target == "" {
|
||||
return nameservers
|
||||
}
|
||||
out := nameservers[:0:0]
|
||||
for _, s := range nameservers {
|
||||
host := s
|
||||
if h, _, err := net.SplitHostPort(s); err == nil {
|
||||
host = h
|
||||
}
|
||||
if host != target {
|
||||
out = append(out, s)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,234 @@
|
||||
//go:build darwin
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net"
|
||||
"os"
|
||||
|
||||
"tailscale.com/net/netmon"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
// interceptDNSTargetStateFile persists which service/value ctrld set, so a
|
||||
// daemon restart (crash, upgrade, plain restart) does not orphan the entry:
|
||||
// without it a restarted daemon would not know the entry is ctrld's own and
|
||||
// could neither remove it on shutdown nor keep its bookkeeping consistent.
|
||||
const interceptDNSTargetStateFile = ".intercept_dns_target"
|
||||
|
||||
var (
|
||||
interceptDNSTargetStatePathFn = func() string { return absHomeDir(interceptDNSTargetStateFile) }
|
||||
interceptDefaultRouteInterfaceFn = netmon.DefaultRouteInterface
|
||||
interceptInterfaceByNameFn = net.InterfaceByName
|
||||
interceptPatchNetIfaceNameFn = patchNetIfaceName
|
||||
interceptCurrentStaticDNSFn = currentStaticDNS
|
||||
interceptSaveCurrentStaticDNSFn = saveCurrentStaticDNS
|
||||
interceptSetDNSFn = setDNS
|
||||
interceptSavedStaticNameserversFn = savedStaticNameservers
|
||||
interceptResetDNSIgnoreUnusableIfaceFn = resetDnsIgnoreUnusableInterface
|
||||
interceptDHCPNameserversForInterfaceFn = ctrld.DHCPNameserversForInterface
|
||||
)
|
||||
|
||||
type interceptDNSTargetState struct {
|
||||
Service string `json:"service"`
|
||||
Value string `json:"value"`
|
||||
}
|
||||
|
||||
// loadInterceptDNSTargetStateLocked hydrates in-memory tracking from the
|
||||
// state file once (only when memory is empty). Callers must hold
|
||||
// interceptDNSTargetMu.
|
||||
func (p *prog) loadInterceptDNSTargetStateLocked() {
|
||||
if p.interceptDNSTargetService != "" || p.interceptDNSTargetLoaded {
|
||||
return
|
||||
}
|
||||
p.interceptDNSTargetLoaded = true
|
||||
data, err := os.ReadFile(interceptDNSTargetStatePathFn())
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
var st interceptDNSTargetState
|
||||
if err := json.Unmarshal(data, &st); err != nil || st.Service == "" || st.Value == "" {
|
||||
return
|
||||
}
|
||||
p.interceptDNSTargetService = st.Service
|
||||
p.interceptDNSTargetSetValue = st.Value
|
||||
mainLog.Load().Debug().Msgf("intercept DNS target: restored tracking of %s on %q from previous run", st.Value, st.Service)
|
||||
}
|
||||
|
||||
// persistInterceptDNSTargetStateLocked writes (or clears) the state file to
|
||||
// match in-memory tracking. Callers must hold interceptDNSTargetMu.
|
||||
func (p *prog) persistInterceptDNSTargetStateLocked() {
|
||||
file := interceptDNSTargetStatePathFn()
|
||||
if p.interceptDNSTargetService == "" {
|
||||
_ = os.Remove(file)
|
||||
return
|
||||
}
|
||||
data, err := json.Marshal(interceptDNSTargetState{Service: p.interceptDNSTargetService, Value: p.interceptDNSTargetSetValue})
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if err := os.WriteFile(file, data, 0600); err != nil {
|
||||
mainLog.Load().Debug().Err(err).Msg("intercept DNS target: could not persist state file")
|
||||
}
|
||||
}
|
||||
|
||||
// ensureInterceptDNSTarget guarantees macOS always has an emittable DNS
|
||||
// target while DNS intercept mode is active.
|
||||
//
|
||||
// Intercept mode deliberately never manages interface DNS: pf redirects DNS
|
||||
// packets in flight. But pf can only redirect packets macOS actually sends,
|
||||
// and mDNSResponder emits none when the active network service has no DNS
|
||||
// configured. IPv6-only networks (e.g. iPhone tethering with 464XLAT) supply
|
||||
// no IPv4 DNS, and the pf ruleset blocks all outbound IPv6 port 53, so such
|
||||
// networks otherwise end in a total DNS outage with a healthy upstream
|
||||
// (issue #533).
|
||||
//
|
||||
// Only when the default-route service has no usable IPv4 DNS at all does
|
||||
// ctrld set a loopback DNS value on it — chosen by interceptDNSTargetValue to
|
||||
// respect the configured listener: the listener IP directly when it serves
|
||||
// port 53, else a distinct loopback address so the pf lo0 rdr rule rewrites
|
||||
// to the listener's real port. The entry is removed when the network regains
|
||||
// IPv4 DNS and on intercept shutdown. Networks that provide IPv4 DNS are
|
||||
// never modified.
|
||||
//
|
||||
// Callers pass a non-nil raw system discovery result to prove discovery ran;
|
||||
// an empty slice is a valid DNS-less result. The decision itself uses static
|
||||
// DNS plus DHCP option 6 from the default-route interface, so resolvers on a
|
||||
// second physical interface cannot suppress the target. Invoked during
|
||||
// startup, debounced network recovery, and periodic pf watchdog reconciliation.
|
||||
func (p *prog) ensureInterceptDNSTarget(systemDiscovery []string) {
|
||||
if !dnsIntercept || p.dnsInterceptState == nil {
|
||||
return
|
||||
}
|
||||
if systemDiscovery == nil {
|
||||
mainLog.Load().Debug().Msg("intercept DNS target: system DNS discovery was not performed; not changing DNS")
|
||||
return
|
||||
}
|
||||
p.interceptDNSTargetMu.Lock()
|
||||
defer p.interceptDNSTargetMu.Unlock()
|
||||
p.loadInterceptDNSTargetStateLocked()
|
||||
|
||||
drIfaceName, err := interceptDefaultRouteInterfaceFn()
|
||||
if err != nil || drIfaceName == "" {
|
||||
// Mid-transition with no default route; the next recovery decides.
|
||||
return
|
||||
}
|
||||
iface, err := interceptInterfaceByNameFn(drIfaceName)
|
||||
if err != nil || iface == nil {
|
||||
return
|
||||
}
|
||||
// Resolve the network service name (e.g. en5 -> "iPhone USB") so
|
||||
// networksetup operates on the right service.
|
||||
if _, err := interceptPatchNetIfaceNameFn(iface); err != nil {
|
||||
mainLog.Load().Debug().Err(err).Msgf("intercept DNS target: could not resolve network service for %s", drIfaceName)
|
||||
return
|
||||
}
|
||||
|
||||
staticDNS, err := interceptCurrentStaticDNSFn(iface)
|
||||
if err != nil {
|
||||
// Interfaces without a network service (utun/VPN tunnels) land here:
|
||||
// networksetup cannot address them, ctrld never writes to them, and
|
||||
// any target set on the underlying physical service stays in place —
|
||||
// still correct while ctrld runs.
|
||||
mainLog.Load().Debug().Err(err).Msgf("intercept DNS target: could not read static DNS for %q", iface.Name)
|
||||
return
|
||||
}
|
||||
// Never count ctrld's own previously-set entry as network-provided DNS,
|
||||
// or the next recovery on the same DNS-less network would remove it and
|
||||
// the one after re-add it.
|
||||
if p.interceptDNSTargetService == iface.Name {
|
||||
staticDNS = filterOwnTarget(staticDNS, p.interceptDNSTargetSetValue)
|
||||
}
|
||||
if hasIPv4DNS(staticDNS) {
|
||||
p.removeInterceptDNSTargetLocked("network has usable static IPv4 DNS")
|
||||
return
|
||||
}
|
||||
|
||||
routeDHCPDNS, err := interceptDHCPNameserversForInterfaceFn(drIfaceName)
|
||||
if err != nil {
|
||||
mainLog.Load().Debug().Err(err).Msgf("intercept DNS target: could not read DHCP DNS for default-route service %q", iface.Name)
|
||||
return
|
||||
}
|
||||
if hasIPv4DNS(routeDHCPDNS) {
|
||||
// The default-route service regained DHCP option 6. Remove a target
|
||||
// previously set on this or another service.
|
||||
p.removeInterceptDNSTargetLocked("network has usable DHCP IPv4 DNS")
|
||||
return
|
||||
}
|
||||
|
||||
target := p.interceptDNSTargetValue()
|
||||
if p.interceptDNSTargetService == iface.Name && p.interceptDNSTargetSetValue == target {
|
||||
return // already set on this service
|
||||
}
|
||||
// Default route moved to a different DNS-less service (or the listener
|
||||
// config changed): clear the stale entry first.
|
||||
p.removeInterceptDNSTargetLocked("default route service changed")
|
||||
|
||||
// Preserve any existing (IPv6-only) static entries for later restore.
|
||||
// saveCurrentStaticDNS filters loopback on write, and
|
||||
// savedStaticNameservers filters loopback on read, so ctrld's own
|
||||
// loopback target can never be recorded or restored as user DNS.
|
||||
if err := interceptSaveCurrentStaticDNSFn(iface); err != nil {
|
||||
mainLog.Load().Debug().Err(err).Msgf("intercept DNS target: could not save static DNS for %q", iface.Name)
|
||||
}
|
||||
if err := interceptSetDNSFn(iface, []string{target}); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msgf("intercept DNS target: could not set %s on %q", target, iface.Name)
|
||||
return
|
||||
}
|
||||
p.interceptDNSTargetService = iface.Name
|
||||
p.interceptDNSTargetSetValue = target
|
||||
p.persistInterceptDNSTargetStateLocked()
|
||||
mainLog.Load().Warn().Msgf("intercept DNS target: service %q provides no usable IPv4 DNS; set %s so macOS can emit DNS queries (removed automatically when the network provides IPv4 DNS)", iface.Name, target)
|
||||
}
|
||||
|
||||
// removeInterceptDNSTarget removes a previously set intercept DNS target,
|
||||
// restoring the service's saved static DNS (or empty). Safe no-op when no
|
||||
// target was set.
|
||||
func (p *prog) removeInterceptDNSTarget(reason string) {
|
||||
p.interceptDNSTargetMu.Lock()
|
||||
defer p.interceptDNSTargetMu.Unlock()
|
||||
p.loadInterceptDNSTargetStateLocked()
|
||||
p.removeInterceptDNSTargetLocked(reason)
|
||||
}
|
||||
|
||||
// removeInterceptDNSTargetLocked is removeInterceptDNSTarget without locking;
|
||||
// callers must hold interceptDNSTargetMu.
|
||||
func (p *prog) removeInterceptDNSTargetLocked(reason string) {
|
||||
svc := p.interceptDNSTargetService
|
||||
val := p.interceptDNSTargetSetValue
|
||||
if svc == "" {
|
||||
return
|
||||
}
|
||||
iface := &net.Interface{Name: svc}
|
||||
// Only remove what ctrld set. If the service's DNS changed externally,
|
||||
// leave that value alone and discard our stale ownership record.
|
||||
cur, err := interceptCurrentStaticDNSFn(iface)
|
||||
if err != nil {
|
||||
mainLog.Load().Debug().Err(err).Msgf("intercept DNS target: could not read %q DNS; retaining cleanup state (%s)", svc, reason)
|
||||
return
|
||||
}
|
||||
if !isInterceptDNSTargetOnly(cur, val) {
|
||||
mainLog.Load().Debug().Msgf("intercept DNS target: %q DNS changed externally; not removing (%s)", svc, reason)
|
||||
p.clearInterceptDNSTargetStateLocked()
|
||||
return
|
||||
}
|
||||
if saved := interceptSavedStaticNameserversFn(iface); len(saved) > 0 {
|
||||
if err := interceptSetDNSFn(iface, saved); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msgf("intercept DNS target: could not restore saved DNS on %q; retaining cleanup state", svc)
|
||||
return
|
||||
}
|
||||
} else if err := interceptResetDNSIgnoreUnusableIfaceFn(iface); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msgf("intercept DNS target: could not reset DNS on %q; retaining cleanup state", svc)
|
||||
return
|
||||
}
|
||||
p.clearInterceptDNSTargetStateLocked()
|
||||
mainLog.Load().Info().Msgf("intercept DNS target: removed %s from %q (%s)", val, svc, reason)
|
||||
}
|
||||
|
||||
func (p *prog) clearInterceptDNSTargetStateLocked() {
|
||||
p.interceptDNSTargetService = ""
|
||||
p.interceptDNSTargetSetValue = ""
|
||||
p.persistInterceptDNSTargetStateLocked()
|
||||
}
|
||||
@@ -0,0 +1,248 @@
|
||||
//go:build darwin
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"testing"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
type interceptTargetHarness struct {
|
||||
dns map[string][]string
|
||||
saved map[string][]string
|
||||
serviceByDev map[string]string
|
||||
dhcp []string
|
||||
dhcpErr error
|
||||
readErr error
|
||||
setErr error
|
||||
resetErr error
|
||||
setCalls []string
|
||||
resetCalls []string
|
||||
statePath string
|
||||
}
|
||||
|
||||
func newInterceptTargetHarness(t *testing.T) *interceptTargetHarness {
|
||||
t.Helper()
|
||||
h := &interceptTargetHarness{
|
||||
dns: make(map[string][]string),
|
||||
saved: make(map[string][]string),
|
||||
serviceByDev: map[string]string{"en1": "Wi-Fi"},
|
||||
statePath: filepath.Join(t.TempDir(), interceptDNSTargetStateFile),
|
||||
}
|
||||
|
||||
origPath := interceptDNSTargetStatePathFn
|
||||
origRoute := interceptDefaultRouteInterfaceFn
|
||||
origIface := interceptInterfaceByNameFn
|
||||
origPatch := interceptPatchNetIfaceNameFn
|
||||
origCurrent := interceptCurrentStaticDNSFn
|
||||
origSave := interceptSaveCurrentStaticDNSFn
|
||||
origSet := interceptSetDNSFn
|
||||
origSaved := interceptSavedStaticNameserversFn
|
||||
origReset := interceptResetDNSIgnoreUnusableIfaceFn
|
||||
origDHCP := interceptDHCPNameserversForInterfaceFn
|
||||
origIntercept := dnsIntercept
|
||||
t.Cleanup(func() {
|
||||
interceptDNSTargetStatePathFn = origPath
|
||||
interceptDefaultRouteInterfaceFn = origRoute
|
||||
interceptInterfaceByNameFn = origIface
|
||||
interceptPatchNetIfaceNameFn = origPatch
|
||||
interceptCurrentStaticDNSFn = origCurrent
|
||||
interceptSaveCurrentStaticDNSFn = origSave
|
||||
interceptSetDNSFn = origSet
|
||||
interceptSavedStaticNameserversFn = origSaved
|
||||
interceptResetDNSIgnoreUnusableIfaceFn = origReset
|
||||
interceptDHCPNameserversForInterfaceFn = origDHCP
|
||||
dnsIntercept = origIntercept
|
||||
})
|
||||
|
||||
dnsIntercept = true
|
||||
interceptDNSTargetStatePathFn = func() string { return h.statePath }
|
||||
interceptDefaultRouteInterfaceFn = func() (string, error) { return "en1", nil }
|
||||
interceptInterfaceByNameFn = func(name string) (*net.Interface, error) { return &net.Interface{Name: name}, nil }
|
||||
interceptPatchNetIfaceNameFn = func(iface *net.Interface) (bool, error) {
|
||||
service, ok := h.serviceByDev[iface.Name]
|
||||
if !ok {
|
||||
return false, errors.New("unknown network service")
|
||||
}
|
||||
iface.Name = service
|
||||
return true, nil
|
||||
}
|
||||
interceptCurrentStaticDNSFn = func(iface *net.Interface) ([]string, error) {
|
||||
if h.readErr != nil {
|
||||
return nil, h.readErr
|
||||
}
|
||||
return slices.Clone(h.dns[iface.Name]), nil
|
||||
}
|
||||
interceptSaveCurrentStaticDNSFn = func(iface *net.Interface) error {
|
||||
h.saved[iface.Name] = slices.Clone(h.dns[iface.Name])
|
||||
return nil
|
||||
}
|
||||
interceptSetDNSFn = func(iface *net.Interface, nameservers []string) error {
|
||||
h.setCalls = append(h.setCalls, iface.Name)
|
||||
if h.setErr != nil {
|
||||
return h.setErr
|
||||
}
|
||||
h.dns[iface.Name] = slices.Clone(nameservers)
|
||||
return nil
|
||||
}
|
||||
interceptSavedStaticNameserversFn = func(iface *net.Interface) []string {
|
||||
return slices.Clone(h.saved[iface.Name])
|
||||
}
|
||||
interceptResetDNSIgnoreUnusableIfaceFn = func(iface *net.Interface) error {
|
||||
h.resetCalls = append(h.resetCalls, iface.Name)
|
||||
if h.resetErr != nil {
|
||||
return h.resetErr
|
||||
}
|
||||
h.dns[iface.Name] = nil
|
||||
return nil
|
||||
}
|
||||
interceptDHCPNameserversForInterfaceFn = func(iface string) ([]string, error) {
|
||||
if iface != "en1" {
|
||||
return nil, errors.New("DHCP lookup used a non-default interface")
|
||||
}
|
||||
return slices.Clone(h.dhcp), h.dhcpErr
|
||||
}
|
||||
return h
|
||||
}
|
||||
|
||||
func newInterceptTargetProg() *prog {
|
||||
return &prog{
|
||||
cfg: &ctrld.Config{Listener: map[string]*ctrld.ListenerConfig{
|
||||
"0": {IP: "127.0.0.1", Port: 5354},
|
||||
}},
|
||||
dnsInterceptState: &interceptStateStub{},
|
||||
}
|
||||
}
|
||||
|
||||
func persistInterceptTargetForTest(t *testing.T, p *prog, service, value string) {
|
||||
t.Helper()
|
||||
p.interceptDNSTargetMu.Lock()
|
||||
defer p.interceptDNSTargetMu.Unlock()
|
||||
p.interceptDNSTargetLoaded = true
|
||||
p.interceptDNSTargetService = service
|
||||
p.interceptDNSTargetSetValue = value
|
||||
p.persistInterceptDNSTargetStateLocked()
|
||||
}
|
||||
|
||||
func TestEnsureInterceptDNSTargetRequiresCompletedDiscovery(t *testing.T) {
|
||||
h := newInterceptTargetHarness(t)
|
||||
p := newInterceptTargetProg()
|
||||
p.ensureInterceptDNSTarget(nil)
|
||||
if len(h.setCalls) != 0 || len(h.resetCalls) != 0 {
|
||||
t.Fatal("nil system discovery changed service DNS")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureInterceptDNSTargetMigratesService(t *testing.T) {
|
||||
h := newInterceptTargetHarness(t)
|
||||
p := newInterceptTargetProg()
|
||||
persistInterceptTargetForTest(t, p, "iPhone USB", "127.0.0.53")
|
||||
h.dns["iPhone USB"] = []string{"127.0.0.53"}
|
||||
h.dns["Wi-Fi"] = nil
|
||||
|
||||
p.ensureInterceptDNSTarget([]string{})
|
||||
|
||||
if len(h.dns["iPhone USB"]) != 0 {
|
||||
t.Fatalf("old service DNS = %v, want empty", h.dns["iPhone USB"])
|
||||
}
|
||||
if got := h.dns["Wi-Fi"]; !slices.Equal(got, []string{"127.0.0.53"}) {
|
||||
t.Fatalf("new service DNS = %v, want [127.0.0.53]", got)
|
||||
}
|
||||
if p.interceptDNSTargetService != "Wi-Fi" || p.interceptDNSTargetSetValue != "127.0.0.53" {
|
||||
t.Fatalf("tracking = %q/%q, want Wi-Fi/127.0.0.53", p.interceptDNSTargetService, p.interceptDNSTargetSetValue)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureInterceptDNSTargetUsesDefaultRouteDHCPOnly(t *testing.T) {
|
||||
t.Run("other interface IPv4 does not suppress target", func(t *testing.T) {
|
||||
h := newInterceptTargetHarness(t)
|
||||
p := newInterceptTargetProg()
|
||||
p.ensureInterceptDNSTarget([]string{"10.10.10.1"})
|
||||
if got := h.dns["Wi-Fi"]; !slices.Equal(got, []string{"127.0.0.53"}) {
|
||||
t.Fatalf("other interface DNS suppressed target: %v", got)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("returned default route DHCP removes target", func(t *testing.T) {
|
||||
h := newInterceptTargetHarness(t)
|
||||
p := newInterceptTargetProg()
|
||||
persistInterceptTargetForTest(t, p, "Wi-Fi", "127.0.0.53")
|
||||
h.dns["Wi-Fi"] = []string{"127.0.0.53"}
|
||||
h.dhcp = []string{"192.168.10.1"}
|
||||
|
||||
p.ensureInterceptDNSTarget([]string{"10.10.10.1"})
|
||||
|
||||
if len(h.dns["Wi-Fi"]) != 0 || p.interceptDNSTargetService != "" {
|
||||
t.Fatalf("returned default-route DHCP DNS did not remove target: dns=%v service=%q", h.dns["Wi-Fi"], p.interceptDNSTargetService)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestRemoveInterceptDNSTargetRestoresStateFileAfterRestart(t *testing.T) {
|
||||
h := newInterceptTargetHarness(t)
|
||||
h.dns["iPhone USB"] = []string{"127.0.0.53"}
|
||||
if err := os.WriteFile(h.statePath, []byte(`{"service":"iPhone USB","value":"127.0.0.53"}`), 0600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
p := newInterceptTargetProg()
|
||||
|
||||
p.removeInterceptDNSTarget("intercept mode inactive")
|
||||
|
||||
if len(h.dns["iPhone USB"]) != 0 || p.interceptDNSTargetService != "" {
|
||||
t.Fatalf("restart cleanup failed: dns=%v service=%q", h.dns["iPhone USB"], p.interceptDNSTargetService)
|
||||
}
|
||||
if _, err := os.Stat(h.statePath); !os.IsNotExist(err) {
|
||||
t.Fatalf("state file still exists after cleanup: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveInterceptDNSTargetKeepsExternalDNS(t *testing.T) {
|
||||
h := newInterceptTargetHarness(t)
|
||||
p := newInterceptTargetProg()
|
||||
persistInterceptTargetForTest(t, p, "Wi-Fi", "127.0.0.53")
|
||||
h.dns["Wi-Fi"] = []string{"8.8.8.8"}
|
||||
|
||||
p.removeInterceptDNSTarget("test")
|
||||
|
||||
if !slices.Equal(h.dns["Wi-Fi"], []string{"8.8.8.8"}) || len(h.setCalls) != 0 || len(h.resetCalls) != 0 {
|
||||
t.Fatalf("external DNS was changed: dns=%v set=%v reset=%v", h.dns["Wi-Fi"], h.setCalls, h.resetCalls)
|
||||
}
|
||||
if p.interceptDNSTargetService != "" {
|
||||
t.Fatal("external change left stale ownership tracking")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveInterceptDNSTargetRetainsStateOnFailure(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
readErr error
|
||||
resetErr error
|
||||
}{
|
||||
{"read failure", errors.New("networksetup read failed"), nil},
|
||||
{"restore failure", nil, errors.New("networksetup reset failed")},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
h := newInterceptTargetHarness(t)
|
||||
p := newInterceptTargetProg()
|
||||
persistInterceptTargetForTest(t, p, "Wi-Fi", "127.0.0.53")
|
||||
h.dns["Wi-Fi"] = []string{"127.0.0.53"}
|
||||
h.readErr = tc.readErr
|
||||
h.resetErr = tc.resetErr
|
||||
|
||||
p.removeInterceptDNSTarget("test")
|
||||
|
||||
if p.interceptDNSTargetService != "Wi-Fi" || p.interceptDNSTargetSetValue != "127.0.0.53" {
|
||||
t.Fatal("failed cleanup discarded retry state")
|
||||
}
|
||||
if _, err := os.Stat(h.statePath); err != nil {
|
||||
t.Fatalf("failed cleanup removed persisted retry state: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
package cli
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestFilterOwnTarget(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
in []string
|
||||
target string
|
||||
wantLen int
|
||||
}{
|
||||
// The oscillation guard (MR !997 review): the second recovery on the
|
||||
// same DNS-less network must not count ctrld's own entry as
|
||||
// network-provided IPv4 DNS.
|
||||
{"removes own entry", []string{"127.0.0.1"}, "127.0.0.1", 0},
|
||||
{"removes own entry with resolver port", []string{"127.0.0.53:53"}, "127.0.0.53", 0},
|
||||
{"keeps user entries", []string{"127.0.0.1", "1.1.1.1"}, "127.0.0.1", 1},
|
||||
{"empty target keeps all", []string{"127.0.0.1"}, "", 1},
|
||||
{"no match keeps all", []string{"1.1.1.1"}, "127.0.0.53", 1},
|
||||
{"nil input", nil, "127.0.0.1", 0},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got := filterOwnTarget(tc.in, tc.target)
|
||||
if len(got) != tc.wantLen {
|
||||
t.Errorf("filterOwnTarget(%v, %q) = %v, want len %d", tc.in, tc.target, got, tc.wantLen)
|
||||
}
|
||||
for _, s := range got {
|
||||
if tc.target != "" && s == tc.target {
|
||||
t.Errorf("filterOwnTarget(%v, %q) retained the target entry", tc.in, tc.target)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestFilterOwnTargetStability pins the recovery-cycle contract: on a
|
||||
// DNS-less network where ctrld already set its target, needsInterceptDNSTarget
|
||||
// over the filtered list must still report true (entry kept, no oscillation),
|
||||
// while a genuine user-added IPv4 server must report false (entry removed).
|
||||
func TestFilterOwnTargetStability(t *testing.T) {
|
||||
target := "127.0.0.1"
|
||||
|
||||
// Second recovery, same tether: only our own entry present. The OS resolver
|
||||
// reports it with :53, while networksetup reports the bare address.
|
||||
static := filterOwnTarget([]string{target}, target)
|
||||
discovered := filterOwnTarget([]string{target + ":53"}, target)
|
||||
if !needsInterceptDNSTarget(static, discovered) {
|
||||
t.Error("second recovery on the same DNS-less network would remove the target (oscillation)")
|
||||
}
|
||||
|
||||
// User manually added a public server meanwhile: target no longer needed.
|
||||
static = filterOwnTarget([]string{target, "1.1.1.1"}, target)
|
||||
if needsInterceptDNSTarget(static, nil) {
|
||||
t.Error("user-added IPv4 DNS not recognized; target would be kept unnecessarily")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
//go:build !darwin
|
||||
|
||||
package cli
|
||||
|
||||
// ensureInterceptDNSTarget is a no-op on non-Darwin platforms: the DNS-less
|
||||
// network problem it solves is specific to macOS pf interception blocking
|
||||
// IPv6 port 53 with no IPv4 fallback (issue #533). Windows intercept mode
|
||||
// uses NRPT, which routes queries regardless of adapter DNS configuration.
|
||||
func (p *prog) ensureInterceptDNSTarget(_ []string) {}
|
||||
|
||||
// removeInterceptDNSTarget is a no-op on non-Darwin platforms.
|
||||
//
|
||||
//lint:ignore U1000 called from Darwin-only intercept shutdown; kept for API symmetry.
|
||||
func (p *prog) removeInterceptDNSTarget(_ string) {}
|
||||
@@ -0,0 +1,113 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
func TestHasIPv4DNS(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
in []string
|
||||
want bool
|
||||
}{
|
||||
{"empty", nil, false},
|
||||
{"ipv4", []string{"8.8.8.8"}, true},
|
||||
{"ipv4 with port", []string{"192.168.1.1:53"}, true},
|
||||
{"loopback counts", []string{"127.0.0.1"}, true},
|
||||
{"ipv6 only", []string{"2001:4860:4860::8888"}, false},
|
||||
{"ipv6 with port", []string{"[2001:4860:4860::8888]:53"}, false},
|
||||
{"mixed", []string{"2001:4860:4860::8888", "9.9.9.9"}, true},
|
||||
{"garbage ignored", []string{"not-an-ip", ""}, false},
|
||||
{"garbage plus v4", []string{"not-an-ip", "1.1.1.1"}, true},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := hasIPv4DNS(tc.in); got != tc.want {
|
||||
t.Errorf("hasIPv4DNS(%v) = %v, want %v", tc.in, got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNeedsInterceptDNSTarget(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
static, discovered []string
|
||||
want bool
|
||||
}{
|
||||
{"no dns at all", nil, nil, true},
|
||||
{"ipv6-only tether (464XLAT, issue #533)", nil, []string{"2605:8d80::1"}, true},
|
||||
{"static v4 present", []string{"1.1.1.1"}, nil, false},
|
||||
{"discovered v4 present", nil, []string{"192.168.1.1:53"}, false},
|
||||
{"existing ctrld target satisfies", []string{"127.0.0.1"}, nil, false},
|
||||
{"ipv6 static, v4 discovered", []string{"2001:db8::1"}, []string{"10.0.0.1"}, false},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := needsInterceptDNSTarget(tc.static, tc.discovered); got != tc.want {
|
||||
t.Errorf("needsInterceptDNSTarget(%v, %v) = %v, want %v", tc.static, tc.discovered, got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsInterceptDNSTargetOnly(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
in []string
|
||||
target string
|
||||
want bool
|
||||
}{
|
||||
{"exactly ours (direct listener)", []string{"127.0.0.1"}, "127.0.0.1", true},
|
||||
{"exactly ours (rdr target)", []string{"127.0.0.53"}, "127.0.0.53", true},
|
||||
{"empty list", nil, "127.0.0.1", false},
|
||||
{"empty target never matches", []string{"127.0.0.1"}, "", false},
|
||||
{"ours plus user entry", []string{"127.0.0.1", "1.1.1.1"}, "127.0.0.1", false},
|
||||
{"user entry only", []string{"1.1.1.1"}, "127.0.0.1", false},
|
||||
{"different loopback than ours", []string{"127.0.0.53"}, "127.0.0.1", false},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := isInterceptDNSTargetOnly(tc.in, tc.target); got != tc.want {
|
||||
t.Errorf("isInterceptDNSTargetOnly(%v, %q) = %v, want %v", tc.in, tc.target, got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestInterceptDNSTargetValue(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
ip string
|
||||
port int
|
||||
want string
|
||||
}{
|
||||
{"default direct listener :53", "127.0.0.1", 53, "127.0.0.1"},
|
||||
{"custom loopback listener :53", "127.0.0.2", 53, "127.0.0.2"},
|
||||
{"non-53 port uses rdr target", "127.0.0.1", 5354, "127.0.0.53"},
|
||||
{"listener on rdr target with non-53 port", "127.0.0.53", 5354, "127.0.0.54"},
|
||||
{"wildcard ip :53 falls back to loopback", "0.0.0.0", 53, "127.0.0.1"},
|
||||
{"wildcard ip non-53 uses rdr target", "0.0.0.0", 5354, "127.0.0.53"},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
p := &prog{cfg: &ctrld.Config{
|
||||
Listener: map[string]*ctrld.ListenerConfig{
|
||||
"0": {IP: tc.ip, Port: tc.port},
|
||||
},
|
||||
}}
|
||||
if got := p.interceptDNSTargetValue(); got != tc.want {
|
||||
t.Errorf("interceptDNSTargetValue() with listener %s:%d = %q, want %q", tc.ip, tc.port, got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestInterceptDNSTargetValue_NoListener(t *testing.T) {
|
||||
p := &prog{cfg: &ctrld.Config{}}
|
||||
if got := p.interceptDNSTargetValue(); got != "127.0.0.1" {
|
||||
t.Errorf("interceptDNSTargetValue() with no listener = %q, want 127.0.0.1", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
func TestUpdateConfigInterceptMode(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
current string
|
||||
mode string
|
||||
want string
|
||||
wantUpdated bool
|
||||
}{
|
||||
{name: "empty flag preserves config", current: "dns", mode: "", want: "dns"},
|
||||
{name: "dns is persisted", mode: "dns", want: "dns", wantUpdated: true},
|
||||
{name: "hard is persisted", current: "dns", mode: "hard", want: "hard", wantUpdated: true},
|
||||
{name: "off clears persisted mode", current: "dns", mode: "off", want: "", wantUpdated: true},
|
||||
{name: "off is idempotent", mode: "off", want: ""},
|
||||
{name: "invalid flag preserves config", current: "hard", mode: "invalid", want: "hard"},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
cfg := &ctrld.Config{}
|
||||
cfg.Service.InterceptMode = tc.current
|
||||
updated := updateConfigInterceptMode(cfg, tc.mode)
|
||||
if updated != tc.wantUpdated {
|
||||
t.Fatalf("updateConfigInterceptMode() updated = %v, want %v", updated, tc.wantUpdated)
|
||||
}
|
||||
if cfg.Service.InterceptMode != tc.want {
|
||||
t.Fatalf("service.intercept_mode = %q, want %q", cfg.Service.InterceptMode, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfiguredInterceptMode(t *testing.T) {
|
||||
oldInterceptMode := interceptMode
|
||||
t.Cleanup(func() { interceptMode = oldInterceptMode })
|
||||
|
||||
p := &prog{cfg: &ctrld.Config{}}
|
||||
p.cfg.Service.InterceptMode = "dns"
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
flag string
|
||||
want string
|
||||
}{
|
||||
{name: "empty flag falls back to config", flag: "", want: "dns"},
|
||||
{name: "explicit off is final", flag: "off", want: "off"},
|
||||
{name: "explicit hard wins over config", flag: "hard", want: "hard"},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
interceptMode = tc.flag
|
||||
if got := p.configuredInterceptMode(); got != tc.want {
|
||||
t.Fatalf("configuredInterceptMode() = %q, want %q", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
package cli
|
||||
|
||||
// Interception probe registry.
|
||||
//
|
||||
// A probe sends a DNS query for a unique synthetic domain through the OS resolver and
|
||||
// waits for ctrld's own handler to receive it. That is the only way to tell "the rules are
|
||||
// present" from "the rules are actually redirecting packets", and both the macOS pf path
|
||||
// and the Windows NRPT path use it.
|
||||
//
|
||||
// Each attempt registers its own domain, so overlapping probes cannot cancel each other,
|
||||
// and deregistration only removes the entry it owns.
|
||||
|
||||
// registerInterceptProbe registers domain and returns the channel it will be signalled on
|
||||
// plus the function that removes the registration.
|
||||
//
|
||||
//lint:ignore U1000 used on darwin (pf probes) and windows (NRPT probes)
|
||||
func (p *prog) registerInterceptProbe(domain string) (<-chan struct{}, func()) {
|
||||
ch := make(chan struct{}, 1)
|
||||
|
||||
p.interceptProbeMu.Lock()
|
||||
current, _ := p.interceptProbes.Load().(map[string]chan struct{})
|
||||
next := make(map[string]chan struct{}, len(current)+1)
|
||||
for k, v := range current {
|
||||
next[k] = v
|
||||
}
|
||||
next[domain] = ch
|
||||
p.interceptProbes.Store(next)
|
||||
p.interceptProbeMu.Unlock()
|
||||
|
||||
return ch, func() {
|
||||
p.interceptProbeMu.Lock()
|
||||
defer p.interceptProbeMu.Unlock()
|
||||
current, _ := p.interceptProbes.Load().(map[string]chan struct{})
|
||||
// Only drop the entry while it is still this attempt's channel. A later probe
|
||||
// that reused the domain owns the slot now, and clearing it would make that one
|
||||
// wait out its timeout for a query it already received.
|
||||
if existing, ok := current[domain]; !ok || existing != ch {
|
||||
return
|
||||
}
|
||||
next := make(map[string]chan struct{}, len(current))
|
||||
for k, v := range current {
|
||||
if k != domain {
|
||||
next[k] = v
|
||||
}
|
||||
}
|
||||
p.interceptProbes.Store(next)
|
||||
}
|
||||
}
|
||||
|
||||
// signalInterceptProbe reports whether domain is a pending probe, signalling its waiter
|
||||
// when it is. Called from the DNS handler for every query, so the common case is a nil or
|
||||
// empty map and no allocation.
|
||||
func (p *prog) signalInterceptProbe(domain string) bool {
|
||||
probes, _ := p.interceptProbes.Load().(map[string]chan struct{})
|
||||
if len(probes) == 0 {
|
||||
return false
|
||||
}
|
||||
ch, ok := probes[domain]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
select {
|
||||
case ch <- struct{}{}:
|
||||
default:
|
||||
// Buffered channel already holds a signal: the waiter has what it needs.
|
||||
}
|
||||
return true
|
||||
}
|
||||
+17
-6
@@ -50,8 +50,13 @@ func httpClientWithFallback(timeout time.Duration) *http.Client {
|
||||
|
||||
// doWithRetry performs an HTTP request with retries
|
||||
func doWithRetry(req *http.Request, maxRetries int, ip string) (*http.Response, error) {
|
||||
return doWithRetryClient(httpClientWithFallback(defaultHTTPTimeout), req, maxRetries, ip)
|
||||
}
|
||||
|
||||
// doWithRetryClient is doWithRetry with an injectable client, so the retry and
|
||||
// error-composition behaviour can be tested without real network access.
|
||||
func doWithRetryClient(client *http.Client, req *http.Request, maxRetries int, ip string) (*http.Response, error) {
|
||||
var lastErr error
|
||||
client := httpClientWithFallback(defaultHTTPTimeout)
|
||||
var ipReq *http.Request
|
||||
if ip != "" {
|
||||
ipReq = req.Clone(req.Context())
|
||||
@@ -67,22 +72,28 @@ func doWithRetry(req *http.Request, maxRetries int, ip string) (*http.Response,
|
||||
if err == nil {
|
||||
return resp, nil
|
||||
}
|
||||
// Keep the hostname attempt's error: it carries the diagnosis (on Windows,
|
||||
// a local firewall denying the socket shows up here as WSAEACCES), while the
|
||||
// direct-IP fallback often fails for an unrelated reason such as an
|
||||
// unreachable IPv6 route.
|
||||
attemptErr := err
|
||||
if ipReq != nil {
|
||||
mainLog.Load().Warn().Err(err).Msgf("dial to %q failed", req.Host)
|
||||
mainLog.Load().Warn().Msgf("fallback to direct IP to download prod version: %q", ip)
|
||||
resp, err = client.Do(ipReq)
|
||||
if err == nil {
|
||||
resp, fallbackErr := client.Do(ipReq)
|
||||
if fallbackErr == nil {
|
||||
return resp, nil
|
||||
}
|
||||
attemptErr = fmt.Errorf("%w; fallback to direct ip %s failed: %w", attemptErr, ip, fallbackErr)
|
||||
}
|
||||
|
||||
lastErr = err
|
||||
mainLog.Load().Debug().Err(err).
|
||||
lastErr = attemptErr
|
||||
mainLog.Load().Debug().Err(attemptErr).
|
||||
Str("method", req.Method).
|
||||
Str("url", req.URL.String()).
|
||||
Msgf("HTTP request attempt %d/%d failed", attempt+1, maxRetries)
|
||||
}
|
||||
return nil, fmt.Errorf("failed after %d attempts to %s %s: %v", maxRetries, req.Method, req.URL, lastErr)
|
||||
return nil, fmt.Errorf("failed after %d attempts to %s %s: %w", maxRetries, req.Method, req.URL, lastErr)
|
||||
}
|
||||
|
||||
// Helper for making GET requests with retries
|
||||
|
||||
@@ -0,0 +1,241 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"syscall"
|
||||
"testing"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld/internal/controld"
|
||||
)
|
||||
|
||||
// wsaEACCES is WSAEACCES (10013): "An attempt was made to access a socket in a way
|
||||
// forbidden by its access permissions." This is what Windows reports when a WFP
|
||||
// filter denies the connect. Used as a plain errno so the test runs everywhere.
|
||||
const wsaEACCES = syscall.Errno(10013)
|
||||
|
||||
// denyingRoundTripper denies the hostname attempt with firstErr and the direct-ip
|
||||
// attempt with fbErr, the shape seen during the Firewall Mode incident: the
|
||||
// hostname attempt was denied by ctrld's own stale block-all filters, while the
|
||||
// direct-ip fallback failed on an unreachable IPv6 route.
|
||||
type denyingRoundTripper struct {
|
||||
hostname string
|
||||
firstErr error
|
||||
fbErr error
|
||||
}
|
||||
|
||||
func (rt *denyingRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
if req.URL.Host == rt.hostname {
|
||||
return nil, &net.OpError{Op: "dial", Net: "tcp4", Err: rt.firstErr}
|
||||
}
|
||||
return nil, &net.OpError{Op: "dial", Net: "tcp6", Err: rt.fbErr}
|
||||
}
|
||||
|
||||
func TestDoWithRetryPreservesHostnameError(t *testing.T) {
|
||||
const hostname = "dl.controld.dev"
|
||||
req, err := http.NewRequest(http.MethodGet, "https://"+hostname+"/v2/windows-amd64/ctrld.exe", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rt := &denyingRoundTripper{
|
||||
hostname: hostname,
|
||||
firstErr: wsaEACCES,
|
||||
fbErr: syscall.EHOSTUNREACH,
|
||||
}
|
||||
|
||||
_, err = doWithRetryClient(&http.Client{Transport: rt}, req, 1, "23.171.240.151")
|
||||
if err == nil {
|
||||
t.Fatal("expected doWithRetry to fail when both attempts are denied")
|
||||
}
|
||||
if !errors.Is(err, wsaEACCES) {
|
||||
t.Errorf("hostname-attempt error (WSAEACCES) was lost, got: %v", err)
|
||||
}
|
||||
if !errors.Is(err, syscall.EHOSTUNREACH) {
|
||||
t.Errorf("fallback error was lost, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// composedAttemptErrors builds the error shape the two-attempt paths return: each
|
||||
// attempt's *url.Error (as produced by http.Client.Do) wrapped by a single fmt.Errorf
|
||||
// with two %w verbs, hostname attempt first. Mirrors doWithFallback in
|
||||
// internal/controld and doWithRetryClient above.
|
||||
func composedAttemptErrors(first, fallback error) error {
|
||||
attempt := func(network string, cause error) error {
|
||||
return &url.Error{
|
||||
Op: "Post",
|
||||
URL: "https://api.controld.com/utility",
|
||||
Err: &net.OpError{Op: "dial", Net: network, Err: cause},
|
||||
}
|
||||
}
|
||||
return fmt.Errorf("request failed: %w; fallback to direct ip %s failed: %w",
|
||||
attempt("tcp4", first), "147.185.34.1", attempt("tcp6", fallback))
|
||||
}
|
||||
|
||||
// TestComposedFallbackErrorRetryClassification pins which attempt decides whether
|
||||
// preflight keeps retrying.
|
||||
//
|
||||
// Reporting both attempt errors is not purely diagnostic: processCDFlags decides
|
||||
// retryability with errUrlNetworkError, which uses errors.As, and errors.As is
|
||||
// order-sensitive - it returns the *first* matching error in the tree. Composing the
|
||||
// hostname attempt first therefore hands the retry predicate the hostname failure,
|
||||
// where previously only the fallback's error survived to be classified.
|
||||
//
|
||||
// The consequence is deliberate: a locally denied socket (WSAEACCES, a firewall
|
||||
// blocking ctrld) is no longer treated as a transient network error, so preflight fails
|
||||
// fast and reports instead of backing off - the incident logged 256 retry cycles
|
||||
// against filters that were never going to clear on their own. The boot case that
|
||||
// justifies the indefinite retry, a network unreachable on both attempts, is preserved.
|
||||
//
|
||||
// If the wrap order is ever reversed, this test fails rather than silently restoring
|
||||
// indefinite retries against a host that is actively refusing.
|
||||
func TestComposedFallbackErrorRetryClassification(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
hostname error
|
||||
fallback error
|
||||
wantRetryable bool
|
||||
}{
|
||||
{
|
||||
// The incident's pair: denied locally, IPv6 route unusable.
|
||||
name: "denied socket then unreachable fallback fails fast",
|
||||
hostname: wsaEACCES,
|
||||
fallback: syscall.EHOSTUNREACH,
|
||||
wantRetryable: false,
|
||||
},
|
||||
{
|
||||
// Boot with no network yet: must still retry indefinitely.
|
||||
name: "network unreachable on both attempts still retries",
|
||||
hostname: syscall.ENETUNREACH,
|
||||
fallback: syscall.ENETUNREACH,
|
||||
wantRetryable: true,
|
||||
},
|
||||
{
|
||||
name: "connection refused still retries",
|
||||
hostname: syscall.ECONNREFUSED,
|
||||
fallback: syscall.EHOSTUNREACH,
|
||||
wantRetryable: true,
|
||||
},
|
||||
{
|
||||
name: "permission denied on both attempts fails fast",
|
||||
hostname: syscall.EACCES,
|
||||
fallback: syscall.EACCES,
|
||||
wantRetryable: false,
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
err := composedAttemptErrors(tc.hostname, tc.fallback)
|
||||
if got := errUrlNetworkError(err); got != tc.wantRetryable {
|
||||
t.Errorf("errUrlNetworkError() = %v, want %v", got, tc.wantRetryable)
|
||||
}
|
||||
// Both attempts remain reportable regardless of classification.
|
||||
if !errors.Is(err, tc.hostname) {
|
||||
t.Error("hostname attempt error was lost")
|
||||
}
|
||||
if !errors.Is(err, tc.fallback) {
|
||||
t.Error("fallback attempt error was lost")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestUnresolvedHostnameDefersToFallbackAttempt covers the asymmetric pair.
|
||||
//
|
||||
// Only the hostname attempt resolves DNS, and Go marks a *net.DNSError as temporary only
|
||||
// for socket failures that reached the server - so a SERVFAIL or "no such host" answer is
|
||||
// not temporary. At boot behind a captive portal, or before a router's forwarder is up,
|
||||
// that is exactly how the hostname attempt fails while the network is merely not ready.
|
||||
// Before the composed error existed only the fallback decided, so this pair retried;
|
||||
// classifying the hostname attempt alone would fail it fast and reach Fatal.
|
||||
//
|
||||
// A name-resolution failure therefore carries no verdict: the fallback attempt decides.
|
||||
// The locally-denied case above still fails fast, because a denied socket is definitive.
|
||||
func TestUnresolvedHostnameDefersToFallbackAttempt(t *testing.T) {
|
||||
dnsFailure := &url.Error{
|
||||
Op: "Post",
|
||||
URL: "https://api.controld.com/utility",
|
||||
Err: &net.DNSError{Err: "server misbehaving", Name: "api.controld.com", IsTemporary: false},
|
||||
}
|
||||
attempt := func(cause error) error {
|
||||
return &url.Error{
|
||||
Op: "Post",
|
||||
URL: "https://api.controld.com/utility",
|
||||
Err: &net.OpError{Op: "dial", Net: "tcp6", Err: cause},
|
||||
}
|
||||
}
|
||||
|
||||
retryable := fmt.Errorf("request failed: %w; fallback to direct ip %s failed: %w",
|
||||
dnsFailure, "147.185.34.1", attempt(syscall.ECONNREFUSED))
|
||||
if !errUrlNetworkError(retryable) {
|
||||
t.Error("an unresolved hostname with a retryable fallback must keep retrying: at boot the network is simply not up yet")
|
||||
}
|
||||
|
||||
denied := fmt.Errorf("request failed: %w; fallback to direct ip %s failed: %w",
|
||||
dnsFailure, "147.185.34.1", attempt(wsaEACCES))
|
||||
if errUrlNetworkError(denied) {
|
||||
t.Error("an unresolved hostname with a denied fallback must fail fast: nothing here clears on its own")
|
||||
}
|
||||
|
||||
// A resolution failure alone still says nothing, so it must not be read as retryable.
|
||||
if errUrlNetworkError(dnsFailure) {
|
||||
t.Error("a bare name-resolution failure must not be classified as retryable")
|
||||
}
|
||||
}
|
||||
|
||||
// TestDoWithFallbackClassificationEndToEnd drives the real composition in
|
||||
// internal/controld through the real predicate, instead of asserting a hand-written copy
|
||||
// of its error shape against another hand-written copy. A change to either side's format
|
||||
// string or wrap order is caught here.
|
||||
func TestDoWithFallbackClassificationEndToEnd(t *testing.T) {
|
||||
const hostname = "api.controld.com"
|
||||
req, err := http.NewRequest(http.MethodPost, "https://"+hostname+"/utility", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rt := &denyingRoundTripper{
|
||||
hostname: hostname,
|
||||
firstErr: wsaEACCES,
|
||||
fbErr: syscall.EHOSTUNREACH,
|
||||
}
|
||||
|
||||
_, gotErr := controld.DoWithFallbackForTest(&http.Client{Transport: rt}, req, "147.185.34.1")
|
||||
if gotErr == nil {
|
||||
t.Fatal("expected both attempts to fail")
|
||||
}
|
||||
if errUrlNetworkError(gotErr) {
|
||||
t.Errorf("the real composed error was classified as retryable: %v", gotErr)
|
||||
}
|
||||
if !errors.Is(gotErr, wsaEACCES) || !errors.Is(gotErr, syscall.EHOSTUNREACH) {
|
||||
t.Errorf("the real composed error lost an attempt: %v", gotErr)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDoWithRetryComposesHostnameAttemptFirst anchors the ordering assumption above to
|
||||
// the real composition, so a reordering of the wrap in doWithRetryClient is caught here
|
||||
// and not only in the hand-built shape.
|
||||
func TestDoWithRetryComposesHostnameAttemptFirst(t *testing.T) {
|
||||
const hostname = "dl.controld.dev"
|
||||
req, err := http.NewRequest(http.MethodGet, "https://"+hostname+"/v2/windows-amd64/ctrld.exe", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rt := &denyingRoundTripper{hostname: hostname, firstErr: wsaEACCES, fbErr: syscall.EHOSTUNREACH}
|
||||
|
||||
_, gotErr := doWithRetryClient(&http.Client{Transport: rt}, req, 1, "23.171.240.151")
|
||||
if gotErr == nil {
|
||||
t.Fatal("expected both attempts to fail")
|
||||
}
|
||||
|
||||
// errors.As must reach the hostname attempt first: that is what the retry
|
||||
// predicate classifies.
|
||||
var opErr *net.OpError
|
||||
if !errors.As(gotErr, &opErr) {
|
||||
t.Fatalf("no net.OpError in the chain: %v", gotErr)
|
||||
}
|
||||
if !errors.Is(opErr.Err, wsaEACCES) {
|
||||
t.Errorf("first OpError in the chain is %v, want the hostname attempt (%v)", opErr.Err, wsaEACCES)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
func TestListenerInterceptModeExplicitOff(t *testing.T) {
|
||||
oldIntercept := interceptMode
|
||||
t.Cleanup(func() { interceptMode = oldIntercept })
|
||||
|
||||
cfg := &ctrld.Config{}
|
||||
cfg.Service.InterceptMode = "dns"
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
flag string
|
||||
want string
|
||||
}{
|
||||
{name: "explicit off is final", flag: "off", want: "off"},
|
||||
{name: "empty flag falls back to config", flag: "", want: "dns"},
|
||||
{name: "explicit dns wins over config", flag: "dns", want: "dns"},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
interceptMode = tc.flag
|
||||
if got := listenerInterceptMode(cfg); got != tc.want {
|
||||
t.Fatalf("listenerInterceptMode() = %q, want %q", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -318,6 +318,12 @@ func (p *prog) initInternalLogging(writers []io.Writer) {
|
||||
|
||||
// needInternalLogging reports whether prog needs to run internal logging.
|
||||
func (p *prog) needInternalLogging() bool {
|
||||
// Do not run in silent mode: the user explicitly asked for no logging, so
|
||||
// ctrld must not create or write the persisted internal log file (nor reset
|
||||
// the global level back to debug). See https://github.com/Control-D-Inc/ctrld/issues/320.
|
||||
if silent {
|
||||
return false
|
||||
}
|
||||
// Do not run in non-cd mode.
|
||||
if cdUID == "" {
|
||||
return false
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
// Test_needInternalLogging_silent is a regression test for
|
||||
// https://github.com/Control-D-Inc/ctrld/issues/320: running with --silent must
|
||||
// not enable internal logging, otherwise ctrld creates and writes
|
||||
// <homedir>/ctrld.log (and, when verbose==0, resets the global level back to
|
||||
// debug) despite the user asking for silence.
|
||||
func Test_needInternalLogging_silent(t *testing.T) {
|
||||
origSilent, origCdUID := silent, cdUID
|
||||
t.Cleanup(func() { silent, cdUID = origSilent, origCdUID })
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
silent bool
|
||||
cdUID string
|
||||
logPath string
|
||||
want bool
|
||||
}{
|
||||
{"silent suppresses internal logging in cd mode", true, "test-uid", "", false},
|
||||
{"cd mode enables internal logging", false, "test-uid", "", true},
|
||||
{"non-cd mode disabled", false, "", "", false},
|
||||
{"explicit log path disables internal logging", false, "test-uid", "/var/log/ctrld.log", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
tt := tt
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
silent = tt.silent
|
||||
cdUID = tt.cdUID
|
||||
p := &prog{cfg: &ctrld.Config{}}
|
||||
p.cfg.Service.LogPath = tt.logPath
|
||||
if got := p.needInternalLogging(); got != tt.want {
|
||||
t.Fatalf("needInternalLogging() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Test_initInternalLogging_silentCreatesNoFile drives the real initInternalLogging
|
||||
// path and asserts that a --silent --cd run does not create <homedir>/ctrld.log,
|
||||
// which is the observable failure reported in
|
||||
// https://github.com/Control-D-Inc/ctrld/issues/320.
|
||||
func Test_initInternalLogging_silentCreatesNoFile(t *testing.T) {
|
||||
origSilent, origCdUID, origHomedir := silent, cdUID, homedir
|
||||
t.Cleanup(func() { silent, cdUID, homedir = origSilent, origCdUID, origHomedir })
|
||||
|
||||
dir := t.TempDir()
|
||||
homedir = dir
|
||||
cdUID = "test-uid" // cd mode, which would otherwise enable internal logging
|
||||
silent = true
|
||||
|
||||
p := &prog{cfg: &ctrld.Config{}}
|
||||
p.initInternalLogging(nil)
|
||||
|
||||
logPath := filepath.Join(dir, logFileName)
|
||||
if _, err := os.Stat(logPath); !os.IsNotExist(err) {
|
||||
t.Fatalf("silent mode must not create %s (stat err = %v)", logPath, err)
|
||||
}
|
||||
}
|
||||
+4
-1
@@ -42,7 +42,7 @@ var (
|
||||
cleanup bool
|
||||
startOnly bool
|
||||
rfc1918 bool
|
||||
interceptMode string // "", "dns", or "hard" — set via --intercept-mode flag or config
|
||||
interceptMode string // "", "off", "dns", or "hard" — set via --intercept-mode flag or config
|
||||
dnsIntercept bool // derived: interceptMode == "dns" || interceptMode == "hard"
|
||||
hardIntercept bool // derived: interceptMode == "hard"
|
||||
|
||||
@@ -56,6 +56,9 @@ const (
|
||||
cdOrgFlagName = "cd-org"
|
||||
customHostnameFlagName = "custom-hostname"
|
||||
nextdnsFlagName = "nextdns"
|
||||
|
||||
// autoIface is the sentinel --iface value meaning "use the default gateway interface".
|
||||
autoIface = "auto"
|
||||
)
|
||||
|
||||
func init() {
|
||||
|
||||
+60
-1
@@ -1,17 +1,76 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/rs/zerolog"
|
||||
)
|
||||
|
||||
var logOutput strings.Builder
|
||||
// logOutput is the log sink for the whole test binary. Tests share it with any
|
||||
// background goroutine the code under test starts (watchdogs, timers), so it
|
||||
// must tolerate concurrent writes.
|
||||
var logOutput syncBuffer
|
||||
|
||||
// syncBuffer is a strings.Builder guarded by a mutex.
|
||||
type syncBuffer struct {
|
||||
mu sync.Mutex
|
||||
sb strings.Builder
|
||||
}
|
||||
|
||||
func (b *syncBuffer) Write(p []byte) (int, error) {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
return b.sb.Write(p)
|
||||
}
|
||||
|
||||
func (b *syncBuffer) String() string {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
return b.sb.String()
|
||||
}
|
||||
|
||||
// envFakeVersionOutput makes this test binary impersonate a ctrld executable: when
|
||||
// set, the process writes the value to stdout and exits without running any test, so
|
||||
// binaryVersion() can be exercised on every platform without building or shipping a
|
||||
// fixture binary. The value envFakeVersionSilent produces no output at all, which
|
||||
// reproduces a ctrld.exe_previous that exists but reports no version.
|
||||
//
|
||||
// This must be handled before m.Run(), which is what parses the test flags: the child
|
||||
// is invoked as "<binary> --version" and would otherwise die on an unknown flag.
|
||||
const (
|
||||
envFakeVersionOutput = "CTRLD_TEST_FAKE_VERSION_OUTPUT"
|
||||
envFakeVersionSilent = "<silent>"
|
||||
)
|
||||
|
||||
func TestMain(m *testing.M) {
|
||||
if out := os.Getenv(envFakeVersionOutput); out != "" {
|
||||
if out != envFakeVersionSilent {
|
||||
fmt.Println(out)
|
||||
}
|
||||
os.Exit(0)
|
||||
}
|
||||
|
||||
l := zerolog.New(&logOutput)
|
||||
mainLog.Store(&l)
|
||||
|
||||
// Stub the self-upgrade command builder for the whole test binary. The real
|
||||
// builder execs os.Executable() — which under `go test` IS this test binary
|
||||
// — with positional args ("upgrade", ...). `go test` stops flag parsing at
|
||||
// the first positional arg and ignores the rest, so the child just re-runs
|
||||
// the entire suite, hits the upgrade tests again, and spawns more children:
|
||||
// a fork bomb of detached processes that stalls the host and (on Windows)
|
||||
// holds the test binary's image locked, breaking CI artifact cleanup.
|
||||
// Point it at the test binary with a no-match -test.run so any test that
|
||||
// reaches performUpgrade still exercises the cmd.Start() success path while
|
||||
// the child exits immediately without recursing.
|
||||
newUpgradeCmd = func(exe string) *exec.Cmd {
|
||||
return exec.Command(exe, "-test.run=^$")
|
||||
}
|
||||
|
||||
os.Exit(m.Run())
|
||||
}
|
||||
|
||||
@@ -113,6 +113,22 @@ func (p *prog) runMetricsServer(ctx context.Context, reloadCh chan struct{}) {
|
||||
}
|
||||
|
||||
addr := p.cfg.Service.MetricsListener
|
||||
if addr != "" {
|
||||
host, port, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msgf("Invalid metrics listener address (%s); expected host:port", addr)
|
||||
} else {
|
||||
if host == "" {
|
||||
host = "127.0.0.1"
|
||||
addr = net.JoinHostPort(host, port)
|
||||
}
|
||||
ip := net.ParseIP(host)
|
||||
if (ip != nil && !ip.IsLoopback()) || (ip == nil && host != "localhost") {
|
||||
mainLog.Load().Warn().Msgf("Metrics server is bound to a non-loopback address (%s). This exposes sensitive data without authentication.", addr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
ms, err := newMetricsServer(addr, reg)
|
||||
if err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("could not create new metrics server")
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/netip"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const nrptRuleName = `CtrldCatchAll`
|
||||
|
||||
// errGPNRPTVerified marks an intercept startup failure that happened while an externally
|
||||
// managed (Group Policy) NRPT catch-all was proved - by probe, not by registry shape
|
||||
// alone - to be routing DNS to this listener. It is the difference between "intercept
|
||||
// failed but DNS still reaches ctrld" and "intercept failed and nothing is filtering",
|
||||
// which is what decides whether the interface-DNS fallback must run.
|
||||
//
|
||||
// Only the Windows path produces it, but setDNS is shared, so the sentinel and its
|
||||
// predicate live here with the other platform-neutral NRPT helpers.
|
||||
var errGPNRPTVerified = errors.New("GP-managed NRPT verified routing to ctrld")
|
||||
|
||||
// errGPNRPTIneffective marks a startup that ends with externally managed NRPT owning the
|
||||
// namespace while no probe has proved it routes to ctrld. DNS is not reaching ctrld, but
|
||||
// adapter DNS was deliberately preserved and no ctrld rule may be written beside an
|
||||
// administrator's catch-all - so this is a failed start that must not take the
|
||||
// interface-DNS fallback either.
|
||||
var errGPNRPTIneffective = errors.New("GP-managed NRPT owns the namespace but no probe reached ctrld")
|
||||
|
||||
// interceptFailedWithVerifiedExternalDNS reports whether an intercept startup failure
|
||||
// happened while externally managed DNS policy was verified to be routing to ctrld.
|
||||
func interceptFailedWithVerifiedExternalDNS(err error) bool {
|
||||
return errors.Is(err, errGPNRPTVerified)
|
||||
}
|
||||
|
||||
// interceptFailedUnderExternalDNSPolicy reports whether an intercept startup failure
|
||||
// happened while externally managed DNS policy owned the namespace, whether or not it was
|
||||
// proved to route. Either way the interface-DNS fallback must not run: adapter DNS was
|
||||
// preserved on purpose, and rewriting it would violate the policy ctrld just deferred to.
|
||||
// Only the verified case is a successful start.
|
||||
func interceptFailedUnderExternalDNSPolicy(err error) bool {
|
||||
return errors.Is(err, errGPNRPTVerified) || errors.Is(err, errGPNRPTIneffective)
|
||||
}
|
||||
|
||||
// isExternalGPCatchAll recognizes only a single catch-all namespace that is not
|
||||
// ctrld's deterministic GP key. Registry access stays in the Windows file; this
|
||||
// pure classifier is shared with host-runnable tests.
|
||||
func isExternalGPCatchAll(ruleName string, namespaces []string) bool {
|
||||
return ruleName != "" && !strings.EqualFold(ruleName, nrptRuleName) && len(namespaces) == 1 && strings.TrimSpace(namespaces[0]) == "."
|
||||
}
|
||||
|
||||
func isMatchingGPNRPTRule(ruleName string, namespaces []string, dnsServers, listenerIP string) bool {
|
||||
if !isExternalGPCatchAll(ruleName, namespaces) {
|
||||
return false
|
||||
}
|
||||
server, err := netip.ParseAddr(strings.TrimSpace(dnsServers))
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
listener, err := netip.ParseAddr(strings.TrimSpace(listenerIP))
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return server.Unmap() == listener.Unmap()
|
||||
}
|
||||
@@ -0,0 +1,129 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestIsMatchingGPNRPTRule(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
ruleName string
|
||||
namespaces []string
|
||||
servers string
|
||||
listener string
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "exact IPv4 catch-all",
|
||||
ruleName: "{A1B2C3D4}",
|
||||
namespaces: []string{"."},
|
||||
servers: "127.0.0.1",
|
||||
listener: "127.0.0.1",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "normalized IPv4-mapped listener",
|
||||
ruleName: "{A1B2C3D4}",
|
||||
namespaces: []string{"."},
|
||||
servers: "::ffff:127.0.0.1",
|
||||
listener: "127.0.0.1",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "ctrld GP key is not external",
|
||||
ruleName: "ctrldcatchall",
|
||||
namespaces: []string{"."},
|
||||
servers: "127.0.0.1",
|
||||
listener: "127.0.0.1",
|
||||
},
|
||||
{
|
||||
name: "partial namespace",
|
||||
ruleName: "{A1B2C3D4}",
|
||||
namespaces: []string{"corp.example"},
|
||||
servers: "127.0.0.1",
|
||||
listener: "127.0.0.1",
|
||||
},
|
||||
{
|
||||
name: "multiple namespaces",
|
||||
ruleName: "{A1B2C3D4}",
|
||||
namespaces: []string{".", "corp.example"},
|
||||
servers: "127.0.0.1",
|
||||
listener: "127.0.0.1",
|
||||
},
|
||||
{
|
||||
name: "wrong listener",
|
||||
ruleName: "{A1B2C3D4}",
|
||||
namespaces: []string{"."},
|
||||
servers: "127.0.0.2",
|
||||
listener: "127.0.0.1",
|
||||
},
|
||||
{
|
||||
name: "multiple nameservers",
|
||||
ruleName: "{A1B2C3D4}",
|
||||
namespaces: []string{"."},
|
||||
servers: "127.0.0.1;127.0.0.2",
|
||||
listener: "127.0.0.1",
|
||||
},
|
||||
{
|
||||
name: "malformed nameserver",
|
||||
ruleName: "{A1B2C3D4}",
|
||||
namespaces: []string{"."},
|
||||
servers: "localhost",
|
||||
listener: "127.0.0.1",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := isMatchingGPNRPTRule(tt.ruleName, tt.namespaces, tt.servers, tt.listener); got != tt.want {
|
||||
t.Fatalf("isMatchingGPNRPTRule() = %t, want %t", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsExternalGPCatchAll(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
ruleName string
|
||||
namespaces []string
|
||||
want bool
|
||||
}{
|
||||
{name: "external catch-all", ruleName: "{GP-RULE}", namespaces: []string{"."}, want: true},
|
||||
{name: "ctrld key", ruleName: nrptRuleName, namespaces: []string{"."}},
|
||||
{name: "partial namespace", ruleName: "{GP-RULE}", namespaces: []string{"corp.example"}},
|
||||
{name: "multiple namespaces", ruleName: "{GP-RULE}", namespaces: []string{".", "corp.example"}},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := isExternalGPCatchAll(tt.ruleName, tt.namespaces); got != tt.want {
|
||||
t.Fatalf("isExternalGPCatchAll() = %t, want %t", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestInterceptFailedWithVerifiedExternalDNS covers the distinction the interface-DNS
|
||||
// fallback turns on. "A GP rule exists" is not enough: if it is not actually routing and
|
||||
// intercept failed too, skipping the fallback leaves the machine with no NRPT, no WFP and
|
||||
// no adapter DNS - that is, unfiltered. Only a probe-verified route earns the skip.
|
||||
func TestInterceptFailedWithVerifiedExternalDNS(t *testing.T) {
|
||||
wfpErr := errors.New("FwpmEngineOpen0 failed: HRESULT 0x5")
|
||||
|
||||
verified := fmt.Errorf("dns intercept: WFP setup failed: %w: %w", wfpErr, errGPNRPTVerified)
|
||||
if !interceptFailedWithVerifiedExternalDNS(verified) {
|
||||
t.Error("a failure carrying errGPNRPTVerified must skip the interface-DNS fallback")
|
||||
}
|
||||
if !errors.Is(verified, wfpErr) {
|
||||
t.Error("the underlying cause must stay inspectable for logs and callers")
|
||||
}
|
||||
|
||||
if interceptFailedWithVerifiedExternalDNS(fmt.Errorf("dns intercept: WFP setup failed: %w", wfpErr)) {
|
||||
t.Error("an unverified failure must take the interface-DNS fallback rather than leave the machine unfiltered")
|
||||
}
|
||||
if interceptFailedWithVerifiedExternalDNS(nil) {
|
||||
t.Error("no error must not read as a verified external route")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
//go:build windows
|
||||
|
||||
package cli
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestWFPStateNRPTPolicyOwner(t *testing.T) {
|
||||
state := &wfpState{}
|
||||
state.setNRPTPolicyOwner(nrptRuleOwnerGroupPolicy, "{GP-RULE}")
|
||||
owner, ruleName := state.nrptPolicyOwner()
|
||||
if owner != nrptRuleOwnerGroupPolicy || ruleName != "{GP-RULE}" {
|
||||
t.Fatalf("owner = %v, rule = %q", owner, ruleName)
|
||||
}
|
||||
|
||||
state.setNRPTPolicyOwner(nrptRuleOwnerCtrld, "")
|
||||
owner, ruleName = state.nrptPolicyOwner()
|
||||
if owner != nrptRuleOwnerCtrld || ruleName != "" {
|
||||
t.Fatalf("owner = %v, rule = %q", owner, ruleName)
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,101 @@
|
||||
//go:build windows
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
const (
|
||||
// Default to current behavior: keep recovering indefinitely unless configured.
|
||||
defaultNRPTRecoveryMaxAttempts = 0
|
||||
defaultNRPTRecoveryCooldown = 30 * time.Minute
|
||||
|
||||
// Require more than one good health tick before clearing the circuit. A probe can
|
||||
// pass briefly after delete/re-add even when another agent recreates broken NRPT state.
|
||||
nrptRecoveryStableSuccessesToReset = 2
|
||||
)
|
||||
|
||||
type nrptRecoveryLimiter struct {
|
||||
mu sync.Mutex
|
||||
attempts int
|
||||
stableSuccesses int
|
||||
cooldownUntil time.Time
|
||||
lastSkipLog time.Time
|
||||
}
|
||||
|
||||
func nrptRecoveryMaxAttempts(cfg *ctrld.Config) int {
|
||||
if cfg != nil && cfg.Service.NRPTRecoveryMaxAttempts != nil {
|
||||
return *cfg.Service.NRPTRecoveryMaxAttempts
|
||||
}
|
||||
return defaultNRPTRecoveryMaxAttempts
|
||||
}
|
||||
|
||||
func nrptRecoveryCooldown(cfg *ctrld.Config) time.Duration {
|
||||
if cfg != nil && cfg.Service.NRPTRecoveryCooldown != nil {
|
||||
return *cfg.Service.NRPTRecoveryCooldown
|
||||
}
|
||||
return defaultNRPTRecoveryCooldown
|
||||
}
|
||||
|
||||
func (l *nrptRecoveryLimiter) allow(now time.Time, cfg *ctrld.Config) (bool, time.Duration) {
|
||||
maxAttempts := nrptRecoveryMaxAttempts(cfg)
|
||||
if maxAttempts <= 0 {
|
||||
return true, 0
|
||||
}
|
||||
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
if now.Before(l.cooldownUntil) {
|
||||
return false, l.cooldownUntil.Sub(now)
|
||||
}
|
||||
return true, 0
|
||||
}
|
||||
|
||||
func (l *nrptRecoveryLimiter) recordRecoveryFlow(now time.Time, cfg *ctrld.Config) {
|
||||
maxAttempts := nrptRecoveryMaxAttempts(cfg)
|
||||
if maxAttempts <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
cooldown := nrptRecoveryCooldown(cfg)
|
||||
if cooldown <= 0 {
|
||||
cooldown = defaultNRPTRecoveryCooldown
|
||||
}
|
||||
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
l.stableSuccesses = 0
|
||||
l.attempts++
|
||||
if l.attempts >= maxAttempts {
|
||||
l.cooldownUntil = now.Add(cooldown)
|
||||
}
|
||||
}
|
||||
|
||||
func (l *nrptRecoveryLimiter) recordStableSuccess() {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
l.stableSuccesses++
|
||||
if l.stableSuccesses >= nrptRecoveryStableSuccessesToReset {
|
||||
l.attempts = 0
|
||||
l.cooldownUntil = time.Time{}
|
||||
l.lastSkipLog = time.Time{}
|
||||
}
|
||||
}
|
||||
|
||||
func (l *nrptRecoveryLimiter) shouldLogSkip(now time.Time) bool {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
if l.lastSkipLog.IsZero() || now.Sub(l.lastSkipLog) >= 5*time.Minute {
|
||||
l.lastSkipLog = now
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
//go:build windows
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
func TestNRPTRecoveryLimiterCooldownAndStableReset(t *testing.T) {
|
||||
maxAttempts := 2
|
||||
cooldown := 10 * time.Minute
|
||||
cfg := &ctrld.Config{}
|
||||
cfg.Service.NRPTRecoveryMaxAttempts = &maxAttempts
|
||||
cfg.Service.NRPTRecoveryCooldown = &cooldown
|
||||
|
||||
limiter := &nrptRecoveryLimiter{}
|
||||
now := time.Unix(100, 0)
|
||||
|
||||
if ok, wait := limiter.allow(now, cfg); !ok || wait != 0 {
|
||||
t.Fatalf("initial allow = %v, %v; want true, 0", ok, wait)
|
||||
}
|
||||
|
||||
limiter.recordRecoveryFlow(now, cfg)
|
||||
if ok, wait := limiter.allow(now.Add(time.Second), cfg); !ok || wait != 0 {
|
||||
t.Fatalf("allow after first flow = %v, %v; want true, 0", ok, wait)
|
||||
}
|
||||
|
||||
limiter.recordRecoveryFlow(now.Add(2*time.Second), cfg)
|
||||
if ok, wait := limiter.allow(now.Add(3*time.Second), cfg); ok || wait <= 0 {
|
||||
t.Fatalf("allow after max flows = %v, %v; want false, positive wait", ok, wait)
|
||||
}
|
||||
|
||||
limiter.recordStableSuccess()
|
||||
if ok, _ := limiter.allow(now.Add(4*time.Second), cfg); ok {
|
||||
t.Fatal("one stable success cleared cooldown; want cooldown to remain")
|
||||
}
|
||||
|
||||
limiter.recordStableSuccess()
|
||||
if ok, wait := limiter.allow(now.Add(5*time.Second), cfg); !ok || wait != 0 {
|
||||
t.Fatalf("allow after stable reset = %v, %v; want true, 0", ok, wait)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNRPTRecoveryLimiterDefaultIsUnlimited(t *testing.T) {
|
||||
cfg := &ctrld.Config{}
|
||||
limiter := &nrptRecoveryLimiter{}
|
||||
now := time.Unix(100, 0)
|
||||
|
||||
for i := 0; i < 10; i++ {
|
||||
limiter.recordRecoveryFlow(now.Add(time.Duration(i)*time.Second), cfg)
|
||||
}
|
||||
if ok, wait := limiter.allow(now.Add(time.Hour), cfg); !ok || wait != 0 {
|
||||
t.Fatalf("default allow after recovery flows = %v, %v; want true, 0", ok, wait)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNRPTRecoveryLimiterUnlimited(t *testing.T) {
|
||||
maxAttempts := 0
|
||||
cfg := &ctrld.Config{}
|
||||
cfg.Service.NRPTRecoveryMaxAttempts = &maxAttempts
|
||||
|
||||
limiter := &nrptRecoveryLimiter{}
|
||||
now := time.Unix(100, 0)
|
||||
|
||||
for i := 0; i < 10; i++ {
|
||||
limiter.recordRecoveryFlow(now.Add(time.Duration(i)*time.Second), cfg)
|
||||
}
|
||||
if ok, wait := limiter.allow(now.Add(time.Hour), cfg); !ok || wait != 0 {
|
||||
t.Fatalf("unlimited allow = %v, %v; want true, 0", ok, wait)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// pfNoRulesMarker is what pfctl prints for a ruleset that contains nothing.
|
||||
const pfNoRulesMarker = "(no rules)"
|
||||
|
||||
// pfFilterRuleLines reduces pfctl output to the lines that are actually pf rules.
|
||||
//
|
||||
// It exists because every pfctl reader here uses CombinedOutput, and pfctl on macOS
|
||||
// writes "No ALTQ support in kernel" and "ALTQ related functions disabled" to stderr on
|
||||
// essentially every show command, so raw output is never a clean rule list. An empty
|
||||
// ruleset can also report "(no rules)", which is a status line rather than a rule.
|
||||
//
|
||||
// Two consequences follow from getting this wrong, and both have bitten this file:
|
||||
// callers that test the output for emptiness can never see empty, and callers that feed
|
||||
// the lines back into "pfctl -f -" would splice non-rule text into a ruleset and have
|
||||
// the reload rejected.
|
||||
//
|
||||
// Registry access and platform specifics stay elsewhere; this is pure string handling
|
||||
// so it can be tested on any host.
|
||||
func pfFilterRuleLines(output string) []string {
|
||||
var rules []string
|
||||
for _, line := range strings.Split(output, "\n") {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
// pfctl stderr warnings, merged in by CombinedOutput.
|
||||
if strings.Contains(line, "ALTQ") {
|
||||
continue
|
||||
}
|
||||
// Status line for an empty ruleset, not a rule.
|
||||
if line == pfNoRulesMarker {
|
||||
continue
|
||||
}
|
||||
rules = append(rules, line)
|
||||
}
|
||||
return rules
|
||||
}
|
||||
|
||||
// pfRulesetEmpty reports whether pfctl output describes a ruleset with no rules.
|
||||
//
|
||||
// Use this rather than testing the raw output for emptiness: the merged stderr warnings
|
||||
// described above mean a raw test is always false, so the condition it guards - an
|
||||
// anchor whose contents were flushed - would never be detected.
|
||||
func pfRulesetEmpty(output string) bool {
|
||||
return len(pfFilterRuleLines(output)) == 0
|
||||
}
|
||||
|
||||
// pfContainsRule checks if any line in the slice contains the given rule string.
|
||||
// Uses substring matching because pfctl may append extra tokens like " all" to rules
|
||||
// (e.g., `rdr-anchor "com.controld.ctrld" all`), which would fail exact matching.
|
||||
func pfContainsRule(lines []string, rule string) bool {
|
||||
for _, line := range lines {
|
||||
if strings.Contains(line, rule) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// pfAnchorReferencesPresent reports whether ctrld's anchor references appear in the
|
||||
// running ruleset, given the output of "pfctl -sn" and "pfctl -sr".
|
||||
//
|
||||
// Removing the references means reloading the entire main ruleset, and that reload
|
||||
// carries no options section - so it resets system-wide pf options, including any
|
||||
// third-party "set skip" directives. Doing that when there is nothing of ours to
|
||||
// remove is pure collateral damage, which is what a startup rollback would otherwise
|
||||
// cause after failing before the references were ever added.
|
||||
func pfAnchorReferencesPresent(natOutput, filterOutput, anchorName string) bool {
|
||||
rdrAnchorRef := fmt.Sprintf("rdr-anchor %q", anchorName)
|
||||
anchorRef := fmt.Sprintf("anchor %q", anchorName)
|
||||
return pfContainsRule(pfFilterRuleLines(natOutput), rdrAnchorRef) ||
|
||||
pfContainsRule(pfFilterRuleLines(filterOutput), anchorRef)
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
package cli
|
||||
|
||||
import "testing"
|
||||
|
||||
// altqNoise is what macOS pfctl writes to stderr on show commands. Because every
|
||||
// pfctl reader here uses CombinedOutput, it lands in the middle of the data being
|
||||
// parsed — which is why these helpers exist.
|
||||
const altqNoise = "No ALTQ support in kernel\nALTQ related functions disabled\n"
|
||||
|
||||
// TestPFRulesetEmpty is the regression guard for a flushed anchor being undetectable.
|
||||
//
|
||||
// The anchor-content checks in verifyPFState and ensurePFAnchorActive decide whether pf
|
||||
// still has ctrld's rules. Testing the raw pfctl output for emptiness can never be true
|
||||
// on macOS, because the merged ALTQ warnings are always present — so a genuinely flushed
|
||||
// anchor reads as healthy and neither the startup gate nor the watchdog restore fires.
|
||||
func TestPFRulesetEmpty(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
output string
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
// The case that was broken: nothing but merged stderr.
|
||||
name: "only ALTQ warnings",
|
||||
output: altqNoise,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
// As captured on macOS 26.6 from "pfctl -sn -a com.controld.ctrld".
|
||||
name: "ALTQ warnings plus the empty-ruleset marker",
|
||||
output: altqNoise + "(no rules)\n",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "empty output",
|
||||
output: "",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "whitespace only",
|
||||
output: "\n \n\t\n",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "a real rdr rule behind the warnings",
|
||||
output: altqNoise + "rdr on lo0 inet proto udp from any to ! 127.0.0.1 port = 53 -> 127.0.0.1 port 5354\n",
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "a real filter rule behind the warnings",
|
||||
output: altqNoise + "pass in quick on lo0 reply-to lo0 inet proto udp from any to 127.0.0.1 port = 5354\n",
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "rule with no warnings at all",
|
||||
output: "anchor \"com.controld.ctrld\" all\n",
|
||||
want: false,
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := pfRulesetEmpty(tc.output); got != tc.want {
|
||||
t.Errorf("pfRulesetEmpty() = %v, want %v\noutput:\n%s", got, tc.want, tc.output)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestPFFilterRuleLines checks what survives filtering, since these lines are fed back
|
||||
// into "pfctl -f -" by the ruleset-rebuild paths. Splicing a warning or the
|
||||
// empty-ruleset marker into a ruleset would have the reload rejected outright.
|
||||
func TestPFFilterRuleLines(t *testing.T) {
|
||||
got := pfFilterRuleLines(altqNoise + "(no rules)\nrdr-anchor \"com.controld.ctrld\" all\n\nanchor \"com.controld.ctrld\" all\n")
|
||||
want := []string{
|
||||
`rdr-anchor "com.controld.ctrld" all`,
|
||||
`anchor "com.controld.ctrld" all`,
|
||||
}
|
||||
if len(got) != len(want) {
|
||||
t.Fatalf("got %d lines %q, want %d %q", len(got), got, len(want), want)
|
||||
}
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Errorf("line %d = %q, want %q", i, got[i], want[i])
|
||||
}
|
||||
}
|
||||
|
||||
if lines := pfFilterRuleLines(altqNoise); lines != nil {
|
||||
t.Errorf("warnings alone must yield no rule lines, got %q", lines)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPFAnchorReferencesPresent guards when the main ruleset may be rewritten.
|
||||
//
|
||||
// Removing our anchor references means reloading the whole main ruleset, and that
|
||||
// reload carries no options section — so it resets system-wide pf options, including
|
||||
// third-party "set skip" directives. Startup rollback runs after failures that happen
|
||||
// before the references were ever added, so without this check it would reset another
|
||||
// application's pf options while removing nothing of ours.
|
||||
func TestPFAnchorReferencesPresent(t *testing.T) {
|
||||
const anchor = "com.controld.ctrld"
|
||||
const otherAppRules = "scrub-anchor \"com.apple/*\" all fragment reassemble\nanchor \"com.vendor.vpn\" all\n"
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
nat string
|
||||
filter string
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "both references present",
|
||||
nat: altqNoise + "rdr-anchor \"com.controld.ctrld\" all\n",
|
||||
filter: altqNoise + "anchor \"com.controld.ctrld\" all\n",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
// pfctl appends tokens like " all", so matching is substring-based.
|
||||
name: "rdr reference only",
|
||||
nat: altqNoise + "rdr-anchor \"com.controld.ctrld\" all\n",
|
||||
filter: altqNoise + otherAppRules,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "filter reference only",
|
||||
nat: altqNoise,
|
||||
filter: altqNoise + "anchor \"com.controld.ctrld\"\n",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
// The rollback case: we failed before adding anything, and another
|
||||
// application owns the ruleset. Rewriting it would be pure collateral.
|
||||
name: "someone else's ruleset, none of ours",
|
||||
nat: altqNoise,
|
||||
filter: altqNoise + otherAppRules,
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "empty ruleset",
|
||||
nat: altqNoise + "(no rules)\n",
|
||||
filter: altqNoise + "(no rules)\n",
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
// A different anchor whose name merely contains ours must not count.
|
||||
name: "another anchor with a similar name",
|
||||
nat: altqNoise,
|
||||
filter: altqNoise + "anchor \"com.vendor.controld-shim\" all\n",
|
||||
want: false,
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := pfAnchorReferencesPresent(tc.nat, tc.filter, anchor); got != tc.want {
|
||||
t.Errorf("pfAnchorReferencesPresent() = %v, want %v", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
+332
-38
@@ -34,6 +34,7 @@ import (
|
||||
"github.com/Control-D-Inc/ctrld/internal/clientinfo"
|
||||
"github.com/Control-D-Inc/ctrld/internal/controld"
|
||||
"github.com/Control-D-Inc/ctrld/internal/dnscache"
|
||||
ctrldnet "github.com/Control-D-Inc/ctrld/internal/net"
|
||||
"github.com/Control-D-Inc/ctrld/internal/router"
|
||||
"github.com/Control-D-Inc/ctrld/internal/router/dnsmasq"
|
||||
)
|
||||
@@ -91,6 +92,16 @@ var svcConfig = &service.Config{
|
||||
|
||||
var useSystemdResolved = false
|
||||
|
||||
type pfAnchorCheckResult uint8
|
||||
|
||||
const (
|
||||
pfAnchorCheckSkipped pfAnchorCheckResult = iota
|
||||
pfAnchorCheckIntact
|
||||
pfAnchorCheckRestored
|
||||
pfAnchorCheckDeferred
|
||||
pfAnchorCheckFailed
|
||||
)
|
||||
|
||||
type prog struct {
|
||||
mu sync.Mutex
|
||||
waitCh chan struct{}
|
||||
@@ -145,6 +156,10 @@ type prog struct {
|
||||
recoveryCancelMu sync.Mutex
|
||||
recoveryCancel context.CancelFunc
|
||||
recoveryRunning atomic.Bool
|
||||
// recoveryGen counts handleRecovery invocations that reached the
|
||||
// recovery-context stage; each recovery captures its own generation and
|
||||
// only touches shared recovery state if it is still the newest (#597).
|
||||
recoveryGen atomic.Uint64
|
||||
|
||||
// recoveryDebounceTimer coalesces rapid NetworkChange recovery triggers
|
||||
// into a single handleRecovery call. Only handleRecovery is debounced —
|
||||
@@ -157,15 +172,47 @@ type prog struct {
|
||||
// instead of using the normal upstream flow.
|
||||
recoveryBypass atomic.Bool
|
||||
|
||||
// interceptDNSTargetService names the macOS network service on which
|
||||
// ctrld set a loopback DNS value because the service provided no usable
|
||||
// IPv4 DNS while DNS intercept mode was active (issue #533);
|
||||
// interceptDNSTargetSetValue records the exact value set. Both empty when
|
||||
// no target is set. Guarded by interceptDNSTargetMu.
|
||||
//
|
||||
//lint:ignore U1000 used in Darwin code.
|
||||
interceptDNSTargetMu sync.Mutex
|
||||
//lint:ignore U1000 used in Darwin code.
|
||||
interceptDNSTargetService string
|
||||
//lint:ignore U1000 used in Darwin code.
|
||||
interceptDNSTargetSetValue string
|
||||
//lint:ignore U1000 used in Darwin code.
|
||||
interceptDNSTargetLoaded bool
|
||||
|
||||
// DNS intercept mode state (platform-specific).
|
||||
// On Windows: *wfpState, on macOS: *pfState, nil on other platforms.
|
||||
dnsInterceptState any
|
||||
|
||||
// lastTunnelIfaces tracks the set of active VPN/tunnel interfaces (utun*, ipsec*, etc.)
|
||||
// discovered during the last pf anchor rule build. When the set changes (e.g., a VPN
|
||||
// connects and creates utun420), we rebuild the pf anchor to add interface-specific
|
||||
// intercept rules for the new interface. Protected by mu.
|
||||
lastTunnelIfaces []string //lint:ignore U1000 used on darwin
|
||||
// dnsInterceptMu serializes DNS intercept lifecycle transitions - start, stop and
|
||||
// the health monitor's rebuild - and guards every write to dnsInterceptState, so a
|
||||
// service stop can never interleave with a monitor-driven rebuild.
|
||||
dnsInterceptMu sync.Mutex //lint:ignore U1000 used on windows
|
||||
|
||||
// dnsInterceptStopRequested is set while a stop waits for dnsInterceptMu. The
|
||||
// health and recovery flows read it as a shutdown signal and abandon their work,
|
||||
// rather than making the stop wait out their probe backoffs.
|
||||
dnsInterceptStopRequested atomic.Bool //lint:ignore U1000 used on windows
|
||||
|
||||
// nrptTransitionMu makes one NRPT ownership transition - observe, mutate, signal,
|
||||
// record owner - atomic against shutdown and against another transition. It is
|
||||
// deliberately finer-grained than dnsInterceptMu: it is taken for the duration of a
|
||||
// single transition, never across the recovery flows' probe backoffs.
|
||||
nrptTransitionMu sync.Mutex //lint:ignore U1000 used on windows
|
||||
|
||||
// lastTunnelIfaces tracks the tunnel set included in the last successfully loaded
|
||||
// pf anchor. Pending tunnel state is kept separately so failed PF work is retried
|
||||
// instead of being mistaken for an applied update. Protected by mu.
|
||||
lastTunnelIfaces []string //lint:ignore U1000 used on darwin
|
||||
pendingTunnelIfaces []string //lint:ignore U1000 used on darwin
|
||||
hasPendingTunnelIfaces bool //lint:ignore U1000 used on darwin
|
||||
|
||||
// pfStabilizing is true while we're waiting for a VPN's pf ruleset to settle.
|
||||
// While true, the watchdog and network change callbacks do NOT restore our rules.
|
||||
@@ -188,15 +235,42 @@ type prog struct {
|
||||
// interception with exponential backoff and auto-heals if broken.
|
||||
pfMonitorRunning atomic.Bool //lint:ignore U1000 used on darwin
|
||||
|
||||
// pfProbeExpected holds the domain name of a pending pf interception probe.
|
||||
// When non-empty, the DNS handler checks incoming queries against this value
|
||||
// and signals pfProbeCh if matched. The probe verifies that pf's rdr rules
|
||||
// are actually translating packets (not just present in rule text).
|
||||
pfProbeExpected atomic.Value // string
|
||||
// pfEnsureRunning ensures only one pf validation or mutation runs at a time.
|
||||
// Network callbacks, VPN exemption updates, delayed rechecks, probes, and the
|
||||
// watchdog can converge during macOS churn; concurrent pfctl/scutil work can
|
||||
// exhaust process/file limits or interleave anchor snapshots.
|
||||
pfEnsureRunning atomic.Bool //lint:ignore U1000 used on darwin
|
||||
|
||||
// pfProbeCh is signaled when the DNS handler receives the expected probe query.
|
||||
// The channel is created by probePFIntercept() and closed when the probe arrives.
|
||||
pfProbeCh atomic.Value // *chan struct{}
|
||||
// pfExecBackoffUntil suppresses pf anchor validation after pfctl/scutil execs
|
||||
// fail due host resource exhaustion (fork unavailable, too many open files).
|
||||
pfExecBackoffUntil atomic.Int64 //lint:ignore U1000 used on darwin
|
||||
|
||||
// pfDelayedRecheckTimers coalesces delayed DNS-intercept rechecks after noisy
|
||||
// network changes. Protected by pfDelayedRecheckMu.
|
||||
pfDelayedRecheckMu sync.Mutex //lint:ignore U1000 used on darwin
|
||||
pfDelayedRecheckTimers []*time.Timer //lint:ignore U1000 used on darwin
|
||||
|
||||
// pfIgnoredChangeLastReconcile bounds immediate pf/VPN-DNS work for noisy
|
||||
// ignored macOS network deltas. Tunnel changes bypass this limit, and the
|
||||
// existing delayed checks provide a trailing reconciliation after churn.
|
||||
pfIgnoredChangeLastReconcile atomic.Int64 //lint:ignore U1000 used on darwin
|
||||
|
||||
// interceptProbes maps the domain of each pending interception probe to the channel
|
||||
// that probe waits on. A probe verifies that interception is actually translating or
|
||||
// redirecting packets, not merely present in rule text: the DNS handler looks up
|
||||
// incoming queries here and signals the matching waiter.
|
||||
//
|
||||
// It holds one entry per in-flight probe rather than a single slot, because probes do
|
||||
// overlap - the health monitor, a handback and a heal cycle can each have one out at
|
||||
// the same time - and a single slot means the last registration wins and the loser
|
||||
// waits out its timeout for a query that was answered. A false failure then triggers
|
||||
// recovery work that was not needed.
|
||||
//
|
||||
// Registrations are rare and lookups happen on every query, so the map is stored as
|
||||
// an immutable snapshot behind an atomic: readers never take a lock, writers copy
|
||||
// under interceptProbeMu.
|
||||
interceptProbes atomic.Value // map[string]chan struct{}
|
||||
interceptProbeMu sync.Mutex //lint:ignore U1000 written only by registerInterceptProbe, used on darwin/windows
|
||||
|
||||
// VPN DNS manager for split DNS routing when intercept mode is active.
|
||||
vpnDNS *vpnDNSManager
|
||||
@@ -269,7 +343,7 @@ func (p *prog) runWait() {
|
||||
continue
|
||||
}
|
||||
if cdUID != "" {
|
||||
rc, err := processCDFlags(newCfg)
|
||||
rc, err := p.fetchCDConfigBoundedByLifetime(newCfg)
|
||||
if err != nil {
|
||||
logger.Err(err).Msg("could not fetch ControlD config")
|
||||
waitOldRunDone()
|
||||
@@ -321,6 +395,18 @@ func (p *prog) runWait() {
|
||||
|
||||
p.mu.Lock()
|
||||
*p.cfg = *newCfg
|
||||
// In DNS-intercept mode on macOS, the DNS listener is bound once at startup and is
|
||||
// NOT re-bound on reload (see prog.run: serveDNS is started only when !reload). When
|
||||
// the configured/generated port (e.g. 127.0.0.1:53) is unavailable at startup because
|
||||
// mDNSResponder owns *:53, ctrld falls back to an alternate local port (e.g. 5354).
|
||||
// The on-disk config still declares 53, so adopting it here would revert p.cfg to a
|
||||
// port nothing is listening on, and the pf rdr rules/probes rebuilt from p.cfg would
|
||||
// target a dead port. Since a reload cannot move the running listener anyway, keep
|
||||
// p.cfg pointing at the actual bound listener. The on-disk config (written above) is
|
||||
// left unchanged. See #551.
|
||||
if dnsIntercept && runtime.GOOS == "darwin" {
|
||||
preserveBoundListeners(p.cfg.Listener, curListener)
|
||||
}
|
||||
p.mu.Unlock()
|
||||
|
||||
logger.Notice().Msg("reloading config successfully")
|
||||
@@ -332,8 +418,41 @@ func (p *prog) runWait() {
|
||||
}
|
||||
}
|
||||
|
||||
// preserveBoundListeners overrides the IP/Port of each listener in newListeners with the
|
||||
// actual bound address from curListeners when they differ, logging the divergence. It is used
|
||||
// on config reload in DNS-intercept mode where the running listener is never re-bound, so a
|
||||
// port change on disk (e.g. reverting a fallback 5354 back to the generated 53) must not be
|
||||
// applied to the in-memory config that drives pf rdr rules and probes.
|
||||
//
|
||||
// Preservation is limited to fallback-eligible (default/unset, i.e. 127.0.0.1:53) listeners.
|
||||
// An explicit, non-default listener in the reloaded config is an intentional change that must
|
||||
// be applied: tryUpdateListenerConfigIntercept binds explicit listeners exactly (no fallback),
|
||||
// and the control-server reload handler detects the IP/port diff to trigger a restart that
|
||||
// re-binds. Reverting an explicit change here would make that comparison return 200 instead of
|
||||
// 201, silently dropping the new listener. See #551.
|
||||
func preserveBoundListeners(newListeners, curListeners map[string]*ctrld.ListenerConfig) {
|
||||
for n, curLc := range curListeners {
|
||||
newLc := newListeners[n]
|
||||
if newLc == nil || curLc == nil {
|
||||
continue
|
||||
}
|
||||
if newLc.IP == curLc.IP && newLc.Port == curLc.Port {
|
||||
continue
|
||||
}
|
||||
if isExplicitInterceptListener(newLc.IP, newLc.Port) {
|
||||
continue
|
||||
}
|
||||
mainLog.Load().Info().
|
||||
Str("configured", net.JoinHostPort(newLc.IP, strconv.Itoa(newLc.Port))).
|
||||
Str("actual", net.JoinHostPort(curLc.IP, strconv.Itoa(curLc.Port))).
|
||||
Msg("DNS intercept: preserving actual bound listener across reload; on-disk config port not applied to running listener")
|
||||
newLc.IP = curLc.IP
|
||||
newLc.Port = curLc.Port
|
||||
}
|
||||
}
|
||||
|
||||
func (p *prog) preRun() {
|
||||
if iface == "auto" {
|
||||
if iface == autoIface {
|
||||
iface = defaultIfaceName()
|
||||
p.requiredMultiNICsConfig = requiredMultiNICsConfig()
|
||||
}
|
||||
@@ -347,10 +466,15 @@ func (p *prog) postRun() {
|
||||
p.runningOnDomainController = isDC
|
||||
mainLog.Load().Debug().Msgf("running on domain controller: %t, role: %d", p.runningOnDomainController, roleInt)
|
||||
}
|
||||
p.resetDNS(false, false)
|
||||
ns := ctrld.InitializeOsResolver(false)
|
||||
// A Windows organization can install a GP-owned NRPT catch-all before
|
||||
// starting ctrld. Detect that policy before resetDNS touches adapter DNS;
|
||||
// startDNSIntercept will then prove the rule functionally before adopting it.
|
||||
if !p.skipInitialDNSReset() {
|
||||
p.resetDNS(false, false)
|
||||
}
|
||||
ns, systemNameservers := initializeOsResolverWithSystemNameserversFn(false)
|
||||
mainLog.Load().Debug().Msgf("initialized OS resolver with nameservers: %v", ns)
|
||||
p.setDNS()
|
||||
p.setDNS(systemNameservers)
|
||||
p.csSetDnsDone <- struct{}{}
|
||||
close(p.csSetDnsDone)
|
||||
p.logInterfacesState()
|
||||
@@ -386,7 +510,7 @@ func (p *prog) apiConfigReload() {
|
||||
Version: rootCmd.Version,
|
||||
Metadata: ctrld.SystemMetadataRuntime(context.Background()),
|
||||
}
|
||||
resolverConfig, err := controld.FetchResolverConfig(req, cdDev)
|
||||
resolverConfig, err := controld.FetchResolverConfig(context.Background(), req, cdDev)
|
||||
selfUninstallCheck(err, p, logger)
|
||||
if err != nil {
|
||||
logger.Warn().Err(err).Msg("could not fetch resolver config")
|
||||
@@ -444,7 +568,7 @@ func (p *prog) apiConfigReload() {
|
||||
}
|
||||
if cfgErr != nil {
|
||||
logger.Warn().Err(err).Msg("skipping invalid custom config")
|
||||
if _, err := controld.UpdateCustomLastFailed(cdUID, rootCmd.Version, cdDev, true); err != nil {
|
||||
if _, err := controld.UpdateCustomLastFailed(context.Background(), cdUID, rootCmd.Version, cdDev, true); err != nil {
|
||||
logger.Error().Err(err).Msg("could not mark custom last update failed")
|
||||
}
|
||||
return
|
||||
@@ -778,7 +902,49 @@ func (p *prog) deAllocateIP() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *prog) setDNS() {
|
||||
// Seams for the intercept-start failure lifecycle. Choosing between the interface-DNS
|
||||
// fallback and refusing it has side effects - restoring the host's DNS, then
|
||||
// terminating - which a test has to observe without reconfiguring the host or exiting
|
||||
// the test binary. The intercept start itself is indirected for the same reason: it is
|
||||
// the real platform interceptor, which on macOS mutates pf and on Windows installs an
|
||||
// NRPT rule, so a test of what happens *after* it fails must not be the thing that
|
||||
// runs it.
|
||||
var (
|
||||
localResolverIPFn = router.LocalResolverIP
|
||||
startDNSInterceptFn = (*prog).startDNSIntercept
|
||||
ensureInterceptDNSTargetFn = (*prog).ensureInterceptDNSTarget
|
||||
removeInterceptDNSTargetFn = (*prog).removeInterceptDNSTarget
|
||||
initializeOsResolverWithSystemNameserversFn = ctrld.InitializeOsResolverWithSystemNameservers
|
||||
setDnsForRunningIfaceFn = (*prog).setDnsForRunningIface
|
||||
resetDNSFn = (*prog).resetDNS
|
||||
refuseFallbackFatal = func(format string, v ...any) {
|
||||
mainLog.Load().Fatal().Msgf(format, v...)
|
||||
}
|
||||
)
|
||||
|
||||
// interfaceDNSFallbackViable reports whether the interface-DNS fallback can actually
|
||||
// direct queries to ctrld's listener.
|
||||
//
|
||||
// Interface DNS names a resolver by IP and has no port field - true of macOS interface
|
||||
// settings and of Windows NRPT rules - so pointing the system straight at a listener
|
||||
// that did not bind :53 sends queries to whatever owns :53 instead, and that resolver's
|
||||
// upstream is ctrld's address: a loop, not a fallback.
|
||||
//
|
||||
// A nil or portless listener is treated as viable: the port is resolved elsewhere and
|
||||
// defaults to 53, so there is nothing to refuse yet.
|
||||
//
|
||||
// A non-53 listener is still viable where a local resolver owns :53 and forwards to
|
||||
// ctrld's port. That is the arrangement on the router platforms with a dnsmasq of their
|
||||
// own: ctrld writes "server=<listener ip>#<listener port>", so the forward follows
|
||||
// whatever port ctrld actually bound. setDNS then points the interface at that resolver
|
||||
// rather than at the listener - see the lc.Port != 53 case there, which this mirrors.
|
||||
// Refusing on port alone would turn a working configuration into a startup failure on
|
||||
// those routers.
|
||||
func interfaceDNSFallbackViable(lc *ctrld.ListenerConfig, localResolverIP string) bool {
|
||||
return lc == nil || lc.Port == 0 || lc.Port == 53 || localResolverIP != ""
|
||||
}
|
||||
|
||||
func (p *prog) setDNS(systemNameservers []string) {
|
||||
setDnsOK := false
|
||||
defer func() {
|
||||
p.csSetDnsOk = setDnsOK
|
||||
@@ -786,12 +952,13 @@ func (p *prog) setDNS() {
|
||||
|
||||
// Validate and resolve intercept mode.
|
||||
// CLI flag (--intercept-mode) takes priority over config file.
|
||||
// Valid values: "" (off), "dns" (with VPN split routing), "hard" (all DNS through ctrld).
|
||||
// Valid values: "" (use config), "off" (explicitly disable), "dns" (with VPN
|
||||
// split routing), and "hard" (all DNS through ctrld).
|
||||
if interceptMode != "" && !validInterceptMode(interceptMode) {
|
||||
mainLog.Load().Fatal().Msgf("invalid --intercept-mode value %q: must be 'off', 'dns', or 'hard'", interceptMode)
|
||||
}
|
||||
if interceptMode == "" || interceptMode == "off" {
|
||||
interceptMode = cfg.Service.InterceptMode
|
||||
if interceptMode == "" {
|
||||
interceptMode = p.configuredInterceptMode()
|
||||
if interceptMode != "" && interceptMode != "off" {
|
||||
mainLog.Load().Info().Msgf("Intercept mode enabled via config (intercept_mode = %q)", interceptMode)
|
||||
}
|
||||
@@ -810,10 +977,64 @@ func (p *prog) setDNS() {
|
||||
// modifying interface DNS settings. This eliminates race conditions with VPN
|
||||
// software that also manages DNS. See issue #489.
|
||||
if dnsIntercept {
|
||||
if err := p.startDNSIntercept(); err != nil {
|
||||
if err := startDNSInterceptFn(p); err != nil {
|
||||
removeInterceptDNSTargetFn(p, "DNS intercept unavailable")
|
||||
// This check comes first: it is the one failure where DNS already works
|
||||
// without ctrld touching anything else, so neither the refusal below nor the
|
||||
// fallback applies.
|
||||
//
|
||||
// An externally managed rule was proved - by probe, not by registry shape -
|
||||
// to be routing DNS to this listener. Falling through would rewrite adapter
|
||||
// DNS after explicitly preserving it, and DNS still works, so stop here.
|
||||
//
|
||||
// Only a verified route earns this. A rule that merely exists does not: if it
|
||||
// is not actually routing and intercept failed too, the machine would be left
|
||||
// with no NRPT, no WFP and no adapter fallback - that is, unfiltered - so
|
||||
// every other failure takes the paths below.
|
||||
if interceptFailedUnderExternalDNSPolicy(err) {
|
||||
if interceptFailedWithVerifiedExternalDNS(err) {
|
||||
mainLog.Load().Error().Err(err).Msg("DNS intercept mode failed but externally managed DNS policy is verified routing to ctrld — not falling back to interface DNS settings")
|
||||
} else {
|
||||
// Owned by external policy but not proved to route: DNS is not
|
||||
// reaching ctrld. Adapter DNS still stays as the organization set it,
|
||||
// and setDnsOK stays false, so this start reports as failed until a
|
||||
// probe succeeds.
|
||||
mainLog.Load().Error().Err(err).Msg("DNS intercept mode failed and externally managed DNS policy is not routing to ctrld — leaving interface DNS settings untouched; the service is not ready")
|
||||
}
|
||||
return
|
||||
}
|
||||
// Interface DNS cannot express a port: macOS interface settings and Windows
|
||||
// NRPT rules both name a resolver by IP alone. So it is only a usable
|
||||
// fallback when the listener actually bound :53. When something else owns
|
||||
// :53 - mDNSResponder on macOS, which is the whole reason the :5354 fallback
|
||||
// exists - pointing the system at 127.0.0.1 hands queries to that other
|
||||
// resolver, whose own upstream is now ctrld's address. That is a resolution
|
||||
// loop, not degraded operation: a healthy ctrld listener nothing on the host
|
||||
// can reach, no working DNS, and no recovery short of stopping the service.
|
||||
//
|
||||
// Refuse instead, after putting the host's own DNS back. A visible startup
|
||||
// failure beats DNS that is broken by design, and it stops a fallback that
|
||||
// cannot work from quietly undoing the fail-closed verification above.
|
||||
if lc := cfg.FirstListener(); !interfaceDNSFallbackViable(lc, localResolverIPFn()) {
|
||||
mainLog.Load().Error().Err(err).Msgf("DNS intercept mode failed with the listener on port %d", lc.Port)
|
||||
// Leave the host resolvable: restore static settings or DHCP rather than
|
||||
// exiting with an interface still pointed at a ctrld that is not serving.
|
||||
resetDNSFn(p, false, true)
|
||||
refuseFallbackFatal("Refusing to fall back to interface DNS: it cannot direct queries to %s:%d, which would leave this host with no working resolver. Free port 53 for ctrld, or resolve the intercept failure, then start again.", lc.IP, lc.Port)
|
||||
// Unreachable in production - the line above exits - but returning
|
||||
// explicitly keeps the refusal from depending on that, so nothing can
|
||||
// fall through to installing the fallback this just rejected.
|
||||
return
|
||||
}
|
||||
mainLog.Load().Error().Err(err).Msg("DNS intercept mode failed — falling back to interface DNS settings")
|
||||
// Fall through to traditional setDNS behavior.
|
||||
} else {
|
||||
// Intercept installation alone is insufficient on a DNS-less network:
|
||||
// without an IPv4 DNS target macOS emits no packet for pf to redirect.
|
||||
// Do this on startup as well as network-change recovery so starting or
|
||||
// restarting while already tethered cannot leave DNS offline.
|
||||
ensureInterceptDNSTargetFn(p, systemNameservers)
|
||||
|
||||
if hardIntercept {
|
||||
mainLog.Load().Info().Msg("Hard intercept mode active — all DNS through ctrld, no VPN split routing")
|
||||
} else {
|
||||
@@ -831,6 +1052,9 @@ func (p *prog) setDNS() {
|
||||
return
|
||||
}
|
||||
}
|
||||
if !dnsIntercept {
|
||||
removeInterceptDNSTargetFn(p, "intercept mode inactive")
|
||||
}
|
||||
|
||||
if cfg.Listener == nil {
|
||||
return
|
||||
@@ -846,7 +1070,7 @@ func (p *prog) setDNS() {
|
||||
ns = "127.0.0.1"
|
||||
case lc.Port != 53:
|
||||
ns = "127.0.0.1"
|
||||
if resolver := router.LocalResolverIP(); resolver != "" {
|
||||
if resolver := localResolverIPFn(); resolver != "" {
|
||||
ns = resolver
|
||||
}
|
||||
default:
|
||||
@@ -865,7 +1089,7 @@ func (p *prog) setDNS() {
|
||||
slices.Sort(nameservers)
|
||||
|
||||
netIfaceName := ""
|
||||
netIface := p.setDnsForRunningIface(nameservers)
|
||||
netIface := setDnsForRunningIfaceFn(p, nameservers)
|
||||
if netIface != nil {
|
||||
netIfaceName = netIface.Name
|
||||
}
|
||||
@@ -898,6 +1122,17 @@ func (p *prog) setDNS() {
|
||||
}
|
||||
}
|
||||
|
||||
// configuredInterceptMode resolves the service's effective intercept mode without
|
||||
// mutating package state. An explicit flag value, including "off", takes priority
|
||||
// over the persisted config value.
|
||||
func (p *prog) configuredInterceptMode() string {
|
||||
im := interceptMode
|
||||
if im == "" {
|
||||
im = p.cfg.Service.InterceptMode
|
||||
}
|
||||
return im
|
||||
}
|
||||
|
||||
func (p *prog) setDnsForRunningIface(nameservers []string) (runningIface *net.Interface) {
|
||||
if p.runningIface == "" {
|
||||
return
|
||||
@@ -1055,6 +1290,10 @@ func (p *prog) dnsWatchdog(iface *net.Interface, nameservers []string) {
|
||||
// resetDNS performs a DNS reset for all interfaces.
|
||||
// In DNS intercept mode, this tears down the WFP/pf filters instead.
|
||||
func (p *prog) resetDNS(isStart bool, restoreStatic bool) {
|
||||
// A previous crash can leave a persisted macOS intercept target even when
|
||||
// no live interceptor state exists. Cleanup must run for stop/uninstall and
|
||||
// traditional-mode startup as well as the normal intercept shutdown path.
|
||||
removeInterceptDNSTargetFn(p, "DNS reset")
|
||||
if dnsIntercept && p.dnsInterceptState != nil {
|
||||
if err := p.stopDNSIntercept(); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("Failed to stop DNS intercept mode during reset")
|
||||
@@ -1328,37 +1567,80 @@ func errAddrInUse(err error) bool {
|
||||
|
||||
var _ = errAddrInUse
|
||||
|
||||
// The unreachable winsock errnos (ENETUNREACH/EHOSTUNREACH) are matched via
|
||||
// ctrldnet.IsUnreachable, which owns their definitions.
|
||||
//
|
||||
// https://learn.microsoft.com/en-us/windows/win32/winsock/windows-sockets-error-codes-2
|
||||
var (
|
||||
windowsECONNREFUSED = syscall.Errno(10061)
|
||||
windowsENETUNREACH = syscall.Errno(10051)
|
||||
windowsEINVAL = syscall.Errno(10022)
|
||||
windowsEADDRINUSE = syscall.Errno(10048)
|
||||
windowsEHOSTUNREACH = syscall.Errno(10065)
|
||||
)
|
||||
|
||||
// errUrlNetworkError reports whether a failed HTTP attempt is worth retrying.
|
||||
//
|
||||
// The two-attempt paths compose one *url.Error per attempt - hostname first, then the
|
||||
// direct-IP fallback - so this walks them in order rather than classifying only the first
|
||||
// one errors.As happens to find. Each attempt can say one of three things:
|
||||
//
|
||||
// - retryable (unreachable, refused, temporary): retry, whichever attempt said it;
|
||||
// - a name-resolution failure: no verdict. Only the hostname attempt resolves DNS, and
|
||||
// at boot behind a captive portal or before the router's forwarder is up it fails
|
||||
// this way while the network is merely not ready yet. Consult the next attempt;
|
||||
// - anything else, notably a locally denied socket (WSAEACCES from a firewall blocking
|
||||
// ctrld): definitive. Stop, because retrying cannot clear it - the Firewall Mode
|
||||
// incident spent 256 retry cycles against filters that were never going to clear.
|
||||
func errUrlNetworkError(err error) bool {
|
||||
var urlErr *url.Error
|
||||
if errors.As(err, &urlErr) {
|
||||
return errNetworkError(urlErr.Err)
|
||||
for _, attempt := range attemptErrors(err) {
|
||||
var urlErr *url.Error
|
||||
if !errors.As(attempt, &urlErr) {
|
||||
continue
|
||||
}
|
||||
switch {
|
||||
case errNetworkError(urlErr.Err):
|
||||
return true
|
||||
case errDNSResolutionFailure(urlErr.Err):
|
||||
// Neutral; let a later attempt decide.
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// attemptErrors returns the per-attempt errors recorded in err, in the order they were
|
||||
// tried. A composed fallback error wraps one per attempt; anything else is a single
|
||||
// attempt.
|
||||
func attemptErrors(err error) []error {
|
||||
if multi, ok := err.(interface{ Unwrap() []error }); ok {
|
||||
return multi.Unwrap()
|
||||
}
|
||||
return []error{err}
|
||||
}
|
||||
|
||||
// errDNSResolutionFailure reports whether err is a name-resolution failure. Go marks a
|
||||
// *net.DNSError as temporary only for socket failures that reached the server, so a
|
||||
// SERVFAIL or "no such host" answer is not temporary - but it is also not evidence that
|
||||
// retrying is pointless, which is why callers treat it as no verdict.
|
||||
func errDNSResolutionFailure(err error) bool {
|
||||
var dnsErr *net.DNSError
|
||||
return errors.As(err, &dnsErr)
|
||||
}
|
||||
|
||||
func errNetworkError(err error) bool {
|
||||
var opErr *net.OpError
|
||||
if errors.As(err, &opErr) {
|
||||
if opErr.Temporary() {
|
||||
return true
|
||||
}
|
||||
if ctrldnet.IsUnreachable(err) {
|
||||
return true
|
||||
}
|
||||
switch {
|
||||
case errors.Is(opErr.Err, syscall.ECONNREFUSED),
|
||||
errors.Is(opErr.Err, syscall.EINVAL),
|
||||
errors.Is(opErr.Err, syscall.ENETUNREACH),
|
||||
errors.Is(opErr.Err, windowsENETUNREACH),
|
||||
errors.Is(opErr.Err, windowsEINVAL),
|
||||
errors.Is(opErr.Err, windowsECONNREFUSED),
|
||||
errors.Is(opErr.Err, windowsEHOSTUNREACH):
|
||||
errors.Is(opErr.Err, windowsECONNREFUSED):
|
||||
return true
|
||||
}
|
||||
}
|
||||
@@ -1651,6 +1933,19 @@ func shouldUpgrade(vt string, cv *semver.Version, logger *zerolog.Logger) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// newUpgradeCmd builds the detached command used to self-upgrade. It is a
|
||||
// package-level variable so tests can stub it. With the real implementation a
|
||||
// *test* binary would re-exec itself — os.Executable() is the test binary, and
|
||||
// because `go test` stops flag parsing at the first positional arg ("upgrade")
|
||||
// it ignores the args and re-runs the entire suite. That child hits the same
|
||||
// upgrade test and spawns another child, recursively: a fork bomb of detached
|
||||
// processes that pins the host and locks the test binary's image file.
|
||||
var newUpgradeCmd = func(exe string) *exec.Cmd {
|
||||
cmd := exec.Command(exe, "upgrade", "prod", "-vv")
|
||||
cmd.SysProcAttr = sysProcAttrForDetachedChildProcess()
|
||||
return cmd
|
||||
}
|
||||
|
||||
// performUpgrade executes the self-upgrade command.
|
||||
// Returns true if upgrade was initiated successfully, false otherwise.
|
||||
func performUpgrade(vt string) bool {
|
||||
@@ -1659,8 +1954,7 @@ func performUpgrade(vt string) bool {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to get executable path, skipped self-upgrade")
|
||||
return false
|
||||
}
|
||||
cmd := exec.Command(exe, "upgrade", "prod", "-vv")
|
||||
cmd.SysProcAttr = sysProcAttrForDetachedChildProcess()
|
||||
cmd := newUpgradeCmd(exe)
|
||||
if err := cmd.Start(); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to start self-upgrade")
|
||||
return false
|
||||
|
||||
@@ -0,0 +1,275 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
// TestInterfaceDNSFallbackViable covers when the interface-DNS fallback may be used
|
||||
// after DNS intercept fails to start.
|
||||
//
|
||||
// The fallback names a resolver by IP with no port, so it can only reach a listener on
|
||||
// :53. Taking it with the listener on a redirect-dependent port produced a total DNS
|
||||
// outage on macOS: the interface points at 127.0.0.1, mDNSResponder answers there, and
|
||||
// its upstream is ctrld's own address - a resolution loop with a healthy ctrld listener
|
||||
// nothing can reach. Intercept startup refuses the fallback in that case rather than
|
||||
// creating it.
|
||||
func TestInterfaceDNSFallbackViable(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
lc *ctrld.ListenerConfig
|
||||
localResolver string
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "listener on 53 can be reached by interface DNS",
|
||||
lc: &ctrld.ListenerConfig{IP: "127.0.0.1", Port: 53},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
// The reported outage: no local resolver, so the :5354 fallback port
|
||||
// cannot be expressed by interface DNS.
|
||||
name: "listener on the fallback port cannot",
|
||||
lc: &ctrld.ListenerConfig{IP: "127.0.0.1", Port: 5354},
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "any other non-53 port cannot",
|
||||
lc: &ctrld.ListenerConfig{IP: "127.0.0.1", Port: 5300},
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
// Router platforms with their own dnsmasq: it owns :53 and forwards to
|
||||
// ctrld's port, so interface DNS reaches the listener through it.
|
||||
// Refusing here would break a working EdgeOS/Firewalla setup.
|
||||
name: "non-53 listener behind a forwarding local resolver",
|
||||
lc: &ctrld.ListenerConfig{IP: "127.0.0.1", Port: 5354},
|
||||
localResolver: "192.168.1.1",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
// Port is resolved elsewhere and defaults to 53; nothing to refuse yet.
|
||||
name: "unset port is not refused",
|
||||
lc: &ctrld.ListenerConfig{IP: "127.0.0.1"},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "no listener is not refused",
|
||||
lc: nil,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
// A non-loopback listener on 53 is still reachable by IP.
|
||||
name: "non-loopback listener on 53",
|
||||
lc: &ctrld.ListenerConfig{IP: "192.168.1.10", Port: 53},
|
||||
want: true,
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := interfaceDNSFallbackViable(tc.lc, tc.localResolver); got != tc.want {
|
||||
t.Errorf("interfaceDNSFallbackViable() = %v, want %v", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// interceptFallbackHarness drives setDNS() through the intercept-start failure path and
|
||||
// records the side effects that decide whether the host ends up with a working
|
||||
// resolver.
|
||||
//
|
||||
// Every host-touching step is stubbed, including the intercept start itself: this test
|
||||
// runs untagged on Linux, macOS and Windows runners, where the real startDNSIntercept
|
||||
// would set up pf or install an NRPT rule on the machine running the tests. Stubbing it
|
||||
// also makes the precondition deterministic - the failure under test is injected rather
|
||||
// than depending on the runner denying a privileged operation.
|
||||
type interceptFallbackHarness struct {
|
||||
interceptCalls int
|
||||
ensureTargetCalls int
|
||||
ensuredNameservers []string
|
||||
installedNameservers []string
|
||||
installCalls int
|
||||
resetCalls int
|
||||
removeTargetCalls int
|
||||
refusals []string
|
||||
}
|
||||
|
||||
func newInterceptFallbackHarness(t *testing.T, lc *ctrld.ListenerConfig) *interceptFallbackHarness {
|
||||
t.Helper()
|
||||
h := &interceptFallbackHarness{}
|
||||
|
||||
origStart, origEnsure, origRemove, origInstall := startDNSInterceptFn, ensureInterceptDNSTargetFn, removeInterceptDNSTargetFn, setDnsForRunningIfaceFn
|
||||
origReset, origFatal := resetDNSFn, refuseFallbackFatal
|
||||
origResolver := localResolverIPFn
|
||||
origCfg, origMode, origIntercept, origHard := cfg, interceptMode, dnsIntercept, hardIntercept
|
||||
t.Cleanup(func() {
|
||||
startDNSInterceptFn, ensureInterceptDNSTargetFn, removeInterceptDNSTargetFn, setDnsForRunningIfaceFn = origStart, origEnsure, origRemove, origInstall
|
||||
resetDNSFn, refuseFallbackFatal = origReset, origFatal
|
||||
localResolverIPFn = origResolver
|
||||
cfg, interceptMode, dnsIntercept, hardIntercept = origCfg, origMode, origIntercept, origHard
|
||||
})
|
||||
|
||||
// Default to no local resolver: the desktop case. Router cases set it per test.
|
||||
localResolverIPFn = func() string { return "" }
|
||||
|
||||
// Never reach the real interceptor: it would configure pf on macOS and NRPT on
|
||||
// Windows, on the machine running the tests.
|
||||
startDNSInterceptFn = func(_ *prog) error {
|
||||
h.interceptCalls++
|
||||
return errors.New("dns intercept: injected start failure")
|
||||
}
|
||||
ensureInterceptDNSTargetFn = func(_ *prog, nameservers []string) {
|
||||
h.ensureTargetCalls++
|
||||
h.ensuredNameservers = slices.Clone(nameservers)
|
||||
}
|
||||
removeInterceptDNSTargetFn = func(_ *prog, _ string) { h.removeTargetCalls++ }
|
||||
setDnsForRunningIfaceFn = func(_ *prog, nameservers []string) *net.Interface {
|
||||
h.installCalls++
|
||||
h.installedNameservers = nameservers
|
||||
return nil
|
||||
}
|
||||
resetDNSFn = func(_ *prog, _ bool, _ bool) { h.resetCalls++ }
|
||||
refuseFallbackFatal = func(format string, v ...any) {
|
||||
h.refusals = append(h.refusals, fmt.Sprintf(format, v...))
|
||||
}
|
||||
|
||||
cfg = ctrld.Config{}
|
||||
cfg.Service.InterceptMode = "dns"
|
||||
cfg.Listener = map[string]*ctrld.ListenerConfig{"0": lc}
|
||||
watchdogOff := false
|
||||
cfg.Service.DnsWatchdogEnabled = &watchdogOff
|
||||
interceptMode, dnsIntercept, hardIntercept = "dns", false, false
|
||||
return h
|
||||
}
|
||||
|
||||
func (h *interceptFallbackHarness) run(t *testing.T) {
|
||||
t.Helper()
|
||||
p := &prog{cfg: &cfg}
|
||||
p.setDNS(nil)
|
||||
}
|
||||
|
||||
func TestSetDNSEnsuresInterceptTargetAfterSuccessfulStart(t *testing.T) {
|
||||
h := newInterceptFallbackHarness(t, &ctrld.ListenerConfig{IP: "127.0.0.1", Port: 5354})
|
||||
startDNSInterceptFn = func(_ *prog) error {
|
||||
h.interceptCalls++
|
||||
return nil
|
||||
}
|
||||
want := []string{"fe80::1"}
|
||||
p := &prog{cfg: &cfg}
|
||||
p.setDNS(want)
|
||||
|
||||
if h.interceptCalls != 1 {
|
||||
t.Fatalf("intercept start called %d time(s), want 1", h.interceptCalls)
|
||||
}
|
||||
if h.ensureTargetCalls != 1 {
|
||||
t.Fatalf("intercept DNS target ensured %d time(s), want 1 after successful start", h.ensureTargetCalls)
|
||||
}
|
||||
if !slices.Equal(h.ensuredNameservers, want) {
|
||||
t.Fatalf("system nameservers = %v, want %v", h.ensuredNameservers, want)
|
||||
}
|
||||
if h.installCalls != 0 {
|
||||
t.Fatalf("interface-DNS fallback installed %d time(s) after successful intercept start", h.installCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetDNSExplicitOffOverridesConfig(t *testing.T) {
|
||||
h := newInterceptFallbackHarness(t, &ctrld.ListenerConfig{IP: "127.0.0.1", Port: 53})
|
||||
interceptMode = "off"
|
||||
dnsIntercept = false
|
||||
hardIntercept = false
|
||||
|
||||
h.run(t)
|
||||
|
||||
if h.interceptCalls != 0 {
|
||||
t.Fatalf("intercept start called %d time(s), want 0: explicit off must override service.intercept_mode", h.interceptCalls)
|
||||
}
|
||||
if h.installCalls != 1 {
|
||||
t.Fatalf("interface DNS installed %d time(s), want 1", h.installCalls)
|
||||
}
|
||||
if h.removeTargetCalls != 1 {
|
||||
t.Fatalf("stale intercept DNS target cleanup called %d time(s), want 1", h.removeTargetCalls)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSetDNSRefusesUnreachableFallback is the behaviour test for the reported outage: it
|
||||
// drives the real setDNS() lifecycle rather than the classification helper alone.
|
||||
//
|
||||
// Deleting or bypassing the guard in setDNS makes the first case fail, because interface
|
||||
// DNS then gets installed pointing at a listener that cannot answer on :53 - which is
|
||||
// the resolution loop this refuses to create.
|
||||
func TestSetDNSRefusesUnreachableFallback(t *testing.T) {
|
||||
t.Run("non-53 listener refuses the fallback and restores DNS", func(t *testing.T) {
|
||||
h := newInterceptFallbackHarness(t, &ctrld.ListenerConfig{IP: "127.0.0.1", Port: 5354})
|
||||
h.run(t)
|
||||
|
||||
if h.interceptCalls != 1 {
|
||||
t.Fatalf("intercept start called %d time(s) through the seam, want 1 — the real platform interceptor must never run here", h.interceptCalls)
|
||||
}
|
||||
if h.installCalls != 0 {
|
||||
t.Errorf("interface DNS was installed %d time(s) for a listener on :5354 — that is the resolver loop", h.installCalls)
|
||||
}
|
||||
if h.resetCalls == 0 {
|
||||
t.Error("host DNS was not restored before refusing, leaving the interface pointed at a ctrld that is not serving")
|
||||
}
|
||||
if h.removeTargetCalls != 1 {
|
||||
t.Errorf("stale intercept DNS target cleanup called %d time(s), want 1 after intercept failure", h.removeTargetCalls)
|
||||
}
|
||||
if len(h.refusals) == 0 {
|
||||
t.Fatal("refusal was not surfaced: startup must fail loudly rather than silently skip the fallback")
|
||||
}
|
||||
if !strings.Contains(h.refusals[0], "5354") {
|
||||
t.Errorf("refusal does not name the unreachable port: %q", h.refusals[0])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-53 listener behind a local resolver still falls back", func(t *testing.T) {
|
||||
// EdgeOS/Firewalla: dnsmasq owns :53 and forwards to ctrld's port, so the
|
||||
// fallback works and must not be refused. setDNS points the interface at the
|
||||
// resolver rather than at the listener.
|
||||
h := newInterceptFallbackHarness(t, &ctrld.ListenerConfig{IP: "127.0.0.1", Port: 5354})
|
||||
localResolverIPFn = func() string { return "192.168.1.1" }
|
||||
h.run(t)
|
||||
|
||||
if h.installCalls != 1 {
|
||||
t.Errorf("interface DNS installed %d time(s), want 1: a forwarding local resolver makes the fallback usable", h.installCalls)
|
||||
}
|
||||
if len(h.refusals) != 0 {
|
||||
t.Errorf("refused a fallback that a local resolver can serve: %v", h.refusals)
|
||||
}
|
||||
// Assert on membership, not on the exact set: setDNS appends platform-dependent
|
||||
// entries beside the chosen nameserver - "::1" on Windows for the local IPv6
|
||||
// listener, the RFC1918 addresses where those listeners are needed. What matters
|
||||
// is that the interface points at the resolver and not at the listener IP, whose
|
||||
// port the interface cannot express.
|
||||
if !slices.Contains(h.installedNameservers, "192.168.1.1") {
|
||||
t.Errorf("nameservers = %v, want the local resolver among them so queries reach ctrld through it", h.installedNameservers)
|
||||
}
|
||||
if slices.Contains(h.installedNameservers, "127.0.0.1") {
|
||||
t.Errorf("nameservers = %v, must not name the listener IP: interface DNS cannot reach it on :5354", h.installedNameservers)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("listener on 53 still reaches the interface-DNS fallback", func(t *testing.T) {
|
||||
h := newInterceptFallbackHarness(t, &ctrld.ListenerConfig{IP: "127.0.0.1", Port: 53})
|
||||
h.run(t)
|
||||
|
||||
if h.interceptCalls != 1 {
|
||||
t.Fatalf("intercept start called %d time(s) through the seam, want 1", h.interceptCalls)
|
||||
}
|
||||
if h.installCalls != 1 {
|
||||
t.Errorf("interface DNS installed %d time(s), want 1: a listener on :53 is reachable, so the fallback must still apply", h.installCalls)
|
||||
}
|
||||
if len(h.refusals) != 0 {
|
||||
t.Errorf("unexpected refusal for a reachable listener: %v", h.refusals)
|
||||
}
|
||||
if len(h.installedNameservers) == 0 {
|
||||
t.Error("fallback installed no nameservers")
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -14,6 +14,9 @@ import (
|
||||
)
|
||||
|
||||
func init() {
|
||||
if isAndroid() {
|
||||
return
|
||||
}
|
||||
if r, err := newLoopbackOSConfigurator(); err == nil {
|
||||
useSystemdResolved = r.Mode() == "systemd-resolved"
|
||||
}
|
||||
|
||||
@@ -1,7 +1,11 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/url"
|
||||
"runtime"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -12,6 +16,32 @@ import (
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
func TestErrNetworkErrorTreatsNoRouteAsNetworkError(t *testing.T) {
|
||||
err := &net.OpError{Op: "dial", Net: "tcp", Err: syscall.EHOSTUNREACH}
|
||||
assert.True(t, errNetworkError(err))
|
||||
assert.True(t, errUrlNetworkError(&url.Error{Op: "Get", URL: "https://dns.controld.com", Err: err}))
|
||||
}
|
||||
|
||||
func TestSleepWithContext(t *testing.T) {
|
||||
assert.True(t, sleepWithContext(context.Background(), time.Millisecond))
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
start := time.Now()
|
||||
assert.False(t, sleepWithContext(ctx, time.Minute))
|
||||
assert.Less(t, time.Since(start), 100*time.Millisecond)
|
||||
}
|
||||
|
||||
func TestUnreachableRecoveryBackoff(t *testing.T) {
|
||||
// Streak starts at the base cadence and doubles each attempt, capped at the max.
|
||||
assert.Equal(t, checkUpstreamBackoffSleep, unreachableRecoveryBackoff(0))
|
||||
assert.Equal(t, checkUpstreamBackoffSleep, unreachableRecoveryBackoff(1))
|
||||
assert.Equal(t, 2*checkUpstreamBackoffSleep, unreachableRecoveryBackoff(2))
|
||||
assert.Equal(t, 4*checkUpstreamBackoffSleep, unreachableRecoveryBackoff(3))
|
||||
assert.Equal(t, checkUpstreamUnreachableBackoffMax, unreachableRecoveryBackoff(100))
|
||||
}
|
||||
|
||||
func Test_prog_dnsWatchdogEnabled(t *testing.T) {
|
||||
p := &prog{cfg: &ctrld.Config{}}
|
||||
|
||||
@@ -262,6 +292,8 @@ func Test_performUpgrade(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
// newUpgradeCmd is stubbed in TestMain so performUpgrade does not re-exec
|
||||
// (and fork-bomb) the test binary; see the comment there.
|
||||
for _, tc := range tests {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,247 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
// A terminal provisioning failure reports the same stable code on three
|
||||
// surfaces: a persisted result file, one fixed-format output line, and a
|
||||
// stage-scoped process exit code. docs/provisioning-failure-codes.md maps
|
||||
// each code to its scenario and must stay in sync with the constants below.
|
||||
// Codes are append-only once released; renaming or reusing one breaks the
|
||||
// support contract.
|
||||
|
||||
type provisionStage string
|
||||
|
||||
const (
|
||||
provisionStageBootstrap provisionStage = "bootstrap"
|
||||
provisionStageListener provisionStage = "listener"
|
||||
provisionStageService provisionStage = "service"
|
||||
)
|
||||
|
||||
type provisionFailureCode string
|
||||
|
||||
const (
|
||||
provisionCodeAPIUnreachable provisionFailureCode = "API_UNREACHABLE"
|
||||
provisionCodeAPIRejected provisionFailureCode = "API_REJECTED"
|
||||
provisionCodeAPIDeviceInvalid provisionFailureCode = "API_DEVICE_INVALID"
|
||||
provisionCodeListenerBindFailed provisionFailureCode = "LISTENER_BIND_FAILED"
|
||||
provisionCodeListenerAddrUnavail provisionFailureCode = "LISTENER_CONFIGURED_ADDR_UNAVAILABLE"
|
||||
provisionCodeServiceInstall provisionFailureCode = "SERVICE_INSTALL_FAILED"
|
||||
provisionCodeServiceStartFailed provisionFailureCode = "SERVICE_START_FAILED"
|
||||
provisionCodeServiceSelfCheck provisionFailureCode = "SERVICE_SELFCHECK_FAILED"
|
||||
)
|
||||
|
||||
var allProvisionFailureCodes = []provisionFailureCode{
|
||||
provisionCodeAPIUnreachable,
|
||||
provisionCodeAPIRejected,
|
||||
provisionCodeAPIDeviceInvalid,
|
||||
provisionCodeListenerBindFailed,
|
||||
provisionCodeListenerAddrUnavail,
|
||||
provisionCodeServiceInstall,
|
||||
provisionCodeServiceStartFailed,
|
||||
provisionCodeServiceSelfCheck,
|
||||
}
|
||||
|
||||
var provisionStageForCode = map[provisionFailureCode]provisionStage{
|
||||
provisionCodeAPIUnreachable: provisionStageBootstrap,
|
||||
provisionCodeAPIRejected: provisionStageBootstrap,
|
||||
provisionCodeAPIDeviceInvalid: provisionStageBootstrap,
|
||||
provisionCodeListenerBindFailed: provisionStageListener,
|
||||
provisionCodeListenerAddrUnavail: provisionStageListener,
|
||||
provisionCodeServiceInstall: provisionStageService,
|
||||
provisionCodeServiceStartFailed: provisionStageService,
|
||||
provisionCodeServiceSelfCheck: provisionStageService,
|
||||
}
|
||||
|
||||
// Exit codes are grouped by stage (bootstrap 30-39, listener 40-49, service
|
||||
// 50-59) so the exit code alone names the failed stage. 0-3 belong to
|
||||
// "ctrld status" and 126 to the deactivation pin check; never reuse those.
|
||||
var provisionExitCodeForCode = map[provisionFailureCode]int{
|
||||
provisionCodeAPIUnreachable: 30,
|
||||
provisionCodeAPIRejected: 31,
|
||||
provisionCodeAPIDeviceInvalid: 32,
|
||||
provisionCodeListenerBindFailed: 41,
|
||||
provisionCodeListenerAddrUnavail: 42,
|
||||
provisionCodeServiceInstall: 51,
|
||||
provisionCodeServiceStartFailed: 52,
|
||||
provisionCodeServiceSelfCheck: 53,
|
||||
}
|
||||
|
||||
const (
|
||||
provisionResultFileName = "provision_result.json"
|
||||
// Detail identifies a failure, it is not a log. Caps keep the artifact
|
||||
// small and predictable.
|
||||
maxProvisionBindAttempts = 12
|
||||
maxProvisionStringLen = 256
|
||||
)
|
||||
|
||||
type provisionBindAttempt struct {
|
||||
Addr string `json:"addr"`
|
||||
Proto string `json:"proto"`
|
||||
OSError string `json:"os_error"`
|
||||
}
|
||||
|
||||
type provisionDetail struct {
|
||||
Attempts []provisionBindAttempt `json:"attempts,omitempty"`
|
||||
}
|
||||
|
||||
type provisionResult struct {
|
||||
Version int `json:"version"`
|
||||
Timestamp string `json:"timestamp"`
|
||||
Stage string `json:"stage"`
|
||||
Code string `json:"code"`
|
||||
ExitCode int `json:"exit_code"`
|
||||
Message string `json:"message"`
|
||||
Detail *provisionDetail `json:"detail,omitempty"`
|
||||
}
|
||||
|
||||
// provisionResultPath is a var so tests can point it at a temp dir.
|
||||
var provisionResultPath = func() string {
|
||||
return absHomeDir(provisionResultFileName)
|
||||
}
|
||||
|
||||
// provisionExit is a var so tests can observe the exit code instead of dying.
|
||||
var provisionExit = os.Exit
|
||||
|
||||
// newProvisionResult builds a result with every field bounded and the given
|
||||
// secrets stripped. The artifact reaches installer logs and support tickets,
|
||||
// so callers pass every secret in scope (provision token, cd UID).
|
||||
func newProvisionResult(code provisionFailureCode, message string, attempts []provisionBindAttempt, secrets ...string) *provisionResult {
|
||||
sanitize := func(s string) string {
|
||||
s = redactSecrets(s, secrets...)
|
||||
if len(s) > maxProvisionStringLen {
|
||||
// Cut on a rune boundary so a localized OS error does not end in
|
||||
// a broken multi-byte sequence.
|
||||
cut := maxProvisionStringLen
|
||||
for cut > 0 && !utf8.RuneStart(s[cut]) {
|
||||
cut--
|
||||
}
|
||||
s = s[:cut]
|
||||
}
|
||||
return s
|
||||
}
|
||||
r := &provisionResult{
|
||||
Version: 1,
|
||||
Timestamp: time.Now().UTC().Format(time.RFC3339),
|
||||
Stage: string(provisionStageForCode[code]),
|
||||
Code: string(code),
|
||||
ExitCode: provisionExitCodeForCode[code],
|
||||
Message: sanitize(message),
|
||||
}
|
||||
if len(attempts) > 0 {
|
||||
if len(attempts) > maxProvisionBindAttempts {
|
||||
attempts = attempts[:maxProvisionBindAttempts]
|
||||
}
|
||||
detail := &provisionDetail{Attempts: make([]provisionBindAttempt, 0, len(attempts))}
|
||||
for _, a := range attempts {
|
||||
detail.Attempts = append(detail.Attempts, provisionBindAttempt{
|
||||
Addr: sanitize(a.Addr),
|
||||
Proto: sanitize(a.Proto),
|
||||
OSError: sanitize(a.OSError),
|
||||
})
|
||||
}
|
||||
r.Detail = detail
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// redactSecrets removes every non-empty secret from s.
|
||||
func redactSecrets(s string, secrets ...string) string {
|
||||
for _, secret := range secrets {
|
||||
if secret == "" {
|
||||
continue
|
||||
}
|
||||
s = strings.ReplaceAll(s, secret, "[redacted]")
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// provisionResultTrusted rejects a result whose code, stage, or exit code is
|
||||
// not part of the known contract, so a corrupt or planted file cannot drive
|
||||
// what "ctrld start" logs and exits with.
|
||||
func provisionResultTrusted(r *provisionResult) bool {
|
||||
code := provisionFailureCode(r.Code)
|
||||
stage, ok := provisionStageForCode[code]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return r.Stage == string(stage) && r.ExitCode == provisionExitCodeForCode[code]
|
||||
}
|
||||
|
||||
func (r *provisionResult) failureLine() string {
|
||||
return fmt.Sprintf("provisioning failed: stage=%s code=%s (exit %d)", r.Stage, r.Code, r.ExitCode)
|
||||
}
|
||||
|
||||
// writeProvisionResult persists the result atomically (temp file + rename in
|
||||
// the same directory) so a reader never sees a partial file.
|
||||
func writeProvisionResult(r *provisionResult) error {
|
||||
path := provisionResultPath()
|
||||
buf, err := json.MarshalIndent(r, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tmp, err := os.CreateTemp(filepath.Dir(path), provisionResultFileName+".tmp*")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tmpName := tmp.Name()
|
||||
if _, err := tmp.Write(buf); err != nil {
|
||||
_ = tmp.Close()
|
||||
_ = os.Remove(tmpName)
|
||||
return err
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
_ = os.Remove(tmpName)
|
||||
return err
|
||||
}
|
||||
if err := os.Chmod(tmpName, 0o600); err != nil {
|
||||
_ = os.Remove(tmpName)
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(tmpName, path); err != nil {
|
||||
_ = os.Remove(tmpName)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func readProvisionResult() (*provisionResult, error) {
|
||||
buf, err := os.ReadFile(provisionResultPath())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r := &provisionResult{}
|
||||
if err := json.Unmarshal(buf, r); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return r, nil
|
||||
}
|
||||
|
||||
// clearProvisionResult removes a stale result once provisioning succeeds, so
|
||||
// support never diagnoses a healthy install from an old failure.
|
||||
func clearProvisionResult() {
|
||||
if err := os.Remove(provisionResultPath()); err != nil && !os.IsNotExist(err) {
|
||||
mainLog.Load().Debug().Err(err).Msg("could not remove provision result file")
|
||||
}
|
||||
}
|
||||
|
||||
// failProvision persists the result, prints the identifier line, unblocks a
|
||||
// waiting "ctrld start" via notify, then exits with the stage code. The write
|
||||
// comes first so the file survives even if logging or notify misbehaves.
|
||||
func failProvision(r *provisionResult, notify func()) {
|
||||
if err := writeProvisionResult(r); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("could not persist provision result")
|
||||
}
|
||||
mainLog.Load().Error().Msg(r.failureLine())
|
||||
if notify != nil {
|
||||
notify()
|
||||
}
|
||||
provisionExit(r.ExitCode)
|
||||
}
|
||||
@@ -0,0 +1,278 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
func overrideProvisionResultPath(t *testing.T) string {
|
||||
t.Helper()
|
||||
path := filepath.Join(t.TempDir(), provisionResultFileName)
|
||||
old := provisionResultPath
|
||||
provisionResultPath = func() string { return path }
|
||||
t.Cleanup(func() { provisionResultPath = old })
|
||||
return path
|
||||
}
|
||||
|
||||
func TestProvisionCodesMapToOneStageAndInRangeExit(t *testing.T) {
|
||||
stageRanges := map[provisionStage][2]int{
|
||||
provisionStageBootstrap: {30, 39},
|
||||
provisionStageListener: {40, 49},
|
||||
provisionStageService: {50, 59},
|
||||
}
|
||||
reservedExits := map[int]string{
|
||||
statusExitRunning: "ctrld status running",
|
||||
statusExitStopped: "ctrld status stopped",
|
||||
statusExitUnknown: "ctrld status unknown",
|
||||
statusExitNotReady: "ctrld status not ready",
|
||||
deactivationPinInvalidExitCode: "deactivation pin invalid",
|
||||
}
|
||||
seenExits := make(map[int]provisionFailureCode)
|
||||
for _, code := range allProvisionFailureCodes {
|
||||
stage, ok := provisionStageForCode[code]
|
||||
if !ok {
|
||||
t.Fatalf("code %s has no stage", code)
|
||||
}
|
||||
exit, ok := provisionExitCodeForCode[code]
|
||||
if !ok {
|
||||
t.Fatalf("code %s has no exit code", code)
|
||||
}
|
||||
r := stageRanges[stage]
|
||||
if exit < r[0] || exit > r[1] {
|
||||
t.Errorf("code %s exit %d outside stage %s range %v", code, exit, stage, r)
|
||||
}
|
||||
if owner, ok := reservedExits[exit]; ok {
|
||||
t.Errorf("code %s exit %d collides with %s", code, exit, owner)
|
||||
}
|
||||
if prev, dup := seenExits[exit]; dup {
|
||||
t.Errorf("codes %s and %s share exit %d", prev, code, exit)
|
||||
}
|
||||
seenExits[exit] = code
|
||||
}
|
||||
if len(allProvisionFailureCodes) != 8 {
|
||||
t.Errorf("expected 8 codes, got %d", len(allProvisionFailureCodes))
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewProvisionResultRedactsSecrets(t *testing.T) {
|
||||
token := "org-secret-token-12345"
|
||||
cdUIDValue := "abcdef123456"
|
||||
attempts := []provisionBindAttempt{
|
||||
{Addr: "127.0.0.1:53", Proto: "udp", OSError: "bind failed for " + token},
|
||||
}
|
||||
r := newProvisionResult(
|
||||
provisionCodeListenerBindFailed,
|
||||
"could not bind, token="+token+" uid="+cdUIDValue,
|
||||
attempts,
|
||||
token, cdUIDValue,
|
||||
)
|
||||
raw, err := json.Marshal(r)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, secret := range []string{token, cdUIDValue} {
|
||||
if strings.Contains(string(raw), secret) {
|
||||
t.Errorf("serialized result contains secret %q: %s", secret, raw)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewProvisionResultBoundsDetail(t *testing.T) {
|
||||
long := strings.Repeat("x", 1000)
|
||||
var attempts []provisionBindAttempt
|
||||
for i := 0; i < 50; i++ {
|
||||
attempts = append(attempts, provisionBindAttempt{Addr: long, Proto: "udp", OSError: long})
|
||||
}
|
||||
r := newProvisionResult(provisionCodeListenerBindFailed, long, attempts)
|
||||
if got := len(r.Detail.Attempts); got > maxProvisionBindAttempts {
|
||||
t.Errorf("attempts not capped: %d > %d", got, maxProvisionBindAttempts)
|
||||
}
|
||||
if len(r.Message) > maxProvisionStringLen {
|
||||
t.Errorf("message not capped: %d", len(r.Message))
|
||||
}
|
||||
for _, a := range r.Detail.Attempts {
|
||||
if len(a.Addr) > maxProvisionStringLen || len(a.OSError) > maxProvisionStringLen {
|
||||
t.Error("attempt fields not capped")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestProvisionResultFields(t *testing.T) {
|
||||
r := newProvisionResult(provisionCodeAPIRejected, "the API rejected this configuration", nil)
|
||||
if r.Version != 1 {
|
||||
t.Errorf("version = %d, want 1", r.Version)
|
||||
}
|
||||
if r.Stage != string(provisionStageBootstrap) {
|
||||
t.Errorf("stage = %q, want bootstrap", r.Stage)
|
||||
}
|
||||
if r.ExitCode != provisionExitCodeForCode[provisionCodeAPIRejected] {
|
||||
t.Errorf("exit = %d", r.ExitCode)
|
||||
}
|
||||
if _, err := time.Parse(time.RFC3339, r.Timestamp); err != nil {
|
||||
t.Errorf("timestamp %q not RFC3339: %v", r.Timestamp, err)
|
||||
}
|
||||
if r.Detail != nil {
|
||||
t.Error("nil attempts should give nil detail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProvisionResultTrusted(t *testing.T) {
|
||||
good := newProvisionResult(provisionCodeListenerBindFailed, "x", nil)
|
||||
if !provisionResultTrusted(good) {
|
||||
t.Error("constructor-built result must be trusted")
|
||||
}
|
||||
bogusCode := newProvisionResult(provisionCodeListenerBindFailed, "x", nil)
|
||||
bogusCode.Code = "TOTALLY_MADE_UP"
|
||||
if provisionResultTrusted(bogusCode) {
|
||||
t.Error("unknown code must not be trusted")
|
||||
}
|
||||
wrongExit := newProvisionResult(provisionCodeListenerBindFailed, "x", nil)
|
||||
wrongExit.ExitCode = 126
|
||||
if provisionResultTrusted(wrongExit) {
|
||||
t.Error("exit code not matching the contract must not be trusted")
|
||||
}
|
||||
wrongStage := newProvisionResult(provisionCodeListenerBindFailed, "x", nil)
|
||||
wrongStage.Stage = string(provisionStageService)
|
||||
if provisionResultTrusted(wrongStage) {
|
||||
t.Error("stage not matching the code must not be trusted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewProvisionResultTruncatesOnRuneBoundary(t *testing.T) {
|
||||
msg := strings.Repeat("é", maxProvisionStringLen) // 2 bytes per rune
|
||||
r := newProvisionResult(provisionCodeListenerBindFailed, msg, nil)
|
||||
if len(r.Message) > maxProvisionStringLen {
|
||||
t.Errorf("message not capped: %d bytes", len(r.Message))
|
||||
}
|
||||
if !utf8.ValidString(r.Message) {
|
||||
t.Error("truncation split a multi-byte rune")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFailureCodeDocTableMatchesConstants(t *testing.T) {
|
||||
buf, err := os.ReadFile(filepath.Join("..", "..", "docs", "provisioning-failure-codes.md"))
|
||||
if os.IsNotExist(err) {
|
||||
// The Windows CI runner executes prebuilt test binaries outside the
|
||||
// repo; the sync guarantee is still enforced on runners with a checkout.
|
||||
t.Skip("failure-code doc not available in this test environment")
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("could not read the failure-code doc: %v", err)
|
||||
}
|
||||
doc := string(buf)
|
||||
rows := 0
|
||||
for _, line := range strings.Split(doc, "\n") {
|
||||
if strings.HasPrefix(line, "| `") {
|
||||
rows++
|
||||
}
|
||||
}
|
||||
if rows != len(allProvisionFailureCodes) {
|
||||
t.Errorf("doc table has %d code rows, want %d", rows, len(allProvisionFailureCodes))
|
||||
}
|
||||
for _, code := range allProvisionFailureCodes {
|
||||
row := "| `" + string(code) + "` | " + string(provisionStageForCode[code]) + " | " + strconv.Itoa(provisionExitCodeForCode[code]) + " |"
|
||||
if !strings.Contains(doc, row) {
|
||||
t.Errorf("doc table missing row for %s (want prefix %q)", code, row)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestProvisionFailureLineFormat(t *testing.T) {
|
||||
r := newProvisionResult(provisionCodeListenerBindFailed, "could not find available listen ip and port", nil)
|
||||
want := "provisioning failed: stage=listener code=LISTENER_BIND_FAILED (exit 41)"
|
||||
if got := r.failureLine(); got != want {
|
||||
t.Errorf("failureLine() = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProvisionResultRoundTrip(t *testing.T) {
|
||||
overrideProvisionResultPath(t)
|
||||
in := newProvisionResult(provisionCodeServiceStartFailed, "service failed to start", nil)
|
||||
if err := writeProvisionResult(in); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
out, err := readProvisionResult()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if out.Code != in.Code || out.Stage != in.Stage || out.ExitCode != in.ExitCode || out.Message != in.Message {
|
||||
t.Errorf("round trip mismatch: in=%+v out=%+v", in, out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteProvisionResultOverwritesAtomically(t *testing.T) {
|
||||
path := overrideProvisionResultPath(t)
|
||||
first := newProvisionResult(provisionCodeAPIUnreachable, "first", nil)
|
||||
if err := writeProvisionResult(first); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
second := newProvisionResult(provisionCodeListenerBindFailed, "second", nil)
|
||||
if err := writeProvisionResult(second); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
out, err := readProvisionResult()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if out.Code != string(provisionCodeListenerBindFailed) || out.Message != "second" {
|
||||
t.Errorf("overwrite failed: %+v", out)
|
||||
}
|
||||
entries, err := os.ReadDir(filepath.Dir(path))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(entries) != 1 {
|
||||
t.Errorf("temp files left behind: %v", entries)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClearProvisionResult(t *testing.T) {
|
||||
path := overrideProvisionResultPath(t)
|
||||
clearProvisionResult() // missing file must not panic or error loudly
|
||||
if err := writeProvisionResult(newProvisionResult(provisionCodeAPIUnreachable, "x", nil)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
clearProvisionResult()
|
||||
if _, err := os.Stat(path); !os.IsNotExist(err) {
|
||||
t.Errorf("result file still present after clear: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadProvisionResultMissing(t *testing.T) {
|
||||
overrideProvisionResultPath(t)
|
||||
if _, err := readProvisionResult(); err == nil {
|
||||
t.Error("expected error reading missing result file")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFailProvisionWritesLogsNotifiesAndExits(t *testing.T) {
|
||||
overrideProvisionResultPath(t)
|
||||
exitCode := -1
|
||||
oldExit := provisionExit
|
||||
provisionExit = func(code int) { exitCode = code }
|
||||
t.Cleanup(func() { provisionExit = oldExit })
|
||||
|
||||
notified := false
|
||||
r := newProvisionResult(provisionCodeListenerBindFailed, "no listen addr", nil)
|
||||
failProvision(r, func() { notified = true })
|
||||
|
||||
if !notified {
|
||||
t.Error("notify func not called")
|
||||
}
|
||||
if exitCode != provisionExitCodeForCode[provisionCodeListenerBindFailed] {
|
||||
t.Errorf("exit code = %d", exitCode)
|
||||
}
|
||||
out, err := readProvisionResult()
|
||||
if err != nil {
|
||||
t.Fatalf("result not persisted: %v", err)
|
||||
}
|
||||
if out.Code != string(provisionCodeListenerBindFailed) {
|
||||
t.Errorf("persisted code = %q", out.Code)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
package cli
|
||||
|
||||
import "context"
|
||||
|
||||
// beginRecovery atomically transfers ownership of shared recovery state. A
|
||||
// network change cancels and replaces the current owner without exposing a nil
|
||||
// recoveryCancel gap; other triggers are coalesced while an owner exists.
|
||||
func (p *prog) beginRecovery(reason RecoveryReason) (ctx context.Context, gen uint64, intercept bool, ok bool) {
|
||||
p.recoveryCancelMu.Lock()
|
||||
defer p.recoveryCancelMu.Unlock()
|
||||
|
||||
if reason != RecoveryReasonNetworkChange && p.recoveryCancel != nil {
|
||||
return nil, 0, false, false
|
||||
}
|
||||
if p.recoveryCancel != nil {
|
||||
p.recoveryCancel()
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
gen = p.recoveryGen.Add(1)
|
||||
intercept = dnsIntercept && p.dnsInterceptState != nil
|
||||
p.recoveryCancel = cancel
|
||||
p.recoveryRunning.Store(true)
|
||||
p.recoveryBypass.Store(intercept)
|
||||
return ctx, gen, intercept, true
|
||||
}
|
||||
|
||||
func (p *prog) recoveryOwnsState(gen uint64) bool {
|
||||
p.recoveryCancelMu.Lock()
|
||||
defer p.recoveryCancelMu.Unlock()
|
||||
return p.recoveryGen.Load() == gen && p.recoveryCancel != nil
|
||||
}
|
||||
|
||||
func systemNameserversForInterceptRetry() []string {
|
||||
_, system := initializeOsResolverWithSystemNameserversFn(true)
|
||||
if system == nil {
|
||||
return []string{}
|
||||
}
|
||||
return system
|
||||
}
|
||||
|
||||
// completeRecovery releases shared state only if gen still owns it. The bypass
|
||||
// reset is unconditional because live intercept state can disappear while a
|
||||
// recovery is running, but a stale true flag still affects proxy routing.
|
||||
func (p *prog) completeRecovery(gen uint64) bool {
|
||||
p.recoveryCancelMu.Lock()
|
||||
defer p.recoveryCancelMu.Unlock()
|
||||
if p.recoveryGen.Load() != gen || p.recoveryCancel == nil {
|
||||
return false
|
||||
}
|
||||
p.recoveryBypass.Store(false)
|
||||
p.recoveryRunning.Store(false)
|
||||
p.recoveryCancel = nil
|
||||
return true
|
||||
}
|
||||
|
||||
// recoveryCanceledCleanup resets shared recovery state after a canceled or
|
||||
// failed recovery, but only when the recovery identified by gen was NOT
|
||||
// superseded by a newer one (issue #597).
|
||||
//
|
||||
// A network-change cancellation is normally followed immediately by a new
|
||||
// handleRecovery that owns recoveryBypass/recoveryRunning/recoveryCancel;
|
||||
// clearing them here would disable the successor's bypass mid-flight and
|
||||
// make it uncancellable. But when the canceled recovery is the LAST one
|
||||
// (e.g. the tail of a network flap burst), nothing else will ever clear the
|
||||
// flags: the daemon would stay in recovery bypass forever — every query
|
||||
// detouring to the OS resolver — and the DNS-settings watchdog would stay
|
||||
// permanently disabled.
|
||||
func (p *prog) recoveryCanceledCleanup(gen uint64) {
|
||||
if !p.completeRecovery(gen) {
|
||||
// Superseded: the newer recovery owns the shared state.
|
||||
return
|
||||
}
|
||||
mainLog.Load().Info().Msg("Recovery canceled with no successor; cleared recovery state and DHCP bypass")
|
||||
}
|
||||
@@ -0,0 +1,161 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// interceptStateStub stands in for the platform pfState/wfpState; the
|
||||
// recovery cleanup path only checks dnsInterceptState != nil.
|
||||
type interceptStateStub struct{}
|
||||
|
||||
// setupInterceptRecovery puts p into "intercept-mode recovery in flight"
|
||||
// state and restores the package-level dnsIntercept flag on cleanup.
|
||||
func setupInterceptRecovery(t *testing.T, p *prog) {
|
||||
t.Helper()
|
||||
oldIntercept := dnsIntercept
|
||||
dnsIntercept = true
|
||||
t.Cleanup(func() { dnsIntercept = oldIntercept })
|
||||
p.dnsInterceptState = &interceptStateStub{}
|
||||
p.recoveryBypass.Store(true)
|
||||
p.recoveryRunning.Store(true)
|
||||
p.recoveryCancel = func() {}
|
||||
}
|
||||
|
||||
// TestRecoveryCanceledCleanup_LastRecoveryResetsState pins issue #597: a
|
||||
// canceled recovery with no successor must clear recoveryBypass and
|
||||
// recoveryRunning, or the daemon stays in bypass forever (every query
|
||||
// detours to the OS resolver) and the DNS watchdog stays disabled.
|
||||
func TestRecoveryCanceledCleanup_LastRecoveryResetsState(t *testing.T) {
|
||||
p := &prog{}
|
||||
setupInterceptRecovery(t, p)
|
||||
gen := p.recoveryGen.Add(1)
|
||||
|
||||
p.recoveryCanceledCleanup(gen)
|
||||
|
||||
if p.recoveryBypass.Load() {
|
||||
t.Error("recoveryBypass still set after canceled recovery with no successor")
|
||||
}
|
||||
if p.recoveryRunning.Load() {
|
||||
t.Error("recoveryRunning still set after canceled recovery with no successor")
|
||||
}
|
||||
p.recoveryCancelMu.Lock()
|
||||
cancelCleared := p.recoveryCancel == nil
|
||||
p.recoveryCancelMu.Unlock()
|
||||
if !cancelCleared {
|
||||
t.Error("recoveryCancel not cleared after canceled recovery with no successor")
|
||||
}
|
||||
}
|
||||
|
||||
// TestRecoveryCanceledCleanup_SupersededKeepsSuccessorState pins the
|
||||
// captive-portal/network-flap contract: when a newer recovery superseded the
|
||||
// canceled one, the canceled recovery must NOT clear shared state — the
|
||||
// successor owns bypass for its own duration.
|
||||
func TestRecoveryCanceledCleanup_SupersededKeepsSuccessorState(t *testing.T) {
|
||||
p := &prog{}
|
||||
setupInterceptRecovery(t, p)
|
||||
gen := p.recoveryGen.Add(1)
|
||||
// A successor recovery started.
|
||||
p.recoveryGen.Add(1)
|
||||
|
||||
p.recoveryCanceledCleanup(gen)
|
||||
|
||||
if !p.recoveryBypass.Load() {
|
||||
t.Error("superseded canceled recovery cleared recoveryBypass owned by its successor")
|
||||
}
|
||||
if !p.recoveryRunning.Load() {
|
||||
t.Error("superseded canceled recovery cleared recoveryRunning owned by its successor")
|
||||
}
|
||||
p.recoveryCancelMu.Lock()
|
||||
cancelKept := p.recoveryCancel != nil
|
||||
p.recoveryCancelMu.Unlock()
|
||||
if !cancelKept {
|
||||
t.Error("superseded canceled recovery cleared the successor's recoveryCancel")
|
||||
}
|
||||
}
|
||||
|
||||
// TestRecoveryCanceledCleanup_NonInterceptResetsRunning covers traditional
|
||||
// (non-intercept) mode: recoveryRunning must still be reset so watchdogs
|
||||
// resume, while bypass is untouched (it is never set in that mode).
|
||||
func TestRecoveryCanceledCleanup_NonInterceptResetsRunning(t *testing.T) {
|
||||
oldIntercept := dnsIntercept
|
||||
dnsIntercept = false
|
||||
t.Cleanup(func() { dnsIntercept = oldIntercept })
|
||||
|
||||
p := &prog{}
|
||||
p.recoveryRunning.Store(true)
|
||||
p.recoveryCancel = func() {}
|
||||
gen := p.recoveryGen.Add(1)
|
||||
|
||||
p.recoveryCanceledCleanup(gen)
|
||||
|
||||
if p.recoveryRunning.Load() {
|
||||
t.Error("recoveryRunning still set after canceled non-intercept recovery")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBeginRecoveryTransfersOwnershipAtomically(t *testing.T) {
|
||||
oldIntercept := dnsIntercept
|
||||
dnsIntercept = true
|
||||
t.Cleanup(func() { dnsIntercept = oldIntercept })
|
||||
|
||||
p := &prog{dnsInterceptState: &interceptStateStub{}}
|
||||
firstCtx, firstGen, _, ok := p.beginRecovery(RecoveryReasonRegularFailure)
|
||||
if !ok {
|
||||
t.Fatal("first recovery did not acquire ownership")
|
||||
}
|
||||
if _, _, _, ok := p.beginRecovery(RecoveryReasonRegularFailure); ok {
|
||||
t.Fatal("duplicate upstream recovery acquired ownership")
|
||||
}
|
||||
|
||||
_, successorGen, intercept, ok := p.beginRecovery(RecoveryReasonNetworkChange)
|
||||
if !ok || !intercept || successorGen <= firstGen {
|
||||
t.Fatalf("network recovery did not replace owner: first=%d successor=%d intercept=%v ok=%v", firstGen, successorGen, intercept, ok)
|
||||
}
|
||||
select {
|
||||
case <-firstCtx.Done():
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("successor did not cancel the previous recovery")
|
||||
}
|
||||
|
||||
p.recoveryCanceledCleanup(firstGen)
|
||||
if !p.recoveryRunning.Load() || !p.recoveryBypass.Load() || !p.recoveryOwnsState(successorGen) {
|
||||
t.Fatal("stale cleanup changed successor-owned recovery state")
|
||||
}
|
||||
if !p.completeRecovery(successorGen) {
|
||||
t.Fatal("successor could not complete its own recovery state")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecoveryCleanupClearsBypassAfterInterceptStateDisappears(t *testing.T) {
|
||||
p := &prog{}
|
||||
p.recoveryBypass.Store(true)
|
||||
p.recoveryRunning.Store(true)
|
||||
p.recoveryCancel = func() {}
|
||||
gen := p.recoveryGen.Add(1)
|
||||
|
||||
p.recoveryCanceledCleanup(gen)
|
||||
if p.recoveryBypass.Load() || p.recoveryRunning.Load() {
|
||||
t.Fatal("cleanup retained recovery flags after intercept state disappeared")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemNameserversForInterceptRetryNormalizesEmptyDiscovery(t *testing.T) {
|
||||
original := initializeOsResolverWithSystemNameserversFn
|
||||
called := false
|
||||
initializeOsResolverWithSystemNameserversFn = func(guard bool) ([]string, []string) {
|
||||
called = true
|
||||
if !guard {
|
||||
t.Error("intercept retry discovery did not guard the existing resolver")
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
t.Cleanup(func() { initializeOsResolverWithSystemNameserversFn = original })
|
||||
|
||||
if got := systemNameserversForInterceptRetry(); got == nil || len(got) != 0 {
|
||||
t.Fatalf("system discovery = %#v, want non-nil empty slice", got)
|
||||
}
|
||||
if !called {
|
||||
t.Fatal("system discovery was not called")
|
||||
}
|
||||
}
|
||||
+20
-10
@@ -162,6 +162,8 @@ func (s *systemd) Start() error {
|
||||
// This is necessary for running self-upgrade flow.
|
||||
func ensureSystemdKillMode(r io.Reader) (opts []*unit.UnitOption, change bool) {
|
||||
opts, err := unit.DeserializeOptions(r)
|
||||
// On success the lexer sends nothing and closes the channel, so the receive
|
||||
// yields a nil error and this branch is not taken.
|
||||
if err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to deserialize options")
|
||||
return
|
||||
@@ -216,22 +218,30 @@ type task struct {
|
||||
Name string
|
||||
}
|
||||
|
||||
func doTasks(tasks []task) bool {
|
||||
for _, task := range tasks {
|
||||
mainLog.Load().Debug().Msgf("Running task %s", task.Name)
|
||||
if err := task.f(); err != nil {
|
||||
if task.abortOnError {
|
||||
mainLog.Load().Error().Msgf("error running task %s: %v", task.Name, err)
|
||||
return false
|
||||
// doTasksE runs tasks in order and reports which abortOnError task, if any,
|
||||
// stopped the run. Use it over doTasks when the failure must be attributed
|
||||
// to a specific task.
|
||||
func doTasksE(tasks []task) (failedTaskName string, err error) {
|
||||
for _, t := range tasks {
|
||||
mainLog.Load().Debug().Msgf("Running task %s", t.Name)
|
||||
if taskErr := t.f(); taskErr != nil {
|
||||
if t.abortOnError {
|
||||
mainLog.Load().Error().Msgf("error running task %s: %v", t.Name, taskErr)
|
||||
return t.Name, taskErr
|
||||
}
|
||||
// if this is darwin stop command, dont print debug
|
||||
// since launchctl complains on every start
|
||||
if runtime.GOOS != "darwin" || task.Name != "Stop" {
|
||||
mainLog.Load().Debug().Msgf("error running task %s: %v", task.Name, err)
|
||||
if runtime.GOOS != "darwin" || t.Name != "Stop" {
|
||||
mainLog.Load().Debug().Msgf("error running task %s: %v", t.Name, taskErr)
|
||||
}
|
||||
}
|
||||
}
|
||||
return true
|
||||
return "", nil
|
||||
}
|
||||
|
||||
func doTasks(tasks []task) bool {
|
||||
_, err := doTasksE(tasks)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
func checkHasElevatedPrivilege() {
|
||||
|
||||
@@ -24,19 +24,19 @@ func serviceConfigFileExists() bool {
|
||||
// to intercept mode without losing the existing --cd flag and other arguments.
|
||||
//
|
||||
// On macOS, this modifies the launchd plist at /Library/LaunchDaemons/ctrld.plist
|
||||
// using the "defaults" command, which is the standard way to edit plists.
|
||||
// using PlistBuddy for exact array reads and writes.
|
||||
//
|
||||
// The function is idempotent: if the flag already exists, it's a no-op.
|
||||
func appendServiceFlag(flag string) error {
|
||||
// Read current ProgramArguments from plist.
|
||||
out, err := exec.Command("defaults", "read", launchdPlistPath, "ProgramArguments").CombinedOutput()
|
||||
out, err := exec.Command("/usr/libexec/PlistBuddy", "-c", "Print :ProgramArguments", launchdPlistPath).CombinedOutput()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read plist ProgramArguments: %w (output: %s)", err, strings.TrimSpace(string(out)))
|
||||
}
|
||||
|
||||
// Check if the flag is already present (idempotent).
|
||||
args := string(out)
|
||||
if strings.Contains(args, flag) {
|
||||
// Check exact array entries. A substring match can confuse a mode such as "off"
|
||||
// with an unrelated path or argument and leave the flag without its value.
|
||||
if serviceArgumentPresent(out, flag) {
|
||||
mainLog.Load().Debug().Msgf("Service flag %q already present in plist, skipping", flag)
|
||||
return nil
|
||||
}
|
||||
@@ -61,9 +61,8 @@ func verifyServiceRegistration() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// removeServiceFlag removes a CLI flag (and its value, if the next argument is not
|
||||
// a flag) from the installed service's launch arguments. For example, removing
|
||||
// "--intercept-mode" also removes the following "dns" or "hard" value argument.
|
||||
// removeServiceFlag removes both "--flag value" and "--flag=value" forms from the
|
||||
// installed service's launch arguments.
|
||||
//
|
||||
// The function is idempotent: if the flag doesn't exist, it's a no-op.
|
||||
func removeServiceFlag(flag string) error {
|
||||
@@ -92,22 +91,14 @@ func removeServiceFlag(flag string) error {
|
||||
entries = append(entries, trimmed)
|
||||
}
|
||||
|
||||
index := -1
|
||||
for i, entry := range entries {
|
||||
if entry == flag {
|
||||
index = i
|
||||
break
|
||||
}
|
||||
}
|
||||
index, hasValue := serviceFlagPosition(entries, flag)
|
||||
|
||||
if index < 0 {
|
||||
mainLog.Load().Debug().Msgf("Service flag %q not present in plist, skipping removal", flag)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Check if the next entry is a value (not a flag). If so, delete it first
|
||||
// (deleting by index shifts subsequent entries down, so delete value before flag).
|
||||
hasValue := index+1 < len(entries) && !strings.HasPrefix(entries[index+1], "-")
|
||||
// Delete a separate value first. An inline --flag=value entry is one array item.
|
||||
if hasValue {
|
||||
delVal := exec.Command(
|
||||
"/usr/libexec/PlistBuddy",
|
||||
@@ -132,3 +123,24 @@ func removeServiceFlag(flag string) error {
|
||||
mainLog.Load().Info().Msgf("Removed %q from service launch arguments", flag)
|
||||
return nil
|
||||
}
|
||||
|
||||
func serviceArgumentPresent(out []byte, argument string) bool {
|
||||
for _, line := range strings.Split(string(out), "\n") {
|
||||
if strings.TrimSpace(line) == argument {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func serviceFlagPosition(entries []string, flag string) (index int, hasValue bool) {
|
||||
for i, entry := range entries {
|
||||
switch {
|
||||
case entry == flag:
|
||||
return i, i+1 < len(entries) && !strings.HasPrefix(entries[i+1], "-")
|
||||
case strings.HasPrefix(entry, flag+"="):
|
||||
return i, false
|
||||
}
|
||||
}
|
||||
return -1, false
|
||||
}
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
//go:build darwin
|
||||
|
||||
package cli
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestServiceArgumentPresent(t *testing.T) {
|
||||
out := []byte("Array {\n /usr/local/bin/ctrld\n run\n --config=/Users/officer/ctrld.toml\n --intercept-mode=dns\n}\n")
|
||||
if !serviceArgumentPresent(out, "--intercept-mode=dns") {
|
||||
t.Fatal("exact inline argument was not found")
|
||||
}
|
||||
if serviceArgumentPresent(out, "--intercept-mode") {
|
||||
t.Fatal("inline flag was mistaken for a separate flag argument")
|
||||
}
|
||||
if serviceArgumentPresent(out, "off") {
|
||||
t.Fatal("substring in an unrelated path was mistaken for the off argument")
|
||||
}
|
||||
|
||||
splitOut := []byte("Array {\n /usr/local/bin/ctrld\n run\n --intercept-mode\n dns\n}\n")
|
||||
if !serviceArgumentPresent(splitOut, "--intercept-mode") {
|
||||
t.Fatal("standalone flag argument was not found")
|
||||
}
|
||||
}
|
||||
|
||||
func TestServiceFlagPosition(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
entries []string
|
||||
wantIndex int
|
||||
wantHasValue bool
|
||||
}{
|
||||
{
|
||||
name: "split form",
|
||||
entries: []string{"run", "--cd=uid", "--intercept-mode", "dns"},
|
||||
wantIndex: 2,
|
||||
wantHasValue: true,
|
||||
},
|
||||
{
|
||||
name: "inline form",
|
||||
entries: []string{"run", "--cd=uid", "--intercept-mode=dns"},
|
||||
wantIndex: 2,
|
||||
},
|
||||
{
|
||||
name: "flag followed by another flag",
|
||||
entries: []string{"run", "--intercept-mode", "--config=/etc/ctrld.toml"},
|
||||
wantIndex: 1,
|
||||
},
|
||||
{
|
||||
name: "absent",
|
||||
entries: []string{"run", "--cd=uid"},
|
||||
wantIndex: -1,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
index, hasValue := serviceFlagPosition(tc.entries, "--intercept-mode")
|
||||
if index != tc.wantIndex || hasValue != tc.wantHasValue {
|
||||
t.Fatalf("serviceFlagPosition() = (%d, %v), want (%d, %v)", index, hasValue, tc.wantIndex, tc.wantHasValue)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -3,10 +3,14 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"errors"
|
||||
"os"
|
||||
)
|
||||
|
||||
// errServiceFlagsUnsupported is returned by the service-argument helpers on
|
||||
// platforms that do not store service arguments in a file ctrld can rewrite.
|
||||
var errServiceFlagsUnsupported = errors.New("modifying service flags is not supported on this platform; use intercept_mode in config instead")
|
||||
|
||||
// serviceConfigFileExists checks common service config file locations on Linux.
|
||||
func serviceConfigFileExists() bool {
|
||||
// systemd unit file
|
||||
@@ -24,7 +28,7 @@ func serviceConfigFileExists() bool {
|
||||
// Linux services (systemd) store args in unit files; intercept mode
|
||||
// should be set via the config file (intercept_mode) on these platforms.
|
||||
func appendServiceFlag(flag string) error {
|
||||
return fmt.Errorf("appending service flags is not supported on this platform; use intercept_mode in config instead")
|
||||
return errServiceFlagsUnsupported
|
||||
}
|
||||
|
||||
// verifyServiceRegistration is a no-op on this platform.
|
||||
@@ -34,5 +38,5 @@ func verifyServiceRegistration() error {
|
||||
|
||||
// removeServiceFlag is not yet implemented on this platform.
|
||||
func removeServiceFlag(flag string) error {
|
||||
return fmt.Errorf("removing service flags is not supported on this platform; use intercept_mode in config instead")
|
||||
return errServiceFlagsUnsupported
|
||||
}
|
||||
|
||||
@@ -47,8 +47,9 @@ func appendServiceFlag(flag string) error {
|
||||
return fmt.Errorf("failed to read service config: %w", err)
|
||||
}
|
||||
|
||||
// Check if flag already present (idempotent).
|
||||
if strings.Contains(config.BinaryPathName, flag) {
|
||||
// Check exact arguments so a short mode such as "off" is not confused with
|
||||
// an unrelated path or value.
|
||||
if binaryPathArgumentPresent(config.BinaryPathName, flag) {
|
||||
mainLog.Load().Debug().Msgf("Service flag %q already present in BinPath, skipping", flag)
|
||||
return nil
|
||||
}
|
||||
@@ -88,7 +89,7 @@ func verifyServiceRegistration() error {
|
||||
mainLog.Load().Debug().Msgf("Service registry: BinaryPathName = %q", config.BinaryPathName)
|
||||
|
||||
// If intercept mode is set, verify the flag is present in BinPath.
|
||||
if interceptMode == "dns" || interceptMode == "hard" {
|
||||
if interceptMode == "off" || interceptMode == "dns" || interceptMode == "hard" {
|
||||
if !strings.Contains(config.BinaryPathName, "--intercept-mode") {
|
||||
return fmt.Errorf("service registry: --intercept-mode flag missing from BinaryPathName (expected mode %q)", interceptMode)
|
||||
}
|
||||
@@ -103,9 +104,8 @@ func verifyServiceRegistration() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// removeServiceFlag removes a CLI flag (and its value, if present) from the installed
|
||||
// Windows service's BinPath. For example, removing "--intercept-mode" also removes
|
||||
// the following "dns" or "hard" value. The function is idempotent.
|
||||
// removeServiceFlag removes both "--flag value" and "--flag=value" forms from the
|
||||
// installed Windows service's BinPath. The function is idempotent.
|
||||
func removeServiceFlag(flag string) error {
|
||||
m, err := mgr.Connect()
|
||||
if err != nil {
|
||||
@@ -124,25 +124,12 @@ func removeServiceFlag(flag string) error {
|
||||
return fmt.Errorf("failed to read service config: %w", err)
|
||||
}
|
||||
|
||||
if !strings.Contains(config.BinaryPathName, flag) {
|
||||
updatedPath, removed := removeBinaryPathFlag(config.BinaryPathName, flag)
|
||||
if !removed {
|
||||
mainLog.Load().Debug().Msgf("Service flag %q not present in BinPath, skipping removal", flag)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Split BinPath into parts, find and remove the flag + its value (if any).
|
||||
parts := strings.Fields(config.BinaryPathName)
|
||||
var newParts []string
|
||||
for i := 0; i < len(parts); i++ {
|
||||
if parts[i] == flag {
|
||||
// Skip the flag. Also skip the next part if it's a value (not a flag).
|
||||
if i+1 < len(parts) && !strings.HasPrefix(parts[i+1], "-") {
|
||||
i++ // skip value too
|
||||
}
|
||||
continue
|
||||
}
|
||||
newParts = append(newParts, parts[i])
|
||||
}
|
||||
config.BinaryPathName = strings.Join(newParts, " ")
|
||||
config.BinaryPathName = updatedPath
|
||||
|
||||
if err := s.UpdateConfig(config); err != nil {
|
||||
return fmt.Errorf("failed to update service config: %w", err)
|
||||
@@ -151,3 +138,32 @@ func removeServiceFlag(flag string) error {
|
||||
mainLog.Load().Info().Msgf("Removed %q from service BinPath", flag)
|
||||
return nil
|
||||
}
|
||||
|
||||
func binaryPathArgumentPresent(binaryPath, argument string) bool {
|
||||
for _, part := range strings.Fields(binaryPath) {
|
||||
if part == argument {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func removeBinaryPathFlag(binaryPath, flag string) (string, bool) {
|
||||
parts := strings.Fields(binaryPath)
|
||||
newParts := make([]string, 0, len(parts))
|
||||
removed := false
|
||||
for i := 0; i < len(parts); i++ {
|
||||
switch {
|
||||
case parts[i] == flag:
|
||||
removed = true
|
||||
if i+1 < len(parts) && !strings.HasPrefix(parts[i+1], "-") {
|
||||
i++
|
||||
}
|
||||
case strings.HasPrefix(parts[i], flag+"="):
|
||||
removed = true
|
||||
default:
|
||||
newParts = append(newParts, parts[i])
|
||||
}
|
||||
}
|
||||
return strings.Join(newParts, " "), removed
|
||||
}
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
//go:build windows
|
||||
|
||||
package cli
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestBinaryPathArgumentPresent(t *testing.T) {
|
||||
path := `C:\ControlD\ctrld.exe run --config=C:\Users\officer\ctrld.toml --intercept-mode=dns`
|
||||
if !binaryPathArgumentPresent(path, "--intercept-mode=dns") {
|
||||
t.Fatal("exact inline argument was not found")
|
||||
}
|
||||
if binaryPathArgumentPresent(path, "--intercept-mode") {
|
||||
t.Fatal("inline flag was mistaken for a separate flag argument")
|
||||
}
|
||||
if binaryPathArgumentPresent(path, "off") {
|
||||
t.Fatal("substring in an unrelated path was mistaken for the off argument")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveBinaryPathFlag(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
binaryPath string
|
||||
wantPath string
|
||||
wantRemoved bool
|
||||
}{
|
||||
{
|
||||
name: "split form",
|
||||
binaryPath: `ctrld.exe run --cd=uid --intercept-mode dns --config=ctrld.toml`,
|
||||
wantPath: `ctrld.exe run --cd=uid --config=ctrld.toml`,
|
||||
wantRemoved: true,
|
||||
},
|
||||
{
|
||||
name: "inline form",
|
||||
binaryPath: `ctrld.exe run --cd=uid --intercept-mode=dns --config=ctrld.toml`,
|
||||
wantPath: `ctrld.exe run --cd=uid --config=ctrld.toml`,
|
||||
wantRemoved: true,
|
||||
},
|
||||
{
|
||||
name: "absent",
|
||||
binaryPath: `ctrld.exe run --cd=uid`,
|
||||
wantPath: `ctrld.exe run --cd=uid`,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
path, removed := removeBinaryPathFlag(tc.binaryPath, "--intercept-mode")
|
||||
if path != tc.wantPath || removed != tc.wantRemoved {
|
||||
t.Fatalf("removeBinaryPathFlag() = (%q, %v), want (%q, %v)", path, removed, tc.wantPath, tc.wantRemoved)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package cli
|
||||
|
||||
import "strings"
|
||||
|
||||
// serviceBinaryFromImagePath extracts the executable path from a Windows service
|
||||
// ImagePath value, which carries the command line rather than a bare path: it may be
|
||||
// quoted and is usually followed by arguments, e.g.
|
||||
//
|
||||
// "C:\Program Files\Control D\ctrld.exe" run --config C:\...\ctrld.toml
|
||||
//
|
||||
// It returns "" when no path can be read, which callers must treat as "cannot tell"
|
||||
// rather than "does not match".
|
||||
func serviceBinaryFromImagePath(imagePath string) string {
|
||||
imagePath = strings.TrimSpace(imagePath)
|
||||
if imagePath == "" {
|
||||
return ""
|
||||
}
|
||||
if imagePath[0] == '"' {
|
||||
// Quoted form: everything up to the closing quote is the path, so a directory
|
||||
// containing spaces stays intact.
|
||||
if end := strings.IndexByte(imagePath[1:], '"'); end >= 0 {
|
||||
return strings.TrimSpace(imagePath[1 : 1+end])
|
||||
}
|
||||
return strings.TrimSpace(imagePath[1:])
|
||||
}
|
||||
// Unquoted form: the path cannot contain spaces, so the first field is it.
|
||||
if idx := strings.IndexByte(imagePath, ' '); idx >= 0 {
|
||||
return strings.TrimSpace(imagePath[:idx])
|
||||
}
|
||||
return imagePath
|
||||
}
|
||||
|
||||
// sameExecutableDir reports whether two Windows executable paths live in the same
|
||||
// directory, compared case-insensitively because Windows paths are.
|
||||
//
|
||||
// The separator handling is explicit rather than filepath's, because filepath follows the
|
||||
// *host* rules: off Windows it does not treat "\\" as a separator, so every backslash path
|
||||
// would reduce to the same directory and any two paths would compare equal. Doing it here
|
||||
// keeps the comparison correct and testable on any host.
|
||||
//
|
||||
// A path with no directory part answers false, which callers read as "cannot tell".
|
||||
func sameExecutableDir(a, b string) bool {
|
||||
dirA, dirB := windowsExecutableDir(a), windowsExecutableDir(b)
|
||||
if dirA == "" || dirB == "" {
|
||||
return false
|
||||
}
|
||||
return strings.EqualFold(dirA, dirB)
|
||||
}
|
||||
|
||||
// windowsExecutableDir returns the directory part of a Windows path, accepting either
|
||||
// separator and normalising to a backslash. It returns "" when there is no directory part.
|
||||
func windowsExecutableDir(path string) string {
|
||||
path = strings.TrimSpace(path)
|
||||
idx := strings.LastIndexAny(path, `\/`)
|
||||
if idx <= 0 {
|
||||
return ""
|
||||
}
|
||||
return strings.ReplaceAll(path[:idx], "/", `\`)
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
//go:build !windows
|
||||
|
||||
package cli
|
||||
|
||||
// installedServiceDirMatches is Windows-only: it exists because socketDir() there is
|
||||
// relative to the running executable. Other platforms answer this question through
|
||||
// hasElevatedPrivilege in readinessVerifiable.
|
||||
func installedServiceDirMatches() bool { return true }
|
||||
@@ -0,0 +1,103 @@
|
||||
package cli
|
||||
|
||||
import "testing"
|
||||
|
||||
// TestServiceBinaryFromImagePath covers the ImagePath shapes Windows stores. Getting this
|
||||
// wrong makes readinessVerifiable compare the wrong directories, and "ctrld status" would
|
||||
// then report a healthy service as not-ready - the false positive the readiness exit code
|
||||
// exists to avoid.
|
||||
func TestServiceBinaryFromImagePath(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
imagePath string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
// The installed form: quoted because the directory contains a space, with the
|
||||
// service arguments following it.
|
||||
name: "quoted path with arguments",
|
||||
imagePath: `"C:\Program Files\Control D\ctrld.exe" run --config "C:\ProgramData\Control D\ctrld.toml"`,
|
||||
want: `C:\Program Files\Control D\ctrld.exe`,
|
||||
},
|
||||
{
|
||||
name: "quoted path without arguments",
|
||||
imagePath: `"C:\Program Files\Control D\ctrld.exe"`,
|
||||
want: `C:\Program Files\Control D\ctrld.exe`,
|
||||
},
|
||||
{
|
||||
name: "unquoted path with arguments",
|
||||
imagePath: `C:\ctrld\ctrld.exe run --cd abc123`,
|
||||
want: `C:\ctrld\ctrld.exe`,
|
||||
},
|
||||
{
|
||||
name: "unquoted path alone",
|
||||
imagePath: `C:\ctrld\ctrld.exe`,
|
||||
want: `C:\ctrld\ctrld.exe`,
|
||||
},
|
||||
{
|
||||
name: "surrounding whitespace",
|
||||
imagePath: ` "C:\ctrld\ctrld.exe" run `,
|
||||
want: `C:\ctrld\ctrld.exe`,
|
||||
},
|
||||
{
|
||||
// Unterminated quote: take what is there rather than returning nothing, since
|
||||
// "" means "cannot tell" and would silently disable the check.
|
||||
name: "unterminated quote",
|
||||
imagePath: `"C:\ctrld\ctrld.exe run`,
|
||||
want: `C:\ctrld\ctrld.exe run`,
|
||||
},
|
||||
{
|
||||
name: "empty",
|
||||
imagePath: "",
|
||||
want: "",
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := serviceBinaryFromImagePath(tc.imagePath); got != tc.want {
|
||||
t.Errorf("serviceBinaryFromImagePath(%q) = %q, want %q", tc.imagePath, got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestSameExecutableDir pins the comparison itself: Windows paths are case-insensitive, and
|
||||
// an empty side means "cannot tell", which must never read as a match.
|
||||
func TestSameExecutableDir(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
a string
|
||||
b string
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "same directory",
|
||||
a: `C:\Program Files\Control D\ctrld.exe`,
|
||||
b: `C:\Program Files\Control D\ctrld.exe`,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "same directory different case",
|
||||
a: `C:\Program Files\Control D\ctrld.exe`,
|
||||
b: `c:\program files\control d\ctrld.exe`,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
// The case the check exists for: a copy run from a download directory
|
||||
// resolves a different control socket than the installed service.
|
||||
name: "different directory",
|
||||
a: `C:\Program Files\Control D\ctrld.exe`,
|
||||
b: `C:\Users\admin\Downloads\ctrld.exe`,
|
||||
want: false,
|
||||
},
|
||||
{name: "unknown installed path", a: "", b: `C:\ctrld\ctrld.exe`, want: false},
|
||||
{name: "unknown self path", a: `C:\ctrld\ctrld.exe`, b: "", want: false},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := sameExecutableDir(tc.a, tc.b); got != tc.want {
|
||||
t.Errorf("sameExecutableDir(%q, %q) = %v, want %v", tc.a, tc.b, got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
//go:build windows
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"os"
|
||||
|
||||
"golang.org/x/sys/windows/registry"
|
||||
)
|
||||
|
||||
// installedServiceDirMatches reports whether this executable is the installed service
|
||||
// binary, by comparing its directory with the one in the service's registered ImagePath.
|
||||
//
|
||||
// socketDir() on Windows is relative to the running executable, so a ctrld.exe run from
|
||||
// somewhere else - a download directory, a build tree - looks for the control socket in
|
||||
// its own directory and never finds the installed daemon's. A failed probe from there
|
||||
// says nothing about the service's health, and reporting "not ready" for it would tell
|
||||
// monitoring to restart a healthy service.
|
||||
//
|
||||
// Anything unreadable answers true, keeping the previous behaviour: readiness stays
|
||||
// verifiable unless there is positive evidence of a different install.
|
||||
func installedServiceDirMatches() bool {
|
||||
self, err := os.Executable()
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
key, err := registry.OpenKey(registry.LOCAL_MACHINE, `SYSTEM\CurrentControlSet\Services\`+ctrldServiceName, registry.QUERY_VALUE)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
defer key.Close()
|
||||
imagePath, _, err := key.GetStringValue("ImagePath")
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
installed := serviceBinaryFromImagePath(imagePath)
|
||||
if installed == "" {
|
||||
return true
|
||||
}
|
||||
return sameExecutableDir(installed, self)
|
||||
}
|
||||
@@ -0,0 +1,177 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Exit codes reported by "ctrld status".
|
||||
const (
|
||||
statusExitRunning = 0
|
||||
statusExitStopped = 1
|
||||
statusExitUnknown = 2
|
||||
// statusExitNotReady means the service manager considers the service running,
|
||||
// but the process has not finished starting up, so it is not serving DNS or
|
||||
// applying policy. This is a distinct code because it needs a distinct response:
|
||||
// the process exists, so restarting the service is what recovers it, while a
|
||||
// stopped service needs starting and an unknown state needs investigation.
|
||||
statusExitNotReady = 3
|
||||
)
|
||||
|
||||
// serviceReadinessTimeout bounds the control-socket probe. Status must answer
|
||||
// quickly, and a service that cannot respond within this window is not usefully
|
||||
// "running" from a caller's point of view either way.
|
||||
const serviceReadinessTimeout = 3 * time.Second
|
||||
|
||||
// statusCmdLong documents what the reported states mean, including that a service the
|
||||
// OS calls running is not necessarily serving.
|
||||
const statusCmdLong = `Show status of the ctrld service.
|
||||
|
||||
Reports both what the OS service manager thinks and whether ctrld has finished
|
||||
starting up, since a service can be registered as running while its process is
|
||||
still in startup and serving nothing.
|
||||
|
||||
Exit codes:
|
||||
0 running and serving, or running with startup not verified
|
||||
1 stopped
|
||||
2 status unknown
|
||||
3 registered as running, but startup has not completed
|
||||
|
||||
Verifying startup requires reaching ctrld's control socket. On Linux, BSD and macOS
|
||||
that socket lives in a directory only the privileged user resolves, so an
|
||||
unprivileged "ctrld status" reports the service manager's view and says startup was
|
||||
not verified rather than claiming the service is unhealthy. Exit 3 is only reported
|
||||
when the check could actually be made.`
|
||||
|
||||
// readiness is what "ctrld status" reports for a service the service manager
|
||||
// considers running.
|
||||
type readiness struct {
|
||||
messages []string
|
||||
exitCode int
|
||||
}
|
||||
|
||||
// readinessVerifiable reports whether a failed control-socket probe can be trusted to
|
||||
// mean "the service has not finished starting up".
|
||||
//
|
||||
// It can only mean that if this process resolves the same socket path the daemon
|
||||
// created, and socketDir() is caller-relative on unix: it returns the system directory
|
||||
// only when that is writable, and the caller's home directory otherwise. So a
|
||||
// root-owned daemon listens on /var/run/ctrld_control.sock while an unprivileged
|
||||
// "ctrld status" looks under $HOME, finds nothing, and gets ENOENT - which means "wrong
|
||||
// path", not "not ready". Reporting exit 3 there would tell a monitoring check to
|
||||
// restart a perfectly healthy daemon.
|
||||
//
|
||||
// On Windows and mobile socketDir() is the install/home directory for every caller, so
|
||||
// the probe is comparable - which matters because Windows is where the hung-start this
|
||||
// exit code exists for was seen. On Windows that only holds while this binary is the
|
||||
// installed one: a copy run from elsewhere resolves a different socket directory, so its
|
||||
// failed probe would say nothing about the service. installedServiceDirMatches() checks
|
||||
// that, and answers true when it cannot tell, preserving the previous behaviour.
|
||||
func readinessVerifiable() bool {
|
||||
if isMobile() {
|
||||
return true
|
||||
}
|
||||
if runtime.GOOS == "windows" {
|
||||
return installedServiceDirMatches()
|
||||
}
|
||||
elevated, err := hasElevatedPrivilege()
|
||||
return err == nil && elevated
|
||||
}
|
||||
|
||||
// classifyReadiness turns a control-socket probe result into the report for a service
|
||||
// the service manager calls running.
|
||||
//
|
||||
// verifiable comes from readinessVerifiable: when it is false a failed probe says
|
||||
// nothing about the service, so the report falls back to the service manager's view.
|
||||
// A *successful* probe is still conclusive either way - reaching the socket at all is
|
||||
// positive evidence, whoever the caller is.
|
||||
func classifyReadiness(ready bool, err error, verifiable bool) readiness {
|
||||
switch {
|
||||
case ready:
|
||||
return readiness{
|
||||
messages: []string{"Service is running"},
|
||||
exitCode: statusExitRunning,
|
||||
}
|
||||
case !verifiable:
|
||||
return readiness{
|
||||
messages: []string{"Service is running (startup not verified: re-run with elevated privileges to check readiness)"},
|
||||
exitCode: statusExitRunning,
|
||||
}
|
||||
case errors.Is(err, errReadinessNotReported):
|
||||
// The service answered, just not with a verdict - an older daemon without the
|
||||
// /started route. It is alive and reachable, so the service manager's view is
|
||||
// the best available answer.
|
||||
return readiness{
|
||||
messages: []string{"Service is running (startup not verified: this ctrld build does not report readiness)"},
|
||||
exitCode: statusExitRunning,
|
||||
}
|
||||
case errors.Is(err, fs.ErrPermission):
|
||||
// Without access to the control socket there is nothing to report beyond the
|
||||
// service manager's view. Do not call a service unhealthy because the caller
|
||||
// lacks privilege.
|
||||
return readiness{
|
||||
messages: []string{"Service is running (startup not verified: control socket requires elevated privileges)"},
|
||||
exitCode: statusExitRunning,
|
||||
}
|
||||
default:
|
||||
return readiness{
|
||||
messages: []string{
|
||||
"Service is registered as running, but has not completed startup: it is not serving DNS",
|
||||
"Check the ctrld log for why startup did not finish, then restart the service",
|
||||
},
|
||||
exitCode: statusExitNotReady,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// serviceReady reports whether a running ctrld has finished starting up, by asking
|
||||
// its control server. The control server answers /started only once the onStarted
|
||||
// hooks have completed, which is after the DNS listeners are up, so a successful
|
||||
// probe means the process is actually serving rather than merely alive.
|
||||
//
|
||||
// An error means "could not confirm readiness" and is returned for the caller to
|
||||
// classify: a refused connection or missing socket is a process that never got that
|
||||
// far, while a permission error says nothing about the service's health.
|
||||
func serviceReady() (bool, error) {
|
||||
dir, err := socketDir()
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return serviceReadyAt(filepath.Join(dir, ControlSocketName()), serviceReadinessTimeout)
|
||||
}
|
||||
|
||||
// errReadinessNotReported marks a control server that answered without a readiness
|
||||
// verdict.
|
||||
//
|
||||
// http.Client.Post returns (resp, nil) for any status, so a daemon with no /started
|
||||
// route answers 404 and an internal failure answers 5xx - neither says the service has
|
||||
// not started. Reporting "not ready" there tells a monitoring check to restart a healthy
|
||||
// service, and it happens in normal operation: after an upgrade replaces the binary on
|
||||
// disk but before the service restarts, and throughout a mixed-version rollout.
|
||||
var errReadinessNotReported = errors.New("control server did not report readiness")
|
||||
|
||||
// serviceReadyAt is serviceReady against an explicit socket path and timeout.
|
||||
func serviceReadyAt(sockPath string, timeout time.Duration) (bool, error) {
|
||||
cc := newControlClient(sockPath)
|
||||
cc.c.Timeout = timeout
|
||||
resp, err := cc.post(startedPath, nil)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
switch resp.StatusCode {
|
||||
case http.StatusOK:
|
||||
return true, nil
|
||||
case http.StatusRequestTimeout:
|
||||
// The daemon's own verdict: its onStarted hooks have not completed. This is the
|
||||
// hung start statusExitNotReady exists for.
|
||||
return false, nil
|
||||
default:
|
||||
return false, fmt.Errorf("%w: HTTP %d", errReadinessNotReported, resp.StatusCode)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,281 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io/fs"
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// startControlSocket serves handler on a unix socket and returns its path.
|
||||
func startControlSocket(t *testing.T, handler http.HandlerFunc) string {
|
||||
t.Helper()
|
||||
// Keep the path short: unix socket paths have a low length limit.
|
||||
dir, err := os.MkdirTemp("", "ctrldsock")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = os.RemoveAll(dir) })
|
||||
|
||||
sockPath := filepath.Join(dir, "s.sock")
|
||||
ln, err := net.Listen("unix", sockPath)
|
||||
if err != nil {
|
||||
t.Skipf("cannot listen on a unix socket: %v", err)
|
||||
}
|
||||
mux := http.NewServeMux()
|
||||
mux.Handle(startedPath, handler)
|
||||
srv := &http.Server{Handler: mux}
|
||||
go func() { _ = srv.Serve(ln) }()
|
||||
t.Cleanup(func() { _ = srv.Close() })
|
||||
return sockPath
|
||||
}
|
||||
|
||||
func TestServiceReadyAt(t *testing.T) {
|
||||
t.Run("ready when the control server reports started", func(t *testing.T) {
|
||||
sock := startControlSocket(t, func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
})
|
||||
ready, err := serviceReadyAt(sock, time.Second)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !ready {
|
||||
t.Error("ready = false, want true")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("not ready when startup has not finished", func(t *testing.T) {
|
||||
// What /started returns when the onStarted hooks have not completed.
|
||||
sock := startControlSocket(t, func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusRequestTimeout)
|
||||
})
|
||||
ready, err := serviceReadyAt(sock, time.Second)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if ready {
|
||||
t.Error("ready = true for a control server that has not finished startup")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("not ready when there is no control socket", func(t *testing.T) {
|
||||
// The incident: the process was alive but had never created the socket, so
|
||||
// every control request was refused.
|
||||
ready, err := serviceReadyAt(filepath.Join(t.TempDir(), "absent.sock"), time.Second)
|
||||
if ready {
|
||||
t.Error("ready = true with no control socket")
|
||||
}
|
||||
if err == nil {
|
||||
t.Error("expected an error when the control socket does not exist")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("not ready when the probe times out", func(t *testing.T) {
|
||||
sock := startControlSocket(t, func(w http.ResponseWriter, r *http.Request) {
|
||||
time.Sleep(2 * time.Second)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
})
|
||||
ready, err := serviceReadyAt(sock, 50*time.Millisecond)
|
||||
if ready {
|
||||
t.Error("ready = true for a probe that timed out")
|
||||
}
|
||||
if err == nil {
|
||||
t.Error("expected an error when the probe times out")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestClassifyReadiness(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
ready bool
|
||||
err error
|
||||
verifiable bool
|
||||
wantCode int
|
||||
}{
|
||||
{
|
||||
name: "ready",
|
||||
ready: true,
|
||||
verifiable: true,
|
||||
wantCode: statusExitRunning,
|
||||
},
|
||||
{
|
||||
// The service manager says running, the process is not serving. This
|
||||
// must not report success.
|
||||
name: "running but never finished startup",
|
||||
err: errors.New("connect: connection refused"),
|
||||
verifiable: true,
|
||||
wantCode: statusExitNotReady,
|
||||
},
|
||||
{
|
||||
// A caller without privilege cannot probe; that is not evidence of a
|
||||
// broken service, so it must not be reported as one.
|
||||
name: "probe not permitted",
|
||||
err: fs.ErrPermission,
|
||||
verifiable: true,
|
||||
wantCode: statusExitRunning,
|
||||
},
|
||||
{
|
||||
name: "wrapped permission error",
|
||||
err: &net.OpError{Op: "dial", Err: fs.ErrPermission},
|
||||
verifiable: true,
|
||||
wantCode: statusExitRunning,
|
||||
},
|
||||
{
|
||||
// The P2: an unprivileged caller on unix resolves a socket path the
|
||||
// daemon never used, so the probe fails with ENOENT rather than a
|
||||
// permission error. That says nothing about the service and must not be
|
||||
// reported as unhealthy - a monitoring check acting on exit 3 would
|
||||
// restart a healthy daemon.
|
||||
name: "missing socket at an unverifiable path",
|
||||
err: &net.OpError{Op: "dial", Err: os.ErrNotExist},
|
||||
verifiable: false,
|
||||
wantCode: statusExitRunning,
|
||||
},
|
||||
{
|
||||
name: "connection refused at an unverifiable path",
|
||||
err: errors.New("connect: connection refused"),
|
||||
verifiable: false,
|
||||
wantCode: statusExitRunning,
|
||||
},
|
||||
{
|
||||
// A probe that actually reached the socket is conclusive whoever ran it.
|
||||
name: "successful probe is trusted even when unverifiable",
|
||||
ready: true,
|
||||
verifiable: false,
|
||||
wantCode: statusExitRunning,
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got := classifyReadiness(tc.ready, tc.err, tc.verifiable)
|
||||
if got.exitCode != tc.wantCode {
|
||||
t.Errorf("exitCode = %d, want %d", got.exitCode, tc.wantCode)
|
||||
}
|
||||
if len(got.messages) == 0 {
|
||||
t.Error("no message to report")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestReadinessVerifiableMatchesSocketVisibility is the closure test for the P2: the
|
||||
// not-ready verdict must only be reachable when this process resolves the same socket
|
||||
// directory the daemon uses.
|
||||
//
|
||||
// On unix that is the privileged user's path, so an unprivileged run - which is how
|
||||
// "ctrld status" is normally invoked, since only darwin has an elevation PreRun and the
|
||||
// root-level alias has none - must not be able to reach exit 3.
|
||||
func TestReadinessVerifiableMatchesSocketVisibility(t *testing.T) {
|
||||
verifiable := readinessVerifiable()
|
||||
|
||||
if runtime.GOOS == "windows" {
|
||||
if !verifiable {
|
||||
t.Error("on Windows every caller resolves the install directory, so the probe is always verifiable")
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
elevated, err := hasElevatedPrivilege()
|
||||
if err != nil {
|
||||
t.Skipf("cannot determine privilege: %v", err)
|
||||
}
|
||||
if verifiable != elevated {
|
||||
t.Errorf("readinessVerifiable() = %v, want %v (elevated)", verifiable, elevated)
|
||||
}
|
||||
|
||||
if !elevated {
|
||||
// The shape the review asked to assert: unprivileged, healthy daemon, and a
|
||||
// probe that cannot see its socket must still report running.
|
||||
dir, err := socketDir()
|
||||
if err != nil {
|
||||
t.Fatalf("socketDir(): %v", err)
|
||||
}
|
||||
if dir == "/var/run" {
|
||||
t.Skip("unprivileged but /var/run is writable, so the probe path does match")
|
||||
}
|
||||
r := classifyReadiness(false, &net.OpError{Op: "dial", Err: os.ErrNotExist}, verifiable)
|
||||
if r.exitCode == statusExitNotReady {
|
||||
t.Errorf("unprivileged status probing %q reported not-ready (exit %d) for a healthy service", dir, r.exitCode)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Every status must map to its own exit code: a caller that cannot tell a hung
|
||||
// service from a healthy or a stopped one is back to the incident's diagnostics.
|
||||
//
|
||||
// The literal values are the contract. statusCmdLong documents them and monitoring
|
||||
// scripts key off them, so asserting the constants against each other would let a
|
||||
// renumbering keep the suite green while silently breaking every caller.
|
||||
func TestStatusExitCodesAreDistinct(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
got int
|
||||
want int
|
||||
}{
|
||||
{"running", statusExitRunning, 0},
|
||||
{"stopped", statusExitStopped, 1},
|
||||
{"unknown", statusExitUnknown, 2},
|
||||
{"not ready", statusExitNotReady, 3},
|
||||
} {
|
||||
if tc.got != tc.want {
|
||||
t.Errorf("%s exit code = %d, want %d: statusCmdLong and monitoring scripts document this value", tc.name, tc.got, tc.want)
|
||||
}
|
||||
}
|
||||
|
||||
codes := map[int]string{
|
||||
statusExitRunning: "running",
|
||||
statusExitStopped: "stopped",
|
||||
statusExitUnknown: "unknown",
|
||||
statusExitNotReady: "not ready",
|
||||
}
|
||||
if len(codes) != 4 {
|
||||
t.Errorf("status exit codes collide, only %d distinct: %v", len(codes), codes)
|
||||
}
|
||||
}
|
||||
|
||||
// TestReadinessProbeStatusHandling covers what each control-server answer means.
|
||||
//
|
||||
// http.Client.Post returns (resp, nil) for any status code, so a daemon without the
|
||||
// /started route answers 404 and the probe must report "cannot confirm" rather than "not
|
||||
// started". That state is reached in normal operation - after an upgrade replaces the
|
||||
// binary but before the service restarts, and throughout a mixed-version rollout - and
|
||||
// reporting exit 3 there tells monitoring to restart a healthy service.
|
||||
func TestReadinessProbeStatusHandling(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
status int
|
||||
wantReady bool
|
||||
wantReported bool // whether the answer carries a readiness verdict
|
||||
wantExitCode int
|
||||
}{
|
||||
{"started", http.StatusOK, true, true, statusExitRunning},
|
||||
{"still starting", http.StatusRequestTimeout, false, true, statusExitNotReady},
|
||||
{"no readiness route", http.StatusNotFound, false, false, statusExitRunning},
|
||||
{"control server error", http.StatusInternalServerError, false, false, statusExitRunning},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
status := tc.status
|
||||
sock := startControlSocket(t, func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(status)
|
||||
})
|
||||
|
||||
ready, err := serviceReadyAt(sock, time.Second)
|
||||
if ready != tc.wantReady {
|
||||
t.Errorf("ready = %v, want %v", ready, tc.wantReady)
|
||||
}
|
||||
if reported := !errors.Is(err, errReadinessNotReported); reported != tc.wantReported {
|
||||
t.Errorf("readiness reported = %v, want %v (err: %v)", reported, tc.wantReported, err)
|
||||
}
|
||||
if got := classifyReadiness(ready, err, true).exitCode; got != tc.wantExitCode {
|
||||
t.Errorf("exit code = %d, want %d", got, tc.wantExitCode)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
@@ -26,3 +27,59 @@ func Test_ensureSystemdKillMode(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoTasksESuccess(t *testing.T) {
|
||||
var ran []string
|
||||
tasks := []task{
|
||||
{func() error { ran = append(ran, "a"); return nil }, false, "a"},
|
||||
{func() error { ran = append(ran, "b"); return nil }, true, "b"},
|
||||
}
|
||||
failedTask, err := doTasksE(tasks)
|
||||
if failedTask != "" || err != nil {
|
||||
t.Errorf("doTasksE() = (%q, %v), want (\"\", nil)", failedTask, err)
|
||||
}
|
||||
if got := strings.Join(ran, ","); got != "a,b" {
|
||||
t.Errorf("ran tasks %q, want all tasks run in order", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoTasksEAbortsOnAbortOnErrorTask(t *testing.T) {
|
||||
wantErr := errors.New("install failed")
|
||||
var ran []string
|
||||
tasks := []task{
|
||||
{func() error { ran = append(ran, "Stop"); return nil }, false, "Stop"},
|
||||
{func() error { ran = append(ran, "Install"); return wantErr }, true, "Install"},
|
||||
{func() error { ran = append(ran, "Start"); return nil }, true, "Start"},
|
||||
}
|
||||
failedTask, err := doTasksE(tasks)
|
||||
if failedTask != "Install" || !errors.Is(err, wantErr) {
|
||||
t.Errorf("doTasksE() = (%q, %v), want (\"Install\", %v)", failedTask, err, wantErr)
|
||||
}
|
||||
if got := strings.Join(ran, ","); got != "Stop,Install" {
|
||||
t.Errorf("ran tasks %q, want the run to stop right after the abort", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoTasksENonAbortFailureContinues(t *testing.T) {
|
||||
var ran []string
|
||||
tasks := []task{
|
||||
{func() error { ran = append(ran, "a"); return errors.New("a failed") }, false, "a"},
|
||||
{func() error { ran = append(ran, "b"); return nil }, true, "b"},
|
||||
}
|
||||
failedTask, err := doTasksE(tasks)
|
||||
if failedTask != "" || err != nil {
|
||||
t.Errorf("doTasksE() = (%q, %v), want (\"\", nil) since the failing task did not abort", failedTask, err)
|
||||
}
|
||||
if got := strings.Join(ran, ","); got != "a,b" {
|
||||
t.Errorf("ran tasks %q, want the run to continue past the non-abort failure", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoTasksDelegatesToDoTasksE(t *testing.T) {
|
||||
if !doTasks([]task{{func() error { return nil }, true, "ok"}}) {
|
||||
t.Error("doTasks() = false, want true on success")
|
||||
}
|
||||
if doTasks([]task{{func() error { return errors.New("boom") }, true, "boom"}}) {
|
||||
t.Error("doTasks() = true, want false when an abortOnError task fails")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
package cli
|
||||
|
||||
import "testing"
|
||||
|
||||
// Test_ensureRunningIfaceForInvalidUninstall is a regression test for issue-556:
|
||||
// after a reboot, the invalid-device self-uninstall path could run before the
|
||||
// running interface was known. Because resetDNS (via resetDNSForRunningIface)
|
||||
// silently skips DNS restoration when p.runningIface is empty, the OS was left
|
||||
// pointed at ctrld's local listener with no internet after the service was
|
||||
// removed. ensureRunningIfaceForInvalidUninstall must populate p.runningIface
|
||||
// before resetDNS runs.
|
||||
func Test_ensureRunningIfaceForInvalidUninstall(t *testing.T) {
|
||||
// preRun mutates the package-level iface global; restore it after the test.
|
||||
origIface := iface
|
||||
t.Cleanup(func() { iface = origIface })
|
||||
|
||||
// newService needs the package service config; it is safe to build here
|
||||
// because ensureRunningIfaceForInvalidUninstall only queries the (absent)
|
||||
// control socket via runningIface, which returns nil when ctrld is not
|
||||
// running, and performs no DNS or service mutation.
|
||||
s, err := newService(&prog{}, svcConfig)
|
||||
if err != nil {
|
||||
t.Fatalf("newService: %v", err)
|
||||
}
|
||||
|
||||
t.Run("iface flag already resolved but not yet copied", func(t *testing.T) {
|
||||
iface = "eth-test"
|
||||
p := &prog{}
|
||||
|
||||
// Precondition mirrors the buggy post-reboot state: an empty running
|
||||
// interface would make resetDNS skip restoration entirely.
|
||||
if p.runningIface != "" {
|
||||
t.Fatalf("precondition: runningIface = %q, want empty", p.runningIface)
|
||||
}
|
||||
|
||||
ensureRunningIfaceForInvalidUninstall(p, s)
|
||||
|
||||
if p.runningIface == "" {
|
||||
t.Fatal("runningIface still empty after prepare: resetDNS would skip DNS " +
|
||||
"restoration and leave the OS pointed at ctrld's local listener")
|
||||
}
|
||||
if p.runningIface != "eth-test" {
|
||||
t.Fatalf("runningIface = %q, want the resolved iface %q", p.runningIface, "eth-test")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("iface unset falls back to auto-detected interface", func(t *testing.T) {
|
||||
iface = ""
|
||||
p := &prog{}
|
||||
|
||||
ensureRunningIfaceForInvalidUninstall(p, s)
|
||||
|
||||
// With iface unset the prep resolves "auto" to the default interface
|
||||
// (defaultIfaceName never returns empty on the supported platforms), so
|
||||
// resetDNS has a concrete interface to restore.
|
||||
if p.runningIface == "" {
|
||||
t.Fatal("runningIface still empty after prepare with iface unset: " +
|
||||
"resetDNS would skip DNS restoration")
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,195 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/kardianos/service"
|
||||
)
|
||||
|
||||
const (
|
||||
// upgradeStopTimeout bounds how long rollback waits for the replacement process to
|
||||
// exit, and for Windows to release the lock on its image afterwards.
|
||||
upgradeStopTimeout = 30 * time.Second
|
||||
// upgradeStopPollInterval is how often the service status is re-checked while
|
||||
// waiting for the process to exit.
|
||||
upgradeStopPollInterval = 500 * time.Millisecond
|
||||
// binaryVersionTimeout bounds the "--version" probe, so a binary that hangs on
|
||||
// startup cannot hang the upgrade.
|
||||
binaryVersionTimeout = 10 * time.Second
|
||||
)
|
||||
|
||||
// rollbackToPreviousBinary restores oldBin over bin after the replacement failed to
|
||||
// become ready, and restarts the service on the restored binary.
|
||||
//
|
||||
// stop must leave the replacement's process gone, because every step here modifies
|
||||
// the executable that process is running from. It is called first for that reason:
|
||||
// readiness failing does not mean the process exited - the service manager can report
|
||||
// a started service whose process never became operational. Windows holds an
|
||||
// exclusive lock on a running executable's image, so the previous code's
|
||||
// os.Remove(bin) failed there with "Access is denied", and because that was fatal the
|
||||
// restore never ran: the broken binary stayed installed with the previous one
|
||||
// stranded at its _previous name.
|
||||
//
|
||||
// Stopping first also puts the host back in a known state, since a stopped ctrld
|
||||
// holds no DNS or intercept enforcement.
|
||||
func rollbackToPreviousBinary(bin, oldBin string, stop func() error, restart func() bool) error {
|
||||
if err := stop(); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("Could not confirm the service stopped; not modifying its binary")
|
||||
return err
|
||||
}
|
||||
|
||||
// Only restore a previous binary that actually runs: a _previous file that exists
|
||||
// but reports no version would replace a service that starts and hangs with one
|
||||
// that cannot start at all.
|
||||
//
|
||||
// The probe is retried for the same reason removeBinaryWithRetry is: on Windows a
|
||||
// single exec can fail transiently while antivirus scans the file or the disk is
|
||||
// busy, and treating that as "no usable previous binary" leaves the host stopped
|
||||
// with the broken binary installed - an end state worse than restoring a binary
|
||||
// that turns out to be bad, which the restart check below catches.
|
||||
//
|
||||
// Running "--version" proves the file executes. It is not an authenticity check:
|
||||
// nothing here compares a signature or checksum before a file becomes the installed
|
||||
// service binary. That is acceptable only because the install directory is writable
|
||||
// by administrators alone, which is this command's standing assumption.
|
||||
prevVer, err := binaryVersionWithRetry(oldBin, upgradeStopTimeout)
|
||||
if err != nil {
|
||||
mainLog.Load().Error().Err(err).Msgf("Previous binary at %s is not usable, keeping it for inspection", oldBin)
|
||||
mainLog.Load().Notice().Msgf("Service is stopped and %s is still the installed binary", bin)
|
||||
return fmt.Errorf("upgrade failed and no usable previous binary to restore: %w", err)
|
||||
}
|
||||
|
||||
mainLog.Load().Warn().Msgf("Restoring previous binary: %s (%s)", oldBin, prevVer)
|
||||
if err := removeBinaryWithRetry(bin, upgradeStopTimeout); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("Failed to remove new binary")
|
||||
mainLog.Load().Notice().Msg("Service is stopped")
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(oldBin, bin); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("Failed to restore old binary")
|
||||
mainLog.Load().Notice().Msgf("Service is stopped and %s is missing; reinstall ctrld to recover", bin)
|
||||
return err
|
||||
}
|
||||
if restart() {
|
||||
mainLog.Load().Notice().Msgf("Restored previous binary successfully - %s", prevVer)
|
||||
return nil
|
||||
}
|
||||
|
||||
mainLog.Load().Error().Msg("Restored the previous binary but it did not become ready either")
|
||||
return errors.New("upgrade failed and the restored binary did not become ready")
|
||||
}
|
||||
|
||||
// stopServiceAndWait stops the service and waits until the service manager reports
|
||||
// it stopped. Rollback needs the process gone, not merely asked to stop: a stop
|
||||
// request returns before the process exits, and on Windows the executable stays
|
||||
// locked until it does.
|
||||
func stopServiceAndWait(s service.Service, timeout time.Duration) error {
|
||||
if err := s.Stop(); err != nil {
|
||||
// Not fatal: the service may already be stopped, or stopping may fail while
|
||||
// the process is exiting anyway. The status poll below decides.
|
||||
mainLog.Load().Debug().Err(err).Msg("Stop request failed, waiting for the process to exit anyway")
|
||||
}
|
||||
deadline := time.Now().Add(timeout)
|
||||
statusReadable := false
|
||||
var lastErr error
|
||||
for {
|
||||
status, err := s.Status()
|
||||
switch {
|
||||
case errors.Is(err, service.ErrNotInstalled):
|
||||
return nil
|
||||
case err == nil:
|
||||
statusReadable = true
|
||||
if status == service.StatusStopped {
|
||||
return nil
|
||||
}
|
||||
default:
|
||||
lastErr = err
|
||||
}
|
||||
if !time.Now().Before(deadline) {
|
||||
if !statusReadable {
|
||||
// The status was never readable, so "did not stop" was never observed -
|
||||
// only "could not be observed". Refusing to continue here would leave the
|
||||
// broken binary installed with the service stopped, which is the outcome
|
||||
// rollback exists to avoid. Let the caller proceed: the remove is retried
|
||||
// while the image is locked, and the restart check still has to pass
|
||||
// before this reports success.
|
||||
mainLog.Load().Warn().Err(lastErr).Msgf("Could not read service status within %s; continuing with rollback", timeout)
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("service did not stop within %s", timeout)
|
||||
}
|
||||
time.Sleep(upgradeStopPollInterval)
|
||||
}
|
||||
}
|
||||
|
||||
// binaryVersionWithRetry probes a binary's version, retrying transient exec failures
|
||||
// until timeout. Only the last error is reported: the earlier attempts are noise once a
|
||||
// retry has been made.
|
||||
func binaryVersionWithRetry(path string, timeout time.Duration) (string, error) {
|
||||
deadline := time.Now().Add(timeout)
|
||||
for {
|
||||
version, err := binaryVersionFn(path)
|
||||
if err == nil {
|
||||
return version, nil
|
||||
}
|
||||
if !time.Now().Before(deadline) {
|
||||
return "", err
|
||||
}
|
||||
mainLog.Load().Debug().Err(err).Msgf("Version probe of %s failed, retrying", path)
|
||||
time.Sleep(upgradeStopPollInterval)
|
||||
}
|
||||
}
|
||||
|
||||
// removeBinaryWithRetry removes path, retrying while it is still locked. Windows
|
||||
// releases the lock on an executable's image asynchronously after its process exits,
|
||||
// so a remove issued immediately after the service reports stopped can still fail
|
||||
// with "Access is denied".
|
||||
func removeBinaryWithRetry(path string, timeout time.Duration) error {
|
||||
deadline := time.Now().Add(timeout)
|
||||
for {
|
||||
err := os.Remove(path)
|
||||
if err == nil || errors.Is(err, os.ErrNotExist) {
|
||||
return nil
|
||||
}
|
||||
if !time.Now().Before(deadline) {
|
||||
return fmt.Errorf("could not remove %s within %s: %w", path, timeout, err)
|
||||
}
|
||||
time.Sleep(upgradeStopPollInterval)
|
||||
}
|
||||
}
|
||||
|
||||
// binaryVersionFn is indirected so rollback can be tested without staging a runnable
|
||||
// executable per platform. The probe itself is covered directly against the test
|
||||
// binary; see TestBinaryVersion.
|
||||
var binaryVersionFn = binaryVersion
|
||||
|
||||
// binaryVersion runs path with "--version" and returns the version it reports. It
|
||||
// answers "can this binary actually run on this host", which is what rollback needs
|
||||
// to know before making a file the installed ctrld.
|
||||
//
|
||||
// On Windows path is ctrld.exe_previous, whose extension is not in PATHEXT. That
|
||||
// resolves because os/exec only falls back to appending PATHEXT entries when the path
|
||||
// has no extension at all (lp_windows.go findExecutable): with one present and the
|
||||
// file on disk, it is used as-is. A suffix that left no extension - renaming
|
||||
// oldBinSuffix such that the result is "ctrld_previous" - would break this probe with
|
||||
// "executable file not found in %PATH%", and rollback would then refuse to restore a
|
||||
// perfectly good binary.
|
||||
func binaryVersion(path string) (string, error) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), binaryVersionTimeout)
|
||||
defer cancel()
|
||||
out, err := exec.CommandContext(ctx, path, "--version").CombinedOutput()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("running %s --version: %w", path, err)
|
||||
}
|
||||
ver, found := strings.CutPrefix(strings.TrimSpace(string(out)), "ctrld version ")
|
||||
if !found {
|
||||
return "", fmt.Errorf("unexpected --version output from %s: %q", path, strings.TrimSpace(string(out)))
|
||||
}
|
||||
return ver, nil
|
||||
}
|
||||
@@ -0,0 +1,288 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/kardianos/service"
|
||||
)
|
||||
|
||||
// fakeService implements the parts of service.Service that rollback uses. Any other
|
||||
// method panics, which keeps accidental dependencies visible.
|
||||
type fakeService struct {
|
||||
service.Service
|
||||
|
||||
stopErr error
|
||||
stopCalls int
|
||||
statuses []service.Status // consumed one per Status() call; the last repeats
|
||||
statusErr error
|
||||
onStopCall func()
|
||||
}
|
||||
|
||||
func (f *fakeService) Stop() error {
|
||||
f.stopCalls++
|
||||
if f.onStopCall != nil {
|
||||
f.onStopCall()
|
||||
}
|
||||
return f.stopErr
|
||||
}
|
||||
|
||||
func (f *fakeService) Status() (service.Status, error) {
|
||||
if f.statusErr != nil {
|
||||
return service.StatusUnknown, f.statusErr
|
||||
}
|
||||
if len(f.statuses) == 0 {
|
||||
return service.StatusStopped, nil
|
||||
}
|
||||
st := f.statuses[0]
|
||||
if len(f.statuses) > 1 {
|
||||
f.statuses = f.statuses[1:]
|
||||
}
|
||||
return st, nil
|
||||
}
|
||||
|
||||
func TestStopServiceAndWait(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
svc *fakeService
|
||||
timeout time.Duration
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "stops after a few polls",
|
||||
svc: &fakeService{statuses: []service.Status{service.StatusRunning, service.StatusRunning, service.StatusStopped}},
|
||||
timeout: 5 * time.Second,
|
||||
},
|
||||
{
|
||||
name: "already stopped",
|
||||
svc: &fakeService{statuses: []service.Status{service.StatusStopped}},
|
||||
timeout: 5 * time.Second,
|
||||
},
|
||||
{
|
||||
// A stop request that errors is not fatal on its own: the process may be
|
||||
// exiting anyway, so the status poll decides.
|
||||
name: "stop errors but service is stopped",
|
||||
svc: &fakeService{stopErr: errors.New("already stopped"), statuses: []service.Status{service.StatusStopped}},
|
||||
timeout: 5 * time.Second,
|
||||
},
|
||||
{
|
||||
name: "not installed",
|
||||
svc: &fakeService{statusErr: service.ErrNotInstalled},
|
||||
timeout: 5 * time.Second,
|
||||
},
|
||||
{
|
||||
// The process never exits. Rollback must be told so, because modifying a
|
||||
// running executable is what produced "Access is denied".
|
||||
name: "never stops",
|
||||
svc: &fakeService{statuses: []service.Status{service.StatusRunning}},
|
||||
timeout: time.Millisecond,
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
err := stopServiceAndWait(tc.svc, tc.timeout)
|
||||
if tc.wantErr && err == nil {
|
||||
t.Fatal("expected an error, got nil")
|
||||
}
|
||||
if !tc.wantErr && err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if tc.svc.stopCalls != 1 {
|
||||
t.Errorf("Stop() called %d times, want 1", tc.svc.stopCalls)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveBinaryWithRetry(t *testing.T) {
|
||||
t.Run("removes an existing file", func(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "ctrld")
|
||||
if err := os.WriteFile(path, []byte("binary"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := removeBinaryWithRetry(path, time.Second); err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if _, err := os.Stat(path); !errors.Is(err, os.ErrNotExist) {
|
||||
t.Errorf("file still exists after removal: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing file is not an error", func(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "absent")
|
||||
if err := removeBinaryWithRetry(path, time.Second); err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("gives up and reports when the path cannot be removed", func(t *testing.T) {
|
||||
// A non-empty directory stands in for a locked executable: os.Remove keeps
|
||||
// failing, so the retry loop must surface the error rather than hang.
|
||||
dir := filepath.Join(t.TempDir(), "locked")
|
||||
if err := os.Mkdir(dir, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(dir, "child"), nil, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := removeBinaryWithRetry(dir, time.Millisecond); err == nil {
|
||||
t.Fatal("expected an error for a path that cannot be removed")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestBinaryVersion(t *testing.T) {
|
||||
t.Run("reports the version", func(t *testing.T) {
|
||||
t.Setenv(envFakeVersionOutput, "ctrld version dev-94fbd3f")
|
||||
got, err := binaryVersion(os.Args[0])
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if got != "dev-94fbd3f" {
|
||||
t.Errorf("binaryVersion() = %q, want %q", got, "dev-94fbd3f")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("rejects a binary that prints no version", func(t *testing.T) {
|
||||
// A ctrld.exe_previous that exists and runs, but produces no version output.
|
||||
// Restoring it would replace a hung service with one that cannot start at all.
|
||||
t.Setenv(envFakeVersionOutput, envFakeVersionSilent)
|
||||
if _, err := binaryVersion(os.Args[0]); err == nil {
|
||||
t.Fatal("expected an error for a binary with no version output")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("rejects a missing binary", func(t *testing.T) {
|
||||
if _, err := binaryVersion(filepath.Join(t.TempDir(), "absent")); err == nil {
|
||||
t.Fatal("expected an error for a missing binary")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// stubBinaryVersion makes the version probe report ver for any path, so a rollback
|
||||
// test does not have to stage a runnable executable.
|
||||
//
|
||||
// Staging one is not portable: oldBin is bin+"_previous", so a fixture named "ctrld"
|
||||
// yields the extension-less "ctrld_previous", which Windows refuses to execute
|
||||
// ("executable file not found in %PATH%"), and a symlink to the test binary needs a
|
||||
// privilege Windows does not grant by default. The probe itself is covered against the
|
||||
// real test binary in TestBinaryVersion; these tests are about rollback's ordering.
|
||||
func stubBinaryVersion(t *testing.T, ver string, err error) {
|
||||
t.Helper()
|
||||
prev := binaryVersionFn
|
||||
binaryVersionFn = func(string) (string, error) { return ver, err }
|
||||
t.Cleanup(func() { binaryVersionFn = prev })
|
||||
}
|
||||
|
||||
func TestRollbackToPreviousBinaryStopsBeforeTouchingTheBinary(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
bin := filepath.Join(dir, "ctrld")
|
||||
oldBin := bin + oldBinSuffix
|
||||
if err := os.WriteFile(bin, []byte("replacement"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(oldBin, []byte("previous"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
stubBinaryVersion(t, "dev-a75d669", nil)
|
||||
|
||||
// The invariant: when stop runs, the replacement's executable is still untouched.
|
||||
// Reversing these two is exactly the "Access is denied" defect.
|
||||
var stopped bool
|
||||
var binExistedAtStop bool
|
||||
stop := func() error {
|
||||
stopped = true
|
||||
_, err := os.Stat(bin)
|
||||
binExistedAtStop = err == nil
|
||||
return nil
|
||||
}
|
||||
restarted := false
|
||||
restart := func() bool { restarted = true; return true }
|
||||
|
||||
if err := rollbackToPreviousBinary(bin, oldBin, stop, restart); err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !stopped {
|
||||
t.Error("rollback did not stop the service")
|
||||
}
|
||||
if !binExistedAtStop {
|
||||
t.Error("the binary was modified before the service was stopped")
|
||||
}
|
||||
if !restarted {
|
||||
t.Error("rollback did not restart the service")
|
||||
}
|
||||
if _, err := os.Stat(oldBin); !errors.Is(err, os.ErrNotExist) {
|
||||
t.Errorf("previous binary was not moved into place: %v", err)
|
||||
}
|
||||
if _, err := os.Stat(bin); err != nil {
|
||||
t.Errorf("restored binary is missing: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRollbackToPreviousBinaryKeepsUnusablePrevious(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
bin := filepath.Join(dir, "ctrld")
|
||||
oldBin := bin + oldBinSuffix
|
||||
if err := os.WriteFile(bin, []byte("replacement"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// A previous binary that exists but does not report a version.
|
||||
if err := os.WriteFile(oldBin, []byte("not a working binary"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Stubbed rather than left to the real probe: that would fail here for the right
|
||||
// reason on unix (not an executable) but the wrong one on Windows (the fixture's
|
||||
// name has no extension), so the assertion would not be about usability at all.
|
||||
stubBinaryVersion(t, "", errors.New("unexpected --version output"))
|
||||
|
||||
stopped := false
|
||||
restarted := false
|
||||
err := rollbackToPreviousBinary(bin, oldBin,
|
||||
func() error { stopped = true; return nil },
|
||||
func() bool { restarted = true; return true },
|
||||
)
|
||||
if err == nil {
|
||||
t.Fatal("expected an error when the previous binary is unusable")
|
||||
}
|
||||
if !stopped {
|
||||
t.Error("the service must still be stopped: a broken replacement holds enforcement")
|
||||
}
|
||||
if restarted {
|
||||
t.Error("must not restart the service with an unusable binary")
|
||||
}
|
||||
// Nothing was swapped, and the previous file is kept for inspection.
|
||||
if _, err := os.Stat(oldBin); err != nil {
|
||||
t.Errorf("unusable previous binary was not preserved: %v", err)
|
||||
}
|
||||
if _, err := os.Stat(bin); err != nil {
|
||||
t.Errorf("installed binary was removed despite having nothing to restore: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRollbackToPreviousBinaryAbortsWhenStopFails(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
bin := filepath.Join(dir, "ctrld")
|
||||
oldBin := bin + oldBinSuffix
|
||||
for _, p := range []string{bin, oldBin} {
|
||||
if err := os.WriteFile(p, []byte("binary"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
stopErr := errors.New("service did not stop within 30s")
|
||||
err := rollbackToPreviousBinary(bin, oldBin,
|
||||
func() error { return stopErr },
|
||||
func() bool { t.Error("must not restart after a failed stop"); return false },
|
||||
)
|
||||
if !errors.Is(err, stopErr) {
|
||||
t.Fatalf("error = %v, want %v", err, stopErr)
|
||||
}
|
||||
// The executable of a process that may still be running must be left alone.
|
||||
if _, err := os.Stat(bin); err != nil {
|
||||
t.Errorf("binary was modified even though the stop failed: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -12,8 +12,27 @@ const (
|
||||
maxFailureRequest = 50
|
||||
// checkUpstreamBackoffSleep is the time interval between each upstream checks.
|
||||
checkUpstreamBackoffSleep = 2 * time.Second
|
||||
// checkUpstreamUnreachableBackoffMax caps the recovery retry interval for an
|
||||
// endpoint that keeps failing with a network-unreachable error. It bounds
|
||||
// the backoff so an unroutable endpoint is still re-probed periodically and
|
||||
// recovers once the route returns.
|
||||
checkUpstreamUnreachableBackoffMax = 60 * time.Second
|
||||
)
|
||||
|
||||
// unreachableRecoveryBackoff returns the retry interval for the given streak of
|
||||
// consecutive network-unreachable failures. It starts at checkUpstreamBackoffSleep
|
||||
// and doubles each attempt, capped at checkUpstreamUnreachableBackoffMax.
|
||||
func unreachableRecoveryBackoff(streak int) time.Duration {
|
||||
d := checkUpstreamBackoffSleep
|
||||
for i := 1; i < streak; i++ {
|
||||
d *= 2
|
||||
if d >= checkUpstreamUnreachableBackoffMax {
|
||||
return checkUpstreamUnreachableBackoffMax
|
||||
}
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
// upstreamMonitor performs monitoring upstreams health.
|
||||
type upstreamMonitor struct {
|
||||
cfg *ctrld.Config
|
||||
|
||||
+237
-26
@@ -3,14 +3,18 @@ package cli
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/rs/zerolog"
|
||||
"tailscale.com/net/netmon"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
var vpnDNSSettlingEnabled = runtime.GOOS == "windows"
|
||||
|
||||
// vpnDNSExemption represents a VPN DNS server that needs pf/WFP exemption,
|
||||
// including the interface it was discovered on. The interface is used on macOS
|
||||
// to create interface-scoped pf exemptions that allow the VPN's local DNS
|
||||
@@ -38,6 +42,24 @@ type vpnDNSManager struct {
|
||||
// as additional nameservers for queries that match split-DNS rules
|
||||
// (from ctrld config, AD domain, or VPN suffix config).
|
||||
domainlessServers []string
|
||||
// appliedExemptions advances only after the platform PF/WFP callback succeeds.
|
||||
// Keeping it separate from discovered configs makes failed rule updates retryable.
|
||||
appliedExemptions []vpnDNSExemption
|
||||
// retainedAfterEmptyDiscovery means Windows reported an empty VPN DNS
|
||||
// snapshot once while previous VPN DNS state existed. We keep that last-known
|
||||
// state for one guarded refresh cycle because Windows can briefly report an
|
||||
// intermediate empty adapter/DNS state after sleep/wake or reconnect.
|
||||
retainedAfterEmptyDiscovery bool
|
||||
// discoverVPNDNS is injected for tests so Refresh does not depend on the
|
||||
// runner host's real VPN/virtual adapter state.
|
||||
discoverVPNDNS func(context.Context) []ctrld.VPNDNSConfig
|
||||
// refreshStateMu keeps noisy network-change storms from running overlapping
|
||||
// full VPN DNS refreshes and retains one trailing refresh when an event arrives
|
||||
// during discovery so the newest OS state is not lost.
|
||||
refreshStateMu sync.Mutex
|
||||
refreshRunning bool
|
||||
refreshPending bool
|
||||
discoveryMu sync.Mutex
|
||||
// Called when VPN DNS server list changes, to update intercept exemptions.
|
||||
onServersChanged vpnDNSExemptFunc
|
||||
}
|
||||
@@ -48,17 +70,52 @@ type vpnDNSManager struct {
|
||||
func newVPNDNSManager(exemptFunc vpnDNSExemptFunc) *vpnDNSManager {
|
||||
return &vpnDNSManager{
|
||||
routes: make(map[string][]string),
|
||||
discoverVPNDNS: ctrld.DiscoverVPNDNS,
|
||||
onServersChanged: exemptFunc,
|
||||
}
|
||||
}
|
||||
|
||||
// Refresh re-discovers VPN DNS configs from the OS.
|
||||
// Called on network change events.
|
||||
// Called on network change events. Overlapping calls are coalesced into one
|
||||
// trailing refresh so a newer OS snapshot is never silently discarded.
|
||||
func (m *vpnDNSManager) Refresh(guardAgainstNoNameservers bool) {
|
||||
m.refreshStateMu.Lock()
|
||||
if m.refreshRunning {
|
||||
m.refreshPending = true
|
||||
m.refreshStateMu.Unlock()
|
||||
mainLog.Load().Debug().Msg("VPN DNS refresh already running, coalescing trailing refresh")
|
||||
return
|
||||
}
|
||||
m.refreshRunning = true
|
||||
m.refreshStateMu.Unlock()
|
||||
|
||||
for {
|
||||
m.refreshOnce(guardAgainstNoNameservers)
|
||||
|
||||
m.refreshStateMu.Lock()
|
||||
if m.refreshPending {
|
||||
m.refreshPending = false
|
||||
m.refreshStateMu.Unlock()
|
||||
guardAgainstNoNameservers = true
|
||||
continue
|
||||
}
|
||||
m.refreshRunning = false
|
||||
m.refreshStateMu.Unlock()
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
func (m *vpnDNSManager) refreshOnce(guardAgainstNoNameservers bool) {
|
||||
logger := mainLog.Load()
|
||||
m.discoveryMu.Lock()
|
||||
defer m.discoveryMu.Unlock()
|
||||
|
||||
logger.Debug().Msg("Refreshing VPN DNS configurations")
|
||||
configs := ctrld.DiscoverVPNDNS(context.Background())
|
||||
discoverVPNDNS := m.discoverVPNDNS
|
||||
if discoverVPNDNS == nil {
|
||||
discoverVPNDNS = ctrld.DiscoverVPNDNS
|
||||
}
|
||||
configs := discoverVPNDNS(context.Background())
|
||||
|
||||
// Detect exit mode: if the default route goes through a VPN DNS interface,
|
||||
// the VPN is routing ALL traffic (exit node / full tunnel). This is more
|
||||
@@ -78,6 +135,31 @@ func (m *vpnDNSManager) Refresh(guardAgainstNoNameservers bool) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if vpnDNSSettlingEnabled && len(configs) == 0 && guardAgainstNoNameservers && m.hasVPNDNSStateLocked() {
|
||||
if !m.retainedAfterEmptyDiscovery {
|
||||
exemptions := m.currentExemptionsLocked()
|
||||
m.retainedAfterEmptyDiscovery = true
|
||||
logger.Debug().Msgf(
|
||||
"VPN DNS discovery empty; retaining last-known VPN DNS state for one guarded refresh (%d domainless servers, %d exemptions)",
|
||||
len(m.domainlessServers), len(exemptions))
|
||||
if m.onServersChanged != nil {
|
||||
if err := m.onServersChanged(exemptions); err != nil {
|
||||
logger.Error().Err(err).Msg("Failed to re-apply retained VPN DNS exemptions")
|
||||
} else {
|
||||
m.appliedExemptions = append([]vpnDNSExemption(nil), exemptions...)
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
logger.Debug().Msgf(
|
||||
"VPN DNS discovery still empty on next guarded refresh; clearing retained VPN DNS state (%d domainless servers)",
|
||||
len(m.domainlessServers))
|
||||
}
|
||||
|
||||
// Any discovery path that does not return with retained state clears the
|
||||
// settling marker: non-empty discovery replaces old servers immediately, and
|
||||
// an unguarded/second empty discovery clears stale state below.
|
||||
m.retainedAfterEmptyDiscovery = false
|
||||
m.configs = configs
|
||||
m.routes = make(map[string][]string)
|
||||
|
||||
@@ -141,14 +223,160 @@ func (m *vpnDNSManager) Refresh(guardAgainstNoNameservers bool) {
|
||||
logger.Debug().Msgf("VPN DNS refresh completed: %d configs, %d routes, %d domainless servers, %d unique exemptions",
|
||||
len(m.configs), len(m.routes), len(m.domainlessServers), len(exemptions))
|
||||
|
||||
// Update intercept rules to permit VPN DNS traffic.
|
||||
// Always call onServersChanged — including when exemptions is empty — so that
|
||||
// stale exemptions from a previous VPN session get cleared on disconnect.
|
||||
if m.onServersChanged != nil {
|
||||
if err := m.onServersChanged(exemptions); err != nil {
|
||||
logger.Error().Err(err).Msg("Failed to update intercept exemptions for VPN DNS servers")
|
||||
// Update intercept rules only when desired exemptions differ from the last
|
||||
// successfully applied set. Failed PF/WFP callbacks remain retryable on the
|
||||
// next refresh even when discovery returns the same VPN DNS state.
|
||||
m.updateInterceptExemptionsIfChanged(logger, exemptions, "VPN DNS")
|
||||
}
|
||||
|
||||
func (m *vpnDNSManager) updateInterceptExemptionsIfChanged(logger *zerolog.Logger, desired []vpnDNSExemption, reason string) {
|
||||
if m.onServersChanged == nil {
|
||||
return
|
||||
}
|
||||
if vpnDNSExemptionsEqual(m.appliedExemptions, desired) {
|
||||
logger.Debug().Msgf("VPN DNS exemptions unchanged after %s refresh; skipping intercept rule update", reason)
|
||||
return
|
||||
}
|
||||
if err := m.onServersChanged(desired); err != nil {
|
||||
logger.Error().Err(err).Msg("Failed to update intercept exemptions for VPN DNS servers")
|
||||
return
|
||||
}
|
||||
m.appliedExemptions = append([]vpnDNSExemption(nil), desired...)
|
||||
}
|
||||
|
||||
// RefreshRoutesOnly re-discovers VPN DNS configs and updates ctrld's
|
||||
// in-memory split-DNS routes. It applies intercept exemptions only when that set
|
||||
// changes, while holding the shared discovery lane so a concurrent full refresh
|
||||
// cannot commit a newer snapshot and then be overwritten by this one.
|
||||
func (m *vpnDNSManager) RefreshRoutesOnly() (routes, domainlessServers, exemptions int) {
|
||||
logger := mainLog.Load()
|
||||
|
||||
m.discoveryMu.Lock()
|
||||
defer m.discoveryMu.Unlock()
|
||||
|
||||
logger.Debug().Msg("Refreshing VPN DNS route state only")
|
||||
discoverVPNDNS := m.discoverVPNDNS
|
||||
if discoverVPNDNS == nil {
|
||||
discoverVPNDNS = ctrld.DiscoverVPNDNS
|
||||
}
|
||||
configs := discoverVPNDNS(context.Background())
|
||||
|
||||
if dri, err := netmon.DefaultRouteInterface(); err == nil && dri != "" {
|
||||
for i := range configs {
|
||||
if configs[i].InterfaceName == dri {
|
||||
configs[i].IsExitMode = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
m.retainedAfterEmptyDiscovery = false
|
||||
m.configs = configs
|
||||
m.routes = make(map[string][]string)
|
||||
|
||||
for _, config := range configs {
|
||||
for _, domain := range config.Domains {
|
||||
domain = strings.TrimPrefix(domain, "~")
|
||||
domain = strings.TrimPrefix(domain, ".")
|
||||
domain = strings.ToLower(domain)
|
||||
if domain != "" {
|
||||
m.routes[domain] = append([]string{}, config.Servers...)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var domainless []string
|
||||
seenDomainless := make(map[string]bool)
|
||||
for _, config := range configs {
|
||||
if len(config.Domains) == 0 && len(config.Servers) > 0 {
|
||||
for _, server := range config.Servers {
|
||||
if !seenDomainless[server] {
|
||||
seenDomainless[server] = true
|
||||
domainless = append(domainless, server)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
m.domainlessServers = domainless
|
||||
currentExemptions := m.currentExemptionsLocked()
|
||||
|
||||
logger.Debug().Msgf("VPN DNS route-only refresh completed: %d configs, %d routes, %d domainless servers, %d exemptions",
|
||||
len(m.configs), len(m.routes), len(m.domainlessServers), len(currentExemptions))
|
||||
m.updateInterceptExemptionsIfChanged(logger, currentExemptions, "route-only VPN DNS")
|
||||
return len(m.routes), len(m.domainlessServers), len(currentExemptions)
|
||||
}
|
||||
|
||||
func (m *vpnDNSManager) markInterceptExemptionsApplied(applied []vpnDNSExemption) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if vpnDNSExemptionsEqual(m.currentExemptionsLocked(), applied) {
|
||||
m.appliedExemptions = append([]vpnDNSExemption(nil), applied...)
|
||||
}
|
||||
}
|
||||
|
||||
func (m *vpnDNSManager) interceptExemptionsPending() bool {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return !vpnDNSExemptionsEqual(m.appliedExemptions, m.currentExemptionsLocked())
|
||||
}
|
||||
|
||||
func (m *vpnDNSManager) hasVPNDNSStateLocked() bool {
|
||||
return len(m.configs) > 0 || len(m.routes) > 0 || len(m.domainlessServers) > 0
|
||||
}
|
||||
|
||||
func (m *vpnDNSManager) currentExemptionsLocked() []vpnDNSExemption {
|
||||
type key struct{ server, iface string }
|
||||
seen := make(map[key]bool)
|
||||
var exemptions []vpnDNSExemption
|
||||
for _, config := range m.configs {
|
||||
for _, server := range config.Servers {
|
||||
k := key{server, config.InterfaceName}
|
||||
if seen[k] {
|
||||
continue
|
||||
}
|
||||
seen[k] = true
|
||||
exemptions = append(exemptions, vpnDNSExemption{
|
||||
Server: server,
|
||||
Interface: config.InterfaceName,
|
||||
IsExitMode: config.IsExitMode,
|
||||
})
|
||||
}
|
||||
}
|
||||
return exemptions
|
||||
}
|
||||
|
||||
// ShouldFailClosedAfterVPNDNSTransportFailure reports whether split-rule
|
||||
// queries should fail closed instead of falling back to OS/public DNS after
|
||||
// every candidate VPN DNS server failed before returning a DNS packet. This is
|
||||
// Windows-only and only active while serving retained VPN DNS state from a
|
||||
// guarded empty discovery, which is the short window where Windows can report
|
||||
// VPN DNS before routes to those servers are usable after wake/reconnect.
|
||||
func (m *vpnDNSManager) ShouldFailClosedAfterVPNDNSTransportFailure(domain string, servers []string) bool {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
if !vpnDNSSettlingEnabled || len(servers) == 0 || !m.retainedAfterEmptyDiscovery || !m.hasVPNDNSStateLocked() {
|
||||
return false
|
||||
}
|
||||
|
||||
mainLog.Load().Debug().Msgf(
|
||||
"VPN DNS transport failed for %s while retained VPN DNS state is active; suppressing OS fallback for this query (servers=%v)",
|
||||
domain, servers)
|
||||
return true
|
||||
}
|
||||
|
||||
// VPNDNSReachable records that a VPN DNS server returned a DNS response. The
|
||||
// response may be negative (NXDOMAIN/SERVFAIL); the important signal is that
|
||||
// the VPN DNS transport is reachable again.
|
||||
func (m *vpnDNSManager) VPNDNSReachable() {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.retainedAfterEmptyDiscovery {
|
||||
mainLog.Load().Debug().Msg("VPN DNS transport recovered; clearing retained-empty-discovery state")
|
||||
}
|
||||
m.retainedAfterEmptyDiscovery = false
|
||||
}
|
||||
|
||||
// UpstreamForDomain checks if the domain matches any VPN search domain.
|
||||
@@ -208,24 +436,7 @@ func (m *vpnDNSManager) CurrentServers() []string {
|
||||
func (m *vpnDNSManager) CurrentExemptions() []vpnDNSExemption {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
type key struct{ server, iface string }
|
||||
seen := make(map[key]bool)
|
||||
var exemptions []vpnDNSExemption
|
||||
for _, config := range m.configs {
|
||||
for _, server := range config.Servers {
|
||||
k := key{server, config.InterfaceName}
|
||||
if !seen[k] {
|
||||
seen[k] = true
|
||||
exemptions = append(exemptions, vpnDNSExemption{
|
||||
Server: server,
|
||||
Interface: config.InterfaceName,
|
||||
IsExitMode: config.IsExitMode,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
return exemptions
|
||||
return m.currentExemptionsLocked()
|
||||
}
|
||||
|
||||
// Routes returns a copy of the current VPN DNS routes for debugging.
|
||||
|
||||
@@ -0,0 +1,298 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
func withVPNDNSSettlingEnabled(t *testing.T) {
|
||||
t.Helper()
|
||||
old := vpnDNSSettlingEnabled
|
||||
vpnDNSSettlingEnabled = true
|
||||
t.Cleanup(func() { vpnDNSSettlingEnabled = old })
|
||||
}
|
||||
|
||||
func TestVPNDNSRefreshCoalescesConcurrentTrailingRefresh(t *testing.T) {
|
||||
m := newVPNDNSManager(nil)
|
||||
started := make(chan struct{})
|
||||
release := make(chan struct{})
|
||||
done := make(chan struct{})
|
||||
var once sync.Once
|
||||
var calls atomic.Int32
|
||||
|
||||
m.discoverVPNDNS = func(context.Context) []ctrld.VPNDNSConfig {
|
||||
call := calls.Add(1)
|
||||
once.Do(func() { close(started) })
|
||||
<-release
|
||||
if call == 2 {
|
||||
return []ctrld.VPNDNSConfig{{
|
||||
InterfaceName: "utun-latest",
|
||||
Servers: []string{"10.0.0.2"},
|
||||
Domains: []string{"latest.internal"},
|
||||
}}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
go func() {
|
||||
defer close(done)
|
||||
m.Refresh(true)
|
||||
}()
|
||||
|
||||
<-started
|
||||
m.Refresh(true)
|
||||
close(release)
|
||||
<-done
|
||||
|
||||
if calls.Load() != 2 {
|
||||
t.Fatalf("expected one active and one trailing discovery call, got %d", calls.Load())
|
||||
}
|
||||
if got := m.Routes()["latest.internal"]; len(got) != 1 || got[0] != "10.0.0.2" {
|
||||
t.Fatalf("trailing refresh did not publish latest OS snapshot: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVPNDNSRefreshRetainsStateForOneGuardedEmptyDiscovery(t *testing.T) {
|
||||
withVPNDNSSettlingEnabled(t)
|
||||
var gotExemptions []vpnDNSExemption
|
||||
m := newVPNDNSManager(func(exemptions []vpnDNSExemption) error {
|
||||
gotExemptions = exemptions
|
||||
return nil
|
||||
})
|
||||
m.discoverVPNDNS = func(context.Context) []ctrld.VPNDNSConfig { return nil }
|
||||
m.configs = []ctrld.VPNDNSConfig{{
|
||||
InterfaceName: "Ethernet 6",
|
||||
Servers: []string{"10.25.37.21", "10.25.37.22"},
|
||||
}}
|
||||
m.domainlessServers = []string{"10.25.37.21", "10.25.37.22"}
|
||||
|
||||
m.Refresh(true)
|
||||
|
||||
if got := m.DomainlessServers(); len(got) != 2 {
|
||||
t.Fatalf("expected retained domainless servers, got %v", got)
|
||||
}
|
||||
if len(gotExemptions) != 2 {
|
||||
t.Fatalf("expected retained exemptions to be re-applied, got %v", gotExemptions)
|
||||
}
|
||||
if !m.retainedAfterEmptyDiscovery {
|
||||
t.Fatal("expected empty discovery retention to be marked")
|
||||
}
|
||||
}
|
||||
|
||||
func TestVPNDNSRefreshClearsOnSecondGuardedEmptyDiscovery(t *testing.T) {
|
||||
withVPNDNSSettlingEnabled(t)
|
||||
var gotExemptions []vpnDNSExemption
|
||||
updates := 0
|
||||
m := newVPNDNSManager(func(exemptions []vpnDNSExemption) error {
|
||||
updates++
|
||||
gotExemptions = exemptions
|
||||
return nil
|
||||
})
|
||||
m.discoverVPNDNS = func(context.Context) []ctrld.VPNDNSConfig { return nil }
|
||||
m.configs = []ctrld.VPNDNSConfig{{
|
||||
InterfaceName: "Ethernet 6",
|
||||
Servers: []string{"10.25.37.21"},
|
||||
}}
|
||||
m.domainlessServers = []string{"10.25.37.21"}
|
||||
m.appliedExemptions = []vpnDNSExemption{{Server: "10.25.37.21", Interface: "Ethernet 6"}}
|
||||
m.retainedAfterEmptyDiscovery = true
|
||||
|
||||
m.Refresh(true)
|
||||
|
||||
if got := m.DomainlessServers(); len(got) != 0 {
|
||||
t.Fatalf("expected domainless servers to be cleared on second empty discovery, got %v", got)
|
||||
}
|
||||
if updates != 1 || len(gotExemptions) != 0 {
|
||||
t.Fatalf("expected one empty exemption update after clearing stale state, calls=%d exemptions=%v", updates, gotExemptions)
|
||||
}
|
||||
if m.retainedAfterEmptyDiscovery {
|
||||
t.Fatal("expected retained empty-discovery marker to be cleared with stale state")
|
||||
}
|
||||
}
|
||||
|
||||
func TestVPNDNSRefreshSkipsUnchangedInterceptExemptions(t *testing.T) {
|
||||
var updates [][]vpnDNSExemption
|
||||
m := newVPNDNSManager(func(exemptions []vpnDNSExemption) error {
|
||||
updates = append(updates, append([]vpnDNSExemption{}, exemptions...))
|
||||
return nil
|
||||
})
|
||||
m.discoverVPNDNS = func(context.Context) []ctrld.VPNDNSConfig {
|
||||
return []ctrld.VPNDNSConfig{{
|
||||
InterfaceName: "utun-test",
|
||||
Servers: []string{"10.102.26.10"},
|
||||
Domains: []string{"example.internal"},
|
||||
}}
|
||||
}
|
||||
|
||||
m.Refresh(true)
|
||||
m.Refresh(true)
|
||||
|
||||
if len(updates) != 1 {
|
||||
t.Fatalf("expected exactly one intercept exemption update for unchanged VPN DNS state, got %d", len(updates))
|
||||
}
|
||||
if len(updates[0]) != 1 || updates[0][0].Server != "10.102.26.10" || updates[0][0].Interface != "utun-test" {
|
||||
t.Fatalf("unexpected exemption update: %+v", updates[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestVPNDNSRefreshRetriesFailedInterceptExemptionUpdate(t *testing.T) {
|
||||
attempts := 0
|
||||
m := newVPNDNSManager(func([]vpnDNSExemption) error {
|
||||
attempts++
|
||||
if attempts == 1 {
|
||||
return errors.New("pf update failed")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
m.discoverVPNDNS = func(context.Context) []ctrld.VPNDNSConfig {
|
||||
return []ctrld.VPNDNSConfig{{
|
||||
InterfaceName: "utun-test",
|
||||
Servers: []string{"10.102.26.10"},
|
||||
Domains: []string{"internal.test"},
|
||||
}}
|
||||
}
|
||||
|
||||
m.Refresh(true)
|
||||
if !m.interceptExemptionsPending() {
|
||||
t.Fatal("failed intercept exemption update was not retained for retry")
|
||||
}
|
||||
m.Refresh(true)
|
||||
if m.interceptExemptionsPending() {
|
||||
t.Fatal("successful intercept exemption retry did not advance applied state")
|
||||
}
|
||||
m.Refresh(true)
|
||||
|
||||
if attempts != 2 {
|
||||
t.Fatalf("intercept exemption update attempts = %d, want failed attempt plus one retry", attempts)
|
||||
}
|
||||
if len(m.appliedExemptions) != 1 || m.appliedExemptions[0].Server != "10.102.26.10" {
|
||||
t.Fatalf("applied exemptions = %+v, want successful retry state", m.appliedExemptions)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVPNDNSMarkAppliedExemptionsRejectsStaleSnapshot(t *testing.T) {
|
||||
m := newVPNDNSManager(nil)
|
||||
m.configs = []ctrld.VPNDNSConfig{{InterfaceName: "utun-new", Servers: []string{"10.0.0.2"}}}
|
||||
|
||||
m.markInterceptExemptionsApplied([]vpnDNSExemption{{Server: "10.0.0.1", Interface: "utun-old"}})
|
||||
if !m.interceptExemptionsPending() {
|
||||
t.Fatal("stale PF snapshot incorrectly advanced applied exemptions")
|
||||
}
|
||||
|
||||
m.markInterceptExemptionsApplied([]vpnDNSExemption{{Server: "10.0.0.2", Interface: "utun-new"}})
|
||||
if m.interceptExemptionsPending() {
|
||||
t.Fatal("current PF snapshot did not advance applied exemptions")
|
||||
}
|
||||
}
|
||||
|
||||
func TestVPNDNSTransportFailureSuppressesFallbackOnlyWhileRetainingState(t *testing.T) {
|
||||
withVPNDNSSettlingEnabled(t)
|
||||
m := newVPNDNSManager(nil)
|
||||
m.domainlessServers = []string{"10.25.37.21"}
|
||||
|
||||
if m.ShouldFailClosedAfterVPNDNSTransportFailure("splunk.aws.arena.net.", []string{"10.25.37.21"}) {
|
||||
t.Fatal("did not expect transport failure to suppress OS fallback outside retained empty-discovery state")
|
||||
}
|
||||
|
||||
m.retainedAfterEmptyDiscovery = true
|
||||
if !m.ShouldFailClosedAfterVPNDNSTransportFailure("splunk.aws.arena.net.", []string{"10.25.37.21"}) {
|
||||
t.Fatal("expected transport failure to suppress OS fallback while retained state is active")
|
||||
}
|
||||
|
||||
m.VPNDNSReachable()
|
||||
if m.retainedAfterEmptyDiscovery {
|
||||
t.Fatal("expected reachable DNS response to clear retained empty-discovery state")
|
||||
}
|
||||
}
|
||||
|
||||
func TestVPNDNSFullAndRouteOnlyDiscoveryAreSerialized(t *testing.T) {
|
||||
var updateMu sync.Mutex
|
||||
var exemptionUpdates []string
|
||||
m := newVPNDNSManager(func(exemptions []vpnDNSExemption) error {
|
||||
updateMu.Lock()
|
||||
defer updateMu.Unlock()
|
||||
if len(exemptions) == 0 {
|
||||
exemptionUpdates = append(exemptionUpdates, "")
|
||||
} else {
|
||||
exemptionUpdates = append(exemptionUpdates, exemptions[0].Server)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
firstStarted := make(chan struct{})
|
||||
releaseFirst := make(chan struct{})
|
||||
secondStarted := make(chan struct{})
|
||||
var calls atomic.Int32
|
||||
|
||||
m.discoverVPNDNS = func(context.Context) []ctrld.VPNDNSConfig {
|
||||
switch calls.Add(1) {
|
||||
case 1:
|
||||
close(firstStarted)
|
||||
<-releaseFirst
|
||||
return []ctrld.VPNDNSConfig{{
|
||||
InterfaceName: "utun-old",
|
||||
Servers: []string{"10.0.0.1"},
|
||||
Domains: []string{"old.internal"},
|
||||
}}
|
||||
case 2:
|
||||
close(secondStarted)
|
||||
return []ctrld.VPNDNSConfig{{
|
||||
InterfaceName: "utun-new",
|
||||
Servers: []string{"10.0.0.2"},
|
||||
Domains: []string{"new.internal"},
|
||||
}}
|
||||
default:
|
||||
t.Fatalf("unexpected discovery call %d", calls.Load())
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
routesDone := make(chan struct{})
|
||||
go func() {
|
||||
defer close(routesDone)
|
||||
m.RefreshRoutesOnly()
|
||||
}()
|
||||
<-firstStarted
|
||||
|
||||
fullDone := make(chan struct{})
|
||||
go func() {
|
||||
defer close(fullDone)
|
||||
m.Refresh(false)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-secondStarted:
|
||||
t.Fatal("full and route-only VPN DNS discovery overlapped")
|
||||
case <-time.After(50 * time.Millisecond):
|
||||
}
|
||||
close(releaseFirst)
|
||||
|
||||
select {
|
||||
case <-routesDone:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("route-only refresh did not finish")
|
||||
}
|
||||
select {
|
||||
case <-fullDone:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("full refresh did not finish")
|
||||
}
|
||||
|
||||
routes := m.Routes()
|
||||
if _, ok := routes["old.internal"]; ok {
|
||||
t.Fatalf("older route-only snapshot overwrote newer full refresh: %v", routes)
|
||||
}
|
||||
if got := routes["new.internal"]; len(got) != 1 || got[0] != "10.0.0.2" {
|
||||
t.Fatalf("final VPN DNS routes = %v, want new.internal -> 10.0.0.2", routes)
|
||||
}
|
||||
updateMu.Lock()
|
||||
defer updateMu.Unlock()
|
||||
if len(exemptionUpdates) != 2 || exemptionUpdates[0] != "10.0.0.1" || exemptionUpdates[1] != "10.0.0.2" {
|
||||
t.Fatalf("serialized exemption updates = %v, want old then new", exemptionUpdates)
|
||||
}
|
||||
}
|
||||
@@ -241,6 +241,8 @@ type ServiceConfig struct {
|
||||
ForceRefetchWaitTime *int `mapstructure:"force_refetch_wait_time" toml:"force_refetch_wait_time,omitempty"`
|
||||
LeakOnUpstreamFailure *bool `mapstructure:"leak_on_upstream_failure" toml:"leak_on_upstream_failure,omitempty"`
|
||||
InterceptMode string `mapstructure:"intercept_mode" toml:"intercept_mode,omitempty" validate:"omitempty,oneof=off dns hard"`
|
||||
NRPTRecoveryMaxAttempts *int `mapstructure:"nrpt_recovery_max_attempts" toml:"nrpt_recovery_max_attempts,omitempty" validate:"omitempty,gte=0"`
|
||||
NRPTRecoveryCooldown *time.Duration `mapstructure:"nrpt_recovery_cooldown" toml:"nrpt_recovery_cooldown,omitempty"`
|
||||
Daemon bool `mapstructure:"-" toml:"-"`
|
||||
AllocateIP bool `mapstructure:"-" toml:"-"`
|
||||
}
|
||||
@@ -640,6 +642,7 @@ func (uc *UpstreamConfig) newDOHTransport(addrs []string) *http.Transport {
|
||||
transport.TLSClientConfig = &tls.Config{
|
||||
RootCAs: uc.certPool,
|
||||
ClientSessionCache: tls.NewLRUClientSessionCache(0),
|
||||
MinVersion: tls.VersionTLS12,
|
||||
}
|
||||
|
||||
// Prevent bad tcp connection hanging the requests for too long.
|
||||
|
||||
+29
-7
@@ -18,7 +18,7 @@ func (uc *UpstreamConfig) newDOH3Transport(addrs []string) http.RoundTripper {
|
||||
return nil
|
||||
}
|
||||
rt := &http3.Transport{}
|
||||
rt.TLSClientConfig = &tls.Config{RootCAs: uc.certPool}
|
||||
rt.TLSClientConfig = &tls.Config{RootCAs: uc.certPool, MinVersion: tls.VersionTLS12}
|
||||
rt.Dial = func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
||||
_, port, _ := net.SplitHostPort(addr)
|
||||
// if we have a bootstrap ip set, use it to avoid DNS lookup
|
||||
@@ -77,7 +77,17 @@ type parallelDialerResult struct {
|
||||
err error
|
||||
}
|
||||
|
||||
type quicParallelDialer struct{}
|
||||
// quicParallelDialer races DialEarly across a list of remote addresses and
|
||||
// returns the first successful connection. When transport is non-nil, all
|
||||
// dials share that transport's UDP socket, which removes both the per-dial
|
||||
// socket allocation and the winner-path socket leak that an owner-of-the-conn
|
||||
// receiver cannot clean up. When transport is nil, the dialer falls back to a
|
||||
// fresh UDP socket per attempt (compat path used where no shared transport is
|
||||
// available yet); the loser paths close their sockets, and the winner path's
|
||||
// socket is owned by quic.DialEarly's internal transport.
|
||||
type quicParallelDialer struct {
|
||||
transport *quic.Transport
|
||||
}
|
||||
|
||||
// Dial performs parallel dialing to the given address list.
|
||||
func (d *quicParallelDialer) Dial(ctx context.Context, addrs []string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
||||
@@ -105,12 +115,24 @@ func (d *quicParallelDialer) Dial(ctx context.Context, addrs []string, tlsCfg *t
|
||||
ch <- ¶llelDialerResult{conn: nil, err: err}
|
||||
return
|
||||
}
|
||||
udpConn, err := net.ListenUDP("udp", nil)
|
||||
if err != nil {
|
||||
ch <- ¶llelDialerResult{conn: nil, err: err}
|
||||
return
|
||||
var (
|
||||
conn *quic.Conn
|
||||
udpConn *net.UDPConn
|
||||
)
|
||||
if d.transport != nil {
|
||||
conn, err = d.transport.DialEarly(ctx, remoteAddr, tlsCfg, cfg)
|
||||
} else {
|
||||
udpConn, err = net.ListenUDP("udp", nil)
|
||||
if err != nil {
|
||||
ch <- ¶llelDialerResult{conn: nil, err: err}
|
||||
return
|
||||
}
|
||||
conn, err = quic.DialEarly(ctx, udpConn, remoteAddr, tlsCfg, cfg)
|
||||
if err != nil {
|
||||
udpConn.Close()
|
||||
udpConn = nil
|
||||
}
|
||||
}
|
||||
conn, err := quic.DialEarly(ctx, udpConn, remoteAddr, tlsCfg, cfg)
|
||||
select {
|
||||
case ch <- ¶llelDialerResult{conn: conn, err: err}:
|
||||
case <-done:
|
||||
|
||||
@@ -14,6 +14,12 @@ func SetCacheReply(answer, msg *dns.Msg, code int) {
|
||||
// See https://datatracker.ietf.org/doc/html/rfc7873#section-4
|
||||
sCookie.Cookie = cCookie.Cookie[:16] + sCookie.Cookie[16:]
|
||||
}
|
||||
// NOTE: the answer's EDNS Client Subnet (ECS) is intentionally left as the
|
||||
// upstream returned it. Correctness across clients is guaranteed by
|
||||
// partitioning the cache and singleflight keys by ECS (see
|
||||
// dnscache.CanonicalECS), so a cache hit only ever serves an answer that was
|
||||
// resolved for the requester's own subnet. Rewriting the ECS option here
|
||||
// without re-scoping the Answer records would violate RFC 7871 §7.3.
|
||||
}
|
||||
|
||||
// getEdns0Cookie returns Edns0 cookie from *dns.OPT if present.
|
||||
|
||||
+57
@@ -0,0 +1,57 @@
|
||||
package ctrld
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
// Test_SetCacheReply_DoesNotRewriteECS documents the post-#564 contract: cross-client
|
||||
// correctness is guaranteed by partitioning the cache/singleflight keys by ECS
|
||||
// (dnscache.CanonicalECS), NOT by rewriting the cached answer's ECS option. Rewriting the
|
||||
// ECS metadata while leaving the Answer records scoped to another subnet would violate
|
||||
// RFC 7871 §7.3 and make forwarders accept a wrong-subnet answer. SetCacheReply must
|
||||
// therefore leave the answer's ECS untouched.
|
||||
func Test_SetCacheReply_DoesNotRewriteECS(t *testing.T) {
|
||||
answer := new(dns.Msg)
|
||||
answer.SetQuestion(dns.Fqdn("controld.com"), dns.TypeA)
|
||||
answer.SetEdns0(4096, false)
|
||||
cachedSubnet := &dns.EDNS0_SUBNET{
|
||||
Code: dns.EDNS0SUBNET,
|
||||
Family: 2,
|
||||
SourceNetmask: 64,
|
||||
SourceScope: 64,
|
||||
Address: net.ParseIP("2001:db8:1::"),
|
||||
}
|
||||
answer.IsEdns0().Option = append(answer.IsEdns0().Option, cachedSubnet)
|
||||
|
||||
req := new(dns.Msg)
|
||||
req.SetQuestion(dns.Fqdn("controld.com"), dns.TypeA)
|
||||
req.SetEdns0(4096, true)
|
||||
req.IsEdns0().Option = append(req.IsEdns0().Option, &dns.EDNS0_SUBNET{
|
||||
Code: dns.EDNS0SUBNET,
|
||||
Family: 2,
|
||||
SourceNetmask: 64,
|
||||
Address: net.ParseIP("2001:db8:2::"),
|
||||
})
|
||||
|
||||
SetCacheReply(answer, req, dns.RcodeSuccess)
|
||||
|
||||
var got *dns.EDNS0_SUBNET
|
||||
for _, o := range answer.IsEdns0().Option {
|
||||
if e, ok := o.(*dns.EDNS0_SUBNET); ok {
|
||||
got = e
|
||||
break
|
||||
}
|
||||
}
|
||||
if got == nil {
|
||||
t.Fatal("SetCacheReply dropped the answer's ECS option")
|
||||
}
|
||||
if want := net.ParseIP("2001:db8:1::"); !got.Address.Equal(want) {
|
||||
t.Fatalf("SetCacheReply rewrote the answer ECS to the requester's subnet: got %v, want %v (unchanged)", got.Address, want)
|
||||
}
|
||||
if got.SourceScope != 64 {
|
||||
t.Fatalf("SetCacheReply altered the answer ECS scope: got %d, want 64 (unchanged)", got.SourceScope)
|
||||
}
|
||||
}
|
||||
+4
-3
@@ -1,4 +1,4 @@
|
||||
# Using Debian bullseye for building regular image.
|
||||
# Using Debian bookworm for building regular image.
|
||||
# Using scratch image for minimal image size.
|
||||
# The final image has:
|
||||
#
|
||||
@@ -8,11 +8,12 @@
|
||||
# - Non-cgo ctrld binary.
|
||||
#
|
||||
# CI_COMMIT_TAG is used to set the version of ctrld binary.
|
||||
FROM golang:1.20-bullseye as base
|
||||
FROM golang:1.25-bookworm AS base
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
RUN apt-get update && apt-get install -y upx-ucl
|
||||
RUN echo "deb http://deb.debian.org/debian bookworm-backports main" | tee /etc/apt/sources.list.d/backports.list
|
||||
RUN apt update && apt install -t bookworm-backports upx-ucl
|
||||
|
||||
COPY . .
|
||||
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
# Using Debian bullseye for building regular image.
|
||||
# Using Debian bookworm for building regular image.
|
||||
# Using scratch image for minimal image size.
|
||||
# The final image has:
|
||||
#
|
||||
@@ -8,11 +8,12 @@
|
||||
# - Non-cgo ctrld binary.
|
||||
#
|
||||
# CI_COMMIT_TAG is used to set the version of ctrld binary.
|
||||
FROM golang:bullseye as base
|
||||
FROM golang:1.25-bookworm AS base
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
RUN apt-get update && apt-get install -y upx-ucl
|
||||
RUN echo "deb http://deb.debian.org/debian bookworm-backports main" | tee /etc/apt/sources.list.d/backports.list
|
||||
RUN apt update && apt install -t bookworm-backports upx-ucl
|
||||
|
||||
COPY . .
|
||||
|
||||
|
||||
@@ -295,6 +295,22 @@ If a remote upstream fails to resolve a query or is unreachable, `ctrld` will fo
|
||||
- Required: no
|
||||
- Default: true on Windows, MacOS and non-router Linux.
|
||||
|
||||
### nrpt_recovery_max_attempts
|
||||
Windows DNS intercept mode uses NRPT health probes and recovery when Windows stops routing queries to the local `ctrld` listener. This limits how many consecutive recovery flows can run before `ctrld` enters a cooldown and stops making policy/Dnscache changes.
|
||||
|
||||
Set to `0` to disable this circuit breaker and keep retrying indefinitely.
|
||||
|
||||
- Type: integer
|
||||
- Required: no
|
||||
- Default: 0 (unlimited, current behavior)
|
||||
|
||||
### nrpt_recovery_cooldown
|
||||
Cooldown duration after `nrpt_recovery_max_attempts` consecutive Windows NRPT recovery flows. During cooldown, `ctrld` logs the suppressed recovery and avoids additional `RefreshPolicyEx`, Dnscache `paramchange`, and DNS cache flush calls.
|
||||
|
||||
- Type: time duration string
|
||||
- Required: no
|
||||
- Default: 30m
|
||||
|
||||
## Upstream
|
||||
The `[upstream]` section specifies the DNS upstream servers that `ctrld` will forward DNS requests to.
|
||||
|
||||
|
||||
+44
-23
@@ -67,28 +67,27 @@ Separating them into modes means most users get `dns` mode (safe, can never brea
|
||||
|
||||
#### Startup Sequence (dns mode)
|
||||
|
||||
1. Creates NRPT catch-all registry rule (`.` → `127.0.0.1`) under `HKLM\...\DnsPolicyConfig\CtrldCatchAll`
|
||||
2. Triggers Group Policy refresh via `RefreshPolicyEx` (userenv.dll) so DNS Client loads NRPT immediately
|
||||
3. Flushes DNS cache to clear stale entries
|
||||
4. **Activates loopback WFP protect** — adds 4 permit filters (IPv4/IPv6 × UDP/TCP) for DNS to localhost with `FWPM_FILTER_FLAG_CLEAR_ACTION_RIGHT`. These prevent third-party WFP block filters from blocking the NRPT → `127.0.0.1` path (see [Loopback WFP Protect](#loopback-wfp-protect) below). Non-fatal if this fails.
|
||||
5. Starts NRPT health monitor (30s periodic check)
|
||||
6. Launches async NRPT probe-and-heal to verify NRPT is actually routing queries
|
||||
1. Checks for a non-ctrld GP child whose only namespace is `.` and whose only nameserver is ctrld's actual listener IP.
|
||||
2. When that candidate exists, preserves adapter DNS, sends a DNS Client probe before any NRPT mutation, and re-reads the same GP child. A matching before/after rule plus a received probe enters **GP-managed mode**; ctrld does not write NRPT, call `RefreshPolicyEx`/`paramchange`, or flush DNS for policy activation.
|
||||
3. Without a still-matching GP candidate, creates the normal ctrld-owned catch-all, signals DNS Client, and flushes stale cache entries.
|
||||
4. **Activates loopback WFP protect** — adds 4 permit filters (IPv4/IPv6 × UDP/TCP) for DNS to localhost with `FWPM_FILTER_FLAG_CLEAR_ACTION_RIGHT`. These prevent third-party WFP block filters from blocking the NRPT → listener path (see [Loopback WFP Protect](#loopback-wfp-protect) below). Non-fatal if this fails.
|
||||
5. Starts the 30-second ownership-aware NRPT health monitor.
|
||||
6. Re-verifies an initially ineffective GP candidate synchronously after WFP setup; ctrld-owned NRPT uses the asynchronous probe-and-heal sequence.
|
||||
|
||||
#### Startup Sequence (hard mode)
|
||||
|
||||
1. Creates NRPT catch-all rule + GP refresh + DNS flush (same as dns mode)
|
||||
2. Opens WFP engine with `RPC_C_AUTHN_DEFAULT` (0xFFFFFFFF)
|
||||
3. Cleans up any stale sublayer from a previous unclean shutdown
|
||||
4. Creates sublayer with maximum weight (0xFFFF)
|
||||
5. Adds **permit** filters (weight 10) for DNS to localhost (`127.0.0.1`/`::1` port 53)
|
||||
6. Adds **permit** filters (weight 10) for DNS to RFC1918 + CGNAT subnets (10/8, 172.16/12, 192.168/16, 100.64/10)
|
||||
7. Adds **block** filters (weight 1) for all other outbound DNS (port 53 UDP+TCP)
|
||||
8. Starts NRPT health monitor (also verifies WFP sublayer in hard mode)
|
||||
9. Launches async NRPT probe-and-heal
|
||||
1. Establishes NRPT routing using the same GP-managed adoption or ctrld-owned fallback sequence as `dns` mode.
|
||||
2. Opens WFP engine with `RPC_C_AUTHN_DEFAULT` (0xFFFFFFFF).
|
||||
3. Cleans up any stale sublayer from a previous unclean shutdown.
|
||||
4. Creates sublayer with maximum weight (0xFFFF).
|
||||
5. Adds **permit** filters (weight 10) for DNS to localhost (`127.0.0.1`/`::1` port 53).
|
||||
6. Adds **permit** filters (weight 10) for DNS to RFC1918 + CGNAT subnets (10/8, 172.16/12, 192.168/16, 100.64/10).
|
||||
7. Adds **block** filters (weight 1) for all other outbound DNS (port 53 UDP+TCP).
|
||||
8. Starts the NRPT/WFP health monitor.
|
||||
|
||||
**Atomic guarantee:** NRPT must succeed before WFP starts. If NRPT fails, WFP is not attempted. If WFP fails, NRPT is rolled back. This prevents DNS blackholes where WFP blocks everything but nothing routes to ctrld.
|
||||
**Atomic guarantee:** NRPT routing must exist before WFP starts. If WFP setup fails, ctrld rolls back only a rule it owns. A GP-managed child is never deleted, rewritten, or replaced with interface DNS merely because ctrld's WFP setup failed.
|
||||
|
||||
On shutdown: stops health monitor, removes NRPT rule, flushes DNS, then (hard mode only) removes all WFP filters and closes engine.
|
||||
On shutdown, ctrld stops its monitor and WFP session. It removes and signals only ctrld-owned NRPT state; a GP-managed catch-all remains untouched.
|
||||
|
||||
#### NRPT Details
|
||||
|
||||
@@ -103,7 +102,27 @@ The **Name Resolution Policy Table** is a Windows feature (originally for Direct
|
||||
|
||||
**Registry path**: `HKLM\SOFTWARE\Policies\Microsoft\Windows NT\DNSClient\DnsPolicyConfig\CtrldCatchAll`
|
||||
|
||||
**Group Policy refresh**: The DNS Client service only reads NRPT from registry during Group Policy processing cycles (default: every 90 minutes). ctrld calls `RefreshPolicyEx(bMachine=TRUE, dwOptions=RP_FORCE)` from `userenv.dll` to trigger an immediate refresh. Falls back to `gpupdate /target:computer /force` if the DLL call fails.
|
||||
**Group Policy refresh**: The DNS Client service only reads NRPT from registry during Group Policy processing cycles (default: every 90 minutes). ctrld calls `RefreshPolicyEx(bMachine=TRUE, dwOptions=RP_FORCE)` when activating or repairing rules it owns. While Group Policy remains the owner, ctrld does not run NRPT activation/heal signaling; the one transition that removes a ctrld fallback is signaled after the external rule has been proven.
|
||||
|
||||
#### GP-managed NRPT ownership
|
||||
|
||||
Enterprise deployments may install a computer-scoped GP child before starting ctrld with:
|
||||
|
||||
- exactly one namespace: `.`;
|
||||
- exactly one `GenericDNSServers` value; and
|
||||
- that nameserver equal to ctrld's actual loopback listener (`127.0.0.1` or the alternate loopback selected on an AD DNS server).
|
||||
|
||||
At service startup ctrld reads that candidate before the normal adapter reset, probes through Windows DNS Client while its listener is already bound, and re-reads the same child. When the rule remains present and the probe arrives, ctrld records **Group Policy** as the NRPT owner. Adapter DNS stays on the organization's resolvers, and ctrld does not create, delete, refresh, or flush NRPT policy.
|
||||
|
||||
The health monitor keeps using functional probes:
|
||||
|
||||
- matching GP rule + successful probe: observe only;
|
||||
- matching GP rule + failed probe: retry loopback WFP protection, then report the external policy as ineffective without running NRPT heal signals;
|
||||
- matching GP rule disappears: create the normal ctrld-owned fallback and verify it, unless another GP catch-all targets a different resolver;
|
||||
- GP catch-all targets another resolver: report the conflict and do not create a second ambiguous catch-all;
|
||||
- matching GP rule returns: prove it with a probe, remove only ctrld's deterministic fallback keys, and return ownership to Group Policy.
|
||||
|
||||
Deploy the GPO **before** starting or restarting ctrld if adapter DNS must remain completely untouched. Remove or unlink the GP rule before intentionally removing the ctrld service. A GP catch-all that remains pointed at loopback while no listener is running causes DNS failure by design; ctrld cannot safely delete an administrator-owned policy during uninstall.
|
||||
|
||||
#### WFP Filter Architecture
|
||||
|
||||
@@ -145,17 +164,19 @@ See: [Issue #526](https://gitlab.int.windscribe.com/controld/clients/ctrld/-/iss
|
||||
|
||||
ctrld verifies NRPT is actually working by sending a probe DNS query (`_nrpt-probe-<hex>.nrpt-probe.ctrld.test`) through Go's `net.Resolver` (which calls `GetAddrInfoW` → DNS Client → NRPT path). If ctrld receives the probe on its listener, NRPT is active.
|
||||
|
||||
**Startup probe (async, non-blocking):** After NRPT setup, an async goroutine probes with escalating remediation: (1) immediate probe, (2) GP refresh + retry, (3) DNS Client service restart + retry, (4) final retry. Only one probe sequence runs at a time.
|
||||
**Startup probes:** A matching GP candidate is probed synchronously before any NRPT mutation and re-read afterward. ctrld-owned rules keep the asynchronous activation/heal sequence: immediate probe, bounded policy signaling retries, then two-phase delete/re-add recovery. Only one probe sequence runs at a time.
|
||||
|
||||
**DNS Client restart (nuclear option):** If GP refresh alone isn't enough, ctrld restarts the `Dnscache` service to force full NRPT re-initialization. This briefly interrupts all DNS (~100ms) but only fires when NRPT is already not working.
|
||||
**Ownership boundary:** When the active owner is Group Policy, a failed probe never enters ctrld's NRPT refresh/delete/re-add sequence. ctrld may repair its narrowly scoped loopback WFP permits, but leaves the external registry child and DNS Client policy signaling to the administrator.
|
||||
|
||||
#### NRPT Health Monitor
|
||||
|
||||
A dedicated background goroutine (`nrptHealthMonitor`) runs every 30 seconds and now performs active probing:
|
||||
|
||||
1. **Registry check:** If the NRPT catch-all rule is missing from the registry, restore it + GP refresh + probe-and-heal
|
||||
2. **Active probe:** If the rule exists, send a probe query to verify it's actually routing — catches cases where the registry key is present but DNS Client hasn't loaded it
|
||||
3. **(hard mode)** Verify WFP sublayer exists; full restart on loss
|
||||
1. **Ownership check:** Distinguish a matching external GP child from ctrld's deterministic local/GP keys.
|
||||
2. **Active probe:** Verify Windows DNS Client still routes to the listener.
|
||||
3. **Transition:** If the external child disappears, activate ctrld's normal fallback. If it returns while the fallback is active, prove it before removing only ctrld's keys.
|
||||
4. **Owned recovery:** Restore/heal only when ctrld owns the NRPT rule.
|
||||
5. **(hard mode)** Verify the WFP sublayer exists and fully restart intercept state on loss.
|
||||
|
||||
This is periodic (not just network-event-driven) because VPN software can clear NRPT at any time. Additionally, `scheduleDelayedRechecks()` (called on network change events) performs immediate NRPT verification at 2s and 4s after changes.
|
||||
|
||||
|
||||
@@ -298,11 +298,22 @@ The full pf reload is VPN-safe: it reassembles from `pfctl -sr` + `pfctl -sn`
|
||||
### What about `set skip on lo0`?
|
||||
Some pf.conf files include `set skip on lo0` which tells pf to skip ALL processing on loopback. **This would break our approach** since both the `rdr on lo0` and `pass in on lo0` rules would be skipped.
|
||||
|
||||
**Mitigation:** When injecting anchor references via `ensurePFAnchorReference()`,
|
||||
we strip `lo0` from any `set skip on` directives before reloading. The watchdog
|
||||
also checks for `set skip on lo0` and triggers a restore if detected. The
|
||||
interception probe provides an additional safety net — if `set skip on lo0` gets
|
||||
re-applied by another program, the probe will fail and trigger a full reload.
|
||||
**Mitigation:** the interception probe. `probePFIntercept()` sends a real query from
|
||||
outside the `_ctrld` group and confirms the listener received the redirect, which cannot
|
||||
succeed while pf is bypassing loopback — so a skip on `lo0` shows up as a probe failure
|
||||
and triggers a full reload.
|
||||
|
||||
**Not implemented, contrary to earlier versions of this document:** ctrld does *not*
|
||||
strip `lo0` from `set skip on` directives, and the watchdog does *not* inspect skip
|
||||
state. Apple's `pfctl` offers no way to read it — `pfctl(8)` accepts `-s` nat, queue,
|
||||
rules, Anchors, states, Sources, info, References, labels, timeouts, memory, Tables,
|
||||
osfp, Interfaces, all, with no options or skip modifier — so text-based detection is not
|
||||
available on macOS.
|
||||
|
||||
Adding an explicit check is tracked as follow-up: `pfctl(8)` documents
|
||||
`-s Interfaces -v` as additionally listing which interfaces have skip rules activated,
|
||||
which is the query to build on once its output shape is confirmed on a host that has a
|
||||
skip configured.
|
||||
|
||||
## Cleanup
|
||||
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
# Provisioning failure codes
|
||||
|
||||
When ctrld hits a terminal failure during provisioning, it reports the same
|
||||
stable code on three surfaces:
|
||||
|
||||
- **Result file** — `provision_result.json` in the ctrld home directory
|
||||
(next to the persisted internal `ctrld.log`). JSON with `stage`, `code`,
|
||||
`exit_code`, `message`, and for listener failures a bounded
|
||||
`detail.attempts` list of `{addr, proto, os_error}`. Written atomically,
|
||||
removed on the next successful provisioning. Never contains provision
|
||||
tokens, resolver/device IDs, or configuration contents.
|
||||
- **Output line** — one fixed-format line on the CLI output:
|
||||
`provisioning failed: stage=<stage> code=<CODE> (exit <N>)`.
|
||||
The macOS pkg `postinstall` extracts exactly this line into the installer
|
||||
log, so MDM consoles see it without any ctrld log configuration.
|
||||
- **Exit code** — stage-scoped: bootstrap 30–39, listener 40–49,
|
||||
service 50–59. Unrelated existing contracts are unchanged
|
||||
(`ctrld status` exits 0–3; invalid deactivation pin exits 126).
|
||||
|
||||
A customer or administrator only needs to report the code (or the whole
|
||||
output line). The table below is the maintained support mapping; it must
|
||||
stay in sync with `cmd/cli/provision_result.go` and changes in the same MR.
|
||||
|
||||
## Codes
|
||||
|
||||
| Code | Stage | Exit | Failure scenario | Next action / evidence |
|
||||
|---|---|---|---|---|
|
||||
| `API_UNREACHABLE` | bootstrap | 30 | The Control D API could not be reached or answered with a retryable error (network failure, proxy interference, 5xx, timeout) and retries ran out. The service manager may retry the service later. | Check the device's network path to `api.controld.com` (DNS, proxy, firewall, captive portal). Ask for the result file's `message` and whether other TLS traffic works. |
|
||||
| `API_REJECTED` | bootstrap | 31 | The API answered and permanently rejected the configuration (4xx other than 408/429): bad or revoked token, malformed request. ctrld exits without burning service-manager restarts because retrying cannot change the answer. | Verify the provision token / org configuration in the Control D dashboard. Re-push after fixing credentials. Evidence: HTTP status in the result file `message`. |
|
||||
| `API_DEVICE_INVALID` | bootstrap | 32 | The API reports the device/resolver no longer exists (error code 40402). ctrld self-uninstalls its service because the identity is gone server-side. | Confirm the device was deleted or re-provisioned in the dashboard; re-provision with a current token. No local evidence needed beyond the code. |
|
||||
| `LISTENER_BIND_FAILED` | listener | 41 | No listen address could be bound after all fallbacks (configured address, 0.0.0.0:53, localhost:53, port 5354, random) were exhausted. `detail.attempts` records each tried address with the UDP/TCP OS error, e.g. `address already in use` (another DNS service owns the port) or `can't assign requested address` (address not on any interface). | Read `detail.attempts`: `address already in use` → find the process owning the port (`sudo lsof -i :53 -nP`); `can't assign requested address` → the configured IP is not present on the device. Then fix the conflict or the listener config. |
|
||||
| `LISTENER_CONFIGURED_ADDR_UNAVAILABLE` | listener | 42 | An explicitly configured listener address could not be bound and configuration checks forbid falling back to another address, or (macOS intercept mode) the required explicit address is unavailable. | The configured `ip:port` in the listener config is wrong for this device or occupied. Verify the address exists on an interface and nothing else binds it; correct the config rather than expecting fallback. |
|
||||
| `SERVICE_INSTALL_FAILED` | service | 51 | The OS service manager refused to install the service (launchd/systemd/SCM registration failed). | Check OS-level constraints: permissions/elevation, MDM policy blocking daemon installation, corrupted previous install. Evidence: result file `message` (service manager error), plus `launchctl print system/ctrld` / `systemctl status ctrld` / SCM state. |
|
||||
| `SERVICE_START_FAILED` | service | 52 | The service installed but the service manager could not start it. | Check the service manager's own log for the start error, then the ctrld home dir `ctrld.log`. Often permissions or a binary quarantined by security tooling. |
|
||||
| `SERVICE_SELFCHECK_FAILED` | service | 53 | The service started but never became healthy: no fresher failure was reported by the daemon, and the post-install DNS self-check failed. The just-installed service is rolled back (uninstalled). If the daemon itself recorded a more specific failure (e.g. a listener code), that code is reported instead of this one. | Ask for the drained service log printed by `ctrld start` and the result file. If the service was running but unreachable, check host firewall rules intercepting DNS to the listener. |
|
||||
|
||||
## Reading the result file
|
||||
|
||||
macOS and Linux (default service home is `/etc/controld`):
|
||||
|
||||
```sh
|
||||
sudo cat /etc/controld/provision_result.json
|
||||
```
|
||||
|
||||
On Windows the file sits next to `ctrld.exe` in the install directory. A
|
||||
custom `homedir` config moves it accordingly; routers and mobile use their
|
||||
platform home directory.
|
||||
|
||||
The file sits in the same directory as the persisted internal log
|
||||
(`ctrld.log`) for the user the service runs as. On a healthy install the
|
||||
file is absent.
|
||||
|
||||
## Rules for maintainers
|
||||
|
||||
- Codes are append-only once released. Never rename, renumber, or reuse a
|
||||
code or exit number; add a new one and note the deprecation here.
|
||||
- Every code added in `cmd/cli/provision_result.go` needs a row here in the
|
||||
same MR. Tests enforce the code/stage/exit maps and that this table has
|
||||
exactly one row per code.
|
||||
- Detail must stay bounded and free of secrets: the constructor strips the
|
||||
provision token and cd UID and caps sizes; do not bypass it.
|
||||
+65
-55
@@ -7,9 +7,7 @@ On Windows, DNS intercept mode uses a two-layer architecture:
|
||||
- **`dns` mode (default)**: NRPT only — graceful DNS routing via the Windows DNS Client service
|
||||
- **`hard` mode**: NRPT + WFP — full enforcement with kernel-level block filters
|
||||
|
||||
This dual-mode design ensures that `dns` mode can never break DNS (at worst, a VPN
|
||||
overwrites NRPT and queries bypass ctrld temporarily), while `hard` mode provides
|
||||
the same enforcement guarantees as macOS pf.
|
||||
`dns` mode avoids ctrld's outbound block filters and therefore degrades more gracefully when owned NRPT is removed. Resolution still depends on the active NRPT target being reachable; an administrator-owned GP catch-all intentionally remains fail-closed if its loopback listener is stopped.
|
||||
|
||||
## Architecture: dns vs hard Mode
|
||||
|
||||
@@ -24,9 +22,8 @@ the same enforcement guarantees as macOS pf.
|
||||
│ localhost, CLEAR_ACTION_RIGHT) prevent third-party VPN WFP │
|
||||
│ blocks (e.g., OpenVPN block-outside-dns) from breaking NRPT. │
|
||||
│ │
|
||||
│ If VPN clears NRPT: health monitor re-adds within 30s │
|
||||
│ Worst case: queries go to VPN DNS until NRPT restored │
|
||||
│ DNS never breaks — graceful degradation │
|
||||
│ Owned rule missing → restore; GP missing → owned fallback │
|
||||
│ GP rule + dead listener remains intentionally fail-closed │
|
||||
└─────────────────────────────────────────────────────────────────┘
|
||||
|
||||
┌─────────────────────────────────────────────────────────────────┐
|
||||
@@ -37,9 +34,8 @@ the same enforcement guarantees as macOS pf.
|
||||
│ Bypass attempt (raw 8.8.8.8:53) → WFP BLOCK filter │
|
||||
│ VPN DNS on private IP → WFP subnet PERMIT filter → allowed │
|
||||
│ │
|
||||
│ NRPT must be active before WFP starts (atomic guarantee) │
|
||||
│ If NRPT fails → WFP not started (avoids DNS blackhole) │
|
||||
│ If WFP fails → NRPT rolled back (all-or-nothing) │
|
||||
│ NRPT route is established before WFP starts │
|
||||
│ WFP failure rolls back only ctrld-owned NRPT; GP is untouched │
|
||||
└─────────────────────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
@@ -79,7 +75,7 @@ unreliable. If we write to the GP path, DNS Client enters GP mode but the rules
|
||||
never activate — resulting in `Get-DnsClientNrptPolicy` returning empty even though
|
||||
`Get-DnsClientNrptRule` shows the rule in registry.
|
||||
|
||||
ctrld uses an adaptive strategy (matching [Tailscale's approach](https://github.com/tailscale/tailscale/blob/main/net/dns/nrpt_windows.go)):
|
||||
ctrld uses an adaptive strategy (matching [Tailscale's approach](https://github.com/tailscale/tailscale/blob/main/net/dns/nrpt_windows.go)) when it owns NRPT:
|
||||
|
||||
1. **Always write to the local path** using a deterministic GUID key name
|
||||
(`{B2E9A3C1-7F4D-4A8E-9D6B-5C1E0F3A2B8D}`). This is the baseline that works
|
||||
@@ -91,6 +87,37 @@ ctrld uses an adaptive strategy (matching [Tailscale's approach](https://github.
|
||||
the empty GP parent key. This ensures DNS Client stays in "local mode" where
|
||||
the local-path rule activates immediately via `paramchange`.
|
||||
|
||||
### Adopting an Organization-Owned GP Catch-All
|
||||
|
||||
Before applying that ctrld-owned strategy, service startup looks for a non-ctrld GP child with exactly `Name=["."]` and one `GenericDNSServers` value equal to the actual listener IP. The registry match is only an ownership candidate: ctrld sends its unique DNS Client probe before any NRPT write and re-reads the same child afterward.
|
||||
|
||||
Intercept state stays unpublished while that happens. Startup publishes nothing until it has fully succeeded, which is why the probe and heal flows take the state as an argument instead of reading the published field — publishing early would expose a half-built `wfpState`, with no engine handle and filter IDs still being assigned, to callers such as the VPN DNS exemption path. The one deliberate exception is hard mode when WFP setup fails while GP-managed NRPT is verified routing: that state is published so the health monitor can keep retrying WFP, and the service start is still reported as failed.
|
||||
|
||||
When the rule remains unchanged and the probe arrives, Group Policy owns NRPT:
|
||||
|
||||
- the startup adapter reset is skipped;
|
||||
- no ctrld NRPT key is created;
|
||||
- `RefreshPolicyEx`, Dnscache `paramchange`, and cache flush are not used for that policy;
|
||||
- shutdown/uninstall leave the GP child untouched.
|
||||
|
||||
A matching GP child that remains present but fails its probe is reported as ineffective and is not rewritten. If the child disappears, ctrld activates its normal owned fallback. A GP catch-all that instead changes to another resolver is reported as a conflict; ctrld does not create a second ambiguous catch-all. If the matching rule later returns, ctrld probes first, removes only its deterministic fallback keys, signals the ownership transition once, and resumes observing Group Policy.
|
||||
|
||||
**Deployment ordering:** apply the GPO before starting ctrld to guarantee adapter DNS is never reset. Remove/unlink it before intentionally stopping or uninstalling ctrld. A GP catch-all still targeting loopback with no listener running is a deliberate fail-closed state and will break DNS.
|
||||
|
||||
### Reproducing the Empty GP Parent Case
|
||||
|
||||
This is a production code reference, so the temporary repro script is not kept in
|
||||
the repository. For MR !942 review, the test script and exact before/after steps
|
||||
are posted in the MR discussion. The scenario to compare is:
|
||||
|
||||
1. Run the same approved PowerShell repro script against a pre-fix build and this
|
||||
branch with the same ctrld config.
|
||||
2. Create an empty GP NRPT parent key while ctrld is running in DNS intercept mode.
|
||||
3. Confirm pre-fix logs can spend policy refresh/paramchange retries while the GP
|
||||
parent remains empty.
|
||||
4. Confirm post-fix logs clean the empty GP parent, send one NRPT-change signal,
|
||||
and re-probe before normal retries.
|
||||
|
||||
### VPN Coexistence
|
||||
|
||||
NRPT uses most-specific-match. VPN NRPT rules for specific domains (e.g.,
|
||||
@@ -195,28 +222,28 @@ by the VPN's own WFP rules.
|
||||
|
||||
**Startup (hard mode):**
|
||||
```
|
||||
1. Add NRPT catch-all rule + GP refresh + DNS flush
|
||||
1. Adopt a proven matching GP catch-all, or install ctrld-owned NRPT
|
||||
2. FwpmEngineOpen0() with RPC_C_AUTHN_DEFAULT (0xFFFFFFFF)
|
||||
3. Delete stale sublayer (crash recovery)
|
||||
4. FwpmSubLayerAdd0() — weight 0xFFFF
|
||||
5. Add 4 localhost permit filters
|
||||
6. Add 4 block filters
|
||||
7. Add RFC1918 + CGNAT subnet permits
|
||||
8. Start NRPT health monitor goroutine
|
||||
8. Start ownership-aware NRPT/WFP health monitor
|
||||
```
|
||||
|
||||
**Startup (dns mode):**
|
||||
```
|
||||
1. Add NRPT catch-all rule + GP refresh + DNS flush
|
||||
1. Adopt a proven matching GP catch-all, or install ctrld-owned NRPT
|
||||
2. Activate loopback WFP protect (4 hard-permit filters for localhost DNS)
|
||||
3. Start NRPT health monitor goroutine
|
||||
3. Start ownership-aware NRPT health monitor
|
||||
```
|
||||
|
||||
**Shutdown:**
|
||||
```
|
||||
1. Stop NRPT health monitor
|
||||
2. Remove NRPT catch-all rule + DNS flush
|
||||
3. (hard mode only) Clean up all WFP filters, sublayer, close engine
|
||||
2. Remove + signal only ctrld-owned NRPT; leave GP-managed policy untouched
|
||||
3. Clean up ctrld WFP filters, sublayer, and engine session
|
||||
```
|
||||
|
||||
**Crash Recovery:**
|
||||
@@ -245,42 +272,24 @@ Windows DNS Client path to verify NRPT is actually working:
|
||||
4. ctrld's DNS handler recognizes the probe prefix and signals success
|
||||
5. If the probe times out (2s), NRPT isn't loaded yet → retry with remediation
|
||||
|
||||
### Startup Probe (Async)
|
||||
### Startup Probes
|
||||
|
||||
After NRPT setup, an async goroutine runs the probe-and-heal sequence without
|
||||
blocking startup:
|
||||
For a matching GP candidate, startup blocks for one 2-second probe before any NRPT mutation and then re-reads the same child. A received query plus the unchanged rule proves both routing and external ownership. If that first probe fails, ctrld installs its normal WFP protection and performs one more ownership-safe probe before advertising startup readiness.
|
||||
|
||||
ctrld-owned NRPT retains the asynchronous sequence:
|
||||
|
||||
```
|
||||
Probe attempt 1 (2s timeout)
|
||||
Immediate probe
|
||||
├─ Success → "NRPT verified working", done
|
||||
└─ Timeout → GP refresh + DNS flush, sleep 1s
|
||||
Probe attempt 2 (2s timeout)
|
||||
├─ Success → done
|
||||
└─ Timeout → Restart DNS Client service (nuclear), sleep 2s
|
||||
Re-add NRPT + GP refresh + DNS flush
|
||||
Probe attempt 3 (2s timeout)
|
||||
├─ Success → done
|
||||
└─ Timeout → GP refresh + DNS flush, sleep 4s
|
||||
Probe attempt 4 (2s timeout)
|
||||
├─ Success → done
|
||||
└─ Timeout → log error, continue
|
||||
└─ Timeout
|
||||
├─ Empty GP parent → clean once, signal once, re-probe
|
||||
└─ Otherwise → bounded 1s/2s/4s signal + probe retries
|
||||
└─ Still failing → two-phase remove/signal/re-add/final probe
|
||||
```
|
||||
|
||||
### DNS Client Restart (Nuclear Option)
|
||||
### GP-Managed Probe Failure
|
||||
|
||||
If GP refresh alone isn't enough, ctrld restarts the Windows DNS Client service
|
||||
(`Dnscache`). This forces the DNS Client to fully re-initialize, including
|
||||
re-reading all NRPT rules from the registry. This is the equivalent of macOS
|
||||
`forceReloadPFMainRuleset()`.
|
||||
|
||||
**Trade-offs:**
|
||||
- Briefly interrupts ALL DNS resolution (few hundred ms during restart)
|
||||
- Clears the system DNS cache (all apps need to re-resolve)
|
||||
- VPN NRPT rules survive (they're in registry, re-read on restart)
|
||||
- Enterprise security tools may log the service restart event
|
||||
|
||||
This only fires as attempt #3 after two GP refresh attempts fail — at that point
|
||||
DNS isn't working through ctrld anyway, so a brief DNS blip is acceptable.
|
||||
A matching GP child still owns Windows' effective NRPT store even when its probe fails. Writing a local rule cannot override that precedence, and rewriting the GP child would violate administrator ownership. ctrld therefore retries only its loopback WFP permit protection, reports the ineffective external policy, and leaves NRPT registry values and policy signals untouched.
|
||||
|
||||
### Health Monitor Integration
|
||||
|
||||
@@ -288,19 +297,20 @@ The 30s periodic health monitor now does actual probing, not just registry check
|
||||
|
||||
```
|
||||
Every 30s:
|
||||
├─ Registry check: nrptCatchAllRuleExists()?
|
||||
│ ├─ Missing → re-add + GP refresh + flush + probe-and-heal
|
||||
│ └─ Present → probe to verify it's actually routing
|
||||
│ ├─ Probe success → OK
|
||||
│ └─ Probe failure → probe-and-heal cycle
|
||||
├─ GP-managed owner
|
||||
│ ├─ Matching child + probe success → observe only
|
||||
│ ├─ Matching child + probe failure → WFP-only retry; no NRPT mutation
|
||||
│ └─ Matching child gone → activate ctrld-owned fallback + verify
|
||||
│
|
||||
└─ (hard mode only) Check: wfpSublayerExists()?
|
||||
├─ Missing → full restart (stopDNSIntercept + startDNSIntercept)
|
||||
└─ Present → OK
|
||||
├─ ctrld-owned owner
|
||||
│ ├─ Working matching GP child returns → remove only ctrld keys; adopt GP
|
||||
│ ├─ ctrld key missing → restore + signal + verify
|
||||
│ └─ ctrld key present → probe; run owned heal sequence on failure
|
||||
│
|
||||
└─ (hard mode) Check WFP sublayer; full intercept restart if missing
|
||||
```
|
||||
|
||||
**Singleton guard:** Only one probe-and-heal sequence runs at a time (atomic bool).
|
||||
The startup probe and health monitor cannot overlap.
|
||||
**Singleton guard:** Only one asynchronous probe-and-heal sequence runs at a time (atomic bool). Startup's GP-candidate probe completes before the health monitor starts; direct periodic probes finish before they schedule a heal sequence.
|
||||
|
||||
**Why periodic, not just network-event?** VPN software or Group Policy updates can
|
||||
clear NRPT at any time, not just during network changes. A 30s periodic check ensures
|
||||
|
||||
@@ -25,6 +25,16 @@ const (
|
||||
dohOsHeader = "x-cd-os"
|
||||
dohClientIDPrefHeader = "x-cd-cpref"
|
||||
headerApplicationDNS = "application/dns-message"
|
||||
|
||||
// dohMaxResponseSize caps the response body read from a DoH/DoH3
|
||||
// upstream. A DNS message is bounded by the protocol's 16-bit length
|
||||
// field; anything larger cannot be a valid response. The cap stops a
|
||||
// malicious or compromised upstream from driving ctrld into unbounded
|
||||
// memory growth via io.ReadAll on attacker-controlled bytes.
|
||||
dohMaxResponseSize = dns.MaxMsgSize
|
||||
// dohMaxErrorBodySize bounds how much of a non-200 response body is
|
||||
// read for inclusion in the returned error.
|
||||
dohMaxErrorBodySize = 1024
|
||||
)
|
||||
|
||||
// EncodeOsNameMap provides mapping from OS name to a shorter string, used for encoding x-cd-os value.
|
||||
@@ -130,13 +140,17 @@ func (r *dohResolver) Resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, erro
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
buf, err := io.ReadAll(resp.Body)
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, dohMaxErrorBodySize))
|
||||
return nil, fmt.Errorf("wrong response from DOH server, got: %s, status: %d", string(body), resp.StatusCode)
|
||||
}
|
||||
|
||||
buf, err := io.ReadAll(io.LimitReader(resp.Body, dohMaxResponseSize+1))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("could not read message from response: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("wrong response from DOH server, got: %s, status: %d", string(buf), resp.StatusCode)
|
||||
if len(buf) > dohMaxResponseSize {
|
||||
return nil, fmt.Errorf("DoH response exceeds %d-byte maximum DNS message size", dohMaxResponseSize)
|
||||
}
|
||||
|
||||
answer := new(dns.Msg)
|
||||
|
||||
+214
@@ -196,6 +196,7 @@ func testTLSServer(t *testing.T, handler http.Handler) (*httptest.Server, *x509.
|
||||
server := httptest.NewUnstartedServer(handler)
|
||||
server.TLS = &tls.Config{
|
||||
Certificates: []tls.Certificate{testCert.tlsCert},
|
||||
MinVersion: tls.VersionTLS12,
|
||||
}
|
||||
server.StartTLS()
|
||||
|
||||
@@ -232,6 +233,7 @@ func newTestHTTP3Server(t *testing.T, handler http.Handler) *testHTTP3Server {
|
||||
tlsConfig := &tls.Config{
|
||||
Certificates: []tls.Certificate{testCert.tlsCert},
|
||||
NextProtos: []string{"h3"}, // HTTP/3 protocol identifier
|
||||
MinVersion: tls.VersionTLS12,
|
||||
}
|
||||
|
||||
// Create HTTP/3 server
|
||||
@@ -264,3 +266,215 @@ func newTestHTTP3Server(t *testing.T, handler http.Handler) *testHTTP3Server {
|
||||
|
||||
return h3Server
|
||||
}
|
||||
|
||||
// blockingBodyHandler writes exactly nbytes of body with the given status,
|
||||
// flushes them, then blocks until release is closed WITHOUT ever returning.
|
||||
// Because the handler does not return, the response stream is never terminated
|
||||
// (no EOF/FIN). A client that stops after a bounded prefix therefore completes,
|
||||
// while a client that reads to EOF blocks. Tests set nbytes to the exact read
|
||||
// cap so the client consumes the whole written body (no half-written frame is
|
||||
// left blocking on flow control) yet still never sees EOF.
|
||||
func blockingBodyHandler(status, nbytes int, release <-chan struct{}) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", headerApplicationDNS)
|
||||
w.WriteHeader(status)
|
||||
if _, err := w.Write(make([]byte, nbytes)); err != nil {
|
||||
return
|
||||
}
|
||||
if f, ok := w.(http.Flusher); ok {
|
||||
f.Flush()
|
||||
}
|
||||
<-release
|
||||
}
|
||||
}
|
||||
|
||||
// requireBoundedResolve asserts that r.Resolve returns the expected size/status
|
||||
// error while the server is still withholding EOF (the handler is blocked in
|
||||
// blockingBodyHandler). Returning under those conditions proves ctrld read only
|
||||
// a bounded prefix of the body: a resolver that instead read to EOF would block
|
||||
// on the withheld stream and trip the deadline. This is the deterministic
|
||||
// regression guard for the issue-312 OOM protections, replacing the earlier
|
||||
// flaky server-side byte counter (issue-561).
|
||||
func requireBoundedResolve(t *testing.T, r Resolver, msg *dns.Msg, wantErrSubstr string) {
|
||||
t.Helper()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
type result struct {
|
||||
answer *dns.Msg
|
||||
err error
|
||||
}
|
||||
done := make(chan result, 1)
|
||||
go func() {
|
||||
answer, err := r.Resolve(ctx, msg)
|
||||
done <- result{answer, err}
|
||||
}()
|
||||
|
||||
select {
|
||||
case res := <-done:
|
||||
if res.err == nil {
|
||||
t.Fatalf("Resolve unexpectedly succeeded; answer=%v", res.answer)
|
||||
}
|
||||
if !strings.Contains(res.err.Error(), wantErrSubstr) {
|
||||
t.Fatalf("error %q does not contain %q", res.err, wantErrSubstr)
|
||||
}
|
||||
if res.answer != nil {
|
||||
t.Fatalf("Resolve returned non-nil answer alongside error: %v", res.answer)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("Resolve did not return while the server withheld EOF: the body is being read to EOF instead of a bounded prefix (issue-312 OOM protection missing)")
|
||||
}
|
||||
}
|
||||
|
||||
// dohUpstreamForTLSServer wires an UpstreamConfig at a local httptest TLS
|
||||
// server, trusting its self-signed certificate. BootstrapIP is set so no
|
||||
// real DNS lookup runs.
|
||||
func dohUpstreamForTLSServer(t *testing.T, srv *httptest.Server) *UpstreamConfig {
|
||||
t.Helper()
|
||||
pool := x509.NewCertPool()
|
||||
pool.AddCert(srv.Certificate())
|
||||
u, err := url.Parse(srv.URL)
|
||||
if err != nil {
|
||||
t.Fatalf("parse server URL: %v", err)
|
||||
}
|
||||
uc := &UpstreamConfig{
|
||||
Name: "doh-oversize",
|
||||
Type: ResolverTypeDOH,
|
||||
Endpoint: srv.URL + "/dns-query",
|
||||
BootstrapIP: u.Hostname(),
|
||||
Timeout: 2000,
|
||||
}
|
||||
uc.SetCertPool(pool)
|
||||
uc.Init()
|
||||
return uc
|
||||
}
|
||||
|
||||
// doh3UpstreamForAddr wires an UpstreamConfig at a local HTTP/3 server,
|
||||
// trusting its self-signed certificate.
|
||||
func doh3UpstreamForAddr(t *testing.T, addr string, cert *x509.Certificate) *UpstreamConfig {
|
||||
t.Helper()
|
||||
pool := x509.NewCertPool()
|
||||
pool.AddCert(cert)
|
||||
host, _, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
t.Fatalf("split host/port %q: %v", addr, err)
|
||||
}
|
||||
uc := &UpstreamConfig{
|
||||
Name: "doh3-oversize",
|
||||
Type: ResolverTypeDOH3,
|
||||
Endpoint: "h3://" + addr + "/dns-query",
|
||||
BootstrapIP: host,
|
||||
Timeout: 5000,
|
||||
}
|
||||
uc.SetCertPool(pool)
|
||||
uc.Init()
|
||||
return uc
|
||||
}
|
||||
|
||||
// TestDoHResolve_OversizedBody_Rejected locks in the fix for
|
||||
// github.com/Control-D-Inc/ctrld/issues/312: a malicious DoH upstream
|
||||
// returning a body larger than the DNS protocol allows must be rejected
|
||||
// with an explicit size error rather than buffered into ctrld memory.
|
||||
func TestDoHResolve_OversizedBody_Rejected(t *testing.T) {
|
||||
// Write exactly the LimitReader cap, then withhold EOF. ctrld's bounded
|
||||
// read (io.LimitReader of dohMaxResponseSize+1) returns after this prefix;
|
||||
// an unbounded read would block on the missing EOF and trip the deadline.
|
||||
release := make(chan struct{})
|
||||
defer close(release)
|
||||
srv := httptest.NewUnstartedServer(blockingBodyHandler(http.StatusOK, dohMaxResponseSize+1, release))
|
||||
testCert := generateTestCertificate(t)
|
||||
srv.TLS = &tls.Config{
|
||||
Certificates: []tls.Certificate{testCert.tlsCert},
|
||||
NextProtos: []string{"h2", "http/1.1"},
|
||||
MinVersion: tls.VersionTLS12,
|
||||
}
|
||||
srv.StartTLS()
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
uc := dohUpstreamForTLSServer(t, srv)
|
||||
r, err := NewResolver(uc)
|
||||
if err != nil {
|
||||
t.Fatalf("NewResolver: %v", err)
|
||||
}
|
||||
|
||||
msg := new(dns.Msg)
|
||||
msg.SetQuestion("example.com.", dns.TypeA)
|
||||
msg.RecursionDesired = true
|
||||
|
||||
requireBoundedResolve(t, r, msg, "maximum DNS message size")
|
||||
}
|
||||
|
||||
// TestDoHResolve_NonOKStatus_BoundedErrorBody locks in that a non-200
|
||||
// response with a huge body does not pull the body fully into ctrld
|
||||
// memory just to format an error string.
|
||||
func TestDoHResolve_NonOKStatus_BoundedErrorBody(t *testing.T) {
|
||||
// Same synchronization as the oversized-body test, but at the error-body
|
||||
// cap: the non-200 path reads through an io.LimitReader of
|
||||
// dohMaxErrorBodySize, so it must return after this prefix without EOF.
|
||||
release := make(chan struct{})
|
||||
defer close(release)
|
||||
srv := httptest.NewUnstartedServer(blockingBodyHandler(http.StatusBadGateway, dohMaxErrorBodySize, release))
|
||||
testCert := generateTestCertificate(t)
|
||||
srv.TLS = &tls.Config{
|
||||
Certificates: []tls.Certificate{testCert.tlsCert},
|
||||
NextProtos: []string{"h2", "http/1.1"},
|
||||
MinVersion: tls.VersionTLS12,
|
||||
}
|
||||
srv.StartTLS()
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
uc := dohUpstreamForTLSServer(t, srv)
|
||||
r, err := NewResolver(uc)
|
||||
if err != nil {
|
||||
t.Fatalf("NewResolver: %v", err)
|
||||
}
|
||||
|
||||
msg := new(dns.Msg)
|
||||
msg.SetQuestion("example.com.", dns.TypeA)
|
||||
msg.RecursionDesired = true
|
||||
|
||||
requireBoundedResolve(t, r, msg, "status: 502")
|
||||
}
|
||||
|
||||
// TestDoHResolve_OversizedBody_DoH3 mirrors the DoH oversized-body check
|
||||
// on the HTTP/3 transport, since github-312 specifically reproduced the
|
||||
// OOM via DoH3.
|
||||
func TestDoHResolve_OversizedBody_DoH3(t *testing.T) {
|
||||
release := make(chan struct{})
|
||||
defer close(release)
|
||||
testCert := generateTestCertificate(t)
|
||||
udpConn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 0})
|
||||
if err != nil {
|
||||
t.Fatalf("udp listen: %v", err)
|
||||
}
|
||||
h3 := &http3.Server{
|
||||
Handler: blockingBodyHandler(http.StatusOK, dohMaxResponseSize+1, release),
|
||||
TLSConfig: &tls.Config{
|
||||
Certificates: []tls.Certificate{testCert.tlsCert},
|
||||
NextProtos: []string{"h3"},
|
||||
MinVersion: tls.VersionTLS12,
|
||||
},
|
||||
}
|
||||
go func() {
|
||||
if err := h3.Serve(udpConn); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||
t.Logf("h3 server: %v", err)
|
||||
}
|
||||
}()
|
||||
t.Cleanup(func() {
|
||||
_ = h3.Close()
|
||||
_ = udpConn.Close()
|
||||
})
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
uc := doh3UpstreamForAddr(t, udpConn.LocalAddr().String(), testCert.cert)
|
||||
r, err := NewResolver(uc)
|
||||
if err != nil {
|
||||
t.Fatalf("NewResolver: %v", err)
|
||||
}
|
||||
|
||||
msg := new(dns.Msg)
|
||||
msg.SetQuestion("example.com.", dns.TypeA)
|
||||
msg.RecursionDesired = true
|
||||
|
||||
requireBoundedResolve(t, r, msg, "maximum DNS message size")
|
||||
}
|
||||
|
||||
@@ -6,15 +6,23 @@ import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"runtime"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
"github.com/quic-go/quic-go"
|
||||
)
|
||||
|
||||
// doqMaxResponseSize caps the bytes read from a DoQ stream: a 2-byte
|
||||
// length prefix plus a DNS message bounded by dns.MaxMsgSize. Anything
|
||||
// larger cannot be a valid response and is rejected before buffering more
|
||||
// data from the upstream.
|
||||
const doqMaxResponseSize = 2 + dns.MaxMsgSize
|
||||
|
||||
type doqResolver struct {
|
||||
uc *UpstreamConfig
|
||||
}
|
||||
@@ -41,6 +49,10 @@ func (r *doqResolver) Resolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, erro
|
||||
const doqPoolSize = 16
|
||||
|
||||
// doqConnPool manages a pool of QUIC connections for DoQ queries using a buffered channel.
|
||||
// A single quic.Transport (and its UDP socket) is shared by every connection in the pool,
|
||||
// so the OS socket lifecycle is tied to the pool rather than to each dial. Without this
|
||||
// ownership model, a strict DoQ upstream that triggers reconnect churn would leak one
|
||||
// caller-owned UDP socket per dial — see github.com/Control-D-Inc/ctrld/issues/309.
|
||||
type doqConnPool struct {
|
||||
uc *UpstreamConfig
|
||||
addrs []string
|
||||
@@ -48,6 +60,13 @@ type doqConnPool struct {
|
||||
tlsConfig *tls.Config
|
||||
quicConfig *quic.Config
|
||||
conns chan *doqConn
|
||||
|
||||
transportMu sync.Mutex
|
||||
transport *quic.Transport
|
||||
transportConn *net.UDPConn
|
||||
transportErr error
|
||||
transportInit bool
|
||||
closed bool
|
||||
}
|
||||
|
||||
type doqConn struct {
|
||||
@@ -64,6 +83,7 @@ func newDOQConnPool(uc *UpstreamConfig, addrs []string) *doqConnPool {
|
||||
NextProtos: []string{"doq"},
|
||||
RootCAs: uc.certPool,
|
||||
ServerName: uc.Domain,
|
||||
MinVersion: tls.VersionTLS12,
|
||||
}
|
||||
|
||||
quicConfig := &quic.Config{
|
||||
@@ -167,29 +187,70 @@ func (p *doqConnPool) doResolve(ctx context.Context, msg *dns.Msg) (*dns.Msg, er
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Read response
|
||||
buf, err := io.ReadAll(stream)
|
||||
stream.Close()
|
||||
|
||||
// Return connection to pool (mark as potentially bad if error occurred)
|
||||
isGood := err == nil && len(buf) > 0
|
||||
p.putConn(conn, isGood)
|
||||
|
||||
if err != nil {
|
||||
// RFC 9250 section 4.2 requires the client to indicate end-of-request by
|
||||
// closing the send side of the stream (STREAM FIN). Servers may defer
|
||||
// processing until FIN arrives, so the close must happen before reading.
|
||||
// Stream.Close closes only the send direction; the receive direction
|
||||
// remains open for the response.
|
||||
if err := stream.Close(); err != nil {
|
||||
p.putConn(conn, false)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// io.ReadAll hides io.EOF error, so check for empty buffer
|
||||
// A DoQ response is a 2-byte length prefix followed by a DNS message.
|
||||
// The DNS message is bounded by the protocol at dns.MaxMsgSize, so a
|
||||
// well-formed response is at most doqMaxResponseSize bytes. Read one
|
||||
// byte past that cap to distinguish "at limit" from "over limit" and
|
||||
// reject oversized responses before they can drive memory growth from
|
||||
// a malicious or compromised upstream.
|
||||
buf, err := io.ReadAll(io.LimitReader(stream, doqMaxResponseSize+1))
|
||||
if err != nil {
|
||||
p.putConn(conn, false)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// io.ReadAll hides io.EOF error, so check for empty buffer.
|
||||
if len(buf) == 0 {
|
||||
p.putConn(conn, false)
|
||||
return nil, io.EOF
|
||||
}
|
||||
|
||||
// Unpack DNS response (skip 2-byte length prefix)
|
||||
if len(buf) > doqMaxResponseSize {
|
||||
p.putConn(conn, false)
|
||||
return nil, fmt.Errorf("DoQ response exceeds %d-byte maximum", doqMaxResponseSize)
|
||||
}
|
||||
|
||||
// RFC 9250: each DoQ DNS message is encoded as a 2-octet length field
|
||||
// followed by the DNS message. Reject responses that are shorter than
|
||||
// the prefix or whose prefix declares more bytes than were received,
|
||||
// and retire the misbehaving connection. Without this guard, buf[2:]
|
||||
// would panic when len(buf) < 2.
|
||||
if len(buf) < 2 {
|
||||
p.putConn(conn, false)
|
||||
return nil, fmt.Errorf("malformed DoQ response: %d byte(s), need >= 2 for length prefix", len(buf))
|
||||
}
|
||||
respLen := int(buf[0])<<8 | int(buf[1])
|
||||
if 2+respLen > len(buf) {
|
||||
p.putConn(conn, false)
|
||||
return nil, fmt.Errorf("malformed DoQ response: length prefix %d exceeds payload %d", respLen, len(buf)-2)
|
||||
}
|
||||
|
||||
p.putConn(conn, true)
|
||||
|
||||
// Unpack DNS response (skip 2-byte length prefix).
|
||||
answer := new(dns.Msg)
|
||||
if err := answer.Unpack(buf[2:]); err != nil {
|
||||
if err := answer.Unpack(buf[2 : 2+respLen]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
answer.SetReply(msg)
|
||||
// RFC 9250 section 4.2.1 requires the DNS Message ID to be 0 on the wire,
|
||||
// so restore the downstream transaction ID for the client. Do NOT use
|
||||
// SetReply here: it rewrites the RCODE to NOERROR and overwrites the
|
||||
// Question with the request's, which would mask upstream failures from the
|
||||
// failover logic (a SERVFAIL would look like success) and let a
|
||||
// wrong-question answer pass validation and poison the cache. Preserve the
|
||||
// upstream RCODE, Question, and answer sections untouched so the proxy can
|
||||
// evaluate them. See github.com/Control-D-Inc/ctrld/issues/322.
|
||||
answer.Id = msg.Id
|
||||
return answer, nil
|
||||
}
|
||||
|
||||
@@ -233,25 +294,26 @@ func (p *doqConnPool) putConn(conn *quic.Conn, isGood bool) {
|
||||
}
|
||||
|
||||
// dialConn creates a new QUIC connection using parallel dialing like DoH3.
|
||||
// All connections from the pool multiplex on a single pool-owned UDP socket,
|
||||
// so reconnect churn cannot grow the host's FD count.
|
||||
func (p *doqConnPool) dialConn(ctx context.Context) (string, *quic.Conn, error) {
|
||||
logger := ProxyLogger.Load()
|
||||
|
||||
tr, err := p.getOrInitTransport()
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
|
||||
// If we have a bootstrap IP, use it directly
|
||||
if p.uc.BootstrapIP != "" {
|
||||
addr := net.JoinHostPort(p.uc.BootstrapIP, p.port)
|
||||
Log(ctx, logger.Debug(), "Sending DoQ request to: %s", addr)
|
||||
udpConn, err := net.ListenUDP("udp", nil)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
remoteAddr, err := net.ResolveUDPAddr("udp", addr)
|
||||
if err != nil {
|
||||
udpConn.Close()
|
||||
return "", nil, err
|
||||
}
|
||||
conn, err := quic.DialEarly(ctx, udpConn, remoteAddr, p.tlsConfig, p.quicConfig)
|
||||
conn, err := tr.DialEarly(ctx, remoteAddr, p.tlsConfig, p.quicConfig)
|
||||
if err != nil {
|
||||
udpConn.Close()
|
||||
return "", nil, err
|
||||
}
|
||||
return addr, conn, nil
|
||||
@@ -263,7 +325,7 @@ func (p *doqConnPool) dialConn(ctx context.Context) (string, *quic.Conn, error)
|
||||
dialAddrs[i] = net.JoinHostPort(p.addrs[i], p.port)
|
||||
}
|
||||
|
||||
pd := &quicParallelDialer{}
|
||||
pd := &quicParallelDialer{transport: tr}
|
||||
conn, err := pd.Dial(ctx, dialAddrs, p.tlsConfig, p.quicConfig)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
@@ -274,9 +336,35 @@ func (p *doqConnPool) dialConn(ctx context.Context) (string, *quic.Conn, error)
|
||||
return addr, conn, nil
|
||||
}
|
||||
|
||||
// CloseIdleConnections closes all connections in the pool.
|
||||
// Connections currently checked out (in use) are not closed.
|
||||
// getOrInitTransport returns the pool's shared quic.Transport, initialising it
|
||||
// on first call. Once the pool has been closed it permanently returns an error
|
||||
// so that callers cannot resurrect a dead pool.
|
||||
func (p *doqConnPool) getOrInitTransport() (*quic.Transport, error) {
|
||||
p.transportMu.Lock()
|
||||
defer p.transportMu.Unlock()
|
||||
if p.closed {
|
||||
return nil, errors.New("doq pool closed")
|
||||
}
|
||||
if p.transportInit {
|
||||
return p.transport, p.transportErr
|
||||
}
|
||||
p.transportInit = true
|
||||
udpConn, err := net.ListenUDP("udp", nil)
|
||||
if err != nil {
|
||||
p.transportErr = err
|
||||
return nil, err
|
||||
}
|
||||
p.transportConn = udpConn
|
||||
p.transport = &quic.Transport{Conn: udpConn}
|
||||
return p.transport, nil
|
||||
}
|
||||
|
||||
// CloseIdleConnections closes all idle connections, the shared quic.Transport,
|
||||
// and the pool's UDP socket. Connections currently checked out (in use) get
|
||||
// terminated by the transport close as well — without that, the OS socket
|
||||
// would remain bound to a goroutine that the caller cannot reach to clean up.
|
||||
func (p *doqConnPool) CloseIdleConnections() {
|
||||
drain:
|
||||
for {
|
||||
select {
|
||||
case dc := <-p.conns:
|
||||
@@ -284,7 +372,22 @@ func (p *doqConnPool) CloseIdleConnections() {
|
||||
dc.conn.CloseWithError(quic.ApplicationErrorCode(quic.NoError), "")
|
||||
}
|
||||
default:
|
||||
return
|
||||
break drain
|
||||
}
|
||||
}
|
||||
p.transportMu.Lock()
|
||||
if p.closed {
|
||||
p.transportMu.Unlock()
|
||||
return
|
||||
}
|
||||
p.closed = true
|
||||
tr := p.transport
|
||||
udpConn := p.transportConn
|
||||
p.transportMu.Unlock()
|
||||
if tr != nil {
|
||||
_ = tr.Close()
|
||||
}
|
||||
if udpConn != nil {
|
||||
_ = udpConn.Close()
|
||||
}
|
||||
}
|
||||
|
||||
+596
-1
@@ -1,4 +1,3 @@
|
||||
// test_helpers.go
|
||||
package ctrld
|
||||
|
||||
import (
|
||||
@@ -8,8 +7,11 @@ import (
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"io"
|
||||
"math/big"
|
||||
"net"
|
||||
"os"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -99,6 +101,7 @@ func newTestQUICServer(t *testing.T) *testQUICServer {
|
||||
tlsConfig := &tls.Config{
|
||||
Certificates: []tls.Certificate{testCert.tlsCert},
|
||||
NextProtos: []string{"doq"},
|
||||
MinVersion: tls.VersionTLS12,
|
||||
}
|
||||
|
||||
// Create QUIC listener
|
||||
@@ -221,3 +224,595 @@ func (s *testQUICServer) handleStream(t *testing.T, stream *quic.Stream) {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// malformedDoQServer is a test QUIC server that drains the client's DoQ
|
||||
// request and writes caller-supplied raw bytes back. The bytes are not
|
||||
// required to be a well-framed DoQ response, which is what lets the
|
||||
// regression tests exercise malformed-response handling.
|
||||
type malformedDoQServer struct {
|
||||
listener *quic.Listener
|
||||
cert *x509.Certificate
|
||||
addr string
|
||||
response []byte
|
||||
}
|
||||
|
||||
func newMalformedDoQServer(t *testing.T, response []byte) *malformedDoQServer {
|
||||
t.Helper()
|
||||
|
||||
testCert := generateTestCertificate(t)
|
||||
tlsConfig := &tls.Config{
|
||||
Certificates: []tls.Certificate{testCert.tlsCert},
|
||||
NextProtos: []string{"doq"},
|
||||
}
|
||||
|
||||
listener, err := quic.ListenAddr("127.0.0.1:0", tlsConfig, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create QUIC listener: %v", err)
|
||||
}
|
||||
|
||||
s := &malformedDoQServer{
|
||||
listener: listener,
|
||||
cert: testCert.cert,
|
||||
addr: listener.Addr().String(),
|
||||
response: response,
|
||||
}
|
||||
|
||||
go s.serve()
|
||||
t.Cleanup(func() { _ = listener.Close() })
|
||||
return s
|
||||
}
|
||||
|
||||
func (s *malformedDoQServer) serve() {
|
||||
for {
|
||||
conn, err := s.listener.Accept(context.Background())
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
go s.handleConn(conn)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *malformedDoQServer) handleConn(conn *quic.Conn) {
|
||||
for {
|
||||
stream, err := conn.AcceptStream(context.Background())
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
go s.handleStream(stream)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *malformedDoQServer) handleStream(stream *quic.Stream) {
|
||||
defer stream.Close()
|
||||
|
||||
// Drain the client's DoQ-framed request so the client's writes complete
|
||||
// cleanly before we reply with our attacker-controlled bytes. Using
|
||||
// io.ReadFull because a single Read on a QUIC stream may return short.
|
||||
lenBuf := make([]byte, 2)
|
||||
if _, err := io.ReadFull(stream, lenBuf); err != nil {
|
||||
return
|
||||
}
|
||||
msgLen := uint16(lenBuf[0])<<8 | uint16(lenBuf[1])
|
||||
if msgLen > 0 {
|
||||
discard := make([]byte, msgLen)
|
||||
if _, err := io.ReadFull(stream, discard); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if len(s.response) > 0 {
|
||||
_, _ = stream.Write(s.response)
|
||||
}
|
||||
}
|
||||
|
||||
// newMalformedDoQUpstream builds an UpstreamConfig wired to a local
|
||||
// malformed test server with the test certificate trusted via a custom
|
||||
// cert pool. We bypass SetupBootstrapIP by setting BootstrapIP directly,
|
||||
// so the pool dials 127.0.0.1 without any DNS lookup.
|
||||
func newMalformedDoQUpstream(t *testing.T, cert *x509.Certificate, addr string) *UpstreamConfig {
|
||||
t.Helper()
|
||||
|
||||
pool := x509.NewCertPool()
|
||||
pool.AddCert(cert)
|
||||
|
||||
host, _, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
t.Fatalf("split host/port %q: %v", addr, err)
|
||||
}
|
||||
|
||||
uc := &UpstreamConfig{
|
||||
Name: "doq-malformed",
|
||||
Type: ResolverTypeDOQ,
|
||||
Endpoint: addr,
|
||||
Domain: host,
|
||||
BootstrapIP: host,
|
||||
Timeout: 2000,
|
||||
}
|
||||
uc.SetCertPool(pool)
|
||||
return uc
|
||||
}
|
||||
|
||||
// TestDoQResolve_MalformedResponse verifies that DoQ upstream
|
||||
// responses violating RFC 9250 framing — fewer than 2 bytes, or a
|
||||
// length prefix declaring more payload than was received — return a
|
||||
// handled error instead of panicking on the length-prefix slice.
|
||||
func TestDoQResolve_MalformedResponse(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
response []byte
|
||||
}{
|
||||
// Empty stream is already handled via io.EOF; locked in so a
|
||||
// future change that drops that branch is caught.
|
||||
{"empty response", nil},
|
||||
|
||||
// One byte: too short to hold the 2-octet length prefix.
|
||||
{"single byte response", []byte{0x00}},
|
||||
|
||||
// Length prefix declares 16 bytes; payload is absent.
|
||||
{"length prefix only", []byte{0x00, 0x10}},
|
||||
|
||||
// Length prefix declares 65535 bytes; only 1 byte of payload
|
||||
// arrived.
|
||||
{"length prefix larger than payload", []byte{0xFF, 0xFF, 0x00}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
tt := tt
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server := newMalformedDoQServer(t, tt.response)
|
||||
uc := newMalformedDoQUpstream(t, server.cert, server.addr)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
pool := newDOQConnPool(uc, []string{"127.0.0.1"})
|
||||
t.Cleanup(pool.CloseIdleConnections)
|
||||
|
||||
msg := new(dns.Msg)
|
||||
msg.SetQuestion("example.com.", dns.TypeA)
|
||||
msg.RecursionDesired = true
|
||||
|
||||
answer, err := pool.Resolve(ctx, msg)
|
||||
if err == nil {
|
||||
t.Fatalf("Resolve unexpectedly succeeded for malformed response %v; answer=%v", tt.response, answer)
|
||||
}
|
||||
if answer != nil {
|
||||
t.Fatalf("Resolve returned non-nil answer alongside error: answer=%v err=%v", answer, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// strictDoQServer accepts DoQ queries but defers the response until the
|
||||
// client signals end-of-request with STREAM FIN, as required by RFC 9250
|
||||
// section 4.2. It exists to lock in the fix for
|
||||
// github.com/Control-D-Inc/ctrld/issues/309 where a client
|
||||
// that never closes its send side caused the server to wait forever and the
|
||||
// client to churn through reconnects.
|
||||
type strictDoQServer struct {
|
||||
listener *quic.Listener
|
||||
cert *x509.Certificate
|
||||
addr string
|
||||
}
|
||||
|
||||
func newStrictDoQServer(t *testing.T) *strictDoQServer {
|
||||
t.Helper()
|
||||
|
||||
testCert := generateTestCertificate(t)
|
||||
tlsConfig := &tls.Config{
|
||||
Certificates: []tls.Certificate{testCert.tlsCert},
|
||||
NextProtos: []string{"doq"},
|
||||
MinVersion: tls.VersionTLS12,
|
||||
}
|
||||
|
||||
listener, err := quic.ListenAddr("127.0.0.1:0", tlsConfig, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create QUIC listener: %v", err)
|
||||
}
|
||||
|
||||
s := &strictDoQServer{
|
||||
listener: listener,
|
||||
cert: testCert.cert,
|
||||
addr: listener.Addr().String(),
|
||||
}
|
||||
go s.serve()
|
||||
t.Cleanup(func() { _ = listener.Close() })
|
||||
return s
|
||||
}
|
||||
|
||||
func (s *strictDoQServer) serve() {
|
||||
for {
|
||||
conn, err := s.listener.Accept(context.Background())
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
go s.handleConn(conn)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *strictDoQServer) handleConn(conn *quic.Conn) {
|
||||
for {
|
||||
stream, err := conn.AcceptStream(context.Background())
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
go s.handleStream(stream)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *strictDoQServer) handleStream(stream *quic.Stream) {
|
||||
defer stream.Close()
|
||||
|
||||
// Drain until the client closes the send side. This is the behaviour
|
||||
// that triggered the bug: if the client never sends STREAM FIN, this
|
||||
// read blocks until the stream's deadline fires.
|
||||
body, err := io.ReadAll(stream)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if len(body) < 2 {
|
||||
return
|
||||
}
|
||||
msgLen := uint16(body[0])<<8 | uint16(body[1])
|
||||
if int(msgLen) != len(body)-2 {
|
||||
return
|
||||
}
|
||||
|
||||
msg := new(dns.Msg)
|
||||
if err := msg.Unpack(body[2:]); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
response := new(dns.Msg)
|
||||
response.SetReply(msg)
|
||||
response.Authoritative = true
|
||||
if len(msg.Question) > 0 && msg.Question[0].Qtype == dns.TypeA {
|
||||
response.Answer = append(response.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{
|
||||
Name: msg.Question[0].Name,
|
||||
Rrtype: dns.TypeA,
|
||||
Class: dns.ClassINET,
|
||||
Ttl: 300,
|
||||
},
|
||||
A: net.ParseIP("192.0.2.1"),
|
||||
})
|
||||
}
|
||||
|
||||
respBytes, err := response.Pack()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
respLen := uint16(len(respBytes))
|
||||
if _, err := stream.Write([]byte{byte(respLen >> 8), byte(respLen & 0xFF)}); err != nil {
|
||||
return
|
||||
}
|
||||
if _, err := stream.Write(respBytes); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
func newStrictDoQUpstream(t *testing.T, cert *x509.Certificate, addr string, useBootstrap bool) *UpstreamConfig {
|
||||
t.Helper()
|
||||
|
||||
pool := x509.NewCertPool()
|
||||
pool.AddCert(cert)
|
||||
|
||||
host, _, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
t.Fatalf("split host/port %q: %v", addr, err)
|
||||
}
|
||||
|
||||
uc := &UpstreamConfig{
|
||||
Name: "doq-strict",
|
||||
Type: ResolverTypeDOQ,
|
||||
Endpoint: addr,
|
||||
Domain: host,
|
||||
Timeout: 3000,
|
||||
}
|
||||
if useBootstrap {
|
||||
uc.BootstrapIP = host
|
||||
}
|
||||
uc.SetCertPool(pool)
|
||||
return uc
|
||||
}
|
||||
|
||||
// TestDoQResolve_StrictServerWaitsForFIN exercises the RFC 9250 client-FIN
|
||||
// requirement. With the bug present, the server's io.ReadAll blocks until
|
||||
// the stream deadline expires and the client sees a timeout, so a successful
|
||||
// resolve here proves that the client now sends STREAM FIN before reading.
|
||||
func TestDoQResolve_StrictServerWaitsForFIN(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server := newStrictDoQServer(t)
|
||||
uc := newStrictDoQUpstream(t, server.cert, server.addr, true)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
host, _, _ := net.SplitHostPort(server.addr)
|
||||
pool := newDOQConnPool(uc, []string{host})
|
||||
t.Cleanup(pool.CloseIdleConnections)
|
||||
|
||||
msg := new(dns.Msg)
|
||||
msg.SetQuestion("example.com.", dns.TypeA)
|
||||
msg.RecursionDesired = true
|
||||
|
||||
answer, err := pool.Resolve(ctx, msg)
|
||||
if err != nil {
|
||||
t.Fatalf("Resolve failed against strict DoQ server: %v", err)
|
||||
}
|
||||
if answer == nil || len(answer.Answer) == 0 {
|
||||
t.Fatalf("Resolve returned no answer records: %+v", answer)
|
||||
}
|
||||
a, ok := answer.Answer[0].(*dns.A)
|
||||
if !ok || !a.A.Equal(net.ParseIP("192.0.2.1")) {
|
||||
t.Fatalf("unexpected answer: %+v", answer.Answer[0])
|
||||
}
|
||||
}
|
||||
|
||||
// TestDoQResolve_ParallelDialPathStrictFIN exercises the parallel-dial path
|
||||
// (no BootstrapIP) against the same FIN-strict server, so that both the
|
||||
// single-dial branch and the parallel-dial branch are covered.
|
||||
func TestDoQResolve_ParallelDialPathStrictFIN(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server := newStrictDoQServer(t)
|
||||
uc := newStrictDoQUpstream(t, server.cert, server.addr, false)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
host, _, _ := net.SplitHostPort(server.addr)
|
||||
pool := newDOQConnPool(uc, []string{host})
|
||||
t.Cleanup(pool.CloseIdleConnections)
|
||||
|
||||
msg := new(dns.Msg)
|
||||
msg.SetQuestion("example.com.", dns.TypeA)
|
||||
msg.RecursionDesired = true
|
||||
|
||||
answer, err := pool.Resolve(ctx, msg)
|
||||
if err != nil {
|
||||
t.Fatalf("Resolve (parallel-dial path) failed against strict DoQ server: %v", err)
|
||||
}
|
||||
if answer == nil || len(answer.Answer) == 0 {
|
||||
t.Fatalf("Resolve (parallel-dial path) returned no answer records: %+v", answer)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDoQPool_ChurnDoesNotGrowFDs exercises the reconnect-churn scenario
|
||||
// described in github.com/Control-D-Inc/ctrld/issues/309: repeated dials
|
||||
// against a server that closes existing connections must not grow the process
|
||||
// FD count, because the pool now shares one UDP socket via quic.Transport instead
|
||||
// of allocating one per dial. Linux-only because /proc/self/fd is the cheapest
|
||||
// portable proxy for "what's still open."
|
||||
func TestDoQPool_ChurnDoesNotGrowFDs(t *testing.T) {
|
||||
if runtime.GOOS != "linux" {
|
||||
t.Skip("FD accounting via /proc/self/fd is linux-only")
|
||||
}
|
||||
t.Parallel()
|
||||
|
||||
server := newStrictDoQServer(t)
|
||||
uc := newStrictDoQUpstream(t, server.cert, server.addr, true)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
host, _, _ := net.SplitHostPort(server.addr)
|
||||
pool := newDOQConnPool(uc, []string{host})
|
||||
t.Cleanup(pool.CloseIdleConnections)
|
||||
|
||||
makeQuery := func(i int) *dns.Msg {
|
||||
msg := new(dns.Msg)
|
||||
// Vary the question so any caching layer cannot short-circuit.
|
||||
msg.SetQuestion(dns.Fqdn(strings.Repeat("a", 1+i%8)+".example.com"), dns.TypeA)
|
||||
msg.RecursionDesired = true
|
||||
return msg
|
||||
}
|
||||
|
||||
// Warm the pool so the steady-state transport and at least one
|
||||
// connection are open. Without this, the first resolve in the measured
|
||||
// loop would inflate the baseline.
|
||||
if _, err := pool.Resolve(ctx, makeQuery(0)); err != nil {
|
||||
t.Fatalf("warm-up Resolve failed: %v", err)
|
||||
}
|
||||
|
||||
baseline := countOpenFDs(t)
|
||||
|
||||
// Force reconnect churn by closing the connection between each query.
|
||||
// Without the fix this would leak one UDP socket per round; with the
|
||||
// fix the pool's shared transport keeps a single socket open.
|
||||
const rounds = 20
|
||||
for i := 1; i <= rounds; i++ {
|
||||
// Drain any pooled connection so the next Resolve has to redial.
|
||||
drainPooledConns(pool)
|
||||
|
||||
if _, err := pool.Resolve(ctx, makeQuery(i)); err != nil {
|
||||
t.Fatalf("Resolve in churn loop iteration %d failed: %v", i, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Give quic-go a moment to drop any background goroutines that hold
|
||||
// references to closed sockets.
|
||||
time.Sleep(200 * time.Millisecond)
|
||||
|
||||
after := countOpenFDs(t)
|
||||
|
||||
// Allow a small slack for transient FDs (goroutine wake-ups, qlog,
|
||||
// etc.) but reject anything that scales with the number of rounds.
|
||||
const slack = 5
|
||||
if after > baseline+slack {
|
||||
t.Fatalf("FD count grew under DoQ churn: baseline=%d after=%d rounds=%d (slack=%d)", baseline, after, rounds, slack)
|
||||
}
|
||||
}
|
||||
|
||||
// drainPooledConns removes any idle pooled connections so the next Resolve
|
||||
// is forced to dial a fresh one. It does not close the pool's transport.
|
||||
func drainPooledConns(p *doqConnPool) {
|
||||
for {
|
||||
select {
|
||||
case dc := <-p.conns:
|
||||
if dc.conn != nil {
|
||||
dc.conn.CloseWithError(quic.ApplicationErrorCode(quic.NoError), "")
|
||||
}
|
||||
default:
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func countOpenFDs(t *testing.T) int {
|
||||
t.Helper()
|
||||
entries, err := os.ReadDir("/proc/self/fd")
|
||||
if err != nil {
|
||||
t.Fatalf("read /proc/self/fd: %v", err)
|
||||
}
|
||||
return len(entries)
|
||||
}
|
||||
|
||||
// TestDoQResolve_OversizedResponse_Rejected locks in the fix for
|
||||
// github.com/Control-D-Inc/ctrld/issues/312 on the DoQ transport: a
|
||||
// malicious upstream that writes a response larger than the DNS protocol
|
||||
// allows must be rejected with an explicit size error, not buffered
|
||||
// without bound into ctrld memory.
|
||||
func TestDoQResolve_OversizedResponse_Rejected(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// doqMaxResponseSize is 2 + dns.MaxMsgSize. Send something well past
|
||||
// that. 256 KiB is enough to exceed the cap while keeping the test
|
||||
// fast on loopback.
|
||||
response := make([]byte, 256*1024)
|
||||
// A well-formed length prefix isn't required: the size cap should
|
||||
// fire before any framing check runs. Use a non-zero prefix so the
|
||||
// test also documents that the order of validation is "size first,
|
||||
// framing later."
|
||||
response[0] = 0xFF
|
||||
response[1] = 0xFF
|
||||
|
||||
server := newMalformedDoQServer(t, response)
|
||||
uc := newMalformedDoQUpstream(t, server.cert, server.addr)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
pool := newDOQConnPool(uc, []string{"127.0.0.1"})
|
||||
t.Cleanup(pool.CloseIdleConnections)
|
||||
|
||||
msg := new(dns.Msg)
|
||||
msg.SetQuestion("example.com.", dns.TypeA)
|
||||
msg.RecursionDesired = true
|
||||
|
||||
answer, err := pool.Resolve(ctx, msg)
|
||||
if err == nil {
|
||||
t.Fatalf("Resolve unexpectedly succeeded for oversized response; answer=%v", answer)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "exceeds") {
|
||||
t.Fatalf("error %q does not surface the size cap", err)
|
||||
}
|
||||
if answer != nil {
|
||||
t.Fatalf("Resolve returned non-nil answer alongside error: %v", answer)
|
||||
}
|
||||
}
|
||||
|
||||
// frameDoQResponse packs msg and prepends the RFC 9250 2-octet length prefix,
|
||||
// producing the exact bytes a DoQ server writes on the wire.
|
||||
func frameDoQResponse(t *testing.T, msg *dns.Msg) []byte {
|
||||
t.Helper()
|
||||
b, err := msg.Pack()
|
||||
if err != nil {
|
||||
t.Fatalf("pack response: %v", err)
|
||||
}
|
||||
n := uint16(len(b))
|
||||
return append([]byte{byte(n >> 8), byte(n & 0xFF)}, b...)
|
||||
}
|
||||
|
||||
// TestDoQResolve_PreservesRcode locks in the fix for
|
||||
// github.com/Control-D-Inc/ctrld/issues/322: the DoQ resolver must not rewrite
|
||||
// an upstream response with SetReply, which would clobber a SERVFAIL into
|
||||
// NOERROR and hide the failure from the proxy's failover logic. The upstream
|
||||
// RCODE must survive; only the transaction ID is restored for the client.
|
||||
func TestDoQResolve_PreservesRcode(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// RFC 9250 puts the DNS Message ID at 0 on the wire.
|
||||
resp := new(dns.Msg)
|
||||
resp.SetQuestion("example.com.", dns.TypeA)
|
||||
resp.Response = true
|
||||
resp.Id = 0
|
||||
resp.Rcode = dns.RcodeServerFailure
|
||||
|
||||
server := newMalformedDoQServer(t, frameDoQResponse(t, resp))
|
||||
uc := newMalformedDoQUpstream(t, server.cert, server.addr)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
pool := newDOQConnPool(uc, []string{"127.0.0.1"})
|
||||
t.Cleanup(pool.CloseIdleConnections)
|
||||
|
||||
msg := new(dns.Msg)
|
||||
msg.SetQuestion("example.com.", dns.TypeA)
|
||||
msg.RecursionDesired = true
|
||||
|
||||
answer, err := pool.Resolve(ctx, msg)
|
||||
if err != nil {
|
||||
t.Fatalf("Resolve failed: %v", err)
|
||||
}
|
||||
if answer.Rcode != dns.RcodeServerFailure {
|
||||
t.Fatalf("upstream SERVFAIL was rewritten to %s; failover would be bypassed",
|
||||
dns.RcodeToString[answer.Rcode])
|
||||
}
|
||||
if answer.Id != msg.Id {
|
||||
t.Fatalf("transaction ID not restored: got %d, want %d", answer.Id, msg.Id)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDoQResolve_PreservesWrongQuestion locks in the fix for
|
||||
// github.com/Control-D-Inc/ctrld/issues/322: when an upstream answers a
|
||||
// different name than asked, the resolver must preserve the upstream's
|
||||
// question rather than rewriting it to the request's question (as SetReply
|
||||
// did). Rewriting would hide the mismatch and let wrong-domain records poison
|
||||
// the shared cache.
|
||||
func TestDoQResolve_PreservesWrongQuestion(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
resp := new(dns.Msg)
|
||||
resp.SetQuestion("attacker.example.", dns.TypeA)
|
||||
resp.Response = true
|
||||
resp.Id = 0
|
||||
resp.Answer = append(resp.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{
|
||||
Name: "attacker.example.",
|
||||
Rrtype: dns.TypeA,
|
||||
Class: dns.ClassINET,
|
||||
Ttl: 300,
|
||||
},
|
||||
A: net.ParseIP("192.0.2.1"),
|
||||
})
|
||||
|
||||
server := newMalformedDoQServer(t, frameDoQResponse(t, resp))
|
||||
uc := newMalformedDoQUpstream(t, server.cert, server.addr)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
pool := newDOQConnPool(uc, []string{"127.0.0.1"})
|
||||
t.Cleanup(pool.CloseIdleConnections)
|
||||
|
||||
msg := new(dns.Msg)
|
||||
msg.SetQuestion("victim.example.", dns.TypeA)
|
||||
msg.RecursionDesired = true
|
||||
|
||||
answer, err := pool.Resolve(ctx, msg)
|
||||
if err != nil {
|
||||
t.Fatalf("Resolve failed: %v", err)
|
||||
}
|
||||
if len(answer.Question) == 0 || !strings.EqualFold(answer.Question[0].Name, "attacker.example.") {
|
||||
t.Fatalf("upstream question was rewritten; got %v, want the upstream's attacker.example.",
|
||||
answer.Question)
|
||||
}
|
||||
if answer.Id != msg.Id {
|
||||
t.Fatalf("transaction ID not restored: got %d, want %d", answer.Id, msg.Id)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -64,7 +64,8 @@ func newDOTClientPool(uc *UpstreamConfig, addrs []string) *dotConnPool {
|
||||
dialer := newDialer(net.JoinHostPort(controldPublicDns, "53"))
|
||||
|
||||
tlsConfig := &tls.Config{
|
||||
RootCAs: uc.certPool,
|
||||
RootCAs: uc.certPool,
|
||||
MinVersion: tls.VersionTLS12,
|
||||
}
|
||||
|
||||
if uc.BootstrapIP != "" {
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
module github.com/Control-D-Inc/ctrld
|
||||
|
||||
go 1.24
|
||||
go 1.25.0
|
||||
|
||||
require (
|
||||
github.com/Masterminds/semver/v3 v3.2.1
|
||||
@@ -15,7 +15,7 @@ require (
|
||||
github.com/godbus/dbus/v5 v5.1.1-0.20230522191255-76236955d466
|
||||
github.com/hashicorp/golang-lru/v2 v2.0.1
|
||||
github.com/illarion/gonotify/v2 v2.0.3
|
||||
github.com/insomniacslk/dhcp v0.0.0-20231206064809-8c70d406f6d2
|
||||
github.com/insomniacslk/dhcp v0.0.0-20260719225207-c76316d4aa82
|
||||
github.com/jaypipes/ghw v0.21.0
|
||||
github.com/jaytaylor/go-hostsfile v0.0.0-20220426042432-61485ac1fa6c
|
||||
github.com/josharian/native v1.1.1-0.20230202152459-5c7d0dd6ab86
|
||||
@@ -29,16 +29,16 @@ require (
|
||||
github.com/prometheus/client_golang v1.19.1
|
||||
github.com/prometheus/client_model v0.5.0
|
||||
github.com/prometheus/prom2json v1.3.3
|
||||
github.com/quic-go/quic-go v0.57.1
|
||||
github.com/quic-go/quic-go v0.59.1
|
||||
github.com/rs/zerolog v1.28.0
|
||||
github.com/spf13/cobra v1.9.1
|
||||
github.com/spf13/pflag v1.0.6
|
||||
github.com/spf13/viper v1.16.0
|
||||
github.com/stretchr/testify v1.11.1
|
||||
github.com/vishvananda/netlink v1.3.1
|
||||
golang.org/x/net v0.43.0
|
||||
golang.org/x/sync v0.16.0
|
||||
golang.org/x/sys v0.35.0
|
||||
golang.org/x/net v0.56.0
|
||||
golang.org/x/sync v0.22.0
|
||||
golang.org/x/sys v0.46.0
|
||||
golang.zx2c4.com/wireguard/windows v0.5.3
|
||||
tailscale.com v1.74.0
|
||||
)
|
||||
@@ -92,11 +92,11 @@ require (
|
||||
github.com/yusufpapurcu/wmi v1.2.4 // indirect
|
||||
go4.org/mem v0.0.0-20220726221520-4f986261bf13 // indirect
|
||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba // indirect
|
||||
golang.org/x/crypto v0.41.0 // indirect
|
||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 // indirect
|
||||
golang.org/x/mod v0.27.0 // indirect
|
||||
golang.org/x/text v0.28.0 // indirect
|
||||
golang.org/x/tools v0.36.0 // indirect
|
||||
golang.org/x/crypto v0.53.0 // indirect
|
||||
golang.org/x/exp v0.0.0-20240119083558-1b970713d09a // indirect
|
||||
golang.org/x/mod v0.37.0 // indirect
|
||||
golang.org/x/text v0.40.0 // indirect
|
||||
golang.org/x/tools v0.47.0 // indirect
|
||||
google.golang.org/protobuf v1.33.0 // indirect
|
||||
gopkg.in/ini.v1 v1.67.0 // indirect
|
||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||
|
||||
@@ -184,8 +184,8 @@ github.com/illarion/gonotify/v2 v2.0.3 h1:B6+SKPo/0Sw8cRJh1aLzNEeNVFfzE3c6N+o+vy
|
||||
github.com/illarion/gonotify/v2 v2.0.3/go.mod h1:38oIJTgFqupkEydkkClkbL6i5lXV/bxdH9do5TALPEE=
|
||||
github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8=
|
||||
github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw=
|
||||
github.com/insomniacslk/dhcp v0.0.0-20231206064809-8c70d406f6d2 h1:9K06NfxkBh25x56yVhWWlKFE8YpicaSfHwoV8SFbueA=
|
||||
github.com/insomniacslk/dhcp v0.0.0-20231206064809-8c70d406f6d2/go.mod h1:3A9PQ1cunSDF/1rbTq99Ts4pVnycWg+vlPkfeD2NLFI=
|
||||
github.com/insomniacslk/dhcp v0.0.0-20260719225207-c76316d4aa82 h1:y5aU8Uvl7eyM5WNgdQvRxbMJb+zo7pD+S72/Yo4pvnQ=
|
||||
github.com/insomniacslk/dhcp v0.0.0-20260719225207-c76316d4aa82/go.mod h1:qfvBmyDNp+/liLEYWRvqny/PEz9hGe2Dz833eXILSmo=
|
||||
github.com/jaypipes/ghw v0.21.0 h1:ClG2xWtYY0c1ud9jZYwVGdSgfCI7AbmZmZyw3S5HHz8=
|
||||
github.com/jaypipes/ghw v0.21.0/go.mod h1:GPrvwbtPoxYUenr74+nAnWbardIZq600vJDD5HnPsPE=
|
||||
github.com/jaypipes/pcidb v1.1.1 h1:QmPhpsbmmnCwZmHeYAATxEaoRuiMAJusKYkUncMC0ro=
|
||||
@@ -271,8 +271,8 @@ github.com/prometheus/prom2json v1.3.3 h1:IYfSMiZ7sSOfliBoo89PcufjWO4eAR0gznGcET
|
||||
github.com/prometheus/prom2json v1.3.3/go.mod h1:Pv4yIPktEkK7btWsrUTWDDDrnpUrAELaOCj+oFwlgmc=
|
||||
github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8=
|
||||
github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII=
|
||||
github.com/quic-go/quic-go v0.57.1 h1:25KAAR9QR8KZrCZRThWMKVAwGoiHIrNbT72ULHTuI10=
|
||||
github.com/quic-go/quic-go v0.57.1/go.mod h1:ly4QBAjHA2VhdnxhojRsCUOeJwKYg+taDlos92xb1+s=
|
||||
github.com/quic-go/quic-go v0.59.1 h1:0Gmua0HW1Tv7ANR7hUYwRyD0MG5OJfgvYSZasGZzBic=
|
||||
github.com/quic-go/quic-go v0.59.1/go.mod h1:upnsH4Ju1YkqpLXC305eW3yDZ4NfnNbmQRCMWS58IKU=
|
||||
github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc=
|
||||
github.com/rivo/uniseg v0.4.4 h1:8TfxU8dW6PdqD27gjM8MVNuicgxIjxpm4K7x4jp8sis=
|
||||
github.com/rivo/uniseg v0.4.4/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88=
|
||||
@@ -349,8 +349,8 @@ golang.org/x/crypto v0.0.0-20210421170649-83a5a9bb288b/go.mod h1:T9bdIzuCu7OtxOm
|
||||
golang.org/x/crypto v0.0.0-20211209193657-4570a0811e8b/go.mod h1:IxCIyHEi3zRg3s0A5j5BB6A9Jmi73HwBIUl50j+osU4=
|
||||
golang.org/x/crypto v0.0.0-20211215153901-e495a2d5b3d3/go.mod h1:IxCIyHEi3zRg3s0A5j5BB6A9Jmi73HwBIUl50j+osU4=
|
||||
golang.org/x/crypto v0.0.0-20220722155217-630584e8d5aa/go.mod h1:IxCIyHEi3zRg3s0A5j5BB6A9Jmi73HwBIUl50j+osU4=
|
||||
golang.org/x/crypto v0.41.0 h1:WKYxWedPGCTVVl5+WHSSrOBT0O8lx32+zxmHxijgXp4=
|
||||
golang.org/x/crypto v0.41.0/go.mod h1:pO5AFd7FA68rFak7rOAGVuygIISepHftHnr8dr6+sUc=
|
||||
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
|
||||
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
|
||||
golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA=
|
||||
golang.org/x/exp v0.0.0-20190306152737-a1d7652674e8/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA=
|
||||
golang.org/x/exp v0.0.0-20190510132918-efd6b22b2522/go.mod h1:ZjyILWgesfNpC6sMxTJOJm9Kp84zZh5NQWvqDGG3Qr8=
|
||||
@@ -361,8 +361,8 @@ golang.org/x/exp v0.0.0-20191227195350-da58074b4299/go.mod h1:2RIsYlXP63K8oxa1u0
|
||||
golang.org/x/exp v0.0.0-20200119233911-0405dc783f0a/go.mod h1:2RIsYlXP63K8oxa1u096TMicItID8zy7Y6sNkU49FU4=
|
||||
golang.org/x/exp v0.0.0-20200207192155-f17229e696bd/go.mod h1:J/WKrq2StrnmMY6+EHIKF9dgMWnmCNThgcyBT1FY9mM=
|
||||
golang.org/x/exp v0.0.0-20200224162631-6cc2880d07d6/go.mod h1:3jZMyOhIsHpP37uCMkUooju7aAi5cS1Q23tOzKc+0MU=
|
||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM=
|
||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc=
|
||||
golang.org/x/exp v0.0.0-20240119083558-1b970713d09a h1:Q8/wZp0KX97QFTc2ywcOE0YRjZPVIx+MXInMzdvQqcA=
|
||||
golang.org/x/exp v0.0.0-20240119083558-1b970713d09a/go.mod h1:idGWGoKP1toJGkd5/ig9ZLuPcZBC3ewk7SzmH0uou08=
|
||||
golang.org/x/image v0.0.0-20190227222117-0694c2d4d067/go.mod h1:kZ7UVZpmo3dzQBMxlp+ypCbDeSB+sBbTgSJuh5dn5js=
|
||||
golang.org/x/image v0.0.0-20190802002840-cff245a6509b/go.mod h1:FeLwcggjj3mMvU+oOTbSwawSJRM1uh48EjtB4UJZlP0=
|
||||
golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE=
|
||||
@@ -386,8 +386,8 @@ golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
||||
golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
||||
golang.org/x/mod v0.4.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
||||
golang.org/x/mod v0.4.1/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
||||
golang.org/x/mod v0.27.0 h1:kb+q2PyFnEADO2IEF935ehFUXlWiNjJWtRNgBLSfbxQ=
|
||||
golang.org/x/mod v0.27.0/go.mod h1:rWI627Fq0DEoudcK+MBkNkCe0EetEaDSwJJkCcjpazc=
|
||||
golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ=
|
||||
golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0=
|
||||
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||
golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||
golang.org/x/net v0.0.0-20190108225652-1e06a53dbb7e/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||
@@ -420,8 +420,8 @@ golang.org/x/net v0.0.0-20201209123823-ac852fbbde11/go.mod h1:m0MpNAwzfU5UDzcl9v
|
||||
golang.org/x/net v0.0.0-20201224014010-6772e930b67b/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||
golang.org/x/net v0.0.0-20211112202133-69e39bad7dc2/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
|
||||
golang.org/x/net v0.43.0 h1:lat02VYK2j4aLzMzecihNvTlJNQUq316m2Mr9rnM6YE=
|
||||
golang.org/x/net v0.43.0/go.mod h1:vhO1fvI4dGsIjh73sWfUVjj3N7CA9WkKJNQm2svM6Jg=
|
||||
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
|
||||
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
|
||||
golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U=
|
||||
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
||||
golang.org/x/oauth2 v0.0.0-20190604053449-0f29369cfe45/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
||||
@@ -441,8 +441,8 @@ golang.org/x/sync v0.0.0-20200317015054-43a5402ce75a/go.mod h1:RxMgew5VJxzue5/jJ
|
||||
golang.org/x/sync v0.0.0-20200625203802-6e8e738ad208/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.16.0 h1:ycBJEhp9p4vXvUZNszeOq0kGTPghopOL8q0fq3vstxw=
|
||||
golang.org/x/sync v0.16.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA=
|
||||
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
||||
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20190312061237-fead79001313/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
@@ -492,8 +492,8 @@ golang.org/x/sys v0.4.1-0.20230131160137-e7d7f63158de/go.mod h1:oPkhp1MJrh7nUepC
|
||||
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.35.0 h1:vz1N37gP5bs89s7He8XuIYXpyY0+QlsKmzipCbUtyxI=
|
||||
golang.org/x/sys v0.35.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
|
||||
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
|
||||
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/term v0.0.0-20201117132131-f5c789dd3221/go.mod h1:Nr5EML6q2oocZ2LXRh80K7BxOlk5/8JxuGnuhpl+muw=
|
||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||
golang.org/x/text v0.0.0-20170915032832-14c0d48ead0c/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||
@@ -504,13 +504,11 @@ golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||
golang.org/x/text v0.3.4/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
|
||||
golang.org/x/text v0.28.0 h1:rhazDwis8INMIwQ4tpjLDzUhx6RlXqZNPEM0huQojng=
|
||||
golang.org/x/text v0.28.0/go.mod h1:U8nCwOR8jO/marOQ0QbDiOngZVEBB7MAiitBuMjXiNU=
|
||||
golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs=
|
||||
golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY=
|
||||
golang.org/x/time v0.0.0-20181108054448-85acf8d2951c/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ=
|
||||
golang.org/x/time v0.0.0-20190308202827-9d24e82272b4/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ=
|
||||
golang.org/x/time v0.0.0-20191024005414-555d28b269f0/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ=
|
||||
golang.org/x/time v0.12.0 h1:ScB/8o8olJvc+CQPWrK3fPZNfh7qgwCrY0zJmoEQLSE=
|
||||
golang.org/x/time v0.12.0/go.mod h1:CDIdPxbZBQxdj6cxyCIdrNogrJKMJ7pr37NYpMcMDSg=
|
||||
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
||||
golang.org/x/tools v0.0.0-20190114222345-bf090417da8b/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
||||
golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY=
|
||||
@@ -558,8 +556,8 @@ golang.org/x/tools v0.0.0-20201208233053-a543418bbed2/go.mod h1:emZCQorbCU4vsT4f
|
||||
golang.org/x/tools v0.0.0-20210105154028-b0ab187a4818/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA=
|
||||
golang.org/x/tools v0.0.0-20210108195828-e2f9c7f1fc8e/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA=
|
||||
golang.org/x/tools v0.1.0/go.mod h1:xkSsbof2nBLbhDlRMhhhyNLN/zl3eTqcnHD5viDpcZ0=
|
||||
golang.org/x/tools v0.36.0 h1:kWS0uv/zsvHEle1LbV5LE8QujrxB3wfQyxHfhOk0Qkg=
|
||||
golang.org/x/tools v0.36.0/go.mod h1:WBDiHKJK8YgLHlcQPYQzNCkUxUypCaa5ZegCVutKm+s=
|
||||
golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q=
|
||||
golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA=
|
||||
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||
|
||||
@@ -11,7 +11,8 @@ func TestCACertPool(t *testing.T) {
|
||||
c := &http.Client{
|
||||
Transport: &http.Transport{
|
||||
TLSClientConfig: &tls.Config{
|
||||
RootCAs: CACertPool(),
|
||||
RootCAs: CACertPool(),
|
||||
MinVersion: tls.VersionTLS12,
|
||||
},
|
||||
},
|
||||
Timeout: 2 * time.Second,
|
||||
|
||||
+68
-22
@@ -63,12 +63,35 @@ type ErrorResponse struct {
|
||||
Message string `json:"message"`
|
||||
Code int `json:"code"`
|
||||
} `json:"error"`
|
||||
// StatusCode is the HTTP status the API answered with. It is not part of the JSON
|
||||
// body: this type is built for *any* non-200 whose body decodes, so the body alone
|
||||
// cannot tell a permanent rejection of the request from a transient server-side
|
||||
// failure, and callers that act differently on the two need the status to tell them
|
||||
// apart. Zero means the status was not recorded.
|
||||
StatusCode int `json:"-"`
|
||||
}
|
||||
|
||||
func (u ErrorResponse) Error() string {
|
||||
return u.ErrorField.Message
|
||||
}
|
||||
|
||||
// apiErrorFromResponse builds the error for a non-200 API answer, recording the HTTP
|
||||
// status alongside the decoded body.
|
||||
//
|
||||
// The status is what tells a caller whether the answer will change on a retry: this type
|
||||
// is built for every non-200 whose body decodes, so a 502 from a load balancer and a 404
|
||||
// for a deleted device are otherwise indistinguishable. Both response paths go through
|
||||
// here so neither can decode a body and forget to record it.
|
||||
func apiErrorFromResponse(statusCode int, d *json.Decoder) (*ErrorResponse, error) {
|
||||
errResp := &ErrorResponse{StatusCode: statusCode}
|
||||
if err := d.Decode(errResp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Decode fills exported fields from the body; StatusCode is json:"-", so it survives.
|
||||
errResp.StatusCode = statusCode
|
||||
return errResp, nil
|
||||
}
|
||||
|
||||
type utilityRequest struct {
|
||||
UID string `json:"uid"`
|
||||
ClientID string `json:"client_id,omitempty"`
|
||||
@@ -96,7 +119,7 @@ type LogsRequest struct {
|
||||
}
|
||||
|
||||
// FetchResolverConfig fetch Control D config for given uid.
|
||||
func FetchResolverConfig(req *ResolverConfigRequest, cdDev bool) (*ResolverConfig, error) {
|
||||
func FetchResolverConfig(ctx context.Context, req *ResolverConfigRequest, cdDev bool) (*ResolverConfig, error) {
|
||||
uid, clientID := ParseRawUID(req.RawUID)
|
||||
uReq := utilityRequest{
|
||||
UID: uid,
|
||||
@@ -106,11 +129,11 @@ func FetchResolverConfig(req *ResolverConfigRequest, cdDev bool) (*ResolverConfi
|
||||
uReq.ClientID = clientID
|
||||
}
|
||||
body, _ := json.Marshal(uReq)
|
||||
return postUtilityAPI(req.Version, cdDev, false, bytes.NewReader(body))
|
||||
return postUtilityAPI(ctx, req.Version, cdDev, false, bytes.NewReader(body))
|
||||
}
|
||||
|
||||
// FetchResolverUID fetch resolver uid from a given request.
|
||||
func FetchResolverUID(req *UtilityOrgRequest, version string, cdDev bool) (*ResolverConfig, error) {
|
||||
func FetchResolverUID(ctx context.Context, req *UtilityOrgRequest, version string, cdDev bool) (*ResolverConfig, error) {
|
||||
if req == nil {
|
||||
return nil, errors.New("invalid request")
|
||||
}
|
||||
@@ -131,26 +154,29 @@ func FetchResolverUID(req *UtilityOrgRequest, version string, cdDev bool) (*Reso
|
||||
ctrld.ProxyLogger.Load().Debug().Msgf("Sending UID request to ControlD API")
|
||||
|
||||
body, _ := json.Marshal(req)
|
||||
return postUtilityAPI(version, cdDev, false, bytes.NewReader(body))
|
||||
return postUtilityAPI(ctx, version, cdDev, false, bytes.NewReader(body))
|
||||
}
|
||||
|
||||
// UpdateCustomLastFailed calls API to mark custom config is bad.
|
||||
func UpdateCustomLastFailed(rawUID, version string, cdDev, lastUpdatedFailed bool) (*ResolverConfig, error) {
|
||||
func UpdateCustomLastFailed(ctx context.Context, rawUID, version string, cdDev, lastUpdatedFailed bool) (*ResolverConfig, error) {
|
||||
uid, clientID := ParseRawUID(rawUID)
|
||||
req := utilityRequest{UID: uid}
|
||||
if clientID != "" {
|
||||
req.ClientID = clientID
|
||||
}
|
||||
body, _ := json.Marshal(req)
|
||||
return postUtilityAPI(version, cdDev, true, bytes.NewReader(body))
|
||||
return postUtilityAPI(ctx, version, cdDev, true, bytes.NewReader(body))
|
||||
}
|
||||
|
||||
func postUtilityAPI(version string, cdDev, lastUpdatedFailed bool, body io.Reader) (*ResolverConfig, error) {
|
||||
func postUtilityAPI(ctx context.Context, version string, cdDev, lastUpdatedFailed bool, body io.Reader) (*ResolverConfig, error) {
|
||||
apiUrl := resolverDataURLCom
|
||||
if cdDev {
|
||||
apiUrl = resolverDataURLDev
|
||||
}
|
||||
req, err := http.NewRequest("POST", apiUrl, body)
|
||||
// Context-bound so an in-flight request is abandoned when the caller is
|
||||
// cancelled - a service stop during API preflight must not wait out the
|
||||
// request timeout, let alone keep retrying.
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", apiUrl, body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("http.NewRequest: %w", err)
|
||||
}
|
||||
@@ -174,8 +200,8 @@ func postUtilityAPI(version string, cdDev, lastUpdatedFailed bool, body io.Reade
|
||||
defer resp.Body.Close()
|
||||
d := json.NewDecoder(resp.Body)
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
errResp := &ErrorResponse{}
|
||||
if err := d.Decode(errResp); err != nil {
|
||||
errResp, err := apiErrorFromResponse(resp.StatusCode, d)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return nil, errResp
|
||||
@@ -189,13 +215,13 @@ func postUtilityAPI(version string, cdDev, lastUpdatedFailed bool, body io.Reade
|
||||
}
|
||||
|
||||
// SendLogs sends runtime log to ControlD API.
|
||||
func SendLogs(lr *LogsRequest, cdDev bool) error {
|
||||
func SendLogs(ctx context.Context, lr *LogsRequest, cdDev bool) error {
|
||||
defer lr.Data.Close()
|
||||
apiUrl := logURLCom
|
||||
if cdDev {
|
||||
apiUrl = logURLDev
|
||||
}
|
||||
req, err := http.NewRequest("POST", apiUrl, lr.Data)
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", apiUrl, lr.Data)
|
||||
if err != nil {
|
||||
return fmt.Errorf("http.NewRequest: %w", err)
|
||||
}
|
||||
@@ -215,8 +241,8 @@ func SendLogs(lr *LogsRequest, cdDev bool) error {
|
||||
defer resp.Body.Close()
|
||||
d := json.NewDecoder(resp.Body)
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
errResp := &ErrorResponse{}
|
||||
if err := d.Decode(errResp); err != nil {
|
||||
errResp, err := apiErrorFromResponse(resp.StatusCode, d)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return errResp
|
||||
@@ -293,7 +319,7 @@ func apiTransport(cdDev bool) *http.Transport {
|
||||
return dial(ctx, "tcp6", addrsFromPort(apiIpsV6, port))
|
||||
}
|
||||
if router.Name() == ddwrt.Name || runtime.GOOS == "android" {
|
||||
transport.TLSClientConfig = &tls.Config{RootCAs: certs.CACertPool()}
|
||||
transport.TLSClientConfig = &tls.Config{RootCAs: certs.CACertPool(), MinVersion: tls.VersionTLS12}
|
||||
}
|
||||
return transport
|
||||
}
|
||||
@@ -306,16 +332,29 @@ func addrsFromPort(ips []string, port string) []string {
|
||||
return addrs
|
||||
}
|
||||
|
||||
// doWithFallback sends req, retrying against apiIp directly if the first attempt
|
||||
// fails (typically because DNS is not usable yet).
|
||||
//
|
||||
// Both failures are reported. The first attempt carries the diagnosis - on Windows
|
||||
// a local firewall denying the socket surfaces there as WSAEACCES ("An attempt was
|
||||
// made to access a socket in a way forbidden by its access permissions"), which
|
||||
// says the host is blocking ctrld rather than that the network is down. Returning
|
||||
// only the fallback error hid that behind a bare "no route to host" from the IPv6
|
||||
// attempt.
|
||||
func doWithFallback(client *http.Client, req *http.Request, apiIp string) (*http.Response, error) {
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
ctrld.ProxyLogger.Load().Warn().Err(err).Msgf("failed to send request, fallback to direct IP: %s", apiIp)
|
||||
ipReq := req.Clone(req.Context())
|
||||
ipReq.Host = apiIp
|
||||
ipReq.URL.Host = apiIp
|
||||
resp, err = client.Do(ipReq)
|
||||
if err == nil {
|
||||
return resp, nil
|
||||
}
|
||||
return resp, err
|
||||
ctrld.ProxyLogger.Load().Warn().Err(err).Msgf("failed to send request, fallback to direct IP: %s", apiIp)
|
||||
ipReq := req.Clone(req.Context())
|
||||
ipReq.Host = apiIp
|
||||
ipReq.URL.Host = apiIp
|
||||
resp, fallbackErr := client.Do(ipReq)
|
||||
if fallbackErr != nil {
|
||||
return nil, fmt.Errorf("request failed: %w; fallback to direct ip %s failed: %w", err, apiIp, fallbackErr)
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// apiServerIP returns the direct IP to connect to API server.
|
||||
@@ -325,3 +364,10 @@ func apiServerIP(cdDev bool) string {
|
||||
}
|
||||
return apiDomainComIPv4
|
||||
}
|
||||
|
||||
// DoWithFallbackForTest exposes doWithFallback so tests outside this package can drive
|
||||
// the real two-attempt composition through the real retry predicate, rather than
|
||||
// asserting a copy of this error shape against another copy of it.
|
||||
func DoWithFallbackForTest(client *http.Client, req *http.Request, apiIp string) (*http.Response, error) {
|
||||
return doWithFallback(client, req, apiIp)
|
||||
}
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
package controld
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -29,3 +32,67 @@ func Test_parseUID(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestAPIErrorRecordsHTTPStatus pins the plumbing the caller's exit decision rests on.
|
||||
//
|
||||
// cmd/cli treats a 4xx as "this configuration is refused, restarting cannot help" and
|
||||
// exits cleanly, while a 5xx keeps the abnormal exit so the service manager retries. Both
|
||||
// readings need the status, and it is not in the JSON body - so a decode path that
|
||||
// forgets to record it would quietly send every API error down the retry branch,
|
||||
// including a deleted device that should self-uninstall and stop.
|
||||
func TestAPIErrorRecordsHTTPStatus(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
statusCode int
|
||||
body string
|
||||
wantCode int
|
||||
wantMsg string
|
||||
}{
|
||||
{
|
||||
name: "deleted device",
|
||||
statusCode: http.StatusNotFound,
|
||||
body: `{"error":{"message":"device does not exist","code":40402}}`,
|
||||
wantCode: InvalidConfigCode,
|
||||
wantMsg: "device does not exist",
|
||||
},
|
||||
{
|
||||
// A gateway error body carries no error object at all, which decodes
|
||||
// cleanly into the zero value - so the status is the only thing that
|
||||
// distinguishes it from a real rejection.
|
||||
name: "gateway error with an empty body",
|
||||
statusCode: http.StatusBadGateway,
|
||||
body: `{}`,
|
||||
},
|
||||
{
|
||||
name: "service unavailable",
|
||||
statusCode: http.StatusServiceUnavailable,
|
||||
body: `{"error":{"message":"try again later","code":0}}`,
|
||||
wantMsg: "try again later",
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
d := json.NewDecoder(strings.NewReader(tc.body))
|
||||
errResp, err := apiErrorFromResponse(tc.statusCode, d)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected decode error: %v", err)
|
||||
}
|
||||
if errResp.StatusCode != tc.statusCode {
|
||||
t.Errorf("StatusCode = %d, want %d: the caller cannot tell a permanent rejection from a transient failure without it", errResp.StatusCode, tc.statusCode)
|
||||
}
|
||||
if errResp.ErrorField.Code != tc.wantCode {
|
||||
t.Errorf("code = %d, want %d", errResp.ErrorField.Code, tc.wantCode)
|
||||
}
|
||||
if errResp.Error() != tc.wantMsg {
|
||||
t.Errorf("message = %q, want %q", errResp.Error(), tc.wantMsg)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
t.Run("an undecodable body is reported as a decode failure", func(t *testing.T) {
|
||||
d := json.NewDecoder(strings.NewReader("<html>502 Bad Gateway</html>"))
|
||||
if _, err := apiErrorFromResponse(http.StatusBadGateway, d); err == nil {
|
||||
t.Error("expected a decode error for a non-JSON body")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -0,0 +1,155 @@
|
||||
package controld
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"syscall"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// errRoundTripper fails the hostname attempt and the direct-ip attempt with
|
||||
// different errors, mimicking the Firewall Mode incident: the hostname attempt is
|
||||
// denied by a local firewall (WSAEACCES on Windows) while the direct-ip fallback
|
||||
// reports an unreachable IPv6 route.
|
||||
type errRoundTripper struct {
|
||||
hostname string
|
||||
firstErr error
|
||||
fbErr error
|
||||
fbCalled bool
|
||||
}
|
||||
|
||||
func (rt *errRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
if req.URL.Host == rt.hostname {
|
||||
return nil, &net.OpError{Op: "dial", Net: "tcp4", Err: rt.firstErr}
|
||||
}
|
||||
rt.fbCalled = true
|
||||
if rt.fbErr == nil {
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: http.NoBody,
|
||||
Request: req,
|
||||
}, nil
|
||||
}
|
||||
return nil, &net.OpError{Op: "dial", Net: "tcp6", Err: rt.fbErr}
|
||||
}
|
||||
|
||||
// wsaEACCES is WSAEACCES (10013): "An attempt was made to access a socket in a way
|
||||
// forbidden by its access permissions." The value is what Windows reports when a
|
||||
// WFP filter denies the connect; it is used here as a plain errno so the test runs
|
||||
// on every platform.
|
||||
const wsaEACCES = syscall.Errno(10013)
|
||||
|
||||
func TestDoWithFallbackPreservesFirstError(t *testing.T) {
|
||||
const (
|
||||
hostname = "api.controld.com"
|
||||
apiIP = "147.185.34.1"
|
||||
)
|
||||
rt := &errRoundTripper{
|
||||
hostname: hostname,
|
||||
firstErr: wsaEACCES,
|
||||
fbErr: syscall.EHOSTUNREACH,
|
||||
}
|
||||
req, err := http.NewRequest(http.MethodPost, "https://"+hostname+"/utility", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
resp, err := doWithFallback(&http.Client{Transport: rt}, req, apiIP)
|
||||
if err == nil {
|
||||
t.Fatalf("expected an error, got response %v", resp)
|
||||
}
|
||||
if !rt.fbCalled {
|
||||
t.Error("direct-ip fallback was not attempted")
|
||||
}
|
||||
|
||||
// The actionable failure must survive: an operator reading this error has to be
|
||||
// able to tell "the host is blocking us" from "the network is down".
|
||||
if !errors.Is(err, wsaEACCES) {
|
||||
t.Errorf("first-attempt error (WSAEACCES) was lost, got: %v", err)
|
||||
}
|
||||
if !errors.Is(err, syscall.EHOSTUNREACH) {
|
||||
t.Errorf("fallback error was lost, got: %v", err)
|
||||
}
|
||||
if got := err.Error(); !strings.Contains(got, apiIP) {
|
||||
t.Errorf("error does not mention the fallback ip %q: %v", apiIP, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoWithFallbackSucceedsOnFallback(t *testing.T) {
|
||||
const hostname = "api.controld.com"
|
||||
rt := &errRoundTripper{hostname: hostname, firstErr: wsaEACCES}
|
||||
req, err := http.NewRequest(http.MethodPost, "https://"+hostname+"/utility", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
resp, err := doWithFallback(&http.Client{Transport: rt}, req, "147.185.34.1")
|
||||
if err != nil {
|
||||
t.Fatalf("expected the fallback to succeed, got: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Errorf("StatusCode = %d, want %d", resp.StatusCode, http.StatusOK)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoWithFallbackNoFallbackOnSuccess(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
req, err := http.NewRequest(http.MethodPost, srv.URL, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp, err := doWithFallback(srv.Client(), req, "127.0.0.2")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Errorf("StatusCode = %d, want %d", resp.StatusCode, http.StatusOK)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDoWithFallbackComposesHostnameAttemptFirst pins the order of the composed error.
|
||||
//
|
||||
// The order is not cosmetic. cmd/cli's preflight retry predicate classifies this error
|
||||
// with errors.As, which returns the first match in the tree, so whichever attempt is
|
||||
// wrapped first decides whether processCDFlags keeps backing off or fails fast. That
|
||||
// predicate lives in another package and cannot be called from here, so this test
|
||||
// guards the property it depends on: the hostname attempt - the one that carries the
|
||||
// diagnosis - must come first.
|
||||
func TestDoWithFallbackComposesHostnameAttemptFirst(t *testing.T) {
|
||||
const hostname = "api.controld.com"
|
||||
rt := &errRoundTripper{
|
||||
hostname: hostname,
|
||||
firstErr: wsaEACCES,
|
||||
fbErr: syscall.EHOSTUNREACH,
|
||||
}
|
||||
req, err := http.NewRequest(http.MethodPost, "https://"+hostname+"/utility", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, gotErr := doWithFallback(&http.Client{Transport: rt}, req, "147.185.34.1")
|
||||
if gotErr == nil {
|
||||
t.Fatal("expected both attempts to fail")
|
||||
}
|
||||
|
||||
var opErr *net.OpError
|
||||
if !errors.As(gotErr, &opErr) {
|
||||
t.Fatalf("no net.OpError in the chain: %v", gotErr)
|
||||
}
|
||||
if !errors.Is(opErr.Err, wsaEACCES) {
|
||||
t.Errorf("first OpError in the chain is %v, want the hostname attempt (%v)", opErr.Err, wsaEACCES)
|
||||
}
|
||||
// The tcp4/tcp6 split distinguishes the two attempts in the fake transport.
|
||||
if opErr.Net != "tcp4" {
|
||||
t.Errorf("first OpError is from the %s attempt, want tcp4 (hostname)", opErr.Net)
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,8 @@
|
||||
package dnscache
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -16,11 +18,17 @@ type Cacher interface {
|
||||
}
|
||||
|
||||
// Key is the caching key for DNS message.
|
||||
//
|
||||
// ECS partitions the cache by EDNS Client Subnet so an answer resolved for one
|
||||
// subnet is never served to a client in a different subnet. Answer records are
|
||||
// scoped to the network that generated them (RFC 7871 §7.3), so they must not
|
||||
// be shared across subnets even when the question is otherwise identical.
|
||||
type Key struct {
|
||||
Qtype uint16
|
||||
Qclass uint16
|
||||
Name string
|
||||
Upstream string
|
||||
ECS string
|
||||
}
|
||||
|
||||
type Value struct {
|
||||
@@ -60,7 +68,58 @@ func NewLRUCache(size int) (*LRUCache, error) {
|
||||
// NewKey creates a new cache key for given DNS message.
|
||||
func NewKey(msg *dns.Msg, upstream string) Key {
|
||||
q := msg.Question[0]
|
||||
return Key{Qtype: q.Qtype, Qclass: q.Qclass, Name: normalizeQname(q.Name), Upstream: upstream}
|
||||
return Key{Qtype: q.Qtype, Qclass: q.Qclass, Name: normalizeQname(q.Name), Upstream: upstream, ECS: CanonicalECS(msg)}
|
||||
}
|
||||
|
||||
// CanonicalECS returns a canonical string form of the EDNS Client Subnet (ECS,
|
||||
// EDNS option 8) carried by msg, suitable for partitioning cache and
|
||||
// singleflight keys. A request with no ECS option returns "", so all ECS-less
|
||||
// queries share one partition and behave as before.
|
||||
//
|
||||
// A request that DOES carry an ECS option is never mapped to "", even at SOURCE
|
||||
// PREFIX-LENGTH 0: the /0 query is forwarded with an ECS option, may draw
|
||||
// different upstream data than an ECS-less query, and its answer (with the
|
||||
// echoed option) must not be served to a client that sent no ECS — RFC 7871
|
||||
// §7.3.1 requires /0-cached data to remain distinguishable. The family is part
|
||||
// of the token, so an IPv4 /0 and an IPv6 /0 stay distinct too.
|
||||
//
|
||||
// The address is masked to its source prefix length so only the significant
|
||||
// subnet bits contribute to the key: two clients in the same subnet share a
|
||||
// partition, while different subnets (or address families) never do. The
|
||||
// response-only SCOPE PREFIX-LENGTH is deliberately excluded — it is not part of
|
||||
// what the client asked for.
|
||||
func CanonicalECS(msg *dns.Msg) string {
|
||||
opt := msg.IsEdns0()
|
||||
if opt == nil {
|
||||
return ""
|
||||
}
|
||||
for _, o := range opt.Option {
|
||||
e, ok := o.(*dns.EDNS0_SUBNET)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
bits := int(e.SourceNetmask)
|
||||
total := 128
|
||||
zero := net.IPv6zero
|
||||
if e.Family == 1 {
|
||||
total = 32
|
||||
zero = net.IPv4zero
|
||||
}
|
||||
addr := e.Address
|
||||
if len(addr) == 0 {
|
||||
// A /0 request commonly carries an empty address; normalize it to
|
||||
// the family zero so its token is stable regardless of encoding.
|
||||
addr = zero
|
||||
}
|
||||
masked := addr
|
||||
if bits <= total {
|
||||
if m := addr.Mask(net.CIDRMask(bits, total)); m != nil {
|
||||
masked = m
|
||||
}
|
||||
}
|
||||
return fmt.Sprintf("%d/%d/%s", e.Family, e.SourceNetmask, masked.String())
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// NewValue creates a new cache value for given DNS message.
|
||||
|
||||
@@ -0,0 +1,186 @@
|
||||
package dnscache
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
// msgWithECS builds an A query for name carrying an EDNS Client Subnet option, or none
|
||||
// when family == 0.
|
||||
func msgWithECS(name string, family uint16, prefix uint8, addr string) *dns.Msg {
|
||||
m := new(dns.Msg)
|
||||
m.SetQuestion(dns.Fqdn(name), dns.TypeA)
|
||||
if family == 0 {
|
||||
return m
|
||||
}
|
||||
m.SetEdns0(4096, true)
|
||||
m.IsEdns0().Option = append(m.IsEdns0().Option, &dns.EDNS0_SUBNET{
|
||||
Code: dns.EDNS0SUBNET,
|
||||
Family: family,
|
||||
SourceNetmask: prefix,
|
||||
Address: net.ParseIP(addr),
|
||||
})
|
||||
return m
|
||||
}
|
||||
|
||||
// answerWithA builds a cached answer holding a single A record with the given address.
|
||||
func answerWithA(name, a string) *dns.Msg {
|
||||
m := new(dns.Msg)
|
||||
m.SetQuestion(dns.Fqdn(name), dns.TypeA)
|
||||
rr, err := dns.NewRR(dns.Fqdn(name) + " 300 IN A " + a)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
m.Answer = []dns.RR{rr}
|
||||
return m
|
||||
}
|
||||
|
||||
func firstA(msg *dns.Msg) string {
|
||||
for _, rr := range msg.Answer {
|
||||
if a, ok := rr.(*dns.A); ok {
|
||||
return a.A.String()
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// TestCanonicalECS covers the key-partitioning helper: only a request with no ECS option
|
||||
// collapses to the shared empty partition, while a carried /0 keeps its own family-scoped
|
||||
// token (distinct from no-ECS, per RFC 7871 §7.3.1); host bits below the source prefix do
|
||||
// not fragment the key, and different subnets / families produce different keys.
|
||||
func TestCanonicalECS(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
msg *dns.Msg
|
||||
want string
|
||||
}{
|
||||
{"no ecs", msgWithECS("controld.com", 0, 0, ""), ""},
|
||||
// A carried ECS option is never "" (RFC 7871 §7.3.1), and IPv4 /0 vs
|
||||
// IPv6 /0 stay distinct; the address encoding must not change the token.
|
||||
{"ipv4 /0", msgWithECS("controld.com", 1, 0, "0.0.0.0"), "1/0/0.0.0.0"},
|
||||
{"ipv4 /0 empty addr", msgWithECS("controld.com", 1, 0, ""), "1/0/0.0.0.0"},
|
||||
{"ipv6 /0", msgWithECS("controld.com", 2, 0, "::"), "2/0/::"},
|
||||
{"ipv6 /0 empty addr", msgWithECS("controld.com", 2, 0, ""), "2/0/::"},
|
||||
{"ipv4 /24", msgWithECS("controld.com", 1, 24, "203.0.113.0"), "1/24/203.0.113.0"},
|
||||
{"ipv4 host bits masked", msgWithECS("controld.com", 1, 24, "203.0.113.7"), "1/24/203.0.113.0"},
|
||||
{"ipv6 /64", msgWithECS("controld.com", 2, 64, "2001:db8:1::"), "2/64/2001:db8:1::"},
|
||||
{"ipv6 host bits masked", msgWithECS("controld.com", 2, 64, "2001:db8:1::dead:beef"), "2/64/2001:db8:1::"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := CanonicalECS(tt.msg); got != tt.want {
|
||||
t.Fatalf("CanonicalECS = %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewKey_ECSPartition verifies the cache key distinguishes subnets while collapsing
|
||||
// same-subnet requests, so the LRU cache cannot serve one subnet's records to another.
|
||||
func TestNewKey_ECSPartition(t *testing.T) {
|
||||
const up = "https://dns.example/dns-query"
|
||||
subnetA := NewKey(msgWithECS("controld.com", 2, 64, "2001:db8:1::"), up)
|
||||
subnetB := NewKey(msgWithECS("controld.com", 2, 64, "2001:db8:2::"), up)
|
||||
subnetAHost := NewKey(msgWithECS("controld.com", 2, 64, "2001:db8:1::5"), up)
|
||||
ipv4 := NewKey(msgWithECS("controld.com", 1, 24, "203.0.113.0"), up)
|
||||
noECS := NewKey(msgWithECS("controld.com", 0, 0, ""), up)
|
||||
|
||||
if subnetA == subnetB {
|
||||
t.Fatal("different subnets must not share a cache key")
|
||||
}
|
||||
if subnetA != subnetAHost {
|
||||
t.Fatal("same subnet (different host bits) must share a cache key")
|
||||
}
|
||||
if subnetA == ipv4 || subnetA == noECS || ipv4 == noECS {
|
||||
t.Fatal("different families / no-ECS must not collide")
|
||||
}
|
||||
}
|
||||
|
||||
// TestLRUCache_ECSZeroPrefixDistinctFromNoECS is the regression test for the /0-vs-no-ECS
|
||||
// collision (RFC 7871 §7.3.1). A query carrying an ECS /0 option is forwarded WITH ECS and
|
||||
// may draw different upstream data, so its cached answer must never be served to a client
|
||||
// that sent no ECS at all, nor may IPv4 /0 and IPv6 /0 cross-serve each other.
|
||||
func TestLRUCache_ECSZeroPrefixDistinctFromNoECS(t *testing.T) {
|
||||
c, err := NewLRUCache(16)
|
||||
if err != nil {
|
||||
t.Fatalf("NewLRUCache: %v", err)
|
||||
}
|
||||
const up = "https://dns.example/dns-query"
|
||||
expire := time.Now().Add(time.Minute)
|
||||
|
||||
noECS := msgWithECS("controld.com", 0, 0, "")
|
||||
zeroV4 := msgWithECS("controld.com", 1, 0, "0.0.0.0")
|
||||
zeroV6 := msgWithECS("controld.com", 2, 0, "::")
|
||||
|
||||
// Only the IPv4 /0 query's answer is cached.
|
||||
c.Add(NewKey(zeroV4, up), NewValue(answerWithA("controld.com", "192.0.2.1"), expire))
|
||||
|
||||
// A no-ECS client must NOT receive the ECS /0 cached answer.
|
||||
if got := c.Get(NewKey(noECS, up)); got != nil {
|
||||
t.Fatalf("no-ECS client received the ECS /0 cached record %q", firstA(got.Msg))
|
||||
}
|
||||
// An IPv6 /0 client must NOT receive the IPv4 /0 answer.
|
||||
if got := c.Get(NewKey(zeroV6, up)); got != nil {
|
||||
t.Fatalf("IPv6 /0 client received the IPv4 /0 cached record %q", firstA(got.Msg))
|
||||
}
|
||||
// The IPv4 /0 client still hits its own entry.
|
||||
if got := c.Get(NewKey(zeroV4, up)); got == nil || firstA(got.Msg) != "192.0.2.1" {
|
||||
t.Fatalf("IPv4 /0 lost its own cached record: %v", got)
|
||||
}
|
||||
|
||||
// The no-ECS partition is independent and does not corrupt the /0 entry.
|
||||
c.Add(NewKey(noECS, up), NewValue(answerWithA("controld.com", "198.51.100.1"), expire))
|
||||
if got := c.Get(NewKey(noECS, up)); got == nil || firstA(got.Msg) != "198.51.100.1" {
|
||||
t.Fatalf("no-ECS partition wrong record: %v", got)
|
||||
}
|
||||
if got := c.Get(NewKey(zeroV4, up)); got == nil || firstA(got.Msg) != "192.0.2.1" {
|
||||
t.Fatalf("IPv4 /0 record corrupted by no-ECS insert: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLRUCache_ECSNoCrossSubnetServe is the real cache-path regression test for #564:
|
||||
// an entry populated for subnet A must never be returned to a client in subnet B, even
|
||||
// though the question (name/type/class/upstream) is identical. The two subnets carry
|
||||
// different A records; the second client must get its own record or a miss, never A's.
|
||||
func TestLRUCache_ECSNoCrossSubnetServe(t *testing.T) {
|
||||
c, err := NewLRUCache(16)
|
||||
if err != nil {
|
||||
t.Fatalf("NewLRUCache: %v", err)
|
||||
}
|
||||
const up = "https://dns.example/dns-query"
|
||||
expire := time.Now().Add(time.Minute)
|
||||
|
||||
reqA := msgWithECS("controld.com", 2, 64, "2001:db8:1::")
|
||||
reqB := msgWithECS("controld.com", 2, 64, "2001:db8:2::")
|
||||
|
||||
// Only subnet A's answer (A record 192.0.2.1) is cached.
|
||||
c.Add(NewKey(reqA, up), NewValue(answerWithA("controld.com", "192.0.2.1"), expire))
|
||||
|
||||
// A client in subnet B must NOT hit subnet A's entry.
|
||||
if got := c.Get(NewKey(reqB, up)); got != nil {
|
||||
t.Fatalf("subnet B received subnet A's cached record %q; cache is not ECS-partitioned", firstA(got.Msg))
|
||||
}
|
||||
|
||||
// Subnet A still hits its own entry.
|
||||
if got := c.Get(NewKey(reqA, up)); got == nil || firstA(got.Msg) != "192.0.2.1" {
|
||||
t.Fatalf("subnet A lost its own cached record: %+v", got)
|
||||
}
|
||||
|
||||
// Now cache subnet B's distinct answer and confirm the two never cross.
|
||||
c.Add(NewKey(reqB, up), NewValue(answerWithA("controld.com", "198.51.100.1"), expire))
|
||||
if got := c.Get(NewKey(reqB, up)); got == nil || firstA(got.Msg) != "198.51.100.1" {
|
||||
t.Fatalf("subnet B got the wrong record: %v", got)
|
||||
}
|
||||
if got := c.Get(NewKey(reqA, up)); got == nil || firstA(got.Msg) != "192.0.2.1" {
|
||||
t.Fatalf("subnet A record corrupted by subnet B insert: %v", got)
|
||||
}
|
||||
|
||||
// A second host within subnet A shares the partition (cache hit with A's record).
|
||||
reqAHost := msgWithECS("controld.com", 2, 64, "2001:db8:1::9")
|
||||
if got := c.Get(NewKey(reqAHost, up)); got == nil || firstA(got.Msg) != "192.0.2.1" {
|
||||
t.Fatalf("same-subnet host missed the shared cache entry: %v", got)
|
||||
}
|
||||
}
|
||||
+135
-3
@@ -157,6 +157,101 @@ type parallelDialerResult struct {
|
||||
err error
|
||||
}
|
||||
|
||||
const (
|
||||
// unreachableBackoffBase is the initial suppression window applied to a
|
||||
// dial address after it returns a network-unreachable error (e.g.
|
||||
// "connect: no route to host"). The window grows exponentially up to
|
||||
// unreachableBackoffMax on repeated failures, and is cleared as soon as
|
||||
// the address dials successfully.
|
||||
unreachableBackoffBase = 5 * time.Second
|
||||
// unreachableBackoffMax caps the suppression window so an address is
|
||||
// always re-probed within a bounded interval, preserving recovery when
|
||||
// the route comes back.
|
||||
unreachableBackoffMax = 60 * time.Second
|
||||
)
|
||||
|
||||
// Windows winsock codes for the unreachable errnos. A failing connect on
|
||||
// Windows surfaces these raw WSA codes (WSAENETUNREACH/WSAEHOSTUNREACH), whereas
|
||||
// syscall.ENETUNREACH/EHOSTUNREACH are Go's portable "invented" values
|
||||
// (APPLICATION_ERROR + iota) that never equal them. Matching these explicitly is
|
||||
// therefore required for the classifier to detect unreachable errors on Windows;
|
||||
// errors.Is against the syscall.* constants alone would not.
|
||||
//
|
||||
// https://learn.microsoft.com/en-us/windows/win32/winsock/windows-sockets-error-codes-2
|
||||
var (
|
||||
windowsENETUNREACH = syscall.Errno(10051)
|
||||
windowsEHOSTUNREACH = syscall.Errno(10065)
|
||||
)
|
||||
|
||||
// IsUnreachable reports whether err indicates the destination network or host
|
||||
// has no route (ENETUNREACH/EHOSTUNREACH). These are the errors produced when
|
||||
// an endpoint's address family is available locally but unroutable, e.g. an
|
||||
// IPv6 DoH endpoint while the host has IPv6 but no route to it.
|
||||
func IsUnreachable(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
var opErr *net.OpError
|
||||
if errors.As(err, &opErr) {
|
||||
return errors.Is(opErr.Err, syscall.ENETUNREACH) ||
|
||||
errors.Is(opErr.Err, syscall.EHOSTUNREACH) ||
|
||||
errors.Is(opErr.Err, windowsENETUNREACH) ||
|
||||
errors.Is(opErr.Err, windowsEHOSTUNREACH)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
type unreachableEntry struct {
|
||||
until time.Time
|
||||
backoff time.Duration
|
||||
}
|
||||
|
||||
// unreachableTracker records dial addresses that recently failed with a
|
||||
// network-unreachable error so ParallelDialer can temporarily stop hammering
|
||||
// them. This prevents an unroutable endpoint from generating a sustained dial
|
||||
// /health-check storm, while still re-probing each address once its bounded
|
||||
// backoff window expires so genuine recovery is never permanently blocked.
|
||||
type unreachableTracker struct {
|
||||
mu sync.Mutex
|
||||
entries map[string]unreachableEntry
|
||||
}
|
||||
|
||||
var unreachable = &unreachableTracker{entries: make(map[string]unreachableEntry)}
|
||||
|
||||
// suppressed reports whether addr is currently within its unreachable backoff
|
||||
// window and should be skipped.
|
||||
func (t *unreachableTracker) suppressed(addr string, now time.Time) bool {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
e, ok := t.entries[addr]
|
||||
return ok && now.Before(e.until)
|
||||
}
|
||||
|
||||
// markUnreachable extends the suppression window for addr using a bounded
|
||||
// exponential backoff.
|
||||
func (t *unreachableTracker) markUnreachable(addr string, now time.Time) {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
e := t.entries[addr]
|
||||
if e.backoff == 0 {
|
||||
e.backoff = unreachableBackoffBase
|
||||
} else {
|
||||
e.backoff *= 2
|
||||
if e.backoff > unreachableBackoffMax {
|
||||
e.backoff = unreachableBackoffMax
|
||||
}
|
||||
}
|
||||
e.until = now.Add(e.backoff)
|
||||
t.entries[addr] = e
|
||||
}
|
||||
|
||||
// markReachable clears any suppression for addr after a successful dial.
|
||||
func (t *unreachableTracker) markReachable(addr string) {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
delete(t.entries, addr)
|
||||
}
|
||||
|
||||
type ParallelDialer struct {
|
||||
net.Dialer
|
||||
}
|
||||
@@ -165,26 +260,63 @@ func (d *ParallelDialer) DialContext(ctx context.Context, network string, addrs
|
||||
if len(addrs) == 0 {
|
||||
return nil, errors.New("empty addresses")
|
||||
}
|
||||
|
||||
// Skip addresses that recently returned a network-unreachable error so an
|
||||
// unroutable endpoint (e.g. an IPv6 DoH address while the host has IPv6
|
||||
// but no route to it) does not generate a sustained dial storm. Suppression
|
||||
// is bounded: once an address's backoff window expires it is re-probed, so
|
||||
// genuine recovery is preserved.
|
||||
now := time.Now()
|
||||
live := make([]string, 0, len(addrs))
|
||||
var suppressed int
|
||||
for _, addr := range addrs {
|
||||
if unreachable.suppressed(addr, now) {
|
||||
suppressed++
|
||||
continue
|
||||
}
|
||||
live = append(live, addr)
|
||||
}
|
||||
if len(live) == 0 {
|
||||
// Every candidate is within its unreachable backoff window. Fail fast
|
||||
// and quietly instead of re-dialing known-unroutable addresses; the
|
||||
// windows expire and re-probe shortly, so recovery still happens.
|
||||
logger.Debug().Msgf("skipping %d unreachable address(es), all in backoff", suppressed)
|
||||
// TODO: the errno here is hardcoded to EHOSTUNREACH, but the actual
|
||||
// failure that triggered suppression may have been ENETUNREACH. This is
|
||||
// harmless today (IsUnreachable treats both the same and nothing else
|
||||
// inspects the errno), but if these errors are ever recorded/reported we
|
||||
// should retain the real error in unreachableEntry and surface it here.
|
||||
return nil, &net.OpError{Op: "dial", Net: network, Err: syscall.EHOSTUNREACH}
|
||||
}
|
||||
if suppressed > 0 {
|
||||
logger.Debug().Msgf("skipping %d unreachable address(es) in backoff", suppressed)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
defer cancel()
|
||||
|
||||
done := make(chan struct{})
|
||||
defer close(done)
|
||||
ch := make(chan *parallelDialerResult, len(addrs))
|
||||
ch := make(chan *parallelDialerResult, len(live))
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(len(addrs))
|
||||
wg.Add(len(live))
|
||||
go func() {
|
||||
wg.Wait()
|
||||
close(ch)
|
||||
}()
|
||||
|
||||
for _, addr := range addrs {
|
||||
for _, addr := range live {
|
||||
go func(addr string) {
|
||||
defer wg.Done()
|
||||
logger.Debug().Msgf("dialing to %s", addr)
|
||||
conn, err := d.Dialer.DialContext(ctx, network, addr)
|
||||
if err != nil {
|
||||
logger.Debug().Msgf("failed to dial %s: %v", addr, err)
|
||||
if IsUnreachable(err) {
|
||||
unreachable.markUnreachable(addr, time.Now())
|
||||
}
|
||||
} else {
|
||||
unreachable.markReachable(addr)
|
||||
}
|
||||
select {
|
||||
case ch <- ¶llelDialerResult{conn: conn, err: err}:
|
||||
|
||||
@@ -2,10 +2,77 @@ package net
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestIsUnreachable(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
err error
|
||||
want bool
|
||||
}{
|
||||
{"nil", nil, false},
|
||||
{"enetunreach", &net.OpError{Op: "dial", Net: "tcp", Err: syscall.ENETUNREACH}, true},
|
||||
{"ehostunreach", &net.OpError{Op: "dial", Net: "tcp", Err: syscall.EHOSTUNREACH}, true},
|
||||
{"windows enetunreach", &net.OpError{Op: "dial", Net: "tcp", Err: windowsENETUNREACH}, true},
|
||||
{"windows ehostunreach", &net.OpError{Op: "dial", Net: "tcp", Err: windowsEHOSTUNREACH}, true},
|
||||
{"connection refused", &net.OpError{Op: "dial", Net: "tcp", Err: syscall.ECONNREFUSED}, false},
|
||||
{"not an opError", syscall.ENETUNREACH, false},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := IsUnreachable(tc.err); got != tc.want {
|
||||
t.Errorf("IsUnreachable(%v) = %v, want %v", tc.err, got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnreachableTracker(t *testing.T) {
|
||||
tr := &unreachableTracker{entries: make(map[string]unreachableEntry)}
|
||||
const addr = "[2606:1a40::22]:443"
|
||||
now := time.Unix(0, 0)
|
||||
|
||||
// Not suppressed before any failure.
|
||||
if tr.suppressed(addr, now) {
|
||||
t.Fatal("addr suppressed before any failure")
|
||||
}
|
||||
|
||||
// First failure suppresses for the base window.
|
||||
tr.markUnreachable(addr, now)
|
||||
if !tr.suppressed(addr, now.Add(unreachableBackoffBase-time.Millisecond)) {
|
||||
t.Fatal("addr not suppressed within base backoff window")
|
||||
}
|
||||
if tr.suppressed(addr, now.Add(unreachableBackoffBase)) {
|
||||
t.Fatal("addr still suppressed at end of base backoff window")
|
||||
}
|
||||
|
||||
// Backoff grows exponentially and is capped at the max.
|
||||
tr.markUnreachable(addr, now)
|
||||
if got := tr.entries[addr].backoff; got != 2*unreachableBackoffBase {
|
||||
t.Fatalf("backoff after second failure = %v, want %v", got, 2*unreachableBackoffBase)
|
||||
}
|
||||
for i := 0; i < 10; i++ {
|
||||
tr.markUnreachable(addr, now)
|
||||
}
|
||||
if got := tr.entries[addr].backoff; got != unreachableBackoffMax {
|
||||
t.Fatalf("backoff not capped: got %v, want %v", got, unreachableBackoffMax)
|
||||
}
|
||||
|
||||
// A successful dial clears suppression entirely.
|
||||
tr.markReachable(addr)
|
||||
if tr.suppressed(addr, now) {
|
||||
t.Fatal("addr still suppressed after markReachable")
|
||||
}
|
||||
if _, ok := tr.entries[addr]; ok {
|
||||
t.Fatal("entry not removed after markReachable")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProbeStackTimeout(t *testing.T) {
|
||||
done := make(chan struct{})
|
||||
started := make(chan struct{})
|
||||
|
||||
@@ -2,11 +2,11 @@ package dnsmasq
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"html/template"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"text/template"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
+41
-15
@@ -10,7 +10,7 @@ import (
|
||||
"io"
|
||||
"net"
|
||||
"os/exec"
|
||||
"regexp"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -25,6 +25,12 @@ func dnsFns() []dnsFn {
|
||||
func getDNSFromScutil() []string {
|
||||
logger := *ProxyLogger.Load()
|
||||
|
||||
// Skip scutil on mobile platforms - not available in sandbox
|
||||
if isMobile() {
|
||||
Log(context.Background(), logger.Debug(), "skipping scutil DNS discovery on mobile platform")
|
||||
return nil
|
||||
}
|
||||
|
||||
const (
|
||||
maxRetries = 10
|
||||
retryInterval = 100 * time.Millisecond
|
||||
@@ -89,24 +95,34 @@ func getDNSFromScutil() []string {
|
||||
}
|
||||
|
||||
func getDHCPNameservers(iface string) ([]string, error) {
|
||||
// Run the ipconfig command for the given interface.
|
||||
cmd := exec.Command("ipconfig", "getpacket", iface)
|
||||
output, err := cmd.Output()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("error running ipconfig: %v", err)
|
||||
// Skip ipconfig on mobile platforms - not available in sandbox
|
||||
if isMobile() {
|
||||
return nil, fmt.Errorf("ipconfig not available on mobile")
|
||||
}
|
||||
|
||||
// Look for a line like:
|
||||
// domain_name_servers = 192.168.1.1 8.8.8.8;
|
||||
re := regexp.MustCompile(`domain_name_servers\s*=\s*(.*);`)
|
||||
matches := re.FindStringSubmatch(string(output))
|
||||
if len(matches) < 2 {
|
||||
return nil, fmt.Errorf("no DHCP nameservers found")
|
||||
// getoption returns the selected interface's DHCP option directly and does
|
||||
// not expose unrelated packet addresses to the parser.
|
||||
output, err := exec.Command("ipconfig", "getoption", iface, "domain_name_server").Output()
|
||||
if err == nil {
|
||||
return parseDHCPOptionNameservers(output), nil
|
||||
}
|
||||
|
||||
// Split the nameservers by whitespace.
|
||||
nameservers := strings.Fields(matches[1])
|
||||
return nameservers, nil
|
||||
// Older macOS releases can fail getoption while still exposing the packet.
|
||||
// Parse the real macOS field shape, for example:
|
||||
// domain_name_server (ip_mult): {192.168.1.1, 8.8.8.8}
|
||||
output, packetErr := exec.Command("ipconfig", "getpacket", iface).Output()
|
||||
if packetErr != nil {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("error reading DHCP DNS option: getoption: %v; getpacket: %v", err, packetErr)
|
||||
}
|
||||
return nil, fmt.Errorf("error reading DHCP packet: %v", packetErr)
|
||||
}
|
||||
return parseDHCPPacketNameservers(output), nil
|
||||
}
|
||||
|
||||
// DHCPNameserversForInterface returns DHCP option 6 for exactly iface.
|
||||
func DHCPNameserversForInterface(iface string) ([]string, error) {
|
||||
return getDHCPNameservers(iface)
|
||||
}
|
||||
|
||||
func getAllDHCPNameservers() []string {
|
||||
@@ -201,6 +217,11 @@ func getAllDHCPNameservers() []string {
|
||||
}
|
||||
|
||||
func patchNetIfaceName(iface *net.Interface) (bool, error) {
|
||||
// Skip networksetup on mobile platforms - not available in sandbox
|
||||
if isMobile() {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
b, err := exec.Command("networksetup", "-listnetworkserviceorder").Output()
|
||||
if err != nil {
|
||||
return false, err
|
||||
@@ -234,3 +255,8 @@ func networkServiceName(ifaceName string, r io.Reader) string {
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// isMobile reports whether the current OS is a mobile platform.
|
||||
func isMobile() bool {
|
||||
return runtime.GOOS == "ios"
|
||||
}
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
package ctrld
|
||||
|
||||
import (
|
||||
"net"
|
||||
"strings"
|
||||
"unicode"
|
||||
)
|
||||
|
||||
func parseDHCPOptionNameservers(output []byte) []string {
|
||||
return parseIPv4Nameservers(string(output))
|
||||
}
|
||||
|
||||
func parseDHCPPacketNameservers(output []byte) []string {
|
||||
for _, line := range strings.Split(string(output), "\n") {
|
||||
field := strings.TrimSpace(line)
|
||||
if strings.HasPrefix(field, "domain_name_server ") ||
|
||||
strings.HasPrefix(field, "domain_name_server:") ||
|
||||
strings.HasPrefix(field, "domain_name_servers ") ||
|
||||
strings.HasPrefix(field, "domain_name_servers:") {
|
||||
return parseIPv4Nameservers(field)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseIPv4Nameservers(value string) []string {
|
||||
seen := make(map[string]struct{})
|
||||
var nameservers []string
|
||||
for _, token := range strings.FieldsFunc(value, func(r rune) bool {
|
||||
return r != '.' && !unicode.IsDigit(r)
|
||||
}) {
|
||||
ip := net.ParseIP(token)
|
||||
if ip == nil || ip.To4() == nil {
|
||||
continue
|
||||
}
|
||||
ns := ip.String()
|
||||
if _, ok := seen[ns]; ok {
|
||||
continue
|
||||
}
|
||||
seen[ns] = struct{}{}
|
||||
nameservers = append(nameservers, ns)
|
||||
}
|
||||
return nameservers
|
||||
}
|
||||
@@ -0,0 +1,64 @@
|
||||
package ctrld
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestParseDHCPOptionNameservers(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
output string
|
||||
want []string
|
||||
}{
|
||||
{"single", "192.168.10.1\n", []string{"192.168.10.1"}},
|
||||
{"multiple", "192.168.10.1\n1.1.1.1\n", []string{"192.168.10.1", "1.1.1.1"}},
|
||||
{"deduplicate and reject invalid", "192.168.10.1 999.1.1.1 192.168.10.1", []string{"192.168.10.1"}},
|
||||
{"empty", "", nil},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := parseDHCPOptionNameservers([]byte(tc.output)); !slices.Equal(got, tc.want) {
|
||||
t.Fatalf("parseDHCPOptionNameservers() = %v, want %v", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseDHCPPacketNameservers(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
output string
|
||||
want []string
|
||||
}{
|
||||
{
|
||||
name: "macos singular ip_mult",
|
||||
output: `op = BOOTREPLY
|
||||
` +
|
||||
`yiaddr = 192.168.10.155
|
||||
` +
|
||||
`domain_name_server (ip_mult): {192.168.10.1, 1.1.1.1}
|
||||
` +
|
||||
`server_identifier (ip): 192.168.10.1
|
||||
`,
|
||||
want: []string{"192.168.10.1", "1.1.1.1"},
|
||||
},
|
||||
{
|
||||
name: "legacy plural equals",
|
||||
output: "domain_name_servers = 192.168.1.1 8.8.8.8;\n",
|
||||
want: []string{"192.168.1.1", "8.8.8.8"},
|
||||
},
|
||||
{
|
||||
name: "packet addresses without option are ignored",
|
||||
output: "yiaddr = 192.168.10.155\nserver_identifier (ip): 192.168.10.1\n",
|
||||
want: nil,
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := parseDHCPPacketNameservers([]byte(tc.output)); !slices.Equal(got, tc.want) {
|
||||
t.Fatalf("parseDHCPPacketNameservers() = %v, want %v", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user