From 46de2651863ec786b8e01922550df13cc830638f Mon Sep 17 00:00:00 2001 From: hpower2 Date: Fri, 12 Jun 2026 01:46:34 +0700 Subject: [PATCH 1/5] feat(seek): add two-column seek-predicate builder with golden SQL tests Add a typed composite sort key (SeekKey/Column/Direction) that emits the row-wise seek WHERE predicate and the matching ORDER BY from a key with per-column ASC/DESC. The predicate expands into the portable lexicographic form ((a > ?) OR (a = ? AND b > ?)) so it works on every SQL engine, and supports both database/sql "?" and pgx "$N" placeholders via Dialect. Backward paging flips each column's effective direction (and comparison operator) so the database walks toward the previous page. Args() expands cursor values to match the placeholders. Golden tests assert exact SQL strings for forward/backward and ASC/DESC combinations. Closes #1 --- seek.go | 160 +++++++++++++++++++++++++++++++++++++++++++++++++++ seek_test.go | 138 ++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 298 insertions(+) create mode 100644 seek.go create mode 100644 seek_test.go 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, "") + } +} From 7849b58dd83b04d6c504ba829c63cc3b0c7beb9e Mon Sep 17 00:00:00 2001 From: hpower2 Date: Fri, 12 Jun 2026 01:46:39 +0700 Subject: [PATCH 2/5] feat(sign): add HMAC-SHA256 signed cursor codec Add SignedCodec, which encodes cursors as base64url(payload).base64url(hmac) using crypto/hmac with SHA-256 so tokens cannot be forged or probed. Decode verifies the signature with a constant-time hmac.Equal before unmarshalling, returning ErrTampered for altered or wrong-key tokens and ErrInvalidCursor for structurally malformed ones. The key is copied on construction so later mutation of the caller's slice is harmless. Tests cover round-trip, payload/signature tampering, wrong-key rejection, malformed input, and key isolation. Closes #2 --- sign.go | 90 +++++++++++++++++++++++++++++++++++++++++ sign_test.go | 112 +++++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 202 insertions(+) create mode 100644 sign.go create mode 100644 sign_test.go 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) + } +} From 015a0b911a65bf685c849e6fbf003447845ca0d7 Mon Sep 17 00:00:00 2001 From: hpower2 Date: Fri, 12 Jun 2026 01:46:45 +0700 Subject: [PATCH 3/5] feat(http): add net/http query-param helper with limit clamping Add ParseParams, which reads ?cursor=&limit= from an *http.Request and returns a PageParams with the limit clamped per LimitConfig: missing, empty, non-numeric, or non-positive limits fall back to the default; values above the max are clamped down. A zero-value LimitConfig stays safe (default 20, max 100) and a default larger than max is itself clamped. Param names are configurable. httptest-based tests cover over-cap, missing, zero, negative, non-numeric, and custom-param cases. Closes #3 --- httpparam.go | 80 +++++++++++++++++++++++++++++++++++++++++++++++ httpparam_test.go | 75 ++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 155 insertions(+) create mode 100644 httpparam.go create mode 100644 httpparam_test.go 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) + } +} From 77daaab9988600de5227ea3509f68c6a02fa6548 Mon Sep 17 00:00:00 2001 From: hpower2 Date: Fri, 12 Jun 2026 01:46:51 +0700 Subject: [PATCH 4/5] feat(graphql): add Relay pageInfo + edges connection adapter Add Connect, which builds a Relay-style Connection (edges + pageInfo) from a page of items, a per-item cursor function, and hasNext/hasPrev flags. PageInfo carries startCursor, endCursor, hasNextPage, and hasPreviousPage with JSON tags matching the GraphQL Cursor Connections spec; start/end are derived from the first/last edges and blank for an empty page. Pure structs, no deps. Tests cover basic, empty, single, multi-item, and JSON-shape (spec field names). Closes #4 --- graphql.go | 52 +++++++++++++++++++++++++ graphql_test.go | 100 ++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 152 insertions(+) create mode 100644 graphql.go create mode 100644 graphql_test.go 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) + } +} From f3ac5fedf976afae693bfaff63fba39b0165b0e4 Mon Sep 17 00:00:00 2001 From: hpower2 Date: Fri, 12 Jun 2026 01:46:57 +0700 Subject: [PATCH 5/5] docs(example): add REST-paginating-Postgres example + v0.2 docs/CI Add example_postgres_test.go: a runnable Example that uses the seek-predicate builder to paginate a (created_at, id)-sorted table behind an httptest handler, walking three pages against an in-memory fake data source (no live DB) and printing the exact pgx SQL the builder emits. Document all v0.2 helpers in the README, add an Unreleased CHANGELOG entry, and extend CI with a "go mod tidy is clean" check while keeping the Go 1.22 + 1.23 matrix and the gofmt/vet/test -race gates. Closes #5 --- .github/workflows/ci.yml | 10 +++ CHANGELOG.md | 19 +++++ README.md | 43 ++++++++++ example_postgres_test.go | 164 +++++++++++++++++++++++++++++++++++++++ 4 files changed, 236 insertions(+) create mode 100644 example_postgres_test.go 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 +}