From ead8eda3c542b8a37af6a4b524cd6cfb39f1febc Mon Sep 17 00:00:00 2001 From: slasher Date: Tue, 3 Sep 2024 10:07:36 +0800 Subject: [PATCH] feat(rpc2): fixup some code of rpc2 @formatter:off Signed-off-by: slasher --- blobstore/api/shardnode/client.go | 10 ++++-- blobstore/api/shardnode/shard.go | 2 +- blobstore/cmd/cmd.go | 20 ++++++----- blobstore/common/rpc/auth/handler.go | 6 +--- blobstore/common/rpc2/checksum.go | 8 ++--- blobstore/common/rpc2/client.go | 45 +++++++++++++++++++++++-- blobstore/common/rpc2/error.go | 4 +++ blobstore/common/rpc2/example/client.go | 9 ++--- blobstore/common/rpc2/example/main.go | 8 +++-- blobstore/common/rpc2/example/server.go | 12 ++++++- blobstore/common/rpc2/example_test.go | 3 +- blobstore/common/rpc2/request.go | 6 +--- blobstore/common/rpc2/response_test.go | 2 +- blobstore/common/rpc2/router.go | 12 ++----- blobstore/common/rpc2/rpc2.go | 15 ++++++++- blobstore/common/rpc2/rpc2_test.go | 43 ++++++++++++++++++++--- blobstore/common/rpc2/server.go | 7 ++-- blobstore/common/rpc2/stream.go | 6 +--- blobstore/common/rpc2/stream_test.go | 3 ++ blobstore/shardnode/rpcservice.go | 4 +-- 20 files changed, 156 insertions(+), 69 deletions(-) diff --git a/blobstore/api/shardnode/client.go b/blobstore/api/shardnode/client.go index d41d81bdb..31d4cc73b 100644 --- a/blobstore/api/shardnode/client.go +++ b/blobstore/api/shardnode/client.go @@ -16,8 +16,10 @@ package shardnode import ( "context" + "time" "github.com/cubefs/cubefs/blobstore/common/rpc2" + "github.com/cubefs/cubefs/blobstore/util/defaulter" ) type Config = rpc2.Client @@ -26,10 +28,12 @@ type Client struct { rpc2.Client } -func New(cli Config) Client { - return Client{ - cli, +func New(cli Config) *Client { + defaulter.Empty(&cli.ConnectorConfig.Network, "tcp") + if cli.ConnectorConfig.DialTimeout.Duration <= 0 { + cli.ConnectorConfig.DialTimeout.Duration = 200 * time.Millisecond } + return &Client{Client: cli} } func (c *Client) doRequest(ctx context.Context, host, path string, args rpc2.Marshaler, ret rpc2.Unmarshaler) (err error) { diff --git a/blobstore/api/shardnode/shard.go b/blobstore/api/shardnode/shard.go index a9a8e6157..b82bcdd8d 100644 --- a/blobstore/api/shardnode/shard.go +++ b/blobstore/api/shardnode/shard.go @@ -47,7 +47,7 @@ func (c *Client) ListShards(ctx context.Context, host string, args ListShardArgs return } -func (c *Client) ListVolume(ctx context.Context, host string, args ListShardArgs) (ret ListVolumeRet, err error) { +func (c *Client) ListVolume(ctx context.Context, host string, args ListVolumeArgs) (ret ListVolumeRet, err error) { err = c.doRequest(ctx, host, "/volume/list", &args, &ret) return } diff --git a/blobstore/cmd/cmd.go b/blobstore/cmd/cmd.go index 476e6894c..55bcc6d93 100644 --- a/blobstore/cmd/cmd.go +++ b/blobstore/cmd/cmd.go @@ -140,8 +140,7 @@ func Main(args []string) { // new profile handler firstly profileHandler := profile.NewProfileHandler(cfg.BindAddr) - isMod2 := mod.SetUp2 != nil - if !isMod2 && mod.graceful { + if mod.SetUp != nil && mod.graceful { programEntry := func(state *graceful.State) { router, handlers := mod.SetUp() httpServer := &http.Server{ @@ -176,9 +175,8 @@ func Main(args []string) { return } - var shutdown interface{ Shutdown(context.Context) error } - - if isMod2 { + var shutdowns []func(context.Context) + if mod.SetUp2 != nil { router, interceptors := mod.SetUp2() rpc2Server := cfg.Rpc2Server rpc2Server.Handler = rpc2Handler(router, lh, cfg.Auth, interceptors) @@ -188,8 +186,10 @@ func Main(args []string) { log.Fatalf("rpc2 Server exits, err: %v", err) } }() - shutdown = rpc2Server - } else { + shutdowns = append(shutdowns, func(ctx context.Context) { rpc2Server.Shutdown(ctx) }) + } + + if mod.SetUp != nil { router, handlers := mod.SetUp() httpServer := &http.Server{ Addr: cfg.BindAddr, @@ -204,7 +204,7 @@ func Main(args []string) { log.Fatalf("Server exits, err: %v", err) } }() - shutdown = httpServer + shutdowns = append(shutdowns, func(ctx context.Context) { httpServer.Shutdown(ctx) }) } // wait for signal @@ -214,7 +214,9 @@ func Main(args []string) { log.Infof("receive signal: %s, stop service...", sig.String()) ctx, cancel := context.WithTimeout(context.Background(), time.Duration(cfg.ShutdownTimeoutS)*time.Second) defer cancel() - shutdown.Shutdown(ctx) + for _, shutdown := range shutdowns { + shutdown(ctx) + } if mod.TearDown != nil { mod.TearDown() diff --git a/blobstore/common/rpc/auth/handler.go b/blobstore/common/rpc/auth/handler.go index 647cbf2a7..8a4cdc2cd 100644 --- a/blobstore/common/rpc/auth/handler.go +++ b/blobstore/common/rpc/auth/handler.go @@ -46,11 +46,7 @@ func (h *handler) Handler(w http.ResponseWriter, req *http.Request, f func(http. func (h *handler) Handle(w rpc2.ResponseWriter, req *rpc2.Request, f rpc2.Handle) error { if err := proto.Decode(req.Header.Get(proto.TokenHeaderKey), []byte(req.RemotePath), h.Secret); err != nil { - return &rpc2.Error{ - Status: http.StatusForbidden, - Reason: "Auth", - Detail: err.Error(), - } + return rpc2.NewError(http.StatusForbidden, "Auth", err.Error()) } return f(w, req) } diff --git a/blobstore/common/rpc2/checksum.go b/blobstore/common/rpc2/checksum.go index 89cc6b80a..db87b5711 100644 --- a/blobstore/common/rpc2/checksum.go +++ b/blobstore/common/rpc2/checksum.go @@ -78,12 +78,10 @@ func unmarshalBlock(b []byte) (*ChecksumBlock, error) { } func checksumError(block ChecksumBlock, exp, act []byte) *Error { - return &Error{ - Status: 400, - Reason: "Checksum", - Detail: fmt.Sprintf("rpc2: internal checksum algorithm(%s) direction(%s) exp(%v) act(%v)", + return NewError(400, "Checksum", + fmt.Sprintf("rpc2: internal checksum algorithm(%s) direction(%s) exp(%v) act(%v)", block.Algorithm.String(), block.Direction.String(), block.Readable(exp), block.Readable(act)), - } + ) } // body encoder and decoder diff --git a/blobstore/common/rpc2/client.go b/blobstore/common/rpc2/client.go index ae68100f5..2dd381657 100644 --- a/blobstore/common/rpc2/client.go +++ b/blobstore/common/rpc2/client.go @@ -19,6 +19,7 @@ import ( "context" "io" "strings" + "sync/atomic" "time" "github.com/cubefs/cubefs/blobstore/common/rpc" @@ -44,7 +45,27 @@ type Client struct { Auth auth_proto.Config `json:"auth"` Selector rpc.Selector `json:"-"` // lb client - LbConfig rpc.LbConfig `json:"lb"` + LbConfig struct { + Hosts []string `json:"hosts"` + BackupHosts []string `json:"backup_hosts"` + HostTryTimes int `json:"host_try_times"` + FailRetryIntervalS int `json:"fail_retry_interval_s"` + MaxFailsPeriodS int `json:"max_fails_period_s"` + } `json:"lb"` + + // dead-lock copied Client when initOnce == 1 + initOnce uint32 // 0 uninitialised, 1 doing, 2 done +} + +// Request simple request, parameter and result both in body. +func (c *Client) Request(ctx context.Context, addr, path string, + para Marshaler, ret Unmarshaler, +) error { + req, err := NewRequest(ctx, addr, path, nil, Codec2Reader(para)) + if err != nil { + return err + } + return c.DoWith(req, ret) } func (c *Client) DoWith(req *Request, ret Unmarshaler) error { @@ -56,10 +77,16 @@ func (c *Client) DoWith(req *Request, ret Unmarshaler) error { } func (c *Client) Do(req *Request, ret Unmarshaler) (resp *Response, err error) { - if c.Connector == nil { // init + if c.lockInit() { defaulter.LessOrEqual(&c.Retry, 3) - c.Connector = defaultConnector(c.ConnectorConfig) c.newSelector() + if c.Connector == nil { + c.Connector = defaultConnector(c.ConnectorConfig) + } + if c.RetryOn == nil { + c.RetryOn = func(err error) bool { return DetectStatusCode(err) >= 500 } + } + atomic.StoreUint32(&c.initOnce, 2) } var lbHost rpc.UniqueHost @@ -124,6 +151,18 @@ func (c *Client) Close() error { return c.Connector.Close() } +func (c *Client) lockInit() bool { + if atomic.LoadUint32(&c.initOnce) >= 2 { + return false + } + for !atomic.CompareAndSwapUint32(&c.initOnce, 0, 1) { + if atomic.LoadUint32(&c.initOnce) >= 2 { + return false + } + } + return true +} + func (c *Client) do(req *Request, ret Unmarshaler) (*Response, error) { req.Header.SetStable() req.Trailer.SetStable() diff --git a/blobstore/common/rpc2/error.go b/blobstore/common/rpc2/error.go index 0cb90454a..c56b52767 100644 --- a/blobstore/common/rpc2/error.go +++ b/blobstore/common/rpc2/error.go @@ -36,6 +36,10 @@ func DetectError(err error) (int, string, error) { var _ rpc.HTTPError = (*Error)(nil) +func NewError(status int32, reason, detail string) *Error { + return &Error{Status: status, Reason: reason, Detail: detail} +} + func (m *Error) Unwrap() error { return errors.New(m.Error()) } func (m *Error) StatusCode() int { return int(m.GetStatus()) } func (m *Error) ErrorCode() string { return m.GetReason() } diff --git a/blobstore/common/rpc2/example/client.go b/blobstore/common/rpc2/example/client.go index ee81b6b2b..132a54eb9 100644 --- a/blobstore/common/rpc2/example/client.go +++ b/blobstore/common/rpc2/example/client.go @@ -53,13 +53,10 @@ func runClient() { } { para := pingPara{I: 7, S: "ping string"} - req, _ := rpc2.NewRequest(context.Background(), - listenon[int(time.Now().UnixNano())%len(listenon)], - "/kick", nil, rpc2.Codec2Reader(¶)) - req.ContentLength = int64(para.Size()) - log.Infof("before request para : %+v", para) - if err := client.DoWith(req, ¶); err != nil { + if err := client.Request(context.Background(), + listenon[int(time.Now().UnixNano())%len(listenon)], + "/kick", ¶, ¶); err != nil { panic(rpc2.ErrorString(err)) } log.Infof("after request result : %+v", para) diff --git a/blobstore/common/rpc2/example/main.go b/blobstore/common/rpc2/example/main.go index 715206a61..5d82f6412 100644 --- a/blobstore/common/rpc2/example/main.go +++ b/blobstore/common/rpc2/example/main.go @@ -6,11 +6,13 @@ import ( "github.com/cubefs/cubefs/blobstore/util/log" ) -var listenon = []string{"localhost:9998", "localhost:9999"} +var ( + listenrpc = "localhost:9997" + listenon = []string{"localhost:9998", "localhost:9999"} -var mode = flag.String("mode", "server", "run mode") + mode = flag.String("mode", "server", "run mode") +) -// main: go run main.go server.go client.go func main() { flag.Parse() diff --git a/blobstore/common/rpc2/example/server.go b/blobstore/common/rpc2/example/server.go index 4ab6de80d..f2592de16 100644 --- a/blobstore/common/rpc2/example/server.go +++ b/blobstore/common/rpc2/example/server.go @@ -3,11 +3,13 @@ package main import ( "bytes" "io" + "net/http" "os" "path" "github.com/cubefs/cubefs/blobstore/cmd" "github.com/cubefs/cubefs/blobstore/common/config" + "github.com/cubefs/cubefs/blobstore/common/rpc" "github.com/cubefs/cubefs/blobstore/common/rpc2" "github.com/cubefs/cubefs/blobstore/util/log" ) @@ -20,10 +22,11 @@ func init() { mod := &cmd.Module{ Name: "example_rpc2", InitConfig: initConfig, + SetUp: setUp, SetUp2: setUp2, TearDown: func() {}, } - cmd.RegisterGracefulModule(mod) + cmd.RegisterModule(mod) } func initConfig(args []string) (*cmd.Config, error) { @@ -36,6 +39,7 @@ func initConfig(args []string) (*cmd.Config, error) { os.MkdirAll(logDir, 0o644) conf.AuditLog.LogDir = logDir conf.LogConf.Filename = path.Join(logDir, "rpc2.log") + conf.BindAddr = listenrpc conf.Rpc2Server.Addresses = []rpc2.NetworkAddress{ {Network: "tcp", Address: listenon[0]}, {Network: "tcp", Address: listenon[1]}, @@ -43,6 +47,12 @@ func initConfig(args []string) (*cmd.Config, error) { return &conf.Config, nil } +func setUp() (*rpc.Router, []rpc.ProgressHandler) { + router := rpc.New() + router.Handle(http.MethodGet, "/rpc", func(c *rpc.Context) { c.Respond() }) + return router, nil +} + func setUp2() (*rpc2.Router, []rpc2.Interceptor) { router := &rpc2.Router{} router.Middleware(handleMiddleware1, handleMiddleware2) diff --git a/blobstore/common/rpc2/example_test.go b/blobstore/common/rpc2/example_test.go index 906559f0b..836a976dc 100644 --- a/blobstore/common/rpc2/example_test.go +++ b/blobstore/common/rpc2/example_test.go @@ -91,8 +91,7 @@ func ExampleServer_request_message() { args := &strMessage{str: "request message"} // message in request & response body - req, _ := NewRequest(testCtx, server.Name, "/", nil, Codec2Reader(args)) - if err := cli.DoWith(req, args); err != nil { + if err := cli.Request(testCtx, server.Name, "/", args, args); err != nil { fmt.Println(err) } fmt.Println(args.str) diff --git a/blobstore/common/rpc2/request.go b/blobstore/common/rpc2/request.go index 70d3ed8ea..a42ef6342 100644 --- a/blobstore/common/rpc2/request.go +++ b/blobstore/common/rpc2/request.go @@ -132,11 +132,7 @@ func (req *Request) request(deadline time.Time) (*Response, error) { } if resp.Status < 200 || resp.Status >= 300 { frame.Close() - return nil, &Error{ - Status: resp.Status, - Reason: resp.Reason, - Detail: resp.Error, - } + return nil, NewError(resp.Status, resp.Reason, resp.Error) } decode := req.checksum != nil && req.checksum.Direction.IsDownload() diff --git a/blobstore/common/rpc2/response_test.go b/blobstore/common/rpc2/response_test.go index 17527c6ff..2350be443 100644 --- a/blobstore/common/rpc2/response_test.go +++ b/blobstore/common/rpc2/response_test.go @@ -40,7 +40,7 @@ func handleResponseDoubleStatus(w ResponseWriter, req *Request) error { // response has wrote 200 OK func handleResponseAfterError(w ResponseWriter, req *Request) error { - w.AfterBody(func() error { return &Error{Status: 511, Detail: "after body"} }) + w.AfterBody(func() error { return NewError(511, "", "after body") }) return w.WriteOK(nil) } diff --git a/blobstore/common/rpc2/router.go b/blobstore/common/rpc2/router.go index 375cff20e..e438c1c2c 100644 --- a/blobstore/common/rpc2/router.go +++ b/blobstore/common/rpc2/router.go @@ -42,11 +42,7 @@ var defaultPanicHandler = func(_ ResponseWriter, req *Request, err interface{}, span := req.Span() span.Errorf("panic fired in path:%s -> %v\n", req.RemotePath, err) span.Error(string(stack)) - return &Error{ - Status: DefaultStatusPanic, - Reason: "HandlePanic", - Detail: fmt.Sprintf("panic(%v)", err), - } + return NewError(DefaultStatusPanic, "HandlePanic", fmt.Sprintf("panic(%v)", err)) } return nil } @@ -129,11 +125,7 @@ func (r *Router) handleWithPanic(h Handle) Handle { func (r *Router) handle(w ResponseWriter, req *Request) (err error) { handle, exist := r.handlers[req.RemotePath] if !exist { - err = &Error{ - Status: 404, - Reason: "NoRouter", - Detail: fmt.Sprintf("no router for path(%s)", req.RemotePath), - } + err = NewError(404, "NoRouter", fmt.Sprintf("no router for path(%s)", req.RemotePath)) return } diff --git a/blobstore/common/rpc2/rpc2.go b/blobstore/common/rpc2/rpc2.go index ce2828fb8..a9d0c2788 100644 --- a/blobstore/common/rpc2/rpc2.go +++ b/blobstore/common/rpc2/rpc2.go @@ -148,6 +148,7 @@ func (noParameter) Unmarshal([]byte) error { return nil } type codecReadWriter struct { once sync.Once + reader io.Reader marshaler Marshaler unmarshaler Unmarshaler } @@ -162,7 +163,19 @@ func (c *codecReadWriter) Size() int { // Read reader marshal to func (c *codecReadWriter) Read(p []byte) (n int, err error) { n, err = 0, io.EOF - c.once.Do(func() { n, err = c.marshaler.MarshalTo(p) }) + c.once.Do(func() { + if len(p) < c.marshaler.Size() { + var buff []byte + if buff, err = c.marshaler.Marshal(); err == nil { + c.reader = bytes.NewReader(buff) + } + } else { + n, err = c.marshaler.MarshalTo(p) + } + }) + if c.reader != nil { + n, err = c.reader.Read(p) + } return } diff --git a/blobstore/common/rpc2/rpc2_test.go b/blobstore/common/rpc2/rpc2_test.go index ac138cdd4..3bfe44895 100644 --- a/blobstore/common/rpc2/rpc2_test.go +++ b/blobstore/common/rpc2/rpc2_test.go @@ -180,6 +180,43 @@ func BenchmarkUploadDownload(b *testing.B) { } } +func TestRpc2CodecReader(t *testing.T) { + var req RequestHeader + req.TraceID = "test rpc2 codec reader" + + size := req.Size() + { + buff := make([]byte, size-1) + r := Codec2Reader(&req) + n, err := r.Read(buff) + require.NoError(t, err) + require.Equal(t, size-1, n) + n, err = r.Read(buff) + require.NoError(t, err) + require.Equal(t, 1, n) + _, err = r.Read(buff) + require.ErrorIs(t, io.EOF, err) + } + { + buff := make([]byte, size) + r := Codec2Reader(&req) + n, err := r.Read(buff) + require.NoError(t, err) + require.Equal(t, size, n) + _, err = r.Read(buff) + require.ErrorIs(t, io.EOF, err) + } + { + buff := make([]byte, size+1) + r := Codec2Reader(&req) + n, err := r.Read(buff) + require.NoError(t, err) + require.Equal(t, size, n) + _, err = r.Read(buff) + require.ErrorIs(t, io.EOF, err) + } +} + func TestRpc2None(t *testing.T) { { var x noneCodec @@ -360,11 +397,7 @@ func TestRpc2Pb(t *testing.T) { v.GetReason() v.GetDetail() run(v, false, true) - v = &Error{ - Status: 100, - Reason: "R", - Detail: "E", - } + v = NewError(100, "R", "E") run(v, false, false) v.GetStatus() v.GetReason() diff --git a/blobstore/common/rpc2/server.go b/blobstore/common/rpc2/server.go index 20cec092f..502b9c92f 100644 --- a/blobstore/common/rpc2/server.go +++ b/blobstore/common/rpc2/server.go @@ -100,7 +100,8 @@ func (s *Server) stating() { log.Debugf("server has %d listeners", len(s.listeners)) log.Debugf("server has %d sessions", len(s.sessions)) for sess := range s.sessions { - log.Debugf("session %v has %d streams", sess.LocalAddr(), sess.NumStreams()) + log.Debugf("session (%v - %v) has %d streams", + sess.LocalAddr(), sess.RemoteAddr(), sess.NumStreams()) } s.mu.Unlock() } @@ -299,7 +300,9 @@ func (s *Server) handleStream(stream *transport.Stream) { } } - resp.WriteOK(nil) + if err = resp.WriteOK(nil); err != nil { + return err + } if err = resp.Flush(); err != nil { return err } diff --git a/blobstore/common/rpc2/stream.go b/blobstore/common/rpc2/stream.go index 62066ae67..3d1615fd6 100644 --- a/blobstore/common/rpc2/stream.go +++ b/blobstore/common/rpc2/stream.go @@ -185,11 +185,7 @@ func (cs *clientStream) RecvMsg(a any) (err error) { cs.trailer.Merge(resp.Trailer.ToHeader()) cs.req.client.Connector.Put(cs.req.Context(), cs.req.conn, true) if resp.Status != 200 { - return &Error{ - Status: resp.Status, - Reason: resp.Reason, - Detail: resp.Error, - } + return NewError(resp.Status, resp.Reason, resp.Error) } return io.EOF } diff --git a/blobstore/common/rpc2/stream_test.go b/blobstore/common/rpc2/stream_test.go index 5054e558d..5cbf906da 100644 --- a/blobstore/common/rpc2/stream_test.go +++ b/blobstore/common/rpc2/stream_test.go @@ -62,6 +62,9 @@ func handleStreamFull(_ ResponseWriter, req *Request) error { } func TestStreamBase(t *testing.T) { + var tc *TransportConfig + require.Nil(t, tc.Transport()) + handler := &Router{} handler.Register("/", handleStreamFull) server, cli, shutdown := newServer("tcp", handler) diff --git a/blobstore/shardnode/rpcservice.go b/blobstore/shardnode/rpcservice.go index 9d5dc90d7..21bc5d6f4 100644 --- a/blobstore/shardnode/rpcservice.go +++ b/blobstore/shardnode/rpcservice.go @@ -135,7 +135,7 @@ func (s *RpcService) UpdateItem(w rpc2.ResponseWriter, req *rpc2.Request) error span := req.Span() args := &shardnode.UpdateItemArgs{} - if err := args.Unmarshal(req.Parameter); err != nil { + if err := req.ParseParameter(args); err != nil { return err } span.Debugf("receive UpdateItem request, args:%+v", args) @@ -148,7 +148,7 @@ func (s *RpcService) DeleteItem(w rpc2.ResponseWriter, req *rpc2.Request) error span := req.Span() args := &shardnode.DeleteItemArgs{} - if err := args.Unmarshal(req.Parameter); err != nil { + if err := req.ParseParameter(args); err != nil { return err } span.Debugf("receive DeleteItem request, args:%+v", args)