diff --git a/internal/pkg/utils/message.go b/internal/pkg/utils/message.go index c38d98c..c6944e0 100644 --- a/internal/pkg/utils/message.go +++ b/internal/pkg/utils/message.go @@ -37,22 +37,28 @@ func SanitizeMessageHTML(content string) string { return stripHTMLImageSrcIfBound(strings.TrimSpace(policy.Sanitize(content))) } -func NormalizeMessageHTMLAssets(content string) string { +func NormalizeMessageHTMLAssets(content string) (string, error) { content = strings.TrimSpace(content) if content == "" { - return "" + return "", nil } doc, err := html.Parse(strings.NewReader("
" + content + "
")) if err != nil { - return content + return content, nil } + var walkErr error var walk func(*html.Node) walk = func(node *html.Node) { - if node == nil { + if node == nil || walkErr != nil { return } if node.Type == html.ElementNode && node.Data == "img" { - if asset := findImageAsset(node); asset != nil { + asset, err := normalizeHTMLImageAsset(node) + if err != nil { + walkErr = err + return + } + if 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)) @@ -64,7 +70,10 @@ func NormalizeMessageHTMLAssets(content string) string { } } walk(doc) - return renderHTMLFragment(doc) + if walkErr != nil { + return "", walkErr + } + return renderHTMLFragment(doc), nil } func BuildHTMLSummary(content string) string { @@ -319,27 +328,39 @@ func hydrateIMMessageAssetPayload(payload *imMessageAssetPayload) *imMessageAsse return payload } -func findImageAsset(node *html.Node) *models.Asset { +func normalizeHTMLImageAsset(node *html.Node) (*models.Asset, error) { 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 - } + return nil, nil } + assetID := strings.TrimSpace(findHTMLAttr(node, "data-asset-id")) 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 + + hasAssetID := assetID != "" + hasProvider := provider != "" + hasStorageKey := storageKey != "" + if hasAssetID || hasProvider || hasStorageKey { + if !(hasAssetID && hasProvider && hasStorageKey) { + return nil, fmt.Errorf("html message image asset attributes are incomplete") + } + asset := repositories.AssetRepository.GetByAssetID(sqls.DB(), assetID) + if asset == nil { + return nil, fmt.Errorf("html message image asset not found") + } + if asset.Provider != provider || strings.TrimSpace(asset.StorageKey) != storageKey { + return nil, fmt.Errorf("html message image asset attributes mismatch") + } + return asset, nil } - return findAssetByMessageImageURL(src) + if src == "" { + return nil, fmt.Errorf("html message image is missing asset metadata") + } + asset := findAssetByMessageImageURL(src) + if asset == nil { + return nil, fmt.Errorf("html message image must reference an uploaded asset") + } + return asset, nil } func findAssetByMessageImageURL(rawURL string) *models.Asset { diff --git a/internal/pkg/utils/message_test.go b/internal/pkg/utils/message_test.go index 125530d..04e7110 100644 --- a/internal/pkg/utils/message_test.go +++ b/internal/pkg/utils/message_test.go @@ -120,7 +120,10 @@ func TestNormalizeMessageHTMLAssetsEnrichesImageDataAttrs(t *testing.T) { Status: enums.AssetStatusSuccess, }) - got := NormalizeMessageHTMLAssets(`

demo

`) + got, err := NormalizeMessageHTMLAssets(`

demo

`) + if err != nil { + t.Fatalf("expected normalization success, got error: %v", err) + } if !strings.Contains(got, `data-asset-id="asset_local_1"`) { t.Fatalf("expected data-asset-id added, got: %s", got) @@ -147,13 +150,52 @@ func TestNormalizeMessageHTMLAssetsKeepsUnknownImageSrc(t *testing.T) { }, }) - got := NormalizeMessageHTMLAssets(`

demo

`) - - if !strings.Contains(got, `src="https://unknown.example.com/demo.png"`) { - t.Fatalf("expected unknown image src kept, got: %s", got) + _, err := NormalizeMessageHTMLAssets(`

demo

`) + if err == nil { + t.Fatalf("expected unknown image src rejected") } - 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 TestNormalizeMessageHTMLAssetsRejectsIncompleteAttrs(t *testing.T) { + setupMessageTestDB(t) + config.SetCurrent(&config.Config{ + Storage: config.StorageConfig{ + Default: enums.AssetProviderLocal, + Local: config.LocalStorageConfig{ + BaseURL: "https://files.example.com", + }, + }, + }) + + _, err := NormalizeMessageHTMLAssets(`

demo

`) + if err == nil { + t.Fatalf("expected incomplete asset attrs rejected") + } +} + +func TestNormalizeMessageHTMLAssetsRejectsMismatchedAttrs(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_2", + Provider: enums.AssetProviderLocal, + StorageKey: "images/real.png", + Filename: "real.png", + FileSize: 456, + MimeType: "image/png", + Status: enums.AssetStatusSuccess, + }) + + _, err := NormalizeMessageHTMLAssets(`

demo

`) + if err == nil { + t.Fatalf("expected mismatched asset attrs rejected") } } diff --git a/internal/services/message_service.go b/internal/services/message_service.go index e2a9ddf..476ee08 100644 --- a/internal/services/message_service.go +++ b/internal/services/message_service.go @@ -457,7 +457,10 @@ func (s *messageService) normalizeMessageContent(conversationID int64, messageTy switch messageType { case enums.IMMessageTypeHTML: sanitized := utils.SanitizeMessageHTML(content) - normalized := utils.NormalizeMessageHTMLAssets(sanitized) + normalized, err := utils.NormalizeMessageHTMLAssets(sanitized) + if err != nil { + return "", "", "", errorsx.InvalidParam("HTML消息中的图片必须使用已上传文件") + } summary := utils.BuildHTMLSummary(normalized) if summary == "" { return "", "", "", errorsx.InvalidParam("消息内容不能为空")