diff --git a/backend/app/store/remote/client.go b/backend/app/store/remote/client.go index 1fb8ed2f..9e5eb796 100644 --- a/backend/app/store/remote/client.go +++ b/backend/app/store/remote/client.go @@ -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) diff --git a/backend/app/store/remote/client_test.go b/backend/app/store/remote/client_test.go index 11105bef..98a2c794 100644 --- a/backend/app/store/remote/client_test.go +++ b/backend/app/store/remote/client_test.go @@ -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) diff --git a/backend/app/store/remote/remote.go b/backend/app/store/remote/remote.go index e72a41a9..e766efcb 100644 --- a/backend/app/store/remote/remote.go +++ b/backend/app/store/remote/remote.go @@ -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"` } diff --git a/backend/app/store/remote/server.go b/backend/app/store/remote/server.go index bfabef2a..4e70ad01 100644 --- a/backend/app/store/remote/server.go +++ b/backend/app/store/remote/server.go @@ -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) + }) +} diff --git a/backend/app/store/remote/server_test.go b/backend/app/store/remote/server_test.go index 4d6a11dc..ae0f38a0 100644 --- a/backend/app/store/remote/server_test.go +++ b/backend/app/store/remote/server_test.go @@ -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") +}