diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 2eca5cc..1a39776 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -18,6 +18,16 @@ jobs: - uses: actions/setup-go@v5 with: go-version: ${{ matrix.go }} + # The cursor project is a single dependency-free core module (stdlib only); + # all v0.2 features live here, so there are no nested sub-modules to verify. + - name: go mod tidy (verify clean) + run: | + go mod tidy + if [ -n "$(git status --porcelain go.mod go.sum)" ]; then + echo "go.mod/go.sum are not tidy; run 'go mod tidy' and commit:" + git --no-pager diff go.mod go.sum + exit 1 + fi - name: gofmt run: | unformatted="$(gofmt -l .)" diff --git a/CHANGELOG.md b/CHANGELOG.md index 2cb0833..2c80354 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,25 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] ### Added +- **Seek-predicate builder** (`SeekKey`, `Column`, `Direction`, `Dialect`): + generate the row-wise `WHERE (a,b) > (?,?)`-style seek predicate and matching + `ORDER BY` from a typed composite key, with per-column `ASC`/`DESC`, forward and + backward paging, and dialect-aware placeholders for `database/sql` (`?`) and + `pgx`/`lib-pq` (`$N`). Golden tests assert exact SQL strings. +- **HMAC-signed cursors** (`SignedCodec`): sign cursor tokens with + HMAC-SHA256 so they cannot be forged or probed; tampered or wrong-key tokens are + rejected with `ErrTampered` on decode. +- **`net/http` query-param helper** (`ParseParams`, `LimitConfig`, `PageParams`): + parse `?cursor=&limit=` with default and max caps, clamping the limit safely + (including a usable zero-value config). +- **GraphQL Relay adapter** (`Connect`, `Connection`, `Edge`, `PageInfo`): build a + spec-compliant `pageInfo` (`startCursor`, `endCursor`, `hasNextPage`, + `hasPreviousPage`) and `edges` from a page result. +- **Runnable example**: a REST endpoint paginating a `(created_at, id)`-sorted + table via the seek builder against an in-memory fake data source, showing the + exact SQL the builder emits (`example_postgres_test.go`). + +### Previously (v0.1) - Initial v0.1 implementation: keyset (cursor) pagination with opaque cursor encode/decode and a `Slice` batch-to-page helper. See [ROADMAP.md](ROADMAP.md) for what comes next. diff --git a/README.md b/README.md index 151a8e4..4f1eef4 100644 --- a/README.md +++ b/README.md @@ -46,6 +46,49 @@ Fetch one more row than you intend to return. If it comes back, there's another Cursors are base64url-encoded JSON. They are opaque, not encrypted — don't put secrets in the sort key. +### v0.2 helpers + +- **Seek-predicate builder** — generate the seek `WHERE` and matching `ORDER BY` + from a typed composite key, with per-column `ASC`/`DESC` and dialect-aware + placeholders (`?` for `database/sql`, `$N` for `pgx`): + + ```go + seek := cursor.NewSeekKey( + cursor.Column{Name: "created_at", Dir: cursor.Asc}, + cursor.Column{Name: "id", Dir: cursor.Asc}, + ) + pred, nextArg := seek.Predicate(cursor.Dollar, false, 1) + // pred -> "(created_at > $1) OR (created_at = $2 AND id > $3)" + // order -> seek.OrderBy(false) == "created_at ASC, id ASC" + // args -> seek.Args(lastCreatedAt, lastID) // expands to ($1,$2,$3) values + ``` + +- **HMAC-signed cursors** — sign tokens so they can't be forged or probed; a + tampered token decodes to `ErrTampered`: + + ```go + codec := cursor.NewSignedCodec(secretKey) + tok, _ := codec.Encode(k) + err := codec.Decode(tok, &k) // ErrTampered if altered + ``` + +- **`net/http` param helper** — parse and clamp `?cursor=&limit=`: + + ```go + p := cursor.ParseParams(r, cursor.LimitConfig{Default: 20, Max: 100}) + // p.Cursor, p.Limit (clamped to [1, Max]) + ``` + +- **GraphQL Relay adapter** — build spec-compliant `edges` + `pageInfo`: + + ```go + conn := cursor.Connect(items, func(u User) string { return enc(u) }, hasNext, hasPrev) + // conn.Edges, conn.PageInfo.{StartCursor,EndCursor,HasNextPage,HasPreviousPage} + ``` + +See `example_postgres_test.go` for an end-to-end REST endpoint paginating a +`(created_at, id)`-sorted table with the seek builder. + ## Install ```sh diff --git a/example_postgres_test.go b/example_postgres_test.go new file mode 100644 index 0000000..741f17a --- /dev/null +++ b/example_postgres_test.go @@ -0,0 +1,164 @@ +package cursor_test + +import ( + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "sort" + "time" + + "github.com/hpower2/cursor" +) + +// event is a row in our imaginary `events` table, sorted on (created_at, id). +type event struct { + ID int64 `json:"id"` + CreatedAt time.Time `json:"created_at"` + Title string `json:"title"` +} + +// seekPos is the cursor payload: the (created_at, id) of the last row on a page. +type seekPos struct { + CreatedAt time.Time `json:"c"` + ID int64 `json:"i"` +} + +// fakeDB is an in-memory stand-in for Postgres. It holds rows already ordered by +// (created_at, id) and applies the seek predicate the same way the database would, +// so the example runs with no live DB while exercising the real builder output. +type fakeDB struct { + rows []event +} + +// query mimics executing the SELECT the builder helps assemble. after is nil for +// the first page. It returns up to limit+1 rows so the caller can detect HasMore. +func (db *fakeDB) query(after *seekPos, limit int) []event { + var out []event + for _, r := range db.rows { + if after != nil { + // Forward seek on (created_at ASC, id ASC): + // (created_at > c) OR (created_at = c AND id > i) + past := r.CreatedAt.After(after.CreatedAt) || + (r.CreatedAt.Equal(after.CreatedAt) && r.ID > after.ID) + if !past { + continue + } + } + out = append(out, r) + if len(out) == limit+1 { + break + } + } + return out +} + +// Example_postgresREST shows a REST endpoint paginating an (created_at, id)-sorted +// table with the seek-predicate builder, and prints the exact SQL the builder +// emits for the pgx dialect. +func Example_postgresREST() { + // The seek key the endpoint paginates on, most-significant column first. + seek := cursor.NewSeekKey( + cursor.Column{Name: "created_at", Dir: cursor.Asc}, + cursor.Column{Name: "id", Dir: cursor.Asc}, + ) + + // Show the exact SQL fragments the builder produces (pgx "$N" placeholders). + // startArg=1 leaves the LIMIT placeholder to follow the seek args. + pred, nextArg := seek.Predicate(cursor.Dollar, false, 1) + orderBy := seek.OrderBy(false) + fullSQL := fmt.Sprintf( + "SELECT id, created_at, title FROM events WHERE %s ORDER BY %s LIMIT %s", + pred, orderBy, cursor.Dollar.Placeholder(nextArg), + ) + fmt.Println("first-page SQL:") + fmt.Println(" SELECT id, created_at, title FROM events ORDER BY", orderBy, "LIMIT $1") + fmt.Println("next-page SQL:") + fmt.Println(" ", fullSQL) + + // Seed the fake data source (already ordered by created_at, id). + base := time.Date(2026, 6, 12, 9, 0, 0, 0, time.UTC) + db := &fakeDB{rows: []event{ + {1, base.Add(0 * time.Minute), "signup"}, + {2, base.Add(1 * time.Minute), "login"}, + {3, base.Add(1 * time.Minute), "click"}, // same timestamp as id=2; tie broken by id + {4, base.Add(2 * time.Minute), "logout"}, + {5, base.Add(3 * time.Minute), "purchase"}, + }} + + // The HTTP handler: parse ?cursor=&limit=, run the seek query, return a page. + handler := func(w http.ResponseWriter, r *http.Request) { + params := cursor.ParseParams(r, cursor.LimitConfig{Default: 2, Max: 50}) + + var after *seekPos + if params.Cursor != "" { + var pos seekPos + if err := cursor.Decode(params.Cursor, &pos); err != nil { + http.Error(w, "bad cursor", http.StatusBadRequest) + return + } + after = &pos + } + + batch := db.query(after, params.Limit) + page, err := cursor.Slice(batch, params.Limit, func(e event) (any, error) { + return seekPos{CreatedAt: e.CreatedAt, ID: e.ID}, nil + }) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + + // Stable output for the example: ids of the page plus the next cursor flag. + ids := make([]int64, len(page.Items)) + for i, e := range page.Items { + ids[i] = e.ID + } + sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] }) + _ = json.NewEncoder(w).Encode(map[string]any{ + "ids": ids, + "has_more": page.HasMore, + "next": page.Next, + }) + } + + srv := httptest.NewServer(http.HandlerFunc(handler)) + defer srv.Close() + + // Page 1. + resp1 := fetch(srv.URL + "/events?limit=2") + fmt.Printf("page1: ids=%v has_more=%v\n", resp1["ids"], resp1["has_more"]) + + // Page 2, using the cursor returned by page 1. + next := resp1["next"].(string) + resp2 := fetch(srv.URL + "/events?limit=2&cursor=" + next) + fmt.Printf("page2: ids=%v has_more=%v\n", resp2["ids"], resp2["has_more"]) + + // Page 3 (final). + next2 := resp2["next"].(string) + resp3 := fetch(srv.URL + "/events?limit=2&cursor=" + next2) + fmt.Printf("page3: ids=%v has_more=%v\n", resp3["ids"], resp3["has_more"]) + + // Output: + // first-page SQL: + // SELECT id, created_at, title FROM events ORDER BY created_at ASC, id ASC LIMIT $1 + // next-page SQL: + // SELECT id, created_at, title FROM events WHERE (created_at > $1) OR (created_at = $2 AND id > $3) ORDER BY created_at ASC, id ASC LIMIT $4 + // page1: ids=[1 2] has_more=true + // page2: ids=[3 4] has_more=true + // page3: ids=[5] has_more=false +} + +// fetch is a tiny helper that GETs a URL and decodes the JSON body. +func fetch(url string) map[string]any { + resp, err := http.Get(url) + if err != nil { + panic(err) + } + defer resp.Body.Close() + var out map[string]any + if err := json.NewDecoder(resp.Body).Decode(&out); err != nil { + panic(err) + } + return out +} diff --git a/graphql.go b/graphql.go new file mode 100644 index 0000000..fea45c8 --- /dev/null +++ b/graphql.go @@ -0,0 +1,52 @@ +package cursor + +// PageInfo is the Relay-style connection metadata describing where a page sits +// within the full result set. Field names and JSON tags follow the GraphQL Cursor +// Connections specification so the struct serializes directly into a `pageInfo`. +type PageInfo struct { + HasNextPage bool `json:"hasNextPage"` + HasPreviousPage bool `json:"hasPreviousPage"` + StartCursor string `json:"startCursor"` + EndCursor string `json:"endCursor"` +} + +// Edge wraps a node with its cursor, as required by a Relay connection's `edges`. +type Edge[T any] struct { + Cursor string `json:"cursor"` + Node T `json:"node"` +} + +// Connection is a Relay-style connection: a list of edges plus page metadata. +type Connection[T any] struct { + Edges []Edge[T] `json:"edges"` + PageInfo PageInfo `json:"pageInfo"` +} + +// Connect builds a Relay Connection from a slice of items, a per-item cursor +// function, and flags describing neighbouring pages. +// +// - cursorOf returns the opaque cursor for an item (e.g. the result of +// SignedCodec.Encode or cursor.Encode over the item's sort key). +// - hasNext / hasPrev describe whether pages exist after / before this one; +// callers typically derive hasNext from a fetched limit+1 sentinel and +// hasPrev from whether an incoming cursor was supplied. +// +// startCursor and endCursor are taken from the first and last edges; for an empty +// page both are "" and both hasNextPage/hasPreviousPage are reported as given. +func Connect[T any](items []T, cursorOf func(T) string, hasNext, hasPrev bool) Connection[T] { + edges := make([]Edge[T], len(items)) + for i, it := range items { + edges[i] = Edge[T]{Cursor: cursorOf(it), Node: it} + } + + info := PageInfo{ + HasNextPage: hasNext, + HasPreviousPage: hasPrev, + } + if len(edges) > 0 { + info.StartCursor = edges[0].Cursor + info.EndCursor = edges[len(edges)-1].Cursor + } + + return Connection[T]{Edges: edges, PageInfo: info} +} diff --git a/graphql_test.go b/graphql_test.go new file mode 100644 index 0000000..ab75c3d --- /dev/null +++ b/graphql_test.go @@ -0,0 +1,100 @@ +package cursor + +import ( + "encoding/json" + "strconv" + "testing" +) + +type node struct { + ID string + Name string +} + +func TestConnectBasic(t *testing.T) { + items := []node{{"1", "a"}, {"2", "b"}, {"3", "c"}} + conn := Connect(items, func(n node) string { return "cur_" + n.ID }, true, false) + + if len(conn.Edges) != 3 { + t.Fatalf("got %d edges, want 3", len(conn.Edges)) + } + if conn.Edges[0].Cursor != "cur_1" || conn.Edges[0].Node.Name != "a" { + t.Errorf("edge0 = %+v", conn.Edges[0]) + } + if conn.PageInfo.StartCursor != "cur_1" { + t.Errorf("StartCursor = %q, want cur_1", conn.PageInfo.StartCursor) + } + if conn.PageInfo.EndCursor != "cur_3" { + t.Errorf("EndCursor = %q, want cur_3", conn.PageInfo.EndCursor) + } + if !conn.PageInfo.HasNextPage { + t.Error("HasNextPage = false, want true") + } + if conn.PageInfo.HasPreviousPage { + t.Error("HasPreviousPage = true, want false") + } +} + +func TestConnectEmpty(t *testing.T) { + conn := Connect([]node{}, func(n node) string { return n.ID }, false, true) + if len(conn.Edges) != 0 { + t.Fatalf("got %d edges, want 0", len(conn.Edges)) + } + if conn.PageInfo.StartCursor != "" || conn.PageInfo.EndCursor != "" { + t.Errorf("empty page cursors should be blank, got %+v", conn.PageInfo) + } + if conn.PageInfo.HasNextPage { + t.Error("HasNextPage should be false") + } + if !conn.PageInfo.HasPreviousPage { + t.Error("HasPreviousPage should be true") + } +} + +func TestConnectSingle(t *testing.T) { + conn := Connect([]node{{"7", "g"}}, func(n node) string { return n.ID }, false, false) + if conn.PageInfo.StartCursor != "7" || conn.PageInfo.EndCursor != "7" { + t.Errorf("single page start==end==7, got %+v", conn.PageInfo) + } +} + +func TestConnectJSONShape(t *testing.T) { + items := []node{{"1", "a"}} + conn := Connect(items, func(n node) string { return "c" + n.ID }, true, true) + b, err := json.Marshal(conn) + if err != nil { + t.Fatal(err) + } + + var raw map[string]json.RawMessage + if err := json.Unmarshal(b, &raw); err != nil { + t.Fatal(err) + } + if _, ok := raw["edges"]; !ok { + t.Error("missing edges in JSON") + } + pi, ok := raw["pageInfo"] + if !ok { + t.Fatal("missing pageInfo in JSON") + } + var info map[string]json.RawMessage + if err := json.Unmarshal(pi, &info); err != nil { + t.Fatal(err) + } + for _, k := range []string{"hasNextPage", "hasPreviousPage", "startCursor", "endCursor"} { + if _, ok := info[k]; !ok { + t.Errorf("pageInfo missing spec field %q", k) + } + } +} + +func TestConnectManyCursors(t *testing.T) { + items := make([]node, 5) + for i := range items { + items[i] = node{ID: strconv.Itoa(i)} + } + conn := Connect(items, func(n node) string { return n.ID }, true, false) + if conn.PageInfo.StartCursor != "0" || conn.PageInfo.EndCursor != "4" { + t.Errorf("start/end = %q/%q, want 0/4", conn.PageInfo.StartCursor, conn.PageInfo.EndCursor) + } +} diff --git a/httpparam.go b/httpparam.go new file mode 100644 index 0000000..a179b88 --- /dev/null +++ b/httpparam.go @@ -0,0 +1,80 @@ +package cursor + +import ( + "net/http" + "strconv" +) + +// LimitConfig configures how the limit query parameter is parsed and clamped. +type LimitConfig struct { + // Default is used when the limit parameter is absent or empty. + Default int + // Max is the inclusive upper bound; a requested limit above it is clamped down. + Max int + // Param is the query-parameter name for the limit (defaults to "limit"). + Param string + // CursorParam is the query-parameter name for the cursor (defaults to "cursor"). + CursorParam string +} + +func (c LimitConfig) limitParam() string { + if c.Param != "" { + return c.Param + } + return "limit" +} + +func (c LimitConfig) cursorParam() string { + if c.CursorParam != "" { + return c.CursorParam + } + return "cursor" +} + +// PageParams holds the parsed, validated pagination inputs from a request. +type PageParams struct { + // Cursor is the raw cursor token from the query string ("" when absent). + Cursor string + // Limit is the clamped page size, guaranteed to be in [1, Max]. + Limit int +} + +// ParseParams reads the cursor and limit query parameters from r and returns them +// clamped according to cfg. +// +// Clamping rules for the limit: +// - missing, empty, non-numeric, or <= 0 -> cfg.Default +// - greater than cfg.Max -> cfg.Max +// - otherwise -> the requested value +// +// If cfg.Default is non-positive it falls back to 20, and if cfg.Max is +// non-positive it falls back to 100, so a zero-value LimitConfig is still safe. +func ParseParams(r *http.Request, cfg LimitConfig) PageParams { + def := cfg.Default + if def <= 0 { + def = 20 + } + max := cfg.Max + if max <= 0 { + max = 100 + } + if def > max { + def = max + } + + q := r.URL.Query() + limit := def + if raw := q.Get(cfg.limitParam()); raw != "" { + if n, err := strconv.Atoi(raw); err == nil && n > 0 { + limit = n + } + } + if limit > max { + limit = max + } + + return PageParams{ + Cursor: q.Get(cfg.cursorParam()), + Limit: limit, + } +} diff --git a/httpparam_test.go b/httpparam_test.go new file mode 100644 index 0000000..52faeb8 --- /dev/null +++ b/httpparam_test.go @@ -0,0 +1,75 @@ +package cursor + +import ( + "net/http/httptest" + "testing" +) + +func TestParseParams(t *testing.T) { + cfg := LimitConfig{Default: 25, Max: 100} + + tests := []struct { + name string + url string + wantCursor string + wantLimit int + }{ + {"both present", "/items?cursor=abc&limit=50", "abc", 50}, + {"missing both", "/items", "", 25}, + {"missing limit uses default", "/items?cursor=xyz", "xyz", 25}, + {"missing cursor", "/items?limit=10", "", 10}, + {"over cap clamped", "/items?limit=1000", "", 100}, + {"exactly cap", "/items?limit=100", "", 100}, + {"zero falls to default", "/items?limit=0", "", 25}, + {"negative falls to default", "/items?limit=-5", "", 25}, + {"non-numeric falls to default", "/items?limit=abc", "", 25}, + {"empty limit falls to default", "/items?limit=", "", 25}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + req := httptest.NewRequest("GET", tt.url, nil) + got := ParseParams(req, cfg) + if got.Cursor != tt.wantCursor { + t.Errorf("Cursor = %q, want %q", got.Cursor, tt.wantCursor) + } + if got.Limit != tt.wantLimit { + t.Errorf("Limit = %d, want %d", got.Limit, tt.wantLimit) + } + }) + } +} + +func TestParseParamsZeroConfigDefaults(t *testing.T) { + // Zero-value config must still be safe: default 20, max 100. + req := httptest.NewRequest("GET", "/items", nil) + if got := ParseParams(req, LimitConfig{}); got.Limit != 20 { + t.Errorf("zero-config default Limit = %d, want 20", got.Limit) + } + + req = httptest.NewRequest("GET", "/items?limit=99999", nil) + if got := ParseParams(req, LimitConfig{}); got.Limit != 100 { + t.Errorf("zero-config max Limit = %d, want 100", got.Limit) + } +} + +func TestParseParamsCustomParamNames(t *testing.T) { + cfg := LimitConfig{Default: 10, Max: 50, Param: "page_size", CursorParam: "page_token"} + req := httptest.NewRequest("GET", "/items?page_token=tok&page_size=5", nil) + got := ParseParams(req, cfg) + if got.Cursor != "tok" { + t.Errorf("Cursor = %q, want %q", got.Cursor, "tok") + } + if got.Limit != 5 { + t.Errorf("Limit = %d, want 5", got.Limit) + } +} + +func TestParseParamsDefaultClampedToMax(t *testing.T) { + // If Default exceeds Max, the effective default must not exceed Max. + cfg := LimitConfig{Default: 500, Max: 30} + req := httptest.NewRequest("GET", "/items", nil) + if got := ParseParams(req, cfg); got.Limit != 30 { + t.Errorf("Limit = %d, want 30 (default clamped to max)", got.Limit) + } +} diff --git a/seek.go b/seek.go new file mode 100644 index 0000000..08f06d0 --- /dev/null +++ b/seek.go @@ -0,0 +1,160 @@ +package cursor + +import ( + "strconv" + "strings" +) + +// Direction is the sort order of a single key column. +type Direction int + +const ( + // Asc sorts ascending (the default zero value). + Asc Direction = iota + // Desc sorts descending. + Desc Direction = iota +) + +func (d Direction) String() string { + if d == Desc { + return "DESC" + } + return "ASC" +} + +// Dialect controls how positional placeholders are rendered in the generated SQL. +type Dialect int + +const ( + // Question renders database/sql style "?" placeholders. + Question Dialect = iota + // Dollar renders pgx/lib-pq style "$1", "$2", ... placeholders. + Dollar Dialect = iota +) + +// Column is one component of a composite sort key: the SQL column name and the +// direction it is sorted in. +type Column struct { + Name string + Dir Direction +} + +// SeekKey describes an ordered, typed composite sort key used to drive keyset +// pagination. The columns are listed most-significant first, exactly as they +// appear in the ORDER BY. +// +// A SeekKey is intentionally small and dependency-free: it knows how to emit the +// WHERE seek predicate and the matching ORDER BY for a given dialect and paging +// direction, but it owns only the pagination clause — never your whole SELECT. +type SeekKey struct { + Columns []Column +} + +// NewSeekKey builds a SeekKey from the given columns (most-significant first). +func NewSeekKey(cols ...Column) SeekKey { + return SeekKey{Columns: cols} +} + +// Placeholder renders the n-th (1-based) positional placeholder for the dialect: +// "?" for Question and "$n" for Dollar. It is exported so callers assembling the +// rest of their SELECT (e.g. the LIMIT placeholder) can stay dialect-consistent. +func (d Dialect) Placeholder(n int) string { + if d == Dollar { + return "$" + strconv.Itoa(n) + } + return "?" +} + +// effectiveDir flips the column direction when paging backward, because reading +// the previous page means walking the index in the opposite order. +func effectiveDir(dir Direction, backward bool) Direction { + if !backward { + return dir + } + if dir == Asc { + return Desc + } + return Asc +} + +// OrderBy returns the "col1 DIR, col2 DIR, ..." clause (without the leading +// "ORDER BY") for paging in the requested direction. When backward is true every +// column's direction is reversed so the database walks toward the previous page; +// callers must then re-reverse the returned rows to restore natural order. +func (k SeekKey) OrderBy(backward bool) string { + parts := make([]string, len(k.Columns)) + for i, c := range k.Columns { + parts[i] = c.Name + " " + effectiveDir(c.Dir, backward).String() + } + return strings.Join(parts, ", ") +} + +// Predicate generates the row-wise seek WHERE clause that selects every row +// strictly after the cursor position, honoring each column's ASC/DESC direction +// and the paging direction. +// +// It returns the SQL fragment (without a leading "WHERE") and the ordered list of +// 1-based placeholder argument positions. The fragment is expanded into the +// classic lexicographic comparison so it works on every SQL engine, not just +// those supporting native row-value comparison: +// +// (a > ?) OR (a = ? AND b > ?) +// +// startArg is the 1-based index of the first placeholder to emit (use 1 when the +// seek values are the only bound parameters). The returned int is the next free +// placeholder index, so callers can chain further bound values (e.g. LIMIT). +func (k SeekKey) Predicate(dialect Dialect, backward bool, startArg int) (sql string, nextArg int) { + if len(k.Columns) == 0 { + return "", startArg + } + + arg := startArg + // ph emits the next placeholder and advances the counter. + ph := func() string { + s := dialect.Placeholder(arg) + arg++ + return s + } + + // cmpOp picks the comparison operator for a column given its effective + // direction: ascending seeks forward with ">", descending with "<". + cmpOp := func(dir Direction) string { + if effectiveDir(dir, backward) == Desc { + return "<" + } + return ">" + } + + var terms []string + // Term i: all columns before i are equal, and column i is strictly past the + // cursor. The union of these disjoint terms is the full seek predicate. + for i := 0; i < len(k.Columns); i++ { + var conds []string + for j := 0; j < i; j++ { + conds = append(conds, k.Columns[j].Name+" = "+ph()) + } + conds = append(conds, k.Columns[i].Name+" "+cmpOp(k.Columns[i].Dir)+" "+ph()) + terms = append(terms, "("+strings.Join(conds, " AND ")+")") + } + + return strings.Join(terms, " OR "), arg +} + +// Args returns the bound parameter values in the exact order the placeholders +// produced by Predicate expect them. values must hold one entry per column, in +// the same column order as the SeekKey; they are the sort-key values of the last +// row on the current page (the cursor position). +// +// The expansion mirrors Predicate's lexicographic terms: +// +// columns (a, b) -> args (a, a, b) +// columns (a,b,c) -> args (a, a,b, a,b,c) +func (k SeekKey) Args(values ...any) []any { + var out []any + for i := 0; i < len(k.Columns) && i < len(values); i++ { + for j := 0; j <= i; j++ { + out = append(out, values[j]) + } + } + return out +} diff --git a/seek_test.go b/seek_test.go new file mode 100644 index 0000000..5cd36cf --- /dev/null +++ b/seek_test.go @@ -0,0 +1,138 @@ +package cursor + +import ( + "reflect" + "testing" +) + +func TestSeekPredicateGolden(t *testing.T) { + tests := []struct { + name string + key SeekKey + dialect Dialect + backward bool + startArg int + wantSQL string + wantNext int + }{ + { + name: "two col asc/asc forward question", + key: NewSeekKey(Column{"created_at", Asc}, Column{"id", Asc}), + dialect: Question, + backward: false, + startArg: 1, + wantSQL: "(created_at > ?) OR (created_at = ? AND id > ?)", + wantNext: 4, + }, + { + name: "two col asc/asc forward dollar", + key: NewSeekKey(Column{"created_at", Asc}, Column{"id", Asc}), + dialect: Dollar, + backward: false, + startArg: 1, + wantSQL: "(created_at > $1) OR (created_at = $2 AND id > $3)", + wantNext: 4, + }, + { + name: "two col desc/asc forward dollar", + key: NewSeekKey(Column{"created_at", Desc}, Column{"id", Asc}), + dialect: Dollar, + backward: false, + startArg: 1, + wantSQL: "(created_at < $1) OR (created_at = $2 AND id > $3)", + wantNext: 4, + }, + { + name: "two col asc/asc backward dollar flips ops", + key: NewSeekKey(Column{"created_at", Asc}, Column{"id", Asc}), + dialect: Dollar, + backward: true, + startArg: 1, + wantSQL: "(created_at < $1) OR (created_at = $2 AND id < $3)", + wantNext: 4, + }, + { + name: "two col desc/desc backward question flips to gt", + key: NewSeekKey(Column{"created_at", Desc}, Column{"id", Desc}), + dialect: Question, + backward: true, + startArg: 1, + wantSQL: "(created_at > ?) OR (created_at = ? AND id > ?)", + wantNext: 4, + }, + { + name: "single col asc forward question", + key: NewSeekKey(Column{"id", Asc}), + dialect: Question, + backward: false, + startArg: 1, + wantSQL: "(id > ?)", + wantNext: 2, + }, + { + name: "dollar with nonzero start arg", + key: NewSeekKey(Column{"created_at", Asc}, Column{"id", Asc}), + dialect: Dollar, + backward: false, + startArg: 5, + wantSQL: "(created_at > $5) OR (created_at = $6 AND id > $7)", + wantNext: 8, + }, + { + name: "three col asc/desc/asc forward dollar", + key: NewSeekKey(Column{"a", Asc}, Column{"b", Desc}, Column{"c", Asc}), + dialect: Dollar, + backward: false, + startArg: 1, + wantSQL: "(a > $1) OR (a = $2 AND b < $3) OR (a = $4 AND b = $5 AND c > $6)", + wantNext: 7, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gotSQL, gotNext := tt.key.Predicate(tt.dialect, tt.backward, tt.startArg) + if gotSQL != tt.wantSQL { + t.Errorf("SQL mismatch\n got: %s\nwant: %s", gotSQL, tt.wantSQL) + } + if gotNext != tt.wantNext { + t.Errorf("nextArg = %d, want %d", gotNext, tt.wantNext) + } + }) + } +} + +func TestSeekOrderByGolden(t *testing.T) { + k := NewSeekKey(Column{"created_at", Desc}, Column{"id", Asc}) + + if got, want := k.OrderBy(false), "created_at DESC, id ASC"; got != want { + t.Errorf("forward OrderBy = %q, want %q", got, want) + } + if got, want := k.OrderBy(true), "created_at ASC, id DESC"; got != want { + t.Errorf("backward OrderBy = %q, want %q", got, want) + } +} + +func TestSeekArgsExpansion(t *testing.T) { + k := NewSeekKey(Column{"created_at", Asc}, Column{"id", Asc}) + got := k.Args(int64(100), "u_7") + want := []any{int64(100), int64(100), "u_7"} + if !reflect.DeepEqual(got, want) { + t.Errorf("Args = %v, want %v", got, want) + } + + k3 := NewSeekKey(Column{"a", Asc}, Column{"b", Asc}, Column{"c", Asc}) + got3 := k3.Args(1, 2, 3) + want3 := []any{1, 1, 2, 1, 2, 3} + if !reflect.DeepEqual(got3, want3) { + t.Errorf("Args3 = %v, want %v", got3, want3) + } +} + +func TestSeekEmptyKey(t *testing.T) { + var k SeekKey + sql, next := k.Predicate(Question, false, 1) + if sql != "" || next != 1 { + t.Errorf("empty key: got (%q, %d), want (%q, 1)", sql, next, "") + } +} diff --git a/sign.go b/sign.go new file mode 100644 index 0000000..89171b1 --- /dev/null +++ b/sign.go @@ -0,0 +1,90 @@ +package cursor + +import ( + "crypto/hmac" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "errors" +) + +// ErrTampered is returned by SignedCodec.Decode when a token's signature does not +// match its payload — i.e. the cursor was forged, truncated, or otherwise altered. +var ErrTampered = errors.New("cursor: signature mismatch") + +// SignedCodec encodes and decodes cursors that carry an HMAC-SHA256 signature so +// they cannot be forged or probed by a client. The wire format is: +// +// base64url( payloadJSON ) "." base64url( hmac-sha256(payloadJSON) ) +// +// Both halves use RawURLEncoding, so the whole token is URL-safe and free of +// padding. The key never appears in the token; only a MAC over the payload does. +// +// A SignedCodec is safe for concurrent use: it holds an immutable key and uses no +// shared mutable state. +type SignedCodec struct { + key []byte +} + +// NewSignedCodec returns a codec that signs cursors with the given secret key. +// The key should be at least 32 random bytes; it is copied so later mutation of +// the caller's slice has no effect. +func NewSignedCodec(key []byte) *SignedCodec { + k := make([]byte, len(key)) + copy(k, key) + return &SignedCodec{key: k} +} + +func (c *SignedCodec) mac(payload []byte) []byte { + m := hmac.New(sha256.New, c.key) + m.Write(payload) + return m.Sum(nil) +} + +// Encode serializes v to JSON, appends an HMAC-SHA256 signature, and returns the +// opaque URL-safe token. +func (c *SignedCodec) Encode(v any) (string, error) { + payload, err := json.Marshal(v) + if err != nil { + return "", err + } + sig := c.mac(payload) + return base64.RawURLEncoding.EncodeToString(payload) + "." + + base64.RawURLEncoding.EncodeToString(sig), nil +} + +// Decode verifies the token's signature against the configured key and, only if +// it matches, unmarshals the payload into v (a pointer). +// +// It returns ErrInvalidCursor when the token is structurally malformed and +// ErrTampered when the structure is valid but the signature does not verify. +// Verification uses hmac.Equal for constant-time comparison. +func (c *SignedCodec) Decode(token string, v any) error { + dot := -1 + for i := 0; i < len(token); i++ { + if token[i] == '.' { + dot = i + break + } + } + if dot < 0 { + return ErrInvalidCursor + } + + payload, err := base64.RawURLEncoding.DecodeString(token[:dot]) + if err != nil { + return ErrInvalidCursor + } + sig, err := base64.RawURLEncoding.DecodeString(token[dot+1:]) + if err != nil { + return ErrInvalidCursor + } + + if !hmac.Equal(sig, c.mac(payload)) { + return ErrTampered + } + if err := json.Unmarshal(payload, v); err != nil { + return ErrInvalidCursor + } + return nil +} diff --git a/sign_test.go b/sign_test.go new file mode 100644 index 0000000..94f5e3e --- /dev/null +++ b/sign_test.go @@ -0,0 +1,112 @@ +package cursor + +import ( + "strings" + "testing" +) + +func TestSignedRoundTrip(t *testing.T) { + c := NewSignedCodec([]byte("super-secret-signing-key-0123456789")) + in := key{CreatedAt: 1718200000, ID: "u_42"} + + tok, err := c.Encode(in) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(tok, ".") { + t.Fatalf("expected payload.sig format, got %q", tok) + } + + var out key + if err := c.Decode(tok, &out); err != nil { + t.Fatal(err) + } + if out != in { + t.Fatalf("round trip mismatch: %+v vs %+v", out, in) + } +} + +func TestSignedTamperPayload(t *testing.T) { + c := NewSignedCodec([]byte("super-secret-signing-key-0123456789")) + tok, err := c.Encode(key{CreatedAt: 1, ID: "a"}) + if err != nil { + t.Fatal(err) + } + + // Flip a character in the payload half (before the dot). + b := []byte(tok) + if b[0] == 'A' { + b[0] = 'B' + } else { + b[0] = 'A' + } + tampered := string(b) + + var out key + if err := c.Decode(tampered, &out); err != ErrTampered { + t.Fatalf("want ErrTampered, got %v", err) + } +} + +func TestSignedTamperSignature(t *testing.T) { + c := NewSignedCodec([]byte("super-secret-signing-key-0123456789")) + tok, err := c.Encode(key{CreatedAt: 1, ID: "a"}) + if err != nil { + t.Fatal(err) + } + + b := []byte(tok) + last := len(b) - 1 + if b[last] == 'A' { + b[last] = 'B' + } else { + b[last] = 'A' + } + + var out key + if err := c.Decode(string(b), &out); err != ErrTampered { + t.Fatalf("want ErrTampered, got %v", err) + } +} + +func TestSignedWrongKeyRejected(t *testing.T) { + signer := NewSignedCodec([]byte("key-A-key-A-key-A-key-A-key-A-key-A")) + verifier := NewSignedCodec([]byte("key-B-key-B-key-B-key-B-key-B-key-B")) + + tok, err := signer.Encode(key{CreatedAt: 7, ID: "z"}) + if err != nil { + t.Fatal(err) + } + var out key + if err := verifier.Decode(tok, &out); err != ErrTampered { + t.Fatalf("want ErrTampered for wrong key, got %v", err) + } +} + +func TestSignedMalformed(t *testing.T) { + c := NewSignedCodec([]byte("k")) + var out key + if err := c.Decode("no-dot-here", &out); err != ErrInvalidCursor { + t.Fatalf("want ErrInvalidCursor for missing dot, got %v", err) + } + if err := c.Decode("!!!.###", &out); err != ErrInvalidCursor { + t.Fatalf("want ErrInvalidCursor for bad base64, got %v", err) + } +} + +func TestSignedKeyCopied(t *testing.T) { + raw := []byte("mutable-key-mutable-key-mutable!!") + c := NewSignedCodec(raw) + tok, err := c.Encode(key{CreatedAt: 1, ID: "a"}) + if err != nil { + t.Fatal(err) + } + // Mutate the caller's slice; the codec must be unaffected. + for i := range raw { + raw[i] = 0 + } + var out key + if err := c.Decode(tok, &out); err != nil { + t.Fatalf("decode failed after caller mutated key slice: %v", err) + } +}