This commit is contained in:
binwiederhier
2026-03-16 09:48:26 -04:00
parent 4487299a80
commit 790ba243c7
10 changed files with 1917 additions and 226 deletions
+42 -144
View File
@@ -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
}
+252 -46
View File
@@ -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(`<?xml version="1.0" encoding="UTF-8"?><Error><Code>NoSuchKey</Code><Message>The specified key does not exist.</Message></Error>`))
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(`<?xml version="1.0" encoding="UTF-8"?><DeleteResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/"></DeleteResult>`))
}
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"`
}