From 2631689958369b4415ea669b1ab084482a689adc Mon Sep 17 00:00:00 2001 From: mlogclub Date: Fri, 17 Apr 2026 12:50:06 +0800 Subject: [PATCH] feat(message): enhance message normalization and asset handling in HTML content --- internal/pkg/utils/message.go | 140 +++++++++++++++++- internal/pkg/utils/message_test.go | 88 +++++++++++ internal/services/message_service.go | 5 +- .../services/wxwork_kf_outbound_service.go | 39 +---- 4 files changed, 231 insertions(+), 41 deletions(-) diff --git a/internal/pkg/utils/message.go b/internal/pkg/utils/message.go index fecddcf..c38d98c 100644 --- a/internal/pkg/utils/message.go +++ b/internal/pkg/utils/message.go @@ -3,10 +3,13 @@ package utils import ( "bytes" "cs-agent/internal/models" + "cs-agent/internal/pkg/config" "cs-agent/internal/pkg/enums" "cs-agent/internal/repositories" "cs-agent/internal/services/storage" "encoding/json" + "fmt" + "net/url" "strings" "github.com/microcosm-cc/bluemonday" @@ -27,13 +30,43 @@ type imMessageAssetPayload struct { func SanitizeMessageHTML(content string) string { policy := bluemonday.UGCPolicy() policy.AllowElements("img") - policy.AllowAttrs("src", "alt", "title", "data-provider", "data-storage-key").OnElements("img") + policy.AllowAttrs("src", "alt", "title", "data-asset-id", "data-provider", "data-storage-key").OnElements("img") policy.AllowURLSchemes("http", "https") policy.AllowStandardURLs() policy.AllowElements("p", "br") return stripHTMLImageSrcIfBound(strings.TrimSpace(policy.Sanitize(content))) } +func NormalizeMessageHTMLAssets(content string) string { + content = strings.TrimSpace(content) + if content == "" { + return "" + } + doc, err := html.Parse(strings.NewReader("
" + content + "
")) + if err != nil { + return content + } + var walk func(*html.Node) + walk = func(node *html.Node) { + if node == nil { + return + } + if node.Type == html.ElementNode && node.Data == "img" { + if asset := findImageAsset(node); asset != nil { + setHTMLAttr(node, "data-asset-id", strings.TrimSpace(asset.AssetID)) + setHTMLAttr(node, "data-provider", strings.TrimSpace(string(asset.Provider))) + setHTMLAttr(node, "data-storage-key", strings.TrimSpace(asset.StorageKey)) + removeHTMLAttr(node, "src") + } + } + for child := node.FirstChild; child != nil; child = child.NextSibling { + walk(child) + } + } + walk(doc) + return renderHTMLFragment(doc) +} + func BuildHTMLSummary(content string) string { if strings.TrimSpace(content) == "" { return "" @@ -285,3 +318,108 @@ func hydrateIMMessageAssetPayload(payload *imMessageAssetPayload) *imMessageAsse } return payload } + +func findImageAsset(node *html.Node) *models.Asset { + if node == nil { + return nil + } + if assetID := strings.TrimSpace(findHTMLAttr(node, "data-asset-id")); assetID != "" { + if asset := repositories.AssetRepository.GetByAssetID(sqls.DB(), assetID); asset != nil { + return asset + } + } + provider := enums.AssetProvider(strings.TrimSpace(findHTMLAttr(node, "data-provider"))) + storageKey := strings.TrimSpace(findHTMLAttr(node, "data-storage-key")) + if provider != "" && storageKey != "" { + if asset := repositories.AssetRepository.GetByStorageKey(sqls.DB(), storageKey); asset != nil { + return asset + } + } + src := strings.TrimSpace(findHTMLAttr(node, "src")) + if src == "" { + return nil + } + return findAssetByMessageImageURL(src) +} + +func findAssetByMessageImageURL(rawURL string) *models.Asset { + storageKey, err := resolveStorageKeyFromMessageImageURL(rawURL) + if err != nil { + return nil + } + return repositories.AssetRepository.GetByStorageKey(sqls.DB(), storageKey) +} + +func FindAssetByMessageImageURL(rawURL string) *models.Asset { + return findAssetByMessageImageURL(rawURL) +} + +func resolveStorageKeyFromMessageImageURL(rawURL string) (string, error) { + cfg := config.Current().Storage + candidates := make([]string, 0, 3) + if baseURL := strings.TrimSpace(cfg.Local.BaseURL); baseURL != "" { + candidates = append(candidates, baseURL) + } + if baseURL := strings.TrimSpace(cfg.OSS.BaseURL); baseURL != "" { + candidates = append(candidates, baseURL) + } + if ossBucketBaseURL := buildOSSBucketBaseURL(cfg.OSS); ossBucketBaseURL != "" { + candidates = append(candidates, ossBucketBaseURL) + } + for _, baseURL := range candidates { + if storageKey, err := resolveStorageKeyFromAssetURL(baseURL, rawURL); err == nil && storageKey != "" { + return storageKey, nil + } + } + return "", fmt.Errorf("image url does not match any storage base url") +} + +func buildOSSBucketBaseURL(cfg config.OSSStorageConfig) string { + endpoint := strings.TrimSpace(cfg.Endpoint) + bucket := strings.TrimSpace(cfg.Bucket) + if endpoint == "" || bucket == "" { + return "" + } + if !strings.Contains(endpoint, "://") { + endpoint = "https://" + endpoint + } + u, err := url.Parse(endpoint) + if err != nil || strings.TrimSpace(u.Host) == "" { + return "" + } + scheme := strings.TrimSpace(u.Scheme) + if scheme == "" { + scheme = "https" + } + return fmt.Sprintf("%s://%s.%s", scheme, bucket, u.Host) +} + +func resolveStorageKeyFromAssetURL(baseURL, rawURL string) (string, error) { + baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/") + rawURL = strings.TrimSpace(rawURL) + if baseURL == "" || rawURL == "" { + return "", fmt.Errorf("invalid image url") + } + if strings.HasPrefix(rawURL, baseURL+"/") { + return strings.TrimLeft(strings.TrimPrefix(rawURL, baseURL), "/"), nil + } + + baseParsed, baseErr := url.Parse(baseURL) + rawParsed, rawErr := url.Parse(rawURL) + if baseErr != nil || rawErr != nil { + return "", fmt.Errorf("invalid image url") + } + if !strings.EqualFold(baseParsed.Host, rawParsed.Host) { + return "", fmt.Errorf("image url host mismatch") + } + basePath := strings.TrimRight(baseParsed.Path, "/") + rawPath := strings.TrimLeft(rawParsed.Path, "/") + if basePath == "" { + return rawPath, nil + } + basePath = strings.TrimLeft(basePath, "/") + if !strings.HasPrefix(rawPath, basePath+"/") { + return "", fmt.Errorf("image url path mismatch") + } + return strings.TrimLeft(strings.TrimPrefix(rawPath, basePath), "/"), nil +} diff --git a/internal/pkg/utils/message_test.go b/internal/pkg/utils/message_test.go index eab1572..125530d 100644 --- a/internal/pkg/utils/message_test.go +++ b/internal/pkg/utils/message_test.go @@ -6,6 +6,11 @@ import ( "cs-agent/internal/pkg/enums" "strings" "testing" + "time" + + "github.com/glebarez/sqlite" + "github.com/mlogclub/simple/sqls" + "gorm.io/gorm" ) func TestBuildIMMessageAssetPayloadForResponseAddsSignedURL(t *testing.T) { @@ -94,3 +99,86 @@ func TestBuildRenderableMessageTransformsPayloadAndHTML(t *testing.T) { t.Fatalf("expected html content signed src, got: %s", htmlContent) } } + +func TestNormalizeMessageHTMLAssetsEnrichesImageDataAttrs(t *testing.T) { + setupMessageTestDB(t) + config.SetCurrent(&config.Config{ + Storage: config.StorageConfig{ + Default: enums.AssetProviderLocal, + Local: config.LocalStorageConfig{ + BaseURL: "https://files.example.com", + }, + }, + }) + createTestAsset(t, &models.Asset{ + AssetID: "asset_local_1", + Provider: enums.AssetProviderLocal, + StorageKey: "images/demo.png", + Filename: "demo.png", + FileSize: 123, + MimeType: "image/png", + Status: enums.AssetStatusSuccess, + }) + + got := NormalizeMessageHTMLAssets(`

demo

`) + + if !strings.Contains(got, `data-asset-id="asset_local_1"`) { + t.Fatalf("expected data-asset-id added, got: %s", got) + } + if !strings.Contains(got, `data-provider="local"`) { + t.Fatalf("expected data-provider added, got: %s", got) + } + if !strings.Contains(got, `data-storage-key="images/demo.png"`) { + t.Fatalf("expected data-storage-key added, got: %s", got) + } + if strings.Contains(got, `src=`) { + t.Fatalf("expected src removed after asset binding, got: %s", got) + } +} + +func TestNormalizeMessageHTMLAssetsKeepsUnknownImageSrc(t *testing.T) { + setupMessageTestDB(t) + config.SetCurrent(&config.Config{ + Storage: config.StorageConfig{ + Default: enums.AssetProviderLocal, + Local: config.LocalStorageConfig{ + BaseURL: "https://files.example.com", + }, + }, + }) + + got := NormalizeMessageHTMLAssets(`

demo

`) + + if !strings.Contains(got, `src="https://unknown.example.com/demo.png"`) { + t.Fatalf("expected unknown image src kept, got: %s", got) + } + if strings.Contains(got, `data-asset-id=`) || strings.Contains(got, `data-provider=`) || strings.Contains(got, `data-storage-key=`) { + t.Fatalf("expected no asset attrs added for unknown image, got: %s", got) + } +} + +func setupMessageTestDB(t *testing.T) { + t.Helper() + db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) + if err != nil { + t.Fatalf("open sqlite failed: %v", err) + } + if err := db.AutoMigrate(&models.Asset{}); err != nil { + t.Fatalf("auto migrate asset failed: %v", err) + } + sqls.SetDB(db) +} + +func createTestAsset(t *testing.T, item *models.Asset) { + t.Helper() + now := time.Now() + if item.CreatedAt.IsZero() { + item.CreatedAt = now + } + if item.UpdatedAt.IsZero() { + item.UpdatedAt = now + } + if err := sqls.DB().Create(item).Error; err != nil { + t.Fatalf("create asset failed: %v", err) + } +} diff --git a/internal/services/message_service.go b/internal/services/message_service.go index fad0856..e2a9ddf 100644 --- a/internal/services/message_service.go +++ b/internal/services/message_service.go @@ -457,11 +457,12 @@ func (s *messageService) normalizeMessageContent(conversationID int64, messageTy switch messageType { case enums.IMMessageTypeHTML: sanitized := utils.SanitizeMessageHTML(content) - summary := utils.BuildHTMLSummary(sanitized) + normalized := utils.NormalizeMessageHTMLAssets(sanitized) + summary := utils.BuildHTMLSummary(normalized) if summary == "" { return "", "", "", errorsx.InvalidParam("消息内容不能为空") } - return sanitized, "", summary, nil + return normalized, "", summary, nil case enums.IMMessageTypeImage, enums.IMMessageTypeAttachment: assetPayload, err := parseIMMessageAssetPayload(payload) if err != nil { diff --git a/internal/services/wxwork_kf_outbound_service.go b/internal/services/wxwork_kf_outbound_service.go index bccb106..c2b816f 100644 --- a/internal/services/wxwork_kf_outbound_service.go +++ b/internal/services/wxwork_kf_outbound_service.go @@ -4,12 +4,10 @@ import ( "encoding/json" "fmt" "log/slog" - "net/url" "strings" "time" "cs-agent/internal/models" - "cs-agent/internal/pkg/config" "cs-agent/internal/pkg/enums" "cs-agent/internal/pkg/utils" "cs-agent/internal/repositories" @@ -477,44 +475,9 @@ func (s *wxWorkKFOutboundService) buildHTMLChunks(content string) ([]wxWorkKFOut } func (s *wxWorkKFOutboundService) resolveAssetIDFromImageSrc(src string) (string, error) { - cfg := config.Current() - storageKey, err := resolveStorageKeyFromAssetURL(strings.TrimSpace(cfg.Storage.Local.BaseURL), src) - if err != nil { - return "", err - } - asset := AssetService.GetByStorageKey(storageKey) + asset := utils.FindAssetByMessageImageURL(src) if asset == nil { return "", fmt.Errorf("未找到图片资源") } return strings.TrimSpace(asset.AssetID), nil } - -func resolveStorageKeyFromAssetURL(baseURL, rawURL string) (string, error) { - baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/") - rawURL = strings.TrimSpace(rawURL) - if baseURL == "" || rawURL == "" { - return "", fmt.Errorf("图片URL不合法") - } - if strings.HasPrefix(rawURL, baseURL+"/") { - return strings.TrimLeft(strings.TrimPrefix(rawURL, baseURL), "/"), nil - } - - baseParsed, baseErr := url.Parse(baseURL) - rawParsed, rawErr := url.Parse(rawURL) - if baseErr != nil || rawErr != nil { - return "", fmt.Errorf("图片URL不合法") - } - if !strings.EqualFold(baseParsed.Host, rawParsed.Host) { - return "", fmt.Errorf("图片URL不属于当前存储域名") - } - basePath := strings.TrimRight(baseParsed.Path, "/") - rawPath := strings.TrimLeft(rawParsed.Path, "/") - if basePath == "" { - return rawPath, nil - } - basePath = strings.TrimLeft(basePath, "/") - if !strings.HasPrefix(rawPath, basePath+"/") { - return "", fmt.Errorf("图片URL不属于当前存储目录") - } - return strings.TrimLeft(strings.TrimPrefix(rawPath, basePath), "/"), nil -}