From 48818927bb4071e8690e0dd5d533be1c663f0802 Mon Sep 17 00:00:00 2001 From: jonaustin09 Date: Fri, 27 Oct 2023 12:38:47 -0400 Subject: [PATCH 1/2] feat: Fixes #286, Created a struct which handles s3 select event streaming and event message construction --- backend/backend.go | 19 ++++++++++-- s3api/controllers/backend_moq_test.go | 33 ++++++++++---------- s3api/controllers/base.go | 7 +++-- s3api/controllers/base_test.go | 5 +-- s3select/message-handler.go | 44 +++++++++++++++++++++++++++ 5 files changed, 85 insertions(+), 23 deletions(-) create mode 100644 s3select/message-handler.go diff --git a/backend/backend.go b/backend/backend.go index e78f5ddc..9a70ca14 100644 --- a/backend/backend.go +++ b/backend/backend.go @@ -15,6 +15,7 @@ package backend import ( + "bufio" "context" "fmt" "io" @@ -22,6 +23,7 @@ import ( "github.com/aws/aws-sdk-go-v2/service/s3" "github.com/versity/versitygw/s3err" "github.com/versity/versitygw/s3response" + "github.com/versity/versitygw/s3select" ) //go:generate moq -out ../s3api/controllers/backend_moq_test.go -pkg controllers . Backend @@ -61,7 +63,7 @@ type Backend interface { // special case object operations RestoreObject(context.Context, *s3.RestoreObjectInput) error - SelectObjectContent(context.Context, *s3.SelectObjectContentInput) (s3response.SelectObjectContentResult, error) + SelectObjectContent(ctx context.Context, input *s3.SelectObjectContentInput) func(w *bufio.Writer) // object tags operations GetObjectTagging(_ context.Context, bucket, object string) (map[string]string, error) @@ -162,8 +164,19 @@ func (BackendUnsupported) PutObjectAcl(context.Context, *s3.PutObjectAclInput) e func (BackendUnsupported) RestoreObject(context.Context, *s3.RestoreObjectInput) error { return s3err.GetAPIError(s3err.ErrNotImplemented) } -func (BackendUnsupported) SelectObjectContent(context.Context, *s3.SelectObjectContentInput) (s3response.SelectObjectContentResult, error) { - return s3response.SelectObjectContentResult{}, s3err.GetAPIError(s3err.ErrNotImplemented) +func (BackendUnsupported) SelectObjectContent(ctx context.Context, input *s3.SelectObjectContentInput) func(w *bufio.Writer) { + return func(w *bufio.Writer) { + var getProgress s3select.GetProgress + progress := input.RequestProgress + if progress != nil && *progress.Enabled { + getProgress = func() (bytesScanned int64, bytesProcessed int64) { + return -1, -1 + } + } + mh := s3select.NewMessageHandler(ctx, w, getProgress) + apiErr := s3err.GetAPIError(s3err.ErrNotImplemented) + mh.FinishWithError(apiErr.Code, apiErr.Description) + } } func (BackendUnsupported) GetObjectTagging(_ context.Context, bucket, object string) (map[string]string, error) { diff --git a/s3api/controllers/backend_moq_test.go b/s3api/controllers/backend_moq_test.go index 27a1bb2f..80eb64df 100644 --- a/s3api/controllers/backend_moq_test.go +++ b/s3api/controllers/backend_moq_test.go @@ -4,6 +4,7 @@ package controllers import ( + "bufio" "context" "github.com/aws/aws-sdk-go-v2/service/s3" "github.com/versity/versitygw/backend" @@ -106,7 +107,7 @@ var _ backend.Backend = &BackendMock{} // RestoreObjectFunc: func(contextMoqParam context.Context, restoreObjectInput *s3.RestoreObjectInput) error { // panic("mock out the RestoreObject method") // }, -// SelectObjectContentFunc: func(contextMoqParam context.Context, selectObjectContentInput *s3.SelectObjectContentInput) (s3response.SelectObjectContentResult, error) { +// SelectObjectContentFunc: func(ctx context.Context, input *s3.SelectObjectContentInput) func(w *bufio.Writer) { // panic("mock out the SelectObjectContent method") // }, // ShutdownFunc: func() { @@ -213,7 +214,7 @@ type BackendMock struct { RestoreObjectFunc func(contextMoqParam context.Context, restoreObjectInput *s3.RestoreObjectInput) error // SelectObjectContentFunc mocks the SelectObjectContent method. - SelectObjectContentFunc func(contextMoqParam context.Context, selectObjectContentInput *s3.SelectObjectContentInput) (s3response.SelectObjectContentResult, error) + SelectObjectContentFunc func(ctx context.Context, input *s3.SelectObjectContentInput) func(w *bufio.Writer) // ShutdownFunc mocks the Shutdown method. ShutdownFunc func() @@ -441,10 +442,10 @@ type BackendMock struct { } // SelectObjectContent holds details about calls to the SelectObjectContent method. SelectObjectContent []struct { - // ContextMoqParam is the contextMoqParam argument value. - ContextMoqParam context.Context - // SelectObjectContentInput is the selectObjectContentInput argument value. - SelectObjectContentInput *s3.SelectObjectContentInput + // Ctx is the ctx argument value. + Ctx context.Context + // Input is the input argument value. + Input *s3.SelectObjectContentInput } // Shutdown holds details about calls to the Shutdown method. Shutdown []struct { @@ -1539,21 +1540,21 @@ func (mock *BackendMock) RestoreObjectCalls() []struct { } // SelectObjectContent calls SelectObjectContentFunc. -func (mock *BackendMock) SelectObjectContent(contextMoqParam context.Context, selectObjectContentInput *s3.SelectObjectContentInput) (s3response.SelectObjectContentResult, error) { +func (mock *BackendMock) SelectObjectContent(ctx context.Context, input *s3.SelectObjectContentInput) func(w *bufio.Writer) { if mock.SelectObjectContentFunc == nil { panic("BackendMock.SelectObjectContentFunc: method is nil but Backend.SelectObjectContent was just called") } callInfo := struct { - ContextMoqParam context.Context - SelectObjectContentInput *s3.SelectObjectContentInput + Ctx context.Context + Input *s3.SelectObjectContentInput }{ - ContextMoqParam: contextMoqParam, - SelectObjectContentInput: selectObjectContentInput, + Ctx: ctx, + Input: input, } mock.lockSelectObjectContent.Lock() mock.calls.SelectObjectContent = append(mock.calls.SelectObjectContent, callInfo) mock.lockSelectObjectContent.Unlock() - return mock.SelectObjectContentFunc(contextMoqParam, selectObjectContentInput) + return mock.SelectObjectContentFunc(ctx, input) } // SelectObjectContentCalls gets all the calls that were made to SelectObjectContent. @@ -1561,12 +1562,12 @@ func (mock *BackendMock) SelectObjectContent(contextMoqParam context.Context, se // // len(mockedBackend.SelectObjectContentCalls()) func (mock *BackendMock) SelectObjectContentCalls() []struct { - ContextMoqParam context.Context - SelectObjectContentInput *s3.SelectObjectContentInput + Ctx context.Context + Input *s3.SelectObjectContentInput } { var calls []struct { - ContextMoqParam context.Context - SelectObjectContentInput *s3.SelectObjectContentInput + Ctx context.Context + Input *s3.SelectObjectContentInput } mock.lockSelectObjectContent.RLock() calls = mock.calls.SelectObjectContent diff --git a/s3api/controllers/base.go b/s3api/controllers/base.go index 02c15458..d7275c5a 100644 --- a/s3api/controllers/base.go +++ b/s3api/controllers/base.go @@ -906,7 +906,7 @@ func (c S3ApiController) CreateActions(ctx *fiber.Ctx) error { return SendXMLResponse(ctx, nil, err, &MetaOpts{Logger: c.logger, Action: "SelectObjectContent", BucketOwner: parsedAcl.Owner}) } - res, err := c.be.SelectObjectContent(ctx.Context(), &s3.SelectObjectContentInput{ + sw := c.be.SelectObjectContent(ctx.Context(), &s3.SelectObjectContentInput{ Bucket: &bucket, Key: &key, Expression: payload.Expression, @@ -916,7 +916,10 @@ func (c S3ApiController) CreateActions(ctx *fiber.Ctx) error { RequestProgress: payload.RequestProgress, ScanRange: payload.ScanRange, }) - return SendXMLResponse(ctx, res, err, &MetaOpts{Logger: c.logger, Action: "SelectObjectContent", BucketOwner: parsedAcl.Owner}) + + ctx.Context().SetBodyStreamWriter(sw) + + return nil } if uploadId != "" { diff --git a/s3api/controllers/base_test.go b/s3api/controllers/base_test.go index 74d71c91..d4d7531a 100644 --- a/s3api/controllers/base_test.go +++ b/s3api/controllers/base_test.go @@ -15,6 +15,7 @@ package controllers import ( + "bufio" "context" "encoding/json" "fmt" @@ -1308,8 +1309,8 @@ func TestS3ApiController_CreateActions(t *testing.T) { CreateMultipartUploadFunc: func(context.Context, *s3.CreateMultipartUploadInput) (*s3.CreateMultipartUploadOutput, error) { return &s3.CreateMultipartUploadOutput{}, nil }, - SelectObjectContentFunc: func(contextMoqParam context.Context, selectObjectContentInput *s3.SelectObjectContentInput) (s3response.SelectObjectContentResult, error) { - return s3response.SelectObjectContentResult{}, nil + SelectObjectContentFunc: func(context.Context, *s3.SelectObjectContentInput) func(w *bufio.Writer) { + return func(w *bufio.Writer) {} }, }, } diff --git a/s3select/message-handler.go b/s3select/message-handler.go new file mode 100644 index 00000000..ecde975f --- /dev/null +++ b/s3select/message-handler.go @@ -0,0 +1,44 @@ +// Copyright 2023 Versity Software +// This file is licensed under the Apache License, Version 2.0 +// (the "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package s3select + +import ( + "bufio" + "context" +) + +type GetProgress func() (bytesScanned int64, bytesProcessed int64) + +type MessageHandler struct{} + +// Creates a new MessageHandler instance and starts the event streaming +func NewMessageHandler(ctx context.Context, w *bufio.Writer, getProgressFunc GetProgress) *MessageHandler { + return &MessageHandler{} +} + +// SendRecord sends a single Records message +func (mh *MessageHandler) SendRecord(payload []byte) error { + return nil +} + +// Finish terminates message stream with Stat and End message +func (mh *MessageHandler) Finish() error { + return nil +} + +// FinishWithError terminates event stream with error +func (mh *MessageHandler) FinishWithError(errorCode, errorMessage string) error { + return nil +} From bed1691a93e82f96f4f88acf80a88ce874189c36 Mon Sep 17 00:00:00 2001 From: Ben McClelland Date: Mon, 4 Dec 2023 09:02:53 -0800 Subject: [PATCH 2/2] feat: implement logic for s3 select object content stream --- s3select/message-handler.go | 324 +++++++++++++++++++++++++++++++++++- 1 file changed, 319 insertions(+), 5 deletions(-) diff --git a/s3select/message-handler.go b/s3select/message-handler.go index ecde975f..fcefc0bb 100644 --- a/s3select/message-handler.go +++ b/s3select/message-handler.go @@ -17,28 +17,342 @@ package s3select import ( "bufio" "context" + "encoding/binary" + "encoding/xml" + "fmt" + "hash/crc32" + "sync" + "sync/atomic" + "time" ) +// Protocol definition for messages can be found here: +// https://docs.aws.amazon.com/AmazonS3/latest/API/RESTSelectObjectAppendix.html + +var ( + // From ptotocol def: + // Enum indicating the header value type. + // For Amazon S3 Select, this is always 7. + headerValueType = byte(7) +) + +func intToTwoBytes(i int) []byte { + return []byte{byte(i >> 8), byte(i)} +} + +func generateHeader(messages ...string) []byte { + var header []byte + + for i, message := range messages { + if i%2 == 1 { + header = append(header, headerValueType) + header = append(header, intToTwoBytes(len(message))...) + } else { + header = append(header, byte(len(message))) + } + header = append(header, message...) + } + + return header +} + +func generateOctetHeader(message string) []byte { + return generateHeader( + ":message-type", + "event", + ":content-type", + "application/octet-stream", + ":event-type", + message) +} + +func generateTextHeader(message string) []byte { + return generateHeader( + ":message-type", + "event", + ":content-type", + "text/xml", + ":event-type", + message) +} + +func generateNoContentHeader(message string) []byte { + return generateHeader( + ":message-type", + "event", + ":event-type", + message) +} + +const ( + // 4 bytes total byte len + + // 4 bytes headers bytes len + + // 4 bytes prelude CRC + preludeLen = 12 + // CRC is uint32 + msgCrcLen = 4 +) + +var ( + recordsHeader = generateOctetHeader("Records") + continuationHeader = generateNoContentHeader("Cont") + continuationMessage = genMessage(continuationHeader, []byte{}) + progressHeader = generateTextHeader("Progress") + statsHeader = generateTextHeader("Stats") + endHeader = generateNoContentHeader("End") + endMessage = genMessage(endHeader, []byte{}) +) + +func uintToBytes(n uint32) []byte { + b := make([]byte, 4) + binary.BigEndian.PutUint32(b, n) + return b +} + +func generatePrelude(msgLen int, headerLen int) []byte { + prelude := make([]byte, 0, preludeLen) + + // 4 bytes total byte len + prelude = append(prelude, uintToBytes(uint32(msgLen+headerLen+preludeLen+msgCrcLen))...) + // 4 bytes headers bytes len + prelude = append(prelude, uintToBytes(uint32(headerLen))...) + // 4 bytes prelude CRC + prelude = append(prelude, uintToBytes(crc32.ChecksumIEEE(prelude))...) + + return prelude +} + +const ( + maxHeaderSize = 1024 * 1024 + maxMessageSize = 5 * 1024 * 1024 * 1024 +) + +func genMessage(header, payload []byte) []byte { + var msg []byte + // below is always true since the size is validated + // in the send record + if len(header) <= maxHeaderSize && len(payload) <= maxMessageSize { + msglen := preludeLen + len(header) + len(payload) + msgCrcLen + msg = make([]byte, 0, msglen) + } + + msg = append(msg, generatePrelude(len(payload), len(header))...) + msg = append(msg, header...) + msg = append(msg, payload...) + msg = append(msg, uintToBytes(crc32.ChecksumIEEE(msg))...) + + return msg +} + +func genRecordsMessage(payload []byte) []byte { + return genMessage(recordsHeader, payload) +} + +type progress struct { + XMLName xml.Name `xml:"Progress"` + BytesScanned int64 `xml:"BytesScanned"` + BytesProcessed int64 `xml:"BytesProcessed"` + BytesReturned int64 `xml:"BytesReturned"` +} + +func genProgressMessage(bytesScanned, bytesProcessed, bytesReturned int64) []byte { + progress := progress{ + BytesScanned: bytesScanned, + BytesProcessed: bytesProcessed, + BytesReturned: bytesReturned, + } + + xmlData, _ := xml.MarshalIndent(progress, "", " ") + payload := []byte(xml.Header + string(xmlData)) + return genMessage(progressHeader, payload) +} + +type stats struct { + XMLName xml.Name `xml:"Stats"` + BytesScanned int64 `xml:"BytesScanned"` + BytesProcessed int64 `xml:"BytesProcessed"` + BytesReturned int64 `xml:"BytesReturned"` +} + +func genStatsMessage(bytesScanned, bytesProcessed, bytesReturned int64) []byte { + stats := stats{ + BytesScanned: bytesScanned, + BytesProcessed: bytesProcessed, + BytesReturned: bytesReturned, + } + + xmlData, _ := xml.MarshalIndent(stats, "", " ") + payload := []byte(xml.Header + string(xmlData)) + return genMessage(statsHeader, payload) +} + +func genErrorMessage(errorCode, errorMessage string) []byte { + return genMessage(generateHeader( + ":error-code", + errorCode, + ":error-message", + errorMessage, + ":message-type", + "error", + ), []byte{}) +} + +// GetProgress is a callback function that periodically retrieves the current +// values for the following if not nil. This is used to send Progress +// messages back to client. +// BytesScanned => Number of bytes that have been processed before being uncompressed (if the file is compressed). +// BytesProcessed => Number of bytes that have been processed after being uncompressed (if the file is compressed). type GetProgress func() (bytesScanned int64, bytesProcessed int64) -type MessageHandler struct{} +type MessageHandler struct { + sync.Mutex + ctx context.Context + cancel context.CancelFunc + writer *bufio.Writer + data chan []byte + getProgress GetProgress + stopCh chan bool + resetCh chan bool + bytesReturned int64 +} -// Creates a new MessageHandler instance and starts the event streaming +// NewMessageHandler creates a new MessageHandler instance and starts the event streaming func NewMessageHandler(ctx context.Context, w *bufio.Writer, getProgressFunc GetProgress) *MessageHandler { - return &MessageHandler{} + ctx, cancel := context.WithCancel(ctx) + + mh := &MessageHandler{ + ctx: ctx, + cancel: cancel, + writer: w, + data: make(chan []byte), + getProgress: getProgressFunc, + resetCh: make(chan bool), + stopCh: make(chan bool), + } + + go mh.sendBackgroundMessages(mh.resetCh, mh.stopCh) + return mh +} + +func (mh *MessageHandler) write(data []byte) error { + mh.Lock() + defer mh.Unlock() + + mh.stopCh <- true + defer func() { mh.resetCh <- true }() + + _, err := mh.writer.Write(data) + if err != nil { + return err + } + + return mh.writer.Flush() +} + +const ( + continuationInterval = time.Second + progressInterval = time.Minute +) + +func (mh *MessageHandler) sendBackgroundMessages(resetCh, stopCh <-chan bool) { + continuationTicker := time.NewTicker(continuationInterval) + defer continuationTicker.Stop() + + var progressTicker *time.Ticker + var progressTickerChan <-chan time.Time + if mh.getProgress != nil { + progressTicker = time.NewTicker(progressInterval) + progressTickerChan = progressTicker.C + defer progressTicker.Stop() + } + +Loop: + for { + select { + case <-mh.ctx.Done(): + break Loop + + case <-continuationTicker.C: + err := mh.write(continuationMessage) + if err != nil { + mh.cancel() + break Loop + } + + case <-resetCh: + continuationTicker.Reset(continuationInterval) + + case <-stopCh: + continuationTicker.Stop() + + case <-progressTickerChan: + var bytesScanned, bytesProcessed int64 + if mh.getProgress != nil { + bytesScanned, bytesProcessed = mh.getProgress() + } + bytesReturned := atomic.LoadInt64(&mh.bytesReturned) + err := mh.write(genProgressMessage(bytesScanned, bytesProcessed, bytesReturned)) + if err != nil { + mh.cancel() + break Loop + } + } + } } // SendRecord sends a single Records message func (mh *MessageHandler) SendRecord(payload []byte) error { + if mh.ctx.Err() != nil { + return mh.ctx.Err() + } + + if len(payload) > maxMessageSize { + return fmt.Errorf("record max size exceeded") + } + + err := mh.write(genRecordsMessage(payload)) + if err != nil { + return err + } + + atomic.AddInt64(&mh.bytesReturned, int64(len(payload))) return nil } -// Finish terminates message stream with Stat and End message -func (mh *MessageHandler) Finish() error { +// Finish terminates message stream with Stats and End message +// generates stats and end message using function args based on: +// BytesScanned => Number of bytes that have been processed before being uncompressed (if the file is compressed). +// BytesProcessed => Number of bytes that have been processed after being uncompressed (if the file is compressed). +func (mh *MessageHandler) Finish(bytesScanned, bytesProcessed int64) error { + if mh.ctx.Err() != nil { + return mh.ctx.Err() + } + + bytesReturned := atomic.LoadInt64(&mh.bytesReturned) + err := mh.write(genStatsMessage(bytesScanned, bytesProcessed, bytesReturned)) + if err != nil { + return err + } + + err = mh.write(endMessage) + if err != nil { + return err + } + + mh.cancel() return nil } // FinishWithError terminates event stream with error func (mh *MessageHandler) FinishWithError(errorCode, errorMessage string) error { + if mh.ctx.Err() != nil { + return mh.ctx.Err() + } + err := mh.write(genErrorMessage(errorCode, errorMessage)) + if err != nil { + return err + } + + mh.cancel() return nil }