Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 18 additions & 0 deletions gateway.go
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,8 @@ type Gateway struct {
locationPriorities []string
preExecutionHook PreExecutionStepHook
postExecutionHook PostExecutionStepHook
preOperationHook PreOperationHook
postOperationHook PostOperationHook
Comment on lines +38 to +39

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@OlfaKaroui you once implemented pre- and post execution hooks into the gateway as well, didn't you? Was that in the same area or somewhere else?


// group up the list of middlewares at startup to avoid it during execution
requestMiddlewares []graphql.NetworkMiddleware
Expand Down Expand Up @@ -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) {
Expand Down
39 changes: 32 additions & 7 deletions http.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
}
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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{}{
Expand All @@ -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)
}
Expand Down
170 changes: 170 additions & 0 deletions http_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"])
}
Loading