diff --git a/catalog/rest/options.go b/catalog/rest/options.go index bf25b499d..884d0a759 100644 --- a/catalog/rest/options.go +++ b/catalog/rest/options.go @@ -23,7 +23,6 @@ import ( "net/url" "github.com/apache/iceberg-go" - "github.com/aws/aws-sdk-go-v2/aws" ) type Option func(*options) @@ -132,10 +131,35 @@ func WithPrefix(prefix string) Option { } } -func WithAwsConfig(cfg aws.Config) Option { +// WithSigner installs a fully constructed RequestSigner that signs every +// catalog request in place (for example AWS SigV4). It is the escape hatch for +// callers that build their own signer: it is used verbatim and therefore +// bypasses the SigV4 machinery entirely, taking precedence over +// WithSignerFactory, WithSigV4 / WithSigV4RegionSvc, the rest.sigv4-enabled +// 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. +func WithSigner(signer RequestSigner) Option { return func(o *options) { - o.awsConfig = cfg - o.awsConfigSet = true + o.signer = signer + } +} + +// WithSignerFactory installs a factory that builds the RequestSigner from the +// resolved signing configuration. Unlike WithSigner, which takes an +// already-built signer, a factory lets WithSigV4 / WithSigV4RegionSvc and any +// server-provided /v1/config overrides remain the single source of the signing +// region and service: they are folded into the SignerConfig passed to the +// factory. Optional backends such as +// github.com/apache/iceberg-go/catalog/rest/sigv4 (via sigv4.WithAwsConfig) use +// this so the AWS SDK stays out of the core package. +// +// The factory is consulted only when SigV4 signing is enabled (WithSigV4, +// WithSigV4RegionSvc, or the rest.sigv4-enabled property); on its own it does +// not enable signing. An explicit WithSigner still takes precedence over it. +func WithSignerFactory(factory SignerFactory) Option { + return func(o *options) { + o.signerFactory = factory } } @@ -185,8 +209,8 @@ func WithTransportFactory(factory TransportFactory) Option { } type options struct { - awsConfig aws.Config - awsConfigSet bool + signer RequestSigner + signerFactory SignerFactory tlsConfig *tls.Config oauthToken string credential string diff --git a/catalog/rest/rest.go b/catalog/rest/rest.go index be4b80d79..77c40c631 100644 --- a/catalog/rest/rest.go +++ b/catalog/rest/rest.go @@ -20,13 +20,10 @@ package rest import ( "bytes" "context" - "crypto/sha256" "crypto/tls" - "encoding/hex" "encoding/json" "errors" "fmt" - "hash" "io" "iter" "log/slog" @@ -41,16 +38,11 @@ import ( "github.com/apache/iceberg-go" "github.com/apache/iceberg-go/catalog" - internalaws "github.com/apache/iceberg-go/internal/awsconfig" iceio "github.com/apache/iceberg-go/io" "github.com/apache/iceberg-go/metrics" "github.com/apache/iceberg-go/table" "github.com/apache/iceberg-go/udf" "github.com/apache/iceberg-go/view" - "github.com/aws/aws-sdk-go-v2/aws" - v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4" - "github.com/aws/aws-sdk-go-v2/config" - "github.com/aws/aws-sdk-go-v2/credentials" "golang.org/x/oauth2" "golang.org/x/oauth2/clientcredentials" "golang.org/x/sync/semaphore" @@ -94,12 +86,6 @@ const ( keyRestSigV4Region = "rest.signing-region" keyRestSigV4Service = "rest.signing-name" keyAuthUrl = "rest.authorization-url" - // keyRestAccessKeyID and friends are the Java-client property names for the - // SigV4 signing credentials. They are accepted as aliases for the s3.* - // properties; the s3.* keys take precedence when both are set. - keyRestAccessKeyID = "rest.access-key-id" - keyRestSecretAccessKey = "rest.secret-access-key" - keyRestSessionToken = "rest.session-token" // keyOAuth2ServerURI is the portable, spec-aligned property for the OAuth2 // token endpoint used by Java, PyIceberg and iceberg-rust. It is the // preferred key; keyAuthUrl is retained as a compatibility alias. When both @@ -247,13 +233,10 @@ type sessionTransport struct { authManager AuthManager defaultHeaders http.Header - signer v4.HTTPSigner - cfg aws.Config - service string - newHash func() hash.Hash - // signingOrigin is the configured catalog origin. Requests to a different - // origin (e.g. a redirect hop) are not signed, so the SigV4 Authorization - // header and session token never reach an unconfigured host. + 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 } @@ -278,9 +261,6 @@ func defaultedPort(u *url.URL) string { } } -// from https://pkg.go.dev/github.com/aws/aws-sdk-go-v2/aws/signer/v4#Signer.SignHTTP -const emptyStringHash = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855" - func (s *sessionTransport) RoundTrip(r *http.Request) (*http.Response, error) { // A session default is applied unless the request already carries that // header (a per-request override of any default, not just Content-Type @@ -326,44 +306,7 @@ func (s *sessionTransport) RoundTrip(r *http.Request) (*http.Response, error) { } if s.signer != nil && (s.signingOrigin == nil || sameOrigin(s.signingOrigin, r.URL)) { - var payloadHash string - if r.Body == nil { - payloadHash = emptyStringHash - } else { - rdr, err := r.GetBody() - if err != nil { - return nil, err - } - - h := s.newHash() - if _, err = io.Copy(h, rdr); err != nil { - if closeErr := rdr.Close(); closeErr != nil { - err = errors.Join(err, closeErr) - } - - return nil, err - } - - if err = rdr.Close(); err != nil { - return nil, err - } - - payloadHash = hex.EncodeToString(h.Sum(nil)) - } - - creds, err := s.cfg.Credentials.Retrieve(r.Context()) - if err != nil { - return nil, err - } - - // Set the x-amz-content-sha256 header before signing. - // This header is required for AWS SigV4 signature verification. - r.Header.Set("x-amz-content-sha256", payloadHash) - - // modifies the request in place - err = s.signer.SignHTTP(r.Context(), creds, r, payloadHash, - s.service, s.cfg.Region, time.Now()) - if err != nil { + if err := s.signer.SignRequest(r); err != nil { return nil, err } } @@ -1137,67 +1080,25 @@ func (r *Catalog) createSession(ctx context.Context, opts *options) (*http.Clien session.authManager = authManager } - if opts.enableSigv4 { - cfg := opts.awsConfig - if !opts.awsConfigSet { - creds, err := staticCredsFromProps(opts.additionalProps) - if err != nil { - cleanup() - - return nil, nil, err - } - // If no config provided, load defaults from environment. - cfg, err = config.LoadDefaultConfig(ctx) - if err != nil { - cleanup() - - return nil, nil, err - } - // Sign with the S3 credentials carried in the catalog properties when - // present, rather than only the AWS default credential chain. - if creds != nil { - cfg.Credentials = creds - } - } - if opts.sigv4Region != "" { - cfg.Region = opts.sigv4Region - } + signer, err := resolveSigner(ctx, opts) + if err != nil { + cleanup() - session.cfg, session.service = cfg, opts.sigv4Service - session.signer, session.newHash = v4.NewSigner(), sha256.New + 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 } -// staticCredsFromProps returns a static credentials provider built from the -// signing-credential properties. It prefers the s3.* keys and falls back to the -// Java-compatible rest.* aliases, resolving the tuple atomically from a single -// namespace so a partial pair is never completed with fields from the other one. -// It returns (nil, nil) when neither namespace sets any credential property, so -// the caller falls back to the default credential chain, and an -// ErrIncompleteStaticCredentials error when the chosen namespace is incomplete. -func staticCredsFromProps(props iceberg.Properties) (aws.CredentialsProvider, error) { - namespaces := [][3]string{ - {iceio.S3AccessKeyID, iceio.S3SecretAccessKey, iceio.S3SessionToken}, - {keyRestAccessKeyID, keyRestSecretAccessKey, keyRestSessionToken}, - } - for _, ns := range namespaces { - accessKey, secretKey, token := props[ns[0]], props[ns[1]], props[ns[2]] - if accessKey == "" && secretKey == "" && token == "" { - continue - } - if err := internalaws.ValidateStaticCredentials(ns[0], ns[1], ns[2], accessKey, secretKey, token); err != nil { - return nil, err - } - - return credentials.NewStaticCredentialsProvider(accessKey, secretKey, token), nil - } - - return nil, nil -} - func (r *Catalog) fetchConfig(ctx context.Context, opts *options) (*options, error) { params := url.Values{} if opts.warehouseLocation != "" { diff --git a/catalog/rest/rest_internal_test.go b/catalog/rest/rest_internal_test.go index 1018c1965..92da26362 100644 --- a/catalog/rest/rest_internal_test.go +++ b/catalog/rest/rest_internal_test.go @@ -23,10 +23,8 @@ import ( "crypto/ecdsa" "crypto/elliptic" "crypto/rand" - "crypto/sha256" "crypto/tls" "crypto/x509" - "encoding/hex" "encoding/json" "errors" "fmt" @@ -36,134 +34,36 @@ import ( "net/http" "net/http/httptest" "net/url" + "sync" "sync/atomic" "testing" "time" "github.com/apache/iceberg-go" "github.com/apache/iceberg-go/catalog" - internalaws "github.com/apache/iceberg-go/internal/awsconfig" - iceio "github.com/apache/iceberg-go/io" "github.com/apache/iceberg-go/table" - "github.com/aws/aws-sdk-go-v2/aws" - v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4" - "github.com/aws/aws-sdk-go-v2/config" - "github.com/aws/aws-sdk-go-v2/credentials" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "golang.org/x/sync/errgroup" ) -func TestStaticCredsFromProps(t *testing.T) { - creds, err := staticCredsFromProps(iceberg.Properties{ - iceio.S3AccessKeyID: "AK", - iceio.S3SecretAccessKey: "SK", - iceio.S3SessionToken: "ST", - }) - require.NoError(t, err) - require.NotNil(t, creds) - got, err := creds.Retrieve(context.Background()) - require.NoError(t, err) - require.Equal(t, "AK", got.AccessKeyID) - require.Equal(t, "SK", got.SecretAccessKey) - require.Equal(t, "ST", got.SessionToken) - - creds, err = staticCredsFromProps(iceberg.Properties{}) - require.NoError(t, err, "no creds must fall back to the default chain") - require.Nil(t, creds) - - _, err = staticCredsFromProps(iceberg.Properties{iceio.S3AccessKeyID: "AK"}) - require.ErrorIs(t, err, internalaws.ErrIncompleteStaticCredentials, "a lone access key must be an error, not the ambient identity") +// markingSigner marks every request it signs with sentinel headers, so a test +// can assert exactly which origins the session transport chose to sign without +// pulling in a cloud SDK. +type markingSigner struct{} - _, err = staticCredsFromProps(iceberg.Properties{iceio.S3SecretAccessKey: "SK"}) - require.ErrorIs(t, err, internalaws.ErrIncompleteStaticCredentials, "a lone secret key must be an error, not the ambient identity") +func (markingSigner) SignRequest(r *http.Request) error { + r.Header.Set("Authorization", "AWS4-HMAC-SHA256 Credential=SENTINELKEY/scope") + r.Header.Set("X-Amz-Security-Token", "SENTINELTOKEN") - _, err = staticCredsFromProps(iceberg.Properties{iceio.S3SessionToken: "ST"}) - require.ErrorIs(t, err, internalaws.ErrIncompleteStaticCredentials, "a lone session token must be an error, not the ambient identity") - - creds, err = staticCredsFromProps(iceberg.Properties{ - keyRestAccessKeyID: "RAK", - keyRestSecretAccessKey: "RSK", - keyRestSessionToken: "RST", - }) - require.NoError(t, err) - require.NotNil(t, creds) - got, err = creds.Retrieve(context.Background()) - require.NoError(t, err) - require.Equal(t, "RAK", got.AccessKeyID) - require.Equal(t, "RSK", got.SecretAccessKey) - require.Equal(t, "RST", got.SessionToken) - - creds, err = staticCredsFromProps(iceberg.Properties{ - iceio.S3AccessKeyID: "AK", - iceio.S3SecretAccessKey: "SK", - keyRestAccessKeyID: "RAK", - keyRestSecretAccessKey: "RSK", - }) - require.NoError(t, err) - got, err = creds.Retrieve(context.Background()) - require.NoError(t, err) - require.Equal(t, "AK", got.AccessKeyID, "s3.* keys take precedence over rest.* aliases") - require.Equal(t, "SK", got.SecretAccessKey) - - _, err = staticCredsFromProps(iceberg.Properties{keyRestAccessKeyID: "RAK"}) - require.ErrorIs(t, err, internalaws.ErrIncompleteStaticCredentials, "a lone rest.* access key must be an error") - - _, err = staticCredsFromProps(iceberg.Properties{ - iceio.S3AccessKeyID: "AK", - keyRestSecretAccessKey: "RSK", - }) - require.ErrorIs(t, err, internalaws.ErrIncompleteStaticCredentials, "a partial pair must not be completed with a field from the other namespace") - - creds, err = staticCredsFromProps(iceberg.Properties{ - iceio.S3AccessKeyID: "AK", - iceio.S3SecretAccessKey: "SK", - keyRestSessionToken: "RST", - }) - require.NoError(t, err) - got, err = creds.Retrieve(context.Background()) - require.NoError(t, err) - require.Equal(t, "AK", got.AccessKeyID) - require.Empty(t, got.SessionToken, "a complete s3.* pair must not inherit an unrelated rest.* session token") + return nil } -// TestSigV4SignsWithPropsCredentials pins the wiring: the SigV4 Authorization -// header must be signed with the credentials carried in the catalog properties. -func TestSigV4SignsWithPropsCredentials(t *testing.T) { - var authHeader 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("/test", func(w http.ResponseWriter, r *http.Request) { - authHeader = r.Header.Get("Authorization") - w.WriteHeader(http.StatusOK) - }) - srv := httptest.NewServer(mux) - defer srv.Close() - - cat, err := NewCatalog(context.Background(), "rest", srv.URL, - WithSigV4RegionSvc("us-east-1", "s3"), - WithAdditionalProps(iceberg.Properties{ - iceio.S3AccessKeyID: "AKIDEXAMPLEPROPS", - iceio.S3SecretAccessKey: "secretexample", - })) - require.NoError(t, err) - - req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, srv.URL+"/test", nil) - require.NoError(t, err) - resp, err := cat.cl.Do(req) - require.NoError(t, err) - require.NoError(t, resp.Body.Close()) - - require.Contains(t, authHeader, "Credential=AKIDEXAMPLEPROPS/", - "SigV4 must sign with the credentials from catalog properties, not the default chain") -} - -// TestSigv4DoesNotSignCrossOriginRedirect pins that a redirect to a different -// origin is not re-signed, so the SigV4 Authorization header and session token -// never reach an unconfigured host. -func TestSigv4DoesNotSignCrossOriginRedirect(t *testing.T) { +// TestSignerDoesNotSignCrossOriginRedirect pins the origin guard in +// sessionTransport.RoundTrip: a redirect 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 var gotAuth, gotToken string second := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -177,7 +77,7 @@ func TestSigv4DoesNotSignCrossOriginRedirect(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{}}) + _ = 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) { firstAuth = r.Header.Get("Authorization") @@ -186,14 +86,9 @@ func TestSigv4DoesNotSignCrossOriginRedirect(t *testing.T) { first := httptest.NewServer(mux) defer first.Close() - cat, err := NewCatalog(context.Background(), "rest", first.URL, - WithSigV4RegionSvc("us-east-1", "s3"), - WithAdditionalProps(iceberg.Properties{ - iceio.S3AccessKeyID: "AKIDEXAMPLEPROPS", - iceio.S3SecretAccessKey: "secretexample", - iceio.S3SessionToken: "SESSIONTOKENEXAMPLE", - })) + cat, err := NewCatalog(context.Background(), "rest", first.URL, WithSigner(markingSigner{})) require.NoError(t, err) + t.Cleanup(func() { _ = cat.Close() }) req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, first.URL+"/redirect", nil) require.NoError(t, err) @@ -201,9 +96,9 @@ func TestSigv4DoesNotSignCrossOriginRedirect(t *testing.T) { require.NoError(t, err) require.NoError(t, resp.Body.Close()) - require.Contains(t, firstAuth, "Credential=AKIDEXAMPLEPROPS/", "the configured origin must still be signed") + 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 SigV4 Authorization header") + 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") } @@ -934,134 +829,6 @@ func TestAuthUriHeader(t *testing.T) { assert.Equal(t, "Bearer some_jwt_token", capturedAuthHeader) } -func TestSigv4EmptyStringHash(t *testing.T) { - t.Parallel() - hash := sha256.New() - payloadHash := hex.EncodeToString(hash.Sum(nil)) - // Sanity check the constant. - require.Equal(t, payloadHash, emptyStringHash) -} - -func TestSigv4ContentSha256Header(t *testing.T) { - t.Parallel() - - cfg, err := config.LoadDefaultConfig(context.Background(), func(opts *config.LoadOptions) error { - opts.Credentials = credentials.StaticCredentialsProvider{ - Value: aws.Credentials{ - AccessKeyID: "test-access-key", - SecretAccessKey: "test-secret-key", - }, - } - - return nil - }) - require.NoError(t, err) - - t.Run("header set when sigv4 enabled", func(t *testing.T) { - t.Parallel() - var capturedHeader string - mux := http.NewServeMux() - srv := httptest.NewServer(mux) - defer srv.Close() - - 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("/test", func(w http.ResponseWriter, r *http.Request) { - capturedHeader = r.Header.Get("x-amz-content-sha256") - w.WriteHeader(http.StatusOK) - }) - - cat, err := NewCatalog(context.Background(), "rest", srv.URL, - WithSigV4(), - WithSigV4RegionSvc("us-east-1", "s3"), - WithAwsConfig(cfg)) - require.NoError(t, err) - - req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, srv.URL+"/test", nil) - require.NoError(t, err) - - _, err = cat.cl.Do(req) - require.NoError(t, err) - - assert.NotEmpty(t, capturedHeader, "x-amz-content-sha256 header should be set when sigv4 is enabled") - assert.Equal(t, emptyStringHash, capturedHeader, "header should contain hash of empty body") - }) - - t.Run("header not set when sigv4 disabled", func(t *testing.T) { - t.Parallel() - var capturedHeader string - headerPresent := false - mux := http.NewServeMux() - srv := httptest.NewServer(mux) - defer srv.Close() - - 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("/test", func(w http.ResponseWriter, r *http.Request) { - capturedHeader = r.Header.Get("x-amz-content-sha256") - _, headerPresent = r.Header["X-Amz-Content-Sha256"] - w.WriteHeader(http.StatusOK) - }) - - cat, err := NewCatalog(context.Background(), "rest", srv.URL) - require.NoError(t, err) - - req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, srv.URL+"/test", nil) - require.NoError(t, err) - - _, err = cat.cl.Do(req) - require.NoError(t, err) - - assert.Empty(t, capturedHeader, "x-amz-content-sha256 header should not be set when sigv4 is disabled") - assert.False(t, headerPresent, "x-amz-content-sha256 header should not be present when sigv4 is disabled") - }) - - t.Run("header contains correct hash for request body", func(t *testing.T) { - t.Parallel() - var capturedHeader string - mux := http.NewServeMux() - srv := httptest.NewServer(mux) - defer srv.Close() - - 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("/test", func(w http.ResponseWriter, r *http.Request) { - capturedHeader = r.Header.Get("x-amz-content-sha256") - w.WriteHeader(http.StatusOK) - }) - - cat, err := NewCatalog(context.Background(), "rest", srv.URL, - WithSigV4(), - WithSigV4RegionSvc("us-east-1", "s3"), - WithAwsConfig(cfg)) - require.NoError(t, err) - - body := []byte(`{"test": "data"}`) - expectedHash := sha256.Sum256(body) - expectedHashStr := hex.EncodeToString(expectedHash[:]) - - req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, srv.URL+"/test", bytes.NewReader(body)) - require.NoError(t, err) - - _, err = cat.cl.Do(req) - require.NoError(t, err) - - assert.Equal(t, expectedHashStr, capturedHeader, "header should contain correct hash of request body") - }) -} - type roundTripFunc func(*http.Request) (*http.Response, error) func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { @@ -1173,176 +940,6 @@ func TestReqOptionsCompose(t *testing.T) { assert.Equal(t, []string{"X-First", "X-Second", "X-Third"}, cfg.suppressHeaders) } -type closeTrackingReadCloser struct { - *bytes.Reader - closeErr error - closed bool -} - -func (r *closeTrackingReadCloser) Close() error { - r.closed = true - - return r.closeErr -} - -func newSigV4TestTransport(rt http.RoundTripper) *sessionTransport { - return &sessionTransport{ - RoundTripper: rt, - signer: v4.NewSigner(), - cfg: aws.Config{ - Region: "us-east-1", - Credentials: credentials.StaticCredentialsProvider{ - Value: aws.Credentials{ - AccessKeyID: "test-access-key", - SecretAccessKey: "test-secret-key", - }, - }, - }, - service: "s3", - newHash: sha256.New, - } -} - -func TestSigv4ClosesClonedRequestBody(t *testing.T) { - t.Parallel() - - body := []byte(`{"test": "data"}`) - var clonedBody *closeTrackingReadCloser - - transport := newSigV4TestTransport(roundTripFunc(func(_ *http.Request) (*http.Response, error) { - return &http.Response{ - StatusCode: http.StatusOK, - Body: io.NopCloser(bytes.NewReader(nil)), - Header: make(http.Header), - }, nil - })) - - req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, - "https://example.com/test", bytes.NewReader(body)) - require.NoError(t, err) - req.GetBody = func() (io.ReadCloser, error) { - clonedBody = &closeTrackingReadCloser{Reader: bytes.NewReader(body)} - - return clonedBody, nil - } - - resp, err := transport.RoundTrip(req) - require.NoError(t, err) - require.NoError(t, resp.Body.Close()) - require.NotNil(t, clonedBody) - assert.True(t, clonedBody.closed) -} - -func TestSigv4ReturnsClonedRequestBodyCloseError(t *testing.T) { - t.Parallel() - - closeErr := errors.New("close failed") - body := []byte(`{"test": "data"}`) - var clonedBody *closeTrackingReadCloser - - transport := newSigV4TestTransport(roundTripFunc(func(_ *http.Request) (*http.Response, error) { - t.Fatal("request should not be sent when the signing body clone fails to close") - - return nil, nil - })) - - req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, - "https://example.com/test", bytes.NewReader(body)) - require.NoError(t, err) - req.GetBody = func() (io.ReadCloser, error) { - clonedBody = &closeTrackingReadCloser{ - Reader: bytes.NewReader(body), - closeErr: closeErr, - } - - return clonedBody, nil - } - - _, err = transport.RoundTrip(req) - require.ErrorIs(t, err, closeErr) - require.NotNil(t, clonedBody) - assert.True(t, clonedBody.closed) -} - -func TestSigv4ConcurrentSigners(t *testing.T) { - t.Parallel() - mux := http.NewServeMux() - srv := httptest.NewUnstartedServer(mux) - // If we use HTTP 1.1, this test can try to make too many connections - // and exhaust ephemeral ports. - srv.EnableHTTP2 = true - srv.StartTLS() // Using TLS to easily support HTTP/2 - rootCAs := x509.NewCertPool() - rootCAs.AddCert(srv.Certificate()) - - 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{}, - }) - }) - - cfg, err := config.LoadDefaultConfig(context.Background(), func(opts *config.LoadOptions) error { - opts.Credentials = credentials.StaticCredentialsProvider{ - Value: aws.Credentials{ - AccessKeyID: "abcdefghjklmnop", - SecretAccessKey: "01234567abcdefgh01234567abcdefgh01234567abcdefgh01234567abcdefgh", - }, - } - - return nil - }) - require.NoError(t, err) - - cat, err := NewCatalog(context.Background(), "rest", srv.URL, - WithSigV4(), - WithSigV4RegionSvc("abc", "def"), - WithAwsConfig(cfg), - WithTLSConfig(&tls.Config{ - RootCAs: rootCAs, - })) - require.NoError(t, err) - assert.NotNil(t, cat) - - // We aren't recreating the signature logic to verify on the server. We're - // just running many concurrent requests to make sure the race detector - // doesn't find any data races with how the session transport and signer - // are used from concurrent goroutines. - ctx, cancel := context.WithCancel(context.Background()) - grp, ctx := errgroup.WithContext(ctx) - var count atomic.Uint64 - for range 10 { - grp.Go(func() error { - for { - if err := ctx.Err(); err != nil { - return nil - } - body := make([]byte, 1024) - if _, err := rand.Read(body); err != nil { - return err - } - // Intentionally using context.Background instead of ctx so that we - // don't get interrupted when context is cancelled. - req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, srv.URL, bytes.NewReader(body)) - if err != nil { - return err - } - resp, err := cat.cl.Do(req) - if err != nil { - return err - } - // We don't actually care about the response, only that it actually made it to the server. - _, _ = io.Copy(io.Discard, resp.Body) - _ = resp.Body.Close() - count.Add(1) - } - }) - } - time.Sleep(5 * time.Second) - cancel() - require.NoError(t, grp.Wait()) - t.Logf("issued %d requests", count.Load()) -} - func TestCredentialRefreshOnExpiry(t *testing.T) { t.Parallel() @@ -2224,3 +1821,164 @@ func TestNamespaceSeparatorDefaultsToUnitSeparator(t *testing.T) { require.NoError(t, err) assert.Equal(t, "/v1/namespaces/a%1Fb", gotPath) } + +// signerFunc adapts a function to the RequestSigner interface for tests. +type signerFunc func(*http.Request) error + +func (f signerFunc) SignRequest(r *http.Request) error { return f(r) } + +func TestSessionTransportInvokesSigner(t *testing.T) { + t.Parallel() + + var signed bool + 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 + }), + defaultHeaders: http.Header{}, + signer: signerFunc(func(r *http.Request) error { + signed = true + 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.True(t, signed, "signer.SignRequest should be invoked") + assert.Equal(t, "yes", got.Get("X-Signed")) +} + +func TestSessionTransportSignerErrorAbortsRequest(t *testing.T) { + t.Parallel() + + wantErr := errors.New("sign failed") + s := &sessionTransport{ + RoundTripper: roundTripFunc(func(_ *http.Request) (*http.Response, error) { + t.Fatal("request must not be sent when signing fails") + + return nil, nil + }), + defaultHeaders: http.Header{}, + signer: signerFunc(func(_ *http.Request) error { return wantErr }), + } + + req, err := http.NewRequest(http.MethodGet, "http://example.com", nil) + require.NoError(t, err) + _, err = s.RoundTrip(req) + require.ErrorIs(t, err, wantErr) +} + +// TestSessionTransportConcurrentRoundTrip drives one shared sessionTransport +// from many goroutines (each with its own request) to confirm the signing and +// default-header path holds no per-request state that races. Meaningful under +// the race detector; complements the sigv4 package's real-signer concurrency +// test. +func TestSessionTransportConcurrentRoundTrip(t *testing.T) { + t.Parallel() + + var count atomic.Int64 + s := &sessionTransport{ + RoundTripper: roundTripFunc(func(_ *http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil + }), + defaultHeaders: http.Header{"X-Default": {"v"}}, + signer: signerFunc(func(r *http.Request) error { + count.Add(1) + r.Header.Set("X-Signed", "yes") + + return nil + }), + } + + var wg sync.WaitGroup + for range 50 { + wg.Go(func() { + req, err := http.NewRequest(http.MethodGet, "http://example.com", nil) + if err != nil { + t.Error(err) + + return + } + if _, err := s.RoundTrip(req); err != nil { + t.Error(err) + } + }) + } + wg.Wait() + + assert.Equal(t, int64(50), count.Load()) +} + +// TestSigV4WithoutBackendReturnsHelpfulError verifies that requesting SigV4 +// without a registered signer backend fails with guidance naming the import to +// add. This internal test binary cannot import catalog/rest/sigv4 (that would be +// an import cycle), so no init registers the "sigv4" backend and the registry is +// always empty here; the "no backend" path is therefore exercised without any +// registry manipulation. +func TestSigV4WithoutBackendReturnsHelpfulError(t *testing.T) { + t.Parallel() + + mux := http.NewServeMux() + srv := httptest.NewServer(mux) + defer srv.Close() + + mux.HandleFunc("/v1/config", func(w http.ResponseWriter, _ *http.Request) { + json.NewEncoder(w).Encode(map[string]any{ + "defaults": map[string]any{}, "overrides": map[string]any{}, + }) + }) + + _, err := NewCatalog(context.Background(), "rest", srv.URL, WithSigV4()) + require.Error(t, err) + assert.Contains(t, err.Error(), "catalog/rest/sigv4") +} + +// TestSignerFactoryReceivesResolvedRegionService verifies that a +// WithSignerFactory is handed the signing region and service resolved from +// WithSigV4RegionSvc and any server-provided /v1/config overrides, rather than +// values frozen when the option was constructed. This is the seam a pre-built +// WithSigner (which sigv4.WithAwsConfig used before it became a factory) got +// wrong: it ignored both WithSigV4RegionSvc and server overrides. +func TestSignerFactoryReceivesResolvedRegionService(t *testing.T) { + t.Parallel() + + mux := http.NewServeMux() + srv := httptest.NewServer(mux) + defer srv.Close() + + mux.HandleFunc("/v1/config", func(w http.ResponseWriter, _ *http.Request) { + json.NewEncoder(w).Encode(map[string]any{ + "defaults": map[string]any{}, + "overrides": map[string]any{ + keyRestSigV4Region: "ap-south-1", + keyRestSigV4Service: "s3tables", + }, + }) + }) + + // Sequential within NewCatalog: the factory is called for the bootstrap + // /v1/config request (pre-override) and again for the real session + // (post-override), so the last call carries the merged values. + var last SignerConfig + factory := func(_ context.Context, cfg SignerConfig) (RequestSigner, error) { + last = cfg + + return signerFunc(func(*http.Request) error { return nil }), nil + } + + cat, err := NewCatalog(context.Background(), "rest", srv.URL, + WithSigV4RegionSvc("us-east-1", "execute-api"), + WithSignerFactory(factory)) + require.NoError(t, err) + t.Cleanup(func() { _ = cat.Close() }) + + assert.Equal(t, "ap-south-1", last.Region, "server signing-region override must reach the factory") + assert.Equal(t, "s3tables", last.Service, "server signing-name override must reach the factory") +} diff --git a/catalog/rest/signer.go b/catalog/rest/signer.go new file mode 100644 index 000000000..0d330f396 --- /dev/null +++ b/catalog/rest/signer.go @@ -0,0 +1,133 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package rest + +import ( + "context" + "fmt" + "net/http" + "sync" + + "github.com/apache/iceberg-go" +) + +// SignerNameSigV4 is the scheme under which the AWS SigV4 backend registers +// itself (see catalog/rest/sigv4). It is also the value carried by the +// rest.sigv4-enabled property path. It is exported so the backend can register +// under the exact same name the core looks up, rather than a duplicated string +// literal that could silently drift. +const SignerNameSigV4 = "sigv4" + +// RequestSigner signs an outgoing catalog HTTP request in place, for example +// with AWS SigV4. Implementations live in optional sub-packages (see +// github.com/apache/iceberg-go/catalog/rest/sigv4) so the core REST client +// depends on no cloud SDK. Signing runs inside the session transport's +// RoundTrip, after auth headers are applied, and may use req.Context. +type RequestSigner interface { + SignRequest(req *http.Request) error +} + +// SignerConfig carries the signing parameters the REST core knows about, +// deliberately free of any cloud-SDK types. A SignerFactory turns it into a +// concrete RequestSigner. +type SignerConfig struct { + // Region is the signing region (SigV4 signing-region). + Region string + // Service is the signing service name (SigV4 signing-name), e.g. + // "execute-api", "s3tables". + Service string + // Props carries the catalog properties after /v1/config overrides have been + // folded in. A backend may read signing credentials from it (e.g. the sigv4 + // backend builds a static credentials provider from the s3.* / rest.* keys) + // without the core knowing any cloud-SDK property names. It may be nil or + // empty when no properties were supplied. + Props iceberg.Properties +} + +// SignerFactory builds a RequestSigner from core configuration. A signing +// backend registers one under a scheme name via RegisterSigner, so that a +// blank import of the backend package is enough to enable property- or +// option-driven signing (e.g. WithSigV4 / rest.sigv4-enabled). +type SignerFactory func(ctx context.Context, cfg SignerConfig) (RequestSigner, error) + +var ( + signerMu sync.RWMutex + signerRegistry = map[string]SignerFactory{} +) + +// RegisterSigner registers a signer factory under name (e.g. "sigv4"). It is +// intended to be called from a backend package's init function. Registering +// the same name twice replaces the previous factory. It panics if factory is +// nil, since a nil factory is a programming error that would otherwise surface +// only later, as a failure when signing is resolved. +func RegisterSigner(name string, factory SignerFactory) { + if factory == nil { + panic(fmt.Sprintf("rest: RegisterSigner: nil factory for %q", name)) + } + + signerMu.Lock() + defer signerMu.Unlock() + signerRegistry[name] = factory +} + +func lookupSigner(name string) (SignerFactory, bool) { + signerMu.RLock() + defer signerMu.RUnlock() + f, ok := signerRegistry[name] + + return f, ok +} + +// resolveSigner determines the request signer for a session. It returns +// (nil, nil) when signing is not configured, and a helpful error when SigV4 is +// requested but no backend has been imported. +// +// Precedence: +// - An explicit WithSigner is used verbatim (the caller fully built it), so it +// bypasses the sigv4-enabled / signing-region / signing-name settings and any +// server-provided overrides. +// - Otherwise, when SigV4 is enabled (WithSigV4 / WithSigV4RegionSvc or the +// rest.sigv4-enabled property), a WithSignerFactory (e.g. sigv4.WithAwsConfig) +// builds the signer if one was supplied, else the registered "sigv4" backend +// does. Both receive the resolved region/service, which by this point already +// include any /v1/config overrides fetchConfig folded into opts. +func resolveSigner(ctx context.Context, opts *options) (RequestSigner, error) { + if opts.signer != nil { + return opts.signer, nil + } + + if !opts.enableSigv4 { + return nil, nil + } + + factory := opts.signerFactory + if factory == nil { + var ok bool + if factory, ok = lookupSigner(SignerNameSigV4); !ok { + return nil, fmt.Errorf( + "rest: SigV4 signing was requested (%s) but no signer backend is registered; add a blank import: import _ %q", + keyRestSigV4, "github.com/apache/iceberg-go/catalog/rest/sigv4") + } + } + + return factory(ctx, SignerConfig{ + Region: opts.sigv4Region, + Service: opts.sigv4Service, + Props: opts.additionalProps, + }) +} diff --git a/catalog/rest/sigv4/sigv4.go b/catalog/rest/sigv4/sigv4.go new file mode 100644 index 000000000..9e9b42e24 --- /dev/null +++ b/catalog/rest/sigv4/sigv4.go @@ -0,0 +1,223 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +// Package sigv4 provides AWS Signature Version 4 signing for the REST catalog +// client. It is an optional backend: importing it (even as a blank import) +// registers the "sigv4" signer so that WithSigV4 / the rest.sigv4-enabled +// property work, and it keeps the AWS SDK out of the core catalog/rest package +// for consumers that do not need it. +// +// import ( +// "github.com/apache/iceberg-go/catalog/rest" +// _ "github.com/apache/iceberg-go/catalog/rest/sigv4" +// ) +// +// cat, err := rest.NewCatalog(ctx, "c", uri, rest.WithSigV4RegionSvc("us-east-1", "s3tables")) +// +// To sign with an explicit aws.Config, pass sigv4.WithAwsConfig alongside +// WithSigV4 / WithSigV4RegionSvc (no blank import needed); the option supplies +// the config while those options still drive the signing region and service. +package sigv4 + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "io" + "net/http" + "time" + + "github.com/apache/iceberg-go" + "github.com/apache/iceberg-go/catalog/rest" + internalaws "github.com/apache/iceberg-go/internal/awsconfig" + iceio "github.com/apache/iceberg-go/io" + "github.com/aws/aws-sdk-go-v2/aws" + v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4" + "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" +) + +// emptyStringHash is the SHA-256 of the empty string, used as the payload hash +// for requests without a body. +// from https://pkg.go.dev/github.com/aws/aws-sdk-go-v2/aws/signer/v4#Signer.SignHTTP +const emptyStringHash = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855" + +// defaultSigningService is the SigV4 service name used when none is supplied, +// matching the Iceberg Java reference default (AwsProperties.REST_SIGNING_NAME_DEFAULT). +const defaultSigningService = "execute-api" + +// keyRestAccessKeyID and friends are the Java-client property names for the +// SigV4 signing credentials. They are accepted as aliases for the s3.* +// properties; the s3.* keys take precedence when both are set. +const ( + keyRestAccessKeyID = "rest.access-key-id" + keyRestSecretAccessKey = "rest.secret-access-key" + keyRestSessionToken = "rest.session-token" +) + +func init() { + rest.RegisterSigner(rest.SignerNameSigV4, func(ctx context.Context, cfg rest.SignerConfig) (rest.RequestSigner, error) { + creds, err := staticCredsFromProps(cfg.Props) + if err != nil { + return nil, err + } + + awscfg, err := config.LoadDefaultConfig(ctx) + if err != nil { + return nil, fmt.Errorf("sigv4: load AWS config: %w", err) + } + // Sign with the credentials carried in the catalog properties when present, + // rather than only the AWS default credential chain. + if creds != nil { + awscfg.Credentials = creds + } + + return buildSigner(awscfg, cfg), nil + }) +} + +// staticCredsFromProps returns a static credentials provider built from the +// signing-credential properties carried in the catalog config. It prefers the +// s3.* keys and falls back to the Java-compatible rest.* aliases, resolving the +// tuple atomically from a single namespace so a partial pair is never completed +// with fields from the other one. It returns (nil, nil) when neither namespace +// sets any credential property, so the caller falls back to the default +// credential chain, and an ErrIncompleteStaticCredentials error when the chosen +// namespace is incomplete. +// +// An explicit aws.Config supplied via WithAwsConfig bypasses this: that path +// already fully specifies the credentials, so it never consults the properties. +func staticCredsFromProps(props iceberg.Properties) (aws.CredentialsProvider, error) { + namespaces := [][3]string{ + {iceio.S3AccessKeyID, iceio.S3SecretAccessKey, iceio.S3SessionToken}, + {keyRestAccessKeyID, keyRestSecretAccessKey, keyRestSessionToken}, + } + for _, ns := range namespaces { + accessKey, secretKey, token := props[ns[0]], props[ns[1]], props[ns[2]] + if accessKey == "" && secretKey == "" && token == "" { + continue + } + if err := internalaws.ValidateStaticCredentials(ns[0], ns[1], ns[2], accessKey, secretKey, token); err != nil { + return nil, err + } + + return credentials.NewStaticCredentialsProvider(accessKey, secretKey, token), nil + } + + return nil, nil +} + +// buildSigner applies the REST-core signing config to an aws.Config and returns +// a ready signer. A non-empty cfg.Region overrides awscfg.Region; an empty +// service falls back to defaultSigningService. The latter covers the +// property-only path (rest.sigv4-enabled without rest.signing-name), matching +// WithSigV4, so signatures stay valid. +func buildSigner(awscfg aws.Config, cfg rest.SignerConfig) *signer { + if cfg.Region != "" { + awscfg.Region = cfg.Region + } + + service := cfg.Service + if service == "" { + service = defaultSigningService + } + + return newSigner(awscfg, service) +} + +// signer signs HTTP requests with AWS Signature Version 4. It implements +// rest.RequestSigner. +type signer struct { + httpSigner v4.HTTPSigner + cfg aws.Config + service string +} + +func newSigner(cfg aws.Config, service string) *signer { + return &signer{ + httpSigner: v4.NewSigner(), + cfg: cfg, + service: service, + } +} + +// SignRequest signs r in place with SigV4: it computes the payload hash, sets +// the required x-amz-content-sha256 header, and applies the signature. +func (s *signer) SignRequest(r *http.Request) error { + var payloadHash string + if r.Body == nil { + payloadHash = emptyStringHash + } else { + // SigV4 needs the payload hash, so the body must be re-readable. + // http.NewRequest* sets GetBody for the standard body types; a + // hand-built request with a raw io.ReadCloser body and no GetBody cannot + // be signed, so fail with a clear error instead of a nil-func panic. + if r.GetBody == nil { + return errors.New("sigv4: cannot sign request whose body is not re-readable (Request.GetBody is nil)") + } + + rdr, err := r.GetBody() + if err != nil { + return err + } + + h := sha256.New() + if _, err = io.Copy(h, rdr); err != nil { + if closeErr := rdr.Close(); closeErr != nil { + err = errors.Join(err, closeErr) + } + + return err + } + + if err = rdr.Close(); err != nil { + return err + } + + payloadHash = hex.EncodeToString(h.Sum(nil)) + } + + creds, err := s.cfg.Credentials.Retrieve(r.Context()) + if err != nil { + return err + } + + // Set the x-amz-content-sha256 header before signing; SigV4 signature + // verification requires it. + r.Header.Set("x-amz-content-sha256", payloadHash) + + return s.httpSigner.SignHTTP(r.Context(), creds, r, payloadHash, s.service, s.cfg.Region, time.Now()) +} + +// WithAwsConfig returns a rest.Option that signs catalog requests with AWS +// SigV4 using the supplied aws.Config, without consulting the AWS default +// credential chain. It replaces the former rest.WithAwsConfig, whose AWS SDK +// dependency now lives only in this optional sub-package. +// +// Like the blank-import path, it does not by itself enable signing: pair it +// with WithSigV4 or WithSigV4RegionSvc (or the rest.sigv4-enabled property), +// which remain the single source of the signing region and service. Those +// values, plus any server-provided /v1/config overrides, are applied on top of +// cfg when the signer is built (a non-empty region overrides cfg.Region; an +// empty service defaults to "execute-api"). +func WithAwsConfig(cfg aws.Config) rest.Option { + return rest.WithSignerFactory(func(_ context.Context, sc rest.SignerConfig) (rest.RequestSigner, error) { + return buildSigner(cfg, sc), nil + }) +} diff --git a/catalog/rest/sigv4/sigv4_internal_test.go b/catalog/rest/sigv4/sigv4_internal_test.go new file mode 100644 index 000000000..94c14b61d --- /dev/null +++ b/catalog/rest/sigv4/sigv4_internal_test.go @@ -0,0 +1,415 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package sigv4 + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + + "github.com/apache/iceberg-go" + "github.com/apache/iceberg-go/catalog/rest" + internalaws "github.com/apache/iceberg-go/internal/awsconfig" + iceio "github.com/apache/iceberg-go/io" + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func newTestSigner(t *testing.T) *signer { + t.Helper() + + cfg, err := config.LoadDefaultConfig(context.Background(), func(o *config.LoadOptions) error { + o.Credentials = credentials.StaticCredentialsProvider{ + Value: aws.Credentials{ + AccessKeyID: "test-access-key", + SecretAccessKey: "test-secret-key", + }, + } + + return nil + }) + require.NoError(t, err) + cfg.Region = "us-east-1" + + return newSigner(cfg, "s3") +} + +func TestEmptyStringHash(t *testing.T) { + t.Parallel() + + h := sha256.New() + assert.Equal(t, hex.EncodeToString(h.Sum(nil)), emptyStringHash) +} + +func TestSignRequestEmptyBodyContentHash(t *testing.T) { + t.Parallel() + + s := newTestSigner(t) + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "https://example.com/test", nil) + require.NoError(t, err) + require.NoError(t, s.SignRequest(req)) + + assert.Equal(t, emptyStringHash, req.Header.Get("x-amz-content-sha256")) + assert.NotEmpty(t, req.Header.Get("Authorization"), "SigV4 should set the Authorization header") +} + +func TestSignRequestBodyContentHash(t *testing.T) { + t.Parallel() + + s := newTestSigner(t) + body := []byte(`{"test": "data"}`) + sum := sha256.Sum256(body) + + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, "https://example.com/test", bytes.NewReader(body)) + require.NoError(t, err) + require.NoError(t, s.SignRequest(req)) + + assert.Equal(t, hex.EncodeToString(sum[:]), req.Header.Get("x-amz-content-sha256")) +} + +func TestSignRequestNilGetBodyReturnsError(t *testing.T) { + t.Parallel() + + s := newTestSigner(t) + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, "https://example.com/test", bytes.NewReader([]byte(`{}`))) + require.NoError(t, err) + req.GetBody = nil // a hand-built request whose body cannot be re-read + + err = s.SignRequest(req) + require.Error(t, err) + assert.Contains(t, err.Error(), "GetBody", "error should explain the body is not re-readable") +} + +type closeTrackingReadCloser struct { + *bytes.Reader + closeErr error + closed bool +} + +func (r *closeTrackingReadCloser) Close() error { + r.closed = true + + return r.closeErr +} + +func TestSignRequestClosesClonedBody(t *testing.T) { + t.Parallel() + + s := newTestSigner(t) + body := []byte(`{"test": "data"}`) + var cloned *closeTrackingReadCloser + + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, "https://example.com/test", bytes.NewReader(body)) + require.NoError(t, err) + req.GetBody = func() (io.ReadCloser, error) { + cloned = &closeTrackingReadCloser{Reader: bytes.NewReader(body)} + + return cloned, nil + } + + require.NoError(t, s.SignRequest(req)) + require.NotNil(t, cloned) + assert.True(t, cloned.closed) +} + +func TestSignRequestReturnsClonedBodyCloseError(t *testing.T) { + t.Parallel() + + s := newTestSigner(t) + closeErr := errors.New("close failed") + body := []byte(`{"test": "data"}`) + var cloned *closeTrackingReadCloser + + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, "https://example.com/test", bytes.NewReader(body)) + require.NoError(t, err) + req.GetBody = func() (io.ReadCloser, error) { + cloned = &closeTrackingReadCloser{Reader: bytes.NewReader(body), closeErr: closeErr} + + return cloned, nil + } + + err = s.SignRequest(req) + require.ErrorIs(t, err, closeErr) + require.NotNil(t, cloned) + assert.True(t, cloned.closed) +} + +func TestSignRequestConcurrent(t *testing.T) { + t.Parallel() + + // POSTs with a body so the payload-hashing path (GetBody clone + SHA-256) + // runs concurrently on a single shared signer, exercising the shared v4 + // signer and aws.Config under the race detector. + s := newTestSigner(t) + body := []byte(`{"test":"data"}`) + var wg sync.WaitGroup + for range 20 { + wg.Go(func() { + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, "https://example.com/test", bytes.NewReader(body)) + if err == nil { + err = s.SignRequest(req) + } + if err != nil { + t.Error(err) + } + }) + } + wg.Wait() +} + +// TestRegisteredBackendEnablesSigV4 verifies that importing this package (its +// init registers the sigv4 signer) lets rest.WithSigV4RegionSvc resolve without +// the caller supplying an explicit signer, and that the resolved signer +// actually signs outbound catalog requests end to end. +func TestRegisteredBackendEnablesSigV4(t *testing.T) { + // t.Setenv precludes t.Parallel; static credentials let the registered + // factory's config.LoadDefaultConfig resolve offline (no EC2 IMDS lookup). + t.Setenv("AWS_ACCESS_KEY_ID", "test-access-key") + t.Setenv("AWS_SECRET_ACCESS_KEY", "test-secret-key") + t.Setenv("AWS_REGION", "us-east-1") + + var gotAuth, gotSHA string + mux := http.NewServeMux() + srv := httptest.NewServer(mux) + defer srv.Close() + + mux.HandleFunc("/v1/config", func(w http.ResponseWriter, r *http.Request) { + gotAuth = r.Header.Get("Authorization") + gotSHA = r.Header.Get("x-amz-content-sha256") + json.NewEncoder(w).Encode(map[string]any{ + "defaults": map[string]any{}, "overrides": map[string]any{}, + }) + }) + + cat, err := rest.NewCatalog(context.Background(), "rest", srv.URL, + rest.WithSigV4RegionSvc("us-east-1", "s3")) + require.NoError(t, err) + require.NotNil(t, cat) + t.Cleanup(func() { _ = cat.Close() }) + + // The registered backend must actually sign the bootstrap /v1/config request, + // not merely let catalog construction succeed. + assert.Contains(t, gotAuth, "AWS4-HMAC-SHA256", "request should carry a SigV4 Authorization header") + assert.NotEmpty(t, gotSHA, "request should carry x-amz-content-sha256") +} + +// TestWithAwsConfigUsesConfiguredRegionService verifies that sigv4.WithAwsConfig +// signs with the region/service from WithSigV4RegionSvc, not the aws.Config's +// own region. Before WithAwsConfig became a signer factory it froze the scope at +// option-construction time, so the natural migration (keep WithSigV4RegionSvc, +// swap rest.WithAwsConfig for sigv4.WithAwsConfig) silently signed for the wrong +// scope and the server rejected it with a 403. +func TestWithAwsConfigUsesConfiguredRegionService(t *testing.T) { + t.Parallel() + + var gotAuth string + mux := http.NewServeMux() + srv := httptest.NewServer(mux) + defer srv.Close() + + mux.HandleFunc("/v1/config", func(w http.ResponseWriter, r *http.Request) { + gotAuth = r.Header.Get("Authorization") + json.NewEncoder(w).Encode(map[string]any{ + "defaults": map[string]any{}, "overrides": map[string]any{}, + }) + }) + + cfg := aws.Config{ + Region: "eu-central-1", // must be overridden by WithSigV4RegionSvc below + Credentials: credentials.StaticCredentialsProvider{Value: aws.Credentials{ + AccessKeyID: "test-access-key", + SecretAccessKey: "test-secret-key", + }}, + } + + cat, err := rest.NewCatalog(context.Background(), "rest", srv.URL, + rest.WithSigV4RegionSvc("us-west-2", "s3tables"), + WithAwsConfig(cfg)) + require.NoError(t, err) + require.NotNil(t, cat) + t.Cleanup(func() { _ = cat.Close() }) + + assert.Contains(t, gotAuth, "/us-west-2/s3tables/aws4_request", + "signature scope must come from WithSigV4RegionSvc, not aws.Config.Region") + assert.NotContains(t, gotAuth, "eu-central-1", + "aws.Config.Region must not leak into the signature scope") +} + +// TestConcurrentSignedCatalogRequests drives the full transport+signer stack +// from concurrent goroutines so the race detector can observe the shared signer +// registry lookups and the signing path. Every request the server sees must be +// SigV4-signed. This restores the end-to-end concurrency coverage the previous +// in-package TestSigv4ConcurrentSigners provided before signing moved here. +func TestConcurrentSignedCatalogRequests(t *testing.T) { + t.Setenv("AWS_ACCESS_KEY_ID", "test-access-key") + t.Setenv("AWS_SECRET_ACCESS_KEY", "test-secret-key") + t.Setenv("AWS_REGION", "us-east-1") + + var total, signed atomic.Int64 + mux := http.NewServeMux() + srv := httptest.NewServer(mux) + defer srv.Close() + + mux.HandleFunc("/v1/config", func(w http.ResponseWriter, r *http.Request) { + total.Add(1) + if strings.Contains(r.Header.Get("Authorization"), "AWS4-HMAC-SHA256") && + r.Header.Get("x-amz-content-sha256") != "" { + signed.Add(1) + } + json.NewEncoder(w).Encode(map[string]any{ + "defaults": map[string]any{}, "overrides": map[string]any{}, + }) + }) + + var wg sync.WaitGroup + for range 10 { + wg.Go(func() { + cat, err := rest.NewCatalog(context.Background(), "rest", srv.URL, + rest.WithSigV4RegionSvc("us-east-1", "s3")) + if err != nil { + t.Error(err) + + return + } + _ = cat.Close() + }) + } + wg.Wait() + + require.Positive(t, total.Load()) + assert.Equal(t, total.Load(), signed.Load(), "every request must be SigV4-signed") +} + +func TestStaticCredsFromProps(t *testing.T) { + t.Parallel() + + creds, err := staticCredsFromProps(iceberg.Properties{ + iceio.S3AccessKeyID: "AK", + iceio.S3SecretAccessKey: "SK", + iceio.S3SessionToken: "ST", + }) + require.NoError(t, err) + require.NotNil(t, creds) + got, err := creds.Retrieve(context.Background()) + require.NoError(t, err) + require.Equal(t, "AK", got.AccessKeyID) + require.Equal(t, "SK", got.SecretAccessKey) + require.Equal(t, "ST", got.SessionToken) + + creds, err = staticCredsFromProps(iceberg.Properties{}) + require.NoError(t, err, "no creds must fall back to the default chain") + require.Nil(t, creds) + + _, err = staticCredsFromProps(iceberg.Properties{iceio.S3AccessKeyID: "AK"}) + require.ErrorIs(t, err, internalaws.ErrIncompleteStaticCredentials, "a lone access key must be an error, not the ambient identity") + + _, err = staticCredsFromProps(iceberg.Properties{iceio.S3SecretAccessKey: "SK"}) + require.ErrorIs(t, err, internalaws.ErrIncompleteStaticCredentials, "a lone secret key must be an error, not the ambient identity") + + _, err = staticCredsFromProps(iceberg.Properties{iceio.S3SessionToken: "ST"}) + require.ErrorIs(t, err, internalaws.ErrIncompleteStaticCredentials, "a lone session token must be an error, not the ambient identity") + + creds, err = staticCredsFromProps(iceberg.Properties{ + keyRestAccessKeyID: "RAK", + keyRestSecretAccessKey: "RSK", + keyRestSessionToken: "RST", + }) + require.NoError(t, err) + require.NotNil(t, creds) + got, err = creds.Retrieve(context.Background()) + require.NoError(t, err) + require.Equal(t, "RAK", got.AccessKeyID) + require.Equal(t, "RSK", got.SecretAccessKey) + require.Equal(t, "RST", got.SessionToken) + + creds, err = staticCredsFromProps(iceberg.Properties{ + iceio.S3AccessKeyID: "AK", + iceio.S3SecretAccessKey: "SK", + keyRestAccessKeyID: "RAK", + keyRestSecretAccessKey: "RSK", + }) + require.NoError(t, err) + got, err = creds.Retrieve(context.Background()) + require.NoError(t, err) + require.Equal(t, "AK", got.AccessKeyID, "s3.* keys take precedence over rest.* aliases") + require.Equal(t, "SK", got.SecretAccessKey) + + _, err = staticCredsFromProps(iceberg.Properties{keyRestAccessKeyID: "RAK"}) + require.ErrorIs(t, err, internalaws.ErrIncompleteStaticCredentials, "a lone rest.* access key must be an error") + + _, err = staticCredsFromProps(iceberg.Properties{ + iceio.S3AccessKeyID: "AK", + keyRestSecretAccessKey: "RSK", + }) + require.ErrorIs(t, err, internalaws.ErrIncompleteStaticCredentials, "a partial pair must not be completed with a field from the other namespace") + + creds, err = staticCredsFromProps(iceberg.Properties{ + iceio.S3AccessKeyID: "AK", + iceio.S3SecretAccessKey: "SK", + keyRestSessionToken: "RST", + }) + require.NoError(t, err) + got, err = creds.Retrieve(context.Background()) + require.NoError(t, err) + require.Equal(t, "AK", got.AccessKeyID) + require.Empty(t, got.SessionToken, "a complete s3.* pair must not inherit an unrelated rest.* session token") +} + +// TestSigV4SignsWithPropsCredentials pins the wiring end to end: the SigV4 +// Authorization header on the bootstrap request must be signed with the +// credentials carried in the catalog properties (via SignerConfig.Props), not +// the AWS default chain. +func TestSigV4SignsWithPropsCredentials(t *testing.T) { + t.Parallel() + + var gotAuth string + mux := http.NewServeMux() + srv := httptest.NewServer(mux) + defer srv.Close() + + mux.HandleFunc("/v1/config", func(w http.ResponseWriter, r *http.Request) { + gotAuth = r.Header.Get("Authorization") + _ = json.NewEncoder(w).Encode(map[string]any{"defaults": map[string]any{}, "overrides": map[string]any{}}) + }) + + cat, err := rest.NewCatalog(context.Background(), "rest", srv.URL, + rest.WithSigV4RegionSvc("us-east-1", "s3"), + rest.WithAdditionalProps(iceberg.Properties{ + iceio.S3AccessKeyID: "AKIDEXAMPLEPROPS", + iceio.S3SecretAccessKey: "secretexample", + })) + require.NoError(t, err) + require.NotNil(t, cat) + t.Cleanup(func() { _ = cat.Close() }) + + require.Contains(t, gotAuth, "Credential=AKIDEXAMPLEPROPS/", + "SigV4 must sign with the credentials from catalog properties, not the default chain") +} diff --git a/cmd/iceberg/main.go b/cmd/iceberg/main.go index 3969e111e..760dfe934 100644 --- a/cmd/iceberg/main.go +++ b/cmd/iceberg/main.go @@ -36,6 +36,11 @@ import ( "github.com/apache/iceberg-go/catalog/hadoop" "github.com/apache/iceberg-go/catalog/hive" "github.com/apache/iceberg-go/catalog/rest" + + // Register the AWS SigV4 signer backend so the --sigv4 / signing-region / + // signing-name flags work. catalog/rest itself no longer links the AWS SDK; + // the CLI already pulls it in via catalog/glue, so this adds no new weight. + _ "github.com/apache/iceberg-go/catalog/rest/sigv4" sqlcat "github.com/apache/iceberg-go/catalog/sql" "github.com/apache/iceberg-go/config" _ "github.com/apache/iceberg-go/io/gocloud" diff --git a/website/src/configuration.md b/website/src/configuration.md index 76807f304..0faebce78 100644 --- a/website/src/configuration.md +++ b/website/src/configuration.md @@ -71,11 +71,34 @@ The most option-rich surface. Source: [`catalog/rest/options.go`](https://github | Group | Options | |---|---| | Authentication | `WithCredential`, `WithOAuthToken`, `WithAuthManager`, `WithAuthURI`, `WithScope`, `WithAudience`, `WithResource` | -| AWS SigV4 | `WithSigV4`, `WithSigV4RegionSvc`, `WithAwsConfig` | +| AWS SigV4 | `WithSigV4`, `WithSigV4RegionSvc`, `WithSigner` | | HTTP | `WithHeaders`, `WithTLSConfig`, `WithOAuthTLSConfig`, `WithCustomTransport` | | Catalog routing | `WithPrefix`, `WithWarehouseLocation`, `WithMetadataLocation` | | Pass-through | `WithAdditionalProps` | +> **AWS SigV4 is an optional backend.** `catalog/rest` no longer links the AWS +> SDK. To sign REST requests with SigV4, pull in the +> [`catalog/rest/sigv4`](https://github.com/apache/iceberg-go/blob/main/catalog/rest/sigv4/sigv4.go) +> sub-package one of two ways: +> +> - Property or ambient credentials: add a blank import +> `_ "github.com/apache/iceberg-go/catalog/rest/sigv4"` and enable signing with +> `WithSigV4` / `WithSigV4RegionSvc` or the `rest.sigv4-enabled` property. The +> backend signs with the `s3.*` credential properties when set (falling back to +> the `rest.*` aliases), otherwise with the AWS default credential chain. +> - Explicit config: pass `sigv4.WithAwsConfig(cfg)` to `NewCatalog` alongside +> `WithSigV4` / `WithSigV4RegionSvc` (no blank import needed). It supplies the +> `aws.Config`; the signing service comes from those options (or a server +> `/v1/config` override), and so does the region when one is set. Otherwise +> the region falls back to `cfg.Region`. +> +> The former `rest.WithAwsConfig(aws.Config)` has been removed in favor of these +> paths. Because signing is enabled by `WithSigV4` / `WithSigV4RegionSvc` (or the +> `rest.sigv4-enabled` property) rather than by supplying a config, a program +> that loads a REST catalog with `rest.sigv4-enabled=true` now fails at startup +> until the `catalog/rest/sigv4` backend is blank-imported; the error names the +> import to add. + #### Metrics reporting The REST catalog can POST scan and commit metrics to the catalog's