@@ -2,7 +2,9 @@ package awss3
22
33import (
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+
43103func TestParseURI (t * testing.T ) {
44104 tests := []struct {
45105 name string
@@ -90,8 +150,11 @@ func TestParseURI(t *testing.T) {
90150
91151type 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
97160func (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+
117203func (s * fakeS3 ) GetObject (input * s3.GetObjectInput ) (* s3.GetObjectOutput , error ) {
118204 bucket := aws .StringValue (input .Bucket )
119205 if bucket != s .bucket {
0 commit comments