mirror of
https://github.com/cubefs/cubefs.git
synced 2026-08-02 02:00:56 +00:00
feat(rpc2): close recv and send if error in stream client
. #1000219604 Signed-off-by: slasher <shenjie1@oppo.com>
This commit is contained in:
parent
cab906aff9
commit
2de4f09154
@ -11,8 +11,8 @@ import (
|
||||
)
|
||||
|
||||
type (
|
||||
streamReq rpc2.AnyCodec[string]
|
||||
streamResp rpc2.AnyCodec[string]
|
||||
streamReq struct{ rpc2.AnyCodec[string] }
|
||||
streamResp struct{ rpc2.AnyCodec[string] }
|
||||
)
|
||||
|
||||
func runStream() {
|
||||
|
||||
@ -18,6 +18,7 @@ import (
|
||||
"context"
|
||||
"io"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
)
|
||||
|
||||
type ClientStream interface {
|
||||
@ -125,6 +126,8 @@ type clientStream struct {
|
||||
req *Request
|
||||
header Header // response header
|
||||
trailer Header // response trailer
|
||||
|
||||
closed atomic.Value
|
||||
}
|
||||
|
||||
var _ ClientStream = (*clientStream)(nil)
|
||||
@ -142,9 +145,12 @@ func (cs *clientStream) Trailer() Header {
|
||||
}
|
||||
|
||||
func (cs *clientStream) CloseSend() error {
|
||||
if err := cs.closedError(); err != nil {
|
||||
return err
|
||||
}
|
||||
req := cs.newRequest()
|
||||
req.StreamCmd = StreamCmd_FIN
|
||||
return req.write(req.client.requestDeadline(req.ctx))
|
||||
return cs.closeIfError(req.write(req.client.requestDeadline(req.ctx)))
|
||||
}
|
||||
|
||||
func (cs *clientStream) SendMsg(a any) error {
|
||||
@ -152,15 +158,18 @@ func (cs *clientStream) SendMsg(a any) error {
|
||||
if !is {
|
||||
panic("rpc2: stream send message must implement rpc2.Codec")
|
||||
}
|
||||
if err := cs.closedError(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
req := cs.newRequest()
|
||||
req.StreamCmd = StreamCmd_PSH
|
||||
if _headerCell+req.RequestHeader.Size()+msg.Size() > req.conn.MaxPayloadSize() {
|
||||
return ErrFrameHeader
|
||||
return cs.closeIfError(ErrFrameHeader)
|
||||
}
|
||||
req.ContentLength = int64(msg.Size())
|
||||
req.Body = clientNopBody(NopCloser(Codec2Reader(msg)))
|
||||
return req.write(req.client.requestDeadline(req.ctx))
|
||||
return cs.closeIfError(req.write(req.client.requestDeadline(req.ctx)))
|
||||
}
|
||||
|
||||
func (cs *clientStream) RecvMsg(a any) (err error) {
|
||||
@ -169,31 +178,34 @@ func (cs *clientStream) RecvMsg(a any) (err error) {
|
||||
panic("rpc2: stream recv message must implement rpc2.Codec")
|
||||
}
|
||||
conn := cs.req.conn
|
||||
if err = cs.closedError(); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
var resp ResponseHeader
|
||||
frame, err := readHeaderFrame(cs.Context(), conn, &resp)
|
||||
if err != nil {
|
||||
return err
|
||||
return cs.closeIfError(err)
|
||||
}
|
||||
defer func() {
|
||||
if errClose := frame.Close(); err == nil {
|
||||
err = errClose
|
||||
err = cs.closeIfError(errClose)
|
||||
}
|
||||
}()
|
||||
|
||||
if resp.Status > 0 { // end
|
||||
cs.trailer.Merge(resp.Trailer.ToHeader())
|
||||
cs.req.client.Connector.Put(cs.req.Context(), cs.req.conn, true)
|
||||
if resp.Status != 200 {
|
||||
return NewError(resp.Status, resp.Reason, resp.Error)
|
||||
err = NewError(resp.Status, resp.Reason, resp.Error)
|
||||
}
|
||||
return io.EOF
|
||||
err = io.EOF
|
||||
return cs.closeIfError(err)
|
||||
}
|
||||
|
||||
if int64(frame.Len()) < resp.ContentLength {
|
||||
return ErrFrameHeader
|
||||
return cs.closeIfError(ErrFrameHeader)
|
||||
}
|
||||
return msg.Unmarshal(frame.Bytes(int(resp.ContentLength)))
|
||||
return cs.closeIfError(msg.Unmarshal(frame.Bytes(int(resp.ContentLength))))
|
||||
}
|
||||
|
||||
func (cs *clientStream) newRequest() *Request {
|
||||
@ -210,6 +222,26 @@ func (cs *clientStream) newRequest() *Request {
|
||||
return req
|
||||
}
|
||||
|
||||
func (cs *clientStream) closedError() error {
|
||||
if val := cs.closed.Load(); val != nil {
|
||||
return val.(error)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (cs *clientStream) closeIfError(err error) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
if succ := cs.closed.CompareAndSwap(nil, err); succ {
|
||||
// close connection once
|
||||
cs.req.client.Connector.Put(cs.req.Context(), cs.req.conn, true)
|
||||
} else {
|
||||
err = cs.closed.Load().(error)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
type serverStream struct {
|
||||
req *Request
|
||||
|
||||
|
||||
@ -15,6 +15,7 @@
|
||||
package rpc2
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"io"
|
||||
"testing"
|
||||
|
||||
@ -116,3 +117,44 @@ func TestStreamBase(t *testing.T) {
|
||||
<-waitc
|
||||
require.Equal(t, "bbb", trailer.Get("stream-trailer"))
|
||||
}
|
||||
|
||||
func TestStreamClientClose(t *testing.T) {
|
||||
handler := &Router{}
|
||||
handler.Register("/", handleStreamFull)
|
||||
server, cli, shutdown := newServer("tcp", handler)
|
||||
defer shutdown()
|
||||
sc := StreamClient[streamReq, streamResp]{Client: cli}
|
||||
|
||||
var para strMessage
|
||||
req, err := NewStreamRequest(testCtx, server.Name, "/", ¶)
|
||||
require.NoError(t, err)
|
||||
|
||||
cc, err := sc.Streaming(req, ¶)
|
||||
require.NoError(t, err)
|
||||
|
||||
errrecv := make(chan error, 1)
|
||||
go func() {
|
||||
for {
|
||||
if _, errx := cc.Recv(); errx != nil {
|
||||
errrecv <- errx
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
for idx := range [10]struct{}{} {
|
||||
var req streamReq
|
||||
req.Value = fmt.Sprintf("request-%d", idx)
|
||||
if idx == 7 {
|
||||
buff := make([]byte, 2<<20)
|
||||
rand.Read(buff)
|
||||
req.Value = string(buff)
|
||||
require.ErrorIs(t, cc.Send(&req), ErrFrameHeader)
|
||||
break
|
||||
}
|
||||
require.NoError(t, cc.Send(&req))
|
||||
}
|
||||
err = cc.CloseSend()
|
||||
require.ErrorIs(t, err, ErrFrameHeader)
|
||||
err = <-errrecv
|
||||
require.ErrorIs(t, err, ErrFrameHeader)
|
||||
}
|
||||
|
||||
Loading…
Reference in New Issue
Block a user