Skip to content

Commit 5bcd008

Browse files
authored
fix(source/aws_s3): paginate ListObjects to load >1000 migrations (#1412)
Signed-off-by: eylon <eylon@vybs.co>
1 parent 15c4690 commit 5bcd008

2 files changed

Lines changed: 108 additions & 14 deletions

File tree

‎source/aws_s3/s3.go‎

Lines changed: 20 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -76,25 +76,33 @@ func parseURI(uri string) (*Config, error) {
7676
}
7777

7878
func (s *s3Driver) loadMigrations() error {
79-
output, err := s.s3client.ListObjects(&s3.ListObjectsInput{
79+
input := &s3.ListObjectsInput{
8080
Bucket: aws.String(s.config.Bucket),
8181
Prefix: aws.String(s.config.Prefix),
8282
Delimiter: aws.String("/"),
83+
}
84+
// ListObjects returns at most 1000 keys per response, so paginate over
85+
// every page; otherwise migrations beyond the first 1000 objects are
86+
// silently dropped and never applied.
87+
var appendErr error
88+
err := s.s3client.ListObjectsPages(input, func(output *s3.ListObjectsOutput, _ bool) bool {
89+
for _, object := range output.Contents {
90+
_, fileName := path.Split(aws.StringValue(object.Key))
91+
m, err := source.DefaultParse(fileName)
92+
if err != nil {
93+
continue
94+
}
95+
if !s.migrations.Append(m) {
96+
appendErr = fmt.Errorf("unable to parse file %v", aws.StringValue(object.Key))
97+
return false
98+
}
99+
}
100+
return true
83101
})
84102
if err != nil {
85103
return err
86104
}
87-
for _, object := range output.Contents {
88-
_, fileName := path.Split(aws.StringValue(object.Key))
89-
m, err := source.DefaultParse(fileName)
90-
if err != nil {
91-
continue
92-
}
93-
if !s.migrations.Append(m) {
94-
return fmt.Errorf("unable to parse file %v", aws.StringValue(object.Key))
95-
}
96-
}
97-
return nil
105+
return appendErr
98106
}
99107

100108
func (s *s3Driver) Close() error {

‎source/aws_s3/s3_test.go‎

Lines changed: 88 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,9 @@ package awss3
22

33
import (
44
"errors"
5+
"fmt"
56
"io"
7+
"os"
68
"strings"
79
"testing"
810

@@ -40,6 +42,64 @@ func Test(t *testing.T) {
4042
st.Test(t, driver)
4143
}
4244

45+
func TestLoadMigrationsPaginates(t *testing.T) {
46+
// A single ListObjects response is capped at 1000 keys by S3. Spread the
47+
// migrations across several pages (via pageSize) to ensure loadMigrations
48+
// walks every page instead of silently stopping after the first one.
49+
const migrationCount = 300
50+
objects := make(map[string]string, migrationCount*2)
51+
for i := 1; i <= migrationCount; i++ {
52+
objects[fmt.Sprintf("prod/migrations/%d_foobar.up.sql", i)] = fmt.Sprintf("%d up", i)
53+
objects[fmt.Sprintf("prod/migrations/%d_foobar.down.sql", i)] = fmt.Sprintf("%d down", i)
54+
}
55+
s3Client := fakeS3{
56+
bucket: "some-bucket",
57+
pageSize: 50,
58+
objects: objects,
59+
}
60+
driver, err := WithInstance(&s3Client, &Config{
61+
Bucket: "some-bucket",
62+
Prefix: "prod/migrations/",
63+
})
64+
if err != nil {
65+
t.Fatal(err)
66+
}
67+
68+
first, err := driver.First()
69+
if err != nil {
70+
t.Fatal(err)
71+
}
72+
assert.Equal(t, uint(1), first)
73+
74+
version := first
75+
count := 1
76+
for {
77+
next, err := driver.Next(version)
78+
if errors.Is(err, os.ErrNotExist) {
79+
break
80+
}
81+
if err != nil {
82+
t.Fatal(err)
83+
}
84+
version = next
85+
count++
86+
}
87+
assert.Equal(t, migrationCount, count, "every migration across all pages should be loaded")
88+
assert.Equal(t, uint(migrationCount), version, "the highest-numbered migration should be loaded")
89+
90+
r, identifier, err := driver.ReadUp(uint(migrationCount))
91+
if err != nil {
92+
t.Fatal(err)
93+
}
94+
defer func() { _ = r.Close() }()
95+
assert.Equal(t, "foobar", identifier)
96+
body, err := io.ReadAll(r)
97+
if err != nil {
98+
t.Fatal(err)
99+
}
100+
assert.Equal(t, fmt.Sprintf("%d up", migrationCount), string(body))
101+
}
102+
43103
func TestParseURI(t *testing.T) {
44104
tests := []struct {
45105
name string
@@ -90,8 +150,11 @@ func TestParseURI(t *testing.T) {
90150

91151
type fakeS3 struct {
92152
s3.S3
93-
bucket string
94-
objects map[string]string
153+
bucket string
154+
// pageSize caps how many objects each ListObjectsPages page returns so
155+
// tests can exercise the multi-page path; 0 means a single page.
156+
pageSize int
157+
objects map[string]string
95158
}
96159

97160
func (s *fakeS3) ListObjects(input *s3.ListObjectsInput) (*s3.ListObjectsOutput, error) {
@@ -114,6 +177,29 @@ func (s *fakeS3) ListObjects(input *s3.ListObjectsInput) (*s3.ListObjectsOutput,
114177
return &output, nil
115178
}
116179

180+
func (s *fakeS3) ListObjectsPages(input *s3.ListObjectsInput, fn func(*s3.ListObjectsOutput, bool) bool) error {
181+
output, err := s.ListObjects(input)
182+
if err != nil {
183+
return err
184+
}
185+
contents := output.Contents
186+
pageSize := s.pageSize
187+
if pageSize <= 0 {
188+
pageSize = len(contents)
189+
}
190+
for start := 0; start < len(contents); start += pageSize {
191+
end := start + pageSize
192+
if end > len(contents) {
193+
end = len(contents)
194+
}
195+
lastPage := end == len(contents)
196+
if !fn(&s3.ListObjectsOutput{Contents: contents[start:end]}, lastPage) {
197+
break
198+
}
199+
}
200+
return nil
201+
}
202+
117203
func (s *fakeS3) GetObject(input *s3.GetObjectInput) (*s3.GetObjectOutput, error) {
118204
bucket := aws.StringValue(input.Bucket)
119205
if bucket != s.bucket {

0 commit comments

Comments
 (0)