mirror of
https://github.com/Control-D-Inc/ctrld.git
synced 2026-07-16 13:17:19 +02:00
Compare commits
439 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 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 | |||
| 5b9ccc5065 | |||
| 6ca1a7ccc7 | |||
| 9d666be5d4 | |||
| 65de7edcde | |||
| 0cdff0d368 | |||
| f87220a908 | |||
| 30ea0c6499 | |||
| 9501e35c60 | |||
| 5ac9d17bdf | |||
| cb14992ddc | |||
| e88372fc8c | |||
| b320662d67 | |||
| ce353cd4d9 | |||
| 4befd33866 | |||
| 4b36e3ac44 | |||
| f507bc8f9e | |||
| 14c88f4a6d | |||
| 3e388c2857 | |||
| cfe1209d61 | |||
| 5a88a7c22c | |||
| 8c661c4401 | |||
| e6f256d640 | |||
| ede354166b | |||
| 282a8ce78e | |||
| 08fe04f1ee | |||
| 082d14a9ba | |||
| 617674ce43 | |||
| 7088df58dd | |||
| 9cbd9b3e44 | |||
| e6586fd360 | |||
| 33a6db2599 | |||
| 70b0c4f7b9 | |||
| 5af3ec4f7b | |||
| 79476add12 | |||
| 1634a06330 | |||
| a007394f60 | |||
| 62a0ba8731 | |||
| e8d3ed1acd | |||
| 8b98faa441 | |||
| 30320ec9c7 | |||
| 5f4a399850 | |||
| 82e0d4b0c4 | |||
| 95a9df826d | |||
| 3b71d26cf3 | |||
| c233ad9b1b | |||
| 12d6484b1c | |||
| bc7b1cc6d8 | |||
| ec684348ed | |||
| 18a19a3aa2 | |||
| 905f2d08c5 | |||
| 04947b4d87 | |||
| 72bf80533e | |||
| 9ddedf926e | |||
| 139dd62ff3 | |||
| 50ef00526e | |||
| 80cf79b9cb | |||
| e6ad39b070 | |||
| 56f9c72569 | |||
| dc48c908b8 | |||
| 9b0f0e792a | |||
| b3eebb19b6 | |||
| c24589a5be | |||
| 1e1c5a4dc8 | |||
| 339023421a | |||
| a00d2a431a | |||
| 5aca118dbb | |||
| 411f7434f4 | |||
| 34801382f5 | |||
| b9f2259ae4 | |||
| 19020a96bf | |||
| 96085147ff | |||
| f3dd344026 | |||
| 486096416f | |||
| 5710f2e984 | |||
| 09936f1f07 | |||
| 0d6ca57536 | |||
| 3ddcb84db8 | |||
| 1012bf063f | |||
| b8155e6182 | |||
| 9a34df61bb | |||
| fbb879edf9 | |||
| ac97c88876 | |||
| a1fda2c0de | |||
| f499770d45 | |||
| 4769da4ef4 | |||
| c2556a8e39 | |||
| 29bf329f6a | |||
| 1dee4305bc | |||
| 429a98b690 | |||
| da01a146d2 | |||
| dd9f2465be | |||
| b5cf0e2b31 | |||
| 1db159ad34 | |||
| 6604f973ac | |||
| 69ee6582e2 | |||
| 6f12667e8c | |||
| b002dff624 | |||
| affef963c1 | |||
| 56b2056190 | |||
| c1e6f5126a | |||
| 1a8c1ec73d | |||
| 52954b8ceb | |||
| a5025e35ea | |||
| 07f80c9ebf | |||
| 13db23553d | |||
| 3963fce43b | |||
| ea4e5147bd | |||
| 7a491a4cc5 | |||
| 5ba90748f6 | |||
| 20f8f22bae | |||
| b50cccac85 | |||
| 34ebe9b054 | |||
| 43d82cf1a7 | |||
| ab88174091 | |||
| ebcbf85373 | |||
| 87513cba6d | |||
| 64bcd2f00d | |||
| cc6ae290f8 | |||
| 3e62bd3dbd | |||
| 8491f9c455 | |||
| 3ca754b438 | |||
| 8c7c3901e8 | |||
| a9672dfff5 | |||
| 203a2ec8b8 | |||
| 810cbd1f4f | |||
| 49eebcdcbc | |||
| e89021ec3a | |||
| 73a697b2fa | |||
| 9319d08046 | |||
| 7dc5138e91 | |||
| 8f189c919a | |||
| 906479a15c | |||
| dabbf2037b | |||
| b496147ce7 | |||
| 583718f234 | |||
| fdb82f6ec3 | |||
| 5145729ab1 | |||
| 4d810261a4 | |||
| 18e8616834 | |||
| d55563cac5 | |||
| bb481d9bcc | |||
| a163be3584 | |||
| 891b7cb2c6 | |||
| 176c22f229 | |||
| faa0ed06b6 | |||
| 9515db7faf | |||
| d822bf4257 | |||
| 0826671809 | |||
| 67d74774a9 | |||
| 5d65416227 | |||
| 49441f62f3 | |||
| 99651f6e5b | |||
| edca1f4f89 | |||
| 3d834f00f6 | |||
| 6bb9e7a766 | |||
| 61fb71b1fa | |||
| f8967c376f |
@@ -9,18 +9,18 @@ jobs:
|
|||||||
fail-fast: false
|
fail-fast: false
|
||||||
matrix:
|
matrix:
|
||||||
os: ["windows-latest", "ubuntu-latest", "macOS-latest"]
|
os: ["windows-latest", "ubuntu-latest", "macOS-latest"]
|
||||||
go: ["1.20.x"]
|
go: ["1.25.x"]
|
||||||
runs-on: ${{ matrix.os }}
|
runs-on: ${{ matrix.os }}
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v3
|
- uses: actions/checkout@v3
|
||||||
with:
|
with:
|
||||||
fetch-depth: 1
|
fetch-depth: 1
|
||||||
- uses: WillAbides/setup-go-faster@v1.8.0
|
- uses: actions/setup-go@v6
|
||||||
with:
|
with:
|
||||||
go-version: ${{ matrix.go }}
|
go-version: ${{ matrix.go }}
|
||||||
- run: "go test -race ./..."
|
- run: "go test -race ./..."
|
||||||
- uses: dominikh/staticcheck-action@v1.2.0
|
- uses: dominikh/staticcheck-action@v1.4.0
|
||||||
with:
|
with:
|
||||||
version: "2023.1.2"
|
version: "2026.1"
|
||||||
install-go: false
|
install-go: false
|
||||||
cache-key: ${{ matrix.go }}
|
cache-key: ${{ matrix.go }}
|
||||||
|
|||||||
+11
@@ -3,3 +3,14 @@ gon.hcl
|
|||||||
|
|
||||||
/Build
|
/Build
|
||||||
.DS_Store
|
.DS_Store
|
||||||
|
|
||||||
|
# Release folder
|
||||||
|
dist/
|
||||||
|
|
||||||
|
# Binaries
|
||||||
|
ctrld-*
|
||||||
|
|
||||||
|
# generated file
|
||||||
|
cmd/cli/rsrc_*.syso
|
||||||
|
ctrld
|
||||||
|
ctrld.exe
|
||||||
|
|||||||
@@ -4,13 +4,16 @@
|
|||||||
[](https://pkg.go.dev/github.com/Control-D-Inc/ctrld)
|
[](https://pkg.go.dev/github.com/Control-D-Inc/ctrld)
|
||||||
[](https://goreportcard.com/report/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:
|
A highly configurable DNS forwarding proxy with support for:
|
||||||
- Multiple listeners for incoming queries
|
- Multiple listeners for incoming queries
|
||||||
- Multiple upstreams with fallbacks
|
- 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
|
- Policy driven domain based "split horizon" DNS with wildcard support
|
||||||
- Integrations with common router vendors and firmware
|
- Integrations with common router vendors and firmware
|
||||||
- LAN client discovery via DHCP, mDNS, and ARP
|
- LAN client discovery via DHCP, mDNS, ARP, NDP, hosts file parsing
|
||||||
|
- Prometheus metrics exporter
|
||||||
|
|
||||||
## TLDR
|
## TLDR
|
||||||
Proxy legacy DNS traffic to secure DNS upstreams in highly configurable ways.
|
Proxy legacy DNS traffic to secure DNS upstreams in highly configurable ways.
|
||||||
@@ -32,13 +35,29 @@ All DNS protocols are supported, including:
|
|||||||
|
|
||||||
## OS Support
|
## OS Support
|
||||||
- Windows (386, amd64, arm)
|
- Windows (386, amd64, arm)
|
||||||
- Mac (amd64, arm64)
|
- Windows Server (386, amd64)
|
||||||
|
- MacOS (amd64, arm64)
|
||||||
- Linux (386, amd64, arm, mips)
|
- Linux (386, amd64, arm, mips)
|
||||||
- FreeBSD
|
- FreeBSD (386, amd64, arm)
|
||||||
- Common routers (See Router Mode below)
|
- 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
|
# Install
|
||||||
There are several ways to download and install `ctrld.
|
There are several ways to download and install `ctrld`.
|
||||||
|
|
||||||
## Quick Install
|
## Quick Install
|
||||||
The simplest way to download and install `ctrld` is to use the following installer command on any UNIX-like platform:
|
The simplest way to download and install `ctrld` is to use the following installer command on any UNIX-like platform:
|
||||||
@@ -47,42 +66,41 @@ 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)"'
|
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
|
```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)
|
Or you can pull and run a Docker container from [Docker Hub](https://hub.docker.com/r/controldns/ctrld)
|
||||||
```
|
```shell
|
||||||
$ docker pull controldns/ctrld
|
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
|
## Download Manually
|
||||||
Alternatively, if you know what you're doing you can download pre-compiled binaries from the [Releases](https://github.com/Control-D-Inc/ctrld/releases) section for the appropriate platform.
|
Alternatively, if you know what you're doing you can download pre-compiled binaries from the [Releases](https://github.com/Control-D-Inc/ctrld/releases) section for the appropriate platform.
|
||||||
|
|
||||||
## Build
|
## Build
|
||||||
Lastly, you can build `ctrld` from source which requires `go1.19+`:
|
Lastly, you can build `ctrld` from source which requires `go1.21+`:
|
||||||
|
|
||||||
```shell
|
```shell
|
||||||
$ go build ./cmd/ctrld
|
go build ./cmd/ctrld
|
||||||
```
|
```
|
||||||
|
|
||||||
or
|
or
|
||||||
|
|
||||||
```shell
|
```shell
|
||||||
$ go install github.com/Control-D-Inc/ctrld/cmd/ctrld@latest
|
go install github.com/Control-D-Inc/ctrld/cmd/ctrld@latest
|
||||||
```
|
```
|
||||||
|
|
||||||
or
|
or
|
||||||
|
|
||||||
```
|
```shell
|
||||||
$ docker build -t controldns/ctrld . -f docker/Dockerfile
|
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
|
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
||||||
# Usage
|
# 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
|
## Arguments
|
||||||
```
|
```
|
||||||
@@ -98,13 +116,16 @@ Usage:
|
|||||||
|
|
||||||
Available Commands:
|
Available Commands:
|
||||||
run Run the DNS proxy server
|
run Run the DNS proxy server
|
||||||
service Manage ctrld service
|
|
||||||
start Quick start service and configure DNS on interface
|
start Quick start service and configure DNS on interface
|
||||||
stop Quick stop service and remove DNS from interface
|
stop Quick stop service and remove DNS from interface
|
||||||
restart Restart the ctrld service
|
restart Restart the ctrld service
|
||||||
|
reload Reload the ctrld service
|
||||||
status Show status of the ctrld service
|
status Show status of the ctrld service
|
||||||
uninstall Stop and uninstall the ctrld service
|
uninstall Stop and uninstall the ctrld service
|
||||||
|
service Manage ctrld service
|
||||||
clients Manage clients
|
clients Manage clients
|
||||||
|
upgrade Upgrading ctrld to latest version
|
||||||
|
log Manage runtime debug logs
|
||||||
|
|
||||||
Flags:
|
Flags:
|
||||||
-h, --help help for ctrld
|
-h, --help help for ctrld
|
||||||
@@ -116,81 +137,99 @@ Use "ctrld [command] --help" for more information about a command.
|
|||||||
```
|
```
|
||||||
|
|
||||||
## Basic Run Mode
|
## 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.
|
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.
|
||||||
1. Start the server
|
|
||||||
```
|
### Command
|
||||||
$ sudo ./ctrld run
|
|
||||||
|
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
|
$ dig verify.controld.com @127.0.0.1 +short
|
||||||
api.controld.com.
|
api.controld.com.
|
||||||
147.185.34.1
|
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
|
## 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
|
Linux or Macos
|
||||||
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
|
sudo ctrld start
|
||||||
- DD-WRT
|
```
|
||||||
- Firewalla
|
|
||||||
- FreshTomato
|
|
||||||
- GL.iNet
|
|
||||||
- OpenWRT
|
|
||||||
- pfSense / OPNsense
|
|
||||||
- Synology
|
|
||||||
- Ubiquiti (UniFi, EdgeOS)
|
|
||||||
|
|
||||||
`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
|
## Unmanaged Service Mode
|
||||||
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.
|
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
|
Windows (Admin Shell)
|
||||||
./ctrld run --cd p2
|
```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.
|
Linux or Macos
|
||||||
|
```shell
|
||||||
```shell
|
sudo ctrld service start
|
||||||
./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
|
|
||||||
|
|
||||||
# Configuration
|
# 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
|
## API Based Auto Configuration
|
||||||
- Start `listener.0` on 127.0.0.1:53
|
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.
|
||||||
- Accept queries from any source address
|
|
||||||
- Send all queries to `upstream.0` via DoH protocol
|
|
||||||
|
|
||||||
### 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
|
```toml
|
||||||
[listener]
|
[listener]
|
||||||
|
|
||||||
[listener.0]
|
[listener.0]
|
||||||
ip = ""
|
ip = '0.0.0.0'
|
||||||
port = 0
|
port = 53
|
||||||
restricted = false
|
|
||||||
|
|
||||||
[network]
|
[network]
|
||||||
|
|
||||||
@@ -198,10 +237,6 @@ See [Configuration Docs](docs/config.md).
|
|||||||
cidrs = ["0.0.0.0/0"]
|
cidrs = ["0.0.0.0/0"]
|
||||||
name = "Network 0"
|
name = "Network 0"
|
||||||
|
|
||||||
[service]
|
|
||||||
log_level = "info"
|
|
||||||
log_path = ""
|
|
||||||
|
|
||||||
[upstream]
|
[upstream]
|
||||||
|
|
||||||
[upstream.0]
|
[upstream.0]
|
||||||
@@ -210,29 +245,88 @@ See [Configuration Docs](docs/config.md).
|
|||||||
name = "Control D - Anti-Malware"
|
name = "Control D - Anti-Malware"
|
||||||
timeout = 5000
|
timeout = 5000
|
||||||
type = "doh"
|
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
|
## CLI Args
|
||||||
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.
|
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
|
## Contributing
|
||||||
See [Contribution Guideline](./docs/contributing.md)
|
See [Contribution Guideline](./docs/contributing.md)
|
||||||
|
|
||||||
## Roadmap
|
|
||||||
The following functionality is on the roadmap and will be available in future releases.
|
|
||||||
- Prometheus metrics exporter
|
|
||||||
- DNS intercept mode
|
|
||||||
- Direct listener mode
|
|
||||||
- Support for more routers (let us know which ones)
|
|
||||||
|
|||||||
@@ -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()
|
||||||
|
}
|
||||||
@@ -0,0 +1,15 @@
|
|||||||
|
//go:build !windows
|
||||||
|
|
||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/Control-D-Inc/ctrld"
|
||||||
|
)
|
||||||
|
|
||||||
|
// addExtraSplitDnsRule adds split DNS rule if present.
|
||||||
|
func addExtraSplitDnsRule(_ *ctrld.Config) bool { return false }
|
||||||
|
|
||||||
|
// getActiveDirectoryDomain returns AD domain name of this computer.
|
||||||
|
func getActiveDirectoryDomain() (string, error) {
|
||||||
|
return "", nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,74 @@
|
|||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"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) bool {
|
||||||
|
domain, err := system.GetActiveDirectoryDomain()
|
||||||
|
if err != nil {
|
||||||
|
mainLog.Load().Debug().Msgf("unable to get active directory domain: %v", err)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if domain == "" {
|
||||||
|
mainLog.Load().Debug().Msg("no active directory domain found")
|
||||||
|
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{}
|
||||||
|
}
|
||||||
|
for _, rule := range lc.Policy.Rules {
|
||||||
|
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 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) {
|
||||||
|
log.SetOutput(io.Discard)
|
||||||
|
defer log.SetOutput(os.Stderr)
|
||||||
|
whost := host.NewWmiLocalHost()
|
||||||
|
cs, err := hh.GetComputerSystem(whost)
|
||||||
|
if cs != nil {
|
||||||
|
defer cs.Close()
|
||||||
|
}
|
||||||
|
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
|
||||||
+923
-883
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,28 @@
|
|||||||
|
package cli
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
+24
-1
@@ -16,8 +16,31 @@ func Test_writeConfigFile(t *testing.T) {
|
|||||||
_, err := os.Stat(configPath)
|
_, err := os.Stat(configPath)
|
||||||
assert.True(t, os.IsNotExist(err))
|
assert.True(t, os.IsNotExist(err))
|
||||||
|
|
||||||
assert.NoError(t, writeConfigFile())
|
assert.NoError(t, writeConfigFile(&cfg))
|
||||||
|
|
||||||
_, err = os.Stat(configPath)
|
_, err = os.Stat(configPath)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func Test_isStableVersion(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
ver string
|
||||||
|
isStable bool
|
||||||
|
}{
|
||||||
|
{"stable", "v1.3.5", true},
|
||||||
|
{"pre", "v1.3.5-next", false},
|
||||||
|
{"pre with commit hash", "v1.3.5-next-asdf", false},
|
||||||
|
{"dev", "dev", false},
|
||||||
|
{"empty", "dev", false},
|
||||||
|
}
|
||||||
|
for _, tc := range tests {
|
||||||
|
tc := tc
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
if got := isStableVersion(tc.ver); got != tc.isStable {
|
||||||
|
t.Errorf("unexpected result for %s, want: %v, got: %v", tc.ver, tc.isStable, got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+1606
File diff suppressed because it is too large
Load Diff
@@ -25,5 +25,20 @@ func newControlClient(addr string) *controlClient {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *controlClient) post(path string, data io.Reader) (*http.Response, error) {
|
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)
|
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)
|
||||||
|
}
|
||||||
|
|
||||||
|
// deactivationRequest represents request for validating deactivation pin.
|
||||||
|
type deactivationRequest struct {
|
||||||
|
Pin int64 `json:"pin"`
|
||||||
|
}
|
||||||
|
|||||||
+393
-11
@@ -3,25 +3,43 @@ package cli
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"os"
|
"os"
|
||||||
"reflect"
|
"reflect"
|
||||||
"sort"
|
"sort"
|
||||||
|
"strconv"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/kardianos/service"
|
||||||
dto "github.com/prometheus/client_model/go"
|
dto "github.com/prometheus/client_model/go"
|
||||||
|
|
||||||
"github.com/Control-D-Inc/ctrld"
|
"github.com/Control-D-Inc/ctrld"
|
||||||
|
"github.com/Control-D-Inc/ctrld/internal/controld"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
contentTypeJson = "application/json"
|
contentTypeJson = "application/json"
|
||||||
listClientsPath = "/clients"
|
listClientsPath = "/clients"
|
||||||
startedPath = "/started"
|
startedPath = "/started"
|
||||||
reloadPath = "/reload"
|
reloadPath = "/reload"
|
||||||
|
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 {
|
type controlServer struct {
|
||||||
server *http.Server
|
server *http.Server
|
||||||
mux *http.ServeMux
|
mux *http.ServeMux
|
||||||
@@ -41,12 +59,18 @@ func newControlServer(addr string) (*controlServer, error) {
|
|||||||
func (s *controlServer) start() error {
|
func (s *controlServer) start() error {
|
||||||
_ = os.Remove(s.addr)
|
_ = os.Remove(s.addr)
|
||||||
unixListener, err := net.Listen("unix", s.addr)
|
unixListener, err := net.Listen("unix", s.addr)
|
||||||
if l, ok := unixListener.(*net.UnixListener); ok {
|
|
||||||
l.SetUnlinkOnClose(true)
|
|
||||||
}
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
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)
|
go s.server.Serve(unixListener)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -64,33 +88,81 @@ func (s *controlServer) register(pattern string, handler http.Handler) {
|
|||||||
|
|
||||||
func (p *prog) registerControlServerHandler() {
|
func (p *prog) registerControlServerHandler() {
|
||||||
p.cs.register(listClientsPath, http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
|
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()
|
clients := p.ciTable.ListClients()
|
||||||
|
mainLog.Load().Debug().Int("client_count", len(clients)).Msg("retrieved clients list")
|
||||||
|
|
||||||
sort.Slice(clients, func(i, j int) bool {
|
sort.Slice(clients, func(i, j int) bool {
|
||||||
return clients[i].IP.Less(clients[j].IP)
|
return clients[i].IP.Less(clients[j].IP)
|
||||||
})
|
})
|
||||||
if p.cfg.Service.MetricsQueryStats {
|
mainLog.Load().Debug().Msg("sorted clients by IP address")
|
||||||
for _, client := range clients {
|
|
||||||
|
if p.metricsQueryStats.Load() {
|
||||||
|
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
|
client.IncludeQueryCount = true
|
||||||
dm := &dto.Metric{}
|
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(
|
m, err := statsClientQueriesCount.MetricVec.GetMetricWithLabelValues(
|
||||||
client.IP.String(),
|
client.IP.String(),
|
||||||
client.Mac,
|
client.Mac,
|
||||||
client.Hostname,
|
client.Hostname,
|
||||||
)
|
)
|
||||||
if err != nil {
|
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
|
continue
|
||||||
}
|
}
|
||||||
if err := m.Write(dm); err == nil {
|
|
||||||
|
if err := m.Write(dm); err == nil && dm.Counter != nil {
|
||||||
client.QueryCount = int64(dm.Counter.GetValue())
|
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 {
|
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)
|
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||||
return
|
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) {
|
p.cs.register(startedPath, http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
|
||||||
select {
|
select {
|
||||||
@@ -146,6 +218,316 @@ func (p *prog) registerControlServerHandler() {
|
|||||||
// Otherwise, reload is done.
|
// Otherwise, reload is done.
|
||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
}))
|
}))
|
||||||
|
p.cs.register(deactivationPath, http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
|
||||||
|
// 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(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
|
||||||
|
}
|
||||||
|
|
||||||
|
var req deactivationRequest
|
||||||
|
if err := json.NewDecoder(request.Body).Decode(&req); err != nil {
|
||||||
|
w.WriteHeader(http.StatusPreconditionFailed)
|
||||||
|
mainLog.Load().Err(err).Msg("invalid deactivation request")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
code := http.StatusForbidden
|
||||||
|
switch req.Pin {
|
||||||
|
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)
|
||||||
|
}))
|
||||||
|
p.cs.register(cdPath, http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
|
||||||
|
if cdUID != "" {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
w.Write([]byte(cdUID))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
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 {
|
||||||
|
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
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
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(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 {
|
func jsonResponse(next http.Handler) http.Handler {
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,188 @@
|
|||||||
|
//go:build darwin
|
||||||
|
|
||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/Control-D-Inc/ctrld"
|
||||||
|
)
|
||||||
|
|
||||||
|
// =============================================================================
|
||||||
|
// buildPFAnchorRules tests
|
||||||
|
// =============================================================================
|
||||||
|
|
||||||
|
func TestPFBuildAnchorRules_Basic(t *testing.T) {
|
||||||
|
p := &prog{cfg: &ctrld.Config{Listener: map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.1", Port: 53}}}}
|
||||||
|
rules := p.buildPFAnchorRules(nil)
|
||||||
|
|
||||||
|
// rdr (translation) must come before pass (filtering)
|
||||||
|
rdrIdx := strings.Index(rules, "rdr on lo0 inet proto udp")
|
||||||
|
passRouteIdx := strings.Index(rules, "pass out quick on ! lo0 route-to lo0 inet proto udp")
|
||||||
|
passInIdx := strings.Index(rules, "pass in quick on lo0 reply-to lo0")
|
||||||
|
|
||||||
|
if rdrIdx < 0 {
|
||||||
|
t.Fatal("missing rdr rule")
|
||||||
|
}
|
||||||
|
if passRouteIdx < 0 {
|
||||||
|
t.Fatal("missing pass out route-to rule")
|
||||||
|
}
|
||||||
|
if passInIdx < 0 {
|
||||||
|
t.Fatal("missing pass in on lo0 rule")
|
||||||
|
}
|
||||||
|
if rdrIdx >= passRouteIdx {
|
||||||
|
t.Error("rdr rules must come before pass out route-to rules")
|
||||||
|
}
|
||||||
|
if passRouteIdx >= passInIdx {
|
||||||
|
t.Error("pass out route-to must come before pass in on lo0")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Both UDP and TCP rdr rules
|
||||||
|
if !strings.Contains(rules, "proto udp") || !strings.Contains(rules, "proto tcp") {
|
||||||
|
t.Error("must have both UDP and TCP rdr rules")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPFBuildAnchorRules_WithVPNServers(t *testing.T) {
|
||||||
|
p := &prog{cfg: &ctrld.Config{Listener: map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.1", Port: 53}}}}
|
||||||
|
vpnServers := []vpnDNSExemption{
|
||||||
|
{Server: "10.8.0.1"},
|
||||||
|
{Server: "10.8.0.2"},
|
||||||
|
}
|
||||||
|
rules := p.buildPFAnchorRules(vpnServers)
|
||||||
|
|
||||||
|
// VPN exemption rules must appear
|
||||||
|
for _, s := range vpnServers {
|
||||||
|
if !strings.Contains(rules, s.Server) {
|
||||||
|
t.Errorf("missing VPN exemption for %s", s.Server)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// VPN exemptions must come before route-to
|
||||||
|
exemptIdx := strings.Index(rules, "10.8.0.1 port 53 group")
|
||||||
|
routeIdx := strings.Index(rules, "pass out quick on ! lo0 route-to lo0 inet proto udp")
|
||||||
|
if exemptIdx < 0 {
|
||||||
|
t.Fatal("missing VPN exemption rule for 10.8.0.1")
|
||||||
|
}
|
||||||
|
if routeIdx < 0 {
|
||||||
|
t.Fatal("missing route-to rule")
|
||||||
|
}
|
||||||
|
if exemptIdx >= routeIdx {
|
||||||
|
t.Error("VPN exemptions must come before route-to rules")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPFBuildAnchorRules_IPv4AndIPv6VPN(t *testing.T) {
|
||||||
|
p := &prog{cfg: &ctrld.Config{Listener: map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.1", Port: 53}}}}
|
||||||
|
vpnServers := []vpnDNSExemption{
|
||||||
|
{Server: "10.8.0.1"},
|
||||||
|
{Server: "fd00::1"},
|
||||||
|
}
|
||||||
|
rules := p.buildPFAnchorRules(vpnServers)
|
||||||
|
|
||||||
|
// IPv4 server should use "inet"
|
||||||
|
lines := strings.Split(rules, "\n")
|
||||||
|
for _, line := range lines {
|
||||||
|
if strings.Contains(line, "10.8.0.1") && strings.HasPrefix(line, "pass") {
|
||||||
|
if !strings.Contains(line, "inet ") {
|
||||||
|
t.Error("IPv4 VPN server rule should contain 'inet'")
|
||||||
|
}
|
||||||
|
if strings.Contains(line, "inet6") {
|
||||||
|
t.Error("IPv4 VPN server rule should not contain 'inet6'")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if strings.Contains(line, "fd00::1") && strings.HasPrefix(line, "pass") {
|
||||||
|
if !strings.Contains(line, "inet6") {
|
||||||
|
t.Error("IPv6 VPN server rule should contain 'inet6'")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPFBuildAnchorRules_Ordering(t *testing.T) {
|
||||||
|
p := &prog{cfg: &ctrld.Config{Listener: map[string]*ctrld.ListenerConfig{"0": {IP: "127.0.0.1", Port: 53}}}}
|
||||||
|
vpnServers := []vpnDNSExemption{
|
||||||
|
{Server: "10.8.0.1"},
|
||||||
|
}
|
||||||
|
rules := p.buildPFAnchorRules(vpnServers)
|
||||||
|
|
||||||
|
// Verify ordering: rdr → exemptions → route-to → pass in on lo0
|
||||||
|
rdrIdx := strings.Index(rules, "rdr on lo0 inet proto udp")
|
||||||
|
exemptIdx := strings.Index(rules, "pass out quick on ! lo0 inet proto { udp, tcp } from any to 10.8.0.1 port 53 group _ctrld")
|
||||||
|
routeIdx := strings.Index(rules, "pass out quick on ! lo0 route-to lo0 inet proto udp")
|
||||||
|
passInIdx := strings.Index(rules, "pass in quick on lo0 reply-to lo0")
|
||||||
|
|
||||||
|
if rdrIdx < 0 || exemptIdx < 0 || routeIdx < 0 || passInIdx < 0 {
|
||||||
|
t.Fatalf("missing expected rules: rdr=%d exempt=%d route=%d passIn=%d", rdrIdx, exemptIdx, routeIdx, passInIdx)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !(rdrIdx < exemptIdx && exemptIdx < routeIdx && routeIdx < passInIdx) {
|
||||||
|
t.Errorf("incorrect rule ordering: rdr(%d) < exempt(%d) < route(%d) < passIn(%d)", rdrIdx, exemptIdx, routeIdx, passInIdx)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestPFAddressFamily tests the pfAddressFamily helper.
|
||||||
|
func TestPFAddressFamily(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
ip string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"10.0.0.1", "inet"},
|
||||||
|
{"192.168.1.1", "inet"},
|
||||||
|
{"127.0.0.1", "inet"},
|
||||||
|
{"::1", "inet6"},
|
||||||
|
{"fd00::1", "inet6"},
|
||||||
|
{"2001:db8::1", "inet6"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := pfAddressFamily(tt.ip); got != tt.want {
|
||||||
|
t.Errorf("pfAddressFamily(%q) = %q, want %q", tt.ip, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsResourceExhaustion(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
err error
|
||||||
|
output []byte
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "exec start failure",
|
||||||
|
err: errors.New("fork/exec /sbin/pfctl: resource temporarily unavailable"),
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "fd exhaustion from stderr output",
|
||||||
|
err: errors.New("exit status 1"),
|
||||||
|
output: []byte("pfctl: Pipe: Too many open files"),
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "process exhaustion from wrapped restore error",
|
||||||
|
err: errors.New("failed to dump running filter rules: exit status 1 (output: too many processes)"),
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "ordinary pf syntax failure",
|
||||||
|
err: errors.New("exit status 1"),
|
||||||
|
output: []byte("pfctl: syntax error"),
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "nil error and empty output",
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := isResourceExhaustion(tt.err, tt.output); got != tt.want {
|
||||||
|
t.Fatalf("isResourceExhaustion() = %v, want %v", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,43 @@
|
|||||||
|
//go:build !windows && !darwin
|
||||||
|
|
||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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
|
||||||
|
}
|
||||||
|
|
||||||
|
// exemptVPNDNSServers is a no-op on unsupported platforms.
|
||||||
|
func (p *prog) exemptVPNDNSServers(exemptions []vpnDNSExemption) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ensurePFAnchorActive is a no-op on unsupported platforms.
|
||||||
|
func (p *prog) ensurePFAnchorActive() bool {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// checkTunnelInterfaceChanges is a no-op on unsupported platforms.
|
||||||
|
func (p *prog) checkTunnelInterfaceChanges() 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,50 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
beforeExemptions := p.vpnDNS.CurrentExemptions()
|
||||||
|
routes, domainlessServers, exemptions = p.vpnDNS.RefreshRoutesOnly()
|
||||||
|
afterExemptions := p.vpnDNS.CurrentExemptions()
|
||||||
|
|
||||||
|
if vpnDNSExemptionsEqual(beforeExemptions, afterExemptions) {
|
||||||
|
mainLog.Load().Info().Msgf("DNS intercept: post-settle VPN DNS route refresh completed — %d routes, %d domainless servers, %d exemptions (pf unchanged)",
|
||||||
|
routes, domainlessServers, exemptions)
|
||||||
|
return routes, domainlessServers, exemptions
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := p.exemptVPNDNSServers(afterExemptions); err != nil {
|
||||||
|
mainLog.Load().Warn().Err(err).Msg("DNS intercept: post-settle VPN DNS exemption update failed")
|
||||||
|
} else {
|
||||||
|
mainLog.Load().Info().Msgf("DNS intercept: post-settle VPN DNS exemptions changed — updated pf/WFP with %d exemptions", len(afterExemptions))
|
||||||
|
}
|
||||||
|
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,49 @@
|
|||||||
|
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) != 0 {
|
||||||
|
t.Fatalf("expected route-only refresh to avoid pf exemption updates, got %+v", exemptionUpdates)
|
||||||
|
}
|
||||||
|
}
|
||||||
File diff suppressed because it is too large
Load Diff
+1224
-76
File diff suppressed because it is too large
Load Diff
+54
-11
@@ -22,14 +22,22 @@ func Test_wildcardMatches(t *testing.T) {
|
|||||||
domain string
|
domain string
|
||||||
match bool
|
match bool
|
||||||
}{
|
}{
|
||||||
{"prefix parent should not match", "*.windscribe.com", "windscribe.com", false},
|
{"domain - prefix parent should not match", "*.example.com", "example.com", false},
|
||||||
{"prefix", "*.windscribe.com", "anything.windscribe.com", true},
|
{"domain - prefix", "*.example.com", "anything.example.com", true},
|
||||||
{"prefix not match other domain", "*.windscribe.com", "example.com", false},
|
{"domain - prefix not match other s", "*.example.com", "other.org", false},
|
||||||
{"prefix not match domain in name", "*.windscribe.com", "wwindscribe.com", false},
|
{"domain - prefix not match s in name", "*.example.com", "eexample.com", false},
|
||||||
{"suffix", "suffix.*", "suffix.windscribe.com", true},
|
{"domain - suffix", "suffix.*", "suffix.example.com", true},
|
||||||
{"suffix not match other", "suffix.*", "suffix1.windscribe.com", false},
|
{"domain - suffix not match other", "suffix.*", "suffix1.example.com", false},
|
||||||
{"both", "suffix.*.windscribe.com", "suffix.anything.windscribe.com", true},
|
{"domain - both", "suffix.*.example.com", "suffix.anything.example.com", true},
|
||||||
{"both not match", "suffix.*.windscribe.com", "suffix1.suffix.windscribe.com", false},
|
{"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},
|
||||||
|
{"mac - suffix", "d4:67:98:*", "d4:67:98:05:b4:2b", true},
|
||||||
|
{"mac - suffix not match other", "d4:67:98:*", "d4:67:97:15:b4:2b", false},
|
||||||
|
{"mac - both", "d4:67:98:*:b4:2b", "d4:67:98:05:b4:2b", true},
|
||||||
|
{"mac - both not match", "d4:67:98:*:b4:2b", "d4:67:97:05:c4:2b", false},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tc := range tests {
|
for _, tc := range tests {
|
||||||
@@ -49,9 +57,9 @@ func Test_canonicalName(t *testing.T) {
|
|||||||
domain string
|
domain string
|
||||||
canonical string
|
canonical string
|
||||||
}{
|
}{
|
||||||
{"fqdn to canonical", "windscribe.com.", "windscribe.com"},
|
{"fqdn to canonical", "example.com.", "example.com"},
|
||||||
{"already canonical", "windscribe.com", "windscribe.com"},
|
{"already canonical", "example.com", "example.com"},
|
||||||
{"case insensitive", "Windscribe.Com.", "windscribe.com"},
|
{"case insensitive", "Example.Com.", "example.com"},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tc := range tests {
|
for _, tc := range tests {
|
||||||
@@ -67,6 +75,7 @@ func Test_canonicalName(t *testing.T) {
|
|||||||
|
|
||||||
func Test_prog_upstreamFor(t *testing.T) {
|
func Test_prog_upstreamFor(t *testing.T) {
|
||||||
cfg := testhelper.SampleConfig(t)
|
cfg := testhelper.SampleConfig(t)
|
||||||
|
cfg.Service.LeakOnUpstreamFailure = func(v bool) *bool { return &v }(false)
|
||||||
p := &prog{cfg: cfg}
|
p := &prog{cfg: cfg}
|
||||||
p.um = newUpstreamMonitor(p.cfg)
|
p.um = newUpstreamMonitor(p.cfg)
|
||||||
p.lanLoopGuard = newLoopGuard()
|
p.lanLoopGuard = newLoopGuard()
|
||||||
@@ -357,6 +366,9 @@ func Test_isLanHostnameQuery(t *testing.T) {
|
|||||||
{"A not LAN", newDnsMsgWithHostname("example.com", dns.TypeA), false},
|
{"A not LAN", newDnsMsgWithHostname("example.com", dns.TypeA), false},
|
||||||
{"AAAA not LAN", newDnsMsgWithHostname("example.com", dns.TypeAAAA), false},
|
{"AAAA not LAN", newDnsMsgWithHostname("example.com", dns.TypeAAAA), false},
|
||||||
{"Not A or AAAA", newDnsMsgWithHostname("foo", dns.TypeTXT), 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 {
|
for _, tc := range tests {
|
||||||
tc := tc
|
tc := tc
|
||||||
@@ -406,6 +418,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) {
|
func Test_isWanClient(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
@@ -431,3 +464,13 @@ 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")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,14 @@
|
|||||||
|
package cli
|
||||||
|
|
||||||
|
import "regexp"
|
||||||
|
|
||||||
|
// validHostname reports whether hostname is a valid hostname.
|
||||||
|
// A valid hostname contains 3 -> 64 characters and conform to RFC1123.
|
||||||
|
func validHostname(hostname string) bool {
|
||||||
|
hostnameLen := len(hostname)
|
||||||
|
if hostnameLen < 3 || hostnameLen > 64 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
validHostnameRfc1123 := regexp.MustCompile(`^(([a-zA-Z0-9]|[a-zA-Z0-9][a-zA-Z0-9\-]*[a-zA-Z0-9])\.)*([A-Za-z0-9]|[A-Za-z0-9][A-Za-z0-9\-]*[A-Za-z0-9])$`)
|
||||||
|
return validHostnameRfc1123.MatchString(hostname)
|
||||||
|
}
|
||||||
@@ -0,0 +1,35 @@
|
|||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
)
|
||||||
|
|
||||||
|
func Test_validHostname(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
hostname string
|
||||||
|
valid bool
|
||||||
|
}{
|
||||||
|
{"localhost", "localhost", true},
|
||||||
|
{"localdomain", "localhost.localdomain", true},
|
||||||
|
{"localhost6", "localhost6.localdomain6", true},
|
||||||
|
{"ip6", "ip6-localhost", true},
|
||||||
|
{"non-domain", "controld", true},
|
||||||
|
{"domain", "controld.com", true},
|
||||||
|
{"empty", "", false},
|
||||||
|
{"min length", "fo", false},
|
||||||
|
{"max length", strings.Repeat("a", 65), false},
|
||||||
|
{"special char", "foo!", false},
|
||||||
|
{"non-ascii", "fooΩ", false},
|
||||||
|
}
|
||||||
|
for _, tc := range tests {
|
||||||
|
tc := tc
|
||||||
|
t.Run(tc.hostname, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
assert.True(t, validHostname(tc.hostname) == tc.valid)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
+81
-5
@@ -1,5 +1,12 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
// AppCallback provides hooks for injecting certain functionalities
|
// AppCallback provides hooks for injecting certain functionalities
|
||||||
// from mobile platforms to main ctrld cli.
|
// from mobile platforms to main ctrld cli.
|
||||||
type AppCallback struct {
|
type AppCallback struct {
|
||||||
@@ -11,9 +18,78 @@ type AppCallback struct {
|
|||||||
|
|
||||||
// AppConfig allows overwriting ctrld cli flags from mobile platforms.
|
// AppConfig allows overwriting ctrld cli flags from mobile platforms.
|
||||||
type AppConfig struct {
|
type AppConfig struct {
|
||||||
CdUID string
|
CdUID string
|
||||||
HomeDir string
|
ProvisionID string
|
||||||
UpstreamProto string
|
CustomHostname string
|
||||||
Verbose int
|
HomeDir string
|
||||||
LogPath 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) {
|
||||||
|
var lastErr error
|
||||||
|
client := httpClientWithFallback(defaultHTTPTimeout)
|
||||||
|
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
|
||||||
|
}
|
||||||
|
if ipReq != nil {
|
||||||
|
mainLog.Load().Warn().Err(err).Msgf("dial to %q failed", req.Host)
|
||||||
|
mainLog.Load().Warn().Msgf("fallback to direct IP to download prod version: %q", ip)
|
||||||
|
resp, err = client.Do(ipReq)
|
||||||
|
if err == nil {
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
lastErr = err
|
||||||
|
mainLog.Load().Debug().Err(err).
|
||||||
|
Str("method", req.Method).
|
||||||
|
Str("url", req.URL.String()).
|
||||||
|
Msgf("HTTP request attempt %d/%d failed", attempt+1, maxRetries)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("failed after %d attempts to %s %s: %v", maxRetries, req.Method, req.URL, lastErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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,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,467 @@
|
|||||||
|
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 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,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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -105,6 +105,10 @@ func (p *prog) checkDnsLoop() {
|
|||||||
for uid := range p.loop {
|
for uid := range p.loop {
|
||||||
msg := loopTestMsg(uid)
|
msg := loopTestMsg(uid)
|
||||||
uc := upstream[uid]
|
uc := upstream[uid]
|
||||||
|
// Skipping upstream which is being marked as down.
|
||||||
|
if uc == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
resolver, err := ctrld.NewResolver(uc)
|
resolver, err := ctrld.NewResolver(uc)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
mainLog.Load().Warn().Err(err).Msgf("could not perform loop check for upstream: %q, endpoint: %q", uc.Name, uc.Endpoint)
|
mainLog.Load().Warn().Err(err).Msgf("could not perform loop check for upstream: %q, endpoint: %q", uc.Name, uc.Endpoint)
|
||||||
|
|||||||
+69
-13
@@ -1,7 +1,9 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"encoding/hex"
|
||||||
"io"
|
"io"
|
||||||
|
"net"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
@@ -29,11 +31,20 @@ var (
|
|||||||
silent bool
|
silent bool
|
||||||
cdUID string
|
cdUID string
|
||||||
cdOrg string
|
cdOrg string
|
||||||
|
customHostname string
|
||||||
cdDev bool
|
cdDev bool
|
||||||
iface string
|
iface string
|
||||||
ifaceStartStop string
|
ifaceStartStop string
|
||||||
nextdns string
|
nextdns string
|
||||||
cdUpstreamProto string
|
cdUpstreamProto string
|
||||||
|
deactivationPin int64
|
||||||
|
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]
|
mainLog atomic.Pointer[zerolog.Logger]
|
||||||
consoleWriter zerolog.ConsoleWriter
|
consoleWriter zerolog.ConsoleWriter
|
||||||
@@ -41,9 +52,10 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
cdUidFlagName = "cd"
|
cdUidFlagName = "cd"
|
||||||
cdOrgFlagName = "cd-org"
|
cdOrgFlagName = "cd-org"
|
||||||
nextdnsFlagName = "nextdns"
|
customHostnameFlagName = "custom-hostname"
|
||||||
|
nextdnsFlagName = "nextdns"
|
||||||
)
|
)
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
@@ -52,6 +64,16 @@ func init() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func Main() {
|
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")
|
ctrld.InitConfig(v, "ctrld")
|
||||||
initCLI()
|
initCLI()
|
||||||
if err := rootCmd.Execute(); err != nil {
|
if err := rootCmd.Execute(); err != nil {
|
||||||
@@ -82,22 +104,33 @@ func initConsoleLogging() {
|
|||||||
multi := zerolog.MultiLevelWriter(consoleWriter)
|
multi := zerolog.MultiLevelWriter(consoleWriter)
|
||||||
l := mainLog.Load().Output(multi).With().Timestamp().Logger()
|
l := mainLog.Load().Output(multi).With().Timestamp().Logger()
|
||||||
mainLog.Store(&l)
|
mainLog.Store(&l)
|
||||||
|
|
||||||
switch {
|
switch {
|
||||||
case silent:
|
case silent:
|
||||||
zerolog.SetGlobalLevel(zerolog.NoLevel)
|
zerolog.SetGlobalLevel(zerolog.NoLevel)
|
||||||
case verbose == 1:
|
case verbose == 1:
|
||||||
|
ctrld.ProxyLogger.Store(&l)
|
||||||
zerolog.SetGlobalLevel(zerolog.InfoLevel)
|
zerolog.SetGlobalLevel(zerolog.InfoLevel)
|
||||||
case verbose > 1:
|
case verbose > 1:
|
||||||
|
ctrld.ProxyLogger.Store(&l)
|
||||||
zerolog.SetGlobalLevel(zerolog.DebugLevel)
|
zerolog.SetGlobalLevel(zerolog.DebugLevel)
|
||||||
default:
|
default:
|
||||||
zerolog.SetGlobalLevel(zerolog.NoticeLevel)
|
zerolog.SetGlobalLevel(zerolog.NoticeLevel)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// initLogging initializes global logging setup.
|
// initInteractiveLogging is like initLogging, but the ProxyLogger is discarded
|
||||||
func initLogging() {
|
// 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"
|
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.
|
// initLoggingWithBackup initializes log setup base on current config.
|
||||||
@@ -106,8 +139,8 @@ func initLogging() {
|
|||||||
// This is only used in runCmd for special handling in case of logging config
|
// 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
|
// change in cd mode. Without special reason, the caller should use initLogging
|
||||||
// wrapper instead of calling this function directly.
|
// wrapper instead of calling this function directly.
|
||||||
func initLoggingWithBackup(doBackup bool) {
|
func initLoggingWithBackup(doBackup bool) []io.Writer {
|
||||||
writers := []io.Writer{io.Discard}
|
var writers []io.Writer
|
||||||
if logFilePath := normalizeLogFilePath(cfg.Service.LogPath); logFilePath != "" {
|
if logFilePath := normalizeLogFilePath(cfg.Service.LogPath); logFilePath != "" {
|
||||||
// Create parent directory if necessary.
|
// Create parent directory if necessary.
|
||||||
if err := os.MkdirAll(filepath.Dir(logFilePath), 0750); err != nil {
|
if err := os.MkdirAll(filepath.Dir(logFilePath), 0750); err != nil {
|
||||||
@@ -119,14 +152,14 @@ func initLoggingWithBackup(doBackup bool) {
|
|||||||
flags := os.O_CREATE | os.O_RDWR | os.O_APPEND
|
flags := os.O_CREATE | os.O_RDWR | os.O_APPEND
|
||||||
if doBackup {
|
if doBackup {
|
||||||
// Backup old log file with .1 suffix.
|
// Backup old log file with .1 suffix.
|
||||||
if err := os.Rename(logFilePath, logFilePath+".1"); err != nil && !os.IsNotExist(err) {
|
if err := os.Rename(logFilePath, logFilePath+oldLogSuffix); err != nil && !os.IsNotExist(err) {
|
||||||
mainLog.Load().Error().Msgf("could not backup old log file: %v", err)
|
mainLog.Load().Error().Msgf("could not backup old log file: %v", err)
|
||||||
} else {
|
} else {
|
||||||
// Backup was created, set flags for truncating old log file.
|
// Backup was created, set flags for truncating old log file.
|
||||||
flags = os.O_CREATE | os.O_RDWR
|
flags = os.O_CREATE | os.O_RDWR
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
logFile, err := os.OpenFile(logFilePath, flags, os.FileMode(0o600))
|
logFile, err := openLogFile(logFilePath, flags)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
mainLog.Load().Error().Msgf("failed to create log file: %v", err)
|
mainLog.Load().Error().Msgf("failed to create log file: %v", err)
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
@@ -145,21 +178,22 @@ func initLoggingWithBackup(doBackup bool) {
|
|||||||
switch {
|
switch {
|
||||||
case silent:
|
case silent:
|
||||||
zerolog.SetGlobalLevel(zerolog.NoLevel)
|
zerolog.SetGlobalLevel(zerolog.NoLevel)
|
||||||
return
|
return writers
|
||||||
case verbose == 1:
|
case verbose == 1:
|
||||||
logLevel = "info"
|
logLevel = "info"
|
||||||
case verbose > 1:
|
case verbose > 1:
|
||||||
logLevel = "debug"
|
logLevel = "debug"
|
||||||
}
|
}
|
||||||
if logLevel == "" {
|
if logLevel == "" {
|
||||||
return
|
return writers
|
||||||
}
|
}
|
||||||
level, err := zerolog.ParseLevel(logLevel)
|
level, err := zerolog.ParseLevel(logLevel)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
mainLog.Load().Warn().Err(err).Msg("could not set log level")
|
mainLog.Load().Warn().Err(err).Msg("could not set log level")
|
||||||
return
|
return writers
|
||||||
}
|
}
|
||||||
zerolog.SetGlobalLevel(level)
|
zerolog.SetGlobalLevel(level)
|
||||||
|
return writers
|
||||||
}
|
}
|
||||||
|
|
||||||
func initCache() {
|
func initCache() {
|
||||||
@@ -170,3 +204,25 @@ func initCache() {
|
|||||||
cfg.Service.CacheSize = 4096
|
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)
|
||||||
|
}
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package cli
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"os"
|
"os"
|
||||||
|
"os/exec"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -13,5 +14,20 @@ var logOutput strings.Builder
|
|||||||
func TestMain(m *testing.M) {
|
func TestMain(m *testing.M) {
|
||||||
l := zerolog.New(&logOutput)
|
l := zerolog.New(&logOutput)
|
||||||
mainLog.Store(&l)
|
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())
|
os.Exit(m.Run())
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-1
@@ -107,7 +107,7 @@ func (p *prog) runMetricsServer(ctx context.Context, reloadCh chan struct{}) {
|
|||||||
|
|
||||||
reg := prometheus.NewRegistry()
|
reg := prometheus.NewRegistry()
|
||||||
// Register queries count stats if enabled.
|
// Register queries count stats if enabled.
|
||||||
if cfg.Service.MetricsQueryStats {
|
if p.metricsQueryStats.Load() {
|
||||||
reg.MustRegister(statsQueriesCount)
|
reg.MustRegister(statsQueriesCount)
|
||||||
reg.MustRegister(statsClientQueriesCount)
|
reg.MustRegister(statsClientQueriesCount)
|
||||||
}
|
}
|
||||||
|
|||||||
+36
-4
@@ -9,17 +9,18 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
)
|
)
|
||||||
|
|
||||||
func patchNetIfaceName(iface *net.Interface) error {
|
func patchNetIfaceName(iface *net.Interface) (bool, error) {
|
||||||
b, err := exec.Command("networksetup", "-listnetworkserviceorder").Output()
|
b, err := exec.Command("networksetup", "-listnetworkserviceorder").Output()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return false, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
patched := false
|
||||||
if name := networkServiceName(iface.Name, bytes.NewReader(b)); name != "" {
|
if name := networkServiceName(iface.Name, bytes.NewReader(b)); name != "" {
|
||||||
|
patched = true
|
||||||
iface.Name = name
|
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 {
|
func networkServiceName(ifaceName string, r io.Reader) string {
|
||||||
@@ -42,3 +43,34 @@ func networkServiceName(ifaceName string, r io.Reader) string {
|
|||||||
}
|
}
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// validInterface reports whether the *net.Interface is a valid one.
|
||||||
|
func validInterface(iface *net.Interface, validIfacesMap map[string]struct{}) bool {
|
||||||
|
_, ok := validIfacesMap[iface.Name]
|
||||||
|
return ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// validInterfacesMap returns a set of all valid hardware ports.
|
||||||
|
func validInterfacesMap() map[string]struct{} {
|
||||||
|
b, err := exec.Command("networksetup", "-listallhardwareports").Output()
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return parseListAllHardwarePorts(bytes.NewReader(b))
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseListAllHardwarePorts parses output of "networksetup -listallhardwareports"
|
||||||
|
// and returns map presents all hardware ports.
|
||||||
|
func parseListAllHardwarePorts(r io.Reader) map[string]struct{} {
|
||||||
|
m := make(map[string]struct{})
|
||||||
|
scanner := bufio.NewScanner(r)
|
||||||
|
for scanner.Scan() {
|
||||||
|
line := scanner.Text()
|
||||||
|
after, ok := strings.CutPrefix(line, "Device: ")
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
m[after] = struct{}{}
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
+18
-3
@@ -1,7 +1,22 @@
|
|||||||
//go:build !darwin
|
//go:build !darwin && !windows && !linux
|
||||||
|
|
||||||
package cli
|
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 }
|
||||||
|
|
||||||
|
// 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: {}}
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,93 @@
|
|||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
"log"
|
||||||
|
"net"
|
||||||
|
"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) (bool, error) {
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// validInterface reports whether the *net.Interface is a valid one.
|
||||||
|
// On Windows, only physical interfaces are considered valid.
|
||||||
|
func validInterface(iface *net.Interface, validIfacesMap map[string]struct{}) bool {
|
||||||
|
_, ok := validIfacesMap[iface.Name]
|
||||||
|
return ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// validInterfacesMap returns a set of all physical interfaces.
|
||||||
|
func validInterfacesMap() map[string]struct{} {
|
||||||
|
m := make(map[string]struct{})
|
||||||
|
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,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)
|
||||||
|
}
|
||||||
|
}
|
||||||
+63
-8
@@ -1,8 +1,12 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bufio"
|
||||||
|
"bytes"
|
||||||
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
|
"strings"
|
||||||
|
|
||||||
"github.com/Control-D-Inc/ctrld/internal/resolvconffile"
|
"github.com/Control-D-Inc/ctrld/internal/resolvconffile"
|
||||||
)
|
)
|
||||||
@@ -27,16 +31,41 @@ func deAllocateIP(ip string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// setDnsIgnoreUnusableInterface likes setDNS, but return a nil error if the interface is not usable.
|
||||||
|
func setDnsIgnoreUnusableInterface(iface *net.Interface, nameservers []string) error {
|
||||||
|
if err := setDNS(iface, nameservers); err != nil {
|
||||||
|
// TODO: investiate whether we can detect this without relying on error message.
|
||||||
|
if strings.Contains(err.Error(), " is not a recognized network service") {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// set the dns server for the provided network interface
|
// set the dns server for the provided network interface
|
||||||
// networksetup -setdnsservers Wi-Fi 8.8.8.8 1.1.1.1
|
// networksetup -setdnsservers Wi-Fi 8.8.8.8 1.1.1.1
|
||||||
// TODO(cuonglm): use system API
|
// TODO(cuonglm): use system API
|
||||||
func setDNS(iface *net.Interface, nameservers []string) error {
|
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"
|
cmd := "networksetup"
|
||||||
args := []string{"-setdnsservers", iface.Name}
|
args := []string{"-setdnsservers", iface.Name}
|
||||||
args = append(args, nameservers...)
|
args = append(args, nameservers...)
|
||||||
|
if out, err := exec.Command(cmd, args...).CombinedOutput(); err != nil {
|
||||||
|
return fmt.Errorf("%v: %w", string(out), err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
if err := exec.Command(cmd, args...).Run(); err != nil {
|
// resetDnsIgnoreUnusableInterface likes resetDNS, but return a nil error if the interface is not usable.
|
||||||
mainLog.Load().Error().Err(err).Msgf("setDNS failed, ips = %q", nameservers)
|
func resetDnsIgnoreUnusableInterface(iface *net.Interface) error {
|
||||||
|
if err := resetDNS(iface); err != nil {
|
||||||
|
// TODO: investiate whether we can detect this without relying on error message.
|
||||||
|
if strings.Contains(err.Error(), " is not a recognized network service") {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
@@ -46,14 +75,40 @@ func setDNS(iface *net.Interface, nameservers []string) error {
|
|||||||
func resetDNS(iface *net.Interface) error {
|
func resetDNS(iface *net.Interface) error {
|
||||||
cmd := "networksetup"
|
cmd := "networksetup"
|
||||||
args := []string{"-setdnsservers", iface.Name, "empty"}
|
args := []string{"-setdnsservers", iface.Name, "empty"}
|
||||||
|
if out, err := exec.Command(cmd, args...).CombinedOutput(); err != nil {
|
||||||
if err := exec.Command(cmd, args...).Run(); err != nil {
|
return fmt.Errorf("%v: %w", string(out), err)
|
||||||
mainLog.Load().Error().Err(err).Msgf("resetDNS failed")
|
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func currentDNS(_ *net.Interface) []string {
|
// restoreDNS restores the DNS settings of the given interface.
|
||||||
return resolvconffile.NameServers("")
|
// 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()
|
||||||
|
}
|
||||||
|
|
||||||
|
// currentStaticDNS returns the current static DNS settings of given interface.
|
||||||
|
func currentStaticDNS(iface *net.Interface) ([]string, error) {
|
||||||
|
cmd := "networksetup"
|
||||||
|
args := []string{"-getdnsservers", iface.Name}
|
||||||
|
out, err := exec.Command(cmd, args...).Output()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
scanner := bufio.NewScanner(bytes.NewReader(out))
|
||||||
|
var ns []string
|
||||||
|
for scanner.Scan() {
|
||||||
|
line := scanner.Text()
|
||||||
|
if ip := net.ParseIP(line); ip != nil {
|
||||||
|
ns = append(ns, ip.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ns, nil
|
||||||
}
|
}
|
||||||
|
|||||||
+40
-5
@@ -5,6 +5,10 @@ import (
|
|||||||
"net/netip"
|
"net/netip"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
|
|
||||||
|
"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/dns"
|
||||||
"github.com/Control-D-Inc/ctrld/internal/resolvconffile"
|
"github.com/Control-D-Inc/ctrld/internal/resolvconffile"
|
||||||
)
|
)
|
||||||
@@ -29,9 +33,14 @@ func deAllocateIP(ip string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// setDnsIgnoreUnusableInterface likes setDNS, but return a nil error if the interface is not usable.
|
||||||
|
func setDnsIgnoreUnusableInterface(iface *net.Interface, nameservers []string) error {
|
||||||
|
return setDNS(iface, nameservers)
|
||||||
|
}
|
||||||
|
|
||||||
// set the dns server for the provided network interface
|
// set the dns server for the provided network interface
|
||||||
func setDNS(iface *net.Interface, nameservers []string) error {
|
func setDNS(iface *net.Interface, nameservers []string) error {
|
||||||
r, err := dns.NewOSConfigurator(logf, iface.Name)
|
r, err := dns.NewOSConfigurator(logf, &health.Tracker{}, &controlknobs.Knobs{}, iface.Name)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
mainLog.Load().Error().Err(err).Msg("failed to create DNS OS configurator")
|
mainLog.Load().Error().Err(err).Msg("failed to create DNS OS configurator")
|
||||||
return err
|
return err
|
||||||
@@ -42,15 +51,30 @@ func setDNS(iface *net.Interface, nameservers []string) error {
|
|||||||
ns = append(ns, netip.MustParseAddr(nameserver))
|
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")
|
mainLog.Load().Error().Err(err).Msg("failed to set DNS")
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// resetDnsIgnoreUnusableInterface likes resetDNS, but return a nil error if the interface is not usable.
|
||||||
|
func resetDnsIgnoreUnusableInterface(iface *net.Interface) error {
|
||||||
|
return resetDNS(iface)
|
||||||
|
}
|
||||||
|
|
||||||
func resetDNS(iface *net.Interface) error {
|
func resetDNS(iface *net.Interface) error {
|
||||||
r, err := dns.NewOSConfigurator(logf, iface.Name)
|
r, err := dns.NewOSConfigurator(logf, &health.Tracker{}, &controlknobs.Knobs{}, iface.Name)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
mainLog.Load().Error().Err(err).Msg("failed to create DNS OS configurator")
|
mainLog.Load().Error().Err(err).Msg("failed to create DNS OS configurator")
|
||||||
return err
|
return err
|
||||||
@@ -63,6 +87,17 @@ func resetDNS(iface *net.Interface) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func currentDNS(_ *net.Interface) []string {
|
// restoreDNS restores the DNS settings of the given interface.
|
||||||
return resolvconffile.NameServers("")
|
// 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()
|
||||||
|
}
|
||||||
|
|
||||||
|
// currentStaticDNS returns the current static DNS settings of given interface.
|
||||||
|
func currentStaticDNS(iface *net.Interface) ([]string, error) {
|
||||||
|
return currentDNS(iface), nil
|
||||||
}
|
}
|
||||||
|
|||||||
+61
-114
@@ -9,15 +9,16 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
"path/filepath"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"syscall"
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/fsnotify/fsnotify"
|
|
||||||
"github.com/insomniacslk/dhcp/dhcpv4/nclient4"
|
"github.com/insomniacslk/dhcp/dhcpv4/nclient4"
|
||||||
"github.com/insomniacslk/dhcp/dhcpv6"
|
"github.com/insomniacslk/dhcp/dhcpv6"
|
||||||
"github.com/insomniacslk/dhcp/dhcpv6/client6"
|
"github.com/insomniacslk/dhcp/dhcpv6/client6"
|
||||||
|
"tailscale.com/control/controlknobs"
|
||||||
|
"tailscale.com/health"
|
||||||
"tailscale.com/util/dnsname"
|
"tailscale.com/util/dnsname"
|
||||||
|
|
||||||
"github.com/Control-D-Inc/ctrld/internal/dns"
|
"github.com/Control-D-Inc/ctrld/internal/dns"
|
||||||
@@ -25,10 +26,7 @@ import (
|
|||||||
"github.com/Control-D-Inc/ctrld/internal/resolvconffile"
|
"github.com/Control-D-Inc/ctrld/internal/resolvconffile"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const resolvConfBackupFailedMsg = "open /etc/resolv.pre-ctrld-backup.conf: read-only file system"
|
||||||
resolvConfPath = "/etc/resolv.conf"
|
|
||||||
resolvConfBackupFailedMsg = "open /etc/resolv.pre-ctrld-backup.conf: read-only file system"
|
|
||||||
)
|
|
||||||
|
|
||||||
// allocate loopback ip
|
// allocate loopback ip
|
||||||
// sudo ip a add 127.0.0.2/24 dev lo
|
// sudo ip a add 127.0.0.2/24 dev lo
|
||||||
@@ -52,9 +50,13 @@ func deAllocateIP(ip string) error {
|
|||||||
|
|
||||||
const maxSetDNSAttempts = 5
|
const maxSetDNSAttempts = 5
|
||||||
|
|
||||||
// set the dns server for the provided network interface
|
// setDnsIgnoreUnusableInterface likes setDNS, but return a nil error if the interface is not usable.
|
||||||
|
func setDnsIgnoreUnusableInterface(iface *net.Interface, nameservers []string) error {
|
||||||
|
return setDNS(iface, nameservers)
|
||||||
|
}
|
||||||
|
|
||||||
func setDNS(iface *net.Interface, nameservers []string) error {
|
func setDNS(iface *net.Interface, nameservers []string) error {
|
||||||
r, err := dns.NewOSConfigurator(logf, iface.Name)
|
r, err := dns.NewOSConfigurator(logf, &health.Tracker{}, &controlknobs.Knobs{}, iface.Name)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
mainLog.Load().Error().Err(err).Msg("failed to create DNS OS configurator")
|
mainLog.Load().Error().Err(err).Msg("failed to create DNS OS configurator")
|
||||||
return err
|
return err
|
||||||
@@ -69,41 +71,39 @@ func setDNS(iface *net.Interface, nameservers []string) error {
|
|||||||
Nameservers: ns,
|
Nameservers: ns,
|
||||||
SearchDomains: []dnsname.FQDN{},
|
SearchDomains: []dnsname.FQDN{},
|
||||||
}
|
}
|
||||||
defer func() {
|
if sds, err := searchDomains(); err == nil {
|
||||||
if r.Mode() == "direct" {
|
// Filter the root domain, since it's not allowed by systemd.
|
||||||
go watchResolveConf(osConfig)
|
// 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
|
trySystemdResolve := false
|
||||||
for i := 0; i < maxSetDNSAttempts; i++ {
|
if err := r.SetDNS(osConfig); err != nil {
|
||||||
if err := r.SetDNS(osConfig); err != nil {
|
if strings.Contains(err.Error(), "Rejected send message") &&
|
||||||
if strings.Contains(err.Error(), "Rejected send message") &&
|
strings.Contains(err.Error(), "org.freedesktop.network1.Manager") {
|
||||||
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")
|
||||||
mainLog.Load().Warn().Msg("Interfaces are managed by systemd-networkd, switch to systemd-resolve for setting DNS")
|
trySystemdResolve = true
|
||||||
trySystemdResolve = true
|
goto systemdResolve
|
||||||
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 useSystemdResolved {
|
// This error happens on read-only file system, which causes ctrld failed to create backup
|
||||||
if out, err := exec.Command("systemctl", "restart", "systemd-resolved").CombinedOutput(); err != nil {
|
// for /etc/resolv.conf file. It is ok, because the DNS is still set anyway, and restore
|
||||||
mainLog.Load().Warn().Err(err).Msgf("could not restart systemd-resolved: %s", string(out))
|
// 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
|
||||||
currentNS := currentDNS(iface)
|
if r.Mode() == "direct" && strings.Contains(err.Error(), resolvConfBackupFailedMsg) {
|
||||||
if isSubSet(nameservers, currentNS) {
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
systemdResolve:
|
||||||
if trySystemdResolve {
|
if trySystemdResolve {
|
||||||
// Stop systemd-networkd and retry setting DNS.
|
// Stop systemd-networkd and retry setting DNS.
|
||||||
if out, err := exec.Command("systemctl", "stop", "systemd-networkd").CombinedOutput(); err != nil {
|
if out, err := exec.Command("systemctl", "stop", "systemd-networkd").CombinedOutput(); err != nil {
|
||||||
@@ -123,11 +123,16 @@ func setDNS(iface *net.Interface, nameservers []string) error {
|
|||||||
}
|
}
|
||||||
time.Sleep(time.Second)
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// resetDnsIgnoreUnusableInterface likes resetDNS, but return a nil error if the interface is not usable.
|
||||||
|
func resetDnsIgnoreUnusableInterface(iface *net.Interface) error {
|
||||||
|
return resetDNS(iface)
|
||||||
|
}
|
||||||
|
|
||||||
func resetDNS(iface *net.Interface) (err error) {
|
func resetDNS(iface *net.Interface) (err error) {
|
||||||
defer func() {
|
defer func() {
|
||||||
if err == nil {
|
if err == nil {
|
||||||
@@ -137,7 +142,7 @@ func resetDNS(iface *net.Interface) (err error) {
|
|||||||
if exe, _ := exec.LookPath("/lib/systemd/systemd-networkd"); exe != "" {
|
if exe, _ := exec.LookPath("/lib/systemd/systemd-networkd"); exe != "" {
|
||||||
_ = exec.Command("systemctl", "start", "systemd-networkd").Run()
|
_ = exec.Command("systemctl", "start", "systemd-networkd").Run()
|
||||||
}
|
}
|
||||||
if r, oerr := dns.NewOSConfigurator(logf, iface.Name); oerr == nil {
|
if r, oerr := dns.NewOSConfigurator(logf, &health.Tracker{}, &controlknobs.Knobs{}, iface.Name); oerr == nil {
|
||||||
_ = r.SetDNS(dns.OSConfig{})
|
_ = r.SetDNS(dns.OSConfig{})
|
||||||
if err := r.Close(); err != nil {
|
if err := r.Close(); err != nil {
|
||||||
mainLog.Load().Error().Err(err).Msg("failed to rollback DNS setting")
|
mainLog.Load().Error().Err(err).Msg("failed to rollback DNS setting")
|
||||||
@@ -168,6 +173,7 @@ func resetDNS(iface *net.Interface) (err error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// TODO(cuonglm): handle DHCPv6 properly.
|
// TODO(cuonglm): handle DHCPv6 properly.
|
||||||
|
mainLog.Load().Debug().Msg("checking for IPv6 availability")
|
||||||
if ctrldnet.IPv6Available(ctx) {
|
if ctrldnet.IPv6Available(ctx) {
|
||||||
c := client6.NewClient()
|
c := client6.NewClient()
|
||||||
conversation, err := c.Exchange(iface.Name)
|
conversation, err := c.Exchange(iface.Name)
|
||||||
@@ -187,6 +193,8 @@ func resetDNS(iface *net.Interface) (err error) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
} else {
|
||||||
|
mainLog.Load().Debug().Msg("IPv6 is not available")
|
||||||
}
|
}
|
||||||
|
|
||||||
return ignoringEINTR(func() error {
|
return ignoringEINTR(func() error {
|
||||||
@@ -194,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 {
|
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 {
|
if ns := fn(iface.Name); len(ns) > 0 {
|
||||||
return ns
|
return ns
|
||||||
}
|
}
|
||||||
@@ -203,6 +218,11 @@ func currentDNS(iface *net.Interface) []string {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// currentStaticDNS returns the current static DNS settings of given interface.
|
||||||
|
func currentStaticDNS(iface *net.Interface) ([]string, error) {
|
||||||
|
return currentDNS(iface), nil
|
||||||
|
}
|
||||||
|
|
||||||
func getDNSByResolvectl(iface string) []string {
|
func getDNSByResolvectl(iface string) []string {
|
||||||
b, err := exec.Command("resolvectl", "dns", "-i", iface).Output()
|
b, err := exec.Command("resolvectl", "dns", "-i", iface).Output()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -284,8 +304,7 @@ func ignoringEINTR(fn func() error) error {
|
|||||||
func isSubSet(s1, s2 []string) bool {
|
func isSubSet(s1, s2 []string) bool {
|
||||||
ok := true
|
ok := true
|
||||||
for _, ns := range s1 {
|
for _, ns := range s1 {
|
||||||
// TODO(cuonglm): use slices.Contains once upgrading to go1.21
|
if slices.Contains(s2, ns) {
|
||||||
if sliceContains(s2, ns) {
|
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
ok = false
|
ok = false
|
||||||
@@ -293,75 +312,3 @@ func isSubSet(s1, s2 []string) bool {
|
|||||||
}
|
}
|
||||||
return ok
|
return ok
|
||||||
}
|
}
|
||||||
|
|
||||||
// sliceContains reports whether v is present in s.
|
|
||||||
func sliceContains[S ~[]E, E comparable](s S, v E) bool {
|
|
||||||
return sliceIndex(s, v) >= 0
|
|
||||||
}
|
|
||||||
|
|
||||||
// sliceIndex returns the index of the first occurrence of v in s,
|
|
||||||
// or -1 if not present.
|
|
||||||
func sliceIndex[S ~[]E, E comparable](s S, v E) int {
|
|
||||||
for i := range s {
|
|
||||||
if v == s[i] {
|
|
||||||
return i
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
|
|
||||||
// watchResolveConf watches any changes to /etc/resolv.conf file,
|
|
||||||
// and reverting to the original config set by ctrld.
|
|
||||||
func watchResolveConf(oc dns.OSConfig) {
|
|
||||||
mainLog.Load().Debug().Msg("start watching /etc/resolv.conf file")
|
|
||||||
watcher, err := fsnotify.NewWatcher()
|
|
||||||
if err != nil {
|
|
||||||
mainLog.Load().Warn().Err(err).Msg("could not create watcher for /etc/resolv.conf")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// We watch /etc instead of /etc/resolv.conf directly,
|
|
||||||
// see: https://github.com/fsnotify/fsnotify#watching-a-file-doesnt-work-well
|
|
||||||
watchDir := filepath.Dir(resolvConfPath)
|
|
||||||
if err := watcher.Add(watchDir); err != nil {
|
|
||||||
mainLog.Load().Warn().Err(err).Msg("could not add /etc/resolv.conf to watcher list")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
r, err := dns.NewOSConfigurator(func(format string, args ...any) {}, "lo") // interface name does not matter.
|
|
||||||
if err != nil {
|
|
||||||
mainLog.Load().Error().Err(err).Msg("failed to create DNS OS configurator")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case event, ok := <-watcher.Events:
|
|
||||||
if !ok {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if event.Name != resolvConfPath { // skip if not /etc/resolv.conf changes.
|
|
||||||
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
|
|
||||||
}
|
|
||||||
if err := r.SetDNS(oc); 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:
|
|
||||||
if !ok {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
mainLog.Load().Err(err).Msg("could not get event for /etc/resolv.conf")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
+288
-51
@@ -1,79 +1,207 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
"errors"
|
"errors"
|
||||||
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
"strconv"
|
"slices"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
"golang.org/x/sys/windows/registry"
|
||||||
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
||||||
|
|
||||||
ctrldnet "github.com/Control-D-Inc/ctrld/internal/net"
|
ctrldnet "github.com/Control-D-Inc/ctrld/internal/net"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
v4InterfaceKeyPathFormat = `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces\`
|
||||||
|
v6InterfaceKeyPathFormat = `SYSTEM\CurrentControlSet\Services\Tcpip6\Parameters\Interfaces\`
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
setDNSOnce sync.Once
|
||||||
|
resetDNSOnce sync.Once
|
||||||
|
)
|
||||||
|
|
||||||
|
// setDnsIgnoreUnusableInterface likes setDNS, but return a nil error if the interface is not usable.
|
||||||
|
func setDnsIgnoreUnusableInterface(iface *net.Interface, nameservers []string) error {
|
||||||
|
return setDNS(iface, nameservers)
|
||||||
|
}
|
||||||
|
|
||||||
|
// setDNS sets the dns server for the provided network interface
|
||||||
func setDNS(iface *net.Interface, nameservers []string) error {
|
func setDNS(iface *net.Interface, nameservers []string) error {
|
||||||
if len(nameservers) == 0 {
|
if len(nameservers) == 0 {
|
||||||
return errors.New("empty DNS nameservers")
|
return errors.New("empty DNS nameservers")
|
||||||
}
|
}
|
||||||
primaryDNS := nameservers[0]
|
setDNSOnce.Do(func() {
|
||||||
if err := setPrimaryDNS(iface, primaryDNS); err != nil {
|
// If there's a Dns server running, that means we are on AD with Dns feature enabled.
|
||||||
return err
|
// Configuring the Dns server to forward queries to ctrld instead.
|
||||||
|
if hasLocalDnsServerRunning() {
|
||||||
|
mainLog.Load().Debug().Msg("Local DNS server detected, configuring forwarders")
|
||||||
|
|
||||||
|
file := absHomeDir(windowsForwardersFilename)
|
||||||
|
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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
luid, err := winipcfg.LUIDFromIndex(uint32(iface.Index))
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("setDNS: %w", err)
|
||||||
}
|
}
|
||||||
if len(nameservers) > 1 {
|
var (
|
||||||
secondaryDNS := nameservers[1]
|
serversV4 []netip.Addr
|
||||||
_ = addSecondaryDNS(iface, secondaryDNS)
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// resetDnsIgnoreUnusableInterface likes resetDNS, but return a nil error if the interface is not usable.
|
||||||
|
func resetDnsIgnoreUnusableInterface(iface *net.Interface) error {
|
||||||
|
return resetDNS(iface)
|
||||||
|
}
|
||||||
|
|
||||||
// TODO(cuonglm): should we use system API?
|
// TODO(cuonglm): should we use system API?
|
||||||
func resetDNS(iface *net.Interface) error {
|
func resetDNS(iface *net.Interface) error {
|
||||||
if ctrldnet.SupportsIPv6ListenLocal() {
|
resetDNSOnce.Do(func() {
|
||||||
if output, err := netsh("interface", "ipv6", "set", "dnsserver", strconv.Itoa(iface.Index), "dhcp"); err != nil {
|
// See corresponding comment in setDNS.
|
||||||
mainLog.Load().Warn().Err(err).Msgf("failed to reset ipv6 DNS: %s", string(output))
|
if hasLocalDnsServerRunning() {
|
||||||
|
file := absHomeDir(windowsForwardersFilename)
|
||||||
|
content, err := os.ReadFile(file)
|
||||||
|
if err != nil {
|
||||||
|
mainLog.Load().Error().Err(err).Msg("could not read forwarders settings")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
nameservers := strings.Split(string(content), ",")
|
||||||
|
if err := removeDnsServerForwarders(nameservers); err != nil {
|
||||||
|
mainLog.Load().Error().Err(err).Msg("could not remove forwarders settings")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
luid, err := winipcfg.LUIDFromIndex(uint32(iface.Index))
|
||||||
|
if err != nil {
|
||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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)
|
||||||
|
for _, ns := range nss {
|
||||||
|
if ctrldnet.IsIPv6(ns) {
|
||||||
|
v6ns = append(v6ns, ns)
|
||||||
|
} else {
|
||||||
|
v4ns = append(v4ns, ns)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
} 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)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
output, err := netsh("interface", "ipv4", "set", "dnsserver", strconv.Itoa(iface.Index), "dhcp")
|
return err
|
||||||
if err != nil {
|
|
||||||
mainLog.Load().Error().Err(err).Msgf("failed to reset ipv4 DNS: %s", string(output))
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func setPrimaryDNS(iface *net.Interface, dns string) error {
|
|
||||||
ipVer := "ipv4"
|
|
||||||
if ctrldnet.IsIPv6(dns) {
|
|
||||||
ipVer = "ipv6"
|
|
||||||
}
|
|
||||||
idx := strconv.Itoa(iface.Index)
|
|
||||||
output, err := netsh("interface", ipVer, "set", "dnsserver", idx, "static", dns)
|
|
||||||
if err != nil {
|
|
||||||
mainLog.Load().Error().Err(err).Msgf("failed to set primary DNS: %s", string(output))
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if ipVer == "ipv4" && ctrldnet.SupportsIPv6ListenLocal() {
|
|
||||||
// Disable IPv6 DNS, so the query will be fallback to IPv4.
|
|
||||||
_, _ = netsh("interface", "ipv6", "set", "dnsserver", idx, "static", "::1", "primary")
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func addSecondaryDNS(iface *net.Interface, dns string) error {
|
|
||||||
ipVer := "ipv4"
|
|
||||||
if ctrldnet.IsIPv6(dns) {
|
|
||||||
ipVer = "ipv6"
|
|
||||||
}
|
|
||||||
output, err := netsh("interface", ipVer, "add", "dns", strconv.Itoa(iface.Index), dns, "index=2")
|
|
||||||
if err != nil {
|
|
||||||
mainLog.Load().Warn().Err(err).Msgf("failed to add secondary DNS: %s", string(output))
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func netsh(args ...string) ([]byte, error) {
|
|
||||||
return exec.Command("netsh", args...).Output()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func currentDNS(iface *net.Interface) []string {
|
func currentDNS(iface *net.Interface) []string {
|
||||||
@@ -93,3 +221,112 @@ func currentDNS(iface *net.Interface) []string {
|
|||||||
}
|
}
|
||||||
return ns
|
return ns
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 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, fmt.Errorf("fallback winipcfg.LUIDFromIndex: %w", err)
|
||||||
|
}
|
||||||
|
guid, err := luid.GUID()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("fallback luid.GUID: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var ns []string
|
||||||
|
keyPaths := []string{v4InterfaceKeyPathFormat, v6InterfaceKeyPathFormat}
|
||||||
|
for _, path := range keyPaths {
|
||||||
|
interfaceKeyPath := path + guid.String()
|
||||||
|
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 {
|
||||||
|
newForwardersMap := make(map[string]struct{})
|
||||||
|
newForwarders := make([]string, len(nameservers))
|
||||||
|
for i := range nameservers {
|
||||||
|
newForwardersMap[nameservers[i]] = struct{}{}
|
||||||
|
newForwarders[i] = fmt.Sprintf("%q", nameservers[i])
|
||||||
|
}
|
||||||
|
oldForwarders := old[:0]
|
||||||
|
for _, fwd := range old {
|
||||||
|
if _, ok := newForwardersMap[fwd]; !ok {
|
||||||
|
oldForwarders = append(oldForwarders, fwd)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// NOTE: It is important to add new forwarder before removing old one.
|
||||||
|
// Testing on Windows Server 2022 shows that removing forwarder1
|
||||||
|
// then adding forwarder2 sometimes ends up adding both of them
|
||||||
|
// to the forwarders list.
|
||||||
|
cmd := fmt.Sprintf("Add-DnsServerForwarder -IPAddress %s", strings.Join(newForwarders, ","))
|
||||||
|
if len(oldForwarders) > 0 {
|
||||||
|
cmd = fmt.Sprintf("%s ; Remove-DnsServerForwarder -IPAddress %s -Force", cmd, strings.Join(oldForwarders, ","))
|
||||||
|
}
|
||||||
|
if out, err := powershell(cmd); err != nil {
|
||||||
|
return fmt.Errorf("%w: %s", err, string(out))
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// removeDnsServerForwarders removes given nameservers from DNS server forwarders list.
|
||||||
|
func removeDnsServerForwarders(nameservers []string) error {
|
||||||
|
for _, ns := range nameservers {
|
||||||
|
cmd := fmt.Sprintf("Remove-DnsServerForwarder -IPAddress %s -Force", ns)
|
||||||
|
if out, err := powershell(cmd); err != nil {
|
||||||
|
return fmt.Errorf("%w: %s", err, string(out))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
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
|
||||||
|
}
|
||||||
+1198
-104
File diff suppressed because it is too large
Load Diff
+39
-3
@@ -1,15 +1,26 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bufio"
|
||||||
|
"bytes"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"strings"
|
||||||
|
|
||||||
"github.com/kardianos/service"
|
"github.com/kardianos/service"
|
||||||
|
|
||||||
"github.com/Control-D-Inc/ctrld/internal/dns"
|
"github.com/Control-D-Inc/ctrld/internal/router"
|
||||||
)
|
)
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
if r, err := dns.NewOSConfigurator(func(format string, args ...any) {}, "lo"); err == nil {
|
if r, err := newLoopbackOSConfigurator(); err == nil {
|
||||||
useSystemdResolved = r.Mode() == "systemd-resolved"
|
useSystemdResolved = r.Mode() == "systemd-resolved"
|
||||||
}
|
}
|
||||||
|
// Disable quic-go's ECN support by default, see https://github.com/quic-go/quic-go/issues/3911
|
||||||
|
if os.Getenv("QUIC_GO_DISABLE_ECN") == "" {
|
||||||
|
os.Setenv("QUIC_GO_DISABLE_ECN", "true")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func setDependencies(svc *service.Config) {
|
func setDependencies(svc *service.Config) {
|
||||||
@@ -18,12 +29,37 @@ func setDependencies(svc *service.Config) {
|
|||||||
"After=network-online.target",
|
"After=network-online.target",
|
||||||
"Wants=NetworkManager-wait-online.service",
|
"Wants=NetworkManager-wait-online.service",
|
||||||
"After=NetworkManager-wait-online.service",
|
"After=NetworkManager-wait-online.service",
|
||||||
"Wants=systemd-networkd-wait-online.service",
|
|
||||||
"Wants=nss-lookup.target",
|
"Wants=nss-lookup.target",
|
||||||
"After=nss-lookup.target",
|
"After=nss-lookup.target",
|
||||||
}
|
}
|
||||||
|
if out, _ := exec.Command("networkctl", "--no-pager").CombinedOutput(); len(out) > 0 {
|
||||||
|
if wantsSystemDNetworkdWaitOnline(bytes.NewReader(out)) {
|
||||||
|
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) {
|
func setWorkingDirectory(svc *service.Config, dir string) {
|
||||||
svc.WorkingDirectory = dir
|
svc.WorkingDirectory = dir
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// wantsSystemDNetworkdWaitOnline reports whether "systemd-networkd-wait-online" service
|
||||||
|
// is required to be added to ctrld dependencies services.
|
||||||
|
// The input reader r is the output of "networkctl --no-pager" command.
|
||||||
|
func wantsSystemDNetworkdWaitOnline(r io.Reader) bool {
|
||||||
|
scanner := bufio.NewScanner(r)
|
||||||
|
// Skip header
|
||||||
|
scanner.Scan()
|
||||||
|
configured := false
|
||||||
|
for scanner.Scan() {
|
||||||
|
fields := strings.Fields(scanner.Text())
|
||||||
|
if len(fields) > 0 && fields[len(fields)-1] == "configured" {
|
||||||
|
configured = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return configured
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,48 @@
|
|||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
networkctlUnmanagedOutput = `IDX LINK TYPE OPERATIONAL SETUP
|
||||||
|
1 lo loopback carrier unmanaged
|
||||||
|
2 wlp0s20f3 wlan routable unmanaged
|
||||||
|
3 tailscale0 none routable unmanaged
|
||||||
|
4 br-9ac33145e060 bridge no-carrier unmanaged
|
||||||
|
5 docker0 bridge no-carrier unmanaged
|
||||||
|
|
||||||
|
5 links listed.
|
||||||
|
`
|
||||||
|
networkctlManagedOutput = `IDX LINK TYPE OPERATIONAL SETUP
|
||||||
|
1 lo loopback carrier unmanaged
|
||||||
|
2 wlp0s20f3 wlan routable configured
|
||||||
|
3 tailscale0 none routable unmanaged
|
||||||
|
4 br-9ac33145e060 bridge no-carrier unmanaged
|
||||||
|
5 docker0 bridge no-carrier unmanaged
|
||||||
|
|
||||||
|
5 links listed.
|
||||||
|
`
|
||||||
|
)
|
||||||
|
|
||||||
|
func Test_wantsSystemDNetworkdWaitOnline(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
r io.Reader
|
||||||
|
required bool
|
||||||
|
}{
|
||||||
|
{"unmanaged", strings.NewReader(networkctlUnmanagedOutput), false},
|
||||||
|
{"managed", strings.NewReader(networkctlManagedOutput), true},
|
||||||
|
{"empty", strings.NewReader(""), false},
|
||||||
|
}
|
||||||
|
for _, tc := range tests {
|
||||||
|
tc := tc
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
if required := wantsSystemDNetworkdWaitOnline(tc.r); required != tc.required {
|
||||||
|
t.Errorf("wants %v got %v", tc.required, required)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
//go:build !linux && !freebsd && !darwin
|
//go:build !linux && !freebsd && !darwin && !windows
|
||||||
|
|
||||||
package cli
|
package cli
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,305 @@
|
|||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net"
|
||||||
|
"net/url"
|
||||||
|
"runtime"
|
||||||
|
"syscall"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"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{}}
|
||||||
|
|
||||||
|
// Default value is true.
|
||||||
|
assert.True(t, p.dnsWatchdogEnabled())
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
enabled bool
|
||||||
|
}{
|
||||||
|
{"enabled", true},
|
||||||
|
{"disabled", false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
tc := tc
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
p.cfg.Service.DnsWatchdogEnabled = &tc.enabled
|
||||||
|
assert.Equal(t, tc.enabled, p.dnsWatchdogEnabled())
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func Test_prog_dnsWatchdogInterval(t *testing.T) {
|
||||||
|
p := &prog{cfg: &ctrld.Config{}}
|
||||||
|
|
||||||
|
// Default value is 20s.
|
||||||
|
assert.Equal(t, dnsWatchdogDefaultInterval, p.dnsWatchdogDuration())
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
duration time.Duration
|
||||||
|
expected time.Duration
|
||||||
|
}{
|
||||||
|
{"valid", time.Minute, time.Minute},
|
||||||
|
{"zero", 0, dnsWatchdogDefaultInterval},
|
||||||
|
{"nagative", time.Duration(-1 * time.Minute), dnsWatchdogDefaultInterval},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
tc := tc
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
p.cfg.Service.DnsWatchdogInvterval = &tc.duration
|
||||||
|
assert.Equal(t, tc.expected, p.dnsWatchdogDuration())
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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
|
||||||
|
}
|
||||||
@@ -51,7 +51,7 @@ var statsClientQueriesCount = prometheus.NewCounterVec(prometheus.CounterOpts{
|
|||||||
|
|
||||||
// WithLabelValuesInc increases prometheus counter by 1 if query stats is enabled.
|
// WithLabelValuesInc increases prometheus counter by 1 if query stats is enabled.
|
||||||
func (p *prog) WithLabelValuesInc(c *prometheus.CounterVec, lvs ...string) {
|
func (p *prog) WithLabelValuesInc(c *prometheus.CounterVec, lvs ...string) {
|
||||||
if p.cfg.Service.MetricsQueryStats {
|
if p.metricsQueryStats.Load() {
|
||||||
c.WithLabelValues(lvs...).Inc()
|
c.WithLabelValues(lvs...).Inc()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,164 @@
|
|||||||
|
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) {
|
||||||
|
resolvConfPath := "/etc/resolv.conf"
|
||||||
|
// Evaluating symbolics link to watch the target file that /etc/resolv.conf point to.
|
||||||
|
if rp, _ := filepath.EvalSymlinks(resolvConfPath); rp != "" {
|
||||||
|
resolvConfPath = rp
|
||||||
|
}
|
||||||
|
mainLog.Load().Debug().Msgf("start watching %s file", resolvConfPath)
|
||||||
|
watcher, err := fsnotify.NewWatcher()
|
||||||
|
if err != nil {
|
||||||
|
mainLog.Load().Warn().Err(err).Msg("could not create watcher for /etc/resolv.conf")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer watcher.Close()
|
||||||
|
|
||||||
|
// We watch /etc instead of /etc/resolv.conf directly,
|
||||||
|
// see: https://github.com/fsnotify/fsnotify#watching-a-file-doesnt-work-well
|
||||||
|
watchDir := filepath.Dir(resolvConfPath)
|
||||||
|
if err := watcher.Add(watchDir); err != nil {
|
||||||
|
mainLog.Load().Warn().Err(err).Msgf("could not add %s to watcher list", watchDir)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-p.dnsWatcherStopCh:
|
||||||
|
return
|
||||||
|
case <-p.stopCh:
|
||||||
|
mainLog.Load().Debug().Msgf("stopping watcher for %s", resolvConfPath)
|
||||||
|
return
|
||||||
|
case event, ok := <-watcher.Events:
|
||||||
|
if p.recoveryRunning.Load() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if event.Name != resolvConfPath { // skip if not /etc/resolv.conf changes.
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if event.Has(fsnotify.Write) || event.Has(fsnotify.Create) {
|
||||||
|
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()
|
||||||
|
}
|
||||||
|
|
||||||
|
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 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:
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
mainLog.Load().Err(err).Msg("could not get event for /etc/resolv.conf")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,49 @@
|
|||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"slices"
|
||||||
|
|
||||||
|
"github.com/Control-D-Inc/ctrld/internal/dns/resolvconffile"
|
||||||
|
)
|
||||||
|
|
||||||
|
const resolvConfPath = "/etc/resolv.conf"
|
||||||
|
|
||||||
|
// setResolvConf sets the content of resolv.conf file using the given nameservers list.
|
||||||
|
func setResolvConf(iface *net.Interface, ns []netip.Addr) error {
|
||||||
|
servers := make([]string, len(ns))
|
||||||
|
for i := range ns {
|
||||||
|
servers[i] = ns[i].String()
|
||||||
|
}
|
||||||
|
if err := setDNS(iface, servers); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
slices.Sort(servers)
|
||||||
|
curNs := currentDNS(iface)
|
||||||
|
slices.Sort(curNs)
|
||||||
|
if !slices.Equal(curNs, servers) {
|
||||||
|
c, err := resolvconffile.ParseFile(resolvConfPath)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
c.Nameservers = ns
|
||||||
|
f, err := os.Create(resolvConfPath)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer f.Close()
|
||||||
|
|
||||||
|
if err := c.Write(f); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return f.Close()
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// shouldWatchResolvconf reports whether ctrld should watch changes to resolv.conf file with given OS configurator.
|
||||||
|
func shouldWatchResolvconf() bool {
|
||||||
|
return true
|
||||||
|
}
|
||||||
@@ -0,0 +1,52 @@
|
|||||||
|
//go:build unix && !darwin
|
||||||
|
|
||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
|
||||||
|
"tailscale.com/control/controlknobs"
|
||||||
|
"tailscale.com/health"
|
||||||
|
"tailscale.com/util/dnsname"
|
||||||
|
|
||||||
|
"github.com/Control-D-Inc/ctrld/internal/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
// setResolvConf sets the content of the resolv.conf file using the given nameservers list.
|
||||||
|
func setResolvConf(iface *net.Interface, ns []netip.Addr) error {
|
||||||
|
r, err := newLoopbackOSConfigurator()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
oc := dns.OSConfig{
|
||||||
|
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 {
|
||||||
|
r, err := newLoopbackOSConfigurator()
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
switch r.Mode() {
|
||||||
|
case "direct", "resolvconf":
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
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,16 @@
|
|||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
)
|
||||||
|
|
||||||
|
// setResolvConf sets the content of resolv.conf file using the given nameservers list.
|
||||||
|
func setResolvConf(_ *net.Interface, _ []netip.Addr) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// shouldWatchResolvconf reports whether ctrld should watch changes to resolv.conf file with given OS configurator.
|
||||||
|
func shouldWatchResolvconf() bool {
|
||||||
|
return false
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
//go:build !windows
|
||||||
|
|
||||||
|
package cli
|
||||||
|
|
||||||
|
var supportedSelfDelete = true
|
||||||
|
|
||||||
|
func selfDeleteExe() error { return nil }
|
||||||
@@ -0,0 +1,134 @@
|
|||||||
|
// Copied from https://github.com/secur30nly/go-self-delete
|
||||||
|
// with modification to suitable for ctrld usage.
|
||||||
|
|
||||||
|
/*
|
||||||
|
License: MIT Licence
|
||||||
|
|
||||||
|
References:
|
||||||
|
- https://github.com/LloydLabs/delete-self-poc
|
||||||
|
- https://twitter.com/jonasLyk/status/1350401461985955840
|
||||||
|
*/
|
||||||
|
|
||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
)
|
||||||
|
|
||||||
|
var supportedSelfDelete = false
|
||||||
|
|
||||||
|
type FILE_RENAME_INFO struct {
|
||||||
|
Union struct {
|
||||||
|
ReplaceIfExists bool
|
||||||
|
Flags uint32
|
||||||
|
}
|
||||||
|
RootDirectory windows.Handle
|
||||||
|
FileNameLength uint32
|
||||||
|
FileName [1]uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
type FILE_DISPOSITION_INFO struct {
|
||||||
|
DeleteFile bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func dsOpenHandle(pwPath *uint16) (windows.Handle, error) {
|
||||||
|
handle, err := windows.CreateFile(
|
||||||
|
pwPath,
|
||||||
|
windows.DELETE,
|
||||||
|
0,
|
||||||
|
nil,
|
||||||
|
windows.OPEN_EXISTING,
|
||||||
|
windows.FILE_ATTRIBUTE_NORMAL,
|
||||||
|
0,
|
||||||
|
)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return handle, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func dsRenameHandle(hHandle windows.Handle) error {
|
||||||
|
var fRename FILE_RENAME_INFO
|
||||||
|
DS_STREAM_RENAME, err := windows.UTF16FromString(":deadbeef")
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
lpwStream := &DS_STREAM_RENAME[0]
|
||||||
|
fRename.FileNameLength = uint32(unsafe.Sizeof(lpwStream))
|
||||||
|
|
||||||
|
windows.NewLazyDLL("kernel32.dll").NewProc("RtlCopyMemory").Call(
|
||||||
|
uintptr(unsafe.Pointer(&fRename.FileName[0])),
|
||||||
|
uintptr(unsafe.Pointer(lpwStream)),
|
||||||
|
unsafe.Sizeof(lpwStream),
|
||||||
|
)
|
||||||
|
|
||||||
|
err = windows.SetFileInformationByHandle(
|
||||||
|
hHandle,
|
||||||
|
windows.FileRenameInfo,
|
||||||
|
(*byte)(unsafe.Pointer(&fRename)),
|
||||||
|
uint32(unsafe.Sizeof(fRename)+unsafe.Sizeof(lpwStream)),
|
||||||
|
)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func dsDepositeHandle(hHandle windows.Handle) error {
|
||||||
|
var fDelete FILE_DISPOSITION_INFO
|
||||||
|
fDelete.DeleteFile = true
|
||||||
|
|
||||||
|
err := windows.SetFileInformationByHandle(
|
||||||
|
hHandle,
|
||||||
|
windows.FileDispositionInfo,
|
||||||
|
(*byte)(unsafe.Pointer(&fDelete)),
|
||||||
|
uint32(unsafe.Sizeof(fDelete)),
|
||||||
|
)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func selfDeleteExe() error {
|
||||||
|
var wcPath [windows.MAX_PATH + 1]uint16
|
||||||
|
var hCurrent windows.Handle
|
||||||
|
|
||||||
|
_, err := windows.GetModuleFileName(0, &wcPath[0], windows.MAX_PATH)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
hCurrent, err = dsOpenHandle(&wcPath[0])
|
||||||
|
if err != nil || hCurrent == windows.InvalidHandle {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := dsRenameHandle(hCurrent); err != nil {
|
||||||
|
_ = windows.CloseHandle(hCurrent)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
_ = windows.CloseHandle(hCurrent)
|
||||||
|
|
||||||
|
hCurrent, err = dsOpenHandle(&wcPath[0])
|
||||||
|
if err != nil || hCurrent == windows.InvalidHandle {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := dsDepositeHandle(hCurrent); err != nil {
|
||||||
|
_ = windows.CloseHandle(hCurrent)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return windows.CloseHandle(hCurrent)
|
||||||
|
}
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
//go:build !unix
|
||||||
|
|
||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
|
||||||
|
"github.com/rs/zerolog"
|
||||||
|
)
|
||||||
|
|
||||||
|
func selfUninstall(p *prog, logger zerolog.Logger) {
|
||||||
|
if uninstallInvalidCdUID(p, logger, false) {
|
||||||
|
logger.Warn().Msgf("service was uninstalled because device %q does not exist", cdUID)
|
||||||
|
os.Exit(0)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,45 @@
|
|||||||
|
//go:build unix
|
||||||
|
|
||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"runtime"
|
||||||
|
"syscall"
|
||||||
|
|
||||||
|
"github.com/rs/zerolog"
|
||||||
|
)
|
||||||
|
|
||||||
|
func selfUninstall(p *prog, logger zerolog.Logger) {
|
||||||
|
if runtime.GOOS == "linux" {
|
||||||
|
selfUninstallLinux(p, logger)
|
||||||
|
}
|
||||||
|
|
||||||
|
bin, err := os.Executable()
|
||||||
|
if err != nil {
|
||||||
|
logger.Fatal().Err(err).Msg("could not determine executable")
|
||||||
|
}
|
||||||
|
args := []string{"uninstall"}
|
||||||
|
if deactivationPinSet() {
|
||||||
|
args = append(args, fmt.Sprintf("--pin=%d", cdDeactivationPin.Load()))
|
||||||
|
}
|
||||||
|
cmd := exec.Command(bin, args...)
|
||||||
|
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
|
||||||
|
if err := cmd.Start(); err != nil {
|
||||||
|
logger.Fatal().Err(err).Msg("could not start self uninstall command")
|
||||||
|
}
|
||||||
|
cmd.Stdout = os.Stdout
|
||||||
|
cmd.Stderr = os.Stderr
|
||||||
|
logger.Warn().Msgf("service was uninstalled because device %q does not exist", cdUID)
|
||||||
|
_ = cmd.Wait()
|
||||||
|
os.Exit(0)
|
||||||
|
}
|
||||||
|
|
||||||
|
func selfUninstallLinux(p *prog, logger zerolog.Logger) {
|
||||||
|
if uninstallInvalidCdUID(p, logger, true) {
|
||||||
|
logger.Warn().Msgf("service was uninstalled because device %q does not exist", cdUID)
|
||||||
|
os.Exit(0)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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,
|
||||||
|
}
|
||||||
|
}
|
||||||
+111
-10
@@ -4,12 +4,16 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
|
"runtime"
|
||||||
|
|
||||||
|
"github.com/coreos/go-systemd/v22/unit"
|
||||||
"github.com/kardianos/service"
|
"github.com/kardianos/service"
|
||||||
|
|
||||||
"github.com/Control-D-Inc/ctrld/internal/router"
|
"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
|
// newService wraps service.New call to return service.Service
|
||||||
@@ -20,14 +24,17 @@ func newService(i service.Interface, c *service.Config) (service.Service, error)
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
switch {
|
switch {
|
||||||
case router.IsOldOpenwrt():
|
case router.IsOldOpenwrt(), router.IsNetGearOrbi():
|
||||||
return &procd{&sysV{s}}, nil
|
return &procd{sysV: &sysV{s}, svcConfig: c}, nil
|
||||||
case router.IsGLiNet():
|
case router.IsGLiNet():
|
||||||
return &sysV{s}, nil
|
return &sysV{s}, nil
|
||||||
case s.Platform() == "unix-systemv":
|
case s.Platform() == "unix-systemv":
|
||||||
return &sysV{s}, nil
|
return &sysV{s}, nil
|
||||||
case s.Platform() == "linux-systemd":
|
case s.Platform() == "linux-systemd":
|
||||||
return &systemd{s}, nil
|
return &systemd{s}, nil
|
||||||
|
case s.Platform() == "darwin-launchd":
|
||||||
|
return newLaunchd(s), nil
|
||||||
|
|
||||||
}
|
}
|
||||||
return s, nil
|
return s, nil
|
||||||
}
|
}
|
||||||
@@ -89,25 +96,31 @@ func (s *sysV) Status() (service.Status, error) {
|
|||||||
// like old GL.iNET Opal router.
|
// like old GL.iNET Opal router.
|
||||||
type procd struct {
|
type procd struct {
|
||||||
*sysV
|
*sysV
|
||||||
|
svcConfig *service.Config
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *procd) Status() (service.Status, error) {
|
func (s *procd) Status() (service.Status, error) {
|
||||||
if !s.installed() {
|
if !s.installed() {
|
||||||
return service.StatusUnknown, service.ErrNotInstalled
|
return service.StatusUnknown, service.ErrNotInstalled
|
||||||
}
|
}
|
||||||
exe, err := os.Executable()
|
bin := s.svcConfig.Executable
|
||||||
if err != nil {
|
if bin == "" {
|
||||||
return service.StatusUnknown, nil
|
exe, err := os.Executable()
|
||||||
|
if err != nil {
|
||||||
|
return service.StatusUnknown, nil
|
||||||
|
}
|
||||||
|
bin = exe
|
||||||
}
|
}
|
||||||
|
|
||||||
// Looking for something like "/sbin/ctrld run ".
|
// Looking for something like "/sbin/ctrld run ".
|
||||||
shellCmd := fmt.Sprintf("ps | grep -q %q", exe+" [r]un ")
|
shellCmd := fmt.Sprintf("ps | grep -q %q", bin+" [r]un ")
|
||||||
if err := exec.Command("sh", "-c", shellCmd).Run(); err != nil {
|
if err := exec.Command("sh", "-c", shellCmd).Run(); err != nil {
|
||||||
return service.StatusStopped, nil
|
return service.StatusStopped, nil
|
||||||
}
|
}
|
||||||
return service.StatusRunning, nil
|
return service.StatusRunning, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// procd wraps a service.Service, and provide status command to
|
// systemd wraps a service.Service, and provide status command to
|
||||||
// report the status correctly.
|
// report the status correctly.
|
||||||
type systemd struct {
|
type systemd struct {
|
||||||
service.Service
|
service.Service
|
||||||
@@ -121,20 +134,101 @@ func (s *systemd) Status() (service.Status, error) {
|
|||||||
return s.Service.Status()
|
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)
|
||||||
|
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,
|
||||||
|
statusErrMsg: "Permission denied",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// launchd wraps a service.Service, and provide status command to
|
||||||
|
// report the status correctly when not running as root on Darwin.
|
||||||
|
//
|
||||||
|
// TODO: remove this wrapper once https://github.com/kardianos/service/issues/400 fixed.
|
||||||
|
type launchd struct {
|
||||||
|
service.Service
|
||||||
|
statusErrMsg string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (l *launchd) Status() (service.Status, error) {
|
||||||
|
if os.Geteuid() != 0 {
|
||||||
|
return service.StatusUnknown, errors.New(l.statusErrMsg)
|
||||||
|
}
|
||||||
|
return l.Service.Status()
|
||||||
|
}
|
||||||
|
|
||||||
type task struct {
|
type task struct {
|
||||||
f func() error
|
f func() error
|
||||||
abortOnError bool
|
abortOnError bool
|
||||||
|
Name string
|
||||||
}
|
}
|
||||||
|
|
||||||
func doTasks(tasks []task) bool {
|
func doTasks(tasks []task) bool {
|
||||||
var prevErr error
|
|
||||||
for _, task := range tasks {
|
for _, task := range tasks {
|
||||||
|
mainLog.Load().Debug().Msgf("Running task %s", task.Name)
|
||||||
if err := task.f(); err != nil {
|
if err := task.f(); err != nil {
|
||||||
if task.abortOnError {
|
if task.abortOnError {
|
||||||
mainLog.Load().Error().Msg(errors.Join(prevErr, err).Error())
|
mainLog.Load().Error().Msgf("error running task %s: %v", task.Name, err)
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
prevErr = err
|
// if this is darwin stop command, dont print debug
|
||||||
|
// since launchctl complains on every start
|
||||||
|
if runtime.GOOS != "darwin" || task.Name != "Stop" {
|
||||||
|
mainLog.Load().Debug().Msgf("error running task %s: %v", task.Name, err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return true
|
return true
|
||||||
@@ -155,6 +249,13 @@ func checkHasElevatedPrivilege() {
|
|||||||
func unixSystemVServiceStatus() (service.Status, error) {
|
func unixSystemVServiceStatus() (service.Status, error) {
|
||||||
out, err := exec.Command("/etc/init.d/ctrld", "status").CombinedOutput()
|
out, err := exec.Command("/etc/init.d/ctrld", "status").CombinedOutput()
|
||||||
if err != nil {
|
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
|
return service.StatusUnknown, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,134 @@
|
|||||||
|
//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 the "defaults" command, which is the standard way to edit plists.
|
||||||
|
//
|
||||||
|
// The function is idempotent: if the flag already exists, it's a no-op.
|
||||||
|
func appendServiceFlag(flag string) error {
|
||||||
|
// Read current ProgramArguments from plist.
|
||||||
|
out, err := exec.Command("defaults", "read", launchdPlistPath, "ProgramArguments").CombinedOutput()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to read plist ProgramArguments: %w (output: %s)", err, strings.TrimSpace(string(out)))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if the flag is already present (idempotent).
|
||||||
|
args := string(out)
|
||||||
|
if strings.Contains(args, flag) {
|
||||||
|
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 a CLI flag (and its value, if the next argument is not
|
||||||
|
// a flag) from the installed service's launch arguments. For example, removing
|
||||||
|
// "--intercept-mode" also removes the following "dns" or "hard" value argument.
|
||||||
|
//
|
||||||
|
// 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 := -1
|
||||||
|
for i, entry := range entries {
|
||||||
|
if entry == flag {
|
||||||
|
index = i
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if index < 0 {
|
||||||
|
mainLog.Load().Debug().Msgf("Service flag %q not present in plist, skipping removal", flag)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if the next entry is a value (not a flag). If so, delete it first
|
||||||
|
// (deleting by index shifts subsequent entries down, so delete value before flag).
|
||||||
|
hasValue := index+1 < len(entries) && !strings.HasPrefix(entries[index+1], "-")
|
||||||
|
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
|
||||||
|
}
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
//go:build !darwin && !windows
|
||||||
|
|
||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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 fmt.Errorf("appending service flags is not supported on this platform; use intercept_mode in config instead")
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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 fmt.Errorf("removing service flags is not supported on this platform; use intercept_mode in config instead")
|
||||||
|
}
|
||||||
@@ -0,0 +1,153 @@
|
|||||||
|
//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 if flag already present (idempotent).
|
||||||
|
if strings.Contains(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 a CLI flag (and its value, if present) from the installed
|
||||||
|
// Windows service's BinPath. For example, removing "--intercept-mode" also removes
|
||||||
|
// the following "dns" or "hard" value. The function is idempotent.
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !strings.Contains(config.BinaryPathName, flag) {
|
||||||
|
mainLog.Load().Debug().Msgf("Service flag %q not present in BinPath, skipping removal", flag)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Split BinPath into parts, find and remove the flag + its value (if any).
|
||||||
|
parts := strings.Fields(config.BinaryPathName)
|
||||||
|
var newParts []string
|
||||||
|
for i := 0; i < len(parts); i++ {
|
||||||
|
if parts[i] == flag {
|
||||||
|
// Skip the flag. Also skip the next part if it's a value (not a flag).
|
||||||
|
if i+1 < len(parts) && !strings.HasPrefix(parts[i+1], "-") {
|
||||||
|
i++ // skip value too
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
newParts = append(newParts, parts[i])
|
||||||
|
}
|
||||||
|
config.BinaryPathName = strings.Join(newParts, " ")
|
||||||
|
|
||||||
|
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
|
||||||
|
}
|
||||||
@@ -9,3 +9,14 @@ import (
|
|||||||
func hasElevatedPrivilege() (bool, error) {
|
func hasElevatedPrivilege() (bool, error) {
|
||||||
return os.Geteuid() == 0, nil
|
return os.Geteuid() == 0, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
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,28 @@
|
|||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
+205
-1
@@ -1,6 +1,22 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import "golang.org/x/sys/windows"
|
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) {
|
func hasElevatedPrivilege() (bool, error) {
|
||||||
var sid *windows.SID
|
var sid *windows.SID
|
||||||
@@ -22,3 +38,191 @@ func hasElevatedPrivilege() (bool, error) {
|
|||||||
token := windows.Token(0)
|
token := windows.Token(0)
|
||||||
return token.IsMember(sid)
|
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}
|
||||||
|
}
|
||||||
|
|
||||||
|
pathP, err := syscall.UTF16PtrFromString(path)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var access uint32
|
||||||
|
switch mode & (os.O_RDONLY | os.O_WRONLY | os.O_RDWR) {
|
||||||
|
case os.O_RDONLY:
|
||||||
|
access = windows.GENERIC_READ
|
||||||
|
case os.O_WRONLY:
|
||||||
|
access = windows.GENERIC_WRITE
|
||||||
|
case os.O_RDWR:
|
||||||
|
access = windows.GENERIC_READ | windows.GENERIC_WRITE
|
||||||
|
}
|
||||||
|
if mode&os.O_CREATE != 0 {
|
||||||
|
access |= windows.GENERIC_WRITE
|
||||||
|
}
|
||||||
|
if mode&os.O_APPEND != 0 {
|
||||||
|
access &^= windows.GENERIC_WRITE
|
||||||
|
access |= windows.FILE_APPEND_DATA
|
||||||
|
}
|
||||||
|
|
||||||
|
shareMode := uint32(syscall.FILE_SHARE_READ | syscall.FILE_SHARE_WRITE | syscall.FILE_SHARE_DELETE)
|
||||||
|
|
||||||
|
var sa *syscall.SecurityAttributes
|
||||||
|
|
||||||
|
var createMode uint32
|
||||||
|
switch {
|
||||||
|
case mode&(os.O_CREATE|os.O_EXCL) == (os.O_CREATE | os.O_EXCL):
|
||||||
|
createMode = windows.CREATE_NEW
|
||||||
|
case mode&(os.O_CREATE|os.O_TRUNC) == (os.O_CREATE | os.O_TRUNC):
|
||||||
|
createMode = windows.CREATE_ALWAYS
|
||||||
|
case mode&os.O_CREATE == os.O_CREATE:
|
||||||
|
createMode = windows.OPEN_ALWAYS
|
||||||
|
case mode&os.O_TRUNC == os.O_TRUNC:
|
||||||
|
createMode = windows.TRUNCATE_EXISTING
|
||||||
|
default:
|
||||||
|
createMode = windows.OPEN_EXISTING
|
||||||
|
}
|
||||||
|
|
||||||
|
handle, err := syscall.CreateFile(pathP, access, shareMode, sa, createMode, syscall.FILE_ATTRIBUTE_NORMAL, 0)
|
||||||
|
if err != nil {
|
||||||
|
return nil, &os.PathError{Path: path, Op: "open", Err: err}
|
||||||
|
}
|
||||||
|
|
||||||
|
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
|
||||||
|
}
|
||||||
+84
-51
@@ -1,38 +1,60 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/miekg/dns"
|
|
||||||
|
|
||||||
"github.com/Control-D-Inc/ctrld"
|
"github.com/Control-D-Inc/ctrld"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// maxFailureRequest is the maximum failed queries allowed before an upstream is marked as down.
|
// 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 is the time interval between each upstream checks.
|
||||||
checkUpstreamBackoffSleep = 2 * time.Second
|
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.
|
// upstreamMonitor performs monitoring upstreams health.
|
||||||
type upstreamMonitor struct {
|
type upstreamMonitor struct {
|
||||||
cfg *ctrld.Config
|
cfg *ctrld.Config
|
||||||
|
|
||||||
mu sync.Mutex
|
mu sync.RWMutex
|
||||||
checking map[string]bool
|
checking map[string]bool
|
||||||
down map[string]bool
|
down map[string]bool
|
||||||
failureReq map[string]uint64
|
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 {
|
func newUpstreamMonitor(cfg *ctrld.Config) *upstreamMonitor {
|
||||||
um := &upstreamMonitor{
|
um := &upstreamMonitor{
|
||||||
cfg: cfg,
|
cfg: cfg,
|
||||||
checking: make(map[string]bool),
|
checking: make(map[string]bool),
|
||||||
down: make(map[string]bool),
|
down: make(map[string]bool),
|
||||||
failureReq: make(map[string]uint64),
|
failureReq: make(map[string]uint64),
|
||||||
|
recovered: make(map[string]bool),
|
||||||
|
failureTimerActive: make(map[string]bool),
|
||||||
}
|
}
|
||||||
for n := range cfg.Upstream {
|
for n := range cfg.Upstream {
|
||||||
upstream := upstreamPrefix + n
|
upstream := upstreamPrefix + n
|
||||||
@@ -42,14 +64,47 @@ func newUpstreamMonitor(cfg *ctrld.Config) *upstreamMonitor {
|
|||||||
return um
|
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) {
|
func (um *upstreamMonitor) increaseFailureCount(upstream string) {
|
||||||
um.mu.Lock()
|
um.mu.Lock()
|
||||||
defer um.mu.Unlock()
|
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
|
um.failureReq[upstream] += 1
|
||||||
failedCount := um.failureReq[upstream]
|
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.
|
// isDown reports whether the given upstream is being marked as down.
|
||||||
@@ -63,50 +118,28 @@ func (um *upstreamMonitor) isDown(upstream string) bool {
|
|||||||
// reset marks an upstream as up and set failed queries counter to zero.
|
// reset marks an upstream as up and set failed queries counter to zero.
|
||||||
func (um *upstreamMonitor) reset(upstream string) {
|
func (um *upstreamMonitor) reset(upstream string) {
|
||||||
um.mu.Lock()
|
um.mu.Lock()
|
||||||
defer um.mu.Unlock()
|
|
||||||
|
|
||||||
um.failureReq[upstream] = 0
|
um.failureReq[upstream] = 0
|
||||||
um.down[upstream] = false
|
um.down[upstream] = false
|
||||||
}
|
um.recovered[upstream] = true
|
||||||
|
|
||||||
// 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 (um *upstreamMonitor) checkUpstream(upstream string, uc *ctrld.UpstreamConfig) {
|
|
||||||
um.mu.Lock()
|
|
||||||
isChecking := um.checking[upstream]
|
|
||||||
if isChecking {
|
|
||||||
um.mu.Unlock()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
um.checking[upstream] = true
|
|
||||||
um.mu.Unlock()
|
um.mu.Unlock()
|
||||||
defer func() {
|
go func() {
|
||||||
|
// debounce the recovery to avoid incrementing failure counts already in flight
|
||||||
|
time.Sleep(1 * time.Second)
|
||||||
um.mu.Lock()
|
um.mu.Lock()
|
||||||
um.checking[upstream] = false
|
um.recovered[upstream] = false
|
||||||
um.mu.Unlock()
|
um.mu.Unlock()
|
||||||
}()
|
}()
|
||||||
|
}
|
||||||
resolver, err := ctrld.NewResolver(uc)
|
|
||||||
if err != nil {
|
// countHealthy returns the number of upstreams in the provided map that are considered healthy.
|
||||||
mainLog.Load().Warn().Err(err).Msg("could not check upstream")
|
func (um *upstreamMonitor) countHealthy(upstreams []string) int {
|
||||||
return
|
var count int
|
||||||
}
|
um.mu.RLock()
|
||||||
msg := new(dns.Msg)
|
for _, upstream := range upstreams {
|
||||||
msg.SetQuestion(".", dns.TypeNS)
|
if !um.down[upstream] {
|
||||||
|
count++
|
||||||
check := func() error {
|
}
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
}
|
||||||
defer cancel()
|
um.mu.RUnlock()
|
||||||
uc.ReBootstrap()
|
return count
|
||||||
_, err := resolver.Resolve(ctx, msg)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
for {
|
|
||||||
if err := check(); err == nil {
|
|
||||||
mainLog.Load().Debug().Msgf("upstream %q is online", uc.Endpoint)
|
|
||||||
um.reset(upstream)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
time.Sleep(checkUpstreamBackoffSleep)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,420 @@
|
|||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net"
|
||||||
|
"runtime"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
|
||||||
|
"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
|
||||||
|
// 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
|
||||||
|
// refreshRunning keeps noisy network-change storms from running overlapping
|
||||||
|
// scutil/networksetup VPN DNS discovery work.
|
||||||
|
refreshRunning atomic.Bool
|
||||||
|
// 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.
|
||||||
|
func (m *vpnDNSManager) Refresh(guardAgainstNoNameservers bool) {
|
||||||
|
logger := mainLog.Load()
|
||||||
|
if !m.refreshRunning.CompareAndSwap(false, true) {
|
||||||
|
logger.Debug().Msg("VPN DNS refresh already running, skipping duplicate")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer m.refreshRunning.Store(false)
|
||||||
|
|
||||||
|
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()
|
||||||
|
|
||||||
|
previousExemptions := m.currentExemptionsLocked()
|
||||||
|
|
||||||
|
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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
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 to permit VPN DNS traffic only when the exemption set
|
||||||
|
// actually changes. Network-change events can fire repeatedly while macOS/VPN
|
||||||
|
// state is otherwise identical; rewriting pf for identical exemptions can feed
|
||||||
|
// a self-triggering network-change loop. Empty exemptions are still applied
|
||||||
|
// when they differ from the previous set, so stale VPN exemptions are cleared
|
||||||
|
// on disconnect.
|
||||||
|
m.updateInterceptExemptionsIfChanged(logger, previousExemptions, exemptions, "VPN DNS")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *vpnDNSManager) updateInterceptExemptionsIfChanged(logger *zerolog.Logger, before, after []vpnDNSExemption, reason string) {
|
||||||
|
if m.onServersChanged == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if vpnDNSExemptionsEqual(before, after) {
|
||||||
|
logger.Debug().Msgf("VPN DNS exemptions unchanged after %s refresh; skipping intercept rule update", reason)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := m.onServersChanged(after); err != nil {
|
||||||
|
logger.Error().Err(err).Msg("Failed to update intercept exemptions for VPN DNS servers")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// RefreshRoutesOnly re-discovers VPN DNS configs and updates only ctrld's
|
||||||
|
// in-memory split-DNS routes. It intentionally does not call onServersChanged,
|
||||||
|
// so it does not rewrite/reload pf/WFP rules. Use this for post-settle discovery
|
||||||
|
// checks where we only need to learn late-published VPN search domains.
|
||||||
|
func (m *vpnDNSManager) RefreshRoutesOnly() (routes, domainlessServers, exemptions int) {
|
||||||
|
logger := mainLog.Load()
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
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(m.currentExemptionsLocked()))
|
||||||
|
return len(m.routes), len(m.domainlessServers), len(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,147 @@
|
|||||||
|
package cli
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/Control-D-Inc/ctrld"
|
||||||
|
)
|
||||||
|
|
||||||
|
func withVPNDNSSettlingEnabled(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
old := vpnDNSSettlingEnabled
|
||||||
|
vpnDNSSettlingEnabled = true
|
||||||
|
t.Cleanup(func() { vpnDNSSettlingEnabled = old })
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestVPNDNSRefreshSkipsConcurrentDuplicate(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 {
|
||||||
|
calls.Add(1)
|
||||||
|
once.Do(func() { close(started) })
|
||||||
|
<-release
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
defer close(done)
|
||||||
|
m.Refresh(true)
|
||||||
|
}()
|
||||||
|
|
||||||
|
<-started
|
||||||
|
m.Refresh(true)
|
||||||
|
close(release)
|
||||||
|
<-done
|
||||||
|
|
||||||
|
if calls.Load() != 1 {
|
||||||
|
t.Fatalf("expected overlapping refresh to be skipped, got %d discovery calls", calls.Load())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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
|
||||||
|
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"},
|
||||||
|
}}
|
||||||
|
m.domainlessServers = []string{"10.25.37.21"}
|
||||||
|
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 len(gotExemptions) != 0 {
|
||||||
|
t.Fatalf("expected empty exemptions after clearing stale state, got %v", 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 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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
{
|
||||||
|
"RT_VERSION": {
|
||||||
|
"#1": {
|
||||||
|
"0000": {
|
||||||
|
"fixed": {
|
||||||
|
"file_version": "0.0.0.1"
|
||||||
|
},
|
||||||
|
"info": {
|
||||||
|
"0409": {
|
||||||
|
"CompanyName": "ControlD Inc",
|
||||||
|
"FileDescription": "Control D DNS daemon",
|
||||||
|
"ProductName": "ctrld",
|
||||||
|
"InternalName": "ctrld",
|
||||||
|
"LegalCopyright": "ControlD Inc 2024"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,4 @@
|
|||||||
|
//go:generate go-winres make --product-version=git-tag --file-version=git-tag
|
||||||
|
package cli
|
||||||
|
|
||||||
|
// Placeholder file for windows builds.
|
||||||
+7
-1
@@ -1,7 +1,13 @@
|
|||||||
package main
|
package main
|
||||||
|
|
||||||
import "github.com/Control-D-Inc/ctrld/cmd/cli"
|
import (
|
||||||
|
"os"
|
||||||
|
|
||||||
|
"github.com/Control-D-Inc/ctrld/cmd/cli"
|
||||||
|
)
|
||||||
|
|
||||||
func main() {
|
func main() {
|
||||||
cli.Main()
|
cli.Main()
|
||||||
|
// make sure we exit with 0 if there are no errors
|
||||||
|
os.Exit(0)
|
||||||
}
|
}
|
||||||
|
|||||||
+17
-10
@@ -28,15 +28,17 @@ type AppCallback interface {
|
|||||||
// Start configures utility with config.toml from provided directory.
|
// Start configures utility with config.toml from provided directory.
|
||||||
// This function will block until Stop is called
|
// This function will block until Stop is called
|
||||||
// Check port availability prior to calling it.
|
// Check port availability prior to calling it.
|
||||||
func (c *Controller) Start(CdUID string, HomeDir string, UpstreamProto string, logLevel int, logPath string) {
|
func (c *Controller) Start(CdUID string, ProvisionID string, CustomHostname string, HomeDir string, UpstreamProto string, logLevel int, logPath string) {
|
||||||
if c.stopCh == nil {
|
if c.stopCh == nil {
|
||||||
c.stopCh = make(chan struct{})
|
c.stopCh = make(chan struct{})
|
||||||
c.Config = cli.AppConfig{
|
c.Config = cli.AppConfig{
|
||||||
CdUID: CdUID,
|
CdUID: CdUID,
|
||||||
HomeDir: HomeDir,
|
ProvisionID: ProvisionID,
|
||||||
UpstreamProto: UpstreamProto,
|
CustomHostname: CustomHostname,
|
||||||
Verbose: logLevel,
|
HomeDir: HomeDir,
|
||||||
LogPath: logPath,
|
UpstreamProto: UpstreamProto,
|
||||||
|
Verbose: logLevel,
|
||||||
|
LogPath: logPath,
|
||||||
}
|
}
|
||||||
appCallback := mapCallback(c.AppCallback)
|
appCallback := mapCallback(c.AppCallback)
|
||||||
cli.RunMobile(&c.Config, &appCallback, c.stopCh)
|
cli.RunMobile(&c.Config, &appCallback, c.stopCh)
|
||||||
@@ -61,13 +63,18 @@ func mapCallback(callback AppCallback) cli.AppCallback {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Controller) Stop() bool {
|
func (c *Controller) Stop(restart bool, pin int64) int {
|
||||||
if c.stopCh != nil {
|
var errorCode = 0
|
||||||
|
// Force disconnect without checking pin.
|
||||||
|
// In iOS restart is required if vpn detects no connectivity after network change.
|
||||||
|
if !restart {
|
||||||
|
errorCode = cli.CheckDeactivationPin(pin, c.stopCh)
|
||||||
|
}
|
||||||
|
if errorCode == 0 && c.stopCh != nil {
|
||||||
close(c.stopCh)
|
close(c.stopCh)
|
||||||
c.stopCh = nil
|
c.stopCh = nil
|
||||||
return true
|
|
||||||
}
|
}
|
||||||
return false
|
return errorCode
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Controller) IsRunning() bool {
|
func (c *Controller) IsRunning() bool {
|
||||||
|
|||||||
@@ -7,8 +7,8 @@ import (
|
|||||||
"crypto/x509"
|
"crypto/x509"
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"errors"
|
"errors"
|
||||||
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"math/rand"
|
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
@@ -22,9 +22,11 @@ import (
|
|||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/ameshkov/dnsstamps"
|
||||||
"github.com/go-playground/validator/v10"
|
"github.com/go-playground/validator/v10"
|
||||||
"github.com/miekg/dns"
|
"github.com/miekg/dns"
|
||||||
"github.com/spf13/viper"
|
"github.com/spf13/viper"
|
||||||
|
"golang.org/x/net/http2"
|
||||||
"golang.org/x/sync/singleflight"
|
"golang.org/x/sync/singleflight"
|
||||||
"tailscale.com/logtail/backoff"
|
"tailscale.com/logtail/backoff"
|
||||||
"tailscale.com/net/tsaddr"
|
"tailscale.com/net/tsaddr"
|
||||||
@@ -46,9 +48,44 @@ const (
|
|||||||
// depending on the record type of the DNS query.
|
// depending on the record type of the DNS query.
|
||||||
IpStackSplit = "split"
|
IpStackSplit = "split"
|
||||||
|
|
||||||
|
// FreeDnsDomain is the domain name of free ControlD service.
|
||||||
|
FreeDnsDomain = "freedns.controld.com"
|
||||||
|
// FreeDNSBoostrapIP is the IP address of freedns.controld.com.
|
||||||
|
FreeDNSBoostrapIP = "76.76.2.11"
|
||||||
|
// FreeDNSBoostrapIPv6 is the IPv6 address of freedns.controld.com.
|
||||||
|
FreeDNSBoostrapIPv6 = "2606:1a40::11"
|
||||||
|
// PremiumDnsDomain is the domain name of premium ControlD service.
|
||||||
|
PremiumDnsDomain = "dns.controld.com"
|
||||||
|
// PremiumDNSBoostrapIP is the IP address of dns.controld.com.
|
||||||
|
PremiumDNSBoostrapIP = "76.76.2.22"
|
||||||
|
// PremiumDNSBoostrapIPv6 is the IPv6 address of dns.controld.com.
|
||||||
|
PremiumDNSBoostrapIPv6 = "2606:1a40::22"
|
||||||
|
|
||||||
|
// freeDnsDomainDev is the domain name of free ControlD service on dev env.
|
||||||
|
freeDnsDomainDev = "freedns.controld.dev"
|
||||||
|
// freeDNSBoostrapIP is the IP address of freedns.controld.dev.
|
||||||
|
freeDNSBoostrapIP = "176.125.239.11"
|
||||||
|
// freeDNSBoostrapIPv6 is the IPv6 address of freedns.controld.com.
|
||||||
|
freeDNSBoostrapIPv6 = "2606:1a40:f000::11"
|
||||||
|
// premiumDnsDomainDev is the domain name of premium ControlD service on dev env.
|
||||||
|
premiumDnsDomainDev = "dns.controld.dev"
|
||||||
|
// premiumDNSBoostrapIP is the IP address of dns.controld.dev.
|
||||||
|
premiumDNSBoostrapIP = "176.125.239.22"
|
||||||
|
// premiumDNSBoostrapIPv6 is the IPv6 address of dns.controld.dev.
|
||||||
|
premiumDNSBoostrapIPv6 = "2606:1a40:f000::22"
|
||||||
|
|
||||||
controlDComDomain = "controld.com"
|
controlDComDomain = "controld.com"
|
||||||
controlDNetDomain = "controld.net"
|
controlDNetDomain = "controld.net"
|
||||||
controlDDevDomain = "controld.dev"
|
controlDDevDomain = "controld.dev"
|
||||||
|
|
||||||
|
endpointPrefixHTTPS = "https://"
|
||||||
|
endpointPrefixQUIC = "quic://"
|
||||||
|
endpointPrefixH3 = "h3://"
|
||||||
|
endpointPrefixSdns = "sdns://"
|
||||||
|
|
||||||
|
rebootstrapNotStarted = 0
|
||||||
|
rebootstrapStarted = 1
|
||||||
|
rebootstrapInProgress = 2
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
@@ -104,14 +141,14 @@ func InitConfig(v *viper.Viper, name string) {
|
|||||||
})
|
})
|
||||||
v.SetDefault("upstream", map[string]*UpstreamConfig{
|
v.SetDefault("upstream", map[string]*UpstreamConfig{
|
||||||
"0": {
|
"0": {
|
||||||
BootstrapIP: "76.76.2.11",
|
BootstrapIP: FreeDNSBoostrapIP,
|
||||||
Name: "Control D - Anti-Malware",
|
Name: "Control D - Anti-Malware",
|
||||||
Type: ResolverTypeDOH,
|
Type: ResolverTypeDOH,
|
||||||
Endpoint: "https://freedns.controld.com/p1",
|
Endpoint: "https://freedns.controld.com/p1",
|
||||||
Timeout: 5000,
|
Timeout: 5000,
|
||||||
},
|
},
|
||||||
"1": {
|
"1": {
|
||||||
BootstrapIP: "76.76.2.11",
|
BootstrapIP: FreeDNSBoostrapIP,
|
||||||
Name: "Control D - No Ads",
|
Name: "Control D - No Ads",
|
||||||
Type: ResolverTypeDOQ,
|
Type: ResolverTypeDOQ,
|
||||||
Endpoint: "p2.freedns.controld.com",
|
Endpoint: "p2.freedns.controld.com",
|
||||||
@@ -179,26 +216,35 @@ func (c *Config) FirstUpstream() *UpstreamConfig {
|
|||||||
|
|
||||||
// ServiceConfig specifies the general ctrld config.
|
// ServiceConfig specifies the general ctrld config.
|
||||||
type ServiceConfig struct {
|
type ServiceConfig struct {
|
||||||
LogLevel string `mapstructure:"log_level" toml:"log_level,omitempty"`
|
LogLevel string `mapstructure:"log_level" toml:"log_level,omitempty"`
|
||||||
LogPath string `mapstructure:"log_path" toml:"log_path,omitempty"`
|
LogPath string `mapstructure:"log_path" toml:"log_path,omitempty"`
|
||||||
CacheEnable bool `mapstructure:"cache_enable" toml:"cache_enable,omitempty"`
|
CacheEnable bool `mapstructure:"cache_enable" toml:"cache_enable,omitempty"`
|
||||||
CacheSize int `mapstructure:"cache_size" toml:"cache_size,omitempty"`
|
CacheSize int `mapstructure:"cache_size" toml:"cache_size,omitempty"`
|
||||||
CacheTTLOverride int `mapstructure:"cache_ttl_override" toml:"cache_ttl_override,omitempty"`
|
CacheTTLOverride int `mapstructure:"cache_ttl_override" toml:"cache_ttl_override,omitempty"`
|
||||||
CacheServeStale bool `mapstructure:"cache_serve_stale" toml:"cache_serve_stale,omitempty"`
|
CacheServeStale bool `mapstructure:"cache_serve_stale" toml:"cache_serve_stale,omitempty"`
|
||||||
MaxConcurrentRequests *int `mapstructure:"max_concurrent_requests" toml:"max_concurrent_requests,omitempty" validate:"omitempty,gte=0"`
|
CacheFlushDomains []string `mapstructure:"cache_flush_domains" toml:"cache_flush_domains" validate:"max=256"`
|
||||||
DHCPLeaseFile string `mapstructure:"dhcp_lease_file_path" toml:"dhcp_lease_file_path" validate:"omitempty,file"`
|
MaxConcurrentRequests *int `mapstructure:"max_concurrent_requests" toml:"max_concurrent_requests,omitempty" validate:"omitempty,gte=0"`
|
||||||
DHCPLeaseFileFormat string `mapstructure:"dhcp_lease_file_format" toml:"dhcp_lease_file_format" validate:"required_unless=DHCPLeaseFile '',omitempty,oneof=dnsmasq isc-dhcp"`
|
DHCPLeaseFile string `mapstructure:"dhcp_lease_file_path" toml:"dhcp_lease_file_path" validate:"omitempty,file"`
|
||||||
DiscoverMDNS *bool `mapstructure:"discover_mdns" toml:"discover_mdns,omitempty"`
|
DHCPLeaseFileFormat string `mapstructure:"dhcp_lease_file_format" toml:"dhcp_lease_file_format" validate:"required_unless=DHCPLeaseFile '',omitempty,oneof=dnsmasq isc-dhcp kea-dhcp4"`
|
||||||
DiscoverARP *bool `mapstructure:"discover_arp" toml:"discover_arp,omitempty"`
|
DiscoverMDNS *bool `mapstructure:"discover_mdns" toml:"discover_mdns,omitempty"`
|
||||||
DiscoverDHCP *bool `mapstructure:"discover_dhcp" toml:"discover_dhcp,omitempty"`
|
DiscoverARP *bool `mapstructure:"discover_arp" toml:"discover_arp,omitempty"`
|
||||||
DiscoverPtr *bool `mapstructure:"discover_ptr" toml:"discover_ptr,omitempty"`
|
DiscoverDHCP *bool `mapstructure:"discover_dhcp" toml:"discover_dhcp,omitempty"`
|
||||||
DiscoverHosts *bool `mapstructure:"discover_hosts" toml:"discover_hosts,omitempty"`
|
DiscoverPtr *bool `mapstructure:"discover_ptr" toml:"discover_ptr,omitempty"`
|
||||||
DiscoverRefreshInterval int `mapstructure:"discover_refresh_interval" toml:"discover_refresh_interval,omitempty"`
|
DiscoverHosts *bool `mapstructure:"discover_hosts" toml:"discover_hosts,omitempty"`
|
||||||
ClientIDPref string `mapstructure:"client_id_preference" toml:"client_id_preference,omitempty" validate:"omitempty,oneof=host mac"`
|
DiscoverRefreshInterval int `mapstructure:"discover_refresh_interval" toml:"discover_refresh_interval,omitempty"`
|
||||||
MetricsQueryStats bool `mapstructure:"metrics_query_stats" toml:"metrics_query_stats,omitempty"`
|
ClientIDPref string `mapstructure:"client_id_preference" toml:"client_id_preference,omitempty" validate:"omitempty,oneof=host mac"`
|
||||||
MetricsListener string `mapstructure:"metrics_listener" toml:"metrics_listener,omitempty"`
|
MetricsQueryStats bool `mapstructure:"metrics_query_stats" toml:"metrics_query_stats,omitempty"`
|
||||||
Daemon bool `mapstructure:"-" toml:"-"`
|
MetricsListener string `mapstructure:"metrics_listener" toml:"metrics_listener,omitempty"`
|
||||||
AllocateIP bool `mapstructure:"-" toml:"-"`
|
DnsWatchdogEnabled *bool `mapstructure:"dns_watchdog_enabled" toml:"dns_watchdog_enabled,omitempty"`
|
||||||
|
DnsWatchdogInvterval *time.Duration `mapstructure:"dns_watchdog_interval" toml:"dns_watchdog_interval,omitempty"`
|
||||||
|
RefetchTime *int `mapstructure:"refetch_time" toml:"refetch_time,omitempty"`
|
||||||
|
ForceRefetchWaitTime *int `mapstructure:"force_refetch_wait_time" toml:"force_refetch_wait_time,omitempty"`
|
||||||
|
LeakOnUpstreamFailure *bool `mapstructure:"leak_on_upstream_failure" toml:"leak_on_upstream_failure,omitempty"`
|
||||||
|
InterceptMode string `mapstructure:"intercept_mode" toml:"intercept_mode,omitempty" validate:"omitempty,oneof=off dns hard"`
|
||||||
|
NRPTRecoveryMaxAttempts *int `mapstructure:"nrpt_recovery_max_attempts" toml:"nrpt_recovery_max_attempts,omitempty" validate:"omitempty,gte=0"`
|
||||||
|
NRPTRecoveryCooldown *time.Duration `mapstructure:"nrpt_recovery_cooldown" toml:"nrpt_recovery_cooldown,omitempty"`
|
||||||
|
Daemon bool `mapstructure:"-" toml:"-"`
|
||||||
|
AllocateIP bool `mapstructure:"-" toml:"-"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// NetworkConfig specifies configuration for networks where ctrld will handle requests.
|
// NetworkConfig specifies configuration for networks where ctrld will handle requests.
|
||||||
@@ -211,7 +257,7 @@ type NetworkConfig struct {
|
|||||||
// UpstreamConfig specifies configuration for upstreams that ctrld will forward requests to.
|
// UpstreamConfig specifies configuration for upstreams that ctrld will forward requests to.
|
||||||
type UpstreamConfig struct {
|
type UpstreamConfig struct {
|
||||||
Name string `mapstructure:"name" toml:"name,omitempty"`
|
Name string `mapstructure:"name" toml:"name,omitempty"`
|
||||||
Type string `mapstructure:"type" toml:"type,omitempty" validate:"oneof=doh doh3 dot doq os legacy"`
|
Type string `mapstructure:"type" toml:"type,omitempty" validate:"oneof=doh doh3 dot doq os legacy sdns ''"`
|
||||||
Endpoint string `mapstructure:"endpoint" toml:"endpoint,omitempty"`
|
Endpoint string `mapstructure:"endpoint" toml:"endpoint,omitempty"`
|
||||||
BootstrapIP string `mapstructure:"bootstrap_ip" toml:"bootstrap_ip,omitempty"`
|
BootstrapIP string `mapstructure:"bootstrap_ip" toml:"bootstrap_ip,omitempty"`
|
||||||
Domain string `mapstructure:"-" toml:"-"`
|
Domain string `mapstructure:"-" toml:"-"`
|
||||||
@@ -225,7 +271,7 @@ type UpstreamConfig struct {
|
|||||||
Discoverable *bool `mapstructure:"discoverable" toml:"discoverable"`
|
Discoverable *bool `mapstructure:"discoverable" toml:"discoverable"`
|
||||||
|
|
||||||
g singleflight.Group
|
g singleflight.Group
|
||||||
rebootstrap atomic.Bool
|
rebootstrap atomic.Int64
|
||||||
bootstrapIPs []string
|
bootstrapIPs []string
|
||||||
bootstrapIPs4 []string
|
bootstrapIPs4 []string
|
||||||
bootstrapIPs6 []string
|
bootstrapIPs6 []string
|
||||||
@@ -236,8 +282,15 @@ type UpstreamConfig struct {
|
|||||||
http3RoundTripper http.RoundTripper
|
http3RoundTripper http.RoundTripper
|
||||||
http3RoundTripper4 http.RoundTripper
|
http3RoundTripper4 http.RoundTripper
|
||||||
http3RoundTripper6 http.RoundTripper
|
http3RoundTripper6 http.RoundTripper
|
||||||
|
doqConnPool *doqConnPool
|
||||||
|
doqConnPool4 *doqConnPool
|
||||||
|
doqConnPool6 *doqConnPool
|
||||||
|
dotClientPool *dotConnPool
|
||||||
|
dotClientPool4 *dotConnPool
|
||||||
|
dotClientPool6 *dotConnPool
|
||||||
certPool *x509.CertPool
|
certPool *x509.CertPool
|
||||||
u *url.URL
|
u *url.URL
|
||||||
|
fallbackOnce sync.Once
|
||||||
uid string
|
uid string
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -285,9 +338,13 @@ type Rule map[string][]string
|
|||||||
|
|
||||||
// Init initialized necessary values for an UpstreamConfig.
|
// Init initialized necessary values for an UpstreamConfig.
|
||||||
func (uc *UpstreamConfig) Init() {
|
func (uc *UpstreamConfig) Init() {
|
||||||
|
if err := uc.initDnsStamps(); err != nil {
|
||||||
|
ProxyLogger.Load().Fatal().Err(err).Msg("invalid DNS Stamps")
|
||||||
|
}
|
||||||
|
uc.initDoHScheme()
|
||||||
uc.uid = upstreamUID()
|
uc.uid = upstreamUID()
|
||||||
if u, err := url.Parse(uc.Endpoint); err == nil {
|
if u, err := url.Parse(uc.Endpoint); err == nil {
|
||||||
uc.Domain = u.Host
|
uc.Domain = u.Hostname()
|
||||||
switch uc.Type {
|
switch uc.Type {
|
||||||
case ResolverTypeDOH, ResolverTypeDOH3:
|
case ResolverTypeDOH, ResolverTypeDOH3:
|
||||||
uc.u = u
|
uc.u = u
|
||||||
@@ -305,7 +362,7 @@ func (uc *UpstreamConfig) Init() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if uc.IPStack == "" {
|
if uc.IPStack == "" {
|
||||||
if uc.isControlD() {
|
if uc.IsControlD() {
|
||||||
uc.IPStack = IpStackSplit
|
uc.IPStack = IpStackSplit
|
||||||
} else {
|
} else {
|
||||||
uc.IPStack = IpStackBoth
|
uc.IPStack = IpStackBoth
|
||||||
@@ -313,6 +370,15 @@ func (uc *UpstreamConfig) Init() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// VerifyMsg creates and returns a new DNS message could be used for testing upstream health.
|
||||||
|
func (uc *UpstreamConfig) VerifyMsg() *dns.Msg {
|
||||||
|
msg := new(dns.Msg)
|
||||||
|
msg.RecursionDesired = true
|
||||||
|
msg.SetQuestion(".", dns.TypeNS)
|
||||||
|
msg.SetEdns0(4096, false) // ensure handling of large DNS response
|
||||||
|
return msg
|
||||||
|
}
|
||||||
|
|
||||||
// VerifyDomain returns the domain name that could be resolved by the upstream endpoint.
|
// VerifyDomain returns the domain name that could be resolved by the upstream endpoint.
|
||||||
// It returns empty for non-ControlD upstream endpoint.
|
// It returns empty for non-ControlD upstream endpoint.
|
||||||
func (uc *UpstreamConfig) VerifyDomain() string {
|
func (uc *UpstreamConfig) VerifyDomain() string {
|
||||||
@@ -343,7 +409,7 @@ func (uc *UpstreamConfig) UpstreamSendClientInfo() bool {
|
|||||||
}
|
}
|
||||||
switch uc.Type {
|
switch uc.Type {
|
||||||
case ResolverTypeDOH, ResolverTypeDOH3:
|
case ResolverTypeDOH, ResolverTypeDOH3:
|
||||||
if uc.isControlD() || uc.isNextDNS() {
|
if uc.IsControlD() || uc.isNextDNS() {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -357,7 +423,7 @@ func (uc *UpstreamConfig) IsDiscoverable() bool {
|
|||||||
return *uc.Discoverable
|
return *uc.Discoverable
|
||||||
}
|
}
|
||||||
switch uc.Type {
|
switch uc.Type {
|
||||||
case ResolverTypeOS, ResolverTypeLegacy, ResolverTypePrivate:
|
case ResolverTypeOS, ResolverTypeLegacy, ResolverTypePrivate, ResolverTypeLocal:
|
||||||
if ip, err := netip.ParseAddr(uc.Domain); err == nil {
|
if ip, err := netip.ParseAddr(uc.Domain); err == nil {
|
||||||
return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || tsaddr.CGNATRange().Contains(ip)
|
return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || tsaddr.CGNATRange().Contains(ip)
|
||||||
}
|
}
|
||||||
@@ -375,12 +441,6 @@ func (uc *UpstreamConfig) SetCertPool(cp *x509.CertPool) {
|
|||||||
uc.certPool = cp
|
uc.certPool = cp
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetupBootstrapIP manually find all available IPs of the upstream.
|
|
||||||
// The first usable IP will be used as bootstrap IP of the upstream.
|
|
||||||
func (uc *UpstreamConfig) SetupBootstrapIP() {
|
|
||||||
uc.setupBootstrapIP(true)
|
|
||||||
}
|
|
||||||
|
|
||||||
// UID returns the unique identifier of the upstream.
|
// UID returns the unique identifier of the upstream.
|
||||||
func (uc *UpstreamConfig) UID() string {
|
func (uc *UpstreamConfig) UID() string {
|
||||||
return uc.uid
|
return uc.uid
|
||||||
@@ -388,11 +448,19 @@ func (uc *UpstreamConfig) UID() string {
|
|||||||
|
|
||||||
// SetupBootstrapIP manually find all available IPs of the upstream.
|
// SetupBootstrapIP manually find all available IPs of the upstream.
|
||||||
// The first usable IP will be used as bootstrap IP of the upstream.
|
// The first usable IP will be used as bootstrap IP of the upstream.
|
||||||
func (uc *UpstreamConfig) setupBootstrapIP(withBootstrapDNS bool) {
|
// The upstream domain will be looked up using following orders:
|
||||||
|
//
|
||||||
|
// - Current system DNS settings.
|
||||||
|
// - Direct IPs table for ControlD upstreams.
|
||||||
|
// - ControlD Bootstrap DNS 76.76.2.22
|
||||||
|
//
|
||||||
|
// The setup process will block until there's usable IPs found.
|
||||||
|
func (uc *UpstreamConfig) SetupBootstrapIP() {
|
||||||
b := backoff.NewBackoff("setupBootstrapIP", func(format string, args ...any) {}, 10*time.Second)
|
b := backoff.NewBackoff("setupBootstrapIP", func(format string, args ...any) {}, 10*time.Second)
|
||||||
isControlD := uc.isControlD()
|
isControlD := uc.IsControlD()
|
||||||
|
nss := initDefaultOsResolver()
|
||||||
for {
|
for {
|
||||||
uc.bootstrapIPs = lookupIP(uc.Domain, uc.Timeout, withBootstrapDNS)
|
uc.bootstrapIPs = lookupIP(uc.Domain, uc.Timeout, nss)
|
||||||
// For ControlD upstream, the bootstrap IPs could not be RFC 1918 addresses,
|
// For ControlD upstream, the bootstrap IPs could not be RFC 1918 addresses,
|
||||||
// filtering them out here to prevent weird behavior.
|
// filtering them out here to prevent weird behavior.
|
||||||
if isControlD {
|
if isControlD {
|
||||||
@@ -405,6 +473,15 @@ func (uc *UpstreamConfig) setupBootstrapIP(withBootstrapDNS bool) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
uc.bootstrapIPs = uc.bootstrapIPs[:n]
|
uc.bootstrapIPs = uc.bootstrapIPs[:n]
|
||||||
|
if len(uc.bootstrapIPs) == 0 {
|
||||||
|
uc.bootstrapIPs = bootstrapIPsFromControlDDomain(uc.Domain)
|
||||||
|
ProxyLogger.Load().Warn().Msgf("no record found for %q, lookup from direct IP table", uc.Domain)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(uc.bootstrapIPs) == 0 {
|
||||||
|
ProxyLogger.Load().Warn().Msgf("no record found for %q, using bootstrap server: %s", uc.Domain, PremiumDNSBoostrapIP)
|
||||||
|
uc.bootstrapIPs = lookupIP(uc.Domain, uc.Timeout, []string{net.JoinHostPort(PremiumDNSBoostrapIP, "53")})
|
||||||
|
|
||||||
}
|
}
|
||||||
if len(uc.bootstrapIPs) > 0 {
|
if len(uc.bootstrapIPs) > 0 {
|
||||||
break
|
break
|
||||||
@@ -425,54 +502,154 @@ func (uc *UpstreamConfig) setupBootstrapIP(withBootstrapDNS bool) {
|
|||||||
// ReBootstrap re-setup the bootstrap IP and the transport.
|
// ReBootstrap re-setup the bootstrap IP and the transport.
|
||||||
func (uc *UpstreamConfig) ReBootstrap() {
|
func (uc *UpstreamConfig) ReBootstrap() {
|
||||||
switch uc.Type {
|
switch uc.Type {
|
||||||
case ResolverTypeDOH, ResolverTypeDOH3:
|
case ResolverTypeDOH, ResolverTypeDOH3, ResolverTypeDOQ, ResolverTypeDOT:
|
||||||
default:
|
default:
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
_, _, _ = uc.g.Do("ReBootstrap", func() (any, error) {
|
_, _, _ = uc.g.Do("ReBootstrap", func() (any, error) {
|
||||||
if uc.rebootstrap.CompareAndSwap(false, true) {
|
if uc.rebootstrap.CompareAndSwap(rebootstrapNotStarted, rebootstrapStarted) {
|
||||||
ProxyLogger.Load().Debug().Msg("re-bootstrapping upstream ip")
|
ProxyLogger.Load().Debug().Msgf("re-bootstrapping upstream ip for %v", uc)
|
||||||
}
|
}
|
||||||
return true, nil
|
return true, nil
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetupTransport initializes the network transport used to connect to upstream server.
|
// ForceReBootstrap immediately replaces the upstream transport, closing old
|
||||||
// For now, only DoH upstream is supported.
|
// connections and creating new ones synchronously. Unlike ReBootstrap() which
|
||||||
func (uc *UpstreamConfig) SetupTransport() {
|
// sets a lazy flag (new transport created on next query), this ensures the
|
||||||
|
// transport is ready before any queries arrive. Use when external events
|
||||||
|
// (e.g. firewall state flush) are known to have killed existing connections.
|
||||||
|
func (uc *UpstreamConfig) ForceReBootstrap() {
|
||||||
switch uc.Type {
|
switch uc.Type {
|
||||||
case ResolverTypeDOH:
|
case ResolverTypeDOH, ResolverTypeDOH3, ResolverTypeDOQ, ResolverTypeDOT:
|
||||||
uc.setupDOHTransport()
|
default:
|
||||||
case ResolverTypeDOH3:
|
return
|
||||||
uc.setupDOH3Transport()
|
}
|
||||||
|
ProxyLogger.Load().Debug().Msgf("force re-bootstrapping upstream transport for %v", uc)
|
||||||
|
uc.SetupTransport()
|
||||||
|
// Clear any pending lazy re-bootstrap flag so ensureSetupTransport()
|
||||||
|
// doesn't redundantly recreate the transport we just built.
|
||||||
|
uc.rebootstrap.Store(rebootstrapNotStarted)
|
||||||
|
}
|
||||||
|
|
||||||
|
// closeTransports closes idle connections on all existing transports.
|
||||||
|
// This is called before creating new transports during re-bootstrap to
|
||||||
|
// force in-flight requests on stale connections to fail quickly, rather
|
||||||
|
// than waiting for the full context deadline (e.g. 5s) after a firewall
|
||||||
|
// state table flush kills the underlying TCP/QUIC connections.
|
||||||
|
func (uc *UpstreamConfig) closeTransports() {
|
||||||
|
if t := uc.transport; t != nil {
|
||||||
|
t.CloseIdleConnections()
|
||||||
|
}
|
||||||
|
if t := uc.transport4; t != nil {
|
||||||
|
t.CloseIdleConnections()
|
||||||
|
}
|
||||||
|
if t := uc.transport6; t != nil {
|
||||||
|
t.CloseIdleConnections()
|
||||||
|
}
|
||||||
|
if p := uc.doqConnPool; p != nil {
|
||||||
|
p.CloseIdleConnections()
|
||||||
|
}
|
||||||
|
if p := uc.doqConnPool4; p != nil {
|
||||||
|
p.CloseIdleConnections()
|
||||||
|
}
|
||||||
|
if p := uc.doqConnPool6; p != nil {
|
||||||
|
p.CloseIdleConnections()
|
||||||
|
}
|
||||||
|
if p := uc.dotClientPool; p != nil {
|
||||||
|
p.CloseIdleConnections()
|
||||||
|
}
|
||||||
|
if p := uc.dotClientPool4; p != nil {
|
||||||
|
p.CloseIdleConnections()
|
||||||
|
}
|
||||||
|
if p := uc.dotClientPool6; p != nil {
|
||||||
|
p.CloseIdleConnections()
|
||||||
|
}
|
||||||
|
// http3RoundTripper is stored as http.RoundTripper but the concrete type
|
||||||
|
// (*http3.Transport) exposes CloseIdleConnections via this interface.
|
||||||
|
type idleCloser interface {
|
||||||
|
CloseIdleConnections()
|
||||||
|
}
|
||||||
|
for _, rt := range []http.RoundTripper{uc.http3RoundTripper, uc.http3RoundTripper4, uc.http3RoundTripper6} {
|
||||||
|
if c, ok := rt.(idleCloser); ok {
|
||||||
|
c.CloseIdleConnections()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) setupDOHTransport() {
|
// SetupTransport initializes the network transport used to connect to upstream servers.
|
||||||
|
// For now, DoH/DoH3/DoQ/DoT upstreams are supported.
|
||||||
|
func (uc *UpstreamConfig) SetupTransport() {
|
||||||
|
switch uc.Type {
|
||||||
|
case ResolverTypeDOH, ResolverTypeDOH3, ResolverTypeDOQ, ResolverTypeDOT:
|
||||||
|
default:
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close existing transport connections before creating new ones.
|
||||||
|
// This forces in-flight requests on stale connections (e.g. after a
|
||||||
|
// firewall state table flush) to fail fast instead of waiting for
|
||||||
|
// the full context deadline timeout.
|
||||||
|
uc.closeTransports()
|
||||||
|
|
||||||
|
ips := uc.bootstrapIPs
|
||||||
switch uc.IPStack {
|
switch uc.IPStack {
|
||||||
case IpStackBoth, "":
|
|
||||||
uc.transport = uc.newDOHTransport(uc.bootstrapIPs)
|
|
||||||
case IpStackV4:
|
case IpStackV4:
|
||||||
uc.transport = uc.newDOHTransport(uc.bootstrapIPs4)
|
ips = uc.bootstrapIPs4
|
||||||
case IpStackV6:
|
case IpStackV6:
|
||||||
uc.transport = uc.newDOHTransport(uc.bootstrapIPs6)
|
ips = uc.bootstrapIPs6
|
||||||
case IpStackSplit:
|
}
|
||||||
|
|
||||||
|
uc.transport = uc.newDOHTransport(ips)
|
||||||
|
uc.http3RoundTripper = uc.newDOH3Transport(ips)
|
||||||
|
uc.doqConnPool = uc.newDOQConnPool(ips)
|
||||||
|
uc.dotClientPool = uc.newDOTClientPool(ips)
|
||||||
|
if uc.IPStack == IpStackSplit {
|
||||||
uc.transport4 = uc.newDOHTransport(uc.bootstrapIPs4)
|
uc.transport4 = uc.newDOHTransport(uc.bootstrapIPs4)
|
||||||
if hasIPv6() {
|
uc.http3RoundTripper4 = uc.newDOH3Transport(uc.bootstrapIPs4)
|
||||||
|
uc.doqConnPool4 = uc.newDOQConnPool(uc.bootstrapIPs4)
|
||||||
|
uc.dotClientPool4 = uc.newDOTClientPool(uc.bootstrapIPs4)
|
||||||
|
if HasIPv6() {
|
||||||
uc.transport6 = uc.newDOHTransport(uc.bootstrapIPs6)
|
uc.transport6 = uc.newDOHTransport(uc.bootstrapIPs6)
|
||||||
|
uc.http3RoundTripper6 = uc.newDOH3Transport(uc.bootstrapIPs6)
|
||||||
|
uc.doqConnPool6 = uc.newDOQConnPool(uc.bootstrapIPs6)
|
||||||
|
uc.dotClientPool6 = uc.newDOTClientPool(uc.bootstrapIPs6)
|
||||||
} else {
|
} else {
|
||||||
uc.transport6 = uc.transport4
|
uc.transport6 = uc.transport4
|
||||||
|
uc.http3RoundTripper6 = uc.http3RoundTripper4
|
||||||
|
uc.doqConnPool6 = uc.doqConnPool4
|
||||||
|
uc.dotClientPool6 = uc.dotClientPool4
|
||||||
}
|
}
|
||||||
uc.transport = uc.newDOHTransport(uc.bootstrapIPs)
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (uc *UpstreamConfig) ensureSetupTransport() {
|
||||||
|
uc.transportOnce.Do(func() {
|
||||||
|
uc.SetupTransport()
|
||||||
|
})
|
||||||
|
if uc.rebootstrap.CompareAndSwap(rebootstrapStarted, rebootstrapInProgress) {
|
||||||
|
uc.SetupTransport()
|
||||||
|
uc.rebootstrap.Store(rebootstrapNotStarted)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) newDOHTransport(addrs []string) *http.Transport {
|
func (uc *UpstreamConfig) newDOHTransport(addrs []string) *http.Transport {
|
||||||
|
if uc.Type != ResolverTypeDOH {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
transport := http.DefaultTransport.(*http.Transport).Clone()
|
transport := http.DefaultTransport.(*http.Transport).Clone()
|
||||||
transport.MaxIdleConnsPerHost = 100
|
transport.MaxIdleConnsPerHost = 100
|
||||||
transport.TLSClientConfig = &tls.Config{
|
transport.TLSClientConfig = &tls.Config{
|
||||||
RootCAs: uc.certPool,
|
RootCAs: uc.certPool,
|
||||||
ClientSessionCache: tls.NewLRUClientSessionCache(0),
|
ClientSessionCache: tls.NewLRUClientSessionCache(0),
|
||||||
|
MinVersion: tls.VersionTLS12,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Prevent bad tcp connection hanging the requests for too long.
|
||||||
|
// See: https://github.com/golang/go/issues/36026
|
||||||
|
if t2, err := http2.ConfigureTransports(transport); err == nil {
|
||||||
|
t2.ReadIdleTimeout = 10 * time.Second
|
||||||
|
t2.PingTimeout = 5 * time.Second
|
||||||
}
|
}
|
||||||
|
|
||||||
dialerTimeoutMs := 2000
|
dialerTimeoutMs := 2000
|
||||||
@@ -495,7 +672,7 @@ func (uc *UpstreamConfig) newDOHTransport(addrs []string) *http.Transport {
|
|||||||
for i := range addrs {
|
for i := range addrs {
|
||||||
dialAddrs[i] = net.JoinHostPort(addrs[i], port)
|
dialAddrs[i] = net.JoinHostPort(addrs[i], port)
|
||||||
}
|
}
|
||||||
conn, err := pd.DialContext(ctx, network, dialAddrs)
|
conn, err := pd.DialContext(ctx, network, dialAddrs, ProxyLogger.Load())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -510,38 +687,69 @@ func (uc *UpstreamConfig) newDOHTransport(addrs []string) *http.Transport {
|
|||||||
|
|
||||||
// Ping warms up the connection to DoH/DoH3 upstream.
|
// Ping warms up the connection to DoH/DoH3 upstream.
|
||||||
func (uc *UpstreamConfig) Ping() {
|
func (uc *UpstreamConfig) Ping() {
|
||||||
|
if err := uc.ping(); err != nil {
|
||||||
|
ProxyLogger.Load().Debug().Err(err).Msgf("upstream ping failed: %s", uc.Endpoint)
|
||||||
|
_ = uc.FallbackToDirectIP()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ErrorPing is like Ping, but return an error if any.
|
||||||
|
func (uc *UpstreamConfig) ErrorPing() error {
|
||||||
|
return uc.ping()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (uc *UpstreamConfig) ping() error {
|
||||||
switch uc.Type {
|
switch uc.Type {
|
||||||
case ResolverTypeDOH, ResolverTypeDOH3:
|
case ResolverTypeDOH, ResolverTypeDOH3, ResolverTypeDOQ:
|
||||||
default:
|
default:
|
||||||
return
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
ping := func(t http.RoundTripper) {
|
ping := func(t http.RoundTripper) error {
|
||||||
if t == nil {
|
if t == nil {
|
||||||
return
|
return nil
|
||||||
}
|
}
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
req, _ := http.NewRequestWithContext(ctx, "HEAD", uc.Endpoint, nil)
|
req, err := http.NewRequestWithContext(ctx, "HEAD", uc.Endpoint, nil)
|
||||||
resp, _ := t.RoundTrip(req)
|
if err != nil {
|
||||||
if resp == nil {
|
return err
|
||||||
return
|
}
|
||||||
|
resp, err := t.RoundTrip(req)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
defer resp.Body.Close()
|
defer resp.Body.Close()
|
||||||
_, _ = io.Copy(io.Discard, resp.Body)
|
_, _ = io.Copy(io.Discard, resp.Body)
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, typ := range []uint16{dns.TypeA, dns.TypeAAAA} {
|
for _, typ := range []uint16{dns.TypeA, dns.TypeAAAA} {
|
||||||
switch uc.Type {
|
switch uc.Type {
|
||||||
case ResolverTypeDOH:
|
case ResolverTypeDOH:
|
||||||
ping(uc.dohTransport(typ))
|
if err := ping(uc.dohTransport(typ)); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
case ResolverTypeDOH3:
|
case ResolverTypeDOH3:
|
||||||
ping(uc.doh3Transport(typ))
|
if err := ping(uc.doh3Transport(typ)); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
case ResolverTypeDOQ:
|
||||||
|
// For DoQ, we just ensure transport is set up by calling doqTransport
|
||||||
|
// DoQ doesn't use HTTP, so we can't ping it the same way
|
||||||
|
_ = uc.doqTransport(typ)
|
||||||
|
case ResolverTypeDOT:
|
||||||
|
// For DoT, we just ensure transport is set up by calling dotTransport
|
||||||
|
// DoT doesn't use HTTP, so we can't ping it the same way
|
||||||
|
_ = uc.dotTransport(typ)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) isControlD() bool {
|
// IsControlD reports whether this is a ControlD upstream.
|
||||||
|
func (uc *UpstreamConfig) IsControlD() bool {
|
||||||
domain := uc.Domain
|
domain := uc.Domain
|
||||||
if domain == "" {
|
if domain == "" {
|
||||||
if u, err := url.Parse(uc.Endpoint); err == nil {
|
if u, err := url.Parse(uc.Endpoint); err == nil {
|
||||||
@@ -567,46 +775,8 @@ func (uc *UpstreamConfig) isNextDNS() bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) dohTransport(dnsType uint16) http.RoundTripper {
|
func (uc *UpstreamConfig) dohTransport(dnsType uint16) http.RoundTripper {
|
||||||
uc.transportOnce.Do(func() {
|
uc.ensureSetupTransport()
|
||||||
uc.SetupTransport()
|
return transportByIpStack(uc.IPStack, dnsType, uc.transport, uc.transport4, uc.transport6)
|
||||||
})
|
|
||||||
if uc.rebootstrap.CompareAndSwap(true, false) {
|
|
||||||
uc.SetupTransport()
|
|
||||||
}
|
|
||||||
switch uc.IPStack {
|
|
||||||
case IpStackBoth, IpStackV4, IpStackV6:
|
|
||||||
return uc.transport
|
|
||||||
case IpStackSplit:
|
|
||||||
switch dnsType {
|
|
||||||
case dns.TypeA:
|
|
||||||
return uc.transport4
|
|
||||||
default:
|
|
||||||
return uc.transport6
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return uc.transport
|
|
||||||
}
|
|
||||||
|
|
||||||
func (uc *UpstreamConfig) bootstrapIPForDNSType(dnsType uint16) string {
|
|
||||||
switch uc.IPStack {
|
|
||||||
case IpStackBoth:
|
|
||||||
return pick(uc.bootstrapIPs)
|
|
||||||
case IpStackV4:
|
|
||||||
return pick(uc.bootstrapIPs4)
|
|
||||||
case IpStackV6:
|
|
||||||
return pick(uc.bootstrapIPs6)
|
|
||||||
case IpStackSplit:
|
|
||||||
switch dnsType {
|
|
||||||
case dns.TypeA:
|
|
||||||
return pick(uc.bootstrapIPs4)
|
|
||||||
default:
|
|
||||||
if hasIPv6() {
|
|
||||||
return pick(uc.bootstrapIPs6)
|
|
||||||
}
|
|
||||||
return pick(uc.bootstrapIPs4)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return pick(uc.bootstrapIPs)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) netForDNSType(dnsType uint16) (string, string) {
|
func (uc *UpstreamConfig) netForDNSType(dnsType uint16) (string, string) {
|
||||||
@@ -622,7 +792,7 @@ func (uc *UpstreamConfig) netForDNSType(dnsType uint16) (string, string) {
|
|||||||
case dns.TypeA:
|
case dns.TypeA:
|
||||||
return "tcp4-tls", "udp4"
|
return "tcp4-tls", "udp4"
|
||||||
default:
|
default:
|
||||||
if hasIPv6() {
|
if HasIPv6() {
|
||||||
return "tcp6-tls", "udp6"
|
return "tcp6-tls", "udp6"
|
||||||
}
|
}
|
||||||
return "tcp4-tls", "udp4"
|
return "tcp4-tls", "udp4"
|
||||||
@@ -631,6 +801,104 @@ func (uc *UpstreamConfig) netForDNSType(dnsType uint16) (string, string) {
|
|||||||
return "tcp-tls", "udp"
|
return "tcp-tls", "udp"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// initDoHScheme initializes the endpoint scheme for DoH/DoH3 upstream if not present.
|
||||||
|
func (uc *UpstreamConfig) initDoHScheme() {
|
||||||
|
if strings.HasPrefix(uc.Endpoint, endpointPrefixH3) && uc.Type == "" {
|
||||||
|
uc.Type = ResolverTypeDOH3
|
||||||
|
}
|
||||||
|
switch uc.Type {
|
||||||
|
case ResolverTypeDOH:
|
||||||
|
case ResolverTypeDOH3:
|
||||||
|
if after, found := strings.CutPrefix(uc.Endpoint, endpointPrefixH3); found {
|
||||||
|
uc.Endpoint = endpointPrefixHTTPS + after
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !strings.HasPrefix(uc.Endpoint, endpointPrefixHTTPS) {
|
||||||
|
uc.Endpoint = endpointPrefixHTTPS + uc.Endpoint
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// initDnsStamps initializes upstream config based on encoded DNS Stamps Endpoint.
|
||||||
|
func (uc *UpstreamConfig) initDnsStamps() error {
|
||||||
|
if strings.HasPrefix(uc.Endpoint, endpointPrefixSdns) && uc.Type == "" {
|
||||||
|
uc.Type = ResolverTypeSDNS
|
||||||
|
}
|
||||||
|
if uc.Type != ResolverTypeSDNS {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
sdns, err := dnsstamps.NewServerStampFromString(uc.Endpoint)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
ip, port, _ := net.SplitHostPort(sdns.ServerAddrStr)
|
||||||
|
providerName, port2, _ := net.SplitHostPort(sdns.ProviderName)
|
||||||
|
if port2 != "" {
|
||||||
|
port = port2
|
||||||
|
}
|
||||||
|
if providerName == "" {
|
||||||
|
providerName = sdns.ProviderName
|
||||||
|
}
|
||||||
|
switch sdns.Proto {
|
||||||
|
case dnsstamps.StampProtoTypeDoH:
|
||||||
|
uc.Type = ResolverTypeDOH
|
||||||
|
host := sdns.ProviderName
|
||||||
|
if port != "" && port != defaultPortFor(uc.Type) {
|
||||||
|
host = net.JoinHostPort(providerName, port)
|
||||||
|
}
|
||||||
|
uc.Endpoint = "https://" + host + sdns.Path
|
||||||
|
case dnsstamps.StampProtoTypeTLS:
|
||||||
|
uc.Type = ResolverTypeDOT
|
||||||
|
uc.Endpoint = net.JoinHostPort(providerName, port)
|
||||||
|
case dnsstamps.StampProtoTypeDoQ:
|
||||||
|
uc.Type = ResolverTypeDOQ
|
||||||
|
uc.Endpoint = net.JoinHostPort(providerName, port)
|
||||||
|
case dnsstamps.StampProtoTypePlain:
|
||||||
|
uc.Type = ResolverTypeLegacy
|
||||||
|
uc.Endpoint = sdns.ServerAddrStr
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("unsupported stamp protocol %q", sdns.Proto)
|
||||||
|
}
|
||||||
|
uc.BootstrapIP = ip
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Context returns a new context with timeout set from upstream config.
|
||||||
|
func (uc *UpstreamConfig) Context(ctx context.Context) (context.Context, context.CancelFunc) {
|
||||||
|
if uc.Timeout > 0 {
|
||||||
|
return context.WithTimeout(ctx, time.Millisecond*time.Duration(uc.Timeout))
|
||||||
|
}
|
||||||
|
return context.WithCancel(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// FallbackToDirectIP changes ControlD upstream endpoint to use direct IP instead of domain.
|
||||||
|
func (uc *UpstreamConfig) FallbackToDirectIP() bool {
|
||||||
|
if !uc.IsControlD() {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if uc.u == nil || uc.Domain == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
done := false
|
||||||
|
uc.fallbackOnce.Do(func() {
|
||||||
|
var ip string
|
||||||
|
switch {
|
||||||
|
case dns.IsSubDomain(PremiumDnsDomain, uc.Domain):
|
||||||
|
ip = PremiumDNSBoostrapIP
|
||||||
|
case dns.IsSubDomain(FreeDnsDomain, uc.Domain):
|
||||||
|
ip = FreeDNSBoostrapIP
|
||||||
|
default:
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ProxyLogger.Load().Warn().Msgf("using direct IP for %q: %s", uc.Endpoint, ip)
|
||||||
|
uc.u.Host = ip
|
||||||
|
done = true
|
||||||
|
})
|
||||||
|
return done
|
||||||
|
}
|
||||||
|
|
||||||
// Init initialized necessary values for an ListenerConfig.
|
// Init initialized necessary values for an ListenerConfig.
|
||||||
func (lc *ListenerConfig) Init() {
|
func (lc *ListenerConfig) Init() {
|
||||||
if lc.Policy != nil {
|
if lc.Policy != nil {
|
||||||
@@ -683,6 +951,24 @@ func upstreamConfigStructLevelValidation(sl validator.StructLevel) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Empty type is ok only for endpoints starts with "h3://" and "sdns://".
|
||||||
|
if uc.Type == "" && !strings.HasPrefix(uc.Endpoint, endpointPrefixH3) && !strings.HasPrefix(uc.Endpoint, endpointPrefixSdns) {
|
||||||
|
sl.ReportError(uc.Endpoint, "type", "type", "oneof", "doh doh3 dot doq os legacy sdns")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// initDoHScheme/initDnsStamps may change upstreams information,
|
||||||
|
// so restoring changed values after validation to keep original one.
|
||||||
|
defer func(ep, typ string) {
|
||||||
|
uc.Endpoint = ep
|
||||||
|
uc.Type = typ
|
||||||
|
}(uc.Endpoint, uc.Type)
|
||||||
|
|
||||||
|
if err := uc.initDnsStamps(); err != nil {
|
||||||
|
sl.ReportError(uc.Endpoint, "endpoint", "Endpoint", "http_url", "")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
uc.initDoHScheme()
|
||||||
// DoH/DoH3 requires endpoint is an HTTP url.
|
// DoH/DoH3 requires endpoint is an HTTP url.
|
||||||
if uc.Type == ResolverTypeDOH || uc.Type == ResolverTypeDOH3 {
|
if uc.Type == ResolverTypeDOH || uc.Type == ResolverTypeDOH3 {
|
||||||
u, err := url.Parse(uc.Endpoint)
|
u, err := url.Parse(uc.Endpoint)
|
||||||
@@ -690,10 +976,6 @@ func upstreamConfigStructLevelValidation(sl validator.StructLevel) {
|
|||||||
sl.ReportError(uc.Endpoint, "endpoint", "Endpoint", "http_url", "")
|
sl.ReportError(uc.Endpoint, "endpoint", "Endpoint", "http_url", "")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if u.Scheme != "http" && u.Scheme != "https" {
|
|
||||||
sl.ReportError(uc.Endpoint, "endpoint", "Endpoint", "http_url", "")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -715,13 +997,19 @@ func defaultPortFor(typ string) string {
|
|||||||
// - If endpoint is an IP address -> ResolverTypeLegacy
|
// - If endpoint is an IP address -> ResolverTypeLegacy
|
||||||
// - If endpoint starts with "https://" -> ResolverTypeDOH
|
// - If endpoint starts with "https://" -> ResolverTypeDOH
|
||||||
// - If endpoint starts with "quic://" -> ResolverTypeDOQ
|
// - If endpoint starts with "quic://" -> ResolverTypeDOQ
|
||||||
|
// - If endpoint starts with "h3://" -> ResolverTypeDOH3
|
||||||
|
// - If endpoint starts with "sdns://" -> ResolverTypeSDNS
|
||||||
// - For anything else -> ResolverTypeDOT
|
// - For anything else -> ResolverTypeDOT
|
||||||
func ResolverTypeFromEndpoint(endpoint string) string {
|
func ResolverTypeFromEndpoint(endpoint string) string {
|
||||||
switch {
|
switch {
|
||||||
case strings.HasPrefix(endpoint, "https://"):
|
case strings.HasPrefix(endpoint, endpointPrefixHTTPS):
|
||||||
return ResolverTypeDOH
|
return ResolverTypeDOH
|
||||||
case strings.HasPrefix(endpoint, "quic://"):
|
case strings.HasPrefix(endpoint, endpointPrefixQUIC):
|
||||||
return ResolverTypeDOQ
|
return ResolverTypeDOQ
|
||||||
|
case strings.HasPrefix(endpoint, endpointPrefixH3):
|
||||||
|
return ResolverTypeDOH3
|
||||||
|
case strings.HasPrefix(endpoint, endpointPrefixSdns):
|
||||||
|
return ResolverTypeSDNS
|
||||||
}
|
}
|
||||||
host := endpoint
|
host := endpoint
|
||||||
if strings.Contains(endpoint, ":") {
|
if strings.Contains(endpoint, ":") {
|
||||||
@@ -733,10 +1021,6 @@ func ResolverTypeFromEndpoint(endpoint string) string {
|
|||||||
return ResolverTypeDOT
|
return ResolverTypeDOT
|
||||||
}
|
}
|
||||||
|
|
||||||
func pick(s []string) string {
|
|
||||||
return s[rand.Intn(len(s))]
|
|
||||||
}
|
|
||||||
|
|
||||||
// upstreamUID generates an unique identifier for an upstream.
|
// upstreamUID generates an unique identifier for an upstream.
|
||||||
func upstreamUID() string {
|
func upstreamUID() string {
|
||||||
b := make([]byte, 4)
|
b := make([]byte, 4)
|
||||||
@@ -748,3 +1032,42 @@ func upstreamUID() string {
|
|||||||
return hex.EncodeToString(b)
|
return hex.EncodeToString(b)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// String returns a string representation of the UpstreamConfig for logging.
|
||||||
|
func (uc *UpstreamConfig) String() string {
|
||||||
|
if uc == nil {
|
||||||
|
return "<nil>"
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("{name: %q, type: %q, endpoint: %q, bootstrap_ip: %q, domain: %q, ip_stack: %q}",
|
||||||
|
uc.Name, uc.Type, uc.Endpoint, uc.BootstrapIP, uc.Domain, uc.IPStack)
|
||||||
|
}
|
||||||
|
|
||||||
|
// bootstrapIPsFromControlDDomain returns bootstrap IPs for ControlD domain.
|
||||||
|
func bootstrapIPsFromControlDDomain(domain string) []string {
|
||||||
|
switch {
|
||||||
|
case dns.IsSubDomain(PremiumDnsDomain, domain):
|
||||||
|
return []string{PremiumDNSBoostrapIP, PremiumDNSBoostrapIPv6}
|
||||||
|
case dns.IsSubDomain(FreeDnsDomain, domain):
|
||||||
|
return []string{FreeDNSBoostrapIP, FreeDNSBoostrapIPv6}
|
||||||
|
case dns.IsSubDomain(premiumDnsDomainDev, domain):
|
||||||
|
return []string{premiumDNSBoostrapIP, premiumDNSBoostrapIPv6}
|
||||||
|
case dns.IsSubDomain(freeDnsDomainDev, domain):
|
||||||
|
return []string{freeDNSBoostrapIP, freeDNSBoostrapIPv6}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func transportByIpStack[T any](ipStack string, dnsType uint16, transport, transport4, transport6 T) T {
|
||||||
|
switch ipStack {
|
||||||
|
case IpStackBoth, IpStackV4, IpStackV6:
|
||||||
|
return transport
|
||||||
|
case IpStackSplit:
|
||||||
|
switch dnsType {
|
||||||
|
case dns.TypeA:
|
||||||
|
return transport4
|
||||||
|
default:
|
||||||
|
return transport6
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return transport
|
||||||
|
}
|
||||||
|
|||||||
+229
-11
@@ -2,30 +2,56 @@ package ctrld
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"net/url"
|
"net/url"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestUpstreamConfig_SetupBootstrapIP(t *testing.T) {
|
func TestUpstreamConfig_SetupBootstrapIP(t *testing.T) {
|
||||||
uc := &UpstreamConfig{
|
tests := []struct {
|
||||||
Name: "test",
|
name string
|
||||||
Type: ResolverTypeDOH,
|
uc *UpstreamConfig
|
||||||
Endpoint: "https://freedns.controld.com/p2",
|
}{
|
||||||
Timeout: 5000,
|
{
|
||||||
|
name: "doh/doh3",
|
||||||
|
uc: &UpstreamConfig{
|
||||||
|
Name: "doh",
|
||||||
|
Type: ResolverTypeDOH,
|
||||||
|
Endpoint: "https://freedns.controld.com/p2",
|
||||||
|
Timeout: 5000,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "doq/dot",
|
||||||
|
uc: &UpstreamConfig{
|
||||||
|
Name: "dot",
|
||||||
|
Type: ResolverTypeDOT,
|
||||||
|
Endpoint: "p2.freedns.controld.com",
|
||||||
|
Timeout: 5000,
|
||||||
|
},
|
||||||
|
},
|
||||||
}
|
}
|
||||||
uc.Init()
|
for _, tc := range tests {
|
||||||
uc.setupBootstrapIP(false)
|
tc := tc
|
||||||
if len(uc.bootstrapIPs) == 0 {
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
t.Log(nameservers())
|
// Enable parallel tests once https://github.com/microsoft/wmi/issues/165 fixed.
|
||||||
t.Fatal("could not bootstrap ip without bootstrap DNS")
|
// t.Parallel()
|
||||||
|
tc.uc.Init()
|
||||||
|
tc.uc.SetupBootstrapIP()
|
||||||
|
if len(tc.uc.bootstrapIPs) == 0 {
|
||||||
|
t.Log(defaultNameservers())
|
||||||
|
t.Fatalf("could not bootstrap ip: %s", tc.uc.String())
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
t.Log(uc)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestUpstreamConfig_Init(t *testing.T) {
|
func TestUpstreamConfig_Init(t *testing.T) {
|
||||||
u1, _ := url.Parse("https://example.com")
|
u1, _ := url.Parse("https://example.com")
|
||||||
u2, _ := url.Parse("https://example.com?k=v")
|
u2, _ := url.Parse("https://example.com?k=v")
|
||||||
|
u3, _ := url.Parse("https://freedns.controld.com/p1")
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
uc *UpstreamConfig
|
uc *UpstreamConfig
|
||||||
@@ -178,6 +204,152 @@ func TestUpstreamConfig_Init(t *testing.T) {
|
|||||||
u: u2,
|
u: u2,
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
"h3",
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "doh3",
|
||||||
|
Type: "doh3",
|
||||||
|
Endpoint: "h3://example.com",
|
||||||
|
BootstrapIP: "",
|
||||||
|
Domain: "",
|
||||||
|
Timeout: 0,
|
||||||
|
},
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "doh3",
|
||||||
|
Type: "doh3",
|
||||||
|
Endpoint: "https://example.com",
|
||||||
|
BootstrapIP: "",
|
||||||
|
Domain: "example.com",
|
||||||
|
Timeout: 0,
|
||||||
|
IPStack: IpStackBoth,
|
||||||
|
u: u1,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"h3 without type",
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "doh3",
|
||||||
|
Endpoint: "h3://example.com",
|
||||||
|
BootstrapIP: "",
|
||||||
|
Domain: "",
|
||||||
|
Timeout: 0,
|
||||||
|
},
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "doh3",
|
||||||
|
Type: "doh3",
|
||||||
|
Endpoint: "https://example.com",
|
||||||
|
BootstrapIP: "",
|
||||||
|
Domain: "example.com",
|
||||||
|
Timeout: 0,
|
||||||
|
IPStack: IpStackBoth,
|
||||||
|
u: u1,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"sdns -> doh",
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "sdns",
|
||||||
|
Type: "sdns",
|
||||||
|
Endpoint: "sdns://AgMAAAAAAAAACjc2Ljc2LjIuMTEAFGZyZWVkbnMuY29udHJvbGQuY29tAy9wMQ",
|
||||||
|
BootstrapIP: "",
|
||||||
|
Domain: "",
|
||||||
|
Timeout: 0,
|
||||||
|
IPStack: IpStackBoth,
|
||||||
|
},
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "sdns",
|
||||||
|
Type: "doh",
|
||||||
|
Endpoint: "https://freedns.controld.com/p1",
|
||||||
|
BootstrapIP: "76.76.2.11",
|
||||||
|
Domain: "freedns.controld.com",
|
||||||
|
Timeout: 0,
|
||||||
|
IPStack: IpStackBoth,
|
||||||
|
u: u3,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"sdns -> dot",
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "sdns",
|
||||||
|
Type: "sdns",
|
||||||
|
Endpoint: "sdns://AwcAAAAAAAAACjc2Ljc2LjIuMTEAFGZyZWVkbnMuY29udHJvbGQuY29t",
|
||||||
|
BootstrapIP: "",
|
||||||
|
Domain: "",
|
||||||
|
Timeout: 0,
|
||||||
|
IPStack: IpStackBoth,
|
||||||
|
},
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "sdns",
|
||||||
|
Type: "dot",
|
||||||
|
Endpoint: "freedns.controld.com:843",
|
||||||
|
BootstrapIP: "76.76.2.11",
|
||||||
|
Domain: "freedns.controld.com",
|
||||||
|
Timeout: 0,
|
||||||
|
IPStack: IpStackBoth,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"sdns -> doq",
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "sdns",
|
||||||
|
Type: "sdns",
|
||||||
|
Endpoint: "sdns://BAcAAAAAAAAACjc2Ljc2LjIuMTEAFGZyZWVkbnMuY29udHJvbGQuY29t",
|
||||||
|
BootstrapIP: "",
|
||||||
|
Domain: "",
|
||||||
|
Timeout: 0,
|
||||||
|
IPStack: IpStackBoth,
|
||||||
|
},
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "sdns",
|
||||||
|
Type: "doq",
|
||||||
|
Endpoint: "freedns.controld.com:784",
|
||||||
|
BootstrapIP: "76.76.2.11",
|
||||||
|
Domain: "freedns.controld.com",
|
||||||
|
Timeout: 0,
|
||||||
|
IPStack: IpStackBoth,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"sdns -> legacy",
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "sdns",
|
||||||
|
Type: "sdns",
|
||||||
|
Endpoint: "sdns://AAcAAAAAAAAACjc2Ljc2LjIuMTE",
|
||||||
|
BootstrapIP: "",
|
||||||
|
Domain: "",
|
||||||
|
Timeout: 0,
|
||||||
|
IPStack: IpStackBoth,
|
||||||
|
},
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "sdns",
|
||||||
|
Type: "legacy",
|
||||||
|
Endpoint: "76.76.2.11:53",
|
||||||
|
BootstrapIP: "76.76.2.11",
|
||||||
|
Domain: "76.76.2.11",
|
||||||
|
Timeout: 0,
|
||||||
|
IPStack: IpStackBoth,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"sdns without type",
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "sdns",
|
||||||
|
Endpoint: "sdns://AAcAAAAAAAAACjc2Ljc2LjIuMTE",
|
||||||
|
BootstrapIP: "",
|
||||||
|
Domain: "",
|
||||||
|
Timeout: 0,
|
||||||
|
IPStack: IpStackBoth,
|
||||||
|
},
|
||||||
|
&UpstreamConfig{
|
||||||
|
Name: "sdns",
|
||||||
|
Type: "legacy",
|
||||||
|
Endpoint: "76.76.2.11:53",
|
||||||
|
BootstrapIP: "76.76.2.11",
|
||||||
|
Domain: "76.76.2.11",
|
||||||
|
Timeout: 0,
|
||||||
|
IPStack: IpStackBoth,
|
||||||
|
},
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tc := range tests {
|
for _, tc := range tests {
|
||||||
@@ -334,6 +506,52 @@ func TestUpstreamConfig_IsDiscoverable(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRebootstrapRace(t *testing.T) {
|
||||||
|
uc := &UpstreamConfig{
|
||||||
|
Name: "test-doh",
|
||||||
|
Type: ResolverTypeDOH,
|
||||||
|
Endpoint: "https://example.com/dns-query",
|
||||||
|
Domain: "example.com",
|
||||||
|
bootstrapIPs: []string{"1.1.1.1", "1.0.0.1"},
|
||||||
|
}
|
||||||
|
|
||||||
|
uc.SetupTransport()
|
||||||
|
|
||||||
|
if uc.transport == nil {
|
||||||
|
t.Fatal("initial transport should be set")
|
||||||
|
}
|
||||||
|
|
||||||
|
const goroutines = 100
|
||||||
|
|
||||||
|
uc.ReBootstrap()
|
||||||
|
|
||||||
|
started := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
close(started)
|
||||||
|
for {
|
||||||
|
switch uc.rebootstrap.Load() {
|
||||||
|
case rebootstrapStarted, rebootstrapInProgress:
|
||||||
|
uc.ReBootstrap()
|
||||||
|
default:
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
<-started
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
wg.Add(goroutines)
|
||||||
|
for range goroutines {
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
uc.ensureSetupTransport()
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
wg.Wait()
|
||||||
|
}
|
||||||
|
|
||||||
func ptrBool(b bool) *bool {
|
func ptrBool(b bool) *bool {
|
||||||
return &b
|
return &b
|
||||||
}
|
}
|
||||||
|
|||||||
+63
-52
@@ -1,5 +1,3 @@
|
|||||||
//go:build !qf
|
|
||||||
|
|
||||||
package ctrld
|
package ctrld
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -11,34 +9,17 @@ import (
|
|||||||
"runtime"
|
"runtime"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/miekg/dns"
|
|
||||||
"github.com/quic-go/quic-go"
|
"github.com/quic-go/quic-go"
|
||||||
"github.com/quic-go/quic-go/http3"
|
"github.com/quic-go/quic-go/http3"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (uc *UpstreamConfig) setupDOH3Transport() {
|
|
||||||
switch uc.IPStack {
|
|
||||||
case IpStackBoth, "":
|
|
||||||
uc.http3RoundTripper = uc.newDOH3Transport(uc.bootstrapIPs)
|
|
||||||
case IpStackV4:
|
|
||||||
uc.http3RoundTripper = uc.newDOH3Transport(uc.bootstrapIPs4)
|
|
||||||
case IpStackV6:
|
|
||||||
uc.http3RoundTripper = uc.newDOH3Transport(uc.bootstrapIPs6)
|
|
||||||
case IpStackSplit:
|
|
||||||
uc.http3RoundTripper4 = uc.newDOH3Transport(uc.bootstrapIPs4)
|
|
||||||
if hasIPv6() {
|
|
||||||
uc.http3RoundTripper6 = uc.newDOH3Transport(uc.bootstrapIPs6)
|
|
||||||
} else {
|
|
||||||
uc.http3RoundTripper6 = uc.http3RoundTripper4
|
|
||||||
}
|
|
||||||
uc.http3RoundTripper = uc.newDOH3Transport(uc.bootstrapIPs)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (uc *UpstreamConfig) newDOH3Transport(addrs []string) http.RoundTripper {
|
func (uc *UpstreamConfig) newDOH3Transport(addrs []string) http.RoundTripper {
|
||||||
rt := &http3.RoundTripper{}
|
if uc.Type != ResolverTypeDOH3 {
|
||||||
rt.TLSClientConfig = &tls.Config{RootCAs: uc.certPool}
|
return nil
|
||||||
rt.Dial = func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (quic.EarlyConnection, error) {
|
}
|
||||||
|
rt := &http3.Transport{}
|
||||||
|
rt.TLSClientConfig = &tls.Config{RootCAs: uc.certPool, MinVersion: tls.VersionTLS12}
|
||||||
|
rt.Dial = func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
||||||
_, port, _ := net.SplitHostPort(addr)
|
_, port, _ := net.SplitHostPort(addr)
|
||||||
// if we have a bootstrap ip set, use it to avoid DNS lookup
|
// if we have a bootstrap ip set, use it to avoid DNS lookup
|
||||||
if uc.BootstrapIP != "" {
|
if uc.BootstrapIP != "" {
|
||||||
@@ -66,31 +47,25 @@ func (uc *UpstreamConfig) newDOH3Transport(addrs []string) http.RoundTripper {
|
|||||||
ProxyLogger.Load().Debug().Msgf("sending doh3 request to: %s", conn.RemoteAddr())
|
ProxyLogger.Load().Debug().Msgf("sending doh3 request to: %s", conn.RemoteAddr())
|
||||||
return conn, err
|
return conn, err
|
||||||
}
|
}
|
||||||
runtime.SetFinalizer(rt, func(rt *http3.RoundTripper) {
|
runtime.SetFinalizer(rt, func(rt *http3.Transport) {
|
||||||
rt.CloseIdleConnections()
|
rt.CloseIdleConnections()
|
||||||
})
|
})
|
||||||
return rt
|
return rt
|
||||||
}
|
}
|
||||||
|
|
||||||
func (uc *UpstreamConfig) doh3Transport(dnsType uint16) http.RoundTripper {
|
func (uc *UpstreamConfig) doh3Transport(dnsType uint16) http.RoundTripper {
|
||||||
uc.transportOnce.Do(func() {
|
uc.ensureSetupTransport()
|
||||||
uc.SetupTransport()
|
return transportByIpStack(uc.IPStack, dnsType, uc.http3RoundTripper, uc.http3RoundTripper4, uc.http3RoundTripper6)
|
||||||
})
|
}
|
||||||
if uc.rebootstrap.CompareAndSwap(true, false) {
|
|
||||||
uc.SetupTransport()
|
func (uc *UpstreamConfig) doqTransport(dnsType uint16) *doqConnPool {
|
||||||
}
|
uc.ensureSetupTransport()
|
||||||
switch uc.IPStack {
|
return transportByIpStack(uc.IPStack, dnsType, uc.doqConnPool, uc.doqConnPool4, uc.doqConnPool6)
|
||||||
case IpStackBoth, IpStackV4, IpStackV6:
|
}
|
||||||
return uc.http3RoundTripper
|
|
||||||
case IpStackSplit:
|
func (uc *UpstreamConfig) dotTransport(dnsType uint16) *dotConnPool {
|
||||||
switch dnsType {
|
uc.ensureSetupTransport()
|
||||||
case dns.TypeA:
|
return transportByIpStack(uc.IPStack, dnsType, uc.dotClientPool, uc.dotClientPool4, uc.dotClientPool6)
|
||||||
return uc.http3RoundTripper4
|
|
||||||
default:
|
|
||||||
return uc.http3RoundTripper6
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return uc.http3RoundTripper
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Putting the code for quic parallel dialer here:
|
// Putting the code for quic parallel dialer here:
|
||||||
@@ -98,14 +73,24 @@ func (uc *UpstreamConfig) doh3Transport(dnsType uint16) http.RoundTripper {
|
|||||||
// - quic dialer is different with net.Dialer
|
// - quic dialer is different with net.Dialer
|
||||||
// - simplification for quic free version
|
// - simplification for quic free version
|
||||||
type parallelDialerResult struct {
|
type parallelDialerResult struct {
|
||||||
conn quic.EarlyConnection
|
conn *quic.Conn
|
||||||
err error
|
err error
|
||||||
}
|
}
|
||||||
|
|
||||||
type quicParallelDialer struct{}
|
// quicParallelDialer races DialEarly across a list of remote addresses and
|
||||||
|
// returns the first successful connection. When transport is non-nil, all
|
||||||
|
// dials share that transport's UDP socket, which removes both the per-dial
|
||||||
|
// socket allocation and the winner-path socket leak that an owner-of-the-conn
|
||||||
|
// receiver cannot clean up. When transport is nil, the dialer falls back to a
|
||||||
|
// fresh UDP socket per attempt (compat path used where no shared transport is
|
||||||
|
// available yet); the loser paths close their sockets, and the winner path's
|
||||||
|
// socket is owned by quic.DialEarly's internal transport.
|
||||||
|
type quicParallelDialer struct {
|
||||||
|
transport *quic.Transport
|
||||||
|
}
|
||||||
|
|
||||||
// Dial performs parallel dialing to the given address list.
|
// Dial performs parallel dialing to the given address list.
|
||||||
func (d *quicParallelDialer) Dial(ctx context.Context, addrs []string, tlsCfg *tls.Config, cfg *quic.Config) (quic.EarlyConnection, error) {
|
func (d *quicParallelDialer) Dial(ctx context.Context, addrs []string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
||||||
if len(addrs) == 0 {
|
if len(addrs) == 0 {
|
||||||
return nil, errors.New("empty addresses")
|
return nil, errors.New("empty addresses")
|
||||||
}
|
}
|
||||||
@@ -130,12 +115,24 @@ func (d *quicParallelDialer) Dial(ctx context.Context, addrs []string, tlsCfg *t
|
|||||||
ch <- ¶llelDialerResult{conn: nil, err: err}
|
ch <- ¶llelDialerResult{conn: nil, err: err}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
udpConn, err := net.ListenUDP("udp", nil)
|
var (
|
||||||
if err != nil {
|
conn *quic.Conn
|
||||||
ch <- ¶llelDialerResult{conn: nil, err: err}
|
udpConn *net.UDPConn
|
||||||
return
|
)
|
||||||
|
if d.transport != nil {
|
||||||
|
conn, err = d.transport.DialEarly(ctx, remoteAddr, tlsCfg, cfg)
|
||||||
|
} else {
|
||||||
|
udpConn, err = net.ListenUDP("udp", nil)
|
||||||
|
if err != nil {
|
||||||
|
ch <- ¶llelDialerResult{conn: nil, err: err}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
conn, err = quic.DialEarly(ctx, udpConn, remoteAddr, tlsCfg, cfg)
|
||||||
|
if err != nil {
|
||||||
|
udpConn.Close()
|
||||||
|
udpConn = nil
|
||||||
|
}
|
||||||
}
|
}
|
||||||
conn, err := quic.DialEarly(ctx, udpConn, remoteAddr, tlsCfg, cfg)
|
|
||||||
select {
|
select {
|
||||||
case ch <- ¶llelDialerResult{conn: conn, err: err}:
|
case ch <- ¶llelDialerResult{conn: conn, err: err}:
|
||||||
case <-done:
|
case <-done:
|
||||||
@@ -160,3 +157,17 @@ func (d *quicParallelDialer) Dial(ctx context.Context, addrs []string, tlsCfg *t
|
|||||||
|
|
||||||
return nil, errors.Join(errs...)
|
return nil, errors.Join(errs...)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (uc *UpstreamConfig) newDOQConnPool(addrs []string) *doqConnPool {
|
||||||
|
if uc.Type != ResolverTypeDOQ {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return newDOQConnPool(uc, addrs)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (uc *UpstreamConfig) newDOTClientPool(addrs []string) *dotConnPool {
|
||||||
|
if uc.Type != ResolverTypeDOT {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return newDOTClientPool(uc, addrs)
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,9 +0,0 @@
|
|||||||
//go:build qf
|
|
||||||
|
|
||||||
package ctrld
|
|
||||||
|
|
||||||
import "net/http"
|
|
||||||
|
|
||||||
func (uc *UpstreamConfig) setupDOH3Transport() {}
|
|
||||||
|
|
||||||
func (uc *UpstreamConfig) doh3Transport(dnsType uint16) http.RoundTripper { return nil }
|
|
||||||
+68
-1
@@ -1,9 +1,11 @@
|
|||||||
package ctrld_test
|
package ctrld_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/go-playground/validator/v10"
|
"github.com/go-playground/validator/v10"
|
||||||
"github.com/spf13/viper"
|
"github.com/spf13/viper"
|
||||||
@@ -21,6 +23,8 @@ func TestLoadConfig(t *testing.T) {
|
|||||||
|
|
||||||
assert.Equal(t, "info", cfg.Service.LogLevel)
|
assert.Equal(t, "info", cfg.Service.LogLevel)
|
||||||
assert.Equal(t, "/path/to/log.log", cfg.Service.LogPath)
|
assert.Equal(t, "/path/to/log.log", cfg.Service.LogPath)
|
||||||
|
assert.Equal(t, false, *cfg.Service.DnsWatchdogEnabled)
|
||||||
|
assert.Equal(t, time.Duration(20*time.Second), *cfg.Service.DnsWatchdogInvterval)
|
||||||
|
|
||||||
assert.Len(t, cfg.Network, 2)
|
assert.Len(t, cfg.Network, 2)
|
||||||
assert.Contains(t, cfg.Network, "0")
|
assert.Contains(t, cfg.Network, "0")
|
||||||
@@ -102,6 +106,12 @@ func TestConfigValidation(t *testing.T) {
|
|||||||
{"invalid lease file format", configWithInvalidLeaseFileFormat(t), true},
|
{"invalid lease file format", configWithInvalidLeaseFileFormat(t), true},
|
||||||
{"invalid doh/doh3 endpoint", configWithInvalidDoHEndpoint(t), true},
|
{"invalid doh/doh3 endpoint", configWithInvalidDoHEndpoint(t), true},
|
||||||
{"invalid client id pref", configWithInvalidClientIDPref(t), true},
|
{"invalid client id pref", configWithInvalidClientIDPref(t), true},
|
||||||
|
{"doh endpoint without scheme", dohUpstreamEndpointWithoutScheme(t), false},
|
||||||
|
{"doh endpoint without type", dohUpstreamEndpointWithoutType(t), true},
|
||||||
|
{"doh3 endpoint without type", doh3UpstreamEndpointWithoutType(t), false},
|
||||||
|
{"sdns endpoint without type", sdnsUpstreamEndpointWithoutType(t), false},
|
||||||
|
{"maximum number of flush cache domains", configWithInvalidFlushCacheDomain(t), true},
|
||||||
|
{"kea dhcp4 format", configWithDhcp4KeaFormat(t), false},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tc := range tests {
|
for _, tc := range tests {
|
||||||
@@ -121,6 +131,21 @@ func TestConfigValidation(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestConfigValidationDoNotChangeEndpoint(t *testing.T) {
|
||||||
|
cfg := configWithInvalidDoHEndpoint(t)
|
||||||
|
endpointMap := map[string]struct{}{}
|
||||||
|
for _, uc := range cfg.Upstream {
|
||||||
|
endpointMap[uc.Endpoint] = struct{}{}
|
||||||
|
}
|
||||||
|
validate := validator.New()
|
||||||
|
_ = ctrld.ValidateConfig(validate, cfg)
|
||||||
|
for _, uc := range cfg.Upstream {
|
||||||
|
if _, ok := endpointMap[uc.Endpoint]; !ok {
|
||||||
|
t.Fatalf("expected endpoint '%s' to exist", uc.Endpoint)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestConfigDiscoverOverride(t *testing.T) {
|
func TestConfigDiscoverOverride(t *testing.T) {
|
||||||
v := viper.NewWithOptions(viper.KeyDelimiter("::"))
|
v := viper.NewWithOptions(viper.KeyDelimiter("::"))
|
||||||
ctrld.InitConfig(v, "test_config_discover_override")
|
ctrld.InitConfig(v, "test_config_discover_override")
|
||||||
@@ -167,6 +192,33 @@ func invalidUpstreamType(t *testing.T) *ctrld.Config {
|
|||||||
return cfg
|
return cfg
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func dohUpstreamEndpointWithoutScheme(t *testing.T) *ctrld.Config {
|
||||||
|
cfg := defaultConfig(t)
|
||||||
|
cfg.Upstream["0"].Endpoint = "freedns.controld.com/p1"
|
||||||
|
return cfg
|
||||||
|
}
|
||||||
|
|
||||||
|
func dohUpstreamEndpointWithoutType(t *testing.T) *ctrld.Config {
|
||||||
|
cfg := defaultConfig(t)
|
||||||
|
cfg.Upstream["0"].Endpoint = "https://freedns.controld.com/p1"
|
||||||
|
cfg.Upstream["0"].Type = ""
|
||||||
|
return cfg
|
||||||
|
}
|
||||||
|
|
||||||
|
func doh3UpstreamEndpointWithoutType(t *testing.T) *ctrld.Config {
|
||||||
|
cfg := defaultConfig(t)
|
||||||
|
cfg.Upstream["0"].Endpoint = "h3://freedns.controld.com/p1"
|
||||||
|
cfg.Upstream["0"].Type = ""
|
||||||
|
return cfg
|
||||||
|
}
|
||||||
|
|
||||||
|
func sdnsUpstreamEndpointWithoutType(t *testing.T) *ctrld.Config {
|
||||||
|
cfg := defaultConfig(t)
|
||||||
|
cfg.Upstream["0"].Endpoint = "sdns://AgMAAAAAAAAACjc2Ljc2LjIuMTEAFGZyZWVkbnMuY29udHJvbGQuY29tAy9wMQ"
|
||||||
|
cfg.Upstream["0"].Type = ""
|
||||||
|
return cfg
|
||||||
|
}
|
||||||
|
|
||||||
func invalidUpstreamTimeout(t *testing.T) *ctrld.Config {
|
func invalidUpstreamTimeout(t *testing.T) *ctrld.Config {
|
||||||
cfg := defaultConfig(t)
|
cfg := defaultConfig(t)
|
||||||
cfg.Upstream["0"].Timeout = -1
|
cfg.Upstream["0"].Timeout = -1
|
||||||
@@ -256,9 +308,15 @@ func configWithInvalidLeaseFileFormat(t *testing.T) *ctrld.Config {
|
|||||||
return cfg
|
return cfg
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func configWithDhcp4KeaFormat(t *testing.T) *ctrld.Config {
|
||||||
|
cfg := defaultConfig(t)
|
||||||
|
cfg.Service.DHCPLeaseFileFormat = "kea-dhcp4"
|
||||||
|
return cfg
|
||||||
|
}
|
||||||
|
|
||||||
func configWithInvalidDoHEndpoint(t *testing.T) *ctrld.Config {
|
func configWithInvalidDoHEndpoint(t *testing.T) *ctrld.Config {
|
||||||
cfg := defaultConfig(t)
|
cfg := defaultConfig(t)
|
||||||
cfg.Upstream["0"].Endpoint = "1.1.1.1"
|
cfg.Upstream["0"].Endpoint = "/1.1.1.1"
|
||||||
cfg.Upstream["0"].Type = ctrld.ResolverTypeDOH
|
cfg.Upstream["0"].Type = ctrld.ResolverTypeDOH
|
||||||
return cfg
|
return cfg
|
||||||
}
|
}
|
||||||
@@ -268,3 +326,12 @@ func configWithInvalidClientIDPref(t *testing.T) *ctrld.Config {
|
|||||||
cfg.Service.ClientIDPref = "foo"
|
cfg.Service.ClientIDPref = "foo"
|
||||||
return cfg
|
return cfg
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func configWithInvalidFlushCacheDomain(t *testing.T) *ctrld.Config {
|
||||||
|
cfg := defaultConfig(t)
|
||||||
|
cfg.Service.CacheFlushDomains = make([]string, 257)
|
||||||
|
for i := range cfg.Service.CacheFlushDomains {
|
||||||
|
cfg.Service.CacheFlushDomains[i] = fmt.Sprintf("%d.com", i)
|
||||||
|
}
|
||||||
|
return cfg
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,7 @@
|
|||||||
|
package ctrld
|
||||||
|
|
||||||
|
// IsDesktopPlatform indicates if ctrld is running on a desktop platform,
|
||||||
|
// currently defined as macOS or Windows workstation.
|
||||||
|
func IsDesktopPlatform() bool {
|
||||||
|
return true
|
||||||
|
}
|
||||||
@@ -0,0 +1,9 @@
|
|||||||
|
//go:build !windows && !darwin
|
||||||
|
|
||||||
|
package ctrld
|
||||||
|
|
||||||
|
// IsDesktopPlatform indicates if ctrld is running on a desktop platform,
|
||||||
|
// currently defined as macOS or Windows workstation.
|
||||||
|
func IsDesktopPlatform() bool {
|
||||||
|
return false
|
||||||
|
}
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
package ctrld
|
||||||
|
|
||||||
|
// IsDesktopPlatform indicates if ctrld is running on a desktop platform,
|
||||||
|
// currently defined as macOS or Windows workstation.
|
||||||
|
func IsDesktopPlatform() bool {
|
||||||
|
return isWindowsWorkStation()
|
||||||
|
}
|
||||||
@@ -0,0 +1,135 @@
|
|||||||
|
//go:build darwin
|
||||||
|
|
||||||
|
package ctrld
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"os/exec"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DiscoverMainUser attempts to find the primary user on macOS systems.
|
||||||
|
// This is designed to work reliably under RMM deployments where traditional
|
||||||
|
// environment variables and session detection may not be available.
|
||||||
|
//
|
||||||
|
// Priority chain (deterministic, lowest UID wins among candidates):
|
||||||
|
// 1. Console user from stat -f %Su /dev/console
|
||||||
|
// 2. Active console session user via scutil
|
||||||
|
// 3. First user with UID >= 501 from dscl (standard macOS user range)
|
||||||
|
func DiscoverMainUser(ctx context.Context) string {
|
||||||
|
logger := ProxyLogger.Load().Debug()
|
||||||
|
|
||||||
|
// Method 1: Check console owner via stat
|
||||||
|
logger.Msg("attempting to discover user via console stat")
|
||||||
|
if user := getConsoleUser(ctx); user != "" && user != "root" {
|
||||||
|
logger.Str("method", "stat").Str("user", user).Msg("found user via console stat")
|
||||||
|
return user
|
||||||
|
}
|
||||||
|
|
||||||
|
// Method 2: Check active console session via scutil
|
||||||
|
logger.Msg("attempting to discover user via scutil ConsoleUser")
|
||||||
|
if user := getScutilConsoleUser(ctx); user != "" && user != "root" {
|
||||||
|
logger.Str("method", "scutil").Str("user", user).Msg("found user via scutil ConsoleUser")
|
||||||
|
return user
|
||||||
|
}
|
||||||
|
|
||||||
|
// Method 3: Find lowest UID >= 501 from directory services
|
||||||
|
logger.Msg("attempting to discover user via dscl directory scan")
|
||||||
|
if user := getLowestRegularUser(ctx); user != "" {
|
||||||
|
logger.Str("method", "dscl").Str("user", user).Msg("found user via dscl scan")
|
||||||
|
return user
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Msg("all user discovery methods failed")
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
// getConsoleUser uses stat to find the owner of /dev/console
|
||||||
|
func getConsoleUser(ctx context.Context) string {
|
||||||
|
cmd := exec.CommandContext(ctx, "stat", "-f", "%Su", "/dev/console")
|
||||||
|
out, err := cmd.Output()
|
||||||
|
if err != nil {
|
||||||
|
ProxyLogger.Load().Debug().Err(err).Msg("failed to stat /dev/console")
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return strings.TrimSpace(string(out))
|
||||||
|
}
|
||||||
|
|
||||||
|
// getScutilConsoleUser uses scutil to get the current console user
|
||||||
|
func getScutilConsoleUser(ctx context.Context) string {
|
||||||
|
cmd := exec.CommandContext(ctx, "scutil", "-r", "ConsoleUser")
|
||||||
|
out, err := cmd.Output()
|
||||||
|
if err != nil {
|
||||||
|
ProxyLogger.Load().Debug().Err(err).Msg("failed to get ConsoleUser via scutil")
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
lines := strings.Split(string(out), "\n")
|
||||||
|
for _, line := range lines {
|
||||||
|
if strings.Contains(line, "Name :") {
|
||||||
|
parts := strings.Fields(line)
|
||||||
|
if len(parts) >= 3 {
|
||||||
|
return strings.TrimSpace(parts[2])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// getLowestRegularUser finds the user with the lowest UID >= 501
|
||||||
|
func getLowestRegularUser(ctx context.Context) string {
|
||||||
|
// Get list of all users with UID >= 501
|
||||||
|
cmd := exec.CommandContext(ctx, "dscl", ".", "list", "/Users", "UniqueID")
|
||||||
|
out, err := cmd.Output()
|
||||||
|
if err != nil {
|
||||||
|
ProxyLogger.Load().Debug().Err(err).Msg("failed to list users via dscl")
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
var candidates []struct {
|
||||||
|
name string
|
||||||
|
uid int
|
||||||
|
}
|
||||||
|
|
||||||
|
lines := strings.Split(string(out), "\n")
|
||||||
|
for _, line := range lines {
|
||||||
|
fields := strings.Fields(line)
|
||||||
|
if len(fields) != 2 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
username := fields[0]
|
||||||
|
uidStr := fields[1]
|
||||||
|
|
||||||
|
uid, err := strconv.Atoi(uidStr)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only consider regular users (UID >= 501 on macOS)
|
||||||
|
if uid >= 501 {
|
||||||
|
candidates = append(candidates, struct {
|
||||||
|
name string
|
||||||
|
uid int
|
||||||
|
}{username, uid})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(candidates) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Find the candidate with the lowest UID (deterministic choice)
|
||||||
|
lowestUID := candidates[0].uid
|
||||||
|
result := candidates[0].name
|
||||||
|
|
||||||
|
for _, candidate := range candidates[1:] {
|
||||||
|
if candidate.uid < lowestUID {
|
||||||
|
lowestUID = candidate.uid
|
||||||
|
result = candidate.name
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
@@ -0,0 +1,238 @@
|
|||||||
|
//go:build linux
|
||||||
|
|
||||||
|
package ctrld
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"context"
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DiscoverMainUser attempts to find the primary user on Linux systems.
|
||||||
|
// This is designed to work reliably under RMM deployments where traditional
|
||||||
|
// environment variables and session detection may not be available.
|
||||||
|
//
|
||||||
|
// Priority chain (deterministic, lowest UID wins among candidates):
|
||||||
|
// 1. Active users from loginctl list-users
|
||||||
|
// 2. Parse /etc/passwd for users with UID >= 1000, prefer admin group members
|
||||||
|
// 3. Fallback to lowest UID >= 1000 from /etc/passwd
|
||||||
|
func DiscoverMainUser(ctx context.Context) string {
|
||||||
|
logger := ProxyLogger.Load().Debug()
|
||||||
|
|
||||||
|
// Method 1: Check active users via loginctl
|
||||||
|
logger.Msg("attempting to discover user via loginctl")
|
||||||
|
if user := getLoginctlUser(ctx); user != "" {
|
||||||
|
logger.Str("method", "loginctl").Str("user", user).Msg("found user via loginctl")
|
||||||
|
return user
|
||||||
|
}
|
||||||
|
|
||||||
|
// Method 2: Parse /etc/passwd and find admin users first
|
||||||
|
logger.Msg("attempting to discover user via /etc/passwd with admin preference")
|
||||||
|
if user := getPasswdUserWithAdminPreference(ctx); user != "" {
|
||||||
|
logger.Str("method", "passwd+admin").Str("user", user).Msg("found admin user via /etc/passwd")
|
||||||
|
return user
|
||||||
|
}
|
||||||
|
|
||||||
|
// Method 3: Fallback to lowest UID >= 1000 from /etc/passwd
|
||||||
|
logger.Msg("attempting to discover user via /etc/passwd lowest UID")
|
||||||
|
if user := getLowestPasswdUser(ctx); user != "" {
|
||||||
|
logger.Str("method", "passwd").Str("user", user).Msg("found user via /etc/passwd")
|
||||||
|
return user
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Msg("all user discovery methods failed")
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
// getLoginctlUser uses loginctl to find active users
|
||||||
|
func getLoginctlUser(ctx context.Context) string {
|
||||||
|
cmd := exec.CommandContext(ctx, "loginctl", "list-users", "--no-legend")
|
||||||
|
out, err := cmd.Output()
|
||||||
|
if err != nil {
|
||||||
|
ProxyLogger.Load().Debug().Err(err).Msg("failed to run loginctl list-users")
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
var candidates []struct {
|
||||||
|
name string
|
||||||
|
uid int
|
||||||
|
}
|
||||||
|
|
||||||
|
lines := strings.Split(string(out), "\n")
|
||||||
|
for _, line := range lines {
|
||||||
|
fields := strings.Fields(line)
|
||||||
|
if len(fields) < 2 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
uidStr := fields[0]
|
||||||
|
username := fields[1]
|
||||||
|
|
||||||
|
uid, err := strconv.Atoi(uidStr)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only consider regular users (UID >= 1000 on Linux)
|
||||||
|
if uid >= 1000 {
|
||||||
|
candidates = append(candidates, struct {
|
||||||
|
name string
|
||||||
|
uid int
|
||||||
|
}{username, uid})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(candidates) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Return user with lowest UID (deterministic choice)
|
||||||
|
lowestUID := candidates[0].uid
|
||||||
|
result := candidates[0].name
|
||||||
|
|
||||||
|
for _, candidate := range candidates[1:] {
|
||||||
|
if candidate.uid < lowestUID {
|
||||||
|
lowestUID = candidate.uid
|
||||||
|
result = candidate.name
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// getPasswdUserWithAdminPreference parses /etc/passwd and prefers admin group members
|
||||||
|
func getPasswdUserWithAdminPreference(ctx context.Context) string {
|
||||||
|
users := parsePasswdFile()
|
||||||
|
if len(users) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
var adminUsers []struct {
|
||||||
|
name string
|
||||||
|
uid int
|
||||||
|
}
|
||||||
|
var regularUsers []struct {
|
||||||
|
name string
|
||||||
|
uid int
|
||||||
|
}
|
||||||
|
|
||||||
|
// Separate admin and regular users
|
||||||
|
for _, user := range users {
|
||||||
|
if isUserInAdminGroups(ctx, user.name) {
|
||||||
|
adminUsers = append(adminUsers, user)
|
||||||
|
} else {
|
||||||
|
regularUsers = append(regularUsers, user)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Prefer admin users, then regular users
|
||||||
|
candidates := adminUsers
|
||||||
|
if len(candidates) == 0 {
|
||||||
|
candidates = regularUsers
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(candidates) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Return user with lowest UID (deterministic choice)
|
||||||
|
lowestUID := candidates[0].uid
|
||||||
|
result := candidates[0].name
|
||||||
|
|
||||||
|
for _, candidate := range candidates[1:] {
|
||||||
|
if candidate.uid < lowestUID {
|
||||||
|
lowestUID = candidate.uid
|
||||||
|
result = candidate.name
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// getLowestPasswdUser returns the user with lowest UID >= 1000 from /etc/passwd
|
||||||
|
func getLowestPasswdUser(ctx context.Context) string {
|
||||||
|
users := parsePasswdFile()
|
||||||
|
if len(users) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Return user with lowest UID (deterministic choice)
|
||||||
|
lowestUID := users[0].uid
|
||||||
|
result := users[0].name
|
||||||
|
|
||||||
|
for _, user := range users[1:] {
|
||||||
|
if user.uid < lowestUID {
|
||||||
|
lowestUID = user.uid
|
||||||
|
result = user.name
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// parsePasswdFile parses /etc/passwd and returns users with UID >= 1000
|
||||||
|
func parsePasswdFile() []struct {
|
||||||
|
name string
|
||||||
|
uid int
|
||||||
|
} {
|
||||||
|
file, err := os.Open("/etc/passwd")
|
||||||
|
if err != nil {
|
||||||
|
ProxyLogger.Load().Debug().Err(err).Msg("failed to open /etc/passwd")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
defer file.Close()
|
||||||
|
|
||||||
|
var users []struct {
|
||||||
|
name string
|
||||||
|
uid int
|
||||||
|
}
|
||||||
|
|
||||||
|
scanner := bufio.NewScanner(file)
|
||||||
|
for scanner.Scan() {
|
||||||
|
line := scanner.Text()
|
||||||
|
fields := strings.Split(line, ":")
|
||||||
|
if len(fields) < 3 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
username := fields[0]
|
||||||
|
uidStr := fields[2]
|
||||||
|
|
||||||
|
uid, err := strconv.Atoi(uidStr)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only consider regular users (UID >= 1000 on Linux)
|
||||||
|
if uid >= 1000 {
|
||||||
|
users = append(users, struct {
|
||||||
|
name string
|
||||||
|
uid int
|
||||||
|
}{username, uid})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return users
|
||||||
|
}
|
||||||
|
|
||||||
|
// isUserInAdminGroups checks if a user is in common admin groups
|
||||||
|
func isUserInAdminGroups(ctx context.Context, username string) bool {
|
||||||
|
adminGroups := []string{"sudo", "wheel", "admin"}
|
||||||
|
|
||||||
|
for _, group := range adminGroups {
|
||||||
|
cmd := exec.CommandContext(ctx, "groups", username)
|
||||||
|
out, err := cmd.Output()
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Contains(string(out), group) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
@@ -0,0 +1,13 @@
|
|||||||
|
//go:build !windows && !linux && !darwin
|
||||||
|
|
||||||
|
package ctrld
|
||||||
|
|
||||||
|
import "context"
|
||||||
|
|
||||||
|
// DiscoverMainUser returns "unknown" for unsupported platforms.
|
||||||
|
// This is a stub implementation for platforms where username detection
|
||||||
|
// is not yet implemented.
|
||||||
|
func DiscoverMainUser(ctx context.Context) string {
|
||||||
|
ProxyLogger.Load().Debug().Msg("username discovery not implemented for this platform")
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
@@ -0,0 +1,292 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package ctrld
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"syscall"
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
"golang.org/x/sys/windows/registry"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
wtsapi32 = windows.NewLazySystemDLL("wtsapi32.dll")
|
||||||
|
procWTSGetActiveConsoleSessionId = wtsapi32.NewProc("WTSGetActiveConsoleSessionId")
|
||||||
|
procWTSQuerySessionInformation = wtsapi32.NewProc("WTSQuerySessionInformationW")
|
||||||
|
procWTSFreeMemory = wtsapi32.NewProc("WTSFreeMemory")
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
WTSUserName = 5
|
||||||
|
)
|
||||||
|
|
||||||
|
// DiscoverMainUser attempts to find the primary user on Windows systems.
|
||||||
|
// This is designed to work reliably under RMM deployments where traditional
|
||||||
|
// environment variables and session detection may not be available.
|
||||||
|
//
|
||||||
|
// Priority chain (deterministic, lowest RID wins among candidates):
|
||||||
|
// 1. Active console session user via WTSGetActiveConsoleSessionId
|
||||||
|
// 2. Registry ProfileList scan for Administrators group members
|
||||||
|
// 3. Fallback to lowest RID from ProfileList
|
||||||
|
func DiscoverMainUser(ctx context.Context) string {
|
||||||
|
logger := ProxyLogger.Load().Debug()
|
||||||
|
|
||||||
|
// Method 1: Check active console session
|
||||||
|
logger.Msg("attempting to discover user via active console session")
|
||||||
|
if user := getActiveConsoleUser(ctx); user != "" {
|
||||||
|
logger.Str("method", "console").Str("user", user).Msg("found user via active console session")
|
||||||
|
return user
|
||||||
|
}
|
||||||
|
|
||||||
|
// Method 2: Scan registry for admin users
|
||||||
|
logger.Msg("attempting to discover user via registry with admin preference")
|
||||||
|
if user := getRegistryUserWithAdminPreference(ctx); user != "" {
|
||||||
|
logger.Str("method", "registry+admin").Str("user", user).Msg("found admin user via registry")
|
||||||
|
return user
|
||||||
|
}
|
||||||
|
|
||||||
|
// Method 3: Fallback to lowest RID from registry
|
||||||
|
logger.Msg("attempting to discover user via registry lowest RID")
|
||||||
|
if user := getLowestRegistryUser(ctx); user != "" {
|
||||||
|
logger.Str("method", "registry").Str("user", user).Msg("found user via registry")
|
||||||
|
return user
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.Msg("all user discovery methods failed")
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
|
||||||
|
// getActiveConsoleUser gets the username of the active console session
|
||||||
|
func getActiveConsoleUser(ctx context.Context) string {
|
||||||
|
// Guard against missing WTS procedures (e.g., Windows Server Core).
|
||||||
|
if err := procWTSGetActiveConsoleSessionId.Find(); err != nil {
|
||||||
|
ProxyLogger.Load().Debug().Err(err).Msg("WTSGetActiveConsoleSessionId not available, skipping console session check")
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
sessionId, _, _ := procWTSGetActiveConsoleSessionId.Call()
|
||||||
|
if sessionId == 0xFFFFFFFF { // Invalid session
|
||||||
|
ProxyLogger.Load().Debug().Msg("no active console session found")
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
var buffer uintptr
|
||||||
|
var bytesReturned uint32
|
||||||
|
|
||||||
|
if err := procWTSQuerySessionInformation.Find(); err != nil {
|
||||||
|
ProxyLogger.Load().Debug().Err(err).Msg("WTSQuerySessionInformationW not available")
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
ret, _, _ := procWTSQuerySessionInformation.Call(
|
||||||
|
0, // WTS_CURRENT_SERVER_HANDLE
|
||||||
|
sessionId,
|
||||||
|
uintptr(WTSUserName),
|
||||||
|
uintptr(unsafe.Pointer(&buffer)),
|
||||||
|
uintptr(unsafe.Pointer(&bytesReturned)),
|
||||||
|
)
|
||||||
|
|
||||||
|
if ret == 0 {
|
||||||
|
ProxyLogger.Load().Debug().Msg("failed to query session information")
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
defer procWTSFreeMemory.Call(buffer)
|
||||||
|
|
||||||
|
// Convert buffer to string
|
||||||
|
username := windows.UTF16PtrToString((*uint16)(unsafe.Pointer(buffer)))
|
||||||
|
if username == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
return username
|
||||||
|
}
|
||||||
|
|
||||||
|
// getRegistryUserWithAdminPreference scans registry profiles and prefers admin users
|
||||||
|
func getRegistryUserWithAdminPreference(ctx context.Context) string {
|
||||||
|
profiles := getRegistryProfiles()
|
||||||
|
if len(profiles) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
var adminProfiles []registryProfile
|
||||||
|
var regularProfiles []registryProfile
|
||||||
|
|
||||||
|
// Separate admin and regular users
|
||||||
|
for _, profile := range profiles {
|
||||||
|
if isUserInAdministratorsGroup(profile.username) {
|
||||||
|
adminProfiles = append(adminProfiles, profile)
|
||||||
|
} else {
|
||||||
|
regularProfiles = append(regularProfiles, profile)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Prefer admin users, then regular users
|
||||||
|
candidates := adminProfiles
|
||||||
|
if len(candidates) == 0 {
|
||||||
|
candidates = regularProfiles
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(candidates) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Return user with lowest RID (deterministic choice)
|
||||||
|
lowestRID := candidates[0].rid
|
||||||
|
result := candidates[0].username
|
||||||
|
|
||||||
|
for _, candidate := range candidates[1:] {
|
||||||
|
if candidate.rid < lowestRID {
|
||||||
|
lowestRID = candidate.rid
|
||||||
|
result = candidate.username
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// getLowestRegistryUser returns the user with lowest RID from registry
|
||||||
|
func getLowestRegistryUser(ctx context.Context) string {
|
||||||
|
profiles := getRegistryProfiles()
|
||||||
|
if len(profiles) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Return user with lowest RID (deterministic choice)
|
||||||
|
lowestRID := profiles[0].rid
|
||||||
|
result := profiles[0].username
|
||||||
|
|
||||||
|
for _, profile := range profiles[1:] {
|
||||||
|
if profile.rid < lowestRID {
|
||||||
|
lowestRID = profile.rid
|
||||||
|
result = profile.username
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
type registryProfile struct {
|
||||||
|
username string
|
||||||
|
rid uint32
|
||||||
|
sid string
|
||||||
|
}
|
||||||
|
|
||||||
|
// getRegistryProfiles scans the registry ProfileList for user profiles
|
||||||
|
func getRegistryProfiles() []registryProfile {
|
||||||
|
key, err := registry.OpenKey(registry.LOCAL_MACHINE, `SOFTWARE\Microsoft\Windows NT\CurrentVersion\ProfileList`, registry.ENUMERATE_SUB_KEYS)
|
||||||
|
if err != nil {
|
||||||
|
ProxyLogger.Load().Debug().Err(err).Msg("failed to open ProfileList registry key")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
defer key.Close()
|
||||||
|
|
||||||
|
subkeys, err := key.ReadSubKeyNames(-1)
|
||||||
|
if err != nil {
|
||||||
|
ProxyLogger.Load().Debug().Err(err).Msg("failed to read ProfileList subkeys")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var profiles []registryProfile
|
||||||
|
|
||||||
|
for _, subkey := range subkeys {
|
||||||
|
// Only process SIDs that start with S-1-5-21 (domain/local user accounts)
|
||||||
|
if !strings.HasPrefix(subkey, "S-1-5-21-") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
profileKey, err := registry.OpenKey(key, subkey, registry.QUERY_VALUE)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
profileImagePath, _, err := profileKey.GetStringValue("ProfileImagePath")
|
||||||
|
profileKey.Close()
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract username from profile path (e.g., C:\Users\username)
|
||||||
|
pathParts := strings.Split(profileImagePath, `\`)
|
||||||
|
if len(pathParts) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
username := pathParts[len(pathParts)-1]
|
||||||
|
|
||||||
|
// Extract RID from SID (last component after final hyphen)
|
||||||
|
sidParts := strings.Split(subkey, "-")
|
||||||
|
if len(sidParts) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
ridStr := sidParts[len(sidParts)-1]
|
||||||
|
rid, err := strconv.ParseUint(ridStr, 10, 32)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only consider regular users (RID >= 1000, excludes built-in accounts).
|
||||||
|
// rid == 500 is the default Administrator account (DOMAIN_USER_RID_ADMIN).
|
||||||
|
// See: https://learn.microsoft.com/en-us/windows/win32/secauthz/well-known-sids
|
||||||
|
if rid == 500 || rid >= 1000 {
|
||||||
|
profiles = append(profiles, registryProfile{
|
||||||
|
username: username,
|
||||||
|
rid: uint32(rid),
|
||||||
|
sid: subkey,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return profiles
|
||||||
|
}
|
||||||
|
|
||||||
|
// isUserInAdministratorsGroup checks if a user is in the Administrators group
|
||||||
|
func isUserInAdministratorsGroup(username string) bool {
|
||||||
|
// Open the user account
|
||||||
|
usernamePtr, err := syscall.UTF16PtrFromString(username)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
var userSID *windows.SID
|
||||||
|
var domain *uint16
|
||||||
|
var userSIDSize, domainSize uint32
|
||||||
|
var use uint32
|
||||||
|
|
||||||
|
// First call to get buffer sizes
|
||||||
|
err = windows.LookupAccountName(nil, usernamePtr, userSID, &userSIDSize, domain, &domainSize, &use)
|
||||||
|
if err != nil && err != windows.ERROR_INSUFFICIENT_BUFFER {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Allocate buffers and make actual call
|
||||||
|
userSID = (*windows.SID)(unsafe.Pointer(&make([]byte, userSIDSize)[0]))
|
||||||
|
domain = (*uint16)(unsafe.Pointer(&make([]uint16, domainSize)[0]))
|
||||||
|
|
||||||
|
err = windows.LookupAccountName(nil, usernamePtr, userSID, &userSIDSize, domain, &domainSize, &use)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if user is member of Administrators group (S-1-5-32-544)
|
||||||
|
adminSID, err := windows.CreateWellKnownSid(windows.WinBuiltinAdministratorsSid)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Open user token (this is a simplified check)
|
||||||
|
var token windows.Token
|
||||||
|
err = windows.OpenProcessToken(windows.CurrentProcess(), windows.TOKEN_QUERY, &token)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
defer token.Close()
|
||||||
|
|
||||||
|
// Check group membership
|
||||||
|
member, err := token.IsMember(adminSID)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
return member
|
||||||
|
}
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
package ctrld
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/miekg/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
// SetCacheReply extracts and stores the necessary data from the message for a cached answer.
|
||||||
|
func SetCacheReply(answer, msg *dns.Msg, code int) {
|
||||||
|
answer.SetRcode(msg, code)
|
||||||
|
cCookie := getEdns0Cookie(msg.IsEdns0())
|
||||||
|
sCookie := getEdns0Cookie(answer.IsEdns0())
|
||||||
|
if cCookie != nil && sCookie != nil {
|
||||||
|
// Client cookie is fixed size 8 bytes, Server cookie is variable size 8 -> 32 bytes.
|
||||||
|
// See https://datatracker.ietf.org/doc/html/rfc7873#section-4
|
||||||
|
sCookie.Cookie = cCookie.Cookie[:16] + sCookie.Cookie[16:]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// getEdns0Cookie returns Edns0 cookie from *dns.OPT if present.
|
||||||
|
func getEdns0Cookie(opt *dns.OPT) *dns.EDNS0_COOKIE {
|
||||||
|
if opt == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
for _, o := range opt.Option {
|
||||||
|
if e, ok := o.(*dns.EDNS0_COOKIE); ok {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+4
-3
@@ -1,4 +1,4 @@
|
|||||||
# Using Debian bullseye for building regular image.
|
# Using Debian bookworm for building regular image.
|
||||||
# Using scratch image for minimal image size.
|
# Using scratch image for minimal image size.
|
||||||
# The final image has:
|
# The final image has:
|
||||||
#
|
#
|
||||||
@@ -8,11 +8,12 @@
|
|||||||
# - Non-cgo ctrld binary.
|
# - Non-cgo ctrld binary.
|
||||||
#
|
#
|
||||||
# CI_COMMIT_TAG is used to set the version of ctrld binary.
|
# CI_COMMIT_TAG is used to set the version of ctrld binary.
|
||||||
FROM golang:1.20-bullseye as base
|
FROM golang:1.25-bookworm AS base
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
RUN apt-get update && apt-get install -y upx-ucl
|
RUN echo "deb http://deb.debian.org/debian bookworm-backports main" | tee /etc/apt/sources.list.d/backports.list
|
||||||
|
RUN apt update && apt install -t bookworm-backports upx-ucl
|
||||||
|
|
||||||
COPY . .
|
COPY . .
|
||||||
|
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
# Using Debian bullseye for building regular image.
|
# Using Debian bookworm for building regular image.
|
||||||
# Using scratch image for minimal image size.
|
# Using scratch image for minimal image size.
|
||||||
# The final image has:
|
# The final image has:
|
||||||
#
|
#
|
||||||
@@ -8,11 +8,12 @@
|
|||||||
# - Non-cgo ctrld binary.
|
# - Non-cgo ctrld binary.
|
||||||
#
|
#
|
||||||
# CI_COMMIT_TAG is used to set the version of ctrld binary.
|
# CI_COMMIT_TAG is used to set the version of ctrld binary.
|
||||||
FROM golang:1.20-bullseye as base
|
FROM golang:1.25-bookworm AS base
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
RUN apt-get update && apt-get install -y upx-ucl
|
RUN echo "deb http://deb.debian.org/debian bookworm-backports main" | tee /etc/apt/sources.list.d/backports.list
|
||||||
|
RUN apt update && apt install -t bookworm-backports upx-ucl
|
||||||
|
|
||||||
COPY . .
|
COPY . .
|
||||||
|
|
||||||
|
|||||||
+89
-11
@@ -14,7 +14,7 @@ The config file allows for advanced configuration of the `ctrld` utility to cove
|
|||||||
|
|
||||||
|
|
||||||
## Config Location
|
## Config Location
|
||||||
`ctrld` uses [TOML](toml_link) format for its configuration file. Default configuration file is `ctrld.toml` found in following order:
|
`ctrld` uses [TOML][toml_link] format for its configuration file. Default configuration file is `ctrld.toml` found in following order:
|
||||||
|
|
||||||
- `/etc/controld` on *nix.
|
- `/etc/controld` on *nix.
|
||||||
- User's home directory on Windows.
|
- User's home directory on Windows.
|
||||||
@@ -157,9 +157,15 @@ stale cached records (regardless of their TTLs) until upstream comes online.
|
|||||||
- Required: no
|
- Required: no
|
||||||
- Default: false
|
- Default: false
|
||||||
|
|
||||||
|
### cache_flush_domains
|
||||||
|
When `ctrld` receives query with domain name in `cache_flush_domains`, the local cache will be discarded
|
||||||
|
before serving the query.
|
||||||
|
|
||||||
|
- Type: array of strings
|
||||||
|
- Required: no
|
||||||
|
|
||||||
### max_concurrent_requests
|
### max_concurrent_requests
|
||||||
The number of concurrent requests that will be handled, must be a non-negative integer.
|
The number of concurrent requests that will be handled, must be a non-negative integer.
|
||||||
Tweaking this value depends on the capacity of your system.
|
|
||||||
|
|
||||||
- Type: number
|
- Type: number
|
||||||
- Required: no
|
- Required: no
|
||||||
@@ -172,6 +178,8 @@ Perform LAN client discovery using mDNS. This will spawn a listener on port 5353
|
|||||||
- Required: no
|
- Required: no
|
||||||
- Default: true
|
- Default: true
|
||||||
|
|
||||||
|
This config is ignored, and always set to `false` on Windows Desktop and Macos.
|
||||||
|
|
||||||
### discover_arp
|
### discover_arp
|
||||||
Perform LAN client discovery using ARP.
|
Perform LAN client discovery using ARP.
|
||||||
|
|
||||||
@@ -179,6 +187,8 @@ Perform LAN client discovery using ARP.
|
|||||||
- Required: no
|
- Required: no
|
||||||
- Default: true
|
- Default: true
|
||||||
|
|
||||||
|
This config is ignored, and always set to `false` on Windows Desktop and Macos.
|
||||||
|
|
||||||
### discover_dhcp
|
### discover_dhcp
|
||||||
Perform LAN client discovery using DHCP leases files. Common file locations are auto-discovered.
|
Perform LAN client discovery using DHCP leases files. Common file locations are auto-discovered.
|
||||||
|
|
||||||
@@ -186,6 +196,8 @@ Perform LAN client discovery using DHCP leases files. Common file locations are
|
|||||||
- Required: no
|
- Required: no
|
||||||
- Default: true
|
- Default: true
|
||||||
|
|
||||||
|
This config is ignored, and always set to `false` on Windows Desktop and Macos.
|
||||||
|
|
||||||
### discover_ptr
|
### discover_ptr
|
||||||
Perform LAN client discovery using PTR queries.
|
Perform LAN client discovery using PTR queries.
|
||||||
|
|
||||||
@@ -193,6 +205,8 @@ Perform LAN client discovery using PTR queries.
|
|||||||
- Required: no
|
- Required: no
|
||||||
- Default: true
|
- Default: true
|
||||||
|
|
||||||
|
This config is ignored, and always set to `false` on Windows Desktop and Macos.
|
||||||
|
|
||||||
### discover_hosts
|
### discover_hosts
|
||||||
Perform LAN client discovery using hosts file.
|
Perform LAN client discovery using hosts file.
|
||||||
|
|
||||||
@@ -200,6 +214,8 @@ Perform LAN client discovery using hosts file.
|
|||||||
- Required: no
|
- Required: no
|
||||||
- Default: true
|
- Default: true
|
||||||
|
|
||||||
|
This config is ignored, and always set to `false` on Windows Desktop and Macos.
|
||||||
|
|
||||||
### discover_refresh_interval
|
### discover_refresh_interval
|
||||||
Time in seconds between each discovery refresh loop to update new client information data.
|
Time in seconds between each discovery refresh loop to update new client information data.
|
||||||
The default value is 120 seconds, lower this value to make the discovery process run more aggressively.
|
The default value is 120 seconds, lower this value to make the discovery process run more aggressively.
|
||||||
@@ -220,15 +236,12 @@ DHCP leases file format.
|
|||||||
|
|
||||||
- Type: string
|
- Type: string
|
||||||
- Required: no
|
- Required: no
|
||||||
- Valid values: `dnsmasq`, `isc-dhcp`
|
- Valid values: `dnsmasq`, `isc-dhcp`, `kea-dhcp4`
|
||||||
- Default: ""
|
- Default: ""
|
||||||
|
|
||||||
### client_id_preference
|
### client_id_preference
|
||||||
Decide how the client ID is generated
|
Decide how the client ID is generated. By default client ID will use both MAC address and Hostname i.e. `hash(mac + host)`. To override this behavior, select one of the 2 allowed values to scope client ID to just MAC address OR Hostname.
|
||||||
|
|
||||||
If `host` -> client id will only use the hostname i.e.`hash(hostname)`.
|
|
||||||
If `mac` -> client id will only use the MAC address `hash(mac)`.
|
|
||||||
Else -> client ID will use both Mac and Hostname i.e. `hash(mac + host)
|
|
||||||
- Type: string
|
- Type: string
|
||||||
- Required: no
|
- Required: no
|
||||||
- Valid values: `mac`, `host`
|
- Valid values: `mac`, `host`
|
||||||
@@ -242,12 +255,62 @@ If set to `true`, collect and export the query counters, and show them in `clien
|
|||||||
- Default: false
|
- Default: false
|
||||||
|
|
||||||
### metrics_listener
|
### metrics_listener
|
||||||
Specifying the `ip` and `port` of the metrics server.
|
Specifying the `ip` and `port` of the Prometheus metrics server. The Prometheus metrics will be available on: `http://ip:port/metrics`. You can also append `/metrics/json` to get the same data in json format.
|
||||||
|
|
||||||
- Type: string
|
- Type: string
|
||||||
- Required: no
|
- Required: no
|
||||||
- Default: ""
|
- Default: ""
|
||||||
|
|
||||||
|
### dns_watchdog_enabled
|
||||||
|
Watches all physical interfaces for DNS changes and reverts them to ctrld's settings.The DNS watchdog process only runs on Windows and MacOS.
|
||||||
|
|
||||||
|
- Type: boolean
|
||||||
|
- Required: no
|
||||||
|
- Default: true
|
||||||
|
|
||||||
|
### dns_watchdog_interval
|
||||||
|
Time duration between each DNS watchdog iteration.
|
||||||
|
|
||||||
|
A duration string is a possibly signed sequence of decimal numbers, each with optional fraction and a unit suffix,
|
||||||
|
such as "300ms", "-1.5h" or "2h45m". Valid time units are "ns", "us" (or "µs"), "ms", "s", "m", "h".
|
||||||
|
|
||||||
|
If the time duration is non-positive, default value will be used.
|
||||||
|
|
||||||
|
- Type: time duration string
|
||||||
|
- Required: no
|
||||||
|
- Default: 20s
|
||||||
|
|
||||||
|
### refetch_time
|
||||||
|
Time in seconds between each iteration that reloads custom config from the API.
|
||||||
|
|
||||||
|
The value must be a positive number, any invalid value will be ignored and default value will be used.
|
||||||
|
- Type: number
|
||||||
|
- Required: no
|
||||||
|
- Default: 3600
|
||||||
|
|
||||||
|
### leak_on_upstream_failure
|
||||||
|
If a remote upstream fails to resolve a query or is unreachable, `ctrld` will forward the queries to the default DNS resolver on the network. If failures persist, `ctrld` will remove itself from all networking interfaces until connectivity is restored.
|
||||||
|
|
||||||
|
- Type: boolean
|
||||||
|
- Required: no
|
||||||
|
- Default: true on Windows, MacOS and non-router Linux.
|
||||||
|
|
||||||
|
### nrpt_recovery_max_attempts
|
||||||
|
Windows DNS intercept mode uses NRPT health probes and recovery when Windows stops routing queries to the local `ctrld` listener. This limits how many consecutive recovery flows can run before `ctrld` enters a cooldown and stops making policy/Dnscache changes.
|
||||||
|
|
||||||
|
Set to `0` to disable this circuit breaker and keep retrying indefinitely.
|
||||||
|
|
||||||
|
- Type: integer
|
||||||
|
- Required: no
|
||||||
|
- Default: 0 (unlimited, current behavior)
|
||||||
|
|
||||||
|
### nrpt_recovery_cooldown
|
||||||
|
Cooldown duration after `nrpt_recovery_max_attempts` consecutive Windows NRPT recovery flows. During cooldown, `ctrld` logs the suppressed recovery and avoids additional `RefreshPolicyEx`, Dnscache `paramchange`, and DNS cache flush calls.
|
||||||
|
|
||||||
|
- Type: time duration string
|
||||||
|
- Required: no
|
||||||
|
- Default: 30m
|
||||||
|
|
||||||
## Upstream
|
## Upstream
|
||||||
The `[upstream]` section specifies the DNS upstream servers that `ctrld` will forward DNS requests to.
|
The `[upstream]` section specifies the DNS upstream servers that `ctrld` will forward DNS requests to.
|
||||||
|
|
||||||
@@ -332,7 +395,7 @@ The protocol that `ctrld` will use to send DNS requests to upstream.
|
|||||||
|
|
||||||
- Type: string
|
- Type: string
|
||||||
- Required: yes
|
- Required: yes
|
||||||
- Valid values: `doh`, `doh3`, `dot`, `doq`, `legacy`, `os`
|
- Valid values: `doh`, `doh3`, `dot`, `doq`, `legacy`
|
||||||
|
|
||||||
### ip_stack
|
### ip_stack
|
||||||
Specifying what kind of ip stack that `ctrld` will use to connect to upstream.
|
Specifying what kind of ip stack that `ctrld` will use to connect to upstream.
|
||||||
@@ -491,6 +554,15 @@ rules = [
|
|||||||
]
|
]
|
||||||
```
|
```
|
||||||
|
|
||||||
|
If there is no explicitly defined rules, LAN queries will be handled solely by the OS resolver.
|
||||||
|
|
||||||
|
These following domains are considered LAN queries:
|
||||||
|
|
||||||
|
- Queries does not have dot `.` in domain name, like `machine1`, `example`, ... (1)
|
||||||
|
- Queries have domain ends with: `.domain`, `.lan`, `.local`. (2)
|
||||||
|
- All `SRV` queries of LAN hostname (1) + (2).
|
||||||
|
- `PTR` queries with private IPs.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
Note that the order of matching preference:
|
Note that the order of matching preference:
|
||||||
@@ -524,6 +596,12 @@ And within each policy, the rules are processed from top to bottom.
|
|||||||
- Required: no
|
- Required: no
|
||||||
- Default: []
|
- Default: []
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
Note that the domain comparisons are done in case in-sensitive manner following [RFC 1034](https://datatracker.ietf.org/doc/html/rfc1034#section-3.1)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
### macs:
|
### macs:
|
||||||
`macs` is the list of mac rules within the policy. Mac address value is case-insensitive.
|
`macs` is the list of mac rules within the policy. Mac address value is case-insensitive.
|
||||||
|
|
||||||
@@ -534,7 +612,7 @@ And within each policy, the rules are processed from top to bottom.
|
|||||||
### failover_rcodes
|
### failover_rcodes
|
||||||
For non success response, `failover_rcodes` allows the request to be forwarded to next upstream, if the response `RCODE` matches any value defined in `failover_rcodes`.
|
For non success response, `failover_rcodes` allows the request to be forwarded to next upstream, if the response `RCODE` matches any value defined in `failover_rcodes`.
|
||||||
|
|
||||||
- Type: array of string
|
- Type: array of strings
|
||||||
- Required: no
|
- Required: no
|
||||||
- Default: []
|
- Default: []
|
||||||
-
|
-
|
||||||
@@ -551,7 +629,7 @@ networks = [
|
|||||||
|
|
||||||
If `upstream.0` returns a NXDOMAIN response, the request will be forwarded to `upstream.1` instead of returning immediately to the client.
|
If `upstream.0` returns a NXDOMAIN response, the request will be forwarded to `upstream.1` instead of returning immediately to the client.
|
||||||
|
|
||||||
See all available DNS Rcodes value [here](rcode_link).
|
See all available DNS Rcodes value [here][rcode_link].
|
||||||
|
|
||||||
[toml_link]: https://toml.io/en
|
[toml_link]: https://toml.io/en
|
||||||
[rcode_link]: https://www.iana.org/assignments/dns-parameters/dns-parameters.xhtml#dns-parameters-6
|
[rcode_link]: https://www.iana.org/assignments/dns-parameters/dns-parameters.xhtml#dns-parameters-6
|
||||||
|
|||||||
Binary file not shown.
|
After Width: | Height: | Size: 458 KiB |
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user