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