From dc0d9805b2c343c8b05dd755bf4d74f664365844 Mon Sep 17 00:00:00 2001 From: Simon Rozet Date: Sat, 3 Dec 2022 17:42:54 +0100 Subject: [PATCH] lean into magicDNS Create one tsnet for each upstream, which gets us DNS and Funnel (not yet implemented) for free. --- dns.go | 178 --- dns_test.go | 70 -- go.mod | 11 +- go.sum | 11 +- internal/autocert/LICENSE | 24 - internal/autocert/autocert.go | 1225 --------------------- internal/autocert/autocert_test.go | 1063 ------------------ internal/autocert/cache.go | 135 --- internal/autocert/cache_test.go | 66 -- internal/autocert/example_test.go | 35 - internal/autocert/internal/acmetest/ca.go | 817 -------------- internal/autocert/listener.go | 155 --- internal/autocert/renewal.go | 156 --- internal/autocert/renewal_test.go | 270 ----- main.go | 271 ++--- main_test.go | 66 -- tsproxy.go | 106 +- tsproxy_test.go | 174 +-- 18 files changed, 282 insertions(+), 4551 deletions(-) delete mode 100644 dns.go delete mode 100644 dns_test.go delete mode 100644 internal/autocert/LICENSE delete mode 100644 internal/autocert/autocert.go delete mode 100644 internal/autocert/autocert_test.go delete mode 100644 internal/autocert/cache.go delete mode 100644 internal/autocert/cache_test.go delete mode 100644 internal/autocert/example_test.go delete mode 100644 internal/autocert/internal/acmetest/ca.go delete mode 100644 internal/autocert/listener.go delete mode 100644 internal/autocert/renewal.go delete mode 100644 internal/autocert/renewal_test.go delete mode 100644 main_test.go diff --git a/dns.go b/dns.go deleted file mode 100644 index 288da73..0000000 --- a/dns.go +++ /dev/null @@ -1,178 +0,0 @@ -package main - -import ( - "context" - "errors" - "fmt" - "net" - "net/netip" - "strings" - - "github.com/sr/tsproxy/internal/autocert" - - "github.com/cenkalti/backoff/v4" - "github.com/dnsimple/dnsimple-go/dnsimple" - "golang.org/x/exp/slog" - "golang.org/x/sync/errgroup" -) - -type dnsimpleClient interface { - ListRecords(context.Context, string, string, *dnsimple.ZoneRecordListOptions) (*dnsimple.ZoneRecordsResponse, error) - CreateRecord(context.Context, string, string, dnsimple.ZoneRecordAttributes) (*dnsimple.ZoneRecordResponse, error) - UpdateRecord(context.Context, string, string, int64, dnsimple.ZoneRecordAttributes) (*dnsimple.ZoneRecordResponse, error) -} - -type dnsResolver interface { - LookupIPAddr(context.Context, string) ([]net.IPAddr, error) -} - -func configureDNS(ctx context.Context, cli dnsimpleClient, resolv dnsResolver, accountID string, zone string, ups []upstream, ips []netip.Addr, hostname string) error { - g, ctx := errgroup.WithContext(ctx) - - // Create A and AAAA records for our Tailscale IPs. - for _, tsIP := range ips { - typ := "A" - if tsIP.Is6() { - typ = "AAAA" - } - tsIP := tsIP - g.Go(func() error { - return upsertZoneRecord(ctx, cli, accountID, zone, dnsimple.ZoneRecord{ - Name: hostname, - Type: typ, - TTL: dnsTTL, - Content: tsIP.String(), - }) - }) - } - - // Create ALIAS records for each upstream. - for _, u := range ups { - u := u - g.Go(func() error { - return upsertZoneRecord(ctx, cli, accountID, zone, dnsimple.ZoneRecord{ - Name: u.name, - Type: "ALIAS", - TTL: dnsTTL, - Content: hostname + "." + zone, - }) - }) - } - - // Wait for A and AAAA records to resolve. - for _, ip := range ips { - ip := ip - g.Go(func() error { - host := fqdn(zone, hostname) - if err := waitDNSResolveToIP(ctx, resolv, host, ip.String()); err != nil { - return fmt.Errorf("dns: wait for %s %s: %w", host, ip, err) - } - return nil - }) - } - - return g.Wait() -} - -// dnsimpleSolver is a DNS-01 challenge solver for autocert. -func dnsimpleDNS01Solver(logger *slog.Logger, cli dnsimpleClient, aid, zone string) autocert.DNS01ChallengeSolver { - return func(ctx context.Context, domain string, record string) (func() error, error) { - name := strings.TrimSuffix(domain, "."+zone) - txt := "_acme-challenge." + name - logger.Info("dnsmimple: dns-01 challenge: create TXT record", slog.String("domain", domain), slog.String("TXT", txt)) - err := upsertZoneRecord(ctx, cli, aid, zone, dnsimple.ZoneRecord{ - Type: "TXT", - TTL: 300, - Name: txt, - Content: record, - }) - if err != nil { - return nil, fmt.Errorf("create TXT record in %s: %w", zone, err) - } - - logger.Info("dnsmimple: dns-01 challenge: wait for TXT record", slog.String("domain", domain), slog.String("TXT", txt)) - err = backoff.Retry(func() error { - if err := ctx.Err(); err != nil { - return err - } - logger.Info("dnsimple: lookup TXT", slog.String("TXT", txt+"."+zone)) - recs, err := net.DefaultResolver.LookupTXT(ctx, txt+"."+zone) - if err != nil { - return err - } - for _, r := range recs { - if r == record { - return nil - } - } - return fmt.Errorf("TXT %s on %s does not resolve to expected value", name, zone) - }, backoff.WithContext(backoff.NewExponentialBackOff(), ctx)) - if err != nil { - return nil, fmt.Errorf("dnsimple: dns-01: wait for TXT record: %w", err) - } - return func() error { return nil }, nil // TODO(sr) Implement cleanup. } - } -} - -func waitDNSResolveToIP(ctx context.Context, resolver dnsResolver, name string, ip string) error { - return backoff.Retry(func() error { - if err := ctx.Err(); err != nil { - return err - } - ips, err := resolver.LookupIPAddr(ctx, name) - if err != nil { - return err - } - for _, v := range ips { - if v.String() == ip { - return nil - } - } - return fmt.Errorf("%s does not resolve to %s", name, ip) - }, backoff.WithContext(backoff.NewExponentialBackOff(), ctx)) -} - -// TODO(sr) Write a test. -func upsertZoneRecord(ctx context.Context, cli dnsimpleClient, aid, zone string, rec dnsimple.ZoneRecord) error { - resp, err := cli.ListRecords(ctx, aid, zone, &dnsimple.ZoneRecordListOptions{ - Name: &rec.Name, - Type: &rec.Type, - }) - if err != nil { - return fmt.Errorf("list records: %w", err) - } - if resp != nil && resp.Pagination != nil && resp.Pagination.TotalPages > 1 { - return errors.New("list records: pagination not implemented") - } - if resp == nil || len(resp.Data) == 0 { - _, err := cli.CreateRecord(ctx, aid, zone, dnsimple.ZoneRecordAttributes{ - Type: rec.Type, - Name: &rec.Name, - Content: rec.Content, - TTL: rec.TTL, - }) - if err != nil { - return fmt.Errorf("create record: %w", err) - } - return nil - } - var found dnsimple.ZoneRecord - for _, r := range resp.Data { - if r.Name == rec.Name && r.Type == rec.Type { - found = r - break - } - } - // This should not happen? - if found.ID == 0 { - return errors.New("matching record not found") - } - if _, err := cli.UpdateRecord(ctx, aid, zone, found.ID, dnsimple.ZoneRecordAttributes{ - Name: &rec.Name, - TTL: rec.TTL, - Content: rec.Content, - }); err != nil { - return fmt.Errorf("update record: %w", err) - } - return nil -} diff --git a/dns_test.go b/dns_test.go deleted file mode 100644 index d23c424..0000000 --- a/dns_test.go +++ /dev/null @@ -1,70 +0,0 @@ -package main - -import ( - "context" - "fmt" - "net" - "net/netip" - "sync" - "testing" - - "github.com/dnsimple/dnsimple-go/dnsimple" - "github.com/google/go-cmp/cmp" - "github.com/google/go-cmp/cmp/cmpopts" -) - -type fakeDNSimpleClient struct { - dnsimpleClient - records []dnsimple.ZoneRecord - mu sync.Mutex -} - -func (c *fakeDNSimpleClient) ListRecords(context.Context, string, string, *dnsimple.ZoneRecordListOptions) (*dnsimple.ZoneRecordsResponse, error) { - return nil, nil -} - -func (c *fakeDNSimpleClient) CreateRecord(ctx context.Context, aid string, zone string, rec dnsimple.ZoneRecordAttributes) (*dnsimple.ZoneRecordResponse, error) { - c.mu.Lock() - defer c.mu.Unlock() - c.records = append(c.records, dnsimple.ZoneRecord{ - Name: *rec.Name, - Type: rec.Type, - Content: rec.Content, - TTL: rec.TTL, - }) - return nil, nil -} - -type fakeDNSResolver struct { - ips []net.IPAddr -} - -func (r *fakeDNSResolver) LookupIPAddr(context.Context, string) ([]net.IPAddr, error) { - return r.ips, nil -} - -func TestConfigureDNS(t *testing.T) { - t.Parallel() - - dns := &fakeDNSimpleClient{} - resolv := &fakeDNSResolver{ips: []net.IPAddr{{IP: net.ParseIP("127.0.0.1")}, {IP: net.ParseIP("::1")}}} - ips := []netip.Addr{netip.MustParseAddr("127.0.0.1"), netip.MustParseAddr("::1")} - ups := []upstream{{name: "test1"}, {name: "test2"}} - err := configureDNS(context.TODO(), dns, resolv, "1234", "example.com", ups, ips, "self") - if err != nil { - t.Fatal(err) - } - - want := []dnsimple.ZoneRecord{ - {Type: "A", Name: "self", Content: "127.0.0.1"}, - {Type: "AAAA", Name: "self", Content: "::1"}, - {Type: "ALIAS", Name: "test2", Content: "self.example.com"}, - {Type: "ALIAS", Name: "test1", Content: "self.example.com"}, - } - less := func(a, b dnsimple.ZoneRecord) bool { - return fmt.Sprintf("%+v", a) < fmt.Sprintf("%+v", b) - } - if diff := cmp.Diff(want, dns.records, cmpopts.SortSlices(less), cmpopts.IgnoreFields(dnsimple.ZoneRecord{}, "TTL")); diff != "" { - t.Errorf("dns records mismatch (-want +got):\n%s", diff) - } -} diff --git a/go.mod b/go.mod index 0a54125..3d38da5 100644 --- a/go.mod +++ b/go.mod @@ -4,15 +4,10 @@ go 1.19 require ( github.com/cenkalti/backoff/v4 v4.1.3 - github.com/dnsimple/dnsimple-go v1.0.0 github.com/google/go-cmp v0.5.8 github.com/oklog/run v1.1.0 github.com/prometheus/client_golang v1.14.0 - golang.org/x/crypto v0.0.0-20220427172511-eb4f295cb31f golang.org/x/exp v0.0.0-20221114191408-850992195362 - golang.org/x/net v0.0.0-20221002022538-bcab6841153b - golang.org/x/oauth2 v0.0.0-20220223155221-ee480838109b - golang.org/x/sync v0.0.0-20220601150217-0de741cfad7f tailscale.com v1.32.3 ) @@ -42,7 +37,6 @@ require ( github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da // indirect github.com/golang/protobuf v1.5.2 // indirect github.com/google/btree v1.0.1 // indirect - github.com/google/go-querystring v1.1.0 // indirect github.com/hdevalence/ed25519consensus v0.0.0-20220222234857-c00d1f31bab3 // indirect github.com/insomniacslk/dhcp v0.0.0-20211209223715-7d93572ebe8e // indirect github.com/jmespath/go-jmespath v0.4.0 // indirect @@ -60,6 +54,7 @@ require ( github.com/prometheus/client_model v0.3.0 // indirect github.com/prometheus/common v0.37.0 // indirect github.com/prometheus/procfs v0.8.0 // indirect + github.com/stretchr/testify v1.8.0 // indirect github.com/tailscale/certstore v0.1.1-0.20220316223106-78d6e1c49d8d // indirect github.com/tailscale/golang-x-crypto v0.0.0-20221009170451-62f465106986 // indirect github.com/tailscale/goupnp v1.0.1-0.20210804011211-c64d0f06ea05 // indirect @@ -71,6 +66,9 @@ require ( github.com/x448/float16 v0.8.4 // indirect go4.org/mem v0.0.0-20210711025021-927187094b94 // indirect go4.org/netipx v0.0.0-20220725152314-7e7bdc8411bf // indirect + golang.org/x/crypto v0.0.0-20220427172511-eb4f295cb31f // indirect + golang.org/x/net v0.0.0-20221002022538-bcab6841153b // indirect + golang.org/x/sync v0.0.0-20220601150217-0de741cfad7f // indirect golang.org/x/sys v0.1.0 // indirect golang.org/x/term v0.0.0-20210927222741-03fcf44c2211 // indirect golang.org/x/text v0.3.8-0.20211105212822-18b340fc7af2 // indirect @@ -78,7 +76,6 @@ require ( golang.zx2c4.com/wintun v0.0.0-20211104114900-415007cec224 // indirect golang.zx2c4.com/wireguard v0.0.0-20220904105730-b51010ba13f0 // indirect golang.zx2c4.com/wireguard/windows v0.5.3 // indirect - google.golang.org/appengine v1.6.7 // indirect google.golang.org/protobuf v1.28.1 // indirect gvisor.dev/gvisor v0.0.0-20220817001344-846276b3dbc5 // indirect nhooyr.io/websocket v1.8.7 // indirect diff --git a/go.sum b/go.sum index 9c08618..22b0f3d 100644 --- a/go.sum +++ b/go.sum @@ -93,8 +93,6 @@ github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ3 github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= -github.com/dnsimple/dnsimple-go v1.0.0 h1:x9UalQ0tHR68+sQxJYJmq746LdJou4OLTK+cZLR2Z9I= -github.com/dnsimple/dnsimple-go v1.0.0/go.mod h1:oaAtPP8bIROK3QXUdc8rMlTN7SyvCBAogw2I31WVNnU= github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4= github.com/envoyproxy/go-control-plane v0.9.1-0.20191026205805-5f8ba28d4473/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4= github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98= @@ -188,8 +186,6 @@ github.com/google/go-cmp v0.5.6/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/ github.com/google/go-cmp v0.5.7/go.mod h1:n+brtR0CgQNWTVd5ZUFpTBC8YFBDLK/h/bpaJ8/DtOE= github.com/google/go-cmp v0.5.8 h1:e6P7q2lk1O+qJJb4BtCQXlK8vWEO8V1ZeuEdJNOqZyg= github.com/google/go-cmp v0.5.8/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= -github.com/google/go-querystring v1.1.0 h1:AnCroh3fv4ZBgVIf1Iwtovgjaw/GiKJo8M8yD/fhyJ8= -github.com/google/go-querystring v1.1.0/go.mod h1:Kcdr2DB4koayq7X8pmAG4sNG59So17icRSOU623lUBU= github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= github.com/google/martian v2.1.0+incompatible/go.mod h1:9I4somxYTbIHy5NJKHRl3wXiIaQGbYVAs8BPL6v8lEs= github.com/google/martian/v3 v3.0.0/go.mod h1:y5Zk1BBys9G+gd6Jrk0W3cC1+ELVxBWuIGO+w/tUAp0= @@ -332,11 +328,14 @@ github.com/smartystreets/assertions v0.0.0-20180927180507-b2de0cb4f26d/go.mod h1 github.com/smartystreets/goconvey v1.6.4/go.mod h1:syvi0/a8iFYH4r/RixwvyeAJjdLS9QV7WQ/tjFTllLA= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/objx v0.1.1/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw= github.com/stretchr/testify v1.2.2/go.mod h1:a8OnRcib4nhh0OaRAV+Yts87kKdq0PP7pXfy6kDkUVs= github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4= github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.8.0 h1:pSgiaMZlXftHpm5L7V1+rVB+AZJydKsMxsQBIJw4PKk= +github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU= github.com/tailscale/certstore v0.1.1-0.20220316223106-78d6e1c49d8d h1:K3j02b5j2Iw1xoggN9B2DIEkhWGheqFOeDkdJdBrJI8= github.com/tailscale/certstore v0.1.1-0.20220316223106-78d6e1c49d8d/go.mod h1:2P+hpOwd53e7JMX/L4f3VXkv1G+33ES6IWZSrkIeWNs= github.com/tailscale/golang-x-crypto v0.0.0-20221009170451-62f465106986 h1:jWSwTR9CY13oa2oxhR3FInk1ybqC1NbF9cFeoWrrx+E= @@ -462,7 +461,6 @@ golang.org/x/oauth2 v0.0.0-20190604053449-0f29369cfe45/go.mod h1:gOpvHmFTYa4Iltr golang.org/x/oauth2 v0.0.0-20191202225959-858c2ad4c8b6/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw= golang.org/x/oauth2 v0.0.0-20200107190931-bf48bf16ab8d/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw= golang.org/x/oauth2 v0.0.0-20210514164344-f6687ab2804c/go.mod h1:KelEdhl1UZF7XfJ4dDtk6s++YSgaE7mD/BuKKDLBl4A= -golang.org/x/oauth2 v0.0.0-20220223155221-ee480838109b h1:clP8eMhB30EHdc0bd2Twtq6kgU7yl5ub2cQLSdrv1Dg= golang.org/x/oauth2 v0.0.0-20220223155221-ee480838109b/go.mod h1:DAh4E804XQdzx2j+YRIaUnCqCV2RuMz24cGBJ5QYIrc= golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= @@ -623,8 +621,6 @@ google.golang.org/appengine v1.5.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7 google.golang.org/appengine v1.6.1/go.mod h1:i06prIuMbXzDqacNJfV5OdTW448YApPu5ww/cMBSeb0= google.golang.org/appengine v1.6.5/go.mod h1:8WjMMxjGQR8xUklV/ARdw2HLXBOI7O7uCIDZVag1xfc= google.golang.org/appengine v1.6.6/go.mod h1:8WjMMxjGQR8xUklV/ARdw2HLXBOI7O7uCIDZVag1xfc= -google.golang.org/appengine v1.6.7 h1:FZR1q0exgwxzPzp/aF+VccGrSfxfPpkBqjIIEq3ru6c= -google.golang.org/appengine v1.6.7/go.mod h1:8WjMMxjGQR8xUklV/ARdw2HLXBOI7O7uCIDZVag1xfc= google.golang.org/genproto v0.0.0-20180817151627-c66870c02cf8/go.mod h1:JiN7NxoALGmiZfu7CAH4rXhgtRTLTxftemlI0sWmxmc= google.golang.org/genproto v0.0.0-20190307195333-5fe7a883aa19/go.mod h1:VzzqZJRnGkLBvHegQrXjBqPurQTc5/KpmUdxsrq26oE= google.golang.org/genproto v0.0.0-20190418145605-e7d98fc518a7/go.mod h1:VzzqZJRnGkLBvHegQrXjBqPurQTc5/KpmUdxsrq26oE= @@ -695,6 +691,7 @@ gopkg.in/yaml.v2 v2.4.0 h1:D8xgwECY7CYvx+Y2n4sBz93Jn9JRvxdiyyo8CTfuKaY= gopkg.in/yaml.v2 v2.4.0/go.mod h1:RDklbk79AGWmwhnvt/jBztapEOGDOx6ZbXqjP6csGnQ= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gvisor.dev/gvisor v0.0.0-20220817001344-846276b3dbc5 h1:cv/zaNV0nr1mJzaeo4S5mHIm5va1W0/9J3/5prlsuRM= gvisor.dev/gvisor v0.0.0-20220817001344-846276b3dbc5/go.mod h1:TIvkJD0sxe8pIob3p6T8IzxXunlp6yfgktvTNp+DGNM= honnef.co/go/tools v0.0.0-20190102054323-c2f93a96b099/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= diff --git a/internal/autocert/LICENSE b/internal/autocert/LICENSE deleted file mode 100644 index 12ec31b..0000000 --- a/internal/autocert/LICENSE +++ /dev/null @@ -1,24 +0,0 @@ -Copyright (c) 2009 The Go Authors. All rights reserved. -Redistribution and use in source and binary forms, with or without -modification, are permitted provided that the following conditions are -met: - * Redistributions of source code must retain the above copyright -notice, this list of conditions and the following disclaimer. - * Redistributions in binary form must reproduce the above -copyright notice, this list of conditions and the following disclaimer -in the documentation and/or other materials provided with the -distribution. - * Neither the name of Google Inc. nor the names of its -contributors may be used to endorse or promote products derived from -this software without specific prior written permission. -THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS -"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT -LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR -A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT -OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, -SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT -LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, -DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY -THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT -(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE -OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. diff --git a/internal/autocert/autocert.go b/internal/autocert/autocert.go deleted file mode 100644 index ddd82e2..0000000 --- a/internal/autocert/autocert.go +++ /dev/null @@ -1,1225 +0,0 @@ -// Copyright 2016 The Go Authors. All rights reserved. -// Use of this source code is governed by a BSD-style -// license that can be found in the LICENSE file. - -// Package autocert is a fork of golang.org/x/crypto/acme/autocert with support for -// DNS-01 adapted from https://go-review.googlesource.com/c/crypto/+/381994. -// -// Original package documentation: -// -// Package autocert provides automatic access to certificates from Let's Encrypt -// and any other ACME-based CA. -// -// This package is a work in progress and makes no API stability promises. -package autocert - -import ( - "bytes" - "context" - "crypto" - "crypto/ecdsa" - "crypto/elliptic" - "crypto/rand" - "crypto/rsa" - "crypto/tls" - "crypto/x509" - "crypto/x509/pkix" - "encoding/pem" - "errors" - "fmt" - "io" - mathrand "math/rand" - "net" - "net/http" - "path" - "strings" - "sync" - "time" - - "golang.org/x/crypto/acme" - "golang.org/x/net/idna" -) - -// DefaultACMEDirectory is the default ACME Directory URL used when the Manager's Client is nil. -const DefaultACMEDirectory = "https://acme-v02.api.letsencrypt.org/directory" - -// createCertRetryAfter is how much time to wait before removing a failed state -// entry due to an unsuccessful createCert call. -// This is a variable instead of a const for testing. -// TODO: Consider making it configurable or an exp backoff? -var createCertRetryAfter = time.Minute - -// pseudoRand is safe for concurrent use. -var pseudoRand *lockedMathRand - -var errPreRFC = errors.New("autocert: ACME server doesn't support RFC 8555") - -func init() { - src := mathrand.NewSource(time.Now().UnixNano()) - pseudoRand = &lockedMathRand{rnd: mathrand.New(src)} -} - -// AcceptTOS is a Manager.Prompt function that always returns true to -// indicate acceptance of the CA's Terms of Service during account -// registration. -func AcceptTOS(tosURL string) bool { return true } - -// HostPolicy specifies which host names the Manager is allowed to respond to. -// It returns a non-nil error if the host should be rejected. -// The returned error is accessible via tls.Conn.Handshake and its callers. -// See Manager's HostPolicy field and GetCertificate method docs for more details. -type HostPolicy func(ctx context.Context, host string) error - -// HostWhitelist returns a policy where only the specified host names are allowed. -// Only exact matches are currently supported. Subdomains, regexp or wildcard -// will not match. -// -// Note that all hosts will be converted to Punycode via idna.Lookup.ToASCII so that -// Manager.GetCertificate can handle the Unicode IDN and mixedcase hosts correctly. -// Invalid hosts will be silently ignored. -func HostWhitelist(hosts ...string) HostPolicy { - whitelist := make(map[string]bool, len(hosts)) - for _, h := range hosts { - if h, err := idna.Lookup.ToASCII(h); err == nil { - whitelist[h] = true - } - } - return func(_ context.Context, host string) error { - if !whitelist[host] { - return fmt.Errorf("acme/autocert: host %q not configured in HostWhitelist", host) - } - return nil - } -} - -// defaultHostPolicy is used when Manager.HostPolicy is not set. -func defaultHostPolicy(context.Context, string) error { - return nil -} - -type DNS01ChallengeSolver func(ctx context.Context, domain string, record string) (cleaner func() error, err error) - -// Manager is a stateful certificate manager built on top of acme.Client. -// It obtains and refreshes certificates automatically using "tls-alpn-01", -// "http-01", or "dns-01" challenge types, as well as providing them to a -// TLS server via tls.Config. -// -// You must specify a cache implementation, such as DirCache, -// to reuse obtained certificates across program restarts. -// Otherwise your server is very likely to exceed the certificate -// issuer's request rate limits. -type Manager struct { - // Prompt specifies a callback function to conditionally accept a CA's Terms of Service (TOS). - // The registration may require the caller to agree to the CA's TOS. - // If so, Manager calls Prompt with a TOS URL provided by the CA. Prompt should report - // whether the caller agrees to the terms. - // - // To always accept the terms, the callers can use AcceptTOS. - Prompt func(tosURL string) bool - - // Cache optionally stores and retrieves previously-obtained certificates - // and other state. If nil, certs will only be cached for the lifetime of - // the Manager. Multiple Managers can share the same Cache. - // - // Using a persistent Cache, such as DirCache, is strongly recommended. - Cache Cache - - // HostPolicy controls which domains the Manager will attempt - // to retrieve new certificates for. It does not affect cached certs. - // - // If non-nil, HostPolicy is called before requesting a new cert. - // If nil, all hosts are currently allowed. This is not recommended, - // as it opens a potential attack where clients connect to a server - // by IP address and pretend to be asking for an incorrect host name. - // Manager will attempt to obtain a certificate for that host, incorrectly, - // eventually reaching the CA's rate limit for certificate requests - // and making it impossible to obtain actual certificates. - // - // See GetCertificate for more details. - HostPolicy HostPolicy - - // RenewBefore optionally specifies how early certificates should - // be renewed before they expire. - // - // If zero, they're renewed 30 days before expiration. - RenewBefore time.Duration - - // Client is used to perform low-level operations, such as account registration - // and requesting new certificates. - // - // If Client is nil, a zero-value acme.Client is used with DefaultACMEDirectory - // as the directory endpoint. - // If the Client.Key is nil, a new ECDSA P-256 key is generated and, - // if Cache is not nil, stored in cache. - // - // Mutating the field after the first call of GetCertificate method will have no effect. - Client *acme.Client - - // Email optionally specifies a contact email address. - // This is used by CAs, such as Let's Encrypt, to notify about problems - // with issued certificates. - // - // If the Client's account key is already registered, Email is not used. - Email string - - // ForceRSA used to make the Manager generate RSA certificates. It is now ignored. - // - // Deprecated: the Manager will request the correct type of certificate based - // on what each client supports. - ForceRSA bool - - // ExtraExtensions are used when generating a new CSR (Certificate Request), - // thus allowing customization of the resulting certificate. - // For instance, TLS Feature Extension (RFC 7633) can be used - // to prevent an OCSP downgrade attack. - // - // The field value is passed to crypto/x509.CreateCertificateRequest - // in the template's ExtraExtensions field as is. - ExtraExtensions []pkix.Extension - - // Sesolver is used to respond to dns-01 challenges returned from the CA. - // If this field is nil then DNS challenges will not be requested from the - // CA. - DNS01 DNS01ChallengeSolver - - // ExternalAccountBinding optionally represents an arbitrary binding to an - // account of the CA to which the ACME server is tied. - // See RFC 8555, Section 7.3.4 for more details. - ExternalAccountBinding *acme.ExternalAccountBinding - - clientMu sync.Mutex - client *acme.Client // initialized by acmeClient method - - stateMu sync.Mutex - state map[certKey]*certState - - // renewal tracks the set of domains currently running renewal timers. - renewalMu sync.Mutex - renewal map[certKey]*domainRenewal - - // challengeMu guards tryHTTP01, certTokens and httpTokens. - challengeMu sync.RWMutex - // tryHTTP01 indicates whether the Manager should try "http-01" challenge type - // during the authorization flow. - tryHTTP01 bool - // httpTokens contains response body values for http-01 challenges - // and is keyed by the URL path at which a challenge response is expected - // to be provisioned. - // The entries are stored for the duration of the authorization flow. - httpTokens map[string][]byte - // certTokens contains temporary certificates for tls-alpn-01 challenges - // and is keyed by the domain name which matches the ClientHello server name. - // The entries are stored for the duration of the authorization flow. - certTokens map[string]*tls.Certificate - - // nowFunc, if not nil, returns the current time. This may be set for - // testing purposes. - nowFunc func() time.Time -} - -// certKey is the key by which certificates are tracked in state, renewal and cache. -type certKey struct { - domain string // without trailing dot - isRSA bool // RSA cert for legacy clients (as opposed to default ECDSA) - isToken bool // tls-based challenge token cert; key type is undefined regardless of isRSA -} - -func (c certKey) String() string { - if c.isToken { - return c.domain + "+token" - } - if c.isRSA { - return c.domain + "+rsa" - } - return c.domain -} - -// TLSConfig creates a new TLS config suitable for net/http.Server servers, -// supporting HTTP/2 and the tls-alpn-01 ACME challenge type. -func (m *Manager) TLSConfig() *tls.Config { - return &tls.Config{ - GetCertificate: m.GetCertificate, - NextProtos: []string{ - "h2", "http/1.1", // enable HTTP/2 - acme.ALPNProto, // enable tls-alpn ACME challenges - }, - } -} - -// GetCertificate implements the tls.Config.GetCertificate hook. -// It provides a TLS certificate for hello.ServerName host, including answering -// tls-alpn-01 challenges. -// All other fields of hello are ignored. -// -// If m.HostPolicy is non-nil, GetCertificate calls the policy before requesting -// a new cert. A non-nil error returned from m.HostPolicy halts TLS negotiation. -// The error is propagated back to the caller of GetCertificate and is user-visible. -// This does not affect cached certs. See HostPolicy field description for more details. -// -// If GetCertificate is used directly, instead of via Manager.TLSConfig, package users will -// also have to add acme.ALPNProto to NextProtos for tls-alpn-01, or use HTTPHandler for http-01. -// If DNSManager is specified, no additional configuration is required and GetCertificate -// can be used directly for dns-01 challenges. -func (m *Manager) GetCertificate(hello *tls.ClientHelloInfo) (*tls.Certificate, error) { - if m.Prompt == nil { - return nil, errors.New("acme/autocert: Manager.Prompt not set") - } - - name := hello.ServerName - if name == "" { - return nil, errors.New("acme/autocert: missing server name") - } - if !strings.Contains(strings.Trim(name, "."), ".") { - return nil, errors.New("acme/autocert: server name component count invalid") - } - - // Note that this conversion is necessary because some server names in the handshakes - // started by some clients (such as cURL) are not converted to Punycode, which will - // prevent us from obtaining certificates for them. In addition, we should also treat - // example.com and EXAMPLE.COM as equivalent and return the same certificate for them. - // Fortunately, this conversion also helped us deal with this kind of mixedcase problems. - // - // Due to the "σςΣ" problem (see https://unicode.org/faq/idn.html#22), we can't use - // idna.Punycode.ToASCII (or just idna.ToASCII) here. - name, err := idna.Lookup.ToASCII(name) - if err != nil { - return nil, errors.New("acme/autocert: server name contains invalid character") - } - - // In the worst-case scenario, the timeout needs to account for caching, host policy, - // domain ownership verification and certificate issuance. - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) - defer cancel() - - // Check whether this is a token cert requested for TLS-ALPN challenge. - if wantsTokenCert(hello) { - m.challengeMu.RLock() - defer m.challengeMu.RUnlock() - if cert := m.certTokens[name]; cert != nil { - return cert, nil - } - if cert, err := m.cacheGet(ctx, certKey{domain: name, isToken: true}); err == nil { - return cert, nil - } - // TODO: cache error results? - return nil, fmt.Errorf("acme/autocert: no token cert for %q", name) - } - - // regular domain - ck := certKey{ - domain: strings.TrimSuffix(name, "."), // golang.org/issue/18114 - isRSA: !supportsECDSA(hello), - } - cert, err := m.cert(ctx, ck) - if err == nil { - return cert, nil - } - if err != ErrCacheMiss { - return nil, err - } - - // first-time - if err := m.hostPolicy()(ctx, name); err != nil { - return nil, err - } - cert, err = m.createCert(ctx, ck) - if err != nil { - return nil, err - } - m.cachePut(ctx, ck, cert) - return cert, nil -} - -// wantsTokenCert reports whether a TLS request with SNI is made by a CA server -// for a challenge verification. -func wantsTokenCert(hello *tls.ClientHelloInfo) bool { - // tls-alpn-01 - if len(hello.SupportedProtos) == 1 && hello.SupportedProtos[0] == acme.ALPNProto { - return true - } - return false -} - -func supportsECDSA(hello *tls.ClientHelloInfo) bool { - // The "signature_algorithms" extension, if present, limits the key exchange - // algorithms allowed by the cipher suites. See RFC 5246, section 7.4.1.4.1. - if hello.SignatureSchemes != nil { - ecdsaOK := false - schemeLoop: - for _, scheme := range hello.SignatureSchemes { - const tlsECDSAWithSHA1 tls.SignatureScheme = 0x0203 // constant added in Go 1.10 - switch scheme { - case tlsECDSAWithSHA1, tls.ECDSAWithP256AndSHA256, - tls.ECDSAWithP384AndSHA384, tls.ECDSAWithP521AndSHA512: - ecdsaOK = true - break schemeLoop - } - } - if !ecdsaOK { - return false - } - } - if hello.SupportedCurves != nil { - ecdsaOK := false - for _, curve := range hello.SupportedCurves { - if curve == tls.CurveP256 { - ecdsaOK = true - break - } - } - if !ecdsaOK { - return false - } - } - for _, suite := range hello.CipherSuites { - switch suite { - case tls.TLS_ECDHE_ECDSA_WITH_RC4_128_SHA, - tls.TLS_ECDHE_ECDSA_WITH_AES_128_CBC_SHA, - tls.TLS_ECDHE_ECDSA_WITH_AES_256_CBC_SHA, - tls.TLS_ECDHE_ECDSA_WITH_AES_128_CBC_SHA256, - tls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256, - tls.TLS_ECDHE_ECDSA_WITH_AES_256_GCM_SHA384, - tls.TLS_ECDHE_ECDSA_WITH_CHACHA20_POLY1305: - return true - } - } - return false -} - -// HTTPHandler configures the Manager to provision ACME "http-01" challenge responses. -// It returns an http.Handler that responds to the challenges and must be -// running on port 80. If it receives a request that is not an ACME challenge, -// it delegates the request to the optional fallback handler. -// -// If fallback is nil, the returned handler redirects all GET and HEAD requests -// to the default TLS port 443 with 302 Found status code, preserving the original -// request path and query. It responds with 400 Bad Request to all other HTTP methods. -// The fallback is not protected by the optional HostPolicy. -// -// Because the fallback handler is run with unencrypted port 80 requests, -// the fallback should not serve TLS-only requests. -// -// If HTTPHandler is never called, the Manager will use the "tls-alpn-01" -// challenge for domain verification and/or "dns-01" if DNSManager is specified. -func (m *Manager) HTTPHandler(fallback http.Handler) http.Handler { - m.challengeMu.Lock() - defer m.challengeMu.Unlock() - m.tryHTTP01 = true - - if fallback == nil { - fallback = http.HandlerFunc(handleHTTPRedirect) - } - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if !strings.HasPrefix(r.URL.Path, "/.well-known/acme-challenge/") { - fallback.ServeHTTP(w, r) - return - } - // A reasonable context timeout for cache and host policy only, - // because we don't wait for a new certificate issuance here. - ctx, cancel := context.WithTimeout(r.Context(), time.Minute) - defer cancel() - if err := m.hostPolicy()(ctx, r.Host); err != nil { - http.Error(w, err.Error(), http.StatusForbidden) - return - } - data, err := m.httpToken(ctx, r.URL.Path) - if err != nil { - http.Error(w, err.Error(), http.StatusNotFound) - return - } - w.Write(data) - }) -} - -func handleHTTPRedirect(w http.ResponseWriter, r *http.Request) { - if r.Method != "GET" && r.Method != "HEAD" { - http.Error(w, "Use HTTPS", http.StatusBadRequest) - return - } - target := "https://" + stripPort(r.Host) + r.URL.RequestURI() - http.Redirect(w, r, target, http.StatusFound) -} - -func stripPort(hostport string) string { - host, _, err := net.SplitHostPort(hostport) - if err != nil { - return hostport - } - return net.JoinHostPort(host, "443") -} - -// cert returns an existing certificate either from m.state or cache. -// If a certificate is found in cache but not in m.state, the latter will be filled -// with the cached value. -func (m *Manager) cert(ctx context.Context, ck certKey) (*tls.Certificate, error) { - m.stateMu.Lock() - if s, ok := m.state[ck]; ok { - m.stateMu.Unlock() - s.RLock() - defer s.RUnlock() - return s.tlscert() - } - defer m.stateMu.Unlock() - cert, err := m.cacheGet(ctx, ck) - if err != nil { - return nil, err - } - signer, ok := cert.PrivateKey.(crypto.Signer) - if !ok { - return nil, errors.New("acme/autocert: private key cannot sign") - } - if m.state == nil { - m.state = make(map[certKey]*certState) - } - s := &certState{ - key: signer, - cert: cert.Certificate, - leaf: cert.Leaf, - } - m.state[ck] = s - m.startRenew(ck, s.key, s.leaf.NotAfter) - return cert, nil -} - -// cacheGet always returns a valid certificate, or an error otherwise. -// If a cached certificate exists but is not valid, ErrCacheMiss is returned. -func (m *Manager) cacheGet(ctx context.Context, ck certKey) (*tls.Certificate, error) { - if m.Cache == nil { - return nil, ErrCacheMiss - } - data, err := m.Cache.Get(ctx, ck.String()) - if err != nil { - return nil, err - } - - // private - priv, pub := pem.Decode(data) - if priv == nil || !strings.Contains(priv.Type, "PRIVATE") { - return nil, ErrCacheMiss - } - privKey, err := parsePrivateKey(priv.Bytes) - if err != nil { - return nil, err - } - - // public - var pubDER [][]byte - for len(pub) > 0 { - var b *pem.Block - b, pub = pem.Decode(pub) - if b == nil { - break - } - pubDER = append(pubDER, b.Bytes) - } - if len(pub) > 0 { - // Leftover content not consumed by pem.Decode. Corrupt. Ignore. - return nil, ErrCacheMiss - } - - // verify and create TLS cert - leaf, err := validCert(ck, pubDER, privKey, m.now()) - if err != nil { - return nil, ErrCacheMiss - } - tlscert := &tls.Certificate{ - Certificate: pubDER, - PrivateKey: privKey, - Leaf: leaf, - } - return tlscert, nil -} - -func (m *Manager) cachePut(ctx context.Context, ck certKey, tlscert *tls.Certificate) error { - if m.Cache == nil { - return nil - } - - // contains PEM-encoded data - var buf bytes.Buffer - - // private - switch key := tlscert.PrivateKey.(type) { - case *ecdsa.PrivateKey: - if err := encodeECDSAKey(&buf, key); err != nil { - return err - } - case *rsa.PrivateKey: - b := x509.MarshalPKCS1PrivateKey(key) - pb := &pem.Block{Type: "RSA PRIVATE KEY", Bytes: b} - if err := pem.Encode(&buf, pb); err != nil { - return err - } - default: - return errors.New("acme/autocert: unknown private key type") - } - - // public - for _, b := range tlscert.Certificate { - pb := &pem.Block{Type: "CERTIFICATE", Bytes: b} - if err := pem.Encode(&buf, pb); err != nil { - return err - } - } - - return m.Cache.Put(ctx, ck.String(), buf.Bytes()) -} - -func encodeECDSAKey(w io.Writer, key *ecdsa.PrivateKey) error { - b, err := x509.MarshalECPrivateKey(key) - if err != nil { - return err - } - pb := &pem.Block{Type: "EC PRIVATE KEY", Bytes: b} - return pem.Encode(w, pb) -} - -// createCert starts the domain ownership verification and returns a certificate -// for that domain upon success. -// -// If the domain is already being verified, it waits for the existing verification to complete. -// Either way, createCert blocks for the duration of the whole process. -func (m *Manager) createCert(ctx context.Context, ck certKey) (*tls.Certificate, error) { - // TODO: maybe rewrite this whole piece using sync.Once - state, err := m.certState(ck) - if err != nil { - return nil, err - } - // state may exist if another goroutine is already working on it - // in which case just wait for it to finish - if !state.locked { - state.RLock() - defer state.RUnlock() - return state.tlscert() - } - - // We are the first; state is locked. - // Unblock the readers when domain ownership is verified - // and we got the cert or the process failed. - defer state.Unlock() - state.locked = false - - der, leaf, err := m.authorizedCert(ctx, state.key, ck) - if err != nil { - // Remove the failed state after some time, - // making the manager call createCert again on the following TLS hello. - didRemove := testDidRemoveState // The lifetime of this timer is untracked, so copy mutable local state to avoid races. - time.AfterFunc(createCertRetryAfter, func() { - defer didRemove(ck) - m.stateMu.Lock() - defer m.stateMu.Unlock() - // Verify the state hasn't changed and it's still invalid - // before deleting. - s, ok := m.state[ck] - if !ok { - return - } - if _, err := validCert(ck, s.cert, s.key, m.now()); err == nil { - return - } - delete(m.state, ck) - }) - return nil, err - } - state.cert = der - state.leaf = leaf - m.startRenew(ck, state.key, state.leaf.NotAfter) - return state.tlscert() -} - -// certState returns a new or existing certState. -// If a new certState is returned, state.exist is false and the state is locked. -// The returned error is non-nil only in the case where a new state could not be created. -func (m *Manager) certState(ck certKey) (*certState, error) { - m.stateMu.Lock() - defer m.stateMu.Unlock() - if m.state == nil { - m.state = make(map[certKey]*certState) - } - // existing state - if state, ok := m.state[ck]; ok { - return state, nil - } - - // new locked state - var ( - err error - key crypto.Signer - ) - if ck.isRSA { - key, err = rsa.GenerateKey(rand.Reader, 2048) - } else { - key, err = ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - } - if err != nil { - return nil, err - } - - state := &certState{ - key: key, - locked: true, - } - state.Lock() // will be unlocked by m.certState caller - m.state[ck] = state - return state, nil -} - -// authorizedCert starts the domain ownership verification process and requests a new cert upon success. -// The key argument is the certificate private key. -func (m *Manager) authorizedCert(ctx context.Context, key crypto.Signer, ck certKey) (der [][]byte, leaf *x509.Certificate, err error) { - csr, err := certRequest(key, ck.domain, m.ExtraExtensions) - if err != nil { - return nil, nil, err - } - - client, err := m.acmeClient(ctx) - if err != nil { - return nil, nil, err - } - dir, err := client.Discover(ctx) - if err != nil { - return nil, nil, err - } - if dir.OrderURL == "" { - return nil, nil, errPreRFC - } - - o, err := m.verifyRFC(ctx, client, ck.domain) - if err != nil { - return nil, nil, err - } - chain, _, err := client.CreateOrderCert(ctx, o.FinalizeURL, csr, true) - if err != nil { - return nil, nil, err - } - - leaf, err = validCert(ck, chain, key, m.now()) - if err != nil { - return nil, nil, err - } - return chain, leaf, nil -} - -// verifyRFC runs the identifier (domain) order-based authorization flow for RFC compliant CAs -// using each applicable ACME challenge type. -func (m *Manager) verifyRFC(ctx context.Context, client *acme.Client, domain string) (*acme.Order, error) { - // Try each supported challenge type starting with a new order each time. - // The nextTyp index of the next challenge type to try is shared across - // all order authorizations: if we've tried a challenge type once and it didn't work, - // it will most likely not work on another order's authorization either. - challengeTypes := m.supportedChallengeTypes() - nextTyp := 0 // challengeTypes index -AuthorizeOrderLoop: - for { - o, err := client.AuthorizeOrder(ctx, acme.DomainIDs(domain)) - if err != nil { - return nil, err - } - // Remove all hanging authorizations to reduce rate limit quotas - // after we're done. - defer func(urls []string) { - go m.deactivatePendingAuthz(urls) - }(o.AuthzURLs) - - // Check if there's actually anything we need to do. - switch o.Status { - case acme.StatusReady: - // Already authorized. - return o, nil - case acme.StatusPending: - // Continue normal Order-based flow. - default: - return nil, fmt.Errorf("acme/autocert: invalid new order status %q; order URL: %q", o.Status, o.URI) - } - - // Satisfy all pending authorizations. - for _, zurl := range o.AuthzURLs { - z, err := client.GetAuthorization(ctx, zurl) - if err != nil { - return nil, err - } - if z.Status != acme.StatusPending { - // We are interested only in pending authorizations. - continue - } - // Pick the next preferred challenge. - var chal *acme.Challenge - for chal == nil && nextTyp < len(challengeTypes) { - chal = pickChallenge(challengeTypes[nextTyp], z.Challenges) - nextTyp++ - } - if chal == nil { - return nil, fmt.Errorf("acme/autocert: unable to satisfy %q for domain %q: no viable challenge type found", z.URI, domain) - } - // Respond to the challenge and wait for validation result. - cleanup, err := m.fulfill(ctx, client, chal, domain) - if err != nil { - continue AuthorizeOrderLoop - } - defer cleanup() - if _, err := client.Accept(ctx, chal); err != nil { - continue AuthorizeOrderLoop - } - if _, err := client.WaitAuthorization(ctx, z.URI); err != nil { - continue AuthorizeOrderLoop - } - } - - // All authorizations are satisfied. - // Wait for the CA to update the order status. - o, err = client.WaitOrder(ctx, o.URI) - if err != nil { - continue AuthorizeOrderLoop - } - return o, nil - } -} - -func pickChallenge(typ string, chal []*acme.Challenge) *acme.Challenge { - for _, c := range chal { - if c.Type == typ { - return c - } - } - return nil -} - -func (m *Manager) supportedChallengeTypes() []string { - m.challengeMu.RLock() - defer m.challengeMu.RUnlock() - typ := []string{"tls-alpn-01"} - if m.tryHTTP01 { - typ = append(typ, "http-01") - } - if m.DNS01 != nil { - typ = append(typ, "dns-01") - } - return typ -} - -// deactivatePendingAuthz relinquishes all authorizations identified by the elements -// of the provided uri slice which are in "pending" state. -// It ignores revocation errors. -// -// deactivatePendingAuthz takes no context argument and instead runs with its own -// "detached" context because deactivations are done in a goroutine separate from -// that of the main issuance or renewal flow. -func (m *Manager) deactivatePendingAuthz(uri []string) { - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) - defer cancel() - client, err := m.acmeClient(ctx) - if err != nil { - return - } - for _, u := range uri { - z, err := client.GetAuthorization(ctx, u) - if err == nil && z.Status == acme.StatusPending { - client.RevokeAuthorization(ctx, u) - } - } -} - -// fulfill provisions a response to the challenge chal. -// The cleanup is non-nil only if provisioning succeeded. -func (m *Manager) fulfill(ctx context.Context, client *acme.Client, chal *acme.Challenge, domain string) (cleanup func(), err error) { - switch chal.Type { - case "tls-alpn-01": - cert, err := client.TLSALPN01ChallengeCert(chal.Token, domain) - if err != nil { - return nil, err - } - m.putCertToken(ctx, domain, &cert) - return func() { go m.deleteCertToken(domain) }, nil - case "http-01": - resp, err := client.HTTP01ChallengeResponse(chal.Token) - if err != nil { - return nil, err - } - p := client.HTTP01ChallengePath(chal.Token) - m.putHTTPToken(ctx, p, resp) - return func() { go m.deleteHTTPToken(p) }, nil - case "dns-01": - rec, err := client.DNS01ChallengeRecord(chal.Token) - if err != nil { - return nil, err - } - cleanup, err := m.DNS01(ctx, domain, rec) - if err != nil { - return nil, err - } - return func() { _ = cleanup() }, err - } - return nil, fmt.Errorf("acme/autocert: unknown challenge type %q", chal.Type) -} - -// putCertToken stores the token certificate with the specified name -// in both m.certTokens map and m.Cache. -func (m *Manager) putCertToken(ctx context.Context, name string, cert *tls.Certificate) { - m.challengeMu.Lock() - defer m.challengeMu.Unlock() - if m.certTokens == nil { - m.certTokens = make(map[string]*tls.Certificate) - } - m.certTokens[name] = cert - m.cachePut(ctx, certKey{domain: name, isToken: true}, cert) -} - -// deleteCertToken removes the token certificate with the specified name -// from both m.certTokens map and m.Cache. -func (m *Manager) deleteCertToken(name string) { - m.challengeMu.Lock() - defer m.challengeMu.Unlock() - delete(m.certTokens, name) - if m.Cache != nil { - ck := certKey{domain: name, isToken: true} - m.Cache.Delete(context.Background(), ck.String()) - } -} - -// httpToken retrieves an existing http-01 token value from an in-memory map -// or the optional cache. -func (m *Manager) httpToken(ctx context.Context, tokenPath string) ([]byte, error) { - m.challengeMu.RLock() - defer m.challengeMu.RUnlock() - if v, ok := m.httpTokens[tokenPath]; ok { - return v, nil - } - if m.Cache == nil { - return nil, fmt.Errorf("acme/autocert: no token at %q", tokenPath) - } - return m.Cache.Get(ctx, httpTokenCacheKey(tokenPath)) -} - -// putHTTPToken stores an http-01 token value using tokenPath as key -// in both in-memory map and the optional Cache. -// -// It ignores any error returned from Cache.Put. -func (m *Manager) putHTTPToken(ctx context.Context, tokenPath, val string) { - m.challengeMu.Lock() - defer m.challengeMu.Unlock() - if m.httpTokens == nil { - m.httpTokens = make(map[string][]byte) - } - b := []byte(val) - m.httpTokens[tokenPath] = b - if m.Cache != nil { - m.Cache.Put(ctx, httpTokenCacheKey(tokenPath), b) - } -} - -// deleteHTTPToken removes an http-01 token value from both in-memory map -// and the optional Cache, ignoring any error returned from the latter. -// -// If m.Cache is non-nil, it blocks until Cache.Delete returns without a timeout. -func (m *Manager) deleteHTTPToken(tokenPath string) { - m.challengeMu.Lock() - defer m.challengeMu.Unlock() - delete(m.httpTokens, tokenPath) - if m.Cache != nil { - m.Cache.Delete(context.Background(), httpTokenCacheKey(tokenPath)) - } -} - -// httpTokenCacheKey returns a key at which an http-01 token value may be stored -// in the Manager's optional Cache. -func httpTokenCacheKey(tokenPath string) string { - return path.Base(tokenPath) + "+http-01" -} - -// startRenew starts a cert renewal timer loop, one per domain. -// -// The loop is scheduled in two cases: -// - a cert was fetched from cache for the first time (wasn't in m.state) -// - a new cert was created by m.createCert -// -// The key argument is a certificate private key. -// The exp argument is the cert expiration time (NotAfter). -func (m *Manager) startRenew(ck certKey, key crypto.Signer, exp time.Time) { - m.renewalMu.Lock() - defer m.renewalMu.Unlock() - if m.renewal[ck] != nil { - // another goroutine is already on it - return - } - if m.renewal == nil { - m.renewal = make(map[certKey]*domainRenewal) - } - dr := &domainRenewal{m: m, ck: ck, key: key} - m.renewal[ck] = dr - dr.start(exp) -} - -// stopRenew stops all currently running cert renewal timers. -// The timers are not restarted during the lifetime of the Manager. -func (m *Manager) stopRenew() { - m.renewalMu.Lock() - defer m.renewalMu.Unlock() - for name, dr := range m.renewal { - delete(m.renewal, name) - dr.stop() - } -} - -func (m *Manager) accountKey(ctx context.Context) (crypto.Signer, error) { - const keyName = "acme_account+key" - - // Previous versions of autocert stored the value under a different key. - const legacyKeyName = "acme_account.key" - - genKey := func() (*ecdsa.PrivateKey, error) { - return ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - } - - if m.Cache == nil { - return genKey() - } - - data, err := m.Cache.Get(ctx, keyName) - if err == ErrCacheMiss { - data, err = m.Cache.Get(ctx, legacyKeyName) - } - if err == ErrCacheMiss { - key, err := genKey() - if err != nil { - return nil, err - } - var buf bytes.Buffer - if err := encodeECDSAKey(&buf, key); err != nil { - return nil, err - } - if err := m.Cache.Put(ctx, keyName, buf.Bytes()); err != nil { - return nil, err - } - return key, nil - } - if err != nil { - return nil, err - } - - priv, _ := pem.Decode(data) - if priv == nil || !strings.Contains(priv.Type, "PRIVATE") { - return nil, errors.New("acme/autocert: invalid account key found in cache") - } - return parsePrivateKey(priv.Bytes) -} - -func (m *Manager) acmeClient(ctx context.Context) (*acme.Client, error) { - m.clientMu.Lock() - defer m.clientMu.Unlock() - if m.client != nil { - return m.client, nil - } - - client := m.Client - if client == nil { - client = &acme.Client{DirectoryURL: DefaultACMEDirectory} - } - if client.Key == nil { - var err error - client.Key, err = m.accountKey(ctx) - if err != nil { - return nil, err - } - } - if client.UserAgent == "" { - client.UserAgent = "autocert" - } - var contact []string - if m.Email != "" { - contact = []string{"mailto:" + m.Email} - } - a := &acme.Account{Contact: contact, ExternalAccountBinding: m.ExternalAccountBinding} - _, err := client.Register(ctx, a, m.Prompt) - if err == nil || isAccountAlreadyExist(err) { - m.client = client - err = nil - } - return m.client, err -} - -// isAccountAlreadyExist reports whether the err, as returned from acme.Client.Register, -// indicates the account has already been registered. -func isAccountAlreadyExist(err error) bool { - if err == acme.ErrAccountAlreadyExists { - return true - } - ae, ok := err.(*acme.Error) - return ok && ae.StatusCode == http.StatusConflict -} - -func (m *Manager) hostPolicy() HostPolicy { - if m.HostPolicy != nil { - return m.HostPolicy - } - return defaultHostPolicy -} - -func (m *Manager) renewBefore() time.Duration { - if m.RenewBefore > renewJitter { - return m.RenewBefore - } - return 720 * time.Hour // 30 days -} - -func (m *Manager) now() time.Time { - if m.nowFunc != nil { - return m.nowFunc() - } - return time.Now() -} - -// certState is ready when its mutex is unlocked for reading. -type certState struct { - sync.RWMutex - locked bool // locked for read/write - key crypto.Signer // private key for cert - cert [][]byte // DER encoding - leaf *x509.Certificate // parsed cert[0]; always non-nil if cert != nil -} - -// tlscert creates a tls.Certificate from s.key and s.cert. -// Callers should wrap it in s.RLock() and s.RUnlock(). -func (s *certState) tlscert() (*tls.Certificate, error) { - if s.key == nil { - return nil, errors.New("acme/autocert: missing signer") - } - if len(s.cert) == 0 { - return nil, errors.New("acme/autocert: missing certificate") - } - return &tls.Certificate{ - PrivateKey: s.key, - Certificate: s.cert, - Leaf: s.leaf, - }, nil -} - -// certRequest generates a CSR for the given common name. -func certRequest(key crypto.Signer, name string, ext []pkix.Extension) ([]byte, error) { - req := &x509.CertificateRequest{ - Subject: pkix.Name{CommonName: name}, - DNSNames: []string{name}, - ExtraExtensions: ext, - } - return x509.CreateCertificateRequest(rand.Reader, req, key) -} - -// Attempt to parse the given private key DER block. OpenSSL 0.9.8 generates -// PKCS#1 private keys by default, while OpenSSL 1.0.0 generates PKCS#8 keys. -// OpenSSL ecparam generates SEC1 EC private keys for ECDSA. We try all three. -// -// Inspired by parsePrivateKey in crypto/tls/tls.go. -func parsePrivateKey(der []byte) (crypto.Signer, error) { - if key, err := x509.ParsePKCS1PrivateKey(der); err == nil { - return key, nil - } - if key, err := x509.ParsePKCS8PrivateKey(der); err == nil { - switch key := key.(type) { - case *rsa.PrivateKey: - return key, nil - case *ecdsa.PrivateKey: - return key, nil - default: - return nil, errors.New("acme/autocert: unknown private key type in PKCS#8 wrapping") - } - } - if key, err := x509.ParseECPrivateKey(der); err == nil { - return key, nil - } - - return nil, errors.New("acme/autocert: failed to parse private key") -} - -// validCert parses a cert chain provided as der argument and verifies the leaf and der[0] -// correspond to the private key, the domain and key type match, and expiration dates -// are valid. It doesn't do any revocation checking. -// -// The returned value is the verified leaf cert. -func validCert(ck certKey, der [][]byte, key crypto.Signer, now time.Time) (leaf *x509.Certificate, err error) { - // parse public part(s) - var n int - for _, b := range der { - n += len(b) - } - pub := make([]byte, n) - n = 0 - for _, b := range der { - n += copy(pub[n:], b) - } - x509Cert, err := x509.ParseCertificates(pub) - if err != nil || len(x509Cert) == 0 { - return nil, errors.New("acme/autocert: no public key found") - } - // verify the leaf is not expired and matches the domain name - leaf = x509Cert[0] - if now.Before(leaf.NotBefore) { - return nil, errors.New("acme/autocert: certificate is not valid yet") - } - if now.After(leaf.NotAfter) { - return nil, errors.New("acme/autocert: expired certificate") - } - if err := leaf.VerifyHostname(ck.domain); err != nil { - return nil, err - } - // renew certificates revoked by Let's Encrypt in January 2022 - if isRevokedLetsEncrypt(leaf) { - return nil, errors.New("acme/autocert: certificate was probably revoked by Let's Encrypt") - } - // ensure the leaf corresponds to the private key and matches the certKey type - switch pub := leaf.PublicKey.(type) { - case *rsa.PublicKey: - prv, ok := key.(*rsa.PrivateKey) - if !ok { - return nil, errors.New("acme/autocert: private key type does not match public key type") - } - if pub.N.Cmp(prv.N) != 0 { - return nil, errors.New("acme/autocert: private key does not match public key") - } - if !ck.isRSA && !ck.isToken { - return nil, errors.New("acme/autocert: key type does not match expected value") - } - case *ecdsa.PublicKey: - prv, ok := key.(*ecdsa.PrivateKey) - if !ok { - return nil, errors.New("acme/autocert: private key type does not match public key type") - } - if pub.X.Cmp(prv.X) != 0 || pub.Y.Cmp(prv.Y) != 0 { - return nil, errors.New("acme/autocert: private key does not match public key") - } - if ck.isRSA && !ck.isToken { - return nil, errors.New("acme/autocert: key type does not match expected value") - } - default: - return nil, errors.New("acme/autocert: unknown public key algorithm") - } - return leaf, nil -} - -// https://community.letsencrypt.org/t/2022-01-25-issue-with-tls-alpn-01-validation-method/170450 -var letsEncryptFixDeployTime = time.Date(2022, time.January, 26, 00, 48, 0, 0, time.UTC) - -// isRevokedLetsEncrypt returns whether the certificate is likely to be part of -// a batch of certificates revoked by Let's Encrypt in January 2022. This check -// can be safely removed from May 2022. -func isRevokedLetsEncrypt(cert *x509.Certificate) bool { - O := cert.Issuer.Organization - return len(O) == 1 && O[0] == "Let's Encrypt" && - cert.NotBefore.Before(letsEncryptFixDeployTime) -} - -type lockedMathRand struct { - sync.Mutex - rnd *mathrand.Rand -} - -func (r *lockedMathRand) int63n(max int64) int64 { - r.Lock() - n := r.rnd.Int63n(max) - r.Unlock() - return n -} - -// For easier testing. -var ( - // Called when a state is removed. - testDidRemoveState = func(certKey) {} -) diff --git a/internal/autocert/autocert_test.go b/internal/autocert/autocert_test.go deleted file mode 100644 index aab375c..0000000 --- a/internal/autocert/autocert_test.go +++ /dev/null @@ -1,1063 +0,0 @@ -// Copyright 2016 The Go Authors. All rights reserved. -// Use of this source code is governed by a BSD-style -// license that can be found in the LICENSE file. - -package autocert - -import ( - "bytes" - "context" - "crypto" - "crypto/ecdsa" - "crypto/elliptic" - "crypto/rand" - "crypto/rsa" - "crypto/tls" - "crypto/x509" - "crypto/x509/pkix" - "encoding/asn1" - "fmt" - "io" - "math/big" - "net/http" - "net/http/httptest" - "reflect" - "strings" - "sync" - "testing" - "time" - - "github.com/sr/tsproxy/internal/autocert/internal/acmetest" - - "golang.org/x/crypto/acme" -) - -var ( - exampleDomain = "example.org" - exampleCertKey = certKey{domain: exampleDomain} - exampleCertKeyRSA = certKey{domain: exampleDomain, isRSA: true} -) - -type memCache struct { - t *testing.T - mu sync.Mutex - keyData map[string][]byte -} - -func (m *memCache) Get(ctx context.Context, key string) ([]byte, error) { - m.mu.Lock() - defer m.mu.Unlock() - - v, ok := m.keyData[key] - if !ok { - return nil, ErrCacheMiss - } - return v, nil -} - -// filenameSafe returns whether all characters in s are printable ASCII -// and safe to use in a filename on most filesystems. -func filenameSafe(s string) bool { - for _, c := range s { - if c < 0x20 || c > 0x7E { - return false - } - switch c { - case '\\', '/', ':', '*', '?', '"', '<', '>', '|': - return false - } - } - return true -} - -func (m *memCache) Put(ctx context.Context, key string, data []byte) error { - if !filenameSafe(key) { - m.t.Errorf("invalid characters in cache key %q", key) - } - - m.mu.Lock() - defer m.mu.Unlock() - - m.keyData[key] = data - return nil -} - -func (m *memCache) Delete(ctx context.Context, key string) error { - m.mu.Lock() - defer m.mu.Unlock() - - delete(m.keyData, key) - return nil -} - -func newMemCache(t *testing.T) *memCache { - return &memCache{ - t: t, - keyData: make(map[string][]byte), - } -} - -func (m *memCache) numCerts() int { - m.mu.Lock() - defer m.mu.Unlock() - - res := 0 - for key := range m.keyData { - if strings.HasSuffix(key, "+token") || - strings.HasSuffix(key, "+key") || - strings.HasSuffix(key, "+http-01") { - continue - } - res++ - } - return res -} - -func dummyCert(pub interface{}, san ...string) ([]byte, error) { - return dateDummyCert(pub, time.Now(), time.Now().Add(90*24*time.Hour), san...) -} - -func dateDummyCert(pub interface{}, start, end time.Time, san ...string) ([]byte, error) { - // use EC key to run faster on 386 - key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - return nil, err - } - t := &x509.Certificate{ - SerialNumber: randomSerial(), - NotBefore: start, - NotAfter: end, - BasicConstraintsValid: true, - KeyUsage: x509.KeyUsageKeyEncipherment, - DNSNames: san, - } - if pub == nil { - pub = &key.PublicKey - } - return x509.CreateCertificate(rand.Reader, t, t, pub, key) -} - -func randomSerial() *big.Int { - serial, err := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 32)) - if err != nil { - panic(err) - } - return serial -} - -type algorithmSupport int - -const ( - algRSA algorithmSupport = iota - algECDSA -) - -func clientHelloInfo(sni string, alg algorithmSupport) *tls.ClientHelloInfo { - hello := &tls.ClientHelloInfo{ - ServerName: sni, - CipherSuites: []uint16{tls.TLS_ECDHE_RSA_WITH_CHACHA20_POLY1305}, - } - if alg == algECDSA { - hello.CipherSuites = append(hello.CipherSuites, tls.TLS_ECDHE_ECDSA_WITH_CHACHA20_POLY1305) - } - return hello -} - -func testManager(t *testing.T) *Manager { - man := &Manager{ - Prompt: AcceptTOS, - Cache: newMemCache(t), - } - t.Cleanup(man.stopRenew) - return man -} - -func TestGetCertificate(t *testing.T) { - tests := []struct { - name string - hello *tls.ClientHelloInfo - domain string - expectError string - prepare func(t *testing.T, man *Manager, s *acmetest.CAServer) - verify func(t *testing.T, man *Manager, leaf *x509.Certificate) - disableALPN bool - disableHTTP bool - }{ - { - name: "ALPN", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - disableHTTP: true, - }, - { - name: "HTTP", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - disableALPN: true, - }, - { - name: "nilPrompt", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - prepare: func(t *testing.T, man *Manager, s *acmetest.CAServer) { - man.Prompt = nil - }, - expectError: "Manager.Prompt not set", - }, - { - name: "trailingDot", - hello: clientHelloInfo("example.org.", algECDSA), - domain: "example.org", - }, - { - name: "unicodeIDN", - hello: clientHelloInfo("éé.com", algECDSA), - domain: "xn--9caa.com", - }, - { - name: "unicodeIDN/mixedCase", - hello: clientHelloInfo("éÉ.com", algECDSA), - domain: "xn--9caa.com", - }, - { - name: "upperCase", - hello: clientHelloInfo("EXAMPLE.ORG", algECDSA), - domain: "example.org", - }, - { - name: "goodCache", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - prepare: func(t *testing.T, man *Manager, s *acmetest.CAServer) { - // Make a valid cert and cache it. - c := s.Start().LeafCert(exampleDomain, "ECDSA", - // Use a time before the Let's Encrypt revocation cutoff to also test - // that non-Let's Encrypt certificates are not renewed. - time.Date(2022, time.January, 1, 0, 0, 0, 0, time.UTC), - time.Date(2122, time.January, 1, 0, 0, 0, 0, time.UTC), - ) - if err := man.cachePut(context.Background(), exampleCertKey, c); err != nil { - t.Fatalf("man.cachePut: %v", err) - } - }, - // Break the server to check that the cache is used. - disableALPN: true, disableHTTP: true, - }, - { - name: "expiredCache", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - prepare: func(t *testing.T, man *Manager, s *acmetest.CAServer) { - // Make an expired cert and cache it. - c := s.Start().LeafCert(exampleDomain, "ECDSA", time.Now().Add(-10*time.Minute), time.Now().Add(-5*time.Minute)) - if err := man.cachePut(context.Background(), exampleCertKey, c); err != nil { - t.Fatalf("man.cachePut: %v", err) - } - }, - }, - { - name: "forceRSA", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - prepare: func(t *testing.T, man *Manager, s *acmetest.CAServer) { - man.ForceRSA = true - }, - verify: func(t *testing.T, man *Manager, leaf *x509.Certificate) { - if _, ok := leaf.PublicKey.(*ecdsa.PublicKey); !ok { - t.Errorf("leaf.PublicKey is %T; want *ecdsa.PublicKey", leaf.PublicKey) - } - }, - }, - { - name: "goodLetsEncrypt", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - prepare: func(t *testing.T, man *Manager, s *acmetest.CAServer) { - // Make a valid certificate issued after the TLS-ALPN-01 - // revocation window and cache it. - s.IssuerName(pkix.Name{Country: []string{"US"}, - Organization: []string{"Let's Encrypt"}, CommonName: "R3"}) - c := s.Start().LeafCert(exampleDomain, "ECDSA", - time.Date(2022, time.January, 26, 12, 0, 0, 0, time.UTC), - time.Date(2122, time.January, 1, 0, 0, 0, 0, time.UTC), - ) - if err := man.cachePut(context.Background(), exampleCertKey, c); err != nil { - t.Fatalf("man.cachePut: %v", err) - } - }, - // Break the server to check that the cache is used. - disableALPN: true, disableHTTP: true, - }, - { - name: "revokedLetsEncrypt", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - prepare: func(t *testing.T, man *Manager, s *acmetest.CAServer) { - // Make a certificate issued during the TLS-ALPN-01 - // revocation window and cache it. - s.IssuerName(pkix.Name{Country: []string{"US"}, - Organization: []string{"Let's Encrypt"}, CommonName: "R3"}) - c := s.Start().LeafCert(exampleDomain, "ECDSA", - time.Date(2022, time.January, 1, 0, 0, 0, 0, time.UTC), - time.Date(2122, time.January, 1, 0, 0, 0, 0, time.UTC), - ) - if err := man.cachePut(context.Background(), exampleCertKey, c); err != nil { - t.Fatalf("man.cachePut: %v", err) - } - }, - verify: func(t *testing.T, man *Manager, leaf *x509.Certificate) { - if leaf.NotBefore.Before(time.Now().Add(-10 * time.Minute)) { - t.Error("certificate was not reissued") - } - }, - }, - { - // TestGetCertificate/tokenCache tests the fallback of token - // certificate fetches to cache when Manager.certTokens misses. - name: "tokenCacheALPN", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - prepare: func(t *testing.T, man *Manager, s *acmetest.CAServer) { - // Make a separate manager with a shared cache, simulating - // separate nodes that serve requests for the same domain. - man2 := testManager(t) - man2.Cache = man.Cache - // Redirect the verification request to man2, although the - // client request will hit man, testing that they can complete a - // verification by communicating through the cache. - s.ResolveGetCertificate("example.org", man2.GetCertificate) - }, - // Drop the default verification paths. - disableALPN: true, - }, - { - name: "tokenCacheHTTP", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - prepare: func(t *testing.T, man *Manager, s *acmetest.CAServer) { - man2 := testManager(t) - man2.Cache = man.Cache - s.ResolveHandler("example.org", man2.HTTPHandler(nil)) - }, - disableHTTP: true, - }, - { - name: "ecdsa", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - verify: func(t *testing.T, man *Manager, leaf *x509.Certificate) { - if _, ok := leaf.PublicKey.(*ecdsa.PublicKey); !ok { - t.Error("an ECDSA client was served a non-ECDSA certificate") - } - }, - }, - { - name: "rsa", - hello: clientHelloInfo("example.org", algRSA), - domain: "example.org", - verify: func(t *testing.T, man *Manager, leaf *x509.Certificate) { - if _, ok := leaf.PublicKey.(*rsa.PublicKey); !ok { - t.Error("an RSA client was served a non-RSA certificate") - } - }, - }, - { - name: "wrongCacheKeyType", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - prepare: func(t *testing.T, man *Manager, s *acmetest.CAServer) { - // Make an RSA cert and cache it without suffix. - c := s.Start().LeafCert(exampleDomain, "RSA", time.Now(), time.Now().Add(90*24*time.Hour)) - if err := man.cachePut(context.Background(), exampleCertKey, c); err != nil { - t.Fatalf("man.cachePut: %v", err) - } - }, - verify: func(t *testing.T, man *Manager, leaf *x509.Certificate) { - // The RSA cached cert should be silently ignored and replaced. - if _, ok := leaf.PublicKey.(*ecdsa.PublicKey); !ok { - t.Error("an ECDSA client was served a non-ECDSA certificate") - } - if numCerts := man.Cache.(*memCache).numCerts(); numCerts != 1 { - t.Errorf("found %d certificates in cache; want %d", numCerts, 1) - } - }, - }, - { - name: "almostExpiredCache", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - prepare: func(t *testing.T, man *Manager, s *acmetest.CAServer) { - man.RenewBefore = 24 * time.Hour - // Cache an almost expired cert. - c := s.Start().LeafCert(exampleDomain, "ECDSA", time.Now(), time.Now().Add(10*time.Minute)) - if err := man.cachePut(context.Background(), exampleCertKey, c); err != nil { - t.Fatalf("man.cachePut: %v", err) - } - }, - }, - { - name: "provideExternalAuth", - hello: clientHelloInfo("example.org", algECDSA), - domain: "example.org", - prepare: func(t *testing.T, man *Manager, s *acmetest.CAServer) { - s.ExternalAccountRequired() - - man.ExternalAccountBinding = &acme.ExternalAccountBinding{ - KID: "test-key", - Key: make([]byte, 32), - } - }, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - man := testManager(t) - s := acmetest.NewCAServer(t) - if !tt.disableALPN { - s.ResolveGetCertificate(tt.domain, man.GetCertificate) - } - if !tt.disableHTTP { - s.ResolveHandler(tt.domain, man.HTTPHandler(nil)) - } - - if tt.prepare != nil { - tt.prepare(t, man, s) - } - - s.Start() - - man.Client = &acme.Client{DirectoryURL: s.URL()} - - var tlscert *tls.Certificate - var err error - done := make(chan struct{}) - go func() { - tlscert, err = man.GetCertificate(tt.hello) - close(done) - }() - select { - case <-time.After(time.Minute): - t.Fatal("man.GetCertificate took too long to return") - case <-done: - } - if tt.expectError != "" { - if err == nil { - t.Fatal("expected error, got certificate") - } - if !strings.Contains(err.Error(), tt.expectError) { - t.Errorf("got %q, expected %q", err, tt.expectError) - } - return - } - if err != nil { - t.Fatalf("man.GetCertificate: %v", err) - } - - leaf, err := x509.ParseCertificate(tlscert.Certificate[0]) - if err != nil { - t.Fatal(err) - } - opts := x509.VerifyOptions{ - DNSName: tt.domain, - Intermediates: x509.NewCertPool(), - Roots: s.Roots(), - } - for _, cert := range tlscert.Certificate[1:] { - c, err := x509.ParseCertificate(cert) - if err != nil { - t.Fatal(err) - } - opts.Intermediates.AddCert(c) - } - if _, err := leaf.Verify(opts); err != nil { - t.Error(err) - } - - if san := leaf.DNSNames[0]; san != tt.domain { - t.Errorf("got SAN %q, expected %q", san, tt.domain) - } - - if tt.verify != nil { - tt.verify(t, man, leaf) - } - }) - } -} - -func TestGetCertificate_failedAttempt(t *testing.T) { - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusBadRequest) - })) - defer ts.Close() - - d := createCertRetryAfter - f := testDidRemoveState - defer func() { - createCertRetryAfter = d - testDidRemoveState = f - }() - createCertRetryAfter = 0 - done := make(chan struct{}) - testDidRemoveState = func(ck certKey) { - if ck != exampleCertKey { - t.Errorf("testDidRemoveState: domain = %v; want %v", ck, exampleCertKey) - } - close(done) - } - - man := &Manager{ - Prompt: AcceptTOS, - Client: &acme.Client{ - DirectoryURL: ts.URL, - }, - } - defer man.stopRenew() - hello := clientHelloInfo(exampleDomain, algECDSA) - if _, err := man.GetCertificate(hello); err == nil { - t.Error("GetCertificate: err is nil") - } - select { - case <-time.After(5 * time.Second): - t.Errorf("took too long to remove the %q state", exampleCertKey) - case <-done: - man.stateMu.Lock() - defer man.stateMu.Unlock() - if v, exist := man.state[exampleCertKey]; exist { - t.Errorf("state exists for %v: %+v", exampleCertKey, v) - } - } -} - -func TestRevokeFailedAuthz(t *testing.T) { - ca := acmetest.NewCAServer(t) - // Make the authz unfulfillable on the client side, so it will be left - // pending at the end of the verification attempt. - ca.ChallengeTypes("fake-01", "fake-02") - ca.Start() - - m := testManager(t) - m.Client = &acme.Client{DirectoryURL: ca.URL()} - - _, err := m.GetCertificate(clientHelloInfo("example.org", algECDSA)) - if err == nil { - t.Fatal("expected GetCertificate to fail") - } - - start := time.Now() - for time.Since(start) < 3*time.Second { - authz, err := m.Client.GetAuthorization(context.Background(), ca.URL()+"/authz/0") - if err != nil { - t.Fatal(err) - } - if authz.Status == acme.StatusDeactivated { - return - } - time.Sleep(50 * time.Millisecond) - } - t.Error("revocations took too long") - -} - -func TestHTTPHandlerDefaultFallback(t *testing.T) { - tt := []struct { - method, url string - wantCode int - wantLocation string - }{ - {"GET", "http://example.org", 302, "https://example.org/"}, - {"GET", "http://example.org/foo", 302, "https://example.org/foo"}, - {"GET", "http://example.org/foo/bar/", 302, "https://example.org/foo/bar/"}, - {"GET", "http://example.org/?a=b", 302, "https://example.org/?a=b"}, - {"GET", "http://example.org/foo?a=b", 302, "https://example.org/foo?a=b"}, - {"GET", "http://example.org:80/foo?a=b", 302, "https://example.org:443/foo?a=b"}, - {"GET", "http://example.org:80/foo%20bar", 302, "https://example.org:443/foo%20bar"}, - {"GET", "http://[2602:d1:xxxx::c60a]:1234", 302, "https://[2602:d1:xxxx::c60a]:443/"}, - {"GET", "http://[2602:d1:xxxx::c60a]", 302, "https://[2602:d1:xxxx::c60a]/"}, - {"GET", "http://[2602:d1:xxxx::c60a]/foo?a=b", 302, "https://[2602:d1:xxxx::c60a]/foo?a=b"}, - {"HEAD", "http://example.org", 302, "https://example.org/"}, - {"HEAD", "http://example.org/foo", 302, "https://example.org/foo"}, - {"HEAD", "http://example.org/foo/bar/", 302, "https://example.org/foo/bar/"}, - {"HEAD", "http://example.org/?a=b", 302, "https://example.org/?a=b"}, - {"HEAD", "http://example.org/foo?a=b", 302, "https://example.org/foo?a=b"}, - {"POST", "http://example.org", 400, ""}, - {"PUT", "http://example.org", 400, ""}, - {"GET", "http://example.org/.well-known/acme-challenge/x", 404, ""}, - } - var m Manager - h := m.HTTPHandler(nil) - for i, test := range tt { - r := httptest.NewRequest(test.method, test.url, nil) - w := httptest.NewRecorder() - h.ServeHTTP(w, r) - if w.Code != test.wantCode { - t.Errorf("%d: w.Code = %d; want %d", i, w.Code, test.wantCode) - t.Errorf("%d: body: %s", i, w.Body.Bytes()) - } - if v := w.Header().Get("Location"); v != test.wantLocation { - t.Errorf("%d: Location = %q; want %q", i, v, test.wantLocation) - } - } -} - -func TestAccountKeyCache(t *testing.T) { - m := Manager{Cache: newMemCache(t)} - ctx := context.Background() - k1, err := m.accountKey(ctx) - if err != nil { - t.Fatal(err) - } - k2, err := m.accountKey(ctx) - if err != nil { - t.Fatal(err) - } - if !reflect.DeepEqual(k1, k2) { - t.Errorf("account keys don't match: k1 = %#v; k2 = %#v", k1, k2) - } -} - -func TestCache(t *testing.T) { - ecdsaKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - t.Fatal(err) - } - cert, err := dummyCert(ecdsaKey.Public(), exampleDomain) - if err != nil { - t.Fatal(err) - } - ecdsaCert := &tls.Certificate{ - Certificate: [][]byte{cert}, - PrivateKey: ecdsaKey, - } - - rsaKey, err := rsa.GenerateKey(rand.Reader, 512) - if err != nil { - t.Fatal(err) - } - cert, err = dummyCert(rsaKey.Public(), exampleDomain) - if err != nil { - t.Fatal(err) - } - rsaCert := &tls.Certificate{ - Certificate: [][]byte{cert}, - PrivateKey: rsaKey, - } - - man := &Manager{Cache: newMemCache(t)} - defer man.stopRenew() - ctx := context.Background() - - if err := man.cachePut(ctx, exampleCertKey, ecdsaCert); err != nil { - t.Fatalf("man.cachePut: %v", err) - } - if err := man.cachePut(ctx, exampleCertKeyRSA, rsaCert); err != nil { - t.Fatalf("man.cachePut: %v", err) - } - - res, err := man.cacheGet(ctx, exampleCertKey) - if err != nil { - t.Fatalf("man.cacheGet: %v", err) - } - if res == nil || !bytes.Equal(res.Certificate[0], ecdsaCert.Certificate[0]) { - t.Errorf("man.cacheGet = %+v; want %+v", res, ecdsaCert) - } - - res, err = man.cacheGet(ctx, exampleCertKeyRSA) - if err != nil { - t.Fatalf("man.cacheGet: %v", err) - } - if res == nil || !bytes.Equal(res.Certificate[0], rsaCert.Certificate[0]) { - t.Errorf("man.cacheGet = %+v; want %+v", res, rsaCert) - } -} - -func TestHostWhitelist(t *testing.T) { - policy := HostWhitelist("example.com", "EXAMPLE.ORG", "*.example.net", "éÉ.com") - tt := []struct { - host string - allow bool - }{ - {"example.com", true}, - {"example.org", true}, - {"xn--9caa.com", true}, // éé.com - {"one.example.com", false}, - {"two.example.org", false}, - {"three.example.net", false}, - {"dummy", false}, - } - for i, test := range tt { - err := policy(nil, test.host) - if err != nil && test.allow { - t.Errorf("%d: policy(%q): %v; want nil", i, test.host, err) - } - if err == nil && !test.allow { - t.Errorf("%d: policy(%q): nil; want an error", i, test.host) - } - } -} - -func TestValidCert(t *testing.T) { - key1, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - t.Fatal(err) - } - key2, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - t.Fatal(err) - } - key3, err := rsa.GenerateKey(rand.Reader, 512) - if err != nil { - t.Fatal(err) - } - cert1, err := dummyCert(key1.Public(), "example.org") - if err != nil { - t.Fatal(err) - } - cert2, err := dummyCert(key2.Public(), "example.org") - if err != nil { - t.Fatal(err) - } - cert3, err := dummyCert(key3.Public(), "example.org") - if err != nil { - t.Fatal(err) - } - now := time.Now() - early, err := dateDummyCert(key1.Public(), now.Add(time.Hour), now.Add(2*time.Hour), "example.org") - if err != nil { - t.Fatal(err) - } - expired, err := dateDummyCert(key1.Public(), now.Add(-2*time.Hour), now.Add(-time.Hour), "example.org") - if err != nil { - t.Fatal(err) - } - - tt := []struct { - ck certKey - key crypto.Signer - cert [][]byte - ok bool - }{ - {certKey{domain: "example.org"}, key1, [][]byte{cert1}, true}, - {certKey{domain: "example.org", isRSA: true}, key3, [][]byte{cert3}, true}, - {certKey{domain: "example.org"}, key1, [][]byte{cert1, cert2, cert3}, true}, - {certKey{domain: "example.org"}, key1, [][]byte{cert1, {1}}, false}, - {certKey{domain: "example.org"}, key1, [][]byte{{1}}, false}, - {certKey{domain: "example.org"}, key1, [][]byte{cert2}, false}, - {certKey{domain: "example.org"}, key2, [][]byte{cert1}, false}, - {certKey{domain: "example.org"}, key1, [][]byte{cert3}, false}, - {certKey{domain: "example.org"}, key3, [][]byte{cert1}, false}, - {certKey{domain: "example.net"}, key1, [][]byte{cert1}, false}, - {certKey{domain: "example.org"}, key1, [][]byte{early}, false}, - {certKey{domain: "example.org"}, key1, [][]byte{expired}, false}, - {certKey{domain: "example.org", isRSA: true}, key1, [][]byte{cert1}, false}, - {certKey{domain: "example.org"}, key3, [][]byte{cert3}, false}, - } - for i, test := range tt { - leaf, err := validCert(test.ck, test.cert, test.key, now) - if err != nil && test.ok { - t.Errorf("%d: err = %v", i, err) - } - if err == nil && !test.ok { - t.Errorf("%d: err is nil", i) - } - if err == nil && test.ok && leaf == nil { - t.Errorf("%d: leaf is nil", i) - } - } -} - -type cacheGetFunc func(ctx context.Context, key string) ([]byte, error) - -func (f cacheGetFunc) Get(ctx context.Context, key string) ([]byte, error) { - return f(ctx, key) -} - -func (f cacheGetFunc) Put(ctx context.Context, key string, data []byte) error { - return fmt.Errorf("unsupported Put of %q = %q", key, data) -} - -func (f cacheGetFunc) Delete(ctx context.Context, key string) error { - return fmt.Errorf("unsupported Delete of %q", key) -} - -func TestManagerGetCertificateBogusSNI(t *testing.T) { - m := Manager{ - Prompt: AcceptTOS, - Cache: cacheGetFunc(func(ctx context.Context, key string) ([]byte, error) { - return nil, fmt.Errorf("cache.Get of %s", key) - }), - } - tests := []struct { - name string - wantErr string - }{ - {"foo.com", "cache.Get of foo.com"}, - {"foo.com.", "cache.Get of foo.com"}, - {`a\b.com`, "acme/autocert: server name contains invalid character"}, - {`a/b.com`, "acme/autocert: server name contains invalid character"}, - {"", "acme/autocert: missing server name"}, - {"foo", "acme/autocert: server name component count invalid"}, - {".foo", "acme/autocert: server name component count invalid"}, - {"foo.", "acme/autocert: server name component count invalid"}, - {"fo.o", "cache.Get of fo.o"}, - } - for _, tt := range tests { - _, err := m.GetCertificate(clientHelloInfo(tt.name, algECDSA)) - got := fmt.Sprint(err) - if got != tt.wantErr { - t.Errorf("GetCertificate(SNI = %q) = %q; want %q", tt.name, got, tt.wantErr) - } - } -} - -func TestCertRequest(t *testing.T) { - key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - t.Fatal(err) - } - // An extension from RFC7633. Any will do. - ext := pkix.Extension{ - Id: asn1.ObjectIdentifier{1, 3, 6, 1, 5, 5, 7, 1}, - Value: []byte("dummy"), - } - b, err := certRequest(key, "example.org", []pkix.Extension{ext}) - if err != nil { - t.Fatalf("certRequest: %v", err) - } - r, err := x509.ParseCertificateRequest(b) - if err != nil { - t.Fatalf("ParseCertificateRequest: %v", err) - } - var found bool - for _, v := range r.Extensions { - if v.Id.Equal(ext.Id) { - found = true - break - } - } - if !found { - t.Errorf("want %v in Extensions: %v", ext, r.Extensions) - } -} - -func TestSupportsECDSA(t *testing.T) { - tests := []struct { - CipherSuites []uint16 - SignatureSchemes []tls.SignatureScheme - SupportedCurves []tls.CurveID - ecdsaOk bool - }{ - {[]uint16{ - tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256, - }, nil, nil, false}, - {[]uint16{ - tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256, tls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256, - }, nil, nil, true}, - - // SignatureSchemes limits, not extends, CipherSuites - {[]uint16{ - tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256, - }, []tls.SignatureScheme{ - tls.PKCS1WithSHA256, tls.ECDSAWithP256AndSHA256, - }, nil, false}, - {[]uint16{ - tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256, tls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256, - }, []tls.SignatureScheme{ - tls.PKCS1WithSHA256, - }, nil, false}, - {[]uint16{ - tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256, tls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256, - }, []tls.SignatureScheme{ - tls.PKCS1WithSHA256, tls.ECDSAWithP256AndSHA256, - }, nil, true}, - - {[]uint16{ - tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256, tls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256, - }, []tls.SignatureScheme{ - tls.PKCS1WithSHA256, tls.ECDSAWithP256AndSHA256, - }, []tls.CurveID{ - tls.CurveP521, - }, false}, - {[]uint16{ - tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256, tls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256, - }, []tls.SignatureScheme{ - tls.PKCS1WithSHA256, tls.ECDSAWithP256AndSHA256, - }, []tls.CurveID{ - tls.CurveP256, - tls.CurveP521, - }, true}, - } - for i, tt := range tests { - result := supportsECDSA(&tls.ClientHelloInfo{ - CipherSuites: tt.CipherSuites, - SignatureSchemes: tt.SignatureSchemes, - SupportedCurves: tt.SupportedCurves, - }) - if result != tt.ecdsaOk { - t.Errorf("%d: supportsECDSA = %v; want %v", i, result, tt.ecdsaOk) - } - } -} - -func TestEndToEndALPN(t *testing.T) { - const domain = "example.org" - - // ACME CA server - ca := acmetest.NewCAServer(t).Start() - - // User HTTPS server. - m := &Manager{ - Prompt: AcceptTOS, - Client: &acme.Client{DirectoryURL: ca.URL()}, - } - us := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Write([]byte("OK")) - })) - us.TLS = &tls.Config{ - NextProtos: []string{"http/1.1", acme.ALPNProto}, - GetCertificate: func(hello *tls.ClientHelloInfo) (*tls.Certificate, error) { - cert, err := m.GetCertificate(hello) - if err != nil { - t.Errorf("m.GetCertificate: %v", err) - } - return cert, err - }, - } - us.StartTLS() - defer us.Close() - // In TLS-ALPN challenge verification, CA connects to the domain:443 in question. - // Because the domain won't resolve in tests, we need to tell the CA - // where to dial to instead. - ca.Resolve(domain, strings.TrimPrefix(us.URL, "https://")) - - // A client visiting user's HTTPS server. - tr := &http.Transport{ - TLSClientConfig: &tls.Config{ - RootCAs: ca.Roots(), - ServerName: domain, - }, - } - client := &http.Client{Transport: tr} - res, err := client.Get(us.URL) - if err != nil { - t.Fatal(err) - } - defer res.Body.Close() - b, err := io.ReadAll(res.Body) - if err != nil { - t.Fatal(err) - } - if v := string(b); v != "OK" { - t.Errorf("user server response: %q; want 'OK'", v) - } -} - -func TestEndToEndHTTP(t *testing.T) { - const domain = "example.org" - - // ACME CA server. - ca := acmetest.NewCAServer(t).ChallengeTypes("http-01").Start() - - // User HTTP server for the ACME challenge. - m := testManager(t) - m.Client = &acme.Client{DirectoryURL: ca.URL()} - s := httptest.NewServer(m.HTTPHandler(nil)) - defer s.Close() - - // User HTTPS server. - ss := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Write([]byte("OK")) - })) - ss.TLS = &tls.Config{ - NextProtos: []string{"http/1.1", acme.ALPNProto}, - GetCertificate: func(hello *tls.ClientHelloInfo) (*tls.Certificate, error) { - cert, err := m.GetCertificate(hello) - if err != nil { - t.Errorf("m.GetCertificate: %v", err) - } - return cert, err - }, - } - ss.StartTLS() - defer ss.Close() - - // Redirect the CA requests to the HTTP server. - ca.Resolve(domain, strings.TrimPrefix(s.URL, "http://")) - - // A client visiting user's HTTPS server. - tr := &http.Transport{ - TLSClientConfig: &tls.Config{ - RootCAs: ca.Roots(), - ServerName: domain, - }, - } - client := &http.Client{Transport: tr} - res, err := client.Get(ss.URL) - if err != nil { - t.Fatal(err) - } - defer res.Body.Close() - b, err := io.ReadAll(res.Body) - if err != nil { - t.Fatal(err) - } - if v := string(b); v != "OK" { - t.Errorf("user server response: %q; want 'OK'", v) - } -} - -func dnsSolver(ca *acmetest.CAServer) DNS01ChallengeSolver { - return func(ctx context.Context, domain, record string) (func() error, error) { - ca.PutDNSResponse(domain, record) - return func() error { return nil }, nil - } -} - -func TestEndToEndDNS(t *testing.T) { - const domain = "example.org" - - // ACME CA server - ca := acmetest.NewCAServer(t).ChallengeTypes("dns-01") - ca.Start() - - // User HTTPS server. - m := &Manager{ - Prompt: AcceptTOS, - Client: &acme.Client{DirectoryURL: ca.URL()}, - DNS01: dnsSolver(ca), - } - us := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Write([]byte("OK")) - })) - us.TLS = &tls.Config{ - GetCertificate: func(hello *tls.ClientHelloInfo) (*tls.Certificate, error) { - cert, err := m.GetCertificate(hello) - if err != nil { - t.Errorf("m.GetCertificate: %v", err) - } - return cert, err - }, - } - us.StartTLS() - defer us.Close() - - // A client visiting user's HTTPS server. - tr := &http.Transport{ - TLSClientConfig: &tls.Config{ - RootCAs: ca.Roots(), - ServerName: domain, - }, - } - client := &http.Client{Transport: tr} - res, err := client.Get(us.URL) - if err != nil { - t.Fatal(err) - } - defer res.Body.Close() - b, err := io.ReadAll(res.Body) - if err != nil { - t.Fatal(err) - } - if v := string(b); v != "OK" { - t.Errorf("user server response: %q; want 'OK'", v) - } -} diff --git a/internal/autocert/cache.go b/internal/autocert/cache.go deleted file mode 100644 index 758ab12..0000000 --- a/internal/autocert/cache.go +++ /dev/null @@ -1,135 +0,0 @@ -// Copyright 2016 The Go Authors. All rights reserved. -// Use of this source code is governed by a BSD-style -// license that can be found in the LICENSE file. - -package autocert - -import ( - "context" - "errors" - "os" - "path/filepath" -) - -// ErrCacheMiss is returned when a certificate is not found in cache. -var ErrCacheMiss = errors.New("acme/autocert: certificate cache miss") - -// Cache is used by Manager to store and retrieve previously obtained certificates -// and other account data as opaque blobs. -// -// Cache implementations should not rely on the key naming pattern. Keys can -// include any printable ASCII characters, except the following: \/:*?"<>| -type Cache interface { - // Get returns a certificate data for the specified key. - // If there's no such key, Get returns ErrCacheMiss. - Get(ctx context.Context, key string) ([]byte, error) - - // Put stores the data in the cache under the specified key. - // Underlying implementations may use any data storage format, - // as long as the reverse operation, Get, results in the original data. - Put(ctx context.Context, key string, data []byte) error - - // Delete removes a certificate data from the cache under the specified key. - // If there's no such key in the cache, Delete returns nil. - Delete(ctx context.Context, key string) error -} - -// DirCache implements Cache using a directory on the local filesystem. -// If the directory does not exist, it will be created with 0700 permissions. -type DirCache string - -// Get reads a certificate data from the specified file name. -func (d DirCache) Get(ctx context.Context, name string) ([]byte, error) { - name = filepath.Join(string(d), filepath.Clean("/"+name)) - var ( - data []byte - err error - done = make(chan struct{}) - ) - go func() { - data, err = os.ReadFile(name) - close(done) - }() - select { - case <-ctx.Done(): - return nil, ctx.Err() - case <-done: - } - if os.IsNotExist(err) { - return nil, ErrCacheMiss - } - return data, err -} - -// Put writes the certificate data to the specified file name. -// The file will be created with 0600 permissions. -func (d DirCache) Put(ctx context.Context, name string, data []byte) error { - if err := os.MkdirAll(string(d), 0700); err != nil { - return err - } - - done := make(chan struct{}) - var err error - go func() { - defer close(done) - var tmp string - if tmp, err = d.writeTempFile(name, data); err != nil { - return - } - defer os.Remove(tmp) - select { - case <-ctx.Done(): - // Don't overwrite the file if the context was canceled. - default: - newName := filepath.Join(string(d), filepath.Clean("/"+name)) - err = os.Rename(tmp, newName) - } - }() - select { - case <-ctx.Done(): - return ctx.Err() - case <-done: - } - return err -} - -// Delete removes the specified file name. -func (d DirCache) Delete(ctx context.Context, name string) error { - name = filepath.Join(string(d), filepath.Clean("/"+name)) - var ( - err error - done = make(chan struct{}) - ) - go func() { - err = os.Remove(name) - close(done) - }() - select { - case <-ctx.Done(): - return ctx.Err() - case <-done: - } - if err != nil && !os.IsNotExist(err) { - return err - } - return nil -} - -// writeTempFile writes b to a temporary file, closes the file and returns its path. -func (d DirCache) writeTempFile(prefix string, b []byte) (name string, reterr error) { - // TempFile uses 0600 permissions - f, err := os.CreateTemp(string(d), prefix) - if err != nil { - return "", err - } - defer func() { - if reterr != nil { - os.Remove(f.Name()) - } - }() - if _, err := f.Write(b); err != nil { - f.Close() - return "", err - } - return f.Name(), f.Close() -} diff --git a/internal/autocert/cache_test.go b/internal/autocert/cache_test.go deleted file mode 100644 index 582e6b0..0000000 --- a/internal/autocert/cache_test.go +++ /dev/null @@ -1,66 +0,0 @@ -// Copyright 2016 The Go Authors. All rights reserved. -// Use of this source code is governed by a BSD-style -// license that can be found in the LICENSE file. - -package autocert - -import ( - "context" - "os" - "path/filepath" - "reflect" - "testing" -) - -// make sure DirCache satisfies Cache interface -var _ Cache = DirCache("/") - -func TestDirCache(t *testing.T) { - dir, err := os.MkdirTemp("", "autocert") - if err != nil { - t.Fatal(err) - } - defer os.RemoveAll(dir) - dir = filepath.Join(dir, "certs") // a nonexistent dir - cache := DirCache(dir) - ctx := context.Background() - - // test cache miss - if _, err := cache.Get(ctx, "nonexistent"); err != ErrCacheMiss { - t.Errorf("get: %v; want ErrCacheMiss", err) - } - - // test put/get - b1 := []byte{1} - if err := cache.Put(ctx, "dummy", b1); err != nil { - t.Fatalf("put: %v", err) - } - b2, err := cache.Get(ctx, "dummy") - if err != nil { - t.Fatalf("get: %v", err) - } - if !reflect.DeepEqual(b1, b2) { - t.Errorf("b1 = %v; want %v", b1, b2) - } - name := filepath.Join(dir, "dummy") - if _, err := os.Stat(name); err != nil { - t.Error(err) - } - - // test put deletes temp file - tmp, err := filepath.Glob(name + "?*") - if err != nil { - t.Error(err) - } - if tmp != nil { - t.Errorf("temp file exists: %s", tmp) - } - - // test delete - if err := cache.Delete(ctx, "dummy"); err != nil { - t.Fatalf("delete: %v", err) - } - if _, err := cache.Get(ctx, "dummy"); err != ErrCacheMiss { - t.Errorf("get: %v; want ErrCacheMiss", err) - } -} diff --git a/internal/autocert/example_test.go b/internal/autocert/example_test.go deleted file mode 100644 index 6c7458b..0000000 --- a/internal/autocert/example_test.go +++ /dev/null @@ -1,35 +0,0 @@ -// Copyright 2017 The Go Authors. All rights reserved. -// Use of this source code is governed by a BSD-style -// license that can be found in the LICENSE file. - -package autocert_test - -import ( - "fmt" - "log" - "net/http" - - "golang.org/x/crypto/acme/autocert" -) - -func ExampleNewListener() { - mux := http.NewServeMux() - mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { - fmt.Fprintf(w, "Hello, TLS user! Your config: %+v", r.TLS) - }) - log.Fatal(http.Serve(autocert.NewListener("example.com"), mux)) -} - -func ExampleManager() { - m := &autocert.Manager{ - Cache: autocert.DirCache("secret-dir"), - Prompt: autocert.AcceptTOS, - Email: "example@example.org", - HostPolicy: autocert.HostWhitelist("example.org", "www.example.org"), - } - s := &http.Server{ - Addr: ":https", - TLSConfig: m.TLSConfig(), - } - s.ListenAndServeTLS("", "") -} diff --git a/internal/autocert/internal/acmetest/ca.go b/internal/autocert/internal/acmetest/ca.go deleted file mode 100644 index 8fc4273..0000000 --- a/internal/autocert/internal/acmetest/ca.go +++ /dev/null @@ -1,817 +0,0 @@ -// Copyright 2018 The Go Authors. All rights reserved. -// Use of this source code is governed by a BSD-style -// license that can be found in the LICENSE file. - -// Package acmetest provides types for testing acme and autocert packages. -// -// TODO: Consider moving this to x/crypto/acme/internal/acmetest for acme tests as well. -package acmetest - -import ( - "context" - "crypto" - "crypto/ecdsa" - "crypto/elliptic" - "crypto/rand" - "crypto/rsa" - "crypto/tls" - "crypto/x509" - "crypto/x509/pkix" - "encoding/asn1" - "encoding/base64" - "encoding/json" - "encoding/pem" - "fmt" - "io" - "math/big" - "net" - "net/http" - "net/http/httptest" - "path" - "strconv" - "strings" - "sync" - "testing" - "time" - - "golang.org/x/crypto/acme" -) - -// CAServer is a simple test server which implements ACME spec bits needed for testing. -type CAServer struct { - rootKey crypto.Signer - rootCert []byte // DER encoding - rootTemplate *x509.Certificate - - t *testing.T - server *httptest.Server - issuer pkix.Name - challengeTypes []string - url string - roots *x509.CertPool - eabRequired bool - - mu sync.Mutex - certCount int // number of issued certs - acctRegistered bool // set once an account has been registered - domainAddr map[string]string // domain name to addr:port resolution - dnsResponses map[string]string // responses to dns challenges - domainGetCert map[string]getCertificateFunc // domain name to GetCertificate function - domainHandler map[string]http.Handler // domain name to Handle function - validAuthz map[string]*authorization // valid authz, keyed by domain name - authorizations []*authorization // all authz, index is used as ID - orders []*order // index is used as order ID - errors []error // encountered client errors -} - -type getCertificateFunc func(hello *tls.ClientHelloInfo) (*tls.Certificate, error) - -// NewCAServer creates a new ACME test server. The returned CAServer issues -// certs signed with the CA roots available in the Roots field. -func NewCAServer(t *testing.T) *CAServer { - ca := &CAServer{t: t, - challengeTypes: []string{"fake-01", "tls-alpn-01", "http-01", "dns-01"}, - domainAddr: make(map[string]string), - domainGetCert: make(map[string]getCertificateFunc), - domainHandler: make(map[string]http.Handler), - validAuthz: make(map[string]*authorization), - dnsResponses: make(map[string]string), - } - - ca.server = httptest.NewUnstartedServer(http.HandlerFunc(ca.handle)) - - r, err := rand.Int(rand.Reader, big.NewInt(1000000)) - if err != nil { - panic(fmt.Sprintf("rand.Int: %v", err)) - } - ca.issuer = pkix.Name{ - Organization: []string{"Test Acme Co"}, - CommonName: "Root CA " + r.String(), - } - - return ca -} - -func (ca *CAServer) generateRoot() { - key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - panic(fmt.Sprintf("ecdsa.GenerateKey: %v", err)) - } - tmpl := &x509.Certificate{ - SerialNumber: big.NewInt(1), - Subject: ca.issuer, - NotBefore: time.Now(), - NotAfter: time.Now().Add(365 * 24 * time.Hour), - KeyUsage: x509.KeyUsageCertSign, - BasicConstraintsValid: true, - IsCA: true, - } - der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &key.PublicKey, key) - if err != nil { - panic(fmt.Sprintf("x509.CreateCertificate: %v", err)) - } - cert, err := x509.ParseCertificate(der) - if err != nil { - panic(fmt.Sprintf("x509.ParseCertificate: %v", err)) - } - ca.roots = x509.NewCertPool() - ca.roots.AddCert(cert) - ca.rootKey = key - ca.rootCert = der - ca.rootTemplate = tmpl -} - -func (ca *CAServer) PutDNSResponse(domain, record string) { - ca.mu.Lock() - defer ca.mu.Unlock() - ca.dnsResponses[domain] = record -} - -// IssuerName sets the name of the issuing CA. -func (ca *CAServer) IssuerName(name pkix.Name) *CAServer { - if ca.url != "" { - panic("IssuerName must be called before Start") - } - ca.issuer = name - return ca -} - -// ChallengeTypes sets the supported challenge types. -func (ca *CAServer) ChallengeTypes(types ...string) *CAServer { - if ca.url != "" { - panic("ChallengeTypes must be called before Start") - } - ca.challengeTypes = types - return ca -} - -// URL returns the server address, after Start has been called. -func (ca *CAServer) URL() string { - if ca.url == "" { - panic("URL called before Start") - } - return ca.url -} - -// Roots returns a pool cointaining the CA root. -func (ca *CAServer) Roots() *x509.CertPool { - if ca.url == "" { - panic("Roots called before Start") - } - return ca.roots -} - -// ExternalAccountRequired makes an EAB JWS required for account registration. -func (ca *CAServer) ExternalAccountRequired() *CAServer { - if ca.url != "" { - panic("ExternalAccountRequired must be called before Start") - } - ca.eabRequired = true - return ca -} - -// Start starts serving requests. The server address becomes available in the -// URL field. -func (ca *CAServer) Start() *CAServer { - if ca.url == "" { - ca.generateRoot() - ca.server.Start() - ca.t.Cleanup(ca.server.Close) - ca.url = ca.server.URL - } - return ca -} - -func (ca *CAServer) serverURL(format string, arg ...interface{}) string { - return ca.server.URL + fmt.Sprintf(format, arg...) -} - -func (ca *CAServer) addr(domain string) (string, bool) { - ca.mu.Lock() - defer ca.mu.Unlock() - addr, ok := ca.domainAddr[domain] - return addr, ok -} - -func (ca *CAServer) getCert(domain string) (getCertificateFunc, bool) { - ca.mu.Lock() - defer ca.mu.Unlock() - f, ok := ca.domainGetCert[domain] - return f, ok -} - -func (ca *CAServer) getHandler(domain string) (http.Handler, bool) { - ca.mu.Lock() - defer ca.mu.Unlock() - h, ok := ca.domainHandler[domain] - return h, ok -} - -func (ca *CAServer) httpErrorf(w http.ResponseWriter, code int, format string, a ...interface{}) { - s := fmt.Sprintf(format, a...) - ca.t.Errorf(format, a...) - http.Error(w, s, code) -} - -// Resolve adds a domain to address resolution for the ca to dial to -// when validating challenges for the domain authorization. -func (ca *CAServer) Resolve(domain, addr string) { - ca.mu.Lock() - defer ca.mu.Unlock() - ca.domainAddr[domain] = addr -} - -// ResolveGetCertificate redirects TLS connections for domain to f when -// validating challenges for the domain authorization. -func (ca *CAServer) ResolveGetCertificate(domain string, f getCertificateFunc) { - ca.mu.Lock() - defer ca.mu.Unlock() - ca.domainGetCert[domain] = f -} - -// ResolveHandler redirects HTTP requests for domain to f when -// validating challenges for the domain authorization. -func (ca *CAServer) ResolveHandler(domain string, h http.Handler) { - ca.mu.Lock() - defer ca.mu.Unlock() - ca.domainHandler[domain] = h -} - -type discovery struct { - NewNonce string `json:"newNonce"` - NewAccount string `json:"newAccount"` - NewOrder string `json:"newOrder"` - NewAuthz string `json:"newAuthz"` - - Meta discoveryMeta `json:"meta,omitempty"` -} - -type discoveryMeta struct { - ExternalAccountRequired bool `json:"externalAccountRequired,omitempty"` -} - -type challenge struct { - URI string `json:"uri"` - Type string `json:"type"` - Token string `json:"token"` -} - -type authorization struct { - Status string `json:"status"` - Challenges []challenge `json:"challenges"` - - domain string - id int -} - -type order struct { - Status string `json:"status"` - AuthzURLs []string `json:"authorizations"` - FinalizeURL string `json:"finalize"` // CSR submit URL - CertURL string `json:"certificate"` // already issued cert - - leaf []byte // issued cert in DER format -} - -func (ca *CAServer) handle(w http.ResponseWriter, r *http.Request) { - ca.t.Logf("%s %s", r.Method, r.URL) - w.Header().Set("Replay-Nonce", "nonce") - // TODO: Verify nonce header for all POST requests. - - switch { - default: - ca.httpErrorf(w, http.StatusBadRequest, "unrecognized r.URL.Path: %s", r.URL.Path) - - // Discovery request. - case r.URL.Path == "/": - resp := &discovery{ - NewNonce: ca.serverURL("/new-nonce"), - NewAccount: ca.serverURL("/new-account"), - NewOrder: ca.serverURL("/new-order"), - Meta: discoveryMeta{ - ExternalAccountRequired: ca.eabRequired, - }, - } - if err := json.NewEncoder(w).Encode(resp); err != nil { - panic(fmt.Sprintf("discovery response: %v", err)) - } - - // Nonce requests. - case r.URL.Path == "/new-nonce": - // Nonce values are always set. Nothing else to do. - return - - // Client key registration request. - case r.URL.Path == "/new-account": - ca.mu.Lock() - defer ca.mu.Unlock() - if ca.acctRegistered { - ca.httpErrorf(w, http.StatusServiceUnavailable, "multiple accounts are not implemented") - return - } - ca.acctRegistered = true - - var req struct { - ExternalAccountBinding json.RawMessage - } - - if err := decodePayload(&req, r.Body); err != nil { - ca.httpErrorf(w, http.StatusBadRequest, err.Error()) - return - } - - if ca.eabRequired && len(req.ExternalAccountBinding) == 0 { - ca.httpErrorf(w, http.StatusBadRequest, "registration failed: no JWS for EAB") - return - } - - // TODO: Check the user account key against a ca.accountKeys? - w.Header().Set("Location", ca.serverURL("/accounts/1")) - w.WriteHeader(http.StatusCreated) - w.Write([]byte("{}")) - - // New order request. - case r.URL.Path == "/new-order": - var req struct { - Identifiers []struct{ Value string } - } - if err := decodePayload(&req, r.Body); err != nil { - ca.httpErrorf(w, http.StatusBadRequest, err.Error()) - return - } - ca.mu.Lock() - defer ca.mu.Unlock() - o := &order{Status: acme.StatusPending} - for _, id := range req.Identifiers { - z := ca.authz(id.Value) - o.AuthzURLs = append(o.AuthzURLs, ca.serverURL("/authz/%d", z.id)) - } - orderID := len(ca.orders) - ca.orders = append(ca.orders, o) - w.Header().Set("Location", ca.serverURL("/orders/%d", orderID)) - w.WriteHeader(http.StatusCreated) - if err := json.NewEncoder(w).Encode(o); err != nil { - panic(err) - } - - // Existing order status requests. - case strings.HasPrefix(r.URL.Path, "/orders/"): - ca.mu.Lock() - defer ca.mu.Unlock() - o, err := ca.storedOrder(strings.TrimPrefix(r.URL.Path, "/orders/")) - if err != nil { - ca.httpErrorf(w, http.StatusBadRequest, err.Error()) - return - } - if err := json.NewEncoder(w).Encode(o); err != nil { - panic(err) - } - - // Accept challenge requests. - case strings.HasPrefix(r.URL.Path, "/challenge/"): - parts := strings.Split(r.URL.Path, "/") - typ, id := parts[len(parts)-2], parts[len(parts)-1] - ca.mu.Lock() - supported := false - for _, suppTyp := range ca.challengeTypes { - if suppTyp == typ { - supported = true - } - } - a, err := ca.storedAuthz(id) - ca.mu.Unlock() - if !supported { - ca.httpErrorf(w, http.StatusBadRequest, "unsupported challenge: %v", typ) - return - } - if err != nil { - ca.httpErrorf(w, http.StatusBadRequest, "challenge accept: %v", err) - return - } - ca.validateChallenge(a, typ) - w.Write([]byte("{}")) - - // Get authorization status requests. - case strings.HasPrefix(r.URL.Path, "/authz/"): - var req struct{ Status string } - decodePayload(&req, r.Body) - deactivate := req.Status == "deactivated" - ca.mu.Lock() - defer ca.mu.Unlock() - authz, err := ca.storedAuthz(strings.TrimPrefix(r.URL.Path, "/authz/")) - if err != nil { - ca.httpErrorf(w, http.StatusNotFound, "%v", err) - return - } - if deactivate { - // Note we don't invalidate authorized orders as we should. - authz.Status = "deactivated" - ca.t.Logf("authz %d is now %s", authz.id, authz.Status) - ca.updatePendingOrders() - } - if err := json.NewEncoder(w).Encode(authz); err != nil { - panic(fmt.Sprintf("encoding authz %d: %v", authz.id, err)) - } - - // Certificate issuance request. - case strings.HasPrefix(r.URL.Path, "/new-cert/"): - ca.mu.Lock() - defer ca.mu.Unlock() - orderID := strings.TrimPrefix(r.URL.Path, "/new-cert/") - o, err := ca.storedOrder(orderID) - if err != nil { - ca.httpErrorf(w, http.StatusBadRequest, err.Error()) - return - } - if o.Status != acme.StatusReady { - ca.httpErrorf(w, http.StatusForbidden, "order status: %s", o.Status) - return - } - // Validate CSR request. - var req struct { - CSR string `json:"csr"` - } - decodePayload(&req, r.Body) - b, _ := base64.RawURLEncoding.DecodeString(req.CSR) - csr, err := x509.ParseCertificateRequest(b) - if err != nil { - ca.httpErrorf(w, http.StatusBadRequest, err.Error()) - return - } - // Issue the certificate. - der, err := ca.leafCert(csr) - if err != nil { - ca.httpErrorf(w, http.StatusBadRequest, "new-cert response: ca.leafCert: %v", err) - return - } - o.leaf = der - o.CertURL = ca.serverURL("/issued-cert/%s", orderID) - o.Status = acme.StatusValid - if err := json.NewEncoder(w).Encode(o); err != nil { - panic(err) - } - - // Already issued cert download requests. - case strings.HasPrefix(r.URL.Path, "/issued-cert/"): - ca.mu.Lock() - defer ca.mu.Unlock() - o, err := ca.storedOrder(strings.TrimPrefix(r.URL.Path, "/issued-cert/")) - if err != nil { - ca.httpErrorf(w, http.StatusBadRequest, err.Error()) - return - } - if o.Status != acme.StatusValid { - ca.httpErrorf(w, http.StatusForbidden, "order status: %s", o.Status) - return - } - w.Header().Set("Content-Type", "application/pem-certificate-chain") - pem.Encode(w, &pem.Block{Type: "CERTIFICATE", Bytes: o.leaf}) - pem.Encode(w, &pem.Block{Type: "CERTIFICATE", Bytes: ca.rootCert}) - } -} - -// storedOrder retrieves a previously created order at index i. -// It requires ca.mu to be locked. -func (ca *CAServer) storedOrder(i string) (*order, error) { - idx, err := strconv.Atoi(i) - if err != nil { - return nil, fmt.Errorf("storedOrder: %v", err) - } - if idx < 0 { - return nil, fmt.Errorf("storedOrder: invalid order index %d", idx) - } - if idx > len(ca.orders)-1 { - return nil, fmt.Errorf("storedOrder: no such order %d", idx) - } - - ca.updatePendingOrders() - return ca.orders[idx], nil -} - -// storedAuthz retrieves a previously created authz at index i. -// It requires ca.mu to be locked. -func (ca *CAServer) storedAuthz(i string) (*authorization, error) { - idx, err := strconv.Atoi(i) - if err != nil { - return nil, fmt.Errorf("storedAuthz: %v", err) - } - if idx < 0 { - return nil, fmt.Errorf("storedAuthz: invalid authz index %d", idx) - } - if idx > len(ca.authorizations)-1 { - return nil, fmt.Errorf("storedAuthz: no such authz %d", idx) - } - return ca.authorizations[idx], nil -} - -// authz returns an existing valid authorization for the identifier or creates a -// new one. It requires ca.mu to be locked. -func (ca *CAServer) authz(identifier string) *authorization { - authz, ok := ca.validAuthz[identifier] - if !ok { - authzId := len(ca.authorizations) - authz = &authorization{ - id: authzId, - domain: identifier, - Status: acme.StatusPending, - } - for _, typ := range ca.challengeTypes { - authz.Challenges = append(authz.Challenges, challenge{ - Type: typ, - URI: ca.serverURL("/challenge/%s/%d", typ, authzId), - Token: challengeToken(authz.domain, typ, authzId), - }) - } - ca.authorizations = append(ca.authorizations, authz) - } - return authz -} - -// leafCert issues a new certificate. -// It requires ca.mu to be locked. -func (ca *CAServer) leafCert(csr *x509.CertificateRequest) (der []byte, err error) { - ca.certCount++ // next leaf cert serial number - leaf := &x509.Certificate{ - SerialNumber: big.NewInt(int64(ca.certCount)), - Subject: pkix.Name{Organization: []string{"Test Acme Co"}}, - NotBefore: time.Now(), - NotAfter: time.Now().Add(90 * 24 * time.Hour), - KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment, - ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, - DNSNames: csr.DNSNames, - BasicConstraintsValid: true, - } - if len(csr.DNSNames) == 0 { - leaf.DNSNames = []string{csr.Subject.CommonName} - } - return x509.CreateCertificate(rand.Reader, leaf, ca.rootTemplate, csr.PublicKey, ca.rootKey) -} - -// LeafCert issues a leaf certificate. -func (ca *CAServer) LeafCert(name, keyType string, notBefore, notAfter time.Time) *tls.Certificate { - if ca.url == "" { - panic("LeafCert called before Start") - } - - ca.mu.Lock() - defer ca.mu.Unlock() - var pk crypto.Signer - switch keyType { - case "RSA": - var err error - pk, err = rsa.GenerateKey(rand.Reader, 1024) - if err != nil { - ca.t.Fatal(err) - } - case "ECDSA": - var err error - pk, err = ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - ca.t.Fatal(err) - } - default: - panic("LeafCert: unknown key type") - } - ca.certCount++ // next leaf cert serial number - leaf := &x509.Certificate{ - SerialNumber: big.NewInt(int64(ca.certCount)), - Subject: pkix.Name{Organization: []string{"Test Acme Co"}}, - NotBefore: notBefore, - NotAfter: notAfter, - KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment, - ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, - DNSNames: []string{name}, - BasicConstraintsValid: true, - } - der, err := x509.CreateCertificate(rand.Reader, leaf, ca.rootTemplate, pk.Public(), ca.rootKey) - if err != nil { - ca.t.Fatal(err) - } - return &tls.Certificate{ - Certificate: [][]byte{der}, - PrivateKey: pk, - } -} - -func (ca *CAServer) validateChallenge(authz *authorization, typ string) { - var err error - switch typ { - case "tls-alpn-01": - err = ca.verifyALPNChallenge(authz) - case "http-01": - err = ca.verifyHTTPChallenge(authz) - case "dns-01": - err = ca.verifyDNSChallenge(authz) - default: - panic(fmt.Sprintf("validation of %q is not implemented", typ)) - } - ca.mu.Lock() - defer ca.mu.Unlock() - if err != nil { - authz.Status = "invalid" - } else { - authz.Status = "valid" - ca.validAuthz[authz.domain] = authz - } - ca.t.Logf("validated %q for %q, err: %v", typ, authz.domain, err) - ca.t.Logf("authz %d is now %s", authz.id, authz.Status) - - ca.updatePendingOrders() -} - -func (ca *CAServer) updatePendingOrders() { - // Update all pending orders. - // An order becomes "ready" if all authorizations are "valid". - // An order becomes "invalid" if any authorization is "invalid". - // Status changes: https://tools.ietf.org/html/rfc8555#section-7.1.6 - for i, o := range ca.orders { - if o.Status != acme.StatusPending { - continue - } - - countValid, countInvalid := ca.validateAuthzURLs(o.AuthzURLs, i) - if countInvalid > 0 { - o.Status = acme.StatusInvalid - ca.t.Logf("order %d is now invalid", i) - continue - } - if countValid == len(o.AuthzURLs) { - o.Status = acme.StatusReady - o.FinalizeURL = ca.serverURL("/new-cert/%d", i) - ca.t.Logf("order %d is now ready", i) - } - } -} - -func (ca *CAServer) validateAuthzURLs(urls []string, orderNum int) (countValid, countInvalid int) { - for _, zurl := range urls { - z, err := ca.storedAuthz(path.Base(zurl)) - if err != nil { - ca.t.Logf("no authz %q for order %d", zurl, orderNum) - continue - } - if z.Status == acme.StatusInvalid { - countInvalid++ - } - if z.Status == acme.StatusValid { - countValid++ - } - } - return countValid, countInvalid -} - -func (ca *CAServer) verifyALPNChallenge(a *authorization) error { - const acmeALPNProto = "acme-tls/1" - - addr, haveAddr := ca.addr(a.domain) - getCert, haveGetCert := ca.getCert(a.domain) - if !haveAddr && !haveGetCert { - return fmt.Errorf("no resolution information for %q", a.domain) - } - if haveAddr && haveGetCert { - return fmt.Errorf("overlapping resolution information for %q", a.domain) - } - - var crt *x509.Certificate - switch { - case haveAddr: - conn, err := tls.Dial("tcp", addr, &tls.Config{ - ServerName: a.domain, - InsecureSkipVerify: true, - NextProtos: []string{acmeALPNProto}, - MinVersion: tls.VersionTLS12, - }) - if err != nil { - return err - } - if v := conn.ConnectionState().NegotiatedProtocol; v != acmeALPNProto { - return fmt.Errorf("CAServer: verifyALPNChallenge: negotiated proto is %q; want %q", v, acmeALPNProto) - } - if n := len(conn.ConnectionState().PeerCertificates); n != 1 { - return fmt.Errorf("len(PeerCertificates) = %d; want 1", n) - } - crt = conn.ConnectionState().PeerCertificates[0] - case haveGetCert: - hello := &tls.ClientHelloInfo{ - ServerName: a.domain, - // TODO: support selecting ECDSA. - CipherSuites: []uint16{tls.TLS_ECDHE_RSA_WITH_CHACHA20_POLY1305}, - SupportedProtos: []string{acme.ALPNProto}, - SupportedVersions: []uint16{tls.VersionTLS12}, - } - c, err := getCert(hello) - if err != nil { - return err - } - crt, err = x509.ParseCertificate(c.Certificate[0]) - if err != nil { - return err - } - } - - if err := crt.VerifyHostname(a.domain); err != nil { - return fmt.Errorf("verifyALPNChallenge: VerifyHostname: %v", err) - } - // See RFC 8737, Section 6.1. - oid := asn1.ObjectIdentifier{1, 3, 6, 1, 5, 5, 7, 1, 31} - for _, x := range crt.Extensions { - if x.Id.Equal(oid) { - // TODO: check the token. - return nil - } - } - return fmt.Errorf("verifyTokenCert: no id-pe-acmeIdentifier extension found") -} - -func (ca *CAServer) verifyDNSChallenge(a *authorization) error { - ca.mu.Lock() - defer ca.mu.Unlock() - - if _, ok := ca.dnsResponses[a.domain]; !ok { - return fmt.Errorf("verifyDNSChallenge: no DNS response registered for domain") - } - - return nil -} - -func (ca *CAServer) verifyHTTPChallenge(a *authorization) error { - addr, haveAddr := ca.addr(a.domain) - handler, haveHandler := ca.getHandler(a.domain) - if !haveAddr && !haveHandler { - return fmt.Errorf("no resolution information for %q", a.domain) - } - if haveAddr && haveHandler { - return fmt.Errorf("overlapping resolution information for %q", a.domain) - } - - token := challengeToken(a.domain, "http-01", a.id) - path := "/.well-known/acme-challenge/" + token - - var body string - switch { - case haveAddr: - t := &http.Transport{ - DialContext: func(ctx context.Context, network, _ string) (net.Conn, error) { - return (&net.Dialer{}).DialContext(ctx, network, addr) - }, - } - req, err := http.NewRequest("GET", "http://"+a.domain+path, nil) - if err != nil { - return err - } - res, err := t.RoundTrip(req) - if err != nil { - return err - } - if res.StatusCode != http.StatusOK { - return fmt.Errorf("http token: w.Code = %d; want %d", res.StatusCode, http.StatusOK) - } - b, err := io.ReadAll(res.Body) - if err != nil { - return err - } - body = string(b) - case haveHandler: - r := httptest.NewRequest("GET", path, nil) - r.Host = a.domain - w := httptest.NewRecorder() - handler.ServeHTTP(w, r) - if w.Code != http.StatusOK { - return fmt.Errorf("http token: w.Code = %d; want %d", w.Code, http.StatusOK) - } - body = w.Body.String() - } - - if !strings.HasPrefix(body, token) { - return fmt.Errorf("http token value = %q; want 'token-http-01.' prefix", body) - } - return nil -} - -func decodePayload(v interface{}, r io.Reader) error { - var req struct{ Payload string } - if err := json.NewDecoder(r).Decode(&req); err != nil { - return err - } - payload, err := base64.RawURLEncoding.DecodeString(req.Payload) - if err != nil { - return err - } - return json.Unmarshal(payload, v) -} - -func challengeToken(domain, challType string, authzID int) string { - return fmt.Sprintf("token-%s-%s-%d", domain, challType, authzID) -} - -func unique(a []string) []string { - seen := make(map[string]bool) - var res []string - for _, s := range a { - if s != "" && !seen[s] { - seen[s] = true - res = append(res, s) - } - } - return res -} diff --git a/internal/autocert/listener.go b/internal/autocert/listener.go deleted file mode 100644 index 9d62f8c..0000000 --- a/internal/autocert/listener.go +++ /dev/null @@ -1,155 +0,0 @@ -// Copyright 2017 The Go Authors. All rights reserved. -// Use of this source code is governed by a BSD-style -// license that can be found in the LICENSE file. - -package autocert - -import ( - "crypto/tls" - "log" - "net" - "os" - "path/filepath" - "runtime" - "time" -) - -// NewListener returns a net.Listener that listens on the standard TLS -// port (443) on all interfaces and returns *tls.Conn connections with -// LetsEncrypt certificates for the provided domain or domains. -// -// It enables one-line HTTPS servers: -// -// log.Fatal(http.Serve(autocert.NewListener("example.com"), handler)) -// -// NewListener is a convenience function for a common configuration. -// More complex or custom configurations can use the autocert.Manager -// type instead. -// -// Use of this function implies acceptance of the LetsEncrypt Terms of -// Service. If domains is not empty, the provided domains are passed -// to HostWhitelist. If domains is empty, the listener will do -// LetsEncrypt challenges for any requested domain, which is not -// recommended. -// -// Certificates are cached in a "golang-autocert" directory under an -// operating system-specific cache or temp directory. This may not -// be suitable for servers spanning multiple machines. -// -// The returned listener uses a *tls.Config that enables HTTP/2, and -// should only be used with servers that support HTTP/2. -// -// The returned Listener also enables TCP keep-alives on the accepted -// connections. The returned *tls.Conn are returned before their TLS -// handshake has completed. -func NewListener(domains ...string) net.Listener { - m := &Manager{ - Prompt: AcceptTOS, - } - if len(domains) > 0 { - m.HostPolicy = HostWhitelist(domains...) - } - dir := cacheDir() - if err := os.MkdirAll(dir, 0700); err != nil { - log.Printf("warning: autocert.NewListener not using a cache: %v", err) - } else { - m.Cache = DirCache(dir) - } - return m.Listener() -} - -// Listener listens on the standard TLS port (443) on all interfaces -// and returns a net.Listener returning *tls.Conn connections. -// -// The returned listener uses a *tls.Config that enables HTTP/2, and -// should only be used with servers that support HTTP/2. -// -// The returned Listener also enables TCP keep-alives on the accepted -// connections. The returned *tls.Conn are returned before their TLS -// handshake has completed. -// -// Unlike NewListener, it is the caller's responsibility to initialize -// the Manager m's Prompt, Cache, HostPolicy, and other desired options. -func (m *Manager) Listener() net.Listener { - ln := &listener{ - conf: m.TLSConfig(), - } - ln.tcpListener, ln.tcpListenErr = net.Listen("tcp", ":443") - return ln -} - -type listener struct { - conf *tls.Config - - tcpListener net.Listener - tcpListenErr error -} - -func (ln *listener) Accept() (net.Conn, error) { - if ln.tcpListenErr != nil { - return nil, ln.tcpListenErr - } - conn, err := ln.tcpListener.Accept() - if err != nil { - return nil, err - } - tcpConn := conn.(*net.TCPConn) - - // Because Listener is a convenience function, help out with - // this too. This is not possible for the caller to set once - // we return a *tcp.Conn wrapping an inaccessible net.Conn. - // If callers don't want this, they can do things the manual - // way and tweak as needed. But this is what net/http does - // itself, so copy that. If net/http changes, we can change - // here too. - tcpConn.SetKeepAlive(true) - tcpConn.SetKeepAlivePeriod(3 * time.Minute) - - return tls.Server(tcpConn, ln.conf), nil -} - -func (ln *listener) Addr() net.Addr { - if ln.tcpListener != nil { - return ln.tcpListener.Addr() - } - // net.Listen failed. Return something non-nil in case callers - // call Addr before Accept: - return &net.TCPAddr{IP: net.IP{0, 0, 0, 0}, Port: 443} -} - -func (ln *listener) Close() error { - if ln.tcpListenErr != nil { - return ln.tcpListenErr - } - return ln.tcpListener.Close() -} - -func homeDir() string { - if runtime.GOOS == "windows" { - return os.Getenv("HOMEDRIVE") + os.Getenv("HOMEPATH") - } - if h := os.Getenv("HOME"); h != "" { - return h - } - return "/" -} - -func cacheDir() string { - const base = "golang-autocert" - switch runtime.GOOS { - case "darwin": - return filepath.Join(homeDir(), "Library", "Caches", base) - case "windows": - for _, ev := range []string{"APPDATA", "CSIDL_APPDATA", "TEMP", "TMP"} { - if v := os.Getenv(ev); v != "" { - return filepath.Join(v, base) - } - } - // Worst case: - return filepath.Join(homeDir(), base) - } - if xdg := os.Getenv("XDG_CACHE_HOME"); xdg != "" { - return filepath.Join(xdg, base) - } - return filepath.Join(homeDir(), ".cache", base) -} diff --git a/internal/autocert/renewal.go b/internal/autocert/renewal.go deleted file mode 100644 index 0df7da7..0000000 --- a/internal/autocert/renewal.go +++ /dev/null @@ -1,156 +0,0 @@ -// Copyright 2016 The Go Authors. All rights reserved. -// Use of this source code is governed by a BSD-style -// license that can be found in the LICENSE file. - -package autocert - -import ( - "context" - "crypto" - "sync" - "time" -) - -// renewJitter is the maximum deviation from Manager.RenewBefore. -const renewJitter = time.Hour - -// domainRenewal tracks the state used by the periodic timers -// renewing a single domain's cert. -type domainRenewal struct { - m *Manager - ck certKey - key crypto.Signer - - timerMu sync.Mutex - timer *time.Timer - timerClose chan struct{} // if non-nil, renew closes this channel (and nils out the timer fields) instead of running -} - -// start starts a cert renewal timer at the time -// defined by the certificate expiration time exp. -// -// If the timer is already started, calling start is a noop. -func (dr *domainRenewal) start(exp time.Time) { - dr.timerMu.Lock() - defer dr.timerMu.Unlock() - if dr.timer != nil { - return - } - dr.timer = time.AfterFunc(dr.next(exp), dr.renew) -} - -// stop stops the cert renewal timer and waits for any in-flight calls to renew -// to complete. If the timer is already stopped, calling stop is a noop. -func (dr *domainRenewal) stop() { - dr.timerMu.Lock() - defer dr.timerMu.Unlock() - for { - if dr.timer == nil { - return - } - if dr.timer.Stop() { - dr.timer = nil - return - } else { - // dr.timer fired, and we acquired dr.timerMu before the renew callback did. - // (We know this because otherwise the renew callback would have reset dr.timer!) - timerClose := make(chan struct{}) - dr.timerClose = timerClose - dr.timerMu.Unlock() - <-timerClose - dr.timerMu.Lock() - } - } -} - -// renew is called periodically by a timer. -// The first renew call is kicked off by dr.start. -func (dr *domainRenewal) renew() { - dr.timerMu.Lock() - defer dr.timerMu.Unlock() - if dr.timerClose != nil { - close(dr.timerClose) - dr.timer, dr.timerClose = nil, nil - return - } - - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Minute) - defer cancel() - // TODO: rotate dr.key at some point? - next, err := dr.do(ctx) - if err != nil { - next = renewJitter / 2 - next += time.Duration(pseudoRand.int63n(int64(next))) - } - testDidRenewLoop(next, err) - dr.timer = time.AfterFunc(next, dr.renew) -} - -// updateState locks and replaces the relevant Manager.state item with the given -// state. It additionally updates dr.key with the given state's key. -func (dr *domainRenewal) updateState(state *certState) { - dr.m.stateMu.Lock() - defer dr.m.stateMu.Unlock() - dr.key = state.key - dr.m.state[dr.ck] = state -} - -// do is similar to Manager.createCert but it doesn't lock a Manager.state item. -// Instead, it requests a new certificate independently and, upon success, -// replaces dr.m.state item with a new one and updates cache for the given domain. -// -// It may lock and update the Manager.state if the expiration date of the currently -// cached cert is far enough in the future. -// -// The returned value is a time interval after which the renewal should occur again. -func (dr *domainRenewal) do(ctx context.Context) (time.Duration, error) { - // a race is likely unavoidable in a distributed environment - // but we try nonetheless - if tlscert, err := dr.m.cacheGet(ctx, dr.ck); err == nil { - next := dr.next(tlscert.Leaf.NotAfter) - if next > dr.m.renewBefore()+renewJitter { - signer, ok := tlscert.PrivateKey.(crypto.Signer) - if ok { - state := &certState{ - key: signer, - cert: tlscert.Certificate, - leaf: tlscert.Leaf, - } - dr.updateState(state) - return next, nil - } - } - } - - der, leaf, err := dr.m.authorizedCert(ctx, dr.key, dr.ck) - if err != nil { - return 0, err - } - state := &certState{ - key: dr.key, - cert: der, - leaf: leaf, - } - tlscert, err := state.tlscert() - if err != nil { - return 0, err - } - if err := dr.m.cachePut(ctx, dr.ck, tlscert); err != nil { - return 0, err - } - dr.updateState(state) - return dr.next(leaf.NotAfter), nil -} - -func (dr *domainRenewal) next(expiry time.Time) time.Duration { - d := expiry.Sub(dr.m.now()) - dr.m.renewBefore() - // add a bit of randomness to renew deadline - n := pseudoRand.int63n(int64(renewJitter)) - d -= time.Duration(n) - if d < 0 { - return 0 - } - return d -} - -var testDidRenewLoop = func(next time.Duration, err error) {} diff --git a/internal/autocert/renewal_test.go b/internal/autocert/renewal_test.go deleted file mode 100644 index 3b1062b..0000000 --- a/internal/autocert/renewal_test.go +++ /dev/null @@ -1,270 +0,0 @@ -// Copyright 2016 The Go Authors. All rights reserved. -// Use of this source code is governed by a BSD-style -// license that can be found in the LICENSE file. - -package autocert - -import ( - "context" - "crypto" - "crypto/ecdsa" - "testing" - "time" - - "github.com/sr/tsproxy/internal/autocert/internal/acmetest" - - "golang.org/x/crypto/acme" -) - -func TestRenewalNext(t *testing.T) { - now := time.Now() - man := &Manager{ - RenewBefore: 7 * 24 * time.Hour, - nowFunc: func() time.Time { return now }, - } - defer man.stopRenew() - tt := []struct { - expiry time.Time - min, max time.Duration - }{ - {now.Add(90 * 24 * time.Hour), 83*24*time.Hour - renewJitter, 83 * 24 * time.Hour}, - {now.Add(time.Hour), 0, 1}, - {now, 0, 1}, - {now.Add(-time.Hour), 0, 1}, - } - - dr := &domainRenewal{m: man} - for i, test := range tt { - next := dr.next(test.expiry) - if next < test.min || test.max < next { - t.Errorf("%d: next = %v; want between %v and %v", i, next, test.min, test.max) - } - } -} - -func TestRenewFromCache(t *testing.T) { - man := testManager(t) - man.RenewBefore = 24 * time.Hour - - ca := acmetest.NewCAServer(t).Start() - ca.ResolveGetCertificate(exampleDomain, man.GetCertificate) - - man.Client = &acme.Client{ - DirectoryURL: ca.URL(), - } - - // cache an almost expired cert - now := time.Now() - c := ca.LeafCert(exampleDomain, "ECDSA", now.Add(-2*time.Hour), now.Add(time.Minute)) - if err := man.cachePut(context.Background(), exampleCertKey, c); err != nil { - t.Fatal(err) - } - - // verify the renewal happened - defer func() { - // Stop the timers that read and execute testDidRenewLoop before restoring it. - // Otherwise the timer callback may race with the deferred write. - man.stopRenew() - testDidRenewLoop = func(next time.Duration, err error) {} - }() - renewed := make(chan bool, 1) - testDidRenewLoop = func(next time.Duration, err error) { - defer func() { - select { - case renewed <- true: - default: - // The renewal timer uses a random backoff. If the first renewal fails for - // some reason, we could end up with multiple calls here before the test - // stops the timer. - } - }() - - if err != nil { - t.Errorf("testDidRenewLoop: %v", err) - } - // Next should be about 90 days: - // CaServer creates 90days expiry + account for man.RenewBefore. - // Previous expiration was within 1 min. - future := 88 * 24 * time.Hour - if next < future { - t.Errorf("testDidRenewLoop: next = %v; want >= %v", next, future) - } - - // ensure the new cert is cached - after := time.Now().Add(future) - tlscert, err := man.cacheGet(context.Background(), exampleCertKey) - if err != nil { - t.Errorf("man.cacheGet: %v", err) - return - } - if !tlscert.Leaf.NotAfter.After(after) { - t.Errorf("cache leaf.NotAfter = %v; want > %v", tlscert.Leaf.NotAfter, after) - } - - // verify the old cert is also replaced in memory - man.stateMu.Lock() - defer man.stateMu.Unlock() - s := man.state[exampleCertKey] - if s == nil { - t.Errorf("m.state[%q] is nil", exampleCertKey) - return - } - tlscert, err = s.tlscert() - if err != nil { - t.Errorf("s.tlscert: %v", err) - return - } - if !tlscert.Leaf.NotAfter.After(after) { - t.Errorf("state leaf.NotAfter = %v; want > %v", tlscert.Leaf.NotAfter, after) - } - } - - // trigger renew - hello := clientHelloInfo(exampleDomain, algECDSA) - if _, err := man.GetCertificate(hello); err != nil { - t.Fatal(err) - } - <-renewed -} - -func TestRenewFromCacheAlreadyRenewed(t *testing.T) { - ca := acmetest.NewCAServer(t).Start() - man := testManager(t) - man.RenewBefore = 24 * time.Hour - man.Client = &acme.Client{ - DirectoryURL: "invalid", - } - - // cache a recently renewed cert with a different private key - now := time.Now() - newCert := ca.LeafCert(exampleDomain, "ECDSA", now.Add(-2*time.Hour), now.Add(time.Hour*24*90)) - if err := man.cachePut(context.Background(), exampleCertKey, newCert); err != nil { - t.Fatal(err) - } - newLeaf, err := validCert(exampleCertKey, newCert.Certificate, newCert.PrivateKey.(crypto.Signer), now) - if err != nil { - t.Fatal(err) - } - - // set internal state to an almost expired cert - oldCert := ca.LeafCert(exampleDomain, "ECDSA", now.Add(-2*time.Hour), now.Add(time.Minute)) - if err != nil { - t.Fatal(err) - } - oldLeaf, err := validCert(exampleCertKey, oldCert.Certificate, oldCert.PrivateKey.(crypto.Signer), now) - if err != nil { - t.Fatal(err) - } - man.stateMu.Lock() - if man.state == nil { - man.state = make(map[certKey]*certState) - } - s := &certState{ - key: oldCert.PrivateKey.(crypto.Signer), - cert: oldCert.Certificate, - leaf: oldLeaf, - } - man.state[exampleCertKey] = s - man.stateMu.Unlock() - - // verify the renewal accepted the newer cached cert - defer func() { - // Stop the timers that read and execute testDidRenewLoop before restoring it. - // Otherwise the timer callback may race with the deferred write. - man.stopRenew() - testDidRenewLoop = func(next time.Duration, err error) {} - }() - renewed := make(chan bool, 1) - testDidRenewLoop = func(next time.Duration, err error) { - defer func() { - select { - case renewed <- true: - default: - // The renewal timer uses a random backoff. If the first renewal fails for - // some reason, we could end up with multiple calls here before the test - // stops the timer. - } - }() - - if err != nil { - t.Errorf("testDidRenewLoop: %v", err) - } - // Next should be about 90 days - // Previous expiration was within 1 min. - future := 88 * 24 * time.Hour - if next < future { - t.Errorf("testDidRenewLoop: next = %v; want >= %v", next, future) - } - - // ensure the cached cert was not modified - tlscert, err := man.cacheGet(context.Background(), exampleCertKey) - if err != nil { - t.Errorf("man.cacheGet: %v", err) - return - } - if !tlscert.Leaf.NotAfter.Equal(newLeaf.NotAfter) { - t.Errorf("cache leaf.NotAfter = %v; want == %v", tlscert.Leaf.NotAfter, newLeaf.NotAfter) - } - - // verify the old cert is also replaced in memory - man.stateMu.Lock() - defer man.stateMu.Unlock() - s := man.state[exampleCertKey] - if s == nil { - t.Errorf("m.state[%q] is nil", exampleCertKey) - return - } - stateKey := s.key.Public().(*ecdsa.PublicKey) - if !stateKey.Equal(newLeaf.PublicKey) { - t.Error("state key was not updated from cache") - return - } - tlscert, err = s.tlscert() - if err != nil { - t.Errorf("s.tlscert: %v", err) - return - } - if !tlscert.Leaf.NotAfter.Equal(newLeaf.NotAfter) { - t.Errorf("state leaf.NotAfter = %v; want == %v", tlscert.Leaf.NotAfter, newLeaf.NotAfter) - } - } - - // assert the expiring cert is returned from state - hello := clientHelloInfo(exampleDomain, algECDSA) - tlscert, err := man.GetCertificate(hello) - if err != nil { - t.Fatal(err) - } - if !oldLeaf.NotAfter.Equal(tlscert.Leaf.NotAfter) { - t.Errorf("state leaf.NotAfter = %v; want == %v", tlscert.Leaf.NotAfter, oldLeaf.NotAfter) - } - - // trigger renew - man.startRenew(exampleCertKey, s.key, s.leaf.NotAfter) - <-renewed - func() { - man.renewalMu.Lock() - defer man.renewalMu.Unlock() - - // verify the private key is replaced in the renewal state - r := man.renewal[exampleCertKey] - if r == nil { - t.Errorf("m.renewal[%q] is nil", exampleCertKey) - return - } - renewalKey := r.key.Public().(*ecdsa.PublicKey) - if !renewalKey.Equal(newLeaf.PublicKey) { - t.Error("renewal private key was not updated from cache") - } - }() - - // assert the new cert is returned from state after renew - hello = clientHelloInfo(exampleDomain, algECDSA) - tlscert, err = man.GetCertificate(hello) - if err != nil { - t.Fatal(err) - } - if !newLeaf.NotAfter.Equal(tlscert.Leaf.NotAfter) { - t.Errorf("state leaf.NotAfter = %v; want == %v", tlscert.Leaf.NotAfter, newLeaf.NotAfter) - } -} diff --git a/main.go b/main.go index b319d4d..17f31f2 100644 --- a/main.go +++ b/main.go @@ -2,10 +2,10 @@ package main import ( "context" - "crypto/tls" "errors" "flag" "fmt" + "html/template" "net" "net/http" "net/url" @@ -16,25 +16,17 @@ import ( "syscall" "time" - "github.com/sr/tsproxy/internal/autocert" - - "github.com/cenkalti/backoff/v4" - "github.com/dnsimple/dnsimple-go/dnsimple" "github.com/oklog/run" "github.com/prometheus/client_golang/prometheus" "github.com/prometheus/client_golang/prometheus/promauto" "github.com/prometheus/client_golang/prometheus/promhttp" "golang.org/x/exp/slog" - "golang.org/x/oauth2" - "tailscale.com/ipn/ipnstate" + "tailscale.com/client/tailscale" "tailscale.com/tsnet" tslogger "tailscale.com/types/logger" ) const ( - // 5 minutes TTL. - dnsTTL = 5 * 60 - // keep this below systemd's DefaultTimeoutStopSec (90 seconds) stopTimeout = 80 * time.Second ) @@ -87,21 +79,23 @@ type upstream struct { prometheus bool } -func fqdn(zone, name string) string { - return name + "." + zone +type target struct { + name string + magicDNS string + prometheus bool } func parseUpstreamFlag(fval string) (upstream, error) { - kv := strings.Split(fval, "=") - if len(kv) != 2 { + k, v, ok := strings.Cut(fval, "=") + if !ok { return upstream{}, errors.New("format: name=http://backend") } - val := strings.Split(kv[1], ";") + val := strings.Split(v, ";") be, err := url.Parse(val[0]) if err != nil { return upstream{}, err } - up := upstream{name: kv[0], backend: be} + up := upstream{name: k, backend: be} if len(val) > 1 { for _, opt := range val[1:] { switch opt { @@ -124,27 +118,15 @@ func main() { func tsproxy(ctx context.Context) error { var ( - tok = flag.String("access-token", "", "DNSimple API Access Token. (Environment: DNSIMPLE_ACCESS_TOKEN)") - zone = flag.String("zone", "", "DNSimple Zone.") - email = flag.String("email", "", "Optional ACME registration email.") - state = flag.String("state", "", "Optional directory for storing Tailscale and autocert state.") - hostname = flag.String("hostname", os.Getenv("HOSTNAME"), "Tailscale machine name.") - tslog = flag.Bool("tailscale-logger", false, "If true, log Tailscale output.") + 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, "Port of the proxy's own HTTP server.") ) - var ups upstreamFlag - flag.Var(&ups, "upstream", "Repeated for each upstream. Format: name=http://backend:8000") + var upstreams upstreamFlag + flag.Var(&upstreams, "upstream", "Repeated for each upstream. Format: name=http://backend:8000") flag.Parse() - if v := os.Getenv("DNSIMPLE_ACCESS_TOKEN"); v != "" { - tok = &v - } - if *tok == "" { - return fmt.Errorf("required flag missing: access-token") - } - if *zone == "" { - return fmt.Errorf("required flag missing: zone") - } - if len(ups) == 0 { + if len(upstreams) == 0 { return fmt.Errorf("required flag missing: upstream") } if *state == "" { @@ -158,178 +140,129 @@ func tsproxy(ctx context.Context) error { } state = &dir } - if *hostname == "" { - if v, err := os.Hostname(); err == nil { - hostname = &v - } else { - return fmt.Errorf("required flag missing: hostname") - } - } logger := slog.New(slog.NewJSONHandler(os.Stderr)) - dnscli := dnsimple.NewClient(oauth2.NewClient(ctx, oauth2.StaticTokenSource(&oauth2.Token{AccessToken: *tok}))) - resp, err := dnscli.Identity.Whoami(ctx) + st, err := tsWaitStatusReady(ctx, &tailscale.LocalClient{}) if err != nil { - return fmt.Errorf("dnsimple: whoami request: %w", err) + return fmt.Errorf("tailscale: wait for node to be ready: %w", err) } - aid := strconv.FormatInt(resp.Data.Account.ID, 10) - // This is our DNS name. It will resolve to our tailscale IPs (A and AAAA records). - self := fqdn(*zone, strings.ToLower(*hostname)) + // service discovery targets (self + all upstreams) + targets := make([]target, len(upstreams)+1) - cache := filepath.Join(*state, "autocert") - acm := &autocert.Manager{ - Prompt: autocert.AcceptTOS, - Cache: autocert.DirCache(cache), - Email: *email, - DNS01: dnsimpleDNS01Solver(logger, dnscli.Zones, aid, *zone), - } - var names []string - for _, u := range ups { - names = append(names, fqdn(*zone, u.name)) - } - names = append(names, self) - acm.HostPolicy = autocert.HostWhitelist(names...) - - if err := os.MkdirAll(cache, 0700); err != nil { - return err - } - - ts := &tsnet.Server{ - Hostname: strings.ToLower(*hostname), - Dir: filepath.Join(*state, "tailscale"), - } - if *tslog { - ts.Logf = func(format string, args ...any) { - logger.LogAttrs(slog.InfoLevel, fmt.Sprintf(format, args...), slog.String("logger", "tailscale")) - } - } else { - ts.Logf = tslogger.Discard - } - if err := os.MkdirAll(ts.Dir, 0700); err != nil { - return err - } - lc, err := ts.LocalClient() - if err != nil { - return fmt.Errorf("tailscale: init local client: %w", err) - } - // ts.LocalClient() implicitly starts the server, make sure it gets closed. - defer ts.Close() + var g run.Group + ctx, cancel := context.WithCancel(ctx) + defer cancel() + g.Add(run.SignalHandler(ctx, os.Interrupt, syscall.SIGTERM)) - var st *ipnstate.Status - err = backoff.Retry(func() error { - if err := ctx.Err(); err != nil { + { + t, err := template.New("index.html").Parse(` + tsproxy + +

tsproxy

+

Metrics

+

Discovery

+

Upstreams:

+ + + `) + if err != nil { return err } - loopCtx, cancel := context.WithTimeout(ctx, time.Second) - st, err = lc.Status(loopCtx) - cancel() + ln, err := net.Listen("tcp", net.JoinHostPort(st.Self.TailscaleIPs[0].String(), strconv.Itoa(*port))) if err != nil { - return fmt.Errorf("get status: %w", err) + return fmt.Errorf("listen on %d: %w", *port, err) } + defer ln.Close() - if st.BackendState != "Running" { - return fmt.Errorf("backend not running: %s", st.BackendState) - } - if len(st.TailscaleIPs) != 2 { - return fmt.Errorf("IPs not yet assigned") + _, p, err := net.SplitHostPort(ln.Addr().String()) + if err != nil { + return err } - return nil - }, backoff.WithContext(backoff.NewExponentialBackOff(), ctx)) - if err != nil { - return fmt.Errorf("tailscale: wait for backend to be ready: %w", err) + + http.Handle("/metrics", promhttp.Handler()) + http.Handle("/sd", serveDiscovery(net.JoinHostPort(st.Self.DNSName, p), targets)) + http.Handle("/", serveIndex(t, targets)) + + srv := &http.Server{} + g.Add(func() error { + logger.Info("server ready", slog.String("addr", ln.Addr().String())) + + return srv.Serve(ln) + }, func(err error) { + if err := srv.Close(); err != nil { + logger.Error("shutdown server", err) + } + }) } - var g run.Group - ctx, cancel := context.WithCancel(ctx) - defer cancel() - g.Add(run.SignalHandler(ctx, os.Interrupt, syscall.SIGTERM)) - { - var ( - // SingleHostReverseProxy for each upstream. - rpx = make(map[string]http.Handler) - - // targets returned by the http_sd discovery endpoint. - targets []string - ) - for _, u := range ups { - fqdn := fqdn(*zone, u.name) - - rpx[fqdn] = tsSingleHostReverseProxy(logger, lc, u.backend) - if u.prometheus { - targets = append(targets, fqdn) + for i, upstream := range upstreams { + // https://go.dev/doc/faq#closures_and_goroutines + i := i + up := upstream + + log := logger.With(slog.String("upstream", up.name)) + + ts := &tsnet.Server{ + Hostname: up.name, + Dir: filepath.Join(*state, "tailscale-"+up.name), + } + if *tslog { + ts.Logf = func(format string, args ...any) { + log.LogAttrs(slog.InfoLevel, fmt.Sprintf(format, args...), slog.String("logger", "tailscale")) } + } else { + ts.Logf = tslogger.Discard + } + if err := os.MkdirAll(ts.Dir, 0700); err != nil { + return err + } + + lc, err := ts.LocalClient() + if err != nil { + return fmt.Errorf("tailscale: get local client for %s: %w", up.name, err) } - // Add self to service discovery. - targets = append(targets, self) srv := &http.Server{ - TLSConfig: &tls.Config{GetCertificate: acm.GetCertificate}, Handler: promhttp.InstrumentHandlerInFlight(requestsInFlight, promhttp.InstrumentHandlerDuration(duration, promhttp.InstrumentHandlerCounter(requests, - tsReverseProxy(rpx, promhttp.Handler(), targets, self)))), + newReverseProxy(log, lc, up.backend)))), } g.Add(func() error { - ln, err := ts.Listen("tcp", ":443") + defer ts.Close() + + ln, err := ts.Listen("tcp", ":80") if err != nil { - return fmt.Errorf("tailscale listen on :443: %w", err) + return fmt.Errorf("tailscale: listen for %s on port 80: %w", up.name, err) } defer ln.Close() - logger.Info("proxy server ready", slog.String("addr", ln.Addr().String())) - return srv.ServeTLS(ln, "", "") - }, func(err error) { - defer cancel() - logger.Info("shutting down proxy server") - shutdownCtx, cancel := context.WithTimeout(ctx, stopTimeout) - defer cancel() - if err := srv.Shutdown(shutdownCtx); err != nil { - logger.Error("proxy server shutdown", err) - } - }) - } - { - srv := &http.Server{ - Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - // Strip the port. - var host string - if h, _, err := net.SplitHostPort(r.Host); err != nil { - host = r.Host - } else { - host = h - } - http.Redirect(w, r, "https://"+host+r.URL.RequestURI(), http.StatusPermanentRedirect) - }), - } - g.Add(func() error { - ln, err := ts.Listen("tcp", ":80") + st, err := tsWaitStatusReady(ctx, lc) if err != nil { - return fmt.Errorf("tailscale listen on :80: %w", err) + return fmt.Errorf("tailscale: wait for node %s to be ready: %w", up.name, err) } - defer ln.Close() - logger.Info("HTTPS redirect server ready", slog.String("addr", ln.Addr().String())) + + // register in service discovery when we're ready. + targets[i] = target{name: up.name, prometheus: up.prometheus, magicDNS: st.Self.DNSName} + + log.Info("server ready", slog.String("addr", ln.Addr().String())) + return srv.Serve(ln) }, func(err error) { - defer cancel() - if err := srv.Close(); err != nil { - logger.Error("shutdown HTTPS redirect server", err) + log.Info("shutting down server") + + sctx, sc := context.WithTimeout(ctx, stopTimeout) + defer sc() + if err := srv.Shutdown(sctx); err != nil && err != http.ErrServerClosed { + log.Error("server shutdown", err) } }) - } - // Configure DNS in the background. - go func() { - start := time.Now() - if err := configureDNS(ctx, dnscli.Zones, net.DefaultResolver, aid, *zone, ups, st.TailscaleIPs, ts.Hostname); err != nil { - logger.Error("configure DNS", err) - } else { - logger.Info("DNS configured", slog.Duration("timer", time.Since(start))) - } - }() + } - return fmt.Errorf("server group exited: %w", g.Run()) + return g.Run() } diff --git a/main_test.go b/main_test.go deleted file mode 100644 index 3dd3439..0000000 --- a/main_test.go +++ /dev/null @@ -1,66 +0,0 @@ -package main - -import ( - "errors" - "net/url" - "reflect" - "strings" - "testing" - - "github.com/google/go-cmp/cmp" -) - -func TestParseUpstream(t *testing.T) { - 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;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 -} diff --git a/tsproxy.go b/tsproxy.go index c858114..8f6434d 100644 --- a/tsproxy.go +++ b/tsproxy.go @@ -4,20 +4,57 @@ import ( "context" "encoding/json" "errors" + "fmt" + "html/template" "net/http" "net/http/httputil" "net/url" "sort" + "strings" + "time" + "github.com/cenkalti/backoff/v4" "golang.org/x/exp/slog" + "tailscale.com/client/tailscale" "tailscale.com/client/tailscale/apitype" + "tailscale.com/ipn/ipnstate" ) type tailscaleLocalClient interface { WhoIs(context.Context, string) (*apitype.WhoIsResponse, error) } -func tsSingleHostReverseProxy(logger *slog.Logger, lc tailscaleLocalClient, url *url.URL) http.Handler { +func tsWaitStatusReady(ctx context.Context, lc *tailscale.LocalClient) (*ipnstate.Status, error) { + var st *ipnstate.Status + + err := backoff.Retry(func() error { + if err := ctx.Err(); err != nil { + return err + } + + loopCtx, cancel := context.WithTimeout(ctx, time.Second) + var err error + st, err = lc.Status(loopCtx) + cancel() + if err != nil { + return fmt.Errorf("get status: %w", err) + } + + if st.BackendState != "Running" { + return fmt.Errorf("backend not running: %s", st.BackendState) + } + if len(st.TailscaleIPs) != 2 { + return fmt.Errorf("IPs not yet assigned") + } + return nil + }, backoff.WithContext(backoff.NewExponentialBackOff(), ctx)) + if err != nil { + return nil, err + } + return st, nil +} + +func newReverseProxy(logger *slog.Logger, lc tailscaleLocalClient, url *url.URL) http.HandlerFunc { // TODO(sr) Instrument proxy.Transport proxy := httputil.NewSingleHostReverseProxy(url) orig := proxy.Director @@ -29,6 +66,7 @@ func tsSingleHostReverseProxy(logger *slog.Logger, lc tailscaleLocalClient, url logger.Error("tailscale whois", err) return } + // TODO(sr) No tags? if whois.UserProfile == nil { logger.Error("tailscale whois", errors.New("response did not include a user profile")) return @@ -46,41 +84,51 @@ func tsSingleHostReverseProxy(logger *slog.Logger, lc tailscaleLocalClient, url }) } -func tsReverseProxy(rpx map[string]http.Handler, metrics http.Handler, targets []string, self string) http.Handler { - sort.Strings(targets) +func serveDiscovery(self string, targets []target) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.TLS == nil { - http.Error(w, http.StatusText(http.StatusInternalServerError), http.StatusInternalServerError) + var tgs []string + tgs = append(tgs, self) + for _, t := range targets { + if t.magicDNS == "" { + continue + } + if !t.prometheus { + continue + } + tgs = append(tgs, t.magicDNS) + } + sort.Strings(tgs) + buf, err := json.Marshal([]struct { + Targets []string `json:"targets"` + }{ + {Targets: tgs}, + }) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) return } - if r.TLS.ServerName == self { - if r.RequestURI == "/sd" { - resp := []struct { - Targets []string `json:"targets"` - }{ - {Targets: targets}, - } - buf, err := json.Marshal(resp) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - w.Header().Set("Content-Type", "application/json; charset=utf-8") - _, _ = w.Write(buf) - return + w.Header().Set("Content-Type", "application/json; charset=utf-8") + _, _ = w.Write(buf) + }) +} + +func serveIndex(t *template.Template, targets []target) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var tgs []string + for _, t := range targets { + if t.magicDNS == "" { + continue } - if r.RequestURI == "/metrics" { - metrics.ServeHTTP(w, r) - return + h, _, ok := strings.Cut(t.magicDNS, ".") // strip the magicDNS suffix. + if !ok { + continue } - http.Error(w, http.StatusText(http.StatusNotFound), http.StatusNotFound) - return + tgs = append(tgs, h) } - px, ok := rpx[r.TLS.ServerName] - if !ok { - http.Error(w, http.StatusText(http.StatusNotFound), http.StatusNotFound) + sort.Strings(tgs) + if err := t.Execute(w, tgs); err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) return } - px.ServeHTTP(w, r) }) } diff --git a/tsproxy_test.go b/tsproxy_test.go index 7c9573d..ee5b048 100644 --- a/tsproxy_test.go +++ b/tsproxy_test.go @@ -1,9 +1,7 @@ package main import ( - "bytes" "context" - "crypto/tls" "errors" "fmt" "io" @@ -11,6 +9,8 @@ import ( "net/http" "net/http/httptest" "net/url" + "reflect" + "strings" "testing" "github.com/google/go-cmp/cmp" @@ -29,7 +29,64 @@ func (c *fakeLocalClient) WhoIs(ctx context.Context, remoteAddr string) (*apityp return c.whois(ctx, remoteAddr) } -func TestNewReverseProxy(t *testing.T) { +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;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() for _, tc := range []struct { @@ -79,7 +136,7 @@ func TestNewReverseProxy(t *testing.T) { if err != nil { log.Fatal(err) } - px := httptest.NewServer(tsSingleHostReverseProxy(slog.New(slog.NewTextHandler(io.Discard)), lc, beURL)) + px := httptest.NewServer(newReverseProxy(slog.New(slog.NewTextHandler(io.Discard)), lc, beURL)) defer px.Close() resp, err := http.Get(px.URL) @@ -104,94 +161,53 @@ func TestNewReverseProxy(t *testing.T) { } } -func TestNewTSProxyHandler(t *testing.T) { +func TestServeDiscovery(t *testing.T) { t.Parallel() - for _, tc := range []struct { - name string - h http.Handler - req *http.Request - want *http.Response - }{ - { - name: "no tls", - h: tsReverseProxy(nil, nil, nil, ""), - req: &http.Request{}, - want: &http.Response{StatusCode: http.StatusInternalServerError}, - }, - { - name: "upstream not found", - h: tsReverseProxy(nil, nil, nil, ""), - req: &http.Request{TLS: &tls.ConnectionState{ServerName: "example.com"}}, - want: &http.Response{StatusCode: http.StatusNotFound}, - }, - { - name: "upstream found", - h: tsReverseProxy(map[string]http.Handler{"example.com": http.RedirectHandler("http://redirect.net", http.StatusMovedPermanently)}, nil, nil, ""), - req: &http.Request{TLS: &tls.ConnectionState{ServerName: "example.com"}}, - want: &http.Response{StatusCode: http.StatusMovedPermanently, Header: http.Header{"Location": []string{"http://redirect.net"}}}, - }, - { - name: "self not found", - h: tsReverseProxy(map[string]http.Handler{"example.com": http.RedirectHandler("http://redirect.net", http.StatusMovedPermanently)}, nil, nil, "example.com"), - req: &http.Request{RequestURI: "/", TLS: &tls.ConnectionState{ServerName: "example.com"}}, - want: &http.Response{StatusCode: http.StatusNotFound}, - }, - { - name: "self metrics", - h: tsReverseProxy(nil, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { fmt.Fprintf(w, "metrics") }), nil, "example.com"), - req: &http.Request{RequestURI: "/metrics", TLS: &tls.ConnectionState{ServerName: "example.com"}}, - want: &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(bytes.NewReader([]byte("metrics")))}, - }, - { - name: "self service discovery", - h: tsReverseProxy(nil, nil, []string{"zzz:80", "localhost:8000"}, "example.com"), - req: &http.Request{RequestURI: "/sd", TLS: &tls.ConnectionState{ServerName: "example.com"}}, - want: &http.Response{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{`application/json; charset=utf-8`}}, Body: io.NopCloser(bytes.NewReader([]byte(`[{"targets":["localhost:8000","zzz:80"]}]`)))}, - }, - } { - tc := tc - t.Run(tc.name, func(t *testing.T) { - t.Parallel() - r := httptest.NewRecorder() - tc.h.ServeHTTP(r, tc.req) - resp := r.Result() - if want, got := tc.want.StatusCode, resp.StatusCode; want != got { - t.Errorf("want status %d, got: %d", want, got) - } - if len(tc.want.Header) > 0 { - if diff := cmp.Diff(tc.want.Header, resp.Header); diff != "" { - t.Errorf("headers mismatch (-want +got):\n%s", diff) - } - } - if tc.want.Body != nil { - want, err := io.ReadAll(tc.want.Body) - if err != nil { - t.Fatal(err) - } - got, err := io.ReadAll(resp.Body) - if err != nil { - t.Fatal(err) - } - if diff := cmp.Diff(string(want), string(got)); diff != "" { - t.Errorf("body mismatch (-want +got):\n%s", diff) - } - } - }) + ts := httptest.NewServer(serveDiscovery("self", []target{ + {magicDNS: "b", prometheus: true}, + {magicDNS: "x"}, + {}, + {magicDNS: "a", prometheus: true}, + })) + defer ts.Close() + + resp, err := http.Get(ts.URL) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + if want, got := http.StatusOK, resp.StatusCode; want != got { + t.Errorf("want status %d, got: %d", want, got) + } + b, err := io.ReadAll(resp.Body) + if err != nil { + t.Fatal(err) + } + if diff := cmp.Diff(`[{"targets":["a","b","self"]}]`, string(b)); diff != "" { + t.Errorf("body mismatch (-want +got):\n%s", diff) } } func TestMetrics(t *testing.T) { t.Parallel() + c, err := testutil.GatherAndCount(prometheus.DefaultGatherer) + if err != nil { + t.Fatalf("GatherAndCount: %v", err) + } + if c == 0 { + t.Fatalf("no metrics collected") + } + lint, err := testutil.GatherAndLint(prometheus.DefaultGatherer) if err != nil { t.Fatalf("CollectAndLint: %v", err) } + if len(lint) > 0 { + t.Error("lint problems detected") + } for _, prob := range lint { t.Errorf("lint: %s: %s", prob.Metric, prob.Text) } - if len(lint) > 0 { - t.Fatal("lint problems detected") - } } -- 2.51.2