internal/postgres: paginate versions query

Add start and limit args to the versions query so it can be
paginated.

Change-Id: I2f96911f74f91da4726127f87fff75f6fc7141f1
Reviewed-on: https://go-review.googlesource.com/c/pkgsite/+/780340
Reviewed-by: Hyang-Ah Hana Kim <hyangah@gmail.com>
Auto-Submit: Jonathan Amsterdam <jba@google.com>
LUCI-TryBot-Result: golang-scoped@luci-project-accounts.iam.gserviceaccount.com <golang-scoped@luci-project-accounts.iam.gserviceaccount.com>
kokoro-CI: kokoro <noreply+kokoro@google.com>
diff --git a/internal/postgres/delete_test.go b/internal/postgres/delete_test.go
index 47c085b..78d3a90 100644
--- a/internal/postgres/delete_test.go
+++ b/internal/postgres/delete_test.go
@@ -178,14 +178,14 @@
 	if err := testDB.DeletePseudoversionsExcept(ctx, sample.ModulePath, pseudo1); err != nil {
 		t.Fatal(err)
 	}
-	mods, err := getPathVersions(ctx, testDB, sample.ModulePath, version.TypeRelease)
+	mods, _, err := getPathVersions(ctx, testDB, sample.ModulePath, "", 800, version.TypeRelease)
 	if err != nil {
 		t.Fatal(err)
 	}
 	if len(mods) != 1 && mods[0].Version != sample.VersionString {
 		t.Errorf("module version %q was not found", sample.VersionString)
 	}
-	mods, err = getPathVersions(ctx, testDB, sample.ModulePath, version.TypePseudo)
+	mods, _, err = getPathVersions(ctx, testDB, sample.ModulePath, "", 10, version.TypePseudo)
 	if err != nil {
 		t.Fatal(err)
 	}
diff --git a/internal/postgres/details.go b/internal/postgres/details.go
index 1a3593e..59fab73 100644
--- a/internal/postgres/details.go
+++ b/internal/postgres/details.go
@@ -188,11 +188,15 @@
 	return nil
 }
 
