fix(agent): gate image model on current-turn media only
This commit is contained in:
parent
d750bfcca5
commit
89ff10ceb0
2 changed files with 114 additions and 7 deletions
|
|
@ -1683,10 +1683,19 @@ func (al *AgentLoop) runTurn(ctx context.Context, ts *turnState) (turnResult, er
|
||||||
ts.recordPersistedMessage(rootMsg)
|
ts.recordPersistedMessage(rootMsg)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
currentTurnResolved := resolveMediaRefs([]providers.Message{
|
||||||
|
{Role: "user", Media: append([]string(nil), ts.media...)},
|
||||||
|
}, al.mediaStore, maxMediaSize)
|
||||||
|
var currentTurnMedia []string
|
||||||
|
if len(currentTurnResolved) > 0 {
|
||||||
|
currentTurnMedia = currentTurnResolved[0].Media
|
||||||
|
}
|
||||||
|
|
||||||
activeCandidates, activeModel, usedLight, useImageFallback := al.selectCandidates(
|
activeCandidates, activeModel, usedLight, useImageFallback := al.selectCandidates(
|
||||||
ts.agent,
|
ts.agent,
|
||||||
ts.userMessage,
|
ts.userMessage,
|
||||||
messages,
|
messages,
|
||||||
|
currentTurnMedia,
|
||||||
)
|
)
|
||||||
activeProvider := ts.agent.Provider
|
activeProvider := ts.agent.Provider
|
||||||
if usedLight && ts.agent.LightProvider != nil {
|
if usedLight && ts.agent.LightProvider != nil {
|
||||||
|
|
@ -2764,8 +2773,9 @@ func (al *AgentLoop) selectCandidates(
|
||||||
agent *AgentInstance,
|
agent *AgentInstance,
|
||||||
userMsg string,
|
userMsg string,
|
||||||
history []providers.Message,
|
history []providers.Message,
|
||||||
|
currentTurnMedia []string,
|
||||||
) (candidates []providers.FallbackCandidate, model string, usedLight bool, useImageFallback bool) {
|
) (candidates []providers.FallbackCandidate, model string, usedLight bool, useImageFallback bool) {
|
||||||
if hasImageMedia(history) && len(agent.ImageCandidates) > 0 {
|
if hasImageMediaRefs(currentTurnMedia) && len(agent.ImageCandidates) > 0 {
|
||||||
logger.InfoCF("agent", "Image model selected",
|
logger.InfoCF("agent", "Image model selected",
|
||||||
map[string]any{
|
map[string]any{
|
||||||
"agent_id": agent.ID,
|
"agent_id": agent.ID,
|
||||||
|
|
@ -2798,14 +2808,12 @@ func (al *AgentLoop) selectCandidates(
|
||||||
return agent.LightCandidates, resolvedCandidateModel(agent.LightCandidates, agent.Router.LightModel()), true, false
|
return agent.LightCandidates, resolvedCandidateModel(agent.LightCandidates, agent.Router.LightModel()), true, false
|
||||||
}
|
}
|
||||||
|
|
||||||
func hasImageMedia(messages []providers.Message) bool {
|
func hasImageMediaRefs(mediaRefs []string) bool {
|
||||||
for _, msg := range messages {
|
for _, ref := range mediaRefs {
|
||||||
for _, ref := range msg.Media {
|
|
||||||
if strings.HasPrefix(strings.ToLower(ref), "data:image/") {
|
if strings.HasPrefix(strings.ToLower(ref), "data:image/") {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -462,6 +462,7 @@ func TestSelectCandidates_UsesImageModelWhenImageMediaPresent(t *testing.T) {
|
||||||
[]providers.Message{
|
[]providers.Message{
|
||||||
{Role: "user", Content: "describe this image", Media: []string{"data:image/png;base64,AAAA"}},
|
{Role: "user", Content: "describe this image", Media: []string{"data:image/png;base64,AAAA"}},
|
||||||
},
|
},
|
||||||
|
[]string{"data:image/png;base64,AAAA"},
|
||||||
)
|
)
|
||||||
|
|
||||||
if usedLight {
|
if usedLight {
|
||||||
|
|
@ -481,6 +482,48 @@ func TestSelectCandidates_UsesImageModelWhenImageMediaPresent(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSelectCandidates_DoesNotUseImageModelForHistoricalImages(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: tmpDir,
|
||||||
|
ModelName: "text-main",
|
||||||
|
ImageModel: "vision-main",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
ModelList: []*config.ModelConfig{
|
||||||
|
{ModelName: "text-main", Model: "openai/gpt-5.4"},
|
||||||
|
{ModelName: "vision-main", Model: "gemini/gemini-2.5-flash-lite"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, &mockProvider{})
|
||||||
|
candidates, model, useImageFallback := (&AgentLoop{}).selectCandidates(
|
||||||
|
agent,
|
||||||
|
"text-only follow-up",
|
||||||
|
[]providers.Message{
|
||||||
|
{Role: "user", Content: "earlier image", Media: []string{"data:image/png;base64,AAAA"}},
|
||||||
|
{Role: "assistant", Content: "previous response"},
|
||||||
|
{Role: "user", Content: "current text-only turn"},
|
||||||
|
},
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
|
||||||
|
if useImageFallback {
|
||||||
|
t.Fatal("expected image fallback to be disabled for text-only current turn")
|
||||||
|
}
|
||||||
|
if model != "gpt-5.4" {
|
||||||
|
t.Fatalf("model = %q, want %q", model, "gpt-5.4")
|
||||||
|
}
|
||||||
|
if len(candidates) != 1 {
|
||||||
|
t.Fatalf("len(candidates) = %d, want 1", len(candidates))
|
||||||
|
}
|
||||||
|
if candidates[0].Provider != "openai" || candidates[0].Model != "gpt-5.4" {
|
||||||
|
t.Fatalf("candidate = %+v, want openai/gpt-5.4", candidates[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestRunAgentLoop_UsesImageModelForImageMessages(t *testing.T) {
|
func TestRunAgentLoop_UsesImageModelForImageMessages(t *testing.T) {
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
cfg := &config.Config{
|
cfg := &config.Config{
|
||||||
|
|
@ -521,6 +564,62 @@ func TestRunAgentLoop_UsesImageModelForImageMessages(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRunAgentLoop_DoesNotUseImageModelWhenOnlyHistoryHasImages(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
Workspace: tmpDir,
|
||||||
|
ModelName: "text-main",
|
||||||
|
ImageModel: "vision-main",
|
||||||
|
MaxTokens: 4096,
|
||||||
|
MaxToolIterations: 5,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
ModelList: []*config.ModelConfig{
|
||||||
|
{ModelName: "text-main", Model: "openai/gpt-5.4"},
|
||||||
|
{ModelName: "vision-main", Model: "gemini/gemini-2.5-flash-lite"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
provider := &mockProvider{}
|
||||||
|
al := NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
agent := al.registry.GetDefaultAgent()
|
||||||
|
|
||||||
|
sessionKey := "image-history-session"
|
||||||
|
_, err := al.runAgentLoop(context.Background(), agent, processOptions{
|
||||||
|
SessionKey: sessionKey,
|
||||||
|
Channel: "telegram",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
UserMessage: "describe this image",
|
||||||
|
Media: []string{"data:image/png;base64,AAAA"},
|
||||||
|
DefaultResponse: "fallback",
|
||||||
|
SendResponse: false,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("first runAgentLoop returned error: %v", err)
|
||||||
|
}
|
||||||
|
if provider.lastModel != "gemini-2.5-flash-lite" {
|
||||||
|
t.Fatalf("first run provider lastModel = %q, want %q", provider.lastModel, "gemini-2.5-flash-lite")
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = al.runAgentLoop(context.Background(), agent, processOptions{
|
||||||
|
SessionKey: sessionKey,
|
||||||
|
Channel: "telegram",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
UserMessage: "now answer text only",
|
||||||
|
DefaultResponse: "fallback",
|
||||||
|
SendResponse: false,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("second runAgentLoop returned error: %v", err)
|
||||||
|
}
|
||||||
|
if provider.lastModel != "gpt-5.4" {
|
||||||
|
t.Fatalf("second run provider lastModel = %q, want %q", provider.lastModel, "gpt-5.4")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestNewAgentLoop_StateInitialized(t *testing.T) {
|
func TestNewAgentLoop_StateInitialized(t *testing.T) {
|
||||||
// Create temp workspace
|
// Create temp workspace
|
||||||
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue