Files
gochat/backend/internal/handler/ws/ws_test.go
T
Rogeeandrogee 18ecee3e46 H-116: close visitor payload trust boundaries (#19)
* H-116: close visitor payload trust boundaries

* H-129: unblock SQLite backend tests

* H-129: remove stale last-seen response assertions

---------

Co-authored-by: Rogee <rogee@ipao.vip>
2026-08-15 01:15:04 +08:00

949 lines
26 KiB
Go

package ws
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/gochat/gochat/internal/auth"
"github.com/gochat/gochat/internal/config"
"github.com/gochat/gochat/internal/model"
wspkg "github.com/gochat/gochat/internal/ws"
)
// --- 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 TestPingFrame_JSONRoundTrip(t *testing.T) {
pf := PingFrame{
Type: ServerPing,
Message: "2026-01-01T00:00:00Z",
}
data, err := json.Marshal(pf)
require.NoError(t, err)
var decoded PingFrame
err = json.Unmarshal(data, &decoded)
require.NoError(t, err)
assert.Equal(t, pf.Message, decoded.Message)
}
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)
}
func TestPingInterval_Constant(t *testing.T) {
assert.Equal(t, 5, PingInterval)
}
// --- 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 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.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))
}
// --- 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 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
conn.ReadMessage()
// 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
conn.ReadMessage()
// 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
conn.ReadMessage()
// 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
conn.ReadMessage()
// 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
conn.ReadMessage()
// 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
conn.ReadMessage()
// 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
conn.ReadMessage()
// 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)
}