Files
2026-06-02 23:14:41 +08:00

674 lines
18 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package wecom
import (
"context"
"crypto/md5"
"encoding/base64"
"encoding/hex"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"sync"
"testing"
"time"
"github.com/chenhg5/cc-connect/core"
"github.com/gorilla/websocket"
)
// ---------------------------------------------------------------------------
// splitByBytes
// ---------------------------------------------------------------------------
func TestSplitByBytes_ShortString(t *testing.T) {
parts := splitByBytes("hello", 100)
if len(parts) != 1 || parts[0] != "hello" {
t.Fatalf("expected single chunk, got %v", parts)
}
}
func TestSplitByBytes_ExactBoundary(t *testing.T) {
s := "abcdef"
parts := splitByBytes(s, 6)
if len(parts) != 1 || parts[0] != s {
t.Fatalf("expected single chunk at exact boundary, got %v", parts)
}
}
func TestSplitByBytes_SplitASCII(t *testing.T) {
s := "abcdef"
parts := splitByBytes(s, 4)
if len(parts) != 2 {
t.Fatalf("expected 2 chunks, got %d: %v", len(parts), parts)
}
if parts[0] != "abcd" || parts[1] != "ef" {
t.Fatalf("unexpected chunks: %v", parts)
}
}
func TestSplitByBytes_UTF8NeverSplitsMidRune(t *testing.T) {
// "你好世界" = 4 runes × 3 bytes = 12 bytes
s := "你好世界"
parts := splitByBytes(s, 5) // 5 < 6, so only one 3-byte rune fits? Actually 3 fits, 4 doesn't → first chunk = "你" (3 bytes)
// With maxBytes=5: first iteration end=5, s[5] is a continuation byte → back off to 3 → "你", next end=5 but only 9 left, s[5] continuation → 6 → "好世" wait...
// Let's just verify no chunk contains a partial rune.
reassembled := ""
for _, p := range parts {
reassembled += p
// Each chunk must be valid UTF-8 (no partial rune)
for i := 0; i < len(p); i++ {
if p[i]>>6 == 0b10 && (i == 0 || p[i-1] < 0x80) {
t.Fatalf("chunk contains orphaned continuation byte: %q", p)
}
}
}
if reassembled != s {
t.Fatalf("reassembled %q != original %q", reassembled, s)
}
}
func TestSplitByBytes_EmptyString(t *testing.T) {
parts := splitByBytes("", 100)
if len(parts) != 1 || parts[0] != "" {
t.Fatalf("expected single empty chunk, got %v", parts)
}
}
func TestSplitByBytes_ReassemblesLargeContent(t *testing.T) {
var s string
for i := 0; i < 500; i++ {
s += fmt.Sprintf("line %d: 这是一段中文\n", i)
}
parts := splitByBytes(s, 2000)
reassembled := ""
for _, p := range parts {
if len(p) > 2000 {
t.Fatalf("chunk exceeds maxBytes: %d", len(p))
}
reassembled += p
}
if reassembled != s {
t.Fatalf("reassembled content does not match original (len %d vs %d)", len(reassembled), len(s))
}
}
// ---------------------------------------------------------------------------
// handleMsgCallback — chatID fallback to userID for single chats
// ---------------------------------------------------------------------------
func newCapturedWSPlatform() (*WSPlatform, <-chan *core.Message) {
p := &WSPlatform{allowFrom: "*"}
captured := make(chan *core.Message, 1)
p.handler = func(_ core.Platform, msg *core.Message) {
captured <- msg
}
return p, captured
}
func wsCallbackFrame(t *testing.T, reqID string, body wsMsgCallbackBody) wsFrame {
t.Helper()
bodyBytes, err := json.Marshal(body)
if err != nil {
t.Fatalf("marshal callback body: %v", err)
}
return wsFrame{
Cmd: "aibot_msg_callback",
Headers: wsFrameHeaders{ReqID: reqID},
Body: bodyBytes,
}
}
func TestHandleMsgCallback_SingleChat_ChatIDFallback(t *testing.T) {
p, captured := newCapturedWSPlatform()
body := wsMsgCallbackBody{
MsgID: "msg_001",
ChatID: "", // single chat: no chatID from server
ChatType: "single",
MsgType: "text",
}
body.From.UserID = "zhangsan"
body.Text.Content = "hello"
body.CreateTime = time.Now().Unix()
p.handleMsgCallback(wsCallbackFrame(t, "req_123", body))
select {
case msg := <-captured:
if msg.SessionKey != "wecom:zhangsan:zhangsan" {
t.Fatalf("expected sessionKey 'wecom:zhangsan:zhangsan', got %q", msg.SessionKey)
}
rc := msg.ReplyCtx.(wsReplyContext)
if rc.chatID != "zhangsan" {
t.Fatalf("expected chatID to fall back to userID 'zhangsan', got %q", rc.chatID)
}
case <-time.After(1 * time.Second):
t.Fatal("handler not called")
}
}
func TestHandleMsgCallback_GroupChat_ChatIDPreserved(t *testing.T) {
p, captured := newCapturedWSPlatform()
body := wsMsgCallbackBody{
MsgID: "msg_002",
ChatID: "group_chat_id_123",
ChatType: "group",
MsgType: "text",
}
body.From.UserID = "zhangsan"
body.Text.Content = "hi group"
body.CreateTime = time.Now().Unix()
p.handleMsgCallback(wsCallbackFrame(t, "req_456", body))
select {
case msg := <-captured:
if msg.SessionKey != "wecom:group_chat_id_123:zhangsan" {
t.Fatalf("expected sessionKey 'wecom:group_chat_id_123:zhangsan', got %q", msg.SessionKey)
}
rc := msg.ReplyCtx.(wsReplyContext)
if rc.chatID != "group_chat_id_123" {
t.Fatalf("expected chatID 'group_chat_id_123', got %q", rc.chatID)
}
case <-time.After(1 * time.Second):
t.Fatal("handler not called")
}
}
func TestHandleMsgCallback_StripsBotMention(t *testing.T) {
p, captured := newCapturedWSPlatform()
p.botID = "robot01"
body := wsMsgCallbackBody{
MsgID: "msg_mention",
ChatID: "grp1",
ChatType: "group",
MsgType: "text",
AibotID: "robot01",
}
body.From.UserID = "u1"
body.Text.Content = "允许 @Robot01"
body.CreateTime = time.Now().Unix()
p.handleMsgCallback(wsCallbackFrame(t, "req_m", body))
select {
case msg := <-captured:
if msg.Content != "允许" {
t.Fatalf("expected stripped content %q, got %q", "允许", msg.Content)
}
case <-time.After(1 * time.Second):
t.Fatal("handler not called")
}
}
// ---------------------------------------------------------------------------
// ReconstructReplyCtx
// ---------------------------------------------------------------------------
func TestReconstructReplyCtx_Valid(t *testing.T) {
p := &WSPlatform{}
rctx, err := p.ReconstructReplyCtx("wecom:chatid123:user456")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
rc := rctx.(wsReplyContext)
if rc.chatID != "chatid123" || rc.userID != "user456" {
t.Fatalf("unexpected context: %+v", rc)
}
}
func TestReconstructReplyCtx_InvalidPrefix(t *testing.T) {
p := &WSPlatform{}
_, err := p.ReconstructReplyCtx("slack:chatid123:user456")
if err == nil {
t.Fatal("expected error for invalid prefix")
}
}
func TestReconstructReplyCtx_TooFewParts(t *testing.T) {
p := &WSPlatform{}
_, err := p.ReconstructReplyCtx("wecom:only")
if err == nil {
t.Fatal("expected error for too few parts")
}
}
// ---------------------------------------------------------------------------
// writeAndWaitAck
// ---------------------------------------------------------------------------
func TestWriteAndWaitAck_SuccessfulAck(t *testing.T) {
p := &WSPlatform{}
reqID := "send_1"
ch := make(chan wsAckResult, 1)
p.pendingAcks.Store(reqID, ch)
// Simulate receiving ack in another goroutine
go func() {
time.Sleep(10 * time.Millisecond)
p.dispatchAck(reqID, wsAckResult{})
}()
assertAckResult(t, ch, func(result wsAckResult) {
if result.err != nil {
t.Fatalf("expected nil ack error, got %v", result.err)
}
})
}
func TestWriteAndWaitAck_AckWithError(t *testing.T) {
p := &WSPlatform{}
reqID := "send_2"
ch := make(chan wsAckResult, 1)
p.pendingAcks.Store(reqID, ch)
ackErr := fmt.Errorf("wecom-ws: ack error: errcode=40001 errmsg=invalid token")
go func() {
time.Sleep(10 * time.Millisecond)
p.dispatchAck(reqID, wsAckResult{err: ackErr})
}()
assertAckResult(t, ch, func(result wsAckResult) {
if result.err == nil {
t.Fatal("expected ack error, got nil")
}
if result.err.Error() != ackErr.Error() {
t.Fatalf("unexpected error: %v", result.err)
}
})
}
func TestWriteAndWaitAck_Timeout(t *testing.T) {
p := &WSPlatform{}
reqID := "send_timeout"
ch := make(chan wsAckResult, 1)
p.pendingAcks.Store(reqID, ch)
// Nobody sends ack → should timeout
start := time.Now()
select {
case <-ch:
t.Fatal("should not receive from channel without ack")
case <-time.After(100 * time.Millisecond):
// Expected: timed out without blocking forever
}
elapsed := time.Since(start)
if elapsed > 1*time.Second {
t.Fatalf("timeout took too long: %v", elapsed)
}
// Clean up
p.pendingAcks.Delete(reqID)
}
func TestWriteAndWaitAck_ContextCancelled(t *testing.T) {
p := &WSPlatform{}
reqID := "send_cancel"
ch := make(chan wsAckResult, 1)
p.pendingAcks.Store(reqID, ch)
ctx, cancel := context.WithCancel(context.Background())
go func() {
time.Sleep(20 * time.Millisecond)
cancel()
}()
select {
case <-ch:
t.Fatal("should not receive ack")
case <-ctx.Done():
// Expected: context cancelled
case <-time.After(1 * time.Second):
t.Fatal("timed out")
}
p.pendingAcks.Delete(reqID)
}
// ---------------------------------------------------------------------------
// handleFrame — ACK dispatch
// ---------------------------------------------------------------------------
func TestHandleFrame_AckDispatch(t *testing.T) {
p := &WSPlatform{}
reqID := "aibot_send_msg_1"
ch := make(chan wsAckResult, 1)
p.pendingAcks.Store(reqID, ch)
errCode := 0
frame := wsFrame{
Cmd: "",
Headers: wsFrameHeaders{ReqID: reqID},
ErrCode: &errCode,
ErrMsg: "ok",
}
p.handleFrame(frame)
assertAckResult(t, ch, func(result wsAckResult) {
if result.err != nil {
t.Fatalf("expected nil error for successful ack, got %v", result.err)
}
})
}
func TestHandleFrame_AckDispatch_WithError(t *testing.T) {
p := &WSPlatform{}
reqID := "aibot_send_msg_2"
ch := make(chan wsAckResult, 1)
p.pendingAcks.Store(reqID, ch)
errCode := 40001
frame := wsFrame{
Cmd: "",
Headers: wsFrameHeaders{ReqID: reqID},
ErrCode: &errCode,
ErrMsg: "invalid token",
}
p.handleFrame(frame)
assertAckResult(t, ch, func(result wsAckResult) {
if result.err == nil {
t.Fatal("expected error for failed ack, got nil")
}
})
}
func assertAckResult(t *testing.T, ch <-chan wsAckResult, check func(wsAckResult)) {
t.Helper()
select {
case result := <-ch:
check(result)
case <-time.After(100 * time.Millisecond):
t.Fatal("ack not dispatched")
}
}
func TestHandleFrame_PingAck_ResetsMissedPong(t *testing.T) {
p := &WSPlatform{}
p.missedPong.Store(2)
frame := wsFrame{
Cmd: "",
Headers: wsFrameHeaders{ReqID: "ping_1"},
}
p.handleFrame(frame)
if p.missedPong.Load() != 0 {
t.Fatalf("expected missedPong to be reset to 0, got %d", p.missedPong.Load())
}
}
// ---------------------------------------------------------------------------
// generateReqID
// ---------------------------------------------------------------------------
func TestGenerateReqID_Monotonic(t *testing.T) {
p := &WSPlatform{}
ids := make(map[string]bool)
for i := 0; i < 100; i++ {
id := p.generateReqID("test")
if ids[id] {
t.Fatalf("duplicate req_id: %s", id)
}
ids[id] = true
}
}
func TestGenerateReqID_Format(t *testing.T) {
p := &WSPlatform{}
id := p.generateReqID("ping")
if id != "ping_1" {
t.Fatalf("expected ping_1, got %s", id)
}
id2 := p.generateReqID("aibot_send_msg")
if id2 != "aibot_send_msg_2" {
t.Fatalf("expected aibot_send_msg_2, got %s", id2)
}
}
// ---------------------------------------------------------------------------
// SendImage
// ---------------------------------------------------------------------------
func TestWSPlatformSendImage_UploadsAndSendsMedia(t *testing.T) {
imageData := []byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n', 1, 2, 3}
serverDone := make(chan error, 1)
upgrader := websocket.Upgrader{}
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
serverDone <- err
return
}
defer conn.Close()
serverDone <- assertWeComWSSendImageFrames(conn, imageData)
}))
defer server.Close()
wsURL := "ws" + server.URL[len("http"):]
conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil)
if err != nil {
t.Fatalf("dial test websocket: %v", err)
}
defer conn.Close()
p := &WSPlatform{conn: conn}
go func() {
for {
var frame wsFrame
if err := conn.ReadJSON(&frame); err != nil {
return
}
p.handleFrame(frame)
}
}()
err = p.SendImage(context.Background(), wsReplyContext{chatID: "chat1", userID: "u1"}, core.ImageAttachment{
MimeType: "image/png",
Data: imageData,
FileName: "chart.png",
})
if err != nil {
t.Fatalf("SendImage returned error: %v", err)
}
select {
case err := <-serverDone:
if err != nil {
t.Fatal(err)
}
case <-time.After(time.Second):
t.Fatal("server did not observe all expected frames")
}
}
func assertWeComWSSendImageFrames(conn *websocket.Conn, imageData []byte) error {
var initFrame struct {
Cmd string `json:"cmd"`
Headers wsFrameHeaders `json:"headers"`
Body struct {
Type string `json:"type"`
Filename string `json:"filename"`
TotalSize int `json:"total_size"`
TotalChunks int `json:"total_chunks"`
MD5 string `json:"md5"`
} `json:"body"`
}
if err := conn.ReadJSON(&initFrame); err != nil {
return fmt.Errorf("read init frame: %w", err)
}
sum := md5.Sum(imageData)
if initFrame.Cmd != "aibot_upload_media_init" ||
initFrame.Body.Type != "image" ||
initFrame.Body.Filename != "chart.png" ||
initFrame.Body.TotalSize != len(imageData) ||
initFrame.Body.TotalChunks != 1 ||
initFrame.Body.MD5 != hex.EncodeToString(sum[:]) {
return fmt.Errorf("unexpected init frame: %#v", initFrame)
}
if err := conn.WriteJSON(map[string]any{
"headers": initFrame.Headers,
"errcode": 0,
"errmsg": "ok",
"body": map[string]string{"upload_id": "upload-1"},
}); err != nil {
return fmt.Errorf("write init ack: %w", err)
}
var chunkFrame struct {
Cmd string `json:"cmd"`
Headers wsFrameHeaders `json:"headers"`
Body struct {
UploadID string `json:"upload_id"`
ChunkIndex int `json:"chunk_index"`
Base64Data string `json:"base64_data"`
} `json:"body"`
}
if err := conn.ReadJSON(&chunkFrame); err != nil {
return fmt.Errorf("read chunk frame: %w", err)
}
if chunkFrame.Cmd != "aibot_upload_media_chunk" ||
chunkFrame.Body.UploadID != "upload-1" ||
chunkFrame.Body.ChunkIndex != 0 ||
chunkFrame.Body.Base64Data != base64.StdEncoding.EncodeToString(imageData) {
return fmt.Errorf("unexpected chunk frame: %#v", chunkFrame)
}
if err := conn.WriteJSON(map[string]any{
"headers": chunkFrame.Headers,
"errcode": 0,
"errmsg": "ok",
}); err != nil {
return fmt.Errorf("write chunk ack: %w", err)
}
var finishFrame struct {
Cmd string `json:"cmd"`
Headers wsFrameHeaders `json:"headers"`
Body struct {
UploadID string `json:"upload_id"`
} `json:"body"`
}
if err := conn.ReadJSON(&finishFrame); err != nil {
return fmt.Errorf("read finish frame: %w", err)
}
if finishFrame.Cmd != "aibot_upload_media_finish" || finishFrame.Body.UploadID != "upload-1" {
return fmt.Errorf("unexpected finish frame: %#v", finishFrame)
}
if err := conn.WriteJSON(map[string]any{
"headers": finishFrame.Headers,
"errcode": 0,
"errmsg": "ok",
"body": map[string]string{"media_id": "media-1"},
}); err != nil {
return fmt.Errorf("write finish ack: %w", err)
}
var sendFrame struct {
Cmd string `json:"cmd"`
Headers wsFrameHeaders `json:"headers"`
Body struct {
ChatID string `json:"chatid"`
MsgType string `json:"msgtype"`
Image struct {
MediaID string `json:"media_id"`
} `json:"image"`
} `json:"body"`
}
if err := conn.ReadJSON(&sendFrame); err != nil {
return fmt.Errorf("read send frame: %w", err)
}
if sendFrame.Cmd != "aibot_send_msg" ||
sendFrame.Body.ChatID != "chat1" ||
sendFrame.Body.MsgType != "image" ||
sendFrame.Body.Image.MediaID != "media-1" {
return fmt.Errorf("unexpected send frame: %#v", sendFrame)
}
if err := conn.WriteJSON(map[string]any{
"headers": sendFrame.Headers,
"errcode": 0,
"errmsg": "ok",
}); err != nil {
return fmt.Errorf("write send ack: %w", err)
}
return nil
}
// ---------------------------------------------------------------------------
// generateReqID — concurrency safety
// ---------------------------------------------------------------------------
func TestGenerateReqID_ConcurrentSafety(t *testing.T) {
p := &WSPlatform{}
var wg sync.WaitGroup
ids := sync.Map{}
for i := 0; i < 50; i++ {
wg.Add(1)
go func() {
defer wg.Done()
id := p.generateReqID("concurrent")
if _, loaded := ids.LoadOrStore(id, true); loaded {
t.Errorf("duplicate req_id: %s", id)
}
}()
}
wg.Wait()
}
// ---------------------------------------------------------------------------
// newWebSocket
// ---------------------------------------------------------------------------
func TestNewWebSocket_MissingCredentials(t *testing.T) {
tests := []struct {
name string
opts map[string]any
}{
{"empty opts", map[string]any{}},
{"missing bot_secret", map[string]any{"bot_id": "aib123"}},
{"missing bot_id", map[string]any{"bot_secret": "secret"}},
{"both empty strings", map[string]any{"bot_id": "", "bot_secret": ""}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
_, err := newWebSocket(tt.opts)
if err == nil {
t.Fatal("expected error for missing credentials")
}
})
}
}
func TestNewWebSocket_ValidConfig(t *testing.T) {
p, err := newWebSocket(map[string]any{
"bot_id": "aibTest",
"bot_secret": "secretXYZ",
"allow_from": "user1,user2",
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
ws := p.(*WSPlatform)
if ws.botID != "aibTest" || ws.secret != "secretXYZ" || ws.allowFrom != "user1,user2" {
t.Fatalf("unexpected config: botID=%s secret=%s allowFrom=%s", ws.botID, ws.secret, ws.allowFrom)
}
}