Merge remote-tracking branch 'origin/main' into refactor/cmd
This commit is contained in:
commit
693bdf5e3c
20 changed files with 462 additions and 187 deletions
|
|
@ -66,7 +66,6 @@ linters:
|
||||||
- testifylint
|
- testifylint
|
||||||
- thelper
|
- thelper
|
||||||
- unparam
|
- unparam
|
||||||
- unused
|
|
||||||
- usestdlibvars
|
- usestdlibvars
|
||||||
- usetesting
|
- usetesting
|
||||||
- wastedassign
|
- wastedassign
|
||||||
|
|
@ -152,6 +151,9 @@ linters:
|
||||||
- gocognit
|
- gocognit
|
||||||
- gocyclo
|
- gocyclo
|
||||||
path: _test\.go$
|
path: _test\.go$
|
||||||
|
- linters:
|
||||||
|
- nolintlint
|
||||||
|
path: 'pkg/tools/(i2c\.go|spi\.go)$'
|
||||||
|
|
||||||
issues:
|
issues:
|
||||||
max-issues-per-linter: 0
|
max-issues-per-linter: 0
|
||||||
|
|
|
||||||
8
Makefile
8
Makefile
|
|
@ -15,7 +15,7 @@ INTERNAL=github.com/sipeed/picoclaw/cmd/picoclaw/internal
|
||||||
LDFLAGS=-ldflags "-X $(INTERNAL).version=$(VERSION) -X $(INTERNAL).gitCommit=$(GIT_COMMIT) -X $(INTERNAL).buildTime=$(BUILD_TIME) -X $(INTERNAL).goVersion=$(GO_VERSION) -s -w"
|
LDFLAGS=-ldflags "-X $(INTERNAL).version=$(VERSION) -X $(INTERNAL).gitCommit=$(GIT_COMMIT) -X $(INTERNAL).buildTime=$(BUILD_TIME) -X $(INTERNAL).goVersion=$(GO_VERSION) -s -w"
|
||||||
|
|
||||||
# Go variables
|
# Go variables
|
||||||
GO?=go
|
GO?=CGO_ENABLED=0 go
|
||||||
GOFLAGS?=-v -tags stdjson
|
GOFLAGS?=-v -tags stdjson
|
||||||
|
|
||||||
# Golangci-lint
|
# Golangci-lint
|
||||||
|
|
@ -145,6 +145,10 @@ fmt:
|
||||||
lint:
|
lint:
|
||||||
@$(GOLANGCI_LINT) run
|
@$(GOLANGCI_LINT) run
|
||||||
|
|
||||||
|
## fix: Fix linting issues
|
||||||
|
fix:
|
||||||
|
@$(GOLANGCI_LINT) run --fix
|
||||||
|
|
||||||
## deps: Download dependencies
|
## deps: Download dependencies
|
||||||
deps:
|
deps:
|
||||||
@$(GO) mod download
|
@$(GO) mod download
|
||||||
|
|
@ -170,7 +174,7 @@ help:
|
||||||
@echo " make [target]"
|
@echo " make [target]"
|
||||||
@echo ""
|
@echo ""
|
||||||
@echo "Targets:"
|
@echo "Targets:"
|
||||||
@grep -E '^## ' $(MAKEFILE_LIST) | sed 's/## / /'
|
@grep -E '^## ' $(MAKEFILE_LIST) | sort | awk -F': ' '{printf " %-16s %s\n", substr($$1, 4), $$2}'
|
||||||
@echo ""
|
@echo ""
|
||||||
@echo "Examples:"
|
@echo "Examples:"
|
||||||
@echo " make build # Build for current platform"
|
@echo " make build # Build for current platform"
|
||||||
|
|
|
||||||
|
|
@ -203,6 +203,9 @@ func gatewayCmd(debug bool) error {
|
||||||
<-sigChan
|
<-sigChan
|
||||||
|
|
||||||
fmt.Println("\nShutting down...")
|
fmt.Println("\nShutting down...")
|
||||||
|
if cp, ok := provider.(providers.StatefulProvider); ok {
|
||||||
|
cp.Close()
|
||||||
|
}
|
||||||
cancel()
|
cancel()
|
||||||
healthServer.Stop(context.Background())
|
healthServer.Stop(context.Background())
|
||||||
deviceService.Stop()
|
deviceService.Stop()
|
||||||
|
|
|
||||||
|
|
@ -217,7 +217,8 @@
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"api_key": "pplx-xxx",
|
"api_key": "pplx-xxx",
|
||||||
"max_results": 5
|
"max_results": 5
|
||||||
}
|
},
|
||||||
|
"proxy": ""
|
||||||
},
|
},
|
||||||
"cron": {
|
"cron": {
|
||||||
"exec_timeout_minutes": 5
|
"exec_timeout_minutes": 5
|
||||||
|
|
|
||||||
|
|
@ -288,25 +288,6 @@ func (cb *ContextBuilder) AddAssistantMessage(
|
||||||
return messages
|
return messages
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cb *ContextBuilder) loadSkills() string {
|
|
||||||
allSkills := cb.skillsLoader.ListSkills()
|
|
||||||
if len(allSkills) == 0 {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
var skillNames []string
|
|
||||||
for _, s := range allSkills {
|
|
||||||
skillNames = append(skillNames, s.Name)
|
|
||||||
}
|
|
||||||
|
|
||||||
content := cb.skillsLoader.LoadSkillsForContext(skillNames)
|
|
||||||
if content == "" {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
return "# Skill Definitions\n\n" + content
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetSkillsInfo returns information about loaded skills.
|
// GetSkillsInfo returns information about loaded skills.
|
||||||
func (cb *ContextBuilder) GetSkillsInfo() map[string]any {
|
func (cb *ContextBuilder) GetSkillsInfo() map[string]any {
|
||||||
allSkills := cb.skillsLoader.ListSkills()
|
allSkills := cb.skillsLoader.ListSkills()
|
||||||
|
|
|
||||||
|
|
@ -106,10 +106,11 @@ func registerSharedTools(
|
||||||
PerplexityAPIKey: cfg.Tools.Web.Perplexity.APIKey,
|
PerplexityAPIKey: cfg.Tools.Web.Perplexity.APIKey,
|
||||||
PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults,
|
PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults,
|
||||||
PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled,
|
PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled,
|
||||||
|
Proxy: cfg.Tools.Web.Proxy,
|
||||||
}); searchTool != nil {
|
}); searchTool != nil {
|
||||||
agent.Tools.Register(searchTool)
|
agent.Tools.Register(searchTool)
|
||||||
}
|
}
|
||||||
agent.Tools.Register(tools.NewWebFetchTool(50000))
|
agent.Tools.Register(tools.NewWebFetchToolWithProxy(50000, cfg.Tools.Web.Proxy))
|
||||||
|
|
||||||
// Hardware tools (I2C, SPI) - Linux only, returns error on other platforms
|
// Hardware tools (I2C, SPI) - Linux only, returns error on other platforms
|
||||||
agent.Tools.Register(tools.NewI2CTool())
|
agent.Tools.Register(tools.NewI2CTool())
|
||||||
|
|
|
||||||
|
|
@ -571,61 +571,6 @@ func (c *WeComAppChannel) sendTextMessage(ctx context.Context, accessToken, user
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// sendMarkdownMessage sends a markdown message to a user
|
|
||||||
func (c *WeComAppChannel) sendMarkdownMessage(ctx context.Context, accessToken, userID, content string) error {
|
|
||||||
apiURL := fmt.Sprintf("%s/cgi-bin/message/send?access_token=%s", wecomAPIBase, accessToken)
|
|
||||||
|
|
||||||
msg := WeComMarkdownMessage{
|
|
||||||
ToUser: userID,
|
|
||||||
MsgType: "markdown",
|
|
||||||
AgentID: c.config.AgentID,
|
|
||||||
}
|
|
||||||
msg.Markdown.Content = content
|
|
||||||
|
|
||||||
jsonData, err := json.Marshal(msg)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to marshal message: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Use configurable timeout (default 5 seconds)
|
|
||||||
timeout := c.config.ReplyTimeout
|
|
||||||
if timeout <= 0 {
|
|
||||||
timeout = 5
|
|
||||||
}
|
|
||||||
|
|
||||||
reqCtx, cancel := context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(reqCtx, http.MethodPost, apiURL, bytes.NewBuffer(jsonData))
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to create request: %w", err)
|
|
||||||
}
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
|
||||||
|
|
||||||
client := &http.Client{Timeout: time.Duration(timeout) * time.Second}
|
|
||||||
resp, err := client.Do(req)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to send message: %w", err)
|
|
||||||
}
|
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
body, err := io.ReadAll(resp.Body)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("failed to read response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var sendResp WeComSendMessageResponse
|
|
||||||
if err := json.Unmarshal(body, &sendResp); err != nil {
|
|
||||||
return fmt.Errorf("failed to parse response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if sendResp.ErrCode != 0 {
|
|
||||||
return fmt.Errorf("API error: %s (code: %d)", sendResp.ErrMsg, sendResp.ErrCode)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleHealth handles health check requests
|
// handleHealth handles health check requests
|
||||||
func (c *WeComAppChannel) handleHealth(w http.ResponseWriter, r *http.Request) {
|
func (c *WeComAppChannel) handleHealth(w http.ResponseWriter, r *http.Request) {
|
||||||
status := map[string]any{
|
status := map[string]any{
|
||||||
|
|
|
||||||
|
|
@ -453,6 +453,9 @@ type WebToolsConfig struct {
|
||||||
Tavily TavilyConfig `json:"tavily"`
|
Tavily TavilyConfig `json:"tavily"`
|
||||||
DuckDuckGo DuckDuckGoConfig `json:"duckduckgo"`
|
DuckDuckGo DuckDuckGoConfig `json:"duckduckgo"`
|
||||||
Perplexity PerplexityConfig `json:"perplexity"`
|
Perplexity PerplexityConfig `json:"perplexity"`
|
||||||
|
// Proxy is an optional proxy URL for web tools (http/https/socks5/socks5h).
|
||||||
|
// For authenticated proxies, prefer HTTP_PROXY/HTTPS_PROXY env vars instead of embedding credentials in config.
|
||||||
|
Proxy string `json:"proxy,omitempty" env:"PICOCLAW_TOOLS_WEB_PROXY"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type CronToolsConfig struct {
|
type CronToolsConfig struct {
|
||||||
|
|
@ -509,6 +512,20 @@ func LoadConfig(path string) (*Config, error) {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Pre-scan the JSON to check how many model_list entries the user provided.
|
||||||
|
// Go's JSON decoder reuses existing slice backing-array elements rather than
|
||||||
|
// zero-initializing them, so fields absent from the user's JSON (e.g. api_base)
|
||||||
|
// would silently inherit values from the DefaultConfig template at the same
|
||||||
|
// index position. We only reset cfg.ModelList when the user actually provides
|
||||||
|
// entries; when count is 0 we keep DefaultConfig's built-in list as fallback.
|
||||||
|
var tmp Config
|
||||||
|
if err := json.Unmarshal(data, &tmp); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if len(tmp.ModelList) > 0 {
|
||||||
|
cfg.ModelList = nil
|
||||||
|
}
|
||||||
|
|
||||||
if err := json.Unmarshal(data, cfg); err != nil {
|
if err := json.Unmarshal(data, cfg); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -392,3 +392,24 @@ func TestLoadConfig_OpenAIWebSearchCanBeDisabled(t *testing.T) {
|
||||||
t.Fatal("OpenAI codex web search should be false when disabled in config file")
|
t.Fatal("OpenAI codex web search should be false when disabled in config file")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLoadConfig_WebToolsProxy(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
configPath := filepath.Join(tmpDir, "config.json")
|
||||||
|
configJSON := `{
|
||||||
|
"agents": {"defaults":{"workspace":"./workspace","model":"gpt4","max_tokens":8192,"max_tool_iterations":20}},
|
||||||
|
"model_list": [{"model_name":"gpt4","model":"openai/gpt-5.2","api_key":"x"}],
|
||||||
|
"tools": {"web":{"proxy":"http://127.0.0.1:7890"}}
|
||||||
|
}`
|
||||||
|
if err := os.WriteFile(configPath, []byte(configJSON), 0o600); err != nil {
|
||||||
|
t.Fatalf("os.WriteFile() error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := LoadConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LoadConfig() error: %v", err)
|
||||||
|
}
|
||||||
|
if cfg.Tools.Web.Proxy != "http://127.0.0.1:7890" {
|
||||||
|
t.Fatalf("Tools.Web.Proxy = %q, want %q", cfg.Tools.Web.Proxy, "http://127.0.0.1:7890")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -35,9 +35,8 @@ var usbClassToCapability = map[string]string{
|
||||||
}
|
}
|
||||||
|
|
||||||
type USBMonitor struct {
|
type USBMonitor struct {
|
||||||
cmd *exec.Cmd
|
cmd *exec.Cmd
|
||||||
cancel context.CancelFunc
|
mu sync.Mutex
|
||||||
mu sync.Mutex
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewUSBMonitor() *USBMonitor {
|
func NewUSBMonitor() *USBMonitor {
|
||||||
|
|
|
||||||
|
|
@ -404,64 +404,6 @@ type antigravityJSONResponse struct {
|
||||||
} `json:"usageMetadata"`
|
} `json:"usageMetadata"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *AntigravityProvider) parseJSONResponse(body []byte) (*LLMResponse, error) {
|
|
||||||
var resp antigravityJSONResponse
|
|
||||||
if err := json.Unmarshal(body, &resp); err != nil {
|
|
||||||
return nil, fmt.Errorf("parsing antigravity response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(resp.Candidates) == 0 {
|
|
||||||
return nil, fmt.Errorf("antigravity: no candidates in response")
|
|
||||||
}
|
|
||||||
|
|
||||||
candidate := resp.Candidates[0]
|
|
||||||
var contentParts []string
|
|
||||||
var toolCalls []ToolCall
|
|
||||||
|
|
||||||
for _, part := range candidate.Content.Parts {
|
|
||||||
if part.Text != "" {
|
|
||||||
contentParts = append(contentParts, part.Text)
|
|
||||||
}
|
|
||||||
if part.FunctionCall != nil {
|
|
||||||
argumentsJSON, _ := json.Marshal(part.FunctionCall.Args)
|
|
||||||
toolCalls = append(toolCalls, ToolCall{
|
|
||||||
ID: fmt.Sprintf("call_%s_%d", part.FunctionCall.Name, time.Now().UnixNano()),
|
|
||||||
Name: part.FunctionCall.Name,
|
|
||||||
Arguments: part.FunctionCall.Args,
|
|
||||||
Function: &FunctionCall{
|
|
||||||
Name: part.FunctionCall.Name,
|
|
||||||
Arguments: string(argumentsJSON),
|
|
||||||
ThoughtSignature: extractPartThoughtSignature(part.ThoughtSignature, part.ThoughtSignatureSnake),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
finishReason := "stop"
|
|
||||||
if len(toolCalls) > 0 {
|
|
||||||
finishReason = "tool_calls"
|
|
||||||
}
|
|
||||||
if candidate.FinishReason == "MAX_TOKENS" {
|
|
||||||
finishReason = "length"
|
|
||||||
}
|
|
||||||
|
|
||||||
var usage *UsageInfo
|
|
||||||
if resp.UsageMetadata.TotalTokenCount > 0 {
|
|
||||||
usage = &UsageInfo{
|
|
||||||
PromptTokens: resp.UsageMetadata.PromptTokenCount,
|
|
||||||
CompletionTokens: resp.UsageMetadata.CandidatesTokenCount,
|
|
||||||
TotalTokens: resp.UsageMetadata.TotalTokenCount,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return &LLMResponse{
|
|
||||||
Content: strings.Join(contentParts, ""),
|
|
||||||
ToolCalls: toolCalls,
|
|
||||||
FinishReason: finishReason,
|
|
||||||
Usage: usage,
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (p *AntigravityProvider) parseSSEResponse(body string) (*LLMResponse, error) {
|
func (p *AntigravityProvider) parseSSEResponse(body string) (*LLMResponse, error) {
|
||||||
var contentParts []string
|
var contentParts []string
|
||||||
var toolCalls []ToolCall
|
var toolCalls []ToolCall
|
||||||
|
|
|
||||||
|
|
@ -17,12 +17,6 @@ func successRun(content string) func(ctx context.Context, provider, model string
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func failRun(err error) func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
|
||||||
return func(ctx context.Context, provider, model string) (*LLMResponse, error) {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFallback_SingleCandidate_Success(t *testing.T) {
|
func TestFallback_SingleCandidate_Success(t *testing.T) {
|
||||||
ct := NewCooldownTracker()
|
ct := NewCooldownTracker()
|
||||||
fc := NewFallbackChain(ct)
|
fc := NewFallbackChain(ct)
|
||||||
|
|
|
||||||
|
|
@ -4,60 +4,84 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
|
||||||
copilot "github.com/github/copilot-sdk/go"
|
copilot "github.com/github/copilot-sdk/go"
|
||||||
)
|
)
|
||||||
|
|
||||||
type GitHubCopilotProvider struct {
|
type GitHubCopilotProvider struct {
|
||||||
uri string
|
uri string
|
||||||
connectMode string // `stdio` or `grpc``
|
connectMode string // "stdio" or "grpc"
|
||||||
|
|
||||||
|
client *copilot.Client
|
||||||
session *copilot.Session
|
session *copilot.Session
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewGitHubCopilotProvider(uri string, connectMode string, model string) (*GitHubCopilotProvider, error) {
|
func NewGitHubCopilotProvider(uri string, connectMode string, model string) (*GitHubCopilotProvider, error) {
|
||||||
var session *copilot.Session
|
|
||||||
if connectMode == "" {
|
if connectMode == "" {
|
||||||
connectMode = "grpc"
|
connectMode = "grpc"
|
||||||
}
|
}
|
||||||
switch connectMode {
|
|
||||||
|
|
||||||
|
switch connectMode {
|
||||||
case "stdio":
|
case "stdio":
|
||||||
// todo
|
// TODO:
|
||||||
|
return nil, fmt.Errorf("stdio mode not implemented")
|
||||||
case "grpc":
|
case "grpc":
|
||||||
client := copilot.NewClient(&copilot.ClientOptions{
|
client := copilot.NewClient(&copilot.ClientOptions{
|
||||||
CLIUrl: uri,
|
CLIUrl: uri,
|
||||||
})
|
})
|
||||||
if err := client.Start(context.Background()); err != nil {
|
if err := client.Start(context.Background()); err != nil {
|
||||||
return nil, fmt.Errorf(
|
return nil, fmt.Errorf(
|
||||||
"Can't connect to Github Copilot, https://github.com/github/copilot-sdk/blob/main/docs/getting-started.md#connecting-to-an-external-cli-server for details",
|
"can't connect to Github Copilot: %w; `https://github.com/github/copilot-sdk/blob/main/docs/getting-started.md#connecting-to-an-external-cli-server` for details",
|
||||||
|
err,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
defer client.Stop()
|
|
||||||
session, _ = client.CreateSession(context.Background(), &copilot.SessionConfig{
|
session, err := client.CreateSession(context.Background(), &copilot.SessionConfig{
|
||||||
Model: model,
|
Model: model,
|
||||||
Hooks: &copilot.SessionHooks{},
|
Hooks: &copilot.SessionHooks{},
|
||||||
})
|
})
|
||||||
|
if err != nil {
|
||||||
|
|
||||||
|
client.Stop()
|
||||||
|
return nil, fmt.Errorf("create session failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &GitHubCopilotProvider{
|
||||||
|
uri: uri,
|
||||||
|
connectMode: connectMode,
|
||||||
|
client: client,
|
||||||
|
session: session,
|
||||||
|
}, nil
|
||||||
|
default:
|
||||||
|
return nil, fmt.Errorf("unknown connect mode: %s", connectMode)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *GitHubCopilotProvider) Close() {
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
if p.client != nil {
|
||||||
|
p.client.Stop()
|
||||||
|
p.client = nil
|
||||||
|
p.session = nil
|
||||||
}
|
}
|
||||||
|
|
||||||
return &GitHubCopilotProvider{
|
|
||||||
uri: uri,
|
|
||||||
connectMode: connectMode,
|
|
||||||
session: session,
|
|
||||||
}, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Chat sends a chat request to GitHub Copilot
|
|
||||||
func (p *GitHubCopilotProvider) Chat(
|
func (p *GitHubCopilotProvider) Chat(
|
||||||
ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]any,
|
ctx context.Context,
|
||||||
|
messages []Message,
|
||||||
|
tools []ToolDefinition,
|
||||||
|
model string,
|
||||||
|
options map[string]any,
|
||||||
) (*LLMResponse, error) {
|
) (*LLMResponse, error) {
|
||||||
type tempMessage struct {
|
type tempMessage struct {
|
||||||
Role string `json:"role"`
|
Role string `json:"role"`
|
||||||
Content string `json:"content"`
|
Content string `json:"content"`
|
||||||
}
|
}
|
||||||
out := make([]tempMessage, 0, len(messages))
|
out := make([]tempMessage, 0, len(messages))
|
||||||
|
|
||||||
for _, msg := range messages {
|
for _, msg := range messages {
|
||||||
out = append(out, tempMessage{
|
out = append(out, tempMessage{
|
||||||
Role: msg.Role,
|
Role: msg.Role,
|
||||||
|
|
@ -65,12 +89,30 @@ func (p *GitHubCopilotProvider) Chat(
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
fullcontent, _ := json.Marshal(out)
|
fullcontent, err := json.Marshal(out)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("marshal messages: %w", err)
|
||||||
|
}
|
||||||
|
p.mu.Lock()
|
||||||
|
session := p.session
|
||||||
|
p.mu.Unlock()
|
||||||
|
|
||||||
content, _ := p.session.Send(ctx, copilot.MessageOptions{
|
if session == nil {
|
||||||
|
return nil, fmt.Errorf("provider closed")
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := session.SendAndWait(ctx, copilot.MessageOptions{
|
||||||
Prompt: string(fullcontent),
|
Prompt: string(fullcontent),
|
||||||
})
|
})
|
||||||
|
|
||||||
|
if resp == nil {
|
||||||
|
return nil, fmt.Errorf("empty response from copilot")
|
||||||
|
}
|
||||||
|
if resp.Data.Content == nil {
|
||||||
|
return nil, fmt.Errorf("no content in copilot response")
|
||||||
|
}
|
||||||
|
content := *resp.Data.Content
|
||||||
|
|
||||||
return &LLMResponse{
|
return &LLMResponse{
|
||||||
FinishReason: "stop",
|
FinishReason: "stop",
|
||||||
Content: content,
|
Content: content,
|
||||||
|
|
|
||||||
|
|
@ -30,6 +30,11 @@ type LLMProvider interface {
|
||||||
GetDefaultModel() string
|
GetDefaultModel() string
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type StatefulProvider interface {
|
||||||
|
LLMProvider
|
||||||
|
Close()
|
||||||
|
}
|
||||||
|
|
||||||
// FailoverReason classifies why an LLM request failed for fallback decisions.
|
// FailoverReason classifies why an LLM request failed for fallback decisions.
|
||||||
type FailoverReason string
|
type FailoverReason string
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -117,13 +117,19 @@ func (t *I2CTool) detect() *ToolResult {
|
||||||
return SilentResult(fmt.Sprintf("Found %d I2C bus(es):\n%s", len(buses), string(result)))
|
return SilentResult(fmt.Sprintf("Found %d I2C bus(es):\n%s", len(buses), string(result)))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Helper functions for I2C operations (used by platform-specific implementations)
|
||||||
|
|
||||||
// isValidBusID checks that a bus identifier is a simple number (prevents path injection)
|
// isValidBusID checks that a bus identifier is a simple number (prevents path injection)
|
||||||
|
//
|
||||||
|
//nolint:unused // Used by i2c_linux.go
|
||||||
func isValidBusID(id string) bool {
|
func isValidBusID(id string) bool {
|
||||||
matched, _ := regexp.MatchString(`^\d+$`, id)
|
matched, _ := regexp.MatchString(`^\d+$`, id)
|
||||||
return matched
|
return matched
|
||||||
}
|
}
|
||||||
|
|
||||||
// parseI2CAddress extracts and validates an I2C address from args
|
// parseI2CAddress extracts and validates an I2C address from args
|
||||||
|
//
|
||||||
|
//nolint:unused // Used by i2c_linux.go
|
||||||
func parseI2CAddress(args map[string]any) (int, *ToolResult) {
|
func parseI2CAddress(args map[string]any) (int, *ToolResult) {
|
||||||
addrFloat, ok := args["address"].(float64)
|
addrFloat, ok := args["address"].(float64)
|
||||||
if !ok {
|
if !ok {
|
||||||
|
|
@ -137,6 +143,8 @@ func parseI2CAddress(args map[string]any) (int, *ToolResult) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// parseI2CBus extracts and validates an I2C bus from args
|
// parseI2CBus extracts and validates an I2C bus from args
|
||||||
|
//
|
||||||
|
//nolint:unused // Used by i2c_linux.go
|
||||||
func parseI2CBus(args map[string]any) (string, *ToolResult) {
|
func parseI2CBus(args map[string]any) (string, *ToolResult) {
|
||||||
bus, ok := args["bus"].(string)
|
bus, ok := args["bus"].(string)
|
||||||
if !ok || bus == "" {
|
if !ok || bus == "" {
|
||||||
|
|
|
||||||
|
|
@ -3,6 +3,7 @@ package tools
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"strings"
|
||||||
)
|
)
|
||||||
|
|
||||||
type SpawnTool struct {
|
type SpawnTool struct {
|
||||||
|
|
@ -66,8 +67,8 @@ func (t *SpawnTool) SetAllowlistChecker(check func(targetAgentID string) bool) {
|
||||||
|
|
||||||
func (t *SpawnTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
func (t *SpawnTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
task, ok := args["task"].(string)
|
task, ok := args["task"].(string)
|
||||||
if !ok {
|
if !ok || strings.TrimSpace(task) == "" {
|
||||||
return ErrorResult("task is required")
|
return ErrorResult("task is required and must be a non-empty string")
|
||||||
}
|
}
|
||||||
|
|
||||||
label, _ := args["label"].(string)
|
label, _ := args["label"].(string)
|
||||||
|
|
|
||||||
79
pkg/tools/spawn_test.go
Normal file
79
pkg/tools/spawn_test.go
Normal file
|
|
@ -0,0 +1,79 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSpawnTool_Execute_EmptyTask(t *testing.T) {
|
||||||
|
provider := &MockLLMProvider{}
|
||||||
|
manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil)
|
||||||
|
tool := NewSpawnTool(manager)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
args map[string]any
|
||||||
|
}{
|
||||||
|
{"empty string", map[string]any{"task": ""}},
|
||||||
|
{"whitespace only", map[string]any{"task": " "}},
|
||||||
|
{"tabs and newlines", map[string]any{"task": "\t\n "}},
|
||||||
|
{"missing task key", map[string]any{"label": "test"}},
|
||||||
|
{"wrong type", map[string]any{"task": 123}},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
result := tool.Execute(ctx, tt.args)
|
||||||
|
if result == nil {
|
||||||
|
t.Fatal("Result should not be nil")
|
||||||
|
}
|
||||||
|
if !result.IsError {
|
||||||
|
t.Error("Expected error for invalid task parameter")
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForLLM, "task is required") {
|
||||||
|
t.Errorf("Error message should mention 'task is required', got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSpawnTool_Execute_ValidTask(t *testing.T) {
|
||||||
|
provider := &MockLLMProvider{}
|
||||||
|
manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil)
|
||||||
|
tool := NewSpawnTool(manager)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
args := map[string]any{
|
||||||
|
"task": "Write a haiku about coding",
|
||||||
|
"label": "haiku-task",
|
||||||
|
}
|
||||||
|
|
||||||
|
result := tool.Execute(ctx, args)
|
||||||
|
if result == nil {
|
||||||
|
t.Fatal("Result should not be nil")
|
||||||
|
}
|
||||||
|
if result.IsError {
|
||||||
|
t.Errorf("Expected success for valid task, got error: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
if !result.Async {
|
||||||
|
t.Error("SpawnTool should return async result")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSpawnTool_Execute_NilManager(t *testing.T) {
|
||||||
|
tool := NewSpawnTool(nil)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
args := map[string]any{"task": "test task"}
|
||||||
|
|
||||||
|
result := tool.Execute(ctx, args)
|
||||||
|
if !result.IsError {
|
||||||
|
t.Error("Expected error for nil manager")
|
||||||
|
}
|
||||||
|
if !strings.Contains(result.ForLLM, "Subagent manager not configured") {
|
||||||
|
t.Errorf("Error message should mention manager not configured, got: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -119,7 +119,11 @@ func (t *SPITool) list() *ToolResult {
|
||||||
return SilentResult(fmt.Sprintf("Found %d SPI device(s):\n%s", len(devices), string(result)))
|
return SilentResult(fmt.Sprintf("Found %d SPI device(s):\n%s", len(devices), string(result)))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Helper function for SPI operations (used by platform-specific implementations)
|
||||||
|
|
||||||
// parseSPIArgs extracts and validates common SPI parameters
|
// parseSPIArgs extracts and validates common SPI parameters
|
||||||
|
//
|
||||||
|
//nolint:unused // Used by spi_linux.go
|
||||||
func parseSPIArgs(args map[string]any) (device string, speed uint32, mode uint8, bits uint8, errMsg string) {
|
func parseSPIArgs(args map[string]any) (device string, speed uint32, mode uint8, bits uint8, errMsg string) {
|
||||||
dev, ok := args["device"].(string)
|
dev, ok := args["device"].(string)
|
||||||
if !ok || dev == "" {
|
if !ok || dev == "" {
|
||||||
|
|
|
||||||
101
pkg/tools/web.go
101
pkg/tools/web.go
|
|
@ -17,12 +17,50 @@ const (
|
||||||
userAgent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36"
|
userAgent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// createHTTPClient creates an HTTP client with optional proxy support
|
||||||
|
func createHTTPClient(proxyURL string, timeout time.Duration) (*http.Client, error) {
|
||||||
|
client := &http.Client{
|
||||||
|
Timeout: timeout,
|
||||||
|
Transport: &http.Transport{
|
||||||
|
MaxIdleConns: 10,
|
||||||
|
IdleConnTimeout: 30 * time.Second,
|
||||||
|
DisableCompression: false,
|
||||||
|
TLSHandshakeTimeout: 15 * time.Second,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
if proxyURL != "" {
|
||||||
|
proxy, err := url.Parse(proxyURL)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid proxy URL: %w", err)
|
||||||
|
}
|
||||||
|
scheme := strings.ToLower(proxy.Scheme)
|
||||||
|
switch scheme {
|
||||||
|
case "http", "https", "socks5", "socks5h":
|
||||||
|
default:
|
||||||
|
return nil, fmt.Errorf(
|
||||||
|
"unsupported proxy scheme %q (supported: http, https, socks5, socks5h)",
|
||||||
|
proxy.Scheme,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if proxy.Host == "" {
|
||||||
|
return nil, fmt.Errorf("invalid proxy URL: missing host")
|
||||||
|
}
|
||||||
|
client.Transport.(*http.Transport).Proxy = http.ProxyURL(proxy)
|
||||||
|
} else {
|
||||||
|
client.Transport.(*http.Transport).Proxy = http.ProxyFromEnvironment
|
||||||
|
}
|
||||||
|
|
||||||
|
return client, nil
|
||||||
|
}
|
||||||
|
|
||||||
type SearchProvider interface {
|
type SearchProvider interface {
|
||||||
Search(ctx context.Context, query string, count int) (string, error)
|
Search(ctx context.Context, query string, count int) (string, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
type BraveSearchProvider struct {
|
type BraveSearchProvider struct {
|
||||||
apiKey string
|
apiKey string
|
||||||
|
proxy string
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *BraveSearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
|
func (p *BraveSearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
|
||||||
|
|
@ -37,7 +75,10 @@ func (p *BraveSearchProvider) Search(ctx context.Context, query string, count in
|
||||||
req.Header.Set("Accept", "application/json")
|
req.Header.Set("Accept", "application/json")
|
||||||
req.Header.Set("X-Subscription-Token", p.apiKey)
|
req.Header.Set("X-Subscription-Token", p.apiKey)
|
||||||
|
|
||||||
client := &http.Client{Timeout: 10 * time.Second}
|
client, err := createHTTPClient(p.proxy, 10*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to create HTTP client: %w", err)
|
||||||
|
}
|
||||||
resp, err := client.Do(req)
|
resp, err := client.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("request failed: %w", err)
|
return "", fmt.Errorf("request failed: %w", err)
|
||||||
|
|
@ -167,7 +208,9 @@ func (p *TavilySearchProvider) Search(ctx context.Context, query string, count i
|
||||||
return strings.Join(lines, "\n"), nil
|
return strings.Join(lines, "\n"), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type DuckDuckGoSearchProvider struct{}
|
type DuckDuckGoSearchProvider struct {
|
||||||
|
proxy string
|
||||||
|
}
|
||||||
|
|
||||||
func (p *DuckDuckGoSearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
|
func (p *DuckDuckGoSearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
|
||||||
searchURL := fmt.Sprintf("https://html.duckduckgo.com/html/?q=%s", url.QueryEscape(query))
|
searchURL := fmt.Sprintf("https://html.duckduckgo.com/html/?q=%s", url.QueryEscape(query))
|
||||||
|
|
@ -179,7 +222,10 @@ func (p *DuckDuckGoSearchProvider) Search(ctx context.Context, query string, cou
|
||||||
|
|
||||||
req.Header.Set("User-Agent", userAgent)
|
req.Header.Set("User-Agent", userAgent)
|
||||||
|
|
||||||
client := &http.Client{Timeout: 10 * time.Second}
|
client, err := createHTTPClient(p.proxy, 10*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to create HTTP client: %w", err)
|
||||||
|
}
|
||||||
resp, err := client.Do(req)
|
resp, err := client.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("request failed: %w", err)
|
return "", fmt.Errorf("request failed: %w", err)
|
||||||
|
|
@ -261,6 +307,7 @@ func stripTags(content string) string {
|
||||||
|
|
||||||
type PerplexitySearchProvider struct {
|
type PerplexitySearchProvider struct {
|
||||||
apiKey string
|
apiKey string
|
||||||
|
proxy string
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
|
func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
|
||||||
|
|
@ -295,7 +342,10 @@ func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, cou
|
||||||
req.Header.Set("Authorization", "Bearer "+p.apiKey)
|
req.Header.Set("Authorization", "Bearer "+p.apiKey)
|
||||||
req.Header.Set("User-Agent", userAgent)
|
req.Header.Set("User-Agent", userAgent)
|
||||||
|
|
||||||
client := &http.Client{Timeout: 30 * time.Second}
|
client, err := createHTTPClient(p.proxy, 30*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to create HTTP client: %w", err)
|
||||||
|
}
|
||||||
resp, err := client.Do(req)
|
resp, err := client.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("request failed: %w", err)
|
return "", fmt.Errorf("request failed: %w", err)
|
||||||
|
|
@ -348,6 +398,7 @@ type WebSearchToolOptions struct {
|
||||||
PerplexityAPIKey string
|
PerplexityAPIKey string
|
||||||
PerplexityMaxResults int
|
PerplexityMaxResults int
|
||||||
PerplexityEnabled bool
|
PerplexityEnabled bool
|
||||||
|
Proxy string
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewWebSearchTool(opts WebSearchToolOptions) *WebSearchTool {
|
func NewWebSearchTool(opts WebSearchToolOptions) *WebSearchTool {
|
||||||
|
|
@ -356,12 +407,12 @@ func NewWebSearchTool(opts WebSearchToolOptions) *WebSearchTool {
|
||||||
|
|
||||||
// Priority: Perplexity > Brave > Tavily > DuckDuckGo
|
// Priority: Perplexity > Brave > Tavily > DuckDuckGo
|
||||||
if opts.PerplexityEnabled && opts.PerplexityAPIKey != "" {
|
if opts.PerplexityEnabled && opts.PerplexityAPIKey != "" {
|
||||||
provider = &PerplexitySearchProvider{apiKey: opts.PerplexityAPIKey}
|
provider = &PerplexitySearchProvider{apiKey: opts.PerplexityAPIKey, proxy: opts.Proxy}
|
||||||
if opts.PerplexityMaxResults > 0 {
|
if opts.PerplexityMaxResults > 0 {
|
||||||
maxResults = opts.PerplexityMaxResults
|
maxResults = opts.PerplexityMaxResults
|
||||||
}
|
}
|
||||||
} else if opts.BraveEnabled && opts.BraveAPIKey != "" {
|
} else if opts.BraveEnabled && opts.BraveAPIKey != "" {
|
||||||
provider = &BraveSearchProvider{apiKey: opts.BraveAPIKey}
|
provider = &BraveSearchProvider{apiKey: opts.BraveAPIKey, proxy: opts.Proxy}
|
||||||
if opts.BraveMaxResults > 0 {
|
if opts.BraveMaxResults > 0 {
|
||||||
maxResults = opts.BraveMaxResults
|
maxResults = opts.BraveMaxResults
|
||||||
}
|
}
|
||||||
|
|
@ -374,7 +425,7 @@ func NewWebSearchTool(opts WebSearchToolOptions) *WebSearchTool {
|
||||||
maxResults = opts.TavilyMaxResults
|
maxResults = opts.TavilyMaxResults
|
||||||
}
|
}
|
||||||
} else if opts.DuckDuckGoEnabled {
|
} else if opts.DuckDuckGoEnabled {
|
||||||
provider = &DuckDuckGoSearchProvider{}
|
provider = &DuckDuckGoSearchProvider{proxy: opts.Proxy}
|
||||||
if opts.DuckDuckGoMaxResults > 0 {
|
if opts.DuckDuckGoMaxResults > 0 {
|
||||||
maxResults = opts.DuckDuckGoMaxResults
|
maxResults = opts.DuckDuckGoMaxResults
|
||||||
}
|
}
|
||||||
|
|
@ -441,6 +492,7 @@ func (t *WebSearchTool) Execute(ctx context.Context, args map[string]any) *ToolR
|
||||||
|
|
||||||
type WebFetchTool struct {
|
type WebFetchTool struct {
|
||||||
maxChars int
|
maxChars int
|
||||||
|
proxy string
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewWebFetchTool(maxChars int) *WebFetchTool {
|
func NewWebFetchTool(maxChars int) *WebFetchTool {
|
||||||
|
|
@ -452,6 +504,16 @@ func NewWebFetchTool(maxChars int) *WebFetchTool {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func NewWebFetchToolWithProxy(maxChars int, proxy string) *WebFetchTool {
|
||||||
|
if maxChars <= 0 {
|
||||||
|
maxChars = 50000
|
||||||
|
}
|
||||||
|
return &WebFetchTool{
|
||||||
|
maxChars: maxChars,
|
||||||
|
proxy: proxy,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (t *WebFetchTool) Name() string {
|
func (t *WebFetchTool) Name() string {
|
||||||
return "web_fetch"
|
return "web_fetch"
|
||||||
}
|
}
|
||||||
|
|
@ -511,20 +573,17 @@ func (t *WebFetchTool) Execute(ctx context.Context, args map[string]any) *ToolRe
|
||||||
|
|
||||||
req.Header.Set("User-Agent", userAgent)
|
req.Header.Set("User-Agent", userAgent)
|
||||||
|
|
||||||
client := &http.Client{
|
client, err := createHTTPClient(t.proxy, 60*time.Second)
|
||||||
Timeout: 60 * time.Second,
|
if err != nil {
|
||||||
Transport: &http.Transport{
|
return ErrorResult(fmt.Sprintf("failed to create HTTP client: %v", err))
|
||||||
MaxIdleConns: 10,
|
}
|
||||||
IdleConnTimeout: 30 * time.Second,
|
|
||||||
DisableCompression: false,
|
// Configure redirect handling
|
||||||
TLSHandshakeTimeout: 15 * time.Second,
|
client.CheckRedirect = func(req *http.Request, via []*http.Request) error {
|
||||||
},
|
if len(via) >= 5 {
|
||||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
return fmt.Errorf("stopped after 5 redirects")
|
||||||
if len(via) >= 5 {
|
}
|
||||||
return fmt.Errorf("stopped after 5 redirects")
|
return nil
|
||||||
}
|
|
||||||
return nil
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
resp, err := client.Do(req)
|
resp, err := client.Do(req)
|
||||||
|
|
|
||||||
|
|
@ -7,6 +7,7 @@ import (
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
// TestWebTool_WebFetch_Success verifies successful URL fetching
|
// TestWebTool_WebFetch_Success verifies successful URL fetching
|
||||||
|
|
@ -334,6 +335,172 @@ func TestWebTool_WebFetch_MissingDomain(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestCreateHTTPClient_ProxyConfigured(t *testing.T) {
|
||||||
|
client, err := createHTTPClient("http://127.0.0.1:7890", 12*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("createHTTPClient() error: %v", err)
|
||||||
|
}
|
||||||
|
if client.Timeout != 12*time.Second {
|
||||||
|
t.Fatalf("client.Timeout = %v, want %v", client.Timeout, 12*time.Second)
|
||||||
|
}
|
||||||
|
|
||||||
|
tr, ok := client.Transport.(*http.Transport)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("client.Transport type = %T, want *http.Transport", client.Transport)
|
||||||
|
}
|
||||||
|
if tr.Proxy == nil {
|
||||||
|
t.Fatal("transport.Proxy is nil, want non-nil")
|
||||||
|
}
|
||||||
|
|
||||||
|
req, err := http.NewRequest("GET", "https://example.com", nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("http.NewRequest() error: %v", err)
|
||||||
|
}
|
||||||
|
proxyURL, err := tr.Proxy(req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("transport.Proxy(req) error: %v", err)
|
||||||
|
}
|
||||||
|
if proxyURL == nil || proxyURL.String() != "http://127.0.0.1:7890" {
|
||||||
|
t.Fatalf("proxy URL = %v, want %q", proxyURL, "http://127.0.0.1:7890")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateHTTPClient_InvalidProxy(t *testing.T) {
|
||||||
|
_, err := createHTTPClient("://bad-proxy", 10*time.Second)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("createHTTPClient() expected error for invalid proxy URL, got nil")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateHTTPClient_Socks5ProxyConfigured(t *testing.T) {
|
||||||
|
client, err := createHTTPClient("socks5://127.0.0.1:1080", 8*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("createHTTPClient() error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
tr, ok := client.Transport.(*http.Transport)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("client.Transport type = %T, want *http.Transport", client.Transport)
|
||||||
|
}
|
||||||
|
req, err := http.NewRequest("GET", "https://example.com", nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("http.NewRequest() error: %v", err)
|
||||||
|
}
|
||||||
|
proxyURL, err := tr.Proxy(req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("transport.Proxy(req) error: %v", err)
|
||||||
|
}
|
||||||
|
if proxyURL == nil || proxyURL.String() != "socks5://127.0.0.1:1080" {
|
||||||
|
t.Fatalf("proxy URL = %v, want %q", proxyURL, "socks5://127.0.0.1:1080")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateHTTPClient_UnsupportedProxyScheme(t *testing.T) {
|
||||||
|
_, err := createHTTPClient("ftp://127.0.0.1:21", 10*time.Second)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("createHTTPClient() expected error for unsupported scheme, got nil")
|
||||||
|
}
|
||||||
|
if !strings.Contains(err.Error(), "unsupported proxy scheme") {
|
||||||
|
t.Fatalf("error = %q, want to contain %q", err.Error(), "unsupported proxy scheme")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateHTTPClient_ProxyFromEnvironmentWhenConfigEmpty(t *testing.T) {
|
||||||
|
t.Setenv("HTTP_PROXY", "http://127.0.0.1:8888")
|
||||||
|
t.Setenv("http_proxy", "http://127.0.0.1:8888")
|
||||||
|
t.Setenv("HTTPS_PROXY", "http://127.0.0.1:8888")
|
||||||
|
t.Setenv("https_proxy", "http://127.0.0.1:8888")
|
||||||
|
t.Setenv("ALL_PROXY", "")
|
||||||
|
t.Setenv("all_proxy", "")
|
||||||
|
t.Setenv("NO_PROXY", "")
|
||||||
|
t.Setenv("no_proxy", "")
|
||||||
|
|
||||||
|
client, err := createHTTPClient("", 10*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("createHTTPClient() error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
tr, ok := client.Transport.(*http.Transport)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("client.Transport type = %T, want *http.Transport", client.Transport)
|
||||||
|
}
|
||||||
|
if tr.Proxy == nil {
|
||||||
|
t.Fatal("transport.Proxy is nil, want proxy function from environment")
|
||||||
|
}
|
||||||
|
|
||||||
|
req, err := http.NewRequest("GET", "https://example.com", nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("http.NewRequest() error: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := tr.Proxy(req); err != nil {
|
||||||
|
t.Fatalf("transport.Proxy(req) error: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewWebFetchToolWithProxy(t *testing.T) {
|
||||||
|
tool := NewWebFetchToolWithProxy(1024, "http://127.0.0.1:7890")
|
||||||
|
if tool.maxChars != 1024 {
|
||||||
|
t.Fatalf("maxChars = %d, want %d", tool.maxChars, 1024)
|
||||||
|
}
|
||||||
|
if tool.proxy != "http://127.0.0.1:7890" {
|
||||||
|
t.Fatalf("proxy = %q, want %q", tool.proxy, "http://127.0.0.1:7890")
|
||||||
|
}
|
||||||
|
|
||||||
|
tool = NewWebFetchToolWithProxy(0, "http://127.0.0.1:7890")
|
||||||
|
if tool.maxChars != 50000 {
|
||||||
|
t.Fatalf("default maxChars = %d, want %d", tool.maxChars, 50000)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewWebSearchTool_PropagatesProxy(t *testing.T) {
|
||||||
|
t.Run("perplexity", func(t *testing.T) {
|
||||||
|
tool := NewWebSearchTool(WebSearchToolOptions{
|
||||||
|
PerplexityEnabled: true,
|
||||||
|
PerplexityAPIKey: "k",
|
||||||
|
PerplexityMaxResults: 3,
|
||||||
|
Proxy: "http://127.0.0.1:7890",
|
||||||
|
})
|
||||||
|
p, ok := tool.provider.(*PerplexitySearchProvider)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("provider type = %T, want *PerplexitySearchProvider", tool.provider)
|
||||||
|
}
|
||||||
|
if p.proxy != "http://127.0.0.1:7890" {
|
||||||
|
t.Fatalf("provider proxy = %q, want %q", p.proxy, "http://127.0.0.1:7890")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("brave", func(t *testing.T) {
|
||||||
|
tool := NewWebSearchTool(WebSearchToolOptions{
|
||||||
|
BraveEnabled: true,
|
||||||
|
BraveAPIKey: "k",
|
||||||
|
BraveMaxResults: 3,
|
||||||
|
Proxy: "http://127.0.0.1:7890",
|
||||||
|
})
|
||||||
|
p, ok := tool.provider.(*BraveSearchProvider)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("provider type = %T, want *BraveSearchProvider", tool.provider)
|
||||||
|
}
|
||||||
|
if p.proxy != "http://127.0.0.1:7890" {
|
||||||
|
t.Fatalf("provider proxy = %q, want %q", p.proxy, "http://127.0.0.1:7890")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("duckduckgo", func(t *testing.T) {
|
||||||
|
tool := NewWebSearchTool(WebSearchToolOptions{
|
||||||
|
DuckDuckGoEnabled: true,
|
||||||
|
DuckDuckGoMaxResults: 3,
|
||||||
|
Proxy: "http://127.0.0.1:7890",
|
||||||
|
})
|
||||||
|
p, ok := tool.provider.(*DuckDuckGoSearchProvider)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("provider type = %T, want *DuckDuckGoSearchProvider", tool.provider)
|
||||||
|
}
|
||||||
|
if p.proxy != "http://127.0.0.1:7890" {
|
||||||
|
t.Fatalf("provider proxy = %q, want %q", p.proxy, "http://127.0.0.1:7890")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
// TestWebTool_TavilySearch_Success verifies successful Tavily search
|
// TestWebTool_TavilySearch_Success verifies successful Tavily search
|
||||||
func TestWebTool_TavilySearch_Success(t *testing.T) {
|
func TestWebTool_TavilySearch_Success(t *testing.T) {
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue