diff --git a/pkg/agent/agent_test.go b/pkg/agent/agent_test.go index dd13995b..a4f651e6 100644 --- a/pkg/agent/agent_test.go +++ b/pkg/agent/agent_test.go @@ -4484,6 +4484,80 @@ func (p *visionAnswerProvider) GetDefaultModel() string { return "vision-answer-model" } +type resolvedImagePathVisionProvider struct { + expectedPath string + calls int + models []string + pathTagSeen []bool + mediaSeen []bool +} + +func (p *resolvedImagePathVisionProvider) Chat( + ctx context.Context, + messages []providers.Message, + tools []providers.ToolDefinition, + model string, + opts map[string]any, +) (*providers.LLMResponse, error) { + p.calls++ + p.models = append(p.models, model) + + pathTag := "[image:" + p.expectedPath + "]" + hasPathTag := false + hasMedia := false + for _, msg := range messages { + if strings.Contains(msg.Content, pathTag) { + hasPathTag = true + } + for _, ref := range msg.Media { + if strings.TrimSpace(ref) != "" { + hasMedia = true + break + } + } + } + p.pathTagSeen = append(p.pathTagSeen, hasPathTag) + p.mediaSeen = append(p.mediaSeen, hasMedia) + + if !hasPathTag { + return nil, fmt.Errorf("vision provider expected resolved image path tag %q", pathTag) + } + if hasMedia { + return nil, fmt.Errorf("vision provider expected resolved attachment turn without raw media refs") + } + + return &providers.LLMResponse{ + Content: "vision direct answer", + ToolCalls: []providers.ToolCall{}, + }, nil +} + +func (p *resolvedImagePathVisionProvider) GetDefaultModel() string { + return "resolved-image-path-vision-model" +} + +type unexpectedTextAttachmentProvider struct { + calls int +} + +func (p *unexpectedTextAttachmentProvider) Chat( + ctx context.Context, + messages []providers.Message, + tools []providers.ToolDefinition, + model string, + opts map[string]any, +) (*providers.LLMResponse, error) { + p.calls++ + return &providers.LLMResponse{ + Content: "text model response", + ToolCalls: []providers.ToolCall{}, + }, nil +} + +func (p *unexpectedTextAttachmentProvider) GetDefaultModel() string { + return "unexpected-text-attachment-model" +} + func TestAgentLoop_VisionUnsupportedErrorReturnsClearFailure(t *testing.T) { workspace := t.TempDir() @@ -4548,6 +4622,95 @@ func TestAgentLoop_VisionUnsupportedErrorReturnsClearFailure(t *testing.T) { } } +func TestAgentLoop_UserAttachmentRoutesToImageModelAfterMediaResolution(t *testing.T) { + workspace := t.TempDir() + pngPath := filepath.Join(workspace, "sample.png") + pngBytes := []byte{ + 0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, + 0x00, 0x00, 0x00, 0x0D, 0x49, 0x48, 0x44, 0x52, + 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x08, 0x02, + 0x00, 0x00, 0x00, 0x90, 0x77, 0x53, 0xDE, + } + if err := os.WriteFile(pngPath, pngBytes, 0o644); err != nil { + t.Fatalf("WriteFile() error = %v", err) + } + + cfg := &config.Config{ + Agents: config.AgentsConfig{ + Defaults: config.AgentDefaults{ + Workspace: workspace, + ModelName: "text-model", + ImageModel: "vision-model", + MaxTokens: 4096, + MaxToolIterations: 3, + }, + }, + ModelList: []*config.ModelConfig{ + {ModelName: "text-model", Model: "openai/text-model"}, + {ModelName: "vision-model", Model: "openai/vision-model"}, + }, + } + + msgBus := bus.NewMessageBus() + textProvider := &unexpectedTextAttachmentProvider{} + al := NewAgentLoop(cfg, msgBus, textProvider) + + store := media.NewFileMediaStore() + al.SetMediaStore(store) + ref, err := store.Store(pngPath, media.MediaMeta{ContentType: "image/png"}, "test:user-attachment") + if err != nil { + t.Fatalf("Store() error = %v", err) + } + + agent := al.registry.GetDefaultAgent() + if agent == nil { + t.Fatal("expected default agent") + } + if len(agent.ImageCandidates) != 1 { + t.Fatalf("len(ImageCandidates) = %d, want 1", len(agent.ImageCandidates)) + } + + visionProvider := &resolvedImagePathVisionProvider{expectedPath: pngPath} + agent.CandidateProviders[providers.ModelKey("openai", "vision-model")] = visionProvider + + timeoutCtx, cancel := context.WithTimeout(context.Background(), responseTimeout) + defer cancel() + + resp, err := al.processMessage(timeoutCtx, testInboundMessage(bus.InboundMessage{ + Context: bus.InboundContext{ + Channel: "telegram", + ChatID: "chat1", + ChatType: "direct", + SenderID: "user1", + MessageID: "m1", + }, + Content: "describe this image", + Media: []string{ref}, + SessionKey: "agent:main:telegram:direct:user1", + })) + if err != nil { + t.Fatalf("processMessage() error = %v", err) + } + if resp != "vision direct answer" { + t.Fatalf("response = %q, want %q", resp, "vision direct answer") + } + if textProvider.calls != 0 { + t.Fatalf("textProvider calls = %d, want 0", textProvider.calls) + } + if visionProvider.calls != 1 { + t.Fatalf("visionProvider calls = %d, want %d", visionProvider.calls, 1) + } + if !slices.Equal(visionProvider.models, []string{"vision-model"}) { + t.Fatalf("visionProvider models = %v, want %v", visionProvider.models, []string{"vision-model"}) + } + if !slices.Equal(visionProvider.pathTagSeen, []bool{true}) { + t.Fatalf("visionProvider pathTagSeen = %v, want %v", visionProvider.pathTagSeen, []bool{true}) + } + if !slices.Equal(visionProvider.mediaSeen, []bool{false}) { + t.Fatalf("visionProvider mediaSeen = %v, want %v", visionProvider.mediaSeen, []bool{false}) + } +} + func TestAgentLoop_LoadImageFollowUpRoutesToImageModel(t *testing.T) { workspace := t.TempDir() pngPath := filepath.Join(workspace, "sample.png") diff --git a/pkg/agent/llm_media.go b/pkg/agent/llm_media.go index 043df9eb..54e1802f 100644 --- a/pkg/agent/llm_media.go +++ b/pkg/agent/llm_media.go @@ -2,12 +2,15 @@ package agent import ( "fmt" + "regexp" "strings" "github.com/sipeed/picoclaw/pkg/logger" "github.com/sipeed/picoclaw/pkg/providers" ) +var resolvedImagePathTagRegex = regexp.MustCompile(`\[image:[^\s\]][^\]]*\]`) + func messagesContainMedia(messages []providers.Message) bool { for _, msg := range messages { for _, ref := range msg.Media { @@ -104,8 +107,20 @@ func sameCandidateSet(a, b []providers.FallbackCandidate) bool { return true } +func messagesContainMediaTurn(messages []providers.Message) bool { + for _, msg := range messages { + if len(msg.Media) > 0 { + return true + } + if resolvedImagePathTagRegex.MatchString(msg.Content) { + return true + } + } + return false +} + func (p *Pipeline) routeMediaTurn(ts *turnState, exec *turnExecution) error { - if p == nil || ts == nil || ts.agent == nil || exec == nil || !hasMediaRefs(exec.callMessages) { + if p == nil || ts == nil || ts.agent == nil || exec == nil || !messagesContainMediaTurn(exec.callMessages) { return nil }