feat(mcp): tool search tools
This commit is contained in:
parent
4768edc67b
commit
770f141cbf
12 changed files with 1262 additions and 47 deletions
|
|
@ -176,8 +176,13 @@
|
||||||
"nickserv_password": "",
|
"nickserv_password": "",
|
||||||
"sasl_user": "",
|
"sasl_user": "",
|
||||||
"sasl_password": "",
|
"sasl_password": "",
|
||||||
"channels": ["#mychannel"],
|
"channels": [
|
||||||
"request_caps": ["server-time", "message-tags"],
|
"#mychannel"
|
||||||
|
],
|
||||||
|
"request_caps": [
|
||||||
|
"server-time",
|
||||||
|
"message-tags"
|
||||||
|
],
|
||||||
"allow_from": [],
|
"allow_from": [],
|
||||||
"group_trigger": {
|
"group_trigger": {
|
||||||
"mention_only": true
|
"mention_only": true
|
||||||
|
|
@ -298,6 +303,13 @@
|
||||||
},
|
},
|
||||||
"mcp": {
|
"mcp": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
|
"discovery": {
|
||||||
|
"enabled": false,
|
||||||
|
"ttl": 5,
|
||||||
|
"max_search_results": 5,
|
||||||
|
"use_bm25": true,
|
||||||
|
"use_regex": false
|
||||||
|
},
|
||||||
"servers": {
|
"servers": {
|
||||||
"context7": {
|
"context7": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
|
|
|
||||||
|
|
@ -7,11 +7,21 @@ PicoClaw's tools configuration is located in the `tools` field of `config.json`.
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"tools": {
|
"tools": {
|
||||||
"web": { ... },
|
"web": {
|
||||||
"mcp": { ... },
|
...
|
||||||
"exec": { ... },
|
},
|
||||||
"cron": { ... },
|
"mcp": {
|
||||||
"skills": { ... }
|
...
|
||||||
|
},
|
||||||
|
"exec": {
|
||||||
|
...
|
||||||
|
},
|
||||||
|
"cron": {
|
||||||
|
...
|
||||||
|
},
|
||||||
|
"skills": {
|
||||||
|
...
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
@ -23,7 +33,7 @@ Web tools are used for web search and fetching.
|
||||||
### Brave
|
### Brave
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| ------------- | ------ | ------- | ------------------------- |
|
|---------------|--------|---------|---------------------------|
|
||||||
| `enabled` | bool | false | Enable Brave search |
|
| `enabled` | bool | false | Enable Brave search |
|
||||||
| `api_key` | string | - | Brave Search API key |
|
| `api_key` | string | - | Brave Search API key |
|
||||||
| `max_results` | int | 5 | Maximum number of results |
|
| `max_results` | int | 5 | Maximum number of results |
|
||||||
|
|
@ -31,14 +41,14 @@ Web tools are used for web search and fetching.
|
||||||
### DuckDuckGo
|
### DuckDuckGo
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| ------------- | ---- | ------- | ------------------------- |
|
|---------------|------|---------|---------------------------|
|
||||||
| `enabled` | bool | true | Enable DuckDuckGo search |
|
| `enabled` | bool | true | Enable DuckDuckGo search |
|
||||||
| `max_results` | int | 5 | Maximum number of results |
|
| `max_results` | int | 5 | Maximum number of results |
|
||||||
|
|
||||||
### Perplexity
|
### Perplexity
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| ------------- | ------ | ------- | ------------------------- |
|
|---------------|--------|---------|---------------------------|
|
||||||
| `enabled` | bool | false | Enable Perplexity search |
|
| `enabled` | bool | false | Enable Perplexity search |
|
||||||
| `api_key` | string | - | Perplexity API key |
|
| `api_key` | string | - | Perplexity API key |
|
||||||
| `max_results` | int | 5 | Maximum number of results |
|
| `max_results` | int | 5 | Maximum number of results |
|
||||||
|
|
@ -48,7 +58,7 @@ Web tools are used for web search and fetching.
|
||||||
The exec tool is used to execute shell commands.
|
The exec tool is used to execute shell commands.
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| ---------------------- | ----- | ------- | ------------------------------------------ |
|
|------------------------|-------|---------|--------------------------------------------|
|
||||||
| `enable_deny_patterns` | bool | true | Enable default dangerous command blocking |
|
| `enable_deny_patterns` | bool | true | Enable default dangerous command blocking |
|
||||||
| `custom_deny_patterns` | array | [] | Custom deny patterns (regular expressions) |
|
| `custom_deny_patterns` | array | [] | Custom deny patterns (regular expressions) |
|
||||||
|
|
||||||
|
|
@ -81,7 +91,10 @@ By default, PicoClaw blocks the following dangerous commands:
|
||||||
"tools": {
|
"tools": {
|
||||||
"exec": {
|
"exec": {
|
||||||
"enable_deny_patterns": true,
|
"enable_deny_patterns": true,
|
||||||
"custom_deny_patterns": ["\\brm\\s+-r\\b", "\\bkillall\\s+python"]
|
"custom_deny_patterns": [
|
||||||
|
"\\brm\\s+-r\\b",
|
||||||
|
"\\bkillall\\s+python"
|
||||||
|
]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -92,24 +105,47 @@ By default, PicoClaw blocks the following dangerous commands:
|
||||||
The cron tool is used for scheduling periodic tasks.
|
The cron tool is used for scheduling periodic tasks.
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| ---------------------- | ---- | ------- | ---------------------------------------------- |
|
|------------------------|------|---------|------------------------------------------------|
|
||||||
| `exec_timeout_minutes` | int | 5 | Execution timeout in minutes, 0 means no limit |
|
| `exec_timeout_minutes` | int | 5 | Execution timeout in minutes, 0 means no limit |
|
||||||
|
|
||||||
## MCP Tool
|
## MCP Tool
|
||||||
|
|
||||||
The MCP tool enables integration with external Model Context Protocol servers.
|
The MCP tool enables integration with external Model Context Protocol servers.
|
||||||
|
|
||||||
|
### Tool Discovery (Lazy Loading)
|
||||||
|
|
||||||
|
When connecting to multiple MCP servers, exposing hundreds of tools simultaneously can exhaust the LLM's context window
|
||||||
|
and increase API costs. The **Discovery** feature solves this by keeping MCP tools *hidden* by default.
|
||||||
|
|
||||||
|
Instead of loading all tools, the LLM is provided with a lightweight search tool (using BM25 keyword matching or Regex).
|
||||||
|
When the LLM needs a specific capability, it searches the hidden library. Matching tools are then temporarily "unlocked"
|
||||||
|
and injected into the context for a configured number of turns (`ttl`).
|
||||||
|
|
||||||
### Global Config
|
### Global Config
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| --------- | ------ | ------- | ----------------------------------- |
|
|-------------|--------|---------|----------------------------------------------|
|
||||||
| `enabled` | bool | false | Enable MCP integration globally |
|
| `enabled` | bool | false | Enable MCP integration globally |
|
||||||
|
| `discovery` | object | `{}` | Configuration for Tool Discovery (see below) |
|
||||||
| `servers` | object | `{}` | Map of server name to server config |
|
| `servers` | object | `{}` | Map of server name to server config |
|
||||||
|
|
||||||
|
### Discovery Config (`discovery`)
|
||||||
|
|
||||||
|
| Config | Type | Default | Description |
|
||||||
|
|----------------------|------|---------|-----------------------------------------------------------------------------------------------------------------------------------|
|
||||||
|
| `enabled` | bool | false | If true, MCP tools are hidden and loaded on-demand via search. If false, all tools are loaded |
|
||||||
|
| `ttl` | int | 5 | Number of conversational turns a discovered tool remains unlocked |
|
||||||
|
| `max_search_results` | int | 5 | Maximum number of tools returned per search query |
|
||||||
|
| `use_bm25` | bool | true | Enable the natural language/keyword search tool (`tool_search_tool_bm25`). **Warning**: consumes more resources than regex search |
|
||||||
|
| `use_regex` | bool | false | Enable the regex pattern search tool (`tool_search_tool_regex`) |
|
||||||
|
|
||||||
|
> **Note:** If `discovery.enabled` is `true`, you MUST enable at least one search engine (`use_bm25` or `use_regex`),
|
||||||
|
> otherwise the application will fail to start.
|
||||||
|
|
||||||
### Per-Server Config
|
### Per-Server Config
|
||||||
|
|
||||||
| Config | Type | Required | Description |
|
| Config | Type | Required | Description |
|
||||||
| ---------- | ------ | -------- | ------------------------------------------ |
|
|------------|--------|----------|--------------------------------------------|
|
||||||
| `enabled` | bool | yes | Enable this MCP server |
|
| `enabled` | bool | yes | Enable this MCP server |
|
||||||
| `type` | string | no | Transport type: `stdio`, `sse`, `http` |
|
| `type` | string | no | Transport type: `stdio`, `sse`, `http` |
|
||||||
| `command` | string | stdio | Executable command for stdio transport |
|
| `command` | string | stdio | Executable command for stdio transport |
|
||||||
|
|
@ -140,7 +176,11 @@ The MCP tool enables integration with external Model Context Protocol servers.
|
||||||
"filesystem": {
|
"filesystem": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"command": "npx",
|
"command": "npx",
|
||||||
"args": ["-y", "@modelcontextprotocol/server-filesystem", "/tmp"]
|
"args": [
|
||||||
|
"-y",
|
||||||
|
"@modelcontextprotocol/server-filesystem",
|
||||||
|
"/tmp"
|
||||||
|
]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -170,6 +210,62 @@ The MCP tool enables integration with external Model Context Protocol servers.
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
#### 3) Massive MCP setup with Tool Discovery enabled
|
||||||
|
|
||||||
|
*In this example, the LLM will only see the `tool_search_tool_bm25`. It will search and unlock Github or Postgres tools
|
||||||
|
dynamically only when requested by the user.*
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"tools": {
|
||||||
|
"mcp": {
|
||||||
|
"enabled": true,
|
||||||
|
"discovery": {
|
||||||
|
"enabled": true,
|
||||||
|
"ttl": 5,
|
||||||
|
"max_search_results": 5,
|
||||||
|
"use_bm25": true,
|
||||||
|
"use_regex": false
|
||||||
|
},
|
||||||
|
"servers": {
|
||||||
|
"github": {
|
||||||
|
"enabled": true,
|
||||||
|
"command": "npx",
|
||||||
|
"args": [
|
||||||
|
"-y",
|
||||||
|
"@modelcontextprotocol/server-github"
|
||||||
|
],
|
||||||
|
"env": {
|
||||||
|
"GITHUB_PERSONAL_ACCESS_TOKEN": "YOUR_GITHUB_TOKEN"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"postgres": {
|
||||||
|
"enabled": true,
|
||||||
|
"command": "npx",
|
||||||
|
"args": [
|
||||||
|
"-y",
|
||||||
|
"@modelcontextprotocol/server-postgres",
|
||||||
|
"postgresql://user:password@localhost/dbname"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"slack": {
|
||||||
|
"enabled": true,
|
||||||
|
"command": "npx",
|
||||||
|
"args": [
|
||||||
|
"-y",
|
||||||
|
"@modelcontextprotocol/server-slack"
|
||||||
|
],
|
||||||
|
"env": {
|
||||||
|
"SLACK_BOT_TOKEN": "YOUR_SLACK_BOT_TOKEN",
|
||||||
|
"SLACK_TEAM_ID": "YOUR_SLACK_TEAM_ID"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
## Skills Tool
|
## Skills Tool
|
||||||
|
|
||||||
The skills tool configures skill discovery and installation via registries like ClawHub.
|
The skills tool configures skill discovery and installation via registries like ClawHub.
|
||||||
|
|
@ -177,7 +273,7 @@ The skills tool configures skill discovery and installation via registries like
|
||||||
### Registries
|
### Registries
|
||||||
|
|
||||||
| Config | Type | Default | Description |
|
| Config | Type | Default | Description |
|
||||||
| ---------------------------------- | ------ | -------------------- | ----------------------- |
|
|------------------------------------|--------|----------------------|----------------------------------------------|
|
||||||
| `registries.clawhub.enabled` | bool | true | Enable ClawHub registry |
|
| `registries.clawhub.enabled` | bool | true | Enable ClawHub registry |
|
||||||
| `registries.clawhub.base_url` | string | `https://clawhub.ai` | ClawHub base URL |
|
| `registries.clawhub.base_url` | string | `https://clawhub.ai` | ClawHub base URL |
|
||||||
| `registries.clawhub.auth_token` | string | `""` | Optional Bearer token for higher rate limits |
|
| `registries.clawhub.auth_token` | string | `""` | Optional Bearer token for higher rate limits |
|
||||||
|
|
@ -217,4 +313,5 @@ For example:
|
||||||
- `PICOCLAW_TOOLS_CRON_EXEC_TIMEOUT_MINUTES=10`
|
- `PICOCLAW_TOOLS_CRON_EXEC_TIMEOUT_MINUTES=10`
|
||||||
- `PICOCLAW_TOOLS_MCP_ENABLED=true`
|
- `PICOCLAW_TOOLS_MCP_ENABLED=true`
|
||||||
|
|
||||||
Note: Nested map-style config (for example `tools.mcp.servers.<name>.*`) is configured in `config.json` rather than environment variables.
|
Note: Nested map-style config (for example `tools.mcp.servers.<name>.*`) is configured in `config.json` rather than
|
||||||
|
environment variables.
|
||||||
|
|
|
||||||
|
|
@ -21,6 +21,7 @@ type ContextBuilder struct {
|
||||||
workspace string
|
workspace string
|
||||||
skillsLoader *skills.SkillsLoader
|
skillsLoader *skills.SkillsLoader
|
||||||
memory *MemoryStore
|
memory *MemoryStore
|
||||||
|
toolDiscovery bool
|
||||||
|
|
||||||
// Cache for system prompt to avoid rebuilding on every call.
|
// Cache for system prompt to avoid rebuilding on every call.
|
||||||
// This fixes issue #607: repeated reprocessing of the entire context.
|
// This fixes issue #607: repeated reprocessing of the entire context.
|
||||||
|
|
@ -41,6 +42,11 @@ type ContextBuilder struct {
|
||||||
skillFilesAtCache map[string]time.Time
|
skillFilesAtCache map[string]time.Time
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (cb *ContextBuilder) WithToolDiscovery(enabled bool) *ContextBuilder {
|
||||||
|
cb.toolDiscovery = enabled
|
||||||
|
return cb
|
||||||
|
}
|
||||||
|
|
||||||
func getGlobalConfigDir() string {
|
func getGlobalConfigDir() string {
|
||||||
if home := os.Getenv("PICOCLAW_HOME"); home != "" {
|
if home := os.Getenv("PICOCLAW_HOME"); home != "" {
|
||||||
return home
|
return home
|
||||||
|
|
@ -71,6 +77,7 @@ func NewContextBuilder(workspace string) *ContextBuilder {
|
||||||
|
|
||||||
func (cb *ContextBuilder) getIdentity() string {
|
func (cb *ContextBuilder) getIdentity() string {
|
||||||
workspacePath, _ := filepath.Abs(filepath.Join(cb.workspace))
|
workspacePath, _ := filepath.Abs(filepath.Join(cb.workspace))
|
||||||
|
toolDiscovery := cb.getDiscoveryRule()
|
||||||
|
|
||||||
return fmt.Sprintf(`# picoclaw 🦞
|
return fmt.Sprintf(`# picoclaw 🦞
|
||||||
|
|
||||||
|
|
@ -90,8 +97,17 @@ Your workspace is at: %s
|
||||||
|
|
||||||
3. **Memory** - When interacting with me if something seems memorable, update %s/memory/MEMORY.md
|
3. **Memory** - When interacting with me if something seems memorable, update %s/memory/MEMORY.md
|
||||||
|
|
||||||
4. **Context summaries** - Conversation summaries provided as context are approximate references only. They may be incomplete or outdated. Always defer to explicit user instructions over summary content.`,
|
4. **Context summaries** - Conversation summaries provided as context are approximate references only. They may be incomplete or outdated. Always defer to explicit user instructions over summary content.
|
||||||
workspacePath, workspacePath, workspacePath, workspacePath, workspacePath)
|
|
||||||
|
%s`,
|
||||||
|
workspacePath, workspacePath, workspacePath, workspacePath, workspacePath, toolDiscovery)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (cb *ContextBuilder) getDiscoveryRule() string {
|
||||||
|
if cb.toolDiscovery {
|
||||||
|
return `5. **Tool Discovery** - Your visible tools are limited to save memory, but a vast hidden library exists. If you lack the right tool for a task, BEFORE giving up, you MUST search using the "tool_search_tool_bm25" or "tool_search_tool_regex" tool. Do not refuse a request unless the search returns nothing. Found tools will temporarily unlock for your next turn.`
|
||||||
|
}
|
||||||
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cb *ContextBuilder) BuildSystemPrompt() string {
|
func (cb *ContextBuilder) BuildSystemPrompt() string {
|
||||||
|
|
|
||||||
|
|
@ -96,7 +96,7 @@ func NewAgentInstance(
|
||||||
sessionsDir := filepath.Join(workspace, "sessions")
|
sessionsDir := filepath.Join(workspace, "sessions")
|
||||||
sessionsManager := session.NewSessionManager(sessionsDir)
|
sessionsManager := session.NewSessionManager(sessionsDir)
|
||||||
|
|
||||||
contextBuilder := NewContextBuilder(workspace)
|
contextBuilder := NewContextBuilder(workspace).WithToolDiscovery(cfg.Tools.MCP.ToolConfig.Discovery.Enabled)
|
||||||
|
|
||||||
agentID := routing.DefaultAgentID
|
agentID := routing.DefaultAgentID
|
||||||
agentName := ""
|
agentName := ""
|
||||||
|
|
|
||||||
|
|
@ -11,6 +11,7 @@ import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"log"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"regexp"
|
"regexp"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
@ -283,7 +284,13 @@ func (al *AgentLoop) Run(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
mcpTool := tools.NewMCPTool(mcpManager, serverName, tool)
|
mcpTool := tools.NewMCPTool(mcpManager, serverName, tool)
|
||||||
|
|
||||||
|
if al.cfg.Tools.MCP.ToolConfig.Enabled {
|
||||||
|
agent.Tools.RegisterHidden(mcpTool)
|
||||||
|
} else {
|
||||||
agent.Tools.Register(mcpTool)
|
agent.Tools.Register(mcpTool)
|
||||||
|
}
|
||||||
|
|
||||||
totalRegistrations++
|
totalRegistrations++
|
||||||
logger.DebugCF("agent", "Registered MCP tool",
|
logger.DebugCF("agent", "Registered MCP tool",
|
||||||
map[string]any{
|
map[string]any{
|
||||||
|
|
@ -302,6 +309,43 @@ func (al *AgentLoop) Run(ctx context.Context) error {
|
||||||
"total_registrations": totalRegistrations,
|
"total_registrations": totalRegistrations,
|
||||||
"agent_count": agentCount,
|
"agent_count": agentCount,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
// Initializes Discovery Tools only if enabled by configuration
|
||||||
|
if al.cfg.Tools.MCP.Enabled && al.cfg.Tools.MCP.Discovery.Enabled {
|
||||||
|
useBM25 := al.cfg.Tools.MCP.Discovery.UseBM25
|
||||||
|
useRegex := al.cfg.Tools.MCP.Discovery.UseRegex
|
||||||
|
|
||||||
|
// Fail fast: If discovery is enabled but no tool is turned on, break the app
|
||||||
|
if !useBM25 && !useRegex {
|
||||||
|
log.Fatalf(
|
||||||
|
"Critical error: tool discovery is enabled but neither 'use_bm25' nor 'use_regex' is set to true in the configuration.",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
ttl := al.cfg.Tools.MCP.Discovery.TTL
|
||||||
|
if ttl <= 0 {
|
||||||
|
ttl = 5 // Default value
|
||||||
|
}
|
||||||
|
|
||||||
|
maxSearchResults := al.cfg.Tools.MCP.Discovery.MaxSearchResults
|
||||||
|
|
||||||
|
for _, agentID := range agentIDs {
|
||||||
|
agent, ok := al.registry.GetAgent(agentID)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if useRegex {
|
||||||
|
agent.Tools.Register(tools.NewRegexSearchTool(agent.Tools, ttl, maxSearchResults))
|
||||||
|
}
|
||||||
|
if useBM25 {
|
||||||
|
agent.Tools.Register(tools.NewBM25SearchTool(agent.Tools, ttl, maxSearchResults))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Always record the fallback
|
||||||
|
agent.Tools.Register(tools.NewCallDiscoveredTool(agent.Tools, ttl))
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -890,6 +934,9 @@ func (al *AgentLoop) runLLMIteration(
|
||||||
"max": agent.MaxIterations,
|
"max": agent.MaxIterations,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
// Scale down the TTL of the discovered tools with each new round of the LLM
|
||||||
|
agent.Tools.TickTTL()
|
||||||
|
|
||||||
// Build tool definitions
|
// Build tool definitions
|
||||||
providerToolDefs := agent.Tools.ToProviderDefs()
|
providerToolDefs := agent.Tools.ToProviderDefs()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -559,8 +559,17 @@ type GatewayConfig struct {
|
||||||
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type ToolDiscoveryConfig struct {
|
||||||
|
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_DISCOVERY_ENABLED"`
|
||||||
|
TTL int `json:"ttl" env:"PICOCLAW_TOOLS_DISCOVERY_TTL"`
|
||||||
|
MaxSearchResults int `json:"max_search_results" env:"PICOCLAW_MAX_SEARCH_RESULTS"`
|
||||||
|
UseBM25 bool `json:"use_bm25" env:"PICOCLAW_TOOLS_DISCOVERY_USE_BM25"`
|
||||||
|
UseRegex bool `json:"use_regex" env:"PICOCLAW_TOOLS_DISCOVERY_USE_REGEX"`
|
||||||
|
}
|
||||||
|
|
||||||
type ToolConfig struct {
|
type ToolConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"ENABLED"`
|
Enabled bool `json:"enabled" env:"ENABLED"`
|
||||||
|
Discovery ToolDiscoveryConfig `json:"discovery"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type BraveConfig struct {
|
type BraveConfig struct {
|
||||||
|
|
|
||||||
|
|
@ -410,6 +410,13 @@ func DefaultConfig() *Config {
|
||||||
MCP: MCPConfig{
|
MCP: MCPConfig{
|
||||||
ToolConfig: ToolConfig{
|
ToolConfig: ToolConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
|
Discovery: ToolDiscoveryConfig{
|
||||||
|
Enabled: false,
|
||||||
|
TTL: 5,
|
||||||
|
MaxSearchResults: 5,
|
||||||
|
UseBM25: true,
|
||||||
|
UseRegex: false,
|
||||||
|
},
|
||||||
},
|
},
|
||||||
Servers: map[string]MCPServerConfig{},
|
Servers: map[string]MCPServerConfig{},
|
||||||
},
|
},
|
||||||
|
|
|
||||||
|
|
@ -11,14 +11,20 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type ToolEntry struct {
|
||||||
|
Tool Tool
|
||||||
|
IsCore bool
|
||||||
|
TTL int
|
||||||
|
}
|
||||||
|
|
||||||
type ToolRegistry struct {
|
type ToolRegistry struct {
|
||||||
tools map[string]Tool
|
tools map[string]*ToolEntry
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewToolRegistry() *ToolRegistry {
|
func NewToolRegistry() *ToolRegistry {
|
||||||
return &ToolRegistry{
|
return &ToolRegistry{
|
||||||
tools: make(map[string]Tool),
|
tools: make(map[string]*ToolEntry),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -30,14 +36,55 @@ func (r *ToolRegistry) Register(tool Tool) {
|
||||||
logger.WarnCF("tools", "Tool registration overwrites existing tool",
|
logger.WarnCF("tools", "Tool registration overwrites existing tool",
|
||||||
map[string]any{"name": name})
|
map[string]any{"name": name})
|
||||||
}
|
}
|
||||||
r.tools[name] = tool
|
r.tools[name] = &ToolEntry{
|
||||||
|
Tool: tool,
|
||||||
|
IsCore: true,
|
||||||
|
TTL: 0, // Core tools do not use TTL
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterHidden saves hidden tools (visible only via TTL)
|
||||||
|
func (r *ToolRegistry) RegisterHidden(tool Tool) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
name := tool.Name()
|
||||||
|
r.tools[name] = &ToolEntry{
|
||||||
|
Tool: tool,
|
||||||
|
IsCore: false,
|
||||||
|
TTL: 0,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// PromoteTool imposta il TTL solo se il tool NON è un core tool
|
||||||
|
func (r *ToolRegistry) PromoteTool(name string, ttl int) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
if entry, exists := r.tools[name]; exists {
|
||||||
|
if !entry.IsCore {
|
||||||
|
entry.TTL = ttl
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TickTTL decreases TTL only for non-core tools
|
||||||
|
func (r *ToolRegistry) TickTTL() {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
for _, entry := range r.tools {
|
||||||
|
if !entry.IsCore && entry.TTL > 0 {
|
||||||
|
entry.TTL--
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *ToolRegistry) Get(name string) (Tool, bool) {
|
func (r *ToolRegistry) Get(name string) (Tool, bool) {
|
||||||
r.mu.RLock()
|
r.mu.RLock()
|
||||||
defer r.mu.RUnlock()
|
defer r.mu.RUnlock()
|
||||||
tool, ok := r.tools[name]
|
entry, ok := r.tools[name]
|
||||||
return tool, ok
|
if !ok {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
return entry.Tool, true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *ToolRegistry) Execute(ctx context.Context, name string, args map[string]any) *ToolResult {
|
func (r *ToolRegistry) Execute(ctx context.Context, name string, args map[string]any) *ToolResult {
|
||||||
|
|
@ -135,7 +182,13 @@ func (r *ToolRegistry) GetDefinitions() []map[string]any {
|
||||||
sorted := r.sortedToolNames()
|
sorted := r.sortedToolNames()
|
||||||
definitions := make([]map[string]any, 0, len(sorted))
|
definitions := make([]map[string]any, 0, len(sorted))
|
||||||
for _, name := range sorted {
|
for _, name := range sorted {
|
||||||
definitions = append(definitions, ToolToSchema(r.tools[name]))
|
entry := r.tools[name]
|
||||||
|
|
||||||
|
if !entry.IsCore && entry.TTL <= 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
definitions = append(definitions, ToolToSchema(r.tools[name].Tool))
|
||||||
}
|
}
|
||||||
return definitions
|
return definitions
|
||||||
}
|
}
|
||||||
|
|
@ -149,8 +202,13 @@ func (r *ToolRegistry) ToProviderDefs() []providers.ToolDefinition {
|
||||||
sorted := r.sortedToolNames()
|
sorted := r.sortedToolNames()
|
||||||
definitions := make([]providers.ToolDefinition, 0, len(sorted))
|
definitions := make([]providers.ToolDefinition, 0, len(sorted))
|
||||||
for _, name := range sorted {
|
for _, name := range sorted {
|
||||||
tool := r.tools[name]
|
entry := r.tools[name]
|
||||||
schema := ToolToSchema(tool)
|
|
||||||
|
if !entry.IsCore && entry.TTL <= 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
schema := ToolToSchema(entry.Tool)
|
||||||
|
|
||||||
// Safely extract nested values with type checks
|
// Safely extract nested values with type checks
|
||||||
fn, ok := schema["function"].(map[string]any)
|
fn, ok := schema["function"].(map[string]any)
|
||||||
|
|
@ -198,8 +256,13 @@ func (r *ToolRegistry) GetSummaries() []string {
|
||||||
sorted := r.sortedToolNames()
|
sorted := r.sortedToolNames()
|
||||||
summaries := make([]string, 0, len(sorted))
|
summaries := make([]string, 0, len(sorted))
|
||||||
for _, name := range sorted {
|
for _, name := range sorted {
|
||||||
tool := r.tools[name]
|
entry := r.tools[name]
|
||||||
summaries = append(summaries, fmt.Sprintf("- `%s` - %s", tool.Name(), tool.Description()))
|
|
||||||
|
if !entry.IsCore && entry.TTL <= 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
summaries = append(summaries, fmt.Sprintf("- `%s` - %s", entry.Tool.Name(), entry.Tool.Description()))
|
||||||
}
|
}
|
||||||
return summaries
|
return summaries
|
||||||
}
|
}
|
||||||
|
|
|
||||||
281
pkg/tools/search_tool.go
Normal file
281
pkg/tools/search_tool.go
Normal file
|
|
@ -0,0 +1,281 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
)
|
||||||
|
|
||||||
|
type RegexSearchTool struct {
|
||||||
|
registry *ToolRegistry
|
||||||
|
ttl int
|
||||||
|
maxSearchResults int
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewRegexSearchTool(r *ToolRegistry, ttl int, maxSearchResults int) *RegexSearchTool {
|
||||||
|
return &RegexSearchTool{registry: r, ttl: ttl, maxSearchResults: maxSearchResults}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *RegexSearchTool) Name() string {
|
||||||
|
return "tool_search_tool_regex"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *RegexSearchTool) Description() string {
|
||||||
|
return "Search available hidden tools on-demand using a regex pattern. Returns JSON schemas of discovered tools."
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *RegexSearchTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"pattern": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Regex pattern to match tool name or description",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"pattern"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *RegexSearchTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
|
pattern, ok := args["pattern"].(string)
|
||||||
|
if !ok || strings.TrimSpace(pattern) == "" {
|
||||||
|
// An empty string regex (?i) will match every hidden tool,
|
||||||
|
// dumping massive payloads into the context and burning tokens.
|
||||||
|
return ErrorResult("Missing or invalid 'pattern' argument. Must be a non-empty string.")
|
||||||
|
}
|
||||||
|
|
||||||
|
res, err := t.registry.SearchRegex(pattern, t.maxSearchResults)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("Invalid regex pattern syntax: %v. Please fix your regex and try again.", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
return formatDiscoveryResponse(t.registry, res, t.ttl)
|
||||||
|
}
|
||||||
|
|
||||||
|
type BM25SearchTool struct {
|
||||||
|
registry *ToolRegistry
|
||||||
|
ttl int
|
||||||
|
maxSearchResults int
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewBM25SearchTool(r *ToolRegistry, ttl int, maxSearchResults int) *BM25SearchTool {
|
||||||
|
return &BM25SearchTool{registry: r, ttl: ttl, maxSearchResults: maxSearchResults}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *BM25SearchTool) Name() string {
|
||||||
|
return "tool_search_tool_bm25"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *BM25SearchTool) Description() string {
|
||||||
|
return "Search available hidden tools on-demand using natural language query describing the action you need to perform. Returns JSON schemas of discovered tools."
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *BM25SearchTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"query": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Search query",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"query"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *BM25SearchTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
|
query, ok := args["query"].(string)
|
||||||
|
if !ok || strings.TrimSpace(query) == "" {
|
||||||
|
// An empty string query will match every hidden tool,
|
||||||
|
// dumping massive payloads into the context and burning tokens.
|
||||||
|
return ErrorResult("Missing or invalid 'query' argument. Must be a non-empty string.")
|
||||||
|
}
|
||||||
|
|
||||||
|
return formatDiscoveryResponse(t.registry, t.registry.SearchBM25(query, t.maxSearchResults), t.ttl)
|
||||||
|
}
|
||||||
|
|
||||||
|
type CallDiscoveredTool struct {
|
||||||
|
registry *ToolRegistry
|
||||||
|
ttl int
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewCallDiscoveredTool(r *ToolRegistry, ttl int) *CallDiscoveredTool {
|
||||||
|
return &CallDiscoveredTool{registry: r, ttl: ttl}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *CallDiscoveredTool) Name() string {
|
||||||
|
return "call_discovered_tool"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *CallDiscoveredTool) Description() string {
|
||||||
|
return "Fallback tool. Execute a tool found via search by passing its required arguments as a JSON object."
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *CallDiscoveredTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"tool_name": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
},
|
||||||
|
"arguments": map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"description": "Arguments to pass to the tool",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"tool_name"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *CallDiscoveredTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
|
name, ok := args["tool_name"].(string)
|
||||||
|
if !ok || name == "" {
|
||||||
|
return ErrorResult("Missing or invalid 'tool_name' argument")
|
||||||
|
}
|
||||||
|
|
||||||
|
parsedArgs := make(map[string]any)
|
||||||
|
|
||||||
|
// Check whether the key "arguments" exists in the payload
|
||||||
|
if argVal, exists := args["arguments"]; exists && argVal != nil {
|
||||||
|
// If it exists, we try to map cast it
|
||||||
|
var valid bool
|
||||||
|
parsedArgs, valid = argVal.(map[string]any)
|
||||||
|
if !valid {
|
||||||
|
// The LLM has passed something, but it is NOT a JSON object!
|
||||||
|
// We have to tell him clearly to get him to correct.
|
||||||
|
return ErrorResult(fmt.Sprintf(
|
||||||
|
"Invalid 'arguments' format for tool '%s'. Expected a JSON object, but got %T. Please fix and try again.",
|
||||||
|
name,
|
||||||
|
argVal,
|
||||||
|
))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Renew the TTL to keep it visible if it is actively used
|
||||||
|
t.registry.PromoteTool(name, t.ttl)
|
||||||
|
|
||||||
|
return t.registry.Execute(ctx, name, parsedArgs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ToolSearchResult represents the result returned to the LLM.
|
||||||
|
type ToolSearchResult struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Description string `json:"description"`
|
||||||
|
Parameters map[string]any `json:"parameters"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *ToolRegistry) SearchRegex(pattern string, maxSearchResults int) ([]ToolSearchResult, error) {
|
||||||
|
regex, err := regexp.Compile("(?i)" + pattern)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to compile regex pattern %q: %w", pattern, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
|
||||||
|
var results []ToolSearchResult
|
||||||
|
|
||||||
|
for name, entry := range r.tools {
|
||||||
|
// Search only among the hidden tools (Core tools are already visible)
|
||||||
|
if !entry.IsCore {
|
||||||
|
// Directly call interface methods! No reflection/unmarshalling needed.
|
||||||
|
desc := entry.Tool.Description()
|
||||||
|
|
||||||
|
if regex.MatchString(name) || regex.MatchString(desc) {
|
||||||
|
results = append(results, ToolSearchResult{
|
||||||
|
Name: name,
|
||||||
|
Description: desc,
|
||||||
|
Parameters: entry.Tool.Parameters(),
|
||||||
|
})
|
||||||
|
if len(results) >= maxSearchResults {
|
||||||
|
break // Stop searching once we hit the max! Saves CPU.
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return results, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func formatDiscoveryResponse(registry *ToolRegistry, results []ToolSearchResult, ttl int) *ToolResult {
|
||||||
|
if len(results) == 0 {
|
||||||
|
return SilentResult("No tools found matching the query.")
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, r := range results {
|
||||||
|
registry.PromoteTool(r.Name, ttl)
|
||||||
|
}
|
||||||
|
|
||||||
|
b, err := json.Marshal(results)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult("Failed to format search results: " + err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
msg := fmt.Sprintf(
|
||||||
|
"Found %d tools:\n%s\n\nSUCCESS: These tools have been temporarily UNLOCKED as native tools! In your next response, you can call them directly just like any normal tool, without needing 'call_discovered_tool'.",
|
||||||
|
len(results),
|
||||||
|
string(b),
|
||||||
|
)
|
||||||
|
|
||||||
|
return SilentResult(msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Lightweight internal type
|
||||||
|
type searchDoc struct {
|
||||||
|
Name string
|
||||||
|
Description string
|
||||||
|
Tool Tool // Hold the interface reference
|
||||||
|
}
|
||||||
|
|
||||||
|
// SearchBM25 ranks hidden tools against query using BM25 via utils.BM25Engine.
|
||||||
|
// The corpus snapshot is built under the registry read-lock, then released
|
||||||
|
// before scoring so the lock is not held during CPU-intensive work.
|
||||||
|
func (r *ToolRegistry) SearchBM25(query string, maxSearchResults int) []ToolSearchResult {
|
||||||
|
// We copy only the lightweight searchDoc values (name, description,
|
||||||
|
// Tool interface reference). This keeps the lock window short and avoids
|
||||||
|
// holding it during BM25 indexing and scoring.
|
||||||
|
r.mu.RLock()
|
||||||
|
snapshot := make([]searchDoc, 0, len(r.tools))
|
||||||
|
for name, entry := range r.tools {
|
||||||
|
if !entry.IsCore {
|
||||||
|
snapshot = append(snapshot, searchDoc{
|
||||||
|
Name: name,
|
||||||
|
Description: entry.Tool.Description(),
|
||||||
|
Tool: entry.Tool,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
r.mu.RUnlock()
|
||||||
|
|
||||||
|
if len(snapshot) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delegate scoring to the generic BM25 engine
|
||||||
|
engine := utils.NewBM25Engine(
|
||||||
|
snapshot,
|
||||||
|
func(doc searchDoc) string {
|
||||||
|
return doc.Name + " " + doc.Description
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
ranked := engine.Search(query, maxSearchResults)
|
||||||
|
if len(ranked) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
out := make([]ToolSearchResult, len(ranked))
|
||||||
|
for i, r := range ranked {
|
||||||
|
out[i] = ToolSearchResult{
|
||||||
|
Name: r.Document.Name,
|
||||||
|
Description: r.Document.Description,
|
||||||
|
Parameters: r.Document.Tool.Parameters(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
236
pkg/tools/search_tools_test.go
Normal file
236
pkg/tools/search_tools_test.go
Normal file
|
|
@ -0,0 +1,236 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Dummy tool to fill the registry in our tests.
|
||||||
|
type mockSearchableTool struct {
|
||||||
|
name string
|
||||||
|
desc string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockSearchableTool) Name() string { return m.name }
|
||||||
|
func (m *mockSearchableTool) Description() string { return m.desc }
|
||||||
|
func (m *mockSearchableTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{"type": "object"}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockSearchableTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
|
return SilentResult("mock executed: " + m.name)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper to initialize a populated ToolRegistry
|
||||||
|
func setupPopulatedRegistry() *ToolRegistry {
|
||||||
|
reg := NewToolRegistry()
|
||||||
|
|
||||||
|
// A core tool (NOT to be found by searches)
|
||||||
|
reg.Register(&mockSearchableTool{
|
||||||
|
name: "core_search",
|
||||||
|
desc: "I am a visible core tool for searching files",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Hidden tools (must be found by searches)
|
||||||
|
reg.RegisterHidden(&mockSearchableTool{
|
||||||
|
name: "mcp_read_file",
|
||||||
|
desc: "Read the contents of a system file",
|
||||||
|
})
|
||||||
|
reg.RegisterHidden(&mockSearchableTool{
|
||||||
|
name: "mcp_list_dir",
|
||||||
|
desc: "List directories and files in the system",
|
||||||
|
})
|
||||||
|
reg.RegisterHidden(&mockSearchableTool{
|
||||||
|
name: "mcp_fetch_net",
|
||||||
|
desc: "Fetch data from a network database",
|
||||||
|
})
|
||||||
|
|
||||||
|
return reg
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegexSearchTool_Execute(t *testing.T) {
|
||||||
|
reg := setupPopulatedRegistry()
|
||||||
|
tool := NewRegexSearchTool(reg, 5, 10)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
t.Run("Empty Pattern Error", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{})
|
||||||
|
if !res.IsError || !strings.Contains(res.ForLLM, "Missing or invalid 'pattern'") {
|
||||||
|
t.Errorf("Expected missing pattern error, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Invalid Regex Syntax", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{"pattern": "[unclosed"})
|
||||||
|
if !res.IsError || !strings.Contains(res.ForLLM, "Invalid regex pattern syntax") {
|
||||||
|
t.Errorf("Expected regex syntax error, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("No Match Found", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{"pattern": "alien"})
|
||||||
|
if res.IsError || !strings.Contains(res.ForLLM, "No tools found matching") {
|
||||||
|
t.Errorf("Expected 'no tools found' message, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Successful Match & Promotion", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{"pattern": "system"})
|
||||||
|
|
||||||
|
if res.IsError {
|
||||||
|
t.Fatalf("Unexpected error: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
if !strings.Contains(res.ForLLM, "SUCCESS: These tools have been temporarily UNLOCKED") {
|
||||||
|
t.Errorf("Expected success string, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
if !strings.Contains(res.ForLLM, "mcp_read_file") {
|
||||||
|
t.Errorf("Expected 'mcp_read_file' in results")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify that the TTL has been updated for the tools found
|
||||||
|
reg.mu.RLock()
|
||||||
|
defer reg.mu.RUnlock()
|
||||||
|
if reg.tools["mcp_read_file"].TTL != 5 {
|
||||||
|
t.Errorf("Expected TTL of 'mcp_read_file' to be promoted to 5, got %d", reg.tools["mcp_read_file"].TTL)
|
||||||
|
}
|
||||||
|
if reg.tools["mcp_fetch_net"].TTL != 0 {
|
||||||
|
t.Errorf("Expected 'mcp_fetch_net' to NOT be promoted (TTL=0)")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25SearchTool_Execute(t *testing.T) {
|
||||||
|
reg := setupPopulatedRegistry()
|
||||||
|
tool := NewBM25SearchTool(reg, 3, 10)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
t.Run("Empty Query Error", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{"query": " "})
|
||||||
|
if !res.IsError || !strings.Contains(res.ForLLM, "Missing or invalid 'query'") {
|
||||||
|
t.Errorf("Expected missing query error, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("No Match Found", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{"query": "aliens spaceships"})
|
||||||
|
if res.IsError || !strings.Contains(res.ForLLM, "No tools found matching") {
|
||||||
|
t.Errorf("Expected 'no tools found', got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Successful Match & Promotion", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{"query": "read files"})
|
||||||
|
|
||||||
|
if res.IsError {
|
||||||
|
t.Fatalf("Unexpected error: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
if !strings.Contains(res.ForLLM, "mcp_read_file") {
|
||||||
|
t.Errorf("Expected 'mcp_read_file' in BM25 results")
|
||||||
|
}
|
||||||
|
|
||||||
|
reg.mu.RLock()
|
||||||
|
defer reg.mu.RUnlock()
|
||||||
|
if reg.tools["mcp_read_file"].TTL != 3 {
|
||||||
|
t.Errorf("Expected TTL of 'mcp_read_file' to be promoted to 3")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCallDiscoveredTool_Execute(t *testing.T) {
|
||||||
|
reg := setupPopulatedRegistry()
|
||||||
|
tool := NewCallDiscoveredTool(reg, 8)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
t.Run("Missing Name", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{"arguments": map[string]any{}})
|
||||||
|
if !res.IsError {
|
||||||
|
t.Error("Expected error for missing tool_name")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Invalid Arguments Type Fallback", func(t *testing.T) {
|
||||||
|
// If the LLM hallucinates and passes a string instead of an object/map,
|
||||||
|
// does a graceful fallback to empty map.
|
||||||
|
res := tool.Execute(ctx, map[string]any{
|
||||||
|
"tool_name": "mcp_read_file",
|
||||||
|
"arguments": "invalid-string-instead-of-object",
|
||||||
|
})
|
||||||
|
// It must be an error
|
||||||
|
if !res.IsError {
|
||||||
|
t.Fatalf("Expected an error for invalid argument type, but got success: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
// The error message should contain the explanation that we have added
|
||||||
|
if !strings.Contains(res.ForLLM, "Invalid 'arguments' format") {
|
||||||
|
t.Errorf("Expected instructional error message, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Successful Passthrough", func(t *testing.T) {
|
||||||
|
res := tool.Execute(ctx, map[string]any{
|
||||||
|
"tool_name": "mcp_read_file",
|
||||||
|
"arguments": map[string]any{"path": "/tmp/test.txt"},
|
||||||
|
})
|
||||||
|
|
||||||
|
if res.IsError {
|
||||||
|
t.Fatalf("Unexpected error: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
if !strings.Contains(res.ForLLM, "mock executed: mcp_read_file") {
|
||||||
|
t.Errorf("Expected underlying tool to be executed, got: %v", res.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The tool should renew the TTL of the tool called
|
||||||
|
reg.mu.RLock()
|
||||||
|
defer reg.mu.RUnlock()
|
||||||
|
if reg.tools["mcp_read_file"].TTL != 8 {
|
||||||
|
t.Errorf("Expected TTL to be renewed to 8")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolRegistry_SearchLimitsAndCoreFiltering(t *testing.T) {
|
||||||
|
reg := NewToolRegistry()
|
||||||
|
|
||||||
|
// Add 1 Core and 10 Hidden, all containing the word "match"
|
||||||
|
reg.Register(&mockSearchableTool{"core_match", "I am core with match"})
|
||||||
|
for i := 0; i < 10; i++ {
|
||||||
|
reg.RegisterHidden(&mockSearchableTool{
|
||||||
|
name: fmt.Sprintf("hidden_match_%d", i),
|
||||||
|
desc: "this has a match",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("Regex limits and core filtering", func(t *testing.T) {
|
||||||
|
// Search with Regex and a limit of maxSearchResults = 4
|
||||||
|
res, err := reg.SearchRegex("match", 4)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("SearchRegex failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(res) != 4 {
|
||||||
|
t.Errorf("Expected exactly 4 results due to limit, got %d", len(res))
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, r := range res {
|
||||||
|
if r.Name == "core_match" {
|
||||||
|
t.Errorf("SearchRegex returned a Core tool, which should be excluded")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("BM25 limits and core filtering", func(t *testing.T) {
|
||||||
|
// Search with BM25 and a limit of maxSearchResults = 3
|
||||||
|
res := reg.SearchBM25("match", 3)
|
||||||
|
|
||||||
|
if len(res) != 3 {
|
||||||
|
t.Errorf("Expected exactly 3 results due to limit, got %d", len(res))
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, r := range res {
|
||||||
|
if r.Name == "core_match" {
|
||||||
|
t.Errorf("SearchBM25 returned a Core tool, which should be excluded")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
272
pkg/utils/bm25.go
Normal file
272
pkg/utils/bm25.go
Normal file
|
|
@ -0,0 +1,272 @@
|
||||||
|
// Package utils provides shared, reusable algorithms.
|
||||||
|
// This file implements a generic BM25 search engine.
|
||||||
|
//
|
||||||
|
// Usage:
|
||||||
|
//
|
||||||
|
// type MyDoc struct { ID string; Body string }
|
||||||
|
//
|
||||||
|
// corpus := []MyDoc{...}
|
||||||
|
// engine := bm25.New(corpus, func(d MyDoc) string {
|
||||||
|
// return d.ID + " " + d.Body
|
||||||
|
// })
|
||||||
|
// results := engine.Search("my query", 5)
|
||||||
|
package utils
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ── Tuning defaults ───────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
const (
|
||||||
|
// DefaultBM25K1 is the term-frequency saturation factor (typical range 1.2–2.0).
|
||||||
|
// Higher values give more weight to repeated terms.
|
||||||
|
DefaultBM25K1 = 1.2
|
||||||
|
|
||||||
|
// DefaultBM25B is the document-length normalization factor (0 = none, 1 = full).
|
||||||
|
DefaultBM25B = 0.75
|
||||||
|
)
|
||||||
|
|
||||||
|
// BM25Engine is a query-time BM25 search engine over a generic corpus.
|
||||||
|
// T is the document type; the caller supplies a TextFunc that extracts the
|
||||||
|
// searchable text from each document.
|
||||||
|
//
|
||||||
|
// The engine is stateless between queries: no caching, no invalidation logic.
|
||||||
|
// All indexing work is performed inside Search() on every call, making it
|
||||||
|
// safe to use on corpora that change frequently.
|
||||||
|
type BM25Engine[T any] struct {
|
||||||
|
corpus []T
|
||||||
|
textFunc func(T) string
|
||||||
|
k1 float64
|
||||||
|
b float64
|
||||||
|
}
|
||||||
|
|
||||||
|
// BM25Option is a functional option to configure a BM25Engine.
|
||||||
|
type BM25Option func(*bm25Config)
|
||||||
|
|
||||||
|
type bm25Config struct {
|
||||||
|
k1 float64
|
||||||
|
b float64
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithK1 overrides the term-frequency saturation constant (default 1.2).
|
||||||
|
func WithK1(k1 float64) BM25Option {
|
||||||
|
return func(c *bm25Config) { c.k1 = k1 }
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithB overrides the document-length normalization factor (default 0.75).
|
||||||
|
func WithB(b float64) BM25Option {
|
||||||
|
return func(c *bm25Config) { c.b = b }
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewBM25Engine creates a BM25Engine for the given corpus.
|
||||||
|
//
|
||||||
|
// - corpus : slice of documents of any type T.
|
||||||
|
// - textFunc : function that returns the searchable text for a document.
|
||||||
|
// - opts : optional tuning (WithK1, WithB).
|
||||||
|
//
|
||||||
|
// The corpus slice is referenced, not copied. Callers must not mutate it
|
||||||
|
// concurrently with Search().
|
||||||
|
func NewBM25Engine[T any](corpus []T, textFunc func(T) string, opts ...BM25Option) *BM25Engine[T] {
|
||||||
|
cfg := bm25Config{k1: DefaultBM25K1, b: DefaultBM25B}
|
||||||
|
for _, o := range opts {
|
||||||
|
o(&cfg)
|
||||||
|
}
|
||||||
|
return &BM25Engine[T]{
|
||||||
|
corpus: corpus,
|
||||||
|
textFunc: textFunc,
|
||||||
|
k1: cfg.k1,
|
||||||
|
b: cfg.b,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BM25Result is a single ranked result from a Search call.
|
||||||
|
type BM25Result[T any] struct {
|
||||||
|
Document T
|
||||||
|
Score float32
|
||||||
|
}
|
||||||
|
|
||||||
|
// Search ranks the corpus against query and returns the top-k results.
|
||||||
|
// Returns an empty slice (not nil) when there are no matches.
|
||||||
|
//
|
||||||
|
// Complexity: O(N×L) for indexing + O(|Q|×avgPostingLen) for scoring,
|
||||||
|
// where N = corpus size, L = average document length, Q = query terms.
|
||||||
|
// Top-k extraction uses a fixed-size min-heap: O(candidates × log k).
|
||||||
|
func (e *BM25Engine[T]) Search(query string, topK int) []BM25Result[T] {
|
||||||
|
if topK <= 0 {
|
||||||
|
return []BM25Result[T]{}
|
||||||
|
}
|
||||||
|
|
||||||
|
queryTerms := bm25Tokenize(query)
|
||||||
|
if len(queryTerms) == 0 {
|
||||||
|
return []BM25Result[T]{}
|
||||||
|
}
|
||||||
|
|
||||||
|
N := len(e.corpus)
|
||||||
|
if N == 0 {
|
||||||
|
return []BM25Result[T]{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 1: build per-document tf + raw doc lengths
|
||||||
|
type docEntry struct {
|
||||||
|
tf map[string]uint32
|
||||||
|
rawLen int
|
||||||
|
}
|
||||||
|
|
||||||
|
entries := make([]docEntry, N)
|
||||||
|
df := make(map[string]int, 64)
|
||||||
|
totalLen := 0
|
||||||
|
|
||||||
|
for i, doc := range e.corpus {
|
||||||
|
tokens := bm25Tokenize(e.textFunc(doc))
|
||||||
|
totalLen += len(tokens)
|
||||||
|
|
||||||
|
tf := make(map[string]uint32, len(tokens))
|
||||||
|
for _, t := range tokens {
|
||||||
|
tf[t]++
|
||||||
|
}
|
||||||
|
// df: each term counts once per document (iterate the map, keys are unique)
|
||||||
|
for t := range tf {
|
||||||
|
df[t]++
|
||||||
|
}
|
||||||
|
|
||||||
|
entries[i] = docEntry{tf: tf, rawLen: len(tokens)}
|
||||||
|
}
|
||||||
|
|
||||||
|
avgDocLen := float64(totalLen) / float64(N)
|
||||||
|
|
||||||
|
// Step 2: pre-compute IDF and per-doc length normalization
|
||||||
|
// IDF (Robertson smoothing): log( (N - df(t) + 0.5) / (df(t) + 0.5) + 1 )
|
||||||
|
idf := make(map[string]float32, len(df))
|
||||||
|
for term, freq := range df {
|
||||||
|
idf[term] = float32(math.Log(
|
||||||
|
(float64(N)-float64(freq)+0.5)/(float64(freq)+0.5) + 1,
|
||||||
|
))
|
||||||
|
}
|
||||||
|
|
||||||
|
// docLenNorm[i] = k1 * (1 - b + b * |doc_i| / avgDocLen)
|
||||||
|
// Stored as float32 — sufficient precision for ranking.
|
||||||
|
docLenNorm := make([]float32, N)
|
||||||
|
for i, entry := range entries {
|
||||||
|
docLenNorm[i] = float32(e.k1 * (1 - e.b + e.b*float64(entry.rawLen)/avgDocLen))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 3: build inverted index (posting lists)
|
||||||
|
// Iterate the tf map directly — map keys are already unique, no seen-set needed.
|
||||||
|
posting := make(map[string][]int32, len(df))
|
||||||
|
for i, entry := range entries {
|
||||||
|
for term := range entry.tf {
|
||||||
|
posting[term] = append(posting[term], int32(i))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 4: score via posting lists
|
||||||
|
// Deduplicate query terms to avoid double-weighting the same term.
|
||||||
|
unique := bm25Dedupe(queryTerms)
|
||||||
|
|
||||||
|
scores := make(map[int32]float32)
|
||||||
|
for _, term := range unique {
|
||||||
|
termIDF, ok := idf[term]
|
||||||
|
if !ok {
|
||||||
|
continue // term not in vocabulary → zero contribution
|
||||||
|
}
|
||||||
|
for _, docID := range posting[term] {
|
||||||
|
freq := float32(entries[docID].tf[term])
|
||||||
|
// TF_norm = freq * (k1+1) / (freq + docLenNorm)
|
||||||
|
tfNorm := freq * float32(e.k1+1) / (freq + docLenNorm[docID])
|
||||||
|
scores[docID] += termIDF * tfNorm
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(scores) == 0 {
|
||||||
|
return []BM25Result[T]{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Step 5: top-K via fixed-size min-heap
|
||||||
|
heap := make([]bm25ScoredDoc, 0, topK)
|
||||||
|
|
||||||
|
for docID, sc := range scores {
|
||||||
|
switch {
|
||||||
|
case len(heap) < topK:
|
||||||
|
heap = append(heap, bm25ScoredDoc{docID: docID, score: sc})
|
||||||
|
if len(heap) == topK {
|
||||||
|
bm25MinHeapify(heap)
|
||||||
|
}
|
||||||
|
case sc > heap[0].score:
|
||||||
|
heap[0] = bm25ScoredDoc{docID: docID, score: sc}
|
||||||
|
bm25SiftDown(heap, 0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
sort.Slice(heap, func(i, j int) bool { return heap[i].score > heap[j].score })
|
||||||
|
|
||||||
|
out := make([]BM25Result[T], len(heap))
|
||||||
|
for i, h := range heap {
|
||||||
|
out[i] = BM25Result[T]{
|
||||||
|
Document: e.corpus[h.docID],
|
||||||
|
Score: h.score,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// bm25Tokenize splits s into lowercase tokens, stripping edge punctuation.
|
||||||
|
func bm25Tokenize(s string) []string {
|
||||||
|
raw := strings.Fields(strings.ToLower(s))
|
||||||
|
out := raw[:0] // reuse backing array to avoid extra allocation
|
||||||
|
for _, t := range raw {
|
||||||
|
t = strings.Trim(t, ".,;:!?\"'()/\\-_")
|
||||||
|
if t != "" {
|
||||||
|
out = append(out, t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// bm25Dedupe returns a new slice with duplicate tokens removed,
|
||||||
|
// preserving first-occurrence order.
|
||||||
|
func bm25Dedupe(tokens []string) []string {
|
||||||
|
seen := make(map[string]struct{}, len(tokens))
|
||||||
|
out := make([]string, 0, len(tokens))
|
||||||
|
for _, t := range tokens {
|
||||||
|
if _, ok := seen[t]; !ok {
|
||||||
|
seen[t] = struct{}{}
|
||||||
|
out = append(out, t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
type bm25ScoredDoc struct {
|
||||||
|
docID int32
|
||||||
|
score float32
|
||||||
|
}
|
||||||
|
|
||||||
|
// bm25MinHeapify builds a min-heap in-place using Floyd's algorithm: O(k).
|
||||||
|
func bm25MinHeapify(h []bm25ScoredDoc) {
|
||||||
|
for i := len(h)/2 - 1; i >= 0; i-- {
|
||||||
|
bm25SiftDown(h, i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// bm25SiftDown restores the min-heap property starting at node i: O(log k).
|
||||||
|
func bm25SiftDown(h []bm25ScoredDoc, i int) {
|
||||||
|
n := len(h)
|
||||||
|
for {
|
||||||
|
smallest := i
|
||||||
|
l, r := 2*i+1, 2*i+2
|
||||||
|
if l < n && h[l].score < h[smallest].score {
|
||||||
|
smallest = l
|
||||||
|
}
|
||||||
|
if r < n && h[r].score < h[smallest].score {
|
||||||
|
smallest = r
|
||||||
|
}
|
||||||
|
if smallest == i {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
h[i], h[smallest] = h[smallest], h[i]
|
||||||
|
i = smallest
|
||||||
|
}
|
||||||
|
}
|
||||||
175
pkg/utils/bm25_test.go
Normal file
175
pkg/utils/bm25_test.go
Normal file
|
|
@ -0,0 +1,175 @@
|
||||||
|
package utils
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// testDoc is a generic structure for use in tests.
|
||||||
|
type testDoc struct {
|
||||||
|
ID int
|
||||||
|
Text string
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractText(d testDoc) string {
|
||||||
|
return d.Text
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Search_EdgeCases(t *testing.T) {
|
||||||
|
corpus := []testDoc{
|
||||||
|
{1, "hello world"},
|
||||||
|
{2, "foo bar"},
|
||||||
|
}
|
||||||
|
engine := NewBM25Engine(corpus, extractText)
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
query string
|
||||||
|
topK int
|
||||||
|
}{
|
||||||
|
{"Zero topK", "hello", 0},
|
||||||
|
{"Negative topK", "hello", -1},
|
||||||
|
{"Empty query", "", 5},
|
||||||
|
{"Query with only punctuation", "...,,,!!!", 5},
|
||||||
|
{"No matches found", "golang", 5},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
results := engine.Search(tt.query, tt.topK)
|
||||||
|
if len(results) != 0 {
|
||||||
|
t.Errorf("expected 0 results, got %d", len(results))
|
||||||
|
}
|
||||||
|
// Check that it never returns nil, but an empty slice
|
||||||
|
if results == nil {
|
||||||
|
t.Errorf("expected empty slice, got nil")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Search_EmptyCorpus(t *testing.T) {
|
||||||
|
engine := NewBM25Engine([]testDoc{}, extractText)
|
||||||
|
results := engine.Search("hello", 5)
|
||||||
|
if len(results) != 0 || results == nil {
|
||||||
|
t.Errorf("expected empty slice from empty corpus, got %v", results)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Search_RankingLogic(t *testing.T) {
|
||||||
|
corpus := []testDoc{
|
||||||
|
{1, "the quick brown fox jumps over the lazy dog"},
|
||||||
|
{2, "quick fox"},
|
||||||
|
{3, "quick quick quick fox"}, // High Term Frequency (TF)
|
||||||
|
{4, "completely irrelevant document here"},
|
||||||
|
}
|
||||||
|
engine := NewBM25Engine(corpus, extractText)
|
||||||
|
|
||||||
|
t.Run("Term Frequency (TF) boosts score", func(t *testing.T) {
|
||||||
|
results := engine.Search("quick", 5)
|
||||||
|
if len(results) < 3 {
|
||||||
|
t.Fatalf("expected at least 3 results, got %d", len(results))
|
||||||
|
}
|
||||||
|
// Doc 3 has the word "quick" repeated 3 times, it should beat Doc 2
|
||||||
|
if results[0].Document.ID != 3 {
|
||||||
|
t.Errorf("expected doc 3 to rank first due to high TF, got doc %d", results[0].Document.ID)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("Document Length penalty", func(t *testing.T) {
|
||||||
|
results := engine.Search("fox", 5)
|
||||||
|
if len(results) < 3 {
|
||||||
|
t.Fatalf("expected at least 3 results, got %d", len(results))
|
||||||
|
}
|
||||||
|
// Doc 2 ("quick fox") is much shorter than Doc 1 ("the quick brown fox..."),
|
||||||
|
// so, with equal Term Frequency for the word "fox" (1 time), Doc 2 wins.
|
||||||
|
if results[0].Document.ID != 2 {
|
||||||
|
t.Errorf("expected doc 2 to rank first due to shorter length, got doc %d", results[0].Document.ID)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("TopK limits results", func(t *testing.T) {
|
||||||
|
results := engine.Search("quick", 2)
|
||||||
|
if len(results) != 2 {
|
||||||
|
t.Errorf("expected exactly 2 results, got %d", len(results))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Tokenize(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
input string
|
||||||
|
expected []string
|
||||||
|
}{
|
||||||
|
{"Hello World", []string{"hello", "world"}},
|
||||||
|
{" spaces everywhere ", []string{"spaces", "everywhere"}},
|
||||||
|
{"punctuation... test!!!", []string{"punctuation", "test"}},
|
||||||
|
{"(parentheses) and-hyphens", []string{"parentheses", "and-hyphens"}}, // hyphens trimmed from edges
|
||||||
|
{"internal-hyphen is kept", []string{"internal-hyphen", "is", "kept"}},
|
||||||
|
{".,;?!", []string{}}, // Becomes empty after trim
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.input, func(t *testing.T) {
|
||||||
|
got := bm25Tokenize(tt.input)
|
||||||
|
if len(got) == 0 && len(tt.expected) == 0 {
|
||||||
|
return // Both empty
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(got, tt.expected) {
|
||||||
|
t.Errorf("bm25Tokenize(%q) = %v, want %v", tt.input, got, tt.expected)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Dedupe(t *testing.T) {
|
||||||
|
input := []string{"apple", "banana", "apple", "orange", "banana"}
|
||||||
|
expected := []string{"apple", "banana", "orange"}
|
||||||
|
|
||||||
|
got := bm25Dedupe(input)
|
||||||
|
if !reflect.DeepEqual(got, expected) {
|
||||||
|
t.Errorf("bm25Dedupe() = %v, want %v", got, expected)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Options(t *testing.T) {
|
||||||
|
corpus := []testDoc{{1, "test"}}
|
||||||
|
|
||||||
|
engine := NewBM25Engine(
|
||||||
|
corpus,
|
||||||
|
extractText,
|
||||||
|
WithK1(2.5),
|
||||||
|
WithB(0.9),
|
||||||
|
)
|
||||||
|
|
||||||
|
if engine.k1 != 2.5 {
|
||||||
|
t.Errorf("expected k1 to be 2.5, got %v", engine.k1)
|
||||||
|
}
|
||||||
|
if engine.b != 0.9 {
|
||||||
|
t.Errorf("expected b to be 0.9, got %v", engine.b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBM25Search_SortingStability(t *testing.T) {
|
||||||
|
// Ensure that sorting by heap returns in correct descending order
|
||||||
|
corpus := []testDoc{
|
||||||
|
{1, "golang is good"},
|
||||||
|
{2, "golang golang"},
|
||||||
|
{3, "golang golang golang"},
|
||||||
|
{4, "golang golang golang golang"},
|
||||||
|
}
|
||||||
|
engine := NewBM25Engine(corpus, extractText)
|
||||||
|
results := engine.Search("golang", 10)
|
||||||
|
|
||||||
|
if len(results) != 4 {
|
||||||
|
t.Fatalf("expected 4 results, got %d", len(results))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Score should be strictly decreasing
|
||||||
|
for i := 1; i < len(results); i++ {
|
||||||
|
if results[i].Score > results[i-1].Score {
|
||||||
|
t.Errorf("results not sorted correctly: result %d score (%v) > result %d score (%v)",
|
||||||
|
i, results[i].Score, i-1, results[i-1].Score)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue