discord_test.go 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355
  1. package discord
  2. import (
  3. "context"
  4. "encoding/json"
  5. "io"
  6. "net/http"
  7. "net/http/httptest"
  8. "path/filepath"
  9. "strings"
  10. "testing"
  11. "time"
  12. "github.com/mhsanaei/3x-ui/v3/internal/database"
  13. "github.com/mhsanaei/3x-ui/v3/internal/web/service"
  14. )
  15. func setupTestDB(t *testing.T) service.SettingService {
  16. t.Helper()
  17. dbPath := filepath.Join(t.TempDir(), "x-ui.db")
  18. if err := database.InitDB(dbPath); err != nil {
  19. t.Fatalf("init db: %v", err)
  20. }
  21. t.Cleanup(func() { _ = database.CloseDB() })
  22. return service.SettingService{}
  23. }
  24. func TestSendMessage_Success(t *testing.T) {
  25. settingService := setupTestDB(t)
  26. if err := settingService.SetDiscordBotToken("test-bot-token"); err != nil {
  27. t.Fatal(err)
  28. }
  29. if err := settingService.SetDiscordChannelId("123456789012345678"); err != nil {
  30. t.Fatal(err)
  31. }
  32. var reqMethod, reqPath, reqAuth, reqUA, reqCT string
  33. var reqPayload MessagePayload
  34. server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  35. reqMethod = r.Method
  36. reqPath = r.URL.Path
  37. reqAuth = r.Header.Get("Authorization")
  38. reqUA = r.Header.Get("User-Agent")
  39. reqCT = r.Header.Get("Content-Type")
  40. body, _ := io.ReadAll(r.Body)
  41. _ = json.Unmarshal(body, &reqPayload)
  42. w.WriteHeader(http.StatusOK)
  43. _, _ = w.Write([]byte(`{"id": "msg-123"}`))
  44. }))
  45. defer server.Close()
  46. svc := NewDiscordService(settingService)
  47. svc.SetBaseURL(server.URL)
  48. svc.SetHTTPClient(server.Client())
  49. payload := MessagePayload{
  50. Content: "Hello Discord!",
  51. Embeds: []Embed{
  52. {
  53. Title: "Test Embed",
  54. Description: "Desc",
  55. Color: ColorGreen,
  56. },
  57. },
  58. }
  59. if err := svc.SendMessage(context.Background(), payload); err != nil {
  60. t.Fatalf("SendMessage failed: %v", err)
  61. }
  62. if reqMethod != http.MethodPost {
  63. t.Errorf("expected POST, got %s", reqMethod)
  64. }
  65. expectedPath := "/channels/123456789012345678/messages"
  66. if reqPath != expectedPath {
  67. t.Errorf("expected path %s, got %s", expectedPath, reqPath)
  68. }
  69. if reqAuth != "Bot test-bot-token" {
  70. t.Errorf("expected auth 'Bot test-bot-token', got %s", reqAuth)
  71. }
  72. if reqUA != discordUserAgent {
  73. t.Errorf("expected User-Agent %s, got %s", discordUserAgent, reqUA)
  74. }
  75. if reqCT != "application/json" {
  76. t.Errorf("expected Content-Type application/json, got %s", reqCT)
  77. }
  78. if reqPayload.Content != "Hello Discord!" || len(reqPayload.Embeds) != 1 {
  79. t.Errorf("payload mismatch: %+v", reqPayload)
  80. }
  81. }
  82. func TestSendMessage_CreatedStatus(t *testing.T) {
  83. settingService := setupTestDB(t)
  84. _ = settingService.SetDiscordBotToken("token")
  85. _ = settingService.SetDiscordChannelId("ch-1")
  86. server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  87. w.WriteHeader(http.StatusCreated)
  88. _, _ = w.Write([]byte(`{}`))
  89. }))
  90. defer server.Close()
  91. svc := NewDiscordService(settingService)
  92. svc.SetBaseURL(server.URL)
  93. svc.SetHTTPClient(server.Client())
  94. if err := svc.SendEmbed(context.Background(), Embed{Title: "Title"}); err != nil {
  95. t.Fatalf("SendEmbed failed: %v", err)
  96. }
  97. }
  98. func TestSendMessage_StatusCodes(t *testing.T) {
  99. cases := []struct {
  100. name string
  101. statusCode int
  102. respBody string
  103. wantErrSub string
  104. }{
  105. {"Bad Request", http.StatusBadRequest, `{"message": "Invalid Form Body"}`, "discord bad request (400)"},
  106. {"Unauthorized", http.StatusUnauthorized, `{"message": "401: Unauthorized"}`, "discord unauthorized (401)"},
  107. {"Forbidden", http.StatusForbidden, `{"message": "Missing Permissions"}`, "discord forbidden (403)"},
  108. {"NotFound", http.StatusNotFound, `{"message": "Unknown Channel"}`, "discord not found (404)"},
  109. {"RateLimited", http.StatusTooManyRequests, `{"retry_after": 1.5}`, "discord rate limited (429)"},
  110. {"InternalError", http.StatusInternalServerError, `{"message": "Server Error"}`, "discord API error (500)"},
  111. }
  112. for _, tc := range cases {
  113. t.Run(tc.name, func(t *testing.T) {
  114. settingService := setupTestDB(t)
  115. _ = settingService.SetDiscordBotToken("token")
  116. _ = settingService.SetDiscordChannelId("ch-1")
  117. server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  118. w.WriteHeader(tc.statusCode)
  119. _, _ = w.Write([]byte(tc.respBody))
  120. }))
  121. defer server.Close()
  122. svc := NewDiscordService(settingService)
  123. svc.SetBaseURL(server.URL)
  124. svc.SetHTTPClient(server.Client())
  125. err := svc.SendEmbed(context.Background(), Embed{Title: "Test"})
  126. if err == nil {
  127. t.Fatalf("expected error for status %d, got nil", tc.statusCode)
  128. }
  129. if !strings.Contains(err.Error(), tc.wantErrSub) {
  130. t.Errorf("expected error containing %q, got %q", tc.wantErrSub, err.Error())
  131. }
  132. })
  133. }
  134. }
  135. func TestSendMessage_MissingConfig(t *testing.T) {
  136. settingService := setupTestDB(t)
  137. svc := NewDiscordService(settingService)
  138. // Both empty
  139. err := svc.SendMessage(context.Background(), MessagePayload{Content: "Hi"})
  140. if err == nil || !strings.Contains(err.Error(), "token is not configured") {
  141. t.Fatalf("expected token not configured error, got %v", err)
  142. }
  143. // Token set, channel empty
  144. _ = settingService.SetDiscordBotToken("some-token")
  145. err = svc.SendMessage(context.Background(), MessagePayload{Content: "Hi"})
  146. if err == nil || !strings.Contains(err.Error(), "channel id is not configured") {
  147. t.Fatalf("expected channel id not configured error, got %v", err)
  148. }
  149. }
  150. func TestSendTest(t *testing.T) {
  151. settingService := setupTestDB(t)
  152. _ = settingService.SetDiscordBotToken("test-bot-token")
  153. _ = settingService.SetDiscordChannelId("999888777")
  154. var receivedPayload MessagePayload
  155. server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  156. body, _ := io.ReadAll(r.Body)
  157. _ = json.Unmarshal(body, &receivedPayload)
  158. w.WriteHeader(http.StatusOK)
  159. }))
  160. defer server.Close()
  161. svc := NewDiscordService(settingService)
  162. svc.SetBaseURL(server.URL)
  163. svc.SetHTTPClient(server.Client())
  164. if err := svc.SendTest(context.Background()); err != nil {
  165. t.Fatalf("SendTest failed: %v", err)
  166. }
  167. if len(receivedPayload.Embeds) != 1 {
  168. t.Fatalf("expected 1 embed, got %d", len(receivedPayload.Embeds))
  169. }
  170. embed := receivedPayload.Embeds[0]
  171. if embed.Color != ColorGreen {
  172. t.Errorf("expected ColorGreen (0x%X), got 0x%X", ColorGreen, embed.Color)
  173. }
  174. if embed.Timestamp == "" {
  175. t.Error("expected non-empty timestamp")
  176. } else {
  177. parsed, err := time.Parse(time.RFC3339, embed.Timestamp)
  178. if err != nil {
  179. t.Errorf("timestamp is not RFC3339: %v", err)
  180. }
  181. if parsed.Location() != time.UTC {
  182. t.Errorf("expected UTC timestamp location, got %v", parsed.Location())
  183. }
  184. }
  185. if len(embed.Fields) == 0 {
  186. t.Error("expected test embed to have fields")
  187. }
  188. }
  189. func TestSendMessage_BotPrefixHandling(t *testing.T) {
  190. settingService := setupTestDB(t)
  191. _ = settingService.SetDiscordBotToken("Bot prefixed-token")
  192. _ = settingService.SetDiscordChannelId("ch-100")
  193. var receivedAuth string
  194. server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  195. receivedAuth = r.Header.Get("Authorization")
  196. w.WriteHeader(http.StatusOK)
  197. _, _ = w.Write([]byte(`{}`))
  198. }))
  199. defer server.Close()
  200. svc := NewDiscordService(settingService)
  201. svc.SetBaseURL(server.URL)
  202. svc.SetHTTPClient(server.Client())
  203. if err := svc.SendEmbed(context.Background(), Embed{Title: "Prefix Test"}); err != nil {
  204. t.Fatalf("SendEmbed failed: %v", err)
  205. }
  206. if receivedAuth != "Bot prefixed-token" {
  207. t.Errorf("expected 'Bot prefixed-token', got %q", receivedAuth)
  208. }
  209. }
  210. func TestSendMessage_NoContentStatus(t *testing.T) {
  211. settingService := setupTestDB(t)
  212. _ = settingService.SetDiscordBotToken("token")
  213. _ = settingService.SetDiscordChannelId("ch-204")
  214. server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  215. w.WriteHeader(http.StatusNoContent)
  216. }))
  217. defer server.Close()
  218. svc := NewDiscordService(settingService)
  219. svc.SetBaseURL(server.URL)
  220. svc.SetHTTPClient(server.Client())
  221. if err := svc.SendEmbed(context.Background(), Embed{Title: "204 Test"}); err != nil {
  222. t.Fatalf("SendEmbed failed for 204: %v", err)
  223. }
  224. }
  225. func TestSendMessage_ContextCancelled(t *testing.T) {
  226. settingService := setupTestDB(t)
  227. _ = settingService.SetDiscordBotToken("token")
  228. _ = settingService.SetDiscordChannelId("ch-1")
  229. server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  230. time.Sleep(100 * time.Millisecond)
  231. w.WriteHeader(http.StatusOK)
  232. }))
  233. defer server.Close()
  234. svc := NewDiscordService(settingService)
  235. svc.SetBaseURL(server.URL)
  236. svc.SetHTTPClient(server.Client())
  237. ctx, cancel := context.WithCancel(context.Background())
  238. cancel()
  239. err := svc.SendMessage(ctx, MessagePayload{Content: "Cancelled"})
  240. if err == nil {
  241. t.Fatal("expected error with cancelled context, got nil")
  242. }
  243. }
  244. func TestSendMessageWithFiles_Success(t *testing.T) {
  245. settingService := setupTestDB(t)
  246. _ = settingService.SetDiscordBotToken("test-bot-token")
  247. _ = settingService.SetDiscordChannelId("ch-multipart")
  248. var receivedCT string
  249. var receivedPayload MessagePayload
  250. receivedFiles := make(map[string][]byte)
  251. server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  252. receivedCT = r.Header.Get("Content-Type")
  253. mr, err := r.MultipartReader()
  254. if err != nil {
  255. t.Fatalf("MultipartReader error: %v", err)
  256. }
  257. for {
  258. part, err := mr.NextPart()
  259. if err == io.EOF {
  260. break
  261. }
  262. if err != nil {
  263. t.Fatalf("NextPart error: %v", err)
  264. }
  265. data, _ := io.ReadAll(part)
  266. formName := part.FormName()
  267. if formName == "payload_json" {
  268. _ = json.Unmarshal(data, &receivedPayload)
  269. } else {
  270. receivedFiles[part.FileName()] = data
  271. }
  272. }
  273. w.WriteHeader(http.StatusOK)
  274. _, _ = w.Write([]byte(`{"id": "msg-files"}`))
  275. }))
  276. defer server.Close()
  277. svc := NewDiscordService(settingService)
  278. svc.SetBaseURL(server.URL)
  279. svc.SetHTTPClient(server.Client())
  280. payload := MessagePayload{
  281. Content: "Report message",
  282. Embeds: []Embed{{Title: "Report Embed"}},
  283. }
  284. files := []FileAttachment{
  285. {Filename: "x-ui.db", Data: []byte("sqlite-db-binary")},
  286. {Filename: "config.json", Data: []byte(`{"log":{}}`)},
  287. }
  288. if err := svc.SendMessageWithFiles(context.Background(), payload, files...); err != nil {
  289. t.Fatalf("SendMessageWithFiles failed: %v", err)
  290. }
  291. if !strings.HasPrefix(receivedCT, "multipart/form-data; boundary=") {
  292. t.Errorf("expected multipart/form-data content type, got %s", receivedCT)
  293. }
  294. if receivedPayload.Content != "Report message" || len(receivedPayload.Embeds) != 1 {
  295. t.Errorf("payload mismatch: %+v", receivedPayload)
  296. }
  297. if string(receivedFiles["x-ui.db"]) != "sqlite-db-binary" {
  298. t.Errorf("x-ui.db mismatch: %s", string(receivedFiles["x-ui.db"]))
  299. }
  300. if string(receivedFiles["config.json"]) != `{"log":{}}` {
  301. t.Errorf("config.json mismatch: %s", string(receivedFiles["config.json"]))
  302. }
  303. }