diff --git a/internal/spark/default_handler.go b/internal/spark/default_handler.go index 57efb88..ae81d56 100644 --- a/internal/spark/default_handler.go +++ b/internal/spark/default_handler.go @@ -21,6 +21,7 @@ package spark import ( "net/http" "net/url" + "strings" "github.com/gin-gonic/gin" log "github.com/okdp/spark-web-proxy/internal/logging" @@ -67,7 +68,16 @@ func (c DefaultSparkHandler) ModifyRequest(upstreamURL *url.URL) func(*http.Requ // ModifyResponse returns a function that rewrites redirect Location headers so // they remain relative when responses pass through the reverse proxy. -func (c DefaultSparkHandler) ModifyResponse() func(*http.Response) error { +// +// Only redirects targeting the upstream are rewritten. A redirect to another +// origin, such as an authentication filter sending the browser to an identity +// provider, is passed through unchanged: stripping its scheme and host would +// make the browser resolve it against the proxy. +func (c DefaultSparkHandler) ModifyResponse(upstreamURL *url.URL) func(*http.Response) error { + // Read here rather than in the closure: the director mutates upstreamURL + // while serving a request. + upstreamScheme, upstreamHost := upstreamURL.Scheme, upstreamURL.Host + return func(resp *http.Response) error { if resp.StatusCode == http.StatusFound { location := resp.Header.Get("Location") @@ -81,6 +91,17 @@ func (c DefaultSparkHandler) ModifyResponse() func(*http.Response) error { return nil } + // Already relative: there is nothing to strip. + if parsedURL.Host == "" { + return nil + } + + // A redirect to another origin keeps its scheme and host. + if !sameOrigin(parsedURL, upstreamScheme, upstreamHost) { + log.Debug("Location header '%s' targets another origin than the upstream '%s://%s', left untouched", location, upstreamScheme, upstreamHost) + return nil + } + parsedURL.Scheme = "" parsedURL.Host = "" @@ -94,3 +115,64 @@ func (c DefaultSparkHandler) ModifyResponse() func(*http.Response) error { return nil } } + +// sameOrigin reports whether the redirect target and the upstream share the +// same origin, that is the same scheme, host and port. "http://spark-history" +// and "https://spark-history" are distinct origins, so a redirect to the +// latter must not be rewritten as if it were the former. +func sameOrigin(target *url.URL, upstreamScheme, upstreamHost string) bool { + // A scheme-relative Location such as //spark-history:18080/jobs/ carries no + // scheme of its own and adopts the one it is served over, here the scheme + // of the upstream that produced the response. + if target.Scheme == "" { + target = &url.URL{Scheme: upstreamScheme, Host: target.Host} + } + + upstream := &url.URL{Scheme: upstreamScheme, Host: upstreamHost} + + return strings.EqualFold(target.Scheme, upstream.Scheme) && + strings.EqualFold(target.Hostname(), upstream.Hostname()) && + portOrDefault(target) == portOrDefault(upstream) +} + +// portOrDefault returns the port the URL is addressed on, falling back to the +// default port of its scheme when the authority leaves it out, so that +// "spark-history" and "spark-history:80" compare equal over http. Hostname and +// Port handle bracketed IPv6 authorities. +func portOrDefault(u *url.URL) string { + if port := u.Port(); port != "" { + return port + } + return schemePort(strings.ToLower(u.Scheme)) +} + +// schemePort returns the default port of a URI scheme, or an empty string when +// the scheme defines none. +// +// Only the HTTP and HTTPS default ports need special handling. +// +// A concrete Kubernetes example is a Spark History Service exposed as: +// +// spec: +// ports: +// - port: 80 +// targetPort: 18080 +// +// With an HTTP upstream, GetSparkHistoryBaseURL builds +// "http://spark-history-server:80". Jetty may omit the default port in an +// absolute redirect and return "http://spark-history-server/history/...". +// +// The equivalent case exists for an HTTPS upstream explicitly configured on +// port 443. Without default-port normalization, these redirects would not be +// recognised as targeting the upstream and an internal cluster address could +// be exposed to the browser. +func schemePort(scheme string) string { + switch scheme { + case "http": + return "80" + case "https": + return "443" + default: + return "" + } +} diff --git a/internal/spark/default_handler_test.go b/internal/spark/default_handler_test.go new file mode 100644 index 0000000..b2d1efb --- /dev/null +++ b/internal/spark/default_handler_test.go @@ -0,0 +1,217 @@ +/* + * Copyright 2025 The OKDP Authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package spark + +import ( + "net/http" + "net/url" + "os" + "testing" + + "github.com/okdp/spark-web-proxy/internal/config" + log "github.com/okdp/spark-web-proxy/internal/logging" +) + +// TestMain installs the global logger the handlers under test log through. It +// stays nil until SetupGlobalLogger is called, which otherwise only happens +// when the command starts the server. The error level keeps the debug lines +// the rewriting emits out of the test output. +func TestMain(m *testing.M) { + log.SetupGlobalLogger(config.Logging{Level: "error", Format: "console"}) + os.Exit(m.Run()) +} + +func TestDefaultSparkHandlerModifyResponse(t *testing.T) { + const upstream = "http://spark-history.spark.svc.cluster.local:18080" + + upstreamURL, err := url.Parse(upstream) + if err != nil { + t.Fatalf("invalid upstream URL: %v", err) + } + + tests := []struct { + name string + status int + location string + want string + }{ + { + name: "absolute redirect to the upstream is made relative", + status: http.StatusFound, + location: upstream + "/history/app-123/jobs/", + want: "/history/app-123/jobs/", + }, + { + name: "query string is preserved when the redirect is made relative", + status: http.StatusFound, + location: upstream + "/history/app-123/jobs/?status=running&sort=%2Bid", + want: "/history/app-123/jobs/?status=running&sort=%2Bid", + }, + { + // A redirect to the bare upstream carries no path, so stripping + // the authority leaves nothing behind. Spark addresses its + // redirects to a path, so the two cases below record that + // behaviour rather than guard against it. + name: "redirect to the bare upstream is left without a path", + status: http.StatusFound, + location: upstream, + want: "", + }, + { + name: "redirect to the bare upstream keeps only its query", + status: http.StatusFound, + location: upstream + "?status=completed", + want: "?status=completed", + }, + { + name: "relative redirect is left untouched", + status: http.StatusFound, + location: "/history/app-456/stages/", + want: "/history/app-456/stages/", + }, + { + // Regression test for issue #33: the auth filter answers an + // unauthenticated request with a 302 to the identity provider. + // Stripping its scheme and host made the browser resolve the path + // against the proxy, so the OIDC flow never started. + name: "cross origin redirect to the identity provider is preserved", + status: http.StatusFound, + location: "https://keycloak.example.com/realms/okdp/protocol/openid-connect/auth?client_id=spark-history&response_type=code&redirect_uri=https%3A%2F%2Fspark-web-proxy.example.com%2Fhome", + want: "https://keycloak.example.com/realms/okdp/protocol/openid-connect/auth?client_id=spark-history&response_type=code&redirect_uri=https%3A%2F%2Fspark-web-proxy.example.com%2Fhome", + }, + { + name: "same host but different scheme is a different origin", + status: http.StatusFound, + location: "https://spark-history.spark.svc.cluster.local:18080/home/", + want: "https://spark-history.spark.svc.cluster.local:18080/home/", + }, + { + name: "non redirect response is left untouched", + status: http.StatusOK, + location: "https://keycloak.example.com/realms/okdp", + want: "https://keycloak.example.com/realms/okdp", + }, + { + name: "status codes other than 302 are out of scope", + status: http.StatusMovedPermanently, + location: upstream + "/history/", + want: upstream + "/history/", + }, + } + + modify := DefaultSparkHandler{}.ModifyResponse(upstreamURL) + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + resp := &http.Response{StatusCode: tc.status, Header: http.Header{}} + resp.Header.Set("Location", tc.location) + + if err := modify(resp); err != nil { + t.Fatalf("ModifyResponse returned an error: %v", err) + } + + if got := resp.Header.Get("Location"); got != tc.want { + t.Errorf("Location = %q, want %q", got, tc.want) + } + }) + } +} + +func TestSameOrigin(t *testing.T) { + const upstream = "http://spark-history:18080" + + tests := []struct { + name string + target string + upstream string + want bool + }{ + { + name: "identical origin", + target: upstream + "/home", + upstream: upstream, + want: true, + }, + { + name: "http default port spelled out on one side only", + target: "http://spark-history:80/home", + upstream: "http://spark-history", + want: true, + }, + { + name: "https default port spelled out on one side only", + target: "https://spark-history:443/home", + upstream: "https://spark-history", + want: true, + }, + { + name: "host comparison is case insensitive", + target: "http://Spark-History:18080/home", + upstream: upstream, + want: true, + }, + { + name: "scheme relative target inherits the upstream scheme", + target: "//spark-history:18080/home", + upstream: upstream, + want: true, + }, + { + name: "ipv6 authority", + target: "http://[fd00::1]:18080/home", + upstream: "http://[fd00::1]:18080", + want: true, + }, + { + name: "different scheme is a different origin", + target: "https://spark-history:18080/home", + upstream: upstream, + want: false, + }, + { + name: "different port is a different origin", + target: "http://spark-history:8080/home", + upstream: upstream, + want: false, + }, + { + name: "different host is a different origin", + target: "https://idp.example.com/realms/okdp", + upstream: upstream, + want: false, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + target, err := url.Parse(tc.target) + if err != nil { + t.Fatalf("invalid target URL: %v", err) + } + + upstreamURL, err := url.Parse(tc.upstream) + if err != nil { + t.Fatalf("invalid upstream URL: %v", err) + } + + got := sameOrigin(target, upstreamURL.Scheme, upstreamURL.Host) + if got != tc.want { + t.Errorf("sameOrigin(%q, upstream %q) = %v, want %v", tc.target, tc.upstream, got, tc.want) + } + }) + } +} diff --git a/internal/spark/incomplete_apps_handler.go b/internal/spark/incomplete_apps_handler.go index 287636f..ca74ee8 100644 --- a/internal/spark/incomplete_apps_handler.go +++ b/internal/spark/incomplete_apps_handler.go @@ -71,7 +71,7 @@ func (c IncompleteAppsHandler) ModifyRequest(upstreamURL *url.URL) func(*http.Re // ModifyResponse returns a function that rewrites the Spark History incomplete // applications page when it contains the "No incomplete applications found!" message. -func (c IncompleteAppsHandler) ModifyResponse() func(*http.Response) error { +func (c IncompleteAppsHandler) ModifyResponse(_ *url.URL) func(*http.Response) error { return func(resp *http.Response) error { resp.TransferEncoding = []string{"identity"} // spark.history.ui.maxApplications = math.MaxInt32 diff --git a/internal/spark/proxy/handler.go b/internal/spark/proxy/handler.go index 8cd930c..5211744 100644 --- a/internal/spark/proxy/handler.go +++ b/internal/spark/proxy/handler.go @@ -36,7 +36,7 @@ import ( // request and response processing. type ReverseProxyHandler interface { ModifyRequest(upstreamURL *url.URL) func(*http.Request) - ModifyResponse() func(*http.Response) error + ModifyResponse(upstreamURL *url.URL) func(*http.Response) error } // DefaultErrorHandler returns a function that handles errors by logging the diff --git a/internal/spark/proxy/proxy.go b/internal/spark/proxy/proxy.go index 1cb8d46..555b520 100644 --- a/internal/spark/proxy/proxy.go +++ b/internal/spark/proxy/proxy.go @@ -34,7 +34,7 @@ type SparkReverseProxy struct { func NewSparkReverseProxy(c ReverseProxyHandler, upstreamURL *url.URL, appID string) *SparkReverseProxy { proxy := httputil.NewSingleHostReverseProxy(upstreamURL) proxy.Director = c.ModifyRequest(upstreamURL) - proxy.ModifyResponse = c.ModifyResponse() + proxy.ModifyResponse = c.ModifyResponse(upstreamURL) proxy.ErrorHandler = DefaultErrorHandler(appID) return &SparkReverseProxy{proxy, appID} }