diff --git a/README.md b/README.md index 98fb9b1..f5fb398 100644 --- a/README.md +++ b/README.md @@ -75,6 +75,12 @@ reaches a different upstream with no change to the Worker. responses, so the upstream only ever sees ciphertext while local Workers keep exchanging cleartext. DEKs are wrapped by a KMS key (AWS KMS, Azure Key Vault, or GCP KMS), rotate automatically, and can be overridden per Namespace. Set `encryption.failures` to seal failure messages and stack traces too, as the SDK's `EncodeCommonAttributes` does. + + Payloads a Worker's own codec already encrypted, or that an earlier proxy sealed, can be forwarded as they are by + listing their encoding under `encryption.skipEncodings`. Listing an encoding trusts every Worker that sets it: the + proxy cannot tell ciphertext from plaintext labeled that way. Listing `binary/encrypted` also returns responses sealed + under a KEK this proxy doesn't hold, including its own if a key was removed from config, rather than failing the + call, so alert on `vault_ops_total{result="unknown_key"}`. - **Pluggable key management.** For a backend the proxy has no built-in support for, such as an on-prem HSM or an internal key service, point it at an extension server you run and it wraps DEKs through that instead. Only key material is exchanged; payloads never reach it. @@ -234,25 +240,30 @@ label reports blank there even though the same series populates it for a request fixed label is added to every series in the table, and to nothing registered outside the proxy's own collectors, so the runtime's `go_*` and `process_*` series stay as they are. -| Subsystem | Metric | Type | Labels | -| -------------- | ---------------------------- | --------- | ---------------------------------- | -| `server` | `requests_total` | counter | `method`, `code` | -| `server` | `request_duration_seconds` | histogram | `method` | -| `server` | `panics_total` | counter | `method` | -| `router` | `decisions_total` | counter | `upstream`, `outcome` | -| `router` | `forwarding_errors_total` | counter | `upstream`, `reason` | -| `encryption` | `vault_ops_total` | counter | `operation`, `result`, `namespace` | -| `encryption` | `vault_ops_duration_seconds` | histogram | `operation`, `namespace` | -| `encryption` | `kek_ops_total` | counter | `provider`, `operation`, `result` | -| `encryption` | `kek_ops_duration_seconds` | histogram | `provider`, `operation` | -| `encryption` | `dek_ops_total` | counter | `operation`, `result` | -| `encryption` | `dek_ops_duration_seconds` | histogram | `operation` | -| `encryption` | `dek_rotations_total` | counter | `reason` | -| `encryption` | `dek_cache_hits_total` | counter | none | -| `encryption` | `dek_cache_misses_total` | counter | none | -| `encryption` | `dek_cache_size` | gauge | none | -| `codec_server` | `requests_total` | counter | `route`, `code` | -| `codec_server` | `request_duration_seconds` | histogram | `route` | +`payloads_skipped_total` carries the metadata labels too (blank, as with `vault_ops`, for operations from the codec +server, since an HTTP request carries no gRPC metadata), and `vault_ops_total` reports `result="unknown_key"` for a +payload sealed under a KEK the proxy doesn't hold. + +| Subsystem | Metric | Type | Labels | +| -------------- | ---------------------------- | --------- | ------------------------------------ | +| `server` | `requests_total` | counter | `method`, `code` | +| `server` | `request_duration_seconds` | histogram | `method` | +| `server` | `panics_total` | counter | `method` | +| `router` | `decisions_total` | counter | `upstream`, `outcome` | +| `router` | `forwarding_errors_total` | counter | `upstream`, `reason` | +| `encryption` | `vault_ops_total` | counter | `operation`, `result`, `namespace` | +| `encryption` | `vault_ops_duration_seconds` | histogram | `operation`, `namespace` | +| `encryption` | `payloads_skipped_total` | counter | `operation`, `encoding`, `namespace` | +| `encryption` | `kek_ops_total` | counter | `provider`, `operation`, `result` | +| `encryption` | `kek_ops_duration_seconds` | histogram | `provider`, `operation` | +| `encryption` | `dek_ops_total` | counter | `operation`, `result` | +| `encryption` | `dek_ops_duration_seconds` | histogram | `operation` | +| `encryption` | `dek_rotations_total` | counter | `reason` | +| `encryption` | `dek_cache_hits_total` | counter | none | +| `encryption` | `dek_cache_misses_total` | counter | none | +| `encryption` | `dek_cache_size` | gauge | none | +| `codec_server` | `requests_total` | counter | `route`, `code` | +| `codec_server` | `request_duration_seconds` | histogram | `route` | The `encryption` subsystem only reports once encryption keys are configured, and `codec_server` only reports once the codec server is enabled. Its `route` label is the matched pattern (`/decode`, and so on), never the request path, diff --git a/dev/config.yaml b/dev/config.yaml index b402a3d..2b64584 100644 --- a/dev/config.yaml +++ b/dev/config.yaml @@ -60,6 +60,10 @@ routing: encryption: enabled: true cacheSize: 200 + # Encodings already encrypted before they reach the proxy, forwarded unsealed. + # List binary/encrypted on a proxy chained behind another one. + # skipEncodings: + # - binary/encrypted default: uri: testing://-ynIaZzFbAjp9VPgu0Ohk9YeQSLS9ta0m9mtnOnGZqo= diff --git a/e2e/encryption_test.go b/e2e/encryption_test.go index ca6f7bb..71fd01d 100644 --- a/e2e/encryption_test.go +++ b/e2e/encryption_test.go @@ -90,15 +90,113 @@ func TestEndToEndPayloadEncryption(t *testing.T) { require.NotEmpty(t, sealed.GetMetadata()[wireDEK], "sealed payload must carry the wrapped DEK") } +// TestEndToEndSkipsWorkerEncryptedPayloads drives a payload a worker's own codec +// already encrypted through the full stack. Not copying SkipEncodings into +// CodecOptions in dataplane.New fails the byte-for-byte check. +func TestEndToEndSkipsWorkerEncryptedPayloads(t *testing.T) { + t.Parallel() + + up := dataplanetest.NewUpstream(t) + + cfg := dataplanetest.Config(up) + cfg.Encryption = config.Encryption{ + Enabled: true, + SkipEncodings: []string{"acme/aes-gcm"}, + Default: &config.KeyPolicy{URI: testingKeyURI(t), Duration: time.Hour}, + } + + f := dataplanetest.StartApp(t, cfg) + + workerSealed := &common.Payload{ + Metadata: map[string][]byte{wireEncoding: []byte("acme/aes-gcm")}, + Data: []byte("ciphertext-from-a-worker"), + } + + resp, err := f.Client().QueryWorkflow(f.Context(), queryWith(workerSealed), grpc.WaitForReady(true)) + require.NoError(t, err) + require.True(t, proto.Equal(workerSealed, resp.GetQueryResult().GetPayloads()[0])) + + reqs := up.Requests() + require.Len(t, reqs, 1) + sent := reqs[0].(*workflowservice.QueryWorkflowRequest).GetQuery().GetQueryArgs().GetPayloads() + require.Len(t, sent, 1) + require.True(t, proto.Equal(workerSealed, sent[0]), "upstream must receive the worker's payload unchanged") +} + +// TestEndToEndChainedProxiesSealOnce runs client -> hop A -> hop B -> upstream, +// where only A holds A's key and B lists binary/encrypted. A round trip alone +// can't prove B skipped, since B would decode its own extra layer, so B's skip +// metric is the evidence. Removing the unknown-key pass-through in +// Encryptor.Decode fails the QueryWorkflow call. +func TestEndToEndChainedProxiesSealOnce(t *testing.T) { + t.Parallel() + + up := dataplanetest.NewUpstream(t) + + cfgB := dataplanetest.Config(up) + cfgB.Encryption = config.Encryption{ + Enabled: true, + SkipEncodings: []string{wireEncryptedMarker}, + Default: &config.KeyPolicy{URI: testingKeyURIOf(t, 0x0b), Duration: time.Hour}, + } + hopB := dataplanetest.StartApp(t, cfgB) + + cfgA := dataplanetest.Config(up) + cfgA.Upstreams[0].Listen = config.ListenConfig{HostPort: hopB.Addr(), Insecure: true} + cfgA.Encryption = config.Encryption{ + Enabled: true, + Default: &config.KeyPolicy{URI: testingKeyURIOf(t, 0x0a), Duration: time.Hour}, + } + hopA := dataplanetest.StartApp(t, cfgA) + + secret := &common.Payload{ + Metadata: map[string][]byte{wireEncoding: []byte("json/plain")}, + Data: []byte(`"the-answer-is-42"`), + } + + resp, err := hopA.Client().QueryWorkflow(hopA.Context(), queryWith(secret), grpc.WaitForReady(true)) + require.NoError(t, err) + require.True(t, proto.Equal(secret, resp.GetQueryResult().GetPayloads()[0])) + + reqs := up.Requests() + require.Len(t, reqs, 1) + payloads := reqs[0].(*workflowservice.QueryWorkflowRequest).GetQuery().GetQueryArgs().GetPayloads() + require.Len(t, payloads, 1) + require.Equal(t, wireEncryptedMarker, string(payloads[0].GetMetadata()[wireEncoding])) + + requireLabel(t, hopB, "test_encryption_payloads_skipped_total", "operation", "encrypt") + requireLabel(t, hopB, "test_encryption_payloads_skipped_total", "operation", "decrypt") +} + // testingKeyURI builds a local testing:// key URI with a fixed 32-byte key. The // kms module rewrites testing:// to gocloud's base64key:// local keeper, so no // cloud KMS is needed. func testingKeyURI(t *testing.T) url.URL { t.Helper() - key := base64.StdEncoding.EncodeToString(bytes.Repeat([]byte{0x2a}, 32)) + return testingKeyURIOf(t, 0x2a) +} + +// testingKeyURIOf is testingKeyURI with the key filled with b, so two proxies in +// one test can hold different keys. +func testingKeyURIOf(t *testing.T, b byte) url.URL { + t.Helper() + + key := base64.StdEncoding.EncodeToString(bytes.Repeat([]byte{b}, 32)) u, err := url.Parse("testing://" + key) require.NoError(t, err) return *u } + +// queryWith is a QueryWorkflow request carrying p as its only query argument. +func queryWith(p *common.Payload) *workflowservice.QueryWorkflowRequest { + return &workflowservice.QueryWorkflowRequest{ + Namespace: "ns1", + Execution: &common.WorkflowExecution{WorkflowId: "wf-1"}, + Query: &query.WorkflowQuery{ + QueryType: "state", + QueryArgs: &common.Payloads{Payloads: []*common.Payload{p}}, + }, + } +} diff --git a/internal/config/encryption.go b/internal/config/encryption.go index 1a04295..350be71 100644 --- a/internal/config/encryption.go +++ b/internal/config/encryption.go @@ -29,13 +29,15 @@ type ( // (local) namespace names, matching the namespace the vault seals under at // request time. Failures also seals the message and stack trace of outbound // failures, the way the Temporal SDK's EncodeCommonAttributes does; it requires - // Enabled. + // Enabled. SkipEncodings lists payload encodings already encrypted before they + // reach the proxy, which are forwarded unsealed. Encryption struct { - Enabled bool `yaml:"enabled"` - Failures bool `yaml:"failures"` - CacheSize *int `yaml:"cacheSize"` - Default *KeyPolicy `yaml:"default"` - Overrides map[string]KeyPolicy `yaml:"overrides"` + Enabled bool `yaml:"enabled"` + Failures bool `yaml:"failures"` + CacheSize *int `yaml:"cacheSize"` + Default *KeyPolicy `yaml:"default"` + Overrides map[string]KeyPolicy `yaml:"overrides"` + SkipEncodings []string `yaml:"skipEncodings"` } // KeyPolicy describes the KMS key backing a DEK and its rotation schedule. @@ -95,6 +97,11 @@ func (e *Encryption) Validate() error { ) } + rules = append(rules, + validation.Field("skipEncodings", e.SkipEncodings, validation.Unique[string]()), + validation.Children("skipEncodings", e.SkipEncodings, nonBlankEncoding()), + ) + return validation.Validate("", rules...) } @@ -195,3 +202,15 @@ func validKeyURIRef() validation.Check[*url.URL] { return nil } } + +// nonBlankEncoding rejects a blank skipEncodings entry, which would match every +// payload that carries no encoding at all. +func nonBlankEncoding() validation.Check[*string] { + return func(s *string) error { + if *s == "" { + return errors.New("must not be blank") + } + + return nil + } +} diff --git a/internal/config/encryption_test.go b/internal/config/encryption_test.go index 8df3aaa..49b5c04 100644 --- a/internal/config/encryption_test.go +++ b/internal/config/encryption_test.go @@ -170,6 +170,35 @@ func TestEncryptionValidate(t *testing.T) { }, wantErr: "overrides", }, + { + name: "skip encodings", + cfg: config.Encryption{ + Enabled: true, + Default: &valid, + SkipEncodings: []string{"binary/encrypted", "acme/aes-gcm"}, + }, + }, + { + // A blank entry would match every payload with no encoding. Dropping the + // Children rule fails this row. + name: "blank skip encoding", + cfg: config.Encryption{ + Enabled: true, + Default: &valid, + SkipEncodings: []string{"acme/aes-gcm", ""}, + }, + wantErr: "skipEncodings[1]", + }, + { + // Dropping the Unique check fails this row. + name: "duplicate skip encoding", + cfg: config.Encryption{ + Enabled: true, + Default: &valid, + SkipEncodings: []string{"acme/aes-gcm", "acme/aes-gcm"}, + }, + wantErr: "skipEncodings", + }, } for _, tt := range tests { diff --git a/internal/dataplane/dataplane.go b/internal/dataplane/dataplane.go index 40e0a1f..488bed2 100644 --- a/internal/dataplane/dataplane.go +++ b/internal/dataplane/dataplane.go @@ -114,7 +114,11 @@ func New(ctx context.Context, cfg *config.Config, opts ...Option) (*Dataplane, e // Every upstream applies the same chain: the vault and the encryption switch // are global, so nothing here varies per upstream. Building it once is also // what lets the codec server apply the identical chain. - codecOpts := proxy.CodecOptions{Encrypt: cfg.Encryption.Enabled, EncodeFailures: cfg.Encryption.Failures} + codecOpts := proxy.CodecOptions{ + Encrypt: cfg.Encryption.Enabled, + EncodeFailures: cfg.Encryption.Failures, + SkipEncodings: cfg.Encryption.SkipEncodings, + } // Only assign the vault once it is known to be there. o.vault is a concrete // pointer and the field is an interface, so assigning unconditionally would @@ -144,6 +148,11 @@ func New(ctx context.Context, cfg *config.Config, opts ...Option) (*Dataplane, e ) } + // With no keys there is no encryption codec, so the list has nothing to skip. + if len(cfg.Encryption.SkipEncodings) > 0 && o.vault == nil { + o.logger.Warn("encryption.skipEncodings is set but no encryption keys are configured, so it has no effect") + } + dp := &Dataplane{ ctx: ctx, hostPort: cfg.Listen.HostPort, diff --git a/internal/dataplane/dataplane_test.go b/internal/dataplane/dataplane_test.go index e71fef7..1017f96 100644 --- a/internal/dataplane/dataplane_test.go +++ b/internal/dataplane/dataplane_test.go @@ -222,6 +222,39 @@ func TestNewTwiceOverOneMetricsFactoryDoesNotPanic(t *testing.T) { } } +// TestNewWarnsWhenSkipEncodingsHasNoKeys covers a list with nothing to apply it +// to. Removing the warning fails the first case; dropping the length check fails +// the second. +func TestNewWarnsWhenSkipEncodingsHasNoKeys(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + skip []string + wantWarn bool + }{ + {name: "list without keys", skip: []string{"acme/aes-gcm"}, wantWarn: true}, + {name: "no list"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + cfg := testConfig() + cfg.Encryption.SkipEncodings = tc.skip + + log := logger.NewTestLogger() + deps := newTestDeps(t, cfg) + deps.logger = log + + _, err := dataplane.New(deps.ctx, cfg, deps.opts()...) + require.NoError(t, err) + require.Equal(t, tc.wantWarn, log.Contains(skipEncodingsUnusedWarning)) + }) + } +} + // opts returns d as the options [dataplane.New] takes, less any named in omit, // so a caller can prove New reports one as missing. Names are the ones New // reports. @@ -303,6 +336,11 @@ func testingKeyURL(t *testing.T) url.URL { const cloudAPIUnusedWarning = "apiTranslations is configured but no upstream is Temporal Cloud, so no method " + "will be translated" +// skipEncodingsUnusedWarning is the message New logs for a skip list with no +// encryption codec to act on. +const skipEncodingsUnusedWarning = "encryption.skipEncodings is set but no encryption keys are configured, so it " + + "has no effect" + // stagingCloudAPI is an override that says something, which is what an inert // block has to be to be worth warning about: one that says nothing is // indistinguishable from no block at all, and describes the defaults anyway. diff --git a/internal/proxy/codec.go b/internal/proxy/codec.go index ec74fb4..81477a6 100644 --- a/internal/proxy/codec.go +++ b/internal/proxy/codec.go @@ -32,6 +32,12 @@ type ( // are sealed with every other payload. It requires Encrypt. EncodeFailures bool + // SkipEncodings lists payload encodings already encrypted before they + // reach the proxy. Outbound payloads under one are forwarded unsealed, and + // listing codec.EncryptionEncoding also returns inbound payloads sealed + // under a KEK the vault doesn't hold, rather than failing the call. + SkipEncodings []string + // Reporter records the duration and result of each vault operation. It is // required whenever a Vault is set. Reporter *Reporter @@ -83,6 +89,21 @@ func NewCodecs(opts CodecOptions) (*Codecs, error) { if opts.Encrypt { c.outbound = append(c.outbound, enc) } + + // The set is built once here; only the observer is bound per request. + if len(opts.SkipEncodings) > 0 { + encodings := codec.WithSkipEncodings(opts.SkipEncodings...) + skip := func(ctx context.Context, ns string) codec.Option { + return codec.WithEncryptorOptions(encodings, codec.WithSkipObserver(func(op, encoding string) { + opts.Reporter.PayloadSkipped(ctx, op, encoding, ns) + })) + } + + c.inbound = append(c.inbound, skip) + if opts.Encrypt { + c.outbound = append(c.outbound, skip) + } + } } return c, nil diff --git a/internal/proxy/encryption.go b/internal/proxy/encryption.go index 7dba566..df05dea 100644 --- a/internal/proxy/encryption.go +++ b/internal/proxy/encryption.go @@ -2,6 +2,7 @@ package proxy import ( "context" + "errors" "time" "github.com/temporalio/temporal-proxy/pkg/crypto" @@ -46,11 +47,16 @@ func (c *cipher) Decrypt(m *crypto.Message) ([]byte, error) { return pt, err } -// resultLabel maps an error to the "result" metric label value. +// resultLabel maps an error to the "result" metric label value. An unknown key +// gets its own value: on a chained hop it is expected, and on a standalone proxy +// it means a KEK went missing, neither of which is a vault failure. func resultLabel(err error) string { - if err != nil { + switch { + case err == nil: + return "success" + case errors.Is(err, crypto.ErrUnknownKey): + return "unknown_key" + default: return "error" } - - return "success" } diff --git a/internal/proxy/encryption_test.go b/internal/proxy/encryption_test.go index 17aabe0..82dd4ae 100644 --- a/internal/proxy/encryption_test.go +++ b/internal/proxy/encryption_test.go @@ -363,6 +363,132 @@ func TestEncryptionSkipsMetricsForPassThrough(t *testing.T) { require.False(t, hasLabels(ops, map[string]string{"operation": "decrypt", "result": "success", "namespace": "ns1"})) } +// TestEncryptionSkipsListedEncodings pins the interceptor's outbound skip and its +// metric. Not passing SkipEncodings into the chain in NewCodecs fails the Same +// check; not binding the observer fails the metric check. +func TestEncryptionSkipsListedEncodings(t *testing.T) { + t.Parallel() + + reg := prometheus.NewRegistry() + reporter := proxy.NewReporter( + metrics.New("proxy", promauto.With(reg)).ForSubsystem("encryption"), + proxy.WithNamespaceLabels(true), + ) + + v := &fakeVault{} + interceptor, err := proxy.CodecInterceptor(proxy.CodecOptions{ + Vault: v, + Encrypt: true, + Reporter: reporter, + SkipEncodings: []string{"acme/aes-gcm"}, + }) + require.NoError(t, err) + + workerSealed := testPayload("acme/aes-gcm", "ciphertext-from-a-worker") + invoker := func(_ context.Context, _ string, gotReq, _ any, _ *grpc.ClientConn, _ ...grpc.CallOption) error { + sent := gotReq.(*workflowservice.StartWorkflowExecutionRequest).Input.Payloads + require.Len(t, sent, 1) + require.Same(t, workerSealed, sent[0]) + return nil + } + + ctx := metadata.AppendToOutgoingContext(t.Context(), meta.NamespaceHeader, "ns1") + resp := &workflowservice.StartWorkflowExecutionRequest{} + require.NoError(t, interceptor(ctx, "/method", startRequest(workerSealed), resp, nil, invoker)) + + require.Empty(t, v.namespaces, "a listed encoding must not reach the vault") + + skipped := gatherFamily(t, reg, "proxy_encryption_payloads_skipped_total") + require.NotNil(t, skipped) + require.True(t, hasLabels(skipped, map[string]string{ + "operation": "encrypt", "encoding": "acme/aes-gcm", "namespace": "ns1", + })) +} + +// TestEncryptionUnknownKey pins both sides of the decode rule at the interceptor, +// and the unknown_key result label. Removing the ErrUnknownKey case from +// resultLabel fails both rows' vault_ops check. +func TestEncryptionUnknownKey(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + skip []string + wantPass bool + }{ + {name: "passes through when binary/encrypted is listed", skip: []string{codec.EncryptionEncoding}, wantPass: true}, + {name: "fails the call otherwise"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + reg := prometheus.NewRegistry() + reporter := proxy.NewReporter( + metrics.New("proxy", promauto.With(reg)).ForSubsystem("encryption"), + proxy.WithNamespaceLabels(true), + ) + + v := &fakeVault{openErr: fmt.Errorf("%w: other-proxy-kek", crypto.ErrUnknownKey)} + interceptor, err := proxy.CodecInterceptor(proxy.CodecOptions{ + Vault: v, + Reporter: reporter, + SkipEncodings: tc.skip, + }) + require.NoError(t, err) + + theirs := sealedPayload(t, testPayload("json/plain", `"x"`)) + ctx := metadata.AppendToOutgoingContext(t.Context(), meta.NamespaceHeader, "ns1") + resp := &workflowservice.StartWorkflowExecutionRequest{} + err = interceptor(ctx, "/method", startRequest(), resp, nil, respondWith(theirs)) + + ops := gatherFamily(t, reg, "proxy_encryption_vault_ops_total") + require.NotNil(t, ops) + require.True(t, hasLabels(ops, map[string]string{ + "operation": "decrypt", "result": "unknown_key", "namespace": "ns1", + })) + + if !tc.wantPass { + // ErrorContains rather than ErrorIs: whether the payload visitor wraps + // with %w is go.temporal.io/api's business, not ours. + require.ErrorContains(t, err, "unknown key: other-proxy-kek") + return + } + + require.NoError(t, err) + require.True(t, proto.Equal(theirs, resp.Input.Payloads[0])) + + skipped := gatherFamily(t, reg, "proxy_encryption_payloads_skipped_total") + require.NotNil(t, skipped) + require.True(t, hasLabels(skipped, map[string]string{ + "operation": "decrypt", "encoding": codec.EncryptionEncoding, "namespace": "ns1", + })) + }) + } +} + +// TestCodecsEncodeSkipsListedEncodings covers the codec server's path, which calls +// Codecs.Encode rather than the interceptor. Building the skip option only inside +// Interceptor fails it. +func TestCodecsEncodeSkipsListedEncodings(t *testing.T) { + t.Parallel() + + c, err := proxy.NewCodecs(proxy.CodecOptions{ + Vault: &fakeVault{}, + Encrypt: true, + Reporter: newTestReporter(t), + SkipEncodings: []string{"acme/aes-gcm"}, + }) + require.NoError(t, err) + + p := testPayload("acme/aes-gcm", "ciphertext-from-a-worker") + got, err := c.Encode(t.Context(), "ns1", []*common.Payload{p}) + require.NoError(t, err) + require.Len(t, got, 1) + require.Same(t, p, got[0]) +} + func (f *fakeVault) Seal(_ context.Context, ns string, data []byte) (*crypto.Message, error) { f.mu.Lock() defer f.mu.Unlock() diff --git a/internal/proxy/reporter.go b/internal/proxy/reporter.go index 84d60dd..2215c87 100644 --- a/internal/proxy/reporter.go +++ b/internal/proxy/reporter.go @@ -12,15 +12,16 @@ type ( // Reporter records envelope-operation telemetry to Prometheus: each seal // (encrypt) and open (decrypt) the encryption codec performs, on both the // gRPC path and the codec server's HTTP path, timed end to end, including - // any KEK wrap or unwrap and any DEK cache lookup along the way. The AES-step - // duration alone is owned by internal/kms. The namespace label is always - // declared but only carries a value when namespace labels are enabled, since - // the set of namespaces is unbounded; handles are resolved per call via - // WithLabelValues rather than pre-computed. A Reporter is safe for concurrent - // use. + // any KEK wrap or unwrap and any DEK cache lookup along the way. It also + // counts payloads the encryption codec skipped. The AES-step duration alone + // is owned by internal/kms. The namespace label is always declared but only + // carries a value when namespace labels are enabled, since the set of + // namespaces is unbounded; handles are resolved per call via WithLabelValues + // rather than pre-computed. A Reporter is safe for concurrent use. Reporter struct { ops *prometheus.CounterVec duration *prometheus.HistogramVec + skipped *prometheus.CounterVec namespaceLabels bool labels metrics.MetadataLabels } @@ -47,6 +48,11 @@ func NewReporter(f *metrics.Factory, opts ...ReporterOption) *Reporter { Help: "Duration of envelope operations in seconds end to end, including any KEK wrap or unwrap, labeled by operation and namespace.", }, append([]string{"operation", "namespace"}, r.labels.Names()...)) + r.skipped = f.NewCounter(prometheus.CounterOpts{ + Name: "payloads_skipped_total", + Help: "Total payloads the encryption codec left as they were, labeled by operation, encoding, and namespace.", + }, append([]string{"operation", "encoding", "namespace"}, r.labels.Names()...)) + return r } @@ -74,6 +80,14 @@ func (r *Reporter) VaultOp(ctx context.Context, operation, result, namespace str r.duration.WithLabelValues(append([]string{operation, ns}, labels...)...).Observe(seconds) } +// PayloadSkipped records one payload the encryption codec forwarded without +// sealing or opening it. encoding only ever holds a configured skipEncodings +// entry, so it stays bounded. +func (r *Reporter) PayloadSkipped(ctx context.Context, operation, encoding, namespace string) { + labels := r.labels.AppendValues(ctx, nil) + r.skipped.WithLabelValues(append([]string{operation, encoding, r.nsLabel(namespace)}, labels...)...).Inc() +} + // nsLabel returns ns when namespace labels are enabled and blank when they are // not, which Prometheus reads as the label not being there. func (r *Reporter) nsLabel(ns string) string { diff --git a/pkg/codec/chain.go b/pkg/codec/chain.go index ffcccee..8945eca 100644 --- a/pkg/codec/chain.go +++ b/pkg/codec/chain.go @@ -29,7 +29,8 @@ type ( Option func(*options) options struct { - cipher Cipher + cipher Cipher + encryptor []EncryptorOption } ) @@ -45,7 +46,7 @@ func NewChain(opts ...Option) Chain { // because ciphertext does not compress. var codecs []Codec if o.cipher != nil { - codecs = append(codecs, NewEncryptor(o.cipher)) + codecs = append(codecs, NewEncryptor(o.cipher, o.encryptor...)) } return Chain{codecs: codecs} @@ -61,6 +62,14 @@ func WithCipher(c Cipher) Option { } } +// WithEncryptorOptions passes opts to the [Encryptor] the chain builds. Without +// [WithCipher] there is no encryptor, so they have no effect. +func WithEncryptorOptions(opts ...EncryptorOption) Option { + return func(o *options) { + o.encryptor = append(o.encryptor, opts...) + } +} + // Encode runs payloads through every codec in order. func (c Chain) Encode(payloads []*common.Payload) ([]*common.Payload, error) { var err error diff --git a/pkg/codec/chain_test.go b/pkg/codec/chain_test.go index b24589c..5f8958c 100644 --- a/pkg/codec/chain_test.go +++ b/pkg/codec/chain_test.go @@ -20,6 +20,10 @@ func TestNewChainNoCodecs(t *testing.T) { }{ {name: "no options"}, {name: "nil cipher is ignored", opts: []codec.Option{codec.WithCipher(nil)}}, + { + name: "encryptor options alone add no codec", + opts: []codec.Option{codec.WithEncryptorOptions(codec.WithSkipEncodings("json/plain"))}, + }, } for _, tc := range tests { @@ -92,3 +96,21 @@ func TestChainCipherErrors(t *testing.T) { require.ErrorContains(t, err, "failed to decrypt payload") }) } + +// TestNewChainForwardsEncryptorOptions shows WithEncryptorOptions reaches the +// encryptor NewChain builds. Dropping o.encryptor from the NewEncryptor call in +// NewChain fails it. +func TestNewChainForwardsEncryptorOptions(t *testing.T) { + t.Parallel() + + chain := codec.NewChain( + codec.WithCipher(&fakeCipher{}), + codec.WithEncryptorOptions(codec.WithSkipEncodings("acme/aes-gcm")), + ) + + p := testPayload("acme/aes-gcm", "ciphertext-from-a-worker") + got, err := chain.Encode([]*common.Payload{p}) + require.NoError(t, err) + require.Len(t, got, 1) + require.Same(t, p, got[0]) +} diff --git a/pkg/codec/doc.go b/pkg/codec/doc.go index a9b6351..4f9f900 100644 --- a/pkg/codec/doc.go +++ b/pkg/codec/doc.go @@ -4,6 +4,8 @@ // them again. A sealed payload is self-describing: the ciphertext travels with // the ID of the key that wrapped its DEK and the wrapped DEK itself, so opening // one needs nothing but the payload and a Cipher that can reach that key. +// Encodings listed through [WithSkipEncodings] are treated as already encrypted +// and forwarded unchanged. // // [NewChain] assembles the codecs its options enable into a single [Chain] in the // order they have to be applied, so callers say what they want enabled rather diff --git a/pkg/codec/encryptor.go b/pkg/codec/encryptor.go index afacfbb..0a75cbf 100644 --- a/pkg/codec/encryptor.go +++ b/pkg/codec/encryptor.go @@ -1,6 +1,7 @@ package codec import ( + "errors" "fmt" "go.temporal.io/api/common/v1" @@ -24,14 +25,31 @@ const ( // EncryptionEncoding is the encoding a sealed payload is marked with. EncryptionEncoding = "binary/encrypted" + + // SkipOpEncrypt is the op a SkipObserver is told when Encode forwards a + // payload whose encoding is listed. + SkipOpEncrypt = "encrypt" + + // SkipOpDecrypt is the op a SkipObserver is told when Decode passes through a + // payload sealed under a KEK the cipher doesn't hold. + SkipOpDecrypt = "decrypt" ) type ( // Encryptor seals payloads with envelope encryption and opens them again. Encryptor struct { cipher Cipher + skip map[string]struct{} + onSkip SkipObserver } + // EncryptorOption configures an [Encryptor]. + EncryptorOption func(*Encryptor) + + // SkipObserver is told about each payload an [Encryptor] leaves as it was: + // op is SkipOpEncrypt or SkipOpDecrypt, and encoding is the payload's. + SkipObserver func(op, encoding string) + // Cipher encrypts and decrypts bytes. It is what [Encryptor] depends on, // typically a wrapper that binds a [crypto.Vault] to a namespace. Cipher interface { @@ -41,17 +59,51 @@ type ( ) // NewEncryptor returns an [Encryptor] that seals and opens payloads through c. -func NewEncryptor(c Cipher) *Encryptor { - return &Encryptor{cipher: c} +func NewEncryptor(c Cipher, opts ...EncryptorOption) *Encryptor { + e := &Encryptor{cipher: c} + for _, opt := range opts { + opt(e) + } + + return e +} + +// WithSkipEncodings lists encodings Encode treats as already encrypted and +// forwards unchanged. Blank entries are dropped, since one would match every +// payload that carries no encoding. A later use replaces an earlier one. +func WithSkipEncodings(encodings ...string) EncryptorOption { + set := make(map[string]struct{}, len(encodings)) + for _, enc := range encodings { + if enc != "" { + set[enc] = struct{}{} + } + } + + return func(e *Encryptor) { e.skip = set } } -// Encode seals every payload in payloads, returning payloads whose data is the -// ciphertext and whose metadata carries the wrapped DEK needed to open it. Each -// original payload is sealed whole, metadata included, so [Encryptor.Decode] -// restores it exactly. +// WithSkipObserver reports each payload the [Encryptor] skips to fn. A nil fn +// reports nothing. +func WithSkipObserver(fn SkipObserver) EncryptorOption { + return func(e *Encryptor) { e.onSkip = fn } +} + +// Encode seals every payload in payloads, returning payloads whose data is +// the ciphertext and whose metadata carries the wrapped DEK needed to open +// it. Each original payload is sealed whole, metadata included, so +// [Encryptor.Decode] restores it exactly. Payloads whose encoding is listed +// through [WithSkipEncodings] are returned as they are. func (c *Encryptor) Encode(payloads []*common.Payload) ([]*common.Payload, error) { res := make([]*common.Payload, len(payloads)) for i, p := range payloads { + if len(c.skip) > 0 { + if enc := string(p.GetMetadata()[MetadataEncoding]); c.skips(enc) { + res[i] = p + c.skipped(SkipOpEncrypt, enc) + continue + } + } + data, err := p.Marshal() if err != nil { return nil, fmt.Errorf("failed to marshal payload: %w", err) @@ -79,7 +131,9 @@ func (c *Encryptor) Encode(payloads []*common.Payload) ([]*common.Payload, error // contract, the EncryptionEncoding marker plus both key-material entries, is opened // and restored to its original form. Anything else passes through unchanged so // payloads produced elsewhere survive the round trip, including ones that use the -// same encoding name without our key material. +// same encoding name without our key material. When EncryptionEncoding is listed +// through [WithSkipEncodings], a payload sealed under a KEK the cipher doesn't hold +// is passed through too. func (c *Encryptor) Decode(payloads []*common.Payload) ([]*common.Payload, error) { res := make([]*common.Payload, len(payloads)) for i, p := range payloads { @@ -105,6 +159,15 @@ func (c *Encryptor) Decode(payloads []*common.Payload) ([]*common.Payload, error }, }) if err != nil { + // A chained hop sees payloads an earlier proxy sealed under keys it + // doesn't hold. Listing EncryptionEncoding is the operator saying + // those belong to someone else. + if c.skips(EncryptionEncoding) && errors.Is(err, crypto.ErrUnknownKey) { + res[i] = p + c.skipped(SkipOpDecrypt, EncryptionEncoding) + continue + } + return nil, fmt.Errorf("failed to decrypt payload: %w", err) } @@ -118,3 +181,16 @@ func (c *Encryptor) Decode(payloads []*common.Payload) ([]*common.Payload, error return res, nil } + +// skips reports whether enc is listed through WithSkipEncodings. +func (c *Encryptor) skips(enc string) bool { + _, ok := c.skip[enc] + return ok +} + +// skipped tells the observer, if there is one, that a payload was left alone. +func (c *Encryptor) skipped(op, enc string) { + if c.onSkip != nil { + c.onSkip(op, enc) + } +} diff --git a/pkg/codec/encryptor_test.go b/pkg/codec/encryptor_test.go index 9bcf333..d4a20fe 100644 --- a/pkg/codec/encryptor_test.go +++ b/pkg/codec/encryptor_test.go @@ -32,6 +32,7 @@ type fakeCipher struct { encErr error // when set, Encrypt returns it decErr error // when set, Decrypt returns it decReturn []byte // when set, Decrypt returns these bytes instead of the unsealed plaintext + unknown string // when set, Decrypt fails with crypto.ErrUnknownKey for messages under this KEK ID } func TestEncryptorRoundtrip(t *testing.T) { @@ -98,6 +99,86 @@ func TestEncryptorEncodeEmpty(t *testing.T) { require.Empty(t, c.encrypted) } +// TestEncryptorEncodeSkipsListedEncodings pins which payloads are forwarded as-is. +// Dropping the skip check in Encode fails every wantSkip row; matching with +// strings.EqualFold fails the case-sensitivity row; keeping blank entries in +// WithSkipEncodings fails the blank row. +func TestEncryptorEncodeSkipsListedEncodings(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + encoding string + skip []string + wantSkip bool + }{ + {name: "listed custom encoding", encoding: "acme/aes-gcm", skip: []string{"acme/aes-gcm"}, wantSkip: true}, + { + // Matching is on the marker alone, so samples-style codec output + // without our wrapped DEK is skipped too. + name: "listed binary/encrypted without our key material", + encoding: codec.EncryptionEncoding, + skip: []string{codec.EncryptionEncoding}, + wantSkip: true, + }, + {name: "unlisted encoding", encoding: "json/plain", skip: []string{"acme/aes-gcm"}}, + {name: "match is case-sensitive", encoding: "ACME/AES-GCM", skip: []string{"acme/aes-gcm"}}, + {name: "empty list", encoding: "acme/aes-gcm"}, + {name: "blank entry never matches a blank encoding", encoding: "", skip: []string{""}}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + c := &fakeCipher{} + var skipped []string + enc := codec.NewEncryptor(c, + codec.WithSkipEncodings(tc.skip...), + codec.WithSkipObserver(func(op, encoding string) { skipped = append(skipped, op+" "+encoding) }), + ) + + p := testPayload(tc.encoding, `"data"`) + got, err := enc.Encode([]*common.Payload{p}) + require.NoError(t, err) + require.Len(t, got, 1) + + if tc.wantSkip { + require.Same(t, p, got[0]) + require.Empty(t, c.encrypted) + require.Equal(t, []string{codec.SkipOpEncrypt + " " + tc.encoding}, skipped) + return + } + + require.Len(t, c.encrypted, 1) + require.True(t, bytes.HasPrefix(got[0].Data, []byte(sealPrefix))) + require.Empty(t, skipped) + }) + } +} + +// TestEncryptorEncodeMixedBatch shows each payload is decided on its own and order +// is kept. Returning early from Encode on the first skip fails the length check; +// reading p.Metadata on a payload with none must not match anything. +func TestEncryptorEncodeMixedBatch(t *testing.T) { + t.Parallel() + + c := &fakeCipher{} + plain := testPayload("json/plain", `"plain"`) + workerSealed := testPayload("acme/aes-gcm", "ciphertext-from-a-worker") + bare := &common.Payload{Data: []byte("no metadata")} + + got, err := codec.NewEncryptor(c, codec.WithSkipEncodings("acme/aes-gcm")). + Encode([]*common.Payload{plain, workerSealed, bare}) + require.NoError(t, err) + require.Len(t, got, 3) + + require.True(t, bytes.HasPrefix(got[0].Data, []byte(sealPrefix))) + require.Same(t, workerSealed, got[1]) + require.True(t, bytes.HasPrefix(got[2].Data, []byte(sealPrefix))) + require.Len(t, c.encrypted, 2) +} + func TestEncryptorDecodePassesThrough(t *testing.T) { t.Parallel() @@ -252,6 +333,100 @@ func TestEncryptorErrors(t *testing.T) { }) } +// TestEncryptorDecodeUnknownKey pins when a payload sealed elsewhere is passed +// through. Dropping the skips(EncryptionEncoding) guard fails the two "fails" +// rows; matching any error instead of crypto.ErrUnknownKey fails the "other +// errors" row. +func TestEncryptorDecodeUnknownKey(t *testing.T) { + t.Parallel() + + const otherKEK = "other-proxy-kek" + errKMS := errors.New("kms unavailable") + + tests := []struct { + name string + skip []string + cipher *fakeCipher + wantErr error + }{ + { + name: "passes through when binary/encrypted is listed", + skip: []string{codec.EncryptionEncoding}, + cipher: &fakeCipher{unknown: otherKEK}, + }, + { + name: "fails when nothing is listed", + cipher: &fakeCipher{unknown: otherKEK}, + wantErr: crypto.ErrUnknownKey, + }, + { + name: "fails when only custom encodings are listed", + skip: []string{"acme/aes-gcm"}, + cipher: &fakeCipher{unknown: otherKEK}, + wantErr: crypto.ErrUnknownKey, + }, + { + name: "other decrypt errors stay fatal", + skip: []string{codec.EncryptionEncoding}, + cipher: &fakeCipher{decErr: errKMS}, + wantErr: errKMS, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + theirs := sealedUnder(t, otherKEK, testPayload("json/plain", `"x"`)) + var skipped []string + enc := codec.NewEncryptor(tc.cipher, + codec.WithSkipEncodings(tc.skip...), + codec.WithSkipObserver(func(op, encoding string) { skipped = append(skipped, op+" "+encoding) }), + ) + + got, err := enc.Decode([]*common.Payload{theirs}) + if tc.wantErr != nil { + require.ErrorIs(t, err, tc.wantErr) + require.Nil(t, got) + require.Empty(t, skipped) + return + } + + require.NoError(t, err) + require.Len(t, got, 1) + require.Same(t, theirs, got[0]) + require.Equal(t, []string{codec.SkipOpDecrypt + " " + codec.EncryptionEncoding}, skipped) + }) + } +} + +// TestEncryptorDecodeMixedOwnership shows a chained hop opens what it sealed and +// passes through what it didn't, in one batch. Breaking or returning on the +// pass-through instead of continuing the loop fails it. +func TestEncryptorDecodeMixedOwnership(t *testing.T) { + t.Parallel() + + original := testPayload("json/plain", `"ours"`) + want := proto.Clone(original).(*common.Payload) + + ours, err := codec.NewEncryptor(&fakeCipher{}).Encode([]*common.Payload{original}) + require.NoError(t, err) + theirs := sealedUnder(t, "other-proxy-kek", testPayload("json/plain", `"theirs"`)) + + var skipped []string + got, err := codec.NewEncryptor( + &fakeCipher{unknown: "other-proxy-kek"}, + codec.WithSkipEncodings(codec.EncryptionEncoding), + codec.WithSkipObserver(func(op, encoding string) { skipped = append(skipped, op+" "+encoding) }), + ).Decode([]*common.Payload{theirs, ours[0]}) + require.NoError(t, err) + require.Len(t, got, 2) + + require.Same(t, theirs, got[0]) + require.True(t, proto.Equal(want, got[1])) + require.Equal(t, []string{codec.SkipOpDecrypt + " " + codec.EncryptionEncoding}, skipped) +} + func (f *fakeCipher) Encrypt(data []byte) (*crypto.Message, error) { if f.encErr != nil { return nil, f.encErr @@ -268,6 +443,10 @@ func (f *fakeCipher) Encrypt(data []byte) (*crypto.Message, error) { func (f *fakeCipher) Decrypt(msg *crypto.Message) ([]byte, error) { f.decrypted = append(f.decrypted, msg) + if f.unknown != "" && msg.KeyMaterial.KEKID == f.unknown { + return nil, fmt.Errorf("%w: %s", crypto.ErrUnknownKey, f.unknown) + } + if f.decErr != nil { return nil, f.decErr } @@ -289,6 +468,20 @@ func testPayload(encoding, data string) *common.Payload { } } +// sealedUnder returns p sealed by fakeCipher, then relabeled as if a cipher +// holding kekID had sealed it, the way another proxy's output would arrive. +func sealedUnder(t *testing.T, kekID string, p *common.Payload) *common.Payload { + t.Helper() + + sealed, err := codec.NewEncryptor(&fakeCipher{}).Encode([]*common.Payload{p}) + require.NoError(t, err) + + out := proto.Clone(sealed[0]).(*common.Payload) + out.Metadata[codec.MetadataEncryptionKeyID] = []byte(kekID) + + return out +} + // markedPayload builds a payload carrying the sealed-payload encoding marker // plus md, so a test can vary the key material a marked payload claims to have. func markedPayload(md map[string][]byte) *common.Payload { diff --git a/pkg/crypto/kek.go b/pkg/crypto/kek.go index 7976938..97ffae2 100644 --- a/pkg/crypto/kek.go +++ b/pkg/crypto/kek.go @@ -10,6 +10,10 @@ import ( "sync" ) +// ErrUnknownKey is returned when a DEK names a KEK the registry holds neither as +// an active nor as a decrypt-only key. +var ErrUnknownKey = errors.New("unknown key") + type ( // KEK defines an interface for a Key Encryption Key. // These keys are used to encrypt/decrypt DEKs and are customer-managed (e.g. via AWS/GCP KMS). @@ -182,7 +186,7 @@ func (r *KEKRegistry) Decrypt(ctx context.Context, m *DEKMaterial) (*DEK, error) k, ok := r.keyIDs[m.KEKID] if !ok { - return nil, fmt.Errorf("unknown key: %s", m.KEKID) + return nil, fmt.Errorf("%w: %s", ErrUnknownKey, m.KEKID) } ct, err := base64.StdEncoding.DecodeString(m.EncryptedDEK) diff --git a/pkg/crypto/kek_test.go b/pkg/crypto/kek_test.go index 828fa97..157d87d 100644 --- a/pkg/crypto/kek_test.go +++ b/pkg/crypto/kek_test.go @@ -296,6 +296,20 @@ func TestKEKRegistryDecrypt(t *testing.T) { } } +// TestKEKRegistryDecryptUnknownKey pins the sentinel callers branch on. Returning +// fmt.Errorf without %w at the unknown-key site fails the ErrorIs check; changing +// the format string fails the EqualError check. +func TestKEKRegistryDecryptUnknownKey(t *testing.T) { + t.Parallel() + + r, err := crypto.NewKEKRegistry(crypto.WithDefaultKey(&fakeKEK{id: "default"})) + require.NoError(t, err) + + _, err = r.Decrypt(t.Context(), &crypto.DEKMaterial{KEKID: "missing"}) + require.ErrorIs(t, err, crypto.ErrUnknownKey) + require.EqualError(t, err, "unknown key: missing") +} + func TestKEKRegistryRoundtrip(t *testing.T) { t.Parallel() diff --git a/pkg/crypto/vault_test.go b/pkg/crypto/vault_test.go index 6bdc873..f531dc1 100644 --- a/pkg/crypto/vault_test.go +++ b/pkg/crypto/vault_test.go @@ -339,6 +339,24 @@ func TestVaultOpenErrors(t *testing.T) { }) } +// TestVaultOpenUnknownKey shows the sentinel survives Vault.Open, which is what +// the codec actually calls. Wrapping the registry error without %w anywhere in +// Open or open fails it. +func TestVaultOpenUnknownKey(t *testing.T) { + t.Parallel() + + ctx := t.Context() + cfg := crypto.WithKeyConfig("ns1", crypto.KeyConfig{Duration: time.Hour}) + sealer := newVault(t, &countingKEK{id: "sealer"}, cfg) + opener := newVault(t, &countingKEK{id: "opener"}, cfg) + + msg, err := sealer.Seal(ctx, "ns1", []byte("data")) + require.NoError(t, err) + + _, err = opener.Open(ctx, msg) + require.ErrorIs(t, err, crypto.ErrUnknownKey) +} + func TestNamespacedVault(t *testing.T) { t.Parallel()