diff --git a/attachment/backend_s3.go b/attachment/backend_s3.go index 44f946f6..9a2d4bef 100644 --- a/attachment/backend_s3.go +++ b/attachment/backend_s3.go @@ -8,8 +8,6 @@ import ( "heckel.io/ntfy/v2/s3" ) -const deleteBatchSize = 1000 - type s3Backend struct { client *s3.Client } @@ -45,17 +43,7 @@ func (b *s3Backend) List() ([]object, error) { } func (b *s3Backend) Delete(ids ...string) error { - // S3 DeleteObjects supports up to 1000 keys per call - for i := 0; i < len(ids); i += deleteBatchSize { - end := i + deleteBatchSize - if end > len(ids) { - end = len(ids) - } - if err := b.client.DeleteObjects(context.Background(), ids[i:end]); err != nil { - return err - } - } - return nil + return b.client.DeleteObjects(context.Background(), ids) } func (b *s3Backend) DeleteIncomplete(cutoff time.Time) error { diff --git a/attachment/store_s3_test.go b/attachment/store_s3_test.go index a41c6f8b..6615f4e9 100644 --- a/attachment/store_s3_test.go +++ b/attachment/store_s3_test.go @@ -14,20 +14,20 @@ import ( ) func TestS3Store_WriteWithPrefix(t *testing.T) { - s3URL := os.Getenv("NTFY_TEST_ATTACHMENT_S3_URL") + s3URL := os.Getenv("NTFY_TEST_S3_URL") if s3URL == "" { - t.Skip("NTFY_TEST_ATTACHMENT_S3_URL not set") + t.Skip("NTFY_TEST_S3_URL not set") } cfg, err := s3.ParseURL(s3URL) require.Nil(t, err) cfg.Prefix = "test-prefix" client := s3.New(cfg) - deleteAllObjects(client) + deleteAllObjects(t, client) backend := newS3Backend(client) cache, err := newStore(backend, 10*1024, nil) require.Nil(t, err) t.Cleanup(func() { - deleteAllObjects(client) + deleteAllObjects(t, client) cache.Close() }) @@ -47,34 +47,46 @@ func TestS3Store_WriteWithPrefix(t *testing.T) { func newTestRealS3Store(t *testing.T, totalSizeLimit int64) (*Store, *modTimeOverrideBackend) { t.Helper() - s3URL := os.Getenv("NTFY_TEST_ATTACHMENT_S3_URL") + s3URL := os.Getenv("NTFY_TEST_S3_URL") if s3URL == "" { - t.Skip("NTFY_TEST_ATTACHMENT_S3_URL not set") + t.Skip("NTFY_TEST_S3_URL not set") } cfg, err := s3.ParseURL(s3URL) require.Nil(t, err) + if cfg.Prefix != "" { + cfg.Prefix = cfg.Prefix + "/testpkg-attachment" + } else { + cfg.Prefix = "testpkg-attachment" + } client := s3.New(cfg) inner := newS3Backend(client) wrapper := &modTimeOverrideBackend{backend: inner, modTimes: make(map[string]time.Time)} - deleteAllObjects(client) + deleteAllObjects(t, client) store, err := newStore(wrapper, totalSizeLimit, nil) require.Nil(t, err) t.Cleanup(func() { - deleteAllObjects(client) + deleteAllObjects(t, client) store.Close() }) return store, wrapper } -func deleteAllObjects(client *s3.Client) { - objects, _ := client.ListObjectsV2(context.Background()) - keys := make([]string, 0, len(objects)) - for _, obj := range objects { - keys = append(keys, obj.Key) - } - if len(keys) > 0 { - client.DeleteObjects(context.Background(), keys) //nolint:errcheck +func deleteAllObjects(t *testing.T, client *s3.Client) { + t.Helper() + for i := 0; i < 20; i++ { + objects, err := client.ListObjectsV2(context.Background()) + require.Nil(t, err) + if len(objects) == 0 { + return + } + keys := make([]string, len(objects)) + for j, obj := range objects { + keys[j] = obj.Key + } + require.Nil(t, client.DeleteObjects(context.Background(), keys)) + time.Sleep(200 * time.Millisecond) } + t.Fatal("timed out waiting for bucket to be empty") } // modTimeOverrideBackend wraps a backend and allows overriding LastModified times returned by List(). diff --git a/attachment/store_test.go b/attachment/store_test.go index 7ac7cddb..11d0b244 100644 --- a/attachment/store_test.go +++ b/attachment/store_test.go @@ -332,7 +332,7 @@ func TestStore_Sync_SkipsRecentFiles(t *testing.T) { // callback that makes a specific object's timestamp old enough for orphan cleanup (> 1 hour). // For the file backend, this uses os.Chtimes; for the S3 backend, it overrides the object's // LastModified time via a modTimeOverrideBackend wrapper. Objects start with recent timestamps -// by default. The S3 subtest is skipped if NTFY_TEST_ATTACHMENT_S3_URL is not set. +// by default. The S3 subtest is skipped if NTFY_TEST_S3_URL is not set. func forEachBackend(t *testing.T, totalSizeLimit int64, f func(t *testing.T, s *Store, makeOld func(string))) { t.Run("file", func(t *testing.T) { dir, s := newTestFileStore(t, totalSizeLimit) diff --git a/s3/client.go b/s3/client.go index d9ec1ab8..e06ff5c9 100644 --- a/s3/client.go +++ b/s3/client.go @@ -11,7 +11,6 @@ import ( "io" "net/http" "net/url" - "strconv" "strings" "time" @@ -125,7 +124,7 @@ func (c *Client) ListObjectsV2(ctx context.Context) ([]*Object, error) { var all []*Object var token string for page := 0; page < maxPages; page++ { - result, err := c.listObjectsV2(ctx, token, 0) + result, err := c.listObjectsV2(ctx, token) if err != nil { return nil, err } @@ -149,8 +148,7 @@ func (c *Client) ListObjectsV2(ctx context.Context) ([]*Object, error) { } // listObjectsV2 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) listObjectsV2(ctx context.Context, continuationToken string, maxKeys int) (*listObjectsV2Result, error) { +func (c *Client) listObjectsV2(ctx context.Context, continuationToken string) (*listObjectsV2Result, error) { log.Tag(tagS3Client).Debug("Listing remote objects with continuation token '%s'", continuationToken) query := url.Values{"list-type": {"2"}} if prefix := c.config.ListPrefix(); prefix != "" { @@ -159,9 +157,6 @@ func (c *Client) listObjectsV2(ctx context.Context, continuationToken string, ma if continuationToken != "" { query.Set("continuation-token", continuationToken) } - if maxKeys > 0 { - query.Set("max-keys", strconv.Itoa(maxKeys)) - } respBody, err := c.do(ctx, "ListObjects", http.MethodGet, c.config.BucketURL()+"?"+query.Encode(), nil, nil) if err != nil { return nil, err @@ -182,6 +177,20 @@ func (c *Client) listObjectsV2(ctx context.Context, continuationToken string, ma // // See https://docs.aws.amazon.com/AmazonS3/latest/API/API_DeleteObjects.html func (c *Client) DeleteObjects(ctx context.Context, keys []string) error { + // S3 DeleteObjects supports up to 1000 keys per call + for i := 0; i < len(keys); i += maxDeleteBatchSize { + end := i + maxDeleteBatchSize + if end > len(keys) { + end = len(keys) + } + if err := c.deleteObjects(ctx, keys[i:end]); err != nil { + return err + } + } + return nil +} + +func (c *Client) deleteObjects(ctx context.Context, keys []string) error { log.Tag(tagS3Client).Debug("Deleting %d object(s)", len(keys)) req := &deleteObjectsRequest{ Quiet: true, diff --git a/s3/client_test.go b/s3/client_test.go index 652db3e7..f4b85089 100644 --- a/s3/client_test.go +++ b/s3/client_test.go @@ -3,13 +3,9 @@ package s3 import ( "bytes" "context" - "encoding/xml" "fmt" "io" - "net/http" - "net/http/httptest" "os" - "sort" "strings" "sync" "testing" @@ -18,271 +14,6 @@ import ( "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 - uploads map[string]map[int][]byte // uploadID -> partNumber -> data - nextID int // counter for generating upload IDs - mu sync.RWMutex -} - -func newMockS3Server() (*httptest.Server, *mockS3Server) { - m := &mockS3Server{ - objects: make(map[string][]byte), - uploads: make(map[string]map[int][]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, "/") - q := r.URL.Query() - - switch { - case r.Method == http.MethodPut && q.Has("partNumber"): - m.handleUploadPart(w, r, path) - case r.Method == http.MethodPut: - m.handlePut(w, r, path) - case r.Method == http.MethodPost && q.Has("uploads"): - m.handleInitiateMultipart(w, r, path) - case r.Method == http.MethodPost && q.Has("uploadId"): - m.handleCompleteMultipart(w, r, path) - case r.Method == http.MethodDelete && q.Has("uploadId"): - m.handleAbortMultipart(w, r, path) - case r.Method == http.MethodGet && q.Get("list-type") == "2": - m.handleList(w, r, path) - case r.Method == http.MethodGet: - m.handleGet(w, r, path) - case r.Method == http.MethodPost && q.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) handleInitiateMultipart(w http.ResponseWriter, r *http.Request, path string) { - m.mu.Lock() - m.nextID++ - uploadID := fmt.Sprintf("upload-%d", m.nextID) - m.uploads[uploadID] = make(map[int][]byte) - m.mu.Unlock() - - w.Header().Set("Content-Type", "application/xml") - w.WriteHeader(http.StatusOK) - fmt.Fprintf(w, `%s`, uploadID) -} - -func (m *mockS3Server) handleUploadPart(w http.ResponseWriter, r *http.Request, path string) { - uploadID := r.URL.Query().Get("uploadId") - var partNumber int - fmt.Sscanf(r.URL.Query().Get("partNumber"), "%d", &partNumber) - - body, err := io.ReadAll(r.Body) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - - m.mu.Lock() - parts, ok := m.uploads[uploadID] - if !ok { - m.mu.Unlock() - http.Error(w, "NoSuchUpload", http.StatusNotFound) - return - } - parts[partNumber] = body - m.mu.Unlock() - - etag := fmt.Sprintf(`"etag-part-%d"`, partNumber) - w.Header().Set("ETag", etag) - w.WriteHeader(http.StatusOK) -} - -func (m *mockS3Server) handleCompleteMultipart(w http.ResponseWriter, r *http.Request, path string) { - uploadID := r.URL.Query().Get("uploadId") - - m.mu.Lock() - parts, ok := m.uploads[uploadID] - if !ok { - m.mu.Unlock() - http.Error(w, "NoSuchUpload", http.StatusNotFound) - return - } - - // Assemble parts in order - var assembled []byte - for i := 1; i <= len(parts); i++ { - assembled = append(assembled, parts[i]...) - } - m.objects[path] = assembled - delete(m.uploads, uploadID) - m.mu.Unlock() - - w.Header().Set("Content-Type", "application/xml") - w.WriteHeader(http.StatusOK) - fmt.Fprintf(w, `%s`, path) -} - -func (m *mockS3Server) handleAbortMultipart(w http.ResponseWriter, r *http.Request, path string) { - uploadID := r.URL.Query().Get("uploadId") - m.mu.Lock() - delete(m.uploads, uploadID) - m.mu.Unlock() - w.WriteHeader(http.StatusNoContent) -} - -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)), LastModified: time.Now().Format(time.RFC3339)}) - } - 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 New(&Config{ - 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) @@ -409,13 +140,10 @@ func TestConfig_ListPrefix(t *testing.T) { require.Equal(t, "", c2.ListPrefix()) } -// --- Integration tests using mock S3 server --- +// --- Integration tests using real S3 --- func TestClient_PutGetObject(t *testing.T) { - server, _ := newMockS3Server() - defer server.Close() - client := newTestClient(server, "my-bucket", "") - + client := newTestClient(t) ctx := context.Background() // Put @@ -432,152 +160,85 @@ func TestClient_PutGetObject(t *testing.T) { 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"), 0) - 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", "") + client := newTestClient(t) _, _, 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", "") - + client := newTestClient(t) 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")), 0) + err := client.PutObject(ctx, fmt.Sprintf("del-%d", i), bytes.NewReader([]byte("data")), 0) require.Nil(t, err) } - require.Equal(t, 5, mock.objectCount()) + waitForCount(t, client, 5) // Delete some - err := client.DeleteObjects(ctx, []string{"key-1", "key-3"}) + err := client.DeleteObjects(ctx, []string{"del-1", "del-3"}) require.Nil(t, err) - require.Equal(t, 3, mock.objectCount()) + waitForCount(t, client, 3) // Verify deleted ones are gone - _, _, err = client.GetObject(ctx, "key-1") + _, _, err = client.GetObject(ctx, "del-1") require.Error(t, err) - _, _, err = client.GetObject(ctx, "key-3") + _, _, err = client.GetObject(ctx, "del-3") require.Error(t, err) // Verify remaining ones are still there - reader, _, err := client.GetObject(ctx, "key-0") - require.Nil(t, err) - reader.Close() + for _, key := range []string{"del-0", "del-2", "del-4"} { + reader, _, err := client.GetObject(ctx, key) + require.Nil(t, err) + reader.Close() + } } func TestClient_ListObjects(t *testing.T) { - server, _ := newMockS3Server() - defer server.Close() - + client := newTestClient(t) 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")), 0) + err := client.PutObject(ctx, fmt.Sprintf("list-%d", i), bytes.NewReader([]byte("x")), 0) 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")), 0) - require.Nil(t, err) - - // List with prefix client: should only see 3 - result, err := client.listObjectsV2(ctx, "", 0) - require.Nil(t, err) - require.Len(t, result.Contents, 3) - require.False(t, result.IsTruncated) - - // List with no-prefix client: should see all 4 - result, err = clientNoPrefix.listObjectsV2(ctx, "", 0) - require.Nil(t, err) - require.Len(t, result.Contents, 4) + waitForCount(t, client, 3) } func TestClient_ListObjects_Pagination(t *testing.T) { - server, _ := newMockS3Server() - defer server.Close() - client := newTestClient(server, "my-bucket", "") - + client := newTestClient(t) 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")), 0) + // Create 1010 objects in parallel (5 goroutines) + const total = 1010 + const workers = 5 + var wg sync.WaitGroup + errs := make(chan error, total) + for w := 0; w < workers; w++ { + wg.Add(1) + go func(start int) { + defer wg.Done() + for i := start; i < total; i += workers { + if err := client.PutObject(ctx, fmt.Sprintf("pg-%04d", i), bytes.NewReader([]byte("x")), 0); err != nil { + errs <- err + return + } + } + }(w) + } + wg.Wait() + close(errs) + for err := range errs { require.Nil(t, err) } - - // List with max-keys=2 - result, err := client.listObjectsV2(ctx, "", 2) - require.Nil(t, err) - require.Len(t, result.Contents, 2) - require.True(t, result.IsTruncated) - require.NotEmpty(t, result.NextContinuationToken) - - // Get next page - result2, err := client.listObjectsV2(ctx, result.NextContinuationToken, 2) - require.Nil(t, err) - require.Len(t, result2.Contents, 2) - require.True(t, result2.IsTruncated) - - // Get last page - result3, err := client.listObjectsV2(ctx, result2.NextContinuationToken, 2) - require.Nil(t, err) - require.Len(t, result3.Contents, 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")), 0) - require.Nil(t, err) - } - - objects, err := client.ListObjectsV2(ctx) - require.Nil(t, err) - require.Len(t, objects, 10) + waitForCount(t, client, total) } func TestClient_PutObject_LargeBody(t *testing.T) { - server, _ := newMockS3Server() - defer server.Close() - client := newTestClient(server, "my-bucket", "") - + client := newTestClient(t) ctx := context.Background() // 1 MB object @@ -598,10 +259,7 @@ func TestClient_PutObject_LargeBody(t *testing.T) { } func TestClient_PutObject_ChunkedUpload(t *testing.T) { - server, _ := newMockS3Server() - defer server.Close() - client := newTestClient(server, "my-bucket", "") - + client := newTestClient(t) ctx := context.Background() // 12 MB object, exceeds 5 MB partSize, triggers multipart upload path @@ -622,10 +280,7 @@ func TestClient_PutObject_ChunkedUpload(t *testing.T) { } func TestClient_PutObject_ExactPartSize(t *testing.T) { - server, _ := newMockS3Server() - defer server.Close() - client := newTestClient(server, "my-bucket", "") - + client := newTestClient(t) ctx := context.Background() // Exactly 5 MB (partSize), should use the simple put path (ReadFull succeeds fully) @@ -646,10 +301,7 @@ func TestClient_PutObject_ExactPartSize(t *testing.T) { } func TestClient_PutObject_StreamingExactLength(t *testing.T) { - server, _ := newMockS3Server() - defer server.Close() - client := newTestClient(server, "my-bucket", "pfx") - + client := newTestClient(t) ctx := context.Background() // untrustedLength matches body exactly — streams directly via putObject @@ -666,10 +318,7 @@ func TestClient_PutObject_StreamingExactLength(t *testing.T) { } func TestClient_PutObject_StreamingBodyLongerThanClaimed(t *testing.T) { - server, _ := newMockS3Server() - defer server.Close() - client := newTestClient(server, "my-bucket", "pfx") - + client := newTestClient(t) ctx := context.Background() // Body has 11 bytes, but we claim 5 — only first 5 bytes should be stored @@ -686,16 +335,12 @@ func TestClient_PutObject_StreamingBodyLongerThanClaimed(t *testing.T) { } func TestClient_PutObject_StreamingBodyShorterThanClaimed(t *testing.T) { - server, _ := newMockS3Server() - defer server.Close() - client := newTestClient(server, "my-bucket", "pfx") - + client := newTestClient(t) ctx := context.Background() // Body has 5 bytes, but we claim 100 — should fail err := client.PutObject(ctx, "stream-short", strings.NewReader("hello"), 100) require.Error(t, err) - require.Contains(t, err.Error(), "ContentLength") // Object should not exist _, _, err = client.GetObject(ctx, "stream-short") @@ -703,10 +348,7 @@ func TestClient_PutObject_StreamingBodyShorterThanClaimed(t *testing.T) { } func TestClient_PutObject_NestedKey(t *testing.T) { - server, _ := newMockS3Server() - defer server.Close() - client := newTestClient(server, "my-bucket", "") - + client := newTestClient(t) ctx := context.Background() err := client.PutObject(ctx, "deep/nested/prefix/file.txt", strings.NewReader("nested"), 0) @@ -719,199 +361,54 @@ func TestClient_PutObject_NestedKey(t *testing.T) { 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") +func newTestClient(t *testing.T) *Client { + t.Helper() + s3URL := os.Getenv("NTFY_TEST_S3_URL") + if s3URL == "" { + t.Skip("NTFY_TEST_S3_URL not set") } - - 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")), 0) - require.Nil(t, err) - } - } - - // List all 20k objects with pagination - objects, err := client.ListObjectsV2(ctx) + cfg, err := ParseURL(s3URL) require.Nil(t, err) - require.Len(t, objects, numObjects) - - // Verify total size - var totalSize int64 - for _, obj := range objects { - totalSize += obj.Size + // Use per-test prefix to isolate objects between tests + if cfg.Prefix != "" { + cfg.Prefix = cfg.Prefix + "/testpkg-s3/" + t.Name() + } else { + cfg.Prefix = "testpkg-s3/" + t.Name() } - 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.ListObjectsV2(ctx) - require.Nil(t, err) - require.Len(t, objects, numObjects-1000) + client := New(cfg) + deleteAllObjects(t, client) + t.Cleanup(func() { deleteAllObjects(t, client) }) + return client } -// --- 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") +func deleteAllObjects(t *testing.T, client *Client) { + t.Helper() + for i := 0; i < 20; i++ { + objects, err := client.ListObjectsV2(context.Background()) + require.Nil(t, err) + if len(objects) == 0 { + return + } + keys := make([]string, len(objects)) + for j, obj := range objects { + keys[j] = obj.Key + } + require.Nil(t, client.DeleteObjects(context.Background(), keys)) + time.Sleep(200 * time.Millisecond) } + t.Fatal("timed out waiting for bucket to be empty") +} - 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 := New(&Config{ - 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.ListObjectsV2(ctx) - require.Nil(t, err) - if len(existing) > 0 { - keys := make([]string, len(existing)) - for i, obj := range existing { - keys[i] = obj.Key - } - // 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) +func waitForCount(t *testing.T, client *Client, expected int) { + t.Helper() + for i := 0; i < 20; i++ { + objects, err := client.ListObjectsV2(context.Background()) + require.Nil(t, err) + if len(objects) == expected { + return } + time.Sleep(200 * time.Millisecond) } - - 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), 0) - 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 := New(&Config{ - 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"), 0) - require.Nil(t, err) - } - - // List - objects, err := listClient.ListObjectsV2(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), 0) - 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) - }) + objects, _ := client.ListObjectsV2(context.Background()) + t.Fatalf("timed out waiting for %d objects, got %d", expected, len(objects)) } diff --git a/s3/util.go b/s3/util.go index 1f4c2dd9..ae692735 100644 --- a/s3/util.go +++ b/s3/util.go @@ -34,6 +34,9 @@ const ( // maxPages is the max number of pages to iterate through when listing objects maxPages = 500 + + // maxDeleteBatchSize is the maximum number of keys per S3 DeleteObjects call + maxDeleteBatchSize = 1000 ) // ParseURL parses an S3 URL of the form: