remote server covered with tests
This commit is contained in:
@@ -4,6 +4,8 @@ import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"reflect"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
@@ -14,20 +16,32 @@ type Client struct {
|
||||
Client http.Client
|
||||
AuthUser string
|
||||
AuthPasswd string
|
||||
|
||||
id uint64
|
||||
}
|
||||
|
||||
// Call remote server with given method and arguments
|
||||
func (r *Client) Call(method string, args ...interface{}) (*Response, error) {
|
||||
|
||||
b, err := json.Marshal(Request{Method: method, Params: args})
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "marshaling failed for %s", method)
|
||||
var b []byte
|
||||
var err error
|
||||
if len(args) == 1 && reflect.TypeOf(args[0]).Kind() == reflect.Struct {
|
||||
b, err = json.Marshal(Request{Method: method, Params: args[0], ID: atomic.AddUint64(&r.id, 1)})
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "marshaling failed for %s", method)
|
||||
}
|
||||
} else {
|
||||
b, err = json.Marshal(Request{Method: method, Params: args, ID: atomic.AddUint64(&r.id, 1)})
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "marshaling failed for %s", method)
|
||||
}
|
||||
}
|
||||
|
||||
req, err := http.NewRequest("POST", r.API, bytes.NewBuffer(b))
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to make request for %s", method)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json; charset=utf-8")
|
||||
|
||||
if r.AuthUser != "" && r.AuthPasswd != "" {
|
||||
req.SetBasicAuth(r.AuthUser, r.AuthPasswd)
|
||||
|
||||
@@ -14,7 +14,7 @@ import (
|
||||
)
|
||||
|
||||
func TestClient_Call(t *testing.T) {
|
||||
ts := testServer(t, `{"method":"test","params":[123,"abc"]}`, `{"result":"12345"}`)
|
||||
ts := testServer(t, `{"method":"test","params":[123,"abc"],"id":1}`, `{"result":"12345"}`)
|
||||
defer ts.Close()
|
||||
c := Client{API: ts.URL, Client: http.Client{}}
|
||||
resp, err := c.Call("test", 123, "abc")
|
||||
@@ -25,8 +25,30 @@ func TestClient_Call(t *testing.T) {
|
||||
t.Logf("%v %T", res, res)
|
||||
}
|
||||
|
||||
func TestClient_CallWithObject(t *testing.T) {
|
||||
ts := testServer(t, `{"method":"test","params":{"F1":123,"F2":"abc","F3":"2019-06-09T23:03:55Z"},"id":1}`, `{"result":"12345"}`)
|
||||
defer ts.Close()
|
||||
c := Client{API: ts.URL, Client: http.Client{}}
|
||||
obj := struct {
|
||||
F1 int
|
||||
F2 string
|
||||
F3 time.Time
|
||||
}{
|
||||
F1: 123,
|
||||
F2: "abc",
|
||||
F3: time.Date(2019, 6, 9, 23, 3, 55, 0, time.UTC),
|
||||
}
|
||||
|
||||
resp, err := c.Call("test", obj)
|
||||
assert.NoError(t, err)
|
||||
res := ""
|
||||
err = json.Unmarshal(*resp.Result, &res)
|
||||
assert.Equal(t, "12345", res)
|
||||
t.Logf("%v %T", res, res)
|
||||
}
|
||||
|
||||
func TestClient_CallError(t *testing.T) {
|
||||
ts := testServer(t, `{"method":"test","params":[123,"abc"]}`, `{"error":"some error"}`)
|
||||
ts := testServer(t, `{"method":"test","params":[123,"abc"],"id":1}`, `{"error":"some error"}`)
|
||||
defer ts.Close()
|
||||
c := Client{API: ts.URL, Client: http.Client{}}
|
||||
_, err := c.Call("test", 123, "abc")
|
||||
@@ -34,7 +56,7 @@ func TestClient_CallError(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestClient_CallBadResponse(t *testing.T) {
|
||||
ts := testServer(t, `{"method":"test","params":[123,"abc"]}`, `{"result":"12345 invalid}`)
|
||||
ts := testServer(t, `{"method":"test","params":[123,"abc"],"id":1}`, `{"result":"12345 invalid}`)
|
||||
defer ts.Close()
|
||||
c := Client{API: ts.URL, Client: http.Client{}}
|
||||
_, err := c.Call("test", 123, "abc")
|
||||
@@ -42,7 +64,7 @@ func TestClient_CallBadResponse(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestClient_CallBadRemote(t *testing.T) {
|
||||
ts := testServer(t, `{"method":"test","params":[123,"abc"]}`, `{"result":"12345"}`)
|
||||
ts := testServer(t, `{"method":"test","params":[123,"abc"],"id":1}`, `{"result":"12345"}`)
|
||||
defer ts.Close()
|
||||
c := Client{API: "http://127.0.0.2", Client: http.Client{Timeout: 10 * time.Millisecond}}
|
||||
_, err := c.Call("test", 123)
|
||||
|
||||
@@ -8,15 +8,16 @@ import (
|
||||
"encoding/json"
|
||||
)
|
||||
|
||||
|
||||
// Request encloses method name and all params
|
||||
type Request struct {
|
||||
Method string `json:"method"`
|
||||
Params interface{} `json:"params"`
|
||||
Params interface{} `json:"params,omitempty"`
|
||||
ID uint64 `json:"id"`
|
||||
}
|
||||
|
||||
// Response encloses result and error received from remote server
|
||||
type Response struct {
|
||||
Result *json.RawMessage `json:"result,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
ID uint64 `json:"id"`
|
||||
}
|
||||
|
||||
@@ -4,20 +4,28 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/didip/tollbooth"
|
||||
"github.com/didip/tollbooth_chi"
|
||||
"github.com/go-chi/chi"
|
||||
"github.com/go-chi/chi/middleware"
|
||||
"github.com/go-chi/render"
|
||||
log "github.com/go-pkgz/lgr"
|
||||
R "github.com/go-pkgz/rest"
|
||||
"github.com/go-pkgz/rest/logger"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
// Server is json-rpc server with an optional basic auth
|
||||
type Server struct {
|
||||
CommandURL string
|
||||
API string
|
||||
AuthUser string
|
||||
AuthPasswd string
|
||||
Version string
|
||||
AppName string
|
||||
|
||||
funcs struct {
|
||||
m map[string]ServerFn
|
||||
@@ -30,33 +38,25 @@ type Server struct {
|
||||
}
|
||||
}
|
||||
|
||||
type ServerFn func(params *json.RawMessage) Response
|
||||
// ServerFn handler registered for each method with Add
|
||||
type ServerFn func(id uint64, params json.RawMessage) Response
|
||||
|
||||
// Run http server on given port
|
||||
func (s *Server) Run(port int) error {
|
||||
|
||||
if s.funcs.m == nil && len(s.funcs.m) == 0 {
|
||||
return errors.Errorf("nothing mapped for dispatch, Add has to be called prior to Run")
|
||||
}
|
||||
|
||||
router := chi.NewRouter()
|
||||
router.Use(middleware.Throttle(1000), middleware.RealIP, R.Recoverer(log.Default()))
|
||||
router.Use(R.AppInfo(s.AppName, "umputun", s.Version), R.Ping)
|
||||
logInfoWithBody := logger.New(logger.Log(log.Default()), logger.WithBody, logger.Prefix("[INFO]")).Handler
|
||||
router.Use(middleware.Timeout(5 * time.Second))
|
||||
router.Use(logInfoWithBody, tollbooth_chi.LimitHandler(tollbooth.NewLimiter(5, nil)), middleware.NoCache)
|
||||
router.Use(s.basicAuth)
|
||||
|
||||
type request struct {
|
||||
Method string `json:"method"`
|
||||
Params *json.RawMessage `json:"params"`
|
||||
}
|
||||
|
||||
router.Post(s.CommandURL, func(w http.ResponseWriter, r *http.Request) {
|
||||
req := request{}
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
fn, ok := s.funcs.m[req.Method]
|
||||
if !ok {
|
||||
w.WriteHeader(http.StatusNotImplemented)
|
||||
return
|
||||
}
|
||||
render.JSON(w, r, fn(req.Params))
|
||||
})
|
||||
router.Post(s.API, s.handler)
|
||||
|
||||
s.httpServer.Lock()
|
||||
s.httpServer.Server = &http.Server{
|
||||
@@ -72,18 +72,23 @@ func (s *Server) Run(port int) error {
|
||||
return s.httpServer.ListenAndServe()
|
||||
}
|
||||
|
||||
func (s *Server) EncodeResponse(resp interface{}) (Response, error) {
|
||||
// EncodeResponse convert anything to Response
|
||||
func (s *Server) EncodeResponse(id uint64, resp interface{}, e error) (Response, error) {
|
||||
v, err := json.Marshal(&resp)
|
||||
if err != nil {
|
||||
return Response{}, err
|
||||
}
|
||||
if e != nil {
|
||||
return Response{ID: id, Result: nil, Error: e.Error()}, nil
|
||||
}
|
||||
raw := json.RawMessage{}
|
||||
if err = raw.UnmarshalJSON(v); err != nil {
|
||||
return Response{}, err
|
||||
}
|
||||
return Response{Result: &raw}, nil
|
||||
return Response{ID: id, Result: &raw}, nil
|
||||
}
|
||||
|
||||
// Shutdown http server
|
||||
func (s *Server) Shutdown() error {
|
||||
s.httpServer.Lock()
|
||||
defer s.httpServer.Unlock()
|
||||
@@ -93,9 +98,47 @@ func (s *Server) Shutdown() error {
|
||||
return s.httpServer.Shutdown(context.TODO())
|
||||
}
|
||||
|
||||
// Add method handler
|
||||
func (s *Server) Add(method string, fn ServerFn) {
|
||||
s.funcs.once.Do(func() {
|
||||
s.funcs.m = map[string]ServerFn{}
|
||||
})
|
||||
s.funcs.m[method] = fn
|
||||
}
|
||||
|
||||
func (s *Server) handler(w http.ResponseWriter, r *http.Request) {
|
||||
req := struct {
|
||||
ID uint64 `json:"id"`
|
||||
Method string `json:"method"`
|
||||
Params *json.RawMessage `json:"params"`
|
||||
}{}
|
||||
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
fn, ok := s.funcs.m[req.Method]
|
||||
if !ok {
|
||||
w.WriteHeader(http.StatusNotImplemented)
|
||||
return
|
||||
}
|
||||
render.JSON(w, r, fn(req.ID, *req.Params))
|
||||
}
|
||||
|
||||
func (s *Server) basicAuth(h http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
if s.AuthUser == "" || s.AuthPasswd == "" {
|
||||
h.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
user, pass, ok := r.BasicAuth()
|
||||
if user != s.AuthUser || pass != s.AuthPasswd || !ok {
|
||||
w.Header().Set("WWW-Authenticate", `Basic realm="Restricted"`)
|
||||
http.Error(w, "Unauthorized.", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
h.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -5,35 +5,36 @@ import (
|
||||
"encoding/json"
|
||||
"io/ioutil"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestServer(t *testing.T) {
|
||||
s := Server{CommandURL: "/v1/cmd"}
|
||||
func TestServerPrimitiveTypes(t *testing.T) {
|
||||
s := Server{API: "/v1/cmd"}
|
||||
|
||||
type respData struct {
|
||||
Res1 string
|
||||
Res2 bool
|
||||
}
|
||||
|
||||
s.Add("test", func(params *json.RawMessage) Response {
|
||||
s.Add("test", func(id uint64, params json.RawMessage) Response {
|
||||
args := []interface{}{}
|
||||
if err := json.Unmarshal(*params, &args); err != nil {
|
||||
if err := json.Unmarshal(params, &args); err != nil {
|
||||
return Response{Error: err.Error()}
|
||||
}
|
||||
t.Logf("%+v", args)
|
||||
|
||||
assert.Equal(t, 4, len(args))
|
||||
assert.Equal(t, 3, len(args))
|
||||
assert.Equal(t, "blah", args[0].(string))
|
||||
assert.Equal(t, 42., args[1].(float64))
|
||||
assert.Equal(t, true, args[2].(bool))
|
||||
assert.Equal(t, "", args[3].(time.Time))
|
||||
|
||||
r, err := s.EncodeResponse(respData{"res blah", true})
|
||||
r, err := s.EncodeResponse(id, respData{"res blah", true}, nil)
|
||||
assert.NoError(t, err)
|
||||
return r
|
||||
})
|
||||
@@ -42,7 +43,7 @@ func TestServer(t *testing.T) {
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
// check with direct http call
|
||||
clientReq := Request{Method: "test", Params: []interface{}{"blah", 42, true, time.Date(2018, 6, 9, 16, 7, 25, 0, time.UTC)}}
|
||||
clientReq := Request{Method: "test", Params: []interface{}{"blah", 42, true}, ID: 123}
|
||||
b := bytes.Buffer{}
|
||||
require.NoError(t, json.NewEncoder(&b).Encode(clientReq))
|
||||
resp, err := http.Post("http://127.0.0.1:9091/v1/cmd", "application/json", &b)
|
||||
@@ -51,11 +52,52 @@ func TestServer(t *testing.T) {
|
||||
assert.Equal(t, 200, resp.StatusCode)
|
||||
data, err := ioutil.ReadAll(resp.Body)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, `{"result":{"Res1":"res blah","Res2":true}}`+"\n", string(data))
|
||||
assert.Equal(t, `{"result":{"Res1":"res blah","Res2":true},"id":123}`+"\n", string(data))
|
||||
|
||||
// check with client call
|
||||
c := Client{API: "http://127.0.0.1:9091/v1/cmd", Client: http.Client{}}
|
||||
r, err := c.Call("test", "blah", 42, true, time.Date(2018, 6, 9, 16, 7, 25, 0, time.UTC))
|
||||
r, err := c.Call("test", "blah", 42, true)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "", r.Error)
|
||||
|
||||
res := respData{}
|
||||
err = json.Unmarshal(*r.Result, &res)
|
||||
assert.Equal(t, respData{Res1: "res blah", Res2: true}, res)
|
||||
assert.Equal(t, uint64(1), r.ID)
|
||||
assert.NoError(t, s.Shutdown())
|
||||
}
|
||||
|
||||
func TestServerWithObject(t *testing.T) {
|
||||
s := Server{API: "/v1/cmd"}
|
||||
|
||||
type respData struct {
|
||||
Res1 string
|
||||
Res2 bool
|
||||
}
|
||||
|
||||
type reqData struct {
|
||||
Time time.Time
|
||||
F1 string
|
||||
F2 time.Duration
|
||||
}
|
||||
|
||||
s.Add("test", func(id uint64, params json.RawMessage) Response {
|
||||
arg := reqData{}
|
||||
if err := json.Unmarshal(params, &arg); err != nil {
|
||||
return Response{Error: err.Error()}
|
||||
}
|
||||
t.Logf("%+v", arg)
|
||||
|
||||
r, err := s.EncodeResponse(id, respData{"res blah", true}, nil)
|
||||
assert.NoError(t, err)
|
||||
return r
|
||||
})
|
||||
|
||||
go func() { s.Run(9091) }()
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
c := Client{API: "http://127.0.0.1:9091/v1/cmd", Client: http.Client{}}
|
||||
r, err := c.Call("test", reqData{Time: time.Now(), F1: "sawert", F2: time.Minute})
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "", r.Error)
|
||||
|
||||
@@ -65,3 +107,95 @@ func TestServer(t *testing.T) {
|
||||
|
||||
assert.NoError(t, s.Shutdown())
|
||||
}
|
||||
|
||||
func TestServerMethodNotImplemented(t *testing.T) {
|
||||
s := Server{}
|
||||
ts := httptest.NewServer(http.HandlerFunc(s.handler))
|
||||
defer ts.Close()
|
||||
s.Add("test", func(id uint64, params json.RawMessage) Response {
|
||||
return Response{}
|
||||
})
|
||||
|
||||
r := Request{Method: "blah"}
|
||||
buf := bytes.Buffer{}
|
||||
assert.NoError(t, json.NewEncoder(&buf).Encode(r))
|
||||
resp, err := http.Post(ts.URL, "application/json", &buf)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, http.StatusNotImplemented, resp.StatusCode)
|
||||
|
||||
assert.EqualError(t, s.Shutdown(), "http server is not running")
|
||||
}
|
||||
|
||||
func TestServerWithAuth(t *testing.T) {
|
||||
s := Server{API: "/v1/cmd", AuthUser: "user", AuthPasswd: "passwd"}
|
||||
|
||||
s.Add("test", func(id uint64, params json.RawMessage) Response {
|
||||
args := []interface{}{}
|
||||
if err := json.Unmarshal(params, &args); err != nil {
|
||||
return Response{Error: err.Error()}
|
||||
}
|
||||
t.Logf("%+v", args)
|
||||
|
||||
assert.Equal(t, 3, len(args))
|
||||
assert.Equal(t, "blah", args[0].(string))
|
||||
assert.Equal(t, 42., args[1].(float64))
|
||||
assert.Equal(t, true, args[2].(bool))
|
||||
|
||||
r, err := s.EncodeResponse(id, "res blah", nil)
|
||||
assert.NoError(t, err)
|
||||
return r
|
||||
})
|
||||
|
||||
go func() { s.Run(9091) }()
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
c := Client{API: "http://127.0.0.1:9091/v1/cmd", Client: http.Client{}, AuthUser: "user", AuthPasswd: "passwd"}
|
||||
r, err := c.Call("test", "blah", 42, true)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "", r.Error)
|
||||
val := ""
|
||||
err = json.Unmarshal(*r.Result, &val)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "res blah", val)
|
||||
|
||||
c = Client{API: "http://127.0.0.1:9091/v1/cmd", Client: http.Client{}}
|
||||
_, err = c.Call("test", "blah", 42, true)
|
||||
assert.EqualError(t, err, "bad status 401 for test")
|
||||
|
||||
assert.NoError(t, s.Shutdown())
|
||||
}
|
||||
|
||||
func TestServerErrReturn(t *testing.T) {
|
||||
s := Server{API: "/v1/cmd", AuthUser: "user", AuthPasswd: "passwd"}
|
||||
|
||||
s.Add("test", func(id uint64, params json.RawMessage) Response {
|
||||
args := []interface{}{}
|
||||
if err := json.Unmarshal(params, &args); err != nil {
|
||||
return Response{Error: err.Error()}
|
||||
}
|
||||
t.Logf("%+v", args)
|
||||
|
||||
assert.Equal(t, 3, len(args))
|
||||
assert.Equal(t, "blah", args[0].(string))
|
||||
assert.Equal(t, 42., args[1].(float64))
|
||||
assert.Equal(t, true, args[2].(bool))
|
||||
|
||||
r, err := s.EncodeResponse(id, "res blah", errors.New("some error"))
|
||||
assert.NoError(t, err)
|
||||
return r
|
||||
})
|
||||
|
||||
go func() { s.Run(9091) }()
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
c := Client{API: "http://127.0.0.1:9091/v1/cmd", Client: http.Client{}, AuthUser: "user", AuthPasswd: "passwd"}
|
||||
_, err := c.Call("test", "blah", 42, true)
|
||||
assert.EqualError(t, err, "some error")
|
||||
|
||||
assert.NoError(t, s.Shutdown())
|
||||
}
|
||||
|
||||
func TestServerNoHandlers(t *testing.T) {
|
||||
s := Server{API: "/v1/cmd", AuthUser: "user", AuthPasswd: "passwd"}
|
||||
assert.EqualError(t, s.Run(9091), "nothing mapped for dispatch, Add has to be called prior to Run")
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user