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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion internal/dataplane/dataplane.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down
5 changes: 2 additions & 3 deletions internal/proxy/doc.go
Original file line number Diff line number Diff line change
@@ -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
38 changes: 13 additions & 25 deletions internal/proxy/forward.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand All @@ -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 {
Expand All @@ -82,23 +77,16 @@ 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)
if err != nil {
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)
Expand Down
121 changes: 49 additions & 72 deletions internal/proxy/forward_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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)
}
Expand Down Expand Up @@ -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)
Expand All @@ -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"
Expand All @@ -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
Expand All @@ -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",
},
}

Expand All @@ -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
Expand All @@ -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{
Expand Down
22 changes: 9 additions & 13 deletions internal/proxy/forwarding_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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(),
Expand Down Expand Up @@ -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
Expand All @@ -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)

Expand Down
3 changes: 1 addition & 2 deletions internal/proxy/translation_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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")
Expand Down
Loading