package cursor import ( "database/sql/driver" "encoding/json" "errors" "fmt" "strings" "sync" "time" "github.com/go-jet/jet/v2/mysql" ) var ( // ErrBadCursor is returned when a cursor is invalid. ErrBadCursorString = errors.New("bad cursor string") columnRegistry = make(map[ColumnKey]mysql.Column) columnRegistryLock sync.RWMutex ) // CursorOrderCol is an interface for SQL expression that support the ORDER BY clause. type CursorOrderCol interface { mysql.Column ASC() mysql.OrderByClause DESC() mysql.OrderByClause } // ColumnKey identifies a column and table in the database with a given value. type ColumnKey struct { Table string `json:"table"` Column string `json:"column"` } // NewColumnKey creates a new ColumnKey with the specified table and column names. func NewColumnKey(col mysql.Column) ColumnKey { return ColumnKey{ Table: col.TableName(), Column: col.Name(), } } // IsEmpty returns true if the ColumnKey is empty. func (k *ColumnKey) IsEmpty() bool { return k == nil || k.Table == "" || k.Column == "" } // String returns a string representation of the ColumnKey. func (k ColumnKey) String() string { return fmt.Sprintf("%s.%s", k.Table, k.Column) } // RegisterColumn registers one or more columns with the cursor registry. func RegisterColumn(columns ...mysql.Column) { columnRegistryLock.Lock() defer columnRegistryLock.Unlock() for _, col := range columns { columnRegistry[NewColumnKey(col)] = col } } // RegisterColumnList registers a list of columns with the cursor registry. func RegisterColumnList(columns mysql.ColumnList) { RegisterColumn([]mysql.Column(columns)...) } // GetColumn retrieves a column from the cursor registry. func GetColumn(table, column string) (mysql.Column, error) { columnRegistryLock.RLock() defer columnRegistryLock.RUnlock() key := ColumnKey{ Table: table, Column: column, } col, ok := columnRegistry[key] if !ok { return nil, fmt.Errorf("column %s.%s not registered", table, column) } return col, nil } // GetColumnByKey retrieves a column from the cursor registry by its ColumnKey. func GetColumnByKey(key ColumnKey) (mysql.Column, error) { columnRegistryLock.RLock() defer columnRegistryLock.RUnlock() col, ok := columnRegistry[key] if !ok { return nil, fmt.Errorf("column %s.%s not registered", key.Table, key.Column) } return col, nil } // OrderDirection represents the order direction of a cursor. type OrderDirection int const ( // OrderAscending sorts results in ascending order. OrderAscending OrderDirection = iota // OrderDescending sorts results in descending order. OrderDescending ) // MarshalJSON marshals the OrderDirection to JSON. func (od OrderDirection) MarshalJSON() ([]byte, error) { switch od { case OrderAscending: return []byte(`"ASC"`), nil case OrderDescending: return []byte(`"DESC"`), nil default: return nil, fmt.Errorf("invalid order direction: %d", od) } } // UnmarshalJSON unmarshals the OrderDirection from JSON. func (od *OrderDirection) UnmarshalJSON(data []byte) error { var dir string if err := json.Unmarshal(data, &dir); err != nil { return err } switch strings.ToUpper(dir) { case "ASC": *od = OrderAscending case "DESC": *od = OrderDescending default: return fmt.Errorf("invalid order direction: %s", dir) } return nil } // Compile-time checks to ensure value types implement ExprMarshaler var _ ExprMarshaler[mysql.StringExpression, mysql.ColumnString] = (*StringValue)(nil) var _ ExprMarshaler[mysql.IntegerExpression, mysql.ColumnInteger] = (*Int64Value)(nil) var _ ExprMarshaler[mysql.IntegerExpression, mysql.ColumnInteger] = (*Uint64Value)(nil) var _ ExprMarshaler[mysql.TimestampExpression, mysql.ColumnTimestamp] = (*TimestampValue)(nil) type ExprMarshaler[E mysql.Expression, C mysql.Column] interface { GenericExpr Expr() E Col() C } // GenericExpr is the untyped cursor value contract shared by index and order values. type GenericExpr interface { ColumnKey() ColumnKey IsEmpty() bool driver.Valuer } var _ GenericExpr = (*StringValue)(nil) var _ GenericExpr = (*Int64Value)(nil) var _ GenericExpr = (*Uint64Value)(nil) var _ GenericExpr = (*TimestampValue)(nil) type StringValue struct { Key ColumnKey `json:"key"` Val string `json:"val"` } func NewStringValue(val string, col mysql.ColumnString) *StringValue { return &StringValue{ Key: NewColumnKey(col), Val: val, } } func (s StringValue) Expr() mysql.StringExpression { return mysql.String(s.Val) } func (s StringValue) Col() mysql.ColumnString { col, err := GetColumnByKey(s.Key) if err != nil { panic(err) } colStr, ok := col.(mysql.ColumnString) if !ok { panic(fmt.Errorf("column %s.%s is not a string column", s.Key.Table, s.Key.Column)) } return colStr } func (s StringValue) Value() (driver.Value, error) { return s.Val, nil } func (s *StringValue) ColumnKey() ColumnKey { return s.Key } func (s *StringValue) IsEmpty() bool { return s == nil || s.Key.Table == "" || s.Key.Column == "" || s.Val == "" } type Int64Value struct { Key ColumnKey `json:"key"` Val int64 `json:"val"` } func NewInt64Value(val int64, col mysql.ColumnInteger) *Int64Value { return &Int64Value{ Key: NewColumnKey(col), Val: val, } } func (i Int64Value) Expr() mysql.IntegerExpression { return mysql.Int64(i.Val) } func (i Int64Value) Col() mysql.ColumnInteger { col, err := GetColumnByKey(i.Key) if err != nil { panic(err) } colInt, ok := col.(mysql.ColumnInteger) if !ok { panic(fmt.Errorf("column %s.%s is not an integer column", i.Key.Table, i.Key.Column)) } return colInt } func (i Int64Value) Value() (driver.Value, error) { return i.Val, nil } func (i *Int64Value) ColumnKey() ColumnKey { return i.Key } func (i *Int64Value) IsEmpty() bool { return i == nil || i.Key.Table == "" || i.Key.Column == "" } type Uint64Value struct { Key ColumnKey `json:"key"` Val uint64 `json:"val"` } func NewUint64Value(val uint64, col mysql.ColumnInteger) *Uint64Value { return &Uint64Value{ Key: NewColumnKey(col), Val: val, } } func (i Uint64Value) Expr() mysql.IntegerExpression { return mysql.Uint64(i.Val) } func (i Uint64Value) Col() mysql.ColumnInteger { col, err := GetColumnByKey(i.Key) if err != nil { panic(err) } colInt, ok := col.(mysql.ColumnInteger) if !ok { panic(fmt.Errorf("column %s.%s is not an integer column", i.Key.Table, i.Key.Column)) } return colInt } func (i Uint64Value) Value() (driver.Value, error) { return i.Val, nil } func (i *Uint64Value) ColumnKey() ColumnKey { return i.Key } func (i *Uint64Value) IsEmpty() bool { return i == nil || i.Key.Table == "" || i.Key.Column == "" } // TimestampValue stores a timestamp cursor position. type TimestampValue struct { Key ColumnKey `json:"key"` Val time.Time `json:"val"` } // NewTimestampValue creates a timestamp cursor value for the given column. func NewTimestampValue(val time.Time, col mysql.ColumnTimestamp) *TimestampValue { return &TimestampValue{ Key: NewColumnKey(col), Val: val, } } func (t TimestampValue) Expr() mysql.TimestampExpression { return mysql.TimestampT(t.Val) } func (t TimestampValue) Col() mysql.ColumnTimestamp { col, err := GetColumnByKey(t.Key) if err != nil { panic(err) } colTS, ok := col.(mysql.ColumnTimestamp) if !ok { panic(fmt.Errorf("column %s.%s is not a timestamp column", t.Key.Table, t.Key.Column)) } return colTS } func (t TimestampValue) Value() (driver.Value, error) { return t.Val, nil } func (t *TimestampValue) ColumnKey() ColumnKey { return t.Key } func (t *TimestampValue) IsEmpty() bool { return t == nil || t.Key.Table == "" || t.Key.Column == "" || t.Val.IsZero() } // GenericCursor is an interface for a cursor that can be used with any type of // expression and column. Only the methods that do not depend on the specific // types of expressions and columns are defined here. type GenericCursor interface { IsEmpty() bool IsComposite() bool UsesTupleOrdering() bool OrderCol() CursorOrderCol Encode() (string, error) Decode(src string) error String() string GenericIndex() GenericExpr GenericOrderValue() GenericExpr Direction() OrderDirection } // Compile-time check to ensure that Cursor implements the GenericCursor interface. var _ GenericCursor = (*Cursor[mysql.Expression, mysql.Column])(nil) // Cursor identifies a location and order in the database. // // Index is the stable position column (typically the primary key). OrderColumnKey // is the primary sort column shown to users. When they differ and OrderValue is // set, pagination uses lexicographic tuple comparison on (order_col, index_col). type Cursor[IE mysql.Expression, IC mysql.Column] struct { // Index identifies the column used for positioning the cursor and its current value. Index ExprMarshaler[IE, IC] `json:"index"` // OrderValue holds the order column value at this cursor position when using // composite tuple pagination. OrderValue GenericExpr `json:"order_val,omitempty"` // OrderColumnKey identifies the column used for ordering the results. OrderColumnKey ColumnKey `json:"order_col"` // OrderDir is the direction of the order (ASC or DESC). OrderDir OrderDirection `json:"order_dir"` } type cursorJSON struct { Index json.RawMessage `json:"index"` OrderValue json.RawMessage `json:"order_val"` OrderColumnKey ColumnKey `json:"order_col"` OrderDir OrderDirection `json:"order_dir"` } // NewCursor creates a new Cursor with the specified parameters. func NewCursor[IE mysql.Expression, IC mysql.Column](index ExprMarshaler[IE, IC], orderCol CursorOrderCol, orderDir OrderDirection) *Cursor[IE, IC] { return &Cursor[IE, IC]{ Index: index, OrderColumnKey: NewColumnKey(orderCol), OrderDir: orderDir, } } // NewCursorFromAfterPtr decodes a Relay-style after cursor. Nil or empty after // returns (nil, nil). func NewCursorFromAfterPtr[IE mysql.Expression, IC mysql.Column]( newZero func() *Cursor[IE, IC], after *string, ) (*Cursor[IE, IC], error) { if after == nil || *after == "" { return nil, nil } c := newZero() return c, c.Decode(*after) } // NewCursorFromJSON returns a Cursor from a JSON representation. func NewCursorFromJSON[IE mysql.Expression, IC mysql.Column](zeroIndex ExprMarshaler[IE, IC], src []byte) (*Cursor[IE, IC], error) { cursor := Cursor[IE, IC]{ Index: zeroIndex, } if err := decodeCursorJSON(&cursor, src, false, OrderAscending); err != nil { return nil, err } return &cursor, nil } // CopyWithVal returns a new cursor with the specified value and the current ordering. func (c *Cursor[IE, IC]) CopyWithVal(val ExprMarshaler[IE, IC]) *Cursor[IE, IC] { return NewCursor(val, c.OrderCol(), c.OrderDir) } // CopyWithVals returns a new cursor with the specified index and order values. // Use this when IsComposite() is true. func (c *Cursor[IE, IC]) CopyWithVals(index ExprMarshaler[IE, IC], orderVal GenericExpr) *Cursor[IE, IC] { result := NewCursor(index, c.OrderCol(), c.OrderDir) result.OrderValue = orderVal return result } // UsesTupleOrdering reports whether results are sorted by (order_col, index_col). // Unlike IsComposite, this does not require OrderValue and applies to default // cursors on the first page. func (c *Cursor[IE, IC]) UsesTupleOrdering() bool { if c == nil || c.Index == nil { return false } indexKey := c.Index.ColumnKey() if indexKey.IsEmpty() || c.OrderColumnKey.IsEmpty() { return false } return c.OrderColumnKey != indexKey } // IsComposite reports whether pagination uses (order_col, index_col) tuple comparison. func (c *Cursor[IE, IC]) IsComposite() bool { if c == nil || c.Index == nil || c.Index.IsEmpty() { return false } if c.OrderColumnKey == c.Index.ColumnKey() { return false } return c.OrderValue != nil && !c.OrderValue.IsEmpty() } // IsEmpty returns true if the cursor or any of its keys are empty or unmapped. func (c *Cursor[IE, IC]) IsEmpty() bool { if c == nil || c.Index == nil || c.Index.IsEmpty() { return true } _, err := GetColumnByKey(c.OrderColumnKey) return err != nil } // GenericIndex returns the Index ExprMarshaler as a GenericExpr. func (c *Cursor[IE, IC]) GenericIndex() GenericExpr { return c.Index } // GenericOrderValue returns the order column value when using composite pagination. func (c *Cursor[IE, IC]) GenericOrderValue() GenericExpr { return c.OrderValue } // Direction returns the order direction expected by the cursor. func (c *Cursor[IE, IC]) Direction() OrderDirection { return c.OrderDir } // OrderCol returns the column used for ordering the results. Panics if the // column is not registered. func (c *Cursor[IE, IC]) OrderCol() CursorOrderCol { col, err := GetColumnByKey(c.OrderColumnKey) if err != nil { panic(err) } return col } // Encode returns a stringified JSON representation of the cursor. func (c *Cursor[IE, IC]) Encode() (string, error) { bytes, err := json.Marshal(c) if err != nil { return "", fmt.Errorf("failed to marshal cursor: %w", err) } return string(bytes), nil } // Decode decodes a stringified JSON representation of the cursor into this object. func (c *Cursor[IE, IC]) Decode(src string) error { if c == nil { return fmt.Errorf("cursor is nil") } return decodeCursorJSON(c, []byte(src), false, OrderAscending) } // DecodeAndOrder decodes a stringified JSON representation of the cursor into // this object and applies a new order direction. func (c *Cursor[IE, IC]) DecodeAndOrder(src string, orderDir OrderDirection) error { return decodeCursorJSON(c, []byte(src), true, orderDir) } func decodeCursorJSON[IE mysql.Expression, IC mysql.Column]( c *Cursor[IE, IC], src []byte, overrideDir bool, orderDir OrderDirection, ) error { var raw cursorJSON if err := json.Unmarshal(src, &raw); err != nil { return fmt.Errorf("failed to unmarshal cursor: %w", err) } if err := json.Unmarshal(raw.Index, c.Index); err != nil { return fmt.Errorf("failed to unmarshal cursor index: %w", err) } c.OrderColumnKey = raw.OrderColumnKey c.OrderDir = raw.OrderDir if overrideDir { c.OrderDir = orderDir } if len(raw.OrderValue) > 0 { orderVal, err := unmarshalGenericExpr(raw.OrderValue, raw.OrderColumnKey) if err != nil { return fmt.Errorf("failed to unmarshal cursor order value: %w", err) } c.OrderValue = orderVal } else { c.OrderValue = nil } return nil } func unmarshalGenericExpr(data []byte, key ColumnKey) (GenericExpr, error) { col, err := GetColumnByKey(key) if err != nil { return nil, err } switch col.(type) { case mysql.ColumnTimestamp: var v TimestampValue if err := json.Unmarshal(data, &v); err != nil { return nil, err } return &v, nil case mysql.ColumnInteger: var v Uint64Value if err := json.Unmarshal(data, &v); err != nil { return nil, err } return &v, nil case mysql.ColumnString: var v StringValue if err := json.Unmarshal(data, &v); err != nil { return nil, err } return &v, nil default: return nil, fmt.Errorf("unsupported order column type for %s", key) } } // String returns a string representation of the Cursor. func (c *Cursor[IE, IC]) String() string { bytes, err := json.Marshal(c) if err != nil { return "" } return string(bytes) }