-// scanModuleInfo constructs an *internal.ModuleInfo from the given scanner.
-func scanModuleInfo(scan func(dest ...any) error) (*internal.ModuleInfo, error) {
+// scanModuleInfo constructs an *internal.ModuleInfo from the
+// given scanner. The extras argument holds the destinations for
+// additional columns.
+func scanModuleInfo(scan func(dest ...any) error, extras ...any) (*internal.ModuleInfo, error) {
 	var mi internal.ModuleInfo
-	if err := scan(&mi.ModulePath, &mi.Version, &mi.CommitTime,
-		&mi.IsRedistributable, &mi.HasGoMod, jsonbScanner{&mi.SourceInfo}); err != nil {
+	args := []any{&mi.ModulePath, &mi.Version, &mi.CommitTime,
+		&mi.IsRedistributable, &mi.HasGoMod, jsonbScanner{&mi.SourceInfo}}
+	args = append(args, extras...)
+	if err := scan(args...); err != nil {
 		return nil, err
 	}
 	return &mi, nil
diff --git a/internal/postgres/version.go b/internal/postgres/version.go
index 602c322..9dcd664 100644
--- a/internal/postgres/version.go
+++ b/internal/postgres/version.go
@@ -9,6 +9,8 @@
 	"database/sql"
 	"errors"
 	"fmt"
+	"io"
+	"strconv"
 	"strings"
 
 	"github.com/Masterminds/squirrel"
@@ -30,25 +32,43 @@
 	defer derrors.WrapStack(&err, "GetVersionsForPath(ctx, %q)", path)
 	defer stats.Elapsed(ctx, "GetVersionsForPath")()
 
-	versions, err := getPathVersions(ctx, db, path, version.TypeRelease, version.TypePrerelease)
+	// When a page shows too many versions, it can result in a Chrome CSS
+	// bug: https://bugs.chromium.org/p/chromium/issues/detail?id=688640.
+	// For example,
+	// https://pkg.go.dev/github.com/aws/aws-sdk-go/aws/signer/v4?tab=versions.
+	// It's not that useful to see that many versions on a page anyway, so
+	// just limit to 800 versions.
+	versions, lsv, err := getPathVersions(ctx, db, path, "", 800, version.TypeRelease, version.TypePrerelease)
 	if err != nil {
 		return nil, err
 	}
 	if len(versions) != 0 {
 		return versions, nil
 	}
-	versions, err = getPathVersions(ctx, db, path, version.TypePseudo)
+	versions, _, err = getPathVersions(ctx, db, path, "", 10, version.TypePseudo)
 	if err != nil {
 		return nil, err
 	}
+	// To satisfy unparam until a subsequent CL.
+	// TODO(jba); remove this.
+	fmt.Fprint(io.Discard, lsv)
 	return versions, nil
 }
 
 // getPathVersions returns a list of versions sorted in descending semver
 // order. The version types included in the list are specified by a list of
 // VersionTypes.
-func getPathVersions(ctx context.Context, db *DB, path string, versionTypes ...version.Type) (_ []*internal.ModuleInfo, err error) {
-	defer derrors.WrapStack(&err, "getPathVersions(ctx, db, %q, %v)", path, versionTypes)
+// The result can be paginated by passing pageToken, which should be either the empty
+// string or a value returned from a previous call.
+// Each subsequent page will begin with the last value of the previous page.
+func getPathVersions(ctx context.Context, db *DB, path string, startPageToken string, limit int, versionTypes ...version.Type) (_ []*internal.ModuleInfo, nextPageToken string, err error) {
+	defer derrors.WrapStack(&err, "getPathVersions(ctx, db, %q, %q, %d, %v)", path, startPageToken, limit, versionTypes)
+
+	// Get previous values from pageToken.
+	var pageTokenArgs = []any{false, "", ""}
+	if startPageToken != "" {
+		pageTokenArgs, err = parsePageToken(startPageToken)
+	}
 
 	baseQuery := `
 	SELECT
@@ -57,7 +77,10 @@
 		m.commit_time,
 		m.redistributable,
 		m.has_go_mod,
-		m.source_info
+		m.source_info,
+		-- to construct the page token
+		m.incompatible,
+		m.sort_version
 	FROM modules m
 	INNER JOIN units u
 		ON u.module_id = m.id
@@ -71,42 +94,69 @@
 			LIMIT 1
 		)
 		AND version_type in (%s)
+		AND ($3 = '' OR (NOT m.incompatible, m.module_path, m.sort_version) <= (NOT $2, $3, $4))
 	ORDER BY
 		m.incompatible,
 		m.module_path DESC,
 		m.sort_version DESC %s`
 
-	queryEnd := `;`
 	if len(versionTypes) == 0 {
-		return nil, fmt.Errorf("error: must specify at least one version type")
-	} else if len(versionTypes) == 1 && versionTypes[0] == version.TypePseudo {
-		queryEnd = `LIMIT 10;`
-	} else {
-		// When a page shows too many versions, it can result in a Chrome CSS
-		// bug: https://bugs.chromium.org/p/chromium/issues/detail?id=688640.
-		// For example,
-		// https://pkg.go.dev/github.com/aws/aws-sdk-go/aws/signer/v4?tab=versions.
-		// It's not that useful to see that many versions on a page anyway, so
-		// just limit to 800 versions.
-		queryEnd = `LIMIT 800;`
+		return nil, "", fmt.Errorf("error: must specify at least one version type")
+	}
+	queryEnd := ";"
+	if limit > 0 {
+		queryEnd = fmt.Sprintf("LIMIT %d;", limit)
 	}
 	query := fmt.Sprintf(baseQuery, versionTypeExpr(versionTypes), queryEnd)
-	var versions []*internal.ModuleInfo
+
+	var (
+		versions         []*internal.ModuleInfo
+		lastIncompatible bool
+		lastSortVersion  string
+	)
 	collect := func(rows *sql.Rows) error {
-		mi, err := scanModuleInfo(rows.Scan)
+		mi, err := scanModuleInfo(rows.Scan, &lastIncompatible, &lastSortVersion)
 		if err != nil {
 			return fmt.Errorf("row.Scan(): %v", err)
 		}
 		versions = append(versions, mi)
 		return nil
 	}
-	if err := db.db.RunQuery(ctx, query, collect, path); err != nil {
-		return nil, err
+	args := append([]any{path}, pageTokenArgs...)
+	if err := db.db.RunQuery(ctx, query, collect, args...); err != nil {
+		return nil, "", err
 	}
 	if err := populateLatestInfos(ctx, db, versions); err != nil {
-		return nil, err
+		return nil, "", err
 	}
-	return versions, nil
+	// Construct the page token for the next page.
+	// See the comment near the top of this function for the format.
+	if len(versions) > 0 {
+		nextPageToken = makePageToken(lastIncompatible, versions[len(versions)-1].ModulePath, lastSortVersion)
+	}
+	return versions, nextPageToken, nil
+}
+
+// parsePageToken parses a page token for getPathVersions.
+// It return a slice of query args.
+func parsePageToken(s string) (queryArgs []any, err error) {
+	// A page token has the form "I P S"
+	// where I is a bool for incompatible version, P is a module path, and S
+	// is a sort version. Spaces suffice to separate these since none can contain a space.
+	parts := strings.Fields(s)
+	if len(parts) != 3 {
+		return nil, errors.New("invalid page token (wrong # parts)")
+	}
+	startIncompatible, err := strconv.ParseBool(parts[0])
+	if err != nil {
+		return nil, fmt.Errorf("invalid page token: %v", err)
+	}
+	return []any{startIncompatible, parts[1], parts[2]}, nil
+}
+
+// makePageToken constructs a page token for getPathVersions.
+func makePageToken(inc bool, mpath, version string) string {
+	return fmt.Sprintf("%t %s %s", inc, mpath, version)
 }
 
 // versionTypeExpr returns a comma-separated list of version types,
diff --git a/internal/postgres/version_test.go b/internal/postgres/version_test.go
index 7b29cea..68c3343 100644
--- a/internal/postgres/version_test.go
+++ b/internal/postgres/version_test.go
@@ -5,16 +5,18 @@
 package postgres
 
 import (
-	"context"
 	"database/sql"
 	"fmt"
+	"slices"
 	"testing"
 
 	"github.com/google/go-cmp/cmp"
+	"golang.org/x/mod/semver"
 	"golang.org/x/pkgsite/internal"
 	"golang.org/x/pkgsite/internal/source"
 	"golang.org/x/pkgsite/internal/stdlib"
 	"golang.org/x/pkgsite/internal/testing/sample"
+	"golang.org/x/pkgsite/internal/version"
 )
 
 func TestGetVersions(t *testing.T) {
@@ -58,7 +60,6 @@
 
 	testDB, release := acquire(t)
 	defer release()
-	ctx := context.Background()
 
 	for _, m := range testModules {
 		goMod := "module " + m.ModulePath
@@ -68,7 +69,7 @@
 				retract v1.0.3 // security flaw
 			`
 		}
-		testDB.MustInsertModuleGoMod(ctx, t, m, goMod)
+		testDB.MustInsertModuleGoMod(t.Context(), t, m, goMod)
 	}
 
 	stdModuleVersions := []*internal.ModuleInfo{
@@ -252,7 +253,7 @@
 				w.LatestVersion = latestVersions[w.ModulePath]
 			}
 
-			got, err := testDB.GetVersionsForPath(ctx, test.path)
+			got, err := testDB.GetVersionsForPath(t.Context(), test.path)
 			if err != nil {
 				t.Fatal(err)
 			}
@@ -269,7 +270,6 @@
 	t.Parallel()
 	testDB, release := acquire(t)
 	defer release()
-	ctx := context.Background()
 
 	for _, m := range []*internal.Module{
 		sample.Module("a.com/M", "v99.0.0+incompatible", "all", "most"),
@@ -399,7 +399,7 @@
 		},
 	} {
 		t.Run(test.unit, func(t *testing.T) {
-			got, err := testDB.GetLatestInfo(ctx, test.unit, test.module, nil)
+			got, err := testDB.GetLatestInfo(t.Context(), test.unit, test.module, nil)
 			if err != nil {
 				t.Fatal(err)
 			}
@@ -432,7 +432,6 @@
 	t.Parallel()
 	testDB, release := acquire(t)
 	defer release()
-	ctx := context.Background()
 
 	const (
 		modulePath = "example.com/m"
@@ -482,7 +481,7 @@
 				lmv.CookedVersion = test.cooked
 				lm = lmv
 			}
-			got, err := getLatestGoodVersion(ctx, testDB.db, modulePath, lm)
+			got, err := getLatestGoodVersion(t.Context(), testDB.db, modulePath, lm)
 			if err != nil {
 				t.Fatal(err)
 			}
@@ -497,7 +496,6 @@
 	t.Parallel()
 	testDB, release := acquire(t)
 	defer release()
-	ctx := context.Background()
 
 	const (
 		modulePath = "example.com/m"
@@ -548,7 +546,7 @@
 			if err != nil {
 				t.Fatal(err)
 			}
-			vGot, err := testDB.UpdateLatestModuleVersions(ctx, vNew)
+			vGot, err := testDB.UpdateLatestModuleVersions(t.Context(), vNew)
 			if err != nil {
 				t.Fatal(err)
 			}
@@ -567,13 +565,12 @@
 	t.Parallel()
 	testDB, release := acquire(t)
 	defer release()
-	ctx := context.Background()
 
 	const modulePath = "example.com/m"
 
 	getStatus := func() int {
 		var s int
-		err := testDB.db.QueryRow(ctx, `
+		err := testDB.db.QueryRow(t.Context(), `
 				SELECT status
 				FROM latest_module_versions l
 				INNER JOIN paths p ON (p.id=l.module_path_id)
@@ -587,7 +584,7 @@
 
 	// Insert a failure status.
 	newStatus := 410
-	if err := testDB.UpdateLatestModuleVersionsStatus(ctx, modulePath, newStatus); err != nil {
+	if err := testDB.UpdateLatestModuleVersionsStatus(t.Context(), modulePath, newStatus); err != nil {
 		t.Fatal(err)
 	}
 	if got := getStatus(); got != newStatus {
@@ -595,7 +592,7 @@
 	}
 
 	// GetLatestModuleVersions should return nil.
-	got, err := testDB.GetLatestModuleVersions(ctx, modulePath)
+	got, err := testDB.GetLatestModuleVersions(t.Context(), modulePath)
 	if err != nil {
 		t.Fatal(err)
 	}
@@ -605,7 +602,7 @@
 
 	// A new failure status should overwrite.
 	newStatus = 404
-	if err := testDB.UpdateLatestModuleVersionsStatus(ctx, modulePath, newStatus); err != nil {
+	if err := testDB.UpdateLatestModuleVersionsStatus(t.Context(), modulePath, newStatus); err != nil {
 		t.Fatal(err)
 	}
 	if got := getStatus(); got != newStatus {
@@ -617,7 +614,7 @@
 	if err != nil {
 		t.Fatal(err)
 	}
-	if _, err := testDB.UpdateLatestModuleVersions(ctx, lmv); err != nil {
+	if _, err := testDB.UpdateLatestModuleVersions(t.Context(), lmv); err != nil {
 		t.Fatal(err)
 	}
 	if got := getStatus(); got != 200 {
@@ -625,10 +622,10 @@
 	}
 
 	// Once we have good information, a bad status won't remove it.
-	if err := testDB.UpdateLatestModuleVersionsStatus(ctx, modulePath, 500); err != nil {
+	if err := testDB.UpdateLatestModuleVersionsStatus(t.Context(), modulePath, 500); err != nil {
 		t.Fatal(err)
 	}
-	got, err = testDB.GetLatestModuleVersions(ctx, modulePath)
+	got, err = testDB.GetLatestModuleVersions(t.Context(), modulePath)
 	if err != nil {
 		t.Fatal(err)
 	}
@@ -642,13 +639,12 @@
 	t.Parallel()
 	testDB, release := acquire(t)
 	defer release()
-	ctx := context.Background()
 
 	const modulePath = "example.com/m"
 
 	check := func(want string) {
 		t.Helper()
-		got, err := testDB.GetLatestModuleVersions(ctx, modulePath)
+		got, err := testDB.GetLatestModuleVersions(t.Context(), modulePath)
 		if err != nil {
 			t.Fatal(err)
 		}
@@ -668,10 +664,91 @@
 	check(v2)
 
 	// New latest-version info retracts v2 (and itself); good version should switch to v1.
-	testDB.MustInsertModuleGoMod(ctx, t, sample.Module(modulePath, "v1.3.0", "pkg"), fmt.Sprintf(`
+	testDB.MustInsertModuleGoMod(t.Context(), t, sample.Module(modulePath, "v1.3.0", "pkg"), fmt.Sprintf(`
 		module %s
 		retract v1.3.0
 		retract %s
 	`, modulePath, v2))
 	check(v1)
 }
+
+func TestGetPathVersionsPagination(t *testing.T) {
+	t.Parallel()
+	testDB, release := acquire(t)
+	defer release()
+
+	modulePath := "pagination.co/module"
+	versions := []string{
+		"v2.0.0+incompatible",
+		"v1.0.0",
+		"v1.2.0",
+		"v1.1.0",
+		"v1.1.0-20200330121822-37ff63d4418a",
+		"v2.0.0",
+		"v1.1.5",
+	}
+	for _, v := range versions {
+		testDB.MustInsertModule(t, sample.Module(modulePath, v, "pkg"))
+	}
+
+	packagePath := "pagination.co/module/pkg"
+
+	// Sort in memory DESC, with incompatibles last.
+	sortedVersions := slices.Clone(versions)
+	slices.SortFunc(sortedVersions, func(v1, v2 string) int {
+		if version.IsIncompatible(v1) != version.IsIncompatible(v2) {
+			if version.IsIncompatible(v1) {
+				return 1
+			}
+			return -1
+		}
+		// semver.Canonical removes build information, and thus +incompatible.
+		// So this works if both are incompatible, or both are compatible.
+		return -semver.Compare(semver.Canonical(v1), semver.Canonical(v2))
+	})
+	pageSize := 2
+	limit := pageSize + 1
+	start := ""
+	versionIndex := 0
+
+	for {
+		mods, nextStart, err := getPathVersions(t.Context(), testDB, packagePath, start, limit, version.TypeRelease, version.TypePrerelease, version.TypePseudo)
+		if err != nil {
+			t.Fatal(err)
+		}
+		if len(mods) == 0 {
+			if versionIndex < len(sortedVersions) {
+				t.Errorf("stopped paginating early: got %d versions, want %d", versionIndex, len(sortedVersions))
+			}
+			break
+		}
+
+		consume := len(mods)
+		hasLookahead := len(mods) == limit
+		if hasLookahead {
+			consume = pageSize
+		}
+
+		// Expected sub-slice
+		end := min(versionIndex+consume, len(sortedVersions))
+		wantSubSlice := sortedVersions[versionIndex:end]
+
+		var gotSubSlice []string
+		for i := range consume {
+			gotSubSlice = append(gotSubSlice, mods[i].Version)
+		}
+		if diff := cmp.Diff(wantSubSlice, gotSubSlice); diff != "" {
+			t.Errorf("page content mismatch at index %d (-want +got):\n%s", versionIndex, diff)
+		}
+
+		versionIndex += consume
+
+		if !hasLookahead {
+			if versionIndex < len(sortedVersions) {
+				t.Errorf("stopped paginating early (no lookahead): got %d versions, want %d", versionIndex, len(sortedVersions))
+			}
+			break
+		}
+		start = nextStart
+	}
+}