feat(agent): thread the new turn through Loop as llm.Message
Run/RunStream took the new turn as a bare string, which had nowhere
to carry ContentPart attachments. Both now take an llm.Message
(Role is forced to RoleUser regardless of what the caller sets), so a
caller building a multimodal turn just fills in Content/Parts on it
instead of the loop needing a second, parallel parameter.
subagent.go and every test call site are updated to wrap their string
prompt as llm.Message{Role: llm.RoleUser, Content: ...} — SubAgent.Run
itself is untouched, it still takes a plain task string.
This commit is contained in:
parent
2f6f5fab1c
commit
49353485e5
6 changed files with 32 additions and 29 deletions
|
|
@ -103,7 +103,7 @@ func TestIntegration_AgentLoop_Generate(t *testing.T) {
|
||||||
MaxIters: 3,
|
MaxIters: 3,
|
||||||
})
|
})
|
||||||
|
|
||||||
resp, err := loop.Run(context.Background(), "What is 20+22? Use the add_numbers tool.")
|
resp, err := loop.Run(context.Background(), llm.Message{Role: llm.RoleUser, Content: "What is 20+22? Use the add_numbers tool."})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("run failed: %v", err)
|
t.Fatalf("run failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -130,7 +130,7 @@ func TestIntegration_AgentLoop_Stream(t *testing.T) {
|
||||||
MaxIters: 3,
|
MaxIters: 3,
|
||||||
})
|
})
|
||||||
|
|
||||||
stream := loop.RunStream(context.Background(), "Say something interesting.")
|
stream := loop.RunStream(context.Background(), llm.Message{Role: llm.RoleUser, Content: "Say something interesting."})
|
||||||
|
|
||||||
var chunks []llm.StreamChunk
|
var chunks []llm.StreamChunk
|
||||||
for chunk, err := range stream {
|
for chunk, err := range stream {
|
||||||
|
|
|
||||||
|
|
@ -77,8 +77,9 @@ func New(cfg Config) *Loop {
|
||||||
|
|
||||||
// Run executes the agent loop and returns the final response.
|
// Run executes the agent loop and returns the final response.
|
||||||
// Optional history messages are appended after the system prompt and before
|
// Optional history messages are appended after the system prompt and before
|
||||||
// the new user input.
|
// the new user input. input's Role is overwritten to RoleUser regardless of
|
||||||
func (l *Loop) Run(ctx context.Context, input string, history ...llm.Message) (Response, error) {
|
// what the caller sets, so callers only need to fill in Content/Parts.
|
||||||
|
func (l *Loop) Run(ctx context.Context, input llm.Message, history ...llm.Message) (Response, error) {
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
|
|
||||||
messages := l.buildInitialMessages(input, history)
|
messages := l.buildInitialMessages(input, history)
|
||||||
|
|
@ -174,8 +175,9 @@ func (l *Loop) Run(ctx context.Context, input string, history ...llm.Message) (R
|
||||||
|
|
||||||
// RunStream executes the agent loop with streaming output.
|
// RunStream executes the agent loop with streaming output.
|
||||||
// Optional history messages are appended after the system prompt and before
|
// Optional history messages are appended after the system prompt and before
|
||||||
// the new user input.
|
// the new user input. input's Role is overwritten to RoleUser regardless of
|
||||||
func (l *Loop) RunStream(ctx context.Context, input string, history ...llm.Message) iter.Seq2[llm.StreamChunk, error] {
|
// what the caller sets, so callers only need to fill in Content/Parts.
|
||||||
|
func (l *Loop) RunStream(ctx context.Context, input llm.Message, history ...llm.Message) iter.Seq2[llm.StreamChunk, error] {
|
||||||
return func(yield func(llm.StreamChunk, error) bool) {
|
return func(yield func(llm.StreamChunk, error) bool) {
|
||||||
messages := l.buildInitialMessages(input, history)
|
messages := l.buildInitialMessages(input, history)
|
||||||
// Same as Run: the schemas are identical on every iteration.
|
// Same as Run: the schemas are identical on every iteration.
|
||||||
|
|
@ -362,12 +364,13 @@ func containsUnparsedToolCall(s string) bool {
|
||||||
return strings.Contains(s, "<tool_call") || strings.Contains(s, "<function=")
|
return strings.Contains(s, "<tool_call") || strings.Contains(s, "<function=")
|
||||||
}
|
}
|
||||||
|
|
||||||
func (l *Loop) buildInitialMessages(input string, history []llm.Message) []llm.Message {
|
func (l *Loop) buildInitialMessages(input llm.Message, history []llm.Message) []llm.Message {
|
||||||
systemPrompt := persona.AssembleSystemPrompt(l.cfg.Persona, l.cfg.AgentsMD)
|
systemPrompt := persona.AssembleSystemPrompt(l.cfg.Persona, l.cfg.AgentsMD)
|
||||||
messages := make([]llm.Message, 0, len(history)+2)
|
messages := make([]llm.Message, 0, len(history)+2)
|
||||||
messages = append(messages, llm.Message{Role: llm.RoleSystem, Content: systemPrompt})
|
messages = append(messages, llm.Message{Role: llm.RoleSystem, Content: systemPrompt})
|
||||||
messages = append(messages, history...)
|
messages = append(messages, history...)
|
||||||
messages = append(messages, llm.Message{Role: llm.RoleUser, Content: input})
|
input.Role = llm.RoleUser
|
||||||
|
messages = append(messages, input)
|
||||||
return messages
|
return messages
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -102,7 +102,7 @@ func TestRun_NoToolCalls(t *testing.T) {
|
||||||
Tools: tools.NewRegistry(),
|
Tools: tools.NewRegistry(),
|
||||||
})
|
})
|
||||||
|
|
||||||
resp, err := loop.Run(context.Background(), "Hello")
|
resp, err := loop.Run(context.Background(), llm.Message{Role: llm.RoleUser, Content: "Hello"})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -133,7 +133,7 @@ func TestRun_IncludesAgentsMD(t *testing.T) {
|
||||||
AgentsMD: "Never edit go.mod directly.",
|
AgentsMD: "Never edit go.mod directly.",
|
||||||
})
|
})
|
||||||
|
|
||||||
if _, err := loop.Run(context.Background(), "Hello"); err != nil {
|
if _, err := loop.Run(context.Background(), llm.Message{Role: llm.RoleUser, Content: "Hello"}); err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
if !strings.Contains(capturedSystemPrompt, "Never edit go.mod directly.") {
|
if !strings.Contains(capturedSystemPrompt, "Never edit go.mod directly.") {
|
||||||
|
|
@ -176,7 +176,7 @@ func TestRun_ToolCalls(t *testing.T) {
|
||||||
Tools: registry,
|
Tools: registry,
|
||||||
})
|
})
|
||||||
|
|
||||||
resp, err := loop.Run(context.Background(), "Say hi")
|
resp, err := loop.Run(context.Background(), llm.Message{Role: llm.RoleUser, Content: "Say hi"})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -233,7 +233,7 @@ func TestRun_ToolCalls_RecordsAssistantTurnAndToolCallID(t *testing.T) {
|
||||||
Tools: registry,
|
Tools: registry,
|
||||||
})
|
})
|
||||||
|
|
||||||
if _, err := loop.Run(context.Background(), "Say hi"); err != nil {
|
if _, err := loop.Run(context.Background(), llm.Message{Role: llm.RoleUser, Content: "Say hi"}); err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -315,7 +315,7 @@ func TestRun_Stream_ToolCalls_RecordsAssistantTurnAndToolCallID(t *testing.T) {
|
||||||
Tools: registry,
|
Tools: registry,
|
||||||
})
|
})
|
||||||
|
|
||||||
for _, err := range loop.RunStream(context.Background(), "test") {
|
for _, err := range loop.RunStream(context.Background(), llm.Message{Role: llm.RoleUser, Content: "test"}) {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected stream error: %v", err)
|
t.Fatalf("unexpected stream error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -379,7 +379,7 @@ func TestRun_MaxIterations(t *testing.T) {
|
||||||
MaxIters: 3,
|
MaxIters: 3,
|
||||||
})
|
})
|
||||||
|
|
||||||
_, err := loop.Run(context.Background(), "test")
|
_, err := loop.Run(context.Background(), llm.Message{Role: llm.RoleUser, Content: "test"})
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Fatal("expected error, got nil")
|
t.Fatal("expected error, got nil")
|
||||||
}
|
}
|
||||||
|
|
@ -408,7 +408,7 @@ func TestRun_ToolNotFound(t *testing.T) {
|
||||||
Tools: tools.NewRegistry(),
|
Tools: tools.NewRegistry(),
|
||||||
})
|
})
|
||||||
|
|
||||||
resp, err := loop.Run(context.Background(), "test")
|
resp, err := loop.Run(context.Background(), llm.Message{Role: llm.RoleUser, Content: "test"})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -451,7 +451,7 @@ func TestRun_ApprovalDenied(t *testing.T) {
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
|
|
||||||
resp, err := loop.Run(context.Background(), "test")
|
resp, err := loop.Run(context.Background(), llm.Message{Role: llm.RoleUser, Content: "test"})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -496,7 +496,7 @@ func TestRun_SandboxViolation(t *testing.T) {
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
|
|
||||||
resp, err := loop.Run(context.Background(), "test")
|
resp, err := loop.Run(context.Background(), llm.Message{Role: llm.RoleUser, Content: "test"})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -540,7 +540,7 @@ func TestRun_OnIterationHook(t *testing.T) {
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
|
|
||||||
_, err := loop.Run(context.Background(), "test")
|
_, err := loop.Run(context.Background(), llm.Message{Role: llm.RoleUser, Content: "test"})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -570,7 +570,7 @@ func TestRun_Stream_NoToolCalls(t *testing.T) {
|
||||||
})
|
})
|
||||||
|
|
||||||
var chunks []llm.StreamChunk
|
var chunks []llm.StreamChunk
|
||||||
stream := loop.RunStream(context.Background(), "test")
|
stream := loop.RunStream(context.Background(), llm.Message{Role: llm.RoleUser, Content: "test"})
|
||||||
for chunk, err := range stream {
|
for chunk, err := range stream {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
|
@ -608,7 +608,7 @@ func TestRun_Stream_ForwardsTrailingUsageOnlyChunk(t *testing.T) {
|
||||||
})
|
})
|
||||||
|
|
||||||
var chunks []llm.StreamChunk
|
var chunks []llm.StreamChunk
|
||||||
stream := loop.RunStream(context.Background(), "test")
|
stream := loop.RunStream(context.Background(), llm.Message{Role: llm.RoleUser, Content: "test"})
|
||||||
for chunk, err := range stream {
|
for chunk, err := range stream {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
|
@ -661,7 +661,7 @@ func TestRun_Stream_WithToolCalls(t *testing.T) {
|
||||||
})
|
})
|
||||||
|
|
||||||
var chunks []llm.StreamChunk
|
var chunks []llm.StreamChunk
|
||||||
stream := loop.RunStream(context.Background(), "test")
|
stream := loop.RunStream(context.Background(), llm.Message{Role: llm.RoleUser, Content: "test"})
|
||||||
for chunk, err := range stream {
|
for chunk, err := range stream {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
|
@ -713,7 +713,7 @@ func TestRun_Stream_MaxIterations(t *testing.T) {
|
||||||
})
|
})
|
||||||
|
|
||||||
var chunks []llm.StreamChunk
|
var chunks []llm.StreamChunk
|
||||||
stream := loop.RunStream(context.Background(), "test")
|
stream := loop.RunStream(context.Background(), llm.Message{Role: llm.RoleUser, Content: "test"})
|
||||||
for chunk, err := range stream {
|
for chunk, err := range stream {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// expect max iterations error
|
// expect max iterations error
|
||||||
|
|
@ -748,7 +748,7 @@ func TestRun_Timeout(t *testing.T) {
|
||||||
Tools: tools.NewRegistry(),
|
Tools: tools.NewRegistry(),
|
||||||
})
|
})
|
||||||
|
|
||||||
_, err := loop.Run(ctx, "test")
|
_, err := loop.Run(ctx, llm.Message{Role: llm.RoleUser, Content: "test"})
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Fatal("expected timeout error, got nil")
|
t.Fatal("expected timeout error, got nil")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -33,7 +33,7 @@ func (s SubAgent) Run(ctx context.Context, llmClient llm.LLMClient, agentsMD str
|
||||||
if cfg.MaxIters == 0 {
|
if cfg.MaxIters == 0 {
|
||||||
cfg.MaxIters = DefaultMaxIterations
|
cfg.MaxIters = DefaultMaxIterations
|
||||||
}
|
}
|
||||||
return New(cfg).Run(ctx, task)
|
return New(cfg).Run(ctx, llm.Message{Role: llm.RoleUser, Content: task})
|
||||||
}
|
}
|
||||||
|
|
||||||
// SubAgentRegistry looks up SubAgents by name for the delegate tool.
|
// SubAgentRegistry looks up SubAgents by name for the delegate tool.
|
||||||
|
|
|
||||||
|
|
@ -28,7 +28,7 @@ func TestRunStream_RecoversFromThinkingBudgetCut(t *testing.T) {
|
||||||
loop := New(Config{LLM: stub, Tools: editTestRegistry(t, &executed), MaxIters: 10})
|
loop := New(Config{LLM: stub, Tools: editTestRegistry(t, &executed), MaxIters: 10})
|
||||||
|
|
||||||
var final strings.Builder
|
var final strings.Builder
|
||||||
for chunk, err := range loop.RunStream(context.Background(), "arregla x.py") {
|
for chunk, err := range loop.RunStream(context.Background(), llm.Message{Role: llm.RoleUser, Content: "arregla x.py"}) {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -62,7 +62,7 @@ func TestRunStream_ThinkingBudgetNudgeGivesUpAfterLimit(t *testing.T) {
|
||||||
|
|
||||||
loop := New(Config{LLM: stub, Tools: editTestRegistry(t, &executed), MaxIters: 10})
|
loop := New(Config{LLM: stub, Tools: editTestRegistry(t, &executed), MaxIters: 10})
|
||||||
|
|
||||||
for _, err := range loop.RunStream(context.Background(), "haz algo") {
|
for _, err := range loop.RunStream(context.Background(), llm.Message{Role: llm.RoleUser, Content: "haz algo"}) {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -97,7 +97,7 @@ func TestRunStream_RecoversFromUnparsedToolCall(t *testing.T) {
|
||||||
loop := New(Config{LLM: stub, Tools: editTestRegistry(t, &executed), MaxIters: 10})
|
loop := New(Config{LLM: stub, Tools: editTestRegistry(t, &executed), MaxIters: 10})
|
||||||
|
|
||||||
var final strings.Builder
|
var final strings.Builder
|
||||||
for chunk, err := range loop.RunStream(context.Background(), "arregla x.py") {
|
for chunk, err := range loop.RunStream(context.Background(), llm.Message{Role: llm.RoleUser, Content: "arregla x.py"}) {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -133,7 +133,7 @@ func TestRunStream_NudgeGivesUpAfterLimit(t *testing.T) {
|
||||||
loop := New(Config{LLM: stub, Tools: editTestRegistry(t, &executed), MaxIters: 10})
|
loop := New(Config{LLM: stub, Tools: editTestRegistry(t, &executed), MaxIters: 10})
|
||||||
|
|
||||||
rounds := 0
|
rounds := 0
|
||||||
for _, err := range loop.RunStream(context.Background(), "haz algo") {
|
for _, err := range loop.RunStream(context.Background(), llm.Message{Role: llm.RoleUser, Content: "haz algo"}) {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -158,7 +158,7 @@ func TestRun_RecoversFromUnparsedToolCall(t *testing.T) {
|
||||||
}}
|
}}
|
||||||
|
|
||||||
loop := New(Config{LLM: stub, Tools: editTestRegistry(t, &executed), MaxIters: 10})
|
loop := New(Config{LLM: stub, Tools: editTestRegistry(t, &executed), MaxIters: 10})
|
||||||
resp, err := loop.Run(context.Background(), "arregla x.py")
|
resp, err := loop.Run(context.Background(), llm.Message{Role: llm.RoleUser, Content: "arregla x.py"})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue