diff --git a/pkg/agent/agent_tool.go b/pkg/agent/agent_tool.go index 0a8fda8f3..0fb112712 100644 --- a/pkg/agent/agent_tool.go +++ b/pkg/agent/agent_tool.go @@ -75,8 +75,22 @@ func (t *agentTool) Execute(ctx context.Context, arguments string) (ToolResult, return ToolResult{}, &MaxToolDepthExceededError{MaxDepth: t.agent.maxToolDepth} } - var params agentToolParams + var fields map[string]json.RawMessage + if err := json.Unmarshal([]byte(arguments), &fields); err != nil { + return ToolResult{ + Content: fmt.Sprintf("Invalid parameters: %s", err.Error()), + IsError: true, + }, nil + } + if _, ok := fields["input"]; !ok { + return ToolResult{ + Content: "Missing required parameters: input", + IsError: true, + }, nil + } + + var params agentToolParams if err := json.Unmarshal([]byte(arguments), ¶ms); err != nil { return ToolResult{ Content: fmt.Sprintf("Invalid parameters: %s", err.Error()), diff --git a/pkg/agent/agent_tool_test.go b/pkg/agent/agent_tool_test.go index 73ca93028..4a16e139e 100644 --- a/pkg/agent/agent_tool_test.go +++ b/pkg/agent/agent_tool_test.go @@ -195,19 +195,13 @@ func TestAgentTool_Execute(t *testing.T) { ) t.Run( - "empty JSON object returns tool error", + "empty JSON object returns tool error for missing input", func(t *testing.T) { t.Parallel() - provider := &mockProvider{ - responses: []*llm.ChatCompletionResponse{ - stopResponse("ok"), - }, - } - ag := agent.New( "sub", - newTestClient(provider), + newTestClient(&mockProvider{}), agent.WithModel("test-model"), ) @@ -215,7 +209,9 @@ func TestAgentTool_Execute(t *testing.T) { result, err := tool.Execute(context.Background(), `{}`) require.NoError(t, err) - assert.False(t, result.IsError) + assert.True(t, result.IsError) + assert.Contains(t, result.Content, "Missing required parameters") + assert.Contains(t, result.Content, "input") }, )