discord_test.go 10 KB

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