mirror of
https://github.com/cubefs/cubefs.git
synced 2026-08-02 02:00:56 +00:00
feat(rpc2): transport with context cancel
. #22548427 Signed-off-by: slasher <shenjie1@oppo.com>
This commit is contained in:
parent
e41f7a497b
commit
946ec8516f
@ -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
|
||||
}
|
||||
|
||||
@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@ -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
|
||||
}
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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 {
|
||||
|
||||
@ -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
|
||||
}
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
}
|
||||
|
||||
Loading…
Reference in New Issue
Block a user