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.