mirror of
https://github.com/geektutu/7days-golang.git
synced 2024-04-21 12:32:11 +00:00
geerpc day2, implement a rpc client
This commit is contained in:
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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) }
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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,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)
|
||||
}
|
||||
}
|
||||
@@ -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) }
|
||||
Reference in New Issue
Block a user