id.go

  1package id
  2
  3import (
  4	"crypto/rand"
  5	"encoding/binary"
  6	"encoding/json"
  7	"fmt"
  8	"math"
  9	"strconv"
 10	"sync"
 11	"sync/atomic"
 12	"time"
 13)
 14
 15// holds a reference to either the default generator or a custom generator set through SetDefault
 16var gen atomic.Pointer[Generator]
 17
 18var NotExist = Key(-1)
 19
 20var (
 21	// NodeBits holds the number of bits to use for a node
 22	nodeBits uint8 = 10
 23	// StepBits holds the number of bits to use for a step. Total of node + step bits
 24	// should be <= 22.
 25	stepBits uint8 = 12
 26
 27	// Maximum number of nodes representable with the number of bits allocated to a node.
 28	// Default is 1024 with a default node bit length of 10 bits.
 29	MaxNodes int64 = pow(2, nodeBits)
 30)
 31
 32func init() {
 33	g, _ := NewGenerator()
 34	gen.Store(g)
 35}
 36
 37// SetDefault sets the package-level generator to a custom configuration
 38func SetDefault(g *Generator) {
 39	gen.Store(g)
 40}
 41
 42// Default returns the package-level generator
 43func Default() *Generator {
 44	return gen.Load()
 45}
 46
 47// New generates a new key using the default generator
 48func New() Key {
 49	return gen.Load().New()
 50}
 51
 52// Generator holds state for the ID sequence
 53type Generator struct {
 54	mu    sync.Mutex
 55	epoch time.Time
 56	time  int64
 57	node  int64
 58	step  int64
 59
 60	nodeBits uint8
 61	randBits uint8
 62}
 63
 64// Key is an 64-bit integer key based on the Snowflake algorithm
 65type Key int64
 66
 67// Option configures the generator
 68type Option func(g *Generator) error
 69
 70// WithNodeNumber identifies node n of some total number of nodes.  Total number of nodes must be less than
 71// MaxNodes.  If total nodes are smaller, the extra bits will be used to inject randomness into the generated
 72// key.  This can help prevent enumeration attacks when the server has low generation load and few nodes.
 73func WithNodeNumber(n int, ofTotal int) Option {
 74	return func(g *Generator) error {
 75		if n >= ofTotal {
 76			return fmt.Errorf("nodes must be numbered from 0...%d for %d total: got %d", ofTotal-1, ofTotal, n)
 77		}
 78		bits := uint8(math.Ceil(math.Log2(float64(ofTotal))))
 79		if bits > nodeBits {
 80			return fmt.Errorf("exceeds max nodes of %d", pow(2, nodeBits))
 81		}
 82		g.node = int64(n)
 83		g.nodeBits = bits
 84		g.randBits = 10 - bits
 85		return nil
 86	}
 87}
 88
 89// WithStartTime adjusts the start time of the generator.  Keys will be good for approximately 69 years
 90// before overflow.
 91func WithStartTime(t time.Time) Option {
 92	return func(g *Generator) error {
 93		now := time.Now()
 94		if t.After(now) {
 95			return fmt.Errorf("start time is in the future")
 96		}
 97		// force monotonic time
 98		g.epoch = t
 99		return nil
100	}
101}
102
103// NewGenerator returns a generator with the configured options.  If no special options
104// are needed, it returns the same generator as the package-level generator and with
105// parameters equivalent to a Snowflake ID generator with a custom start time
106// of December 15, 2024 midnight UTC.
107func NewGenerator(opts ...Option) (*Generator, error) {
108
109	g := &Generator{}
110
111	// defaults give you a standard snowflake generator with a more recent start time
112	defaults := []Option{WithNodeNumber(0, int(pow(2, nodeBits))), WithStartTime(time.Date(2024, time.December, 15, 0, 0, 0, 0, time.UTC))}
113
114	opts = append(defaults, opts...)
115	for _, opt := range opts {
116		if err := opt(g); err != nil {
117			return nil, fmt.Errorf("error configuring ID generator: %w", err)
118		}
119	}
120	return g, nil
121}
122
123// New creates and returns a unique snowflake ID
124// To help guarantee uniqueness
125// - Make sure your system is keeping accurate system time
126// - Make sure you never have multiple nodes running with the same node ID
127func (n *Generator) New() Key {
128	return n.new(time.Since(n.epoch).Milliseconds())
129}
130
131// new allows passing in a custom time for testing
132func (n *Generator) new(now int64) Key {
133	n.mu.Lock()
134	defer n.mu.Unlock()
135
136	if now == n.time {
137		// did step overflow allowable bits
138		n.step = (n.step + 1) & (-1 ^ (-1 << stepBits))
139
140		// overflow, wait until next millisecond
141		if n.step == 0 {
142			for now <= n.time {
143				now = time.Since(n.epoch).Milliseconds()
144			}
145		}
146	} else {
147		n.step = 0
148	}
149
150	n.time = now
151
152	return Key((now)<<(stepBits+nodeBits) | ((getRandBits(n.randBits, n.nodeBits) | n.node) << stepBits) | (n.step))
153}
154
155// getRandBits fills the unused node bits with random numbers for additional key diversity to
156// help stop enumeration attacks
157func getRandBits(randBits uint8, nodeBits uint8) int64 {
158	if randBits == 0 {
159		return 0
160	}
161	numBytes := int(math.Ceil(float64(randBits+nodeBits) / 8.0))
162	r := make([]byte, numBytes)
163	rand.Read(r)
164	out, _ := binary.Uvarint(r)
165
166	mask := int64(-1 ^ (-1 << nodeBits))
167	truncate := int64(-1 ^ (-1 << (randBits + nodeBits)))
168	return int64(out) & truncate &^ mask
169}
170
171// String representation of the key in base10
172func (id Key) String() string {
173	return strconv.FormatInt(int64(id), 10)
174}
175
176// Compact returns a base36 string
177func (id Key) Compact() string {
178	return strconv.FormatInt(int64(id), 36)
179}
180
181// Int64 returns the key as int64
182func (id Key) Int64() int64 {
183	return int64(id)
184}
185
186// MarshalJSON serializes the key as a base36-encoded string for compact,
187// precision-safe JSON transport.
188func (id Key) MarshalJSON() ([]byte, error) {
189	return json.Marshal(id.Compact())
190}
191
192// UnmarshalJSON parses a base36-encoded string back into a key.
193func (id *Key) UnmarshalJSON(data []byte) error {
194	var s string
195	if err := json.Unmarshal(data, &s); err != nil {
196		return err
197	}
198	k, err := FromCompact(s)
199	if err != nil {
200		return err
201	}
202	*id = k
203	return nil
204}
205
206// Time extracts the creation time of the key. If no generator is provided,
207// it will load the default generator.  The time will not be accurate if the key
208// was created by a generator with a different start time.
209func (id Key) Time(g ...*Generator) time.Time {
210	var ge *Generator
211	if len(g) > 0 {
212		ge = g[0]
213	} else {
214		ge = gen.Load()
215	}
216	t := id >> (stepBits + nodeBits)
217	return ge.epoch.Add(time.Duration(t * 1_000_000))
218}
219
220// FromCompact converts a Base36-encoded string to an ID
221func FromCompact(id string) (Key, error) {
222	i, err := strconv.ParseInt(id, 36, 64)
223	if err != nil {
224		return NotExist, err
225	}
226	return Key(i), nil
227}
228
229func pow(n int64, m uint8) int64 {
230	if m == 0 {
231		return 1
232	}
233
234	if m == 1 {
235		return n
236	}
237
238	result := n
239	for i := 2; i <= int(m); i++ {
240		result *= n
241	}
242	return result
243}