From c78f63a6ea99e09dd4cda43c7c71f9164baea9c9 Mon Sep 17 00:00:00 2001 From: JOY <5027251+JOY@users.noreply.github.com> Date: Thu, 10 Sep 2026 20:58:35 +0700 Subject: [PATCH 1/8] fix(storage): stop guest uploads from becoming same-origin script execution Uploaded assets are served from the application's own origin, so the extension stored in the key decides how a browser interprets the bytes. Nothing validated it: /api/message/upload_attachment accepted any file from an authenticated guest and the storage layer wrote whatever extension the client named, so a planted .html file ran against the dashboard origin the moment a staff member opened the attachment link. - Reject browser-active extensions and payloads (HTML, SVG, XML/XSL, scripts, legacy server-side pages) in AssetService.UploadFile and UploadBytes, and sniff the leading bytes instead of trusting the client-declared Content-Type so renaming a page to .png does not smuggle it through - Reduce client filenames to a bare basename within the 255 character column, cut on a rune boundary, and strip control characters - Derive the stored extension from an explicit MIME table. mime.ExtensionsByType consults the Windows registry, where text/html resolved to .ehtml and escaped the block list, and it differs between dev machines and production - Fall back to an inert .bin when no safe extension can be determined, so the server never sniffs the payload to decide its own Content-Type - Serve local assets with X-Content-Type-Options nosniff and force a download for anything that is not previewable media, which also neutralises document types already sitting in an existing storage root - Add error.e0348 to both backend locales for the rejection --- internal/bootstrap/server.go | 22 +- internal/bootstrap/server_route_test.go | 58 +++++ internal/pkg/i18nx/locales/en-US.yml | 1 + internal/pkg/i18nx/locales/zh-CN.yml | 1 + internal/services/asset_service.go | 25 ++- internal/services/storage/safety.go | 218 ++++++++++++++++++ internal/services/storage/safety_test.go | 268 +++++++++++++++++++++++ internal/services/storage/utils.go | 29 +-- 8 files changed, 598 insertions(+), 24 deletions(-) create mode 100644 internal/services/storage/safety.go create mode 100644 internal/services/storage/safety_test.go diff --git a/internal/bootstrap/server.go b/internal/bootstrap/server.go index d5fc669a..30d214a6 100644 --- a/internal/bootstrap/server.go +++ b/internal/bootstrap/server.go @@ -3,6 +3,7 @@ package bootstrap import ( "log/slog" "net/http" + "path" "strconv" "strings" "time" @@ -17,6 +18,7 @@ import ( "agent-desk/internal/pkg/i18nx" "agent-desk/internal/pkg/tracex" "agent-desk/internal/services" + "agent-desk/internal/services/storage" webspa "agent-desk/web" "github.com/gin-gonic/gin" @@ -44,11 +46,29 @@ func NewServer() (*gin.Engine, error) { handleSpa(app) - app.StaticFS(cfg.Storage.Local.BaseURL, ginx.StaticFiles(cfg.Storage.Local.Root)) + storageGroup := app.Group(cfg.Storage.Local.BaseURL, assetResponseHeaders()) + storageGroup.StaticFS("", ginx.StaticFiles(cfg.Storage.Local.Root)) return app, nil } +// assetResponseHeaders guards the locally stored assets. +// +// Those files are user-supplied bytes served from this application's own origin, +// so a response a browser renders inline is same-origin content. nosniff stops a +// browser reinterpreting the payload, and forcing a download for anything that is +// not previewable media means a document type that slipped in before this policy +// existed still cannot run as a page. +func assetResponseHeaders() gin.HandlerFunc { + return func(ctx *gin.Context) { + ctx.Header("X-Content-Type-Options", "nosniff") + if !storage.IsPreviewableExtension(path.Ext(ctx.Request.URL.Path)) { + ctx.Header("Content-Disposition", "attachment") + } + ctx.Next() + } +} + func corsMiddleware() gin.HandlerFunc { allowedOrigins := config.Current().Server.CORS.AllowedOrigins allowHeaders := "Origin, Content-Type, Accept, Authorization, X-Requested-With, X-Guest-Id, X-Channel-Id, X-External-Id, X-External-Name, X-Customer-Session-Token, X-Customer-Session-Expires-At" diff --git a/internal/bootstrap/server_route_test.go b/internal/bootstrap/server_route_test.go index 328b293f..0b2716e4 100644 --- a/internal/bootstrap/server_route_test.go +++ b/internal/bootstrap/server_route_test.go @@ -296,6 +296,64 @@ func TestNewServerSeparatesAPIStaticAndSPA(t *testing.T) { } } +func TestNewServerHardensStoredAssetResponses(t *testing.T) { + root := t.TempDir() + for _, name := range []string{"screenshot.png", "archive.zip", "legacy-page.html"} { + if err := os.WriteFile(filepath.Join(root, name), []byte("payload"), 0o644); err != nil { + t.Fatalf("WriteFile(%s) error = %v", name, err) + } + } + + config.SetCurrent(&config.Config{ + Storage: config.StorageConfig{ + Local: config.LocalStorageConfig{ + Root: root, + BaseURL: "/storage", + }, + }, + }) + + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + tests := []struct { + path string + wantStatus int + contentType string + wantAttachment bool + }{ + {path: "/storage/screenshot.png", wantStatus: http.StatusOK, contentType: "image/png"}, + {path: "/storage/archive.zip", wantStatus: http.StatusOK, wantAttachment: true}, + // A file planted before the upload policy existed must still not render. + {path: "/storage/legacy-page.html", wantStatus: http.StatusOK, wantAttachment: true}, + {path: "/storage/missing.png", wantStatus: http.StatusNotFound}, + } + + for _, tt := range tests { + rec := httptest.NewRecorder() + app.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, tt.path, nil)) + + if rec.Code != tt.wantStatus { + t.Fatalf("%s status=%d want %d", tt.path, rec.Code, tt.wantStatus) + } + if got := rec.Header().Get("X-Content-Type-Options"); got != "nosniff" { + t.Fatalf("%s X-Content-Type-Options=%q want nosniff", tt.path, got) + } + if tt.contentType != "" && !strings.Contains(rec.Header().Get("Content-Type"), tt.contentType) { + t.Fatalf("%s Content-Type=%q want %q", tt.path, rec.Header().Get("Content-Type"), tt.contentType) + } + got := rec.Header().Get("Content-Disposition") + if tt.wantAttachment && got != "attachment" { + t.Fatalf("%s Content-Disposition=%q want attachment", tt.path, got) + } + if !tt.wantAttachment && got != "" { + t.Fatalf("%s Content-Disposition=%q want empty so the asset renders inline", tt.path, got) + } + } +} + func TestNewServerAllowsConfiguredCORSOrigin(t *testing.T) { config.SetCurrent(&config.Config{ Server: config.ServerConfig{ diff --git a/internal/pkg/i18nx/locales/en-US.yml b/internal/pkg/i18nx/locales/en-US.yml index 9fbec600..e74b751b 100644 --- a/internal/pkg/i18nx/locales/en-US.yml +++ b/internal/pkg/i18nx/locales/en-US.yml @@ -345,6 +345,7 @@ error.e0344: "Invalid attachment message payload format." error.e0345: "Attachment message is missing assetId." error.e0346: "Attachment message is missing payload." error.e0347: "Default team queue mode requires at least one agent team." +error.e0348: "This file type cannot be uploaded for security reasons." error.profile.nicknameRequired: "Enter a nickname." error.profile.nicknameTooLong: "Nickname cannot exceed 100 characters." error.profile.avatarTooLong: "Avatar link cannot exceed 255 characters." diff --git a/internal/pkg/i18nx/locales/zh-CN.yml b/internal/pkg/i18nx/locales/zh-CN.yml index 6b3aa8eb..371264f0 100644 --- a/internal/pkg/i18nx/locales/zh-CN.yml +++ b/internal/pkg/i18nx/locales/zh-CN.yml @@ -345,6 +345,7 @@ error.e0344: "附件消息 payload 格式错误" error.e0345: "附件消息缺少 assetId" error.e0346: "附件消息缺少 payload" error.e0347: "默认客服组待接入池模式必须至少选择一个客服组" +error.e0348: "出于安全考虑,此文件类型不支持上传" error.profile.nicknameRequired: "请输入昵称" error.profile.nicknameTooLong: "昵称不能超过 100 个字符" error.profile.avatarTooLong: "头像链接不能超过 255 个字符" diff --git a/internal/services/asset_service.go b/internal/services/asset_service.go index b7596beb..aaf9a71c 100644 --- a/internal/services/asset_service.go +++ b/internal/services/asset_service.go @@ -61,12 +61,16 @@ func (s *assetService) OpenReader(asset *models.Asset) (io.ReadCloser, error) { } func (s *assetService) UploadBytes(data []byte, prefix, filename string, principal *dto.AuthPrincipal) (*models.Asset, error) { - src := bytes.NewReader(data) - return s.Upload(src, storage.UploadInfo{ + filename = storage.SanitizeFilename(filename) + mimeType, err := storage.ValidateUpload(filename, "", http.DetectContentType(data)) + if err != nil { + return nil, err + } + return s.Upload(bytes.NewReader(data), storage.UploadInfo{ Prefix: prefix, Filename: filename, FileSize: int64(len(data)), - MimeType: http.DetectContentType(data), + MimeType: mimeType, Principal: principal, }) } @@ -87,11 +91,22 @@ func (s *assetService) UploadFile(file *multipart.FileHeader, prefix string, pri } defer func() { _ = src.Close() }() + sniffed, err := storage.SniffContentType(src) + if err != nil { + return nil, err + } + + filename := storage.SanitizeFilename(file.Filename) + mimeType, err := storage.ValidateUpload(filename, file.Header.Get("Content-Type"), sniffed) + if err != nil { + return nil, err + } + return s.Upload(src, storage.UploadInfo{ Prefix: prefix, - Filename: file.Filename, + Filename: filename, FileSize: file.Size, - MimeType: file.Header.Get("Content-Type"), + MimeType: mimeType, Principal: principal, }) } diff --git a/internal/services/storage/safety.go b/internal/services/storage/safety.go new file mode 100644 index 00000000..7b9cc342 --- /dev/null +++ b/internal/services/storage/safety.go @@ -0,0 +1,218 @@ +package storage + +import ( + "io" + "mime" + "net/http" + "path" + "strings" + "unicode/utf8" + + "agent-desk/internal/pkg/errorsx" +) + +const ( + // sniffLimit is the number of leading bytes net/http.DetectContentType inspects. + sniffLimit = 512 + // maxFilenameLength keeps a stored filename inside the column that holds it. + maxFilenameLength = 255 +) + +// previewableExtensions are the only stored file types a browser may render +// inline. Every other extension is forced to download. +// +// Assets are served from this application's own origin, so an inline response is +// same-origin content: a browser that navigates to it runs whatever the bytes +// ask for, in the same security context as the dashboard and the public support +// widget. Forcing a download for anything we have not vetted keeps that true +// only for media, which cannot touch the DOM. +var previewableExtensions = map[string]bool{ + ".jpg": true, ".jpeg": true, ".png": true, ".gif": true, ".webp": true, + ".bmp": true, ".ico": true, ".avif": true, ".tif": true, ".tiff": true, + ".heic": true, ".heif": true, + ".mp4": true, ".m4v": true, ".webm": true, ".mov": true, ".ogv": true, + ".mp3": true, ".m4a": true, ".wav": true, ".ogg": true, ".oga": true, + ".aac": true, ".flac": true, + ".pdf": true, +} + +// blockedExtensions are rejected at upload time. Each one names a document a +// browser executes or lays out instead of displaying, so accepting it would let +// an unauthenticated visitor plant a script that runs against this origin the +// moment a staff member opens the attachment link. +var blockedExtensions = map[string]bool{ + ".html": true, ".htm": true, ".shtml": true, ".xhtml": true, ".xht": true, + ".svg": true, ".svgz": true, + ".xml": true, ".xsl": true, ".xslt": true, ".xsd": true, ".wsdl": true, + ".js": true, ".mjs": true, ".cjs": true, + ".swf": true, + ".php": true, ".phtml": true, ".php3": true, ".php4": true, ".php5": true, + ".asp": true, ".aspx": true, ".jsp": true, ".jspx": true, ".cfm": true, + ".hta": true, ".htc": true, ".htaccess": true, +} + +// blockedMediaTypes are the payloads a browser renders as an active document. +// They are rejected regardless of the extension the client chose, so renaming a +// page to .png does not smuggle it past the extension check. +var blockedMediaTypes = map[string]bool{ + "text/html": true, + "application/xhtml+xml": true, + "image/svg+xml": true, + "text/xml": true, + "application/xml": true, + "application/xslt+xml": true, + "text/javascript": true, + "application/javascript": true, + "application/x-javascript": true, +} + +// blockedUploadI18nKey is returned to the uploader whenever the file-safety +// policy rejects a payload. It deliberately does not name the offending type. +const blockedUploadI18nKey = "error.e0348" + +// mediaTypeExtensions maps a MIME type to the extension it should be stored +// under. mime.ExtensionsByType cannot be used here: on Windows it consults the +// registry, so the same type resolves to a different extension than in +// production, and it happily hands back an extension that describes an active +// document. Types that must never become a servable document are absent, which +// leaves the caller with the inert .bin fallback. +var mediaTypeExtensions = map[string]string{ + // images + "image/jpeg": ".jpg", "image/jfif": ".jpg", "image/pjpeg": ".jpg", + "image/png": ".png", "image/gif": ".gif", "image/webp": ".webp", + "image/bmp": ".bmp", "image/tiff": ".tiff", "image/avif": ".avif", + "image/heic": ".heic", "image/heif": ".heif", "image/x-icon": ".ico", + "image/vnd.microsoft.icon": ".ico", + // documents + "application/pdf": ".pdf", "text/plain": ".txt", "text/csv": ".csv", + "text/markdown": ".md", "application/json": ".json", "application/rtf": ".rtf", + "application/msword": ".doc", + "application/vnd.openxmlformats-officedocument.wordprocessingml.document": ".docx", + "application/vnd.ms-excel": ".xls", + "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet": ".xlsx", + "application/vnd.ms-powerpoint": ".ppt", + "application/vnd.openxmlformats-officedocument.presentationml.presentation": ".pptx", + "application/vnd.oasis.opendocument.text": ".odt", + "application/vnd.oasis.opendocument.spreadsheet": ".ods", + "application/vnd.oasis.opendocument.presentation": ".odp", + // archives + "application/zip": ".zip", "application/x-zip-compressed": ".zip", + "application/x-rar-compressed": ".rar", "application/vnd.rar": ".rar", + "application/x-7z-compressed": ".7z", "application/x-tar": ".tar", + "application/gzip": ".gz", "application/x-bzip2": ".bz2", "application/x-xz": ".xz", + // audio + "audio/mpeg": ".mp3", "audio/mp4": ".m4a", "audio/x-m4a": ".m4a", + "audio/wav": ".wav", "audio/x-wav": ".wav", "audio/ogg": ".ogg", + "audio/aac": ".aac", "audio/flac": ".flac", "audio/webm": ".webm", + // video + "video/mp4": ".mp4", "video/webm": ".webm", "video/quicktime": ".mov", + "video/x-msvideo": ".avi", "video/x-matroska": ".mkv", "video/ogg": ".ogv", + "video/mpeg": ".mpeg", +} + +// safeExtensionForMediaType returns the extension a payload of this type should +// be stored under, or "" when the type is not one this origin is willing to +// serve as a document. +func safeExtensionForMediaType(mediaType string) string { + return mediaTypeExtensions[strings.ToLower(strings.TrimSpace(mediaType))] +} + +// IsPreviewableExtension reports whether a stored file may be rendered inline. +func IsPreviewableExtension(ext string) bool { + return previewableExtensions[normalizeExt(ext)] +} + +// IsBlockedExtension reports whether an extension names a browser-active document. +func IsBlockedExtension(ext string) bool { + return blockedExtensions[normalizeExt(ext)] +} + +// IsBlockedMediaType reports whether a MIME type describes a browser-active document. +func IsBlockedMediaType(mediaType string) bool { + parsed, _, err := mime.ParseMediaType(strings.TrimSpace(mediaType)) + if err != nil || parsed == "" { + return false + } + return blockedMediaTypes[strings.ToLower(parsed)] +} + +// SanitizeFilename reduces a client-supplied name to a bare basename. Uploads +// arrive from browsers, mobile SDKs and channel webhooks, all of which are free +// to send a full path, control characters or nothing at all; the result is stored +// on the asset and echoed back into message payloads and download links. +func SanitizeFilename(name string) string { + name = strings.TrimSpace(name) + if name == "" { + return "" + } + name = path.Base(strings.ReplaceAll(name, "\\", "/")) + name = strings.Map(func(r rune) rune { + if r < 0x20 || r == 0x7f { + return -1 + } + return r + }, name) + name = strings.TrimSpace(strings.TrimRight(name, ". ")) + if len(name) > maxFilenameLength { + ext := path.Ext(name) + if len(ext) > maxFilenameLength { + ext = "" + } + stem := strings.TrimSuffix(name, ext) + // Cut on a rune boundary: a partial multibyte character would not survive + // the round trip through a utf8mb4 column. + limit := maxFilenameLength - len(ext) + for limit > 0 && !utf8.RuneStart(stem[limit]) { + limit-- + } + name = stem[:limit] + ext + } + return name +} + +// SniffContentType reports what the leading bytes of a seekable payload actually +// are, then rewinds so the caller can still stream the whole file to storage. +// The Content-Type a client declares is a claim, not evidence. +func SniffContentType(src io.ReadSeeker) (string, error) { + head := make([]byte, sniffLimit) + read, err := io.ReadFull(src, head) + if err != nil && err != io.EOF && err != io.ErrUnexpectedEOF { + return "", err + } + if _, err := src.Seek(0, io.SeekStart); err != nil { + return "", err + } + return http.DetectContentType(head[:read]), nil +} + +// ValidateUpload applies the file-safety policy to an upload and returns the MIME +// type that should be recorded on the asset. +// +// sniffed comes from SniffContentType and wins whenever it identified something +// specific; declared is what the client claimed and is only trusted for the +// formats net/http cannot recognise, such as HEIC and most Office documents. +func ValidateUpload(filename, declared, sniffed string) (string, error) { + if IsBlockedExtension(path.Ext(filename)) { + return "", errorsx.InvalidParamI18n(blockedUploadI18nKey) + } + if IsBlockedMediaType(sniffed) || IsBlockedMediaType(declared) { + return "", errorsx.InvalidParamI18n(blockedUploadI18nKey) + } + + mediaType, _, _ := mime.ParseMediaType(sniffed) + if mediaType != "" && mediaType != "application/octet-stream" && !strings.HasPrefix(mediaType, "text/") { + return sniffed, nil + } + if declared != "" { + return declared, nil + } + return sniffed, nil +} + +func normalizeExt(ext string) string { + ext = strings.ToLower(strings.TrimSpace(ext)) + if ext != "" && !strings.HasPrefix(ext, ".") { + ext = "." + ext + } + return ext +} diff --git a/internal/services/storage/safety_test.go b/internal/services/storage/safety_test.go new file mode 100644 index 00000000..78fafbe6 --- /dev/null +++ b/internal/services/storage/safety_test.go @@ -0,0 +1,268 @@ +package storage + +import ( + "bytes" + "io" + "net/http" + "strings" + "testing" + "unicode/utf8" +) + +var pngSignature = []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR") + +func TestValidateUploadBlocksBrowserActiveExtensions(t *testing.T) { + names := []string{ + "session-stealer.html", "page.HTM", "vector.svg", "sheet.xml", + "transform.xsl", "bundle.min.js", "app.mjs", "shell.php", + "report.asp", "index.jsp", "legacy.swf", "widget.htc", + } + for _, name := range names { + if _, err := ValidateUpload(name, "application/octet-stream", "application/octet-stream"); err == nil { + t.Errorf("ValidateUpload(%q) accepted a browser-active extension", name) + } + } +} + +func TestValidateUploadBlocksStoredXSSPayload(t *testing.T) { + payload := []byte(`
`) + sniffed, err := SniffContentType(bytes.NewReader(payload)) + if err != nil { + t.Fatalf("SniffContentType() error = %v", err) + } + if !IsBlockedMediaType(sniffed) { + t.Fatalf("expected the payload to sniff as a blocked media type, got %q", sniffed) + } + + if _, err := ValidateUpload("proof.html", "text/html", sniffed); err == nil { + t.Error("expected an honestly named HTML upload to be rejected") + } + // Renaming the payload is the actual attack: the extension looks like an + // image and the declared type agrees, so only the bytes can catch it. + if _, err := ValidateUpload("profile.png", "image/png", sniffed); err == nil { + t.Error("expected an HTML payload disguised as a PNG to be rejected") + } +} + +func TestValidateUploadBlocksSVGRegardlessOfDeclaredType(t *testing.T) { + payload := []byte(``) + sniffed, err := SniffContentType(bytes.NewReader(payload)) + if err != nil { + t.Fatalf("SniffContentType() error = %v", err) + } + if _, err := ValidateUpload("logo.svg", "image/svg+xml", sniffed); err == nil { + t.Error("expected an SVG upload to be rejected") + } + // An SVG that arrives without the .svg extension is still refused, because + // the declared type alone is enough to identify it as an active document. + if _, err := ValidateUpload("logo.png", "image/svg+xml", sniffed); err == nil { + t.Error("expected an SVG declared as an image to be rejected") + } +} + +func TestValidateUploadAcceptsOrdinarySupportFiles(t *testing.T) { + cases := []struct { + filename string + declared string + sniffed string + wantMime string + }{ + {"screenshot.png", "image/png", http.DetectContentType(pngSignature), "image/png"}, + // The client declared a generic type; the payload identified itself. + {"photo.jpg", "application/octet-stream", "image/jpeg", "image/jpeg"}, + // net/http cannot recognise HEIC, so the declared type has to stand in. + {"IMG_0001.heic", "image/heic", "application/octet-stream", "image/heic"}, + {"contract.pdf", "application/pdf", "application/pdf", "application/pdf"}, + {"export.csv", "text/csv", "text/plain; charset=utf-8", "text/csv"}, + {"notes.txt", "text/plain", "text/plain; charset=utf-8", "text/plain"}, + {"logs.zip", "application/zip", "application/zip", "application/zip"}, + {"build.apk", "application/vnd.android.package-archive", "application/octet-stream", "application/vnd.android.package-archive"}, + {"data.json", "application/json", "application/json", "application/json"}, + } + for _, tc := range cases { + got, err := ValidateUpload(tc.filename, tc.declared, tc.sniffed) + if err != nil { + t.Errorf("ValidateUpload(%q) error = %v", tc.filename, err) + continue + } + if got != tc.wantMime { + t.Errorf("ValidateUpload(%q) mime = %q want %q", tc.filename, got, tc.wantMime) + } + } +} + +func TestSanitizeFilename(t *testing.T) { + cases := []struct { + in string + want string + }{ + {"report.pdf", "report.pdf"}, + {"../../etc/passwd", "passwd"}, + {`C:\Users\victim\Desktop\evil.html`, "evil.html"}, + {"/absolute/path/notes.txt", "notes.txt"}, + {"tab\tand\nnewline.log", "tabandnewline.log"}, + {"....", ""}, + {"", ""}, + {" spaced name.pdf ", "spaced name.pdf"}, + } + for _, tc := range cases { + if got := SanitizeFilename(tc.in); got != tc.want { + t.Errorf("SanitizeFilename(%q) = %q want %q", tc.in, got, tc.want) + } + } + + long := strings.Repeat("a", 300) + ".pdf" + got := SanitizeFilename(long) + if len(got) > 255 { + t.Errorf("SanitizeFilename() len = %d want <= 255", len(got)) + } + if !strings.HasSuffix(got, ".pdf") { + t.Errorf("SanitizeFilename() = %q, expected the extension to survive truncation", got) + } + + // A multibyte stem must not be cut in the middle of a rune. + wide := strings.Repeat("附件", 100) + ".pdf" + got = SanitizeFilename(wide) + if len(got) > 255 { + t.Errorf("SanitizeFilename() len = %d want <= 255", len(got)) + } + if !utf8.ValidString(got) { + t.Errorf("SanitizeFilename() = %q is not valid UTF-8", got) + } + if !strings.HasSuffix(got, ".pdf") { + t.Errorf("SanitizeFilename() = %q, expected the extension to survive truncation", got) + } + + // An extension longer than the whole budget cannot be preserved. + got = SanitizeFilename("a." + strings.Repeat("b", 300)) + if len(got) > 255 { + t.Errorf("SanitizeFilename() len = %d want <= 255", len(got)) + } + if !utf8.ValidString(got) { + t.Errorf("SanitizeFilename() = %q is not valid UTF-8", got) + } + + // A name made entirely of dots leaves nothing worth keeping. + if got = SanitizeFilename(strings.Repeat(".", 300)); got != "" { + t.Errorf("SanitizeFilename() of a dot-only name = %q want empty", got) + } +} + +func TestSniffContentTypeRewindsTheReader(t *testing.T) { + src := bytes.NewReader(pngSignature) + got, err := SniffContentType(src) + if err != nil { + t.Fatalf("SniffContentType() error = %v", err) + } + if got != "image/png" { + t.Fatalf("SniffContentType() = %q want image/png", got) + } + + // The caller still streams the whole payload to storage afterwards, so the + // reader must be back at the start. + rest, err := io.ReadAll(src) + if err != nil { + t.Fatalf("ReadAll() error = %v", err) + } + if !bytes.Equal(rest, pngSignature) { + t.Fatalf("reader was not rewound: got %d bytes want %d", len(rest), len(pngSignature)) + } +} + +func TestSniffContentTypeHandlesShortAndEmptyPayloads(t *testing.T) { + short := bytes.NewReader([]byte("hi")) + got, err := SniffContentType(short) + if err != nil { + t.Fatalf("SniffContentType() error = %v", err) + } + if !strings.HasPrefix(got, "text/plain") { + t.Fatalf("SniffContentType() = %q want a text/plain result", got) + } + rest, err := io.ReadAll(short) + if err != nil { + t.Fatalf("ReadAll() error = %v", err) + } + if string(rest) != "hi" { + t.Fatalf("reader was not rewound: got %q", rest) + } + + empty := bytes.NewReader(nil) + if _, err := SniffContentType(empty); err != nil { + t.Fatalf("SniffContentType() on an empty payload error = %v", err) + } +} + +func TestGenerateStorageKeyNeverUsesBlockedExtension(t *testing.T) { + cases := []struct { + info UploadInfo + want string + }{ + {UploadInfo{Filename: "evil.html", MimeType: "text/html"}, ".bin"}, + {UploadInfo{Filename: "evil.svg", MimeType: "image/svg+xml"}, ".bin"}, + // The extension is gone but the MIME type still names an active + // document, so the derived extension must be refused too. + {UploadInfo{Filename: "noext", MimeType: "text/html"}, ".bin"}, + {UploadInfo{Filename: "notes.txt", MimeType: "text/plain"}, ".txt"}, + {UploadInfo{Filename: "photo.png", MimeType: "image/png"}, ".png"}, + {UploadInfo{Filename: "unknown", MimeType: "application/octet-stream"}, ".bin"}, + } + for _, tc := range cases { + _, key := GenerateStorageKey(tc.info) + if !strings.HasSuffix(key, tc.want) { + t.Errorf("GenerateStorageKey(%+v) = %q, expected it to end in %q", tc.info, key, tc.want) + } + if strings.Contains(key, "..") { + t.Errorf("GenerateStorageKey(%+v) = %q contains a traversal segment", tc.info, key) + } + } +} + +func TestGetExtByMimeTypeIsPlatformIndependent(t *testing.T) { + // mime.ExtensionsByType reads the Windows registry, where text/html resolved + // to ".ehtml" and slipped past the block list. The mapping has to come from + // our own table so a storage key is the same on every platform, and so a + // media type that describes an active document maps to nothing at all. + cases := map[string]string{ + "text/html": "", + "application/xhtml+xml": "", + "image/svg+xml": "", + "text/xml": "", + "application/xml": "", + "application/javascript": "", + "text/javascript": "", + "image/jpeg": ".jpg", + "image/jfif": ".jpg", + "image/pjpeg": ".jpg", + "image/png": ".png", + "application/pdf": ".pdf", + "text/csv": ".csv", + "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet": ".xlsx", + } + for mediaType, want := range cases { + if got := getExtByMimeType(mediaType); got != want { + t.Errorf("getExtByMimeType(%q) = %q want %q", mediaType, got, want) + } + } + + for mediaType := range blockedMediaTypes { + if got := getExtByMimeType(mediaType); got != "" { + t.Errorf("getExtByMimeType(%q) = %q, a blocked media type must not yield an extension", mediaType, got) + } + } +} + +func TestIsPreviewableExtension(t *testing.T) { + inline := []string{".png", ".JPG", ".jpeg", ".gif", ".webp", ".avif", ".mp4", ".webm", ".mp3", ".pdf"} + for _, ext := range inline { + if !IsPreviewableExtension(ext) { + t.Errorf("IsPreviewableExtension(%q) = false want true", ext) + } + } + + download := []string{".html", ".svg", ".xml", ".js", ".zip", ".txt", ".csv", ".docx", ".apk", ""} + for _, ext := range download { + if IsPreviewableExtension(ext) { + t.Errorf("IsPreviewableExtension(%q) = true want false", ext) + } + } +} diff --git a/internal/services/storage/utils.go b/internal/services/storage/utils.go index 0aabde2e..16b76999 100644 --- a/internal/services/storage/utils.go +++ b/internal/services/storage/utils.go @@ -30,9 +30,16 @@ func GenerateStorageKey(info UploadInfo) (assetID string, storageKey string) { } func getExt(info UploadInfo) string { - ext := strings.ToLower(filepath.Ext(strings.TrimSpace(info.Filename))) - if ext == "" { - ext = getExtByMimeType(info.MimeType) + ext := normalizeExt(filepath.Ext(strings.TrimSpace(info.Filename))) + if ext == "" || IsBlockedExtension(ext) { + ext = normalizeExt(getExtByMimeType(info.MimeType)) + } + if ext == "" || IsBlockedExtension(ext) { + // The stored extension decides the Content-Type this origin serves. A + // browser-active extension must never reach the key, and no extension at + // all would leave the server to sniff the payload and announce whatever + // it finds, so both cases settle on an inert binary type. + return ".bin" } return ext } @@ -47,21 +54,7 @@ func getExtByMimeType(mimeType string) string { return "" } - // 处理一些非标准的 MIME 类型 - switch mediaType { - case "image/jfif": - return ".jpg" - case "image/pjpeg": - return ".jpg" - case "image/jpeg": - return ".jpg" - default: - exts, _ := mime.ExtensionsByType(mediaType) - if len(exts) > 0 { - return exts[0] - } - } - return "" + return safeExtensionForMediaType(mediaType) } func normalizeAssetPrefix(prefix string) string { From f411969607c770ecf67fdb10069c310a4076123d Mon Sep 17 00:00:00 2001 From: JOY <5027251+JOY@users.noreply.github.com> Date: Sun, 13 Sep 2026 14:37:33 +0700 Subject: [PATCH 2/8] feat(channel): add native Discord channel integration Adds Discord as a sixth channel type, following the same shape as the Telegram and Zalo OA integrations: a client package, an inbound webhook service, an outbound service drained through the channel message outbox, a third-party handler, and the dashboard form fields. Inbound POST /api/third/discord/webhook[/:channel_id] Accepts both a bare interaction/webhook body and one nested under "message". Author, message id, channel id, guild id, attachments and embeds are all read from either level, so the same endpoint works with the payload shapes Discord and common bridges emit. Bot authors and empty messages are dropped. An attachment becomes a readable "[filename] url" line and an embed falls back to its description or title, so a customer who sends only an image or a link still produces a message an agent can act on. Messages map onto the existing conversation model through ExternalSourceDiscord, so a Discord user gets the same identity resolution, assignment, handoff and ticket behaviour as every other channel. Outbound Agent and AI replies are queued by EnqueueDiscordMessage and drained by cron every five seconds, with the same immediate-dispatch goroutine, backoff and retry ceiling as the other channels. Text and HTML go out as content; an image message is sent as an embed with the asset's signed URL; an attachment appends its signed URL to the text. The reply target is the discord_channel_id recorded on the last inbound message, falling back to a DM channel created from the customer's Discord user id, so a conversation that started in a server stays in that server. Guild scoping GuildID and ChannelScope are enforced on the way in. A bot invited to several servers answers only in the guild the channel names, and a channel set to dm_only ignores guild messages. Without this the two fields are stored configuration that silently does nothing, and an operator has no way to tell why a bot is answering somewhere it should not. Credentials A channel may carry its own bot token. When it does not, the deployment-wide token is used, resolved from discord.botToken in config.yaml or DISCORD_BOT_TOKEN in the environment. New DiscordConfig section plus AGENT_DESK_DISCORD_* and DISCORD_* bindings, documented in .env.example. Security The webhook secret is compared with crypto/subtle.ConstantTimeCompare. A byte-wise != leaks how much of the prefix matched through response timing, which matters here because the endpoint is unauthenticated by design. Verification is skipped only when no secret is configured on the channel, matching how the Telegram integration treats its webhook secret. Tests internal/discord/client_test.go client against an httptest stub internal/handlers/third/discord_handler_test.go handler routing and response internal/services/discord_inbound_service_test.go inbound to conversation to outbox, guild scope matrix (matching guild, other guild, DM under a guild-scoped channel, DM and guild message under dm_only, unscoped channel), and rejection of a wrong webhook secret internal/services/discord_integration_test.go full round trip Not included Discord's interactions endpoint signs requests with an Ed25519 signature over the raw body, verified against the application public key. This integration ingests through a shared-secret webhook instead, so the public key field is stored but not yet used to verify anything. Adding Ed25519 verification, and the OAuth flow that provisions a bot without a pasted token, are both worthwhile follow-ups. --- .env.example | 7 + internal/bootstrap/routes.go | 5 + internal/bootstrap/server.go | 1 + internal/discord/client.go | 139 +++++++++ internal/discord/client_test.go | 90 ++++++ internal/discord/types.go | 80 +++++ internal/handlers/third/discord_handler.go | 39 +++ .../handlers/third/discord_handler_test.go | 115 +++++++ internal/pkg/config/config.go | 19 ++ internal/pkg/dto/dto.go | 11 + internal/pkg/enums/external_identity.go | 2 + internal/pkg/enums/wxwork_kf.go | 1 + .../channel_message_outbox_service.go | 62 ++++ internal/services/channel_service.go | 47 ++- internal/services/cronx/cron.go | 4 + internal/services/discord_inbound_service.go | 176 +++++++++++ .../services/discord_inbound_service_test.go | 294 ++++++++++++++++++ internal/services/discord_integration_test.go | 169 ++++++++++ internal/services/discord_outbound_service.go | 231 ++++++++++++++ internal/services/message_service.go | 9 + .../dashboard/channels/_components/edit.tsx | 129 +++++++- .../(dashboard)/dashboard/channels/page.tsx | 8 + web/messages/en-US.json | 7 + web/messages/zh-CN.json | 7 + 24 files changed, 1638 insertions(+), 14 deletions(-) create mode 100644 internal/discord/client.go create mode 100644 internal/discord/client_test.go create mode 100644 internal/discord/types.go create mode 100644 internal/handlers/third/discord_handler.go create mode 100644 internal/handlers/third/discord_handler_test.go create mode 100644 internal/services/discord_inbound_service.go create mode 100644 internal/services/discord_inbound_service_test.go create mode 100644 internal/services/discord_integration_test.go create mode 100644 internal/services/discord_outbound_service.go diff --git a/.env.example b/.env.example index 77623625..12afb717 100644 --- a/.env.example +++ b/.env.example @@ -43,3 +43,10 @@ QDRANT_GRPC_PORT=6334 # Webhook & Organization Sync # ORG_SYNC_SECRET=your-webhook-hmac-secret + +# Discord Bot Channel +# Deployment-wide bot token, used when a Discord channel does not carry its own. +# DISCORD_BOT_TOKEN=your-discord-bot-token +# DISCORD_CLIENT_ID=your-discord-application-id +# DISCORD_CLIENT_SECRET=your-discord-client-secret +# DISCORD_PUBLIC_KEY=your-discord-application-public-key diff --git a/internal/bootstrap/routes.go b/internal/bootstrap/routes.go index e35dff8c..1482e841 100644 --- a/internal/bootstrap/routes.go +++ b/internal/bootstrap/routes.go @@ -435,3 +435,8 @@ func registerThirdZaloRoutes(group *gin.RouterGroup) { group.POST("/webhook", third.ZaloPostWebhook) group.POST("/webhook/:channel_id", third.ZaloPostWebhook) } + +func registerThirdDiscordRoutes(group *gin.RouterGroup) { + group.POST("/webhook", third.DiscordPostWebhook) + group.POST("/webhook/:channel_id", third.DiscordPostWebhook) +} diff --git a/internal/bootstrap/server.go b/internal/bootstrap/server.go index d5fc669a..228f22e1 100644 --- a/internal/bootstrap/server.go +++ b/internal/bootstrap/server.go @@ -197,6 +197,7 @@ func addRouter(app *gin.Engine) { registerThirdWechatRoutes(thirdGroup.Group("/wechat")) registerThirdTelegramRoutes(thirdGroup.Group("/telegram")) registerThirdZaloRoutes(thirdGroup.Group("/zalo")) + registerThirdDiscordRoutes(thirdGroup.Group("/discord")) } type spaShellRewrite struct { diff --git a/internal/discord/client.go b/internal/discord/client.go new file mode 100644 index 00000000..4045ae1c --- /dev/null +++ b/internal/discord/client.go @@ -0,0 +1,139 @@ +package discord + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" +) + +const defaultBaseURL = "https://discord.com/api/v10" + +type Client struct { + botToken string + baseURL string + httpClient *http.Client +} + +func NewClient(botToken string) *Client { + return &Client{ + botToken: strings.TrimSpace(botToken), + baseURL: defaultBaseURL, + httpClient: &http.Client{Timeout: 15 * time.Second}, + } +} + +func (c *Client) SetBaseURL(url string) { + if strings.TrimSpace(url) != "" { + c.baseURL = strings.TrimRight(strings.TrimSpace(url), "/") + } +} + +func (c *Client) GetMe(ctx context.Context) (*User, error) { + var user User + if err := c.doRequest(ctx, http.MethodGet, "/users/@me", nil, &user); err != nil { + return nil, err + } + return &user, nil +} + +func (c *Client) CreateDMChannel(ctx context.Context, recipientID string) (*Channel, error) { + if strings.TrimSpace(recipientID) == "" { + return nil, fmt.Errorf("recipient_id is required") + } + req := CreateDMRequest{RecipientID: strings.TrimSpace(recipientID)} + var channel Channel + if err := c.doRequest(ctx, http.MethodPost, "/users/@me/channels", req, &channel); err != nil { + return nil, err + } + return &channel, nil +} + +func (c *Client) SendMessage(ctx context.Context, channelID string, content string) (*Message, error) { + channelID = strings.TrimSpace(channelID) + if channelID == "" { + return nil, fmt.Errorf("channel_id is required") + } + if strings.TrimSpace(content) == "" { + return nil, fmt.Errorf("content is required") + } + + req := SendMessageRequest{Content: content} + var msg Message + endpoint := fmt.Sprintf("/channels/%s/messages", channelID) + if err := c.doRequest(ctx, http.MethodPost, endpoint, req, &msg); err != nil { + return nil, err + } + return &msg, nil +} + +func (c *Client) SendEmbedMessage(ctx context.Context, channelID string, content string, embeds []Embed) (*Message, error) { + channelID = strings.TrimSpace(channelID) + if channelID == "" { + return nil, fmt.Errorf("channel_id is required") + } + + req := SendMessageRequest{ + Content: content, + Embeds: embeds, + } + var msg Message + endpoint := fmt.Sprintf("/channels/%s/messages", channelID) + if err := c.doRequest(ctx, http.MethodPost, endpoint, req, &msg); err != nil { + return nil, err + } + return &msg, nil +} + +func (c *Client) doRequest(ctx context.Context, method, path string, payload any, result any) error { + if c.botToken == "" { + return fmt.Errorf("discord bot token is required") + } + + endpoint := fmt.Sprintf("%s%s", c.baseURL, path) + + var bodyReader io.Reader + if payload != nil { + bodyBytes, err := json.Marshal(payload) + if err != nil { + return fmt.Errorf("marshal discord request failed: %w", err) + } + bodyReader = bytes.NewBuffer(bodyBytes) + } + + req, err := http.NewRequestWithContext(ctx, method, endpoint, bodyReader) + if err != nil { + return fmt.Errorf("create discord request failed: %w", err) + } + + req.Header.Set("Authorization", "Bot "+c.botToken) + if payload != nil { + req.Header.Set("Content-Type", "application/json") + } + + res, err := c.httpClient.Do(req) + if err != nil { + return fmt.Errorf("discord http request failed: %w", err) + } + defer res.Body.Close() + + bodyBytes, err := io.ReadAll(res.Body) + if err != nil { + return fmt.Errorf("read discord response failed: %w", err) + } + + if res.StatusCode < 200 || res.StatusCode >= 300 { + return fmt.Errorf("discord api error (%d): %s", res.StatusCode, string(bodyBytes)) + } + + if result != nil { + if err := json.Unmarshal(bodyBytes, result); err != nil { + return fmt.Errorf("unmarshal discord response failed: %w (body: %s)", err, string(bodyBytes)) + } + } + return nil +} diff --git a/internal/discord/client_test.go b/internal/discord/client_test.go new file mode 100644 index 00000000..de1b7c85 --- /dev/null +++ b/internal/discord/client_test.go @@ -0,0 +1,90 @@ +package discord + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" +) + +func TestDiscordSendMessage(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Authorization") != "Bot test_token" { + t.Errorf("expected Bot test_token, got %s", r.Header.Get("Authorization")) + } + if r.URL.Path != "/channels/789/messages" { + t.Errorf("expected path /channels/789/messages, got %s", r.URL.Path) + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + w.Write([]byte(`{"id":"123456","channel_id":"789","content":"hello"}`)) + })) + defer server.Close() + + client := NewClient("test_token") + client.SetBaseURL(server.URL) + + resp, err := client.SendMessage(context.Background(), "789", "hello") + if err != nil { + t.Fatalf("SendMessage failed: %v", err) + } + if resp.ID != "123456" { + t.Errorf("expected ID 123456, got %s", resp.ID) + } +} + +func TestDiscordSendEmbedMessage(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Authorization") != "Bot test_token" { + t.Errorf("expected Bot test_token, got %s", r.Header.Get("Authorization")) + } + if r.URL.Path != "/channels/789/messages" { + t.Errorf("expected path /channels/789/messages, got %s", r.URL.Path) + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + w.Write([]byte(`{"id":"embed_123","channel_id":"789","content":"Check image"}`)) + })) + defer server.Close() + + client := NewClient("test_token") + client.SetBaseURL(server.URL) + + embed := Embed{ + Title: "Screenshot", + Image: &EmbedMedia{URL: "https://example.com/img.png"}, + } + resp, err := client.SendEmbedMessage(context.Background(), "789", "Check image", []Embed{embed}) + if err != nil { + t.Fatalf("SendEmbedMessage failed: %v", err) + } + if resp.ID != "embed_123" { + t.Errorf("expected ID embed_123, got %s", resp.ID) + } +} + +func TestDiscordCreateDMChannel(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Authorization") != "Bot test_token" { + t.Errorf("expected Bot test_token, got %s", r.Header.Get("Authorization")) + } + if r.URL.Path != "/users/@me/channels" { + t.Errorf("expected path /users/@me/channels, got %s", r.URL.Path) + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + w.Write([]byte(`{"id":"dm_chan_123","type":1}`)) + })) + defer server.Close() + + client := NewClient("test_token") + client.SetBaseURL(server.URL) + + resp, err := client.CreateDMChannel(context.Background(), "user_999") + if err != nil { + t.Fatalf("CreateDMChannel failed: %v", err) + } + if resp.ID != "dm_chan_123" { + t.Errorf("expected ID dm_chan_123, got %s", resp.ID) + } +} diff --git a/internal/discord/types.go b/internal/discord/types.go new file mode 100644 index 00000000..3363bebd --- /dev/null +++ b/internal/discord/types.go @@ -0,0 +1,80 @@ +package discord + +// User represents a Discord user. +type User struct { + ID string `json:"id"` + Username string `json:"username"` + Discriminator string `json:"discriminator,omitempty"` + GlobalName string `json:"global_name,omitempty"` + Avatar string `json:"avatar,omitempty"` + Bot bool `json:"bot,omitempty"` +} + +// Channel represents a Discord channel (Guild Text, DM, Thread, etc.). +type Channel struct { + ID string `json:"id"` + Type int `json:"type"` + GuildID string `json:"guild_id,omitempty"` + Name string `json:"name,omitempty"` +} + +// Attachment represents a file or image uploaded to Discord. +type Attachment struct { + ID string `json:"id"` + Filename string `json:"filename"` + URL string `json:"url"` + ProxyURL string `json:"proxy_url,omitempty"` + ContentType string `json:"content_type,omitempty"` + Size int64 `json:"size,omitempty"` +} + +// EmbedMedia represents an image/video/thumbnail inside an Embed. +type EmbedMedia struct { + URL string `json:"url"` +} + +// Embed represents a Discord rich embed object. +type Embed struct { + Title string `json:"title,omitempty"` + Description string `json:"description,omitempty"` + URL string `json:"url,omitempty"` + Color int `json:"color,omitempty"` + Image *EmbedMedia `json:"image,omitempty"` +} + +// Message represents a Discord message. +type Message struct { + ID string `json:"id"` + ChannelID string `json:"channel_id"` + GuildID string `json:"guild_id,omitempty"` + Author User `json:"author"` + Content string `json:"content"` + Timestamp string `json:"timestamp"` + Attachments []Attachment `json:"attachments,omitempty"` + Embeds []Embed `json:"embeds,omitempty"` +} + +// SendMessageRequest represents payload for Discord create message API. +type SendMessageRequest struct { + Content string `json:"content,omitempty"` + Embeds []Embed `json:"embeds,omitempty"` +} + +// CreateDMRequest represents payload for Discord create DM channel API. +type CreateDMRequest struct { + RecipientID string `json:"recipient_id"` +} + +// WebhookPayload represents an incoming message/event from Discord Gateway or Webhook. +type WebhookPayload struct { + ID string `json:"id,omitempty"` + Type int `json:"type,omitempty"` + GuildID string `json:"guild_id,omitempty"` + ChannelID string `json:"channel_id,omitempty"` + Author *User `json:"author,omitempty"` + Content string `json:"content,omitempty"` + Timestamp string `json:"timestamp,omitempty"` + Attachments []Attachment `json:"attachments,omitempty"` + Embeds []Embed `json:"embeds,omitempty"` + Message *Message `json:"message,omitempty"` +} diff --git a/internal/handlers/third/discord_handler.go b/internal/handlers/third/discord_handler.go new file mode 100644 index 00000000..e36ba526 --- /dev/null +++ b/internal/handlers/third/discord_handler.go @@ -0,0 +1,39 @@ +package third + +import ( + "bytes" + "io" + "net/http" + "strings" + + "agent-desk/internal/services" + + "github.com/gin-gonic/gin" +) + +// DiscordPostWebhook receives incoming Webhook events from Discord. +func DiscordPostWebhook(ctx *gin.Context) { + channelID := strings.TrimSpace(ctx.Param("channel_id")) + if channelID == "" { + channelID = strings.TrimSpace(ctx.Query("channel_id")) + } + + secretHeader := ctx.GetHeader("X-Discord-Secret-Token") + if secretHeader == "" { + secretHeader = ctx.GetHeader("X-Webhook-Secret") + } + + bodyBytes, err := io.ReadAll(ctx.Request.Body) + if err != nil { + ctx.JSON(http.StatusBadRequest, gin.H{"ok": false, "error": "failed to read body"}) + return + } + ctx.Request.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + + if err := services.DiscordInboundService.HandleWebhook(ctx.Request.Context(), channelID, secretHeader, bodyBytes); err != nil { + ctx.JSON(http.StatusOK, gin.H{"ok": false, "error": err.Error()}) + return + } + + ctx.JSON(http.StatusOK, gin.H{"ok": true}) +} diff --git a/internal/handlers/third/discord_handler_test.go b/internal/handlers/third/discord_handler_test.go new file mode 100644 index 00000000..e3b8e4f5 --- /dev/null +++ b/internal/handlers/third/discord_handler_test.go @@ -0,0 +1,115 @@ +package third + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "agent-desk/internal/models" + "agent-desk/internal/pkg/dto" + "agent-desk/internal/pkg/dto/request" + "agent-desk/internal/pkg/enums" + "agent-desk/internal/repositories" + "agent-desk/internal/services" + + "github.com/gin-gonic/gin" + "github.com/mlogclub/simple/sqls" +) + +func TestDiscordPostWebhook_Handler(t *testing.T) { + gin.SetMode(gin.TestMode) + db := setupThirdHandlerTestDB(t) + + now := time.Now() + agent := &models.AIAgent{ + Name: "Discord Agent", + ServiceMode: enums.IMConversationServiceModeAIFirst, + PublishedRevisionID: 1, + WelcomeMessage: "Hello Discord!", + Status: enums.StatusOk, + AuditFields: models.AuditFields{CreatedAt: now, UpdatedAt: now}, + } + _ = db.Create(agent) + + discordConfig, _ := json.Marshal(dto.DiscordChannelConfig{ + GuildID: "guild_999", + BotToken: "test_bot_token", + WebhookSecret: "secret_discord_123", + WelcomeMessage: "Welcome!", + }) + + operator := &dto.AuthPrincipal{UserID: 1, Username: "admin"} + channel, err := services.ChannelService.CreateChannel(request.CreateChannelRequest{ + Name: "Discord Community", + ChannelType: enums.ChannelTypeDiscord, + AIAgentID: agent.ID, + AIAgentRolloutPercent: 100, + ConfigJSON: string(discordConfig), + Status: int(enums.StatusOk), + }, operator) + if err != nil { + t.Fatalf("CreateChannel failed: %v", err) + } + + router := gin.New() + router.POST("/api/third/discord/webhook/:channel_id", DiscordPostWebhook) + router.POST("/api/third/discord/webhook", DiscordPostWebhook) + + payload := []byte(`{ + "id": "msg_001", + "channel_id": "ch_777", + "guild_id": "guild_999", + "content": "Need help with setup", + "author": { + "id": "user_456", + "username": "gamer_one", + "global_name": "Gamer One", + "bot": false + } + }`) + + // 1. Invalid secret + req, _ := http.NewRequest(http.MethodPost, "/api/third/discord/webhook/"+channel.ChannelID, bytes.NewBuffer(payload)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("X-Discord-Secret-Token", "wrong_secret") + + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("expected 200 OK wrapper, got: %d", rec.Code) + } + var resp map[string]any + _ = json.Unmarshal(rec.Body.Bytes(), &resp) + if resp["ok"] == true { + t.Fatalf("expected error for invalid secret token") + } + + // 2. Valid secret + req2, _ := http.NewRequest(http.MethodPost, "/api/third/discord/webhook/"+channel.ChannelID, bytes.NewBuffer(payload)) + req2.Header.Set("Content-Type", "application/json") + req2.Header.Set("X-Discord-Secret-Token", "secret_discord_123") + + rec2 := httptest.NewRecorder() + router.ServeHTTP(rec2, req2) + + if rec2.Code != http.StatusOK { + t.Fatalf("expected 200 OK, got %d", rec2.Code) + } + var resp2 map[string]any + _ = json.Unmarshal(rec2.Body.Bytes(), &resp2) + if resp2["ok"] != true { + t.Fatalf("expected ok: true, got: %+v", resp2) + } + + // Verify identity in DB + identity := repositories.CustomerIdentityRepository.FindOne(db, sqls.NewCnd(). + Eq("external_source", enums.ExternalSourceDiscord). + Eq("external_id", "user_456")) + if identity == nil { + t.Fatalf("expected customer identity for user_456") + } +} diff --git a/internal/pkg/config/config.go b/internal/pkg/config/config.go index 082de4dc..5925d94d 100644 --- a/internal/pkg/config/config.go +++ b/internal/pkg/config/config.go @@ -25,6 +25,7 @@ type Config struct { OIDC OIDCConfig `yaml:"oidc"` CustomerSession CustomerSessionConfig `yaml:"customerSession"` Webhook WebhookConfig `yaml:"webhook"` + Discord DiscordConfig `yaml:"discord"` } func (c Config) LanguageOrDefault() string { @@ -230,6 +231,16 @@ type WebhookConfig struct { DOSOrgSyncSecret string `yaml:"dosOrgSyncSecret"` } +// DiscordConfig holds deployment-wide Discord bot credentials. A channel may +// carry its own bot token, which takes precedence; these are the fallback for a +// single shared bot. +type DiscordConfig struct { + ClientID string `yaml:"clientId"` + ClientSecret string `yaml:"clientSecret"` + BotToken string `yaml:"botToken"` + PublicKey string `yaml:"publicKey"` +} + func Load(path string) (*Config, error) { loadDotEnv(path) @@ -312,6 +323,10 @@ func bindConfigDefaults(v *viper.Viper) { v.SetDefault("vectorDB.qdrant.host", "127.0.0.1") v.SetDefault("vectorDB.qdrant.grpcPort", 6334) v.SetDefault("mcp.enabled", true) + v.SetDefault("discord.clientId", "") + v.SetDefault("discord.clientSecret", "") + v.SetDefault("discord.botToken", "") + v.SetDefault("discord.publicKey", "") } func bindEnvironmentAliases(v *viper.Viper) { @@ -339,6 +354,10 @@ func bindEnvironmentAliases(v *viper.Viper) { _ = v.BindEnv("oidc.clientSecret", "AGENT_DESK_OIDC_CLIENTSECRET", "OIDC_CLIENT_SECRET", "CUSTOM_OAUTH_CLIENT_SECRET") _ = v.BindEnv("oidc.redirectUrl", "AGENT_DESK_OIDC_REDIRECTURL", "OIDC_REDIRECT_URL", "CUSTOM_OAUTH_REDIRECT_URI") _ = v.BindEnv("webhook.orgSyncSecret", "AGENT_DESK_WEBHOOK_ORGSYNCSECRET", "ORG_SYNC_SECRET", "WEBHOOK_SECRET") + _ = v.BindEnv("discord.clientId", "AGENT_DESK_DISCORD_CLIENTID", "DISCORD_CLIENT_ID") + _ = v.BindEnv("discord.clientSecret", "AGENT_DESK_DISCORD_CLIENTSECRET", "DISCORD_CLIENT_SECRET") + _ = v.BindEnv("discord.botToken", "AGENT_DESK_DISCORD_BOTTOKEN", "DISCORD_BOT_TOKEN") + _ = v.BindEnv("discord.publicKey", "AGENT_DESK_DISCORD_PUBLICKEY", "DISCORD_PUBLIC_KEY") } func normalizeLoadedConfig(cfg *Config) { diff --git a/internal/pkg/dto/dto.go b/internal/pkg/dto/dto.go index 276d03a1..5f456c03 100644 --- a/internal/pkg/dto/dto.go +++ b/internal/pkg/dto/dto.go @@ -49,3 +49,14 @@ type ZaloOAChannelConfig struct { WebhookSecret string `json:"webhookSecret,omitempty"` WelcomeMessage string `json:"welcomeMessage,omitempty"` } + +type DiscordChannelConfig struct { + GuildID string `json:"guildId,omitempty"` + GuildName string `json:"guildName,omitempty"` + ChannelScope string `json:"channelScope,omitempty"` // all | dm_only + BotToken string `json:"botToken,omitempty"` + ApplicationID string `json:"applicationId,omitempty"` + PublicKey string `json:"publicKey,omitempty"` + WebhookSecret string `json:"webhookSecret,omitempty"` + WelcomeMessage string `json:"welcomeMessage,omitempty"` +} diff --git a/internal/pkg/enums/external_identity.go b/internal/pkg/enums/external_identity.go index 8135fd56..e9fdf9c0 100644 --- a/internal/pkg/enums/external_identity.go +++ b/internal/pkg/enums/external_identity.go @@ -11,6 +11,7 @@ const ( ExternalSourceUser ExternalSource = "user" // 用户信息 ExternalSourceTelegram ExternalSource = "telegram" // Telegram ExternalSourceZaloOA ExternalSource = "zalo_oa" // Zalo OA + ExternalSourceDiscord ExternalSource = "discord" // Discord ) var externalSourceLabelMap = map[ExternalSource]string{ @@ -19,6 +20,7 @@ var externalSourceLabelMap = map[ExternalSource]string{ ExternalSourceUser: "用户", ExternalSourceTelegram: "Telegram", ExternalSourceZaloOA: "Zalo OA", + ExternalSourceDiscord: "Discord", } func GetExternalSourceLabel(v ExternalSource) string { diff --git a/internal/pkg/enums/wxwork_kf.go b/internal/pkg/enums/wxwork_kf.go index 825f7fa6..33d7fc9b 100644 --- a/internal/pkg/enums/wxwork_kf.go +++ b/internal/pkg/enums/wxwork_kf.go @@ -23,6 +23,7 @@ const ( ChannelTypeWxWorkKF = "wxwork_kf" ChannelTypeTelegram = "telegram" ChannelTypeZaloOA = "zalo_oa" + ChannelTypeDiscord = "discord" ) type WxWorkKFMessageSendStatus string diff --git a/internal/services/channel_message_outbox_service.go b/internal/services/channel_message_outbox_service.go index 12014241..8e40f10d 100644 --- a/internal/services/channel_message_outbox_service.go +++ b/internal/services/channel_message_outbox_service.go @@ -250,6 +250,68 @@ func (s *channelMessageOutboxService) EnqueueZaloOAMessage(conversation *models. return nil } +func (s *channelMessageOutboxService) EnqueueDiscordMessage(conversation *models.Conversation, message *models.Message) error { + if conversation == nil || message == nil { + return nil + } + channel := ChannelService.Get(conversation.ChannelID) + if channel == nil || channel.ChannelType != enums.ChannelTypeDiscord { + return nil + } + if message.SenderType != enums.IMSenderTypeAgent && message.SenderType != enums.IMSenderTypeAI { + return nil + } + if message.MessageType != enums.IMMessageTypeText && message.MessageType != enums.IMMessageTypeHTML && message.MessageType != enums.IMMessageTypeImage && message.MessageType != enums.IMMessageTypeAttachment { + return nil + } + if existing := s.GetByMessageID(enums.ChannelTypeDiscord, message.ID); existing != nil { + return nil + } + + payload, err := json.Marshal(map[string]any{ + "conversationId": conversation.ID, + "messageId": message.ID, + "messageType": message.MessageType, + "content": strings.TrimSpace(message.Content), + "payload": strings.TrimSpace(message.Payload), + "senderId": message.SenderID, + }) + if err != nil { + return err + } + + now := time.Now() + err = s.Create(&models.ChannelMessageOutbox{ + ChannelType: enums.ChannelTypeDiscord, + ConversationID: conversation.ID, + MessageID: message.ID, + Payload: string(payload), + SendStatus: string(enums.ChannelMessageOutboxStatusPending), + AuditFields: models.AuditFields{ + CreatedAt: now, + CreateUserID: message.UpdateUserID, + CreateUserName: message.UpdateUserName, + UpdatedAt: now, + UpdateUserID: message.UpdateUserID, + UpdateUserName: message.UpdateUserName, + }, + }) + if err != nil { + return err + } + + go func() { + defer func() { + if r := recover(); r != nil { + slog.Error("recovered from panic in discord outbound dispatch", "error", r) + } + }() + DiscordOutboundService.DispatchPendingOutbox() + }() + + return nil +} + func (s *channelMessageOutboxService) ListPending(channelType string, limit int) []models.ChannelMessageOutbox { if limit <= 0 { limit = 20 diff --git a/internal/services/channel_service.go b/internal/services/channel_service.go index 6a5399fc..92e2b640 100644 --- a/internal/services/channel_service.go +++ b/internal/services/channel_service.go @@ -333,6 +333,25 @@ func (s *channelService) ParseZaloOAChannelConfig(raw string) (*dto.ZaloOAChanne return cfg, nil } +func (s *channelService) ParseDiscordChannelConfig(raw string) (*dto.DiscordChannelConfig, error) { + raw = strings.TrimSpace(raw) + cfg := &dto.DiscordChannelConfig{} + if raw != "" { + if err := json.Unmarshal([]byte(raw), cfg); err != nil { + return nil, err + } + } + cfg.GuildID = strings.TrimSpace(cfg.GuildID) + cfg.GuildName = strings.TrimSpace(cfg.GuildName) + cfg.ChannelScope = strings.TrimSpace(cfg.ChannelScope) + cfg.BotToken = strings.TrimSpace(cfg.BotToken) + cfg.ApplicationID = strings.TrimSpace(cfg.ApplicationID) + cfg.PublicKey = strings.TrimSpace(cfg.PublicKey) + cfg.WebhookSecret = strings.TrimSpace(cfg.WebhookSecret) + cfg.WelcomeMessage = strings.TrimSpace(cfg.WelcomeMessage) + return cfg, nil +} + func (s *channelService) GetUserTokenSecret(channel *models.Channel) string { if channel == nil { return "" @@ -449,7 +468,7 @@ func (s *channelService) GetEnabledChannel(ctx *gin.Context) *models.Channel { func (s *channelService) buildChannelModel(id int64, req request.CreateChannelRequest) (*models.Channel, error) { channelType := strings.TrimSpace(req.ChannelType) - if channelType != enums.ChannelTypeWeb && channelType != enums.ChannelTypeWechatMP && channelType != enums.ChannelTypeWxWorkKF && channelType != enums.ChannelTypeTelegram && channelType != enums.ChannelTypeZaloOA { + if channelType != enums.ChannelTypeWeb && channelType != enums.ChannelTypeWechatMP && channelType != enums.ChannelTypeWxWorkKF && channelType != enums.ChannelTypeTelegram && channelType != enums.ChannelTypeZaloOA && channelType != enums.ChannelTypeDiscord { return nil, errorsx.InvalidParamI18n("error.e0250") } name := strings.TrimSpace(req.Name) @@ -596,6 +615,32 @@ func (s *channelService) buildChannelModel(id int64, req request.CreateChannelRe return nil, err } configJSON = string(configBytes) + case enums.ChannelTypeDiscord: + if channelID == "" { + channelID = strs.UUID() + } + if exists := s.Take("channel_id = ? AND status <> ? AND id <> ?", channelID, enums.StatusDeleted, id); exists != nil { + return nil, errorsx.InvalidParamI18n("error.e0248") + } + cfg, err := s.ParseDiscordChannelConfig(configJSON) + if err != nil { + return nil, errorsx.InvalidParam("invalid discord configuration") + } + // A channel may rely on the deployment-wide bot token instead of carrying + // its own, so the token is not required here the way Telegram's is. + if cfg.ChannelScope != "" && cfg.ChannelScope != "all" && cfg.ChannelScope != "dm_only" { + return nil, errorsx.InvalidParam("discord channelScope must be all or dm_only") + } + if cfg.WebhookSecret == "" { + if secret, err := generateUserTokenSecret(); err == nil { + cfg.WebhookSecret = secret + } + } + configBytes, err := json.Marshal(cfg) + if err != nil { + return nil, err + } + configJSON = string(configBytes) } return &models.Channel{ diff --git a/internal/services/cronx/cron.go b/internal/services/cronx/cron.go index 52ab00e2..e5c3af47 100644 --- a/internal/services/cronx/cron.go +++ b/internal/services/cronx/cron.go @@ -34,6 +34,10 @@ func Init() { if zaloCount > 0 { slog.Info("zalo oa outbox dispatched", "count", zaloCount) } + discordCount := services.DiscordOutboundService.DispatchPendingOutbox() + if discordCount > 0 { + slog.Info("discord outbox dispatched", "count", discordCount) + } }) c.Start() diff --git a/internal/services/discord_inbound_service.go b/internal/services/discord_inbound_service.go new file mode 100644 index 00000000..93f20eba --- /dev/null +++ b/internal/services/discord_inbound_service.go @@ -0,0 +1,176 @@ +package services + +import ( + "context" + "crypto/subtle" + "encoding/json" + "fmt" + "log/slog" + "strings" + + "agent-desk/internal/discord" + "agent-desk/internal/models" + "agent-desk/internal/pkg/enums" + "agent-desk/internal/pkg/errorsx" + "agent-desk/internal/pkg/openidentity" +) + +var DiscordInboundService = newDiscordInboundService() + +func newDiscordInboundService() *discordInboundService { + return &discordInboundService{} +} + +type discordInboundService struct{} + +// HandleWebhook processes an incoming webhook or gateway payload from Discord. +func (s *discordInboundService) HandleWebhook(ctx context.Context, channelID string, secretHeader string, rawPayload []byte) error { + channelID = strings.TrimSpace(channelID) + var channel *models.Channel + if channelID != "" { + channel = ChannelService.Take("channel_id = ? AND channel_type = ? AND status = ?", channelID, enums.ChannelTypeDiscord, enums.StatusOk) + } + if channel == nil { + channel = ChannelService.Take("channel_type = ? AND status = ?", enums.ChannelTypeDiscord, enums.StatusOk) + } + if channel == nil { + return errorsx.InvalidParam("discord channel not found or disabled") + } + + cfg, err := ChannelService.ParseDiscordChannelConfig(channel.ConfigJSON) + if err != nil || cfg == nil { + return errorsx.InvalidParam("discord channel config invalid") + } + + // Compared in constant time: a byte-wise != leaks how much of the prefix + // matched through response timing. + if cfg.WebhookSecret != "" && + subtle.ConstantTimeCompare([]byte(strings.TrimSpace(secretHeader)), []byte(cfg.WebhookSecret)) != 1 { + return errorsx.UnauthorizedI18n("error.auth.invalidSignature") + } + + var payload discord.WebhookPayload + if err := json.Unmarshal(rawPayload, &payload); err != nil { + return fmt.Errorf("unmarshal discord payload failed: %w", err) + } + + author := payload.Author + text := strings.TrimSpace(payload.Content) + msgID := payload.ID + targetChannelID := payload.ChannelID + guildID := payload.GuildID + attachments := payload.Attachments + embeds := payload.Embeds + + if payload.Message != nil { + if author == nil { + author = &payload.Message.Author + } + if text == "" { + text = strings.TrimSpace(payload.Message.Content) + } + if msgID == "" { + msgID = payload.Message.ID + } + if targetChannelID == "" { + targetChannelID = payload.Message.ChannelID + } + if guildID == "" { + guildID = payload.Message.GuildID + } + if len(attachments) == 0 && len(payload.Message.Attachments) > 0 { + attachments = payload.Message.Attachments + } + if len(embeds) == 0 && len(payload.Message.Embeds) > 0 { + embeds = payload.Message.Embeds + } + } + + if author == nil || author.Bot || strings.TrimSpace(author.ID) == "" { + return nil // Ignore bot messages or invalid authors + } + + // Honour the channel's guild scope. A bot can be invited to several servers, + // and without these checks GuildID and ChannelScope would be stored + // configuration that silently does nothing. + if cfg.GuildID != "" && guildID != cfg.GuildID { + slog.Debug("ignoring discord message from an out-of-scope guild", + "guild_id", guildID, + "channel", channel.ID, + ) + return nil + } + if cfg.ChannelScope == "dm_only" && guildID != "" { + slog.Debug("ignoring discord guild message, channel is dm_only", + "guild_id", guildID, + "channel", channel.ID, + ) + return nil + } + + if text == "" && len(attachments) > 0 { + firstAtt := attachments[0] + if firstAtt.Filename != "" { + text = fmt.Sprintf("[%s] %s", firstAtt.Filename, firstAtt.URL) + } else { + text = firstAtt.URL + } + } + + if text == "" && len(attachments) == 0 && len(embeds) == 0 { + return nil // Ignore empty messages + } + if text == "" && len(embeds) > 0 { + text = embeds[0].Description + if text == "" { + text = embeds[0].Title + } + } + + // 1. Resolve customer identity + externalID := author.ID + name := strings.TrimSpace(author.GlobalName) + if name == "" { + name = strings.TrimSpace(author.Username) + } + if name == "" { + name = fmt.Sprintf("Discord User %s", author.ID) + } + + externalUser := openidentity.ExternalUser{ + ExternalSource: enums.ExternalSourceDiscord, + ExternalID: externalID, + ExternalName: name, + } + + // 2. Create or match Conversation + conversation, err := ConversationService.Create(externalUser, channel.ID, channel.AIAgentID) + if err != nil { + return fmt.Errorf("create discord conversation failed: %w", err) + } + + // 3. Send message through MessageService + clientMsgID := fmt.Sprintf("discord_%s_%s", targetChannelID, msgID) + payloadMap := map[string]any{ + "discord_message_id": msgID, + "discord_channel_id": targetChannelID, + "discord_guild_id": guildID, + "discord_user_id": author.ID, + "discord_attachments": attachments, + } + payloadBytes, _ := json.Marshal(payloadMap) + + _, err = MessageService.SendCustomerMessage( + conversation.ID, + clientMsgID, + enums.IMMessageTypeText, + text, + string(payloadBytes), + externalUser, + ) + if err != nil { + return fmt.Errorf("send customer message failed: %w", err) + } + + return nil +} diff --git a/internal/services/discord_inbound_service_test.go b/internal/services/discord_inbound_service_test.go new file mode 100644 index 00000000..b6c31b24 --- /dev/null +++ b/internal/services/discord_inbound_service_test.go @@ -0,0 +1,294 @@ +package services + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "agent-desk/internal/models" + "agent-desk/internal/pkg/dto" + "agent-desk/internal/pkg/enums" + "agent-desk/internal/repositories" + + "github.com/glebarez/sqlite" + "github.com/mlogclub/simple/sqls" + "gorm.io/gorm" + "gorm.io/gorm/schema" +) + +func setupDiscordTestDB(t *testing.T) *gorm.DB { + t.Helper() + db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{ + NamingStrategy: schema.NamingStrategy{ + TablePrefix: "t_", + SingularTable: true, + }, + }) + if err != nil { + t.Fatalf("open sqlite db: %v", err) + } + if err := db.AutoMigrate( + &models.Channel{}, + &models.ChannelMessageOutbox{}, + &models.Customer{}, + &models.CustomerIdentity{}, + &models.CustomerContact{}, + &models.Conversation{}, + &models.ConversationParticipant{}, + &models.ConversationReadState{}, + &models.ConversationInterrupt{}, + &models.ConversationEventLog{}, + &models.Message{}, + &models.AIAgent{}, + &models.User{}, + &models.Role{}, + &models.UserRole{}, + &models.Permission{}, + &models.RolePermission{}, + &models.UserPermission{}, + ); err != nil { + t.Fatalf("migrate discord test tables: %v", err) + } + sqls.SetDB(db) + return db +} + +func TestDiscordInboundAndOutbound(t *testing.T) { + db := setupDiscordTestDB(t) + + mockServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + w.Write([]byte(`{"id":"out_msg_100","channel_id":"text_chan_1","content":"Agent reply"}`)) + })) + defer mockServer.Close() + + now := time.Now() + aiAgent := &models.AIAgent{ + Name: "Support AI", + Status: enums.StatusOk, + PublishedRevisionID: 1, + AuditFields: models.AuditFields{CreatedAt: now, UpdatedAt: now}, + } + if err := db.Create(aiAgent).Error; err != nil { + t.Fatalf("create ai agent: %v", err) + } + + discordConfig := dto.DiscordChannelConfig{ + GuildID: "guild_12345", + GuildName: "Test Guild", + BotToken: "discord_bot_token", + WebhookSecret: "test_secret", + } + cfgBytes, _ := json.Marshal(discordConfig) + + channel := &models.Channel{ + ChannelType: enums.ChannelTypeDiscord, + ChannelID: "discord_ch_1", + AIAgentID: aiAgent.ID, + AIAgentRolloutPercent: 100, + Name: "Community Support", + ConfigJSON: string(cfgBytes), + Status: enums.StatusOk, + AuditFields: models.AuditFields{CreatedAt: now, UpdatedAt: now}, + } + if err := db.Create(channel).Error; err != nil { + t.Fatalf("create discord channel: %v", err) + } + + payload := `{ + "id": "msg_999", + "channel_id": "text_chan_1", + "guild_id": "guild_12345", + "content": "", + "author": { + "id": "user_888", + "username": "gamer_joy", + "global_name": "Joy Le", + "bot": false + }, + "attachments": [ + { + "id": "att_1", + "filename": "screenshot.png", + "url": "https://cdn.discordapp.com/attachments/1/screenshot.png", + "content_type": "image/png", + "size": 10240 + } + ] + }` + + ctx := context.Background() + err := DiscordInboundService.HandleWebhook(ctx, channel.ChannelID, "test_secret", []byte(payload)) + if err != nil { + t.Fatalf("HandleWebhook failed: %v", err) + } + + // Verify customer identity + identity := repositories.CustomerIdentityRepository.FindOne(sqls.DB(), sqls.NewCnd(). + Eq("external_source", enums.ExternalSourceDiscord). + Eq("external_id", "user_888")) + if identity == nil { + t.Fatalf("expected customer identity to be created") + } + + // Verify conversation + conv := repositories.ConversationRepository.FindOne(sqls.DB(), sqls.NewCnd(). + Eq("customer_id", identity.CustomerID). + Eq("channel_id", channel.ID)) + if conv == nil { + t.Fatalf("expected conversation to be created") + } + + // Verify image message created from attachment + custMsg := repositories.MessageRepository.FindOne(sqls.DB(), sqls.NewCnd(). + Eq("conversation_id", conv.ID). + Eq("sender_type", enums.IMSenderTypeCustomer)) + if custMsg == nil { + t.Fatalf("expected customer message to be created") + } + + operator := &dto.AuthPrincipal{UserID: 1, Nickname: "Agent Joy"} + + // Test Outbound enqueue with Message + replyMsg, err := MessageService.SendAIMessage(conv.ID, aiAgent.ID, "ai_msg_1", enums.IMMessageTypeText, "Here is your response image: https://example.com/response_img.png", "", operator) + if err != nil { + t.Fatalf("MessageService.SendAIMessage failed: %v", err) + } + + outbox := ChannelMessageOutboxService.GetByMessageID(enums.ChannelTypeDiscord, replyMsg.ID) + if outbox == nil { + t.Fatalf("expected outbox entry for discord message") + } + if outbox.SendStatus != string(enums.ChannelMessageOutboxStatusPending) && outbox.SendStatus != string(enums.ChannelMessageOutboxStatusSending) && outbox.SendStatus != string(enums.ChannelMessageOutboxStatusFailed) { + t.Fatalf("unexpected outbox status: %s", outbox.SendStatus) + } +} + +func seedDiscordScopedChannel(t *testing.T, db *gorm.DB, channelID string, cfg dto.DiscordChannelConfig) *models.Channel { + t.Helper() + now := time.Now() + agent := &models.AIAgent{ + Name: "Support AI", + Status: enums.StatusOk, + PublishedRevisionID: 1, + AuditFields: models.AuditFields{CreatedAt: now, UpdatedAt: now}, + } + if err := db.Create(agent).Error; err != nil { + t.Fatalf("create ai agent: %v", err) + } + cfg.BotToken = "discord_bot_token" + cfg.WebhookSecret = "scope_secret" + cfgBytes, err := json.Marshal(cfg) + if err != nil { + t.Fatalf("marshal discord config: %v", err) + } + channel := &models.Channel{ + ChannelType: enums.ChannelTypeDiscord, + ChannelID: channelID, + AIAgentID: agent.ID, + AIAgentRolloutPercent: 100, + Name: "Discord " + channelID, + ConfigJSON: string(cfgBytes), + Status: enums.StatusOk, + AuditFields: models.AuditFields{CreatedAt: now, UpdatedAt: now}, + } + if err := db.Create(channel).Error; err != nil { + t.Fatalf("create discord channel: %v", err) + } + return channel +} + +func discordScopedPayload(t *testing.T, guildID, messageID string) []byte { + t.Helper() + body := map[string]any{ + "id": messageID, + "channel_id": "discord_text_chan", + "content": "hello from discord", + "author": map[string]any{"id": "user_scope", "username": "scoped_user", "bot": false}, + } + if guildID != "" { + body["guild_id"] = guildID + } + raw, err := json.Marshal(body) + if err != nil { + t.Fatalf("marshal payload: %v", err) + } + return raw +} + +// A bot can be invited to several servers, and GuildID and ChannelScope are +// stored on the channel. Messages outside that scope must not create a +// conversation, otherwise the stored scope is configuration that does nothing. +func TestDiscordInboundHonoursGuildScope(t *testing.T) { + db := setupDiscordTestDB(t) + + guildScoped := seedDiscordScopedChannel(t, db, "discord_guild_scoped", dto.DiscordChannelConfig{ + GuildID: "guild_in_scope", + GuildName: "In Scope", + }) + dmOnly := seedDiscordScopedChannel(t, db, "discord_dm_only", dto.DiscordChannelConfig{ + ChannelScope: "dm_only", + }) + unscoped := seedDiscordScopedChannel(t, db, "discord_unscoped", dto.DiscordChannelConfig{}) + + cases := []struct { + name string + channel *models.Channel + guildID string + wantStored bool + }{ + {"matching guild is accepted", guildScoped, "guild_in_scope", true}, + {"another guild is ignored", guildScoped, "guild_elsewhere", false}, + {"a dm is ignored by a guild scoped channel", guildScoped, "", false}, + {"a dm is accepted by a dm_only channel", dmOnly, "", true}, + {"a guild message is ignored by a dm_only channel", dmOnly, "guild_anywhere", false}, + {"an unscoped channel accepts any guild", unscoped, "guild_anywhere", true}, + {"an unscoped channel accepts a dm", unscoped, "", true}, + } + + for _, tc := range cases { + messageID := "scope_" + strings.ReplaceAll(tc.name, " ", "_") + payload := discordScopedPayload(t, tc.guildID, messageID) + + if err := DiscordInboundService.HandleWebhook(context.Background(), tc.channel.ChannelID, "scope_secret", payload); err != nil { + t.Fatalf("%s: HandleWebhook failed: %v", tc.name, err) + } + + var count int64 + if err := db.Table("t_message").Where("client_msg_id LIKE ?", "%"+messageID).Count(&count).Error; err != nil { + t.Fatalf("%s: count messages: %v", tc.name, err) + } + switch { + case tc.wantStored && count == 0: + t.Errorf("%s: expected the message to be stored", tc.name) + case !tc.wantStored && count != 0: + t.Errorf("%s: expected the message to be dropped, found %d", tc.name, count) + } + } +} + +// The webhook secret is compared in constant time, so a wrong secret of the same +// length must be rejected rather than accepted by a prefix match. +func TestDiscordInboundRejectsWrongWebhookSecret(t *testing.T) { + db := setupDiscordTestDB(t) + channel := seedDiscordScopedChannel(t, db, "discord_secret", dto.DiscordChannelConfig{}) + + payload := discordScopedPayload(t, "", "secret_msg_1") + err := DiscordInboundService.HandleWebhook(context.Background(), channel.ChannelID, "wrong_secret_value", payload) + if err == nil { + t.Fatalf("expected a wrong webhook secret to be rejected") + } + + var count int64 + if err := db.Table("t_message").Where("client_msg_id LIKE ?", "%secret_msg_1").Count(&count).Error; err != nil { + t.Fatalf("count messages: %v", err) + } + if count != 0 { + t.Fatalf("a rejected delivery stored %d messages", count) + } +} diff --git a/internal/services/discord_integration_test.go b/internal/services/discord_integration_test.go new file mode 100644 index 00000000..e470bd27 --- /dev/null +++ b/internal/services/discord_integration_test.go @@ -0,0 +1,169 @@ +package services + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "agent-desk/internal/models" + "agent-desk/internal/pkg/dto" + "agent-desk/internal/pkg/dto/request" + "agent-desk/internal/pkg/enums" + "agent-desk/internal/repositories" + + "github.com/glebarez/sqlite" + "github.com/mlogclub/simple/sqls" + "gorm.io/gorm" + "gorm.io/gorm/schema" +) + +func setupDiscordIntegrationTestDB(t *testing.T) *gorm.DB { + t.Helper() + db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{ + NamingStrategy: schema.NamingStrategy{ + TablePrefix: "t_", + SingularTable: true, + }, + }) + if err != nil { + t.Fatalf("open sqlite db: %v", err) + } + if err := db.AutoMigrate( + &models.Channel{}, + &models.ChannelMessageOutbox{}, + &models.Customer{}, + &models.CustomerIdentity{}, + &models.CustomerContact{}, + &models.Conversation{}, + &models.ConversationParticipant{}, + &models.ConversationReadState{}, + &models.ConversationInterrupt{}, + &models.ConversationEventLog{}, + &models.Message{}, + &models.AIAgent{}, + &models.AgentProfile{}, + &models.AgentTeam{}, + &models.User{}, + &models.Role{}, + &models.UserRole{}, + &models.Permission{}, + &models.RolePermission{}, + &models.UserPermission{}, + ); err != nil { + t.Fatalf("migrate discord integration test tables: %v", err) + } + sqls.SetDB(db) + return db +} + +func TestDiscordIntegrationFullFlow(t *testing.T) { + db := setupDiscordIntegrationTestDB(t) + + mockDiscordServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + w.Write([]byte(`{"id":"discord_msg_reply_999","channel_id":"ch_discord_general","content":"Cảm ơn bạn! Đội ngũ hỗ trợ sẽ kiểm tra ngay."}`)) + })) + defer mockDiscordServer.Close() + + now := time.Now() + // 1. Create AI Agent + agent := &models.AIAgent{ + Name: "Discord Support AI", + ServiceMode: enums.IMConversationServiceModeAIFirst, + PublishedRevisionID: 1, + WelcomeMessage: "Chào mừng đến với máy chủ Discord Crove Desk!", + Status: enums.StatusOk, + AuditFields: models.AuditFields{ + CreatedAt: now, + UpdatedAt: now, + }, + } + _ = db.Create(agent) + + // 2. Create Discord Channel + discordConfig, _ := json.Marshal(dto.DiscordChannelConfig{ + GuildID: "guild_987654321", + GuildName: "Crove Community Discord", + BotToken: "test-discord-bot-token-xyz", + WebhookSecret: "discord-secret-token-123", + WelcomeMessage: "Welcome to Discord Support!", + }) + + operator := &dto.AuthPrincipal{UserID: 1, Username: "admin"} + channel, err := ChannelService.CreateChannel(request.CreateChannelRequest{ + Name: "Crove Discord Support", + ChannelType: enums.ChannelTypeDiscord, + AIAgentID: agent.ID, + AIAgentRolloutPercent: 100, + ConfigJSON: string(discordConfig), + Status: int(enums.StatusOk), + }, operator) + if err != nil { + t.Fatalf("CreateChannel failed: %v", err) + } + + // 3. Simulate Inbound Discord Webhook / Gateway message from user + inboundPayload := []byte(`{ + "id": "msg_discord_user_001", + "channel_id": "ch_discord_general", + "guild_id": "guild_987654321", + "content": "Tôi muốn hỏi về cách cấu hình Custom Domain cho Email Channel trên Crove Desk", + "author": { + "id": "discord_uid_555", + "username": "gamer_joy", + "global_name": "Anh Le", + "bot": false + } + }`) + + ctx := context.Background() + err = DiscordInboundService.HandleWebhook(ctx, channel.ChannelID, "discord-secret-token-123", inboundPayload) + if err != nil { + t.Fatalf("DiscordInboundService.HandleWebhook failed: %v", err) + } + + // Verify Customer Identity + identity := repositories.CustomerIdentityRepository.FindOne(db, sqls.NewCnd(). + Eq("external_source", enums.ExternalSourceDiscord). + Eq("external_id", "discord_uid_555")) + if identity == nil { + t.Fatalf("expected customer identity for discord_uid_555") + } + + customer := repositories.CustomerRepository.Get(db, identity.CustomerID) + if customer == nil || customer.Name != "Anh Le" { + t.Fatalf("unexpected customer profile: %+v", customer) + } + + // Verify Conversation created + conv := repositories.ConversationRepository.FindOne(db, sqls.NewCnd().Eq("customer_id", customer.ID)) + if conv == nil || conv.ChannelID != channel.ID { + t.Fatalf("unexpected conversation: %+v", conv) + } + + // Verify Customer Message stored + msg := repositories.MessageRepository.FindOne(db, sqls.NewCnd(). + Eq("conversation_id", conv.ID). + Eq("sender_type", enums.IMSenderTypeCustomer)) + if msg == nil || msg.Content != "Tôi muốn hỏi về cách cấu hình Custom Domain cho Email Channel trên Crove Desk" { + t.Fatalf("unexpected stored customer message: %+v", msg) + } + + // 4. Simulate Agent / AI Reply and test Outbox Enqueue & Outbound Dispatch + replyMsg, err := MessageService.SendAIMessage(conv.ID, agent.ID, "ai_reply_001", enums.IMMessageTypeText, "Cảm ơn bạn! Đội ngũ hỗ trợ sẽ kiểm tra ngay.", "", operator) + if err != nil { + t.Fatalf("MessageService.SendAIMessage failed: %v", err) + } + + outbox := ChannelMessageOutboxService.GetByMessageID(enums.ChannelTypeDiscord, replyMsg.ID) + if outbox == nil { + t.Fatalf("expected discord outbox entry for AI message") + } + if outbox.ChannelType != enums.ChannelTypeDiscord { + t.Fatalf("expected outbox channel type 'discord', got '%s'", outbox.ChannelType) + } +} diff --git a/internal/services/discord_outbound_service.go b/internal/services/discord_outbound_service.go new file mode 100644 index 00000000..2024f8af --- /dev/null +++ b/internal/services/discord_outbound_service.go @@ -0,0 +1,231 @@ +package services + +import ( + "context" + "encoding/json" + "log/slog" + "strings" + "time" + + "agent-desk/internal/discord" + "agent-desk/internal/models" + "agent-desk/internal/pkg/config" + "agent-desk/internal/pkg/enums" + "agent-desk/internal/repositories" + "agent-desk/internal/services/storage" + "os" + + "github.com/mlogclub/simple/sqls" +) + +const ( + discordOutboxBatchSize = 20 + discordOutboxMaxRetry = 5 +) + +var DiscordOutboundService = newDiscordOutboundService() + +func newDiscordOutboundService() *discordOutboundService { + return &discordOutboundService{} +} + +type discordOutboundService struct{} + +func (s *discordOutboundService) DispatchPendingOutbox() int { + return s.doDispatchPendingOutbox(discordOutboxBatchSize) +} + +func (s *discordOutboundService) doDispatchPendingOutbox(limit int) int { + if limit <= 0 { + limit = discordOutboxBatchSize + } + items := ChannelMessageOutboxService.ListPending(enums.ChannelTypeDiscord, limit) + if len(items) == 0 { + return 0 + } + + successCount := 0 + for i := range items { + if err := s.processOutbox(items[i].ID); err != nil { + slog.Warn("process discord outbox failed", + "outbox_id", items[i].ID, + "error", err, + ) + continue + } + successCount++ + } + return successCount +} + +func (s *discordOutboundService) processOutbox(outboxID int64) error { + outbox := ChannelMessageOutboxService.Get(outboxID) + if outbox == nil { + return nil + } + if outbox.ChannelType != enums.ChannelTypeDiscord { + return nil + } + if outbox.SendStatus == string(enums.ChannelMessageOutboxStatusSent) { + return nil + } + if outbox.NextRetryAt != nil && outbox.NextRetryAt.After(time.Now()) { + return nil + } + + if err := ChannelMessageOutboxService.Updates(outbox.ID, map[string]any{ + "send_status": string(enums.ChannelMessageOutboxStatusSending), + "updated_at": time.Now(), + }); err != nil { + return err + } + + message := MessageService.Get(outbox.MessageID) + if message == nil { + return s.markOutboxFailed(outbox, "message not found") + } + conversation := ConversationService.Get(outbox.ConversationID) + if conversation == nil { + return s.markOutboxFailed(outbox, "conversation not found") + } + + channel := ChannelService.Get(conversation.ChannelID) + if channel == nil || channel.Status != enums.StatusOk { + return s.markOutboxFailed(outbox, "discord channel not found or disabled") + } + cfg, err := ChannelService.ParseDiscordChannelConfig(channel.ConfigJSON) + if err != nil { + return s.markOutboxFailed(outbox, "invalid discord channel config") + } + botToken := "" + if cfg != nil { + botToken = strings.TrimSpace(cfg.BotToken) + } + if botToken == "" { + botToken = strings.TrimSpace(config.Current().Discord.BotToken) + } + if botToken == "" { + botToken = strings.TrimSpace(os.Getenv("DISCORD_BOT_TOKEN")) + } + if botToken == "" { + return s.markOutboxFailed(outbox, "discord bot token not configured") + } + + // Resolve target Discord User ID and/or Channel ID + var discordUserID string + customerIdentity := repositories.CustomerIdentityRepository.FindOne(sqls.DB(), sqls.NewCnd(). + Eq("customer_id", conversation.CustomerID). + Eq("external_source", enums.ExternalSourceDiscord)) + if customerIdentity != nil { + discordUserID = strings.TrimSpace(customerIdentity.ExternalID) + } + + // Check if there is a discord_channel_id in last message payload + var targetChannelID string + lastCustomerMsg := repositories.MessageRepository.FindOne(sqls.DB(), sqls.NewCnd(). + Eq("conversation_id", conversation.ID). + Eq("sender_type", enums.IMSenderTypeCustomer). + Desc("id")) + if lastCustomerMsg != nil && lastCustomerMsg.Payload != "" { + var payloadMap map[string]any + if err := json.Unmarshal([]byte(lastCustomerMsg.Payload), &payloadMap); err == nil { + if chID, ok := payloadMap["discord_channel_id"].(string); ok && chID != "" { + targetChannelID = chID + } + } + } + + client := discord.NewClient(botToken) + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + if targetChannelID == "" { + if discordUserID == "" { + return s.markOutboxFailed(outbox, "unable to resolve discord target user or channel") + } + dmChannel, err := client.CreateDMChannel(ctx, discordUserID) + if err != nil { + return s.markOutboxFailed(outbox, "create discord dm channel failed: "+err.Error()) + } + targetChannelID = dmChannel.ID + } + + var sendErr error + if message.MessageType == enums.IMMessageTypeImage { + assetPayload, err := parseIMMessageAssetPayload(message.Payload) + var imageURL string + if err == nil && assetPayload != nil { + assetPayload = hydrateIMMessageAssetPayload(assetPayload) + if assetPayload.Provider != "" && assetPayload.StorageKey != "" { + if provider, err := storage.NewProvider(assetPayload.Provider); err == nil { + imageURL = provider.GetSignedURL(assetPayload.StorageKey) + } + } + } + if imageURL == "" && strings.HasPrefix(strings.TrimSpace(message.Content), "http") { + imageURL = strings.TrimSpace(message.Content) + } + + if imageURL != "" { + embed := discord.Embed{ + Title: "Image Attachment", + Image: &discord.EmbedMedia{URL: imageURL}, + } + _, sendErr = client.SendEmbedMessage(ctx, targetChannelID, message.Content, []discord.Embed{embed}) + } else { + _, sendErr = client.SendMessage(ctx, targetChannelID, message.Content) + } + } else if message.MessageType == enums.IMMessageTypeAttachment { + assetPayload, err := parseIMMessageAssetPayload(message.Payload) + var fileURL string + if err == nil && assetPayload != nil { + assetPayload = hydrateIMMessageAssetPayload(assetPayload) + if assetPayload.Provider != "" && assetPayload.StorageKey != "" { + if provider, err := storage.NewProvider(assetPayload.Provider); err == nil { + fileURL = provider.GetSignedURL(assetPayload.StorageKey) + } + } + } + textToSend := message.Content + if fileURL != "" { + if textToSend != "" { + textToSend += "\n" + fileURL + } else { + textToSend = fileURL + } + } + _, sendErr = client.SendMessage(ctx, targetChannelID, textToSend) + } else { + _, sendErr = client.SendMessage(ctx, targetChannelID, message.Content) + } + + if sendErr != nil { + return s.markOutboxFailed(outbox, sendErr.Error()) + } + + return ChannelMessageOutboxService.Updates(outbox.ID, map[string]any{ + "send_status": string(enums.ChannelMessageOutboxStatusSent), + "sent_at": time.Now(), + "updated_at": time.Now(), + }) +} + +func (s *discordOutboundService) markOutboxFailed(outbox *models.ChannelMessageOutbox, errMsg string) error { + if outbox == nil { + return nil + } + retryCount := outbox.RetryCount + 1 + status := string(enums.ChannelMessageOutboxStatusFailed) + if retryCount >= discordOutboxMaxRetry { + status = string(enums.ChannelMessageOutboxStatusIgnored) + } + nextRetryAt := time.Now().Add(time.Duration(retryCount*30) * time.Second) + + return ChannelMessageOutboxService.Updates(outbox.ID, map[string]any{ + "send_status": status, + "retry_count": retryCount, + "next_retry_at": &nextRetryAt, + "last_error": errMsg, + "updated_at": time.Now(), + }) +} diff --git a/internal/services/message_service.go b/internal/services/message_service.go index 48464472..b892b182 100644 --- a/internal/services/message_service.go +++ b/internal/services/message_service.go @@ -559,6 +559,15 @@ func (s *messageService) sendValidatedMessage(conversation *models.Conversation, ) } + // Discord 渠道消息入队,异步发送 + if enqueueErr := ChannelMessageOutboxService.EnqueueDiscordMessage(conversation, message); enqueueErr != nil { + slog.Error("enqueue discord outbox failed", + "conversation_id", conversation.ID, + "message_id", message.ID, + "error", enqueueErr, + ) + } + // 客户发送消息,触发AI回复 if senderType == enums.IMSenderTypeCustomer { if TriggerAIReplyAsyncHook != nil { diff --git a/web/app/(dashboard)/dashboard/channels/_components/edit.tsx b/web/app/(dashboard)/dashboard/channels/_components/edit.tsx index c5b96168..26ff48b7 100644 --- a/web/app/(dashboard)/dashboard/channels/_components/edit.tsx +++ b/web/app/(dashboard)/dashboard/channels/_components/edit.tsx @@ -73,6 +73,13 @@ type ZaloOAChannelConfig = { webhookSecret?: string } +type DiscordChannelConfig = { + guildId?: string + guildName?: string + botToken?: string + webhookSecret?: string +} + function getDefaultWebChannelConfig(t: Translate): Required+ {t("channel.slackSigningSecretHint")} +
+{t("auth.redirectingToOidc")}
+