diff --git a/cmd/autocar/server.go b/cmd/autocar/server.go index f0e199d..566da3b 100644 --- a/cmd/autocar/server.go +++ b/cmd/autocar/server.go @@ -31,6 +31,7 @@ func runServer(parent context.Context, args []string) error { disableFallback := fs.Bool("disable-tcp-fallback", false, "disable the TCP/TLS fallback listener") coverRoot := fs.String("cover-root", "", "web protocol: directory served as the public cover origin") coverUpstream := fs.String("cover-upstream", "", "web protocol: fixed http(s) origin used as the public cover") + coverPublicOrigin := fs.String("cover-public-origin", "", "web upstream only: fixed public HTTPS origin for same-origin website sessions (opt-in)") certFile := fs.String("cert", "", "server certificate PEM (required)") keyFile := fs.String("key", "", "server private key PEM (required)") clientCAFile := fs.String("client-ca", "", "optional PEM CA that enables mandatory mTLS") @@ -64,6 +65,19 @@ func runServer(parent context.Context, args []string) error { if err := validateServerProtocolOptions(serverProtocol, *clientCAFile, *coverRoot, *coverUpstream, *disableFallback); err != nil { return err } + var coverHandler http.Handler + if *coverPublicOrigin != "" { + if serverProtocol != "web" || strings.TrimSpace(*coverRoot) != "" || strings.TrimSpace(*coverUpstream) == "" { + return errors.New("--cover-public-origin requires --protocol=web and --cover-upstream, without --cover-root") + } + // This constructor validates only local URL configuration; it performs + // no DNS, upstream requests or listening, including during --check. + var err error + coverHandler, err = buildCoverHandlerWithPublicOrigin(*coverRoot, *coverUpstream, *coverPublicOrigin) + if err != nil { + return err + } + } if *certFile == "" || *keyFile == "" { return errors.New("--cert and --key are required") } @@ -160,8 +174,7 @@ func runServer(parent context.Context, args []string) error { if err != nil { return err } - var coverHandler http.Handler - if serverProtocol == "web" { + if serverProtocol == "web" && coverHandler == nil { coverHandler, err = buildCoverHandler(*coverRoot, *coverUpstream) if err != nil { return err @@ -337,7 +350,14 @@ func validateServerProtocolOptions(protocolMode, clientCAFile, coverRoot, coverU } func buildCoverHandler(root, upstream string) (http.Handler, error) { + return buildCoverHandlerWithPublicOrigin(root, upstream, "") +} + +func buildCoverHandlerWithPublicOrigin(root, upstream, publicOrigin string) (http.Handler, error) { if strings.TrimSpace(root) != "" { + if publicOrigin != "" { + return nil, errors.New("--cover-public-origin requires --cover-upstream, without --cover-root") + } handler, err := cover.NewStaticHandler(root) if err != nil { return nil, fmt.Errorf("configure static cover: %w", err) @@ -348,7 +368,16 @@ func buildCoverHandler(root, upstream string) (http.Handler, error) { if err != nil { return nil, fmt.Errorf("parse --cover-upstream: %w", err) } - handler, err := cover.NewReverseProxyHandler(origin, nil) + var handler http.Handler + if publicOrigin == "" { + handler, err = cover.NewReverseProxyHandler(origin, nil) + } else { + public, parseErr := url.Parse(strings.TrimSpace(publicOrigin)) + if parseErr != nil { + return nil, errors.New("parse --cover-public-origin: invalid URL") + } + handler, err = cover.NewReverseProxyHandlerWithPublicOrigin(origin, public, nil) + } if err != nil { return nil, fmt.Errorf("configure upstream cover: %w", err) } diff --git a/cmd/autocar/web_public_origin_test.go b/cmd/autocar/web_public_origin_test.go new file mode 100644 index 0000000..9fb5550 --- /dev/null +++ b/cmd/autocar/web_public_origin_test.go @@ -0,0 +1,85 @@ +package main + +import ( + "context" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" +) + +func TestRunServerPublicOriginValidationBeforeCredentials(t *testing.T) { + for _, test := range []struct { + name string + args []string + want string + }{ + {"native", []string{"--cover-public-origin", "https://site.test"}, "requires --protocol=web"}, + {"static", []string{"--protocol", "web", "--cover-root", "/unused", "--cover-public-origin", "https://site.test"}, "requires --protocol=web"}, + {"http public", []string{"--protocol", "web", "--cover-upstream", "http://origin.test", "--cover-public-origin", "http://site.test"}, "configure upstream cover"}, + {"whitespace public", []string{"--protocol", "web", "--cover-upstream", "http://origin.test", "--cover-public-origin", " "}, "configure upstream cover"}, + {"bad public URL", []string{"--protocol", "web", "--cover-upstream", "http://origin.test", "--cover-public-origin", "https://site.test:%"}, "parse --cover-public-origin"}, + {"upstream base path", []string{"--protocol", "web", "--cover-upstream", "http://origin.test/base", "--cover-public-origin", "https://site.test"}, "configure upstream cover"}, + } { + t.Run(test.name, func(t *testing.T) { + err := runServer(context.Background(), test.args) + if err == nil || !strings.Contains(err.Error(), test.want) { + t.Fatalf("error = %v; want %q", err, test.want) + } + }) + } +} + +func TestPublicOriginPreflightStaysOffline(t *testing.T) { + clearPreflightEnvironment(t) + files := newPreflightFiles(t) + dns := denyPreflightDNS(t) + var requests atomic.Int32 + upstream := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { requests.Add(1) })) + defer upstream.Close() + for _, target := range []string{"https://origin.invalid", upstream.URL} { + args := append(files.serverArgs(), "--protocol", "web", "--cover-upstream", target, "--cover-public-origin", "https://site.invalid:8443/") + if err := runServer(context.Background(), args); err != nil { + t.Fatal(err) + } + } + if dns.Load() != 0 || requests.Load() != 0 { + t.Fatalf("preflight DNS=%d upstream=%d", dns.Load(), requests.Load()) + } +} + +func TestPublicOriginJSONAndCLIOverride(t *testing.T) { + clearPreflightEnvironment(t) + files := newPreflightFiles(t) + path := writeTestCommandConfig(t, `{"protocol":"web","cover-upstream":"https://origin.invalid","cover-public-origin":"https://site.invalid"}`) + args := append(files.serverArgs(), "--config", path) + if err := runServer(context.Background(), args); err != nil { + t.Fatalf("JSON public origin: %v", err) + } + // Explicit empty CLI opt-out restores the old upstream base-path policy. + args = append(args, "--cover-public-origin=", "--cover-upstream=https://origin.invalid/base?fixed=1") + if err := runServer(context.Background(), args); err != nil { + t.Fatalf("CLI opt-out: %v", err) + } + args = append(args, "--cover-public-origin=http://site.invalid") + if err := runServer(context.Background(), args); err == nil { + t.Fatal("invalid CLI override was accepted") + } +} + +func TestPublicOriginCoverFactory(t *testing.T) { + if _, err := buildCoverHandlerWithPublicOrigin(t.TempDir(), "", "https://site.test"); err == nil { + t.Fatal("static public-origin combination accepted") + } + handler, err := buildCoverHandlerWithPublicOrigin("", "https://origin.invalid", "https://site.test") + if err != nil { + t.Fatal(err) + } + request := httptest.NewRequest(http.MethodGet, "https://other.test/", nil) + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + if response.Code != http.StatusMisdirectedRequest { + t.Fatalf("wrong Host status=%d", response.Code) + } +} diff --git a/docs/DEPLOYMENT.md b/docs/DEPLOYMENT.md index 7740984..5394929 100644 --- a/docs/DEPLOYMENT.md +++ b/docs/DEPLOYMENT.md @@ -155,6 +155,19 @@ headers, and preserves the request path and query. Treat that origin as an Internet-facing application; do not point it at metadata or control-plane services. +For an owned website that must see its public virtual host, source builds +after v1.0.1 additionally accept `--cover-public-origin https://www.example.com` +(JSON: `"cover-public-origin":"https://www.example.com"`). This optional mode +retains fixed upstream dialing/TLS, uses the public HTTP Host, generates only +trusted public Host/HTTPS forwarding metadata, and rejects mismatched cover +Host or present Origin. Both URLs must be root origins; the public URL must +be HTTPS. Configure the backend's public virtual host, canonical URLs and +CSRF/session policy first: cookies, redirects and Origin are not rewritten. +This is not available with static cover or in the v1.0.1 binary. Offline +`--check` validates the configuration without contacting the upstream. +Remove the key/flag to roll back; existing default behavior is unchanged. +See [public-origin requirements and boundaries](WEB_COVER.md#optional-fixed-public-origin-in-source-builds). + Source builds after v1.0.1 can forward narrowly validated H1.1 WebSocket GET upgrades to this same fixed origin; no new flag or arbitrary Upgrade proxy is introduced. Authorization and nominated hop fields are still stripped. diff --git a/docs/WEB_COVER.md b/docs/WEB_COVER.md index 7114d02..3b1a0f7 100644 --- a/docs/WEB_COVER.md +++ b/docs/WEB_COVER.md @@ -202,6 +202,62 @@ the upstream cannot be reached. Consequently, an upstream that requires an `Authorization` request header is not suitable without a separate authorized front end. +#### Optional fixed public origin in source builds + +Source builds after v1.0.1 can separate the website's public HTTP identity +from the fixed upstream connection address: + +```bash +./autocar server \ + --protocol web --listen :443 --tcp-listen :443 \ + --cover-upstream https://origin.example.net \ + --cover-public-origin https://www.example.com \ + --cert /etc/autocar/server.crt --key /etc/autocar/server.key \ + --token-file /etc/autocar/relay-token +``` + +This is opt-in, upstream-only, and not part of the v1.0.1 binary. Without +`--cover-public-origin`, the existing upstream-Host and base-path/query behavior +is unchanged. With it, both configured URLs must be root origins (empty path +or `/`, no query, user information or fragment); the public origin must use +HTTPS. Use canonical ASCII DNS names (including already encoded punycode), +IPv4, or bracketed IPv6, with an optional numeric port from 1 to 65535. +Trailing dots, zone IDs, nonstandard numeric IP aliases and leading-zero ports +are rejected; no DNS or IDNA conversion is performed during validation. +DNS case and explicit default ports are equivalent for the public guard. + +The upstream URL still controls dialing, TLS certificate verification and +SNI. Only the outgoing HTTP `Host` changes to the configured public authority. +Incoming `Forwarded`, all `X-Forwarded-*` headers and `X-Real-IP` are removed +from headers and declared request trailers; +the proxy generates only fixed `X-Forwarded-Host` and `X-Forwarded-Proto: https`. +It does not forward the visitor's IP address. Configure the backend to accept +this public virtual host, generate public canonical URLs, and trust **only** +these two relay-generated headers from this relay. Other backend-specific +trust headers are not a universal allowlist: do not trust arbitrary client +headers or expose an internal control plane as the website. + +Cover requests with a different Host receive ordinary `421 Misdirected +Request`. If an Origin header is present it must contain exactly one valid +same-public HTTPS origin; foreign, null, empty or ambiguous values receive +ordinary `403 Forbidden`. A Connection nomination of Origin or Referer also +receives `403`, before hop-header stripping can hide it. This opt-in mode is +for same-origin websites, not cross-origin CORS applications. Authenticated +tunnel dispatch is unchanged: these guards apply only to the website handler. + +Cookie, Set-Cookie, Location, Origin, Referer and HTML are **not rewritten**. +The website remains responsible for CSRF tokens, session authentication, +cookie domains/attributes and its own Origin policy, including requests with +no Origin. Origin guards inspect HTTP headers, not trailer values; the backend +must not merge security or trust trailers into request headers. A backend that +compares Origin and Host lexically can still reject +equivalent noncanonical spellings; the proxy deliberately preserves Origin +bytes. Third-party redirects are passed through, not followed by the proxy. +This mode cannot make an arbitrary third-party site compatible or establish +browser-like traffic fingerprints. `--check` validates it offline; JSON uses +the same `cover-public-origin` key. Remove the flag/key (or explicitly pass +`--cover-public-origin=`) to restore the default mode. + For an actual `OPTIONS *` request, source builds after v1.0.1 preserve the asterisk request-target at the fixed upstream authority. The configured base path and query are not added: this request concerns the origin as a whole, not diff --git a/internal/cover/handler.go b/internal/cover/handler.go index cfb5b6f..034656b 100644 --- a/internal/cover/handler.go +++ b/internal/cover/handler.go @@ -36,6 +36,10 @@ func NewReverseProxyHandler(origin *url.URL, transport http.RoundTripper) (http. if err != nil { return nil, err } + return newReverseProxyHandler(target, transport, ""), nil +} + +func newReverseProxyHandler(target *url.URL, transport http.RoundTripper, publicAuthority string) *httputil.ReverseProxy { if transport == nil { defaultTransport := http.DefaultTransport.(*http.Transport).Clone() defaultTransport.Proxy = nil @@ -62,6 +66,19 @@ func NewReverseProxyHandler(origin *url.URL, transport http.RoundTripper) (http. // nomination boundary, and restore only our validated H1 WebSocket. removeConnectionNominatedHeaders(request.Out.Header, request.In.Header) removeUnsafeHeaders(request.Out.Header) + if publicAuthority != "" { + // The URL still chooses the fixed upstream dial and TLS target. + // Only this opt-in mode supplies a configured public vhost and + // trusted forwarding metadata, after all request nominations. + removePublicOriginForwardingHeaders(request.Out.Header) + // Replayed bodies can already have populated Trailer values. + // ReverseProxy owns this cloned map; do not trust those fields + // or mutate the original request's body or trailer map. + removePublicOriginForwardingHeaders(request.Out.Trailer) + request.Out.Host = publicAuthority + request.Out.Header.Set("X-Forwarded-Host", publicAuthority) + request.Out.Header.Set("X-Forwarded-Proto", "https") + } if upgrade.eligible { request.Out.Header.Set("Connection", "Upgrade") request.Out.Header.Set("Upgrade", "websocket") @@ -99,7 +116,7 @@ func NewReverseProxyHandler(origin *url.URL, transport http.RoundTripper) (http. }, ErrorLog: log.New(io.Discard, "", 0), } - return proxy, nil + return proxy } // informationalHeaderTransport filters upstream headers before ReverseProxy's diff --git a/internal/cover/public_origin.go b/internal/cover/public_origin.go new file mode 100644 index 0000000..1a58d3f --- /dev/null +++ b/internal/cover/public_origin.go @@ -0,0 +1,251 @@ +package cover + +import ( + "errors" + "net/http" + "net/netip" + "net/url" + "strconv" + "strings" +) + +// NewReverseProxyHandlerWithPublicOrigin returns an opt-in reverse proxy for +// an upstream site configured to serve public's HTTPS virtual host. Routing +// and TLS still use upstream; public supplies only the outbound Host and the +// trusted X-Forwarded-Host/Proto values. Both URLs must be root origins. +// +// Origin, Referer, cookies, redirects, and content are never rewritten. The +// upstream must generate its own public URLs and retain its CSRF/Origin checks; +// an absent Origin is not evidence that a request is safe. Its trusted-header +// whitelist must use only the proxy-owned forwarding fields, not arbitrary +// client headers or backend-specific client-IP aliases. +func NewReverseProxyHandlerWithPublicOrigin(upstream, public *url.URL, transport http.RoundTripper) (http.Handler, error) { + target, err := normalizePublicOriginURL(upstream, false) + if err != nil { + return nil, err + } + publicURL, err := normalizePublicOriginURL(public, true) + if err != nil { + return nil, err + } + return &publicOriginHandler{ + proxy: newReverseProxyHandler(target, transport, publicURL.Host), + authority: publicURL.Host, + }, nil +} + +// A fresh URL and canonical strings own all configuration used after the +// constructor returns. Caller mutations cannot change policy or routing. +func normalizePublicOriginURL(origin *url.URL, public bool) (*url.URL, error) { + invalid := errors.New("public-origin reverse proxy requires valid root origins") + if origin == nil { + return nil, invalid + } + scheme := strings.ToLower(origin.Scheme) + if scheme != "https" && (public || scheme != "http") { + return nil, invalid + } + if origin.Path != "" && origin.Path != "/" || origin.RawPath != "" || + origin.RawQuery != "" || origin.ForceQuery || origin.User != nil || origin.Opaque != "" || + origin.Fragment != "" || origin.RawFragment != "" || origin.OmitHost { + return nil, invalid + } + authority, err := canonicalOriginAuthority(origin.Host, scheme) + if err != nil { + return nil, invalid + } + return &url.URL{Scheme: scheme, Host: authority, Path: origin.Path}, nil +} + +type publicOriginHandler struct { + proxy http.Handler + authority string +} + +func (h *publicOriginHandler) ServeHTTP(w http.ResponseWriter, request *http.Request) { + if request == nil { + http.Error(w, http.StatusText(http.StatusMisdirectedRequest), http.StatusMisdirectedRequest) + return + } + authority, err := canonicalOriginAuthority(request.Host, "https") + if err != nil || authority != h.authority { + http.Error(w, http.StatusText(http.StatusMisdirectedRequest), http.StatusMisdirectedRequest) + return + } + // Inspect the original request before ReverseProxy's hop filtering can + // erase security metadata, even if the nominated field is absent. + nominations := collectConnectionNominations(request.Header) + _, originNominated := nominations["origin"] + _, refererNominated := nominations["referer"] + if originNominated || refererNominated || !h.acceptsOrigin(request.Header) { + http.Error(w, http.StatusText(http.StatusForbidden), http.StatusForbidden) + return + } + // Do not wrap the ResponseWriter: its streaming/Hijacker capabilities and + // the existing WebSocket/response-body owners remain unchanged. + h.proxy.ServeHTTP(w, request) +} + +func (h *publicOriginHandler) acceptsOrigin(header http.Header) bool { + present := false + count := 0 + value := "" + for field, values := range header { + if !httpToken(field) || !strings.EqualFold(field, "Origin") { + continue + } + present = true + for _, entry := range values { + count++ + if count > 1 { + return false + } + value = entry + } + } + if !present { + return true + } + value = strings.Trim(value, " \t") + if count != 1 || value == "" || len(value) > len("https://")+maxOriginAuthorityLength || strings.ContainsAny(value, "?#") { + return false + } + origin, err := url.Parse(value) + if err != nil || !strings.EqualFold(origin.Scheme, "https") || origin.Host == "" || + origin.Path != "" || origin.RawPath != "" || origin.RawQuery != "" || origin.ForceQuery || + origin.User != nil || origin.Opaque != "" || origin.Fragment != "" || origin.RawFragment != "" || origin.OmitHost { + return false + } + authority, err := canonicalOriginAuthority(origin.Host, "https") + return err == nil && authority == h.authority +} + +// A DNS name is at most 253 bytes, with an optional colon and five port +// digits. Accept only conventional ASCII DNS or exact IP literals: legacy +// numeric/hex IPv4 aliases and DNS resolution are deliberately not involved. +const maxOriginAuthorityLength = 259 + +func canonicalOriginAuthority(authority, scheme string) (string, error) { + invalid := errors.New("invalid origin authority") + if authority == "" || len(authority) > maxOriginAuthorityLength { + return "", invalid + } + for index := range authority { + if authority[index] < '!' || authority[index] >= 127 { + return "", invalid + } + } + if strings.ContainsAny(authority, "%\\/?#@") { + return "", invalid + } + host, port := authority, "" + if authority[0] == '[' { + end := strings.IndexByte(authority, ']') + if end < 2 { + return "", invalid + } + addr, err := netip.ParseAddr(authority[1:end]) + if err != nil || addr.Is4() || addr.Zone() != "" { + return "", invalid + } + host = "[" + addr.String() + "]" + if suffix := authority[end+1:]; suffix != "" { + if suffix[0] != ':' || len(suffix) == 1 { + return "", invalid + } + port = suffix[1:] + } + } else { + if strings.ContainsAny(authority, "[]") || strings.Count(authority, ":") > 1 { + return "", invalid + } + if index := strings.IndexByte(authority, ':'); index >= 0 { + host, port = authority[:index], authority[index+1:] + if port == "" { + return "", invalid + } + } + if addr, err := netip.ParseAddr(host); err == nil && addr.Is4() { + host = addr.String() + } else { + host = strings.ToLower(host) + if !validOriginDNSName(host) { + return "", invalid + } + } + } + if port != "" { + if len(port) > 5 || port[0] == '0' { + return "", invalid + } + for index := range port { + if port[index] < '0' || port[index] > '9' { + return "", invalid + } + } + number, err := strconv.Atoi(port) + if err != nil || number < 1 || number > 65535 { + return "", invalid + } + if scheme == "https" && number == 443 || scheme == "http" && number == 80 { + port = "" + } + } + if port != "" { + return host + ":" + port, nil + } + return host, nil +} + +func validOriginDNSName(host string) bool { + if host == "" || len(host) > 253 { + return false + } + labels := strings.Split(host, ".") + for _, label := range labels { + if len(label) == 0 || len(label) > 63 || label[0] == '-' || label[len(label)-1] == '-' { + return false + } + for index := range label { + char := label[index] + if char >= 'a' && char <= 'z' || char >= '0' && char <= '9' || char == '-' { + continue + } + return false + } + } + terminal := labels[len(labels)-1] + decimal := true + for index := range terminal { + decimal = decimal && terminal[index] >= '0' && terminal[index] <= '9' + } + if decimal { + return false + } + if strings.HasPrefix(terminal, "0x") { + hexadecimal := true + for index := 2; index < len(terminal); index++ { + char := terminal[index] + hexadecimal = hexadecimal && (char >= '0' && char <= '9' || char >= 'a' && char <= 'f') + } + if hexadecimal { + return false + } + } + return true +} + +func removePublicOriginForwardingHeaders(header http.Header) { + var scratch [64]byte + folded := scratch[:0] + for field := range header { + if !httpToken(field) { + continue + } + folded = foldASCIIHeaderName(folded, field) + name := string(folded) + if name == "forwarded" || name == "x-real-ip" || strings.HasPrefix(name, "x-forwarded-") { + delete(header, field) + } + } +} diff --git a/internal/cover/public_origin_test.go b/internal/cover/public_origin_test.go new file mode 100644 index 0000000..4664723 --- /dev/null +++ b/internal/cover/public_origin_test.go @@ -0,0 +1,524 @@ +package cover + +import ( + "fmt" + "io" + "net/http" + "net/http/httptest" + "net/http/httputil" + "net/url" + "reflect" + "strings" + "sync" + "sync/atomic" + "testing" +) + +// These tests exercise the opt-in public API, not its normalization helpers. +// Real TLS/SNI and site sessions have separate owned-loopback controls. +func TestPublicOriginConfiguredURLPolicy(t *testing.T) { + mutations := []struct { + name string + edit func(*url.URL) + }{ + {"nil", nil}, + {"empty_scheme", func(u *url.URL) { u.Scheme = "" }}, + {"other_scheme", func(u *url.URL) { u.Scheme = "wss" }}, + {"empty_host", func(u *url.URL) { u.Host = "" }}, + {"base_path", func(u *url.URL) { u.Path = "/site" }}, + {"raw_path", func(u *url.URL) { u.RawPath = "/" }}, + {"query", func(u *url.URL) { u.RawQuery = "next=site" }}, + {"force_query", func(u *url.URL) { u.ForceQuery = true }}, + {"user", func(u *url.URL) { u.User = url.UserPassword("dummy", "dummy") }}, + {"opaque", func(u *url.URL) { u.Opaque = "//example.com" }}, + {"fragment", func(u *url.URL) { u.Fragment = "section" }}, + {"raw_fragment", func(u *url.URL) { u.RawFragment = "section" }}, + {"omit_host", func(u *url.URL) { u.OmitHost = true }}, + } + for _, mode := range []string{"public", "upstream"} { + t.Run(mode, func(t *testing.T) { + for _, tc := range mutations { + t.Run(tc.name, func(t *testing.T) { + upstream, public := publicOriginUnitURLs() + selected := public + if mode == "upstream" { + selected = upstream + } + if tc.edit == nil { + if mode == "public" { + public = nil + } else { + upstream = nil + } + } else { + tc.edit(selected) + } + handler, err := NewReverseProxyHandlerWithPublicOrigin(upstream, public, publicOriginUnitTransport()) + if err == nil || handler != nil { + t.Fatalf("invalid %s URL accepted: handler=%T err=%v", mode, handler, err) + } + }) + } + }) + } + t.Run("public_http_forbidden", func(t *testing.T) { + upstream, public := publicOriginUnitURLs() + public.Scheme = "http" + if handler, err := NewReverseProxyHandlerWithPublicOrigin(upstream, public, publicOriginUnitTransport()); err == nil || handler != nil { + t.Fatalf("public HTTP accepted: handler=%T err=%v", handler, err) + } + }) + for _, tc := range []struct{ name, upstreamScheme, publicScheme, path string }{ + {"empty_root_http", "http", "https", ""}, + {"slash_root_https", "https", "https", "/"}, + {"scheme_case", "HTTP", "HTTPS", ""}, + } { + t.Run(tc.name, func(t *testing.T) { + upstream, public := publicOriginUnitURLs() + upstream.Scheme, public.Scheme = tc.upstreamScheme, tc.publicScheme + upstream.Path, public.Path = tc.path, tc.path + if _, err := NewReverseProxyHandlerWithPublicOrigin(upstream, public, publicOriginUnitTransport()); err != nil { + t.Fatalf("valid root origins rejected: %v", err) + } + }) + } +} + +func TestPublicOriginConfiguredAuthorityPolicy(t *testing.T) { + longDNS := strings.Repeat("a", 63) + "." + strings.Repeat("b", 63) + "." + strings.Repeat("c", 63) + "." + strings.Repeat("d", 61) + valid := []struct{ authority, canonical string }{ + {"PUBLIC.Example", "public.example"}, {"PUBLIC.Example:443", "public.example"}, + {"public.example:8443", "public.example:8443"}, {"127.0.0.1:443", "127.0.0.1"}, + {"[0:0:0:0:0:0:0:1]:443", "[::1]"}, {"[2001:db8::1]:8443", "[2001:db8::1]:8443"}, + {longDNS, longDNS}, {"a-b.example", "a-b.example"}, + } + for index, tc := range valid { + t.Run(fmt.Sprintf("valid_%02d", index), func(t *testing.T) { + upstream, public := publicOriginUnitURLs() + upstream.Host = "UPSTREAM.Example:80" + public.Host = tc.authority + var gotHost, gotTarget string + handler, err := NewReverseProxyHandlerWithPublicOrigin(upstream, public, roundTripFunc(func(r *http.Request) (*http.Response, error) { + gotHost = r.Host + gotTarget = r.URL.Host + return publicOriginUnitResponse(nil), nil + })) + if err != nil { + t.Fatal(err) + } + request := publicOriginUnitRequest() + request.Host = tc.authority + request.Header.Set("Origin", "https://"+tc.authority) + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + if response.Code != http.StatusOK || gotHost != tc.canonical || gotTarget != "upstream.example" { + t.Fatalf("canonical authority: status=%d got=%q want=%q fixedtarget=%q", response.Code, gotHost, tc.canonical, gotTarget) + } + }) + } + invalid := []struct{ name, authority string }{ + {"empty", ""}, {"trailing_dot", "public.example."}, {"empty_label", "public..example"}, + {"leading_hyphen", "-public.example"}, {"trailing_hyphen", "public-.example"}, + {"underscore", "public_site.example"}, {"unicode", "püblic.example"}, + {"numeric_last_label", "public.123"}, {"single_decimal", "2130706433"}, + {"short_ipv4", "127.1"}, {"octal_ipv4", "0177.0.0.1"}, {"hex_alias", "0x7f000001"}, + {"hex_last_label", "public.0xdead"}, {"label_too_long", strings.Repeat("a", 64) + ".example"}, + {"dns_too_long", longDNS + "e"}, {"space", "public.example "}, {"tab", "public.example\t"}, + {"newline", "public.example\n"}, {"at", "user@public.example"}, {"slash", "public.example/path"}, + {"backslash", "public.example\\path"}, {"percent", "public%2eexample"}, + {"empty_port", "public.example:"}, {"leading_zero_port", "public.example:0443"}, + {"zero_port", "public.example:0"}, {"port_overflow", "public.example:65536"}, + {"signed_port", "public.example:+443"}, {"named_port", "public.example:https"}, + {"ipv6_unbracketed", "::1"}, {"ipv6_zone", "[fe80::1%lo0]"}, + {"bracketed_ipv4", "[127.0.0.1]"}, {"ipv6_bad_suffix", "[::1]x"}, + } + for _, tc := range invalid { + t.Run(tc.name, func(t *testing.T) { + for _, mode := range []string{"public", "upstream"} { + upstream, public := publicOriginUnitURLs() + if mode == "public" { + public.Host = tc.authority + } else { + upstream.Host = tc.authority + } + if handler, err := NewReverseProxyHandlerWithPublicOrigin(upstream, public, publicOriginUnitTransport()); err == nil || handler != nil { + t.Fatalf("invalid %s authority %q accepted", mode, tc.authority) + } + } + }) + } +} + +func TestPublicOriginOriginalHostGuard(t *testing.T) { + for _, tc := range []struct { + name, host string + status int + }{ + {"exact", "public.example", 200}, {"case", "PUBLIC.EXAMPLE", 200}, {"default_port", "public.example:443", 200}, + {"foreign", "foreign.example", 421}, {"nondefault_port", "public.example:8443", 421}, + {"missing", "", 421}, {"trailing_dot", "public.example.", 421}, {"leading_zero_port", "public.example:0443", 421}, + {"whitespace", " public.example", 421}, {"userinfo", "user@public.example", 421}, + {"unbracketed_ip", "::1", 421}, {"zone", "[fe80::1%lo0]", 421}, + } { + t.Run(tc.name, func(t *testing.T) { + var calls int + upstream, public := publicOriginUnitURLs() + handler, err := NewReverseProxyHandlerWithPublicOrigin(upstream, public, roundTripFunc(func(r *http.Request) (*http.Response, error) { + calls++ + return publicOriginUnitResponse(nil), nil + })) + if err != nil { + t.Fatal(err) + } + request := publicOriginUnitRequest() + request.Host = tc.host + request.Header.Set("X-Forwarded-Host", "public.example") // Cannot rescue a foreign original Host. + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + wantCalls := 0 + if tc.status == 200 { + wantCalls = 1 + } + if response.Code != tc.status || calls != wantCalls { + t.Fatalf("status=%d calls=%d want=%d/%d", response.Code, calls, tc.status, wantCalls) + } + }) + } +} + +func TestPublicOriginOptionsAsteriskAndPortIsolation(t *testing.T) { + for _, tc := range []struct { + name, authority, origin string + status int + }{ + {"ipv6_canonical_origin", "[::1]", "https://[0:0:0:0:0:0:0:1]:443", 200}, + {"nondefault_exact", "public.example:8443", "https://PUBLIC.EXAMPLE:8443", 200}, + {"nondefault_wrong_origin_port", "public.example:8443", "https://public.example:443", 403}, + } { + t.Run(tc.name, func(t *testing.T) { + upstream, public := publicOriginUnitURLs() + public.Host = tc.authority + var calls int + handler, err := NewReverseProxyHandlerWithPublicOrigin(upstream, public, roundTripFunc(func(r *http.Request) (*http.Response, error) { + calls++ + if r.Header.Get("Origin") != tc.origin || r.Host != tc.authority { + t.Errorf("accepted original Origin/vhost changed: Host=%s Origin=%q", r.Host, r.Header.Get("Origin")) + } + return publicOriginUnitResponse(nil), nil + })) + if err != nil { + t.Fatal(err) + } + request := publicOriginUnitRequest() + request.Host = tc.authority + request.Header.Set("Origin", tc.origin) + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + wantCalls := 0 + if tc.status == 200 { + wantCalls = 1 + } + if response.Code != tc.status || calls != wantCalls { + t.Fatalf("status=%d calls=%d want=%d/%d", response.Code, calls, tc.status, wantCalls) + } + }) + } + for _, tc := range []struct { + name, host string + status int + }{ + {"options_public", "public.example", 200}, + {"options_foreign_host", "foreign.example", 421}, + } { + t.Run(tc.name, func(t *testing.T) { + upstream, public := publicOriginUnitURLs() + upstream.Path = "/" + var calls int + handler, err := NewReverseProxyHandlerWithPublicOrigin(upstream, public, roundTripFunc(func(r *http.Request) (*http.Response, error) { + calls++ + if r.Method != http.MethodOptions || r.URL.Path != "*" || r.URL.RawPath != "" || + r.URL.RawQuery != "" || r.URL.ForceQuery || r.URL.Scheme != "http" || r.URL.Host != "upstream.example" || + r.Host != "public.example" || r.Header.Get("X-Forwarded-Host") != "public.example" || r.Header.Get("X-Forwarded-Proto") != "https" { + t.Errorf("OPTIONS* route/metadata changed: request=%v URL=%s header=%v", r, r.URL, r.Header) + } + return publicOriginUnitResponse(nil), nil + })) + if err != nil { + t.Fatal(err) + } + request := httptest.NewRequest(http.MethodOptions, "*", nil) + request.Host = tc.host + request.Header.Set("Origin", "https://public.example") + request.Header.Set("X-Forwarded-Host", "public.example") + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + wantCalls := 0 + if tc.status == 200 { + wantCalls = 1 + } + if response.Code != tc.status || calls != wantCalls { + t.Fatalf("status=%d calls=%d want=%d/%d", response.Code, calls, tc.status, wantCalls) + } + }) + } +} + +func TestPublicOriginOriginalOriginAndNominations(t *testing.T) { + cases := []struct { + name string + header http.Header + status int + }{ + {"absent", http.Header{}, 200}, + {"exact", http.Header{"Origin": {"https://public.example"}}, 200}, + {"case_default_port", http.Header{"oRiGiN": {"HTTPS://PUBLIC.EXAMPLE:443"}}, 200}, + {"outer_ows", http.Header{"Origin": {" \thttps://public.example\t "}}, 200}, + {"foreign", http.Header{"Origin": {"https://foreign.example"}}, 403}, + {"null", http.Header{"Origin": {"null"}}, 403}, + {"empty", http.Header{"Origin": {""}}, 403}, + {"no_values", http.Header{"Origin": nil}, 403}, + {"multiple", http.Header{"Origin": {"https://public.example", "https://public.example"}}, 403}, + {"case_alias_multiple", http.Header{"Origin": {"https://public.example"}, "origin": {"https://public.example"}}, 403}, + {"comma_list", http.Header{"Origin": {"https://public.example, https://foreign.example"}}, 403}, + {"space_list", http.Header{"Origin": {"https://public.example https://foreign.example"}}, 403}, + {"http", http.Header{"Origin": {"http://public.example"}}, 403}, + {"slash_path", http.Header{"Origin": {"https://public.example/"}}, 403}, + {"resource_path", http.Header{"Origin": {"https://public.example/path"}}, 403}, + {"encoded_path", http.Header{"Origin": {"https://public.example/%2f"}}, 403}, + {"query", http.Header{"Origin": {"https://public.example?"}}, 403}, + {"fragment", http.Header{"Origin": {"https://public.example#"}}, 403}, + {"userinfo", http.Header{"Origin": {"https://user@public.example"}}, 403}, + {"opaque", http.Header{"Origin": {"https:public.example"}}, 403}, + {"non_ows", http.Header{"Origin": {"\u00a0https://public.example"}}, 403}, + {"newline", http.Header{"Origin": {"\nhttps://public.example"}}, 403}, + {"origin_nominated_absent", http.Header{"Connection": {"Origin"}}, 403}, + {"origin_nominated_present", http.Header{"Connection": {"keep-alive, oRiGiN"}, "Origin": {"https://public.example"}}, 403}, + {"referer_nominated_absent", http.Header{"connection": {"keep-alive, Referer"}}, 403}, + {"referer_nominated_present", http.Header{"Connection": {"\tREFERER "}, "Referer": {"https://public.example/page"}}, 403}, + {"alias_connection_nomination", http.Header{"Connection": {"keep-alive"}, "cOnNeCtIoN": {"Origin"}}, 403}, + {"ordinary_nomination", http.Header{"Connection": {"X-Hop"}, "X-Hop": {"dummy"}}, 200}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + upstream, public := publicOriginUnitURLs() + var calls int + var observed http.Header + handler, err := NewReverseProxyHandlerWithPublicOrigin(upstream, public, roundTripFunc(func(r *http.Request) (*http.Response, error) { + calls++ + observed = r.Header.Clone() + return publicOriginUnitResponse(nil), nil + })) + if err != nil { + t.Fatal(err) + } + request := publicOriginUnitRequest() + request.Header = tc.header.Clone() + before := request.Header.Clone() + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + wantCalls := 0 + if tc.status == 200 { + wantCalls = 1 + } + if response.Code != tc.status || calls != wantCalls { + t.Fatalf("status=%d calls=%d want=%d/%d", response.Code, calls, tc.status, wantCalls) + } + if !reflect.DeepEqual(request.Header, before) { + t.Fatal("original inbound headers mutated") + } + if tc.status == 200 { + for field, values := range tc.header { + if strings.EqualFold(field, "Origin") && !reflect.DeepEqual(observed[field], values) { + t.Fatalf("original Origin bytes changed: got=%q want=%q", observed[field], values) + } + } + if tc.name == "ordinary_nomination" && observed.Get("X-Hop") != "" { + t.Fatal("ordinary hop nomination survived") + } + } + }) + } +} + +func TestPublicOriginFixedRouteAndWebsiteMetadata(t *testing.T) { + upstream, public := publicOriginUnitURLs() + upstream.Scheme, upstream.Host = "https", "UPSTREAM.Example:443" + public.Host = "PUBLIC.Example:443" + upstreamBefore, publicBefore := *upstream, *public + websiteHeaders := http.Header{ + "Location": {"https://login.external.example/callback?next=%2faccount"}, + "Set-Cookie": {"session=owned-dummy; Path=/; Secure; HttpOnly; SameSite=Strict; Future-Flag=opaque", "oauth=owned-dummy; Domain=public.example; Path=/callback; Secure; Partitioned"}, + "Authentication-Info": {"ordinary-site-proof"}, "Www-Authenticate": {"Basic realm=ordinary-site"}, + "Content-Location": {"https://upstream.example/original"}, + } + var observed *http.Request + handler, err := NewReverseProxyHandlerWithPublicOrigin(upstream, public, roundTripFunc(func(r *http.Request) (*http.Response, error) { + observed = r.Clone(r.Context()) + return publicOriginUnitResponse(websiteHeaders.Clone()), nil + })) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(*upstream, upstreamBefore) || !reflect.DeepEqual(*public, publicBefore) { + t.Fatal("constructor mutated caller URL inputs") + } + request := publicOriginUnitRequest() + request.Header = http.Header{ + "Origin": {"https://public.example"}, "Referer": {"https://public.example/from?verbatim=%2f"}, + "Cookie": {"session=owned-dummy; opaque=%2F"}, "X-End-To-End": {"ordinary"}, + "Forwarded": {"host=attacker.example;proto=http"}, "forwarded": {"host=other.example"}, + "X-Forwarded-Host": {"attacker.example"}, "x-forwarded-host": {"other.example"}, + "X-Forwarded-Proto": {"http"}, "x-forwarded-proto": {"ftp"}, + "X-Forwarded-For": {"192.0.2.1"}, "x-forwarded-for": {"192.0.2.2"}, + "X-Forwarded-Port": {"80"}, "X-Forwarded-Prefix": {"/attacker"}, "x-forwarded-Whatever": {"dummy"}, + "X-Real-IP": {"192.0.2.3"}, "x-real-ip": {"192.0.2.4"}, + "Authorization": {"dummy-private"}, "Proxy-Authorization": {"dummy-private"}, + "Connection": {"X-Hop, X-Forwarded-Host"}, "X-Hop": {"dummy-hop"}, + } + request.Trailer = http.Header{ + "Forwarded": {"host=trailer-attacker.example"}, "forwarded": {"proto=http"}, + "X-Forwarded-Host": {"trailer-attacker.example"}, "x-forwarded-proto": {"http"}, + "X-Forwarded-For": {"192.0.2.5"}, "X-Forwarded-Prefix": {"/trailer"}, + "X-Real-IP": {"192.0.2.6"}, "x-real-ip": {"192.0.2.7"}, + "X-Ordinary-Trailer": {"site-trailer-verbatim"}, + } + beforeHeader, beforeURL := request.Header.Clone(), *request.URL + beforeTrailer := request.Trailer.Clone() + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + if response.Code != 200 || observed == nil { + t.Fatalf("ordinary response failed: %d", response.Code) + } + if observed.URL.Scheme != "https" || observed.URL.Host != "upstream.example" || observed.Host != "public.example" || + observed.URL.EscapedPath() != "/account%2Fentry" || observed.URL.RawQuery != "next=%2f&x=1" { + t.Fatalf("route/vhost changed: URL=%s Host=%q", observed.URL, observed.Host) + } + for field, values := range observed.Header { + lower := strings.ToLower(field) + if lower == "forwarded" || lower == "x-real-ip" || strings.HasPrefix(lower, "x-forwarded-") { + if field == "X-Forwarded-Host" && reflect.DeepEqual(values, []string{"public.example"}) || + field == "X-Forwarded-Proto" && reflect.DeepEqual(values, []string{"https"}) { + continue + } + t.Errorf("untrusted forwarding field survived: %s=%q", field, values) + } + } + if observed.Header.Get("X-Forwarded-Host") != "public.example" || observed.Header.Get("X-Forwarded-Proto") != "https" { + t.Fatal("fixed forwarding fields missing") + } + if !reflect.DeepEqual(observed.Trailer, http.Header{"X-Ordinary-Trailer": {"site-trailer-verbatim"}}) { + t.Errorf("forwarding trailers survived or ordinary trailer changed: %v", observed.Trailer) + } + for _, name := range []string{"Origin", "Referer", "Cookie", "X-End-To-End"} { + if !reflect.DeepEqual(observed.Header.Values(name), beforeHeader.Values(name)) { + t.Errorf("ordinary %s bytes changed", name) + } + } + for _, name := range []string{"Authorization", "Proxy-Authorization", "X-Hop"} { + if observed.Header.Get(name) != "" { + t.Errorf("unsafe %s survived", name) + } + } + for name, values := range websiteHeaders { + if !reflect.DeepEqual(response.Header().Values(name), values) { + t.Errorf("site metadata %s changed: got=%q want=%q", name, response.Header().Values(name), values) + } + } + if response.Body.String() != "ordinary-site-body" || !reflect.DeepEqual(request.Header, beforeHeader) || + !reflect.DeepEqual(request.Trailer, beforeTrailer) || !reflect.DeepEqual(*request.URL, beforeURL) { + t.Fatal("body or original request changed") + } +} + +func TestPublicOriginConfigurationIsImmutableAndConcurrent(t *testing.T) { + upstream, public := publicOriginUnitURLs() + var calls atomic.Int32 + violations := make(chan string, 32) // At most two independent oracle failures per request. + handler, err := NewReverseProxyHandlerWithPublicOrigin(upstream, public, roundTripFunc(func(r *http.Request) (*http.Response, error) { + calls.Add(1) + if r.URL.Host != "upstream.example" || r.URL.Scheme != "http" || r.Host != "public.example" || + r.Header.Get("X-Forwarded-Host") != "public.example" || r.Header.Get("X-Forwarded-Proto") != "https" || + r.Header.Get("Origin") != "https://public.example" { + violations <- fmt.Sprintf("URL=%s Host=%s header=%v", r.URL, r.Host, r.Header) + } + return publicOriginUnitResponse(http.Header{"X-Request-Marker": {r.Header.Get("X-Request-Marker")}}), nil + })) + if err != nil { + t.Fatal(err) + } + // Mutate only after construction, not concurrently with construction. The + // handler must own copied policy, not read these caller objects later. + upstream.Scheme, upstream.Host, upstream.Path, upstream.RawQuery = "https", "foreign.example", "/changed", "changed=1" + public.Host, public.Scheme = "foreign.example", "http" + var joined sync.WaitGroup + for index := range 16 { + joined.Add(1) + go func() { + defer joined.Done() + request := publicOriginUnitRequest() + marker := fmt.Sprint(index) + request.Header.Set("Origin", "https://public.example") + request.Header.Set("X-Request-Marker", marker) + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + if response.Code != 200 || response.Header().Get("X-Request-Marker") != marker { + violations <- fmt.Sprintf("request=%s status=%d marker=%q", marker, response.Code, response.Header().Get("X-Request-Marker")) + } + }() + } + joined.Wait() // The fake RoundTripper performs no blocking I/O. + close(violations) + for failure := range violations { + t.Error(failure) + } + if calls.Load() != 16 { + t.Fatalf("calls=%d want=16", calls.Load()) + } +} + +func TestPublicOriginDefaultConstructorPolicyUnchanged(t *testing.T) { + origin := &url.URL{Scheme: "http", Host: "upstream.example", Path: "/base", RawQuery: "fixed=1"} + var observed *http.Request + handler, err := NewReverseProxyHandler(origin, roundTripFunc(func(r *http.Request) (*http.Response, error) { + observed = r.Clone(r.Context()) + return publicOriginUnitResponse(nil), nil + })) + if err != nil { + t.Fatal(err) + } + if _, ok := handler.(*httputil.ReverseProxy); !ok { + t.Fatalf("old constructor return type changed: %T", handler) + } + request := publicOriginUnitRequest() + request.Host = "unconfigured.example" + request.Header.Set("Origin", "https://foreign.example") + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + if response.Code != 200 || observed == nil || observed.Host != "upstream.example" || + observed.URL.EscapedPath() != "/base/account%2Fentry" || observed.URL.RawQuery != "fixed=1&next=%2f&x=1" || + observed.Header.Get("Origin") != "https://foreign.example" || observed.Header.Get("X-Forwarded-Host") != "" { + t.Fatalf("old default/base-path policy changed: status=%d request=%v", response.Code, observed) + } +} + +func publicOriginUnitURLs() (*url.URL, *url.URL) { + return &url.URL{Scheme: "http", Host: "upstream.example"}, &url.URL{Scheme: "https", Host: "public.example"} +} + +func publicOriginUnitRequest() *http.Request { + request := httptest.NewRequest(http.MethodGet, "https://public.example/account%2Fentry?next=%2f&x=1", nil) + request.Host = "public.example" + return request +} + +func publicOriginUnitTransport() http.RoundTripper { + return roundTripFunc(func(*http.Request) (*http.Response, error) { return publicOriginUnitResponse(nil), nil }) +} + +func publicOriginUnitResponse(header http.Header) *http.Response { + if header == nil { + header = make(http.Header) + } + return &http.Response{StatusCode: 200, Proto: "HTTP/1.1", ProtoMajor: 1, ProtoMinor: 1, Header: header, + Body: io.NopCloser(strings.NewReader("ordinary-site-body"))} +} diff --git a/internal/tunnel/web_cover_public_origin_test.go b/internal/tunnel/web_cover_public_origin_test.go new file mode 100644 index 0000000..5a4dc24 --- /dev/null +++ b/internal/tunnel/web_cover_public_origin_test.go @@ -0,0 +1,741 @@ +package tunnel + +import ( + "bufio" + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/sha1" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/base64" + "errors" + "fmt" + "io" + "math/big" + "net" + "net/http" + "net/http/cookiejar" + "net/netip" + "net/url" + "reflect" + "strconv" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/cppla/autocar/internal/cover" + "github.com/cppla/autocar/internal/transport" +) + +const ( + webPublicOriginTLSName = "fixed-origin.fixture.test" + webPublicOriginCSRF = "fictional-local-csrf-token" + webPublicOriginSession = "fictional-local-session" + webPublicOriginForm = `
` + webPublicOriginWelcome = "ordinary authenticated website welcome" + webPublicOriginThirdParty = "https://thirdparty.invalid/callback?opaque=one%2Ftwo" +) + +var webPublicOriginRawCookies = []string{ + "origin_cookie=one; Domain=fixed-origin.fixture.test; Path=/; Secure; HttpOnly; SameSite=Lax; Priority=High", + "odd_cookie=two; Path=/; Secure; Experimental=kept", +} + +// These native-client exchanges test a website's independent cookie and CSRF +// policy. They do not claim browser SameSite enforcement or a browser campaign. +func TestWebCoverPublicOriginLoginOnWire(t *testing.T) { + for _, proto := range []int{1, 2, 3} { + t.Run(fmt.Sprintf("h%d", proto), func(t *testing.T) { + for _, mode := range []bool{false, true} { + name := "default_host_mismatch" + if mode { + name = "configured_public_vhost" + } + t.Run(name, func(t *testing.T) { + f := newWebPublicOriginFixture(t, proto, mode) + response, body := f.exchange(t, http.MethodGet, "/form", nil, nil, "") + if response.StatusCode != 200 || body != webPublicOriginForm { + t.Fatalf("ordinary form = %d/%q", response.StatusCode, body) + } + f.observed(t, "", "") + assertWebPublicOriginCookie(t, response, "csrf", webPublicOriginCSRF, false) + postHeaders := http.Header{"Origin": {f.public.String()}, "Referer": {f.public.String() + "/form"}} + form := url.Values{"csrf": {webPublicOriginCSRF}, "username": {"ordinary-visitor"}, "password": {"fictional-local-password"}} + response, _ = f.exchange(t, http.MethodPost, "/login", form, postHeaders, "") + seen := f.observed(t, f.public.String(), f.public.String()+"/form") + if !strings.Contains(seen.header.Get("Cookie"), "csrf="+webPublicOriginCSRF) { + t.Error("actual login POST did not return the form's CSRF cookie") + } + if !mode { + // The existing constructor intentionally uses the upstream + // Host. This is a supported default-policy control, not a + // compile-failure or a claim that the default is broken. + if response.StatusCode != 403 || seen.host != f.upstream.Host { + t.Errorf("default same-origin rejection = %d, actual Host=%q", response.StatusCode, seen.host) + } + if response.Header.Get("Location") != "" { + t.Error("rejected login redirected") + } + return + } + assertWebPublicOriginRedirect(t, response, f.public.String()+"/welcome") + assertWebPublicOriginCookie(t, response, "site_session", webPublicOriginSession, false) + location, err := url.Parse(response.Header.Get("Location")) + if err != nil || location.Scheme != f.public.Scheme || location.Host != f.public.Host { + t.Fatalf("refusing any non-owned login redirect: %q, %v", response.Header.Get("Location"), err) + } + response, body = f.exchange(t, http.MethodGet, location.RequestURI(), nil, nil, "") + seen = f.observed(t, "", "") + if response.StatusCode != 200 || body != webPublicOriginWelcome || !strings.Contains(seen.header.Get("Cookie"), "site_session="+webPublicOriginSession) { + t.Fatalf("actual cookie-authenticated welcome = %d/%q, Cookie=%q", response.StatusCode, body, seen.header.Get("Cookie")) + } + response, _ = f.exchange(t, http.MethodPost, "/logout", url.Values{"csrf": {webPublicOriginCSRF}}, postHeaders, "") + f.observed(t, f.public.String(), f.public.String()+"/form") + assertWebPublicOriginRedirect(t, response, f.public.String()+"/form") + assertWebPublicOriginCookie(t, response, "site_session", "", true) + response, _ = f.exchange(t, http.MethodGet, "/welcome", nil, nil, "") + seen = f.observed(t, "", "") + if response.StatusCode != 401 || strings.Contains(seen.header.Get("Cookie"), "site_session=") { + t.Errorf("post-logout welcome = %d, Cookie=%q", response.StatusCode, seen.header.Get("Cookie")) + } + t.Log("actual website CSRF/cookie login and logout; TLS target/SNI remains fixed origin, HTTP Host is public vhost") + }) + } + }) + } +} + +func TestWebCoverPublicOriginGuardOnWire(t *testing.T) { + for _, proto := range []int{1, 2, 3} { + t.Run(fmt.Sprintf("h%d", proto), func(t *testing.T) { + f := newWebPublicOriginFixture(t, proto, true) + for _, test := range []struct { + name, host string + header http.Header + status int + }{ + {"wrong_host", "evil.invalid", nil, 421}, + {"wrong_host_before_origin", "evil.invalid", http.Header{"Origin": {"https://evil.invalid"}}, 421}, + {"wrong_port", "127.0.0.1:1", nil, 421}, + {"evil_origin", "", http.Header{"Origin": {"https://evil.invalid"}}, 403}, + {"null_origin", "", http.Header{"Origin": {"null"}}, 403}, + {"empty_origin", "", http.Header{"Origin": {""}}, 403}, + {"duplicate_origin", "", http.Header{"Origin": {f.public.String(), f.public.String()}}, 403}, + {"origin_list", "", http.Header{"Origin": {f.public.String() + " https://evil.invalid"}}, 403}, + {"origin_with_path", "", http.Header{"Origin": {f.public.String() + "/login"}}, 403}, + {"wrong_scheme", "", http.Header{"Origin": {"http://" + f.public.Host}}, 403}, + } { + t.Run(test.name, func(t *testing.T) { + before := f.appCalls.Load() + response, _ := f.exchange(t, http.MethodGet, "/form", nil, test.header, test.host) + if response.StatusCode != test.status { + t.Errorf("guard status = %d, want %d", response.StatusCode, test.status) + } + if f.appCalls.Load() != before || f.originDials.Load() != 0 { + t.Error("guard rejection reached the fixed website or dialed its origin") + } + }) + } + // H2 native clients reject Connection: Origin before transmission; + // connection-specific headers are also forbidden in H3. These are + // real H1 guard probes, not claimed H2/H3 guard observations. + if proto == 1 { + for _, field := range []string{"Origin", "Referer"} { + t.Run("nominated_missing_"+strings.ToLower(field), func(t *testing.T) { + response, _ := f.exchange(t, http.MethodGet, "/form", nil, http.Header{"Connection": {field}}, "") + if response.StatusCode != 403 || f.appCalls.Load() != 0 || f.originDials.Load() != 0 { + t.Errorf("original missing %s nomination = %d, site/dial=%d/%d", field, response.StatusCode, f.appCalls.Load(), f.originDials.Load()) + } + }) + } + } + t.Run("absent_origin_keeps_site_csrf_policy", func(t *testing.T) { + before := f.appCalls.Load() + response, _ := f.exchange(t, http.MethodPost, "/login", url.Values{"csrf": {webPublicOriginCSRF}}, http.Header{"Referer": {"https://evil.invalid/form"}}, "") + f.observed(t, "", "https://evil.invalid/form") + if response.StatusCode != 403 || f.appCalls.Load() != before+1 { + t.Errorf("site's missing-Origin rejection = %d, site call delta=%d", response.StatusCode, f.appCalls.Load()-before) + } + }) + }) + } +} + +func TestWebCoverPublicOriginForwardingAndMetadataOnWire(t *testing.T) { + for _, proto := range []int{1, 2, 3} { + t.Run(fmt.Sprintf("h%d", proto), func(t *testing.T) { + f := newWebPublicOriginFixture(t, proto, true) + header := http.Header{ + "Origin": {f.public.String()}, "Referer": {f.public.String() + "/ordinary?q=kept"}, + "Cookie": {"ordinary=kept; another=two"}, "Forwarded": {"host=evil.invalid;proto=http"}, + "X-Forwarded-Host": {"evil.invalid", "second.invalid"}, "X-Forwarded-Proto": {"http"}, + "X-Forwarded-For": {"192.0.2.99"}, "X-Forwarded-Mystery": {"poisoned"}, "X-Real-IP": {"192.0.2.98"}, + } + response, body := f.exchange(t, http.MethodGet, "/metadata?plain=one%2Ftwo&empty=", nil, header, "") + seen := f.observed(t, f.public.String(), f.public.String()+"/ordinary?q=kept") + if response.StatusCode != 302 || response.Header.Get("Location") != webPublicOriginThirdParty || body != "location and cookie bytes unchanged" { + t.Errorf("unmodified metadata response = %d/%q/%q", response.StatusCode, response.Header.Get("Location"), body) + } + if got := response.Header.Values("Set-Cookie"); !reflect.DeepEqual(got, webPublicOriginRawCookies) { + t.Errorf("Set-Cookie bytes = %q, want %q", got, webPublicOriginRawCookies) + } + if seen.header.Get("Cookie") != "ordinary=kept; another=two" || seen.rawQuery != "plain=one%2Ftwo&empty=" { + t.Errorf("ordinary request Cookie/query changed = %q/%q", seen.header.Get("Cookie"), seen.rawQuery) + } + for field, values := range seen.header { + lower := strings.ToLower(field) + if lower == "forwarded" || lower == "x-real-ip" || (strings.HasPrefix(lower, "x-forwarded-") && lower != "x-forwarded-host" && lower != "x-forwarded-proto") { + t.Errorf("untrusted forwarding field reached website: %s=%q", field, values) + } + } + if !reflect.DeepEqual(seen.header.Values("X-Forwarded-Host"), []string{f.public.Host}) || !reflect.DeepEqual(seen.header.Values("X-Forwarded-Proto"), []string{"https"}) { + t.Errorf("trusted forwarding values = %v", seen.header) + } + if f.appCalls.Load() != 1 { + t.Errorf("redirect-disabled website calls=%d, want 1", f.appCalls.Load()) + } + }) + } +} + +func TestWebCoverPublicOriginWebSocketOnWire(t *testing.T) { + f := newWebPublicOriginFixture(t, 1, true) + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + raw, err := (&net.Dialer{Timeout: time.Second}).DialContext(ctx, "tcp", f.public.Host) + if err != nil { + t.Fatal(err) + } + defer raw.Close() + if err := raw.SetDeadline(time.Now().Add(3 * time.Second)); err != nil { + t.Fatal(err) + } + clientTLS := f.frontTLS.Clone() + clientTLS.NextProtos = []string{webHTTP11ALPN} + conn := tls.Client(raw, clientTLS) + defer conn.Close() + if err := conn.HandshakeContext(ctx); err != nil { + t.Fatal(err) + } + if len(conn.ConnectionState().VerifiedChains) == 0 { + t.Fatal("WebSocket frontend TLS was not verified") + } + request, err := http.NewRequest(http.MethodGet, f.public.String()+"/websocket", nil) + if err != nil { + t.Fatal(err) + } + const key = "dGhlIHNhbXBsZSBub25jZQ==" + request.Header.Set("Connection", "Upgrade") + request.Header.Set("Upgrade", "websocket") + request.Header.Set("Sec-WebSocket-Version", "13") + request.Header.Set("Sec-WebSocket-Key", key) + request.Header.Set("Origin", f.public.String()) + request.Header.Set("X-Public-Vhost-Fixture", t.Name()) + request.Header.Set("Proxy-Authorization", "Bearer fictional-invalid-ticket") + if err := request.Write(conn); err != nil { + t.Fatal(err) + } + reader := bufio.NewReader(conn) + response, err := http.ReadResponse(reader, request) + if err != nil { + t.Fatal(err) + } + accept := sha1.Sum([]byte(key + "258EAFA5-E914-47DA-95CA-C5AB0DC85B11")) + if response.StatusCode != 101 || response.ProtoMajor != 1 || response.ProtoMinor != 1 || response.Header.Get("Connection") != "Upgrade" || response.Header.Get("Upgrade") != "websocket" || response.Header.Get("Sec-WebSocket-Accept") != base64.StdEncoding.EncodeToString(accept[:]) { + t.Fatalf("actual public WebSocket handshake=%s/%d/%v", response.Proto, response.StatusCode, response.Header) + } + const payload = "echo" + mask := [4]byte{1, 2, 3, 4} + frame := []byte{0x81, 0x80 | byte(len(payload))} + frame = append(frame, mask[:]...) + for index := range payload { + frame = append(frame, payload[index]^mask[index%4]) + } + if n, err := conn.Write(frame); err != nil || n != len(frame) { + t.Fatalf("masked client frame write=%d/%v", n, err) + } + echo := make([]byte, 2+len(payload)) + if _, err := io.ReadFull(reader, echo); err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(echo, append([]byte{0x81, byte(len(payload))}, []byte(payload)...)) { + t.Fatalf("actual unmasked origin echo frame=%x", echo) + } + _ = conn.Close() + f.observed(t, f.public.String(), "") + select { + case marker := <-f.completed: + if marker != t.Name() { + t.Errorf("WebSocket completion marker=%q", marker) + } + case <-time.After(2 * time.Second): + t.Fatal("WebSocket proxy worker did not independently join") + } + if f.appCalls.Load() != 1 { + t.Errorf("actual WebSocket website calls=%d", f.appCalls.Load()) + } + t.Log("real H1 masked-client/unmasked-origin echo after website default same-origin Host predicate; no extended CONNECT or browser claim") +} + +type webPublicOriginObserved struct { + host, sni, rawQuery string + proto int + header http.Header +} + +type webPublicOriginFixture struct { + public, upstream *url.URL + client *http.Client + frontTLS *tls.Config + proto int + mode bool + appCalls, originDials, invalidDials, verifiedOriginResponses, targetDials, targetResolutions atomic.Int64 + records chan webPublicOriginObserved + completed chan string + stateMu sync.Mutex + session bool + registerHijacked func(net.Conn) bool + unregisterHijacked func(net.Conn) + asyncErrors chan error +} + +func newWebPublicOriginFixture(t *testing.T, proto int, mode bool) *webPublicOriginFixture { + t.Helper() + originTLS, originClientTLS := webPublicOriginTLSConfigs(t) + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = listener.Close() }) + upstream, err := url.Parse("https://" + net.JoinHostPort(webPublicOriginTLSName, strconv.Itoa(listener.Addr().(*net.TCPAddr).Port))) + if err != nil { + t.Fatal(err) + } + f := &webPublicOriginFixture{upstream: upstream, proto: proto, mode: mode, records: make(chan webPublicOriginObserved, 64), completed: make(chan string, 64), asyncErrors: make(chan error, 64)} + var admissionMu sync.Mutex + var workers sync.WaitGroup + closing := false + hijacked := make(map[net.Conn]struct{}) + f.registerHijacked = func(conn net.Conn) bool { + admissionMu.Lock() + defer admissionMu.Unlock() + if closing { + return false + } + hijacked[conn] = struct{}{} + return true + } + f.unregisterHijacked = func(conn net.Conn) { admissionMu.Lock(); defer admissionMu.Unlock(); delete(hijacked, conn) } + admit := func() bool { + admissionMu.Lock() + defer admissionMu.Unlock() + if closing { + return false + } + workers.Add(1) + return true + } + origin := &http.Server{ReadHeaderTimeout: 2 * time.Second, IdleTimeout: time.Second, Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !admit() { + http.Error(w, "closing", 503) + return + } + defer workers.Done() + f.appCalls.Add(1) + sni := "" + if r.TLS != nil { + sni = r.TLS.ServerName + } + f.records <- webPublicOriginObserved{r.Host, sni, r.URL.RawQuery, r.ProtoMajor, r.Header.Clone()} + f.website(w, r) + })} + originRT := &http.Transport{Proxy: nil, TLSClientConfig: originClientTLS, TLSNextProto: make(map[string]func(string, *tls.Conn) http.RoundTripper), TLSHandshakeTimeout: 2 * time.Second, ResponseHeaderTimeout: 2 * time.Second, IdleConnTimeout: time.Second, + DialContext: func(ctx context.Context, network, address string) (net.Conn, error) { + if network != "tcp" || address != upstream.Host { + f.invalidDials.Add(1) + return nil, errors.New("public-vhost fixture forbids every non-owned upstream") + } + f.originDials.Add(1) + return (&net.Dialer{Timeout: time.Second}).DialContext(ctx, "tcp", listener.Addr().String()) + }, + } + t.Cleanup(originRT.CloseIdleConnections) + var proxy http.Handler // Assigned once, before either server starts. + website := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !admit() { + http.Error(w, "closing", 503) + return + } + defer workers.Done() + defer func() { f.completed <- r.Header.Get("X-Public-Vhost-Fixture") }() + proxy.ServeHTTP(w, r) + }) + serverTLS, clientTLS := testTLSConfigs(t) + f.frontTLS = clientTLS.Clone() + server, err := ListenWeb(WebServerConfig{TCPAddress: "127.0.0.1:0", UDPAddress: "127.0.0.1:0", Token: webTestToken, TLSConfig: serverTLS, Cover: website, + Dialer: transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { + f.targetDials.Add(1) + return nil, errors.New("public vhost fixture forbids tunnel dial") + }), + UDPResolver: UDPResolverFunc(func(context.Context, string) ([]netip.AddrPort, error) { + f.targetResolutions.Add(1) + return nil, errors.New("public vhost fixture forbids UDP resolution") + }), + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = server.Close() }) + f.public, err = url.Parse("https://" + server.TCPAddr().String()) + if err != nil { + t.Fatal(err) + } + checkedRT := webPublicOriginRoundTripFunc(func(r *http.Request) (*http.Response, error) { + response, err := originRT.RoundTrip(r) + if err == nil && response != nil { + if response.ProtoMajor != 1 || response.TLS == nil || len(response.TLS.VerifiedChains) == 0 { + _ = response.Body.Close() + return nil, errors.New("owned origin was not actual verified-TLS H1") + } + f.verifiedOriginResponses.Add(1) + } + return response, err + }) + if mode { + proxy, err = cover.NewReverseProxyHandlerWithPublicOrigin(upstream, f.public, checkedRT) + } else { + proxy, err = cover.NewReverseProxyHandler(upstream, checkedRT) + } + if err != nil { + t.Fatal(err) + } + jar, err := cookiejar.New(nil) + if err != nil { + t.Fatal(err) + } + f.client = &http.Client{Transport: webCoverMetadataPublicTransport(t, clientTLS, proto), Jar: jar, Timeout: 3 * time.Second, CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }} + serveCtx, cancel := context.WithCancel(context.Background()) + originDone, frontDone := make(chan struct{}), make(chan struct{}) + originResult, frontResult := make(chan error, 1), make(chan error, 1) + go func() { + defer close(originDone) + originResult <- origin.Serve(tls.NewListener(webPublicOriginDeadlineListener{listener}, originTLS)) + }() + go func() { defer close(frontDone); frontResult <- server.Serve(serveCtx) }() + t.Cleanup(func() { + admissionMu.Lock() + closing = true + connections := make([]net.Conn, 0, len(hijacked)) + for conn := range hijacked { + connections = append(connections, conn) + } + admissionMu.Unlock() + for _, conn := range connections { + _ = conn.Close() + } + cancel() + _ = server.Close() + _ = origin.Close() + originRT.CloseIdleConnections() + joined := make(chan struct{}) + go func() { defer close(joined); workers.Wait() }() + webCoverMetadataJoin(t, "public-vhost request workers", joined) + webCoverMetadataJoin(t, "public-vhost fixed origin Serve", originDone) + webCoverMetadataJoin(t, "public-vhost combined Serve", frontDone) + select { + case err := <-originResult: + if err != nil && !errors.Is(err, http.ErrServerClosed) { + t.Errorf("origin Serve: %v", err) + } + default: + } + select { + case err := <-frontResult: + if err != nil { + t.Errorf("combined Serve: %v", err) + } + default: + } + if f.invalidDials.Load() != 0 || f.targetDials.Load() != 0 || f.targetResolutions.Load() != 0 { + t.Errorf("non-owned/tunnel dial/UDP resolution=%d/%d/%d", f.invalidDials.Load(), f.targetDials.Load(), f.targetResolutions.Load()) + } + if f.verifiedOriginResponses.Load() != f.appCalls.Load() { + t.Errorf("verified TLS H1 responses/site calls=%d/%d", f.verifiedOriginResponses.Load(), f.appCalls.Load()) + } + for len(f.asyncErrors) > 0 { + t.Errorf("website worker: %v", <-f.asyncErrors) + } + t.Logf("actual site calls=%d, verified fixed-origin TLS H1 responses=%d; non-owned/tunnel dial/UDP resolution=%d/%d/%d", f.appCalls.Load(), f.verifiedOriginResponses.Load(), f.invalidDials.Load(), f.targetDials.Load(), f.targetResolutions.Load()) + }) + return f +} + +func (f *webPublicOriginFixture) exchange(t *testing.T, method, path string, form url.Values, header http.Header, host string) (*http.Response, string) { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + var body io.Reader + if form != nil { + body = strings.NewReader(form.Encode()) + } + request, err := http.NewRequestWithContext(ctx, method, f.public.String()+path, body) + if err != nil { + t.Fatal(err) + } + request.Header = header.Clone() + if request.Header == nil { + request.Header = make(http.Header) + } + request.Header.Set("X-Public-Vhost-Fixture", t.Name()) + request.Header.Set("Proxy-Authorization", "Bearer fictional-invalid-ticket") + if form != nil { + request.Header.Set("Content-Type", "application/x-www-form-urlencoded") + } + if host != "" { + request.Host = host + } + response, err := f.client.Do(request) + if err != nil { + t.Fatal(err) + } + data, err := io.ReadAll(io.LimitReader(response.Body, 32<<10)) + _ = response.Body.Close() + if err != nil { + t.Fatal(err) + } + if response.ProtoMajor != f.proto || response.TLS == nil || len(response.TLS.VerifiedChains) == 0 { + t.Fatalf("actual public response=%s verifiedTLS=%t, want H%d", response.Proto, response.TLS != nil && len(response.TLS.VerifiedChains) != 0, f.proto) + } + if response.Header.Get(webAuthResponseHeader) != "" || response.Header.Get("Proxy-Authenticate") != "" { + t.Error("ordinary website exposed relay authentication metadata") + } + select { + case marker := <-f.completed: + if marker != t.Name() { + t.Errorf("proxy completion marker=%q, want %q", marker, t.Name()) + } + case <-time.After(2 * time.Second): + t.Fatal("public-vhost handler did not complete within independent budget") + } + return response, string(data) +} + +func (f *webPublicOriginFixture) observed(t *testing.T, origin, referer string) webPublicOriginObserved { + t.Helper() + var seen webPublicOriginObserved + select { + case seen = <-f.records: + case <-time.After(2 * time.Second): + t.Fatal("fixed website did not publish an actual request") + } + wantHost := f.upstream.Host + if f.mode { + wantHost = f.public.Host + } + if seen.host != wantHost || seen.sni != webPublicOriginTLSName || seen.proto != 1 { + t.Errorf("fixed website actual Host/SNI/proto=%q/%q/%d, want %q/%q/1", seen.host, seen.sni, seen.proto, wantHost, webPublicOriginTLSName) + } + if seen.header.Get("Origin") != origin || seen.header.Get("Referer") != referer { + t.Errorf("original Origin/Referer changed=%q/%q, want %q/%q", seen.header.Get("Origin"), seen.header.Get("Referer"), origin, referer) + } + return seen +} + +func (f *webPublicOriginFixture) website(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/websocket" { + f.websiteWebSocket(w, r) + return + } + if r.URL.Path == "/metadata" { + for _, value := range webPublicOriginRawCookies { + w.Header().Add("Set-Cookie", value) + } + w.Header().Set("Location", webPublicOriginThirdParty) + w.WriteHeader(302) + _, _ = io.WriteString(w, "location and cookie bytes unchanged") + return + } + if r.Method == http.MethodGet && r.URL.Path == "/form" { + http.SetCookie(w, &http.Cookie{Name: "csrf", Value: webPublicOriginCSRF, Path: "/", Secure: true, HttpOnly: true, SameSite: http.SameSiteLaxMode}) + w.Header().Set("Content-Type", "text/html; charset=utf-8") + _, _ = io.WriteString(w, webPublicOriginForm) + return + } + if r.Method == http.MethodPost && (r.URL.Path == "/login" || r.URL.Path == "/logout") { + // This is independent application policy based on its actual Host, + // not a reimplementation of the proxy's configured-origin guard. + if r.Header.Get("Origin") != "https://"+r.Host { + http.Error(w, "website same-origin policy", 403) + return + } + r.Body = http.MaxBytesReader(w, r.Body, 4096) + csrf, err := r.Cookie("csrf") + if r.ParseForm() != nil || err != nil || csrf.Value != webPublicOriginCSRF || r.PostForm.Get("csrf") != webPublicOriginCSRF { + http.Error(w, "website CSRF policy", 403) + return + } + f.stateMu.Lock() + defer f.stateMu.Unlock() + if r.URL.Path == "/login" { + if r.PostForm.Get("username") != "ordinary-visitor" || r.PostForm.Get("password") != "fictional-local-password" { + http.Error(w, "website credentials", 403) + return + } + f.session = true + http.SetCookie(w, &http.Cookie{Name: "site_session", Value: webPublicOriginSession, Path: "/", Secure: true, HttpOnly: true, SameSite: http.SameSiteLaxMode}) + http.Redirect(w, r, f.public.String()+"/welcome", 303) + return + } + f.session = false + http.SetCookie(w, &http.Cookie{Name: "site_session", Value: "", Path: "/", Secure: true, HttpOnly: true, SameSite: http.SameSiteLaxMode, MaxAge: -1}) + http.Redirect(w, r, f.public.String()+"/form", 303) + return + } + if r.Method == http.MethodGet && r.URL.Path == "/welcome" { + cookie, err := r.Cookie("site_session") + f.stateMu.Lock() + valid := f.session + f.stateMu.Unlock() + if !valid || err != nil || cookie.Value != webPublicOriginSession { + http.Error(w, "website session absent", 401) + return + } + _, _ = io.WriteString(w, webPublicOriginWelcome) + return + } + http.NotFound(w, r) +} + +func (f *webPublicOriginFixture) websiteWebSocket(w http.ResponseWriter, r *http.Request) { + // This is the standard website same-origin predicate (as in Gorilla's + // default): compare the parsed Origin Host to the application's actual + // request Host, with no knowledge of the proxy's public configuration. + if origins := r.Header.Values("Origin"); len(origins) != 0 { + origin, err := url.Parse(origins[0]) + if err != nil || !strings.EqualFold(origin.Host, r.Host) { + http.Error(w, "website WebSocket origin policy", 403) + return + } + } + if r.Method != http.MethodGet || r.ProtoMajor != 1 || r.ProtoMinor != 1 || r.Header.Get("Connection") != "Upgrade" || r.Header.Get("Upgrade") != "websocket" || r.Header.Get("Sec-WebSocket-Version") != "13" || r.Header.Get("Sec-WebSocket-Key") != "dGhlIHNhbXBsZSBub25jZQ==" { + http.Error(w, "website requires a valid observed WebSocket handshake", 400) + return + } + hijacker, ok := w.(http.Hijacker) + if !ok { + http.Error(w, "website cannot hijack", 500) + return + } + conn, rw, err := hijacker.Hijack() + if err != nil { + f.asyncErrors <- err + return + } + if !f.registerHijacked(conn) { + _ = conn.Close() + return + } + defer f.unregisterHijacked(conn) + defer conn.Close() + if err := conn.SetDeadline(time.Now().Add(3 * time.Second)); err != nil { + f.asyncErrors <- err + return + } + accept := sha1.Sum([]byte(r.Header.Get("Sec-WebSocket-Key") + "258EAFA5-E914-47DA-95CA-C5AB0DC85B11")) + _, err = fmt.Fprintf(rw, "HTTP/1.1 101 Switching Protocols\r\nConnection: Upgrade\r\nUpgrade: websocket\r\nSec-WebSocket-Accept: %s\r\n\r\n", base64.StdEncoding.EncodeToString(accept[:])) + if err == nil { + err = rw.Flush() + } + if err != nil { + f.asyncErrors <- err + return + } + header := make([]byte, 6) + if _, err = io.ReadFull(rw, header); err != nil { + f.asyncErrors <- err + return + } + if header[0] != 0x81 || header[1] != 0x84 { + f.asyncErrors <- fmt.Errorf("actual client frame header=%x, want complete masked text frame of four bytes", header[:2]) + return + } + payload := make([]byte, 4) + if _, err = io.ReadFull(rw, payload); err != nil { + f.asyncErrors <- err + return + } + for index := range payload { + payload[index] ^= header[2+index%4] + } + if string(payload) != "echo" { + f.asyncErrors <- fmt.Errorf("actual decoded client frame payload=%q", payload) + return + } + frame := append([]byte{0x81, byte(len(payload))}, payload...) + if n, err := conn.Write(frame); err != nil || n != len(frame) { + f.asyncErrors <- fmt.Errorf("unmasked origin frame write=%d/%v", n, err) + } +} + +func assertWebPublicOriginRedirect(t *testing.T, response *http.Response, location string) { + t.Helper() + if response.StatusCode != 303 || response.Header.Get("Location") != location { + t.Errorf("website redirect=%d/%q, want303/%q", response.StatusCode, response.Header.Get("Location"), location) + } +} +func assertWebPublicOriginCookie(t *testing.T, response *http.Response, name, value string, deleted bool) { + t.Helper() + for _, cookie := range response.Cookies() { + if cookie.Name != name { + continue + } + if cookie.Value != value || cookie.Domain != "" || cookie.Path != "/" || !cookie.Secure || !cookie.HttpOnly || cookie.SameSite != http.SameSiteLaxMode || ((cookie.MaxAge < 0) != deleted) { + t.Errorf("website cookie=%#v", cookie) + } + return + } + t.Errorf("website Set-Cookie missing %s", name) +} + +type webPublicOriginRoundTripFunc func(*http.Request) (*http.Response, error) + +func (f webPublicOriginRoundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } + +type webPublicOriginDeadlineListener struct{ net.Listener } + +func (l webPublicOriginDeadlineListener) Accept() (net.Conn, error) { + conn, err := l.Listener.Accept() + if err == nil { + if err = conn.SetDeadline(time.Now().Add(4 * time.Second)); err != nil { + _ = conn.Close() + return nil, err + } + } + return conn, err +} + +func webPublicOriginTLSConfigs(t *testing.T) (*tls.Config, *tls.Config) { + t.Helper() + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatal(err) + } + now := time.Now() + template := &x509.Certificate{SerialNumber: big.NewInt(now.UnixNano()), Subject: pkix.Name{CommonName: "owned private origin fixture"}, NotBefore: now.Add(-time.Minute), NotAfter: now.Add(time.Hour), KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, BasicConstraintsValid: true, IsCA: true, DNSNames: []string{webPublicOriginTLSName}} + der, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key) + if err != nil { + t.Fatal(err) + } + parsed, err := x509.ParseCertificate(der) + if err != nil { + t.Fatal(err) + } + pool := x509.NewCertPool() + pool.AddCert(parsed) + return &tls.Config{Certificates: []tls.Certificate{{Certificate: [][]byte{der}, PrivateKey: key}}, NextProtos: []string{webHTTP11ALPN}}, &tls.Config{RootCAs: pool, ServerName: webPublicOriginTLSName, NextProtos: []string{webHTTP11ALPN}} +}