cache oracle root CAs (#60858)

This commit is contained in:
Nic Klaassen
2025-11-03 18:42:18 +00:00
committed by GitHub
parent d4e6f1974e
commit b551d373ed
6 changed files with 262 additions and 16 deletions
+4
View File
@@ -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)
}
+32 -14
View File
@@ -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)
+112
View File
@@ -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 <http://www.gnu.org/licenses/>.
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
}
+108
View File
@@ -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 <http://www.gnu.org/licenses/>.
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())
}
+5 -2
View File
@@ -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(),
}
}
+1
View File
@@ -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.