683b0ddbf4
Prepare the extracted library for external consumption with grouped godoc, extension point docs, and setup requirements. Wire ErrBadCursorString into decode paths and drop redundant Paginate*Conds aliases. Co-authored-by: Cursor <cursoragent@cursor.com>
262 lines
7.4 KiB
Go
262 lines
7.4 KiB
Go
package cursor
|
|
|
|
import (
|
|
"errors"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"gitea.auvem.com/go-toolkit/dbx"
|
|
"github.com/go-jet/jet/v2/mysql"
|
|
"github.com/stretchr/testify/assert"
|
|
"github.com/stretchr/testify/require"
|
|
)
|
|
|
|
func TestConnectionFromRelayArgs(t *testing.T) {
|
|
assert := assert.New(t)
|
|
require := require.New(t)
|
|
RegisterColumn(User.ID)
|
|
|
|
newZero := func() *Cursor[mysql.IntegerExpression, mysql.ColumnInteger] {
|
|
return NewCursor(NewInt64Value(0, User.ID), User.ID, OrderDescending)
|
|
}
|
|
|
|
var gotCursor *Cursor[mysql.IntegerExpression, mysql.ColumnInteger]
|
|
var gotLimit int
|
|
|
|
list := func(c *Cursor[mysql.IntegerExpression, mysql.ColumnInteger], limit int) (*Connection[int], error) {
|
|
gotCursor = c
|
|
gotLimit = limit
|
|
return &Connection[int]{}, nil
|
|
}
|
|
|
|
conn, err := ConnectionFromRelayArgs(nil, nil, newZero, list)
|
|
require.NoError(err)
|
|
require.NotNil(conn)
|
|
assert.Nil(gotCursor)
|
|
assert.Equal(0, gotLimit)
|
|
|
|
first := 25
|
|
conn, err = ConnectionFromRelayArgs(nil, &first, newZero, list)
|
|
require.NoError(err)
|
|
require.NotNil(conn)
|
|
assert.Nil(gotCursor)
|
|
assert.Equal(25, gotLimit)
|
|
|
|
encoded, err := NewCursor(
|
|
NewInt64Value(7, User.ID),
|
|
User.ID,
|
|
OrderDescending,
|
|
).Encode()
|
|
require.NoError(err)
|
|
|
|
conn, err = ConnectionFromRelayArgs(&encoded, &first, newZero, list)
|
|
require.NoError(err)
|
|
require.NotNil(conn)
|
|
require.NotNil(gotCursor)
|
|
assert.Equal(int64(7), gotCursor.Index.(*Int64Value).Val)
|
|
assert.Equal(25, gotLimit)
|
|
|
|
bad := "{not-json"
|
|
_, err = ConnectionFromRelayArgs(&bad, nil, newZero, list)
|
|
require.Error(err)
|
|
require.ErrorIs(err, ErrBadCursorString)
|
|
|
|
listErr := errors.New("list failed")
|
|
_, err = ConnectionFromRelayArgs(nil, nil, newZero, func(*Cursor[mysql.IntegerExpression, mysql.ColumnInteger], int) (*Connection[int], error) {
|
|
return nil, listErr
|
|
})
|
|
require.ErrorIs(err, listErr)
|
|
}
|
|
|
|
func TestPageQueryRunSQLComposite(t *testing.T) {
|
|
assert := assert.New(t)
|
|
registerMeetingCursorColumns(t)
|
|
|
|
when := time.Date(2026, 6, 1, 9, 0, 0, 0, time.UTC)
|
|
active := NewCursor(
|
|
NewStringValue("", Meeting.ID),
|
|
Meeting.StartTime,
|
|
OrderDescending,
|
|
).CopyWithVals(
|
|
NewStringValue("meeting-uuid-1", Meeting.ID),
|
|
NewTimestampValue(when, Meeting.StartTime),
|
|
)
|
|
|
|
stmt := mysql.SELECT(Meeting.AllColumns).FROM(Meeting)
|
|
stmt.WHERE(mysql.Bool(true).AND(PaginateConds(active)))
|
|
stmt.ORDER_BY(OrderByClauses(active)...)
|
|
sql := stmt.DebugSql()
|
|
|
|
assert.Contains(sql, "(meeting.start_time, meeting.id) < ('")
|
|
assert.True(
|
|
strings.Contains(sql, "start_time DESC") && strings.Contains(sql, "id DESC"),
|
|
"expected composite order by, got: %s", sql,
|
|
)
|
|
}
|
|
|
|
func TestPageQueryRunSQLSimple(t *testing.T) {
|
|
assert := assert.New(t)
|
|
RegisterColumn(MeetingRoom.ID)
|
|
|
|
active := NewCursor(
|
|
NewStringValue("room-1", MeetingRoom.ID),
|
|
MeetingRoom.ID,
|
|
OrderDescending,
|
|
)
|
|
|
|
stmt := mysql.SELECT(MeetingRoom.AllColumns).FROM(MeetingRoom)
|
|
stmt.WHERE(mysql.Bool(true).AND(PaginateConds(active)))
|
|
stmt.ORDER_BY(OrderByClauses(active)...)
|
|
sql := stmt.DebugSql()
|
|
|
|
assert.Contains(sql, "meeting_room.id < 'room-1'")
|
|
assert.Contains(sql, "id DESC")
|
|
}
|
|
|
|
func TestPageQueryRunSQLFirstPageNilCursor(t *testing.T) {
|
|
assert := assert.New(t)
|
|
registerMeetingCursorColumns(t)
|
|
|
|
defaultCursor := NewCursor(
|
|
NewStringValue("", Meeting.ID),
|
|
Meeting.StartTime,
|
|
OrderDescending,
|
|
)
|
|
|
|
var nilCursor *Cursor[mysql.StringExpression, mysql.ColumnString]
|
|
stmt := mysql.SELECT(Meeting.AllColumns).FROM(Meeting)
|
|
stmt.WHERE(mysql.Bool(true).AND(PaginateConds(nilCursor)))
|
|
stmt.ORDER_BY(OrderByClauses(defaultCursor)...)
|
|
sql := stmt.DebugSql()
|
|
|
|
assert.NotContains(sql, "meeting.id < ''")
|
|
assert.Contains(sql, "start_time DESC")
|
|
assert.Contains(sql, "id DESC")
|
|
}
|
|
|
|
func TestPageQueryRunSQLFirstPageDefaultCursor(t *testing.T) {
|
|
assert := assert.New(t)
|
|
registerMeetingCursorColumns(t)
|
|
|
|
defaultCursor := NewCursor(
|
|
NewStringValue("", Meeting.ID),
|
|
Meeting.StartTime,
|
|
OrderDescending,
|
|
)
|
|
|
|
stmt := mysql.SELECT(Meeting.AllColumns).FROM(Meeting)
|
|
stmt.WHERE(mysql.Bool(true).AND(PaginateConds(defaultCursor)))
|
|
stmt.ORDER_BY(OrderByClauses(defaultCursor)...)
|
|
sql := stmt.DebugSql()
|
|
|
|
assert.NotContains(sql, "meeting.id < ''")
|
|
assert.Contains(sql, "start_time DESC")
|
|
assert.Contains(sql, "id DESC")
|
|
}
|
|
|
|
func TestPageQueryRunSQLFirstPageRoomDefaultCursor(t *testing.T) {
|
|
assert := assert.New(t)
|
|
RegisterColumn(MeetingRoom.ID)
|
|
|
|
defaultCursor := NewCursor(
|
|
NewStringValue("", MeetingRoom.ID),
|
|
MeetingRoom.ID,
|
|
OrderDescending,
|
|
)
|
|
|
|
stmt := mysql.SELECT(MeetingRoom.AllColumns).FROM(MeetingRoom)
|
|
stmt.WHERE(mysql.Bool(true).AND(PaginateConds(defaultCursor)))
|
|
stmt.ORDER_BY(OrderByClauses(defaultCursor)...)
|
|
sql := stmt.DebugSql()
|
|
|
|
assert.NotContains(sql, "meeting_room.id < ''")
|
|
assert.Contains(sql, "id DESC")
|
|
}
|
|
|
|
func TestPageQueryRun(t *testing.T) {
|
|
RegisterColumn(User.ID)
|
|
|
|
type row struct {
|
|
ID int64
|
|
}
|
|
|
|
defaultCursor := func() *Cursor[mysql.IntegerExpression, mysql.ColumnInteger] {
|
|
return NewCursor(NewInt64Value(0, User.ID), User.ID, OrderDescending)
|
|
}
|
|
|
|
items := []*row{{ID: 1}, {ID: 2}}
|
|
scanErr := errors.New("scan failed")
|
|
afterScanErr := errors.New("after scan failed")
|
|
|
|
t.Run("success with default cursor and after scan", func(t *testing.T) {
|
|
var scannedStmt mysql.SelectStatement
|
|
q := PageQuery[row, mysql.IntegerExpression, mysql.ColumnInteger]{
|
|
Stmt: User.SELECT(User.AllColumns).FROM(User),
|
|
Conds: mysql.Bool(true),
|
|
Default: defaultCursor,
|
|
Limit: 10,
|
|
CountFn: func(_ dbx.Queryable, _, _ GenericCursor) (QueryCountResult, error) {
|
|
return QueryCountResult{Total: 2, After: 1, Before: 0}, nil
|
|
},
|
|
Scan: func(stmt mysql.SelectStatement, dest *[]*row) error {
|
|
scannedStmt = stmt
|
|
*dest = items
|
|
return nil
|
|
},
|
|
AfterScan: func(rows []*row) ([]*row, error) {
|
|
return rows[:1], nil
|
|
},
|
|
ToEdge: func(_ *Cursor[mysql.IntegerExpression, mysql.ColumnInteger], item *row) (GenericCursor, error) {
|
|
return NewCursor(NewInt64Value(item.ID, User.ID), User.ID, OrderDescending), nil
|
|
},
|
|
}
|
|
|
|
conn, err := q.Run()
|
|
require.NoError(t, err)
|
|
require.NotNil(t, conn)
|
|
assert.Len(t, conn.Edges, 1)
|
|
assert.True(t, conn.PageInfo.HasNextPage)
|
|
assert.False(t, conn.PageInfo.HasPreviousPage)
|
|
assert.Equal(t, 2, conn.TotalCount)
|
|
assert.Contains(t, scannedStmt.DebugSql(), "LIMIT 10")
|
|
})
|
|
|
|
t.Run("scan error", func(t *testing.T) {
|
|
q := PageQuery[row, mysql.IntegerExpression, mysql.ColumnInteger]{
|
|
Stmt: User.SELECT(User.AllColumns).FROM(User),
|
|
Conds: mysql.Bool(true),
|
|
Default: defaultCursor,
|
|
Scan: func(_ mysql.SelectStatement, _ *[]*row) error {
|
|
return scanErr
|
|
},
|
|
ToEdge: func(_ *Cursor[mysql.IntegerExpression, mysql.ColumnInteger], item *row) (GenericCursor, error) {
|
|
return NewCursor(NewInt64Value(item.ID, User.ID), User.ID, OrderDescending), nil
|
|
},
|
|
}
|
|
_, err := q.Run()
|
|
require.Error(t, err)
|
|
assert.Contains(t, err.Error(), "failed to run paginated query")
|
|
})
|
|
|
|
t.Run("after scan error", func(t *testing.T) {
|
|
q := PageQuery[row, mysql.IntegerExpression, mysql.ColumnInteger]{
|
|
Stmt: User.SELECT(User.AllColumns).FROM(User),
|
|
Conds: mysql.Bool(true),
|
|
Default: defaultCursor,
|
|
Scan: func(_ mysql.SelectStatement, dest *[]*row) error {
|
|
*dest = items
|
|
return nil
|
|
},
|
|
AfterScan: func(_ []*row) ([]*row, error) {
|
|
return nil, afterScanErr
|
|
},
|
|
ToEdge: func(_ *Cursor[mysql.IntegerExpression, mysql.ColumnInteger], item *row) (GenericCursor, error) {
|
|
return NewCursor(NewInt64Value(item.ID, User.ID), User.ID, OrderDescending), nil
|
|
},
|
|
}
|
|
_, err := q.Run()
|
|
require.ErrorIs(t, err, afterScanErr)
|
|
})
|
|
}
|