diff --git a/internal/dataplane/dataplane.go b/internal/dataplane/dataplane.go index 9b3217a..a9286c3 100644 --- a/internal/dataplane/dataplane.go +++ b/internal/dataplane/dataplane.go @@ -397,7 +397,7 @@ func newUpstreamForwarder( return nil, nil, err } - fw, err := proxy.NewForwarder(conn, o.allowlist, proxy.WithProtoTypes(o.types)) + fw, err := proxy.NewForwarder(conn, proxy.WithProtoTypes(o.types)) if err != nil { return nil, nil, err } diff --git a/internal/proxy/doc.go b/internal/proxy/doc.go index 9a44a4c..3591299 100644 --- a/internal/proxy/doc.go +++ b/internal/proxy/doc.go @@ -1,4 +1,3 @@ -// Package proxy forwards every allowlisted service to an upstream Temporal -// Service over gRPC, applying namespace translation and payload codecs on the -// way. +// Package proxy forwards requests to an upstream Temporal Service over gRPC, +// applying namespace translation and payload codecs on the way. package proxy diff --git a/internal/proxy/forward.go b/internal/proxy/forward.go index bd81427..e74fc21 100644 --- a/internal/proxy/forward.go +++ b/internal/proxy/forward.go @@ -23,15 +23,15 @@ import ( var transportHeaders = []string{"user-agent", ":authority", "content-type"} type ( - // Forwarder forwards any allowlisted method to a single upstream, typing each - // request and response from the proto registry rather than being generated per - // service. The typing is load-bearing: namespace translation and payload + // Forwarder forwards any method to a single upstream, typing each request and + // response from the proto registry rather than being generated per service. + // Deciding which services may be forwarded is the caller's job: the gateway + // rejects the rest before routing. The typing is load-bearing: namespace translation and payload // encryption are client interceptors on cc that operate on proto messages, so // an opaque byte passthrough (as the router uses) would silently skip both. // Resolved methods are cached, and a Forwarder is safe for concurrent use. Forwarder struct { cc grpc.ClientConnInterface - allowed services.Allowlist methods sync.Map // fullName -> *methodInfo types protoutil.Types } @@ -47,22 +47,17 @@ type ( } ) -// NewForwarder builds a Forwarder that forwards over cc every method belonging -// to a service a admits. It fails when cc or a is nil. By default methods are -// typed against the global proto registry; use [WithProtoTypes] to override it. -func NewForwarder(cc grpc.ClientConnInterface, a services.Allowlist, opts ...ForwarderOption) (*Forwarder, error) { +// NewForwarder builds a Forwarder that forwards over cc. It fails when cc is nil. +// By default methods are typed against the global proto registry; use +// [WithProtoTypes] to override it. +func NewForwarder(cc grpc.ClientConnInterface, opts ...ForwarderOption) (*Forwarder, error) { if cc == nil { return nil, fmt.Errorf("proxy: nil client connection passed to forwarder") } - if a == nil { - return nil, fmt.Errorf("proxy: nil allowlist passed to forwarder") - } - f := &Forwarder{ - cc: cc, - allowed: a, - types: protoregistry.GlobalTypes, + cc: cc, + types: protoregistry.GlobalTypes, } for _, opt := range opts { @@ -82,12 +77,9 @@ func WithProtoTypes(t protoutil.Types) ForwarderOption { } } -// Handle forwards one stream to the upstream, and suits -// [google.golang.org/grpc.UnknownServiceHandler]. A method whose service the -// [services.Allowlist] does not admit is rejected with Unimplemented before any -// upstream work, so the proxy answers as a server that does not implement it -// rather than revealing that an upstream might. Only methods present in the -// compiled descriptors can be forwarded; anything else is Unimplemented too. +// Handle forwards one stream to the upstream. Only methods present in the +// compiled descriptors can be forwarded; anything else is rejected with +// Unimplemented before any upstream work. func (f *Forwarder) Handle(_ any, ss grpc.ServerStream) error { ctx := ss.Context() method, err := rpc.FullMethod(ctx) @@ -95,10 +87,6 @@ func (f *Forwarder) Handle(_ any, ss grpc.ServerStream) error { return err } - if service := rpc.Service(method); !f.allowed.Allows(service) { - return status.Errorf(codes.Unimplemented, "unknown service %s", service) - } - info := f.lookup(method) if info == nil { return status.Errorf(codes.Unimplemented, "unknown method %s", method) diff --git a/internal/proxy/forward_test.go b/internal/proxy/forward_test.go index d349dde..648d12c 100644 --- a/internal/proxy/forward_test.go +++ b/internal/proxy/forward_test.go @@ -82,25 +82,20 @@ func TestForwardContextWithoutIncomingMetadata(t *testing.T) { func TestNewForwarderValidation(t *testing.T) { t.Parallel() - allowed := services.NewAllowlist(services.Default()) - tests := []struct { - name string - cc grpc.ClientConnInterface - allowed services.Allowlist - err string + name string + cc grpc.ClientConnInterface + err string }{ - {name: "nil client connection", allowed: allowed, err: "nil client connection"}, - {name: "nil allowlist", cc: &testutil.ClientConn{}, err: "nil allowlist"}, - {name: "both nil reports the connection first", err: "nil client connection"}, - {name: "a connection and an allowlist is enough", cc: &testutil.ClientConn{}, allowed: allowed}, + {name: "nil client connection", err: "nil client connection"}, + {name: "a connection is enough", cc: &testutil.ClientConn{}}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { t.Parallel() - fw, err := NewForwarder(tt.cc, tt.allowed) + fw, err := NewForwarder(tt.cc) if tt.err != "" { require.ErrorContains(t, err, tt.err) require.Nil(t, fw) @@ -119,12 +114,12 @@ func TestWithProtoTypes(t *testing.T) { custom := partialTypes{} - fw, err := NewForwarder(&testutil.ClientConn{}, services.NewAllowlist(services.Default()), WithProtoTypes(custom)) + fw, err := NewForwarder(&testutil.ClientConn{}, WithProtoTypes(custom)) require.NoError(t, err) require.Equal(t, custom, fw.types) // A nil registry leaves the default in place rather than disabling resolution. - fw, err = NewForwarder(&testutil.ClientConn{}, services.NewAllowlist(services.Default()), WithProtoTypes(nil)) + fw, err = NewForwarder(&testutil.ClientConn{}, WithProtoTypes(nil)) require.NoError(t, err) require.Equal(t, protoregistry.GlobalTypes, fw.types) } @@ -164,7 +159,7 @@ func TestResolveMethod(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - fw, err := NewForwarder(&testutil.ClientConn{}, services.NewAllowlist(services.Default()), WithProtoTypes(tt.types)) + fw, err := NewForwarder(&testutil.ClientConn{}, WithProtoTypes(tt.types)) require.NoError(t, err) got := fw.resolveMethod(tt.in) @@ -183,7 +178,7 @@ func TestResolveMethod(t *testing.T) { func TestLookupCachesResolvedMethods(t *testing.T) { t.Parallel() - fw, err := NewForwarder(&testutil.ClientConn{}, services.NewAllowlist(services.Default())) + fw, err := NewForwarder(&testutil.ClientConn{}) require.NoError(t, err) const method = "/" + services.WorkflowService + "/GetSystemInfo" @@ -210,7 +205,6 @@ func TestHandleErrors(t *testing.T) { name string method string noStream bool - allowed []string conn *testutil.ClientConn ss testutil.ServerStream code codes.Code @@ -223,71 +217,57 @@ func TestHandleErrors(t *testing.T) { msg: "no server transport stream", }, { - name: "a service that is not on the allowlist", - method: "/" + services.OperatorService + "/DeleteNamespace", - allowed: []string{services.WorkflowService}, - code: codes.Unimplemented, - msg: "unknown service " + services.OperatorService, - }, - { - name: "a method the compiled descriptors do not have", - method: "/" + services.WorkflowService + "/NoSuchMethod", - allowed: services.Default(), - code: codes.Unimplemented, - msg: "unknown method", + name: "a method the compiled descriptors do not have", + method: "/" + services.WorkflowService + "/NoSuchMethod", + code: codes.Unimplemented, + msg: "unknown method", }, { - name: "an upstream that rejects the unary call", - method: unaryMethod, - allowed: services.Default(), - conn: &testutil.ClientConn{InvokeErr: status.Error(codes.PermissionDenied, "denied")}, - code: codes.PermissionDenied, - msg: "denied", + name: "an upstream that rejects the unary call", + method: unaryMethod, + conn: &testutil.ClientConn{InvokeErr: status.Error(codes.PermissionDenied, "denied")}, + code: codes.PermissionDenied, + msg: "denied", }, { - name: "a caller that cannot receive the relayed header", - method: unaryMethod, - allowed: services.Default(), - conn: &testutil.ClientConn{Header: metadata.Pairs("x-upstream", "1")}, - ss: testutil.ServerStream{HeaderErr: errors.New("header refused")}, - code: codes.Internal, - msg: "header refused", + name: "a caller that cannot receive the relayed header", + method: unaryMethod, + conn: &testutil.ClientConn{Header: metadata.Pairs("x-upstream", "1")}, + ss: testutil.ServerStream{HeaderErr: errors.New("header refused")}, + code: codes.Internal, + msg: "header refused", }, { - name: "a caller that cannot receive the response", - method: unaryMethod, - allowed: services.Default(), - ss: testutil.ServerStream{SendErr: errors.New("broken pipe")}, - code: codes.Internal, - msg: "broken pipe", + name: "a caller that cannot receive the response", + method: unaryMethod, + ss: testutil.ServerStream{SendErr: errors.New("broken pipe")}, + code: codes.Internal, + msg: "broken pipe", }, { - name: "an upstream stream that cannot be opened", - method: streamMethod, - allowed: []string{services.Reflection}, - conn: &testutil.ClientConn{StreamErr: status.Error(codes.Unavailable, "upstream down")}, - code: codes.Unavailable, - msg: "upstream down", + name: "an upstream stream that cannot be opened", + method: streamMethod, + conn: &testutil.ClientConn{StreamErr: status.Error(codes.Unavailable, "upstream down")}, + code: codes.Unavailable, + msg: "upstream down", }, { - name: "an upstream stream that fails before its header", - method: streamMethod, - allowed: []string{services.Reflection}, - conn: &testutil.ClientConn{Stream: &testutil.ClientStream{HeaderErr: status.Error(codes.Aborted, "header failed")}}, - ss: testutil.ServerStream{RecvErr: io.EOF}, - code: codes.Aborted, - msg: "header failed", + name: "an upstream stream that fails before its header", + method: streamMethod, + conn: &testutil.ClientConn{Stream: &testutil.ClientStream{HeaderErr: status.Error(codes.Aborted, "header failed")}}, + ss: testutil.ServerStream{RecvErr: io.EOF}, + code: codes.Aborted, + msg: "header failed", }, { // The request pump reports first because the response pump is parked in // RecvMsg, so this pins the mapping of a caller-side stream failure. - name: "a caller whose request stream breaks", - method: streamMethod, - allowed: []string{services.Reflection}, - conn: &testutil.ClientConn{Stream: &testutil.ClientStream{BlockRecv: make(chan struct{})}}, - ss: testutil.ServerStream{RecvErr: errors.New("caller vanished")}, - code: codes.Internal, - msg: "caller vanished", + name: "a caller whose request stream breaks", + method: streamMethod, + conn: &testutil.ClientConn{Stream: &testutil.ClientStream{BlockRecv: make(chan struct{})}}, + ss: testutil.ServerStream{RecvErr: errors.New("caller vanished")}, + code: codes.Internal, + msg: "caller vanished", }, } @@ -304,7 +284,7 @@ func TestHandleErrors(t *testing.T) { t.Cleanup(func() { close(cs.BlockRecv) }) } - fw, err := NewForwarder(conn, services.NewAllowlist(tt.allowed)) + fw, err := NewForwarder(conn) require.NoError(t, err) ss := tt.ss @@ -329,10 +309,7 @@ func TestStreamTreatsWrappedEOFAsHalfClose(t *testing.T) { // otherwise successful call. wrapped := fmt.Errorf("transport closed: %w", io.EOF) - fw, err := NewForwarder( - &testutil.ClientConn{Stream: &testutil.ClientStream{RecvErr: wrapped}}, - services.NewAllowlist([]string{services.Reflection}), - ) + fw, err := NewForwarder(&testutil.ClientConn{Stream: &testutil.ClientStream{RecvErr: wrapped}}) require.NoError(t, err) ss := testutil.ServerStream{ diff --git a/internal/proxy/forwarding_test.go b/internal/proxy/forwarding_test.go index 851d9fa..21fb4de 100644 --- a/internal/proxy/forwarding_test.go +++ b/internal/proxy/forwarding_test.go @@ -68,7 +68,7 @@ func TestStreamForwardsBidiReflection(t *testing.T) { workflowservice.RegisterWorkflowServiceServer(s, &metadataStampingService{}) reflection.Register(s) }) - conn := startProxy(t, addr, services.Reflection) + conn := startProxy(t, addr) stream, err := grpc_reflection_v1.NewServerReflectionClient(conn).ServerReflectionInfo( t.Context(), @@ -139,20 +139,16 @@ func (*metadataStampingService) GetSystemInfo( return &workflowservice.GetSystemInfoResponse{}, nil } -// forwarder builds a forwarder for the named services over a plain client conn -// to upstream, standing in for the pool-backed connection used in production. -func forwarder(t *testing.T, upstream string, allowed ...string) *proxy.Forwarder { +// forwarder builds a forwarder over a plain client conn to upstream, standing in +// for the pool-backed connection used in production. +func forwarder(t *testing.T, upstream string) *proxy.Forwarder { t.Helper() - if len(allowed) == 0 { - allowed = services.Default() - } - conn, err := grpc.NewClient(upstream, grpc.WithTransportCredentials(insecure.NewCredentials())) require.NoError(t, err) t.Cleanup(func() { _ = conn.Close() }) - fw, err := proxy.NewForwarder(conn, services.NewAllowlist(allowed)) + fw, err := proxy.NewForwarder(conn) require.NoError(t, err) return fw @@ -177,15 +173,15 @@ func serveUpstream(t *testing.T, register ...func(*grpc.Server)) string { return lis.Addr().String() } -// startProxy serves a forwarder for the named services to upstream on a -// loopback port and returns a client connection to it. -func startProxy(t *testing.T, upstream string, allowed ...string) *grpc.ClientConn { +// startProxy serves a forwarder to upstream on a loopback port and returns a +// client connection to it. +func startProxy(t *testing.T, upstream string) *grpc.ClientConn { t.Helper() lis, err := net.Listen("tcp", "127.0.0.1:0") require.NoError(t, err) - svr := grpc.NewServer(grpc.UnknownServiceHandler(forwarder(t, upstream, allowed...).Handle)) + svr := grpc.NewServer(grpc.UnknownServiceHandler(forwarder(t, upstream).Handle)) go func() { _ = svr.Serve(lis) }() t.Cleanup(svr.Stop) diff --git a/internal/proxy/translation_test.go b/internal/proxy/translation_test.go index 8778727..69b93eb 100644 --- a/internal/proxy/translation_test.go +++ b/internal/proxy/translation_test.go @@ -19,7 +19,6 @@ import ( "google.golang.org/protobuf/reflect/protoregistry" "github.com/temporalio/temporal-proxy/internal/protoutil" - "github.com/temporalio/temporal-proxy/internal/services" ) type ( @@ -148,7 +147,7 @@ func TestOutboundNamespaceTranslation(t *testing.T) { require.NoError(t, err) t.Cleanup(func() { _ = cc.Close() }) - fw, err := NewForwarder(cc, services.NewAllowlist(services.Default())) + fw, err := NewForwarder(cc) require.NoError(t, err) proxyLis, err := net.Listen("tcp", "127.0.0.1:0")