fix: restore web widget realtime replies (HH-556) (#136)

* fix: restore web widget realtime replies (HH-556)

* fix: make realtime delivery durable (HH-556)

* fix: make realtime message outbox atomic (HH-556)

---------

Co-authored-by: Rogee <rogee@ipao.vip>
This commit is contained in:
Rogee
2026-08-23 22:35:10 +08:00
committed by GitHub
co-authored by rogee
parent 17244bcc9d
commit 60a6ac4785
13 changed files with 987 additions and 136 deletions
+5 -1
View File
@@ -786,11 +786,14 @@ func Bootstrap(env string) (*App, error) {
}
})
eventPublisher := wspkg.NewEventPublisher(wsHub, nil, wsRelay)
eventPublisher.SetWorkerPool(workerPool)
presenceTracker := wspkg.NewPresenceTracker(rdb, wsRelay)
// Bridge channel dispatcher events to the WebSocket EventPublisher so that
// message/conversation/contact events reach connected WS/SSE clients.
channelDispatcher.Register(wsevent.New(eventPublisher))
realtimeBridge := wsevent.New(eventPublisher)
channelDispatcher.Register(realtimeBridge)
messageService.SetRealtimeEventBridge(realtimeBridge)
// Widget service + handler (M11 — WebWidget channel completion)
// hubTypingAdapter delegates typing events to the WS hub's direct broadcast,
@@ -802,6 +805,7 @@ func Bootstrap(env string) (*App, error) {
widgetService := service.NewWidgetService(inboxRepo, contactRepo, contactInboxRepo, conversationRepo, messageRepo, widgetTypingAdapter, widgetThemeConfigRepo, preChatFormRepo, widgetFileUploadRepo, widgetOfflineMessageRepo, inboxMemberRepo, tagRepo, campaignRepo)
widgetService.SetWorkerPool(workerPool)
widgetService.SetDispatcher(channelDispatcher)
widgetService.SetRealtimeEventBridge(realtimeBridge)
widgetHandler := widget.NewHandler(widgetService)
// Upload: DirectUpload repo + service + handler (account-level + widget direct uploads)
@@ -168,17 +168,20 @@ func (s *MessageHandlerTestSuite) SetupTest() {
contact := &model.Contact{AccountID: account.ID, Name: "MsgHandlerTestContact"}
s.Require().NoError(s.db.Create(contact).Error)
s.testContact = contact
contactInbox := &model.ContactInbox{ContactID: contact.ID, InboxID: inbox.ID, SourceID: "message-handler-contact", PubsubToken: "message-handler-token"}
s.Require().NoError(s.db.Create(contactInbox).Error)
displayID := uint(4242)
conv := &model.Conversation{
AccountID: account.ID,
DisplayID: &displayID,
InboxID: inbox.ID,
ContactID: contact.ID,
Status: string(model.ConversationStatusOpen),
Priority: string(model.ConversationPriorityMedium),
ChannelType: "web_widget",
Channel: "web_widget",
AccountID: account.ID,
DisplayID: &displayID,
InboxID: inbox.ID,
ContactID: contact.ID,
ContactInboxID: &contactInbox.ID,
Status: string(model.ConversationStatusOpen),
Priority: string(model.ConversationPriorityMedium),
ChannelType: "web_widget",
Channel: "web_widget",
}
s.Require().NoError(s.db.Create(conv).Error)
s.testConv = conv
@@ -19,6 +19,7 @@ import (
"time"
"github.com/gin-gonic/gin"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/datatypes"
@@ -28,6 +29,7 @@ import (
"github.com/gochat/gochat/internal/automation"
"github.com/gochat/gochat/internal/campaign"
"github.com/gochat/gochat/internal/channel"
"github.com/gochat/gochat/internal/config"
apiv1 "github.com/gochat/gochat/internal/handler/api/v1"
"github.com/gochat/gochat/internal/middleware"
@@ -35,7 +37,9 @@ import (
channelmodel "github.com/gochat/gochat/internal/model/channel"
"github.com/gochat/gochat/internal/repository"
"github.com/gochat/gochat/internal/service"
"github.com/gochat/gochat/internal/worker"
ws "github.com/gochat/gochat/internal/ws"
"github.com/gochat/gochat/internal/wsevent"
)
// noopTypingIndicatorWidget is a stub for handler tests.
@@ -2108,6 +2112,51 @@ func TestWidgetHandler_LegacySendMessageUsesWidgetPayload(t *testing.T) {
assert.Equal(t, float64(account.ID), attachment["account_id"])
}
func TestWidgetHandler_RealtimeFailureReturnsSuccessWithoutDuplicateMessage(t *testing.T) {
db, router, handler := setupWidgetHandlerTest(t)
require.NoError(t, db.AutoMigrate(&model.BackgroundJob{}))
seedWidgetHandlerData(t, db)
pool := worker.NewWorkerPoolWithOptions(db, worker.WithBackoff(func(int) time.Duration { return 0 }))
handler.widgetService.SetWorkerPool(pool)
rdb := redis.NewClient(&redis.Options{Addr: "127.0.0.1:1", MaxRetries: -1, DialTimeout: 10 * time.Millisecond})
t.Cleanup(func() { require.NoError(t, rdb.Close()) })
publisher := ws.NewEventPublisher(nil, nil, ws.NewBroadcastRelay(rdb, nil))
publisher.SetWorkerPool(pool)
dispatcher := channel.NewDispatcher()
dispatcher.Register(wsevent.New(publisher))
handler.widgetService.SetDispatcher(dispatcher)
initBody, err := json.Marshal(map[string]interface{}{"website_token": "handler_ws_token_123"})
require.NoError(t, err)
initResponse := httptest.NewRecorder()
initRequest := httptest.NewRequest(http.MethodPost, "/widget/init", bytes.NewReader(initBody))
initRequest.Header.Set("Content-Type", "application/json")
router.ServeHTTP(initResponse, initRequest)
require.Equal(t, http.StatusOK, initResponse.Code, initResponse.Body.String())
var initPayload map[string]interface{}
require.NoError(t, json.Unmarshal(initResponse.Body.Bytes(), &initPayload))
messageBody, err := json.Marshal(map[string]interface{}{"content": "persist once"})
require.NoError(t, err)
messageResponse := httptest.NewRecorder()
messageRequest := httptest.NewRequest(http.MethodPost, "/widget/messages", bytes.NewReader(messageBody))
messageRequest.Header.Set("Content-Type", "application/json")
messageRequest.Header.Set("X-Widget-Token", initPayload["widget_token"].(string))
router.ServeHTTP(messageResponse, messageRequest)
require.Equal(t, http.StatusOK, messageResponse.Code, messageResponse.Body.String())
processed, err := pool.ProcessOne(context.Background())
require.True(t, processed)
require.Error(t, err)
var messageCount int64
require.NoError(t, db.Model(&model.Message{}).Count(&messageCount).Error)
require.Equal(t, int64(1), messageCount)
var realtimeJob model.BackgroundJob
require.NoError(t, db.Where("queue = ? AND last_error != ''", "events").First(&realtimeJob).Error)
require.Equal(t, model.BackgroundJobStatusRetrying, realtimeJob.Status)
}
func TestWidgetHandler_SendMessage_NoWidgetToken(t *testing.T) {
_, router, _ := setupWidgetHandlerTest(t)
+105 -1
View File
@@ -4,10 +4,12 @@ import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
"time"
@@ -28,10 +30,38 @@ import (
"github.com/gochat/gochat/internal/model"
"github.com/gochat/gochat/internal/repository"
"github.com/gochat/gochat/internal/service"
"github.com/gochat/gochat/internal/worker"
wspkg "github.com/gochat/gochat/internal/ws"
"github.com/gochat/gochat/internal/wsevent"
)
type failTokenPublishOnceHook struct {
mu sync.Mutex
channel string
failed bool
}
func (h *failTokenPublishOnceHook) DialHook(next redis.DialHook) redis.DialHook { return next }
func (h *failTokenPublishOnceHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook {
return func(ctx context.Context, cmd redis.Cmder) error {
if cmd.Name() == "publish" && len(cmd.Args()) >= 2 && fmt.Sprint(cmd.Args()[1]) == h.channel {
h.mu.Lock()
if !h.failed {
h.failed = true
h.mu.Unlock()
return errors.New("token room unavailable")
}
h.mu.Unlock()
}
return next(ctx, cmd)
}
}
func (h *failTokenPublishOnceHook) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook {
return next
}
// --- Protocol Tests ---
func TestCommandType_Constants(t *testing.T) {
@@ -288,6 +318,79 @@ func TestHub_SendToAccount(t *testing.T) {
}
}
func TestDurableWidgetRetryDoesNotDuplicateRealHubSubscribers(t *testing.T) {
db, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s-%d?mode=memory&cache=shared", t.Name(), time.Now().UnixNano())), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.BackgroundJob{}))
sqlDB, err := db.DB()
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, sqlDB.Close()) })
mini := miniredis.RunT(t)
rdb := redis.NewClient(&redis.Options{Addr: mini.Addr()})
rdb.AddHook(&failTokenPublishOnceHook{channel: wspkg.RedisPrefixRoom + "pubsub_token_visitor"})
t.Cleanup(func() { require.NoError(t, rdb.Close()) })
hub := NewHubSimple()
dashboard := NewClient(1, 1, nil, hub)
dashboard.Identifier = `{"channel":"AccountChannel","account_id":1}`
hub.Register(dashboard)
visitor := NewClient(0, 1, nil, hub)
visitor.IsContact = true
visitor.PubsubToken = "visitor"
visitor.Identifier = `{"channel":"RoomChannel","pubsub_token":"visitor"}`
hub.Register(visitor)
relay := wspkg.NewBroadcastRelay(rdb, hub)
ctx, cancel := context.WithCancel(context.Background())
require.NoError(t, relay.Start(ctx))
t.Cleanup(func() {
cancel()
require.NoError(t, relay.Stop())
})
pool := worker.NewWorkerPoolWithOptions(db, worker.WithBackoff(func(int) time.Duration { return 0 }))
publisher := wspkg.NewEventPublisher(hub, nil, relay)
publisher.SetWorkerPool(pool)
require.NoError(t, publisher.PublishWidgetEvent(1, "visitor", wspkg.EventMessageCreated, map[string]any{"id": 7}))
processed, err := pool.ProcessOne(context.Background())
require.True(t, processed)
require.NoError(t, err)
select {
case <-dashboard.Send:
case <-time.After(time.Second):
t.Fatal("account subscriber did not receive message.created")
}
processed, err = pool.ProcessOne(context.Background())
require.True(t, processed)
require.ErrorContains(t, err, "token room unavailable")
select {
case <-visitor.Send:
t.Fatal("token subscriber received failed delivery")
default:
}
processed, err = pool.ProcessOne(context.Background())
require.True(t, processed)
require.NoError(t, err)
select {
case <-visitor.Send:
case <-time.After(time.Second):
t.Fatal("token subscriber did not receive recovered message.created")
}
time.Sleep(20 * time.Millisecond)
select {
case <-dashboard.Send:
t.Fatal("account subscriber received duplicate message.created")
default:
}
select {
case <-visitor.Send:
t.Fatal("token subscriber received duplicate message.created")
default:
}
}
func TestHub_SendToAccountSanitizesVisitorIdentity(t *testing.T) {
hub := NewHubSimple()
agent := NewClient(1, 10, nil, hub)
@@ -721,6 +824,8 @@ func TestDashboardOutgoingReachesDashboardAndReconnectedWidget(t *testing.T) {
))
account := &model.Account{Name: "Realtime account", Active: true}
require.NoError(t, db.Create(account).Error)
user := &model.User{Base: model.Base{ID: 7}, Name: "Agent", Email: "agent@example.com", Provider: "local", Active: true}
require.NoError(t, db.Create(user).Error)
inbox := &model.Inbox{AccountID: account.ID, Name: "Website", ChannelType: "web_widget", Enabled: true}
require.NoError(t, db.Create(inbox).Error)
contact := &model.Contact{AccountID: account.ID, Name: "Visitor"}
@@ -758,7 +863,6 @@ func TestDashboardOutgoingReachesDashboardAndReconnectedWidget(t *testing.T) {
server := httptest.NewServer(router)
t.Cleanup(server.Close)
user := &model.User{Base: model.Base{ID: 7}, Provider: "local"}
tokenPair, err := jwtService.GenerateTokenPair(user, account.ID, "agent")
require.NoError(t, err)
dashboard := dialCable(t, server.URL, "?token="+tokenPair.AccessToken)
+98 -28
View File
@@ -16,6 +16,7 @@ import (
"github.com/gochat/gochat/internal/repository"
"github.com/gochat/gochat/internal/search"
"github.com/gochat/gochat/internal/worker"
"github.com/gochat/gochat/internal/wsevent"
applogger "github.com/gochat/gochat/pkg/logger"
pkgvalidator "github.com/gochat/gochat/pkg/validator"
@@ -32,6 +33,7 @@ type MessageService struct {
searchIndexer SearchIndexer
llmProvider llm.Provider
worker *worker.WorkerPool
realtime *wsevent.BridgeListener
}
// NewMessageService creates a new Message service.
@@ -49,6 +51,10 @@ func (s *MessageService) SetWorkerPool(wp *worker.WorkerPool) {
RegisterMessageDeliverySearchIndexer(wp, s.repo.DB(), s.dispatcher, s.searchIndexer)
}
func (s *MessageService) SetRealtimeEventBridge(bridge *wsevent.BridgeListener) {
s.realtime = bridge
}
func (s *MessageService) DB() *gorm.DB {
if s == nil || s.repo == nil {
return nil
@@ -69,7 +75,15 @@ func (s *MessageService) deleteMessageIndex(ctx context.Context, accountID uint,
}
// dispatchMessageEvent is a helper to build and dispatch a message event.
func (s *MessageService) dispatchMessageEvent(ctx context.Context, eventType channel.EventType, message *model.Message) {
func (s *MessageService) dispatchMessageEvent(ctx context.Context, eventType channel.EventType, message *model.Message) error {
event, err := s.buildMessageEvent(ctx, s.repo.DB(), eventType, message)
if err != nil {
return err
}
return s.dispatchPreparedMessageEvent(ctx, event, message.ID)
}
func (s *MessageService) buildMessageEvent(ctx context.Context, db *gorm.DB, eventType channel.EventType, message *model.Message) (*channel.ChannelEvent, error) {
event := channel.NewChannelEvent(eventType, channel.ChannelAPI, message.AccountID, message.InboxID)
event.ConversationID = message.ConversationID
if message.SenderID != nil {
@@ -77,18 +91,25 @@ func (s *MessageService) dispatchMessageEvent(ctx context.Context, eventType cha
}
event.Data["message"] = message
if eventType == channel.EventMessageCreated || eventType == channel.EventMessageUpdated || eventType == channel.EventMessageDeleted {
s.addMessageEventContext(ctx, event, message)
}
applogger.L().Infof("dispatching event %s for message %d", eventType, message.ID)
if err := s.dispatcher.Dispatch(ctx, event); err != nil {
applogger.L().Errorf("failed to dispatch event %s for message %d: %v", eventType, message.ID, err)
if err := s.addMessageEventContext(ctx, db, event, message); err != nil {
return nil, fmt.Errorf("build %s event for message %d: %w", eventType, message.ID, err)
}
}
return event, nil
}
func (s *MessageService) addMessageEventContext(ctx context.Context, event *channel.ChannelEvent, message *model.Message) {
func (s *MessageService) dispatchPreparedMessageEvent(ctx context.Context, event *channel.ChannelEvent, messageID uint) error {
applogger.L().Infof("dispatching event %s for message %d", event.Type, messageID)
if err := s.dispatcher.Dispatch(ctx, event); err != nil {
return fmt.Errorf("dispatch event %s for message %d: %w", event.Type, messageID, err)
}
return nil
}
func (s *MessageService) addMessageEventContext(ctx context.Context, db *gorm.DB, event *channel.ChannelEvent, message *model.Message) error {
var inbox model.Inbox
if err := s.repo.DB().WithContext(ctx).First(&inbox, message.InboxID).Error; err != nil {
return
if err := db.WithContext(ctx).Where("id = ? AND account_id = ?", message.InboxID, message.AccountID).First(&inbox).Error; err != nil {
return fmt.Errorf("load inbox %d: %w", message.InboxID, err)
}
event.Data["inbox"] = &inbox
isWebWidget := strings.EqualFold(strings.TrimSpace(inbox.ChannelType), string(channel.ChannelWebWidget)) ||
@@ -98,42 +119,60 @@ func (s *MessageService) addMessageEventContext(ctx context.Context, event *chan
}
var conversation model.Conversation
if err := s.repo.DB().WithContext(ctx).First(&conversation, message.ConversationID).Error; err != nil {
return
if err := db.WithContext(ctx).Where("id = ? AND account_id = ?", message.ConversationID, message.AccountID).First(&conversation).Error; err != nil {
return fmt.Errorf("load conversation %d: %w", message.ConversationID, err)
}
event.ContactID = conversation.ContactID
event.Data["conversation"] = &conversation
widgetVisible := isWebWidget && !message.Private && message.MessageType != string(model.MessageTypeActivity)
var contact model.Contact
if err := s.repo.DB().WithContext(ctx).First(&contact, conversation.ContactID).Error; err == nil {
event.Data["contact"] = &contact
if err := db.WithContext(ctx).Where("id = ? AND account_id = ?", conversation.ContactID, message.AccountID).First(&contact).Error; err != nil {
return fmt.Errorf("load contact %d: %w", conversation.ContactID, err)
}
event.Data["contact"] = &contact
if message.SenderID != nil {
switch strings.ToLower(strings.TrimSpace(message.SenderType)) {
case "user":
var sender model.User
if err := s.repo.DB().WithContext(ctx).First(&sender, *message.SenderID).Error; err == nil {
event.Data["sender"] = &sender
if err := db.WithContext(ctx).First(&sender, *message.SenderID).Error; err != nil {
if widgetVisible {
return fmt.Errorf("load user sender %d: %w", *message.SenderID, err)
}
break
}
event.Data["sender"] = &sender
case "agentbot", "agent_bot":
var sender model.AgentBot
if err := s.repo.DB().WithContext(ctx).First(&sender, *message.SenderID).Error; err == nil {
event.Data["sender"] = &sender
if err := db.WithContext(ctx).Where("id = ? AND (account_id IS NULL OR account_id = ?)", *message.SenderID, message.AccountID).First(&sender).Error; err != nil {
return fmt.Errorf("load agent bot sender %d: %w", *message.SenderID, err)
}
event.Data["sender"] = &sender
case "captain::assistant", "captainassistant", "captain_assistant":
var sender model.CaptainAssistant
if err := s.repo.DB().WithContext(ctx).First(&sender, *message.SenderID).Error; err == nil {
event.Data["sender"] = &sender
if err := db.WithContext(ctx).Where("id = ? AND account_id = ?", *message.SenderID, message.AccountID).First(&sender).Error; err != nil {
return fmt.Errorf("load captain sender %d: %w", *message.SenderID, err)
}
event.Data["sender"] = &sender
}
}
if !isWebWidget || message.Private || message.MessageType == string(model.MessageTypeActivity) || conversation.ContactInboxID == nil {
return
if !widgetVisible {
return nil
}
if conversation.ContactInboxID == nil {
return fmt.Errorf("web widget conversation %d has no contact inbox", conversation.ID)
}
var contactInbox model.ContactInbox
if err := s.repo.DB().WithContext(ctx).Select("pubsub_token").First(&contactInbox, *conversation.ContactInboxID).Error; err == nil {
event.Data["widget_token"] = contactInbox.PubsubToken
if err := db.WithContext(ctx).Select("pubsub_token").Where(
"id = ? AND contact_id = ? AND inbox_id = ?", *conversation.ContactInboxID, conversation.ContactID, conversation.InboxID,
).First(&contactInbox).Error; err != nil {
return fmt.Errorf("load contact inbox %d: %w", *conversation.ContactInboxID, err)
}
if strings.TrimSpace(contactInbox.PubsubToken) == "" {
return fmt.Errorf("contact inbox %d has no pubsub token", *conversation.ContactInboxID)
}
event.Data["widget_token"] = contactInbox.PubsubToken
return nil
}
// ListByConversation retrieves all messages for a conversation.
@@ -360,6 +399,8 @@ func (s *MessageService) Create(ctx context.Context, accountID uint, userID uint
var deliveryJob *model.BackgroundJob
var deliveryCreated bool
var realtimeJobs []*model.BackgroundJob
var messageCreatedEvent *channel.ChannelEvent
if err := s.repo.DB().WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if req.ExpectedAITakeoverVersion != 0 {
var active model.Conversation
@@ -445,6 +486,16 @@ func (s *MessageService) Create(ctx context.Context, accountID uint, userID uint
return err
}
}
if s.realtime != nil {
messageCreatedEvent, err = s.buildMessageEvent(ctx, tx, channel.EventMessageCreated, message)
if err != nil {
return err
}
realtimeJobs, err = s.realtime.EnqueueInTransaction(ctx, tx, messageCreatedEvent)
if err != nil {
return err
}
}
return nil
}); err != nil {
if errors.Is(err, errAITakeoverEnded) {
@@ -461,6 +512,9 @@ func (s *MessageService) Create(ctx context.Context, accountID uint, userID uint
if deliveryCreated {
s.worker.Publish(ctx, deliveryJob)
}
if s.realtime != nil {
s.realtime.PublishEnqueued(ctx, realtimeJobs)
}
if message.AITakeoverExited {
event := channel.NewChannelEvent(channel.EventConversationUpdated, channel.ChannelType(conversation.ChannelType), conversation.AccountID, conversation.InboxID)
@@ -473,12 +527,22 @@ func (s *MessageService) Create(ctx context.Context, accountID uint, userID uint
}
// Dispatch EventMessageCreated
s.dispatchMessageEvent(ctx, channel.EventMessageCreated, message)
var dispatchErr error
if messageCreatedEvent != nil {
dispatchErr = s.dispatchPreparedMessageEvent(ctx, messageCreatedEvent, message.ID)
} else {
dispatchErr = s.dispatchMessageEvent(ctx, channel.EventMessageCreated, message)
}
if dispatchErr != nil {
return message, dispatchErr
}
s.indexMessage(ctx, message)
// Dispatch additional event based on message type
if req.MessageType == "incoming" {
s.dispatchMessageEvent(ctx, channel.EventMessageIncoming, message)
if err := s.dispatchMessageEvent(ctx, channel.EventMessageIncoming, message); err != nil {
return message, err
}
if s.worker != nil {
if _, err := EnqueueCaptainConversationResponseForMessage(ctx, s.worker, s.repo.DB(), message.ID); err != nil {
return message, err
@@ -492,7 +556,9 @@ func (s *MessageService) Create(ctx context.Context, accountID uint, userID uint
return message, err
}
} else {
s.dispatchMessageEvent(ctx, channel.EventMessageOutgoing, message)
if err := s.dispatchMessageEvent(ctx, channel.EventMessageOutgoing, message); err != nil {
return message, err
}
}
}
@@ -856,7 +922,9 @@ func (s *MessageService) UpdateInConversation(ctx context.Context, accountID, co
}
// Dispatch EventMessageUpdated
s.dispatchMessageEvent(ctx, channel.EventMessageUpdated, message)
if err := s.dispatchMessageEvent(ctx, channel.EventMessageUpdated, message); err != nil {
return message, err
}
s.indexMessage(ctx, message)
return message, nil
@@ -899,7 +967,9 @@ func (s *MessageService) DeleteInConversation(ctx context.Context, accountID, co
}
// Dispatch EventMessageDeleted
s.dispatchMessageEvent(ctx, channel.EventMessageDeleted, message)
if err := s.dispatchMessageEvent(ctx, channel.EventMessageDeleted, message); err != nil {
return message, err
}
s.deleteMessageIndex(ctx, accountID, message.ID)
return message, nil
@@ -3,10 +3,12 @@ package service
import (
"context"
"encoding/json"
"errors"
"fmt"
"testing"
"time"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/datatypes"
@@ -18,6 +20,8 @@ import (
"github.com/gochat/gochat/internal/repository"
"github.com/gochat/gochat/internal/search"
"github.com/gochat/gochat/internal/worker"
wspkg "github.com/gochat/gochat/internal/ws"
"github.com/gochat/gochat/internal/wsevent"
)
// mockMessageLLMProvider implements llm.Provider for message service testing.
@@ -57,13 +61,14 @@ type mockRetryListener struct {
type messageEventListener struct {
events []*channel.ChannelEvent
err error
}
func (l *messageEventListener) Name() string { return "message-event-listener" }
func (l *messageEventListener) OnEvent(_ context.Context, event *channel.ChannelEvent) error {
l.events = append(l.events, event)
return nil
return l.err
}
func (l *mockRetryListener) Name() string {
@@ -106,6 +111,13 @@ func setupMessageServiceWithDefaultLLM(t *testing.T) (*gorm.DB, *repository.Mess
func TestMessageService_WebWidgetReplyCarriesRealtimeContext(t *testing.T) {
db, _, dispatcher, svc := setupMessageServiceWithDefaultLLM(t)
pool := worker.NewWorkerPool(db, nil)
svc.SetWorkerPool(pool)
publisher := wspkg.NewEventPublisherLocal(nil, nil)
publisher.SetWorkerPool(pool)
bridge := wsevent.New(publisher)
dispatcher.Register(bridge)
svc.SetRealtimeEventBridge(bridge)
account := createTestAccount(t, db)
agent := createTestUser(t, db, account.ID)
inbox := createTestInbox(t, db, account.ID, string(channel.ChannelWebWidget))
@@ -138,6 +150,123 @@ func TestMessageService_WebWidgetReplyCarriesRealtimeContext(t *testing.T) {
assert.IsType(t, &model.Conversation{}, created.Data["conversation"])
assert.IsType(t, &model.Contact{}, created.Data["contact"])
assert.IsType(t, &model.User{}, created.Data["sender"])
var realtimeJobs int64
require.NoError(t, db.Model(&model.BackgroundJob{}).Where("idempotency_key LIKE ?", "realtime:message.created:%").Count(&realtimeJobs).Error)
require.Equal(t, int64(2), realtimeJobs)
}
func TestMessageService_WebWidgetReplyRequiresRealtimeContext(t *testing.T) {
db, _, dispatcher, svc := setupMessageServiceWithDefaultLLM(t)
account := createTestAccount(t, db)
agent := createTestUser(t, db, account.ID)
inbox := createTestInbox(t, db, account.ID, string(channel.ChannelWebWidget))
contact := createTestContact(t, db, account.ID)
conversation := &model.Conversation{
AccountID: account.ID, InboxID: inbox.ID, ContactID: contact.ID,
Status: string(model.ConversationStatusOpen), ChannelType: "web_widget", Channel: "web_widget",
}
require.NoError(t, db.Create(conversation).Error)
listener := &messageEventListener{}
dispatcher.Register(listener)
_, err := svc.Create(context.Background(), account.ID, agent.ID, CreateMessageRequest{
ConversationID: conversation.ID,
Content: "agent reply",
ContentType: "text",
MessageType: "outgoing",
})
require.ErrorContains(t, err, "contact inbox")
assert.Empty(t, listener.events)
}
func TestMessageService_ReturnsMessageDispatchError(t *testing.T) {
db, _, dispatcher, svc := setupMessageServiceWithDefaultLLM(t)
account := createTestAccount(t, db)
agent := createTestUser(t, db, account.ID)
inbox := createTestInbox(t, db, account.ID, string(channel.ChannelAPI))
contact := createTestContact(t, db, account.ID)
conversation := createTestConversation(t, db, account.ID, inbox.ID, contact.ID)
dispatchErr := errors.New("realtime unavailable")
dispatcher.Register(&messageEventListener{err: dispatchErr})
_, err := svc.Create(context.Background(), account.ID, agent.ID, CreateMessageRequest{
ConversationID: conversation.ID,
Content: "agent reply",
ContentType: "text",
MessageType: "outgoing",
})
require.ErrorIs(t, err, dispatchErr)
}
func TestMessageService_RealtimeFailureDoesNotBlockSendReply(t *testing.T) {
db, _, dispatcher, svc := setupMessageServiceWithDefaultLLM(t)
pool := worker.NewWorkerPoolWithOptions(db, worker.WithBackoff(func(int) time.Duration { return 0 }))
svc.SetWorkerPool(pool)
rdb := redis.NewClient(&redis.Options{Addr: "127.0.0.1:1", MaxRetries: -1, DialTimeout: 10 * time.Millisecond})
t.Cleanup(func() { require.NoError(t, rdb.Close()) })
publisher := wspkg.NewEventPublisher(nil, nil, wspkg.NewBroadcastRelay(rdb, nil))
publisher.SetWorkerPool(pool)
bridge := wsevent.New(publisher)
dispatcher.Register(bridge)
svc.SetRealtimeEventBridge(bridge)
account := createTestAccount(t, db)
agent := createTestUser(t, db, account.ID)
inbox := createTestInbox(t, db, account.ID, string(channel.ChannelAPI))
contact := createTestContact(t, db, account.ID)
conversation := createTestConversation(t, db, account.ID, inbox.ID, contact.ID)
message, err := svc.Create(context.Background(), account.ID, agent.ID, CreateMessageRequest{
ConversationID: conversation.ID,
Content: "agent reply",
ContentType: "text",
MessageType: "outgoing",
})
require.NoError(t, err)
require.NotNil(t, message)
processed, err := pool.ProcessOne(context.Background())
require.True(t, processed)
require.ErrorContains(t, err, "connect: connection refused")
var persistedMessages int64
require.NoError(t, db.Model(&model.Message{}).Where("id = ?", message.ID).Count(&persistedMessages).Error)
require.Equal(t, int64(1), persistedMessages)
var realtimeJob, sendReplyJob model.BackgroundJob
require.NoError(t, db.Where("queue = ? AND idempotency_key LIKE ?", "events", "realtime:%").First(&realtimeJob).Error)
require.Equal(t, model.BackgroundJobStatusRetrying, realtimeJob.Status)
require.NotEmpty(t, realtimeJob.LastError)
require.NoError(t, db.Where("job_type = ?", TaskTypeMessageSendReply).First(&sendReplyJob).Error)
require.Equal(t, model.BackgroundJobStatusQueued, sendReplyJob.Status)
}
func TestMessageService_RealtimeJobFailureRollsBackMessage(t *testing.T) {
db, _, dispatcher, svc := setupMessageServiceWithDefaultLLM(t)
pool := worker.NewWorkerPool(db, nil)
svc.SetWorkerPool(pool)
publisher := wspkg.NewEventPublisherLocal(nil, nil)
publisher.SetWorkerPool(pool)
bridge := wsevent.New(publisher)
dispatcher.Register(bridge)
svc.SetRealtimeEventBridge(bridge)
account := createTestAccount(t, db)
agent := createTestUser(t, db, account.ID)
inbox := createTestInbox(t, db, account.ID, string(channel.ChannelAPI))
contact := createTestContact(t, db, account.ID)
conversation := createTestConversation(t, db, account.ID, inbox.ID, contact.ID)
failBackgroundJobCreate(t, db)
_, err := svc.Create(context.Background(), account.ID, agent.ID, CreateMessageRequest{
ConversationID: conversation.ID, Content: "must roll back", ContentType: "text", MessageType: "outgoing",
})
require.ErrorContains(t, err, "forced background job insert failure")
var messages, jobs int64
require.NoError(t, db.Model(&model.Message{}).Where("content = ?", "must roll back").Count(&messages).Error)
require.NoError(t, db.Model(&model.BackgroundJob{}).Count(&jobs).Error)
require.Zero(t, messages)
require.Zero(t, jobs)
}
func TestMessageService_APIReplyCarriesRealtimeContext(t *testing.T) {
@@ -207,8 +336,8 @@ func TestMessageService_MessageEventContextResolvesAutomatedSenders(t *testing.T
t.Run(string(channelType), func(t *testing.T) {
require.NoError(t, db.Model(inbox).Update("channel_type", string(channelType)).Error)
event := channel.NewChannelEvent(channel.EventMessageCreated, channelType, account.ID, inbox.ID)
message := &model.Message{ConversationID: conversation.ID, InboxID: inbox.ID, SenderID: &tt.senderID, SenderType: tt.senderType}
svc.addMessageEventContext(context.Background(), event, message)
message := &model.Message{AccountID: account.ID, ConversationID: conversation.ID, InboxID: inbox.ID, SenderID: &tt.senderID, SenderType: tt.senderType}
require.NoError(t, svc.addMessageEventContext(context.Background(), db, event, message))
assert.IsType(t, &model.Contact{}, event.Data["contact"])
if channelType == channel.ChannelWebWidget {
@@ -84,6 +84,22 @@ func setupServiceTestDB(t *testing.T) *gorm.DB {
return db
}
func failBackgroundJobCreate(t *testing.T, db *gorm.DB) {
t.Helper()
name := "test:fail_background_job_create:" + t.Name()
requireNoError := func(err error) {
if err != nil {
t.Fatalf("configure background job failure: %v", err)
}
}
requireNoError(db.Callback().Create().Before("gorm:create").Register(name, func(tx *gorm.DB) {
if tx.Statement.Schema != nil && tx.Statement.Schema.Table == "background_jobs" {
tx.AddError(fmt.Errorf("forced background job insert failure"))
}
}))
t.Cleanup(func() { requireNoError(db.Callback().Create().Remove(name)) })
}
// marshalAssistantConfig converts an arbitrary value to json.RawMessage for the Config field.
func marshalAssistantConfig(cfg interface{}) json.RawMessage {
b, _ := json.Marshal(cfg)
@@ -343,14 +359,32 @@ func createTestInbox(t *testing.T, db *gorm.DB, accountID uint, channelType stri
// createTestConversation creates a test Conversation in the database.
func createTestConversation(t *testing.T, db *gorm.DB, accountID, inboxID, contactID uint) *model.Conversation {
t.Helper()
var inbox model.Inbox
if err := db.First(&inbox, inboxID).Error; err != nil {
t.Fatalf("failed to load test inbox: %v", err)
}
conv := &model.Conversation{
AccountID: accountID,
InboxID: inboxID,
ContactID: contactID,
Status: string(model.ConversationStatusOpen),
Priority: string(model.ConversationPriorityMedium),
ChannelType: "web_widget",
Channel: "web_widget",
ChannelType: inbox.ChannelType,
Channel: inbox.ChannelType,
}
if inbox.ChannelType == "web_widget" || inbox.ChannelType == string(model.InboxChannelTypeWebWidget) {
var contactInbox model.ContactInbox
if err := db.Where("contact_id = ? AND inbox_id = ?", contactID, inboxID).First(&contactInbox).Error; err != nil {
contactInbox = model.ContactInbox{
ContactID: contactID, InboxID: inboxID,
SourceID: fmt.Sprintf("test-source-%d-%d", contactID, inboxID),
PubsubToken: fmt.Sprintf("test-token-%d-%d", contactID, inboxID),
}
if err := db.Create(&contactInbox).Error; err != nil {
t.Fatalf("failed to create test contact inbox: %v", err)
}
}
conv.ContactInboxID = &contactInbox.ID
}
if err := db.Create(conv).Error; err != nil {
t.Fatalf("failed to create test conversation: %v", err)
+76 -54
View File
@@ -20,6 +20,7 @@ import (
"github.com/gochat/gochat/internal/search"
"github.com/gochat/gochat/internal/worker"
ws "github.com/gochat/gochat/internal/ws"
"github.com/gochat/gochat/internal/wsevent"
applogger "github.com/gochat/gochat/pkg/logger"
"gorm.io/datatypes"
"gorm.io/gorm"
@@ -66,6 +67,7 @@ type WidgetService struct {
campaignRepo *repository.CampaignRepo
worker *worker.WorkerPool
dispatcher *channel.Dispatcher
realtime *wsevent.BridgeListener
}
// NewWidgetService creates a new Widget service.
@@ -114,6 +116,10 @@ func (s *WidgetService) SetDispatcher(dispatcher *channel.Dispatcher) {
s.dispatcher = dispatcher
}
func (s *WidgetService) SetRealtimeEventBridge(bridge *wsevent.BridgeListener) {
s.realtime = bridge
}
// --- DTOs ---
// WidgetInitRequest is the DTO for the /widget/init endpoint.
@@ -407,51 +413,43 @@ func (s *WidgetService) SendMessage(ctx context.Context, req WidgetSendMessageRe
}
}
attachments, err := s.createIncomingMessage(ctx, conversation, &msg, req.AttachmentIDs)
if err != nil {
return nil, err
}
// Dispatch the same complete lifecycle context Chatwoot gives ActionCable.
if s.dispatcher != nil {
inbox, _ := s.inboxRepo.FindByID(ctx, conversation.InboxID)
data := map[string]interface{}{
var messageEvent *channel.ChannelEvent
var eventData map[string]interface{}
if s.dispatcher != nil || s.realtime != nil {
inbox, err := s.inboxRepo.FindByID(ctx, conversation.InboxID)
if err != nil {
return nil, fmt.Errorf("load widget inbox %d: %w", conversation.InboxID, err)
}
eventData = map[string]interface{}{
"inbox": inbox,
"conversation": conversation,
"contact": &contactInbox.Contact,
"widget_token": contactInbox.PubsubToken,
"channel_type": "web_widget",
}
messageEvent = widgetMessageCreatedEvent(channel.ChannelWebWidget, conversation, &msg, eventData)
}
attachments, err := s.createIncomingMessage(ctx, conversation, &msg, req.AttachmentIDs, messageEvent)
if err != nil {
return nil, err
}
// Dispatch the same complete lifecycle context Chatwoot gives ActionCable.
if s.dispatcher != nil {
if conversationCreated {
event := &channel.ChannelEvent{
Type: channel.EventConversationCreated, Channel: channel.ChannelWebWidget,
ConversationID: conversation.ID, InboxID: conversation.InboxID,
AccountID: conversation.AccountID, ContactID: conversation.ContactID,
Timestamp: time.Now().Unix(), Data: data,
Timestamp: time.Now().Unix(), Data: eventData,
}
if dispatchErr := s.dispatcher.Dispatch(ctx, event); dispatchErr != nil {
applogger.L().Warnf("widget conversation event dispatch failed: inbox=%d conv=%d err=%v",
conversation.InboxID, conversation.ID, dispatchErr)
if err := s.dispatcher.Dispatch(ctx, event); err != nil {
return nil, fmt.Errorf("dispatch widget conversation event: %w", err)
}
}
messageData := make(map[string]interface{}, len(data)+1)
for key, value := range data {
messageData[key] = value
}
messageData["message"] = &msg
event := &channel.ChannelEvent{
Type: channel.EventMessageCreated,
Channel: channel.ChannelWebWidget,
ConversationID: conversation.ID,
InboxID: conversation.InboxID,
AccountID: conversation.AccountID,
Timestamp: time.Now().Unix(),
ContactID: conversation.ContactID,
Data: messageData,
}
if dispatchErr := s.dispatcher.Dispatch(ctx, event); dispatchErr != nil {
applogger.L().Warnf("widget message event dispatch failed: inbox=%d conv=%d err=%v",
conversation.InboxID, conversation.ID, dispatchErr)
if err := s.dispatcher.Dispatch(ctx, messageEvent); err != nil {
return nil, fmt.Errorf("dispatch widget message event: %w", err)
}
}
@@ -465,11 +463,26 @@ func (s *WidgetService) SendMessage(ctx context.Context, req WidgetSendMessageRe
}, nil
}
func (s *WidgetService) createIncomingMessage(ctx context.Context, conversation *model.Conversation, message *model.Message, attachmentIDs []string) ([]model.Attachment, error) {
func widgetMessageCreatedEvent(eventChannel channel.ChannelType, conversation *model.Conversation, message *model.Message, data map[string]interface{}) *channel.ChannelEvent {
messageData := make(map[string]interface{}, len(data)+1)
for key, value := range data {
messageData[key] = value
}
messageData["message"] = message
return &channel.ChannelEvent{
Type: channel.EventMessageCreated, Channel: eventChannel,
ConversationID: conversation.ID, InboxID: conversation.InboxID,
AccountID: conversation.AccountID, ContactID: conversation.ContactID,
Timestamp: time.Now().Unix(), Data: messageData,
}
}
func (s *WidgetService) createIncomingMessage(ctx context.Context, conversation *model.Conversation, message *model.Message, attachmentIDs []string, event *channel.ChannelEvent) ([]model.Attachment, error) {
var messageTimestamp int64
var attachments []model.Attachment
var captainJob *model.BackgroundJob
var captainJobCreated bool
var realtimeJobs []*model.BackgroundJob
reopen := conversation != nil && !conversation.Muted &&
(conversation.Status == string(model.ConversationStatusSnoozed) || conversation.Status == string(model.ConversationStatusResolved))
if err := s.messageRepo.DB().WithContext(ctx).Transaction(func(tx *gorm.DB) error {
@@ -496,6 +509,17 @@ func (s *WidgetService) createIncomingMessage(ctx context.Context, conversation
return err
}
captainJob, captainJobCreated, err = enqueueCaptainConversationResponseForMessageInTransaction(ctx, s.worker, tx, message.ID)
if err != nil {
return err
}
if s.realtime != nil && event != nil {
var committedConversation model.Conversation
if err := tx.Where("id = ? AND account_id = ?", conversation.ID, conversation.AccountID).First(&committedConversation).Error; err != nil {
return err
}
event.Data["conversation"] = &committedConversation
realtimeJobs, err = s.realtime.EnqueueInTransaction(ctx, tx, event)
}
return err
}); err != nil {
return nil, fmt.Errorf("failed to create message: %w", err)
@@ -509,6 +533,9 @@ func (s *WidgetService) createIncomingMessage(ctx context.Context, conversation
if captainJobCreated {
s.worker.Publish(ctx, captainJob)
}
if s.realtime != nil {
s.realtime.PublishEnqueued(ctx, realtimeJobs)
}
return attachments, nil
}
@@ -1081,34 +1108,29 @@ func (s *WidgetService) PublicCreateMessage(ctx context.Context, inboxIdentifier
Status: "sent",
SourceID: req.EchoID,
}
attachments, err := s.createIncomingMessage(ctx, conversation, message, req.AttachmentIDs)
var messageEvent *channel.ChannelEvent
if s.dispatcher != nil || s.realtime != nil {
inbox, err := s.inboxRepo.FindByID(ctx, conversation.InboxID)
if err != nil {
return nil, nil, nil, fmt.Errorf("load widget inbox %d: %w", conversation.InboxID, err)
}
messageEvent = widgetMessageCreatedEvent(channel.ChannelAPI, conversation, message, map[string]interface{}{
"inbox": inbox,
"conversation": conversation,
"contact": &contactInbox.Contact,
"widget_token": contactInbox.PubsubToken,
"channel_type": "api",
})
}
attachments, err := s.createIncomingMessage(ctx, conversation, message, req.AttachmentIDs, messageEvent)
if err != nil {
return nil, nil, nil, err
}
// Public API messages use the same committed lifecycle payload as widget messages.
if s.dispatcher != nil {
inbox, _ := s.inboxRepo.FindByID(ctx, conversation.InboxID)
event := &channel.ChannelEvent{
Type: channel.EventMessageCreated,
Channel: channel.ChannelAPI,
ConversationID: conversation.ID,
InboxID: conversation.InboxID,
AccountID: conversation.AccountID,
ContactID: conversation.ContactID,
Timestamp: time.Now().Unix(),
Data: map[string]interface{}{
"inbox": inbox,
"conversation": conversation,
"contact": &contactInbox.Contact,
"widget_token": contactInbox.PubsubToken,
"channel_type": "api",
"message": message,
},
}
if dispatchErr := s.dispatcher.Dispatch(ctx, event); dispatchErr != nil {
applogger.L().Warnf("widget message event dispatch failed: inbox=%d conv=%d err=%v",
conversation.InboxID, conversation.ID, dispatchErr)
if err := s.dispatcher.Dispatch(ctx, messageEvent); err != nil {
return message, conversation, attachments, fmt.Errorf("dispatch widget message event: %w", err)
}
}
+152 -1
View File
@@ -6,10 +6,12 @@ import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"testing"
"time"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/driver/sqlite"
@@ -21,17 +23,20 @@ import (
channelmodel "github.com/gochat/gochat/internal/model/channel"
"github.com/gochat/gochat/internal/repository"
"github.com/gochat/gochat/internal/worker"
wspkg "github.com/gochat/gochat/internal/ws"
"github.com/gochat/gochat/internal/wsevent"
)
type widgetLifecycleListener struct {
events []*channel.ChannelEvent
err error
}
func (l *widgetLifecycleListener) Name() string { return "widget-lifecycle-test" }
func (l *widgetLifecycleListener) OnEvent(_ context.Context, event *channel.ChannelEvent) error {
l.events = append(l.events, event)
return nil
return l.err
}
func TestIsGenericShangwutongName(t *testing.T) {
@@ -136,6 +141,139 @@ func TestWidgetService_SendMessageDispatchesCompleteLifecycle(t *testing.T) {
assert.IsType(t, &model.Message{}, listener.events[1].Data["message"])
}
func TestWidgetService_SendMessageReturnsDispatchError(t *testing.T) {
db, svc := setupWidgetServiceTest(t)
seedWidgetInbox(t, db)
initResp, err := svc.Init(context.Background(), WidgetInitRequest{WebsiteToken: "test_ws_token_123"})
require.NoError(t, err)
dispatchErr := errors.New("realtime unavailable")
dispatcher := channel.NewDispatcher()
dispatcher.Register(&widgetLifecycleListener{err: dispatchErr})
svc.SetDispatcher(dispatcher)
_, err = svc.SendMessage(context.Background(), WidgetSendMessageRequest{
WidgetToken: initResp.WidgetToken,
Content: "hello from widget",
})
require.ErrorIs(t, err, dispatchErr)
}
func TestWidgetService_RealtimeFailureDoesNotBlockCaptain(t *testing.T) {
db, svc := setupWidgetServiceTest(t)
account, inbox := seedWidgetInbox(t, db)
assistant := &model.CaptainAssistant{AccountID: account.ID, Name: "Captain", Status: model.AssistantStatusActive, Config: json.RawMessage(`{}`)}
require.NoError(t, db.Create(assistant).Error)
require.NoError(t, db.Create(&model.CaptainInbox{AccountID: account.ID, InboxID: inbox.ID, AssistantID: assistant.ID}).Error)
require.NoError(t, db.Create(&model.CaptainPreference{AccountID: account.ID, AutoReplyEnabled: true}).Error)
pool := worker.NewWorkerPoolWithOptions(db, worker.WithBackoff(func(int) time.Duration { return 0 }))
svc.SetWorkerPool(pool)
rdb := redis.NewClient(&redis.Options{Addr: "127.0.0.1:1", MaxRetries: -1, DialTimeout: 10 * time.Millisecond})
t.Cleanup(func() { require.NoError(t, rdb.Close()) })
publisher := wspkg.NewEventPublisher(nil, nil, wspkg.NewBroadcastRelay(rdb, nil))
publisher.SetWorkerPool(pool)
dispatcher := channel.NewDispatcher()
bridge := wsevent.New(publisher)
dispatcher.Register(bridge)
svc.SetDispatcher(dispatcher)
svc.SetRealtimeEventBridge(bridge)
initResp, err := svc.Init(context.Background(), WidgetInitRequest{WebsiteToken: "test_ws_token_123"})
require.NoError(t, err)
response, err := svc.SendMessage(context.Background(), WidgetSendMessageRequest{WidgetToken: initResp.WidgetToken, Content: "ask captain"})
require.NoError(t, err)
var captainJob model.BackgroundJob
require.NoError(t, db.Where("job_type = ?", TaskTypeCaptainConversationResponseBuilder).First(&captainJob).Error)
require.Equal(t, model.BackgroundJobStatusQueued, captainJob.Status)
require.NoError(t, db.Model(&captainJob).Update("scheduled_at", time.Now().Add(time.Hour)).Error)
processed, err := pool.ProcessOne(context.Background())
require.True(t, processed)
require.ErrorContains(t, err, "connect: connection refused")
var messageCount int64
require.NoError(t, db.Model(&model.Message{}).Where("id = ?", response.Message.ID).Count(&messageCount).Error)
require.Equal(t, int64(1), messageCount)
require.NoError(t, db.First(&captainJob, captainJob.ID).Error)
require.Equal(t, model.BackgroundJobStatusQueued, captainJob.Status)
var realtimeJob model.BackgroundJob
require.NoError(t, db.Where("queue = ? AND last_error != ''", "events").First(&realtimeJob).Error)
require.Equal(t, model.BackgroundJobStatusRetrying, realtimeJob.Status)
}
func TestWidgetService_RealtimeJobFailureRollsBackMessage(t *testing.T) {
db, svc := setupWidgetServiceTest(t)
seedWidgetInbox(t, db)
pool := worker.NewWorkerPool(db, nil)
svc.SetWorkerPool(pool)
publisher := wspkg.NewEventPublisherLocal(nil, nil)
publisher.SetWorkerPool(pool)
bridge := wsevent.New(publisher)
dispatcher := channel.NewDispatcher()
dispatcher.Register(bridge)
svc.SetDispatcher(dispatcher)
svc.SetRealtimeEventBridge(bridge)
initResp, err := svc.Init(context.Background(), WidgetInitRequest{WebsiteToken: "test_ws_token_123"})
require.NoError(t, err)
failBackgroundJobCreate(t, db)
_, err = svc.SendMessage(context.Background(), WidgetSendMessageRequest{WidgetToken: initResp.WidgetToken, Content: "must roll back"})
require.ErrorContains(t, err, "forced background job insert failure")
var messages, jobs int64
require.NoError(t, db.Model(&model.Message{}).Where("content = ?", "must roll back").Count(&messages).Error)
require.NoError(t, db.Model(&model.BackgroundJob{}).Count(&jobs).Error)
require.Zero(t, messages)
require.Zero(t, jobs)
}
func TestWidgetService_RealtimeJobsRecoverAfterCommitWithoutRedisWakeup(t *testing.T) {
db, svc := setupWidgetServiceTest(t)
seedWidgetInbox(t, db)
pool := worker.NewWorkerPool(db, nil)
svc.SetWorkerPool(pool)
publisher := wspkg.NewEventPublisherLocal(nil, nil)
publisher.SetWorkerPool(pool)
bridge := wsevent.New(publisher)
dispatcher := channel.NewDispatcher()
dispatcher.Register(bridge)
svc.SetDispatcher(dispatcher)
svc.SetRealtimeEventBridge(bridge)
initResp, err := svc.Init(context.Background(), WidgetInitRequest{WebsiteToken: "test_ws_token_123"})
require.NoError(t, err)
response, err := svc.SendMessage(context.Background(), WidgetSendMessageRequest{WidgetToken: initResp.WidgetToken, Content: "recover me"})
require.NoError(t, err)
var queued int64
require.NoError(t, db.Model(&model.BackgroundJob{}).Where("idempotency_key LIKE ? AND status = ?", "realtime:message.created:%", model.BackgroundJobStatusQueued).Count(&queued).Error)
require.Equal(t, int64(2), queued)
sse := wspkg.NewSSERegistry()
accountEvents := sse.Subscribe("restart", initResp.AccountID, 1)
restartPool := worker.NewWorkerPool(db, nil)
restartPublisher := wspkg.NewEventPublisherLocal(nil, sse)
restartPublisher.SetWorkerPool(restartPool)
for range 2 {
processed, processErr := restartPool.ProcessOne(context.Background())
require.True(t, processed)
require.NoError(t, processErr)
}
select {
case event := <-accountEvents.Events:
require.Equal(t, wspkg.EventMessageCreated, event.Type)
case <-time.After(time.Second):
t.Fatal("restarted worker did not recover account event")
}
var messages, completed int64
require.NoError(t, db.Model(&model.Message{}).Where("id = ?", response.Message.ID).Count(&messages).Error)
require.NoError(t, db.Model(&model.BackgroundJob{}).Where("idempotency_key LIKE ? AND status = ?", "realtime:message.created:%", model.BackgroundJobStatusCompleted).Count(&completed).Error)
require.Equal(t, int64(1), messages)
require.Equal(t, int64(2), completed)
}
func TestWidgetService_SendMessageUpdatesExistingConversationActivityBeforeDispatch(t *testing.T) {
db, svc := setupWidgetServiceTest(t)
seedWidgetInbox(t, db)
@@ -214,6 +352,19 @@ func TestWidgetService_PublicCreateMessageUpdatesConversationBeforeDispatch(t *t
assert.Equal(t, wantTimestamp, *eventConversation.LastActivityAt)
}
func TestWidgetService_PublicCreateMessageReturnsDispatchError(t *testing.T) {
db, svc := setupWidgetServiceTest(t)
_, _, _, _, _, displayID, _ := seedPublicMessageTest(t, db, model.ConversationStatusOpen)
dispatchErr := errors.New("realtime unavailable")
dispatcher := channel.NewDispatcher()
dispatcher.Register(&widgetLifecycleListener{err: dispatchErr})
svc.SetDispatcher(dispatcher)
_, _, _, err := svc.PublicCreateMessage(context.Background(), "public-api", "visitor-source", displayID, PublicMessageRequest{Content: "hello"})
require.ErrorIs(t, err, dispatchErr)
}
func TestWidgetService_PublicCreateMessageRollsBackSecondAttachmentFailureAndRetries(t *testing.T) {
db, svc := setupWidgetServiceTest(t)
account, _, _, _, conversation, displayID, oldTimestamp := seedPublicMessageTest(t, db, model.ConversationStatusOpen)
+173 -26
View File
@@ -6,12 +6,31 @@ package ws
import (
"context"
"crypto/sha256"
"encoding/json"
"fmt"
"github.com/gochat/gochat/internal/model"
"github.com/gochat/gochat/internal/worker"
applogger "github.com/gochat/gochat/pkg/logger"
"gorm.io/gorm"
)
const taskTypeRealtimeEventPublish = "realtime:event_publish"
const (
realtimeTargetAccount = "account"
realtimeTargetToken = "pubsub_token"
)
type realtimeEventPublishJob struct {
AccountID uint `json:"account_id"`
PubsubToken string `json:"pubsub_token,omitempty"`
Target string `json:"target,omitempty"`
EventType string `json:"event_type"`
Payload interface{} `json:"payload"`
}
// EventPublisher routes business events to both WebSocket and SSE clients.
// It is the central point that services call after mutations to push real-time
// updates, matching Chatwoot's pattern where controllers broadcast events
@@ -29,9 +48,10 @@ import (
// entry point, with Hub and SSERegistry as local delivery targets, and Redis
// Pub/Sub relay for multi-instance fan-out.
type EventPublisher struct {
hub MessageHandler // WebSocket Hub (local delivery)
sse *SSERegistry // SSE registry (local delivery)
relay *BroadcastRelay // Redis Pub/Sub relay (cross-instance delivery)
hub MessageHandler // WebSocket Hub (local delivery)
sse *SSERegistry // SSE registry (local delivery)
relay *BroadcastRelay // Redis Pub/Sub relay (cross-instance delivery)
worker *worker.WorkerPool
}
// NewEventPublisher creates an EventPublisher with all delivery targets.
@@ -52,6 +72,15 @@ func NewEventPublisherLocal(hub MessageHandler, sse *SSERegistry) *EventPublishe
}
}
// SetWorkerPool makes realtime delivery durable. The request path only stores
// the job; Redis failures are retried and remain visible in background_jobs.
func (p *EventPublisher) SetWorkerPool(wp *worker.WorkerPool) {
p.worker = wp
if wp != nil {
wp.Register(taskTypeRealtimeEventPublish, p.performPublishJob)
}
}
// PublishEvent publishes a real-time event to all delivery targets.
// This is the primary API that services call after business mutations.
//
@@ -67,7 +96,14 @@ func NewEventPublisherLocal(hub MessageHandler, sse *SSERegistry) *EventPublishe
//
// Reference: Chatwoot controllers call broadcast_event after mutations,
// which triggers Wisper → ActionCable → Redis Pub/Sub relay.
func (p *EventPublisher) PublishEvent(accountID uint, eventType string, payload interface{}) {
func (p *EventPublisher) PublishEvent(accountID uint, eventType string, payload interface{}) error {
if p.worker != nil {
return p.enqueuePublishJobs(context.Background(), realtimePublishJobs(accountID, "", eventType, payload))
}
return p.publishEvent(accountID, eventType, payload)
}
func (p *EventPublisher) publishEvent(accountID uint, eventType string, payload interface{}) error {
// Build the wire-format message (matching Chatwoot ActionCable event format)
wsMsg := &WSMessage{
Event: eventType,
@@ -77,8 +113,7 @@ func (p *EventPublisher) PublishEvent(accountID uint, eventType string, payload
data, err := json.Marshal(wsMsg)
if err != nil {
applogger.L().Warnf("event publisher: failed to marshal event %s: %v", eventType, err)
return
return fmt.Errorf("marshal event %s: %w", eventType, err)
}
// 1. Deliver to WebSocket Hub (local clients)
@@ -98,9 +133,10 @@ func (p *EventPublisher) PublishEvent(accountID uint, eventType string, payload
if p.relay != nil {
room := accountRoomNameHelper(accountID)
if err := p.relay.Publish(context.Background(), room, wsMsg); err != nil {
applogger.L().Warnf("event publisher: redis publish failed for %s: %v", eventType, err)
return fmt.Errorf("publish event %s to account room: %w", eventType, err)
}
}
return nil
}
// PublishConversationEvent publishes a real-time event scoped to a specific
@@ -150,7 +186,24 @@ func (p *EventPublisher) PublishConversationEvent(accountID uint, conversationID
// contact's pubsub_token room. Chatwoot's widget ActionCable connector
// subscribes to RoomChannel with pubsub_token, so widget-visible events must be
// available on that token-scoped room in addition to the dashboard account room.
func (p *EventPublisher) PublishWidgetEvent(accountID uint, pubsubToken string, eventType string, payload interface{}) {
func (p *EventPublisher) PublishWidgetEvent(accountID uint, pubsubToken string, eventType string, payload interface{}) error {
if p.worker != nil {
return p.enqueuePublishJobs(context.Background(), realtimePublishJobs(accountID, pubsubToken, eventType, payload))
}
return p.publishWidgetEvent(accountID, pubsubToken, eventType, payload)
}
func (p *EventPublisher) publishWidgetEvent(accountID uint, pubsubToken string, eventType string, payload interface{}) error {
if err := p.publishEvent(accountID, eventType, payload); err != nil {
return err
}
if pubsubToken == "" {
return nil
}
return p.publishTokenEvent(accountID, pubsubToken, eventType, payload)
}
func (p *EventPublisher) publishTokenEvent(accountID uint, pubsubToken string, eventType string, payload interface{}) error {
wsMsg := &WSMessage{
Event: eventType,
Data: payload,
@@ -159,32 +212,126 @@ func (p *EventPublisher) PublishWidgetEvent(accountID uint, pubsubToken string,
data, err := json.Marshal(wsMsg)
if err != nil {
applogger.L().Warnf("event publisher: failed to marshal widget event %s: %v", eventType, err)
return
return fmt.Errorf("marshal widget event %s: %w", eventType, err)
}
if p.hub != nil {
p.hub.SendToAccount(accountID, data)
if pubsubToken != "" {
p.hub.SendToRoom(pubsubTokenRoomNameHelper(pubsubToken), data)
}
}
if p.sse != nil {
p.sse.SendToAccount(accountID, SSEEvent{Type: eventType, Payload: payload})
p.hub.SendToRoom(pubsubTokenRoomNameHelper(pubsubToken), data)
}
if p.relay != nil {
room := accountRoomNameHelper(accountID)
if err := p.relay.Publish(context.Background(), room, wsMsg); err != nil {
applogger.L().Warnf("event publisher: redis publish failed for %s: %v", eventType, err)
}
if pubsubToken != "" {
if err := p.relay.Publish(context.Background(), pubsubTokenRoomNameHelper(pubsubToken), wsMsg); err != nil {
applogger.L().Warnf("event publisher: redis publish failed for widget %s: %v", eventType, err)
}
if err := p.relay.Publish(context.Background(), pubsubTokenRoomNameHelper(pubsubToken), wsMsg); err != nil {
return fmt.Errorf("publish widget event %s to token room: %w", eventType, err)
}
}
return nil
}
func realtimePublishJobs(accountID uint, pubsubToken, eventType string, payload interface{}) []realtimeEventPublishJob {
jobs := []realtimeEventPublishJob{{AccountID: accountID, Target: realtimeTargetAccount, EventType: eventType, Payload: payload}}
if pubsubToken != "" {
jobs = append(jobs, realtimeEventPublishJob{AccountID: accountID, PubsubToken: pubsubToken, Target: realtimeTargetToken, EventType: eventType, Payload: payload})
}
return jobs
}
func realtimePublishJobOptions(job realtimeEventPublishJob) ([]worker.EnqueueOption, error) {
opts := []worker.EnqueueOption{worker.WithQueue("events"), worker.WithMaxAttempts(3)}
if job.EventType == EventMessageCreated {
encoded, err := json.Marshal(job)
if err != nil {
return nil, fmt.Errorf("marshal realtime publish job: %w", err)
}
opts = append(opts, worker.WithIdempotencyKey(fmt.Sprintf("realtime:message.created:%s:%x", job.Target, sha256.Sum256(encoded))))
}
return opts, nil
}
func (p *EventPublisher) enqueuePublishJobs(ctx context.Context, jobs []realtimeEventPublishJob) error {
for _, job := range jobs {
opts, err := realtimePublishJobOptions(job)
if err != nil {
return err
}
if _, err := p.worker.Enqueue(ctx, taskTypeRealtimeEventPublish, job, opts...); err != nil {
applogger.L().Errorf("enqueue realtime event %s failed: %v", job.EventType, err)
return err
}
}
return nil
}
// EnqueueInTransaction atomically persists every target channel with the
// caller's business mutation. The returned jobs are safe to publish only after
// the transaction commits.
func (p *EventPublisher) EnqueueInTransaction(ctx context.Context, tx *gorm.DB, accountID uint, pubsubToken, eventType string, payload interface{}) ([]*model.BackgroundJob, error) {
if p.worker == nil {
return nil, worker.ErrWorkerDatabaseRequired
}
var createdJobs []*model.BackgroundJob
for _, job := range realtimePublishJobs(accountID, pubsubToken, eventType, payload) {
opts, err := realtimePublishJobOptions(job)
if err != nil {
return nil, err
}
backgroundJob, created, err := p.worker.EnqueueInTransaction(ctx, tx, taskTypeRealtimeEventPublish, job, opts...)
if err != nil {
return nil, err
}
if created {
createdJobs = append(createdJobs, backgroundJob)
}
}
return createdJobs, nil
}
// PublishEnqueued triggers Redis after the surrounding transaction commits.
// Database sweep recovery remains authoritative if this process exits here.
func (p *EventPublisher) PublishEnqueued(ctx context.Context, jobs []*model.BackgroundJob) {
for _, job := range jobs {
p.worker.Publish(ctx, job)
}
}
func (p *EventPublisher) performPublishJob(ctx context.Context, backgroundJob *model.BackgroundJob) error {
var job realtimeEventPublishJob
if err := json.Unmarshal(backgroundJob.Payload, &job); err != nil {
return worker.Permanent(fmt.Errorf("unmarshal realtime publish job: %w", err))
}
if job.AccountID == 0 || job.EventType == "" {
return worker.Permanent(fmt.Errorf("invalid realtime publish job: account_id=%d event_type=%q", job.AccountID, job.EventType))
}
var err error
switch job.Target {
case realtimeTargetToken:
err = p.publishTokenJob(job.AccountID, job.PubsubToken, job.EventType, job.Payload)
default:
err = p.publishAccountJob(job.AccountID, job.EventType, job.Payload)
}
if err != nil {
applogger.L().Errorf("realtime event job %d failed for %s: %v", backgroundJob.ID, job.EventType, err)
}
return err
}
func (p *EventPublisher) publishAccountJob(accountID uint, eventType string, payload interface{}) error {
if p.relay == nil {
return p.publishEvent(accountID, eventType, payload)
}
if err := p.relay.Publish(context.Background(), accountRoomNameHelper(accountID), &WSMessage{Event: eventType, Data: payload, AccountID: accountID}); err != nil {
return err
}
if p.sse != nil {
p.sse.SendToAccount(accountID, SSEEvent{Type: eventType, Payload: payload})
}
return nil
}
func (p *EventPublisher) publishTokenJob(accountID uint, pubsubToken, eventType string, payload interface{}) error {
if p.relay == nil {
return p.publishTokenEvent(accountID, pubsubToken, eventType, payload)
}
return p.relay.Publish(context.Background(), pubsubTokenRoomNameHelper(pubsubToken), &WSMessage{Event: eventType, Data: payload, AccountID: accountID})
}
// accountRoomNameHelper generates the room name for an account channel.
+100
View File
@@ -1,15 +1,63 @@
package ws
import (
"context"
"encoding/json"
"errors"
"fmt"
"sync"
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
"github.com/gochat/gochat/internal/model"
"github.com/gochat/gochat/internal/worker"
)
type failPublishChannelOnceHook struct {
mu sync.Mutex
channel string
failed bool
attempts map[string]int
succeeded map[string]int
}
func (h *failPublishChannelOnceHook) DialHook(next redis.DialHook) redis.DialHook { return next }
func (h *failPublishChannelOnceHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook {
return func(ctx context.Context, cmd redis.Cmder) error {
if cmd.Name() != "publish" || len(cmd.Args()) < 2 {
return next(ctx, cmd)
}
channel := fmt.Sprint(cmd.Args()[1])
h.mu.Lock()
h.attempts[channel]++
if channel == h.channel && !h.failed {
h.failed = true
h.mu.Unlock()
return errors.New("token room unavailable")
}
h.mu.Unlock()
err := next(ctx, cmd)
if err == nil {
h.mu.Lock()
h.succeeded[channel]++
h.mu.Unlock()
}
return err
}
}
func (h *failPublishChannelOnceHook) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook {
return next
}
// === EventPublisher Tests ===
// Uses mockHandler (from broadcast_test.go) as the MessageHandler implementation,
// plus SSERegistry for SSE delivery verification.
@@ -709,3 +757,55 @@ func TestEventPublisher_WidgetEvent_PubsubTokenRoomDelivery(t *testing.T) {
assert.Equal(t, EventMessageCreated, roomMsg.Event)
assert.Equal(t, uint(1), roomMsg.AccountID)
}
func TestEventPublisher_DurableWidgetPublishRetriesPartialFailure(t *testing.T) {
db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.BackgroundJob{}))
mini := miniredis.RunT(t)
rdb := redis.NewClient(&redis.Options{Addr: mini.Addr()})
t.Cleanup(func() { require.NoError(t, rdb.Close()) })
tokenChannel := RedisPrefixRoom + "pubsub_token_visitor"
hook := &failPublishChannelOnceHook{
channel: tokenChannel, attempts: map[string]int{}, succeeded: map[string]int{},
}
rdb.AddHook(hook)
pool := worker.NewWorkerPoolWithOptions(db, worker.WithBackoff(func(int) time.Duration { return 0 }))
publisher := NewEventPublisher(nil, nil, NewBroadcastRelay(rdb, nil))
publisher.SetWorkerPool(pool)
payload := map[string]interface{}{"id": 7, "content": "hello widget"}
require.NoError(t, publisher.PublishWidgetEvent(1, "visitor", EventMessageCreated, payload))
require.NoError(t, publisher.PublishWidgetEvent(1, "visitor", EventMessageCreated, payload))
var count int64
require.NoError(t, db.Model(&model.BackgroundJob{}).Where("job_type = ?", taskTypeRealtimeEventPublish).Count(&count).Error)
require.Equal(t, int64(2), count, "account and token jobs must each be idempotent")
processed, err := pool.ProcessOne(context.Background())
require.True(t, processed)
require.NoError(t, err)
processed, err = pool.ProcessOne(context.Background())
require.True(t, processed)
require.ErrorContains(t, err, "token room unavailable")
var job model.BackgroundJob
require.NoError(t, db.Where("idempotency_key LIKE ?", "realtime:message.created:pubsub_token:%").First(&job).Error)
require.Equal(t, model.BackgroundJobStatusRetrying, job.Status)
require.Contains(t, job.LastError, "token room unavailable")
processed, err = pool.ProcessOne(context.Background())
require.True(t, processed)
require.NoError(t, err)
require.NoError(t, db.First(&job, job.ID).Error)
require.Equal(t, model.BackgroundJobStatusCompleted, job.Status)
accountChannel := RedisPrefixRoom + "account_1"
hook.mu.Lock()
defer hook.mu.Unlock()
require.Equal(t, 1, hook.attempts[accountChannel])
require.Equal(t, 1, hook.succeeded[accountChannel])
require.Equal(t, 2, hook.attempts[tokenChannel])
require.Equal(t, 1, hook.succeeded[tokenChannel])
}
+40 -12
View File
@@ -16,6 +16,7 @@ import (
"github.com/gochat/gochat/internal/model"
wspkg "github.com/gochat/gochat/internal/ws"
applogger "github.com/gochat/gochat/pkg/logger"
"gorm.io/gorm"
)
// BridgeListener implements channel.EventListener and forwards events to the
@@ -49,14 +50,48 @@ func (l *BridgeListener) Name() string {
// The event type string is passed through directly (channel.EventType and
// wspkg event constants share the same string values).
func (l *BridgeListener) OnEvent(ctx context.Context, event *channel.ChannelEvent) error {
eventType := string(event.Type)
if !isWSEventType(eventType) {
accountID, pubsubToken, eventType, payload, ok := realtimeEvent(event)
if !ok {
return nil
}
if pubsubToken != "" {
if err := l.publisher.PublishWidgetEvent(accountID, pubsubToken, eventType, payload); err != nil {
return err
}
} else {
if err := l.publisher.PublishEvent(accountID, eventType, payload); err != nil {
return err
}
}
applogger.L().Debugf("ws_bridge: forwarded event %s for account %d", eventType, accountID)
return nil
}
// EnqueueInTransaction stores every realtime channel beside the business
// mutation. Call PublishEnqueued only after the caller's transaction commits.
func (l *BridgeListener) EnqueueInTransaction(ctx context.Context, tx *gorm.DB, event *channel.ChannelEvent) ([]*model.BackgroundJob, error) {
accountID, pubsubToken, eventType, payload, ok := realtimeEvent(event)
if !ok {
return nil, nil
}
return l.publisher.EnqueueInTransaction(ctx, tx, accountID, pubsubToken, eventType, payload)
}
func (l *BridgeListener) PublishEnqueued(ctx context.Context, jobs []*model.BackgroundJob) {
l.publisher.PublishEnqueued(ctx, jobs)
}
func realtimeEvent(event *channel.ChannelEvent) (uint, string, string, map[string]interface{}, bool) {
if event == nil {
return 0, "", "", nil, false
}
eventType := string(event.Type)
if !isWSEventType(eventType) {
return 0, "", "", nil, false
}
payload := wsEventPayload(event)
payload["account_id"] = event.AccountID
if event.ConversationID != 0 {
if _, exists := payload["conversation_id"]; !exists {
payload["conversation_id"] = event.ConversationID
@@ -65,15 +100,8 @@ func (l *BridgeListener) OnEvent(ctx context.Context, event *channel.ChannelEven
if event.InboxID != 0 {
payload["inbox_id"] = event.InboxID
}
if pubsubToken, _ := event.Data["widget_token"].(string); pubsubToken != "" {
l.publisher.PublishWidgetEvent(event.AccountID, pubsubToken, eventType, payload)
} else {
l.publisher.PublishEvent(event.AccountID, eventType, payload)
}
applogger.L().Debugf("ws_bridge: forwarded event %s for account %d", eventType, event.AccountID)
return nil
pubsubToken, _ := event.Data["widget_token"].(string)
return event.AccountID, pubsubToken, eventType, payload, true
}
// wsEventPayload converts internal dispatcher data into the flat push payload
@@ -60,6 +60,16 @@ func TestBridgeListenerRoutesWebWidgetEventsToDashboardAndVisitor(t *testing.T)
}
}
func TestBridgeListenerReturnsPublisherError(t *testing.T) {
event := channel.NewChannelEvent(channel.EventInboxCreated, channel.ChannelAPI, 1, 4)
event.Data["invalid"] = make(chan int)
err := New(wspkg.NewEventPublisherLocal(nil, nil)).OnEvent(context.Background(), event)
if err == nil {
t.Fatal("expected publisher error")
}
}
func TestBridgeListenerPreservesWebWidgetSenderTypeContract(t *testing.T) {
tests := []struct {
senderType string