From 9748088a59b926f78008597fa274ed2e6b438433 Mon Sep 17 00:00:00 2001 From: uwin <0xuwin@proton.me> Date: Mon, 7 Oct 2024 13:44:33 +0800 Subject: [PATCH] rpc: check callback params are either exported or builtin types --- rpc/client_test.go | 24 ++++++++++---------- rpc/service.go | 50 ++++++++++++++++++++++++++++++++++++++--- rpc/testservice_test.go | 8 +++---- 3 files changed, 63 insertions(+), 19 deletions(-) diff --git a/rpc/client_test.go b/rpc/client_test.go index b7607adfce..f18a31e19e 100644 --- a/rpc/client_test.go +++ b/rpc/client_test.go @@ -44,10 +44,10 @@ func TestClientRequest(t *testing.T) { defer client.Close() var resp echoResult - if err := client.Call(&resp, "test_echo", "hello", 10, &echoArgs{"world"}); err != nil { + if err := client.Call(&resp, "test_echo", "hello", 10, &EchoArgs{"world"}); err != nil { t.Fatal(err) } - if !reflect.DeepEqual(resp, echoResult{"hello", 10, &echoArgs{"world"}}) { + if !reflect.DeepEqual(resp, echoResult{"hello", 10, &EchoArgs{"world"}}) { t.Errorf("incorrect result %#v", resp) } } @@ -58,12 +58,12 @@ func TestClientResponseType(t *testing.T) { client := DialInProc(server) defer client.Close() - if err := client.Call(nil, "test_echo", "hello", 10, &echoArgs{"world"}); err != nil { + if err := client.Call(nil, "test_echo", "hello", 10, &EchoArgs{"world"}); err != nil { t.Errorf("Passing nil as result should be fine, but got an error: %v", err) } var resultVar echoResult // Note: passing the var, not a ref - err := client.Call(resultVar, "test_echo", "hello", 10, &echoArgs{"world"}) + err := client.Call(resultVar, "test_echo", "hello", 10, &EchoArgs{"world"}) if err == nil { t.Error("Passing a var as result should be an error") } @@ -129,12 +129,12 @@ func TestClientBatchRequest(t *testing.T) { batch := []BatchElem{ { Method: "test_echo", - Args: []interface{}{"hello", 10, &echoArgs{"world"}}, + Args: []interface{}{"hello", 10, &EchoArgs{"world"}}, Result: new(echoResult), }, { Method: "test_echo", - Args: []interface{}{"hello2", 11, &echoArgs{"world"}}, + Args: []interface{}{"hello2", 11, &EchoArgs{"world"}}, Result: new(echoResult), }, { @@ -149,13 +149,13 @@ func TestClientBatchRequest(t *testing.T) { wantResult := []BatchElem{ { Method: "test_echo", - Args: []interface{}{"hello", 10, &echoArgs{"world"}}, - Result: &echoResult{"hello", 10, &echoArgs{"world"}}, + Args: []interface{}{"hello", 10, &EchoArgs{"world"}}, + Result: &echoResult{"hello", 10, &EchoArgs{"world"}}, }, { Method: "test_echo", - Args: []interface{}{"hello2", 11, &echoArgs{"world"}}, - Result: &echoResult{"hello2", 11, &echoArgs{"world"}}, + Args: []interface{}{"hello2", 11, &EchoArgs{"world"}}, + Result: &echoResult{"hello2", 11, &EchoArgs{"world"}}, }, { Method: "no_such_method", @@ -290,7 +290,7 @@ func TestClientNotify(t *testing.T) { client := DialInProc(server) defer client.Close() - if err := client.Notify(context.Background(), "test_echo", "hello", 10, &echoArgs{"world"}); err != nil { + if err := client.Notify(context.Background(), "test_echo", "hello", 10, &EchoArgs{"world"}); err != nil { t.Fatal(err) } } @@ -773,7 +773,7 @@ func TestClientHTTP(t *testing.T) { var ( results = make([]echoResult, 100) errc = make(chan error, len(results)) - wantResult = echoResult{"a", 1, new(echoArgs)} + wantResult = echoResult{"a", 1, new(EchoArgs)} ) for i := range results { i := i diff --git a/rpc/service.go b/rpc/service.go index d50090e9fb..a5a42c8c3f 100644 --- a/rpc/service.go +++ b/rpc/service.go @@ -24,6 +24,7 @@ import ( "strings" "sync" "unicode" + "unicode/utf8" "github.com/ethereum/go-ethereum/log" ) @@ -136,7 +137,9 @@ func newCallback(receiver, fn reflect.Value) *callback { fntype := fn.Type() c := &callback{fn: fn, rcvr: receiver, errPos: -1, isSubscribe: isPubSub(fntype)} // Determine parameter types. They must all be exported or builtin types. - c.makeArgTypes() + if err := c.makeArgTypes(); err != nil { + return nil + } // Verify return types. The function must return at most one error // and/or one other non-error value. @@ -161,7 +164,7 @@ func newCallback(receiver, fn reflect.Value) *callback { } // makeArgTypes composes the argTypes list. -func (c *callback) makeArgTypes() { +func (c *callback) makeArgTypes() error { fntype := c.fn.Type() // Skip receiver and context.Context parameter (if present). firstArg := 0 @@ -175,8 +178,14 @@ func (c *callback) makeArgTypes() { // Add all remaining parameters. c.argTypes = make([]reflect.Type, fntype.NumIn()-firstArg) for i := firstArg; i < fntype.NumIn(); i++ { - c.argTypes[i-firstArg] = fntype.In(i) + intype := fntype.In(i) + if isBuiltinType(intype) || isExportedType(intype) { + c.argTypes[i-firstArg] = fntype.In(i) + } else { + return fmt.Errorf("param %T is neither a builtin type nor exported type", intype) + } } + return nil } // call invokes the callback. @@ -214,6 +223,41 @@ func (c *callback) call(ctx context.Context, method string, args []reflect.Value return results[0].Interface(), nil } +// Is t a builtin type or nested builtin type? +func isBuiltinType(t reflect.Type) bool { + if t.PkgPath() != "" { + return false + } + + switch k := t.Kind(); k { + case reflect.Array, reflect.Chan, reflect.Pointer, reflect.Slice, reflect.UnsafePointer: + return isBuiltinType(t.Elem()) + case reflect.Map: + return isBuiltinType(t.Key()) && isBuiltinType(t.Elem()) + case reflect.Struct: + for i := 0; i < t.NumField(); i++ { + if isBuiltinType(t.Field(i).Type) { + return true + } + } + default: + if k >= reflect.Invalid && k <= reflect.UnsafePointer { + return true + } + } + return false +} + +// Does t have an exported (upper case) name? +func isExportedType(t reflect.Type) bool { + for t.Kind() == reflect.Ptr { + t = t.Elem() + } + // PkgPath will be non-empty even for an exported type (e.g, unnamed type). + rune, _ := utf8.DecodeRuneInString(t.Name()) + return t.PkgPath() != "" && unicode.IsUpper(rune) +} + // Does t satisfy the error interface? func isErrorType(t reflect.Type) bool { return t.Implements(errorType) diff --git a/rpc/testservice_test.go b/rpc/testservice_test.go index 69199e21b7..413db1962e 100644 --- a/rpc/testservice_test.go +++ b/rpc/testservice_test.go @@ -54,14 +54,14 @@ func sequentialIDGenerator() func() ID { type testService struct{} -type echoArgs struct { +type EchoArgs struct { S string } type echoResult struct { String string Int int - Args *echoArgs + Args *EchoArgs } type testError struct{} @@ -82,11 +82,11 @@ func (s *testService) Null() any { return nil } -func (s *testService) Echo(str string, i int, args *echoArgs) echoResult { +func (s *testService) Echo(str string, i int, args *EchoArgs) echoResult { return echoResult{str, i, args} } -func (s *testService) EchoWithCtx(ctx context.Context, str string, i int, args *echoArgs) echoResult { +func (s *testService) EchoWithCtx(ctx context.Context, str string, i int, args *EchoArgs) echoResult { return echoResult{str, i, args} }