ws_token.go 2.0 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677
  1. package utils
  2. import (
  3. "crypto/hmac"
  4. "crypto/sha256"
  5. "encoding/base64"
  6. "fmt"
  7. "os"
  8. "strconv"
  9. "strings"
  10. "time"
  11. )
  12. func wsTokenSecret() []byte {
  13. // 与现有系统保持一致:优先使用 ENCRYPTION_KEY;未设置时回退固定开发值。
  14. secret := os.Getenv("ENCRYPTION_KEY")
  15. if secret == "" {
  16. secret = "abcdefghijklmnopqrstuvwxyz123456"
  17. }
  18. return []byte(secret)
  19. }
  20. // GenerateWSToken 生成客服 WebSocket 短期令牌。
  21. func GenerateWSToken(userID uint, ttl time.Duration) (token string, expireAt int64, err error) {
  22. if userID == 0 {
  23. return "", 0, fmt.Errorf("invalid user id")
  24. }
  25. if ttl <= 0 {
  26. ttl = 24 * time.Hour
  27. }
  28. expireAt = time.Now().Add(ttl).Unix()
  29. payload := fmt.Sprintf("%d:%d", userID, expireAt)
  30. payloadEnc := base64.RawURLEncoding.EncodeToString([]byte(payload))
  31. mac := hmac.New(sha256.New, wsTokenSecret())
  32. _, _ = mac.Write([]byte(payloadEnc))
  33. signature := base64.RawURLEncoding.EncodeToString(mac.Sum(nil))
  34. return payloadEnc + "." + signature, expireAt, nil
  35. }
  36. // ValidateWSToken 校验客服 WebSocket 令牌是否与用户匹配且未过期。
  37. func ValidateWSToken(token string, expectedUserID uint) bool {
  38. if expectedUserID == 0 || token == "" {
  39. return false
  40. }
  41. parts := strings.Split(token, ".")
  42. if len(parts) != 2 {
  43. return false
  44. }
  45. payloadEnc, signature := parts[0], parts[1]
  46. mac := hmac.New(sha256.New, wsTokenSecret())
  47. _, _ = mac.Write([]byte(payloadEnc))
  48. expectedSig := base64.RawURLEncoding.EncodeToString(mac.Sum(nil))
  49. if !hmac.Equal([]byte(signature), []byte(expectedSig)) {
  50. return false
  51. }
  52. payloadRaw, err := base64.RawURLEncoding.DecodeString(payloadEnc)
  53. if err != nil {
  54. return false
  55. }
  56. payloadParts := strings.Split(string(payloadRaw), ":")
  57. if len(payloadParts) != 2 {
  58. return false
  59. }
  60. uid64, err := strconv.ParseUint(payloadParts[0], 10, 64)
  61. if err != nil || uint(uid64) != expectedUserID {
  62. return false
  63. }
  64. expireAt, err := strconv.ParseInt(payloadParts[1], 10, 64)
  65. if err != nil {
  66. return false
  67. }
  68. return time.Now().Unix() <= expireAt
  69. }