From 3b12207c968b732d59ddc48e35b5a5d9abfd808e Mon Sep 17 00:00:00 2001 From: Haifeng He Date: Thu, 8 Oct 2026 09:17:06 -0700 Subject: [PATCH 1/2] Replace credentials.enabled with credentials.identity local.credentials.identity controls whose credentials the local Temporal server sees, in the "authorization" and "authorization-extras" headers: caller (default) forward whatever the caller sent; add nothing proxy drop forwarded credentials and send the CredentialProvider's instead none drop forwarded credentials and send none With identity "proxy", a replication stream forwarded from the remote side reaches the local server with only the proxy's token, never with a token the remote side sent. Forwarded credentials are removed by a client interceptor that runs before gRPC adds the per-RPC credentials, so the proxy's own token is never stripped. Validation rejects unknown identities, identity "proxy" on a non-TCP local connection, and any remote.credentials block. The proxy refuses to start with identity "proxy" and no CredentialProvider. Co-Authored-By: Claude Opus 5.5 --- auth/credential_provider.go | 4 + config/cluster_conn_config.go | 45 +++++++--- config/credentials_test.go | 42 +++++++--- examples/bearer-token/README.md | 23 ++++-- examples/bearer-token/config.yaml | 16 ++-- examples/bearer-token/main_test.go | 4 +- proxy/cluster_connection.go | 16 +++- proxy/credential_provider_test.go | 127 +++++++++++++++++++++-------- transport/grpcutil/grpc.go | 51 +++++++++++- transport/grpcutil/grpc_test.go | 30 +++++++ 10 files changed, 279 insertions(+), 79 deletions(-) create mode 100644 transport/grpcutil/grpc_test.go diff --git a/auth/credential_provider.go b/auth/credential_provider.go index 9770332..6a538c8 100644 --- a/auth/credential_provider.go +++ b/auth/credential_provider.go @@ -17,6 +17,10 @@ type ( EmptyCredentialProvider struct{} ) +// ForwardedCredentialHeaders are the headers that carry a caller's credentials. With credential identity "proxy" or +// "none", they are removed from every call the proxy forwards to the cluster. +var ForwardedCredentialHeaders = []string{"authorization", "authorization-extras"} + // Module provides the CredentialProvider. It defaults to EmptyCredentialProvider; override it with // WithCredentialProvider. var Module = fx.Options( diff --git a/config/cluster_conn_config.go b/config/cluster_conn_config.go index 2e73fd5..d2ee6d8 100644 --- a/config/cluster_conn_config.go +++ b/config/cluster_conn_config.go @@ -52,16 +52,20 @@ type ( TcpServer TCPTLSInfo `yaml:"tcpServer"` MuxCount int `yaml:"muxCount"` MuxAddressInfo TCPTLSInfo `yaml:"muxAddressInfo"` - // Credentials controls whether calls on this connection carry the auth.CredentialProvider's credentials. + // Credentials controls which identity calls on this connection present to the cluster. // Only the local cluster definition supports it. Credentials *CredentialsConfig `yaml:"credentials"` } CredentialsConfig struct { - // Enabled attaches the CredentialProvider's credentials to every call made to this cluster. - Enabled bool `yaml:"enabled"` + // Identity is whose credentials the cluster sees. Empty means CredentialIdentityCaller. + Identity CredentialIdentity `yaml:"identity"` } + // CredentialIdentity is whose credentials a cluster sees on calls from the proxy. The credential headers are + // "authorization" and "authorization-extras". + CredentialIdentity string + TCPTLSInfo struct { ConnectionString string `yaml:"address"` TLSConfig encryption.TLSConfig `yaml:"tls"` @@ -74,6 +78,15 @@ type ( } ) +const ( + // CredentialIdentityCaller forwards whatever credentials the caller sent, and adds none. This is the default. + CredentialIdentityCaller CredentialIdentity = "caller" + // CredentialIdentityProxy drops forwarded credentials and sends the auth.CredentialProvider's instead. + CredentialIdentityProxy CredentialIdentity = "proxy" + // CredentialIdentityNone drops forwarded credentials and sends none. + CredentialIdentityNone CredentialIdentity = "none" +) + const ( ConnTypeTCP ConnectionType = "tcp" ConnTypeMuxServer ConnectionType = "mux-server" @@ -98,9 +111,12 @@ func (config *StringTranslator) AsLocalToRemoteBiMap() (collect.StaticBiMap[stri return config.cachedBiMap, nil } -// CredentialsEnabled reports whether calls on this connection carry the CredentialProvider's credentials. -func (c ClusterDefinition) CredentialsEnabled() bool { - return c.Credentials != nil && c.Credentials.Enabled +// CredentialIdentity returns whose credentials this cluster sees, defaulting to CredentialIdentityCaller. +func (c ClusterDefinition) CredentialIdentity() CredentialIdentity { + if c.Credentials == nil || c.Credentials.Identity == "" { + return CredentialIdentityCaller + } + return c.Credentials.Identity } // Validate reports problems in this connection's config. Only the encryption @@ -109,11 +125,20 @@ func (c *ClusterConnConfig) Validate() error { return validation.Validate( "", validation.Nested("encryption", &c.EncryptionConfig), - validation.Field("local.credentials", c.Local, func(local ClusterDefinition) error { - if local.CredentialsEnabled() && local.ConnectionType != ConnTypeTCP { - return fmt.Errorf("credentials require connectionType %q, got %q", ConnTypeTCP, local.ConnectionType) + validation.Field("local.credentials.identity", c.Local, func(local ClusterDefinition) error { + switch identity := local.CredentialIdentity(); identity { + case CredentialIdentityCaller, CredentialIdentityNone: + return nil + case CredentialIdentityProxy: + if local.ConnectionType != ConnTypeTCP { + return fmt.Errorf("identity %q requires connectionType %q, got %q", + identity, ConnTypeTCP, local.ConnectionType) + } + return nil + default: + return fmt.Errorf("unsupported identity %q: must be %q, %q or %q", + identity, CredentialIdentityCaller, CredentialIdentityProxy, CredentialIdentityNone) } - return nil }), validation.Field("remote.credentials", c.Remote.Credentials, func(credentials *CredentialsConfig) error { if credentials != nil { diff --git a/config/credentials_test.go b/config/credentials_test.go index b111c2f..be4a295 100644 --- a/config/credentials_test.go +++ b/config/credentials_test.go @@ -7,28 +7,43 @@ import ( ) func TestCredentialsConfigValidate(t *testing.T) { - enabled := &CredentialsConfig{Enabled: true} + identity := func(identity CredentialIdentity) *CredentialsConfig { + return &CredentialsConfig{Identity: identity} + } tests := []struct { name string conn ClusterConnConfig wantError string }{ { - name: "local tcp with credentials", - conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeTCP, Credentials: enabled}}, + name: "no credentials block", + conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeMuxClient}}, + }, + { + name: "local tcp with proxy identity", + conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeTCP, Credentials: identity(CredentialIdentityProxy)}}, + }, + { + name: "local mux with caller identity", + conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeMuxClient, Credentials: identity(CredentialIdentityCaller)}}, + }, + { + name: "local mux with none identity", + conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeMuxClient, Credentials: identity(CredentialIdentityNone)}}, }, { - name: "local mux with credentials disabled", - conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeMuxClient, Credentials: &CredentialsConfig{}}}, + name: "local mux with proxy identity", + conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeMuxClient, Credentials: identity(CredentialIdentityProxy)}}, + wantError: `identity "proxy" requires connectionType "tcp", got "mux-client"`, }, { - name: "local mux with credentials enabled", - conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeMuxClient, Credentials: enabled}}, - wantError: `credentials require connectionType "tcp", got "mux-client"`, + name: "unsupported identity", + conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeTCP, Credentials: identity("everyone")}}, + wantError: `unsupported identity "everyone"`, }, { name: "remote with credentials", - conn: ClusterConnConfig{Remote: ClusterDefinition{ConnectionType: ConnTypeTCP, Credentials: enabled}}, + conn: ClusterConnConfig{Remote: ClusterDefinition{ConnectionType: ConnTypeTCP, Credentials: identity(CredentialIdentityProxy)}}, wantError: "credentials are only supported on the local cluster definition", }, } @@ -44,8 +59,9 @@ func TestCredentialsConfigValidate(t *testing.T) { } } -func TestCredentialsEnabled(t *testing.T) { - require.False(t, ClusterDefinition{}.CredentialsEnabled()) - require.False(t, ClusterDefinition{Credentials: &CredentialsConfig{}}.CredentialsEnabled()) - require.True(t, ClusterDefinition{Credentials: &CredentialsConfig{Enabled: true}}.CredentialsEnabled()) +func TestCredentialIdentityDefaultsToCaller(t *testing.T) { + require.Equal(t, CredentialIdentityCaller, ClusterDefinition{}.CredentialIdentity()) + require.Equal(t, CredentialIdentityCaller, ClusterDefinition{Credentials: &CredentialsConfig{}}.CredentialIdentity()) + require.Equal(t, CredentialIdentityProxy, + ClusterDefinition{Credentials: &CredentialsConfig{Identity: CredentialIdentityProxy}}.CredentialIdentity()) } diff --git a/examples/bearer-token/README.md b/examples/bearer-token/README.md index 692b1c7..0a6de16 100644 --- a/examples/bearer-token/README.md +++ b/examples/bearer-token/README.md @@ -21,7 +21,7 @@ app.New("s2s-proxy-bearer-example", "dev", ) ``` -The provider is used only where the config turns it on, with `local.credentials.enabled`: +The provider is used only where the config turns it on, with `local.credentials.identity: proxy`: ```yaml local: @@ -30,13 +30,24 @@ local: address: temporal-frontend.example.invalid:7233 tls: { ... } credentials: - enabled: true + identity: proxy ``` -The program supplies how to get a token, and the config decides whether to send it. If the config enables -credentials but the binary has no provider, for example the stock `s2s-proxy`, the proxy refuses to start rather than -call the server without a token. A provider in the binary does nothing for a connection that does not enable it. -`credentials` is only accepted on the local cluster definition. +`identity` is whose credentials the local server sees, in the `authorization` and `authorization-extras` headers: + +| `identity` | Credentials the caller sent | Proxy's credentials | +| --- | --- | --- | +| `caller` (default) | forwarded | not sent | +| `proxy` | dropped | sent | +| `none` | dropped | not sent | + +With `proxy`, a request forwarded from the remote side, such as a replication stream, reaches the local server with +the proxy's token and never with a token the remote side sent. + +The program supplies how to get a token, and the config decides whether to send it. If the config sets +`identity: proxy` but the binary has no provider, for example the stock `s2s-proxy`, the proxy refuses to start rather +than call the server without a token. A provider in the binary does nothing for a connection that does not set +`identity: proxy`. `credentials` is only accepted on the local cluster definition. `Get` returns a gRPC `credentials.PerRPCCredentials`. gRPC calls its `GetRequestMetadata` for every unary call and every new stream, so a rotated token takes effect on the next call without restarting the proxy or recreating connections. diff --git a/examples/bearer-token/config.yaml b/examples/bearer-token/config.yaml index de551da..8d8d9a8 100644 --- a/examples/bearer-token/config.yaml +++ b/examples/bearer-token/config.yaml @@ -1,8 +1,8 @@ # s2s-proxy configuration for the bearer-token example. # -# local.credentials.enabled attaches the bearer token to every call the proxy makes to the local Temporal frontend -# (local.tcpClient). It is never sent to the remote side. The token is read from S2S_PROXY_EXAMPLE_BEARER_TOKEN, not -# from this file. +# local.credentials.identity: proxy attaches the bearer token to every call the proxy makes to the local Temporal +# frontend (local.tcpClient). It is never sent to the remote side. The token is read from +# S2S_PROXY_EXAMPLE_BEARER_TOKEN, not from this file. # # Hostnames use the reserved .invalid TLD and certificate paths are placeholders; replace both for a real deployment. clusterConnections: @@ -16,9 +16,15 @@ clusterConnections: remoteCAPath: /etc/s2s-proxy/tls/frontend-ca.pem caServerName: temporal-frontend.example.invalid skipCAVerification: false - # Send the CredentialProvider's credentials on calls to the local Temporal server. + # Whose credentials the local Temporal server sees on calls from the proxy, carried in the "authorization" and + # "authorization-extras" headers: + # caller - (default) forward whatever the caller sent; the proxy adds nothing. + # proxy - drop any credentials the caller sent and send the proxy's own, from its CredentialProvider. The + # proxy refuses to start if the binary has no CredentialProvider. + # none - drop any credentials the caller sent and send none. + # Only the local cluster definition accepts credentials. credentials: - enabled: true + identity: proxy tcpServer: address: 0.0.0.0:9002 tls: diff --git a/examples/bearer-token/main_test.go b/examples/bearer-token/main_test.go index 6978a1c..59cb163 100644 --- a/examples/bearer-token/main_test.go +++ b/examples/bearer-token/main_test.go @@ -49,9 +49,9 @@ func TestExampleConfiguration(t *testing.T) { require.NoError(t, proxyConfig.Validate()) require.Len(t, proxyConfig.ClusterConnections, 1) - // Credentials are enabled on the local client, which must be TCP with TLS because the token requires it. + // The local client presents the proxy's identity, so it must be TCP with TLS because the token requires it. local := proxyConfig.ClusterConnections[0].Local - require.True(t, local.CredentialsEnabled()) + require.Equal(t, config.CredentialIdentityProxy, local.CredentialIdentity()) require.Equal(t, config.ConnTypeTCP, local.ConnectionType) require.True(t, local.TcpClient.TLSConfig.IsEnabled()) } diff --git a/proxy/cluster_connection.go b/proxy/cluster_connection.go index cef2918..5f164aa 100644 --- a/proxy/cluster_connection.go +++ b/proxy/cluster_connection.go @@ -130,7 +130,7 @@ func NewClusterConnection(lifetime context.Context, connConfig config.ClusterCon } // newClusterConnection is NewClusterConnection with a CredentialProvider. Its credentials are attached to calls made to -// the local Temporal server when local.credentials.enabled is set, and never to the remote side. +// the local Temporal server when local.credentials.identity is "proxy", and never to the remote side. func newClusterConnection( lifetime context.Context, connConfig config.ClusterConnConfig, @@ -258,10 +258,16 @@ func createClient( credentialProvider auth.CredentialProvider, ) (closableClientConn, error) { var clientOptions grpcutil.ClientOptions - if transportCfg.CredentialsEnabled() { + switch identity := transportCfg.CredentialIdentity(); identity { + case config.CredentialIdentityCaller: + // Forward whatever the caller sent. + case config.CredentialIdentityNone: + clientOptions.StripOutgoingMetadataKeys = auth.ForwardedCredentialHeaders + case config.CredentialIdentityProxy: if auth.IsEmptyCredentialProvider(credentialProvider) { - return nil, fmt.Errorf("%s client: credentials are enabled but no CredentialProvider is configured", directionLabel) + return nil, fmt.Errorf("%s client: credentials identity %q but no CredentialProvider is configured", directionLabel, identity) } + clientOptions.StripOutgoingMetadataKeys = auth.ForwardedCredentialHeaders clientOptions.PerRPCCredentials = credentialProvider.Get() if clientOptions.PerRPCCredentials == nil { return nil, fmt.Errorf("%s client: credential provider returned no credentials", directionLabel) @@ -272,6 +278,8 @@ func createClient( if clientOptions.PerRPCCredentials.RequireTransportSecurity() && !transportCfg.TcpClient.TLSConfig.IsEnabled() { return nil, fmt.Errorf("%s client: credentials require TLS, but tcpClient.tls is not configured", directionLabel) } + default: + return nil, fmt.Errorf("%s client: unsupported credentials identity %q", directionLabel, identity) } switch transportCfg.ConnectionType { @@ -280,7 +288,7 @@ func createClient( case config.ConnTypeMuxClient, config.ConnTypeMuxServer: return grpcutil.NewMultiClientConn(lifetime, fmt.Sprintf("client-conn-%s", connectionName), // TLS is handled by the mux connection, so tlsConfig will always be nil - grpcutil.MakeDialOptions(nil, metrics.GetGRPCClientMetrics(directionLabel))...) + grpcutil.MakeDialOptions(nil, metrics.GetGRPCClientMetrics(directionLabel), clientOptions)...) default: return nil, errors.New("invalid connection type") } diff --git a/proxy/credential_provider_test.go b/proxy/credential_provider_test.go index 7448ef5..4b18a5c 100644 --- a/proxy/credential_provider_test.go +++ b/proxy/credential_provider_test.go @@ -23,7 +23,12 @@ import ( "github.com/temporalio/s2s-proxy/logging" ) -const testBearerToken = "Bearer local-token" +const ( + testBearerToken = "Bearer local-token" + peerBearerToken = "Bearer peer-token" + + describeClusterMethod = "/temporal.server.api.adminservice.v1.AdminService/DescribeCluster" +) type staticCredentials struct { requireTLS bool @@ -41,22 +46,31 @@ type staticCredentialProvider struct { func (p staticCredentialProvider) Get() credentials.PerRPCCredentials { return p.creds } -// authRecordingServer is a fake Temporal server that records the authorization header of every call it receives, -// for any service, and answers Unimplemented. +// recordedCredentials are the credential headers one call arrived with. +type recordedCredentials struct { + authorization []string + authorizationExtras []string +} + +// authRecordingServer is a fake Temporal server that records the credential headers of every call it receives, for +// any service, and answers Unimplemented. type authRecordingServer struct { mu sync.Mutex - headers map[string][]string // method -> authorization values + headers map[string]recordedCredentials // method -> credential headers } func startAuthRecordingServer(t *testing.T, address string) *authRecordingServer { - s := &authRecordingServer{headers: make(map[string][]string)} + s := &authRecordingServer{headers: make(map[string]recordedCredentials)} listener, err := net.Listen("tcp", address) require.NoError(t, err) server := grpc.NewServer(grpc.UnknownServiceHandler(func(_ any, stream grpc.ServerStream) error { method, _ := grpc.MethodFromServerStream(stream) md, _ := metadata.FromIncomingContext(stream.Context()) s.mu.Lock() - s.headers[method] = md.Get("authorization") + s.headers[method] = recordedCredentials{ + authorization: md.Get("authorization"), + authorizationExtras: md.Get("authorization-extras"), + } s.mu.Unlock() return status.Error(codes.Unimplemented, "recorded") })) @@ -65,36 +79,45 @@ func startAuthRecordingServer(t *testing.T, address string) *authRecordingServer return s } -func (s *authRecordingServer) authorization(method string) ([]string, bool) { +func (s *authRecordingServer) credentials(t *testing.T, method string) recordedCredentials { s.mu.Lock() defer s.mu.Unlock() - values, ok := s.headers[method] - return values, ok + recorded, ok := s.headers[method] + require.True(t, ok, "server never saw %s", method) + return recorded } func newCredentialTestConnection( t *testing.T, a plccAddresses, - credentialsEnabled bool, + identity config.CredentialIdentity, provider auth.CredentialProvider, ) *ClusterConnection { connConfig := makeTCPClusterConfig("creds", localFVI, remoteFVI, "", a.localTemporalAddr, a.localProxyOutbound, a.localProxyInbound, a.remoteTemporalAddr) - connConfig.Local.Credentials = &config.CredentialsConfig{Enabled: credentialsEnabled} + connConfig.Local.Credentials = &config.CredentialsConfig{Identity: identity} loggers := logging.NewLoggerProvider(log.NewTestLogger(), config.NewMockConfigProvider(config.S2SProxyConfig{})) cc, err := newClusterConnection(t.Context(), connConfig, loggers, provider) require.NoError(t, err) return cc } -func TestCredentialsAreAttachedToLocalCallsOnly(t *testing.T) { +// forwardedContext carries the credential headers a peer sent, the way the replication stream forwarder passes them on. +func forwardedContext(t *testing.T) context.Context { + return metadata.NewOutgoingContext(t.Context(), metadata.Pairs( + "authorization", peerBearerToken, + "authorization-extras", "peer-extras", + )) +} + +func TestProxyIdentityReplacesForwardedCredentials(t *testing.T) { a := getDynamicPlccAddresses(t) localTemporal := startAuthRecordingServer(t, a.localTemporalAddr) remoteTemporal := startAuthRecordingServer(t, a.remoteTemporalAddr) - cc := newCredentialTestConnection(t, a, true, staticCredentialProvider{creds: staticCredentials{}}) + cc := newCredentialTestConnection(t, a, config.CredentialIdentityProxy, staticCredentialProvider{creds: staticCredentials{}}) - // Every service, unary and streaming, made on the local client carries the token. - ctx := t.Context() + // Every service, unary and streaming, carries only the proxy's token, whatever the peer sent. + ctx := forwardedContext(t) _, _ = adminservice.NewAdminServiceClient(cc.inboundClient).DescribeCluster(ctx, &adminservice.DescribeClusterRequest{}) _, _ = workflowservice.NewWorkflowServiceClient(cc.inboundClient).GetSystemInfo(ctx, &workflowservice.GetSystemInfoRequest{}) _, _ = operatorservice.NewOperatorServiceClient(cc.inboundClient).ListClusters(ctx, &operatorservice.ListClustersRequest{}) @@ -103,42 +126,57 @@ func TestCredentialsAreAttachedToLocalCallsOnly(t *testing.T) { _, _ = stream.Recv() for _, method := range []string{ - "/temporal.server.api.adminservice.v1.AdminService/DescribeCluster", + describeClusterMethod, "/temporal.api.workflowservice.v1.WorkflowService/GetSystemInfo", "/temporal.api.operatorservice.v1.OperatorService/ListClusters", "/temporal.server.api.adminservice.v1.AdminService/StreamWorkflowReplicationMessages", } { - values, ok := localTemporal.authorization(method) - require.True(t, ok, "local server never saw %s", method) - require.Equal(t, []string{testBearerToken}, values, method) + recorded := localTemporal.credentials(t, method) + require.Equal(t, []string{testBearerToken}, recorded.authorization, method) + require.Empty(t, recorded.authorizationExtras, method) } - // The remote side never sees the local token. - _, _ = adminservice.NewAdminServiceClient(cc.outboundClient).DescribeCluster(ctx, &adminservice.DescribeClusterRequest{}) - values, ok := remoteTemporal.authorization("/temporal.server.api.adminservice.v1.AdminService/DescribeCluster") - require.True(t, ok) - require.Empty(t, values) + // The remote side never sees the proxy's token. + _, _ = adminservice.NewAdminServiceClient(cc.outboundClient).DescribeCluster(t.Context(), &adminservice.DescribeClusterRequest{}) + require.Empty(t, remoteTemporal.credentials(t, describeClusterMethod).authorization) } -func TestCredentialsAreNotAttachedUnlessEnabled(t *testing.T) { +func TestNoneIdentityStripsForwardedCredentials(t *testing.T) { a := getDynamicPlccAddresses(t) localTemporal := startAuthRecordingServer(t, a.localTemporalAddr) startAuthRecordingServer(t, a.remoteTemporalAddr) - cc := newCredentialTestConnection(t, a, false, staticCredentialProvider{creds: staticCredentials{}}) + // A provider is configured but unused: identity "none" never sends credentials. + cc := newCredentialTestConnection(t, a, config.CredentialIdentityNone, staticCredentialProvider{creds: staticCredentials{}}) - _, _ = adminservice.NewAdminServiceClient(cc.inboundClient).DescribeCluster(t.Context(), &adminservice.DescribeClusterRequest{}) - values, ok := localTemporal.authorization("/temporal.server.api.adminservice.v1.AdminService/DescribeCluster") - require.True(t, ok) - require.Empty(t, values) + _, _ = adminservice.NewAdminServiceClient(cc.inboundClient).DescribeCluster(forwardedContext(t), &adminservice.DescribeClusterRequest{}) + recorded := localTemporal.credentials(t, describeClusterMethod) + require.Empty(t, recorded.authorization) + require.Empty(t, recorded.authorizationExtras) +} + +func TestCallerIdentityForwardsCredentials(t *testing.T) { + for _, identity := range []config.CredentialIdentity{"", config.CredentialIdentityCaller} { + t.Run("identity="+string(identity), func(t *testing.T) { + a := getDynamicPlccAddresses(t) + localTemporal := startAuthRecordingServer(t, a.localTemporalAddr) + startAuthRecordingServer(t, a.remoteTemporalAddr) + cc := newCredentialTestConnection(t, a, identity, staticCredentialProvider{creds: staticCredentials{}}) + + _, _ = adminservice.NewAdminServiceClient(cc.inboundClient).DescribeCluster(forwardedContext(t), &adminservice.DescribeClusterRequest{}) + recorded := localTemporal.credentials(t, describeClusterMethod) + require.Equal(t, []string{peerBearerToken}, recorded.authorization) + require.Equal(t, []string{"peer-extras"}, recorded.authorizationExtras) + }) + } } func TestCreateClientRejectsUnusableCredentials(t *testing.T) { - enabled := &config.CredentialsConfig{Enabled: true} + proxyIdentity := &config.CredentialsConfig{Identity: config.CredentialIdentityProxy} tcp := func(tls encryption.TLSConfig) config.ClusterDefinition { return config.ClusterDefinition{ ConnectionType: config.ConnTypeTCP, TcpClient: config.TCPTLSInfo{ConnectionString: "localhost:7233", TLSConfig: tls}, - Credentials: enabled, + Credentials: proxyIdentity, } } tests := []struct { @@ -148,10 +186,10 @@ func TestCreateClientRejectsUnusableCredentials(t *testing.T) { wantError string }{ { - name: "credentials enabled without a provider", + name: "proxy identity without a provider", cluster: tcp(encryption.TLSConfig{}), provider: auth.EmptyCredentialProvider{}, - wantError: "credentials are enabled but no CredentialProvider is configured", + wantError: `credentials identity "proxy" but no CredentialProvider is configured`, }, { name: "provider returns no credentials", @@ -161,7 +199,7 @@ func TestCreateClientRejectsUnusableCredentials(t *testing.T) { }, { name: "mux connection", - cluster: config.ClusterDefinition{ConnectionType: config.ConnTypeMuxClient, Credentials: enabled}, + cluster: config.ClusterDefinition{ConnectionType: config.ConnTypeMuxClient, Credentials: proxyIdentity}, provider: staticCredentialProvider{creds: staticCredentials{}}, wantError: "credentials require a tcp connection", }, @@ -171,6 +209,15 @@ func TestCreateClientRejectsUnusableCredentials(t *testing.T) { provider: staticCredentialProvider{creds: staticCredentials{requireTLS: true}}, wantError: "credentials require TLS", }, + { + name: "unsupported identity", + cluster: config.ClusterDefinition{ + ConnectionType: config.ConnTypeTCP, + Credentials: &config.CredentialsConfig{Identity: "everyone"}, + }, + provider: auth.EmptyCredentialProvider{}, + wantError: `unsupported credentials identity "everyone"`, + }, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { @@ -179,9 +226,17 @@ func TestCreateClientRejectsUnusableCredentials(t *testing.T) { }) } - t.Run("empty provider leaves a mux connection alone", func(t *testing.T) { + t.Run("caller identity leaves a mux connection alone", func(t *testing.T) { _, err := createClient(t.Context(), "test", config.ClusterDefinition{ConnectionType: config.ConnTypeMuxClient}, "inbound", auth.EmptyCredentialProvider{}) require.NoError(t, err) }) + + t.Run("none identity works on a mux connection", func(t *testing.T) { + _, err := createClient(t.Context(), "test", config.ClusterDefinition{ + ConnectionType: config.ConnTypeMuxClient, + Credentials: &config.CredentialsConfig{Identity: config.CredentialIdentityNone}, + }, "inbound", auth.EmptyCredentialProvider{}) + require.NoError(t, err) + }) } diff --git a/transport/grpcutil/grpc.go b/transport/grpcutil/grpc.go index bd1b1a2..a84cda9 100644 --- a/transport/grpcutil/grpc.go +++ b/transport/grpcutil/grpc.go @@ -1,6 +1,7 @@ package grpcutil import ( + "context" "crypto/tls" "time" @@ -10,6 +11,7 @@ import ( "google.golang.org/grpc/credentials" "google.golang.org/grpc/credentials/insecure" "google.golang.org/grpc/encoding" + "google.golang.org/grpc/metadata" "github.com/temporalio/s2s-proxy/proto/compat" ) @@ -34,6 +36,9 @@ const ( type ClientOptions struct { // PerRPCCredentials, when set, are attached to every call made on the connection. PerRPCCredentials credentials.PerRPCCredentials + // StripOutgoingMetadataKeys are removed from the outgoing metadata of every call made on the connection. + // PerRPCCredentials are added after this runs, so they are never removed. + StripOutgoingMetadataKeys []string } func MakeDialOptions(tlsConfig *tls.Config, clientMetrics *grpcprom.ClientMetrics, clientOptions ...ClientOptions) []grpc.DialOption { @@ -64,13 +69,53 @@ func MakeDialOptions(tlsConfig *tls.Config, clientMetrics *grpcprom.ClientMetric grpc.WithDefaultServiceConfig(DefaultServiceConfig), grpc.WithDisableServiceConfig(), grpc.WithConnectParams(cp), - grpc.WithUnaryInterceptor(clientMetrics.UnaryClientInterceptor()), - grpc.WithStreamInterceptor(clientMetrics.StreamClientInterceptor()), } + + var unaryInterceptors []grpc.UnaryClientInterceptor + var streamInterceptors []grpc.StreamClientInterceptor for _, options := range clientOptions { + if len(options.StripOutgoingMetadataKeys) > 0 { + unaryInterceptors = append(unaryInterceptors, stripUnaryInterceptor(options.StripOutgoingMetadataKeys)) + streamInterceptors = append(streamInterceptors, stripStreamInterceptor(options.StripOutgoingMetadataKeys)) + } if options.PerRPCCredentials != nil { dialOptions = append(dialOptions, grpc.WithPerRPCCredentials(options.PerRPCCredentials)) } } - return dialOptions + unaryInterceptors = append(unaryInterceptors, clientMetrics.UnaryClientInterceptor()) + streamInterceptors = append(streamInterceptors, clientMetrics.StreamClientInterceptor()) + return append(dialOptions, + grpc.WithChainUnaryInterceptor(unaryInterceptors...), + grpc.WithChainStreamInterceptor(streamInterceptors...), + ) +} + +// stripOutgoingMetadata returns ctx with keys removed from its outgoing metadata. The caller's metadata is not +// modified. +func stripOutgoingMetadata(ctx context.Context, keys []string) context.Context { + md, ok := metadata.FromOutgoingContext(ctx) + if !ok { + return ctx + } + md = md.Copy() + for _, key := range keys { + md.Delete(key) + } + return metadata.NewOutgoingContext(ctx, md) +} + +func stripUnaryInterceptor(keys []string) grpc.UnaryClientInterceptor { + return func(ctx context.Context, method string, req, reply any, cc *grpc.ClientConn, + invoker grpc.UnaryInvoker, opts ...grpc.CallOption, + ) error { + return invoker(stripOutgoingMetadata(ctx, keys), method, req, reply, cc, opts...) + } +} + +func stripStreamInterceptor(keys []string) grpc.StreamClientInterceptor { + return func(ctx context.Context, desc *grpc.StreamDesc, cc *grpc.ClientConn, method string, + streamer grpc.Streamer, opts ...grpc.CallOption, + ) (grpc.ClientStream, error) { + return streamer(stripOutgoingMetadata(ctx, keys), desc, cc, method, opts...) + } } diff --git a/transport/grpcutil/grpc_test.go b/transport/grpcutil/grpc_test.go new file mode 100644 index 0000000..4262dbc --- /dev/null +++ b/transport/grpcutil/grpc_test.go @@ -0,0 +1,30 @@ +package grpcutil + +import ( + "testing" + + "github.com/stretchr/testify/require" + "google.golang.org/grpc/metadata" +) + +func TestStripOutgoingMetadata(t *testing.T) { + original := metadata.Pairs( + "authorization", "Bearer peer", + "Authorization-Extras", "extras", + "temporal-client-shard-id", "4", + ) + ctx := metadata.NewOutgoingContext(t.Context(), original) + + stripped, ok := metadata.FromOutgoingContext(stripOutgoingMetadata(ctx, []string{"authorization", "authorization-extras"})) + require.True(t, ok) + require.Equal(t, metadata.Pairs("temporal-client-shard-id", "4"), stripped) + + // The caller's metadata is left alone. + require.Equal(t, []string{"Bearer peer"}, original.Get("authorization")) + require.Equal(t, []string{"extras"}, original.Get("authorization-extras")) +} + +func TestStripOutgoingMetadataWithoutMetadata(t *testing.T) { + ctx := t.Context() + require.Equal(t, ctx, stripOutgoingMetadata(ctx, []string{"authorization"})) +} From e1e5a5cfba60a2e9e8ec50f126123b978f24e469 Mon Sep 17 00:00:00 2001 From: Haifeng He Date: Thu, 8 Oct 2026 09:34:53 -0700 Subject: [PATCH 2/2] Rename credential identities to default, proxy and strip Rename CredentialIdentityCaller to CredentialIdentityDefault ("default"), which forwards the caller's credentials and applies when identity is not set, and CredentialIdentityNone to CredentialIdentityStrip ("strip"). Rename auth.ForwardedCredentialHeaders to auth.CredentialHeaders. Co-Authored-By: Claude Opus 5.5 --- auth/credential_provider.go | 6 +++--- config/cluster_conn_config.go | 19 ++++++++++--------- config/credentials_test.go | 14 +++++++------- examples/bearer-token/README.md | 4 ++-- examples/bearer-token/config.yaml | 8 ++++---- proxy/cluster_connection.go | 10 +++++----- proxy/credential_provider_test.go | 16 ++++++++-------- 7 files changed, 39 insertions(+), 38 deletions(-) diff --git a/auth/credential_provider.go b/auth/credential_provider.go index 6a538c8..51e698e 100644 --- a/auth/credential_provider.go +++ b/auth/credential_provider.go @@ -17,9 +17,9 @@ type ( EmptyCredentialProvider struct{} ) -// ForwardedCredentialHeaders are the headers that carry a caller's credentials. With credential identity "proxy" or -// "none", they are removed from every call the proxy forwards to the cluster. -var ForwardedCredentialHeaders = []string{"authorization", "authorization-extras"} +// CredentialHeaders are the headers that carry a caller's credentials. With credential identity "proxy" or +// "strip", they are removed from every call the proxy forwards to the cluster. +var CredentialHeaders = []string{"authorization", "authorization-extras"} // Module provides the CredentialProvider. It defaults to EmptyCredentialProvider; override it with // WithCredentialProvider. diff --git a/config/cluster_conn_config.go b/config/cluster_conn_config.go index d2ee6d8..8dc1a54 100644 --- a/config/cluster_conn_config.go +++ b/config/cluster_conn_config.go @@ -58,7 +58,7 @@ type ( } CredentialsConfig struct { - // Identity is whose credentials the cluster sees. Empty means CredentialIdentityCaller. + // Identity is whose credentials the cluster sees. Empty means CredentialIdentityDefault. Identity CredentialIdentity `yaml:"identity"` } @@ -79,12 +79,13 @@ type ( ) const ( - // CredentialIdentityCaller forwards whatever credentials the caller sent, and adds none. This is the default. - CredentialIdentityCaller CredentialIdentity = "caller" + // CredentialIdentityDefault forwards the caller's credentials: whatever the caller sent is passed through + // unchanged, and the proxy adds none. It applies when identity is not set. + CredentialIdentityDefault CredentialIdentity = "default" // CredentialIdentityProxy drops forwarded credentials and sends the auth.CredentialProvider's instead. CredentialIdentityProxy CredentialIdentity = "proxy" - // CredentialIdentityNone drops forwarded credentials and sends none. - CredentialIdentityNone CredentialIdentity = "none" + // CredentialIdentityStrip drops forwarded credentials and sends none. + CredentialIdentityStrip CredentialIdentity = "strip" ) const ( @@ -111,10 +112,10 @@ func (config *StringTranslator) AsLocalToRemoteBiMap() (collect.StaticBiMap[stri return config.cachedBiMap, nil } -// CredentialIdentity returns whose credentials this cluster sees, defaulting to CredentialIdentityCaller. +// CredentialIdentity returns whose credentials this cluster sees, defaulting to CredentialIdentityDefault. func (c ClusterDefinition) CredentialIdentity() CredentialIdentity { if c.Credentials == nil || c.Credentials.Identity == "" { - return CredentialIdentityCaller + return CredentialIdentityDefault } return c.Credentials.Identity } @@ -127,7 +128,7 @@ func (c *ClusterConnConfig) Validate() error { validation.Nested("encryption", &c.EncryptionConfig), validation.Field("local.credentials.identity", c.Local, func(local ClusterDefinition) error { switch identity := local.CredentialIdentity(); identity { - case CredentialIdentityCaller, CredentialIdentityNone: + case CredentialIdentityDefault, CredentialIdentityStrip: return nil case CredentialIdentityProxy: if local.ConnectionType != ConnTypeTCP { @@ -137,7 +138,7 @@ func (c *ClusterConnConfig) Validate() error { return nil default: return fmt.Errorf("unsupported identity %q: must be %q, %q or %q", - identity, CredentialIdentityCaller, CredentialIdentityProxy, CredentialIdentityNone) + identity, CredentialIdentityDefault, CredentialIdentityProxy, CredentialIdentityStrip) } }), validation.Field("remote.credentials", c.Remote.Credentials, func(credentials *CredentialsConfig) error { diff --git a/config/credentials_test.go b/config/credentials_test.go index be4a295..1c6c5c4 100644 --- a/config/credentials_test.go +++ b/config/credentials_test.go @@ -24,12 +24,12 @@ func TestCredentialsConfigValidate(t *testing.T) { conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeTCP, Credentials: identity(CredentialIdentityProxy)}}, }, { - name: "local mux with caller identity", - conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeMuxClient, Credentials: identity(CredentialIdentityCaller)}}, + name: "local mux with default identity", + conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeMuxClient, Credentials: identity(CredentialIdentityDefault)}}, }, { - name: "local mux with none identity", - conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeMuxClient, Credentials: identity(CredentialIdentityNone)}}, + name: "local mux with strip identity", + conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeMuxClient, Credentials: identity(CredentialIdentityStrip)}}, }, { name: "local mux with proxy identity", @@ -59,9 +59,9 @@ func TestCredentialsConfigValidate(t *testing.T) { } } -func TestCredentialIdentityDefaultsToCaller(t *testing.T) { - require.Equal(t, CredentialIdentityCaller, ClusterDefinition{}.CredentialIdentity()) - require.Equal(t, CredentialIdentityCaller, ClusterDefinition{Credentials: &CredentialsConfig{}}.CredentialIdentity()) +func TestCredentialIdentityDefault(t *testing.T) { + require.Equal(t, CredentialIdentityDefault, ClusterDefinition{}.CredentialIdentity()) + require.Equal(t, CredentialIdentityDefault, ClusterDefinition{Credentials: &CredentialsConfig{}}.CredentialIdentity()) require.Equal(t, CredentialIdentityProxy, ClusterDefinition{Credentials: &CredentialsConfig{Identity: CredentialIdentityProxy}}.CredentialIdentity()) } diff --git a/examples/bearer-token/README.md b/examples/bearer-token/README.md index 0a6de16..4bc8381 100644 --- a/examples/bearer-token/README.md +++ b/examples/bearer-token/README.md @@ -37,9 +37,9 @@ local: | `identity` | Credentials the caller sent | Proxy's credentials | | --- | --- | --- | -| `caller` (default) | forwarded | not sent | +| `default` (when not set) | forwarded | not sent | | `proxy` | dropped | sent | -| `none` | dropped | not sent | +| `strip` | dropped | not sent | With `proxy`, a request forwarded from the remote side, such as a replication stream, reaches the local server with the proxy's token and never with a token the remote side sent. diff --git a/examples/bearer-token/config.yaml b/examples/bearer-token/config.yaml index 8d8d9a8..d8a6ab8 100644 --- a/examples/bearer-token/config.yaml +++ b/examples/bearer-token/config.yaml @@ -18,10 +18,10 @@ clusterConnections: skipCAVerification: false # Whose credentials the local Temporal server sees on calls from the proxy, carried in the "authorization" and # "authorization-extras" headers: - # caller - (default) forward whatever the caller sent; the proxy adds nothing. - # proxy - drop any credentials the caller sent and send the proxy's own, from its CredentialProvider. The - # proxy refuses to start if the binary has no CredentialProvider. - # none - drop any credentials the caller sent and send none. + # default - (when identity is not set) forward the caller's credentials unchanged; the proxy adds nothing. + # proxy - drop any credentials the caller sent and send the proxy's own, from its CredentialProvider. The + # proxy refuses to start if the binary has no CredentialProvider. + # strip - drop any credentials the caller sent and send none. # Only the local cluster definition accepts credentials. credentials: identity: proxy diff --git a/proxy/cluster_connection.go b/proxy/cluster_connection.go index 5f164aa..b4cfef5 100644 --- a/proxy/cluster_connection.go +++ b/proxy/cluster_connection.go @@ -259,15 +259,15 @@ func createClient( ) (closableClientConn, error) { var clientOptions grpcutil.ClientOptions switch identity := transportCfg.CredentialIdentity(); identity { - case config.CredentialIdentityCaller: - // Forward whatever the caller sent. - case config.CredentialIdentityNone: - clientOptions.StripOutgoingMetadataKeys = auth.ForwardedCredentialHeaders + case config.CredentialIdentityDefault: + // Forward the caller's credentials. + case config.CredentialIdentityStrip: + clientOptions.StripOutgoingMetadataKeys = auth.CredentialHeaders case config.CredentialIdentityProxy: if auth.IsEmptyCredentialProvider(credentialProvider) { return nil, fmt.Errorf("%s client: credentials identity %q but no CredentialProvider is configured", directionLabel, identity) } - clientOptions.StripOutgoingMetadataKeys = auth.ForwardedCredentialHeaders + clientOptions.StripOutgoingMetadataKeys = auth.CredentialHeaders clientOptions.PerRPCCredentials = credentialProvider.Get() if clientOptions.PerRPCCredentials == nil { return nil, fmt.Errorf("%s client: credential provider returned no credentials", directionLabel) diff --git a/proxy/credential_provider_test.go b/proxy/credential_provider_test.go index 4b18a5c..099413d 100644 --- a/proxy/credential_provider_test.go +++ b/proxy/credential_provider_test.go @@ -141,12 +141,12 @@ func TestProxyIdentityReplacesForwardedCredentials(t *testing.T) { require.Empty(t, remoteTemporal.credentials(t, describeClusterMethod).authorization) } -func TestNoneIdentityStripsForwardedCredentials(t *testing.T) { +func TestStripIdentityStripsForwardedCredentials(t *testing.T) { a := getDynamicPlccAddresses(t) localTemporal := startAuthRecordingServer(t, a.localTemporalAddr) startAuthRecordingServer(t, a.remoteTemporalAddr) - // A provider is configured but unused: identity "none" never sends credentials. - cc := newCredentialTestConnection(t, a, config.CredentialIdentityNone, staticCredentialProvider{creds: staticCredentials{}}) + // A provider is configured but unused: identity "strip" never sends credentials. + cc := newCredentialTestConnection(t, a, config.CredentialIdentityStrip, staticCredentialProvider{creds: staticCredentials{}}) _, _ = adminservice.NewAdminServiceClient(cc.inboundClient).DescribeCluster(forwardedContext(t), &adminservice.DescribeClusterRequest{}) recorded := localTemporal.credentials(t, describeClusterMethod) @@ -154,8 +154,8 @@ func TestNoneIdentityStripsForwardedCredentials(t *testing.T) { require.Empty(t, recorded.authorizationExtras) } -func TestCallerIdentityForwardsCredentials(t *testing.T) { - for _, identity := range []config.CredentialIdentity{"", config.CredentialIdentityCaller} { +func TestDefaultIdentityForwardsCredentials(t *testing.T) { + for _, identity := range []config.CredentialIdentity{"", config.CredentialIdentityDefault} { t.Run("identity="+string(identity), func(t *testing.T) { a := getDynamicPlccAddresses(t) localTemporal := startAuthRecordingServer(t, a.localTemporalAddr) @@ -226,16 +226,16 @@ func TestCreateClientRejectsUnusableCredentials(t *testing.T) { }) } - t.Run("caller identity leaves a mux connection alone", func(t *testing.T) { + t.Run("default identity leaves a mux connection alone", func(t *testing.T) { _, err := createClient(t.Context(), "test", config.ClusterDefinition{ConnectionType: config.ConnTypeMuxClient}, "inbound", auth.EmptyCredentialProvider{}) require.NoError(t, err) }) - t.Run("none identity works on a mux connection", func(t *testing.T) { + t.Run("strip identity works on a mux connection", func(t *testing.T) { _, err := createClient(t.Context(), "test", config.ClusterDefinition{ ConnectionType: config.ConnTypeMuxClient, - Credentials: &config.CredentialsConfig{Identity: config.CredentialIdentityNone}, + Credentials: &config.CredentialsConfig{Identity: config.CredentialIdentityStrip}, }, "inbound", auth.EmptyCredentialProvider{}) require.NoError(t, err) })