redis_bus.go 3.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149
  1. package websocket
  2. import (
  3. "context"
  4. "encoding/json"
  5. "fmt"
  6. "os"
  7. "strconv"
  8. "time"
  9. "github.com/redis/go-redis/v9"
  10. )
  11. type DistributedBus interface {
  12. Publish(msg *Message) error
  13. Subscribe(handler func(msg *Message))
  14. Close() error
  15. }
  16. type redisWireMessage struct {
  17. ConversationID uint `json:"conversation_id"`
  18. Type string `json:"type"`
  19. Scope string `json:"scope,omitempty"`
  20. Data json.RawMessage `json:"data"`
  21. Source string `json:"source"`
  22. }
  23. type RedisBus struct {
  24. ctx context.Context
  25. client *redis.Client
  26. channel string
  27. nodeID string
  28. pubsub *redis.PubSub
  29. }
  30. func NewRedisBusFromEnv() (DistributedBus, error) {
  31. redisURL := os.Getenv("REDIS_URL")
  32. redisAddr := os.Getenv("REDIS_ADDR")
  33. if redisURL == "" && redisAddr == "" {
  34. return nil, nil
  35. }
  36. var opts *redis.Options
  37. var err error
  38. if redisURL != "" {
  39. opts, err = redis.ParseURL(redisURL)
  40. if err != nil {
  41. return nil, fmt.Errorf("parse REDIS_URL failed: %w", err)
  42. }
  43. } else {
  44. opts = &redis.Options{
  45. Addr: redisAddr,
  46. Password: os.Getenv("REDIS_PASSWORD"),
  47. DB: 0,
  48. }
  49. if dbRaw := os.Getenv("REDIS_DB"); dbRaw != "" {
  50. if db, parseErr := strconv.Atoi(dbRaw); parseErr == nil {
  51. opts.DB = db
  52. }
  53. }
  54. }
  55. client := redis.NewClient(opts)
  56. ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
  57. defer cancel()
  58. if pingErr := client.Ping(ctx).Err(); pingErr != nil {
  59. _ = client.Close()
  60. return nil, fmt.Errorf("redis ping failed: %w", pingErr)
  61. }
  62. channel := os.Getenv("REDIS_WS_CHANNEL")
  63. if channel == "" {
  64. channel = "ai_cs:ws_events"
  65. }
  66. nodeID := fmt.Sprintf("%s-%d", hostnameOrDefault(), time.Now().UnixNano())
  67. return &RedisBus{
  68. ctx: context.Background(),
  69. client: client,
  70. channel: channel,
  71. nodeID: nodeID,
  72. }, nil
  73. }
  74. func (r *RedisBus) Publish(msg *Message) error {
  75. if msg == nil {
  76. return nil
  77. }
  78. dataBytes, err := json.Marshal(msg.Data)
  79. if err != nil {
  80. return err
  81. }
  82. wire := redisWireMessage{
  83. ConversationID: msg.ConversationID,
  84. Type: msg.Type,
  85. Scope: msg.Scope,
  86. Data: dataBytes,
  87. Source: r.nodeID,
  88. }
  89. payload, err := json.Marshal(wire)
  90. if err != nil {
  91. return err
  92. }
  93. return r.client.Publish(r.ctx, r.channel, payload).Err()
  94. }
  95. func (r *RedisBus) Subscribe(handler func(msg *Message)) {
  96. if handler == nil {
  97. return
  98. }
  99. r.pubsub = r.client.Subscribe(r.ctx, r.channel)
  100. ch := r.pubsub.Channel()
  101. go func() {
  102. for item := range ch {
  103. var wire redisWireMessage
  104. if err := json.Unmarshal([]byte(item.Payload), &wire); err != nil {
  105. continue
  106. }
  107. if wire.Source == r.nodeID {
  108. continue
  109. }
  110. var data interface{}
  111. if err := json.Unmarshal(wire.Data, &data); err != nil {
  112. continue
  113. }
  114. handler(&Message{
  115. ConversationID: wire.ConversationID,
  116. Type: wire.Type,
  117. Scope: wire.Scope,
  118. Data: data,
  119. FromRemote: true,
  120. })
  121. }
  122. }()
  123. }
  124. func (r *RedisBus) Close() error {
  125. if r.pubsub != nil {
  126. _ = r.pubsub.Close()
  127. }
  128. return r.client.Close()
  129. }
  130. func hostnameOrDefault() string {
  131. name, err := os.Hostname()
  132. if err != nil || name == "" {
  133. return "node"
  134. }
  135. return name
  136. }