diff --git a/go.mod b/go.mod index 3085e6b..dee251c 100644 --- a/go.mod +++ b/go.mod @@ -8,6 +8,8 @@ require ( github.com/google/go-cmp v0.7.0 github.com/oklog/run v1.1.0 github.com/prometheus/client_golang v1.21.1 + github.com/prometheus/common v0.63.0 + github.com/tailscale/hujson v0.0.0-20221223112325-20486734a56a tailscale.com v1.82.0 ) @@ -64,13 +66,11 @@ require ( github.com/pierrec/lz4/v4 v4.1.21 // indirect github.com/prometheus-community/pro-bing v0.4.0 // indirect github.com/prometheus/client_model v0.6.1 // indirect - github.com/prometheus/common v0.63.0 // indirect github.com/prometheus/procfs v0.15.1 // indirect github.com/safchain/ethtool v0.3.0 // indirect github.com/tailscale/certstore v0.1.1-0.20231202035212-d3fa0460f47e // indirect github.com/tailscale/go-winio v0.0.0-20231025203758-c4f33415bf55 // indirect github.com/tailscale/goupnp v1.0.1-0.20210804011211-c64d0f06ea05 // indirect - github.com/tailscale/hujson v0.0.0-20221223112325-20486734a56a // indirect github.com/tailscale/netlink v1.1.1-0.20240822203006-4d49adab4de7 // indirect github.com/tailscale/peercred v0.0.0-20250107143737-35a0c7bd7edc // indirect github.com/tailscale/web-client-prebuilt v0.0.0-20250124233751-d4cd19a26976 // indirect diff --git a/main.go b/main.go index f810eb5..0d0c503 100644 --- a/main.go +++ b/main.go @@ -4,7 +4,6 @@ import ( "context" "crypto/tls" "encoding/json" - "errors" "flag" "fmt" "log/slog" @@ -16,7 +15,6 @@ import ( "path/filepath" "sort" "strconv" - "strings" "syscall" "github.com/oklog/run" @@ -25,6 +23,7 @@ import ( "github.com/prometheus/client_golang/prometheus/promauto" "github.com/prometheus/client_golang/prometheus/promhttp" "github.com/prometheus/common/version" + "github.com/tailscale/hujson" "tailscale.com/client/local" "tailscale.com/client/tailscale/apitype" "tailscale.com/tsnet" @@ -61,26 +60,11 @@ var ( ) ) -type upstreamFlag []upstream - -func (f *upstreamFlag) String() string { - return fmt.Sprintf("%+v", *f) -} - -func (f *upstreamFlag) Set(val string) error { - up, err := parseUpstreamFlag(val) - if err != nil { - return err - } - *f = append(*f, up) - return nil -} - type upstream struct { - name string - backend *url.URL - prometheus bool - funnel bool + Name string + Backend string + Prometheus bool + Funnel bool } type target struct { @@ -89,32 +73,6 @@ type target struct { prometheus bool } -func parseUpstreamFlag(fval string) (upstream, error) { - k, v, ok := strings.Cut(fval, "=") - if !ok { - return upstream{}, errors.New("format: name=http://backend") - } - val := strings.Split(v, ";") - be, err := url.Parse(val[0]) - if err != nil { - return upstream{}, err - } - up := upstream{name: k, backend: be} - if len(val) > 1 { - for _, opt := range val[1:] { - switch opt { - case "prometheus": - up.prometheus = true - case "funnel": - up.funnel = true - default: - return upstream{}, fmt.Errorf("unsupported option: %v", opt) - } - } - } - return up, nil -} - func main() { if err := tsproxy(context.Background()); err != nil { fmt.Fprintf(os.Stderr, "tsproxy: %v\n", err) @@ -124,13 +82,12 @@ func main() { func tsproxy(ctx context.Context) error { var ( - state = flag.String("state", "", "Optional directory for storing Tailscale state.") - tslog = flag.Bool("tslog", false, "If true, log Tailscale output.") - port = flag.Int("port", 32019, "HTTP port for metrics and service discovery.") - ver = flag.Bool("version", false, "print the version and exit") + state = flag.String("state", "", "Optional directory for storing Tailscale state.") + tslog = flag.Bool("tslog", false, "If true, log Tailscale output.") + port = flag.Int("port", 32019, "HTTP port for metrics and service discovery.") + ver = flag.Bool("version", false, "print the version and exit") + upfile = flag.String("upstream", "", "path to upstreams config file") ) - var upstreams upstreamFlag - flag.Var(&upstreams, "upstream", "Repeated for each upstream. Format: name=http://backend:8000") flag.Parse() if *ver { @@ -138,9 +95,26 @@ func tsproxy(ctx context.Context) error { os.Exit(0) } - if len(upstreams) == 0 { + if *upfile == "" { return fmt.Errorf("required flag missing: upstream") } + + in, err := os.ReadFile(*upfile) + if err != nil { + return err + } + inJSON, err := hujson.Standardize(in) + if err != nil { + return fmt.Errorf("hujson: %w", err) + } + var upstreams []upstream + if err := json.Unmarshal(inJSON, &upstreams); err != nil { + return fmt.Errorf("json: %w", err) + } + if len(upstreams) == 0 { + return fmt.Errorf("file does not contain any upstreams: %s", *upfile) + } + if *state == "" { v, err := os.UserCacheDir() if err != nil { @@ -219,11 +193,11 @@ func tsproxy(ctx context.Context) error { i := i upstream := upstream - log := logger.With(slog.String("upstream", upstream.name)) + log := logger.With(slog.String("upstream", upstream.Name)) ts := &tsnet.Server{ - Hostname: upstream.name, - Dir: filepath.Join(*state, "tailscale-"+upstream.name), + Hostname: upstream.Name, + Dir: filepath.Join(*state, "tailscale-"+upstream.Name), RunWebClient: true, } defer ts.Close() @@ -242,29 +216,34 @@ func tsproxy(ctx context.Context) error { lc, err := ts.LocalClient() if err != nil { - return fmt.Errorf("tailscale: get local client for %s: %w", upstream.name, err) + return fmt.Errorf("tailscale: get local client for %s: %w", upstream.Name, err) + } + + backendURL, err := url.Parse(upstream.Backend) + if err != nil { + return fmt.Errorf("upstream %s: parse backend URL: %w", upstream.Name, err) } srv := &http.Server{ TLSConfig: &tls.Config{GetCertificate: lc.GetCertificate}, - Handler: promhttp.InstrumentHandlerInFlight(requestsInFlight.With(prometheus.Labels{"upstream": upstream.name}), - promhttp.InstrumentHandlerDuration(duration.MustCurryWith(prometheus.Labels{"upstream": upstream.name}), - promhttp.InstrumentHandlerCounter(requests.MustCurryWith(prometheus.Labels{"upstream": upstream.name}), - newReverseProxy(log, lc, upstream.backend)))), + Handler: promhttp.InstrumentHandlerInFlight(requestsInFlight.With(prometheus.Labels{"upstream": upstream.Name}), + promhttp.InstrumentHandlerDuration(duration.MustCurryWith(prometheus.Labels{"upstream": upstream.Name}), + promhttp.InstrumentHandlerCounter(requests.MustCurryWith(prometheus.Labels{"upstream": upstream.Name}), + newReverseProxy(log, lc, backendURL)))), } g.Add(func() error { st, err := ts.Up(ctx) if err != nil { - return fmt.Errorf("tailscale: wait for node %s to be ready: %w", upstream.name, err) + return fmt.Errorf("tailscale: wait for node %s to be ready: %w", upstream.Name, err) } // register in service discovery when we're ready. - targets[i] = target{name: upstream.name, prometheus: upstream.prometheus, magicDNS: st.Self.DNSName} + targets[i] = target{name: upstream.Name, prometheus: upstream.Prometheus, magicDNS: st.Self.DNSName} ln, err := ts.Listen("tcp", ":80") if err != nil { - return fmt.Errorf("tailscale: listen for %s on port 80: %w", upstream.name, err) + return fmt.Errorf("tailscale: listen for %s on port 80: %w", upstream.Name, err) } return srv.Serve(ln) }, func(_ error) { @@ -276,20 +255,20 @@ func tsproxy(ctx context.Context) error { g.Add(func() error { _, err := ts.Up(ctx) if err != nil { - return fmt.Errorf("tailscale: wait for node %s to be ready: %w", upstream.name, err) + return fmt.Errorf("tailscale: wait for node %s to be ready: %w", upstream.Name, err) } - if upstream.funnel { + if upstream.Funnel { ln, err := ts.ListenFunnel("tcp", ":443") if err != nil { - return fmt.Errorf("tailscale: funnel for %s on port 443: %w", upstream.name, err) + return fmt.Errorf("tailscale: funnel for %s on port 443: %w", upstream.Name, err) } return srv.Serve(ln) } ln, err := ts.Listen("tcp", ":443") if err != nil { - return fmt.Errorf("tailscale: listen for %s on port 443: %w", upstream.name, err) + return fmt.Errorf("tailscale: listen for %s on port 443: %w", upstream.Name, err) } return srv.ServeTLS(ln, "", "") }, func(_ error) { diff --git a/tsproxy_test.go b/tsproxy_test.go index 4307237..7619b0e 100644 --- a/tsproxy_test.go +++ b/tsproxy_test.go @@ -10,8 +10,6 @@ import ( "net/http" "net/http/httptest" "net/url" - "reflect" - "strings" "testing" "github.com/google/go-cmp/cmp" @@ -29,67 +27,6 @@ func (c *fakeLocalClient) WhoIs(ctx context.Context, remoteAddr string) (*apityp return c.whois(ctx, remoteAddr) } -func TestParseUpstream(t *testing.T) { - t.Parallel() - - for _, tc := range []struct { - upstream string - want upstream - err error - }{ - { - upstream: "test=http://example.com:-80/", - want: upstream{}, - err: errors.New(`parse "http://`), - }, - { - upstream: "test=http://localhost", - want: upstream{name: "test", backend: mustParseURL("http://localhost")}, - }, - { - upstream: "test=http://localhost;prometheus", - want: upstream{name: "test", backend: mustParseURL("http://localhost"), prometheus: true}, - }, - { - upstream: "test=http://localhost;funnel;prometheus", - want: upstream{name: "test", backend: mustParseURL("http://localhost"), prometheus: true, funnel: true}, - }, - { - upstream: "test=http://localhost;foo", - want: upstream{}, - err: errors.New("unsupported option: foo"), - }, - } { - tc := tc - t.Run(tc.upstream, func(t *testing.T) { - t.Parallel() - up, err := parseUpstreamFlag(tc.upstream) - if tc.err != nil { - if err == nil { - t.Fatalf("want err %v, got nil", tc.err) - } - if !strings.Contains(err.Error(), tc.err.Error()) { - t.Fatalf("want err %v, got %v", tc.err, err) - } - } - if tc.err == nil && err != nil { - t.Fatalf("want no err, got %v", err) - } - if diff := cmp.Diff(tc.want, up, cmp.Exporter(func(_ reflect.Type) bool { return true })); diff != "" { - t.Errorf("mismatch (-want +got):\n%s", diff) - } - }) - } -} - -func mustParseURL(s string) *url.URL { - v, err := url.Parse(s) - if err != nil { - panic(err) - } - return v -} - func TestReverseProxy(t *testing.T) { t.Parallel()