Skip to content

Commit beb4d5d

Browse files
add support to set the value of rpc-encoding header
1 parent 2a6e639 commit beb4d5d

3 files changed

Lines changed: 128 additions & 2 deletions

File tree

options.go

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -89,6 +89,7 @@ type TransportOptions struct {
8989
RoutingKey string `long:"rk" description:"The routing key overrides the service name traffic group for proxies."`
9090
RoutingDelegate string `long:"rd" description:"The routing delegate overrides the routing key traffic group for proxies."`
9191
ShardKey string `long:"sk" description:"The shard key is a transport header that clues where to send a request within a clustered traffic group."`
92+
RPCEncoding string `long:"rpc-encoding" description:"Override the rpc-encoding header/metadata value for gRPC and HTTP transports. This does not re-encode the request body and is intended for development."`
9293
Jaeger bool `long:"jaeger" description:"Use the Jaeger tracing client to send Uber style traces and baggage headers"`
9394
TransportHeaders map[string]string `short:"T" long:"topt" description:"Transport options for TChannel, protocol headers for HTTP"`
9495
HTTPMethod string `long:"http-method" description:"The HTTP method to use"`

rpc_encoding_flag_test.go

Lines changed: 117 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,117 @@
1+
package main
2+
3+
import (
4+
"context"
5+
"net"
6+
"net/http"
7+
"net/http/httptest"
8+
"testing"
9+
10+
"github.com/opentracing/opentracing-go"
11+
"github.com/stretchr/testify/assert"
12+
"github.com/stretchr/testify/require"
13+
"github.com/yarpc/yab/encoding"
14+
yabtransport "github.com/yarpc/yab/transport"
15+
16+
apitransport "go.uber.org/yarpc/api/transport"
17+
yarpcgrpc "go.uber.org/yarpc/transport/grpc"
18+
)
19+
20+
type unaryHandlerFunc func(context.Context, *apitransport.Request, apitransport.ResponseWriter) error
21+
22+
func (f unaryHandlerFunc) Handle(ctx context.Context, req *apitransport.Request, resw apitransport.ResponseWriter) error {
23+
return f(ctx, req, resw)
24+
}
25+
26+
type encodingCaptureRouter struct {
27+
expectedService string
28+
expectedProcedure string
29+
capturedEncoding chan string
30+
}
31+
32+
func (r *encodingCaptureRouter) Procedures() []apitransport.Procedure {
33+
return nil
34+
}
35+
36+
func (r *encodingCaptureRouter) Choose(_ context.Context, req *apitransport.Request) (apitransport.HandlerSpec, error) {
37+
if req.Service == r.expectedService && req.Procedure == r.expectedProcedure {
38+
select {
39+
case r.capturedEncoding <- string(req.Encoding):
40+
default:
41+
}
42+
return apitransport.NewUnaryHandlerSpec(unaryHandlerFunc(func(_ context.Context, _ *apitransport.Request, resw apitransport.ResponseWriter) error {
43+
_, _ = resw.Write([]byte("ok"))
44+
return nil
45+
})), nil
46+
}
47+
return apitransport.HandlerSpec{}, apitransport.UnrecognizedProcedureError(req)
48+
}
49+
50+
func TestRPCEncodingFlagOverridesHTTPHeader(t *testing.T) {
51+
const want = "dev-override"
52+
53+
var got string
54+
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
55+
got = r.Header.Get("RPC-Encoding")
56+
w.WriteHeader(http.StatusOK)
57+
_, _ = w.Write([]byte("ok"))
58+
}))
59+
defer srv.Close()
60+
61+
opts := TransportOptions{
62+
ServiceName: "svc",
63+
CallerName: "caller",
64+
Peers: []string{srv.URL},
65+
HTTPMethod: "POST",
66+
RPCEncoding: want,
67+
}
68+
69+
tp, err := getTransport(opts, resolvedProtocolEncoding{protocol: yabtransport.HTTP, enc: encoding.JSON}, opentracing.NoopTracer{})
70+
require.NoError(t, err)
71+
72+
_, err = tp.Call(context.Background(), &yabtransport.Request{Method: "Foo::Bar", Body: []byte("hello")})
73+
require.NoError(t, err)
74+
assert.Equal(t, want, got)
75+
}
76+
77+
func TestRPCEncodingFlagOverridesGRPCHeader(t *testing.T) {
78+
const want = "dev-override"
79+
80+
gt := yarpcgrpc.NewTransport(yarpcgrpc.Tracer(opentracing.NoopTracer{}))
81+
require.NoError(t, gt.Start())
82+
defer func() { _ = gt.Stop() }()
83+
84+
lis, err := net.Listen("tcp", "127.0.0.1:0")
85+
require.NoError(t, err)
86+
defer func() { _ = lis.Close() }()
87+
88+
router := &encodingCaptureRouter{
89+
expectedService: "svc",
90+
expectedProcedure: "svc::echo",
91+
capturedEncoding: make(chan string, 1),
92+
}
93+
94+
inbound := gt.NewInbound(lis)
95+
inbound.SetRouter(router)
96+
require.NoError(t, inbound.Start())
97+
defer func() { _ = inbound.Stop() }()
98+
99+
opts := TransportOptions{
100+
ServiceName: "svc",
101+
CallerName: "caller",
102+
Peers: []string{lis.Addr().String()},
103+
RPCEncoding: want,
104+
}
105+
tp, err := getTransport(opts, resolvedProtocolEncoding{protocol: yabtransport.GRPC, enc: encoding.Protobuf}, opentracing.NoopTracer{})
106+
require.NoError(t, err)
107+
108+
_, err = tp.Call(context.Background(), &yabtransport.Request{TargetService: "svc", Method: "svc::echo", Body: []byte("hello")})
109+
require.NoError(t, err)
110+
111+
select {
112+
case got := <-router.capturedEncoding:
113+
assert.Equal(t, want, got)
114+
default:
115+
t.Fatal("did not capture inbound encoding")
116+
}
117+
}

transport.go

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -164,25 +164,33 @@ func getTransport(opts TransportOptions, resolved resolvedProtocolEncoding, trac
164164
}
165165

166166
if resolved.protocol == transport.GRPC {
167+
grpcEncoding := resolved.enc.String()
168+
if opts.RPCEncoding != "" {
169+
grpcEncoding = opts.RPCEncoding
170+
}
167171
return transport.NewGRPC(transport.GRPCOptions{
168172
Addresses: getHosts(opts.Peers),
169173
Tracer: tracer,
170174
Caller: opts.CallerName,
171-
Encoding: resolved.enc.String(),
175+
Encoding: grpcEncoding,
172176
RoutingKey: opts.RoutingKey,
173177
RoutingDelegate: opts.RoutingDelegate,
174178
MaxResponseSize: opts.GRPCMaxResponseSize,
175179
})
176180
}
177181

182+
httpEncoding := resolved.enc.String()
183+
if opts.RPCEncoding != "" {
184+
httpEncoding = opts.RPCEncoding
185+
}
178186
hopts := transport.HTTPOptions{
179187
Method: opts.HTTPMethod,
180188
SourceService: opts.CallerName,
181189
TargetService: opts.ServiceName,
182190
RoutingDelegate: opts.RoutingDelegate,
183191
RoutingKey: opts.RoutingKey,
184192
ShardKey: opts.ShardKey,
185-
Encoding: resolved.enc.String(),
193+
Encoding: httpEncoding,
186194
URLs: opts.Peers,
187195
Tracer: tracer,
188196
}

0 commit comments

Comments
 (0)