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(`
`)
+
+ 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(`
`)
+
+ 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
-}