Skip to content
File

Blob: tun/client/config.go

go233 lines
1package client
2 
3import (
4 "errors"
5 "fmt"
6 "net/url"
7 "os"
8 "runtime"
9 "strings"
10 "time"
11 
12 "go.miragespace.co/specter/spec/pki"
13 
14 "github.com/zhangyunhao116/skipmap"
15 "go.yaml.in/yaml/v3"
16)
17 
18type Tunnel struct {
19 parsed *url.URL
20 Target string `yaml:"target" json:"target"`
21 Hostname string `yaml:"hostname,omitempty" json:"hostname,omitempty"`
22 Insecure bool `yaml:"insecure,omitempty" json:"insecure"`
23 ProxyHeaderTimeout time.Duration `yaml:"headerTimeout,omitempty" json:"headerTimeout,omitempty"`
24 ProxyHeaderHost string `yaml:"headerHost,omitempty" json:"headerHost,omitempty"`
25 ProxyHeaderMode string `yaml:"headerMode,omitempty" json:"headerMode,omitempty"`
26}
27 
28func (t *Tunnel) UnmarshalYAML(value *yaml.Node) error {
29 type tunnelAlias Tunnel
30 var aux struct {
31 tunnelAlias `yaml:",inline"`
32 // legacy names for backward compatibility
33 OldHeaderTimeout time.Duration `yaml:"proxyHeaderTimeout,omitempty"`
34 OldHeaderHost string `yaml:"proxyHeaderHost,omitempty"`
35 OldHeaderMode string `yaml:"proxyHeaderMode,omitempty"`
36 }
37 if err := value.Decode(&aux); err != nil {
38 return err
39 }
40 *t = Tunnel(aux.tunnelAlias)
41 
42 // If new header* fields are not set, fall back to legacy proxyHeader* values.
43 // This ensures configs using header* take precedence when both are present.
44 if aux.OldHeaderTimeout != 0 && t.ProxyHeaderTimeout == 0 {
45 t.ProxyHeaderTimeout = aux.OldHeaderTimeout
46 }
47 if aux.OldHeaderHost != "" && t.ProxyHeaderHost == "" {
48 t.ProxyHeaderHost = aux.OldHeaderHost
49 }
50 if aux.OldHeaderMode != "" && t.ProxyHeaderMode == "" {
51 t.ProxyHeaderMode = aux.OldHeaderMode
52 }
53 
54 return nil
55}
56 
57type Config struct {
58 router *skipmap.StringMap[route]
59 path string
60 Version int `yaml:"version" json:"version"`
61 Apex string `yaml:"apex" json:"apex"`
62 Certificate string `yaml:"certificate,omitempty" json:"certificate,omitempty"`
63 PrivKey string `yaml:"privKey,omitempty" json:"privKey,omitempty"`
64 Tunnels []Tunnel `yaml:"tunnels,omitempty" json:"tunnels,omitempty"`
65}
66 
67type route struct {
68 parsed *url.URL
69 insecure bool
70 proxyHeaderReadTimeout time.Duration
71 proxyHeaderHost string
72 proxyHeaderMode string
73}
74 
75func NewConfig(path string) (*Config, error) {
76 cfg := &Config{
77 path: path,
78 router: skipmap.NewString[route](),
79 }
80 if err := cfg.readFile(); err != nil {
81 return nil, err
82 }
83 if err := cfg.checkVersion(); err != nil {
84 return nil, err
85 }
86 if err := cfg.validate(); err != nil {
87 return nil, err
88 }
89 return cfg, nil
90}
91 
92func (c *Config) clone() *Config {
93 cfg := *c
94 cfg.router = skipmap.NewString[route]()
95 cfg.Tunnels = make([]Tunnel, len(c.Tunnels))
96 for i := range c.Tunnels {
97 cfg.Tunnels[i] = c.Tunnels[i]
98 cfg.Tunnels[i].parsed = nil
99 }
100 cfg.validate()
101 return &cfg
102}
103 
104func (c *Config) buildRouter(drop ...Tunnel) {
105 for _, tunnel := range drop {
106 c.router.Delete(tunnel.Hostname)
107 }
108 for _, tunnel := range c.Tunnels {
109 c.router.Store(tunnel.Hostname, route{
110 parsed: tunnel.parsed,
111 insecure: tunnel.Insecure,
112 proxyHeaderReadTimeout: tunnel.ProxyHeaderTimeout,
113 proxyHeaderHost: tunnel.ProxyHeaderHost,
114 proxyHeaderMode: tunnel.ProxyHeaderMode,
115 })
116 }
117}
118 
119func (c *Config) checkVersion() error {
120 if c.Version != 2 {
121 return fmt.Errorf("expecting config version 2, got %v; migration is needed", c.Version)
122 }
123 return nil
124}
125 
126func (c *Config) validate() error {
127 for i, tunnel := range c.Tunnels {
128 u, err := parseTarget(tunnel.Target)
129 if err != nil {
130 return err
131 }
132 
133 switch tunnel.ProxyHeaderMode {
134 case "", "target", "hostname", "custom":
135 default:
136 return fmt.Errorf("unsupported proxyHeaderMode %q", tunnel.ProxyHeaderMode)
137 }
138 
139 if tunnel.ProxyHeaderMode == "custom" && tunnel.ProxyHeaderHost == "" {
140 return fmt.Errorf("proxyHeaderMode 'custom' requires proxyHeaderHost to be set")
141 }
142 
143 if (u.Scheme == "unix" || u.Scheme == "winio") && tunnel.ProxyHeaderMode == "target" {
144 return fmt.Errorf("proxyHeaderMode 'target' is not valid for pipe targets; use 'hostname' or 'custom'")
145 }
146 c.Tunnels[i].parsed = u
147 }
148 if c.PrivKey == "" {
149 _, c.PrivKey = pki.GeneratePrivKey()
150 }
151 return nil
152}
153 
154func (c *Config) reloadFile(callbacks ...func(prev, curr []Tunnel)) error {
155 f, err := os.Open(c.path)
156 if err != nil {
157 return fmt.Errorf("error opening config file for reading: %w", err)
158 }
159 defer f.Close()
160 
161 next := &Config{}
162 if err := yaml.NewDecoder(f).Decode(next); err != nil {
163 return fmt.Errorf("error decoding config file: %w", err)
164 }
165 
166 if err := next.validate(); err != nil {
167 return fmt.Errorf("error validating config file: %w", err)
168 }
169 
170 prev := c.clone()
171 c.Tunnels = next.Tunnels
172 
173 for _, cb := range callbacks {
174 cb(prev.Tunnels, next.Tunnels)
175 }
176 return nil
177}
178 
179func (c *Config) readFile() error {
180 f, err := os.Open(c.path)
181 if err != nil {
182 return fmt.Errorf("error opening config file for reading: %w", err)
183 }
184 defer f.Close()
185 return yaml.NewDecoder(f).Decode(c)
186}
187 
188func (c *Config) writeFile() (err error) {
189 f, err := os.OpenFile(c.path, os.O_RDWR|os.O_CREATE|os.O_TRUNC, 0644)
190 if err != nil {
191 return fmt.Errorf("error opening config file for writing: %w", err)
192 }
193 defer func() { err = errors.Join(err, f.Close()) }()
194 
195 encoder := yaml.NewEncoder(f)
196 encoder.SetIndent(2)
197 if err := encoder.Encode(c); err != nil {
198 return err
199 }
200 if err := encoder.Close(); err != nil {
201 return err
202 }
203 return f.Sync()
204}
205 
206func parseTarget(target string) (*url.URL, error) {
207 u, err := url.Parse(target)
208 if err != nil {
209 return nil, fmt.Errorf("error parsing target %s: %w", target, err)
210 }
211 switch u.Scheme {
212 case "http", "https", "tcp", "unix":
213 default:
214 if strings.HasPrefix(u.Path, "\\\\.\\pipe") {
215 u.Scheme = "winio"
216 break
217 }
218 return nil, fmt.Errorf("unsupported scheme. valid schemes: http, https, tcp, or unix; got %s", u.Scheme)
219 }
220 switch u.Scheme {
221 case "winio":
222 if runtime.GOOS != "windows" {
223 return nil, errors.New("named pipe is not supported on non-Windows platform")
224 }
225 case "unix":
226 if runtime.GOOS == "windows" {
227 return nil, errors.New("unix socket is not supported on non-Unix platform")
228 }
229 }
230 
231 return u, nil
232}