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("")
+ }
+ 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)
+}