fix(agent): route resolved user attachments to image models
This commit is contained in:
parent
9fd0dbda96
commit
58926f76b9
2 changed files with 179 additions and 1 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue