feat(rpc2): router add middlewares

. #22358085

Signed-off-by: slasher <shenjie1@oppo.com>
This commit is contained in:
slasher 2024-08-09 13:05:36 +08:00
parent bcec11314d
commit d3e06e797a
7 changed files with 60 additions and 36 deletions

View File

@ -87,7 +87,7 @@ func (r *bodyAndTrailer) Close() (err error) {
err = r.sr.Close()
if r.req != nil {
if err == nil && r.sr.Finished() {
err = r.req.cli.Connector.Put(r.req.Context(), r.req.conn)
err = r.req.client.Connector.Put(r.req.Context(), r.req.conn)
} else {
r.req.conn.Close()
}

View File

@ -87,7 +87,7 @@ func (c *Client) do(req *Request, ret Unmarshaler) (*Response, error) {
if err != nil {
return nil, err
}
req.cli = c
req.client = c
req.conn = conn
resp, err := req.request(c.requestDeadline(req.Context()))

View File

@ -13,10 +13,21 @@ import (
var handler = &rpc2.Router{}
func init() {
handler.Middleware(handleMiddleware1, handleMiddleware2)
handler.Register("/ping", handlePing)
handler.Register("/stream", handleStream)
}
func handleMiddleware1(w rpc2.ResponseWriter, req *rpc2.Request) error {
log.Info("middleware-1")
return nil
}
func handleMiddleware2(w rpc2.ResponseWriter, req *rpc2.Request) error {
log.Info("middleware-2")
return nil
}
func handlePing(w rpc2.ResponseWriter, req *rpc2.Request) error {
log.Info(req.RequestHeader.ToString())
var para pingPara

View File

@ -58,9 +58,9 @@ type OptionRequest func(*Request)
type Request struct {
RequestHeader
ctx context.Context
cli *Client
conn *transport.Stream
ctx context.Context
client *Client // client side
conn *transport.Stream
stream *serverStream
@ -69,7 +69,7 @@ type Request struct {
Body Body
GetBody func() (io.ReadCloser, error) // client side
// fill trailer header
// fill trailer
AfterBody func() error
}

View File

@ -36,19 +36,27 @@ var defaultPanicHandler = func(_ ResponseWriter, req *Request, err interface{},
type Router struct {
PanicHandler func(w ResponseWriter, req *Request, err interface{}, stack []byte) error
maps map[string]Handle
middlewares []Handle
handlers map[string]Handle
}
var _ Handler = (*Router)(nil)
func (r *Router) Register(handler string, handle Handle) {
if r.maps == nil {
r.maps = make(map[string]Handle)
func (r *Router) Middleware(mws ...Handle) {
if len(r.middlewares)+len(mws) > 1<<10 {
panic("rpc2: too much middlewares (>1024)")
}
if _, exist := r.maps[handler]; exist {
r.middlewares = append(r.middlewares, mws...)
}
func (r *Router) Register(handler string, handle Handle) {
if r.handlers == nil {
r.handlers = make(map[string]Handle)
}
if _, exist := r.handlers[handler]; exist {
panic(fmt.Sprintf("rpc2: handle(%s) has registered", handler))
}
r.maps[handler] = handle
r.handlers[handler] = handle
if r.PanicHandler == nil {
r.PanicHandler = defaultPanicHandler
@ -56,7 +64,7 @@ func (r *Router) Register(handler string, handle Handle) {
}
func (r *Router) Handle(w ResponseWriter, req *Request) (err error) {
handle, exist := r.maps[req.RemoteHandler]
handle, exist := r.handlers[req.RemoteHandler]
if !exist {
err = &Error{
Status: 404,
@ -75,6 +83,11 @@ func (r *Router) Handle(w ResponseWriter, req *Request) (err error) {
}
}()
for idx := range r.middlewares {
if err = r.middlewares[idx](w, req); err != nil {
return
}
}
err = handle(w, req)
return
}

View File

@ -41,6 +41,26 @@ var (
ErrFrameProtocol = errors.New("rpc2: invalid protocol frame")
)
type (
Marshaler interface {
Marshal() ([]byte, error)
}
Unmarshaler interface {
Unmarshal([]byte) error
}
Codec interface {
Marshaler
Unmarshaler
}
)
type noneCodec struct{}
func (*noneCodec) Marshal() ([]byte, error) { return nil, nil }
func (*noneCodec) Unmarshal([]byte) error { return nil }
var _ Codec = (*noneCodec)(nil)
type Body interface {
io.Reader
io.WriterTo

View File

@ -19,19 +19,6 @@ import (
"io"
)
type Marshaler interface {
Marshal() ([]byte, error)
}
type Unmarshaler interface {
Unmarshal([]byte) error
}
type Codec interface {
Marshaler
Unmarshaler
}
type ClientStream interface {
Context() context.Context
@ -97,13 +84,6 @@ type StreamingServer[Req any, Res any] interface {
ServerStream
}
type noneCodec struct{}
func (*noneCodec) Marshal() ([]byte, error) { return nil, nil }
func (*noneCodec) Unmarshal([]byte) error { return nil }
var _ Codec = (*noneCodec)(nil)
type GenericClientStream[Req any, Res any] struct {
ClientStream
}
@ -164,7 +144,7 @@ func (cs *clientStream) CloseSend() error {
req := cs.newRequest()
req.StreamCmd = StreamCmd_FIN
req.Trailer = cs.trailer.ToFixedHeader()
return req.write(req.cli.requestDeadline(req.ctx))
return req.write(req.client.requestDeadline(req.ctx))
}
func (cs *clientStream) SendMsg(a any) error {
@ -181,7 +161,7 @@ func (cs *clientStream) SendMsg(a any) error {
req.StreamCmd = StreamCmd_PSH
req.Parameter = b
req.Body = NoBody
return req.write(req.cli.requestDeadline(req.ctx))
return req.write(req.client.requestDeadline(req.ctx))
}
func (cs *clientStream) RecvMsg(a any) (err error) {
@ -222,7 +202,7 @@ func (cs *clientStream) RecvMsg(a any) (err error) {
func (cs *clientStream) newRequest() *Request {
req := baseRequest()
req.ctx = cs.req.ctx
req.cli = cs.req.cli
req.client = cs.req.client
req.conn = cs.req.conn
return req
}