Skip to content

Commit 86a452f

Browse files
committed
Add gated S3 mutations to edge gateway
1 parent bb173cc commit 86a452f

4 files changed

Lines changed: 595 additions & 48 deletions

File tree

cmd/edge-gateway/main.go

Lines changed: 190 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,7 @@ import (
3434
"github.com/terion-name/air3/internal/s3fetch"
3535
"github.com/terion-name/air3/internal/signing"
3636
"github.com/terion-name/air3/internal/tickets"
37+
"github.com/terion-name/air3/internal/uploadsource"
3738
)
3839

3940
type ticketPublisher interface {
@@ -47,6 +48,7 @@ type objectFetcher interface {
4748
type edgeServer struct {
4849
cfg config.EdgeConfig
4950
registry *pending.Registry
51+
uploadSources *uploadsource.Registry
5052
publisher ticketPublisher
5153
directFetchers map[string]objectFetcher
5254
logger *slog.Logger
@@ -84,11 +86,17 @@ func run(ctx context.Context, logger *slog.Logger) error {
8486
}
8587

8688
reg := pending.NewRegistry(pending.Options{})
87-
edge := newEdgeServer(cfg, reg, publisher, logger, directFetchers)
89+
uploadReg := uploadsource.NewRegistry(uploadsource.Options{})
90+
edge := newEdgeServer(cfg, reg, uploadReg, publisher, logger, directFetchers)
8891
ingestHandler, err := ingest.NewHandler(ingest.Options{Registry: reg, AllowedConnectorIdentities: cfg.AllowedConnectorIdentities, StreamCopyBufferBytes: cfg.StreamCopyBufferBytes})
8992
if err != nil {
9093
return err
9194
}
95+
uploadHandler, err := uploadsource.NewHandler(uploadsource.HandlerOptions{Registry: uploadReg, AllowedConnectorIdentities: cfg.AllowedConnectorIdentities, StreamCopyBufferBytes: cfg.StreamCopyBufferBytes})
96+
if err != nil {
97+
return err
98+
}
99+
privateHandler := newPrivateIngestHandler(ingestHandler, uploadHandler)
92100

93101
publicServer := &http.Server{Addr: cfg.PublicListenAddr, Handler: edge}
94102
var tlsCfg *tls.Config
@@ -101,7 +109,7 @@ func run(ctx context.Context, logger *slog.Logger) error {
101109
if tlsConfigured(cfg.MTLS) {
102110
publicServer.TLSConfig = publicTLSConfig(tlsCfg)
103111
}
104-
ingestServer := newIngestHTTPServer(cfg, ingestHandler, tlsCfg)
112+
ingestServer := newIngestHTTPServer(cfg, privateHandler, tlsCfg)
105113
ingestListener, err := newNonHTTPIngestListener(cfg, reg, tlsCfg)
106114
if err != nil {
107115
return err
@@ -160,15 +168,22 @@ func newDirectFetchers(ctx context.Context, directServers map[string]config.S3Co
160168
return fetchers, nil
161169
}
162170

163-
func newEdgeServer(cfg config.EdgeConfig, reg *pending.Registry, publisher ticketPublisher, logger *slog.Logger, directFetchers ...map[string]objectFetcher) *edgeServer {
171+
func newEdgeServer(cfg config.EdgeConfig, reg *pending.Registry, uploadSources *uploadsource.Registry, publisher ticketPublisher, logger *slog.Logger, directFetchers ...map[string]objectFetcher) *edgeServer {
164172
if logger == nil {
165173
logger = slog.New(slog.NewTextHandler(io.Discard, nil))
166174
}
167175
var fetchers map[string]objectFetcher
168176
if len(directFetchers) > 0 {
169177
fetchers = directFetchers[0]
170178
}
171-
return &edgeServer{cfg: cfg, registry: reg, publisher: publisher, directFetchers: fetchers, logger: logger, now: time.Now, newToken: randomToken}
179+
return &edgeServer{cfg: cfg, registry: reg, uploadSources: uploadSources, publisher: publisher, directFetchers: fetchers, logger: logger, now: time.Now, newToken: randomToken}
180+
}
181+
182+
func newPrivateIngestHandler(ingestHandler, uploadHandler http.Handler) http.Handler {
183+
mux := http.NewServeMux()
184+
mux.Handle(ingest.PathPrefix, ingestHandler)
185+
mux.Handle(uploadsource.PathPrefix, uploadHandler)
186+
return mux
172187
}
173188

174189
func (s *edgeServer) ServeHTTP(w http.ResponseWriter, r *http.Request) {
@@ -255,13 +270,15 @@ func (s *edgeServer) ServeHTTP(w http.ResponseWriter, r *http.Request) {
255270
}
256271

257272
type s3EdgeRequest struct {
258-
server string
259-
bucket string
260-
key string
261-
rangeHeader string
262-
operation tickets.Operation
263-
list *tickets.ListRequest
264-
headBucket bool
273+
server string
274+
bucket string
275+
key string
276+
rangeHeader string
277+
operation tickets.Operation
278+
list *tickets.ListRequest
279+
headBucket bool
280+
contentLength *int64
281+
contentType string
265282
}
266283

267284
func edgeS3APIMode(cfg config.EdgeConfig) s3api.RoutingMode {
@@ -307,6 +324,24 @@ func (s *edgeServer) serveS3API(w http.ResponseWriter, r *http.Request) {
307324
return
308325
}
309326

327+
ticket := tickets.Ticket{Version: tickets.Version, RequestID: reqID, Bucket: req.bucket, Key: req.key, Method: r.Method, Operation: req.operation, Range: req.rangeHeader, List: req.list, Server: req.server, DeadlineUnixMS: deadline.UnixMilli(), IngestURL: ingestURL, IngestToken: ingestToken, TraceID: reqID}
328+
if req.operation == tickets.OperationPutObject {
329+
uploadToken, err := s.newToken()
330+
if err != nil {
331+
writeS3Error(w, r, s3HTTPError{status: http.StatusInternalServerError, code: "InternalError", message: "Request setup failed"})
332+
return
333+
}
334+
uploadURL, err := uploadsource.URLForRequest(s.cfg.IngestURL, reqID)
335+
if err != nil {
336+
writeS3Error(w, r, s3HTTPError{status: http.StatusInternalServerError, code: "InternalError", message: "Request setup failed"})
337+
return
338+
}
339+
ticket.UploadSourceURL = uploadURL
340+
ticket.UploadToken = uploadToken
341+
ticket.ContentLength = req.contentLength
342+
ticket.ContentType = req.contentType
343+
}
344+
310345
sink := newResponseSink(w, r.Method, r.Context())
311346
pendingReq := pending.Request{ID: reqID, Deadline: deadline, IngestToken: ingestToken, Method: r.Method, Operation: req.operation, Bucket: req.bucket, Key: req.key, Range: req.rangeHeader, List: req.list}
312347
if err := s.registry.Register(pendingReq, sink); err != nil {
@@ -315,17 +350,55 @@ func (s *edgeServer) serveS3API(w http.ResponseWriter, r *http.Request) {
315350
}
316351
defer s.registry.Cancel(reqID, pending.ErrCanceled)
317352

318-
ticket := tickets.Ticket{Version: tickets.Version, RequestID: reqID, Bucket: req.bucket, Key: req.key, Method: r.Method, Operation: req.operation, Range: req.rangeHeader, List: req.list, Server: req.server, DeadlineUnixMS: deadline.UnixMilli(), IngestURL: ingestURL, IngestToken: ingestToken, TraceID: reqID}
353+
uploadRegistered := false
354+
if req.operation == tickets.OperationPutObject {
355+
if s.uploadSources == nil {
356+
s.registry.Cancel(reqID, uploadsource.ErrInvalidSource)
357+
writeS3Error(w, r, s3HTTPError{status: http.StatusInternalServerError, code: "InternalError", message: "Request setup failed"})
358+
return
359+
}
360+
source := uploadsource.Source{RequestID: reqID, Token: ticket.UploadToken, Body: r.Body, ContentLength: req.contentLength, ContentType: req.contentType, Deadline: deadline, Context: r.Context()}
361+
if err := s.uploadSources.Register(source); err != nil {
362+
s.registry.Cancel(reqID, err)
363+
writeS3Error(w, r, s3ErrorForUploadSetup(err))
364+
return
365+
}
366+
uploadRegistered = true
367+
defer func() {
368+
if uploadRegistered {
369+
s.uploadSources.Cancel(reqID, uploadsource.ErrCanceled)
370+
}
371+
}()
372+
}
373+
374+
if _, err := tickets.Marshal(ticket, s.now()); err != nil {
375+
s.registry.Cancel(reqID, err)
376+
if uploadRegistered {
377+
s.uploadSources.Cancel(reqID, err)
378+
uploadRegistered = false
379+
}
380+
writeS3Error(w, r, s3ErrorForSetup(err))
381+
return
382+
}
319383
subject, err := s.ticketSubject(req.server)
320384
if err != nil {
321385
s.registry.Cancel(reqID, err)
386+
if uploadRegistered {
387+
s.uploadSources.Cancel(reqID, err)
388+
uploadRegistered = false
389+
}
322390
writeS3Error(w, r, s3HTTPError{status: http.StatusInternalServerError, code: "InternalError", message: "Request setup failed"})
323391
return
324392
}
325393
publishCtx, cancelPublish := context.WithDeadline(r.Context(), deadline)
326394
err = s.publisher.PublishTicketTo(publishCtx, subject, ticket)
327395
cancelPublish()
328396
if err != nil {
397+
s.registry.Cancel(reqID, err)
398+
if uploadRegistered {
399+
s.uploadSources.Cancel(reqID, err)
400+
uploadRegistered = false
401+
}
329402
s.logger.Warn("ticket publish failed", "request_id", reqID, "error", safeLogError(err))
330403
writeS3Error(w, r, s3HTTPError{status: http.StatusServiceUnavailable, code: "ServiceUnavailable", message: "Backend unavailable"})
331404
return
@@ -349,11 +422,12 @@ func (s *edgeServer) serveS3API(w http.ResponseWriter, r *http.Request) {
349422
}
350423

351424
func (s *edgeServer) resolveS3Request(r *http.Request) (s3EdgeRequest, error) {
352-
if _, err := s3api.VerifySigV4(r, s3api.VerifyOptions{
425+
authCtx, err := s3api.VerifySigV4(r, s3api.VerifyOptions{
353426
Credentials: s3api.Credentials{AccessKeyID: s.cfg.S3API.AccessKeyID, SecretAccessKey: s.cfg.S3API.SecretAccessKey},
354427
Region: s.cfg.S3API.Region,
355428
Now: s.now,
356-
}); err != nil {
429+
})
430+
if err != nil {
357431
return s3EdgeRequest{}, s3ErrorForAuth(err)
358432
}
359433

@@ -381,6 +455,27 @@ func (s *edgeServer) resolveS3Request(r *http.Request) (s3EdgeRequest, error) {
381455
if err := validateS3ObjectRequest(req); err != nil {
382456
return s3EdgeRequest{}, err
383457
}
458+
case s3api.OperationPutObject:
459+
req.operation = tickets.OperationPutObject
460+
req.contentType = r.Header.Get("Content-Type")
461+
if r.ContentLength >= 0 {
462+
contentLength := r.ContentLength
463+
req.contentLength = &contentLength
464+
}
465+
if err := validateS3ObjectRequest(req); err != nil {
466+
return s3EdgeRequest{}, err
467+
}
468+
if err := s.validateS3MutationRequest(r, mapping.Operation, authCtx, req); err != nil {
469+
return s3EdgeRequest{}, err
470+
}
471+
case s3api.OperationDeleteObject:
472+
req.operation = tickets.OperationDeleteObject
473+
if err := validateS3ObjectRequest(req); err != nil {
474+
return s3EdgeRequest{}, err
475+
}
476+
if err := s.validateS3MutationRequest(r, mapping.Operation, authCtx, req); err != nil {
477+
return s3EdgeRequest{}, err
478+
}
384479
case s3api.OperationListObjectsV2:
385480
req.operation = tickets.OperationListObjectsV2
386481
req.list = s3ListRequest(mapping)
@@ -401,6 +496,64 @@ func (s *edgeServer) resolveS3Request(r *http.Request) (s3EdgeRequest, error) {
401496
return req, nil
402497
}
403498

499+
func (s *edgeServer) validateS3MutationRequest(r *http.Request, operation s3api.Operation, authCtx s3api.AuthContext, req s3EdgeRequest) error {
500+
if !s.cfg.MutationsEnabled {
501+
return s3HTTPError{status: http.StatusMethodNotAllowed, code: "MethodNotAllowed", message: "Method not allowed"}
502+
}
503+
if strings.TrimSpace(r.Header.Get("Range")) != "" {
504+
return s3HTTPError{status: http.StatusBadRequest, code: "InvalidRequest", message: "Invalid request"}
505+
}
506+
if err := s3api.ValidatePayloadHashForOperation(operation, authCtx); err != nil {
507+
return s3HTTPError{status: http.StatusBadRequest, code: "InvalidRequest", message: "Invalid request"}
508+
}
509+
if operation == s3api.OperationPutObject {
510+
if req.contentLength == nil || *req.contentLength < 0 {
511+
return s3HTTPError{status: http.StatusBadRequest, code: "InvalidRequest", message: "Invalid request"}
512+
}
513+
if hasAWSChunkedMarkers(r, authCtx) {
514+
return s3HTTPError{status: http.StatusBadRequest, code: "InvalidRequest", message: "Invalid request"}
515+
}
516+
return nil
517+
}
518+
if deleteHasBodyOrChunkedMarkers(r, authCtx) {
519+
return s3HTTPError{status: http.StatusBadRequest, code: "InvalidRequest", message: "Invalid request"}
520+
}
521+
return nil
522+
}
523+
524+
func hasAWSChunkedMarkers(r *http.Request, authCtx s3api.AuthContext) bool {
525+
return authCtx.PayloadHashMode() == s3api.PayloadHashModeStreaming ||
526+
headerHasToken(r.Header, "Content-Encoding", "aws-chunked") ||
527+
r.Header.Get("X-Amz-Decoded-Content-Length") != "" ||
528+
r.Header.Get("X-Amz-Trailer") != "" ||
529+
r.Header.Get("Trailer") != "" ||
530+
hasTransferEncoding(r, "chunked")
531+
}
532+
533+
func deleteHasBodyOrChunkedMarkers(r *http.Request, authCtx s3api.AuthContext) bool {
534+
return r.ContentLength != 0 || hasAWSChunkedMarkers(r, authCtx)
535+
}
536+
537+
func headerHasToken(h http.Header, name, token string) bool {
538+
for _, value := range h.Values(name) {
539+
for _, part := range strings.Split(value, ",") {
540+
if strings.EqualFold(strings.TrimSpace(part), token) {
541+
return true
542+
}
543+
}
544+
}
545+
return false
546+
}
547+
548+
func hasTransferEncoding(r *http.Request, encoding string) bool {
549+
for _, value := range r.TransferEncoding {
550+
if strings.EqualFold(strings.TrimSpace(value), encoding) {
551+
return true
552+
}
553+
}
554+
return false
555+
}
556+
404557
func s3HeadBucketBackend(cfg config.EdgeConfig, mapping s3api.RequestMapping) string {
405558
if mapping.BackendBucket != "" {
406559
return mapping.BackendBucket
@@ -479,7 +632,13 @@ func (s *edgeServer) serveS3Direct(w http.ResponseWriter, r *http.Request, req s
479632
return
480633
}
481634

482-
fetched, err := fetcher.Fetch(r.Context(), s3fetch.Request{Method: r.Method, Operation: req.operation, Bucket: req.bucket, Key: req.key, Range: req.rangeHeader, List: req.list})
635+
fetchReq := s3fetch.Request{Method: r.Method, Operation: req.operation, Bucket: req.bucket, Key: req.key, Range: req.rangeHeader, List: req.list}
636+
if req.operation == tickets.OperationPutObject {
637+
fetchReq.Body = r.Body
638+
fetchReq.ContentLength = req.contentLength
639+
fetchReq.ContentType = req.contentType
640+
}
641+
fetched, err := fetcher.Fetch(r.Context(), fetchReq)
483642
if err != nil {
484643
s.writeS3DirectFetchError(w, r, req.server, err)
485644
return
@@ -492,6 +651,15 @@ func (s *edgeServer) serveS3Direct(w http.ResponseWriter, r *http.Request, req s
492651
if fetched.Body != nil {
493652
defer fetched.Body.Close()
494653
}
654+
if req.operation == tickets.OperationPutObject {
655+
setTrimmedHeader(w.Header(), "ETag", fetched.ETag)
656+
w.WriteHeader(http.StatusOK)
657+
return
658+
}
659+
if req.operation == tickets.OperationDeleteObject {
660+
w.WriteHeader(http.StatusNoContent)
661+
return
662+
}
495663

496664
metadata := directMetadata(fetched)
497665
copyPublicMetadata(w.Header(), metadata)
@@ -578,6 +746,13 @@ func s3ErrorForSetup(err error) s3HTTPError {
578746
return s3HTTPError{status: http.StatusInternalServerError, code: "InternalError", message: "Request setup failed"}
579747
}
580748

749+
func s3ErrorForUploadSetup(err error) s3HTTPError {
750+
if errors.Is(err, uploadsource.ErrExpired) || errors.Is(err, uploadsource.ErrInvalidSource) {
751+
return s3HTTPError{status: http.StatusBadRequest, code: "InvalidRequest", message: "Invalid request"}
752+
}
753+
return s3HTTPError{status: http.StatusInternalServerError, code: "InternalError", message: "Request setup failed"}
754+
}
755+
581756
func s3ErrorForWait(err error) s3HTTPError {
582757
if errors.Is(err, context.DeadlineExceeded) || errors.Is(err, pending.ErrExpired) {
583758
return s3HTTPError{status: http.StatusGatewayTimeout, code: "RequestTimeout", message: "Request timeout"}

0 commit comments

Comments
 (0)