Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions e2e/harness_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import (
"testing"

"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/promauto"
"github.com/stretchr/testify/require"
"go.temporal.io/api/workflowservice/v1"
"go.uber.org/fx"
Expand Down Expand Up @@ -124,6 +125,7 @@ func newProxyApp(t *testing.T, cfg *config.Config) *fx.App {
fx.Supply(fx.Annotate(t.Context(), fx.As(new(context.Context)))),
fx.Supply(cfg),
fx.Provide(func() *crypto.Vault { return nil }),
fx.Provide(func() *metrics.Factory { return metrics.New("test", promauto.With(prometheus.NewRegistry())) }),
connect.Module,
protoutil.Module,
proxy.Module,
Expand Down
29 changes: 19 additions & 10 deletions internal/kms/fx.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ import (
"go.uber.org/fx"

"github.com/temporalio/temporal-proxy/internal/config"
"github.com/temporalio/temporal-proxy/internal/metrics"
"github.com/temporalio/temporal-proxy/pkg/crypto"
"github.com/temporalio/temporal-proxy/pkg/logger"
"github.com/temporalio/temporal-proxy/pkg/logger/tag"
Expand All @@ -35,17 +36,20 @@ var Module = fx.Options(
return nil, nil
}

reporter := NewReporter(p.Factory.ForSubsystem("encryption"))

r, err := createKEKRegistry(
p.Context,
p.Lifecycle,
p.Config,
p.Logger,
reporter,
)
if err != nil {
return nil, err
}

v, err := createVault(p.Config, r)
v, err := createVault(p.Config, r, reporter)
if err != nil {
_ = r.Close()
return nil, err
Expand Down Expand Up @@ -83,6 +87,7 @@ type (
Config *config.Config
Lifecycle fx.Lifecycle
Logger logger.Logger
Factory *metrics.Factory
}

// vaultRefresher is the subset of *crypto.Vault the rotation loop depends
Expand Down Expand Up @@ -126,9 +131,12 @@ func runRotation(ctx context.Context, v vaultRefresher, interval time.Duration,
// createVault builds a vault from the registry, applying the configured cache
// size and, when a default key policy is set, its DEK duration and renewal lead
// time.
func createVault(c *config.Config, r *crypto.KEKRegistry) (*crypto.Vault, error) {
opts := make([]crypto.VaultOption, 0, 2+len(c.Encryption.Overrides))
opts = append(opts, crypto.WithCacheSize(c.Encryption.CacheSize))
func createVault(c *config.Config, r *crypto.KEKRegistry, reporter *Reporter) (*crypto.Vault, error) {
opts := make([]crypto.VaultOption, 0, 3+len(c.Encryption.Overrides))
opts = append(opts,
crypto.WithCacheSize(c.Encryption.CacheSize),
crypto.WithObserver(reporter),
)

if dp := c.Encryption.Default; dp != nil {
opts = append(opts, crypto.WithDefaultKeyConfig(crypto.KeyConfig{
Expand Down Expand Up @@ -158,10 +166,10 @@ func createVault(c *config.Config, r *crypto.KEKRegistry) (*crypto.Vault, error)
// createKEKRegistry opens the configured KEKs and assembles a registry,
// registering an fx OnStop hook that closes the registry (and its KEKs) on
// shutdown.
func createKEKRegistry(ctx context.Context, lc fx.Lifecycle, c *config.Config, logger logger.Logger) (*crypto.KEKRegistry, error) {
func createKEKRegistry(ctx context.Context, lc fx.Lifecycle, c *config.Config, logger logger.Logger, reporter *Reporter) (*crypto.KEKRegistry, error) {
opts := []crypto.KEKRegistryOption{}
if dp := c.Encryption.Default; dp != nil {
res, err := keyPolicyRegistryOpts(ctx, dp, logger, defaultNamespace, true)
res, err := keyPolicyRegistryOpts(ctx, dp, logger, defaultNamespace, true, reporter)
if err != nil {
return nil, err
}
Expand All @@ -173,7 +181,7 @@ func createKEKRegistry(ctx context.Context, lc fx.Lifecycle, c *config.Config, l
// logs and, on a partial failure, a repeatable point of failure).
for _, ns := range slices.Sorted(maps.Keys(c.Encryption.Overrides)) {
policy := c.Encryption.Overrides[ns]
res, err := keyPolicyRegistryOpts(ctx, &policy, logger, ns, false)
res, err := keyPolicyRegistryOpts(ctx, &policy, logger, ns, false, reporter)
if err != nil {
return nil, err
}
Expand Down Expand Up @@ -207,9 +215,10 @@ func keyPolicyRegistryOpts(
log logger.Logger,
ns string,
asDefault bool,
reporter *Reporter,
) ([]crypto.KEKRegistryOption, error) {
opts := []crypto.KEKRegistryOption{}
keys, err := createKEKs(ctx, p, log, ns)
keys, err := createKEKs(ctx, p, log, ns, reporter)
if err != nil {
return nil, err
}
Expand All @@ -232,7 +241,7 @@ func keyPolicyRegistryOpts(
// returning the KEKs in that order. The primary key is always element zero. If
// any key fails to open, every key opened so far is closed before returning, so
// a partial failure leaks no keepers.
func createKEKs(ctx context.Context, p *config.KeyPolicy, log logger.Logger, ns string) (_ []crypto.KEK, err error) {
func createKEKs(ctx context.Context, p *config.KeyPolicy, log logger.Logger, ns string, reporter *Reporter) (_ []crypto.KEK, err error) {
log = log.With(tag.String("namespace", ns))
keys := make([]crypto.KEK, 0, len(p.DecryptURIs)+1)

Expand All @@ -251,7 +260,7 @@ func createKEKs(ctx context.Context, p *config.KeyPolicy, log logger.Logger, ns
return err
}

keys = append(keys, k)
keys = append(keys, newMeteredKEK(k, providerForScheme(uri.Scheme), reporter))
return nil
}

Expand Down
39 changes: 27 additions & 12 deletions internal/kms/fx_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,11 +11,14 @@ import (
"testing/synctest"
"time"

"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/promauto"
"github.com/stretchr/testify/require"
"go.uber.org/fx"
"go.uber.org/fx/fxtest"

"github.com/temporalio/temporal-proxy/internal/config"
"github.com/temporalio/temporal-proxy/internal/metrics"
"github.com/temporalio/temporal-proxy/pkg/crypto"
"github.com/temporalio/temporal-proxy/pkg/logger"
)
Expand Down Expand Up @@ -88,7 +91,7 @@ func TestCreateKEKs(t *testing.T) {
t.Parallel()

p := &config.KeyPolicy{URI: primary, DecryptURIs: []url.URL{distinctKeyURL(t, 2), distinctKeyURL(t, 3)}}
keys, err := createKEKs(t.Context(), p, log, defaultNamespace)
keys, err := createKEKs(t.Context(), p, log, defaultNamespace, newFxTestReporter(t))
require.NoError(t, err)
t.Cleanup(func() { closeKEKs(t, keys) })

Expand All @@ -99,7 +102,7 @@ func TestCreateKEKs(t *testing.T) {
t.Run("primary only", func(t *testing.T) {
t.Parallel()

keys, err := createKEKs(t.Context(), &config.KeyPolicy{URI: primary}, log, defaultNamespace)
keys, err := createKEKs(t.Context(), &config.KeyPolicy{URI: primary}, log, defaultNamespace, newFxTestReporter(t))
require.NoError(t, err)
t.Cleanup(func() { closeKEKs(t, keys) })

Expand All @@ -109,15 +112,15 @@ func TestCreateKEKs(t *testing.T) {
t.Run("bad primary uri errors", func(t *testing.T) {
t.Parallel()

_, err := createKEKs(t.Context(), &config.KeyPolicy{URI: url.URL{Scheme: "bogus", Host: "x"}}, log, defaultNamespace)
_, err := createKEKs(t.Context(), &config.KeyPolicy{URI: url.URL{Scheme: "bogus", Host: "x"}}, log, defaultNamespace, newFxTestReporter(t))
require.Error(t, err)
})

t.Run("bad decrypt uri errors", func(t *testing.T) {
t.Parallel()

p := &config.KeyPolicy{URI: primary, DecryptURIs: []url.URL{{Scheme: "bogus", Host: "x"}}}
_, err := createKEKs(t.Context(), p, log, defaultNamespace)
_, err := createKEKs(t.Context(), p, log, defaultNamespace, newFxTestReporter(t))
require.Error(t, err)
})
}
Expand All @@ -134,7 +137,7 @@ func TestKeyPolicyRegistryOpts(t *testing.T) {
t.Run("asDefault registers a usable default key", func(t *testing.T) {
t.Parallel()

opts, err := keyPolicyRegistryOpts(t.Context(), &config.KeyPolicy{URI: distinctKeyURL(t, 1)}, log, defaultNamespace, true)
opts, err := keyPolicyRegistryOpts(t.Context(), &config.KeyPolicy{URI: distinctKeyURL(t, 1)}, log, defaultNamespace, true, newFxTestReporter(t))
require.NoError(t, err)

reg, err := crypto.NewKEKRegistry(opts...)
Expand All @@ -145,7 +148,7 @@ func TestKeyPolicyRegistryOpts(t *testing.T) {
t.Run("non-default registers only a namespace key", func(t *testing.T) {
t.Parallel()

opts, err := keyPolicyRegistryOpts(t.Context(), &config.KeyPolicy{URI: distinctKeyURL(t, 2)}, log, "other", false)
opts, err := keyPolicyRegistryOpts(t.Context(), &config.KeyPolicy{URI: distinctKeyURL(t, 2)}, log, "other", false, newFxTestReporter(t))
require.NoError(t, err)

_, err = crypto.NewKEKRegistry(opts...)
Expand All @@ -158,7 +161,7 @@ func TestKeyPolicyRegistryOpts(t *testing.T) {
t.Run("default namespace with asDefault false is a namespace key", func(t *testing.T) {
t.Parallel()

opts, err := keyPolicyRegistryOpts(t.Context(), &config.KeyPolicy{URI: distinctKeyURL(t, 3)}, log, defaultNamespace, false)
opts, err := keyPolicyRegistryOpts(t.Context(), &config.KeyPolicy{URI: distinctKeyURL(t, 3)}, log, defaultNamespace, false, newFxTestReporter(t))
require.NoError(t, err)

_, err = crypto.NewKEKRegistry(opts...)
Expand Down Expand Up @@ -193,7 +196,7 @@ func TestCreateVault(t *testing.T) {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()

v, err := createVault(&config.Config{Encryption: tt.enc}, reg)
v, err := createVault(&config.Config{Encryption: tt.enc}, reg, newFxTestReporter(t))
require.NoError(t, err)
require.NotNil(t, v)
})
Expand Down Expand Up @@ -221,7 +224,7 @@ func TestCreateVault_AppliesOverrideKeyConfig(t *testing.T) {
},
}}

_, err = createVault(cfg, reg)
_, err = createVault(cfg, reg, newFxTestReporter(t))
require.Error(t, err)
}

Expand All @@ -231,7 +234,7 @@ func TestCreateKEKRegistry(t *testing.T) {
lc := fxtest.NewLifecycle(t)
cfg := &config.Config{Encryption: config.Encryption{Enabled: true, Default: &config.KeyPolicy{URI: distinctKeyURL(t, 1)}}}

reg, err := createKEKRegistry(t.Context(), lc, cfg, logger.NewNoopLogger())
reg, err := createKEKRegistry(t.Context(), lc, cfg, logger.NewNoopLogger(), newFxTestReporter(t))
require.NoError(t, err)
require.NotNil(t, reg)

Expand All @@ -256,12 +259,12 @@ func TestCreateKEKRegistry_OverrideKeySelectedByNamespace(t *testing.T) {
},
}}

reg, err := createKEKRegistry(t.Context(), lc, cfg, logger.NewNoopLogger())
reg, err := createKEKRegistry(t.Context(), lc, cfg, logger.NewNoopLogger(), newFxTestReporter(t))
require.NoError(t, err)
lc.RequireStart()
t.Cleanup(func() { lc.RequireStop() })

vault, err := createVault(cfg, reg)
vault, err := createVault(cfg, reg, newFxTestReporter(t))
require.NoError(t, err)

// The override namespace seals under its own KEK.
Expand All @@ -285,6 +288,7 @@ func TestModule_NoKeys_ProvidesNilVault(t *testing.T) {
fx.Supply(fx.Annotate(t.Context(), fx.As(new(context.Context)))),
fx.Supply(&config.Config{Encryption: config.Encryption{Enabled: false}}),
fx.Provide(func() logger.Logger { return logger.NewNoopLogger() }),
fx.Provide(func() *metrics.Factory { return metrics.New("test", promauto.With(prometheus.NewRegistry())) }),
Module,
fx.Populate(&v),
fx.NopLogger,
Expand Down Expand Up @@ -312,6 +316,7 @@ func TestModule_DisabledWithKeys_ProvidesVault(t *testing.T) {
fx.Supply(fx.Annotate(t.Context(), fx.As(new(context.Context)))),
fx.Supply(cfg),
fx.Provide(func() logger.Logger { return logger.NewNoopLogger() }),
fx.Provide(func() *metrics.Factory { return metrics.New("test", promauto.With(prometheus.NewRegistry())) }),
Module,
fx.Populate(&v),
)
Expand All @@ -336,6 +341,7 @@ func TestModule_Enabled_ProvidesVaultAndRunsCleanly(t *testing.T) {
fx.Supply(fx.Annotate(t.Context(), fx.As(new(context.Context)))),
fx.Supply(cfg),
fx.Provide(func() logger.Logger { return logger.NewNoopLogger() }),
fx.Provide(func() *metrics.Factory { return metrics.New("test", promauto.With(prometheus.NewRegistry())) }),
Module,
fx.Populate(&v),
)
Expand All @@ -360,6 +366,7 @@ func TestModule_Enabled_InvalidURI_FailsConstruction(t *testing.T) {
fx.Supply(fx.Annotate(t.Context(), fx.As(new(context.Context)))),
fx.Supply(cfg),
fx.Provide(func() logger.Logger { return logger.NewNoopLogger() }),
fx.Provide(func() *metrics.Factory { return metrics.New("test", promauto.With(prometheus.NewRegistry())) }),
Module,
fx.NopLogger,
)
Expand Down Expand Up @@ -389,3 +396,11 @@ func closeKEKs(t *testing.T, keys []crypto.KEK) {
require.NoError(t, k.Close())
}
}

// newFxTestReporter builds a Reporter backed by its own private registry, so
// fx.go call sites under test have a Reporter to record into without
// colliding with any other test's metrics.
func newFxTestReporter(t *testing.T) *Reporter {
t.Helper()
return NewReporter(metrics.New("test", promauto.With(prometheus.NewRegistry())).ForSubsystem("encryption"))
}
71 changes: 71 additions & 0 deletions internal/kms/metered.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,71 @@
package kms

import (
"context"
"time"

"github.com/temporalio/temporal-proxy/pkg/crypto"
)

type (
// meteredKEK decorates a crypto.KEK, recording each wrap (Encrypt) and unwrap
// (Decrypt) as a KEK operation. It embeds the wrapped KEK so ID and Close pass
// through unchanged, and returns the wrapped call's result and error verbatim so
// metering never alters behavior.
meteredKEK struct {
crypto.KEK
provider string
recorder kekRecorder
}

// kekRecorder records KEK wrap/unwrap operations. *Reporter satisfies it; the
// interface lets meteredKEK be tested without a Prometheus-backed reporter.
kekRecorder interface {
KEKOp(provider, operation, result string, seconds float64)
}
)

func (m *meteredKEK) Encrypt(ctx context.Context, b []byte) ([]byte, error) {
start := time.Now()
ct, err := m.KEK.Encrypt(ctx, b)
m.recorder.KEKOp(m.provider, "wrap", resultLabel(err), time.Since(start).Seconds())
return ct, err
}

func (m *meteredKEK) Decrypt(ctx context.Context, b []byte) ([]byte, error) {
start := time.Now()
pt, err := m.KEK.Decrypt(ctx, b)
m.recorder.KEKOp(m.provider, "unwrap", resultLabel(err), time.Since(start).Seconds())
return pt, err
}

// newMeteredKEK wraps k so its KMS calls are recorded under provider.
func newMeteredKEK(k crypto.KEK, provider string, r kekRecorder) crypto.KEK {
return &meteredKEK{KEK: k, provider: provider, recorder: r}
}

// providerForScheme maps a KMS URI scheme to a stable, low-cardinality provider
// label. Unknown schemes are returned unchanged so a new backend still produces
// a usable (if unrecognized) label rather than an empty one.
func providerForScheme(scheme string) string {
switch scheme {
case "awskms":
return "aws"
case "gcpkms":
return "gcp"
case "azurekeyvault":
return "azure"
case "testing", "base64key":
return "testing"
default:
return scheme
}
}

// resultLabel maps an error to the "result" metric label value.
func resultLabel(err error) string {
if err != nil {
return "error"
}
return "success"
}
Loading