mirror of
https://github.com/geektutu/7days-golang.git
synced 2024-04-21 12:32:11 +00:00
fix HeaderPool recycle opportunity, add geerpc day3 service
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
module geerpc
|
||||
|
||||
go 1.13
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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) }
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
Reference in New Issue
Block a user