mirror of
https://github.com/geektutu/7days-golang.git
synced 2024-04-21 12:32:11 +00:00
229 lines
6.1 KiB
Go
229 lines
6.1 KiB
Go
// 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"
|
|
"reflect"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
)
|
|
|
|
const MagicNumber = 0x3bef5c
|
|
|
|
type Option struct {
|
|
MagicNumber int // MagicNumber marks this's a geerpc request
|
|
CodecType codec.Type // client may choose different Codec to encode body
|
|
ConnectTimeout time.Duration // 0 means no limit
|
|
HandleTimeout time.Duration
|
|
}
|
|
|
|
var DefaultOption = &Option{
|
|
MagicNumber: MagicNumber,
|
|
CodecType: codec.GobType,
|
|
ConnectTimeout: time.Second * 10,
|
|
}
|
|
|
|
// 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 Option
|
|
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), &opt)
|
|
}
|
|
|
|
// invalidRequest is a placeholder for response argv when error occurs
|
|
var invalidRequest = struct{}{}
|
|
|
|
func (server *Server) serveCodec(cc codec.Codec, opt *Option) {
|
|
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, opt.HandleTimeout)
|
|
}
|
|
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) {
|
|
var h codec.Header
|
|
if err := cc.ReadHeader(&h); err != nil {
|
|
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)
|
|
}
|
|
}
|
|
|
|
func (server *Server) handleRequest(cc codec.Codec, req *request, sending *sync.Mutex, wg *sync.WaitGroup, timeout time.Duration) {
|
|
defer wg.Done()
|
|
called := make(chan struct{})
|
|
sent := make(chan struct{})
|
|
go func() {
|
|
err := req.svc.call(req.mtype, req.argv, req.replyv)
|
|
called <- struct{}{}
|
|
if err != nil {
|
|
req.h.Error = err.Error()
|
|
server.sendResponse(cc, req.h, invalidRequest, sending)
|
|
sent <- struct{}{}
|
|
return
|
|
}
|
|
server.sendResponse(cc, req.h, req.replyv.Interface(), sending)
|
|
sent <- struct{}{}
|
|
}()
|
|
|
|
if timeout == 0 {
|
|
<-called
|
|
<-sent
|
|
return
|
|
}
|
|
select {
|
|
case <-time.After(timeout):
|
|
req.h.Error = fmt.Sprintf("rpc server: request handle timeout: expect within %s", timeout)
|
|
server.sendResponse(cc, req.h, invalidRequest, sending)
|
|
case <-called:
|
|
<-sent
|
|
}
|
|
}
|
|
|
|
// 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) }
|