feat(rpc2): transport with context cancel

. #22548427

Signed-off-by: slasher <shenjie1@oppo.com>
This commit is contained in:
slasher 2024-08-28 19:06:25 +08:00
parent e41f7a497b
commit 946ec8516f
11 changed files with 136 additions and 51 deletions

View File

@ -15,6 +15,7 @@
package rpc2
import (
"context"
"io"
"sync"
@ -142,8 +143,8 @@ func makeBodyWithTrailer(sr *transport.SizedReader, req *Request,
}
// readHeaderFrame try to read request or response header.
func readHeaderFrame(stream *transport.Stream, hdr Unmarshaler) (*transport.FrameRead, error) {
frame, err := stream.ReadFrame()
func readHeaderFrame(ctx context.Context, stream *transport.Stream, hdr Unmarshaler) (*transport.FrameRead, error) {
frame, err := stream.ReadFrame(ctx)
if err != nil {
return nil, err
}

View File

@ -49,7 +49,7 @@ func TestRpc2ReadFrame(t *testing.T) {
frame, _ := conn.AllocFrame(1)
frame.Write([]byte{0xee})
conn.WriteFrame(frame)
_, err = conn.ReadFrame()
_, err = conn.ReadFrame(testCtx)
require.ErrorIs(t, io.EOF, err)
}
{
@ -58,7 +58,7 @@ func TestRpc2ReadFrame(t *testing.T) {
frame, _ := conn.AllocFrame(5)
frame.Write([]byte{0x1, 0x00, 0x00, 0x00})
conn.WriteFrame(frame)
_, err = conn.ReadFrame()
_, err = conn.ReadFrame(testCtx)
require.ErrorIs(t, io.EOF, err)
}
{
@ -67,7 +67,7 @@ func TestRpc2ReadFrame(t *testing.T) {
frame, _ := conn.AllocFrame(5)
frame.Write([]byte{0x1, 0x00, 0x00, 0x00, 0xee})
conn.WriteFrame(frame)
_, err = conn.ReadFrame()
_, err = conn.ReadFrame(testCtx)
require.ErrorIs(t, io.EOF, err)
}
}

View File

@ -37,13 +37,14 @@ type Request struct {
ctx context.Context
client *Client // client side
opts []OptionRequest
conn *transport.Stream
stream *serverStream
checksum *ChecksumBlock
opts []OptionRequest
checksum *ChecksumBlock
// server side
cancel context.CancelFunc
stream *serverStream
readablePara bool
Body Body
@ -112,7 +113,7 @@ func (req *Request) write(deadline time.Time) error {
size := _headerCell + reqHeaderSize + int(encodeLen) + req.Trailer.AllSize()
req.conn.SetDeadline(deadline)
_, err := req.conn.SizedWrite(io.MultiReader(cell.Reader(),
_, err := req.conn.SizedWrite(req.ctx, io.MultiReader(cell.Reader(),
req.RequestHeader.MarshalToReader(),
io.LimitReader(req.Body, encodeLen), // the body was encoded
req.trailerReader(),
@ -125,7 +126,7 @@ func (req *Request) request(deadline time.Time) (*Response, error) {
return nil, err
}
resp := &Response{Request: req}
frame, err := readHeaderFrame(req.conn, &resp.ResponseHeader)
frame, err := readHeaderFrame(req.ctx, req.conn, &resp.ResponseHeader)
if err != nil {
return nil, err
}
@ -145,7 +146,7 @@ func (req *Request) request(deadline time.Time) (*Response, error) {
} else {
payloadSize += int(resp.ContentLength)
}
resp.Body = makeBodyWithTrailer(req.conn.NewSizedReader(payloadSize, frame),
resp.Body = makeBodyWithTrailer(req.conn.NewSizedReader(req.ctx, payloadSize, frame),
req, &resp.Trailer, resp.ContentLength, decode)
return resp, nil
}

View File

@ -49,7 +49,7 @@ func TestRequestTimeout(t *testing.T) {
require.ErrorIs(t, transport.ErrTimeout, err)
cli.RequestTimeout.Duration = 0
ctx, cancel := context.WithDeadline(testCtx, time.Now().Add(100*time.Millisecond))
ctx, cancel := context.WithTimeout(testCtx, 100*time.Millisecond)
req, err = NewRequest(ctx, server.Name, "/", nil, nil)
require.NoError(t, err)
err = cli.DoWith(req, nil)
@ -75,7 +75,7 @@ func TestRequestTimeout(t *testing.T) {
cli.ResponseTimeout.Duration = 0
cli.Timeout.Duration = time.Second
ctx, cancel = context.WithDeadline(testCtx, time.Now().Add(100*time.Millisecond))
ctx, cancel = context.WithTimeout(testCtx, 100*time.Millisecond)
req, err = NewRequest(ctx, server.Name, "/none", nil, bytes.NewReader(buff))
require.NoError(t, err)
resp, err = cli.Do(req, nil)
@ -86,6 +86,23 @@ func TestRequestTimeout(t *testing.T) {
cancel()
}
func TestRequestContextCancel(t *testing.T) {
var handler Router
handler.Register("/", handleRequestTimeout)
server, cli, shutdown := newServer("tcp", &handler)
defer shutdown()
ctx, cancel := context.WithTimeout(testCtx, 100*time.Millisecond)
req, _ := NewRequest(ctx, server.Name, "/", nil, nil)
require.Error(t, cli.DoWith(req, nil))
cancel()
ctx, cancel = context.WithCancel(testCtx)
req, _ = NewRequest(ctx, server.Name, "/", nil, nil)
cancel()
require.Error(t, cli.DoWith(req, nil))
}
func TestRequestErrors(t *testing.T) {
addr, cli, shutdown := newTcpServer()
defer shutdown()

View File

@ -16,6 +16,7 @@ package rpc2
import (
"bytes"
"context"
"io"
"github.com/cubefs/cubefs/blobstore/common/rpc2/transport"
@ -67,6 +68,7 @@ func (resp *Response) ParseResult(ret Unmarshaler) error {
type response struct {
hdr ResponseHeader
ctx context.Context
conn *transport.Stream
connBroken bool
@ -197,7 +199,7 @@ func (resp *response) Flush() error {
if resp.connBroken {
return io.ErrClosedPipe
}
_, err := resp.conn.SizedWrite(io.MultiReader(resp.toList...), resp.toWrite)
_, err := resp.conn.SizedWrite(resp.ctx, io.MultiReader(resp.toList...), resp.toWrite)
if err != nil {
resp.connBroken = true
return err

View File

@ -270,7 +270,7 @@ func (s *Server) handleStream(stream *transport.Stream) {
}
ctx = req.Context()
resp := &response{conn: stream}
resp := &response{ctx: req.ctx, conn: stream}
if ss := req.stream; ss != nil {
if err = s.Handler.Handle(resp, req); err != nil {
status, reason, detail := DetectError(err)
@ -310,6 +310,7 @@ func (s *Server) handleStream(stream *transport.Stream) {
if resp.connBroken {
return errors.New("stream conn has broken")
}
req.cancel()
}
}(); err != nil {
span := getSpan(ctx)
@ -320,7 +321,7 @@ func (s *Server) handleStream(stream *transport.Stream) {
func (s *Server) readRequest(stream *transport.Stream) (*Request, error) {
var hdr RequestHeader
frame, err := readHeaderFrame(stream, &hdr)
frame, err := readHeaderFrame(context.Background(), stream, &hdr)
if err != nil {
return nil, err
}
@ -337,9 +338,12 @@ func (s *Server) readRequest(stream *transport.Stream) (*Request, error) {
if traceID == "" {
traceID = trace.RandomID().String()
}
_, ctx := trace.StartSpanFromContextWithTraceID(context.Background(), "", traceID)
ctx, cancel := context.WithCancel(context.Background())
_, ctx = trace.StartSpanFromContextWithTraceID(ctx, "", traceID)
req := &Request{RequestHeader: hdr, ctx: ctx, conn: stream}
req.cancel = cancel
if sum := hdr.Header.Get(HeaderInternalChecksum); sum != "" {
block, err := unmarshalBlock([]byte(sum))
if err != nil {
@ -356,7 +360,7 @@ func (s *Server) readRequest(stream *transport.Stream) (*Request, error) {
} else {
payloadSize += int(req.ContentLength)
}
req.Body = makeBodyWithTrailer(stream.NewSizedReader(payloadSize, frame),
req.Body = makeBodyWithTrailer(stream.NewSizedReader(req.ctx, payloadSize, frame),
req, &req.Trailer, req.ContentLength, decode)
if hdr.StreamCmd == StreamCmd_SYN {

View File

@ -171,7 +171,7 @@ func (cs *clientStream) RecvMsg(a any) (err error) {
conn := cs.req.conn
var resp ResponseHeader
frame, err := readHeaderFrame(conn, &resp)
frame, err := readHeaderFrame(cs.Context(), conn, &resp)
if err != nil {
return err
}
@ -277,7 +277,7 @@ func (ss *serverStream) RecvMsg(a any) (err error) {
}
var req RequestHeader
frame, err := readHeaderFrame(ss.req.conn, &req)
frame, err := readHeaderFrame(ss.Context(), ss.req.conn, &req)
if err != nil {
return
}
@ -314,7 +314,7 @@ func (ss *serverStream) writeFrameMsg(hdr *ResponseHeader, msg Marshaler) error
}
var cell headerCell
cell.Set(hdr.Size())
_, err := ss.req.conn.SizedWrite(io.MultiReader(cell.Reader(),
_, err := ss.req.conn.SizedWrite(ss.Context(), io.MultiReader(cell.Reader(),
hdr.MarshalToReader(), Codec2Reader(msg)), size)
return err
}

View File

@ -1,6 +1,7 @@
package transport
import (
"context"
"encoding/binary"
"fmt"
"io"
@ -94,6 +95,8 @@ type FrameWrite struct {
closer interface {
Free([]byte) error
}
ctx context.Context
}
func (f *FrameWrite) tryLock() bool {
@ -148,6 +151,17 @@ func (f *FrameWrite) Close() (err error) {
return
}
func (f *FrameWrite) Context() context.Context {
if f.ctx == nil {
return context.Background()
}
return f.ctx
}
func (f *FrameWrite) WithContext(ctx context.Context) {
f.ctx = ctx
}
// FrameRead frame for read
type FrameRead struct {
off int

View File

@ -595,6 +595,7 @@ func (s *Session) writeFrameInternal(f *FrameWrite, deadline writeDealine, class
writeCh = s.ctrl
}
ctx := f.Context()
select {
case <-deadline.wait:
return 0, ErrTimeout
@ -607,6 +608,8 @@ func (s *Session) writeFrameInternal(f *FrameWrite, deadline writeDealine, class
return 0, s.socketWriteError.Load().(error)
case <-deadline.wait:
return 0, ErrTimeout
case <-ctx.Done():
return 0, ctx.Err()
}
}
@ -618,6 +621,8 @@ func (s *Session) writeFrameInternal(f *FrameWrite, deadline writeDealine, class
return 0, io.ErrClosedPipe
case <-s.chSocketWriteError:
return 0, s.socketWriteError.Load().(error)
case <-ctx.Done():
return 0, ctx.Err()
}
}

View File

@ -2,6 +2,7 @@ package transport
import (
"bytes"
"context"
crand "crypto/rand"
"encoding/binary"
"fmt"
@ -18,6 +19,8 @@ import (
"time"
)
var testCtx = context.Background()
func init() {
runtime.GOMAXPROCS(1)
go func() {
@ -67,7 +70,7 @@ func handleConnection(tb testing.TB, conn net.Conn, v2 bool) {
if stream, err := session.AcceptStream(); err == nil {
go func(s *Stream) {
for {
fr, err := s.ReadFrame()
fr, err := s.ReadFrame(testCtx)
if err != nil {
if err != io.EOF {
tb.Error(err)
@ -98,12 +101,12 @@ func handleConnection(tb testing.TB, conn net.Conn, v2 bool) {
func writeThenRead(s *Stream, msg string) {
size := len(msg)
buf := make([]byte, size)
_, err := s.SizedWrite(strings.NewReader(msg), size)
_, err := s.SizedWrite(testCtx, strings.NewReader(msg), size)
if err != nil {
panic(err)
}
fr := s.SizedReader(size)
fr := s.SizedReader(testCtx, size)
defer fr.Close()
if _, err := fr.Read(buf); err != nil {
panic(err)
@ -140,7 +143,7 @@ func TestEcho(t *testing.T) {
stream.WriteFrame(fw)
fw.Close()
sent += msg
fr, err := stream.ReadFrame()
fr, err := stream.ReadFrame(testCtx)
if err != nil {
t.Fatal(err)
}
@ -215,7 +218,7 @@ func TestSpeed(t *testing.T) {
wg.Add(1)
go func() {
w := &sizedWriter{}
rc := stream.SizedReader(4096 * 4096)
rc := stream.SizedReader(testCtx, 4096*4096)
defer rc.Close()
_, err := rc.WriteTo(w)
if err != nil {
@ -231,7 +234,7 @@ func TestSpeed(t *testing.T) {
}()
msg := make([]byte, 8192)
for range [2048]struct{}{} {
stream.SizedWrite(bytes.NewReader(msg), len(msg))
stream.SizedWrite(testCtx, bytes.NewReader(msg), len(msg))
}
wg.Wait()
session.Close()
@ -250,11 +253,11 @@ func TestSizedReadWrite(t *testing.T) {
{
size := k64 / 2
_, err := stream.SizedWrite(bytes.NewReader(make([]byte, size)), size)
_, err := stream.SizedWrite(testCtx, bytes.NewReader(make([]byte, size)), size)
if err != nil {
t.Fatal(err)
}
rc := stream.SizedReader(size)
rc := stream.SizedReader(testCtx, size)
_, err = rc.WriteTo(rwErr{})
if err == nil {
t.Fatal(err)
@ -263,11 +266,11 @@ func TestSizedReadWrite(t *testing.T) {
}
{
size := 2 * k64
_, err := stream.SizedWrite(bytes.NewReader(make([]byte, size)), size)
_, err := stream.SizedWrite(testCtx, bytes.NewReader(make([]byte, size)), size)
if err != nil {
t.Fatal(err)
}
rc := stream.SizedReader(size)
rc := stream.SizedReader(testCtx, size)
n, err := rc.Read(make([]byte, k64))
if err != nil {
t.Fatal(err)
@ -303,11 +306,11 @@ func TestSizedReadWrite(t *testing.T) {
}
{
size := k64 / 2
_, err := stream.SizedWrite(bytes.NewReader(make([]byte, size)), size)
_, err := stream.SizedWrite(testCtx, bytes.NewReader(make([]byte, size)), size)
if err != nil {
t.Fatal(err)
}
rc := stream.SizedReader(size - 1)
rc := stream.SizedReader(testCtx, size-1)
w := &sizedWriter{}
_, err = rc.WriteTo(w)
if err != ErrFrameOdd {
@ -317,11 +320,11 @@ func TestSizedReadWrite(t *testing.T) {
}
{
size := k64 / 2
_, err := stream.SizedWrite(bytes.NewReader(make([]byte, size)), size)
_, err := stream.SizedWrite(testCtx, bytes.NewReader(make([]byte, size)), size)
if err != nil {
t.Fatal(err)
}
rc := stream.SizedReader(size + 1)
rc := stream.SizedReader(testCtx, size+1)
stream.SetDeadline(time.Now().Add(time.Second))
w := &sizedWriter{}
_, err = rc.WriteTo(w)
@ -335,6 +338,32 @@ func TestSizedReadWrite(t *testing.T) {
session.Close()
}
func TestSizedReadWriteContext(t *testing.T) {
_, stop, cli, err := setupServer(t)
if err != nil {
t.Fatal(err)
}
defer stop()
session, _ := Client(cli, nil)
stream, _ := session.OpenStream()
msg := make([]byte, 8)
ctx, cancel := context.WithCancel(testCtx)
cancel()
if _, err = stream.SizedWrite(ctx, bytes.NewReader(msg), len(msg)); err == nil {
t.Fatal("write canceled context")
}
if _, err = stream.ReadFrame(ctx); err == nil {
t.Fatal("read canceled context")
}
ctx, cancel = context.WithTimeout(testCtx, 200*time.Millisecond)
if _, err = stream.ReadFrame(ctx); err == nil {
t.Fatal("read canceled context")
}
cancel()
session.Close()
}
func TestParallel(t *testing.T) {
_, stop, cli, err := setupServer(t)
if err != nil {
@ -505,7 +534,7 @@ func TestTinyReadBuffer(t *testing.T) {
t.Fatal("cannot write")
}
nrecv := 0
fr, err := stream.ReadFrame()
fr, err := stream.ReadFrame(testCtx)
if err != nil {
t.Fatal(err)
}
@ -654,7 +683,7 @@ func TestServerEcho(t *testing.T) {
fw.Write([]byte(msg))
stream.WriteFrame(fw)
fw.Close()
fr, err := stream.ReadFrame()
fr, err := stream.ReadFrame(testCtx)
if err != nil {
return err
}
@ -681,7 +710,7 @@ func TestServerEcho(t *testing.T) {
if session, errx := Client(cli, nil); errx == nil {
if s, erry := session.AcceptStream(); erry == nil {
for {
fr, err := s.ReadFrame()
fr, err := s.ReadFrame(testCtx)
if err != nil {
break
}
@ -748,7 +777,7 @@ func TestReadStreamAfterSessionClose(t *testing.T) {
session, _ := Client(cli, nil)
stream, _ := session.OpenStream()
session.Close()
if _, err := stream.ReadFrame(); err != nil {
if _, err := stream.ReadFrame(testCtx); err != nil {
t.Log(err)
} else {
t.Fatal("read stream after session close succeeded")
@ -991,7 +1020,7 @@ func TestReadDeadline(t *testing.T) {
var readErr error
for i := 0; i < N; i++ {
stream.SetReadDeadline(time.Now().Add(-1 * time.Minute))
if _, readErr = stream.ReadFrame(); readErr != nil {
if _, readErr = stream.ReadFrame(testCtx); readErr != nil {
break
}
}
@ -1047,7 +1076,7 @@ type streamRW struct {
}
func (s streamRW) Read(p []byte) (n int, err error) {
f, _ := s.s.ReadFrame()
f, _ := s.s.ReadFrame(testCtx)
n, err = f.Read(p)
f.Close()
return

View File

@ -1,6 +1,7 @@
package transport
import (
"context"
"encoding/binary"
"io"
"net"
@ -78,16 +79,22 @@ func (s *Stream) AllocFrame(size int) (*FrameWrite, error) {
// SizedReader the size must be some full frames,
// should close it whatever happens.
func (s *Stream) SizedReader(size int) *SizedReader {
return s.NewSizedReader(size, nil)
func (s *Stream) SizedReader(ctx context.Context, size int) *SizedReader {
return s.NewSizedReader(ctx, size, nil)
}
// ReadFrame returns frame data, closed by caller
func (s *Stream) ReadFrame() (*FrameRead, error) {
func (s *Stream) ReadFrame(ctx context.Context) (*FrameRead, error) {
for {
select {
case <-ctx.Done():
return nil, ctx.Err()
default:
}
f, err := s.tryReadFrame()
if err == ErrWouldBlock {
if ew := s.waitRead(); ew != nil {
if ew := s.waitRead(ctx); ew != nil {
return nil, ew
}
} else {
@ -176,7 +183,7 @@ func (s *Stream) sendWindowUpdate(consumed uint32) error {
return err
}
func (s *Stream) waitRead() error {
func (s *Stream) waitRead(ctx context.Context) error {
var timer *time.Timer
var deadline <-chan time.Time
if d, ok := s.readDeadline.Load().(time.Time); ok && !d.IsZero() {
@ -200,10 +207,12 @@ func (s *Stream) waitRead() error {
return ErrTimeout
case <-s.die:
return io.ErrClosedPipe
case <-ctx.Done():
return ctx.Err()
}
}
func (s *Stream) SizedWrite(r io.Reader, size int) (n int, err error) {
func (s *Stream) SizedWrite(ctx context.Context, r io.Reader, size int) (n int, err error) {
maxPayloadSize := s.MaxPayloadSize()
var nn int
var fw *FrameWrite
@ -223,6 +232,7 @@ func (s *Stream) SizedWrite(r io.Reader, size int) (n int, err error) {
return
}
fw.WithContext(ctx)
nn, err = s.WriteFrame(fw)
if err != nil {
fw.Close()
@ -428,6 +438,8 @@ func (s *Stream) fin() {
}
type SizedReader struct {
ctx context.Context
n int
s *Stream
f *FrameRead
@ -438,8 +450,8 @@ type SizedReader struct {
err error
}
func (s *Stream) NewSizedReader(size int, f *FrameRead) *SizedReader {
return &SizedReader{n: size, s: s, f: f}
func (s *Stream) NewSizedReader(ctx context.Context, size int, f *FrameRead) *SizedReader {
return &SizedReader{ctx: ctx, n: size, s: s, f: f}
}
func (r *SizedReader) tryNextFrame() error {
@ -460,7 +472,7 @@ func (r *SizedReader) tryNextFrame() error {
if r.f != nil && r.f.Len() == 0 {
r.f.Close()
}
r.f, r.err = r.s.ReadFrame()
r.f, r.err = r.s.ReadFrame(r.ctx)
if r.err != nil {
return r.err
}