diff --git a/attachment/store_s3.go b/attachment/store_s3.go index 118da4ce..5c47a81b 100644 --- a/attachment/store_s3.go +++ b/attachment/store_s3.go @@ -4,24 +4,20 @@ import ( "context" "fmt" "io" - "net/url" - "strings" + "os" "sync" - "github.com/aws/aws-sdk-go-v2/aws" - "github.com/aws/aws-sdk-go-v2/credentials" - "github.com/aws/aws-sdk-go-v2/service/s3" - s3types "github.com/aws/aws-sdk-go-v2/service/s3/types" "heckel.io/ntfy/v2/log" + "heckel.io/ntfy/v2/s3" "heckel.io/ntfy/v2/util" ) -const tagS3Store = "s3_store" +const ( + tagS3Store = "s3_store" +) type s3Store struct { client *s3.Client - bucket string - prefix string totalSizeCurrent int64 totalSizeLimit int64 mu sync.Mutex @@ -31,14 +27,12 @@ type s3Store struct { // // s3://ACCESS_KEY:SECRET_KEY@BUCKET[/PREFIX]?region=REGION[&endpoint=ENDPOINT] func NewS3Store(s3URL string, totalSizeLimit int64) (Store, error) { - bucket, prefix, client, err := parseS3URL(s3URL) + cfg, err := s3.ParseURL(s3URL) if err != nil { return nil, err } store := &s3Store{ - client: client, - bucket: bucket, - prefix: prefix, + client: s3.New(cfg), totalSizeLimit: totalSizeLimit, } if totalSizeLimit > 0 { @@ -51,98 +45,40 @@ func NewS3Store(s3URL string, totalSizeLimit int64) (Store, error) { return store, nil } -func parseS3URL(s3URL string) (bucket string, prefix string, client *s3.Client, err error) { - u, err := url.Parse(s3URL) - if err != nil { - return "", "", nil, fmt.Errorf("s3 store: invalid URL: %w", err) - } - if u.Scheme != "s3" { - return "", "", nil, fmt.Errorf("s3 store: URL scheme must be 's3', got '%s'", u.Scheme) - } - if u.Host == "" { - return "", "", nil, fmt.Errorf("s3 store: bucket name must be specified as host") - } - bucket = u.Host - prefix = strings.TrimPrefix(u.Path, "/") - - accessKey := u.User.Username() - secretKey, _ := u.User.Password() - if accessKey == "" || secretKey == "" { - return "", "", nil, fmt.Errorf("s3 store: access key and secret key must be specified in URL") - } - - region := u.Query().Get("region") - if region == "" { - return "", "", nil, fmt.Errorf("s3 store: region query parameter is required") - } - endpoint := u.Query().Get("endpoint") - - cfg := aws.Config{ - Region: region, - Credentials: credentials.NewStaticCredentialsProvider(accessKey, secretKey, ""), - } - var opts []func(*s3.Options) - if endpoint != "" { - opts = append(opts, func(o *s3.Options) { - o.BaseEndpoint = aws.String(endpoint) - o.UsePathStyle = true - }) - } - client = s3.NewFromConfig(cfg, opts...) - return bucket, prefix, client, nil -} - -func (c *s3Store) objectKey(id string) string { - if c.prefix != "" { - return c.prefix + "/" + id - } - return id -} - func (c *s3Store) Write(id string, in io.Reader, limiters ...util.Limiter) (int64, error) { if !fileIDRegex.MatchString(id) { return 0, errInvalidFileID } log.Tag(tagS3Store).Field("message_id", id).Debug("Writing attachment to S3") - // Use io.Pipe so we can apply limiters while streaming to S3 - pr, pw := io.Pipe() - var writeErr error - var size int64 - + // Write through limiters into a temp file. This avoids buffering the full attachment in + // memory while still giving us the Content-Length that PutObject requires. limiters = append(limiters, util.NewFixedLimiter(c.Remaining())) - go func() { - limitWriter := util.NewLimitWriter(pw, limiters...) - size, writeErr = io.Copy(limitWriter, in) - if writeErr != nil { - pw.CloseWithError(writeErr) - } else { - pw.Close() - } - }() - - key := c.objectKey(id) - _, err := c.client.PutObject(context.Background(), &s3.PutObjectInput{ - Bucket: aws.String(c.bucket), - Key: aws.String(key), - Body: pr, - }) + tmpFile, err := os.CreateTemp("", "ntfy-s3-upload-*") if err != nil { - // If the limiter caused the error, return the original write error - if writeErr != nil { - return 0, writeErr - } - return 0, fmt.Errorf("s3 store: PutObject failed: %w", err) + return 0, fmt.Errorf("s3 store: failed to create temp file: %w", err) } - if writeErr != nil { - // The write goroutine failed but PutObject somehow succeeded; clean up - _, _ = c.client.DeleteObject(context.Background(), &s3.DeleteObjectInput{ - Bucket: aws.String(c.bucket), - Key: aws.String(key), - }) - return 0, writeErr + tmpPath := tmpFile.Name() + defer os.Remove(tmpPath) + limitWriter := util.NewLimitWriter(tmpFile, limiters...) + size, err := io.Copy(limitWriter, in) + if err != nil { + tmpFile.Close() + return 0, err + } + if err := tmpFile.Close(); err != nil { + return 0, err } + // Re-open the temp file for reading and stream it to S3 + f, err := os.Open(tmpPath) + if err != nil { + return 0, err + } + defer f.Close() + if err := c.client.PutObject(context.Background(), id, f, size); err != nil { + return 0, err + } c.mu.Lock() c.totalSizeCurrent += size c.mu.Unlock() @@ -153,19 +89,7 @@ func (c *s3Store) Read(id string) (io.ReadCloser, int64, error) { if !fileIDRegex.MatchString(id) { return nil, 0, errInvalidFileID } - key := c.objectKey(id) - resp, err := c.client.GetObject(context.Background(), &s3.GetObjectInput{ - Bucket: aws.String(c.bucket), - Key: aws.String(key), - }) - if err != nil { - return nil, 0, fmt.Errorf("s3 store: GetObject failed: %w", err) - } - var size int64 - if resp.ContentLength != nil { - size = *resp.ContentLength - } - return resp.Body, size, nil + return c.client.GetObject(context.Background(), id) } func (c *s3Store) Remove(ids ...string) error { @@ -181,23 +105,11 @@ func (c *s3Store) Remove(ids ...string) error { end = len(ids) } batch := ids[i:end] - objects := make([]s3types.ObjectIdentifier, len(batch)) - for j, id := range batch { + for _, id := range batch { log.Tag(tagS3Store).Field("message_id", id).Debug("Deleting attachment from S3") - key := c.objectKey(id) - objects[j] = s3types.ObjectIdentifier{ - Key: aws.String(key), - } } - _, err := c.client.DeleteObjects(context.Background(), &s3.DeleteObjectsInput{ - Bucket: aws.String(c.bucket), - Delete: &s3types.Delete{ - Objects: objects, - Quiet: aws.Bool(true), - }, - }) - if err != nil { - return fmt.Errorf("s3 store: DeleteObjects failed: %w", err) + if err := c.client.DeleteObjects(context.Background(), batch); err != nil { + return err } } // Recalculate totalSizeCurrent via ListObjectsV2 (matches fileStore's dirSize rescan pattern) @@ -227,29 +139,15 @@ func (c *s3Store) Remaining() int64 { return remaining } +// computeSize uses ListAllObjects to sum up the total size of all objects with our prefix. func (c *s3Store) computeSize() (int64, error) { - var size int64 - paginator := s3.NewListObjectsV2Paginator(c.client, &s3.ListObjectsV2Input{ - Bucket: aws.String(c.bucket), - Prefix: aws.String(c.prefixForList()), - }) - for paginator.HasMorePages() { - page, err := paginator.NextPage(context.Background()) - if err != nil { - return 0, err - } - for _, obj := range page.Contents { - if obj.Size != nil { - size += *obj.Size - } - } + objects, err := c.client.ListAllObjects(context.Background()) + if err != nil { + return 0, err } - return size, nil -} - -func (c *s3Store) prefixForList() string { - if c.prefix != "" { - return c.prefix + "/" + var totalSize int64 + for _, obj := range objects { + totalSize += obj.Size } - return "" + return totalSize, nil } diff --git a/attachment/store_s3_test.go b/attachment/store_s3_test.go index a1808d0c..c898244d 100644 --- a/attachment/store_s3_test.go +++ b/attachment/store_s3_test.go @@ -1,76 +1,282 @@ package attachment import ( + "bytes" + "encoding/xml" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync" "testing" "github.com/stretchr/testify/require" + "heckel.io/ntfy/v2/s3" + "heckel.io/ntfy/v2/util" ) -func TestParseS3URL_Success(t *testing.T) { - bucket, prefix, client, err := parseS3URL("s3://AKID:SECRET@my-bucket/attachments?region=us-east-1") +// --- Integration tests using a mock S3 server --- + +func TestS3Store_WriteReadRemove(t *testing.T) { + server := newMockS3Server() + defer server.Close() + + store := newTestS3Store(t, server, "my-bucket", "pfx", 10*1024) + + // Write + size, err := store.Write("abcdefghijkl", strings.NewReader("hello world")) require.Nil(t, err) - require.Equal(t, "my-bucket", bucket) - require.Equal(t, "attachments", prefix) - require.NotNil(t, client) -} + require.Equal(t, int64(11), size) + require.Equal(t, int64(11), store.Size()) -func TestParseS3URL_NoPrefix(t *testing.T) { - bucket, prefix, client, err := parseS3URL("s3://AKID:SECRET@my-bucket?region=us-east-1") + // Read back + reader, readSize, err := store.Read("abcdefghijkl") require.Nil(t, err) - require.Equal(t, "my-bucket", bucket) - require.Equal(t, "", prefix) - require.NotNil(t, client) -} - -func TestParseS3URL_WithEndpoint(t *testing.T) { - bucket, prefix, client, err := parseS3URL("s3://AKID:SECRET@my-bucket/prefix?region=us-east-1&endpoint=https://s3.example.com") + require.Equal(t, int64(11), readSize) + data, err := io.ReadAll(reader) + reader.Close() require.Nil(t, err) - require.Equal(t, "my-bucket", bucket) - require.Equal(t, "prefix", prefix) - require.NotNil(t, client) + require.Equal(t, "hello world", string(data)) + + // Remove + require.Nil(t, store.Remove("abcdefghijkl")) + require.Equal(t, int64(0), store.Size()) + + // Read after remove should fail + _, _, err = store.Read("abcdefghijkl") + require.Error(t, err) } -func TestParseS3URL_NestedPrefix(t *testing.T) { - bucket, prefix, _, err := parseS3URL("s3://AKID:SECRET@my-bucket/a/b/c?region=us-east-1") +func TestS3Store_WriteNoPrefix(t *testing.T) { + server := newMockS3Server() + defer server.Close() + + store := newTestS3Store(t, server, "my-bucket", "", 10*1024) + + size, err := store.Write("abcdefghijkl", strings.NewReader("test")) require.Nil(t, err) - require.Equal(t, "my-bucket", bucket) - require.Equal(t, "a/b/c", prefix) + require.Equal(t, int64(4), size) + + reader, _, err := store.Read("abcdefghijkl") + require.Nil(t, err) + data, err := io.ReadAll(reader) + reader.Close() + require.Nil(t, err) + require.Equal(t, "test", string(data)) } -func TestParseS3URL_MissingRegion(t *testing.T) { - _, _, _, err := parseS3URL("s3://AKID:SECRET@my-bucket") +func TestS3Store_WriteTotalSizeLimit(t *testing.T) { + server := newMockS3Server() + defer server.Close() + + store := newTestS3Store(t, server, "my-bucket", "pfx", 100) + + // First write fits + _, err := store.Write("abcdefghijk0", bytes.NewReader(make([]byte, 80))) + require.Nil(t, err) + require.Equal(t, int64(80), store.Size()) + require.Equal(t, int64(20), store.Remaining()) + + // Second write exceeds total limit + _, err = store.Write("abcdefghijk1", bytes.NewReader(make([]byte, 50))) + require.Equal(t, util.ErrLimitReached, err) +} + +func TestS3Store_WriteFileSizeLimit(t *testing.T) { + server := newMockS3Server() + defer server.Close() + + store := newTestS3Store(t, server, "my-bucket", "pfx", 10*1024) + + _, err := store.Write("abcdefghijkl", bytes.NewReader(make([]byte, 200)), util.NewFixedLimiter(100)) + require.Equal(t, util.ErrLimitReached, err) +} + +func TestS3Store_WriteRemoveMultiple(t *testing.T) { + server := newMockS3Server() + defer server.Close() + + store := newTestS3Store(t, server, "my-bucket", "pfx", 10*1024) + + for i := 0; i < 5; i++ { + _, err := store.Write(fmt.Sprintf("abcdefghijk%d", i), bytes.NewReader(make([]byte, 100))) + require.Nil(t, err) + } + require.Equal(t, int64(500), store.Size()) + + require.Nil(t, store.Remove("abcdefghijk1", "abcdefghijk3")) + require.Equal(t, int64(300), store.Size()) +} + +func TestS3Store_ReadNotFound(t *testing.T) { + server := newMockS3Server() + defer server.Close() + + store := newTestS3Store(t, server, "my-bucket", "pfx", 10*1024) + + _, _, err := store.Read("abcdefghijkl") require.Error(t, err) - require.Contains(t, err.Error(), "region") } -func TestParseS3URL_MissingCredentials(t *testing.T) { - _, _, _, err := parseS3URL("s3://my-bucket?region=us-east-1") - require.Error(t, err) - require.Contains(t, err.Error(), "access key") +func TestS3Store_InvalidID(t *testing.T) { + server := newMockS3Server() + defer server.Close() + + store := newTestS3Store(t, server, "my-bucket", "pfx", 10*1024) + + _, err := store.Write("bad", strings.NewReader("x")) + require.Equal(t, errInvalidFileID, err) + + _, _, err = store.Read("bad") + require.Equal(t, errInvalidFileID, err) + + err = store.Remove("bad") + require.Equal(t, errInvalidFileID, err) } -func TestParseS3URL_MissingSecretKey(t *testing.T) { - _, _, _, err := parseS3URL("s3://AKID@my-bucket?region=us-east-1") - require.Error(t, err) - require.Contains(t, err.Error(), "secret key") +// --- Helpers --- + +func newTestS3Store(t *testing.T, server *httptest.Server, bucket, prefix string, totalSizeLimit int64) Store { + t.Helper() + // httptest.NewTLSServer URL is like "https://127.0.0.1:PORT" + host := strings.TrimPrefix(server.URL, "https://") + s := &s3Store{ + client: &s3.Client{ + AccessKey: "AKID", + SecretKey: "SECRET", + Region: "us-east-1", + Endpoint: host, + Bucket: bucket, + Prefix: prefix, + PathStyle: true, + HTTPClient: server.Client(), + }, + totalSizeLimit: totalSizeLimit, + } + // Compute initial size (should be 0 for fresh mock) + size, err := s.computeSize() + require.Nil(t, err) + s.totalSizeCurrent = size + return s } -func TestParseS3URL_WrongScheme(t *testing.T) { - _, _, _, err := parseS3URL("http://AKID:SECRET@my-bucket?region=us-east-1") - require.Error(t, err) - require.Contains(t, err.Error(), "scheme") +// --- Mock S3 server --- +// +// A minimal S3-compatible HTTP server that supports PutObject, GetObject, DeleteObjects, and +// ListObjectsV2. Uses path-style addressing: /{bucket}/{key}. Objects are stored in memory. + +type mockS3Server struct { + objects map[string][]byte // full key (bucket/key) -> body + mu sync.RWMutex } -func TestParseS3URL_EmptyBucket(t *testing.T) { - _, _, _, err := parseS3URL("s3://AKID:SECRET@?region=us-east-1") - require.Error(t, err) - require.Contains(t, err.Error(), "bucket") +func newMockS3Server() *httptest.Server { + m := &mockS3Server{objects: make(map[string][]byte)} + return httptest.NewTLSServer(m) } -func TestS3Store_ObjectKey(t *testing.T) { - s := &s3Store{prefix: "attachments"} - require.Equal(t, "attachments/abcdefghijkl", s.objectKey("abcdefghijkl")) +func (m *mockS3Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { + // Path is /{bucket}[/{key...}] + path := strings.TrimPrefix(r.URL.Path, "/") - s2 := &s3Store{prefix: ""} - require.Equal(t, "abcdefghijkl", s2.objectKey("abcdefghijkl")) + switch { + case r.Method == http.MethodPut: + m.handlePut(w, r, path) + case r.Method == http.MethodGet && r.URL.Query().Get("list-type") == "2": + m.handleList(w, r, path) + case r.Method == http.MethodGet: + m.handleGet(w, r, path) + case r.Method == http.MethodPost && r.URL.Query().Has("delete"): + m.handleDelete(w, r, path) + default: + http.Error(w, "not implemented", http.StatusNotImplemented) + } +} + +func (m *mockS3Server) handlePut(w http.ResponseWriter, r *http.Request, path string) { + body, err := io.ReadAll(r.Body) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + m.mu.Lock() + m.objects[path] = body + m.mu.Unlock() + w.WriteHeader(http.StatusOK) +} + +func (m *mockS3Server) handleGet(w http.ResponseWriter, r *http.Request, path string) { + m.mu.RLock() + body, ok := m.objects[path] + m.mu.RUnlock() + if !ok { + w.WriteHeader(http.StatusNotFound) + w.Write([]byte(`NoSuchKeyThe specified key does not exist.`)) + return + } + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(body))) + w.WriteHeader(http.StatusOK) + w.Write(body) +} + +func (m *mockS3Server) handleDelete(w http.ResponseWriter, r *http.Request, bucketPath string) { + // bucketPath is just the bucket name + body, err := io.ReadAll(r.Body) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + var req struct { + Objects []struct { + Key string `xml:"Key"` + } `xml:"Object"` + } + if err := xml.Unmarshal(body, &req); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + m.mu.Lock() + for _, obj := range req.Objects { + delete(m.objects, bucketPath+"/"+obj.Key) + } + m.mu.Unlock() + w.WriteHeader(http.StatusOK) + w.Write([]byte(``)) +} + +func (m *mockS3Server) handleList(w http.ResponseWriter, r *http.Request, bucketPath string) { + prefix := r.URL.Query().Get("prefix") + m.mu.RLock() + var contents []s3ListObject + for key, body := range m.objects { + // key is "bucket/objectkey", strip bucket prefix + objKey := strings.TrimPrefix(key, bucketPath+"/") + if objKey == key { + continue // different bucket + } + if prefix == "" || strings.HasPrefix(objKey, prefix) { + contents = append(contents, s3ListObject{Key: objKey, Size: int64(len(body))}) + } + } + m.mu.RUnlock() + + resp := s3ListResponse{ + Contents: contents, + IsTruncated: false, + } + w.Header().Set("Content-Type", "application/xml") + w.WriteHeader(http.StatusOK) + xml.NewEncoder(w).Encode(resp) +} + +type s3ListResponse struct { + XMLName xml.Name `xml:"ListBucketResult"` + Contents []s3ListObject `xml:"Contents"` + IsTruncated bool `xml:"IsTruncated"` +} + +type s3ListObject struct { + Key string `xml:"Key"` + Size int64 `xml:"Size"` } diff --git a/go.mod b/go.mod index f3cd7791..c073d6aa 100644 --- a/go.mod +++ b/go.mod @@ -30,9 +30,6 @@ require github.com/pkg/errors v0.9.1 // indirect require ( firebase.google.com/go/v4 v4.19.0 github.com/SherClockHolmes/webpush-go v1.4.0 - github.com/aws/aws-sdk-go-v2 v1.41.4 - github.com/aws/aws-sdk-go-v2/credentials v1.19.12 - github.com/aws/aws-sdk-go-v2/service/s3 v1.97.1 github.com/jackc/pgx/v5 v5.8.0 github.com/microcosm-cc/bluemonday v1.0.27 github.com/prometheus/client_golang v1.23.2 @@ -55,15 +52,6 @@ require ( github.com/GoogleCloudPlatform/opentelemetry-operations-go/exporter/metric v0.55.0 // indirect github.com/GoogleCloudPlatform/opentelemetry-operations-go/internal/resourcemapping v0.55.0 // indirect github.com/MicahParks/keyfunc v1.9.0 // indirect - github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.7 // indirect - github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.20 // indirect - github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.20 // indirect - github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.21 // indirect - github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.7 // indirect - github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.12 // indirect - github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.20 // indirect - github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.20 // indirect - github.com/aws/smithy-go v1.24.2 // indirect github.com/aymerick/douceur v0.2.0 // indirect github.com/beorn7/perks v1.0.1 // indirect github.com/cespare/xxhash/v2 v2.3.0 // indirect diff --git a/go.sum b/go.sum index 3f373614..1c6eada9 100644 --- a/go.sum +++ b/go.sum @@ -40,30 +40,6 @@ github.com/MicahParks/keyfunc v1.9.0 h1:lhKd5xrFHLNOWrDc4Tyb/Q1AJ4LCzQ48GVJyVIID github.com/MicahParks/keyfunc v1.9.0/go.mod h1:IdnCilugA0O/99dW+/MkvlyrsX8+L8+x95xuVNtM5jw= github.com/SherClockHolmes/webpush-go v1.4.0 h1:ocnzNKWN23T9nvHi6IfyrQjkIc0oJWv1B1pULsf9i3s= github.com/SherClockHolmes/webpush-go v1.4.0/go.mod h1:XSq8pKX11vNV8MJEMwjrlTkxhAj1zKfxmyhdV7Pd6UA= -github.com/aws/aws-sdk-go-v2 v1.41.4 h1:10f50G7WyU02T56ox1wWXq+zTX9I1zxG46HYuG1hH/k= -github.com/aws/aws-sdk-go-v2 v1.41.4/go.mod h1:mwsPRE8ceUUpiTgF7QmQIJ7lgsKUPQOUl3o72QBrE1o= -github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.7 h1:3kGOqnh1pPeddVa/E37XNTaWJ8W6vrbYV9lJEkCnhuY= -github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.7/go.mod h1:lyw7GFp3qENLh7kwzf7iMzAxDn+NzjXEAGjKS2UOKqI= -github.com/aws/aws-sdk-go-v2/credentials v1.19.12 h1:oqtA6v+y5fZg//tcTWahyN9PEn5eDU/Wpvc2+kJ4aY8= -github.com/aws/aws-sdk-go-v2/credentials v1.19.12/go.mod h1:U3R1RtSHx6NB0DvEQFGyf/0sbrpJrluENHdPy1j/3TE= -github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.20 h1:CNXO7mvgThFGqOFgbNAP2nol2qAWBOGfqR/7tQlvLmc= -github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.20/go.mod h1:oydPDJKcfMhgfcgBUZaG+toBbwy8yPWubJXBVERtI4o= -github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.20 h1:tN6W/hg+pkM+tf9XDkWUbDEjGLb+raoBMFsTodcoYKw= -github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.20/go.mod h1:YJ898MhD067hSHA6xYCx5ts/jEd8BSOLtQDL3iZsvbc= -github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.21 h1:SwGMTMLIlvDNyhMteQ6r8IJSBPlRdXX5d4idhIGbkXA= -github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.21/go.mod h1:UUxgWxofmOdAMuqEsSppbDtGKLfR04HGsD0HXzvhI1k= -github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.7 h1:5EniKhLZe4xzL7a+fU3C2tfUN4nWIqlLesfrjkuPFTY= -github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.7/go.mod h1:x0nZssQ3qZSnIcePWLvcoFisRXJzcTVvYpAAdYX8+GI= -github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.12 h1:qtJZ70afD3ISKWnoX3xB0J2otEqu3LqicRcDBqsj0hQ= -github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.12/go.mod h1:v2pNpJbRNl4vEUWEh5ytQok0zACAKfdmKS51Hotc3pQ= -github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.20 h1:2HvVAIq+YqgGotK6EkMf+KIEqTISmTYh5zLpYyeTo1Y= -github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.20/go.mod h1:V4X406Y666khGa8ghKmphma/7C0DAtEQYhkq9z4vpbk= -github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.20 h1:siU1A6xjUZ2N8zjTHSXFhB9L/2OY8Dqs0xXiLjF30jA= -github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.20/go.mod h1:4TLZCmVJDM3FOu5P5TJP0zOlu9zWgDWU7aUxWbr+rcw= -github.com/aws/aws-sdk-go-v2/service/s3 v1.97.1 h1:csi9NLpFZXb9fxY7rS1xVzgPRGMt7MSNWeQ6eo247kE= -github.com/aws/aws-sdk-go-v2/service/s3 v1.97.1/go.mod h1:qXVal5H0ChqXP63t6jze5LmFalc7+ZE7wOdLtZ0LCP0= -github.com/aws/smithy-go v1.24.2 h1:FzA3bu/nt/vDvmnkg+R8Xl46gmzEDam6mZ1hzmwXFng= -github.com/aws/smithy-go v1.24.2/go.mod h1:YE2RhdIuDbA5E5bTdciG9KrW3+TiEONeUWCqxX9i1Fc= github.com/aymerick/douceur v0.2.0 h1:Mv+mAeH1Q+n9Fr+oyamOlAkUNPWPlA8PPGR0QAaYuPk= github.com/aymerick/douceur v0.2.0/go.mod h1:wlT5vV2O3h55X9m7iVYN0TBM0NH/MmbLnd30/FjWUq4= github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM= diff --git a/s3/client.go b/s3/client.go new file mode 100644 index 00000000..7fdd8093 --- /dev/null +++ b/s3/client.go @@ -0,0 +1,325 @@ +// Package s3 provides a minimal S3-compatible client that works with AWS S3, DigitalOcean Spaces, +// GCP Cloud Storage, MinIO, Backblaze B2, and other S3-compatible providers. It uses raw HTTP +// requests with AWS Signature V4 signing, no AWS SDK dependency required. +package s3 + +import ( + "bytes" + "context" + "crypto/md5" //nolint:gosec // MD5 is required by the S3 protocol for Content-MD5 headers + "encoding/base64" + "encoding/hex" + "encoding/xml" + "fmt" + "io" + "net/http" + "net/url" + "sort" + "strconv" + "strings" + "time" +) + +// Client is a minimal S3-compatible client. It supports PutObject, GetObject, DeleteObjects, +// and ListObjectsV2 operations using AWS Signature V4 signing. The bucket and optional key prefix +// are fixed at construction time. All operations target the same bucket and prefix. +// +// Fields must not be modified after the Client is passed to any method or goroutine. +type Client struct { + AccessKey string // AWS access key ID + SecretKey string // AWS secret access key + Region string // e.g. "us-east-1" + Endpoint string // host[:port] only, e.g. "s3.amazonaws.com" or "nyc3.digitaloceanspaces.com" + Bucket string // S3 bucket name + Prefix string // optional key prefix (e.g. "attachments"); prepended to all keys automatically + PathStyle bool // if true, use path-style addressing; otherwise virtual-hosted-style + HTTPClient *http.Client // if nil, http.DefaultClient is used +} + +// New creates a new S3 client from the given Config. +func New(config *Config) *Client { + return &Client{ + AccessKey: config.AccessKey, + SecretKey: config.SecretKey, + Region: config.Region, + Endpoint: config.Endpoint, + Bucket: config.Bucket, + Prefix: config.Prefix, + PathStyle: config.PathStyle, + } +} + +// PutObject uploads body to the given key. The key is automatically prefixed with the client's +// configured prefix. The body size must be known in advance. The payload is sent as +// UNSIGNED-PAYLOAD, which is supported by all major S3-compatible providers over HTTPS. +func (c *Client) PutObject(ctx context.Context, key string, body io.Reader, size int64) error { + fullKey := c.objectKey(key) + req, err := http.NewRequestWithContext(ctx, http.MethodPut, c.objectURL(fullKey), body) + if err != nil { + return fmt.Errorf("s3: PutObject request: %w", err) + } + req.ContentLength = size + c.signV4(req, unsignedPayload) + resp, err := c.httpClient().Do(req) + if err != nil { + return fmt.Errorf("s3: PutObject: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode/100 != 2 { + return parseError(resp) + } + return nil +} + +// GetObject downloads an object. The key is automatically prefixed with the client's configured +// prefix. The caller must close the returned ReadCloser. +func (c *Client) GetObject(ctx context.Context, key string) (io.ReadCloser, int64, error) { + fullKey := c.objectKey(key) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.objectURL(fullKey), nil) + if err != nil { + return nil, 0, fmt.Errorf("s3: GetObject request: %w", err) + } + c.signV4(req, emptyPayloadHash) + resp, err := c.httpClient().Do(req) + if err != nil { + return nil, 0, fmt.Errorf("s3: GetObject: %w", err) + } + if resp.StatusCode/100 != 2 { + err := parseError(resp) + resp.Body.Close() + return nil, 0, err + } + return resp.Body, resp.ContentLength, nil +} + +// DeleteObjects removes multiple objects in a single batch request. Keys are automatically +// prefixed with the client's configured prefix. S3 supports up to 1000 keys per call; the +// caller is responsible for batching if needed. +// +// Even when S3 returns HTTP 200, individual keys may fail. If any per-key errors are present +// in the response, they are returned as a combined error. +func (c *Client) DeleteObjects(ctx context.Context, keys []string) error { + var body bytes.Buffer + body.WriteString("true") + for _, key := range keys { + body.WriteString("") + xml.EscapeText(&body, []byte(c.objectKey(key))) + body.WriteString("") + } + body.WriteString("") + bodyBytes := body.Bytes() + payloadHash := sha256Hex(bodyBytes) + + // Content-MD5 is required by the S3 protocol for DeleteObjects requests. + md5Sum := md5.Sum(bodyBytes) //nolint:gosec + contentMD5 := base64.StdEncoding.EncodeToString(md5Sum[:]) + + reqURL := c.bucketURL() + "?delete=" + req, err := http.NewRequestWithContext(ctx, http.MethodPost, reqURL, bytes.NewReader(bodyBytes)) + if err != nil { + return fmt.Errorf("s3: DeleteObjects request: %w", err) + } + req.ContentLength = int64(len(bodyBytes)) + req.Header.Set("Content-Type", "application/xml") + req.Header.Set("Content-MD5", contentMD5) + c.signV4(req, payloadHash) + resp, err := c.httpClient().Do(req) + if err != nil { + return fmt.Errorf("s3: DeleteObjects: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode/100 != 2 { + return parseError(resp) + } + + // S3 may return HTTP 200 with per-key errors in the response body + respBody, err := io.ReadAll(io.LimitReader(resp.Body, maxResponseBytes)) + if err != nil { + return fmt.Errorf("s3: DeleteObjects read response: %w", err) + } + var result deleteResult + if err := xml.Unmarshal(respBody, &result); err != nil { + return nil // If we can't parse, assume success (Quiet mode returns empty body on success) + } + if len(result.Errors) > 0 { + var msgs []string + for _, e := range result.Errors { + msgs = append(msgs, fmt.Sprintf("%s: %s", e.Key, e.Message)) + } + return fmt.Errorf("s3: DeleteObjects partial failure: %s", strings.Join(msgs, "; ")) + } + return nil +} + +// ListObjects performs a single ListObjectsV2 request using the client's configured prefix. +// Use continuationToken for pagination. Set maxKeys to 0 for the server default (typically 1000). +func (c *Client) ListObjects(ctx context.Context, continuationToken string, maxKeys int) (*ListResult, error) { + query := url.Values{"list-type": {"2"}} + if prefix := c.prefixForList(); prefix != "" { + query.Set("prefix", prefix) + } + if continuationToken != "" { + query.Set("continuation-token", continuationToken) + } + if maxKeys > 0 { + query.Set("max-keys", strconv.Itoa(maxKeys)) + } + req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.bucketURL()+"?"+query.Encode(), nil) + if err != nil { + return nil, fmt.Errorf("s3: ListObjects request: %w", err) + } + c.signV4(req, emptyPayloadHash) + resp, err := c.httpClient().Do(req) + if err != nil { + return nil, fmt.Errorf("s3: ListObjects: %w", err) + } + respBody, err := io.ReadAll(io.LimitReader(resp.Body, maxResponseBytes)) + resp.Body.Close() + if err != nil { + return nil, fmt.Errorf("s3: ListObjects read: %w", err) + } + if resp.StatusCode/100 != 2 { + return nil, parseErrorFromBytes(resp.StatusCode, respBody) + } + var result listObjectsV2Response + if err := xml.Unmarshal(respBody, &result); err != nil { + return nil, fmt.Errorf("s3: ListObjects XML: %w", err) + } + objects := make([]Object, len(result.Contents)) + for i, obj := range result.Contents { + objects[i] = Object(obj) + } + return &ListResult{ + Objects: objects, + IsTruncated: result.IsTruncated, + NextContinuationToken: result.NextContinuationToken, + }, nil +} + +// ListAllObjects returns all objects under the client's configured prefix by paginating through +// ListObjectsV2 results automatically. It stops after 10,000 pages as a safety valve. +func (c *Client) ListAllObjects(ctx context.Context) ([]Object, error) { + const maxPages = 10000 + var all []Object + var token string + for page := 0; page < maxPages; page++ { + result, err := c.ListObjects(ctx, token, 0) + if err != nil { + return nil, err + } + all = append(all, result.Objects...) + if !result.IsTruncated { + return all, nil + } + token = result.NextContinuationToken + } + return nil, fmt.Errorf("s3: ListAllObjects exceeded %d pages", maxPages) +} + +// signV4 signs req in place using AWS Signature V4. payloadHash is the hex-encoded SHA-256 +// of the request body, or the literal string "UNSIGNED-PAYLOAD" for streaming uploads. +func (c *Client) signV4(req *http.Request, payloadHash string) { + now := time.Now().UTC() + datestamp := now.Format("20060102") + amzDate := now.Format("20060102T150405Z") + + // Required headers + req.Header.Set("Host", c.hostHeader()) + req.Header.Set("X-Amz-Date", amzDate) + req.Header.Set("X-Amz-Content-Sha256", payloadHash) + + // Canonical headers (all headers we set, sorted by lowercase key) + signedKeys := make([]string, 0, len(req.Header)) + canonHeaders := make(map[string]string, len(req.Header)) + for k := range req.Header { + lk := strings.ToLower(k) + signedKeys = append(signedKeys, lk) + canonHeaders[lk] = strings.TrimSpace(req.Header.Get(k)) + } + sort.Strings(signedKeys) + signedHeadersStr := strings.Join(signedKeys, ";") + var chBuf strings.Builder + for _, k := range signedKeys { + chBuf.WriteString(k) + chBuf.WriteByte(':') + chBuf.WriteString(canonHeaders[k]) + chBuf.WriteByte('\n') + } + + // Canonical request + canonicalRequest := strings.Join([]string{ + req.Method, + canonicalURI(req.URL), + canonicalQueryString(req.URL.Query()), + chBuf.String(), + signedHeadersStr, + payloadHash, + }, "\n") + + // String to sign + credentialScope := datestamp + "/" + c.Region + "/s3/aws4_request" + stringToSign := "AWS4-HMAC-SHA256\n" + amzDate + "\n" + credentialScope + "\n" + sha256Hex([]byte(canonicalRequest)) + + // Signing key + signingKey := hmacSHA256(hmacSHA256(hmacSHA256(hmacSHA256( + []byte("AWS4"+c.SecretKey), []byte(datestamp)), + []byte(c.Region)), + []byte("s3")), + []byte("aws4_request")) + + signature := hex.EncodeToString(hmacSHA256(signingKey, []byte(stringToSign))) + req.Header.Set("Authorization", fmt.Sprintf( + "AWS4-HMAC-SHA256 Credential=%s/%s, SignedHeaders=%s, Signature=%s", + c.AccessKey, credentialScope, signedHeadersStr, signature, + )) +} + +func (c *Client) httpClient() *http.Client { + if c.HTTPClient != nil { + return c.HTTPClient + } + return http.DefaultClient +} + +// objectKey prepends the configured prefix to the given key. +func (c *Client) objectKey(key string) string { + if c.Prefix != "" { + return c.Prefix + "/" + key + } + return key +} + +// prefixForList returns the prefix to use in ListObjectsV2 requests, +// with a trailing slash so that only objects under the prefix directory are returned. +func (c *Client) prefixForList() string { + if c.Prefix != "" { + return c.Prefix + "/" + } + return "" +} + +// bucketURL returns the base URL for bucket-level operations. +func (c *Client) bucketURL() string { + if c.PathStyle { + return fmt.Sprintf("https://%s/%s", c.Endpoint, c.Bucket) + } + return fmt.Sprintf("https://%s.%s", c.Bucket, c.Endpoint) +} + +// objectURL returns the full URL for an object (key should already include the prefix). +// Each path segment is URI-encoded to handle special characters in keys. +func (c *Client) objectURL(key string) string { + segments := strings.Split(key, "/") + for i, seg := range segments { + segments[i] = uriEncode(seg) + } + return c.bucketURL() + "/" + strings.Join(segments, "/") +} + +// hostHeader returns the value for the Host header. +func (c *Client) hostHeader() string { + if c.PathStyle { + return c.Endpoint + } + return c.Bucket + "." + c.Endpoint +} diff --git a/s3/client_test.go b/s3/client_test.go new file mode 100644 index 00000000..f4a10213 --- /dev/null +++ b/s3/client_test.go @@ -0,0 +1,727 @@ +package s3 + +import ( + "bytes" + "context" + "encoding/xml" + "fmt" + "io" + "net/http" + "net/http/httptest" + "os" + "sort" + "strings" + "sync" + "testing" + + "github.com/stretchr/testify/require" +) + +// --- Mock S3 server --- +// +// A minimal S3-compatible HTTP server that supports PutObject, GetObject, DeleteObjects, and +// ListObjectsV2. Uses path-style addressing: /{bucket}/{key}. Objects are stored in memory. + +type mockS3Server struct { + objects map[string][]byte // full key (bucket/key) -> body + mu sync.RWMutex +} + +func newMockS3Server() (*httptest.Server, *mockS3Server) { + m := &mockS3Server{objects: make(map[string][]byte)} + return httptest.NewTLSServer(m), m +} + +func (m *mockS3Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { + // Path is /{bucket}[/{key...}] + path := strings.TrimPrefix(r.URL.Path, "/") + + switch { + case r.Method == http.MethodPut: + m.handlePut(w, r, path) + case r.Method == http.MethodGet && r.URL.Query().Get("list-type") == "2": + m.handleList(w, r, path) + case r.Method == http.MethodGet: + m.handleGet(w, r, path) + case r.Method == http.MethodPost && r.URL.Query().Has("delete"): + m.handleDelete(w, r, path) + default: + http.Error(w, "not implemented", http.StatusNotImplemented) + } +} + +func (m *mockS3Server) handlePut(w http.ResponseWriter, r *http.Request, path string) { + body, err := io.ReadAll(r.Body) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + m.mu.Lock() + m.objects[path] = body + m.mu.Unlock() + w.WriteHeader(http.StatusOK) +} + +func (m *mockS3Server) handleGet(w http.ResponseWriter, r *http.Request, path string) { + m.mu.RLock() + body, ok := m.objects[path] + m.mu.RUnlock() + if !ok { + w.WriteHeader(http.StatusNotFound) + w.Write([]byte(`NoSuchKeyThe specified key does not exist.`)) + return + } + w.Header().Set("Content-Length", fmt.Sprintf("%d", len(body))) + w.WriteHeader(http.StatusOK) + w.Write(body) +} + +type listObjectsResponse struct { + XMLName xml.Name `xml:"ListBucketResult"` + Contents []listObject `xml:"Contents"` + // Pagination support + IsTruncated bool `xml:"IsTruncated"` + NextContinuationToken string `xml:"NextContinuationToken"` +} + +func (m *mockS3Server) handleDelete(w http.ResponseWriter, r *http.Request, bucketPath string) { + // bucketPath is just the bucket name + body, err := io.ReadAll(r.Body) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + var req struct { + Objects []struct { + Key string `xml:"Key"` + } `xml:"Object"` + } + if err := xml.Unmarshal(body, &req); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + m.mu.Lock() + for _, obj := range req.Objects { + delete(m.objects, bucketPath+"/"+obj.Key) + } + m.mu.Unlock() + w.WriteHeader(http.StatusOK) + w.Write([]byte(``)) +} + +func (m *mockS3Server) handleList(w http.ResponseWriter, r *http.Request, bucketPath string) { + prefix := r.URL.Query().Get("prefix") + contToken := r.URL.Query().Get("continuation-token") + + m.mu.RLock() + var allKeys []string + for key := range m.objects { + objKey := strings.TrimPrefix(key, bucketPath+"/") + if objKey == key { + continue // different bucket + } + if prefix == "" || strings.HasPrefix(objKey, prefix) { + allKeys = append(allKeys, objKey) + } + } + m.mu.RUnlock() + sort.Strings(allKeys) + + // Simple continuation token: it's the key to start after + startIdx := 0 + if contToken != "" { + for i, k := range allKeys { + if k == contToken { + startIdx = i + 1 + break + } + } + } + + maxKeys := 1000 + if mk := r.URL.Query().Get("max-keys"); mk != "" { + fmt.Sscanf(mk, "%d", &maxKeys) + } + + endIdx := startIdx + maxKeys + truncated := false + nextToken := "" + if endIdx < len(allKeys) { + truncated = true + nextToken = allKeys[endIdx-1] + allKeys = allKeys[startIdx:endIdx] + } else { + allKeys = allKeys[startIdx:] + } + + m.mu.RLock() + var contents []listObject + for _, objKey := range allKeys { + body := m.objects[bucketPath+"/"+objKey] + contents = append(contents, listObject{Key: objKey, Size: int64(len(body))}) + } + m.mu.RUnlock() + + resp := listObjectsResponse{ + Contents: contents, + IsTruncated: truncated, + NextContinuationToken: nextToken, + } + w.Header().Set("Content-Type", "application/xml") + w.WriteHeader(http.StatusOK) + xml.NewEncoder(w).Encode(resp) +} + +func (m *mockS3Server) objectCount() int { + m.mu.RLock() + defer m.mu.RUnlock() + return len(m.objects) +} + +// --- Helper to create a test client pointing at mock server --- + +func newTestClient(server *httptest.Server, bucket, prefix string) *Client { + // httptest.NewTLSServer URL is like "https://127.0.0.1:PORT" + host := strings.TrimPrefix(server.URL, "https://") + return &Client{ + AccessKey: "AKID", + SecretKey: "SECRET", + Region: "us-east-1", + Endpoint: host, + Bucket: bucket, + Prefix: prefix, + PathStyle: true, + HTTPClient: server.Client(), + } +} + +// --- URL parsing tests --- + +func TestParseURL_Success(t *testing.T) { + cfg, err := ParseURL("s3://AKID:SECRET@my-bucket/attachments?region=us-east-1") + require.Nil(t, err) + require.Equal(t, "my-bucket", cfg.Bucket) + require.Equal(t, "attachments", cfg.Prefix) + require.Equal(t, "us-east-1", cfg.Region) + require.Equal(t, "AKID", cfg.AccessKey) + require.Equal(t, "SECRET", cfg.SecretKey) + require.Equal(t, "s3.us-east-1.amazonaws.com", cfg.Endpoint) + require.False(t, cfg.PathStyle) +} + +func TestParseURL_NoPrefix(t *testing.T) { + cfg, err := ParseURL("s3://AKID:SECRET@my-bucket?region=us-east-1") + require.Nil(t, err) + require.Equal(t, "my-bucket", cfg.Bucket) + require.Equal(t, "", cfg.Prefix) +} + +func TestParseURL_WithEndpoint(t *testing.T) { + cfg, err := ParseURL("s3://AKID:SECRET@my-bucket/prefix?region=us-east-1&endpoint=https://s3.example.com") + require.Nil(t, err) + require.Equal(t, "my-bucket", cfg.Bucket) + require.Equal(t, "prefix", cfg.Prefix) + require.Equal(t, "s3.example.com", cfg.Endpoint) + require.True(t, cfg.PathStyle) +} + +func TestParseURL_EndpointHTTP(t *testing.T) { + cfg, err := ParseURL("s3://AKID:SECRET@my-bucket?region=us-east-1&endpoint=http://localhost:9000") + require.Nil(t, err) + require.Equal(t, "localhost:9000", cfg.Endpoint) + require.True(t, cfg.PathStyle) +} + +func TestParseURL_EndpointTrailingSlash(t *testing.T) { + cfg, err := ParseURL("s3://AKID:SECRET@my-bucket?region=us-east-1&endpoint=https://s3.example.com/") + require.Nil(t, err) + require.Equal(t, "s3.example.com", cfg.Endpoint) +} + +func TestParseURL_NestedPrefix(t *testing.T) { + cfg, err := ParseURL("s3://AKID:SECRET@my-bucket/a/b/c?region=us-east-1") + require.Nil(t, err) + require.Equal(t, "my-bucket", cfg.Bucket) + require.Equal(t, "a/b/c", cfg.Prefix) +} + +func TestParseURL_MissingRegion(t *testing.T) { + _, err := ParseURL("s3://AKID:SECRET@my-bucket") + require.Error(t, err) + require.Contains(t, err.Error(), "region") +} + +func TestParseURL_MissingCredentials(t *testing.T) { + _, err := ParseURL("s3://my-bucket?region=us-east-1") + require.Error(t, err) + require.Contains(t, err.Error(), "access key") +} + +func TestParseURL_MissingSecretKey(t *testing.T) { + _, err := ParseURL("s3://AKID@my-bucket?region=us-east-1") + require.Error(t, err) + require.Contains(t, err.Error(), "secret key") +} + +func TestParseURL_WrongScheme(t *testing.T) { + _, err := ParseURL("http://AKID:SECRET@my-bucket?region=us-east-1") + require.Error(t, err) + require.Contains(t, err.Error(), "scheme") +} + +func TestParseURL_EmptyBucket(t *testing.T) { + _, err := ParseURL("s3://AKID:SECRET@?region=us-east-1") + require.Error(t, err) + require.Contains(t, err.Error(), "bucket") +} + +// --- Unit tests: URL construction --- + +func TestClient_BucketURL_PathStyle(t *testing.T) { + c := &Client{Endpoint: "s3.example.com", Bucket: "my-bucket", PathStyle: true} + require.Equal(t, "https://s3.example.com/my-bucket", c.bucketURL()) +} + +func TestClient_BucketURL_VirtualHosted(t *testing.T) { + c := &Client{Endpoint: "s3.us-east-1.amazonaws.com", Bucket: "my-bucket", PathStyle: false} + require.Equal(t, "https://my-bucket.s3.us-east-1.amazonaws.com", c.bucketURL()) +} + +func TestClient_ObjectURL_PathStyle(t *testing.T) { + c := &Client{Endpoint: "s3.example.com", Bucket: "my-bucket", PathStyle: true} + require.Equal(t, "https://s3.example.com/my-bucket/prefix/obj", c.objectURL("prefix/obj")) +} + +func TestClient_ObjectURL_VirtualHosted(t *testing.T) { + c := &Client{Endpoint: "s3.us-east-1.amazonaws.com", Bucket: "my-bucket", PathStyle: false} + require.Equal(t, "https://my-bucket.s3.us-east-1.amazonaws.com/prefix/obj", c.objectURL("prefix/obj")) +} + +func TestClient_HostHeader_PathStyle(t *testing.T) { + c := &Client{Endpoint: "s3.example.com", Bucket: "my-bucket", PathStyle: true} + require.Equal(t, "s3.example.com", c.hostHeader()) +} + +func TestClient_HostHeader_VirtualHosted(t *testing.T) { + c := &Client{Endpoint: "s3.us-east-1.amazonaws.com", Bucket: "my-bucket", PathStyle: false} + require.Equal(t, "my-bucket.s3.us-east-1.amazonaws.com", c.hostHeader()) +} + +func TestClient_ObjectKey(t *testing.T) { + c := &Client{Prefix: "attachments"} + require.Equal(t, "attachments/file123", c.objectKey("file123")) + + c2 := &Client{Prefix: ""} + require.Equal(t, "file123", c2.objectKey("file123")) +} + +func TestClient_PrefixForList(t *testing.T) { + c := &Client{Prefix: "attachments"} + require.Equal(t, "attachments/", c.prefixForList()) + + c2 := &Client{Prefix: ""} + require.Equal(t, "", c2.prefixForList()) +} + +// --- Integration tests using mock S3 server --- + +func TestClient_PutGetObject(t *testing.T) { + server, _ := newMockS3Server() + defer server.Close() + client := newTestClient(server, "my-bucket", "") + + ctx := context.Background() + + // Put + err := client.PutObject(ctx, "test-key", strings.NewReader("hello world"), 11) + require.Nil(t, err) + + // Get + reader, size, err := client.GetObject(ctx, "test-key") + require.Nil(t, err) + require.Equal(t, int64(11), size) + data, err := io.ReadAll(reader) + reader.Close() + require.Nil(t, err) + require.Equal(t, "hello world", string(data)) +} + +func TestClient_PutGetObject_WithPrefix(t *testing.T) { + server, _ := newMockS3Server() + defer server.Close() + client := newTestClient(server, "my-bucket", "pfx") + + ctx := context.Background() + + err := client.PutObject(ctx, "test-key", strings.NewReader("hello"), 5) + require.Nil(t, err) + + reader, _, err := client.GetObject(ctx, "test-key") + require.Nil(t, err) + data, _ := io.ReadAll(reader) + reader.Close() + require.Equal(t, "hello", string(data)) +} + +func TestClient_GetObject_NotFound(t *testing.T) { + server, _ := newMockS3Server() + defer server.Close() + client := newTestClient(server, "my-bucket", "") + + _, _, err := client.GetObject(context.Background(), "nonexistent") + require.Error(t, err) + var errResp *ErrorResponse + require.ErrorAs(t, err, &errResp) + require.Equal(t, 404, errResp.StatusCode) + require.Equal(t, "NoSuchKey", errResp.Code) +} + +func TestClient_DeleteObjects(t *testing.T) { + server, mock := newMockS3Server() + defer server.Close() + client := newTestClient(server, "my-bucket", "") + + ctx := context.Background() + + // Put several objects + for i := 0; i < 5; i++ { + err := client.PutObject(ctx, fmt.Sprintf("key-%d", i), bytes.NewReader([]byte("data")), 4) + require.Nil(t, err) + } + require.Equal(t, 5, mock.objectCount()) + + // Delete some + err := client.DeleteObjects(ctx, []string{"key-1", "key-3"}) + require.Nil(t, err) + require.Equal(t, 3, mock.objectCount()) + + // Verify deleted ones are gone + _, _, err = client.GetObject(ctx, "key-1") + require.Error(t, err) + _, _, err = client.GetObject(ctx, "key-3") + require.Error(t, err) + + // Verify remaining ones are still there + reader, _, err := client.GetObject(ctx, "key-0") + require.Nil(t, err) + reader.Close() +} + +func TestClient_ListObjects(t *testing.T) { + server, _ := newMockS3Server() + defer server.Close() + + ctx := context.Background() + + // Client with prefix "pfx": list should only return objects under pfx/ + client := newTestClient(server, "my-bucket", "pfx") + for i := 0; i < 3; i++ { + err := client.PutObject(ctx, fmt.Sprintf("%d", i), bytes.NewReader([]byte("x")), 1) + require.Nil(t, err) + } + + // Also put an object outside the prefix using a no-prefix client + clientNoPrefix := newTestClient(server, "my-bucket", "") + err := clientNoPrefix.PutObject(ctx, "other", bytes.NewReader([]byte("y")), 1) + require.Nil(t, err) + + // List with prefix client: should only see 3 + result, err := client.ListObjects(ctx, "", 0) + require.Nil(t, err) + require.Len(t, result.Objects, 3) + require.False(t, result.IsTruncated) + + // List with no-prefix client: should see all 4 + result, err = clientNoPrefix.ListObjects(ctx, "", 0) + require.Nil(t, err) + require.Len(t, result.Objects, 4) +} + +func TestClient_ListObjects_Pagination(t *testing.T) { + server, _ := newMockS3Server() + defer server.Close() + client := newTestClient(server, "my-bucket", "") + + ctx := context.Background() + + // Put 5 objects + for i := 0; i < 5; i++ { + err := client.PutObject(ctx, fmt.Sprintf("key-%02d", i), bytes.NewReader([]byte("x")), 1) + require.Nil(t, err) + } + + // List with max-keys=2 + result, err := client.ListObjects(ctx, "", 2) + require.Nil(t, err) + require.Len(t, result.Objects, 2) + require.True(t, result.IsTruncated) + require.NotEmpty(t, result.NextContinuationToken) + + // Get next page + result2, err := client.ListObjects(ctx, result.NextContinuationToken, 2) + require.Nil(t, err) + require.Len(t, result2.Objects, 2) + require.True(t, result2.IsTruncated) + + // Get last page + result3, err := client.ListObjects(ctx, result2.NextContinuationToken, 2) + require.Nil(t, err) + require.Len(t, result3.Objects, 1) + require.False(t, result3.IsTruncated) +} + +func TestClient_ListAllObjects(t *testing.T) { + server, _ := newMockS3Server() + defer server.Close() + client := newTestClient(server, "my-bucket", "pfx") + + ctx := context.Background() + + for i := 0; i < 10; i++ { + err := client.PutObject(ctx, fmt.Sprintf("key-%02d", i), bytes.NewReader([]byte("x")), 1) + require.Nil(t, err) + } + + objects, err := client.ListAllObjects(ctx) + require.Nil(t, err) + require.Len(t, objects, 10) +} + +func TestClient_PutObject_LargeBody(t *testing.T) { + server, _ := newMockS3Server() + defer server.Close() + client := newTestClient(server, "my-bucket", "") + + ctx := context.Background() + + // 1 MB object + data := make([]byte, 1024*1024) + for i := range data { + data[i] = byte(i % 256) + } + err := client.PutObject(ctx, "large", bytes.NewReader(data), int64(len(data))) + require.Nil(t, err) + + reader, size, err := client.GetObject(ctx, "large") + require.Nil(t, err) + require.Equal(t, int64(1024*1024), size) + got, err := io.ReadAll(reader) + reader.Close() + require.Nil(t, err) + require.Equal(t, data, got) +} + +func TestClient_PutObject_NestedKey(t *testing.T) { + server, _ := newMockS3Server() + defer server.Close() + client := newTestClient(server, "my-bucket", "") + + ctx := context.Background() + + err := client.PutObject(ctx, "deep/nested/prefix/file.txt", strings.NewReader("nested"), 6) + require.Nil(t, err) + + reader, _, err := client.GetObject(ctx, "deep/nested/prefix/file.txt") + require.Nil(t, err) + data, _ := io.ReadAll(reader) + reader.Close() + require.Equal(t, "nested", string(data)) +} + +// --- Scale test: 20k objects (ntfy-adjacent) --- + +func TestClient_ListAllObjects_20k(t *testing.T) { + if testing.Short() { + t.Skip("skipping 20k object test in short mode") + } + + server, _ := newMockS3Server() + defer server.Close() + client := newTestClient(server, "my-bucket", "attachments") + + ctx := context.Background() + const numObjects = 20000 + const batchSize = 500 + + // Insert 20k objects in batches to keep it fast + for batch := 0; batch < numObjects/batchSize; batch++ { + for i := 0; i < batchSize; i++ { + idx := batch*batchSize + i + key := fmt.Sprintf("%08d", idx) + err := client.PutObject(ctx, key, bytes.NewReader([]byte("x")), 1) + require.Nil(t, err) + } + } + + // List all 20k objects with pagination + objects, err := client.ListAllObjects(ctx) + require.Nil(t, err) + require.Len(t, objects, numObjects) + + // Verify total size + var totalSize int64 + for _, obj := range objects { + totalSize += obj.Size + } + require.Equal(t, int64(numObjects), totalSize) + + // Delete 1000 objects (simulating attachment expiry cleanup) + keys := make([]string, 1000) + for i := range keys { + keys[i] = fmt.Sprintf("%08d", i) + } + err = client.DeleteObjects(ctx, keys) + require.Nil(t, err) + + // List again: should have 19000 + objects, err = client.ListAllObjects(ctx) + require.Nil(t, err) + require.Len(t, objects, numObjects-1000) +} + +// --- Real S3 integration test --- +// +// Set the following environment variables to run this test against a real S3 bucket: +// +// S3_ACCESS_KEY, S3_SECRET_KEY, S3_REGION, S3_BUCKET +// +// Optional: +// +// S3_ENDPOINT: host[:port] for S3-compatible providers (e.g. "nyc3.digitaloceanspaces.com") +// S3_PATH_STYLE: set to "true" for path-style addressing +// S3_PREFIX: key prefix to use (default: "ntfy-s3-test") +func TestClient_RealBucket(t *testing.T) { + accessKey := os.Getenv("S3_ACCESS_KEY") + secretKey := os.Getenv("S3_SECRET_KEY") + region := os.Getenv("S3_REGION") + bucket := os.Getenv("S3_BUCKET") + + if accessKey == "" || secretKey == "" || region == "" || bucket == "" { + t.Skip("skipping real S3 test: set S3_ACCESS_KEY, S3_SECRET_KEY, S3_REGION, S3_BUCKET") + } + + endpoint := os.Getenv("S3_ENDPOINT") + if endpoint == "" { + endpoint = fmt.Sprintf("s3.%s.amazonaws.com", region) + } + pathStyle := os.Getenv("S3_PATH_STYLE") == "true" + prefix := os.Getenv("S3_PREFIX") + if prefix == "" { + prefix = "ntfy-s3-test" + } + + client := &Client{ + AccessKey: accessKey, + SecretKey: secretKey, + Region: region, + Endpoint: endpoint, + Bucket: bucket, + Prefix: prefix, + PathStyle: pathStyle, + } + + ctx := context.Background() + + // Clean up any leftover objects from previous runs + existing, err := client.ListAllObjects(ctx) + require.Nil(t, err) + if len(existing) > 0 { + keys := make([]string, len(existing)) + for i, obj := range existing { + // Strip the prefix since DeleteObjects will re-add it + keys[i] = strings.TrimPrefix(obj.Key, prefix+"/") + } + // Batch delete in groups of 1000 + for i := 0; i < len(keys); i += 1000 { + end := i + 1000 + if end > len(keys) { + end = len(keys) + } + err := client.DeleteObjects(ctx, keys[i:end]) + require.Nil(t, err) + } + } + + t.Run("PutGetDelete", func(t *testing.T) { + key := "test-object" + content := "hello from ntfy s3 test" + + // Put + err := client.PutObject(ctx, key, strings.NewReader(content), int64(len(content))) + require.Nil(t, err) + + // Get + reader, size, err := client.GetObject(ctx, key) + require.Nil(t, err) + require.Equal(t, int64(len(content)), size) + data, err := io.ReadAll(reader) + reader.Close() + require.Nil(t, err) + require.Equal(t, content, string(data)) + + // Delete + err = client.DeleteObjects(ctx, []string{key}) + require.Nil(t, err) + + // Get after delete should fail + _, _, err = client.GetObject(ctx, key) + require.Error(t, err) + var errResp *ErrorResponse + require.ErrorAs(t, err, &errResp) + require.Equal(t, 404, errResp.StatusCode) + }) + + t.Run("ListObjects", func(t *testing.T) { + // Use a sub-prefix client for isolation + listClient := &Client{ + AccessKey: accessKey, + SecretKey: secretKey, + Region: region, + Endpoint: endpoint, + Bucket: bucket, + Prefix: prefix + "/list-test", + PathStyle: pathStyle, + } + + // Put 10 objects + for i := 0; i < 10; i++ { + err := listClient.PutObject(ctx, fmt.Sprintf("%d", i), strings.NewReader("x"), 1) + require.Nil(t, err) + } + + // List + objects, err := listClient.ListAllObjects(ctx) + require.Nil(t, err) + require.Len(t, objects, 10) + + // Clean up + keys := make([]string, 10) + for i := range keys { + keys[i] = fmt.Sprintf("%d", i) + } + err = listClient.DeleteObjects(ctx, keys) + require.Nil(t, err) + }) + + t.Run("LargeObject", func(t *testing.T) { + key := "large-object" + data := make([]byte, 5*1024*1024) // 5 MB + for i := range data { + data[i] = byte(i % 256) + } + + err := client.PutObject(ctx, key, bytes.NewReader(data), int64(len(data))) + require.Nil(t, err) + + reader, size, err := client.GetObject(ctx, key) + require.Nil(t, err) + require.Equal(t, int64(len(data)), size) + got, err := io.ReadAll(reader) + reader.Close() + require.Nil(t, err) + require.Equal(t, data, got) + + err = client.DeleteObjects(ctx, []string{key}) + require.Nil(t, err) + }) +} diff --git a/s3/types.go b/s3/types.go new file mode 100644 index 00000000..5929ec6c --- /dev/null +++ b/s3/types.go @@ -0,0 +1,65 @@ +package s3 + +import "fmt" + +// Config holds the parsed fields from an S3 URL. Use ParseURL to create one from a URL string. +type Config struct { + Endpoint string // host[:port] only, e.g. "s3.us-east-1.amazonaws.com" + PathStyle bool + Bucket string + Prefix string + Region string + AccessKey string + SecretKey string +} + +// Object represents an S3 object returned by list operations. +type Object struct { + Key string + Size int64 +} + +// ListResult holds the response from a ListObjectsV2 call. +type ListResult struct { + Objects []Object + IsTruncated bool + NextContinuationToken string +} + +// ErrorResponse is returned when S3 responds with a non-2xx status code. +type ErrorResponse struct { + StatusCode int + Code string `xml:"Code"` + Message string `xml:"Message"` + Body string `xml:"-"` // raw response body +} + +func (e *ErrorResponse) Error() string { + if e.Code != "" { + return fmt.Sprintf("s3: %s (HTTP %d): %s", e.Code, e.StatusCode, e.Message) + } + return fmt.Sprintf("s3: HTTP %d: %s", e.StatusCode, e.Body) +} + +// listObjectsV2Response is the XML response from S3 ListObjectsV2 +type listObjectsV2Response struct { + Contents []listObject `xml:"Contents"` + IsTruncated bool `xml:"IsTruncated"` + NextContinuationToken string `xml:"NextContinuationToken"` +} + +type listObject struct { + Key string `xml:"Key"` + Size int64 `xml:"Size"` +} + +// deleteResult is the XML response from S3 DeleteObjects +type deleteResult struct { + Errors []deleteError `xml:"Error"` +} + +type deleteError struct { + Key string `xml:"Key"` + Code string `xml:"Code"` + Message string `xml:"Message"` +} diff --git a/s3/util.go b/s3/util.go new file mode 100644 index 00000000..cf9d4ba8 --- /dev/null +++ b/s3/util.go @@ -0,0 +1,161 @@ +package s3 + +import ( + "crypto/hmac" + "crypto/sha256" + "encoding/hex" + "encoding/xml" + "fmt" + "io" + "net/http" + "net/url" + "sort" + "strings" +) + +const ( + // SHA-256 hash of the empty string, used as the payload hash for bodiless requests + emptyPayloadHash = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855" + + // Sent as the payload hash for streaming uploads where the body is not buffered in memory + unsignedPayload = "UNSIGNED-PAYLOAD" + + // maxResponseBytes caps the size of S3 response bodies we read into memory (10 MB) + maxResponseBytes = 10 * 1024 * 1024 +) + +// ParseURL parses an S3 URL of the form: +// +// s3://ACCESS_KEY:SECRET_KEY@BUCKET[/PREFIX]?region=REGION[&endpoint=ENDPOINT] +// +// When endpoint is specified, path-style addressing is enabled automatically. +func ParseURL(s3URL string) (*Config, error) { + u, err := url.Parse(s3URL) + if err != nil { + return nil, fmt.Errorf("s3: invalid URL: %w", err) + } + if u.Scheme != "s3" { + return nil, fmt.Errorf("s3: URL scheme must be 's3', got '%s'", u.Scheme) + } + if u.Host == "" { + return nil, fmt.Errorf("s3: bucket name must be specified as host") + } + bucket := u.Host + prefix := strings.TrimPrefix(u.Path, "/") + accessKey := u.User.Username() + secretKey, _ := u.User.Password() + if accessKey == "" || secretKey == "" { + return nil, fmt.Errorf("s3: access key and secret key must be specified in URL") + } + region := u.Query().Get("region") + if region == "" { + return nil, fmt.Errorf("s3: region query parameter is required") + } + endpointParam := u.Query().Get("endpoint") + var endpoint string + var pathStyle bool + if endpointParam != "" { + // Custom endpoint: strip scheme prefix to extract host[:port] + ep := strings.TrimRight(endpointParam, "/") + ep = strings.TrimPrefix(ep, "https://") + ep = strings.TrimPrefix(ep, "http://") + endpoint = ep + pathStyle = true + } else { + endpoint = fmt.Sprintf("s3.%s.amazonaws.com", region) + pathStyle = false + } + return &Config{ + Endpoint: endpoint, + PathStyle: pathStyle, + Bucket: bucket, + Prefix: prefix, + Region: region, + AccessKey: accessKey, + SecretKey: secretKey, + }, nil +} + +// parseError reads an S3 error response and returns an *ErrorResponse. +func parseError(resp *http.Response) error { + body, err := io.ReadAll(io.LimitReader(resp.Body, maxResponseBytes)) + if err != nil { + return fmt.Errorf("s3: reading error response: %w", err) + } + return parseErrorFromBytes(resp.StatusCode, body) +} + +func parseErrorFromBytes(statusCode int, body []byte) error { + errResp := &ErrorResponse{ + StatusCode: statusCode, + Body: string(body), + } + // Try to parse XML error; if it fails, we still have StatusCode and Body + _ = xml.Unmarshal(body, errResp) + return errResp +} + +// canonicalURI returns the URI-encoded path for the canonical request. Each path segment is +// percent-encoded per RFC 3986; forward slashes are preserved. +func canonicalURI(u *url.URL) string { + p := u.Path + if p == "" { + return "/" + } + segments := strings.Split(p, "/") + for i, seg := range segments { + segments[i] = uriEncode(seg) + } + return strings.Join(segments, "/") +} + +// canonicalQueryString builds the query string for the canonical request. Keys and values +// are URI-encoded per RFC 3986 (using %20, not +) and sorted lexically by key. +func canonicalQueryString(values url.Values) string { + if len(values) == 0 { + return "" + } + keys := make([]string, 0, len(values)) + for k := range values { + keys = append(keys, k) + } + sort.Strings(keys) + var pairs []string + for _, k := range keys { + ek := uriEncode(k) + vs := make([]string, len(values[k])) + copy(vs, values[k]) + sort.Strings(vs) + for _, v := range vs { + pairs = append(pairs, ek+"="+uriEncode(v)) + } + } + return strings.Join(pairs, "&") +} + +// uriEncode percent-encodes a string per RFC 3986, encoding everything except unreserved +// characters (A-Z a-z 0-9 - _ . ~). +func uriEncode(s string) string { + var buf strings.Builder + for i := 0; i < len(s); i++ { + b := s[i] + if (b >= 'A' && b <= 'Z') || (b >= 'a' && b <= 'z') || (b >= '0' && b <= '9') || + b == '-' || b == '_' || b == '.' || b == '~' { + buf.WriteByte(b) + } else { + fmt.Fprintf(&buf, "%%%02X", b) + } + } + return buf.String() +} + +func sha256Hex(data []byte) string { + h := sha256.Sum256(data) + return hex.EncodeToString(h[:]) +} + +func hmacSHA256(key, data []byte) []byte { + h := hmac.New(sha256.New, key) + h.Write(data) + return h.Sum(nil) +} diff --git a/s3/util_test.go b/s3/util_test.go new file mode 100644 index 00000000..d30c5664 --- /dev/null +++ b/s3/util_test.go @@ -0,0 +1,181 @@ +package s3 + +import ( + "net/http" + "net/url" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestURIEncode(t *testing.T) { + // Unreserved characters are not encoded + require.Equal(t, "abcdefghijklmnopqrstuvwxyz", uriEncode("abcdefghijklmnopqrstuvwxyz")) + require.Equal(t, "ABCDEFGHIJKLMNOPQRSTUVWXYZ", uriEncode("ABCDEFGHIJKLMNOPQRSTUVWXYZ")) + require.Equal(t, "0123456789", uriEncode("0123456789")) + require.Equal(t, "-_.~", uriEncode("-_.~")) + + // Spaces use %20, not + + require.Equal(t, "hello%20world", uriEncode("hello world")) + + // Slashes are encoded (canonicalURI handles slash splitting separately) + require.Equal(t, "a%2Fb", uriEncode("a/b")) + + // Special characters + require.Equal(t, "%2B", uriEncode("+")) + require.Equal(t, "%3D", uriEncode("=")) + require.Equal(t, "%26", uriEncode("&")) + require.Equal(t, "%40", uriEncode("@")) + require.Equal(t, "%23", uriEncode("#")) + + // Mixed + require.Equal(t, "test~file-name_1.txt", uriEncode("test~file-name_1.txt")) + require.Equal(t, "key%20with%20spaces%2Fand%2Fslashes", uriEncode("key with spaces/and/slashes")) + + // Empty string + require.Equal(t, "", uriEncode("")) +} + +func TestCanonicalURI(t *testing.T) { + // Simple path + u, _ := url.Parse("https://example.com/bucket/key") + require.Equal(t, "/bucket/key", canonicalURI(u)) + + // Root path + u, _ = url.Parse("https://example.com/") + require.Equal(t, "/", canonicalURI(u)) + + // Empty path + u, _ = url.Parse("https://example.com") + require.Equal(t, "/", canonicalURI(u)) + + // Path with special characters + u, _ = url.Parse("https://example.com/bucket/key%20with%20spaces") + require.Equal(t, "/bucket/key%20with%20spaces", canonicalURI(u)) + + // Nested path + u, _ = url.Parse("https://example.com/bucket/a/b/c/file.txt") + require.Equal(t, "/bucket/a/b/c/file.txt", canonicalURI(u)) +} + +func TestCanonicalQueryString(t *testing.T) { + // Multiple keys sorted alphabetically + vals := url.Values{ + "prefix": {"test/"}, + "list-type": {"2"}, + } + require.Equal(t, "list-type=2&prefix=test%2F", canonicalQueryString(vals)) + + // Empty values + require.Equal(t, "", canonicalQueryString(url.Values{})) + + // Single key + require.Equal(t, "key=value", canonicalQueryString(url.Values{"key": {"value"}})) + + // Key with multiple values (sorted) + vals = url.Values{"key": {"b", "a"}} + require.Equal(t, "key=a&key=b", canonicalQueryString(vals)) + + // Keys requiring encoding + vals = url.Values{"continuation-token": {"abc+def"}} + require.Equal(t, "continuation-token=abc%2Bdef", canonicalQueryString(vals)) +} + +func TestSHA256Hex(t *testing.T) { + // SHA-256 of empty string + require.Equal(t, emptyPayloadHash, sha256Hex([]byte(""))) + + // SHA-256 of known value + require.Equal(t, "2cf24dba5fb0a30e26e83b2ac5b9e29e1b161e5c1fa7425e73043362938b9824", sha256Hex([]byte("hello"))) +} + +func TestHmacSHA256(t *testing.T) { + // Known test vector: HMAC-SHA256("key", "message") + result := hmacSHA256([]byte("key"), []byte("message")) + require.Len(t, result, 32) // SHA-256 produces 32 bytes + require.NotEqual(t, make([]byte, 32), result) + + // Same inputs should produce same output + result2 := hmacSHA256([]byte("key"), []byte("message")) + require.Equal(t, result, result2) + + // Different inputs should produce different output + result3 := hmacSHA256([]byte("different-key"), []byte("message")) + require.NotEqual(t, result, result3) +} + +func TestSignV4_SetsRequiredHeaders(t *testing.T) { + c := &Client{ + AccessKey: "AKID", + SecretKey: "SECRET", + Region: "us-east-1", + Endpoint: "s3.us-east-1.amazonaws.com", + Bucket: "my-bucket", + } + + req, _ := http.NewRequest(http.MethodGet, "https://my-bucket.s3.us-east-1.amazonaws.com/test-key", nil) + c.signV4(req, emptyPayloadHash) + + // All required SigV4 headers must be set + require.NotEmpty(t, req.Header.Get("Host")) + require.NotEmpty(t, req.Header.Get("X-Amz-Date")) + require.Equal(t, emptyPayloadHash, req.Header.Get("X-Amz-Content-Sha256")) + + // Authorization header must have correct format + auth := req.Header.Get("Authorization") + require.Contains(t, auth, "AWS4-HMAC-SHA256") + require.Contains(t, auth, "Credential=AKID/") + require.Contains(t, auth, "/us-east-1/s3/aws4_request") + require.Contains(t, auth, "SignedHeaders=") + require.Contains(t, auth, "Signature=") +} + +func TestSignV4_UnsignedPayload(t *testing.T) { + c := &Client{ + AccessKey: "AKID", + SecretKey: "SECRET", + Region: "us-east-1", + Endpoint: "s3.us-east-1.amazonaws.com", + Bucket: "my-bucket", + } + + req, _ := http.NewRequest(http.MethodPut, "https://my-bucket.s3.us-east-1.amazonaws.com/test-key", nil) + c.signV4(req, unsignedPayload) + + require.Equal(t, unsignedPayload, req.Header.Get("X-Amz-Content-Sha256")) +} + +func TestSignV4_DifferentRegions(t *testing.T) { + c1 := &Client{AccessKey: "AKID", SecretKey: "SECRET", Region: "us-east-1", Endpoint: "s3.us-east-1.amazonaws.com", Bucket: "b"} + c2 := &Client{AccessKey: "AKID", SecretKey: "SECRET", Region: "eu-west-1", Endpoint: "s3.eu-west-1.amazonaws.com", Bucket: "b"} + + req1, _ := http.NewRequest(http.MethodGet, "https://b.s3.us-east-1.amazonaws.com/key", nil) + c1.signV4(req1, emptyPayloadHash) + + req2, _ := http.NewRequest(http.MethodGet, "https://b.s3.eu-west-1.amazonaws.com/key", nil) + c2.signV4(req2, emptyPayloadHash) + + // Different regions should produce different signatures + require.NotEqual(t, req1.Header.Get("Authorization"), req2.Header.Get("Authorization")) +} + +func TestParseError_XMLResponse(t *testing.T) { + xmlBody := []byte(`NoSuchKeyThe specified key does not exist.`) + err := parseErrorFromBytes(404, xmlBody) + + var errResp *ErrorResponse + require.ErrorAs(t, err, &errResp) + require.Equal(t, 404, errResp.StatusCode) + require.Equal(t, "NoSuchKey", errResp.Code) + require.Equal(t, "The specified key does not exist.", errResp.Message) +} + +func TestParseError_NonXMLResponse(t *testing.T) { + err := parseErrorFromBytes(500, []byte("internal server error")) + + var errResp *ErrorResponse + require.ErrorAs(t, err, &errResp) + require.Equal(t, 500, errResp.StatusCode) + require.Equal(t, "", errResp.Code) // XML parsing failed, no code + require.Contains(t, errResp.Body, "internal server error") +} diff --git a/tools/s3cli/main.go b/tools/s3cli/main.go new file mode 100644 index 00000000..697d4e71 --- /dev/null +++ b/tools/s3cli/main.go @@ -0,0 +1,164 @@ +// Command s3cli is a minimal CLI for testing the s3 package. It supports put, get, rm, and ls. +// +// Usage: +// +// export S3_URL="s3://ACCESS_KEY:SECRET_KEY@BUCKET/PREFIX?region=REGION&endpoint=ENDPOINT" +// +// s3cli put Upload a file +// s3cli put - Upload from stdin +// s3cli get Download to stdout +// s3cli rm [...] Delete one or more objects +// s3cli ls List all objects +package main + +import ( + "context" + "fmt" + "io" + "os" + "text/tabwriter" + + "heckel.io/ntfy/v2/s3" +) + +func main() { + if len(os.Args) < 2 { + usage() + } + s3URL := os.Getenv("S3_URL") + if s3URL == "" { + fail("S3_URL environment variable is required") + } + cfg, err := s3.ParseURL(s3URL) + if err != nil { + fail("invalid S3_URL: %s", err) + } + client := s3.New(cfg) + ctx := context.Background() + + switch os.Args[1] { + case "put": + cmdPut(ctx, client) + case "get": + cmdGet(ctx, client) + case "rm": + cmdRm(ctx, client) + case "ls": + cmdLs(ctx, client) + default: + usage() + } +} + +func cmdPut(ctx context.Context, client *s3.Client) { + if len(os.Args) != 4 { + fail("usage: s3cli put \n") + } + key := os.Args[2] + path := os.Args[3] + + var r io.Reader + var size int64 + if path == "-" { + // Read stdin into a temp file to get the size + tmp, err := os.CreateTemp("", "s3cli-*") + if err != nil { + fail("create temp file: %s", err) + } + defer os.Remove(tmp.Name()) + n, err := io.Copy(tmp, os.Stdin) + if err != nil { + tmp.Close() + fail("read stdin: %s", err) + } + if _, err := tmp.Seek(0, io.SeekStart); err != nil { + tmp.Close() + fail("seek: %s", err) + } + r = tmp + size = n + defer tmp.Close() + } else { + f, err := os.Open(path) + if err != nil { + fail("open %s: %s", path, err) + } + defer f.Close() + info, err := f.Stat() + if err != nil { + fail("stat %s: %s", path, err) + } + r = f + size = info.Size() + } + + if err := client.PutObject(ctx, key, r, size); err != nil { + fail("put: %s", err) + } + fmt.Fprintf(os.Stderr, "uploaded %s (%d bytes)\n", key, size) +} + +func cmdGet(ctx context.Context, client *s3.Client) { + if len(os.Args) != 3 { + fail("usage: s3cli get \n") + } + key := os.Args[2] + + reader, size, err := client.GetObject(ctx, key) + if err != nil { + fail("get: %s", err) + } + defer reader.Close() + n, err := io.Copy(os.Stdout, reader) + if err != nil { + fail("read: %s", err) + } + fmt.Fprintf(os.Stderr, "downloaded %s (%d bytes, content-length: %d)\n", key, n, size) +} + +func cmdRm(ctx context.Context, client *s3.Client) { + if len(os.Args) < 3 { + fail("usage: s3cli rm [...]\n") + } + keys := os.Args[2:] + if err := client.DeleteObjects(ctx, keys); err != nil { + fail("rm: %s", err) + } + fmt.Fprintf(os.Stderr, "deleted %d object(s)\n", len(keys)) +} + +func cmdLs(ctx context.Context, client *s3.Client) { + objects, err := client.ListAllObjects(ctx) + if err != nil { + fail("ls: %s", err) + } + w := tabwriter.NewWriter(os.Stdout, 0, 0, 2, ' ', 0) + var totalSize int64 + for _, obj := range objects { + fmt.Fprintf(w, "%d\t%s\n", obj.Size, obj.Key) + totalSize += obj.Size + } + w.Flush() + fmt.Fprintf(os.Stderr, "%d object(s), %d bytes total\n", len(objects), totalSize) +} + +func usage() { + fmt.Fprintf(os.Stderr, `Usage: s3cli [args...] + +Commands: + put Upload a file (use - for stdin) + get Download to stdout + rm [keys...] Delete objects + ls List all objects + +Environment: + S3_URL S3 connection URL (required) + s3://ACCESS_KEY:SECRET_KEY@BUCKET[/PREFIX]?region=REGION[&endpoint=ENDPOINT] +`) + os.Exit(1) +} + +func fail(format string, args ...any) { + fmt.Fprintf(os.Stderr, format+"\n", args...) + os.Exit(1) +}