From 22ee79fcb8731acbe6f4144a8e0e7d4de57be54b Mon Sep 17 00:00:00 2001 From: Rajat Vig Date: Wed, 22 Dec 2021 14:13:36 +0000 Subject: [PATCH] Add the copyheaders code back --- pkg/api/echo.go | 176 +++++++++++++++++++++++++++--------------------- 1 file changed, 99 insertions(+), 77 deletions(-) diff --git a/pkg/api/echo.go b/pkg/api/echo.go index 8b82444..5126219 100644 --- a/pkg/api/echo.go +++ b/pkg/api/echo.go @@ -1,18 +1,18 @@ package api import ( - "bytes" - "context" - "fmt" - "io/ioutil" - "net/http" - "net/http/httptrace" - "sync" + "bytes" + "context" + "fmt" + "io/ioutil" + "net/http" + "net/http/httptrace" + "sync" - "github.com/stefanprodan/podinfo/pkg/version" - "go.opentelemetry.io/contrib/instrumentation/net/http/httptrace/otelhttptrace" - "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" - "go.uber.org/zap" + "github.com/stefanprodan/podinfo/pkg/version" + "go.opentelemetry.io/contrib/instrumentation/net/http/httptrace/otelhttptrace" + "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" + "go.uber.org/zap" ) // Echo godoc @@ -24,83 +24,105 @@ import ( // @Router /api/echo [post] // @Success 202 {object} api.MapResponse func (s *Server) echoHandler(w http.ResponseWriter, r *http.Request) { - ctx, span := s.tracer.Start(r.Context(), "echoHandler") - defer span.End() + ctx, span := s.tracer.Start(r.Context(), "echoHandler") + defer span.End() - body, err := ioutil.ReadAll(r.Body) - if err != nil { - s.logger.Error("reading the request body failed", zap.Error(err)) - s.ErrorResponse(w, r, span, "invalid request body", http.StatusBadRequest) - return - } - defer r.Body.Close() + body, err := ioutil.ReadAll(r.Body) + if err != nil { + s.logger.Error("reading the request body failed", zap.Error(err)) + s.ErrorResponse(w, r, span, "invalid request body", http.StatusBadRequest) + return + } + defer r.Body.Close() - client := http.Client{Transport: otelhttp.NewTransport(http.DefaultTransport)} + client := http.Client{Transport: otelhttp.NewTransport(http.DefaultTransport)} - if len(s.config.BackendURL) > 0 { - result := make([]string, len(s.config.BackendURL)) - var wg sync.WaitGroup - wg.Add(len(s.config.BackendURL)) - for i, b := range s.config.BackendURL { - go func(index int, backend string) { - defer wg.Done() + if len(s.config.BackendURL) > 0 { + result := make([]string, len(s.config.BackendURL)) + var wg sync.WaitGroup + wg.Add(len(s.config.BackendURL)) + for i, b := range s.config.BackendURL { + go func(index int, backend string) { + defer wg.Done() - ctx = httptrace.WithClientTrace(ctx, otelhttptrace.NewClientTrace(ctx)) - ctx, cancel := context.WithTimeout(ctx, s.config.HttpClientTimeout) - defer cancel() + ctx = httptrace.WithClientTrace(ctx, otelhttptrace.NewClientTrace(ctx)) + ctx, cancel := context.WithTimeout(ctx, s.config.HttpClientTimeout) + defer cancel() - backendReq, err := http.NewRequestWithContext(ctx, "POST", backend, bytes.NewReader(body)) - if err != nil { - s.logger.Error("backend call failed", zap.Error(err), zap.String("url", backend)) - return - } + backendReq, err := http.NewRequestWithContext(ctx, "POST", backend, bytes.NewReader(body)) + if err != nil { + s.logger.Error("backend call failed", zap.Error(err), zap.String("url", backend)) + return + } - backendReq.Header.Set("X-API-Version", version.VERSION) - backendReq.Header.Set("X-API-Revision", version.REVISION) + // forward headers + copyTracingHeaders(r, backendReq) - // call backend - resp, err := client.Do(backendReq) - if err != nil { - s.logger.Error("backend call failed", zap.Error(err), zap.String("url", backend)) - result[index] = fmt.Sprintf("backend %v call failed %v", backend, err) - return - } - defer resp.Body.Close() + backendReq.Header.Set("X-API-Version", version.VERSION) + backendReq.Header.Set("X-API-Revision", version.REVISION) - // copy error status from backend and exit - if resp.StatusCode >= 400 { - s.logger.Error("backend call failed", zap.Int("status", resp.StatusCode), zap.String("url", backend)) - result[index] = fmt.Sprintf("backend %v response status code %v", backend, resp.StatusCode) - return - } + // call backend + resp, err := client.Do(backendReq) + if err != nil { + s.logger.Error("backend call failed", zap.Error(err), zap.String("url", backend)) + result[index] = fmt.Sprintf("backend %v call failed %v", backend, err) + return + } + defer resp.Body.Close() - // forward the received body - rbody, err := ioutil.ReadAll(resp.Body) - if err != nil { - s.logger.Error( - "reading the backend request body failed", - zap.Error(err), - zap.String("url", backend)) - result[index] = fmt.Sprintf("backend %v call failed %v", backend, err) - return - } + // copy error status from backend and exit + if resp.StatusCode >= 400 { + s.logger.Error("backend call failed", zap.Int("status", resp.StatusCode), zap.String("url", backend)) + result[index] = fmt.Sprintf("backend %v response status code %v", backend, resp.StatusCode) + return + } - s.logger.Debug( - "payload received from backend", - zap.String("response", string(rbody)), - zap.String("url", backend)) + // forward the received body + rbody, err := ioutil.ReadAll(resp.Body) + if err != nil { + s.logger.Error( + "reading the backend request body failed", + zap.Error(err), + zap.String("url", backend)) + result[index] = fmt.Sprintf("backend %v call failed %v", backend, err) + return + } - result[index] = string(rbody) - }(i, b) - } - wg.Wait() + s.logger.Debug( + "payload received from backend", + zap.String("response", string(rbody)), + zap.String("url", backend)) - w.Header().Set("X-Color", s.config.UIColor) - s.JSONResponse(w, r, result) + result[index] = string(rbody) + }(i, b) + } + wg.Wait() - } else { - w.Header().Set("X-Color", s.config.UIColor) - w.WriteHeader(http.StatusAccepted) - w.Write(body) - } + w.Header().Set("X-Color", s.config.UIColor) + s.JSONResponse(w, r, result) + + } else { + w.Header().Set("X-Color", s.config.UIColor) + w.WriteHeader(http.StatusAccepted) + w.Write(body) + } +} + +func copyTracingHeaders(from *http.Request, to *http.Request) { + headers := []string{ + "x-request-id", + "x-b3-traceid", + "x-b3-spanid", + "x-b3-parentspanid", + "x-b3-sampled", + "x-b3-flags", + "x-ot-span-context", + } + + for i := range headers { + headerValue := from.Header.Get(headers[i]) + if len(headerValue) > 0 { + to.Header.Set(headers[i], headerValue) + } + } }