diff --git a/pkg/ffapi/apiserver.go b/pkg/ffapi/apiserver.go index c349cc7..0a95920 100644 --- a/pkg/ffapi/apiserver.go +++ b/pkg/ffapi/apiserver.go @@ -87,6 +87,7 @@ type APIServerOptions[T any] struct { VersionedAPIs *VersionedAPIs MonitoringRoutes []*Route EnrichRequest func(r *APIRequest) (T, error) + MapError func(ctx context.Context, err error) error // optional: applied to every error before its status is derived Description string APIConfig config.Section MonitoringConfig config.Section @@ -304,6 +305,7 @@ func (as *apiServer[T]) handlerFactory(logLevel logrus.Level) *HandlerFactory { SupportFieldRedaction: as.SupportFieldRedaction, AlwaysPaginate: as.alwaysPaginate, HandleYAML: as.handleYAML, + MapError: as.MapError, } hf.SetAPIEntryLoggingLevel(logLevel) return hf diff --git a/pkg/ffapi/handler.go b/pkg/ffapi/handler.go index 8c3b40f..e2bdc66 100644 --- a/pkg/ffapi/handler.go +++ b/pkg/ffapi/handler.go @@ -78,6 +78,7 @@ type HandlerFactory struct { SupportFieldRedaction bool BasePath string BasePathParams []*PathParam + MapError func(ctx context.Context, err error) error // optional error mapping function, runs before status code extraction/setting apiEntryLoggingLevel logrus.Level // the log level at which entry/exit logging is enabled at all (does not affect trace logging) } @@ -396,6 +397,9 @@ func (hs *HandlerFactory) APIWrapper(handler HandlerFunction) http.HandlerFunc { status, err := handler(res, req) durationMS := float64(time.Since(startTime)) / float64(time.Millisecond) if err != nil { + if hs.MapError != nil { + err = hs.MapError(ctx, err) + } if ffe, ok := (interface{}(err)).(i18n.FFError); ok { if logrus.IsLevelEnabled(logrus.DebugLevel) { l.Debugf("%s:\n%s", ffe.Error(), ffe.StackTrace()) diff --git a/pkg/ffapi/handler_test.go b/pkg/ffapi/handler_test.go index 90ce1a4..4fced88 100644 --- a/pkg/ffapi/handler_test.go +++ b/pkg/ffapi/handler_test.go @@ -37,6 +37,7 @@ import ( "github.com/gorilla/mux" "github.com/hyperledger-firefly/common/pkg/config" "github.com/hyperledger-firefly/common/pkg/httpserver" + "github.com/hyperledger-firefly/common/pkg/i18n" "github.com/sirupsen/logrus" "github.com/stretchr/testify/assert" ) @@ -749,3 +750,83 @@ func TestGetFormEmptyValue(t *testing.T) { _, err := hs.getFormParams(req) require.NoError(t, err) } + +func TestMapErrorAppliedToEveryError(t *testing.T) { + driverErr := fmt.Errorf("driver: check constraint violated") + hs := newTestHandlerFactory("", nil) + var mapped []error + hs.MapError = func(ctx context.Context, err error) error { + assert.NotNil(t, ctx) + mapped = append(mapped, err) + if err == driverErr { + return i18n.NewError(ctx, i18n.MsgUnknownFieldValue, "thing", "x") // 400 + } + return err + } + jsonRoute := &Route{ + Name: "json", + Path: "/json", + Method: http.MethodPost, + JSONInputValue: func() interface{} { return make(map[string]interface{}) }, + JSONOutputValue: func() interface{} { return make(map[string]interface{}) }, + JSONOutputCodes: []int{200}, + JSONHandler: func(r *APIRequest) (output interface{}, err error) { + return nil, driverErr + }, + } + streamRoute := &Route{ + Name: "stream", + Path: "/stream", + Method: http.MethodGet, + JSONInputValue: nil, + JSONOutputCodes: []int{200}, + StreamHandler: func(r *APIRequest) (output io.ReadCloser, err error) { + return nil, driverErr + }, + } + + for _, tc := range []struct { + name string + route *Route + req *http.Request + status int + match string + }{ + {"json handler", jsonRoute, httptest.NewRequest(http.MethodPost, "/json", strings.NewReader(`{}`)), 400, "FF00111"}, + {"stream handler", streamRoute, httptest.NewRequest(http.MethodGet, "/stream", nil), 400, "FF00111"}, + {"input decode", jsonRoute, httptest.NewRequest(http.MethodPost, "/json", strings.NewReader(`{`)), 400, "unexpected EOF"}, + } { + t.Run(tc.name, func(t *testing.T) { + mapped = nil + tc.req.Header.Set("Content-Type", "application/json") + res := httptest.NewRecorder() + hs.RouteHandler(tc.route)(res, tc.req) + assert.Equal(t, tc.status, res.Code) + assert.Len(t, mapped, 1) + var resJSON map[string]interface{} + require.NoError(t, json.NewDecoder(res.Body).Decode(&resJSON)) + assert.Regexp(t, tc.match, resJSON["error"]) + }) + } +} + +func TestMapErrorNotCalledOnSuccess(t *testing.T) { + hs := newTestHandlerFactory("", nil) + hs.MapError = func(ctx context.Context, err error) error { + assert.Fail(t, "MapError called without an error") + return err + } + res := httptest.NewRecorder() + hs.RouteHandler(&Route{ + Name: "json", + Path: "/json", + Method: http.MethodGet, + JSONInputValue: nil, + JSONOutputValue: func() interface{} { return make(map[string]interface{}) }, + JSONOutputCodes: []int{200}, + JSONHandler: func(r *APIRequest) (output interface{}, err error) { + return map[string]interface{}{}, nil + }, + })(res, httptest.NewRequest(http.MethodGet, "/json", nil)) + assert.Equal(t, 200, res.Code) +}