diff --git a/lib/join/oraclejoin/join_test.go b/lib/join/oraclejoin/join_test.go index 6a08c50fec3..757c1063178 100644 --- a/lib/join/oraclejoin/join_test.go +++ b/lib/join/oraclejoin/join_test.go @@ -416,6 +416,8 @@ func TestInstanceKeyAlgorithms(t *testing.T) { return assert.ErrorAs(t, err, new(*trace.AccessDeniedError), msgAndArgs...) } + rootCACache := oraclejoin.NewRootCACache() + for _, tc := range []struct { desc string instanceKey crypto.Signer @@ -475,6 +477,7 @@ func TestInstanceKeyAlgorithms(t *testing.T) { Solution: solution, ProvisionToken: token, HTTPClient: fakeOracleAPIClient, + RootCACache: rootCACache, } _, err = oraclejoin.CheckChallengeSolution(t.Context(), params) @@ -604,6 +607,7 @@ func (f *fakeOracleAPI) handleRootCACertificates(w http.ResponseWriter, r *http. } resp := rootCAResp{ Certificates: []string{f.rootCABase64}, + RefreshIn: time.Now().Add(time.Hour).Format(time.RFC3339), } json.NewEncoder(w).Encode(resp) } diff --git a/lib/join/oraclejoin/oracle.go b/lib/join/oraclejoin/oracle.go index 734b45ad2f4..0b4d08dafc6 100644 --- a/lib/join/oraclejoin/oracle.go +++ b/lib/join/oraclejoin/oracle.go @@ -30,6 +30,7 @@ import ( "io" "net/http" "strings" + "time" "github.com/gravitational/trace" "github.com/oracle/oci-go-sdk/v65/common" @@ -50,6 +51,8 @@ type CheckChallengeSolutionParams struct { Solution *messages.OracleChallengeSolution // ProvisionToken is the token being used for the request. ProvisionToken provision.Token + // RootCACache caches Oracle root CAs per region. + RootCACache *RootCACache // HTTPClient (optional) is an HTTP client that will be used to send // requests to the Oracle API. HTTPClient utils.HTTPDoClient @@ -69,6 +72,8 @@ func (p *CheckChallengeSolutionParams) checkAndSetDefaults() error { return trace.BadParameter("Signature is required") case len(p.Solution.SignedRootCAReq) == 0: return trace.BadParameter("SignedRootCAReq is required") + case p.RootCACache == nil: + return trace.BadParameter("RootCACache is required") } if p.HTTPClient == nil { httpClient, err := defaults.HTTPClient() @@ -207,23 +212,20 @@ func makeIntermediateCAPool(intermediateCAPEM []byte) (*x509.CertPool, error) { } func makeRootCAPool(ctx context.Context, params *CheckChallengeSolutionParams, instanceID string) (*x509.CertPool, error) { - // TODO(nklaassen): considering caching the root CA pool per region. rootCAReq, err := parseRootCAReq(params.Solution.SignedRootCAReq) if err != nil { return nil, trace.Wrap(err, "parsing signed root CA request") } - if err := validateRootCAReq(rootCAReq, instanceID); err != nil { - return nil, trace.Wrap(err) - } - - resp, err := executeRootCAReq(ctx, params.HTTPClient, rootCAReq) + region, err := validateRootCAReq(rootCAReq, instanceID) if err != nil { - return nil, trace.Wrap(err, "fetching root CAs from Oracle API") + return nil, trace.Wrap(err, "validating root CA request") } - rootCAPool, err := parseRootCAPool(resp) - return rootCAPool, trace.Wrap(err, "parsing Oracle root CA pool") + rootCAPool, err := params.RootCACache.get(ctx, region, func() (*x509.CertPool, time.Time, error) { + return getRootCAPool(ctx, params.HTTPClient, rootCAReq) + }) + return rootCAPool, trace.Wrap(err) } func parseRootCAReq(req []byte) (*http.Request, error) { @@ -243,20 +245,20 @@ func parseRootCAReq(req []byte) (*http.Request, error) { return httpReq, nil } -func validateRootCAReq(req *http.Request, instanceID string) error { +func validateRootCAReq(req *http.Request, instanceID string) (string, error) { const rootCAPath = "/v1/instancePrincipalRootCACertificates" if req.URL.Path != rootCAPath { - return trace.BadParameter("path must be %s, got %s", rootCAPath, req.URL.Path) + return "", trace.BadParameter("path must be %s, got %s", rootCAPath, req.URL.Path) } expectedRegion, err := regionFromInstanceID(instanceID) if err != nil { - return trace.Wrap(err) + return "", trace.Wrap(err) } expectedHostname := expectedRegion.Endpoint("auth") if req.URL.Hostname() != expectedHostname { - return trace.BadParameter("hostname must be %s, got %s", expectedHostname, req.URL.Hostname()) + return "", trace.BadParameter("hostname must be %s, got %s", expectedHostname, req.URL.Hostname()) } - return nil + return string(expectedRegion), nil } // regionFromInstanceID returns an oci region for a given instance ID. It will @@ -282,6 +284,22 @@ func regionFromInstanceID(instanceID string) (common.Region, error) { return region, nil } +func getRootCAPool(ctx context.Context, httpClient utils.HTTPDoClient, rootCAReq *http.Request) (*x509.CertPool, time.Time, error) { + resp, err := executeRootCAReq(ctx, httpClient, rootCAReq) + if err != nil { + return nil, time.Time{}, trace.Wrap(err, "executing Oracle root CA request") + } + rootCAPool, err := parseRootCAPool(resp) + if err != nil { + return nil, time.Time{}, trace.Wrap(err, "parsing Oracle root CA pool") + } + expires, err := time.Parse(time.RFC3339, resp.RefreshIn) + if err != nil { + expires = time.Now().Add(5 * time.Minute) + } + return rootCAPool, expires, nil +} + func executeRootCAReq(ctx context.Context, client utils.HTTPDoClient, req *http.Request) (*rootCAResp, error) { req = req.WithContext(ctx) resp, err := client.Do(req) diff --git a/lib/join/oraclejoin/rootcacache.go b/lib/join/oraclejoin/rootcacache.go new file mode 100644 index 00000000000..903dca436df --- /dev/null +++ b/lib/join/oraclejoin/rootcacache.go @@ -0,0 +1,112 @@ +// Teleport +// Copyright (C) 2025 Gravitational, Inc. +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . + +package oraclejoin + +import ( + "context" + "crypto/x509" + "sync" + "time" +) + +// RootCACache caches Oracle instance root CAs for each region. It is inspired +// by lib/utils.FnCache which is designed for caching backend reads, but this +// cache makes a couple different choices: +// - Each cache entry expiry is determined by the response we get from the +// Oracle API and unknown before making the request. +// - Error results are always retried at the next request. +// - There is no regularly scheduled cleanup, expired entries will be cleaned +// up on the next get. This cache will not be large, max 1 entry per OCI region. +type RootCACache struct { + entries map[string]*rootCACacheEntry + mu sync.Mutex +} + +// NewRootCACache returns a new RootCACache, ready to use. +func NewRootCACache() *RootCACache { + return &RootCACache{ + entries: make(map[string]*rootCACacheEntry), + } +} + +type getRootCAPoolFn func() (*x509.CertPool, time.Time, error) + +// get is the only method consumers of RootCACache should call. Given a key +// (which should be the region of the root CA) and a getRootCAPoolFn, it will +// return a cached root CA pool if one is present and not expired, else it will +// call loadFn to request the root CA pool. +func (c *RootCACache) get(ctx context.Context, key string, loadFn getRootCAPoolFn) (*x509.CertPool, error) { + entry := c.getEntry(key, loadFn) + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-entry.loaded: + return entry.caPool, entry.err + } +} + +func (c *RootCACache) getEntry(key string, loadFn getRootCAPoolFn) *rootCACacheEntry { + c.mu.Lock() + defer c.mu.Unlock() + + now := time.Now() + + c.removeExpiredLocked(now) + + if entry, ok := c.entries[key]; ok { + // There is an existing cache entry, return it. If the entry is expired + // or stores an error it would have been removed in removeExpiredLocked + // above. + return entry + } + + // Create a new cache entry to store and return. + entry := &rootCACacheEntry{ + loaded: make(chan struct{}), + } + c.entries[key] = entry + go func() { + entry.caPool, entry.expires, entry.err = loadFn() + close(entry.loaded) + }() + + return entry +} + +func (c *RootCACache) removeExpiredLocked(now time.Time) { + for key, entry := range c.entries { + select { + case <-entry.loaded: + if entry.expires.Before(now) || entry.err != nil { + delete(c.entries, key) + } + default: + // entry is still being loaded. + continue + } + } +} + +type rootCACacheEntry struct { + // loaded will be closed when the entry is ready. It also acts as a memory barrier: + // the below fields must only be written before loaded is closed, and must + // only be read after it is closed. + loaded chan struct{} + caPool *x509.CertPool + expires time.Time + err error +} diff --git a/lib/join/oraclejoin/rootcacache_test.go b/lib/join/oraclejoin/rootcacache_test.go new file mode 100644 index 00000000000..238189cedb5 --- /dev/null +++ b/lib/join/oraclejoin/rootcacache_test.go @@ -0,0 +1,108 @@ +// Teleport +// Copyright (C) 2025 Gravitational, Inc. +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . + +package oraclejoin + +import ( + "crypto/x509" + "errors" + "strconv" + "sync" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestRootCACache(t *testing.T) { + t.Parallel() + synctest.Test(t, testRootCACache) +} + +func testRootCACache(t *testing.T) { + cache := NewRootCACache() + + caPool := x509.NewCertPool() + const caTTL = 5 * time.Minute + + var loadCount atomic.Int32 + loadFn := func() (*x509.CertPool, time.Time, error) { + loadCount.Add(1) + time.Sleep(10 * time.Millisecond) + expires := time.Now().Add(caTTL) + return caPool, expires, nil + } + + // Getting the same key many times only results in 1 load. + for range 10 { + pool, err := cache.get(t.Context(), "test", loadFn) + require.Equal(t, caPool, pool) + require.NoError(t, err) + } + require.Equal(t, int32(1), loadCount.Load()) + + // Waiting until the entry expires then getting the same key again a number + // of times results in 1 additional load. + time.Sleep(caTTL + time.Minute) + for range 10 { + _, err := cache.get(t.Context(), "test", loadFn) + require.NoError(t, err) + } + require.Equal(t, int32(2), loadCount.Load()) + + // Getting 10 different keys multiple times results in 10 total loads. + loadCount.Store(0) + for range 10 { + for i := range 10 { + _, err := cache.get(t.Context(), strconv.Itoa(i), loadFn) + require.NoError(t, err) + } + } + require.Equal(t, int32(10), loadCount.Load()) + + // Wait until all current entries expire, then do 1 more load, expired + // entries should be cleaned up. + time.Sleep(caTTL + time.Minute) + _, err := cache.get(t.Context(), "test", loadFn) + require.NoError(t, err) + require.Len(t, cache.entries, 1) + + // Error values should not be cached. + _, err = cache.get(t.Context(), "testerror", func() (*x509.CertPool, time.Time, error) { + return nil, time.Time{}, errors.New("test error") + }) + require.Error(t, err) + // An immediate re-fetch does not get the error. + _, err = cache.get(t.Context(), "testerror", loadFn) + require.NoError(t, err) + + // Even with many concurrent calls the value will only be loaded once. + loadCount.Store(0) + var wg sync.WaitGroup + wg.Add(100) + for range 100 { + go func() { + _, err := cache.get(t.Context(), "concurrent", loadFn) + assert.NoError(t, err) + wg.Done() + }() + } + wg.Wait() + require.Equal(t, int32(1), loadCount.Load()) +} diff --git a/lib/join/server.go b/lib/join/server.go index 24d7b8a1dcc..256875be6e4 100644 --- a/lib/join/server.go +++ b/lib/join/server.go @@ -50,6 +50,7 @@ import ( "github.com/gravitational/teleport/lib/join/internal/diagnostic" "github.com/gravitational/teleport/lib/join/internal/messages" "github.com/gravitational/teleport/lib/join/joinutils" + "github.com/gravitational/teleport/lib/join/oraclejoin" "github.com/gravitational/teleport/lib/join/provision" "github.com/gravitational/teleport/lib/scopes/joining" "github.com/gravitational/teleport/lib/services" @@ -94,13 +95,15 @@ type ServerConfig struct { // Server implements cluster joining for nodes and bots. type Server struct { - cfg *ServerConfig + cfg *ServerConfig + oracleRootCACache *oraclejoin.RootCACache } // NewServer returns a new [Server] instance. func NewServer(cfg *ServerConfig) *Server { return &Server{ - cfg: cfg, + cfg: cfg, + oracleRootCACache: oraclejoin.NewRootCACache(), } } diff --git a/lib/join/server_oracle.go b/lib/join/server_oracle.go index df25d76e5d5..b1eccb8b4be 100644 --- a/lib/join/server_oracle.go +++ b/lib/join/server_oracle.go @@ -80,6 +80,7 @@ func (s *Server) handleOracleJoin( Solution: solution, ProvisionToken: provisionToken, HTTPClient: s.cfg.OracleHTTPClient, + RootCACache: s.oracleRootCACache, }) // CheckOracleRequest may return claims even when returning an error, which // may aid in debugging.