import files, add README & LICENSE
This commit is contained in:
@@ -0,0 +1,558 @@
|
||||
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)
|
||||
}
|
||||
Reference in New Issue
Block a user