初始化仓库

This commit is contained in:
2026-06-02 23:14:41 +08:00
commit 0bc3f02670
520 changed files with 191097 additions and 0 deletions
+673
View File
@@ -0,0 +1,673 @@
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)
}
}