diff --git a/gateway.go b/gateway.go index ed5004c..1100476 100644 --- a/gateway.go +++ b/gateway.go @@ -35,6 +35,8 @@ type Gateway struct { locationPriorities []string preExecutionHook PreExecutionStepHook postExecutionHook PostExecutionStepHook + preOperationHook PreOperationHook + postOperationHook PostOperationHook // group up the list of middlewares at startup to avoid it during execution requestMiddlewares []graphql.NetworkMiddleware @@ -352,6 +354,22 @@ func WithPostExecutionHook(hook PostExecutionStepHook) Option { } } +// WithPreOperationHook returns an Option that sets a hook run by GraphQLHandler +// once per operation, before it is planned and executed +func WithPreOperationHook(hook PreOperationHook) Option { + return func(g *Gateway) { + g.preOperationHook = hook + } +} + +// WithPostOperationHook returns an Option that sets a hook run by GraphQLHandler +// once per operation, on the payload about to be written for it +func WithPostOperationHook(hook PostOperationHook) Option { + return func(g *Gateway) { + g.postOperationHook = hook + } +} + // WithLogger returns an Option that sets the logger of the gateway func WithLogger(l Logger) Option { return func(g *Gateway) { diff --git a/http.go b/http.go index 3e39cca..2281923 100644 --- a/http.go +++ b/http.go @@ -30,6 +30,22 @@ type HTTPOperation struct { type setResultFunc func(r map[string]interface{}) +// PreOperationHook is a function that GraphQLHandler runs once for every +// operation in a request, after the RequestContext has been built and before +// the operation is planned. Operations in a batch request share the incoming +// http.Request context, so a hook that needs per-operation state should +// derive a child context and store it in rc.Context; the executor and the +// queryers see that context during execution. +type PreOperationHook func(rc *RequestContext) + +// PostOperationHook is a function that GraphQLHandler runs once for every +// operation in a request, as the last step before the operation's response +// payload is written. The payload is exactly what the gateway would have +// written without the hook, including any keys the gateway adds itself. The +// hook may add or modify keys and is responsible for preserving existing ones +// it does not mean to change. +type PostOperationHook func(rc *RequestContext, payload map[string]interface{}) + func formatErrors(err error) map[string]interface{} { return formatErrorsWithCode(nil, err, "UNKNOWN_ERROR") } @@ -104,6 +120,10 @@ func (g *Gateway) GraphQLHandler(w http.ResponseWriter, r *http.Request) { CacheKey: cacheKey, } + if g.preOperationHook != nil { + g.preOperationHook(requestContext) + } + // Get the plan, and return a 400 if we can't get the plan plan, err := g.GetPlans(requestContext) if err != nil { @@ -161,17 +181,17 @@ func (g *Gateway) executeRequest(requestContext *RequestContext, plan QueryPlanL // fire the query with the request context passed through to execution result, err := g.Execute(requestContext, plan) - if err != nil { - setResult(formatErrorsWithCode(result, err, "INTERNAL_SERVER_ERROR")) - - return - } // the result for this operation - payload := map[string]interface{}{"data": result} + var payload map[string]interface{} + if err != nil { + payload = formatErrorsWithCode(result, err, "INTERNAL_SERVER_ERROR") + } else { + payload = map[string]interface{}{"data": result} + } // if there was a cache key associated with this query - if requestContext.CacheKey != "" { + if err == nil && requestContext.CacheKey != "" { // embed the cache key in the response payload["extensions"] = map[string]interface{}{ "persistedQuery": map[string]interface{}{ @@ -181,6 +201,11 @@ func (g *Gateway) executeRequest(requestContext *RequestContext, plan QueryPlanL } } + // the hook gets the last word on the payload + if g.postOperationHook != nil { + g.postOperationHook(requestContext, payload) + } + // add this result to the list setResult(payload) } diff --git a/http_test.go b/http_test.go index 0280a8a..73460f2 100644 --- a/http_test.go +++ b/http_test.go @@ -1557,3 +1557,173 @@ func TestGraphQLHandler_OptionsMethod(t *testing.T) { ] }`, response.Body.String()) } + +type operationHookCtxKey struct{} + +func TestGraphQLHandler_operationHooks(t *testing.T) { + t.Parallel() + schema, err := graphql.LoadSchema(` + type Query { + queryA: String! + queryB: String! + } + `) + assert.NoError(t, err) + + // the executor reads the per-operation value the pre hook stored in the context + // and echoes it back so the test can verify each operation saw its own value + gateway, err := New([]*graphql.RemoteSchema{ + {Schema: schema, URL: "url1"}, + }, WithExecutor(ExecutorFunc( + func(ec *ExecutionContext) (map[string]interface{}, error) { + seen, _ := ec.RequestContext.Value(operationHookCtxKey{}).(string) + return map[string]interface{}{"seen": seen}, nil + }, + )), WithPreOperationHook(func(rc *RequestContext) { + rc.Context = context.WithValue(rc.Context, operationHookCtxKey{}, "ctx-"+rc.OperationName) + }), WithPostOperationHook(func(rc *RequestContext, payload map[string]interface{}) { + payload["extensions"] = map[string]interface{}{ + "operation": rc.OperationName, + "seen": rc.Context.Value(operationHookCtxKey{}), + } + })) + if err != nil { + t.Error(err.Error()) + return + } + + request := httptest.NewRequestWithContext(t.Context(), http.MethodPost, "/graphql", strings.NewReader(`[ + { "query": "query queryAOperation { queryA }", "operationName": "queryAOperation" }, + { "query": "query queryBOperation { queryB }", "operationName": "queryBOperation" } + ]`)) + responseRecorder := httptest.NewRecorder() + gateway.GraphQLHandler(responseRecorder, request) + + response := responseRecorder.Result() + defer response.Body.Close() + assert.Equal(t, http.StatusOK, response.StatusCode) + + result := []map[string]interface{}{} + if err := json.NewDecoder(response.Body).Decode(&result); err != nil { + t.Error(err.Error()) + return + } + + // each operation in the batch ran under its own context and got its own payload + assert.Equal(t, []map[string]interface{}{ + { + "data": map[string]interface{}{"seen": "ctx-queryAOperation"}, + "extensions": map[string]interface{}{"operation": "queryAOperation", "seen": "ctx-queryAOperation"}, + }, + { + "data": map[string]interface{}{"seen": "ctx-queryBOperation"}, + "extensions": map[string]interface{}{"operation": "queryBOperation", "seen": "ctx-queryBOperation"}, + }, + }, result) +} + +func TestGraphQLHandler_postOperationHookSeesGatewayExtensions(t *testing.T) { + t.Parallel() + schema, err := graphql.LoadSchema(` + type Query { + allUsers: [String!]! + } + `) + assert.NoError(t, err) + + // the hook runs last, so it sees the gateway's own extensions and can add to them + gateway, err := New([]*graphql.RemoteSchema{ + {Schema: schema, URL: "url1"}, + }, WithExecutor(ExecutorFunc( + func(*ExecutionContext) (map[string]interface{}, error) { + return map[string]interface{}{"Hello": "world"}, nil + }, + )), WithAutomaticQueryPlanCache(), WithPostOperationHook(func(_ *RequestContext, payload map[string]interface{}) { + ext, ok := payload["extensions"].(map[string]interface{}) + if assert.True(t, ok, "hook should see the gateway's extensions") { + ext["fromHook"] = "yes" + } + })) + if err != nil { + t.Error(err.Error()) + return + } + + request := httptest.NewRequestWithContext(t.Context(), http.MethodPost, "/graphql", strings.NewReader(` + { + "query": "{ allUsers }", + "extensions": { + "persistedQuery": { + "version": 1, + "sha256Hash": "1234" + } + } + } + `)) + responseRecorder := httptest.NewRecorder() + gateway.GraphQLHandler(responseRecorder, request) + + response := responseRecorder.Result() + defer response.Body.Close() + assert.Equal(t, http.StatusOK, response.StatusCode) + + result := map[string]interface{}{} + if err := json.NewDecoder(response.Body).Decode(&result); err != nil { + t.Error(err.Error()) + return + } + + // both the gateway's key and the hook's key are in the response + assert.Equal(t, map[string]interface{}{ + "data": map[string]interface{}{"Hello": "world"}, + "extensions": map[string]interface{}{ + "fromHook": "yes", + "persistedQuery": map[string]interface{}{ + "sha265Hash": "1234", + "version": "1", + }, + }, + }, result) +} + +func TestGraphQLHandler_postOperationHookRunsOnExecutionError(t *testing.T) { + t.Parallel() + schema, err := graphql.LoadSchema(` + type Query { + allUsers: [String!]! + } + `) + assert.NoError(t, err) + + gateway, err := New([]*graphql.RemoteSchema{ + {Schema: schema, URL: "url1"}, + }, WithExecutor(ExecutorFunc( + func(*ExecutionContext) (map[string]interface{}, error) { + return nil, errors.New("downstream exploded") + }, + )), WithPostOperationHook(func(_ *RequestContext, payload map[string]interface{}) { + payload["extensions"] = map[string]interface{}{"fromHook": "yes"} + })) + if err != nil { + t.Error(err.Error()) + return + } + + request := httptest.NewRequestWithContext(t.Context(), http.MethodPost, "/graphql", strings.NewReader(`{ "query": "{ allUsers }" }`)) + responseRecorder := httptest.NewRecorder() + gateway.GraphQLHandler(responseRecorder, request) + + response := responseRecorder.Result() + defer response.Body.Close() + + result := map[string]interface{}{} + if err := json.NewDecoder(response.Body).Decode(&result); err != nil { + t.Error(err.Error()) + return + } + + // errors are reported as before and the hook's extension is present + assert.Nil(t, result["data"]) + assert.Len(t, result["errors"], 1) + assert.Equal(t, map[string]interface{}{"fromHook": "yes"}, result["extensions"]) +}