diff --git a/cmd/sam-node/main.go b/cmd/sam-node/main.go index ce988382..3ed2b985 100644 --- a/cmd/sam-node/main.go +++ b/cmd/sam-node/main.go @@ -223,6 +223,20 @@ func normalizeControlPlaneURL(url string) string { return strings.TrimSuffix(url, "/") } +// parseRouterAddrs keeps the stored router addresses that still parse. +func parseRouterAddrs(addrs []string) []multiaddr.Multiaddr { + out := make([]multiaddr.Multiaddr, 0, len(addrs)) + for _, s := range addrs { + ma, err := multiaddr.NewMultiaddr(s) + if err != nil { + logger.Warnf("Ignoring stored router address %q: %v", s, err) + continue + } + out = append(out, ma) + } + return out +} + // interactiveJoin discovers a control plane's OIDC settings and completes an // interactive browser/device-code login against it. Shared by "join" and // "run --join"; targetControlPlane must already be a normalized URL. store is @@ -332,13 +346,15 @@ func main() { var controlPlanePubKey ed25519.PublicKey var routerAddrs []multiaddr.Multiaddr - storedPubKey, syncedAddrs, bannedPeerIDs, err := node.SyncMeshConfig(context.Background(), store) + // What the last run left behind; the node itself pulls what the + // control plane knows now before it starts (SyncControlPlane). + storedPubKey, storedAddrs, err := store.LoadMeshConfig() if err != nil { - logger.Warnf("Failed to sync mesh config: %v", err) + logger.Warnf("Failed to load stored mesh config: %v", err) } if len(storedPubKey) > 0 { controlPlanePubKey = storedPubKey - routerAddrs = syncedAddrs + routerAddrs = parseRouterAddrs(storedAddrs) } if controlPlanePublicKeyFlag != "" { @@ -457,7 +473,6 @@ func main() { ControlPlanePubKey: controlPlanePubKey, RouterAddrs: routerAddrs, Store: store, - BannedPeerIDs: bannedPeerIDs, MeshID: meshFlag, DiscoveryInterval: discoveryIntervalFlag, ListenAddrs: listenAddrs, @@ -483,6 +498,9 @@ func main() { if err != nil { logger.Fatalf("Failed to initialize mesh node: %v", err) } + if err := meshNode.SyncControlPlane(ctx); err != nil { + logger.Warnf("Control plane sync before start failed (using stored config): %v", err) + } if err := meshNode.Start(ctx); err != nil { logger.Fatalf("Failed to start mesh node: %v", err) } @@ -525,7 +543,6 @@ func main() { PrivKey: priv, RouterAddrs: initRouterAddrs, Store: store, - BannedPeerIDs: bannedPeerIDs, MeshID: meshFlag, DiscoveryInterval: discoveryIntervalFlag, ListenAddrs: listenAddrs, @@ -583,20 +600,20 @@ func main() { } enrollCancel() - storedPubKey, newRouterAddrs, postEnrollBannedPeerIDs, err := node.SyncMeshConfig(context.Background(), store) + // Enrollment stored the control plane key and the router + // addresses it answered with; the node built from them pulls + // the rest before it starts. + controlPlanePubKey, storedAddrs, err = store.LoadMeshConfig() if err != nil { - logger.Warnf("Failed to sync mesh config post-enrollment: %v", err) + logger.Fatalf("Failed to load mesh config after enrollment: %v", err) } - controlPlanePubKey = storedPubKey - bannedPeerIDs = postEnrollBannedPeerIDs logger.Debugf("listenAddrs: %v, allowLoopback: %v", listenAddrs, allowLoopbackFlag) meshNode, err = node.NewSamNode(node.Options{ PrivKey: priv, ControlPlanePubKey: controlPlanePubKey, - RouterAddrs: newRouterAddrs, + RouterAddrs: parseRouterAddrs(storedAddrs), Store: store, - BannedPeerIDs: bannedPeerIDs, MeshID: meshFlag, DiscoveryInterval: discoveryIntervalFlag, ListenAddrs: listenAddrs, @@ -618,6 +635,9 @@ func main() { if err != nil { logger.Fatalf("Failed to initialize node after enrollment: %v", err) } + if err := meshNode.SyncControlPlane(ctx); err != nil { + logger.Warnf("Control plane sync after enrollment failed (using enrollment config): %v", err) + } if err := meshNode.Start(ctx); err != nil { logger.Fatalf("Failed to start node after enrollment: %v", err) } diff --git a/internal/controlplane/client/client.go b/internal/controlplane/client/client.go new file mode 100644 index 00000000..6f61260b --- /dev/null +++ b/internal/controlplane/client/client.go @@ -0,0 +1,164 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Package client is how a mesh component reads from its control plane. Node +// and router share it, so the body cap, the status handling and the signature +// check on /keys are a single code path. It depends on api/ only: importing it +// pulls in none of the control plane server. +package client + +import ( + "context" + "crypto/ed25519" + "encoding/base64" + "errors" + "fmt" + "io" + "net/http" + "strings" + "time" + + "google.golang.org/protobuf/proto" + + "github.com/google/sam/api" +) + +// MaxBodyBytes caps every response body read from a control plane: a +// misbehaving or impersonated server must not be able to make a client +// buffer arbitrary amounts of memory. It is sized for the largest legitimate +// answer, the ban set in /info at roughly 55 bytes per peer ID, so about +// 150k banned peers fit; the policy is bounded by the control plane's own +// 1 MiB cap on POST /policies, and /keys is a few hundred bytes. +const MaxBodyBytes = 8 << 20 + +// ErrBodyTooLarge marks an answer over MaxBodyBytes. It is an error, never a +// prefix: a protobuf message cut at a field boundary still decodes, so a +// truncated ban set or router list would be read as a smaller, valid one. +var ErrBodyTooLarge = errors.New("control plane answer exceeds the body cap") + +// ReadBody reads a control plane response body of at most MaxBodyBytes and +// reports ErrBodyTooLarge for anything larger. +func ReadBody(r io.Reader) ([]byte, error) { + body, err := io.ReadAll(io.LimitReader(r, MaxBodyBytes+1)) + if err != nil { + return nil, fmt.Errorf("failed to read response body: %w", err) + } + if len(body) > MaxBodyBytes { + return nil, fmt.Errorf("%w (%d bytes)", ErrBodyTooLarge, MaxBodyBytes) + } + return body, nil +} + +// transport applies api.ValidateControlPlaneTransport to every request, +// redirects included, so a plaintext hop is refused wherever the URL came +// from. allowInsecure is read per request: the node learns the operator's +// choice after its clients exist. +type transport struct { + allowInsecure func() bool +} + +func (t transport) RoundTrip(req *http.Request) (*http.Response, error) { + if err := api.ValidateControlPlaneTransport(req.URL.String(), t.allowInsecure()); err != nil { + return nil, err + } + return http.DefaultTransport.RoundTrip(req) +} + +// NewHTTPClient is the HTTP client for every request a mesh component makes +// to its control plane. A nil allowInsecure never allows plaintext. +func NewHTTPClient(timeout time.Duration, allowInsecure func() bool) *http.Client { + if allowInsecure == nil { + allowInsecure = func() bool { return false } + } + return &http.Client{Timeout: timeout, Transport: transport{allowInsecure: allowInsecure}} +} + +// Client reads the pull side of the mesh protocol from one control plane. +type Client struct { + baseURL string + http *http.Client +} + +// New normalizes baseURL, https:// when no scheme is given and no trailing +// slash, and speaks through httpClient, which the caller builds with +// NewHTTPClient so its own transport policy applies. +func New(baseURL string, httpClient *http.Client) *Client { + if !strings.HasPrefix(baseURL, "http://") && !strings.HasPrefix(baseURL, "https://") { + baseURL = "https://" + baseURL + } + return &Client{baseURL: strings.TrimSuffix(baseURL, "/"), http: httpClient} +} + +// FetchInfo is GET /info: the router addresses, the ban set and the OIDC +// details a node needs to enroll. +func (c *Client) FetchInfo(ctx context.Context) (*api.ControlPlaneInfoResponse, error) { + var info api.ControlPlaneInfoResponse + if err := c.get(ctx, "/info", nil, &info); err != nil { + return nil, err + } + return &info, nil +} + +// FetchKeys is GET /keys: the control plane's currently valid signing keys. +// The set is accepted only if signed by a key in trusted +// (api.VerifyKeysResponse): whoever answers the URL must already be the +// control plane, not become it. +func (c *Client) FetchKeys(ctx context.Context, trusted []ed25519.PublicKey) ([]ed25519.PublicKey, error) { + var resp api.KeysResponse + if err := c.get(ctx, "/keys", nil, &resp); err != nil { + return nil, err + } + keys, err := api.VerifyKeysResponse(&resp, trusted, time.Now()) + if err != nil { + return nil, fmt.Errorf("/keys response rejected: %w", err) + } + return keys, nil +} + +// FetchPolicy is GET /policies, authenticated with the caller's biscuit: the +// roles and bindings a node compiles into its authorization rules. +func (c *Client) FetchPolicy(ctx context.Context, biscuit []byte) (*api.PolicyConfigGetResponse, error) { + var policy api.PolicyConfigGetResponse + if err := c.get(ctx, "/policies", biscuit, &policy); err != nil { + return nil, err + } + return &policy, nil +} + +func (c *Client) get(ctx context.Context, path string, biscuit []byte, msg proto.Message) error { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.baseURL+path, nil) + if err != nil { + return fmt.Errorf("failed to create HTTP request: %w", err) + } + if len(biscuit) > 0 { + req.Header.Set("Authorization", "Bearer "+base64.StdEncoding.EncodeToString(biscuit)) + } + resp, err := c.http.Do(req) + if err != nil { + return fmt.Errorf("HTTP request failed: %w", err) + } + defer func() { _ = resp.Body.Close() }() + + body, err := ReadBody(resp.Body) + if err != nil { + return fmt.Errorf("%s: %w", path, err) + } + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("control plane returned status %s: %s", resp.Status, string(body)) + } + if err := proto.Unmarshal(body, msg); err != nil { + return fmt.Errorf("failed to decode %s response: %w", path, err) + } + return nil +} diff --git a/internal/controlplane/client/client_test.go b/internal/controlplane/client/client_test.go new file mode 100644 index 00000000..ef2689e0 --- /dev/null +++ b/internal/controlplane/client/client_test.go @@ -0,0 +1,292 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package client + +import ( + "context" + "crypto/ed25519" + "crypto/rand" + "encoding/base64" + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "google.golang.org/protobuf/proto" + + "github.com/google/sam/api" +) + +func writeProto(t *testing.T, w http.ResponseWriter, msg proto.Message) { + t.Helper() + data, err := proto.Marshal(msg) + if err != nil { + t.Errorf("proto.Marshal(%T): %v", msg, err) + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + w.Header().Set("Content-Type", "application/x-protobuf") + if _, err := w.Write(data); err != nil { + t.Errorf("write %T: %v", msg, err) + } +} + +func TestFetchInfo(t *testing.T) { + want := &api.ControlPlaneInfoResponse{ + RouterAddresses: []string{"/ip4/10.0.0.1/tcp/4501"}, + BannedPeerIds: []string{"12D3KooWBanned"}, + } + var gotPath string + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotPath = r.URL.Path + writeProto(t, w, want) + })) + defer srv.Close() + + // A trailing slash on the base URL must not double up in the path. + c := New(srv.URL+"/", NewHTTPClient(time.Second, nil)) + info, err := c.FetchInfo(context.Background()) + if err != nil { + t.Fatalf("FetchInfo: %v", err) + } + if gotPath != "/info" { + t.Errorf("request path = %q, want /info", gotPath) + } + if !proto.Equal(info, want) { + t.Errorf("FetchInfo = %v, want %v", info, want) + } +} + +func TestFetchKeys(t *testing.T) { + trustedPub, trustedPriv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + successorPub, successorPriv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + strangerPub, strangerPriv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + + serve := func(t *testing.T, pubs []ed25519.PublicKey, privs []ed25519.PrivateKey) *Client { + t.Helper() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + resp := &api.KeysResponse{} + for _, p := range pubs { + resp.PublicKeys = append(resp.PublicKeys, p) + } + if privs != nil { + if err := api.SignKeysResponse(resp, privs, time.Now()); err != nil { + t.Errorf("SignKeysResponse: %v", err) + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + } + writeProto(t, w, resp) + })) + t.Cleanup(srv.Close) + return New(srv.URL, NewHTTPClient(time.Second, nil)) + } + + t.Run("a set vouched for by a trusted key is adopted whole", func(t *testing.T) { + c := serve(t, []ed25519.PublicKey{trustedPub, successorPub}, []ed25519.PrivateKey{trustedPriv, successorPriv}) + keys, err := c.FetchKeys(context.Background(), []ed25519.PublicKey{trustedPub}) + if err != nil { + t.Fatalf("FetchKeys: %v", err) + } + if len(keys) != 2 || !keys[0].Equal(trustedPub) || !keys[1].Equal(successorPub) { + t.Errorf("FetchKeys = %d keys, want trusted then successor", len(keys)) + } + }) + + t.Run("a set signed only by a stranger is rejected", func(t *testing.T) { + c := serve(t, []ed25519.PublicKey{strangerPub}, []ed25519.PrivateKey{strangerPriv}) + if _, err := c.FetchKeys(context.Background(), []ed25519.PublicKey{trustedPub}); err == nil { + t.Fatal("FetchKeys accepted a set no trusted key signed") + } + }) + + t.Run("an unsigned set is rejected", func(t *testing.T) { + c := serve(t, []ed25519.PublicKey{trustedPub, strangerPub}, nil) + if _, err := c.FetchKeys(context.Background(), []ed25519.PublicKey{trustedPub}); err == nil { + t.Fatal("FetchKeys accepted an unsigned set") + } + }) +} + +func TestErrors(t *testing.T) { + t.Run("non-200 carries the status and body", func(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Error(w, "store down", http.StatusServiceUnavailable) + })) + defer srv.Close() + c := New(srv.URL, NewHTTPClient(time.Second, nil)) + _, err := c.FetchInfo(context.Background()) + if err == nil || !strings.Contains(err.Error(), "503") || !strings.Contains(err.Error(), "store down") { + t.Fatalf("FetchInfo error = %v, want status and body", err) + } + }) + + t.Run("a body that is not the message is an error", func(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if _, err := w.Write([]byte("not protobuf")); err != nil { + t.Errorf("write: %v", err) + } + })) + defer srv.Close() + c := New(srv.URL, NewHTTPClient(time.Second, nil)) + if _, err := c.FetchInfo(context.Background()); err == nil || !strings.Contains(err.Error(), "decode /info") { + t.Fatalf("FetchInfo error = %v, want a decode error naming the path", err) + } + }) + + t.Run("an oversized body is an error, not a shorter message", func(t *testing.T) { + // A ban set cut at an entry boundary would decode as a valid, smaller + // ban set; the client must refuse the answer instead. + big := &api.ControlPlaneInfoResponse{} + for len(big.BannedPeerIds) < 200_000 { + big.BannedPeerIds = append(big.BannedPeerIds, "12D3KooWL7xNnc7bdzPGobhPxvunLC1uNhHfSCK8ZTXjR1YSTG9T") + } + if n := proto.Size(big); n <= MaxBodyBytes { + t.Fatalf("test answer is %d bytes, need more than %d", n, MaxBodyBytes) + } + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + writeProto(t, w, big) + })) + defer srv.Close() + c := New(srv.URL, NewHTTPClient(5*time.Second, nil)) + info, err := c.FetchInfo(context.Background()) + if !errors.Is(err, ErrBodyTooLarge) { + t.Fatalf("FetchInfo = (%d bans, %v), want ErrBodyTooLarge", len(info.GetBannedPeerIds()), err) + } + }) +} + +// The cap must never cut a legitimate answer: a large ban set well inside it +// arrives whole, and a body of exactly the cap still decodes. +func TestLargeAnswersArriveWhole(t *testing.T) { + const bans = 100_000 + big := &api.ControlPlaneInfoResponse{RouterAddresses: []string{"/ip4/10.0.0.1/tcp/4501"}} + for i := 0; i < bans; i++ { + big.BannedPeerIds = append(big.BannedPeerIds, "12D3KooWL7xNnc7bdzPGobhPxvunLC1uNhHfSCK8ZTXjR1YSTG9T") + } + if n := proto.Size(big); n >= MaxBodyBytes { + t.Fatalf("%d bans is %d bytes, over the %d cap: the cap is too small for a mesh this size", bans, n, MaxBodyBytes) + } + + t.Run("large ban set", func(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + writeProto(t, w, big) + })) + defer srv.Close() + info, err := New(srv.URL, NewHTTPClient(5*time.Second, nil)).FetchInfo(context.Background()) + if err != nil { + t.Fatalf("FetchInfo: %v", err) + } + if got := len(info.BannedPeerIds); got != bans { + t.Fatalf("FetchInfo returned %d bans, want %d: the answer was cut", got, bans) + } + }) + + t.Run("body of exactly the cap", func(t *testing.T) { + // Pad the router address so the encoded message is exactly MaxBodyBytes. + exact := &api.ControlPlaneInfoResponse{RouterAddresses: []string{""}} + pad := MaxBodyBytes - proto.Size(exact) - 3 // 3: the length prefix grows to three bytes + exact.RouterAddresses[0] = strings.Repeat("a", pad) + for proto.Size(exact) < MaxBodyBytes { + exact.RouterAddresses[0] += "a" + } + if n := proto.Size(exact); n != MaxBodyBytes { + t.Fatalf("test body is %d bytes, want exactly %d", n, MaxBodyBytes) + } + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + writeProto(t, w, exact) + })) + defer srv.Close() + info, err := New(srv.URL, NewHTTPClient(5*time.Second, nil)).FetchInfo(context.Background()) + if err != nil { + t.Fatalf("FetchInfo at exactly the cap: %v", err) + } + if len(info.RouterAddresses) != 1 || len(info.RouterAddresses[0]) != len(exact.RouterAddresses[0]) { + t.Fatal("body of exactly the cap did not arrive whole") + } + }) +} + +func TestFetchPolicy(t *testing.T) { + want := &api.PolicyConfigGetResponse{ + Roles: []*api.PolicyRole{{Name: "developer", AllowedServices: []string{"mcp://*"}}}, + Bindings: []*api.PolicyBinding{{Role: "developer", Members: []string{"group:eng"}}}, + } + var gotAuth string + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotAuth = r.Header.Get("Authorization") + if r.URL.Path != "/policies" { + http.NotFound(w, r) + return + } + writeProto(t, w, want) + })) + defer srv.Close() + + policy, err := New(srv.URL, NewHTTPClient(time.Second, nil)).FetchPolicy(context.Background(), []byte("biscuit")) + if err != nil { + t.Fatalf("FetchPolicy: %v", err) + } + if !proto.Equal(policy, want) { + t.Errorf("FetchPolicy = %v, want %v", policy, want) + } + if gotAuth != "Bearer "+base64.StdEncoding.EncodeToString([]byte("biscuit")) { + t.Errorf("Authorization = %q, want the biscuit as a base64 bearer token", gotAuth) + } +} + +// The control plane is the trust root, so a plaintext hop to it is refused +// unless the operator opted in; loopback is the standalone case and is fine. +// The choice is read per request, since the node learns it after its clients +// exist. +func TestHTTPClientTransportPolicy(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + writeProto(t, w, &api.ControlPlaneInfoResponse{}) + })) + defer srv.Close() + // httptest binds 127.0.0.1; spell it as a non-loopback name that the + // transport must refuse before any connection is attempted. + nonLoopbackURL := strings.Replace(srv.URL, "127.0.0.1", "sam-control-plane.invalid", 1) + + allow := false + httpClient := NewHTTPClient(time.Second, func() bool { return allow }) + + if _, err := New(srv.URL, httpClient).FetchInfo(context.Background()); err != nil { + t.Fatalf("loopback plaintext must be accepted: %v", err) + } + _, err := New(nonLoopbackURL, httpClient).FetchInfo(context.Background()) + if !errors.Is(err, api.ErrInsecureControlPlaneURL) { + t.Fatalf("plaintext to a non-loopback host: err = %v, want %v", err, api.ErrInsecureControlPlaneURL) + } + + // With the opt-in the request is attempted; the name does not resolve, + // which is a dial error, not the policy error. + allow = true + _, err = New(nonLoopbackURL, httpClient).FetchInfo(context.Background()) + if err == nil || errors.Is(err, api.ErrInsecureControlPlaneURL) { + t.Fatalf("with the opt-in the policy must not be what fails: %v", err) + } +} diff --git a/internal/node/controlplane.go b/internal/node/controlplane.go index f5bfb310..99540802 100644 --- a/internal/node/controlplane.go +++ b/internal/node/controlplane.go @@ -26,94 +26,35 @@ import ( "time" "github.com/google/sam/api" - "github.com/multiformats/go-multiaddr" + cpclient "github.com/google/sam/internal/controlplane/client" "google.golang.org/protobuf/proto" ) // maxControlPlaneBodyBytes caps every response body read from the control // plane or an IdP: a misbehaving or impersonated server must not be able to -// make the node buffer arbitrary amounts of memory. -const maxControlPlaneBodyBytes = 1 << 20 +// make the node buffer arbitrary amounts of memory. Bodies that carry a +// message go through cpclient.ReadBody, which turns an oversized answer into +// an error rather than a truncated message. +const maxControlPlaneBodyBytes = cpclient.MaxBodyBytes + +// controlPlaneClient speaks the pull endpoints of controlPlaneURL through the +// node's transport policy. +func controlPlaneClient(controlPlaneURL string) *cpclient.Client { + return cpclient.New(controlPlaneURL, controlPlaneHTTPClient(10*time.Second)) +} // FetchControlPlaneInfo retrieves the latest configuration from the control plane's /info endpoint. func FetchControlPlaneInfo(ctx context.Context, controlPlaneURL string) (*api.ControlPlaneInfoResponse, error) { - if !strings.HasPrefix(controlPlaneURL, "http://") && !strings.HasPrefix(controlPlaneURL, "https://") { - controlPlaneURL = "https://" + controlPlaneURL - } - controlPlaneURL = strings.TrimSuffix(controlPlaneURL, "/") - - urlStr := controlPlaneURL + "/info" - req, err := http.NewRequestWithContext(ctx, "GET", urlStr, nil) - if err != nil { - return nil, fmt.Errorf("failed to create HTTP request: %w", err) - } - - client := controlPlaneHTTPClient(10 * time.Second) - resp, err := client.Do(req) - if err != nil { - return nil, fmt.Errorf("HTTP request failed: %w", err) - } - defer resp.Body.Close() //nolint:errcheck - - body, err := io.ReadAll(io.LimitReader(resp.Body, maxControlPlaneBodyBytes)) - if err != nil { - return nil, fmt.Errorf("failed to read response body: %w", err) - } - - if resp.StatusCode != http.StatusOK { - return nil, fmt.Errorf("control plane returned status %s: %s", resp.Status, string(body)) - } - - var info api.ControlPlaneInfoResponse - if err := proto.Unmarshal(body, &info); err != nil { - return nil, fmt.Errorf("failed to decode /info response: %w", err) - } - - return &info, nil + return controlPlaneClient(controlPlaneURL).FetchInfo(ctx) } // FetchControlPlaneKeys retrieves the full set of currently valid control -// plane public keys from the /keys endpoint — the same catch-up path routers -// use. Enrollment only hands out the newest key, so this is how a node -// learns keys still in their rotation grace period, or rotations it missed -// while offline. The set is accepted only if signed by a key in trusted: -// whoever answers /keys must already be the control plane, not become it. +// plane public keys from /keys, the same catch-up path routers use. +// Enrollment only hands out the newest key, so this is how a node learns +// keys still in their rotation grace period, or rotations it missed while +// offline. The set is accepted only if signed by a key in trusted. func FetchControlPlaneKeys(ctx context.Context, controlPlaneURL string, trusted []ed25519.PublicKey) ([]ed25519.PublicKey, error) { - if !strings.HasPrefix(controlPlaneURL, "http://") && !strings.HasPrefix(controlPlaneURL, "https://") { - controlPlaneURL = "https://" + controlPlaneURL - } - controlPlaneURL = strings.TrimSuffix(controlPlaneURL, "/") - - req, err := http.NewRequestWithContext(ctx, "GET", controlPlaneURL+"/keys", nil) - if err != nil { - return nil, fmt.Errorf("failed to create HTTP request: %w", err) - } - - client := controlPlaneHTTPClient(10 * time.Second) - resp, err := client.Do(req) - if err != nil { - return nil, fmt.Errorf("HTTP request failed: %w", err) - } - defer resp.Body.Close() //nolint:errcheck - - body, err := io.ReadAll(io.LimitReader(resp.Body, maxControlPlaneBodyBytes)) - if err != nil { - return nil, fmt.Errorf("failed to read response body: %w", err) - } - if resp.StatusCode != http.StatusOK { - return nil, fmt.Errorf("control plane returned status %s: %s", resp.Status, string(body)) - } - - var keysResp api.KeysResponse - if err := proto.Unmarshal(body, &keysResp); err != nil { - return nil, fmt.Errorf("failed to decode /keys response: %w", err) - } - - keys, err := api.VerifyKeysResponse(&keysResp, trusted, time.Now()) - if err != nil { - return nil, fmt.Errorf("/keys response rejected: %w", err) - } - return keys, nil + return controlPlaneClient(controlPlaneURL).FetchKeys(ctx, trusted) } // mergeTrustedKeys replaces the stored trust set with the authoritative set @@ -134,88 +75,6 @@ func mergeTrustedKeys(existing []TrustedKey, fetched []ed25519.PublicKey, now ti return merged } -// SyncMeshConfig loads the mesh configuration from the store, attempts to refresh it -// via HTTP from the control plane, and updates the store if successful. -// It returns the control plane public key, the latest multiaddresses, and the -// control plane's current ban set. -// -// The ban set is deliberately not persisted. MeshEvent_BANNED is published once -// and gossip has no replay, so a node that restarted or was offline has to be -// told again; /info is that catch-up, and reading it fresh each start is also -// what makes an unban take effect. Nil means the control plane was not reached, -// which is not the same as "nothing is banned": callers must not treat it as an -// instruction to clear anything. -func SyncMeshConfig(ctx context.Context, s *Store) ([]byte, []multiaddr.Multiaddr, []string, error) { - pubKey, storedAddrsStr, err := s.LoadMeshConfig() - if err != nil { - return nil, nil, nil, fmt.Errorf("failed to load mesh config from store: %w", err) - } - - controlPlaneURL, err := s.LoadControlPlaneURL() - if err != nil { - return nil, nil, nil, fmt.Errorf("failed to load control plane URL from store: %w", err) - } - var bannedPeerIDs []string - var routerAddrs []multiaddr.Multiaddr - - // Parse stored addresses - for _, addrStr := range storedAddrsStr { - ma, err := multiaddr.NewMultiaddr(addrStr) - if err != nil { - logger.Warnf("Failed to parse stored router address %q: %v", addrStr, err) - continue - } - routerAddrs = append(routerAddrs, ma) - } - - // If we have a URL, fetch the latest info - if controlPlaneURL != "" { - logger.Infof("Fetching latest router addresses via HTTP from %s...", controlPlaneURL) - info, err := FetchControlPlaneInfo(ctx, controlPlaneURL) - if err != nil { - logger.Warnf("Failed to fetch updated addresses via HTTP (using cached): %v", err) - } else if bannedPeerIDs = info.GetBannedPeerIds(); len(info.RouterAddresses) > 0 { - logger.Infof("Discovered latest router addresses: %v", info.RouterAddresses) - var newRouterAddrs []multiaddr.Multiaddr - for _, addrStr := range info.RouterAddresses { - ma, parseErr := multiaddr.NewMultiaddr(addrStr) - if parseErr != nil { - logger.Warnf("Failed to parse discovered router address %q: %v", addrStr, parseErr) - continue - } - newRouterAddrs = append(newRouterAddrs, ma) - } - if len(newRouterAddrs) > 0 { - routerAddrs = newRouterAddrs - if len(pubKey) > 0 { - if saveErr := s.SaveMeshConfig(pubKey, info.RouterAddresses); saveErr != nil { - logger.Errorf("Failed to save updated mesh config to store: %v", saveErr) - } - } - } - } - - // Catch up on the valid key set. Verified against what is already - // trusted, so with nothing stored yet there is nothing to do: the - // first key comes from enrollment. An empty result would wipe the - // trust set, so it is ignored like any fetch failure. - existing, loadErr := s.LoadTrustedKeys() - if loadErr != nil { - logger.Warnf("Failed to load stored trusted keys, skipping key sync: %v", loadErr) - } else if len(existing) == 0 { - logger.Debugf("No trusted control plane keys stored yet; skipping /keys sync until enrolled") - } else if keys, keysErr := FetchControlPlaneKeys(ctx, controlPlaneURL, publicKeysOf(existing)); keysErr != nil { - logger.Warnf("Failed to fetch control plane keys via HTTP (using cached): %v", keysErr) - } else if len(keys) > 0 { - if saveErr := s.SaveTrustedKeys(mergeTrustedKeys(existing, keys, time.Now())); saveErr != nil { - logger.Errorf("Failed to save trusted keys to store: %v", saveErr) - } - } - } - - return pubKey, routerAddrs, bannedPeerIDs, nil -} - func publicKeysOf(keys []TrustedKey) []ed25519.PublicKey { out := make([]ed25519.PublicKey, 0, len(keys)) for _, tk := range keys { @@ -226,41 +85,7 @@ func publicKeysOf(keys []TrustedKey) []ed25519.PublicKey { // FetchMeshPolicy retrieves the latest mesh policy from the control plane's /policies endpoint using a biscuit token. func FetchMeshPolicy(ctx context.Context, controlPlaneURL string, biscuitToken []byte) (*api.PolicyConfigGetResponse, error) { - if !strings.HasPrefix(controlPlaneURL, "http://") && !strings.HasPrefix(controlPlaneURL, "https://") { - controlPlaneURL = "https://" + controlPlaneURL - } - controlPlaneURL = strings.TrimSuffix(controlPlaneURL, "/") - - urlStr := controlPlaneURL + "/policies" - req, err := http.NewRequestWithContext(ctx, "GET", urlStr, nil) - if err != nil { - return nil, fmt.Errorf("failed to create HTTP request: %w", err) - } - - req.Header.Set("Authorization", "Bearer "+base64.StdEncoding.EncodeToString(biscuitToken)) - - client := controlPlaneHTTPClient(10 * time.Second) - resp, err := client.Do(req) - if err != nil { - return nil, fmt.Errorf("HTTP request failed: %w", err) - } - defer resp.Body.Close() //nolint:errcheck - - body, err := io.ReadAll(io.LimitReader(resp.Body, maxControlPlaneBodyBytes)) - if err != nil { - return nil, fmt.Errorf("failed to read response body: %w", err) - } - - if resp.StatusCode != http.StatusOK { - return nil, fmt.Errorf("control plane returned status %s: %s", resp.Status, string(body)) - } - - var policyResp api.PolicyConfigGetResponse - if err := proto.Unmarshal(body, &policyResp); err != nil { - return nil, fmt.Errorf("failed to decode /policies response: %w", err) - } - - return &policyResp, nil + return controlPlaneClient(controlPlaneURL).FetchPolicy(ctx, biscuitToken) } // ReportNodeCatalog self-reports this node's locally registered services to diff --git a/internal/node/controlplane_client.go b/internal/node/controlplane_client.go index 8411e89a..83ff75a0 100644 --- a/internal/node/controlplane_client.go +++ b/internal/node/controlplane_client.go @@ -19,7 +19,7 @@ import ( "sync/atomic" "time" - "github.com/google/sam/api" + cpclient "github.com/google/sam/internal/controlplane/client" ) // allowInsecureControlPlane is process-wide because the control-plane URL @@ -33,25 +33,8 @@ func SetAllowInsecureControlPlane(allow bool) { allowInsecureControlPlane.Store(allow) } -// controlPlaneTransport applies api.ValidateControlPlaneTransport to every -// request, including redirects, so a plaintext hop is refused wherever the -// URL came from. -type controlPlaneTransport struct { - base http.RoundTripper -} - -func (t controlPlaneTransport) RoundTrip(req *http.Request) (*http.Response, error) { - if err := api.ValidateControlPlaneTransport(req.URL.String(), allowInsecureControlPlane.Load()); err != nil { - return nil, err - } - return t.base.RoundTrip(req) -} - // controlPlaneHTTPClient is the client for every request the node makes to -// its control plane. +// its control plane; the plaintext policy is re-checked on every hop. func controlPlaneHTTPClient(timeout time.Duration) *http.Client { - return &http.Client{ - Timeout: timeout, - Transport: controlPlaneTransport{base: http.DefaultTransport}, - } + return cpclient.NewHTTPClient(timeout, allowInsecureControlPlane.Load) } diff --git a/internal/node/controlplane_sync.go b/internal/node/controlplane_sync.go index 619e52d2..ab375fbb 100644 --- a/internal/node/controlplane_sync.go +++ b/internal/node/controlplane_sync.go @@ -22,6 +22,7 @@ import ( "time" "github.com/libp2p/go-libp2p/core/peer" + "github.com/multiformats/go-multiaddr" ) // The node reads three things from the control plane while it runs: the @@ -32,10 +33,12 @@ import ( // arrives, because a once-published event with no replay is missed by any // node that was down, partitioned, or simply enrolled later. -// syncControlPlane is the node's single pull from the control plane. Each +// SyncControlPlane is the node's single pull from the control plane. Each // part is attempted even when another fails, so a policy outage does not -// stop a key rotation from landing; the errors are reported together. -func (n *SamNode) syncControlPlane(ctx context.Context) error { +// stop a key rotation from landing; the errors are reported together. Called +// once before Start, so the node dials the routers the control plane knows +// now and enforces the bans it holds now, and then by the sync loop. +func (n *SamNode) SyncControlPlane(ctx context.Context) error { if n.Store == nil { return errors.New("node has no store") } @@ -87,9 +90,11 @@ func (n *SamNode) syncTrustedKeys(ctx context.Context, controlPlaneURL string) e return nil } -// syncMeshInfo reads /info: the router addresses are persisted for the next -// start, and the ban set is reconciled against the revocation cache, which is -// how a running node learns a ban or an unban it got no event for. +// syncMeshInfo reads /info. The router addresses are persisted for the next +// start and, until the host exists, adopted as the static relays Start will +// dial; once it is running they only matter to the next start. The ban set +// is reconciled against the revocation cache, which is how a running node +// learns a ban or an unban it got no event for. func (n *SamNode) syncMeshInfo(ctx context.Context, controlPlaneURL string) error { // Taken before the request: a ban recorded after this instant cannot be // in the answer, so its absence must not be read as an unban. @@ -108,11 +113,30 @@ func (n *SamNode) syncMeshInfo(ctx context.Context, controlPlaneURL string) erro logger.Warnf("Failed to persist router addresses: %v", saveErr) } } + if n.Host == nil { + if addrs := parseMultiaddrs(info.RouterAddresses); len(addrs) > 0 { + n.config.RouterAddrs = addrs + } + } } n.reconcileBannedPeers(info.GetBannedPeerIds(), fetchedAt) return nil } +// parseMultiaddrs keeps the addresses that parse and logs the rest. +func parseMultiaddrs(addrs []string) []multiaddr.Multiaddr { + out := make([]multiaddr.Multiaddr, 0, len(addrs)) + for _, s := range addrs { + ma, err := multiaddr.NewMultiaddr(s) + if err != nil { + logger.Warnf("Ignoring router address %q from the control plane: %v", s, err) + continue + } + out = append(out, ma) + } + return out +} + // reconcileBannedPeers makes the revocation cache match the control plane's // ban set as of fetchedAt. Bans recorded at or after fetchedAt are kept even // when absent from the answer: the answer predates them and cannot speak to @@ -209,7 +233,7 @@ func (n *SamNode) startControlPlaneSyncLoop(ctx context.Context, interval time.D } timer.Reset(interval + time.Duration(rand.Int63n(int64(interval/10)+1))) - err := n.syncControlPlane(ctx) + err := n.SyncControlPlane(ctx) switch { case err != nil && failures == 0: logger.Warnf("Control plane sync failed: %v", err) diff --git a/internal/node/controlplane_test.go b/internal/node/controlplane_test.go index a45b3812..eddb368e 100644 --- a/internal/node/controlplane_test.go +++ b/internal/node/controlplane_test.go @@ -15,6 +15,7 @@ package node import ( + "bytes" "context" "crypto/ed25519" "encoding/base64" @@ -29,7 +30,9 @@ import ( "github.com/google/sam/api" lru "github.com/hashicorp/golang-lru/v2" + "github.com/libp2p/go-libp2p/core/crypto" "github.com/libp2p/go-libp2p/core/peer" + "github.com/multiformats/go-multiaddr" "google.golang.org/protobuf/proto" ) @@ -185,7 +188,8 @@ func TestSyncTrustedKeys(t *testing.T) { // A ban the control plane holds is applied, one it no longer holds is lifted, // and one recorded after the answer was requested is left alone: the answer -// predates it and cannot speak to it. +// predates it and cannot speak to it. The wire form may be any encoding of +// the peer ID; the cache is keyed on the canonical one. func TestReconcileBannedPeers(t *testing.T) { stillBanned := randomPeerID(t) unbanned := randomPeerID(t) @@ -202,13 +206,16 @@ func TestReconcileBannedPeers(t *testing.T) { n.revokedPeers.Add(unbanned.String(), fetchedAt.Add(-time.Hour).UnixMilli()) n.revokedPeers.Add(bannedAfterFetch.String(), fetchedAt.UnixMilli()) - n.reconcileBannedPeers([]string{stillBanned.String(), newlyBanned.String(), "not-a-peer-id"}, fetchedAt) + n.reconcileBannedPeers([]string{stillBanned.String(), peer.ToCid(newlyBanned).String(), "not-a-peer-id"}, fetchedAt) for _, want := range []peer.ID{stillBanned, newlyBanned, bannedAfterFetch} { if !n.revokedPeers.Contains(want.String()) { t.Errorf("%s must be banned after reconciliation", want) } } + if n.revokedPeers.Contains(peer.ToCid(newlyBanned).String()) { + t.Error("the cache must be keyed on the canonical peer ID, not the wire encoding") + } if n.revokedPeers.Contains(unbanned.String()) { t.Error("a peer absent from the control plane's ban set must be unbanned") } @@ -251,9 +258,9 @@ func TestSyncControlPlane(t *testing.T) { n := &SamNode{Store: store, revokedPeers: cache, trustedKeys: []TrustedKey{{Key: oldPub, ReceivedAt: time.Now()}}} n.SetIdentityCache([]byte("identity")) - err = n.syncControlPlane(context.Background()) + err = n.SyncControlPlane(context.Background()) if err == nil || !strings.Contains(err.Error(), "policy") { - t.Fatalf("syncControlPlane = %v, want the policy failure reported", err) + t.Fatalf("SyncControlPlane = %v, want the policy failure reported", err) } if !containsTrustedKey(n.trustedKeys, newPub) { t.Error("rotated key not learned although /keys answered") @@ -268,6 +275,10 @@ func TestSyncControlPlane(t *testing.T) { if len(addrs) != 1 { t.Errorf("router addresses not persisted: got %v, want 1", addrs) } + // Before Start the answer also becomes the static relays Start will dial. + if len(n.config.RouterAddrs) != 1 || n.config.RouterAddrs[0].String() != addrs[0] { + t.Errorf("router addresses not adopted before start: got %v, want %v", n.config.RouterAddrs, addrs) + } } // The loop is what turns a missed gossip event into a delay rather than a @@ -387,121 +398,103 @@ func TestFetchControlPlaneInfo_InvalidProto(t *testing.T) { } } -func TestSyncMeshConfig(t *testing.T) { - expectedInfo := &api.ControlPlaneInfoResponse{ - RouterAddresses: []string{"/ip4/127.0.0.1/tcp/4001"}, - OidcIssuer: "https://issuer.example.com", - ClientId: "client-id", - } +// A node built from stored config, as sam-node run does, must come out of +// its pre-start pull with the control plane's current router addresses and +// the full key set, both in memory and on disk; and an unreachable control +// plane must leave the stored config in place rather than blank it. +func TestSyncControlPlaneBeforeStart(t *testing.T) { + cpPub, cpPriv := mustGenerateKey(t) + gracePub, gracePriv := mustGenerateKey(t) + freshRouter := "/ip4/10.0.0.9/tcp/4501/p2p/" + randomPeerID(t).String() - body, err := proto.Marshal(expectedInfo) - if err != nil { - t.Fatalf("Failed to marshal info: %v", err) - } - - cpPub, cpPriv, _ := ed25519.GenerateKey(nil) - gracePub, gracePriv, _ := ed25519.GenerateKey(nil) - keysResp := &api.KeysResponse{PublicKeys: [][]byte{cpPub, gracePub}} - if err := api.SignKeysResponse(keysResp, []ed25519.PrivateKey{cpPriv, gracePriv}, time.Now()); err != nil { - t.Fatal(err) - } - keysBody, err := proto.Marshal(keysResp) - if err != nil { - t.Fatalf("Failed to marshal keys: %v", err) - } + mux := http.NewServeMux() + mux.HandleFunc("/keys", keysHandler(t, []ed25519.PublicKey{cpPub, gracePub}, []ed25519.PrivateKey{cpPriv, gracePriv})) + mux.HandleFunc("/info", protoHandler(t, &api.ControlPlaneInfoResponse{RouterAddresses: []string{freshRouter}})) + mux.HandleFunc("/policies", protoHandler(t, &api.PolicyConfigGetResponse{})) + srv := httptest.NewServer(mux) + defer srv.Close() - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusOK) - if r.URL.Path == "/keys" { - _, _ = w.Write(keysBody) - return + newStoredNode := func(t *testing.T, controlPlaneURL string) *SamNode { + t.Helper() + store, err := NewStore(t.TempDir()) + if err != nil { + t.Fatal(err) } - _, _ = w.Write(body) - })) - defer server.Close() - - tempDir := t.TempDir() - store, err := NewStore(tempDir) - if err != nil { - t.Fatalf("Failed to create store: %v", err) - } - defer store.Close() //nolint:errcheck - - // Initial store is empty, so SyncMeshConfig should just return empty - pubKey, addrs, bannedPeerIDs, err := SyncMeshConfig(context.Background(), store) - if err != nil { - t.Fatalf("SyncMeshConfig failed: %v", err) - } - if len(bannedPeerIDs) != 0 { - t.Errorf("Expected no banned peers for empty store, got %v", bannedPeerIDs) - } - if len(pubKey) != 0 || len(addrs) != 0 { - t.Errorf("Expected empty result for empty store, got pubKey=%v, addrs=%v", pubKey, addrs) - } - - // Save initial config with explicit control plane URL. The node trusts - // only cpPub, as after enrollment; the grace key must be learned through - // cpPub's signature on the set. - testPubKey := []byte("test-pub-key") - if err := store.SaveMeshConfig(testPubKey, []string{"/ip4/1.2.3.4/tcp/1234"}); err != nil { - t.Fatalf("Failed to save mesh config: %v", err) - } - if err := store.SaveControlPlaneURL(server.URL); err != nil { - t.Fatalf("Failed to save control plane URL: %v", err) - } - if err := store.SaveTrustedKeys([]TrustedKey{{Key: cpPub, ReceivedAt: time.Now()}}); err != nil { - t.Fatalf("SaveTrustedKeys: %v", err) - } - - // Call SyncMeshConfig, it should fetch new addrs from server - pubKey, addrs, _, err = SyncMeshConfig(context.Background(), store) - if err != nil { - t.Fatalf("SyncMeshConfig failed: %v", err) - } - - if string(pubKey) != string(testPubKey) { - t.Errorf("Expected pubKey %s, got %s", testPubKey, pubKey) - } - - if len(addrs) != 1 || addrs[0].String() != expectedInfo.RouterAddresses[0] { - t.Errorf("Expected addrs %v, got %v", expectedInfo.RouterAddresses, addrs) + t.Cleanup(func() { _ = store.Close() }) + if err := store.SaveMeshConfig(cpPub, []string{"/ip4/1.2.3.4/tcp/1234"}); err != nil { + t.Fatal(err) + } + if err := store.SaveControlPlaneURL(controlPlaneURL); err != nil { + t.Fatal(err) + } + if err := store.SaveTrustedKeys([]TrustedKey{{Key: cpPub, ReceivedAt: time.Now()}}); err != nil { + t.Fatal(err) + } + stale, err := multiaddr.NewMultiaddr("/ip4/1.2.3.4/tcp/1234") + if err != nil { + t.Fatal(err) + } + priv, _, err := crypto.GenerateKeyPair(crypto.Ed25519, -1) + if err != nil { + t.Fatal(err) + } + n, err := NewSamNode(Options{PrivKey: priv, Store: store, ControlPlanePubKey: cpPub, RouterAddrs: []multiaddr.Multiaddr{stale}}) + if err != nil { + t.Fatal(err) + } + n.SetIdentityCache([]byte("identity")) + return n } - // Verify the new addrs were saved to the store - savedPubKey, savedAddrsStr, err := store.LoadMeshConfig() - if err != nil { - t.Fatalf("Failed to load mesh config: %v", err) - } - if string(savedPubKey) != string(testPubKey) { - t.Errorf("Expected saved pubKey %s, got %s", testPubKey, savedPubKey) - } + t.Run("reachable control plane", func(t *testing.T) { + n := newStoredNode(t, srv.URL) + if err := n.SyncControlPlane(context.Background()); err != nil { + t.Fatalf("SyncControlPlane: %v", err) + } + if len(n.config.RouterAddrs) != 1 || n.config.RouterAddrs[0].String() != freshRouter { + t.Errorf("RouterAddrs = %v, want the control plane's %s", n.config.RouterAddrs, freshRouter) + } + if len(n.trustedKeys) != 2 || !containsTrustedKey(n.trustedKeys, gracePub) { + t.Errorf("trust set = %d keys, want the enrollment key and the grace key", len(n.trustedKeys)) + } + savedPub, savedAddrs, err := n.Store.LoadMeshConfig() + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(savedPub, cpPub) || len(savedAddrs) != 1 || savedAddrs[0] != freshRouter { + t.Errorf("persisted mesh config = (%x, %v), want (%x, [%s])", savedPub, savedAddrs, cpPub, freshRouter) + } + stored, err := n.Store.LoadTrustedKeys() + if err != nil { + t.Fatal(err) + } + if len(stored) != 2 { + t.Errorf("persisted %d trusted keys, want 2", len(stored)) + } + }) - // The full valid key set from /keys must have been persisted - trusted, err := store.LoadTrustedKeys() - if err != nil { - t.Fatalf("LoadTrustedKeys: %v", err) - } - if len(trusted) != 2 { - t.Fatalf("expected 2 trusted keys from /keys, got %d", len(trusted)) - } - if !trusted[0].Key.Equal(cpPub) || !trusted[1].Key.Equal(gracePub) { - t.Errorf("persisted keys do not match /keys response") - } - if len(savedAddrsStr) != 1 || savedAddrsStr[0] != expectedInfo.RouterAddresses[0] { - t.Errorf("Expected saved addrs %v, got %v", expectedInfo.RouterAddresses, savedAddrsStr) - } + t.Run("unreachable control plane keeps the stored config", func(t *testing.T) { + down := httptest.NewServer(http.NotFoundHandler()) + down.Close() + n := newStoredNode(t, down.URL) + if err := n.SyncControlPlane(context.Background()); err == nil { + t.Fatal("SyncControlPlane against a closed server must report the failure") + } + if len(n.config.RouterAddrs) != 1 || n.config.RouterAddrs[0].String() != "/ip4/1.2.3.4/tcp/1234" { + t.Errorf("RouterAddrs = %v, want the stored address kept", n.config.RouterAddrs) + } + if len(n.trustedKeys) != 1 || !n.trustedKeys[0].Key.Equal(cpPub) { + t.Errorf("trust set = %d keys, want the stored key kept", len(n.trustedKeys)) + } + }) } // Whoever answers /keys must already be the control plane: a set that is not -// signed by a key the node trusts leaves the trust set untouched. -func TestSyncMeshConfigRefusesUntrustedKeySet(t *testing.T) { - cpPub, _, _ := ed25519.GenerateKey(nil) - attackerPub, attackerPriv, _ := ed25519.GenerateKey(nil) - - infoBody, err := proto.Marshal(&api.ControlPlaneInfoResponse{RouterAddresses: []string{"/ip4/127.0.0.1/tcp/4001"}}) - if err != nil { - t.Fatal(err) - } +// signed by a key the node trusts leaves the trust set untouched, in memory +// and on disk, while the rest of the pull still lands. +func TestSyncControlPlaneRefusesUntrustedKeySet(t *testing.T) { + cpPub, _ := mustGenerateKey(t) + attackerPub, attackerPriv := mustGenerateKey(t) for name, keysResp := range map[string]*api.KeysResponse{ "unsigned set": {PublicKeys: [][]byte{cpPub, attackerPub}}, @@ -514,44 +507,46 @@ func TestSyncMeshConfigRefusesUntrustedKeySet(t *testing.T) { }(), } { t.Run(name, func(t *testing.T) { - keysBody, err := proto.Marshal(keysResp) - if err != nil { - t.Fatal(err) - } - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/keys" { - _, _ = w.Write(keysBody) - return - } - _, _ = w.Write(infoBody) - })) - defer server.Close() + mux := http.NewServeMux() + mux.HandleFunc("/keys", protoHandler(t, keysResp)) + mux.HandleFunc("/info", protoHandler(t, &api.ControlPlaneInfoResponse{RouterAddresses: []string{"/ip4/127.0.0.1/tcp/4001"}})) + mux.HandleFunc("/policies", protoHandler(t, &api.PolicyConfigGetResponse{})) + srv := httptest.NewServer(mux) + defer srv.Close() store, err := NewStore(t.TempDir()) if err != nil { t.Fatal(err) } - defer store.Close() //nolint:errcheck + defer func() { _ = store.Close() }() if err := store.SaveMeshConfig(cpPub, nil); err != nil { t.Fatal(err) } - if err := store.SaveControlPlaneURL(server.URL); err != nil { + if err := store.SaveControlPlaneURL(srv.URL); err != nil { t.Fatal(err) } if err := store.SaveTrustedKeys([]TrustedKey{{Key: cpPub, ReceivedAt: time.Now()}}); err != nil { t.Fatal(err) } + n := &SamNode{Store: store, trustedKeys: []TrustedKey{{Key: cpPub, ReceivedAt: time.Now()}}} + n.SetIdentityCache([]byte("identity")) - if _, _, _, err := SyncMeshConfig(context.Background(), store); err != nil { - t.Fatalf("SyncMeshConfig: %v", err) + err = n.SyncControlPlane(context.Background()) + if err == nil || !strings.Contains(err.Error(), "keys") { + t.Fatalf("SyncControlPlane = %v, want the /keys rejection reported", err) + } + if len(n.trustedKeys) != 1 || !n.trustedKeys[0].Key.Equal(cpPub) { + t.Fatalf("in-memory trust set was replaced by an unverified /keys answer: %d keys", len(n.trustedKeys)) } - trusted, err := store.LoadTrustedKeys() if err != nil { t.Fatal(err) } if len(trusted) != 1 || !trusted[0].Key.Equal(cpPub) { - t.Fatalf("trust set was replaced by an unverified /keys answer: %d keys", len(trusted)) + t.Fatalf("persisted trust set was replaced by an unverified /keys answer: %d keys", len(trusted)) + } + if len(n.config.RouterAddrs) != 1 { + t.Errorf("a rejected /keys must not stop /info from landing: RouterAddrs = %v", n.config.RouterAddrs) } }) } diff --git a/internal/node/enroll.go b/internal/node/enroll.go index b1f3efd7..c113a3bb 100644 --- a/internal/node/enroll.go +++ b/internal/node/enroll.go @@ -28,6 +28,7 @@ import ( "time" "github.com/google/sam/api" + cpclient "github.com/google/sam/internal/controlplane/client" "github.com/google/sam/internal/identity" golog "github.com/ipfs/go-log/v2" "github.com/libp2p/go-libp2p/core/crypto" @@ -171,9 +172,9 @@ func (n *SamNode) processEnrollResponse(resp *http.Response) (*api.EnrollRespons return nil, fmt.Errorf("enrollment failed with status %s: %s", resp.Status, string(body)) } - respData, err := io.ReadAll(io.LimitReader(resp.Body, maxControlPlaneBodyBytes)) + respData, err := cpclient.ReadBody(resp.Body) if err != nil { - return nil, fmt.Errorf("failed to read response body: %v", err) + return nil, fmt.Errorf("enrollment response: %w", err) } var enrollResp api.EnrollResponse @@ -289,9 +290,9 @@ func (n *SamNode) EnrollBootstrap(ctx context.Context, controlPlaneURL string, b return fmt.Errorf("enrollment failed with status %s: %s", resp.Status, string(body)) } - respData, err := io.ReadAll(io.LimitReader(resp.Body, maxControlPlaneBodyBytes)) + respData, err := cpclient.ReadBody(resp.Body) if err != nil { - return fmt.Errorf("failed to read response body: %w", err) + return fmt.Errorf("bootstrap enrollment response: %w", err) } enrollResp := &api.BootstrapEnrollResponse{} @@ -351,10 +352,10 @@ func (n *SamNode) EnrollBootstrap(ctx context.Context, controlPlaneURL string, b continue } - hRespData, err := io.ReadAll(io.LimitReader(hResp.Body, maxControlPlaneBodyBytes)) + hRespData, err := cpclient.ReadBody(hResp.Body) _ = hResp.Body.Close() if err != nil { - logger.Warnf("Failed to read status response body: %v", err) + logger.Warnf("Enrollment status response: %v", err) continue } diff --git a/internal/node/gate_test.go b/internal/node/gate_test.go index 121ca3d3..1df88817 100644 --- a/internal/node/gate_test.go +++ b/internal/node/gate_test.go @@ -93,10 +93,10 @@ func TestConnectionGater(t *testing.T) { } } -// The revocation cache is seeded from the control plane's ban set at startup -// (Options.BannedPeerIDs, filled by SyncMeshConfig). Without that a restarted -// node would enforce no ban at all until the next MeshEvent_BANNED, which for -// a ban published while it was down never arrives. +// The revocation cache is filled from the control plane's ban set by the +// pre-start pull (SyncControlPlane -> reconcileBannedPeers). Without that a +// restarted node would enforce no ban at all until the next MeshEvent_BANNED, +// which for a ban published while it was down never arrives. func TestGaterEnforcesSeededBans(t *testing.T) { dir := t.TempDir() store, err := NewStore(dir) @@ -117,20 +117,22 @@ func TestGaterEnforcesSeededBans(t *testing.T) { if err != nil { t.Fatal(err) } - otherPriv, _, _ := crypto.GenerateEd25519Key(nil) + otherPriv, _, err := crypto.GenerateEd25519Key(nil) + if err != nil { + t.Fatal(err) + } allowed, err := peer.IDFromPrivateKey(otherPriv) if err != nil { t.Fatal(err) } - node, err := NewSamNode(Options{ - PrivKey: bannedPriv, - Store: store, - BannedPeerIDs: []string{peer.ToCid(banned).String(), "not-a-peer-id"}, - }) + node, err := NewSamNode(Options{PrivKey: bannedPriv, Store: store}) if err != nil { - t.Fatalf("a ban set with an undecodable entry must not fail startup: %v", err) + t.Fatal(err) } + // The wire form may be any encoding of the peer ID, and one bad entry + // must not stop the rest from being enforced. + node.reconcileBannedPeers([]string{peer.ToCid(banned).String(), "not-a-peer-id"}, time.Now()) gater := &nodeConnGate{node: node} if gater.InterceptPeerDial(banned) { diff --git a/internal/node/identity_evidence.go b/internal/node/identity_evidence.go index cf47474a..f9a40fda 100644 --- a/internal/node/identity_evidence.go +++ b/internal/node/identity_evidence.go @@ -296,8 +296,8 @@ func (n *SamNode) buildPeerEvidence(requested peer.ID, observation peerBiscuitOb } // peerIsRevoked reports whether the peer is in the revocation cache, which is -// seeded from the control plane's ban set at startup (see SyncMeshConfig) and -// updated by MeshEvent_BANNED. +// reconciled against the control plane's ban set by every SyncControlPlane +// and updated by MeshEvent_BANNED. func (n *SamNode) peerIsRevoked(peerID peer.ID) bool { return n.revokedPeers != nil && n.revokedPeers.Contains(peerID.String()) } diff --git a/internal/node/node.go b/internal/node/node.go index 9c4109fb..e726f118 100644 --- a/internal/node/node.go +++ b/internal/node/node.go @@ -37,6 +37,7 @@ import ( "github.com/biscuit-auth/biscuit-go/v2" "github.com/google/sam/api" + cpclient "github.com/google/sam/internal/controlplane/client" "github.com/google/sam/internal/identity" samdiscovery "github.com/google/sam/internal/node/discovery" "github.com/google/sam/internal/ratelimit" @@ -321,17 +322,6 @@ func NewSamNode(cfg Options) (*SamNode, error) { if err != nil { return nil, fmt.Errorf("failed to create revocation cache: %w", err) } - // Seed from the control plane's ban set so the gater enforces existing - // bans from the first connection, instead of waiting for an event that - // was already published while this node was down. - for _, id := range cfg.BannedPeerIDs { - p, err := peer.Decode(id) - if err != nil { - logger.Warnf("Ignoring undecodable banned peer ID %q from the control plane: %v", id, err) - continue - } - node.revokedPeers.Add(p.String(), time.Now().UnixMilli()) - } node.peerLabelGate, err = lru.New[string, time.Time](labelGateCacheSize) if err != nil { return nil, fmt.Errorf("failed to create label gate cache: %w", err) @@ -1189,9 +1179,9 @@ func (n *SamNode) RefreshEnrollment(ctx context.Context) error { } } - respData, err := io.ReadAll(io.LimitReader(resp.Body, maxControlPlaneBodyBytes)) + respData, err := cpclient.ReadBody(resp.Body) if err != nil { - return fmt.Errorf("failed to read response: %w", err) + return fmt.Errorf("refresh response: %w", err) } var refreshResp api.TokenRefreshResponse diff --git a/internal/node/node_test.go b/internal/node/node_test.go index 5459f536..2352db3f 100644 --- a/internal/node/node_test.go +++ b/internal/node/node_test.go @@ -728,41 +728,6 @@ func TestNewSamNode_BiscuitTimeout(t *testing.T) { }) } -func TestNewSamNode_BannedPeerCanonicalisation(t *testing.T) { - bannedPriv, _, err := crypto.GenerateEd25519Key(nil) - if err != nil { - t.Fatalf("failed to generate key: %v", err) - } - p, err := peer.IDFromPrivateKey(bannedPriv) - if err != nil { - t.Fatalf("failed to derive peer ID: %v", err) - } - canonicalID := p.String() - cidv1ID := peer.ToCid(p).String() - - priv, _, err := crypto.GenerateKeyPair(crypto.Ed25519, -1) - if err != nil { - t.Fatalf("failed to generate node key: %v", err) - } - store, err := NewStore(t.TempDir()) - if err != nil { - t.Fatal(err) - } - defer func() { _ = store.Close() }() - - node, err := NewSamNode(Options{ - PrivKey: priv, - Store: store, - BannedPeerIDs: []string{cidv1ID}, - }) - if err != nil { - t.Fatalf("NewSamNode: %v", err) - } - if !node.revokedPeers.Contains(canonicalID) { - t.Errorf("revokedPeers missing canonical ID %q (seeded with %q)", canonicalID, cidv1ID) - } -} - func TestNewSamNode_DHTOptions(t *testing.T) { priv, _, _ := crypto.GenerateKeyPair(crypto.Ed25519, -1) store, _ := NewStore(t.TempDir()) diff --git a/internal/node/options.go b/internal/node/options.go index 40e7ef79..ed19fee8 100644 --- a/internal/node/options.go +++ b/internal/node/options.go @@ -41,11 +41,6 @@ type Options struct { RouterAddrs []multiaddr.Multiaddr Store *Store - // BannedPeerIDs seeds the revocation cache from the control plane's ban - // set (see SyncMeshConfig). Without it a restarted node would enforce no - // ban until the next MeshEvent_BANNED, which for an existing ban never - // comes. - BannedPeerIDs []string MeshID string DiscoveryInterval string ListenAddrs []string diff --git a/internal/node/store.go b/internal/node/store.go index 0356ea1a..4d039d9a 100644 --- a/internal/node/store.go +++ b/internal/node/store.go @@ -307,5 +307,6 @@ func (s *Store) Close() error { // Peer bans are deliberately not kept here. A ban on disk cannot be undone by // the control plane -- there is no unban event -- and it says nothing about a // node that was offline when the ban was published. Both are handled instead by -// reconciling against the ban set in /info on every start (see SyncMeshConfig), -// with MeshEvent_BANNED as the sub-second path for nodes that are already up. +// reconciling against the ban set in /info before start and on every sync +// (see SyncControlPlane), with MeshEvent_BANNED as the sub-second path for +// nodes that are already up. diff --git a/internal/router/router.go b/internal/router/router.go index a8f4a470..b2364023 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -36,6 +36,7 @@ import ( "time" "github.com/google/sam/api" + cpclient "github.com/google/sam/internal/controlplane/client" "github.com/google/sam/internal/identity" "github.com/google/sam/internal/ratelimit" golog "github.com/ipfs/go-log/v2" @@ -443,9 +444,9 @@ func (r *Router) enroll(peerID peer.ID) error { return fmt.Errorf("enrollment response status %s: %s", resp.Status, string(body)) } - body, err := io.ReadAll(resp.Body) + body, err := cpclient.ReadBody(resp.Body) if err != nil { - return err + return fmt.Errorf("enrollment response: %w", err) } var enrollResp api.EnrollResponse @@ -504,9 +505,9 @@ func (r *Router) enrollBootstrap(peerID peer.ID) error { return fmt.Errorf("bootstrap enrollment request response status %s: %s", resp.Status, string(body)) } - respData, err := io.ReadAll(resp.Body) + respData, err := cpclient.ReadBody(resp.Body) if err != nil { - return err + return fmt.Errorf("bootstrap enrollment response: %w", err) } enrollResp := &api.BootstrapEnrollResponse{} @@ -565,10 +566,10 @@ func (r *Router) enrollBootstrap(peerID peer.ID) error { continue } - statusBody, err := io.ReadAll(statusResp.Body) + statusBody, err := cpclient.ReadBody(statusResp.Body) _ = statusResp.Body.Close() if err != nil { - logger.Warnf("failed to read status body: %v", err) + logger.Warnf("enrollment status response: %v", err) continue } @@ -679,32 +680,11 @@ func (r *Router) recoverAfterLease401() error { } func (r *Router) syncKeys() error { - client := r.controlPlaneClient(10 * time.Second) - resp, err := client.Get(r.config.ControlPlaneURL + "/keys") - if err != nil { - return err - } - defer func() { _ = resp.Body.Close() }() - - if resp.StatusCode != http.StatusOK { - return fmt.Errorf("failed to sync keys, status %s", resp.Status) - } - - body, err := io.ReadAll(resp.Body) - if err != nil { - return err - } - - var keysResp api.KeysResponse - if err := proto.Unmarshal(body, &keysResp); err != nil { - return err - } - // Only a set signed by a key this router already trusts may replace // the trust set; anything else is whoever answered the URL. - newKeys, err := api.VerifyKeysResponse(&keysResp, r.getTrustedPublicKeys(), time.Now()) + newKeys, err := r.controlPlane(10*time.Second).FetchKeys(r.ctx, r.getTrustedPublicKeys()) if err != nil { - return fmt.Errorf("/keys response rejected: %w", err) + return err } r.keysMu.Lock() @@ -715,24 +695,17 @@ func (r *Router) syncKeys() error { return nil } +// controlPlane reads the pull endpoints of the control plane. +func (r *Router) controlPlane(timeout time.Duration) *cpclient.Client { + return cpclient.New(r.config.ControlPlaneURL, r.controlPlaneClient(timeout)) +} + // controlPlaneClient is the client for every request to the control plane; // its transport re-checks the plaintext policy on each hop, redirects included. func (r *Router) controlPlaneClient(timeout time.Duration) *http.Client { - return &http.Client{ - Timeout: timeout, - Transport: roundTripperFunc(func(req *http.Request) (*http.Response, error) { - if err := api.ValidateControlPlaneTransport(req.URL.String(), r.config.AllowInsecureControlPlane); err != nil { - return nil, err - } - return http.DefaultTransport.RoundTrip(req) - }), - } + return cpclient.NewHTTPClient(timeout, func() bool { return r.config.AllowInsecureControlPlane }) } -type roundTripperFunc func(*http.Request) (*http.Response, error) - -func (f roundTripperFunc) RoundTrip(req *http.Request) (*http.Response, error) { return f(req) } - func (r *Router) getTrustedPublicKeys() []ed25519.PublicKey { r.keysMu.RLock() defer r.keysMu.RUnlock() @@ -928,8 +901,13 @@ func (r *Router) renewLease() { return } - body, _ := io.ReadAll(resp.Body) + body, readErr := cpclient.ReadBody(resp.Body) _ = resp.Body.Close() + if readErr != nil { + logger.Errorf("Control plane lease renewal response: %v", readErr) + leaseRenewalsTotal.WithLabelValues(leaseRejected).Inc() + return + } if resp.StatusCode == http.StatusUnauthorized && attempt == 0 { logger.Warnf("Control plane lease renewal rejected (401 Unauthorized: %s), attempting recovery...", string(body)) @@ -1045,32 +1023,15 @@ func (r *Router) runFederationLoop() { } func (r *Router) connectBootstrapRouters() { - client := r.controlPlaneClient(10 * time.Second) // Taken before the request: anything banned after this point cannot be // reflected in the answer, so reconciliation must not read its absence as // an unban (see reconcileBannedPeers). fetchedAt := time.Now() - resp, err := client.Get(r.config.ControlPlaneURL + "/info") + info, err := r.controlPlane(10 * time.Second).FetchInfo(r.ctx) if err != nil { logger.Errorf("[Federation] Failed to fetch router info: %v", err) return } - defer func() { _ = resp.Body.Close() }() - - if resp.StatusCode != http.StatusOK { - logger.Errorf("[Federation] GET /info returned status %d", resp.StatusCode) - return - } - - body, err := io.ReadAll(resp.Body) - if err != nil { - return - } - - var info api.ControlPlaneInfoResponse - if err := proto.Unmarshal(body, &info); err != nil { - return - } // /info carries the whole ban set, so this is also where a router that // restarted or missed a MeshEvent_BANNED catches up. @@ -1406,9 +1367,9 @@ func (r *Router) RefreshEnrollment(ctx context.Context) error { return fmt.Errorf("refresh failed with status %s: %s", resp.Status, string(body)) } - respData, err := io.ReadAll(resp.Body) + respData, err := cpclient.ReadBody(resp.Body) if err != nil { - return fmt.Errorf("failed to read response: %w", err) + return fmt.Errorf("refresh response: %w", err) } var refreshResp api.TokenRefreshResponse diff --git a/internal/router/trust_test.go b/internal/router/trust_test.go index a2ef260b..f0f6df7d 100644 --- a/internal/router/trust_test.go +++ b/internal/router/trust_test.go @@ -96,7 +96,7 @@ func TestSyncKeysRequiresTrustedSignature(t *testing.T) { defer srv.Close() newRouter := func() *Router { - return &Router{config: Options{ControlPlaneURL: srv.URL}, trustedPublicKeys: []ed25519.PublicKey{oldPub}} + return &Router{ctx: context.Background(), config: Options{ControlPlaneURL: srv.URL}, trustedPublicKeys: []ed25519.PublicKey{oldPub}} } marshal := func(resp *api.KeysResponse) []byte { data, err := proto.Marshal(resp) diff --git a/mobile/sam-node-ffi/ffi/ffi.go b/mobile/sam-node-ffi/ffi/ffi.go index 62eaefea..2075c128 100644 --- a/mobile/sam-node-ffi/ffi/ffi.go +++ b/mobile/sam-node-ffi/ffi/ffi.go @@ -153,11 +153,16 @@ func StartNode(configJSON string) error { var controlPlanePubKey ed25519.PublicKey var routerAddrs []multiaddr.Multiaddr - // Sync config from stored/synced configuration - storedPubKey, syncedAddrs, bannedPeerIDs, err := node.SyncMeshConfig(context.Background(), store) + // What the last run left behind; the node pulls what the control plane + // knows now before it starts. + storedPubKey, storedAddrs, err := store.LoadMeshConfig() if err == nil && len(storedPubKey) > 0 { controlPlanePubKey = storedPubKey - routerAddrs = syncedAddrs + for _, s := range storedAddrs { + if ma, parseErr := multiaddr.NewMultiaddr(s); parseErr == nil { + routerAddrs = append(routerAddrs, ma) + } + } } priv := node.GetOrGenerateKey(store) @@ -208,7 +213,6 @@ func StartNode(configJSON string) error { ControlPlanePubKey: controlPlanePubKey, RouterAddrs: routerAddrs, Store: store, - BannedPeerIDs: bannedPeerIDs, MeshID: meshID, DiscoveryInterval: discoveryInterval, ListenAddrs: listenAddrs, @@ -230,6 +234,9 @@ func StartNode(configJSON string) error { ctx, cancel := context.WithCancel(context.Background()) cancelFunc = cancel + if err := samNode.SyncControlPlane(ctx); err != nil { + logger.Warnf("Control plane sync before start failed (using stored config): %v", err) + } if err := samNode.Start(ctx); err != nil { cancel() _ = store.Close() @@ -443,11 +450,6 @@ func enrollWith(dataDir string, controlPlaneURL string, allowLoopback bool, labe return fmt.Errorf("failed to save control plane URL: %w", err) } - _, _, _, err = node.SyncMeshConfig(enrollCtx, store) - if err != nil { - return fmt.Errorf("failed to sync mesh config post-enrollment: %w", err) - } - return nil } diff --git a/tests/integration/fallback_test.go b/tests/integration/fallback_test.go index af0faac8..f8ffb9cf 100644 --- a/tests/integration/fallback_test.go +++ b/tests/integration/fallback_test.go @@ -19,6 +19,7 @@ import ( "crypto/ed25519" "crypto/rand" "encoding/json" + "fmt" "io" "os" @@ -26,6 +27,7 @@ import ( "path/filepath" "strings" "sync" + "sync/atomic" "testing" "time" @@ -65,22 +67,33 @@ func (s *safeBuffer) String() string { return s.buf.String() } +// TestSelfHealingHTTPFallback: the routers a node stored can be gone by its +// next start (a redeploy moves every router's address); the node must ask the +// control plane for the current ones and reach one. The assertions are the two +// ends of that path: the control plane saw the request, and the router that +// only exists since the address change completed an auth handshake with the +// node. Neither depends on what the node logs. func TestSelfHealingHTTPFallback(t *testing.T) { nodeBin := buildBinary(t, "./cmd/sam-node") var mu sync.Mutex var currentP2PAddr string + var infoRequests atomic.Int32 pub, priv, err := ed25519.GenerateKey(rand.Reader) if err != nil { t.Fatalf("Failed to generate control plane key: %v", err) } - createNewHost := func() host.Host { + // createNewHost is a router as the node sees it: the auth handshake, and a + // channel closed once a node has completed it. + createNewHost := func() (host.Host, <-chan struct{}) { newH, err := libp2p.New(libp2p.ListenAddrStrings("/ip4/127.0.0.1/tcp/0")) if err != nil { t.Fatal(err) } + authenticated := make(chan struct{}) + var once sync.Once newH.SetStreamHandler(api.AuthProtocolID, func(s network.Stream) { defer func() { _ = s.Close() }() reader := msgio.NewVarintReaderSize(s, 1024*64) @@ -95,13 +108,21 @@ func TestSelfHealingHTTPFallback(t *testing.T) { Success: true, Biscuit: createMockBiscuitToken(t, newH.ID().String(), priv, api.RoleRouter, nil), } - respBytes, _ := proto.Marshal(resp) - _ = writer.WriteMsg(respBytes) + respBytes, err := proto.Marshal(resp) + if err != nil { + t.Errorf("marshal auth response: %v", err) + return + } + if err := writer.WriteMsg(respBytes); err != nil { + t.Errorf("write auth response: %v", err) + return + } + once.Do(func() { close(authenticated) }) }) - return newH + return newH, authenticated } - h := createNewHost() + h, _ := createNewHost() defer func() { _ = h.Close() }() mu.Lock() @@ -159,11 +180,19 @@ func TestSelfHealingHTTPFallback(t *testing.T) { ControlPlanePublicKey: pub, RouterAddresses: []string{addr}, } - data, _ := proto.Marshal(resp) + data, err := proto.Marshal(resp) + if err != nil { + t.Errorf("marshal /register: %v", err) + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } w.Header().Set("Content-Type", "application/x-protobuf") - _, _ = w.Write(data) + if _, err := w.Write(data); err != nil { + t.Errorf("write /register: %v", err) + } }) mux.HandleFunc("/info", func(w http.ResponseWriter, r *http.Request) { + infoRequests.Add(1) mu.Lock() addr := currentP2PAddr mu.Unlock() @@ -173,9 +202,16 @@ func TestSelfHealingHTTPFallback(t *testing.T) { Audience: "sam-mesh-audience", RouterAddresses: []string{addr}, } - data, _ := proto.Marshal(resp) + data, err := proto.Marshal(resp) + if err != nil { + t.Errorf("marshal /info: %v", err) + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } w.Header().Set("Content-Type", "application/x-protobuf") - _, _ = w.Write(data) + if _, err := w.Write(data); err != nil { + t.Errorf("write /info: %v", err) + } }) httpServer := httptest.NewServer(mux) @@ -208,43 +244,54 @@ func TestSelfHealingHTTPFallback(t *testing.T) { t.Fatalf("Join did not succeed:\n%s", out) } - // Step 2: Simulate router changing its P2P port (HTTP URL stays the same) + // Step 2: Simulate router changing its P2P port (HTTP URL stays the same). + // The stored address now points at nothing; only /info knows the new one. _ = h.Close() - h = createNewHost() - defer func() { _ = h.Close() }() + newRouter, authenticated := createNewHost() + defer func() { _ = newRouter.Close() }() mu.Lock() - currentP2PAddr = h.Addrs()[0].String() + "/p2p/" + h.ID().String() + currentP2PAddr = newRouter.Addrs()[0].String() + "/p2p/" + newRouter.ID().String() mu.Unlock() + infoRequests.Store(0) - // Step 3: Start sam-node run - runCmd := exec.Command(nodeBin, "run", "--listen", "/ip4/127.0.0.1/tcp/0", "--bind-addr", "127.0.0.1:0") + // Step 3: Start sam-node run, with everything it needs to stay up: a + // node that reaches the router and then exits has not healed. + runCmd := exec.Command(nodeBin, "run", + "--listen", "/ip4/127.0.0.1/tcp/0", + "--bind-addr", fmt.Sprintf("127.0.0.1:%d", getFreePort(t)), + "--api-token-path", tokenPath(t, "fallback-test-token"), + ) runCmd.Env = env - var stdout safeBuffer - runCmd.Stdout = &stdout - runCmd.Stderr = &stdout + var output safeBuffer + runCmd.Stdout = &output + runCmd.Stderr = &output if err := runCmd.Start(); err != nil { t.Fatal(err) } + exited := make(chan error, 1) + go func() { exited <- runCmd.Wait() }() defer func() { _ = runCmd.Process.Kill() - _ = runCmd.Wait() + <-exited }() - // Wait for successful fallback - success := false - for i := 0; i < 50; i++ { - out = stdout.String() - if strings.Contains(out, "Fetching latest router addresses via HTTP") && - strings.Contains(out, "Successfully authenticated with router via libp2p") { - success = true - break - } - time.Sleep(100 * time.Millisecond) + select { + case <-authenticated: + case err := <-exited: + t.Fatalf("sam-node run exited (%v) before authenticating with the moved router.\nOutput:\n%s", err, output.String()) + case <-time.After(5 * time.Second): + t.Fatalf("sam-node run never authenticated with the moved router.\nOutput:\n%s", output.String()) + } + if infoRequests.Load() == 0 { + t.Fatal("the node reached the moved router without asking the control plane for its address") } - if !success { - t.Fatalf("Failed to detect self-healing HTTP fallback in output.\nOutput:\n%s", out) + // Reaching the router is not the end: the node must stay up on it. + select { + case err := <-exited: + t.Fatalf("sam-node run exited (%v) right after authenticating.\nOutput:\n%s", err, output.String()) + case <-time.After(500 * time.Millisecond): } }