diff --git a/blobstore/common/rpc2/body.go b/blobstore/common/rpc2/body.go index 2e7b00f42..08141ef9a 100644 --- a/blobstore/common/rpc2/body.go +++ b/blobstore/common/rpc2/body.go @@ -107,7 +107,7 @@ func (r *bodyAndTrailer) WriteTo(w io.Writer) (int64, error) { func (r *bodyAndTrailer) Close() error { r.storeError(r.tryReadTrailer()) r.closeOnce.Do(func() { - err := r.sr.Close() + err := r.br.Close() if cli := r.req.client; cli != nil { cli.Connector.Put(r.req.Context(), r.req.conn, err != nil || !r.sr.Finished()) @@ -127,7 +127,7 @@ func makeBodyWithTrailer(sr *transport.SizedReader, req *Request, ) Body { var br Body if decode { - br = newEdBody(*req.checksum, sr, int(l), false) + br = newEdBody(req.checksum, sr, int(l), false) } else { br = sr } diff --git a/blobstore/common/rpc2/checksum.go b/blobstore/common/rpc2/checksum.go index 5065be304..3dcc30f6f 100644 --- a/blobstore/common/rpc2/checksum.go +++ b/blobstore/common/rpc2/checksum.go @@ -24,12 +24,51 @@ import ( "hash" "hash/crc32" "io" + "sync" "github.com/zeebo/xxh3" ) const DefaultBlockSize = 64 << 10 +var ( + // (4) crc32.Size, (8) xxh3.New().Size() + sumPool = sync.Pool{ + New: func() any { + buff := make([]byte, 8) + return &buff + }, + } + bodywtPool = sync.Pool{ + New: func() any { + return &edBodyWriter{} + }, + } + bodyPools = map[ChecksumBlock]*sync.Pool{} +) + +func init() { + for _, alg := range []ChecksumAlgorithm{ + ChecksumAlgorithm_Crc_IEEE, + ChecksumAlgorithm_Hash_xxh3, + } { + for _, size := range []uint32{32 << 10, 64 << 10} { + block := ChecksumBlock{Algorithm: alg, BlockSize: size} + bodyPools[block] = &sync.Pool{ + New: func() any { + hasher := block.Hasher() + return &edBody{ + block: block, + hasher: hasher, + cell: make([]byte, hasher.Size()), + } + }, + } + + } + } +} + var algorithms = map[ChecksumAlgorithm]func() hash.Hash{ ChecksumAlgorithm_Crc_IEEE: func() hash.Hash { return crc32.NewIEEE() }, ChecksumAlgorithm_Hash_xxh3: func() hash.Hash { return xxh3.New() }, @@ -44,7 +83,7 @@ func (cd ChecksumDirection) IsDownload() bool { } func (cb *ChecksumBlock) EncodeSize(originalSize int64) int64 { - if cb == nil { + if cb == nil || *cb == (ChecksumBlock{}) { return originalSize } hasher := algorithms[cb.Algorithm]() @@ -66,15 +105,15 @@ func (cb *ChecksumBlock) Readable(b []byte) any { } } -func unmarshalBlock(b []byte) (*ChecksumBlock, error) { +func unmarshalBlock(b []byte) (ChecksumBlock, error) { var block ChecksumBlock if err := block.Unmarshal(b); err != nil { - return nil, fmt.Errorf("rpc2: internal checksum %s", err.Error()) + return block, fmt.Errorf("rpc2: internal checksum %s", err.Error()) } if _, exist := algorithms[block.Algorithm]; !exist || block.BlockSize == 0 { - return nil, fmt.Errorf("rpc2: checksum(%s) not implements", block.String()) + return block, fmt.Errorf("rpc2: checksum(%s) not implements", block.String()) } - return &block, nil + return block, nil } func checksumError(block ChecksumBlock, exp, act []byte) *Error { @@ -83,6 +122,17 @@ func checksumError(block ChecksumBlock, exp, act []byte) *Error { ) } +func compare(block ChecksumBlock, exp []byte, hasher hash.Hash) (err error) { + pbuff := sumPool.Get().(*[]byte) + act := (*pbuff)[:hasher.Size()] + hasher.Sum(act[:0]) + if !bytes.Equal(exp, act) { + err = checksumError(block, exp, act) + } + sumPool.Put(pbuff) // nolint: staticcheck + return +} + // body encoder and decoder type edBody struct { block ChecksumBlock @@ -99,6 +149,23 @@ type edBody struct { } func newEdBody(block ChecksumBlock, body Body, remain int, encode bool) *edBody { + cacheBlock := ChecksumBlock{ + Algorithm: block.Algorithm, + BlockSize: block.BlockSize, + } + pool, has := bodyPools[cacheBlock] + if has { + r := pool.Get().(*edBody) + r.encode = encode + r.hasher.Reset() + r.remain = remain + r.nx = 0 + r.cx = -1 + r.err = nil + r.Body = body + return r + } + hasher := block.Hasher() return &edBody{ block: block, @@ -151,7 +218,7 @@ func (r *edBody) encodeRead(p []byte) (nn int, err error) { r.remain -= n if r.nx == blockSize || r.remain == 0 { - copy(r.cell, r.hasher.Sum(nil)) + r.hasher.Sum(r.cell[:0]) r.hasher.Reset() r.cx = 0 r.nx = 0 @@ -174,9 +241,7 @@ func (r *edBody) decodeRead(p []byte) (nn int, err error) { return 0, err } - act := r.hasher.Sum(nil) - if !bytes.Equal(r.cell, act) { - r.err = checksumError(r.block, r.cell, act) + if r.err = compare(r.block, r.cell, r.hasher); r.err != nil { return 0, r.err } @@ -208,9 +273,7 @@ func (r *edBody) decodeRead(p []byte) (nn int, err error) { return 0, err } - act := r.hasher.Sum(nil) - if !bytes.Equal(r.cell, act) { - r.err = checksumError(r.block, r.cell, act) + if r.err = compare(r.block, r.cell, r.hasher); r.err != nil { return 0, r.err } @@ -241,7 +304,22 @@ func (r *edBody) WriteTo(w io.Writer) (int64, error) { if r.encode { return r.Body.WriteTo(w) } - return r.Body.WriteTo(&edBodyWriter{edBody: r, w: w}) + wt := bodywtPool.Get().(*edBodyWriter) + wt.edBody = r + wt.w = w + nn, err := r.Body.WriteTo(wt) + bodywtPool.Put(wt) // nolint: staticcheck + return nn, err +} + +func (r *edBody) Close() (err error) { + err = r.Body.Close() + pool, has := bodyPools[r.block] + if has { + r.Body = nil + pool.Put(r) // nolint: staticcheck + } + return } type edBodyWriter struct { @@ -264,9 +342,7 @@ func (r *edBodyWriter) Write(p []byte) (nn int, err error) { return } - act := r.hasher.Sum(nil) - if !bytes.Equal(r.cell, act) { - r.err = checksumError(r.block, r.cell, act) + if r.err = compare(r.block, r.cell, r.hasher); r.err != nil { return 0, r.err } diff --git a/blobstore/common/rpc2/checksum_test.go b/blobstore/common/rpc2/checksum_test.go index 6968432e4..8a31d20a3 100644 --- a/blobstore/common/rpc2/checksum_test.go +++ b/blobstore/common/rpc2/checksum_test.go @@ -163,6 +163,7 @@ func TestEncodeDecodeBodyBase(t *testing.T) { require.Equal(t, int64(len(transBody.data)), nn) _, err = encodeBody.Read(make([]byte, 1)) require.Equal(t, io.EOF, err, logName) + encodeBody.Close() transBody.off = 0 decodeBody := newEdBody(block, transBody, size, false) @@ -389,6 +390,7 @@ func BenchmarkEncodeDecodeBody(b *testing.B) { decodeBody := newEdBody(block, transBody, cs.size, false) serverBody := &noneReadWriter{} decodeBody.WriteTo(serverBody) + decodeBody.Close() } }, ) diff --git a/blobstore/common/rpc2/client.go b/blobstore/common/rpc2/client.go index 9ad5f45ea..6ab9e1c25 100644 --- a/blobstore/common/rpc2/client.go +++ b/blobstore/common/rpc2/client.go @@ -60,12 +60,14 @@ type Client struct { // Request simple request, parameter and result both in body. func (c *Client) Request(ctx context.Context, addr, path string, para Marshaler, ret Unmarshaler, -) error { +) (err error) { req, err := NewRequest(ctx, addr, path, nil, Codec2Reader(para)) if err != nil { return err } - return c.DoWith(req, ret) + err = c.DoWith(req, ret) + req.reuse() + return } func (c *Client) DoWith(req *Request, ret Unmarshaler) error { @@ -248,23 +250,33 @@ func NewRequest(ctx context.Context, addr, path string, para Marshaler, body io. if para == nil { para = NoParameter } - paraData, err := para.Marshal() - if err != nil { - return nil, err - } - req := &Request{ - RemoteAddr: addr, - RequestHeader: RequestHeader{ - Version: Version, - Magic: Magic, - RemotePath: path, - TraceID: getSpan(ctx).TraceID(), - Parameter: paraData, - }, - ctx: ctx, - Body: clientNopBody(rc), - AfterBody: func() error { return nil }, + + req := getRequest() + req.RemotePath = path + req.TraceID = getSpan(ctx).TraceID() + if psize := para.Size(); psize > 0 { + if cap(req.Parameter) >= psize { + nn, err := para.MarshalTo(req.Parameter[:psize]) + if err != nil { + return nil, err + } + if nn != psize { + return nil, io.ErrShortWrite + } + req.Parameter = req.Parameter[:psize] + } else { + paraData, err := para.Marshal() + if err != nil { + return nil, err + } + req.Parameter = paraData + } } + req.RemoteAddr = addr + req.ctx = ctx + req.Body = clientNopBody(rc) + req.AfterBody = func() error { return nil } + if body != nil { switch v := body.(type) { case *bytes.Buffer: diff --git a/blobstore/common/rpc2/client_test.go b/blobstore/common/rpc2/client_test.go index 6fa69ca00..7e4a57c8a 100644 --- a/blobstore/common/rpc2/client_test.go +++ b/blobstore/common/rpc2/client_test.go @@ -27,8 +27,10 @@ import ( type errNoParameter struct{ Codec } -func (errNoParameter) Marshal() ([]byte, error) { return nil, fmt.Errorf("codec") } -func (errNoParameter) Unmarshal([]byte) error { return fmt.Errorf("codec") } +func (errNoParameter) Size() int { return 1 } +func (errNoParameter) MarshalTo([]byte) (int, error) { return 0, fmt.Errorf("codec") } +func (errNoParameter) Marshal() ([]byte, error) { return nil, fmt.Errorf("codec") } +func (errNoParameter) Unmarshal([]byte) error { return fmt.Errorf("codec") } func TestClientRetry(t *testing.T) { var emptyCli Client diff --git a/blobstore/common/rpc2/connector.go b/blobstore/common/rpc2/connector.go index 8c65c0d5e..4897acd90 100644 --- a/blobstore/common/rpc2/connector.go +++ b/blobstore/common/rpc2/connector.go @@ -199,6 +199,11 @@ func (c *connector) get(ctx context.Context, addr string, newSession bool) (*tra if ses, ok = c.sessions[addr]; !ok { c.sessions[addr] = map[*transport.Session]struct{}{sess: {}} } else { + if len(ses) != sesLen { // add by other, try again + c.mu.Unlock() + sess.Close() + return c.get(ctx, addr, newSession) + } if len(ses) >= c.config.MaxSessionPerAddress { c.mu.Unlock() sess.Close() diff --git a/blobstore/common/rpc2/example_test.go b/blobstore/common/rpc2/example_test.go index c36a395ef..fd868b1bd 100644 --- a/blobstore/common/rpc2/example_test.go +++ b/blobstore/common/rpc2/example_test.go @@ -72,7 +72,6 @@ func (s *strMessageUnread) Readable() bool { return false } func handleMessage(w ResponseWriter, req *Request) error { var args strMessageUnread req.ParseParameter(&args) - req.Body.Close() args.Value = "-> " + args.Value return w.WriteOK(&args) } @@ -102,7 +101,6 @@ func handleUpload(w ResponseWriter, req *Request) error { return err } args.Value = fmt.Sprint(req.checksum.Readable(hasher.Sum(nil))) - req.Body.Close() return w.WriteOK(&args) } @@ -120,11 +118,11 @@ func ExampleServer_request_upload() { req, _ := NewRequest(testCtx, server.Name, "/", args, bytes.NewReader(buff)) req.OptionCrcUpload() req.ContentLength = int64(len(buff)) + hasher := req.checksum.Hasher() + hasher.Write(buff) if err := cli.DoWith(req, args); err != nil { fmt.Println(err) } - hasher := req.checksum.Hasher() - hasher.Write(buff) fmt.Println(args.Value == fmt.Sprint(req.checksum.Readable(hasher.Sum(nil)))) // Output: @@ -135,11 +133,9 @@ func handleUpDown(w ResponseWriter, req *Request) error { var args strMessage req.ParseParameter(&args) uhasher := req.checksum.Hasher() - if _, err := req.Body.WriteTo(LimitWriter(uhasher, req.ContentLength)); err != nil { return err } - req.Body.Close() buff := make([]byte, mrand.Intn(4<<20)+1<<20) crand.Read(buff) @@ -175,6 +171,7 @@ func ExampleServer_request_updown() { uhasher := req.checksum.Hasher() uhasher.Write(buff) dhasher := req.checksum.Hasher() + rr := req.checksum.Readable resp, _ := cli.Do(req, args) defer resp.Body.Close() @@ -187,7 +184,6 @@ func ExampleServer_request_updown() { break } } - rr := req.checksum.Readable fmt.Println(args.Value == fmt.Sprintf("%v %v", rr(uhasher.Sum(nil)), rr(dhasher.Sum(nil)))) // Output: @@ -197,7 +193,6 @@ func ExampleServer_request_updown() { func handleTrailer(w ResponseWriter, req *Request) error { hasher := md5.New() req.Body.WriteTo(LimitWriter(hasher, req.ContentLength)) - req.Body.Close() w.Header().Merge(req.Header) w.Header().Add("add", "header-stable") diff --git a/blobstore/common/rpc2/header.go b/blobstore/common/rpc2/header.go index 7e8c4d6fa..a147b35ae 100644 --- a/blobstore/common/rpc2/header.go +++ b/blobstore/common/rpc2/header.go @@ -86,6 +86,13 @@ func (h *Header) Merge(other Header) { } } +func (h *Header) Renew() { + h.stable = false + for key := range h.M { + delete(h.M, key) + } +} + func (fh *FixedHeader) newIfNil() { if fh.M == nil { fh.M = make(map[string]FixedValue) @@ -152,6 +159,13 @@ func (fh *FixedHeader) MergeHeader(h Header) { } } +func (fh *FixedHeader) Renew() { + fh.stable = false + for key := range fh.M { + delete(fh.M, key) + } +} + func (fh *FixedHeader) AllSize() (n int) { for _, v := range fh.M { n += int(v.Len) diff --git a/blobstore/common/rpc2/header_test.go b/blobstore/common/rpc2/header_test.go index b8ae7e03b..1f7c62121 100644 --- a/blobstore/common/rpc2/header_test.go +++ b/blobstore/common/rpc2/header_test.go @@ -55,6 +55,12 @@ func TestRpc2Header(t *testing.T) { header.Set("a", "a") require.True(t, header.Has("a")) + header.Renew() + require.NotNil(t, header.M) + header.Set("b", "b") + require.False(t, header.Has("a")) + require.True(t, header.Has("b")) + header.Reset() require.Nil(t, header.M) } @@ -91,6 +97,9 @@ func TestRpc2FixedHeader(t *testing.T) { header.Del("b") require.True(t, header.Has("b")) + header.Renew() + require.NotNil(t, header.M) + header.Reset() require.Nil(t, header.M) } diff --git a/blobstore/common/rpc2/request.go b/blobstore/common/rpc2/request.go index e2530737f..2c55d2493 100644 --- a/blobstore/common/rpc2/request.go +++ b/blobstore/common/rpc2/request.go @@ -18,6 +18,7 @@ import ( "context" "fmt" "io" + "sync" "time" "github.com/cubefs/cubefs/blobstore/common/rpc2/transport" @@ -36,7 +37,7 @@ type Request struct { opts []OptionRequest conn *transport.Stream - checksum *ChecksumBlock + checksum ChecksumBlock // server side cancel context.CancelFunc @@ -115,8 +116,8 @@ func (req *Request) write(deadline time.Time) error { size := _headerCell + reqHeaderSize + int(encodeLen) + req.Trailer.AllSize() req.conn.SetDeadline(deadline) - _, err := req.conn.SizedWrite(req.ctx, io.MultiReader(cell.Reader(), - Codec2Reader(&req.RequestHeader), + _, err := req.conn.SizedWrite(req.ctx, io.MultiReader( + codec2CellReader(cell, &req.RequestHeader), io.LimitReader(req.Body, encodeLen), // the body was encoded req.trailerReader(), ), size) @@ -137,7 +138,7 @@ func (req *Request) request(deadline time.Time) (*Response, error) { return nil, NewError(resp.Status, resp.Reason, resp.Error) } - decode := req.checksum != nil && req.checksum.Direction.IsDownload() + decode := req.checksum != ChecksumBlock{} && req.checksum.Direction.IsDownload() payloadSize := resp.Trailer.AllSize() if decode { payloadSize += int(req.checksum.EncodeSize(resp.ContentLength)) @@ -177,7 +178,7 @@ func (req *Request) OptionChecksum(block ChecksumBlock) *Request { if _, exist := algorithms[block.Algorithm]; !exist || block.BlockSize == 0 { panic(fmt.Sprintf("rpc2: checksum(%s) not implements", block.String())) } - if req.checksum != nil { + if req.checksum != (ChecksumBlock{}) { return req } cb, err := block.Marshal() @@ -185,7 +186,7 @@ func (req *Request) OptionChecksum(block ChecksumBlock) *Request { return req } - req.checksum = &block + req.checksum = block req.Header.Set(HeaderInternalChecksum, string(cb)) if req.ContentLength == 0 || !block.Direction.IsUpload() { return req @@ -219,3 +220,52 @@ func (req *Request) RemoteAddrString() string { } return "" } + +func (req *Request) reuse() { + putRequest(req) +} + +var poolRequest = sync.Pool{ + New: func() any { + return &Request{ + RequestHeader: RequestHeader{ + Version: Version, + Magic: Magic, + }, + } + }, +} + +func getRequest() *Request { + return poolRequest.Get().(*Request) +} + +func putRequest(req *Request) { + req.StreamCmd = StreamCmd_NOT + req.RemotePath = "" + req.TraceID = "" + req.ContentLength = 0 + req.Header.Renew() + req.Trailer.Renew() + req.Parameter = req.Parameter[:0] + + req.RemoteAddr = "" + req.BodyRead = 0 + + req.ctx = nil + req.client = nil + req.opts = req.opts[:0] + req.conn = nil + + req.checksum = ChecksumBlock{} + + req.cancel = nil + req.stream = nil + req.readablePara = false + + req.Body = nil + req.GetBody = nil + req.AfterBody = nil + + poolRequest.Put(req) // nolint: staticcheck +} diff --git a/blobstore/common/rpc2/request_test.go b/blobstore/common/rpc2/request_test.go index d006ce83d..74514a1d1 100644 --- a/blobstore/common/rpc2/request_test.go +++ b/blobstore/common/rpc2/request_test.go @@ -18,6 +18,7 @@ import ( "bytes" "context" crand "crypto/rand" + "encoding/binary" "errors" "fmt" "io" @@ -134,7 +135,6 @@ func handleBodyReadable(w ResponseWriter, req *Request) error { if err := req.ParseParameter(&args); err != nil { return err } - req.Body.Close() req.GetReadableParameter() if len(req.Parameter) == 0 { return errors.New("copy to parameter") @@ -179,12 +179,16 @@ func TestRequestRetryCrc(t *testing.T) { } require.Error(t, cli.DoWith(req, args)) + req, _ = NewRequest(testCtx, server.Name, "/", args, bytes.NewReader(buff)) + req.OptionCrcUpload() + req.ContentLength = int64(len(buff)) + 1 req.GetBody = func() (io.ReadCloser, error) { req.ContentLength = int64(len(buff)) return io.NopCloser(bytes.NewReader(buff)), nil } - require.NoError(t, cli.DoWith(req, args)) hasher := req.checksum.Hasher() hasher.Write(buff) - require.True(t, args.Value == fmt.Sprint(req.checksum.Readable(hasher.Sum(nil)))) + sum := hasher.Sum(nil) + require.NoError(t, cli.DoWith(req, args)) + require.True(t, args.Value == fmt.Sprintf("%d", binary.BigEndian.Uint32(sum))) } diff --git a/blobstore/common/rpc2/response.go b/blobstore/common/rpc2/response.go index 47f980ccb..12eeb9486 100644 --- a/blobstore/common/rpc2/response.go +++ b/blobstore/common/rpc2/response.go @@ -18,6 +18,7 @@ import ( "bytes" "context" "io" + "sync" "github.com/cubefs/cubefs/blobstore/common/rpc2/transport" ) @@ -140,7 +141,7 @@ func (resp *response) WriteHeader(status int, obj Marshaler) error { var cell headerCell cell.Set(resp.hdr.Size()) resp.toWrite += _headerCell + resp.hdr.Size() - resp.toList = append(resp.toList, cell.Reader(), Codec2Reader(&resp.hdr)) + resp.toList = append(resp.toList, codec2CellReader(cell, &resp.hdr)) return nil } @@ -230,8 +231,8 @@ func (resp *response) AfterBody(fn func() error) { } func (resp *response) options(req *Request) { - if req.checksum != nil && req.checksum.Direction.IsDownload() { - resp.bodyEncoder = newEdBody(*req.checksum, nil, 0, true) + if req.checksum != (ChecksumBlock{}) && req.checksum.Direction.IsDownload() { + resp.bodyEncoder = newEdBody(req.checksum, nil, 0, true) } } @@ -242,3 +243,47 @@ func (resp *response) encodeBody(r io.Reader) (io.Reader, int) { resp.bodyEncoder.Body = clientNopBody(io.NopCloser(r)) return resp.bodyEncoder, int(resp.bodyEncoder.block.EncodeSize(int64(resp.remain))) } + +func (resp *response) reuse() { + putResponse(resp) +} + +var poolResponse = sync.Pool{ + New: func() any { + return &response{ + hdr: ResponseHeader{ + Version: Version, + Magic: Magic, + }, + } + }, +} + +func getResponse() *response { + return poolResponse.Get().(*response) +} + +func putResponse(resp *response) { + resp.hdr.Status = 0 + resp.hdr.Reason = "" + resp.hdr.Error = "" + resp.hdr.ContentLength = 0 + resp.hdr.Header.Renew() + resp.hdr.Trailer.Renew() + resp.hdr.Parameter = resp.hdr.Parameter[:0] + + resp.ctx = nil + resp.conn = nil + resp.connBroken = false + + resp.hasWroteHeader = false + resp.hasWroteBody = false + resp.bodyEncoder = nil + + resp.remain = 0 + resp.toWrite = 0 + resp.toList = resp.toList[:0] + resp.afterBody = nil + + poolResponse.Put(resp) // nolint: staticcheck +} diff --git a/blobstore/common/rpc2/rpc2.go b/blobstore/common/rpc2/rpc2.go index 5764f06f3..c8cbbea07 100644 --- a/blobstore/common/rpc2/rpc2.go +++ b/blobstore/common/rpc2/rpc2.go @@ -173,6 +173,8 @@ type codecReadWriter struct { remain int unmarshaler Unmarshaler recv int + withcell bool + cell headerCell cache *bytes.Buffer } @@ -192,8 +194,21 @@ func (c *codecReadWriter) Read(p []byte) (n int, err error) { err = ErrFrameHeader return } + var nn int + if c.withcell { + if len(p) < len(c.cell) { + err = io.ErrShortBuffer + return + } + nn = copy(p, c.cell[:]) + n += nn + p = p[nn:] + c.withcell = false + } + if len(p) >= size { - n, err = c.marshaler.MarshalTo(p) + nn, err = c.marshaler.MarshalTo(p) + n += nn return } @@ -202,7 +217,6 @@ func (c *codecReadWriter) Read(p []byte) (n int, err error) { cache.Grow(size) cache.ReadFrom(util.DiscardReader(size)) - var nn int buff := cache.Bytes() nn, err = c.marshaler.MarshalTo(buff) if err != nil { @@ -273,6 +287,10 @@ func Codec2Writer(m Unmarshaler, size int) io.Writer { return &codecReadWriter{unmarshaler: m, recv: size} } +func codec2CellReader(cell headerCell, m Marshaler) io.Reader { + return &codecReadWriter{marshaler: m, withcell: true, cell: cell} +} + // LimitedWriter wrap Body with WriteTo type LimitedWriter struct { w io.Writer @@ -323,10 +341,6 @@ func (h *headerCell) Write(p []byte) (int, error) { return _headerCell, nil } -func (h headerCell) Reader() io.Reader { - return bytes.NewReader(h[:]) -} - func beforeContextDeadline(ctx context.Context, t time.Time) time.Time { d, ok := ctx.Deadline() if !ok { diff --git a/blobstore/common/rpc2/rpc2_test.go b/blobstore/common/rpc2/rpc2_test.go index 7bdf3c7d7..4d52be083 100644 --- a/blobstore/common/rpc2/rpc2_test.go +++ b/blobstore/common/rpc2/rpc2_test.go @@ -69,7 +69,6 @@ func (noCopyReadWriter) Write(p []byte) (int, error) { return len(p), nil } func handleNone(w ResponseWriter, req *Request) error { if req.ContentLength == 0 { - req.Body.Close() return w.WriteOK(nil) } req.Body.WriteTo(LimitWriter(noCopyReadWriter{}, req.ContentLength)) @@ -177,9 +176,11 @@ func BenchmarkUploadDownload(b *testing.B) { for i := 0; i < b.N; i++ { req, _ := NewRequest(testCtx, server.Name, "/", nil, noCopyReadWriter{}) req.ContentLength = l + req.OptionCrc() resp, _ := cli.Do(req, nil) resp.Body.WriteTo(LimitWriter(noCopyReadWriter{}, l)) resp.Body.Close() + req.reuse() } } @@ -280,8 +281,8 @@ func TestRpc2CodecWriter(t *testing.T) { func TestRpc2None(t *testing.T) { { var x noneCodec - x.Size() - x.Marshal() + _ = x.Size() + _, _ = x.Marshal() x.MarshalTo(nil) x.Unmarshal(nil) } diff --git a/blobstore/common/rpc2/server.go b/blobstore/common/rpc2/server.go index 7a6ad2de2..f52268540 100644 --- a/blobstore/common/rpc2/server.go +++ b/blobstore/common/rpc2/server.go @@ -273,7 +273,9 @@ func (s *Server) handleStream(stream *transport.Stream) { } ctx = req.Context() - resp := &response{ctx: req.ctx, conn: stream} + resp := getResponse() + resp.ctx = req.ctx + resp.conn = stream if ss := req.stream; ss != nil { if err = s.Handler.Handle(resp, req); err != nil { status, reason, detail := DetectError(err) @@ -316,11 +318,13 @@ func (s *Server) handleStream(stream *transport.Stream) { return errors.New("stream conn has broken") } req.cancel() + req.reuse() + resp.reuse() } }(); err != nil { span := getSpan(ctx) errMsg := fmt.Sprintf("stream(%d, %v, %v) %s", stream.ID(), stream.LocalAddr(), stream.RemoteAddr(), err.Error()) - if errors.Is(io.EOF, err) { + if errors.Is(err, io.EOF) { span.Warn(errMsg) } else { span.Error(errMsg) @@ -330,13 +334,13 @@ func (s *Server) handleStream(stream *transport.Stream) { } func (s *Server) readRequest(stream *transport.Stream) (*Request, error) { - var hdr RequestHeader - frame, err := readHeaderFrame(context.Background(), stream, &hdr) + req := getRequest() + frame, err := readHeaderFrame(context.Background(), stream, &req.RequestHeader) if err != nil { return nil, err } - switch hdr.StreamCmd { + switch req.StreamCmd { case StreamCmd_NOT, StreamCmd_SYN: case StreamCmd_PSH, StreamCmd_FIN: return nil, ErrFrameProtocol @@ -344,7 +348,7 @@ func (s *Server) readRequest(stream *transport.Stream) (*Request, error) { return nil, ErrFrameProtocol } - traceID := hdr.TraceID + traceID := req.TraceID if traceID == "" { traceID = trace.RandomID().String() } @@ -352,9 +356,10 @@ func (s *Server) readRequest(stream *transport.Stream) (*Request, error) { ctx, cancel := context.WithCancel(context.Background()) _, ctx = trace.StartSpanFromContextWithTraceID(ctx, "", traceID) - req := &Request{RequestHeader: hdr, ctx: ctx, conn: stream} + req.ctx = ctx + req.conn = stream req.cancel = cancel - if sum := hdr.Header.Get(HeaderInternalChecksum); sum != "" { + if sum := req.Header.Get(HeaderInternalChecksum); sum != "" { block, err := unmarshalBlock([]byte(sum)) if err != nil { frame.Close() @@ -363,7 +368,7 @@ func (s *Server) readRequest(stream *transport.Stream) (*Request, error) { req.checksum = block } - decode := req.checksum != nil && req.checksum.Direction.IsUpload() + decode := req.checksum != ChecksumBlock{} && req.checksum.Direction.IsUpload() payloadSize := req.Trailer.AllSize() if decode { payloadSize += int(req.checksum.EncodeSize(req.ContentLength)) @@ -373,7 +378,7 @@ func (s *Server) readRequest(stream *transport.Stream) (*Request, error) { req.Body = makeBodyWithTrailer(stream.NewSizedReader(req.ctx, payloadSize, frame), req, &req.Trailer, req.ContentLength, decode) - if hdr.StreamCmd == StreamCmd_SYN { + if req.StreamCmd == StreamCmd_SYN { req.stream = &serverStream{req: req} } diff --git a/blobstore/common/rpc2/stream.go b/blobstore/common/rpc2/stream.go index 771772567..3b86c40ce 100644 --- a/blobstore/common/rpc2/stream.go +++ b/blobstore/common/rpc2/stream.go @@ -310,7 +310,7 @@ func (ss *serverStream) writeFrameMsg(hdr *ResponseHeader, msg Marshaler) error } var cell headerCell cell.Set(hdr.Size()) - _, err := ss.req.conn.SizedWrite(ss.Context(), io.MultiReader(cell.Reader(), - Codec2Reader(hdr), Codec2Reader(msg)), size) + _, err := ss.req.conn.SizedWrite(ss.Context(), io.MultiReader( + codec2CellReader(cell, hdr), Codec2Reader(msg)), size) return err }