refactor(wecom): rename image handling functions to media handling and enhance media type support

This commit is contained in:
Zhang Rui 2026-03-18 21:25:40 +08:00
parent 5b76fb010f
commit a2a733038b
2 changed files with 172 additions and 39 deletions

View file

@ -705,7 +705,7 @@ func (c *WeComAIBotWSChannel) wsHandleMediaMessage(
ctx, cancel := context.WithTimeout(c.ctx, wsImageDownloadTimeout) ctx, cancel := context.WithTimeout(c.ctx, wsImageDownloadTimeout)
defer cancel() defer cancel()
ref, err := c.storeWSImage(ctx, chatID, msg.MsgID, resourceURL, aesKey) ref, err := c.storeWSMedia(ctx, chatID, msg.MsgID, resourceURL, aesKey, wsLabelToDefaultExt(label))
if err != nil { if err != nil {
logger.WarnCF("wecom_aibot", "Failed to download/store WS "+label, logger.WarnCF("wecom_aibot", "Failed to download/store WS "+label,
map[string]any{"error": err.Error(), "url": resourceURL}) map[string]any{"error": err.Error(), "url": resourceURL})
@ -742,8 +742,8 @@ func (c *WeComAIBotWSChannel) handleWSMixedMessage(reqID string, msg WeComAIBotW
} }
case "image": case "image":
if item.Image != nil { if item.Image != nil {
ref, err := c.storeWSImage(ctx, chatID, ref, err := c.storeWSMedia(ctx, chatID,
msg.MsgID+"-"+wsGenerateID(), item.Image.URL, item.Image.AESKey) msg.MsgID+"-"+wsGenerateID(), item.Image.URL, item.Image.AESKey, ".jpg")
if err != nil { if err != nil {
logger.WarnCF("wecom_aibot", "Failed to download/store mixed image", logger.WarnCF("wecom_aibot", "Failed to download/store mixed image",
map[string]any{"error": err.Error()}) map[string]any{"error": err.Error()})
@ -751,6 +751,9 @@ func (c *WeComAIBotWSChannel) handleWSMixedMessage(reqID string, msg WeComAIBotW
mediaRefs = append(mediaRefs, ref) mediaRefs = append(mediaRefs, ref)
} }
} }
default:
logger.WarnCF("wecom_aibot", "Unsupported item type in mixed message",
map[string]any{"msgtype": item.MsgType})
} }
} }
@ -1116,12 +1119,15 @@ func wsGenerateID() string {
return generateRandomID(10) return generateRandomID(10)
} }
// ---- Inbound image download helpers ---- // ---- Inbound media download helpers ----
// storeWSImage downloads the image at imageURL (with optional AES-CBC decryption) and stores it in the MediaStore. // storeWSMedia downloads the resource at resourceURL (with optional AES-CBC
func (c *WeComAIBotWSChannel) storeWSImage( // decryption) and stores it in the MediaStore. The file extension is inferred
// from the HTTP Content-Type response header; defaultExt is used as a fallback
// when the content type is absent or unrecognized.
func (c *WeComAIBotWSChannel) storeWSMedia(
ctx context.Context, ctx context.Context,
chatID, msgID, imageURL, aesKey string, chatID, msgID, resourceURL, aesKey, defaultExt string,
) (string, error) { ) (string, error) {
store := c.GetMediaStore() store := c.GetMediaStore()
if store == nil { if store == nil {
@ -1130,7 +1136,7 @@ func (c *WeComAIBotWSChannel) storeWSImage(
const maxSize = 20 << 20 // 20 MB const maxSize = 20 << 20 // 20 MB
req, err := http.NewRequestWithContext(ctx, http.MethodGet, imageURL, nil) req, err := http.NewRequestWithContext(ctx, http.MethodGet, resourceURL, nil)
if err != nil { if err != nil {
return "", fmt.Errorf("create request: %w", err) return "", fmt.Errorf("create request: %w", err)
} }
@ -1143,13 +1149,19 @@ func (c *WeComAIBotWSChannel) storeWSImage(
return "", fmt.Errorf("download HTTP %d", resp.StatusCode) return "", fmt.Errorf("download HTTP %d", resp.StatusCode)
} }
// Buffer the image in memory, bounded to maxSize. // Infer file extension from the Content-Type response header.
ext := wsMediaExtFromContentType(resp.Header.Get("Content-Type"))
if ext == "" {
ext = defaultExt
}
// Buffer the media in memory, bounded to maxSize.
data, err := io.ReadAll(io.LimitReader(resp.Body, int64(maxSize)+1)) data, err := io.ReadAll(io.LimitReader(resp.Body, int64(maxSize)+1))
if err != nil { if err != nil {
return "", fmt.Errorf("read image: %w", err) return "", fmt.Errorf("read media: %w", err)
} }
if len(data) > maxSize { if len(data) > maxSize {
return "", fmt.Errorf("image too large (> %d MB)", maxSize>>20) return "", fmt.Errorf("media too large (> %d MB)", maxSize>>20)
} }
// AES-CBC decryption if a key is present. // AES-CBC decryption if a key is present.
@ -1158,12 +1170,12 @@ func (c *WeComAIBotWSChannel) storeWSImage(
if decErr != nil || len(key) != 32 { if decErr != nil || len(key) != 32 {
key, decErr = decodeWeComAESKey(aesKey) key, decErr = decodeWeComAESKey(aesKey)
if decErr != nil { if decErr != nil {
return "", fmt.Errorf("decode image AES key: %w", decErr) return "", fmt.Errorf("decode media AES key: %w", decErr)
} }
} }
data, err = decryptAESCBC(key, data) data, err = decryptAESCBC(key, data)
if err != nil { if err != nil {
return "", fmt.Errorf("decrypt image: %w", err) return "", fmt.Errorf("decrypt media: %w", err)
} }
} }
@ -1173,7 +1185,7 @@ func (c *WeComAIBotWSChannel) storeWSImage(
if err = os.MkdirAll(mediaDir, 0o700); err != nil { if err = os.MkdirAll(mediaDir, 0o700); err != nil {
return "", fmt.Errorf("mkdir: %w", err) return "", fmt.Errorf("mkdir: %w", err)
} }
tmpFile, err := os.CreateTemp(mediaDir, msgID+"-*.jpg") tmpFile, err := os.CreateTemp(mediaDir, msgID+"-*"+ext)
if err != nil { if err != nil {
return "", fmt.Errorf("create temp file: %w", err) return "", fmt.Errorf("create temp file: %w", err)
} }
@ -1182,16 +1194,16 @@ func (c *WeComAIBotWSChannel) storeWSImage(
closeErr := tmpFile.Close() closeErr := tmpFile.Close()
if writeErr != nil { if writeErr != nil {
os.Remove(tmpPath) os.Remove(tmpPath)
return "", fmt.Errorf("write image: %w", writeErr) return "", fmt.Errorf("write media: %w", writeErr)
} }
if closeErr != nil { if closeErr != nil {
os.Remove(tmpPath) os.Remove(tmpPath)
return "", fmt.Errorf("close image: %w", closeErr) return "", fmt.Errorf("close media: %w", closeErr)
} }
scope := channels.BuildMediaScope("wecom_aibot", chatID, msgID) scope := channels.BuildMediaScope("wecom_aibot", chatID, msgID)
ref, err := store.Store(tmpPath, media.MediaMeta{ ref, err := store.Store(tmpPath, media.MediaMeta{
Filename: msgID + ".jpg", Filename: msgID + ext,
Source: "wecom_aibot", Source: "wecom_aibot",
}, scope) }, scope)
if err != nil { if err != nil {
@ -1200,3 +1212,71 @@ func (c *WeComAIBotWSChannel) storeWSImage(
} }
return ref, nil return ref, nil
} }
// wsMediaExtFromContentType returns the lowercase file extension (with leading
// dot) for the given Content-Type value, or "" when the type is unrecognized.
func wsMediaExtFromContentType(contentType string) string {
if contentType == "" {
return ""
}
// Strip parameters (e.g. "image/jpeg; charset=utf-8" → "image/jpeg").
mt := strings.ToLower(strings.TrimSpace(strings.SplitN(contentType, ";", 2)[0]))
switch mt {
case "image/jpeg", "image/jpg":
return ".jpg"
case "image/png":
return ".png"
case "image/gif":
return ".gif"
case "image/webp":
return ".webp"
case "video/mp4":
return ".mp4"
case "video/mpeg", "video/x-mpeg":
return ".mpeg"
case "video/quicktime":
return ".mov"
case "video/webm":
return ".webm"
case "audio/mpeg", "audio/mp3":
return ".mp3"
case "audio/ogg":
return ".ogg"
case "audio/wav":
return ".wav"
case "application/pdf":
return ".pdf"
case "application/zip":
return ".zip"
case "application/x-rar-compressed", "application/vnd.rar":
return ".rar"
case "text/plain":
return ".txt"
case "application/msword":
return ".doc"
case "application/vnd.openxmlformats-officedocument.wordprocessingml.document":
return ".docx"
case "application/vnd.ms-excel":
return ".xls"
case "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet":
return ".xlsx"
case "application/vnd.ms-powerpoint":
return ".ppt"
case "application/vnd.openxmlformats-officedocument.presentationml.presentation":
return ".pptx"
}
return ""
}
// wsLabelToDefaultExt returns the default file extension for the given media label
// used in wsHandleMediaMessage. It is the fallback when Content-Type detection fails.
func wsLabelToDefaultExt(label string) string {
switch label {
case "image":
return ".jpg"
case "video":
return ".mp4"
default: // "file" and any future labels
return ".bin"
}
}

View file

@ -29,19 +29,19 @@ func newTestWSChannel(t *testing.T) *WeComAIBotWSChannel {
return ch return ch
} }
// TestStoreWSImage_NilStore verifies that storeWSImage returns an error when no // TestStoreWSMedia_NilStore verifies that storeWSMedia returns an error when no
// MediaStore has been injected. // MediaStore has been injected.
func TestStoreWSImage_NilStore(t *testing.T) { func TestStoreWSMedia_NilStore(t *testing.T) {
ch := newTestWSChannel(t) ch := newTestWSChannel(t)
_, err := ch.storeWSImage(context.Background(), "chat1", "msg1", "http://any", "") _, err := ch.storeWSMedia(context.Background(), "chat1", "msg1", "http://any", "", ".jpg")
if err == nil { if err == nil {
t.Fatal("expected error when no MediaStore is set") t.Fatal("expected error when no MediaStore is set")
} }
} }
// TestStoreWSImage_HTTPError verifies that storeWSImage propagates HTTP errors // TestStoreWSMedia_HTTPError verifies that storeWSMedia propagates HTTP errors
// from the image server. // from the media server.
func TestStoreWSImage_HTTPError(t *testing.T) { func TestStoreWSMedia_HTTPError(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
http.Error(w, "not found", http.StatusNotFound) http.Error(w, "not found", http.StatusNotFound)
})) }))
@ -50,29 +50,29 @@ func TestStoreWSImage_HTTPError(t *testing.T) {
ch := newTestWSChannel(t) ch := newTestWSChannel(t)
ch.SetMediaStore(media.NewFileMediaStore()) ch.SetMediaStore(media.NewFileMediaStore())
_, err := ch.storeWSImage(context.Background(), "chat1", "msg1", srv.URL, "") _, err := ch.storeWSMedia(context.Background(), "chat1", "msg1", srv.URL, "", ".jpg")
if err == nil { if err == nil {
t.Fatal("expected error for HTTP 404") t.Fatal("expected error for HTTP 404")
} }
} }
// TestStoreWSImage_ServerUnavailable verifies that storeWSImage returns a clear // TestStoreWSMedia_ServerUnavailable verifies that storeWSMedia returns a clear
// error when the image server cannot be reached. // error when the media server cannot be reached.
func TestStoreWSImage_ServerUnavailable(t *testing.T) { func TestStoreWSMedia_ServerUnavailable(t *testing.T) {
ch := newTestWSChannel(t) ch := newTestWSChannel(t)
ch.SetMediaStore(media.NewFileMediaStore()) ch.SetMediaStore(media.NewFileMediaStore())
// Port 1 is reserved and will refuse the connection immediately. // Port 1 is reserved and will refuse the connection immediately.
_, err := ch.storeWSImage(context.Background(), "chat1", "msg1", "http://127.0.0.1:1", "") _, err := ch.storeWSMedia(context.Background(), "chat1", "msg1", "http://127.0.0.1:1", "", ".jpg")
if err == nil { if err == nil {
t.Fatal("expected error for unreachable server") t.Fatal("expected error for unreachable server")
} }
} }
// TestStoreWSImage_Success_NoAES verifies the happy path: the image is downloaded, // TestStoreWSMedia_Success_NoAES verifies the happy path: the media is downloaded,
// a media ref is returned, and the file persists and is readable via Resolve until // a media ref is returned, and the file persists and is readable via Resolve until
// ReleaseAll is called. // ReleaseAll is called. The server returns no Content-Type, so the defaultExt is used.
func TestStoreWSImage_Success_NoAES(t *testing.T) { func TestStoreWSMedia_Success_NoAES(t *testing.T) {
imageData := bytes.Repeat([]byte("x"), 256) imageData := bytes.Repeat([]byte("x"), 256)
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK) w.WriteHeader(http.StatusOK)
@ -84,7 +84,7 @@ func TestStoreWSImage_Success_NoAES(t *testing.T) {
store := media.NewFileMediaStore() store := media.NewFileMediaStore()
ch.SetMediaStore(store) ch.SetMediaStore(store)
ref, err := ch.storeWSImage(context.Background(), "chat1", "msg1", srv.URL, "") ref, err := ch.storeWSMedia(context.Background(), "chat1", "msg1", srv.URL, "", ".jpg")
if err != nil { if err != nil {
t.Fatalf("expected no error, got %v", err) t.Fatalf("expected no error, got %v", err)
} }
@ -92,7 +92,7 @@ func TestStoreWSImage_Success_NoAES(t *testing.T) {
t.Fatal("expected non-empty ref") t.Fatal("expected non-empty ref")
} }
// File must be accessible after storeWSImage returns (no premature deletion). // File must be accessible after storeWSMedia returns (no premature deletion).
path, err := store.Resolve(ref) path, err := store.Resolve(ref)
if err != nil { if err != nil {
t.Fatalf("ref should resolve: %v", err) t.Fatalf("ref should resolve: %v", err)
@ -115,9 +115,9 @@ func TestStoreWSImage_Success_NoAES(t *testing.T) {
} }
} }
// TestStoreWSImage_MultipleMessages verifies that concurrent image messages with // TestStoreWSMedia_MultipleMessages verifies that concurrent media messages with
// different msgIDs do not collide and each resolve to distinct files. // different msgIDs do not collide and each resolve to distinct files.
func TestStoreWSImage_MultipleMessages(t *testing.T) { func TestStoreWSMedia_MultipleMessages(t *testing.T) {
imageA := bytes.Repeat([]byte("a"), 64) imageA := bytes.Repeat([]byte("a"), 64)
imageB := bytes.Repeat([]byte("b"), 64) imageB := bytes.Repeat([]byte("b"), 64)
@ -136,13 +136,13 @@ func TestStoreWSImage_MultipleMessages(t *testing.T) {
store := media.NewFileMediaStore() store := media.NewFileMediaStore()
ch.SetMediaStore(store) ch.SetMediaStore(store)
refA, err := ch.storeWSImage(context.Background(), "chat1", "msgA", srvA.URL, "") refA, err := ch.storeWSMedia(context.Background(), "chat1", "msgA", srvA.URL, "", ".jpg")
if err != nil { if err != nil {
t.Fatalf("storeWSImage A: %v", err) t.Fatalf("storeWSMedia A: %v", err)
} }
refB, err := ch.storeWSImage(context.Background(), "chat1", "msgB", srvB.URL, "") refB, err := ch.storeWSMedia(context.Background(), "chat1", "msgB", srvB.URL, "", ".jpg")
if err != nil { if err != nil {
t.Fatalf("storeWSImage B: %v", err) t.Fatalf("storeWSMedia B: %v", err)
} }
if refA == refB { if refA == refB {
t.Fatal("distinct messages must produce distinct refs") t.Fatal("distinct messages must produce distinct refs")
@ -163,3 +163,56 @@ func TestStoreWSImage_MultipleMessages(t *testing.T) {
t.Errorf("content mismatch for message B") t.Errorf("content mismatch for message B")
} }
} }
// TestStoreWSMedia_ContentTypeExt verifies that the file extension is inferred
// from the HTTP Content-Type header and the defaultExt fallback is used when the
// type is absent or unrecognized.
func TestStoreWSMedia_ContentTypeExt(t *testing.T) {
tests := []struct {
contentType string
wantExt string
}{
{"image/jpeg", ".jpg"},
{"image/png", ".png"},
{"video/mp4", ".mp4"},
{"application/pdf", ".pdf"},
{"application/zip", ".zip"},
// With parameters stripped.
{"video/mp4; codecs=avc1", ".mp4"},
// Unknown type → falls back to defaultExt.
{"", ""},
{"application/octet-stream", ""},
}
for _, tc := range tests {
got := wsMediaExtFromContentType(tc.contentType)
if got != tc.wantExt {
t.Errorf("wsMediaExtFromContentType(%q) = %q, want %q", tc.contentType, got, tc.wantExt)
}
}
// End-to-end: server returns Content-Type: video/mp4, defaultExt is .bin.
// The stored file should carry the .mp4 extension, not .bin.
payload := bytes.Repeat([]byte("v"), 128)
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "video/mp4")
w.WriteHeader(http.StatusOK)
_, _ = w.Write(payload)
}))
defer srv.Close()
ch := newTestWSChannel(t)
store := media.NewFileMediaStore()
ch.SetMediaStore(store)
ref, err := ch.storeWSMedia(context.Background(), "chat1", "vid1", srv.URL, "", ".bin")
if err != nil {
t.Fatalf("storeWSMedia: %v", err)
}
path, err := store.Resolve(ref)
if err != nil {
t.Fatalf("resolve: %v", err)
}
if ext := path[len(path)-4:]; ext != ".mp4" {
t.Errorf("expected .mp4 extension from Content-Type, got %q", ext)
}
}