diff --git a/catalog/rest/metrics_reporter_test.go b/catalog/rest/metrics_reporter_test.go index 7485135b7..3814fb639 100644 --- a/catalog/rest/metrics_reporter_test.go +++ b/catalog/rest/metrics_reporter_test.go @@ -344,6 +344,7 @@ func reporterWithSession(t *testing.T, auth AuthManager, tr http.RoundTripper) * RoundTripper: tr, authManager: auth, defaultHeaders: http.Header{}, + catalogOrigin: &url.URL{Scheme: "http", Host: "catalog.invalid"}, } session.defaultHeaders.Set("Content-Type", "application/json") @@ -384,6 +385,7 @@ func TestRESTMetricsReporterRefetchesCredentialPerReport(t *testing.T) { RoundTripper: &captureTransport{ch: received}, authManager: &rotatingAuthManager{}, defaultHeaders: http.Header{}, + catalogOrigin: &url.URL{Scheme: "http", Host: "catalog.invalid"}, } rep := reporterWith(t, session, d, nil) diff --git a/catalog/rest/options.go b/catalog/rest/options.go index 884d0a759..e58b2c63a 100644 --- a/catalog/rest/options.go +++ b/catalog/rest/options.go @@ -139,6 +139,10 @@ func WithPrefix(prefix string) Option { // property, and any server-provided signing-region / signing-name overrides. // Use WithSignerFactory (or, for AWS, sigv4.WithAwsConfig) instead if you want // the region and service to stay driven by those settings. +// +// A signer only runs for requests to the catalog origin. Requests to any other +// origin, including an OAuth token endpoint set with WithAuthURI, are sent +// unsigned. func WithSigner(signer RequestSigner) Option { return func(o *options) { o.signer = signer diff --git a/catalog/rest/rest.go b/catalog/rest/rest.go index 2ad477634..9af16c479 100644 --- a/catalog/rest/rest.go +++ b/catalog/rest/rest.go @@ -235,10 +235,43 @@ type sessionTransport struct { authManager AuthManager defaultHeaders http.Header signer RequestSigner - // signingOrigin is the configured catalog origin. A request to a different - // origin (e.g. a redirect hop) is not signed, so the signer's Authorization - // header and any session token never reach an unconfigured host. - signingOrigin *url.URL + + // builtinHeaders is the subset of defaultHeaders under the built-in keys + // (with any operator override applied), which identify the client and + // carry no credentials. It is all a request to an origin other than + // catalogOrigin or authOrigin (e.g. a redirect hop) receives. + builtinHeaders http.Header + // catalogOrigin is the configured catalog origin: the only origin that + // receives the auth header and the signer's signature. authOrigin is the + // configured OAuth token endpoint, if any, which also receives + // user-supplied default headers. + // These are compared against the request URL rather than derived from + // the redirect chain, so a transport that omits Response.Request cannot + // widen them. A nil catalogOrigin trusts no origin. + catalogOrigin *url.URL + authOrigin *url.URL +} + +// maxRedirects matches net/http's default redirect policy, which a custom +// CheckRedirect replaces. +const maxRedirects = 10 + +// sameOriginRedirectsOnly is the CheckRedirect policy for every client the +// catalog builds. A redirect is followed only while it stays on the origin of +// the request that started the chain. A cross-origin hop would otherwise +// replay a 307/308 body (such as the client_credentials form, client_secret +// included) to the other origin, and a hop that redirects back to a different +// catalog path would be sent with the catalog's credentials. +func sameOriginRedirectsOnly(req *http.Request, via []*http.Request) error { + if len(via) >= maxRedirects { + return fmt.Errorf("stopped after %d redirects", maxRedirects) + } + if from := via[0].URL; !sameOrigin(from, req.URL) { + return fmt.Errorf("%w: refusing redirect from %s://%s to a different origin %s://%s", + ErrRESTError, from.Scheme, from.Host, req.URL.Scheme, req.URL.Host) + } + + return nil } // sameOrigin reports whether two URLs share scheme, host, and effective port. @@ -263,12 +296,25 @@ func defaultedPort(u *url.URL) string { } func (s *sessionTransport) RoundTrip(r *http.Request) (*http.Response, error) { + // This transport adds the auth header, the signature and the session + // defaults to every request it sends, including each redirect hop. + // net/http's redirect stripping only covers Authorization and cookies set + // on the request passed to Client.Do, and never custom headers, so these + // are applied only for the configured origins. sameOriginRedirectsOnly + // keeps the clients from following a cross-origin hop at all; this gate + // is the second line of defense. + toCatalog := s.catalogOrigin != nil && sameOrigin(s.catalogOrigin, r.URL) + defaults := s.builtinHeaders + if toCatalog || (s.authOrigin != nil && sameOrigin(s.authOrigin, r.URL)) { + defaults = s.defaultHeaders + } + // A session default is applied unless the request already carries that // header (a per-request override of any default, not just Content-Type // wins) or explicitly opted out of it via withSuppressedHeaders (carried on // the context as an explicit set, never inferred from header values). suppressed := suppressedHeadersFrom(r.Context()) - for k, v := range s.defaultHeaders { + for k, v := range defaults { ck := http.CanonicalHeaderKey(k) if _, ok := r.Header[ck]; ok { continue @@ -287,7 +333,7 @@ func (s *sessionTransport) RoundTrip(r *http.Request) (*http.Response, error) { // session default of the same key. A caller cannot suppress or spoof the // Authorization header by supplying its own. Do not reorder this before the // default-header loop. - if s.authManager != nil && r.Context().Value(skipOAuth) == nil { + if s.authManager != nil && toCatalog && r.Context().Value(skipOAuth) == nil { var ( k, v string err error @@ -306,7 +352,12 @@ func (s *sessionTransport) RoundTrip(r *http.Request) (*http.Response, error) { r.Header.Set(k, v) } - if s.signer != nil && (s.signingOrigin == nil || sameOrigin(s.signingOrigin, r.URL)) { + // A signer only signs requests to the configured catalog origin: a request + // to a different origin (e.g. a redirect hop) is left unsigned, so the + // signer's Authorization header and any session token never reach an + // unconfigured host. The guard lives here, in core, so it covers every + // signer, including one installed verbatim via WithSigner. + if s.signer != nil && toCatalog { if err := s.signer.SignRequest(r); err != nil { return nil, err } @@ -953,7 +1004,11 @@ func setupOAuthManager(r *Catalog, cl *http.Client, opts *options) (AuthManager, // this reuses the catalog client's transport — preserving its TLS, proxy and // header behavior — but as a distinct *http.Client so the Timeout applies to // refresh alone and not to ordinary catalog requests. - oauthClient := &http.Client{Transport: cl.Transport, Timeout: defaultOAuthTimeout} + oauthClient := &http.Client{ + Transport: cl.Transport, + Timeout: defaultOAuthTimeout, + CheckRedirect: sameOriginRedirectsOnly, + } var closeIdleConnections func() if opts.oauthTLSConfig != nil { transport := &http.Transport{ @@ -961,8 +1016,9 @@ func setupOAuthManager(r *Catalog, cl *http.Client, opts *options) (AuthManager, TLSClientConfig: opts.oauthTLSConfig, } oauthClient = &http.Client{ - Transport: transport, - Timeout: defaultOAuthTimeout, + Transport: transport, + Timeout: defaultOAuthTimeout, + CheckRedirect: sameOriginRedirectsOnly, } closeIdleConnections = transport.CloseIdleConnections } @@ -1020,6 +1076,19 @@ func (r *Catalog) init(ctx context.Context, ops *options, uri string) error { return nil } +// builtinHeaderDefaults returns the headers every session sends by default. +// They identify the client and carry no credentials, so they are the only +// defaults a request to an unconfigured origin receives (see +// sessionTransport.builtinHeaders). +func builtinHeaderDefaults() []struct{ key, value string } { + return []struct{ key, value string }{ + {"X-Client-Version", icebergRestSpecVersion}, + {"Content-Type", "application/json"}, + {"User-Agent", "GoIceberg/" + iceberg.Version()}, + {headerIcebergAccessDelegation, defaultAccessDelegation}, + } +} + // createSession returns a cleanup that closes only transports created by this // function, never transports supplied by the caller. func (r *Catalog) createSession(ctx context.Context, opts *options) (*http.Client, func(), error) { @@ -1047,8 +1116,10 @@ func (r *Catalog) createSession(ctx context.Context, opts *options) (*http.Clien session := &sessionTransport{ RoundTripper: baseTransport, defaultHeaders: http.Header{}, + catalogOrigin: r.baseURI, + authOrigin: opts.authUri, } - cl := &http.Client{Transport: session} + cl := &http.Client{Transport: session, CheckRedirect: sameOriginRedirectsOnly} // If the user does not set an AuthManager, construct one for this session // without storing it in opts. Bootstrap authentication must not leak into @@ -1062,10 +1133,10 @@ func (r *Catalog) createSession(ctx context.Context, opts *options) (*http.Clien } } - session.defaultHeaders.Set("X-Client-Version", icebergRestSpecVersion) - session.defaultHeaders.Set("Content-Type", "application/json") - session.defaultHeaders.Set("User-Agent", "GoIceberg/"+iceberg.Version()) - session.defaultHeaders.Set(headerIcebergAccessDelegation, defaultAccessDelegation) + builtins := builtinHeaderDefaults() + for _, h := range builtins { + session.defaultHeaders.Set(h.key, h.value) + } for k, v := range opts.headers { session.defaultHeaders.Set(k, v) @@ -1077,6 +1148,13 @@ func (r *Catalog) createSession(ctx context.Context, opts *options) (*http.Clien } } + session.builtinHeaders = http.Header{} + for _, h := range builtins { + if v := session.defaultHeaders.Values(h.key); len(v) > 0 { + session.builtinHeaders[http.CanonicalHeaderKey(h.key)] = slices.Clone(v) + } + } + if authManager != nil { session.authManager = authManager } @@ -1088,14 +1166,6 @@ func (r *Catalog) createSession(ctx context.Context, opts *options) (*http.Clien return nil, nil, err } session.signer = signer - // A signer only signs requests to the configured catalog origin: a request - // to a different origin (e.g. a redirect hop) is left unsigned, so the - // signer's Authorization header and any session token never reach an - // unconfigured host. The guard lives here, in core, so it covers every - // signer, including one installed verbatim via WithSigner. - if signer != nil { - session.signingOrigin = r.baseURI - } return cl, cleanup, nil } diff --git a/catalog/rest/rest_internal_test.go b/catalog/rest/rest_internal_test.go index c3eb75da9..057f3e565 100644 --- a/catalog/rest/rest_internal_test.go +++ b/catalog/rest/rest_internal_test.go @@ -58,16 +58,31 @@ func (markingSigner) SignRequest(r *http.Request) error { return nil } -// TestSignerDoesNotSignCrossOriginRedirect pins the origin guard in -// sessionTransport.RoundTrip: a redirect to a different origin is not signed, so +// serveEmptyConfig answers GET /v1/config with no defaults or overrides. +func serveEmptyConfig(w http.ResponseWriter, _ *http.Request) { + _ = json.NewEncoder(w).Encode(map[string]any{"defaults": map[string]any{}, "overrides": map[string]any{}}) +} + +// roundTripTo sends a GET for target through the catalog session's transport +// directly, bypassing the client's redirect policy, so a test can exercise the +// transport's own origin gate. +func roundTripTo(t *testing.T, cat *Catalog, target string) { + t.Helper() + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, target, nil) + require.NoError(t, err) + resp, err := cat.cl.Transport.RoundTrip(req) + require.NoError(t, err) + require.NoError(t, resp.Body.Close()) +} + +// TestSignerDoesNotSignCrossOriginRequest pins the origin guard in +// sessionTransport.RoundTrip: a request to a different origin is not signed, so // the signer's Authorization header and any session token never reach an // unconfigured host. The guard is signer-agnostic, so it covers an explicit // WithSigner too; a marking signer exercises it without a real SigV4 backend. -func TestSignerDoesNotSignCrossOriginRedirect(t *testing.T) { - var secondHit bool +func TestSignerDoesNotSignCrossOriginRequest(t *testing.T) { var gotAuth, gotToken string second := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - secondHit = true gotAuth = r.Header.Get("Authorization") gotToken = r.Header.Get("X-Amz-Security-Token") w.WriteHeader(http.StatusOK) @@ -76,12 +91,10 @@ func TestSignerDoesNotSignCrossOriginRedirect(t *testing.T) { var firstAuth string mux := http.NewServeMux() - mux.HandleFunc("/v1/config", func(w http.ResponseWriter, r *http.Request) { - _ = json.NewEncoder(w).Encode(map[string]any{"defaults": map[string]any{}, "overrides": map[string]any{}}) - }) - mux.HandleFunc("/redirect", func(w http.ResponseWriter, r *http.Request) { + mux.HandleFunc("/v1/config", serveEmptyConfig) + mux.HandleFunc("/ok", func(w http.ResponseWriter, r *http.Request) { firstAuth = r.Header.Get("Authorization") - http.Redirect(w, r, second.URL+"/landing", http.StatusTemporaryRedirect) + w.WriteHeader(http.StatusOK) }) first := httptest.NewServer(mux) defer first.Close() @@ -90,16 +103,365 @@ func TestSignerDoesNotSignCrossOriginRedirect(t *testing.T) { require.NoError(t, err) t.Cleanup(func() { _ = cat.Close() }) - req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, first.URL+"/redirect", nil) + roundTripTo(t, cat, first.URL+"/ok") + roundTripTo(t, cat, second.URL+"/landing") + + require.NotEmpty(t, firstAuth, "the configured origin must be signed") + require.Empty(t, gotAuth, "another origin must not receive the Authorization header") + require.Empty(t, gotToken, "another origin must not receive the session token") +} + +// TestCredentialsNotSentCrossOrigin pins that the OAuth bearer token and +// user-supplied default headers are only added for the catalog origin, while +// built-in client headers go everywhere, including an operator override of +// one of them. +func TestCredentialsNotSentCrossOrigin(t *testing.T) { + var other http.Header + second := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + other = r.Header.Clone() + w.WriteHeader(http.StatusOK) + })) + defer second.Close() + + var first http.Header + mux := http.NewServeMux() + mux.HandleFunc("/v1/config", serveEmptyConfig) + mux.HandleFunc("/ok", func(w http.ResponseWriter, r *http.Request) { + first = r.Header.Clone() + w.WriteHeader(http.StatusOK) + }) + srv := httptest.NewServer(mux) + defer srv.Close() + + cat, err := NewCatalog(context.Background(), "rest", srv.URL, + WithOAuthToken("SECRET-CATALOG-TOKEN"), + WithHeaders(map[string]string{"X-Custom": "SECRET-CUSTOM", "User-Agent": "corp-agent/1"}), + WithAdditionalProps(iceberg.Properties{ + "header.X-Api-Key": "SECRET-API-KEY", + "header." + headerIcebergAccessDelegation: "remote-signing", + })) + require.NoError(t, err) + + roundTripTo(t, cat, srv.URL+"/ok") + roundTripTo(t, cat, second.URL+"/landing") + + assert.Equal(t, "Bearer SECRET-CATALOG-TOKEN", first.Get("Authorization")) + assert.Equal(t, "SECRET-API-KEY", first.Get("X-Api-Key")) + assert.Equal(t, "SECRET-CUSTOM", first.Get("X-Custom")) + + assert.Empty(t, other.Get("Authorization"), "another origin must not receive the bearer token") + assert.Empty(t, other.Get("X-Api-Key"), "another origin must not receive header.* defaults") + assert.Empty(t, other.Get("X-Custom"), "another origin must not receive WithHeaders defaults") + assert.Equal(t, "corp-agent/1", other.Get("User-Agent"), "an overridden built-in header keeps the override") + assert.Equal(t, "remote-signing", other.Get(headerIcebergAccessDelegation)) + assert.Equal(t, icebergRestSpecVersion, other.Get("X-Client-Version"), "built-in client headers are still sent") +} + +// TestCrossOriginRedirectNotFollowed pins that the catalog client refuses a +// redirect to another origin before contacting it, while a same-origin +// redirect is still followed with credentials. +func TestCrossOriginRedirectNotFollowed(t *testing.T) { + var otherHit atomic.Bool + second := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + otherHit.Store(true) + w.WriteHeader(http.StatusOK) + })) + defer second.Close() + + var landingAuth string + mux := http.NewServeMux() + mux.HandleFunc("/v1/config", serveEmptyConfig) + mux.HandleFunc("/cross", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, second.URL+"/landing", http.StatusTemporaryRedirect) + }) + mux.HandleFunc("/same", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "/landing", http.StatusTemporaryRedirect) + }) + // The loop ends after maxRedirects*2 hops, so a missing redirect limit fails + // the test instead of hanging it. + var loopHops atomic.Int32 + mux.HandleFunc("/loop", func(w http.ResponseWriter, r *http.Request) { + if loopHops.Add(1) > maxRedirects*2 { + w.WriteHeader(http.StatusOK) + + return + } + http.Redirect(w, r, "/loop", http.StatusTemporaryRedirect) + }) + mux.HandleFunc("/landing", func(w http.ResponseWriter, r *http.Request) { + landingAuth = r.Header.Get("Authorization") + w.WriteHeader(http.StatusOK) + }) + srv := httptest.NewServer(mux) + defer srv.Close() + + cat, err := NewCatalog(context.Background(), "rest", srv.URL, WithOAuthToken("SECRET-CATALOG-TOKEN")) require.NoError(t, err) - resp, err := cat.cl.Do(req) + + get := func(t *testing.T, path string) (*http.Response, error) { + t.Helper() + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, srv.URL+path, nil) + require.NoError(t, err) + + return cat.cl.Do(req) + } + + t.Run("cross origin", func(t *testing.T) { + _, err := get(t, "/cross") + require.ErrorIs(t, err, ErrRESTError) + assert.False(t, otherHit.Load(), "the redirect target must not be contacted") + }) + + t.Run("same origin", func(t *testing.T) { + resp, err := get(t, "/same") + require.NoError(t, err) + require.NoError(t, resp.Body.Close()) + assert.Equal(t, "Bearer SECRET-CATALOG-TOKEN", landingAuth) + }) + + t.Run("loop", func(t *testing.T) { + _, err := get(t, "/loop") + require.ErrorContains(t, err, "stopped after 10 redirects") + }) +} + +// TestSyntheticCrossOriginRedirectNotFollowed pins that the redirect policy +// does not depend on the transport linking Response.Request: a custom +// transport that returns a bare 307 to another origin is not followed either. +func TestSyntheticCrossOriginRedirectNotFollowed(t *testing.T) { + var otherHit bool + transport := roundTripFunc(func(r *http.Request) (*http.Response, error) { + switch { + case r.URL.Host == "other.test": + otherHit = true + + return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil + case r.URL.Path == "/v1/config": + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader([]byte(`{"defaults":{},"overrides":{}}`))), + }, nil + default: + // A synthetic redirect: no Request backlink on the response. + return &http.Response{ + StatusCode: http.StatusTemporaryRedirect, + Header: http.Header{"Location": {"http://other.test/landing"}}, + Body: http.NoBody, + }, nil + } + }) + + cat, err := NewCatalog(context.Background(), "rest", "http://catalog.test", + WithCustomTransport(transport), + WithOAuthToken("SECRET-CATALOG-TOKEN")) require.NoError(t, err) - require.NoError(t, resp.Body.Close()) - require.NotEmpty(t, firstAuth, "the configured origin must be signed") - require.True(t, secondHit, "the redirect target must be reached") - require.Empty(t, gotAuth, "the redirect target must not receive the Authorization header") - require.Empty(t, gotToken, "the redirect target must not receive the session token") + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "http://catalog.test/redirect", nil) + require.NoError(t, err) + _, err = cat.cl.Do(req) + require.ErrorIs(t, err, ErrRESTError) + assert.False(t, otherHit, "the redirect target must not be contacted") +} + +// TestTokenRequestCrossOriginRedirectNotFollowed pins that a token endpoint +// redirecting to another origin does not get the client_credentials form, +// client_secret included, replayed there, on every OAuth client the catalog +// builds. +func TestTokenRequestCrossOriginRedirectNotFollowed(t *testing.T) { + var evilHit atomic.Bool + evil := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + evilHit.Store(true) + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{"access_token": "EVIL", "token_type": "Bearer", "expires_in": 3600}) + })) + defer evil.Close() + + redirectToEvil := func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, evil.URL+"/steal", http.StatusTemporaryRedirect) + } + + mux := http.NewServeMux() + mux.HandleFunc("/v1/config", serveEmptyConfig) + mux.HandleFunc("/v1/oauth/tokens", redirectToEvil) + srv := httptest.NewServer(mux) + defer srv.Close() + + tokenSrv := httptest.NewServer(http.HandlerFunc(redirectToEvil)) + defer tokenSrv.Close() + tlsTokenSrv := httptest.NewTLSServer(http.HandlerFunc(redirectToEvil)) + defer tlsTokenSrv.Close() + + authURI, err := url.Parse(tokenSrv.URL + "/token") + require.NoError(t, err) + tlsAuthURI, err := url.Parse(tlsTokenSrv.URL + "/token") + require.NoError(t, err) + tlsTransport, ok := tlsTokenSrv.Client().Transport.(*http.Transport) + require.True(t, ok) + + for _, tc := range []struct { + name string + opts []Option + }{ + {"catalog token endpoint", nil}, + {"separate token endpoint", []Option{WithAuthURI(authURI)}}, + {"separate oauth tls config", []Option{ + WithAuthURI(tlsAuthURI), + WithOAuthTLSConfig(tlsTransport.TLSClientConfig.Clone()), + }}, + } { + t.Run(tc.name, func(t *testing.T) { + evilHit.Store(false) + opts := append([]Option{WithCredential("client:SECRET-CLIENT-SECRET")}, tc.opts...) + _, err := NewCatalog(context.Background(), "rest", srv.URL, opts...) + require.ErrorIs(t, err, ErrRESTError) + assert.False(t, evilHit.Load(), "the redirect target must not receive the token request") + }) + } +} + +// TestDropTableCrossOriginRedirectNotFollowed pins that a DELETE redirected to +// another origin is not followed, so that origin cannot bounce it back to the +// catalog as a purge of a different table carrying the bearer. +func TestDropTableCrossOriginRedirectNotFollowed(t *testing.T) { + var srvURL string + var evilHit, victimHit atomic.Bool + evil := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + evilHit.Store(true) + http.Redirect(w, r, srvURL+"/v1/namespaces/ns/tables/victim?purgeRequested=true", http.StatusTemporaryRedirect) + })) + defer evil.Close() + + mux := http.NewServeMux() + mux.HandleFunc("/v1/config", serveEmptyConfig) + mux.HandleFunc("/v1/namespaces/ns/tables/t", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, evil.URL+"/bounce", http.StatusTemporaryRedirect) + }) + mux.HandleFunc("/v1/namespaces/ns/tables/victim", func(w http.ResponseWriter, r *http.Request) { + victimHit.Store(true) + w.WriteHeader(http.StatusNoContent) + }) + srv := httptest.NewServer(mux) + defer srv.Close() + srvURL = srv.URL + + cat, err := NewCatalog(context.Background(), "rest", srv.URL, WithOAuthToken("SECRET-CATALOG-TOKEN")) + require.NoError(t, err) + + err = cat.DropTable(context.Background(), table.Identifier{"ns", "t"}) + require.ErrorIs(t, err, ErrRESTError) + assert.False(t, evilHit.Load(), "the redirect target must not be contacted") + assert.False(t, victimHit.Load(), "a different table must not be dropped") +} + +// TestConfigURIOverrideKeepsCredentials pins that a /v1/config uri override +// onto another origin moves the trusted origin with it. The origin gate relies +// on init creating the final session after fetchConfig applies the override. +func TestConfigURIOverrideKeepsCredentials(t *testing.T) { + var got http.Header + second := http.NewServeMux() + second.HandleFunc("/v1/oauth/tokens", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{"access_token": "SECOND-TOKEN", "token_type": "Bearer", "expires_in": 3600}) + }) + second.HandleFunc("/v1/namespaces", func(w http.ResponseWriter, r *http.Request) { + got = r.Header.Clone() + _ = json.NewEncoder(w).Encode(map[string]any{"namespaces": [][]string{}}) + }) + secondSrv := httptest.NewServer(second) + defer secondSrv.Close() + + first := http.NewServeMux() + first.HandleFunc("/v1/oauth/tokens", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{"access_token": "FIRST-TOKEN", "token_type": "Bearer", "expires_in": 3600}) + }) + first.HandleFunc("/v1/config", func(w http.ResponseWriter, r *http.Request) { + _ = json.NewEncoder(w).Encode(map[string]any{"defaults": map[string]any{}, "overrides": map[string]any{"uri": secondSrv.URL}}) + }) + firstSrv := httptest.NewServer(first) + defer firstSrv.Close() + + for _, tc := range []struct { + name string + opt Option + header string + wantHeader string + }{ + {"static token", WithOAuthToken("STATIC-TOKEN"), "Authorization", "Bearer STATIC-TOKEN"}, + {"client credentials", WithCredential("client:secret"), "Authorization", "Bearer SECOND-TOKEN"}, + {"signer", WithSigner(markingSigner{}), "X-Amz-Security-Token", "SENTINELTOKEN"}, + } { + t.Run(tc.name, func(t *testing.T) { + got = nil + cat, err := NewCatalog(context.Background(), "rest", firstSrv.URL, tc.opt, + WithAdditionalProps(iceberg.Properties{"header.X-Api-Key": "SECRET-API-KEY"})) + require.NoError(t, err) + + _, err = cat.ListNamespaces(context.Background(), nil) + require.NoError(t, err) + require.NotNil(t, got, "the overridden origin must be reached") + assert.Equal(t, tc.wantHeader, got.Get(tc.header)) + assert.Equal(t, "SECRET-API-KEY", got.Get("X-Api-Key")) + }) + } +} + +func TestSameOrigin(t *testing.T) { + for _, tc := range []struct { + a, b string + want bool + }{ + {"https://catalog.example.com/v1", "https://catalog.example.com/other?x=1", true}, + {"https://catalog.example.com", "https://CATALOG.example.com", true}, + {"HTTPS://catalog.example.com", "https://catalog.example.com", true}, + {"https://catalog.example.com", "https://catalog.example.com:443", true}, + {"http://catalog.example.com", "http://catalog.example.com:80", true}, + {"https://catalog.example.com", "http://catalog.example.com", false}, + {"https://catalog.example.com", "https://catalog.example.com:8443", false}, + {"https://example.com", "https://sub.example.com", false}, + {"https://sub.example.com", "https://example.com", false}, + {"https://catalog.example.com", "https://other.example.com", false}, + } { + a, err := url.Parse(tc.a) + require.NoError(t, err) + b, err := url.Parse(tc.b) + require.NoError(t, err) + assert.Equal(t, tc.want, sameOrigin(a, b), "%s vs %s", tc.a, tc.b) + } +} + +// TestHeaderDefaultsReachSeparateTokenEndpoint pins that the configured OAuth +// token endpoint is trusted for header.* defaults even on a different origin +// from the catalog, so the redirect guard does not break IdPs that need them. +func TestHeaderDefaultsReachSeparateTokenEndpoint(t *testing.T) { + var tokenAPIKey string + tokenSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + tokenAPIKey = r.Header.Get("X-Api-Key") + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{"access_token": "TOKEN", "token_type": "Bearer", "expires_in": 3600}) + })) + defer tokenSrv.Close() + + var catalogAuth string + mux := http.NewServeMux() + mux.HandleFunc("/v1/config", func(w http.ResponseWriter, r *http.Request) { + catalogAuth = r.Header.Get("Authorization") + json.NewEncoder(w).Encode(map[string]any{"defaults": map[string]any{}, "overrides": map[string]any{}}) + }) + srv := httptest.NewServer(mux) + defer srv.Close() + + authURI, err := url.Parse(tokenSrv.URL + "/token") + require.NoError(t, err) + + _, err = NewCatalog(context.Background(), "rest", srv.URL, + WithCredential("client:secret"), + WithAuthURI(authURI), + WithAdditionalProps(iceberg.Properties{"header.X-Api-Key": "SECRET-API-KEY"})) + require.NoError(t, err) + + assert.Equal(t, "Bearer TOKEN", catalogAuth) + assert.Equal(t, "SECRET-API-KEY", tokenAPIKey, "the configured token endpoint must still receive header.* defaults") } func TestSplitIdentForPathRequiresNamespaceAndName(t *testing.T) { @@ -891,6 +1253,7 @@ func TestRoundTripDefaultHeaderHandling(t *testing.T) { return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil }), defaultHeaders: http.Header{}, + catalogOrigin: &url.URL{Scheme: "http", Host: "example.com"}, } s.defaultHeaders.Set(headerIcebergAccessDelegation, defaultAccessDelegation) @@ -956,6 +1319,7 @@ func TestRoundTripAuthManagerWinsOverPerRequestHeader(t *testing.T) { return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil }), defaultHeaders: http.Header{}, + catalogOrigin: &url.URL{Scheme: "http", Host: "example.com"}, authManager: staticAuthManager{key: "Authorization", value: "Bearer managed-token"}, } @@ -1884,6 +2248,7 @@ func TestSessionTransportInvokesSigner(t *testing.T) { return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil }), defaultHeaders: http.Header{}, + catalogOrigin: &url.URL{Scheme: "http", Host: "example.com"}, signer: signerFunc(func(r *http.Request) error { signed = true r.Header.Set("X-Signed", "yes") @@ -1900,6 +2265,37 @@ func TestSessionTransportInvokesSigner(t *testing.T) { assert.Equal(t, "yes", got.Get("X-Signed")) } +// TestSessionTransportNilCatalogOriginTrustsNothing pins that the origin gate +// fails closed: a sessionTransport built without a catalogOrigin adds no auth +// header, signature or user-supplied default to any request. +func TestSessionTransportNilCatalogOriginTrustsNothing(t *testing.T) { + t.Parallel() + + var got http.Header + s := &sessionTransport{ + RoundTripper: roundTripFunc(func(r *http.Request) (*http.Response, error) { + got = r.Header.Clone() + + return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil + }), + authManager: staticAuthManager{key: "Authorization", value: "Bearer managed-token"}, + defaultHeaders: http.Header{"X-Api-Key": {"SECRET-API-KEY"}}, + signer: signerFunc(func(r *http.Request) error { + r.Header.Set("X-Signed", "yes") + + return nil + }), + } + + req, err := http.NewRequest(http.MethodGet, "http://example.com", nil) + require.NoError(t, err) + _, err = s.RoundTrip(req) + require.NoError(t, err) + assert.Empty(t, got.Get("Authorization")) + assert.Empty(t, got.Get("X-Signed")) + assert.Empty(t, got.Get("X-Api-Key")) +} + func TestSessionTransportSignerErrorAbortsRequest(t *testing.T) { t.Parallel() @@ -1911,6 +2307,7 @@ func TestSessionTransportSignerErrorAbortsRequest(t *testing.T) { return nil, nil }), defaultHeaders: http.Header{}, + catalogOrigin: &url.URL{Scheme: "http", Host: "example.com"}, signer: signerFunc(func(_ *http.Request) error { return wantErr }), } @@ -1934,6 +2331,7 @@ func TestSessionTransportConcurrentRoundTrip(t *testing.T) { return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil }), defaultHeaders: http.Header{"X-Default": {"v"}}, + catalogOrigin: &url.URL{Scheme: "http", Host: "example.com"}, signer: signerFunc(func(r *http.Request) error { count.Add(1) r.Header.Set("X-Signed", "yes")