* 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>
949 lines
26 KiB
Go
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)
|
|
}
|