diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..9bdcacd --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,64 @@ +name: ci + +on: + push: + # master is covered by the run that release.yml calls before it tags, so pushing + # there does not need a second one + branches-ignore: [master] + workflow_call: + +# A branch that is pushed again supersedes its own in-flight run, but master is left +# alone: its run gates the release. +concurrency: + group: ci-${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: ${{ github.ref != 'refs/heads/master' }} + +permissions: + contents: read + +jobs: + lint: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v7 + + - uses: actions/setup-go@v7 + with: + go-version-file: go.mod + + - name: Tool versions + id: tools + run: | + echo "golangci=$(make -s print-GOLANGCI_VERSION)" >> "$GITHUB_OUTPUT" + echo "actionlint=$(make -s print-ACTIONLINT_VERSION)" >> "$GITHUB_OUTPUT" + + - uses: golangci/golangci-lint-action@v9 + with: + version: ${{ steps.tools.outputs.golangci }} + + # the other half of `make lint`, so CI is never weaker than the local target + - name: Lint the workflows + env: + ACTIONLINT_VERSION: ${{ steps.tools.outputs.actionlint }} + run: | + go install "github.com/rhysd/actionlint/cmd/actionlint@${ACTIONLINT_VERSION}" + "$(go env GOPATH)/bin/actionlint" + + - name: go.mod and go.sum are tidy + run: | + go mod tidy + git diff --exit-code go.mod go.sum + + test: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v7 + + - uses: actions/setup-go@v7 + with: + go-version-file: go.mod + + - run: go build ./... + + # -race because the resolver cache and the registry heartbeat are concurrent + - run: go test -race -count=1 ./... diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml new file mode 100644 index 0000000..e91f7d8 --- /dev/null +++ b/.github/workflows/release.yml @@ -0,0 +1,68 @@ +name: release + +on: + push: + branches: [master] + +# Releases are serialised, never cancelled: two merges landing together must each get +# their turn at the tag rather than one clobbering the other's run. +concurrency: + group: release + cancel-in-progress: false + +permissions: + contents: read + +jobs: + # Nothing is tagged that has not passed the same lint and tests a PR does. + ci: + uses: ./.github/workflows/ci.yml + + tag: + needs: ci + runs-on: ubuntu-latest + permissions: + contents: write + steps: + - uses: actions/checkout@v7 + with: + fetch-depth: 0 # the whole tag history, to find the version to bump + + - name: Work out the next version + id: version + run: | + set -euo pipefail + + # A re-run of this workflow, or a commit tagged by hand, must not publish a + # second version of the same code. + if tag=$(git describe --exact-match --tags HEAD 2>/dev/null); then + echo "HEAD is already released as $tag, nothing to do" + exit 0 + fi + + latest=$(git tag -l 'v*' --sort=-v:refname | grep -E '^v[0-9]+\.[0-9]+\.[0-9]+$' | head -1 || true) + if [ -z "$latest" ]; then + next=v0.1.0 + else + IFS=. read -r major minor patch <<<"${latest#v}" + next="v${major}.${minor}.$((patch + 1))" + fi + + echo "next=$next" >> "$GITHUB_OUTPUT" + echo "Releasing $next (previous: ${latest:-none})" + + - name: Tag and release + if: steps.version.outputs.next != '' + env: + GH_TOKEN: ${{ github.token }} + NEXT: ${{ steps.version.outputs.next }} + run: | + set -euo pipefail + + git config user.name "github-actions[bot]" + git config user.email "41898282+github-actions[bot]@users.noreply.github.com" + + git tag -a "$NEXT" -m "$NEXT" + git push origin "$NEXT" + + gh release create "$NEXT" --title "$NEXT" --generate-notes diff --git a/.golangci.yml b/.golangci.yml new file mode 100644 index 0000000..d9bbe13 --- /dev/null +++ b/.golangci.yml @@ -0,0 +1,30 @@ +version: "2" + +linters: + default: standard # errcheck, govet, ineffassign, staticcheck, unused + enable: + - bodyclose # a leaked response body in an HTTP middleware leaks a connection + - errorlint + - misspell + - wrapcheck # errors crossing a package boundary say which package they came from + + settings: + wrapcheck: + extra-ignore-sigs: + # a middleware passing the next transport's error along is not the one that + # should be naming it + - .RoundTrip( + + exclusions: + rules: + # tests close what they open, and report errors to t rather than to a caller who + # would need to know which package they came from + - path: _test\.go + linters: + - bodyclose + - errcheck + - wrapcheck + +formatters: + enable: + - gofmt diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..7a9df1b --- /dev/null +++ b/Makefile @@ -0,0 +1,20 @@ +GOLANGCI_VERSION ?= v2.13.2 +ACTIONLINT_VERSION ?= v1.7.12 + +GOBIN ?= $(or $(shell go env GOBIN),$(shell go env GOPATH)/bin) + +.PHONY: test lint install-tools + +test: + go test -v ./... + +lint: + $(GOBIN)/golangci-lint run ./... + $(GOBIN)/actionlint + +install-tools: + go install github.com/golangci/golangci-lint/v2/cmd/golangci-lint@$(GOLANGCI_VERSION) + go install github.com/rhysd/actionlint/cmd/actionlint@$(ACTIONLINT_VERSION) + +print-%: + @echo $($*) diff --git a/caller/caller.go b/caller/caller.go new file mode 100644 index 0000000..189e83f --- /dev/null +++ b/caller/caller.go @@ -0,0 +1,50 @@ +package caller + +import ( + "fmt" + "net/http" + + "github.com/draincloud/callpack/caller/middleware" +) + +type Caller struct { + client http.Client + mws []middleware.RoundTripperHandler +} + +func New(client http.Client, mws ...middleware.RoundTripperHandler) *Caller { + return &Caller{ + client: client, + mws: mws, + } +} + +func (r *Caller) Use(middlewares ...middleware.RoundTripperHandler) { + r.mws = append(r.mws, middlewares...) +} + +func (r *Caller) With(middlewares ...middleware.RoundTripperHandler) *Caller { + combined := make([]middleware.RoundTripperHandler, 0, len(r.mws)+len(middlewares)) + combined = append(combined, r.mws...) + combined = append(combined, middlewares...) + return &Caller{client: r.client, mws: combined} +} + +func (r *Caller) Client() *http.Client { + cl := r.client + if cl.Transport == nil { + cl.Transport = http.DefaultTransport + } + for _, handler := range r.mws { + cl.Transport = handler(cl.Transport) + } + return &cl +} + +func (r *Caller) Do(req *http.Request) (*http.Response, error) { + resp, err := r.Client().Do(req) + if err != nil { + return nil, fmt.Errorf("caller: failed to execute request: %w", err) + } + return resp, nil +} diff --git a/caller/middleware/consul/consul.go b/caller/middleware/consul/consul.go new file mode 100644 index 0000000..046643d --- /dev/null +++ b/caller/middleware/consul/consul.go @@ -0,0 +1,104 @@ +package consul + +import ( + "context" + "fmt" + "net" + "net/http" + "strconv" + "sync" + "time" + + "github.com/draincloud/callpack/caller/middleware" + "github.com/hashicorp/consul/api" +) + +const DefaultTTL = time.Second + +func Resolve(client *api.Client, ttl time.Duration) middleware.RoundTripperHandler { + if ttl <= 0 { + ttl = DefaultTTL + } + c := &cache{health: client.Health(), ttl: ttl, services: map[string]*service{}} + + return func(next http.RoundTripper) http.RoundTripper { + return middleware.RoundTripperFunc(func(req *http.Request) (*http.Response, error) { + addr, err := c.address(req.Context(), req.URL.Hostname()) + if err != nil { + return nil, err + } + + out := req.Clone(req.Context()) + if out.Host == "" { + out.Host = req.URL.Host + } + out.URL.Host = addr + + return next.RoundTrip(out) + }) + } +} + +type cache struct { + health *api.Health + ttl time.Duration + + mu sync.Mutex + services map[string]*service +} + +type service struct { + mu sync.Mutex + addresses []string + next uint64 + expires time.Time +} + +func (c *cache) address(ctx context.Context, name string) (string, error) { + if name == "" { + return "", fmt.Errorf("consul: request has no host to resolve") + } + + c.mu.Lock() + s, ok := c.services[name] + if !ok { + s = &service{} + c.services[name] = s + } + c.mu.Unlock() + + s.mu.Lock() + defer s.mu.Unlock() + + if time.Now().After(s.expires) { + addresses, err := c.lookup(ctx, name) + if err != nil { + return "", err + } + s.addresses, s.expires, s.next = addresses, time.Now().Add(c.ttl), 0 + } + + address := s.addresses[s.next%uint64(len(s.addresses))] + s.next++ + return address, nil +} + +func (c *cache) lookup(ctx context.Context, name string) ([]string, error) { + entries, _, err := c.health.Service(name, "", true, (&api.QueryOptions{}).WithContext(ctx)) + if err != nil { + return nil, fmt.Errorf("consul: failed to resolve service %q: %w", name, err) + } + if len(entries) == 0 { + return nil, fmt.Errorf("consul: service %q has no healthy instances", name) + } + + addresses := make([]string, 0, len(entries)) + for _, entry := range entries { + host := entry.Service.Address + if host == "" { + host = entry.Node.Address + } + addresses = append(addresses, net.JoinHostPort(host, strconv.Itoa(entry.Service.Port))) + } + return addresses, nil +} diff --git a/caller/middleware/consul/consul_test.go b/caller/middleware/consul/consul_test.go new file mode 100644 index 0000000..2217a29 --- /dev/null +++ b/caller/middleware/consul/consul_test.go @@ -0,0 +1,301 @@ +package consul_test + +import ( + "encoding/json" + "io" + "net" + "net/http" + "net/http/httptest" + "net/url" + "strconv" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/draincloud/callpack/caller" + "github.com/draincloud/callpack/caller/middleware" + "github.com/draincloud/callpack/caller/middleware/consul" + "github.com/hashicorp/consul/api" +) + +type catalog struct { + mu sync.Mutex + instances map[string][]*api.ServiceEntry + lookups int64 +} + +func (c *catalog) ServeHTTP(w http.ResponseWriter, r *http.Request) { + atomic.AddInt64(&c.lookups, 1) + name := r.PathValue("name") + + c.mu.Lock() + entries := c.instances[name] + c.mu.Unlock() + if entries == nil { + entries = []*api.ServiceEntry{} + } + + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(entries) +} + +func (c *catalog) set(name string, entries ...*api.ServiceEntry) { + c.mu.Lock() + defer c.mu.Unlock() + c.instances[name] = entries +} + +func instance(nodeAddr, serviceAddr string, port int) *api.ServiceEntry { + return &api.ServiceEntry{ + Node: &api.Node{Node: "node", Address: nodeAddr}, + Service: &api.AgentService{Service: "api", Address: serviceAddr, Port: port}, + } +} + +type recorder struct { + mu sync.Mutex + seen []*http.Request + calls int64 +} + +func (t *recorder) RoundTrip(r *http.Request) (*http.Response, error) { + atomic.AddInt64(&t.calls, 1) + t.mu.Lock() + t.seen = append(t.seen, r) + t.mu.Unlock() + return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Request: r}, nil +} + +func (t *recorder) last() *http.Request { + t.mu.Lock() + defer t.mu.Unlock() + return t.seen[len(t.seen)-1] +} + +func setup(t *testing.T, ttl time.Duration) (*catalog, *recorder, http.RoundTripper) { + t.Helper() + + cat := &catalog{instances: map[string][]*api.ServiceEntry{}} + mux := http.NewServeMux() + mux.Handle("/v1/health/service/{name}", cat) + srv := httptest.NewServer(mux) + t.Cleanup(srv.Close) + + client, err := api.NewClient(&api.Config{Address: srv.URL}) + if err != nil { + t.Fatalf("api.NewClient: %v", err) + } + + rec := &recorder{} + return cat, rec, consul.Resolve(client, ttl)(rec) +} + +func request(t *testing.T, rawURL string) *http.Request { + t.Helper() + req, err := http.NewRequest(http.MethodGet, rawURL, nil) + if err != nil { + t.Fatalf("http.NewRequest: %v", err) + } + return req +} + +func TestResolveRewritesTargetAndKeepsServiceNameAsHost(t *testing.T) { + cat, rec, rt := setup(t, 0) + cat.set("api", instance("10.0.0.1", "10.0.0.11", 8080)) + + req := request(t, "http://api/v1/things?q=1") + if _, err := rt.RoundTrip(req); err != nil { + t.Fatalf("RoundTrip: %v", err) + } + + sent := rec.last() + if got, want := sent.URL.Host, "10.0.0.11:8080"; got != want { + t.Errorf("target host = %q, want %q", got, want) + } + if got, want := sent.Host, "api"; got != want { + t.Errorf("Host header = %q, want %q", got, want) + } + if got, want := sent.URL.RequestURI(), "/v1/things?q=1"; got != want { + t.Errorf("path = %q, want %q", got, want) + } + if got, want := req.URL.Host, "api"; got != want { + t.Errorf("caller's request was mutated: host = %q, want %q", got, want) + } +} + +func TestResolveFallsBackToNodeAddressAndIgnoresURLPort(t *testing.T) { + cat, rec, rt := setup(t, 0) + cat.set("api", instance("10.0.0.1", "", 9000)) // no service address: the node's is used + + if _, err := rt.RoundTrip(request(t, "http://api:1234/x")); err != nil { + t.Fatalf("RoundTrip: %v", err) + } + if got, want := rec.last().URL.Host, "10.0.0.1:9000"; got != want { + t.Errorf("target host = %q, want %q", got, want) + } +} + +func TestResolveCachesForTTL(t *testing.T) { + cat, _, rt := setup(t, 200*time.Millisecond) + cat.set("api", instance("10.0.0.1", "10.0.0.11", 8080)) + + for range 5 { + if _, err := rt.RoundTrip(request(t, "http://api/x")); err != nil { + t.Fatalf("RoundTrip: %v", err) + } + } + if got := atomic.LoadInt64(&cat.lookups); got != 1 { + t.Errorf("lookups within the TTL = %d, want 1", got) + } + + time.Sleep(250 * time.Millisecond) + if _, err := rt.RoundTrip(request(t, "http://api/x")); err != nil { + t.Fatalf("RoundTrip: %v", err) + } + if got := atomic.LoadInt64(&cat.lookups); got != 2 { + t.Errorf("lookups after the TTL = %d, want 2", got) + } +} + +func TestResolveCachesPerService(t *testing.T) { + cat, rec, rt := setup(t, time.Minute) + cat.set("api", instance("10.0.0.1", "10.0.0.11", 8080)) + cat.set("billing", instance("10.0.0.2", "10.0.0.22", 9090)) + + for _, name := range []string{"api", "billing", "api", "billing"} { + if _, err := rt.RoundTrip(request(t, "http://"+name+"/x")); err != nil { + t.Fatalf("RoundTrip: %v", err) + } + } + if got := atomic.LoadInt64(&cat.lookups); got != 2 { + t.Errorf("lookups = %d, want 2 (one per service)", got) + } + if got, want := rec.last().URL.Host, "10.0.0.22:9090"; got != want { + t.Errorf("billing resolved to %q, want %q", got, want) + } +} + +func TestResolveRoundRobinsAcrossInstances(t *testing.T) { + cat, rec, rt := setup(t, time.Minute) + cat.set("api", + instance("10.0.0.1", "10.0.0.11", 8080), + instance("10.0.0.2", "10.0.0.12", 8080), + instance("10.0.0.3", "10.0.0.13", 8080), + ) + + var got []string + for range 6 { + if _, err := rt.RoundTrip(request(t, "http://api/x")); err != nil { + t.Fatalf("RoundTrip: %v", err) + } + got = append(got, rec.last().URL.Host) + } + + want := []string{ + "10.0.0.11:8080", "10.0.0.12:8080", "10.0.0.13:8080", + "10.0.0.11:8080", "10.0.0.12:8080", "10.0.0.13:8080", + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("instance order = %v, want %v", got, want) + } + } +} + +func TestResolveFailsWithoutHealthyInstance(t *testing.T) { + cat, rec, rt := setup(t, time.Minute) + cat.set("api") // registered, nothing passing + + if _, err := rt.RoundTrip(request(t, "http://api/x")); err == nil { + t.Fatal("RoundTrip succeeded, want an error") + } + if got := atomic.LoadInt64(&rec.calls); got != 0 { + t.Errorf("transport calls = %d, want 0: the request must not be sent unresolved", got) + } + + cat.set("api", instance("10.0.0.1", "10.0.0.11", 8080)) + if _, err := rt.RoundTrip(request(t, "http://api/x")); err != nil { + t.Fatalf("RoundTrip after recovery: %v", err) + } + if got, want := rec.last().URL.Host, "10.0.0.11:8080"; got != want { + t.Errorf("target host = %q, want %q", got, want) + } +} + +func TestResolveRejectsRequestWithoutHost(t *testing.T) { + _, _, rt := setup(t, 0) + + req := &http.Request{Method: http.MethodGet, URL: &url.URL{Scheme: "http", Path: "/x"}, Header: http.Header{}} + if _, err := rt.RoundTrip(req); err == nil { + t.Fatal("RoundTrip succeeded, want an error") + } +} + +func TestResolveMakesOneLookupForConcurrentRequests(t *testing.T) { + cat, _, rt := setup(t, time.Minute) + cat.set("api", instance("10.0.0.1", "10.0.0.11", 8080)) + + var wg sync.WaitGroup + for range 20 { + wg.Add(1) + go func() { + defer wg.Done() + if _, err := rt.RoundTrip(request(t, "http://api/x")); err != nil { + t.Errorf("RoundTrip: %v", err) + } + }() + } + wg.Wait() + + if got := atomic.LoadInt64(&cat.lookups); got != 1 { + t.Errorf("lookups = %d, want 1: concurrent requests must share one lookup", got) + } +} + +// TestResolveThroughCaller drives the middleware the way a user does: through a Caller, over a real +// transport, to a backend registered in the catalog. +func TestResolveThroughCaller(t *testing.T) { + backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte("served " + r.Host + r.URL.Path)) + })) + t.Cleanup(backend.Close) + + host, port, err := net.SplitHostPort(strings.TrimPrefix(backend.URL, "http://")) + if err != nil { + t.Fatalf("SplitHostPort: %v", err) + } + backendPort, err := strconv.Atoi(port) + if err != nil { + t.Fatalf("Atoi: %v", err) + } + + cat := &catalog{instances: map[string][]*api.ServiceEntry{}} + mux := http.NewServeMux() + mux.Handle("/v1/health/service/{name}", cat) + consulSrv := httptest.NewServer(mux) + t.Cleanup(consulSrv.Close) + cat.set("api", instance(host, "", backendPort)) + + client, err := api.NewClient(&api.Config{Address: consulSrv.URL}) + if err != nil { + t.Fatalf("api.NewClient: %v", err) + } + + c := caller.New(http.Client{}, consul.Resolve(client, 0), middleware.Header("X-Trace", "abc")) + + resp, err := c.Do(request(t, "http://api/things")) + if err != nil { + t.Fatalf("Do: %v", err) + } + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + if err != nil { + t.Fatalf("ReadAll: %v", err) + } + if got, want := string(body), "served api/things"; got != want { + t.Errorf("body = %q, want %q (the backend must see the service name as its Host)", got, want) + } +} diff --git a/caller/middleware/headers.go b/caller/middleware/headers.go new file mode 100644 index 0000000..8af7fa2 --- /dev/null +++ b/caller/middleware/headers.go @@ -0,0 +1,106 @@ +package middleware + +import ( + "net/http" + "strings" +) + +// Header - adds headers to requests. +func Header(key, value string) RoundTripperHandler { + return headerHandler(key, value, false) +} + +// SecretHeader - adds a header carrying a credential to request. +func SecretHeader(key, value string) func(http.RoundTripper) http.RoundTripper { + return headerHandler(key, value, true) +} + +// JSON - sets Content-Type and Accept headers to json. +func JSON(next http.RoundTripper) http.RoundTripper { + fn := func(req *http.Request) (*http.Response, error) { + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json") + return next.RoundTrip(req) + } + return RoundTripperFunc(fn) +} + +// BasicAuth - adds basic auth to request. +func BasicAuth(user, passwd string) func(http.RoundTripper) http.RoundTripper { + return func(next http.RoundTripper) http.RoundTripper { + fn := func(req *http.Request) (*http.Response, error) { + if onOriginalHost(req) { + req.SetBasicAuth(user, passwd) + } + return roundTrip(next, req) + } + return RoundTripperFunc(fn) + } +} + +func headerHandler(key, value string, secret bool) func(http.RoundTripper) http.RoundTripper { + return func(next http.RoundTripper) http.RoundTripper { + fn := func(req *http.Request) (*http.Response, error) { + if secret && !onOriginalHost(req) { + if !credentialHeader(key) { + req.Header.Del(key) + } + return roundTrip(next, req) + } + req.Header.Set(key, value) + return roundTrip(next, req) + } + return RoundTripperFunc(fn) + } +} + +func roundTrip(next http.RoundTripper, req *http.Request) (*http.Response, error) { + resp, err := next.RoundTrip(req) + if resp != nil && resp.Request == nil { + resp.Request = req + } + return resp, err //nolint:wrapcheck +} + +func credentialHeader(key string) bool { + switch http.CanonicalHeaderKey(key) { + case "Authorization", "Www-Authenticate", "Cookie", "Cookie2", "Proxy-Authorization", "Proxy-Authenticate": + return true + } + return false +} + +func onOriginalHost(req *http.Request) bool { + if req.Response == nil { + return true + } + + origin := req + for origin.Response != nil { + if origin.Response.Request == nil { + return false + } + origin = origin.Response.Request + } + + originHost := strings.ToLower(origin.URL.Hostname()) + for r := req; r != origin; r = r.Response.Request { + if !domainOrSubdomain(strings.ToLower(r.URL.Hostname()), originHost) { + return false + } + } + return true +} + +func domainOrSubdomain(sub, parent string) bool { + if sub == parent { + return true + } + if strings.ContainsAny(sub, ":%") { + return false + } + if !strings.HasSuffix(sub, parent) { + return false + } + return sub[len(sub)-len(parent)-1] == '.' +} diff --git a/caller/middleware/logger.go b/caller/middleware/logger.go new file mode 100644 index 0000000..c870d7c --- /dev/null +++ b/caller/middleware/logger.go @@ -0,0 +1 @@ +package middleware diff --git a/caller/middleware/middlewares.go b/caller/middleware/middlewares.go new file mode 100644 index 0000000..bfa7321 --- /dev/null +++ b/caller/middleware/middlewares.go @@ -0,0 +1,12 @@ +package middleware + +import "net/http" + +// RoundTripperHandler is a type for middleware handler +type RoundTripperHandler func(http.RoundTripper) http.RoundTripper + +// RoundTripperFunc is a functional adapter for RoundTripperHandler +type RoundTripperFunc func(*http.Request) (*http.Response, error) + +// RoundTrip adopts function to the type +func (rt RoundTripperFunc) RoundTrip(r *http.Request) (*http.Response, error) { return rt(r) } diff --git a/go.mod b/go.mod index b20a2e5..426bb06 100644 --- a/go.mod +++ b/go.mod @@ -1,11 +1,12 @@ module github.com/draincloud/callpack -go 1.24.1 +go 1.25.0 + +require github.com/hashicorp/consul/api v1.32.1 require ( github.com/armon/go-metrics v0.4.1 // indirect github.com/fatih/color v1.16.0 // indirect - github.com/hashicorp/consul/api v1.32.1 // indirect github.com/hashicorp/errwrap v1.1.0 // indirect github.com/hashicorp/go-cleanhttp v0.5.2 // indirect github.com/hashicorp/go-hclog v1.5.0 // indirect diff --git a/go.sum b/go.sum index 2c9ca5f..8be927d 100644 --- a/go.sum +++ b/go.sum @@ -18,6 +18,8 @@ github.com/circonus-labs/circonus-gometrics v2.3.1+incompatible/go.mod h1:nmEj6D github.com/circonus-labs/circonusllhist v0.1.3/go.mod h1:kMXHVDlOchFAehlya5ePtbp5jckzBHf4XRpQvBOLI+I= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM= +github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/fatih/color v1.7.0/go.mod h1:Zm6kSWBoL9eyXnKyktHP6abPY2pDugNf5KwzbycvMj4= github.com/fatih/color v1.9.0/go.mod h1:eQcE1qtQxscV5RaZvpXrrb8Drkc3/DdQ+uUYCNjL+zU= github.com/fatih/color v1.13.0/go.mod h1:kLAiJbzzSOZDVNGyDpeOxJ47H46qBXwg5ILebYFFOfk= @@ -33,11 +35,17 @@ github.com/golang/protobuf v1.2.0/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5y github.com/golang/protobuf v1.3.1/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U= github.com/golang/protobuf v1.3.2/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U= github.com/google/btree v0.0.0-20180813153112-4030bb1f1f0c/go.mod h1:lNA+9X1NB3Zf8V7Ke586lFgjr2dZNuvo3lPJSGZ5JPQ= +github.com/google/btree v1.0.1 h1:gK4Kx5IaGY9CD5sPJ36FHiBJ6ZXl0kilRiiCj+jdYp4= +github.com/google/btree v1.0.1/go.mod h1:xXMiIv4Fb/0kKde4SpL7qlzvu5cMJDRkFDxJfI9uaxA= github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU= github.com/google/go-cmp v0.4.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= +github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= +github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= github.com/hashicorp/consul/api v1.32.1 h1:0+osr/3t/aZNAdJX558crU3PEjVrG4x6715aZHRgceE= github.com/hashicorp/consul/api v1.32.1/go.mod h1:mXUWLnxftwTmDv4W3lzxYCPD199iNLLUyLfLGFJbtl4= +github.com/hashicorp/consul/sdk v0.16.1 h1:V8TxTnImoPD5cj0U9Spl0TUxcytjcbbJeADFF07KdHg= +github.com/hashicorp/consul/sdk v0.16.1/go.mod h1:fSXvwxB2hmh1FMZCNl6PwX0Q/1wdWtHJcZ7Ea5tns0s= github.com/hashicorp/errwrap v1.0.0/go.mod h1:YH+1FKiLXxHSkmPseP+kNlulaMuP3n2brvKWEqk/Jc4= github.com/hashicorp/errwrap v1.1.0 h1:OxrOeh75EUXMY8TBjag2fzXGZ40LB6IKw45YeGUDY2I= github.com/hashicorp/errwrap v1.1.0/go.mod h1:YH+1FKiLXxHSkmPseP+kNlulaMuP3n2brvKWEqk/Jc4= @@ -50,6 +58,8 @@ github.com/hashicorp/go-immutable-radix v1.0.0/go.mod h1:0y9vanUI8NX6FsYoO3zeMjh github.com/hashicorp/go-immutable-radix v1.3.1 h1:DKHmCUm2hRBK510BaiZlwvpD40f8bJFeZnpfm2KLowc= github.com/hashicorp/go-immutable-radix v1.3.1/go.mod h1:0y9vanUI8NX6FsYoO3zeMjhV/C5i9g4Q3DwcSNZ4P60= github.com/hashicorp/go-msgpack v0.5.3/go.mod h1:ahLV/dePpqEmjfWmKiqvPkv/twdG7iPBM1vqhUKIvfM= +github.com/hashicorp/go-msgpack v0.5.5 h1:i9R9JSrqIz0QVLz3sz+i3YJdT7TTSLcfLLzJi9aZTuI= +github.com/hashicorp/go-msgpack v0.5.5/go.mod h1:ahLV/dePpqEmjfWmKiqvPkv/twdG7iPBM1vqhUKIvfM= github.com/hashicorp/go-multierror v1.0.0/go.mod h1:dHtQlpGsu+cZNNAkkCN/P3hoUDHhCYQXV3UM06sGGrk= github.com/hashicorp/go-multierror v1.1.0/go.mod h1:spPvp8C1qA32ftKqdAHm4hHTbPw+vmowP0z+KUhOZdA= github.com/hashicorp/go-multierror v1.1.1 h1:H5DkEtf6CXdFp0N0Em5UCwQpXMWke8IA0+lD48awMYo= @@ -58,14 +68,21 @@ github.com/hashicorp/go-retryablehttp v0.5.3/go.mod h1:9B5zBasrRhHXnJnui7y6sL7es github.com/hashicorp/go-rootcerts v1.0.2 h1:jzhAVGtqPKbwpyCPELlgNWhE1znq+qwJtW5Oi2viEzc= github.com/hashicorp/go-rootcerts v1.0.2/go.mod h1:pqUvnprVnM5bf7AOirdbb01K4ccR319Vf4pU3K5EGc8= github.com/hashicorp/go-sockaddr v1.0.0/go.mod h1:7Xibr9yA9JjQq1JpNB2Vw7kxv8xerXegt+ozgdvDeDU= +github.com/hashicorp/go-sockaddr v1.0.2 h1:ztczhD1jLxIRjVejw8gFomI1BQZOe2WoVOu0SyteCQc= +github.com/hashicorp/go-sockaddr v1.0.2/go.mod h1:rB4wwRAUzs07qva3c5SdrY/NEtAUjGlgmH/UkBUC97A= github.com/hashicorp/go-syslog v1.0.0/go.mod h1:qPfqrKkXGihmCqbJM2mZgkZGvKG1dFdvsLplgctolz4= github.com/hashicorp/go-uuid v1.0.0/go.mod h1:6SBZvOh/SIDV7/2o3Jml5SYk/TvGqwFJ/bN7x4byOro= github.com/hashicorp/go-uuid v1.0.1/go.mod h1:6SBZvOh/SIDV7/2o3Jml5SYk/TvGqwFJ/bN7x4byOro= +github.com/hashicorp/go-uuid v1.0.3 h1:2gKiV6YVmrJ1i2CKKa9obLvRieoRGviZFL26PcT/Co8= +github.com/hashicorp/go-uuid v1.0.3/go.mod h1:6SBZvOh/SIDV7/2o3Jml5SYk/TvGqwFJ/bN7x4byOro= +github.com/hashicorp/go-version v1.2.1 h1:zEfKbn2+PDgroKdiOzqiE8rsmLqU2uwi5PB5pBJ3TkI= +github.com/hashicorp/go-version v1.2.1/go.mod h1:fltr4n8CU8Ke44wwGCBoEymUuxUHl09ZGVZPK5anwXA= github.com/hashicorp/golang-lru v0.5.0/go.mod h1:/m3WP610KZHVQ1SGc6re/UDhFvYD7pJ4Ao+sR/qLZy8= github.com/hashicorp/golang-lru v0.5.4 h1:YDjusn29QI/Das2iO9M0BHnIbxPeyuCHsjMW+lJfyTc= github.com/hashicorp/golang-lru v0.5.4/go.mod h1:iADmTwqILo4mZ8BN3D2Q6+9jd8WM5uGBxy+E8yxSoD4= github.com/hashicorp/logutils v1.0.0/go.mod h1:QIAnNjmIWmVIIkWDTG1z5v++HQmx9WQRO+LraFDTW64= github.com/hashicorp/mdns v1.0.4/go.mod h1:mtBihi+LeNXGtG8L9dX59gAEa12BDtBQSp4v/YAJqrc= +github.com/hashicorp/memberlist v0.5.0 h1:EtYPN8DpAURiapus508I4n9CzHs2W+8NZGbmmR/prTM= github.com/hashicorp/memberlist v0.5.0/go.mod h1:yvyXLpo0QaGE59Y7hDTsTzDD25JYBZ4mHgHUZ8lrOI0= github.com/hashicorp/serf v0.10.1 h1:Z1H2J60yRKvfDYAOZLd2MU0ND4AH/WDz7xYHDWQsIPY= github.com/hashicorp/serf v0.10.1/go.mod h1:yL2t6BqATOLGc5HF7qbFkTfXoPIY0WZdWHfEvMqbG+4= @@ -94,6 +111,7 @@ github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWE github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/matttproud/golang_protobuf_extensions v1.0.1/go.mod h1:D8He9yQNgCq6Z5Ld7szi9bcBfOoFv/3dc6xSMkL2PC0= github.com/miekg/dns v1.1.26/go.mod h1:bPDLeHnStXmXAq1m/Ch/hvfNHr14JKNPMBo3VZKjuso= +github.com/miekg/dns v1.1.41 h1:WMszZWJG0XmzbK9FEmzH2TVcqYzFesusSIB41b8KHxY= github.com/miekg/dns v1.1.41/go.mod h1:p6aan82bvRIyn+zDIv9xYNUpwa73JcSh9BKwknJysuI= github.com/mitchellh/cli v1.1.0/go.mod h1:xcISNoH86gajksDmfB23e/pu+B+GeFRMYmoHXxx3xhI= github.com/mitchellh/go-homedir v1.1.0 h1:lukF9ziXFxDFPkA1vsr5zpc1XuPDn/wFntq5mG+4E0Y= @@ -107,10 +125,15 @@ github.com/modern-go/reflect2 v0.0.0-20180701023420-4b7aa43c6742/go.mod h1:bx2lN github.com/modern-go/reflect2 v1.0.1/go.mod h1:bx2lNnkwVCuqBIxFjflWJWanXIb3RllmbCylyMrvgv0= github.com/mwitkow/go-conntrack v0.0.0-20161129095857-cc309e4a2223/go.mod h1:qRWi+5nqEBWmkhHvq77mSJWrCKwh8bxhgT7d/eI7P4U= github.com/pascaldekloe/goe v0.0.0-20180627143212-57f6aae5913c/go.mod h1:lzWF7FIEvWOWxwDKqyGYQf6ZUaNfKdP144TG7ZOy1lc= +github.com/pascaldekloe/goe v0.1.0 h1:cBOtyMzM9HTpWjXfbbunk26uA6nG3a8n06Wieeh0MwY= github.com/pascaldekloe/goe v0.1.0/go.mod h1:lzWF7FIEvWOWxwDKqyGYQf6ZUaNfKdP144TG7ZOy1lc= github.com/pkg/errors v0.8.0/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= github.com/pkg/errors v0.8.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= +github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4= +github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U= +github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/posener/complete v1.1.1/go.mod h1:em0nMJCgc9GFtwrmVmEMR/ZL6WyhyjMBndrE9hABlRI= github.com/posener/complete v1.2.3/go.mod h1:WZIdtGGp+qx0sLrYKtIRAruyNpv6hFCicSgv7Sy7s/s= github.com/prometheus/client_golang v0.9.1/go.mod h1:7SWBe2y4D6OKWSNQJUaRYU/AaXPKyh/dDVn+NZz0KFw= @@ -125,15 +148,20 @@ github.com/prometheus/procfs v0.0.0-20181005140218-185b4288413d/go.mod h1:c3At6R github.com/prometheus/procfs v0.0.2/go.mod h1:TjEm7ze935MbeOT/UhFTIMYKhuLP4wbCsTZCD3I8kEA= github.com/prometheus/procfs v0.0.8/go.mod h1:7Qr8sr6344vo1JqZ6HhLceV9o3AJ1Ff+GxbHq6oeK9A= github.com/ryanuber/columnize v0.0.0-20160712163229-9b3edd62028f/go.mod h1:sm1tb6uqfes/u+d4ooFouqFdy9/2g9QGwK3SQygK0Ts= +github.com/sean-/seed v0.0.0-20170313163322-e2103e2c3529 h1:nn5Wsu0esKSJiIVhscUtVbo7ada43DJhG55ua/hjS5I= github.com/sean-/seed v0.0.0-20170313163322-e2103e2c3529/go.mod h1:DxrIzT+xaE7yg65j358z/aeFdxmN0P9QXhEzd20vsDc= github.com/sirupsen/logrus v1.2.0/go.mod h1:LxeOpSwHxABJmUn/MG1IvRgCAasNZTLOkJPxbbu5VWo= github.com/sirupsen/logrus v1.4.2/go.mod h1:tLMulIdttU9McNUspp0xgXVQah82FyeX6MwdIuYE2rE= 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.5.0 h1:1zr/of2m5FGMsad5YfcqgdqdWrIhu+EBEJRhR1U7z/c= +github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= 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.7.2/go.mod h1:R6va5+xMeoiuVRoj+gSkQ7d3FALtqAAGI1FQKckRals= +github.com/stretchr/testify v1.8.4 h1:CcVxjf3Q8PM0mHUKJCdn+eZZtm5yQwehR5yeSVQQcUk= +github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= github.com/tv42/httpunix v0.0.0-20150427012821-b75d8614f926/go.mod h1:9ESjWnEqriFuLhtthL60Sar/7RFoluCcXsuvEwTV5KM= golang.org/x/crypto v0.0.0-20180904163835-0709b304e793/go.mod h1:6SG95UA2DQfeDnfUPMdvaQW0Q7yPrPDi9nlGo2tz2b4= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= @@ -147,6 +175,8 @@ golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLL golang.org/x/net v0.0.0-20190923162816-aa69164e4478/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20210410081132-afb366fc7cd1/go.mod h1:9tjilg8BloeKEkVJvy7fQ90B1CfIiPueXVOjqfkSzI8= +golang.org/x/net v0.38.0 h1:vRMAPTMaeGqVhG5QyLJHqNDwecKTomGeqbnfZyKlBI8= +golang.org/x/net v0.38.0/go.mod h1:ivrbrMbzFq5J41QOQh0siUuly180yBYtLp+CKbEaFx8= golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= @@ -191,4 +221,5 @@ gopkg.in/yaml.v2 v2.2.1/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.2.4/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.2.5/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/registry/address.go b/registry/address.go new file mode 100644 index 0000000..97d2f8b --- /dev/null +++ b/registry/address.go @@ -0,0 +1,50 @@ +package registry + +import ( + "fmt" + "net" + "os" +) + +func localAddress() (string, error) { + if host, err := os.Hostname(); err == nil && host != "" { + if ips, err := net.LookupIP(host); err == nil { + if addr := routable(ips); addr != "" { + return addr, nil + } + } + } + + addrs, err := net.InterfaceAddrs() + if err != nil { + return "", fmt.Errorf("failed to read the local interfaces: %w", err) + } + + ips := make([]net.IP, 0, len(addrs)) + for _, addr := range addrs { + if ipnet, ok := addr.(*net.IPNet); ok { + ips = append(ips, ipnet.IP) + } + } + if addr := routable(ips); addr != "" { + return addr, nil + } + + return "", fmt.Errorf("failed to detect a routable address, set Service.Address") +} + +func routable(ips []net.IP) string { + var v6 string + for _, ip := range ips { + if !ip.IsGlobalUnicast() { + continue + } + if v4 := ip.To4(); v4 != nil { + return v4.String() + } + if v6 == "" { + v6 = ip.String() + } + } + return v6 +} diff --git a/registry/registry.go b/registry/registry.go new file mode 100644 index 0000000..04948dc --- /dev/null +++ b/registry/registry.go @@ -0,0 +1,187 @@ +package registry + +import ( + "context" + "fmt" + "log/slog" + "sync" + "time" + + "github.com/hashicorp/consul/api" +) + +const DefaultTTL = 15 * time.Second + +const deregisterAfter = time.Minute + +type Service struct { + Name string + + // Port is the port this instance serves on. + Port int + + // Address is where callers reach this instance. Detected from the process's own + // network when empty. + Address string + + Tags []string + Meta map[string]string + + // TTL bounds how long the catalog keeps routing to this instance after it stops + // heartbeating. + TTL time.Duration + + // Health, when set, is called before every heartbeat. A non-nil error marks this + // instance critical, which takes it out of discovery without deregistering it, so + // it comes back on its own once the error clears. + Health func(context.Context) error + + // Logger reports heartbeat failures, the only errors raised after Register returns. + // Defaults to slog.Default(). + Logger *slog.Logger +} + +// Registration is a live registration, kept alive by a background heartbeat until Close. +type Registration struct { + agent *api.Agent + service *api.AgentServiceRegistration + health func(context.Context) error + log *slog.Logger + ttl time.Duration + + cancel context.CancelFunc + done chan struct{} + closed sync.Once +} + +// Register adds this instance to the catalog and starts heartbeating for it. +func Register(ctx context.Context, client *api.Client, svc Service) (*Registration, error) { + if svc.Name == "" { + return nil, fmt.Errorf("registry: service name is required") + } + if svc.Port <= 0 { + return nil, fmt.Errorf("registry: service %q needs a port", svc.Name) + } + + address := svc.Address + if address == "" { + var err error + if address, err = localAddress(); err != nil { + return nil, fmt.Errorf("registry: service %q: %w", svc.Name, err) + } + } + + r := &Registration{ + agent: client.Agent(), + health: svc.Health, + log: svc.Logger, + ttl: svc.TTL, + done: make(chan struct{}), + } + if r.log == nil { + r.log = slog.Default() + } + if r.ttl <= 0 { + r.ttl = DefaultTTL + } + + id := fmt.Sprintf("%s-%s-%d", svc.Name, address, svc.Port) + r.service = &api.AgentServiceRegistration{ + ID: id, + Name: svc.Name, + Address: address, + Port: svc.Port, + Tags: svc.Tags, + Meta: svc.Meta, + Check: &api.AgentServiceCheck{ + CheckID: id + "-heartbeat", + Name: "heartbeat", + TTL: r.ttl.String(), + DeregisterCriticalServiceAfter: deregisterAfter.String(), + }, + } + if err := r.register(ctx); err != nil { + return nil, err + } + + loop, cancel := context.WithCancel(context.WithoutCancel(ctx)) + r.cancel = cancel + go r.heartbeat(loop) + + return r, nil +} + +// ID is the catalog ID of this instance, unique among the replicas of the service. +func (r *Registration) ID() string { return r.service.ID } + +// Close stops the heartbeat and removes this instance from the catalog. +func (r *Registration) Close() error { + var err error + r.closed.Do(func() { + r.cancel() + <-r.done + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + if derr := r.agent.ServiceDeregisterOpts(r.service.ID, query(ctx)); derr != nil { + err = fmt.Errorf("registry: failed to deregister %q: %w", r.service.ID, derr) + } + }) + return err +} + +func (r *Registration) heartbeat(ctx context.Context) { + defer close(r.done) + + ticker := time.NewTicker(r.ttl / 2) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + if err := r.beat(ctx); err != nil { + r.log.ErrorContext(ctx, "consul heartbeat failed", "service", r.service.Name, "id", r.service.ID, "error", err) + } + } + } +} + +func (r *Registration) beat(ctx context.Context) error { + status, output := r.probe(ctx) + err := r.agent.UpdateTTLOpts(r.service.Check.CheckID, output, status, query(ctx)) + if err == nil { + return nil + } + + if rerr := r.register(ctx); rerr != nil { + return fmt.Errorf("%w (re-registering: %w)", err, rerr) + } + return nil +} + +func (r *Registration) register(ctx context.Context) error { + r.service.Check.Status, r.service.Check.Notes = r.probe(ctx) + + opts := api.ServiceRegisterOpts{ReplaceExistingChecks: true}.WithContext(ctx) + if err := r.agent.ServiceRegisterOpts(r.service, opts); err != nil { + return fmt.Errorf("registry: failed to register %q: %w", r.service.ID, err) + } + return nil +} + +func (r *Registration) probe(ctx context.Context) (status, output string) { + if r.health == nil { + return api.HealthPassing, "alive" + } + if err := r.health(ctx); err != nil { + return api.HealthCritical, err.Error() + } + return api.HealthPassing, "healthy" +} + +func query(ctx context.Context) *api.QueryOptions { + return (&api.QueryOptions{}).WithContext(ctx) +} diff --git a/registry/registry_test.go b/registry/registry_test.go new file mode 100644 index 0000000..e5593cd --- /dev/null +++ b/registry/registry_test.go @@ -0,0 +1,415 @@ +package registry_test + +import ( + "context" + "encoding/json" + "fmt" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "github.com/draincloud/callpack/caller" + "github.com/draincloud/callpack/caller/middleware/consul" + "github.com/draincloud/callpack/registry" + "github.com/hashicorp/consul/api" +) + +type agent struct { + mu sync.Mutex + services map[string]*api.AgentServiceRegistration + checks map[string]string + forgotten bool + registers int + updates int +} + +func newAgent(t *testing.T) (*agent, *api.Client) { + t.Helper() + + a := &agent{ + services: map[string]*api.AgentServiceRegistration{}, + checks: map[string]string{}, + } + + mux := http.NewServeMux() + mux.HandleFunc("PUT /v1/agent/service/register", a.register) + mux.HandleFunc("PUT /v1/agent/service/deregister/{id}", a.deregister) + mux.HandleFunc("PUT /v1/agent/check/update/{id}", a.update) + mux.HandleFunc("GET /v1/health/service/{name}", a.health) + + srv := httptest.NewServer(mux) + t.Cleanup(srv.Close) + + config := api.DefaultConfig() + config.Address = strings.TrimPrefix(srv.URL, "http://") + client, err := api.NewClient(config) + if err != nil { + t.Fatalf("consul client: %v", err) + } + return a, client +} + +func (a *agent) register(w http.ResponseWriter, r *http.Request) { + var svc api.AgentServiceRegistration + if err := json.NewDecoder(r.Body).Decode(&svc); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + + a.mu.Lock() + defer a.mu.Unlock() + a.registers++ + a.forgotten = false + a.services[svc.ID] = &svc + a.checks[svc.Check.CheckID] = svc.Check.Status +} + +func (a *agent) deregister(_ http.ResponseWriter, r *http.Request) { + a.mu.Lock() + defer a.mu.Unlock() + delete(a.services, r.PathValue("id")) +} + +func (a *agent) update(w http.ResponseWriter, r *http.Request) { + var update struct{ Status, Output string } + if err := json.NewDecoder(r.Body).Decode(&update); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + + a.mu.Lock() + defer a.mu.Unlock() + + id := r.PathValue("id") + if _, ok := a.checks[id]; !ok || a.forgotten { + // what a real agent answers for a check it does not know about + http.Error(w, fmt.Sprintf("CheckID %q does not have associated TTL", id), http.StatusInternalServerError) + return + } + a.updates++ + a.checks[id] = update.Status +} + +func (a *agent) health(w http.ResponseWriter, r *http.Request) { + name := r.PathValue("name") + passingOnly := r.URL.Query().Get("passing") != "" + + a.mu.Lock() + entries := []*api.ServiceEntry{} + for _, svc := range a.services { + if svc.Name != name || (passingOnly && a.checks[svc.Check.CheckID] != api.HealthPassing) { + continue + } + entries = append(entries, &api.ServiceEntry{ + Node: &api.Node{Node: "node", Address: "10.0.0.1"}, + Service: &api.AgentService{ID: svc.ID, Service: svc.Name, Address: svc.Address, Port: svc.Port}, + }) + } + a.mu.Unlock() + + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(entries) +} + +func (a *agent) service(id string) *api.AgentServiceRegistration { + a.mu.Lock() + defer a.mu.Unlock() + return a.services[id] +} + +func (a *agent) status(checkID string) string { + a.mu.Lock() + defer a.mu.Unlock() + return a.checks[checkID] +} + +func (a *agent) counts() (registers, updates int) { + a.mu.Lock() + defer a.mu.Unlock() + return a.registers, a.updates +} + +func (a *agent) forget() { + a.mu.Lock() + defer a.mu.Unlock() + a.forgotten = true +} + +func quiet() *slog.Logger { + return slog.New(slog.NewTextHandler(io.Discard, nil)) +} + +func register(t *testing.T, client *api.Client, svc registry.Service) *registry.Registration { + t.Helper() + + if svc.Logger == nil { + svc.Logger = quiet() + } + reg, err := registry.Register(t.Context(), client, svc) + if err != nil { + t.Fatalf("register %q: %v", svc.Name, err) + } + t.Cleanup(func() { _ = reg.Close() }) + return reg +} + +func waitFor(t *testing.T, what string, cond func() bool) { + t.Helper() + + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + if cond() { + return + } + time.Sleep(2 * time.Millisecond) + } + t.Fatalf("timed out waiting for %s", what) +} + +func TestRegisterAdvertisesInstanceWithATTLCheck(t *testing.T) { + a, client := newAgent(t) + + reg := register(t, client, registry.Service{ + Name: "api", + Address: "10.1.2.3", + Port: 8080, + Tags: []string{"v2"}, + Meta: map[string]string{"commit": "abc"}, + }) + + svc := a.service(reg.ID()) + if svc == nil { + t.Fatalf("service %q was not registered, agent has %v", reg.ID(), a.services) + } + if svc.Name != "api" || svc.Address != "10.1.2.3" || svc.Port != 8080 { + t.Errorf("registered %q %s:%d, want api 10.1.2.3:8080", svc.Name, svc.Address, svc.Port) + } + if len(svc.Tags) != 1 || svc.Tags[0] != "v2" || svc.Meta["commit"] != "abc" { + t.Errorf("tags %v and meta %v were not passed through", svc.Tags, svc.Meta) + } + if svc.Check.TTL != registry.DefaultTTL.String() { + t.Errorf("check TTL is %q, want the default %q", svc.Check.TTL, registry.DefaultTTL) + } + if svc.Check.DeregisterCriticalServiceAfter == "" { + t.Error("a killed replica would linger: DeregisterCriticalServiceAfter is unset") + } + if svc.Check.Status != api.HealthPassing { + t.Errorf("check registered as %q, want %q so the instance is discoverable at once", svc.Check.Status, api.HealthPassing) + } +} + +func TestRegisterDetectsTheLocalAddress(t *testing.T) { + a, client := newAgent(t) + + reg := register(t, client, registry.Service{Name: "api", Port: 8080}) + + svc := a.service(reg.ID()) + if svc.Address == "" { + t.Fatal("no address was detected, callers would fall back to the agent's node") + } + if strings.HasPrefix(svc.Address, "127.") { + t.Errorf("detected %q, which no other host can reach", svc.Address) + } +} + +func TestRegisterRejectsAnIncompleteService(t *testing.T) { + _, client := newAgent(t) + + for _, svc := range []registry.Service{{Port: 8080}, {Name: "api"}} { + if _, err := registry.Register(t.Context(), client, svc); err == nil { + t.Errorf("registering %+v succeeded, want an error", svc) + } + } +} + +func TestReplicasOfAServiceRegisterSeparately(t *testing.T) { + a, client := newAgent(t) + + one := register(t, client, registry.Service{Name: "api", Address: "10.1.2.3", Port: 8080}) + two := register(t, client, registry.Service{Name: "api", Address: "10.1.2.4", Port: 8080}) + + if one.ID() == two.ID() { + t.Fatalf("both replicas registered as %q, so one replaced the other", one.ID()) + } + if a.service(one.ID()) == nil || a.service(two.ID()) == nil { + t.Fatalf("only one replica is in the catalog: %v", a.services) + } + + // a replica that restarts on the same address must replace itself, not pile up + restarted := register(t, client, registry.Service{Name: "api", Address: "10.1.2.3", Port: 8080}) + if restarted.ID() != one.ID() { + t.Errorf("restarted replica registered as %q, want %q", restarted.ID(), one.ID()) + } +} + +func TestHeartbeatKeepsTheCheckPassing(t *testing.T) { + a, client := newAgent(t) + + reg := register(t, client, registry.Service{ + Name: "api", Address: "10.1.2.3", Port: 8080, TTL: 40 * time.Millisecond, + }) + + waitFor(t, "heartbeats", func() bool { + _, updates := a.counts() + return updates >= 3 + }) + if got := a.status(reg.ID() + "-heartbeat"); got != api.HealthPassing { + t.Errorf("check is %q after heartbeating, want %q", got, api.HealthPassing) + } +} + +func TestHeartbeatReportsAnUnhealthyInstance(t *testing.T) { + a, client := newAgent(t) + + var failing bool + var mu sync.Mutex + health := func(context.Context) error { + mu.Lock() + defer mu.Unlock() + if failing { + return fmt.Errorf("database is unreachable") + } + return nil + } + + reg := register(t, client, registry.Service{ + Name: "api", Address: "10.1.2.3", Port: 8080, TTL: 40 * time.Millisecond, Health: health, + }) + check := reg.ID() + "-heartbeat" + + mu.Lock() + failing = true + mu.Unlock() + + waitFor(t, "the check to go critical", func() bool { + return a.status(check) == api.HealthCritical + }) + + // an unhealthy instance stays registered, so it recovers on its own + if a.service(reg.ID()) == nil { + t.Fatal("the instance was deregistered instead of marked critical") + } + mu.Lock() + failing = false + mu.Unlock() + + waitFor(t, "the check to recover", func() bool { + return a.status(check) == api.HealthPassing + }) +} + +func TestHeartbeatRegistersAgainWhenTheAgentForgot(t *testing.T) { + a, client := newAgent(t) + + reg := register(t, client, registry.Service{ + Name: "api", Address: "10.1.2.3", Port: 8080, TTL: 40 * time.Millisecond, + }) + + registers, _ := a.counts() + a.forget() + + waitFor(t, "the instance to register again", func() bool { + again, _ := a.counts() + return again > registers + }) + if a.service(reg.ID()) == nil { + t.Fatal("the instance is not back in the catalog") + } +} + +func TestCloseTakesTheInstanceOutOfTheCatalog(t *testing.T) { + a, client := newAgent(t) + + reg, err := registry.Register(t.Context(), client, registry.Service{ + Name: "api", Address: "10.1.2.3", Port: 8080, Logger: quiet(), + }) + if err != nil { + t.Fatalf("register: %v", err) + } + + if err := reg.Close(); err != nil { + t.Fatalf("close: %v", err) + } + if a.service(reg.ID()) != nil { + t.Error("the instance is still in the catalog after Close") + } + if err := reg.Close(); err != nil { + t.Errorf("second close: %v, want it to be a no-op", err) + } +} + +func TestRegisteredServiceIsDiscoverableThroughTheMiddleware(t *testing.T) { + _, client := newAgent(t) + + replicas := map[string]bool{} + var mu sync.Mutex + for _, name := range []string{"one", "two"} { + backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + mu.Lock() + replicas[name] = true + mu.Unlock() + fmt.Fprintf(w, "served by %s, host %s", name, r.Host) + })) + t.Cleanup(backend.Close) + + host, port := split(t, backend.URL) + register(t, client, registry.Service{Name: "api", Address: host, Port: port}) + } + + call := caller.New(http.Client{}, consul.Resolve(client, 0)) + for i := range 2 { + req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, "http://api/things", nil) + if err != nil { + t.Fatalf("request: %v", err) + } + resp, err := call.Do(req) + if err != nil { + t.Fatalf("request %d: %v", i, err) + } + body, _ := io.ReadAll(resp.Body) + resp.Body.Close() + + if !strings.Contains(string(body), "host api") { + t.Errorf("backend saw %q, want the service name as the Host header", body) + } + } + + mu.Lock() + defer mu.Unlock() + if len(replicas) != 2 { + t.Errorf("only %v were reached, want both replicas load balanced across", replicas) + } +} +func TestUnhealthyServiceIsNotDiscoverable(t *testing.T) { + _, client := newAgent(t) + + register(t, client, registry.Service{ + Name: "api", Address: "10.1.2.3", Port: 8080, TTL: 40 * time.Millisecond, + Health: func(context.Context) error { return fmt.Errorf("not ready") }, + }) + + call := caller.New(http.Client{}, consul.Resolve(client, time.Millisecond)) + req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, "http://api/things", nil) + if err != nil { + t.Fatalf("request: %v", err) + } + if _, err := call.Do(req); err == nil { + t.Fatal("the unhealthy replica was routed to") + } +} + +func split(t *testing.T, rawURL string) (host string, port int) { + t.Helper() + + var err error + host, rawPort, _ := strings.Cut(strings.TrimPrefix(rawURL, "http://"), ":") + if _, err = fmt.Sscanf(rawPort, "%d", &port); err != nil { + t.Fatalf("port of %q: %v", rawURL, err) + } + return host, port +}