diff --git a/blobstore/common/rpc2/connector.go b/blobstore/common/rpc2/connector.go index 461241185..dc6712acd 100644 --- a/blobstore/common/rpc2/connector.go +++ b/blobstore/common/rpc2/connector.go @@ -20,6 +20,7 @@ import ( "sync" "github.com/cubefs/cubefs/blobstore/common/rpc2/transport" + "github.com/cubefs/cubefs/blobstore/util/defaulter" ) type Dialer interface { @@ -29,40 +30,118 @@ type Dialer interface { type Connector interface { Get(ctx context.Context, addr string) (*transport.Stream, error) Put(ctx context.Context, stream *transport.Stream) error -} - -type sessionStreams struct { - session *transport.Session - streamC chan *transport.Stream + Close() } type connector struct { dialer Dialer - config *transport.Config + config ConnectorConfig mu sync.RWMutex - sessions map[string]sessionStreams + sessions map[string]map[*transport.Session]struct{} // remote address + streams map[net.Addr]chan *transport.Stream // local address } -func DefaultConnector(dialer Dialer, config *transport.Config) Connector { +type ConnectorConfig struct { + Transport *transport.Config + + MaxStreamPerSession int +} + +func DefaultConnector(dialer Dialer, config ConnectorConfig) Connector { + defaulter.LessOrEqual(&config.MaxStreamPerSession, int(1024)) return &connector{ dialer: dialer, config: config, - sessions: make(map[string]sessionStreams), + sessions: make(map[string]map[*transport.Session]struct{}), + streams: make(map[net.Addr]chan *transport.Stream), } } func (c *connector) Get(ctx context.Context, addr string) (*transport.Stream, error) { c.mu.RLock() - ss, ok := c.sessions[addr] + ses, ok := c.sessions[addr] c.mu.RUnlock() - if ok { - if ss.session.IsClosed() { + if !ok || len(ses) == 0 { + conn, err := c.dialer.Dial(ctx, addr) + if err != nil { + return nil, err } + sess, err := transport.Client(conn, c.config.Transport) + if err != nil { + conn.Close() + return nil, err + } + stream, err := sess.OpenStream() + if err != nil { + sess.Close() + return nil, err + } + + c.mu.Lock() + c.sessions[addr] = map[*transport.Session]struct{}{sess: {}} + c.streams[sess.LocalAddr()] = make(chan *transport.Stream, c.config.MaxStreamPerSession) + c.mu.Unlock() + return stream, nil } - return nil, nil + + // try to get opened stream + var stream *transport.Stream + c.mu.RLock() + sesCopy := make(map[*transport.Session]struct{}, len(ses)) + for sess := range ses { + select { + case stream = <-c.streams[sess.LocalAddr()]: + default: + } + if stream != nil { + break + } + sesCopy[sess] = struct{}{} + } + c.mu.RUnlock() + if stream != nil { + return stream, nil + } + + // try to open new stream + for sess := range sesCopy { + newStream, err := sess.OpenStream() + if err != nil { + c.mu.Lock() + delete(c.sessions[addr], sess) + delete(c.streams, sess.LocalAddr()) + c.mu.Unlock() + sess.Close() + continue + } + return newStream, nil + } + + return c.Get(ctx, addr) } func (c *connector) Put(ctx context.Context, stream *transport.Stream) error { + c.mu.RLock() + ch, ok := c.streams[stream.LocalAddr()] + c.mu.RUnlock() + if ok { + select { + case ch <- stream: + default: + } + } return nil } + +func (c *connector) Close() { + c.mu.Lock() + for _, sesss := range c.sessions { + for sess := range sesss { + sess.Close() + } + } + c.sessions = make(map[string]map[*transport.Session]struct{}) + c.streams = make(map[net.Addr]chan *transport.Stream) + c.mu.Unlock() +} diff --git a/blobstore/common/rpc2/example/client.go b/blobstore/common/rpc2/example/client.go index 6e1798e1f..20ba1b9c9 100644 --- a/blobstore/common/rpc2/example/client.go +++ b/blobstore/common/rpc2/example/client.go @@ -9,16 +9,18 @@ import ( "time" "github.com/cubefs/cubefs/blobstore/common/rpc2" - "github.com/cubefs/cubefs/blobstore/common/rpc2/transport" "github.com/cubefs/cubefs/blobstore/util/log" ) -type connector struct { - stream *transport.Stream +type dialer struct{} + +func (dialer) Dial(context.Context, string) (net.Conn, error) { + return net.Dial("tcp", listenon[int(time.Now().UnixNano())%len(listenon)]) } -func (c *connector) Get(context.Context, string) (*transport.Stream, error) { return c.stream, nil } -func (c *connector) Put(context.Context, *transport.Stream) error { return nil } +func makeConnector() rpc2.Connector { + return rpc2.DefaultConnector(dialer{}, rpc2.ConnectorConfig{}) +} type pingPara struct { I int @@ -36,25 +38,13 @@ func (p *pingPara) Unmarshal(data []byte) error { } func runClient() { - conn, err := net.Dial("tcp", listenon[int(time.Now().UnixNano())%len(listenon)]) - if err != nil { - panic(err) - } - session, err := transport.Client(conn, nil) - if err != nil { - panic(err) - } - stream, err := session.OpenStream() - if err != nil { - panic(err) - } - cli := rpc2.Client{Connector: &connector{stream: stream}} + client := rpc2.Client{Connector: makeConnector()} - para := pingPara{ I: 7, S: "ping string", } + para := pingPara{I: 7, S: "ping string"} buff := []byte("ping") req, _ := rpc2.NewRequest(context.Background(), "", "/ping", ¶, bytes.NewReader(buff)) var ret pingPara - resp, err := cli.Do(req, &ret) + resp, err := client.Do(req, &ret) if err != nil { panic(err) } @@ -65,5 +55,5 @@ func runClient() { log.Infof("got para: %+v | %v", ret, err) log.Info("got body:", string(got), err) - stream.Close() + client.Connector.Close() } diff --git a/blobstore/common/rpc2/example/server.go b/blobstore/common/rpc2/example/server.go index f6ed6e0fd..ff2ebe4e5 100644 --- a/blobstore/common/rpc2/example/server.go +++ b/blobstore/common/rpc2/example/server.go @@ -10,20 +10,15 @@ import ( "github.com/cubefs/cubefs/blobstore/util/log" ) -type handler struct{} +var handler = &rpc2.Router{} -func (h *handler) Handle(w rpc2.ResponseWriter, req *rpc2.Request) error { - switch req.RequestHeader.RemoteHandler { - case "/ping": - return handlePing(w, req) - case "/stream": - return handleStream(w, req) - } - return nil +func init() { + handler.Register("/ping", handlePing) + handler.Register("/stream", handleStream) } func handlePing(w rpc2.ResponseWriter, req *rpc2.Request) error { - log.Infof("%+v", req.RequestHeader) + log.Info(req.RequestHeader.ToString()) var para pingPara para.Unmarshal(req.GetParameter()) w.SetContentLength(req.ContentLength) @@ -77,7 +72,7 @@ func runServer() { server := rpc2.Server{ Name: ln1.Addr().String() + " | " + ln2.Addr().String(), - Handler: &handler{}, + Handler: handler, StatDuration: 3 * time.Second, } go func() { diff --git a/blobstore/common/rpc2/example/stream.go b/blobstore/common/rpc2/example/stream.go index 9a65e9d8f..89afa9854 100644 --- a/blobstore/common/rpc2/example/stream.go +++ b/blobstore/common/rpc2/example/stream.go @@ -4,11 +4,8 @@ import ( "context" "fmt" "io" - "net" - "time" "github.com/cubefs/cubefs/blobstore/common/rpc2" - "github.com/cubefs/cubefs/blobstore/common/rpc2/transport" "github.com/cubefs/cubefs/blobstore/util/log" ) @@ -28,19 +25,7 @@ var ( ) func runStream() { - conn, err := net.Dial("tcp", listenon[int(time.Now().UnixNano())%len(listenon)]) - if err != nil { - panic(err) - } - session, err := transport.Client(conn, nil) - if err != nil { - panic(err) - } - stream, err := session.OpenStream() - if err != nil { - panic(err) - } - client := rpc2.Client{Connector: &connector{stream: stream}} + client := rpc2.Client{Connector: makeConnector()} streamCli := rpc2.StreamClient[streamReq, streamResp]{Client: &client} ctx := context.Background() diff --git a/blobstore/common/rpc2/request.go b/blobstore/common/rpc2/request.go index 2b4718462..aa7463ddf 100644 --- a/blobstore/common/rpc2/request.go +++ b/blobstore/common/rpc2/request.go @@ -16,6 +16,7 @@ package rpc2 import ( "context" + "fmt" "io" "sync" "time" @@ -42,6 +43,15 @@ func (m *RequestHeader) MarshalToReader() io.Reader { return &headerReader{marshaler: m} } +func (m *RequestHeader) ToString() string { + return fmt.Sprintf("Version:%d Magic:%d"+ + " StreamCmd:%s RemoteAddr:%s RemoteHandler:%s TraceID:%s"+ + " ContentLength:%d Header:%+v Trailer:%+v Parameter:len(%d)", + m.Version, m.Magic, + m.StreamCmd.String(), m.RemoteAddr, m.RemoteHandler, m.TraceID, + m.ContentLength, m.Header.M, m.Trailer.M, len(m.Parameter)) +} + type Request struct { RequestHeader diff --git a/blobstore/common/rpc2/response.go b/blobstore/common/rpc2/response.go index c6711f787..3330e4e63 100644 --- a/blobstore/common/rpc2/response.go +++ b/blobstore/common/rpc2/response.go @@ -16,6 +16,7 @@ package rpc2 import ( "bytes" + "fmt" "io" "github.com/cubefs/cubefs/blobstore/common/rpc2/transport" @@ -39,6 +40,13 @@ func (m *ResponseHeader) MarshalToReader() io.Reader { return &headerReader{marshaler: m} } +func (m *ResponseHeader) ToString() string { + return fmt.Sprintf("Version:%d Magic:%d Status:%d Reason:%s Error:%s"+ + " ContentLength:%d Header:%+v Trailer:%+v Parameter:len(%d)", + m.Version, m.Magic, m.Status, m.Reason, m.Error, + m.ContentLength, m.Header.M, m.Trailer.M, len(m.Parameter)) +} + // client side response type Response struct { ResponseHeader diff --git a/blobstore/common/rpc2/router.go b/blobstore/common/rpc2/router.go new file mode 100644 index 000000000..0f7eb55ef --- /dev/null +++ b/blobstore/common/rpc2/router.go @@ -0,0 +1,45 @@ +// Copyright 2024 The CubeFS Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or +// implied. See the License for the specific language governing +// permissions and limitations under the License. + +package rpc2 + +import "fmt" + +type Router struct { + maps 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) + } + if _, exist := r.maps[handler]; exist { + panic(fmt.Sprintf("rpc2: handle(%s) has registered", handler)) + } + r.maps[handler] = handle +} + +func (r *Router) Handle(w ResponseWriter, req *Request) error { + handle, exist := r.maps[req.RemoteHandler] + if !exist { + return &Error{ + Status: 404, + Reason: "NoRouter", + Detail: fmt.Sprintf("no router for handler(%s)", req.RemoteHandler), + } + } + return handle(w, req) +} diff --git a/blobstore/common/rpc2/server.go b/blobstore/common/rpc2/server.go index 3bb57584e..124e63586 100644 --- a/blobstore/common/rpc2/server.go +++ b/blobstore/common/rpc2/server.go @@ -27,6 +27,8 @@ import ( "github.com/cubefs/cubefs/blobstore/util/log" ) +type Handle func(ResponseWriter, *Request) error + type Handler interface { Handle(ResponseWriter, *Request) error }