mirror of
https://github.com/Control-D-Inc/ctrld.git
synced 2026-09-04 13:36:35 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
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 | ||
|
|
7e6f88b4ed | ||
|
|
5dd5846cca | ||
|
|
2b27c148be | ||
|
|
8cb383d87e | ||
|
|
afed925404 | ||
|
|
d1ea70d688 | ||
|
|
ed98104384 | ||
|
|
eaa171f66f | ||
|
|
839b8236e7 | ||
|
|
3f59cdad1a | ||
|
|
c55e2a722c | ||
|
|
22a796f673 | ||
|
|
95dd871e2d | ||
|
|
5c0585b2e8 | ||
|
|
112d1cb5a9 | ||
|
|
bd9bb90dd4 | ||
|
|
82fc628bf3 | ||
|
|
2926c76b76 | ||
|
|
fe08f00746 | ||
|
|
9be15aeec8 | ||
|
|
9b2e51f53a | ||
|
|
e7040bd9f9 | ||
|
|
768cc81855 | ||
|
|
289a46dc2c | ||
|
|
1e8240bd1c | ||
|
|
12715e6f24 | ||
|
|
147106f2b9 | ||
|
|
a4f0418811 | ||
|
|
40c68a13a1 | ||
|
|
3f30ec30d8 | ||
|
|
4790eb2c88 | ||
|
|
da3ea05763 | ||
|
|
209c9211b9 | ||
|
|
acbebcf7c2 | ||
|
|
2e8a0f00a0 | ||
|
|
1f4c47318e | ||
|
|
e8d1a4604e | ||
|
|
8d63a755ba | ||
|
|
f05519d1c8 | ||
|
|
1804e6db67 | ||
|
|
d0341497d1 | ||
|
|
27c5be43c2 | ||
|
|
3beffd0dc8 | ||
|
|
1f9c586444 | ||
|
|
a92e1ca024 | ||
|
|
705df72110 | ||
|
|
22122c45b2 | ||
|
|
57a9bb9fab | ||
|
|
78ea2d6361 | ||
|
|
df3cf7ef62 | ||
|
|
80e652b8d9 | ||
|
|
091c7edb19 | ||
|
|
6c550b1d74 | ||
|
|
3ca559e5a4 | ||
|
|
0e3f764299 | ||
|
|
e52402eb0c | ||
|
|
2133f31854 | ||
|
|
a198a5cd65 | ||
|
|
eb2b231bd2 | ||
|
|
7af29cfbc0 | ||
|
|
ce1a165348 | ||
|
|
fd48e6d795 | ||
|
|
d71d1341b6 | ||
|
|
21855df4af | ||
|
|
66e2d3a40a | ||
|
|
26257cf24a | ||
|
|
36a7423634 | ||
|
|
e616091249 | ||
|
|
0948161529 | ||
|
|
ce29b5d217 | ||
|
|
de24fa293e | ||
|
|
6663925c4d | ||
|
|
b9ece6d7b9 | ||
|
|
c4efa1ab97 | ||
|
|
7cea5305e1 | ||
|
|
a20fbf95de | ||
|
|
628c4302aa | ||
|
|
8dc34f8bf5 | ||
|
|
b4faf82f76 | ||
|
|
a983dfaee2 | ||
|
|
62f73bcaa2 | ||
|
|
00e9d2bdd3 | ||
|
|
ace3b1e66e | ||
|
|
d1ea1ba08c | ||
|
|
c06c8aa859 | ||
|
|
0c2cc00c4f | ||
|
|
8d6ea91f35 | ||
|
|
7dfb77228f | ||
|
|
24910f1fa6 | ||
|
|
433a61d2ee | ||
|
|
3937e885f0 | ||
|
|
c651003cc4 | ||
|
|
b7ccfcb8b4 | ||
|
|
a9ed70200b | ||
|
|
c6365e6b74 | ||
|
|
dacc67e50f | ||
|
|
c60cf33af3 | ||
|
|
f27cbe3525 | ||
|
|
2de1b9929a | ||
|
|
8bf654aece | ||
|
|
84376ed719 | ||
|
|
7a136b8874 | ||
|
|
58c0e4f15a | ||
|
|
e0d35d8ba2 | ||
|
|
3b2e48761e | ||
|
|
b27064008e | ||
|
|
1ad63827e1 | ||
|
|
20e61550c2 | ||
|
|
020b814402 | ||
|
|
e578867118 | ||
|
|
46a1039f21 | ||
|
|
cc9e27de5f | ||
|
|
6ab3ab9faf | ||
|
|
e68bfa795a | ||
|
|
e60a92e93e | ||
|
|
62fe14f76b | ||
|
|
a0c5062e3a | ||
|
|
49eb152d02 | ||
|
|
b05056423a | ||
|
|
c7168739c7 | ||
|
|
5b1faf1ce3 | ||
|
|
513a6f9ec7 | ||
|
|
8db6fa4232 | ||
|
|
5036de2602 | ||
|
|
332f8ccc37 | ||
|
|
a582195cec | ||
|
|
9fe36ae984 | ||
|
|
54cb455522 | ||
|
|
8bd3b9e474 | ||
|
|
eff5ff580b | ||
|
|
c45f863ed8 | ||
|
|
414d4e356d | ||
|
|
ef697eb781 | ||
|
|
0631ffe831 | ||
|
|
7444d8517a | ||
|
|
3480043e40 | ||
|
|
619b6e7516 | ||
|
|
0123ca44fb | ||
|
|
7929aafe2a | ||
|
|
dc433f8dc9 | ||
|
|
8ccaeeab60 | ||
|
|
043a28eb33 | ||
|
|
c329402f5d | ||
|
|
23e6ad6e1f | ||
|
|
e6de78c1fa | ||
|
|
a670708f93 | ||
|
|
4ebe2fb5f4 | ||
|
|
3403b2039d | ||
|
|
e30ad31e0f | ||
|
|
81e0bad739 | ||
|
|
7d07d738dc | ||
|
|
0fae584e65 | ||
|
|
9e83085f2a | ||
|
|
41a00c68ac | ||
|
|
e3b99bf339 | ||
|
|
5007a87d3a | ||
|
|
60e65a37a6 | ||
|
|
d37d0e942c | ||
|
|
98042d8dbd | ||
|
|
af4b826b68 | ||
|
|
253a57ca01 | ||
|
|
caf98b4dfe | ||
|
|
398f71fd00 | ||
|
|
e1301ade96 | ||
|
|
7a23f82192 | ||
|
|
715bcc4aa1 | ||
|
|
0c74838740 | ||
|
|
4b05b6da7b | ||
|
|
375844ff1a | ||
|
|
1d207379cb | ||
|
|
fb49cb71e3 | ||
|
|
9618efbcde | ||
|
|
bb2210b06a | ||
|
|
917052723d | ||
|
|
fef85cadeb | ||
|
|
4a05fb6b28 | ||
|
|
6644ce53f2 | ||
|
|
72f0b89fdc | ||
|
|
41a97a6609 | ||
|
|
38064d6ad5 | ||
|
|
ae6945cedf | ||
|
|
3132d1b032 | ||
|
|
2716ae29bd | ||
|
|
1c50c2b6af | ||
|
|
cf6d16b439 | ||
|
|
60686f55ff | ||
|
|
47d7ace3a7 | ||
|
|
2d3779ec27 | ||
|
|
595071b608 | ||
|
|
57ef717080 | ||
|
|
eb27d1482b | ||
|
|
f57972ead7 | ||
|
|
168eaf538b | ||
|
|
1560455ca3 | ||
|
|
028475a193 | ||
|
|
f7a6dbe39b | ||
|
|
e573a490c9 | ||
|
|
ce3281e70d | ||
|
|
0fbfd160c9 | ||
|
|
20759017e6 | ||
|
|
69e0aab73e | ||
|
|
7ed6733fb7 | ||
|
|
9718ab8579 | ||
|
|
2687a4a018 | ||
|
|
2d9c60dea1 | ||
|
|
841be069b7 | ||
|
|
7833132917 | ||
|
|
e9e63b0983 | ||
|
|
4df470b869 | ||
|
|
89600f6091 | ||
|
|
f986a575e8 | ||
|
|
9c2fe8d21f | ||
|
|
8bcbb9249e | ||
|
|
a95d50c0af | ||
|
|
5db7d3577b | ||
|
|
c53a0ca1c4 | ||
|
|
6fd3d1788a | ||
|
|
087c1975e5 | ||
|
|
3713cbecc3 | ||
|
|
6046789fa4 | ||
|
|
3ea69b180c | ||
|
|
db6e977e3a | ||
|
|
a5c776c846 | ||
|
|
5a566c028a | ||
|
|
ff43c74d8d | ||
|
|
3c7255569c | ||
|
|
4a92ec4d2d | ||
|
|
9bbccb4082 | ||
|
|
4f62314646 | ||
|
|
cb49d0d947 | ||
|
|
89f7874fc6 | ||
|
|
221917e80b | ||
|
|
37d41bd215 | ||
|
|
8a96b8bec4 | ||
|
|
02ee113b95 | ||
|
|
f71dd78915 | ||
|
|
cd5619a05b | ||
|
|
a63a30c76b | ||
|
|
f5ba8be182 | ||
|
|
a9f76322bd | ||
|
|
ed39269c80 | ||
|
|
09426dcd36 | ||
|
|
17941882a9 | ||
|
|
70ab8032a0 | ||
|
|
8360bdc50a | ||
|
|
6837176ec7 | ||
|
|
5e9b4244e7 | ||
|
|
9b6a308958 | ||
|
|
71e327653a | ||
|
|
a56711796f | ||
|
|
09495f2a7c | ||
|
|
484643e114 | ||
|
|
da91aabc35 | ||
|
|
c654398981 | ||
|
|
47a90ec2a1 | ||
|
|
2875e22d0b | ||
|
|
c5d14e0075 | ||
|
|
84e06c363c |
@@ -9,18 +9,18 @@ jobs:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
os: ["windows-latest", "ubuntu-latest", "macOS-latest"]
|
||||
go: ["1.23.x"]
|
||||
go: ["1.26.x"]
|
||||
runs-on: ${{ matrix.os }}
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
with:
|
||||
fetch-depth: 1
|
||||
- uses: WillAbides/setup-go-faster@v1.8.0
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: ${{ matrix.go }}
|
||||
- run: "go test -race ./..."
|
||||
- uses: dominikh/staticcheck-action@v1.2.0
|
||||
- uses: dominikh/staticcheck-action@v1.4.1
|
||||
with:
|
||||
version: "2024.1.1"
|
||||
version: "2026.2"
|
||||
install-go: false
|
||||
cache-key: ${{ matrix.go }}
|
||||
|
||||
@@ -12,3 +12,5 @@ ctrld-*
|
||||
|
||||
# generated file
|
||||
cmd/cli/rsrc_*.syso
|
||||
ctrld
|
||||
ctrld.exe
|
||||
|
||||
@@ -4,12 +4,12 @@
|
||||
[](https://pkg.go.dev/github.com/Control-D-Inc/ctrld)
|
||||
[](https://goreportcard.com/report/github.com/Control-D-Inc/ctrld)
|
||||
|
||||

|
||||

|
||||
|
||||
A highly configurable DNS forwarding proxy with support for:
|
||||
- Multiple listeners for incoming queries
|
||||
- Multiple upstreams with fallbacks
|
||||
- Multiple network policy driven DNS query steering
|
||||
- Multiple network policy driven DNS query steering (via network cidr, MAC address or FQDN)
|
||||
- Policy driven domain based "split horizon" DNS with wildcard support
|
||||
- Integrations with common router vendors and firmware
|
||||
- LAN client discovery via DHCP, mDNS, ARP, NDP, hosts file parsing
|
||||
@@ -35,13 +35,29 @@ All DNS protocols are supported, including:
|
||||
|
||||
## OS Support
|
||||
- Windows (386, amd64, arm)
|
||||
- Mac (amd64, arm64)
|
||||
- Windows Server (386, amd64)
|
||||
- MacOS (amd64, arm64)
|
||||
- Linux (386, amd64, arm, mips)
|
||||
- FreeBSD
|
||||
- Common routers (See Router Mode below)
|
||||
- FreeBSD (386, amd64, arm)
|
||||
- Common routers (See below)
|
||||
|
||||
|
||||
### Supported Routers
|
||||
You can run `ctrld` on any supported router. The list of supported routers and firmware includes:
|
||||
- Asus Merlin
|
||||
- DD-WRT
|
||||
- Firewalla
|
||||
- FreshTomato
|
||||
- GL.iNet
|
||||
- OpenWRT
|
||||
- pfSense / OPNsense
|
||||
- Synology
|
||||
- Ubiquiti (UniFi, EdgeOS)
|
||||
|
||||
`ctrld` will attempt to interface with dnsmasq (or Windows Server) whenever possible and set itself as the upstream, while running on port 5354. On FreeBSD based OSes, `ctrld` will terminate dnsmasq and unbound in order to be able to listen on port 53 directly.
|
||||
|
||||
# Install
|
||||
There are several ways to download and install `ctrld.
|
||||
There are several ways to download and install `ctrld`.
|
||||
|
||||
## Quick Install
|
||||
The simplest way to download and install `ctrld` is to use the following installer command on any UNIX-like platform:
|
||||
@@ -50,14 +66,14 @@ The simplest way to download and install `ctrld` is to use the following install
|
||||
sh -c 'sh -c "$(curl -sL https://api.controld.com/dl)"'
|
||||
```
|
||||
|
||||
Windows user and prefer Powershell (who doesn't)? No problem, execute this command instead in administrative cmd:
|
||||
Windows user and prefer Powershell (who doesn't)? No problem, execute this command instead in administrative PowerShell:
|
||||
```shell
|
||||
powershell -Command "(Invoke-WebRequest -Uri 'https://api.controld.com/dl' -UseBasicParsing).Content | Set-Content 'ctrld_install.bat'" && ctrld_install.bat
|
||||
(Invoke-WebRequest -Uri 'https://api.controld.com/dl/ps1' -UseBasicParsing).Content | Set-Content "$env:TEMPctrld_install.ps1"; Invoke-Expression "& '$env:TEMPctrld_install.ps1'"
|
||||
```
|
||||
|
||||
Or you can pull and run a Docker container from [Docker Hub](https://hub.docker.com/r/controldns/ctrld)
|
||||
```
|
||||
$ docker pull controldns/ctrld
|
||||
```shell
|
||||
docker run -d --name=ctrld -p 127.0.0.1:53:53/tcp -p 127.0.0.1:53:53/udp controldns/ctrld:latest
|
||||
```
|
||||
|
||||
## Download Manually
|
||||
@@ -67,25 +83,24 @@ Alternatively, if you know what you're doing you can download pre-compiled binar
|
||||
Lastly, you can build `ctrld` from source which requires `go1.21+`:
|
||||
|
||||
```shell
|
||||
$ go build ./cmd/ctrld
|
||||
go build ./cmd/ctrld
|
||||
```
|
||||
|
||||
or
|
||||
|
||||
```shell
|
||||
$ go install github.com/Control-D-Inc/ctrld/cmd/ctrld@latest
|
||||
go install github.com/Control-D-Inc/ctrld/cmd/ctrld@latest
|
||||
```
|
||||
|
||||
or
|
||||
|
||||
```
|
||||
$ docker build -t controldns/ctrld . -f docker/Dockerfile
|
||||
$ docker run -d --name=ctrld -p 53:53/tcp -p 53:53/udp controldns/ctrld --cd=RESOLVER_ID_GOES_HERE -vv
|
||||
```shell
|
||||
docker build -t controldns/ctrld . -f docker/Dockerfile
|
||||
```
|
||||
|
||||
|
||||
# Usage
|
||||
The cli is self documenting, so free free to run `--help` on any sub-command to get specific usages.
|
||||
The cli is self documenting, so feel free to run `--help` on any sub-command to get specific usages.
|
||||
|
||||
## Arguments
|
||||
```
|
||||
@@ -101,15 +116,16 @@ Usage:
|
||||
|
||||
Available Commands:
|
||||
run Run the DNS proxy server
|
||||
service Manage ctrld service
|
||||
start Quick start service and configure DNS on interface
|
||||
stop Quick stop service and remove DNS from interface
|
||||
restart Restart the ctrld service
|
||||
reload Reload the ctrld service
|
||||
status Show status of the ctrld service
|
||||
uninstall Stop and uninstall the ctrld service
|
||||
service Manage ctrld service
|
||||
clients Manage clients
|
||||
upgrade Upgrading ctrld to latest version
|
||||
log Manage runtime debug logs
|
||||
|
||||
Flags:
|
||||
-h, --help help for ctrld
|
||||
@@ -121,81 +137,99 @@ Use "ctrld [command] --help" for more information about a command.
|
||||
```
|
||||
|
||||
## Basic Run Mode
|
||||
To start the server with default configuration, simply run: `./ctrld run`. This will create a generic `ctrld.toml` file in the **working directory** and start the application in foreground.
|
||||
1. Start the server
|
||||
```
|
||||
$ sudo ./ctrld run
|
||||
This is the most basic way to run `ctrld`, in foreground mode. Unless you already have a config file, a default one will be generated.
|
||||
|
||||
### Command
|
||||
|
||||
Windows (Admin Shell)
|
||||
```shell
|
||||
ctrld.exe run
|
||||
```
|
||||
|
||||
2. Run a test query using a DNS client, for example, `dig`:
|
||||
Linux or Macos
|
||||
```shell
|
||||
sudo ctrld run
|
||||
```
|
||||
|
||||
You can then run a test query using a DNS client, for example, `dig`:
|
||||
```
|
||||
$ dig verify.controld.com @127.0.0.1 +short
|
||||
api.controld.com.
|
||||
147.185.34.1
|
||||
```
|
||||
|
||||
If `verify.controld.com` resolves, you're successfully using the default Control D upstream. From here, you can start editing the config file and go nuts with it. To enforce a new config, restart the server.
|
||||
If `verify.controld.com` resolves, you're successfully using the default Control D upstream. From here, you can start editing the config file that was generated. To enforce a new config, restart the server.
|
||||
|
||||
## Service Mode
|
||||
To run the application in service mode on any Windows, MacOS, Linux distibution or supported router, simply run: `./ctrld start` as system/root user. This will create a generic `ctrld.toml` file in the **user home** directory (on Windows) or `/etc/controld/` (almost everywhere else), start the system service, and configure the listener on the default network interface. Service will start on OS boot.
|
||||
This mode will run the application as a background system service on any Windows, MacOS, Linux, FreeBSD distribution or supported router. This will create a generic `ctrld.toml` file in the **C:\ControlD** directory (on Windows) or `/etc/controld/` (almost everywhere else), start the system service, and **configure the listener on all physical network interface**. Service will start on OS boot.
|
||||
|
||||
When Control D upstreams are used, `ctrld` willl [relay your network topology](https://docs.controld.com/docs/device-clients) to Control D (LAN IPs, MAC addresses, and hostnames), and you will be able to see your LAN devices in the web panel, view analytics and apply unique profiles to them.
|
||||
When Control D upstreams are used on a router type device, `ctrld` will [relay your network topology](https://docs.controld.com/docs/device-clients) to Control D (LAN IPs, MAC addresses, and hostnames), and you will be able to see your LAN devices in the web panel, view analytics and apply unique profiles to them.
|
||||
|
||||
In order to stop the service, and restore your DNS to original state, simply run `./ctrld stop`. If you wish to stop and uninstall the service permanently, run `./ctrld uninstall`.
|
||||
### Command
|
||||
|
||||
Windows (Admin Shell)
|
||||
```shell
|
||||
ctrld.exe start
|
||||
```
|
||||
|
||||
### Supported Routers
|
||||
You can run `ctrld` on any supported router, which will function similarly to the Service Mode mentioned above. The list of supported routers and firmware includes:
|
||||
- Asus Merlin
|
||||
- DD-WRT
|
||||
- Firewalla
|
||||
- FreshTomato
|
||||
- GL.iNet
|
||||
- OpenWRT
|
||||
- pfSense / OPNsense
|
||||
- Synology
|
||||
- Ubiquiti (UniFi, EdgeOS)
|
||||
Linux or Macos
|
||||
```
|
||||
sudo ctrld start
|
||||
```
|
||||
|
||||
`ctrld` will attempt to interface with dnsmasq whenever possible and set itself as the upstream, while running on port 5354. On FreeBSD based OSes, `ctrld` will terminate dnsmasq and unbound in order to be able to listen on port 53 directly.
|
||||
If `ctrld` is not in your system path (you installed it manually), you will need to run the above commands from the directory where you installed `ctrld`.
|
||||
|
||||
In order to stop the service, and restore your DNS to original state, simply run `ctrld stop`. If you wish to stop and uninstall the service permanently, run `ctrld uninstall`.
|
||||
|
||||
### Control D Auto Configuration
|
||||
Application can be started with a specific resolver config, instead of the default one. Simply supply your Resolver ID with a `--cd` flag, when using the `run` (foreground) or `start` (service) modes.
|
||||
## Unmanaged Service Mode
|
||||
This mode functions similarly to the "Service Mode" above except it will simply start a system service and the config defined listeners, but **will not make any changes to any network interfaces**. You can then set the `ctrld` listener(s) IP on the desired network interfaces manually.
|
||||
|
||||
The following command will start the application in foreground mode, using the free "p2" resolver, which blocks Ads & Trackers.
|
||||
### Command
|
||||
|
||||
```shell
|
||||
./ctrld run --cd p2
|
||||
```
|
||||
Windows (Admin Shell)
|
||||
```shell
|
||||
ctrld.exe service start
|
||||
```
|
||||
|
||||
Alternatively, you can use your own personal Control D Device resolver, and start the application in service mode. Your resolver ID is displayed on the "Show Resolvers" screen for the relevant Control D Device.
|
||||
|
||||
```shell
|
||||
./ctrld start --cd abcd1234
|
||||
```
|
||||
|
||||
Once you run the above commands (in service mode only), the following things will happen:
|
||||
- You resolver configuration will be fetched from the API, and config file templated with the resolver data
|
||||
- Application will start as a service, and keep running (even after reboot) until you run the `stop` or `uninstall` sub-commands
|
||||
- Your default network interface will be updated to use the listener started by the service
|
||||
- All OS DNS queries will be sent to the listener
|
||||
Linux or Macos
|
||||
```shell
|
||||
sudo ctrld service start
|
||||
```
|
||||
|
||||
# Configuration
|
||||
See [Configuration Docs](docs/config.md).
|
||||
`ctrld` can be configured in variety of different ways, which include: API, local config file or via cli launch args.
|
||||
|
||||
## Example
|
||||
- Start `listener.0` on 127.0.0.1:53
|
||||
- Accept queries from any source address
|
||||
- Send all queries to `upstream.0` via DoH protocol
|
||||
## API Based Auto Configuration
|
||||
Application can be started with a specific Control D resolver config, instead of the default one. Simply supply your Resolver ID with a `--cd` flag, when using the `start` (service) mode. In this mode, the application will automatically choose a non-conflicting IP and/or port and configure itself as the upstream to whatever process is running on port 53 (like dnsmasq or Windows DNS Server). This mode is used when the 1 liner installer command from the Control D onboarding guide is executed.
|
||||
|
||||
### Default Config
|
||||
The following command will use your own personal Control D Device resolver, and start the application in service mode. Your resolver ID is displayed on the "Show Resolvers" screen for the relevant Control D Endpoint.
|
||||
|
||||
Windows (Admin Shell)
|
||||
```shell
|
||||
ctrld.exe start --cd abcd1234
|
||||
```
|
||||
|
||||
Linux or Macos
|
||||
```shell
|
||||
sudo ctrld start --cd abcd1234
|
||||
```
|
||||
|
||||
Once you run the above command, the following things will happen:
|
||||
- You resolver configuration will be fetched from the API, and config file templated with the resolver data
|
||||
- Application will start as a service, and keep running (even after reboot) until you run the `stop` or `uninstall` sub-commands
|
||||
- All physical network interface will be updated to use the listener started by the service or dnsmasq upstream will be switched to `ctrld`
|
||||
- All DNS queries will be sent to the listener
|
||||
|
||||
## Manual Configuration
|
||||
`ctrld` is entirely config driven and can be configured in many different ways, please see [Configuration Docs](docs/config.md).
|
||||
|
||||
### Example
|
||||
```toml
|
||||
[listener]
|
||||
|
||||
[listener.0]
|
||||
ip = ""
|
||||
port = 0
|
||||
restricted = false
|
||||
ip = '0.0.0.0'
|
||||
port = 53
|
||||
|
||||
[network]
|
||||
|
||||
@@ -203,10 +237,6 @@ See [Configuration Docs](docs/config.md).
|
||||
cidrs = ["0.0.0.0/0"]
|
||||
name = "Network 0"
|
||||
|
||||
[service]
|
||||
log_level = "info"
|
||||
log_path = ""
|
||||
|
||||
[upstream]
|
||||
|
||||
[upstream.0]
|
||||
@@ -215,28 +245,88 @@ See [Configuration Docs](docs/config.md).
|
||||
name = "Control D - Anti-Malware"
|
||||
timeout = 5000
|
||||
type = "doh"
|
||||
|
||||
[upstream.1]
|
||||
bootstrap_ip = "76.76.2.11"
|
||||
endpoint = "p2.freedns.controld.com"
|
||||
name = "Control D - No Ads"
|
||||
timeout = 3000
|
||||
type = "doq"
|
||||
|
||||
```
|
||||
|
||||
`ctrld` will pick a working config for `listener.0` then writing the default config to disk for the first run.
|
||||
The above basic config will:
|
||||
- Start listener on 0.0.0.0:53
|
||||
- Accept queries from any source address
|
||||
- Send all queries to `https://freedns.controld.com/p1` using DoH protocol
|
||||
|
||||
## Advanced Configuration
|
||||
The above is the most basic example, which will work out of the box. If you're looking to do advanced configurations using policies, see [Configuration Docs](docs/config.md) for complete documentation of the config file.
|
||||
## CLI Args
|
||||
If you're unable to use a config file, `ctrld` can be be supplied with basic configuration via launch arguments, in [Ephemeral Mode](docs/ephemeral_mode.md).
|
||||
|
||||
You can also supply configuration via launch argeuments, in [Ephemeral Mode](docs/ephemeral_mode.md).
|
||||
### Example
|
||||
```
|
||||
ctrld run --listen=127.0.0.1:53 --primary_upstream=https://freedns.controld.com/p2 --secondary_upstream=10.0.10.1:53 --domains=*.company.int,very-secure.local --log /path/to/log.log
|
||||
```
|
||||
|
||||
The above will start a foreground process and:
|
||||
- Listen on `127.0.0.1:53` for DNS queries
|
||||
- Forward all queries to `https://freedns.controld.com/p2` using DoH protocol, while...
|
||||
- Excluding `*.company.int` and `very-secure.local` matching queries, that are forwarded to `10.0.10.1:53`
|
||||
- Write a debug log to `/path/to/log.log`
|
||||
|
||||
## DNS Intercept Mode
|
||||
When running `ctrld` alongside VPN software, DNS conflicts can cause intermittent failures, bypassed filtering, or configuration loops. DNS Intercept Mode prevents these issues by transparently capturing all DNS traffic on the system and routing it through `ctrld`, without modifying network adapter DNS settings.
|
||||
|
||||
### When to Use
|
||||
Enable DNS Intercept Mode if you:
|
||||
- Use corporate VPN software (F5, Cisco AnyConnect, Palo Alto GlobalProtect, Zscaler)
|
||||
- Run overlay networks like Tailscale or WireGuard
|
||||
- Experience random DNS failures when VPN connects/disconnects
|
||||
- See gaps in your Control D analytics when VPN is active
|
||||
- Have endpoint security software that also manages DNS
|
||||
|
||||
### Command
|
||||
|
||||
Windows (Admin Shell)
|
||||
```shell
|
||||
ctrld.exe start --intercept-mode dns --cd RESOLVER_ID_HERE
|
||||
```
|
||||
|
||||
macOS
|
||||
```shell
|
||||
sudo ctrld start --intercept-mode dns --cd RESOLVER_ID_HERE
|
||||
```
|
||||
|
||||
`--intercept-mode dns` automatically detects VPN internal domains and routes them to the VPN's DNS server, while Control D handles everything else.
|
||||
|
||||
To disable intercept mode on a service that already has it enabled:
|
||||
|
||||
Windows (Admin Shell)
|
||||
```shell
|
||||
ctrld.exe start --intercept-mode off
|
||||
```
|
||||
|
||||
macOS
|
||||
```shell
|
||||
sudo ctrld start --intercept-mode off
|
||||
```
|
||||
|
||||
This removes the intercept rules and reverts to standard interface-based DNS configuration.
|
||||
|
||||
### Platform Support
|
||||
| Platform | Supported | Mechanism |
|
||||
|----------|-----------|-----------|
|
||||
| Windows | ✅ | NRPT (Name Resolution Policy Table) |
|
||||
| macOS | ✅ | pf (packet filter) redirect |
|
||||
| Linux | ❌ | Not currently supported |
|
||||
|
||||
### Features
|
||||
- **VPN split routing** — VPN-specific domains are automatically detected and forwarded to the VPN's DNS server
|
||||
- **Captive portal recovery** — Wi-Fi login pages (hotels, airports, coffee shops) work automatically
|
||||
- **No network adapter changes** — DNS settings stay untouched, eliminating conflicts entirely
|
||||
- **Automatic port 53 conflict resolution** — if another process (e.g., `mDNSResponder` on macOS) is already using port 53, `ctrld` automatically listens on a different port. OS-level packet interception redirects all DNS traffic to `ctrld` transparently, so no manual configuration is needed. This only applies to intercept mode.
|
||||
|
||||
### Tested VPN Software
|
||||
- F5 BIG-IP APM
|
||||
- Cisco AnyConnect
|
||||
- Palo Alto GlobalProtect
|
||||
- Tailscale (including Exit Nodes)
|
||||
- Windscribe
|
||||
- WireGuard
|
||||
|
||||
For more details, see the [DNS Intercept Mode documentation](https://docs.controld.com/docs/dns-intercept).
|
||||
|
||||
## Contributing
|
||||
See [Contribution Guideline](./docs/contributing.md)
|
||||
|
||||
## Roadmap
|
||||
The following functionality is on the roadmap and will be available in future releases.
|
||||
- DNS intercept mode
|
||||
- Direct listener mode
|
||||
- Support for more routers (let us know which ones)
|
||||
|
||||
@@ -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.
|
||||
@@ -0,0 +1,4 @@
|
||||
package ctrld
|
||||
|
||||
// SelfDiscover reports whether ctrld should only do self discover.
|
||||
func SelfDiscover() bool { return true }
|
||||
@@ -0,0 +1,6 @@
|
||||
//go:build !windows && !darwin
|
||||
|
||||
package ctrld
|
||||
|
||||
// SelfDiscover reports whether ctrld should only do self discover.
|
||||
func SelfDiscover() bool { return false }
|
||||
@@ -0,0 +1,18 @@
|
||||
package ctrld
|
||||
|
||||
import (
|
||||
"golang.org/x/sys/windows"
|
||||
)
|
||||
|
||||
// isWindowsWorkStation reports whether ctrld was run on a Windows workstation machine.
|
||||
func isWindowsWorkStation() bool {
|
||||
// From https://learn.microsoft.com/en-us/windows/win32/api/winnt/ns-winnt-osversioninfoexa
|
||||
const VER_NT_WORKSTATION = 0x0000001
|
||||
osvi := windows.RtlGetVersion()
|
||||
return osvi.ProductType == VER_NT_WORKSTATION
|
||||
}
|
||||
|
||||
// SelfDiscover reports whether ctrld should only do self discover.
|
||||
func SelfDiscover() bool {
|
||||
return isWindowsWorkStation()
|
||||
}
|
||||
@@ -7,4 +7,9 @@ import (
|
||||
)
|
||||
|
||||
// addExtraSplitDnsRule adds split DNS rule if present.
|
||||
func addExtraSplitDnsRule(_ *ctrld.Config) {}
|
||||
func addExtraSplitDnsRule(_ *ctrld.Config) bool { return false }
|
||||
|
||||
// getActiveDirectoryDomain returns AD domain name of this computer.
|
||||
func getActiveDirectoryDomain() (string, error) {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
+45
-16
@@ -1,45 +1,74 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"github.com/microsoft/wmi/pkg/base/host"
|
||||
hh "github.com/microsoft/wmi/pkg/hardware/host"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
"github.com/Control-D-Inc/ctrld/internal/system"
|
||||
)
|
||||
|
||||
// addExtraSplitDnsRule adds split DNS rule for domain if it's part of active directory.
|
||||
func addExtraSplitDnsRule(cfg *ctrld.Config) {
|
||||
domain, err := getActiveDirectoryDomain()
|
||||
func addExtraSplitDnsRule(cfg *ctrld.Config) bool {
|
||||
domain, err := system.GetActiveDirectoryDomain()
|
||||
if err != nil {
|
||||
mainLog.Load().Debug().Msgf("unable to get active directory domain: %v", err)
|
||||
return
|
||||
return false
|
||||
}
|
||||
if domain == "" {
|
||||
mainLog.Load().Debug().Msg("no active directory domain found")
|
||||
return
|
||||
return false
|
||||
}
|
||||
// Network rules are lowercase during toml config marshaling,
|
||||
// lowercase the domain here too for consistency.
|
||||
domain = strings.ToLower(domain)
|
||||
domainRuleAdded := addSplitDnsRule(cfg, domain)
|
||||
wildcardDomainRuleRuleAdded := addSplitDnsRule(cfg, "*."+strings.TrimPrefix(domain, "."))
|
||||
return domainRuleAdded || wildcardDomainRuleRuleAdded
|
||||
}
|
||||
|
||||
// addSplitDnsRule adds split-rule for given domain if there's no existed rule.
|
||||
// The return value indicates whether the split-rule was added or not.
|
||||
func addSplitDnsRule(cfg *ctrld.Config, domain string) bool {
|
||||
for n, lc := range cfg.Listener {
|
||||
if lc.Policy == nil {
|
||||
lc.Policy = &ctrld.ListenerPolicyConfig{}
|
||||
}
|
||||
domainRule := "*." + strings.TrimPrefix(domain, ".")
|
||||
for _, rule := range lc.Policy.Rules {
|
||||
if _, ok := rule[domainRule]; ok {
|
||||
mainLog.Load().Debug().Msgf("domain rule already exist for listener.%s", n)
|
||||
return
|
||||
if _, ok := rule[domain]; ok {
|
||||
mainLog.Load().Debug().Msgf("split-rule %q already existed for listener.%s", domain, n)
|
||||
return false
|
||||
}
|
||||
}
|
||||
mainLog.Load().Debug().Msgf("adding active directory domain for listener.%s", n)
|
||||
lc.Policy.Rules = append(lc.Policy.Rules, ctrld.Rule{domainRule: []string{}})
|
||||
mainLog.Load().Debug().Msgf("adding split-rule %q for listener.%s", domain, n)
|
||||
lc.Policy.Rules = append(lc.Policy.Rules, ctrld.Rule{domain: []string{}})
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// getActiveDirectoryDomain returns AD domain name of this computer.
|
||||
func getActiveDirectoryDomain() (string, error) {
|
||||
cmd := "$obj = Get-WmiObject Win32_ComputerSystem; if ($obj.PartOfDomain) { $obj.Domain }"
|
||||
output, err := powershell(cmd)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to get domain name: %w, output:\n\n%s", err, string(output))
|
||||
log.SetOutput(io.Discard)
|
||||
defer log.SetOutput(os.Stderr)
|
||||
whost := host.NewWmiLocalHost()
|
||||
cs, err := hh.GetComputerSystem(whost)
|
||||
if cs != nil {
|
||||
defer cs.Close()
|
||||
}
|
||||
return string(output), nil
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
pod, err := cs.GetPropertyPartOfDomain()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if pod {
|
||||
return cs.GetPropertyDomain()
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
"github.com/Control-D-Inc/ctrld/internal/system"
|
||||
"github.com/Control-D-Inc/ctrld/testhelper"
|
||||
)
|
||||
|
||||
func Test_getActiveDirectoryDomain(t *testing.T) {
|
||||
start := time.Now()
|
||||
domain, err := system.GetActiveDirectoryDomain()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Logf("Using Windows API takes: %d", time.Since(start).Milliseconds())
|
||||
|
||||
start = time.Now()
|
||||
domainPowershell, err := getActiveDirectoryDomainPowershell()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Logf("Using Powershell takes: %d", time.Since(start).Milliseconds())
|
||||
|
||||
if domain != domainPowershell {
|
||||
t.Fatalf("result mismatch, want: %v, got: %v", domainPowershell, domain)
|
||||
}
|
||||
}
|
||||
|
||||
func getActiveDirectoryDomainPowershell() (string, error) {
|
||||
cmd := "$obj = Get-WmiObject Win32_ComputerSystem; if ($obj.PartOfDomain) { $obj.Domain }"
|
||||
output, err := powershell(cmd)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to get domain name: %w, output:\n\n%s", err, string(output))
|
||||
}
|
||||
return string(output), nil
|
||||
}
|
||||
|
||||
func Test_addSplitDnsRule(t *testing.T) {
|
||||
newCfg := func(domains ...string) *ctrld.Config {
|
||||
cfg := testhelper.SampleConfig(t)
|
||||
lc := cfg.Listener["0"]
|
||||
for _, domain := range domains {
|
||||
lc.Policy.Rules = append(lc.Policy.Rules, ctrld.Rule{domain: []string{}})
|
||||
}
|
||||
return cfg
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
cfg *ctrld.Config
|
||||
domain string
|
||||
added bool
|
||||
}{
|
||||
{"added", newCfg(), "example.com", true},
|
||||
{"TLD existed", newCfg("example.com"), "*.example.com", true},
|
||||
{"wildcard existed", newCfg("*.example.com"), "example.com", true},
|
||||
{"not added TLD", newCfg("example.com", "*.example.com"), "example.com", false},
|
||||
{"not added wildcard", newCfg("example.com", "*.example.com"), "*.example.com", false},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
added := addSplitDnsRule(tc.cfg, tc.domain)
|
||||
assert.Equal(t, tc.added, added)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
//go:build cgo
|
||||
|
||||
package cli
|
||||
|
||||
const cgoEnabled = true
|
||||
+727
-1126
File diff suppressed because it is too large
Load Diff
@@ -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.
|
||||
+1710
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -25,6 +25,16 @@ func newControlClient(addr string) *controlClient {
|
||||
}
|
||||
|
||||
func (c *controlClient) post(path string, data io.Reader) (*http.Response, error) {
|
||||
// for log/send, set the timeout to 5 minutes
|
||||
if path == sendLogsPath {
|
||||
c.c.Timeout = time.Minute * 5
|
||||
}
|
||||
return c.c.Post("http://unix"+path, contentTypeJson, data)
|
||||
}
|
||||
|
||||
// postStream sends a POST request with no timeout, suitable for long-lived streaming connections.
|
||||
func (c *controlClient) postStream(path string, data io.Reader) (*http.Response, error) {
|
||||
c.c.Timeout = 0
|
||||
return c.c.Post("http://unix"+path, contentTypeJson, data)
|
||||
}
|
||||
|
||||
|
||||
+347
-13
@@ -3,18 +3,21 @@ package cli
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/kardianos/service"
|
||||
|
||||
dto "github.com/prometheus/client_model/go"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
"github.com/Control-D-Inc/ctrld/internal/controld"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -25,8 +28,18 @@ const (
|
||||
deactivationPath = "/deactivation"
|
||||
cdPath = "/cd"
|
||||
ifacePath = "/iface"
|
||||
viewLogsPath = "/log/view"
|
||||
sendLogsPath = "/log/send"
|
||||
tailLogsPath = "/log/tail"
|
||||
)
|
||||
|
||||
type ifaceResponse struct {
|
||||
Name string `json:"name"`
|
||||
All bool `json:"all"`
|
||||
OK bool `json:"ok"`
|
||||
InterceptMode string `json:"intercept_mode,omitempty"` // "dns", "hard", or "" (not intercepting)
|
||||
}
|
||||
|
||||
type controlServer struct {
|
||||
server *http.Server
|
||||
mux *http.ServeMux
|
||||
@@ -46,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
|
||||
}
|
||||
@@ -69,33 +88,81 @@ func (s *controlServer) register(pattern string, handler http.Handler) {
|
||||
|
||||
func (p *prog) registerControlServerHandler() {
|
||||
p.cs.register(listClientsPath, http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
|
||||
mainLog.Load().Debug().Msg("handling list clients request")
|
||||
|
||||
clients := p.ciTable.ListClients()
|
||||
mainLog.Load().Debug().Int("client_count", len(clients)).Msg("retrieved clients list")
|
||||
|
||||
sort.Slice(clients, func(i, j int) bool {
|
||||
return clients[i].IP.Less(clients[j].IP)
|
||||
})
|
||||
mainLog.Load().Debug().Msg("sorted clients by IP address")
|
||||
|
||||
if p.metricsQueryStats.Load() {
|
||||
for _, client := range clients {
|
||||
mainLog.Load().Debug().Msg("metrics query stats enabled, collecting query counts")
|
||||
|
||||
for idx, client := range clients {
|
||||
mainLog.Load().Debug().
|
||||
Int("index", idx).
|
||||
Str("ip", client.IP.String()).
|
||||
Str("mac", client.Mac).
|
||||
Str("hostname", client.Hostname).
|
||||
Msg("processing client metrics")
|
||||
|
||||
client.IncludeQueryCount = true
|
||||
dm := &dto.Metric{}
|
||||
|
||||
if statsClientQueriesCount.MetricVec == nil {
|
||||
mainLog.Load().Debug().
|
||||
Str("client_ip", client.IP.String()).
|
||||
Msg("skipping metrics collection: MetricVec is nil")
|
||||
continue
|
||||
}
|
||||
|
||||
m, err := statsClientQueriesCount.MetricVec.GetMetricWithLabelValues(
|
||||
client.IP.String(),
|
||||
client.Mac,
|
||||
client.Hostname,
|
||||
)
|
||||
if err != nil {
|
||||
mainLog.Load().Debug().Err(err).Msgf("could not get metrics for client: %v", client)
|
||||
mainLog.Load().Debug().
|
||||
Err(err).
|
||||
Str("client_ip", client.IP.String()).
|
||||
Str("mac", client.Mac).
|
||||
Str("hostname", client.Hostname).
|
||||
Msg("failed to get metrics for client")
|
||||
continue
|
||||
}
|
||||
if err := m.Write(dm); err == nil {
|
||||
|
||||
if err := m.Write(dm); err == nil && dm.Counter != nil {
|
||||
client.QueryCount = int64(dm.Counter.GetValue())
|
||||
mainLog.Load().Debug().
|
||||
Str("client_ip", client.IP.String()).
|
||||
Int64("query_count", client.QueryCount).
|
||||
Msg("successfully collected query count")
|
||||
} else if err != nil {
|
||||
mainLog.Load().Debug().
|
||||
Err(err).
|
||||
Str("client_ip", client.IP.String()).
|
||||
Msg("failed to write metric")
|
||||
}
|
||||
}
|
||||
} else {
|
||||
mainLog.Load().Debug().Msg("metrics query stats disabled, skipping query counts")
|
||||
}
|
||||
|
||||
if err := json.NewEncoder(w).Encode(&clients); err != nil {
|
||||
mainLog.Load().Error().
|
||||
Err(err).
|
||||
Int("client_count", len(clients)).
|
||||
Msg("failed to encode clients response")
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
mainLog.Load().Debug().
|
||||
Int("client_count", len(clients)).
|
||||
Msg("successfully sent clients list response")
|
||||
}))
|
||||
p.cs.register(startedPath, http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
|
||||
select {
|
||||
@@ -152,8 +219,36 @@ func (p *prog) registerControlServerHandler() {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
p.cs.register(deactivationPath, http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
|
||||
// Non-cd mode or pin code not set, always allowing deactivation.
|
||||
if cdUID == "" || deactivationPinNotSet() {
|
||||
// Non-cd mode always allowing deactivation.
|
||||
if cdUID == "" {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
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(context.Background(), rcReq, cdDev); rc != nil {
|
||||
if rc.DeactivationPin != nil {
|
||||
cdDeactivationPin.Store(*rc.DeactivationPin)
|
||||
} else {
|
||||
cdDeactivationPin.Store(defaultDeactivationPin)
|
||||
}
|
||||
} else {
|
||||
mainLog.Load().Warn().Err(err).Msg("could not re-fetch deactivation pin code")
|
||||
}
|
||||
|
||||
// If pin code not set, allowing deactivation.
|
||||
if !deactivationPinSet() {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
return
|
||||
}
|
||||
@@ -167,11 +262,21 @@ func (p *prog) registerControlServerHandler() {
|
||||
|
||||
code := http.StatusForbidden
|
||||
switch req.Pin {
|
||||
case cdDeactivationPin:
|
||||
case cdDeactivationPin.Load():
|
||||
code = http.StatusOK
|
||||
deactivationFailedAttempts.Store(0)
|
||||
select {
|
||||
case p.pinCodeValidCh <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
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)
|
||||
}))
|
||||
@@ -184,16 +289,245 @@ func (p *prog) registerControlServerHandler() {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
}))
|
||||
p.cs.register(ifacePath, http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
|
||||
res := &ifaceResponse{Name: iface}
|
||||
// p.setDNS is only called when running as a service
|
||||
if !service.Interactive() {
|
||||
<-p.csSetDnsDone
|
||||
if p.csSetDnsOk {
|
||||
w.Write([]byte(iface))
|
||||
return
|
||||
res.Name = p.runningIface
|
||||
res.All = p.requiredMultiNICsConfig
|
||||
res.OK = true
|
||||
// Report intercept mode to the start command for proper log output.
|
||||
if interceptMode == "dns" || interceptMode == "hard" {
|
||||
res.InterceptMode = interceptMode
|
||||
}
|
||||
}
|
||||
}
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
if err := json.NewEncoder(w).Encode(res); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
http.Error(w, fmt.Sprintf("could not marshal iface data: %v", err), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
}))
|
||||
p.cs.register(viewLogsPath, http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
|
||||
lr, err := p.logReader()
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
defer lr.r.Close()
|
||||
if lr.size == 0 {
|
||||
w.WriteHeader(http.StatusMovedPermanently)
|
||||
return
|
||||
}
|
||||
data, err := io.ReadAll(lr.r)
|
||||
if err != nil {
|
||||
http.Error(w, fmt.Sprintf("could not read log: %v", err), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
if err := json.NewEncoder(w).Encode(&logViewResponse{Data: string(data)}); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
http.Error(w, fmt.Sprintf("could not marshal log data: %v", err), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
}))
|
||||
p.cs.register(sendLogsPath, http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
|
||||
if time.Since(p.internalLogSent) < logWriterSentInterval {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
r, err := p.logReader()
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
if r.size == 0 {
|
||||
w.WriteHeader(http.StatusMovedPermanently)
|
||||
return
|
||||
}
|
||||
req := &controld.LogsRequest{
|
||||
UID: cdUID,
|
||||
Data: r.r,
|
||||
}
|
||||
mainLog.Load().Debug().Msg("sending log file to ControlD server")
|
||||
resp := logSentResponse{Size: r.size}
|
||||
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)
|
||||
} else {
|
||||
mainLog.Load().Debug().Msg("sending log file successfully")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}
|
||||
if err := json.NewEncoder(w).Encode(&resp); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
}
|
||||
p.internalLogSent = time.Now()
|
||||
}))
|
||||
p.cs.register(tailLogsPath, http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
|
||||
flusher, ok := w.(http.Flusher)
|
||||
if !ok {
|
||||
http.Error(w, "streaming unsupported", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
// Determine logging mode and validate before starting the stream.
|
||||
var lw *logWriter
|
||||
useInternalLog := p.needInternalLogging()
|
||||
if useInternalLog {
|
||||
p.mu.Lock()
|
||||
lw = p.internalLogWriter
|
||||
p.mu.Unlock()
|
||||
if lw == nil {
|
||||
w.WriteHeader(http.StatusMovedPermanently)
|
||||
return
|
||||
}
|
||||
} else if p.cfg.Service.LogPath == "" {
|
||||
// No logging configured at all.
|
||||
w.WriteHeader(http.StatusMovedPermanently)
|
||||
return
|
||||
}
|
||||
|
||||
// Parse optional "lines" query param for initial context.
|
||||
numLines := 10
|
||||
if v := request.URL.Query().Get("lines"); v != "" {
|
||||
if n, err := strconv.Atoi(v); err == nil && n >= 0 {
|
||||
numLines = n
|
||||
}
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||
w.Header().Set("Transfer-Encoding", "chunked")
|
||||
w.Header().Set("X-Content-Type-Options", "nosniff")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
|
||||
if useInternalLog {
|
||||
// Internal logging mode: subscribe to the logWriter.
|
||||
|
||||
// Send last N lines as initial context.
|
||||
if numLines > 0 {
|
||||
if tail := lw.tailLastLines(numLines); len(tail) > 0 {
|
||||
w.Write(tail)
|
||||
flusher.Flush()
|
||||
}
|
||||
}
|
||||
|
||||
ch, unsub := lw.Subscribe()
|
||||
defer unsub()
|
||||
for {
|
||||
select {
|
||||
case data, ok := <-ch:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if _, err := w.Write(data); err != nil {
|
||||
return
|
||||
}
|
||||
flusher.Flush()
|
||||
case <-request.Context().Done():
|
||||
return
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// File-based logging mode: tail the log file.
|
||||
logFile := normalizeLogFilePath(p.cfg.Service.LogPath)
|
||||
f, err := os.Open(logFile)
|
||||
if err != nil {
|
||||
// Already committed 200, just return.
|
||||
return
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
// Seek to show last N lines.
|
||||
if numLines > 0 {
|
||||
if tail := tailFileLastLines(f, numLines); len(tail) > 0 {
|
||||
w.Write(tail)
|
||||
flusher.Flush()
|
||||
}
|
||||
} else {
|
||||
// Seek to end.
|
||||
f.Seek(0, io.SeekEnd)
|
||||
}
|
||||
|
||||
// Poll for new data.
|
||||
buf := make([]byte, 4096)
|
||||
ticker := time.NewTicker(200 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
n, err := f.Read(buf)
|
||||
if n > 0 {
|
||||
if _, werr := w.Write(buf[:n]); werr != nil {
|
||||
return
|
||||
}
|
||||
flusher.Flush()
|
||||
}
|
||||
if err != nil && err != io.EOF {
|
||||
return
|
||||
}
|
||||
case <-request.Context().Done():
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}))
|
||||
}
|
||||
|
||||
// tailFileLastLines reads the last n lines from a file and returns them.
|
||||
// The file position is left at the end of the file after this call.
|
||||
func tailFileLastLines(f *os.File, n int) []byte {
|
||||
stat, err := f.Stat()
|
||||
if err != nil || stat.Size() == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Read from the end in chunks to find the last n lines.
|
||||
const chunkSize = 4096
|
||||
fileSize := stat.Size()
|
||||
var lines []byte
|
||||
offset := fileSize
|
||||
count := 0
|
||||
|
||||
for offset > 0 && count <= n {
|
||||
readSize := int64(chunkSize)
|
||||
if readSize > offset {
|
||||
readSize = offset
|
||||
}
|
||||
offset -= readSize
|
||||
buf := make([]byte, readSize)
|
||||
nRead, err := f.ReadAt(buf, offset)
|
||||
if err != nil && err != io.EOF {
|
||||
break
|
||||
}
|
||||
buf = buf[:nRead]
|
||||
lines = append(buf, lines...)
|
||||
|
||||
// Count newlines in this chunk.
|
||||
for _, b := range buf {
|
||||
if b == '\n' {
|
||||
count++
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Trim to last n lines.
|
||||
idx := 0
|
||||
nlCount := 0
|
||||
for i := len(lines) - 1; i >= 0; i-- {
|
||||
if lines[i] == '\n' {
|
||||
nlCount++
|
||||
if nlCount == n+1 {
|
||||
idx = i + 1
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
lines = lines[idx:]
|
||||
|
||||
// Seek to end of file for subsequent reads.
|
||||
f.Seek(0, io.SeekEnd)
|
||||
return lines
|
||||
}
|
||||
|
||||
func jsonResponse(next http.Handler) http.Handler {
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
//go:build !windows && !darwin
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// startDNSIntercept is not supported on this platform.
|
||||
// DNS intercept mode is only available on Windows (via WFP) and macOS (via pf).
|
||||
func (p *prog) startDNSIntercept() error {
|
||||
return fmt.Errorf("dns intercept: not supported on this platform (only Windows and macOS)")
|
||||
}
|
||||
|
||||
// stopDNSIntercept is a no-op on unsupported platforms.
|
||||
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() pfAnchorCheckResult {
|
||||
return pfAnchorCheckSkipped
|
||||
}
|
||||
|
||||
// checkTunnelInterfaceChanges is a no-op on unsupported platforms.
|
||||
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() {}
|
||||
|
||||
// pfInterceptMonitor is a no-op on unsupported platforms.
|
||||
func (p *prog) pfInterceptMonitor() {}
|
||||
|
||||
// osHealthcheckSuppressed always returns false on non-Windows platforms —
|
||||
// WFP loopback protect (the trigger for suppression) is Windows-only.
|
||||
func (p *prog) osHealthcheckSuppressed() bool { return false }
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
+1146
-83
File diff suppressed because it is too large
Load Diff
+84
-11
@@ -22,14 +22,15 @@ func Test_wildcardMatches(t *testing.T) {
|
||||
domain string
|
||||
match bool
|
||||
}{
|
||||
{"domain - prefix parent should not match", "*.windscribe.com", "windscribe.com", false},
|
||||
{"domain - prefix", "*.windscribe.com", "anything.windscribe.com", true},
|
||||
{"domain - prefix not match other s", "*.windscribe.com", "example.com", false},
|
||||
{"domain - prefix not match s in name", "*.windscribe.com", "wwindscribe.com", false},
|
||||
{"domain - suffix", "suffix.*", "suffix.windscribe.com", true},
|
||||
{"domain - suffix not match other", "suffix.*", "suffix1.windscribe.com", false},
|
||||
{"domain - both", "suffix.*.windscribe.com", "suffix.anything.windscribe.com", true},
|
||||
{"domain - both not match", "suffix.*.windscribe.com", "suffix1.suffix.windscribe.com", false},
|
||||
{"domain - prefix parent should not match", "*.example.com", "example.com", false},
|
||||
{"domain - prefix", "*.example.com", "anything.example.com", true},
|
||||
{"domain - prefix not match other s", "*.example.com", "other.org", false},
|
||||
{"domain - prefix not match s in name", "*.example.com", "eexample.com", false},
|
||||
{"domain - suffix", "suffix.*", "suffix.example.com", true},
|
||||
{"domain - suffix not match other", "suffix.*", "suffix1.example.com", false},
|
||||
{"domain - both", "suffix.*.example.com", "suffix.anything.example.com", true},
|
||||
{"domain - both not match", "suffix.*.example.com", "suffix1.suffix.example.com", false},
|
||||
{"domain - case-insensitive", "*.EXAMPLE.com", "anything.example.com", true},
|
||||
{"mac - prefix", "*:98:05:b4:2b", "d4:67:98:05:b4:2b", true},
|
||||
{"mac - prefix not match other s", "*:98:05:b4:2b", "0d:ba:54:09:94:2c", false},
|
||||
{"mac - prefix not match s in name", "*:98:05:b4:2b", "e4:67:97:05:b4:2b", false},
|
||||
@@ -56,9 +57,9 @@ func Test_canonicalName(t *testing.T) {
|
||||
domain string
|
||||
canonical string
|
||||
}{
|
||||
{"fqdn to canonical", "windscribe.com.", "windscribe.com"},
|
||||
{"already canonical", "windscribe.com", "windscribe.com"},
|
||||
{"case insensitive", "Windscribe.Com.", "windscribe.com"},
|
||||
{"fqdn to canonical", "example.com.", "example.com"},
|
||||
{"already canonical", "example.com", "example.com"},
|
||||
{"case insensitive", "Example.Com.", "example.com"},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
@@ -74,6 +75,7 @@ func Test_canonicalName(t *testing.T) {
|
||||
|
||||
func Test_prog_upstreamFor(t *testing.T) {
|
||||
cfg := testhelper.SampleConfig(t)
|
||||
cfg.Service.LeakOnUpstreamFailure = func(v bool) *bool { return &v }(false)
|
||||
p := &prog{cfg: cfg}
|
||||
p.um = newUpstreamMonitor(p.cfg)
|
||||
p.lanLoopGuard = newLoopGuard()
|
||||
@@ -364,6 +366,9 @@ func Test_isLanHostnameQuery(t *testing.T) {
|
||||
{"A not LAN", newDnsMsgWithHostname("example.com", dns.TypeA), false},
|
||||
{"AAAA not LAN", newDnsMsgWithHostname("example.com", dns.TypeAAAA), false},
|
||||
{"Not A or AAAA", newDnsMsgWithHostname("foo", dns.TypeTXT), false},
|
||||
{".domain", newDnsMsgWithHostname("foo.domain", dns.TypeA), true},
|
||||
{".lan", newDnsMsgWithHostname("foo.lan", dns.TypeA), true},
|
||||
{".local", newDnsMsgWithHostname("foo.local", dns.TypeA), true},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
tc := tc
|
||||
@@ -400,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 {
|
||||
@@ -413,6 +420,27 @@ func Test_isPrivatePtrLookup(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func Test_isSrvLanLookup(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
msg *dns.Msg
|
||||
isSrvLookup bool
|
||||
}{
|
||||
{"SRV LAN", newDnsMsgWithHostname("foo", dns.TypeSRV), true},
|
||||
{"Not SRV", newDnsMsgWithHostname("foo", dns.TypeNone), false},
|
||||
{"Not SRV LAN", newDnsMsgWithHostname("controld.com", dns.TypeSRV), false},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
if got := isSrvLanLookup(tc.msg); tc.isSrvLookup != got {
|
||||
t.Errorf("unexpected result, want: %v, got: %v", tc.isSrvLookup, got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func Test_isWanClient(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -426,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 {
|
||||
@@ -438,3 +471,43 @@ func Test_isWanClient(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func Test_prog_queryFromSelf(t *testing.T) {
|
||||
p := &prog{}
|
||||
require.NotPanics(t, func() {
|
||||
p.queryFromSelf("")
|
||||
})
|
||||
require.NotPanics(t, func() {
|
||||
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,38 @@
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
+92
-5
@@ -1,5 +1,12 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
// AppCallback provides hooks for injecting certain functionalities
|
||||
// from mobile platforms to main ctrld cli.
|
||||
type AppCallback struct {
|
||||
@@ -11,9 +18,89 @@ type AppCallback struct {
|
||||
|
||||
// AppConfig allows overwriting ctrld cli flags from mobile platforms.
|
||||
type AppConfig struct {
|
||||
CdUID string
|
||||
HomeDir string
|
||||
UpstreamProto string
|
||||
Verbose int
|
||||
LogPath string
|
||||
CdUID string
|
||||
ProvisionID string
|
||||
CustomHostname string
|
||||
HomeDir string
|
||||
UpstreamProto string
|
||||
Verbose int
|
||||
LogPath string
|
||||
}
|
||||
|
||||
const (
|
||||
defaultHTTPTimeout = 30 * time.Second
|
||||
defaultMaxRetries = 3
|
||||
downloadServerIp = "23.171.240.151"
|
||||
)
|
||||
|
||||
// httpClientWithFallback returns an HTTP client configured with timeout and IPv4 fallback
|
||||
func httpClientWithFallback(timeout time.Duration) *http.Client {
|
||||
return &http.Client{
|
||||
Timeout: timeout,
|
||||
Transport: &http.Transport{
|
||||
// Prefer IPv4 over IPv6
|
||||
DialContext: (&net.Dialer{
|
||||
Timeout: 10 * time.Second,
|
||||
KeepAlive: 30 * time.Second,
|
||||
FallbackDelay: 1 * time.Millisecond, // Very small delay to prefer IPv4
|
||||
}).DialContext,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// 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
|
||||
var ipReq *http.Request
|
||||
if ip != "" {
|
||||
ipReq = req.Clone(req.Context())
|
||||
ipReq.Host = ip
|
||||
ipReq.URL.Host = ip
|
||||
}
|
||||
for attempt := 0; attempt < maxRetries; attempt++ {
|
||||
if attempt > 0 {
|
||||
time.Sleep(time.Second * time.Duration(attempt+1)) // Exponential backoff
|
||||
}
|
||||
|
||||
resp, err := client.Do(req)
|
||||
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, 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 = 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: %w", maxRetries, req.Method, req.URL, lastErr)
|
||||
}
|
||||
|
||||
// Helper for making GET requests with retries
|
||||
func getWithRetry(url string, ip string) (*http.Response, error) {
|
||||
req, err := http.NewRequest(http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return doWithRetry(req, defaultMaxRetries, ip)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,339 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"io"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// =============================================================================
|
||||
// logWriter.tailLastLines tests
|
||||
// =============================================================================
|
||||
|
||||
func Test_logWriter_tailLastLines_Empty(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
if got := lw.tailLastLines(10); got != nil {
|
||||
t.Fatalf("expected nil for empty buffer, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_tailLastLines_ZeroLines(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
lw.Write([]byte("line1\nline2\n"))
|
||||
if got := lw.tailLastLines(0); got != nil {
|
||||
t.Fatalf("expected nil for n=0, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_tailLastLines_NegativeLines(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
lw.Write([]byte("line1\nline2\n"))
|
||||
if got := lw.tailLastLines(-1); got != nil {
|
||||
t.Fatalf("expected nil for n=-1, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_tailLastLines_FewerThanN(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
lw.Write([]byte("line1\nline2\n"))
|
||||
got := string(lw.tailLastLines(10))
|
||||
want := "line1\nline2\n"
|
||||
if got != want {
|
||||
t.Fatalf("got %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_tailLastLines_ExactN(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
lw.Write([]byte("line1\nline2\nline3\n"))
|
||||
got := string(lw.tailLastLines(3))
|
||||
want := "line1\nline2\nline3\n"
|
||||
if got != want {
|
||||
t.Fatalf("got %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_tailLastLines_MoreThanN(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
lw.Write([]byte("line1\nline2\nline3\nline4\nline5\n"))
|
||||
got := string(lw.tailLastLines(2))
|
||||
want := "line4\nline5\n"
|
||||
if got != want {
|
||||
t.Fatalf("got %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_tailLastLines_NoTrailingNewline(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
lw.Write([]byte("line1\nline2\nline3"))
|
||||
// Without trailing newline, "line3" is a partial line.
|
||||
// Asking for 1 line returns the last newline-terminated line plus the partial.
|
||||
got := string(lw.tailLastLines(1))
|
||||
want := "line2\nline3"
|
||||
if got != want {
|
||||
t.Fatalf("got %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_tailLastLines_SingleLineNoNewline(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
lw.Write([]byte("only line"))
|
||||
got := string(lw.tailLastLines(5))
|
||||
want := "only line"
|
||||
if got != want {
|
||||
t.Fatalf("got %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_tailLastLines_SingleLineWithNewline(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
lw.Write([]byte("only line\n"))
|
||||
got := string(lw.tailLastLines(1))
|
||||
want := "only line\n"
|
||||
if got != want {
|
||||
t.Fatalf("got %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// logWriter.Subscribe tests
|
||||
// =============================================================================
|
||||
|
||||
func Test_logWriter_Subscribe_Basic(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
ch, unsub := lw.Subscribe()
|
||||
defer unsub()
|
||||
|
||||
msg := []byte("hello world\n")
|
||||
lw.Write(msg)
|
||||
|
||||
select {
|
||||
case got := <-ch:
|
||||
if string(got) != string(msg) {
|
||||
t.Fatalf("got %q, want %q", got, msg)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("timed out waiting for subscriber data")
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_Subscribe_MultipleSubscribers(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
ch1, unsub1 := lw.Subscribe()
|
||||
defer unsub1()
|
||||
ch2, unsub2 := lw.Subscribe()
|
||||
defer unsub2()
|
||||
|
||||
msg := []byte("broadcast\n")
|
||||
lw.Write(msg)
|
||||
|
||||
for i, ch := range []<-chan []byte{ch1, ch2} {
|
||||
select {
|
||||
case got := <-ch:
|
||||
if string(got) != string(msg) {
|
||||
t.Fatalf("subscriber %d: got %q, want %q", i, got, msg)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatalf("subscriber %d: timed out", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_Subscribe_Unsubscribe(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
ch, unsub := lw.Subscribe()
|
||||
|
||||
// Verify subscribed.
|
||||
lw.Write([]byte("before unsub\n"))
|
||||
select {
|
||||
case <-ch:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("timed out before unsub")
|
||||
}
|
||||
|
||||
unsub()
|
||||
|
||||
// Channel should be closed after unsub.
|
||||
if _, ok := <-ch; ok {
|
||||
t.Fatal("channel should be closed after unsubscribe")
|
||||
}
|
||||
|
||||
// Verify subscriber list is empty.
|
||||
lw.mu.Lock()
|
||||
count := len(lw.subscribers)
|
||||
lw.mu.Unlock()
|
||||
if count != 0 {
|
||||
t.Fatalf("expected 0 subscribers after unsub, got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_Subscribe_UnsubscribeIdempotent(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
_, unsub := lw.Subscribe()
|
||||
unsub()
|
||||
// Second unsub should not panic.
|
||||
unsub()
|
||||
}
|
||||
|
||||
func Test_logWriter_Subscribe_SlowSubscriberDropped(t *testing.T) {
|
||||
lw := newLogWriterWithSize(4096)
|
||||
ch, unsub := lw.Subscribe()
|
||||
defer unsub()
|
||||
|
||||
// Fill the subscriber channel (buffer size is 256).
|
||||
for i := 0; i < 300; i++ {
|
||||
lw.Write([]byte("msg\n"))
|
||||
}
|
||||
|
||||
// Should have 256 buffered messages, rest dropped.
|
||||
count := 0
|
||||
for {
|
||||
select {
|
||||
case <-ch:
|
||||
count++
|
||||
default:
|
||||
goto done
|
||||
}
|
||||
}
|
||||
done:
|
||||
if count != 256 {
|
||||
t.Fatalf("expected 256 buffered messages, got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_Subscribe_ConcurrentWriteAndRead(t *testing.T) {
|
||||
lw := newLogWriterWithSize(64 * 1024)
|
||||
ch, unsub := lw.Subscribe()
|
||||
defer unsub()
|
||||
|
||||
const numWrites = 100
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for i := 0; i < numWrites; i++ {
|
||||
lw.Write([]byte("concurrent write\n"))
|
||||
}
|
||||
}()
|
||||
|
||||
received := 0
|
||||
timeout := time.After(5 * time.Second)
|
||||
for received < numWrites {
|
||||
select {
|
||||
case <-ch:
|
||||
received++
|
||||
case <-timeout:
|
||||
t.Fatalf("timed out after receiving %d/%d messages", received, numWrites)
|
||||
}
|
||||
}
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// tailFileLastLines tests
|
||||
// =============================================================================
|
||||
|
||||
func writeTempFile(t *testing.T, content string) *os.File {
|
||||
t.Helper()
|
||||
f, err := os.CreateTemp(t.TempDir(), "tail-test-*")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := f.WriteString(content); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return f
|
||||
}
|
||||
|
||||
func Test_tailFileLastLines_Empty(t *testing.T) {
|
||||
f := writeTempFile(t, "")
|
||||
defer f.Close()
|
||||
if got := tailFileLastLines(f, 10); got != nil {
|
||||
t.Fatalf("expected nil for empty file, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_tailFileLastLines_FewerThanN(t *testing.T) {
|
||||
f := writeTempFile(t, "line1\nline2\n")
|
||||
defer f.Close()
|
||||
got := string(tailFileLastLines(f, 10))
|
||||
want := "line1\nline2\n"
|
||||
if got != want {
|
||||
t.Fatalf("got %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_tailFileLastLines_ExactN(t *testing.T) {
|
||||
f := writeTempFile(t, "a\nb\nc\n")
|
||||
defer f.Close()
|
||||
got := string(tailFileLastLines(f, 3))
|
||||
want := "a\nb\nc\n"
|
||||
if got != want {
|
||||
t.Fatalf("got %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_tailFileLastLines_MoreThanN(t *testing.T) {
|
||||
f := writeTempFile(t, "line1\nline2\nline3\nline4\nline5\n")
|
||||
defer f.Close()
|
||||
got := string(tailFileLastLines(f, 2))
|
||||
want := "line4\nline5\n"
|
||||
if got != want {
|
||||
t.Fatalf("got %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_tailFileLastLines_NoTrailingNewline(t *testing.T) {
|
||||
f := writeTempFile(t, "line1\nline2\nline3")
|
||||
defer f.Close()
|
||||
// Without trailing newline, partial last line comes with the previous line.
|
||||
got := string(tailFileLastLines(f, 1))
|
||||
want := "line2\nline3"
|
||||
if got != want {
|
||||
t.Fatalf("got %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_tailFileLastLines_LargerThanChunk(t *testing.T) {
|
||||
// Build content larger than the 4096 chunk size to exercise multi-chunk reads.
|
||||
var sb strings.Builder
|
||||
for i := 0; i < 200; i++ {
|
||||
sb.WriteString(strings.Repeat("x", 50))
|
||||
sb.WriteByte('\n')
|
||||
}
|
||||
f := writeTempFile(t, sb.String())
|
||||
defer f.Close()
|
||||
got := string(tailFileLastLines(f, 3))
|
||||
lines := strings.Split(strings.TrimRight(got, "\n"), "\n")
|
||||
if len(lines) != 3 {
|
||||
t.Fatalf("expected 3 lines, got %d: %q", len(lines), got)
|
||||
}
|
||||
expectedLine := strings.Repeat("x", 50)
|
||||
for _, line := range lines {
|
||||
if line != expectedLine {
|
||||
t.Fatalf("unexpected line content: %q", line)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func Test_tailFileLastLines_SeeksToEnd(t *testing.T) {
|
||||
f := writeTempFile(t, "line1\nline2\nline3\n")
|
||||
defer f.Close()
|
||||
tailFileLastLines(f, 1)
|
||||
|
||||
// After tailFileLastLines, file position should be at the end.
|
||||
pos, err := f.Seek(0, io.SeekCurrent)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
stat, err := f.Stat()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if pos != stat.Size() {
|
||||
t.Fatalf("expected file position at end (%d), got %d", stat.Size(), pos)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,473 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/rs/zerolog"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
const (
|
||||
logWriterSize = 1024 * 1024 * 5 // 5 MB
|
||||
logWriterSmallSize = 1024 * 1024 * 1 // 1 MB
|
||||
logWriterInitialSize = 32 * 1024 // 32 KB
|
||||
logWriterSentInterval = time.Minute
|
||||
logWriterInitEndMarker = "\n\n=== INIT_END ===\n\n"
|
||||
logWriterLogEndMarker = "\n\n=== LOG_END ===\n\n"
|
||||
|
||||
logFileName = "ctrld.log"
|
||||
logFileMaxSize = 1024 * 1024 * 5 // 5 MB
|
||||
)
|
||||
|
||||
type logViewResponse struct {
|
||||
Data string `json:"data"`
|
||||
}
|
||||
|
||||
type logSentResponse struct {
|
||||
Size int64 `json:"size"`
|
||||
Error string `json:"error"`
|
||||
}
|
||||
|
||||
type logReader struct {
|
||||
r io.ReadCloser
|
||||
size int64
|
||||
}
|
||||
|
||||
// logSubscriber represents a subscriber to live log output.
|
||||
type logSubscriber struct {
|
||||
ch chan []byte
|
||||
}
|
||||
|
||||
// logWriter is an internal buffer to keep track of runtime log when no logging is enabled.
|
||||
// When a file path is configured via setLogFile, writes are also persisted to
|
||||
// a rotated file on disk (max logFileMaxSize, 1 backup) so logs survive restarts.
|
||||
type logWriter struct {
|
||||
mu sync.Mutex
|
||||
buf bytes.Buffer
|
||||
size int
|
||||
subscribers []*logSubscriber
|
||||
|
||||
// File persistence fields.
|
||||
logFile *os.File
|
||||
logFilePath string
|
||||
logFileSize int64
|
||||
}
|
||||
|
||||
// newLogWriter creates an internal log writer.
|
||||
func newLogWriter() *logWriter {
|
||||
return newLogWriterWithSize(logWriterSize)
|
||||
}
|
||||
|
||||
// newSmallLogWriter creates an internal log writer with small buffer size.
|
||||
func newSmallLogWriter() *logWriter {
|
||||
return newLogWriterWithSize(logWriterSmallSize)
|
||||
}
|
||||
|
||||
// newLogWriterWithSize creates an internal log writer with a given buffer size.
|
||||
func newLogWriterWithSize(size int) *logWriter {
|
||||
lw := &logWriter{size: size}
|
||||
return lw
|
||||
}
|
||||
|
||||
// setLogFile configures file-backed persistence for the log writer.
|
||||
// The directory is created if it does not exist. An existing file is
|
||||
// opened in append mode and its current size is tracked for rotation.
|
||||
func (lw *logWriter) setLogFile(path string) error {
|
||||
dir := filepath.Dir(path)
|
||||
if err := os.MkdirAll(dir, 0750); err != nil {
|
||||
return fmt.Errorf("creating log directory: %w", err)
|
||||
}
|
||||
f, err := os.OpenFile(path, os.O_CREATE|os.O_RDWR|os.O_APPEND, 0600)
|
||||
if err != nil {
|
||||
return fmt.Errorf("opening log file: %w", err)
|
||||
}
|
||||
st, err := f.Stat()
|
||||
if err != nil {
|
||||
f.Close()
|
||||
return fmt.Errorf("stat log file: %w", err)
|
||||
}
|
||||
lw.mu.Lock()
|
||||
defer lw.mu.Unlock()
|
||||
lw.logFile = f
|
||||
lw.logFilePath = path
|
||||
lw.logFileSize = st.Size()
|
||||
return nil
|
||||
}
|
||||
|
||||
// rotateLogFile rotates the current log file to a .1 backup.
|
||||
// It returns true if lw.logFile is usable after the call, false otherwise.
|
||||
// Must be called with lw.mu held.
|
||||
func (lw *logWriter) rotateLogFile() bool {
|
||||
if lw.logFile == nil {
|
||||
return false
|
||||
}
|
||||
lw.logFile.Close()
|
||||
backupPath := lw.logFilePath + ".1"
|
||||
// Best effort: rename current to backup (overwrites old backup).
|
||||
os.Rename(lw.logFilePath, backupPath)
|
||||
f, err := os.OpenFile(lw.logFilePath, os.O_CREATE|os.O_RDWR|os.O_TRUNC, 0600)
|
||||
if err != nil {
|
||||
// If we can't reopen, disable file logging.
|
||||
lw.logFile = nil
|
||||
lw.logFileSize = 0
|
||||
return false
|
||||
}
|
||||
lw.logFile = f
|
||||
lw.logFileSize = 0
|
||||
return true
|
||||
}
|
||||
|
||||
// closeLogFile closes the backing file if open.
|
||||
func (lw *logWriter) closeLogFile() {
|
||||
lw.mu.Lock()
|
||||
defer lw.mu.Unlock()
|
||||
if lw.logFile != nil {
|
||||
lw.logFile.Close()
|
||||
lw.logFile = nil
|
||||
}
|
||||
}
|
||||
|
||||
// logFilePaths returns the paths to the current log file and its backup
|
||||
// (if they exist) for inclusion in log send payloads.
|
||||
func (lw *logWriter) logFilePaths() (current, backup string) {
|
||||
lw.mu.Lock()
|
||||
defer lw.mu.Unlock()
|
||||
if lw.logFilePath == "" {
|
||||
return "", ""
|
||||
}
|
||||
current = lw.logFilePath
|
||||
bp := lw.logFilePath + ".1"
|
||||
if _, err := os.Stat(bp); err == nil {
|
||||
backup = bp
|
||||
}
|
||||
return current, backup
|
||||
}
|
||||
|
||||
// Subscribe returns a channel that receives new log data as it's written,
|
||||
// and an unsubscribe function to clean up when done.
|
||||
func (lw *logWriter) Subscribe() (<-chan []byte, func()) {
|
||||
lw.mu.Lock()
|
||||
defer lw.mu.Unlock()
|
||||
sub := &logSubscriber{ch: make(chan []byte, 256)}
|
||||
lw.subscribers = append(lw.subscribers, sub)
|
||||
unsub := func() {
|
||||
lw.mu.Lock()
|
||||
defer lw.mu.Unlock()
|
||||
for i, s := range lw.subscribers {
|
||||
if s == sub {
|
||||
lw.subscribers = append(lw.subscribers[:i], lw.subscribers[i+1:]...)
|
||||
close(sub.ch)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
return sub.ch, unsub
|
||||
}
|
||||
|
||||
// tailLastLines returns the last n lines from the current buffer.
|
||||
func (lw *logWriter) tailLastLines(n int) []byte {
|
||||
lw.mu.Lock()
|
||||
defer lw.mu.Unlock()
|
||||
data := lw.buf.Bytes()
|
||||
if n <= 0 || len(data) == 0 {
|
||||
return nil
|
||||
}
|
||||
// Find the last n newlines from the end.
|
||||
count := 0
|
||||
pos := len(data)
|
||||
for pos > 0 {
|
||||
pos--
|
||||
if data[pos] == '\n' {
|
||||
count++
|
||||
if count == n+1 {
|
||||
pos++ // move past this newline
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
result := make([]byte, len(data)-pos)
|
||||
copy(result, data[pos:])
|
||||
return result
|
||||
}
|
||||
|
||||
func (lw *logWriter) Write(p []byte) (int, error) {
|
||||
lw.mu.Lock()
|
||||
defer lw.mu.Unlock()
|
||||
|
||||
// Fan-out to subscribers (non-blocking).
|
||||
if len(lw.subscribers) > 0 {
|
||||
cp := make([]byte, len(p))
|
||||
copy(cp, p)
|
||||
for _, sub := range lw.subscribers {
|
||||
select {
|
||||
case sub.ch <- cp:
|
||||
default:
|
||||
// Drop if subscriber is slow to avoid blocking the logger.
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Write to backing file if configured.
|
||||
if lw.logFile != nil {
|
||||
needsRotation := lw.logFileSize+int64(len(p)) > logFileMaxSize
|
||||
if !needsRotation || lw.rotateLogFile() {
|
||||
if n, err := lw.logFile.Write(p); err == nil {
|
||||
lw.logFileSize += int64(n)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// If writing p causes overflows, discard old data.
|
||||
if lw.buf.Len()+len(p) > lw.size {
|
||||
buf := lw.buf.Bytes()
|
||||
haveEndMarker := false
|
||||
// If there's init end marker already, preserve the data til the marker.
|
||||
if idx := bytes.LastIndex(buf, []byte(logWriterInitEndMarker)); idx >= 0 {
|
||||
buf = buf[:idx+len(logWriterInitEndMarker)]
|
||||
haveEndMarker = true
|
||||
} else {
|
||||
// Otherwise, preserve the initial size data.
|
||||
buf = buf[:logWriterInitialSize]
|
||||
if idx := bytes.LastIndex(buf, []byte("\n")); idx != -1 {
|
||||
buf = buf[:idx]
|
||||
}
|
||||
}
|
||||
lw.buf.Reset()
|
||||
lw.buf.Write(buf)
|
||||
if !haveEndMarker {
|
||||
lw.buf.WriteString(logWriterInitEndMarker) // indicate that the log was truncated.
|
||||
}
|
||||
}
|
||||
// If p is bigger than buffer size, truncate p by half until its size is smaller.
|
||||
for len(p)+lw.buf.Len() > lw.size {
|
||||
p = p[len(p)/2:]
|
||||
}
|
||||
return lw.buf.Write(p)
|
||||
}
|
||||
|
||||
// initLogging initializes global logging setup.
|
||||
func (p *prog) initLogging(backup bool) {
|
||||
zerolog.TimeFieldFormat = time.RFC3339 + ".000"
|
||||
logWriters := initLoggingWithBackup(backup)
|
||||
|
||||
// Initializing internal logging after global logging.
|
||||
p.initInternalLogging(logWriters)
|
||||
}
|
||||
|
||||
// internalLogFilePath returns the path for persisted internal logs.
|
||||
// The file lives in the ctrld home directory alongside other runtime state.
|
||||
func internalLogFilePath() string {
|
||||
return absHomeDir(logFileName)
|
||||
}
|
||||
|
||||
// initInternalLogging performs internal logging if there's no log enabled.
|
||||
func (p *prog) initInternalLogging(writers []io.Writer) {
|
||||
if !p.needInternalLogging() {
|
||||
return
|
||||
}
|
||||
p.initInternalLogWriterOnce.Do(func() {
|
||||
mainLog.Load().Notice().Msg("internal logging enabled")
|
||||
p.internalLogWriter = newLogWriter()
|
||||
p.internalLogSent = time.Now().Add(-logWriterSentInterval)
|
||||
p.internalWarnLogWriter = newSmallLogWriter()
|
||||
// Persist internal logs to disk so they survive restarts.
|
||||
if path := internalLogFilePath(); path != "" {
|
||||
if err := p.internalLogWriter.setLogFile(path); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("could not enable persistent internal logging")
|
||||
} else {
|
||||
mainLog.Load().Notice().Msgf("internal log file: %s", path)
|
||||
}
|
||||
}
|
||||
})
|
||||
p.mu.Lock()
|
||||
lw := p.internalLogWriter
|
||||
wlw := p.internalWarnLogWriter
|
||||
p.mu.Unlock()
|
||||
// If ctrld was run without explicit verbose level,
|
||||
// run the internal logging at debug level, so we could
|
||||
// have enough information for troubleshooting.
|
||||
if verbose == 0 {
|
||||
for i := range writers {
|
||||
w := &zerolog.FilteredLevelWriter{
|
||||
Writer: zerolog.LevelWriterAdapter{Writer: writers[i]},
|
||||
Level: zerolog.NoticeLevel,
|
||||
}
|
||||
writers[i] = w
|
||||
}
|
||||
zerolog.SetGlobalLevel(zerolog.DebugLevel)
|
||||
}
|
||||
writers = append(writers, lw)
|
||||
writers = append(writers, &zerolog.FilteredLevelWriter{
|
||||
Writer: zerolog.LevelWriterAdapter{Writer: wlw},
|
||||
Level: zerolog.WarnLevel,
|
||||
})
|
||||
multi := zerolog.MultiLevelWriter(writers...)
|
||||
l := mainLog.Load().Output(multi).With().Logger()
|
||||
mainLog.Store(&l)
|
||||
ctrld.ProxyLogger.Store(&l)
|
||||
}
|
||||
|
||||
// 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
|
||||
}
|
||||
// Do not run if there's already log file.
|
||||
if p.cfg.Service.LogPath != "" {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (p *prog) logReader() (*logReader, error) {
|
||||
if p.needInternalLogging() {
|
||||
p.mu.Lock()
|
||||
lw := p.internalLogWriter
|
||||
wlw := p.internalWarnLogWriter
|
||||
p.mu.Unlock()
|
||||
if lw == nil {
|
||||
return nil, errors.New("nil internal log writer")
|
||||
}
|
||||
if wlw == nil {
|
||||
return nil, errors.New("nil internal warn log writer")
|
||||
}
|
||||
|
||||
// If we have a persisted log file, read from disk (includes data
|
||||
// from previous runs that the in-memory buffer wouldn't have).
|
||||
current, backup := lw.logFilePaths()
|
||||
if current != "" {
|
||||
return p.logReaderFromFiles(current, backup, wlw)
|
||||
}
|
||||
|
||||
// Fall back to in-memory buffer.
|
||||
lw.mu.Lock()
|
||||
lwReader := bytes.NewReader(lw.buf.Bytes())
|
||||
lwSize := lw.buf.Len()
|
||||
lw.mu.Unlock()
|
||||
// Warn log content.
|
||||
wlw.mu.Lock()
|
||||
wlwReader := bytes.NewReader(wlw.buf.Bytes())
|
||||
wlwSize := wlw.buf.Len()
|
||||
wlw.mu.Unlock()
|
||||
reader := io.MultiReader(lwReader, bytes.NewReader([]byte(logWriterLogEndMarker)), wlwReader)
|
||||
lr := &logReader{r: io.NopCloser(reader)}
|
||||
lr.size = int64(lwSize + wlwSize)
|
||||
if lr.size == 0 {
|
||||
return nil, errors.New("internal log is empty")
|
||||
}
|
||||
return lr, nil
|
||||
}
|
||||
if p.cfg.Service.LogPath == "" {
|
||||
return &logReader{r: io.NopCloser(strings.NewReader(""))}, nil
|
||||
}
|
||||
f, err := os.Open(normalizeLogFilePath(p.cfg.Service.LogPath))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
lr := &logReader{r: f}
|
||||
if st, err := f.Stat(); err == nil {
|
||||
lr.size = st.Size()
|
||||
} else {
|
||||
return nil, fmt.Errorf("f.Stat: %w", err)
|
||||
}
|
||||
if lr.size == 0 {
|
||||
return nil, errors.New("log file is empty")
|
||||
}
|
||||
return lr, nil
|
||||
}
|
||||
|
||||
// logReaderFromFiles builds a logReader that concatenates the backup file
|
||||
// (if it exists), the current log file, and the in-memory warn log buffer.
|
||||
func (p *prog) logReaderFromFiles(current, backup string, wlw *logWriter) (*logReader, error) {
|
||||
var rcs []io.ReadCloser
|
||||
var totalSize int64
|
||||
|
||||
closeAll := func() {
|
||||
for _, rc := range rcs {
|
||||
rc.Close()
|
||||
}
|
||||
}
|
||||
|
||||
// Read backup file first (older entries).
|
||||
if backup != "" {
|
||||
if bf, err := os.Open(backup); err == nil {
|
||||
if st, err := bf.Stat(); err == nil {
|
||||
totalSize += st.Size()
|
||||
}
|
||||
rcs = append(rcs, bf)
|
||||
}
|
||||
}
|
||||
|
||||
// Read current file.
|
||||
cf, err := os.Open(current)
|
||||
if err != nil {
|
||||
closeAll()
|
||||
return nil, fmt.Errorf("opening current log file: %w", err)
|
||||
}
|
||||
if st, err := cf.Stat(); err == nil {
|
||||
totalSize += st.Size()
|
||||
}
|
||||
rcs = append(rcs, cf)
|
||||
|
||||
// Append warn log content from memory.
|
||||
wlw.mu.Lock()
|
||||
warnData := make([]byte, wlw.buf.Len())
|
||||
copy(warnData, wlw.buf.Bytes())
|
||||
wlw.mu.Unlock()
|
||||
|
||||
if len(warnData) > 0 {
|
||||
rcs = append(rcs, io.NopCloser(bytes.NewReader([]byte(logWriterLogEndMarker))))
|
||||
rcs = append(rcs, io.NopCloser(bytes.NewReader(warnData)))
|
||||
totalSize += int64(len(logWriterLogEndMarker) + len(warnData))
|
||||
}
|
||||
|
||||
if totalSize == 0 {
|
||||
closeAll()
|
||||
return nil, errors.New("internal log is empty")
|
||||
}
|
||||
|
||||
readers := make([]io.Reader, len(rcs))
|
||||
closers := make([]io.Closer, len(rcs))
|
||||
for i, rc := range rcs {
|
||||
readers[i] = rc
|
||||
closers[i] = rc
|
||||
}
|
||||
combined := io.MultiReader(readers...)
|
||||
lr := &logReader{
|
||||
r: &multiCloser{Reader: combined, closers: closers},
|
||||
size: totalSize,
|
||||
}
|
||||
return lr, nil
|
||||
}
|
||||
|
||||
// multiCloser wraps an io.Reader and closes multiple underlying closers.
|
||||
type multiCloser struct {
|
||||
io.Reader
|
||||
closers []io.Closer
|
||||
}
|
||||
|
||||
func (mc *multiCloser) Close() error {
|
||||
var firstErr error
|
||||
for _, c := range mc.closers {
|
||||
if err := c.Close(); err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
return firstErr
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,210 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func Test_logWriter_Write(t *testing.T) {
|
||||
size := 64 * 1024
|
||||
lw := &logWriter{size: size}
|
||||
lw.buf.Grow(lw.size)
|
||||
data := strings.Repeat("A", size)
|
||||
lw.Write([]byte(data))
|
||||
if lw.buf.String() != data {
|
||||
t.Fatalf("unexpected buf content: %v", lw.buf.String())
|
||||
}
|
||||
newData := "B"
|
||||
halfData := strings.Repeat("A", len(data)/2) + logWriterInitEndMarker
|
||||
lw.Write([]byte(newData))
|
||||
if lw.buf.String() != halfData+newData {
|
||||
t.Fatalf("unexpected new buf content: %v", lw.buf.String())
|
||||
}
|
||||
|
||||
bigData := strings.Repeat("B", 256*1024)
|
||||
expected := halfData + strings.Repeat("B", 16*1024)
|
||||
lw.Write([]byte(bigData))
|
||||
if lw.buf.String() != expected {
|
||||
t.Fatalf("unexpected big buf content: %v", lw.buf.String())
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_ConcurrentWrite(t *testing.T) {
|
||||
size := 64 * 1024
|
||||
lw := &logWriter{size: size}
|
||||
n := 10
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(n)
|
||||
for i := 0; i < n; i++ {
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
lw.Write([]byte(strings.Repeat("A", i)))
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
if lw.buf.Len() > lw.size {
|
||||
t.Fatalf("unexpected buf size: %v, content: %q", lw.buf.Len(), lw.buf.String())
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_MarkerInitEnd(t *testing.T) {
|
||||
size := 64 * 1024
|
||||
lw := &logWriter{size: size}
|
||||
lw.buf.Grow(lw.size)
|
||||
|
||||
paddingSize := 10
|
||||
// Writing half of the size, minus len(end marker) and padding size.
|
||||
dataSize := size/2 - len(logWriterInitEndMarker) - paddingSize
|
||||
data := strings.Repeat("A", dataSize)
|
||||
// Inserting newline for making partial init data
|
||||
data += "\n"
|
||||
// Filling left over buffer to make the log full.
|
||||
// The data length: len(end marker) + padding size - 1 (for newline above) + size/2
|
||||
data += strings.Repeat("A", len(logWriterInitEndMarker)+paddingSize-1+(size/2))
|
||||
lw.Write([]byte(data))
|
||||
if lw.buf.String() != data {
|
||||
t.Fatalf("unexpected buf content: %v", lw.buf.String())
|
||||
}
|
||||
lw.Write([]byte("B"))
|
||||
lw.Write([]byte(strings.Repeat("B", 256*1024)))
|
||||
firstIdx := strings.Index(lw.buf.String(), logWriterInitEndMarker)
|
||||
lastIdx := strings.LastIndex(lw.buf.String(), logWriterInitEndMarker)
|
||||
// Check if init end marker present.
|
||||
if firstIdx == -1 || lastIdx == -1 {
|
||||
t.Fatalf("missing init end marker: %s", lw.buf.String())
|
||||
}
|
||||
// Check if init end marker appears only once.
|
||||
if firstIdx != lastIdx {
|
||||
t.Fatalf("log init end marker appears more than once: %s", lw.buf.String())
|
||||
}
|
||||
// Ensure that we have the correct init log data.
|
||||
if !strings.Contains(lw.buf.String(), strings.Repeat("A", dataSize)+logWriterInitEndMarker) {
|
||||
t.Fatalf("unexpected log content: %s", lw.buf.String())
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_SetLogFile(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "test.log")
|
||||
lw := newLogWriterWithSize(logWriterSize)
|
||||
if err := lw.setLogFile(path); err != nil {
|
||||
t.Fatalf("setLogFile: %v", err)
|
||||
}
|
||||
defer lw.closeLogFile()
|
||||
|
||||
msg := "hello file\n"
|
||||
lw.Write([]byte(msg))
|
||||
|
||||
// Verify data in memory buffer.
|
||||
if lw.buf.String() != msg {
|
||||
t.Fatalf("buffer: got %q, want %q", lw.buf.String(), msg)
|
||||
}
|
||||
// Verify data on disk.
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile: %v", err)
|
||||
}
|
||||
if string(data) != msg {
|
||||
t.Fatalf("file: got %q, want %q", data, msg)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_FileRotation(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "test.log")
|
||||
// Use a tiny max size to trigger rotation quickly.
|
||||
lw := newLogWriterWithSize(logWriterSize)
|
||||
if err := lw.setLogFile(path); err != nil {
|
||||
t.Fatalf("setLogFile: %v", err)
|
||||
}
|
||||
defer lw.closeLogFile()
|
||||
|
||||
// Write enough to exceed logFileMaxSize.
|
||||
chunk := strings.Repeat("X", 1024) + "\n"
|
||||
written := 0
|
||||
for written < logFileMaxSize+1024 {
|
||||
lw.Write([]byte(chunk))
|
||||
written += len(chunk)
|
||||
}
|
||||
|
||||
// Backup file should exist.
|
||||
backupPath := path + ".1"
|
||||
if _, err := os.Stat(backupPath); os.IsNotExist(err) {
|
||||
t.Fatal("expected backup file to exist after rotation")
|
||||
}
|
||||
|
||||
// Current file should be smaller than max (it was rotated).
|
||||
st, err := os.Stat(path)
|
||||
if err != nil {
|
||||
t.Fatalf("stat current: %v", err)
|
||||
}
|
||||
if st.Size() > logFileMaxSize {
|
||||
t.Fatalf("current file too large after rotation: %d", st.Size())
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_FilePaths(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "test.log")
|
||||
lw := newLogWriterWithSize(logWriterSize)
|
||||
|
||||
// No file configured.
|
||||
c, b := lw.logFilePaths()
|
||||
if c != "" || b != "" {
|
||||
t.Fatalf("expected empty paths, got %q %q", c, b)
|
||||
}
|
||||
|
||||
if err := lw.setLogFile(path); err != nil {
|
||||
t.Fatalf("setLogFile: %v", err)
|
||||
}
|
||||
defer lw.closeLogFile()
|
||||
|
||||
// Current exists, no backup yet.
|
||||
c, b = lw.logFilePaths()
|
||||
if c != path {
|
||||
t.Fatalf("current: got %q, want %q", c, path)
|
||||
}
|
||||
if b != "" {
|
||||
t.Fatalf("backup should be empty, got %q", b)
|
||||
}
|
||||
|
||||
// Create a backup file manually.
|
||||
os.WriteFile(path+".1", []byte("old"), 0600)
|
||||
_, b = lw.logFilePaths()
|
||||
if b != path+".1" {
|
||||
t.Fatalf("backup: got %q, want %q", b, path+".1")
|
||||
}
|
||||
}
|
||||
|
||||
func Test_logWriter_FileAppendOnRestart(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "test.log")
|
||||
|
||||
// Simulate first run.
|
||||
lw1 := newLogWriterWithSize(logWriterSize)
|
||||
if err := lw1.setLogFile(path); err != nil {
|
||||
t.Fatalf("setLogFile: %v", err)
|
||||
}
|
||||
lw1.Write([]byte("run1\n"))
|
||||
lw1.closeLogFile()
|
||||
|
||||
// Simulate second run (restart) — file should be appended.
|
||||
lw2 := newLogWriterWithSize(logWriterSize)
|
||||
if err := lw2.setLogFile(path); err != nil {
|
||||
t.Fatalf("setLogFile: %v", err)
|
||||
}
|
||||
lw2.Write([]byte("run2\n"))
|
||||
lw2.closeLogFile()
|
||||
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile: %v", err)
|
||||
}
|
||||
want := "run1\nrun2\n"
|
||||
if string(data) != want {
|
||||
t.Fatalf("file: got %q, want %q", data, want)
|
||||
}
|
||||
}
|
||||
+63
-13
@@ -1,7 +1,9 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"encoding/hex"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync/atomic"
|
||||
@@ -39,6 +41,10 @@ var (
|
||||
skipSelfChecks bool
|
||||
cleanup bool
|
||||
startOnly bool
|
||||
rfc1918 bool
|
||||
interceptMode string // "", "dns", or "hard" — set via --intercept-mode flag or config
|
||||
dnsIntercept bool // derived: interceptMode == "dns" || interceptMode == "hard"
|
||||
hardIntercept bool // derived: interceptMode == "hard"
|
||||
|
||||
mainLog atomic.Pointer[zerolog.Logger]
|
||||
consoleWriter zerolog.ConsoleWriter
|
||||
@@ -50,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() {
|
||||
@@ -58,6 +67,16 @@ func init() {
|
||||
}
|
||||
|
||||
func Main() {
|
||||
// Fast path for pf interception probe subprocess. This runs before cobra
|
||||
// initialization to minimize startup time. The parent process spawns us with
|
||||
// "pf-probe-send <host> <hex-dns-packet>" and a non-_ctrld GID so pf
|
||||
// intercepts the DNS query. If pf rdr is working, the query reaches ctrld's
|
||||
// listener; if not, it goes to the real DNS server and ctrld detects the miss.
|
||||
if len(os.Args) >= 4 && os.Args[1] == "pf-probe-send" {
|
||||
pfProbeSend(os.Args[2], os.Args[3])
|
||||
return
|
||||
}
|
||||
|
||||
ctrld.InitConfig(v, "ctrld")
|
||||
initCLI()
|
||||
if err := rootCmd.Execute(); err != nil {
|
||||
@@ -67,11 +86,8 @@ func Main() {
|
||||
}
|
||||
|
||||
func normalizeLogFilePath(logFilePath string) string {
|
||||
// In cleanup mode, we always want the full log file path.
|
||||
if !cleanup {
|
||||
if logFilePath == "" || filepath.IsAbs(logFilePath) || service.Interactive() {
|
||||
return logFilePath
|
||||
}
|
||||
if logFilePath == "" || filepath.IsAbs(logFilePath) || service.Interactive() {
|
||||
return logFilePath
|
||||
}
|
||||
if homedir != "" {
|
||||
return filepath.Join(homedir, logFilePath)
|
||||
@@ -91,22 +107,33 @@ func initConsoleLogging() {
|
||||
multi := zerolog.MultiLevelWriter(consoleWriter)
|
||||
l := mainLog.Load().Output(multi).With().Timestamp().Logger()
|
||||
mainLog.Store(&l)
|
||||
|
||||
switch {
|
||||
case silent:
|
||||
zerolog.SetGlobalLevel(zerolog.NoLevel)
|
||||
case verbose == 1:
|
||||
ctrld.ProxyLogger.Store(&l)
|
||||
zerolog.SetGlobalLevel(zerolog.InfoLevel)
|
||||
case verbose > 1:
|
||||
ctrld.ProxyLogger.Store(&l)
|
||||
zerolog.SetGlobalLevel(zerolog.DebugLevel)
|
||||
default:
|
||||
zerolog.SetGlobalLevel(zerolog.NoticeLevel)
|
||||
}
|
||||
}
|
||||
|
||||
// initLogging initializes global logging setup.
|
||||
func initLogging() {
|
||||
// initInteractiveLogging is like initLogging, but the ProxyLogger is discarded
|
||||
// to be used for all interactive commands.
|
||||
//
|
||||
// Current log file config will also be ignored.
|
||||
func initInteractiveLogging() {
|
||||
old := cfg.Service.LogPath
|
||||
cfg.Service.LogPath = ""
|
||||
zerolog.TimeFieldFormat = time.RFC3339 + ".000"
|
||||
initLoggingWithBackup(true)
|
||||
initLoggingWithBackup(false)
|
||||
cfg.Service.LogPath = old
|
||||
l := zerolog.New(io.Discard)
|
||||
ctrld.ProxyLogger.Store(&l)
|
||||
}
|
||||
|
||||
// initLoggingWithBackup initializes log setup base on current config.
|
||||
@@ -115,8 +142,8 @@ func initLogging() {
|
||||
// This is only used in runCmd for special handling in case of logging config
|
||||
// change in cd mode. Without special reason, the caller should use initLogging
|
||||
// wrapper instead of calling this function directly.
|
||||
func initLoggingWithBackup(doBackup bool) {
|
||||
writers := []io.Writer{io.Discard}
|
||||
func initLoggingWithBackup(doBackup bool) []io.Writer {
|
||||
var writers []io.Writer
|
||||
if logFilePath := normalizeLogFilePath(cfg.Service.LogPath); logFilePath != "" {
|
||||
// Create parent directory if necessary.
|
||||
if err := os.MkdirAll(filepath.Dir(logFilePath), 0750); err != nil {
|
||||
@@ -154,21 +181,22 @@ func initLoggingWithBackup(doBackup bool) {
|
||||
switch {
|
||||
case silent:
|
||||
zerolog.SetGlobalLevel(zerolog.NoLevel)
|
||||
return
|
||||
return writers
|
||||
case verbose == 1:
|
||||
logLevel = "info"
|
||||
case verbose > 1:
|
||||
logLevel = "debug"
|
||||
}
|
||||
if logLevel == "" {
|
||||
return
|
||||
return writers
|
||||
}
|
||||
level, err := zerolog.ParseLevel(logLevel)
|
||||
if err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("could not set log level")
|
||||
return
|
||||
return writers
|
||||
}
|
||||
zerolog.SetGlobalLevel(level)
|
||||
return writers
|
||||
}
|
||||
|
||||
func initCache() {
|
||||
@@ -179,3 +207,25 @@ func initCache() {
|
||||
cfg.Service.CacheSize = 4096
|
||||
}
|
||||
}
|
||||
|
||||
// pfProbeSend is a minimal subprocess that sends a pre-built DNS query packet
|
||||
// to the specified host on port 53. It's invoked by probePFIntercept() with a
|
||||
// non-_ctrld GID so pf interception applies to the query.
|
||||
//
|
||||
// Usage: ctrld pf-probe-send <host> <hex-encoded-dns-packet>
|
||||
func pfProbeSend(host, hexPacket string) {
|
||||
packet, err := hex.DecodeString(hexPacket)
|
||||
if err != nil {
|
||||
os.Exit(1)
|
||||
}
|
||||
conn, err := net.DialTimeout("udp", net.JoinHostPort(host, "53"), time.Second)
|
||||
if err != nil {
|
||||
os.Exit(1)
|
||||
}
|
||||
defer conn.Close()
|
||||
conn.SetDeadline(time.Now().Add(time.Second))
|
||||
_, _ = conn.Write(packet)
|
||||
// Read response (don't care about result, just need the send to happen)
|
||||
buf := make([]byte, 512)
|
||||
_, _ = conn.Read(buf)
|
||||
}
|
||||
|
||||
+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")
|
||||
|
||||
@@ -9,17 +9,18 @@ import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
func patchNetIfaceName(iface *net.Interface) error {
|
||||
func patchNetIfaceName(iface *net.Interface) (bool, error) {
|
||||
b, err := exec.Command("networksetup", "-listnetworkserviceorder").Output()
|
||||
if err != nil {
|
||||
return err
|
||||
return false, err
|
||||
}
|
||||
|
||||
patched := false
|
||||
if name := networkServiceName(iface.Name, bytes.NewReader(b)); name != "" {
|
||||
patched = true
|
||||
iface.Name = name
|
||||
mainLog.Load().Debug().Str("network_service", name).Msg("found network service name for interface")
|
||||
}
|
||||
return nil
|
||||
return patched, nil
|
||||
}
|
||||
|
||||
func networkServiceName(ifaceName string, r io.Reader) string {
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"tailscale.com/net/netmon"
|
||||
)
|
||||
|
||||
func patchNetIfaceName(iface *net.Interface) (bool, error) { return true, nil }
|
||||
|
||||
// validInterface reports whether the *net.Interface is a valid one.
|
||||
// Only non-virtual interfaces are considered valid.
|
||||
func validInterface(iface *net.Interface, validIfacesMap map[string]struct{}) bool {
|
||||
_, ok := validIfacesMap[iface.Name]
|
||||
return ok
|
||||
}
|
||||
|
||||
// validInterfacesMap returns a set containing non virtual interfaces.
|
||||
func validInterfacesMap() map[string]struct{} {
|
||||
m := make(map[string]struct{})
|
||||
vis := virtualInterfaces()
|
||||
netmon.ForeachInterface(func(i netmon.Interface, prefixes []netip.Prefix) {
|
||||
if _, existed := vis[i.Name]; existed {
|
||||
return
|
||||
}
|
||||
m[i.Name] = struct{}{}
|
||||
})
|
||||
// Fallback to default route interface if found nothing.
|
||||
if len(m) == 0 {
|
||||
defaultRoute, err := netmon.DefaultRoute()
|
||||
if err != nil {
|
||||
return m
|
||||
}
|
||||
m[defaultRoute.InterfaceName] = struct{}{}
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// virtualInterfaces returns a map of virtual interfaces on current machine.
|
||||
func virtualInterfaces() map[string]struct{} {
|
||||
s := make(map[string]struct{})
|
||||
entries, _ := os.ReadDir("/sys/devices/virtual/net")
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() {
|
||||
s[strings.TrimSpace(entry.Name())] = struct{}{}
|
||||
}
|
||||
}
|
||||
return s
|
||||
}
|
||||
+15
-4
@@ -1,11 +1,22 @@
|
||||
//go:build !darwin && !windows
|
||||
//go:build !darwin && !windows && !linux
|
||||
|
||||
package cli
|
||||
|
||||
import "net"
|
||||
import (
|
||||
"net"
|
||||
|
||||
func patchNetIfaceName(iface *net.Interface) error { return nil }
|
||||
"tailscale.com/net/netmon"
|
||||
)
|
||||
|
||||
func patchNetIfaceName(iface *net.Interface) (bool, error) { return true, nil }
|
||||
|
||||
func validInterface(iface *net.Interface, validIfacesMap map[string]struct{}) bool { return true }
|
||||
|
||||
func validInterfacesMap() map[string]struct{} { return nil }
|
||||
// validInterfacesMap returns a set containing only default route interfaces.
|
||||
func validInterfacesMap() map[string]struct{} {
|
||||
defaultRoute, err := netmon.DefaultRoute()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return map[string]struct{}{defaultRoute.InterfaceName: {}}
|
||||
}
|
||||
|
||||
+71
-12
@@ -1,14 +1,20 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"io"
|
||||
"log"
|
||||
"net"
|
||||
"strings"
|
||||
"os"
|
||||
|
||||
"github.com/microsoft/wmi/pkg/base/host"
|
||||
"github.com/microsoft/wmi/pkg/base/instance"
|
||||
"github.com/microsoft/wmi/pkg/base/query"
|
||||
"github.com/microsoft/wmi/pkg/constant"
|
||||
"github.com/microsoft/wmi/pkg/hardware/network/netadapter"
|
||||
)
|
||||
|
||||
func patchNetIfaceName(iface *net.Interface) error {
|
||||
return nil
|
||||
func patchNetIfaceName(iface *net.Interface) (bool, error) {
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// validInterface reports whether the *net.Interface is a valid one.
|
||||
@@ -20,15 +26,68 @@ func validInterface(iface *net.Interface, validIfacesMap map[string]struct{}) bo
|
||||
|
||||
// validInterfacesMap returns a set of all physical interfaces.
|
||||
func validInterfacesMap() map[string]struct{} {
|
||||
out, err := powershell("Get-NetAdapter -Physical | Select-Object -ExpandProperty Name")
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
m := make(map[string]struct{})
|
||||
scanner := bufio.NewScanner(bytes.NewReader(out))
|
||||
for scanner.Scan() {
|
||||
ifaceName := strings.TrimSpace(scanner.Text())
|
||||
for _, ifaceName := range validInterfaces() {
|
||||
m[ifaceName] = struct{}{}
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// validInterfaces returns a list of all physical interfaces.
|
||||
func validInterfaces() []string {
|
||||
log.SetOutput(io.Discard)
|
||||
defer log.SetOutput(os.Stderr)
|
||||
whost := host.NewWmiLocalHost()
|
||||
q := query.NewWmiQuery("MSFT_NetAdapter")
|
||||
instances, err := instance.GetWmiInstancesFromHost(whost, string(constant.StadardCimV2), q)
|
||||
if instances != nil {
|
||||
defer instances.Close()
|
||||
}
|
||||
if err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("failed to get wmi network adapter")
|
||||
return nil
|
||||
}
|
||||
var adapters []string
|
||||
for _, i := range instances {
|
||||
adapter, err := netadapter.NewNetworkAdapter(i)
|
||||
if err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("failed to get network adapter")
|
||||
continue
|
||||
}
|
||||
|
||||
name, err := adapter.GetPropertyName()
|
||||
if err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("failed to get interface name")
|
||||
continue
|
||||
}
|
||||
|
||||
// From: https://learn.microsoft.com/en-us/previous-versions/windows/desktop/legacy/hh968170(v=vs.85)
|
||||
//
|
||||
// "Indicates if a connector is present on the network adapter. This value is set to TRUE
|
||||
// if this is a physical adapter or FALSE if this is not a physical adapter."
|
||||
physical, err := adapter.GetPropertyConnectorPresent()
|
||||
if err != nil {
|
||||
mainLog.Load().Debug().Str("method", "validInterfaces").Str("interface", name).Msg("failed to get network adapter connector present property")
|
||||
continue
|
||||
}
|
||||
if !physical {
|
||||
mainLog.Load().Debug().Str("method", "validInterfaces").Str("interface", name).Msg("skipping non-physical adapter")
|
||||
continue
|
||||
}
|
||||
|
||||
// Check if it's a hardware interface. Checking only for connector present is not enough
|
||||
// because some interfaces are not physical but have a connector.
|
||||
hardware, err := adapter.GetPropertyHardwareInterface()
|
||||
if err != nil {
|
||||
mainLog.Load().Debug().Str("method", "validInterfaces").Str("interface", name).Msg("failed to get network adapter hardware interface property")
|
||||
continue
|
||||
}
|
||||
if !hardware {
|
||||
mainLog.Load().Debug().Str("method", "validInterfaces").Str("interface", name).Msg("skipping non-hardware interface")
|
||||
continue
|
||||
}
|
||||
|
||||
adapters = append(adapters, name)
|
||||
}
|
||||
return adapters
|
||||
}
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func Test_validInterfaces(t *testing.T) {
|
||||
verbose = 3
|
||||
initConsoleLogging()
|
||||
start := time.Now()
|
||||
ifaces := validInterfaces()
|
||||
t.Logf("Using Windows API takes: %d", time.Since(start).Milliseconds())
|
||||
|
||||
start = time.Now()
|
||||
ifacesPowershell := validInterfacesPowershell()
|
||||
t.Logf("Using Powershell takes: %d", time.Since(start).Milliseconds())
|
||||
|
||||
slices.Sort(ifaces)
|
||||
slices.Sort(ifacesPowershell)
|
||||
if !slices.Equal(ifaces, ifacesPowershell) {
|
||||
t.Fatalf("result mismatch, want: %v, got: %v", ifacesPowershell, ifaces)
|
||||
}
|
||||
}
|
||||
|
||||
func validInterfacesPowershell() []string {
|
||||
out, err := powershell("Get-NetAdapter -Physical | Select-Object -ExpandProperty Name")
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
var res []string
|
||||
scanner := bufio.NewScanner(bytes.NewReader(out))
|
||||
for scanner.Scan() {
|
||||
ifaceName := strings.TrimSpace(scanner.Text())
|
||||
res = append(res, ifaceName)
|
||||
}
|
||||
return res
|
||||
}
|
||||
@@ -1,34 +0,0 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/vishvananda/netlink"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
func (p *prog) watchLinkState(ctx context.Context) {
|
||||
ch := make(chan netlink.LinkUpdate)
|
||||
done := make(chan struct{})
|
||||
defer close(done)
|
||||
if err := netlink.LinkSubscribe(ch, done); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("could not subscribe link")
|
||||
return
|
||||
}
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case lu := <-ch:
|
||||
if lu.Change == 0xFFFFFFFF {
|
||||
continue
|
||||
}
|
||||
if lu.Change&unix.IFF_UP != 0 {
|
||||
mainLog.Load().Debug().Msgf("link state changed, re-bootstrapping")
|
||||
for _, uc := range p.cfg.Upstream {
|
||||
uc.ReBootstrap()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,7 +0,0 @@
|
||||
//go:build !linux
|
||||
|
||||
package cli
|
||||
|
||||
import "context"
|
||||
|
||||
func (p *prog) watchLinkState(ctx context.Context) {}
|
||||
@@ -0,0 +1,5 @@
|
||||
//go:build !cgo
|
||||
|
||||
package cli
|
||||
|
||||
const cgoEnabled = false
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+13
-6
@@ -47,6 +47,9 @@ func setDnsIgnoreUnusableInterface(iface *net.Interface, nameservers []string) e
|
||||
// networksetup -setdnsservers Wi-Fi 8.8.8.8 1.1.1.1
|
||||
// TODO(cuonglm): use system API
|
||||
func setDNS(iface *net.Interface, nameservers []string) error {
|
||||
// Note that networksetup won't modify search domains settings,
|
||||
// This assignment is just a placeholder to silent linter.
|
||||
_ = searchDomains
|
||||
cmd := "networksetup"
|
||||
args := []string{"-setdnsservers", iface.Name}
|
||||
args = append(args, nameservers...)
|
||||
@@ -70,11 +73,6 @@ func resetDnsIgnoreUnusableInterface(iface *net.Interface) error {
|
||||
|
||||
// TODO(cuonglm): use system API
|
||||
func resetDNS(iface *net.Interface) error {
|
||||
if ns := savedStaticNameservers(iface); len(ns) > 0 {
|
||||
if err := setDNS(iface, ns); err == nil {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
cmd := "networksetup"
|
||||
args := []string{"-setdnsservers", iface.Name, "empty"}
|
||||
if out, err := exec.Command(cmd, args...).CombinedOutput(); err != nil {
|
||||
@@ -83,8 +81,17 @@ func resetDNS(iface *net.Interface) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// restoreDNS restores the DNS settings of the given interface.
|
||||
// this should only be executed upon turning off the ctrld service.
|
||||
func restoreDNS(iface *net.Interface) (err error) {
|
||||
if ns := savedStaticNameservers(iface); len(ns) > 0 {
|
||||
err = setDNS(iface, ns)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func currentDNS(_ *net.Interface) []string {
|
||||
return resolvconffile.NameServers("")
|
||||
return resolvconffile.NameServers()
|
||||
}
|
||||
|
||||
// currentStaticDNS returns the current static DNS settings of given interface.
|
||||
|
||||
+23
-7
@@ -5,7 +5,9 @@ import (
|
||||
"net/netip"
|
||||
"os/exec"
|
||||
|
||||
"tailscale.com/tsd"
|
||||
"tailscale.com/control/controlknobs"
|
||||
"tailscale.com/health"
|
||||
"tailscale.com/util/dnsname"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld/internal/dns"
|
||||
"github.com/Control-D-Inc/ctrld/internal/resolvconffile"
|
||||
@@ -38,8 +40,7 @@ func setDnsIgnoreUnusableInterface(iface *net.Interface, nameservers []string) e
|
||||
|
||||
// set the dns server for the provided network interface
|
||||
func setDNS(iface *net.Interface, nameservers []string) error {
|
||||
sys := new(tsd.System)
|
||||
r, err := dns.NewOSConfigurator(logf, sys.HealthTracker(), sys.ControlKnobs(), iface.Name)
|
||||
r, err := dns.NewOSConfigurator(logf, &health.Tracker{}, &controlknobs.Knobs{}, iface.Name)
|
||||
if err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to create DNS OS configurator")
|
||||
return err
|
||||
@@ -50,7 +51,17 @@ func setDNS(iface *net.Interface, nameservers []string) error {
|
||||
ns = append(ns, netip.MustParseAddr(nameserver))
|
||||
}
|
||||
|
||||
if err := r.SetDNS(dns.OSConfig{Nameservers: ns}); err != nil {
|
||||
osConfig := dns.OSConfig{
|
||||
Nameservers: ns,
|
||||
SearchDomains: []dnsname.FQDN{},
|
||||
}
|
||||
if sds, err := searchDomains(); err == nil {
|
||||
osConfig.SearchDomains = sds
|
||||
} else {
|
||||
mainLog.Load().Debug().Err(err).Msg("failed to get search domains list")
|
||||
}
|
||||
|
||||
if err := r.SetDNS(osConfig); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to set DNS")
|
||||
return err
|
||||
}
|
||||
@@ -63,8 +74,7 @@ func resetDnsIgnoreUnusableInterface(iface *net.Interface) error {
|
||||
}
|
||||
|
||||
func resetDNS(iface *net.Interface) error {
|
||||
sys := new(tsd.System)
|
||||
r, err := dns.NewOSConfigurator(logf, sys.HealthTracker(), sys.ControlKnobs(), iface.Name)
|
||||
r, err := dns.NewOSConfigurator(logf, &health.Tracker{}, &controlknobs.Knobs{}, iface.Name)
|
||||
if err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to create DNS OS configurator")
|
||||
return err
|
||||
@@ -77,8 +87,14 @@ func resetDNS(iface *net.Interface) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// restoreDNS restores the DNS settings of the given interface.
|
||||
// this should only be executed upon turning off the ctrld service.
|
||||
func restoreDNS(iface *net.Interface) (err error) {
|
||||
return err
|
||||
}
|
||||
|
||||
func currentDNS(_ *net.Interface) []string {
|
||||
return resolvconffile.NameServers("")
|
||||
return resolvconffile.NameServers()
|
||||
}
|
||||
|
||||
// currentStaticDNS returns the current static DNS settings of given interface.
|
||||
|
||||
+44
-32
@@ -14,11 +14,11 @@ import (
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"tailscale.com/tsd"
|
||||
|
||||
"github.com/insomniacslk/dhcp/dhcpv4/nclient4"
|
||||
"github.com/insomniacslk/dhcp/dhcpv6"
|
||||
"github.com/insomniacslk/dhcp/dhcpv6/client6"
|
||||
"tailscale.com/control/controlknobs"
|
||||
"tailscale.com/health"
|
||||
"tailscale.com/util/dnsname"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld/internal/dns"
|
||||
@@ -56,8 +56,7 @@ func setDnsIgnoreUnusableInterface(iface *net.Interface, nameservers []string) e
|
||||
}
|
||||
|
||||
func setDNS(iface *net.Interface, nameservers []string) error {
|
||||
sys := new(tsd.System)
|
||||
r, err := dns.NewOSConfigurator(logf, sys.HealthTracker(), sys.ControlKnobs(), iface.Name)
|
||||
r, err := dns.NewOSConfigurator(logf, &health.Tracker{}, &controlknobs.Knobs{}, iface.Name)
|
||||
if err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to create DNS OS configurator")
|
||||
return err
|
||||
@@ -72,35 +71,39 @@ func setDNS(iface *net.Interface, nameservers []string) error {
|
||||
Nameservers: ns,
|
||||
SearchDomains: []dnsname.FQDN{},
|
||||
}
|
||||
if sds, err := searchDomains(); err == nil {
|
||||
// Filter the root domain, since it's not allowed by systemd.
|
||||
// See https://github.com/systemd/systemd/issues/9515
|
||||
filteredSds := slices.DeleteFunc(sds, func(s dnsname.FQDN) bool {
|
||||
return s == "" || s == "."
|
||||
})
|
||||
if len(filteredSds) != len(sds) {
|
||||
mainLog.Load().Debug().Msg(`Removed root domain "." from search domains list`)
|
||||
}
|
||||
osConfig.SearchDomains = filteredSds
|
||||
} else {
|
||||
mainLog.Load().Debug().Err(err).Msg("failed to get search domains list")
|
||||
}
|
||||
trySystemdResolve := false
|
||||
for i := 0; i < maxSetDNSAttempts; i++ {
|
||||
if err := r.SetDNS(osConfig); err != nil {
|
||||
if strings.Contains(err.Error(), "Rejected send message") &&
|
||||
strings.Contains(err.Error(), "org.freedesktop.network1.Manager") {
|
||||
mainLog.Load().Warn().Msg("Interfaces are managed by systemd-networkd, switch to systemd-resolve for setting DNS")
|
||||
trySystemdResolve = true
|
||||
break
|
||||
}
|
||||
// This error happens on read-only file system, which causes ctrld failed to create backup
|
||||
// for /etc/resolv.conf file. It is ok, because the DNS is still set anyway, and restore
|
||||
// DNS will fallback to use DHCP if there's no backup /etc/resolv.conf file.
|
||||
// The error format is controlled by us, so checking for error string is fine.
|
||||
// See: ../../internal/dns/direct.go:L278
|
||||
if r.Mode() == "direct" && strings.Contains(err.Error(), resolvConfBackupFailedMsg) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
if err := r.SetDNS(osConfig); err != nil {
|
||||
if strings.Contains(err.Error(), "Rejected send message") &&
|
||||
strings.Contains(err.Error(), "org.freedesktop.network1.Manager") {
|
||||
mainLog.Load().Warn().Msg("Interfaces are managed by systemd-networkd, switch to systemd-resolve for setting DNS")
|
||||
trySystemdResolve = true
|
||||
goto systemdResolve
|
||||
}
|
||||
if useSystemdResolved {
|
||||
if out, err := exec.Command("systemctl", "restart", "systemd-resolved").CombinedOutput(); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msgf("could not restart systemd-resolved: %s", string(out))
|
||||
}
|
||||
}
|
||||
currentNS := currentDNS(iface)
|
||||
if isSubSet(nameservers, currentNS) {
|
||||
// This error happens on read-only file system, which causes ctrld failed to create backup
|
||||
// for /etc/resolv.conf file. It is ok, because the DNS is still set anyway, and restore
|
||||
// DNS will fallback to use DHCP if there's no backup /etc/resolv.conf file.
|
||||
// The error format is controlled by us, so checking for error string is fine.
|
||||
// See: ../../internal/dns/direct.go:L278
|
||||
if r.Mode() == "direct" && strings.Contains(err.Error(), resolvConfBackupFailedMsg) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
systemdResolve:
|
||||
if trySystemdResolve {
|
||||
// Stop systemd-networkd and retry setting DNS.
|
||||
if out, err := exec.Command("systemctl", "stop", "systemd-networkd").CombinedOutput(); err != nil {
|
||||
@@ -120,8 +123,8 @@ func setDNS(iface *net.Interface, nameservers []string) error {
|
||||
}
|
||||
time.Sleep(time.Second)
|
||||
}
|
||||
mainLog.Load().Debug().Msg("DNS was not set for some reason")
|
||||
}
|
||||
mainLog.Load().Debug().Msg("DNS was not set for some reason")
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -139,8 +142,7 @@ func resetDNS(iface *net.Interface) (err error) {
|
||||
if exe, _ := exec.LookPath("/lib/systemd/systemd-networkd"); exe != "" {
|
||||
_ = exec.Command("systemctl", "start", "systemd-networkd").Run()
|
||||
}
|
||||
sys := new(tsd.System)
|
||||
if r, oerr := dns.NewOSConfigurator(logf, sys.HealthTracker(), sys.ControlKnobs(), iface.Name); oerr == nil {
|
||||
if r, oerr := dns.NewOSConfigurator(logf, &health.Tracker{}, &controlknobs.Knobs{}, iface.Name); oerr == nil {
|
||||
_ = r.SetDNS(dns.OSConfig{})
|
||||
if err := r.Close(); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to rollback DNS setting")
|
||||
@@ -171,6 +173,7 @@ func resetDNS(iface *net.Interface) (err error) {
|
||||
}
|
||||
|
||||
// TODO(cuonglm): handle DHCPv6 properly.
|
||||
mainLog.Load().Debug().Msg("checking for IPv6 availability")
|
||||
if ctrldnet.IPv6Available(ctx) {
|
||||
c := client6.NewClient()
|
||||
conversation, err := c.Exchange(iface.Name)
|
||||
@@ -190,6 +193,8 @@ func resetDNS(iface *net.Interface) (err error) {
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
mainLog.Load().Debug().Msg("IPv6 is not available")
|
||||
}
|
||||
|
||||
return ignoringEINTR(func() error {
|
||||
@@ -197,8 +202,15 @@ func resetDNS(iface *net.Interface) (err error) {
|
||||
})
|
||||
}
|
||||
|
||||
// restoreDNS restores the DNS settings of the given interface.
|
||||
// this should only be executed upon turning off the ctrld service.
|
||||
func restoreDNS(iface *net.Interface) (err error) {
|
||||
return err
|
||||
}
|
||||
|
||||
func currentDNS(iface *net.Interface) []string {
|
||||
for _, fn := range []getDNS{getDNSByResolvectl, getDNSBySystemdResolved, getDNSByNmcli, resolvconffile.NameServers} {
|
||||
resolvconfFunc := func(_ string) []string { return resolvconffile.NameServers() }
|
||||
for _, fn := range []getDNS{getDNSByResolvectl, getDNSBySystemdResolved, getDNSByNmcli, resolvconfFunc} {
|
||||
if ns := fn(iface.Name); len(ns) > 0 {
|
||||
return ns
|
||||
}
|
||||
|
||||
+159
-45
@@ -1,23 +1,27 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"os/exec"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"golang.org/x/sys/windows"
|
||||
"golang.org/x/sys/windows/registry"
|
||||
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
||||
|
||||
ctrldnet "github.com/Control-D-Inc/ctrld/internal/net"
|
||||
)
|
||||
|
||||
const (
|
||||
v4InterfaceKeyPathFormat = `HKLM:\SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces\`
|
||||
v6InterfaceKeyPathFormat = `HKLM:\SYSTEM\CurrentControlSet\Services\Tcpip6\Parameters\Interfaces\`
|
||||
v4InterfaceKeyPathFormat = `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces\`
|
||||
v6InterfaceKeyPathFormat = `SYSTEM\CurrentControlSet\Services\Tcpip6\Parameters\Interfaces\`
|
||||
)
|
||||
|
||||
var (
|
||||
@@ -30,14 +34,6 @@ func setDnsIgnoreUnusableInterface(iface *net.Interface, nameservers []string) e
|
||||
return setDNS(iface, nameservers)
|
||||
}
|
||||
|
||||
func setDnsPowershellCmd(iface *net.Interface, nameservers []string) string {
|
||||
nss := make([]string, 0, len(nameservers))
|
||||
for _, ns := range nameservers {
|
||||
nss = append(nss, strconv.Quote(ns))
|
||||
}
|
||||
return fmt.Sprintf("Set-DnsClientServerAddress -InterfaceIndex %d -ServerAddresses (%s)", iface.Index, strings.Join(nss, ","))
|
||||
}
|
||||
|
||||
// setDNS sets the dns server for the provided network interface
|
||||
func setDNS(iface *net.Interface, nameservers []string) error {
|
||||
if len(nameservers) == 0 {
|
||||
@@ -46,28 +42,80 @@ func setDNS(iface *net.Interface, nameservers []string) error {
|
||||
setDNSOnce.Do(func() {
|
||||
// If there's a Dns server running, that means we are on AD with Dns feature enabled.
|
||||
// Configuring the Dns server to forward queries to ctrld instead.
|
||||
if windowsHasLocalDnsServerRunning() {
|
||||
if hasLocalDnsServerRunning() {
|
||||
mainLog.Load().Debug().Msg("Local DNS server detected, configuring forwarders")
|
||||
|
||||
file := absHomeDir(windowsForwardersFilename)
|
||||
oldForwardersContent, _ := os.ReadFile(file)
|
||||
hasLocalIPv6Listener := needLocalIPv6Listener()
|
||||
mainLog.Load().Debug().Msgf("Using forwarders file: %s", file)
|
||||
|
||||
oldForwardersContent, err := os.ReadFile(file)
|
||||
if err != nil {
|
||||
mainLog.Load().Debug().Err(err).Msg("Could not read existing forwarders file")
|
||||
} else {
|
||||
mainLog.Load().Debug().Msgf("Existing forwarders content: %s", string(oldForwardersContent))
|
||||
}
|
||||
|
||||
hasLocalIPv6Listener := needLocalIPv6Listener(interceptMode)
|
||||
mainLog.Load().Debug().Bool("has_ipv6_listener", hasLocalIPv6Listener).Msg("IPv6 listener status")
|
||||
|
||||
forwarders := slices.DeleteFunc(slices.Clone(nameservers), func(s string) bool {
|
||||
if !hasLocalIPv6Listener {
|
||||
return false
|
||||
}
|
||||
return s == "::1"
|
||||
})
|
||||
mainLog.Load().Debug().Strs("forwarders", forwarders).Msg("Filtered forwarders list")
|
||||
|
||||
if err := os.WriteFile(file, []byte(strings.Join(forwarders, ",")), 0600); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("could not save forwarders settings")
|
||||
} else {
|
||||
mainLog.Load().Debug().Msg("Successfully wrote new forwarders file")
|
||||
}
|
||||
|
||||
oldForwarders := strings.Split(string(oldForwardersContent), ",")
|
||||
mainLog.Load().Debug().Strs("old_forwarders", oldForwarders).Msg("Previous forwarders")
|
||||
|
||||
if err := addDnsServerForwarders(forwarders, oldForwarders); err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("could not set forwarders settings")
|
||||
} else {
|
||||
mainLog.Load().Debug().Msg("Successfully configured DNS server forwarders")
|
||||
}
|
||||
}
|
||||
})
|
||||
out, err := powershell(setDnsPowershellCmd(iface, nameservers))
|
||||
luid, err := winipcfg.LUIDFromIndex(uint32(iface.Index))
|
||||
if err != nil {
|
||||
return fmt.Errorf("%w: %s", err, string(out))
|
||||
return fmt.Errorf("setDNS: %w", err)
|
||||
}
|
||||
var (
|
||||
serversV4 []netip.Addr
|
||||
serversV6 []netip.Addr
|
||||
)
|
||||
for _, ns := range nameservers {
|
||||
if addr, err := netip.ParseAddr(ns); err == nil {
|
||||
if addr.Is4() {
|
||||
serversV4 = append(serversV4, addr)
|
||||
} else {
|
||||
serversV6 = append(serversV6, addr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Note that Windows won't modify the current search domains if passing nil to luid.SetDNS function.
|
||||
// searchDomains is still implemented for Windows just in case Windows API changes in future versions.
|
||||
_ = searchDomains
|
||||
|
||||
if len(serversV4) == 0 && len(serversV6) == 0 {
|
||||
return errors.New("invalid DNS nameservers")
|
||||
}
|
||||
if len(serversV4) > 0 {
|
||||
if err := luid.SetDNS(windows.AF_INET, serversV4, nil); err != nil {
|
||||
return fmt.Errorf("could not set DNS ipv4: %w", err)
|
||||
}
|
||||
}
|
||||
if len(serversV6) > 0 {
|
||||
if err := luid.SetDNS(windows.AF_INET6, serversV6, nil); err != nil {
|
||||
return fmt.Errorf("could not set DNS ipv6: %w", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -81,7 +129,7 @@ func resetDnsIgnoreUnusableInterface(iface *net.Interface) error {
|
||||
func resetDNS(iface *net.Interface) error {
|
||||
resetDNSOnce.Do(func() {
|
||||
// See corresponding comment in setDNS.
|
||||
if windowsHasLocalDnsServerRunning() {
|
||||
if hasLocalDnsServerRunning() {
|
||||
file := absHomeDir(windowsForwardersFilename)
|
||||
content, err := os.ReadFile(file)
|
||||
if err != nil {
|
||||
@@ -96,14 +144,23 @@ func resetDNS(iface *net.Interface) error {
|
||||
}
|
||||
})
|
||||
|
||||
// Restoring DHCP settings.
|
||||
cmd := fmt.Sprintf("Set-DnsClientServerAddress -InterfaceIndex %d -ResetServerAddresses", iface.Index)
|
||||
out, err := powershell(cmd)
|
||||
luid, err := winipcfg.LUIDFromIndex(uint32(iface.Index))
|
||||
if err != nil {
|
||||
return fmt.Errorf("%w: %s", err, string(out))
|
||||
return fmt.Errorf("resetDNS: %w", err)
|
||||
}
|
||||
// Restoring DHCP settings.
|
||||
if err := luid.SetDNS(windows.AF_INET, nil, nil); err != nil {
|
||||
return fmt.Errorf("could not reset DNS ipv4: %w", err)
|
||||
}
|
||||
if err := luid.SetDNS(windows.AF_INET6, nil, nil); err != nil {
|
||||
return fmt.Errorf("could not reset DNS ipv6: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// If there's static DNS saved, restoring it.
|
||||
// restoreDNS restores the DNS settings of the given interface.
|
||||
// this should only be executed upon turning off the ctrld service.
|
||||
func restoreDNS(iface *net.Interface) (err error) {
|
||||
if nss := savedStaticNameservers(iface); len(nss) > 0 {
|
||||
v4ns := make([]string, 0, 2)
|
||||
v6ns := make([]string, 0, 2)
|
||||
@@ -115,17 +172,36 @@ func resetDNS(iface *net.Interface) error {
|
||||
}
|
||||
}
|
||||
|
||||
for _, ns := range [][]string{v4ns, v6ns} {
|
||||
if len(ns) == 0 {
|
||||
continue
|
||||
luid, err := winipcfg.LUIDFromIndex(uint32(iface.Index))
|
||||
if err != nil {
|
||||
return fmt.Errorf("restoreDNS: %w", err)
|
||||
}
|
||||
|
||||
if len(v4ns) > 0 {
|
||||
mainLog.Load().Debug().Msgf("restoring IPv4 static DNS for interface %q: %v", iface.Name, v4ns)
|
||||
if err := setDNS(iface, v4ns); err != nil {
|
||||
return fmt.Errorf("restoreDNS (IPv4): %w", err)
|
||||
}
|
||||
mainLog.Load().Debug().Msgf("setting static DNS for interface %q", iface.Name)
|
||||
if err := setDNS(iface, ns); err != nil {
|
||||
return err
|
||||
} else {
|
||||
mainLog.Load().Debug().Msgf("restoring IPv4 DHCP for interface %q", iface.Name)
|
||||
if err := luid.SetDNS(windows.AF_INET, nil, nil); err != nil {
|
||||
return fmt.Errorf("restoreDNS (IPv4 clear): %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if len(v6ns) > 0 {
|
||||
mainLog.Load().Debug().Msgf("restoring IPv6 static DNS for interface %q: %v", iface.Name, v6ns)
|
||||
if err := setDNS(iface, v6ns); err != nil {
|
||||
return fmt.Errorf("restoreDNS (IPv6): %w", err)
|
||||
}
|
||||
} else {
|
||||
mainLog.Load().Debug().Msgf("restoring IPv6 DHCP for interface %q", iface.Name)
|
||||
if err := luid.SetDNS(windows.AF_INET6, nil, nil); err != nil {
|
||||
return fmt.Errorf("restoreDNS (IPv6 clear): %w", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
return err
|
||||
}
|
||||
|
||||
func currentDNS(iface *net.Interface) []string {
|
||||
@@ -146,37 +222,69 @@ func currentDNS(iface *net.Interface) []string {
|
||||
return ns
|
||||
}
|
||||
|
||||
// currentStaticDNS returns the current static DNS settings of given interface.
|
||||
// currentStaticDNS checks both the IPv4 and IPv6 paths for static DNS values using keys
|
||||
// like "NameServer" and "ProfileNameServer".
|
||||
func currentStaticDNS(iface *net.Interface) ([]string, error) {
|
||||
luid, err := winipcfg.LUIDFromIndex(uint32(iface.Index))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, fmt.Errorf("fallback winipcfg.LUIDFromIndex: %w", err)
|
||||
}
|
||||
guid, err := luid.GUID()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, fmt.Errorf("fallback luid.GUID: %w", err)
|
||||
}
|
||||
|
||||
var ns []string
|
||||
for _, path := range []string{v4InterfaceKeyPathFormat, v6InterfaceKeyPathFormat} {
|
||||
keyPaths := []string{v4InterfaceKeyPathFormat, v6InterfaceKeyPathFormat}
|
||||
for _, path := range keyPaths {
|
||||
interfaceKeyPath := path + guid.String()
|
||||
found := false
|
||||
for _, key := range []string{"NameServer", "ProfileNameServer"} {
|
||||
if found {
|
||||
continue
|
||||
}
|
||||
cmd := fmt.Sprintf(`Get-ItemPropertyValue -Path "%s" -Name "%s"`, interfaceKeyPath, key)
|
||||
out, err := powershell(cmd)
|
||||
if err == nil && len(out) > 0 {
|
||||
found = true
|
||||
for _, e := range strings.Split(string(out), ",") {
|
||||
ns = append(ns, strings.TrimRight(e, "\x00"))
|
||||
k, err := registry.OpenKey(registry.LOCAL_MACHINE, interfaceKeyPath, registry.QUERY_VALUE)
|
||||
if err != nil {
|
||||
mainLog.Load().Debug().Err(err).Msgf("failed to open registry key %q for interface %q; trying next key", interfaceKeyPath, iface.Name)
|
||||
continue
|
||||
}
|
||||
func() {
|
||||
defer k.Close()
|
||||
for _, keyName := range []string{"NameServer", "ProfileNameServer"} {
|
||||
value, _, err := k.GetStringValue(keyName)
|
||||
if err != nil && !errors.Is(err, registry.ErrNotExist) {
|
||||
mainLog.Load().Debug().Err(err).Msgf("error reading %s registry key", keyName)
|
||||
continue
|
||||
}
|
||||
if len(value) > 0 {
|
||||
mainLog.Load().Debug().Msgf("found static DNS for interface %q: %s", iface.Name, value)
|
||||
parsed := parseDNSServers(value)
|
||||
for _, pns := range parsed {
|
||||
if !slices.Contains(ns, pns) {
|
||||
ns = append(ns, pns)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
if len(ns) == 0 {
|
||||
mainLog.Load().Debug().Msgf("no static DNS values found for interface %q", iface.Name)
|
||||
}
|
||||
return ns, nil
|
||||
}
|
||||
|
||||
// parseDNSServers splits a DNS server string that may be comma- or space-separated,
|
||||
// and trims any extraneous whitespace or null characters.
|
||||
func parseDNSServers(val string) []string {
|
||||
fields := strings.FieldsFunc(val, func(r rune) bool {
|
||||
return r == ' ' || r == ','
|
||||
})
|
||||
var servers []string
|
||||
for _, f := range fields {
|
||||
trimmed := strings.TrimSpace(f)
|
||||
if len(trimmed) > 0 {
|
||||
servers = append(servers, trimmed)
|
||||
}
|
||||
}
|
||||
return servers
|
||||
}
|
||||
|
||||
// addDnsServerForwarders adds given nameservers to DNS server forwarders list,
|
||||
// and also removing old forwarders if provided.
|
||||
func addDnsServerForwarders(nameservers, old []string) error {
|
||||
@@ -216,3 +324,9 @@ func removeDnsServerForwarders(nameservers []string) error {
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// powershell runs the given powershell command.
|
||||
func powershell(cmd string) ([]byte, error) {
|
||||
out, err := exec.Command("powershell", "-Command", cmd).CombinedOutput()
|
||||
return bytes.TrimSpace(out), err
|
||||
}
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
||||
)
|
||||
|
||||
func Test_currentStaticDNS(t *testing.T) {
|
||||
iface, err := net.InterfaceByName(defaultIfaceName())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
start := time.Now()
|
||||
staticDns, err := currentStaticDNS(iface)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Logf("Using Windows API takes: %d", time.Since(start).Milliseconds())
|
||||
|
||||
start = time.Now()
|
||||
staticDnsPowershell, err := currentStaticDnsPowershell(iface)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Logf("Using Powershell takes: %d", time.Since(start).Milliseconds())
|
||||
|
||||
slices.Sort(staticDns)
|
||||
slices.Sort(staticDnsPowershell)
|
||||
if !slices.Equal(staticDns, staticDnsPowershell) {
|
||||
t.Fatalf("result mismatch, want: %v, got: %v", staticDnsPowershell, staticDns)
|
||||
}
|
||||
}
|
||||
|
||||
func currentStaticDnsPowershell(iface *net.Interface) ([]string, error) {
|
||||
luid, err := winipcfg.LUIDFromIndex(uint32(iface.Index))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
guid, err := luid.GUID()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var ns []string
|
||||
for _, path := range []string{"HKLM:\\" + v4InterfaceKeyPathFormat, "HKLM:\\" + v6InterfaceKeyPathFormat} {
|
||||
interfaceKeyPath := path + guid.String()
|
||||
found := false
|
||||
for _, key := range []string{"NameServer", "ProfileNameServer"} {
|
||||
if found {
|
||||
continue
|
||||
}
|
||||
cmd := fmt.Sprintf(`Get-ItemPropertyValue -Path "%s" -Name "%s"`, interfaceKeyPath, key)
|
||||
out, err := powershell(cmd)
|
||||
if err == nil && len(out) > 0 {
|
||||
found = true
|
||||
for _, e := range strings.Split(string(out), ",") {
|
||||
ns = append(ns, strings.TrimRight(e, "\x00"))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return ns, nil
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
+1063
-128
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,237 @@
|
||||
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
|
||||
installedNameservers []string
|
||||
installCalls int
|
||||
resetCalls int
|
||||
refusals []string
|
||||
}
|
||||
|
||||
func newInterceptFallbackHarness(t *testing.T, lc *ctrld.ListenerConfig) *interceptFallbackHarness {
|
||||
t.Helper()
|
||||
h := &interceptFallbackHarness{}
|
||||
|
||||
origStart, origInstall := startDNSInterceptFn, setDnsForRunningIfaceFn
|
||||
origReset, origFatal := resetDNSFn, refuseFallbackFatal
|
||||
origResolver := localResolverIPFn
|
||||
origCfg, origMode, origIntercept, origHard := cfg, interceptMode, dnsIntercept, hardIntercept
|
||||
t.Cleanup(func() {
|
||||
startDNSInterceptFn, setDnsForRunningIfaceFn = origStart, 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")
|
||||
}
|
||||
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()
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
// 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 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")
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -7,16 +7,17 @@ import (
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"tailscale.com/tsd"
|
||||
|
||||
"github.com/kardianos/service"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld/internal/dns"
|
||||
"github.com/Control-D-Inc/ctrld/internal/router"
|
||||
)
|
||||
|
||||
func init() {
|
||||
sys := new(tsd.System)
|
||||
if r, err := dns.NewOSConfigurator(func(format string, args ...any) {}, sys.HealthTracker(), sys.ControlKnobs(), "lo"); err == nil {
|
||||
if isAndroid() {
|
||||
return
|
||||
}
|
||||
if r, err := newLoopbackOSConfigurator(); err == nil {
|
||||
useSystemdResolved = r.Mode() == "systemd-resolved"
|
||||
}
|
||||
// Disable quic-go's ECN support by default, see https://github.com/quic-go/quic-go/issues/3911
|
||||
@@ -39,6 +40,9 @@ func setDependencies(svc *service.Config) {
|
||||
svc.Dependencies = append(svc.Dependencies, "Wants=systemd-networkd-wait-online.service")
|
||||
}
|
||||
}
|
||||
if routerDeps := router.ServiceDependencies(); len(routerDeps) > 0 {
|
||||
svc.Dependencies = append(svc.Dependencies, routerDeps...)
|
||||
}
|
||||
}
|
||||
|
||||
func setWorkingDirectory(svc *service.Config, dir string) {
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
//go:build !linux && !freebsd && !darwin
|
||||
//go:build !linux && !freebsd && !darwin && !windows
|
||||
|
||||
package cli
|
||||
|
||||
|
||||
+249
-1
@@ -1,13 +1,47 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/url"
|
||||
"runtime"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
"github.com/Masterminds/semver/v3"
|
||||
"github.com/rs/zerolog"
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
"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{}}
|
||||
|
||||
@@ -55,3 +89,217 @@ func Test_prog_dnsWatchdogInterval(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func Test_shouldUpgrade(t *testing.T) {
|
||||
// Helper function to create a version
|
||||
makeVersion := func(v string) *semver.Version {
|
||||
ver, err := semver.NewVersion(v)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create version %s: %v", v, err)
|
||||
}
|
||||
return ver
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
versionTarget string
|
||||
currentVersion *semver.Version
|
||||
shouldUpgrade bool
|
||||
description string
|
||||
}{
|
||||
{
|
||||
name: "empty version target",
|
||||
versionTarget: "",
|
||||
currentVersion: makeVersion("v1.0.0"),
|
||||
shouldUpgrade: false,
|
||||
description: "should skip upgrade when version target is empty",
|
||||
},
|
||||
{
|
||||
name: "invalid version target",
|
||||
versionTarget: "invalid-version",
|
||||
currentVersion: makeVersion("v1.0.0"),
|
||||
shouldUpgrade: false,
|
||||
description: "should skip upgrade when version target is invalid",
|
||||
},
|
||||
{
|
||||
name: "same version",
|
||||
versionTarget: "v1.0.0",
|
||||
currentVersion: makeVersion("v1.0.0"),
|
||||
shouldUpgrade: false,
|
||||
description: "should skip upgrade when target version equals current version",
|
||||
},
|
||||
{
|
||||
name: "older version",
|
||||
versionTarget: "v1.0.0",
|
||||
currentVersion: makeVersion("v1.1.0"),
|
||||
shouldUpgrade: false,
|
||||
description: "should skip upgrade when target version is older than current version",
|
||||
},
|
||||
{
|
||||
name: "patch upgrade allowed",
|
||||
versionTarget: "v1.0.1",
|
||||
currentVersion: makeVersion("v1.0.0"),
|
||||
shouldUpgrade: true,
|
||||
description: "should allow patch version upgrade within same major version",
|
||||
},
|
||||
{
|
||||
name: "minor upgrade allowed",
|
||||
versionTarget: "v1.1.0",
|
||||
currentVersion: makeVersion("v1.0.0"),
|
||||
shouldUpgrade: true,
|
||||
description: "should allow minor version upgrade within same major version",
|
||||
},
|
||||
{
|
||||
name: "major upgrade blocked",
|
||||
versionTarget: "v2.0.0",
|
||||
currentVersion: makeVersion("v1.0.0"),
|
||||
shouldUpgrade: false,
|
||||
description: "should block major version upgrade",
|
||||
},
|
||||
{
|
||||
name: "major downgrade blocked",
|
||||
versionTarget: "v1.0.0",
|
||||
currentVersion: makeVersion("v2.0.0"),
|
||||
shouldUpgrade: false,
|
||||
description: "should block major version downgrade",
|
||||
},
|
||||
{
|
||||
name: "version without v prefix",
|
||||
versionTarget: "1.0.1",
|
||||
currentVersion: makeVersion("v1.0.0"),
|
||||
shouldUpgrade: true,
|
||||
description: "should handle version target without v prefix",
|
||||
},
|
||||
{
|
||||
name: "complex version upgrade allowed",
|
||||
versionTarget: "v1.5.3",
|
||||
currentVersion: makeVersion("v1.4.2"),
|
||||
shouldUpgrade: true,
|
||||
description: "should allow complex version upgrade within same major version",
|
||||
},
|
||||
{
|
||||
name: "complex major upgrade blocked",
|
||||
versionTarget: "v3.1.0",
|
||||
currentVersion: makeVersion("v2.5.3"),
|
||||
shouldUpgrade: false,
|
||||
description: "should block complex major version upgrade",
|
||||
},
|
||||
{
|
||||
name: "pre-release version upgrade allowed",
|
||||
versionTarget: "v1.0.1-beta.1",
|
||||
currentVersion: makeVersion("v1.0.0"),
|
||||
shouldUpgrade: true,
|
||||
description: "should allow pre-release version upgrade within same major version",
|
||||
},
|
||||
{
|
||||
name: "pre-release major upgrade blocked",
|
||||
versionTarget: "v2.0.0-alpha.1",
|
||||
currentVersion: makeVersion("v1.0.0"),
|
||||
shouldUpgrade: false,
|
||||
description: "should block pre-release major version upgrade",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
// Create test logger
|
||||
testLogger := zerolog.New(zerolog.NewTestWriter(t)).With().Logger()
|
||||
|
||||
// Call the function and capture the result
|
||||
result := shouldUpgrade(tc.versionTarget, tc.currentVersion, &testLogger)
|
||||
|
||||
// Assert the expected result
|
||||
assert.Equal(t, tc.shouldUpgrade, result, tc.description)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func Test_selfUpgradeCheck(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("skipped due to Windows file locking issue on Github Action runners")
|
||||
}
|
||||
|
||||
// Helper function to create a version
|
||||
makeVersion := func(v string) *semver.Version {
|
||||
ver, err := semver.NewVersion(v)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create version %s: %v", v, err)
|
||||
}
|
||||
return ver
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
versionTarget string
|
||||
currentVersion *semver.Version
|
||||
shouldUpgrade bool
|
||||
description string
|
||||
}{
|
||||
{
|
||||
name: "upgrade allowed",
|
||||
versionTarget: "v1.0.1",
|
||||
currentVersion: makeVersion("v1.0.0"),
|
||||
shouldUpgrade: true,
|
||||
description: "should allow upgrade and attempt to perform it",
|
||||
},
|
||||
{
|
||||
name: "upgrade blocked",
|
||||
versionTarget: "v2.0.0",
|
||||
currentVersion: makeVersion("v1.0.0"),
|
||||
shouldUpgrade: false,
|
||||
description: "should block upgrade and not attempt to perform it",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
// Create test logger
|
||||
testLogger := zerolog.New(zerolog.NewTestWriter(t)).With().Logger()
|
||||
|
||||
// Call the function and capture the result
|
||||
result := selfUpgradeCheck(tc.versionTarget, tc.currentVersion, &testLogger)
|
||||
|
||||
// Assert the expected result
|
||||
assert.Equal(t, tc.shouldUpgrade, result, tc.description)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func Test_performUpgrade(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("skipped due to Windows file locking issue on Github Action runners")
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
versionTarget string
|
||||
expectedResult bool
|
||||
description string
|
||||
}{
|
||||
{
|
||||
name: "valid version target",
|
||||
versionTarget: "v1.0.1",
|
||||
expectedResult: true,
|
||||
description: "should attempt to perform upgrade with valid version target",
|
||||
},
|
||||
{
|
||||
name: "empty version target",
|
||||
versionTarget: "",
|
||||
expectedResult: true,
|
||||
description: "should attempt to perform upgrade even with empty version target",
|
||||
},
|
||||
}
|
||||
|
||||
// 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) {
|
||||
// Call the function and capture the result
|
||||
result := performUpgrade(tc.versionTarget)
|
||||
assert.Equal(t, tc.expectedResult, result, tc.description)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
package cli
|
||||
|
||||
import "github.com/kardianos/service"
|
||||
|
||||
func setDependencies(svc *service.Config) {
|
||||
if hasLocalDnsServerRunning() {
|
||||
svc.Dependencies = []string{"DNS"}
|
||||
}
|
||||
}
|
||||
|
||||
func setWorkingDirectory(svc *service.Config, dir string) {
|
||||
// WorkingDirectory is not supported on Windows.
|
||||
svc.WorkingDirectory = dir
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+101
-10
@@ -3,11 +3,38 @@ package cli
|
||||
import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/fsnotify/fsnotify"
|
||||
)
|
||||
|
||||
// parseResolvConfNameservers reads the resolv.conf file and returns the nameservers found.
|
||||
// Returns nil if no nameservers are found.
|
||||
func (p *prog) parseResolvConfNameservers(path string) ([]string, error) {
|
||||
content, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Parse the file for "nameserver" lines
|
||||
var currentNS []string
|
||||
lines := strings.Split(string(content), "\n")
|
||||
for _, line := range lines {
|
||||
trimmed := strings.TrimSpace(line)
|
||||
if strings.HasPrefix(trimmed, "nameserver") {
|
||||
parts := strings.Fields(trimmed)
|
||||
if len(parts) >= 2 {
|
||||
currentNS = append(currentNS, parts[1])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return currentNS, nil
|
||||
}
|
||||
|
||||
// watchResolvConf watches any changes to /etc/resolv.conf file,
|
||||
// and reverting to the original config set by ctrld.
|
||||
func (p *prog) watchResolvConf(iface *net.Interface, ns []netip.Addr, setDnsFn func(iface *net.Interface, ns []netip.Addr) error) {
|
||||
@@ -40,7 +67,7 @@ func (p *prog) watchResolvConf(iface *net.Interface, ns []netip.Addr, setDnsFn f
|
||||
mainLog.Load().Debug().Msgf("stopping watcher for %s", resolvConfPath)
|
||||
return
|
||||
case event, ok := <-watcher.Events:
|
||||
if p.leakingQuery.Load() {
|
||||
if p.recoveryRunning.Load() {
|
||||
return
|
||||
}
|
||||
if !ok {
|
||||
@@ -50,17 +77,81 @@ func (p *prog) watchResolvConf(iface *net.Interface, ns []netip.Addr, setDnsFn f
|
||||
continue
|
||||
}
|
||||
if event.Has(fsnotify.Write) || event.Has(fsnotify.Create) {
|
||||
mainLog.Load().Debug().Msg("/etc/resolv.conf changes detected, reverting to ctrld setting")
|
||||
if err := watcher.Remove(watchDir); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to pause watcher")
|
||||
continue
|
||||
mainLog.Load().Debug().Msgf("/etc/resolv.conf changes detected, reading changes...")
|
||||
|
||||
// Convert expected nameservers to strings for comparison
|
||||
expectedNS := make([]string, len(ns))
|
||||
for i, addr := range ns {
|
||||
expectedNS[i] = addr.String()
|
||||
}
|
||||
if err := setDnsFn(iface, ns); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to revert /etc/resolv.conf changes")
|
||||
|
||||
var foundNS []string
|
||||
var err error
|
||||
|
||||
maxRetries := 1
|
||||
for retry := 0; retry < maxRetries; retry++ {
|
||||
foundNS, err = p.parseResolvConfNameservers(resolvConfPath)
|
||||
if err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to read resolv.conf content")
|
||||
break
|
||||
}
|
||||
|
||||
// If we found nameservers, break out of retry loop
|
||||
if len(foundNS) > 0 {
|
||||
break
|
||||
}
|
||||
|
||||
// Only retry if we found no nameservers
|
||||
if retry < maxRetries-1 {
|
||||
mainLog.Load().Debug().Msgf("resolv.conf has no nameserver entries, retry %d/%d in 2 seconds", retry+1, maxRetries)
|
||||
select {
|
||||
case <-p.stopCh:
|
||||
return
|
||||
case <-p.dnsWatcherStopCh:
|
||||
return
|
||||
case <-time.After(2 * time.Second):
|
||||
continue
|
||||
}
|
||||
} else {
|
||||
mainLog.Load().Debug().Msg("resolv.conf remained empty after all retries")
|
||||
}
|
||||
}
|
||||
if err := watcher.Add(watchDir); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to continue running watcher")
|
||||
return
|
||||
|
||||
// If we found nameservers, check if they match what we expect
|
||||
if len(foundNS) > 0 {
|
||||
// Check if the nameservers match exactly what we expect
|
||||
matches := len(foundNS) == len(expectedNS)
|
||||
if matches {
|
||||
for i := range foundNS {
|
||||
if foundNS[i] != expectedNS[i] {
|
||||
matches = false
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
mainLog.Load().Debug().
|
||||
Strs("found", foundNS).
|
||||
Strs("expected", expectedNS).
|
||||
Bool("matches", matches).
|
||||
Msg("checking nameservers")
|
||||
|
||||
// Only revert if the nameservers don't match
|
||||
if !matches {
|
||||
if err := watcher.Remove(watchDir); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to pause watcher")
|
||||
continue
|
||||
}
|
||||
|
||||
if err := setDnsFn(iface, ns); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to revert /etc/resolv.conf changes")
|
||||
}
|
||||
|
||||
if err := watcher.Add(watchDir); err != nil {
|
||||
mainLog.Load().Error().Err(err).Msg("failed to continue running watcher")
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
case err, ok := <-watcher.Errors:
|
||||
|
||||
@@ -6,16 +6,16 @@ import (
|
||||
"net"
|
||||
"net/netip"
|
||||
|
||||
"tailscale.com/tsd"
|
||||
"tailscale.com/control/controlknobs"
|
||||
"tailscale.com/health"
|
||||
"tailscale.com/util/dnsname"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld/internal/dns"
|
||||
)
|
||||
|
||||
// setResolvConf sets the content of resolv.conf file using the given nameservers list.
|
||||
// setResolvConf sets the content of the resolv.conf file using the given nameservers list.
|
||||
func setResolvConf(iface *net.Interface, ns []netip.Addr) error {
|
||||
sys := new(tsd.System)
|
||||
r, err := dns.NewOSConfigurator(func(format string, args ...any) {}, sys.HealthTracker(), sys.ControlKnobs(), "lo") // interface name does not matter.
|
||||
r, err := newLoopbackOSConfigurator()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -24,13 +24,17 @@ func setResolvConf(iface *net.Interface, ns []netip.Addr) error {
|
||||
Nameservers: ns,
|
||||
SearchDomains: []dnsname.FQDN{},
|
||||
}
|
||||
if sds, err := searchDomains(); err == nil {
|
||||
oc.SearchDomains = sds
|
||||
} else {
|
||||
mainLog.Load().Debug().Err(err).Msg("failed to get search domains list when reverting resolv.conf file")
|
||||
}
|
||||
return r.SetDNS(oc)
|
||||
}
|
||||
|
||||
// shouldWatchResolvconf reports whether ctrld should watch changes to resolv.conf file with given OS configurator.
|
||||
func shouldWatchResolvconf() bool {
|
||||
sys := new(tsd.System)
|
||||
r, err := dns.NewOSConfigurator(func(format string, args ...any) {}, sys.HealthTracker(), sys.ControlKnobs(), "lo") // interface name does not matter.
|
||||
r, err := newLoopbackOSConfigurator()
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
@@ -41,3 +45,8 @@ func shouldWatchResolvconf() bool {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// newLoopbackOSConfigurator creates an OSConfigurator for DNS management using the "lo" interface.
|
||||
func newLoopbackOSConfigurator() (dns.OSConfigurator, error) {
|
||||
return dns.NewOSConfigurator(noopLogf, &health.Tracker{}, &controlknobs.Knobs{}, "lo")
|
||||
}
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
//go:build unix
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"tailscale.com/util/dnsname"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld/internal/resolvconffile"
|
||||
)
|
||||
|
||||
// searchDomains returns the current search domains config.
|
||||
func searchDomains() ([]dnsname.FQDN, error) {
|
||||
return resolvconffile.SearchDomains()
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"syscall"
|
||||
|
||||
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
||||
"tailscale.com/util/dnsname"
|
||||
)
|
||||
|
||||
// searchDomains returns the current search domains config.
|
||||
func searchDomains() ([]dnsname.FQDN, error) {
|
||||
flags := winipcfg.GAAFlagIncludeGateways |
|
||||
winipcfg.GAAFlagIncludePrefix
|
||||
|
||||
aas, err := winipcfg.GetAdaptersAddresses(syscall.AF_UNSPEC, flags)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("winipcfg.GetAdaptersAddresses: %w", err)
|
||||
}
|
||||
|
||||
var sds []dnsname.FQDN
|
||||
for _, aa := range aas {
|
||||
if aa.OperStatus != winipcfg.IfOperStatusUp {
|
||||
continue
|
||||
}
|
||||
|
||||
// Skip if software loopback or other non-physical types
|
||||
// This is to avoid the "Loopback Pseudo-Interface 1" issue we see on windows
|
||||
if aa.IfType == winipcfg.IfTypeSoftwareLoopback {
|
||||
continue
|
||||
}
|
||||
|
||||
for a := aa.FirstDNSSuffix; a != nil; a = a.Next {
|
||||
d, err := dnsname.ToFQDN(a.String())
|
||||
if err != nil {
|
||||
mainLog.Load().Debug().Err(err).Msgf("failed to parse domain: %s", a.String())
|
||||
continue
|
||||
}
|
||||
sds = append(sds, d)
|
||||
}
|
||||
}
|
||||
return sds, nil
|
||||
}
|
||||
@@ -22,8 +22,8 @@ func selfUninstall(p *prog, logger zerolog.Logger) {
|
||||
logger.Fatal().Err(err).Msg("could not determine executable")
|
||||
}
|
||||
args := []string{"uninstall"}
|
||||
if !deactivationPinNotSet() {
|
||||
args = append(args, fmt.Sprintf("--pin=%d", cdDeactivationPin))
|
||||
if deactivationPinSet() {
|
||||
args = append(args, fmt.Sprintf("--pin=%d", cdDeactivationPin.Load()))
|
||||
}
|
||||
cmd := exec.Command(bin, args...)
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
|
||||
|
||||
@@ -0,0 +1,12 @@
|
||||
//go:build !windows
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"syscall"
|
||||
)
|
||||
|
||||
// sysProcAttrForDetachedChildProcess returns *syscall.SysProcAttr instance for running a detached child command.
|
||||
func sysProcAttrForDetachedChildProcess() *syscall.SysProcAttr {
|
||||
return &syscall.SysProcAttr{Setsid: true}
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"syscall"
|
||||
)
|
||||
|
||||
// From: https://learn.microsoft.com/en-us/windows/win32/procthread/process-creation-flags?redirectedfrom=MSDN
|
||||
|
||||
// SYSCALL_CREATE_NO_WINDOW set flag to run process without a console window.
|
||||
const SYSCALL_CREATE_NO_WINDOW = 0x08000000
|
||||
|
||||
// sysProcAttrForDetachedChildProcess returns *syscall.SysProcAttr instance for running self-upgrade command.
|
||||
func sysProcAttrForDetachedChildProcess() *syscall.SysProcAttr {
|
||||
return &syscall.SysProcAttr{
|
||||
CreationFlags: syscall.CREATE_NEW_PROCESS_GROUP | SYSCALL_CREATE_NO_WINDOW,
|
||||
HideWindow: true,
|
||||
}
|
||||
}
|
||||
+92
-11
@@ -4,12 +4,16 @@ import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"os/exec"
|
||||
"runtime"
|
||||
|
||||
"github.com/coreos/go-systemd/v22/unit"
|
||||
"github.com/kardianos/service"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld/internal/router"
|
||||
"github.com/Control-D-Inc/ctrld/internal/router/openwrt"
|
||||
)
|
||||
|
||||
// newService wraps service.New call to return service.Service
|
||||
@@ -130,6 +134,63 @@ func (s *systemd) Status() (service.Status, error) {
|
||||
return s.Service.Status()
|
||||
}
|
||||
|
||||
func (s *systemd) Start() error {
|
||||
const systemdUnitFile = "/etc/systemd/system/ctrld.service"
|
||||
f, err := os.Open(systemdUnitFile)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer f.Close()
|
||||
if opts, change := ensureSystemdKillMode(f); change {
|
||||
mode := os.FileMode(0644)
|
||||
buf, err := io.ReadAll(unit.Serialize(opts))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.WriteFile(systemdUnitFile, buf, mode); err != nil {
|
||||
return err
|
||||
}
|
||||
if out, err := exec.Command("systemctl", "daemon-reload").CombinedOutput(); err != nil {
|
||||
return fmt.Errorf("systemctl daemon-reload failed: %w\n%s", err, string(out))
|
||||
}
|
||||
mainLog.Load().Debug().Msg("set KillMode=process successfully")
|
||||
}
|
||||
return s.Service.Start()
|
||||
}
|
||||
|
||||
// ensureSystemdKillMode ensure systemd unit file is configured with KillMode=process.
|
||||
// This is necessary for running self-upgrade flow.
|
||||
func ensureSystemdKillMode(r io.Reader) (opts []*unit.UnitOption, change bool) {
|
||||
opts, err := unit.DeserializeOptions(r)
|
||||
// staticcheck sees only the explicit non-nil sends on the lexer's error
|
||||
// channel, so it reports this comparison as always true. 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
|
||||
}
|
||||
change = true
|
||||
needKillModeOpt := true
|
||||
killModeOpt := unit.NewUnitOption("Service", "KillMode", "process")
|
||||
for _, opt := range opts {
|
||||
if opt.Match(killModeOpt) {
|
||||
needKillModeOpt = false
|
||||
change = false
|
||||
break
|
||||
}
|
||||
if opt.Section == killModeOpt.Section && opt.Name == killModeOpt.Name {
|
||||
opt.Value = killModeOpt.Value
|
||||
needKillModeOpt = false
|
||||
break
|
||||
}
|
||||
}
|
||||
if needKillModeOpt {
|
||||
opts = append(opts, killModeOpt)
|
||||
}
|
||||
return opts, change
|
||||
}
|
||||
|
||||
func newLaunchd(s service.Service) *launchd {
|
||||
return &launchd{
|
||||
Service: s,
|
||||
@@ -156,20 +217,33 @@ func (l *launchd) Status() (service.Status, error) {
|
||||
type task struct {
|
||||
f func() error
|
||||
abortOnError bool
|
||||
Name string
|
||||
}
|
||||
|
||||
// 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" || t.Name != "Stop" {
|
||||
mainLog.Load().Debug().Msgf("error running task %s: %v", t.Name, taskErr)
|
||||
}
|
||||
}
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
func doTasks(tasks []task) bool {
|
||||
var prevErr error
|
||||
for _, task := range tasks {
|
||||
if err := task.f(); err != nil {
|
||||
if task.abortOnError {
|
||||
mainLog.Load().Error().Msg(errors.Join(prevErr, err).Error())
|
||||
return false
|
||||
}
|
||||
prevErr = err
|
||||
}
|
||||
}
|
||||
return true
|
||||
_, err := doTasksE(tasks)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
func checkHasElevatedPrivilege() {
|
||||
@@ -187,6 +261,13 @@ func checkHasElevatedPrivilege() {
|
||||
func unixSystemVServiceStatus() (service.Status, error) {
|
||||
out, err := exec.Command("/etc/init.d/ctrld", "status").CombinedOutput()
|
||||
if err != nil {
|
||||
// Specific case for openwrt >= 24.10, it returns non-success code
|
||||
// for above status command, which may not right.
|
||||
if router.Name() == openwrt.Name {
|
||||
if string(bytes.ToLower(bytes.TrimSpace(out))) == "inactive" {
|
||||
return service.StatusStopped, nil
|
||||
}
|
||||
}
|
||||
return service.StatusUnknown, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,146 @@
|
||||
//go:build darwin
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const launchdPlistPath = "/Library/LaunchDaemons/ctrld.plist"
|
||||
|
||||
// serviceConfigFileExists returns true if the launchd plist for ctrld exists on disk.
|
||||
// This is more reliable than checking launchctl status, which may report "not found"
|
||||
// if the service was unloaded but the plist file still exists.
|
||||
func serviceConfigFileExists() bool {
|
||||
_, err := os.Stat(launchdPlistPath)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// appendServiceFlag appends a CLI flag (e.g., "--intercept-mode") to the installed
|
||||
// service's launch arguments. This is used when upgrading an existing installation
|
||||
// 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 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("/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 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
|
||||
}
|
||||
|
||||
// Use PlistBuddy to append the flag to ProgramArguments array.
|
||||
// PlistBuddy is more reliable than "defaults" for array manipulation.
|
||||
addCmd := exec.Command(
|
||||
"/usr/libexec/PlistBuddy",
|
||||
"-c", fmt.Sprintf("Add :ProgramArguments: string %s", flag),
|
||||
launchdPlistPath,
|
||||
)
|
||||
if out, err := addCmd.CombinedOutput(); err != nil {
|
||||
return fmt.Errorf("failed to append %q to plist ProgramArguments: %w (output: %s)", flag, err, strings.TrimSpace(string(out)))
|
||||
}
|
||||
|
||||
mainLog.Load().Info().Msgf("Appended %q to service launch arguments", flag)
|
||||
return nil
|
||||
}
|
||||
|
||||
// verifyServiceRegistration is a no-op on macOS (launchd plist verification not needed).
|
||||
func verifyServiceRegistration() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 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 {
|
||||
// Read current ProgramArguments to find the index.
|
||||
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)))
|
||||
}
|
||||
|
||||
// Parse the PlistBuddy output to find the flag's index.
|
||||
// PlistBuddy prints arrays as:
|
||||
// Array {
|
||||
// /path/to/ctrld
|
||||
// run
|
||||
// --cd=xxx
|
||||
// --intercept-mode
|
||||
// dns
|
||||
// }
|
||||
lines := strings.Split(string(out), "\n")
|
||||
var entries []string
|
||||
for _, line := range lines {
|
||||
trimmed := strings.TrimSpace(line)
|
||||
if trimmed == "Array {" || trimmed == "}" || trimmed == "" {
|
||||
continue
|
||||
}
|
||||
entries = append(entries, trimmed)
|
||||
}
|
||||
|
||||
index, hasValue := serviceFlagPosition(entries, flag)
|
||||
|
||||
if index < 0 {
|
||||
mainLog.Load().Debug().Msgf("Service flag %q not present in plist, skipping removal", flag)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Delete a separate value first. An inline --flag=value entry is one array item.
|
||||
if hasValue {
|
||||
delVal := exec.Command(
|
||||
"/usr/libexec/PlistBuddy",
|
||||
"-c", fmt.Sprintf("Delete :ProgramArguments:%d", index+1),
|
||||
launchdPlistPath,
|
||||
)
|
||||
if out, err := delVal.CombinedOutput(); err != nil {
|
||||
return fmt.Errorf("failed to remove value for %q from plist: %w (output: %s)", flag, err, strings.TrimSpace(string(out)))
|
||||
}
|
||||
}
|
||||
|
||||
// Delete the flag itself.
|
||||
delCmd := exec.Command(
|
||||
"/usr/libexec/PlistBuddy",
|
||||
"-c", fmt.Sprintf("Delete :ProgramArguments:%d", index),
|
||||
launchdPlistPath,
|
||||
)
|
||||
if out, err := delCmd.CombinedOutput(); err != nil {
|
||||
return fmt.Errorf("failed to remove %q from plist ProgramArguments: %w (output: %s)", flag, err, strings.TrimSpace(string(out)))
|
||||
}
|
||||
|
||||
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,58 @@
|
||||
//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")
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
//go:build !darwin && !windows
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"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
|
||||
if _, err := os.Stat("/etc/systemd/system/ctrld.service"); err == nil {
|
||||
return true
|
||||
}
|
||||
// SysV init script
|
||||
if _, err := os.Stat("/etc/init.d/ctrld"); err == nil {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// appendServiceFlag is not yet implemented on this platform.
|
||||
// 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 errServiceFlagsUnsupported
|
||||
}
|
||||
|
||||
// verifyServiceRegistration is a no-op on this platform.
|
||||
func verifyServiceRegistration() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// removeServiceFlag is not yet implemented on this platform.
|
||||
func removeServiceFlag(flag string) error {
|
||||
return errServiceFlagsUnsupported
|
||||
}
|
||||
@@ -0,0 +1,169 @@
|
||||
//go:build windows
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/sys/windows/svc/mgr"
|
||||
)
|
||||
|
||||
// serviceConfigFileExists returns true if the ctrld Windows service is registered.
|
||||
func serviceConfigFileExists() bool {
|
||||
m, err := mgr.Connect()
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
defer m.Disconnect()
|
||||
s, err := m.OpenService(ctrldServiceName)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
s.Close()
|
||||
return true
|
||||
}
|
||||
|
||||
// appendServiceFlag appends a CLI flag (e.g., "--intercept-mode") to the installed
|
||||
// Windows service's BinPath arguments. This is used when upgrading an existing
|
||||
// installation to intercept mode without losing the existing --cd flag.
|
||||
//
|
||||
// The function is idempotent: if the flag already exists, it's a no-op.
|
||||
func appendServiceFlag(flag string) error {
|
||||
m, err := mgr.Connect()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to connect to Windows SCM: %w", err)
|
||||
}
|
||||
defer m.Disconnect()
|
||||
|
||||
s, err := m.OpenService(ctrldServiceName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to open service %q: %w", ctrldServiceName, err)
|
||||
}
|
||||
defer s.Close()
|
||||
|
||||
config, err := s.Config()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read service config: %w", err)
|
||||
}
|
||||
|
||||
// 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
|
||||
}
|
||||
|
||||
// Append the flag to BinPath.
|
||||
config.BinaryPathName = strings.TrimSpace(config.BinaryPathName) + " " + flag
|
||||
|
||||
if err := s.UpdateConfig(config); err != nil {
|
||||
return fmt.Errorf("failed to update service config with %q: %w", flag, err)
|
||||
}
|
||||
|
||||
mainLog.Load().Info().Msgf("Appended %q to service BinPath", flag)
|
||||
return nil
|
||||
}
|
||||
|
||||
// verifyServiceRegistration opens the Windows Service Control Manager and verifies
|
||||
// that the ctrld service is correctly registered: logs the BinaryPathName, checks
|
||||
// that --intercept-mode is present if expected, and verifies SERVICE_AUTO_START.
|
||||
func verifyServiceRegistration() error {
|
||||
m, err := mgr.Connect()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to connect to Windows SCM: %w", err)
|
||||
}
|
||||
defer m.Disconnect()
|
||||
|
||||
s, err := m.OpenService(ctrldServiceName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to open service %q: %w", ctrldServiceName, err)
|
||||
}
|
||||
defer s.Close()
|
||||
|
||||
config, err := s.Config()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read service config: %w", err)
|
||||
}
|
||||
|
||||
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 !strings.Contains(config.BinaryPathName, "--intercept-mode") {
|
||||
return fmt.Errorf("service registry: --intercept-mode flag missing from BinaryPathName (expected mode %q)", interceptMode)
|
||||
}
|
||||
mainLog.Load().Debug().Msgf("Service registry: --intercept-mode flag present in BinaryPathName")
|
||||
}
|
||||
|
||||
// Verify auto-start. mgr.StartAutomatic == 2 == SERVICE_AUTO_START.
|
||||
if config.StartType != mgr.StartAutomatic {
|
||||
return fmt.Errorf("service registry: StartType is %d, expected SERVICE_AUTO_START (%d)", config.StartType, mgr.StartAutomatic)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// 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 {
|
||||
return fmt.Errorf("failed to connect to Windows SCM: %w", err)
|
||||
}
|
||||
defer m.Disconnect()
|
||||
|
||||
s, err := m.OpenService(ctrldServiceName)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to open service %q: %w", ctrldServiceName, err)
|
||||
}
|
||||
defer s.Close()
|
||||
|
||||
config, err := s.Config()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read service config: %w", err)
|
||||
}
|
||||
|
||||
updatedPath, removed := removeBinaryPathFlag(config.BinaryPathName, flag)
|
||||
if !removed {
|
||||
mainLog.Load().Debug().Msgf("Service flag %q not present in BinPath, skipping removal", flag)
|
||||
return nil
|
||||
}
|
||||
config.BinaryPathName = updatedPath
|
||||
|
||||
if err := s.UpdateConfig(config); err != nil {
|
||||
return fmt.Errorf("failed to update service config: %w", err)
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
@@ -13,3 +13,10 @@ func hasElevatedPrivilege() (bool, error) {
|
||||
func openLogFile(path string, flags int) (*os.File, error) {
|
||||
return os.OpenFile(path, flags, os.FileMode(0o600))
|
||||
}
|
||||
|
||||
// hasLocalDnsServerRunning reports whether we are on Windows and having Dns server running.
|
||||
func hasLocalDnsServerRunning() bool { return false }
|
||||
|
||||
func ConfigureWindowsServiceFailureActions(serviceName string) error { return nil }
|
||||
|
||||
func isRunningOnDomainControllerWindows() (bool, int) { return false, 0 }
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func Test_ensureSystemdKillMode(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
unitFile string
|
||||
wantChange bool
|
||||
}{
|
||||
{"no KillMode", "[Service]\nExecStart=/bin/sleep 1", true},
|
||||
{"not KillMode=process", "[Service]\nExecStart=/bin/sleep 1\nKillMode=mixed", true},
|
||||
{"KillMode=process", "[Service]\nExecStart=/bin/sleep 1\nKillMode=process", false},
|
||||
{"invalid unit file", "[Service\nExecStart=/bin/sleep 1\nKillMode=process", false},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
if _, change := ensureSystemdKillMode(strings.NewReader(tc.unitFile)); tc.wantChange != change {
|
||||
t.Errorf("ensureSystemdKillMode(%q) = %v, want %v", tc.unitFile, change, tc.wantChange)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
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")
|
||||
}
|
||||
}
|
||||
@@ -2,9 +2,20 @@ package cli
|
||||
|
||||
import (
|
||||
"os"
|
||||
"reflect"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
"unsafe"
|
||||
|
||||
"github.com/microsoft/wmi/pkg/base/host"
|
||||
"github.com/microsoft/wmi/pkg/base/instance"
|
||||
"github.com/microsoft/wmi/pkg/base/query"
|
||||
"github.com/microsoft/wmi/pkg/constant"
|
||||
"golang.org/x/sys/windows"
|
||||
"golang.org/x/sys/windows/svc/mgr"
|
||||
)
|
||||
|
||||
func hasElevatedPrivilege() (bool, error) {
|
||||
@@ -28,6 +39,67 @@ func hasElevatedPrivilege() (bool, error) {
|
||||
return token.IsMember(sid)
|
||||
}
|
||||
|
||||
// ConfigureWindowsServiceFailureActions checks if the given service
|
||||
// has the correct failure actions configured, and updates them if not.
|
||||
func ConfigureWindowsServiceFailureActions(serviceName string) error {
|
||||
if runtime.GOOS != "windows" {
|
||||
return nil // no-op on non-Windows
|
||||
}
|
||||
|
||||
m, err := mgr.Connect()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer m.Disconnect()
|
||||
|
||||
s, err := m.OpenService(serviceName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer s.Close()
|
||||
|
||||
// 1. Retrieve the current config
|
||||
cfg, err := s.Config()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 2. Update the Description
|
||||
cfg.Description = "A highly configurable, multi-protocol DNS forwarding proxy"
|
||||
|
||||
// 3. Apply the updated config
|
||||
if err := s.UpdateConfig(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Then proceed with existing actions, e.g. setting failure actions
|
||||
actions := []mgr.RecoveryAction{
|
||||
{Type: mgr.ServiceRestart, Delay: time.Second * 5}, // 5 seconds
|
||||
{Type: mgr.ServiceRestart, Delay: time.Second * 5}, // 5 seconds
|
||||
{Type: mgr.ServiceRestart, Delay: time.Second * 5}, // 5 seconds
|
||||
}
|
||||
|
||||
// Set the recovery actions (3 restarts, reset period = 120).
|
||||
err = s.SetRecoveryActions(actions, 120)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Ensure that failure actions are NOT triggered on user-initiated stops.
|
||||
var failureActionsFlag windows.SERVICE_FAILURE_ACTIONS_FLAG
|
||||
failureActionsFlag.FailureActionsOnNonCrashFailures = 0
|
||||
|
||||
if err := windows.ChangeServiceConfig2(
|
||||
s.Handle,
|
||||
windows.SERVICE_CONFIG_FAILURE_ACTIONS_FLAG,
|
||||
(*byte)(unsafe.Pointer(&failureActionsFlag)),
|
||||
); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func openLogFile(path string, mode int) (*os.File, error) {
|
||||
if len(path) == 0 {
|
||||
return nil, &os.PathError{Path: path, Op: "open", Err: syscall.ERROR_FILE_NOT_FOUND}
|
||||
@@ -79,3 +151,78 @@ func openLogFile(path string, mode int) (*os.File, error) {
|
||||
|
||||
return os.NewFile(uintptr(handle), path), nil
|
||||
}
|
||||
|
||||
const processEntrySize = uint32(unsafe.Sizeof(windows.ProcessEntry32{}))
|
||||
|
||||
// hasLocalDnsServerRunning reports whether we are on Windows and having Dns server running.
|
||||
func hasLocalDnsServerRunning() bool {
|
||||
h, e := windows.CreateToolhelp32Snapshot(windows.TH32CS_SNAPPROCESS, 0)
|
||||
if e != nil {
|
||||
return false
|
||||
}
|
||||
defer windows.CloseHandle(h)
|
||||
p := windows.ProcessEntry32{Size: processEntrySize}
|
||||
for {
|
||||
e := windows.Process32Next(h, &p)
|
||||
if e != nil {
|
||||
return false
|
||||
}
|
||||
if strings.ToLower(windows.UTF16ToString(p.ExeFile[:])) == "dns.exe" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func isRunningOnDomainControllerWindows() (bool, int) {
|
||||
whost := host.NewWmiLocalHost()
|
||||
q := query.NewWmiQuery("Win32_ComputerSystem")
|
||||
instances, err := instance.GetWmiInstancesFromHost(whost, string(constant.CimV2), q)
|
||||
if err != nil {
|
||||
mainLog.Load().Debug().Err(err).Msg("WMI query failed")
|
||||
return false, 0
|
||||
}
|
||||
if instances == nil {
|
||||
mainLog.Load().Debug().Msg("WMI query returned nil instances")
|
||||
return false, 0
|
||||
}
|
||||
defer instances.Close()
|
||||
|
||||
if len(instances) == 0 {
|
||||
mainLog.Load().Debug().Msg("no rows returned from Win32_ComputerSystem")
|
||||
return false, 0
|
||||
}
|
||||
|
||||
val, err := instances[0].GetProperty("DomainRole")
|
||||
if err != nil {
|
||||
mainLog.Load().Debug().Err(err).Msg("failed to get DomainRole property")
|
||||
return false, 0
|
||||
}
|
||||
if val == nil {
|
||||
mainLog.Load().Debug().Msg("DomainRole property is nil")
|
||||
return false, 0
|
||||
}
|
||||
|
||||
// Safely handle varied types: string or integer
|
||||
var roleInt int
|
||||
switch v := val.(type) {
|
||||
case string:
|
||||
// "4", "5", etc.
|
||||
parsed, parseErr := strconv.Atoi(v)
|
||||
if parseErr != nil {
|
||||
mainLog.Load().Debug().Err(parseErr).Msgf("failed to parse DomainRole value %q", v)
|
||||
return false, 0
|
||||
}
|
||||
roleInt = parsed
|
||||
case int8, int16, int32, int64:
|
||||
roleInt = int(reflect.ValueOf(v).Int())
|
||||
case uint8, uint16, uint32, uint64:
|
||||
roleInt = int(reflect.ValueOf(v).Uint())
|
||||
default:
|
||||
mainLog.Load().Debug().Msgf("unexpected DomainRole type: %T value=%v", v, v)
|
||||
return false, 0
|
||||
}
|
||||
|
||||
// Check if role indicates a domain controller
|
||||
isDC := roleInt == BackupDomainController || roleInt == PrimaryDomainController
|
||||
return isDC, roleInt
|
||||
}
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func Test_hasLocalDnsServerRunning(t *testing.T) {
|
||||
start := time.Now()
|
||||
hasDns := hasLocalDnsServerRunning()
|
||||
t.Logf("Using Windows API takes: %d", time.Since(start).Milliseconds())
|
||||
|
||||
start = time.Now()
|
||||
hasDnsPowershell := hasLocalDnsServerRunningPowershell()
|
||||
t.Logf("Using Powershell takes: %d", time.Since(start).Milliseconds())
|
||||
|
||||
if hasDns != hasDnsPowershell {
|
||||
t.Fatalf("result mismatch, want: %v, got: %v", hasDnsPowershell, hasDns)
|
||||
}
|
||||
}
|
||||
|
||||
func hasLocalDnsServerRunningPowershell() bool {
|
||||
_, err := powershell("Get-Process -Name DNS")
|
||||
return err == nil
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+84
-57
@@ -1,38 +1,60 @@
|
||||
package cli
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld"
|
||||
)
|
||||
|
||||
const (
|
||||
// maxFailureRequest is the maximum failed queries allowed before an upstream is marked as down.
|
||||
maxFailureRequest = 100
|
||||
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
|
||||
|
||||
mu sync.Mutex
|
||||
mu sync.RWMutex
|
||||
checking map[string]bool
|
||||
down map[string]bool
|
||||
failureReq map[string]uint64
|
||||
recovered map[string]bool
|
||||
|
||||
// failureTimerActive tracks if a timer is already running for a given upstream.
|
||||
failureTimerActive map[string]bool
|
||||
}
|
||||
|
||||
func newUpstreamMonitor(cfg *ctrld.Config) *upstreamMonitor {
|
||||
um := &upstreamMonitor{
|
||||
cfg: cfg,
|
||||
checking: make(map[string]bool),
|
||||
down: make(map[string]bool),
|
||||
failureReq: make(map[string]uint64),
|
||||
cfg: cfg,
|
||||
checking: make(map[string]bool),
|
||||
down: make(map[string]bool),
|
||||
failureReq: make(map[string]uint64),
|
||||
recovered: make(map[string]bool),
|
||||
failureTimerActive: make(map[string]bool),
|
||||
}
|
||||
for n := range cfg.Upstream {
|
||||
upstream := upstreamPrefix + n
|
||||
@@ -42,14 +64,47 @@ func newUpstreamMonitor(cfg *ctrld.Config) *upstreamMonitor {
|
||||
return um
|
||||
}
|
||||
|
||||
// increaseFailureCount increase failed queries count for an upstream by 1.
|
||||
// increaseFailureCount increases failed queries count for an upstream by 1 and logs debug information.
|
||||
// It uses a timer to debounce failure detection, ensuring that an upstream is marked as down
|
||||
// within 10 seconds if failures persist, without spawning duplicate goroutines.
|
||||
func (um *upstreamMonitor) increaseFailureCount(upstream string) {
|
||||
um.mu.Lock()
|
||||
defer um.mu.Unlock()
|
||||
|
||||
if um.recovered[upstream] {
|
||||
mainLog.Load().Debug().Msgf("upstream %q is recovered, skipping failure count increase", upstream)
|
||||
return
|
||||
}
|
||||
|
||||
um.failureReq[upstream] += 1
|
||||
failedCount := um.failureReq[upstream]
|
||||
um.down[upstream] = failedCount >= maxFailureRequest
|
||||
|
||||
// Log the updated failure count.
|
||||
mainLog.Load().Debug().Msgf("upstream %q failure count updated to %d", upstream, failedCount)
|
||||
|
||||
// If this is the first failure and no timer is running, start a 10-second timer.
|
||||
if failedCount == 1 && !um.failureTimerActive[upstream] {
|
||||
um.failureTimerActive[upstream] = true
|
||||
go func(upstream string) {
|
||||
time.Sleep(10 * time.Second)
|
||||
um.mu.Lock()
|
||||
defer um.mu.Unlock()
|
||||
// If no success occurred during the 10-second window (i.e. counter remains > 0)
|
||||
// and the upstream is not in a recovered state, mark it as down.
|
||||
if um.failureReq[upstream] > 0 && !um.recovered[upstream] {
|
||||
um.down[upstream] = true
|
||||
mainLog.Load().Warn().Msgf("upstream %q marked as down after 10 seconds (failure count: %d)", upstream, um.failureReq[upstream])
|
||||
}
|
||||
// Reset the timer flag so that a new timer can be spawned if needed.
|
||||
um.failureTimerActive[upstream] = false
|
||||
}(upstream)
|
||||
}
|
||||
|
||||
// If the failure count quickly reaches the threshold, mark the upstream as down immediately.
|
||||
if failedCount >= maxFailureRequest {
|
||||
um.down[upstream] = true
|
||||
mainLog.Load().Warn().Msgf("upstream %q marked as down immediately (failure count: %d)", upstream, failedCount)
|
||||
}
|
||||
}
|
||||
|
||||
// isDown reports whether the given upstream is being marked as down.
|
||||
@@ -63,56 +118,28 @@ func (um *upstreamMonitor) isDown(upstream string) bool {
|
||||
// reset marks an upstream as up and set failed queries counter to zero.
|
||||
func (um *upstreamMonitor) reset(upstream string) {
|
||||
um.mu.Lock()
|
||||
defer um.mu.Unlock()
|
||||
|
||||
um.failureReq[upstream] = 0
|
||||
um.down[upstream] = false
|
||||
}
|
||||
|
||||
// checkUpstream checks the given upstream status, periodically sending query to upstream
|
||||
// until successfully. An upstream status/counter will be reset once it becomes reachable.
|
||||
func (p *prog) checkUpstream(upstream string, uc *ctrld.UpstreamConfig) {
|
||||
p.um.mu.Lock()
|
||||
isChecking := p.um.checking[upstream]
|
||||
if isChecking {
|
||||
p.um.mu.Unlock()
|
||||
return
|
||||
}
|
||||
p.um.checking[upstream] = true
|
||||
p.um.mu.Unlock()
|
||||
defer func() {
|
||||
p.um.mu.Lock()
|
||||
p.um.checking[upstream] = false
|
||||
p.um.mu.Unlock()
|
||||
um.recovered[upstream] = true
|
||||
um.mu.Unlock()
|
||||
go func() {
|
||||
// debounce the recovery to avoid incrementing failure counts already in flight
|
||||
time.Sleep(1 * time.Second)
|
||||
um.mu.Lock()
|
||||
um.recovered[upstream] = false
|
||||
um.mu.Unlock()
|
||||
}()
|
||||
}
|
||||
|
||||
resolver, err := ctrld.NewResolver(uc)
|
||||
if err != nil {
|
||||
mainLog.Load().Warn().Err(err).Msg("could not check upstream")
|
||||
return
|
||||
}
|
||||
msg := new(dns.Msg)
|
||||
msg.SetQuestion(".", dns.TypeNS)
|
||||
|
||||
check := func() error {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
uc.ReBootstrap()
|
||||
_, err := resolver.Resolve(ctx, msg)
|
||||
return err
|
||||
}
|
||||
for {
|
||||
if err := check(); err == nil {
|
||||
mainLog.Load().Debug().Msgf("upstream %q is online", uc.Endpoint)
|
||||
p.um.reset(upstream)
|
||||
if p.leakingQuery.CompareAndSwap(true, false) {
|
||||
p.leakingQueryMu.Lock()
|
||||
p.leakingQueryWasRun = false
|
||||
p.leakingQueryMu.Unlock()
|
||||
mainLog.Load().Warn().Msg("stop leaking query")
|
||||
}
|
||||
return
|
||||
// countHealthy returns the number of upstreams in the provided map that are considered healthy.
|
||||
func (um *upstreamMonitor) countHealthy(upstreams []string) int {
|
||||
var count int
|
||||
um.mu.RLock()
|
||||
for _, upstream := range upstreams {
|
||||
if !um.down[upstream] {
|
||||
count++
|
||||
}
|
||||
time.Sleep(checkUpstreamBackoffSleep)
|
||||
}
|
||||
um.mu.RUnlock()
|
||||
return count
|
||||
}
|
||||
|
||||
@@ -0,0 +1,469 @@
|
||||
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
|
||||
// handler (e.g., Tailscale's MagicDNS Network Extension) to receive queries
|
||||
// from all processes — not just ctrld.
|
||||
type vpnDNSExemption struct {
|
||||
Server string // DNS server IP (e.g., "100.100.100.100")
|
||||
Interface string // Interface name from scutil (e.g., "utun11"), may be empty
|
||||
IsExitMode bool // True if this VPN is in exit/full-tunnel mode (all traffic routed through VPN)
|
||||
}
|
||||
|
||||
// vpnDNSExemptFunc is called when VPN DNS servers change, to update
|
||||
// the intercept layer (WFP/pf) to permit VPN DNS traffic.
|
||||
type vpnDNSExemptFunc func(exemptions []vpnDNSExemption) error
|
||||
|
||||
// vpnDNSManager tracks active VPN DNS configurations and provides
|
||||
// domain-to-upstream routing for VPN split DNS.
|
||||
type vpnDNSManager struct {
|
||||
mu sync.RWMutex
|
||||
configs []ctrld.VPNDNSConfig
|
||||
// Map of domain suffix → DNS servers for fast lookup
|
||||
routes map[string][]string
|
||||
// DNS servers from VPN interfaces that have no domain/suffix config.
|
||||
// These are NOT added to the global OS resolver. They're only used
|
||||
// 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
|
||||
}
|
||||
|
||||
// newVPNDNSManager creates a new manager. Only call when dnsIntercept is active.
|
||||
// exemptFunc is called whenever VPN DNS servers are discovered/changed, to update
|
||||
// the OS-level intercept rules to permit ctrld's outbound queries to those IPs.
|
||||
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. 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")
|
||||
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
|
||||
// reliable than scutil flag parsing because the routing table is the ground
|
||||
// truth for traffic flow, regardless of how the VPN presents itself in scutil.
|
||||
if dri, err := netmon.DefaultRouteInterface(); err == nil && dri != "" {
|
||||
for i := range configs {
|
||||
if configs[i].InterfaceName == dri {
|
||||
if !configs[i].IsExitMode {
|
||||
logger.Info().Msgf("VPN DNS on %s: default route interface match — EXIT MODE (route-based detection)", dri)
|
||||
}
|
||||
configs[i].IsExitMode = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
|
||||
// Build domain -> DNS servers mapping
|
||||
for _, config := range configs {
|
||||
logger.Debug().Msgf("Processing VPN interface %s with %d domains and %d servers",
|
||||
config.InterfaceName, len(config.Domains), len(config.Servers))
|
||||
|
||||
for _, domain := range config.Domains {
|
||||
// Normalize domain: remove leading dot, Linux routing domain prefix (~),
|
||||
// and convert to lowercase.
|
||||
domain = strings.TrimPrefix(domain, "~")
|
||||
domain = strings.TrimPrefix(domain, ".")
|
||||
domain = strings.ToLower(domain)
|
||||
|
||||
if domain != "" {
|
||||
m.routes[domain] = append([]string{}, config.Servers...)
|
||||
logger.Debug().Msgf("Added VPN DNS route: %s -> %v", domain, config.Servers)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Collect unique VPN DNS exemptions (server + interface) for pf/WFP rules.
|
||||
type exemptionKey struct{ server, iface string }
|
||||
seen := make(map[exemptionKey]bool)
|
||||
var exemptions []vpnDNSExemption
|
||||
for _, config := range configs {
|
||||
for _, server := range config.Servers {
|
||||
key := exemptionKey{server, config.InterfaceName}
|
||||
if !seen[key] {
|
||||
seen[key] = true
|
||||
exemptions = append(exemptions, vpnDNSExemption{
|
||||
Server: server,
|
||||
Interface: config.InterfaceName,
|
||||
IsExitMode: config.IsExitMode,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Collect domain-less VPN DNS servers. These are NOT added to the global
|
||||
// OS resolver (that would pollute captive portal / DHCP flows). Instead,
|
||||
// they're stored separately and only used for queries that match existing
|
||||
// split-DNS rules (from ctrld config, AD domain, or VPN suffix config).
|
||||
var domainlessServers []string
|
||||
seen2 := make(map[string]bool)
|
||||
for _, config := range configs {
|
||||
if len(config.Domains) == 0 && len(config.Servers) > 0 {
|
||||
logger.Debug().Msgf("VPN interface %s has DNS servers but no domains, storing as split-rule fallback: %v",
|
||||
config.InterfaceName, config.Servers)
|
||||
for _, s := range config.Servers {
|
||||
if !seen2[s] {
|
||||
seen2[s] = true
|
||||
domainlessServers = append(domainlessServers, s)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
m.domainlessServers = domainlessServers
|
||||
|
||||
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 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.
|
||||
// Returns VPN DNS servers if matched, nil otherwise.
|
||||
func (m *vpnDNSManager) UpstreamForDomain(domain string) []string {
|
||||
if domain == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
domain = strings.TrimSuffix(domain, ".")
|
||||
domain = strings.ToLower(domain)
|
||||
|
||||
if servers, ok := m.routes[domain]; ok {
|
||||
return append([]string{}, servers...)
|
||||
}
|
||||
|
||||
for vpnDomain, servers := range m.routes {
|
||||
if strings.HasSuffix(domain, "."+vpnDomain) {
|
||||
return append([]string{}, servers...)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DomainlessServers returns VPN DNS servers that have no associated domains.
|
||||
// These should only be used for queries matching split-DNS rules, not for
|
||||
// general OS resolver queries (to avoid polluting captive portal / DHCP flows).
|
||||
func (m *vpnDNSManager) DomainlessServers() []string {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return append([]string{}, m.domainlessServers...)
|
||||
}
|
||||
|
||||
// CurrentServers returns the current set of unique VPN DNS server IPs.
|
||||
func (m *vpnDNSManager) CurrentServers() []string {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
seen := make(map[string]bool)
|
||||
var servers []string
|
||||
for _, ss := range m.routes {
|
||||
for _, s := range ss {
|
||||
if !seen[s] {
|
||||
seen[s] = true
|
||||
servers = append(servers, s)
|
||||
}
|
||||
}
|
||||
}
|
||||
return servers
|
||||
}
|
||||
|
||||
// CurrentExemptions returns VPN DNS server + interface pairs for pf exemption rules.
|
||||
func (m *vpnDNSManager) CurrentExemptions() []vpnDNSExemption {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.currentExemptionsLocked()
|
||||
}
|
||||
|
||||
// Routes returns a copy of the current VPN DNS routes for debugging.
|
||||
func (m *vpnDNSManager) Routes() map[string][]string {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
routes := make(map[string][]string)
|
||||
for domain, servers := range m.routes {
|
||||
routes[domain] = append([]string{}, servers...)
|
||||
}
|
||||
return routes
|
||||
}
|
||||
|
||||
// upstreamConfigFor creates a legacy upstream configuration for the given VPN DNS server.
|
||||
func (m *vpnDNSManager) upstreamConfigFor(server string) *ctrld.UpstreamConfig {
|
||||
// Use net.JoinHostPort to correctly handle both IPv4 and IPv6 addresses.
|
||||
// Previously, the strings.Contains(":") check would skip appending ":53"
|
||||
// for IPv6 addresses (they contain colons), leaving a bare address like
|
||||
// "2a0d:6fc0:9b0:3600::1" which net.Dial rejects with "too many colons".
|
||||
// net.JoinHostPort produces "[2a0d:6fc0:9b0:3600::1]:53" as required.
|
||||
endpoint := net.JoinHostPort(server, "53")
|
||||
|
||||
return &ctrld.UpstreamConfig{
|
||||
Name: "VPN DNS",
|
||||
Type: ctrld.ResolverTypeLegacy,
|
||||
Endpoint: endpoint,
|
||||
Timeout: 2000,
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+7
-1
@@ -1,7 +1,13 @@
|
||||
package main
|
||||
|
||||
import "github.com/Control-D-Inc/ctrld/cmd/cli"
|
||||
import (
|
||||
"os"
|
||||
|
||||
"github.com/Control-D-Inc/ctrld/cmd/cli"
|
||||
)
|
||||
|
||||
func main() {
|
||||
cli.Main()
|
||||
// make sure we exit with 0 if there are no errors
|
||||
os.Exit(0)
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user