fix(pico): preserve image media across pico attachments and client (#2874)

* fix(pico): preserve image media across pico attachments and client

* * fix ci

* fix(pico): preserve text when client media parsing fails

- Skip non-inline Pico attachment URLs instead of treating them as invalid inline media
- Preserve pico_client text messages when malformed media payloads are received
- Add regression coverage for media.create, download attachments, and invalid media payloads

* fix lint
This commit is contained in:
lxowalle 2026-05-15 15:49:07 +08:00 committed by GitHub
parent bfb2b35f74
commit 412705783d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 272 additions and 12 deletions

View file

@ -235,6 +235,8 @@ func (c *PicoClientChannel) handleInbound(pc *picoConn, msg PicoMessage) {
case TypeMessageCreate:
// Server sent us a message — treat as inbound
c.handleServerMessage(pc, msg)
case TypeMediaCreate:
c.handleServerMessage(pc, msg)
default:
logger.DebugCF("pico_client", "Ignoring message type", map[string]any{
"type": msg.Type,
@ -248,7 +250,17 @@ func (c *PicoClientChannel) handleServerMessage(pc *picoConn, msg PicoMessage) {
}
content, _ := msg.Payload[PayloadKeyContent].(string)
if strings.TrimSpace(content) == "" {
media, err := parseInlineImageMedia(msg.Payload)
if err != nil {
logger.WarnCF("pico_client", "Ignoring invalid media payload", map[string]any{
"error": err.Error(),
})
if strings.TrimSpace(content) == "" {
return
}
media = nil
}
if strings.TrimSpace(content) == "" && len(media) == 0 {
return
}
@ -281,7 +293,7 @@ func (c *PicoClientChannel) handleServerMessage(pc *picoConn, msg PicoMessage) {
},
}
c.HandleInboundContext(c.ctx, chatID, content, nil, inboundCtx, sender)
c.HandleInboundContext(c.ctx, chatID, content, media, inboundCtx, sender)
}
// Send sends a message to the remote server.

View file

@ -285,6 +285,24 @@ func TestParseInlineImageMedia_Valid(t *testing.T) {
}
}
func TestParseInlineImageMedia_Attachments(t *testing.T) {
imageURL := "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mP8/x8AAwMCAO+X2ioAAAAASUVORK5CYII="
media, err := parseInlineImageMedia(map[string]any{
"attachments": []any{
map[string]any{
"type": "image",
"url": imageURL,
},
},
})
if err != nil {
t.Fatalf("parseInlineImageMedia() error = %v", err)
}
if len(media) != 1 || media[0] != imageURL {
t.Fatalf("media = %#v, want attachment image payload", media)
}
}
func TestPicoChannel_HandleMessageSend_AllowsMediaOnly(t *testing.T) {
mb := bus.NewMessageBus()
bc := &config.Channel{Type: "pico", Enabled: true}
@ -326,6 +344,178 @@ func TestPicoChannel_HandleMessageSend_AllowsMediaOnly(t *testing.T) {
}
}
func newTestPicoClientChannel(t *testing.T) (*PicoClientChannel, *bus.MessageBus) {
t.Helper()
mb := bus.NewMessageBus()
bc := &config.Channel{Type: config.ChannelPicoClient, Enabled: true}
ch, err := NewPicoClientChannel(bc, &config.PicoClientSettings{
URL: "ws://localhost:8080/ws",
}, mb)
if err != nil {
t.Fatalf("NewPicoClientChannel() error = %v", err)
}
ch.ctx = context.Background()
return ch, mb
}
func assertInboundMessage(
t *testing.T,
mb *bus.MessageBus,
wantContent string,
wantMedia []string,
timeoutMessage string,
) {
t.Helper()
select {
case msg := <-mb.InboundChan():
if msg.Content != wantContent {
t.Fatalf("msg.Content = %q, want %s", msg.Content, wantContent)
}
if len(msg.Media) != len(wantMedia) {
t.Fatalf("msg.Media = %#v, want %#v", msg.Media, wantMedia)
}
for i := range wantMedia {
if msg.Media[i] != wantMedia[i] {
t.Fatalf("msg.Media = %#v, want %#v", msg.Media, wantMedia)
}
}
case <-time.After(time.Second):
t.Fatal(timeoutMessage)
}
}
func TestPicoClientChannel_HandleServerMessage_ForwardsMedia(t *testing.T) {
ch, mb := newTestPicoClientChannel(t)
pc := &picoConn{sessionID: "sess-media"}
imageURL := "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mP8/x8AAwMCAO+X2ioAAAAASUVORK5CYII="
ch.handleServerMessage(pc, PicoMessage{
Type: TypeMessageCreate,
Payload: map[string]any{
PayloadKeyContent: "describe this",
"attachments": []any{
map[string]any{
"type": "image",
"url": imageURL,
},
},
},
})
assertInboundMessage(
t,
mb,
"describe this",
[]string{imageURL},
"timed out waiting for forwarded media message",
)
}
func TestPicoClientChannel_HandleInbound_ForwardsMediaCreate(t *testing.T) {
ch, mb := newTestPicoClientChannel(t)
pc := &picoConn{sessionID: "sess-media-create"}
imageURL := "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mP8/x8AAwMCAO+X2ioAAAAASUVORK5CYII="
ch.handleInbound(pc, PicoMessage{
Type: TypeMediaCreate,
Payload: map[string]any{
PayloadKeyContent: "describe media.create",
"attachments": []any{
map[string]any{
"type": "image",
"url": imageURL,
},
},
},
})
assertInboundMessage(
t,
mb,
"describe media.create",
[]string{imageURL},
"timed out waiting for media.create message",
)
}
func TestPicoClientChannel_HandleServerMessage_ForwardsTextWithDownloadAttachment(t *testing.T) {
mb := bus.NewMessageBus()
bc := &config.Channel{Type: config.ChannelPicoClient, Enabled: true}
ch, err := NewPicoClientChannel(bc, &config.PicoClientSettings{
URL: "ws://localhost:8080/ws",
}, mb)
if err != nil {
t.Fatalf("NewPicoClientChannel() error = %v", err)
}
ch.ctx = context.Background()
pc := &picoConn{sessionID: "sess-download-attachment"}
ch.handleServerMessage(pc, PicoMessage{
Type: TypeMessageCreate,
Payload: map[string]any{
PayloadKeyContent: "see attached",
"attachments": []any{
map[string]any{
"type": "image",
"url": "/pico/media/abc",
"filename": "image.png",
"content_type": "image/png",
},
},
},
})
select {
case msg := <-mb.InboundChan():
if msg.Content != "see attached" {
t.Fatalf("msg.Content = %q, want see attached", msg.Content)
}
if len(msg.Media) != 0 {
t.Fatalf("msg.Media = %#v, want no inline media", msg.Media)
}
case <-time.After(time.Second):
t.Fatal("timed out waiting for text message with download attachment")
}
}
func TestPicoClientChannel_HandleServerMessage_ForwardsTextWithInvalidMediaPayload(t *testing.T) {
mb := bus.NewMessageBus()
bc := &config.Channel{Type: config.ChannelPicoClient, Enabled: true}
ch, err := NewPicoClientChannel(bc, &config.PicoClientSettings{
URL: "ws://localhost:8080/ws",
}, mb)
if err != nil {
t.Fatalf("NewPicoClientChannel() error = %v", err)
}
ch.ctx = context.Background()
pc := &picoConn{sessionID: "sess-invalid-media"}
ch.handleServerMessage(pc, PicoMessage{
Type: TypeMessageCreate,
Payload: map[string]any{
PayloadKeyContent: "hello despite invalid media",
"attachments": "not-an-array",
},
})
select {
case msg := <-mb.InboundChan():
if msg.Content != "hello despite invalid media" {
t.Fatalf("msg.Content = %q, want hello despite invalid media", msg.Content)
}
if len(msg.Media) != 0 {
t.Fatalf("msg.Media = %#v, want no inline media", msg.Media)
}
case <-time.After(time.Second):
t.Fatal("timed out waiting for text message with invalid media payload")
}
}
func TestIsThoughtPayload(t *testing.T) {
tests := []struct {
name string

View file

@ -990,11 +990,24 @@ func parseInlineImageMedia(payload map[string]any) ([]string, error) {
return nil, nil
}
raw, ok := payload["media"]
if !ok || raw == nil {
return nil, nil
media, err := parseInlineImageValues(payload["media"])
if err != nil {
return nil, err
}
attachments, err := parseInlineImageAttachments(payload["attachments"])
if err != nil {
return nil, err
}
media = append(media, attachments...)
return media, nil
}
func parseInlineImageValues(raw any) ([]string, error) {
if raw == nil {
return nil, nil
}
switch values := raw.(type) {
case []any:
media := make([]string, 0, len(values))
@ -1030,6 +1043,47 @@ func parseInlineImageMedia(payload map[string]any) ([]string, error) {
}
}
func parseInlineImageAttachments(raw any) ([]string, error) {
if raw == nil {
return nil, nil
}
values, ok := raw.([]any)
if !ok {
return nil, fmt.Errorf("attachments must be an array")
}
media := make([]string, 0, len(values))
for i, item := range values {
attachment, ok := item.(map[string]any)
if !ok {
return nil, fmt.Errorf("attachments[%d]: attachment must be an object", i)
}
attachmentType, _ := attachment["type"].(string)
attachmentType = strings.ToLower(strings.TrimSpace(attachmentType))
if attachmentType != "" && attachmentType != "image" {
continue
}
value, err := inlineImageValue(attachment)
if err != nil {
if attachmentType == "image" {
return nil, fmt.Errorf("attachments[%d]: %w", i, err)
}
continue
}
if !strings.HasPrefix(value, "data:") {
continue
}
if err := validateInlineImageDataURL(value); err != nil {
return nil, fmt.Errorf("attachments[%d]: %w", i, err)
}
media = append(media, value)
}
return media, nil
}
func inlineImageValue(item any) (string, error) {
switch value := item.(type) {
case string:

View file

@ -137,10 +137,10 @@ func NewProvider(ctx context.Context, opts ...Option) (*Provider, error) {
// converseParams holds the shared request parameters for Converse and ConverseStream.
type converseParams struct {
messages []types.Message
system []types.SystemContentBlock
messages []types.Message
system []types.SystemContentBlock
inferenceConfig *types.InferenceConfiguration
toolConfig *types.ToolConfiguration
toolConfig *types.ToolConfiguration
}
func buildConverseParams(messages []Message, tools []ToolDefinition, options map[string]any) converseParams {
@ -174,10 +174,10 @@ func buildConverseParams(messages []Message, tools []ToolDefinition, options map
}
return converseParams{
messages: bedrockMessages,
system: systemPrompts,
messages: bedrockMessages,
system: systemPrompts,
inferenceConfig: inferenceConfig,
toolConfig: toolConfig,
toolConfig: toolConfig,
}
}
@ -394,7 +394,11 @@ func parseStreamResponse(
usage = &UsageInfo{
PromptTokens: int(aws.ToInt32(e.Value.Usage.InputTokens)),
CompletionTokens: int(aws.ToInt32(e.Value.Usage.OutputTokens)),
TotalTokens: int(aws.ToInt32(e.Value.Usage.InputTokens)) + int(aws.ToInt32(e.Value.Usage.OutputTokens)),
TotalTokens: int(
aws.ToInt32(e.Value.Usage.InputTokens),
) + int(
aws.ToInt32(e.Value.Usage.OutputTokens),
),
}
}
}