diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index e62714a9..5abc626e 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -30,7 +30,11 @@ jobs: env: GH_TOKEN: ${{ github.token }} run: | - gh release create "$GITHUB_REF_NAME" \ - --title "$GITHUB_REF_NAME" \ - --generate-notes \ - goofys goofys-amd64 goofys-arm64 + if gh release view "$GITHUB_REF_NAME" >/dev/null 2>&1; then + gh release upload "$GITHUB_REF_NAME" --clobber goofys goofys-amd64 goofys-arm64 + else + gh release create "$GITHUB_REF_NAME" \ + --title "$GITHUB_REF_NAME" \ + --generate-notes \ + goofys goofys-amd64 goofys-arm64 + fi diff --git a/README.md b/README.md index 564dbb76..5fca2d65 100644 --- a/README.md +++ b/README.md @@ -126,7 +126,7 @@ Additionally, goofys also works with the following non-S3 object stores: # References * Data is stored on [Amazon S3](https://aws.amazon.com/s3/) - * [Amazon SDK for Go](https://github.com/aws/aws-sdk-go) + * [AWS SDK for Go v2](https://github.com/aws/aws-sdk-go-v2) * Other related fuse filesystems * [catfs](https://github.com/kahing/catfs): caching layer that can be used with goofys * [s3fs](https://github.com/s3fs-fuse/s3fs-fuse): another popular filesystem for S3 diff --git a/api/common/conf_s3.go b/api/common/conf_s3.go index 7b2e45eb..2c383b61 100644 --- a/api/common/conf_s3.go +++ b/api/common/conf_s3.go @@ -15,16 +15,20 @@ package common import ( + "context" "crypto/md5" "encoding/base64" "fmt" "net/http" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/aws/client" - "github.com/aws/aws-sdk-go/aws/credentials" - "github.com/aws/aws-sdk-go/aws/credentials/stscreds" - "github.com/aws/aws-sdk-go/aws/session" + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/aws/ratelimit" + "github.com/aws/aws-sdk-go-v2/aws/retry" + "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" + "github.com/aws/aws-sdk-go-v2/credentials/stscreds" + "github.com/aws/aws-sdk-go-v2/service/sts" + "github.com/aws/smithy-go/logging" ) type S3Config struct { @@ -51,14 +55,12 @@ type S3Config struct { Subdomain bool - Credentials *credentials.Credentials - Session *session.Session + Credentials aws.CredentialsProvider + Session *aws.Config BucketOwner string } -var s3Session *session.Session - func (c *S3Config) Init() *S3Config { if c.Region == "" { c.Region = "us-east-1" @@ -70,58 +72,77 @@ func (c *S3Config) Init() *S3Config { } func (c *S3Config) ToAwsConfig(flags *FlagStorage) (*aws.Config, error) { - awsConfig := (&aws.Config{ - Region: &c.Region, - Logger: GetLogger("s3"), - }).WithHTTPClient(&http.Client{ + ctx := context.Background() + httpClient := &http.Client{ Transport: &defaultHTTPTransport, Timeout: flags.HTTPTimeout, + } + log := GetLogger("s3") + sdkLogger := logging.LoggerFunc(func(_ logging.Classification, format string, args ...interface{}) { + log.Debugf(format, args...) }) + var logMode aws.ClientLogMode if flags.DebugS3 { - awsConfig.LogLevel = aws.LogLevel(aws.LogDebug | aws.LogDebugWithRequestErrors) + logMode = aws.LogRequest | aws.LogResponse | aws.LogRetries } - - if c.Credentials == nil { - if c.AccessKey != "" { - c.Credentials = credentials.NewStaticCredentials(c.AccessKey, c.SecretKey, "") - } else if c.Profile != "" { - c.Credentials = newSharedFileCredentials(c.Profile) - } + loadOptions := []func(*config.LoadOptions) error{ + config.WithRegion(c.Region), + config.WithHTTPClient(httpClient), + config.WithLogger(sdkLogger), + config.WithClientLogMode(logMode), } - if flags.Endpoint != "" { - awsConfig.Endpoint = &flags.Endpoint + if c.Credentials == nil && c.AccessKey != "" { + c.Credentials = credentials.NewStaticCredentialsProvider(c.AccessKey, c.SecretKey, "") } - - awsConfig.S3ForcePathStyle = aws.Bool(!c.Subdomain) - if c.Session == nil { - if s3Session == nil { - var err error - s3Session, err = session.NewSessionWithOptions(session.Options{ - Profile: c.Profile, - SharedConfigState: session.SharedConfigEnable, - }) - if err != nil { - return nil, err - } + options := sharedConfigLoadOptions(c.Profile, loadOptions) + if c.Credentials != nil { + options = append(options, config.WithCredentialsProvider(c.Credentials)) } - c.Session = s3Session + loaded, err := config.LoadDefaultConfig(ctx, options...) + if err != nil { + return nil, err + } + if c.Credentials == nil { + loaded.Credentials = &sharedFileProvider{profile: c.Profile, loadOptions: loadOptions} + } + c.Session = &loaded + } else if c.Credentials == nil && c.Profile != "" { + c.Credentials = &sharedFileProvider{profile: c.Profile, loadOptions: loadOptions} } - if c.RoleArn != "" { - c.Credentials = stscreds.NewCredentials(stsConfigProvider{c}, c.RoleArn, - func(p *stscreds.AssumeRoleProvider) { - if c.RoleExternalId != "" { - p.ExternalID = &c.RoleExternalId - } - p.RoleSessionName = c.RoleSessionName + awsConfig := c.Session.Copy() + awsConfig.Region = c.Region + awsConfig.HTTPClient = httpClient + awsConfig.Logger = sdkLogger + awsConfig.ClientLogMode = logMode + if awsConfig.Retryer == nil { + awsConfig.Retryer = func() aws.Retryer { + return retry.NewStandard(func(options *retry.StandardOptions) { + options.MaxAttempts = 4 + options.RateLimiter = ratelimit.None }) + } } - if c.Credentials != nil { awsConfig.Credentials = c.Credentials } + if c.RoleArn != "" { + stsClient := sts.NewFromConfig(awsConfig, func(options *sts.Options) { + if c.StsEndpoint != "" { + options.BaseEndpoint = aws.String(c.StsEndpoint) + } + }) + awsConfig.Credentials = aws.NewCredentialsCache(stscreds.NewAssumeRoleProvider(stsClient, c.RoleArn, + func(options *stscreds.AssumeRoleOptions) { + if c.RoleExternalId != "" { + options.ExternalID = aws.String(c.RoleExternalId) + } + options.RoleSessionName = c.RoleSessionName + })) + } + if c.SseC != "" { key, err := base64.StdEncoding.DecodeString(c.SseC) if err != nil { @@ -133,21 +154,5 @@ func (c *S3Config) ToAwsConfig(flags *FlagStorage) (*aws.Config, error) { c.SseCDigest = base64.StdEncoding.EncodeToString(m[:]) } - return awsConfig, nil -} - -type stsConfigProvider struct { - *S3Config -} - -func (c stsConfigProvider) ClientConfig(serviceName string, cfgs ...*aws.Config) client.Config { - config := c.Session.ClientConfig(serviceName, cfgs...) - if c.Credentials != nil { - config.Config.Credentials = c.Credentials - } - if c.StsEndpoint != "" { - config.Endpoint = c.StsEndpoint - } - - return config + return &awsConfig, nil } diff --git a/api/common/conf_s3_credentials.go b/api/common/conf_s3_credentials.go index b80d1d16..5d5f04c3 100644 --- a/api/common/conf_s3_credentials.go +++ b/api/common/conf_s3_credentials.go @@ -15,13 +15,13 @@ package common import ( + "context" "os" "sync" "time" - "github.com/aws/aws-sdk-go/aws/credentials" - "github.com/aws/aws-sdk-go/aws/defaults" - "github.com/aws/aws-sdk-go/aws/session" + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/config" ) const sharedCredentialsStatInterval = 10 * time.Second @@ -38,45 +38,36 @@ func (s sharedFileState) equal(other sharedFileState) bool { } type sharedFileProvider struct { - mu sync.Mutex - profile string - resolved *credentials.Credentials - value credentials.Value - state sharedFileState - loaded bool - expired bool - unreadable bool - checkedAt time.Time + mu sync.Mutex + profile string + loadOptions []func(*config.LoadOptions) error + value aws.Credentials + state sharedFileState + loaded bool + expired bool + unreadable bool + checkedAt time.Time } -func newSharedFileCredentials(profile string) *credentials.Credentials { - creds := credentials.NewCredentials(&sharedFileProvider{profile: profile}) - if _, err := creds.Get(); err != nil { - credentialsLog.Warnf("cannot resolve credentials for profile %v: %v", profile, err) - } - return creds -} - -func (p *sharedFileProvider) Retrieve() (credentials.Value, error) { +func (p *sharedFileProvider) Retrieve(ctx context.Context) (aws.Credentials, error) { p.mu.Lock() defer p.mu.Unlock() + if !p.isExpired() { + return p.value, nil + } state, stated := p.fileState() - sess, err := session.NewSessionWithOptions(session.Options{ - Profile: p.profile, - SharedConfigState: session.SharedConfigEnable, - }) + cfg, err := config.LoadDefaultConfig(ctx, sharedConfigLoadOptions(p.profile, p.loadOptions)...) if err != nil { return p.keepLoaded(err) } - value, err := sess.Config.Credentials.Get() + value, err := cfg.Credentials.Retrieve(ctx) if err != nil { return p.keepLoaded(err) } - p.resolved = sess.Config.Credentials p.value = value p.loaded = true p.expired = false @@ -91,22 +82,31 @@ func (p *sharedFileProvider) Retrieve() (credentials.Value, error) { func (p *sharedFileProvider) IsExpired() bool { p.mu.Lock() defer p.mu.Unlock() + return p.isExpired() +} + +func (p *sharedFileProvider) Expire() { + p.mu.Lock() + defer p.mu.Unlock() + p.expired = true +} +func (p *sharedFileProvider) isExpired() bool { if !p.loaded || p.expired { return true } + if p.value.Expired() { + p.expired = true + return true + } + now := time.Now() if now.Sub(p.checkedAt) < sharedCredentialsStatInterval { return false } p.checkedAt = now - if p.resolved.IsExpired() { - p.expired = true - return true - } - state, stated := p.fileState() if !stated { if !p.unreadable { @@ -122,9 +122,9 @@ func (p *sharedFileProvider) IsExpired() bool { return p.expired } -func (p *sharedFileProvider) keepLoaded(err error) (credentials.Value, error) { +func (p *sharedFileProvider) keepLoaded(err error) (aws.Credentials, error) { if !p.loaded { - return credentials.Value{}, err + return aws.Credentials{}, err } p.expired = false @@ -134,6 +134,14 @@ func (p *sharedFileProvider) keepLoaded(err error) (credentials.Value, error) { return p.value, nil } +func sharedConfigLoadOptions(profile string, base []func(*config.LoadOptions) error) []func(*config.LoadOptions) error { + options := append([]func(*config.LoadOptions) error{}, base...) + if profile != "" { + options = append(options, config.WithSharedConfigProfile(profile)) + } + return options +} + func (p *sharedFileProvider) fileState() (sharedFileState, bool) { info, err := os.Stat(sharedCredentialsFilename()) if err != nil { @@ -146,5 +154,5 @@ func sharedCredentialsFilename() string { if filename := os.Getenv("AWS_SHARED_CREDENTIALS_FILE"); filename != "" { return filename } - return defaults.SharedCredentialsFilename() + return config.DefaultSharedCredentialsFilename() } diff --git a/api/common/conf_s3_credentials_test.go b/api/common/conf_s3_credentials_test.go index a893e547..ba612a2b 100644 --- a/api/common/conf_s3_credentials_test.go +++ b/api/common/conf_s3_credentials_test.go @@ -15,13 +15,12 @@ package common import ( + "context" "fmt" "os" "path/filepath" "testing" "time" - - "github.com/aws/aws-sdk-go/aws/credentials" ) func writeSharedCredentials(t *testing.T, path, profile, accessKey, secretKey, sessionToken string) { @@ -41,6 +40,14 @@ func sharedCredentialsFile(t *testing.T) string { t.Setenv("AWS_CONFIG_FILE", filepath.Join(dir, "config")) t.Setenv("AWS_EC2_METADATA_DISABLED", "true") t.Setenv("AWS_CONTAINER_CREDENTIALS_RELATIVE_URI", "") + t.Setenv("AWS_CONTAINER_CREDENTIALS_FULL_URI", "") + t.Setenv("AWS_ACCESS_KEY_ID", "") + t.Setenv("AWS_SECRET_ACCESS_KEY", "") + t.Setenv("AWS_SESSION_TOKEN", "") + t.Setenv("AWS_PROFILE", "") + t.Setenv("AWS_DEFAULT_PROFILE", "") + t.Setenv("AWS_WEB_IDENTITY_TOKEN_FILE", "") + t.Setenv("AWS_ROLE_ARN", "") return path } @@ -62,7 +69,7 @@ func TestSharedFileCredentialsRereadAfterExpire(t *testing.T) { t.Fatal("newSharedFileCredentials() = nil, want credentials") } - value, err := creds.Get() + value, err := creds.Retrieve(t.Context()) if err != nil { t.Fatal(err) } @@ -73,7 +80,7 @@ func TestSharedFileCredentialsRereadAfterExpire(t *testing.T) { writeSharedCredentials(t, path, "bucket", "AKIANEW", "secret-new", "token-new") creds.Expire() - value, err = creds.Get() + value, err = creds.Retrieve(t.Context()) if err != nil { t.Fatal(err) } @@ -87,9 +94,9 @@ func TestSharedFileCredentialsRereadAfterRotation(t *testing.T) { writeSharedCredentials(t, path, "bucket", "AKIAOLD", "secret-old", "token-old") provider := &sharedFileProvider{profile: "bucket"} - creds := credentials.NewCredentials(provider) + creds := provider - value, err := creds.Get() + value, err := creds.Retrieve(t.Context()) if err != nil { t.Fatal(err) } @@ -100,7 +107,7 @@ func TestSharedFileCredentialsRereadAfterRotation(t *testing.T) { rotateSharedCredentials(t, path, "bucket", "AKIANEW", "secret-new", "token-new") provider.checkedAt = time.Now().Add(-sharedCredentialsStatInterval) - value, err = creds.Get() + value, err = creds.Retrieve(t.Context()) if err != nil { t.Fatal(err) } @@ -114,7 +121,7 @@ func TestSharedFileProviderIsExpired(t *testing.T) { writeSharedCredentials(t, path, "bucket", "AKIAOLD", "secret-old", "token-old") provider := &sharedFileProvider{profile: "bucket"} - if _, err := provider.Retrieve(); err != nil { + if _, err := provider.Retrieve(t.Context()); err != nil { t.Fatal(err) } @@ -148,7 +155,7 @@ func TestSharedFileProviderIsExpiredIgnoresPreservedTimestamp(t *testing.T) { } provider := &sharedFileProvider{profile: "bucket"} - if _, err := provider.Retrieve(); err != nil { + if _, err := provider.Retrieve(t.Context()); err != nil { t.Fatal(err) } @@ -168,7 +175,7 @@ func TestSharedFileProviderMissingFileKeepsCredentials(t *testing.T) { writeSharedCredentials(t, path, "bucket", "AKIAOLD", "secret-old", "token-old") provider := &sharedFileProvider{profile: "bucket"} - if _, err := provider.Retrieve(); err != nil { + if _, err := provider.Retrieve(t.Context()); err != nil { t.Fatal(err) } @@ -187,13 +194,14 @@ func TestSharedFileProviderFailedReloadKeepsCredentials(t *testing.T) { writeSharedCredentials(t, path, "bucket", "AKIAOLD", "secret-old", "token-old") provider := &sharedFileProvider{profile: "bucket"} - if _, err := provider.Retrieve(); err != nil { + if _, err := provider.Retrieve(t.Context()); err != nil { t.Fatal(err) } rotateSharedCredentials(t, path, "other", "AKIANEW", "secret-new", "token-new") + provider.checkedAt = time.Now().Add(-sharedCredentialsStatInterval) - value, err := provider.Retrieve() + value, err := provider.Retrieve(t.Context()) if err != nil { t.Fatalf("Retrieve() = %v, want the previously loaded credentials", err) } @@ -202,6 +210,57 @@ func TestSharedFileProviderFailedReloadKeepsCredentials(t *testing.T) { } } +func TestSharedFileProviderFailedReloadRecovers(t *testing.T) { + path := sharedCredentialsFile(t) + writeSharedCredentials(t, path, "bucket", "OLD", "old-secret", "old-token") + provider := newSharedFileCredentials("bucket") + if err := os.WriteFile(path, []byte("[invalid"), 0600); err != nil { + t.Fatal(err) + } + provider.checkedAt = time.Now().Add(-sharedCredentialsStatInterval) + value, err := provider.Retrieve(t.Context()) + if err != nil || value.AccessKeyID != "OLD" { + t.Fatalf("failed reload must preserve old credentials: %v, %q", err, value.AccessKeyID) + } + rotateSharedCredentials(t, path, "bucket", "NEW", "new-secret", "new-token") + provider.checkedAt = time.Now().Add(-sharedCredentialsStatInterval) + value, err = provider.Retrieve(t.Context()) + if err != nil || value.AccessKeyID != "NEW" { + t.Fatalf("provider did not recover after file became readable: %v, %q", err, value.AccessKeyID) + } +} + +func TestSharedFileProviderRefreshesExpiredCredentials(t *testing.T) { + for _, profile := range []string{"", "bucket"} { + t.Run("profile="+profile, func(t *testing.T) { + path := sharedCredentialsFile(t) + fileProfile := profile + if fileProfile == "" { + fileProfile = "default" + } + writeSharedCredentials(t, path, fileProfile, "KEY", "secret", "token") + cfg, err := (&S3Config{Profile: profile}).Init().ToAwsConfig(&FlagStorage{}) + if err != nil { + t.Fatal(err) + } + if _, err := cfg.Credentials.Retrieve(t.Context()); err != nil { + t.Fatal(err) + } + provider := cfg.Credentials.(*sharedFileProvider) + provider.value.CanExpire = true + provider.value.Expires = time.Now().Add(-time.Minute) + provider.checkedAt = time.Now() + value, err := cfg.Credentials.Retrieve(t.Context()) + if err != nil { + t.Fatal(err) + } + if value.Expired() || value.CanExpire { + t.Fatal("expired provider credentials were not refreshed within the stat interval") + } + }) + } +} + func TestNewSharedFileCredentialsResolvesConfigProfile(t *testing.T) { path := sharedCredentialsFile(t) writeSharedCredentials(t, path, "other", "AKIAOTHER", "secret-other", "token-other") @@ -216,7 +275,7 @@ func TestNewSharedFileCredentialsResolvesConfigProfile(t *testing.T) { t.Fatal("newSharedFileCredentials() = nil, want the profile resolved from the config file") } - value, err := creds.Get() + value, err := creds.Retrieve(t.Context()) if err != nil { t.Fatal(err) } @@ -234,7 +293,7 @@ func TestNewSharedFileCredentialsUnknownProfile(t *testing.T) { t.Fatal("newSharedFileCredentials() = nil, want credentials that report the resolution error") } - if _, err := creds.Get(); err == nil { + if _, err := creds.Retrieve(t.Context()); err == nil { t.Error("a profile missing from the shared configuration must not resolve to another profile") } } @@ -243,10 +302,6 @@ func TestToAwsConfigProfileCredentialsFollowRotation(t *testing.T) { path := sharedCredentialsFile(t) writeSharedCredentials(t, path, "bucket", "AKIAOLD", "secret-old", "token-old") - previousSession := s3Session - s3Session = nil - t.Cleanup(func() { s3Session = previousSession }) - config := (&S3Config{Profile: "bucket"}).Init() awsConfig, err := config.ToAwsConfig(&FlagStorage{}) if err != nil { @@ -256,7 +311,7 @@ func TestToAwsConfigProfileCredentialsFollowRotation(t *testing.T) { t.Fatal("ToAwsConfig() left credentials unset for a profile mount") } - value, err := awsConfig.Credentials.Get() + value, err := awsConfig.Credentials.Retrieve(t.Context()) if err != nil { t.Fatal(err) } @@ -265,9 +320,9 @@ func TestToAwsConfigProfileCredentialsFollowRotation(t *testing.T) { } writeSharedCredentials(t, path, "bucket", "AKIANEW", "secret-new", "token-new") - awsConfig.Credentials.Expire() + awsConfig.Credentials.(*sharedFileProvider).Expire() - value, err = awsConfig.Credentials.Get() + value, err = awsConfig.Credentials.Retrieve(t.Context()) if err != nil { t.Fatal(err) } @@ -275,3 +330,9 @@ func TestToAwsConfigProfileCredentialsFollowRotation(t *testing.T) { t.Errorf("credentials after expiry = %+v, want the rotated profile", value) } } + +func newSharedFileCredentials(profile string) *sharedFileProvider { + creds := &sharedFileProvider{profile: profile} + creds.Retrieve(context.Background()) + return creds +} diff --git a/api/common/conf_s3_test.go b/api/common/conf_s3_test.go new file mode 100644 index 00000000..a63d0158 --- /dev/null +++ b/api/common/conf_s3_test.go @@ -0,0 +1,287 @@ +package common + +import ( + "crypto/md5" + "encoding/base64" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "os" + "strings" + "sync" + "testing" + "time" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/aws/retry" + "github.com/aws/aws-sdk-go-v2/credentials" + "github.com/aws/aws-sdk-go-v2/service/s3" +) + +func TestToAwsConfigCredentialPrecedence(t *testing.T) { + for _, tc := range []struct { + name string + config S3Config + envProfile string + want string + }{ + {name: "explicit provider", config: S3Config{Credentials: credentials.NewStaticCredentialsProvider("CUSTOM", "custom-secret", ""), AccessKey: "STATIC", SecretKey: "static-secret", Profile: "bucket"}, want: "CUSTOM"}, + {name: "static keys", config: S3Config{AccessKey: "STATIC", SecretKey: "static-secret", Profile: "bucket"}, want: "STATIC"}, + {name: "explicit profile", config: S3Config{Profile: "bucket"}, want: "PROFILE"}, + {name: "environment", want: "ENV"}, + {name: "environment profile with environment keys", envProfile: "bucket", want: "ENV"}, + {name: "supplied session", config: S3Config{Session: &aws.Config{Credentials: credentials.NewStaticCredentialsProvider("SESSION", "session-secret", "")}}, want: "SESSION"}, + } { + t.Run(tc.name, func(t *testing.T) { + path := sharedCredentialsFile(t) + writeSharedCredentials(t, path, "bucket", "PROFILE", "profile-secret", "") + t.Setenv("AWS_ACCESS_KEY_ID", "ENV") + t.Setenv("AWS_SECRET_ACCESS_KEY", "env-secret") + t.Setenv("AWS_PROFILE", tc.envProfile) + cfg, err := tc.config.Init().ToAwsConfig(&FlagStorage{}) + if err != nil { + t.Fatal(err) + } + value, err := cfg.Credentials.Retrieve(t.Context()) + if err != nil { + t.Fatal(err) + } + if value.AccessKeyID != tc.want { + t.Fatalf("credential source = %q, want %q", value.AccessKeyID, tc.want) + } + }) + } +} + +func TestToAwsConfigProfilesDoNotShareSession(t *testing.T) { + path := sharedCredentialsFile(t) + if err := os.WriteFile(path, []byte("[first]\naws_access_key_id=FIRST\naws_secret_access_key=first-secret\n[second]\naws_access_key_id=SECOND\naws_secret_access_key=second-secret\n"), 0600); err != nil { + t.Fatal(err) + } + first := (&S3Config{Profile: "first"}).Init() + second := (&S3Config{Profile: "second"}).Init() + for _, c := range []*S3Config{first, second} { + cfg, err := c.ToAwsConfig(&FlagStorage{}) + if err != nil { + t.Fatal(err) + } + value, err := cfg.Credentials.Retrieve(t.Context()) + if err != nil { + t.Fatal(err) + } + if value.AccessKeyID != strings.ToUpper(c.Profile) { + t.Fatalf("profile %q resolved %q", c.Profile, value.AccessKeyID) + } + } + if first.Session == second.Session { + t.Fatal("independent profiles must not share session configuration") + } +} + +func TestToAwsConfigDefaultCredentialsRotate(t *testing.T) { + for _, profile := range []string{"", "bucket"} { + t.Run("AWS_PROFILE="+profile, func(t *testing.T) { + path := sharedCredentialsFile(t) + t.Setenv("AWS_PROFILE", profile) + if profile == "" { + profile = "default" + } + writeSharedCredentials(t, path, profile, "OLD", "old-secret", "old-token") + cfg, err := (&S3Config{}).Init().ToAwsConfig(&FlagStorage{}) + if err != nil { + t.Fatal(err) + } + if _, err := cfg.Credentials.Retrieve(t.Context()); err != nil { + t.Fatal(err) + } + rotateSharedCredentials(t, path, profile, "NEW", "new-secret", "new-token") + provider := cfg.Credentials.(*sharedFileProvider) + provider.checkedAt = time.Now().Add(-sharedCredentialsStatInterval) + var wg sync.WaitGroup + for range 16 { + wg.Add(1) + go func() { + defer wg.Done() + value, err := cfg.Credentials.Retrieve(t.Context()) + if err != nil { + t.Error(err) + return + } + if value.AccessKeyID != "NEW" || value.SessionToken != "new-token" { + t.Error("default credentials did not rotate") + } + }() + } + wg.Wait() + }) + } +} + +func TestS3ClientCredentialsFollowRotation(t *testing.T) { + for _, profile := range []string{"", "bucket"} { + t.Run("profile="+profile, func(t *testing.T) { + path := sharedCredentialsFile(t) + fileProfile := profile + if fileProfile == "" { + fileProfile = "default" + } + writeSharedCredentials(t, path, fileProfile, "OLD", "old-secret", "old-token") + authorizations := make(chan string, 2) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + authorizations <- r.Header.Get("Authorization") + w.WriteHeader(http.StatusOK) + })) + defer server.Close() + cfg, err := (&S3Config{Profile: profile}).Init().ToAwsConfig(&FlagStorage{HTTPTimeout: time.Second}) + if err != nil { + t.Fatal(err) + } + client := s3.NewFromConfig(*cfg, func(options *s3.Options) { options.BaseEndpoint = aws.String(server.URL); options.UsePathStyle = true }) + for _, key := range []string{"OLD", "NEW"} { + if _, err := client.HeadBucket(t.Context(), &s3.HeadBucketInput{Bucket: aws.String("bucket")}); err != nil { + t.Fatal(err) + } + if auth := <-authorizations; !strings.Contains(auth, "Credential="+key+"/") { + t.Fatalf("request authorization did not use %s: %s", key, auth) + } + if key == "OLD" { + rotateSharedCredentials(t, path, fileProfile, "NEW", "new-secret", "new-token") + cfg.Credentials.(*sharedFileProvider).checkedAt = time.Now().Add(-sharedCredentialsStatInterval) + } + } + }) + } +} + +func TestToAwsConfigAssumeRole(t *testing.T) { + sharedCredentialsFile(t) + var calls int + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + calls++ + if err := r.ParseForm(); err != nil { + t.Error(err) + } + for key, want := range map[string]string{"Action": "AssumeRole", "RoleArn": "arn:aws:iam::123456789012:role/bucket", "ExternalId": "external", "RoleSessionName": "mount"} { + if got := r.Form.Get(key); got != want { + t.Errorf("%s = %q, want %q", key, got, want) + } + } + if !strings.Contains(r.Header.Get("Authorization"), "Credential=SOURCE/") { + t.Errorf("STS used incorrect source: %s", r.Header.Get("Authorization")) + } + w.Header().Set("Content-Type", "text/xml") + fmt.Fprint(w, `ASSUMEDassumed-secretassumed-token2099-01-01T00:00:00Zid:mountarn:aws:sts::123456789012:assumed-role/bucket/mount`) + })) + defer server.Close() + cfg, err := (&S3Config{AccessKey: "SOURCE", SecretKey: "source-secret", RoleArn: "arn:aws:iam::123456789012:role/bucket", RoleExternalId: "external", RoleSessionName: "mount", StsEndpoint: server.URL}).Init().ToAwsConfig(&FlagStorage{Endpoint: "http://must-not-receive-sts.invalid", HTTPTimeout: time.Second}) + if err != nil { + t.Fatal(err) + } + for range 2 { + value, err := cfg.Credentials.Retrieve(t.Context()) + if err != nil { + t.Fatal(err) + } + if value.AccessKeyID != "ASSUMED" || value.SessionToken != "assumed-token" { + t.Fatal("role credentials were not returned") + } + } + if calls != 1 { + t.Fatalf("STS requests = %d, want one cached assumption", calls) + } +} + +func TestToAwsConfigNamedRoleProfile(t *testing.T) { + path := sharedCredentialsFile(t) + writeSharedCredentials(t, path, "source", "SOURCE", "source-secret", "") + t.Setenv("AWS_ACCESS_KEY_ID", "ENV") + t.Setenv("AWS_SECRET_ACCESS_KEY", "env-secret") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if err := r.ParseForm(); err != nil { + t.Error(err) + } + if r.Form.Get("ExternalId") != "profile-external" || r.Form.Get("RoleSessionName") != "profile-session" { + t.Error("named role options were not preserved") + } + if !strings.Contains(r.Header.Get("Authorization"), "Credential=SOURCE/") { + t.Error("explicit role profile must override environment credentials") + } + w.Header().Set("Content-Type", "text/xml") + fmt.Fprint(w, `PROFILE-ROLErole-secretrole-token2099-01-01T00:00:00Z`) + })) + defer server.Close() + t.Setenv("AWS_ENDPOINT_URL_STS", server.URL) + contents := "[profile bucket]\nrole_arn=arn:aws:iam::123456789012:role/bucket\nsource_profile=source\nexternal_id=profile-external\nrole_session_name=profile-session\n" + if err := os.WriteFile(os.Getenv("AWS_CONFIG_FILE"), []byte(contents), 0600); err != nil { + t.Fatal(err) + } + cfg, err := (&S3Config{Profile: "bucket"}).Init().ToAwsConfig(&FlagStorage{HTTPTimeout: time.Second}) + if err != nil { + t.Fatal(err) + } + value, err := cfg.Credentials.Retrieve(t.Context()) + if err != nil { + t.Fatal(err) + } + if value.AccessKeyID != "PROFILE-ROLE" { + t.Fatalf("named role profile resolved %q", value.AccessKeyID) + } +} + +func TestToAwsConfigRetryPolicy(t *testing.T) { + sharedCredentialsFile(t) + cfg, err := (&S3Config{AccessKey: "KEY", SecretKey: "secret"}).Init().ToAwsConfig(&FlagStorage{}) + if err != nil { + t.Fatal(err) + } + r := cfg.Retryer() + if r.MaxAttempts() != 4 { + t.Fatalf("max attempts = %d, want 4", r.MaxAttempts()) + } + for range 200 { + release, err := r.GetRetryToken(t.Context(), errors.New("retry")) + if err != nil { + t.Fatalf("retry quota must not limit requests: %v", err) + } + if err := release(errors.New("failed")); err != nil { + t.Fatal(err) + } + } + custom := retry.NewStandard(func(options *retry.StandardOptions) { options.MaxAttempts = 7 }) + cfg, err = (&S3Config{AccessKey: "KEY", SecretKey: "secret", Session: &aws.Config{Retryer: func() aws.Retryer { return custom }}}).Init().ToAwsConfig(&FlagStorage{}) + if err != nil { + t.Fatal(err) + } + if cfg.Retryer() != custom { + t.Fatal("custom session retryer was replaced") + } +} + +func TestToAwsConfigHTTPLoggingAndSSEC(t *testing.T) { + sharedCredentialsFile(t) + key := "0123456789abcdef0123456789abcdef" + c := (&S3Config{AccessKey: "KEY", SecretKey: "secret", Region: "us-west-2", SseC: base64.StdEncoding.EncodeToString([]byte(key))}).Init() + cfg, err := c.ToAwsConfig(&FlagStorage{HTTPTimeout: 17 * time.Second, DebugS3: true}) + if err != nil { + t.Fatal(err) + } + client, ok := cfg.HTTPClient.(*http.Client) + if !ok || client.Timeout != 17*time.Second || client.Transport != &defaultHTTPTransport { + t.Fatal("configured HTTP client was not preserved") + } + if cfg.Logger == nil || !cfg.ClientLogMode.IsRequest() || !cfg.ClientLogMode.IsResponse() || !cfg.ClientLogMode.IsRetries() { + t.Fatal("SDK debug logging was not configured") + } + if cfg.Region != "us-west-2" { + t.Fatalf("region = %q", cfg.Region) + } + digest := md5.Sum([]byte(key)) + if c.SseC != key || c.SseCDigest != base64.StdEncoding.EncodeToString(digest[:]) { + t.Fatal("SSE-C key or digest changed") + } + invalid := (&S3Config{AccessKey: "KEY", SecretKey: "secret", SseC: "not base64"}).Init() + if _, err := invalid.ToAwsConfig(&FlagStorage{}); err == nil || !strings.Contains(err.Error(), "sse-c is not base64-encoded") { + t.Fatalf("invalid SSE-C error = %v", err) + } +} diff --git a/go.mod b/go.mod index b19550cf..3dc05f5d 100644 --- a/go.mod +++ b/go.mod @@ -11,7 +11,12 @@ require ( github.com/Azure/go-autorest/autorest/adal v0.9.13 github.com/Azure/go-autorest/autorest/azure/auth v0.5.7 github.com/Azure/go-autorest/autorest/azure/cli v0.4.2 - github.com/aws/aws-sdk-go v1.44.37 + github.com/aws/aws-sdk-go-v2 v1.45.1 + github.com/aws/aws-sdk-go-v2/config v1.33.2 + github.com/aws/aws-sdk-go-v2/credentials v1.20.2 + github.com/aws/aws-sdk-go-v2/service/s3 v1.110.0 + github.com/aws/aws-sdk-go-v2/service/sts v1.48.0 + github.com/aws/smithy-go v1.28.1 github.com/gofrs/uuid v4.2.0+incompatible github.com/google/uuid v1.2.0 github.com/jacobsa/fuse v0.0.0-20221016084658-a4cd154343d8 @@ -37,6 +42,18 @@ require ( github.com/Azure/go-autorest/autorest/validation v0.3.1 // indirect github.com/Azure/go-autorest/logger v0.2.1 // indirect github.com/Azure/go-autorest/tracing v0.6.0 // indirect + github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.20 // indirect + github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.19.1 // indirect + github.com/aws/aws-sdk-go-v2/internal/configsources v1.5.1 // indirect + github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.8.1 // indirect + github.com/aws/aws-sdk-go-v2/internal/v4a v1.5.1 // indirect + github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.19 // indirect + github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.11.1 // indirect + github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.14.1 // indirect + github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.20.1 // indirect + github.com/aws/aws-sdk-go-v2/service/signin v1.8.0 // indirect + github.com/aws/aws-sdk-go-v2/service/sso v1.36.0 // indirect + github.com/aws/aws-sdk-go-v2/service/ssooidc v1.41.0 // indirect github.com/dimchansky/utfbom v1.1.1 // indirect github.com/ebitengine/purego v0.8.4 // indirect github.com/form3tech-oss/jwt-go v3.2.2+incompatible // indirect @@ -45,7 +62,6 @@ require ( github.com/golang/protobuf v1.4.3 // indirect github.com/googleapis/gax-go/v2 v2.0.5 // indirect github.com/gopherjs/gopherjs v0.0.0-20210413103415-7d3cbed7d026 // indirect - github.com/jmespath/go-jmespath v0.4.0 // indirect github.com/jstemmer/go-junit-report v0.9.1 // indirect github.com/konsorten/go-windows-terminal-sequences v1.0.1 // indirect github.com/kr/pretty v0.2.1 // indirect @@ -70,5 +86,4 @@ require ( google.golang.org/genproto v0.0.0-20210319143718-93e7006c17a6 // indirect google.golang.org/grpc v1.36.0 // indirect google.golang.org/protobuf v1.25.0 // indirect - gopkg.in/yaml.v2 v2.4.0 // indirect ) diff --git a/go.sum b/go.sum index bb0d6f74..107c1e52 100644 --- a/go.sum +++ b/go.sum @@ -80,8 +80,42 @@ github.com/alecthomas/units v0.0.0-20151022065526-2efee857e7cf/go.mod h1:ybxpYRF github.com/armon/circbuf v0.0.0-20150827004946-bbbad097214e/go.mod h1:3U/XgcO3hCbHZ8TKRvWD2dDTCfh9M9ya+I9JpbB7O8o= github.com/armon/go-metrics v0.0.0-20180917152333-f0300d1749da/go.mod h1:Q73ZrmVTwzkszR9V5SSuryQ31EELlFMUz1kKyl939pY= github.com/armon/go-radix v0.0.0-20180808171621-7fddfc383310/go.mod h1:ufUuZ+zHj4x4TnLV4JWEpy2hxWSpsRywHrMgIH9cCH8= -github.com/aws/aws-sdk-go v1.44.37 h1:KvDxCX6dfJeEDC77U5GPGSP0ErecmNnhDHFxw+NIvlI= -github.com/aws/aws-sdk-go v1.44.37/go.mod h1:y4AeaBuwd2Lk+GepC1E9v0qOiTws0MIWAX4oIKwKHZo= +github.com/aws/aws-sdk-go-v2 v1.45.1 h1:iIoG3NaLhV6UZpPXyPXlDj2I9oS8tV/nMcMnITCC6Ks= +github.com/aws/aws-sdk-go-v2 v1.45.1/go.mod h1:bttEH6JqnUL8LepvDVfdrds/fZ5bCIxzpe3abyUrhDU= +github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.20 h1:GPRlPwz40I2B2VrBEASOA3Bi77NyeqejNLkifosX0rs= +github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.20/go.mod h1:g7PNzKcsOKWb4fkSRBA7BZVAS6Y8IcxzN+nRohhQ1Q8= +github.com/aws/aws-sdk-go-v2/config v1.33.2 h1:Pj4+nF2kc4Z+1BJysVPnX9d5dMN7IYFXR4UJaWK2IpA= +github.com/aws/aws-sdk-go-v2/config v1.33.2/go.mod h1:Igw+HTwbR2tsTU/ydifAS9EHAFJ2s/FCgkwQWFnAdE4= +github.com/aws/aws-sdk-go-v2/credentials v1.20.2 h1:VQjZODPNfdikCX2ZZrltw4zNLkcwjyUFDUl2vT9yTwg= +github.com/aws/aws-sdk-go-v2/credentials v1.20.2/go.mod h1:OmeHCn28vZylsBvalLDf7t8fuJ2rHYQprJs+7WuxniI= +github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.19.1 h1:YIEBqcqRnpi4Pfv0YHImtgi6czGCwKHANC7SwmUAVD0= +github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.19.1/go.mod h1:imEf0oufgAo8KAkCHhrOdqGEC0YWx1PPBQH82shSxGw= +github.com/aws/aws-sdk-go-v2/internal/configsources v1.5.1 h1:pc138gM1CW+XPc60rEwUlwwuwWFQK16CI1T7v1F9Oec= +github.com/aws/aws-sdk-go-v2/internal/configsources v1.5.1/go.mod h1:1+koxpPIbfBdfzP6vojm5/zTpTQ/micYwlxIiNB3TxI= +github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.8.1 h1:K0JsbZQj+1h208Ro1zHeA4l7bMp0NvRffHQ91q8Ol1s= +github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.8.1/go.mod h1:W3/vL6EtCIatICGy9ab29QhMuae+cOKPWcMxv02CO+Q= +github.com/aws/aws-sdk-go-v2/internal/v4a v1.5.1 h1:yhw5KD1phVyP9vijxOUzDfEtJx+bt+L63k+VfuiYFAA= +github.com/aws/aws-sdk-go-v2/internal/v4a v1.5.1/go.mod h1:ZW2e0d7DYlRxlS9hEiMXE47gTdX5KRN4byUiNbUpG+Q= +github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.19 h1:bAdDl/HkGCcGPoe25ToSHEw23VIxt6CT5fLcg111BKg= +github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.19/go.mod h1:KaUzbLxv4CeSxh6ZCl9B4m7CuFenS8kUEaDs+f/DQr4= +github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.11.1 h1:s67hBfG5t9rn1NCvDuB4E3QIep3UFhHPtaIqFDjV3N8= +github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.11.1/go.mod h1:FpvjBMXtSNMLPmDJsWwcY5cRnqJlpS2y1R6n4pvzs4k= +github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.14.1 h1:RmmWQPREQdk9U+PfqeHW3MqZaBaNK7TpV9W3RY+b+7g= +github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.14.1/go.mod h1:0A3W4F+68ZnNk5XcNL/e9HFMwnP8RlEicFfy6eOEDyw= +github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.20.1 h1:ZMbtPZZQRca+3+XYQne9PBvRiYpHZlNJJOZfE9WNfT0= +github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.20.1/go.mod h1:YAGWQdCYlVCoqrzvfv3RLxO6zKwti7gsAULOGWPLYv4= +github.com/aws/aws-sdk-go-v2/service/s3 v1.110.0 h1:He8vaTTqAAJrux/KdpjFXNWueLJZyKqE49QEXoqAu4I= +github.com/aws/aws-sdk-go-v2/service/s3 v1.110.0/go.mod h1:CUr46sCpGAg/rHaclRyhJX0LJAmH73uWSJPPSaMUrSk= +github.com/aws/aws-sdk-go-v2/service/signin v1.8.0 h1:bSvKIoLuRGFqGwASgeCQncCJDi9YKKBDEmCEZzOX1uU= +github.com/aws/aws-sdk-go-v2/service/signin v1.8.0/go.mod h1:9IqUlsJDbUPcg6cgx3WEzXdjrbWzLDQrak0aaSqlTcI= +github.com/aws/aws-sdk-go-v2/service/sso v1.36.0 h1:iivsh357VnfIc18IFWSuoyQEluf8frfWf4cL2Y0JUQw= +github.com/aws/aws-sdk-go-v2/service/sso v1.36.0/go.mod h1:tWuiVBUtPBr8/rgRiYS8Uf85sHcAN+G7XS3D3CEoUh8= +github.com/aws/aws-sdk-go-v2/service/ssooidc v1.41.0 h1:wVxM3QzSKIK8tSN6OGgezp9OK91lCLH2zhmRInN9rFM= +github.com/aws/aws-sdk-go-v2/service/ssooidc v1.41.0/go.mod h1:naFe83jSMuYkH+QjQPX8n1MLhBkeCFM5Lsnh5m5wz3c= +github.com/aws/aws-sdk-go-v2/service/sts v1.48.0 h1:RzZVCzYM19vhJCT5s6vO2wN8ie770Li/TmbAZ9B6N7E= +github.com/aws/aws-sdk-go-v2/service/sts v1.48.0/go.mod h1:mKo/CzaCz8qytGW70NG4vIIGAx1HXTlb5lHNkC5k3lk= +github.com/aws/smithy-go v1.28.1 h1:R/nXH00c8qcfCzQVELtRw+eLQWtzv+VAIEFJ1/xxXlQ= +github.com/aws/smithy-go v1.28.1/go.mod h1:YE2RhdIuDbA5E5bTdciG9KrW3+TiEONeUWCqxX9i1Fc= github.com/beorn7/perks v0.0.0-20180321164747-3a771d992973/go.mod h1:Dwedo/Wpr24TaqPxmxbtue+5NUziq4I4S80YR8gNf3Q= github.com/beorn7/perks v1.0.0/go.mod h1:KWe93zE9D1o94FZ5RNwFwVgaQK1VOXiVxmqh+CedLV8= github.com/bgentry/speakeasy v0.1.0/go.mod h1:+zsyZBPWlz7T6j88CTgSN5bM796AkVf0kBD4zp0CCIs= @@ -237,10 +271,6 @@ github.com/ianlancetaylor/demangle v0.0.0-20200824232613-28f6c0f3b639/go.mod h1: github.com/inconshreveable/mousetrap v1.0.0/go.mod h1:PxqpIevigyE2G7u3NXJIT2ANytuPF1OarO4DADm73n8= github.com/jacobsa/fuse v0.0.0-20221016084658-a4cd154343d8 h1:uv+7DeBJF6+a54ZUihuXR+uTzhv2JslllK5ByILxxbg= github.com/jacobsa/fuse v0.0.0-20221016084658-a4cd154343d8/go.mod h1:liOmRdJd8oTwHCQ5M9JemRE3CebdlYcZWLk+ZjQeuq0= -github.com/jmespath/go-jmespath v0.4.0 h1:BEgLn5cpjn8UN1mAw4NjwDrS35OdebyEtFe+9YPoQUg= -github.com/jmespath/go-jmespath v0.4.0/go.mod h1:T8mJZnbsbmF+m6zOOFylbeCJqk5+pHWvzYPziyZiYoo= -github.com/jmespath/go-jmespath/internal/testify v1.5.1 h1:shLQSRRSCCPj3f2gpwzGwWFoC7ycTf1rcQZHOlsJ6N8= -github.com/jmespath/go-jmespath/internal/testify v1.5.1/go.mod h1:L3OGu8Wl2/fWfCI6z80xFu9LTZmf1ZRjMHUOPmWr69U= github.com/jonboulle/clockwork v0.1.0/go.mod h1:Ii8DK3G1RaLaWxj9trq07+26W01tbo22gdxWY5EU2bo= github.com/json-iterator/go v1.1.6/go.mod h1:+SdeFBvtyEkXs7REEP0seUULqWtbJapLOCVDaaPEHmU= github.com/jstemmer/go-junit-report v0.0.0-20190106144839-af01ea7f8024/go.mod h1:6v2b51hI/fHJwM22ozAgKL4VKDeJcHhJFhtBdhmNjmU= @@ -460,7 +490,6 @@ golang.org/x/net v0.0.0-20201209123823-ac852fbbde11/go.mod h1:m0MpNAwzfU5UDzcl9v golang.org/x/net v0.0.0-20201224014010-6772e930b67b/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20210119194325-5f4716e94777/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= -golang.org/x/net v0.0.0-20220127200216-cd36cc0744dd/go.mod h1:CfG3xpIq0wQ8r1q4Su4UZFWDARRcnwPjda9FqA0JpMk= golang.org/x/net v0.0.0-20220526153639-5463443f8c37 h1:lUkvobShwKsOesNfWWlCS5q7fnbG1MEliIzwu886fn8= golang.org/x/net v0.0.0-20220526153639-5463443f8c37/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U= @@ -536,15 +565,12 @@ golang.org/x/sys v0.0.0-20210225134936-a50acf3fe073/go.mod h1:h1NjWce9XRLGQEsW7w golang.org/x/sys v0.0.0-20210305230114-8fe3ee5dd75b/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20210320140829-1e4c9ba3b0c4/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20210403161142-5e06dd20ab57/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= -golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.0.0-20211216021012-1d35b9e2eb4e/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.8.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.11.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.28.0 h1:Fksou7UEQUWlKvIdsqzJmUmCX3cZuD2+P3XyyzwMhlA= golang.org/x/sys v0.28.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= golang.org/x/term v0.0.0-20201117132131-f5c789dd3221/go.mod h1:Nr5EML6q2oocZ2LXRh80K7BxOlk5/8JxuGnuhpl+muw= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= -golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8= golang.org/x/text v0.0.0-20170915032832-14c0d48ead0c/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.1-0.20180807135948-17ff2d5776d2/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= @@ -726,8 +752,6 @@ gopkg.in/yaml.v2 v2.0.0-20170812160011-eb3733d160e7/go.mod h1:JAlM8MvJe8wmxCU4Bl gopkg.in/yaml.v2 v2.2.1/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.2.4/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= -gopkg.in/yaml.v2 v2.2.8/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= -gopkg.in/yaml.v2 v2.4.0 h1:D8xgwECY7CYvx+Y2n4sBz93Jn9JRvxdiyyo8CTfuKaY= gopkg.in/yaml.v2 v2.4.0/go.mod h1:RDklbk79AGWmwhnvt/jBztapEOGDOx6ZbXqjP6csGnQ= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= diff --git a/internal/aws_error_test.go b/internal/aws_error_test.go new file mode 100644 index 00000000..057c8fa2 --- /dev/null +++ b/internal/aws_error_test.go @@ -0,0 +1,63 @@ +package internal + +import ( + "errors" + "fmt" + "net/http" + "syscall" + "testing" + + "github.com/aws/smithy-go" + smithyhttp "github.com/aws/smithy-go/transport/http" +) + +func TestMapAwsError(t *testing.T) { + for _, test := range []struct { + name string + code string + status int + want error + }{ + {"missing bucket", "NoSuchBucket", 404, syscall.ENXIO}, + {"owned bucket", "BucketAlreadyOwnedByYou", 409, syscall.EEXIST}, + {"invalid request", "InvalidArgument", 400, syscall.EINVAL}, + {"unauthorized", "Unauthorized", 401, syscall.EACCES}, + {"access denied", "AccessDenied", 403, syscall.EACCES}, + {"missing object", "NoSuchKey", 404, syscall.ENOENT}, + {"unsupported", "MethodNotAllowed", 405, syscall.ENOTSUP}, + {"conflict", "Conflict", 409, syscall.EINTR}, + {"throttled", "TooManyRequests", 429, syscall.EAGAIN}, + {"server error", "InternalError", 500, syscall.EAGAIN}, + } { + t.Run(test.name, func(t *testing.T) { + err := &smithy.OperationError{ + ServiceID: "S3", + OperationName: "HeadObject", + Err: &smithyhttp.ResponseError{ + Response: &smithyhttp.Response{Response: &http.Response{StatusCode: test.status}}, + Err: &smithy.GenericAPIError{Code: test.code, Message: test.name}, + }, + } + if got := mapAwsError(fmt.Errorf("wrapped: %w", err)); !errors.Is(got, test.want) { + t.Fatalf("mapAwsError() = %v, want %v", got, test.want) + } + }) + } +} + +func TestMapAwsErrorPreservesUnknownErrors(t *testing.T) { + for _, err := range []error{ + nil, + errors.New("transport failed"), + &smithy.GenericAPIError{Code: "PermanentRedirect"}, + &smithy.GenericAPIError{Code: "UnknownError"}, + &smithyhttp.ResponseError{ + Response: &smithyhttp.Response{Response: &http.Response{StatusCode: 503}}, + Err: errors.New("unavailable"), + }, + } { + if got := mapAwsError(err); got != err { + t.Fatalf("mapAwsError(%v) = %v", err, got) + } + } +} diff --git a/internal/aws_test.go b/internal/aws_test.go index fc259afe..5a7ae4c6 100644 --- a/internal/aws_test.go +++ b/internal/aws_test.go @@ -43,7 +43,7 @@ func (s *AwsTest) TestRegionDetection(t *C) { err, isAws := s.s3.detectBucketLocationByHEAD() t.Assert(err, IsNil) - t.Assert(*s.s3.awsConfig.Region, Equals, "eu-west-1") + t.Assert(s.s3.awsConfig.Region, Equals, "eu-west-1") t.Assert(isAws, Equals, true) } diff --git a/internal/backend_compat_test.go b/internal/backend_compat_test.go new file mode 100644 index 00000000..7bc87b4f --- /dev/null +++ b/internal/backend_compat_test.go @@ -0,0 +1,281 @@ +package internal + +import ( + "bytes" + "crypto/md5" + "encoding/base64" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "syscall" + "testing" + "time" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/aws/retry" + "github.com/aws/aws-sdk-go-v2/credentials" + . "github.com/kahing/goofys/api/common" +) + +func compatBackend(t *testing.T, handler http.HandlerFunc, configure func(*S3Config)) *S3Backend { + t.Helper() + config := (&S3Config{RegionSet: true, Credentials: credentials.NewStaticCredentialsProvider("access", "secret", "token")}).Init() + if configure != nil { + configure(config) + } + server := httptest.NewUnstartedServer(handler) + if config.SseC != "" { + server.StartTLS() + } else { + server.Start() + } + t.Cleanup(server.Close) + backend, err := NewS3("bucket", &FlagStorage{Endpoint: server.URL, HTTPTimeout: time.Second}, config) + if err != nil { + t.Fatal(err) + } + if config.SseC != "" { + backend.awsConfig.HTTPClient = server.Client() + } + retryer := backend.awsConfig.Retryer + backend.awsConfig.Retryer = func() aws.Retryer { + return retry.AddWithMaxBackoffDelay(retryer(), time.Nanosecond) + } + backend.newS3() + return backend +} + +func TestS3CompatibilityHeadersMetadataAndCopy(t *testing.T) { + key := base64.StdEncoding.EncodeToString([]byte("01234567890123456789012345678901")) + backend := compatBackend(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.EscapedPath() != "/bucket/a%20b%2Bc" { + t.Errorf("path = %s", r.URL.EscapedPath()) + } + if !strings.HasPrefix(r.Header.Get("Authorization"), "AWS4-HMAC-SHA256 ") { + t.Errorf("authorization = %q", r.Header.Get("Authorization")) + } + if !strings.Contains(r.Header.Get("User-Agent"), "goofys/") { + t.Errorf("user agent = %q", r.Header.Get("User-Agent")) + } + if r.Header.Get("x-amz-server-side-encryption-customer-key") != key { + t.Errorf("SSE-C key = %q", r.Header.Get("x-amz-server-side-encryption-customer-key")) + } + w.Header().Set("x-amz-request-id", "request") + w.Header().Set("x-amz-id-2", "host") + switch r.Method { + case "HEAD", "GET": + if r.Header.Get("x-amz-request-payer") != "requester" { + t.Error("missing requester pays") + } + w.Header().Set("ETag", "\"etag\"") + w.Header().Set("Content-Length", "3") + w.Header().Set("x-amz-meta-MiXeD", "value") + if r.Method == "GET" { + if r.Header.Get("Accept-Encoding") != "identity" || strings.Contains(r.Header.Get("Authorization"), "accept-encoding") { + t.Error("GET encoding must be identity and unsigned") + } + if r.Header.Get("Range") != "bytes=1-3" { + t.Errorf("range = %q", r.Header.Get("Range")) + } + io.WriteString(w, "abc") + } + case "PUT": + if r.Header.Get("x-amz-copy-source") != "bucket%2Fa+b%2Bc" { + t.Errorf("copy source = %q", r.Header.Get("x-amz-copy-source")) + } + if r.Header.Get("x-amz-storage-class") != "STANDARD" { + t.Errorf("storage class = %q", r.Header.Get("x-amz-storage-class")) + } + if r.Header.Get("x-amz-copy-source-server-side-encryption-customer-key") != key { + t.Error("missing copy SSE-C key") + } + io.WriteString(w, "\"etag\"") + } + }, func(c *S3Config) { c.SseC = key; c.RequesterPays = true; c.StorageClass = "STANDARD_IA" }) + head, err := backend.HeadBlob(&HeadBlobInput{Key: "a b+c"}) + if err != nil { + t.Fatal(err) + } + if head.StorageClass != nil || aws.ToString(head.Metadata["mixed"]) != "value" || head.RequestId != "request: host" { + t.Fatalf("head = %+v", head) + } + got, err := backend.GetBlob(&GetBlobInput{Key: "a b+c", Start: 1, Count: 3}) + if err != nil { + t.Fatal(err) + } + body, err := io.ReadAll(got.Body) + got.Body.Close() + if err != nil || string(body) != "abc" { + t.Fatalf("body = %q, %v", body, err) + } + if _, err := backend.CopyBlob(&CopyBlobInput{Source: "a b+c", Destination: "a b+c"}); err != nil { + t.Fatal(err) + } +} + +func TestS3CompatibilityRetriesAndMetadata(t *testing.T) { + var attempts atomic.Int32 + backend := compatBackend(t, func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + if string(body) != "payload" { + t.Errorf("body = %q", body) + } + digest := md5.Sum(body) + if r.Header.Get("Content-MD5") != base64.StdEncoding.EncodeToString(digest[:]) { + t.Errorf("upload MD5 = %q", r.Header.Get("Content-MD5")) + } + if r.Header.Get("x-amz-meta-mixed") != "value" || r.Header.Get("x-amz-checksum-crc32") != "" { + t.Errorf("headers = %v", r.Header) + } + if _, ok := r.Header["X-Amz-Meta-Absent"]; ok { + t.Error("nil metadata was serialized") + } + if attempts.Add(1) < 4 { + w.WriteHeader(500) + io.WriteString(w, "InternalError") + return + } + w.Header().Set("ETag", "etag") + w.Header().Set("Date", "Fri, 11 Sep 2026 00:00:00 GMT") + }, nil) + result, err := backend.PutBlob(&PutBlobInput{Key: "key", Body: bytes.NewReader([]byte("payload")), Metadata: map[string]*string{"MiXeD": aws.String("value"), "absent": nil}}) + if err != nil || attempts.Load() != 4 { + t.Fatalf("attempts=%d error=%v", attempts.Load(), err) + } + if result.LastModified == nil || aws.ToString(result.ETag) != "etag" { + t.Fatalf("result = %+v", result) + } +} + +func TestS3CompatibilityDeleteMD5AndFallback(t *testing.T) { + var attempts atomic.Int32 + backend := compatBackend(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method == "HEAD" { + attempts.Add(1) + if !strings.HasPrefix(r.Header.Get("Authorization"), "AWS access:") { + w.WriteHeader(403) + return + } + w.WriteHeader(404) + return + } + body, _ := io.ReadAll(r.Body) + digest := md5.Sum(body) + if r.Header.Get("Content-MD5") != base64.StdEncoding.EncodeToString(digest[:]) { + t.Errorf("MD5 = %q", r.Header.Get("Content-MD5")) + } + if r.Header.Get("x-amz-checksum-crc32") != "" { + t.Error("unexpected CRC32") + } + if !strings.HasPrefix(r.Header.Get("Authorization"), "AWS access:") { + t.Error("missing v2 signature") + } + io.WriteString(w, "") + }, nil) + if err := backend.Init("missing"); err != nil { + t.Fatal(err) + } + if !backend.v2Signer || attempts.Load() != 2 { + t.Fatalf("fallback=%v attempts=%d", backend.v2Signer, attempts.Load()) + } + if _, err := backend.DeleteBlobs(&DeleteBlobsInput{Items: []string{"a", "b"}}); err != nil { + t.Fatal(err) + } +} + +func TestS3CompatibilitySSERequiresTLS(t *testing.T) { + var calls atomic.Int32 + backend := compatBackend(t, func(w http.ResponseWriter, r *http.Request) { + calls.Add(1) + }, nil) + backend.config.SseC = "01234567890123456789012345678901" + for _, operation := range []func() error{ + func() error { + _, err := backend.HeadBlob(&HeadBlobInput{Key: "key"}) + return err + }, + func() error { + _, err := backend.PutBlob(&PutBlobInput{Key: "key", Body: strings.NewReader("payload")}) + return err + }, + func() error { + size := uint64(3) + _, err := backend.CopyBlob(&CopyBlobInput{Source: "source", Destination: "destination", Size: &size, ETag: aws.String("etag")}) + return err + }, + } { + if err := operation(); err == nil || !strings.Contains(err.Error(), "cannot send SSE keys over HTTP.") { + t.Fatalf("error = %v", err) + } + } + if calls.Load() != 0 { + t.Fatalf("sent %d unencrypted requests", calls.Load()) + } +} + +func TestS3CompatibilityAnonymousAndErrors(t *testing.T) { + backend := compatBackend(t, func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Authorization") != "" { + t.Error("anonymous request was signed") + } + w.WriteHeader(403) + }, func(c *S3Config) { c.Credentials = aws.AnonymousCredentials{} }) + if _, err := backend.HeadBlob(&HeadBlobInput{Key: "key"}); err != syscall.EACCES { + t.Fatalf("error = %v", err) + } +} + +func TestGCSCompatibilityResumable(t *testing.T) { + var starts, parts, retries atomic.Int32 + backend := compatBackend(t, func(w http.ResponseWriter, r *http.Request) { + if r.Method == "POST" { + starts.Add(1) + if r.URL.RawQuery != "" || r.Header.Get("x-goog-resumable") != "start" || !strings.HasPrefix(r.Header.Get("Authorization"), "AWS access:") { + t.Errorf("start URL=%s headers=%v", r.URL, r.Header) + } + w.Header().Set("Location", "http://"+r.Host+"/resumable?upload_id=token") + w.WriteHeader(201) + return + } + if r.Header.Get("Content-Range") == "bytes 0-262143/*" && retries.Add(1) == 1 { + io.Copy(io.Discard, r.Body) + w.WriteHeader(500) + io.WriteString(w, "InternalError") + return + } + part := parts.Add(1) + if r.URL.Path != "/resumable" || r.URL.Query().Get("upload_id") != "token" || r.Header.Get("Authorization") != "" { + t.Errorf("part URL=%s headers=%v", r.URL, r.Header) + } + body, _ := io.ReadAll(r.Body) + if part == 1 { + if len(body) != 256*1024 || r.Header.Get("Content-Range") != "bytes 0-262143/*" { + t.Errorf("first part len=%d range=%s", len(body), r.Header.Get("Content-Range")) + } + w.WriteHeader(308) + } else { + if string(body) != "last" || r.Header.Get("Content-Range") != "bytes 262144-262147/262148" { + t.Errorf("last part body=%q range=%s", body, r.Header.Get("Content-Range")) + } + w.Header().Set("ETag", "final") + } + }, nil) + gcs := &GCS3{S3Backend: backend} + backend.gcs = true + commit, err := gcs.MultipartBlobBegin(&MultipartBlobBeginInput{Key: "key", ContentType: aws.String("text/plain")}) + if err != nil { + t.Fatal(err) + } + for i, body := range [][]byte{make([]byte, 256*1024), []byte("last")} { + _, err = gcs.MultipartBlobAdd(&MultipartBlobAddInput{Commit: commit, PartNumber: uint32(i + 1), Body: bytes.NewReader(body), Size: uint64(len(body))}) + if err != nil { + t.Fatal(err) + } + } + result, err := gcs.MultipartBlobCommit(commit) + if err != nil || result == nil || aws.ToString(result.ETag) != "final" || starts.Load() != 1 || parts.Load() != 2 { + t.Fatalf("result=%+v error=%v starts=%d parts=%d", result, err, starts.Load(), parts.Load()) + } +} diff --git a/internal/backend_gcs3.go b/internal/backend_gcs3.go index e3b4dc72..6ab8159b 100644 --- a/internal/backend_gcs3.go +++ b/internal/backend_gcs3.go @@ -17,14 +17,20 @@ package internal import ( . "github.com/kahing/goofys/api/common" + "context" "fmt" "io" + "net/http" "net/url" "strconv" "sync" "sync/atomic" - "github.com/aws/aws-sdk-go/service/s3" + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/s3" + "github.com/aws/aws-sdk-go-v2/service/s3/types" + "github.com/aws/smithy-go/middleware" + smithyhttp "github.com/aws/smithy-go/transport/http" "github.com/jacobsa/fuse" ) @@ -32,6 +38,26 @@ import ( // GCS variant of S3 type GCS3 struct { *S3Backend + + resumableMu sync.Mutex + resumableBase *s3.Client + resumableStart *s3.Client + resumableUpload *s3.Client +} + +func (s *GCS3) resumableClients() (start, upload *s3.Client) { + s.resumableMu.Lock() + defer s.resumableMu.Unlock() + if s.resumableBase != s.Client { + options := s.Client.Options() + options.AuthSchemes = nil + s.resumableStart = s3.New(options, V2Signer(s.bucket)) + options = s.Client.Options() + options.Credentials = aws.AnonymousCredentials{} + s.resumableUpload = s3.New(options) + s.resumableBase = s.Client + } + return s.resumableStart, s.resumableUpload } type GCS3MultipartBlobCommitInput struct { @@ -81,39 +107,67 @@ func (s *GCS3) DeleteBlobs(param *DeleteBlobsInput) (*DeleteBlobsOutput, error) return &DeleteBlobsOutput{}, nil } +func gcsRequest(update func(*smithyhttp.Request) error, response **http.Response, upload bool) func(*s3.Options) { + return func(o *s3.Options) { + o.APIOptions = append(o.APIOptions, func(stack *middleware.Stack) error { + if err := stack.Finalize.Insert(middleware.FinalizeMiddlewareFunc("GCSResumableRequest", func(ctx context.Context, in middleware.FinalizeInput, next middleware.FinalizeHandler) (middleware.FinalizeOutput, middleware.Metadata, error) { + if err := update(in.Request.(*smithyhttp.Request)); err != nil { + return middleware.FinalizeOutput{}, middleware.Metadata{}, err + } + return next.HandleFinalize(ctx, in) + }), "Signing", middleware.Before); err != nil { + return err + } + return stack.Deserialize.Add(middleware.DeserializeMiddlewareFunc("GCSResumableResponse", func(ctx context.Context, in middleware.DeserializeInput, next middleware.DeserializeHandler) (middleware.DeserializeOutput, middleware.Metadata, error) { + out, metadata, err := next.HandleDeserialize(ctx, in) + if resp, ok := out.RawResponse.(*smithyhttp.Response); ok { + *response = resp.Response + if upload && resp.StatusCode == 308 { + out.Result = &s3.PutObjectOutput{ETag: aws.String(resp.Header.Get("ETag"))} + err = nil + } else if !upload && resp.StatusCode >= 200 && resp.StatusCode < 300 { + out.Result = &s3.CreateMultipartUploadOutput{} + err = nil + } + } + return out, metadata, err + }), middleware.Before) + }) + } +} + func (s *GCS3) MultipartBlobBegin(param *MultipartBlobBeginInput) (*MultipartBlobCommitInput, error) { mpu := s3.CreateMultipartUploadInput{ Bucket: &s.bucket, Key: ¶m.Key, - StorageClass: &s.config.StorageClass, + StorageClass: types.StorageClass(s.config.StorageClass), ContentType: param.ContentType, } if s.config.UseSSE { - mpu.ServerSideEncryption = &s.sseType + mpu.ServerSideEncryption = s.sseType if s.config.UseKMS && s.config.KMSKeyID != "" { mpu.SSEKMSKeyId = &s.config.KMSKeyID } } if s.config.ACL != "" { - mpu.ACL = &s.config.ACL + mpu.ACL = types.ObjectCannedACL(s.config.ACL) } - req, _ := s.CreateMultipartUploadRequest(&mpu) - // v4 signing of this fails - s.setV2Signer(&req.Handlers) - // get rid of ?uploads= - req.HTTPRequest.URL.RawQuery = "" - req.HTTPRequest.Header.Set("x-goog-resumable", "start") - - err := req.Send() + var response *http.Response + client, _ := s.resumableClients() + _, err := client.CreateMultipartUpload(context.TODO(), &mpu, gcsRequest(func(req *smithyhttp.Request) error { + req.URL.RawQuery = "" + req.Header.Set("x-goog-resumable", "start") + return nil + }, &response, false)) if err != nil { s3Log.Errorf("CreateMultipartUpload %v = %v", param.Key, err) return nil, mapAwsError(err) } - location := req.HTTPResponse.Header.Get("Location") + location := response.Header.Get("Location") _, err = url.Parse(location) if err != nil { s3Log.Errorf("CreateMultipartUpload %v %v = %v", param.Key, location, err) @@ -148,9 +202,10 @@ func (s *GCS3) uploadPart(param *MultipartBlobAddInput, totalSize uint64, last b s3Log.Debug(params) - req, resp := s.PutObjectRequest(params) - req.Handlers.Sign.Clear() - req.HTTPRequest.URL, _ = url.Parse(*param.Commit.UploadId) + location, err := url.Parse(*param.Commit.UploadId) + if err != nil { + return nil, err + } start := totalSize - param.Size end := totalSize - 1 @@ -163,18 +218,16 @@ func (s *GCS3) uploadPart(param *MultipartBlobAddInput, totalSize uint64, last b contentRange := fmt.Sprintf("bytes %v-%v/%v", start, end, size) - req.HTTPRequest.Header.Set("Content-Length", strconv.FormatUint(param.Size, 10)) - req.HTTPRequest.Header.Set("Content-Range", contentRange) - - err = req.Send() + params.ContentLength = aws.Int64(int64(param.Size)) + var response *http.Response + _, client := s.resumableClients() + resp, err := client.PutObject(context.TODO(), params, gcsRequest(func(req *smithyhttp.Request) error { + req.URL = location + req.Header.Set("Content-Range", contentRange) + return nil + }, &response, true)) if err != nil { - // status indicating that we need more parts to finish this - if req.HTTPResponse.StatusCode == 308 { - err = nil - } else { - err = mapAwsError(err) - return - } + return nil, mapAwsError(err) } etag = resp.ETag diff --git a/internal/backend_operations_compat_test.go b/internal/backend_operations_compat_test.go new file mode 100644 index 00000000..d958e2f5 --- /dev/null +++ b/internal/backend_operations_compat_test.go @@ -0,0 +1,173 @@ +package internal + +import ( + "bytes" + "crypto/md5" + "encoding/base64" + "io" + "net/http" + "strings" + "sync/atomic" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/s3" + "github.com/aws/aws-sdk-go-v2/service/s3/types" + . "github.com/kahing/goofys/api/common" +) + +func TestS3CompatibilityListVersions(t *testing.T) { + for _, isAWS := range []bool{false, true} { + t.Run(map[bool]string{false: "custom", true: "aws"}[isAWS], func(t *testing.T) { + backend := compatBackend(t, func(w http.ResponseWriter, r *http.Request) { + q := r.URL.Query() + if q.Get("prefix") != "dir/" || q.Get("delimiter") != "/" || q.Get("max-keys") != "17" { + t.Errorf("query = %v", q) + } + if isAWS { + if q.Get("list-type") != "2" || q.Get("continuation-token") != "next" || q.Get("start-after") != "after" { + t.Errorf("v2 query = %v", q) + } + } else if q.Get("list-type") != "" || q.Get("marker") != "after" { + t.Errorf("v1 query = %v", q) + } + io.WriteString(w, `truemarkertokendir/key3STANDARDdir/sub/`) + }, nil) + backend.aws = isAWS + maxKeys := uint32(17) + got, err := backend.ListBlobs(&ListBlobsInput{Prefix: aws.String("dir/"), Delimiter: aws.String("/"), MaxKeys: &maxKeys, StartAfter: aws.String("after"), ContinuationToken: aws.String("next")}) + if err != nil { + t.Fatal(err) + } + wantToken := "marker" + if isAWS { + wantToken = "token" + } + if !got.IsTruncated || aws.ToString(got.NextContinuationToken) != wantToken || len(got.Items) != 1 || got.Items[0].Size != 3 || len(got.Prefixes) != 1 { + t.Fatalf("list = %+v", got) + } + }) + } +} + +func TestS3CompatibilityMultipart(t *testing.T) { + key := base64.StdEncoding.EncodeToString([]byte("01234567890123456789012345678901")) + var calls atomic.Int32 + backend := compatBackend(t, func(w http.ResponseWriter, r *http.Request) { + calls.Add(1) + if r.Method == "POST" && r.Header.Get("x-amz-request-payer") != "requester" { + t.Error("missing request payer") + } + w.Header().Set("x-amz-request-id", "request") + w.Header().Set("x-amz-id-2", "host") + switch { + case r.URL.Query().Has("uploads"): + if r.Header.Get("x-amz-server-side-encryption-customer-key") != key || r.Header.Get("x-amz-acl") != "private" { + t.Errorf("begin headers = %v", r.Header) + } + io.WriteString(w, `upload+id/=`) + case r.Method == "PUT": + if r.URL.Query().Get("uploadId") != "upload+id/=" || r.URL.Query().Get("partNumber") != "1" || r.Header.Get("x-amz-server-side-encryption-customer-key") != key { + t.Errorf("part URL=%s headers=%v", r.URL, r.Header) + } + body, _ := io.ReadAll(r.Body) + if string(body) != "part" { + t.Errorf("part body = %q", body) + } + w.Header().Set("ETag", "\"part-etag\"") + case r.Method == "POST": + body, _ := io.ReadAll(r.Body) + if !bytes.Contains(body, []byte("part-etag")) || !bytes.Contains(body, []byte("1")) { + t.Errorf("complete body = %q", body) + } + io.WriteString(w, `"final"`) + case r.Method == "DELETE": + w.WriteHeader(204) + } + }, func(c *S3Config) { c.SseC = key; c.RequesterPays = true; c.ACL = "private" }) + commit, err := backend.MultipartBlobBegin(&MultipartBlobBeginInput{Key: "key"}) + if err != nil { + t.Fatal(err) + } + if _, err := backend.MultipartBlobAdd(&MultipartBlobAddInput{Commit: commit, PartNumber: 1, Size: 4, Body: strings.NewReader("part")}); err != nil { + t.Fatal(err) + } + result, err := backend.MultipartBlobCommit(commit) + if err != nil || result == nil || aws.ToString(result.ETag) != `"final"` || result.RequestId != "request: host" || result.LastModified == nil { + t.Fatalf("complete=%+v error=%v", result, err) + } + if _, err := backend.MultipartBlobAbort(commit); err != nil { + t.Fatal(err) + } + if calls.Load() != 4 { + t.Fatalf("calls = %d", calls.Load()) + } +} + +func TestS3CompatibilityCreateBucketRegion(t *testing.T) { + for _, region := range []string{"us-east-1", "eu-west-1"} { + t.Run(region, func(t *testing.T) { + backend := compatBackend(t, func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + if region == "us-east-1" && len(body) != 0 { + t.Errorf("unexpected create body = %q", body) + } + if region != "us-east-1" && !strings.Contains(string(body), ""+region+"") { + t.Errorf("create body = %q", body) + } + }, func(c *S3Config) { c.Region = region }) + if _, err := backend.MakeBucket(&MakeBucketInput{}); err != nil { + t.Fatal(err) + } + }) + } +} + +func TestS3CompatibilityBucketTaggingMD5(t *testing.T) { + backend := compatBackend(t, func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + digest := md5.Sum(body) + if r.Header.Get("Content-MD5") != base64.StdEncoding.EncodeToString(digest[:]) || r.Header.Get("x-amz-checksum-crc32") != "" { + t.Errorf("headers = %v", r.Header) + } + w.WriteHeader(204) + }, nil) + _, err := backend.PutBucketTagging(t.Context(), &s3.PutBucketTaggingInput{Bucket: aws.String("bucket"), Tagging: &types.Tagging{TagSet: []types.Tag{{Key: aws.String("Owner"), Value: aws.String("owner")}}}}) + if err != nil { + t.Fatal(err) + } +} + +func TestS3CompatibilityBucketDetection(t *testing.T) { + for _, anonymous := range []bool{false, true} { + t.Run(map[bool]string{false: "region", true: "anonymous"}[anonymous], func(t *testing.T) { + var calls atomic.Int32 + backend := compatBackend(t, func(w http.ResponseWriter, r *http.Request) { + if calls.Add(1) == 1 { + if r.URL.Path != "/bucket" || r.Header.Get("Authorization") != "" { + t.Errorf("discovery URL=%s headers=%v", r.URL, r.Header) + } + w.Header().Set("Server", "AmazonS3") + w.Header().Set("x-amz-bucket-region", "us-west-2") + if !anonymous { + w.WriteHeader(403) + } + return + } + if anonymous && r.Header.Get("Authorization") != "" { + t.Error("anonymous request was signed") + } + if !anonymous && !strings.Contains(r.Header.Get("Authorization"), "/us-west-2/s3/aws4_request") { + t.Errorf("authorization = %q", r.Header.Get("Authorization")) + } + w.WriteHeader(404) + }, func(c *S3Config) { c.RegionSet = false }) + if err := backend.Init("missing"); err != nil { + t.Fatal(err) + } + if backend.awsConfig.Region != "us-west-2" || !backend.aws || backend.v2Signer { + t.Fatalf("region=%s aws=%v v2=%v", backend.awsConfig.Region, backend.aws, backend.v2Signer) + } + }) + } +} diff --git a/internal/backend_s3.go b/internal/backend_s3.go index 5bb4c3ba..5ba6c270 100644 --- a/internal/backend_s3.go +++ b/internal/backend_s3.go @@ -17,7 +17,11 @@ package internal import ( . "github.com/kahing/goofys/api/common" + "context" + "crypto/md5" + "encoding/base64" "fmt" + "io" "net/http" "net/url" "strconv" @@ -26,25 +30,26 @@ import ( "syscall" "time" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/aws/awserr" - "github.com/aws/aws-sdk-go/aws/corehandlers" - "github.com/aws/aws-sdk-go/aws/credentials" - "github.com/aws/aws-sdk-go/aws/request" - "github.com/aws/aws-sdk-go/service/s3" + "github.com/aws/aws-sdk-go-v2/aws" + awsmiddleware "github.com/aws/aws-sdk-go-v2/aws/middleware" + "github.com/aws/aws-sdk-go-v2/service/s3" + "github.com/aws/aws-sdk-go-v2/service/s3/types" + "github.com/aws/smithy-go" + "github.com/aws/smithy-go/middleware" + smithyhttp "github.com/aws/smithy-go/transport/http" "github.com/jacobsa/fuse" ) type S3Backend struct { - *s3.S3 + *s3.Client cap Capabilities bucket string awsConfig *aws.Config flags *FlagStorage config *S3Config - sseType string + sseType types.ServerSideEncryption aws bool gcs bool @@ -68,15 +73,15 @@ func NewS3(bucket string, flags *FlagStorage, config *S3Config) (*S3Backend, err } if flags.DebugS3 { - awsConfig.LogLevel = aws.LogLevel(aws.LogDebug | aws.LogDebugWithRequestErrors) + awsConfig.ClientLogMode = aws.LogRequest | aws.LogResponse | aws.LogRetries } if config.UseKMS { //SSE header string for KMS server-side encryption (SSE-KMS) - s.sseType = s3.ServerSideEncryptionAwsKms + s.sseType = types.ServerSideEncryptionAwsKms } else if config.UseSSE { //SSE header string for non-KMS server-side encryption (SSE-S3) - s.sseType = s3.ServerSideEncryptionAes256 + s.sseType = types.ServerSideEncryptionAes256 } s.newS3() @@ -91,43 +96,119 @@ func (s *S3Backend) Capabilities() *Capabilities { return &s.cap } -func addAcceptEncoding(req *request.Request) { - if req.HTTPRequest.Method == "GET" { - // we need "Accept-Encoding: identity" so that objects - // with content-encoding won't be automatically - // deflated, but we don't want to sign it because GCS - // doesn't like it - req.HTTPRequest.Header.Set("Accept-Encoding", "identity") - } +const checksumSetupMiddlewareID = "AWSChecksum:SetupInputContext" + +func (s *S3Backend) newS3() { + s.Client = s3.NewFromConfig(*s.awsConfig, func(o *s3.Options) { + o.UsePathStyle = !s.config.Subdomain + if s.flags.Endpoint != "" { + o.BaseEndpoint = &s.flags.Endpoint + } + o.RequestChecksumCalculation = aws.RequestChecksumCalculationWhenRequired + o.ResponseChecksumValidation = aws.ResponseChecksumValidationWhenRequired + if s.v2Signer { + V2Signer(s.bucket)(o) + } + o.APIOptions = append(o.APIOptions, func(stack *middleware.Stack) error { + if stack.ID() == "DeleteObjects" || stack.ID() == "PutBucketTagging" { + if _, ok := stack.Initialize.Get(checksumSetupMiddlewareID); ok { + if _, err := stack.Initialize.Remove(checksumSetupMiddlewareID); err != nil { + return err + } + } + } + if stack.ID() == "DeleteObjects" || stack.ID() == "PutBucketTagging" || stack.ID() == "PutObject" || stack.ID() == "UploadPart" { + if err := stack.Build.Add(middleware.BuildMiddlewareFunc("GoofysContentMD5", func(ctx context.Context, in middleware.BuildInput, next middleware.BuildHandler) (middleware.BuildOutput, middleware.Metadata, error) { + req := in.Request.(*smithyhttp.Request) + if req.Header.Get("Content-MD5") != "" || !req.IsStreamSeekable() { + return next.HandleBuild(ctx, in) + } + hash := md5.New() + if req.GetStream() != nil { + if _, err := io.Copy(hash, req.GetStream()); err != nil { + return middleware.BuildOutput{}, middleware.Metadata{}, err + } + if err := req.RewindStream(); err != nil { + return middleware.BuildOutput{}, middleware.Metadata{}, err + } + } + req.Header.Set("Content-MD5", base64.StdEncoding.EncodeToString(hash.Sum(nil))) + return next.HandleBuild(ctx, in) + }), middleware.After); err != nil { + return err + } + } + if err := stack.Build.Add(middleware.BuildMiddlewareFunc("GoofysHeaders", func(ctx context.Context, in middleware.BuildInput, next middleware.BuildHandler) (middleware.BuildOutput, middleware.Metadata, error) { + req := in.Request.(*smithyhttp.Request) + req.Header.Set("User-Agent", req.Header.Get("User-Agent")+" goofys/"+VersionNumber+"-"+VersionHash) + if s.config.RequesterPays && (req.Method == "GET" || req.Method == "HEAD" || req.Method == "POST") { + req.Header.Set("x-amz-request-payer", "requester") + } + return next.HandleBuild(ctx, in) + }), middleware.After); err != nil { + return err + } + if err := stack.Finalize.Insert(middleware.FinalizeMiddlewareFunc("GoofysUnsignedEncoding", func(ctx context.Context, in middleware.FinalizeInput, next middleware.FinalizeHandler) (middleware.FinalizeOutput, middleware.Metadata, error) { + req := in.Request.(*smithyhttp.Request) + if req.URL.Scheme != "https" && (req.Header.Get("x-amz-server-side-encryption-customer-key") != "" || req.Header.Get("x-amz-copy-source-server-side-encryption-customer-key") != "") { + return middleware.FinalizeOutput{}, middleware.Metadata{}, &smithy.GenericAPIError{Code: "ConfigError", Message: "cannot send SSE keys over HTTP."} + } + req.Header.Del("Accept-Encoding") + return next.HandleFinalize(ctx, in) + }), "Signing", middleware.Before); err != nil { + return err + } + return stack.Finalize.Add(middleware.FinalizeMiddlewareFunc("GoofysAcceptEncoding", func(ctx context.Context, in middleware.FinalizeInput, next middleware.FinalizeHandler) (middleware.FinalizeOutput, middleware.Metadata, error) { + req := in.Request.(*smithyhttp.Request) + if req.Method == "GET" { + req.Header.Set("Accept-Encoding", "identity") + } + return next.HandleFinalize(ctx, in) + }), middleware.After) + }) + }) +} + +func (s *S3Backend) customerKey() *string { + return aws.String(base64.StdEncoding.EncodeToString([]byte(s.config.SseC))) } -func addRequestPayer(req *request.Request) { - // "Requester Pays" is only applicable to these - // see https://docs.aws.amazon.com/AmazonS3/latest/dev/RequesterPaysBuckets.html - if req.HTTPRequest.Method == "GET" || req.HTTPRequest.Method == "HEAD" || req.HTTPRequest.Method == "POST" { - req.HTTPRequest.Header.Set("x-amz-request-payer", "requester") +func metadataToAWS(m map[string]*string) map[string]string { + if m == nil { + return nil } + out := make(map[string]string, len(m)) + for k, v := range m { + if v != nil { + out[strings.ToLower(k)] = *v + } + } + return out } -func (s *S3Backend) setV2Signer(handlers *request.Handlers) { - handlers.Sign.Clear() - handlers.Sign.PushBack(SignV2) - handlers.Sign.PushBackNamed(corehandlers.BuildContentLengthHandler) +func metadataFromAWS(m map[string]string) map[string]*string { + if m == nil { + return nil + } + out := make(map[string]*string, len(m)) + for k, v := range m { + out[strings.ToLower(k)] = aws.String(v) + } + return out } -func (s *S3Backend) newS3() { - s.S3 = s3.New(s.config.Session, s.awsConfig) - if s.config.RequesterPays { - s.S3.Handlers.Build.PushBack(addRequestPayer) +func optionalString(value string) *string { + if value == "" { + return nil } - if s.v2Signer { - s.setV2Signer(&s.S3.Handlers) + return &value +} + +func responseHTTP(metadata middleware.Metadata) *http.Response { + if resp, ok := awsmiddleware.GetRawResponse(metadata).(*smithyhttp.Response); ok { + return resp.Response } - s.S3.Handlers.Sign.PushBack(addAcceptEncoding) - s.S3.Handlers.Build.PushFrontNamed(request.NamedHandler{ - Name: "UserAgentHandler", - Fn: request.MakeAddToUserAgentHandler("goofys", VersionNumber+"-"+VersionHash), - }) + return nil } func (s *S3Backend) detectBucketLocationByHEAD() (err error, isAws bool) { @@ -137,8 +218,8 @@ func (s *S3Backend) detectBucketLocationByHEAD() (err error, isAws bool) { Path: s.bucket, } - if s.awsConfig.Endpoint != nil { - endpoint, err := url.Parse(*s.awsConfig.Endpoint) + if s.flags.Endpoint != "" { + endpoint, err := url.Parse(s.flags.Endpoint) if err != nil { return err, false } @@ -187,7 +268,7 @@ func (s *S3Backend) detectBucketLocationByHEAD() (err error, isAws bool) { case 200: // note that this only happen if the bucket is in us-east-1 if len(s.config.Profile) == 0 { - s.awsConfig.Credentials = credentials.AnonymousCredentials + s.awsConfig.Credentials = aws.AnonymousCredentials{} s3Log.Infof("anonymous bucket detected") } case 400: @@ -199,14 +280,14 @@ func (s *S3Backend) detectBucketLocationByHEAD() (err error, isAws bool) { case 405: err = syscall.ENOTSUP default: - err = awserr.New(strconv.Itoa(resp.StatusCode), resp.Status, nil) + err = &smithy.GenericAPIError{Code: strconv.Itoa(resp.StatusCode), Message: resp.Status} } if len(region) != 0 { - if region[0] != *s.awsConfig.Region { + if region[0] != s.awsConfig.Region { s3Log.Infof("Switching from region '%v' to '%v'", - *s.awsConfig.Region, region[0]) - s.awsConfig.Region = ®ion[0] + s.awsConfig.Region, region[0]) + s.awsConfig.Region = region[0] } // we detected a region, this is aws, the error is irrelevant @@ -286,12 +367,11 @@ func (s *S3Backend) Init(key string) error { func (s *S3Backend) ListObjectsV2(params *s3.ListObjectsV2Input) (*s3.ListObjectsV2Output, string, error) { if s.aws { - req, resp := s.S3.ListObjectsV2Request(params) - err := req.Send() + resp, err := s.Client.ListObjectsV2(context.TODO(), params) if err != nil { return nil, "", err } - return resp, s.getRequestId(req), nil + return resp, s.getRequestId(resp.ResultMetadata), nil } else { v1 := s3.ListObjectsInput{ Bucket: params.Bucket, @@ -307,12 +387,12 @@ func (s *S3Backend) ListObjectsV2(params *s3.ListObjectsV2Input) (*s3.ListObject v1.Marker = params.ContinuationToken } - objs, err := s.S3.ListObjects(&v1) + objs, err := s.Client.ListObjects(context.TODO(), &v1) if err != nil { return nil, "", err } - count := int64(len(objs.Contents)) + count := int32(len(objs.Contents)) v2Objs := s3.ListObjectsV2Output{ CommonPrefixes: objs.CommonPrefixes, Contents: objs.Contents, @@ -349,9 +429,12 @@ func metadataToLower(m map[string]*string) map[string]*string { return m } -func (s *S3Backend) getRequestId(r *request.Request) string { - return r.HTTPResponse.Header.Get("x-amz-request-id") + ": " + - r.HTTPResponse.Header.Get("x-amz-id-2") +func (s *S3Backend) getRequestId(metadata middleware.Metadata) string { + r := responseHTTP(metadata) + if r == nil { + return "" + } + return r.Header.Get("x-amz-request-id") + ": " + r.Header.Get("x-amz-id-2") } func (s *S3Backend) HeadBlob(param *HeadBlobInput) (*HeadBlobOutput, error) { @@ -360,12 +443,11 @@ func (s *S3Backend) HeadBlob(param *HeadBlobInput) (*HeadBlobOutput, error) { } if s.config.SseC != "" { head.SSECustomerAlgorithm = PString("AES256") - head.SSECustomerKey = &s.config.SseC + head.SSECustomerKey = s.customerKey() head.SSECustomerKeyMD5 = &s.config.SseCDigest } - req, resp := s.S3.HeadObjectRequest(&head) - err := req.Send() + resp, err := s.Client.HeadObject(context.TODO(), &head) if err != nil { return nil, mapAwsError(err) } @@ -375,20 +457,20 @@ func (s *S3Backend) HeadBlob(param *HeadBlobInput) (*HeadBlobOutput, error) { ETag: resp.ETag, LastModified: resp.LastModified, Size: uint64(*resp.ContentLength), - StorageClass: resp.StorageClass, + StorageClass: optionalString(string(resp.StorageClass)), }, ContentType: resp.ContentType, - Metadata: metadataToLower(resp.Metadata), + Metadata: metadataFromAWS(resp.Metadata), IsDirBlob: strings.HasSuffix(param.Key, "/"), - RequestId: s.getRequestId(req), + RequestId: s.getRequestId(resp.ResultMetadata), }, nil } func (s *S3Backend) ListBlobs(param *ListBlobsInput) (*ListBlobsOutput, error) { - var maxKeys *int64 + var maxKeys *int32 if param.MaxKeys != nil { - maxKeys = aws.Int64(int64(*param.MaxKeys)) + maxKeys = aws.Int32(int32(*param.MaxKeys)) } resp, reqId, err := s.ListObjectsV2(&s3.ListObjectsV2Input{ @@ -415,7 +497,7 @@ func (s *S3Backend) ListBlobs(param *ListBlobsInput) (*ListBlobsOutput, error) { ETag: i.ETag, LastModified: i.LastModified, Size: uint64(*i.Size), - StorageClass: i.StorageClass, + StorageClass: optionalString(string(i.StorageClass)), }) } @@ -423,46 +505,43 @@ func (s *S3Backend) ListBlobs(param *ListBlobsInput) (*ListBlobsOutput, error) { Prefixes: prefixes, Items: items, NextContinuationToken: resp.NextContinuationToken, - IsTruncated: *resp.IsTruncated, + IsTruncated: aws.ToBool(resp.IsTruncated), RequestId: reqId, }, nil } func (s *S3Backend) DeleteBlob(param *DeleteBlobInput) (*DeleteBlobOutput, error) { - req, _ := s.DeleteObjectRequest(&s3.DeleteObjectInput{ + resp, err := s.DeleteObject(context.TODO(), &s3.DeleteObjectInput{ Bucket: &s.bucket, Key: ¶m.Key, }) - err := req.Send() if err != nil { return nil, mapAwsError(err) } - return &DeleteBlobOutput{s.getRequestId(req)}, nil + return &DeleteBlobOutput{s.getRequestId(resp.ResultMetadata)}, nil } func (s *S3Backend) DeleteBlobs(param *DeleteBlobsInput) (*DeleteBlobsOutput, error) { num_objs := len(param.Items) - var items s3.Delete - var objs = make([]*s3.ObjectIdentifier, num_objs) + var items types.Delete + var objs = make([]types.ObjectIdentifier, num_objs) - for i, _ := range param.Items { - objs[i] = &s3.ObjectIdentifier{Key: ¶m.Items[i]} + for i := range param.Items { + objs[i] = types.ObjectIdentifier{Key: ¶m.Items[i]} } - // Add list of objects to delete to Delete object - items.SetObjects(objs) + items.Objects = objs - req, _ := s.DeleteObjectsRequest(&s3.DeleteObjectsInput{ + resp, err := s.DeleteObjects(context.TODO(), &s3.DeleteObjectsInput{ Bucket: &s.bucket, Delete: &items, }) - err := req.Send() if err != nil { return nil, mapAwsError(err) } - return &DeleteBlobsOutput{s.getRequestId(req)}, nil + return &DeleteBlobsOutput{s.getRequestId(resp.ResultMetadata)}, nil } func (s *S3Backend) RenameBlob(param *RenameBlobInput) (*RenameBlobOutput, error) { @@ -483,20 +562,20 @@ func (s *S3Backend) mpuCopyPart(from string, to string, mpuId string, bytes stri UploadId: &mpuId, CopySourceRange: &bytes, CopySourceIfMatch: srcEtag, - PartNumber: &part, + PartNumber: aws.Int32(int32(part)), } if s.config.SseC != "" { params.SSECustomerAlgorithm = PString("AES256") - params.SSECustomerKey = &s.config.SseC + params.SSECustomerKey = s.customerKey() params.SSECustomerKeyMD5 = &s.config.SseCDigest params.CopySourceSSECustomerAlgorithm = PString("AES256") - params.CopySourceSSECustomerKey = &s.config.SseC + params.CopySourceSSECustomerKey = s.customerKey() params.CopySourceSSECustomerKeyMD5 = &s.config.SseCDigest } s3Log.Debug(params) - resp, err := s.UploadPartCopy(params) + resp, err := s.UploadPartCopy(context.TODO(), params) if err != nil { s3Log.Errorf("UploadPartCopy %v = %v", params, err) *errout = mapAwsError(err) @@ -561,27 +640,27 @@ func (s *S3Backend) copyObjectMultipart(size int64, from string, to string, mpuI params := &s3.CreateMultipartUploadInput{ Bucket: &s.bucket, Key: &to, - StorageClass: storageClass, + StorageClass: types.StorageClass(aws.ToString(storageClass)), ContentType: s.flags.GetMimeType(to), - Metadata: metadataToLower(metadata), + Metadata: metadataToAWS(metadata), } if s.config.UseSSE { - params.ServerSideEncryption = &s.sseType + params.ServerSideEncryption = s.sseType if s.config.UseKMS && s.config.KMSKeyID != "" { params.SSEKMSKeyId = &s.config.KMSKeyID } } else if s.config.SseC != "" { params.SSECustomerAlgorithm = PString("AES256") - params.SSECustomerKey = &s.config.SseC + params.SSECustomerKey = s.customerKey() params.SSECustomerKeyMD5 = &s.config.SseCDigest } if s.config.ACL != "" { - params.ACL = &s.config.ACL + params.ACL = types.ObjectCannedACL(s.config.ACL) } - resp, err := s.CreateMultipartUpload(params) + resp, err := s.CreateMultipartUpload(context.TODO(), params) if err != nil { return "", mapAwsError(err) } @@ -594,11 +673,11 @@ func (s *S3Backend) copyObjectMultipart(size int64, from string, to string, mpuI if err != nil { return } else { - parts := make([]*s3.CompletedPart, nParts) + parts := make([]types.CompletedPart, nParts) for i := 0; i < nParts; i++ { - parts[i] = &s3.CompletedPart{ + parts[i] = types.CompletedPart{ ETag: etags[i], - PartNumber: aws.Int64(int64(i + 1)), + PartNumber: aws.Int32(int32(i + 1)), } } @@ -606,20 +685,20 @@ func (s *S3Backend) copyObjectMultipart(size int64, from string, to string, mpuI Bucket: &s.bucket, Key: &to, UploadId: &mpuId, - MultipartUpload: &s3.CompletedMultipartUpload{ + MultipartUpload: &types.CompletedMultipartUpload{ Parts: parts, }, } s3Log.Debug(params) - req, _ := s.CompleteMultipartUploadRequest(params) - err = req.Send() + resp, completeErr := s.CompleteMultipartUpload(context.TODO(), params) + err = completeErr if err != nil { s3Log.Errorf("Complete MPU %v = %v", params, err) err = mapAwsError(err) } else { - requestId = s.getRequestId(req) + requestId = s.getRequestId(resp.ResultMetadata) } } @@ -627,9 +706,9 @@ func (s *S3Backend) copyObjectMultipart(size int64, from string, to string, mpuI } func (s *S3Backend) CopyBlob(param *CopyBlobInput) (*CopyBlobOutput, error) { - metadataDirective := s3.MetadataDirectiveCopy + metadataDirective := types.MetadataDirectiveCopy if param.Metadata != nil { - metadataDirective = s3.MetadataDirectiveReplace + metadataDirective = types.MetadataDirectiveReplace } COPY_LIMIT := uint64(5 * 1024 * 1024 * 1024) @@ -673,46 +752,45 @@ func (s *S3Backend) CopyBlob(param *CopyBlobInput) (*CopyBlobOutput, error) { Bucket: &s.bucket, CopySource: aws.String(url.QueryEscape(from)), Key: ¶m.Destination, - StorageClass: param.StorageClass, + StorageClass: types.StorageClass(aws.ToString(param.StorageClass)), ContentType: s.flags.GetMimeType(param.Destination), - Metadata: metadataToLower(param.Metadata), - MetadataDirective: &metadataDirective, + Metadata: metadataToAWS(param.Metadata), + MetadataDirective: metadataDirective, } s3Log.Debug(params) if s.config.UseSSE { - params.ServerSideEncryption = &s.sseType + params.ServerSideEncryption = s.sseType if s.config.UseKMS && s.config.KMSKeyID != "" { params.SSEKMSKeyId = &s.config.KMSKeyID } } else if s.config.SseC != "" { params.SSECustomerAlgorithm = PString("AES256") - params.SSECustomerKey = &s.config.SseC + params.SSECustomerKey = s.customerKey() params.SSECustomerKeyMD5 = &s.config.SseCDigest params.CopySourceSSECustomerAlgorithm = PString("AES256") - params.CopySourceSSECustomerKey = &s.config.SseC + params.CopySourceSSECustomerKey = s.customerKey() params.CopySourceSSECustomerKeyMD5 = &s.config.SseCDigest } if s.config.ACL != "" { - params.ACL = &s.config.ACL + params.ACL = types.ObjectCannedACL(s.config.ACL) } - req, _ := s.CopyObjectRequest(params) - // make a shallow copy of the client so we can change the - // timeout only for this request but still re-use the - // connection pool - c := *(req.Config.HTTPClient) - req.Config.HTTPClient = &c - req.Config.HTTPClient.Timeout = 15 * time.Minute - err := req.Send() + resp, err := s.CopyObject(context.TODO(), params, func(o *s3.Options) { + if client, ok := o.HTTPClient.(*http.Client); ok { + c := *client + c.Timeout = 15 * time.Minute + o.HTTPClient = &c + } + }) if err != nil { s3Log.Errorf("CopyObject %v = %v", params, err) return nil, mapAwsError(err) } - return &CopyBlobOutput{s.getRequestId(req)}, nil + return &CopyBlobOutput{s.getRequestId(resp.ResultMetadata)}, nil } func (s *S3Backend) GetBlob(param *GetBlobInput) (*GetBlobOutput, error) { @@ -723,7 +801,7 @@ func (s *S3Backend) GetBlob(param *GetBlobInput) (*GetBlobOutput, error) { if s.config.SseC != "" { get.SSECustomerAlgorithm = PString("AES256") - get.SSECustomerKey = &s.config.SseC + get.SSECustomerKey = s.customerKey() get.SSECustomerKeyMD5 = &s.config.SseCDigest } @@ -738,8 +816,7 @@ func (s *S3Backend) GetBlob(param *GetBlobInput) (*GetBlobOutput, error) { } // TODO handle IfMatch - req, resp := s.GetObjectRequest(&get) - err := req.Send() + resp, err := s.GetObject(context.TODO(), &get) if err != nil { return nil, mapAwsError(err) } @@ -751,17 +828,20 @@ func (s *S3Backend) GetBlob(param *GetBlobInput) (*GetBlobOutput, error) { ETag: resp.ETag, LastModified: resp.LastModified, Size: uint64(*resp.ContentLength), - StorageClass: resp.StorageClass, + StorageClass: optionalString(string(resp.StorageClass)), }, ContentType: resp.ContentType, - Metadata: metadataToLower(resp.Metadata), + Metadata: metadataFromAWS(resp.Metadata), }, Body: resp.Body, - RequestId: s.getRequestId(req), + RequestId: s.getRequestId(resp.ResultMetadata), }, nil } func getDate(resp *http.Response) *time.Time { + if resp == nil { + return nil + } date := resp.Header.Get("Date") if date != "" { t, err := http.ParseTime(date) @@ -783,38 +863,37 @@ func (s *S3Backend) PutBlob(param *PutBlobInput) (*PutBlobOutput, error) { put := &s3.PutObjectInput{ Bucket: &s.bucket, Key: ¶m.Key, - Metadata: metadataToLower(param.Metadata), + Metadata: metadataToAWS(param.Metadata), Body: param.Body, - StorageClass: &storageClass, + StorageClass: types.StorageClass(storageClass), ContentType: param.ContentType, } if s.config.UseSSE { - put.ServerSideEncryption = &s.sseType + put.ServerSideEncryption = s.sseType if s.config.UseKMS && s.config.KMSKeyID != "" { put.SSEKMSKeyId = &s.config.KMSKeyID } } else if s.config.SseC != "" { put.SSECustomerAlgorithm = PString("AES256") - put.SSECustomerKey = &s.config.SseC + put.SSECustomerKey = s.customerKey() put.SSECustomerKeyMD5 = &s.config.SseCDigest } if s.config.ACL != "" { - put.ACL = &s.config.ACL + put.ACL = types.ObjectCannedACL(s.config.ACL) } - req, resp := s.PutObjectRequest(put) - err := req.Send() + resp, err := s.PutObject(context.TODO(), put) if err != nil { return nil, mapAwsError(err) } return &PutBlobOutput{ ETag: resp.ETag, - LastModified: getDate(req.HTTPResponse), + LastModified: getDate(responseHTTP(resp.ResultMetadata)), StorageClass: &storageClass, - RequestId: s.getRequestId(req), + RequestId: s.getRequestId(resp.ResultMetadata), }, nil } @@ -822,26 +901,26 @@ func (s *S3Backend) MultipartBlobBegin(param *MultipartBlobBeginInput) (*Multipa mpu := s3.CreateMultipartUploadInput{ Bucket: &s.bucket, Key: ¶m.Key, - StorageClass: &s.config.StorageClass, + StorageClass: types.StorageClass(s.config.StorageClass), ContentType: param.ContentType, } if s.config.UseSSE { - mpu.ServerSideEncryption = &s.sseType + mpu.ServerSideEncryption = s.sseType if s.config.UseKMS && s.config.KMSKeyID != "" { mpu.SSEKMSKeyId = &s.config.KMSKeyID } } else if s.config.SseC != "" { mpu.SSECustomerAlgorithm = PString("AES256") - mpu.SSECustomerKey = &s.config.SseC + mpu.SSECustomerKey = s.customerKey() mpu.SSECustomerKeyMD5 = &s.config.SseCDigest } if s.config.ACL != "" { - mpu.ACL = &s.config.ACL + mpu.ACL = types.ObjectCannedACL(s.config.ACL) } - resp, err := s.CreateMultipartUpload(&mpu) + resp, err := s.CreateMultipartUpload(context.TODO(), &mpu) if err != nil { s3Log.Errorf("CreateMultipartUpload %v = %v", param.Key, err) return nil, mapAwsError(err) @@ -862,19 +941,18 @@ func (s *S3Backend) MultipartBlobAdd(param *MultipartBlobAddInput) (*MultipartBl params := s3.UploadPartInput{ Bucket: &s.bucket, Key: param.Commit.Key, - PartNumber: aws.Int64(int64(param.PartNumber)), + PartNumber: aws.Int32(int32(param.PartNumber)), UploadId: param.Commit.UploadId, Body: param.Body, } if s.config.SseC != "" { params.SSECustomerAlgorithm = PString("AES256") - params.SSECustomerKey = &s.config.SseC + params.SSECustomerKey = s.customerKey() params.SSECustomerKeyMD5 = &s.config.SseCDigest } s3Log.Debug(params) - req, resp := s.UploadPartRequest(¶ms) - err := req.Send() + resp, err := s.UploadPart(context.TODO(), ¶ms) if err != nil { return nil, mapAwsError(err) } @@ -884,15 +962,15 @@ func (s *S3Backend) MultipartBlobAdd(param *MultipartBlobAddInput) (*MultipartBl } *en = resp.ETag - return &MultipartBlobAddOutput{s.getRequestId(req)}, nil + return &MultipartBlobAddOutput{s.getRequestId(resp.ResultMetadata)}, nil } func (s *S3Backend) MultipartBlobCommit(param *MultipartBlobCommitInput) (*MultipartBlobCommitOutput, error) { - parts := make([]*s3.CompletedPart, param.NumParts) + parts := make([]types.CompletedPart, param.NumParts) for i := uint32(0); i < param.NumParts; i++ { - parts[i] = &s3.CompletedPart{ + parts[i] = types.CompletedPart{ ETag: param.Parts[i], - PartNumber: aws.Int64(int64(i + 1)), + PartNumber: aws.Int32(int32(i + 1)), } } @@ -900,15 +978,14 @@ func (s *S3Backend) MultipartBlobCommit(param *MultipartBlobCommitInput) (*Multi Bucket: &s.bucket, Key: param.Key, UploadId: param.UploadId, - MultipartUpload: &s3.CompletedMultipartUpload{ + MultipartUpload: &types.CompletedMultipartUpload{ Parts: parts, }, } s3Log.Debug(mpu) - req, resp := s.CompleteMultipartUploadRequest(&mpu) - err := req.Send() + resp, err := s.CompleteMultipartUpload(context.TODO(), &mpu) if err != nil { return nil, mapAwsError(err) } @@ -917,8 +994,8 @@ func (s *S3Backend) MultipartBlobCommit(param *MultipartBlobCommitInput) (*Multi return &MultipartBlobCommitOutput{ ETag: resp.ETag, - LastModified: getDate(req.HTTPResponse), - RequestId: s.getRequestId(req), + LastModified: getDate(responseHTTP(resp.ResultMetadata)), + RequestId: s.getRequestId(resp.ResultMetadata), }, nil } @@ -928,16 +1005,15 @@ func (s *S3Backend) MultipartBlobAbort(param *MultipartBlobCommitInput) (*Multip Key: param.Key, UploadId: param.UploadId, } - req, _ := s.AbortMultipartUploadRequest(&mpu) - err := req.Send() + resp, err := s.AbortMultipartUpload(context.TODO(), &mpu) if err != nil { return nil, mapAwsError(err) } - return &MultipartBlobAbortOutput{s.getRequestId(req)}, nil + return &MultipartBlobAbortOutput{s.getRequestId(resp.ResultMetadata)}, nil } func (s *S3Backend) MultipartExpire(param *MultipartExpireInput) (*MultipartExpireOutput, error) { - mpu, err := s.ListMultipartUploads(&s3.ListMultipartUploadsInput{ + mpu, err := s.ListMultipartUploads(context.TODO(), &s3.ListMultipartUploadsInput{ Bucket: &s.bucket, }) if err != nil { @@ -955,7 +1031,7 @@ func (s *S3Backend) MultipartExpire(param *MultipartExpireInput) (*MultipartExpi Key: upload.Key, UploadId: upload.UploadId, } - resp, err := s.AbortMultipartUpload(params) + resp, err := s.AbortMultipartUpload(context.TODO(), params) s3Log.Debug(resp) if mapAwsError(err) == syscall.EACCES { @@ -970,7 +1046,7 @@ func (s *S3Backend) MultipartExpire(param *MultipartExpireInput) (*MultipartExpi } func (s *S3Backend) RemoveBucket(param *RemoveBucketInput) (*RemoveBucketOutput, error) { - _, err := s.DeleteBucket(&s3.DeleteBucketInput{Bucket: &s.bucket}) + _, err := s.DeleteBucket(context.TODO(), &s3.DeleteBucketInput{Bucket: &s.bucket}) if err != nil { return nil, mapAwsError(err) } @@ -978,28 +1054,32 @@ func (s *S3Backend) RemoveBucket(param *RemoveBucketInput) (*RemoveBucketOutput, } func (s *S3Backend) MakeBucket(param *MakeBucketInput) (*MakeBucketOutput, error) { - _, err := s.CreateBucket(&s3.CreateBucketInput{ + input := &s3.CreateBucketInput{ Bucket: &s.bucket, - ACL: &s.config.ACL, - }) + ACL: types.BucketCannedACL(s.config.ACL), + } + if s.awsConfig.Region != "us-east-1" { + input.CreateBucketConfiguration = &types.CreateBucketConfiguration{ + LocationConstraint: types.BucketLocationConstraint(s.awsConfig.Region), + } + } + _, err := s.CreateBucket(context.TODO(), input) if err != nil { return nil, mapAwsError(err) } if s.config.BucketOwner != "" { - var owner s3.Tag - owner.SetKey("Owner") - owner.SetValue(s.config.BucketOwner) + owner := types.Tag{Key: aws.String("Owner"), Value: &s.config.BucketOwner} param := s3.PutBucketTaggingInput{ Bucket: &s.bucket, - Tagging: &s3.Tagging{ - TagSet: []*s3.Tag{&owner}, + Tagging: &types.Tagging{ + TagSet: []types.Tag{owner}, }, } for i := 0; i < 10; i++ { - _, err = s.PutBucketTagging(¶m) + _, err = s.PutBucketTagging(context.TODO(), ¶m) err = mapAwsError((err)) switch err { case nil: diff --git a/internal/backend_signer_compat_test.go b/internal/backend_signer_compat_test.go new file mode 100644 index 00000000..a441114e --- /dev/null +++ b/internal/backend_signer_compat_test.go @@ -0,0 +1,40 @@ +package internal + +import ( + "net/http" + "testing" + "time" + + "github.com/aws/aws-sdk-go-v2/aws" +) + +func TestS3CompatibilityV2Signatures(t *testing.T) { + for _, test := range []struct { + name, endpoint, bucket, signature string + pathStyle bool + }{ + {"path", "https://s3.example/bucket/a%20b%2Bc?uploadId=a%2Bb%2F%3D&partNumber=2&ignored=value", "bucket", "UpDBNEoGs635XNIqfjv9JDFrA00=", true}, + {"virtual", "https://dotted.bucket.s3.example/a%20b%2Bc?uploadId=a%2Bb%2F%3D&partNumber=2", "dotted.bucket", "erwfWpTSRZUfBOPgWwSdoBzny5w=", false}, + {"escaped", "https://s3.example/bucket/a%252Fb", "bucket", "hpO3z8swEuHkgAOCLtzSEgBo/pI=", true}, + {"pathFallback", "https://s3.example/bucket/a%252Fb", "bucket", "hpO3z8swEuHkgAOCLtzSEgBo/pI=", false}, + } { + t.Run(test.name, func(t *testing.T) { + req, err := http.NewRequestWithContext(t.Context(), "PUT", test.endpoint, nil) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Content-MD5", "md5") + req.Header.Set("Content-Type", "text/plain") + req.Header.Set("Authorization", "previous signature") + credentials := aws.Credentials{AccessKeyID: "access", SecretAccessKey: "secret", SessionToken: "token"} + for range 2 { + if err := SignV2(req, credentials, time.Date(2026, 9, 11, 0, 0, 0, 0, time.UTC), test.pathStyle, test.bucket); err != nil { + t.Fatal(err) + } + if got := req.Header.Get("Authorization"); got != "AWS access:"+test.signature { + t.Fatalf("signature = %q", got) + } + } + }) + } +} diff --git a/internal/dir.go b/internal/dir.go index 64e838f2..a1dbeae0 100644 --- a/internal/dir.go +++ b/internal/dir.go @@ -23,7 +23,7 @@ import ( "syscall" "time" - "github.com/aws/aws-sdk-go/aws" + "github.com/aws/aws-sdk-go-v2/aws" "github.com/jacobsa/fuse" "github.com/jacobsa/fuse/fuseops" diff --git a/internal/goofys.go b/internal/goofys.go index 4f4bd76f..b693cf79 100644 --- a/internal/goofys.go +++ b/internal/goofys.go @@ -18,6 +18,7 @@ import ( . "github.com/kahing/goofys/api/common" "context" + "errors" "fmt" "math/rand" "net/url" @@ -28,7 +29,8 @@ import ( "syscall" "time" - "github.com/aws/aws-sdk-go/aws/awserr" + "github.com/aws/smithy-go" + smithyhttp "github.com/aws/smithy-go/transport/http" "github.com/jacobsa/fuse" "github.com/jacobsa/fuse/fuseops" @@ -542,36 +544,26 @@ func mapAwsError(err error) error { return nil } - if awsErr, ok := err.(awserr.Error); ok { - switch awsErr.Code() { - case "BucketRegionError": - // don't need to log anything, we should detect region after - return err + var awsErr smithy.APIError + if errors.As(err, &awsErr) { + switch awsErr.ErrorCode() { case "NoSuchBucket": return syscall.ENXIO case "BucketAlreadyOwnedByYou": return fuse.EEXIST } + } - if reqErr, ok := err.(awserr.RequestFailure); ok { - // A service error occurred - err = mapHttpError(reqErr.StatusCode()) - if err != nil { - return err - } else { - s3Log.Errorf("http=%v %v s3=%v request=%v\n", - reqErr.StatusCode(), reqErr.Message(), - awsErr.Code(), reqErr.RequestID()) - return reqErr - } - } else { - // Generic AWS Error with Code, Message, and original error (if any) - s3Log.Errorf("code=%v msg=%v, err=%v\n", awsErr.Code(), awsErr.Message(), awsErr.OrigErr()) - return awsErr + var reqErr *smithyhttp.ResponseError + if errors.As(err, &reqErr) { + if mapped := mapHttpError(reqErr.HTTPStatusCode()); mapped != nil { + return mapped } - } else { - return err + s3Log.Errorf("http=%v err=%v", reqErr.HTTPStatusCode(), err) + } else if awsErr != nil { + s3Log.Errorf("code=%v msg=%v, err=%v", awsErr.ErrorCode(), awsErr.ErrorMessage(), err) } + return err } func (fs *Goofys) allocateInodeId() (id fuseops.InodeID) { diff --git a/internal/goofys_test.go b/internal/goofys_test.go index 2b6196d1..0f376e75 100644 --- a/internal/goofys_test.go +++ b/internal/goofys_test.go @@ -40,9 +40,7 @@ import ( "context" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/aws/corehandlers" - "github.com/aws/aws-sdk-go/aws/credentials" + "github.com/aws/aws-sdk-go-v2/aws" "github.com/Azure/azure-storage-blob-go/azblob" "github.com/Azure/go-autorest/autorest" @@ -501,11 +499,9 @@ func (s *GoofysTest) SetUpTest(t *C) { } if s.emulator { - s3.Handlers.Sign.Clear() - s3.Handlers.Sign.PushBack(SignV2) - s3.Handlers.Sign.PushBackNamed(corehandlers.BuildContentLengthHandler) + t.Assert(s3.fallbackV2Signer(), IsNil) } - _, err = s3.ListBuckets(nil) + _, err = s3.ListBuckets(context.TODO(), nil) t.Assert(err, IsNil) } else if cloud == "gcs3" { @@ -2057,7 +2053,7 @@ func (s *GoofysTest) anonymous(t *C) { s3, ok = cloud.Delegate().(*S3Backend) t.Assert(ok, Equals, true) - s3.awsConfig.Credentials = credentials.AnonymousCredentials + s3.awsConfig.Credentials = aws.AnonymousCredentials{} s3.newS3() } @@ -2968,7 +2964,7 @@ func (s *GoofysTest) TestRead403(t *C) { fh, err := in.OpenFile(fuseops.OpContext{uint32(os.Getpid())}) t.Assert(err, IsNil) - s3.awsConfig.Credentials = credentials.AnonymousCredentials + s3.awsConfig.Credentials = aws.AnonymousCredentials{} s3.newS3() // fake enable read-ahead @@ -3366,9 +3362,7 @@ func (s *GoofysTest) newBackend(t *C, bucket string, createBucket bool) (cloud S s3.aws = hasEnv("AWS") if s.emulator { - s3.Handlers.Sign.Clear() - s3.Handlers.Sign.PushBack(SignV2) - s3.Handlers.Sign.PushBackNamed(corehandlers.BuildContentLengthHandler) + t.Assert(s3.fallbackV2Signer(), IsNil) } if s3.aws { diff --git a/internal/handles.go b/internal/handles.go index e8b308fb..8610dab9 100644 --- a/internal/handles.go +++ b/internal/handles.go @@ -25,7 +25,7 @@ import ( "syscall" "time" - "github.com/aws/aws-sdk-go/aws" + "github.com/aws/aws-sdk-go-v2/aws" "github.com/jacobsa/fuse" "github.com/jacobsa/fuse/fuseops" diff --git a/internal/minio_test.go b/internal/minio_test.go index 0f8a78e3..4cb56cce 100644 --- a/internal/minio_test.go +++ b/internal/minio_test.go @@ -20,7 +20,7 @@ import ( "context" - "github.com/aws/aws-sdk-go/service/s3" + "github.com/aws/aws-sdk-go-v2/service/s3" ) type MinioTest struct { @@ -47,14 +47,14 @@ func (s *MinioTest) SetUpSuite(t *C) { s.s3, err = NewS3("", &s.flags, conf) t.Assert(err, IsNil) - _, err = s.s3.ListBuckets(nil) + _, err = s.s3.ListBuckets(context.TODO(), nil) t.Assert(err, IsNil) } func (s *MinioTest) SetUpTest(t *C) { bucket := RandStringBytesMaskImprSrc(32) - _, err := s.s3.CreateBucket(&s3.CreateBucketInput{ + _, err := s.s3.CreateBucket(context.TODO(), &s3.CreateBucketInput{ Bucket: &bucket, }) t.Assert(err, IsNil) diff --git a/internal/v2signer.go b/internal/v2signer.go index 53174c68..2adc4950 100644 --- a/internal/v2signer.go +++ b/internal/v2signer.go @@ -15,6 +15,7 @@ package internal import ( + "context" "crypto/hmac" "crypto/sha1" "encoding/base64" @@ -26,10 +27,11 @@ import ( "strings" "time" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/aws/credentials" - "github.com/aws/aws-sdk-go/aws/request" - "github.com/aws/aws-sdk-go/private/protocol/rest" + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/aws/signer/v4" + "github.com/aws/aws-sdk-go-v2/service/s3" + "github.com/aws/smithy-go/encoding/httpbinding" + "github.com/aws/smithy-go/logging" ) var ( @@ -65,9 +67,9 @@ type signer struct { // Values that must be populated from the request Request *http.Request Time time.Time - Credentials *credentials.Credentials - Debug aws.LogLevelType - Logger aws.Logger + Credentials aws.Credentials + Debug aws.ClientLogMode + Logger logging.Logger pathStyle bool bucket string @@ -76,35 +78,34 @@ type signer struct { signature string } -// Sign requests with signature version 2. -// -// Will sign the requests with the service config's Credentials object -// Signing is skipped if the credentials is the credentials.AnonymousCredentials -// object. -func SignV2(req *request.Request) { - // If the request does not need to be signed ignore the signing of the - // request if the AnonymousCredentials object is used. - if req.Config.Credentials == credentials.AnonymousCredentials { - return - } - +func SignV2(req *http.Request, credentials aws.Credentials, signingTime time.Time, pathStyle bool, bucket string) error { v2 := signer{ - Request: req.HTTPRequest, - Time: req.Time, - Credentials: req.Config.Credentials, - Debug: req.Config.LogLevel.Value(), - Logger: req.Config.Logger, - pathStyle: aws.BoolValue(req.Config.S3ForcePathStyle), + Request: req, + Time: signingTime, + Credentials: credentials, + pathStyle: pathStyle, + bucket: bucket, } + return v2.Sign() +} - req.Error = v2.Sign() +type v2HTTPsigner struct { + bucket string + pathStyle bool } -func (v2 *signer) Sign() error { - credValue, err := v2.Credentials.Get() - if err != nil { - return err +func (s v2HTTPsigner) SignHTTP(ctx context.Context, credentials aws.Credentials, req *http.Request, payloadHash, service, region string, signingTime time.Time, optFns ...func(*v4.SignerOptions)) error { + return SignV2(req, credentials, signingTime, s.pathStyle, s.bucket) +} + +func V2Signer(bucket string) func(*s3.Options) { + return func(o *s3.Options) { + o.HTTPSignerV4 = v2HTTPsigner{bucket: bucket, pathStyle: o.UsePathStyle} } +} + +func (v2 *signer) Sign() error { + credValue := v2.Credentials v2.Query = v2.Request.URL.Query() @@ -131,10 +132,9 @@ func (v2 *signer) Sign() error { } else { uri = v2.Request.URL.Path } - path := rest.EscapePath(uri, false) - if !v2.pathStyle { - host := strings.SplitN(v2.Request.URL.Host, ".", 2)[0] - path = "/" + host + uri + path := httpbinding.EscapePath(uri, false) + if !v2.pathStyle && strings.HasPrefix(v2.Request.URL.Host, v2.bucket+".") { + path = "/" + v2.bucket + path } if path == "" { path = "/" @@ -192,7 +192,7 @@ func (v2 *signer) Sign() error { v2.Request.Header.Set("Authorization", "AWS "+credValue.AccessKeyID+":"+v2.signature) - if v2.Debug.Matches(aws.LogDebugWithSigning) { + if v2.Debug.IsSigning() && v2.Logger != nil { v2.logSigningInfo() } @@ -208,5 +208,5 @@ const logSignInfoMsg = `DEBUG: Request Signature: func (v2 *signer) logSigningInfo() { msg := fmt.Sprintf(logSignInfoMsg, v2.stringToSign, v2.Request.Header.Get("Authorization")) - v2.Logger.Log(msg) + v2.Logger.Logf(logging.Debug, "%s", msg) }