update
This commit is contained in:
126
internal/pkg/hashid/hashid.go
Normal file
126
internal/pkg/hashid/hashid.go
Normal file
@ -0,0 +1,126 @@
|
||||
// Package hashid 把自增主键(uint32)编码为无序、URL 安全的短串,
|
||||
// 用于公开接口对外暴露,避免爬虫按 1,2,3... 顺序枚举全部文章/品牌。
|
||||
//
|
||||
// 设计要点:
|
||||
// - 内部仍用数字主键,仅对外序列化时编码、入参时解码,DB 与内部逻辑完全不变。
|
||||
// - 编码基于 32-bit 平衡 Feistel 网络(密钥由部署盐值派生)+ base62,
|
||||
// 是真实双射:Decode(Encode(n)) == n 严格成立,解码即可还原主键。
|
||||
// - 非顺序:相邻 id 的编码结果无规律,无法 +1 遍历;盐值不同编码结果不同。
|
||||
// - 零外部依赖;字母表 0-9a-zA-Z 全部 URL 安全。
|
||||
package hashid
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
alphabet = "0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ"
|
||||
base = 62
|
||||
minLen = 8
|
||||
)
|
||||
|
||||
var (
|
||||
keys [4]uint32
|
||||
inited bool
|
||||
)
|
||||
|
||||
// Init 用部署级盐值初始化混淆密钥。盐值为空时使用内置默认(仅防顺序枚举,不算安全)。
|
||||
// 必须在服务启动时调用一次(main.go 加载配置后)。
|
||||
func Init(secret string) {
|
||||
if secret == "" {
|
||||
secret = "fashion-archive-default-salt-change-me"
|
||||
}
|
||||
h := sha256.Sum256([]byte(secret))
|
||||
for i := 0; i < 4; i++ {
|
||||
keys[i] = binary.BigEndian.Uint32(h[i*4 : i*4+4])
|
||||
}
|
||||
inited = true
|
||||
}
|
||||
|
||||
func ensure() {
|
||||
if !inited {
|
||||
Init("")
|
||||
}
|
||||
}
|
||||
|
||||
// feistel 32-bit 平衡 Feistel 网络。encrypt=true 加密,false 解密。
|
||||
// Feistel 网络的逆只需逆序执行轮函数,因此无论 round function 是否可逆都能精确还原。
|
||||
// 加密轮:f 作用于右半块,左下一 = 右、右下一 = 左 ^ f(右)。
|
||||
// 解密轮:f 作用于左半块(解密时左半块即上一轮的右半块),右下一 = 左、左下一 = 右 ^ f(左)。
|
||||
func feistel(v uint32, encrypt bool) uint32 {
|
||||
const rounds = 8
|
||||
l, r := uint16(v>>16), uint16(v&0xffff)
|
||||
for i := 0; i < rounds; i++ {
|
||||
idx := i
|
||||
if !encrypt {
|
||||
idx = rounds - 1 - i
|
||||
}
|
||||
// round function:乘法扩散 + 密钥混合 + 高地位混淆,输出取低 16 位
|
||||
round := func(h uint16) uint16 {
|
||||
f := uint32(h)*0x9E3779B1 + keys[idx%4]
|
||||
return uint16((f ^ (f >> 16)) & 0xffff)
|
||||
}
|
||||
if encrypt {
|
||||
nl := r
|
||||
nr := l ^ round(r)
|
||||
l, r = nl, nr
|
||||
} else {
|
||||
nl := r ^ round(l)
|
||||
nr := l
|
||||
l, r = nl, nr
|
||||
}
|
||||
}
|
||||
return uint32(l)<<16 | uint32(r)
|
||||
}
|
||||
|
||||
func encodeNum(n uint32) string {
|
||||
x := feistel(n, true)
|
||||
var sb strings.Builder
|
||||
for x > 0 {
|
||||
sb.WriteByte(alphabet[x%base])
|
||||
x /= base
|
||||
}
|
||||
if sb.Len() == 0 {
|
||||
sb.WriteByte(alphabet[0])
|
||||
}
|
||||
// base62 低位在前,反转成高位在前
|
||||
runes := []rune(sb.String())
|
||||
for i, j := 0, len(runes)-1; i < j; i, j = i+1, j-1 {
|
||||
runes[i], runes[j] = runes[j], runes[i]
|
||||
}
|
||||
out := string(runes)
|
||||
if len(out) < minLen {
|
||||
out = strings.Repeat("0", minLen-len(out)) + out
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func decodeStr(s string) (uint32, error) {
|
||||
var x uint32
|
||||
for _, c := range s {
|
||||
idx := strings.IndexRune(alphabet, c)
|
||||
if idx < 0 {
|
||||
return 0, errors.New("invalid hashid: 含非法字符")
|
||||
}
|
||||
x = x*base + uint32(idx)
|
||||
}
|
||||
return feistel(x, false), nil
|
||||
}
|
||||
|
||||
// Encode 把数字主键编码为对外暴露的无序串。
|
||||
func Encode(id uint32) string {
|
||||
ensure()
|
||||
return encodeNum(id)
|
||||
}
|
||||
|
||||
// Decode 把对外串还原为数字主键;非法串返回 error(调用方应视为 404/未找到)。
|
||||
func Decode(s string) (uint32, error) {
|
||||
ensure()
|
||||
if strings.TrimSpace(s) == "" {
|
||||
return 0, errors.New("empty hashid")
|
||||
}
|
||||
return decodeStr(strings.TrimSpace(s))
|
||||
}
|
||||
42
internal/pkg/hashid/hashid_test.go
Normal file
42
internal/pkg/hashid/hashid_test.go
Normal file
@ -0,0 +1,42 @@
|
||||
package hashid
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestRoundtrip(t *testing.T) {
|
||||
// 覆盖边界与大量随机值,验证 Decode(Encode(n)) == n 严格成立
|
||||
cases := []uint32{0, 1, 2, 3, 7, 15, 16, 255, 256, 4095, 4096, 65535, 65536, 1 << 31, (1 << 32) - 1}
|
||||
for _, n := range cases {
|
||||
got, err := Decode(Encode(n))
|
||||
if err != nil {
|
||||
t.Fatalf("Decode(Encode(%d)) 返回错误: %v", n, err)
|
||||
}
|
||||
if got != n {
|
||||
t.Fatalf("往返不一致: 输入 %d, 得到 %d", n, got)
|
||||
}
|
||||
}
|
||||
|
||||
// 随机大批量
|
||||
for n := uint32(1); n < 5000; n++ {
|
||||
got, err := Decode(Encode(n))
|
||||
if err != nil || got != n {
|
||||
t.Fatalf("往返不一致: 输入 %d, 得到 %d (err=%v)", n, got, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNonSequential(t *testing.T) {
|
||||
Init("test-salt")
|
||||
a, b, c := Encode(100), Encode(101), Encode(102)
|
||||
if a == b || b == c || a == c {
|
||||
t.Fatalf("相邻 id 编码结果出现了相等: %s %s %s", a, b, c)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvalid(t *testing.T) {
|
||||
if _, err := Decode(""); err == nil {
|
||||
t.Fatal("空串应返回错误")
|
||||
}
|
||||
if _, err := Decode("!!!"); err == nil {
|
||||
t.Fatal("含非法字符应返回错误")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user