diff --git a/pkg/api/client.go b/pkg/api/client.go index dea61a7bd..ff628295b 100644 --- a/pkg/api/client.go +++ b/pkg/api/client.go @@ -2,15 +2,16 @@ package api import ( "context" - "encoding/json" + "errors" + "fmt" "net/http" rawclient "github.com/Yamashou/gqlgenc/clientv2" - "github.com/pkg/errors" "github.com/pluralsh/gqlclient" "github.com/pluralsh/plural-cli/pkg/config" "github.com/pluralsh/plural-cli/pkg/utils" + clierrors "github.com/pluralsh/plural-cli/pkg/utils/errors" ) type authedTransport struct { @@ -79,7 +80,7 @@ func FromConfig(conf *config.Config) Client { } return &client{ - pluralClient: gqlclient.NewClient(&httpClient, conf.Url(), nil), + pluralClient: gqlclient.NewClient(&httpClient, conf.Url(), nil, clierrors.GraphQLInterceptor), config: *conf, ctx: context.Background(), httpClient: &httpClient, @@ -91,23 +92,9 @@ func GetErrorResponse(err error, methodName string) error { return nil } utils.LogError().Println(err) - errResponse := &rawclient.ErrorResponse{} - newErr := json.Unmarshal([]byte(err.Error()), errResponse) - if newErr != nil { + var errResponse *rawclient.ErrorResponse + if !errors.As(err, &errResponse) { return err } - - errList := errors.New(methodName) - if errResponse.GqlErrors != nil { - for _, err := range *errResponse.GqlErrors { - errList = errors.Wrap(errList, err.Message) - } - errList = errors.Wrap(errList, "GraphQL error") - } - if errResponse.NetworkError != nil { - errList = errors.Wrap(errList, errResponse.NetworkError.Message) - errList = errors.Wrap(errList, "Network error") - } - - return errList + return fmt.Errorf("%s: %w", methodName, clierrors.GraphQL(err)) } diff --git a/pkg/api/client_test.go b/pkg/api/client_test.go new file mode 100644 index 000000000..be98bbd04 --- /dev/null +++ b/pkg/api/client_test.go @@ -0,0 +1,136 @@ +package api_test + +import ( + "context" + "errors" + "fmt" + "io" + "net/http" + "strings" + "testing" + + "github.com/Yamashou/gqlgenc/clientv2" + consoleclient "github.com/pluralsh/console/go/client" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/vektah/gqlparser/v2/gqlerror" + + "github.com/pluralsh/plural-cli/pkg/api" + "github.com/pluralsh/plural-cli/pkg/config" + "github.com/pluralsh/plural-cli/pkg/console" + consoleerrors "github.com/pluralsh/plural-cli/pkg/console/errors" +) + +type roundTripperFunc func(*http.Request) (*http.Response, error) + +func (f roundTripperFunc) RoundTrip(req *http.Request) (*http.Response, error) { + return f(req) +} + +func TestClientErrorMessages(t *testing.T) { + originalTransport := http.DefaultTransport + t.Cleanup(func() { http.DefaultTransport = originalTransport }) + + for _, tt := range []struct { + name string + status int + body string + want string + }{ + {"GraphQL", 200, `{"errors":[{"message":"handle has already been taken","path":["createCluster"],"extensions":{"code":"invalid"}},{"message":"name is required"}]}`, "handle has already been taken; name is required"}, + {"GraphQL with HTTP failure", 400, `{"errors":[{"message":"handle has already been taken"}]}`, "handle has already been taken"}, + {"HTTP failure", 502, `upstream unavailable`, "HTTP 502 Bad Gateway"}, + } { + t.Run(tt.name, func(t *testing.T) { + http.DefaultTransport = roundTripperFunc(func(req *http.Request) (*http.Response, error) { + if req.Header.Get("Authorization") == "Token token" { + assert.NotEmpty(t, req.URL.Query().Get("documentId"), "persisted query interceptor must still run") + } + return &http.Response{StatusCode: tt.status, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(tt.body))}, nil + }) + apiClient := api.FromConfig(&config.Config{Endpoint: "api.example.com", Token: "token"}) + consoleClient, err := console.NewConsoleClient("token", "https://console.example.com") + require.NoError(t, err) + + _, err = apiClient.Me() + assert.EqualError(t, err, tt.want) + _, err = consoleClient.CreateCluster(consoleclient.ClusterAttributes{}) + assert.EqualError(t, err, tt.want) + assert.Equal(t, tt.status != 502, consoleerrors.Like(err, "handle")) + _, err = consoleClient.ListClusters() + assert.EqualError(t, err, "ListClusters: "+tt.want) + var response *clientv2.ErrorResponse + assert.ErrorAs(t, err, &response) + assert.Equal(t, tt.status != 502, consoleerrors.Like(err, "handle")) + }) + } +} + +func TestGetErrorResponse(t *testing.T) { + response := &clientv2.ErrorResponse{GqlErrors: &gqlerror.List{{Message: "permission denied"}}} + wrapped := fmt.Errorf("fetching cluster: %w", response) + got := api.GetErrorResponse(wrapped, "GetCluster") + assert.EqualError(t, got, "GetCluster: fetching cluster: permission denied") + assert.ErrorIs(t, got, wrapped) + var preserved *clientv2.ErrorResponse + assert.ErrorAs(t, got, &preserved) + assert.Same(t, response, preserved) + assert.Nil(t, api.GetErrorResponse(nil, "GetCluster")) + for _, err := range []error{errors.New("connection refused"), errors.New(`{"message":"unrelated"}`)} { + assert.Same(t, err, api.GetErrorResponse(err, "GetCluster")) + } +} + +func TestClientsPreserveTransportErrors(t *testing.T) { + originalTransport := http.DefaultTransport + t.Cleanup(func() { http.DefaultTransport = originalTransport }) + + for _, transportErr := range []error{errors.New("connection refused"), context.Canceled, context.DeadlineExceeded} { + t.Run(transportErr.Error(), func(t *testing.T) { + http.DefaultTransport = roundTripperFunc(func(*http.Request) (*http.Response, error) { + return nil, transportErr + }) + apiClient := api.FromConfig(&config.Config{Endpoint: "api.example.com", Token: "token"}) + consoleClient, err := console.NewConsoleClient("token", "https://console.example.com") + require.NoError(t, err) + + _, apiErr := apiClient.Me() + _, consoleErr := consoleClient.ListClusters() + for _, err := range []error{apiErr, consoleErr} { + require.Error(t, err) + assert.ErrorIs(t, err, transportErr) + assert.Contains(t, err.Error(), transportErr.Error()) + var response *clientv2.ErrorResponse + assert.False(t, errors.As(err, &response)) + assert.Same(t, err, api.GetErrorResponse(err, "ListClusters")) + } + }) + } +} + +func TestClientsSuccessfulResponses(t *testing.T) { + originalTransport := http.DefaultTransport + t.Cleanup(func() { http.DefaultTransport = originalTransport }) + http.DefaultTransport = roundTripperFunc(func(req *http.Request) (*http.Response, error) { + body := `{"data":{"me":{"id":"user-1","email":"user@example.com"}}}` + if req.URL.Host == "console.example.com" { + body = `{"data":{"clusters":{"edges":[]}}}` + } + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(body))}, nil + }) + + apiClient := api.FromConfig(&config.Config{Endpoint: "api.example.com", Token: "token"}) + me, err := apiClient.Me() + require.NoError(t, err) + require.NotNil(t, me) + assert.Equal(t, "user-1", me.Id) + assert.Equal(t, "user@example.com", me.Email) + + consoleClient, err := console.NewConsoleClient("token", "https://console.example.com") + require.NoError(t, err) + clusters, err := consoleClient.ListClusters() + require.NoError(t, err) + require.NotNil(t, clusters) + require.NotNil(t, clusters.Clusters) + assert.Empty(t, clusters.Clusters.Edges) +} diff --git a/pkg/console/console.go b/pkg/console/console.go index 0cc271406..d578609a4 100644 --- a/pkg/console/console.go +++ b/pkg/console/console.go @@ -8,6 +8,7 @@ import ( "time" consoleclient "github.com/pluralsh/console/go/client" + clierrors "github.com/pluralsh/plural-cli/pkg/utils/errors" ) type consoleClient struct { @@ -93,7 +94,7 @@ func NewConsoleClient(token, url string) (ConsoleClient, error) { url: NormalizeUrl(url), extUrl: NormalizeExtUrl(url), token: token, - client: consoleclient.NewClient(&httpClient, NormalizeUrl(url), nil, consoleclient.PersistedQueryInterceptor), + client: consoleclient.NewClient(&httpClient, NormalizeUrl(url), nil, clierrors.GraphQLInterceptor, consoleclient.PersistedQueryInterceptor), ctx: context.Background(), }, nil } diff --git a/pkg/console/errors/like.go b/pkg/console/errors/like.go index f13b3d013..963fefb07 100644 --- a/pkg/console/errors/like.go +++ b/pkg/console/errors/like.go @@ -22,8 +22,11 @@ func Like(err error, msg string) bool { } func isLike(err *client.ErrorResponse, msg string) bool { + if err == nil || err.GqlErrors == nil { + return false + } for _, g := range *err.GqlErrors { - if strings.Contains(g.Message, msg) { + if g != nil && strings.Contains(g.Message, msg) { return true } } diff --git a/pkg/console/errors/like_test.go b/pkg/console/errors/like_test.go new file mode 100644 index 000000000..4de4a4011 --- /dev/null +++ b/pkg/console/errors/like_test.go @@ -0,0 +1,43 @@ +package errors_test + +import ( + "errors" + "fmt" + "testing" + + "github.com/Yamashou/gqlgenc/clientv2" + "github.com/stretchr/testify/assert" + "github.com/vektah/gqlparser/v2/gqlerror" + + consoleerrors "github.com/pluralsh/plural-cli/pkg/console/errors" + clierrors "github.com/pluralsh/plural-cli/pkg/utils/errors" +) + +func TestLike(t *testing.T) { + response := &clientv2.ErrorResponse{GqlErrors: &gqlerror.List{ + nil, + {Message: "name is required"}, + {Message: "handle has already been taken"}, + }} + var nilResponse *clientv2.ErrorResponse + for _, tt := range []struct { + name string + err error + msg string + want bool + }{ + {name: "nil error", msg: "handle"}, + {name: "typed nil response", err: nilResponse, msg: "handle"}, + {name: "unrelated error with matching text", err: errors.New("handle has already been taken"), msg: "handle"}, + {name: "network error only", err: &clientv2.ErrorResponse{NetworkError: &clientv2.HTTPError{Code: 502}}, msg: "handle"}, + {name: "empty list", err: &clientv2.ErrorResponse{GqlErrors: &gqlerror.List{}}, msg: "handle"}, + {name: "nil entries", err: &clientv2.ErrorResponse{GqlErrors: &gqlerror.List{nil}}, msg: "handle"}, + {name: "match after nil and nonmatching entries", err: response, msg: "handle", want: true}, + {name: "no match", err: response, msg: "permission denied"}, + {name: "wrapped formatted response", err: fmt.Errorf("creating cluster: %w", clierrors.GraphQL(response)), msg: "handle", want: true}, + } { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, consoleerrors.Like(tt.err, tt.msg)) + }) + } +} diff --git a/pkg/utils/errors/graphql.go b/pkg/utils/errors/graphql.go new file mode 100644 index 000000000..5658293ad --- /dev/null +++ b/pkg/utils/errors/graphql.go @@ -0,0 +1,65 @@ +package errors + +import ( + "context" + "errors" + "fmt" + "net/http" + "strings" + + "github.com/Yamashou/gqlgenc/clientv2" +) + +type graphQLError struct { + err error + message string +} + +func (e *graphQLError) Error() string { return e.message } +func (e *graphQLError) Unwrap() error { return e.err } + +// GraphQL presents server error messages while retaining the response for errors.As. +// Errors without a usable message are returned unchanged. +func GraphQL(err error) error { + var formatted *graphQLError + if errors.As(err, &formatted) { + return err + } + + var response *clientv2.ErrorResponse + if !errors.As(err, &response) || response == nil { + return err + } + + var messages []string + if response.GqlErrors != nil { + for _, gqlErr := range *response.GqlErrors { + if gqlErr != nil && strings.TrimSpace(gqlErr.Message) != "" { + messages = append(messages, gqlErr.Message) + } + } + } + + // Non-2xx responses can contain the same GraphQL errors in the raw HTTP + // body. Prefer the parsed messages and avoid printing that body again. + if len(messages) == 0 && response.NetworkError != nil { + if code := response.NetworkError.Code; code != 0 { + messages = append(messages, strings.TrimSpace(fmt.Sprintf("HTTP %d %s", code, http.StatusText(code)))) + } else if message := strings.TrimSpace(response.NetworkError.Message); message != "" { + messages = append(messages, message) + } + } + if len(messages) == 0 { + return err + } + + return &graphQLError{ + err: err, + message: strings.Replace(err.Error(), response.Error(), strings.Join(messages, "; "), 1), + } +} + +// GraphQLInterceptor formats errors after all other request interceptors finish. +func GraphQLInterceptor(ctx context.Context, req *http.Request, info *clientv2.GQLRequestInfo, res any, next clientv2.RequestInterceptorFunc) error { + return GraphQL(next(ctx, req, info, res)) +} diff --git a/pkg/utils/errors/graphql_test.go b/pkg/utils/errors/graphql_test.go new file mode 100644 index 000000000..7e42af306 --- /dev/null +++ b/pkg/utils/errors/graphql_test.go @@ -0,0 +1,60 @@ +package errors_test + +import ( + "errors" + "fmt" + "testing" + + "github.com/Yamashou/gqlgenc/clientv2" + "github.com/stretchr/testify/assert" + "github.com/vektah/gqlparser/v2/gqlerror" + + clierrors "github.com/pluralsh/plural-cli/pkg/utils/errors" +) + +func TestGraphQL(t *testing.T) { + response := &clientv2.ErrorResponse{GqlErrors: &gqlerror.List{ + {Message: "handle has already been taken"}, + nil, + {Message: " "}, + {Message: "name is required"}, + }} + want := "handle has already been taken; name is required" + tests := []struct { + name string + err error + want string + }{ + {name: "multiple messages", err: response, want: want}, + {name: "wrapped response", err: fmt.Errorf("creating cluster: %w", response), want: "creating cluster: " + want}, + {name: "already formatted", err: fmt.Errorf("creating cluster: %w", clierrors.GraphQL(response)), want: "creating cluster: " + want}, + {name: "HTTP error", err: &clientv2.ErrorResponse{NetworkError: &clientv2.HTTPError{Code: 502, Message: "Response body Bad gateway"}}, want: "HTTP 502 Bad Gateway"}, + {name: "unknown HTTP status", err: &clientv2.ErrorResponse{NetworkError: &clientv2.HTTPError{Code: 599}}, want: "HTTP 599"}, + {name: "network message without status", err: &clientv2.ErrorResponse{NetworkError: &clientv2.HTTPError{Message: "connection failed"}}, want: "connection failed"}, + {name: "GraphQL over HTTP error", err: &clientv2.ErrorResponse{GqlErrors: response.GqlErrors, NetworkError: &clientv2.HTTPError{Code: 400, Message: "Response body raw JSON"}}, want: want}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := clierrors.GraphQL(tt.err) + assert.EqualError(t, got, tt.want) + assert.ErrorIs(t, got, tt.err) + var original, preserved *clientv2.ErrorResponse + assert.ErrorAs(t, tt.err, &original) + assert.ErrorAs(t, got, &preserved) + assert.Same(t, original, preserved) + }) + } +} + +func TestGraphQLUnchanged(t *testing.T) { + for _, err := range []error{ + nil, + errors.New("connection refused"), + errors.New(`{"message":"unrelated JSON error"}`), + &clientv2.ErrorResponse{}, + &clientv2.ErrorResponse{GqlErrors: &gqlerror.List{}}, + &clientv2.ErrorResponse{GqlErrors: &gqlerror.List{nil, {Message: ""}}}, + } { + assert.Equal(t, err, clierrors.GraphQL(err)) + } +}