added two new providers: NVIDIA and Azure plus Security enhancments to lock down skills if desired and added a configurable chat API
This commit is contained in:
parent
60d7ec20a5
commit
0a5cf5668e
24 changed files with 560 additions and 85 deletions
2
logs/gateway.log
Normal file
2
logs/gateway.log
Normal file
|
|
@ -0,0 +1,2 @@
|
||||||
|
{"level":"warn","path":"/home/stevef/dev/tomerge/github/picoclaw/config.json","time":"2026-03-24T08:13:49+01:00","caller":"/home/stevef/dev/tomerge/github/picoclaw/pkg/config/config.go:1363","message":"config file not found, using default config"}
|
||||||
|
{"level":"warn","path":"/home/stevef/dev/tomerge/github/picoclaw/config.json","time":"2026-03-24T08:15:23+01:00","caller":"/home/stevef/dev/tomerge/github/picoclaw/pkg/config/config.go:1363","message":"config file not found, using default config"}
|
||||||
26
logs/gateway_panic.log
Normal file
26
logs/gateway_panic.log
Normal file
|
|
@ -0,0 +1,26 @@
|
||||||
|
Error: error creating provider: model "" not found in model_list: model "" not found in model_list or providers
|
||||||
|
Usage:
|
||||||
|
picoclaw gateway [flags]
|
||||||
|
|
||||||
|
Aliases:
|
||||||
|
gateway, g
|
||||||
|
|
||||||
|
Flags:
|
||||||
|
-E, --allow-empty Continue starting even when no default model is configured
|
||||||
|
-d, --debug Enable debug logging
|
||||||
|
-h, --help help for gateway
|
||||||
|
-T, --no-truncate Disable string truncation in debug logs
|
||||||
|
|
||||||
|
Error: error creating provider: model "" not found in model_list: model "" not found in model_list or providers
|
||||||
|
Usage:
|
||||||
|
picoclaw gateway [flags]
|
||||||
|
|
||||||
|
Aliases:
|
||||||
|
gateway, g
|
||||||
|
|
||||||
|
Flags:
|
||||||
|
-E, --allow-empty Continue starting even when no default model is configured
|
||||||
|
-d, --debug Enable debug logging
|
||||||
|
-h, --help help for gateway
|
||||||
|
-T, --no-truncate Disable string truncation in debug logs
|
||||||
|
|
||||||
|
|
@ -272,11 +272,11 @@ func registerSharedTools(
|
||||||
cfg.Tools.Skills.SearchCache.MaxSize,
|
cfg.Tools.Skills.SearchCache.MaxSize,
|
||||||
time.Duration(cfg.Tools.Skills.SearchCache.TTLSeconds)*time.Second,
|
time.Duration(cfg.Tools.Skills.SearchCache.TTLSeconds)*time.Second,
|
||||||
)
|
)
|
||||||
agent.Tools.Register(tools.NewFindSkillsTool(registryMgr, searchCache))
|
agent.Tools.Register(tools.NewFindSkillsTool(registryMgr, searchCache, cfg.Tools.Skills.Whitelist, cfg.Tools.Skills.WhitelistEnabled))
|
||||||
}
|
}
|
||||||
|
|
||||||
if install_skills_enable {
|
if install_skills_enable {
|
||||||
agent.Tools.Register(tools.NewInstallSkillTool(registryMgr, agent.Workspace))
|
agent.Tools.Register(tools.NewInstallSkillTool(registryMgr, agent.Workspace, cfg.Tools.Skills.Whitelist, cfg.Tools.Skills.WhitelistEnabled))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -374,6 +374,8 @@ func registerSharedTools(
|
||||||
} else if (spawnEnabled || spawnStatusEnabled) && !cfg.Tools.IsToolEnabled("subagent") {
|
} else if (spawnEnabled || spawnStatusEnabled) && !cfg.Tools.IsToolEnabled("subagent") {
|
||||||
logger.WarnCF("agent", "spawn/spawn_status tools require subagent to be enabled", nil)
|
logger.WarnCF("agent", "spawn/spawn_status tools require subagent to be enabled", nil)
|
||||||
}
|
}
|
||||||
|
// Apply global tools whitelist
|
||||||
|
agent.Tools.Filter(cfg.Tools.Whitelist, cfg.Tools.WhitelistEnabled)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -383,7 +385,7 @@ func (al *AgentLoop) Run(ctx context.Context) error {
|
||||||
if err := al.ensureHooksInitialized(ctx); err != nil {
|
if err := al.ensureHooksInitialized(ctx); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if err := al.ensureMCPInitialized(ctx); err != nil {
|
if err := al.EnsureMCPInitialized(ctx); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -1204,7 +1206,7 @@ func (al *AgentLoop) ProcessDirectWithChannel(
|
||||||
if err := al.ensureHooksInitialized(ctx); err != nil {
|
if err := al.ensureHooksInitialized(ctx); err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
if err := al.ensureMCPInitialized(ctx); err != nil {
|
if err := al.EnsureMCPInitialized(ctx); err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -1228,7 +1230,7 @@ func (al *AgentLoop) ProcessHeartbeat(
|
||||||
if err := al.ensureHooksInitialized(ctx); err != nil {
|
if err := al.ensureHooksInitialized(ctx); err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
if err := al.ensureMCPInitialized(ctx); err != nil {
|
if err := al.EnsureMCPInitialized(ctx); err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -59,7 +59,7 @@ func (r *mcpRuntime) hasManager() bool {
|
||||||
|
|
||||||
// ensureMCPInitialized loads MCP servers/tools once so both Run() and direct
|
// ensureMCPInitialized loads MCP servers/tools once so both Run() and direct
|
||||||
// agent mode share the same initialization path.
|
// agent mode share the same initialization path.
|
||||||
func (al *AgentLoop) ensureMCPInitialized(ctx context.Context) error {
|
func (al *AgentLoop) EnsureMCPInitialized(ctx context.Context) error {
|
||||||
if !al.cfg.Tools.IsToolEnabled("mcp") {
|
if !al.cfg.Tools.IsToolEnabled("mcp") {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -332,7 +332,7 @@ func (al *AgentLoop) Continue(ctx context.Context, sessionKey, channel, chatID s
|
||||||
if err := al.ensureHooksInitialized(ctx); err != nil {
|
if err := al.ensureHooksInitialized(ctx); err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
if err := al.ensureMCPInitialized(ctx); err != nil {
|
if err := al.EnsureMCPInitialized(ctx); err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -464,6 +464,10 @@ type DiscordConfig struct {
|
||||||
ReasoningChannelID string `json:"reasoning_channel_id" yaml:"-" env:"PICOCLAW_CHANNELS_DISCORD_REASONING_CHANNEL_ID"`
|
ReasoningChannelID string `json:"reasoning_channel_id" yaml:"-" env:"PICOCLAW_CHANNELS_DISCORD_REASONING_CHANNEL_ID"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *DiscordConfig) SetToken(token string) {
|
||||||
|
c.Token = *NewSecureString(token)
|
||||||
|
}
|
||||||
|
|
||||||
type MaixCamConfig struct {
|
type MaixCamConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_MAIXCAM_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_CHANNELS_MAIXCAM_ENABLED"`
|
||||||
Host string `json:"host" env:"PICOCLAW_CHANNELS_MAIXCAM_HOST"`
|
Host string `json:"host" env:"PICOCLAW_CHANNELS_MAIXCAM_HOST"`
|
||||||
|
|
@ -504,6 +508,14 @@ type SlackConfig struct {
|
||||||
ReasoningChannelID string `json:"reasoning_channel_id" yaml:"-" env:"PICOCLAW_CHANNELS_SLACK_REASONING_CHANNEL_ID"`
|
ReasoningChannelID string `json:"reasoning_channel_id" yaml:"-" env:"PICOCLAW_CHANNELS_SLACK_REASONING_CHANNEL_ID"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *SlackConfig) SetBotToken(token string) {
|
||||||
|
c.BotToken = *NewSecureString(token)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *SlackConfig) SetAppToken(token string) {
|
||||||
|
c.AppToken = *NewSecureString(token)
|
||||||
|
}
|
||||||
|
|
||||||
type MatrixConfig struct {
|
type MatrixConfig struct {
|
||||||
Enabled bool `json:"enabled" yaml:"-" env:"PICOCLAW_CHANNELS_MATRIX_ENABLED"`
|
Enabled bool `json:"enabled" yaml:"-" env:"PICOCLAW_CHANNELS_MATRIX_ENABLED"`
|
||||||
Homeserver string `json:"homeserver" yaml:"-" env:"PICOCLAW_CHANNELS_MATRIX_HOMESERVER"`
|
Homeserver string `json:"homeserver" yaml:"-" env:"PICOCLAW_CHANNELS_MATRIX_HOMESERVER"`
|
||||||
|
|
@ -680,6 +692,24 @@ type ModelConfig struct {
|
||||||
isVirtual bool
|
isVirtual bool
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *ModelConfig) UnmarshalJSON(data []byte) error {
|
||||||
|
type Alias ModelConfig
|
||||||
|
aux := &struct {
|
||||||
|
APIKey string `json:"api_key"`
|
||||||
|
APIKeys []string `json:"api_keys"`
|
||||||
|
*Alias
|
||||||
|
}{
|
||||||
|
Alias: (*Alias)(c),
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal(data, aux); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
c.apiKeys = MergeAPIKeys(aux.APIKey, aux.APIKeys)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// APIKey returns the first API key from apiKeys
|
// APIKey returns the first API key from apiKeys
|
||||||
func (c *ModelConfig) APIKey() string {
|
func (c *ModelConfig) APIKey() string {
|
||||||
if len(c.APIKeys) > 0 {
|
if len(c.APIKeys) > 0 {
|
||||||
|
|
@ -713,10 +743,12 @@ func (c *ModelConfig) SetAPIKey(value string) {
|
||||||
}
|
}
|
||||||
|
|
||||||
type GatewayConfig struct {
|
type GatewayConfig struct {
|
||||||
Host string `json:"host" env:"PICOCLAW_GATEWAY_HOST"`
|
Host string `json:"host" env:"PICOCLAW_GATEWAY_HOST"`
|
||||||
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
||||||
HotReload bool `json:"hot_reload" env:"PICOCLAW_GATEWAY_HOT_RELOAD"`
|
APIKey string `json:"api_key" env:"PICOCLAW_GATEWAY_API_KEY"`
|
||||||
LogLevel string `json:"log_level,omitempty" env:"PICOCLAW_LOG_LEVEL"`
|
ChatEnabled bool `json:"chat_enabled" env:"PICOCLAW_GATEWAY_CHAT_ENABLED"`
|
||||||
|
HotReload bool `json:"hot_reload" env:"PICOCLAW_GATEWAY_HOT_RELOAD"`
|
||||||
|
LogLevel string `json:"log_level,omitempty" env:"PICOCLAW_LOG_LEVEL"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type ToolDiscoveryConfig struct {
|
type ToolDiscoveryConfig struct {
|
||||||
|
|
@ -873,6 +905,8 @@ type SkillsToolsConfig struct {
|
||||||
Github SkillsGithubConfig `yaml:"github,omitempty" json:"github"`
|
Github SkillsGithubConfig `yaml:"github,omitempty" json:"github"`
|
||||||
MaxConcurrentSearches int `yaml:"-" json:"max_concurrent_searches" env:"PICOCLAW_TOOLS_SKILLS_MAX_CONCURRENT_SEARCHES"`
|
MaxConcurrentSearches int `yaml:"-" json:"max_concurrent_searches" env:"PICOCLAW_TOOLS_SKILLS_MAX_CONCURRENT_SEARCHES"`
|
||||||
SearchCache SearchCacheConfig `yaml:"-" json:"search_cache"`
|
SearchCache SearchCacheConfig `yaml:"-" json:"search_cache"`
|
||||||
|
Whitelist FlexibleStringSlice `json:"whitelist,omitempty" env:"PICOCLAW_TOOLS_SKILLS_WHITELIST"`
|
||||||
|
WhitelistEnabled bool `json:"whitelist_enabled,omitempty" env:"PICOCLAW_TOOLS_SKILLS_WHITELIST_ENABLED"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type MediaCleanupConfig struct {
|
type MediaCleanupConfig struct {
|
||||||
|
|
@ -902,7 +936,9 @@ type ToolsConfig struct {
|
||||||
Exec ExecConfig `json:"exec" yaml:"-"`
|
Exec ExecConfig `json:"exec" yaml:"-"`
|
||||||
Skills SkillsToolsConfig `json:"skills" yaml:"skills,omitempty"`
|
Skills SkillsToolsConfig `json:"skills" yaml:"skills,omitempty"`
|
||||||
MediaCleanup MediaCleanupConfig `json:"media_cleanup" yaml:"-"`
|
MediaCleanup MediaCleanupConfig `json:"media_cleanup" yaml:"-"`
|
||||||
MCP MCPConfig `json:"mcp" yaml:"-"`
|
Whitelist FlexibleStringSlice `json:"whitelist,omitempty" yaml:"-" env:"PICOCLAW_TOOLS_WHITELIST"`
|
||||||
|
WhitelistEnabled bool `json:"whitelist_enabled,omitempty" yaml:"-" env:"PICOCLAW_TOOLS_WHITELIST_ENABLED"`
|
||||||
|
MCP MCPConfig `json:"mcp" yaml:"-""`
|
||||||
AppendFile ToolConfig `json:"append_file" yaml:"-" envPrefix:"PICOCLAW_TOOLS_APPEND_FILE_"`
|
AppendFile ToolConfig `json:"append_file" yaml:"-" envPrefix:"PICOCLAW_TOOLS_APPEND_FILE_"`
|
||||||
EditFile ToolConfig `json:"edit_file" yaml:"-" envPrefix:"PICOCLAW_TOOLS_EDIT_FILE_"`
|
EditFile ToolConfig `json:"edit_file" yaml:"-" envPrefix:"PICOCLAW_TOOLS_EDIT_FILE_"`
|
||||||
FindSkills ToolConfig `json:"find_skills" yaml:"-" envPrefix:"PICOCLAW_TOOLS_FIND_SKILLS_"`
|
FindSkills ToolConfig `json:"find_skills" yaml:"-" envPrefix:"PICOCLAW_TOOLS_FIND_SKILLS_"`
|
||||||
|
|
|
||||||
|
|
@ -354,10 +354,11 @@ func DefaultConfig() *Config {
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
Gateway: GatewayConfig{
|
Gateway: GatewayConfig{
|
||||||
Host: "127.0.0.1",
|
Host: "127.0.0.1",
|
||||||
Port: 18790,
|
Port: 18790,
|
||||||
HotReload: false,
|
ChatEnabled: true,
|
||||||
LogLevel: "warn",
|
HotReload: false,
|
||||||
|
LogLevel: "warn",
|
||||||
},
|
},
|
||||||
Tools: ToolsConfig{
|
Tools: ToolsConfig{
|
||||||
FilterSensitiveData: true,
|
FilterSensitiveData: true,
|
||||||
|
|
|
||||||
|
|
@ -155,8 +155,20 @@ func Run(debug bool, homePath, configPath string, allowEmptyStartup bool) error
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
runningServices.HealthServer.SetReloadFunc(reloadTrigger)
|
runningServices.HealthServer.SetReloadFunc(reloadTrigger)
|
||||||
|
runningServices.HealthServer.SetAPIKey(cfg.Gateway.APIKey)
|
||||||
agentLoop.SetReloadFunc(reloadTrigger)
|
agentLoop.SetReloadFunc(reloadTrigger)
|
||||||
|
|
||||||
|
// Setup synchronous /chat endpoint handler
|
||||||
|
if cfg.Gateway.ChatEnabled {
|
||||||
|
runningServices.HealthServer.SetChatFunc(func(ctx context.Context, message, sessionID string) (string, error) {
|
||||||
|
if sessionID == "" {
|
||||||
|
sessionID = "http-chat"
|
||||||
|
}
|
||||||
|
return agentLoop.ProcessDirectWithChannel(ctx, message, sessionID, "http", "chat")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
fmt.Printf("✓ Gateway started on %s:%d\n", cfg.Gateway.Host, cfg.Gateway.Port)
|
fmt.Printf("✓ Gateway started on %s:%d\n", cfg.Gateway.Host, cfg.Gateway.Port)
|
||||||
fmt.Println("Press Ctrl+C to stop")
|
fmt.Println("Press Ctrl+C to stop")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -9,8 +9,20 @@ import (
|
||||||
"os"
|
"os"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// ChatRequest is the JSON body for POST /chat.
|
||||||
|
type ChatRequest struct {
|
||||||
|
Message string `json:"message"`
|
||||||
|
SessionID string `json:"session_id,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// ChatResponse is the JSON response from POST /chat.
|
||||||
|
type ChatResponse struct {
|
||||||
|
Response string `json:"response"`
|
||||||
|
}
|
||||||
|
|
||||||
type Server struct {
|
type Server struct {
|
||||||
server *http.Server
|
server *http.Server
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
|
|
@ -18,6 +30,8 @@ type Server struct {
|
||||||
checks map[string]Check
|
checks map[string]Check
|
||||||
startTime time.Time
|
startTime time.Time
|
||||||
reloadFunc func() error
|
reloadFunc func() error
|
||||||
|
chatFunc func(ctx context.Context, message, sessionID string) (string, error)
|
||||||
|
apiKey string
|
||||||
}
|
}
|
||||||
|
|
||||||
type Check struct {
|
type Check struct {
|
||||||
|
|
@ -45,13 +59,15 @@ func NewServer(host string, port int) *Server {
|
||||||
mux.HandleFunc("/health", s.healthHandler)
|
mux.HandleFunc("/health", s.healthHandler)
|
||||||
mux.HandleFunc("/ready", s.readyHandler)
|
mux.HandleFunc("/ready", s.readyHandler)
|
||||||
mux.HandleFunc("/reload", s.reloadHandler)
|
mux.HandleFunc("/reload", s.reloadHandler)
|
||||||
|
mux.HandleFunc("/chat", s.chatHandler)
|
||||||
|
|
||||||
addr := fmt.Sprintf("%s:%d", host, port)
|
addr := fmt.Sprintf("%s:%d", host, port)
|
||||||
s.server = &http.Server{
|
s.server = &http.Server{
|
||||||
Addr: addr,
|
Addr: addr,
|
||||||
Handler: mux,
|
Handler: mux,
|
||||||
ReadTimeout: 5 * time.Second,
|
ReadTimeout: 10 * time.Second,
|
||||||
WriteTimeout: 5 * time.Second,
|
// WriteTimeout must be long enough for LLM inference; 5 min is generous.
|
||||||
|
WriteTimeout: 5 * time.Minute,
|
||||||
}
|
}
|
||||||
|
|
||||||
return s
|
return s
|
||||||
|
|
@ -115,7 +131,39 @@ func (s *Server) SetReloadFunc(fn func() error) {
|
||||||
s.reloadFunc = fn
|
s.reloadFunc = fn
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetChatFunc sets the callback that processes /chat requests.
|
||||||
|
// fn receives the user message and an optional session ID and must return the
|
||||||
|
// agent's reply (or an error). It is called synchronously inside the HTTP
|
||||||
|
// handler, so the write timeout on the server governs the maximum duration.
|
||||||
|
func (s *Server) SetChatFunc(fn func(ctx context.Context, message, sessionID string) (string, error)) {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
s.chatFunc = fn
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetAPIKey sets the expected X-API-Key header value.
|
||||||
|
func (s *Server) SetAPIKey(key string) {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
s.apiKey = key
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) verifyAPIKey(r *http.Request) bool {
|
||||||
|
s.mu.RLock()
|
||||||
|
defer s.mu.RUnlock()
|
||||||
|
if s.apiKey == "" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return r.Header.Get("X-API-Key") == s.apiKey
|
||||||
|
}
|
||||||
|
|
||||||
func (s *Server) reloadHandler(w http.ResponseWriter, r *http.Request) {
|
func (s *Server) reloadHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if !s.verifyAPIKey(r) {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusUnauthorized)
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{"error": "unauthorized"})
|
||||||
|
return
|
||||||
|
}
|
||||||
if r.Method != http.MethodPost {
|
if r.Method != http.MethodPost {
|
||||||
w.Header().Set("Content-Type", "application/json")
|
w.Header().Set("Content-Type", "application/json")
|
||||||
w.WriteHeader(http.StatusMethodNotAllowed)
|
w.WriteHeader(http.StatusMethodNotAllowed)
|
||||||
|
|
@ -198,12 +246,72 @@ func (s *Server) readyHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// RegisterOnMux registers /health, /ready and /reload handlers onto the given mux.
|
// RegisterOnMux registers /health, /ready, /reload and /chat handlers onto the
|
||||||
// This allows the health endpoints to be served by a shared HTTP server.
|
// given mux. This allows the health endpoints to be served by a shared HTTP server.
|
||||||
func (s *Server) RegisterOnMux(mux *http.ServeMux) {
|
func (s *Server) RegisterOnMux(mux *http.ServeMux) {
|
||||||
mux.HandleFunc("/health", s.healthHandler)
|
mux.HandleFunc("/health", s.healthHandler)
|
||||||
mux.HandleFunc("/ready", s.readyHandler)
|
mux.HandleFunc("/ready", s.readyHandler)
|
||||||
mux.HandleFunc("/reload", s.reloadHandler)
|
mux.HandleFunc("/reload", s.reloadHandler)
|
||||||
|
mux.HandleFunc("/chat", s.chatHandler)
|
||||||
|
mux.HandleFunc("/v1/chat/completions", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
logger.Error("GATEWAY IS HITTING ITSELF FOR LLM CALLS!")
|
||||||
|
http.Error(w, "GATEWAY LOOP DETECTION", http.StatusLoopDetected)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// chatHandler handles POST /chat — a synchronous HTTP chat API.
|
||||||
|
// Request body: {"message": "...", "session_id": "..." (optional)}
|
||||||
|
// Response body: {"response": "..."}
|
||||||
|
func (s *Server) chatHandler(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if !s.verifyAPIKey(r) {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusUnauthorized)
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{"error": "unauthorized"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if r.Method != http.MethodPost {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusMethodNotAllowed)
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{"error": "method not allowed, use POST"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
s.mu.RLock()
|
||||||
|
chatFunc := s.chatFunc
|
||||||
|
s.mu.RUnlock()
|
||||||
|
|
||||||
|
if chatFunc == nil {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusServiceUnavailable)
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{"error": "chat not configured"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var req ChatRequest
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusBadRequest)
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{"error": "invalid JSON: " + err.Error()})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if req.Message == "" {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusBadRequest)
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{"error": "message field is required"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
reply, err := chatFunc(r.Context(), req.Message, req.SessionID)
|
||||||
|
if err != nil {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusInternalServerError)
|
||||||
|
json.NewEncoder(w).Encode(map[string]string{"error": err.Error()})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
json.NewEncoder(w).Encode(ChatResponse{Response: reply})
|
||||||
}
|
}
|
||||||
|
|
||||||
func statusString(ok bool) string {
|
func statusString(ok bool) string {
|
||||||
|
|
|
||||||
|
|
@ -154,7 +154,7 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
|
||||||
}
|
}
|
||||||
return provider, modelID, nil
|
return provider, modelID, nil
|
||||||
|
|
||||||
case "litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
|
case "litellm", "openrouter", "groq", "zhipu", "gemini",
|
||||||
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
|
||||||
"vivgrid", "volcengine", "vllm", "qwen", "qwen-intl", "qwen-international", "dashscope-intl",
|
"vivgrid", "volcengine", "vllm", "qwen", "qwen-intl", "qwen-international", "dashscope-intl",
|
||||||
"qwen-us", "dashscope-us", "mistral", "avian", "longcat", "modelscope", "novita",
|
"qwen-us", "dashscope-us", "mistral", "avian", "longcat", "modelscope", "novita",
|
||||||
|
|
@ -176,6 +176,37 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
|
||||||
cfg.ExtraBody,
|
cfg.ExtraBody,
|
||||||
), modelID, nil
|
), modelID, nil
|
||||||
|
|
||||||
|
case "nvidia":
|
||||||
|
apiBase := cfg.APIBase
|
||||||
|
if apiBase == "" {
|
||||||
|
apiBase = getDefaultAPIBase(protocol)
|
||||||
|
}
|
||||||
|
p := NewHTTPProviderWithMaxTokensFieldAndRequestTimeout(
|
||||||
|
cfg.APIKey(),
|
||||||
|
apiBase,
|
||||||
|
cfg.Proxy,
|
||||||
|
cfg.MaxTokensField,
|
||||||
|
cfg.RequestTimeout,
|
||||||
|
cfg.ExtraBody,
|
||||||
|
)
|
||||||
|
// NVIDIA sometimes prefers api-key header or has issues with Bearer in some environments
|
||||||
|
p.SetUseAzureHeaders(false) // NVIDIA main gateway prefers standard Bearer headers; api-key causes 404s
|
||||||
|
return p, "nvidia/" + modelID, nil
|
||||||
|
|
||||||
|
case "azure-ai", "azure-foundry":
|
||||||
|
// Azure AI Foundry / Studio compatible with OpenAI API format,
|
||||||
|
// but using api-key header instead of Authorization: Bearer.
|
||||||
|
if cfg.APIKey() == "" && cfg.APIBase == "" {
|
||||||
|
return nil, "", fmt.Errorf("api_key or api_base is required for protocol %q", protocol)
|
||||||
|
}
|
||||||
|
return NewAzureAIProvider(
|
||||||
|
cfg.APIKey(),
|
||||||
|
cfg.APIBase,
|
||||||
|
cfg.Proxy,
|
||||||
|
cfg.RequestTimeout,
|
||||||
|
), modelID, nil
|
||||||
|
|
||||||
|
|
||||||
case "minimax":
|
case "minimax":
|
||||||
// Minimax requires reasoning_split: true in the request body
|
// Minimax requires reasoning_split: true in the request body
|
||||||
if cfg.APIKey() == "" && cfg.APIBase == "" {
|
if cfg.APIKey() == "" && cfg.APIBase == "" {
|
||||||
|
|
|
||||||
|
|
@ -44,6 +44,18 @@ func NewHTTPProviderWithMaxTokensFieldAndRequestTimeout(
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func NewAzureAIProvider(apiKey, apiBase, proxy string, requestTimeoutSeconds int) *HTTPProvider {
|
||||||
|
return &HTTPProvider{
|
||||||
|
delegate: openai_compat.NewProvider(
|
||||||
|
apiKey,
|
||||||
|
apiBase,
|
||||||
|
proxy,
|
||||||
|
openai_compat.WithAzureHeaders(),
|
||||||
|
openai_compat.WithRequestTimeout(time.Duration(requestTimeoutSeconds)*time.Second),
|
||||||
|
),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (p *HTTPProvider) Chat(
|
func (p *HTTPProvider) Chat(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
messages []Message,
|
messages []Message,
|
||||||
|
|
@ -71,6 +83,11 @@ func (p *HTTPProvider) GetDefaultModel() string {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (p *HTTPProvider) SetUseAzureHeaders(use bool) {
|
||||||
|
p.delegate.SetUseAzureHeaders(use)
|
||||||
|
}
|
||||||
|
|
||||||
func (p *HTTPProvider) SupportsNativeSearch() bool {
|
func (p *HTTPProvider) SupportsNativeSearch() bool {
|
||||||
return p.delegate.SupportsNativeSearch()
|
return p.delegate.SupportsNativeSearch()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -36,6 +36,7 @@ type Provider struct {
|
||||||
maxTokensField string // Field name for max tokens (e.g., "max_completion_tokens" for o1/glm models)
|
maxTokensField string // Field name for max tokens (e.g., "max_completion_tokens" for o1/glm models)
|
||||||
httpClient *http.Client
|
httpClient *http.Client
|
||||||
extraBody map[string]any // Additional fields to inject into request body
|
extraBody map[string]any // Additional fields to inject into request body
|
||||||
|
useAzureHeaders bool // Use api-key header instead of Authorization: Bearer
|
||||||
}
|
}
|
||||||
|
|
||||||
type Option func(*Provider)
|
type Option func(*Provider)
|
||||||
|
|
@ -62,6 +63,16 @@ func WithExtraBody(extraBody map[string]any) Option {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func WithAzureHeaders() Option {
|
||||||
|
return func(p *Provider) {
|
||||||
|
p.useAzureHeaders = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Provider) SetUseAzureHeaders(use bool) {
|
||||||
|
p.useAzureHeaders = use
|
||||||
|
}
|
||||||
|
|
||||||
func NewProvider(apiKey, apiBase, proxy string, opts ...Option) *Provider {
|
func NewProvider(apiKey, apiBase, proxy string, opts ...Option) *Provider {
|
||||||
p := &Provider{
|
p := &Provider{
|
||||||
apiKey: apiKey,
|
apiKey: apiKey,
|
||||||
|
|
@ -181,7 +192,11 @@ func (p *Provider) Chat(
|
||||||
|
|
||||||
req.Header.Set("Content-Type", "application/json")
|
req.Header.Set("Content-Type", "application/json")
|
||||||
if p.apiKey != "" {
|
if p.apiKey != "" {
|
||||||
req.Header.Set("Authorization", "Bearer "+p.apiKey)
|
if p.useAzureHeaders {
|
||||||
|
req.Header.Set("api-key", p.apiKey)
|
||||||
|
} else {
|
||||||
|
req.Header.Set("Authorization", "Bearer "+p.apiKey)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
resp, err := p.httpClient.Do(req)
|
resp, err := p.httpClient.Do(req)
|
||||||
|
|
@ -227,7 +242,11 @@ func (p *Provider) ChatStream(
|
||||||
req.Header.Set("Content-Type", "application/json")
|
req.Header.Set("Content-Type", "application/json")
|
||||||
req.Header.Set("Accept", "text/event-stream")
|
req.Header.Set("Accept", "text/event-stream")
|
||||||
if p.apiKey != "" {
|
if p.apiKey != "" {
|
||||||
req.Header.Set("Authorization", "Bearer "+p.apiKey)
|
if p.useAzureHeaders {
|
||||||
|
req.Header.Set("api-key", p.apiKey)
|
||||||
|
} else {
|
||||||
|
req.Header.Set("Authorization", "Bearer "+p.apiKey)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Use a client without Timeout for streaming — the http.Client.Timeout covers
|
// Use a client without Timeout for streaming — the http.Client.Timeout covers
|
||||||
|
|
@ -387,19 +406,30 @@ func parseStreamResponse(
|
||||||
}
|
}
|
||||||
|
|
||||||
func normalizeModel(model, apiBase string) string {
|
func normalizeModel(model, apiBase string) string {
|
||||||
|
if strings.Contains(strings.ToLower(apiBase), "openrouter.ai") {
|
||||||
|
return model
|
||||||
|
}
|
||||||
|
|
||||||
|
// NVIDIA endpoints (integrate.api.nvidia.com) require the provider prefix
|
||||||
|
// (e.g., nvidia/, meta/, mistral/) for routing. Do not strip them.
|
||||||
|
// We also re-add the prefix if it was likely stripped by the agent's protocol resolution logic.
|
||||||
|
if strings.Contains(strings.ToLower(apiBase), ".nvidia.com") {
|
||||||
|
if !strings.Contains(model, "/") {
|
||||||
|
return "nvidia/" + model
|
||||||
|
}
|
||||||
|
return model
|
||||||
|
}
|
||||||
|
|
||||||
before, after, ok := strings.Cut(model, "/")
|
before, after, ok := strings.Cut(model, "/")
|
||||||
if !ok {
|
if !ok {
|
||||||
return model
|
return model
|
||||||
}
|
}
|
||||||
|
|
||||||
if strings.Contains(strings.ToLower(apiBase), "openrouter.ai") {
|
|
||||||
return model
|
|
||||||
}
|
|
||||||
|
|
||||||
prefix := strings.ToLower(before)
|
prefix := strings.ToLower(before)
|
||||||
switch prefix {
|
switch prefix {
|
||||||
case "litellm", "moonshot", "nvidia", "groq", "ollama", "deepseek", "google",
|
case "litellm", "moonshot", "groq", "ollama", "deepseek", "google",
|
||||||
"openrouter", "zhipu", "mistral", "vivgrid", "minimax", "novita":
|
"openrouter", "zhipu", "mistral", "vivgrid", "minimax", "novita",
|
||||||
|
"azure-ai", "azure-foundry":
|
||||||
return after
|
return after
|
||||||
default:
|
default:
|
||||||
return model
|
return model
|
||||||
|
|
@ -430,7 +460,7 @@ func isNativeSearchHost(apiBase string) bool {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
host := u.Hostname()
|
host := u.Hostname()
|
||||||
return host == "api.openai.com" || strings.HasSuffix(host, ".openai.azure.com")
|
return host == "api.openai.com"
|
||||||
}
|
}
|
||||||
|
|
||||||
// supportsPromptCacheKey reports whether the given API base is known to
|
// supportsPromptCacheKey reports whether the given API base is known to
|
||||||
|
|
@ -443,5 +473,7 @@ func supportsPromptCacheKey(apiBase string) bool {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
host := u.Hostname()
|
host := u.Hostname()
|
||||||
return host == "api.openai.com" || strings.HasSuffix(host, ".openai.azure.com")
|
// Strictly limit to OpenAI official. Azure OpenAI often rejects this field
|
||||||
|
// depending on model version and region, causing 400 errors.
|
||||||
|
return host == "api.openai.com"
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -907,8 +907,8 @@ func TestSupportsPromptCacheKey(t *testing.T) {
|
||||||
}{
|
}{
|
||||||
{"https://api.openai.com/v1", true},
|
{"https://api.openai.com/v1", true},
|
||||||
{"https://api.openai.com/v1/", true},
|
{"https://api.openai.com/v1/", true},
|
||||||
{"https://myresource.openai.azure.com/openai/deployments/gpt-4", true},
|
{"https://myresource.openai.azure.com/openai/deployments/gpt-4", false},
|
||||||
{"https://eastus.openai.azure.com/v1", true},
|
{"https://eastus.openai.azure.com/v1", false},
|
||||||
{"https://api.mistral.ai/v1", false},
|
{"https://api.mistral.ai/v1", false},
|
||||||
{"https://generativelanguage.googleapis.com/v1beta", false},
|
{"https://generativelanguage.googleapis.com/v1beta", false},
|
||||||
{"https://api.deepseek.com/v1", false},
|
{"https://api.deepseek.com/v1", false},
|
||||||
|
|
@ -979,7 +979,7 @@ func TestIsNativeSearchHost(t *testing.T) {
|
||||||
want bool
|
want bool
|
||||||
}{
|
}{
|
||||||
{"https://api.openai.com/v1", true},
|
{"https://api.openai.com/v1", true},
|
||||||
{"https://myresource.openai.azure.com/openai/deployments/gpt-4", true},
|
{"https://myresource.openai.azure.com/openai/deployments/gpt-4", false},
|
||||||
{"https://api.mistral.ai/v1", false},
|
{"https://api.mistral.ai/v1", false},
|
||||||
{"https://api.deepseek.com/v1", false},
|
{"https://api.deepseek.com/v1", false},
|
||||||
{"https://api.groq.com/openai/v1", false},
|
{"https://api.groq.com/openai/v1", false},
|
||||||
|
|
|
||||||
|
|
@ -63,6 +63,8 @@ type SkillsLoader struct {
|
||||||
workspaceSkills string // workspace skills (project-level)
|
workspaceSkills string // workspace skills (project-level)
|
||||||
globalSkills string // global skills (~/.picoclaw/skills)
|
globalSkills string // global skills (~/.picoclaw/skills)
|
||||||
builtinSkills string // builtin skills
|
builtinSkills string // builtin skills
|
||||||
|
whitelist []string
|
||||||
|
whitelistEnabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// SkillRoots returns all unique skill root directories used by this loader.
|
// SkillRoots returns all unique skill root directories used by this loader.
|
||||||
|
|
@ -88,12 +90,14 @@ func (sl *SkillsLoader) SkillRoots() []string {
|
||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewSkillsLoader(workspace string, globalSkills string, builtinSkills string) *SkillsLoader {
|
func NewSkillsLoader(workspace string, globalSkills string, builtinSkills string, whitelist []string, whitelistEnabled bool) *SkillsLoader {
|
||||||
return &SkillsLoader{
|
return &SkillsLoader{
|
||||||
workspace: workspace,
|
workspace: workspace,
|
||||||
workspaceSkills: filepath.Join(workspace, "skills"),
|
workspaceSkills: filepath.Join(workspace, "skills"),
|
||||||
globalSkills: globalSkills, // ~/.picoclaw/skills
|
globalSkills: globalSkills, // ~/.picoclaw/skills
|
||||||
builtinSkills: builtinSkills,
|
builtinSkills: builtinSkills,
|
||||||
|
whitelist: whitelist,
|
||||||
|
whitelistEnabled: whitelistEnabled,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -101,6 +105,18 @@ func (sl *SkillsLoader) ListSkills() []SkillInfo {
|
||||||
skills := make([]SkillInfo, 0)
|
skills := make([]SkillInfo, 0)
|
||||||
seen := make(map[string]bool)
|
seen := make(map[string]bool)
|
||||||
|
|
||||||
|
isWhitelisted := func(name string) bool {
|
||||||
|
if !sl.whitelistEnabled {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
for _, w := range sl.whitelist {
|
||||||
|
if w == name {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
addSkills := func(dir, source string) {
|
addSkills := func(dir, source string) {
|
||||||
if dir == "" {
|
if dir == "" {
|
||||||
return
|
return
|
||||||
|
|
@ -113,6 +129,12 @@ func (sl *SkillsLoader) ListSkills() []SkillInfo {
|
||||||
if !d.IsDir() {
|
if !d.IsDir() {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// First check if whitelisted before doing more expensive operations.
|
||||||
|
if !isWhitelisted(d.Name()) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
skillFile := filepath.Join(dir, d.Name(), "SKILL.md")
|
skillFile := filepath.Join(dir, d.Name(), "SKILL.md")
|
||||||
if _, err := os.Stat(skillFile); err != nil {
|
if _, err := os.Stat(skillFile); err != nil {
|
||||||
continue
|
continue
|
||||||
|
|
@ -127,6 +149,12 @@ func (sl *SkillsLoader) ListSkills() []SkillInfo {
|
||||||
info.Description = metadata.Description
|
info.Description = metadata.Description
|
||||||
info.Name = metadata.Name
|
info.Name = metadata.Name
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Double check whitelisted name if metadata name is different from directory name
|
||||||
|
if info.Name != d.Name() && !isWhitelisted(info.Name) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
if err := info.validate(); err != nil {
|
if err := info.validate(); err != nil {
|
||||||
slog.Warn("invalid skill from "+source, "name", info.Name, "error", err)
|
slog.Warn("invalid skill from "+source, "name", info.Name, "error", err)
|
||||||
continue
|
continue
|
||||||
|
|
@ -148,6 +176,19 @@ func (sl *SkillsLoader) ListSkills() []SkillInfo {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (sl *SkillsLoader) LoadSkill(name string) (string, bool) {
|
func (sl *SkillsLoader) LoadSkill(name string) (string, bool) {
|
||||||
|
if sl.whitelistEnabled {
|
||||||
|
whitelisted := false
|
||||||
|
for _, w := range sl.whitelist {
|
||||||
|
if w == name {
|
||||||
|
whitelisted = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !whitelisted {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// 1. load from workspace skills first (project-level)
|
// 1. load from workspace skills first (project-level)
|
||||||
if sl.workspaceSkills != "" {
|
if sl.workspaceSkills != "" {
|
||||||
skillFile := filepath.Join(sl.workspaceSkills, name, "SKILL.md")
|
skillFile := filepath.Join(sl.workspaceSkills, name, "SKILL.md")
|
||||||
|
|
@ -155,6 +196,7 @@ func (sl *SkillsLoader) LoadSkill(name string) (string, bool) {
|
||||||
return sl.stripFrontmatter(string(content)), true
|
return sl.stripFrontmatter(string(content)), true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
// ...
|
||||||
|
|
||||||
// 2. then load from global skills (~/.picoclaw/skills)
|
// 2. then load from global skills (~/.picoclaw/skills)
|
||||||
if sl.globalSkills != "" {
|
if sl.globalSkills != "" {
|
||||||
|
|
@ -204,11 +246,11 @@ func (sl *SkillsLoader) BuildSkillsSummary() string {
|
||||||
escapedDesc := escapeXML(s.Description)
|
escapedDesc := escapeXML(s.Description)
|
||||||
escapedPath := escapeXML(s.Path)
|
escapedPath := escapeXML(s.Path)
|
||||||
|
|
||||||
lines = append(lines, fmt.Sprintf(" <skill>"))
|
lines = append(lines, " <skill>")
|
||||||
lines = append(lines, fmt.Sprintf(" <name>%s</name>", escapedName))
|
lines = append(lines, " <name>"+escapedName+"</name>")
|
||||||
lines = append(lines, fmt.Sprintf(" <description>%s</description>", escapedDesc))
|
lines = append(lines, " <description>"+escapedDesc+"</description>")
|
||||||
lines = append(lines, fmt.Sprintf(" <location>%s</location>", escapedPath))
|
lines = append(lines, " <location>"+escapedPath+"</location>")
|
||||||
lines = append(lines, fmt.Sprintf(" <source>%s</source>", s.Source))
|
lines = append(lines, " <source>"+s.Source+"</source>")
|
||||||
lines = append(lines, " </skill>")
|
lines = append(lines, " </skill>")
|
||||||
}
|
}
|
||||||
lines = append(lines, "</skills>")
|
lines = append(lines, "</skills>")
|
||||||
|
|
|
||||||
|
|
@ -155,7 +155,7 @@ func TestListSkillsWorkspaceOverridesGlobal(t *testing.T) {
|
||||||
createSkillDir(t, filepath.Join(ws, "skills"), "my-skill", "my-skill", "workspace version")
|
createSkillDir(t, filepath.Join(ws, "skills"), "my-skill", "my-skill", "workspace version")
|
||||||
createSkillDir(t, global, "my-skill", "my-skill", "global version")
|
createSkillDir(t, global, "my-skill", "my-skill", "global version")
|
||||||
|
|
||||||
sl := NewSkillsLoader(ws, global, "")
|
sl := NewSkillsLoader(ws, global, "", nil, false)
|
||||||
skills := sl.ListSkills()
|
skills := sl.ListSkills()
|
||||||
|
|
||||||
assert.Len(t, skills, 1)
|
assert.Len(t, skills, 1)
|
||||||
|
|
@ -172,7 +172,7 @@ func TestListSkillsGlobalOverridesBuiltin(t *testing.T) {
|
||||||
createSkillDir(t, global, "my-skill", "my-skill", "global version")
|
createSkillDir(t, global, "my-skill", "my-skill", "global version")
|
||||||
createSkillDir(t, builtin, "my-skill", "my-skill", "builtin version")
|
createSkillDir(t, builtin, "my-skill", "my-skill", "builtin version")
|
||||||
|
|
||||||
sl := NewSkillsLoader(ws, global, builtin)
|
sl := NewSkillsLoader(ws, global, builtin, nil, false)
|
||||||
skills := sl.ListSkills()
|
skills := sl.ListSkills()
|
||||||
|
|
||||||
assert.Len(t, skills, 1)
|
assert.Len(t, skills, 1)
|
||||||
|
|
@ -189,7 +189,7 @@ func TestListSkillsMetadataNameDedup(t *testing.T) {
|
||||||
createSkillDir(t, filepath.Join(ws, "skills"), "dir-a", "shared-name", "workspace version")
|
createSkillDir(t, filepath.Join(ws, "skills"), "dir-a", "shared-name", "workspace version")
|
||||||
createSkillDir(t, global, "dir-b", "shared-name", "global version")
|
createSkillDir(t, global, "dir-b", "shared-name", "global version")
|
||||||
|
|
||||||
sl := NewSkillsLoader(ws, global, "")
|
sl := NewSkillsLoader(ws, global, "", nil, false)
|
||||||
skills := sl.ListSkills()
|
skills := sl.ListSkills()
|
||||||
|
|
||||||
assert.Len(t, skills, 1)
|
assert.Len(t, skills, 1)
|
||||||
|
|
@ -207,7 +207,7 @@ func TestListSkillsMultipleDistinctSkills(t *testing.T) {
|
||||||
createSkillDir(t, global, "skill-b", "skill-b", "desc b")
|
createSkillDir(t, global, "skill-b", "skill-b", "desc b")
|
||||||
createSkillDir(t, builtin, "skill-c", "skill-c", "desc c")
|
createSkillDir(t, builtin, "skill-c", "skill-c", "desc c")
|
||||||
|
|
||||||
sl := NewSkillsLoader(ws, global, builtin)
|
sl := NewSkillsLoader(ws, global, builtin, nil, false)
|
||||||
skills := sl.ListSkills()
|
skills := sl.ListSkills()
|
||||||
|
|
||||||
assert.Len(t, skills, 3)
|
assert.Len(t, skills, 3)
|
||||||
|
|
@ -230,7 +230,7 @@ func TestListSkillsInvalidSkillSkipped(t *testing.T) {
|
||||||
// Valid skill
|
// Valid skill
|
||||||
createSkillDir(t, global, "good-skill", "good-skill", "desc")
|
createSkillDir(t, global, "good-skill", "good-skill", "desc")
|
||||||
|
|
||||||
sl := NewSkillsLoader(ws, global, "")
|
sl := NewSkillsLoader(ws, global, "", nil, false)
|
||||||
skills := sl.ListSkills()
|
skills := sl.ListSkills()
|
||||||
|
|
||||||
assert.Len(t, skills, 1)
|
assert.Len(t, skills, 1)
|
||||||
|
|
@ -243,7 +243,7 @@ func TestListSkillsEmptyAndNonexistentDirs(t *testing.T) {
|
||||||
emptyDir := filepath.Join(tmp, "empty")
|
emptyDir := filepath.Join(tmp, "empty")
|
||||||
require.NoError(t, os.MkdirAll(emptyDir, 0o755))
|
require.NoError(t, os.MkdirAll(emptyDir, 0o755))
|
||||||
|
|
||||||
sl := NewSkillsLoader(ws, emptyDir, filepath.Join(tmp, "nonexistent"))
|
sl := NewSkillsLoader(ws, emptyDir, filepath.Join(tmp, "nonexistent"), nil, false)
|
||||||
skills := sl.ListSkills()
|
skills := sl.ListSkills()
|
||||||
|
|
||||||
assert.Empty(t, skills)
|
assert.Empty(t, skills)
|
||||||
|
|
@ -259,7 +259,7 @@ func TestListSkillsDirWithoutSkillMD(t *testing.T) {
|
||||||
// Valid skill alongside
|
// Valid skill alongside
|
||||||
createSkillDir(t, global, "real-skill", "real-skill", "desc")
|
createSkillDir(t, global, "real-skill", "real-skill", "desc")
|
||||||
|
|
||||||
sl := NewSkillsLoader(ws, global, "")
|
sl := NewSkillsLoader(ws, global, "", nil, false)
|
||||||
skills := sl.ListSkills()
|
skills := sl.ListSkills()
|
||||||
|
|
||||||
assert.Len(t, skills, 1)
|
assert.Len(t, skills, 1)
|
||||||
|
|
@ -333,7 +333,7 @@ func TestSkillRootsTrimsWhitespaceAndDedups(t *testing.T) {
|
||||||
global := filepath.Join(tmp, "global")
|
global := filepath.Join(tmp, "global")
|
||||||
builtin := filepath.Join(tmp, "builtin")
|
builtin := filepath.Join(tmp, "builtin")
|
||||||
|
|
||||||
sl := NewSkillsLoader(workspace, " "+global+" ", "\t"+builtin+"\n")
|
sl := NewSkillsLoader(workspace, " "+global+" ", "\t"+builtin+"\n", nil, false)
|
||||||
roots := sl.SkillRoots()
|
roots := sl.SkillRoots()
|
||||||
|
|
||||||
assert.Equal(t, []string{
|
assert.Equal(t, []string{
|
||||||
|
|
@ -417,3 +417,47 @@ func TestGetSkillMetadata_IgnoresHTMLCommentBlocks(t *testing.T) {
|
||||||
assert.Equal(t, "biomed-skill", meta.Name)
|
assert.Equal(t, "biomed-skill", meta.Name)
|
||||||
assert.Equal(t, "Summarize biomedical papers.", meta.Description)
|
assert.Equal(t, "Summarize biomedical papers.", meta.Description)
|
||||||
}
|
}
|
||||||
|
func TestListSkillsWithWhitelist(t *testing.T) {
|
||||||
|
tmp := t.TempDir()
|
||||||
|
ws := filepath.Join(tmp, "workspace")
|
||||||
|
global := filepath.Join(tmp, "global")
|
||||||
|
builtin := filepath.Join(tmp, "builtin")
|
||||||
|
|
||||||
|
createSkillDir(t, filepath.Join(ws, "skills"), "skill-a", "skill-a", "desc a")
|
||||||
|
createSkillDir(t, global, "skill-b", "skill-b", "desc b")
|
||||||
|
createSkillDir(t, builtin, "skill-c", "skill-c", "desc c")
|
||||||
|
|
||||||
|
t.Run("allow-one", func(t *testing.T) {
|
||||||
|
sl := NewSkillsLoader(ws, global, builtin, []string{"skill-a"}, true)
|
||||||
|
skills := sl.ListSkills()
|
||||||
|
assert.Len(t, skills, 1)
|
||||||
|
assert.Equal(t, "skill-a", skills[0].Name)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("allow-two", func(t *testing.T) {
|
||||||
|
sl := NewSkillsLoader(ws, global, builtin, []string{"skill-a", "skill-c"}, true)
|
||||||
|
skills := sl.ListSkills()
|
||||||
|
assert.Len(t, skills, 2)
|
||||||
|
names := []string{skills[0].Name, skills[1].Name}
|
||||||
|
assert.Contains(t, names, "skill-a")
|
||||||
|
assert.Contains(t, names, "skill-c")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("allow-none", func(t *testing.T) {
|
||||||
|
sl := NewSkillsLoader(ws, global, builtin, []string{"non-existent"}, true)
|
||||||
|
skills := sl.ListSkills()
|
||||||
|
assert.Empty(t, skills)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("empty-whitelist-allows-all", func(t *testing.T) {
|
||||||
|
sl := NewSkillsLoader(ws, global, builtin, []string{}, false)
|
||||||
|
skills := sl.ListSkills()
|
||||||
|
assert.Len(t, skills, 3)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("nil-whitelist-allows-all", func(t *testing.T) {
|
||||||
|
sl := NewSkillsLoader(ws, global, builtin, nil, false)
|
||||||
|
skills := sl.ListSkills()
|
||||||
|
assert.Len(t, skills, 3)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -422,21 +422,32 @@ func (r *ToolRegistry) GetSummaries() []string {
|
||||||
return summaries
|
return summaries
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetAll returns all registered tools (both core and non-core with TTL > 0).
|
// Filter removes tools that are not in the whitelist.
|
||||||
// Used by SubTurn to inherit parent's tool set.
|
// If enabled is false, it does nothing.
|
||||||
func (r *ToolRegistry) GetAll() []Tool {
|
func (r *ToolRegistry) Filter(whitelist []string, enabled bool) {
|
||||||
r.mu.RLock()
|
if !enabled {
|
||||||
defer r.mu.RUnlock()
|
return
|
||||||
|
}
|
||||||
|
|
||||||
sorted := r.sortedToolNames()
|
r.mu.Lock()
|
||||||
tools := make([]Tool, 0, len(sorted))
|
defer r.mu.Unlock()
|
||||||
for _, name := range sorted {
|
|
||||||
entry := r.tools[name]
|
|
||||||
|
|
||||||
// Include core tools and non-core tools with active TTL
|
whitelistMap := make(map[string]struct{}, len(whitelist))
|
||||||
if entry.IsCore || entry.TTL > 0 {
|
for _, name := range whitelist {
|
||||||
tools = append(tools, entry.Tool)
|
whitelistMap[name] = struct{}{}
|
||||||
|
}
|
||||||
|
|
||||||
|
removed := 0
|
||||||
|
for name := range r.tools {
|
||||||
|
if _, allowed := whitelistMap[name]; !allowed {
|
||||||
|
delete(r.tools, name)
|
||||||
|
removed++
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return tools
|
|
||||||
|
if removed > 0 {
|
||||||
|
r.version.Add(1)
|
||||||
|
logger.InfoCF("tools", "Filtered tools based on whitelist",
|
||||||
|
map[string]any{"removed": removed, "remaining": len(r.tools)})
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -15,22 +15,23 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/utils"
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
)
|
)
|
||||||
|
|
||||||
// InstallSkillTool allows the LLM agent to install skills from registries.
|
|
||||||
// It shares the same RegistryManager that FindSkillsTool uses,
|
|
||||||
// so all registries configured in config are available for installation.
|
|
||||||
type InstallSkillTool struct {
|
type InstallSkillTool struct {
|
||||||
registryMgr *skills.RegistryManager
|
registryMgr *skills.RegistryManager
|
||||||
workspace string
|
workspace string
|
||||||
mu sync.Mutex
|
whitelist []string
|
||||||
|
whitelistEnabled bool
|
||||||
|
mu sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewInstallSkillTool creates a new InstallSkillTool.
|
// NewInstallSkillTool creates a new InstallSkillTool.
|
||||||
// registryMgr is the shared registry manager (same instance as FindSkillsTool).
|
// registryMgr is the shared registry manager (same instance as FindSkillsTool).
|
||||||
// workspace is the root workspace directory; skills install to {workspace}/skills/{slug}/.
|
// workspace is the root workspace directory; skills install to {workspace}/skills/{slug}/.
|
||||||
func NewInstallSkillTool(registryMgr *skills.RegistryManager, workspace string) *InstallSkillTool {
|
func NewInstallSkillTool(registryMgr *skills.RegistryManager, workspace string, whitelist []string, whitelistEnabled bool) *InstallSkillTool {
|
||||||
return &InstallSkillTool{
|
return &InstallSkillTool{
|
||||||
registryMgr: registryMgr,
|
registryMgr: registryMgr,
|
||||||
workspace: workspace,
|
workspace: workspace,
|
||||||
|
whitelist: whitelist,
|
||||||
|
whitelistEnabled: whitelistEnabled,
|
||||||
mu: sync.Mutex{},
|
mu: sync.Mutex{},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -80,6 +81,20 @@ func (t *InstallSkillTool) Execute(ctx context.Context, args map[string]any) *To
|
||||||
return ErrorResult(fmt.Sprintf("invalid slug %q: error: %s", slug, err.Error()))
|
return ErrorResult(fmt.Sprintf("invalid slug %q: error: %s", slug, err.Error()))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Check whitelist
|
||||||
|
if t.whitelistEnabled {
|
||||||
|
whitelisted := false
|
||||||
|
for _, w := range t.whitelist {
|
||||||
|
if w == slug {
|
||||||
|
whitelisted = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !whitelisted {
|
||||||
|
return ErrorResult(fmt.Sprintf("skill %q is not in whitelist and cannot be installed", slug))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Validate registry
|
// Validate registry
|
||||||
registryName, _ := args["registry"].(string)
|
registryName, _ := args["registry"].(string)
|
||||||
if err := utils.ValidateSkillIdentifier(registryName); err != nil {
|
if err := utils.ValidateSkillIdentifier(registryName); err != nil {
|
||||||
|
|
|
||||||
|
|
@ -13,19 +13,19 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestInstallSkillToolName(t *testing.T) {
|
func TestInstallSkillToolName(t *testing.T) {
|
||||||
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir())
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir(), nil, false)
|
||||||
assert.Equal(t, "install_skill", tool.Name())
|
assert.Equal(t, "install_skill", tool.Name())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestInstallSkillToolMissingSlug(t *testing.T) {
|
func TestInstallSkillToolMissingSlug(t *testing.T) {
|
||||||
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir())
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir(), nil, false)
|
||||||
result := tool.Execute(context.Background(), map[string]any{})
|
result := tool.Execute(context.Background(), map[string]any{})
|
||||||
assert.True(t, result.IsError)
|
assert.True(t, result.IsError)
|
||||||
assert.Contains(t, result.ForLLM, "identifier is required and must be a non-empty string")
|
assert.Contains(t, result.ForLLM, "identifier is required and must be a non-empty string")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestInstallSkillToolEmptySlug(t *testing.T) {
|
func TestInstallSkillToolEmptySlug(t *testing.T) {
|
||||||
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir())
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir(), nil, false)
|
||||||
result := tool.Execute(context.Background(), map[string]any{
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
"slug": " ",
|
"slug": " ",
|
||||||
})
|
})
|
||||||
|
|
@ -34,7 +34,7 @@ func TestInstallSkillToolEmptySlug(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestInstallSkillToolUnsafeSlug(t *testing.T) {
|
func TestInstallSkillToolUnsafeSlug(t *testing.T) {
|
||||||
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir())
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir(), nil, false)
|
||||||
|
|
||||||
cases := []string{
|
cases := []string{
|
||||||
"../etc/passwd",
|
"../etc/passwd",
|
||||||
|
|
@ -56,7 +56,7 @@ func TestInstallSkillToolAlreadyExists(t *testing.T) {
|
||||||
skillDir := filepath.Join(workspace, "skills", "existing-skill")
|
skillDir := filepath.Join(workspace, "skills", "existing-skill")
|
||||||
require.NoError(t, os.MkdirAll(skillDir, 0o755))
|
require.NoError(t, os.MkdirAll(skillDir, 0o755))
|
||||||
|
|
||||||
tool := NewInstallSkillTool(skills.NewRegistryManager(), workspace)
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), workspace, nil, false)
|
||||||
result := tool.Execute(context.Background(), map[string]any{
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
"slug": "existing-skill",
|
"slug": "existing-skill",
|
||||||
"registry": "clawhub",
|
"registry": "clawhub",
|
||||||
|
|
@ -67,7 +67,7 @@ func TestInstallSkillToolAlreadyExists(t *testing.T) {
|
||||||
|
|
||||||
func TestInstallSkillToolRegistryNotFound(t *testing.T) {
|
func TestInstallSkillToolRegistryNotFound(t *testing.T) {
|
||||||
workspace := t.TempDir()
|
workspace := t.TempDir()
|
||||||
tool := NewInstallSkillTool(skills.NewRegistryManager(), workspace)
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), workspace, nil, false)
|
||||||
result := tool.Execute(context.Background(), map[string]any{
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
"slug": "some-skill",
|
"slug": "some-skill",
|
||||||
"registry": "nonexistent",
|
"registry": "nonexistent",
|
||||||
|
|
@ -78,7 +78,7 @@ func TestInstallSkillToolRegistryNotFound(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestInstallSkillToolParameters(t *testing.T) {
|
func TestInstallSkillToolParameters(t *testing.T) {
|
||||||
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir())
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir(), nil, false)
|
||||||
params := tool.Parameters()
|
params := tool.Parameters()
|
||||||
|
|
||||||
props, ok := params["properties"].(map[string]any)
|
props, ok := params["properties"].(map[string]any)
|
||||||
|
|
@ -95,10 +95,55 @@ func TestInstallSkillToolParameters(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestInstallSkillToolMissingRegistry(t *testing.T) {
|
func TestInstallSkillToolMissingRegistry(t *testing.T) {
|
||||||
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir())
|
tool := NewInstallSkillTool(skills.NewRegistryManager(), t.TempDir(), nil, false)
|
||||||
result := tool.Execute(context.Background(), map[string]any{
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
"slug": "some-skill",
|
"slug": "some-skill",
|
||||||
})
|
})
|
||||||
assert.True(t, result.IsError)
|
assert.True(t, result.IsError)
|
||||||
assert.Contains(t, result.ForLLM, "invalid registry")
|
assert.Contains(t, result.ForLLM, "invalid registry")
|
||||||
}
|
}
|
||||||
|
func TestInstallSkillToolWhitelist(t *testing.T) {
|
||||||
|
workspace := t.TempDir()
|
||||||
|
rm := skills.NewRegistryManager()
|
||||||
|
|
||||||
|
t.Run("blocked-by-whitelist", func(t *testing.T) {
|
||||||
|
tool := NewInstallSkillTool(rm, workspace, []string{"allowed-skill"}, true)
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"slug": "blocked-skill",
|
||||||
|
"registry": "clawhub",
|
||||||
|
})
|
||||||
|
assert.True(t, result.IsError)
|
||||||
|
assert.Contains(t, result.ForLLM, "not in whitelist")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("allowed-by-whitelist", func(t *testing.T) {
|
||||||
|
// This will still fail because registry is not found, but it should pass the whitelist check
|
||||||
|
tool := NewInstallSkillTool(rm, workspace, []string{"allowed-skill"}, true)
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"slug": "allowed-skill",
|
||||||
|
"registry": "clawhub",
|
||||||
|
})
|
||||||
|
assert.True(t, result.IsError)
|
||||||
|
assert.NotContains(t, result.ForLLM, "not in whitelist")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("empty-whitelist-allows-all", func(t *testing.T) {
|
||||||
|
tool := NewInstallSkillTool(rm, workspace, []string{}, false)
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"slug": "any-skill",
|
||||||
|
"registry": "clawhub",
|
||||||
|
})
|
||||||
|
assert.True(t, result.IsError)
|
||||||
|
assert.NotContains(t, result.ForLLM, "not in whitelist")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("nil-whitelist-allows-all", func(t *testing.T) {
|
||||||
|
tool := NewInstallSkillTool(rm, workspace, nil, false)
|
||||||
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
|
"slug": "any-skill",
|
||||||
|
"registry": "clawhub",
|
||||||
|
})
|
||||||
|
assert.True(t, result.IsError)
|
||||||
|
assert.NotContains(t, result.ForLLM, "not in whitelist")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -12,15 +12,19 @@ import (
|
||||||
type FindSkillsTool struct {
|
type FindSkillsTool struct {
|
||||||
registryMgr *skills.RegistryManager
|
registryMgr *skills.RegistryManager
|
||||||
cache *skills.SearchCache
|
cache *skills.SearchCache
|
||||||
|
whitelist []string
|
||||||
|
enabled bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewFindSkillsTool creates a new FindSkillsTool.
|
// NewFindSkillsTool creates a new FindSkillsTool.
|
||||||
// registryMgr is the shared registry manager (built from config in createToolRegistry).
|
// registryMgr is the shared registry manager (built from config in createToolRegistry).
|
||||||
// cache is the search cache for deduplicating similar queries.
|
// cache is the search cache for deduplicating similar queries.
|
||||||
func NewFindSkillsTool(registryMgr *skills.RegistryManager, cache *skills.SearchCache) *FindSkillsTool {
|
func NewFindSkillsTool(registryMgr *skills.RegistryManager, cache *skills.SearchCache, whitelist []string, enabled bool) *FindSkillsTool {
|
||||||
return &FindSkillsTool{
|
return &FindSkillsTool{
|
||||||
registryMgr: registryMgr,
|
registryMgr: registryMgr,
|
||||||
cache: cache,
|
cache: cache,
|
||||||
|
whitelist: whitelist,
|
||||||
|
enabled: enabled,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -79,6 +83,22 @@ func (t *FindSkillsTool) Execute(ctx context.Context, args map[string]any) *Tool
|
||||||
return ErrorResult(fmt.Sprintf("skill search failed: %v", err))
|
return ErrorResult(fmt.Sprintf("skill search failed: %v", err))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Filter by whitelist if enabled
|
||||||
|
if t.enabled {
|
||||||
|
filtered := make([]skills.SearchResult, 0, len(results))
|
||||||
|
whitelistMap := make(map[string]struct{}, len(t.whitelist))
|
||||||
|
for _, w := range t.whitelist {
|
||||||
|
whitelistMap[w] = struct{}{}
|
||||||
|
}
|
||||||
|
for _, r := range results {
|
||||||
|
if _, ok := whitelistMap[r.Slug]; ok {
|
||||||
|
filtered = append(filtered, r)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
results = filtered
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
// Cache the results.
|
// Cache the results.
|
||||||
if t.cache != nil && len(results) > 0 {
|
if t.cache != nil && len(results) > 0 {
|
||||||
t.cache.Put(query, results)
|
t.cache.Put(query, results)
|
||||||
|
|
|
||||||
|
|
@ -10,19 +10,19 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestFindSkillsToolName(t *testing.T) {
|
func TestFindSkillsToolName(t *testing.T) {
|
||||||
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil)
|
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil, nil, false)
|
||||||
assert.Equal(t, "find_skills", tool.Name())
|
assert.Equal(t, "find_skills", tool.Name())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFindSkillsToolMissingQuery(t *testing.T) {
|
func TestFindSkillsToolMissingQuery(t *testing.T) {
|
||||||
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil)
|
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil, nil, false)
|
||||||
result := tool.Execute(context.Background(), map[string]any{})
|
result := tool.Execute(context.Background(), map[string]any{})
|
||||||
assert.True(t, result.IsError)
|
assert.True(t, result.IsError)
|
||||||
assert.Contains(t, result.ForLLM, "query is required")
|
assert.Contains(t, result.ForLLM, "query is required")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFindSkillsToolEmptyQuery(t *testing.T) {
|
func TestFindSkillsToolEmptyQuery(t *testing.T) {
|
||||||
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil)
|
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil, nil, false)
|
||||||
result := tool.Execute(context.Background(), map[string]any{
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
"query": " ",
|
"query": " ",
|
||||||
})
|
})
|
||||||
|
|
@ -35,7 +35,7 @@ func TestFindSkillsToolCacheHit(t *testing.T) {
|
||||||
{Slug: "github", Score: 0.9, RegistryName: "clawhub"},
|
{Slug: "github", Score: 0.9, RegistryName: "clawhub"},
|
||||||
})
|
})
|
||||||
|
|
||||||
tool := NewFindSkillsTool(skills.NewRegistryManager(), cache)
|
tool := NewFindSkillsTool(skills.NewRegistryManager(), cache, nil, false)
|
||||||
result := tool.Execute(context.Background(), map[string]any{
|
result := tool.Execute(context.Background(), map[string]any{
|
||||||
"query": "github",
|
"query": "github",
|
||||||
})
|
})
|
||||||
|
|
@ -46,7 +46,7 @@ func TestFindSkillsToolCacheHit(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFindSkillsToolParameters(t *testing.T) {
|
func TestFindSkillsToolParameters(t *testing.T) {
|
||||||
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil)
|
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil, nil, false)
|
||||||
params := tool.Parameters()
|
params := tool.Parameters()
|
||||||
|
|
||||||
props, ok := params["properties"].(map[string]any)
|
props, ok := params["properties"].(map[string]any)
|
||||||
|
|
@ -60,7 +60,7 @@ func TestFindSkillsToolParameters(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFindSkillsToolDescription(t *testing.T) {
|
func TestFindSkillsToolDescription(t *testing.T) {
|
||||||
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil)
|
tool := NewFindSkillsTool(skills.NewRegistryManager(), nil, nil, false)
|
||||||
assert.NotEmpty(t, tool.Description())
|
assert.NotEmpty(t, tool.Description())
|
||||||
assert.Contains(t, tool.Description(), "skill")
|
assert.Contains(t, tool.Description(), "skill")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
22
workspace/HEARTBEAT.md
Normal file
22
workspace/HEARTBEAT.md
Normal file
|
|
@ -0,0 +1,22 @@
|
||||||
|
# Heartbeat Check List
|
||||||
|
|
||||||
|
This file contains tasks for the heartbeat service to check periodically.
|
||||||
|
|
||||||
|
## Examples
|
||||||
|
|
||||||
|
- Check for unread messages
|
||||||
|
- Review upcoming calendar events
|
||||||
|
- Check device status (e.g., MaixCam)
|
||||||
|
|
||||||
|
## Instructions
|
||||||
|
|
||||||
|
- Execute ALL tasks listed below. Do NOT skip any task.
|
||||||
|
- For simple tasks (e.g., report current time), respond directly.
|
||||||
|
- For complex tasks that may take time, use the spawn tool to create a subagent.
|
||||||
|
- The spawn tool is async - subagent results will be sent to the user automatically.
|
||||||
|
- After spawning a subagent, CONTINUE to process remaining tasks.
|
||||||
|
- Only respond with HEARTBEAT_OK when ALL tasks are done AND nothing needs attention.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
Add your heartbeat tasks below this line:
|
||||||
4
workspace/cron/jobs.json
Normal file
4
workspace/cron/jobs.json
Normal file
|
|
@ -0,0 +1,4 @@
|
||||||
|
{
|
||||||
|
"version": 1,
|
||||||
|
"jobs": []
|
||||||
|
}
|
||||||
1
workspace/heartbeat.log
Normal file
1
workspace/heartbeat.log
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
[2026-03-24 08:15:50] [INFO] Created default HEARTBEAT.md template
|
||||||
4
workspace/state/state.json
Normal file
4
workspace/state/state.json
Normal file
|
|
@ -0,0 +1,4 @@
|
||||||
|
{
|
||||||
|
"last_channel": "telegram:8271300679",
|
||||||
|
"timestamp": "2026-03-24T08:43:14.295101255+01:00"
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue