diff --git a/rpcmiddleware/interface.go b/rpcmiddleware/interface.go index f108b444..bb768b29 100644 --- a/rpcmiddleware/interface.go +++ b/rpcmiddleware/interface.go @@ -18,13 +18,16 @@ var ( // errorType is the reflection type of the error interface. errorType = reflect.TypeOf((*error)(nil)).Elem() + // ctxType is the reflection type of the context.Context interface. + ctxType = reflect.TypeOf((*context.Context)(nil)).Elem() + // protoMessageType is the reflection type of the proto.Message // interface. protoMessageType = reflect.TypeOf((*proto.Message)(nil)).Elem() // passThroughMessageHandler is a messageHandler that does not modify // the message and just passes it through. - passThroughMessageHandler messageHandler = func( + passThroughMessageHandler messageHandler = func(context.Context, proto.Message) (proto.Message, error) { return nil, nil @@ -37,8 +40,8 @@ var ( } // messageDenyHandler disallows the given message. - messageDenyHandler messageHandler = func(req proto.Message) ( - proto.Message, error) { + messageDenyHandler messageHandler = func(context.Context, + proto.Message) (proto.Message, error) { return nil, ErrNotSupported } @@ -72,7 +75,7 @@ type RequestInterceptor interface { // new message of the same type (=return non-nil message, nil error) or abort // the call by returning a non-nil error. If the message is a request, then // returning a non-nil error will reject the request. -type messageHandler func(req proto.Message) (proto.Message, error) +type messageHandler func(context.Context, proto.Message) (proto.Message, error) // ErrorHandler is a function type for a generic gRPC error handler. It can // pass through the error unchanged (=return nil, nil), replace the error with @@ -99,14 +102,14 @@ type RoundTripChecker interface { // a new message of the same type (=return non-nil message, nil) or // refuse (=return non-nil error with rejection reason) an incoming // request. - HandleRequest(proto.Message) (proto.Message, error) + HandleRequest(context.Context, proto.Message) (proto.Message, error) // HandleResponse is called for each outgoing gRPC response message of // the type declared to be handled by HandlesResponse. The handler can // pass through the response (=return nil, nil), replace the response // with a new message of the same type (=return non-nil message, nil // error) or abort the call by returning a non-nil error. - HandleResponse(proto.Message) (proto.Message, error) + HandleResponse(context.Context, proto.Message) (proto.Message, error) // HandleErrorResponse is called for any error response. // The handler can pass through the error (=return nil, nil), replace @@ -142,27 +145,26 @@ func (r *DefaultChecker) HandlesResponse(t protoreflect.MessageType) bool { return t == r.responseType } -// HandleRequest is called for each incoming gRPC request message of the -// type declared to be accepted by HandlesRequest. The handler can -// accept the request as is (=return nil, nil), replace the request with -// a new message of the same type (=return non-nil message, nil) or -// refuse (=return non-nil error with rejection reason) an incoming -// request. -func (r *DefaultChecker) HandleRequest(req proto.Message) (proto.Message, - error) { +// HandleRequest is called for each incoming gRPC request message of the type +// declared to be accepted by HandlesRequest. The handler can accept the request +// as is (=return nil, nil), replace the request with a new message of the same +// type (=return non-nil message, nil) or refuse (=return non-nil error with +// rejection reason) an incoming request. +func (r *DefaultChecker) HandleRequest(ctx context.Context, + req proto.Message) (proto.Message, error) { - return r.requestHandler(req) + return r.requestHandler(ctx, req) } -// HandleResponse is called for each outgoing gRPC response message of -// the type declared to be handled by HandlesResponse. The handler can -// pass through the response (=return nil, nil), replace the response -// with a new message of the same type (=return non-nil message, nil -// error) or abort the call by returning a non-nil error. -func (r *DefaultChecker) HandleResponse(resp proto.Message) (proto.Message, - error) { +// HandleResponse is called for each outgoing gRPC response message of the type +// declared to be handled by HandlesResponse. The handler can pass through the +// response (=return nil, nil), replace the response with a new message of the +// same type (=return non-nil message, nil error) or abort the call by returning +// a non-nil error. +func (r *DefaultChecker) HandleResponse(ctx context.Context, + resp proto.Message) (proto.Message, error) { - return r.responseHandler(resp) + return r.responseHandler(ctx, resp) } // HandleErrorResponse is called for any error response. @@ -310,7 +312,9 @@ func newReflectionRequestCheckHandler(requestSample proto.Message, panic(err) } - return func(req proto.Message) (proto.Message, error) { + return func(ctx context.Context, req proto.Message) (proto.Message, + error) { + if req.ProtoReflect().Type() != requestProtoType { return nil, fmt.Errorf("request handler called for "+ "unsupported type %v (expected %v)", @@ -320,6 +324,7 @@ func newReflectionRequestCheckHandler(requestSample proto.Message, // We made sure this call would succeed when creating the // handler. resp := handlerValue.Call([]reflect.Value{ + reflect.ValueOf(ctx), reflect.ValueOf(req), }) @@ -352,7 +357,9 @@ func newReflectionMessageHandler(messageSample proto.Message, panic(err) } - return func(req proto.Message) (proto.Message, error) { + return func(ctx context.Context, req proto.Message) (proto.Message, + error) { + if req.ProtoReflect().Type() != messageProtoType { return nil, fmt.Errorf("message handler called for "+ "unsupported type %v (expected %v)", @@ -362,6 +369,7 @@ func newReflectionMessageHandler(messageSample proto.Message, // We made sure this call would succeed when creating the // handler. resp := handlerValue.Call([]reflect.Value{ + reflect.ValueOf(ctx), reflect.ValueOf(req), }) @@ -390,12 +398,16 @@ func validateRequestCheckHandler(typedHandlerType reflect.Type, if typedHandlerType.Kind() != reflect.Func { return fmt.Errorf("request handler must be a function") } - if typedHandlerType.NumIn() != 1 || typedHandlerType.NumOut() != 1 { - return fmt.Errorf("request handler must have exactly one " + + if typedHandlerType.NumIn() != 2 || typedHandlerType.NumOut() != 1 { + return fmt.Errorf("request handler must have exactly two " + "parameter and one return value") } - if !typedHandlerType.In(0).ConvertibleTo(requestType) { - return fmt.Errorf("request handler must have one parameter " + + if !typedHandlerType.In(0).ConvertibleTo(ctxType) { + return fmt.Errorf("request handler must have first parameter " + + "with a sub type of context.Context") + } + if !typedHandlerType.In(1).ConvertibleTo(requestType) { + return fmt.Errorf("request handler must have second parameter " + "with a sub type of proto.Message") } if typedHandlerType.Out(0) != errorType { @@ -414,12 +426,16 @@ func validateMessageHandler(typedHandlerType reflect.Type, if typedHandlerType.Kind() != reflect.Func { return fmt.Errorf("message handler must be a function") } - if typedHandlerType.NumIn() != 1 || typedHandlerType.NumOut() != 2 { - return fmt.Errorf("message handler must have exactly one " + + if typedHandlerType.NumIn() != 2 || typedHandlerType.NumOut() != 2 { + return fmt.Errorf("message handler must have exactly two " + "parameter and two return values") } - if !typedHandlerType.In(0).ConvertibleTo(messageType) { - return fmt.Errorf("message handler must have one parameter " + + if !typedHandlerType.In(0).ConvertibleTo(ctxType) { + return fmt.Errorf("request handler must have first parameter " + + "with a sub type of context.Context") + } + if !typedHandlerType.In(1).ConvertibleTo(messageType) { + return fmt.Errorf("message handler must have second parameter " + "with a sub type of proto.Message") } outType0 := typedHandlerType.Out(0) diff --git a/rpcmiddleware/interface_test.go b/rpcmiddleware/interface_test.go index b8163dbb..846706fc 100644 --- a/rpcmiddleware/interface_test.go +++ b/rpcmiddleware/interface_test.go @@ -1,6 +1,7 @@ package rpcmiddleware import ( + "context" "fmt" "testing" @@ -10,6 +11,8 @@ import ( ) var ( + ctxb = context.Background() + listPeersReq = &lnrpc.ListPeersRequest{} listPeersReqType = listPeersReq.ProtoReflect().Type() @@ -37,11 +40,11 @@ func TestPassThrough(t *testing.T) { require.True(t, peersChecker.HandlesRequest(listPeersReqType)) require.True(t, peersChecker.HandlesResponse(listPeersRespType)) - req, err := peersChecker.HandleRequest(listPeersReq) + req, err := peersChecker.HandleRequest(ctxb, listPeersReq) require.NoError(t, err) require.Nil(t, req) - resp, err := peersChecker.HandleResponse(listPeersResp) + resp, err := peersChecker.HandleResponse(ctxb, listPeersResp) require.NoError(t, err) require.Nil(t, resp) @@ -57,11 +60,11 @@ func TestRequestDenier(t *testing.T) { require.True(t, peersChecker.HandlesRequest(listPeersReqType)) require.True(t, peersChecker.HandlesResponse(listPeersRespType)) - req, err := peersChecker.HandleRequest(listPeersReq) + req, err := peersChecker.HandleRequest(ctxb, listPeersReq) require.ErrorIs(t, err, ErrNotSupported) require.Nil(t, req) - resp, err := peersChecker.HandleResponse(listPeersResp) + resp, err := peersChecker.HandleResponse(ctxb, listPeersResp) require.ErrorIs(t, err, ErrNotSupported) require.Nil(t, resp) @@ -74,7 +77,7 @@ func TestRequestDenier(t *testing.T) { func TestRequestChecker(t *testing.T) { peersChecker := NewRequestChecker( listPeersReq, listPeersResp, - func(peer *lnrpc.ListPeersRequest) error { + func(context.Context, *lnrpc.ListPeersRequest) error { return nil }, ) @@ -82,11 +85,11 @@ func TestRequestChecker(t *testing.T) { require.True(t, peersChecker.HandlesRequest(listPeersReqType)) require.True(t, peersChecker.HandlesResponse(listPeersRespType)) - req, err := peersChecker.HandleRequest(listPeersReq) + req, err := peersChecker.HandleRequest(ctxb, listPeersReq) require.NoError(t, err) require.Nil(t, req) - resp, err := peersChecker.HandleResponse(listPeersResp) + resp, err := peersChecker.HandleResponse(ctxb, listPeersResp) require.NoError(t, err) require.Nil(t, resp) @@ -99,7 +102,9 @@ func TestRequestChecker(t *testing.T) { func TestRequestRewriter(t *testing.T) { peersChecker := NewRequestRewriter( listPeersReq, listPeersResp, - func(peer *lnrpc.ListPeersRequest) (proto.Message, error) { + func(ctx context.Context, + peer *lnrpc.ListPeersRequest) (proto.Message, error) { + return peer, nil }, ) @@ -107,11 +112,11 @@ func TestRequestRewriter(t *testing.T) { require.True(t, peersChecker.HandlesRequest(listPeersReqType)) require.True(t, peersChecker.HandlesResponse(listPeersRespType)) - req, err := peersChecker.HandleRequest(listPeersReq) + req, err := peersChecker.HandleRequest(ctxb, listPeersReq) require.NoError(t, err) require.Equal(t, listPeersReq, req) - resp, err := peersChecker.HandleResponse(listPeersResp) + resp, err := peersChecker.HandleResponse(ctxb, listPeersResp) require.NoError(t, err) require.Nil(t, resp) @@ -124,7 +129,9 @@ func TestRequestRewriter(t *testing.T) { func TestResponseRewriter(t *testing.T) { peersChecker := NewResponseRewriter( listPeersReq, listPeersResp, - func(peer *lnrpc.ListPeersResponse) (proto.Message, error) { + func(ctx context.Context, + peer *lnrpc.ListPeersResponse) (proto.Message, error) { + return peer, nil }, PassThroughErrorHandler, ) @@ -132,11 +139,11 @@ func TestResponseRewriter(t *testing.T) { require.True(t, peersChecker.HandlesRequest(listPeersReqType)) require.True(t, peersChecker.HandlesResponse(listPeersRespType)) - req, err := peersChecker.HandleRequest(listPeersReq) + req, err := peersChecker.HandleRequest(ctxb, listPeersReq) require.NoError(t, err) require.Nil(t, req) - resp, err := peersChecker.HandleResponse(listPeersResp) + resp, err := peersChecker.HandleResponse(ctxb, listPeersResp) require.NoError(t, err) require.Equal(t, listPeersResp, resp) @@ -150,10 +157,12 @@ func TestFullChecker(t *testing.T) { myErr := fmt.Errorf("some error happened") peersChecker := NewFullChecker( listPeersReq, listPeersResp, - func(peer *lnrpc.ListPeersRequest) error { + func(ctx context.Context, peer *lnrpc.ListPeersRequest) error { return myErr }, - func(*lnrpc.ListPeersResponse) (proto.Message, error) { + func(context.Context, *lnrpc.ListPeersResponse) (proto.Message, + error) { + return nil, myErr }, PassThroughErrorHandler, ) @@ -161,10 +170,10 @@ func TestFullChecker(t *testing.T) { require.True(t, peersChecker.HandlesRequest(listPeersReqType)) require.True(t, peersChecker.HandlesResponse(listPeersRespType)) - _, err := peersChecker.HandleRequest(listPeersReq) + _, err := peersChecker.HandleRequest(ctxb, listPeersReq) require.Equal(t, myErr, err) - resp, err := peersChecker.HandleResponse(listPeersResp) + resp, err := peersChecker.HandleResponse(ctxb, listPeersResp) require.Error(t, err) require.Equal(t, myErr, err) require.Nil(t, resp) @@ -179,10 +188,14 @@ func TestFullRewriter(t *testing.T) { myErr := fmt.Errorf("some error happened") peersChecker := NewFullRewriter( listPeersReq, listPeersResp, - func(peer *lnrpc.ListPeersRequest) (proto.Message, error) { + func(ctx context.Context, + peer *lnrpc.ListPeersRequest) (proto.Message, error) { + return nil, myErr }, - func(*lnrpc.ListPeersResponse) (proto.Message, error) { + func(context.Context, *lnrpc.ListPeersResponse) (proto.Message, + error) { + return nil, myErr }, PassThroughErrorHandler, ) @@ -190,10 +203,10 @@ func TestFullRewriter(t *testing.T) { require.True(t, peersChecker.HandlesRequest(listPeersReqType)) require.True(t, peersChecker.HandlesResponse(listPeersRespType)) - _, err := peersChecker.HandleRequest(listPeersReq) + _, err := peersChecker.HandleRequest(ctxb, listPeersReq) require.Equal(t, myErr, err) - resp, err := peersChecker.HandleResponse(listPeersResp) + resp, err := peersChecker.HandleResponse(ctxb, listPeersResp) require.Error(t, err) require.Equal(t, myErr, err) require.Nil(t, resp) @@ -213,7 +226,7 @@ func TestImplementationPanics(t *testing.T) { }, ) require.PanicsWithError( - t, "request handler must have exactly one parameter and one "+ + t, "request handler must have exactly two parameter and one "+ "return value", func() { _ = NewRequestChecker( @@ -232,7 +245,7 @@ func TestImplementationPanics(t *testing.T) { }, ) require.PanicsWithError( - t, "message handler must have exactly one parameter and two "+ + t, "message handler must have exactly two parameter and two "+ "return values", func() { _ = NewRequestRewriter( @@ -252,7 +265,7 @@ func TestImplementationPanics(t *testing.T) { }, ) require.PanicsWithError( - t, "message handler must have exactly one parameter and two "+ + t, "message handler must have exactly two parameter and two "+ "return values", func() { _ = NewResponseRewriter( diff --git a/rpcmiddleware/proto.go b/rpcmiddleware/proto.go index 7bc24a74..8a777892 100644 --- a/rpcmiddleware/proto.go +++ b/rpcmiddleware/proto.go @@ -57,7 +57,7 @@ func RPCErr(req *lnrpc.RPCMiddlewareRequest, return RPCErrString(req, err.Error()) } - return RPCErrString(req, "") + return RPCOk(req) } // RPCErrString constructs a middleware response. If an empty format param is