mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
cache oracle root CAs (#60858)
This commit is contained in:
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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(),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user