From d11662aa0cda5f72fac4932bee60b3e2ceff02ac Mon Sep 17 00:00:00 2001 From: Florian Beeres Date: Tue, 8 Sep 2026 14:40:42 +0200 Subject: [PATCH] Add per-operation hooks to GraphQLHandler The handler already exposes pre/post hooks per plan step, but nothing runs once per operation, and the response payload is assembled entirely inside executeRequest. Users who need to touch the payload, e.g. to add response extensions gathered during execution, had to reimplement the handler. WithPreOperationHook runs after the RequestContext is built and before planning; it may replace rc.Context, which is how per-operation state reaches the executor and the queryers even in batch mode, where every operation otherwise shares the incoming request context. WithPostOperationHook runs as the last step before the operation's payload is written, on success and on failure. It sees the payload exactly as the gateway would have written it, including the gateway's own extensions, and may add or modify keys. The gateway itself makes no assumptions about what the hook does with them. Co-Authored-By: Claude Fable 5.1 --- gateway.go | 18 ++++++ http.go | 39 +++++++++--- http_test.go | 170 +++++++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 220 insertions(+), 7 deletions(-) diff --git a/gateway.go b/gateway.go index ed5004cd..11004762 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 3e39ccaa..2281923c 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 0280a8a9..73460f2d 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"]) +}