From 00e37b58da400f3c8007be550640daa4fc95dfdb Mon Sep 17 00:00:00 2001 From: gzdaijie Date: Thu, 1 Oct 2020 15:59:49 +0800 Subject: [PATCH] geerpc day2, implement a rpc client --- gee-rpc/day1-codec/main/main.go | 24 ++-- gee-rpc/day1-codec/server.go | 106 ++++++++------ gee-rpc/day2-client/client.go | 219 +++++++++++++++++++++++++++++ gee-rpc/day2-client/codec/codec.go | 39 +++++ gee-rpc/day2-client/codec/gob.go | 57 ++++++++ gee-rpc/day2-client/go.mod | 3 + gee-rpc/day2-client/main/main.go | 36 +++++ gee-rpc/day2-client/server.go | 170 ++++++++++++++++++++++ 8 files changed, 601 insertions(+), 53 deletions(-) create mode 100644 gee-rpc/day2-client/client.go create mode 100644 gee-rpc/day2-client/codec/codec.go create mode 100644 gee-rpc/day2-client/codec/gob.go create mode 100644 gee-rpc/day2-client/go.mod create mode 100644 gee-rpc/day2-client/main/main.go create mode 100644 gee-rpc/day2-client/server.go diff --git a/gee-rpc/day1-codec/main/main.go b/gee-rpc/day1-codec/main/main.go index 96f2273..95dd9e2 100644 --- a/gee-rpc/day1-codec/main/main.go +++ b/gee-rpc/day1-codec/main/main.go @@ -7,29 +7,28 @@ import ( "geerpc/codec" "log" "net" - "net/http" - "time" ) -func startServer() { - l, err := net.Listen("tcp", ":9999") +func startServer(addr chan string) { + // pick a free port + l, err := net.Listen("tcp", ":0") if err != nil { - log.Panic("network error:", err) + log.Fatal("network error:", err) } + log.Println("start rpc server on", l.Addr()) + addr <- l.Addr().String() geerpc.Accept(l) - } func main() { - log.SetPrefix("") - go startServer() - time.Sleep(time.Second) + addr := make(chan string) + go startServer(addr) - // In fact, following code is like a simple GeeRPC Client - conn, _ := net.Dial("tcp", ":9999") + // in fact, following code is like a simple geerpc client + conn, _ := net.Dial("tcp", <-addr) defer func() { _ = conn.Close() }() - // negotiate options + // send options _ = json.NewEncoder(conn).Encode(&geerpc.Options{ MagicNumber: geerpc.MagicNumber, CodecType: codec.GobType, @@ -48,5 +47,4 @@ func main() { _ = cc.ReadBody(&reply) log.Println("reply:", reply) } - log.Fatal(http.ListenAndServe(":9999", nil)) } diff --git a/gee-rpc/day1-codec/server.go b/gee-rpc/day1-codec/server.go index 3e32f1f..7646f9d 100644 --- a/gee-rpc/day1-codec/server.go +++ b/gee-rpc/day1-codec/server.go @@ -1,3 +1,7 @@ +// Copyright 2009 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + package geerpc import ( @@ -7,11 +11,10 @@ import ( "io" "log" "net" + "reflect" "sync" ) -type Server struct{} - const MagicNumber = 0x3bef5c type Options struct { @@ -19,12 +22,24 @@ type Options struct { CodecType codec.Type // client may choose different Codec to encode body } -func newServer() *Server { +var defaultOptions = &Options{ + MagicNumber: MagicNumber, + CodecType: codec.GobType, +} + +// Server represents an RPC Server. +type Server struct{} + +// NewServer returns a new Server. +func NewServer() *Server { return &Server{} } -var DefaultServer = newServer() +// DefaultServer is the default instance of *Server. +var DefaultServer = NewServer() +// ServeConn runs the server on a single connection. +// ServeConn blocks, serving the connection until the client hangs up. func (server *Server) ServeConn(conn io.ReadWriteCloser) { defer func() { _ = conn.Close() }() var opt Options @@ -41,13 +56,38 @@ func (server *Server) ServeConn(conn io.ReadWriteCloser) { log.Printf("rpc server: invalid codec type %s", opt.CodecType) return } - server.ServeCodec(f(conn)) + server.serveCodec(f(conn)) +} + +// invalidRequest is a placeholder for response argv when error occurs +var invalidRequest = struct{}{} + +func (server *Server) serveCodec(cc codec.Codec) { + sending := new(sync.Mutex) // ensure header and argv is not separated by other response + wg := new(sync.WaitGroup) // wait until all request are handled + for { + req, keepReading, err := server.readRequest(cc) + if err != nil { + if !keepReading { + break // it's not possible to recover, so close the connection + } + if req != nil { + req.h.Error = err.Error() + server.sendResponse(cc, req.h, invalidRequest, sending) + } + continue + } + wg.Add(1) + go server.handleRequest(cc, req, sending, wg) + } + wg.Wait() + _ = cc.Close() } // request stores all information of a call type request struct { - h *codec.Header // header of request argv - argv string // TODO suppose argv is a string + h *codec.Header // header of request + argv, replyv reflect.Value // argv and replyv of request } var requestPool = sync.Pool{ @@ -58,6 +98,10 @@ func (server *Server) readRequestHeader(cc codec.Codec) (req *request, keepReadi req, _ = requestPool.Get().(*request) h, _ := codec.HeaderPool.Get().(*codec.Header) if err = cc.ReadHeader(h); err != nil { + // client closed the connection + if err == io.EOF || err != io.ErrUnexpectedEOF { + return + } log.Println("rpc server: read header error:", err) return } @@ -80,12 +124,12 @@ func (server *Server) readRequest(cc codec.Codec) (req *request, keepReading boo // we can still recover and move on to the next request. keepReading = true - // TODO: suppose argv is a string, now we can't judge the type of request argv - var str string - if err = cc.ReadBody(&str); err != nil { + // TODO: now we can't judge the type of request argv + // day 1, just suppose it's string + req.argv = reflect.New(reflect.TypeOf("")) + if err = cc.ReadBody(req.argv.Interface()); err != nil { log.Println("rpc server: read argv err:", err) } - req.argv = str return } @@ -97,39 +141,19 @@ func (server *Server) sendResponse(cc codec.Codec, h *codec.Header, body interfa } } -func (server *Server) Handle(cc codec.Codec, req *request, sending *sync.Mutex, wg *sync.WaitGroup) { - // TODO, should call registered rpc methods - // day 1 just print argv and send a hello message +func (server *Server) handleRequest(cc codec.Codec, req *request, sending *sync.Mutex, wg *sync.WaitGroup) { + // TODO, should call registered rpc methods to get the right replyv + // day 1, just print argv and send a hello message defer wg.Done() - log.Println(req.h, req.argv) - server.sendResponse(cc, req.h, fmt.Sprintf("geerpc resp %d", req.h.Seq), sending) -} + defer codec.HeaderPool.Put(req.h) // recycle Header object + log.Println(req.h, req.argv.Elem()) + req.replyv = reflect.ValueOf(fmt.Sprintf("geerpc resp %d", req.h.Seq)) + server.sendResponse(cc, req.h, req.replyv.Interface(), sending) -// invalidRequest is a placeholder for response argv when error occurs -var invalidRequest = struct{}{} - -func (server *Server) ServeCodec(cc codec.Codec) { - sending := new(sync.Mutex) // ensure header and argv is not separated by other response - wg := new(sync.WaitGroup) // wait until all request are handled - for { - req, keepReading, err := server.readRequest(cc) - if err != nil { - if !keepReading { - break // it's not possible to recover, so close the connection - } - if req != nil { - req.h.Error = err.Error() - server.sendResponse(cc, req.h, invalidRequest, sending) - } - continue - } - wg.Add(1) - go server.Handle(cc, req, sending, wg) - } - wg.Wait() - _ = cc.Close() } +// Accept accepts connections on the listener and serves requests +// for each incoming connection. func (server *Server) Accept(lis net.Listener) { for { conn, err := lis.Accept() @@ -141,4 +165,6 @@ func (server *Server) Accept(lis net.Listener) { } } +// Accept accepts connections on the listener and serves requests +// for each incoming connection. func Accept(lis net.Listener) { DefaultServer.Accept(lis) } diff --git a/gee-rpc/day2-client/client.go b/gee-rpc/day2-client/client.go new file mode 100644 index 0000000..79ef524 --- /dev/null +++ b/gee-rpc/day2-client/client.go @@ -0,0 +1,219 @@ +// Copyright 2009 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +package geerpc + +import ( + "encoding/json" + "errors" + "fmt" + "geerpc/codec" + "io" + "log" + "net" + "sync" +) + +// Call represents an active RPC. +type Call struct { + ServiceMethod string // format "." + Args interface{} // arguments to the function + Reply interface{} // reply from the function + Error error // if error occurs, it will be set + Done chan *Call // Strobes when call is complete. +} + +func (call *Call) done() { + call.Done <- call +} + +// Client represents an RPC Client. +// There may be multiple outstanding Calls associated +// with a single Client, and a Client may be used by +// multiple goroutines simultaneously. +type Client struct { + cc codec.Codec + sending sync.Mutex // protect sending a complete request + mu sync.Mutex // protect following + seq uint64 + pending map[uint64]*Call + closed bool // user has called Close +} + +var _ io.Closer = (*Client)(nil) + +var ErrShutdown = errors.New("connection is shut down") + +// Close the connection +func (client *Client) Close() error { + client.mu.Lock() + defer client.mu.Unlock() + if client.closed { + return ErrShutdown + } + client.closed = true + return client.cc.Close() +} + +func (client *Client) registerCall(call *Call) (uint64, error) { + client.mu.Lock() + defer client.mu.Unlock() + if client.closed { + return 0, ErrShutdown + } + seq := client.seq + client.pending[seq] = call + client.seq++ + return seq, nil +} + +func (client *Client) removeCall(seq uint64) *Call { + client.mu.Lock() + defer client.mu.Unlock() + call := client.pending[seq] + delete(client.pending, seq) + return call +} + +func (client *Client) terminateCalls(err error) { + client.sending.Lock() + defer client.sending.Unlock() + client.mu.Lock() + defer client.mu.Unlock() + for _, call := range client.pending { + call.Error = err + call.done() + } +} + +func (client *Client) send(call *Call) { + // make sure that the client will send a complete request + client.sending.Lock() + defer client.sending.Unlock() + + // register this call. + seq, err := client.registerCall(call) + if err != nil { + call.Error = err + call.done() + return + } + + // prepare request header + h, _ := codec.HeaderPool.Get().(*codec.Header) + h.ServiceMethod = call.ServiceMethod + h.Seq = seq + h.Error = "" + defer codec.HeaderPool.Put(h) + + // encode and send the request + if err := client.cc.Write(h, call.Args); err != nil { + call := client.removeCall(seq) + // call may be nil, it usually means that Write partially failed, + // client has received the response and handled + if call != nil { + call.Error = err + call.done() + } + } +} + +func (client *Client) receive() { + h, _ := codec.HeaderPool.Get().(*codec.Header) + defer codec.HeaderPool.Put(h) + var err error + for err == nil { + if err = client.cc.ReadHeader(h); err != nil { + break + } + call := client.removeCall(h.Seq) + switch { + case call == nil: + // it usually means that Write partially failed + // and call was already removed. + err = client.cc.ReadBody(nil) + case h.Error != "": + call.Error = fmt.Errorf(h.Error) + err = client.cc.ReadBody(nil) + call.done() + default: + err = client.cc.ReadBody(call.Reply) + call.done() + } + } + // error occurs, so terminateCalls pending calls + client.terminateCalls(err) +} + +// Go invokes the function asynchronously. +// It returns the Call structure representing the invocation. +func (client *Client) Go(serviceMethod string, args, reply interface{}, done chan *Call) *Call { + if done == nil { + done = make(chan *Call, 10) + } else if cap(done) == 0 { + log.Panic("rpc client: done channel is unbuffered") + } + call := &Call{ + ServiceMethod: serviceMethod, + Args: args, + Reply: reply, + Done: done, + } + client.send(call) + return call +} + +// Call invokes the named function, waits for it to complete, +// and returns its error status. +func (client *Client) Call(serviceMethod string, args, reply interface{}) error { + call := <-client.Go(serviceMethod, args, reply, make(chan *Call, 1)).Done + return call.Error +} + +func NewClient(conn io.ReadWriteCloser, opt *Options) (*Client, error) { + var err error + defer func() { + if err != nil { + _ = conn.Close() + } + }() + if opt.MagicNumber == 0 { + opt.MagicNumber = MagicNumber + } + f := codec.NewCodecFuncMap[opt.CodecType] + if f == nil { + err = fmt.Errorf("invalid codec type %s", opt.CodecType) + log.Println("rpc client: codec error:", err) + return nil, err + } + // send options with server + if err = json.NewEncoder(conn).Encode(opt); err != nil { + log.Println("rpc client: options error: ", err) + return nil, err + } + return newClientCodec(f(conn)), nil +} + +func newClientCodec(cc codec.Codec) *Client { + client := &Client{ + cc: cc, + pending: make(map[uint64]*Call), + } + go client.receive() + return client +} + +// DialWithOptions connects to an RPC server at the specified network address +func DialWithOptions(network, address string, opt *Options) (*Client, error) { + conn, err := net.Dial(network, address) + if err != nil { + return nil, err + } + return NewClient(conn, opt) +} + +// Dial connects to an RPC server at the specified network address +func Dial(network, address string) (*Client, error) { + return DialWithOptions(network, address, defaultOptions) +} diff --git a/gee-rpc/day2-client/codec/codec.go b/gee-rpc/day2-client/codec/codec.go new file mode 100644 index 0000000..54d2e56 --- /dev/null +++ b/gee-rpc/day2-client/codec/codec.go @@ -0,0 +1,39 @@ +package codec + +import ( + "io" + "sync" +) + +type Header struct { + ServiceMethod string // format "Service.Method" + Seq uint64 // sequence number chosen by client + Error string +} + +var HeaderPool = sync.Pool{ + New: func() interface{} { return &Header{} }, +} + +type Codec interface { + io.Closer + ReadHeader(*Header) error + ReadBody(interface{}) error + Write(*Header, interface{}) error +} + +type NewCodecFunc func(io.ReadWriteCloser) Codec + +type Type string + +const ( + GobType Type = "application/gob" + JsonType Type = "application/json" +) + +var NewCodecFuncMap map[Type]NewCodecFunc + +func init() { + NewCodecFuncMap = make(map[Type]NewCodecFunc) + NewCodecFuncMap[GobType] = NewGobCodec +} diff --git a/gee-rpc/day2-client/codec/gob.go b/gee-rpc/day2-client/codec/gob.go new file mode 100644 index 0000000..808d97b --- /dev/null +++ b/gee-rpc/day2-client/codec/gob.go @@ -0,0 +1,57 @@ +package codec + +import ( + "bufio" + "encoding/gob" + "io" + "log" +) + +type GobCodec struct { + conn io.ReadWriteCloser + buf *bufio.Writer + dec *gob.Decoder + enc *gob.Encoder +} + +var _ Codec = (*GobCodec)(nil) + +func NewGobCodec(conn io.ReadWriteCloser) Codec { + buf := bufio.NewWriter(conn) + return &GobCodec{ + conn: conn, + buf: buf, + dec: gob.NewDecoder(conn), + enc: gob.NewEncoder(buf), + } +} + +func (c *GobCodec) ReadHeader(h *Header) error { + return c.dec.Decode(h) +} + +func (c *GobCodec) ReadBody(body interface{}) error { + return c.dec.Decode(body) +} + +func (c *GobCodec) Write(h *Header, body interface{}) (err error) { + defer func() { + _ = c.buf.Flush() + if err != nil { + _ = c.Close() + } + }() + if err := c.enc.Encode(h); err != nil { + log.Println("rpc: gob error encoding header:", err) + return err + } + if err := c.enc.Encode(body); err != nil { + log.Println("rpc: gob error encoding body:", err) + return err + } + return nil +} + +func (c *GobCodec) Close() error { + return c.conn.Close() +} diff --git a/gee-rpc/day2-client/go.mod b/gee-rpc/day2-client/go.mod new file mode 100644 index 0000000..0ec8aeb --- /dev/null +++ b/gee-rpc/day2-client/go.mod @@ -0,0 +1,3 @@ +module geerpc + +go 1.13 diff --git a/gee-rpc/day2-client/main/main.go b/gee-rpc/day2-client/main/main.go new file mode 100644 index 0000000..b370e82 --- /dev/null +++ b/gee-rpc/day2-client/main/main.go @@ -0,0 +1,36 @@ +package main + +import ( + "fmt" + "geerpc" + "log" + "net" +) + +func startServer(addr chan string) { + // pick a free port + l, err := net.Listen("tcp", ":0") + if err != nil { + log.Fatal("network error:", err) + } + log.Println("start rpc server on", l.Addr()) + addr <- l.Addr().String() + geerpc.Accept(l) +} + +func main() { + addr := make(chan string) + go startServer(addr) + client, _ := geerpc.Dial("tcp", <-addr) + defer func() { _ = client.Close() }() + + // send request & receive response + for i := 0; i < 3; i++ { + args := fmt.Sprintf("geerpc req %d", i) + var reply string + if err := client.Call("Foo.Sum", args, &reply); err != nil { + log.Fatal("call Foo.Sum error", err) + } + log.Println("reply:", reply) + } +} diff --git a/gee-rpc/day2-client/server.go b/gee-rpc/day2-client/server.go new file mode 100644 index 0000000..7646f9d --- /dev/null +++ b/gee-rpc/day2-client/server.go @@ -0,0 +1,170 @@ +// Copyright 2009 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +package geerpc + +import ( + "encoding/json" + "fmt" + "geerpc/codec" + "io" + "log" + "net" + "reflect" + "sync" +) + +const MagicNumber = 0x3bef5c + +type Options struct { + MagicNumber int // MagicNumber marks this's a geerpc request + CodecType codec.Type // client may choose different Codec to encode body +} + +var defaultOptions = &Options{ + MagicNumber: MagicNumber, + CodecType: codec.GobType, +} + +// Server represents an RPC Server. +type Server struct{} + +// NewServer returns a new Server. +func NewServer() *Server { + return &Server{} +} + +// DefaultServer is the default instance of *Server. +var DefaultServer = NewServer() + +// ServeConn runs the server on a single connection. +// ServeConn blocks, serving the connection until the client hangs up. +func (server *Server) ServeConn(conn io.ReadWriteCloser) { + defer func() { _ = conn.Close() }() + var opt Options + if err := json.NewDecoder(conn).Decode(&opt); err != nil { + log.Println("rpc server: options error: ", err) + return + } + if opt.MagicNumber != MagicNumber { + log.Printf("rpc server: invalid magic number %x", opt.MagicNumber) + return + } + f := codec.NewCodecFuncMap[opt.CodecType] + if f == nil { + log.Printf("rpc server: invalid codec type %s", opt.CodecType) + return + } + server.serveCodec(f(conn)) +} + +// invalidRequest is a placeholder for response argv when error occurs +var invalidRequest = struct{}{} + +func (server *Server) serveCodec(cc codec.Codec) { + sending := new(sync.Mutex) // ensure header and argv is not separated by other response + wg := new(sync.WaitGroup) // wait until all request are handled + for { + req, keepReading, err := server.readRequest(cc) + if err != nil { + if !keepReading { + break // it's not possible to recover, so close the connection + } + if req != nil { + req.h.Error = err.Error() + server.sendResponse(cc, req.h, invalidRequest, sending) + } + continue + } + wg.Add(1) + go server.handleRequest(cc, req, sending, wg) + } + wg.Wait() + _ = cc.Close() +} + +// request stores all information of a call +type request struct { + h *codec.Header // header of request + argv, replyv reflect.Value // argv and replyv of request +} + +var requestPool = sync.Pool{ + New: func() interface{} { return &request{} }, +} + +func (server *Server) readRequestHeader(cc codec.Codec) (req *request, keepReading bool, err error) { + req, _ = requestPool.Get().(*request) + h, _ := codec.HeaderPool.Get().(*codec.Header) + if err = cc.ReadHeader(h); err != nil { + // client closed the connection + if err == io.EOF || err != io.ErrUnexpectedEOF { + return + } + log.Println("rpc server: read header error:", err) + return + } + // We read the header successfully. If we see an error now, + // we can still recover and move on to the next request. + keepReading = true + req.h = h + return +} + +func (server *Server) readRequest(cc codec.Codec) (req *request, keepReading bool, err error) { + req, keepReading, err = server.readRequestHeader(cc) + if err != nil { + // discard argv + _ = cc.ReadBody(nil) + return + } + + // We read the header successfully. If we see an error now, + // we can still recover and move on to the next request. + keepReading = true + + // TODO: now we can't judge the type of request argv + // day 1, just suppose it's string + req.argv = reflect.New(reflect.TypeOf("")) + if err = cc.ReadBody(req.argv.Interface()); err != nil { + log.Println("rpc server: read argv err:", err) + } + return +} + +func (server *Server) sendResponse(cc codec.Codec, h *codec.Header, body interface{}, sending *sync.Mutex) { + sending.Lock() + defer sending.Unlock() + if err := cc.Write(h, body); err != nil { + log.Println("rpc server: write response error:", err) + } +} + +func (server *Server) handleRequest(cc codec.Codec, req *request, sending *sync.Mutex, wg *sync.WaitGroup) { + // TODO, should call registered rpc methods to get the right replyv + // day 1, just print argv and send a hello message + defer wg.Done() + defer codec.HeaderPool.Put(req.h) // recycle Header object + log.Println(req.h, req.argv.Elem()) + req.replyv = reflect.ValueOf(fmt.Sprintf("geerpc resp %d", req.h.Seq)) + server.sendResponse(cc, req.h, req.replyv.Interface(), sending) + +} + +// Accept accepts connections on the listener and serves requests +// for each incoming connection. +func (server *Server) Accept(lis net.Listener) { + for { + conn, err := lis.Accept() + if err != nil { + log.Println("rpc server: accept error:", err) + return + } + go server.ServeConn(conn) + } +} + +// Accept accepts connections on the listener and serves requests +// for each incoming connection. +func Accept(lis net.Listener) { DefaultServer.Accept(lis) }