Files
gochat/backend/internal/handler/ws/ws_test.go
T
Rogeeandrogee f36606a4f2 HH-564: harden durable realtime publish boundaries (#139)
* fix(HH-564): harden durable realtime enqueue

* fix(HH-564): wire production SSE stream

---------

Co-authored-by: Rogee <rogee@ipao.vip>
2026-08-23 23:38:00 +08:00

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)
}