* fix(HH-564): harden durable realtime enqueue * fix(HH-564): wire production SSE stream --------- Co-authored-by: Rogee <rogee@ipao.vip>
1393 lines
45 KiB
Go
1393 lines
45 KiB
Go
package ws
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/alicebob/miniredis/v2"
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/gorilla/websocket"
|
|
"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/auth"
|
|
"github.com/gochat/gochat/internal/channel"
|
|
"github.com/gochat/gochat/internal/config"
|
|
v1 "github.com/gochat/gochat/internal/handler/api/v1"
|
|
widgethandler "github.com/gochat/gochat/internal/handler/widget"
|
|
"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) {
|
|
assert.Equal(t, CommandType("subscribe"), CommandSubscribe)
|
|
assert.Equal(t, CommandType("unsubscribe"), CommandUnsubscribe)
|
|
assert.Equal(t, CommandType("ping"), CommandPing)
|
|
assert.Equal(t, CommandType("message"), CommandMessage)
|
|
}
|
|
|
|
func TestServerMessageType_Constants(t *testing.T) {
|
|
assert.Equal(t, ServerMessageType("event"), ServerEvent)
|
|
assert.Equal(t, ServerMessageType("confirm_subscription"), ServerConfirmSubscribe)
|
|
assert.Equal(t, ServerMessageType("confirm_unsubscribe"), ServerConfirmUnsubscribe)
|
|
assert.Equal(t, ServerMessageType("reject_subscription"), ServerRejectSubscribe)
|
|
assert.Equal(t, ServerMessageType("ping"), ServerPing)
|
|
assert.Equal(t, ServerMessageType("welcome"), ServerWelcome)
|
|
assert.Equal(t, ServerMessageType("disconnect"), ServerDisconnect)
|
|
}
|
|
|
|
func TestChannelName_Constants(t *testing.T) {
|
|
assert.Equal(t, "AccountChannel", ChannelAccount)
|
|
assert.Equal(t, "ConversationChannel", ChannelConversation)
|
|
assert.Equal(t, "RoomChannel", ChannelRoom)
|
|
}
|
|
|
|
func TestCommandEvent_Constants(t *testing.T) {
|
|
assert.Equal(t, "message.created", EventMessageCreated)
|
|
assert.Equal(t, "conversation.resolved", EventConversationResolved)
|
|
assert.Equal(t, "agent.typing_on", EventAgentTypingOn)
|
|
}
|
|
|
|
func TestCommandFrame_JSONRoundTrip(t *testing.T) {
|
|
cmd := CommandFrame{
|
|
Command: CommandSubscribe,
|
|
Identifier: `{"channel":"AccountChannel","account_id":1}`,
|
|
}
|
|
data, err := json.Marshal(cmd)
|
|
require.NoError(t, err)
|
|
|
|
var decoded CommandFrame
|
|
err = json.Unmarshal(data, &decoded)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, cmd.Command, decoded.Command)
|
|
assert.Equal(t, cmd.Identifier, decoded.Identifier)
|
|
}
|
|
|
|
func TestChannelIdentifier_JSONRoundTrip(t *testing.T) {
|
|
ci := ChannelIdentifier{
|
|
Channel: ChannelConversation,
|
|
AccountID: 1,
|
|
ConversationID: 42,
|
|
}
|
|
data, err := json.Marshal(ci)
|
|
require.NoError(t, err)
|
|
|
|
var decoded ChannelIdentifier
|
|
err = json.Unmarshal(data, &decoded)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, ci.Channel, decoded.Channel)
|
|
assert.Equal(t, ci.AccountID, decoded.AccountID)
|
|
assert.Equal(t, ci.ConversationID, decoded.ConversationID)
|
|
}
|
|
|
|
func TestEventFrame_JSONRoundTrip(t *testing.T) {
|
|
ef := EventFrame{
|
|
Type: ServerEvent,
|
|
Event: EventMessageCreated,
|
|
Payload: map[string]interface{}{"id": 1},
|
|
Identifier: `{"channel":"AccountChannel","account_id":1}`,
|
|
}
|
|
data, err := json.Marshal(ef)
|
|
require.NoError(t, err)
|
|
|
|
var decoded EventFrame
|
|
err = json.Unmarshal(data, &decoded)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, ef.Type, decoded.Type)
|
|
assert.Equal(t, ef.Event, decoded.Event)
|
|
}
|
|
|
|
func TestConfirmFrame_JSONRoundTrip(t *testing.T) {
|
|
cf := ConfirmFrame{
|
|
Type: ServerConfirmSubscribe,
|
|
Identifier: `{"channel":"AccountChannel","account_id":1}`,
|
|
}
|
|
data, err := json.Marshal(cf)
|
|
require.NoError(t, err)
|
|
|
|
var decoded ConfirmFrame
|
|
err = json.Unmarshal(data, &decoded)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, cf.Type, decoded.Type)
|
|
}
|
|
|
|
func TestRejectFrame_JSONRoundTrip(t *testing.T) {
|
|
rf := RejectFrame{
|
|
Type: ServerRejectSubscribe,
|
|
Identifier: `{"channel":"AccountChannel","account_id":1}`,
|
|
Reason: "account_id mismatch",
|
|
}
|
|
data, err := json.Marshal(rf)
|
|
require.NoError(t, err)
|
|
|
|
var decoded RejectFrame
|
|
err = json.Unmarshal(data, &decoded)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, rf.Reason, decoded.Reason)
|
|
}
|
|
|
|
func TestWritePump_ActionCableHeartbeatCompatibility(t *testing.T) {
|
|
serverConn, clientConn := newTestWSConn(t)
|
|
defer clientConn.Close()
|
|
|
|
client := NewClient(1, 10, serverConn, NewHubSimple())
|
|
done := make(chan struct{})
|
|
go func() {
|
|
(&Handler{}).writePump(client)
|
|
close(done)
|
|
}()
|
|
defer func() {
|
|
close(client.Send)
|
|
<-done
|
|
}()
|
|
|
|
require.NoError(t, clientConn.SetReadDeadline(time.Now().Add(8*time.Second)))
|
|
receivedAt := make([]time.Time, 2)
|
|
for i := range receivedAt {
|
|
messageType, data, err := clientConn.ReadMessage()
|
|
receivedAt[i] = time.Now()
|
|
require.NoError(t, err)
|
|
require.Equal(t, websocket.TextMessage, messageType)
|
|
|
|
var frame struct {
|
|
Type ServerMessageType `json:"type"`
|
|
Message json.RawMessage `json:"message"`
|
|
}
|
|
require.NoError(t, json.Unmarshal(data, &frame))
|
|
require.Equal(t, ServerPing, frame.Type)
|
|
require.Regexp(t, `^[0-9]+$`, string(frame.Message), "message must be an unquoted Unix timestamp")
|
|
|
|
var unixSeconds int64
|
|
require.NoError(t, json.Unmarshal(frame.Message, &unixSeconds))
|
|
assert.InDelta(t, receivedAt[i].Unix(), unixSeconds, 1)
|
|
}
|
|
|
|
assert.WithinDuration(t, receivedAt[0].Add(3*time.Second), receivedAt[1], time.Second)
|
|
}
|
|
|
|
func TestWelcomeFrame_JSONRoundTrip(t *testing.T) {
|
|
wf := WelcomeFrame{Type: ServerWelcome}
|
|
data, err := json.Marshal(wf)
|
|
require.NoError(t, err)
|
|
|
|
var decoded WelcomeFrame
|
|
err = json.Unmarshal(data, &decoded)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, wf.Type, decoded.Type)
|
|
}
|
|
|
|
func TestDisconnectFrame_JSONRoundTrip(t *testing.T) {
|
|
df := DisconnectFrame{
|
|
Type: ServerDisconnect,
|
|
Reason: "server shutdown",
|
|
Reconnect: false,
|
|
}
|
|
data, err := json.Marshal(df)
|
|
require.NoError(t, err)
|
|
|
|
var decoded DisconnectFrame
|
|
err = json.Unmarshal(data, &decoded)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, df.Reason, decoded.Reason)
|
|
assert.Equal(t, df.Reconnect, decoded.Reconnect)
|
|
}
|
|
|
|
// --- Hub Tests ---
|
|
|
|
func TestNewHubSimple(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
require.NotNil(t, hub)
|
|
assert.NotNil(t, hub.clients)
|
|
assert.NotNil(t, hub.rooms)
|
|
assert.NotNil(t, hub.accounts)
|
|
assert.NotNil(t, hub.commandChan)
|
|
}
|
|
|
|
func TestGenerateClientID(t *testing.T) {
|
|
id1 := generateClientID()
|
|
id2 := generateClientID()
|
|
assert.NotEmpty(t, id1)
|
|
assert.NotEmpty(t, id2)
|
|
assert.NotEqual(t, id1, id2)
|
|
assert.Len(t, id1, 32) // 16 bytes hex = 32 chars
|
|
}
|
|
|
|
func TestNewClient(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
client := NewClient(1, 10, nil, hub)
|
|
require.NotNil(t, client)
|
|
assert.Equal(t, uint(1), client.UserID)
|
|
assert.Equal(t, uint(10), client.AccountID)
|
|
assert.NotNil(t, client.Send)
|
|
assert.NotNil(t, client.SubscribedRooms)
|
|
assert.NotEmpty(t, client.ID)
|
|
}
|
|
|
|
func TestClient_SubscribeUnsubscribe(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
client := NewClient(1, 10, nil, hub)
|
|
|
|
room := "account_10"
|
|
client.Subscribe(room)
|
|
assert.True(t, client.SubscribedRooms[room])
|
|
assert.True(t, hub.rooms[room][client.ID])
|
|
|
|
client.Unsubscribe(room)
|
|
_, exists := client.SubscribedRooms[room]
|
|
assert.False(t, exists)
|
|
}
|
|
|
|
func TestHub_RegisterUnregister(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
client := NewClient(1, 10, nil, hub)
|
|
|
|
hub.Register(client)
|
|
assert.Contains(t, hub.clients, client.ID)
|
|
roomName := accountRoomName(10)
|
|
assert.True(t, client.SubscribedRooms[roomName])
|
|
|
|
hub.Unregister(client)
|
|
_, exists := hub.clients[client.ID]
|
|
assert.False(t, exists)
|
|
}
|
|
|
|
func TestHub_Unregister_AlreadyUnregistered(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
client := NewClient(1, 10, nil, hub)
|
|
hub.Unregister(client) // should not panic
|
|
}
|
|
|
|
func TestHub_SendToAccount(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
client := NewClient(1, 10, nil, hub)
|
|
client.Identifier = `{"channel":"AccountChannel","account_id":10}`
|
|
hub.Register(client)
|
|
|
|
hub.SendToAccount(10, []byte(`{"event":"test"}`))
|
|
|
|
select {
|
|
case msg := <-client.Send:
|
|
assert.NotEmpty(t, msg)
|
|
default:
|
|
t.Fatal("expected to receive message")
|
|
}
|
|
}
|
|
|
|
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)
|
|
sse := wspkg.NewSSERegistry()
|
|
dashboardSSE := sse.Subscribe("dashboard", 1, 1)
|
|
|
|
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, sse, 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")
|
|
}
|
|
select {
|
|
case event := <-dashboardSSE.Events:
|
|
require.Equal(t, wspkg.EventMessageCreated, event.Type)
|
|
case <-time.After(time.Second):
|
|
t.Fatal("SSE 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 <-dashboardSSE.Events:
|
|
t.Fatal("SSE 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)
|
|
agent.Identifier = `{"channel":"AccountChannel","account_id":10}`
|
|
hub.Register(agent)
|
|
visitor := NewClient(2, 10, nil, hub)
|
|
visitor.IsContact = true
|
|
visitor.Identifier = `{"channel":"AccountChannel","account_id":10}`
|
|
hub.Register(visitor)
|
|
hub.subscribeClient(visitor.ID, accountRoomName(10))
|
|
visitor.SubscribedRooms[accountRoomName(10)] = true
|
|
|
|
hub.SendToAccount(10, []byte(`{"event":"message.created","data":{"content":"same reply","sender_type":"AgentBot","sender_id":7,"ai_takeover_active":true,"additional_attributes":{"agent_name":"Captain"}}}`))
|
|
|
|
var agentMessage, visitorMessage []byte
|
|
select {
|
|
case agentMessage = <-agent.Send:
|
|
default:
|
|
t.Fatal("expected agent event")
|
|
}
|
|
select {
|
|
case visitorMessage = <-visitor.Send:
|
|
default:
|
|
t.Fatal("expected visitor event")
|
|
}
|
|
assert.Contains(t, string(agentMessage), "AgentBot")
|
|
assert.NotContains(t, string(visitorMessage), "AgentBot")
|
|
assert.NotContains(t, string(visitorMessage), "sender_id")
|
|
assert.NotContains(t, string(visitorMessage), "sender_type")
|
|
assert.NotContains(t, string(visitorMessage), "ai_takeover_active")
|
|
assert.NotContains(t, string(visitorMessage), "agent_name")
|
|
assert.NotContains(t, string(visitorMessage), "Captain")
|
|
}
|
|
|
|
func TestHub_VisitorRoomAndClientDeliverySanitizeIdentity(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
visitor := NewClient(2, 10, nil, hub)
|
|
visitor.IsContact = true
|
|
visitor.Identifier = `{"channel":"RoomChannel"}`
|
|
hub.Register(visitor)
|
|
hub.subscribeClient(visitor.ID, "visitor-room")
|
|
visitor.SubscribedRooms["visitor-room"] = true
|
|
data := []byte(`{"event":"message.created","data":{"sender_type":"Captain::Assistant","sender_name":"Captain"}}`)
|
|
|
|
hub.SendToRoom("visitor-room", data)
|
|
hub.SendToClient(visitor.ID, data)
|
|
for range 2 {
|
|
select {
|
|
case message := <-visitor.Send:
|
|
assert.NotContains(t, string(message), "Captain")
|
|
assert.NotContains(t, string(message), "sender_type")
|
|
default:
|
|
t.Fatal("expected sanitized visitor event")
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestHub_SendToAccount_NoClients(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
hub.SendToAccount(10, []byte(`{"event":"test"}`)) // should not panic
|
|
}
|
|
|
|
func TestHub_SendToAccountConversation(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
client := NewClient(1, 10, nil, hub)
|
|
client.Identifier = `{"channel":"ConversationChannel","account_id":10,"conversation_id":5}`
|
|
hub.Register(client)
|
|
room := conversationRoomName(10, 5)
|
|
hub.subscribeClient(client.ID, room)
|
|
client.SubscribedRooms[room] = true
|
|
|
|
hub.SendToAccountConversation(10, 5, []byte(`{"event":"test"}`))
|
|
|
|
select {
|
|
case msg := <-client.Send:
|
|
assert.NotEmpty(t, msg)
|
|
default:
|
|
t.Fatal("expected to receive message")
|
|
}
|
|
}
|
|
|
|
func TestHub_SendToRoom(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
client := NewClient(1, 10, nil, hub)
|
|
client.Identifier = `{"channel":"RoomChannel"}`
|
|
hub.Register(client)
|
|
hub.subscribeClient(client.ID, "custom_room")
|
|
client.SubscribedRooms["custom_room"] = true
|
|
|
|
hub.SendToRoom("custom_room", []byte(`{"event":"test"}`))
|
|
|
|
select {
|
|
case msg := <-client.Send:
|
|
assert.NotEmpty(t, msg)
|
|
default:
|
|
t.Fatal("expected to receive message")
|
|
}
|
|
}
|
|
|
|
func TestHub_SendToClient(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
client := NewClient(1, 10, nil, hub)
|
|
client.Identifier = `{"channel":"AccountChannel"}`
|
|
hub.Register(client)
|
|
|
|
hub.SendToClient(client.ID, []byte(`{"event":"test"}`))
|
|
|
|
select {
|
|
case msg := <-client.Send:
|
|
assert.NotEmpty(t, msg)
|
|
default:
|
|
t.Fatal("expected to receive message")
|
|
}
|
|
}
|
|
|
|
func TestHub_SendToClient_NotFound(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
hub.SendToClient("nonexistent", []byte(`{"event":"test"}`)) // should not panic
|
|
}
|
|
|
|
func TestWrapActionCableMessage_WithIdentifier(t *testing.T) {
|
|
identifier := `{"channel":"AccountChannel","account_id":1}`
|
|
data := []byte(`{"event":"test"}`)
|
|
result := wrapActionCableMessage(identifier, data)
|
|
assert.NotEmpty(t, result)
|
|
|
|
var decoded map[string]json.RawMessage
|
|
err := json.Unmarshal(result, &decoded)
|
|
require.NoError(t, err)
|
|
assert.Contains(t, decoded, "identifier")
|
|
assert.Contains(t, decoded, "message")
|
|
}
|
|
|
|
func TestWrapActionCableMessage_NoIdentifier(t *testing.T) {
|
|
result := wrapActionCableMessage("", []byte(`{"event":"test"}`))
|
|
assert.Nil(t, result)
|
|
}
|
|
|
|
func TestAccountRoomName(t *testing.T) {
|
|
assert.Equal(t, "account_1", accountRoomName(1))
|
|
assert.Equal(t, "account_42", accountRoomName(42))
|
|
}
|
|
|
|
func TestConversationRoomName(t *testing.T) {
|
|
assert.Equal(t, "account_1_conversation_5", conversationRoomName(1, 5))
|
|
assert.Equal(t, "account_10_conversation_99", conversationRoomName(10, 99))
|
|
}
|
|
|
|
func TestPubsubTokenRoomName(t *testing.T) {
|
|
assert.Equal(t, "pubsub_token_abc123", pubsubTokenRoomName("abc123"))
|
|
}
|
|
|
|
func TestHub_Run_Shutdown(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
|
|
done := make(chan struct{})
|
|
go func() {
|
|
hub.Run(ctx)
|
|
close(done)
|
|
}()
|
|
|
|
time.Sleep(50 * time.Millisecond)
|
|
cancel()
|
|
|
|
select {
|
|
case <-done:
|
|
case <-time.After(2 * time.Second):
|
|
t.Fatal("hub did not shut down")
|
|
}
|
|
}
|
|
|
|
func TestHub_Shutdown(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
// Do not register a client with nil Conn — shutdown() calls client.Conn.Close()
|
|
// which panics on nil. Just verify Shutdown completes without error.
|
|
hub.Shutdown(context.Background())
|
|
}
|
|
|
|
func TestHub_SubscribeClient_UnsubscribeClient(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
hub.subscribeClient("client1", "room1")
|
|
assert.True(t, hub.rooms["room1"]["client1"])
|
|
|
|
hub.unsubscribeClient("client1", "room1")
|
|
_, exists := hub.rooms["room1"]
|
|
assert.False(t, exists)
|
|
}
|
|
|
|
func TestHub_SubmitCommand(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
client := NewClient(1, 10, nil, hub)
|
|
hub.Register(client)
|
|
|
|
hub.SubmitCommand(client, wspkg.WSCommand{Command: "ping", Data: ""})
|
|
|
|
select {
|
|
case c := <-hub.commandChan:
|
|
assert.NotNil(t, c.Client)
|
|
assert.Equal(t, "ping", c.Cmd.Command)
|
|
default:
|
|
t.Fatal("expected command in channel")
|
|
}
|
|
}
|
|
|
|
// --- Handler Tests ---
|
|
|
|
func TestUintToStr(t *testing.T) {
|
|
assert.Equal(t, "1", uintToStr(1))
|
|
assert.Equal(t, "42", uintToStr(42))
|
|
assert.Equal(t, "0", uintToStr(0))
|
|
}
|
|
|
|
func TestHandleSubscribe_ContactRoomUsesAuthenticatedPubsubToken(t *testing.T) {
|
|
hub := NewHubSimple()
|
|
handler := NewHandler(hub, nil)
|
|
client := NewClient(9, 10, nil, hub)
|
|
client.IsContact = true
|
|
client.PubsubToken = "visitor-token"
|
|
|
|
identifier, err := json.Marshal(ChannelIdentifier{
|
|
Channel: ChannelRoom,
|
|
PubsubToken: "visitor-token",
|
|
})
|
|
require.NoError(t, err)
|
|
handler.handleSubscribe(client, CommandFrame{Command: CommandSubscribe, Identifier: string(identifier)})
|
|
|
|
assert.True(t, client.SubscribedRooms[pubsubTokenRoomName("visitor-token")])
|
|
var confirm ConfirmFrame
|
|
require.NoError(t, json.Unmarshal(<-client.Send, &confirm))
|
|
assert.Equal(t, ServerConfirmSubscribe, confirm.Type)
|
|
|
|
other := NewClient(9, 10, nil, hub)
|
|
other.IsContact = true
|
|
other.PubsubToken = "visitor-token"
|
|
mismatch, err := json.Marshal(ChannelIdentifier{Channel: ChannelRoom, PubsubToken: "other-token"})
|
|
require.NoError(t, err)
|
|
handler.handleSubscribe(other, CommandFrame{Command: CommandSubscribe, Identifier: string(mismatch)})
|
|
|
|
assert.False(t, other.SubscribedRooms[pubsubTokenRoomName("other-token")])
|
|
var reject RejectFrame
|
|
require.NoError(t, json.Unmarshal(<-other.Send, &reject))
|
|
assert.Equal(t, ServerRejectSubscribe, reject.Type)
|
|
}
|
|
|
|
// --- Integration: ServeWS with real WebSocket ---
|
|
|
|
// createTestHandler creates a Handler with a real WSAuthenticator using a test JWT config.
|
|
func createTestHandler(t *testing.T) *Handler {
|
|
t.Helper()
|
|
jwtCfg := &config.JWTConfig{
|
|
Secret: "test-secret-key-for-ws",
|
|
ExpiryHours: 1,
|
|
AccessExpiryMinutes: 60,
|
|
}
|
|
jwtSvc := auth.NewJWTService(jwtCfg)
|
|
authenticator := wspkg.NewWSAuthenticator(jwtSvc, nil)
|
|
|
|
hub := NewHubSimple()
|
|
return NewHandler(hub, authenticator)
|
|
}
|
|
|
|
func TestServeWS_NoToken_Unauthorized(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
h := createTestHandler(t)
|
|
|
|
router := gin.New()
|
|
router.GET("/ws", h.ServeWS)
|
|
|
|
w := httptest.NewRecorder()
|
|
req := httptest.NewRequest("GET", "/ws", nil)
|
|
router.ServeHTTP(w, req)
|
|
|
|
assert.Equal(t, http.StatusUnauthorized, w.Code)
|
|
}
|
|
|
|
func TestServeWS_InvalidToken_Unauthorized(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
h := createTestHandler(t)
|
|
|
|
router := gin.New()
|
|
router.GET("/ws", h.ServeWS)
|
|
|
|
w := httptest.NewRecorder()
|
|
req := httptest.NewRequest("GET", "/ws?token=invalid-token", nil)
|
|
router.ServeHTTP(w, req)
|
|
|
|
assert.Equal(t, http.StatusUnauthorized, w.Code)
|
|
}
|
|
|
|
func TestServeCable_DelegatesToServeWS(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
h := createTestHandler(t)
|
|
|
|
router := gin.New()
|
|
router.GET("/cable", h.ServeCable)
|
|
|
|
w := httptest.NewRecorder()
|
|
req := httptest.NewRequest("GET", "/cable", nil)
|
|
router.ServeHTTP(w, req)
|
|
|
|
assert.Equal(t, http.StatusUnauthorized, w.Code)
|
|
}
|
|
|
|
func TestServeWS_ValidToken_Success(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
jwtCfg := &config.JWTConfig{
|
|
Secret: "test-secret-key-for-ws",
|
|
ExpiryHours: 1,
|
|
AccessExpiryMinutes: 60,
|
|
}
|
|
jwtSvc := auth.NewJWTService(jwtCfg)
|
|
authenticator := wspkg.NewWSAuthenticator(jwtSvc, nil)
|
|
hub := NewHubSimple()
|
|
h := NewHandler(hub, authenticator)
|
|
|
|
// Generate a valid JWT token
|
|
tokenPair, err := jwtSvc.GenerateTokenPair(&model.User{Base: model.Base{ID: 1}, Provider: "local"}, 10, "agent")
|
|
require.NoError(t, err)
|
|
token := tokenPair.AccessToken
|
|
require.NoError(t, err)
|
|
require.NotEmpty(t, token)
|
|
|
|
router := gin.New()
|
|
router.GET("/ws", h.ServeWS)
|
|
|
|
srv := httptest.NewServer(router)
|
|
defer srv.Close()
|
|
|
|
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token
|
|
dialer := websocket.Dialer{HandshakeTimeout: 2 * time.Second}
|
|
conn, resp, err := dialer.Dial(wsURL, nil)
|
|
require.NoError(t, err)
|
|
defer conn.Close()
|
|
defer resp.Body.Close()
|
|
|
|
// Should receive welcome frame
|
|
_, msg, err := conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
var welcome WelcomeFrame
|
|
err = json.Unmarshal(msg, &welcome)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, ServerWelcome, welcome.Type)
|
|
|
|
// Send a ping command
|
|
pingCmd := CommandFrame{Command: CommandPing}
|
|
cmdData, _ := json.Marshal(pingCmd)
|
|
err = conn.WriteMessage(websocket.TextMessage, cmdData)
|
|
require.NoError(t, err)
|
|
|
|
// Should receive a ping response
|
|
_, msg, err = conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
var pingResp PingFrame
|
|
err = json.Unmarshal(msg, &pingResp)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, ServerPing, pingResp.Type)
|
|
}
|
|
|
|
func TestServeCable_WidgetReceivesTokenRoomEvent(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=private"), &gorm.Config{})
|
|
require.NoError(t, err)
|
|
require.NoError(t, db.AutoMigrate(&model.Account{}, &model.Contact{}, &model.Inbox{}, &model.ContactInbox{}))
|
|
account := model.Account{Name: "Widget account"}
|
|
require.NoError(t, db.Create(&account).Error)
|
|
contact := model.Contact{AccountID: account.ID, Name: "Visitor"}
|
|
require.NoError(t, db.Create(&contact).Error)
|
|
inbox := model.Inbox{AccountID: account.ID, Name: "Website", ChannelType: "Channel::WebWidget", Enabled: true}
|
|
require.NoError(t, db.Create(&inbox).Error)
|
|
contactInbox := model.ContactInbox{ContactID: contact.ID, InboxID: inbox.ID, PubsubToken: "visitor-token"}
|
|
require.NoError(t, db.Create(&contactInbox).Error)
|
|
|
|
hub := NewHubSimple()
|
|
authenticator := wspkg.NewWSAuthenticator(nil, repository.NewContactInboxRepo(db), db)
|
|
handler := NewHandler(hub, authenticator)
|
|
router := gin.New()
|
|
router.GET("/cable", handler.ServeCable)
|
|
server := httptest.NewServer(router)
|
|
t.Cleanup(server.Close)
|
|
|
|
conn, _, err := websocket.DefaultDialer.Dial(
|
|
"ws"+strings.TrimPrefix(server.URL, "http")+"/cable?pubsub_token=visitor-token",
|
|
nil,
|
|
)
|
|
require.NoError(t, err)
|
|
t.Cleanup(func() { _ = conn.Close() })
|
|
require.NoError(t, conn.SetReadDeadline(time.Now().Add(2*time.Second)))
|
|
_, _, err = conn.ReadMessage() // welcome
|
|
require.NoError(t, err)
|
|
|
|
identifier, err := json.Marshal(ChannelIdentifier{Channel: ChannelRoom, PubsubToken: "visitor-token"})
|
|
require.NoError(t, err)
|
|
command, err := json.Marshal(CommandFrame{Command: CommandSubscribe, Identifier: string(identifier)})
|
|
require.NoError(t, err)
|
|
require.NoError(t, conn.WriteMessage(websocket.TextMessage, command))
|
|
_, confirmation, err := conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
var confirm ConfirmFrame
|
|
require.NoError(t, json.Unmarshal(confirmation, &confirm))
|
|
assert.Equal(t, ServerConfirmSubscribe, confirm.Type)
|
|
|
|
hub.SendToRoom(pubsubTokenRoomName("visitor-token"), []byte(`{"event":"message.created","data":{"id":12,"content":"Dashboard reply","message_type":1,"conversation_id":42}}`))
|
|
_, message, err := conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
var delivered struct {
|
|
Identifier string `json:"identifier"`
|
|
Message json.RawMessage `json:"message"`
|
|
}
|
|
require.NoError(t, json.Unmarshal(message, &delivered))
|
|
assert.JSONEq(t, string(identifier), delivered.Identifier)
|
|
var event wspkg.WSMessage
|
|
require.NoError(t, json.Unmarshal(delivered.Message, &event))
|
|
assert.Equal(t, wspkg.EventMessageCreated, event.Event)
|
|
payload := event.Data.(map[string]interface{})
|
|
assert.Equal(t, "Dashboard reply", payload["content"])
|
|
assert.Equal(t, float64(1), payload["message_type"])
|
|
assert.Equal(t, float64(42), payload["conversation_id"])
|
|
}
|
|
|
|
func TestDashboardOutgoingReachesDashboardAndReconnectedWidget(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=private"), &gorm.Config{})
|
|
require.NoError(t, err)
|
|
require.NoError(t, db.AutoMigrate(
|
|
&model.Account{}, &model.User{}, &model.Inbox{}, &model.Contact{}, &model.ContactInbox{},
|
|
&model.Conversation{}, &model.Message{}, &model.Attachment{},
|
|
))
|
|
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"}
|
|
require.NoError(t, db.Create(contact).Error)
|
|
contactInbox := &model.ContactInbox{ContactID: contact.ID, InboxID: inbox.ID, PubsubToken: "visitor-token"}
|
|
require.NoError(t, db.Create(contactInbox).Error)
|
|
displayID := uint(42)
|
|
conversation := &model.Conversation{
|
|
AccountID: account.ID, InboxID: inbox.ID, ContactID: contact.ID, ContactInboxID: &contactInbox.ID,
|
|
DisplayID: &displayID, Status: "open", ChannelType: "web_widget", Channel: "web_widget",
|
|
}
|
|
require.NoError(t, db.Create(conversation).Error)
|
|
|
|
hub := NewHubSimple()
|
|
dispatcher := channel.NewDispatcher()
|
|
dispatcher.Register(wsevent.New(wspkg.NewEventPublisherLocal(hub, nil)))
|
|
messageService := service.NewMessageService(repository.NewMessageRepo(db), dispatcher, nil)
|
|
messageHandler := v1.NewMessageHandler(messageService)
|
|
widgetService := service.NewWidgetService(
|
|
repository.NewInboxRepo(db), repository.NewContactRepo(db), repository.NewContactInboxRepo(db),
|
|
repository.NewConversationRepo(db), repository.NewMessageRepo(db), nil, nil, nil, nil, nil, nil, nil, nil,
|
|
)
|
|
widgetHandler := widgethandler.NewHandler(widgetService)
|
|
jwtService := auth.NewJWTService(&config.JWTConfig{Secret: "dashboard-widget-chain", ExpiryHours: 1, AccessExpiryMinutes: 60})
|
|
wsHandler := NewHandler(hub, wspkg.NewWSAuthenticator(jwtService, repository.NewContactInboxRepo(db)))
|
|
|
|
router := gin.New()
|
|
router.Use(func(c *gin.Context) {
|
|
c.Set("user_id", uint(7))
|
|
c.Next()
|
|
})
|
|
router.GET("/cable", wsHandler.ServeCable)
|
|
router.POST("/api/v1/accounts/:account_id/conversations/:conversation_id/messages", messageHandler.Create)
|
|
router.GET("/api/v1/widget/messages", widgetHandler.GetLatestMessages)
|
|
server := httptest.NewServer(router)
|
|
t.Cleanup(server.Close)
|
|
|
|
tokenPair, err := jwtService.GenerateTokenPair(user, account.ID, "agent")
|
|
require.NoError(t, err)
|
|
dashboard := dialCable(t, server.URL, "?token="+tokenPair.AccessToken)
|
|
t.Cleanup(func() { _ = dashboard.Close() })
|
|
dashboardIdentifier := subscribeCable(t, dashboard, ChannelIdentifier{Channel: ChannelRoom, AccountID: account.ID})
|
|
|
|
visitor := dialCable(t, server.URL, "?pubsub_token="+contactInbox.PubsubToken)
|
|
mismatch, err := json.Marshal(ChannelIdentifier{Channel: ChannelRoom, PubsubToken: "other-token"})
|
|
require.NoError(t, err)
|
|
require.NoError(t, visitor.WriteJSON(CommandFrame{Command: CommandSubscribe, Identifier: string(mismatch)}))
|
|
var rejected RejectFrame
|
|
require.NoError(t, visitor.ReadJSON(&rejected))
|
|
assert.Equal(t, ServerRejectSubscribe, rejected.Type)
|
|
visitorIdentifier := subscribeCable(t, visitor, ChannelIdentifier{Channel: ChannelRoom, PubsubToken: contactInbox.PubsubToken})
|
|
|
|
createDashboardMessage(t, server.URL, account.ID, displayID, "dashboard reply one")
|
|
assertCableMessage(t, dashboard, dashboardIdentifier, "dashboard reply one", displayID)
|
|
assertCableMessage(t, visitor, visitorIdentifier, "dashboard reply one", displayID)
|
|
|
|
refreshRequest, err := http.NewRequest(http.MethodGet, server.URL+"/api/v1/widget/messages", nil)
|
|
require.NoError(t, err)
|
|
refreshRequest.Header.Set("X-Auth-Token", contactInbox.PubsubToken)
|
|
refreshResponse, err := http.DefaultClient.Do(refreshRequest)
|
|
require.NoError(t, err)
|
|
defer refreshResponse.Body.Close()
|
|
require.Equal(t, http.StatusOK, refreshResponse.StatusCode)
|
|
var refresh struct {
|
|
Payload []struct {
|
|
Content string `json:"content"`
|
|
ConversationID uint `json:"conversation_id"`
|
|
} `json:"payload"`
|
|
}
|
|
require.NoError(t, json.NewDecoder(refreshResponse.Body).Decode(&refresh))
|
|
require.Len(t, refresh.Payload, 1)
|
|
assert.Equal(t, "dashboard reply one", refresh.Payload[0].Content)
|
|
assert.Equal(t, displayID, refresh.Payload[0].ConversationID)
|
|
|
|
require.NoError(t, visitor.Close())
|
|
visitor = dialCable(t, server.URL, "?pubsub_token="+contactInbox.PubsubToken)
|
|
t.Cleanup(func() { _ = visitor.Close() })
|
|
visitorIdentifier = subscribeCable(t, visitor, ChannelIdentifier{Channel: ChannelRoom, PubsubToken: contactInbox.PubsubToken})
|
|
createDashboardMessage(t, server.URL, account.ID, displayID, "dashboard reply after reconnect")
|
|
assertCableMessage(t, dashboard, dashboardIdentifier, "dashboard reply after reconnect", displayID)
|
|
assertCableMessage(t, visitor, visitorIdentifier, "dashboard reply after reconnect", displayID)
|
|
}
|
|
|
|
func dialCable(t *testing.T, serverURL, query string) *websocket.Conn {
|
|
t.Helper()
|
|
conn, _, err := websocket.DefaultDialer.Dial("ws"+strings.TrimPrefix(serverURL, "http")+"/cable"+query, nil)
|
|
require.NoError(t, err)
|
|
require.NoError(t, conn.SetReadDeadline(time.Now().Add(2*time.Second)))
|
|
var welcome WelcomeFrame
|
|
require.NoError(t, conn.ReadJSON(&welcome))
|
|
require.Equal(t, ServerWelcome, welcome.Type)
|
|
return conn
|
|
}
|
|
|
|
func subscribeCable(t *testing.T, conn *websocket.Conn, identifier ChannelIdentifier) string {
|
|
t.Helper()
|
|
raw, err := json.Marshal(identifier)
|
|
require.NoError(t, err)
|
|
require.NoError(t, conn.WriteJSON(CommandFrame{Command: CommandSubscribe, Identifier: string(raw)}))
|
|
var confirmed ConfirmFrame
|
|
require.NoError(t, conn.ReadJSON(&confirmed))
|
|
require.Equal(t, ServerConfirmSubscribe, confirmed.Type)
|
|
return string(raw)
|
|
}
|
|
|
|
func createDashboardMessage(t *testing.T, serverURL string, accountID, displayID uint, content string) {
|
|
t.Helper()
|
|
body, err := json.Marshal(map[string]any{"content": content, "message_type": "outgoing"})
|
|
require.NoError(t, err)
|
|
request, err := http.NewRequest(http.MethodPost,
|
|
fmt.Sprintf("%s/api/v1/accounts/%d/conversations/%d/messages", serverURL, accountID, displayID), bytes.NewReader(body))
|
|
require.NoError(t, err)
|
|
request.Header.Set("Content-Type", "application/json")
|
|
response, err := http.DefaultClient.Do(request)
|
|
require.NoError(t, err)
|
|
defer response.Body.Close()
|
|
require.Equal(t, http.StatusOK, response.StatusCode)
|
|
}
|
|
|
|
func assertCableMessage(t *testing.T, conn *websocket.Conn, identifier, content string, displayID uint) {
|
|
t.Helper()
|
|
var delivered struct {
|
|
Identifier string `json:"identifier"`
|
|
Message json.RawMessage `json:"message"`
|
|
}
|
|
require.NoError(t, conn.ReadJSON(&delivered))
|
|
assert.JSONEq(t, identifier, delivered.Identifier)
|
|
var event wspkg.WSMessage
|
|
require.NoError(t, json.Unmarshal(delivered.Message, &event))
|
|
require.Equal(t, wspkg.EventMessageCreated, event.Event)
|
|
payload := event.Data.(map[string]interface{})
|
|
assert.Equal(t, content, payload["content"])
|
|
assert.Equal(t, float64(1), payload["message_type"])
|
|
assert.Equal(t, float64(displayID), payload["conversation_id"])
|
|
}
|
|
|
|
func TestRemoteSocketRechecksAccessWhenDisconnectPublishFails(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=private"), &gorm.Config{})
|
|
require.NoError(t, err)
|
|
require.NoError(t, db.AutoMigrate(&model.User{}, &model.UserSession{}))
|
|
user := &model.User{Name: "Agent", Email: "remote-agent@example.com", Provider: "email", Active: true}
|
|
require.NoError(t, db.Create(user).Error)
|
|
require.NoError(t, db.Create(&model.UserSession{UserID: user.ID, ClientID: "browser"}).Error)
|
|
|
|
jwtSvc := auth.NewJWTService(&config.JWTConfig{Secret: "remote-ws-test", ExpiryHours: 1, RefreshExpiryHours: 24})
|
|
pair, err := jwtSvc.GenerateTokenPairForClient(user, 1, "agent", "browser")
|
|
require.NoError(t, err)
|
|
remoteHub := NewHubSimple()
|
|
remoteHandler := NewHandler(remoteHub, wspkg.NewWSAuthenticator(jwtSvc, nil, db))
|
|
router := gin.New()
|
|
router.GET("/ws", remoteHandler.ServeWS)
|
|
server := httptest.NewServer(router)
|
|
t.Cleanup(server.Close)
|
|
conn, _, err := websocket.DefaultDialer.Dial("ws"+strings.TrimPrefix(server.URL, "http")+"/ws?token="+pair.AccessToken, nil)
|
|
require.NoError(t, err)
|
|
t.Cleanup(func() { _ = conn.Close() })
|
|
_, _, err = conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
require.NoError(t, db.Model(user).Update("active", false).Error)
|
|
require.NoError(t, db.Where("user_id = ?", user.ID).Delete(&model.UserSession{}).Error)
|
|
mr := miniredis.RunT(t)
|
|
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
|
|
t.Cleanup(func() { _ = rdb.Close() })
|
|
mr.Close()
|
|
require.Error(t, wspkg.NewBroadcastRelay(rdb, NewHubSimple()).PublishUserDisconnect(context.Background(), user.ID))
|
|
|
|
command, err := json.Marshal(CommandFrame{Command: CommandPing})
|
|
require.NoError(t, err)
|
|
require.NoError(t, conn.WriteMessage(websocket.TextMessage, command))
|
|
require.NoError(t, conn.SetReadDeadline(time.Now().Add(time.Second)))
|
|
_, _, err = conn.ReadMessage()
|
|
require.Error(t, err)
|
|
require.Eventually(t, func() bool {
|
|
remoteHub.mu.RLock()
|
|
defer remoteHub.mu.RUnlock()
|
|
return len(remoteHub.clients) == 0
|
|
}, time.Second, 10*time.Millisecond)
|
|
}
|
|
|
|
func TestServeWS_SubscribeAccount(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
jwtCfg := &config.JWTConfig{
|
|
Secret: "test-secret-key-for-ws",
|
|
ExpiryHours: 1,
|
|
AccessExpiryMinutes: 60,
|
|
}
|
|
jwtSvc := auth.NewJWTService(jwtCfg)
|
|
authenticator := wspkg.NewWSAuthenticator(jwtSvc, nil)
|
|
hub := NewHubSimple()
|
|
h := NewHandler(hub, authenticator)
|
|
|
|
tokenPair, err := jwtSvc.GenerateTokenPair(&model.User{Base: model.Base{ID: 1}, Provider: "local"}, 10, "agent")
|
|
require.NoError(t, err)
|
|
token := tokenPair.AccessToken
|
|
require.NoError(t, err)
|
|
|
|
router := gin.New()
|
|
router.GET("/ws", h.ServeWS)
|
|
|
|
srv := httptest.NewServer(router)
|
|
defer srv.Close()
|
|
|
|
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token
|
|
dialer := websocket.Dialer{HandshakeTimeout: 2 * time.Second}
|
|
conn, _, err := dialer.Dial(wsURL, nil)
|
|
require.NoError(t, err)
|
|
defer conn.Close()
|
|
|
|
// Read welcome
|
|
_, _, err = conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
// Send subscribe command
|
|
identifier := `{"channel":"AccountChannel","account_id":10}`
|
|
subCmd := CommandFrame{Command: CommandSubscribe, Identifier: identifier}
|
|
cmdData, _ := json.Marshal(subCmd)
|
|
err = conn.WriteMessage(websocket.TextMessage, cmdData)
|
|
require.NoError(t, err)
|
|
|
|
// Should receive confirm frame
|
|
_, msg, err := conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
var confirm ConfirmFrame
|
|
err = json.Unmarshal(msg, &confirm)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, ServerConfirmSubscribe, confirm.Type)
|
|
assert.Equal(t, identifier, confirm.Identifier)
|
|
}
|
|
|
|
func TestServeWS_SubscribeAccountMismatch(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
jwtCfg := &config.JWTConfig{
|
|
Secret: "test-secret-key-for-ws",
|
|
ExpiryHours: 1,
|
|
AccessExpiryMinutes: 60,
|
|
}
|
|
jwtSvc := auth.NewJWTService(jwtCfg)
|
|
authenticator := wspkg.NewWSAuthenticator(jwtSvc, nil)
|
|
hub := NewHubSimple()
|
|
h := NewHandler(hub, authenticator)
|
|
|
|
tokenPair, err := jwtSvc.GenerateTokenPair(&model.User{Base: model.Base{ID: 1}, Provider: "local"}, 10, "agent")
|
|
require.NoError(t, err)
|
|
token := tokenPair.AccessToken
|
|
require.NoError(t, err)
|
|
|
|
router := gin.New()
|
|
router.GET("/ws", h.ServeWS)
|
|
|
|
srv := httptest.NewServer(router)
|
|
defer srv.Close()
|
|
|
|
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token
|
|
dialer := websocket.Dialer{HandshakeTimeout: 2 * time.Second}
|
|
conn, _, err := dialer.Dial(wsURL, nil)
|
|
require.NoError(t, err)
|
|
defer conn.Close()
|
|
|
|
// Read welcome
|
|
_, _, err = conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
// Send subscribe with wrong account_id
|
|
identifier := `{"channel":"AccountChannel","account_id":999}`
|
|
subCmd := CommandFrame{Command: CommandSubscribe, Identifier: identifier}
|
|
cmdData, _ := json.Marshal(subCmd)
|
|
err = conn.WriteMessage(websocket.TextMessage, cmdData)
|
|
require.NoError(t, err)
|
|
|
|
_, msg, err := conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
var reject RejectFrame
|
|
err = json.Unmarshal(msg, &reject)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, ServerRejectSubscribe, reject.Type)
|
|
assert.Contains(t, reject.Reason, "account_id mismatch")
|
|
}
|
|
|
|
func TestServeWS_SubscribeConversation(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
jwtCfg := &config.JWTConfig{
|
|
Secret: "test-secret-key-for-ws",
|
|
ExpiryHours: 1,
|
|
AccessExpiryMinutes: 60,
|
|
}
|
|
jwtSvc := auth.NewJWTService(jwtCfg)
|
|
authenticator := wspkg.NewWSAuthenticator(jwtSvc, nil)
|
|
hub := NewHubSimple()
|
|
h := NewHandler(hub, authenticator)
|
|
|
|
tokenPair, err := jwtSvc.GenerateTokenPair(&model.User{Base: model.Base{ID: 1}, Provider: "local"}, 10, "agent")
|
|
require.NoError(t, err)
|
|
token := tokenPair.AccessToken
|
|
require.NoError(t, err)
|
|
|
|
router := gin.New()
|
|
router.GET("/ws", h.ServeWS)
|
|
|
|
srv := httptest.NewServer(router)
|
|
defer srv.Close()
|
|
|
|
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token
|
|
dialer := websocket.Dialer{HandshakeTimeout: 2 * time.Second}
|
|
conn, _, err := dialer.Dial(wsURL, nil)
|
|
require.NoError(t, err)
|
|
defer conn.Close()
|
|
|
|
// Read welcome
|
|
_, _, err = conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
// Subscribe to conversation
|
|
identifier := `{"channel":"ConversationChannel","account_id":10,"conversation_id":5}`
|
|
subCmd := CommandFrame{Command: CommandSubscribe, Identifier: identifier}
|
|
cmdData, _ := json.Marshal(subCmd)
|
|
err = conn.WriteMessage(websocket.TextMessage, cmdData)
|
|
require.NoError(t, err)
|
|
|
|
_, msg, err := conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
var confirm ConfirmFrame
|
|
err = json.Unmarshal(msg, &confirm)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, ServerConfirmSubscribe, confirm.Type)
|
|
|
|
// Unsubscribe
|
|
unsubCmd := CommandFrame{Command: CommandUnsubscribe, Identifier: identifier}
|
|
unsubData, _ := json.Marshal(unsubCmd)
|
|
err = conn.WriteMessage(websocket.TextMessage, unsubData)
|
|
require.NoError(t, err)
|
|
|
|
_, msg, err = conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
var unsubConfirm ConfirmFrame
|
|
err = json.Unmarshal(msg, &unsubConfirm)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, ServerConfirmUnsubscribe, unsubConfirm.Type)
|
|
}
|
|
|
|
func TestServeWS_SubscribeConversationNoID(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
jwtCfg := &config.JWTConfig{
|
|
Secret: "test-secret-key-for-ws",
|
|
ExpiryHours: 1,
|
|
AccessExpiryMinutes: 60,
|
|
}
|
|
jwtSvc := auth.NewJWTService(jwtCfg)
|
|
authenticator := wspkg.NewWSAuthenticator(jwtSvc, nil)
|
|
hub := NewHubSimple()
|
|
h := NewHandler(hub, authenticator)
|
|
|
|
tokenPair, err := jwtSvc.GenerateTokenPair(&model.User{Base: model.Base{ID: 1}, Provider: "local"}, 10, "agent")
|
|
require.NoError(t, err)
|
|
token := tokenPair.AccessToken
|
|
require.NoError(t, err)
|
|
|
|
router := gin.New()
|
|
router.GET("/ws", h.ServeWS)
|
|
|
|
srv := httptest.NewServer(router)
|
|
defer srv.Close()
|
|
|
|
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token
|
|
dialer := websocket.Dialer{HandshakeTimeout: 2 * time.Second}
|
|
conn, _, err := dialer.Dial(wsURL, nil)
|
|
require.NoError(t, err)
|
|
defer conn.Close()
|
|
|
|
// Read welcome
|
|
_, _, err = conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
// Subscribe to conversation without conversation_id
|
|
identifier := `{"channel":"ConversationChannel","account_id":10}`
|
|
subCmd := CommandFrame{Command: CommandSubscribe, Identifier: identifier}
|
|
cmdData, _ := json.Marshal(subCmd)
|
|
err = conn.WriteMessage(websocket.TextMessage, cmdData)
|
|
require.NoError(t, err)
|
|
|
|
_, msg, err := conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
var reject RejectFrame
|
|
err = json.Unmarshal(msg, &reject)
|
|
require.NoError(t, err)
|
|
assert.Contains(t, reject.Reason, "conversation_id required")
|
|
}
|
|
|
|
func TestServeWS_SubscribeUnknownChannel(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
jwtCfg := &config.JWTConfig{
|
|
Secret: "test-secret-key-for-ws",
|
|
ExpiryHours: 1,
|
|
AccessExpiryMinutes: 60,
|
|
}
|
|
jwtSvc := auth.NewJWTService(jwtCfg)
|
|
authenticator := wspkg.NewWSAuthenticator(jwtSvc, nil)
|
|
hub := NewHubSimple()
|
|
h := NewHandler(hub, authenticator)
|
|
|
|
tokenPair, err := jwtSvc.GenerateTokenPair(&model.User{Base: model.Base{ID: 1}, Provider: "local"}, 10, "agent")
|
|
require.NoError(t, err)
|
|
token := tokenPair.AccessToken
|
|
require.NoError(t, err)
|
|
|
|
router := gin.New()
|
|
router.GET("/ws", h.ServeWS)
|
|
|
|
srv := httptest.NewServer(router)
|
|
defer srv.Close()
|
|
|
|
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token
|
|
dialer := websocket.Dialer{HandshakeTimeout: 2 * time.Second}
|
|
conn, _, err := dialer.Dial(wsURL, nil)
|
|
require.NoError(t, err)
|
|
defer conn.Close()
|
|
|
|
// Read welcome
|
|
_, _, err = conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
// Subscribe to unknown channel
|
|
identifier := `{"channel":"UnknownChannel","account_id":10}`
|
|
subCmd := CommandFrame{Command: CommandSubscribe, Identifier: identifier}
|
|
cmdData, _ := json.Marshal(subCmd)
|
|
err = conn.WriteMessage(websocket.TextMessage, cmdData)
|
|
require.NoError(t, err)
|
|
|
|
_, msg, err := conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
var reject RejectFrame
|
|
err = json.Unmarshal(msg, &reject)
|
|
require.NoError(t, err)
|
|
assert.Contains(t, reject.Reason, "unknown channel type")
|
|
}
|
|
|
|
func TestServeWS_SubscribeInvalidIdentifier(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
jwtCfg := &config.JWTConfig{
|
|
Secret: "test-secret-key-for-ws",
|
|
ExpiryHours: 1,
|
|
AccessExpiryMinutes: 60,
|
|
}
|
|
jwtSvc := auth.NewJWTService(jwtCfg)
|
|
authenticator := wspkg.NewWSAuthenticator(jwtSvc, nil)
|
|
hub := NewHubSimple()
|
|
h := NewHandler(hub, authenticator)
|
|
|
|
tokenPair, err := jwtSvc.GenerateTokenPair(&model.User{Base: model.Base{ID: 1}, Provider: "local"}, 10, "agent")
|
|
require.NoError(t, err)
|
|
token := tokenPair.AccessToken
|
|
require.NoError(t, err)
|
|
|
|
router := gin.New()
|
|
router.GET("/ws", h.ServeWS)
|
|
|
|
srv := httptest.NewServer(router)
|
|
defer srv.Close()
|
|
|
|
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token
|
|
dialer := websocket.Dialer{HandshakeTimeout: 2 * time.Second}
|
|
conn, _, err := dialer.Dial(wsURL, nil)
|
|
require.NoError(t, err)
|
|
defer conn.Close()
|
|
|
|
// Read welcome
|
|
_, _, err = conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
// Send invalid JSON identifier
|
|
subCmd := CommandFrame{Command: CommandSubscribe, Identifier: "invalid-json"}
|
|
cmdData, _ := json.Marshal(subCmd)
|
|
err = conn.WriteMessage(websocket.TextMessage, cmdData)
|
|
require.NoError(t, err)
|
|
|
|
_, msg, err := conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
var reject RejectFrame
|
|
err = json.Unmarshal(msg, &reject)
|
|
require.NoError(t, err)
|
|
assert.Contains(t, reject.Reason, "invalid identifier")
|
|
}
|
|
|
|
func TestServeWS_MessageCommand(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
|
|
jwtCfg := &config.JWTConfig{
|
|
Secret: "test-secret-key-for-ws",
|
|
ExpiryHours: 1,
|
|
AccessExpiryMinutes: 60,
|
|
}
|
|
jwtSvc := auth.NewJWTService(jwtCfg)
|
|
authenticator := wspkg.NewWSAuthenticator(jwtSvc, nil)
|
|
hub := NewHubSimple()
|
|
h := NewHandler(hub, authenticator)
|
|
|
|
tokenPair, err := jwtSvc.GenerateTokenPair(&model.User{Base: model.Base{ID: 1}, Provider: "local"}, 10, "agent")
|
|
require.NoError(t, err)
|
|
token := tokenPair.AccessToken
|
|
require.NoError(t, err)
|
|
|
|
router := gin.New()
|
|
router.GET("/ws", h.ServeWS)
|
|
|
|
srv := httptest.NewServer(router)
|
|
defer srv.Close()
|
|
|
|
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token
|
|
dialer := websocket.Dialer{HandshakeTimeout: 2 * time.Second}
|
|
conn, _, err := dialer.Dial(wsURL, nil)
|
|
require.NoError(t, err)
|
|
defer conn.Close()
|
|
|
|
// Read welcome
|
|
_, _, err = conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
// Send message command (should be acknowledged but no specific response)
|
|
msgCmd := CommandFrame{Command: CommandMessage, Data: `{"action":"update_presence"}`}
|
|
cmdData, _ := json.Marshal(msgCmd)
|
|
err = conn.WriteMessage(websocket.TextMessage, cmdData)
|
|
require.NoError(t, err)
|
|
|
|
// Send ping — should get a response
|
|
pingCmd := CommandFrame{Command: CommandPing}
|
|
pingData, _ := json.Marshal(pingCmd)
|
|
err = conn.WriteMessage(websocket.TextMessage, pingData)
|
|
require.NoError(t, err)
|
|
|
|
_, msg, err := conn.ReadMessage()
|
|
require.NoError(t, err)
|
|
|
|
var pingResp PingFrame
|
|
err = json.Unmarshal(msg, &pingResp)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, ServerPing, pingResp.Type)
|
|
}
|