Initial working version
Summary: * Logging is still bad * test output in a lot of placed * FIXMEs every where Test Plan: None, just review Differential Revision: https://phabricator.cschwarz.com/D2
This commit is contained in:
-136
@@ -1,136 +0,0 @@
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"reflect"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
type Client struct {
|
||||
ml *MessageLayer
|
||||
logger Logger
|
||||
}
|
||||
|
||||
func NewClient(rwc io.ReadWriteCloser) *Client {
|
||||
return &Client{NewMessageLayer(rwc), noLogger{}}
|
||||
}
|
||||
|
||||
func (c *Client) SetLogger(logger Logger, logMessageLayer bool) {
|
||||
c.logger = logger
|
||||
if logMessageLayer {
|
||||
c.ml.logger = logger
|
||||
} else {
|
||||
c.ml.logger = noLogger{}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) Close() (err error) {
|
||||
|
||||
c.logger.Printf("sending Close request")
|
||||
header := Header{
|
||||
DataType: DataTypeControl,
|
||||
Endpoint: ControlEndpointClose,
|
||||
Accept: DataTypeControl,
|
||||
}
|
||||
err = c.ml.WriteHeader(&header)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
c.logger.Printf("reading Close ACK")
|
||||
ack, err := c.ml.ReadHeader()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.logger.Printf("received Close ACK: %#v", ack)
|
||||
if ack.Error != StatusOK {
|
||||
err = errors.Errorf("error hanging up: remote error (%s) %s", ack.Error, ack.ErrorMessage)
|
||||
return
|
||||
}
|
||||
|
||||
c.logger.Printf("closing MessageLayer")
|
||||
if err = c.ml.Close(); err != nil {
|
||||
c.logger.Printf("error closing RWC: %+v", err)
|
||||
return
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *Client) recvResponse() (h *Header, err error) {
|
||||
h, err = c.ml.ReadHeader()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "cannot read header")
|
||||
}
|
||||
// TODO validate
|
||||
return
|
||||
}
|
||||
|
||||
func (c *Client) writeRequest(h *Header) (err error) {
|
||||
// TODO validate
|
||||
err = c.ml.WriteHeader(h)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "cannot write header")
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func (c *Client) Call(endpoint string, in, out interface{}) (err error) {
|
||||
|
||||
var accept DataType
|
||||
{
|
||||
outType := reflect.TypeOf(out)
|
||||
if typeIsIOReaderPtr(outType) {
|
||||
accept = DataTypeOctets
|
||||
} else {
|
||||
accept = DataTypeMarshaledJSON
|
||||
}
|
||||
}
|
||||
|
||||
h := Header{
|
||||
Endpoint: endpoint,
|
||||
DataType: DataTypeMarshaledJSON,
|
||||
Accept: accept,
|
||||
}
|
||||
|
||||
if err = c.writeRequest(&h); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var buf bytes.Buffer
|
||||
if err = json.NewEncoder(&buf).Encode(in); err != nil {
|
||||
panic("cannot encode 'in' parameter")
|
||||
}
|
||||
if err = c.ml.WriteData(&buf); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
rh, err := c.recvResponse()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if rh.Error != StatusOK {
|
||||
return &RPCError{rh}
|
||||
}
|
||||
|
||||
rd := c.ml.ReadData()
|
||||
|
||||
switch accept {
|
||||
case DataTypeOctets:
|
||||
c.logger.Printf("setting out to ML data reader")
|
||||
outPtr := out.(*io.Reader) // we checked that above
|
||||
*outPtr = rd
|
||||
case DataTypeMarshaledJSON:
|
||||
c.logger.Printf("decoding marshaled json")
|
||||
if err = json.NewDecoder(c.ml.ReadData()).Decode(out); err != nil {
|
||||
return errors.Wrap(err, "cannot decode marshaled reply")
|
||||
}
|
||||
default:
|
||||
panic("implementation error") // accept is controlled by us
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
@@ -1,17 +0,0 @@
|
||||
// Code generated by "stringer -type=DataType"; DO NOT EDIT.
|
||||
|
||||
package rpc
|
||||
|
||||
import "strconv"
|
||||
|
||||
const _DataType_name = "DataTypeNoneDataTypeControlDataTypeMarshaledJSONDataTypeOctets"
|
||||
|
||||
var _DataType_index = [...]uint8{0, 12, 27, 48, 62}
|
||||
|
||||
func (i DataType) String() string {
|
||||
i -= 1
|
||||
if i >= DataType(len(_DataType_index)-1) {
|
||||
return "DataType(" + strconv.FormatInt(int64(i+1), 10) + ")"
|
||||
}
|
||||
return _DataType_name[_DataType_index[i]:_DataType_index[i+1]]
|
||||
}
|
||||
@@ -1,302 +0,0 @@
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
type Frame struct {
|
||||
Type FrameType
|
||||
NoMoreFrames bool
|
||||
PayloadLength uint32
|
||||
}
|
||||
|
||||
//go:generate stringer -type=FrameType
|
||||
type FrameType uint8
|
||||
|
||||
const (
|
||||
FrameTypeHeader FrameType = 0x01
|
||||
FrameTypeData FrameType = 0x02
|
||||
FrameTypeTrailer FrameType = 0x03
|
||||
FrameTypeRST FrameType = 0xff
|
||||
)
|
||||
|
||||
//go:generate stringer -type=Status
|
||||
type Status uint64
|
||||
|
||||
const (
|
||||
StatusOK Status = 1 + iota
|
||||
StatusRequestError
|
||||
StatusServerError
|
||||
// Returned when an error occurred but the side at fault cannot be determined
|
||||
StatusError
|
||||
)
|
||||
|
||||
type Header struct {
|
||||
// Request-only
|
||||
Endpoint string
|
||||
// Data type of body (request & reply)
|
||||
DataType DataType
|
||||
// Request-only
|
||||
Accept DataType
|
||||
// Reply-only
|
||||
Error Status
|
||||
// Reply-only
|
||||
ErrorMessage string
|
||||
}
|
||||
|
||||
func NewErrorHeader(status Status, format string, args ...interface{}) (h *Header) {
|
||||
h = &Header{}
|
||||
h.Error = status
|
||||
h.ErrorMessage = fmt.Sprintf(format, args...)
|
||||
return
|
||||
}
|
||||
|
||||
//go:generate stringer -type=DataType
|
||||
type DataType uint8
|
||||
|
||||
const (
|
||||
DataTypeNone DataType = 1 + iota
|
||||
DataTypeControl
|
||||
DataTypeMarshaledJSON
|
||||
DataTypeOctets
|
||||
)
|
||||
|
||||
const (
|
||||
MAX_PAYLOAD_LENGTH = 4 * 1024 * 1024
|
||||
MAX_HEADER_LENGTH = 4 * 1024
|
||||
)
|
||||
|
||||
type frameBridgingReader struct {
|
||||
l *MessageLayer
|
||||
frameType FrameType
|
||||
// < 0 means no limit
|
||||
bytesLeftToLimit int
|
||||
f Frame
|
||||
}
|
||||
|
||||
func NewFrameBridgingReader(l *MessageLayer, frameType FrameType, totalLimit int) *frameBridgingReader {
|
||||
return &frameBridgingReader{l, frameType, totalLimit, Frame{}}
|
||||
}
|
||||
|
||||
func (r *frameBridgingReader) Read(b []byte) (n int, err error) {
|
||||
if r.bytesLeftToLimit == 0 {
|
||||
r.l.logger.Printf("limit reached, returning EOF")
|
||||
return 0, io.EOF
|
||||
}
|
||||
log := r.l.logger
|
||||
if r.f.PayloadLength == 0 {
|
||||
|
||||
if r.f.NoMoreFrames {
|
||||
r.l.logger.Printf("no more frames flag set, returning EOF")
|
||||
err = io.EOF
|
||||
return
|
||||
}
|
||||
|
||||
log.Printf("reading frame")
|
||||
r.f, err = r.l.readFrame()
|
||||
if err != nil {
|
||||
log.Printf("error reading frame: %+v", err)
|
||||
return 0, err
|
||||
}
|
||||
log.Printf("read frame: %#v", r.f)
|
||||
if r.f.Type != r.frameType {
|
||||
err = errors.Wrapf(err, "expected frame of type %s", r.frameType)
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
maxread := len(b)
|
||||
if maxread > int(r.f.PayloadLength) {
|
||||
maxread = int(r.f.PayloadLength)
|
||||
}
|
||||
if r.bytesLeftToLimit > 0 && maxread > r.bytesLeftToLimit {
|
||||
maxread = r.bytesLeftToLimit
|
||||
}
|
||||
nb, err := r.l.rwc.Read(b[:maxread])
|
||||
log.Printf("read %v from rwc\n", nb)
|
||||
if nb < 0 {
|
||||
panic("should not return negative number of bytes")
|
||||
}
|
||||
r.f.PayloadLength -= uint32(nb)
|
||||
r.bytesLeftToLimit -= nb
|
||||
return nb, err // TODO io.EOF for maxread = r.f.PayloadLength ?
|
||||
}
|
||||
|
||||
type frameBridgingWriter struct {
|
||||
l *MessageLayer
|
||||
frameType FrameType
|
||||
// < 0 means no limit
|
||||
bytesLeftToLimit int
|
||||
payloadLength int
|
||||
buffer *bytes.Buffer
|
||||
}
|
||||
|
||||
func NewFrameBridgingWriter(l *MessageLayer, frameType FrameType, totalLimit int) *frameBridgingWriter {
|
||||
return &frameBridgingWriter{l, frameType, totalLimit, MAX_PAYLOAD_LENGTH, bytes.NewBuffer(make([]byte, 0, MAX_PAYLOAD_LENGTH))}
|
||||
}
|
||||
|
||||
func (w *frameBridgingWriter) Write(b []byte) (n int, err error) {
|
||||
for n = 0; n < len(b); {
|
||||
i, err := w.writeUntilFrameFull(b[n:])
|
||||
n += i
|
||||
if err != nil {
|
||||
return n, errors.WithStack(err)
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func (w *frameBridgingWriter) writeUntilFrameFull(b []byte) (n int, err error) {
|
||||
if len(b) <= 0 {
|
||||
return
|
||||
}
|
||||
if w.bytesLeftToLimit == 0 {
|
||||
err = errors.Errorf("message exceeds max number of allowed bytes")
|
||||
return
|
||||
}
|
||||
maxwrite := len(b)
|
||||
remainingInFrame := w.payloadLength - w.buffer.Len()
|
||||
|
||||
if maxwrite > remainingInFrame {
|
||||
maxwrite = remainingInFrame
|
||||
}
|
||||
if w.bytesLeftToLimit > 0 && maxwrite > w.bytesLeftToLimit {
|
||||
maxwrite = w.bytesLeftToLimit
|
||||
}
|
||||
w.buffer.Write(b[:maxwrite])
|
||||
w.bytesLeftToLimit -= maxwrite
|
||||
n = maxwrite
|
||||
if w.bytesLeftToLimit == 0 {
|
||||
err = w.flush(true)
|
||||
} else if w.buffer.Len() == w.payloadLength {
|
||||
err = w.flush(false)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func (w *frameBridgingWriter) flush(nomore bool) (err error) {
|
||||
|
||||
f := Frame{w.frameType, nomore, uint32(w.buffer.Len())}
|
||||
err = w.l.writeFrame(f)
|
||||
if err != nil {
|
||||
errors.WithStack(err)
|
||||
}
|
||||
_, err = w.buffer.WriteTo(w.l.rwc)
|
||||
return
|
||||
}
|
||||
|
||||
func (w *frameBridgingWriter) Close() (err error) {
|
||||
return w.flush(true)
|
||||
}
|
||||
|
||||
type MessageLayer struct {
|
||||
rwc io.ReadWriteCloser
|
||||
logger Logger
|
||||
}
|
||||
|
||||
func NewMessageLayer(rwc io.ReadWriteCloser) *MessageLayer {
|
||||
return &MessageLayer{rwc, noLogger{}}
|
||||
}
|
||||
|
||||
func (l *MessageLayer) Close() (err error) {
|
||||
f := Frame{
|
||||
Type: FrameTypeRST,
|
||||
NoMoreFrames: true,
|
||||
}
|
||||
if err = l.writeFrame(f); err != nil {
|
||||
l.logger.Printf("error sending RST frame: %s", err)
|
||||
return errors.WithStack(err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
var RST error = fmt.Errorf("reset frame observed on connection")
|
||||
|
||||
func (l *MessageLayer) readFrame() (f Frame, err error) {
|
||||
err = binary.Read(l.rwc, binary.LittleEndian, &f.Type)
|
||||
if err != nil {
|
||||
err = errors.WithStack(err)
|
||||
return
|
||||
}
|
||||
err = binary.Read(l.rwc, binary.LittleEndian, &f.NoMoreFrames)
|
||||
if err != nil {
|
||||
err = errors.WithStack(err)
|
||||
return
|
||||
}
|
||||
err = binary.Read(l.rwc, binary.LittleEndian, &f.PayloadLength)
|
||||
if err != nil {
|
||||
err = errors.WithStack(err)
|
||||
return
|
||||
}
|
||||
if f.Type == FrameTypeRST {
|
||||
l.logger.Printf("read RST frame")
|
||||
err = RST
|
||||
return
|
||||
}
|
||||
if f.PayloadLength > MAX_PAYLOAD_LENGTH {
|
||||
err = errors.Errorf("frame exceeds max payload length")
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func (l *MessageLayer) writeFrame(f Frame) (err error) {
|
||||
err = binary.Write(l.rwc, binary.LittleEndian, &f.Type)
|
||||
if err != nil {
|
||||
return errors.WithStack(err)
|
||||
}
|
||||
err = binary.Write(l.rwc, binary.LittleEndian, &f.NoMoreFrames)
|
||||
if err != nil {
|
||||
return errors.WithStack(err)
|
||||
}
|
||||
err = binary.Write(l.rwc, binary.LittleEndian, &f.PayloadLength)
|
||||
if err != nil {
|
||||
return errors.WithStack(err)
|
||||
}
|
||||
if f.PayloadLength > MAX_PAYLOAD_LENGTH {
|
||||
err = errors.Errorf("frame exceeds max payload length")
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func (l *MessageLayer) ReadHeader() (h *Header, err error) {
|
||||
|
||||
r := NewFrameBridgingReader(l, FrameTypeHeader, MAX_HEADER_LENGTH)
|
||||
h = &Header{}
|
||||
if err = json.NewDecoder(r).Decode(&h); err != nil {
|
||||
l.logger.Printf("cannot decode marshaled header: %s", err)
|
||||
return nil, err
|
||||
}
|
||||
return h, nil
|
||||
}
|
||||
|
||||
func (l *MessageLayer) WriteHeader(h *Header) (err error) {
|
||||
w := NewFrameBridgingWriter(l, FrameTypeHeader, MAX_HEADER_LENGTH)
|
||||
err = json.NewEncoder(w).Encode(h)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "cannot encode header, probably fatal")
|
||||
}
|
||||
w.Close()
|
||||
return
|
||||
}
|
||||
|
||||
func (l *MessageLayer) ReadData() (reader io.Reader) {
|
||||
r := NewFrameBridgingReader(l, FrameTypeData, -1)
|
||||
return r
|
||||
}
|
||||
|
||||
func (l *MessageLayer) WriteData(source io.Reader) (err error) {
|
||||
w := NewFrameBridgingWriter(l, FrameTypeData, -1)
|
||||
_, err = io.Copy(w, source)
|
||||
if err != nil {
|
||||
return errors.WithStack(err)
|
||||
}
|
||||
err = w.Close()
|
||||
return
|
||||
}
|
||||
@@ -1,26 +0,0 @@
|
||||
// Code generated by "stringer -type=FrameType"; DO NOT EDIT.
|
||||
|
||||
package rpc
|
||||
|
||||
import "strconv"
|
||||
|
||||
const (
|
||||
_FrameType_name_0 = "FrameTypeHeaderFrameTypeDataFrameTypeTrailer"
|
||||
_FrameType_name_1 = "FrameTypeRST"
|
||||
)
|
||||
|
||||
var (
|
||||
_FrameType_index_0 = [...]uint8{0, 15, 28, 44}
|
||||
)
|
||||
|
||||
func (i FrameType) String() string {
|
||||
switch {
|
||||
case 1 <= i && i <= 3:
|
||||
i -= 1
|
||||
return _FrameType_name_0[_FrameType_index_0[i]:_FrameType_index_0[i+1]]
|
||||
case i == 255:
|
||||
return _FrameType_name_1
|
||||
default:
|
||||
return "FrameType(" + strconv.FormatInt(int64(i), 10) + ")"
|
||||
}
|
||||
}
|
||||
@@ -1,63 +0,0 @@
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"github.com/pkg/errors"
|
||||
"reflect"
|
||||
)
|
||||
|
||||
type LocalRPC struct {
|
||||
endpoints map[string]reflect.Value
|
||||
}
|
||||
|
||||
func NewLocalRPC() *LocalRPC {
|
||||
return &LocalRPC{make(map[string]reflect.Value, 0)}
|
||||
}
|
||||
|
||||
func (s *LocalRPC) RegisterEndpoint(name string, handler interface{}) (err error) {
|
||||
_, ok := s.endpoints[name]
|
||||
if ok {
|
||||
return errors.Errorf("already set up an endpoint for '%s'", name)
|
||||
}
|
||||
ep, err := makeEndpointDescr(handler)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.endpoints[name] = ep.handler
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *LocalRPC) Serve() (err error) {
|
||||
panic("local cannot serve")
|
||||
}
|
||||
|
||||
func (c *LocalRPC) Call(endpoint string, in, out interface{}) (err error) {
|
||||
ep, ok := c.endpoints[endpoint]
|
||||
if !ok {
|
||||
panic("implementation error: implementation should not call local RPC without knowing which endpoints exist")
|
||||
}
|
||||
|
||||
args := []reflect.Value{reflect.ValueOf(in), reflect.ValueOf(out)}
|
||||
|
||||
if err = checkRPCParamTypes(args[0].Type(), args[1].Type()); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
rets := ep.Call(args)
|
||||
|
||||
if len(rets) != 1 {
|
||||
panic("implementation error: endpoints must have one error ")
|
||||
}
|
||||
if err = checkRPCReturnType(rets[0].Type()); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
err = nil
|
||||
if !rets[0].IsNil() {
|
||||
err = rets[0].Interface().(error) // we checked that above
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func (c *LocalRPC) Close() (err error) {
|
||||
return nil
|
||||
}
|
||||
-259
@@ -1,259 +0,0 @@
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"reflect"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
type Server struct {
|
||||
ml *MessageLayer
|
||||
logger Logger
|
||||
endpoints map[string]endpointDescr
|
||||
}
|
||||
|
||||
type typeMap struct {
|
||||
local reflect.Type
|
||||
proto DataType
|
||||
}
|
||||
type endpointDescr struct {
|
||||
inType typeMap
|
||||
outType typeMap
|
||||
handler reflect.Value
|
||||
}
|
||||
|
||||
type MarshaledJSONEndpoint func(bodyJSON interface{})
|
||||
|
||||
func NewServer(rwc io.ReadWriteCloser) *Server {
|
||||
ml := NewMessageLayer(rwc)
|
||||
return &Server{
|
||||
ml, noLogger{}, make(map[string]endpointDescr),
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) SetLogger(logger Logger, logMessageLayer bool) {
|
||||
s.logger = logger
|
||||
if logMessageLayer {
|
||||
s.ml.logger = logger
|
||||
} else {
|
||||
s.ml.logger = noLogger{}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) RegisterEndpoint(name string, handler interface{}) (err error) {
|
||||
_, ok := s.endpoints[name]
|
||||
if ok {
|
||||
return errors.Errorf("already set up an endpoint for '%s'", name)
|
||||
}
|
||||
s.endpoints[name], err = makeEndpointDescr(handler)
|
||||
return
|
||||
}
|
||||
|
||||
func checkResponseHeader(h *Header) (err error) {
|
||||
var statusNotSet Status
|
||||
if h.Error == statusNotSet {
|
||||
return errors.Errorf("status has zero-value")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Server) writeResponse(h *Header) (err error) {
|
||||
// TODO validate
|
||||
return s.ml.WriteHeader(h)
|
||||
}
|
||||
|
||||
func (s *Server) recvRequest() (h *Header, err error) {
|
||||
h, err = s.ml.ReadHeader()
|
||||
if err != nil {
|
||||
s.logger.Printf("error reading header: %s", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
s.logger.Printf("validating request")
|
||||
err = nil // TODO validate
|
||||
if err == nil {
|
||||
return h, nil
|
||||
}
|
||||
s.logger.Printf("request validation error: %s", err)
|
||||
|
||||
r := NewErrorHeader(StatusRequestError, "%s", err)
|
||||
return nil, s.writeResponse(r)
|
||||
}
|
||||
|
||||
var doneServeNext error = errors.New("this should not cause a HangUp() in the server")
|
||||
var doneStopServing error = errors.New("this should cause the server to close the connection")
|
||||
|
||||
var ProtocolError error = errors.New("protocol error, server should hang up")
|
||||
|
||||
const ControlEndpointClose string = "Close"
|
||||
|
||||
// Serve the connection until failure or the client hangs up
|
||||
func (s *Server) Serve() (err error) {
|
||||
for {
|
||||
|
||||
err = s.ServeRequest()
|
||||
|
||||
if err == nil {
|
||||
continue
|
||||
}
|
||||
if err == doneServeNext {
|
||||
s.logger.Printf("subroutine returned pseudo-error indicating early-exit")
|
||||
err = nil
|
||||
continue
|
||||
}
|
||||
|
||||
if err == doneStopServing {
|
||||
s.logger.Printf("subroutine returned pseudo-error indicating close request")
|
||||
err = nil
|
||||
break
|
||||
}
|
||||
|
||||
break
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
s.logger.Printf("an error occurred that could not be handled on PRC protocol level: %+v", err)
|
||||
}
|
||||
|
||||
s.logger.Printf("cloing MessageLayer")
|
||||
if mlErr := s.ml.Close(); mlErr != nil {
|
||||
s.logger.Printf("error closing MessageLayer: %+v", mlErr)
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
// Serve a single request
|
||||
// * wait for request to come in
|
||||
// * call handler
|
||||
// * reply
|
||||
//
|
||||
// The connection is left open, the next bytes on the conn should be
|
||||
// the next request header.
|
||||
//
|
||||
// Returns an err != nil if the error is bad enough to hang up on the client.
|
||||
// Examples: protocol version mismatches, protocol errors in general, ...
|
||||
// Non-Examples: a handler error
|
||||
func (s *Server) ServeRequest() (err error) {
|
||||
|
||||
ml := s.ml
|
||||
|
||||
s.logger.Printf("reading header")
|
||||
h, err := s.recvRequest()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if h.DataType == DataTypeControl {
|
||||
switch h.Endpoint {
|
||||
case ControlEndpointClose:
|
||||
ack := Header{Error: StatusOK, DataType: DataTypeControl}
|
||||
err = s.writeResponse(&ack)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return doneStopServing
|
||||
default:
|
||||
r := NewErrorHeader(StatusRequestError, "unregistered control endpoint %s", h.Endpoint)
|
||||
return s.writeResponse(r)
|
||||
}
|
||||
}
|
||||
|
||||
ep, ok := s.endpoints[h.Endpoint]
|
||||
if !ok {
|
||||
r := NewErrorHeader(StatusRequestError, "unregistered endpoint %s", h.Endpoint)
|
||||
return s.writeResponse(r)
|
||||
}
|
||||
|
||||
if ep.inType.proto != h.DataType {
|
||||
r := NewErrorHeader(StatusRequestError, "wrong DataType for endpoint %s (has %s, you provided %s)", h.Endpoint, ep.inType.proto, h.DataType)
|
||||
return s.writeResponse(r)
|
||||
}
|
||||
|
||||
if ep.outType.proto != h.Accept {
|
||||
r := NewErrorHeader(StatusRequestError, "wrong Accept for endpoint %s (has %s, you provided %s)", h.Endpoint, ep.outType.proto, h.Accept)
|
||||
return s.writeResponse(r)
|
||||
}
|
||||
|
||||
dr := ml.ReadData()
|
||||
|
||||
// Determine inval
|
||||
var inval reflect.Value
|
||||
switch ep.inType.proto {
|
||||
case DataTypeMarshaledJSON:
|
||||
// Unmarshal input
|
||||
inval = reflect.New(ep.inType.local.Elem())
|
||||
invalIface := inval.Interface()
|
||||
err = json.NewDecoder(dr).Decode(invalIface)
|
||||
if err != nil {
|
||||
r := NewErrorHeader(StatusRequestError, "cannot decode marshaled JSON: %s", err)
|
||||
return s.writeResponse(r)
|
||||
}
|
||||
case DataTypeOctets:
|
||||
// Take data as is
|
||||
inval = reflect.ValueOf(dr)
|
||||
default:
|
||||
panic("not implemented")
|
||||
}
|
||||
|
||||
outval := reflect.New(ep.outType.local.Elem()) // outval is a double pointer
|
||||
|
||||
s.logger.Printf("before handler, inval=%v outval=%v", inval, outval)
|
||||
|
||||
// Call the handler
|
||||
errs := ep.handler.Call([]reflect.Value{inval, outval})
|
||||
|
||||
if !errs[0].IsNil() {
|
||||
he := errs[0].Interface().(error) // we checked that before...
|
||||
s.logger.Printf("handler returned error: %s", err)
|
||||
r := NewErrorHeader(StatusError, "%s", he.Error())
|
||||
return s.writeResponse(r)
|
||||
}
|
||||
|
||||
switch ep.outType.proto {
|
||||
|
||||
case DataTypeMarshaledJSON:
|
||||
|
||||
var dataBuf bytes.Buffer
|
||||
// Marshal output
|
||||
err = json.NewEncoder(&dataBuf).Encode(outval.Interface())
|
||||
if err != nil {
|
||||
r := NewErrorHeader(StatusServerError, "cannot marshal response: %s", err)
|
||||
return s.writeResponse(r)
|
||||
}
|
||||
|
||||
replyHeader := Header{
|
||||
Error: StatusOK,
|
||||
DataType: ep.outType.proto,
|
||||
}
|
||||
if err = s.writeResponse(&replyHeader); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err = ml.WriteData(&dataBuf); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
case DataTypeOctets:
|
||||
|
||||
h := Header{
|
||||
Error: StatusOK,
|
||||
DataType: DataTypeOctets,
|
||||
}
|
||||
if err = s.writeResponse(&h); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
reader := outval.Interface().(*io.Reader) // we checked that when adding the endpoint
|
||||
err = ml.WriteData(*reader)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
-111
@@ -1,111 +0,0 @@
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"github.com/pkg/errors"
|
||||
"io"
|
||||
"reflect"
|
||||
)
|
||||
|
||||
type RPCServer interface {
|
||||
Serve() (err error)
|
||||
RegisterEndpoint(name string, handler interface{}) (err error)
|
||||
}
|
||||
|
||||
type RPCClient interface {
|
||||
Call(endpoint string, in, out interface{}) (err error)
|
||||
Close() (err error)
|
||||
}
|
||||
|
||||
type Logger interface {
|
||||
Printf(format string, args ...interface{})
|
||||
}
|
||||
|
||||
type noLogger struct{}
|
||||
|
||||
func (l noLogger) Printf(format string, args ...interface{}) {}
|
||||
func typeIsIOReader(t reflect.Type) bool {
|
||||
return t == reflect.TypeOf((*io.Reader)(nil)).Elem()
|
||||
}
|
||||
|
||||
func typeIsIOReaderPtr(t reflect.Type) bool {
|
||||
return t == reflect.TypeOf((*io.Reader)(nil))
|
||||
}
|
||||
|
||||
// An error returned by the Client if the response indicated a status code other than StatusOK
|
||||
type RPCError struct {
|
||||
ResponseHeader *Header
|
||||
}
|
||||
|
||||
func (e *RPCError) Error() string {
|
||||
return fmt.Sprintf("%s: %s", e.ResponseHeader.Error, e.ResponseHeader.ErrorMessage)
|
||||
}
|
||||
|
||||
type RPCProtoError struct {
|
||||
Message string
|
||||
UnderlyingError error
|
||||
}
|
||||
|
||||
func (e *RPCProtoError) Error() string {
|
||||
return e.Message
|
||||
}
|
||||
|
||||
func checkRPCParamTypes(in, out reflect.Type) (err error) {
|
||||
if !(in.Kind() == reflect.Ptr || typeIsIOReader(in)) {
|
||||
err = errors.Errorf("input parameter must be a pointer or an io.Reader, is of kind %s, type %s", in.Kind(), in)
|
||||
return
|
||||
}
|
||||
if !(out.Kind() == reflect.Ptr) {
|
||||
err = errors.Errorf("second input parameter (the non-error output parameter) must be a pointer or an *io.Reader")
|
||||
return
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func checkRPCReturnType(rt reflect.Type) (err error) {
|
||||
errInterfaceType := reflect.TypeOf((*error)(nil)).Elem()
|
||||
if !rt.Implements(errInterfaceType) {
|
||||
err = errors.Errorf("handler must return an error")
|
||||
return
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func makeEndpointDescr(handler interface{}) (descr endpointDescr, err error) {
|
||||
|
||||
ht := reflect.TypeOf(handler)
|
||||
|
||||
if ht.Kind() != reflect.Func {
|
||||
err = errors.Errorf("handler must be of kind reflect.Func")
|
||||
return
|
||||
}
|
||||
|
||||
if ht.NumIn() != 2 || ht.NumOut() != 1 {
|
||||
err = errors.Errorf("handler must have exactly two input parameters and one output parameter")
|
||||
return
|
||||
}
|
||||
if err = checkRPCParamTypes(ht.In(0), ht.In(1)); err != nil {
|
||||
return
|
||||
}
|
||||
if err = checkRPCReturnType(ht.Out(0)); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
descr.handler = reflect.ValueOf(handler)
|
||||
descr.inType.local = ht.In(0)
|
||||
descr.outType.local = ht.In(1)
|
||||
|
||||
if typeIsIOReader(ht.In(0)) {
|
||||
descr.inType.proto = DataTypeOctets
|
||||
} else {
|
||||
descr.inType.proto = DataTypeMarshaledJSON
|
||||
}
|
||||
|
||||
if typeIsIOReaderPtr(ht.In(1)) {
|
||||
descr.outType.proto = DataTypeOctets
|
||||
} else {
|
||||
descr.outType.proto = DataTypeMarshaledJSON
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
@@ -1,17 +0,0 @@
|
||||
// Code generated by "stringer -type=Status"; DO NOT EDIT.
|
||||
|
||||
package rpc
|
||||
|
||||
import "strconv"
|
||||
|
||||
const _Status_name = "StatusOKStatusRequestErrorStatusServerErrorStatusError"
|
||||
|
||||
var _Status_index = [...]uint8{0, 8, 26, 43, 54}
|
||||
|
||||
func (i Status) String() string {
|
||||
i -= 1
|
||||
if i >= Status(len(_Status_index)-1) {
|
||||
return "Status(" + strconv.FormatInt(int64(i+1), 10) + ")"
|
||||
}
|
||||
return _Status_name[_Status_index[i]:_Status_index[i+1]]
|
||||
}
|
||||
Reference in New Issue
Block a user