geerpc day2, implement a rpc client

This commit is contained in:
gzdaijie
2020-10-01 15:59:49 +08:00
parent 78521ef688
commit 00e37b58da
8 changed files with 601 additions and 53 deletions
+11 -13
View File
@@ -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))
}
+66 -40
View File
@@ -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) }
+219
View File
@@ -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 "<service>.<method>"
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)
}
+39
View File
@@ -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
}
+57
View File
@@ -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()
}
+3
View File
@@ -0,0 +1,3 @@
module geerpc
go 1.13
+36
View File
@@ -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)
}
}
+170
View File
@@ -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) }