package ws import ( "bytes" "context" "encoding/json" "errors" "fmt" "net/http" "net/http/httptest" "os" "path/filepath" "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","message_type":"outgoing","sender":{"available_name":"Captain","name":"Captain","agent_name":"Captain"},"sender_type":"AgentBot","sender_id":7,"ai_takeover_active":true,"additional_attributes":{"agent_name":"Captain","sender_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 TestEventDataForClient_VisitorMessageKeepsPublicFieldsOnly(t *testing.T) { hub := NewHubSimple() visitor := NewClient(2, 10, nil, hub) visitor.IsContact = true data := eventDataForClient(visitor, []byte(`{"event":"message.created","data":{"id":12,"content":"Dashboard reply","message_type":1,"conversation_id":42,"sender":{"available_name":"Captain","name":"Captain","agent_name":"Captain"},"sender_type":"AgentBot","sender_id":7,"ai_takeover_active":true,"additional_attributes":{"sender_name":"Captain","agent_name":"Captain"}}}`)) fixture, err := os.ReadFile(filepath.Join("..", "..", "..", "..", "frontend", "testdata", "visitor_ws_payload.json")) require.NoError(t, err) assert.JSONEq(t, string(fixture), string(data)) var payload map[string]any require.NoError(t, json.Unmarshal(data, &payload)) assert.Equal(t, "message.created", payload["event"]) message, ok := payload["data"].(map[string]any) require.True(t, ok) assert.Equal(t, "Dashboard reply", message["content"]) assert.Equal(t, float64(1), message["message_type"]) assert.Equal(t, float64(42), message["conversation_id"]) assert.NotContains(t, message, "sender") assert.NotContains(t, message, "sender_type") assert.NotContains(t, message, "sender_id") assert.NotContains(t, message, "ai_takeover_active") assert.Equal(t, map[string]any{}, message["additional_attributes"]) } 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,"sender":{"available_name":"Captain","name":"Captain"},"sender_type":"AgentBot","sender_id":7,"ai_takeover_active":true,"additional_attributes":{"sender_name":"Captain","agent_name":"Captain"}}}`)) _, 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"]) assert.NotContains(t, payload, "sender") assert.NotContains(t, payload, "sender_type") assert.NotContains(t, payload, "sender_id") assert.NotContains(t, payload, "ai_takeover_active") assert.Equal(t, map[string]interface{}{}, payload["additional_attributes"]) } func TestServeCable_UnknownPubsubTokenReturnsUnauthorized(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{})) handler := NewHandler(NewHubSimple(), wspkg.NewWSAuthenticator(nil, repository.NewContactInboxRepo(db), db)) router := gin.New() router.GET("/cable", handler.ServeCable) server := httptest.NewServer(router) t.Cleanup(server.Close) response, err := http.Get(server.URL + "/cable?pubsub_token=stale-token") require.NoError(t, err) t.Cleanup(func() { _ = response.Body.Close() }) require.Equal(t, http.StatusUnauthorized, response.StatusCode) } 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) }