feat(rpc2): close recv and send if error in stream client

. #1000219604

Signed-off-by: slasher <shenjie1@oppo.com>
This commit is contained in:
slasher 2025-07-08 19:19:18 +08:00
parent cab906aff9
commit 2de4f09154
3 changed files with 86 additions and 12 deletions

View File

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

View File

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

View File

@ -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, "/", &para)
require.NoError(t, err)
cc, err := sc.Streaming(req, &para)
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)
}