mirror of
https://github.com/Control-D-Inc/ctrld.git
synced 2026-09-29 13:51:51 +02:00
cmd/ctrld: update config when "--cd" present
This commit is contained in:
1 parent
6edd42629e
commit
9f90811567
3 files changed
+50
-30
No files matched your search
+36
-30
@@ -3,6 +3,7 @@ package main
|
|||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
"net"
|
"net"
|
||||||
"os"
|
"os"
|
||||||
@@ -13,7 +14,6 @@ import (
|
|||||||
|
|
||||||
"github.com/go-playground/validator/v10"
|
"github.com/go-playground/validator/v10"
|
||||||
"github.com/kardianos/service"
|
"github.com/kardianos/service"
|
||||||
"github.com/pelletier/go-toml"
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
"github.com/spf13/viper"
|
"github.com/spf13/viper"
|
||||||
|
|
||||||
@@ -64,8 +64,8 @@ func initCLI() {
|
|||||||
log.Fatal("Cannot run in daemon mode. Please install a Windows service.")
|
log.Fatal("Cannot run in daemon mode. Please install a Windows service.")
|
||||||
}
|
}
|
||||||
|
|
||||||
noConfigStart := isNoConfigStart(cmd) && cdUID != ""
|
noConfigStart := isNoConfigStart(cmd)
|
||||||
writeDefaultConfig := !noConfigStart && configBase64 == "" && cdUID == ""
|
writeDefaultConfig := !noConfigStart && configBase64 == ""
|
||||||
configs := []struct {
|
configs := []struct {
|
||||||
name string
|
name string
|
||||||
written bool
|
written bool
|
||||||
@@ -84,10 +84,10 @@ func initCLI() {
|
|||||||
|
|
||||||
readBase64Config()
|
readBase64Config()
|
||||||
processNoConfigFlags(noConfigStart)
|
processNoConfigFlags(noConfigStart)
|
||||||
processCDFlags()
|
|
||||||
if err := v.Unmarshal(&cfg); err != nil {
|
if err := v.Unmarshal(&cfg); err != nil {
|
||||||
log.Fatalf("failed to unmarshal config: %v", err)
|
log.Fatalf("failed to unmarshal config: %v", err)
|
||||||
}
|
}
|
||||||
|
processCDFlags(writeDefaultConfig)
|
||||||
if err := ctrld.ValidateConfig(validator.New(), &cfg); err != nil {
|
if err := ctrld.ValidateConfig(validator.New(), &cfg); err != nil {
|
||||||
log.Fatalf("invalid config: %v", err)
|
log.Fatalf("invalid config: %v", err)
|
||||||
}
|
}
|
||||||
@@ -151,25 +151,35 @@ func initCLI() {
|
|||||||
Short: "Start the ctrld service",
|
Short: "Start the ctrld service",
|
||||||
Args: cobra.NoArgs,
|
Args: cobra.NoArgs,
|
||||||
Run: func(cmd *cobra.Command, args []string) {
|
Run: func(cmd *cobra.Command, args []string) {
|
||||||
cfg := &service.Config{}
|
sc := &service.Config{}
|
||||||
*cfg = *svcConfig
|
*sc = *svcConfig
|
||||||
osArgs := os.Args[2:]
|
osArgs := os.Args[2:]
|
||||||
if os.Args[1] == "service" {
|
if os.Args[1] == "service" {
|
||||||
osArgs = os.Args[3:]
|
osArgs = os.Args[3:]
|
||||||
}
|
}
|
||||||
cfg.Arguments = append([]string{"run"}, osArgs...)
|
sc.Arguments = append([]string{"run"}, osArgs...)
|
||||||
if dir, err := os.UserHomeDir(); err == nil {
|
if dir, err := os.UserHomeDir(); err == nil {
|
||||||
// WorkingDirectory is not supported on Windows.
|
// WorkingDirectory is not supported on Windows.
|
||||||
cfg.WorkingDirectory = dir
|
sc.WorkingDirectory = dir
|
||||||
// No config path, generating config in HOME directory.
|
// No config path, generating config in HOME directory.
|
||||||
noConfigStart := isNoConfigStart(cmd) && cdUID != ""
|
noConfigStart := isNoConfigStart(cmd)
|
||||||
writeDefaultConfig := !noConfigStart && configBase64 == "" && cdUID == ""
|
writeDefaultConfig := !noConfigStart && configBase64 == ""
|
||||||
if configPath == "" && writeDefaultConfig {
|
if configPath == "" && writeDefaultConfig {
|
||||||
defaultConfigFile = filepath.Join(dir, defaultConfigFile)
|
defaultConfigFile = filepath.Join(dir, defaultConfigFile)
|
||||||
readConfigFile(true)
|
readConfigFile(true)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// On Windows, the service will be run as SYSTEM, so if ctrld start as Admin,
|
||||||
|
// the written config won't be writable by SYSTEM account, we have to update
|
||||||
|
// the config here when "--cd" is supplied.
|
||||||
|
if runtime.GOOS == "windows" && cdUID != "" {
|
||||||
|
if err := v.Unmarshal(&cfg); err != nil {
|
||||||
|
log.Fatalf("failed to unmarshal config: %v", err)
|
||||||
|
}
|
||||||
|
processCDFlags(writeDefaultConfig)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
s, err := service.New(&prog{}, cfg)
|
s, err := service.New(&prog{}, sc)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
stderrMsg(err.Error())
|
stderrMsg(err.Error())
|
||||||
return
|
return
|
||||||
@@ -311,12 +321,10 @@ func initCLI() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func writeConfigFile() {
|
func writeConfigFile() {
|
||||||
c := v.AllSettings()
|
if cfu := v.ConfigFileUsed(); cfu != "" {
|
||||||
bs, err := toml.Marshal(c)
|
defaultConfigFile = cfu
|
||||||
if err != nil {
|
|
||||||
log.Fatalf("unable to marshal config to toml: %v", err)
|
|
||||||
}
|
}
|
||||||
if err := os.WriteFile(defaultConfigFile, bs, 0600); err != nil {
|
if err := v.WriteConfigAs(defaultConfigFile); err != nil {
|
||||||
log.Printf("failed to write config file: %v\n", err)
|
log.Printf("failed to write config file: %v\n", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -325,6 +333,7 @@ func readConfigFile(writeDefaultConfig bool) bool {
|
|||||||
// If err == nil, there's a config supplied via `--config`, no default config written.
|
// If err == nil, there's a config supplied via `--config`, no default config written.
|
||||||
err := v.ReadInConfig()
|
err := v.ReadInConfig()
|
||||||
if err == nil {
|
if err == nil {
|
||||||
|
fmt.Println("loading config file from: ", v.ConfigFileUsed())
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -390,7 +399,7 @@ func processNoConfigFlags(noConfigStart bool) {
|
|||||||
processLogAndCacheFlags()
|
processLogAndCacheFlags()
|
||||||
}
|
}
|
||||||
|
|
||||||
func processCDFlags() {
|
func processCDFlags(writeConfig bool) {
|
||||||
if cdUID == "" {
|
if cdUID == "" {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -399,25 +408,22 @@ func processCDFlags() {
|
|||||||
log.Fatalf("failed to fetch resolver config: %v", err)
|
log.Fatalf("failed to fetch resolver config: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
upstream := map[string]*ctrld.UpstreamConfig{
|
u0 := cfg.Upstream["0"]
|
||||||
"0": {
|
u0.Name = resolverConfig.DOH
|
||||||
Name: resolverConfig.DOH,
|
u0.Endpoint = resolverConfig.DOH
|
||||||
Endpoint: resolverConfig.DOH,
|
u0.Type = ctrld.ResolverTypeDOH
|
||||||
Type: ctrld.ResolverTypeDOH,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
v.Set("upstream", upstream)
|
|
||||||
|
|
||||||
processListenFlag()
|
|
||||||
|
|
||||||
rules := make([]ctrld.Rule, 0, len(resolverConfig.Exclude))
|
rules := make([]ctrld.Rule, 0, len(resolverConfig.Exclude))
|
||||||
for _, domain := range resolverConfig.Exclude {
|
for _, domain := range resolverConfig.Exclude {
|
||||||
rules = append(rules, ctrld.Rule{domain: []string{}})
|
rules = append(rules, ctrld.Rule{domain: []string{}})
|
||||||
}
|
}
|
||||||
lc := v.Get("listener").(map[string]*ctrld.ListenerConfig)["0"]
|
cfg.Listener["0"].Policy = &ctrld.ListenerPolicyConfig{Name: "My Policy", Rules: rules}
|
||||||
lc.Policy = &ctrld.ListenerPolicyConfig{Name: "My Policy", Rules: rules}
|
|
||||||
|
|
||||||
processLogAndCacheFlags()
|
if writeConfig {
|
||||||
|
v.Set("listener", cfg.Listener)
|
||||||
|
v.Set("upstream", cfg.Upstream)
|
||||||
|
writeConfigFile()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func processListenFlag() {
|
func processListenFlag() {
|
||||||
|
|||||||
@@ -314,6 +314,17 @@ Above policy will:
|
|||||||
- Forward requests on `listener.0` for `test.com` to `upstream.2`. If timeout is reached, retry on `upstream.1`.
|
- Forward requests on `listener.0` for `test.com` to `upstream.2`. If timeout is reached, retry on `upstream.1`.
|
||||||
- All other requests on `listener.0` that do not match above conditions will be forwarded to `upstream.0`.
|
- All other requests on `listener.0` that do not match above conditions will be forwarded to `upstream.0`.
|
||||||
|
|
||||||
|
An empty upstream would not route the request to any defined upstreams, and use the OS default resolver.
|
||||||
|
|
||||||
|
```toml
|
||||||
|
[listener.0.policy]
|
||||||
|
name = "OS Resolver"
|
||||||
|
|
||||||
|
rules = [
|
||||||
|
{"*.local" = []},
|
||||||
|
]
|
||||||
|
```
|
||||||
|
|
||||||
#### name
|
#### name
|
||||||
`name` is the name for the policy.
|
`name` is the name for the policy.
|
||||||
|
|
||||||
|
|||||||
@@ -41,6 +41,9 @@ func FetchResolverConfig(uid string) (*ResolverConfig, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("http.NewRequest: %w", err)
|
return nil, fmt.Errorf("http.NewRequest: %w", err)
|
||||||
}
|
}
|
||||||
|
q := req.URL.Query()
|
||||||
|
q.Set("platform", "ctrld")
|
||||||
|
req.URL.RawQuery = q.Encode()
|
||||||
req.Header.Add("Content-Type", "application/json")
|
req.Header.Add("Content-Type", "application/json")
|
||||||
client := http.Client{Timeout: 5 * time.Second}
|
client := http.Client{Timeout: 5 * time.Second}
|
||||||
resp, err := client.Do(req)
|
resp, err := client.Do(req)
|
||||||
|
|||||||
Reference in new issue
Block a user