feat(rpc2): cache server request and response in pool

. #23076156

Signed-off-by: slasher <shenjie1@oppo.com>
This commit is contained in:
slasher 2025-02-24 17:08:57 +08:00
parent e0a59713ff
commit f7664eec0f
16 changed files with 313 additions and 79 deletions

View File

@ -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
}

View File

@ -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
}

View File

@ -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()
}
},
)

View File

@ -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:

View File

@ -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

View File

@ -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()

View File

@ -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")

View File

@ -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)

View File

@ -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)
}

View File

@ -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
}

View File

@ -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)))
}

View File

@ -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
}

View File

@ -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 {

View File

@ -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)
}

View File

@ -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}
}

View File

@ -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
}