114 lines
3.9 KiB
Go
114 lines
3.9 KiB
Go
package domain
|
|
|
|
import (
|
|
"crypto/hmac"
|
|
"crypto/sha256"
|
|
"encoding/base64"
|
|
"encoding/binary"
|
|
"errors"
|
|
"fmt"
|
|
)
|
|
|
|
var (
|
|
ErrInvalidCursor = errors.New("invalid pagination cursor")
|
|
ErrCursorFilterMismatch = errors.New("pagination cursor does not match query filters")
|
|
)
|
|
|
|
const (
|
|
cursorVersion = 1
|
|
maxCursorFieldBytes = 1024
|
|
cursorMACBytes = sha256.Size
|
|
)
|
|
|
|
type CursorKind uint8
|
|
|
|
const (
|
|
CursorKindCommands CursorKind = iota + 1
|
|
CursorKindOutput
|
|
CursorKindIncidents
|
|
CursorKindClients
|
|
)
|
|
|
|
type Cursor struct {
|
|
Kind CursorKind
|
|
FilterHash [sha256.Size]byte
|
|
Position []byte
|
|
SnapshotBoundary []byte
|
|
}
|
|
|
|
type CursorCodec struct{ key []byte }
|
|
|
|
func NewCursorCodec(key []byte) (*CursorCodec, error) {
|
|
if len(key) < sha256.Size {
|
|
return nil, fmt.Errorf("%w: signing key must be at least %d bytes", ErrInvalidCursor, sha256.Size)
|
|
}
|
|
return &CursorCodec{key: append([]byte(nil), key...)}, nil
|
|
}
|
|
|
|
func HashCursorFilters(canonical []byte) [sha256.Size]byte { return sha256.Sum256(canonical) }
|
|
|
|
func (codec *CursorCodec) Encode(cursor Cursor) (string, error) {
|
|
if !validCursorKind(cursor.Kind) || len(cursor.Position) == 0 || len(cursor.Position) > maxCursorFieldBytes || len(cursor.SnapshotBoundary) == 0 || len(cursor.SnapshotBoundary) > maxCursorFieldBytes {
|
|
return "", ErrInvalidCursor
|
|
}
|
|
payloadLength := 1 + 1 + sha256.Size + 2 + len(cursor.Position) + 2 + len(cursor.SnapshotBoundary)
|
|
payload := make([]byte, payloadLength, payloadLength+cursorMACBytes)
|
|
payload[0] = cursorVersion
|
|
payload[1] = byte(cursor.Kind)
|
|
copy(payload[2:], cursor.FilterHash[:])
|
|
offset := 2 + sha256.Size
|
|
binary.BigEndian.PutUint16(payload[offset:], uint16(len(cursor.Position)))
|
|
offset += 2
|
|
copy(payload[offset:], cursor.Position)
|
|
offset += len(cursor.Position)
|
|
binary.BigEndian.PutUint16(payload[offset:], uint16(len(cursor.SnapshotBoundary)))
|
|
offset += 2
|
|
copy(payload[offset:], cursor.SnapshotBoundary)
|
|
mac := hmac.New(sha256.New, codec.key)
|
|
_, _ = mac.Write(payload)
|
|
return base64.RawURLEncoding.EncodeToString(append(payload, mac.Sum(nil)...)), nil
|
|
}
|
|
|
|
func (codec *CursorCodec) Decode(token string, expectedKind CursorKind, expectedFilterHash [sha256.Size]byte) (Cursor, error) {
|
|
if token == "" || len(token) > base64.RawURLEncoding.EncodedLen(2*maxCursorFieldBytes+128) {
|
|
return Cursor{}, ErrInvalidCursor
|
|
}
|
|
decoded, err := base64.RawURLEncoding.DecodeString(token)
|
|
if err != nil || len(decoded) < 1+1+sha256.Size+2+1+2+1+cursorMACBytes {
|
|
return Cursor{}, ErrInvalidCursor
|
|
}
|
|
payload, suppliedMAC := decoded[:len(decoded)-cursorMACBytes], decoded[len(decoded)-cursorMACBytes:]
|
|
mac := hmac.New(sha256.New, codec.key)
|
|
_, _ = mac.Write(payload)
|
|
if !hmac.Equal(suppliedMAC, mac.Sum(nil)) {
|
|
return Cursor{}, ErrInvalidCursor
|
|
}
|
|
if payload[0] != cursorVersion || !validCursorKind(CursorKind(payload[1])) || CursorKind(payload[1]) != expectedKind {
|
|
return Cursor{}, ErrInvalidCursor
|
|
}
|
|
var filterHash [sha256.Size]byte
|
|
copy(filterHash[:], payload[2:2+sha256.Size])
|
|
if !hmac.Equal(filterHash[:], expectedFilterHash[:]) {
|
|
return Cursor{}, ErrCursorFilterMismatch
|
|
}
|
|
offset := 2 + sha256.Size
|
|
positionLength := int(binary.BigEndian.Uint16(payload[offset:]))
|
|
offset += 2
|
|
if positionLength == 0 || positionLength > maxCursorFieldBytes || offset+positionLength+2 > len(payload) {
|
|
return Cursor{}, ErrInvalidCursor
|
|
}
|
|
position := append([]byte(nil), payload[offset:offset+positionLength]...)
|
|
offset += positionLength
|
|
boundaryLength := int(binary.BigEndian.Uint16(payload[offset:]))
|
|
offset += 2
|
|
if boundaryLength == 0 || boundaryLength > maxCursorFieldBytes || offset+boundaryLength != len(payload) {
|
|
return Cursor{}, ErrInvalidCursor
|
|
}
|
|
boundary := append([]byte(nil), payload[offset:]...)
|
|
return Cursor{Kind: expectedKind, FilterHash: filterHash, Position: position, SnapshotBoundary: boundary}, nil
|
|
}
|
|
|
|
func validCursorKind(kind CursorKind) bool {
|
|
return kind >= CursorKindCommands && kind <= CursorKindClients
|
|
}
|