diff --git a/backend/internal/service/message_service.go b/backend/internal/service/message_service.go index b4b54015..1a94f721 100644 --- a/backend/internal/service/message_service.go +++ b/backend/internal/service/message_service.go @@ -77,7 +77,7 @@ func (s *MessageService) dispatchMessageEvent(ctx context.Context, eventType cha } event.Data["message"] = message if eventType == channel.EventMessageCreated || eventType == channel.EventMessageUpdated || eventType == channel.EventMessageDeleted { - s.addWebWidgetEventContext(ctx, event, message) + 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 { @@ -85,13 +85,17 @@ func (s *MessageService) dispatchMessageEvent(ctx context.Context, eventType cha } } -func (s *MessageService) addWebWidgetEventContext(ctx context.Context, event *channel.ChannelEvent, message *model.Message) { +func (s *MessageService) addMessageEventContext(ctx context.Context, event *channel.ChannelEvent, message *model.Message) { var inbox model.Inbox - if err := s.repo.DB().WithContext(ctx).First(&inbox, message.InboxID).Error; err != nil || inbox.ChannelType != string(channel.ChannelWebWidget) { + if err := s.repo.DB().WithContext(ctx).First(&inbox, message.InboxID).Error; err != nil { return } - event.Channel = channel.ChannelWebWidget event.Data["inbox"] = &inbox + isWebWidget := strings.EqualFold(strings.TrimSpace(inbox.ChannelType), string(channel.ChannelWebWidget)) || + strings.EqualFold(strings.TrimSpace(inbox.ChannelType), string(model.InboxChannelTypeWebWidget)) + if isWebWidget { + event.Channel = channel.ChannelWebWidget + } var conversation model.Conversation if err := s.repo.DB().WithContext(ctx).First(&conversation, message.ConversationID).Error; err != nil { @@ -123,7 +127,7 @@ func (s *MessageService) addWebWidgetEventContext(ctx context.Context, event *ch } } } - if message.Private || message.MessageType == string(model.MessageTypeActivity) || conversation.ContactInboxID == nil { + if !isWebWidget || message.Private || message.MessageType == string(model.MessageTypeActivity) || conversation.ContactInboxID == nil { return } var contactInbox model.ContactInbox diff --git a/backend/internal/service/message_service_test.go b/backend/internal/service/message_service_test.go index c96303ec..dd9be5a8 100644 --- a/backend/internal/service/message_service_test.go +++ b/backend/internal/service/message_service_test.go @@ -140,7 +140,41 @@ func TestMessageService_WebWidgetReplyCarriesRealtimeContext(t *testing.T) { assert.IsType(t, &model.User{}, created.Data["sender"]) } -func TestMessageService_WebWidgetReplyResolvesAutomatedSenders(t *testing.T) { +func TestMessageService_APIReplyCarriesRealtimeContext(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) + + 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.NoError(t, err) + + var created *channel.ChannelEvent + for _, event := range listener.events { + if event.Type == channel.EventMessageCreated { + created = event + break + } + } + require.NotNil(t, created) + assert.Equal(t, channel.ChannelAPI, created.Channel) + assert.NotContains(t, created.Data, "widget_token") + assert.IsType(t, &model.Inbox{}, created.Data["inbox"]) + assert.IsType(t, &model.Conversation{}, created.Data["conversation"]) + assert.IsType(t, &model.Contact{}, created.Data["contact"]) + assert.IsType(t, &model.User{}, created.Data["sender"]) +} + +func TestMessageService_MessageEventContextResolvesAutomatedSenders(t *testing.T) { db, _, _, svc := setupMessageServiceWithDefaultLLM(t) require.NoError(t, db.AutoMigrate(&model.AgentBot{}, &model.CaptainAssistant{})) account := createTestAccount(t, db) @@ -169,16 +203,25 @@ func TestMessageService_WebWidgetReplyResolvesAutomatedSenders(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - event := channel.NewChannelEvent(channel.EventMessageCreated, channel.ChannelWebWidget, account.ID, inbox.ID) - message := &model.Message{ConversationID: conversation.ID, InboxID: inbox.ID, SenderID: &tt.senderID, SenderType: tt.senderType} - svc.addWebWidgetEventContext(context.Background(), event, message) + for _, channelType := range []channel.ChannelType{channel.ChannelWebWidget, channel.ChannelAPI} { + 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) - assert.Equal(t, "visitor-token", event.Data["widget_token"]) - assert.IsType(t, &model.Contact{}, event.Data["contact"]) - if tt.expected == nil { - assert.NotContains(t, event.Data, "sender") - } else { - assert.IsType(t, tt.expected, event.Data["sender"]) + assert.IsType(t, &model.Contact{}, event.Data["contact"]) + if channelType == channel.ChannelWebWidget { + assert.Equal(t, "visitor-token", event.Data["widget_token"]) + } else { + assert.NotContains(t, event.Data, "widget_token") + } + if tt.expected == nil { + assert.NotContains(t, event.Data, "sender") + } else { + assert.IsType(t, tt.expected, event.Data["sender"]) + } + }) } }) } diff --git a/backend/tests/e2e/websocket_multi_instance_e2e_test.go b/backend/tests/e2e/websocket_multi_instance_e2e_test.go index 626cff47..72126cbe 100644 --- a/backend/tests/e2e/websocket_multi_instance_e2e_test.go +++ b/backend/tests/e2e/websocket_multi_instance_e2e_test.go @@ -114,6 +114,7 @@ func TestWebSocketMultiInstanceFanout(t *testing.T) { t.Cleanup(instanceB.stop) authHeaders := signIn(t, baseURLA, seedData.AdminEmail, seedData.AdminPassword) + apiSeed := createAPIFanoutSeed(t, baseURLA, authHeaders, seedData.AccountID, stamp) connA, identifier := connectAccountCable(t, baseURLA, issueWSTicket(t, baseURLA, authHeaders), seedData.AccountID) defer connA.Close() connB, identifierB := connectAccountCable(t, baseURLB, issueWSTicket(t, baseURLB, authHeaders), seedData.AccountID) @@ -127,31 +128,36 @@ func TestWebSocketMultiInstanceFanout(t *testing.T) { senderObjectType string senderID uint widget bool + apiInbox bool attachment bool }{ {name: "agent_bot", senderType: "AgentBot", senderObjectType: "agent_bot", senderID: agentBotID}, {name: "contact", senderType: "Contact", senderObjectType: "contact", widget: true}, {name: "captain_assistant", senderType: "Captain::Assistant", senderObjectType: "captain_assistant", senderID: seedData.CaptainAssistantID}, - {name: "user", senderType: "User", senderObjectType: "user", attachment: true}, + {name: "user", senderType: "User", senderObjectType: "user", apiInbox: true, attachment: true}, } frames := make(map[string]any, len(senderCases)) for _, senderCase := range senderCases { content := "multi-instance " + senderCase.name + " " + stamp + targetSeed := seedData + if senderCase.apiInbox { + targetSeed = apiSeed + } var created map[string]any if senderCase.widget { - createWidgetMessage(t, baseURLA, seedData, content) + createWidgetMessage(t, baseURLA, targetSeed, content) } else { - created = createDashboardMessage(t, baseURLA, authHeaders, seedData, content, senderCase.senderType, senderCase.senderID, senderCase.attachment) + created = createDashboardMessage(t, baseURLA, authHeaders, targetSeed, content, senderCase.senderType, senderCase.senderID, senderCase.attachment) } frameA := readMessageCreated(t, connA, content) frameB := readMessageCreated(t, connB, content) if created == nil { - created = fetchHTTPMessage(t, baseURLA, authHeaders, seedData, content) + created = fetchHTTPMessage(t, baseURLA, authHeaders, targetSeed, content) } require.Equal(t, frameA, frameB, "%s sender must fan out identically", senderCase.name) - assertFanoutContract(t, frameA, identifier, created, seedData, content, senderCase.senderType, senderCase.senderObjectType, !senderCase.widget, senderCase.attachment) - assertFanoutContract(t, frameB, identifier, created, seedData, content, senderCase.senderType, senderCase.senderObjectType, !senderCase.widget, senderCase.attachment) + assertFanoutContract(t, frameA, identifier, created, targetSeed, content, senderCase.senderType, senderCase.senderObjectType, !senderCase.widget, senderCase.attachment) + assertFanoutContract(t, frameB, identifier, created, targetSeed, content, senderCase.senderType, senderCase.senderObjectType, !senderCase.widget, senderCase.attachment) frames[senderCase.name+"_a"] = frameA frames[senderCase.name+"_b"] = frameB } @@ -298,6 +304,45 @@ func signIn(t *testing.T, baseURL, email, password string) http.Header { return response.Header.Clone() } +func createAPIFanoutSeed(t *testing.T, baseURL string, authHeaders http.Header, accountID uint, stamp string) fanoutSeed { + t.Helper() + post := func(endpoint string, payload any) map[string]any { + body, err := json.Marshal(payload) + require.NoError(t, err) + request, err := http.NewRequest(http.MethodPost, endpoint, bytes.NewReader(body)) + require.NoError(t, err) + request.Header.Set("Content-Type", "application/json") + setAuthHeaders(request, authHeaders) + response, err := http.DefaultClient.Do(request) + require.NoError(t, err) + defer response.Body.Close() + responseBody, err := io.ReadAll(response.Body) + require.NoError(t, err) + require.Equal(t, http.StatusOK, response.StatusCode, "create API fanout fixture response: %s", responseBody) + var result map[string]any + require.NoError(t, json.Unmarshal(responseBody, &result)) + return result + } + + inbox := post(fmt.Sprintf("%s/api/v1/accounts/%d/inboxes", baseURL, accountID), map[string]any{ + "name": "WebSocket API Inbox " + stamp, + "channel": map[string]any{"type": "api"}, + }) + inboxID := uint(inbox["id"].(float64)) + contactResponse := post(fmt.Sprintf("%s/api/v1/accounts/%d/contacts", baseURL, accountID), map[string]any{ + "name": "WebSocket API Contact " + stamp, + "identifier": "ws-api-" + stamp, + "inbox_id": inboxID, + "source_id": "ws-api-source-" + stamp, + }) + contactID := uint(contactResponse["payload"].(map[string]any)["contact"].(map[string]any)["id"].(float64)) + conversation := post(fmt.Sprintf("%s/api/v1/accounts/%d/conversations", baseURL, accountID), map[string]any{ + "inbox_id": inboxID, + "contact_id": contactID, + }) + return fanoutSeed{AccountID: accountID, ConversationDisplayID: uint(conversation["id"].(float64))} +} + func issueWSTicket(t *testing.T, baseURL string, authHeaders http.Header) string { t.Helper() request, err := http.NewRequest(http.MethodPost, baseURL+"/api/v1/auth/ws_ticket", nil)