feat: add default temperature handling and update related tests
This commit is contained in:
parent
989b989fb0
commit
493bb52fff
7 changed files with 95 additions and 32 deletions
|
|
@ -83,6 +83,11 @@ func NewAgentInstance(
|
||||||
maxTokens = 8192
|
maxTokens = 8192
|
||||||
}
|
}
|
||||||
|
|
||||||
|
temperature := defaults.Temperature
|
||||||
|
if temperature == 0 {
|
||||||
|
temperature = 0.7
|
||||||
|
}
|
||||||
|
|
||||||
// Resolve fallback candidates
|
// Resolve fallback candidates
|
||||||
modelCfg := providers.ModelConfig{
|
modelCfg := providers.ModelConfig{
|
||||||
Primary: model,
|
Primary: model,
|
||||||
|
|
@ -98,7 +103,7 @@ func NewAgentInstance(
|
||||||
Workspace: workspace,
|
Workspace: workspace,
|
||||||
MaxIterations: maxIter,
|
MaxIterations: maxIter,
|
||||||
MaxTokens: maxTokens,
|
MaxTokens: maxTokens,
|
||||||
Temperature: defaults.Temperature,
|
Temperature: temperature,
|
||||||
ContextWindow: maxTokens,
|
ContextWindow: maxTokens,
|
||||||
Provider: provider,
|
Provider: provider,
|
||||||
Sessions: sessionsManager,
|
Sessions: sessionsManager,
|
||||||
|
|
|
||||||
|
|
@ -36,3 +36,30 @@ func TestNewAgentInstance_UsesDefaultsTemperatureAndMaxTokens(t *testing.T) {
|
||||||
t.Fatalf("Temperature = %f, want %f", agent.Temperature, 1.0)
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -14,20 +14,6 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/tools"
|
"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) {
|
func TestRecordLastChannel(t *testing.T) {
|
||||||
// Create temp workspace
|
// Create temp workspace
|
||||||
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||||
|
|
@ -603,7 +589,6 @@ func TestAgentLoop_ContextExhaustionRetry(t *testing.T) {
|
||||||
// Call ProcessDirectWithChannel
|
// Call ProcessDirectWithChannel
|
||||||
// Note: ProcessDirectWithChannel calls processMessage which will execute runLLMIteration
|
// Note: ProcessDirectWithChannel calls processMessage which will execute runLLMIteration
|
||||||
response, err := al.ProcessDirectWithChannel(context.Background(), "Trigger message", sessionKey, "test", "test-chat")
|
response, err := al.ProcessDirectWithChannel(context.Background(), "Trigger message", sessionKey, "test", "test-chat")
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Expected success after retry, got error: %v", err)
|
t.Fatalf("Expected success after retry, got error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
20
pkg/agent/mock_provider_test.go
Normal file
20
pkg/agent/mock_provider_test.go
Normal file
|
|
@ -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"
|
||||||
|
}
|
||||||
|
|
@ -23,19 +23,19 @@ type SubagentTask struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type SubagentManager struct {
|
type SubagentManager struct {
|
||||||
tasks map[string]*SubagentTask
|
tasks map[string]*SubagentTask
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
provider providers.LLMProvider
|
provider providers.LLMProvider
|
||||||
defaultModel string
|
defaultModel string
|
||||||
bus *bus.MessageBus
|
bus *bus.MessageBus
|
||||||
workspace string
|
workspace string
|
||||||
tools *ToolRegistry
|
tools *ToolRegistry
|
||||||
maxIterations int
|
maxIterations int
|
||||||
maxTokens int
|
maxTokens int
|
||||||
temperature float64
|
temperature float64
|
||||||
hasMaxTokens bool
|
hasMaxTokens bool
|
||||||
hasTemperature bool
|
hasTemperature bool
|
||||||
nextID int
|
nextID int
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewSubagentManager(provider providers.LLMProvider, defaultModel, workspace string, bus *bus.MessageBus) *SubagentManager {
|
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,
|
MaxIterations: maxIter,
|
||||||
LLMOptions: llmOptions,
|
LLMOptions: llmOptions,
|
||||||
}, messages, t.originChannel, t.originChatID)
|
}, messages, t.originChannel, t.originChatID)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ErrorResult(fmt.Sprintf("Subagent execution failed: %v", err)).WithError(err)
|
return ErrorResult(fmt.Sprintf("Subagent execution failed: %v", err)).WithError(err)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -10,9 +10,12 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
// MockLLMProvider is a test implementation of LLMProvider
|
// 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) {
|
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
|
// Find the last user message to generate a response
|
||||||
for i := len(messages) - 1; i >= 0; i-- {
|
for i := len(messages) - 1; i >= 0; i-- {
|
||||||
if messages[i].Role == "user" {
|
if messages[i].Role == "user" {
|
||||||
|
|
@ -36,6 +39,32 @@ func (m *MockLLMProvider) GetContextWindow() int {
|
||||||
return 4096
|
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
|
// TestSubagentTool_Name verifies tool name
|
||||||
func TestSubagentTool_Name(t *testing.T) {
|
func TestSubagentTool_Name(t *testing.T) {
|
||||||
provider := &MockLLMProvider{}
|
provider := &MockLLMProvider{}
|
||||||
|
|
|
||||||
|
|
@ -57,8 +57,6 @@ func RunToolLoop(ctx context.Context, config ToolLoopConfig, messages []provider
|
||||||
if llmOpts == nil {
|
if llmOpts == nil {
|
||||||
llmOpts = map[string]any{}
|
llmOpts = map[string]any{}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
// 3. Call LLM
|
// 3. Call LLM
|
||||||
response, err := config.Provider.Chat(ctx, messages, providerToolDefs, config.Model, llmOpts)
|
response, err := config.Provider.Chat(ctx, messages, providerToolDefs, config.Model, llmOpts)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue