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 { return { title: t("channel.defaultTitleWeb"), @@ -87,7 +94,7 @@ function getDefaultWebChannelConfig(t: Translate): Required { function createSchema(t: Translate) { return z .object({ - channelType: z.enum(["web", "wechat_mp", "wxwork_kf", "telegram", "zalo_oa"], t("channel.typeRequired")), + channelType: z.enum(["web", "wechat_mp", "wxwork_kf", "telegram", "zalo_oa", "discord"], t("channel.typeRequired")), aiAgentId: z.string().trim().regex(/^\d+$/, t("channel.agentRequired")), aiAgentRolloutPercent: z.coerce.number().int().min(1).max(100), name: z.string().trim().min(1, t("channel.nameRequired")), @@ -99,6 +106,9 @@ function createSchema(t: Translate) { zaloOaId: z.string().trim(), zaloAccessToken: z.string().trim(), zaloSecretKey: z.string().trim(), + discordGuildId: z.string().trim(), + discordGuildName: z.string().trim(), + discordBotToken: z.string().trim(), widgetTitle: z.string().trim(), widgetSubtitle: z.string().trim(), widgetThemeColor: z.string().trim(), @@ -133,7 +143,7 @@ function createSchema(t: Translate) { } type EditForm = { - channelType: "web" | "wechat_mp" | "wxwork_kf" | "telegram" | "zalo_oa" + channelType: "web" | "wechat_mp" | "wxwork_kf" | "telegram" | "zalo_oa" | "discord" aiAgentId: string aiAgentRolloutPercent: number name: string @@ -145,6 +155,9 @@ type EditForm = { zaloOaId: string zaloAccessToken: string zaloSecretKey: string + discordGuildId: string + discordGuildName: string + discordBotToken: string widgetTitle: string widgetSubtitle: string widgetThemeColor: string @@ -169,6 +182,9 @@ function createEmptyForm(t: Translate): EditForm { zaloOaId: "", zaloAccessToken: "", zaloSecretKey: "", + discordGuildId: "", + discordGuildName: "", + discordBotToken: "", widgetTitle: defaultWebChannelConfig.title, widgetSubtitle: defaultWebChannelConfig.subtitle, widgetThemeColor: defaultWebChannelConfig.themeColor, @@ -222,6 +238,21 @@ function parseZaloOAChannelConfig(configJson: string): ZaloOAChannelConfig { } } +function parseDiscordChannelConfig(configJson: string): DiscordChannelConfig { + if (!configJson.trim()) return {} + try { + const parsed = JSON.parse(configJson) as DiscordChannelConfig + return { + guildId: parsed.guildId?.trim() || "", + guildName: parsed.guildName?.trim() || "", + botToken: parsed.botToken?.trim() || "", + webhookSecret: parsed.webhookSecret?.trim() || "", + } + } catch { + return {} + } +} + function parseWebChannelConfig(configJson: string, t: Translate): Required { const defaultWebChannelConfig = getDefaultWebChannelConfig(t) if (!configJson.trim()) { @@ -276,6 +307,7 @@ function buildForm(item: AdminChannel | null, t: Translate): EditForm { const isWechatMP = item.channelType === "wechat_mp" const isTelegram = item.channelType === "telegram" const isZaloOA = item.channelType === "zalo_oa" + const isDiscord = item.channelType === "discord" const webConfig = parseWebChannelConfig(item.configJson, t) const wechatConfig = isWechatMP ? parseWechatMPChannelConfig(item.configJson, t) @@ -286,6 +318,9 @@ function buildForm(item: AdminChannel | null, t: Translate): EditForm { const zaloConfig = isZaloOA ? parseZaloOAChannelConfig(item.configJson) : null + const discordConfig = isDiscord + ? parseDiscordChannelConfig(item.configJson) + : null return { channelType: item.channelType === "wxwork_kf" @@ -294,20 +329,29 @@ function buildForm(item: AdminChannel | null, t: Translate): EditForm { ? "telegram" : item.channelType === "zalo_oa" ? "zalo_oa" - : item.channelType === "wechat_mp" - ? "wechat_mp" - : "web", + : item.channelType === "discord" + ? "discord" + : item.channelType === "wechat_mp" + ? "wechat_mp" + : "web", aiAgentId: item.aiAgentId > 0 ? String(item.aiAgentId) : "", aiAgentRolloutPercent: item.aiAgentRolloutPercent || 100, name: item.name, openKfId: parseOpenKfId(item.configJson), botToken: telegramConfig?.botToken ?? "", botUsername: telegramConfig?.botUsername ?? "", - webhookSecret: telegramConfig?.webhookSecret ?? zaloConfig?.webhookSecret ?? "", + webhookSecret: + telegramConfig?.webhookSecret ?? + zaloConfig?.webhookSecret ?? + discordConfig?.webhookSecret ?? + "", zaloAppId: zaloConfig?.appId ?? "", zaloOaId: zaloConfig?.oaId ?? "", zaloAccessToken: zaloConfig?.accessToken ?? "", zaloSecretKey: zaloConfig?.secretKey ?? "", + discordGuildId: discordConfig?.guildId ?? "", + discordGuildName: discordConfig?.guildName ?? "", + discordBotToken: discordConfig?.botToken ?? "", widgetTitle: wechatConfig?.title ?? webConfig.title, widgetSubtitle: wechatConfig?.subtitle ?? webConfig.subtitle, widgetThemeColor: wechatConfig?.themeColor ?? webConfig.themeColor, @@ -347,14 +391,21 @@ function buildPayload(form: EditForm, status: number, t: Translate): CreateAdmin secretKey: form.zaloSecretKey.trim(), webhookSecret: form.webhookSecret.trim(), }) - : channelType === "wechat_mp" - ? JSON.stringify(webLikeConfig) - : JSON.stringify({ - ...webLikeConfig, - position: form.widgetPosition || defaultWebChannelConfig.position, - width: form.widgetWidth.trim() || defaultWebChannelConfig.width, - userTokenSecret: form.userTokenSecret.trim(), + : channelType === "discord" + ? JSON.stringify({ + guildId: form.discordGuildId.trim(), + guildName: form.discordGuildName.trim(), + botToken: form.discordBotToken.trim(), + webhookSecret: form.webhookSecret.trim(), }) + : channelType === "wechat_mp" + ? JSON.stringify(webLikeConfig) + : JSON.stringify({ + ...webLikeConfig, + position: form.widgetPosition || defaultWebChannelConfig.position, + width: form.widgetWidth.trim() || defaultWebChannelConfig.width, + userTokenSecret: form.userTokenSecret.trim(), + }) return { channelType, aiAgentId: Number(form.aiAgentId), @@ -543,6 +594,7 @@ function ChannelFormBody({ const channelTypeOptions = [ { value: "web", label: t("channel.typeWeb") }, { value: "telegram", label: t("channel.typeTelegram") }, + { value: "discord", label: t("channel.typeDiscord") }, { value: "wechat_mp", label: t("channel.typeWechatMp") }, { value: "wxwork_kf", label: t("channel.typeWxworkKf") }, ] as const @@ -765,6 +817,57 @@ function ChannelFormBody({ ) : null} + {channelType === "discord" ? ( +
+
+ + {t("channel.discordGuildId")} + + + + + + + + {t("channel.discordGuildName")} + + + + + +
+ + + {t("channel.discordBotToken")} + + + + + + +
+
{t("channel.discordSetupTitle")}
+
{t("channel.discordSetupDescription")}
+
+ {t("channel.inboundWebhookUrl")}: /api/third/discord/webhook +
+
+
+ ) : null} + {channelType === "telegram" ? (
diff --git a/web/app/(dashboard)/dashboard/channels/page.tsx b/web/app/(dashboard)/dashboard/channels/page.tsx index aa8c06f7..4ecbc8de 100644 --- a/web/app/(dashboard)/dashboard/channels/page.tsx +++ b/web/app/(dashboard)/dashboard/channels/page.tsx @@ -2,6 +2,7 @@ import { Building2Icon, + Gamepad2Icon, MessagesSquareIcon, MessageSquareMoreIcon, SendIcon, @@ -39,6 +40,9 @@ function getChannelTypeLabel(channelType: string, t: (key: string) => string) { if (channelType === "zalo_oa") { return t("channel.typeZaloOa") } + if (channelType === "discord") { + return t("channel.typeDiscord") + } return t("channel.typeWeb") } @@ -62,6 +66,9 @@ function ChannelIcon({ channelType }: { channelType: string }) { if (channelType === "telegram" || channelType === "zalo_oa") { return } + if (channelType === "discord") { + return + } return } @@ -78,6 +85,7 @@ export default function DashboardChannelsPage() { { value: "all", label: t("channel.allTypes") }, { value: "web", label: t("channel.typeWeb") }, { value: "telegram", label: t("channel.typeTelegram") }, + { value: "discord", label: t("channel.typeDiscord") }, { value: "zalo_oa", label: t("channel.typeZaloOa") }, { value: "wechat_mp", label: t("channel.typeWechatMp") }, { value: "wxwork_kf", label: t("channel.typeWxworkKf") }, diff --git a/web/messages/en-US.json b/web/messages/en-US.json index c83bc000..c961bf00 100644 --- a/web/messages/en-US.json +++ b/web/messages/en-US.json @@ -624,6 +624,13 @@ "zaloAppId": "App ID", "zaloAutoConnectTitle": "Zalo Official Account Connection", "zaloAutoConnectDescription": "Connect your Zalo Official Account using the Access Token to automatically receive customer messages and dispatch replies.", + "typeDiscord": "Discord", + "discordGuildId": "Guild / Server ID", + "discordGuildName": "Guild / Server Name", + "discordBotToken": "Bot Token", + "discordSetupTitle": "Discord Bot Connection", + "discordSetupDescription": "Create a bot in the Discord Developer Portal, enable the Message Content privileged intent, invite it to your server, and point an interaction or bridge endpoint at the webhook URL below. Leave the bot token empty to use the deployment-wide DISCORD_BOT_TOKEN.", + "inboundWebhookUrl": "Inbound Webhook Endpoint", "loadFailed": "Could not load channels.", "created": "Channel created: {name}", "updated": "Channel updated: {name}", diff --git a/web/messages/zh-CN.json b/web/messages/zh-CN.json index 14e67ae3..5472359b 100644 --- a/web/messages/zh-CN.json +++ b/web/messages/zh-CN.json @@ -624,6 +624,13 @@ "zaloAppId": "App ID", "zaloAutoConnectTitle": "Zalo OA 渠道连接", "zaloAutoConnectDescription": "输入 Zalo OA 的 Access Token 即可自动双向同步客户会话与消息。", + "typeDiscord": "Discord", + "discordGuildId": "服务器 ID", + "discordGuildName": "服务器名称", + "discordBotToken": "Bot Token", + "discordSetupTitle": "Discord 机器人接入", + "discordSetupDescription": "在 Discord Developer Portal 创建机器人,开启 Message Content 特权意图,邀请机器人进入服务器,并将交互或转发端点指向下方的 Webhook 地址。Bot Token 留空则使用部署级 DISCORD_BOT_TOKEN。", + "inboundWebhookUrl": "入站 Webhook 地址", "loadFailed": "加载接入渠道失败", "created": "已创建接入渠道:{name}", "updated": "已更新接入渠道:{name}", From a1e3d7c7674175d98ad57d4d9b335140c0576cf3 Mon Sep 17 00:00:00 2001 From: JOY <5027251+JOY@users.noreply.github.com> Date: Sun, 13 Sep 2026 17:11:58 +0700 Subject: [PATCH 3/8] feat(channel): add native Slack channel integration Adds Slack as a sixth channel type, following the same shape as the Telegram and Zalo OA integrations: a client package, an inbound Events API service, an outbound service drained through the channel message outbox, a third-party handler, and the dashboard form fields. Inbound POST /api/third/slack/webhook[/:channel_id] Answers Slack's url_verification handshake before any channel lookup, because Slack sends it once while the endpoint is being configured and does not retry. event_callback messages map onto the existing conversation model through ExternalSourceSlack. Bot authors, bot_message subtypes and empty text are dropped so the bot cannot talk to itself in a loop. The channel is resolved by explicit id, then by team_id, then by the single enabled Slack channel. Outbound Agent and AI replies are queued by EnqueueSlackMessage and drained by cron every five seconds, with the same immediate-dispatch goroutine, backoff and retry ceiling as the other channels. Replies are posted with the inbound thread_ts so an answer stays in the customer's thread instead of landing in the workspace timeline. Text and HTML go out as content; an image is sent with its signed asset URL; an attachment appends the signed URL. Signing secret verification Two changes here, both because the endpoint is unauthenticated by design. A channel with a signing secret now rejects a delivery that arrives without the X-Slack-Signature and X-Slack-Request-Timestamp headers. Previously the check only ran when both headers were present, so omitting them bypassed verification entirely: anyone who learned the webhook URL could post into a workspace conversation and trigger paid AI replies. Slack always signs once a signing secret exists, so a missing header means the sender is not Slack. Request timestamps are now rejected outside a five minute window in either direction. Slack's own verification guide requires this, and without it a captured request replays indefinitely: the signature covers the timestamp and the body, but nothing in it expires. A channel with no signing secret still accepts deliveries, so this is not a breaking change for an existing installation. The dashboard field explains what leaving it empty costs. Credentials Everything is per channel: bot token, signing secret, app id, team id, team name and a default channel. No environment variables are added, so .env.example is untouched. The bot token is required when creating a channel, the same way Telegram's is, because without it the channel can receive events but can never reply. Tests internal/slack/client_test.go threaded reply carries thread_ts, a top-level post omits it, ok:false is surfaced as an error rather than treated as success, and input validation internal/handlers/third/slack_handler_test.go url_verification handshake, an unsigned event is rejected and creates no customer identity, a correctly signed event resolves the sender internal/services/slack_inbound_service_test.go inbound to conversation to outbox; rejection of missing headers, a signature without a timestamp, a timestamp without a signature, a wrong signature and a non-numeric timestamp; replayed timestamps at ten minutes, one hour and ten minutes in the future, each correctly signed for its own timestamp; fresh and four-minute-old signatures still accepted Not included Slack's OAuth install flow, so the bot token is pasted rather than obtained by authorization. Channel and DM history backfill is also not implemented: the channel starts receiving from the moment the Events API subscription is live. --- internal/bootstrap/routes.go | 5 + internal/bootstrap/server.go | 1 + internal/handlers/third/slack_handler.go | 43 +++ internal/handlers/third/slack_handler_test.go | 153 ++++++++ internal/pkg/dto/dto.go | 10 + 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 | 41 +- internal/services/cronx/cron.go | 4 + internal/services/message_service.go | 9 + internal/services/slack_inbound_service.go | 158 ++++++++ .../services/slack_inbound_service_test.go | 349 ++++++++++++++++++ internal/services/slack_outbound_service.go | 180 +++++++++ internal/slack/client.go | 109 ++++++ internal/slack/client_test.go | 120 ++++++ internal/slack/types.go | 37 ++ .../dashboard/channels/_components/edit.tsx | 188 +++++++++- .../(dashboard)/dashboard/channels/page.tsx | 8 + web/messages/en-US.json | 12 + web/messages/zh-CN.json | 12 + 21 files changed, 1491 insertions(+), 13 deletions(-) create mode 100644 internal/handlers/third/slack_handler.go create mode 100644 internal/handlers/third/slack_handler_test.go create mode 100644 internal/services/slack_inbound_service.go create mode 100644 internal/services/slack_inbound_service_test.go create mode 100644 internal/services/slack_outbound_service.go create mode 100644 internal/slack/client.go create mode 100644 internal/slack/client_test.go create mode 100644 internal/slack/types.go diff --git a/internal/bootstrap/routes.go b/internal/bootstrap/routes.go index e35dff8c..b940b353 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 registerThirdSlackRoutes(group *gin.RouterGroup) { + group.POST("/webhook", third.SlackPostWebhook) + group.POST("/webhook/:channel_id", third.SlackPostWebhook) +} diff --git a/internal/bootstrap/server.go b/internal/bootstrap/server.go index d5fc669a..efc58a8a 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")) + registerThirdSlackRoutes(thirdGroup.Group("/slack")) } type spaShellRewrite struct { diff --git a/internal/handlers/third/slack_handler.go b/internal/handlers/third/slack_handler.go new file mode 100644 index 00000000..215c996d --- /dev/null +++ b/internal/handlers/third/slack_handler.go @@ -0,0 +1,43 @@ +package third + +import ( + "bytes" + "io" + "net/http" + "strings" + + "agent-desk/internal/services" + + "github.com/gin-gonic/gin" +) + +// SlackPostWebhook receives incoming Events API payloads from Slack. +func SlackPostWebhook(ctx *gin.Context) { + channelID := strings.TrimSpace(ctx.Param("channel_id")) + if channelID == "" { + channelID = strings.TrimSpace(ctx.Query("channel_id")) + } + + timestampHeader := ctx.GetHeader("X-Slack-Request-Timestamp") + signatureHeader := ctx.GetHeader("X-Slack-Signature") + + 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)) + + challenge, err := services.SlackInboundService.HandleWebhook(ctx.Request.Context(), channelID, timestampHeader, signatureHeader, bodyBytes) + if err != nil { + ctx.JSON(http.StatusOK, gin.H{"ok": false, "error": err.Error()}) + return + } + + if challenge != nil { + ctx.JSON(http.StatusOK, gin.H{"challenge": *challenge}) + return + } + + ctx.JSON(http.StatusOK, gin.H{"ok": true}) +} diff --git a/internal/handlers/third/slack_handler_test.go b/internal/handlers/third/slack_handler_test.go new file mode 100644 index 00000000..9e7b82e2 --- /dev/null +++ b/internal/handlers/third/slack_handler_test.go @@ -0,0 +1,153 @@ +package third + +import ( + "bytes" + "crypto/hmac" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "net/http" + "net/http/httptest" + "strconv" + "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" +) + +const slackHandlerTestSigningSecret = "test_signing_secret" + +// signSlackHandlerPayload builds the X-Slack-Request-Timestamp and +// X-Slack-Signature headers Slack sends for this body right now. +func signSlackHandlerPayload(payload []byte) (string, string) { + timestamp := strconv.FormatInt(time.Now().Unix(), 10) + mac := hmac.New(sha256.New, []byte(slackHandlerTestSigningSecret)) + mac.Write([]byte("v0:" + timestamp + ":" + string(payload))) + return timestamp, "v0=" + hex.EncodeToString(mac.Sum(nil)) +} + +func TestSlackWebhook_Handler(t *testing.T) { + gin.SetMode(gin.TestMode) + db := setupThirdHandlerTestDB(t) + + now := time.Now() + agent := &models.AIAgent{ + Name: "Slack Agent", + ServiceMode: enums.IMConversationServiceModeAIFirst, + PublishedRevisionID: 1, + WelcomeMessage: "Hello Slack User!", + Status: enums.StatusOk, + AuditFields: models.AuditFields{CreatedAt: now, UpdatedAt: now}, + } + if err := db.Create(agent).Error; err != nil { + t.Fatalf("create ai agent: %v", err) + } + + slackConfig, err := json.Marshal(dto.SlackChannelConfig{ + BotToken: "xoxb-test-token", + SigningSecret: slackHandlerTestSigningSecret, + TeamID: "T_SLACK_100", + DefaultChannel: "C_GENERAL", + }) + if err != nil { + t.Fatalf("marshal slack config: %v", err) + } + + operator := &dto.AuthPrincipal{UserID: 1, Username: "admin"} + channel, err := services.ChannelService.CreateChannel(request.CreateChannelRequest{ + Name: "Slack Channel", + ChannelType: enums.ChannelTypeSlack, + AIAgentID: agent.ID, + AIAgentRolloutPercent: 100, + ConfigJSON: string(slackConfig), + Status: int(enums.StatusOk), + }, operator) + if err != nil { + t.Fatalf("CreateChannel failed: %v", err) + } + + router := gin.New() + router.POST("/api/third/slack/webhook/:channel_id", SlackPostWebhook) + router.POST("/api/third/slack/webhook", SlackPostWebhook) + + post := func(path string, payload []byte, timestamp, signature string) *httptest.ResponseRecorder { + req, _ := http.NewRequest(http.MethodPost, path, bytes.NewBuffer(payload)) + req.Header.Set("Content-Type", "application/json") + if timestamp != "" { + req.Header.Set("X-Slack-Request-Timestamp", timestamp) + } + if signature != "" { + req.Header.Set("X-Slack-Signature", signature) + } + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + return rec + } + + webhookPath := "/api/third/slack/webhook/" + channel.ChannelID + + // 1. The url_verification handshake is answered without a signature, because + // Slack sends it once while the endpoint is being configured. + challengePayload := []byte(`{ + "token": "token123", + "challenge": "slack_challenge_string_999", + "type": "url_verification" + }`) + recChallenge := post(webhookPath, challengePayload, "", "") + if recChallenge.Code != http.StatusOK { + t.Fatalf("expected 200 OK for challenge, got: %d", recChallenge.Code) + } + var challengeResp map[string]any + if err := json.Unmarshal(recChallenge.Body.Bytes(), &challengeResp); err != nil { + t.Fatalf("unmarshal challenge response: %v", err) + } + if challengeResp["challenge"] != "slack_challenge_string_999" { + t.Fatalf("expected challenge in body, got: %+v", challengeResp) + } + + // 2. An event the channel's signing secret did not produce is rejected, and + // must not create a customer identity. + eventPayload := []byte(`{ + "token": "token123", + "team_id": "T_SLACK_100", + "type": "event_callback", + "event": { + "type": "message", + "user": "U_USER_777", + "text": "Hello support team on Slack!", + "ts": "1725260000.000100", + "channel": "C_GENERAL" + } + }`) + recUnsigned := post(webhookPath, eventPayload, "", "") + if recUnsigned.Code != http.StatusOK { + t.Fatalf("expected the handler to answer 200 with an error body, got: %d", recUnsigned.Code) + } + if repositories.CustomerIdentityRepository.FindOne(db, sqls.NewCnd(). + Eq("external_source", enums.ExternalSourceSlack). + Eq("external_id", "U_USER_777")) != nil { + t.Fatalf("an unsigned event created a customer identity") + } + + // 3. A correctly signed event is accepted and resolves the sender. + timestamp, signature := signSlackHandlerPayload(eventPayload) + recEvent := post(webhookPath, eventPayload, timestamp, signature) + if recEvent.Code != http.StatusOK { + t.Fatalf("expected 200 OK for event, got: %d", recEvent.Code) + } + + identity := repositories.CustomerIdentityRepository.FindOne(db, sqls.NewCnd(). + Eq("external_source", enums.ExternalSourceSlack). + Eq("external_id", "U_USER_777")) + if identity == nil { + t.Fatalf("expected customer identity for U_USER_777") + } +} diff --git a/internal/pkg/dto/dto.go b/internal/pkg/dto/dto.go index 276d03a1..4124a243 100644 --- a/internal/pkg/dto/dto.go +++ b/internal/pkg/dto/dto.go @@ -49,3 +49,13 @@ type ZaloOAChannelConfig struct { WebhookSecret string `json:"webhookSecret,omitempty"` WelcomeMessage string `json:"welcomeMessage,omitempty"` } + +type SlackChannelConfig struct { + BotToken string `json:"botToken,omitempty"` // xoxb-... Bot Token + SigningSecret string `json:"signingSecret,omitempty"` // Slack Signing Secret + AppID string `json:"appId,omitempty"` // Slack App ID + TeamID string `json:"teamId,omitempty"` // Slack Workspace Team ID + TeamName string `json:"teamName,omitempty"` // Slack Workspace Team Name + DefaultChannel string `json:"defaultChannel,omitempty"` // Default channel to post + WelcomeMessage string `json:"welcomeMessage,omitempty"` +} diff --git a/internal/pkg/enums/external_identity.go b/internal/pkg/enums/external_identity.go index 8135fd56..ad73b6b8 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 + ExternalSourceSlack ExternalSource = "slack" // Slack ) var externalSourceLabelMap = map[ExternalSource]string{ @@ -19,6 +20,7 @@ var externalSourceLabelMap = map[ExternalSource]string{ ExternalSourceUser: "用户", ExternalSourceTelegram: "Telegram", ExternalSourceZaloOA: "Zalo OA", + ExternalSourceSlack: "Slack", } func GetExternalSourceLabel(v ExternalSource) string { diff --git a/internal/pkg/enums/wxwork_kf.go b/internal/pkg/enums/wxwork_kf.go index 825f7fa6..7483bb20 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" + ChannelTypeSlack = "slack" ) type WxWorkKFMessageSendStatus string diff --git a/internal/services/channel_message_outbox_service.go b/internal/services/channel_message_outbox_service.go index 12014241..3e1be941 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) EnqueueSlackMessage(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.ChannelTypeSlack { + 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.ChannelTypeSlack, 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.ChannelTypeSlack, + 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 slack outbound dispatch", "error", r) + } + }() + SlackOutboundService.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..2f282aab 100644 --- a/internal/services/channel_service.go +++ b/internal/services/channel_service.go @@ -333,6 +333,24 @@ func (s *channelService) ParseZaloOAChannelConfig(raw string) (*dto.ZaloOAChanne return cfg, nil } +func (s *channelService) ParseSlackChannelConfig(raw string) (*dto.SlackChannelConfig, error) { + raw = strings.TrimSpace(raw) + cfg := &dto.SlackChannelConfig{} + if raw != "" { + if err := json.Unmarshal([]byte(raw), cfg); err != nil { + return nil, err + } + } + cfg.BotToken = strings.TrimSpace(cfg.BotToken) + cfg.SigningSecret = strings.TrimSpace(cfg.SigningSecret) + cfg.AppID = strings.TrimSpace(cfg.AppID) + cfg.TeamID = strings.TrimSpace(cfg.TeamID) + cfg.TeamName = strings.TrimSpace(cfg.TeamName) + cfg.DefaultChannel = strings.TrimSpace(cfg.DefaultChannel) + cfg.WelcomeMessage = strings.TrimSpace(cfg.WelcomeMessage) + return cfg, nil +} + func (s *channelService) GetUserTokenSecret(channel *models.Channel) string { if channel == nil { return "" @@ -449,7 +467,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.ChannelTypeSlack { return nil, errorsx.InvalidParamI18n("error.e0250") } name := strings.TrimSpace(req.Name) @@ -596,6 +614,27 @@ func (s *channelService) buildChannelModel(id int64, req request.CreateChannelRe return nil, err } configJSON = string(configBytes) + case enums.ChannelTypeSlack: + 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.ParseSlackChannelConfig(configJSON) + if err != nil { + return nil, errorsx.InvalidParam("invalid slack configuration") + } + // Without a bot token the channel can receive events but can never reply, + // so it is required the same way Telegram's bot token is. + if cfg == nil || cfg.BotToken == "" { + return nil, errorsx.InvalidParam("slack botToken is required") + } + 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..0c73c767 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) } + slackCount := services.SlackOutboundService.DispatchPendingOutbox() + if slackCount > 0 { + slog.Info("slack outbox dispatched", "count", slackCount) + } }) c.Start() diff --git a/internal/services/message_service.go b/internal/services/message_service.go index 48464472..f6a03b45 100644 --- a/internal/services/message_service.go +++ b/internal/services/message_service.go @@ -559,6 +559,15 @@ func (s *messageService) sendValidatedMessage(conversation *models.Conversation, ) } + // Slack 渠道消息入队,异步发送 + if enqueueErr := ChannelMessageOutboxService.EnqueueSlackMessage(conversation, message); enqueueErr != nil { + slog.Error("enqueue slack outbox failed", + "conversation_id", conversation.ID, + "message_id", message.ID, + "error", enqueueErr, + ) + } + // 客户发送消息,触发AI回复 if senderType == enums.IMSenderTypeCustomer { if TriggerAIReplyAsyncHook != nil { diff --git a/internal/services/slack_inbound_service.go b/internal/services/slack_inbound_service.go new file mode 100644 index 00000000..f5b444fa --- /dev/null +++ b/internal/services/slack_inbound_service.go @@ -0,0 +1,158 @@ +package services + +import ( + "context" + "crypto/hmac" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "strconv" + "strings" + "time" + + "agent-desk/internal/models" + "agent-desk/internal/pkg/enums" + "agent-desk/internal/pkg/errorsx" + "agent-desk/internal/pkg/openidentity" + "agent-desk/internal/slack" +) + +var SlackInboundService = newSlackInboundService() + +func newSlackInboundService() *slackInboundService { + return &slackInboundService{} +} + +type slackInboundService struct{} + +// HandleWebhook processes an incoming Events API event from Slack. +func (s *slackInboundService) HandleWebhook(ctx context.Context, channelID string, timestampHeader, signatureHeader string, rawPayload []byte) (*string, error) { + var event slack.EventCallback + if err := json.Unmarshal(rawPayload, &event); err != nil { + return nil, fmt.Errorf("unmarshal slack event failed: %w", err) + } + + // 1. URL Verification Challenge + if event.Type == "url_verification" { + return &event.Challenge, nil + } + + if event.Type != "event_callback" || event.Event == nil { + return nil, nil // Ignore non-message callbacks + } + + teamID := strings.TrimSpace(event.TeamID) + ev := event.Event + + if ev.BotID != "" || ev.Subtype == "bot_message" || strings.TrimSpace(ev.User) == "" { + return nil, nil // Ignore bot loops + } + + text := strings.TrimSpace(ev.Text) + if text == "" { + return nil, nil + } + + var channel *models.Channel + channelID = strings.TrimSpace(channelID) + if channelID != "" { + channel = ChannelService.Take("channel_id = ? AND channel_type = ? AND status = ?", channelID, enums.ChannelTypeSlack, enums.StatusOk) + } + if channel == nil && teamID != "" { + channel = ChannelService.Take("channel_type = ? AND status = ? AND (channel_id = ? OR config_json LIKE ?)", + enums.ChannelTypeSlack, enums.StatusOk, teamID, "%"+teamID+"%") + } + if channel == nil { + channel = ChannelService.Take("channel_type = ? AND status = ?", enums.ChannelTypeSlack, enums.StatusOk) + } + if channel == nil { + return nil, errorsx.InvalidParam("slack channel not found or disabled") + } + + cfg, err := ChannelService.ParseSlackChannelConfig(channel.ConfigJSON) + if err != nil || cfg == nil { + return nil, errorsx.InvalidParam("slack channel config invalid") + } + + // Verify the Slack signing secret whenever the channel has one. A delivery + // with no signature headers is rejected rather than waved through: Slack + // always signs once a signing secret exists, so a missing header means the + // sender is not Slack. + if cfg.SigningSecret != "" { + if strings.TrimSpace(signatureHeader) == "" || strings.TrimSpace(timestampHeader) == "" { + return nil, errorsx.UnauthorizedI18n("error.auth.invalidSignature") + } + if !verifySlackSignature(cfg.SigningSecret, timestampHeader, signatureHeader, rawPayload) { + return nil, errorsx.UnauthorizedI18n("error.auth.invalidSignature") + } + } + + // 1. Resolve customer identity + senderID := strings.TrimSpace(ev.User) + externalUser := openidentity.ExternalUser{ + ExternalSource: enums.ExternalSourceSlack, + ExternalID: senderID, + ExternalName: fmt.Sprintf("Slack User %s", senderID), + } + + // 2. Create or match Conversation + conversation, err := ConversationService.Create(externalUser, channel.ID, channel.AIAgentID) + if err != nil { + return nil, fmt.Errorf("create slack conversation failed: %w", err) + } + + // 3. Send message through MessageService + msgTS := strings.TrimSpace(ev.TS) + threadTS := strings.TrimSpace(ev.ThreadTS) + if threadTS == "" { + threadTS = msgTS + } + clientMsgID := fmt.Sprintf("slack_%s_%s", ev.Channel, msgTS) + payloadMap := map[string]any{ + "slack_channel": ev.Channel, + "slack_ts": msgTS, + "slack_thread_ts": threadTS, + "slack_user": senderID, + "slack_team": teamID, + } + payloadBytes, _ := json.Marshal(payloadMap) + + _, err = MessageService.SendCustomerMessage( + conversation.ID, + clientMsgID, + enums.IMMessageTypeText, + text, + string(payloadBytes), + externalUser, + ) + if err != nil { + return nil, fmt.Errorf("send customer message failed: %w", err) + } + + return nil, nil +} + +// slackTimestampTolerance is how far a request timestamp may drift from now. +// +// Slack's own verification guide requires rejecting anything older than five +// minutes. Without the check a captured request replays indefinitely: the +// signature covers the timestamp and the body, but nothing in it expires. +const slackTimestampTolerance = 5 * time.Minute + +func verifySlackSignature(signingSecret, timestampHeader, signatureHeader string, payload []byte) bool { + timestampHeader = strings.TrimSpace(timestampHeader) + timestamp, err := strconv.ParseInt(timestampHeader, 10, 64) + if err != nil || timestamp <= 0 { + return false + } + if drift := time.Since(time.Unix(timestamp, 0)); drift > slackTimestampTolerance || drift < -slackTimestampTolerance { + return false + } + + sigBasestring := fmt.Sprintf("v0:%s:%s", timestampHeader, string(payload)) + mac := hmac.New(sha256.New, []byte(signingSecret)) + mac.Write([]byte(sigBasestring)) + expectedSig := "v0=" + hex.EncodeToString(mac.Sum(nil)) + return hmac.Equal([]byte(strings.TrimSpace(signatureHeader)), []byte(expectedSig)) +} diff --git a/internal/services/slack_inbound_service_test.go b/internal/services/slack_inbound_service_test.go new file mode 100644 index 00000000..bb0c8aa0 --- /dev/null +++ b/internal/services/slack_inbound_service_test.go @@ -0,0 +1,349 @@ +package services + +import ( + "context" + "crypto/hmac" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "strconv" + "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" +) + +// slackTestSigningSecret is the signing secret the test channel is configured with. +const slackTestSigningSecret = "test_signing_secret_999" + +// signSlackPayload builds the X-Slack-Request-Timestamp and X-Slack-Signature +// headers Slack would send for this body right now. +func signSlackPayload(t *testing.T, secret string, payload []byte) (string, string) { + t.Helper() + return signSlackPayloadAt(t, secret, payload, time.Now()) +} + +func signSlackPayloadAt(t *testing.T, secret string, payload []byte, at time.Time) (string, string) { + t.Helper() + timestamp := strconv.FormatInt(at.Unix(), 10) + mac := hmac.New(sha256.New, []byte(secret)) + mac.Write([]byte("v0:" + timestamp + ":" + string(payload))) + return timestamp, "v0=" + hex.EncodeToString(mac.Sum(nil)) +} + +func setupSlackTestDB(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 slack test tables: %v", err) + } + sqls.SetDB(db) + return db +} + +func TestSlackInboundAndOutbound(t *testing.T) { + db := setupSlackTestDB(t) + + now := time.Now() + aiAgent := &models.AIAgent{ + Name: "Slack Bot Agent", + 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) + } + + slackConfig := dto.SlackChannelConfig{ + BotToken: "xoxb-test-bot-token-12345", + SigningSecret: slackTestSigningSecret, + TeamID: "T0123456789", + TeamName: "Acme Corp", + DefaultChannel: "C9876543210", + } + cfgBytes, _ := json.Marshal(slackConfig) + + channel := &models.Channel{ + ChannelType: enums.ChannelTypeSlack, + ChannelID: "T0123456789", + AIAgentID: aiAgent.ID, + AIAgentRolloutPercent: 100, + Name: "Slack Support Channel", + ConfigJSON: string(cfgBytes), + Status: enums.StatusOk, + AuditFields: models.AuditFields{CreatedAt: now, UpdatedAt: now}, + } + if err := db.Create(channel).Error; err != nil { + t.Fatalf("create slack channel: %v", err) + } + + payload := `{ + "token": "verification_token", + "team_id": "T0123456789", + "api_app_id": "A01234567", + "type": "event_callback", + "event": { + "type": "message", + "user": "U12345678", + "text": "Help with API key generation", + "ts": "1725260000.000200", + "channel": "C9876543210", + "channel_type": "channel" + } + }` + + ctx := context.Background() + timestamp, signature := signSlackPayload(t, slackTestSigningSecret, []byte(payload)) + _, err := SlackInboundService.HandleWebhook(ctx, "", timestamp, signature, []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.ExternalSourceSlack). + Eq("external_id", "U12345678")) + if identity == nil { + t.Fatalf("expected customer identity for U12345678") + } + + // 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 message + msg := repositories.MessageRepository.FindOne(sqls.DB(), sqls.NewCnd(). + Eq("conversation_id", conv.ID). + Eq("sender_type", enums.IMSenderTypeCustomer)) + if msg == nil { + t.Fatalf("expected message to be created") + } + if msg.Content != "Help with API key generation" { + t.Fatalf("expected message content 'Help with API key generation', got %s", msg.Content) + } + + operator := &dto.AuthPrincipal{UserID: 1, Nickname: "Agent Joy"} + + // Test Outbound enqueue + replyMsg, err := MessageService.SendAIMessage(conv.ID, aiAgent.ID, "ai_slack_reply_1", enums.IMMessageTypeText, "You can generate your API key under Settings > API Keys.", "", operator) + if err != nil { + t.Fatalf("MessageService.SendAIMessage failed: %v", err) + } + + outbox := ChannelMessageOutboxService.GetByMessageID(enums.ChannelTypeSlack, replyMsg.ID) + if outbox == nil { + t.Fatalf("expected outbox entry for slack message") + } + if outbox.ChannelType != enums.ChannelTypeSlack { + t.Fatalf("expected outbox channel type 'slack', got %s", outbox.ChannelType) + } +} + +func seedSlackChannel(t *testing.T, db *gorm.DB, signingSecret string) *models.Channel { + t.Helper() + 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) + } + cfgBytes, err := json.Marshal(dto.SlackChannelConfig{ + BotToken: "xoxb-test-bot-token-12345", + SigningSecret: signingSecret, + TeamID: "T0123456789", + }) + if err != nil { + t.Fatalf("marshal slack config: %v", err) + } + channel := &models.Channel{ + ChannelType: enums.ChannelTypeSlack, + ChannelID: "slack_sig_channel", + AIAgentID: aiAgent.ID, + AIAgentRolloutPercent: 100, + Name: "Slack 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 slack channel: %v", err) + } + return channel +} + +func slackEventPayload(messageID string) []byte { + return []byte(`{ + "team_id": "T0123456789", + "type": "event_callback", + "event": { + "type": "message", + "user": "U_SIG_TEST", + "text": "signature probe", + "ts": "` + messageID + `", + "channel": "C9876543210", + "channel_type": "channel" + } + }`) +} + +func countSlackMessages(t *testing.T, db *gorm.DB, messageTS string) int64 { + t.Helper() + var count int64 + if err := db.Table("t_message").Where("client_msg_id = ?", "slack_C9876543210_"+messageTS).Count(&count).Error; err != nil { + t.Fatalf("count messages: %v", err) + } + return count +} + +// A channel with a signing secret configured must reject anything Slack did not +// sign. Accepting an unsigned delivery would let anyone who learns the webhook +// URL write into a workspace conversation and trigger paid AI replies. +func TestSlackInboundRejectsUnauthenticatedDelivery(t *testing.T) { + db := setupSlackTestDB(t) + channel := seedSlackChannel(t, db, slackTestSigningSecret) + + cases := []struct { + name string + messageTS string + timestamp string + signature string + }{ + {"no signature headers at all", "1725260000.000001", "", ""}, + {"signature without a timestamp", "1725260000.000002", "", "v0=deadbeef"}, + {"timestamp without a signature", "1725260000.000003", "1725260000", ""}, + {"wrong signature", "1725260000.000004", "1725260000", "v0=deadbeef"}, + {"non-numeric timestamp", "1725260000.000005", "not-a-timestamp", "v0=deadbeef"}, + } + + for _, tc := range cases { + payload := slackEventPayload(tc.messageTS) + _, err := SlackInboundService.HandleWebhook(context.Background(), channel.ChannelID, tc.timestamp, tc.signature, payload) + if err == nil { + t.Errorf("%s: expected the delivery to be rejected", tc.name) + continue + } + if countSlackMessages(t, db, tc.messageTS) != 0 { + t.Errorf("%s: a rejected delivery stored a message", tc.name) + } + } +} + +// Slack signs the timestamp and the body but nothing in the signature expires, so +// a captured request replays forever unless the timestamp is checked. Slack's own +// guide requires rejecting anything older than five minutes. +func TestSlackInboundRejectsReplayedTimestamp(t *testing.T) { + db := setupSlackTestDB(t) + channel := seedSlackChannel(t, db, slackTestSigningSecret) + + cases := []struct { + name string + age time.Duration + }{ + {"ten minutes old", 10 * time.Minute}, + {"one hour old", time.Hour}, + {"ten minutes in the future", -10 * time.Minute}, + } + + for i, tc := range cases { + messageTS := "1725260000.0000" + strconv.Itoa(10+i) + payload := slackEventPayload(messageTS) + // Correctly signed for its own timestamp, which is exactly what a replayed + // capture looks like on the wire. + timestamp, signature := signSlackPayloadAt(t, slackTestSigningSecret, payload, time.Now().Add(-tc.age)) + + if _, err := SlackInboundService.HandleWebhook(context.Background(), channel.ChannelID, timestamp, signature, payload); err == nil { + t.Errorf("%s: expected a stale timestamp to be rejected", tc.name) + } + if countSlackMessages(t, db, messageTS) != 0 { + t.Errorf("%s: a replayed delivery stored a message", tc.name) + } + } +} + +// A correctly signed, fresh delivery is still accepted, and one signed just +// inside the tolerance window is not rejected for clock drift. +func TestSlackInboundAcceptsFreshValidSignature(t *testing.T) { + db := setupSlackTestDB(t) + channel := seedSlackChannel(t, db, slackTestSigningSecret) + + cases := []struct { + name string + age time.Duration + }{ + {"signed now", 0}, + {"signed four minutes ago", 4 * time.Minute}, + } + + for i, tc := range cases { + messageTS := "1725260000.0000" + strconv.Itoa(20+i) + payload := slackEventPayload(messageTS) + timestamp, signature := signSlackPayloadAt(t, slackTestSigningSecret, payload, time.Now().Add(-tc.age)) + + if _, err := SlackInboundService.HandleWebhook(context.Background(), channel.ChannelID, timestamp, signature, payload); err != nil { + t.Fatalf("%s: HandleWebhook failed: %v", tc.name, err) + } + if countSlackMessages(t, db, messageTS) != 1 { + t.Errorf("%s: expected the signed message to be stored", tc.name) + } + } +} + +// The url_verification handshake has to be answered before any channel lookup or +// signature check, because Slack sends it once while the endpoint is being +// configured and will not retry. +func TestSlackInboundAnswersURLVerificationChallenge(t *testing.T) { + setupSlackTestDB(t) + + payload := []byte(`{"type":"url_verification","challenge":"challenge_token_abc","token":"verification_token"}`) + challenge, err := SlackInboundService.HandleWebhook(context.Background(), "", "", "", payload) + if err != nil { + t.Fatalf("HandleWebhook failed: %v", err) + } + if challenge == nil || *challenge != "challenge_token_abc" { + t.Fatalf("challenge = %v, want challenge_token_abc", challenge) + } +} diff --git a/internal/services/slack_outbound_service.go b/internal/services/slack_outbound_service.go new file mode 100644 index 00000000..68951df3 --- /dev/null +++ b/internal/services/slack_outbound_service.go @@ -0,0 +1,180 @@ +package services + +import ( + "context" + "encoding/json" + "log/slog" + "strings" + "time" + + "agent-desk/internal/models" + "agent-desk/internal/pkg/enums" + "agent-desk/internal/repositories" + "agent-desk/internal/services/storage" + "agent-desk/internal/slack" + + "github.com/mlogclub/simple/sqls" +) + +const ( + slackOutboxBatchSize = 20 + slackOutboxMaxRetry = 5 +) + +var SlackOutboundService = newSlackOutboundService() + +func newSlackOutboundService() *slackOutboundService { + return &slackOutboundService{} +} + +type slackOutboundService struct{} + +func (s *slackOutboundService) DispatchPendingOutbox() int { + return s.doDispatchPendingOutbox(slackOutboxBatchSize) +} + +func (s *slackOutboundService) doDispatchPendingOutbox(limit int) int { + if limit <= 0 { + limit = slackOutboxBatchSize + } + items := ChannelMessageOutboxService.ListPending(enums.ChannelTypeSlack, 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 slack outbox failed", + "outbox_id", items[i].ID, + "error", err, + ) + continue + } + successCount++ + } + return successCount +} + +func (s *slackOutboundService) processOutbox(outboxID int64) error { + outbox := ChannelMessageOutboxService.Get(outboxID) + if outbox == nil { + return nil + } + if outbox.ChannelType != enums.ChannelTypeSlack { + 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, "slack channel not found or disabled") + } + cfg, err := ChannelService.ParseSlackChannelConfig(channel.ConfigJSON) + if err != nil || cfg == nil || strings.TrimSpace(cfg.BotToken) == "" { + return s.markOutboxFailed(outbox, "slack bot token not configured") + } + + // Resolve target Slack Channel ID and Thread TS + var targetChannel string + var threadTS 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 ch, ok := payloadMap["slack_channel"].(string); ok && ch != "" { + targetChannel = ch + } + if ts, ok := payloadMap["slack_thread_ts"].(string); ok && ts != "" { + threadTS = ts + } + } + } + + if targetChannel == "" { + targetChannel = cfg.DefaultChannel + } + if targetChannel == "" { + return s.markOutboxFailed(outbox, "unable to resolve target slack channel") + } + + client := slack.NewClient(cfg.BotToken) + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + textToSend := message.Content + if message.MessageType == enums.IMMessageTypeImage || message.MessageType == enums.IMMessageTypeAttachment { + assetPayload, err := parseIMMessageAssetPayload(message.Payload) + 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) + if fileURL != "" { + if textToSend != "" { + textToSend += "\n" + fileURL + } else { + textToSend = fileURL + } + } + } + } + } + } + + _, sendErr := client.PostMessage(ctx, targetChannel, textToSend, threadTS) + 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 *slackOutboundService) markOutboxFailed(outbox *models.ChannelMessageOutbox, errMsg string) error { + if outbox == nil { + return nil + } + retryCount := outbox.RetryCount + 1 + status := string(enums.ChannelMessageOutboxStatusFailed) + if retryCount >= slackOutboxMaxRetry { + 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/slack/client.go b/internal/slack/client.go new file mode 100644 index 00000000..4ed2fdc7 --- /dev/null +++ b/internal/slack/client.go @@ -0,0 +1,109 @@ +package slack + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" +) + +const defaultBaseURL = "https://slack.com/api" + +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) PostMessage(ctx context.Context, channel string, text string, threadTS string) (*SendMessageResponse, error) { + channel = strings.TrimSpace(channel) + if channel == "" { + return nil, fmt.Errorf("slack channel is required") + } + text = strings.TrimSpace(text) + if text == "" { + return nil, fmt.Errorf("message text is required") + } + + payload := SendMessageRequest{ + Channel: channel, + Text: text, + ThreadTS: threadTS, + } + + var resp SendMessageResponse + if err := c.doRequest(ctx, "/chat.postMessage", payload, &resp); err != nil { + return nil, err + } + if !resp.OK { + return nil, fmt.Errorf("slack api error: %s", resp.Error) + } + return &resp, nil +} + +func (c *Client) doRequest(ctx context.Context, path string, payload any, result any) error { + if c.botToken == "" { + return fmt.Errorf("slack 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 slack request failed: %w", err) + } + bodyReader = bytes.NewBuffer(bodyBytes) + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bodyReader) + if err != nil { + return fmt.Errorf("create slack request failed: %w", err) + } + + req.Header.Set("Authorization", "Bearer "+c.botToken) + if payload != nil { + req.Header.Set("Content-Type", "application/json; charset=utf-8") + } + + res, err := c.httpClient.Do(req) + if err != nil { + return fmt.Errorf("slack http request failed: %w", err) + } + defer res.Body.Close() + + bodyBytes, err := io.ReadAll(res.Body) + if err != nil { + return fmt.Errorf("read slack response failed: %w", err) + } + + if res.StatusCode < 200 || res.StatusCode >= 300 { + return fmt.Errorf("slack api error (%d): %s", res.StatusCode, string(bodyBytes)) + } + + if result != nil { + if err := json.Unmarshal(bodyBytes, result); err != nil { + return fmt.Errorf("unmarshal slack response failed: %w (body: %s)", err, string(bodyBytes)) + } + } + return nil +} diff --git a/internal/slack/client_test.go b/internal/slack/client_test.go new file mode 100644 index 00000000..96fe56e9 --- /dev/null +++ b/internal/slack/client_test.go @@ -0,0 +1,120 @@ +package slack + +import ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "testing" +) + +func TestPostMessageSendsThreadedReply(t *testing.T) { + var ( + gotPath string + gotAuth string + gotBody map[string]any + ) + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotPath = r.URL.Path + gotAuth = r.Header.Get("Authorization") + raw, _ := io.ReadAll(r.Body) + _ = json.Unmarshal(raw, &gotBody) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"ok":true,"channel":"C0123","ts":"1725260000.000100"}`)) + })) + defer server.Close() + + client := NewClient("xoxb-test-token") + client.SetBaseURL(server.URL) + + resp, err := client.PostMessage(context.Background(), "C0123", "hello there", "1725260000.000200") + if err != nil { + t.Fatalf("PostMessage failed: %v", err) + } + if resp == nil || resp.TS != "1725260000.000100" { + t.Fatalf("response = %+v, want the posted message ts", resp) + } + + if gotPath != "/chat.postMessage" { + t.Errorf("path = %q, want /chat.postMessage", gotPath) + } + if gotAuth != "Bearer xoxb-test-token" { + t.Errorf("authorization = %q, want the bot token as a bearer", gotAuth) + } + if gotBody["channel"] != "C0123" { + t.Errorf("channel = %v, want C0123", gotBody["channel"]) + } + if gotBody["text"] != "hello there" { + t.Errorf("text = %v, want 'hello there'", gotBody["text"]) + } + // A reply has to stay in the customer's thread, otherwise the workspace + // timeline fills up with support answers. + if gotBody["thread_ts"] != "1725260000.000200" { + t.Errorf("thread_ts = %v, want the parent ts", gotBody["thread_ts"]) + } +} + +func TestPostMessageOmitsThreadTSForTopLevel(t *testing.T) { + var gotRaw string + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + raw, _ := io.ReadAll(r.Body) + gotRaw = string(raw) + _, _ = w.Write([]byte(`{"ok":true}`)) + })) + defer server.Close() + + client := NewClient("xoxb-test-token") + client.SetBaseURL(server.URL) + + if _, err := client.PostMessage(context.Background(), "C0123", "top level", ""); err != nil { + t.Fatalf("PostMessage failed: %v", err) + } + // thread_ts is omitempty on the request struct, so a top-level post must not + // carry an empty one. + var body map[string]any + if err := json.Unmarshal([]byte(gotRaw), &body); err != nil { + t.Fatalf("unmarshal request body: %v", err) + } + if _, present := body["thread_ts"]; present { + t.Errorf("request body carried thread_ts for a top-level message: %s", gotRaw) + } +} + +func TestPostMessageSurfacesSlackAPIError(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Slack returns HTTP 200 with ok:false for application-level failures, so + // the client has to read the envelope rather than the status code. + _, _ = w.Write([]byte(`{"ok":false,"error":"channel_not_found"}`)) + })) + defer server.Close() + + client := NewClient("xoxb-test-token") + client.SetBaseURL(server.URL) + + _, err := client.PostMessage(context.Background(), "C_MISSING", "hello", "") + if err == nil { + t.Fatalf("expected an error when Slack answers ok:false") + } + if got := err.Error(); got != "slack api error: channel_not_found" { + t.Errorf("error = %q, want the slack error code to be carried through", got) + } +} + +func TestPostMessageValidatesInputs(t *testing.T) { + client := NewClient("xoxb-test-token") + + if _, err := client.PostMessage(context.Background(), "", "hello", ""); err == nil { + t.Errorf("expected an error for an empty channel") + } + if _, err := client.PostMessage(context.Background(), "C0123", " ", ""); err == nil { + t.Errorf("expected an error for blank text") + } + + noToken := NewClient("") + if _, err := noToken.PostMessage(context.Background(), "C0123", "hello", ""); err == nil { + t.Errorf("expected an error when no bot token is configured") + } +} diff --git a/internal/slack/types.go b/internal/slack/types.go new file mode 100644 index 00000000..270e4e48 --- /dev/null +++ b/internal/slack/types.go @@ -0,0 +1,37 @@ +package slack + +// SendMessageRequest represents payload for Slack chat.postMessage API. +type SendMessageRequest struct { + Channel string `json:"channel"` + Text string `json:"text"` + ThreadTS string `json:"thread_ts,omitempty"` + ParseMode string `json:"parse,omitempty"` +} + +// SendMessageResponse represents response from Slack Web API. +type SendMessageResponse struct { + OK bool `json:"ok"` + Channel string `json:"channel,omitempty"` + TS string `json:"ts,omitempty"` + Error string `json:"error,omitempty"` +} + +// EventCallback represents incoming Slack Events API payload. +type EventCallback struct { + Token string `json:"token"` + TeamID string `json:"team_id"` + APIAppID string `json:"api_app_id"` + Type string `json:"type"` // url_verification | event_callback + Challenge string `json:"challenge"` // for url_verification + Event *struct { + Type string `json:"type"` // message | app_mention + User string `json:"user"` + Text string `json:"text"` + TS string `json:"ts"` + ThreadTS string `json:"thread_ts,omitempty"` + Channel string `json:"channel"` + ChannelType string `json:"channel_type"` // im | channel | group + BotID string `json:"bot_id,omitempty"` + Subtype string `json:"subtype,omitempty"` + } `json:"event,omitempty"` +} diff --git a/web/app/(dashboard)/dashboard/channels/_components/edit.tsx b/web/app/(dashboard)/dashboard/channels/_components/edit.tsx index c5b96168..f3b91262 100644 --- a/web/app/(dashboard)/dashboard/channels/_components/edit.tsx +++ b/web/app/(dashboard)/dashboard/channels/_components/edit.tsx @@ -73,6 +73,15 @@ type ZaloOAChannelConfig = { webhookSecret?: string } +type SlackChannelConfig = { + botToken?: string + signingSecret?: string + appId?: string + teamId?: string + teamName?: string + defaultChannel?: string +} + function getDefaultWebChannelConfig(t: Translate): Required { return { title: t("channel.defaultTitleWeb"), @@ -87,7 +96,7 @@ function getDefaultWebChannelConfig(t: Translate): Required { function createSchema(t: Translate) { return z .object({ - channelType: z.enum(["web", "wechat_mp", "wxwork_kf", "telegram", "zalo_oa"], t("channel.typeRequired")), + channelType: z.enum(["web", "wechat_mp", "wxwork_kf", "telegram", "zalo_oa", "slack"], t("channel.typeRequired")), aiAgentId: z.string().trim().regex(/^\d+$/, t("channel.agentRequired")), aiAgentRolloutPercent: z.coerce.number().int().min(1).max(100), name: z.string().trim().min(1, t("channel.nameRequired")), @@ -99,6 +108,12 @@ function createSchema(t: Translate) { zaloOaId: z.string().trim(), zaloAccessToken: z.string().trim(), zaloSecretKey: z.string().trim(), + slackBotToken: z.string().trim(), + slackSigningSecret: z.string().trim(), + slackAppId: z.string().trim(), + slackTeamId: z.string().trim(), + slackTeamName: z.string().trim(), + slackDefaultChannel: z.string().trim(), widgetTitle: z.string().trim(), widgetSubtitle: z.string().trim(), widgetThemeColor: z.string().trim(), @@ -129,11 +144,18 @@ function createSchema(t: Translate) { message: "Zalo OA Access Token is required", }) } + if (values.channelType === "slack" && !values.slackBotToken.trim()) { + ctx.addIssue({ + code: "custom", + path: ["slackBotToken"], + message: t("channel.slackBotTokenRequired"), + }) + } }) } type EditForm = { - channelType: "web" | "wechat_mp" | "wxwork_kf" | "telegram" | "zalo_oa" + channelType: "web" | "wechat_mp" | "wxwork_kf" | "telegram" | "zalo_oa" | "slack" aiAgentId: string aiAgentRolloutPercent: number name: string @@ -145,6 +167,12 @@ type EditForm = { zaloOaId: string zaloAccessToken: string zaloSecretKey: string + slackBotToken: string + slackSigningSecret: string + slackAppId: string + slackTeamId: string + slackTeamName: string + slackDefaultChannel: string widgetTitle: string widgetSubtitle: string widgetThemeColor: string @@ -169,6 +197,12 @@ function createEmptyForm(t: Translate): EditForm { zaloOaId: "", zaloAccessToken: "", zaloSecretKey: "", + slackBotToken: "", + slackSigningSecret: "", + slackAppId: "", + slackTeamId: "", + slackTeamName: "", + slackDefaultChannel: "", widgetTitle: defaultWebChannelConfig.title, widgetSubtitle: defaultWebChannelConfig.subtitle, widgetThemeColor: defaultWebChannelConfig.themeColor, @@ -222,6 +256,23 @@ function parseZaloOAChannelConfig(configJson: string): ZaloOAChannelConfig { } } +function parseSlackChannelConfig(configJson: string): SlackChannelConfig { + if (!configJson.trim()) return {} + try { + const parsed = JSON.parse(configJson) as SlackChannelConfig + return { + botToken: parsed.botToken?.trim() || "", + signingSecret: parsed.signingSecret?.trim() || "", + appId: parsed.appId?.trim() || "", + teamId: parsed.teamId?.trim() || "", + teamName: parsed.teamName?.trim() || "", + defaultChannel: parsed.defaultChannel?.trim() || "", + } + } catch { + return {} + } +} + function parseWebChannelConfig(configJson: string, t: Translate): Required { const defaultWebChannelConfig = getDefaultWebChannelConfig(t) if (!configJson.trim()) { @@ -276,6 +327,7 @@ function buildForm(item: AdminChannel | null, t: Translate): EditForm { const isWechatMP = item.channelType === "wechat_mp" const isTelegram = item.channelType === "telegram" const isZaloOA = item.channelType === "zalo_oa" + const isSlack = item.channelType === "slack" const webConfig = parseWebChannelConfig(item.configJson, t) const wechatConfig = isWechatMP ? parseWechatMPChannelConfig(item.configJson, t) @@ -286,6 +338,9 @@ function buildForm(item: AdminChannel | null, t: Translate): EditForm { const zaloConfig = isZaloOA ? parseZaloOAChannelConfig(item.configJson) : null + const slackConfig = isSlack + ? parseSlackChannelConfig(item.configJson) + : null return { channelType: item.channelType === "wxwork_kf" @@ -294,9 +349,11 @@ function buildForm(item: AdminChannel | null, t: Translate): EditForm { ? "telegram" : item.channelType === "zalo_oa" ? "zalo_oa" - : item.channelType === "wechat_mp" - ? "wechat_mp" - : "web", + : item.channelType === "slack" + ? "slack" + : item.channelType === "wechat_mp" + ? "wechat_mp" + : "web", aiAgentId: item.aiAgentId > 0 ? String(item.aiAgentId) : "", aiAgentRolloutPercent: item.aiAgentRolloutPercent || 100, name: item.name, @@ -308,6 +365,12 @@ function buildForm(item: AdminChannel | null, t: Translate): EditForm { zaloOaId: zaloConfig?.oaId ?? "", zaloAccessToken: zaloConfig?.accessToken ?? "", zaloSecretKey: zaloConfig?.secretKey ?? "", + slackBotToken: slackConfig?.botToken ?? "", + slackSigningSecret: slackConfig?.signingSecret ?? "", + slackAppId: slackConfig?.appId ?? "", + slackTeamId: slackConfig?.teamId ?? "", + slackTeamName: slackConfig?.teamName ?? "", + slackDefaultChannel: slackConfig?.defaultChannel ?? "", widgetTitle: wechatConfig?.title ?? webConfig.title, widgetSubtitle: wechatConfig?.subtitle ?? webConfig.subtitle, widgetThemeColor: wechatConfig?.themeColor ?? webConfig.themeColor, @@ -347,14 +410,23 @@ function buildPayload(form: EditForm, status: number, t: Translate): CreateAdmin secretKey: form.zaloSecretKey.trim(), webhookSecret: form.webhookSecret.trim(), }) - : channelType === "wechat_mp" - ? JSON.stringify(webLikeConfig) - : JSON.stringify({ - ...webLikeConfig, - position: form.widgetPosition || defaultWebChannelConfig.position, - width: form.widgetWidth.trim() || defaultWebChannelConfig.width, - userTokenSecret: form.userTokenSecret.trim(), + : channelType === "slack" + ? JSON.stringify({ + botToken: form.slackBotToken.trim(), + signingSecret: form.slackSigningSecret.trim(), + appId: form.slackAppId.trim(), + teamId: form.slackTeamId.trim(), + teamName: form.slackTeamName.trim(), + defaultChannel: form.slackDefaultChannel.trim(), }) + : channelType === "wechat_mp" + ? JSON.stringify(webLikeConfig) + : JSON.stringify({ + ...webLikeConfig, + position: form.widgetPosition || defaultWebChannelConfig.position, + width: form.widgetWidth.trim() || defaultWebChannelConfig.width, + userTokenSecret: form.userTokenSecret.trim(), + }) return { channelType, aiAgentId: Number(form.aiAgentId), @@ -543,6 +615,7 @@ function ChannelFormBody({ const channelTypeOptions = [ { value: "web", label: t("channel.typeWeb") }, { value: "telegram", label: t("channel.typeTelegram") }, + { value: "slack", label: t("channel.typeSlack") }, { value: "wechat_mp", label: t("channel.typeWechatMp") }, { value: "wxwork_kf", label: t("channel.typeWxworkKf") }, ] as const @@ -765,6 +838,97 @@ function ChannelFormBody({
) : null} + {channelType === "slack" ? ( +
+ + {t("channel.slackBotToken")} * + + + + + + + + {t("channel.slackSigningSecret")} + + + +

+ {t("channel.slackSigningSecretHint")} +

+
+
+ +
+ + {t("channel.slackTeamId")} + + + + + + + + {t("channel.slackDefaultChannel")} + + + + + + + + {t("channel.slackAppId")} + + + + + + + + {t("channel.slackTeamName")} + + + + + +
+ +
+
{t("channel.slackSetupTitle")}
+
{t("channel.slackSetupDescription")}
+
+ {t("channel.slackRequestUrl")}: /api/third/slack/webhook +
+
+
+ ) : null} + {channelType === "telegram" ? (
diff --git a/web/app/(dashboard)/dashboard/channels/page.tsx b/web/app/(dashboard)/dashboard/channels/page.tsx index aa8c06f7..990bfc8b 100644 --- a/web/app/(dashboard)/dashboard/channels/page.tsx +++ b/web/app/(dashboard)/dashboard/channels/page.tsx @@ -5,6 +5,7 @@ import { MessagesSquareIcon, MessageSquareMoreIcon, SendIcon, + SlackIcon, } from "lucide-react" import { @@ -39,6 +40,9 @@ function getChannelTypeLabel(channelType: string, t: (key: string) => string) { if (channelType === "zalo_oa") { return t("channel.typeZaloOa") } + if (channelType === "slack") { + return t("channel.typeSlack") + } return t("channel.typeWeb") } @@ -62,6 +66,9 @@ function ChannelIcon({ channelType }: { channelType: string }) { if (channelType === "telegram" || channelType === "zalo_oa") { return } + if (channelType === "slack") { + return + } return } @@ -78,6 +85,7 @@ export default function DashboardChannelsPage() { { value: "all", label: t("channel.allTypes") }, { value: "web", label: t("channel.typeWeb") }, { value: "telegram", label: t("channel.typeTelegram") }, + { value: "slack", label: t("channel.typeSlack") }, { value: "zalo_oa", label: t("channel.typeZaloOa") }, { value: "wechat_mp", label: t("channel.typeWechatMp") }, { value: "wxwork_kf", label: t("channel.typeWxworkKf") }, diff --git a/web/messages/en-US.json b/web/messages/en-US.json index c83bc000..06cb4a51 100644 --- a/web/messages/en-US.json +++ b/web/messages/en-US.json @@ -624,6 +624,18 @@ "zaloAppId": "App ID", "zaloAutoConnectTitle": "Zalo Official Account Connection", "zaloAutoConnectDescription": "Connect your Zalo Official Account using the Access Token to automatically receive customer messages and dispatch replies.", + "typeSlack": "Slack", + "slackBotToken": "Bot User OAuth Token", + "slackBotTokenRequired": "Slack Bot User OAuth Token is required", + "slackSigningSecret": "Signing Secret", + "slackSigningSecretHint": "Required to verify inbound events. Without it the webhook accepts unsigned requests, so anyone who learns the URL can post into a conversation.", + "slackAppId": "App ID", + "slackTeamId": "Workspace ID", + "slackTeamName": "Workspace Name", + "slackDefaultChannel": "Default Channel ID", + "slackSetupTitle": "Slack App Connection", + "slackSetupDescription": "Create a Slack app, install it to the workspace, and copy the Bot User OAuth Token and Signing Secret from Basic Information. Subscribe to message.events, then set the Request URL below in Event Subscriptions.", + "slackRequestUrl": "Request URL", "loadFailed": "Could not load channels.", "created": "Channel created: {name}", "updated": "Channel updated: {name}", diff --git a/web/messages/zh-CN.json b/web/messages/zh-CN.json index 14e67ae3..f73f838f 100644 --- a/web/messages/zh-CN.json +++ b/web/messages/zh-CN.json @@ -624,6 +624,18 @@ "zaloAppId": "App ID", "zaloAutoConnectTitle": "Zalo OA 渠道连接", "zaloAutoConnectDescription": "输入 Zalo OA 的 Access Token 即可自动双向同步客户会话与消息。", + "typeSlack": "Slack", + "slackBotToken": "Bot User OAuth Token", + "slackBotTokenRequired": "请填写 Slack Bot User OAuth Token", + "slackSigningSecret": "Signing Secret", + "slackSigningSecretHint": "用于校验入站事件,建议填写。留空时该 Webhook 会接受未签名请求,任何知道地址的人都能向会话发消息。", + "slackAppId": "App ID", + "slackTeamId": "工作区 ID", + "slackTeamName": "工作区名称", + "slackDefaultChannel": "默认频道 ID", + "slackSetupTitle": "Slack 应用接入", + "slackSetupDescription": "创建 Slack 应用并安装到工作区,在 Basic Information 页面复制 Bot User OAuth Token 和 Signing Secret。订阅 message.events 事件,再把下方的 Request URL 填入 Event Subscriptions。", + "slackRequestUrl": "Request URL", "loadFailed": "加载接入渠道失败", "created": "已创建接入渠道:{name}", "updated": "已更新接入渠道:{name}", From cfdc56321db30d0fc4e812cd16b559406b3e011d Mon Sep 17 00:00:00 2001 From: JOY <5027251+JOY@users.noreply.github.com> Date: Tue, 22 Sep 2026 09:44:48 +0700 Subject: [PATCH 4/8] fix(auth): stop granting admin to every first-time OIDC login ensureDefaultOIDCRole assigned the admin role (falling back to super_admin) to any user that arrived without a role, which in practice means every first-time OIDC login: a stranger with an account at the identity provider became an administrator of the support desk. Give the lowest staff role (cs_user) instead; administrative access must be granted deliberately. The regression test seeds all three candidate roles and asserts a first-time OIDC user lands on cs_user and nothing else. --- internal/services/oidc_login_service.go | 9 ++-- internal/services/oidc_login_service_test.go | 57 ++++++++++++++++++++ 2 files changed, 62 insertions(+), 4 deletions(-) diff --git a/internal/services/oidc_login_service.go b/internal/services/oidc_login_service.go index 9827c4d7..e764500f 100644 --- a/internal/services/oidc_login_service.go +++ b/internal/services/oidc_login_service.go @@ -257,6 +257,10 @@ func shortSubjectHash(subject string) string { return hex.EncodeToString(sum[:])[:16] } +// ensureDefaultOIDCRole gives a first-time OIDC user the lowest staff role. +// The provider vouches for who the user is, not what they may administer, so +// a user arriving without any role must never escalate to an administrative +// one. func (s *oidcLoginService) ensureDefaultOIDCRole(tx *gorm.DB, user *models.User) { if user == nil || user.ID <= 0 { return @@ -265,10 +269,7 @@ func (s *oidcLoginService) ensureDefaultOIDCRole(tx *gorm.DB, user *models.User) if existingRole != nil { return } - defaultRole := repositories.RoleRepository.GetByCode(tx, constants.RoleCodeAdmin) - if defaultRole == nil { - defaultRole = repositories.RoleRepository.GetByCode(tx, constants.RoleCodeSuperAdmin) - } + defaultRole := repositories.RoleRepository.GetByCode(tx, constants.RoleCodeCsUser) if defaultRole == nil { return } diff --git a/internal/services/oidc_login_service_test.go b/internal/services/oidc_login_service_test.go index 4e0749d7..daa1909e 100644 --- a/internal/services/oidc_login_service_test.go +++ b/internal/services/oidc_login_service_test.go @@ -3,9 +3,11 @@ package services import ( "strings" "testing" + "time" "agent-desk/internal/models" "agent-desk/internal/pkg/config" + "agent-desk/internal/pkg/constants" "agent-desk/internal/pkg/enums" ) @@ -94,3 +96,58 @@ func TestOIDCLoginReusesExistingIdentity(t *testing.T) { t.Fatalf("expected existing identity to reuse user, got %d users", count) } } + +// A first-time OIDC user is a stranger the provider vouched for, nothing +// more: they must land on the lowest staff role, never on an administrative +// one. Seeds every candidate role so the assignment cannot pass by accident +// of a missing row. +func TestOIDCLoginFirstUserGetsLowestStaffRole(t *testing.T) { + db := setupAuthServiceTestDB(t) + now := time.Now() + for _, role := range []struct{ name, code string }{ + {"Super Admin", constants.RoleCodeSuperAdmin}, + {"Admin", constants.RoleCodeAdmin}, + {"Support Agent", constants.RoleCodeCsUser}, + } { + if err := db.Create(&models.Role{ + Name: role.name, + Code: role.code, + Status: enums.StatusOk, + AuditFields: models.AuditFields{ + CreatedAt: now, + UpdatedAt: now, + }, + }).Error; err != nil { + t.Fatalf("seed role %s: %v", role.code, err) + } + } + + if _, err := newOIDCLoginService().loginWithOIDCProfile(&oidcLoginProfile{ + Subject: "sub-777", + Email: "stranger@example.com", + PreferredUsername: "stranger", + Name: "Stranger", + RawProfile: `{"sub":"sub-777"}`, + }, config.AuthConfig{TokenTTLHours: 2}, "127.0.0.1", "go-test"); err != nil { + t.Fatalf("loginWithOIDCProfile() error = %v", err) + } + + var user models.User + if err := db.Take(&user, "username = ?", "stranger").Error; err != nil { + t.Fatalf("expected OIDC user to be created: %v", err) + } + + var roles []models.Role + if err := db. + Joins("JOIN t_user_role ON t_user_role.role_id = t_role.id"). + Where("t_user_role.user_id = ?", user.ID). + Find(&roles).Error; err != nil { + t.Fatalf("query user roles: %v", err) + } + if len(roles) != 1 { + t.Fatalf("expected exactly one role for a first-time OIDC user, got %d", len(roles)) + } + if roles[0].Code != constants.RoleCodeCsUser { + t.Fatalf("first-time OIDC user role = %q, want %q", roles[0].Code, constants.RoleCodeCsUser) + } +} From fd82111fa3c76395fb8154d34199199d98f4a372 Mon Sep 17 00:00:00 2001 From: JOY <5027251+JOY@users.noreply.github.com> Date: Sun, 13 Sep 2026 16:35:07 +0700 Subject: [PATCH 5/8] fix(security): make the client IP trustworthy and scope credential lockout to it Two coupled defects. The second could not be fixed safely before the first. Gin was never told which proxies to trust. Nothing in the repository calls SetTrustedProxies or sets TrustedPlatform, so the engine ran on Gin's own default of 0.0.0.0/0 and ::/0 - every peer is a trusted proxy. validateHeader then walks X-Forwarded-For right to left and stops at the first address that is *not* a trusted proxy; with everything trusted it never stops, and returns the leftmost value, which is the one the caller supplied. Any client could therefore choose the address recorded in t_login_credential_log.client_ip and t_user.last_login_ip by sending one header. X-Real-IP is in Gin's default RemoteIPHeaders too, so it was a second route to the same result. - ServerConfig gains trustedProxies and trustedPlatform. Left unset, trustedProxies falls back to the loopback, RFC1918, IPv6 unique-local and link-local ranges rather than to Gin's trust-everything default. That covers a sidecar tunnel, a compose network and a local nginx, and a client on the public internet cannot present one of those addresses as its direct peer. - trustedPlatform maps "cloudflare", "fly.io" and "google-app-engine" onto Gin's own constants and passes any other value through as a literal header name. It defaults to empty on purpose: an edge that appends rather than overwrites would turn the header back into a client-controlled value. - NewServer applies both before any middleware is registered, and fails startup on an unparseable CIDR rather than quietly dropping the trust boundary. - Documented in config/config.example.yaml, and .env.example sets TRUSTED_PLATFORM=cloudflare because this deployment sits behind a Cloudflare Tunnel. With an address that can be believed, credential lockout stops being a denial of service. isCredentialLocked keyed on the username alone, so anybody who knew a username could lock the real account out for the whole window, from anywhere, as often as they liked - and the window was renewable indefinitely. - The per-account window is now keyed on principal *and* client address, so one source grinding on one account is still stopped. - A second window counts failures from one address across all principals, which is what credential stuffing looks like. auth.maxFailedAttemptsPerIP configures it and defaults to four times maxFailedAttempts; disabling maxFailedAttempts disables both, matching the existing convention. - An address the server could not determine is normalised to a single "unknown" bucket and excluded from the per-address window, so a missing IP cannot pool every such caller into one lockout. - createLoginCredentialLog normalises on write so the read and write sides cannot drift apart. Tests. TestGinDefaultTrustsEveryProxy pins the vulnerable default, so the fix stays load-bearing and a future Gin that changes the default says so instead of silently making the configuration decorative. Three more cover an untrusted peer's forged header being ignored, a trusted proxy's appended address beating client-prepended junk, and the platform header taking precedence over X-Forwarded-For. Two lockout tests cover the denial of service being closed and the per-address window catching a username spray that no single account trips. Four existing lockout tests seeded credential logs without a client address, which is exactly the assumption this removes, so they now seed one. --- .env.example | 11 ++ config/config.example.yaml | 12 ++ internal/bootstrap/server.go | 13 ++ internal/bootstrap/server_clientip_test.go | 148 +++++++++++++++++++++ internal/pkg/config/config.go | 84 +++++++++++- internal/pkg/config/config_test.go | 104 +++++++++++++++ internal/services/auth_service.go | 67 ++++++++-- internal/services/auth_service_test.go | 79 +++++++++++ 8 files changed, 505 insertions(+), 13 deletions(-) create mode 100644 internal/bootstrap/server_clientip_test.go diff --git a/.env.example b/.env.example index 12afb717..edbe2b24 100644 --- a/.env.example +++ b/.env.example @@ -9,6 +9,17 @@ COMPANY_NAME= COMPANY_LOGO_URL= # AGENT_DESK_SERVER_CORS_ALLOWEDORIGINS="http://localhost:3000,http://127.0.0.1:8083" +# Client IP resolution. +# Gin defaults to trusting every proxy, which makes X-Forwarded-For authoritative: +# any caller could then choose the IP recorded in the login credential log and in +# user.last_login_ip, and any IP-keyed rate limit or lockout would be bypassable. +# Unset, trustedProxies falls back to the loopback / RFC1918 / IPv6 unique-local / +# link-local ranges, which already covers a cloudflared sidecar on a compose +# network. Set trustedPlatform when an edge overwrites the header instead of +# appending to it - Cloudflare does, and CF-Connecting-IP is not client-forgeable. +# TRUSTED_PROXIES="10.0.0.0/8,172.16.0.0/12" +TRUSTED_PLATFORM=cloudflare + # Database Configuration # Driver options: sqlite, mysql, postgres DB_TYPE=sqlite diff --git a/config/config.example.yaml b/config/config.example.yaml index 4dd9f45f..d719f25d 100644 --- a/config/config.example.yaml +++ b/config/config.example.yaml @@ -5,6 +5,18 @@ server: # Custom platform / instance branding (optional) companyName: "" companyLogoUrl: "" + # Reverse proxies in front of the application, as CIDR blocks. Gin's own default + # trusts every peer, which makes X-Forwarded-For authoritative and lets any + # caller choose the client address the server records in the login credential + # log and in user.last_login_ip. Left empty this falls back to the loopback, + # RFC1918, IPv6 unique-local and link-local ranges, which covers a sidecar + # tunnel, a compose network and a local nginx. Set it explicitly when the proxy + # sits on a public address, and never set it to 0.0.0.0/0. + trustedProxies: [] + # Set to "cloudflare" (also accepted: "fly.io", "google-app-engine", or a literal + # header name) when an edge overwrites rather than appends the real client + # address. When set it takes precedence over X-Forwarded-For entirely. + trustedPlatform: "" cors: # Browser CORS allowlist. In production, replace this with the actual frontend or embedded-site domains, such as https://support.example.com. # Leave it empty to reject cross-origin browser requests. Same-origin and non-browser calls are still supported. diff --git a/internal/bootstrap/server.go b/internal/bootstrap/server.go index 93dcaf50..16416620 100644 --- a/internal/bootstrap/server.go +++ b/internal/bootstrap/server.go @@ -1,6 +1,7 @@ package bootstrap import ( + "fmt" "log/slog" "net/http" "path" @@ -35,6 +36,18 @@ func NewServer() (*gin.Engine, error) { printBanner() app := gin.New() + + // Gin defaults to trusting every proxy, which makes ClientIP() return the + // leftmost X-Forwarded-For value - a header any caller can set. Everything + // keyed on a client address depends on this being settled first: the login + // credential log, the user's last login IP, and any abuse control. + if platform := cfg.Server.TrustedPlatformHeader(); platform != "" { + app.TrustedPlatform = platform + } + if err := app.SetTrustedProxies(cfg.Server.TrustedProxiesOrDefault()); err != nil { + return nil, fmt.Errorf("invalid server.trustedProxies: %w", err) + } + app.Use(requestIDMiddleware()) app.Use(corsMiddleware()) app.Use(gin.Recovery()) diff --git a/internal/bootstrap/server_clientip_test.go b/internal/bootstrap/server_clientip_test.go new file mode 100644 index 00000000..2711eb24 --- /dev/null +++ b/internal/bootstrap/server_clientip_test.go @@ -0,0 +1,148 @@ +package bootstrap + +import ( + "bytes" + "log/slog" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "agent-desk/internal/pkg/config" + + "github.com/gin-gonic/gin" +) + +// captureRequestLog swaps the slog default for one writing into buf, because +// requestLogMiddleware records the address Gin resolved. That is the only place +// the server exposes its ClientIP() decision, and it is the value that ends up in +// t_login_credential_log and t_user.last_login_ip. +func captureRequestLog(t *testing.T) *bytes.Buffer { + t.Helper() + var buf bytes.Buffer + previous := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&buf, nil))) + t.Cleanup(func() { slog.SetDefault(previous) }) + return &buf +} + +func setTestServerConfig(server config.ServerConfig) { + config.SetCurrent(&config.Config{ + Server: server, + Storage: config.StorageConfig{Local: config.LocalStorageConfig{Root: "storage", BaseURL: "/storage"}}, + }) +} + +// TestGinDefaultTrustsEveryProxy pins the vulnerability the configuration above +// exists to close. A bare Gin engine - which is what NewServer used to build - +// trusts 0.0.0.0/0 and ::/0, so validateHeader walks X-Forwarded-For right to +// left, finds no untrusted proxy to stop at, and returns the leftmost value the +// caller chose. If this assertion ever starts failing, Gin's default has changed +// and the trusted-proxy configuration should be revisited rather than assumed +// necessary. +func TestGinDefaultTrustsEveryProxy(t *testing.T) { + app := gin.New() + var resolved string + app.GET("/probe", func(ctx *gin.Context) { resolved = ctx.ClientIP() }) + + req := httptest.NewRequest(http.MethodGet, "/probe", nil) + req.RemoteAddr = "203.0.113.7:52000" + req.Header.Set("X-Forwarded-For", "198.51.100.9") + app.ServeHTTP(httptest.NewRecorder(), req) + + if resolved != "198.51.100.9" { + t.Fatalf("a bare engine resolved ClientIP to %q, expected the forged 198.51.100.9", resolved) + } +} + +func TestNewServerIgnoresForgedForwardedForFromAnUntrustedPeer(t *testing.T) { + buf := captureRequestLog(t) + setTestServerConfig(config.ServerConfig{}) + + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/health", nil) + // A public peer is outside every default trusted range, so nothing it claims + // about the originating address may be believed. + req.RemoteAddr = "203.0.113.7:52000" + req.Header.Set("X-Forwarded-For", "198.51.100.9") + req.Header.Set("X-Real-IP", "198.51.100.9") + app.ServeHTTP(rec, req) + + logged := buf.String() + if !strings.Contains(logged, "clientIp=203.0.113.7") { + t.Fatalf("expected the real peer address to be logged, got: %s", logged) + } + if strings.Contains(logged, "198.51.100.9") { + t.Fatalf("a forged forwarding header was trusted: %s", logged) + } +} + +func TestNewServerReadsForwardedForFromATrustedProxy(t *testing.T) { + buf := captureRequestLog(t) + setTestServerConfig(config.ServerConfig{TrustedProxies: []string{"10.0.0.0/8"}}) + + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/health", nil) + req.RemoteAddr = "10.0.0.5:52000" + // A hostile client prepends a value it chose; the proxy appends the address it + // actually saw. Gin walks the list right to left and stops at the first + // address that is not itself a trusted proxy, which is the appended one. + req.Header.Set("X-Forwarded-For", "198.51.100.9, 203.0.113.7") + app.ServeHTTP(rec, req) + + logged := buf.String() + if !strings.Contains(logged, "clientIp=203.0.113.7") { + t.Fatalf("expected the proxy-appended address, got: %s", logged) + } + if strings.Contains(logged, "198.51.100.9") { + t.Fatalf("the client-prepended address was trusted: %s", logged) + } +} + +func TestNewServerPrefersTheConfiguredTrustedPlatformHeader(t *testing.T) { + buf := captureRequestLog(t) + setTestServerConfig(config.ServerConfig{TrustedPlatform: "cloudflare"}) + + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/health", nil) + req.RemoteAddr = "10.0.0.5:52000" + req.Header.Set("X-Forwarded-For", "198.51.100.9") + // The edge overwrites this header rather than appending to it, which is what + // makes it usable as an identity at all. + req.Header.Set("CF-Connecting-IP", "203.0.113.7") + app.ServeHTTP(rec, req) + + logged := buf.String() + if !strings.Contains(logged, "clientIp=203.0.113.7") { + t.Fatalf("expected the platform header to win, got: %s", logged) + } + if strings.Contains(logged, "198.51.100.9") { + t.Fatalf("X-Forwarded-For was consulted despite a trusted platform: %s", logged) + } +} + +func TestNewServerRejectsAnInvalidTrustedProxyCIDR(t *testing.T) { + captureRequestLog(t) + setTestServerConfig(config.ServerConfig{TrustedProxies: []string{"10.0.0.0/8", "not-a-cidr"}}) + + if _, err := NewServer(); err == nil { + t.Fatal("NewServer() accepted an unparseable CIDR; a typo here would silently disable the trust boundary") + } else if !strings.Contains(err.Error(), "server.trustedProxies") { + t.Fatalf("NewServer() error = %v, expected it to name server.trustedProxies", err) + } +} diff --git a/internal/pkg/config/config.go b/internal/pkg/config/config.go index 5925d94d..8d224a64 100644 --- a/internal/pkg/config/config.go +++ b/internal/pkg/config/config.go @@ -52,6 +52,65 @@ type ServerConfig struct { CompanyName string `yaml:"companyName"` CompanyLogoURL string `yaml:"companyLogoUrl"` CORS CORSConfig `yaml:"cors"` + // TrustedProxies are the CIDR blocks of the reverse proxies that sit in front + // of the application. Gin's own default is 0.0.0.0/0 and ::/0, which trusts + // every peer and makes ClientIP() return the leftmost X-Forwarded-For value - + // a header any caller can set. + TrustedProxies []string `yaml:"trustedProxies"` + // TrustedPlatform names an edge that overwrites rather than appends the real + // client address, for example "cloudflare". When set it takes precedence over + // X-Forwarded-For entirely. + TrustedPlatform string `yaml:"trustedPlatform"` +} + +// defaultTrustedProxies covers loopback, RFC1918, IPv6 unique-local and +// link-local ranges. That is the shape of almost every real deployment - a +// sidecar tunnel, a compose network, a local nginx - and a client on the public +// internet cannot present one of these addresses as its direct peer, so +// X-Forwarded-For stays honest. When the app is exposed directly the peer is a +// public address, is not trusted, and Gin falls back to it. +var defaultTrustedProxies = []string{ + "127.0.0.0/8", + "10.0.0.0/8", + "172.16.0.0/12", + "192.168.0.0/16", + "::1/128", + "fc00::/7", + "fe80::/10", +} + +func (s ServerConfig) TrustedProxiesOrDefault() []string { + proxies := make([]string, 0, len(s.TrustedProxies)) + for _, proxy := range s.TrustedProxies { + if proxy = strings.TrimSpace(proxy); proxy != "" { + proxies = append(proxies, proxy) + } + } + if len(proxies) == 0 { + return defaultTrustedProxies + } + return proxies +} + +// TrustedPlatformHeader resolves the configured platform name to the header Gin +// should read the client address from. Recognised names map to Gin's own +// constants; any other non-empty value is passed through as a literal header +// name, which is what Gin's TrustedPlatform field expects. +func (s ServerConfig) TrustedPlatformHeader() string { + platform := strings.TrimSpace(s.TrustedPlatform) + if platform == "" { + return "" + } + switch strings.ToLower(platform) { + case "cloudflare", "cf": + return "CF-Connecting-IP" + case "fly.io", "flyio", "fly-io": + return "Fly-Client-IP" + case "google-app-engine", "appengine", "gae": + return "X-Appengine-Remote-Addr" + default: + return platform + } } func (s ServerConfig) Address() string { @@ -86,7 +145,25 @@ type AuthConfig struct { PasswordLoginEnabled *bool `yaml:"passwordLoginEnabled"` TokenTTLHours int `yaml:"tokenTTLHours"` MaxFailedAttempts int `yaml:"maxFailedAttempts"` - CredentialLockMinute int `yaml:"credentialLockMinute"` + // MaxFailedAttemptsPerIP bounds failures from one client address across every + // username, which is what credential stuffing looks like. Zero or unset + // derives four times MaxFailedAttempts; it is disabled when MaxFailedAttempts + // is disabled. + MaxFailedAttemptsPerIP int `yaml:"maxFailedAttemptsPerIP"` + CredentialLockMinute int `yaml:"credentialLockMinute"` +} + +// MaxFailedAttemptsPerIPOrDefault derives the per-address threshold from the +// per-account one so that a deployment which only tunes MaxFailedAttempts still +// gets a coherent pair of limits. +func (a AuthConfig) MaxFailedAttemptsPerIPOrDefault() int { + if a.MaxFailedAttemptsPerIP > 0 { + return a.MaxFailedAttemptsPerIP + } + if a.MaxFailedAttempts <= 0 { + return 0 + } + return a.MaxFailedAttempts * 4 } func (a AuthConfig) IsPasswordLoginEnabled() bool { @@ -301,6 +378,8 @@ func bindConfigDefaults(v *viper.Viper) { v.SetDefault("server.companyName", "") v.SetDefault("server.companyLogoUrl", "") v.SetDefault("server.cors.allowedOrigins", []string{}) + v.SetDefault("server.trustedProxies", []string{}) + v.SetDefault("server.trustedPlatform", "") v.SetDefault("db.type", "sqlite") v.SetDefault("db.dsn", "file:./data/app.db?_busy_timeout=5000") v.SetDefault("db.maxIdleConns", 5) @@ -312,6 +391,7 @@ func bindConfigDefaults(v *viper.Viper) { v.SetDefault("logger.addSource", false) v.SetDefault("auth.tokenTTLHours", 12) v.SetDefault("auth.maxFailedAttempts", 5) + v.SetDefault("auth.maxFailedAttemptsPerIP", 0) v.SetDefault("auth.credentialLockMinute", 15) v.SetDefault("customerSession.ttlMinutes", 120) v.SetDefault("customerSession.refreshThresholdMinutes", 30) @@ -336,6 +416,8 @@ func bindEnvironmentAliases(v *viper.Viper) { _ = v.BindEnv("server.port", "AGENT_DESK_SERVER_PORT", "PORT", "SERVER_PORT") _ = v.BindEnv("server.companyName", "AGENT_DESK_SERVER_COMPANYNAME", "COMPANY_NAME", "NEXT_PUBLIC_COMPANY_NAME", "BRAND_NAME", "BRAND_COMPANY_NAME") _ = v.BindEnv("server.companyLogoUrl", "AGENT_DESK_SERVER_COMPANYLOGOURL", "COMPANY_LOGO_URL", "NEXT_PUBLIC_COMPANY_LOGO_URL", "BRAND_LOGO_URL") + _ = v.BindEnv("server.trustedProxies", "AGENT_DESK_SERVER_TRUSTEDPROXIES", "TRUSTED_PROXIES") + _ = v.BindEnv("server.trustedPlatform", "AGENT_DESK_SERVER_TRUSTEDPLATFORM", "TRUSTED_PLATFORM") _ = v.BindEnv("db.type", "AGENT_DESK_DB_TYPE", "DATABASE_TYPE", "DB_TYPE") _ = v.BindEnv("db.dsn", "AGENT_DESK_DB_DSN", "DATABASE_URL", "DB_DSN") _ = v.BindEnv("auth.passwordLoginEnabled", "AGENT_DESK_AUTH_PASSWORDLOGINENABLED", "PASSWORD_LOGIN_ENABLED") diff --git a/internal/pkg/config/config_test.go b/internal/pkg/config/config_test.go index 00a17f95..de1defe5 100644 --- a/internal/pkg/config/config_test.go +++ b/internal/pkg/config/config_test.go @@ -156,3 +156,107 @@ ORG_SYNC_SECRET=webhook-secret-789 t.Fatalf("Webhook.OrgSyncSecret=%q", cfg.Webhook.OrgSyncSecret) } } + +func TestAuthConfigMaxFailedAttemptsPerIPOrDefault(t *testing.T) { + cases := []struct { + name string + cfg AuthConfig + want int + }{ + {"explicit value wins", AuthConfig{MaxFailedAttempts: 5, MaxFailedAttemptsPerIP: 30}, 30}, + {"unset derives four times the per-account limit", AuthConfig{MaxFailedAttempts: 5}, 20}, + {"disabled per-account limit disables both", AuthConfig{MaxFailedAttempts: 0}, 0}, + {"negative per-account limit disables both", AuthConfig{MaxFailedAttempts: -1}, 0}, + } + for _, tc := range cases { + if got := tc.cfg.MaxFailedAttemptsPerIPOrDefault(); got != tc.want { + t.Errorf("%s: MaxFailedAttemptsPerIPOrDefault() = %d want %d", tc.name, got, tc.want) + } + } +} + +func TestServerConfigTrustedProxiesOrDefault(t *testing.T) { + // An unset list must not fall through to Gin's own default of 0.0.0.0/0 and + // ::/0, which trusts every peer and makes X-Forwarded-For authoritative. + got := ServerConfig{}.TrustedProxiesOrDefault() + for _, want := range []string{"127.0.0.0/8", "10.0.0.0/8", "172.16.0.0/12", "192.168.0.0/16", "::1/128", "fc00::/7", "fe80::/10"} { + found := false + for _, entry := range got { + if entry == want { + found = true + } + } + if !found { + t.Errorf("TrustedProxiesOrDefault() = %v, missing %s", got, want) + } + } + for _, entry := range got { + if entry == "0.0.0.0/0" || entry == "::/0" { + t.Errorf("TrustedProxiesOrDefault() includes %s, which trusts every peer", entry) + } + } + + // Blank entries from a comma-separated environment variable must not become + // empty CIDRs, which Gin would reject at startup. + got = ServerConfig{TrustedProxies: []string{" 10.1.0.0/16 ", "", " "}}.TrustedProxiesOrDefault() + if len(got) != 1 || got[0] != "10.1.0.0/16" { + t.Errorf("TrustedProxiesOrDefault() = %v want [10.1.0.0/16]", got) + } +} + +func TestServerConfigTrustedPlatformHeader(t *testing.T) { + cases := []struct { + platform string + want string + }{ + {"", ""}, + {"cloudflare", "CF-Connecting-IP"}, + {"Cloudflare", "CF-Connecting-IP"}, + {" CF ", "CF-Connecting-IP"}, + {"fly.io", "Fly-Client-IP"}, + {"google-app-engine", "X-Appengine-Remote-Addr"}, + // Anything unrecognised is a literal header name, which is what Gin's + // TrustedPlatform field expects. + {"X-CDN-IP", "X-CDN-IP"}, + } + for _, tc := range cases { + if got := (ServerConfig{TrustedPlatform: tc.platform}).TrustedPlatformHeader(); got != tc.want { + t.Errorf("TrustedPlatformHeader(%q) = %q want %q", tc.platform, got, tc.want) + } + } +} + +// TestLoadReadsTrustedProxySettings covers the environment spelling, including +// the comma-separated list. Gin rejects an unparseable CIDR at startup, so a +// value that arrives as one string instead of a slice would take the process +// down rather than degrade quietly. +func TestLoadReadsTrustedProxySettings(t *testing.T) { + path := filepath.Join(t.TempDir(), "config.yaml") + if err := os.WriteFile(path, []byte("server:\n port: 8083\n"), 0600); err != nil { + t.Fatalf("WriteFile() error = %v", err) + } + + t.Setenv("ENV_FILE", os.DevNull) + t.Setenv("AGENT_DESK_ENV_FILE", os.DevNull) + t.Setenv("TRUSTED_PROXIES", "10.0.0.0/8,172.16.0.0/12") + t.Setenv("TRUSTED_PLATFORM", "cloudflare") + + cfg, err := Load(path) + if err != nil { + t.Fatalf("Load() error = %v", err) + } + + want := []string{"10.0.0.0/8", "172.16.0.0/12"} + got := cfg.Server.TrustedProxies + if len(got) != len(want) { + t.Fatalf("TrustedProxies = %v want %v", got, want) + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("TrustedProxies[%d] = %q want %q", i, got[i], want[i]) + } + } + if cfg.Server.TrustedPlatformHeader() != "CF-Connecting-IP" { + t.Errorf("TrustedPlatformHeader() = %q want CF-Connecting-IP", cfg.Server.TrustedPlatformHeader()) + } +} diff --git a/internal/services/auth_service.go b/internal/services/auth_service.go index 5765768c..4b259459 100644 --- a/internal/services/auth_service.go +++ b/internal/services/auth_service.go @@ -90,7 +90,7 @@ func (s *authService) Login(req request.LoginRequest, authCfg config.AuthConfig, return nil, errorsx.InvalidParamI18n("error.e0258") } - if s.isCredentialLocked(principal, authCfg) { + if s.isCredentialLocked(principal, clientIP, authCfg) { _ = s.createLoginCredentialLog(principal, 0, false, clientIP, userAgent, "credential locked") return nil, errorsx.CredentialLockedI18n("error.e0270") } @@ -455,28 +455,71 @@ func (s *authService) createLoginCredentialLog(principal string, userID int64, s Principal: principal, UserID: userID, Success: success, - ClientIP: clientIP, + ClientIP: normalizeClientIP(clientIP), UserAgent: userAgent, Reason: reason, CreatedAt: time.Now(), }) } -func (s *authService) isCredentialLocked(principal string, authCfg config.AuthConfig) bool { - maxFailedAttempts := authCfg.MaxFailedAttempts - if maxFailedAttempts <= 0 { - return false - } +// isCredentialLocked reports whether this attempt should be refused before the +// password is even checked. +// +// Both windows are keyed on the client address as well as the username. Keying on +// the username alone - which is what this used to do - means anyone who knows a +// username can lock the real account out for the whole window, from anywhere, as +// often as they like: that is a denial of service dressed up as a protection. The +// principal-and-address window still stops one source grinding on one account, and +// the address-only window across all principals still stops credential stuffing, +// but a legitimate user arriving from a different address is no longer locked out +// by somebody else's failures. +// +// This is only sound because the address itself is trustworthy; see the trusted +// proxy configuration in bootstrap.NewServer. +func (s *authService) isCredentialLocked(principal, clientIP string, authCfg config.AuthConfig) bool { lockMinute := authCfg.CredentialLockMinute if lockMinute <= 0 { lockMinute = 15 } since := time.Now().Add(-time.Duration(lockMinute) * time.Minute) - return LoginCredentialLogService.Count(sqls.NewCnd(). - Eq("principal", normalizeLoginPrincipal(principal)). - Eq("success", false). - NotEq("reason", "credential locked"). - Where("created_at >= ?", since)) >= int64(maxFailedAttempts) + ip := normalizeClientIP(clientIP) + + base := func() *sqls.Cnd { + return sqls.NewCnd(). + Eq("client_ip", ip). + Eq("success", false). + NotEq("reason", "credential locked"). + Where("created_at >= ?", since) + } + + if maxFailedAttempts := authCfg.MaxFailedAttempts; maxFailedAttempts > 0 { + count := LoginCredentialLogService.Count(base().Eq("principal", normalizeLoginPrincipal(principal))) + if count >= int64(maxFailedAttempts) { + return true + } + } + + // An undetermined address would otherwise pool every such caller into one + // bucket and lock them all out together. + if maxPerIP := authCfg.MaxFailedAttemptsPerIPOrDefault(); maxPerIP > 0 && ip != unknownClientIP { + if count := LoginCredentialLogService.Count(base()); count >= int64(maxPerIP) { + return true + } + } + + return false +} + +// unknownClientIP stands in for an address the server could not determine, so +// those attempts share one bucket on the principal window instead of matching +// nothing at all. +const unknownClientIP = "unknown" + +func normalizeClientIP(clientIP string) string { + if clientIP = strings.TrimSpace(clientIP); clientIP == "" { + return unknownClientIP + } + return clientIP } func normalizeLoginPrincipal(principal string) string { diff --git a/internal/services/auth_service_test.go b/internal/services/auth_service_test.go index 156db6d5..d137323d 100644 --- a/internal/services/auth_service_test.go +++ b/internal/services/auth_service_test.go @@ -112,6 +112,7 @@ func TestAuthServiceLoginCredentialLockout(t *testing.T) { Principal: "admin", UserID: user.ID, Success: false, + ClientIP: "127.0.0.1", Reason: "password mismatch", CreatedAt: now.Add(-time.Duration(i+1) * time.Minute), }).Error; err != nil { @@ -122,6 +123,7 @@ func TestAuthServiceLoginCredentialLockout(t *testing.T) { Principal: "admin", UserID: user.ID, Success: false, + ClientIP: "127.0.0.1", Reason: "password mismatch", CreatedAt: now.Add(-30 * time.Minute), }).Error; err != nil { @@ -164,6 +166,7 @@ func TestAuthServiceCredentialLockoutDoesNotExtendWhileLocked(t *testing.T) { Principal: "admin", UserID: 1, Success: false, + ClientIP: "127.0.0.1", Reason: "password mismatch", CreatedAt: now.Add(-2 * time.Minute), }, @@ -171,6 +174,7 @@ func TestAuthServiceCredentialLockoutDoesNotExtendWhileLocked(t *testing.T) { Principal: "admin", UserID: 0, Success: false, + ClientIP: "127.0.0.1", Reason: "credential locked", CreatedAt: now.Add(-1 * time.Minute), }, @@ -199,6 +203,7 @@ func TestAuthServiceCredentialLockoutNormalizesPrincipalCase(t *testing.T) { Principal: "admin", UserID: 1, Success: false, + ClientIP: "127.0.0.1", Reason: "password mismatch", CreatedAt: time.Now().Add(-time.Minute), }).Error; err != nil { @@ -223,6 +228,79 @@ func TestAuthServiceCredentialLockoutNormalizesPrincipalCase(t *testing.T) { } } +// TestAuthServiceCredentialLockoutIsScopedToTheClientAddress is the regression +// test for the denial of service. Keyed on the username alone, anybody who knew a +// username could lock the real account out for the whole window, from anywhere, +// as often as they liked. +func TestAuthServiceCredentialLockoutIsScopedToTheClientAddress(t *testing.T) { + db := setupAuthServiceTestDB(t) + user := createAuthTestUser(t, db, "admin", "secret") + now := time.Now() + for i := 0; i < 5; i++ { + if err := db.Create(&models.LoginCredentialLog{ + Principal: "admin", + UserID: user.ID, + Success: false, + ClientIP: "203.0.113.66", + Reason: "password mismatch", + CreatedAt: now.Add(-time.Duration(i+1) * time.Minute), + }).Error; err != nil { + t.Fatalf("seed credential log: %v", err) + } + } + + authCfg := config.AuthConfig{TokenTTLHours: 2, MaxFailedAttempts: 3, CredentialLockMinute: 15} + + if _, err := newAuthService().Login(request.LoginRequest{Username: "admin", Password: "secret"}, authCfg, "203.0.113.66", "go-test"); !hasCode(err, errorsx.CodeAuthCredentialLocked) { + t.Fatalf("expected the attacking address to be locked, got %v", err) + } + + ret, err := newAuthService().Login(request.LoginRequest{Username: "admin", Password: "secret"}, authCfg, "198.51.100.7", "go-test") + if err != nil { + t.Fatalf("the legitimate owner was locked out by somebody else's failures: %v", err) + } + if ret == nil || ret.AccessToken == "" { + t.Fatalf("expected a session for the legitimate owner, got %+v", ret) + } +} + +// TestAuthServiceCredentialLockoutByAddressAcrossPrincipals covers the window that +// replaces the removed account-wide lock: one address spraying many usernames. +func TestAuthServiceCredentialLockoutByAddressAcrossPrincipals(t *testing.T) { + db := setupAuthServiceTestDB(t) + createAuthTestUser(t, db, "admin", "secret") + now := time.Now() + for i, principal := range []string{"admin", "root", "operator", "support"} { + if err := db.Create(&models.LoginCredentialLog{ + Principal: principal, + UserID: 0, + Success: false, + ClientIP: "203.0.113.66", + Reason: "user not found", + CreatedAt: now.Add(-time.Duration(i+1) * time.Minute), + }).Error; err != nil { + t.Fatalf("seed credential log: %v", err) + } + } + + authCfg := config.AuthConfig{ + TokenTTLHours: 2, + MaxFailedAttempts: 10, + MaxFailedAttemptsPerIP: 4, + CredentialLockMinute: 15, + } + + // No single username has reached MaxFailedAttempts, so only the per-address + // window can catch this. + if _, err := newAuthService().Login(request.LoginRequest{Username: "admin", Password: "secret"}, authCfg, "203.0.113.66", "go-test"); !hasCode(err, errorsx.CodeAuthCredentialLocked) { + t.Fatalf("expected credential stuffing from one address to be locked, got %v", err) + } + + if _, err := newAuthService().Login(request.LoginRequest{Username: "admin", Password: "secret"}, authCfg, "198.51.100.7", "go-test"); err != nil { + t.Fatalf("an unrelated address was locked by another address's failures: %v", err) + } +} + func TestAuthServiceCredentialLockoutDisabledWhenMaxAttemptsNonPositive(t *testing.T) { db := setupAuthServiceTestDB(t) user := createAuthTestUser(t, db, "admin", "secret") @@ -232,6 +310,7 @@ func TestAuthServiceCredentialLockoutDisabledWhenMaxAttemptsNonPositive(t *testi Principal: "admin", UserID: user.ID, Success: false, + ClientIP: "127.0.0.1", Reason: "credential locked", CreatedAt: now.Add(-time.Duration(i+1) * time.Minute), }).Error; err != nil { From c42d2de28c1848051395d4caf4553216a5997d74 Mon Sep 17 00:00:00 2001 From: JOY <5027251+JOY@users.noreply.github.com> Date: Sun, 13 Sep 2026 17:15:37 +0700 Subject: [PATCH 6/8] feat(security): rate limit the public endpoints an anonymous caller can drive There was no rate limiting anywhere in the application. Every unauthenticated endpoint could be called in a tight loop, which made three things cheap: flooding the message upload endpoints with 20 MB bodies, spraying login and support registration, and burning LLM spend through the conversation endpoints. Add a small in-process fixed-window limiter and attach it to six public routes: login (20 per window), support registration (10), the widget session exchange (120), message attachment and image upload (30, shared between the two), and documentation feedback (20). The window is 60 seconds by default, and both it and the on/off switch are configurable. The limits are sized for a human driving a browser, not for a machine. The session exchange budget in particular has to absorb a whole support office loading the widget from behind one NAT address, which is why it sits an order of magnitude above the others. Deliberately not limited: - /api/third/* channel webhooks. A platform that receives a 429 from its webhook stops retrying and eventually disables the delivery, which takes an entire channel offline - a far worse outcome than the flood the limit was meant to stop. Those endpoints authenticate by signature instead. - /api/ws/* websockets, and /api/webhooks/* which are HMAC verified. - /api/dashboard/* - already behind AuthMiddleware, and staff sharing one office address would throttle each other. Rejections return 429 with a Retry-After header and an ordinary JsonResult body. The status is a real 429 rather than the 200-with-error-code the auth middleware uses, because web/lib/api/client.ts parses the payload before it inspects response.ok and surfaces payload.message, so the localized text still reaches the user - and Retry-After is only meaningful on a 429 or 503. The message is error.e0354 in both backend locales and does not disclose the limit or the remaining budget, which would only tell a caller how much room they have left. Retry-After rounds the remaining window up rather than down. Telling a caller to come back sooner than the window actually resets just earns another 429. Rejections are not logged separately. requestLogMiddleware already records path, status and client address for every request, so a 429 is visible there without giving a flood a second way to fill the log. The limiter keys on ctx.ClientIP(), which is only trustworthy because of the trusted-proxy configuration that landed immediately before this. Without it a caller would pick their own bucket with an X-Forwarded-For header. Counters are per process, and expired buckets are swept lazily from inside Allow rather than by a goroutine, so there is no background lifetime to own and no unbounded growth. Running several replicas gives each its own budget, weakening the bound by the replica count; config.example.yaml says so explicitly. These limits exist to make flooding expensive, not to meter a quota. Tests: seven for the limiter, including an exact-count check across 16 goroutines making 3200 calls against one key; three for the middleware, covering the 429 shape, per-address isolation and a nil limiter allowing everything; and five end to end against a server built by NewServer, including one that fires 1000 requests at the health, config, org-sync and two channel webhook routes and fails if any of them returns 429. The concurrency test could not be run under -race: this repository builds with CGO disabled and go test -race requires cgo. It still has teeth, because an unguarded concurrent map write panics rather than merely miscounting. --- .env.example | 7 + config/config.example.yaml | 15 ++ internal/bootstrap/routes.go | 75 ++++++-- internal/bootstrap/server.go | 10 +- internal/bootstrap/server_ratelimit_test.go | 162 ++++++++++++++++++ internal/middleware/ratelimit_middleware.go | 56 ++++++ .../middleware/ratelimit_middleware_test.go | 103 +++++++++++ internal/pkg/config/config.go | 33 +++- internal/pkg/config/config_test.go | 51 ++++++ internal/pkg/i18nx/locales/en-US.yml | 1 + internal/pkg/i18nx/locales/zh-CN.yml | 1 + internal/pkg/ratelimit/ratelimit.go | 101 +++++++++++ internal/pkg/ratelimit/ratelimit_test.go | 132 ++++++++++++++ 13 files changed, 732 insertions(+), 15 deletions(-) create mode 100644 internal/bootstrap/server_ratelimit_test.go create mode 100644 internal/middleware/ratelimit_middleware.go create mode 100644 internal/middleware/ratelimit_middleware_test.go create mode 100644 internal/pkg/ratelimit/ratelimit.go create mode 100644 internal/pkg/ratelimit/ratelimit_test.go diff --git a/.env.example b/.env.example index edbe2b24..0e8968c1 100644 --- a/.env.example +++ b/.env.example @@ -20,6 +20,13 @@ COMPANY_LOGO_URL= # TRUSTED_PROXIES="10.0.0.0/8,172.16.0.0/12" TRUSTED_PLATFORM=cloudflare +# Rate limiting for the public, unauthenticated endpoints (login, support +# registration, widget session exchange, message uploads, doc feedback). Enabled +# by default; the counters are per process. Channel webhooks, websockets and the +# authenticated dashboard are deliberately exempt. +# RATE_LIMIT_ENABLED=true +# RATE_LIMIT_WINDOW_SECONDS=60 + # Database Configuration # Driver options: sqlite, mysql, postgres DB_TYPE=sqlite diff --git a/config/config.example.yaml b/config/config.example.yaml index d719f25d..91eb2943 100644 --- a/config/config.example.yaml +++ b/config/config.example.yaml @@ -17,6 +17,21 @@ server: # header name) when an edge overwrites rather than appends the real client # address. When set it takes precedence over X-Forwarded-For entirely. trustedPlatform: "" + rateLimit: + # Bounds how often one client address may call the public, unauthenticated + # endpoints: login, support registration, the widget session exchange, message + # uploads and documentation feedback. Enabled by default. + # + # Channel webhooks under /api/third, the websocket routes, /api/webhooks and + # the whole authenticated dashboard are deliberately NOT covered. A platform + # that receives a 429 from its webhook stops retrying and eventually disables + # the delivery, which takes an entire channel offline. + # + # The counters live in this process. Running several replicas gives each one + # its own budget, which weakens the bound by the replica count; these limits + # exist to make flooding expensive, not to meter a quota. + enabled: true + windowSeconds: 60 cors: # Browser CORS allowlist. In production, replace this with the actual frontend or embedded-site domains, such as https://support.example.com. # Leave it empty to reject cross-origin browser requests. Same-origin and non-browser calls are still supported. diff --git a/internal/bootstrap/routes.go b/internal/bootstrap/routes.go index 1482e841..f02307b6 100644 --- a/internal/bootstrap/routes.go +++ b/internal/bootstrap/routes.go @@ -1,15 +1,66 @@ package bootstrap import ( + "time" + "agent-desk/internal/handlers/api" "agent-desk/internal/handlers/dashboard" "agent-desk/internal/handlers/third" + "agent-desk/internal/middleware" + "agent-desk/internal/pkg/config" + "agent-desk/internal/pkg/ratelimit" "github.com/gin-gonic/gin" ) -func registerApiAuthRoutes(group *gin.RouterGroup) { - group.POST("/login", api.Login) +// publicRateLimits holds one limiter per abuse-prone unauthenticated endpoint. +// They are built once per server rather than per request, and the middleware keys +// them by client address. +// +// Channel webhooks under /api/third, the websocket routes and /api/webhooks are +// deliberately absent. A platform that receives a 429 from its webhook stops +// retrying and eventually disables the delivery, which would take a whole channel +// offline; those endpoints authenticate by signature instead. The dashboard is +// absent too - it is already behind AuthMiddleware, and staff sharing one office +// address would throttle each other. +type publicRateLimits struct { + login *ratelimit.Limiter + supportRegister *ratelimit.Limiter + sessionExchange *ratelimit.Limiter + upload *ratelimit.Limiter + docFeedback *ratelimit.Limiter +} + +// Limits are sized for a human driving a browser, not for a machine. They are +// deliberately generous: the point is to make flooding expensive, not to police +// legitimate use. The session exchange budget in particular has to absorb a whole +// support office loading the widget from behind one NAT address. +const ( + limitLogin = 20 + limitSupportRegister = 10 + limitSessionExchange = 120 + limitUpload = 30 + limitDocFeedback = 20 +) + +func newPublicRateLimits(cfg config.RateLimitConfig) publicRateLimits { + if !cfg.IsEnabled() { + // Every field stays nil and a nil limiter allows everything, so "disabled" + // needs no branch at any call site. + return publicRateLimits{} + } + window := time.Duration(cfg.WindowSecondsOrDefault()) * time.Second + return publicRateLimits{ + login: ratelimit.New(limitLogin, window), + supportRegister: ratelimit.New(limitSupportRegister, window), + sessionExchange: ratelimit.New(limitSessionExchange, window), + upload: ratelimit.New(limitUpload, window), + docFeedback: ratelimit.New(limitDocFeedback, window), + } +} + +func registerApiAuthRoutes(group *gin.RouterGroup, limits publicRateLimits) { + group.POST("/login", middleware.RateLimit(limits.login), api.Login) group.POST("/logout", api.Logout) group.GET("/profile", api.Profile) group.POST("/profile/update", api.UpdateProfile) @@ -33,8 +84,8 @@ func registerApiWebhookRoutes(group *gin.RouterGroup) { group.POST("/dos-org-sync", api.DOSOrgSyncWebhook) } -func registerApiCustomerRoutes(group *gin.RouterGroup) { - group.POST("/session_exchange", api.CustomerPostSession_exchange) +func registerApiCustomerRoutes(group *gin.RouterGroup, limits publicRateLimits) { + group.POST("/session_exchange", middleware.RateLimit(limits.sessionExchange), api.CustomerPostSession_exchange) } func registerApiConversationRoutes(group *gin.RouterGroup) { @@ -43,22 +94,26 @@ func registerApiConversationRoutes(group *gin.RouterGroup) { group.POST("/create_or_match", api.ConversationPostCreate_or_match) } -func registerApiMessageRoutes(group *gin.RouterGroup) { +func registerApiMessageRoutes(group *gin.RouterGroup, limits publicRateLimits) { group.Any("/list", api.MessageAnyList) group.POST("/read", api.MessagePostRead) group.POST("/send", api.MessagePostSend) - group.POST("/upload_attachment", api.MessagePostUpload_attachment) - group.POST("/upload_image", api.MessagePostUpload_image) + // One shared budget for both upload routes: what matters is how many bytes an + // unauthenticated caller can push at the storage layer, not which of the two + // endpoints they used. + uploadLimit := middleware.RateLimit(limits.upload) + group.POST("/upload_attachment", uploadLimit, api.MessagePostUpload_attachment) + group.POST("/upload_image", uploadLimit, api.MessagePostUpload_image) } -func registerApiSupportRoutes(group *gin.RouterGroup) { +func registerApiSupportRoutes(group *gin.RouterGroup, limits publicRateLimits) { group.GET("/config", api.SupportConfigGetConfig) - group.POST("/auth/register", api.SupportAuthPostRegister) + group.POST("/auth/register", middleware.RateLimit(limits.supportRegister), api.SupportAuthPostRegister) group.GET("/me", api.SupportGetMe) group.Any("/doc-page/list", api.DocPageAnyList) group.GET("/doc-page/navigation", api.DocPageGetNavigation) group.GET("/doc-page/:id", api.DocPageGetBy) - group.POST("/doc-page/feedback", api.DocPagePostFeedback) + group.POST("/doc-page/feedback", middleware.RateLimit(limits.docFeedback), api.DocPagePostFeedback) group.Any("/community/categories/list", api.CategoryAnyList) group.Any("/community/posts/list", api.PostAnyList) group.GET("/community/posts/:id", api.PostGetBy) diff --git a/internal/bootstrap/server.go b/internal/bootstrap/server.go index 16416620..49d79580 100644 --- a/internal/bootstrap/server.go +++ b/internal/bootstrap/server.go @@ -171,17 +171,19 @@ func isWebsocketUpgrade(ctx *gin.Context) bool { } func addRouter(app *gin.Engine) { + limits := newPublicRateLimits(config.Current().Server.RateLimit) + app.Any("/api/mcp", gin.WrapH(mcps.NewHTTPHandler())) apiGroup := app.Group("/api") apiGroup.GET("/health", api.Health) apiGroup.GET("/config", api.PublicConfig) - registerApiAuthRoutes(apiGroup.Group("/auth")) + registerApiAuthRoutes(apiGroup.Group("/auth"), limits) registerApiChannelRoutes(apiGroup.Group("/channel")) - registerApiCustomerRoutes(apiGroup.Group("/customer")) + registerApiCustomerRoutes(apiGroup.Group("/customer"), limits) registerApiConversationRoutes(apiGroup.Group("/conversation", middleware.ExternalUserMiddleware)) - registerApiMessageRoutes(apiGroup.Group("/message", middleware.ExternalUserMiddleware)) - registerApiSupportRoutes(apiGroup.Group("/support")) + registerApiMessageRoutes(apiGroup.Group("/message", middleware.ExternalUserMiddleware), limits) + registerApiSupportRoutes(apiGroup.Group("/support"), limits) registerApiWebhookRoutes(apiGroup.Group("/webhooks")) wsGroup := app.Group("/api/ws") diff --git a/internal/bootstrap/server_ratelimit_test.go b/internal/bootstrap/server_ratelimit_test.go new file mode 100644 index 00000000..f1ae3f14 --- /dev/null +++ b/internal/bootstrap/server_ratelimit_test.go @@ -0,0 +1,162 @@ +package bootstrap + +import ( + "net/http" + "net/http/httptest" + "strconv" + "strings" + "testing" + + "agent-desk/internal/pkg/config" + + "github.com/gin-gonic/gin" +) + +func parseInt(t *testing.T, value string) int { + t.Helper() + parsed, err := strconv.Atoi(value) + if err != nil { + t.Fatalf("%q is not an integer", value) + } + return parsed +} + +// setRateLimitTestConfig disables password login so /api/auth/login returns +// before it reaches the database. The limiter runs ahead of the handler either +// way, which is what these tests are about. +func setRateLimitTestConfig(rateLimit config.RateLimitConfig) { + disabled := false + config.SetCurrent(&config.Config{ + Server: config.ServerConfig{ + RateLimit: rateLimit, + CORS: config.CORSConfig{AllowedOrigins: []string{}}, + }, + Auth: config.AuthConfig{PasswordLoginEnabled: &disabled}, + Storage: config.StorageConfig{Local: config.LocalStorageConfig{Root: "storage", BaseURL: "/storage"}}, + }) +} + +func postJSON(app *gin.Engine, path, clientIP string) *httptest.ResponseRecorder { + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(`{"username":"admin","password":"secret"}`)) + req.Header.Set("Content-Type", "application/json") + req.RemoteAddr = clientIP + ":52000" + app.ServeHTTP(rec, req) + return rec +} + +func TestNewServerRateLimitsThePublicLoginEndpoint(t *testing.T) { + setRateLimitTestConfig(config.RateLimitConfig{}) + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + for i := 1; i <= limitLogin; i++ { + if rec := postJSON(app, "/api/auth/login", "203.0.113.7"); rec.Code == http.StatusTooManyRequests { + t.Fatalf("request %d was throttled although the login limit is %d per window", i, limitLogin) + } + } + rec := postJSON(app, "/api/auth/login", "203.0.113.7") + if rec.Code != http.StatusTooManyRequests { + t.Fatalf("request %d got status %d, want 429", limitLogin+1, rec.Code) + } + if rec.Header().Get("Retry-After") == "" { + t.Error("the 429 carried no Retry-After header") + } + + // A second address must keep working, otherwise one attacker could lock the + // login page for every user behind their own NAT. + if rec := postJSON(app, "/api/auth/login", "198.51.100.9"); rec.Code == http.StatusTooManyRequests { + t.Fatal("an unrelated address was throttled by another address's requests") + } +} + +func TestNewServerRateLimitsTheSupportRegistrationEndpoint(t *testing.T) { + setRateLimitTestConfig(config.RateLimitConfig{}) + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + for i := 1; i <= limitSupportRegister; i++ { + postJSON(app, "/api/support/auth/register", "203.0.113.7") + } + if rec := postJSON(app, "/api/support/auth/register", "203.0.113.7"); rec.Code != http.StatusTooManyRequests { + t.Fatalf("registration attempt %d got status %d, want 429", limitSupportRegister+1, rec.Code) + } +} + +// TestNewServerLeavesWebhooksAndPublicReadsUnthrottled guards the exemption that +// matters most. A channel platform that receives a 429 from its webhook stops +// retrying and eventually disables the delivery, which takes a whole channel +// offline - a far worse outcome than the flood the limit was meant to stop. +func TestNewServerLeavesWebhooksAndPublicReadsUnthrottled(t *testing.T) { + setRateLimitTestConfig(config.RateLimitConfig{}) + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + requests := []struct { + method string + path string + }{ + {http.MethodGet, "/api/health"}, + {http.MethodGet, "/api/config"}, + {http.MethodPost, "/api/webhooks/org-sync"}, + {http.MethodGet, "/api/third/telegram/webhook"}, + {http.MethodPost, "/api/third/whatsapp/webhook"}, + } + for _, r := range requests { + for i := 0; i < 200; i++ { + rec := httptest.NewRecorder() + req := httptest.NewRequest(r.method, r.path, strings.NewReader(`{}`)) + req.Header.Set("Content-Type", "application/json") + req.RemoteAddr = "203.0.113.7:52000" + app.ServeHTTP(rec, req) + // Any status is acceptable here except 429: these routes either + // succeed, fail validation, or error out, but they must never be + // throttled. + if rec.Code == http.StatusTooManyRequests { + t.Fatalf("%s %s returned 429 on request %d; this route must stay unthrottled", r.method, r.path, i+1) + } + } + } +} + +func TestNewServerRateLimitCanBeDisabled(t *testing.T) { + disabled := false + setRateLimitTestConfig(config.RateLimitConfig{Enabled: &disabled}) + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + for i := 0; i < limitLogin*4; i++ { + if rec := postJSON(app, "/api/auth/login", "203.0.113.7"); rec.Code == http.StatusTooManyRequests { + t.Fatalf("request %d was throttled although rate limiting is disabled", i+1) + } + } +} + +func TestNewServerRateLimitWindowIsConfigurable(t *testing.T) { + setRateLimitTestConfig(config.RateLimitConfig{WindowSeconds: 3600}) + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + for i := 1; i <= limitLogin; i++ { + postJSON(app, "/api/auth/login", "203.0.113.7") + } + rec := postJSON(app, "/api/auth/login", "203.0.113.7") + if rec.Code != http.StatusTooManyRequests { + t.Fatalf("got status %d, want 429", rec.Code) + } + if retryAfter := rec.Header().Get("Retry-After"); retryAfter == "" { + t.Fatal("no Retry-After header") + } else if seconds := parseInt(t, retryAfter); seconds < 3500 || seconds > 3600 { + t.Errorf("Retry-After = %s, want roughly the configured 3600 second window", retryAfter) + } +} diff --git a/internal/middleware/ratelimit_middleware.go b/internal/middleware/ratelimit_middleware.go new file mode 100644 index 00000000..506b31b7 --- /dev/null +++ b/internal/middleware/ratelimit_middleware.go @@ -0,0 +1,56 @@ +package middleware + +import ( + "net/http" + "strconv" + "time" + + "agent-desk/internal/pkg/errorsx" + "agent-desk/internal/pkg/httpx" + "agent-desk/internal/pkg/ratelimit" + + "github.com/gin-gonic/gin" +) + +// rateLimitMessageKey is the localized explanation returned with a 429. It +// deliberately does not disclose the limit or the remaining budget, which would +// only tell a caller how much room they have left. +const rateLimitMessageKey = "error.e0354" + +// RateLimit bounds how often one client address may call a single public +// endpoint. +// +// The address comes from ctx.ClientIP(), which is only meaningful because +// NewServer configures Gin's trusted proxies and platform before any middleware +// is registered. Without that a caller picks their own bucket by setting +// X-Forwarded-For, and the limit is decoration. +// +// A nil limiter allows everything, so "disabled" is expressed by handing out nil +// limiters rather than by branching at every call site. +// +// Rejections are not logged here. requestLogMiddleware already records the path, +// status and client address of every request, so a 429 is visible there without +// giving a flood a second way to fill the log. +func RateLimit(limiter *ratelimit.Limiter) gin.HandlerFunc { + return func(ctx *gin.Context) { + allowed, retryAfter := limiter.Allow(ctx.ClientIP()) + if allowed { + ctx.Next() + return + } + // Retry-After is a whole number of seconds, so round up: telling a caller + // to come back sooner than the window actually resets just earns another + // 429. Rounding down and adding one overshoots when the remaining time is + // already an exact number of seconds. + seconds := int64(retryAfter / time.Second) + if time.Duration(seconds)*time.Second < retryAfter { + seconds++ + } + if seconds < 1 { + seconds = 1 + } + ctx.Header("Retry-After", strconv.FormatInt(seconds, 10)) + httpx.WriteHttpStatusJSON(ctx, http.StatusTooManyRequests, errorsx.InvalidParamI18n(rateLimitMessageKey)) + ctx.Abort() + } +} diff --git a/internal/middleware/ratelimit_middleware_test.go b/internal/middleware/ratelimit_middleware_test.go new file mode 100644 index 00000000..e2aa1cbf --- /dev/null +++ b/internal/middleware/ratelimit_middleware_test.go @@ -0,0 +1,103 @@ +package middleware + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strconv" + "testing" + "time" + + "agent-desk/internal/pkg/ratelimit" + + "github.com/gin-gonic/gin" +) + +func newRateLimitTestEngine(limiter *ratelimit.Limiter) *gin.Engine { + gin.SetMode(gin.TestMode) + app := gin.New() + app.POST("/probe", RateLimit(limiter), func(ctx *gin.Context) { + ctx.JSON(http.StatusOK, gin.H{"success": true}) + }) + return app +} + +func postProbe(app *gin.Engine, clientIP string) *httptest.ResponseRecorder { + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/probe", nil) + req.RemoteAddr = clientIP + ":52000" + app.ServeHTTP(rec, req) + return rec +} + +func TestRateLimitRefusesAfterTheLimitAndSetsRetryAfter(t *testing.T) { + app := newRateLimitTestEngine(ratelimit.New(2, time.Minute)) + + for i := 1; i <= 2; i++ { + if rec := postProbe(app, "203.0.113.7"); rec.Code != http.StatusOK { + t.Fatalf("request %d got status %d, want 200", i, rec.Code) + } + } + + rec := postProbe(app, "203.0.113.7") + if rec.Code != http.StatusTooManyRequests { + t.Fatalf("the third request got status %d, want 429", rec.Code) + } + + retryAfter := rec.Header().Get("Retry-After") + if retryAfter == "" { + t.Fatal("a 429 without Retry-After gives the caller nothing to act on") + } + seconds, err := strconv.Atoi(retryAfter) + if err != nil { + t.Fatalf("Retry-After = %q is not a number of seconds", retryAfter) + } + // The window is one minute and the requests are microseconds apart, so the + // answer is 60 give or take a clock tick. What must not happen is a value + // above the window, which would tell the caller to wait longer than needed, or + // zero, which would have them retry immediately and collect another 429. + if seconds < 1 || seconds > 60 { + t.Errorf("Retry-After = %d seconds, want between 1 and the 60 second window", seconds) + } + + // The body has to stay a JsonResult, because web/lib/api/client.ts parses the + // payload before it looks at response.ok and surfaces payload.message. + var body struct { + Success bool `json:"success"` + Message string `json:"message"` + } + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("429 body is not JSON: %v (%s)", err, rec.Body.String()) + } + if body.Success { + t.Error("429 body reported success=true") + } + if body.Message == "" { + t.Error("429 body carried no message, so the caller would see a blank error") + } +} + +func TestRateLimitIsKeyedByClientAddress(t *testing.T) { + app := newRateLimitTestEngine(ratelimit.New(1, time.Minute)) + + if rec := postProbe(app, "203.0.113.7"); rec.Code != http.StatusOK { + t.Fatalf("first request got %d, want 200", rec.Code) + } + if rec := postProbe(app, "203.0.113.7"); rec.Code != http.StatusTooManyRequests { + t.Fatalf("second request from the same address got %d, want 429", rec.Code) + } + // Somebody else must not inherit the first caller's exhaustion. Behind a NAT + // this is the difference between a working product and a shared lockout. + if rec := postProbe(app, "198.51.100.9"); rec.Code != http.StatusOK { + t.Fatalf("a different address got %d, want 200", rec.Code) + } +} + +func TestRateLimitWithNilLimiterAllowsEverything(t *testing.T) { + app := newRateLimitTestEngine(nil) + for i := 0; i < 50; i++ { + if rec := postProbe(app, "203.0.113.7"); rec.Code != http.StatusOK { + t.Fatalf("request %d got %d with a nil limiter, want 200", i, rec.Code) + } + } +} diff --git a/internal/pkg/config/config.go b/internal/pkg/config/config.go index 8d224a64..ef17610e 100644 --- a/internal/pkg/config/config.go +++ b/internal/pkg/config/config.go @@ -60,7 +60,8 @@ type ServerConfig struct { // TrustedPlatform names an edge that overwrites rather than appends the real // client address, for example "cloudflare". When set it takes precedence over // X-Forwarded-For entirely. - TrustedPlatform string `yaml:"trustedPlatform"` + TrustedPlatform string `yaml:"trustedPlatform"` + RateLimit RateLimitConfig `yaml:"rateLimit"` } // defaultTrustedProxies covers loopback, RFC1918, IPv6 unique-local and @@ -126,6 +127,33 @@ type CORSConfig struct { AllowedOrigins []string `yaml:"allowedOrigins"` } +// RateLimitConfig bounds how often one client address may call the public, +// unauthenticated endpoints. It deliberately does not cover channel webhooks, +// websockets or authenticated dashboard routes: a platform that receives a 429 +// from a webhook endpoint stops retrying and eventually disables the webhook. +type RateLimitConfig struct { + // Enabled defaults to true. Set it to false to switch the limits off without + // taking them out of the route table. + Enabled *bool `yaml:"enabled"` + // WindowSeconds is the length of the counting window. Zero or negative means + // one minute. + WindowSeconds int `yaml:"windowSeconds"` +} + +func (r RateLimitConfig) IsEnabled() bool { + if r.Enabled == nil { + return true + } + return *r.Enabled +} + +func (r RateLimitConfig) WindowSecondsOrDefault() int { + if r.WindowSeconds <= 0 { + return 60 + } + return r.WindowSeconds +} + type DBConfig struct { Type string `yaml:"type"` DSN string `yaml:"dsn"` @@ -380,6 +408,7 @@ func bindConfigDefaults(v *viper.Viper) { v.SetDefault("server.cors.allowedOrigins", []string{}) v.SetDefault("server.trustedProxies", []string{}) v.SetDefault("server.trustedPlatform", "") + v.SetDefault("server.rateLimit.windowSeconds", 60) v.SetDefault("db.type", "sqlite") v.SetDefault("db.dsn", "file:./data/app.db?_busy_timeout=5000") v.SetDefault("db.maxIdleConns", 5) @@ -418,6 +447,8 @@ func bindEnvironmentAliases(v *viper.Viper) { _ = v.BindEnv("server.companyLogoUrl", "AGENT_DESK_SERVER_COMPANYLOGOURL", "COMPANY_LOGO_URL", "NEXT_PUBLIC_COMPANY_LOGO_URL", "BRAND_LOGO_URL") _ = v.BindEnv("server.trustedProxies", "AGENT_DESK_SERVER_TRUSTEDPROXIES", "TRUSTED_PROXIES") _ = v.BindEnv("server.trustedPlatform", "AGENT_DESK_SERVER_TRUSTEDPLATFORM", "TRUSTED_PLATFORM") + _ = v.BindEnv("server.rateLimit.enabled", "AGENT_DESK_SERVER_RATELIMIT_ENABLED", "RATE_LIMIT_ENABLED") + _ = v.BindEnv("server.rateLimit.windowSeconds", "AGENT_DESK_SERVER_RATELIMIT_WINDOWSECONDS", "RATE_LIMIT_WINDOW_SECONDS") _ = v.BindEnv("db.type", "AGENT_DESK_DB_TYPE", "DATABASE_TYPE", "DB_TYPE") _ = v.BindEnv("db.dsn", "AGENT_DESK_DB_DSN", "DATABASE_URL", "DB_DSN") _ = v.BindEnv("auth.passwordLoginEnabled", "AGENT_DESK_AUTH_PASSWORDLOGINENABLED", "PASSWORD_LOGIN_ENABLED") diff --git a/internal/pkg/config/config_test.go b/internal/pkg/config/config_test.go index de1defe5..e6874258 100644 --- a/internal/pkg/config/config_test.go +++ b/internal/pkg/config/config_test.go @@ -260,3 +260,54 @@ func TestLoadReadsTrustedProxySettings(t *testing.T) { t.Errorf("TrustedPlatformHeader() = %q want CF-Connecting-IP", cfg.Server.TrustedPlatformHeader()) } } + +func TestRateLimitConfigDefaultsToEnabledWithAMinuteWindow(t *testing.T) { + var cfg RateLimitConfig + if !cfg.IsEnabled() { + t.Error("rate limiting must default to enabled; a zero-value config should not silently turn protection off") + } + if got := cfg.WindowSecondsOrDefault(); got != 60 { + t.Errorf("WindowSecondsOrDefault() = %d want 60", got) + } + + off := false + cfg = RateLimitConfig{Enabled: &off, WindowSeconds: -5} + if cfg.IsEnabled() { + t.Error("Enabled=false must disable the limits") + } + if got := cfg.WindowSecondsOrDefault(); got != 60 { + t.Errorf("a non-positive window must fall back to 60, got %d", got) + } + + on := true + cfg = RateLimitConfig{Enabled: &on, WindowSeconds: 30} + if !cfg.IsEnabled() { + t.Error("Enabled=true must enable the limits") + } + if got := cfg.WindowSecondsOrDefault(); got != 30 { + t.Errorf("WindowSecondsOrDefault() = %d want 30", got) + } +} + +func TestLoadReadsRateLimitSettings(t *testing.T) { + path := filepath.Join(t.TempDir(), "config.yaml") + if err := os.WriteFile(path, []byte("server:\n port: 8083\n"), 0600); err != nil { + t.Fatalf("WriteFile() error = %v", err) + } + + t.Setenv("ENV_FILE", os.DevNull) + t.Setenv("AGENT_DESK_ENV_FILE", os.DevNull) + t.Setenv("RATE_LIMIT_ENABLED", "false") + t.Setenv("RATE_LIMIT_WINDOW_SECONDS", "30") + + cfg, err := Load(path) + if err != nil { + t.Fatalf("Load() error = %v", err) + } + if cfg.Server.RateLimit.IsEnabled() { + t.Error("RATE_LIMIT_ENABLED=false did not disable the limits") + } + if got := cfg.Server.RateLimit.WindowSecondsOrDefault(); got != 30 { + t.Errorf("WindowSecondsOrDefault() = %d want 30", got) + } +} diff --git a/internal/pkg/i18nx/locales/en-US.yml b/internal/pkg/i18nx/locales/en-US.yml index e74b751b..c07c8d97 100644 --- a/internal/pkg/i18nx/locales/en-US.yml +++ b/internal/pkg/i18nx/locales/en-US.yml @@ -346,6 +346,7 @@ 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.e0354: "Too many requests. Please wait a moment and try again." 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 371264f0..457540e5 100644 --- a/internal/pkg/i18nx/locales/zh-CN.yml +++ b/internal/pkg/i18nx/locales/zh-CN.yml @@ -346,6 +346,7 @@ error.e0345: "附件消息缺少 assetId" error.e0346: "附件消息缺少 payload" error.e0347: "默认客服组待接入池模式必须至少选择一个客服组" error.e0348: "出于安全考虑,此文件类型不支持上传" +error.e0354: "请求过于频繁,请稍后再试。" error.profile.nicknameRequired: "请输入昵称" error.profile.nicknameTooLong: "昵称不能超过 100 个字符" error.profile.avatarTooLong: "头像链接不能超过 255 个字符" diff --git a/internal/pkg/ratelimit/ratelimit.go b/internal/pkg/ratelimit/ratelimit.go new file mode 100644 index 00000000..47f4e6ea --- /dev/null +++ b/internal/pkg/ratelimit/ratelimit.go @@ -0,0 +1,101 @@ +// Package ratelimit provides a small in-process fixed-window counter for the +// public endpoints that an unauthenticated caller can drive. +// +// It is deliberately not distributed. A deployment running several replicas gets +// the configured limit per replica rather than per cluster, which weakens the +// bound by a factor of the replica count. That is the right trade here because +// these limits exist to make flooding expensive, not to meter a quota, and +// because reaching for a shared store would put a network round trip on the +// hottest unauthenticated paths in the application. If this ever needs to be +// exact across replicas, the Limiter interface is small enough to swap. +package ratelimit + +import ( + "sync" + "time" +) + +type bucket struct { + count int + resetAt time.Time +} + +// Limiter counts requests per key inside a fixed window. The zero value is not +// usable; call New. A nil *Limiter allows everything, so a caller can represent +// "disabled" by holding nil instead of branching at every use site. +type Limiter struct { + limit int + window time.Duration + + mu sync.Mutex + buckets map[string]bucket + sweptAt time.Time +} + +// New returns a Limiter that allows limit requests per window for each distinct +// key. A limit of zero or less, or a non-positive window, disables the limiter. +func New(limit int, window time.Duration) *Limiter { + if limit <= 0 || window <= 0 { + return &Limiter{limit: limit, window: window, buckets: make(map[string]bucket)} + } + return &Limiter{limit: limit, window: window, buckets: make(map[string]bucket)} +} + +// Allow records one request for key and reports whether it may proceed. When it +// returns false, retryAfter is how long the caller should wait before the window +// resets, and is meant for a Retry-After header. +func (l *Limiter) Allow(key string) (bool, time.Duration) { + if l == nil || l.limit <= 0 || l.window <= 0 { + return true, 0 + } + + now := time.Now() + + l.mu.Lock() + defer l.mu.Unlock() + + l.sweep(now) + + current, ok := l.buckets[key] + if !ok || !now.Before(current.resetAt) { + l.buckets[key] = bucket{count: 1, resetAt: now.Add(l.window)} + return true, 0 + } + + current.count++ + l.buckets[key] = current + if current.count > l.limit { + retryAfter := current.resetAt.Sub(now) + if retryAfter < 0 { + retryAfter = 0 + } + return false, retryAfter + } + return true, 0 +} + +// Len reports how many keys are currently tracked. It exists for tests. +func (l *Limiter) Len() int { + if l == nil { + return 0 + } + l.mu.Lock() + defer l.mu.Unlock() + return len(l.buckets) +} + +// sweep drops buckets whose window has elapsed, so a long-running process does +// not accumulate one entry per address it has ever seen. It runs at most once per +// window and under the same lock as Allow, which keeps the cost amortised and +// avoids a background goroutine whose lifetime somebody would have to own. +func (l *Limiter) sweep(now time.Time) { + if !l.sweptAt.IsZero() && now.Sub(l.sweptAt) < l.window { + return + } + for key, current := range l.buckets { + if !now.Before(current.resetAt) { + delete(l.buckets, key) + } + } + l.sweptAt = now +} diff --git a/internal/pkg/ratelimit/ratelimit_test.go b/internal/pkg/ratelimit/ratelimit_test.go new file mode 100644 index 00000000..dfa78ea6 --- /dev/null +++ b/internal/pkg/ratelimit/ratelimit_test.go @@ -0,0 +1,132 @@ +package ratelimit + +import ( + "fmt" + "sync" + "sync/atomic" + "testing" + "time" +) + +func TestLimiterAllowsUpToTheLimitThenRefuses(t *testing.T) { + l := New(3, time.Minute) + for i := 1; i <= 3; i++ { + if ok, retryAfter := l.Allow("203.0.113.7"); !ok { + t.Fatalf("request %d was refused although the limit is 3 (retryAfter=%v)", i, retryAfter) + } + } + + ok, retryAfter := l.Allow("203.0.113.7") + if ok { + t.Fatal("the fourth request was allowed although the limit is 3") + } + if retryAfter <= 0 || retryAfter > time.Minute { + t.Fatalf("retryAfter = %v, want a positive duration within the window", retryAfter) + } +} + +func TestLimiterKeysAreIndependent(t *testing.T) { + l := New(1, time.Minute) + if ok, _ := l.Allow("203.0.113.7"); !ok { + t.Fatal("the first request for a key was refused") + } + if ok, _ := l.Allow("203.0.113.7"); ok { + t.Fatal("the second request for the same key was allowed") + } + // A different caller must not inherit somebody else's exhaustion. Without key + // isolation one blocked address would lock out every user behind it. + if ok, _ := l.Allow("198.51.100.9"); !ok { + t.Fatal("a different key was refused because another key was exhausted") + } +} + +func TestLimiterWindowResets(t *testing.T) { + l := New(1, 20*time.Millisecond) + if ok, _ := l.Allow("k"); !ok { + t.Fatal("the first request was refused") + } + if ok, _ := l.Allow("k"); ok { + t.Fatal("the second request inside the window was allowed") + } + time.Sleep(40 * time.Millisecond) + if ok, _ := l.Allow("k"); !ok { + t.Fatal("the window did not reset") + } +} + +func TestNilLimiterAllowsEverything(t *testing.T) { + var l *Limiter + for i := 0; i < 100; i++ { + if ok, retryAfter := l.Allow("k"); !ok { + t.Fatalf("a nil limiter refused request %d (retryAfter=%v)", i, retryAfter) + } + } + if l.Len() != 0 { + t.Fatalf("Len() on a nil limiter = %d want 0", l.Len()) + } +} + +func TestDisabledLimiterAllowsEverything(t *testing.T) { + cases := []struct { + name string + limiter *Limiter + }{ + {"zero limit", New(0, time.Minute)}, + {"negative limit", New(-1, time.Minute)}, + {"zero window", New(5, 0)}, + {"negative window", New(5, -time.Second)}, + } + for _, tc := range cases { + for i := 0; i < 50; i++ { + if ok, _ := tc.limiter.Allow("k"); !ok { + t.Fatalf("%s: request %d was refused although the limiter is disabled", tc.name, i) + } + } + } +} + +func TestSweepReleasesExpiredKeys(t *testing.T) { + l := New(10, 20*time.Millisecond) + for i := 0; i < 200; i++ { + l.Allow(fmt.Sprintf("key-%d", i)) + } + if got := l.Len(); got != 200 { + t.Fatalf("Len() = %d before the sweep, want 200", got) + } + + time.Sleep(40 * time.Millisecond) + // The sweep is driven by Allow rather than by a goroutine, so one more call + // has to reclaim everything that expired. + l.Allow("fresh") + if got := l.Len(); got != 1 { + t.Fatalf("Len() = %d after the sweep, want only the fresh key", got) + } +} + +// TestLimiterCountsExactlyUnderConcurrency is the reason Allow holds a single +// mutex across the read-modify-write rather than checking and then recording. +func TestLimiterCountsExactlyUnderConcurrency(t *testing.T) { + const goroutines = 16 + const perGoroutine = 200 + const limit = 1000 + + l := New(limit, time.Minute) + var allowed atomic.Int64 + var wg sync.WaitGroup + for i := 0; i < goroutines; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for j := 0; j < perGoroutine; j++ { + if ok, _ := l.Allow("shared"); ok { + allowed.Add(1) + } + } + }() + } + wg.Wait() + + if got := allowed.Load(); got != limit { + t.Fatalf("allowed %d of %d requests, want exactly the limit of %d", got, goroutines*perGoroutine, limit) + } +} From 7ebed455964f78473c668f8f37345e6505bb71d6 Mon Sep 17 00:00:00 2001 From: JOY <5027251+JOY@users.noreply.github.com> Date: Tue, 22 Sep 2026 10:42:16 +0700 Subject: [PATCH 7/8] feat(auth): auto-redirect to the SSO provider when it is the only transport When OIDC is enabled, password login is disabled, and no other provider is configured, /dashboard/login now sends the visitor straight to the provider instead of rendering a page whose only working control is the "sign in with SSO" button. The redirect waits for the session probe, so an already-signed-in visitor still lands on their destination, and it is suppressed when the provider bounced back with ?oidcError= so a failing IdP cannot create a redirect loop (the page then renders with its error toast and the manual button), and in WxWork-only environments. --- web/components/login-form.tsx | 37 ++++++++++++++++++++++++++++++++++- web/messages/en-US.json | 1 + web/messages/zh-CN.json | 1 + 3 files changed, 38 insertions(+), 1 deletion(-) diff --git a/web/components/login-form.tsx b/web/components/login-form.tsx index 34bb9ad4..499bb365 100644 --- a/web/components/login-form.tsx +++ b/web/components/login-form.tsx @@ -38,7 +38,7 @@ export function LoginForm({ const t = useI18n() const router = useRouter() const searchParams = useSearchParams() - const { session } = useAuth() + const { session, ready } = useAuth() const [isPending, setIsPending] = useState(false) const [isWxWorkEnv, setIsWxWorkEnv] = useState(false) const [publicConfig, setPublicConfig] = useState(null) @@ -52,6 +52,28 @@ export function LoginForm({ Number(publicConfig?.wxworkEnabled) + Number(publicConfig?.oidcEnabled) const isPasswordLoginEnabled = publicConfig?.passwordLoginEnabled !== false + // Redirect-only mode: when OIDC is the only login transport, skip the + // chooser and go straight to the provider. Suppressed while the session + // probe is in flight (an already-signed-in visitor goes to their + // destination instead of the IdP), in WxWork-only environments, and when + // the provider bounced back with ?oidcError= - otherwise this would loop. + const shouldRedirectToOIDC = Boolean( + ready && + !session && + publicConfig && + publicConfig.oidcEnabled && + !publicConfig.passwordLoginEnabled && + !publicConfig.wxworkEnabled && + !oidcError, + ) + + useEffect(() => { + if (!shouldRedirectToOIDC) { + return + } + window.location.href = `/api/auth/oidc_login?next=${encodeURIComponent(redirectPath)}` + }, [shouldRedirectToOIDC, redirectPath]) + useEffect(() => { if (session) { router.replace(redirectPath) @@ -151,6 +173,19 @@ export function LoginForm({ ) } + if (shouldRedirectToOIDC) { + return ( +
+ + + +

{t("auth.redirectingToOidc")}

+
+
+
+ ) + } + return (
Date: Tue, 22 Sep 2026 23:40:07 +0700 Subject: [PATCH 8/8] fix(sync): dedupe the duplicated Slack handler test and restore upstream coverage The merge left TestSlackWebhook_Handler declared twice (the fork's whatsapp_slack_handler_test.go and upstream's slack_handler_test.go). Drop the fork's copy - upstream's version is the newer signed-payload form and already covers the unsigned-rejection case - and keep the WhatsApp test in the fork file. Also restore the upstream-only tests the merge had dropped (slack inbound rejection/replay/challenge, discord guild-scope and wrong-secret, OIDC first-login lowest role) and the slack/discord setup i18n keys the channel dialog renders. --- internal/handlers/third/slack_handler_test.go | 4 + .../third/whatsapp_slack_handler_test.go | 139 ------------------ 2 files changed, 4 insertions(+), 139 deletions(-) diff --git a/internal/handlers/third/slack_handler_test.go b/internal/handlers/third/slack_handler_test.go index 9e7b82e2..72bf72a9 100644 --- a/internal/handlers/third/slack_handler_test.go +++ b/internal/handlers/third/slack_handler_test.go @@ -150,4 +150,8 @@ func TestSlackWebhook_Handler(t *testing.T) { if identity == nil { t.Fatalf("expected customer identity for U_USER_777") } + + if identity == nil { + t.Fatalf("expected customer identity for U_USER_777") + } } diff --git a/internal/handlers/third/whatsapp_slack_handler_test.go b/internal/handlers/third/whatsapp_slack_handler_test.go index 73dbb8ff..d28f20a9 100644 --- a/internal/handlers/third/whatsapp_slack_handler_test.go +++ b/internal/handlers/third/whatsapp_slack_handler_test.go @@ -8,7 +8,6 @@ import ( "encoding/json" "net/http" "net/http/httptest" - "strconv" "testing" "time" @@ -33,19 +32,6 @@ func signWhatsAppTestPayload(payload []byte) string { return "sha256=" + hex.EncodeToString(mac.Sum(nil)) } -// slackTestSigningSecret is the Slack signing secret the test channel is -// configured with. -const slackTestSigningSecret = "test_signing_secret" - -// signSlackTestPayload builds the X-Slack-Request-Timestamp and -// X-Slack-Signature headers Slack would send for this body right now. -func signSlackTestPayload(payload []byte) (string, string) { - timestamp := strconv.FormatInt(time.Now().Unix(), 10) - mac := hmac.New(sha256.New, []byte(slackTestSigningSecret)) - mac.Write([]byte("v0:" + timestamp + ":" + string(payload))) - return timestamp, "v0=" + hex.EncodeToString(mac.Sum(nil)) -} - func TestWhatsAppWebhook_Handler(t *testing.T) { gin.SetMode(gin.TestMode) db := setupThirdHandlerTestDB(t) @@ -171,128 +157,3 @@ func TestWhatsAppWebhook_Handler(t *testing.T) { t.Fatalf("expected customer identity for 1234567890") } } - -func TestSlackWebhook_Handler(t *testing.T) { - gin.SetMode(gin.TestMode) - db := setupThirdHandlerTestDB(t) - - now := time.Now() - agent := &models.AIAgent{ - Name: "Slack Agent", - ServiceMode: enums.IMConversationServiceModeAIFirst, - PublishedRevisionID: 1, - WelcomeMessage: "Hello Slack User!", - Status: enums.StatusOk, - AuditFields: models.AuditFields{CreatedAt: now, UpdatedAt: now}, - } - _ = db.Create(agent) - - slackConfig, _ := json.Marshal(dto.SlackChannelConfig{ - BotToken: "xoxb-test-token", - SigningSecret: slackTestSigningSecret, - TeamID: "T_SLACK_100", - DefaultChannel: "C_GENERAL", - }) - - operator := &dto.AuthPrincipal{UserID: 1, Username: "admin"} - channel, err := services.ChannelService.CreateChannel(request.CreateChannelRequest{ - Name: "Slack Channel", - ChannelType: enums.ChannelTypeSlack, - AIAgentID: agent.ID, - AIAgentRolloutPercent: 100, - ConfigJSON: string(slackConfig), - Status: int(enums.StatusOk), - }, operator) - if err != nil { - t.Fatalf("CreateChannel failed: %v", err) - } - - router := gin.New() - router.POST("/api/third/slack/webhook/:channel_id", SlackPostWebhook) - router.POST("/api/third/slack/webhook", SlackPostWebhook) - - // 1. URL Verification - challengePayload := []byte(`{ - "token": "token123", - "challenge": "slack_challenge_string_999", - "type": "url_verification" - }`) - reqChallenge, _ := http.NewRequest(http.MethodPost, "/api/third/slack/webhook/"+channel.ChannelID, bytes.NewBuffer(challengePayload)) - reqChallenge.Header.Set("Content-Type", "application/json") - recChallenge := httptest.NewRecorder() - router.ServeHTTP(recChallenge, reqChallenge) - - if recChallenge.Code != http.StatusOK { - t.Fatalf("expected 200 OK for challenge, got: %d", recChallenge.Code) - } - var challengeResp map[string]any - _ = json.Unmarshal(recChallenge.Body.Bytes(), &challengeResp) - if challengeResp["challenge"] != "slack_challenge_string_999" { - t.Fatalf("expected challenge in body, got: %+v", challengeResp) - } - - // 2. Event Callback - eventPayload := []byte(`{ - "token": "token123", - "team_id": "T_SLACK_100", - "type": "event_callback", - "event": { - "type": "message", - "user": "U_USER_777", - "text": "Hello support team on Slack!", - "ts": "1725260000.000100", - "channel": "C_GENERAL" - } - }`) - reqEvent, _ := http.NewRequest(http.MethodPost, "/api/third/slack/webhook/"+channel.ChannelID, bytes.NewBuffer(eventPayload)) - reqEvent.Header.Set("Content-Type", "application/json") - slackTimestamp, slackSignature := signSlackTestPayload(eventPayload) - reqEvent.Header.Set("X-Slack-Request-Timestamp", slackTimestamp) - reqEvent.Header.Set("X-Slack-Signature", slackSignature) - recEvent := httptest.NewRecorder() - router.ServeHTTP(recEvent, reqEvent) - - if recEvent.Code != http.StatusOK { - t.Fatalf("expected 200 OK for event, got: %d", recEvent.Code) - } - - // Verify identity - identity := repositories.CustomerIdentityRepository.FindOne(db, sqls.NewCnd(). - Eq("external_source", enums.ExternalSourceSlack). - Eq("external_id", "U_USER_777")) - if identity == nil { - t.Fatalf("expected customer identity for U_USER_777") - } - - // 3. Unsigned delivery is rejected: once a signing secret resolves for the - // channel, a payload without Slack signature headers must not provision - // anything. The handler still answers 200 ok=false so Slack does not retry. - unsignedPayload := []byte(`{ - "token": "token123", - "team_id": "T_SLACK_100", - "type": "event_callback", - "event": { - "type": "message", - "user": "U_USER_888", - "text": "Unsigned spoof attempt", - "ts": "1725260001.000100", - "channel": "C_GENERAL" - } - }`) - reqUnsigned, _ := http.NewRequest(http.MethodPost, "/api/third/slack/webhook/"+channel.ChannelID, bytes.NewBuffer(unsignedPayload)) - reqUnsigned.Header.Set("Content-Type", "application/json") - recUnsigned := httptest.NewRecorder() - router.ServeHTTP(recUnsigned, reqUnsigned) - - var unsignedResp map[string]any - _ = json.Unmarshal(recUnsigned.Body.Bytes(), &unsignedResp) - if unsignedResp["ok"] != false { - t.Fatalf("expected unsigned delivery to be rejected with ok=false, got: %+v", unsignedResp) - } - unsignedIdentity := repositories.CustomerIdentityRepository.FindOne(db, sqls.NewCnd(). - Eq("external_source", enums.ExternalSourceSlack). - Eq("external_id", "U_USER_888")) - if unsignedIdentity != nil { - t.Fatalf("unsigned delivery must not provision an identity for U_USER_888") - } -}