diff --git a/cli/apiserver/provider.go b/cli/apiserver/provider.go index 07afe6803..b4fd4b394 100644 --- a/cli/apiserver/provider.go +++ b/cli/apiserver/provider.go @@ -4,9 +4,11 @@ import ( "bytes" "encoding/json" "fmt" + "io" "io/ioutil" "net/http" "net/url" + "strings" "time" "github.com/up9inc/mizu/shared/kubernetes" @@ -57,10 +59,8 @@ func (provider *Provider) TestConnection() error { func (provider *Provider) isReachable() (bool, error) { echoUrl := fmt.Sprintf("%s/echo", provider.url) - if response, err := provider.client.Get(echoUrl); err != nil { + if _, err := provider.get(echoUrl); err != nil { return false, err - } else if response.StatusCode != 200 { - return false, fmt.Errorf("invalid status code %v", response.StatusCode) } else { return true, nil } @@ -72,10 +72,8 @@ func (provider *Provider) ReportTapperStatus(tapperStatus shared.TapperStatus) e if jsonValue, err := json.Marshal(tapperStatus); err != nil { return fmt.Errorf("failed Marshal the tapper status %w", err) } else { - if response, err := provider.client.Post(tapperStatusUrl, "application/json", bytes.NewBuffer(jsonValue)); err != nil { + if _, err := provider.post(tapperStatusUrl, "application/json", bytes.NewBuffer(jsonValue)); err != nil { return fmt.Errorf("failed sending to API server the tapped pods %w", err) - } else if response.StatusCode != 200 { - return fmt.Errorf("failed sending to API server the tapper status, response status code %v", response.StatusCode) } else { logger.Log.Debugf("Reported to server API about tapper status: %v", tapperStatus) return nil @@ -91,10 +89,8 @@ func (provider *Provider) ReportTappedPods(pods []core.Pod) error { if jsonValue, err := json.Marshal(podInfos); err != nil { return fmt.Errorf("failed Marshal the tapped pods %w", err) } else { - if response, err := provider.client.Post(tappedPodsUrl, "application/json", bytes.NewBuffer(jsonValue)); err != nil { + if _, err := provider.post(tappedPodsUrl, "application/json", bytes.NewBuffer(jsonValue)); err != nil { return fmt.Errorf("failed sending to API server the tapped pods %w", err) - } else if response.StatusCode != 200 { - return fmt.Errorf("failed sending to API server the tapped pods, response status code %v", response.StatusCode) } else { logger.Log.Debugf("Reported to server API about %d taped pods successfully", len(podInfos)) return nil @@ -105,11 +101,9 @@ func (provider *Provider) ReportTappedPods(pods []core.Pod) error { func (provider *Provider) GetGeneralStats() (map[string]interface{}, error) { generalStatsUrl := fmt.Sprintf("%s/status/general", provider.url) - response, requestErr := provider.client.Get(generalStatsUrl) + response, requestErr := provider.get(generalStatsUrl) if requestErr != nil { return nil, fmt.Errorf("failed to get general stats for telemetry, err: %w", requestErr) - } else if response.StatusCode != 200 { - return nil, fmt.Errorf("failed to get general stats for telemetry, status code: %v", response.StatusCode) } defer response.Body.Close() @@ -132,7 +126,7 @@ func (provider *Provider) GetVersion() (string, error) { Method: http.MethodGet, URL: versionUrl, } - statusResp, err := provider.client.Do(req) + statusResp, err := provider.do(req) if err != nil { return "", err } @@ -145,3 +139,40 @@ func (provider *Provider) GetVersion() (string, error) { return versionResponse.Ver, nil } + +// When err is nil, resp always contains a non-nil resp.Body. +// Caller should close resp.Body when done reading from it. +func (provider *Provider) get(url string) (*http.Response, error) { + return provider.checkError(provider.client.Get(url)) +} + +// When err is nil, resp always contains a non-nil resp.Body. +// Caller should close resp.Body when done reading from it. +func (provider *Provider) post(url, contentType string, body io.Reader) (*http.Response, error) { + return provider.checkError(provider.client.Post(url, contentType, body)) +} + +// When err is nil, resp always contains a non-nil resp.Body. +// Caller should close resp.Body when done reading from it. +func (provider *Provider) do(req *http.Request) (*http.Response, error) { + return provider.checkError(provider.client.Do(req)) +} + +func (provider *Provider) checkError(response *http.Response, errInOperation error) (*http.Response, error) { + if (errInOperation != nil) { + return response, errInOperation + // Check only if status != 200 (and not status >= 300). Agent APIs return only 200 on success. + } else if response.StatusCode != http.StatusOK { + body, err := ioutil.ReadAll(response.Body) + response.Body.Close() + response.Body = io.NopCloser(bytes.NewBuffer(body)) // rewind + if err != nil { + return response, err + } + + errorMsg := strings.ReplaceAll((string(body)), "\n", ";") + return response, fmt.Errorf("got response with status code: %d, body: %s", response.StatusCode, errorMsg) + } + + return response, nil +}