diff --git a/auth/credential_provider.go b/auth/credential_provider.go index 9770332..51e698e 100644 --- a/auth/credential_provider.go +++ b/auth/credential_provider.go @@ -17,6 +17,10 @@ type ( EmptyCredentialProvider struct{} ) +// 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. var Module = fx.Options( diff --git a/config/cluster_conn_config.go b/config/cluster_conn_config.go index 2e73fd5..8dc1a54 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 CredentialIdentityDefault. + 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,16 @@ type ( } ) +const ( + // 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" + // CredentialIdentityStrip drops forwarded credentials and sends none. + CredentialIdentityStrip CredentialIdentity = "strip" +) + const ( ConnTypeTCP ConnectionType = "tcp" ConnTypeMuxServer ConnectionType = "mux-server" @@ -98,9 +112,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 CredentialIdentityDefault. +func (c ClusterDefinition) CredentialIdentity() CredentialIdentity { + if c.Credentials == nil || c.Credentials.Identity == "" { + return CredentialIdentityDefault + } + return c.Credentials.Identity } // Validate reports problems in this connection's config. Only the encryption @@ -109,11 +126,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 CredentialIdentityDefault, CredentialIdentityStrip: + 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, CredentialIdentityDefault, CredentialIdentityProxy, CredentialIdentityStrip) } - 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..1c6c5c4 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 default identity", + conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeMuxClient, Credentials: identity(CredentialIdentityDefault)}}, + }, + { + name: "local mux with strip identity", + conn: ClusterConnConfig{Local: ClusterDefinition{ConnectionType: ConnTypeMuxClient, Credentials: identity(CredentialIdentityStrip)}}, }, { - 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 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 692b1c7..4bc8381 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 | +| --- | --- | --- | +| `default` (when not set) | forwarded | not sent | +| `proxy` | dropped | 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. + +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..d8a6ab8 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: + # 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: - 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..b4cfef5 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.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 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.CredentialHeaders 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..099413d 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 TestStripIdentityStripsForwardedCredentials(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 "strip" never sends credentials. + cc := newCredentialTestConnection(t, a, config.CredentialIdentityStrip, 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 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) + 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("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("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.CredentialIdentityStrip}, + }, "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"})) +}