fix HeaderPool recycle opportunity, add geerpc day3 service

This commit is contained in:
gzdaijie
2020-10-01 18:56:32 +08:00
parent 00e37b58da
commit 2763266fef
15 changed files with 697 additions and 101 deletions
+2 -2
View File
@@ -42,11 +42,11 @@ func (c *GobCodec) Write(h *Header, body interface{}) (err error) {
}
}()
if err := c.enc.Encode(h); err != nil {
log.Println("rpc: gob error encoding header:", err)
log.Println("rpc codec: gob error encoding header:", err)
return err
}
if err := c.enc.Encode(body); err != nil {
log.Println("rpc: gob error encoding body:", err)
log.Println("rpc codec: gob error encoding body:", err)
return err
}
return nil
+1 -1
View File
@@ -36,7 +36,7 @@ func main() {
cc := codec.NewGobCodec(conn)
// send request & receive response
for i := 0; i < 3; i++ {
for i := 0; i < 5; i++ {
h := &codec.Header{
ServiceMethod: "Foo.Sum",
Seq: uint64(i),
+19 -38
View File
@@ -63,18 +63,16 @@ func (server *Server) ServeConn(conn io.ReadWriteCloser) {
var invalidRequest = struct{}{}
func (server *Server) serveCodec(cc codec.Codec) {
sending := new(sync.Mutex) // ensure header and argv is not separated by other response
sending := new(sync.Mutex) // make sure to send a complete response
wg := new(sync.WaitGroup) // wait until all request are handled
for {
req, keepReading, err := server.readRequest(cc)
req, err := server.readRequest(cc)
if err != nil {
if !keepReading {
if req == nil {
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)
}
req.h.Error = err.Error()
server.sendResponse(cc, req.h, invalidRequest, sending)
continue
}
wg.Add(1)
@@ -90,47 +88,31 @@ type request struct {
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)
func (server *Server) readRequestHeader(cc codec.Codec) (*codec.Header, error) {
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
if err := cc.ReadHeader(h); err != nil {
codec.HeaderPool.Put(h)
if err != io.EOF && err != io.ErrUnexpectedEOF {
log.Println("rpc server: read header error:", err)
}
log.Println("rpc server: read header error:", err)
return
return nil, err
}
// 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
return h, nil
}
func (server *Server) readRequest(cc codec.Codec) (req *request, keepReading bool, err error) {
req, keepReading, err = server.readRequestHeader(cc)
func (server *Server) readRequest(cc codec.Codec) (*request, error) {
h, err := server.readRequestHeader(cc)
if err != nil {
// discard argv
_ = cc.ReadBody(nil)
return
return nil, err
}
// 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
req := &request{h: h}
// TODO: now we don't know 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
return req, nil
}
func (server *Server) sendResponse(cc codec.Codec, h *codec.Header, body interface{}, sending *sync.Mutex) {
@@ -139,17 +121,16 @@ func (server *Server) sendResponse(cc codec.Codec, h *codec.Header, body interfa
if err := cc.Write(h, body); err != nil {
log.Println("rpc server: write response error:", err)
}
codec.HeaderPool.Put(h) // recycle Header object
}
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
+5 -3
View File
@@ -120,11 +120,10 @@ func (client *Client) send(call *Call) {
}
func (client *Client) receive() {
h, _ := codec.HeaderPool.Get().(*codec.Header)
defer codec.HeaderPool.Put(h)
var h codec.Header
var err error
for err == nil {
if err = client.cc.ReadHeader(h); err != nil {
if err = client.cc.ReadHeader(&h); err != nil {
break
}
call := client.removeCall(h.Seq)
@@ -139,6 +138,9 @@ func (client *Client) receive() {
call.done()
default:
err = client.cc.ReadBody(call.Reply)
if err != nil {
call.Error = errors.New("reading body " + err.Error())
}
call.done()
}
}
+2 -2
View File
@@ -42,11 +42,11 @@ func (c *GobCodec) Write(h *Header, body interface{}) (err error) {
}
}()
if err := c.enc.Encode(h); err != nil {
log.Println("rpc: gob error encoding header:", err)
log.Println("rpc codec: gob error encoding header:", err)
return err
}
if err := c.enc.Encode(body); err != nil {
log.Println("rpc: gob error encoding body:", err)
log.Println("rpc codec: gob error encoding body:", err)
return err
}
return nil
+14 -7
View File
@@ -5,6 +5,7 @@ import (
"geerpc"
"log"
"net"
"sync"
)
func startServer(addr chan string) {
@@ -25,12 +26,18 @@ func main() {
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)
var wg sync.WaitGroup
for i := 0; i < 5; i++ {
wg.Add(1)
go func(i int) {
defer wg.Done()
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)
}(i)
}
wg.Wait()
}
+19 -38
View File
@@ -63,18 +63,16 @@ func (server *Server) ServeConn(conn io.ReadWriteCloser) {
var invalidRequest = struct{}{}
func (server *Server) serveCodec(cc codec.Codec) {
sending := new(sync.Mutex) // ensure header and argv is not separated by other response
sending := new(sync.Mutex) // make sure to send a complete response
wg := new(sync.WaitGroup) // wait until all request are handled
for {
req, keepReading, err := server.readRequest(cc)
req, err := server.readRequest(cc)
if err != nil {
if !keepReading {
if req == nil {
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)
}
req.h.Error = err.Error()
server.sendResponse(cc, req.h, invalidRequest, sending)
continue
}
wg.Add(1)
@@ -90,47 +88,31 @@ type request struct {
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)
func (server *Server) readRequestHeader(cc codec.Codec) (*codec.Header, error) {
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
if err := cc.ReadHeader(h); err != nil {
codec.HeaderPool.Put(h)
if err != io.EOF && err != io.ErrUnexpectedEOF {
log.Println("rpc server: read header error:", err)
}
log.Println("rpc server: read header error:", err)
return
return nil, err
}
// 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
return h, nil
}
func (server *Server) readRequest(cc codec.Codec) (req *request, keepReading bool, err error) {
req, keepReading, err = server.readRequestHeader(cc)
func (server *Server) readRequest(cc codec.Codec) (*request, error) {
h, err := server.readRequestHeader(cc)
if err != nil {
// discard argv
_ = cc.ReadBody(nil)
return
return nil, err
}
// 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
req := &request{h: h}
// TODO: now we don't know 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
return req, nil
}
func (server *Server) sendResponse(cc codec.Codec, h *codec.Header, body interface{}, sending *sync.Mutex) {
@@ -139,17 +121,16 @@ func (server *Server) sendResponse(cc codec.Codec, h *codec.Header, body interfa
if err := cc.Write(h, body); err != nil {
log.Println("rpc server: write response error:", err)
}
codec.HeaderPool.Put(h) // recycle Header object
}
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
+221
View File
@@ -0,0 +1,221 @@
// 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() {
var h codec.Header
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)
if err != nil {
call.Error = errors.New("reading body " + err.Error())
}
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
+55
View File
@@ -0,0 +1,55 @@
package main
import (
"geerpc"
"log"
"net"
"sync"
)
type Foo int
type Args struct{ Num1, Num2 int }
func (f Foo) Sum(args Args, reply *int) error {
*reply = args.Num1 + args.Num2
return nil
}
func startServer(addr chan string) {
var foo Foo
if err := geerpc.Register(&foo); err != nil {
log.Fatal("register error:", err)
}
// 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
var wg sync.WaitGroup
for i := 0; i < 5; i++ {
wg.Add(1)
go func(i int) {
defer wg.Done()
args := &Args{Num1: i, Num2: i * i}
var reply int
if err := client.Call("Foo.Sum", args, &reply); err != nil {
log.Fatal("call Foo.Sum error:", err)
}
log.Printf("%d + %d = %d", args.Num1, args.Num2, reply)
}(i)
}
wg.Wait()
}
+203
View File
@@ -0,0 +1,203 @@
// 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"
"geerpc/codec"
"io"
"log"
"net"
"reflect"
"strings"
"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 {
serviceMap sync.Map
}
// 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) // make sure to send a complete response
wg := new(sync.WaitGroup) // wait until all request are handled
for {
req, err := server.readRequest(cc)
if err != nil {
if req == nil {
break // it's not possible to recover, so close the connection
}
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
mtype *methodType
svc *service
}
func (server *Server) readRequestHeader(cc codec.Codec) (*codec.Header, error) {
h, _ := codec.HeaderPool.Get().(*codec.Header)
if err := cc.ReadHeader(h); err != nil {
codec.HeaderPool.Put(h)
if err != io.EOF && err != io.ErrUnexpectedEOF {
log.Println("rpc server: read header error:", err)
}
return nil, err
}
return h, nil
}
func (server *Server) findService(serviceMethod string) (svc *service, mtype *methodType, err error) {
dot := strings.LastIndex(serviceMethod, ".")
if dot < 0 {
err = errors.New("rpc server: service/method request ill-formed: " + serviceMethod)
return
}
serviceName, methodName := serviceMethod[:dot], serviceMethod[dot+1:]
svci, ok := server.serviceMap.Load(serviceName)
if !ok {
err = errors.New("rpc server: can't find service " + serviceName)
return
}
svc = svci.(*service)
mtype = svc.method[methodName]
if mtype == nil {
err = errors.New("rpc server: can't find method " + methodName)
}
return
}
func (server *Server) readRequest(cc codec.Codec) (*request, error) {
h, err := server.readRequestHeader(cc)
if err != nil {
return nil, err
}
req := &request{h: h}
req.svc, req.mtype, err = server.findService(h.ServiceMethod)
if err != nil {
return req, err
}
req.argv = req.mtype.newArgv()
req.replyv = req.mtype.newReplyv()
// make sure that argvi is a pointer, ReadBody need a pointer as parameter
argvi := req.argv.Interface()
if req.argv.Type().Kind() != reflect.Ptr {
argvi = req.argv.Addr().Interface()
}
if err = cc.ReadBody(argvi); err != nil {
log.Println("rpc server: read body err:", err)
return req, err
}
return req, nil
}
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)
}
codec.HeaderPool.Put(h) // recycle Header object
}
func (server *Server) handleRequest(cc codec.Codec, req *request, sending *sync.Mutex, wg *sync.WaitGroup) {
defer wg.Done()
err := req.svc.call(req.mtype, req.argv, req.replyv)
if err != nil {
req.h.Error = err.Error()
}
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) }
// Register publishes in the server the set of methods of the
// receiver value that satisfy the following conditions:
// - exported method of exported type
// - two arguments, both of exported type
// - the second argument is a pointer
// - one return value, of type error
func (server *Server) Register(rcvr interface{}) error {
s := newService(rcvr)
if _, dup := server.serviceMap.LoadOrStore(s.name, s); dup {
return errors.New("rpc: service already defined: " + s.name)
}
return nil
}
// Register publishes the receiver's methods in the DefaultServer.
func Register(rcvr interface{}) error { return DefaultServer.Register(rcvr) }
+9 -10
View File
@@ -4,26 +4,28 @@ import (
"go/ast"
"log"
"reflect"
"sync/atomic"
)
type methodType struct {
method reflect.Method
argType reflect.Type
replyType reflect.Type
numCalls uint64
}
func (m *methodType) NewArg() reflect.Value {
func (m *methodType) newArgv() reflect.Value {
var argv reflect.Value
// arg may be a pointer type, or a value type
if m.argType.Kind() == reflect.Ptr {
argv = reflect.New(m.argType.Elem())
} else {
argv = reflect.New(m.argType)
argv = reflect.New(m.argType).Elem()
}
return argv
}
func (m *methodType) NewReply() interface{} {
func (m *methodType) newReplyv() reflect.Value {
// reply must be a pointer type
replyv := reflect.New(m.replyType.Elem())
switch m.replyType.Elem().Kind() {
@@ -50,6 +52,7 @@ func newService(rcvr interface{}) *service {
if !ast.IsExported(s.name) {
log.Fatalf("rpc server: %s is not a valid service name", s.name)
}
s.registerMethods()
return s
}
@@ -64,7 +67,6 @@ func (s *service) registerMethods() {
if mType.Out(0) != reflect.TypeOf((*error)(nil)).Elem() {
continue
}
argType, replyType := mType.In(1), mType.In(2)
if !isExportedOrBuiltinType(argType) || !isExportedOrBuiltinType(replyType) {
continue
@@ -78,12 +80,9 @@ func (s *service) registerMethods() {
}
}
func (s *service) call(mType methodType, argv, replyv reflect.Value) error {
f := mType.method.Func
// if argv is not a ptr, need to indirect before calling.
if mType.argType.Kind() != reflect.Ptr {
argv = argv.Elem()
}
func (s *service) call(m *methodType, argv, replyv reflect.Value) error {
atomic.AddUint64(&m.numCalls, 1)
f := m.method.Func
returnValues := f.Call([]reflect.Value{s.rcvr, argv, replyv})
if errInter := returnValues[0].Interface(); errInter != nil {
return errInter.(error)
+48
View File
@@ -0,0 +1,48 @@
package geerpc
import (
"fmt"
"reflect"
"testing"
)
type Foo int
type Args struct{ Num1, Num2 int }
func (f Foo) Sum(args Args, reply *int) error {
*reply = args.Num1 + args.Num2
return nil
}
// it's not a exported method
func (f Foo) sum(args Args, reply *int) error {
*reply = args.Num1 + args.Num2
return nil
}
func _assert(condition bool, msg string, v ...interface{}) {
if !condition {
panic(fmt.Sprintf("assertion failed: "+msg, v...))
}
}
func TestNewService(t *testing.T) {
var foo Foo
s := newService(&foo)
_assert(len(s.method) == 1, "wrong service method, expect 1, but got %d", len(s.method))
mType := s.method["Sum"]
_assert(mType != nil, "wrong method, Sum shouldn't nil")
}
func TestMethodType_Call(t *testing.T) {
var foo Foo
s := newService(&foo)
mType := s.method["Sum"]
argv := mType.newArgv()
replyv := mType.newReplyv()
argv.Set(reflect.ValueOf(Args{Num1: 1, Num2: 3}))
err := s.call(mType, argv, replyv)
_assert(err == nil && *replyv.Interface().(*int) == 4 && mType.numCalls == 1, "failed to call Foo.Sum")
}