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}