From 493bb52fff0a6d9f97c5cadb64cddfdd448cd2c0 Mon Sep 17 00:00:00 2001 From: cointem Date: Thu, 19 Feb 2026 23:41:11 +0800 Subject: [PATCH] feat: add default temperature handling and update related tests --- pkg/agent/instance.go | 7 ++++++- pkg/agent/instance_test.go | 27 +++++++++++++++++++++++++++ pkg/agent/loop_test.go | 15 --------------- pkg/agent/mock_provider_test.go | 20 ++++++++++++++++++++ pkg/tools/subagent.go | 25 ++++++++++++------------- pkg/tools/subagent_tool_test.go | 31 ++++++++++++++++++++++++++++++- pkg/tools/toolloop.go | 2 -- 7 files changed, 95 insertions(+), 32 deletions(-) create mode 100644 pkg/agent/mock_provider_test.go diff --git a/pkg/agent/instance.go b/pkg/agent/instance.go index a992945fa..55fced78b 100644 --- a/pkg/agent/instance.go +++ b/pkg/agent/instance.go @@ -83,6 +83,11 @@ func NewAgentInstance( maxTokens = 8192 } + temperature := defaults.Temperature + if temperature == 0 { + temperature = 0.7 + } + // Resolve fallback candidates modelCfg := providers.ModelConfig{ Primary: model, @@ -98,7 +103,7 @@ func NewAgentInstance( Workspace: workspace, MaxIterations: maxIter, MaxTokens: maxTokens, - Temperature: defaults.Temperature, + Temperature: temperature, ContextWindow: maxTokens, Provider: provider, Sessions: sessionsManager, diff --git a/pkg/agent/instance_test.go b/pkg/agent/instance_test.go index 2500d1d88..baf824911 100644 --- a/pkg/agent/instance_test.go +++ b/pkg/agent/instance_test.go @@ -36,3 +36,30 @@ func TestNewAgentInstance_UsesDefaultsTemperatureAndMaxTokens(t *testing.T) { t.Fatalf("Temperature = %f, want %f", agent.Temperature, 1.0) } } + +func TestNewAgentInstance_DefaultsTemperatureWhenZero(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "agent-instance-test-*") + if err != nil { + t.Fatalf("Failed to create temp dir: %v", err) + } + defer os.RemoveAll(tmpDir) + + cfg := &config.Config{ + Agents: config.AgentsConfig{ + Defaults: config.AgentDefaults{ + Workspace: tmpDir, + Model: "test-model", + MaxTokens: 1234, + Temperature: 0, + MaxToolIterations: 5, + }, + }, + } + + provider := &mockProvider{} + agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider) + + if agent.Temperature != 0.7 { + t.Fatalf("Temperature = %f, want %f", agent.Temperature, 0.7) + } +} diff --git a/pkg/agent/loop_test.go b/pkg/agent/loop_test.go index f2257973c..360685eca 100644 --- a/pkg/agent/loop_test.go +++ b/pkg/agent/loop_test.go @@ -14,20 +14,6 @@ import ( "github.com/sipeed/picoclaw/pkg/tools" ) -// mockProvider is a simple mock LLM provider for testing -type mockProvider struct{} - -func (m *mockProvider) Chat(ctx context.Context, messages []providers.Message, tools []providers.ToolDefinition, model string, opts map[string]interface{}) (*providers.LLMResponse, error) { - return &providers.LLMResponse{ - Content: "Mock response", - ToolCalls: []providers.ToolCall{}, - }, nil -} - -func (m *mockProvider) GetDefaultModel() string { - return "mock-model" -} - func TestRecordLastChannel(t *testing.T) { // Create temp workspace tmpDir, err := os.MkdirTemp("", "agent-test-*") @@ -603,7 +589,6 @@ func TestAgentLoop_ContextExhaustionRetry(t *testing.T) { // Call ProcessDirectWithChannel // Note: ProcessDirectWithChannel calls processMessage which will execute runLLMIteration response, err := al.ProcessDirectWithChannel(context.Background(), "Trigger message", sessionKey, "test", "test-chat") - if err != nil { t.Fatalf("Expected success after retry, got error: %v", err) } diff --git a/pkg/agent/mock_provider_test.go b/pkg/agent/mock_provider_test.go new file mode 100644 index 000000000..ccbecbafe --- /dev/null +++ b/pkg/agent/mock_provider_test.go @@ -0,0 +1,20 @@ +package agent + +import ( + "context" + + "github.com/sipeed/picoclaw/pkg/providers" +) + +type mockProvider struct{} + +func (m *mockProvider) Chat(ctx context.Context, messages []providers.Message, tools []providers.ToolDefinition, model string, opts map[string]interface{}) (*providers.LLMResponse, error) { + return &providers.LLMResponse{ + Content: "Mock response", + ToolCalls: []providers.ToolCall{}, + }, nil +} + +func (m *mockProvider) GetDefaultModel() string { + return "mock-model" +} diff --git a/pkg/tools/subagent.go b/pkg/tools/subagent.go index 821b75446..294ba6ea8 100644 --- a/pkg/tools/subagent.go +++ b/pkg/tools/subagent.go @@ -23,19 +23,19 @@ type SubagentTask struct { } type SubagentManager struct { - tasks map[string]*SubagentTask - mu sync.RWMutex - provider providers.LLMProvider - defaultModel string - bus *bus.MessageBus - workspace string - tools *ToolRegistry - maxIterations int - maxTokens int - temperature float64 - hasMaxTokens bool + tasks map[string]*SubagentTask + mu sync.RWMutex + provider providers.LLMProvider + defaultModel string + bus *bus.MessageBus + workspace string + tools *ToolRegistry + maxIterations int + maxTokens int + temperature float64 + hasMaxTokens bool hasTemperature bool - nextID int + nextID int } func NewSubagentManager(provider providers.LLMProvider, defaultModel, workspace string, bus *bus.MessageBus) *SubagentManager { @@ -333,7 +333,6 @@ func (t *SubagentTool) Execute(ctx context.Context, args map[string]interface{}) MaxIterations: maxIter, LLMOptions: llmOptions, }, messages, t.originChannel, t.originChatID) - if err != nil { return ErrorResult(fmt.Sprintf("Subagent execution failed: %v", err)).WithError(err) } diff --git a/pkg/tools/subagent_tool_test.go b/pkg/tools/subagent_tool_test.go index 8a7d22f24..55f3a2528 100644 --- a/pkg/tools/subagent_tool_test.go +++ b/pkg/tools/subagent_tool_test.go @@ -10,9 +10,12 @@ import ( ) // MockLLMProvider is a test implementation of LLMProvider -type MockLLMProvider struct{} +type MockLLMProvider struct{ + lastOptions map[string]interface{} +} func (m *MockLLMProvider) Chat(ctx context.Context, messages []providers.Message, tools []providers.ToolDefinition, model string, options map[string]interface{}) (*providers.LLMResponse, error) { + m.lastOptions = options // Find the last user message to generate a response for i := len(messages) - 1; i >= 0; i-- { if messages[i].Role == "user" { @@ -36,6 +39,32 @@ func (m *MockLLMProvider) GetContextWindow() int { return 4096 } +func TestSubagentManager_SetLLMOptions_AppliesToRunToolLoop(t *testing.T) { + provider := &MockLLMProvider{} + manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil) + manager.SetLLMOptions(2048, 0.6) + tool := NewSubagentTool(manager) + tool.SetContext("cli", "direct") + + ctx := context.Background() + args := map[string]interface{}{"task": "Do something"} + result := tool.Execute(ctx, args) + + if result == nil || result.IsError { + t.Fatalf("Expected successful result, got: %+v", result) + } + + if provider.lastOptions == nil { + t.Fatal("Expected LLM options to be passed, got nil") + } + if provider.lastOptions["max_tokens"] != 2048 { + t.Fatalf("max_tokens = %v, want %d", provider.lastOptions["max_tokens"], 2048) + } + if provider.lastOptions["temperature"] != 0.6 { + t.Fatalf("temperature = %v, want %v", provider.lastOptions["temperature"], 0.6) + } +} + // TestSubagentTool_Name verifies tool name func TestSubagentTool_Name(t *testing.T) { provider := &MockLLMProvider{} diff --git a/pkg/tools/toolloop.go b/pkg/tools/toolloop.go index afc1150c4..e893217d3 100644 --- a/pkg/tools/toolloop.go +++ b/pkg/tools/toolloop.go @@ -57,8 +57,6 @@ func RunToolLoop(ctx context.Context, config ToolLoopConfig, messages []provider if llmOpts == nil { llmOpts = map[string]any{} } - - // 3. Call LLM response, err := config.Provider.Chat(ctx, messages, providerToolDefs, config.Model, llmOpts) if err != nil {