memos/internal/storage/s3/s3.go

227 lines
7.9 KiB
Go

package s3
import (
"context"
"crypto/tls"
"fmt"
"io"
"net/http"
"github.com/aws/aws-sdk-go-v2/aws"
v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4" //nolint:revive // goimports insists on aliasing versioned import paths
awshttp "github.com/aws/aws-sdk-go-v2/aws/transport/http"
"github.com/aws/aws-sdk-go-v2/config"
"github.com/aws/aws-sdk-go-v2/credentials"
"github.com/aws/aws-sdk-go-v2/service/s3"
"github.com/aws/smithy-go"
"github.com/aws/smithy-go/middleware"
smithyhttp "github.com/aws/smithy-go/transport/http"
"github.com/pkg/errors"
storepb "github.com/usememos/memos/proto/gen/store"
)
// ErrRangeNotSatisfiable reports a ranged read whose byte range falls outside
// the object, so HTTP handlers can answer 416 instead of 500.
var ErrRangeNotSatisfiable = errors.New("requested range not satisfiable")
// RangeNotSatisfiableError carries response metadata for an unsatisfied range.
type RangeNotSatisfiableError struct {
ContentRange string
}
func (*RangeNotSatisfiableError) Error() string {
return ErrRangeNotSatisfiable.Error()
}
func (*RangeNotSatisfiableError) Unwrap() error {
return ErrRangeNotSatisfiable
}
// ObjectStream is object content with the metadata needed to answer HTTP
// range requests.
type ObjectStream struct {
Body io.ReadCloser
// ContentLength is the number of bytes in Body, or -1 when unknown.
ContentLength int64
// ContentRange echoes the backend's Content-Range header for partial reads;
// empty when the whole object is returned.
ContentRange string
}
// Driver stores attachment objects in an S3-compatible object store.
type Driver struct {
Client *s3.Client
Bucket *string
}
// NewDriver creates an S3 storage driver from the supplied configuration.
func NewDriver(ctx context.Context, s3Config *storepb.StorageS3Config) (*Driver, error) {
loadOptions := []func(*config.LoadOptions) error{
config.WithCredentialsProvider(credentials.NewStaticCredentialsProvider(s3Config.AccessKeyId, s3Config.AccessKeySecret, "")),
config.WithRegion(s3Config.Region),
}
if s3Config.InsecureSkipTlsVerify {
// Skip TLS certificate verification for endpoints using self-signed certificates.
// This is opt-in and removes protection against man-in-the-middle attacks.
httpClient := awshttp.NewBuildableClient().WithTransportOptions(func(tr *http.Transport) {
tr.TLSClientConfig = &tls.Config{InsecureSkipVerify: true} // #nosec G402 -- opt-in for self-signed S3 endpoints
})
loadOptions = append(loadOptions, config.WithHTTPClient(httpClient))
}
cfg, err := config.LoadDefaultConfig(ctx, loadOptions...)
if err != nil {
return nil, errors.Wrap(err, "failed to load s3 config")
}
client := s3.NewFromConfig(cfg, func(o *s3.Options) {
o.BaseEndpoint = aws.String(s3Config.Endpoint)
o.UsePathStyle = s3Config.UsePathStyle
o.RequestChecksumCalculation = aws.RequestChecksumCalculationWhenRequired
o.ResponseChecksumValidation = aws.ResponseChecksumValidationWhenRequired
o.APIOptions = append(o.APIOptions, excludeAcceptEncodingFromSigning, forceSignedPayload)
})
return &Driver{
Client: client,
Bucket: aws.String(s3Config.Bucket),
}, nil
}
type acceptEncodingKey struct{}
type dropAcceptEncoding struct{}
func (dropAcceptEncoding) ID() string { return "MemosDropAcceptEncoding" }
func (dropAcceptEncoding) HandleFinalize(ctx context.Context, in middleware.FinalizeInput, next middleware.FinalizeHandler) (middleware.FinalizeOutput, middleware.Metadata, error) {
if req, ok := in.Request.(*smithyhttp.Request); ok {
if values := req.Header.Values("Accept-Encoding"); len(values) > 0 {
ctx = context.WithValue(ctx, acceptEncodingKey{}, values)
req.Header.Del("Accept-Encoding")
}
}
return next.HandleFinalize(ctx, in)
}
type restoreAcceptEncoding struct{}
func (restoreAcceptEncoding) ID() string { return "MemosRestoreAcceptEncoding" }
func (restoreAcceptEncoding) HandleFinalize(ctx context.Context, in middleware.FinalizeInput, next middleware.FinalizeHandler) (middleware.FinalizeOutput, middleware.Metadata, error) {
if req, ok := in.Request.(*smithyhttp.Request); ok {
if values, ok := ctx.Value(acceptEncodingKey{}).([]string); ok {
for _, value := range values {
req.Header.Add("Accept-Encoding", value)
}
}
}
return next.HandleFinalize(ctx, in)
}
// excludeAcceptEncodingFromSigning keeps the Accept-Encoding header out of the
// SigV4 signature while still sending it on the wire. Some S3-compatible
// providers (notably Google Cloud Storage) rewrite the header in transit, so a
// signature covering it never verifies and requests fail with
// SignatureDoesNotMatch.
func excludeAcceptEncodingFromSigning(stack *middleware.Stack) error {
if err := stack.Finalize.Insert(dropAcceptEncoding{}, "Signing", middleware.Before); err != nil {
return err
}
return stack.Finalize.Insert(restoreAcceptEncoding{}, "Signing", middleware.After)
}
// forceSignedPayload signs request payloads with their real SHA-256 instead of
// the UNSIGNED-PAYLOAD marker the SDK uses over TLS, which some S3-compatible
// providers (notably Google Cloud Storage) reject.
func forceSignedPayload(stack *middleware.Stack) error {
_, err := stack.Finalize.Swap((*v4.ComputePayloadSHA256)(nil).ID(), &v4.ComputePayloadSHA256{})
return err
}
// UploadObject uploads an object to S3.
func (c *Driver) UploadObject(ctx context.Context, key string, fileType string, content io.Reader) (string, error) {
putInput := s3.PutObjectInput{
Bucket: c.Bucket,
Key: aws.String(key),
ContentType: aws.String(fileType),
Body: content,
}
if _, err := c.Client.PutObject(ctx, &putInput); err != nil {
return "", errors.Wrap(err, "failed to upload object")
}
return key, nil
}
// GetObject retrieves an object from S3.
func (c *Driver) GetObject(ctx context.Context, key string) ([]byte, error) {
stream, err := c.GetObjectStream(ctx, key, "")
if err != nil {
return nil, err
}
defer stream.Body.Close()
data, err := io.ReadAll(stream.Body)
if err != nil {
return nil, errors.Wrap(err, "failed to read object body")
}
return data, nil
}
// GetObjectStream retrieves an object from S3 as a stream. A non-empty
// byteRange is forwarded as an HTTP Range header (e.g. "bytes=0-1023") and
// yields a partial object with ContentRange set. Callers must supply at most
// one range because S3 does not support multipart range responses.
func (c *Driver) GetObjectStream(ctx context.Context, key string, byteRange string) (*ObjectStream, error) {
input := &s3.GetObjectInput{
Bucket: c.Bucket,
Key: aws.String(key),
}
if byteRange != "" {
input.Range = aws.String(byteRange)
}
output, err := c.Client.GetObject(ctx, input)
if err != nil {
var apiErr smithy.APIError
if errors.As(err, &apiErr) && apiErr.ErrorCode() == "InvalidRange" {
rangeErr := &RangeNotSatisfiableError{}
var responseErr *smithyhttp.ResponseError
if errors.As(err, &responseErr) && responseErr.Response != nil && responseErr.Response.Response != nil {
rangeErr.ContentRange = responseErr.Response.Header.Get("Content-Range")
}
if rangeErr.ContentRange == "" {
head, headErr := c.Client.HeadObject(ctx, &s3.HeadObjectInput{
Bucket: c.Bucket,
Key: aws.String(key),
})
if headErr == nil && head.ContentLength != nil {
rangeErr.ContentRange = fmt.Sprintf("bytes */%d", *head.ContentLength)
}
}
return nil, rangeErr
}
return nil, errors.Wrap(err, "failed to get object")
}
stream := &ObjectStream{
Body: output.Body,
ContentLength: -1,
}
if output.ContentLength != nil {
stream.ContentLength = *output.ContentLength
}
if output.ContentRange != nil {
stream.ContentRange = *output.ContentRange
}
return stream, nil
}
// DeleteObject deletes an object in S3.
func (c *Driver) DeleteObject(ctx context.Context, key string) error {
_, err := c.Client.DeleteObject(ctx, &s3.DeleteObjectInput{
Bucket: c.Bucket,
Key: aws.String(key),
})
if err != nil {
return errors.Wrap(err, "failed to delete object")
}
return nil
}