diff --git a/pkg/infra/httpclient/httpclientprovider/host_redirect_validation_middleware.go b/pkg/infra/httpclient/httpclientprovider/host_redirect_validation_middleware.go new file mode 100644 index 00000000000..42bad8708b4 --- /dev/null +++ b/pkg/infra/httpclient/httpclientprovider/host_redirect_validation_middleware.go @@ -0,0 +1,33 @@ +package httpclientprovider + +import ( + "net/http" + + "github.com/grafana/grafana/pkg/models" + + sdkhttpclient "github.com/grafana/grafana-plugin-sdk-go/backend/httpclient" +) + +const HostRedirectValidationMiddlewareName = "host-redirect-validation" + +func RedirectLimitMiddleware(reqValidator models.PluginRequestValidator) sdkhttpclient.Middleware { + return sdkhttpclient.NamedMiddlewareFunc(HostRedirectValidationMiddlewareName, func(opts sdkhttpclient.Options, next http.RoundTripper) http.RoundTripper { + return sdkhttpclient.RoundTripperFunc(func(req *http.Request) (*http.Response, error) { + res, err := next.RoundTrip(req) + if err != nil { + return nil, err + } + if res.StatusCode >= 300 && res.StatusCode < 400 { + location, err := res.Location() + if err != nil { + return nil, err + } + + if err := reqValidator.Validate(location.String(), nil); err != nil { + return nil, err + } + } + return res, err + }) + }) +} diff --git a/pkg/infra/httpclient/httpclientprovider/http_client_provider.go b/pkg/infra/httpclient/httpclientprovider/http_client_provider.go index 26b9b79c4b4..a011258afc2 100644 --- a/pkg/infra/httpclient/httpclientprovider/http_client_provider.go +++ b/pkg/infra/httpclient/httpclientprovider/http_client_provider.go @@ -5,6 +5,8 @@ import ( "net/http" "time" + "github.com/grafana/grafana/pkg/models" + sdkhttpclient "github.com/grafana/grafana-plugin-sdk-go/backend/httpclient" "github.com/grafana/grafana/pkg/infra/log" "github.com/grafana/grafana/pkg/infra/metrics/metricutil" @@ -16,7 +18,7 @@ import ( var newProviderFunc = sdkhttpclient.NewProvider // New creates a new HTTP client provider with pre-configured middlewares. -func New(cfg *setting.Cfg, tracer tracing.Tracer) *sdkhttpclient.Provider { +func New(cfg *setting.Cfg, validator models.PluginRequestValidator, tracer tracing.Tracer) *sdkhttpclient.Provider { logger := log.New("httpclient") userAgent := fmt.Sprintf("Grafana/%s", cfg.BuildVersion) @@ -27,6 +29,7 @@ func New(cfg *setting.Cfg, tracer tracing.Tracer) *sdkhttpclient.Provider { sdkhttpclient.BasicAuthenticationMiddleware(), sdkhttpclient.CustomHeadersMiddleware(), ResponseLimitMiddleware(cfg.ResponseLimit), + RedirectLimitMiddleware(validator), } if cfg.SigV4AuthEnabled {