package app import ( "context" "database/sql" "encoding/json" "fmt" "net" "net/http" "net/http/httptest" "os" "path/filepath" "regexp" "strings" "testing" "time" "github.com/go-sql-driver/mysql" "github.com/gorilla/websocket" ) // Never run against IM or a remote database: create and clean up a unique local // schema, apply the real migrations, and exercise the actual handlers/SQL. func isolatedIMDatabase(t *testing.T) *sql.DB { t.Helper() dsn := os.Getenv("IM_TEST_MYSQL_DSN") if dsn == "" { t.Skip("set IM_TEST_MYSQL_DSN to enable isolated local IM integration tests") } cfg, err := mysql.ParseDSN(dsn) if err != nil { t.Fatal("invalid test DSN") } host, _, err := net.SplitHostPort(cfg.Addr) if err != nil || cfg.Net != "tcp" || (host != "127.0.0.1" && host != "localhost" && host != "::1") || cfg.DBName != "" { t.Fatal("IM tests require a loopback TCP DSN without a database name") } cfg.ParseTime, cfg.MultiStatements = true, true admin, err := sql.Open("mysql", cfg.FormatDSN()) if err != nil { t.Fatal(err) } t.Cleanup(func() { admin.Close() }) name := fmt.Sprintf("im_unread_test_%d", time.Now().UnixNano()) if !regexp.MustCompile(`^im_unread_test_[0-9]+$`).MatchString(name) { t.Fatal("unsafe test schema") } if _, err := admin.Exec("CREATE DATABASE `" + name + "` CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci"); err != nil { t.Fatal(err) } t.Cleanup(func() { if _, err := admin.Exec("DROP DATABASE `" + name + "`"); err != nil { t.Errorf("test database cleanup: %v", err) } }) cfg.DBName = name db, err := sql.Open("mysql", cfg.FormatDSN()) if err != nil { t.Fatal(err) } t.Cleanup(func() { db.Close() }) files, err := filepath.Glob(filepath.Join("..", "..", "migrations", "*.sql")) if err != nil { t.Fatal(err) } for _, file := range files { body, err := os.ReadFile(file) if err != nil { t.Fatal(err) } if _, err := db.Exec(string(body)); err != nil { t.Fatalf("migration %s: %v", file, err) } } for id := 1; id <= 3; id++ { for _, query := range []string{ `INSERT INTO users(id,public_id,password_hash) VALUES(?,CONCAT('IMTEST',?),'!NO_LOGIN')`, `INSERT INTO user_profiles(user_id,nickname) VALUES(?,CONCAT('IM测试',?))`, `INSERT INTO user_privacy_settings(user_id) VALUES(?)`, } { args := []any{id} if strings.Count(query, "?") == 2 { args = append(args, id) } if _, err := db.Exec(query, args...); err != nil { t.Fatal(err) } } } return db } func imTestRequest(user int64, method, path, body string) *http.Request { r := httptest.NewRequest(method, path, strings.NewReader(body)) return r.WithContext(context.WithValue(r.Context(), identityKey{}, identity{ID: user, Role: "user"})) } func imTestCall(t *testing.T, handler http.HandlerFunc, user int64, method, path, body string, wantStatus int) json.RawMessage { t.Helper() w := httptest.NewRecorder() handler(w, imTestRequest(user, method, path, body)) if w.Code != wantStatus { t.Fatalf("%s %s user %d: HTTP %d: %s", method, path, user, w.Code, w.Body.String()) } var response struct { Data json.RawMessage `json:"data"` } if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil { t.Fatal(err) } return response.Data } func TestIMUnreadMySQL(t *testing.T) { db := isolatedIMDatabase(t) a := &App{db: db, hub: NewHub(), config: Config{Environment: "development", JWTSecret: "isolated-im-regression-secret-only"}} data := imTestCall(t, a.directConversation, 1, "POST", "/api/v1/im/conversations", `{"userId":2}`, 200) var conversation struct { ID int64 `json:"id"` } if err := json.Unmarshal(data, &conversation); err != nil { t.Fatal(err) } id := conversation.ID messagePath := fmt.Sprintf("/api/v1/im/conversations/%d/messages", id) settingsPath := fmt.Sprintf("/api/v1/im/conversations/%d/settings", id) unread := func(user int64, want int, empty bool) { t.Helper() var list struct { Items []struct { Unread int `json:"unread"` LastMessageAt *string `json:"lastMessageAt"` } `json:"items"` } if err := json.Unmarshal(imTestCall(t, a.conversations, user, "GET", "/api/v1/im/conversations", "", 200), &list); err != nil { t.Fatal(err) } if len(list.Items) != 1 || list.Items[0].Unread != want { t.Fatalf("user %d unread: %+v, want %d", user, list, want) } if empty { if list.Items[0].LastMessageAt != nil { t.Fatal("empty conversation time must be null") } return } if list.Items[0].LastMessageAt == nil { t.Fatal("missing message timestamp") } if _, err := time.Parse(time.RFC3339Nano, *list.Items[0].LastMessageAt); err != nil { t.Fatal("message time must be RFC3339, not sql.NullTime JSON") } } send := func(user int64, key string) messageView { t.Helper() var message messageView payload := fmt.Sprintf(`{"clientMsgId":%q,"type":1,"content":{"text":"hello"}}`, key) if err := json.Unmarshal(imTestCall(t, a.sendMessageHTTP, user, "POST", messagePath, payload, 200), &message); err != nil { t.Fatal(err) } return message } t.Run("empty conversation has no unread and null time", func(t *testing.T) { unread(1, 0, true); unread(2, 0, true) }) t.Run("sender zero recipient one", func(t *testing.T) { send(1, "a1"); unread(1, 0, false); unread(2, 1, false) }) t.Run("reply does not mark earlier incoming messages read", func(t *testing.T) { first := send(2, "b1") again := send(2, "b1") if first.ID != again.ID { t.Fatal("retry must be idempotent") } unread(1, 1, false) unread(2, 1, false) send(1, "a2") unread(1, 1, false) unread(2, 2, false) }) t.Run("incremental history returns only the sequence gap", func(t *testing.T) { var delta struct { Items []messageView `json:"items"` HasMore bool `json:"hasMore"` NextAfterSeq int64 `json:"nextAfterSeq"` } body := imTestCall(t, a.messages, 2, "GET", messagePath+"?afterSeq=1&limit=1", "", 200) if err := json.Unmarshal(body, &delta); err != nil { t.Fatal(err) } if len(delta.Items) != 1 || delta.Items[0].Seq != 2 || !delta.HasMore || delta.NextAfterSeq != 2 { t.Fatalf("unexpected incremental page: %+v", delta) } imTestCall(t, a.messages, 2, "GET", messagePath+"?beforeSeq=3&afterSeq=1", "", 400) }) t.Run("older history cannot read newer incoming messages", func(t *testing.T) { imTestCall(t, a.messages, 1, "GET", messagePath+"?beforeSeq=2&limit=1", "", 200) unread(1, 1, false) }) t.Run("latest history clears only the reader", func(t *testing.T) { imTestCall(t, a.messages, 2, "GET", messagePath+"?limit=1", "", 200) unread(1, 1, false) unread(2, 0, false) }) t.Run("read permissions and future sequence clamp", func(t *testing.T) { imTestCall(t, a.conversationSettings, 3, "PUT", settingsPath, `{"readSeq":999}`, 403) imTestCall(t, a.conversationSettings, 1, "PUT", settingsPath, `{"readSeq":-1}`, 400) var state struct { ReadSeq int64 `json:"readSeq"` Unread int64 `json:"unread"` } if err := json.Unmarshal(imTestCall(t, a.conversationSettings, 1, "PUT", settingsPath, `{"readSeq":999}`, 200), &state); err != nil { t.Fatal(err) } if state.ReadSeq != 3 || state.Unread != 0 { t.Fatalf("authoritative read state: %+v", state) } var read int64 if err := db.QueryRow(`SELECT read_seq FROM im_conversation_members WHERE conversation_id=? AND user_id=1`, id).Scan(&read); err != nil || read != 3 { t.Fatalf("read marker %d: %v", read, err) } unread(1, 0, false) m := send(2, "b2") unread(1, 1, false) unread(2, 0, false) imTestCall(t, a.recallMessage, 2, "POST", fmt.Sprintf("/api/v1/im/messages/%d/recall", m.ID), "", 200) unread(1, 0, false) }) t.Run("moderated and cleared messages are not unread", func(t *testing.T) { m := send(1, "a3") unread(2, 1, false) if _, err := db.Exec(`UPDATE im_messages SET admin_removed_at=NOW(3) WHERE id=?`, m.ID); err != nil { t.Fatal(err) } unread(2, 0, false) if _, err := db.Exec(`UPDATE im_messages SET admin_removed_at=NULL WHERE id=?`, m.ID); err != nil { t.Fatal(err) } unread(2, 1, false) if _, err := db.Exec(`UPDATE im_conversation_members SET clear_seq=? WHERE conversation_id=? AND user_id=2`, m.Seq, id); err != nil { t.Fatal(err) } unread(2, 0, false) }) t.Run("websocket send echoes are persisted once and HTTP read broadcasts ack", func(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(a.websocket)) defer server.Close() token, err := a.token(1, "user", "fixture", time.Minute) if err != nil { t.Fatal(err) } dialer := websocket.Dialer{Subprotocols: []string{"xingyu.jwt." + token}} conn, _, err := dialer.Dial("ws"+strings.TrimPrefix(server.URL, "http"), nil) if err != nil { t.Fatal(err) } defer conn.Close() readFrame := func() map[string]any { t.Helper() conn.SetReadDeadline(time.Now().Add(3 * time.Second)) var frame map[string]any if err := conn.ReadJSON(&frame); err != nil { t.Fatal(err) } return frame } if readFrame()["command"] != "AUTH_ACK" { t.Fatal("missing auth ack") } if err := conn.WriteJSON(map[string]any{"command": "SEND_MESSAGE", "conversationId": id, "clientMsgId": "ws-a4", "type": 1, "content": map[string]any{"text": "ws hello"}}); err != nil { t.Fatal(err) } if readFrame()["command"] != "MESSAGE_PUSH" { t.Fatal("missing sender echo") } unread(1, 0, false) unread(2, 1, false) imTestCall(t, a.messages, 1, "GET", messagePath, "", 200) readAck := readFrame() if readAck["command"] != "READ_ACK" { t.Fatal("HTTP read must sync this user's other device") } if data, ok := readAck["data"].(map[string]any); !ok || data["unread"] != float64(0) { t.Fatalf("HTTP read ack must include authoritative unread: %+v", readAck) } if err := conn.WriteJSON(map[string]any{"command": "READ", "conversationId": id, "readSeq": 999}); err != nil { t.Fatal(err) } if readFrame()["command"] != "READ_ACK" { t.Fatal("missing ws read ack") } incoming := send(2, "b3") if readFrame()["command"] != "MESSAGE_PUSH" { t.Fatal("missing incoming push") } unread(1, 1, false) own := send(1, "a5") if readFrame()["command"] != "MESSAGE_PUSH" { t.Fatal("missing outgoing echo") } var readState struct { ReadSeq int64 `json:"readSeq"` Unread int64 `json:"unread"` } body := imTestCall(t, a.conversationSettings, 1, "PUT", settingsPath, fmt.Sprintf(`{"readSeq":%d}`, incoming.Seq), 200) if err := json.Unmarshal(body, &readState); err != nil { t.Fatal(err) } if readState.ReadSeq != own.Seq || readState.Unread != 0 { t.Fatalf("own sequence must not keep a legacy badge alive: %+v, own seq %d", readState, own.Seq) } ack := readFrame() if data, ok := ack["data"].(map[string]any); ack["command"] != "READ_ACK" || !ok || data["readSeq"] != float64(own.Seq) || data["unread"] != float64(0) { t.Fatalf("legacy-compatible read ack: %+v", ack) } unread(1, 0, false) }) }